mirror of
https://github.com/d3vyce/fastapi-toolsets.git
synced 2026-08-04 23:54:09 +00:00
feat: add security module (#111)
* feat: add security module * feat(security): add oauth helpers * docs: add authentication example * fix: cleanup + simplify * docs: update module and reference * fix: multiple security bugs + remove example for now * feat: use async_lru for caching * fix: rename nonce by state_token
This commit is contained in:
@@ -0,0 +1,197 @@
|
||||
"""OAuth 2.0 / OIDC helper utilities."""
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import hmac
|
||||
import json
|
||||
import secrets
|
||||
from typing import Any
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import httpx
|
||||
from async_lru import alru_cache
|
||||
from fastapi.responses import RedirectResponse
|
||||
|
||||
|
||||
@alru_cache(maxsize=32)
|
||||
async def oauth_resolve_provider_urls(
|
||||
discovery_url: str,
|
||||
) -> tuple[str, str, str | None]:
|
||||
"""Fetch the OIDC discovery document and return endpoint URLs.
|
||||
|
||||
Args:
|
||||
discovery_url: URL of the provider's ``/.well-known/openid-configuration``.
|
||||
|
||||
Returns:
|
||||
A ``(authorization_url, token_url, userinfo_url)`` tuple.
|
||||
*userinfo_url* is ``None`` when the provider does not advertise one.
|
||||
"""
|
||||
async with httpx.AsyncClient() as client:
|
||||
resp = await client.get(discovery_url)
|
||||
resp.raise_for_status()
|
||||
cfg = resp.json()
|
||||
return (
|
||||
cfg["authorization_endpoint"],
|
||||
cfg["token_endpoint"],
|
||||
cfg.get("userinfo_endpoint"),
|
||||
)
|
||||
|
||||
|
||||
async def oauth_fetch_userinfo(
|
||||
*,
|
||||
token_url: str,
|
||||
userinfo_url: str,
|
||||
code: str,
|
||||
client_id: str,
|
||||
client_secret: str,
|
||||
redirect_uri: str,
|
||||
required_scopes: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Exchange an authorization code for tokens and return the userinfo payload.
|
||||
|
||||
Args:
|
||||
token_url: Provider's token endpoint.
|
||||
userinfo_url: Provider's userinfo endpoint.
|
||||
code: Authorization code received from the provider's callback.
|
||||
client_id: OAuth application client ID.
|
||||
client_secret: OAuth application client secret.
|
||||
redirect_uri: Redirect URI that was used in the authorization request.
|
||||
required_scopes: Space-separated scopes that must be present in the token
|
||||
response ``scope`` field (RFC 6749 §3.3). Raises ``ValueError`` if
|
||||
the provider granted fewer scopes than requested.
|
||||
|
||||
Returns:
|
||||
The JSON payload returned by the userinfo endpoint as a plain ``dict``.
|
||||
|
||||
Raises:
|
||||
ValueError: If the provider granted a different token type than ``bearer``
|
||||
or did not grant all ``required_scopes``.
|
||||
"""
|
||||
async with httpx.AsyncClient() as client:
|
||||
token_resp = await client.post(
|
||||
token_url,
|
||||
data={
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"redirect_uri": redirect_uri,
|
||||
},
|
||||
headers={"Accept": "application/json"},
|
||||
)
|
||||
token_resp.raise_for_status()
|
||||
token_data = token_resp.json()
|
||||
|
||||
if token_data.get("token_type", "bearer").lower() != "bearer":
|
||||
raise ValueError(
|
||||
f"unsupported token_type: {token_data.get('token_type')!r}"
|
||||
)
|
||||
|
||||
if required_scopes is not None:
|
||||
granted = set(token_data.get("scope", "").split())
|
||||
missing = set(required_scopes.split()) - granted
|
||||
if missing:
|
||||
raise ValueError(f"provider did not grant required scopes: {missing}")
|
||||
|
||||
access_token = token_data["access_token"]
|
||||
|
||||
userinfo_resp = await client.get(
|
||||
userinfo_url,
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
)
|
||||
userinfo_resp.raise_for_status()
|
||||
return userinfo_resp.json()
|
||||
|
||||
|
||||
def oauth_generate_state_token() -> str:
|
||||
"""Generate a cryptographically random CSRF token for the OAuth ``state`` parameter."""
|
||||
return secrets.token_urlsafe(32)
|
||||
|
||||
|
||||
def oauth_build_authorization_redirect(
|
||||
authorization_url: str,
|
||||
*,
|
||||
client_id: str,
|
||||
scopes: str,
|
||||
redirect_uri: str,
|
||||
destination: str,
|
||||
state_token: str,
|
||||
) -> RedirectResponse:
|
||||
"""Return an OAuth 2.0 authorization ``RedirectResponse``.
|
||||
|
||||
Args:
|
||||
authorization_url: Provider's authorization endpoint.
|
||||
client_id: OAuth application client ID.
|
||||
scopes: Space-separated list of requested scopes.
|
||||
redirect_uri: URI the provider should redirect back to after authorization.
|
||||
destination: URL the user should be sent to after the full OAuth flow
|
||||
completes (embedded in ``state``).
|
||||
state_token: CSRF token generated by :func:`oauth_generate_state_token`.
|
||||
Must be stored server-side (session or signed cookie) and verified via
|
||||
:func:`oauth_decode_state` on the callback endpoint (RFC 6749 §10.12).
|
||||
|
||||
Returns:
|
||||
A :class:`~fastapi.responses.RedirectResponse` to the provider's
|
||||
authorization page.
|
||||
"""
|
||||
params = urlencode(
|
||||
{
|
||||
"client_id": client_id,
|
||||
"response_type": "code",
|
||||
"scope": scopes,
|
||||
"redirect_uri": redirect_uri,
|
||||
"state": oauth_encode_state(destination, state_token),
|
||||
}
|
||||
)
|
||||
return RedirectResponse(f"{authorization_url}?{params}")
|
||||
|
||||
|
||||
def oauth_encode_state(url: str, state_token: str) -> str:
|
||||
"""Encode a destination URL and CSRF token into an OAuth ``state`` parameter.
|
||||
|
||||
Args:
|
||||
url: Post-login destination URL.
|
||||
state_token: CSRF token from :func:`oauth_generate_state_token`.
|
||||
"""
|
||||
payload = json.dumps({"n": state_token, "d": url}, separators=(",", ":"))
|
||||
return base64.urlsafe_b64encode(payload.encode()).decode()
|
||||
|
||||
|
||||
def oauth_decode_state(
|
||||
state: str | None, *, expected_state_token: str, fallback: str
|
||||
) -> str:
|
||||
"""Decode and CSRF-verify an OAuth ``state`` parameter.
|
||||
|
||||
Uses a constant-time comparison for the CSRF token to prevent timing attacks.
|
||||
|
||||
Args:
|
||||
state: Raw ``state`` query parameter from the provider's callback.
|
||||
expected_state_token: The token stored before the authorization redirect.
|
||||
If it does not match the decoded value, ``fallback`` is returned.
|
||||
fallback: URL to return when ``state`` is absent, malformed, or fails
|
||||
CSRF verification.
|
||||
|
||||
Returns:
|
||||
The destination URL embedded in ``state``, or ``fallback``.
|
||||
|
||||
Important:
|
||||
**Single-use**: delete the stored token from the session immediately
|
||||
after calling this function — whether it matched or not — so that a
|
||||
captured callback URL cannot be replayed.
|
||||
|
||||
**Open-redirect**: validate the returned URL against a known-good
|
||||
origin or relative-path allowlist before issuing the final redirect.
|
||||
Do not forward arbitrary URLs to ``RedirectResponse``.
|
||||
"""
|
||||
if not state or state == "null": # "null" guards against JS JSON.stringify(null)
|
||||
return fallback
|
||||
try:
|
||||
padded = state + "=" * (-len(state) % 4)
|
||||
payload = json.loads(base64.urlsafe_b64decode(padded).decode("utf-8"))
|
||||
if not isinstance(payload, dict) or not hmac.compare_digest(
|
||||
payload.get("n", "").encode(), expected_state_token.encode()
|
||||
):
|
||||
return fallback
|
||||
return str(payload["d"])
|
||||
except (UnicodeDecodeError, ValueError, binascii.Error, KeyError):
|
||||
return fallback
|
||||
Reference in New Issue
Block a user