aboutsummaryrefslogtreecommitdiff
path: root/jb/models
diff options
context:
space:
mode:
authorMax Nanis2026-09-10 01:00:39 -0700
committerMax Nanis2026-09-10 01:00:39 -0700
commit4dca7296742b607e74f16e2f6484c51163a41ace (patch)
tree0839c16d0905deb587d3ee25ff93fd3ccf7eeeb6 /jb/models
parent832aecaddce80e312095ecdb572d7756eb9df5e9 (diff)
downloadamt-jb-4dca7296742b607e74f16e2f6484c51163a41ace.tar.gz
amt-jb-4dca7296742b607e74f16e2f6484c51163a41ace.zip
using model_validator on GRLSettings. Allows null default values, then to asser them on load. Required so pydantic_settings can be loaded in tests without params
Diffstat (limited to 'jb/models')
-rw-r--r--jb/models/assignment.py5
-rw-r--r--jb/models/auth.py1
-rw-r--r--jb/models/hit.py110
3 files changed, 61 insertions, 55 deletions
diff --git a/jb/models/assignment.py b/jb/models/assignment.py
index 775cd63..1f7033d 100644
--- a/jb/models/assignment.py
+++ b/jb/models/assignment.py
@@ -1,4 +1,3 @@
-import logging
from datetime import datetime, timezone
from typing import Any, TypedDict
from xml.etree import ElementTree
@@ -15,6 +14,7 @@ from pydantic import (
)
from typing_extensions import Self
+from jb.decorators import LOG
from jb.models.custom_types import AMTBoto3ID, AwareDatetimeISO, UUIDStr
from jb.models.definitions import AssignmentStatus
@@ -140,7 +140,8 @@ class Assignment(AssignmentStub):
values["tsid"] = TypeAdapter(UUIDStr).validate_python(tsid)
except ValidationError as e:
# Don't break the model validation if a baddie messes with the tsid in the answer.
- logging.warning(e)
+ LOG.warning(e)
+
values["tsid"] = None
return values
diff --git a/jb/models/auth.py b/jb/models/auth.py
index fef5070..886a607 100644
--- a/jb/models/auth.py
+++ b/jb/models/auth.py
@@ -22,6 +22,7 @@ def email_to_product_user_id(email: str) -> str:
The same normalized email and secret salt always produce the same ID. Keep
the salt private and stable; changing it changes every generated ID.
"""
+ assert settings.magic_token_salt
salt_bytes = settings.magic_token_salt.get_secret_value().encode("utf-8")
if len(salt_bytes) < 32:
raise ValueError("salt must be at least 32 bytes")
diff --git a/jb/models/hit.py b/jb/models/hit.py
index f6c854f..a091c83 100644
--- a/jb/models/hit.py
+++ b/jb/models/hit.py
@@ -88,15 +88,19 @@ class HitType(HitTypeCommon):
# --- 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")
@@ -109,7 +113,7 @@ class HitType(HitTypeCommon):
return cls.model_validate(data)
def generate_hit_amt_request(self, question: HitQuestion) -> dict[str, Any]:
- d = dict()
+ d = {}
d["HITTypeId"] = self.amt_hit_type_id
d["MaxAssignments"] = 1
d["LifetimeInSeconds"] = round(timedelta(days=14).total_seconds())
@@ -175,27 +179,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
@@ -203,27 +207,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
@@ -246,7 +250,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