aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorMax Nanis2026-09-10 01:00:39 -0700
committerMax Nanis2026-09-10 01:00:39 -0700
commit4dca7296742b607e74f16e2f6484c51163a41ace (patch)
tree0839c16d0905deb587d3ee25ff93fd3ccf7eeeb6
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
-rw-r--r--jb/api/auth.py3
-rw-r--r--jb/api/magic_token.py2
-rw-r--r--jb/config.py13
-rw-r--r--jb/decorators.py11
-rw-r--r--jb/flow/assignment_tasks.py17
-rw-r--r--jb/flow/events.py9
-rw-r--r--jb/flow/tasks.py17
-rw-r--r--jb/main.py5
-rw-r--r--jb/managers/amt.py16
-rw-r--r--jb/managers/email_manager.py7
-rw-r--r--jb/managers/gr_api.py12
-rw-r--r--jb/managers/hit.py2
-rw-r--r--jb/models/assignment.py5
-rw-r--r--jb/models/auth.py1
-rw-r--r--jb/models/hit.py110
-rw-r--r--jb/settings.py86
-rw-r--r--jb/views/auth.py4
-rw-r--r--tests/conftest.py125
-rw-r--r--tests/http/test_auth.py4
19 files changed, 275 insertions, 174 deletions
diff --git a/jb/api/auth.py b/jb/api/auth.py
index 411b8f1..1542e70 100644
--- a/jb/api/auth.py
+++ b/jb/api/auth.py
@@ -1,4 +1,3 @@
-import logging
from datetime import datetime, timedelta, timezone
from typing import Annotated
from uuid import uuid4
@@ -13,7 +12,6 @@ from jb.managers.gr_api import GRApiManager
from jb.models.auth import User
bearer = HTTPBearer(auto_error=False)
-logger = logging.getLogger(__name__)
SESSION_COOKIE_NAME = "jb_session"
JWT_ISSUER = "jamesbillings67"
@@ -74,6 +72,7 @@ def get_authenticated_user(
def create_session(product_user_id: str) -> str:
now = datetime.now(timezone.utc)
+ assert settings.session_jwt_secret
return jwt.encode(
{
"sub": product_user_id,
diff --git a/jb/api/magic_token.py b/jb/api/magic_token.py
index b7f7575..5f78996 100644
--- a/jb/api/magic_token.py
+++ b/jb/api/magic_token.py
@@ -40,7 +40,7 @@ def consume_magic_token(token: str) -> str:
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid or expired magic token",
)
- return user_email
+ return str(user_email)
def create_amt_account_link_token(email: str, amt_worker_id: str) -> str:
diff --git a/jb/config.py b/jb/config.py
index 359f108..7993f53 100644
--- a/jb/config.py
+++ b/jb/config.py
@@ -1,17 +1,8 @@
import logging
-from generalresearch.config import is_debug
+from jb.settings import get_settings
-from jb.settings import get_settings, get_test_settings
-
-if is_debug():
- print("running using TEST settings")
- settings = get_test_settings()
- assert settings.debug is True
-else:
- print("running using PROD settings")
- settings = get_settings()
- assert settings.debug is False
+settings = get_settings()
if settings.debug:
LOG_LEVEL = logging.DEBUG
diff --git a/jb/decorators.py b/jb/decorators.py
index cafd7a9..1a7a145 100644
--- a/jb/decorators.py
+++ b/jb/decorators.py
@@ -1,3 +1,5 @@
+import logging
+
import boto3
from botocore.config import Config
from generalresearch.pg_helper import PostgresConfig
@@ -22,6 +24,15 @@ redis_config = RedisConfig(
)
REDIS = redis_config.create_redis_client()
+# --- Logging ---
+
+logging.basicConfig(
+ level=logging.INFO,
+ format="%(asctime)s - %(levelname)s:%(name)s:%(message)s",
+ datefmt="%Y-%m-%d %H:%M:%S",
+)
+LOG = logging.getLogger("amtjb")
+
CLIENT_CONFIG = Config(
# connect_timeout (float or int) – The time in seconds till a timeout
# exception is thrown when attempting to make a connection. The default
diff --git a/jb/flow/assignment_tasks.py b/jb/flow/assignment_tasks.py
index bdebb3d..b3c820a 100644
--- a/jb/flow/assignment_tasks.py
+++ b/jb/flow/assignment_tasks.py
@@ -1,5 +1,4 @@
-import logging
-
+from jb.decorators import LOG
from jb.flow.monitoring import emit_assignment_event, emit_error_event
from jb.managers.amt import (
REJECT_MESSAGE_UNKNOWN_ASSIGNMENT,
@@ -27,7 +26,7 @@ def process_assignment_submitted(
#
# Step 1: Attempt to get the Assignment out of the API
#
- logging.info(f"{event=}")
+ LOG.info(f"{event=}")
# This is the assignment model from AMT. In the DB, we should only have
# the AssignmentStub
@@ -42,7 +41,7 @@ def process_assignment_submitted(
# It is not found in amt, either it is invalid, not yet submitted, or
# already been approved/rejected, so we just do nothing ...
# todo: maybe we confirm its state matches what we have in the db
- logging.warning(f"No assignment found on AMT: {event.amt_assignment_id}")
+ LOG.warning(f"No assignment found on AMT: {event.amt_assignment_id}")
emit_error_event(
event_type="assignment_not_found_in_amt",
amt_hit_type_id=event.amt_hit_type_id,
@@ -70,9 +69,7 @@ def review_hit(amtm: AMTManager, hm: HitManager, assignment: Assignment) -> None
hit, _ = amtm.get_hit_if_exists(amt_hit_id=assignment.amt_hit_id)
if hit is None:
- logging.warning(
- f"Hit not found when trying to review hit: {assignment.amt_hit_id}"
- )
+ LOG.warning(f"Hit not found when trying to review hit: {assignment.amt_hit_id}")
return
# Update the db
@@ -101,7 +98,7 @@ def reject_assignment(
event_type="failed_to_reject_assignment",
amt_hit_type_id=amt_hit_type_id,
)
- logging.exception(f"Failed to reject assignment: {amt_assignment_id}")
+ LOG.exception(f"Failed to reject assignment: {amt_assignment_id}")
# We just rejected this assignment, get it from amazon again
assignment = amtm.get_assignment(amt_assignment_id=amt_assignment_id)
@@ -112,7 +109,7 @@ def reject_assignment(
# need to create as assignment first ...
stub = am.get_stub_if_exists(amt_assignment_id=assignment.amt_assignment_id)
if stub is None:
- logging.warning(
+ LOG.warning(
f"Rejected assignment doesn't exist in DB. Creating ... : {amt_assignment_id}"
)
# Even if the assignment doesn't exist, the hit must ...
@@ -123,5 +120,5 @@ def reject_assignment(
emit_assignment_event(
status=AssignmentStatus.Rejected, amt_hit_type_id=amt_hit_type_id, reason=msg
)
- logging.warning(f"Rejected assignment: {amt_assignment_id}")
+ LOG.warning(f"Rejected assignment: {amt_assignment_id}")
return assignment
diff --git a/jb/flow/events.py b/jb/flow/events.py
index 0eb91fd..f224c01 100644
--- a/jb/flow/events.py
+++ b/jb/flow/events.py
@@ -1,4 +1,3 @@
-import logging
import time
from concurrent import futures
from concurrent.futures import Executor, ThreadPoolExecutor
@@ -11,7 +10,7 @@ from jb.config import (
CONSUMER_NAME,
JB_EVENTS_STREAM,
)
-from jb.decorators import REDIS
+from jb.decorators import LOG, REDIS
from jb.flow.assignment_tasks import process_assignment_submitted
from jb.flow.monitoring import emit_error_event
from jb.models.event import MTurkEvent
@@ -26,7 +25,7 @@ def process_mturk_events_task():
try:
process_mturk_events(executor=executor)
except Exception as e:
- logging.exception(e)
+ LOG.exception(e)
finally:
time.sleep(1)
@@ -71,7 +70,7 @@ def process_mturk_events_chunk(executor: Executor) -> int | None:
executor.submit(process_assignment_submitted_event, event, str(msg_id))
)
else:
- logging.info(f"Discarding {event}")
+ LOG.info(f"Discarding {event}")
REDIS.xdel(JB_EVENTS_STREAM, msg_id)
futures.wait(fs, timeout=60)
@@ -84,7 +83,7 @@ def process_assignment_submitted_event(event: MTurkEvent, msg_id: str):
try:
process_assignment_submitted(amtm=AMTM, am=AM, hm=HM, bm=BM, event=event)
except Exception as e:
- logging.exception(f"{event.amt_assignment_id=}, {e=}")
+ LOG.exception(f"{event.amt_assignment_id=}, {e=}")
emit_error_event(
event_type="failed_process_assignment_submitted",
amt_hit_type_id=event.amt_hit_type_id,
diff --git a/jb/flow/tasks.py b/jb/flow/tasks.py
index 24e96d4..c555021 100644
--- a/jb/flow/tasks.py
+++ b/jb/flow/tasks.py
@@ -1,19 +1,14 @@
-import logging
import time
from typing import TypedDict, cast
from generalresearch.config import is_debug
-from jb.decorators import AMTM, HM, HQM, HTM, pg_config
+from jb.decorators import AMTM, HM, HQM, HTM, LOG, pg_config
from jb.flow.maintenance import check_hit_status
from jb.flow.monitoring import emit_hit_event, write_hit_gauge
from jb.models.definitions import HitStatus
from jb.models.hit import Hit, HitQuestion, HitType
-logging.basicConfig()
-logger = logging.getLogger()
-logger.setLevel(logging.INFO)
-
class HitRow(TypedDict):
amt_hit_id: str
@@ -34,7 +29,7 @@ def check_stale_hits():
params={"status": HitStatus.Assignable.value},
)
for hit in cast(list[HitRow], res):
- logging.info(f"check_stale_hits: {hit["amt_hit_id"]}")
+ LOG.info(f"check_stale_hits: {hit["amt_hit_id"]}")
check_hit_status(
amtm=AMTM,
amt_hit_id=hit["amt_hit_id"],
@@ -56,7 +51,7 @@ def check_expired_hits():
params={"status": HitStatus.Assignable.value},
)
for hit in cast(list[HitRow], res):
- logging.info(f"check_expired_hits: {hit["amt_hit_id"]}")
+ LOG.info(f"check_expired_hits: {hit["amt_hit_id"]}")
check_hit_status(
amtm=AMTM,
amt_hit_id=hit["amt_hit_id"],
@@ -87,7 +82,7 @@ def refill_hits() -> None:
assert hit_type.amt_hit_type_id
active_count = HM.get_active_count(hit_type_id=hit_type.id)
- logging.info(
+ LOG.info(
f"HitType: {hit_type.amt_hit_type_id}, {hit_type.min_active=}, active_count={active_count}"
)
write_hit_gauge(
@@ -97,7 +92,7 @@ def refill_hits() -> None:
)
if active_count < hit_type.min_active:
cnt_todo = hit_type.min_active - active_count
- logging.info(f"Refilling {cnt_todo} hits")
+ LOG.info(f"Refilling {cnt_todo} hits")
for _ in range(cnt_todo):
create_hit_from_hittype(hit_type)
@@ -109,6 +104,6 @@ def refill_hits_task():
check_stale_hits()
refill_hits()
except Exception as e:
- logging.exception(e)
+ LOG.exception(e)
finally:
time.sleep(5 * 60)
diff --git a/jb/main.py b/jb/main.py
index 9f4f000..cbbda98 100644
--- a/jb/main.py
+++ b/jb/main.py
@@ -3,13 +3,12 @@ from typing import Any
from fastapi import FastAPI
from fastapi.responses import HTMLResponse
-from starlette.middleware.cors import CORSMiddleware
-from starlette.middleware.trustedhost import TrustedHostMiddleware
-
from jb.config import settings
from jb.settings import BASE_HTML
from jb.views.auth import auth_router
from jb.views.common import common_router
+from starlette.middleware.cors import CORSMiddleware
+from starlette.middleware.trustedhost import TrustedHostMiddleware
app = FastAPI(
servers=[
diff --git a/jb/managers/amt.py b/jb/managers/amt.py
index e2c7e90..2411080 100644
--- a/jb/managers/amt.py
+++ b/jb/managers/amt.py
@@ -1,4 +1,3 @@
-import logging
from datetime import datetime, timezone
from typing import Any
@@ -15,6 +14,7 @@ from mypy_boto3_mturk.type_defs import (
from pydantic import ValidationError
from jb.config import TOPIC_ARN
+from jb.decorators import LOG
from jb.models import AMTAccount
from jb.models.assignment import Assignment
from jb.models.bonus import Bonus
@@ -78,7 +78,7 @@ class AMTManager:
return HitStatus.Disposed
else:
- logging.warning(msg)
+ LOG.warning(msg)
return HitStatus.Unassignable
return res.status
@@ -137,7 +137,7 @@ class AMTManager:
# Baddies have been known to submit assignments with purposely
# malformed "answer" (xml) section, which will raise
# a pydantic validation error. Try to parse again with no Answer.
- logging.exception(e)
+ LOG.exception(e)
ass_res["Answer"] = None
# If it wasn't the Answer that caused the ValidationError, it'll raise again
assignment = Assignment.from_amt_get_assignment(ass_res)
@@ -152,7 +152,7 @@ class AMTManager:
try:
return self.get_assignment(amt_assignment_id=amt_assignment_id)
except botocore.exceptions.ClientError as e:
- logging.warning(e)
+ LOG.warning(e)
error_code = e.response["Error"]["Code"]
error_msg = e.response["Error"]["Message"]
if error_code == "RequestError" and expected_err_msg in error_msg:
@@ -170,7 +170,7 @@ class AMTManager:
)
except botocore.exceptions.ClientError as e:
- logging.warning(e)
+ LOG.warning(e)
return None
def approve_assignment_if_possible(
@@ -189,7 +189,7 @@ class AMTManager:
)
except botocore.exceptions.ClientError as e:
- logging.warning(e)
+ LOG.warning(e)
return None
def update_hit_review_status(self, amt_hit_id: str, revert: bool = False) -> None:
@@ -198,7 +198,7 @@ class AMTManager:
self.amt_client.update_hit_review_status(HITId=amt_hit_id, Revert=revert)
except botocore.exceptions.ClientError as e:
- logging.warning(f"{amt_hit_id=}, {e}")
+ LOG.warning(f"{amt_hit_id=}, {e}")
error_msg = e.response["Error"]["Message"]
if "does not exist" in error_msg:
@@ -225,7 +225,7 @@ class AMTManager:
)
except botocore.exceptions.ClientError as e:
- logging.warning(f"{amt_worker_id=} {amt_assignment_id=}, {e}")
+ LOG.warning(f"{amt_worker_id=} {amt_assignment_id=}, {e}")
return None
def get_bonus(self, amt_assignment_id: str, payout_event_id: str) -> Bonus | None:
diff --git a/jb/managers/email_manager.py b/jb/managers/email_manager.py
index e740e20..78acc9d 100644
--- a/jb/managers/email_manager.py
+++ b/jb/managers/email_manager.py
@@ -1,9 +1,12 @@
import requests
+from generalresearch.config import is_debug
from jb.config import settings
MAUTIC_BASE_URL = "https://mail.jamesbillings67.com"
EMAIL_TEMPLATE_ID = 1
+
+assert settings.mautic_api_key
auth_headers = {"Authorization": f"Basic {settings.mautic_api_key.get_secret_value()}"}
@@ -22,6 +25,10 @@ def get_or_create_contact(email: str, amt_worker_id: str | None = None):
def send_login_email_from_url(mautic_url: str, magic_link: str) -> None:
+ if is_debug():
+ print("MAGIC_LINK: ", magic_link)
+ return
+
email_tokens = {
"magic_link": magic_link,
}
diff --git a/jb/managers/gr_api.py b/jb/managers/gr_api.py
index 494ceea..70c6a60 100644
--- a/jb/managers/gr_api.py
+++ b/jb/managers/gr_api.py
@@ -1,14 +1,12 @@
"""Client for General Research's product-user API."""
-import logging
from typing import Any
import requests
+from jb.decorators import LOG
from jb.models.auth import User
-logger = logging.getLogger(__name__)
-
class GRApiError(RuntimeError):
"""The General Research API could not satisfy a request."""
@@ -60,7 +58,7 @@ class GRApiManager:
f"General Research API request failed: {method} {url}"
) from exc
- def _parse_user_response(self, res: dict) -> User:
+ def _parse_user_response(self, res: dict[str, Any]) -> User:
return User.model_validate(
{
"product_user_id": res["product_user_id"],
@@ -98,7 +96,7 @@ class GRApiManager:
"""This should only be called once per user upon account creation.
A user cannot change their email address."""
url = f"{self.base_url}/{self.product_id}/user/{user.product_user_id}/metadata/"
- res = self._request(
+ _ = self._request(
"PATCH",
url,
json={"email_address": str(user.email)},
@@ -110,7 +108,7 @@ class GRApiManager:
if user.display_name is None:
return user
url = f"{self.base_url}/{self.product_id}/user/{user.product_user_id}/metadata/"
- res = self._request(
+ _ = self._request(
"PATCH",
url,
json={"display_name": user.display_name},
@@ -145,7 +143,7 @@ class GRApiManager:
raise
self.set_user_email(user)
transitioned_user = self.get_user(user.product_user_id)
- logger.warning(
+ LOG.warning(
"Transitioned product user from AMT worker %s to %s with email %s",
amt_worker_id,
transitioned_user.product_user_id,
diff --git a/jb/managers/hit.py b/jb/managers/hit.py
index a178d4d..2c6067b 100644
--- a/jb/managers/hit.py
+++ b/jb/managers/hit.py
@@ -306,7 +306,7 @@ class HitManager(PostgresManager):
def get_active_count(self, hit_type_id: int) -> int:
return self.pg_config.execute_sql_query(
- """
+ query="""
SELECT COUNT(1) as active_count
FROM mtwerk_hit
WHERE status = %(status)s
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
diff --git a/jb/settings.py b/jb/settings.py
index 7747afc..f0851a5 100644
--- a/jb/settings.py
+++ b/jb/settings.py
@@ -1,14 +1,15 @@
-import os
from functools import lru_cache
+from os.path import abspath
+from os.path import dirname as pdirname
+from os.path import join as pjoin
from pathlib import Path
from generalresearch.models.custom_types import InfluxDsn
-from pydantic import Field, HttpUrl, PostgresDsn, RedisDsn, SecretStr
-from pydantic_settings import BaseSettings, SettingsConfigDict
-
from jb.models.custom_types import UUIDStr
+from pydantic import Field, HttpUrl, PostgresDsn, RedisDsn, SecretStr, model_validator
+from pydantic_settings import BaseSettings, SettingsConfigDict
-BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+BASE_DIR = pdirname(pdirname(abspath(__file__)))
BASE_HTML_PATH = Path(BASE_DIR) / "templates" / "base.html"
BASE_HTML = BASE_HTML_PATH.read_text()
@@ -20,24 +21,28 @@ class AmtJbBaseSettings(BaseSettings):
redis: RedisDsn | None = Field(default=None)
redis_timeout: float = Field(default=0.10)
- amt_jb_db: PostgresDsn = Field()
+ amt_jb_db: PostgresDsn | None = Field(default=None)
amt_endpoint: HttpUrl | None = Field(default=None)
amt_access_id: str | None = Field(default=None)
amt_secret_key: str | None = Field(default=None)
- aws_owner_id: str = Field()
- aws_subscription_arn: str = Field()
+ aws_owner_id: str | None = Field(default=None)
+ aws_subscription_arn: str | None = Field(default=None)
class Settings(AmtJbBaseSettings):
model_config = SettingsConfigDict(
- env_prefix="",
+ env_file=(
+ pjoin(BASE_DIR, x)
+ for x in [".env.test", ".env.testing", ".env.staging", ".env.prod"]
+ ),
+ env_file_encoding="utf-8",
case_sensitive=False,
- env_file=os.path.join(BASE_DIR, ".env"),
extra="allow",
cli_parse_args=False,
)
+
debug: bool = False
app_name: str = "AMT JB API"
base_url: HttpUrl = Field(default=HttpUrl("https://jamesbillings67.com/"))
@@ -46,41 +51,58 @@ class Settings(AmtJbBaseSettings):
# Needed for admin function on fsb w/o authentication
fsb_host_private_route: str | None = Field(default=None)
- product_id: UUIDStr = Field()
+ product_id: UUIDStr | None = Field(default=None)
influx_db: InfluxDsn | None = Field(default=None)
- sns_path: str = Field()
+ sns_path: str | None = Field(default=None)
session_token_ttl_seconds: int = Field(default=30 * 24 * 60 * 60, gt=0)
- session_jwt_secret: SecretStr = Field(min_length=32)
+ session_jwt_secret: SecretStr | None = Field(default=None, min_length=32)
- magic_token_salt: SecretStr = Field(min_length=32)
+ magic_token_salt: SecretStr | None = Field(default=None, min_length=32)
gr_api_host: HttpUrl = Field(default=HttpUrl("https://generalresearch.com/api/v2/"))
- gr_api_token: SecretStr = Field(min_length=1)
+ gr_api_token: SecretStr | None = Field(default=None, min_length=1)
- mautic_api_key: SecretStr = Field(min_length=32)
+ mautic_api_key: SecretStr | None = Field(default=None, min_length=32)
+ @model_validator(mode="after")
+ def validate_host_and_key(self) -> "Settings":
-class TestSettings(Settings):
- model_config = SettingsConfigDict(
- env_prefix="",
- case_sensitive=False,
- env_file=os.path.join(BASE_DIR, ".env.test"),
- extra="allow",
- cli_parse_args=False,
- )
- debug: bool = True
- app_name: str = "AMT JB API Test"
- base_url: HttpUrl = Field(default=HttpUrl("http://127.0.0.1:8081/"))
+ if not self.amt_jb_db:
+ raise ValueError("amt_jb_db is required")
+ if not self.aws_owner_id:
+ raise ValueError("aws_owner_id is required")
-@lru_cache
-def get_settings():
- return Settings()
+ if not self.aws_subscription_arn:
+ raise ValueError("aws_subscription_arn is required")
+
+ if not self.product_id:
+ raise ValueError("product_id is required")
+
+ if not self.sns_path:
+ raise ValueError("sns_path is required")
+
+ if not self.session_jwt_secret:
+ raise ValueError("session_jwt_secret is required")
+
+ if not self.magic_token_salt:
+ raise ValueError("magic_token_salt is required")
+
+ if self.session_jwt_secret == self.magic_token_salt:
+ raise ValueError("JWT Secret must be different than Magic Token Salt")
+
+ if not self.gr_api_token:
+ raise ValueError("gr_api_token is required")
+
+ if not self.mautic_api_key:
+ raise ValueError("mautic_api_key is required")
+
+ return self
@lru_cache
-def get_test_settings():
- return TestSettings()
+def get_settings():
+ return Settings()
diff --git a/jb/views/auth.py b/jb/views/auth.py
index 6f8a3e2..44fef10 100644
--- a/jb/views/auth.py
+++ b/jb/views/auth.py
@@ -1,4 +1,3 @@
-import logging
from typing import Annotated
from urllib.parse import urlencode
@@ -17,6 +16,7 @@ from jb.api.magic_token import (
create_magic_token,
)
from jb.config import settings
+from jb.decorators import LOG
from jb.dependencies import get_gr_api_manager
from jb.managers.email_manager import (
get_or_create_contact,
@@ -147,7 +147,7 @@ def exchange_amt_account_link(
try:
_exchange_amt_account_link(body.token, response, gr_api)
except ValueError as e:
- logging.error(f"Failed to exchange AMT account link: {e}")
+ LOG.error(f"Failed to exchange AMT account link: {e}")
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e))
diff --git a/tests/conftest.py b/tests/conftest.py
index 2a3a580..002aced 100644
--- a/tests/conftest.py
+++ b/tests/conftest.py
@@ -1,12 +1,18 @@
+from __future__ import annotations
+
import os
+import subprocess
+import sys
+from collections.abc import Callable
+from pathlib import Path
from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from _pytest.config import Config
-from dotenv import load_dotenv
-from generalresearch.pg_helper import PostgresConfig
+from generalresearch.models.custom_types import PostgresDict
+from generalresearch.pg_helper import PostgresConfig, PostgresDsn
from mypy_boto3_mturk import MTurkClient
+from pytest import TempPathFactory
from jb.decorators import CLIENT_CONFIG
from tests import generate_amt_id
@@ -77,26 +83,108 @@ def pe_id() -> str:
@pytest.fixture(scope="session")
-def env_file_path(pytestconfig: Config) -> str:
- root_path = pytestconfig.rootpath
- env_path = os.path.join(root_path, ".env.test")
+def settings() -> "Settings":
+ from jb.settings import Settings as JBSettings
+
+ return JBSettings()
- if os.path.exists(env_path):
- load_dotenv(dotenv_path=env_path, override=True)
- return env_path
+# --- Database Connectors ---
@pytest.fixture(scope="session")
-def settings(env_file_path: str) -> "Settings":
- from jb.settings import Settings as JBSettings
+def django_db_factory(
+ postgres_instance: PostgresDsn,
+ gr_repo: Callable[..., Path],
+ django_settings_file: Callable[..., tuple[str, Path]],
+ postgres_instance_dict: PostgresDict,
+ tmp_path_factory: TempPathFactory,
+) -> Callable[..., PostgresDsn | None]:
+
+ _ran = {}
+
+ def _inner(
+ django_project: str = "generalresearch.thl_django",
+ ) -> PostgresDsn | None:
+
+ if _ran.get(django_project, False):
+ print(f"Already ran django_db_factory:{django_project}")
+ return postgres_instance
+ _ran[django_project] = True
+
+ _cwd = None
+ _manage_path = "generalresearch.thl_django.app.manage"
+ _settings_module, _settings_dir = django_settings_file(
+ extra_installed_apps=[
+ "generalresearch.thl_django",
+ ],
+ )
+
+ pythonpath = str(_settings_dir)
+ if existing_pythonpath := os.environ.get("PYTHONPATH"):
+ pythonpath += os.pathsep + existing_pythonpath
+
+ env = {
+ **os.environ,
+ "DJANGO_SETTINGS_MODULE": _settings_module,
+ "PYTHONPATH": pythonpath,
+ }
+
+ # we check right after. if we check now, we won't print if bad
+ res1 = subprocess.run( # noqa: PLW1510
+ [
+ sys.executable,
+ "-m",
+ _manage_path,
+ "makemigrations",
+ f"--settings={_settings_module}",
+ ],
+ cwd=str(_cwd) if _cwd is not None else None,
+ env=env,
+ capture_output=True,
+ text=True,
+ )
+
+ if res1.returncode != 0:
+ print("STDOUT:", res1.stdout)
+ print("STDERR:", res1.stderr)
+ res1.check_returncode()
+
+ res2 = subprocess.run( # noqa: PLW1510
+ [
+ sys.executable,
+ "-m",
+ _manage_path,
+ "migrate",
+ f"--settings={_settings_module}",
+ ],
+ env=env,
+ cwd=str(_cwd) if _cwd is not None else None,
+ capture_output=True,
+ text=True,
+ )
+
+ if res2.returncode != 0:
+ print("STDOUT:", res2.stdout)
+ print("STDERR:", res2.stderr)
+ res2.check_returncode()
+
+ # 3. Return the Dsn so the factory gives a way to connect
+ return postgres_instance
+
+ return _inner
- s = JBSettings(_env_file=env_file_path)
- return s
+@pytest.fixture(scope="session")
+def pg_config(settings: "Settings") -> PostgresConfig:
+ return PostgresConfig(
+ dsn=settings.amt_jb_db,
+ connect_timeout=1,
+ statement_timeout=1,
+ )
-# --- Database Connectors ---
+# --- Redis ---
@pytest.fixture(scope="session")
@@ -112,15 +200,6 @@ def redis(settings: "Settings"):
return redis_config.create_redis_client()
-@pytest.fixture(scope="session")
-def pg_config(settings: "Settings") -> PostgresConfig:
- return PostgresConfig(
- dsn=settings.amt_jb_db,
- connect_timeout=1,
- statement_timeout=1,
- )
-
-
# --- Connectors ---
@pytest.fixture(scope="session")
def amt_client(settings: "Settings") -> MTurkClient:
diff --git a/tests/http/test_auth.py b/tests/http/test_auth.py
index ebda742..1fa1335 100644
--- a/tests/http/test_auth.py
+++ b/tests/http/test_auth.py
@@ -1,4 +1,3 @@
-import secrets
from urllib.parse import parse_qs, urlparse
import pytest
@@ -38,8 +37,7 @@ class FakeGRApiManager:
@pytest.fixture
def email() -> str:
- email = secrets.token_urlsafe(16) + "@gmail.com"
- return email.lower()
+ return "unittest@generalresearch.com"
@pytest.fixture