From 9383e9a2556a533d289a8aa041c7a6660e806581 Mon Sep 17 00:00:00 2001 From: JainamShah-22 <80155538+JainamShah-22@users.noreply.github.com> Date: Mon, 20 Jul 2026 00:07:50 +0530 Subject: [PATCH] fix(server): scope auth DB sessions to prevent connection-pool exhaustion (#6237) --- server/auth.py | 23 ++++++++++++++--------- server/routers/auth.py | 26 +++++++++++++++++--------- 2 files changed, 31 insertions(+), 18 deletions(-) diff --git a/server/auth.py b/server/auth.py index f1f9f8513..ee4959b54 100644 --- a/server/auth.py +++ b/server/auth.py @@ -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.") diff --git a/server/routers/auth.py b/server/routers/auth.py index 0645a5292..86741ea01 100644 --- a/server/routers/auth.py +++ b/server/routers/auth.py @@ -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.")