feat(auth v2): implement JWT access + refresh cookie; update deps and tests
This commit is contained in:
parent
538538d909
commit
b5e41dc023
7 changed files with 245 additions and 27 deletions
|
|
@ -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
|
||||||
|
|
|
||||||
34
api/deps.py
34
api/deps.py
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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/"):
|
||||||
|
|
|
||||||
|
|
@ -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
92
security.py
Normal 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)
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue