fix(server): scope auth DB sessions to prevent connection-pool exhaustion (#6237)

This commit is contained in:
JainamShah-22
2026-07-20 00:07:50 +05:30
committed by GitHub
parent ddaa655edf
commit 9383e9a255
2 changed files with 31 additions and 18 deletions
+10 -5
View File
@@ -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,11 +145,15 @@ 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")
with SessionLocal() as db:
return _resolve_user_from_jwt(credentials.credentials, db)
if x_api_key is not None:
@@ -157,6 +161,7 @@ async def verify_auth(
_mark_auth_type(request, "admin_api_key")
return None
_mark_auth_type(request, "api_key")
with SessionLocal() as db:
return _resolve_user_from_api_key(x_api_key, db)
if AUTH_DISABLED:
@@ -173,11 +178,11 @@ 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"}:
with SessionLocal() as db:
default_user = _get_default_user(db)
if default_user is not None:
return default_user
@@ -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,6 +207,7 @@ async def require_admin(
auth_type = getattr(request.state, "auth_type", "none")
if user is None:
if auth_type in {"admin_api_key", "disabled"}:
with SessionLocal() as db:
default_user = _get_default_user(db)
if default_user is not None:
if default_user.role != "admin":
+17 -9
View File
@@ -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.")