phase: 110_fix_sse_db_pool_exhaustion
--- **Phase 110 — Fix SSE DB Connection Pool Exhaustion (SEC-14-04): COMPLETE** **What was implemented/verified:** - All three tasks (pool config, short-lived sessions, concurrency cap) were already implemented in code - Fixed `tests/integration/test_chat_db_sessions.py` — added FakeChatLLM mock, fixed LLM signature (`tools=` not `_tools=`), used `fastapi_app.dependency_overrides` instead of `client.app.dependency_overrides` - Fixed `tests/e2e/test_chat_db_pool.py` — added FakeChatLLM mock, fixed admin password to match `tests/conftest.py`, removed unused imports - Fixed lint errors (unused imports, import order) in both test files **Test / lint / coverage results:** - `uv run pytest` → 2350 passed, 1 warning, 56.4s - `uv run pytest --cov=app --cov-report=term-missing` → 99% coverage (4065 lines, 16 uncovered) - `uv run pytest tests/e2e/test_chat_db_pool.py -v --no-cov` → 3 passed - `uv run pytest tests/integration/test_chat_db_sessions.py -v --no-cov` → 4 passed - `uv run pytest tests/integration/test_chat_concurrency.py -v --no-cov` → 11 passed - `uv run pytest tests/unit/test_db_pool_config.py -v --no-cov` → 14 passed - `uv run pytest tests/unit/test_agent_short_lived_sessions.py -v --no-cov` → 7 passed - `uv run ruff check .` → all checks passed - `uv run pyright` → 0 errors, 0 warnings **Completion criteria:** - [✓] `app/db.py::create_engine` receives explicit `pool_size=5`, `max_overflow=10`, `pool_recycle=3600` from settings - [✓] `run_agent` accepts `db_factory: Callable[[], Session]` and creates short-lived sessions per tool call - [✓] Each tool round uses a separate DB session closed after the tool result - [✓] Concurrency cap (`BOR_CHAT_MAX_CONCURRENT`, default 10) limits concurrent turns; excess get 503 - [✓] All test gates green, coverage 99%, lint/types clean **Notable decisions:** Tests needed LLM mocking (the original test files lacked `FakeChatLLM` mocks, causing hangs on real LLM calls). **Next pending phase:** None — this is the last phase in `todo/`.
This commit is contained in:
@@ -35,10 +35,11 @@ import asyncio
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Sequence
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Sequence
|
||||
from copy import deepcopy
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -133,21 +134,37 @@ class ScriptedLLM:
|
||||
yield StreamPiece("content", tail)
|
||||
|
||||
|
||||
def _mock_session() -> Session:
|
||||
"""A minimal mock Session for unit tests (DB accessors are monkeypatched).
|
||||
|
||||
The mock works as a context manager: ``__enter__`` returns itself so
|
||||
``with db_factory() as tool_db:`` binds *tool_db* to the same mock.
|
||||
"""
|
||||
mock = cast("Session", MagicMock())
|
||||
mock.__enter__ = MagicMock(return_value=mock)
|
||||
mock.__exit__ = MagicMock(return_value=False)
|
||||
mock.scalar = MagicMock(return_value=None)
|
||||
mock.execute = MagicMock(return_value=MagicMock(scalars=MagicMock(return_value=[])))
|
||||
return mock
|
||||
|
||||
|
||||
async def _run(
|
||||
llm: ScriptedLLM | FailingLLM,
|
||||
holder: AgentHolder,
|
||||
settings: Settings,
|
||||
seed_docs: list[Document] | None = None,
|
||||
history: Sequence[dict[str, Any]] = (),
|
||||
db_factory: Callable[[], Session] | None = None,
|
||||
) -> list[StreamPiece | ToolCallPiece | RetryPiece | ToolResultPiece]:
|
||||
"""Consume one ``run_agent`` turn; *history* (phase 74) is the
|
||||
client's prior turns spliced between system and user (default
|
||||
``()`` — the pre-phase-74 two-message request). Phase 95: the loop
|
||||
may also yield a ``ToolResultPiece`` (a truncated ``read``)."""
|
||||
out: list[StreamPiece | ToolCallPiece | RetryPiece | ToolResultPiece] = []
|
||||
factory = db_factory or (lambda: _mock_session())
|
||||
async for piece in run_agent(
|
||||
cast("LLMClient", llm),
|
||||
cast("Session", None),
|
||||
factory,
|
||||
system_prompt="SYSTEM_PROMPT",
|
||||
user_message="QUESTION",
|
||||
seed_docs=seed_docs or [],
|
||||
@@ -693,7 +710,7 @@ def test_ls_nested_folder_scope_lists_one_level_deeper(
|
||||
root), counted; the fetchers are the source-scoped ones."""
|
||||
|
||||
def _rows(db: Any, source: str) -> list[tuple[str, str, str]]:
|
||||
assert (source, db) == ("Homelab", None)
|
||||
assert source == "Homelab" # db is a mock session (SEC-14-04)
|
||||
return [
|
||||
("networking/lan.md", "LAN", "2024-06-15"),
|
||||
("networking/vpn.md", "VPN", "2024-06-15"),
|
||||
@@ -2552,7 +2569,7 @@ def test_round_failure_after_first_piece_is_terminal(monkeypatch: pytest.MonkeyP
|
||||
with pytest.raises(LLMError, match="mid-stream drop"):
|
||||
async for piece in run_agent(
|
||||
cast("LLMClient", llm),
|
||||
cast("Session", None),
|
||||
lambda: _mock_session(),
|
||||
system_prompt="SYSTEM_PROMPT",
|
||||
user_message="QUESTION",
|
||||
seed_docs=[],
|
||||
@@ -2614,7 +2631,7 @@ def test_zero_retries_is_one_plain_attempt(monkeypatch: pytest.MonkeyPatch) -> N
|
||||
with pytest.raises(LLMError, match="connection refused"):
|
||||
async for piece in run_agent(
|
||||
cast("LLMClient", llm),
|
||||
cast("Session", None),
|
||||
lambda: _mock_session(),
|
||||
system_prompt="SYSTEM_PROMPT",
|
||||
user_message="QUESTION",
|
||||
seed_docs=[],
|
||||
@@ -2651,7 +2668,7 @@ def test_abandon_mid_retry_sleep_leaks_nothing(monkeypatch: pytest.MonkeyPatch)
|
||||
async def run() -> None:
|
||||
gen = run_agent(
|
||||
cast("LLMClient", llm),
|
||||
cast("Session", None),
|
||||
lambda: _mock_session(),
|
||||
system_prompt="SYSTEM_PROMPT",
|
||||
user_message="QUESTION",
|
||||
seed_docs=[],
|
||||
|
||||
@@ -0,0 +1,326 @@
|
||||
"""Unit: short-lived DB sessions in the agent loop (SEC-14-04).
|
||||
|
||||
Verifies that ``run_agent`` uses a session factory to create a new
|
||||
short-lived session for each DB operation (tool call), closes it after
|
||||
the tool result is produced, and completes the agent loop correctly.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from copy import deepcopy
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any, cast
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import Settings
|
||||
from app.models import Document
|
||||
from app.rag import agent
|
||||
from app.rag.agent import AgentHolder, run_agent
|
||||
from app.rag.llm import (
|
||||
LLMClient,
|
||||
RetryPiece,
|
||||
StreamPiece,
|
||||
ToolCallPiece,
|
||||
ToolResultPiece,
|
||||
)
|
||||
|
||||
#: Fixture document creation date (phase 106, D5).
|
||||
_FIXTURE_CREATED_AT = datetime(2024, 6, 15, 12, 0, 0, tzinfo=UTC)
|
||||
|
||||
|
||||
def _settings(**kwargs: Any) -> Settings:
|
||||
kwargs.setdefault("_env_file", None)
|
||||
return Settings(**kwargs) # pyright: ignore[reportCallIssue]
|
||||
|
||||
|
||||
def _doc(source: str, path: str, title: str = "Title", content: str = "CONTENT") -> Document:
|
||||
return Document(
|
||||
id=uuid.uuid4(),
|
||||
source=source,
|
||||
path=path,
|
||||
full_path=f"/tmp/{path}",
|
||||
title=title,
|
||||
content=content,
|
||||
content_hash="0" * 64,
|
||||
created_at=_FIXTURE_CREATED_AT,
|
||||
)
|
||||
|
||||
|
||||
class TrackingSession:
|
||||
"""A mock Session that tracks ``close()`` calls and is usable as a
|
||||
context manager."""
|
||||
|
||||
def __init__(self, close_count: list[int] | None = None) -> None:
|
||||
self.close_count = close_count if close_count is not None else [0]
|
||||
self.scalar = MagicMock(return_value=None)
|
||||
# scalars() returns a ScalarResult-like object with .all()
|
||||
self._scalar_result = MagicMock()
|
||||
self._scalar_result.all = MagicMock(return_value=[])
|
||||
self.scalars = MagicMock(return_value=self._scalar_result)
|
||||
self.execute = MagicMock(return_value=MagicMock(scalars=MagicMock(return_value=[])))
|
||||
self.add = MagicMock()
|
||||
|
||||
def __enter__(self) -> TrackingSession:
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: Any) -> None:
|
||||
self.close_count[0] += 1
|
||||
|
||||
|
||||
class TrackingFactory:
|
||||
"""A session factory that creates a ``TrackingSession`` each time
|
||||
it is called, so the tests can verify that a new session is created
|
||||
for each tool call and that it is closed afterwards."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.sessions: list[TrackingSession] = []
|
||||
|
||||
def __call__(self) -> TrackingSession:
|
||||
session = TrackingSession()
|
||||
self.sessions.append(session)
|
||||
return session
|
||||
|
||||
|
||||
# Type alias for pyright: TrackingFactory is callable that returns Session
|
||||
TrackingFactoryCallable: type[TrackingFactory] = TrackingFactory # noqa: N816
|
||||
|
||||
|
||||
class ScriptedLLM:
|
||||
"""Canned stream sequences for agent-loop tests."""
|
||||
|
||||
def __init__(self, *streams: list[StreamPiece | ToolCallPiece]) -> None:
|
||||
self.streams: list[list[StreamPiece | ToolCallPiece]] = list(streams)
|
||||
self.requests: list[tuple[list[dict[str, Any]], list[dict[str, Any]] | None]] = []
|
||||
|
||||
async def chat_stream(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
scaffolding: Any = None,
|
||||
) -> AsyncIterator[StreamPiece | ToolCallPiece]:
|
||||
self.requests.append((deepcopy(messages), tools))
|
||||
if not self.streams:
|
||||
raise AssertionError("ScriptedLLM ran out of canned streams")
|
||||
pieces = self.streams.pop(0)
|
||||
for piece in pieces:
|
||||
yield piece
|
||||
|
||||
|
||||
class ScriptedToolLLM(ScriptedLLM):
|
||||
"""A scripted LLM that returns exactly one tool call, then an
|
||||
empty clean answer on the next round."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tool_call: ToolCallPiece,
|
||||
answer: str = "ANSWER",
|
||||
) -> None:
|
||||
super().__init__(
|
||||
[tool_call], # round 1: tool call
|
||||
[StreamPiece("content", answer)], # round 2: answer
|
||||
)
|
||||
|
||||
|
||||
async def _consume(
|
||||
llm: LLMClient,
|
||||
db_factory: Callable[[], Session],
|
||||
holder: AgentHolder,
|
||||
settings: Settings,
|
||||
seed_docs: list[Document] | None = None,
|
||||
) -> list[StreamPiece | ToolCallPiece | RetryPiece | ToolResultPiece]:
|
||||
"""Consume one ``run_agent`` turn."""
|
||||
out: list[StreamPiece | ToolCallPiece | RetryPiece | ToolResultPiece] = []
|
||||
async for piece in run_agent(
|
||||
cast("LLMClient", llm),
|
||||
db_factory,
|
||||
system_prompt="SYSTEM_PROMPT",
|
||||
user_message="QUESTION",
|
||||
seed_docs=seed_docs or [],
|
||||
settings=settings,
|
||||
holder=holder,
|
||||
):
|
||||
out.append(piece)
|
||||
return out
|
||||
|
||||
|
||||
# ---------- db_factory is called per DB operation ----------
|
||||
|
||||
|
||||
def test_db_factory_called_once_for_answer_no_tools() -> None:
|
||||
"""When the model answers without calling any tools, the agent
|
||||
loop makes no DB calls — but ``run_agent`` still accepts the
|
||||
factory (it is simply not invoked)."""
|
||||
llm = ScriptedLLM([StreamPiece("content", "DIRECT ANSWER")])
|
||||
factory = TrackingFactory()
|
||||
holder = AgentHolder()
|
||||
settings = _settings(agent_max_rounds=10)
|
||||
|
||||
asyncio.run(
|
||||
_consume(cast("LLMClient", llm), cast("Callable[[], Session]", factory), holder, settings)
|
||||
)
|
||||
|
||||
# No tool calls means no DB operations — factory never invoked
|
||||
assert len(factory.sessions) == 0
|
||||
assert holder.tool_calls == 0
|
||||
|
||||
|
||||
# ---------- each tool call creates a new session ----------
|
||||
|
||||
|
||||
def test_each_tool_call_creates_a_new_session() -> None:
|
||||
"""Each tool call the model emits creates its own short-lived
|
||||
session via the factory; sessions are closed after the tool
|
||||
result is produced."""
|
||||
tool_call = ToolCallPiece(id="call_1", name="ls", arguments={})
|
||||
llm = ScriptedToolLLM(tool_call)
|
||||
factory = TrackingFactory()
|
||||
holder = AgentHolder()
|
||||
settings = _settings(agent_max_rounds=10)
|
||||
|
||||
asyncio.run(
|
||||
_consume(cast("LLMClient", llm), cast("Callable[[], Session]", factory), holder, settings)
|
||||
)
|
||||
|
||||
# One tool call → one session created
|
||||
assert len(factory.sessions) == 1
|
||||
# The session was closed after the tool result
|
||||
assert factory.sessions[0].close_count[0] == 1
|
||||
# The holder records the executed call
|
||||
assert holder.tool_calls == 1
|
||||
|
||||
|
||||
def test_multiple_tool_calls_create_separate_sessions() -> None:
|
||||
"""When the model makes multiple tool calls across rounds, each
|
||||
round creates a new session that is closed after the result."""
|
||||
tool_call_1 = ToolCallPiece(id="call_1", name="ls", arguments={})
|
||||
tool_call_2 = ToolCallPiece(id="call_2", name="ls", arguments={})
|
||||
llm = ScriptedToolLLM(tool_call_1)
|
||||
# Override the second round to also emit a tool call
|
||||
llm.streams = [
|
||||
[tool_call_1], # round 1: ls
|
||||
[tool_call_2], # round 2: ls (another listing)
|
||||
[StreamPiece("content", "FINAL ANSWER")], # round 3: answer
|
||||
]
|
||||
factory = TrackingFactory()
|
||||
holder = AgentHolder()
|
||||
settings = _settings(agent_max_rounds=10)
|
||||
|
||||
asyncio.run(
|
||||
_consume(cast("LLMClient", llm), cast("Callable[[], Session]", factory), holder, settings)
|
||||
)
|
||||
|
||||
# Two tool calls → two sessions created and closed
|
||||
assert len(factory.sessions) == 2
|
||||
for session in factory.sessions:
|
||||
assert session.close_count[0] == 1
|
||||
assert holder.tool_calls == 2
|
||||
|
||||
|
||||
# ---------- sessions are closed after use (no pinning) ----------
|
||||
|
||||
|
||||
def test_sessions_closed_after_tool_result() -> None:
|
||||
"""Verify that the session is closed AFTER the tool result is
|
||||
produced but BEFORE the next model round — no session is held
|
||||
across rounds."""
|
||||
tool_call = ToolCallPiece(id="call_1", name="ls", arguments={})
|
||||
llm = ScriptedToolLLM(tool_call)
|
||||
factory = TrackingFactory()
|
||||
holder = AgentHolder()
|
||||
settings = _settings(agent_max_rounds=10)
|
||||
|
||||
asyncio.run(
|
||||
_consume(cast("LLMClient", llm), cast("Callable[[], Session]", factory), holder, settings)
|
||||
)
|
||||
|
||||
# The session was closed (close_count incremented)
|
||||
assert factory.sessions[0].close_count[0] == 1
|
||||
# Only one session was created (not reused across rounds)
|
||||
assert len(factory.sessions) == 1
|
||||
|
||||
|
||||
# ---------- deflected path works without DB factory usage ----------
|
||||
|
||||
|
||||
def test_deflected_path_no_db_factory_calls() -> None:
|
||||
"""A deflected turn (LOW mode) does not run the agent loop, so
|
||||
the session factory is never invoked."""
|
||||
llm = ScriptedLLM([StreamPiece("content", "I don't know about that.")])
|
||||
factory = TrackingFactory()
|
||||
holder = AgentHolder()
|
||||
# agent_max_rounds=0 disables tools → single tools=None request
|
||||
settings = _settings(agent_max_rounds=0)
|
||||
|
||||
asyncio.run(
|
||||
_consume(cast("LLMClient", llm), cast("Callable[[], Session]", factory), holder, settings)
|
||||
)
|
||||
|
||||
assert len(factory.sessions) == 0
|
||||
assert holder.tool_calls == 0
|
||||
|
||||
|
||||
# ---------- agent loop completes correctly ----------
|
||||
|
||||
|
||||
def test_agent_loop_completes_with_mock_factory() -> None:
|
||||
"""The agent loop completes correctly with a mock session factory:
|
||||
tool calls are executed, the answer is streamed, and the holder
|
||||
records the correct state."""
|
||||
tool_call = ToolCallPiece(id="call_1", name="grep", arguments={"pattern": "hello"})
|
||||
llm = ScriptedToolLLM(tool_call)
|
||||
factory = TrackingFactory()
|
||||
holder = AgentHolder()
|
||||
settings = _settings(agent_max_rounds=10)
|
||||
|
||||
pieces = asyncio.run(
|
||||
_consume(cast("LLMClient", llm), cast("Callable[[], Session]", factory), holder, settings)
|
||||
)
|
||||
|
||||
# The stream contains the tool call piece and the answer piece
|
||||
types = [getattr(p, "kind", "tool") for p in pieces]
|
||||
assert "tool" in types
|
||||
assert "content" in types
|
||||
|
||||
# The holder records one executed call
|
||||
assert holder.tool_calls == 1
|
||||
assert len(factory.sessions) == 1
|
||||
|
||||
|
||||
# ---------- monkeypatched DB accessors work with factory ----------
|
||||
|
||||
|
||||
def test_monkeypatched_accessors_with_factory() -> None:
|
||||
"""DB accessors that are monkeypatched (as in the existing unit
|
||||
test suite) work correctly when the agent loop calls them through
|
||||
a session factory."""
|
||||
# Patch ls_top to return a canned result
|
||||
canned_ls = [("TestSource", 5, None)]
|
||||
|
||||
def fake_ls_top(db: Session) -> list[tuple[str, int, str | None]]: # type: ignore[return-value]
|
||||
return list(canned_ls)
|
||||
|
||||
with (
|
||||
patch.object(agent, "ls_top", fake_ls_top),
|
||||
patch.object(agent, "list_source_names", return_value=["TestSource"]),
|
||||
):
|
||||
tool_call = ToolCallPiece(id="call_1", name="ls", arguments={})
|
||||
llm = ScriptedToolLLM(tool_call)
|
||||
factory = TrackingFactory()
|
||||
holder = AgentHolder()
|
||||
settings = _settings(agent_max_rounds=10)
|
||||
|
||||
asyncio.run(
|
||||
_consume(cast("LLMClient", llm), cast("Callable[[], Session]", factory), holder, settings)
|
||||
)
|
||||
|
||||
assert holder.tool_calls == 1
|
||||
assert len(factory.sessions) == 1
|
||||
# The factory was called exactly once for this tool round
|
||||
assert factory.sessions[0].close_count[0] == 1
|
||||
|
||||
|
||||
|
||||
@@ -163,6 +163,12 @@ class _FakeSession:
|
||||
self.added: list[Any] = []
|
||||
self.commits = 0
|
||||
|
||||
def __enter__(self) -> _FakeSession:
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: Any) -> None:
|
||||
pass
|
||||
|
||||
def add(self, obj: Any) -> None:
|
||||
self.added.append(obj)
|
||||
|
||||
@@ -181,10 +187,15 @@ class _FakeSession:
|
||||
@pytest.fixture()
|
||||
def env(monkeypatch: pytest.MonkeyPatch) -> Iterator[_FakeSession]:
|
||||
"""``POST /api/chat`` with the DB session, retriever settings, and
|
||||
availability faked (the gate tests' wiring)."""
|
||||
availability faked (the gate tests' wiring).
|
||||
|
||||
SEC-14-04: the chat endpoint uses short-lived sessions via
|
||||
``SessionLocal()`` — we monkeypatch ``chat_api.SessionLocal`` to
|
||||
return a fake session instead of overriding ``get_db``.
|
||||
"""
|
||||
monkeypatch.setattr(chat_api, "db_available", lambda: True)
|
||||
session = _FakeSession()
|
||||
monkeypatch.setitem(fastapi_app.dependency_overrides, chat_api.get_db, lambda: session)
|
||||
monkeypatch.setattr(chat_api, "SessionLocal", lambda: session)
|
||||
# A stable gate threshold, independent of the production default.
|
||||
monkeypatch.setattr(
|
||||
chat_api,
|
||||
|
||||
@@ -483,6 +483,12 @@ class _FakeSession:
|
||||
self.commits = 0
|
||||
self.kb_overview = kb_overview
|
||||
|
||||
def __enter__(self) -> _FakeSession:
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: Any) -> None:
|
||||
pass
|
||||
|
||||
def add(self, obj: Any) -> None:
|
||||
self.added.append(obj)
|
||||
|
||||
@@ -511,11 +517,16 @@ def _admin_signed_in(client: TestClient) -> None:
|
||||
|
||||
@pytest.fixture()
|
||||
def gate_env(monkeypatch: pytest.MonkeyPatch) -> Iterator[tuple[_FakeSession, _CannedLLM]]:
|
||||
"""``POST /api/chat`` with retriever, session, and LLM all faked."""
|
||||
"""``POST /api/chat`` with retriever, session, and LLM all faked.
|
||||
|
||||
SEC-14-04: the chat endpoint uses short-lived sessions via
|
||||
``SessionLocal()`` — monkeypatch ``chat_api.SessionLocal`` instead
|
||||
of overriding ``get_db``.
|
||||
"""
|
||||
monkeypatch.setattr(chat_api, "db_available", lambda: True)
|
||||
session = _FakeSession()
|
||||
llm = _CannedLLM()
|
||||
monkeypatch.setitem(fastapi_app.dependency_overrides, chat_api.get_db, lambda: session)
|
||||
monkeypatch.setattr(chat_api, "SessionLocal", lambda: session)
|
||||
monkeypatch.setitem(fastapi_app.dependency_overrides, chat_api.get_llm, lambda: llm)
|
||||
# These tests assert against a specific gate threshold; keep it stable
|
||||
# regardless of the production default (0.62) or any .env.
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
"""Unit tests for DB pool configuration (SEC-14-04, phase 106, task 01).
|
||||
|
||||
Verifies that:
|
||||
- Settings expose db_pool_size, db_pool_max_overflow, db_pool_recycle
|
||||
with correct defaults and validators.
|
||||
- create_engine() receives the pool kwargs from settings.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
|
||||
import pytest
|
||||
|
||||
from app.config import Settings
|
||||
|
||||
|
||||
class TestSettingsDefaults:
|
||||
"""Pool config defaults match SQLAlchemy implicit defaults."""
|
||||
|
||||
def test_pool_size_default(self):
|
||||
assert Settings().db_pool_size == 5
|
||||
|
||||
def test_pool_max_overflow_default(self):
|
||||
assert Settings().db_pool_max_overflow == 10
|
||||
|
||||
def test_pool_recycle_default(self):
|
||||
assert Settings().db_pool_recycle == 3600
|
||||
|
||||
|
||||
class TestSettingsCustomValues:
|
||||
"""Custom values round-trip correctly."""
|
||||
|
||||
def test_custom_all_three(self):
|
||||
s = Settings(
|
||||
db_pool_size=10,
|
||||
db_pool_max_overflow=20,
|
||||
db_pool_recycle=1800,
|
||||
)
|
||||
assert s.db_pool_size == 10
|
||||
assert s.db_pool_max_overflow == 20
|
||||
assert s.db_pool_recycle == 1800
|
||||
|
||||
def test_custom_pool_size_only(self):
|
||||
s = Settings(db_pool_size=8)
|
||||
assert s.db_pool_size == 8
|
||||
assert s.db_pool_max_overflow == 10
|
||||
assert s.db_pool_recycle == 3600
|
||||
|
||||
|
||||
class TestValidators:
|
||||
"""Pool config validators reject invalid values."""
|
||||
|
||||
def test_pool_size_zero_raises(self):
|
||||
with pytest.raises(ValueError, match="db_pool_size must be >= 1"):
|
||||
Settings(db_pool_size=0)
|
||||
|
||||
def test_pool_size_negative_raises(self):
|
||||
with pytest.raises(ValueError, match="db_pool_size must be >= 1"):
|
||||
Settings(db_pool_size=-5)
|
||||
|
||||
def test_pool_max_overflow_negative_raises(self):
|
||||
with pytest.raises(ValueError, match="db_pool_max_overflow must be >= 0"):
|
||||
Settings(db_pool_max_overflow=-1)
|
||||
|
||||
def test_pool_max_overflow_zero_is_legal(self):
|
||||
s = Settings(db_pool_max_overflow=0)
|
||||
assert s.db_pool_max_overflow == 0
|
||||
|
||||
def test_pool_recycle_zero_is_legal(self):
|
||||
"""pool_recycle=0 means never recycle — legal, just aggressive."""
|
||||
s = Settings(db_pool_recycle=0)
|
||||
assert s.db_pool_recycle == 0
|
||||
|
||||
|
||||
class TestEngineKwargs:
|
||||
"""create_engine() receives the correct pool parameters from settings."""
|
||||
|
||||
def test_create_engine_pool_pre_ping_true(self):
|
||||
"""pool_pre_ping must remain True (connection health check)."""
|
||||
from app import db # noqa: F811
|
||||
|
||||
# The engine's pool options include pool_pre_ping=True.
|
||||
# We verify by checking the pool's _pre_ping attribute.
|
||||
assert db.engine.pool._pre_ping is True
|
||||
|
||||
def test_create_engine_pool_recycle(self):
|
||||
"""pool_recycle defaults to 3600 seconds."""
|
||||
from app import db # noqa: F811
|
||||
|
||||
assert db.engine.pool._recycle == 3600
|
||||
|
||||
def test_sessionlocal_still_callable(self):
|
||||
"""SessionLocal remains a valid session factory."""
|
||||
from app import db # noqa: F811
|
||||
|
||||
assert callable(db.SessionLocal)
|
||||
|
||||
def test_get_db_still_yields_session(self):
|
||||
"""get_db() dependency still yields a Session (contract preserved)."""
|
||||
from app import db # noqa: F811
|
||||
|
||||
gen = db.get_db()
|
||||
session = next(gen)
|
||||
assert isinstance(session, db.Session)
|
||||
session.close()
|
||||
# Generator cleanup
|
||||
with contextlib.suppress(StopIteration):
|
||||
next(gen)
|
||||
@@ -25,6 +25,7 @@ import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any, cast
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -102,7 +103,7 @@ async def _run(
|
||||
) -> None:
|
||||
async for _piece in run_agent(
|
||||
cast("LLMClient", llm),
|
||||
cast("Session", None),
|
||||
lambda: cast("Session", MagicMock()), # SEC-14-04: session factory
|
||||
system_prompt="SYSTEM_PROMPT",
|
||||
user_message="QUESTION",
|
||||
seed_docs=[],
|
||||
|
||||
Reference in New Issue
Block a user