diff options
Diffstat (limited to 'jb/models/hit.py')
| -rw-r--r-- | jb/models/hit.py | 153 |
1 files changed, 78 insertions, 75 deletions
diff --git a/jb/models/hit.py b/jb/models/hit.py index fba2ecf..a550943 100644 --- a/jb/models/hit.py +++ b/jb/models/hit.py @@ -1,25 +1,26 @@ -from datetime import datetime, timezone, timedelta -from typing import Optional, List, Dict, Any +from datetime import datetime, timedelta, timezone +from typing import Any from uuid import uuid4 from xml.etree import ElementTree +from generalresearch.currency import USDCent +from generalresearch.models.custom_types import AwareDatetimeISO, HttpsUrlStr from mypy_boto3_mturk.type_defs import HITTypeDef from pydantic import ( BaseModel, - Field, - PositiveInt, ConfigDict, + Field, NonNegativeInt, + PositiveInt, ) from typing_extensions import Self -from generalresearchutils.currency import USDCent -from jb.models.custom_types import AMTBoto3ID, HttpsUrlStr, AwareDatetimeISO -from jb.models.definitions import HitStatus, HitReviewStatus +from jb.models.custom_types import AMTBoto3ID +from jb.models.definitions import HitReviewStatus, HitStatus class HitQuestion(BaseModel): - id: Optional[PositiveInt] = Field(default=None) + id: PositiveInt | None = Field(default=None) url: HttpsUrlStr = Field() height: PositiveInt = Field(default=1_200, ge=100, le=4_000) @@ -33,7 +34,7 @@ class HitQuestion(BaseModel): def xml(self) -> str: return f"""<?xml version="1.0" encoding="UTF-8"?> <ExternalQuestion xmlns="http://mechanicalturk.amazonaws.com/AWSMechanicalTurkDataSchemas/2006-07-14/ExternalQuestion.xsd"> - <ExternalURL>{str(self.url)}</ExternalURL> + <ExternalURL>{self.url!s}</ExternalURL> <FrameHeight>{self.height}</FrameHeight> </ExternalQuestion>""" @@ -82,21 +83,25 @@ class HitType(HitTypeCommon): https://docs.aws.amazon.com/AWSMechTurk/latest/AWSMturkAPI/ApiReference_CreateHITTypeOperation.html """ - id: Optional[PositiveInt] = Field(default=None) - amt_hit_type_id: Optional[AMTBoto3ID] = Field(default=None) + id: PositiveInt | None = Field(default=None) + amt_hit_type_id: AMTBoto3ID | None = Field(default=None) # --- GRL Specific --- min_active: NonNegativeInt = Field(default=0, le=100_000) - def to_api_request_body(self): - return dict( - AutoApprovalDelayInSeconds=round(self.auto_approval_delay.total_seconds()), - AssignmentDurationInSeconds=round(self.assignment_duration.total_seconds()), - Reward=str(self.reward.to_usd()), - Title=self.title, - Keywords=self.keywords, - Description=self.description, - ) + def to_api_request_body(self) -> dict[str, Any]: + return { + "AutoApprovalDelayInSeconds": round( + self.auto_approval_delay.total_seconds() + ), + "AssignmentDurationInSeconds": round( + self.assignment_duration.total_seconds() + ), + "Reward": str(self.reward.to_usd()), + "Title": self.title, + "Keywords": self.keywords, + "Description": self.description, + } def to_postgres(self): d = self.model_dump(mode="json") @@ -104,12 +109,12 @@ class HitType(HitTypeCommon): return d @classmethod - def from_postgres(cls, data: Dict[str, Any]) -> Self: + def from_postgres(cls, data: dict[str, Any]) -> Self: data["reward"] = USDCent(round(data["reward"] * 100)) return cls.model_validate(data) - def generate_hit_amt_request(self, question: HitQuestion) -> Dict[str, Any]: - d = dict() + def generate_hit_amt_request(self, question: HitQuestion) -> dict[str, Any]: + d = {} d["HITTypeId"] = self.amt_hit_type_id d["MaxAssignments"] = 1 d["LifetimeInSeconds"] = round(timedelta(days=14).total_seconds()) @@ -124,9 +129,9 @@ class Hit(HitTypeCommon): validate_assignment=True, ) - id: Optional[PositiveInt] = Field(default=None) - hit_type_id: Optional[PositiveInt] = Field(default=None) - question_id: Optional[PositiveInt] = Field(default=None) + id: PositiveInt | None = Field(default=None) + hit_type_id: PositiveInt | None = Field(default=None) + question_id: PositiveInt | None = Field(default=None) amt_hit_id: AMTBoto3ID = Field() amt_hit_type_id: AMTBoto3ID = Field() @@ -138,10 +143,8 @@ class Hit(HitTypeCommon): # TODO: Check if this is actually ever going to be None. I type fixed it, # but I don't have anything to suggest it isn't requred. -- Max 2026-02-24 - creation_time: Optional[AwareDatetimeISO] = Field( - default=None, description="From aws" - ) - expiration: Optional[AwareDatetimeISO] = Field(default=None) + creation_time: AwareDatetimeISO | None = Field(default=None, description="From aws") + expiration: AwareDatetimeISO | None = Field(default=None) # GRL Specific created_at: AwareDatetimeISO = Field( @@ -155,7 +158,7 @@ class Hit(HitTypeCommon): # -- Hit specific - qualification_requirements: Optional[List[Dict[str, Any]]] = Field(default=None) + qualification_requirements: list[dict[str, Any]] | None = Field(default=None) max_assignments: int = Field() # # this comes back as expiration. only for the request @@ -177,27 +180,27 @@ class Hit(HitTypeCommon): assert hit_type.amt_hit_type_id is not None h = cls.model_validate( - dict( - amt_hit_id=data["HITId"], - amt_hit_type_id=data["HITTypeId"], - amt_group_id=data["HITGroupId"], - status=HitStatus[data["HITStatus"]], - review_status=HitReviewStatus[data["HITReviewStatus"]], - creation_time=data["CreationTime"].astimezone(tz=timezone.utc), - expiration=data["Expiration"].astimezone(tz=timezone.utc), - hit_question_xml=data["Question"], - qualification_requirements=data["QualificationRequirements"], - max_assignments=data["MaxAssignments"], - assignment_pending_count=data["NumberOfAssignmentsPending"], - assignment_available_count=data["NumberOfAssignmentsAvailable"], - assignment_completed_count=data["NumberOfAssignmentsCompleted"], - description=data["Description"], - keywords=data["Keywords"], - reward=USDCent(round(float(data["Reward"]) * 100)), - title=data["Title"], - question_id=question.id, - hit_type_id=hit_type.id, - ) + { + "amt_hit_id": data["HITId"], + "amt_hit_type_id": data["HITTypeId"], + "amt_group_id": data["HITGroupId"], + "status": HitStatus[data["HITStatus"]], + "review_status": HitReviewStatus[data["HITReviewStatus"]], + "creation_time": data["CreationTime"].astimezone(tz=timezone.utc), + "expiration": data["Expiration"].astimezone(tz=timezone.utc), + "hit_question_xml": data["Question"], + "qualification_requirements": data["QualificationRequirements"], + "max_assignments": data["MaxAssignments"], + "assignment_pending_count": data["NumberOfAssignmentsPending"], + "assignment_available_count": data["NumberOfAssignmentsAvailable"], + "assignment_completed_count": data["NumberOfAssignmentsCompleted"], + "description": data["Description"], + "keywords": data["Keywords"], + "reward": USDCent(round(float(data["Reward"]) * 100)), + "title": data["Title"], + "question_id": question.id, + "hit_type_id": hit_type.id, + } ) return h @@ -205,27 +208,27 @@ class Hit(HitTypeCommon): @classmethod def from_amt_get_hit(cls, data: HITTypeDef) -> Self: h = cls.model_validate( - dict( - amt_hit_id=data["HITId"], - amt_hit_type_id=data["HITTypeId"], - amt_group_id=data["HITGroupId"], - status=HitStatus[data["HITStatus"]], - review_status=HitReviewStatus[data["HITReviewStatus"]], - creation_time=data["CreationTime"].astimezone(tz=timezone.utc), - expiration=data["Expiration"].astimezone(tz=timezone.utc), - hit_question_xml=data["Question"], - qualification_requirements=data["QualificationRequirements"], - max_assignments=data["MaxAssignments"], - assignment_pending_count=data["NumberOfAssignmentsPending"], - assignment_available_count=data["NumberOfAssignmentsAvailable"], - assignment_completed_count=data["NumberOfAssignmentsCompleted"], - description=data["Description"], - keywords=data["Keywords"], - reward=USDCent(round(float(data["Reward"]) * 100)), - title=data["Title"], - question_id=None, - hit_type_id=None, - ) + { + "amt_hit_id": data["HITId"], + "amt_hit_type_id": data["HITTypeId"], + "amt_group_id": data["HITGroupId"], + "status": HitStatus[data["HITStatus"]], + "review_status": HitReviewStatus[data["HITReviewStatus"]], + "creation_time": data["CreationTime"].astimezone(tz=timezone.utc), + "expiration": data["Expiration"].astimezone(tz=timezone.utc), + "hit_question_xml": data["Question"], + "qualification_requirements": data["QualificationRequirements"], + "max_assignments": data["MaxAssignments"], + "assignment_pending_count": data["NumberOfAssignmentsPending"], + "assignment_available_count": data["NumberOfAssignmentsAvailable"], + "assignment_completed_count": data["NumberOfAssignmentsCompleted"], + "description": data["Description"], + "keywords": data["Keywords"], + "reward": USDCent(round(float(data["Reward"]) * 100)), + "title": data["Title"], + "question_id": None, + "hit_type_id": None, + } ) return h @@ -235,7 +238,7 @@ class Hit(HitTypeCommon): return d @classmethod - def from_postgres(cls, data: Dict[str, Any]) -> Self: + def from_postgres(cls, data: dict[str, Any]) -> Self: data["reward"] = USDCent(round(data["reward"] * 100)) return cls.model_validate(data) @@ -248,7 +251,7 @@ class Hit(HitTypeCommon): } res = {} - lookup_table = dict(ExternalURL="url", FrameHeight="height") + lookup_table = {"ExternalURL": "url", "FrameHeight": "height"} for a in root.findall("mt:*", ns): key = lookup_table[a.tag.split("}")[1]] val = a.text |
