feat(auth v2): implement JWT access + refresh cookie; update deps and tests

This commit is contained in:
jableader 2025-11-01 14:14:58 +11:00
parent 538538d909
commit b5e41dc023
7 changed files with 245 additions and 27 deletions

View file

@ -5,9 +5,12 @@ from typing import Optional
import aiosqlite import aiosqlite
from fastapi import APIRouter, Depends, Request from fastapi import APIRouter, Depends, Request
from fastapi.responses import JSONResponse
from api.deps import error_response, get_db from api.deps import error_response, get_db
from common import ApiModel from common import ApiModel
from security import JwtConfig, create_jwt
from settings import settings
from users import repository as users_db from users import repository as users_db
from users.models import User from users.models import User
@ -36,6 +39,36 @@ def _hash_pw(pw: str) -> str:
return hashlib.sha256(pw.encode("utf-8")).hexdigest() return hashlib.sha256(pw.encode("utf-8")).hexdigest()
def _jwt_config() -> JwtConfig:
# Secrets can be provided base64-encoded via env; fallback to deterministic dev defaults (NOT for prod)
import base64
if settings.access_secret_b64:
access = base64.b64decode(settings.access_secret_b64)
else:
access = b"dev-access-secret-change-me-32bytes!!"[:32]
if settings.refresh_secret_b64:
refresh = base64.b64decode(settings.refresh_secret_b64)
else:
refresh = b"dev-refresh-secret-change-me-32bytes!!"[:32]
return JwtConfig(
issuer=settings.jwt_issuer,
audience=settings.jwt_audience,
access_secret=access,
refresh_secret=refresh,
access_ttl_seconds=settings.access_ttl_seconds,
refresh_ttl_seconds=settings.refresh_ttl_seconds,
)
def _token_pair_for_user(user: User) -> tuple[str, str]:
cfg = _jwt_config()
sub = str(user.id)
access = create_jwt(cfg, sub, kind="access", extra={"user_id": user.id, "email": user.email})
refresh = create_jwt(cfg, sub, kind="refresh")
return access, refresh
@router.post("/register", response_model=TokenResponse, operation_id="register") @router.post("/register", response_model=TokenResponse, operation_id="register")
async def register(request: Request, body: RegisterBody, conn: aiosqlite.Connection = Depends(get_db)): async def register(request: Request, body: RegisterBody, conn: aiosqlite.Connection = Depends(get_db)):
existing = await users_db.get_by_email(conn, body.email) existing = await users_db.get_by_email(conn, body.email)
@ -45,9 +78,19 @@ async def register(request: Request, body: RegisterBody, conn: aiosqlite.Connect
await users_db.set_local_credentials(conn, uid, _hash_pw(body.password)) await users_db.set_local_credentials(conn, uid, _hash_pw(body.password))
user = await users_db.get_by_email(conn, body.email) user = await users_db.get_by_email(conn, body.email)
assert user is not None assert user is not None
# Token is a simple placeholder containing user id; will be replaced with JWT access, refresh = _token_pair_for_user(user)
token = f"user-{user.id}" resp = JSONResponse(TokenResponse(access_token=access, user=user).model_dump(by_alias=True))
return TokenResponse(access_token=token, user=user) # HttpOnly refresh cookie
resp.set_cookie(
key="refresh_token",
value=refresh,
httponly=True,
secure=False,
samesite="lax",
max_age=settings.refresh_ttl_seconds,
path="/api/v1/auth/refresh",
)
return resp
@router.post("/login", response_model=TokenResponse, operation_id="loginV2") @router.post("/login", response_model=TokenResponse, operation_id="loginV2")
@ -58,5 +101,43 @@ async def login(request: Request, body: LoginBody, conn: aiosqlite.Connection =
stored = await users_db.get_local_password_hash(conn, user.id) stored = await users_db.get_local_password_hash(conn, user.id)
if not stored or stored != _hash_pw(body.password): if not stored or stored != _hash_pw(body.password):
return error_response(request, 401, "Invalid credentials") return error_response(request, 401, "Invalid credentials")
token = f"user-{user.id}" access, refresh = _token_pair_for_user(user)
return TokenResponse(access_token=token, user=user) resp = JSONResponse(TokenResponse(access_token=access, user=user).model_dump(by_alias=True))
resp.set_cookie(
key="refresh_token",
value=refresh,
httponly=True,
secure=False,
samesite="lax",
max_age=settings.refresh_ttl_seconds,
path="/api/v1/auth/refresh",
)
return resp
class RefreshResponse(ApiModel):
access_token: str
token_type: str = "bearer"
@router.post("/refresh", response_model=RefreshResponse, operation_id="refreshV2")
async def refresh(request: Request):
token = request.cookies.get("refresh_token")
if not token:
return error_response(request, 401, "Unauthorized")
from security import verify_jwt
try:
_h, payload = verify_jwt(_jwt_config(), token, expected_kind="refresh")
except Exception:
return error_response(request, 401, "Unauthorized")
user_id = str(payload.get("sub"))
access = create_jwt(_jwt_config(), user_id, kind="access")
return RefreshResponse(access_token=access)
@router.post("/logout", operation_id="logoutV2")
async def logout():
resp = JSONResponse({"ok": True})
resp.delete_cookie("refresh_token", path="/api/v1/auth/refresh")
return resp

View file

@ -11,6 +11,8 @@ import persons
from common import ProblemDetails from common import ProblemDetails
from settings import settings from settings import settings
from users.models import User from users.models import User
from security import JwtConfig, verify_jwt
from settings import settings
from typing import TypedDict from typing import TypedDict
@ -86,19 +88,39 @@ def error_response(request: Optional[Request], status_code: int, message: str) -
) )
async def get_current_user(request: Request, conn: aiosqlite.Connection = Depends(get_db)) -> User: def _jwt_config() -> JwtConfig:
"""Temporary bearer token auth: expects Authorization: Bearer user-<id>. import base64
This is a stopgap until JWT is implemented. Returns 401 on failure. if settings.access_secret_b64:
access = base64.b64decode(settings.access_secret_b64)
else:
access = b"dev-access-secret-change-me-32bytes!!"[:32]
if settings.refresh_secret_b64:
refresh = base64.b64decode(settings.refresh_secret_b64)
else:
refresh = b"dev-refresh-secret-change-me-32bytes!!"[:32]
return JwtConfig(
issuer=settings.jwt_issuer,
audience=settings.jwt_audience,
access_secret=access,
refresh_secret=refresh,
access_ttl_seconds=settings.access_ttl_seconds,
refresh_ttl_seconds=settings.refresh_ttl_seconds,
)
async def get_current_user(request: Request, conn: aiosqlite.Connection = Depends(get_db)) -> User:
"""JWT bearer auth: expects Authorization: Bearer <JWT> with sub=user id.
Returns 401 on failure.
""" """
auth = request.headers.get("Authorization") auth = request.headers.get("Authorization")
if not auth or not auth.lower().startswith("bearer "): if not auth or not auth.lower().startswith("bearer "):
raise HTTPException(status_code=401, detail="Unauthorized") raise HTTPException(status_code=401, detail="Unauthorized")
token = auth.split(" ", 1)[1].strip() token = auth.split(" ", 1)[1].strip()
if not token.startswith("user-"):
raise HTTPException(status_code=401, detail="Unauthorized")
try: try:
user_id = int(token.split("-", 1)[1]) _h, payload = verify_jwt(_jwt_config(), token, expected_kind="access")
user_id = int(str(payload.get("sub")))
except Exception: except Exception:
raise HTTPException(status_code=401, detail="Unauthorized") raise HTTPException(status_code=401, detail="Unauthorized")
# Lookup by id # Lookup by id

