diff --git a/apps/api/tests/contract/test_retrieval_contract.py b/apps/api/tests/contract/test_retrieval_contract.py index 06cbc173b..2c9db2ff7 100644 --- a/apps/api/tests/contract/test_retrieval_contract.py +++ b/apps/api/tests/contract/test_retrieval_contract.py @@ -1,7 +1,7 @@ -from collections.abc import Callable +from collections.abc import Callable, Coroutine, Sequence from contextlib import AbstractAsyncContextManager from datetime import datetime, timezone -from typing import cast +from typing import Any, cast from uuid import uuid4 import pytest @@ -14,6 +14,9 @@ from shared.services.retrieval.workflow.run_request import WorkflowRunRequest from shared.services.retrieval.workflow.types import PlannedStep, QueryPlan, WorkflowResult +LLMFnInput = str | Sequence[dict[str, Any]] +LLMFn = Callable[[LLMFnInput], Coroutine[Any, Any, str]] + async def _seed_retrieval_document( *, @@ -23,6 +26,9 @@ async def _seed_retrieval_document( section_path: str, content: str, chunk_id: str | None = None, + chunk_type: str = "text", + file_path: str | None = None, + chunk_metadata: dict[str, Any] | None = None, ) -> dict[str, str]: document_id = f"doc_{uuid4().hex[:12]}" job_id = f"job_{uuid4().hex[:12]}" @@ -80,9 +86,11 @@ async def _seed_retrieval_document( document_id=document_id, job_result_id=job_result_id, section_id=section_id, - chunk_type="text", + chunk_type=chunk_type, content=content, section_path=section_path, + file_path=file_path, + chunk_metadata=chunk_metadata, ) return { @@ -600,6 +608,121 @@ async def test_agentic_retrieval_should_reference_discovery_content_when_navigat } +@pytest.mark.asyncio +async def test_agentic_retrieval_should_not_send_table_artifacts_to_vlm( + developer_api_client_factory: Callable[ + [], AbstractAsyncContextManager[AsyncClient] + ], + monkeypatch: MonkeyPatch, +) -> None: + monkeypatch.setenv("LLM_MOCK_ENABLED", "true") + vlm_calls: list[LLMFnInput] = [] + + async def fake_vlm(prompt: LLMFnInput) -> str: + vlm_calls.append(prompt) + return '{"status":"DONE","answer":"unexpected table VLM answer"}' + + def fake_create_retrieval_vlm_fn(**_kwargs: object) -> LLMFn: + return fake_vlm + + class FakeResultStorage: + def generate_artifact_url( + self, + *, + job_id: str, + artifact_ref: str, + expires_in: int = 3600, + ) -> str | None: + del expires_in + return f"https://assets.example.com/{job_id}/{artifact_ref}" + + def normalize_artifact_ref(self, artifact_ref: str | None) -> str | None: + if not artifact_ref: + return None + normalized = artifact_ref.strip().replace("\\", "/").lstrip("/") + if not normalized: + return None + root_dir = normalized.split("/", 1)[0] + if root_dir not in {"images", "tables"}: + return None + return normalized + + def fake_get_result_storage() -> FakeResultStorage: + return FakeResultStorage() + + async with developer_api_client_factory() as api_client: + monkeypatch.setattr( + "shared.services.retrieval.llm_adapter.create_retrieval_vlm_fn", + fake_create_retrieval_vlm_fn, + ) + monkeypatch.setattr( + "shared.services.retrieval.hydration.assets.get_result_storage", + fake_get_result_storage, + ) + table_document = await _seed_retrieval_document( + user_id="local-dev-user", + namespace="contract-agentic-table-vlm-filter", + source_file_name="table-report.md", + section_path="Realdata Results Summary / Main Metrics", + content=( + "" + "" + "
budgetmetricvalue
1000Flat inspect_evidence_score_mean0.5674
" + ), + chunk_type="table", + file_path="tables/table-0-main-metrics.html", + ) + await _seed_retrieval_document( + user_id="local-dev-user", + namespace="contract-agentic-table-vlm-filter", + source_file_name="filler-table-report.md", + section_path="Appendix / Filler Metrics", + content=( + "" + "
metricvalue
unrelated filler metric999
" + ), + chunk_type="table", + file_path="tables/table-1-filler-metrics.html", + ) + + response = await api_client.post( + "/api/v1/retrieval/query", + json={ + "namespace": "contract-agentic-table-vlm-filter", + "query": "budget 1000 Flat inspect_evidence_score_mean", + "top_k": 1, + "data_type": 4, + "use_agentic": True, + }, + ) + + assert response.status_code == 200 + + response_json = cast(dict[str, object], response.json()) + referenced_chunks = cast(list[dict[str, object]], response_json["referenced_chunks"]) + results = cast(list[dict[str, object]], response_json["results"]) + + assert response_json["router_used"] == "workflow_single_step" + assert vlm_calls == [] + matching_references = [ + reference + for reference in referenced_chunks + if reference["chunk_id"] == table_document["chunk_id"] + ] + assert len(matching_references) == 1 + assert matching_references[0]["document_id"] == table_document["document_id"] + assert matching_references[0]["chunk_type"] == "table" + assert matching_references[0]["section_path"] == table_document["section_path"] + assert matching_references[0]["file_path"] == "tables/table-0-main-metrics.html" + assert matching_references[0]["job_id"] == table_document["job_id"] + assert str(matching_references[0]["asset_url"]).startswith( + "https://assets.example.com/" + ) + assert len(results) == 1 + assert results[0]["chunk_type"] == "table" + assert _result_source(results[0])["document_id"] == table_document["document_id"] + + @pytest.mark.asyncio async def test_agentic_retrieval_should_not_hydrate_references_outside_request_scope( developer_api_client_factory: Callable[ diff --git a/packages/shared-python/shared/services/retrieval/agentic/evidence/builder.py b/packages/shared-python/shared/services/retrieval/agentic/evidence/builder.py index 6da147f42..9e158191a 100644 --- a/packages/shared-python/shared/services/retrieval/agentic/evidence/builder.py +++ b/packages/shared-python/shared/services/retrieval/agentic/evidence/builder.py @@ -52,18 +52,29 @@ def with_context_prompt_projection( return projected -def collect_media_chunks(node: DocTreeNode) -> list[dict[str, Any]]: - media: list[dict[str, Any]] = [] +def _collect_chunks_by_type( + node: DocTreeNode, + chunk_types: set[str], +) -> list[dict[str, Any]]: + collected_chunks: list[dict[str, Any]] = [] for chunks in node.leaf_content.values(): for chunk in chunks: chunk_type = ( chunk.get("chunk_type") or chunk.get("type") or "" ).strip().lower() - if chunk_type in ("image", "table"): - media.append(chunk) + if chunk_type in chunk_types: + collected_chunks.append(chunk) for child in node.children.values(): - media.extend(collect_media_chunks(child)) - return media + collected_chunks.extend(_collect_chunks_by_type(child, chunk_types)) + return collected_chunks + + +def collect_media_chunks(node: DocTreeNode) -> list[dict[str, Any]]: + return _collect_chunks_by_type(node, {"image", "table"}) + + +def collect_image_chunks(node: DocTreeNode) -> list[dict[str, Any]]: + return _collect_chunks_by_type(node, {"image"}) def collect_media_chunks_all( @@ -75,6 +86,15 @@ def collect_media_chunks_all( return media +def collect_image_chunks_all( + doc_trees: dict[str, DocTreeNode], +) -> list[dict[str, Any]]: + image_chunks: list[dict[str, Any]] = [] + for tree in doc_trees.values(): + image_chunks.extend(collect_image_chunks(tree)) + return image_chunks + + async def build_asset_url_map( media_chunks: list[dict[str, Any]], ) -> dict[str, str]: @@ -84,6 +104,13 @@ async def build_asset_url_map( ) +async def build_vlm_image_urls( + doc_trees: dict[str, DocTreeNode], +) -> list[str]: + asset_url_map = await build_asset_url_map(collect_image_chunks_all(doc_trees)) + return [url for url in asset_url_map.values() if url] + + def _collect_all_leaf_paths(node: DocTreeNode) -> set[str]: paths = set(node.leaf_content.keys()) for child in node.children.values(): diff --git a/packages/shared-python/shared/services/retrieval/agentic/orchestrator.py b/packages/shared-python/shared/services/retrieval/agentic/orchestrator.py index 4ffdda5d8..0e1cc51b2 100644 --- a/packages/shared-python/shared/services/retrieval/agentic/orchestrator.py +++ b/packages/shared-python/shared/services/retrieval/agentic/orchestrator.py @@ -31,8 +31,7 @@ ) from shared.services.retrieval.agentic.navigation.document import DocumentNavigationRunner from shared.services.retrieval.agentic.evidence.builder import ( - build_asset_url_map as _build_asset_url_map, - collect_media_chunks_all as _collect_media_chunks_all, + build_vlm_image_urls as _build_vlm_image_urls, render_evidence as _render_evidence, trim_evidence_to_budget as _trim_evidence_to_budget, with_context_prompt_projection as _with_context_prompt_projection, @@ -321,12 +320,7 @@ async def run( # Collect image URLs from evidence for VLM switch evidence_image_urls: list[str] = [] if vlm_fn: - asset_url_map = await _build_asset_url_map( - _collect_media_chunks_all(state.doc_trees), - ) - evidence_image_urls = [ - url for url in asset_url_map.values() if url - ] + evidence_image_urls = await _build_vlm_image_urls(state.doc_trees) async def vlm_context_call(prompt, _vlm_fn=vlm_fn): return await llm_budget.call(cast(LLMFn, _vlm_fn), prompt, pool='context')