diff options
| -rw-r--r-- | jb/api/magic_token.py | 34 | ||||
| -rw-r--r-- | jb/managers/gr_api.py | 59 | ||||
| -rw-r--r-- | jb/models/auth.py | 11 | ||||
| -rw-r--r-- | jb/views/auth.py | 58 | ||||
| -rw-r--r-- | tests/conftest.py | 4 | ||||
| -rw-r--r-- | tests/fixtures/flow.py | 10 | ||||
| -rw-r--r-- | tests/fixtures/managers.py | 2 | ||||
| -rw-r--r-- | tests/fixtures/models.py | 4 | ||||
| -rw-r--r-- | tests/flow/test_tasks.py | 2 | ||||
| -rw-r--r-- | tests/http/test_auth.py | 58 |
10 files changed, 215 insertions, 27 deletions
diff --git a/jb/api/magic_token.py b/jb/api/magic_token.py index 136e1b4..e0a1cca 100644 --- a/jb/api/magic_token.py +++ b/jb/api/magic_token.py @@ -4,15 +4,17 @@ import secrets from fastapi import HTTPException, status from jb.decorators import REDIS +from jb.models.auth import AmtAccountLink, User MAGIC_TOKEN_PREFIX = "auth:magic:" +AMT_ACCOUNT_LINK_TOKEN_PREFIX = "auth:amt-account-link:" MAGIC_TOKEN_TTL: int = 5 * 60 # 5 minutes, in seconds -def redis_token_key(token: str) -> str: +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"{MAGIC_TOKEN_PREFIX}{digest}" + return f"{prefix}{digest}" def create_magic_token(user_email: str) -> str: @@ -39,3 +41,31 @@ def consume_magic_token(token: str) -> str: detail="Invalid or expired magic token", ) return user_email + + +def create_amt_account_link_token(user: User, amt_worker_id: str) -> str: + """Bind an email and AMT worker ID to an opaque, short-lived token.""" + data = AmtAccountLink( + email=user.email, + amt_worker_id=amt_worker_id, + ) + token = secrets.token_urlsafe(32) + REDIS.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) -> AmtAccountLink: + """Atomically consume and validate an AMT account-link token.""" + 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, + detail="Invalid or expired account-link token", + ) + return AmtAccountLink.model_validate_json(raw_data) diff --git a/jb/managers/gr_api.py b/jb/managers/gr_api.py index 8ffeaf0..a31f7da 100644 --- a/jb/managers/gr_api.py +++ b/jb/managers/gr_api.py @@ -1,16 +1,23 @@ """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(__name__) + 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, @@ -40,6 +47,14 @@ class GRApiManager: 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}" @@ -101,12 +116,46 @@ class GRApiManager: ).json() return self.get_user(user.product_user_id) - def transition_product_user_id(self, user: User, amt_worker_id: str) -> User: + 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}/" - res = self._request( - "PATCH", url, json={"product_user_id": user.product_user_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) - return self.get_user(user.product_user_id) + 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/models/auth.py b/jb/models/auth.py index 8b48088..d0e8386 100644 --- a/jb/models/auth.py +++ b/jb/models/auth.py @@ -68,7 +68,9 @@ class User(BaseModel): 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}") + 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) @@ -92,6 +94,13 @@ class MagicLinkExchangeRequest(BaseModel): 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" diff --git a/jb/views/auth.py b/jb/views/auth.py index 33c9452..dc63ca3 100644 --- a/jb/views/auth.py +++ b/jb/views/auth.py @@ -1,10 +1,3 @@ -"""Redis-backed magic-link authentication. - -The user service is deliberately not coupled to this module. Once that service -has resolved an email address to its stable user identifier, call -``create_magic_token`` and put the returned token in the emailed login URL. -""" - from typing import Annotated from urllib.parse import urlencode @@ -16,12 +9,18 @@ from jb.api.auth import ( create_session, get_authenticated_user, ) -from jb.api.magic_token import consume_magic_token, create_magic_token +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.dependencies import get_gr_api_manager from jb.managers.gr_api import GRApiManager from jb.models.auth import ( AccountLogin, + AmtAccountLink, MagicLinkExchangeRequest, User, ) @@ -44,6 +43,20 @@ def request_mock_magic_link(body: AccountLogin) -> dict[str, str]: return {"magic_link": f"/auth/magic-link/?{query}"} +@auth_router.post("/link-amt/request") +def link_amt_account(body: AmtAccountLink) -> dict[str, str]: + """Create a mock AMT account-link email in development.""" + if not settings.debug: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND) + + # TODO: Derive amt_worker_id from a server-validated AMT assignment and + # send this link by email instead of returning it. + user = User(email=body.email) + token = create_amt_account_link_token(user, body.amt_worker_id) + query = urlencode({"token": token}) + return {"magic_link": f"/auth/link-amt/?{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.""" @@ -57,6 +70,12 @@ def magic_link_landing_page() -> HTMLResponse: ) +@auth_router.get("/link-amt/", response_class=HTMLResponse, include_in_schema=False) +def link_amt_account_landing_page() -> HTMLResponse: + """Serve the account-link SPA without consuming the one-time token.""" + return magic_link_landing_page() + + @auth_router.post("/magic-link/exchange", status_code=status.HTTP_204_NO_CONTENT) def exchange_magic_link( body: MagicLinkExchangeRequest, @@ -82,6 +101,29 @@ def exchange_magic_link( ) +@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.""" + token_data = consume_amt_account_link_token(body.token) + user = User(email=token_data.email) + user = gr_api.transition_user_from_amt(user, token_data.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)], diff --git a/tests/conftest.py b/tests/conftest.py index 25eb457..c056821 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,7 +3,7 @@ from typing import TYPE_CHECKING from uuid import uuid4 from dotenv import load_dotenv import pytest -from generalresearchutils.pg_helper import PostgresConfig +from generalresearch.pg_helper import PostgresConfig from tests import generate_amt_id from _pytest.config import Config from jb.decorators import CLIENT_CONFIG @@ -99,7 +99,7 @@ def settings(env_file_path: str) -> "Settings": @pytest.fixture(scope="session") def redis(settings: "Settings"): - from generalresearchutils.redis_helper import RedisConfig + from generalresearch.redis_helper import RedisConfig redis_config = RedisConfig( dsn=settings.redis, diff --git a/tests/fixtures/flow.py b/tests/fixtures/flow.py index dd2f83e..08ec49e 100644 --- a/tests/fixtures/flow.py +++ b/tests/fixtures/flow.py @@ -4,9 +4,9 @@ from uuid import uuid4 import pytest import requests -from generalresearchutils.models.thl.payout import UserPayoutEvent -from generalresearchutils.models.thl.wallet import PayoutType -from generalresearchutils.models.thl.wallet.cashout_method import ( +from generalresearch.models.thl.payout import UserPayoutEvent +from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.cashout_method import ( CashoutRequestResponse, CashoutRequestInfo, ) @@ -20,8 +20,8 @@ from jb.managers.amt import ( APPROVAL_MESSAGE, BONUS_MESSAGE, ) -from generalresearchutils.currency import USDCent -from generalresearchutils.models.thl.definitions import PayoutStatus +from generalresearch.currency import USDCent +from generalresearch.models.thl.definitions import PayoutStatus @pytest.fixture diff --git a/tests/fixtures/managers.py b/tests/fixtures/managers.py index a3187d7..d10b542 100644 --- a/tests/fixtures/managers.py +++ b/tests/fixtures/managers.py @@ -1,7 +1,7 @@ from typing import TYPE_CHECKING import pytest from jb.managers import Permission -from generalresearchutils.pg_helper import PostgresConfig +from generalresearch.pg_helper import PostgresConfig from mypy_boto3_mturk import MTurkClient if TYPE_CHECKING: diff --git a/tests/fixtures/models.py b/tests/fixtures/models.py index 671c7b3..b818caa 100644 --- a/tests/fixtures/models.py +++ b/tests/fixtures/models.py @@ -3,13 +3,13 @@ from datetime import timezone, datetime import pytest from jb.models.event import MTurkEvent -from generalresearchutils.pg_helper import PostgresConfig +from generalresearch.pg_helper import PostgresConfig 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 generalresearchutils.currency import USDCent +from generalresearch.currency import USDCent from jb.models.definitions import HitStatus, HitReviewStatus, AssignmentStatus from jb.models.hit import HitType, HitQuestion, Hit from tests import generate_amt_id diff --git a/tests/flow/test_tasks.py b/tests/flow/test_tasks.py index 3a71504..f939d20 100644 --- a/tests/flow/test_tasks.py +++ b/tests/flow/test_tasks.py @@ -16,7 +16,7 @@ from jb.managers.amt import ( from mypy_boto3_mturk.type_defs import ( GetAssignmentResponseTypeDef, ) -from generalresearchutils.currency import USDCent +from generalresearch.currency import USDCent from jb.managers.assignment import AssignmentManager from jb.managers.bonus import BonusManager from jb.managers.hit import HitManager diff --git a/tests/http/test_auth.py b/tests/http/test_auth.py index 8a625bf..02ac88a 100644 --- a/tests/http/test_auth.py +++ b/tests/http/test_auth.py @@ -13,6 +13,8 @@ from jb.models.auth import User class FakeGRApiManager: def __init__(self): self.users: dict[str, User] = {} + self.amt_worker_ids: set[str] = set() + self.transition_calls: list[tuple[User, str]] = [] def ensure_user_exists(self, user: User) -> User: self.users.setdefault(user.product_user_id, user) @@ -21,6 +23,18 @@ class FakeGRApiManager: def get_user(self, product_user_id: str) -> User: return self.users[product_user_id] + def add_amt_user(self, amt_worker_id: str) -> None: + self.amt_worker_ids.add(amt_worker_id) + + def transition_user_from_amt(self, user: User, amt_worker_id: str) -> User: + self.transition_calls.append((user, amt_worker_id)) + if amt_worker_id not in self.amt_worker_ids: + raise ValueError(f"User {amt_worker_id} does not exist") + + self.amt_worker_ids.remove(amt_worker_id) + self.users[user.product_user_id] = user + return user + @pytest.fixture def email() -> str: @@ -64,3 +78,47 @@ class TestAuth: res = await client.get("/auth/session") assert res.status_code == 200 assert res.json()["email"] == email + + @pytest.mark.anyio + async def test_amt_account_link( + self, + httpxclient: AsyncClient, + fake_gr_api_manager: FakeGRApiManager, + email: str, + amt_worker_id: str, + ): + client = httpxclient + fake_gr_api_manager.add_amt_user(amt_worker_id) + + res = await client.post( + "/auth/link-amt/request", + json={"email": email, "amt_worker_id": amt_worker_id}, + ) + d = res.json() + assert res.status_code == 200 + assert d["magic_link"] + assert fake_gr_api_manager.transition_calls == [] + + token = parse_qs(urlparse(d["magic_link"]).query)["token"][0] + + url = "/auth/link-amt/exchange" + body = {"token": token} + res = await client.post(url, json=body) + assert res.status_code == 204 + assert client.cookies.get(SESSION_COOKIE_NAME) + assert len(fake_gr_api_manager.transition_calls) == 1 + transitioned_user, transitioned_amt_worker_id = ( + fake_gr_api_manager.transition_calls[0] + ) + assert transitioned_user.email == email + assert transitioned_amt_worker_id == amt_worker_id + assert amt_worker_id not in fake_gr_api_manager.amt_worker_ids + + res = await client.get("/auth/session") + assert res.status_code == 200 + assert res.json()["email"] == email + + # Make sure we can't do it again + res = await client.post(url, json=body) + assert res.status_code == 401 + assert len(fake_gr_api_manager.transition_calls) == 1 |