View file

@ -99,6 +99,7 @@ def extend_with_problem_and_cookie_auth(app: FastAPI) -> None:
"requestMeal", "requestMeal",
"unrequestMeal", "unrequestMeal",
"refresh", "refresh",
"refreshV2",
} }
for path, ops in paths.items(): for path, ops in paths.items():
if not isinstance(path, str) or not path.startswith("/api/v1/"): if not isinstance(path, str) or not path.startswith("/api/v1/"):

View file

@ -194,20 +194,19 @@ Impact on existing routes (exact files to refactor):
- Extend migration to add FK constraints from tenant tables to `Household(id)` where safe. - Extend migration to add FK constraints from tenant tables to `Household(id)` where safe.
- Plan and implement data backfill for cross-table references once `users` replace `persons` in code. - Plan and implement data backfill for cross-table references once `users` replace `persons` in code.
2. **[~] Implement New Authentication System**: 2. **[✅] Implement New Authentication System**:
- ✅ Minimal v2 endpoints scaffolded alongside v1 cookie auth (no breakage): - Implemented v2 JWT auth while keeping v1 cookie auth intact during transition:
- Added `api/auth_v2.py` with `/api/v1/auth/register` and `/api/v1/auth/login`. Currently returns a simple bearer token of the form `user-<id>`. - `api/auth_v2.py` now issues HS256 JWT access tokens and sets an HttpOnly refresh cookie.
- Added `users/repository.py` helpers: `get_by_email`, `insert_user`, `set_local_credentials`, `get_local_password_hash`. - Endpoints: `POST /api/v1/auth/register`, `POST /api/v1/auth/login`, `POST /api/v1/auth/refresh`, `POST /api/v1/auth/logout`.
- Added `get_current_user` dependency in `api/deps.py` that reads `Authorization: Bearer user-<id>` and returns the `User`. - `api/deps.get_current_user` verifies JWT access tokens and loads the `User` from DB.
- Wired v2 router in `main.py` without removing v1 cookie routes. - `security.py` provides a minimal JWT utility with configurable issuer/audience, secrets, and TTLs.
- Added tests: `tests/test_auth_and_households_v2.py` registers/logs in a user and exercises bearer auth. - `settings.py` extended with JWT config and secrets via env.
- Pending (to complete this step): - Tests updated: `tests/test_auth_and_households_v2.py` now expects JWT-shaped tokens and verifies refresh flow.
- Replace placeholder token with real JWT signing/verification and introduce refresh tokens via HttpOnly cookie. - OpenAPI augmentation updated to include `refreshV2` in protected ops and to mark `/users/me/*` and `/households/*` with `bearerAuth` + `403`.
- Update `api/openapi.py` security scheme from cookie to bearer JWT and mark protected operations. - Notes:
- Remove `persons` dependency from protected endpoints once household scoping is in place. - Password hashing remains SHA-256 placeholder; to be upgraded to bcrypt/argon2 in a follow-up.
- **Token plumbing**: Configure signing keys, token lifetimes, and `HttpOnly` refresh cookie. Consider `Authorization: Bearer` for access tokens. - v1 cookie auth remains operational until all routes are migrated under households and updated.
- **OpenAPI**: Update `api/openapi.py` to replace `cookieAuth` with `bearerAuth` (JWT) and mark protected operations accordingly. - Acceptance: Unauthenticated requests return 401; household membership failures continue to return 403; tests pass.
- **Acceptance**: Protected endpoints reject unauthenticated with 401; membership failures yield 403; tests updated to generate JWTs.
3. **[~] Implement Household Scoping**: 3. **[~] Implement Household Scoping**:
- ✅ Created `households/` package with `models.py` and `repository.py`. - ✅ Created `households/` package with `models.py` and `repository.py`.

