diff --git a/src/contract_check/api/app.py b/src/contract_check/api/app.py index e032238..5152b32 100644 --- a/src/contract_check/api/app.py +++ b/src/contract_check/api/app.py @@ -14,6 +14,7 @@ from fastapi import FastAPI from fastapi.responses import RedirectResponse from ..core.config import get_settings +from ..core.db.session import dispose_engine from ..core.llm import port as llm_port # noqa: F401 — package loaded from ..core.logging import bind_context, configure_logging, get_logger from ..core.metrics import redis_connected @@ -92,6 +93,7 @@ async def lifespan(app: FastAPI) -> Any: await publisher.close() await notification_publisher.close() + await dispose_engine() if redis_client is not None: try: await redis_client.aclose() diff --git a/src/contract_check/api/deps.py b/src/contract_check/api/deps.py index f974ad9..f721c29 100644 --- a/src/contract_check/api/deps.py +++ b/src/contract_check/api/deps.py @@ -17,7 +17,7 @@ from ..core.auth import AuthError, TokenExpiredError, TokenInvalidError, verify_ from ..core.auth_refresh import RefreshTokenStore from ..core.config import get_settings from ..core.db.models import PasskeyCredential, User -from ..core.db.session import create_session_factory +from ..core.db.session import get_session from ..core.logging import get_logger from ..core.mq.publisher import Publisher from ..core.notifications.publisher import NotificationPublisher @@ -31,8 +31,7 @@ log = get_logger(__name__) async def get_db_session() -> AsyncIterator[AsyncSession]: - factory = create_session_factory() - async with factory() as session: + async for session in get_session(): yield session @@ -513,38 +512,6 @@ async def require_current_user( CurrentUserDep = Annotated[CurrentUser, Depends(require_current_user)] -async def fetch_document_status( - session: AsyncSession, document_id: UUID, user_id: UUID -) -> dict[str, object] | None: - result = await session.execute( - text( - "SELECT d.id, d.status, d.stage, d.filename, d.created_at, " - " r.markdown, r.content_json, r.model_used, " - " r.prompt_tokens, r.eval_tokens, r.latency_ms " - "FROM documents d " - "LEFT JOIN reports r ON r.document_id = d.id " - "WHERE d.id = :d AND d.user_id = :u" - ), - {"d": document_id, "u": user_id}, - ) - row = result.first() - if row is None: - return None - return { - "id": row[0], - "status": row[1], - "stage": row[2], - "filename": row[3], - "created_at": row[4], - "markdown": row[5], - "content_json": row[6], - "model_used": row[7], - "prompt_tokens": row[8], - "eval_tokens": row[9], - "latency_ms": row[10], - } - - async def get_credits(session: AsyncSession, user_id: UUID) -> int: result = await session.execute( text("SELECT credits_left FROM users WHERE id = :u"), diff --git a/src/contract_check/api/routes/reports.py b/src/contract_check/api/routes/reports.py index 6d67db1..d858bac 100644 --- a/src/contract_check/api/routes/reports.py +++ b/src/contract_check/api/routes/reports.py @@ -6,7 +6,7 @@ from uuid import UUID from fastapi import APIRouter, HTTPException -from ..deps import AsyncSessionDep, CurrentUserDep, fetch_document_status +from ..deps import AsyncSessionDep, CurrentUserDep, fetch_document_status_for_user router = APIRouter(tags=["reports"]) @@ -17,7 +17,7 @@ async def get_report( document_id: UUID, user: CurrentUserDep, ) -> dict[str, object]: - row = await fetch_document_status(session, document_id, user.user_id) + row = await fetch_document_status_for_user(session, document_id, user.user_id) if row is None: raise HTTPException(status_code=404, detail="report not found") diff --git a/src/contract_check/core/db/session.py b/src/contract_check/core/db/session.py index 082daa9..ee4e9d7 100644 --- a/src/contract_check/core/db/session.py +++ b/src/contract_check/core/db/session.py @@ -19,6 +19,27 @@ from sqlalchemy.ext.asyncio import ( from ..config import get_settings +_engine: AsyncEngine | None = None +_session_factory: async_sessionmaker[AsyncSession] | None = None + + +def _get_engine() -> AsyncEngine: + """Get or create the singleton engine.""" + global _engine + if _engine is None: + settings = get_settings() + _engine = create_async_engine(settings.database_url, pool_pre_ping=True) + return _engine + + +def _get_session_factory() -> async_sessionmaker[AsyncSession]: + """Get or create the singleton session factory.""" + global _session_factory + if _session_factory is None: + engine = _get_engine() + _session_factory = async_sessionmaker(engine, expire_on_commit=False) + return _session_factory + def create_engine(url: str | None = None, **kwargs: Any) -> AsyncEngine: """Build an async engine from DATABASE_URL (or an explicit url).""" @@ -38,6 +59,15 @@ async def get_session( factory: async_sessionmaker[AsyncSession] | None = None, ) -> AsyncIterator[AsyncSession]: """FastAPI/yield-style session dependency.""" - factory = factory or create_session_factory() + factory = factory or _get_session_factory() async with factory() as session: yield session + + +async def dispose_engine() -> None: + """Dispose the singleton engine and reset the session factory.""" + global _engine, _session_factory + if _engine is not None: + await _engine.dispose() + _engine = None + _session_factory = None diff --git a/src/contract_check/core/llm/ollama_cloud.py b/src/contract_check/core/llm/ollama_cloud.py index 0873f53..0347b98 100644 --- a/src/contract_check/core/llm/ollama_cloud.py +++ b/src/contract_check/core/llm/ollama_cloud.py @@ -23,8 +23,13 @@ from ..analysis.analyzer import build_user_prompt, dedupe_findings, sort_finding from ..analysis.checklist import checklist_for_prompt from ..analysis.chunker import chunk_text from ..analysis.report_schema import ReportPayload -from ..errors import TerminalError from ..logging import get_logger +from .errors import ( + LLMConfigError, + LLMError, + LLMQuotaError, + LLMUnavailableError, +) from .port import AnalysisResult from .prescreen import PRESCREEN_MAX_CHARS_DEFAULT, PRESCREEN_SYSTEM, PrescreenExtraction @@ -57,22 +62,6 @@ SYSTEM_PROMPT = ( ) -class LLMError(Exception): - """Unrecoverable LLM failure (after all retries).""" - - -class LLMQuotaError(LLMError): - """429 / quota on both primary and fallback — refundable failure.""" - - -class LLMUnavailableError(LLMError): - """Ollama server unreachable or returns non-200 status — retryable.""" - - -class LLMConfigError(LLMError, TerminalError): - """Misconfigured Ollama host/model/endpoint — terminal, do not retry.""" - - class _QuotaSignal(Exception): """Internal: 429 triggers fallback within the same call.""" diff --git a/src/contract_check/core/llm/yandex_gpt.py b/src/contract_check/core/llm/yandex_gpt.py index ba1d83c..1ca91b9 100644 --- a/src/contract_check/core/llm/yandex_gpt.py +++ b/src/contract_check/core/llm/yandex_gpt.py @@ -28,8 +28,13 @@ from ..analysis.analyzer import build_user_prompt, dedupe_findings, sort_finding from ..analysis.checklist import checklist_for_prompt from ..analysis.chunker import chunk_markdown from ..analysis.report_schema import ReportPayload -from ..errors import TerminalError from ..logging import get_logger +from .errors import ( + LLMConfigError, + LLMError, + LLMQuotaError, + LLMUnavailableError, +) from .port import AnalysisResult from .prescreen import PRESCREEN_MAX_CHARS_DEFAULT, PRESCREEN_SYSTEM, PrescreenExtraction @@ -78,22 +83,6 @@ SYSTEM_PROMPT = ( ) -class LLMError(Exception): - """Unrecoverable LLM failure (after all retries).""" - - -class LLMQuotaError(LLMError): - """Quota/rate-limit on both primary and fallback — refundable failure.""" - - -class LLMUnavailableError(LLMError): - """Yandex API unreachable or returns non-200 status — retryable.""" - - -class LLMConfigError(LLMError, TerminalError): - """Misconfigured folder ID / API key / endpoint — terminal, do not retry.""" - - class _QuotaSignal(Exception): """Internal: 429 triggers fallback within the same call.""" diff --git a/src/contract_check/worker_analyze/handler.py b/src/contract_check/worker_analyze/handler.py index cfa2aab..e6abcf9 100644 --- a/src/contract_check/worker_analyze/handler.py +++ b/src/contract_check/worker_analyze/handler.py @@ -23,6 +23,12 @@ from ..core.analysis.checklist import checklist_for_prompt from ..core.config import get_settings from ..core.credits import refund_credit from ..core.db.enums import DOC_TERMINAL, FailureClass +from ..core.llm.errors import ( + LLMConfigError, + LLMError, + LLMQuotaError, + LLMUnavailableError, +) from ..core.llm.factory import build_llm_provider from ..core.llm.port import LLMProvider from ..core.logging import get_logger @@ -180,13 +186,6 @@ class AnalyzeHandler: ) def classify(self, exc: BaseException) -> FailureClass: - from ..core.llm.ollama_cloud import ( - LLMConfigError, - LLMError, - LLMQuotaError, - LLMUnavailableError, - ) - if isinstance(exc, LLMQuotaError): return "llm_quota" if isinstance(exc, LLMConfigError):