fix(server): scope auth DB sessions to prevent connection-pool exhaustion (#6237)
This commit is contained in:
+14
-9
@@ -3,7 +3,7 @@ import secrets
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from db import get_db
|
||||
from db import SessionLocal
|
||||
from fastapi import Depends, HTTPException, Request
|
||||
from fastapi.security import APIKeyHeader, HTTPAuthorizationCredentials, HTTPBearer
|
||||
from jose import JWTError, jwt
|
||||
@@ -145,19 +145,24 @@ async def verify_auth(
|
||||
request: Request,
|
||||
credentials: HTTPAuthorizationCredentials | None = Depends(bearer_scheme),
|
||||
x_api_key: str | None = Depends(api_key_header),
|
||||
db: Session = Depends(get_db),
|
||||
) -> User | None:
|
||||
"""Authenticate via JWT, X-API-Key, or legacy ADMIN_API_KEY. Returns User or None."""
|
||||
"""Authenticate via JWT, X-API-Key, or legacy ADMIN_API_KEY. Returns User or None.
|
||||
|
||||
A short-lived session is opened only on the branches that query the DB, so no
|
||||
pooled connection is held for the lifetime of the (possibly long-running) request.
|
||||
"""
|
||||
if credentials is not None:
|
||||
_mark_auth_type(request, "bearer")
|
||||
return _resolve_user_from_jwt(credentials.credentials, db)
|
||||
with SessionLocal() as db:
|
||||
return _resolve_user_from_jwt(credentials.credentials, db)
|
||||
|
||||
if x_api_key is not None:
|
||||
if ADMIN_API_KEY and secrets.compare_digest(x_api_key, ADMIN_API_KEY):
|
||||
_mark_auth_type(request, "admin_api_key")
|
||||
return None
|
||||
_mark_auth_type(request, "api_key")
|
||||
return _resolve_user_from_api_key(x_api_key, db)
|
||||
with SessionLocal() as db:
|
||||
return _resolve_user_from_api_key(x_api_key, db)
|
||||
|
||||
if AUTH_DISABLED:
|
||||
_mark_auth_type(request, "disabled")
|
||||
@@ -173,12 +178,12 @@ async def verify_auth(
|
||||
async def require_auth(
|
||||
request: Request,
|
||||
user: User | None = Depends(verify_auth),
|
||||
db: Session = Depends(get_db),
|
||||
) -> User:
|
||||
"""Like verify_auth but guarantees a non-None User. Use for endpoints that require auth."""
|
||||
if user is None:
|
||||
if getattr(request.state, "auth_type", "none") in {"admin_api_key", "disabled"}:
|
||||
default_user = _get_default_user(db)
|
||||
with SessionLocal() as db:
|
||||
default_user = _get_default_user(db)
|
||||
if default_user is not None:
|
||||
return default_user
|
||||
raise HTTPException(status_code=401, detail="Authentication required.")
|
||||
@@ -193,7 +198,6 @@ _BOOTSTRAP_ADMIN = User(
|
||||
async def require_admin(
|
||||
request: Request,
|
||||
user: User | None = Depends(verify_auth),
|
||||
db: Session = Depends(get_db),
|
||||
) -> User:
|
||||
"""Like require_auth but also enforces admin role.
|
||||
|
||||
@@ -203,7 +207,8 @@ async def require_admin(
|
||||
auth_type = getattr(request.state, "auth_type", "none")
|
||||
if user is None:
|
||||
if auth_type in {"admin_api_key", "disabled"}:
|
||||
default_user = _get_default_user(db)
|
||||
with SessionLocal() as db:
|
||||
default_user = _get_default_user(db)
|
||||
if default_user is not None:
|
||||
if default_user.role != "admin":
|
||||
raise HTTPException(status_code=403, detail="Admin role required.")
|
||||
|
||||
+17
-9
@@ -175,18 +175,23 @@ def update_me(
|
||||
user: User = Depends(require_auth),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
if body.name is not None and body.name.strip():
|
||||
user.name = body.name.strip()
|
||||
# require_auth resolves the user in its own short-lived session, so `user` is
|
||||
# detached from this request's `db`. Load a session-managed copy to mutate.
|
||||
db_user = db.get(User, user.id)
|
||||
if db_user is None:
|
||||
raise HTTPException(status_code=404, detail="User not found.")
|
||||
|
||||
if body.email is not None and body.email != user.email:
|
||||
collision = db.scalar(select(User).where(User.email == body.email, User.id != user.id))
|
||||
if body.name is not None and body.name.strip():
|
||||
db_user.name = body.name.strip()
|
||||
|
||||
if body.email is not None and body.email != db_user.email:
|
||||
collision = db.scalar(select(User).where(User.email == body.email, User.id != db_user.id))
|
||||
if collision is not None:
|
||||
raise HTTPException(status_code=409, detail="Email is already in use.")
|
||||
user.email = body.email
|
||||
db_user.email = body.email
|
||||
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
return user
|
||||
return db_user
|
||||
|
||||
|
||||
@router.post("/change-password", response_model=MessageResponse)
|
||||
@@ -195,12 +200,15 @@ def change_password(
|
||||
user: User = Depends(require_auth),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
if not verify_password(body.current_password, user.password_hash):
|
||||
# require_auth resolves the user in its own short-lived session, so `user` is
|
||||
# detached from this request's `db`. Load a session-managed copy to mutate.
|
||||
db_user = db.get(User, user.id)
|
||||
if db_user is None or not verify_password(body.current_password, db_user.password_hash):
|
||||
raise HTTPException(status_code=401, detail="Current password is incorrect.")
|
||||
|
||||
_require_password_length(body.new_password)
|
||||
|
||||
user.password_hash = hash_password(body.new_password)
|
||||
db_user.password_hash = hash_password(body.new_password)
|
||||
db.commit()
|
||||
return MessageResponse(message="Password updated.")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user