diff options
| author | Max Nanis | 2026-09-13 19:03:55 +0000 |
|---|---|---|
| committer | Max Nanis | 2026-09-13 19:03:55 +0000 |
| commit | 8fbe8d439b418796932aaa88297e006c714e4d13 (patch) | |
| tree | 6c800476edc1e771fc570b559d483df8d59d4324 /jb | |
| parent | 12f6fee851b68e86af658dfa17e4a0daed457dd1 (diff) | |
| parent | 2c94f248d2438071a918fa9a30bf114ef9aa29b4 (diff) | |
| download | amt-jb-8fbe8d439b418796932aaa88297e006c714e4d13.tar.gz amt-jb-8fbe8d439b418796932aaa88297e006c714e4d13.zip | |
Merges pull request #3
Off of Amazon!!!
Diffstat (limited to 'jb')
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) @@ -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 |
