mirror of
https://github.com/d3vyce/fastapi-toolsets.git
synced 2026-08-13 19:42:59 +00:00
wip2
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user