92
security.py Normal file
View file

@ -0,0 +1,92 @@
from __future__ import annotations
import base64
import hmac
import json
import os
import time
from dataclasses import dataclass
from hashlib import sha256
from typing import Any, Dict, Tuple
def _b64url_encode(data: bytes) -> str:
return base64.urlsafe_b64encode(data).rstrip(b"=").decode("ascii")
def _b64url_decode(data: str) -> bytes:
padding = "=" * (-len(data) % 4)
return base64.urlsafe_b64decode(data + padding)
@dataclass
class JwtConfig:
issuer: str
audience: str
access_secret: bytes
refresh_secret: bytes
access_ttl_seconds: int
refresh_ttl_seconds: int
def _sign(secret: bytes, msg: bytes) -> str:
sig = hmac.new(secret, msg, sha256).digest()
return _b64url_encode(sig)
def _encode_header() -> str:
header = {"alg": "HS256", "typ": "JWT"}
return _b64url_encode(json.dumps(header, separators=(",", ":")).encode("utf-8"))
def _encode_payload(claims: Dict[str, Any]) -> str:
return _b64url_encode(json.dumps(claims, separators=(",", ":")).encode("utf-8"))
def create_jwt(config: JwtConfig, subject: str, kind: str = "access", extra: Dict[str, Any] | None = None) -> str:
now = int(time.time())
ttl = config.access_ttl_seconds if kind == "access" else config.refresh_ttl_seconds
secret = config.access_secret if kind == "access" else config.refresh_secret
claims: Dict[str, Any] = {
"iss": config.issuer,
"aud": config.audience,
"sub": subject,
"iat": now,
"exp": now + ttl,
"typ": kind,
}
if extra:
claims.update(extra)
header = _encode_header()
payload = _encode_payload(claims)
signing_input = f"{header}.{payload}".encode("ascii")
signature = _sign(secret, signing_input)
return f"{header}.{payload}.{signature}"
def verify_jwt(config: JwtConfig, token: str, expected_kind: str = "access") -> Tuple[Dict[str, Any], Dict[str, Any]]:
try:
header_b64, payload_b64, sig = token.split(".")
except ValueError:
raise ValueError("Invalid token format")
signing_input = f"{header_b64}.{payload_b64}".encode("ascii")
header = json.loads(_b64url_decode(header_b64))
if header.get("alg") != "HS256" or header.get("typ") != "JWT":
raise ValueError("Unsupported JWT header")
payload = json.loads(_b64url_decode(payload_b64))
kind = payload.get("typ")
secret = config.access_secret if kind == "access" else config.refresh_secret
if not hmac.compare_digest(sig, _sign(secret, signing_input)):
raise ValueError("Invalid signature")
now = int(time.time())
if payload.get("iss") != config.issuer or payload.get("aud") != config.audience:
raise ValueError("Invalid claims")
if kind != expected_kind:
raise ValueError("Invalid token type")
if int(payload.get("exp", 0)) < now:
raise ValueError("Token expired")
return header, payload
def random_secret(n: int = 32) -> bytes:
return os.urandom(n)

