diff --git a/.env.example b/.env.example index e0e7a05..f274cca 100644 --- a/.env.example +++ b/.env.example @@ -13,6 +13,24 @@ KEYCLOAK_REALM=cell-explorer KEYCLOAK_CLIENT_ID=cell-explorer-app KEYCLOAK_CLIENT_SECRET= +# OIDC provider (generic). AUTH_PROVIDER picks preset defaults: keycloak (default) | entra | oidc. +# For Keycloak the KEYCLOAK_* vars above are sufficient — no OIDC_* needed. +# AUTH_PROVIDER=keycloak +# +# Microsoft Entra: +# AUTH_PROVIDER=entra +# OIDC_ISSUER=https://login.microsoftonline.com//v2.0 +# OIDC_CLIENT_ID= +# OIDC_CLIENT_SECRET= +# Define App Roles in the Entra app registration; they arrive in the `roles` claim. +# offline_access is added automatically for refresh tokens. +# If your Entra access-token `aud` differs from the client id (often api://), +# set OIDC_AUDIENCE to match, or token validation will 401. +# +# Generic OIDC: AUTH_PROVIDER=oidc plus OIDC_ISSUER, OIDC_CLIENT_ID/SECRET, and +# OIDC_ROLES_CLAIMS (comma-separated dotted claim-paths, e.g. "roles"). +# OIDC_SCOPES / OIDC_AUDIENCE / OIDC_ROLES_CLAIMS override the preset defaults. + # Session cookie lifetimes (seconds). Tune REFRESH_COOKIE_MAX_AGE to match the # realm's ssoSessionMaxLifespan (currently 24h = 86400). Setting it higher than # Keycloak allows just causes refresh to fail before the cookie expires. diff --git a/packages/api/src/cell_explorer_api/auth/admin.py b/packages/api/src/cell_explorer_api/auth/admin.py index 41decd4..cf775ad 100644 --- a/packages/api/src/cell_explorer_api/auth/admin.py +++ b/packages/api/src/cell_explorer_api/auth/admin.py @@ -27,15 +27,15 @@ async def require_admin( if credentials and credentials.credentials == settings.admin_api_key: return - # Try Keycloak JWT with admin role + # Try OIDC JWT with admin role if settings.auth_enabled: access_token = request.cookies.get("cce_access") if access_token: try: - from cell_explorer_api.auth.keycloak import KeycloakClient + from cell_explorer_api.auth.oidc import OidcClient - keycloak: KeycloakClient = request.app.state.keycloak - user = keycloak.decode_token(access_token) + oidc: OidcClient = request.app.state.oidc + user = oidc.decode_token(access_token) if "admin" in user.roles: return except Exception: diff --git a/packages/api/src/cell_explorer_api/auth/dependencies.py b/packages/api/src/cell_explorer_api/auth/dependencies.py index 3f0520d..e13e8cd 100644 --- a/packages/api/src/cell_explorer_api/auth/dependencies.py +++ b/packages/api/src/cell_explorer_api/auth/dependencies.py @@ -4,7 +4,7 @@ from fastapi import HTTPException, Request -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.auth.models import User logger = logging.getLogger(__name__) @@ -15,19 +15,19 @@ async def require_auth(request: Request) -> User: access_token = request.cookies.get("cce_access") refresh_token = request.cookies.get("cce_refresh") - # No credentials at all — short-circuit before hitting Keycloak config. + # No credentials at all — short-circuit before hitting OIDC config. if not access_token and not refresh_token: raise HTTPException(status_code=401, detail="Not authenticated") if not request.app.state.settings.auth_enabled: raise HTTPException(status_code=501, detail="Authentication is not configured") - keycloak: KeycloakClient = request.app.state.keycloak + oidc: OidcClient = request.app.state.oidc # Happy path: try to decode the access token if we have one. if access_token: try: - return keycloak.decode_token(access_token) + return oidc.decode_token(access_token) except Exception as e: logger.warning("Access token decode failed: %s", e) @@ -38,8 +38,8 @@ async def require_auth(request: Request) -> User: try: logger.info("Attempting token refresh") - tokens = await keycloak.refresh_token(refresh_token) - user = keycloak.decode_token(tokens["access_token"]) + tokens = await oidc.refresh_token(refresh_token) + user = oidc.decode_token(tokens["access_token"]) request.state.new_access_token = tokens["access_token"] request.state.new_refresh_token = tokens.get("refresh_token", refresh_token) return user diff --git a/packages/api/src/cell_explorer_api/auth/keycloak.py b/packages/api/src/cell_explorer_api/auth/keycloak.py deleted file mode 100644 index 86508e2..0000000 --- a/packages/api/src/cell_explorer_api/auth/keycloak.py +++ /dev/null @@ -1,131 +0,0 @@ -"""Keycloak OIDC client — token validation, auth URLs, token exchange.""" - -import logging -from urllib.parse import urlencode - -import httpx -import jwt -from cryptography.hazmat.primitives.serialization import ( - Encoding, - PublicFormat, - load_pem_public_key, -) - -from cell_explorer_api.auth.models import User -from cell_explorer_api.config import Settings - -logger = logging.getLogger(__name__) - - -class KeycloakClient: - """Handles Keycloak OIDC operations.""" - - def __init__(self, settings: Settings) -> None: - self._settings = settings - self._base = f"{settings.keycloak_url}/realms/{settings.keycloak_realm}" - self._oidc = f"{self._base}/protocol/openid-connect" - self._jwks: dict[str, bytes] = {} # kid → PEM public key bytes - - def authorization_url(self, redirect_uri: str, state: str) -> str: - """Build the Keycloak authorization endpoint URL.""" - params = { - "client_id": self._settings.keycloak_client_id, - "response_type": "code", - "redirect_uri": redirect_uri, - "state": state, - "scope": "openid profile email", - } - if self._settings.keycloak_idp_hint: - params["kc_idp_hint"] = self._settings.keycloak_idp_hint - return f"{self._oidc}/auth?{urlencode(params)}" - - async def exchange_code(self, code: str, redirect_uri: str) -> dict: - """Exchange an authorization code for tokens.""" - async with httpx.AsyncClient() as client: - response = await client.post( - f"{self._oidc}/token", - data={ - "grant_type": "authorization_code", - "code": code, - "redirect_uri": redirect_uri, - "client_id": self._settings.keycloak_client_id, - "client_secret": self._settings.keycloak_client_secret, - }, - ) - response.raise_for_status() - return response.json() - - async def refresh_token(self, refresh_token: str) -> dict: - """Use a refresh token to get a new access token.""" - async with httpx.AsyncClient() as client: - response = await client.post( - f"{self._oidc}/token", - data={ - "grant_type": "refresh_token", - "refresh_token": refresh_token, - "client_id": self._settings.keycloak_client_id, - "client_secret": self._settings.keycloak_client_secret, - }, - ) - response.raise_for_status() - return response.json() - - async def fetch_jwks(self) -> None: - """Fetch and cache Keycloak's JSON Web Key Set.""" - async with httpx.AsyncClient() as client: - response = await client.get(f"{self._oidc}/certs") - response.raise_for_status() - jwks = response.json() - - self._jwks = {} - for key_data in jwks.get("keys", []): - kid = key_data.get("kid") - if kid: - public_key = jwt.algorithms.RSAAlgorithm.from_jwk(key_data) - self._jwks[kid] = public_key.public_bytes( - encoding=Encoding.PEM, - format=PublicFormat.SubjectPublicKeyInfo, - ) - - def decode_token(self, token: str) -> User: - """Decode and validate a JWT access token. Raises on invalid/expired.""" - header = jwt.get_unverified_header(token) - kid = header.get("kid") - if kid not in self._jwks: - raise jwt.InvalidTokenError(f"Unknown kid: {kid}") - - public_key = load_pem_public_key(self._jwks[kid]) - # leeway=30s tolerates small clock skew between Keycloak and this - # container. Without it, a token Keycloak just issued can be rejected - # as "not yet valid (iat)" if Keycloak's clock is even a second ahead - # — which manifests as a chronic refresh failure once the access - # cookie expires (see issue #131). - claims = jwt.decode( - token, - public_key, - algorithms=["RS256"], - audience=self._settings.keycloak_client_id, - issuer=self._settings.oidc_issuer_url, - leeway=30, - ) - - realm_roles = claims.get("realm_access", {}).get("roles", []) - client_roles = ( - claims.get("resource_access", {}) - .get(self._settings.keycloak_client_id, {}) - .get("roles", []) - ) - return User( - sub=claims["sub"], - name=claims.get("name"), - email=claims.get("email"), - roles=sorted(set(realm_roles + client_roles)), - ) - - def logout_url(self, redirect_uri: str) -> str: - """Build the Keycloak logout endpoint URL.""" - params = { - "client_id": self._settings.keycloak_client_id, - "post_logout_redirect_uri": redirect_uri, - } - return f"{self._oidc}/logout?{urlencode(params)}" diff --git a/packages/api/src/cell_explorer_api/auth/oidc.py b/packages/api/src/cell_explorer_api/auth/oidc.py new file mode 100644 index 0000000..438520e --- /dev/null +++ b/packages/api/src/cell_explorer_api/auth/oidc.py @@ -0,0 +1,152 @@ +"""Generic OIDC client — discovery, token validation, auth URLs, token exchange. + +Works with any OIDC provider (Keycloak, Microsoft Entra, ...). Endpoints are +resolved from the provider's .well-known/openid-configuration; roles are read +from configurable dotted claim-paths. Provider presets live in Settings. +""" + +import logging +from urllib.parse import urlencode + +import httpx +import jwt +from cryptography.hazmat.primitives.serialization import ( + Encoding, + PublicFormat, + load_pem_public_key, +) + +from cell_explorer_api.auth.models import User +from cell_explorer_api.config import Settings + +logger = logging.getLogger(__name__) + + +def extract_roles(claims: dict, paths: list[str]) -> list[str]: + """Merge role lists found at each dotted JSON-path into a sorted, deduped list.""" + roles: set[str] = set() + for path in paths: + node = claims + for part in path.split("."): + if not isinstance(node, dict) or part not in node: + node = None + break + node = node[part] + if isinstance(node, list): + roles.update(str(r) for r in node) + return sorted(roles) + + +class OidcClient: + """Provider-agnostic OIDC operations, driven by resolved Settings + discovery.""" + + def __init__(self, settings: Settings) -> None: + self._settings = settings + self._jwks: dict[str, bytes] = {} # kid → PEM public key bytes + self._authorization_endpoint: str | None = None + self._token_endpoint: str | None = None + self._jwks_uri: str | None = None + self._issuer: str | None = None + self._end_session_endpoint: str | None = None + + def _apply_discovery(self, doc: dict) -> None: + self._authorization_endpoint = doc.get("authorization_endpoint") + self._token_endpoint = doc.get("token_endpoint") + self._jwks_uri = doc.get("jwks_uri") + self._issuer = doc.get("issuer") + self._end_session_endpoint = doc.get("end_session_endpoint") + + async def discover(self) -> None: + """Fetch .well-known/openid-configuration and cache the endpoints.""" + url = self._settings.discovery_url + async with httpx.AsyncClient() as client: + response = await client.get(url) + response.raise_for_status() + self._apply_discovery(response.json()) + + def authorization_url(self, redirect_uri: str, state: str) -> str: + params = { + "client_id": self._settings.resolved_client_id, + "response_type": "code", + "redirect_uri": redirect_uri, + "state": state, + "scope": self._settings.resolved_scopes, + } + if self._settings.resolved_idp_hint: + params["kc_idp_hint"] = self._settings.resolved_idp_hint + return f"{self._authorization_endpoint}?{urlencode(params)}" + + async def exchange_code(self, code: str, redirect_uri: str) -> dict: + async with httpx.AsyncClient() as client: + response = await client.post( + self._token_endpoint, + data={ + "grant_type": "authorization_code", + "code": code, + "redirect_uri": redirect_uri, + "client_id": self._settings.resolved_client_id, + "client_secret": self._settings.resolved_client_secret, + }, + ) + response.raise_for_status() + return response.json() + + async def refresh_token(self, refresh_token: str) -> dict: + async with httpx.AsyncClient() as client: + response = await client.post( + self._token_endpoint, + data={ + "grant_type": "refresh_token", + "refresh_token": refresh_token, + "client_id": self._settings.resolved_client_id, + "client_secret": self._settings.resolved_client_secret, + }, + ) + response.raise_for_status() + return response.json() + + async def fetch_jwks(self) -> None: + async with httpx.AsyncClient() as client: + response = await client.get(self._jwks_uri) + response.raise_for_status() + jwks = response.json() + self._jwks = {} + for key_data in jwks.get("keys", []): + kid = key_data.get("kid") + if kid: + public_key = jwt.algorithms.RSAAlgorithm.from_jwk(key_data) + self._jwks[kid] = public_key.public_bytes( + encoding=Encoding.PEM, format=PublicFormat.SubjectPublicKeyInfo, + ) + + def decode_token(self, token: str) -> User: + """Decode and validate a JWT access token. Raises on invalid/expired.""" + header = jwt.get_unverified_header(token) + kid = header.get("kid") + if kid not in self._jwks: + raise jwt.InvalidTokenError(f"Unknown kid: {kid}") + + public_key = load_pem_public_key(self._jwks[kid]) + # leeway=30s tolerates small clock skew between the IdP and this + # container (see issue #131). + claims = jwt.decode( + token, + public_key, + algorithms=["RS256"], + audience=self._settings.resolved_audience, + issuer=self._issuer or self._settings.resolved_issuer, + leeway=30, + ) + return User( + sub=claims["sub"], + name=claims.get("name"), + email=claims.get("email"), + roles=extract_roles(claims, self._settings.resolved_roles_claims), + ) + + def logout_url(self, redirect_uri: str) -> str: + params = { + "client_id": self._settings.resolved_client_id, + "post_logout_redirect_uri": redirect_uri, + } + return f"{self._end_session_endpoint}?{urlencode(params)}" diff --git a/packages/api/src/cell_explorer_api/auth/optional.py b/packages/api/src/cell_explorer_api/auth/optional.py index ca0578c..88cfef2 100644 --- a/packages/api/src/cell_explorer_api/auth/optional.py +++ b/packages/api/src/cell_explorer_api/auth/optional.py @@ -21,10 +21,10 @@ async def optional_auth(request: Request) -> User | None: return None try: - from cell_explorer_api.auth.keycloak import KeycloakClient + from cell_explorer_api.auth.oidc import OidcClient - keycloak: KeycloakClient = request.app.state.keycloak - return keycloak.decode_token(access_token) + oidc: OidcClient = request.app.state.oidc + return oidc.decode_token(access_token) except Exception: logger.debug("Optional auth: token decode failed, treating as anonymous") return None diff --git a/packages/api/src/cell_explorer_api/cli/main.py b/packages/api/src/cell_explorer_api/cli/main.py index a44289c..d98d94f 100644 --- a/packages/api/src/cell_explorer_api/cli/main.py +++ b/packages/api/src/cell_explorer_api/cli/main.py @@ -31,7 +31,7 @@ def _decode_username(access_token: str) -> str: """Best-effort extract username from the JWT access token (no signature check). Authoritative validation happens inside the agent loop when the user is - constructed from the Settings/KeycloakClient. Here we only want a display + constructed from the Settings/OidcClient. Here we only want a display value for the login banner. """ import jwt as _jwt @@ -115,13 +115,14 @@ class _User: async def _load_user_from_auth(settings: Settings) -> _User: """Load auth.json, refresh if needed, decode the access token, return _User.""" - from cell_explorer_api.auth.keycloak import KeycloakClient + from cell_explorer_api.auth.oidc import OidcClient cfg = load_auth_config() - keycloak = KeycloakClient(settings) - await keycloak.fetch_jwks() - cfg = await ensure_fresh_access_token(cfg, keycloak) - user_obj = keycloak.decode_token(cfg.access_token) + oidc = OidcClient(settings) + await oidc.discover() + await oidc.fetch_jwks() + cfg = await ensure_fresh_access_token(cfg, oidc) + user_obj = oidc.decode_token(cfg.access_token) # The auth.models.User has sub/name/email/roles but no `username` field. # Compute a display value from whichever fields are populated. display = user_obj.email or user_obj.name or user_obj.sub diff --git a/packages/api/src/cell_explorer_api/config.py b/packages/api/src/cell_explorer_api/config.py index 78da367..e1c2878 100644 --- a/packages/api/src/cell_explorer_api/config.py +++ b/packages/api/src/cell_explorer_api/config.py @@ -47,6 +47,17 @@ class Settings(BaseSettings): keycloak_client_id: str | None = None keycloak_client_secret: str | None = None keycloak_idp_hint: str | None = None + + # Generic OIDC (provider-agnostic). auth_provider selects preset defaults; + # KEYCLOAK_* remain the zero-config path for auth_provider="keycloak". + auth_provider: str = "keycloak" # keycloak | entra | oidc + oidc_issuer: str | None = None + oidc_client_id: str | None = None + oidc_client_secret: str | None = None + oidc_scopes: str = "openid profile email" + oidc_roles_claims: str | None = None + oidc_audience: str | None = None + cors_origins: str = "" # Session cookie lifetimes (seconds). Defaults match a typical Keycloak @@ -90,21 +101,57 @@ def effective_database_url(self) -> str: return f"sqlite+aiosqlite:///{self.app_data_dir / 'cell_explorer.db'}" @property - def auth_enabled(self) -> bool: - """Auth is enabled when all required Keycloak fields are set.""" - return all([ - self.keycloak_url, - self.keycloak_realm, - self.keycloak_client_id, - self.keycloak_client_secret, - ]) + def resolved_client_id(self) -> str | None: + return self.oidc_client_id or self.keycloak_client_id + + @property + def resolved_client_secret(self) -> str | None: + return self.oidc_client_secret or self.keycloak_client_secret + + @property + def resolved_issuer(self) -> str | None: + """OIDC issuer URL. Keycloak derives it from url+realm; others use OIDC_ISSUER.""" + if self.oidc_issuer: + return self.oidc_issuer + if self.auth_provider == "keycloak" and self.keycloak_url and self.keycloak_realm: + return f"{self.keycloak_url}/realms/{self.keycloak_realm}" + return None + + @property + def discovery_url(self) -> str | None: + issuer = self.resolved_issuer + return f"{issuer}/.well-known/openid-configuration" if issuer else None @property - def oidc_issuer_url(self) -> str | None: - """Keycloak OIDC issuer URL, constructed from base URL and realm.""" - if not self.keycloak_url or not self.keycloak_realm: - return None - return f"{self.keycloak_url}/realms/{self.keycloak_realm}" + def resolved_audience(self) -> str | None: + return self.oidc_audience or self.resolved_client_id + + @property + def resolved_scopes(self) -> str: + scopes = self.oidc_scopes + if self.auth_provider == "entra" and "offline_access" not in scopes.split(): + scopes = f"{scopes} offline_access" + return scopes + + @property + def resolved_idp_hint(self) -> str | None: + return self.keycloak_idp_hint if self.auth_provider == "keycloak" else None + + @property + def resolved_roles_claims(self) -> list[str]: + """Dotted claim-paths whose role lists are merged into User.roles.""" + if self.oidc_roles_claims: + return [p.strip() for p in self.oidc_roles_claims.split(",") if p.strip()] + if self.auth_provider == "keycloak": + return ["realm_access.roles", f"resource_access.{self.resolved_client_id}.roles"] + if self.auth_provider == "entra": + return ["roles"] + return [] + + @property + def auth_enabled(self) -> bool: + """Auth is enabled when issuer + client id + secret resolve.""" + return all([self.resolved_issuer, self.resolved_client_id, self.resolved_client_secret]) @property def cors_origin_list(self) -> list[str]: diff --git a/packages/api/src/cell_explorer_api/main.py b/packages/api/src/cell_explorer_api/main.py index 382328a..08a1a24 100644 --- a/packages/api/src/cell_explorer_api/main.py +++ b/packages/api/src/cell_explorer_api/main.py @@ -75,18 +75,19 @@ def create_app(settings: Settings | None = None) -> FastAPI: app.state.db_engine = create_engine(settings.effective_database_url) - # Auth — routes always registered (return 501 when disabled); Keycloak client only when configured + # Auth — routes always registered (return 501 when disabled); OIDC client only when configured if settings.auth_enabled: - from cell_explorer_api.auth.keycloak import KeycloakClient + from cell_explorer_api.auth.oidc import OidcClient - keycloak = KeycloakClient(settings) - app.state.keycloak = keycloak + oidc = OidcClient(settings) + app.state.oidc = oidc # Unified lifespan: always dispose db engine; fetch JWKS only when auth enabled @asynccontextmanager async def lifespan(app): if settings.auth_enabled: - await app.state.keycloak.fetch_jwks() + await app.state.oidc.discover() + await app.state.oidc.fetch_jwks() yield # Flush any queued Langfuse traces before tearing down. The flush # is a best-effort call that swallows errors; never raises. diff --git a/packages/api/src/cell_explorer_api/routes/auth.py b/packages/api/src/cell_explorer_api/routes/auth.py index bc84d51..717d572 100644 --- a/packages/api/src/cell_explorer_api/routes/auth.py +++ b/packages/api/src/cell_explorer_api/routes/auth.py @@ -9,7 +9,7 @@ from fastapi.responses import RedirectResponse from cell_explorer_api.auth.dependencies import require_auth -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.auth.models import User router = APIRouter(prefix="/auth", tags=["auth"]) @@ -130,10 +130,10 @@ def _clear_token_cookies(response: Response) -> None: async def login(request: Request): """Redirect to Keycloak login page.""" _require_auth_enabled(request) - keycloak: KeycloakClient = request.app.state.keycloak + oidc: OidcClient = request.app.state.oidc state = secrets.token_urlsafe(32) redirect_uri = _callback_uri(request) - auth_url = keycloak.authorization_url(redirect_uri=redirect_uri, state=state) + auth_url = oidc.authorization_url(redirect_uri=redirect_uri, state=state) response = RedirectResponse(url=auth_url, status_code=307) defaults = _cookie_defaults(request) response.set_cookie("cce_state", state, max_age=600, httponly=True, secure=defaults["secure"], samesite="lax", path="/api/auth") @@ -159,9 +159,9 @@ async def cli_login(request: Request, redirect_uri: str): ) signed_state = _sign_cli_state(secret, redirect_uri) - keycloak: KeycloakClient = request.app.state.keycloak + oidc: OidcClient = request.app.state.oidc api_callback_uri = _callback_uri(request) - auth_url = keycloak.authorization_url(redirect_uri=api_callback_uri, state=signed_state) + auth_url = oidc.authorization_url(redirect_uri=api_callback_uri, state=signed_state) return RedirectResponse(url=auth_url, status_code=307) @@ -169,7 +169,7 @@ async def cli_login(request: Request, redirect_uri: str): async def callback(request: Request, code: str, state: str): """Handle Keycloak callback — exchange code for tokens.""" _require_auth_enabled(request) - keycloak: KeycloakClient = request.app.state.keycloak + oidc: OidcClient = request.app.state.oidc # Try CLI flow first (state is a signed JWT). settings = request.app.state.settings @@ -177,7 +177,7 @@ async def callback(request: Request, code: str, state: str): cli_state = _decode_cli_state(state, settings.cli_state_secret) if cli_state is not None: redirect_uri = cli_state["redirect_uri"] - tokens = await keycloak.exchange_code(code, _callback_uri(request)) + tokens = await oidc.exchange_code(code, _callback_uri(request)) import json as _json tokens_payload = { "access_token": tokens["access_token"], @@ -195,7 +195,7 @@ async def callback(request: Request, code: str, state: str): if not expected_state or state != expected_state: return Response(status_code=400, content="Invalid state parameter") redirect_uri = _callback_uri(request) - tokens = await keycloak.exchange_code(code, redirect_uri) + tokens = await oidc.exchange_code(code, redirect_uri) response = RedirectResponse(url="/", status_code=302) _set_token_cookies(request, response, tokens["access_token"], tokens["refresh_token"]) response.delete_cookie("cce_state", path="/api/auth") @@ -231,9 +231,9 @@ async def refresh(request: Request): refresh_token = request.cookies.get("cce_refresh") if not refresh_token: raise HTTPException(status_code=401, detail="No refresh cookie") - keycloak: KeycloakClient = request.app.state.keycloak + oidc: OidcClient = request.app.state.oidc try: - tokens = await keycloak.refresh_token(refresh_token) + tokens = await oidc.refresh_token(refresh_token) except Exception: raise HTTPException(status_code=401, detail="Refresh failed") response = Response(status_code=204) @@ -250,13 +250,13 @@ async def refresh(request: Request): async def token_exchange(request: Request): """Exchange an external access token for session cookies.""" _require_auth_enabled(request) - keycloak: KeycloakClient = request.app.state.keycloak + oidc: OidcClient = request.app.state.oidc body = await request.json() access_token = body.get("accessToken") if not access_token: return Response(status_code=400, content="accessToken required") try: - user = keycloak.decode_token(access_token) + user = oidc.decode_token(access_token) except Exception: return Response(status_code=401, content="Invalid token") response = Response(status_code=200) diff --git a/packages/api/tests/routes/conftest.py b/packages/api/tests/routes/conftest.py index 567712a..a9bf0b1 100644 --- a/packages/api/tests/routes/conftest.py +++ b/packages/api/tests/routes/conftest.py @@ -16,7 +16,7 @@ from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.config import Settings from cell_explorer_api.db.models import Dataset, Datasource, DatasourceType from cell_explorer_api.main import create_app @@ -81,11 +81,18 @@ def _set_auth_cookie(client: TestClient, app, *, sub="user-1", roles=None): settings.keycloak_realm = "test-realm" settings.keycloak_client_id = "test-client" settings.keycloak_client_secret = "test-secret" - keycloak = KeycloakClient(settings) - keycloak._jwks = { + oidc = OidcClient(settings) + oidc._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) + oidc._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } - app.state.keycloak = keycloak + app.state.oidc = oidc token = jwt.encode( { diff --git a/packages/api/tests/routes/test_auth_cli_login.py b/packages/api/tests/routes/test_auth_cli_login.py index 67f5049..10c6c2d 100644 --- a/packages/api/tests/routes/test_auth_cli_login.py +++ b/packages/api/tests/routes/test_auth_cli_login.py @@ -21,11 +21,11 @@ def _settings_with_keycloak() -> Settings: def _make_client() -> tuple[TestClient, MagicMock, Settings]: settings = _settings_with_keycloak() app = create_app(settings=settings) - # Inject a fake KeycloakClient whose authorization_url we can assert on. + # Inject a fake OidcClient whose authorization_url we can assert on. fake_kc = MagicMock() fake_kc.authorization_url = MagicMock(return_value="http://kc.test/auth?...signed_state") fake_kc.fetch_jwks = AsyncMock() - app.state.keycloak = fake_kc + app.state.oidc = fake_kc return TestClient(app), fake_kc, settings @@ -105,7 +105,7 @@ def test_cli_login_requires_cli_state_secret(): app = create_app(settings=settings) fake_kc = MagicMock() fake_kc.fetch_jwks = AsyncMock() - app.state.keycloak = fake_kc + app.state.oidc = fake_kc client = TestClient(app) resp = client.get( "/api/auth/cli-login", diff --git a/packages/api/tests/test_admin_auth.py b/packages/api/tests/test_admin_auth.py index 35a7be4..678e2d4 100644 --- a/packages/api/tests/test_admin_auth.py +++ b/packages/api/tests/test_admin_auth.py @@ -10,7 +10,7 @@ from fastapi.testclient import TestClient from cell_explorer_api.auth.admin import require_admin -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.config import Settings @@ -83,8 +83,15 @@ def test_admin_accepts_keycloak_admin_role(): keycloak_client_secret="test-secret", ) app = _make_test_app(settings) - keycloak: KeycloakClient = app.state.keycloak - keycloak._jwks = { + oidc: OidcClient = app.state.oidc + oidc._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) + oidc._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } diff --git a/packages/api/tests/test_auth_dependency.py b/packages/api/tests/test_auth_dependency.py index 60c9903..51e136e 100644 --- a/packages/api/tests/test_auth_dependency.py +++ b/packages/api/tests/test_auth_dependency.py @@ -11,7 +11,7 @@ from fastapi.testclient import TestClient from cell_explorer_api.auth.dependencies import require_auth -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.auth.models import User from cell_explorer_api.config import Settings @@ -30,10 +30,10 @@ def _make_settings() -> Settings: ) -def _make_app_with_protected_route(keycloak: KeycloakClient) -> FastAPI: +def _make_app_with_protected_route(oidc: OidcClient) -> FastAPI: app = FastAPI() app.state.settings = _make_settings() - app.state.keycloak = keycloak + app.state.oidc = oidc @app.get("/protected") async def protected(user: User = Depends(require_auth)): @@ -51,7 +51,14 @@ def rsa_keys(): def keycloak(rsa_keys): _, public_key = rsa_keys settings = _make_settings() - client = KeycloakClient(settings) + client = OidcClient(settings) + client._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) client._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } diff --git a/packages/api/tests/test_auth_routes.py b/packages/api/tests/test_auth_routes.py index 4cbc354..f3f5567 100644 --- a/packages/api/tests/test_auth_routes.py +++ b/packages/api/tests/test_auth_routes.py @@ -8,7 +8,7 @@ from cryptography.hazmat.primitives.serialization import Encoding, PublicFormat from fastapi.testclient import TestClient -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.config import Settings from cell_explorer_api.main import create_app @@ -51,8 +51,15 @@ def auth_client(rsa_keys): _, public_key = rsa_keys settings = _make_settings() app = create_app(settings) - keycloak: KeycloakClient = app.state.keycloak - keycloak._jwks = { + oidc: OidcClient = app.state.oidc + oidc._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) + oidc._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } return TestClient(app) @@ -129,10 +136,10 @@ def test_me_sets_refreshed_cookies(auth_client, rsa_keys): # Create a fresh token that the refresh flow will return fresh_token = _make_token(private_key) - # Mock the refresh_token method on the keycloak client - keycloak: KeycloakClient = auth_client.app.state.keycloak + # Mock the refresh_token method on the oidc client + oidc: OidcClient = auth_client.app.state.oidc import unittest.mock as mock - keycloak.refresh_token = mock.AsyncMock(return_value={ + oidc.refresh_token = mock.AsyncMock(return_value={ "access_token": fresh_token, "refresh_token": "new-refresh-token-value", }) @@ -167,9 +174,9 @@ def test_refresh_endpoint_rotates_cookies(auth_client, rsa_keys): private_key, _ = rsa_keys fresh_access = _make_token(private_key) - keycloak: KeycloakClient = auth_client.app.state.keycloak + oidc: OidcClient = auth_client.app.state.oidc import unittest.mock as mock - keycloak.refresh_token = mock.AsyncMock(return_value={ + oidc.refresh_token = mock.AsyncMock(return_value={ "access_token": fresh_access, "refresh_token": "rotated-refresh", }) @@ -194,10 +201,10 @@ def test_refresh_endpoint_401_when_no_refresh_cookie(auth_client): def test_refresh_endpoint_401_when_keycloak_rejects(auth_client): - """When Keycloak rejects the refresh token, /refresh returns 401.""" - keycloak: KeycloakClient = auth_client.app.state.keycloak + """When the OIDC provider rejects the refresh token, /refresh returns 401.""" + oidc: OidcClient = auth_client.app.state.oidc import unittest.mock as mock - keycloak.refresh_token = mock.AsyncMock(side_effect=Exception("invalid_grant")) + oidc.refresh_token = mock.AsyncMock(side_effect=Exception("invalid_grant")) auth_client.cookies.set("cce_refresh", "expired-refresh") response = auth_client.post("/api/auth/refresh") diff --git a/packages/api/tests/test_config.py b/packages/api/tests/test_config.py index fbe2092..67bb929 100644 --- a/packages/api/tests/test_config.py +++ b/packages/api/tests/test_config.py @@ -158,3 +158,88 @@ def test_chat_disabled_when_no_key(): settings = Settings() assert settings.chat_enabled is False + + +def test_keycloak_backward_compat_resolution(): + # Only KEYCLOAK_* set (existing prod shape) — provider defaults to keycloak. + from cell_explorer_api.config import Settings + + s = Settings( + keycloak_url="https://auth.example.com", + keycloak_realm="cell-explorer", + keycloak_client_id="ce-app", + keycloak_client_secret="secret", + ) + assert s.auth_provider == "keycloak" + assert s.auth_enabled is True + assert s.resolved_issuer == "https://auth.example.com/realms/cell-explorer" + assert s.resolved_client_id == "ce-app" + assert s.resolved_client_secret == "secret" + assert s.resolved_audience == "ce-app" + assert s.resolved_scopes == "openid profile email" + assert s.resolved_roles_claims == [ + "realm_access.roles", + "resource_access.ce-app.roles", + ] + assert s.resolved_idp_hint is None + assert s.discovery_url == ( + "https://auth.example.com/realms/cell-explorer/.well-known/openid-configuration" + ) + + +def test_keycloak_idp_hint_flows_through(): + from cell_explorer_api.config import Settings + + s = Settings( + keycloak_url="https://a", keycloak_realm="r", + keycloak_client_id="c", keycloak_client_secret="x", + keycloak_idp_hint="pingId", + ) + assert s.resolved_idp_hint == "pingId" + + +def test_entra_resolution(): + from cell_explorer_api.config import Settings + + s = Settings( + auth_provider="entra", + oidc_issuer="https://login.microsoftonline.com/TENANT/v2.0", + oidc_client_id="app-guid", + oidc_client_secret="secret", + ) + assert s.auth_enabled is True + assert s.resolved_issuer == "https://login.microsoftonline.com/TENANT/v2.0" + assert s.resolved_roles_claims == ["roles"] + assert s.resolved_audience == "app-guid" + assert "offline_access" in s.resolved_scopes.split() + assert s.resolved_idp_hint is None + + +def test_roles_claims_override(): + from cell_explorer_api.config import Settings + + s = Settings( + auth_provider="oidc", + oidc_issuer="https://idp.example.com", + oidc_client_id="c", oidc_client_secret="x", + oidc_roles_claims="groups, my.custom.path", + ) + assert s.resolved_roles_claims == ["groups", "my.custom.path"] + + +def test_oidc_client_creds_fall_back_to_keycloak_vars(): + from cell_explorer_api.config import Settings + + s = Settings( + keycloak_url="https://a", keycloak_realm="r", + keycloak_client_id="kc", keycloak_client_secret="kx", + ) + assert s.resolved_client_id == "kc" + assert s.resolved_client_secret == "kx" + + +def test_auth_disabled_when_incomplete(): + from cell_explorer_api.config import Settings + + assert Settings().auth_enabled is False + assert Settings(auth_provider="entra", oidc_issuer="https://i").auth_enabled is False diff --git a/packages/api/tests/test_dataset_access.py b/packages/api/tests/test_dataset_access.py index 5d69cf2..9537161 100644 --- a/packages/api/tests/test_dataset_access.py +++ b/packages/api/tests/test_dataset_access.py @@ -12,7 +12,7 @@ from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.config import Settings from cell_explorer_api.db.models import Dataset, Datasource, DatasourceType from cell_explorer_api.main import create_app @@ -53,8 +53,15 @@ async def access_app(monkeypatch, tmp_path): ) app = create_app(settings) - keycloak: KeycloakClient = app.state.keycloak - keycloak._jwks = { + oidc: OidcClient = app.state.oidc + oidc._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) + oidc._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } diff --git a/packages/api/tests/test_datasets.py b/packages/api/tests/test_datasets.py index 0f1f4ef..b0c1563 100644 --- a/packages/api/tests/test_datasets.py +++ b/packages/api/tests/test_datasets.py @@ -13,7 +13,7 @@ from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.config import Settings from cell_explorer_api.db.models import Dataset, Datasource, DatasourceType from cell_explorer_api.main import create_app @@ -94,17 +94,24 @@ def test_list_datasets_anonymous_sees_public_only(seeded_app): def test_list_datasets_authenticated_sees_authorized(seeded_app): private_key, public_key = _generate_rsa_keypair() settings = seeded_app.state.settings - # Enable Keycloak for this test + # Enable OIDC for this test settings.keycloak_url = "https://auth.example.com" settings.keycloak_realm = "test-realm" settings.keycloak_client_id = "test-client" settings.keycloak_client_secret = "test-secret" - keycloak = KeycloakClient(settings) - keycloak._jwks = { + oidc = OidcClient(settings) + oidc._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) + oidc._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } - seeded_app.state.keycloak = keycloak + seeded_app.state.oidc = oidc token = jwt.encode( { @@ -138,11 +145,18 @@ def test_list_datasets_authenticated_no_matching_role(seeded_app): settings.keycloak_client_id = "test-client" settings.keycloak_client_secret = "test-secret" - keycloak = KeycloakClient(settings) - keycloak._jwks = { + oidc = OidcClient(settings) + oidc._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) + oidc._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } - seeded_app.state.keycloak = keycloak + seeded_app.state.oidc = oidc token = jwt.encode( { @@ -191,11 +205,18 @@ def test_private_dataset_url_is_null_for_anonymous(seeded_app): settings.keycloak_client_id = "test-client" settings.keycloak_client_secret = "test-secret" - keycloak = KeycloakClient(settings) - keycloak._jwks = { + oidc = OidcClient(settings) + oidc._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) + oidc._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } - seeded_app.state.keycloak = keycloak + seeded_app.state.oidc = oidc # Authenticated user with matching role token = jwt.encode( diff --git a/packages/api/tests/test_keycloak.py b/packages/api/tests/test_keycloak.py deleted file mode 100644 index fdbe2dd..0000000 --- a/packages/api/tests/test_keycloak.py +++ /dev/null @@ -1,142 +0,0 @@ -"""Tests for Keycloak OIDC client.""" - -import time - -import jwt -import pytest -from cryptography.hazmat.primitives.asymmetric import rsa -from cryptography.hazmat.primitives.serialization import Encoding, PublicFormat - -from cell_explorer_api.auth.keycloak import KeycloakClient -from cell_explorer_api.config import Settings - - -def _generate_rsa_keypair(): - """Generate an RSA key pair for testing.""" - private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) - return private_key, private_key.public_key() - - -def _make_settings(**overrides) -> Settings: - defaults = { - "keycloak_url": "https://auth.example.com", - "keycloak_realm": "test-realm", - "keycloak_client_id": "test-client", - "keycloak_client_secret": "test-secret", - } - return Settings(**(defaults | overrides)) - - -def _encode_token(claims: dict, private_key, kid: str = "test-kid") -> str: - return jwt.encode(claims, private_key, algorithm="RS256", headers={"kid": kid}) - - -@pytest.fixture() -def rsa_keys(): - return _generate_rsa_keypair() - - -@pytest.fixture() -def keycloak(rsa_keys): - settings = _make_settings() - client = KeycloakClient(settings) - _, public_key = rsa_keys - client._jwks = { - "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) - } - return client - - -def test_authorization_url(): - settings = _make_settings() - client = KeycloakClient(settings) - url = client.authorization_url( - redirect_uri="https://app.example.com/api/auth/callback", - state="abc123", - ) - assert "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth" in url - assert "client_id=test-client" in url - assert "redirect_uri=" in url - assert "state=abc123" in url - assert "response_type=code" in url - assert "kc_idp_hint" not in url - - -def test_authorization_url_with_idp_hint(): - settings = _make_settings(keycloak_idp_hint="pingId") - client = KeycloakClient(settings) - url = client.authorization_url( - redirect_uri="https://app.example.com/api/auth/callback", - state="abc123", - ) - assert "kc_idp_hint=pingId" in url - - -def test_decode_valid_token(keycloak, rsa_keys): - private_key, _ = rsa_keys - claims = { - "sub": "user-123", - "name": "Test User", - "email": "test@example.com", - "realm_access": {"roles": ["viewer"]}, - "iss": "https://auth.example.com/realms/test-realm", - "aud": "test-client", - "exp": int(time.time()) + 300, - "iat": int(time.time()), - } - token = _encode_token(claims, private_key) - user = keycloak.decode_token(token) - assert user.sub == "user-123" - assert user.name == "Test User" - assert user.email == "test@example.com" - assert user.roles == ["viewer"] - - -def test_decode_merges_realm_and_client_roles(keycloak, rsa_keys): - private_key, _ = rsa_keys - claims = { - "sub": "user-123", - "realm_access": {"roles": ["realm-role", "shared"]}, - "resource_access": { - "test-client": {"roles": ["client-role", "shared"]}, - "other-client": {"roles": ["should-not-include"]}, - }, - "iss": "https://auth.example.com/realms/test-realm", - "aud": "test-client", - "exp": int(time.time()) + 300, - "iat": int(time.time()), - } - token = _encode_token(claims, private_key) - user = keycloak.decode_token(token) - # Merged, deduplicated, sorted - assert user.roles == ["client-role", "realm-role", "shared"] - assert "should-not-include" not in user.roles - - -def test_decode_expired_token_raises(keycloak, rsa_keys): - private_key, _ = rsa_keys - claims = { - "sub": "user-123", - "iss": "https://auth.example.com/realms/test-realm", - "aud": "test-client", - # Past the 30s leeway window so the decoder genuinely rejects. - "exp": int(time.time()) - 60, - "iat": int(time.time()) - 300, - } - token = _encode_token(claims, private_key) - with pytest.raises(jwt.ExpiredSignatureError): - keycloak.decode_token(token) - - -def test_decode_invalid_signature_raises(keycloak): - other_key, _ = _generate_rsa_keypair() - claims = { - "sub": "user-123", - "iss": "https://auth.example.com/realms/test-realm", - "aud": "test-client", - "exp": int(time.time()) + 300, - "iat": int(time.time()), - } - token = _encode_token(claims, other_key) - with pytest.raises(jwt.InvalidSignatureError): - keycloak.decode_token(token) diff --git a/packages/api/tests/test_oidc.py b/packages/api/tests/test_oidc.py new file mode 100644 index 0000000..4bc49ae --- /dev/null +++ b/packages/api/tests/test_oidc.py @@ -0,0 +1,134 @@ +"""Tests for the generic OIDC client.""" + +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.hazmat.primitives.serialization import Encoding, PublicFormat + +from cell_explorer_api.auth.oidc import OidcClient, extract_roles +from cell_explorer_api.config import Settings + +DISCOVERY = { + "issuer": "https://idp.example.com", + "authorization_endpoint": "https://idp.example.com/authorize", + "token_endpoint": "https://idp.example.com/token", + "jwks_uri": "https://idp.example.com/jwks", + "end_session_endpoint": "https://idp.example.com/logout", +} + + +def _keys(): + pk = rsa.generate_private_key(public_exponent=65537, key_size=2048) + return pk, pk.public_key() + + +def _client(settings: Settings, public_key=None, kid="kid1") -> OidcClient: + c = OidcClient(settings) + c._apply_discovery(DISCOVERY) + if public_key is not None: + c._jwks = {kid: public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo)} + return c + + +def _entra_settings(**o): + return Settings(auth_provider="entra", oidc_issuer="https://idp.example.com", + oidc_client_id="app", oidc_client_secret="s", **o) + + +def _kc_settings(**o): + return Settings(keycloak_url="https://idp.example.com", keycloak_realm="r", + keycloak_client_id="app", keycloak_client_secret="s", **o) + + +class TestExtractRoles: + def test_keycloak_merges_realm_and_client_roles(self): + claims = { + "realm_access": {"roles": ["viewer", "admin"]}, + "resource_access": {"app": {"roles": ["curator"]}}, + } + roles = extract_roles(claims, ["realm_access.roles", "resource_access.app.roles"]) + assert roles == ["admin", "curator", "viewer"] + + def test_entra_single_roles_claim(self): + assert extract_roles({"roles": ["viewer"]}, ["roles"]) == ["viewer"] + + def test_missing_path_contributes_nothing(self): + assert extract_roles({"roles": ["viewer"]}, ["nope.here", "roles"]) == ["viewer"] + + def test_no_paths_yields_no_roles(self): + assert extract_roles({"roles": ["viewer"]}, []) == [] + + +class TestAuthorizationUrl: + def test_includes_scopes_and_no_idp_hint_for_entra(self): + url = _client(_entra_settings()).authorization_url("https://cb/cb", "st8") + assert url.startswith("https://idp.example.com/authorize?") + assert "scope=openid+profile+email+offline_access" in url + assert "kc_idp_hint" not in url + assert "state=st8" in url and "response_type=code" in url + + def test_keycloak_idp_hint_included(self): + url = _client(_kc_settings(keycloak_idp_hint="pingId")).authorization_url("https://cb", "s") + assert "kc_idp_hint=pingId" in url + + +class TestDecodeToken: + def test_entra_roles_and_aud_iss(self): + pk, pub = _keys() + c = _client(_entra_settings(), public_key=pub) + token = jwt.encode( + {"sub": "u1", "name": "N", "email": "e@x", "roles": ["viewer"], + "aud": "app", "iss": "https://idp.example.com"}, + pk, algorithm="RS256", headers={"kid": "kid1"}, + ) + user = c.decode_token(token) + assert user.sub == "u1" and user.roles == ["viewer"] + + def test_wrong_audience_rejected(self): + pk, pub = _keys() + c = _client(_entra_settings(), public_key=pub) + token = jwt.encode( + {"sub": "u", "aud": "WRONG", "iss": "https://idp.example.com"}, + pk, algorithm="RS256", headers={"kid": "kid1"}, + ) + with pytest.raises(jwt.InvalidAudienceError): + c.decode_token(token) + + def test_wrong_issuer_rejected(self): + pk, pub = _keys() + c = _client(_entra_settings(), public_key=pub) + token = jwt.encode( + {"sub": "u", "aud": "app", "iss": "https://attacker.example.com"}, + pk, algorithm="RS256", headers={"kid": "kid1"}, + ) + with pytest.raises(jwt.InvalidIssuerError): + c.decode_token(token) + + def test_keycloak_dual_role_merge(self): + pk, pub = _keys() + c = _client(_kc_settings(), public_key=pub) + token = jwt.encode( + {"sub": "u", "aud": "app", "iss": "https://idp.example.com", + "realm_access": {"roles": ["viewer"]}, + "resource_access": {"app": {"roles": ["curator"]}}}, + pk, algorithm="RS256", headers={"kid": "kid1"}, + ) + assert c.decode_token(token).roles == ["curator", "viewer"] + + +class TestDiscoveryParsing: + def test_apply_discovery_sets_endpoints(self): + c = OidcClient(_entra_settings()) + c._apply_discovery(DISCOVERY) + assert c._authorization_endpoint == "https://idp.example.com/authorize" + assert c._token_endpoint == "https://idp.example.com/token" + assert c._jwks_uri == "https://idp.example.com/jwks" + assert c._issuer == "https://idp.example.com" + assert c._end_session_endpoint == "https://idp.example.com/logout" + + +class TestLogoutUrl: + def test_uses_end_session_endpoint(self): + url = _client(_entra_settings()).logout_url("https://cb/home") + assert url.startswith("https://idp.example.com/logout?") + assert "post_logout_redirect_uri=https%3A%2F%2Fcb%2Fhome" in url diff --git a/packages/api/tests/test_optional_auth.py b/packages/api/tests/test_optional_auth.py index c42e853..ecb7587 100644 --- a/packages/api/tests/test_optional_auth.py +++ b/packages/api/tests/test_optional_auth.py @@ -9,7 +9,7 @@ from fastapi import APIRouter, Depends, FastAPI from fastapi.testclient import TestClient -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.auth.models import User from cell_explorer_api.auth.optional import optional_auth from cell_explorer_api.config import Settings @@ -54,8 +54,15 @@ def test_optional_auth_returns_user_when_authenticated(): keycloak_client_secret="test-secret", ) app = _make_test_app(settings) - keycloak: KeycloakClient = app.state.keycloak - keycloak._jwks = { + oidc: OidcClient = app.state.oidc + oidc._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) + oidc._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } token = jwt.encode( diff --git a/packages/api/tests/test_token_refresh_middleware.py b/packages/api/tests/test_token_refresh_middleware.py index 6f25e27..ee9db3b 100644 --- a/packages/api/tests/test_token_refresh_middleware.py +++ b/packages/api/tests/test_token_refresh_middleware.py @@ -22,7 +22,7 @@ from fastapi.testclient import TestClient from cell_explorer_api.auth.dependencies import require_auth -from cell_explorer_api.auth.keycloak import KeycloakClient +from cell_explorer_api.auth.oidc import OidcClient from cell_explorer_api.auth.middleware import TokenRefreshMiddleware from cell_explorer_api.auth.models import User from cell_explorer_api.config import Settings @@ -65,17 +65,24 @@ def rsa_keys(): def keycloak(rsa_keys): _, public_key = rsa_keys settings = _make_settings() - client = KeycloakClient(settings) + client = OidcClient(settings) + client._apply_discovery({ + "issuer": "https://auth.example.com/realms/test-realm", + "authorization_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/auth", + "token_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/token", + "jwks_uri": "https://auth.example.com/realms/test-realm/protocol/openid-connect/certs", + "end_session_endpoint": "https://auth.example.com/realms/test-realm/protocol/openid-connect/logout", + }) client._jwks = { "test-kid": public_key.public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo) } return client -def _app_with_middleware(keycloak: KeycloakClient) -> FastAPI: +def _app_with_middleware(keycloak: OidcClient) -> FastAPI: app = FastAPI() app.state.settings = _make_settings() - app.state.keycloak = keycloak + app.state.oidc = keycloak app.add_middleware(TokenRefreshMiddleware) @app.get("/json")