diff --git a/src/fastapi_toolsets/security/base.py b/src/fastapi_toolsets/security/base.py index 7d82db7..6640bf8 100644 --- a/src/fastapi_toolsets/security/base.py +++ b/src/fastapi_toolsets/security/base.py @@ -10,66 +10,20 @@ from fastapi.security import SecurityScopes from fastapi_toolsets.exceptions import UnauthorizedError +async def _call_validator( + validator: Callable[..., Any], *args: Any, **kwargs: Any +) -> Any: + """Call *validator* with *args* and *kwargs*, awaiting it if it is a coroutine function.""" + if inspect.iscoroutinefunction(validator): + return await validator(*args, **kwargs) + return validator(*args, **kwargs) + + 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. - """ + """Abstract base class for authentication sources.""" 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. - """ + """Set up the default FastAPI dependency signature.""" source = self async def _call( @@ -86,27 +40,12 @@ class AuthSource(ABC): @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. - """ + """Extract the raw credential from the request without validating.""" @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. - """ + """Validate a credential and return the authenticated identity.""" 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. - """ + """FastAPI dependency dispatch.""" return await self._call_fn(**kwargs) diff --git a/src/fastapi_toolsets/security/sources/bearer.py b/src/fastapi_toolsets/security/sources/bearer.py index 4db0c49..c486e1c 100644 --- a/src/fastapi_toolsets/security/sources/bearer.py +++ b/src/fastapi_toolsets/security/sources/bearer.py @@ -9,7 +9,7 @@ from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer, SecurityS from fastapi_toolsets.exceptions import UnauthorizedError -from ..base import AuthSource +from ..base import AuthSource, _call_validator class BearerTokenAuth(AuthSource): @@ -20,9 +20,9 @@ class BearerTokenAuth(AuthSource): 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 + validator: Sync or 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 @@ -63,14 +63,11 @@ class BearerTokenAuth(AuthSource): token = credentials.credentials if _prefix is not None and not token.startswith(_prefix): raise UnauthorizedError() - return await _validator(token, **_kwargs) + return await _call_validator(_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: Any) -> str | None: """Extract the raw credential from the request without validating. @@ -94,7 +91,7 @@ class BearerTokenAuth(AuthSource): Calls ``await validator(credential, **kwargs)`` where ``kwargs`` are the extra keyword arguments provided at instantiation. """ - return await self._validator(credential, **self._kwargs) + return await _call_validator(self._validator, credential, **self._kwargs) def require(self, **kwargs: Any) -> "BearerTokenAuth": """Return a new instance with additional (or overriding) validator kwargs.""" diff --git a/src/fastapi_toolsets/security/sources/cookie.py b/src/fastapi_toolsets/security/sources/cookie.py index d01d306..573c58a 100644 --- a/src/fastapi_toolsets/security/sources/cookie.py +++ b/src/fastapi_toolsets/security/sources/cookie.py @@ -13,7 +13,7 @@ from fastapi.security import APIKeyCookie, SecurityScopes from fastapi_toolsets.exceptions import UnauthorizedError -from ..base import AuthSource +from ..base import AuthSource, _call_validator class CookieAuth(AuthSource): @@ -25,9 +25,10 @@ class CookieAuth(AuthSource): Args: name: Cookie name. - validator: Async callable that receives the cookie value (plain, after - signature verification when ``secret_key`` is set) and any extra - keyword arguments, and returns the authenticated identity. + validator: Sync or async callable that receives the cookie value + (plain, after signature verification when ``secret_key`` is set) + and any extra keyword arguments, and returns the authenticated + identity. secret_key: When provided, the cookie is HMAC-SHA256 signed. :meth:`set_cookie` embeds an expiry and signs the payload; :meth:`extract` verifies the signature and expiry before handing @@ -66,14 +67,11 @@ class CookieAuth(AuthSource): if value is None: raise UnauthorizedError() plain = _self._verify(value) - return await _self._validator(plain, **_kwargs) + return await _call_validator(_self._validator, plain, **_kwargs) self._call_fn = _call self.__signature__ = inspect.signature(_call) - async def __call__(self, **kwargs: Any) -> Any: - return await self._call_fn(**kwargs) - def _hmac(self, data: str) -> str: assert self._secret_key is not None return hmac.new( @@ -116,7 +114,7 @@ class CookieAuth(AuthSource): async def authenticate(self, credential: str) -> Any: plain = self._verify(credential) - return await self._validator(plain, **self._kwargs) + return await _call_validator(self._validator, plain, **self._kwargs) def require(self, **kwargs: Any) -> "CookieAuth": """Return a new instance with additional (or overriding) validator kwargs.""" diff --git a/src/fastapi_toolsets/security/sources/oauth2.py b/src/fastapi_toolsets/security/sources/oauth2.py index c11be0c..8759062 100644 --- a/src/fastapi_toolsets/security/sources/oauth2.py +++ b/src/fastapi_toolsets/security/sources/oauth2.py @@ -8,7 +8,7 @@ from fastapi.security import OAuth2PasswordBearer, SecurityScopes from fastapi_toolsets.exceptions import UnauthorizedError -from ..base import AuthSource +from ..base import AuthSource, _call_validator class OAuth2Auth(AuthSource): @@ -19,8 +19,8 @@ class OAuth2Auth(AuthSource): 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. + validator: Sync or 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. """ @@ -46,14 +46,11 @@ class OAuth2Auth(AuthSource): ) -> Any: if token is None: raise UnauthorizedError() - return await _validator(token, **_kwargs) + return await _call_validator(_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", "") @@ -64,7 +61,7 @@ class OAuth2Auth(AuthSource): async def authenticate(self, credential: str) -> Any: """Validate a credential and return the identity.""" - return await self._validator(credential, **self._kwargs) + return await _call_validator(self._validator, credential, **self._kwargs) def require(self, **kwargs: Any) -> "OAuth2Auth": """Return a new instance with additional (or overriding) validator kwargs.""" diff --git a/src/fastapi_toolsets/security/sources/openid.py b/src/fastapi_toolsets/security/sources/openid.py index 6ea44cd..9df66d7 100644 --- a/src/fastapi_toolsets/security/sources/openid.py +++ b/src/fastapi_toolsets/security/sources/openid.py @@ -8,7 +8,7 @@ from fastapi.security import OpenIdConnect, SecurityScopes from fastapi_toolsets.exceptions import UnauthorizedError -from ..base import AuthSource +from ..base import AuthSource, _call_validator class OpenIDAuth(AuthSource): @@ -23,10 +23,10 @@ class OpenIDAuth(AuthSource): 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. + validator: Sync or 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"``). """ @@ -59,14 +59,11 @@ class OpenIDAuth(AuthSource): token = authorization[7:] if not token: raise UnauthorizedError() - return await _validator(token, **_kwargs) + return await _call_validator(_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", "") @@ -76,7 +73,7 @@ class OpenIDAuth(AuthSource): async def authenticate(self, credential: str) -> Any: """Validate a credential and return the identity.""" - return await self._validator(credential, **self._kwargs) + return await _call_validator(self._validator, credential, **self._kwargs) def require(self, **kwargs: Any) -> "OpenIDAuth": """Return a new instance with additional (or overriding) validator kwargs.""" diff --git a/tests/test_security.py b/tests/test_security.py index 2d3eac0..9ad13eb 100644 --- a/tests/test_security.py +++ b/tests/test_security.py @@ -806,6 +806,204 @@ class TestOpenIDAuth: assert await oidc.extract(request) == "mytoken" +class TestSyncValidators: + """Sync (non-async) validators — covers the sync path in _call_validator.""" + + def test_bearer_sync_validator(self): + def sync_validator(credential: str) -> dict: + if credential != VALID_TOKEN: + raise UnauthorizedError() + return {"user": "alice"} + + bearer = BearerTokenAuth(sync_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_sync_validator_via_authenticate(self): + """authenticate() with sync validator (MultiAuth path).""" + + def sync_validator(credential: str) -> dict: + if credential != VALID_TOKEN: + raise UnauthorizedError() + return {"user": "alice"} + + bearer = BearerTokenAuth(sync_validator) + multi = MultiAuth(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": f"Bearer {VALID_TOKEN}"}) + assert response.status_code == 200 + assert response.json() == {"user": "alice"} + + +class TestOpenIDAuthEdgeCases: + DISCOVERY_URL = "https://accounts.example.com/.well-known/openid-configuration" + + def test_non_bearer_authorization_returns_401(self): + """Authorization header present but not Bearer scheme.""" + 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": "Basic dXNlcjpwYXNz"}) + assert response.status_code == 401 + + def test_empty_bearer_token_returns_401(self): + """Authorization: Bearer with no token after the scheme prefix.""" + 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 "}) + assert response.status_code == 401 + + +class TestCookieAuthSigned: + """CookieAuth with HMAC-SHA256 signed cookies (secret_key path).""" + + SECRET = "test-hmac-secret" + + def test_valid_signed_cookie_via_set_cookie(self): + """set_cookie signs the value; the signed cookie is verified on read.""" + from fastapi import Response + + auth = CookieAuth("session", cookie_validator, secret_key=self.SECRET) + + def setup(app: FastAPI): + @app.get("/login") + async def login(response: Response): + auth.set_cookie(response, VALID_COOKIE) + return {"ok": True} + + @app.get("/me") + async def me(user=Security(auth)): + return user + + with TestClient(_app(setup)) as client: + client.get("/login") + response = client.get("/me") + assert response.status_code == 200 + assert response.json() == {"session": VALID_COOKIE} + + def test_tampered_signature_returns_401(self): + """A cookie whose HMAC signature has been modified is rejected.""" + import base64 as _b64 + import json as _json + import time as _time + + auth = CookieAuth("session", cookie_validator, secret_key=self.SECRET) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(auth)): + return user + + client = TestClient(_app(setup)) + data = _b64.urlsafe_b64encode( + _json.dumps({"v": VALID_COOKIE, "exp": int(_time.time()) + 9999}).encode() + ).decode() + response = client.get("/me", cookies={"session": f"{data}.invalidsig"}) + assert response.status_code == 401 + + def test_expired_signed_cookie_returns_401(self): + """A signed cookie past its expiry timestamp is rejected.""" + import base64 as _b64 + import hashlib as _hashlib + import hmac as _hmac + import json as _json + import time as _time + + auth = CookieAuth("session", cookie_validator, secret_key=self.SECRET) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(auth)): + return user + + client = TestClient(_app(setup)) + data = _b64.urlsafe_b64encode( + _json.dumps({"v": VALID_COOKIE, "exp": int(_time.time()) - 1}).encode() + ).decode() + sig = _hmac.new( + self.SECRET.encode(), data.encode(), _hashlib.sha256 + ).hexdigest() + response = client.get("/me", cookies={"session": f"{data}.{sig}"}) + assert response.status_code == 401 + + def test_invalid_json_payload_returns_401(self): + """A signed cookie whose payload is not valid JSON is rejected.""" + import base64 as _b64 + import hashlib as _hashlib + import hmac as _hmac + + auth = CookieAuth("session", cookie_validator, secret_key=self.SECRET) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(auth)): + return user + + client = TestClient(_app(setup)) + data = _b64.urlsafe_b64encode(b"not-valid-json").decode() + sig = _hmac.new( + self.SECRET.encode(), data.encode(), _hashlib.sha256 + ).hexdigest() + response = client.get("/me", cookies={"session": f"{data}.{sig}"}) + assert response.status_code == 401 + + def test_malformed_cookie_no_dot_returns_401(self): + """A signed cookie without the dot separator is rejected.""" + auth = CookieAuth("session", cookie_validator, secret_key=self.SECRET) + + def setup(app: FastAPI): + @app.get("/me") + async def me(user=Security(auth)): + return user + + client = TestClient(_app(setup)) + response = client.get("/me", cookies={"session": "nodothere"}) + assert response.status_code == 401 + + def test_set_cookie_without_secret(self): + """set_cookie without secret_key writes the raw value.""" + from starlette.responses import Response as StarletteResponse + + auth = CookieAuth("session", cookie_validator) + response = StarletteResponse() + auth.set_cookie(response, "rawvalue") + assert "session=rawvalue" in response.headers["set-cookie"] + + def test_delete_cookie(self): + """delete_cookie produces a Set-Cookie header that clears the session.""" + from starlette.responses import Response as StarletteResponse + + auth = CookieAuth("session", cookie_validator) + response = StarletteResponse() + auth.delete_cookie(response) + assert "session" in response.headers["set-cookie"] + + # Minimal concrete subclass used only in tests below. class _HeaderAuth(AuthSource): """Reads a custom X-Token header — no FastAPI security scheme."""