aboutsummaryrefslogtreecommitdiff
path: root/jb
diff options
context:
space:
mode:
Diffstat (limited to 'jb')
-rw-r--r--jb/api/__init__.py0
-rw-r--r--jb/api/auth.py88
-rw-r--r--jb/api/magic_token.py96
-rw-r--r--jb/config.py14
-rw-r--r--jb/decorators.py36
-rw-r--r--jb/dependencies.py8
-rw-r--r--jb/flow/assignment_tasks.py468
-rw-r--r--jb/flow/events.py94
-rw-r--r--jb/flow/maintenance.py4
-rw-r--r--jb/flow/monitoring.py23
-rw-r--r--jb/flow/setup_tasks.py5
-rw-r--r--jb/flow/tasks.py23
-rw-r--r--jb/main.py12
-rw-r--r--jb/managers/__init__.py23
-rw-r--r--jb/managers/amt.py47
-rw-r--r--jb/managers/assignment.py87
-rw-r--r--jb/managers/base.py16
-rw-r--r--jb/managers/bonus.py24
-rw-r--r--jb/managers/email_manager.py65
-rw-r--r--jb/managers/gr_api.py160
-rw-r--r--jb/managers/hit.py119
-rw-r--r--jb/managers/thl.py74
-rw-r--r--jb/managers/worker.py6
-rw-r--r--jb/models/__init__.py40
-rw-r--r--jb/models/amt.py19
-rw-r--r--jb/models/assignment.py32
-rw-r--r--jb/models/auth.py110
-rw-r--r--jb/models/bonus.py15
-rw-r--r--jb/models/custom_types.py99
-rw-r--r--jb/models/errors.py4
-rw-r--r--jb/models/event.py7
-rw-r--r--jb/models/hit.py153
-rw-r--r--jb/models/response.py21
-rw-r--r--jb/settings.py128
-rw-r--r--jb/views/auth.py203
-rw-r--r--jb/views/common.py68
-rw-r--r--jb/views/tasks.py78
37 files changed, 1237 insertions, 1232 deletions
diff --git a/jb/api/__init__.py b/jb/api/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/jb/api/__init__.py
diff --git a/jb/api/auth.py b/jb/api/auth.py
new file mode 100644
index 0000000..1542e70
--- /dev/null
+++ b/jb/api/auth.py
@@ -0,0 +1,88 @@
+from datetime import datetime, timedelta, timezone
+from typing import Annotated
+from uuid import uuid4
+
+import jwt
+from fastapi import Depends, HTTPException, Request, status
+from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
+
+from jb.config import settings
+from jb.dependencies import get_gr_api_manager
+from jb.managers.gr_api import GRApiManager
+from jb.models.auth import User
+
+bearer = HTTPBearer(auto_error=False)
+
+SESSION_COOKIE_NAME = "jb_session"
+JWT_ISSUER = "jamesbillings67"
+JWT_AUDIENCE = "jamesbillings67"
+
+
+def get_authenticated_user(
+ request: Request,
+ credentials: Annotated[HTTPAuthorizationCredentials | None, Depends(bearer)],
+ gr_api: Annotated[GRApiManager, Depends(get_gr_api_manager)],
+) -> User:
+ """FastAPI dependency for endpoints requiring a valid session."""
+ if settings.session_jwt_secret is None:
+ raise HTTPException(
+ status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ detail="Account session signing key is not configured",
+ )
+
+ token = request.cookies.get(SESSION_COOKIE_NAME)
+ if credentials is not None and credentials.scheme.lower() == "bearer":
+ token = credentials.credentials
+
+ if token is None:
+ raise HTTPException(
+ status_code=status.HTTP_401_UNAUTHORIZED,
+ detail="Missing session token",
+ headers={"WWW-Authenticate": "Bearer"},
+ )
+
+ try:
+ claims = jwt.decode(
+ token,
+ settings.session_jwt_secret.get_secret_value(),
+ algorithms=["HS256"],
+ issuer=JWT_ISSUER,
+ audience=JWT_AUDIENCE,
+ options={"require": ["sub", "iat", "exp", "jti", "iss", "aud", "type"]},
+ )
+ except jwt.PyJWTError:
+ raise HTTPException(
+ status_code=status.HTTP_401_UNAUTHORIZED,
+ detail="Invalid or expired session token",
+ headers={"WWW-Authenticate": "Bearer"},
+ )
+
+ if claims["type"] != "session" or not isinstance(claims["sub"], str):
+ raise HTTPException(
+ status_code=status.HTTP_401_UNAUTHORIZED,
+ detail="Invalid or expired session token",
+ headers={"WWW-Authenticate": "Bearer"},
+ )
+ product_user_id = claims.get("sub")
+
+ user = gr_api.get_user(product_user_id=product_user_id)
+
+ return 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,
+ "iat": now,
+ "exp": now + timedelta(seconds=settings.session_token_ttl_seconds),
+ "jti": uuid4().hex,
+ "iss": JWT_ISSUER,
+ "aud": JWT_AUDIENCE,
+ "type": "session",
+ },
+ settings.session_jwt_secret.get_secret_value(),
+ algorithm="HS256",
+ )
diff --git a/jb/api/magic_token.py b/jb/api/magic_token.py
new file mode 100644
index 0000000..b60d5e9
--- /dev/null
+++ b/jb/api/magic_token.py
@@ -0,0 +1,96 @@
+import hashlib
+import secrets
+
+from fastapi import HTTPException, status
+from generalresearch.redis_helper import RedisConfig
+
+from jb.decorators import get_redis_config
+from jb.models.auth import AmtAccountLink
+
+MAGIC_TOKEN_PREFIX = "auth:magic:"
+AMT_ACCOUNT_LINK_TOKEN_PREFIX = "auth:amt-account-link:"
+MAGIC_TOKEN_TTL: int = 10 * 60 # 10 minutes, in seconds
+
+
+def redis_token_key(token: str, prefix: str = MAGIC_TOKEN_PREFIX) -> str:
+ # Redis never contains a usable credential, even if its keys are exposed.
+ digest = hashlib.sha256(token.encode("utf-8")).hexdigest()
+ return f"{prefix}{digest}"
+
+
+def create_magic_token(user_email: str, redis_config: RedisConfig | None = None) -> str:
+ """Create a short-lived, single-use token for a user.
+ The raw token can then be sent by email.
+ """
+ if not user_email or not user_email.strip():
+ raise ValueError("user_email must not be empty")
+
+ token = secrets.token_urlsafe(32)
+
+ if redis_config is None:
+ redis_config = get_redis_config()
+
+ redis_client = redis_config.create_redis_client()
+
+ redis_client.set(
+ redis_token_key(token),
+ user_email,
+ ex=MAGIC_TOKEN_TTL,
+ )
+ return token
+
+
+def consume_magic_token(token: str, redis_config: RedisConfig | None = None) -> str:
+ if redis_config is None:
+ redis_config = get_redis_config()
+
+ redis_client = redis_config.create_redis_client()
+
+ user_email = redis_client.getdel(redis_token_key(token))
+ if user_email is None:
+ raise HTTPException(
+ status_code=status.HTTP_401_UNAUTHORIZED,
+ detail="Invalid or expired magic token",
+ )
+ return str(user_email)
+
+
+def create_amt_account_link_token(
+ email: str, amt_worker_id: str, redis_config: RedisConfig | None = None
+) -> str:
+ """Bind an email and AMT worker ID to an opaque, short-lived token."""
+ if redis_config is None:
+ redis_config = get_redis_config()
+
+ redis_client = redis_config.create_redis_client()
+
+ data = AmtAccountLink(
+ email=email,
+ amt_worker_id=amt_worker_id,
+ )
+ token = secrets.token_urlsafe(32)
+ redis_client.set(
+ redis_token_key(token, AMT_ACCOUNT_LINK_TOKEN_PREFIX),
+ data.model_dump_json(),
+ ex=MAGIC_TOKEN_TTL,
+ )
+ return token
+
+
+def consume_amt_account_link_token(
+ token: str, redis_config: RedisConfig | None = None
+) -> AmtAccountLink:
+ """Atomically consume and validate an AMT account-link token."""
+ if redis_config is None:
+ redis_config = get_redis_config()
+ redis_client = redis_config.create_redis_client()
+
+ raw_data = redis_client.getdel(
+ redis_token_key(token, AMT_ACCOUNT_LINK_TOKEN_PREFIX)
+ )
+ if raw_data is None:
+ raise HTTPException(
+ status_code=status.HTTP_401_UNAUTHORIZED,
+ detail="Invalid or expired account-link token",
+ )
+ return AmtAccountLink.model_validate_json(raw_data)
diff --git a/jb/config.py b/jb/config.py
index c7d07e5..7993f53 100644
--- a/jb/config.py
+++ b/jb/config.py
@@ -1,17 +1,8 @@
import logging
-from generalresearchutils.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
@@ -30,7 +21,6 @@ SUBSCRIPTION = {
}
JB_EVENTS_STREAM = "amt_jb_events"
-JB_EVENTS_FAILED_STREAM = "amt_jb_events_failed"
CONSUMER_GROUP = "amt-jb-0"
# We'll only have 1 consumer atm, change this if we don't
CONSUMER_NAME = "amt-jb-0"
diff --git a/jb/decorators.py b/jb/decorators.py
index 5c1b1f5..9c7a31c 100644
--- a/jb/decorators.py
+++ b/jb/decorators.py
@@ -1,17 +1,20 @@
+import logging
+
import boto3
from botocore.config import Config
-from generalresearchutils.pg_helper import PostgresConfig
-from generalresearchutils.redis_helper import RedisConfig
+from generalresearch.managers.base import Permission
+from generalresearch.pg_helper import PostgresConfig
+from generalresearch.redis_helper import RedisConfig
from influxdb import InfluxDBClient
from mypy_boto3_mturk import MTurkClient
from mypy_boto3_sns import SNSClient
from jb.config import settings
-from jb.managers import Permission
from jb.managers.amt import AMTManager
from jb.managers.assignment import AssignmentManager
from jb.managers.bonus import BonusManager
-from jb.managers.hit import HitTypeManager, HitManager, HitQuestionManager
+from jb.managers.gr_api import GRApiManager
+from jb.managers.hit import HitManager, HitQuestionManager, HitTypeManager
redis_config = RedisConfig(
dsn=settings.redis,
@@ -19,7 +22,24 @@ redis_config = RedisConfig(
socket_timeout=settings.redis_timeout,
socket_connect_timeout=settings.redis_timeout,
)
-REDIS = redis_config.create_redis_client()
+
+
+def get_redis_config():
+ return redis_config
+
+
+def get_redis():
+ return 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
@@ -32,6 +52,12 @@ CLIENT_CONFIG = Config(
read_timeout=2.5,
)
+gr_api_manager = GRApiManager(
+ base_url=str(settings.gr_api_host),
+ token=settings.gr_api_token.get_secret_value(),
+ product_id=settings.product_id,
+)
+
# We shouldn't use this directly. Use our AMTManager wrapper
AMT_CLIENT: MTurkClient = boto3.client(
service_name="mturk",
diff --git a/jb/dependencies.py b/jb/dependencies.py
new file mode 100644
index 0000000..b9577bc
--- /dev/null
+++ b/jb/dependencies.py
@@ -0,0 +1,8 @@
+"""FastAPI dependency providers shared by application routes."""
+
+from jb.decorators import gr_api_manager
+from jb.managers.gr_api import GRApiManager
+
+
+def get_gr_api_manager() -> GRApiManager:
+ return gr_api_manager
diff --git a/jb/flow/assignment_tasks.py b/jb/flow/assignment_tasks.py
index 18a844e..b3c820a 100644
--- a/jb/flow/assignment_tasks.py
+++ b/jb/flow/assignment_tasks.py
@@ -1,37 +1,15 @@
-import logging
-import math
-from datetime import timedelta
-from typing import Optional
-
-from generalresearchutils.models.thl.definitions import PayoutStatus, StatusCode1
-from generalresearchutils.models.thl.wallet.cashout_method import CashoutRequestInfo
-from generalresearchutils.currency import USDCent
-
-from jb.flow.monitoring import emit_error_event, emit_assignment_event, emit_bonus_event
+from jb.decorators import LOG
+from jb.flow.monitoring import emit_assignment_event, emit_error_event
from jb.managers.amt import (
- AMTManager,
REJECT_MESSAGE_UNKNOWN_ASSIGNMENT,
- REJECT_MESSAGE_NO_WORK,
- NO_WORK_APPROVAL_MESSAGE,
- REJECT_MESSAGE_BADDIE,
- APPROVAL_MESSAGE,
- BONUS_MESSAGE,
-)
-from jb.managers.thl import (
- get_user_blocked,
- get_task_status,
- user_cashout_request,
- manage_pending_cashout,
- get_user_blocked_or_not_exists,
- get_wallet_balance_if_non_negative,
+ AMTManager,
)
+from jb.managers.assignment import AssignmentManager
+from jb.managers.bonus import BonusManager
+from jb.managers.hit import HitManager
from jb.models.assignment import Assignment
from jb.models.definitions import AssignmentStatus
from jb.models.event import MTurkEvent
-from jb.managers.assignment import AssignmentManager
-from jb.managers.hit import HitManager
-from jb.managers.bonus import BonusManager
-from jb.config import settings
def process_assignment_submitted(
@@ -42,16 +20,13 @@ def process_assignment_submitted(
event: MTurkEvent,
) -> None:
"""
- Called either directly or from the SNS Notification that a
- HIT was submitted
-
- :return: None
+ Reject any submitted assignments
"""
#
# 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
@@ -66,95 +41,26 @@ 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,
)
- return None
+ return
# Even if the assignment doesn't exist, the hit must ...
hit = hm.get_from_amt_id(amt_hit_id=assignment.amt_hit_id)
- #
- # Step 2: Attempt to get the Assignment out of the DB
- #
- # Now, we need to confirm it is something that we have in the db. If not,
- # that means either something broke, or some funny business is happening
- # (maybe a baddie is submitting an assignment without doing any work).
- stub = am.get_stub_if_exists(amt_assignment_id=assignment.amt_assignment_id)
- if stub is None:
- # When they visited the "work" page, it should have created an
- # AssignmentStub in the db. If that doesn't exist, something bad
- # happened.
- logging.warning(f"No assignment found in DB: {event.amt_assignment_id}")
- emit_error_event(
- event_type="assignment_stub_not_found_in_db",
- amt_hit_type_id=event.amt_hit_type_id,
- )
- reject_assignment(
- amtm=amtm,
- am=am,
- hm=hm,
- amt_assignment_id=assignment.amt_assignment_id,
- msg=REJECT_MESSAGE_UNKNOWN_ASSIGNMENT,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- review_hit(amtm=amtm, hm=hm, assignment=assignment)
- return None
-
- assert assignment.amt_assignment_id == event.amt_assignment_id
- assert assignment.amt_hit_id == event.amt_hit_id
- assert assignment.amt_hit_id == stub.amt_hit_id
- assert assignment.amt_worker_id == stub.amt_worker_id
- amt_assignment_id = assignment.amt_assignment_id
- amt_worker_id = assignment.amt_worker_id
-
- # We don't have a TSID associated with the assignment until we the
- # assignment is submitted.
- am.update_answer(assignment=assignment)
-
- # check if the user is blocked by thl
- if get_user_blocked_or_not_exists(amt_worker_id=amt_worker_id):
- logging.warning(
- f"User {amt_worker_id} blocked or not exists. Rejecting: {amt_assignment_id}"
- )
- emit_error_event(
- event_type="assignment_submitted_user_blocked_or_not_exists",
- amt_hit_type_id=event.amt_hit_type_id,
- )
- reject_assignment(
- amtm=amtm,
- am=am,
- hm=hm,
- amt_assignment_id=amt_assignment_id,
- msg=REJECT_MESSAGE_BADDIE,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- review_hit(amtm=amtm, hm=hm, assignment=assignment)
- return None
-
- if assignment.tsid is None:
- assignment = handle_assignment_w_no_work(
- amtm=amtm, am=am, hm=hm, assignment=assignment
- )
- else:
- # We need to validate the work exists on thl, and if so, approve
- assignment = handle_assignment_w_work(
- amtm=amtm, am=am, hm=hm, assignment=assignment
- )
-
- #
- # Step 4: Tell Amazon we've reviewed the HIT, and update the DB
- #
+ reject_assignment(
+ amtm=amtm,
+ am=am,
+ hm=hm,
+ amt_assignment_id=assignment.amt_assignment_id,
+ msg=REJECT_MESSAGE_UNKNOWN_ASSIGNMENT,
+ amt_hit_type_id=hit.amt_hit_type_id,
+ )
review_hit(amtm=amtm, hm=hm, assignment=assignment)
-
- if (
- assignment.tsid
- and assignment.status == AssignmentStatus.Approved
- and assignment.requester_feedback != NO_WORK_APPROVAL_MESSAGE
- ):
- return issue_worker_payment(amtm=amtm, hm=hm, bm=bm, assignment=assignment)
+ return
def review_hit(amtm: AMTManager, hm: HitManager, assignment: Assignment) -> None:
@@ -163,66 +69,13 @@ 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}"
- )
- return None
+ LOG.warning(f"Hit not found when trying to review hit: {assignment.amt_hit_id}")
+ return
# Update the db
hm.update_hit(hit)
- return None
-
-
-def handle_assignment_w_no_work(
- amtm: AMTManager, am: AssignmentManager, hm: HitManager, assignment: Assignment
-) -> Assignment:
- """
- Called when an assignment is submitted without a wall event.
- Not entirely clear why this happens. I think they accept a HIT, get no work
- available for whatever reason, then report, and submit it.
-
- :return: The Assignment
- """
- logging.warning(
- f"Assignment submitted with no tsid: {assignment.amt_assignment_id}"
- )
- amt_worker_id = assignment.amt_worker_id
- amt_assignment_id = assignment.amt_assignment_id
- hit = hm.get_from_amt_id(amt_hit_id=assignment.amt_hit_id)
- emit_error_event(
- event_type="assignment_submitted_no_work",
- amt_hit_type_id=hit.amt_hit_type_id,
- )
-
- # They get 0 chances due to abuse
- if True:
- # if (am.missing_tsid_count(amt_worker_id=amt_worker_id) >= 3) or (
- # am.rejected_count(amt_worker_id=amt_worker_id) >= 3
- # or get_user_blocked(amt_worker_id=amt_worker_id)
- # ):
- assignment = reject_assignment(
- amtm=amtm,
- am=am,
- hm=hm,
- amt_assignment_id=amt_assignment_id,
- msg=REJECT_MESSAGE_NO_WORK,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- # todo: we don't have a way to "block" a user (i.e. tattle to thl)
- # make_block_worker_decision(user)
- return assignment
-
- # Approve with a message explaining they shouldn't do it.
- assignment = approve_assignment(
- amtm=amtm,
- am=am,
- amt_assignment_id=amt_assignment_id,
- msg=NO_WORK_APPROVAL_MESSAGE,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
-
- return assignment
+ return
def reject_assignment(
@@ -245,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)
@@ -256,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 ...
@@ -267,276 +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
-
-
-def approve_assignment(
- amtm: AMTManager,
- am: AssignmentManager,
- amt_assignment_id: str,
- msg: str,
- amt_hit_type_id: str,
- override_rejection: bool = False,
-) -> Assignment:
- # Approve in AMT, update db
-
- res = amtm.approve_assignment_if_possible(
- amt_assignment_id=amt_assignment_id,
- msg=msg,
- override_rejection=override_rejection,
- )
- if res is None:
- # We failed to approve this assignment. Cannot distinguish between
- # failed b/c assignment is already approved, or it is not possible.
- emit_error_event(
- event_type="failed_to_approve_assignment",
- amt_hit_type_id=amt_hit_type_id,
- )
- # The assignment might already be approved, the error msg is useless, so
- # keep going.
- # raise Exception(f"Failed to approve assignment: {amt_assignment_id}")
-
- # We just approved this assignment, get it from amazon again
- assignment = amtm.get_assignment(amt_assignment_id=amt_assignment_id)
- assert assignment.status == AssignmentStatus.Approved
- # And update the db
- am.approve(assignment=assignment)
- emit_assignment_event(
- status=AssignmentStatus.Approved, amt_hit_type_id=amt_hit_type_id, reason=msg
- )
- logging.warning(f"Approved assignment: {amt_assignment_id}")
- return assignment
-
-
-def handle_assignment_w_work(
- amtm: AMTManager, am: AssignmentManager, hm: HitManager, assignment: Assignment
-) -> Assignment:
- """
- Called when an assignment is submitted with a tsid.
- - Check the tsid (thl status endpoint). Make sure it is finished, and
- stuff matches (doesn't matter if not a complete)
- - Try to submit a cashout request for the HIT payout (e.g. 5c)
- """
-
- amt_worker_id = assignment.amt_worker_id
- amt_assignment_id = assignment.amt_assignment_id
- tsid = assignment.tsid
- assert (
- tsid is not None
- ), "Assignment must have a tsid to be handled in handle_assignment_w_work"
-
- hit = hm.get_from_amt_id(amt_hit_id=assignment.amt_hit_id)
-
- tsr = get_task_status(tsid=tsid)
- if (
- tsr is None
- or tsr.status is None
- or tsr.status_code_1
- in {
- StatusCode1.SESSION_START_FAIL,
- StatusCode1.SESSION_START_QUALITY_FAIL,
- StatusCode1.SESSION_CONTINUE_QUALITY_FAIL,
- }
- ):
- # TSID doesn't exist or work is not finished:
- # Reject the assignment instead
- if tsr is not None and tsr.status_code_1 in {
- StatusCode1.SESSION_START_QUALITY_FAIL,
- StatusCode1.SESSION_CONTINUE_QUALITY_FAIL,
- }:
- event_type = "assignment_submitted_quality_fail"
- else:
- event_type = "assignment_submitted_work_not_complete"
-
- emit_error_event(
- event_type=event_type,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- assignment = reject_assignment(
- amtm=amtm,
- am=am,
- hm=hm,
- amt_assignment_id=amt_assignment_id,
- msg=REJECT_MESSAGE_BADDIE,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- return assignment
-
- assert tsr.product_user_id == amt_worker_id
- assert tsr.finished is not None
- assert (tsr.finished - assignment.created_at) <= timedelta(minutes=90)
-
- # Request an AMT_ASSIGNMENT cashout for 1c
- req = submit_and_approve_amt_assignment_request(
- amt_worker_id=amt_worker_id, amount=hit.reward
- )
- if req is None:
- # Reject the assignment instead
- logging.warning(
- f"submit_and_approve_amt_assignment_request failed: {amt_assignment_id}"
- )
- emit_error_event(
- event_type="assignment_cashout_request_failed",
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- assignment = reject_assignment(
- amtm=amtm,
- am=am,
- hm=hm,
- amt_assignment_id=amt_assignment_id,
- msg=REJECT_MESSAGE_BADDIE,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- return assignment
-
- assert req.id
-
- # We've approved the HIT payment, now update the db to reflect this, and approve the assignment
- assignment = approve_assignment(
- amtm=amtm,
- am=am,
- amt_assignment_id=amt_assignment_id,
- msg=APPROVAL_MESSAGE,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- # We complete after the assignment is approved
- complete_res = manage_pending_cashout(
- cashout_id=req.id, payout_status=PayoutStatus.COMPLETE
- )
- if complete_res.status != PayoutStatus.COMPLETE:
- # unclear wny this would happen
- raise ValueError(f"Failed to complete cashout: {req.id}")
- return assignment
-
-
-def submit_and_approve_amt_assignment_request(
- amt_worker_id: str, amount: USDCent
-) -> Optional[CashoutRequestInfo]:
- # If successful, returns the cashout id, otherwise, returns None
- req = user_cashout_request(
- amt_worker_id=amt_worker_id,
- amount=amount,
- cashout_method_id=settings.amt_assignment_cashout_method,
- )
- assert req.id
-
- if req.status != PayoutStatus.PENDING:
- return None
-
- approve_res = manage_pending_cashout(req.id, PayoutStatus.APPROVED)
- if approve_res.status != PayoutStatus.APPROVED:
- return None
-
- return req
-
-
-def submit_and_approve_amt_bonus_request(
- amt_worker_id: str, amount: USDCent
-) -> Optional[CashoutRequestInfo]:
- # If successful, returns the cashout id, otherwise, returns None
- req = user_cashout_request(
- amt_worker_id=amt_worker_id,
- amount=amount,
- cashout_method_id=settings.amt_bonus_cashout_method,
- )
- assert req.id
-
- if req.status != PayoutStatus.PENDING:
- return None
-
- approve_res = manage_pending_cashout(req.id, PayoutStatus.APPROVED)
- if approve_res.status != PayoutStatus.APPROVED:
- return None
-
- return req
-
-
-def issue_worker_payment(
- amtm: AMTManager, hm: HitManager, bm: BonusManager, assignment: Assignment
-) -> None:
- # For now, since we have no "I want my bonus" request/button. A user's
- # balance will be sent out anytime they get an approved assignment. We
- # don't need the task status, the tsid, nor the amount / user_payout
- # that was paid, or anything.
- # We just get the wallet balance and submit a cashout request if >0
- # then approve it, send the amt bonus, then complete it
- amt_assignment_id = assignment.amt_assignment_id
- hit = hm.get_from_amt_id(amt_hit_id=assignment.amt_hit_id)
- wallet_balance = get_wallet_balance_if_non_negative(
- amt_worker_id=assignment.amt_worker_id
- )
- if not wallet_balance:
- return None
- amount = round_payment(amount=wallet_balance)
- if not amount:
- return None
-
- # Don't send more than $4.97 at a time. If they have a higher wallet balance,
- # they just need to get another hit approved.
- amount = min(amount, USDCent(4_97))
-
- pe = submit_and_approve_amt_bonus_request(
- amt_worker_id=assignment.amt_worker_id, amount=amount
- )
- if pe is None:
- logging.warning(
- f"submit_and_approve_amt_bonus_request failed: {amt_assignment_id}"
- )
- emit_error_event(
- event_type="bonus_cashout_request_failed",
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- return None
- assert pe.id
-
- amtm.send_bonus(
- amt_worker_id=assignment.amt_worker_id,
- amt_assignment_id=assignment.amt_assignment_id,
- amount=amount,
- reason=BONUS_MESSAGE,
- unique_request_token=pe.id,
- )
-
- # Confirm it was sent through amt
- bonus = amtm.get_bonus(
- amt_assignment_id=assignment.amt_assignment_id, payout_event_id=pe.id
- )
-
- if bonus is None:
- logging.warning(
- f"Failed to find bonus after sending it: {amt_assignment_id} {pe.id}"
- )
- emit_error_event(
- event_type="bonus_not_found_after_sending",
- amt_hit_type_id=hit.amt_hit_type_id,
- )
- return None
-
- # Create in DB
- bm.create(bonus=bonus)
- emit_bonus_event(amount=amount, amt_hit_type_id=hit.amt_hit_type_id)
-
- # Complete cashout
- res = manage_pending_cashout(pe.id, PayoutStatus.COMPLETE)
- if res.status != PayoutStatus.COMPLETE:
- raise ValueError(
- f"{assignment.amt_assignment_id} {pe.id=} manage_pending_cashout COMPLETE failed: {res=}"
- )
-
-
-def round_payment(amount: USDCent) -> USDCent:
- """
- Don't pay bonuses less than 7 cents, just add it to their wallet.
- Round down bonuses (>=7 cents) to the nearest multiple of 5
- starting at 2
- """
- if amount < 7:
- return USDCent(0)
-
- amt = (5 * math.floor((int(amount) - 2) / 5)) + 2
-
- payout = USDCent(amt)
- assert 0 <= payout <= 40_00, "Payout must be between $0.00 and $40.00"
-
- return payout
diff --git a/jb/flow/events.py b/jb/flow/events.py
index 2825cb1..bee63fd 100644
--- a/jb/flow/events.py
+++ b/jb/flow/events.py
@@ -1,18 +1,16 @@
-import logging
import time
from concurrent import futures
-from concurrent.futures import ThreadPoolExecutor, Executor
-from typing import Optional, cast, TypedDict
+from concurrent.futures import Executor, ThreadPoolExecutor
+from typing import cast
import redis
from jb.config import (
- JB_EVENTS_STREAM,
CONSUMER_GROUP,
CONSUMER_NAME,
- JB_EVENTS_FAILED_STREAM,
+ JB_EVENTS_STREAM,
)
-from jb.decorators import REDIS
+from jb.decorators import LOG, get_redis
from jb.flow.assignment_tasks import process_assignment_submitted
from jb.flow.monitoring import emit_error_event
from jb.models.event import MTurkEvent
@@ -20,13 +18,6 @@ from jb.models.event import MTurkEvent
StreamMessages = list[tuple[str, list[tuple[bytes, dict[bytes, bytes]]]]]
-class PendingEntry(TypedDict):
- message_id: bytes
- consumer: bytes
- time_since_delivered: int
- times_delivered: int
-
-
def process_mturk_events_task():
executor = ThreadPoolExecutor(max_workers=5)
create_consumer_group()
@@ -34,21 +25,11 @@ 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)
-def handle_pending_msgs_task():
- while True:
- try:
- handle_pending_msgs()
- except Exception as e:
- logging.exception(e)
- finally:
- time.sleep(60)
-
-
def process_mturk_events(executor: Executor):
while True:
n = process_mturk_events_chunk(executor=executor)
@@ -58,7 +39,10 @@ def process_mturk_events(executor: Executor):
def create_consumer_group():
try:
- REDIS.xgroup_create(JB_EVENTS_STREAM, CONSUMER_GROUP, id="0", mkstream=True)
+ redis_client = get_redis()
+ redis_client.xgroup_create(
+ JB_EVENTS_STREAM, CONSUMER_GROUP, id="0", mkstream=True
+ )
except redis.exceptions.ResponseError as e:
if "BUSYGROUP Consumer Group name already exists" in str(e):
pass # group already exists
@@ -66,8 +50,9 @@ def create_consumer_group():
raise
-def process_mturk_events_chunk(executor: Executor) -> Optional[int]:
- msgs_raw = REDIS.xreadgroup(
+def process_mturk_events_chunk(executor: Executor) -> int | None:
+ redis_client = get_redis()
+ msgs_raw = redis_client.xreadgroup(
groupname=CONSUMER_GROUP,
consumername=CONSUMER_NAME,
streams={JB_EVENTS_STREAM: ">"},
@@ -89,66 +74,25 @@ def process_mturk_events_chunk(executor: Executor) -> Optional[int]:
executor.submit(process_assignment_submitted_event, event, str(msg_id))
)
else:
- logging.info(f"Discarding {event}")
- REDIS.xdel(JB_EVENTS_STREAM, msg_id)
+ LOG.info(f"Discarding {event}")
+ redis_client.xdel(JB_EVENTS_STREAM, msg_id)
futures.wait(fs, timeout=60)
return len(msgs)
def process_assignment_submitted_event(event: MTurkEvent, msg_id: str):
- from jb.decorators import AMTM, AM, HM, BM
+ from jb.decorators import AM, AMTM, BM, HM
+
+ redis_client = get_redis()
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,
)
- REDIS.xackdel(JB_EVENTS_STREAM, CONSUMER_GROUP, msg_id)
-
-
-def handle_pending_msgs():
- # TODO!: This doesn't run at all.
-
- # Looks in the redis queue for msgs that
- # are pending (read by a consumer but not ACK). These prob failed.
- # Below is from chatgpt, idk if it works
- pending = cast(
- list[PendingEntry],
- REDIS.xpending_range(
- JB_EVENTS_STREAM, CONSUMER_GROUP, min="-", max="+", count=10
- ),
- )
-
- for entry in pending:
- msg_id = entry["message_id"]
- # Claim message if idle > 10 sec
- if entry["idle"] > 10_000: # milliseconds
- claimed = REDIS.xclaim(
- JB_EVENTS_STREAM,
- CONSUMER_GROUP,
- CONSUMER_NAME,
- min_idle_time=10_000,
- message_ids=[msg_id],
- )
- for cid, data in claimed:
- msg_json = data["data"]
- event = MTurkEvent.model_validate_json(msg_json)
- if event.event_type == "AssignmentSubmitted":
- # Try to process it again. If it fails, add
- # it to the failed stream, so maybe we can fix
- # and try again?
- try:
- process_assignment_submitted_event(event, cid)
- REDIS.xack(JB_EVENTS_STREAM, CONSUMER_GROUP, cid)
- except Exception as e:
- logging.exception(e)
- REDIS.xadd(JB_EVENTS_FAILED_STREAM, data)
- REDIS.xack(JB_EVENTS_STREAM, CONSUMER_GROUP, cid)
- else:
- logging.info(f"Discarding {event}")
- REDIS.xdel(JB_EVENTS_STREAM, msg_id)
+ redis_client.xackdel(JB_EVENTS_STREAM, CONSUMER_GROUP, msg_id)
diff --git a/jb/flow/maintenance.py b/jb/flow/maintenance.py
index f8ca971..4a2fe81 100644
--- a/jb/flow/maintenance.py
+++ b/jb/flow/maintenance.py
@@ -1,5 +1,3 @@
-from typing import Optional
-
from jb.decorators import HM
from jb.flow.monitoring import emit_hit_event
from jb.managers.amt import AMTManager
@@ -10,7 +8,7 @@ def check_hit_status(
amtm: AMTManager,
amt_hit_id: str,
amt_hit_type_id: str,
- reason: Optional[str] = None,
+ reason: str | None = None,
) -> HitStatus:
"""
(this used to be called "process_hit")
diff --git a/jb/flow/monitoring.py b/jb/flow/monitoring.py
index 28f7271..22e9e93 100644
--- a/jb/flow/monitoring.py
+++ b/jb/flow/monitoring.py
@@ -1,12 +1,11 @@
import socket
-from typing import Optional
+from generalresearch.currency import USDCent
from mypy_boto3_mturk.literals import EventTypeType
from jb.config import settings
from jb.decorators import influx_client
-from generalresearchutils.currency import USDCent
-from jb.models.definitions import HitStatus, AssignmentStatus
+from jb.models.definitions import AssignmentStatus, HitStatus
def write_hit_gauge(status: HitStatus, amt_hit_type_id: str, cnt: int) -> None:
@@ -26,8 +25,6 @@ def write_hit_gauge(status: HitStatus, amt_hit_type_id: str, cnt: int) -> None:
if influx_client:
influx_client.write_points(points=[point])
- return None
-
def write_assignment_gauge(
status: AssignmentStatus, amt_hit_type_id: str, cnt: int
@@ -47,11 +44,9 @@ def write_assignment_gauge(
if influx_client:
influx_client.write_points(points=[point])
- return None
-
def emit_hit_event(
- status: HitStatus, amt_hit_type_id: str, reason: Optional[str] = None
+ status: HitStatus, amt_hit_type_id: str, reason: str | None = None
) -> None:
"""
e.g. a HIT was created, Reviewable, etc. We don't have a "created"
@@ -77,11 +72,9 @@ def emit_hit_event(
if influx_client:
influx_client.write_points([point])
- return None
-
def emit_assignment_event(
- status: AssignmentStatus, amt_hit_type_id: str, reason: Optional[str] = None
+ status: AssignmentStatus, amt_hit_type_id: str, reason: str | None = None
) -> None:
"""
e.g. an Assignment was accepted/approved/reject
@@ -106,8 +99,6 @@ def emit_assignment_event(
if influx_client:
influx_client.write_points([point])
- return None
-
def emit_mturk_notification_event(
event_type: EventTypeType, amt_hit_type_id: str
@@ -133,8 +124,6 @@ def emit_mturk_notification_event(
if influx_client:
influx_client.write_points([point])
- return None
-
def emit_error_event(event_type: str, amt_hit_type_id: str) -> None:
"""
@@ -157,8 +146,6 @@ def emit_error_event(event_type: str, amt_hit_type_id: str) -> None:
if influx_client:
influx_client.write_points([point])
- return None
-
def emit_bonus_event(amount: USDCent, amt_hit_type_id: str) -> None:
"""
@@ -179,5 +166,3 @@ def emit_bonus_event(amount: USDCent, amt_hit_type_id: str) -> None:
if influx_client:
influx_client.write_points([point])
-
- return None
diff --git a/jb/flow/setup_tasks.py b/jb/flow/setup_tasks.py
index 4664374..f5cda48 100644
--- a/jb/flow/setup_tasks.py
+++ b/jb/flow/setup_tasks.py
@@ -1,6 +1,5 @@
-from jb.config import TOPIC_ARN, SUBSCRIPTION
-from jb.decorators import SNS_CLIENT, AMT_CLIENT
-from jb.config import settings
+from jb.config import SUBSCRIPTION, TOPIC_ARN, settings
+from jb.decorators import AMT_CLIENT, SNS_CLIENT
def initial_setup():
diff --git a/jb/flow/tasks.py b/jb/flow/tasks.py
index 77825d3..c555021 100644
--- a/jb/flow/tasks.py
+++ b/jb/flow/tasks.py
@@ -1,18 +1,13 @@
-import logging
import time
from typing import TypedDict, cast
-from generalresearchutils.config import is_debug
+from generalresearch.config import is_debug
-from jb.decorators import AMTM, HTM, HM, HQM, 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 write_hit_gauge, emit_hit_event
+from jb.flow.monitoring import emit_hit_event, write_hit_gauge
from jb.models.definitions import HitStatus
-from jb.models.hit import HitType, HitQuestion, Hit
-
-logging.basicConfig()
-logger = logging.getLogger()
-logger.setLevel(logging.INFO)
+from jb.models.hit import Hit, HitQuestion, HitType
class HitRow(TypedDict):
@@ -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 fa59167..e30bb38 100644
--- a/jb/main.py
+++ b/jb/main.py
@@ -1,14 +1,15 @@
from multiprocessing import Process
-from typing import Any, Dict
+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.views.common import common_router
-from jb.settings import BASE_HTML
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
app = FastAPI(
servers=[
@@ -21,6 +22,7 @@ app = FastAPI(
version="1.0.0",
)
+
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
@@ -31,12 +33,13 @@ app.add_middleware(
app.add_middleware(TrustedHostMiddleware, allowed_hosts=["*"])
app.include_router(router=common_router)
+app.include_router(router=auth_router)
@app.get("/robots.txt")
@app.get("/sitemap.xml")
@app.get("/favicon.ico")
-def return_nothing() -> Dict[str, Any]:
+def return_nothing() -> dict[str, Any]:
return {}
@@ -52,7 +55,6 @@ def schedule_tasks():
from jb.flow.tasks import refill_hits_task
Process(target=process_mturk_events_task).start()
- # Process(target=handle_pending_msgs_task).start()
Process(target=refill_hits_task).start()
diff --git a/jb/managers/__init__.py b/jb/managers/__init__.py
index e99569a..e69de29 100644
--- a/jb/managers/__init__.py
+++ b/jb/managers/__init__.py
@@ -1,23 +0,0 @@
-from enum import IntEnum
-from typing import Collection
-
-from generalresearchutils.pg_helper import PostgresConfig
-
-
-class Permission(IntEnum):
- READ = 1
- UPDATE = 2
- CREATE = 3
- DELETE = 4
-
-
-class PostgresManager:
- def __init__(
- self,
- pg_config: PostgresConfig,
- permissions: Collection[Permission] = None, # type: ignore
- **kwargs, # type: ignore
- ):
- super().__init__(**kwargs)
- self.pg_config = pg_config
- self.permissions = set(permissions) if permissions else set()
diff --git a/jb/managers/amt.py b/jb/managers/amt.py
index 88ae07d..17e5630 100644
--- a/jb/managers/amt.py
+++ b/jb/managers/amt.py
@@ -1,25 +1,24 @@
-import logging
-from datetime import timezone, datetime
-from typing import Tuple, Optional, List, Dict, Any
+from datetime import datetime, timezone
+from typing import Any
import botocore.exceptions
-from generalresearchutils.currency import USDCent
+from generalresearch.currency import USDCent
from mypy_boto3_mturk import MTurkClient
from mypy_boto3_mturk.type_defs import (
AssignmentTypeDef,
BonusPaymentTypeDef,
CreateHITTypeResponseTypeDef,
- GetHITResponseTypeDef,
CreateHITWithHITTypeResponseTypeDef,
+ GetHITResponseTypeDef,
)
from pydantic import ValidationError
from jb.config import TOPIC_ARN
-from jb.models import AMTAccount
+from jb.models.amt import AMTAccount
from jb.models.assignment import Assignment
from jb.models.bonus import Bonus
from jb.models.definitions import HitStatus
-from jb.models.hit import HitType, HitQuestion, Hit
+from jb.models.hit import Hit, HitQuestion, HitType
REJECT_MESSAGE_UNKNOWN_ASSIGNMENT = "Unknown assignment"
REJECT_MESSAGE_NO_WORK = "Assignment was submitted with no attempted work."
@@ -56,7 +55,7 @@ class AMTManager:
}
)
- def get_hit_if_exists(self, amt_hit_id: str) -> Tuple[Optional[Hit], Optional[str]]:
+ def get_hit_if_exists(self, amt_hit_id: str) -> tuple[Hit | None, str | None]:
try:
res: GetHITResponseTypeDef = self.amt_client.get_hit(HITId=amt_hit_id)
@@ -78,7 +77,7 @@ class AMTManager:
return HitStatus.Disposed
else:
- logging.warning(msg)
+ LOG.warning(msg)
return HitStatus.Unassignable
return res.status
@@ -137,7 +136,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)
@@ -146,13 +145,13 @@ class AMTManager:
assert assignment.id is None
return assignment
- def get_assignment_if_exists(self, amt_assignment_id: str) -> Optional[Assignment]:
+ def get_assignment_if_exists(self, amt_assignment_id: str) -> Assignment | None:
expected_err_msg = f"Assignment {amt_assignment_id} does not exist"
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:
@@ -161,7 +160,7 @@ class AMTManager:
def reject_assignment_if_possible(
self, amt_assignment_id: str, msg: str = REJECT_MESSAGE_UNKNOWN_ASSIGNMENT
- ) -> Optional[Dict[str, Any]]:
+ ) -> dict[str, Any] | None:
# Unclear to me when this would fail
try:
@@ -170,7 +169,7 @@ class AMTManager:
)
except botocore.exceptions.ClientError as e:
- logging.warning(e)
+ LOG.warning(e)
return None
def approve_assignment_if_possible(
@@ -178,7 +177,7 @@ class AMTManager:
amt_assignment_id: str,
msg: str = APPROVAL_MESSAGE,
override_rejection: bool = False,
- ) -> Optional[Dict[str, Any]]:
+ ) -> dict[str, Any] | None:
# Unclear to me when this would fail
try:
@@ -189,7 +188,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 +197,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:
@@ -207,8 +206,6 @@ class AMTManager:
# elif "This HIT is currently in the state 'Reviewing'" in error_msg:
# logging.warning(error_msg)
- return None
-
def send_bonus(
self,
amt_worker_id: str,
@@ -216,7 +213,7 @@ class AMTManager:
amt_assignment_id: str,
reason: str,
unique_request_token: str,
- ) -> Optional[Dict[str, Any]]:
+ ) -> dict[str, Any] | None:
try:
return self.amt_client.send_bonus(
WorkerId=amt_worker_id,
@@ -227,14 +224,12 @@ 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
- ) -> Optional[Bonus]:
+ def get_bonus(self, amt_assignment_id: str, payout_event_id: str) -> Bonus | None:
- res: List[BonusPaymentTypeDef] = self.amt_client.list_bonus_payments(
+ res: list[BonusPaymentTypeDef] = self.amt_client.list_bonus_payments(
AssignmentId=amt_assignment_id
)["BonusPayments"]
@@ -268,5 +263,3 @@ class AMTManager:
self.amt_client.update_expiration_for_hit(
HITId=hit["HITId"], ExpireAt=now
)
-
- return None
diff --git a/jb/managers/assignment.py b/jb/managers/assignment.py
index f6aa2ce..2740aee 100644
--- a/jb/managers/assignment.py
+++ b/jb/managers/assignment.py
@@ -1,11 +1,10 @@
from datetime import datetime, timezone
-from typing import Optional
from psycopg import sql
from pydantic import NonNegativeInt, PositiveInt
-from jb.managers import PostgresManager
-from jb.models.assignment import AssignmentStub, Assignment
+from jb.managers.base import PostgresManager
+from jb.models.assignment import Assignment, AssignmentStub
from jb.models.definitions import AssignmentStatus
@@ -14,8 +13,7 @@ class AssignmentManager(PostgresManager):
def create_stub(self, stub: AssignmentStub) -> None:
assert stub.id is None
data = stub.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_assignment
(amt_assignment_id, amt_worker_id, status,
created_at, modified_at, hit_id)
@@ -23,16 +21,13 @@ class AssignmentManager(PostgresManager):
(%(amt_assignment_id)s, %(amt_worker_id)s, %(status)s,
%(created_at)s, %(modified_at)s, %(hit_id)s)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
stub.id = pk
- return None
def create(self, assignment: Assignment) -> None:
# Typically this is NOT used (we'd create the stub when HIT is
@@ -42,8 +37,7 @@ class AssignmentManager(PostgresManager):
assert assignment.id is None
data = assignment.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_assignment
(amt_assignment_id, amt_worker_id, status,
created_at, modified_at, hit_id,
@@ -57,16 +51,13 @@ class AssignmentManager(PostgresManager):
%(approval_time)s, %(rejection_time)s, %(requester_feedback)s,
%(tsid)s)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
assignment.id = pk
- return None
def get_stub(self, amt_assignment_id: str) -> AssignmentStub:
res = self.pg_config.execute_sql_query(
@@ -85,7 +76,7 @@ class AssignmentManager(PostgresManager):
assert len(res) == 1
return AssignmentStub.model_validate(res[0])
- def get_stub_if_exists(self, amt_assignment_id: str) -> Optional[AssignmentStub]:
+ def get_stub_if_exists(self, amt_assignment_id: str) -> AssignmentStub | None:
try:
return self.get_stub(amt_assignment_id=amt_assignment_id)
except AssertionError:
@@ -121,8 +112,7 @@ class AssignmentManager(PostgresManager):
"amt_assignment_id": assignment.amt_assignment_id,
"modified_at": now,
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_assignment
SET submit_time = %(submit_time)s,
auto_approval_time = %(auto_approval_time)s,
@@ -130,15 +120,12 @@ class AssignmentManager(PostgresManager):
tsid = %(tsid)s,
modified_at = %(modified_at)s
WHERE amt_assignment_id = %(amt_assignment_id)s
- """
- )
+ """)
# We force this to fail if the assignment doesn't already exist in the db
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
- conn.commit()
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
+ conn.commit()
def reject(self, assignment: Assignment) -> None:
assert assignment.status == AssignmentStatus.Rejected
@@ -156,8 +143,7 @@ class AssignmentManager(PostgresManager):
"accept_time": assignment.accept_time,
"modified_at": now,
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_assignment
SET submit_time = %(submit_time)s,
rejection_time = %(rejection_time)s,
@@ -167,15 +153,12 @@ class AssignmentManager(PostgresManager):
accept_time = %(accept_time)s,
modified_at = %(modified_at)s
WHERE amt_assignment_id = %(amt_assignment_id)s
- """
- )
+ """)
# We force this to fail if the assignment doesn't already exist in the db
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
- conn.commit()
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
+ conn.commit()
def approve(self, assignment: Assignment) -> None:
assert assignment.status == AssignmentStatus.Approved
@@ -194,8 +177,7 @@ class AssignmentManager(PostgresManager):
"accept_time": assignment.accept_time,
"modified_at": now,
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_assignment
SET submit_time = %(submit_time)s,
approval_time = %(approval_time)s,
@@ -205,15 +187,12 @@ class AssignmentManager(PostgresManager):
accept_time = %(accept_time)s,
modified_at = %(modified_at)s
WHERE amt_assignment_id = %(amt_assignment_id)s
- """
- )
+ """)
# We force this to fail if the assignment doesn't already exist in the db
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
- conn.commit()
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
+ conn.commit()
def missing_tsid_count(
self, amt_worker_id: str, lookback_hrs: PositiveInt = 24
diff --git a/jb/managers/base.py b/jb/managers/base.py
new file mode 100644
index 0000000..4b0637f
--- /dev/null
+++ b/jb/managers/base.py
@@ -0,0 +1,16 @@
+from collections.abc import Collection
+
+from generalresearch.managers.base import Permission
+from generalresearch.pg_helper import PostgresConfig
+
+
+class PostgresManager:
+ def __init__(
+ self,
+ pg_config: PostgresConfig,
+ permissions: Collection[Permission] | None = None,
+ **kwargs, # type: ignore
+ ):
+ super().__init__(**kwargs)
+ self.pg_config = pg_config
+ self.permissions = set(permissions) if permissions else set()
diff --git a/jb/managers/bonus.py b/jb/managers/bonus.py
index 89b81f0..b649103 100644
--- a/jb/managers/bonus.py
+++ b/jb/managers/bonus.py
@@ -1,8 +1,8 @@
-from typing import List, Any
+from typing import Any
from psycopg import sql
-from jb.managers import PostgresManager
+from jb.managers.base import PostgresManager
from jb.models.bonus import Bonus
@@ -11,8 +11,7 @@ class BonusManager(PostgresManager):
def create(self, bonus: Bonus) -> None:
assert bonus.id is None
data = bonus.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_bonus
(payout_event_id, amt_worker_id, amount, grant_time, assignment_id, reason)
VALUES (
@@ -29,20 +28,17 @@ class BonusManager(PostgresManager):
%(reason)s
)
RETURNING id, assignment_id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- res = c.fetchone()
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ res = c.fetchone()
+ conn.commit()
bonus.id = res["id"] # type: ignore
bonus.assignment_id = res["assignment_id"] # type: ignore
- return None
- def filter(self, amt_assignment_id: str) -> List[Bonus]:
- res: List[Any] = self.pg_config.execute_sql_query(
+ def filter(self, amt_assignment_id: str) -> list[Bonus]:
+ res: list[Any] = self.pg_config.execute_sql_query(
"""
SELECT mb.*, ma.amt_assignment_id
FROM mtwerk_bonus mb
diff --git a/jb/managers/email_manager.py b/jb/managers/email_manager.py
new file mode 100644
index 0000000..7298c60
--- /dev/null
+++ b/jb/managers/email_manager.py
@@ -0,0 +1,65 @@
+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()}"}
+
+
+def get_or_create_contact(email: str, amt_worker_id: str | None = None) -> int:
+ body = {"email": email}
+ if amt_worker_id:
+ body["amt_worker_id"] = amt_worker_id
+ res = requests.post(
+ url=f"{MAUTIC_BASE_URL}/api/contacts/new",
+ json=body,
+ headers=auth_headers,
+ ).json()
+ return int(res["contact"]["id"])
+
+
+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,
+ }
+ body = {"tokens": email_tokens}
+ response = requests.post(url=mautic_url, json=body, headers=auth_headers)
+
+ try:
+ response.raise_for_status()
+ except requests.exceptions.HTTPError:
+ print(f"Failed to send email. Status code: {response.status_code}")
+ print(response.text)
+ raise
+ d = response.json()
+
+ assert d.get("success"), f"Failed to send email: {d.get('failed')}"
+ print("Email sent successfully")
+
+
+def send_login_email(email: str, magic_token: str) -> None:
+ contact_id = get_or_create_contact(email=email)
+ mautic_url = (
+ f"{MAUTIC_BASE_URL}/api/emails/{EMAIL_TEMPLATE_ID}/contact/{contact_id}/send"
+ )
+ magic_link = f"{settings.base_url}auth/magic-link/?token={magic_token}"
+ send_login_email_from_url(mautic_url, magic_link)
+
+
+def send_amt_link_email(email: str, magic_token: str) -> None:
+ # Don't actually associate the email with the worker ID
+ # until they click the link
+ contact_id = get_or_create_contact(email=email)
+ mautic_url = (
+ f"{MAUTIC_BASE_URL}/api/emails/{EMAIL_TEMPLATE_ID}/contact/{contact_id}/send"
+ )
+ magic_link = f"{settings.base_url}auth/link-amt/?token={magic_token}"
+ send_login_email_from_url(mautic_url, magic_link)
diff --git a/jb/managers/gr_api.py b/jb/managers/gr_api.py
new file mode 100644
index 0000000..f0be1a7
--- /dev/null
+++ b/jb/managers/gr_api.py
@@ -0,0 +1,160 @@
+"""Client for General Research's product-user API."""
+
+import logging
+from typing import Any
+
+import requests
+
+from jb.models.auth import User
+
+logger = logging.getLogger("amtjb")
+
+
+class GRApiError(RuntimeError):
+ """The General Research API could not satisfy a request."""
+
+
+class GRApiNotFoundError(GRApiError):
+ """The requested General Research API resource does not exist."""
+
+
+class GRApiManager:
+ def __init__(
+ self,
+ base_url: str,
+ token: str,
+ product_id: str,
+ timeout: float = 5.0,
+ session: requests.Session | None = None,
+ ) -> None:
+ if not token:
+ raise ValueError("General Research API token must not be empty")
+
+ self.base_url = base_url.rstrip("/")
+ self.product_id = product_id
+ self.timeout = timeout
+ self.session = session or requests.Session()
+ self.session.headers.update(
+ {
+ "Authorization": token,
+ "Accept": "application/json",
+ "Content-Type": "application/json",
+ }
+ )
+
+ def _request(self, method: str, url: str, **kwargs: Any) -> requests.Response:
+ try:
+ response = self.session.request(method, url, timeout=self.timeout, **kwargs)
+ response.raise_for_status()
+ return response
+ except requests.HTTPError as exc:
+ if exc.response is not None and exc.response.status_code == 404:
+ raise GRApiNotFoundError(
+ f"General Research API resource not found: {method} {url}"
+ ) from exc
+ raise GRApiError(
+ f"General Research API request failed: {method} {url}"
+ ) from exc
+ except requests.RequestException as exc:
+ raise GRApiError(
+ f"General Research API request failed: {method} {url}"
+ ) from exc
+
+ def _parse_user_response(self, res: dict[str, Any]) -> User:
+ return User.model_validate(
+ {
+ "product_user_id": res["product_user_id"],
+ "email": res["metadata"].get("email_address"),
+ "display_name": res["metadata"].get("display_name"),
+ "blocked": res["blocked"],
+ }
+ )
+
+ def ensure_user_exists(self, user: User) -> User:
+ """Idempotently create the product user if it does not already exist."""
+ url = f"{self.base_url}/{self.product_id}/user/{user.product_user_id}/"
+ res = self._request("PUT", url).json()
+ if res["metadata"].get("email_address") is None:
+ self.set_user_email(user)
+ return self.get_user(user.product_user_id)
+ user_thl = self._parse_user_response(res)
+ if user_thl.email != user.email:
+ raise ValueError(
+ f"user {user.product_user_id} already exists with email {user_thl.email}"
+ )
+ return user_thl
+
+ def get_user(self, product_user_id: str) -> User:
+ """Retrieve the user's email address and display name."""
+ url = f"{self.base_url}/{self.product_id}/user/{product_user_id}/"
+ res = self._request("GET", url).json()
+ return self._parse_user_response(res)
+
+ def get_user_by_email(self, email: str) -> User:
+ user = User.model_validate({"email": email})
+ return self.get_user(user.product_user_id)
+
+ def set_user_email(self, user: User) -> None:
+ """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/"
+ _ = self._request(
+ "PATCH",
+ url,
+ json={"email_address": str(user.email)},
+ ).json()
+
+ def set_user_display_name(self, user: User) -> User:
+ """Can be called as many times as needed. Multiple
+ users can have the same display name."""
+ if user.display_name is None:
+ return user
+ url = f"{self.base_url}/{self.product_id}/user/{user.product_user_id}/metadata/"
+ _ = self._request(
+ "PATCH",
+ url,
+ json={"display_name": user.display_name},
+ ).json()
+ return self.get_user(user.product_user_id)
+
+ def transition_user_from_amt(self, user: User, amt_worker_id: str) -> User:
+ """This should only be called once upon transition from an
+ AMT account to a General Research account."""
+ url = f"{self.base_url}/{self.product_id}/user/{amt_worker_id}/"
+ try:
+ self._request("PATCH", url, json={"product_user_id": user.product_user_id})
+ except GRApiNotFoundError as exc:
+ raise ValueError(f"User {amt_worker_id} does not exist") from exc
+ except GRApiError as exc:
+ http_error = exc.__cause__
+ response = (
+ http_error.response
+ if isinstance(http_error, requests.HTTPError)
+ else None
+ )
+ if response is not None and response.status_code == 400:
+ try:
+ detail = response.json().get("detail")
+ except (ValueError, AttributeError):
+ detail = None
+ if detail == "Unable to update User":
+ raise ValueError(
+ "unable to update user, probably another user already "
+ "exists with this email"
+ ) from exc
+ raise
+ self.set_user_email(user)
+ transitioned_user = self.get_user(user.product_user_id)
+ logger.warning(
+ "Transitioned product user from AMT worker %s to %s with email %s",
+ amt_worker_id,
+ transitioned_user.product_user_id,
+ transitioned_user.email,
+ extra={
+ "event": "transition_user_from_amt",
+ "amt_worker_id": amt_worker_id,
+ "product_user_id": transitioned_user.product_user_id,
+ "email": str(transitioned_user.email),
+ },
+ )
+ return transitioned_user
diff --git a/jb/managers/hit.py b/jb/managers/hit.py
index 63af2d4..bbae92b 100644
--- a/jb/managers/hit.py
+++ b/jb/managers/hit.py
@@ -1,11 +1,10 @@
from datetime import datetime, timezone
-from typing import Optional, List
from psycopg import sql
-from jb.managers import PostgresManager
+from jb.managers.base import PostgresManager
from jb.models.definitions import HitStatus
-from jb.models.hit import HitQuestion, HitType, Hit
+from jb.models.hit import Hit, HitQuestion, HitType
class HitQuestionManager(PostgresManager):
@@ -13,21 +12,17 @@ class HitQuestionManager(PostgresManager):
def create(self, question: HitQuestion) -> None:
assert question.id is None
data = question.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_question (url, height)
VALUES (%(url)s, %(height)s)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
question.id = pk
- return None
def get_by_id(self, question_id: int) -> HitQuestion:
res = self.pg_config.execute_sql_query(
@@ -42,7 +37,7 @@ class HitQuestionManager(PostgresManager):
assert len(res) == 1
return HitQuestion.model_validate(res[0])
- def get_by_values_if_exists(self, url: str, height: int) -> Optional[HitQuestion]:
+ def get_by_values_if_exists(self, url: str, height: int) -> HitQuestion | None:
res = self.pg_config.execute_sql_query(
"""
SELECT *
@@ -73,8 +68,7 @@ class HitTypeManager(PostgresManager):
assert hit_type.amt_hit_type_id is not None
data = hit_type.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_hittype (
amt_hit_type_id,
title,
@@ -96,27 +90,21 @@ class HitTypeManager(PostgresManager):
%(min_active)s
)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
hit_type.id = pk
- return None
-
- def filter_active(self) -> List[HitType]:
- res = self.pg_config.execute_sql_query(
- """
+ def filter_active(self) -> list[HitType]:
+ res = self.pg_config.execute_sql_query("""
SELECT *
FROM mtwerk_hittype
WHERE min_active > 0
LIMIT 50
- """
- )
+ """)
if len(res) == 50:
raise ValueError("Too many HitTypes!")
@@ -135,7 +123,7 @@ class HitTypeManager(PostgresManager):
assert len(res) == 1
return HitType.from_postgres(res[0])
- def get_if_exists(self, amt_hit_type_id: str) -> Optional[HitType]:
+ def get_if_exists(self, amt_hit_type_id: str) -> HitType | None:
try:
return self.get(amt_hit_type_id=amt_hit_type_id)
except AssertionError:
@@ -151,22 +139,18 @@ class HitTypeManager(PostgresManager):
def set_min_active(self, hit_type: HitType) -> None:
assert hit_type.id, "must be in the db first!"
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_hittype
SET min_active = %(min_active)s
WHERE id = %(id)s
- """
- )
+ """)
data = {"id": hit_type.id, "min_active": hit_type.min_active}
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- conn.commit()
- row_cnt = c.rowcount
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ conn.commit()
+ row_cnt = c.rowcount
assert row_cnt == 1, f"Expected 1 row updated, got {row_cnt}"
- return None
class HitManager(PostgresManager):
@@ -175,8 +159,7 @@ class HitManager(PostgresManager):
assert hit.amt_hit_id is not None
assert hit.id is None
data = hit.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_hit (
amt_hit_id,
hit_type_id,
@@ -210,38 +193,32 @@ class HitManager(PostgresManager):
%(assignment_available_count)s
)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
hit.id = pk
return hit
def update_status(self, amt_hit_id: str, hit_status: HitStatus):
now = datetime.now(tz=timezone.utc)
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_hit
SET status = %(status)s, modified_at = %(modified_at)s
WHERE amt_hit_id = %(amt_hit_id)s;
- """
- )
+ """)
data = {
"amt_hit_id": amt_hit_id,
"status": hit_status.value,
"modified_at": now,
}
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- conn.commit()
- assert c.rowcount == 1, c.rowcount
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ conn.commit()
+ assert c.rowcount == 1, c.rowcount
def update_hit(self, hit: Hit):
hit.modified_at = datetime.now(tz=timezone.utc)
@@ -255,8 +232,7 @@ class HitManager(PostgresManager):
"amt_hit_id",
"modified_at",
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_hit
SET status = %(status)s, review_status = %(review_status)s,
assignment_pending_count = %(assignment_pending_count)s,
@@ -265,17 +241,14 @@ class HitManager(PostgresManager):
modified_at = %(modified_at)s
WHERE amt_hit_id = %(amt_hit_id)s
RETURNING id;
- """
- )
+ """)
data = hit.model_dump(mode="json", include=fields)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- conn.commit()
- assert c.rowcount == 1, c.rowcount
- hit.id = c.fetchone()["id"] # type: ignore
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ conn.commit()
+ assert c.rowcount == 1, c.rowcount
+ hit.id = c.fetchone()["id"] # type: ignore
def get_from_amt_id(self, amt_hit_id: str) -> Hit:
res = self.pg_config.execute_sql_query(
@@ -324,7 +297,7 @@ class HitManager(PostgresManager):
return Hit.from_postgres(res)
- def get_from_amt_id_if_exists(self, amt_hit_id: str) -> Optional[Hit]:
+ def get_from_amt_id_if_exists(self, amt_hit_id: str) -> Hit | None:
try:
return self.get_from_amt_id(amt_hit_id=amt_hit_id)
@@ -333,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/managers/thl.py b/jb/managers/thl.py
index 6e8effc..85e0697 100644
--- a/jb/managers/thl.py
+++ b/jb/managers/thl.py
@@ -1,64 +1,18 @@
-from generalresearchutils.models.thl.payout import UserPayoutEvent
-from generalresearchutils.models.thl.task_status import TaskStatusResponse
-from generalresearchutils.models.thl.wallet.cashout_method import (
- CashoutRequestResponse,
+import requests
+from generalresearch.currency import USDCent
+from generalresearch.models.thl.definitions import PayoutStatus
+from generalresearch.models.thl.payout import UserPayoutEvent
+from generalresearch.models.thl.task_status import TaskStatusResponse
+from generalresearch.models.thl.wallet.cashout_method import (
CashoutRequestInfo,
+ CashoutRequestResponse,
)
-from generalresearchutils.models.thl.user_profile import UserProfile
-from generalresearchutils.currency import USDCent
-
from jb.config import settings
+from jb.models.auth import User
-from generalresearchutils.models.thl.definitions import PayoutStatus
-
-
-from typing import Optional
-import requests
-
-# TODO: Organize this more with other endpoints (offerwall, cashout
-# requests/approvals, etc).
-
-
-def get_user_profile(amt_worker_id: str) -> UserProfile:
- url = f"{settings.fsb_host}{settings.product_id}/user/{amt_worker_id}/profile/"
- res = requests.get(url).json()
-
- if res.get("detail") == "user not found":
- raise ValueError("user not found")
- user_profile = res["user_profile"]
- # todo: these are computed fields, need a UserProfile parser
- user_profile.pop("email_md5", None)
- user_profile.pop("email_sha1", None)
- user_profile.pop("email_sha256", None)
- # todo: this contains computed fields inside each streak object
- user_profile.pop("streaks", None)
- # todo: this shouldn't be in here anyways ---v
- user_profile["user"].pop("id", None)
-
- return UserProfile.model_validate(user_profile)
-
-
-def get_user_blocked(amt_worker_id: str) -> bool:
- # Not blocked if None
- res = get_user_profile(amt_worker_id=amt_worker_id)
- return res.user.blocked if res.user.blocked is not None else False
-
-
-def get_user_blocked_or_not_exists(amt_worker_id: str) -> Optional[bool]:
- try:
- res = get_user_profile(amt_worker_id=amt_worker_id)
- return res.user.blocked if res.user.blocked is not None else False
-
- except ValueError as e:
- if e.args[0] == "user not found":
- return True
-
- return None
-
-
-def get_task_status(tsid: str) -> Optional[TaskStatusResponse]:
+def get_task_status(tsid: str) -> TaskStatusResponse | None:
url = f"{settings.fsb_host}{settings.product_id}/status/{tsid}/"
d = requests.get(url).json()
if d.get("msg") == "invalid tsid":
@@ -68,18 +22,14 @@ def get_task_status(tsid: str) -> Optional[TaskStatusResponse]:
def user_cashout_request(
- amt_worker_id: str, amount: USDCent, cashout_method_id: str
+ user: User, amount: USDCent, cashout_method_id: str
) -> CashoutRequestInfo:
- assert cashout_method_id in {
- settings.amt_assignment_cashout_method,
- settings.amt_bonus_cashout_method,
- }
assert isinstance(amount, USDCent)
assert USDCent(0) < amount < USDCent(10_00)
url = f"{settings.fsb_host}{settings.product_id}/cashout/"
body: dict[str, str | int] = {
- "bpuid": amt_worker_id,
+ "bpuid": user.product_user_id,
"amount": int(amount),
"cashout_method_id": cashout_method_id,
}
@@ -110,7 +60,7 @@ def get_wallet_balance(amt_worker_id: str) -> USDCent:
return USDCent(requests.get(url, params=params).json()["wallet"]["amount"])
-def get_wallet_balance_if_non_negative(amt_worker_id: str) -> Optional[USDCent]:
+def get_wallet_balance_if_non_negative(amt_worker_id: str) -> USDCent | None:
url = f"{settings.fsb_host}{settings.product_id}/wallet/"
params = {"bpuid": amt_worker_id}
amt = requests.get(url, params=params).json()["wallet"]["amount"]
diff --git a/jb/managers/worker.py b/jb/managers/worker.py
index e2d7237..9e7d5e3 100644
--- a/jb/managers/worker.py
+++ b/jb/managers/worker.py
@@ -1,5 +1,3 @@
-from typing import List
-
from mypy_boto3_mturk.type_defs import WorkerBlockTypeDef
from jb.decorators import AMT_CLIENT
@@ -8,9 +6,9 @@ from jb.decorators import AMT_CLIENT
class WorkerManager:
@staticmethod
- def fetch_worker_blocks() -> List[WorkerBlockTypeDef]:
+ def fetch_worker_blocks() -> list[WorkerBlockTypeDef]:
p = AMT_CLIENT.get_paginator("list_worker_blocks")
- res: List[WorkerBlockTypeDef] = []
+ res: list[WorkerBlockTypeDef] = []
for item in p.paginate():
res.extend(item["WorkerBlocks"])
return res
diff --git a/jb/models/__init__.py b/jb/models/__init__.py
index 0aeae14..e69de29 100644
--- a/jb/models/__init__.py
+++ b/jb/models/__init__.py
@@ -1,40 +0,0 @@
-from decimal import Decimal
-from typing import Optional
-
-from pydantic import BaseModel, Field, ConfigDict
-
-
-class HTTPHeaders(BaseModel):
- request_id: str = Field(alias="x-amzn-requestid", min_length=36, max_length=36)
- content_type: str = Field(alias="content-type", min_length=26, max_length=26)
- # 'content-length': '1255',
- content_length: str = Field(alias="content-length", min_length=2)
- # 'Mon, 15 Jan 2024 23:40:32 GMT'
- date: str = Field()
-
- connection: Optional[str] = Field(default=None) # 'close'
-
-
-class ResponseMetadata(BaseModel):
- model_config = ConfigDict(extra="forbid", validate_assignment=True)
-
- request_id: str = Field(alias="RequestId", min_length=36, max_length=36)
- status_code: int = Field(alias="HTTPStatusCode", ge=200, le=599)
- headers: HTTPHeaders = Field(alias="HTTPHeaders")
- retry_attempts: int = Field(alias="RetryAttempts", ge=0)
-
-
-class AMTAccount(BaseModel):
- model_config = ConfigDict(extra="ignore", validate_assignment=True)
-
- # Remaining available AWS Billing usage if you have enabled AWS Billing.
- available_balance: Decimal = Field()
- onhold_balance: Decimal = Field(default=Decimal(0))
-
- # --- Properties ---
-
- @property
- def is_healthy(self) -> bool:
- # A healthy account is one with at least $2,500 worth of
- # credit available to it
- return self.available_balance >= 2_500
diff --git a/jb/models/amt.py b/jb/models/amt.py
new file mode 100644
index 0000000..e012741
--- /dev/null
+++ b/jb/models/amt.py
@@ -0,0 +1,19 @@
+from decimal import Decimal
+
+from pydantic import BaseModel, ConfigDict, Field
+
+
+class AMTAccount(BaseModel):
+ model_config = ConfigDict(extra="ignore", validate_assignment=True)
+
+ # Remaining available AWS Billing usage if you have enabled AWS Billing.
+ available_balance: Decimal = Field()
+ onhold_balance: Decimal = Field(default=Decimal(0))
+
+ # --- Properties ---
+
+ @property
+ def is_healthy(self) -> bool:
+ # A healthy account is one with at least $2,500 worth of
+ # credit available to it
+ return self.available_balance >= 2_500
diff --git a/jb/models/assignment.py b/jb/models/assignment.py
index 92e5a89..fa6ccd5 100644
--- a/jb/models/assignment.py
+++ b/jb/models/assignment.py
@@ -1,23 +1,26 @@
import logging
from datetime import datetime, timezone
-from typing import Optional, TypedDict, Any
+from typing import Any, TypedDict
from xml.etree import ElementTree
+from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
from mypy_boto3_mturk.type_defs import AssignmentTypeDef
from pydantic import (
BaseModel,
- Field,
ConfigDict,
- model_validator,
+ Field,
PositiveInt,
TypeAdapter,
ValidationError,
+ model_validator,
)
from typing_extensions import Self
-from jb.models.custom_types import AMTBoto3ID, AwareDatetimeISO, UUIDStr
+from jb.models.custom_types import AMTBoto3ID
from jb.models.definitions import AssignmentStatus
+logger = logging.getLogger("amtjb")
+
class AnswerDict(TypedDict):
amt_assignment_id: str
@@ -36,8 +39,8 @@ class AssignmentStub(BaseModel):
validate_assignment=True,
)
- id: Optional[PositiveInt] = Field(default=None)
- hit_id: Optional[PositiveInt] = Field(default=None)
+ id: PositiveInt | None = Field(default=None)
+ hit_id: PositiveInt | None = Field(default=None)
amt_assignment_id: AMTBoto3ID = Field()
amt_hit_id: AMTBoto3ID = Field()
amt_worker_id: str = Field(min_length=3, max_length=50)
@@ -50,7 +53,7 @@ class AssignmentStub(BaseModel):
description="When this record was saved in the database",
)
- modified_at: Optional[AwareDatetimeISO] = Field(
+ modified_at: AwareDatetimeISO | None = Field(
default_factory=lambda: datetime.now(tz=timezone.utc),
description="When this record was updated / modified in the database",
)
@@ -96,18 +99,18 @@ class Assignment(AssignmentStub):
"submitted results.",
)
- approval_time: Optional[AwareDatetimeISO] = Field(
+ approval_time: AwareDatetimeISO | None = Field(
default=None,
description="The date and time the Requester approved the results. This "
"value is omitted from the assignment if the Requester has "
"not yet approved the results.",
)
- rejection_time: Optional[AwareDatetimeISO] = Field(
+ rejection_time: AwareDatetimeISO | None = Field(
default=None,
description="The date and time the Requester rejected the results.",
)
- requester_feedback: Optional[str] = Field(
+ requester_feedback: str | None = Field(
# Default: None. This field isn't returned with assignment data by
# default. To request this field, specify a response group of
# AssignmentFeedback. For information about response groups, see
@@ -123,11 +126,11 @@ class Assignment(AssignmentStub):
},
)
- answer_xml: Optional[str] = Field(default=None, exclude=True)
+ answer_xml: str | None = Field(default=None, exclude=True)
# GRL Specific
- tsid: Optional[UUIDStr] = Field(default=None)
+ tsid: UUIDStr | None = Field(default=None)
# --- Validators ---
@@ -140,7 +143,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)
+ logger.warning(e)
+
values["tsid"] = None
return values
@@ -173,7 +177,7 @@ class Assignment(AssignmentStub):
# --- Properties ---
@property
- def answers_dict(self) -> Optional[AnswerDict]:
+ def answers_dict(self) -> AnswerDict | None:
# See https://docs.aws.amazon.com/AWSMechTurk/latest/AWSMturkAPI/ApiReference_AssignmentDataStructureArticle.html
# https://docs.aws.amazon.com/AWSMechTurk/latest/AWSMechanicalTurkRequester/Concepts_NotificationsArticle.html
if self.answer_xml is None:
diff --git a/jb/models/auth.py b/jb/models/auth.py
new file mode 100644
index 0000000..886a607
--- /dev/null
+++ b/jb/models/auth.py
@@ -0,0 +1,110 @@
+import hashlib
+import hmac
+from typing import Any
+
+from pydantic import (
+ BaseModel,
+ ConfigDict,
+ EmailStr,
+ Field,
+ TypeAdapter,
+ computed_field,
+ field_validator,
+ model_validator,
+)
+
+from jb.config import settings
+
+
+def email_to_product_user_id(email: str) -> str:
+ """Return a deterministic, non-reversible product user ID for an email.
+
+ 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")
+
+ if not email.isascii():
+ raise ValueError("email must contain ASCII characters only")
+
+ normalized_email = str(TypeAdapter(EmailStr).validate_python(email)).lower()
+ return hmac.new(
+ key=salt_bytes,
+ msg=normalized_email.encode("utf-8"),
+ digestmod=hashlib.sha1,
+ ).hexdigest()
+
+
+class User(BaseModel):
+ """A user that has been authenticated and exists in THL"""
+
+ email: EmailStr = Field()
+
+ display_name: str | None = Field(default=None, max_length=255)
+
+ blocked: bool = Field(default=False)
+
+ @field_validator("email", mode="before")
+ @classmethod
+ def normalize_email(cls, value: Any) -> Any:
+ if not isinstance(value, str):
+ return value
+ if not value.isascii():
+ raise ValueError("email must contain ASCII characters only")
+ return value.lower()
+
+ @model_validator(mode="before")
+ @classmethod
+ def validate_product_user_id(cls, data: Any) -> Any:
+ if not isinstance(data, dict) or "product_user_id" not in data:
+ return data
+
+ provided_id = data["product_user_id"]
+ email = data.get("email")
+ if email is None:
+ return data
+
+ expected_id = email_to_product_user_id(str(email))
+ if not isinstance(provided_id, str) or not hmac.compare_digest(
+ provided_id, expected_id
+ ):
+ raise ValueError(
+ f"product_user_id {provided_id} does not match email {email}"
+ )
+
+ # The computed field is authoritative; do not retain the input value.
+ validated_data = dict(data)
+ validated_data.pop("product_user_id")
+ return validated_data
+
+ @computed_field
+ @property
+ def product_user_id(self) -> str:
+ return email_to_product_user_id(self.email)
+
+
+class AccountLogin(BaseModel):
+ # There is no practical difference between an Account "login" and "create"
+ email: EmailStr = Field()
+
+
+class MagicLinkExchangeRequest(BaseModel):
+ model_config = ConfigDict(extra="forbid")
+
+ token: str = Field(min_length=1)
+
+
+class AmtAccountLink(BaseModel):
+ model_config = ConfigDict(extra="forbid")
+
+ email: EmailStr
+ amt_worker_id: str = Field(min_length=3, max_length=50)
+
+
+class SessionResponse(BaseModel):
+ session_token: str
+ token_type: str = "bearer"
+ expires_in: int
diff --git a/jb/models/bonus.py b/jb/models/bonus.py
index a536dd1..c6da3c4 100644
--- a/jb/models/bonus.py
+++ b/jb/models/bonus.py
@@ -1,10 +1,11 @@
-from typing import Optional, Dict, Any
+from typing import Any
-from pydantic import BaseModel, Field, ConfigDict, PositiveInt
+from generalresearch.currency import USDCent
+from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
+from pydantic import BaseModel, ConfigDict, Field, PositiveInt
from typing_extensions import Self
-from generalresearchutils.currency import USDCent
-from jb.models.custom_types import AMTBoto3ID, AwareDatetimeISO, UUIDStr
+from jb.models.custom_types import AMTBoto3ID
class Bonus(BaseModel):
@@ -20,8 +21,8 @@ class Bonus(BaseModel):
extra="forbid",
validate_assignment=True,
)
- id: Optional[PositiveInt] = Field(default=None)
- assignment_id: Optional[PositiveInt] = Field(default=None)
+ id: PositiveInt | None = Field(default=None)
+ assignment_id: PositiveInt | None = Field(default=None)
amt_worker_id: str = Field(min_length=3, max_length=50)
amt_assignment_id: AMTBoto3ID = Field()
@@ -40,7 +41,7 @@ class Bonus(BaseModel):
return d
@classmethod
- def from_postgres(cls, data: Dict[str, Any]) -> Self:
+ def from_postgres(cls, data: dict[str, Any]) -> Self:
data["amount"] = USDCent(round(data["amount"] * 100))
fields = set(cls.model_fields.keys())
data = {k: v for k, v in data.items() if k in fields}
diff --git a/jb/models/custom_types.py b/jb/models/custom_types.py
index 10bc9d1..385c1ba 100644
--- a/jb/models/custom_types.py
+++ b/jb/models/custom_types.py
@@ -1,101 +1,8 @@
import re
-from datetime import datetime, timezone
-from typing import Any, Optional
-from uuid import UUID
+from typing import Annotated
-from pydantic import (
- AwareDatetime,
- StringConstraints,
- TypeAdapter,
- HttpUrl,
-)
-from pydantic.functional_serializers import PlainSerializer
-from pydantic.functional_validators import AfterValidator, BeforeValidator
-from pydantic.networks import UrlConstraints
-from pydantic_core import Url
-from typing_extensions import Annotated
-
-
-def convert_datetime_to_iso_8601_with_z_suffix(dt: datetime) -> str:
- # By default, datetimes are serialized with the %f optional. We don't want that because
- # then the deserialization fails if the datetime didn't have microseconds.
- return dt.strftime("%Y-%m-%dT%H:%M:%S.%fZ")
-
-
-def convert_str_dt(v: Any) -> Optional[AwareDatetime]:
- # By default, pydantic is unable to handle tz-aware isoformat str. Attempt to parse a str
- # that was dumped using the iso8601 format with Z suffix.
- if v is not None and type(v) is str:
- assert v.endswith("Z") and "T" in v, "invalid format"
- return datetime.strptime(v, "%Y-%m-%dT%H:%M:%S.%fZ").replace(
- tzinfo=timezone.utc
- )
- return v
-
-
-def assert_utc(v: AwareDatetime) -> AwareDatetime:
- assert v.tzinfo == timezone.utc, "Timezone is not UTC"
- return v
-
-
-# Our custom AwareDatetime that correctly serializes and deserializes
-# to an ISO8601 str with timezone
-AwareDatetimeISO = Annotated[
- AwareDatetime,
- BeforeValidator(convert_str_dt),
- AfterValidator(assert_utc),
- PlainSerializer(
- lambda x: x.strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
- when_used="json-unless-none",
- ),
-]
-
-# ISO 3166-1 alpha-2 (two-letter codes, lowercase)
-# "Like" b/c it matches the format, but we're not explicitly checking
-# it is one of our supported values. See models.thl.locales for that.
-CountryISOLike = Annotated[
- str, StringConstraints(max_length=2, min_length=2, pattern=r"^[a-z]{2}$")
-]
-# 3-char ISO 639-2/B, lowercase
-LanguageISOLike = Annotated[
- str, StringConstraints(max_length=3, min_length=3, pattern=r"^[a-z]{3}$")
-]
-
-
-def check_valid_uuid(v: str) -> str:
- try:
- assert UUID(v).hex == v
- except Exception:
- raise ValueError("Invalid UUID")
- return v
-
-
-# Our custom field that stores a UUID4 as the .hex string representation
-UUIDStr = Annotated[
- str,
- StringConstraints(min_length=32, max_length=32),
- AfterValidator(check_valid_uuid),
-]
-# Accepts the non-hex representation and coerces
-UUIDStrCoerce = Annotated[
- str,
- StringConstraints(min_length=32, max_length=32),
- BeforeValidator(lambda value: TypeAdapter(UUID).validate_python(value).hex),
- AfterValidator(check_valid_uuid),
-]
-
-# Same thing as UUIDStr with HttpUrl field. It is confusing that this
-# is not a str https://github.com/pydantic/pydantic/discussions/6395
-HttpUrlStr = Annotated[
- str,
- BeforeValidator(lambda value: str(TypeAdapter(HttpUrl).validate_python(value))),
-]
-
-HttpsUrl = Annotated[Url, UrlConstraints(max_length=2083, allowed_schemes=["https"])]
-HttpsUrlStr = Annotated[
- str,
- BeforeValidator(lambda value: str(TypeAdapter(HttpsUrl).validate_python(value))),
-]
+from pydantic import StringConstraints
+from pydantic.functional_validators import AfterValidator
def check_valid_amt_boto3_id(v: str) -> str:
diff --git a/jb/models/errors.py b/jb/models/errors.py
index 94f5fbb..1fe71df 100644
--- a/jb/models/errors.py
+++ b/jb/models/errors.py
@@ -1,9 +1,9 @@
import re
from enum import Enum
-from pydantic import BaseModel, Field, ConfigDict, model_validator
+from pydantic import BaseModel, ConfigDict, Field, model_validator
-from jb.models import ResponseMetadata
+from jb.models.response import ResponseMetadata
class BotoRequestErrorOperation(str, Enum):
diff --git a/jb/models/event.py b/jb/models/event.py
index f8867c0..fb5735b 100644
--- a/jb/models/event.py
+++ b/jb/models/event.py
@@ -1,9 +1,10 @@
-from typing import Dict, Any
+from typing import Any
+from generalresearch.models.custom_types import AwareDatetimeISO
from mypy_boto3_mturk.literals import EventTypeType
from pydantic import BaseModel, Field
-from jb.models.custom_types import AwareDatetimeISO, AMTBoto3ID
+from jb.models.custom_types import AMTBoto3ID
class MTurkEvent(BaseModel):
@@ -29,7 +30,7 @@ class MTurkEvent(BaseModel):
)
@classmethod
- def from_sns(cls, data: Dict[str, Any]):
+ def from_sns(cls, data: dict[str, Any]):
return cls.model_validate(
{
"event_type": data["EventType"],
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
diff --git a/jb/models/response.py b/jb/models/response.py
new file mode 100644
index 0000000..22985af
--- /dev/null
+++ b/jb/models/response.py
@@ -0,0 +1,21 @@
+from pydantic import BaseModel, ConfigDict, Field
+
+
+class HTTPHeaders(BaseModel):
+ request_id: str = Field(alias="x-amzn-requestid", min_length=36, max_length=36)
+ content_type: str = Field(alias="content-type", min_length=26, max_length=26)
+ # 'content-length': '1255',
+ content_length: str = Field(alias="content-length", min_length=2)
+ # 'Mon, 15 Jan 2024 23:40:32 GMT'
+ date: str = Field()
+
+ connection: str | None = Field(default=None) # 'close'
+
+
+class ResponseMetadata(BaseModel):
+ model_config = ConfigDict(extra="forbid", validate_assignment=True)
+
+ request_id: str = Field(alias="RequestId", min_length=36, max_length=36)
+ status_code: int = Field(alias="HTTPStatusCode", ge=200, le=599)
+ headers: HTTPHeaders = Field(alias="HTTPHeaders")
+ retry_attempts: int = Field(alias="RetryAttempts", ge=0)
diff --git a/jb/settings.py b/jb/settings.py
index d529591..1c02075 100644
--- a/jb/settings.py
+++ b/jb/settings.py
@@ -1,78 +1,108 @@
-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 typing import Optional
-from generalresearchutils.models.custom_types import InfluxDsn
-from pydantic import Field, PostgresDsn, HttpUrl, RedisDsn
-from pydantic_settings import BaseSettings, SettingsConfigDict
+from generalresearch.config import GRLBaseSettings, is_debug
+from generalresearch.models.custom_types import (
+ InfluxDsn,
+ UUIDStr,
+)
+from pydantic import Field, HttpUrl, PostgresDsn, SecretStr, model_validator
+from pydantic_settings import SettingsConfigDict
-from jb.models.custom_types import UUIDStr
-
-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()
-class AmtJbBaseSettings(BaseSettings):
- debug: bool = Field(default=True)
-
- redis: Optional[RedisDsn] = Field(default=None)
- redis_timeout: float = Field(default=0.10)
-
- amt_jb_db: PostgresDsn = Field()
-
- amt_endpoint: Optional[HttpUrl] = Field(default=None)
- amt_access_id: Optional[str] = Field(default=None)
- amt_secret_key: Optional[str] = Field(default=None)
-
- aws_owner_id: str = Field()
- aws_subscription_arn: str = Field()
+class Settings(GRLBaseSettings):
- amt_bonus_cashout_method: str = Field()
- amt_assignment_cashout_method: str = Field()
-
-
-class Settings(AmtJbBaseSettings):
model_config = SettingsConfigDict(
- env_prefix="",
+ env_file=(".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
+
+ 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 | None = Field(default=None)
+ aws_subscription_arn: str | None = Field(default=None)
+
+ # --- Pytest ---
+
app_name: str = "AMT JB API"
+ base_url: HttpUrl = Field(default=HttpUrl("https://jamesbillings67.com/"))
fsb_host: HttpUrl = Field(default=HttpUrl("https://fsb.generalresearch.com/"))
# Needed for admin function on fsb w/o authentication
- fsb_host_private_route: Optional[str] = Field(default=None)
+ fsb_host_private_route: str | None = Field(default=None)
- product_id: UUIDStr = Field()
+ product_id: UUIDStr | None = Field(default=None)
- influx_db: Optional[InfluxDsn] = 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) # 30 days
+ session_jwt_secret: SecretStr | None = Field(default=None, min_length=32)
-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"
+ 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 | None = Field(default=None, min_length=1)
-@lru_cache
-def get_settings():
- return Settings()
+ mautic_api_key: SecretStr | None = Field(default=None, min_length=32)
+
+ @model_validator(mode="after")
+ def validate_host_and_key(self) -> "Settings":
+
+ self.debug = is_debug()
+
+ if self.debug:
+ return self
+
+ 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")
+
+ 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
new file mode 100644
index 0000000..ce9c1e0
--- /dev/null
+++ b/jb/views/auth.py
@@ -0,0 +1,203 @@
+from typing import Annotated
+from urllib.parse import urlencode
+
+from fastapi import APIRouter, Depends, HTTPException, Response, status
+from fastapi.responses import HTMLResponse, RedirectResponse
+
+from jb.api.auth import (
+ SESSION_COOKIE_NAME,
+ create_session,
+ get_authenticated_user,
+)
+from jb.api.magic_token import (
+ consume_amt_account_link_token,
+ consume_magic_token,
+ create_amt_account_link_token,
+ 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,
+ send_amt_link_email,
+ send_login_email,
+)
+from jb.managers.gr_api import GRApiManager
+from jb.models.auth import (
+ AccountLogin,
+ AmtAccountLink,
+ MagicLinkExchangeRequest,
+ User,
+)
+from jb.settings import BASE_HTML
+
+auth_router = APIRouter(prefix="/auth", tags=["Auth"])
+
+
+@auth_router.post("/magic-link/request")
+def request_magic_link(body: AccountLogin) -> dict[str, str]:
+ """Create a magic link."""
+ email = str(body.email)
+ token = create_magic_token(user_email=email)
+
+ if settings.debug:
+ query = urlencode({"token": token})
+ return {"magic_link": f"{settings.base_url}auth/magic-link/?{query}"}
+
+ send_login_email(email=email, magic_token=token)
+ return {"detail": "Link sent. Check your inbox and follow the link to log in."}
+
+
+@auth_router.get("/magic-link/", response_class=HTMLResponse, include_in_schema=False)
+def magic_link_landing_page(
+ gr_api: Annotated[GRApiManager, Depends(get_gr_api_manager)],
+ token: str | None = None,
+) -> Response:
+ """Serve the SPA without redeeming the token; email prefetches are harmless."""
+ if settings.debug:
+ if token is None:
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail="token is required",
+ )
+ response = RedirectResponse(url="/", status_code=status.HTTP_303_SEE_OTHER)
+ _exchange_magic_link(token, response, gr_api)
+ return response
+ return HTMLResponse(
+ BASE_HTML,
+ headers={
+ "Cache-Control": "no-store",
+ "Referrer-Policy": "no-referrer",
+ "X-Robots-Tag": "noindex, nofollow",
+ },
+ )
+
+
+@auth_router.post("/magic-link/exchange", status_code=status.HTTP_204_NO_CONTENT)
+def exchange_magic_link(
+ body: MagicLinkExchangeRequest,
+ response: Response,
+ gr_api: Annotated[GRApiManager, Depends(get_gr_api_manager)],
+) -> None:
+ """Exchange a magic link only after its landing page makes an explicit POST."""
+ _exchange_magic_link(body.token, response, gr_api)
+
+
+def _exchange_magic_link(token: str, response: Response, gr_api: GRApiManager) -> None:
+ user_email = consume_magic_token(token)
+ user = gr_api.ensure_user_exists(User.model_validate({"email": user_email}))
+ response.set_cookie(
+ key=SESSION_COOKIE_NAME,
+ value=create_session(user.product_user_id),
+ max_age=settings.session_token_ttl_seconds,
+ httponly=True,
+ secure=not settings.debug,
+ samesite="lax",
+ path="/",
+ )
+
+
+@auth_router.post("/link-amt/request")
+def link_amt_account(body: AmtAccountLink) -> dict[str, str]:
+ """Link an AMT account and login."""
+ email = str(body.email)
+ amt_worker_id = body.amt_worker_id
+
+ # TODO! Prevent a user that's already transitioned their account, from being
+ # TODO! able to continuously create this special link token.
+ # TODO! Max Notes: this seems to be handled within the gr-api, and that
+ # TODO! can raise, the following line would / should fail if needed.
+
+ token = create_amt_account_link_token(email=email, amt_worker_id=amt_worker_id)
+
+ if settings.debug:
+ query = urlencode({"token": token})
+ return {"magic_link": f"{settings.base_url}auth/link-amt/?{query}"}
+
+ send_amt_link_email(email=email, magic_token=token)
+ return {}
+
+
+@auth_router.get("/debug/", response_class=HTMLResponse, include_in_schema=False)
+def link_amt_account_landing_page(
+ gr_api: Annotated[GRApiManager, Depends(get_gr_api_manager)],
+ token: str | None = None,
+) -> HTMLResponse:
+ """Serve the account-link SPA without consuming the one-time token."""
+
+ # TODO! Try catch any of this, and if it fails, show the user a
+ # TODO! failed HTML page. As of now, it shows them a failed JSON response.
+
+ if settings.debug:
+ if token is None:
+ raise HTTPException(
+ status_code=status.HTTP_400_BAD_REQUEST,
+ detail="token is required",
+ )
+
+ _response = RedirectResponse(url="/", status_code=status.HTTP_303_SEE_OTHER)
+ _exchange_amt_account_link(token=token, response=_response, gr_api=gr_api)
+
+ return HTMLResponse(
+ BASE_HTML,
+ headers={
+ "Cache-Control": "no-store",
+ "Referrer-Policy": "no-referrer",
+ "X-Robots-Tag": "noindex, nofollow",
+ },
+ )
+
+
+@auth_router.post("/link-amt/exchange", status_code=status.HTTP_204_NO_CONTENT)
+def exchange_amt_account_link(
+ body: MagicLinkExchangeRequest,
+ response: Response,
+ gr_api: Annotated[GRApiManager, Depends(get_gr_api_manager)],
+) -> None:
+ """Validate the email link, then transition the bound AMT account."""
+ try:
+ _exchange_amt_account_link(body.token, response, gr_api)
+ except ValueError as e:
+ LOG.error(f"Failed to exchange AMT account link: {e}")
+ raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e))
+
+
+def _exchange_amt_account_link(token: str, response: Response, gr_api: GRApiManager):
+ token_data = consume_amt_account_link_token(token)
+ email = token_data.email
+ amt_worker_id = token_data.amt_worker_id
+
+ user = User(email=email)
+ user = gr_api.transition_user_from_amt(user=user, amt_worker_id=amt_worker_id)
+
+ # In Mautic, associate the email with the worker ID (AFTER the user has transitioned)
+ get_or_create_contact(email=email, amt_worker_id=amt_worker_id)
+
+ session_token = create_session(user.product_user_id)
+ response.set_cookie(
+ key=SESSION_COOKIE_NAME,
+ value=session_token,
+ max_age=settings.session_token_ttl_seconds,
+ httponly=True,
+ secure=not settings.debug,
+ samesite="lax",
+ path="/",
+ )
+
+
+@auth_router.get("/session", response_model=User)
+def get_session(
+ user: Annotated[User, Depends(get_authenticated_user)],
+) -> User:
+ return user
+
+
+@auth_router.delete("/session", status_code=status.HTTP_204_NO_CONTENT)
+def delete_session(
+ response: Response,
+) -> None:
+ # Logout is idempotent so clients can always discard their local token.
+ # JWT sessions are stateless, so logout discards the browser cookie. A token
+ # copied elsewhere remains valid until its short, configured expiration.
+ response.delete_cookie(key=SESSION_COOKIE_NAME, path="/")
diff --git a/jb/views/common.py b/jb/views/common.py
index 7011557..9eee453 100644
--- a/jb/views/common.py
+++ b/jb/views/common.py
@@ -1,19 +1,19 @@
import json
-from typing import Dict, Any
+from typing import Annotated, Any
import requests
-from fastapi import Request, APIRouter, HTTPException
+from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import HTMLResponse
+from generalresearch.redis_helper import RedisConfig
from starlette.responses import RedirectResponse
-from jb.config import settings, JB_EVENTS_STREAM
-from jb.decorators import REDIS, HM
-from jb.flow.monitoring import emit_assignment_event, emit_mturk_notification_event
-from jb.models.definitions import AssignmentStatus
+from jb.api.auth import get_authenticated_user
+from jb.config import JB_EVENTS_STREAM, settings
+from jb.decorators import get_redis_config
+from jb.flow.monitoring import emit_mturk_notification_event
+from jb.models.auth import User
from jb.models.event import MTurkEvent
from jb.settings import BASE_HTML
-from jb.config import settings
-from jb.views.tasks import process_request
common_router = APIRouter(prefix="", tags=["API"], include_in_schema=True)
@@ -29,37 +29,24 @@ async def work(request: Request):
amt_hit_id = request.query_params.get("hitId", None)
print(f"work: {amt_assignment_id=} {worker_id=} {amt_hit_id=}")
- if not worker_id:
+ if (
+ not worker_id
+ or not amt_assignment_id
+ or amt_assignment_id == "ASSIGNMENT_ID_NOT_AVAILABLE"
+ ):
return RedirectResponse(
url=f"/preview/?{request.url.query}" if request.url.query else "/preview/",
status_code=302,
)
- if amt_assignment_id is None or amt_assignment_id == "ASSIGNMENT_ID_NOT_AVAILABLE":
- # Worker is previewing the HIT
- amt_hit_type_id = "unknown"
- if amt_hit_id:
- hit = HM.get_from_amt_id(amt_hit_id=amt_hit_id)
- amt_hit_type_id = hit.amt_hit_type_id
- emit_assignment_event(
- status=AssignmentStatus.PreviewState, amt_hit_type_id=amt_hit_type_id
- )
- return RedirectResponse(
- url=f"/preview/?{request.url.query}" if request.url.query else "/preview/",
- status_code=302,
- )
+ return HTMLResponse(BASE_HTML)
- try:
- # The Worker has accepted the HIT
- process_request(request)
- except Exception:
- raise HTTPException(status_code=500, detail="Error processing request")
- return HTMLResponse(BASE_HTML)
+RedisConfigDep = Annotated[RedisConfig, Depends(get_redis_config)]
@common_router.post(path=f"/{settings.sns_path}/", include_in_schema=False)
-async def mturk_notifications(request: Request):
+async def mturk_notifications(request: Request, redis_config: RedisConfigDep):
"""
Our SNS topic will POST to this endpoint whenever we get a new message
"""
@@ -77,7 +64,7 @@ async def mturk_notifications(request: Request):
case "Notification":
msg = json.loads(message["Message"])
print("Received MTurk event:", msg)
- enqueue_mturk_notifications(msg)
+ enqueue_mturk_notifications(msg=msg, redis_config=redis_config)
case _:
raise HTTPException(status_code=500, detail="Invalid JSON")
@@ -85,10 +72,27 @@ async def mturk_notifications(request: Request):
return {"status": "ok"}
-def enqueue_mturk_notifications(msg: Dict[str, Any]) -> None:
+def enqueue_mturk_notifications(msg: dict[str, Any], redis_config: RedisConfig) -> None:
+ redis_client = redis_config.create_redis_client()
+
for evt in msg["Events"]:
event = MTurkEvent.from_sns(evt)
emit_mturk_notification_event(
event_type=event.event_type, amt_hit_type_id=event.amt_hit_type_id
)
- REDIS.xadd(JB_EVENTS_STREAM, {"data": event.model_dump_json()})
+
+ print("enqueue_mturk_notifications", event.model_dump_json())
+ redis_client.xadd(JB_EVENTS_STREAM, {"data": event.model_dump_json()})
+
+
+@common_router.get(path="/work/direct/", response_class=HTMLResponse)
+async def work_direct(
+ request: Request,
+ user: Annotated[User, Depends(get_authenticated_user)],
+):
+ """
+ View for direct work (not on AMT).
+ Makes sure user is authenticated.
+ """
+ # todo: emit event
+ return HTMLResponse(BASE_HTML)
diff --git a/jb/views/tasks.py b/jb/views/tasks.py
deleted file mode 100644
index 15857c3..0000000
--- a/jb/views/tasks.py
+++ /dev/null
@@ -1,78 +0,0 @@
-from datetime import datetime, timezone, timedelta
-
-from fastapi import Request
-
-from jb.decorators import AMTM, AM, HM
-from jb.flow.maintenance import check_hit_status
-from jb.flow.monitoring import emit_assignment_event
-from jb.models.assignment import AssignmentStub
-from jb.models.definitions import AssignmentStatus
-
-
-def process_request(request: Request) -> None:
- """
- A worker has loaded the HIT (work) page and (probably) accepted the HIT.
- AMT creates an assignment, tied to this hit and this worker.
- Create it in the DB.
- """
- amt_assignment_id = request.query_params.get("assignmentId", None)
- if amt_assignment_id == "ASSIGNMENT_ID_NOT_AVAILABLE":
- raise ValueError("shouldn't happen")
-
- amt_hit_id = request.query_params.get("hitId", None)
- amt_worker_id = request.query_params.get("workerId", None)
- print(f"process_request: {amt_assignment_id=} {amt_worker_id=} {amt_hit_id=}")
- assert amt_worker_id and amt_hit_id and amt_assignment_id
-
- # Check that the HIT is still valid
- hit = HM.get_from_amt_id_if_exists(amt_hit_id=amt_hit_id)
- if not hit:
- raise ValueError(f"Hit {amt_hit_id} not found in DB")
-
- _ = check_hit_status(
- amtm=AMTM, amt_hit_id=amt_hit_id, amt_hit_type_id=hit.amt_hit_type_id
- )
-
- emit_assignment_event(
- status=AssignmentStatus.Accepted,
- amt_hit_type_id=hit.amt_hit_type_id,
- )
-
- # I think it won't be assignable anymore? idk
- # assert hit_status == HitStatus.Assignable, f"hit {amt_hit_id} {hit_status=}. Expected Assignable"
-
- # I would like to verify in the AMT API that this assignment is valid, but there
- # is no way to do that (until the assignment is submitted)
-
- # # Make an offerwall to create a user account...
- # # todo: GSS: Do we really need to do this???
- # client_ip = get_client_ip(request)
- # url = f"{settings.fsb_host}{settings.product_id}/offerwall/45b7228a7/"
- # _ = requests.get(
- # url,
- # {"bpuid": amt_worker_id, "ip": client_ip, "n_bins": 1, "format": "json"},
- # ).json()
-
- # This assignment shouldn't already exist. If it does, just make sure it
- # is all the same.
- assignment_stub = AM.get_stub_if_exists(amt_assignment_id=amt_assignment_id)
- if assignment_stub:
- print(f"{assignment_stub=}")
- assert assignment_stub.amt_worker_id == amt_worker_id
- assert assignment_stub.amt_assignment_id == amt_assignment_id
- assert assignment_stub.created_at > (
- datetime.now(tz=timezone.utc) - timedelta(minutes=90)
- )
- return None
-
- assignment_stub = AssignmentStub(
- amt_hit_id=amt_hit_id,
- amt_worker_id=amt_worker_id,
- amt_assignment_id=amt_assignment_id,
- status=AssignmentStatus.Accepted,
- hit_id=hit.id,
- )
-
- AM.create_stub(stub=assignment_stub)
-
- return None