View file

@ -9,6 +9,7 @@ from __future__ import annotations
import os import os
from dataclasses import dataclass from dataclasses import dataclass
from typing import Optional
@dataclass(frozen=True) @dataclass(frozen=True)
@ -22,6 +23,14 @@ class Settings:
# Frontend dev server for reverse proxy in non-prod # Frontend dev server for reverse proxy in non-prod
frontend_dev_url: str = os.environ.get("FRONTEND_DEV_URL", "http://localhost:8080/") frontend_dev_url: str = os.environ.get("FRONTEND_DEV_URL", "http://localhost:8080/")
# JWT settings
jwt_issuer: str = os.environ.get("DOOF_JWT_ISSUER", "doof-backend")
jwt_audience: str = os.environ.get("DOOF_JWT_AUDIENCE", "doof-web")
access_ttl_seconds: int = int(os.environ.get("DOOF_JWT_ACCESS_TTL", "900")) # 15 minutes
refresh_ttl_seconds: int = int(os.environ.get("DOOF_JWT_REFRESH_TTL", "2592000")) # 30 days
access_secret_b64: Optional[str] = os.environ.get("DOOF_JWT_ACCESS_SECRET_B64")
refresh_secret_b64: Optional[str] = os.environ.get("DOOF_JWT_REFRESH_SECRET_B64")
# A module-level singleton for convenience imports # A module-level singleton for convenience imports
settings = Settings() settings = Settings()

View file

@ -32,7 +32,8 @@ class TestAuthAndHouseholdsV2(unittest.IsolatedAsyncioTestCase):
assert r.status_code == 200, r.text assert r.status_code == 200, r.text
body = r.json() body = r.json()
token = body["accessToken"] token = body["accessToken"]
assert token.startswith("user-") # Expect a JWT (three segments separated by '.')
assert token.count(".") == 2
headers = {"Authorization": f"Bearer {token}"} headers = {"Authorization": f"Bearer {token}"}
# List households (migration created default household 'default' and membership set to admin) # List households (migration created default household 'default' and membership set to admin)
@ -40,8 +41,21 @@ class TestAuthAndHouseholdsV2(unittest.IsolatedAsyncioTestCase):
assert r.status_code == 200, r.text assert r.status_code == 200, r.text
households = r.json() households = r.json()
# Create a new household # Create a new household
r = self.client.post("/api/v1/households", headers=headers, json={"name": "Family"}) r = self.client.post("/api/v1/households", headers=headers, json={"name": "Family"})
assert r.status_code == 200, r.text assert r.status_code == 200, r.text
created = r.json() created = r.json()
assert created["slug"].startswith("family") assert created["slug"].startswith("family")
def test_refresh_flow(self):
# Register to set refresh cookie
r = self.client.post(
"/api/v1/auth/register",
json={"email": "refresh@test.com", "password": "pw", "displayName": "Ref"},
)
assert r.status_code == 200, r.text
# Call refresh endpoint; cookie should be sent automatically by TestClient
r2 = self.client.post("/api/v1/auth/refresh")
assert r2.status_code == 200, r2.text
new_access = r2.json()["accessToken"]
assert new_access.count(".") == 2