Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
33 commits
Select commit Hold shift + click to select a range
60918ef
feat: allow filtering users by connection when searching by email
marius-mather Aug 24, 2026
b8418f4
feat: updated schemas for AAF info
marius-mather Aug 24, 2026
d36f484
feat: link_identity method for Auth0Client
marius-mather Aug 24, 2026
917e625
feat: record other user ID for users
marius-mather Aug 24, 2026
e07815a
feat: method to link account for BiocommonsUser
marius-mather Aug 24, 2026
fc85223
fix: linking_completed_at should be datetime
marius-mather Aug 24, 2026
ebe6de1
feat: dependency to check Auth0 action tokens
marius-mather Aug 24, 2026
cb208b7
feat: router with initial /aaf/check-link endpoint
marius-mather Aug 24, 2026
3e416e1
chore: include aaf router in main
marius-mather Aug 24, 2026
5902ee9
refactor: use a base URL for Auth0 API calls
marius-mather Aug 24, 2026
809cb51
test: unit tests of new Auth0Client methods
marius-mather Aug 24, 2026
6380271
test: unit tests of AAF link function
marius-mather Aug 24, 2026
7c71e2e
test: tests of require_action_token
marius-mather Aug 25, 2026
04551d5
test: unit tests of account linking endpoint
marius-mather Aug 25, 2026
a7f67c6
fix: get connection correctly when linking accounts
marius-mather Aug 25, 2026
40d8a00
fix: skip linking if already done
marius-mather Aug 25, 2026
63a59a1
fix: update app_metadata to mark user AAF only if no existing account
marius-mather Aug 25, 2026
10fd2e3
test: unit tests of app_metadata for AAF-only accounts
marius-mather Aug 25, 2026
531039c
fix: handle errors when Auth0 already linked
marius-mather Aug 25, 2026
7362c98
feat: create_action_token function to sign tokens for Auth0 actions
marius-mather Sep 1, 2026
539cbc6
fix: need to return a redirect response to Auth0 in check-link
marius-mather Sep 1, 2026
99057e6
test: update tests of check-link
marius-mather Sep 1, 2026
badfa91
Merge remote-tracking branch 'origin/aaf-dev' into feat/aaf-account-l…
marius-mather Sep 1, 2026
83bfd6f
fix: fall back to Auth0 domain if custom domain not defined
marius-mather Sep 1, 2026
c6a5aa7
test: make sure we test redirect location
marius-mather Sep 1, 2026
f9486f6
fix: set required fields on action token payload
marius-mather Sep 3, 2026
8fb9eec
fix: make sure we're including all the required fields for a token re…
marius-mather Sep 3, 2026
5d2b2ad
test: update tests
marius-mather Sep 3, 2026
a9b9468
fix: action token should also contain iat field
marius-mather Sep 3, 2026
e1cf483
fix: don't return 403 response for blocked accounts, return a token t…
marius-mather Sep 3, 2026
0b38c1e
fix: can't link on this side, return identity to Auth0 for linking
marius-mather Sep 4, 2026
59cc94c
test: update tests
marius-mather Sep 4, 2026
cfa75c6
fix: remove response_model - response is a redirect
marius-mather Sep 4, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions auth/validator.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import json
import logging
import weakref
from datetime import UTC, datetime, timedelta

import httpx
import jwt
Expand Down Expand Up @@ -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
124 changes: 106 additions & 18 deletions routers/aaf.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,18 +2,23 @@
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
from db.setup import get_db_session
from dependencies.auth import require_action_token
from schemas.auth0 import Auth0ActionToken
from schemas.biocommons import (
Auth0Identity,
Auth0UserData,
BiocommonsAppMetadataUpdate,
BiocommonsUserAccountType,
Expand All @@ -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(
Expand All @@ -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):
Expand All @@ -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)],
Expand All @@ -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:
Expand All @@ -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,
)
2 changes: 2 additions & 0 deletions schemas/auth0.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,3 +30,5 @@ class Auth0ActionToken(BaseModel):
email: EmailStr
client_id: str
purpose: str
sub: str | None = None
iss: str | None = None
90 changes: 89 additions & 1 deletion tests/auth/test_auth_validator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -19,14 +19,17 @@
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
from auth.user_permissions import user_is_general_admin
from auth.validator import (
KEY_CACHE,
_fetch_rsa_keys,
create_action_token,
get_rsa_key,
verify_action_token,
verify_jwt,
Expand All @@ -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():
Expand Down Expand Up @@ -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):
Expand Down
Loading