# 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()