phase: 110_fix_sse_db_pool_exhaustion
Build and Push Containers / build-and-push-app (push) Successful in 2m14s
Build and Push Containers / build-and-push-db (push) Successful in 13s

---

**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:
2026-09-14 15:55:13 -04:00
parent 35d65d2f25
commit 3a4035fc96
43 changed files with 2886 additions and 444 deletions
+23 -6
View File
@@ -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
+13 -2
View File
@@ -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,
+13 -2
View File
@@ -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.
+108
View File
@@ -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)
+2 -1
View File
@@ -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=[],