Source code for fsh_lib.auth

# NOTE: ``auth_context_dependency`` and the transport ``extract_dep``
# below build inner functions annotated with
# ``Annotated[..., Depends(<closure-local>)]``.  pydantic's
# ``TypeAdapter`` calls ``typing.get_type_hints`` against those
# inner functions when FastAPI builds the OpenAPI schema; closure
# locals aren't in ``__globals__``, so any stringified annotation
# (PEP 563) fails to resolve and 500s the schema build.  PEP 749's
# default deferred-but-lazy evaluation in 3.14 keeps annotations as
# real objects, preserving the closure scope -- but only as long as
# nothing forces them back to strings.  ``collections.abc`` must
# therefore be imported at runtime (not under ``TYPE_CHECKING``) so
# the same closure-local resolution can find ``Awaitable``,
# ``Callable``, and ``Sequence`` at request time.

import os
from collections.abc import (  # noqa: TC003 -- runtime, see NOTE above
    Awaitable,
    Callable,
    Sequence,
)
from datetime import UTC, datetime, timedelta
from typing import Annotated, Any, Literal

import jwt
from fastapi import Cookie, Depends, HTTPException, Response, status
from fastapi.security import OAuth2PasswordBearer
from pydantic import BaseModel
from sqlalchemy.ext.asyncio import (  # noqa: TC002 -- runtime, see NOTE above
    AsyncSession,
)

from fsh_lib.db import set_token

DEFAULT_TOKEN_TTL = timedelta(minutes=30)

TOKEN_ID_CLAIM = "tid"  # noqa: S105 -- claim name, not a secret

FLOW_STATE_PURPOSE = "oauth_flow"

Source = Literal["bearer", "cookie"]
SameSite = Literal["lax", "strict", "none"]


[docs] class LoginResponse(BaseModel): access_token: str token_type: Literal["bearer"] = "bearer" # noqa: S105 -- not a secret
[docs] class OkResponse(BaseModel): ok: Literal[True] = True
def _unauthorized() -> HTTPException: return HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated", headers={"WWW-Authenticate": "Bearer"}, ) def _revoked() -> HTTPException: return HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Session revoked", headers={"WWW-Authenticate": "Bearer"}, ) class _Transport: @classmethod def from_config(cls, **kwargs: Any) -> _Transport: # noqa: ANN401 raise NotImplementedError def extract_dep(self) -> Callable[..., Awaitable[str | None]]: raise NotImplementedError def emit( self, response: Response, token: str, ttl: timedelta, ) -> LoginResponse | None: raise NotImplementedError def clear(self, response: Response) -> None: raise NotImplementedError class _BearerTransport(_Transport): def __init__(self) -> None: self._oauth = OAuth2PasswordBearer(tokenUrl="/auth", auto_error=False) @classmethod def from_config(cls, **kwargs: Any) -> _BearerTransport: # noqa: ANN401, ARG003 return cls() def extract_dep(self) -> Callable[..., Awaitable[str | None]]: async def _extract( bearer: Annotated[str | None, Depends(self._oauth)] = None, ) -> str | None: return bearer return _extract def emit( self, response: Response, # noqa: ARG002 token: str, ttl: timedelta, # noqa: ARG002 ) -> LoginResponse | None: return LoginResponse(access_token=token) def clear(self, response: Response) -> None: # noqa: ARG002 # Bearer logout is client-side -- clients discard the token. return class _CookieTransport(_Transport): def __init__( self, *, name: str = "access_token", secure: bool = True, samesite: SameSite = "lax", ) -> None: self._name = name self._secure = secure self._samesite = samesite @classmethod def from_config(cls, **kwargs: Any) -> _CookieTransport: # noqa: ANN401 return cls( name=kwargs.get("cookie_name", "access_token"), secure=kwargs.get("cookie_secure", True), samesite=kwargs.get("cookie_samesite", "lax"), ) def extract_dep(self) -> Callable[..., Awaitable[str | None]]: name = self._name async def _extract( cookie: Annotated[str | None, Cookie(alias=name)] = None, ) -> str | None: return cookie return _extract def emit( self, response: Response, token: str, ttl: timedelta, ) -> LoginResponse | None: response.set_cookie( key=self._name, value=token, max_age=int(ttl.total_seconds()), httponly=True, secure=self._secure, samesite=self._samesite, ) return None def clear(self, response: Response) -> None: response.delete_cookie( key=self._name, httponly=True, secure=self._secure, samesite=self._samesite, ) _TRANSPORTS: dict[Source, type[_Transport]] = { "bearer": _BearerTransport, "cookie": _CookieTransport, } async def _no_token() -> str | None: return None def _build_transports( sources: Sequence[Source], **config: Any, # noqa: ANN401 ) -> dict[Source, _Transport]: unknown = [src for src in sources if src not in _TRANSPORTS] if unknown: msg = f"unknown source(s): {sorted(set(unknown))}" raise ValueError(msg) if not sources: msg = f"sources must contain at least one of {sorted(_TRANSPORTS)}" raise ValueError(msg) return {src: _TRANSPORTS[src].from_config(**config) for src in sources} def encode_jwt( payload: dict[str, Any], *, secret_env: str, algorithm: str, ttl: timedelta = DEFAULT_TOKEN_TTL, ) -> str: claims = dict(payload) claims.setdefault( "exp", datetime.now(tz=UTC) + ttl, ) return jwt.encode(claims, os.environ[secret_env], algorithm=algorithm) def decode_jwt( token: str, *, secret_env: str, algorithm: str, ) -> dict[str, Any]: if not (secret := os.environ.get(secret_env)): raise _unauthorized() try: return jwt.decode(token, secret, algorithms=[algorithm]) except jwt.InvalidTokenError as exc: raise _unauthorized() from exc def auth_context_dependency[T: BaseModel]( _schema: type[T], sources: Sequence[Source], *, load: Callable[[str], Awaitable[BaseModel | None]], secret_env: str = "JWT_SECRET", # noqa: S107 algorithm: str = "HS256", cookie_name: str = "access_token", public_auth_context: T | None = None, ) -> Callable[..., Awaitable[T]]: transports = _build_transports( sources, cookie_name=cookie_name, ) bearer_transport = transports.get("bearer") cookie_transport = transports.get("cookie") bearer_ext = ( bearer_transport.extract_dep() if bearer_transport else _no_token ) cookie_ext = ( cookie_transport.extract_dep() if cookie_transport else _no_token ) async def resolve(token: str | None) -> T: if token is None: if public_auth_context is not None: return public_auth_context raise _unauthorized() claims = decode_jwt(token, secret_env=secret_env, algorithm=algorithm) if claims.get("purpose") == FLOW_STATE_PURPOSE: raise _unauthorized() token_id = claims.get(TOKEN_ID_CLAIM) if not isinstance(token_id, str): raise _unauthorized() auth = await load(token_id) if not auth: raise _unauthorized() return auth # type: ignore[return-value] async def get_auth_context( bearer: Annotated[str | None, Depends(bearer_ext)] = None, cookie: Annotated[str | None, Depends(cookie_ext)] = None, ) -> T: return await resolve(bearer or cookie) return get_auth_context def construct_audited_auth_context_dependency( inner: Callable[..., Awaitable[BaseModel]], *, db_dep: Callable[..., Any], token_id_attr: str = "token_id", # noqa: S107 ) -> Callable[..., Awaitable[BaseModel]]: async def get_auth_context( db: Annotated[AsyncSession, Depends(db_dep)], principal: Annotated[BaseModel, Depends(inner)], ) -> BaseModel: token_id = getattr(principal, token_id_attr) user_id = getattr(principal, "user_id", None) set_token(db.sync_session, token_id, user_id) return principal return get_auth_context def issue_token( response: Response, token_id: str | None, sources: Sequence[Source], *, secret_env: str = "JWT_SECRET", # noqa: S107 algorithm: str = "HS256", ttl: timedelta = DEFAULT_TOKEN_TTL, cookie_name: str | None = "access_token", cookie_secure: bool = True, cookie_samesite: SameSite = "lax", ) -> LoginResponse | OkResponse: if token_id is None: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid credentials", headers={"WWW-Authenticate": "Bearer"}, ) transports = _build_transports( sources, cookie_name=cookie_name, cookie_secure=cookie_secure, cookie_samesite=cookie_samesite, ) token = encode_jwt( {TOKEN_ID_CLAIM: token_id}, secret_env=secret_env, algorithm=algorithm, ttl=ttl, ) body: LoginResponse | OkResponse = OkResponse() for transport in transports.values(): if emitted := transport.emit(response, token, ttl): body = emitted return body def clear_token( response: Response, sources: Sequence[Source], **config: Any, # noqa: ANN401 ) -> OkResponse: transports = _build_transports(sources, **config) for transport in transports.values(): transport.clear(response) return OkResponse()