munch-ease-backend/api/deps.py

132 lines
4.3 KiB
Python
Raw Normal View History

from __future__ import annotations
2025-10-26 04:14:05 +00:00
from typing import AsyncGenerator, Optional
import aiosqlite
from fastapi import Depends, HTTPException, Request
from fastapi.responses import JSONResponse
import db
from common import ProblemDetails
from settings import settings
from users.models import User
from security import JwtConfig, verify_jwt
from typing import TypedDict
class HouseholdCtx(TypedDict):
id: int
slug: str
# Dependency to create SQLite connection with PRAGMAs and per-request transaction
async def get_db() -> AsyncGenerator[aiosqlite.Connection, None]:
sql_db = await db.connect(settings.database_path)
# Connection-level configuration
try:
# Enable FK enforcement
await sql_db.execute("PRAGMA foreign_keys=ON;")
# Prefer WAL for better concurrency; ignore result
async with sql_db.execute("PRAGMA journal_mode=WAL;") as _:
await _.fetchone()
# Reasonable durability/perf tradeoff
await sql_db.execute("PRAGMA synchronous=NORMAL;")
# Begin a transaction for the whole request
await sql_db.execute("BEGIN;")
try:
yield sql_db
await sql_db.commit()
except Exception:
await sql_db.rollback()
raise
finally:
await sql_db.close()
def error_response(request: Optional[Request], status_code: int, message: str) -> JSONResponse:
body = ProblemDetails(
title=message,
status=status_code,
type=f"https://httpstatuses.com/{status_code}",
instance=str(request.url) if request else None,
)
return JSONResponse(
content=body.model_dump(by_alias=True),
status_code=status_code,
media_type="application/problem+json",
)
def _jwt_config() -> JwtConfig:
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,
)
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")
if not auth or not auth.lower().startswith("bearer "):
raise HTTPException(status_code=401, detail="Unauthorized")
token = auth.split(" ", 1)[1].strip()
try:
_h, payload = verify_jwt(_jwt_config(), token, expected_kind="access")
user_id = int(str(payload.get("sub")))
except Exception:
raise HTTPException(status_code=401, detail="Unauthorized")
# Lookup by id
async with conn.execute(
"SELECT id, email, display_name, profile_photo_url FROM User WHERE id = ?",
(user_id,),
) as c:
row = await c.fetchone()
if not row:
raise HTTPException(status_code=401, detail="Unauthorized")
return User(id=int(row[0]), email=row[1], display_name=row[2], profile_photo_url=row[3])
async def get_household_from_slug(
request: Request,
householdSlug: str, # path parameter
user: User = Depends(get_current_user),
conn: aiosqlite.Connection = Depends(get_db),
) -> HouseholdCtx:
# Find household by slug
async with conn.execute(
"SELECT id, slug FROM Household WHERE slug = ? LIMIT 1",
(householdSlug,),
) as c:
row = await c.fetchone()
if not row:
# 404 to avoid leaking membership existence
raise HTTPException(status_code=404, detail="Household not found")
hid = int(row[0])
# Verify membership
async with conn.execute(
"SELECT 1 FROM HouseholdMember WHERE user_id = ? AND household_id = ? LIMIT 1",
(user.id, hid),
) as c:
m = await c.fetchone()
if not m:
raise HTTPException(status_code=403, detail="Forbidden")
return {"id": hid, "slug": row[1]}