fix: multiple security bugs + remove example for now

This commit is contained in:
2026-05-07 18:21:34 -04:00
parent 95f5a83bd2
commit 6bef88fde6
18 changed files with 336 additions and 587 deletions
-1
View File
@@ -1 +0,0 @@
# Authentication
+60 -75
View File
@@ -47,12 +47,9 @@ async def me(user: User = Security(bearer)):
#### Token prefix #### Token prefix
The optional `prefix` parameter restricts a `BearerTokenAuth` instance to tokens 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.
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 This lets you deploy multiple `BearerTokenAuth` instances in the same application and disambiguate them efficiently in `MultiAuth`:
and disambiguate them efficiently in `MultiAuth`:
```python ```python
user_bearer = BearerTokenAuth(verify_user, prefix="user_") # matches "Bearer user_..." user_bearer = BearerTokenAuth(verify_user, prefix="user_") # matches "Bearer user_..."
@@ -63,9 +60,7 @@ Use [`generate_token()`](#token-generation) to create correctly-prefixed tokens.
#### Token generation #### Token generation
`BearerTokenAuth.generate_token()` produces a secure random token ready to store `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:
in your database and return to the client. If a prefix is configured it is
prepended automatically:
```python ```python
bearer = BearerTokenAuth(verify_token, prefix="user_") bearer = BearerTokenAuth(verify_token, prefix="user_")
@@ -75,18 +70,23 @@ await db.store_token(user_id, token)
return {"access_token": token, "token_type": "bearer"} return {"access_token": token, "token_type": "bearer"}
``` ```
The client sends `Authorization: Bearer user_Xk3mN...` and the validator receives The client sends `Authorization: Bearer user_Xk3mN...` and the validator receives the full token (prefix included) to compare against the stored value.
the full token (prefix included) to compare against the stored value.
### [`CookieAuth`](../reference/security.md#fastapi_toolsets.security.CookieAuth) ### [`CookieAuth`](../reference/security.md#fastapi_toolsets.security.CookieAuth)
Reads a named cookie. Wraps `APIKeyCookie` for OpenAPI. Reads a named cookie. Wraps `APIKeyCookie` for OpenAPI.
Cookies are issued with the `Secure` flag set by default, meaning they are only transmitted over HTTPS. Set `secure=False` when running locally over plain HTTP:
```python ```python
from fastapi_toolsets.security import CookieAuth from fastapi_toolsets.security import CookieAuth
# Production (HTTPS) — default
cookie_auth = CookieAuth("session", validator=verify_session) cookie_auth = CookieAuth("session", validator=verify_session)
# Local development (HTTP only)
cookie_auth = CookieAuth("session", validator=verify_session, secure=False)
@app.get("/me") @app.get("/me")
async def me(user: User = Security(cookie_auth)): async def me(user: User = Security(cookie_auth)):
return user return user
@@ -94,16 +94,17 @@ async def me(user: User = Security(cookie_auth)):
#### Signed cookies #### Signed cookies
Pass `secret_key` to enable HMAC-SHA256 signed, tamper-proof cookies. The cookie Pass `secret_key` to enable HMAC-SHA256 signed, tamper-proof cookies. The cookie payload includes an expiry timestamp (`ttl`, default 24 h). No database entry is required — the signature is self-contained.
payload includes an expiry timestamp (`ttl`, default 24 h). No database entry is
required — the signature is self-contained.
Use `set_cookie()` to issue the signed cookie on login and `delete_cookie()` to Use `set_cookie()` to issue the signed cookie on login and `delete_cookie()` to clear it on logout:
clear it on logout:
```python ```python
# Production
cookie_auth = CookieAuth("session", verify_session, secret_key="your-secret") cookie_auth = CookieAuth("session", verify_session, secret_key="your-secret")
# Local development
cookie_auth = CookieAuth("session", verify_session, secret_key="your-secret", secure=False)
@app.post("/login") @app.post("/login")
async def login(response: Response): async def login(response: Response):
cookie_auth.set_cookie(response, user_id) cookie_auth.set_cookie(response, user_id)
@@ -119,8 +120,7 @@ async def me(user: User = Security(cookie_auth)):
return user return user
``` ```
When `secret_key` is not set, the raw cookie value is passed directly to the When `secret_key` is not set, the raw cookie value is passed directly to the validator (stateful session behaviour — you manage the session store).
validator (stateful session behaviour — you manage the session store).
### [`APIKeyHeaderAuth`](../reference/security.md#fastapi_toolsets.security.APIKeyHeaderAuth) ### [`APIKeyHeaderAuth`](../reference/security.md#fastapi_toolsets.security.APIKeyHeaderAuth)
@@ -136,14 +136,11 @@ async def data(user: User = Security(api_key_auth)):
return user return user
``` ```
The header name is configurable — use any header your API defines (e.g. The header name is configurable — use any header your API defines (e.g. `"X-API-Key"`, `"Authorization"`, `"X-Service-Token"`).
`"X-API-Key"`, `"Authorization"`, `"X-Service-Token"`).
## Typed validator kwargs ## Typed validator kwargs
All auth classes forward extra instantiation keyword arguments to the validator. 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.
Arguments can be any type — enums, strings, integers, etc. The validator returns
the authenticated identity, which FastAPI injects directly into the route handler.
```python ```python
async def verify_token(token: str, *, role: Role, permission: str) -> User: async def verify_token(token: str, *, role: Role, permission: str) -> User:
@@ -155,14 +152,11 @@ async def verify_token(token: str, *, role: Role, permission: str) -> User:
bearer = BearerTokenAuth(verify_token, role=Role.ADMIN, permission="billing:read") bearer = BearerTokenAuth(verify_token, role=Role.ADMIN, permission="billing:read")
``` ```
Each auth instance is self-contained — create a separate instance per distinct Each auth instance is self-contained — create a separate instance per distinct requirement instead of passing requirements through `Security(scopes=[...])`.
requirement instead of passing requirements through `Security(scopes=[...])`.
### Using `.require()` inline ### Using `.require()` inline
If declaring a new top-level variable per role feels verbose, use `.require()` to 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:
create a configured clone directly in the route decorator. The original instance
is not mutated:
```python ```python
bearer = BearerTokenAuth(verify_token) bearer = BearerTokenAuth(verify_token)
@@ -191,13 +185,9 @@ multi = MultiAuth(
## MultiAuth ## MultiAuth
[`MultiAuth`](../reference/security.md#fastapi_toolsets.security.MultiAuth) combines [`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.
multiple auth sources into a single callable. Sources are tried in order; the
first one that finds a credential wins.
If a credential is extracted but the validator raises, the exception propagates If a credential is extracted but the validator raises, the exception propagates immediately — the remaining sources are **not** tried. This prevents silent fallthrough on invalid credentials.
immediately — the remaining sources are **not** tried. This prevents silent
fallthrough on invalid credentials.
```python ```python
from fastapi_toolsets.security import MultiAuth from fastapi_toolsets.security import MultiAuth
@@ -211,9 +201,7 @@ async def data_route(user = Security(multi)):
### Using `.require()` on MultiAuth ### Using `.require()` on MultiAuth
`MultiAuth` also supports `.require()`, which propagates the kwargs to every `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:
source that implements it. Sources that do not (e.g. custom `AuthSource`
subclasses) are passed through unchanged:
```python ```python
multi = MultiAuth(bearer, cookie) multi = MultiAuth(bearer, cookie)
@@ -237,9 +225,7 @@ MultiAuth(
### Prefix-based dispatch ### Prefix-based dispatch
Because `extract()` is pure string matching (no I/O), prefix-based source 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:
selection is essentially free. Only the matching source's validator (which may
involve DB or network I/O) is ever called:
```python ```python
user_bearer = BearerTokenAuth(verify_user, prefix="user_") user_bearer = BearerTokenAuth(verify_user, prefix="user_")
@@ -251,8 +237,7 @@ multi = MultiAuth(user_bearer, org_bearer)
# "Bearer org_acme" → only verify_org runs, receives "org_acme" # "Bearer org_acme" → only verify_org runs, receives "org_acme"
``` ```
Tokens are stored and compared **with their prefix** — use `generate_token()` on Tokens are stored and compared **with their prefix** — use `generate_token()` on each source to issue correctly-prefixed tokens:
each source to issue correctly-prefixed tokens:
```python ```python
user_token = user_bearer.generate_token() # "user_..." user_token = user_bearer.generate_token() # "user_..."
@@ -261,9 +246,7 @@ org_token = org_bearer.generate_token() # "org_..."
## Custom auth sources ## Custom auth sources
Subclass [`AuthSource`](../reference/security.md#fastapi_toolsets.security.AuthSource) Subclass [`AuthSource`](../reference/security.md#fastapi_toolsets.security.AuthSource) to implement any credential extraction strategy. You only need to implement `extract()` and `authenticate()`:
to implement any credential extraction strategy. You only need to implement
`extract()` and `authenticate()`:
```python ```python
from fastapi_toolsets.security import AuthSource from fastapi_toolsets.security import AuthSource
@@ -284,16 +267,11 @@ Custom sources work transparently inside `MultiAuth`.
## OAuth 2.0 / OIDC helpers ## OAuth 2.0 / OIDC helpers
The module provides standalone async utilities for building OAuth 2.0 / OIDC The module provides standalone async utilities for building OAuth 2.0 / OIDC login flows. They handle provider discovery, authorization redirects, token exchange, and state encoding — leaving JWT validation and session management to your application.
login flows. They handle provider discovery, authorization redirects, token
exchange, and state encoding — leaving JWT validation and session management to
your application.
### Provider discovery ### Provider discovery
[`oauth_resolve_provider_urls()`](../reference/security.md#fastapi_toolsets.security.oauth_resolve_provider_urls) [`oauth_resolve_provider_urls()`](../reference/security.md#fastapi_toolsets.security.oauth_resolve_provider_urls) fetches the OIDC discovery document and returns the endpoint URLs. Results are cached in-process to avoid repeated network calls:
fetches the OIDC discovery document and returns the endpoint URLs. Results are
cached in-process to avoid repeated network calls:
```python ```python
from fastapi_toolsets.security import oauth_resolve_provider_urls from fastapi_toolsets.security import oauth_resolve_provider_urls
@@ -303,42 +281,51 @@ auth_url, token_url, userinfo_url = await oauth_resolve_provider_urls(
) )
``` ```
Returns a `(authorization_url, token_url, userinfo_url)` tuple. `userinfo_url` Returns a `(authorization_url, token_url, userinfo_url)` tuple. `userinfo_url` is `None` when the provider does not advertise one.
is `None` when the provider does not advertise one.
### Authorization redirect ### Authorization redirect
[`oauth_build_authorization_redirect()`](../reference/security.md#fastapi_toolsets.security.oauth_build_authorization_redirect) [`oauth_build_authorization_redirect()`](../reference/security.md#fastapi_toolsets.security.oauth_build_authorization_redirect) constructs the redirect to the provider's authorization page. It requires a `nonce` — a random CSRF token generated by [`oauth_generate_nonce()`](../reference/security.md#fastapi_toolsets.security.oauth_generate_nonce) — that must be stored server-side (e.g. in the session) and verified on the callback to prevent login-CSRF attacks (RFC 6749 §10.12):
constructs the redirect to the provider's authorization page. The `destination`
URL (where to send the user after the full flow) is encoded as the `state`
parameter:
```python ```python
from fastapi_toolsets.security import oauth_build_authorization_redirect from fastapi import Request
from fastapi_toolsets.security import oauth_build_authorization_redirect, oauth_generate_nonce
@app.get("/auth/google/login") @app.get("/auth/google/login")
async def google_login(): async def google_login(request: Request):
auth_url, _, _ = await oauth_resolve_provider_urls(GOOGLE_DISCOVERY_URL) auth_url, _, _ = await oauth_resolve_provider_urls(GOOGLE_DISCOVERY_URL)
nonce = oauth_generate_nonce()
request.session["oauth_nonce"] = nonce # requires SessionMiddleware
return oauth_build_authorization_redirect( return oauth_build_authorization_redirect(
auth_url, auth_url,
client_id=GOOGLE_CLIENT_ID, client_id=GOOGLE_CLIENT_ID,
scopes="openid email profile", scopes="openid email profile",
redirect_uri="https://myapp.com/auth/google/callback", redirect_uri="https://myapp.com/auth/google/callback",
destination="/dashboard", destination="/dashboard",
nonce=nonce,
) )
``` ```
### Token exchange and userinfo ### Token exchange and userinfo
[`oauth_fetch_userinfo()`](../reference/security.md#fastapi_toolsets.security.oauth_fetch_userinfo) [`oauth_fetch_userinfo()`](../reference/security.md#fastapi_toolsets.security.oauth_fetch_userinfo) performs the two-step exchange: it POSTs the authorization code to the token endpoint, then GETs the userinfo endpoint with the resulting access token.
performs the two-step exchange: it POSTs the authorization code to the token
endpoint, then GETs the userinfo endpoint with the resulting access token: On the callback, retrieve the stored nonce and pass it to [`oauth_decode_state()`](../reference/security.md#fastapi_toolsets.security.oauth_decode_state) to verify the CSRF token before processing the code:
```python ```python
from fastapi_toolsets.security import oauth_fetch_userinfo from fastapi import HTTPException, Request
from fastapi_toolsets.security import oauth_decode_state, oauth_fetch_userinfo
@app.get("/auth/google/callback") @app.get("/auth/google/callback")
async def google_callback(code: str, state: str): async def google_callback(request: Request, code: str, state: str):
# Pop nonce first — single-use, regardless of whether verification succeeds
nonce = request.session.pop("oauth_nonce", None)
if nonce is None:
raise HTTPException(status_code=400, detail="missing OAuth state")
destination = oauth_decode_state(state, expected_nonce=nonce, fallback="/")
if not destination.startswith("/"): # reject absolute URLs to prevent open-redirect
destination = "/"
_, token_url, userinfo_url = await oauth_resolve_provider_urls(GOOGLE_DISCOVERY_URL) _, token_url, userinfo_url = await oauth_resolve_provider_urls(GOOGLE_DISCOVERY_URL)
userinfo = await oauth_fetch_userinfo( userinfo = await oauth_fetch_userinfo(
token_url=token_url, token_url=token_url,
@@ -347,30 +334,28 @@ async def google_callback(code: str, state: str):
client_id=GOOGLE_CLIENT_ID, client_id=GOOGLE_CLIENT_ID,
client_secret=GOOGLE_CLIENT_SECRET, client_secret=GOOGLE_CLIENT_SECRET,
redirect_uri="https://myapp.com/auth/google/callback", redirect_uri="https://myapp.com/auth/google/callback",
required_scopes="openid email profile",
) )
user = await db.upsert_user(email=userinfo["email"]) user = await db.upsert_user(email=userinfo["email"])
destination = oauth_decode_state(state, fallback="/")
response = RedirectResponse(destination) response = RedirectResponse(destination)
session_cookie.set_cookie(response, str(user.id)) session_cookie.set_cookie(response, str(user.id))
return response return response
``` ```
Pass `required_scopes` to guard against providers silently granting fewer scopes than requested — `oauth_fetch_userinfo` raises `ValueError` if any are missing.
### State encoding ### State encoding
[`oauth_encode_state()`](../reference/security.md#fastapi_toolsets.security.oauth_encode_state) [`oauth_encode_state()`](../reference/security.md#fastapi_toolsets.security.oauth_encode_state) and [`oauth_decode_state()`](../reference/security.md#fastapi_toolsets.security.oauth_decode_state) encode and decode the destination URL together with the CSRF nonce embedded in the OAuth `state` parameter. `oauth_decode_state` returns `fallback` if `state` is absent, malformed, or the nonce does not match:
and
[`oauth_decode_state()`](../reference/security.md#fastapi_toolsets.security.oauth_decode_state)
base64url-encode and decode the destination URL embedded in the OAuth `state`
parameter. `oauth_decode_state` handles missing padding and returns the `fallback`
if `state` is absent, `"null"`, or malformed:
```python ```python
from fastapi_toolsets.security import oauth_encode_state, oauth_decode_state from fastapi_toolsets.security import oauth_encode_state, oauth_decode_state
encoded = oauth_encode_state("/dashboard") # e.g. "L2Rhc2hib2FyZA==" nonce = "my-random-nonce"
decoded = oauth_decode_state(encoded, fallback="/") # "/dashboard" encoded = oauth_encode_state("/dashboard", nonce)
decoded = oauth_decode_state(None, fallback="/") # "/" decoded = oauth_decode_state(encoded, expected_nonce=nonce, fallback="/") # "/dashboard"
decoded = oauth_decode_state("null", fallback="/") # "/" decoded = oauth_decode_state(encoded, expected_nonce="wrong", fallback="/") # "/"
decoded = oauth_decode_state(None, expected_nonce=nonce, fallback="/") # "/"
``` ```
--- ---
+3
View File
@@ -15,6 +15,7 @@ from fastapi_toolsets.security import (
oauth_decode_state, oauth_decode_state,
oauth_encode_state, oauth_encode_state,
oauth_fetch_userinfo, oauth_fetch_userinfo,
oauth_generate_nonce,
oauth_resolve_provider_urls, oauth_resolve_provider_urls,
) )
``` ```
@@ -33,6 +34,8 @@ from fastapi_toolsets.security import (
## ::: fastapi_toolsets.security.oauth_fetch_userinfo ## ::: fastapi_toolsets.security.oauth_fetch_userinfo
## ::: fastapi_toolsets.security.oauth_generate_nonce
## ::: fastapi_toolsets.security.oauth_build_authorization_redirect ## ::: fastapi_toolsets.security.oauth_build_authorization_redirect
## ::: fastapi_toolsets.security.oauth_encode_state ## ::: fastapi_toolsets.security.oauth_encode_state
-9
View File
@@ -1,9 +0,0 @@
from fastapi import FastAPI
from fastapi_toolsets.exceptions import init_exceptions_handlers
from .routes import router
app = FastAPI()
init_exceptions_handlers(app=app)
app.include_router(router=router)
-9
View File
@@ -1,9 +0,0 @@
from fastapi_toolsets.crud import CrudFactory
from .models import OAuthAccount, OAuthProvider, Team, User, UserToken
TeamCrud = CrudFactory(model=Team)
UserCrud = CrudFactory(model=User)
UserTokenCrud = CrudFactory(model=UserToken)
OAuthProviderCrud = CrudFactory(model=OAuthProvider)
OAuthAccountCrud = CrudFactory(model=OAuthAccount)
-15
View File
@@ -1,15 +0,0 @@
from fastapi import Depends
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from fastapi_toolsets.db import create_db_context, create_db_dependency
DATABASE_URL = "postgresql+asyncpg://postgres:postgres@localhost:5432/postgres"
engine = create_async_engine(url=DATABASE_URL, future=True)
async_session_maker = async_sessionmaker(bind=engine, expire_on_commit=False)
get_db = create_db_dependency(session_maker=async_session_maker)
get_db_context = create_db_context(session_maker=async_session_maker)
SessionDep = Depends(get_db)
-105
View File
@@ -1,105 +0,0 @@
import enum
from datetime import datetime
from uuid import UUID
from sqlalchemy import (
Boolean,
DateTime,
Enum,
ForeignKey,
Integer,
String,
UniqueConstraint,
)
from sqlalchemy.dialects.postgresql import UUID as PG_UUID
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship
from fastapi_toolsets.models import TimestampMixin, UUIDMixin
class Base(DeclarativeBase, UUIDMixin):
type_annotation_map = {
str: String(),
int: Integer(),
UUID: PG_UUID(as_uuid=True),
datetime: DateTime(timezone=True),
}
class UserRole(enum.Enum):
admin = "admin"
moderator = "moderator"
user = "user"
class Team(Base, TimestampMixin):
__tablename__ = "teams"
name: Mapped[str] = mapped_column(String, unique=True, index=True)
users: Mapped[list["User"]] = relationship(back_populates="team")
class User(Base, TimestampMixin):
__tablename__ = "users"
username: Mapped[str] = mapped_column(String, unique=True, index=True)
email: Mapped[str | None] = mapped_column(
String, unique=True, index=True, nullable=True
)
hashed_password: Mapped[str | None] = mapped_column(String, nullable=True)
is_active: Mapped[bool] = mapped_column(Boolean, default=True)
role: Mapped[UserRole] = mapped_column(Enum(UserRole), default=UserRole.user)
team_id: Mapped[UUID | None] = mapped_column(ForeignKey("teams.id"), nullable=True)
team: Mapped["Team | None"] = relationship(back_populates="users")
oauth_accounts: Mapped[list["OAuthAccount"]] = relationship(back_populates="user")
tokens: Mapped[list["UserToken"]] = relationship(back_populates="user")
class UserToken(Base, TimestampMixin):
"""API tokens for a user (multiple allowed)."""
__tablename__ = "user_tokens"
user_id: Mapped[UUID] = mapped_column(ForeignKey("users.id"))
# Store hashed token value
token_hash: Mapped[str] = mapped_column(String, unique=True, index=True)
name: Mapped[str | None] = mapped_column(String, nullable=True)
expires_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True
)
user: Mapped["User"] = relationship(back_populates="tokens")
class OAuthProvider(Base, TimestampMixin):
"""Configurable OAuth2 / OpenID Connect provider."""
__tablename__ = "oauth_providers"
slug: Mapped[str] = mapped_column(String, unique=True, index=True)
name: Mapped[str] = mapped_column(String)
client_id: Mapped[str] = mapped_column(String)
client_secret: Mapped[str] = mapped_column(String)
discovery_url: Mapped[str] = mapped_column(String, nullable=False)
scopes: Mapped[str] = mapped_column(String, default="openid email profile")
is_active: Mapped[bool] = mapped_column(Boolean, default=True)
accounts: Mapped[list["OAuthAccount"]] = relationship(back_populates="provider")
class OAuthAccount(Base, TimestampMixin):
"""OAuth2 / OpenID Connect account linked to a user."""
__tablename__ = "oauth_accounts"
__table_args__ = (
UniqueConstraint("provider_id", "subject", name="uq_oauth_provider_subject"),
)
user_id: Mapped[UUID] = mapped_column(ForeignKey("users.id"))
provider_id: Mapped[UUID] = mapped_column(ForeignKey("oauth_providers.id"))
# OAuth `sub` / OpenID subject identifier
subject: Mapped[str] = mapped_column(String)
user: Mapped["User"] = relationship(back_populates="oauth_accounts")
provider: Mapped["OAuthProvider"] = relationship(back_populates="accounts")
-122
View File
@@ -1,122 +0,0 @@
from typing import Annotated
from uuid import UUID
import bcrypt
from fastapi import APIRouter, Form, HTTPException, Response, Security
from fastapi_toolsets.dependencies import PathDependency
from .crud import UserCrud, UserTokenCrud
from .db import SessionDep
from .models import OAuthProvider, User, UserToken
from .schemas import (
ApiTokenCreateRequest,
ApiTokenResponse,
RegisterRequest,
UserCreate,
UserResponse,
)
from .security import auth, cookie_auth, create_api_token
ProviderDep = PathDependency(
model=OAuthProvider,
field=OAuthProvider.slug,
session_dep=SessionDep,
param_name="slug",
)
def hash_password(password: str) -> str:
return bcrypt.hashpw(password.encode(), bcrypt.gensalt()).decode()
def verify_password(plain: str, hashed: str) -> bool:
return bcrypt.checkpw(plain.encode(), hashed.encode())
router = APIRouter(prefix="/auth")
@router.post("/register", response_model=UserResponse, status_code=201)
async def register(body: RegisterRequest, session: SessionDep):
existing = await UserCrud.first(
session=session, filters=[User.username == body.username]
)
if existing:
raise HTTPException(status_code=409, detail="Username already taken")
user = await UserCrud.create(
session=session,
obj=UserCreate(
username=body.username,
email=body.email,
hashed_password=hash_password(body.password),
),
)
return user
@router.post("/token", status_code=204)
async def login(
session: SessionDep,
response: Response,
username: Annotated[str, Form()],
password: Annotated[str, Form()],
):
user = await UserCrud.first(session=session, filters=[User.username == username])
if (
not user
or not user.hashed_password
or not verify_password(password, user.hashed_password)
):
raise HTTPException(status_code=401, detail="Invalid credentials")
if not user.is_active:
raise HTTPException(status_code=403, detail="Account disabled")
cookie_auth.set_cookie(response, str(user.id))
@router.post("/logout", status_code=204)
async def logout(response: Response):
cookie_auth.delete_cookie(response)
@router.get("/me", response_model=UserResponse)
async def me(user: User = Security(auth)):
return user
@router.post("/tokens", response_model=ApiTokenResponse, status_code=201)
async def create_token(
body: ApiTokenCreateRequest,
user: User = Security(auth),
):
raw, token_row = await create_api_token(
user.id, name=body.name, expires_at=body.expires_at
)
return ApiTokenResponse(
id=token_row.id,
name=token_row.name,
expires_at=token_row.expires_at,
created_at=token_row.created_at,
token=raw,
)
@router.delete("/tokens/{token_id}", status_code=204)
async def revoke_token(
session: SessionDep,
token_id: UUID,
user: User = Security(auth),
):
if not await UserTokenCrud.first(
session=session,
filters=[UserToken.id == token_id, UserToken.user_id == user.id],
):
raise HTTPException(status_code=404, detail="Token not found")
await UserTokenCrud.delete(
session=session,
filters=[UserToken.id == token_id, UserToken.user_id == user.id],
)
@@ -1,64 +0,0 @@
from datetime import datetime
from uuid import UUID
from pydantic import EmailStr
from fastapi_toolsets.schemas import PydanticBase
class RegisterRequest(PydanticBase):
username: str
password: str
email: EmailStr | None = None
class UserResponse(PydanticBase):
id: UUID
username: str
email: str | None
role: str
is_active: bool
model_config = {"from_attributes": True}
class ApiTokenCreateRequest(PydanticBase):
name: str | None = None
expires_at: datetime | None = None
class ApiTokenResponse(PydanticBase):
id: UUID
name: str | None
expires_at: datetime | None
created_at: datetime
# Only populated on creation
token: str | None = None
model_config = {"from_attributes": True}
class OAuthProviderResponse(PydanticBase):
slug: str
name: str
model_config = {"from_attributes": True}
class UserCreate(PydanticBase):
username: str
email: str | None = None
hashed_password: str | None = None
class UserTokenCreate(PydanticBase):
user_id: UUID
token_hash: str
name: str | None = None
expires_at: datetime | None = None
class OAuthAccountCreate(PydanticBase):
user_id: UUID
provider_id: UUID
subject: str
@@ -1,100 +0,0 @@
import hashlib
from datetime import datetime, timezone
from uuid import UUID
from fastapi import HTTPException
from sqlalchemy.orm import selectinload
from fastapi_toolsets.exceptions import UnauthorizedError
from fastapi_toolsets.security import (
APIKeyHeaderAuth,
BearerTokenAuth,
CookieAuth,
MultiAuth,
)
from .crud import UserCrud, UserTokenCrud
from .db import get_db_context
from .models import User, UserRole, UserToken
from .schemas import UserTokenCreate
SESSION_COOKIE = "session"
SECRET_KEY = "123456789"
def _hash_token(token: str) -> str:
return hashlib.sha256(token.encode()).hexdigest()
async def _verify_token(token: str, role: UserRole | None = None) -> User:
async with get_db_context() as db:
user_token = await UserTokenCrud.first(
session=db,
filters=[UserToken.token_hash == _hash_token(token)],
load_options=[selectinload(UserToken.user)],
)
if user_token is None or not user_token.user.is_active:
raise UnauthorizedError()
if user_token.expires_at and user_token.expires_at < datetime.now(timezone.utc):
raise UnauthorizedError()
user = user_token.user
if role is not None and user.role != role:
raise HTTPException(status_code=403, detail="Insufficient permissions")
return user
async def _verify_cookie(user_id: str, role: UserRole | None = None) -> User:
async with get_db_context() as db:
user = await UserCrud.first(
session=db,
filters=[User.id == UUID(user_id)],
)
if not user or not user.is_active:
raise UnauthorizedError()
if role is not None and user.role != role:
raise HTTPException(status_code=403, detail="Insufficient permissions")
return user
bearer_auth = BearerTokenAuth(
validator=_verify_token,
prefix="ctf_",
)
header_auth = APIKeyHeaderAuth(
name="X-API-Key",
validator=_verify_token,
)
cookie_auth = CookieAuth(
name=SESSION_COOKIE,
validator=_verify_cookie,
secret_key=SECRET_KEY,
)
auth = MultiAuth(bearer_auth, header_auth, cookie_auth)
async def create_api_token(
user_id: UUID,
*,
name: str | None = None,
expires_at: datetime | None = None,
) -> tuple[str, UserToken]:
raw = bearer_auth.generate_token()
async with get_db_context() as db:
token_row = await UserTokenCrud.create(
session=db,
obj=UserTokenCreate(
user_id=user_id,
token_hash=_hash_token(raw),
name=name,
expires_at=expires_at,
),
)
return raw, token_row
@@ -6,6 +6,7 @@ from .oauth import (
oauth_decode_state, oauth_decode_state,
oauth_encode_state, oauth_encode_state,
oauth_fetch_userinfo, oauth_fetch_userinfo,
oauth_generate_nonce,
oauth_resolve_provider_urls, oauth_resolve_provider_urls,
) )
from .sources import APIKeyHeaderAuth, BearerTokenAuth, CookieAuth, MultiAuth from .sources import APIKeyHeaderAuth, BearerTokenAuth, CookieAuth, MultiAuth
@@ -20,5 +21,6 @@ __all__ = [
"oauth_decode_state", "oauth_decode_state",
"oauth_encode_state", "oauth_encode_state",
"oauth_fetch_userinfo", "oauth_fetch_userinfo",
"oauth_generate_nonce",
"oauth_resolve_provider_urls", "oauth_resolve_provider_urls",
] ]
+2
View File
@@ -1,5 +1,6 @@
"""Abstract base class for authentication sources.""" """Abstract base class for authentication sources."""
import functools
import inspect import inspect
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, Callable from typing import Any, Callable
@@ -15,6 +16,7 @@ def _ensure_async(fn: Callable[..., Any]) -> Callable[..., Any]:
if inspect.iscoroutinefunction(fn): if inspect.iscoroutinefunction(fn):
return fn return fn
@functools.wraps(fn)
async def wrapper(*args: Any, **kwargs: Any) -> Any: async def wrapper(*args: Any, **kwargs: Any) -> Any:
return fn(*args, **kwargs) return fn(*args, **kwargs)
+96 -26
View File
@@ -1,13 +1,20 @@
"""OAuth 2.0 / OIDC helper utilities.""" """OAuth 2.0 / OIDC helper utilities."""
import base64 import base64
import binascii
import hmac
import json
import secrets
import time as _time
from typing import Any from typing import Any
from urllib.parse import urlencode from urllib.parse import urlencode
import httpx import httpx
from fastapi.responses import RedirectResponse from fastapi.responses import RedirectResponse
_discovery_cache: dict[str, dict] = {} _discovery_cache: dict[str, tuple[dict[str, Any], float]] = {}
_DISCOVERY_TTL_SECONDS = 3600 # 1 hour
_DISCOVERY_CACHE_MAX = 32
async def oauth_resolve_provider_urls( async def oauth_resolve_provider_urls(
@@ -22,12 +29,17 @@ async def oauth_resolve_provider_urls(
A ``(authorization_url, token_url, userinfo_url)`` tuple. A ``(authorization_url, token_url, userinfo_url)`` tuple.
*userinfo_url* is ``None`` when the provider does not advertise one. *userinfo_url* is ``None`` when the provider does not advertise one.
""" """
if discovery_url not in _discovery_cache: now = _time.time()
cached = _discovery_cache.get(discovery_url)
if cached is None or now - cached[1] > _DISCOVERY_TTL_SECONDS:
async with httpx.AsyncClient() as client: async with httpx.AsyncClient() as client:
resp = await client.get(discovery_url) resp = await client.get(discovery_url)
resp.raise_for_status() resp.raise_for_status()
_discovery_cache[discovery_url] = resp.json() if len(_discovery_cache) >= _DISCOVERY_CACHE_MAX:
cfg = _discovery_cache[discovery_url] oldest = min(_discovery_cache, key=lambda k: _discovery_cache[k][1])
del _discovery_cache[oldest]
_discovery_cache[discovery_url] = (resp.json(), now)
cfg = _discovery_cache[discovery_url][0]
return ( return (
cfg["authorization_endpoint"], cfg["authorization_endpoint"],
cfg["token_endpoint"], cfg["token_endpoint"],
@@ -43,14 +55,10 @@ async def oauth_fetch_userinfo(
client_id: str, client_id: str,
client_secret: str, client_secret: str,
redirect_uri: str, redirect_uri: str,
required_scopes: str | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
"""Exchange an authorization code for tokens and return the userinfo payload. """Exchange an authorization code for tokens and return the userinfo payload.
Performs the two-step OAuth 2.0 / OIDC token exchange:
1. POSTs the authorization *code* to *token_url* to obtain an access token.
2. GETs *userinfo_url* using that access token as a Bearer credential.
Args: Args:
token_url: Provider's token endpoint. token_url: Provider's token endpoint.
userinfo_url: Provider's userinfo endpoint. userinfo_url: Provider's userinfo endpoint.
@@ -58,9 +66,16 @@ async def oauth_fetch_userinfo(
client_id: OAuth application client ID. client_id: OAuth application client ID.
client_secret: OAuth application client secret. client_secret: OAuth application client secret.
redirect_uri: Redirect URI that was used in the authorization request. 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: Returns:
The JSON payload returned by the userinfo endpoint as a plain ``dict``. 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: async with httpx.AsyncClient() as client:
token_resp = await client.post( token_resp = await client.post(
@@ -75,7 +90,20 @@ async def oauth_fetch_userinfo(
headers={"Accept": "application/json"}, headers={"Accept": "application/json"},
) )
token_resp.raise_for_status() token_resp.raise_for_status()
access_token = token_resp.json()["access_token"] 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_resp = await client.get(
userinfo_url, userinfo_url,
@@ -85,6 +113,16 @@ async def oauth_fetch_userinfo(
return userinfo_resp.json() return userinfo_resp.json()
def oauth_generate_nonce() -> str:
"""Generate a cryptographically random nonce for use as an OAuth CSRF token.
Call this before :func:`oauth_build_authorization_redirect`, persist the
returned value in the user's session or a ``Secure; HttpOnly; SameSite=Lax``
cookie, then verify it with :func:`oauth_decode_state` on the callback.
"""
return secrets.token_urlsafe(32)
def oauth_build_authorization_redirect( def oauth_build_authorization_redirect(
authorization_url: str, authorization_url: str,
*, *,
@@ -92,6 +130,7 @@ def oauth_build_authorization_redirect(
scopes: str, scopes: str,
redirect_uri: str, redirect_uri: str,
destination: str, destination: str,
nonce: str,
) -> RedirectResponse: ) -> RedirectResponse:
"""Return an OAuth 2.0 authorization ``RedirectResponse``. """Return an OAuth 2.0 authorization ``RedirectResponse``.
@@ -101,7 +140,10 @@ def oauth_build_authorization_redirect(
scopes: Space-separated list of requested scopes. scopes: Space-separated list of requested scopes.
redirect_uri: URI the provider should redirect back to after authorization. redirect_uri: URI the provider should redirect back to after authorization.
destination: URL the user should be sent to after the full OAuth flow destination: URL the user should be sent to after the full OAuth flow
completes (encoded as ``state``). completes (embedded in ``state``).
nonce: CSRF token generated by :func:`oauth_generate_nonce`. Must be
stored server-side (session or signed cookie) and verified via
:func:`oauth_decode_state` on the callback endpoint.
Returns: Returns:
A :class:`~fastapi.responses.RedirectResponse` to the provider's A :class:`~fastapi.responses.RedirectResponse` to the provider's
@@ -113,28 +155,56 @@ def oauth_build_authorization_redirect(
"response_type": "code", "response_type": "code",
"scope": scopes, "scope": scopes,
"redirect_uri": redirect_uri, "redirect_uri": redirect_uri,
"state": oauth_encode_state(destination), "state": oauth_encode_state(destination, nonce),
} }
) )
return RedirectResponse(f"{authorization_url}?{params}") return RedirectResponse(f"{authorization_url}?{params}")
def oauth_encode_state(url: str) -> str: def oauth_encode_state(url: str, nonce: str) -> str:
"""Base64url-encode a URL to embed as an OAuth ``state`` parameter.""" """Encode a destination URL and CSRF nonce into an OAuth ``state`` parameter.
return base64.urlsafe_b64encode(url.encode()).decode()
Args:
def oauth_decode_state(state: str | None, *, fallback: str) -> str: url: Post-login destination URL.
"""Decode a base64url OAuth ``state`` parameter. nonce: CSRF token from :func:`oauth_generate_nonce`.
Handles missing padding (some providers strip ``=``).
Returns *fallback* if *state* is absent, the literal string ``"null"``,
or cannot be decoded.
""" """
if not state or state == "null": payload = json.dumps({"n": nonce, "d": url}, separators=(",", ":"))
return base64.urlsafe_b64encode(payload.encode()).decode()
def oauth_decode_state(state: str | None, *, expected_nonce: str, fallback: str) -> str:
"""Decode and CSRF-verify an OAuth ``state`` parameter.
Uses a constant-time comparison for the nonce to prevent timing attacks.
Args:
state: Raw ``state`` query parameter from the provider's callback.
expected_nonce: The nonce stored before the authorization redirect.
If the decoded nonce does not match, ``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 nonce 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 return fallback
try: try:
padded = state + "=" * (4 - len(state) % 4) padded = state + "=" * (-len(state) % 4)
return base64.urlsafe_b64decode(padded).decode() payload = json.loads(base64.urlsafe_b64decode(padded).decode("utf-8"))
except Exception: if not isinstance(payload, dict) or not hmac.compare_digest(
payload.get("n", "").encode(), expected_nonce.encode()
):
return fallback
return str(payload["d"])
except (UnicodeDecodeError, ValueError, binascii.Error, KeyError):
return fallback return fallback
@@ -4,7 +4,7 @@ import inspect
import secrets import secrets
from typing import Annotated, Any, Callable from typing import Annotated, Any, Callable
from fastapi import Depends from fastapi import Depends, Request
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer, SecurityScopes from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer, SecurityScopes
from fastapi_toolsets.exceptions import UnauthorizedError from fastapi_toolsets.exceptions import UnauthorizedError
@@ -66,7 +66,7 @@ class BearerTokenAuth(AuthSource):
raise UnauthorizedError() raise UnauthorizedError()
return await self._validator(token, **self._kwargs) return await self._validator(token, **self._kwargs)
async def extract(self, request: Any) -> str | None: async def extract(self, request: Request) -> str | None:
"""Extract the raw credential from the request without validating. """Extract the raw credential from the request without validating.
Returns ``None`` if no ``Authorization: Bearer`` header is present, Returns ``None`` if no ``Authorization: Bearer`` header is present,
@@ -36,6 +36,9 @@ class CookieAuth(AuthSource):
cookie value is passed to the validator as-is. cookie value is passed to the validator as-is.
ttl: Cookie lifetime in seconds (default 24 h). Only used when ttl: Cookie lifetime in seconds (default 24 h). Only used when
``secret_key`` is set. ``secret_key`` is set.
secure: Set the ``Secure`` flag on the cookie so it is only transmitted
over HTTPS (default ``True``). Set to ``False`` only in local
development environments where HTTPS is unavailable.
**kwargs: Extra keyword arguments forwarded to the validator on every **kwargs: Extra keyword arguments forwarded to the validator on every
call (e.g. ``role=Role.ADMIN``). call (e.g. ``role=Role.ADMIN``).
""" """
@@ -47,12 +50,14 @@ class CookieAuth(AuthSource):
*, *,
secret_key: str | None = None, secret_key: str | None = None,
ttl: int = 86400, ttl: int = 86400,
secure: bool = True,
**kwargs: Any, **kwargs: Any,
) -> None: ) -> None:
self._name = name self._name = name
self._validator = _ensure_async(validator) self._validator = _ensure_async(validator)
self._secret_key = secret_key self._secret_key = secret_key
self._ttl = ttl self._ttl = ttl
self._secure = secure
self._kwargs = kwargs self._kwargs = kwargs
self._scheme = APIKeyCookie(name=name, auto_error=False) self._scheme = APIKeyCookie(name=name, auto_error=False)
@@ -120,6 +125,7 @@ class CookieAuth(AuthSource):
self._validator, self._validator,
secret_key=self._secret_key, secret_key=self._secret_key,
ttl=self._ttl, ttl=self._ttl,
secure=self._secure,
**{**self._kwargs, **kwargs}, **{**self._kwargs, **kwargs},
) )
@@ -131,9 +137,12 @@ class CookieAuth(AuthSource):
cookie_value, cookie_value,
httponly=True, httponly=True,
samesite="lax", samesite="lax",
secure=self._secure,
max_age=self._ttl, max_age=self._ttl,
) )
def delete_cookie(self, response: Response) -> None: def delete_cookie(self, response: Response) -> None:
"""Clear the session cookie (logout).""" """Clear the session cookie (logout)."""
response.delete_cookie(self._name, httponly=True, samesite="lax") response.delete_cookie(
self._name, httponly=True, samesite="lax", secure=self._secure
)
+1 -49
View File
@@ -14,42 +14,8 @@ from ..abc import AuthSource
class MultiAuth: class MultiAuth:
"""Combine multiple authentication sources into a single callable. """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: Args:
*sources: Auth source instances to try in order. *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: def __init__(self, *sources: AuthSource) -> None:
@@ -95,21 +61,7 @@ class MultiAuth:
return await self._call_fn(**kwargs) return await self._call_fn(**kwargs)
def require(self, **kwargs: Any) -> "MultiAuth": def require(self, **kwargs: Any) -> "MultiAuth":
"""Return a new :class:`MultiAuth` with kwargs forwarded to each source. """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( new_sources = tuple(
cast(Any, source).require(**kwargs) cast(Any, source).require(**kwargs)
if hasattr(source, "require") if hasattr(source, "require")
+160 -9
View File
@@ -18,6 +18,7 @@ from fastapi_toolsets.security import (
oauth_decode_state, oauth_decode_state,
oauth_encode_state, oauth_encode_state,
oauth_fetch_userinfo, oauth_fetch_userinfo,
oauth_generate_nonce,
oauth_resolve_provider_urls, oauth_resolve_provider_urls,
) )
@@ -760,7 +761,10 @@ class TestCookieAuthSigned:
"""set_cookie signs the value; the signed cookie is verified on read.""" """set_cookie signs the value; the signed cookie is verified on read."""
from fastapi import Response from fastapi import Response
auth = CookieAuth("session", cookie_validator, secret_key=self.SECRET) # secure=False for test client which runs over plain HTTP
auth = CookieAuth(
"session", cookie_validator, secret_key=self.SECRET, secure=False
)
def setup(app: FastAPI): def setup(app: FastAPI):
@app.get("/login") @app.get("/login")
@@ -778,6 +782,26 @@ class TestCookieAuthSigned:
assert response.status_code == 200 assert response.status_code == 200
assert response.json() == {"session": VALID_COOKIE} assert response.json() == {"session": VALID_COOKIE}
def test_set_cookie_has_secure_flag_by_default(self):
"""set_cookie includes Secure flag when secure=True (the default)."""
from starlette.responses import Response as StarletteResponse
auth = CookieAuth("session", cookie_validator, secret_key=self.SECRET)
response = StarletteResponse()
auth.set_cookie(response, "value")
assert "secure" in response.headers["set-cookie"].lower()
def test_set_cookie_no_secure_flag_when_disabled(self):
"""set_cookie omits Secure flag when secure=False (local dev)."""
from starlette.responses import Response as StarletteResponse
auth = CookieAuth(
"session", cookie_validator, secret_key=self.SECRET, secure=False
)
response = StarletteResponse()
auth.set_cookie(response, "value")
assert "secure" not in response.headers["set-cookie"].lower()
def test_tampered_signature_returns_401(self): def test_tampered_signature_returns_401(self):
"""A cookie whose HMAC signature has been modified is rejected.""" """A cookie whose HMAC signature has been modified is rejected."""
import base64 as _b64 import base64 as _b64
@@ -989,28 +1013,56 @@ def _make_async_client_mock(get_return=None, post_return=None):
class TestEncodeDecodeOAuthState: class TestEncodeDecodeOAuthState:
def test_encode_returns_base64url_string(self): def test_encode_returns_base64url_string(self):
result = oauth_encode_state("https://example.com/dashboard") result = oauth_encode_state("https://example.com/dashboard", "test-nonce")
assert isinstance(result, str) assert isinstance(result, str)
assert "+" not in result assert "+" not in result
assert "/" not in result assert "/" not in result
def test_round_trip(self): def test_round_trip(self):
url = "https://example.com/after-login?next=/home" url = "https://example.com/after-login?next=/home"
assert oauth_decode_state(oauth_encode_state(url), fallback="/") == url nonce = "test-nonce"
assert (
oauth_decode_state(
oauth_encode_state(url, nonce), expected_nonce=nonce, fallback="/"
)
== url
)
def test_decode_none_returns_fallback(self): def test_decode_none_returns_fallback(self):
assert oauth_decode_state(None, fallback="/home") == "/home" assert (
oauth_decode_state(None, expected_nonce="any", fallback="/home") == "/home"
)
def test_decode_null_string_returns_fallback(self): def test_decode_null_string_returns_fallback(self):
assert oauth_decode_state("null", fallback="/home") == "/home" assert (
oauth_decode_state("null", expected_nonce="any", fallback="/home")
== "/home"
)
def test_decode_invalid_base64_returns_fallback(self): def test_decode_invalid_base64_returns_fallback(self):
assert oauth_decode_state("!!!notbase64!!!", fallback="/home") == "/home" assert (
oauth_decode_state(
"!!!notbase64!!!", expected_nonce="any", fallback="/home"
)
== "/home"
)
def test_decode_handles_missing_padding(self): def test_decode_handles_missing_padding(self):
url = "https://example.com/x" url = "https://example.com/x"
encoded = oauth_encode_state(url).rstrip("=") nonce = "test-nonce"
assert oauth_decode_state(encoded, fallback="/") == url encoded = oauth_encode_state(url, nonce).rstrip("=")
assert oauth_decode_state(encoded, expected_nonce=nonce, fallback="/") == url
def test_decode_wrong_nonce_returns_fallback(self):
url = "https://example.com/dashboard"
encoded = oauth_encode_state(url, "correct-nonce")
assert (
oauth_decode_state(encoded, expected_nonce="wrong-nonce", fallback="/")
== "/"
)
def test_generate_nonce_is_random(self):
assert oauth_generate_nonce() != oauth_generate_nonce()
class TestBuildAuthorizationRedirect: class TestBuildAuthorizationRedirect:
@@ -1023,16 +1075,19 @@ class TestBuildAuthorizationRedirect:
scopes="openid email", scopes="openid email",
redirect_uri="https://app.example.com/callback", redirect_uri="https://app.example.com/callback",
destination="https://app.example.com/dashboard", destination="https://app.example.com/dashboard",
nonce="test-nonce",
) )
assert isinstance(response, RedirectResponse) assert isinstance(response, RedirectResponse)
def test_redirect_location_contains_all_params(self): def test_redirect_location_contains_all_params(self):
nonce = "test-nonce"
response = oauth_build_authorization_redirect( response = oauth_build_authorization_redirect(
"https://auth.example.com/authorize", "https://auth.example.com/authorize",
client_id="my-client", client_id="my-client",
scopes="openid email", scopes="openid email",
redirect_uri="https://app.example.com/callback", redirect_uri="https://app.example.com/callback",
destination="https://app.example.com/dashboard", destination="https://app.example.com/dashboard",
nonce=nonce,
) )
location = response.headers["location"] location = response.headers["location"]
parsed = urlparse(location) parsed = urlparse(location)
@@ -1046,7 +1101,7 @@ class TestBuildAuthorizationRedirect:
assert params["scope"] == ["openid email"] assert params["scope"] == ["openid email"]
assert params["redirect_uri"] == ["https://app.example.com/callback"] assert params["redirect_uri"] == ["https://app.example.com/callback"]
assert ( assert (
oauth_decode_state(params["state"][0], fallback="") oauth_decode_state(params["state"][0], expected_nonce=nonce, fallback="")
== "https://app.example.com/dashboard" == "https://app.example.com/dashboard"
) )
@@ -1178,3 +1233,99 @@ class TestFetchUserinfo:
"https://auth.example.com/userinfo", "https://auth.example.com/userinfo",
headers={"Authorization": "Bearer tok123"}, headers={"Authorization": "Bearer tok123"},
) )
@pytest.mark.anyio
async def test_raises_on_unsupported_token_type(self):
token_resp = MagicMock()
token_resp.raise_for_status = MagicMock()
token_resp.json.return_value = {"access_token": "tok123", "token_type": "mac"}
cm, _ = _make_async_client_mock(post_return=token_resp, get_return=MagicMock())
with patch("httpx.AsyncClient", return_value=cm):
with pytest.raises(ValueError, match="unsupported token_type"):
await oauth_fetch_userinfo(
token_url="https://auth.example.com/token",
userinfo_url="https://auth.example.com/userinfo",
code="authcode123",
client_id="client-id",
client_secret="client-secret",
redirect_uri="https://app.example.com/callback",
)
@pytest.mark.anyio
async def test_accepts_bearer_token_type_case_insensitive(self):
token_resp = MagicMock()
token_resp.raise_for_status = MagicMock()
token_resp.json.return_value = {
"access_token": "tok123",
"token_type": "Bearer",
}
userinfo_resp = MagicMock()
userinfo_resp.raise_for_status = MagicMock()
userinfo_resp.json.return_value = {"sub": "user-1"}
cm, _ = _make_async_client_mock(
post_return=token_resp, get_return=userinfo_resp
)
with patch("httpx.AsyncClient", return_value=cm):
result = await oauth_fetch_userinfo(
token_url="https://auth.example.com/token",
userinfo_url="https://auth.example.com/userinfo",
code="authcode123",
client_id="client-id",
client_secret="client-secret",
redirect_uri="https://app.example.com/callback",
)
assert result == {"sub": "user-1"}
@pytest.mark.anyio
async def test_raises_when_required_scopes_not_granted(self):
token_resp = MagicMock()
token_resp.raise_for_status = MagicMock()
token_resp.json.return_value = {"access_token": "tok123", "scope": "openid"}
cm, _ = _make_async_client_mock(post_return=token_resp, get_return=MagicMock())
with patch("httpx.AsyncClient", return_value=cm):
with pytest.raises(ValueError, match="required scopes"):
await oauth_fetch_userinfo(
token_url="https://auth.example.com/token",
userinfo_url="https://auth.example.com/userinfo",
code="authcode123",
client_id="client-id",
client_secret="client-secret",
redirect_uri="https://app.example.com/callback",
required_scopes="openid email profile",
)
@pytest.mark.anyio
async def test_passes_when_all_required_scopes_granted(self):
token_resp = MagicMock()
token_resp.raise_for_status = MagicMock()
token_resp.json.return_value = {
"access_token": "tok123",
"scope": "openid email profile",
}
userinfo_resp = MagicMock()
userinfo_resp.raise_for_status = MagicMock()
userinfo_resp.json.return_value = {"sub": "user-1", "email": "a@b.com"}
cm, _ = _make_async_client_mock(
post_return=token_resp, get_return=userinfo_resp
)
with patch("httpx.AsyncClient", return_value=cm):
result = await oauth_fetch_userinfo(
token_url="https://auth.example.com/token",
userinfo_url="https://auth.example.com/userinfo",
code="authcode123",
client_id="client-id",
client_secret="client-secret",
redirect_uri="https://app.example.com/callback",
required_scopes="openid email",
)
assert result["email"] == "a@b.com"