From 3f8d178d38686c1a5d71d94ebccdb61701469e95 Mon Sep 17 00:00:00 2001 From: stuppie Date: Mon, 31 Aug 2026 16:52:15 -0600 Subject: working on magic token and session token auth --- requirements.txt | 1 + 1 file changed, 1 insertion(+) (limited to 'requirements.txt') diff --git a/requirements.txt b/requirements.txt index 8976073..c23b668 100644 --- a/requirements.txt +++ b/requirements.txt @@ -66,6 +66,7 @@ Pygments==2.19.2 pylibmc==1.6.3 pymemcache==4.0.0 PyMySQL==1.1.2 +PyJWT==2.13.0 pytest==8.4.2 python-dateutil==2.9.0.post0 python-dotenv==1.1.1 -- cgit v1.2.3 From c991d4f51254b1c81fa6e0e19b0ba94b8c5f61c4 Mon Sep 17 00:00:00 2001 From: stuppie Date: Tue, 1 Sep 2026 11:56:05 -0600 Subject: add GRApiManager. py-utils to generalresearch --- jb/api/auth.py | 13 ++--- jb/config.py | 2 +- jb/decorators.py | 13 +++-- jb/flow/assignment_tasks.py | 6 +-- jb/flow/monitoring.py | 2 +- jb/flow/tasks.py | 2 +- jb/managers/__init__.py | 2 +- jb/managers/amt.py | 2 +- jb/managers/gr_api.py | 112 ++++++++++++++++++++++++++++++++++++++++++++ jb/managers/thl.py | 12 ++--- jb/models/auth.py | 47 ++++++++++++++++--- jb/models/bonus.py | 2 +- jb/models/hit.py | 2 +- jb/settings.py | 7 ++- jb/views/auth.py | 36 ++++++++++---- requirements.txt | 9 +++- 16 files changed, 223 insertions(+), 46 deletions(-) create mode 100644 jb/managers/gr_api.py (limited to 'requirements.txt') diff --git a/jb/api/auth.py b/jb/api/auth.py index 7d6c352..3481b30 100644 --- a/jb/api/auth.py +++ b/jb/api/auth.py @@ -8,13 +8,12 @@ from fastapi import Depends, HTTPException, Request, Response, status from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer from jb.config import settings -from jb.models.auth import AuthenticatedUser +from jb.decorators import gr_api_manager +from jb.models.auth import User bearer = HTTPBearer(auto_error=False) logger = logging.getLogger(__name__) -AUTH_CACHE_TTL_SECONDS = 60 - SESSION_COOKIE_NAME = "jb_session" JWT_ISSUER = "jamesbillings67" JWT_AUDIENCE = "jamesbillings67" @@ -23,7 +22,7 @@ JWT_AUDIENCE = "jamesbillings67" def get_authenticated_user( request: Request, credentials: Annotated[HTTPAuthorizationCredentials | None, Depends(bearer)], -) -> AuthenticatedUser: +) -> User: """FastAPI dependency for endpoints requiring a valid session.""" if settings.session_jwt_secret is None: raise HTTPException( @@ -66,11 +65,9 @@ def get_authenticated_user( ) product_user_id = claims.get("sub") - # todo: in here, hit THL by the user's bpuid (email hash), - # in order to 1) be sure user exists & 2) pull the display_name - user_email = ... # lookup in thl by product_user_id + user = gr_api_manager.get_user(product_user_id=product_user_id) - return AuthenticatedUser(email=user_email) + return user def create_session(product_user_id: str) -> str: diff --git a/jb/config.py b/jb/config.py index c7d07e5..a91494d 100644 --- a/jb/config.py +++ b/jb/config.py @@ -1,6 +1,6 @@ import logging -from generalresearchutils.config import is_debug +from generalresearch.config import is_debug from jb.settings import get_settings, get_test_settings diff --git a/jb/decorators.py b/jb/decorators.py index 5c1b1f5..cafd7a9 100644 --- a/jb/decorators.py +++ b/jb/decorators.py @@ -1,7 +1,7 @@ import boto3 from botocore.config import Config -from generalresearchutils.pg_helper import PostgresConfig -from generalresearchutils.redis_helper import RedisConfig +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 @@ -11,7 +11,8 @@ 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, @@ -32,6 +33,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/flow/assignment_tasks.py b/jb/flow/assignment_tasks.py index 18a844e..2db20f8 100644 --- a/jb/flow/assignment_tasks.py +++ b/jb/flow/assignment_tasks.py @@ -3,9 +3,9 @@ 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 generalresearch.models.thl.definitions import PayoutStatus, StatusCode1 +from generalresearch.models.thl.wallet.cashout_method import CashoutRequestInfo +from generalresearch.currency import USDCent from jb.flow.monitoring import emit_error_event, emit_assignment_event, emit_bonus_event from jb.managers.amt import ( diff --git a/jb/flow/monitoring.py b/jb/flow/monitoring.py index 28f7271..334eab3 100644 --- a/jb/flow/monitoring.py +++ b/jb/flow/monitoring.py @@ -5,7 +5,7 @@ from mypy_boto3_mturk.literals import EventTypeType from jb.config import settings from jb.decorators import influx_client -from generalresearchutils.currency import USDCent +from generalresearch.currency import USDCent from jb.models.definitions import HitStatus, AssignmentStatus diff --git a/jb/flow/tasks.py b/jb/flow/tasks.py index 77825d3..6c5d6bb 100644 --- a/jb/flow/tasks.py +++ b/jb/flow/tasks.py @@ -2,7 +2,7 @@ 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.flow.maintenance import check_hit_status diff --git a/jb/managers/__init__.py b/jb/managers/__init__.py index e99569a..ec64f9f 100644 --- a/jb/managers/__init__.py +++ b/jb/managers/__init__.py @@ -1,7 +1,7 @@ from enum import IntEnum from typing import Collection -from generalresearchutils.pg_helper import PostgresConfig +from generalresearch.pg_helper import PostgresConfig class Permission(IntEnum): diff --git a/jb/managers/amt.py b/jb/managers/amt.py index 88ae07d..2cb0cbd 100644 --- a/jb/managers/amt.py +++ b/jb/managers/amt.py @@ -3,7 +3,7 @@ from datetime import timezone, datetime from typing import Tuple, Optional, List, Dict, 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, diff --git a/jb/managers/gr_api.py b/jb/managers/gr_api.py new file mode 100644 index 0000000..8ffeaf0 --- /dev/null +++ b/jb/managers/gr_api.py @@ -0,0 +1,112 @@ +"""Client for General Research's product-user API.""" + +from typing import Any + +import requests + +from jb.models.auth import User + + +class GRApiError(RuntimeError): + """The General Research API could not satisfy a request.""" + + +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.RequestException as exc: + raise GRApiError( + f"General Research API request failed: {method} {url}" + ) from exc + + def _parse_user_response(self, res: dict) -> 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"), + } + ) + + 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/" + res = 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/" + res = self._request( + "PATCH", + url, + json={"display_name": user.display_name}, + ).json() + return self.get_user(user.product_user_id) + + def transition_product_user_id(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}/" + res = self._request( + "PATCH", url, json={"product_user_id": user.product_user_id} + ) + self.set_user_email(user) + return self.get_user(user.product_user_id) diff --git a/jb/managers/thl.py b/jb/managers/thl.py index 6e8effc..f0534db 100644 --- a/jb/managers/thl.py +++ b/jb/managers/thl.py @@ -1,16 +1,16 @@ -from generalresearchutils.models.thl.payout import UserPayoutEvent -from generalresearchutils.models.thl.task_status import TaskStatusResponse -from generalresearchutils.models.thl.wallet.cashout_method import ( +from generalresearch.models.thl.payout import UserPayoutEvent +from generalresearch.models.thl.task_status import TaskStatusResponse +from generalresearch.models.thl.wallet.cashout_method import ( CashoutRequestResponse, CashoutRequestInfo, ) -from generalresearchutils.models.thl.user_profile import UserProfile -from generalresearchutils.currency import USDCent +from generalresearch.models.thl.user_profile import UserProfile +from generalresearch.currency import USDCent from jb.config import settings -from generalresearchutils.models.thl.definitions import PayoutStatus +from generalresearch.models.thl.definitions import PayoutStatus from typing import Optional diff --git a/jb/models/auth.py b/jb/models/auth.py index 5f31bd3..8b48088 100644 --- a/jb/models/auth.py +++ b/jb/models/auth.py @@ -1,5 +1,6 @@ import hashlib import hmac +from typing import Any from pydantic import ( BaseModel, @@ -8,6 +9,8 @@ from pydantic import ( Field, TypeAdapter, computed_field, + field_validator, + model_validator, ) from jb.config import settings @@ -30,29 +33,59 @@ def email_to_product_user_id(email: str) -> str: return hmac.new( key=salt_bytes, msg=normalized_email.encode("utf-8"), - digestmod=hashlib.sha256, + digestmod=hashlib.sha1, ).hexdigest() -class AuthenticatedUser(BaseModel): +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) + + @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 AccountCreate(BaseModel): - email: EmailStr = Field() - - 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") diff --git a/jb/models/bonus.py b/jb/models/bonus.py index a536dd1..5f81add 100644 --- a/jb/models/bonus.py +++ b/jb/models/bonus.py @@ -3,7 +3,7 @@ from typing import Optional, Dict, Any from pydantic import BaseModel, Field, ConfigDict, PositiveInt from typing_extensions import Self -from generalresearchutils.currency import USDCent +from generalresearch.currency import USDCent from jb.models.custom_types import AMTBoto3ID, AwareDatetimeISO, UUIDStr diff --git a/jb/models/hit.py b/jb/models/hit.py index fba2ecf..45478fc 100644 --- a/jb/models/hit.py +++ b/jb/models/hit.py @@ -13,7 +13,7 @@ from pydantic import ( ) from typing_extensions import Self -from generalresearchutils.currency import USDCent +from generalresearch.currency import USDCent from jb.models.custom_types import AMTBoto3ID, HttpsUrlStr, AwareDatetimeISO from jb.models.definitions import HitStatus, HitReviewStatus diff --git a/jb/settings.py b/jb/settings.py index cc1e870..78b4b02 100644 --- a/jb/settings.py +++ b/jb/settings.py @@ -3,7 +3,7 @@ from functools import lru_cache from pathlib import Path from typing import Optional -from generalresearchutils.models.custom_types import InfluxDsn +from generalresearch.models.custom_types import InfluxDsn from pydantic import Field, PostgresDsn, HttpUrl, RedisDsn, SecretStr from pydantic_settings import BaseSettings, SettingsConfigDict @@ -60,6 +60,11 @@ class Settings(AmtJbBaseSettings): magic_token_salt: SecretStr = Field(min_length=32) + gr_api_host: HttpUrl = Field( + default=HttpUrl("https://generalresearch.com/api/v2/") + ) + gr_api_token: SecretStr = Field(min_length=1) + class TestSettings(Settings): model_config = SettingsConfigDict( diff --git a/jb/views/auth.py b/jb/views/auth.py index 0591294..06dbdde 100644 --- a/jb/views/auth.py +++ b/jb/views/auth.py @@ -6,8 +6,9 @@ has resolved an email address to its stable user identifier, call """ from typing import Annotated +from urllib.parse import urlencode -from fastapi import APIRouter, Depends, Response, status +from fastapi import APIRouter, Depends, HTTPException, Response, status from fastapi.responses import HTMLResponse from jb.api.auth import ( @@ -15,11 +16,13 @@ from jb.api.auth import ( create_session, get_authenticated_user, ) -from jb.api.magic_token import consume_magic_token +from jb.api.magic_token import consume_magic_token, create_magic_token from jb.config import settings +from jb.decorators import gr_api_manager from jb.models.auth import ( MagicLinkExchangeRequest, - AuthenticatedUser, + AccountLogin, + User, email_to_product_user_id, ) from jb.settings import BASE_HTML @@ -27,6 +30,20 @@ from jb.settings import BASE_HTML auth_router = APIRouter(prefix="/auth", tags=["Auth"]) +@auth_router.post("/magic-link/request") +def request_mock_magic_link(body: AccountLogin) -> dict[str, str]: + """Create a magic link without sending email in development.""" + if not settings.debug: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND) + + # todo: send email here + + user = User(email=body.email) + token = create_magic_token(str(user.email)) + query = urlencode({"token": token}) + return {"magic_link": f"/auth/magic-link/?{query}"} + + @auth_router.get("/magic-link/", response_class=HTMLResponse, include_in_schema=False) def magic_link_landing_page() -> HTMLResponse: """Serve the SPA without redeeming the token; email prefetches are harmless.""" @@ -44,11 +61,12 @@ def magic_link_landing_page() -> HTMLResponse: def exchange_magic_link(body: MagicLinkExchangeRequest, response: Response) -> None: """Exchange a magic link only after its landing page makes an explicit POST.""" user_email = consume_magic_token(body.token) - product_user_id = email_to_product_user_id(user_email) - # todo: hit thl to make sure this user exists + user = User.model_validate({'email': user_email}) + # hit thl to make sure this user exists + user = gr_api_manager.ensure_user_exists(user) - session_token = create_session(product_user_id) + session_token = create_session(user.product_user_id) response.set_cookie( key=SESSION_COOKIE_NAME, value=session_token, @@ -60,10 +78,10 @@ def exchange_magic_link(body: MagicLinkExchangeRequest, response: Response) -> N ) -@auth_router.get("/session", response_model=AuthenticatedUser) +@auth_router.get("/session", response_model=User) def get_session( - user: Annotated[AuthenticatedUser, Depends(get_authenticated_user)], -) -> AuthenticatedUser: + user: Annotated[User, Depends(get_authenticated_user)], +) -> User: return user diff --git a/requirements.txt b/requirements.txt index c23b668..315a711 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ -git+ssh://code.g-r-l.com:6611/py-utils@v3.3.4 +git+ssh://code.g-r-l.com:6611/generalresearch@v3.4.5 aiohappyeyeballs==2.6.1 aiohttp==3.13.0 aiosignal==1.4.0 @@ -24,6 +24,7 @@ dnspython==2.8.0 email-validator==2.3.0 Faker==37.11.0 fastapi==0.119.0 +filelock==3.32.5 frozenlist==1.8.0 fsspec==2025.9.0 geoip2==5.1.0 @@ -51,22 +52,24 @@ pandas==2.3.3 pandera==0.26.1 partd==1.4.2 pathspec==0.12.1 +phonenumbers==9.0.38 platformdirs==4.5.0 pluggy==1.6.0 propcache==0.4.1 protobuf==6.32.1 psutil==7.1.1 psycopg==3.2.10 +pyarrow==25.0.1 pycountry==24.6.1 pydantic==2.12.0 pydantic-extra-types==2.10.6 pydantic-settings==2.11.0 pydantic_core==2.41.1 Pygments==2.19.2 +PyJWT==2.13.0 pylibmc==1.6.3 pymemcache==4.0.0 PyMySQL==1.1.2 -PyJWT==2.13.0 pytest==8.4.2 python-dateutil==2.9.0.post0 python-dotenv==1.1.1 @@ -75,6 +78,7 @@ pytz==2025.2 PyYAML==6.0.3 redis==6.4.0 requests==2.32.5 +requests-file==3.0.1 requests-mock==1.12.1 s3transfer==0.14.0 scipy==1.16.2 @@ -86,6 +90,7 @@ sniffio==1.3.1 sortedcontainers==2.4.0 starlette==0.48.0 tblib==3.2.0 +tldextract==5.3.2 toolz==1.1.0 tornado==6.5.2 typeguard==4.4.4 -- cgit v1.2.3 From 832aecaddce80e312095ecdb572d7756eb9df5e9 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Wed, 9 Sep 2026 20:32:17 -0700 Subject: Ruff auto fix, and requirements bump --- jb/api/auth.py | 3 +- jb/api/magic_token.py | 6 +- jb/flow/maintenance.py | 4 +- jb/flow/monitoring.py | 23 ++------ jb/flow/setup_tasks.py | 5 +- jb/flow/tasks.py | 6 +- jb/main.py | 10 ++-- jb/managers/__init__.py | 2 +- jb/managers/amt.py | 28 ++++------ jb/managers/assignment.py | 85 +++++++++++------------------ jb/managers/bonus.py | 22 +++----- jb/managers/email_manager.py | 5 +- jb/managers/gr_api.py | 6 +- jb/managers/hit.py | 115 +++++++++++++++------------------------ jb/managers/thl.py | 19 ++----- jb/managers/worker.py | 6 +- jb/models/__init__.py | 5 +- jb/models/assignment.py | 24 ++++---- jb/models/bonus.py | 12 ++-- jb/models/custom_types.py | 7 +-- jb/models/errors.py | 2 +- jb/models/event.py | 6 +- jb/models/hit.py | 42 +++++++------- jb/settings.py | 19 +++---- requirements.txt | 2 +- tests/__init__.py | 8 +-- tests/conftest.py | 10 ++-- tests/fixtures/amt.py | 11 ++-- tests/fixtures/flow.py | 39 ++++++------- tests/fixtures/http.py | 21 ++++--- tests/fixtures/managers.py | 8 ++- tests/fixtures/models.py | 66 ++++++++++------------ tests/flow/test_tasks.py | 52 +++++++++--------- tests/http/test_auth.py | 4 +- tests/http/test_notifications.py | 11 ++-- tests/http/test_preview.py | 3 +- tests/http/test_work.py | 4 +- tests/managers/test_amt.py | 9 ++- tests/managers/test_hit.py | 4 +- tests/models/test_assignment.py | 3 +- tests/models/test_event.py | 1 - tests/models/test_hit.py | 1 + tests_sandbox/__init__.py | 0 43 files changed, 312 insertions(+), 407 deletions(-) delete mode 100644 tests_sandbox/__init__.py (limited to 'requirements.txt') diff --git a/jb/api/auth.py b/jb/api/auth.py index a92515d..411b8f1 100644 --- a/jb/api/auth.py +++ b/jb/api/auth.py @@ -4,7 +4,7 @@ from typing import Annotated from uuid import uuid4 import jwt -from fastapi import Depends, HTTPException, Request, Response, status +from fastapi import Depends, HTTPException, Request, status from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer from jb.config import settings @@ -87,4 +87,3 @@ def create_session(product_user_id: str) -> str: settings.session_jwt_secret.get_secret_value(), algorithm="HS256", ) - diff --git a/jb/api/magic_token.py b/jb/api/magic_token.py index 562ba81..b7f7575 100644 --- a/jb/api/magic_token.py +++ b/jb/api/magic_token.py @@ -4,7 +4,7 @@ import secrets from fastapi import HTTPException, status from jb.decorators import REDIS -from jb.models.auth import AmtAccountLink, User +from jb.models.auth import AmtAccountLink MAGIC_TOKEN_PREFIX = "auth:magic:" AMT_ACCOUNT_LINK_TOKEN_PREFIX = "auth:amt-account-link:" @@ -60,9 +60,7 @@ def create_amt_account_link_token(email: str, amt_worker_id: str) -> str: def consume_amt_account_link_token(token: str) -> AmtAccountLink: """Atomically consume and validate an AMT account-link token.""" - raw_data = REDIS.getdel( - redis_token_key(token, AMT_ACCOUNT_LINK_TOKEN_PREFIX) - ) + raw_data = REDIS.getdel(redis_token_key(token, AMT_ACCOUNT_LINK_TOKEN_PREFIX)) if raw_data is None: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, 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 334eab3..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 generalresearch.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 6c5d6bb..24e96d4 100644 --- a/jb/flow/tasks.py +++ b/jb/flow/tasks.py @@ -4,11 +4,11 @@ from typing import TypedDict, cast from generalresearch.config import is_debug -from jb.decorators import AMTM, HTM, HM, HQM, pg_config +from jb.decorators import AMTM, HM, HQM, HTM, 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 +from jb.models.hit import Hit, HitQuestion, HitType logging.basicConfig() logger = logging.getLogger() diff --git a/jb/main.py b/jb/main.py index 70c98e9..9f4f000 100644 --- a/jb/main.py +++ b/jb/main.py @@ -1,15 +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.views.auth import auth_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=[ @@ -38,7 +38,7 @@ 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 {} diff --git a/jb/managers/__init__.py b/jb/managers/__init__.py index ec64f9f..92ba8bd 100644 --- a/jb/managers/__init__.py +++ b/jb/managers/__init__.py @@ -1,5 +1,5 @@ +from collections.abc import Collection from enum import IntEnum -from typing import Collection from generalresearch.pg_helper import PostgresConfig diff --git a/jb/managers/amt.py b/jb/managers/amt.py index 2cb0cbd..e2c7e90 100644 --- a/jb/managers/amt.py +++ b/jb/managers/amt.py @@ -1,6 +1,6 @@ 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 generalresearch.currency import USDCent @@ -9,8 +9,8 @@ from mypy_boto3_mturk.type_defs import ( AssignmentTypeDef, BonusPaymentTypeDef, CreateHITTypeResponseTypeDef, - GetHITResponseTypeDef, CreateHITWithHITTypeResponseTypeDef, + GetHITResponseTypeDef, ) from pydantic import ValidationError @@ -19,7 +19,7 @@ from jb.models 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 +56,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) @@ -146,7 +146,7 @@ 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: @@ -161,7 +161,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: @@ -178,7 +178,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: @@ -207,8 +207,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 +214,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, @@ -230,11 +228,9 @@ class AMTManager: logging.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 +264,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..089adb1 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.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/bonus.py b/jb/managers/bonus.py index 89b81f0..15d0e5b 100644 --- a/jb/managers/bonus.py +++ b/jb/managers/bonus.py @@ -1,4 +1,4 @@ -from typing import List, Any +from typing import Any from psycopg import sql @@ -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 index dcbe167..e740e20 100644 --- a/jb/managers/email_manager.py +++ b/jb/managers/email_manager.py @@ -21,7 +21,7 @@ def get_or_create_contact(email: str, amt_worker_id: str | None = None): return contact_id -def send_login_email_from_url(mautic_url, magic_link) -> None: +def send_login_email_from_url(mautic_url: str, magic_link: str) -> None: email_tokens = { "magic_link": magic_link, } @@ -48,7 +48,8 @@ def send_login_email(email: str, magic_token: str): def send_amt_link_email(email: str, magic_token: str): - # don't actually associate the email with the worker ID until they click the link + # 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" diff --git a/jb/managers/gr_api.py b/jb/managers/gr_api.py index 20dc051..494ceea 100644 --- a/jb/managers/gr_api.py +++ b/jb/managers/gr_api.py @@ -66,7 +66,7 @@ class GRApiManager: "product_user_id": res["product_user_id"], "email": res["metadata"].get("email_address"), "display_name": res["metadata"].get("display_name"), - 'blocked': res['blocked'], + "blocked": res["blocked"], } ) @@ -122,9 +122,7 @@ class GRApiManager: 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} - ) + 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: diff --git a/jb/managers/hit.py b/jb/managers/hit.py index 63af2d4..a178d4d 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.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) diff --git a/jb/managers/thl.py b/jb/managers/thl.py index e50fe76..85e0697 100644 --- a/jb/managers/thl.py +++ b/jb/managers/thl.py @@ -1,25 +1,18 @@ +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 ( - CashoutRequestResponse, CashoutRequestInfo, + CashoutRequestResponse, ) -from generalresearch.models.thl.user_profile import UserProfile -from generalresearch.currency import USDCent - from jb.config import settings - -from generalresearch.models.thl.definitions import PayoutStatus - - -from typing import Optional -import requests - from jb.models.auth import User -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": @@ -67,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..7fe23a7 100644 --- a/jb/models/__init__.py +++ b/jb/models/__init__.py @@ -1,7 +1,6 @@ from decimal import Decimal -from typing import Optional -from pydantic import BaseModel, Field, ConfigDict +from pydantic import BaseModel, ConfigDict, Field class HTTPHeaders(BaseModel): @@ -12,7 +11,7 @@ class HTTPHeaders(BaseModel): # 'Mon, 15 Jan 2024 23:40:32 GMT' date: str = Field() - connection: Optional[str] = Field(default=None) # 'close' + connection: str | None = Field(default=None) # 'close' class ResponseMetadata(BaseModel): diff --git a/jb/models/assignment.py b/jb/models/assignment.py index 92e5a89..775cd63 100644 --- a/jb/models/assignment.py +++ b/jb/models/assignment.py @@ -1,17 +1,17 @@ import logging from datetime import datetime, timezone -from typing import Optional, TypedDict, Any +from typing import Any, TypedDict from xml.etree import ElementTree 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 @@ -36,8 +36,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 +50,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 +96,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 +123,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 --- @@ -173,7 +173,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/bonus.py b/jb/models/bonus.py index 5f81add..2c1d00c 100644 --- a/jb/models/bonus.py +++ b/jb/models/bonus.py @@ -1,9 +1,9 @@ -from typing import Optional, Dict, Any +from typing import Any -from pydantic import BaseModel, Field, ConfigDict, PositiveInt +from generalresearch.currency import USDCent +from pydantic import BaseModel, ConfigDict, Field, PositiveInt from typing_extensions import Self -from generalresearch.currency import USDCent from jb.models.custom_types import AMTBoto3ID, AwareDatetimeISO, UUIDStr @@ -20,8 +20,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 +40,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..a58dcb7 100644 --- a/jb/models/custom_types.py +++ b/jb/models/custom_types.py @@ -1,19 +1,18 @@ import re from datetime import datetime, timezone -from typing import Any, Optional +from typing import Annotated, Any from uuid import UUID from pydantic import ( AwareDatetime, + HttpUrl, 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: @@ -22,7 +21,7 @@ def convert_datetime_to_iso_8601_with_z_suffix(dt: datetime) -> str: return dt.strftime("%Y-%m-%dT%H:%M:%S.%fZ") -def convert_str_dt(v: Any) -> Optional[AwareDatetime]: +def convert_str_dt(v: Any) -> AwareDatetime | None: # 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: diff --git a/jb/models/errors.py b/jb/models/errors.py index 94f5fbb..c590c6a 100644 --- a/jb/models/errors.py +++ b/jb/models/errors.py @@ -1,7 +1,7 @@ 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 diff --git a/jb/models/event.py b/jb/models/event.py index f8867c0..0016ca7 100644 --- a/jb/models/event.py +++ b/jb/models/event.py @@ -1,9 +1,9 @@ -from typing import Dict, Any +from typing import Any 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, AwareDatetimeISO class MTurkEvent(BaseModel): @@ -29,7 +29,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 45478fc..f6c854f 100644 --- a/jb/models/hit.py +++ b/jb/models/hit.py @@ -1,25 +1,25 @@ -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 mypy_boto3_mturk.type_defs import HITTypeDef from pydantic import ( BaseModel, - Field, - PositiveInt, ConfigDict, + Field, NonNegativeInt, + PositiveInt, ) from typing_extensions import Self -from generalresearch.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, AwareDatetimeISO, HttpsUrlStr +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 +33,7 @@ class HitQuestion(BaseModel): def xml(self) -> str: return f""" - {str(self.url)} + {self.url!s} {self.height} """ @@ -82,8 +82,8 @@ 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) @@ -104,11 +104,11 @@ 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]: + def generate_hit_amt_request(self, question: HitQuestion) -> dict[str, Any]: d = dict() d["HITTypeId"] = self.amt_hit_type_id d["MaxAssignments"] = 1 @@ -124,9 +124,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 +138,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 +153,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 @@ -235,7 +233,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) diff --git a/jb/settings.py b/jb/settings.py index 86c8a36..7747afc 100644 --- a/jb/settings.py +++ b/jb/settings.py @@ -1,10 +1,9 @@ import os from functools import lru_cache from pathlib import Path -from typing import Optional from generalresearch.models.custom_types import InfluxDsn -from pydantic import Field, PostgresDsn, HttpUrl, RedisDsn, SecretStr +from pydantic import Field, HttpUrl, PostgresDsn, RedisDsn, SecretStr from pydantic_settings import BaseSettings, SettingsConfigDict from jb.models.custom_types import UUIDStr @@ -18,14 +17,14 @@ BASE_HTML = BASE_HTML_PATH.read_text() class AmtJbBaseSettings(BaseSettings): debug: bool = Field(default=True) - redis: Optional[RedisDsn] = Field(default=None) + redis: RedisDsn | None = 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) + amt_endpoint: HttpUrl | None = Field(default=None) + amt_access_id: str | None = Field(default=None) + amt_secret_key: str | None = Field(default=None) aws_owner_id: str = Field() aws_subscription_arn: str = Field() @@ -45,11 +44,11 @@ class Settings(AmtJbBaseSettings): 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() - influx_db: Optional[InfluxDsn] = Field(default=None) + influx_db: InfluxDsn | None = Field(default=None) sns_path: str = Field() @@ -58,9 +57,7 @@ class Settings(AmtJbBaseSettings): magic_token_salt: SecretStr = Field(min_length=32) - gr_api_host: HttpUrl = Field( - default=HttpUrl("https://generalresearch.com/api/v2/") - ) + gr_api_host: HttpUrl = Field(default=HttpUrl("https://generalresearch.com/api/v2/")) gr_api_token: SecretStr = Field(min_length=1) mautic_api_key: SecretStr = Field(min_length=32) diff --git a/requirements.txt b/requirements.txt index 315a711..adfe6e9 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ -git+ssh://code.g-r-l.com:6611/generalresearch@v3.4.5 +git+ssh://code.g-r-l.com:6611/generalresearch@v3.4.7 aiohappyeyeballs==2.6.1 aiohttp==3.13.0 aiosignal==1.4.0 diff --git a/tests/__init__.py b/tests/__init__.py index e60faf0..0166072 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -1,7 +1,7 @@ -import random -import string +from random import choices as rand_choices +from string import ascii_uppercase, digits def generate_amt_id(length: int = 30) -> str: - chars = string.ascii_uppercase + string.digits - return "".join(random.choices(chars, k=length)) + chars = ascii_uppercase + digits + return "".join(rand_choices(chars, k=length)) diff --git a/tests/conftest.py b/tests/conftest.py index c056821..2a3a580 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,14 +1,16 @@ import os from typing import TYPE_CHECKING from uuid import uuid4 -from dotenv import load_dotenv + import pytest -from generalresearch.pg_helper import PostgresConfig -from tests import generate_amt_id from _pytest.config import Config -from jb.decorators import CLIENT_CONFIG +from dotenv import load_dotenv +from generalresearch.pg_helper import PostgresConfig from mypy_boto3_mturk import MTurkClient +from jb.decorators import CLIENT_CONFIG +from tests import generate_amt_id + if TYPE_CHECKING: from jb.settings import Settings diff --git a/tests/fixtures/amt.py b/tests/fixtures/amt.py index 65df125..d7cc03d 100644 --- a/tests/fixtures/amt.py +++ b/tests/fixtures/amt.py @@ -1,17 +1,18 @@ -import pytest import copy - +from collections.abc import Callable from datetime import datetime, timedelta -from typing import Callable from uuid import uuid4 + +import pytest from dateutil.tz import tzlocal from mypy_boto3_mturk.type_defs import ( - GetHITResponseTypeDef, CreateHITTypeResponseTypeDef, - ResponseMetadataTypeDef, CreateHITWithHITTypeResponseTypeDef, GetAssignmentResponseTypeDef, + GetHITResponseTypeDef, + ResponseMetadataTypeDef, ) + from jb.managers.amt import APPROVAL_MESSAGE, NO_WORK_APPROVAL_MESSAGE from tests import generate_amt_id diff --git a/tests/fixtures/flow.py b/tests/fixtures/flow.py index 3fcca81..7bf8b4f 100644 --- a/tests/fixtures/flow.py +++ b/tests/fixtures/flow.py @@ -1,18 +1,21 @@ -from datetime import timezone, datetime -from typing import Dict, Callable, Any, Optional +from collections.abc import Callable +from datetime import datetime, timezone +from typing import Any from uuid import uuid4 import pytest 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.wallet.definitions import PayoutType from generalresearch.models.thl.wallet.cashout_method import ( - CashoutRequestResponse, CashoutRequestInfo, + CashoutRequestResponse, ) +from generalresearch.models.thl.wallet.definitions import PayoutType from mypy_boto3_mturk.type_defs import ( - GetHITResponseTypeDef, GetAssignmentResponseTypeDef, + GetHITResponseTypeDef, ) from jb.config import settings @@ -20,8 +23,6 @@ from jb.managers.amt import ( APPROVAL_MESSAGE, BONUS_MESSAGE, ) -from generalresearch.currency import USDCent -from generalresearch.models.thl.definitions import PayoutStatus @pytest.fixture @@ -31,16 +32,16 @@ def approved_assignment_stubs( amt_assignment_id: str, amt_hit_id: str, hit_response_reviewing: GetHITResponseTypeDef, -) -> Callable[..., list[Dict[str, Any]]]: +) -> Callable[..., list[dict[str, Any]]]: # These are the AMT_CLIENT stubs/mocks that need to be set when running # process_assignment_submitted() which will result in an approved # assignment and sent bonus def _inner( feedback: str = APPROVAL_MESSAGE, - override_response: Optional[str] = None, - override_approve_response: Optional[str] = None, - ) -> list[Dict[str, Any]]: + override_response: str | None = None, + override_approve_response: str | None = None, + ) -> list[dict[str, Any]]: response = override_response or assignment_response approve_response = ( @@ -84,11 +85,11 @@ def approved_assignment_stubs( @pytest.fixture def approved_assignment_stubs_w_bonus( - approved_assignment_stubs: Callable[..., list[Dict[str, Any]]], + approved_assignment_stubs: Callable[..., list[dict[str, Any]]], amt_worker_id: str, amt_assignment_id: str, pe_id: str, -) -> list[Dict[str, Any]]: +) -> list[dict[str, Any]]: now = datetime.now(tz=timezone.utc) stubs = approved_assignment_stubs().copy() @@ -132,16 +133,16 @@ def rejected_assignment_stubs( amt_assignment_id: str, amt_hit_id: str, hit_response_reviewing: GetHITResponseTypeDef, -) -> Callable[..., list[Dict[str, Any]]]: +) -> Callable[..., list[dict[str, Any]]]: # These are the AMT_CLIENT stubs/mocks that need to be set when running # process_assignment_submitted() which will result in a rejected # assignment def _inner( reject_reason: str, - override_response: Optional[str] = None, - override_reject_response: Optional[str] = None, - ) -> list[Dict[str, Any]]: + override_response: str | None = None, + override_reject_response: str | None = None, + ) -> list[dict[str, Any]]: response = override_response or assignment_response reject_response = ( @@ -233,7 +234,7 @@ def mock_thl_responses( elif url == wallet_url: class MockThlWalletResponse: - def json(self) -> Dict[str, Any]: + def json(self) -> dict[str, Any]: return { "wallet": { "amount": wallet_redeemable_amount, @@ -246,7 +247,7 @@ def mock_thl_responses( elif url == status_url: class MockThlStatusResponse: - def json(self) -> Dict[str, Any]: + def json(self) -> dict[str, Any]: return { "tsid": tsid, "product_id": str(settings.product_id), diff --git a/tests/fixtures/http.py b/tests/fixtures/http.py index 5f50580..4b0792c 100644 --- a/tests/fixtures/http.py +++ b/tests/fixtures/http.py @@ -1,20 +1,19 @@ +import json +import secrets +from collections.abc import AsyncGenerator +from typing import Any + import httpx -import redis import pytest +import redis import requests_mock from asgi_lifespan import LifespanManager -from httpx import AsyncClient, ASGITransport -from typing import Dict, Any, AsyncGenerator +from httpx import ASGITransport, AsyncClient +from jb.config import JB_EVENTS_STREAM, settings from jb.main import app -import json - -from httpx import AsyncClient -import secrets - -from jb.models.hit import Hit from jb.models.assignment import AssignmentStub -from jb.config import JB_EVENTS_STREAM, settings +from jb.models.hit import Hit from tests import generate_amt_id @@ -69,7 +68,7 @@ def generate_hex_id(length: int = 40) -> str: @pytest.fixture def mturk_event_body_record( hit_record: Hit, assignment_stub_record: AssignmentStub -) -> Dict[str, Any]: +) -> dict[str, Any]: return { "Type": "Notification", "Message": json.dumps( diff --git a/tests/fixtures/managers.py b/tests/fixtures/managers.py index d10b542..22eae5e 100644 --- a/tests/fixtures/managers.py +++ b/tests/fixtures/managers.py @@ -1,14 +1,16 @@ from typing import TYPE_CHECKING + import pytest -from jb.managers import Permission from generalresearch.pg_helper import PostgresConfig from mypy_boto3_mturk import MTurkClient +from jb.managers import Permission + if TYPE_CHECKING: - from jb.managers.hit import HitQuestionManager, HitTypeManager, HitManager + from jb.managers.amt import AMTManager from jb.managers.assignment import AssignmentManager from jb.managers.bonus import BonusManager - from jb.managers.amt import AMTManager + from jb.managers.hit import HitManager, HitQuestionManager, HitTypeManager # --- Managers --- diff --git a/tests/fixtures/models.py b/tests/fixtures/models.py index b818caa..157daec 100644 --- a/tests/fixtures/models.py +++ b/tests/fixtures/models.py @@ -1,23 +1,22 @@ -from datetime import timezone, datetime +from collections.abc import Callable, Generator +from datetime import datetime, timedelta, timezone +from typing import TYPE_CHECKING import pytest - -from jb.models.event import MTurkEvent +from generalresearch.currency import USDCent from generalresearch.pg_helper import PostgresConfig +from psycopg.errors import ForeignKeyViolation -from datetime import datetime, timezone, timedelta -from typing import Optional, TYPE_CHECKING, Callable, Generator from jb.managers.amt import AMTManager -from jb.models.assignment import AssignmentStub, Assignment -from generalresearch.currency import USDCent -from jb.models.definitions import HitStatus, HitReviewStatus, AssignmentStatus -from jb.models.hit import HitType, HitQuestion, Hit +from jb.models.assignment import Assignment, AssignmentStub +from jb.models.definitions import AssignmentStatus, HitReviewStatus, HitStatus +from jb.models.event import MTurkEvent +from jb.models.hit import Hit, HitQuestion, HitType from tests import generate_amt_id -from psycopg.errors import ForeignKeyViolation if TYPE_CHECKING: - from jb.managers.hit import HitQuestionManager, HitTypeManager, HitManager from jb.managers.assignment import AssignmentManager + from jb.managers.hit import HitManager, HitQuestionManager, HitTypeManager # --- MTurk Event --- @@ -76,10 +75,9 @@ def hit_type_record( yield ht try: - with pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("DELETE FROM mtwerk_hittype WHERE id=%s", (ht.id,)) - conn.commit() + with pg_config.make_connection() as conn, conn.cursor() as c: + c.execute("DELETE FROM mtwerk_hittype WHERE id=%s", (ht.id,)) + conn.commit() except ForeignKeyViolation: pass # DB gets dropped anyway, don't care @@ -99,10 +97,9 @@ def hit_type_record_with_amt_id( yield ht try: - with pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("DELETE FROM mtwerk_hittype WHERE id=%s", (ht.id,)) - conn.commit() + with pg_config.make_connection() as conn, conn.cursor() as c: + c.execute("DELETE FROM mtwerk_hittype WHERE id=%s", (ht.id,)) + conn.commit() except ForeignKeyViolation: pass # DB gets dropped anyway, don't care @@ -166,10 +163,9 @@ def hit_record( yield hit try: - with pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("DELETE FROM mtwerk_hit WHERE id=%s", (hit.id,)) - conn.commit() + with pg_config.make_connection() as conn, conn.cursor() as c: + c.execute("DELETE FROM mtwerk_hit WHERE id=%s", (hit.id,)) + conn.commit() except ForeignKeyViolation: pass # DB gets dropped anyway, don't care @@ -228,12 +224,11 @@ def assignment_stub_record( yield assignment_stub try: - with pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - "DELETE FROM mtwerk_assignment WHERE id=%s", (assignment_stub.id,) - ) - conn.commit() + with pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + "DELETE FROM mtwerk_assignment WHERE id=%s", (assignment_stub.id,) + ) + conn.commit() except ForeignKeyViolation: pass # DB gets dropped anyway, don't care @@ -255,19 +250,18 @@ def assignment_record( yield assignment try: - with pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("DELETE FROM mtwerk_assignment WHERE id=%s", (assignment.id,)) - conn.commit() + with pg_config.make_connection() as conn, conn.cursor() as c: + c.execute("DELETE FROM mtwerk_assignment WHERE id=%s", (assignment.id,)) + conn.commit() except ForeignKeyViolation: pass # DB gets dropped anyway, don't care @pytest.fixture -def assignment_factory(hit_record: Hit) -> Callable[[Optional[str]], Assignment]: +def assignment_factory(hit_record: Hit) -> Callable[[str | None], Assignment]: - def _inner(amt_worker_id: Optional[str] = None) -> Assignment: + def _inner(amt_worker_id: str | None = None) -> Assignment: now = datetime.now(tz=timezone.utc) amt_assignment_id = generate_amt_id() amt_worker_id = amt_worker_id or generate_amt_id() @@ -292,7 +286,7 @@ def assignment_record_factory( am: "AssignmentManager", assignment_factory: Callable[..., Assignment] ) -> Callable[..., Assignment]: - def _inner(hit_id: int, amt_worker_id: Optional[str] = None) -> Assignment: + def _inner(hit_id: int, amt_worker_id: str | None = None) -> Assignment: a = assignment_factory(amt_worker_id=amt_worker_id) a.hit_id = hit_id am.create_stub(a) diff --git a/tests/flow/test_tasks.py b/tests/flow/test_tasks.py index f939d20..9cc111f 100644 --- a/tests/flow/test_tasks.py +++ b/tests/flow/test_tasks.py @@ -1,34 +1,36 @@ import logging +from collections.abc import Callable from contextlib import contextmanager -from typing import Callable, Dict, Any +from typing import Any import pytest from botocore.stub import Stubber +from generalresearch.currency import USDCent +from mypy_boto3_mturk import MTurkClient +from mypy_boto3_mturk.type_defs import ( + GetAssignmentResponseTypeDef, +) + from jb.flow.assignment_tasks import process_assignment_submitted from jb.managers.amt import ( - AMTManager, APPROVAL_MESSAGE, + NO_WORK_APPROVAL_MESSAGE, REJECT_MESSAGE_BADDIE, - REJECT_MESSAGE_UNKNOWN_ASSIGNMENT, REJECT_MESSAGE_NO_WORK, - NO_WORK_APPROVAL_MESSAGE, -) -from mypy_boto3_mturk.type_defs import ( - GetAssignmentResponseTypeDef, + REJECT_MESSAGE_UNKNOWN_ASSIGNMENT, + AMTManager, ) -from generalresearch.currency import USDCent from jb.managers.assignment import AssignmentManager from jb.managers.bonus import BonusManager from jb.managers.hit import HitManager +from jb.models.assignment import Assignment, AssignmentStub from jb.models.definitions import AssignmentStatus from jb.models.event import MTurkEvent from jb.models.hit import Hit -from jb.models.assignment import Assignment, AssignmentStub -from mypy_boto3_mturk import MTurkClient @contextmanager -def amt_stub_context(amt_client: MTurkClient, responses: list[Dict[str, Any]]): +def amt_stub_context(amt_client: MTurkClient, responses: list[dict[str, Any]]): # ty chatgpt for this with Stubber(amt_client) as stub: @@ -83,7 +85,7 @@ class TestProcessAssignmentSubmitted: mturk_event: MTurkEvent, amt_assignment_id: str, caplog: pytest.LogCaptureFixture, - rejected_assignment_stubs: Callable[..., list[Dict[str, Any]]], + rejected_assignment_stubs: Callable[..., list[dict[str, Any]]], ): # These records are auto cleaned up, so we need to explicitly create @@ -108,8 +110,8 @@ class TestProcessAssignmentSubmitted: stub.assert_no_pending_responses() assert f"No assignment found in DB: {amt_assignment_id}" in caplog.text - assert f"Rejected assignment doesn't exist in DB. Creating ... " in caplog.text - assert f"Rejected assignment: " in caplog.text + assert "Rejected assignment doesn't exist in DB. Creating ... " in caplog.text + assert "Rejected assignment: " in caplog.text stub.assert_no_pending_responses() ass = am.get(amt_assignment_id=amt_assignment_id) @@ -128,7 +130,7 @@ class TestProcessAssignmentSubmitted: assignment_stub_record: AssignmentStub, caplog: pytest.LogCaptureFixture, mock_thl_responses: Callable[..., None], - rejected_assignment_stubs: Callable[..., list[Dict[str, Any]]], + rejected_assignment_stubs: Callable[..., list[dict[str, Any]]], ): # An assignment is submitted. The hit and AssignmentStub exist in the # DB. We think we're going to approve the Assignment, but the @@ -152,8 +154,8 @@ class TestProcessAssignmentSubmitted: stub.assert_no_pending_responses() assert f"No assignment found in DB: {amt_assignment_id}" not in caplog.text - assert f"blocked or not exists" in caplog.text - assert f"Rejected assignment: " in caplog.text + assert "blocked or not exists" in caplog.text + assert "Rejected assignment: " in caplog.text ass = am.get(amt_assignment_id=amt_assignment_id) assert ass.status == AssignmentStatus.Rejected @@ -171,7 +173,7 @@ class TestProcessAssignmentSubmitted: assignment_stub_record: AssignmentStub, caplog: pytest.LogCaptureFixture, mock_thl_responses: Callable[..., None], - approved_assignment_stubs: Callable[..., list[Dict[str, Any]]], + approved_assignment_stubs: Callable[..., list[dict[str, Any]]], assignment_response_approved_no_tsid: GetAssignmentResponseTypeDef, assignment_response_no_tsid: GetAssignmentResponseTypeDef, ): @@ -202,8 +204,8 @@ class TestProcessAssignmentSubmitted: stub.assert_no_pending_responses() assert f"No assignment found in DB: {amt_assignment_id}" not in caplog.text - assert f"Assignment submitted with no tsid" in caplog.text - assert f"Approved assignment: " in caplog.text + assert "Assignment submitted with no tsid" in caplog.text + assert "Approved assignment: " in caplog.text ass = am.get(amt_assignment_id=amt_assignment_id) assert ass.status == AssignmentStatus.Approved @@ -222,7 +224,7 @@ class TestProcessAssignmentSubmitted: assignment_stub_record: AssignmentStub, caplog: pytest.LogCaptureFixture, mock_thl_responses: Callable[..., None], - rejected_assignment_stubs: Callable[..., list[Dict[str, Any]]], + rejected_assignment_stubs: Callable[..., list[dict[str, Any]]], assignment_response_factory_rejected_no_tsid: Callable[ ..., GetAssignmentResponseTypeDef ], @@ -273,8 +275,8 @@ class TestProcessAssignmentSubmitted: stub.assert_no_pending_responses() assert f"No assignment found in DB: {amt_assignment_id}" not in caplog.text - assert f"Assignment submitted with no tsid" in caplog.text - assert f"Rejected assignment: " in caplog.text + assert "Assignment submitted with no tsid" in caplog.text + assert "Rejected assignment: " in caplog.text # It will exist in the db since we can validate the model. ass = am.get(amt_assignment_id=amt_assignment_id) @@ -293,7 +295,7 @@ class TestProcessAssignmentSubmitted: assignment_stub_record: Assignment, caplog: pytest.LogCaptureFixture, mock_thl_responses: Callable[..., None], - approved_assignment_stubs: Callable[..., list[Dict[str, Any]]], + approved_assignment_stubs: Callable[..., list[dict[str, Any]]], ): _ = assignment_stub_record # we need this to make the assignment stub in the db @@ -330,7 +332,7 @@ class TestProcessAssignmentSubmitted: assignment_stub_record: Assignment, caplog: pytest.LogCaptureFixture, mock_thl_responses: Callable[..., None], - approved_assignment_stubs_w_bonus: list[Dict[str, Any]], + approved_assignment_stubs_w_bonus: list[dict[str, Any]], ): _ = assignment_stub_record # we need this to make the assignment stub in the db mock_thl_responses(status_complete=True, wallet_redeemable_amount=10) diff --git a/tests/http/test_auth.py b/tests/http/test_auth.py index 02ac88a..ebda742 100644 --- a/tests/http/test_auth.py +++ b/tests/http/test_auth.py @@ -60,9 +60,7 @@ class TestAuth: ): client = httpxclient - res = await client.post( - "/auth/magic-link/request", json={"email": email} - ) + res = await client.post("/auth/magic-link/request", json={"email": email}) d = res.json() assert res.status_code == 200 assert d["magic_link"] diff --git a/tests/http/test_notifications.py b/tests/http/test_notifications.py index 508b236..60b94e6 100644 --- a/tests/http/test_notifications.py +++ b/tests/http/test_notifications.py @@ -1,14 +1,15 @@ -import pytest import json +from typing import Any +from uuid import uuid4 + +import pytest import redis -from typing import Dict, Any from httpx import AsyncClient -from uuid import uuid4 from jb.config import JB_EVENTS_STREAM, settings +from jb.models.assignment import AssignmentStub from jb.models.event import MTurkEvent from jb.models.hit import Hit -from jb.models.assignment import AssignmentStub class TestNotifications: @@ -55,7 +56,7 @@ class TestNotifications: httpxclient: AsyncClient, hit_record: Hit, assignment_stub_record: AssignmentStub, - mturk_event_body_record: Dict[str, Any], + mturk_event_body_record: dict[str, Any], ): client = httpxclient diff --git a/tests/http/test_preview.py b/tests/http/test_preview.py index 467c63c..39a6f5b 100644 --- a/tests/http/test_preview.py +++ b/tests/http/test_preview.py @@ -3,8 +3,9 @@ import pytest from httpx import AsyncClient -from jb.models.hit import Hit + from jb.models.assignment import AssignmentStub +from jb.models.hit import Hit class TestPreview: diff --git a/tests/http/test_work.py b/tests/http/test_work.py index 66251f6..7f10b46 100644 --- a/tests/http/test_work.py +++ b/tests/http/test_work.py @@ -1,9 +1,9 @@ import pytest from httpx import AsyncClient -from jb.models.hit import Hit -from jb.models.assignment import AssignmentStub from jb.managers.assignment import AssignmentManager +from jb.models.assignment import AssignmentStub +from jb.models.hit import Hit class TestWork: diff --git a/tests/managers/test_amt.py b/tests/managers/test_amt.py index a20d0d4..6a944a8 100644 --- a/tests/managers/test_amt.py +++ b/tests/managers/test_amt.py @@ -1,14 +1,13 @@ -from jb.managers.amt import AMTManager -from jb.models.hit import HitType, HitQuestion - -from jb.managers.hit import HitQuestionManager, HitTypeManager, HitManager from mypy_boto3_mturk import MTurkClient from mypy_boto3_mturk.type_defs import ( - GetAssignmentResponseTypeDef, GetAccountBalanceResponseTypeDef, ListHITsResponseTypeDef, ) +from jb.managers.amt import AMTManager +from jb.managers.hit import HitManager, HitTypeManager +from jb.models.hit import HitQuestion, HitType + # from jb.decorators import HM # from jb.flow.tasks import refill_hits, check_stale_hits, check_expired_hits diff --git a/tests/managers/test_hit.py b/tests/managers/test_hit.py index 974bd18..7227e4d 100644 --- a/tests/managers/test_hit.py +++ b/tests/managers/test_hit.py @@ -1,5 +1,5 @@ -from jb.models.hit import HitQuestion, HitType, Hit -from jb.managers.hit import HitTypeManager, HitManager +from jb.managers.hit import HitManager, HitTypeManager +from jb.models.hit import Hit, HitQuestion, HitType class TestHitQuestionManager: diff --git a/tests/models/test_assignment.py b/tests/models/test_assignment.py index 2a87364..ecaafe9 100644 --- a/tests/models/test_assignment.py +++ b/tests/models/test_assignment.py @@ -1,8 +1,9 @@ -from jb.models.assignment import Assignment, AssignmentStub from mypy_boto3_mturk.type_defs import ( GetAssignmentResponseTypeDef, ) +from jb.models.assignment import Assignment, AssignmentStub + class TestAssignmentStub: diff --git a/tests/models/test_event.py b/tests/models/test_event.py index 0496574..a4c591d 100644 --- a/tests/models/test_event.py +++ b/tests/models/test_event.py @@ -1,6 +1,5 @@ import pytest - from jb.models.event import MTurkEvent diff --git a/tests/models/test_hit.py b/tests/models/test_hit.py index 3952068..aa48f00 100644 --- a/tests/models/test_hit.py +++ b/tests/models/test_hit.py @@ -1,4 +1,5 @@ import pytest + from jb.models.hit import Hit diff --git a/tests_sandbox/__init__.py b/tests_sandbox/__init__.py deleted file mode 100644 index e69de29..0000000 -- cgit v1.2.3