diff --git a/docs/module/security.md b/docs/module/security.md new file mode 100644 index 0000000..6ddddb5 --- /dev/null +++ b/docs/module/security.md @@ -0,0 +1,267 @@ +# Security + +Composable authentication helpers for FastAPI that use `Security()` for OpenAPI documentation and accept user-provided validator functions with full type flexibility. + +## Overview + +The `security` module provides four auth source classes and a `MultiAuth` factory. Each class wraps a FastAPI security scheme for OpenAPI and accepts a validator function called as: + +```python +await validator(credential, **kwargs) +``` + +where `kwargs` are the extra keyword arguments provided at instantiation (roles, permissions, enums, etc.). The validator returns the authenticated identity (e.g. a `User` model) which becomes the route dependency value. + +```python +from fastapi import Security +from fastapi_toolsets.security import BearerTokenAuth + +async def verify_token(token: str, *, role: str) -> User: + user = await db.get_by_token(token) + if not user or user.role != role: + raise UnauthorizedError() + return user + +bearer_admin = BearerTokenAuth(verify_token, role="admin") + +@app.get("/admin") +async def admin_route(user: User = Security(bearer_admin)): + return user +``` + +## Auth sources + +### [`BearerTokenAuth`](../reference/security.md#fastapi_toolsets.security.BearerTokenAuth) + +Reads the `Authorization: Bearer ` header. Wraps `HTTPBearer` for OpenAPI. + +```python +from fastapi_toolsets.security import BearerTokenAuth + +bearer = BearerTokenAuth(validator=verify_token) + +@app.get("/me") +async def me(user: User = Security(bearer)): + return user +``` + +#### Token prefix + +The optional `prefix` parameter restricts a `BearerTokenAuth` instance to tokens +that start with a given string. The prefix is **kept** in the value passed to the +validator — store and compare tokens with their prefix included. + +This lets you deploy multiple `BearerTokenAuth` instances in the same application +and disambiguate them efficiently in `MultiAuth`: + +```python +user_bearer = BearerTokenAuth(verify_user, prefix="user_") # matches "Bearer user_..." +org_bearer = BearerTokenAuth(verify_org, prefix="org_") # matches "Bearer org_..." +``` + +Use [`generate_token()`](#token-generation) to create correctly-prefixed tokens. + +#### Token generation + +`BearerTokenAuth.generate_token()` produces a secure random token ready to store +in your database and return to the client. If a prefix is configured it is +prepended automatically: + +```python +bearer = BearerTokenAuth(verify_token, prefix="user_") + +token = bearer.generate_token() # e.g. "user_Xk3mN..." +await db.store_token(user_id, token) +return {"access_token": token, "token_type": "bearer"} +``` + +The client sends `Authorization: Bearer user_Xk3mN...` and the validator receives +the full token (prefix included) to compare against the stored value. + +### [`CookieAuth`](../reference/security.md#fastapi_toolsets.security.CookieAuth) + +Reads a named cookie. Wraps `APIKeyCookie` for OpenAPI. + +```python +from fastapi_toolsets.security import CookieAuth + +cookie_auth = CookieAuth("session", validator=verify_session) + +@app.get("/me") +async def me(user: User = Security(cookie_auth)): + return user +``` + +### [`OAuth2Auth`](../reference/security.md#fastapi_toolsets.security.OAuth2Auth) + +Reads the `Authorization: Bearer ` header and registers the token endpoint +in OpenAPI via `OAuth2PasswordBearer`. + +```python +from fastapi_toolsets.security import OAuth2Auth + +oauth2_auth = OAuth2Auth(token_url="/token", validator=verify_token) + +@app.get("/me") +async def me(user: User = Security(oauth2_auth)): + return user +``` + +### [`OpenIDAuth`](../reference/security.md#fastapi_toolsets.security.OpenIDAuth) + +Reads the `Authorization: Bearer ` header and registers the OpenID Connect +discovery URL in OpenAPI via `OpenIdConnect`. Token validation is fully delegated +to your validator — use any OIDC / JWT library (`authlib`, `python-jose`, `PyJWT`). + +```python +from fastapi_toolsets.security import OpenIDAuth + +async def verify_google_token(token: str, *, audience: str) -> User: + payload = jwt.decode(token, google_public_keys, algorithms=["RS256"], + audience=audience) + return User(email=payload["email"], name=payload["name"]) + +google_auth = OpenIDAuth( + "https://accounts.google.com/.well-known/openid-configuration", + verify_google_token, + audience="my-client-id", +) + +@app.get("/me") +async def me(user: User = Security(google_auth)): + return user +``` + +The discovery URL is used **only for OpenAPI documentation** — no requests are made +to it by this class. You are responsible for fetching and caching the provider's +public keys in your validator. + +Multiple providers work naturally with `MultiAuth`: + +```python +multi = MultiAuth(google_auth, github_auth) + +@app.get("/data") +async def data(user: User = Security(multi)): + return user +``` + +## Typed validator kwargs + +All auth classes forward extra instantiation keyword arguments to the validator. +Arguments can be any type — enums, strings, integers, etc. The validator returns +the authenticated identity, which FastAPI injects directly into the route handler. + +```python +async def verify_token(token: str, *, role: Role, permission: str) -> User: + user = await decode_token(token) + if user.role != role or permission not in user.permissions: + raise UnauthorizedError() + return user + +bearer = BearerTokenAuth(verify_token, role=Role.ADMIN, permission="billing:read") +``` + +Each auth instance is self-contained — create a separate instance per distinct +requirement instead of passing requirements through `Security(scopes=[...])`. + +### Using `.require()` inline + +If declaring a new top-level variable per role feels verbose, use `.require()` to +create a configured clone directly in the route decorator. The original instance +is not mutated: + +```python +bearer = BearerTokenAuth(verify_token) + +@app.get("/admin/stats") +async def admin_stats(user: User = Security(bearer.require(role=Role.ADMIN))): + return {"message": f"Hello admin {user.name}"} + +@app.get("/profile") +async def profile(user: User = Security(bearer.require(role=Role.USER))): + return {"id": user.id, "name": user.name} +``` + +`.require()` kwargs are merged over existing ones — new values win on conflict. +The `prefix` (for `BearerTokenAuth`) and cookie name (for `CookieAuth`) are +always preserved. + +`.require()` instances work transparently inside `MultiAuth`: + +```python +multi = MultiAuth( + user_bearer.require(role=Role.USER), + org_bearer.require(role=Role.ADMIN), +) +``` + +## MultiAuth + +[`MultiAuth`](../reference/security.md#fastapi_toolsets.security.MultiAuth) combines +multiple auth sources into a single callable. Sources are tried in order; the +first one that finds a credential wins. + +```python +from fastapi_toolsets.security import MultiAuth + +multi = MultiAuth(user_bearer, org_bearer, cookie_auth) + +@app.get("/data") +async def data_route(user = Security(multi)): + return user +``` + +### Using `.require()` on MultiAuth + +`MultiAuth` also supports `.require()`, which propagates the kwargs to every +source that implements it. Sources that do not (e.g. custom `AuthSource` +subclasses) are passed through unchanged: + +```python +multi = MultiAuth(bearer, cookie) + +@app.get("/admin") +async def admin(user: User = Security(multi.require(role=Role.ADMIN))): + return user +``` + +This is equivalent to calling `.require()` on each source individually: + +```python +# These two are identical +multi.require(role=Role.ADMIN) + +MultiAuth( + bearer.require(role=Role.ADMIN), + cookie.require(role=Role.ADMIN), +) +``` + +### Prefix-based dispatch + +Because `extract()` is pure string matching (no I/O), prefix-based source +selection is essentially free. Only the matching source's validator (which may +involve DB or network I/O) is ever called: + +```python +user_bearer = BearerTokenAuth(verify_user, prefix="user_") +org_bearer = BearerTokenAuth(verify_org, prefix="org_") + +multi = MultiAuth(user_bearer, org_bearer) + +# "Bearer user_alice" → only verify_user runs, receives "user_alice" +# "Bearer org_acme" → only verify_org runs, receives "org_acme" +``` + +Tokens are stored and compared **with their prefix** — use `generate_token()` on +each source to issue correctly-prefixed tokens: + +```python +user_token = user_bearer.generate_token() # "user_..." +org_token = org_bearer.generate_token() # "org_..." +``` + +--- + +[:material-api: API Reference](../reference/security.md) diff --git a/docs/reference/security.md b/docs/reference/security.md new file mode 100644 index 0000000..f38235d --- /dev/null +++ b/docs/reference/security.md @@ -0,0 +1,28 @@ +# `security` + +Here's the reference for the authentication helpers provided by the `security` module. + +You can import them directly from `fastapi_toolsets.security`: + +```python +from fastapi_toolsets.security import ( + AuthSource, + BearerTokenAuth, + CookieAuth, + OAuth2Auth, + OpenIDAuth, + MultiAuth, +) +``` + +## ::: fastapi_toolsets.security.AuthSource + +## ::: fastapi_toolsets.security.BearerTokenAuth + +## ::: fastapi_toolsets.security.CookieAuth + +## ::: fastapi_toolsets.security.OAuth2Auth + +## ::: fastapi_toolsets.security.OpenIDAuth + +## ::: fastapi_toolsets.security.MultiAuth diff --git a/src/fastapi_toolsets/security/__init__.py b/src/fastapi_toolsets/security/__init__.py new file mode 100644 index 0000000..2ff9c43 --- /dev/null +++ b/src/fastapi_toolsets/security/__init__.py @@ -0,0 +1,14 @@ +"""Authentication helpers for FastAPI using Security().""" + +from .base import AuthSource +from .multi import MultiAuth +from .sources import BearerTokenAuth, CookieAuth, OAuth2Auth, OpenIDAuth + +__all__ = [ + "AuthSource", + "BearerTokenAuth", + "CookieAuth", + "OAuth2Auth", + "OpenIDAuth", + "MultiAuth", +] diff --git a/src/fastapi_toolsets/security/base.py b/src/fastapi_toolsets/security/base.py new file mode 100644 index 0000000..7d82db7 --- /dev/null +++ b/src/fastapi_toolsets/security/base.py @@ -0,0 +1,112 @@ +"""Abstract base class for authentication sources.""" + +import inspect +from abc import ABC, abstractmethod +from typing import Any, Callable + +from fastapi import Request +from fastapi.security import SecurityScopes + +from fastapi_toolsets.exceptions import UnauthorizedError + + +class AuthSource(ABC): + """Abstract base class for authentication sources. + + Subclass this to create a custom auth source that works with + :func:`~fastapi_toolsets.security.MultiAuth` and can be used directly + with :func:`fastapi.Security`. + + Concrete subclasses must implement :meth:`extract` and + :meth:`authenticate`. The default :meth:`__call__` (set up in + :meth:`__init__`) wires them together for FastAPI dependency injection. + + Custom subclasses with their own ``__init__`` **must** call + ``super().__init__()`` to activate the default dependency behaviour:: + + class JWTAuth(AuthSource): + def __init__(self, secret: str, *, role: str | None = None): + super().__init__() # required + self._secret = secret + self._role = role + + async def extract(self, request: Request) -> str | None: + auth = request.headers.get("Authorization", "") + if not auth.startswith("Bearer "): + return None + return auth[7:] or None + + async def authenticate(self, credential: str) -> User: + payload = jwt.decode(credential, self._secret) + if self._role and payload.get("role") != self._role: + raise UnauthorizedError() + return User(**payload) + + jwt_auth = JWTAuth(secret="mysecret") + + @app.get("/me") + async def me(user: User = Security(jwt_auth)): + return user + + # Works with MultiAuth too + multi = MultiAuth(jwt_auth, CookieAuth("session", verify_session)) + + .. note:: + The default ``__call__`` does not register a security scheme in the + OpenAPI spec. Built-in sources (``BearerTokenAuth`` etc.) override + ``__call__`` using the ``__signature__`` trick to provide a FastAPI + security scheme for Swagger UI. + """ + + def __init__(self) -> None: + """Set up the default FastAPI dependency signature. + + Creates a closure that FastAPI can introspect to inject + :class:`fastapi.Request` and :class:`fastapi.security.SecurityScopes`. + The :meth:`__signature__` attribute is set so that ``inspect.signature`` + (which FastAPI uses internally) returns the correct parameter list. + + Subclasses with their own ``__init__`` must call ``super().__init__()``. + Built-in subclasses (``BearerTokenAuth`` etc.) skip this and set up + their own ``_call_fn`` / ``__signature__`` directly. + """ + source = self + + async def _call( + request: Request, + security_scopes: SecurityScopes, # noqa: ARG001 + ) -> Any: + credential = await source.extract(request) + if credential is None: + raise UnauthorizedError() + return await source.authenticate(credential) + + self._call_fn: Callable[..., Any] = _call + self.__signature__ = inspect.signature(_call) + + @abstractmethod + async def extract(self, request: Request) -> str | None: + """Extract the raw credential from the request without validating. + + Returns ``None`` if no credential is present for this source. + This method must be fast and free of I/O — it is called by + :func:`~fastapi_toolsets.security.MultiAuth` for every source on + every request. + """ + + @abstractmethod + async def authenticate(self, credential: str) -> Any: + """Validate a credential and return the authenticated identity. + + Should raise :class:`~fastapi_toolsets.exceptions.UnauthorizedError` + (or any exception) when the credential is invalid. The return value + is injected into the route handler as the dependency value. + """ + + async def __call__(self, **kwargs: Any) -> Any: + """FastAPI dependency dispatch. + + Delegates to the closure stored in ``_call_fn``, whose signature + (stored in ``__signature__``) tells FastAPI which parameters to inject. + """ + return await self._call_fn(**kwargs) diff --git a/src/fastapi_toolsets/security/multi.py b/src/fastapi_toolsets/security/multi.py new file mode 100644 index 0000000..9420f40 --- /dev/null +++ b/src/fastapi_toolsets/security/multi.py @@ -0,0 +1,98 @@ +"""MultiAuth: combine multiple authentication sources into a single callable.""" + +import inspect +from typing import Any, cast + +from fastapi import Request +from fastapi.security import SecurityScopes + +from fastapi_toolsets.exceptions import UnauthorizedError + +from .base import AuthSource + + +class MultiAuth: + """Combine multiple authentication sources into a single callable. + + Sources are tried in order; the first one whose + :meth:`~AuthSource.extract` returns a non-``None`` credential wins. + Its :meth:`~AuthSource.authenticate` is called and the result returned. + + If a credential is found but the validator raises, the exception propagates + immediately — the remaining sources are **not** tried. This prevents + silent fallthrough on invalid credentials. + + If no source provides a credential, + :class:`~fastapi_toolsets.exceptions.UnauthorizedError` is raised. + + The :meth:`~AuthSource.extract` method of each source performs only + string matching (no I/O), so prefix-based dispatch is essentially free. + + Any :class:`~AuthSource` subclass — including user-defined ones — can be + passed as a source. + + Args: + *sources: Auth source instances to try in order. + + Example:: + + user_bearer = BearerTokenAuth(verify_user, prefix="user_") + org_bearer = BearerTokenAuth(verify_org, prefix="org_") + cookie = CookieAuth("session", verify_session) + + multi = MultiAuth(user_bearer, org_bearer, cookie) + + @app.get("/data") + async def data_route(user = Security(multi)): + return user + + # Apply a shared requirement to all sources at once + @app.get("/admin") + async def admin_route(user = Security(multi.require(role=Role.ADMIN))): + return user + """ + + def __init__(self, *sources: AuthSource) -> None: + self._sources = sources + + _sources = sources + + async def _call( + request: Request, + security_scopes: SecurityScopes, # noqa: ARG001 + ) -> Any: + for source in _sources: + credential = await source.extract(request) + if credential is not None: + return await source.authenticate(credential) + raise UnauthorizedError() + + self._call_fn = _call + self.__signature__ = inspect.signature(_call) + + async def __call__(self, **kwargs: Any) -> Any: + return await self._call_fn(**kwargs) + + def require(self, **kwargs: Any) -> "MultiAuth": + """Return a new :class:`MultiAuth` with kwargs forwarded to each source. + + Calls ``.require(**kwargs)`` on every source that supports it. Sources + that do not implement ``.require()`` (e.g. custom :class:`~AuthSource` + subclasses) are passed through unchanged. + + New kwargs are merged over each source's existing kwargs — new values + win on conflict:: + + multi = MultiAuth(bearer, cookie) + + @app.get("/admin") + async def admin(user = Security(multi.require(role=Role.ADMIN))): + return user + """ + new_sources = tuple( + cast(Any, source).require(**kwargs) + if hasattr(source, "require") + else source + for source in self._sources + ) + return MultiAuth(*new_sources) diff --git a/src/fastapi_toolsets/security/sources/__init__.py b/src/fastapi_toolsets/security/sources/__init__.py new file mode 100644 index 0000000..5fd948f --- /dev/null +++ b/src/fastapi_toolsets/security/sources/__init__.py @@ -0,0 +1,8 @@ +"""Built-in authentication source implementations.""" + +from .bearer import BearerTokenAuth +from .cookie import CookieAuth +from .oauth2 import OAuth2Auth +from .openid import OpenIDAuth + +__all__ = ["BearerTokenAuth", "CookieAuth", "OAuth2Auth", "OpenIDAuth"] diff --git a/src/fastapi_toolsets/security/sources/bearer.py b/src/fastapi_toolsets/security/sources/bearer.py new file mode 100644 index 0000000..7eefcf7 --- /dev/null +++ b/src/fastapi_toolsets/security/sources/bearer.py @@ -0,0 +1,171 @@ +"""Bearer token authentication source.""" + +import inspect +import secrets +from typing import Annotated, Any, Callable + +from fastapi import Depends +from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer, SecurityScopes + +from fastapi_toolsets.exceptions import UnauthorizedError + +from ..base import AuthSource + + +class BearerTokenAuth(AuthSource): + """Bearer token authentication source. + + Wraps :class:`fastapi.security.HTTPBearer` for OpenAPI documentation. + The validator is called as ``await validator(credential, **kwargs)`` + where ``kwargs`` are the extra keyword arguments provided at instantiation. + + Args: + validator: Async callable that receives the credential and any extra + keyword arguments, and returns the authenticated identity (e.g. a + ``User`` model). Should raise + :class:`~fastapi_toolsets.exceptions.UnauthorizedError` on failure. + prefix: Optional token prefix (e.g. ``"user_"``). If set, only tokens + whose value starts with this prefix are matched. The prefix is + **kept** in the value passed to the validator — store and compare + tokens with their prefix included. Use :meth:`generate_token` to + create correctly-prefixed tokens. This enables multiple + ``BearerTokenAuth`` instances in the same app (e.g. ``"user_"`` + for user tokens, ``"org_"`` for org tokens). + **kwargs: Extra keyword arguments forwarded to the validator on every + call (e.g. ``role=Role.ADMIN``). + + Example:: + + async def verify_token(token: str, *, role: Role) -> User: + user = await db.get_by_token(token) # token includes prefix + if not user or user.role != role: + raise UnauthorizedError() + return user + + bearer_admin = BearerTokenAuth(verify_token, prefix="user_", role=Role.ADMIN) + + # Generate a token to store in DB and return to the client: + token = bearer_admin.generate_token() # e.g. "user_Xk3..." + + @app.get("/admin") + async def admin_route(user: User = Security(bearer_admin)): + return user + """ + + def __init__( + self, + validator: Callable[..., Any], + *, + prefix: str | None = None, + **kwargs: Any, + ) -> None: + self._validator = validator + self._prefix = prefix + self._kwargs = kwargs + self._scheme = HTTPBearer(auto_error=False) + + # Capture locals for the closure — self._scheme cannot be referenced + # inside the Annotated default because annotations are evaluated at + # function-definition time (no `from __future__ import annotations`). + _scheme = self._scheme + _validator = validator + _kwargs = kwargs + _prefix = prefix + + async def _call( + # security_scopes is unused in the body but its presence in the + # signature tells FastAPI to aggregate scopes from Security() calls + # up the dependency chain and expose them in the OpenAPI schema. + security_scopes: SecurityScopes, # noqa: ARG001 + credentials: Annotated[ + HTTPAuthorizationCredentials | None, Depends(_scheme) + ] = None, + ) -> Any: + if credentials is None: + raise UnauthorizedError() + token = credentials.credentials + if _prefix is not None and not token.startswith(_prefix): + raise UnauthorizedError() + return await _validator(token, **_kwargs) + + # __call__ must be defined on the class (not the instance) so that + # callable(self) returns True. We expose the closure's signature via + # __signature__ so FastAPI resolves the correct sub-dependencies. + self._call_fn = _call + self.__signature__ = inspect.signature(_call) + + async def __call__(self, **kwargs: Any) -> Any: + return await self._call_fn(**kwargs) + + async def extract(self, request: Any) -> str | None: + """Extract the raw credential from the request without validating. + + Returns ``None`` if no ``Authorization: Bearer`` header is present, + the token is empty, or the token does not match the configured prefix. + The prefix is included in the returned value. + """ + auth = request.headers.get("Authorization", "") + if not auth.startswith("Bearer "): + return None + token = auth[7:] + if not token: + return None + if self._prefix is not None and not token.startswith(self._prefix): + return None + return token + + async def authenticate(self, credential: str) -> Any: + """Validate a credential and return the identity. + + Calls ``await validator(credential, **kwargs)`` where ``kwargs`` are + the extra keyword arguments provided at instantiation. + """ + return await self._validator(credential, **self._kwargs) + + def require(self, **kwargs: Any) -> "BearerTokenAuth": + """Return a new instance with additional (or overriding) validator kwargs. + + Useful for specifying per-endpoint requirements inline without + declaring a new top-level variable:: + + bearer = BearerTokenAuth(verify_token) + + @app.get("/admin") + async def admin(user: User = Security(bearer.require(role=Role.ADMIN))): + return user + + The ``prefix`` is preserved. New kwargs are merged over existing ones + (new values win on conflict). + """ + return BearerTokenAuth( + self._validator, + prefix=self._prefix, + **{**self._kwargs, **kwargs}, + ) + + def generate_token(self, nbytes: int = 32) -> str: + """Generate a secure random token for this auth source. + + Returns a URL-safe random token. If a prefix is configured it is + prepended — the returned value is what you store in your database + and return to the client as-is. + + Args: + nbytes: Number of random bytes before base64 encoding. The + resulting string is ``ceil(nbytes * 4 / 3)`` characters + (43 chars for the default 32 bytes). Defaults to 32. + + Returns: + A ready-to-use token string (e.g. ``"user_Xk3..."``). + + Example:: + + bearer = BearerTokenAuth(verify_token, prefix="user_") + token = bearer.generate_token() # "user_" + await db.store_token(user_id, token) + return {"access_token": token, "token_type": "bearer"} + """ + token = secrets.token_urlsafe(nbytes) + if self._prefix is not None: + return f"{self._prefix}{token}" + return token diff --git a/src/fastapi_toolsets/security/sources/cookie.py b/src/fastapi_toolsets/security/sources/cookie.py new file mode 100644 index 0000000..37eb9c2 --- /dev/null +++ b/src/fastapi_toolsets/security/sources/cookie.py @@ -0,0 +1,94 @@ +"""Cookie-based authentication source.""" + +import inspect +from typing import Annotated, Any, Callable + +from fastapi import Depends, Request +from fastapi.security import APIKeyCookie, SecurityScopes + +from fastapi_toolsets.exceptions import UnauthorizedError + +from ..base import AuthSource + + +class CookieAuth(AuthSource): + """Cookie-based authentication source. + + Wraps :class:`fastapi.security.APIKeyCookie` for OpenAPI documentation. + + Args: + name: Cookie name to read the credential from. + validator: Async callable that receives the cookie value and any extra + keyword arguments, and returns the authenticated identity. + **kwargs: Extra keyword arguments forwarded to the validator on every + call. + + Example:: + + async def verify_session(session_id: str) -> User: + user = await db.get_by_session(session_id) + if not user: + raise UnauthorizedError() + return user + + cookie_auth = CookieAuth("session", verify_session) + + @app.get("/me") + async def me(user: User = Security(cookie_auth)): + return user + """ + + def __init__( + self, + name: str, + validator: Callable[..., Any], + **kwargs: Any, + ) -> None: + self._name = name + self._validator = validator + self._kwargs = kwargs + self._scheme = APIKeyCookie(name=name, auto_error=False) + + _scheme = self._scheme + _validator = validator + _kwargs = kwargs + + async def _call( + security_scopes: SecurityScopes, # noqa: ARG001 + value: Annotated[str | None, Depends(_scheme)] = None, + ) -> Any: + if value is None: + raise UnauthorizedError() + return await _validator(value, **_kwargs) + + self._call_fn = _call + self.__signature__ = inspect.signature(_call) + + async def __call__(self, **kwargs: Any) -> Any: + return await self._call_fn(**kwargs) + + async def extract(self, request: Request) -> str | None: + """Extract the cookie value from the request without validating.""" + return request.cookies.get(self._name) + + async def authenticate(self, credential: str) -> Any: + """Validate a credential and return the identity.""" + return await self._validator(credential, **self._kwargs) + + def require(self, **kwargs: Any) -> "CookieAuth": + """Return a new instance with additional (or overriding) validator kwargs. + + The cookie name is preserved. New kwargs are merged over existing ones + (new values win on conflict):: + + cookie = CookieAuth("session", verify_session) + + @app.get("/admin") + async def admin(user: User = Security(cookie.require(role=Role.ADMIN))): + return user + """ + return CookieAuth( + self._name, + self._validator, + **{**self._kwargs, **kwargs}, + ) diff --git a/src/fastapi_toolsets/security/sources/oauth2.py b/src/fastapi_toolsets/security/sources/oauth2.py new file mode 100644 index 0000000..e515747 --- /dev/null +++ b/src/fastapi_toolsets/security/sources/oauth2.py @@ -0,0 +1,96 @@ +"""OAuth2 password-bearer authentication source.""" + +import inspect +from typing import Annotated, Any, Callable + +from fastapi import Depends, Request +from fastapi.security import OAuth2PasswordBearer, SecurityScopes + +from fastapi_toolsets.exceptions import UnauthorizedError + +from ..base import AuthSource + + +class OAuth2Auth(AuthSource): + """OAuth2 password-bearer authentication source. + + Wraps :class:`fastapi.security.OAuth2PasswordBearer` for OpenAPI + documentation. + + Args: + token_url: URL of the token endpoint (used in OpenAPI docs). + validator: Async callable that receives the token and any extra keyword + arguments, and returns the authenticated identity. + **kwargs: Extra keyword arguments forwarded to the validator on every + call. + + Example:: + + async def verify_token(token: str) -> User: + ... + + oauth2_auth = OAuth2Auth(token_url="/token", validator=verify_token) + + @app.get("/me") + async def me(user: User = Security(oauth2_auth)): + return user + """ + + def __init__( + self, + token_url: str, + validator: Callable[..., Any], + **kwargs: Any, + ) -> None: + self._token_url = token_url + self._validator = validator + self._kwargs = kwargs + self._scheme = OAuth2PasswordBearer(tokenUrl=token_url, auto_error=False) + + _scheme = self._scheme + _validator = validator + _kwargs = kwargs + + async def _call( + security_scopes: SecurityScopes, # noqa: ARG001 + token: Annotated[str | None, Depends(_scheme)] = None, + ) -> Any: + if token is None: + raise UnauthorizedError() + return await _validator(token, **_kwargs) + + self._call_fn = _call + self.__signature__ = inspect.signature(_call) + + async def __call__(self, **kwargs: Any) -> Any: + return await self._call_fn(**kwargs) + + async def extract(self, request: Request) -> str | None: + """Extract the bearer token from the Authorization header.""" + auth = request.headers.get("Authorization", "") + if not auth.startswith("Bearer "): + return None + token = auth[7:] + return token or None + + async def authenticate(self, credential: str) -> Any: + """Validate a credential and return the identity.""" + return await self._validator(credential, **self._kwargs) + + def require(self, **kwargs: Any) -> "OAuth2Auth": + """Return a new instance with additional (or overriding) validator kwargs. + + The token URL is preserved. New kwargs are merged over existing ones + (new values win on conflict):: + + oauth2 = OAuth2Auth("/token", verify_token) + + @app.get("/admin") + async def admin(user: User = Security(oauth2.require(role=Role.ADMIN))): + return user + """ + return OAuth2Auth( + self._token_url, + self._validator, + **{**self._kwargs, **kwargs}, + ) diff --git a/src/fastapi_toolsets/security/sources/openid.py b/src/fastapi_toolsets/security/sources/openid.py new file mode 100644 index 0000000..e12d887 --- /dev/null +++ b/src/fastapi_toolsets/security/sources/openid.py @@ -0,0 +1,127 @@ +"""OpenID Connect authentication source.""" + +import inspect +from typing import Annotated, Any, Callable + +from fastapi import Depends, Request +from fastapi.security import OpenIdConnect, SecurityScopes + +from fastapi_toolsets.exceptions import UnauthorizedError + +from ..base import AuthSource + + +class OpenIDAuth(AuthSource): + """OpenID Connect authentication source. + + Wraps :class:`fastapi.security.OpenIdConnect` for OpenAPI documentation. + Token extraction reads the ``Authorization: Bearer `` header; + validation is fully delegated to the user-supplied validator (use any + OIDC / JWT library such as ``authlib``, ``python-jose``, or ``PyJWT``). + + Args: + openid_connect_url: URL of the OIDC discovery document + (``/.well-known/openid-configuration``). Used only for OpenAPI + documentation — no requests are made to this URL by this class. + validator: Async callable that receives the raw bearer token and any + extra keyword arguments, and returns the authenticated identity. + Should raise :class:`~fastapi_toolsets.exceptions.UnauthorizedError` + on failure. + **kwargs: Extra keyword arguments forwarded to the validator on every + call (e.g. ``audience="my-app"``). + + Example — Google:: + + import jwt # e.g. PyJWT or python-jose + + async def verify_google_token(token: str, *, audience: str) -> User: + payload = jwt.decode(token, google_public_keys, algorithms=["RS256"], + audience=audience) + return User(email=payload["email"], name=payload["name"]) + + google_auth = OpenIDAuth( + "https://accounts.google.com/.well-known/openid-configuration", + verify_google_token, + audience="my-client-id", + ) + + @app.get("/me") + async def me(user: User = Security(google_auth)): + return user + + Multiple providers with :func:`~fastapi_toolsets.security.MultiAuth`:: + + multi = MultiAuth(google_auth, github_auth) + + @app.get("/data") + async def data(user: User = Security(multi)): + return user + """ + + def __init__( + self, + openid_connect_url: str, + validator: Callable[..., Any], + **kwargs: Any, + ) -> None: + self._openid_connect_url = openid_connect_url + self._validator = validator + self._kwargs = kwargs + self._scheme = OpenIdConnect( + openIdConnectUrl=openid_connect_url, auto_error=False + ) + + _scheme = self._scheme + _validator = validator + _kwargs = kwargs + + async def _call( + security_scopes: SecurityScopes, # noqa: ARG001 + # OpenIdConnect (OAuth2 base) returns the full Authorization header + # value (e.g. "Bearer "), unlike OAuth2PasswordBearer which + # strips the scheme prefix. + authorization: Annotated[str | None, Depends(_scheme)] = None, + ) -> Any: + if authorization is None: + raise UnauthorizedError() + if not authorization.startswith("Bearer "): + raise UnauthorizedError() + token = authorization[7:] + if not token: + raise UnauthorizedError() + return await _validator(token, **_kwargs) + + self._call_fn = _call + self.__signature__ = inspect.signature(_call) + + async def __call__(self, **kwargs: Any) -> Any: + return await self._call_fn(**kwargs) + + async def extract(self, request: Request) -> str | None: + """Extract the bearer token from the Authorization header.""" + auth = request.headers.get("Authorization", "") + if not auth.startswith("Bearer "): + return None + return auth[7:] or None + + async def authenticate(self, credential: str) -> Any: + """Validate a credential and return the identity.""" + return await self._validator(credential, **self._kwargs) + + def require(self, **kwargs: Any) -> "OpenIDAuth": + """Return a new instance with additional (or overriding) validator kwargs. + + The discovery URL is preserved. New kwargs are merged over existing ones + (new values win on conflict):: + + google_auth = OpenIDAuth(discovery_url, verify_google_token, audience="app") + + @app.get("/admin") + async def admin(user: User = Security(google_auth.require(role=Role.ADMIN))): + return user + """ + return OpenIDAuth( + self._openid_connect_url, + self._validator, + **{**self._kwargs, **kwargs}, + ) diff --git a/tests/test_security.py b/tests/test_security.py new file mode 100644 index 0000000..2d3eac0 --- /dev/null +++ b/tests/test_security.py @@ -0,0 +1,902 @@ +"""Tests for fastapi_toolsets.security.""" + +import pytest +from fastapi import FastAPI, Security +from fastapi.testclient import TestClient + +from fastapi_toolsets.exceptions import UnauthorizedError, init_exceptions_handlers +from fastapi_toolsets.security import ( + AuthSource, + BearerTokenAuth, + CookieAuth, + MultiAuth, + OAuth2Auth, + OpenIDAuth, +) + + +def _app(*routes_setup_fns): + """Build a minimal FastAPI test app with exception handlers.""" + app = FastAPI() + init_exceptions_handlers(app) + for fn in routes_setup_fns: + fn(app) + return app + + +VALID_TOKEN = "secret" +VALID_COOKIE = "session123" + + +async def simple_validator(credential: str) -> dict: + if credential != VALID_TOKEN: + raise UnauthorizedError() + return {"user": "alice"} + + +async def role_validator(credential: str, *, role: str) -> dict: + if credential != VALID_TOKEN: + raise UnauthorizedError() + return {"user": "alice", "role": role} + + +async def cookie_validator(value: str) -> dict: + if value != VALID_COOKIE: + raise UnauthorizedError() + return {"session": value} + + +class TestBearerTokenAuth: + def test_valid_token_returns_identity(self): + bearer = BearerTokenAuth(simple_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(bearer)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice"} + + def test_missing_header_returns_401(self): + bearer = BearerTokenAuth(simple_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(bearer)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me") + assert response.status_code == 401 + + def test_invalid_token_returns_401(self): + bearer = BearerTokenAuth(simple_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(bearer)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": "Bearer wrong"}) + assert response.status_code == 401 + + def test_kwargs_forwarded_to_validator(self): + bearer = BearerTokenAuth(role_validator, role="admin") + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(bearer)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice", "role": "admin"} + + def test_prefix_matching_passes_full_token(self): + """Token with matching prefix: full token (with prefix) is passed to validator.""" + received: list[str] = [] + + async def capturing_validator(credential: str) -> dict: + received.append(credential) + return {"user": "alice"} + + bearer = BearerTokenAuth(capturing_validator, prefix="user_") + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(bearer)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": "Bearer user_abc123"}) + assert response.status_code == 200 + # Prefix is kept — validator receives the full token as stored in DB + assert received == ["user_abc123"] + + def test_prefix_mismatch_returns_401(self): + bearer = BearerTokenAuth(simple_validator, prefix="user_") + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(bearer)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": "Bearer org_abc123"}) + assert response.status_code == 401 + + # --- extract() --- + + @pytest.mark.anyio + async def test_extract_no_header(self): + from starlette.requests import Request + + bearer = BearerTokenAuth(simple_validator) + scope = {"type": "http", "method": "GET", "path": "/", "headers": []} + request = Request(scope) + assert await bearer.extract(request) is None + + @pytest.mark.anyio + async def test_extract_empty_token(self): + from starlette.requests import Request + + bearer = BearerTokenAuth(simple_validator) + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"authorization", b"Bearer ")], + } + request = Request(scope) + assert await bearer.extract(request) is None + + @pytest.mark.anyio + async def test_extract_no_prefix(self): + from starlette.requests import Request + + bearer = BearerTokenAuth(simple_validator) + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"authorization", b"Bearer mytoken")], + } + request = Request(scope) + assert await bearer.extract(request) == "mytoken" + + @pytest.mark.anyio + async def test_extract_prefix_match(self): + from starlette.requests import Request + + bearer = BearerTokenAuth(simple_validator, prefix="user_") + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"authorization", b"Bearer user_abc")], + } + request = Request(scope) + assert await bearer.extract(request) == "user_abc" + + @pytest.mark.anyio + async def test_extract_prefix_no_match(self): + from starlette.requests import Request + + bearer = BearerTokenAuth(simple_validator, prefix="user_") + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"authorization", b"Bearer org_abc")], + } + request = Request(scope) + assert await bearer.extract(request) is None + + # --- generate_token() --- + + def test_generate_token_no_prefix(self): + bearer = BearerTokenAuth(simple_validator) + token = bearer.generate_token() + assert isinstance(token, str) + assert len(token) > 0 + + def test_generate_token_with_prefix(self): + bearer = BearerTokenAuth(simple_validator, prefix="user_") + token = bearer.generate_token() + assert token.startswith("user_") + + def test_generate_token_uniqueness(self): + bearer = BearerTokenAuth(simple_validator) + assert bearer.generate_token() != bearer.generate_token() + + def test_generate_token_is_valid_credential(self): + """A generated token (with prefix) is accepted by the same auth source.""" + stored: list[str] = [] + + async def storing_validator(credential: str) -> dict: + stored.append(credential) + return {"token": credential} + + bearer = BearerTokenAuth(storing_validator, prefix="user_") + token = bearer.generate_token() + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(bearer)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {token}"}) + assert response.status_code == 200 + assert stored == [token] + + +class TestCookieAuth: + def test_valid_cookie_returns_identity(self): + cookie_auth = CookieAuth("session", cookie_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(cookie_auth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", cookies={"session": VALID_COOKIE}) + assert response.status_code == 200 + assert response.json() == {"session": VALID_COOKIE} + + def test_missing_cookie_returns_401(self): + cookie_auth = CookieAuth("session", cookie_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(cookie_auth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me") + assert response.status_code == 401 + + def test_invalid_cookie_returns_401(self): + cookie_auth = CookieAuth("session", cookie_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(cookie_auth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", cookies={"session": "wrong"}) + assert response.status_code == 401 + + def test_kwargs_forwarded_to_validator(self): + async def session_validator(value: str, *, scope: str) -> dict: + if value != VALID_COOKIE: + raise UnauthorizedError() + return {"session": value, "scope": scope} + + cookie_auth = CookieAuth("session", session_validator, scope="read") + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(cookie_auth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", cookies={"session": VALID_COOKIE}) + assert response.status_code == 200 + assert response.json() == {"session": VALID_COOKIE, "scope": "read"} + + @pytest.mark.anyio + async def test_extract_no_cookie(self): + from starlette.requests import Request + + auth = CookieAuth("session", cookie_validator) + scope = {"type": "http", "method": "GET", "path": "/", "headers": []} + request = Request(scope) + assert await auth.extract(request) is None + + @pytest.mark.anyio + async def test_extract_cookie_present(self): + from starlette.requests import Request + + auth = CookieAuth("session", cookie_validator) + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"cookie", b"session=abc")], + } + request = Request(scope) + assert await auth.extract(request) == "abc" + + +class TestOAuth2Auth: + def test_valid_token_returns_identity(self): + oauth = OAuth2Auth(token_url="/token", validator=simple_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(oauth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice"} + + def test_missing_token_returns_401(self): + oauth = OAuth2Auth(token_url="/token", validator=simple_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(oauth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me") + assert response.status_code == 401 + + @pytest.mark.anyio + async def test_extract_no_header(self): + from starlette.requests import Request + + auth = OAuth2Auth("/token", simple_validator) + scope = {"type": "http", "method": "GET", "path": "/", "headers": []} + request = Request(scope) + assert await auth.extract(request) is None + + @pytest.mark.anyio + async def test_extract_empty_token(self): + from starlette.requests import Request + + auth = OAuth2Auth("/token", simple_validator) + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"authorization", b"Bearer ")], + } + request = Request(scope) + assert await auth.extract(request) is None + + @pytest.mark.anyio + async def test_extract_token(self): + from starlette.requests import Request + + auth = OAuth2Auth("/token", simple_validator) + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"authorization", b"Bearer mytoken")], + } + request = Request(scope) + assert await auth.extract(request) == "mytoken" + + def test_in_multi_auth(self): + """OAuth2Auth.authenticate() is exercised when used inside MultiAuth.""" + oauth = OAuth2Auth(token_url="/token", validator=simple_validator) + multi = MultiAuth(oauth) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice"} + + +class TestMultiAuth: + def test_first_source_matches(self): + bearer = BearerTokenAuth(simple_validator) + cookie = CookieAuth("session", cookie_validator) + multi = MultiAuth(bearer, cookie) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice"} + + def test_second_source_matches_when_first_absent(self): + bearer = BearerTokenAuth(simple_validator) + cookie = CookieAuth("session", cookie_validator) + multi = MultiAuth(bearer, cookie) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + # No Authorization header — falls through to cookie + response = client.get("/me", cookies={"session": VALID_COOKIE}) + assert response.status_code == 200 + assert response.json() == {"session": VALID_COOKIE} + + def test_no_source_matches_returns_401(self): + bearer = BearerTokenAuth(simple_validator) + cookie = CookieAuth("session", cookie_validator) + multi = MultiAuth(bearer, cookie) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me") + assert response.status_code == 401 + + def test_invalid_credential_does_not_fallthrough(self): + """If a credential is found but invalid, the next source is NOT tried.""" + second_called: list[bool] = [] + + async def tracking_validator(credential: str) -> dict: + second_called.append(True) + return {"from": "second"} + + bearer = BearerTokenAuth(simple_validator) # raises on wrong token + cookie = CookieAuth("session", tracking_validator) + multi = MultiAuth(bearer, cookie) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + # Bearer credential present but wrong — should NOT try cookie + response = client.get( + "/me", + headers={"Authorization": "Bearer wrong"}, + cookies={"session": VALID_COOKIE}, + ) + assert response.status_code == 401 + assert second_called == [] # cookie validator was never called + + def test_prefix_routes_to_correct_source(self): + """Prefix-based dispatch: only the matching source's validator is called.""" + user_calls: list[str] = [] + org_calls: list[str] = [] + + async def user_validator(credential: str) -> dict: + user_calls.append(credential) + return {"type": "user", "id": credential} + + async def org_validator(credential: str) -> dict: + org_calls.append(credential) + return {"type": "org", "id": credential} + + user_bearer = BearerTokenAuth(user_validator, prefix="user_") + org_bearer = BearerTokenAuth(org_validator, prefix="org_") + multi = MultiAuth(user_bearer, org_bearer) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + + response = client.get("/me", headers={"Authorization": "Bearer user_alice"}) + assert response.status_code == 200 + assert response.json() == {"type": "user", "id": "user_alice"} + assert user_calls == ["user_alice"] + assert org_calls == [] + + user_calls.clear() + + response = client.get("/me", headers={"Authorization": "Bearer org_acme"}) + assert response.status_code == 200 + assert response.json() == {"type": "org", "id": "org_acme"} + assert user_calls == [] + assert org_calls == ["org_acme"] + + def test_require_returns_new_multi_auth(self): + from fastapi_toolsets.security.multi import MultiAuth as MultiAuthClass + + bearer = BearerTokenAuth(role_validator) + multi = MultiAuth(bearer) + derived = multi.require(role="admin") + assert isinstance(derived, MultiAuthClass) + assert derived is not multi + + def test_require_forwards_kwargs_to_sources(self): + """multi.require() propagates to all sources that support it.""" + bearer = BearerTokenAuth(role_validator) + multi = MultiAuth(bearer) + + def setup(app: FastAPI): + @app.get("/admin") + async def admin(user=Security(multi.require(role="admin"))): + return user + + client = TestClient(_app(setup)) + response = client.get( + "/admin", headers={"Authorization": f"Bearer {VALID_TOKEN}"} + ) + assert response.status_code == 200 + assert response.json() == {"user": "alice", "role": "admin"} + + def test_require_skips_sources_without_require(self): + """Sources without require() are passed through unchanged.""" + header_auth = _HeaderAuth(secret="s3cr3t") + multi = MultiAuth(header_auth) + derived = multi.require(role="admin") + assert derived._sources[0] is header_auth + + def test_require_does_not_mutate_original(self): + bearer = BearerTokenAuth(role_validator, role="user") + multi = MultiAuth(bearer) + multi.require(role="admin") + assert bearer._kwargs == {"role": "user"} + + def test_require_mixed_sources(self): + """require() applies to sources with require(), skips those without.""" + from typing import cast + + bearer = BearerTokenAuth(role_validator) + header_auth = _HeaderAuth(secret="s3cr3t") + multi = MultiAuth(bearer, header_auth) + derived = multi.require(role="admin") + # bearer got require() applied, header_auth passed through + assert cast(BearerTokenAuth, derived._sources[0])._kwargs == {"role": "admin"} + assert derived._sources[1] is header_auth + + +class TestRequire: + def test_bearer_require_forwards_kwargs(self): + """require() creates a new instance that passes merged kwargs to validator.""" + bearer = BearerTokenAuth(role_validator) + + def setup(app: FastAPI): + @app.get("/admin") + async def admin(user=Security(bearer.require(role="admin"))): + return user + + client = TestClient(_app(setup)) + response = client.get( + "/admin", headers={"Authorization": f"Bearer {VALID_TOKEN}"} + ) + assert response.status_code == 200 + assert response.json() == {"user": "alice", "role": "admin"} + + def test_bearer_require_overrides_existing_kwarg(self): + """require() kwargs override kwargs set at instantiation.""" + bearer = BearerTokenAuth(role_validator, role="user") + + def setup(app: FastAPI): + @app.get("/admin") + async def admin(user=Security(bearer.require(role="admin"))): + return user + + client = TestClient(_app(setup)) + response = client.get( + "/admin", headers={"Authorization": f"Bearer {VALID_TOKEN}"} + ) + assert response.status_code == 200 + assert response.json()["role"] == "admin" + + def test_bearer_require_preserves_prefix(self): + """require() keeps the prefix of the original instance.""" + bearer = BearerTokenAuth(role_validator, prefix="user_") + derived = bearer.require(role="admin") + assert derived._prefix == "user_" + + def test_bearer_require_does_not_mutate_original(self): + """require() returns a new instance — original kwargs are unchanged.""" + bearer = BearerTokenAuth(role_validator, role="user") + bearer.require(role="admin") + assert bearer._kwargs == {"role": "user"} + + def test_cookie_require_forwards_kwargs(self): + async def scoped_validator(value: str, *, scope: str) -> dict: + if value != VALID_COOKIE: + raise UnauthorizedError() + return {"session": value, "scope": scope} + + cookie = CookieAuth("session", scoped_validator) + + def setup(app: FastAPI): + @app.get("/admin") + async def admin(user=Security(cookie.require(scope="admin"))): + return user + + client = TestClient(_app(setup)) + response = client.get("/admin", cookies={"session": VALID_COOKIE}) + assert response.status_code == 200 + assert response.json() == {"session": VALID_COOKIE, "scope": "admin"} + + def test_cookie_require_preserves_name(self): + cookie = CookieAuth("session", cookie_validator) + derived = cookie.require(scope="admin") + assert derived._name == "session" + + def test_oauth2_require_forwards_kwargs(self): + oauth = OAuth2Auth("/token", role_validator) + + def setup(app: FastAPI): + @app.get("/admin") + async def admin(user=Security(oauth.require(role="admin"))): + return user + + client = TestClient(_app(setup)) + response = client.get( + "/admin", headers={"Authorization": f"Bearer {VALID_TOKEN}"} + ) + assert response.status_code == 200 + assert response.json() == {"user": "alice", "role": "admin"} + + def test_oauth2_require_preserves_token_url(self): + oauth = OAuth2Auth("/token", simple_validator) + derived = oauth.require(role="admin") + assert derived._token_url == "/token" + + def test_bearer_require_in_multi_auth(self): + """require() instances work seamlessly inside MultiAuth.""" + PREFIXED_TOKEN = f"user_{VALID_TOKEN}" + + async def prefixed_role_validator(credential: str, *, role: str) -> dict: + if credential != PREFIXED_TOKEN: + raise UnauthorizedError() + return {"user": "alice", "role": role} + + bearer = BearerTokenAuth(prefixed_role_validator, prefix="user_") + multi = MultiAuth(bearer.require(role="admin")) + + def setup(app: FastAPI): + @app.get("/admin") + async def admin(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + response = client.get( + "/admin", headers={"Authorization": f"Bearer {PREFIXED_TOKEN}"} + ) + assert response.status_code == 200 + assert response.json() == {"user": "alice", "role": "admin"} + + +class TestOpenIDAuth: + DISCOVERY_URL = "https://accounts.example.com/.well-known/openid-configuration" + + def test_valid_token_returns_identity(self): + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(oidc)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice"} + + def test_missing_token_returns_401(self): + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(oidc)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me") + assert response.status_code == 401 + + def test_invalid_token_returns_401(self): + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(oidc)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": "Bearer wrong"}) + assert response.status_code == 401 + + def test_kwargs_forwarded_to_validator(self): + oidc = OpenIDAuth(self.DISCOVERY_URL, role_validator, role="admin") + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(oidc)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice", "role": "admin"} + + def test_require_forwards_kwargs(self): + oidc = OpenIDAuth(self.DISCOVERY_URL, role_validator) + + def setup(app: FastAPI): + @app.get("/admin") + async def admin(user=Security(oidc.require(role="admin"))): + return user + + client = TestClient(_app(setup)) + response = client.get( + "/admin", headers={"Authorization": f"Bearer {VALID_TOKEN}"} + ) + assert response.status_code == 200 + assert response.json() == {"user": "alice", "role": "admin"} + + def test_require_preserves_discovery_url(self): + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + derived = oidc.require(role="admin") + assert derived._openid_connect_url == self.DISCOVERY_URL + + def test_require_does_not_mutate_original(self): + oidc = OpenIDAuth(self.DISCOVERY_URL, role_validator, role="user") + oidc.require(role="admin") + assert oidc._kwargs == {"role": "user"} + + def test_is_auth_source(self): + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + assert isinstance(oidc, AuthSource) + + def test_in_multi_auth(self): + """OpenIDAuth works seamlessly inside MultiAuth.""" + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + multi = MultiAuth(oidc) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice"} + + @pytest.mark.anyio + async def test_extract_no_header(self): + from starlette.requests import Request + + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + scope = {"type": "http", "method": "GET", "path": "/", "headers": []} + request = Request(scope) + assert await oidc.extract(request) is None + + @pytest.mark.anyio + async def test_extract_empty_token(self): + from starlette.requests import Request + + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"authorization", b"Bearer ")], + } + request = Request(scope) + assert await oidc.extract(request) is None + + @pytest.mark.anyio + async def test_extract_token(self): + from starlette.requests import Request + + oidc = OpenIDAuth(self.DISCOVERY_URL, simple_validator) + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [(b"authorization", b"Bearer mytoken")], + } + request = Request(scope) + assert await oidc.extract(request) == "mytoken" + + +# Minimal concrete subclass used only in tests below. +class _HeaderAuth(AuthSource): + """Reads a custom X-Token header — no FastAPI security scheme.""" + + def __init__(self, secret: str) -> None: + super().__init__() + self._secret = secret + + async def extract(self, request) -> str | None: + return request.headers.get("X-Token") or None + + async def authenticate(self, credential: str) -> dict: + if credential != self._secret: + raise UnauthorizedError() + return {"token": credential} + + +class TestAuthSource: + def test_cannot_instantiate_abstract_class(self): + with pytest.raises(TypeError): + AuthSource() + + def test_builtin_classes_are_auth_sources(self): + bearer = BearerTokenAuth(simple_validator) + cookie = CookieAuth("session", cookie_validator) + oauth = OAuth2Auth("/token", simple_validator) + oidc = OpenIDAuth( + "https://example.com/.well-known/openid-configuration", simple_validator + ) + assert isinstance(bearer, AuthSource) + assert isinstance(cookie, AuthSource) + assert isinstance(oauth, AuthSource) + assert isinstance(oidc, AuthSource) + + def test_custom_source_standalone_valid(self): + """Default __call__ wires extract + authenticate via Request injection.""" + auth = _HeaderAuth(secret="s3cr3t") + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(auth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"X-Token": "s3cr3t"}) + assert response.status_code == 200 + assert response.json() == {"token": "s3cr3t"} + + def test_custom_source_standalone_missing_credential(self): + auth = _HeaderAuth(secret="s3cr3t") + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(auth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me") # no X-Token header + assert response.status_code == 401 + + def test_custom_source_standalone_invalid_credential(self): + auth = _HeaderAuth(secret="s3cr3t") + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(auth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", headers={"X-Token": "wrong"}) + assert response.status_code == 401 + + def test_custom_source_in_multi_auth(self): + """Custom AuthSource works transparently inside MultiAuth.""" + header_auth = _HeaderAuth(secret="s3cr3t") + bearer = BearerTokenAuth(simple_validator) + multi = MultiAuth(bearer, header_auth) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(multi)): + return user + + client = TestClient(_app(setup)) + + # Bearer matches first + response = client.get("/me", headers={"Authorization": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice"} + + # No bearer → falls through to custom header source + response = client.get("/me", headers={"X-Token": "s3cr3t"}) + assert response.status_code == 200 + assert response.json() == {"token": "s3cr3t"}