"""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