aboutsummaryrefslogtreecommitdiff
path: root/jb/models/hit.py
diff options
context:
space:
mode:
Diffstat (limited to 'jb/models/hit.py')
-rw-r--r--jb/models/hit.py153
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