diff --git a/auth/validator.py b/auth/validator.py index a0491fe4..8bb80025 100644 --- a/auth/validator.py +++ b/auth/validator.py @@ -2,6 +2,7 @@ import json import logging import weakref +from datetime import UTC, datetime, timedelta import httpx import jwt @@ -168,3 +169,23 @@ def verify_action_token(token: str, settings: Settings) -> dict: except InvalidTokenError: raise HTTPException(status_code=401, detail="invalid session_token") return payload + + +def create_action_token(payload: dict, settings: Settings, expires_in_seconds: int = 300) -> dict: + """ + Create a signed JWT that can be passed back to Auth0 actions + """ + required_fields = ["sub", "iss", "state"] + for field in required_fields: + if field not in payload: + raise ValueError( f"Missing required field {field} in action token") + now = datetime.now(tz=UTC) + exp = (now + timedelta(seconds=expires_in_seconds)).timestamp() + payload = {**payload, "exp": int(exp), "iat": int(now.timestamp())} + secret = settings.auth0_management_secret + signed_payload = jwt.encode( + payload, + key=secret, + algorithm="HS256", + ) + return signed_payload diff --git a/routers/aaf.py b/routers/aaf.py index 83d13140..430b9ba3 100644 --- a/routers/aaf.py +++ b/routers/aaf.py @@ -2,11 +2,15 @@ from datetime import datetime, timezone from http import HTTPStatus from typing import Annotated +from urllib.parse import urlparse +import httpx2 from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel from sqlmodel import Session +from starlette.responses import RedirectResponse +from auth.validator import create_action_token from auth0.client import Auth0Client, UpdateUserData, get_auth0_client from config import Settings, get_settings from db.models import BiocommonsUser @@ -14,6 +18,7 @@ from dependencies.auth import require_action_token from schemas.auth0 import Auth0ActionToken from schemas.biocommons import ( + Auth0Identity, Auth0UserData, BiocommonsAppMetadataUpdate, BiocommonsUserAccountType, @@ -28,31 +33,47 @@ class AccountLinkResponse(BaseModel): link: bool aaf_only: bool = False + blocked: bool = False primary_id: str + aaf_identity: Auth0Identity | None = None -def link_aaf_account(db_user_id: str, aaf_user_id: str, auth0_client: Auth0Client, session: Session): - def _get_aaf_identity(aaf_user: Auth0UserData): +class SignedAccountLinkResponse(AccountLinkResponse): + """ + Set the fields we need a signed action token to return + to Auth0 + """ + sub: str + iss: str + state: str + + +def link_aaf_account(db_user_id: str, aaf_user_id: str, auth0_client: Auth0Client, session: Session) -> Auth0Identity: + """ + Update app_metadata and DB record for AAF account linking. + + NOTE: we can't do the actual linking here, it needs to be done + on the Auth0 side. See: https://support.auth0.com/center/s/article/Unable-to-process-redirect-callback + """ + def _get_aaf_identity(aaf_user: Auth0UserData) -> Auth0Identity | None: for identity in aaf_user.identities: if identity.connection == "AAF": return identity return None - db_user = BiocommonsUser.get_by_id_or_404(db_user_id, session=session) - if db_user.account_type == BiocommonsUserAccountType.AAF and db_user.other_user_id == aaf_user_id: - logger.info("AAF account already linked, skipping re-link") - return db_user aaf_user_info = auth0_client.get_user(user_id=aaf_user_id) aaf_identity = _get_aaf_identity(aaf_user_info) if aaf_identity is None: raise HTTPException(status_code=HTTPStatus.NOT_FOUND, detail="Couldn't get AAF provider information") - logger.info("Linking AAF account to existing database account") + db_user = BiocommonsUser.get_by_id_or_404(db_user_id, session=session) + if db_user.account_type == BiocommonsUserAccountType.AAF and db_user.other_user_id == aaf_user_id: + logger.info("AAF account already linked, skipping re-link") + return aaf_identity + + logger.info("Updating DB record and account metadata") now = datetime.now(tz=timezone.utc) try: - auth0_client.link_identity(primary_user_id=db_user_id, secondary_user_id=aaf_user_id, - secondary_provider=aaf_identity.provider, - secondary_connection_name=aaf_identity.connection) auth0_client.update_user( db_user_id, update_data=UpdateUserData( @@ -64,11 +85,11 @@ def _get_aaf_identity(aaf_user: Auth0UserData): ) ) except ValueError as exc: - logger.error(f"Failed to link AAF account in Auth0: {exc}") + logger.error(f"Failed to update account metadata in Auth0: {exc}") raise HTTPException(status_code=HTTPStatus.BAD_GATEWAY, detail="Failed to link account with Auth0") from exc logger.info("Updating DB record") db_user.link_aaf_account(aaf_user_id=aaf_user_id, session=session, updated_by=db_user, commit=True) - return db_user + return aaf_identity def mark_user_aaf_only(user_email: str, aaf_user_id: str, auth0_client: Auth0Client): @@ -83,9 +104,52 @@ def mark_user_aaf_only(user_email: str, aaf_user_id: str, auth0_client: Auth0Cli auth0_client.update_user(user_id=aaf_user_id, update_data=update_data) +def return_signed_response( + state: str, + response: AccountLinkResponse, + action_token: Auth0ActionToken, + settings: Settings, +): + """ + Auth0 Actions need to receive the response as a signed JWT + token, sign the AccountLinkResponse we want to return + and redirect to the continue endpoint + """ + signed_payload = SignedAccountLinkResponse( + **response.model_dump(), + sub=action_token.sub or action_token.user_id, + iss=action_token.iss or settings.auth0_domain, + state=state, + ) + signed_token = create_action_token( + payload=signed_payload.model_dump(mode="json", exclude_none=True), + settings=settings, + ) + auth0_base_url = get_auth0_continue_base_url(action_token, settings) + redirect_url = httpx2.URL( + f"{auth0_base_url}/continue", + params={"state": state, "session_token": signed_token}, + ) + return RedirectResponse(url=redirect_url) + + +def get_auth0_continue_base_url(action_token: Auth0ActionToken, settings: Settings) -> str: + """ + Get domain to continue action from action token's issuer, if possible - want + to ensure we continue on the same domain + """ + if action_token.iss: + issuer = action_token.iss.rstrip("/") + parsed_issuer = urlparse(issuer) + if parsed_issuer.scheme and parsed_issuer.netloc: + return f"{parsed_issuer.scheme}://{parsed_issuer.netloc}" + return f"https://{issuer}" + return settings.auth0_custom_domain or f"https://{settings.auth0_domain}" -@router.get("/check-link", response_model=AccountLinkResponse) + +@router.get("/check-link") def check_aaf_account_link( + state: str, token: Annotated[Auth0ActionToken, Depends(require_action_token(purpose="aaf_link"))], session: Annotated[Session, Depends(get_db_session)], auth0_client: Annotated[Auth0Client, Depends(get_auth0_client)], @@ -100,7 +164,13 @@ def check_aaf_account_link( # No existing account: no need to link if not auth0_matches: mark_user_aaf_only(email, aaf_user_id, auth0_client) - return AccountLinkResponse(link=False, aaf_only=True, primary_id=token.user_id) + resp = AccountLinkResponse(link=False, aaf_only=True, primary_id=aaf_user_id) + return return_signed_response( + state=state, + response=resp, + action_token=token, + settings=settings, + ) existing_account = None for user in auth0_matches: @@ -110,15 +180,33 @@ def check_aaf_account_link( # No exact match: no existing account if not existing_account: mark_user_aaf_only(email, aaf_user_id, auth0_client) - return AccountLinkResponse(link=False, aaf_only=True, primary_id=email) + resp = AccountLinkResponse(link=False, aaf_only=True, primary_id=aaf_user_id) + return return_signed_response( + state=state, + response=resp, + action_token=token, + settings=settings, + ) if existing_account.blocked: - raise HTTPException(status_code=403, detail="Existing account is blocked.") + resp = AccountLinkResponse(link=True, aaf_only=False, primary_id=existing_account.user_id, blocked=True) + return return_signed_response( + state=state, + response=resp, + action_token=token, + settings=settings, + ) - link_aaf_account( + aaf_identity = link_aaf_account( db_user_id=existing_account.user_id, aaf_user_id=aaf_user_id, auth0_client=auth0_client, session=session, ) - return AccountLinkResponse(link=True, aaf_only=False, primary_id=existing_account.user_id) + resp = AccountLinkResponse(link=True, aaf_only=False, primary_id=existing_account.user_id, aaf_identity=aaf_identity) + return return_signed_response( + state=state, + response=resp, + action_token=token, + settings=settings, + ) diff --git a/schemas/auth0.py b/schemas/auth0.py index 57070774..c44f7414 100644 --- a/schemas/auth0.py +++ b/schemas/auth0.py @@ -30,3 +30,5 @@ class Auth0ActionToken(BaseModel): email: EmailStr client_id: str purpose: str + sub: str | None = None + iss: str | None = None diff --git a/tests/auth/test_auth_validator.py b/tests/auth/test_auth_validator.py index 302f0bc5..4e8f2d45 100644 --- a/tests/auth/test_auth_validator.py +++ b/tests/auth/test_auth_validator.py @@ -2,7 +2,7 @@ import json import uuid from dataclasses import dataclass -from datetime import datetime, timedelta +from datetime import UTC, datetime, timedelta from typing import Optional from unittest.mock import AsyncMock, patch @@ -19,7 +19,9 @@ from fastapi import Depends, FastAPI, HTTPException from fastapi.security import HTTPAuthorizationCredentials from fastapi.testclient import TestClient +from freezegun import freeze_time from httpx import Request, Response +from jwt import InvalidSignatureError from jwt.algorithms import RSAAlgorithm from auth import auth0_security, get_auth0_token @@ -27,6 +29,7 @@ from auth.validator import ( KEY_CACHE, _fetch_rsa_keys, + create_action_token, get_rsa_key, verify_action_token, verify_jwt, @@ -39,6 +42,16 @@ TEST_MANAGEMENT_SECRET = "test-management-secret-key-with-32b" TEST_WRONG_MANAGEMENT_SECRET = "wrong-management-secret-key-32b!" TEST_CORRECT_MANAGEMENT_SECRET = "correct-management-secret-key-32" +FROZEN_TIME = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + + +@pytest.fixture +def frozen_time(): + """ + Freeze time so datetime.now() returns FROZEN_TIME. + """ + with freeze_time("2025-01-01 12:00:00") as time: + yield time def generate_public_private_key_pair(): @@ -586,6 +599,81 @@ def test_verify_action_token_missing_exp(mock_settings: Settings): assert excinfo.value.detail == "invalid session_token" +def test_create_action_token_success(mock_settings: Settings, frozen_time): + secret = TEST_MANAGEMENT_SECRET + mock_settings.auth0_management_secret = secret + payload = { + "sub": "auth0|123", + "iss": "https://issuer.example.com", + "state": "test-state", + "data": "test-payload", + } + token = create_action_token(payload, mock_settings) + decoded = jwt.decode(token, secret, algorithms=["HS256"]) + assert decoded["sub"] == "auth0|123" + assert decoded["iss"] == "https://issuer.example.com" + assert decoded["state"] == "test-state" + assert decoded["data"] == "test-payload" + assert decoded["exp"] == (FROZEN_TIME + timedelta(minutes=5)).timestamp() + assert decoded["iat"] == FROZEN_TIME.timestamp() + + +def test_create_action_token_expiry(mock_settings: Settings, frozen_time): + """ + Test we can change expiry time of an action token. + """ + secret = TEST_MANAGEMENT_SECRET + mock_settings.auth0_management_secret = secret + payload = { + "sub": "auth0|123", + "iss": "https://issuer.example.com", + "state": "test-state", + "data": "test-payload", + } + token = create_action_token(payload, mock_settings, expires_in_seconds=10 * 60) + decoded = jwt.decode(token, secret, algorithms=["HS256"]) + assert decoded["data"] == "test-payload" + assert decoded["exp"] == (FROZEN_TIME + timedelta(minutes=10)).timestamp() + + +@pytest.mark.parametrize("missing_field", ["sub", "iss", "state"]) +def test_create_action_token_missing_required_field( + mock_settings: Settings, + frozen_time, + missing_field: str, +): + secret = TEST_MANAGEMENT_SECRET + mock_settings.auth0_management_secret = secret + payload = { + "sub": "auth0|123", + "iss": "https://issuer.example.com", + "state": "test-state", + } + del payload[missing_field] + + with pytest.raises( + ValueError, + match=f"Missing required field {missing_field} in action token", + ): + create_action_token(payload, mock_settings) + + +def test_create_action_token_invalid(mock_settings: Settings, frozen_time): + secret = TEST_MANAGEMENT_SECRET + # Use invalid secret to sign + mock_settings.auth0_management_secret = "invalid-secret" + payload = { + "sub": "auth0|123", + "iss": "https://issuer.example.com", + "state": "test-state", + "data": "test-payload", + } + token = create_action_token(payload, mock_settings, expires_in_seconds=10 * 60) + # Use expected secret to decode + with pytest.raises(InvalidSignatureError, match="Signature verification failed"): + jwt.decode(token, secret, algorithms=["HS256"]) + + @pytest.mark.asyncio @respx.mock async def test_fetch_rsa_keys_only_refreshes_once_when_cache_is_expired(mock_settings: Settings): diff --git a/tests/test_aaf.py b/tests/test_aaf.py index 489ee26a..42478e48 100644 --- a/tests/test_aaf.py +++ b/tests/test_aaf.py @@ -1,6 +1,8 @@ from http import HTTPStatus from unittest.mock import MagicMock +from urllib.parse import parse_qs, urlparse +import jwt import pytest from fastapi import HTTPException @@ -30,13 +32,7 @@ def test_link_aaf_account_links_identity_updates_metadata_and_db(test_db_session ) auth0_client.get_user.assert_called_once_with(user_id=aaf_user_id) - auth0_client.link_identity.assert_called_once_with( - primary_user_id=db_user.id, - secondary_user_id=aaf_user_id, - secondary_provider="samlp", - secondary_connection_name="AAF", - ) - + auth0_client.link_identity.assert_not_called() auth0_client.update_user.assert_called_once() call_args, call_kwargs = auth0_client.update_user.call_args assert call_args[0] == db_user.id @@ -45,9 +41,7 @@ def test_link_aaf_account_links_identity_updates_metadata_and_db(test_db_session assert app_metadata.linking_completed is True assert app_metadata.linking_completed_at is not None - assert result.id == db_user.id - assert result.account_type == BiocommonsUserAccountType.AAF - assert result.other_user_id == aaf_user_id + assert result == aaf_user_data.identities[0] test_db_session.refresh(db_user) assert db_user.account_type == BiocommonsUserAccountType.AAF @@ -89,7 +83,7 @@ def test_link_aaf_account_raises_404_when_no_aaf_identity(test_db_session, persi assert db_user.other_user_id is None -def test_link_aaf_account_raises_502_when_auth0_linking_fails(test_db_session, persistent_factories): +def test_link_aaf_account_raises_502_when_auth0_update_fails(test_db_session, persistent_factories): db_user = BiocommonsUserFactory.create_sync(account_type=BiocommonsUserAccountType.AUTH0, other_user_id=None) test_db_session.commit() @@ -99,7 +93,7 @@ def test_link_aaf_account_raises_502_when_auth0_linking_fails(test_db_session, p ) auth0_client = MagicMock(spec=Auth0Client) auth0_client.get_user.return_value = aaf_user_data - auth0_client.link_identity.side_effect = ValueError("Identity already linked") + auth0_client.update_user.side_effect = ValueError("Auth0 update failed") with pytest.raises(HTTPException) as exc_info: link_aaf_account( @@ -110,7 +104,8 @@ def test_link_aaf_account_raises_502_when_auth0_linking_fails(test_db_session, p ) assert exc_info.value.status_code == HTTPStatus.BAD_GATEWAY - auth0_client.update_user.assert_not_called() + auth0_client.link_identity.assert_not_called() + auth0_client.update_user.assert_called_once() test_db_session.refresh(db_user) assert db_user.account_type == BiocommonsUserAccountType.AUTH0 @@ -124,7 +119,11 @@ def test_link_aaf_account_idempotent_when_already_linked(test_db_session, persis ) test_db_session.commit() + aaf_user_data = Auth0UserDataFactory.build( + identities=[Auth0Identity(connection="AAF", provider="samlp", user_id=aaf_user_id, isSocial=False)] + ) auth0_client = MagicMock(spec=Auth0Client) + auth0_client.get_user.return_value = aaf_user_data result = link_aaf_account( db_user_id=db_user.id, @@ -133,14 +132,31 @@ def test_link_aaf_account_idempotent_when_already_linked(test_db_session, persis session=test_db_session, ) - auth0_client.get_user.assert_not_called() + auth0_client.get_user.assert_called_once_with(user_id=aaf_user_id) auth0_client.link_identity.assert_not_called() auth0_client.update_user.assert_not_called() - assert result.id == db_user.id - - -def _action_token_payload(user_id: str, email: str, purpose: str = "aaf_link", client_id: str = "test-client") -> dict: - return {"user_id": user_id, "email": email, "client_id": client_id, "purpose": purpose} + assert result == aaf_user_data.identities[0] + + +def _action_token_payload( + user_id: str, + email: str, + purpose: str = "aaf_link", + client_id: str = "test-client", + sub: str | None = None, + iss: str | None = "mock-domain", +) -> dict: + payload = { + "user_id": user_id, + "email": email, + "client_id": client_id, + "purpose": purpose, + } + if sub is not None: + payload["sub"] = sub + if iss is not None: + payload["iss"] = iss + return payload def _assert_marked_aaf_only(update_user_mock, aaf_user_id: str, email: str): @@ -154,6 +170,45 @@ def _assert_marked_aaf_only(update_user_mock, aaf_user_id: str, email: str): assert app_metadata.linking_completed_at is not None +def _expected_auth0_continue_base_url(incoming_token: dict, settings) -> str: + issuer = incoming_token.get("iss") + if issuer is None: + return settings.auth0_custom_domain or f"https://{settings.auth0_domain}" + issuer = issuer.rstrip("/") + parsed_issuer = urlparse(issuer) + if parsed_issuer.scheme and parsed_issuer.netloc: + return f"{parsed_issuer.scheme}://{parsed_issuer.netloc}" + return f"https://{issuer}" + + +def _decode_check_link_redirect( + response, + state: str, + settings, + incoming_token: dict, +) -> dict: + assert response.status_code == HTTPStatus.TEMPORARY_REDIRECT + assert response.is_redirect + redirect = urlparse(response.headers["location"]) + expected_base_url = _expected_auth0_continue_base_url(incoming_token, settings) + expected_continue_url = urlparse(f"{expected_base_url}/continue") + assert redirect.scheme == expected_continue_url.scheme + assert redirect.netloc == expected_continue_url.netloc + assert redirect.path == expected_continue_url.path + query_params = parse_qs(redirect.query) + assert query_params["state"] == [state] + assert "session_token" in query_params + decoded_token = jwt.decode( + query_params["session_token"][0], + key=settings.auth0_management_secret, + algorithms=["HS256"], + ) + assert decoded_token["state"] == state + assert decoded_token["sub"] == incoming_token.get("sub", incoming_token["user_id"]) + assert decoded_token["iss"] == incoming_token.get("iss", settings.auth0_domain) + return decoded_token + + def test_mark_user_aaf_only(): aaf_user_id = random_auth0_id() email = "aaf-user@example.com" @@ -164,52 +219,113 @@ def test_mark_user_aaf_only(): _assert_marked_aaf_only(auth0_client.update_user, aaf_user_id, email) -def test_check_link_no_existing_account(test_client, mocker): +def test_check_link_no_existing_account(test_client, mocker, mock_settings): aaf_user_id = random_auth0_id() email = "new-aaf-user@example.com" - mocker.patch("dependencies.auth.verify_action_token", return_value=_action_token_payload(aaf_user_id, email)) + action_token_payload = _action_token_payload( + aaf_user_id, + email, + sub=aaf_user_id, + iss="https://login.example.com/", + ) + mocker.patch("dependencies.auth.verify_action_token", return_value=action_token_payload) mocker.patch("routers.aaf.Auth0Client.search_users_by_email", return_value=[]) link_identity = mocker.patch("routers.aaf.Auth0Client.link_identity") update_user = mocker.patch("routers.aaf.Auth0Client.update_user") - response = test_client.get("/aaf/check-link", params={"session_token": "valid_token"}) + state = "dummy" + response = test_client.get( + "/aaf/check-link", + params={"session_token": "valid_token", "state": state}, + follow_redirects=False, + ) + + decoded_token = _decode_check_link_redirect( + response, + state=state, + settings=mock_settings, + incoming_token=action_token_payload, + ) + assert decoded_token["link"] is False + assert decoded_token["aaf_only"] is True + assert decoded_token["primary_id"] == aaf_user_id - assert response.status_code == 200 - assert response.json() == {"link": False, "aaf_only": True, "primary_id": aaf_user_id} link_identity.assert_not_called() _assert_marked_aaf_only(update_user, aaf_user_id, email) -def test_check_link_no_exact_email_match(test_client, mocker): +def test_check_link_no_exact_email_match(test_client, mocker, mock_settings): aaf_user_id = random_auth0_id() email = "aaf-user@example.com" - mocker.patch("dependencies.auth.verify_action_token", return_value=_action_token_payload(aaf_user_id, email)) + action_token_payload = _action_token_payload( + aaf_user_id, + email, + sub="auth0|original-sub", + ) + mocker.patch("dependencies.auth.verify_action_token", return_value=action_token_payload) near_miss_account = Auth0UserDataFactory.build(email="different-user@example.com") mocker.patch("routers.aaf.Auth0Client.search_users_by_email", return_value=[near_miss_account]) link_identity = mocker.patch("routers.aaf.Auth0Client.link_identity") update_user = mocker.patch("routers.aaf.Auth0Client.update_user") - response = test_client.get("/aaf/check-link", params={"session_token": "valid_token"}) + state = "dummy" + response = test_client.get( + "/aaf/check-link", + params={"session_token": "valid_token", "state": state}, + follow_redirects=False, + ) - assert response.status_code == 200 - assert response.json() == {"link": False, "aaf_only": True, "primary_id": email} + decoded_token = _decode_check_link_redirect( + response, + state=state, + settings=mock_settings, + incoming_token=action_token_payload, + ) + assert decoded_token["link"] is False + assert decoded_token["aaf_only"] is True + assert decoded_token["primary_id"] == aaf_user_id link_identity.assert_not_called() _assert_marked_aaf_only(update_user, aaf_user_id, email) -def test_check_link_existing_account_blocked(test_client, mocker): +def test_check_link_existing_account_blocked(test_client, mocker, mock_settings): aaf_user_id = random_auth0_id() email = "blocked-user@example.com" - mocker.patch("dependencies.auth.verify_action_token", return_value=_action_token_payload(aaf_user_id, email)) + action_token_payload = _action_token_payload(aaf_user_id, email) + mocker.patch( + "dependencies.auth.verify_action_token", + return_value=action_token_payload, + ) existing_account = Auth0UserDataFactory.build(email=email, blocked=True) mocker.patch("routers.aaf.Auth0Client.search_users_by_email", return_value=[existing_account]) + get_user = mocker.patch("routers.aaf.Auth0Client.get_user") + link_identity = mocker.patch("routers.aaf.Auth0Client.link_identity") + update_user = mocker.patch("routers.aaf.Auth0Client.update_user") - response = test_client.get("/aaf/check-link", params={"session_token": "valid_token"}) + state = "dummy" + response = test_client.get( + "/aaf/check-link", + params={"session_token": "valid_token", "state": state}, + follow_redirects=False, + ) - assert response.status_code == 403 + decoded_token = _decode_check_link_redirect( + response, + state=state, + settings=mock_settings, + incoming_token=action_token_payload, + ) + assert decoded_token["link"] is True + assert decoded_token["aaf_only"] is False + assert decoded_token["blocked"] is True + assert decoded_token["primary_id"] == existing_account.user_id + + get_user.assert_not_called() + link_identity.assert_not_called() + update_user.assert_not_called() -def test_check_link_existing_account_links(test_client, test_db_session, persistent_factories, mocker): +def test_check_link_existing_account_links(test_client, test_db_session, persistent_factories, mocker, mock_settings): email = "existing-user@example.com" aaf_user_id = random_auth0_id() db_user = BiocommonsUserFactory.create_sync( @@ -217,7 +333,8 @@ def test_check_link_existing_account_links(test_client, test_db_session, persist ) test_db_session.commit() - mocker.patch("dependencies.auth.verify_action_token", return_value=_action_token_payload(aaf_user_id, email)) + action_token_payload = _action_token_payload(aaf_user_id, email, sub=aaf_user_id) + mocker.patch("dependencies.auth.verify_action_token", return_value=action_token_payload) existing_account = Auth0UserDataFactory.build(user_id=db_user.id, email=email, blocked=False) mocker.patch("routers.aaf.Auth0Client.search_users_by_email", return_value=[existing_account]) aaf_user_data = Auth0UserDataFactory.build( @@ -227,16 +344,29 @@ def test_check_link_existing_account_links(test_client, test_db_session, persist link_identity = mocker.patch("routers.aaf.Auth0Client.link_identity") update_user = mocker.patch("routers.aaf.Auth0Client.update_user") - response = test_client.get("/aaf/check-link", params={"session_token": "valid_token"}) + state = "dummy" + response = test_client.get( + "/aaf/check-link", + params={"session_token": "valid_token", "state": state}, + follow_redirects=False, + ) - assert response.status_code == 200 - assert response.json() == {"link": True, "aaf_only": False, "primary_id": db_user.id} - link_identity.assert_called_once_with( - primary_user_id=db_user.id, - secondary_user_id=aaf_user_id, - secondary_provider="samlp", - secondary_connection_name="AAF", + decoded_token = _decode_check_link_redirect( + response, + state=state, + settings=mock_settings, + incoming_token=action_token_payload, ) + assert decoded_token["link"] is True + assert decoded_token["aaf_only"] is False + assert decoded_token["primary_id"] == db_user.id + assert decoded_token["aaf_identity"] == { + "connection": "AAF", + "provider": "samlp", + "user_id": aaf_user_id, + "isSocial": False, + } + link_identity.assert_not_called() update_user.assert_called_once() test_db_session.refresh(db_user) @@ -244,7 +374,7 @@ def test_check_link_existing_account_links(test_client, test_db_session, persist assert db_user.other_user_id == aaf_user_id -def test_check_link_already_linked_is_idempotent(test_client, test_db_session, persistent_factories, mocker): +def test_check_link_already_linked_is_idempotent(test_client, test_db_session, persistent_factories, mocker, mock_settings): email = "already-linked@example.com" aaf_user_id = random_auth0_id() db_user = BiocommonsUserFactory.create_sync( @@ -252,17 +382,44 @@ def test_check_link_already_linked_is_idempotent(test_client, test_db_session, p ) test_db_session.commit() - mocker.patch("dependencies.auth.verify_action_token", return_value=_action_token_payload(aaf_user_id, email)) + action_token_payload = _action_token_payload( + aaf_user_id, + email, + sub=None, + iss=None, + ) + mocker.patch("dependencies.auth.verify_action_token", return_value=action_token_payload) existing_account = Auth0UserDataFactory.build(user_id=db_user.id, email=email, blocked=False) mocker.patch("routers.aaf.Auth0Client.search_users_by_email", return_value=[existing_account]) - get_user = mocker.patch("routers.aaf.Auth0Client.get_user") + aaf_user_data = Auth0UserDataFactory.build( + identities=[Auth0Identity(connection="AAF", provider="samlp", user_id=aaf_user_id, isSocial=False)] + ) + get_user = mocker.patch("routers.aaf.Auth0Client.get_user", return_value=aaf_user_data) link_identity = mocker.patch("routers.aaf.Auth0Client.link_identity") update_user = mocker.patch("routers.aaf.Auth0Client.update_user") - response = test_client.get("/aaf/check-link", params={"session_token": "valid_token"}) + state = "dummy" + response = test_client.get( + "/aaf/check-link", + params={"session_token": "valid_token", "state": state}, + follow_redirects=False, + ) - assert response.status_code == 200 - assert response.json() == {"link": True, "aaf_only": False, "primary_id": db_user.id} - get_user.assert_not_called() + decoded_token = _decode_check_link_redirect( + response, + state=state, + settings=mock_settings, + incoming_token=action_token_payload, + ) + assert decoded_token["link"] is True + assert decoded_token["aaf_only"] is False + assert decoded_token["primary_id"] == db_user.id + assert decoded_token["aaf_identity"] == { + "connection": "AAF", + "provider": "samlp", + "user_id": aaf_user_id, + "isSocial": False, + } + get_user.assert_called_once_with(user_id=aaf_user_id) link_identity.assert_not_called() update_user.assert_not_called() diff --git a/tests/test_dependencies.py b/tests/test_dependencies.py index 6ee38290..2f417f42 100644 --- a/tests/test_dependencies.py +++ b/tests/test_dependencies.py @@ -16,10 +16,13 @@ def create_signed_action_token(payload: Auth0ActionToken, secret: str) -> str: def test_require_action_token(mock_settings): - payload = Auth0ActionToken(user_id=random_auth0_id(), + user_id = random_auth0_id() + payload = Auth0ActionToken(user_id=user_id, email="test@example.com", client_id="abc123", - purpose="action_purpose") + purpose="action_purpose", + sub=user_id, + iss="mock-domain") signed_token = create_signed_action_token(payload, secret=mock_settings.auth0_management_secret) checker = require_action_token(purpose="action_purpose") checked_payload = checker(signed_token, mock_settings) @@ -27,10 +30,13 @@ def test_require_action_token(mock_settings): def test_require_action_token_wrong_purpose(mock_settings): - payload = Auth0ActionToken(user_id=random_auth0_id(), + user_id = random_auth0_id() + payload = Auth0ActionToken(user_id=user_id, email="test@example.com", client_id="abc123", - purpose="wrong_purpose") + purpose="wrong_purpose", + sub=user_id, + iss="mock-domain") signed_token = create_signed_action_token(payload, secret=mock_settings.auth0_management_secret) checker = require_action_token(purpose="action_purpose") with pytest.raises(HTTPException):