Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
184 changes: 179 additions & 5 deletions apps/api/tests/contract/test_retrieval_contract.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,22 @@
from collections.abc import Callable, Coroutine, Sequence
from contextlib import AbstractAsyncContextManager
import asyncio
from collections.abc import AsyncGenerator, Callable, Coroutine, Sequence
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from typing import Any, cast
from uuid import uuid4

import pytest
from httpx import AsyncClient
from pytest import MonkeyPatch
from sqlalchemy.ext.asyncio import AsyncSession

from tests.support.contract_database import ContractDatabase

from shared.services.retrieval.execution import routes as retrieval_routes
from shared.services.retrieval.execution.reference_resolver import (
ResolvedWorkflowReferences,
)
from shared.services.retrieval.execution.route_types import RetrievalRouteContext

LLMFnInput = str | Sequence[dict[str, Any]]
LLMFn = Callable[[LLMFnInput], Coroutine[Any, Any, str]]

Expand Down Expand Up @@ -529,9 +537,9 @@ async def test_agent_explore_retrieval_should_return_seeded_chunk_via_fake_episo

assert response_json["router_used"] == "agent_explore"
assert response_json["stop_reason"] == "finished"
assert isinstance(response_json.get("decision_trace"), list)
assert response_json["decision_trace"]
assert response_json["decision_trace"][-1]["phase"] == "finish"
decision_trace = cast(list[dict[str, object]], response_json["decision_trace"])
assert decision_trace
assert decision_trace[-1]["phase"] == "finish"
assert any(
ref.get("chunk_id") == target["chunk_id"]
and ref.get("document_id") == target["document_id"]
Expand Down Expand Up @@ -789,3 +797,169 @@ async def test_agent_explore_should_preserve_same_chunk_id_across_documents(
}
assert {ref["chunk_id"] for ref in referenced_chunks} == {shared_chunk_id}
assert len(results) == 2


@pytest.mark.asyncio
async def test_agent_explore_should_release_route_session_before_final_hydration(
monkeypatch: MonkeyPatch,
) -> None:
from shared.services.retrieval.agent_explore.types import AgentStep, EpisodeResult

events: list[str] = []
route_db = _RecordingRouteSession(events)
final_db = object()

@asynccontextmanager
async def fake_open_fresh_database_context() -> AsyncGenerator[AsyncSession, None]:
events.append("final_db_open")
try:
yield cast(AsyncSession, final_db)
finally:
events.append("final_db_close")

class FakeHarness:
async def run_episode(self, **_kwargs: object) -> EpisodeResult:
events.append("episode")
assert events[:2] == ["route_rollback", "episode"]
# Model an episode exceeding a simulated idle-in-transaction
# timeout without making the contract suite wait 60 seconds.
simulated_idle_timeout_seconds = 0.001
await asyncio.sleep(simulated_idle_timeout_seconds + 0.01)
return EpisodeResult(
refs=[
{
"document_id": "doc_contract",
"section_path": "contract/section",
}
],
notes="",
steps=[
AgentStep(
step_index=0,
tool_name="finish",
tool_args={},
observation_text="done",
error=None,
elapsed_ms=1,
tokens_used_delta=1,
tokens_used_total=1,
)
],
)

async def fake_resolve_finish_refs(
db: AsyncSession,
**_kwargs: object,
) -> list[dict[str, str]]:
assert db is final_db
assert _kwargs["user_id"] == "contract-user"
assert _kwargs["namespace"] == "contract-namespace"
events.append("resolve_finish_refs")
return [{"document_id": "doc_contract", "chunk_id": "chunk_contract"}]

async def fake_resolve_workflow_references(
*,
db: AsyncSession,
**_kwargs: object,
) -> ResolvedWorkflowReferences:
assert db is final_db
assert _kwargs["user_id"] == "contract-user"
assert _kwargs["namespace"] == "contract-namespace"
assert _kwargs["revision_pins"] is None
events.append("resolve_references")
row = {
"document_id": "doc_contract",
"chunk_id": "chunk_contract",
"source_file_name": "contract.pdf",
"section_path": "contract/section",
"chunk_type": "text",
"content": "contract content",
}
return ResolvedWorkflowReferences(refs=[row], rows=[row])

async def fake_assemble_retrieval_results(
*,
db: AsyncSession,
rows: list[dict[str, Any]],
**_kwargs: object,
) -> list[dict[str, Any]]:
assert db is final_db
assert _kwargs["exclude_document_ids"] == []
assert _kwargs["exclude_sections"] == []
assert _kwargs["allowed_chunk_types"] is None
assert _kwargs["revision_pins"] is None
events.append("assemble_results")
return rows

class FakeTraceRecorder:
def __init__(self, db: AsyncSession, **_kwargs: object) -> None:
assert db is final_db
events.append("trace_init")

async def create_run(self) -> None:
events.append("trace_create")

def record_decision_trace_step(self, _step: object) -> None:
events.append("trace_step")

async def complete(self, *_args: object, **_kwargs: object) -> None:
events.append("trace_complete")

monkeypatch.setattr(retrieval_routes, "open_fresh_database_context", fake_open_fresh_database_context)
monkeypatch.setattr(
"shared.services.retrieval.agent_explore.harness.resolve_harness",
lambda: FakeHarness(),
)
monkeypatch.setattr(
"shared.services.retrieval.agent_explore.ref_resolution.resolve_finish_refs",
fake_resolve_finish_refs,
)
monkeypatch.setattr(retrieval_routes, "resolve_workflow_references", fake_resolve_workflow_references)
monkeypatch.setattr(retrieval_routes, "assemble_retrieval_results", fake_assemble_retrieval_results)
monkeypatch.setattr("shared.services.retrieval.trace.TraceRecorder", FakeTraceRecorder)

outcome = await retrieval_routes._run_agent_explore_route(
RetrievalRouteContext(
db=cast(AsyncSession, route_db),
user_id="contract-user",
namespace="contract-namespace",
query="session lifetime",
top_k=1,
exclude_document_ids=[],
exclude_sections=[],
allowed_chunk_types=None,
chunk_types=None,
signal_paths=None,
filter_mode="delete",
channels=None,
channel_weights=None,
rerank=False,
threshold=0.0,
internal_recall_k=None,
effective_recall_k=3,
use_agentic=True,
)
)

assert events == [
"route_rollback",
"episode",
"final_db_open",
"resolve_finish_refs",
"resolve_references",
"assemble_results",
"trace_init",
"trace_create",
"trace_step",
"trace_complete",
"final_db_close",
]
assert outcome.response["router_used"] == "agent_explore"


class _RecordingRouteSession:
def __init__(self, events: list[str]) -> None:
self._events = events

async def rollback(self) -> None:
self._events.append("route_rollback")
100 changes: 53 additions & 47 deletions packages/shared-python/shared/services/retrieval/execution/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,6 +199,12 @@ async def _run_agent_explore_route(
)
from shared.services.retrieval.trace import TraceRecorder

# End any read transaction created by the route's pre-episode work before
# handing control to the external agent/LLM. The episode can outlive
# PostgreSQL's idle-in-transaction timeout, so the request session must not
# be reused for post-episode database work.
await context.db.rollback()

harness = resolve_harness()
episode_started = time.perf_counter()
episode = await harness.run_episode(
Expand All @@ -222,55 +228,56 @@ async def _run_agent_explore_route(
# episode.refs are document_id + section_path (what the agent actually
# sees in tool text); resolve_workflow_references requires chunk_id —
# see ref_resolution.py's module docstring for why this bridge exists.
chunk_refs = await resolve_finish_refs(
context.db,
user_id=context.user_id,
namespace=context.namespace,
refs=episode.refs,
)
resolved = await resolve_workflow_references(
db=context.db,
user_id=context.user_id,
namespace=context.namespace,
refs=chunk_refs,
revision_pins=context.revision_pins,
)
assembled_rows = await assemble_retrieval_results(
db=context.db,
rows=resolved.rows,
exclude_document_ids=context.exclude_document_ids,
exclude_sections=context.exclude_sections,
allowed_chunk_types=context.allowed_chunk_types,
revision_pins=context.revision_pins,
)

decision_steps = build_decision_trace(episode.steps)
decision_trace = [step.to_dict() for step in decision_steps]
selected_doc_ids = list(
{row.get("document_id", "") for row in resolved.rows if row.get("document_id")}
)

trace = TraceRecorder(
context.db,
user_id=context.user_id,
namespace=context.namespace,
query=context.query,
top_k=context.top_k,
chunk_types=context.allowed_chunk_types,
policy_name="agent_explore_v1",
)
await trace.create_run()
for step in decision_steps:
trace.record_decision_trace_step(step)
if episode.stop_reason.startswith("budget_"):
trace.record_budget_stop(episode.stop_reason.removeprefix("budget_"))
await trace.complete(
assembled_rows,
"agent_explore",
token_count=episode.tokens_used,
model_name=episode.model_name,
selected_doc_ids=selected_doc_ids,
)
async with open_fresh_database_context() as final_db:
chunk_refs = await resolve_finish_refs(
final_db,
user_id=context.user_id,
namespace=context.namespace,
refs=episode.refs,
)
resolved = await resolve_workflow_references(
db=final_db,
user_id=context.user_id,
namespace=context.namespace,
refs=chunk_refs,
revision_pins=context.revision_pins,
)
assembled_rows = await assemble_retrieval_results(
db=final_db,
rows=resolved.rows,
exclude_document_ids=context.exclude_document_ids,
exclude_sections=context.exclude_sections,
allowed_chunk_types=context.allowed_chunk_types,
revision_pins=context.revision_pins,
)

selected_doc_ids = list(
{row.get("document_id", "") for row in resolved.rows if row.get("document_id")}
)
trace = TraceRecorder(
final_db,
user_id=context.user_id,
namespace=context.namespace,
query=context.query,
top_k=context.top_k,
chunk_types=context.allowed_chunk_types,
policy_name="agent_explore_v1",
)
await trace.create_run()
for step in decision_steps:
trace.record_decision_trace_step(step)
if episode.stop_reason.startswith("budget_"):
trace.record_budget_stop(episode.stop_reason.removeprefix("budget_"))
await trace.complete(
assembled_rows,
"agent_explore",
token_count=episode.tokens_used,
model_name=episode.model_name,
selected_doc_ids=selected_doc_ids,
)

evidence_text = _render_rows_evidence(assembled_rows)
response = {
Expand All @@ -293,4 +300,3 @@ async def _run_agent_explore_route(
f"chunks | evidence={len(evidence_text)} chars | router=agent_explore"
),
)

Loading