aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--jb/api/magic_token.py34
-rw-r--r--jb/managers/gr_api.py59
-rw-r--r--jb/models/auth.py11
-rw-r--r--jb/views/auth.py58
-rw-r--r--tests/conftest.py4
-rw-r--r--tests/fixtures/flow.py10
-rw-r--r--tests/fixtures/managers.py2
-rw-r--r--tests/fixtures/models.py4
-rw-r--r--tests/flow/test_tasks.py2
-rw-r--r--tests/http/test_auth.py58
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