From 3e0cb57ddaf1c9a059054e0dcd17aceec64b8fc7 Mon Sep 17 00:00:00 2001 From: chengke <404835780@qq.com> Date: Fri, 11 Sep 2026 16:05:24 +0800 Subject: [PATCH] feat(retrieval): apply include and exclude document IDs before tool results Keep request-level document scope on Explore tools, classic retrieval, and citations so excluded documents never reach the model. Co-authored-by: Cursor --- README.md | 2 + apps/api/app/api/v1/routes/retrieval.py | 9 +- .../contract/test_retrieval_document_scope.py | 619 ++++++++++++++++++ docs/retrieval-document-scope.md | 38 ++ .../retrieval/agent_explore/dispatch.py | 6 +- .../retrieval/agent_explore/harness/base.py | 3 + .../agent_explore/harness/cursor_harness.py | 4 + .../agent_explore/harness/openai_harness.py | 4 + .../retrieval/agent_explore/ref_resolution.py | 5 + .../retrieval/agent_tools/registry.py | 3 + .../retrieval/agent_tools/tools/assets.py | 1 + .../retrieval/agent_tools/tools/grep.py | 1 + .../agent_tools/tools/list_documents.py | 1 + .../retrieval/agent_tools/tools/neighbors.py | 11 +- .../agent_tools/tools/node_filter.py | 1 + .../retrieval/agent_tools/tools/outline.py | 1 + .../retrieval/agent_tools/tools/read.py | 2 + .../retrieval/agent_tools/tools/recall.py | 36 +- .../shared/services/retrieval/app_service.py | 2 + .../services/retrieval/cache_service.py | 4 + .../services/retrieval/document_scope.py | 46 ++ .../services/retrieval/execution/plan.py | 2 + .../retrieval/execution/query_request.py | 12 + .../retrieval/execution/route_types.py | 9 + .../services/retrieval/execution/routes.py | 8 + .../services/retrieval/hydration/connected.py | 7 + .../retrieval/hydration/result_assembly.py | 5 + .../services/retrieval/hydration/row_utils.py | 7 +- .../retrieval/search/map_unit_discovery.py | 33 +- .../retrieval/search/scoped_corpus.py | 10 +- 30 files changed, 841 insertions(+), 51 deletions(-) create mode 100644 apps/api/tests/contract/test_retrieval_document_scope.py create mode 100644 docs/retrieval-document-scope.md create mode 100644 packages/shared-python/shared/services/retrieval/document_scope.py diff --git a/README.md b/README.md index a9fdf7fc..0d2dfd41 100644 --- a/README.md +++ b/README.md @@ -269,6 +269,8 @@ make check ## Additional Guides +- Retrieval document scope: + [docs/retrieval-document-scope.md](docs/retrieval-document-scope.md) - External dependency guide: [docs/external-services.md](docs/external-services.md) - Architecture decisions: diff --git a/apps/api/app/api/v1/routes/retrieval.py b/apps/api/app/api/v1/routes/retrieval.py index 5d84f57f..06213fdc 100644 --- a/apps/api/app/api/v1/routes/retrieval.py +++ b/apps/api/app/api/v1/routes/retrieval.py @@ -32,7 +32,13 @@ class RetrievalQueryRequest(BaseModel): ) query: str top_k: int = DEFAULT_TOP_K - exclude_document_ids: list[str] = Field(default_factory=list) + include_document_ids: list[str] | None = Field( + None, description="Document allowlist: null means unrestricted, [] means empty; exclusions win." + ) + exclude_document_ids: list[str] = Field( + default_factory=list, + description="Documents excluded from every retrieval path. Exclusions override include_document_ids.", + ) exclude_sections: list[ExcludeSection] = Field(default_factory=list) data_type: int = Field( 1, @@ -190,6 +196,7 @@ async def execute_retrieval_query( namespace=normalize_retrieval_namespace(payload.namespace), query=payload.query, top_k=payload.top_k, + include_document_ids=payload.include_document_ids, exclude_document_ids=payload.exclude_document_ids, exclude_sections=[item.model_dump() for item in payload.exclude_sections], chunk_types=resolved_chunk_types, diff --git a/apps/api/tests/contract/test_retrieval_document_scope.py b/apps/api/tests/contract/test_retrieval_document_scope.py new file mode 100644 index 00000000..b7bd2fd2 --- /dev/null +++ b/apps/api/tests/contract/test_retrieval_document_scope.py @@ -0,0 +1,619 @@ +"""Document boundaries exercised against published PostgreSQL corpus data.""" + +import asyncio +import json +from contextlib import asynccontextmanager +from types import SimpleNamespace +from typing import Any +from uuid import uuid4 + +import pytest +from sqlalchemy import select + +from shared.models.database.document import DocumentChunk, GraphEdge, GraphNode +from shared.services.retrieval.agent_explore.dispatch import dispatch_tool_call +from shared.services.retrieval.agent_explore.budget import EpisodeBudget +from shared.services.retrieval.agent_explore.ref_resolution import resolve_finish_refs +from shared.services.retrieval.agent_explore.types import EpisodeResult +from shared.services.retrieval.cache_service import _cache_shape_digest +from shared.services.retrieval.document_scope import DocumentScope +from shared.services.retrieval.hydration.connected import hydrate_connected_target_rows +from shared.services.retrieval.hydration.result_assembly import ( + assemble_retrieval_results, +) +from tests.contract.test_retrieval_classic_map_unit_contract import _publish_document +from tests.support.retrieval_snapshot_support import contract_db_session + + +async def _corpus(namespace): + docs = [] + for label in ("alpha", "beta", "gamma"): + asset = f"{namespace}-{label}-asset" + doc = await _publish_document( + namespace=namespace, + source_file_name=f"{label}.pdf", + chunks=[ + { + "chunk_id": f"{namespace}-{label}-{index}", + "type": "text", + "content": f"{'scopeprobe' if index <= 2 else 'unrelated filler'} {label} evidence {index} [images/{label}.png]", + "path": f"{label}.pdf/Root/Section{index}/body", + "order": index, + "metadata": { + "connect_to": [ + {"target": asset, "ref": f"[images/{label}.png]"} + ] + }, + } + for index in range(1, 7) + ] + + [ + { + "chunk_id": asset, + "type": "image", + "content": f"{label} asset secret", + "path": f"images/{label}.png", + "order": 3, + "metadata": { + "summary": f"scopeprobe {label}", + "file_path": f"images/{label}.png", + }, + } + ], + ) + docs.append(doc) + async with contract_db_session() as db: + for doc in docs: + await db.merge( + GraphNode( + node_id=f"doc:{doc['document_id']}", + user_id="local-dev-user", + namespace=namespace, + node_kind="document", + owner_document_id=doc["document_id"], + job_result_id=doc["job_result_id"], + properties={"top_summary": doc["document_id"]}, + ) + ) + await db.flush() + for left, right in ((0, 1), (1, 2), (2, 0)): + db.add( + GraphEdge( + edge_id=f"scope-{uuid4().hex}", + user_id="local-dev-user", + namespace=namespace, + edge_kind="related", + source_node_id=f"doc:{docs[left]['document_id']}", + target_node_id=f"doc:{docs[right]['document_id']}", + owner_document_id=docs[left]["document_id"], + job_result_id=docs[left]["job_result_id"], + weight=1.0, + ) + ) + await db.commit() + return docs + + +def _matrix(ids): + a, b, c = ids + return [ + (None, [], {a, b, c}), + ([], [], set()), + ([a], [], {a}), + ([a, b], [b], {a}), + (None, [b], {a, c}), + ([a], [a], set()), + (["missing"], [], set()), + ([a, a, b], [c], {a, b}), + ] + + +async def test_scope_matrix_all_corpus_tools_and_refs(developer_api_client_factory): + namespace = f"scope-tools-{uuid4().hex[:8]}" + async with developer_api_client_factory(): + docs = await _corpus(namespace) + ids = [d["document_id"] for d in docs] + async with contract_db_session() as db: + chunks = ( + ( + await db.execute( + select(DocumentChunk).where(DocumentChunk.document_id.in_(ids)) + ) + ) + .scalars() + .all() + ) + refs = [{"document_id": c.document_id, "chunk_id": c.chunk_id} for c in chunks] + for include, exclude, expected in _matrix(ids): + scope = DocumentScope( + None if include is None else frozenset(include), frozenset(exclude) + ) + + async def call(name, args): + return await dispatch_tool_call( + f"corpus.{name}", + args, + db_factory=contract_db_session, + user_id="local-dev-user", + namespace=namespace, + document_scope=scope, + ) + + for name, args in [ + ("list_documents", {}), + ("grep", {"pattern": "scopeprobe"}), + ("grep", {"pattern": "scopeprobe", "document_ids": ids}), + ("recall", {"query": "scopeprobe", "channels": ["term"], "top_k": 50}), + ( + "recall", + {"query": "scopeprobe", "channels": ["path_content"], "top_k": 50}, + ), + ("recall", {"query": "scopeprobe", "document_ids": ids, "top_k": 50}), + ( + "node_filter", + { + "document_ids": ids, + "predicates": [{"field": "path", "terms": ["Section"]}], + }, + ), + ("assets", {}), + ( + "assets", + { + "host_of": [ + c.chunk_id for c in chunks if c.chunk_type == "image" + ] + }, + ), + ("read", {"refs": refs, "include_assets": False}), + ("read", {"refs": refs}), + ]: + result = await call(name, args) + assert {ref["document_id"] for ref in result.refs} == expected, ( + name, + include, + exclude, + result, + ) + if expected: + assert result.error is None, (name, result.error) + for doc_id in ids: + outline = await call("outline", {"document_id": doc_id}) + assert {r["document_id"] for r in outline.refs} == ({doc_id} & expected) + neighbors = await call("neighbors", {"document_id": doc_id}) + assert {r["document_id"] for r in neighbors.refs} == ( + expected - {doc_id} if doc_id in expected else set() + ) + # Explicit tool selectors cannot widen the request boundary. + if expected: + outside = next((doc for doc in ids if doc not in expected), None) + if outside: + for name, args in [ + ("grep", {"pattern": "scopeprobe", "document_ids": [outside]}), + ("recall", {"query": "scopeprobe", "document_ids": [outside]}), + ("assets", {"document_ids": [outside]}), + ]: + assert not (await call(name, args)).refs + async with contract_db_session() as db: + final_refs = await resolve_finish_refs( + db, + user_id="local-dev-user", + namespace=namespace, + refs=refs, + document_scope=scope, + ) + assert {r["document_id"] for r in final_refs} == expected + padded_refs = await resolve_finish_refs( + db, + user_id="local-dev-user", + namespace=namespace, + refs=[ + { + "document_id": f" {r['document_id']} ", + "chunk_id": r["chunk_id"], + } + for r in refs + ], + document_scope=scope, + ) + assert padded_refs == final_refs + + +@pytest.mark.parametrize("version", ["v1", "v2"]) +async def test_api_scope_matrix_small_classic_and_cache( + developer_api_client_factory, version +): + namespace = f"scope-api-{uuid4().hex[:8]}" + async with developer_api_client_factory() as client: + docs = await _corpus(namespace) + ids = [d["document_id"] for d in docs] + for top_k in (50, 1): + for include, exclude, expected in _matrix(ids): + response = await client.post( + f"/api/{version}/retrieval/query", + json={ + "namespace": namespace, + "query": "scopeprobe", + "top_k": top_k, + "use_agentic": False, + "include_document_ids": include, + "exclude_document_ids": exclude, + }, + ) + assert response.status_code == 200, response.text + body = response.json() + returned = {row["source"]["document_id"] for row in body["results"]} + assert returned <= expected + assert bool(returned) == bool(expected) + if top_k == 50: + assert returned == expected + if top_k == 1 and expected: + assert body["router_used"] == "classic_topk" + else: + assert body["router_used"] == "small_corpus_all" + + +async def test_agent_scope_survives_dispatch_and_untrusted_finish( + developer_api_client_factory, monkeypatch +): + namespace = f"scope-agent-{uuid4().hex[:8]}" + async with developer_api_client_factory() as client: + docs = await _corpus(namespace) + ids = [d["document_id"] for d in docs] + for include, exclude in [ + (None, [ids[1], ids[2]]), + ([ids[0], ids[1]], [ids[1]]), + ]: + calls = [] + + class Harness: + async def run_episode( + self, *, db_factory, user_id, namespace, document_scope, **kwargs + ): + result = await dispatch_tool_call( + "corpus.list_documents", + {}, + db_factory=db_factory, + user_id=user_id, + namespace=namespace, + document_scope=document_scope, + ) + calls.append(result) + assert {r["document_id"] for r in result.refs} == {ids[0]} + return EpisodeResult( + refs=[ + { + "document_id": doc, + "section_path": "Root / Section1 / body", + } + for doc in ids + ] + + [ + { + "document_id": f" {ids[2]} ", + "chunk_id": f"{namespace}-gamma-1", + }, + { + "document_id": f" {ids[0]} ", + "chunk_id": f"{namespace}-alpha-1", + }, + ], + notes="", + ) + + monkeypatch.setattr( + "shared.services.retrieval.agent_explore.harness.resolve_harness", + lambda: Harness(), + ) + response = await client.post( + "/api/v1/retrieval/query", + json={ + "namespace": namespace, + "query": "scopeprobe", + "top_k": 1, + "include_document_ids": include, + "exclude_document_ids": exclude, + }, + ) + assert response.status_code == 200, response.text + body = response.json() + assert calls + assert body["router_used"] == "agent_explore" + assert {r["document_id"] for r in body["referenced_chunks"]} == {ids[0]} + assert {r["source"]["document_id"] for r in body["results"]} == {ids[0]} + assert ( + "beta" not in body["evidence_text"] + and "gamma" not in body["evidence_text"] + ) + + +async def test_connected_hydration_and_assembly_reject_outside_rows( + developer_api_client_factory, +): + namespace = f"scope-assets-{uuid4().hex[:8]}" + async with developer_api_client_factory(): + docs = await _corpus(namespace) + ids = [d["document_id"] for d in docs] + async with contract_db_session() as db: + chunks = ( + ( + await db.execute( + select(DocumentChunk).where(DocumentChunk.document_id.in_(ids)) + ) + ) + .scalars() + .all() + ) + rows = [ + { + "document_id": c.document_id, + "job_result_id": c.job_result_id, + "chunk_id": c.chunk_id, + "chunk_type": c.chunk_type, + "content": c.content, + "chunk_metadata": c.chunk_metadata, + } + for c in chunks + if c.chunk_type == "text" + ] + for include, exclude, expected in _matrix(ids): + scope = DocumentScope( + None if include is None else frozenset(include), frozenset(exclude) + ) + hydrated = await hydrate_connected_target_rows( + db=db, + rows=rows, + exclude_document_ids=[], + exclude_sections=[], + document_scope=scope, + ) + assert {r["document_id"] for r in hydrated} == expected + legacy_filtered = await hydrate_connected_target_rows( + db=db, + rows=rows, + exclude_document_ids=[ids[0]], + exclude_sections=[], + document_scope=scope, + ) + assert {r["document_id"] for r in legacy_filtered} == expected - { + ids[0] + } + assembled = await assemble_retrieval_results( + db=db, + rows=rows, + exclude_document_ids=[], + exclude_sections=[], + document_scope=scope, + ) + assert {r["document_id"] for r in assembled} == expected + for row in assembled: + assert "asset secret" in row["content"] + + +def test_cache_scope_none_empty_and_set_identity(): + def digest(include, exclude=()): + return _cache_shape_digest( + query="same", + top_k=10, + exclude_document_ids=list(exclude), + exclude_sections=[], + include_document_ids=include, + ) + + assert ( + len( + { + digest(None), + digest([]), + digest(["a"]), + digest(["b"]), + digest(["a"], ["a"]), + } + ) + == 5 + ) + assert digest(["b", "a", "a"]) == digest(["a", "b"]) + + +async def test_legacy_excludes_only_narrow_scope_in_database( + developer_api_client_factory, +): + from dataclasses import replace + from shared.services.retrieval.execution.query_request import RetrievalQuery + from shared.services.retrieval.search.map_unit_discovery import map_unit_discovery + from shared.services.retrieval.search.scoped_corpus import ( + count_scoped_chunks, + load_all_scoped_chunks, + ) + + namespace = f"scope-legacy-{uuid4().hex[:8]}" + async with developer_api_client_factory(): + docs = await _corpus(namespace) + ids = [d["document_id"] for d in docs] + scope = DocumentScope(frozenset(ids), frozenset([ids[2]])) + async with contract_db_session() as db: + context = RetrievalQuery.from_parameters( + db=db, + user_id="local-dev-user", + namespace=namespace, + query="scopeprobe", + top_k=1, + exclude_document_ids=[ids[1]], + exclude_sections=[], + ).build_route_context() + context = replace(context, document_scope=scope) + assert [doc for doc in ids if context.document_scope.allows(doc)] == ids[:1] + kwargs: dict[str, Any] = dict( + user_id="local-dev-user", + namespace=namespace, + exclude_document_ids=[ids[1]], + document_scope=scope, + ) + assert ( + await count_scoped_chunks(db, **kwargs, allowed_chunk_types=None) == 7 + ) + rows = await load_all_scoped_chunks( + db, + **kwargs, + exclude_sections=[], + allowed_chunk_types=None, + signal_paths=[], + filter_mode="delete", + ) + assert {row["document_id"] for row in rows} == {ids[0]} + discovery = await map_unit_discovery( + db, + **kwargs, + query="scopeprobe", + top_k=20, + exclude_sections=[], + ) + assert {row["document_id"] for row in discovery.payload["fused_rows"]} == { + ids[0] + } + + +@pytest.mark.parametrize("provider", ["openai", "cursor"]) +@pytest.mark.parametrize("boundary", ["unscoped", "exclude", "include", "empty"]) +async def test_real_harness_provider_loop_scopes_postgresql_tools( + developer_api_client_factory, + monkeypatch, + provider, + boundary, +): + from shared.services.retrieval.agent_explore.harness import ( + cursor_harness, + openai_harness, + ) + + namespace = f"scope-provider-{uuid4().hex[:8]}" + async with developer_api_client_factory(): + docs = await _corpus(namespace) + ids = [doc["document_id"] for doc in docs] + scope = { + "unscoped": DocumentScope(), + "exclude": DocumentScope(exclude=frozenset(ids[1:])), + "include": DocumentScope(frozenset(ids[:2]), frozenset([ids[1]])), + "empty": DocumentScope(frozenset()), + }[boundary] + expected = {doc for doc in ids if scope.allows(doc)} + observed = [] + + def verify_observation(content): + observed.append(content) + for doc in ids: + assert (doc in content) == (doc in expected), (boundary, content) + + tool_calls = [ + SimpleNamespace( + id="list", + function=SimpleNamespace(name="corpus_list_documents", arguments="{}"), + ), + SimpleNamespace( + id="grep", + function=SimpleNamespace( + name="corpus_grep", + arguments=json.dumps( + {"pattern": "scopeprobe", "document_ids": ids} + ), + ), + ), + ] + + class CompletionClient: + def chat_completion_raw_with_usage(self, *, messages, **kwargs): + tool_messages = [m for m in messages if m["role"] == "tool"] + if tool_messages: + for message in tool_messages: + verify_observation(message["content"]) + calls = [ + SimpleNamespace( + id="finish", + function=SimpleNamespace( + name="finish", arguments='{"refs": []}' + ), + ) + ] + else: + calls = tool_calls + return SimpleNamespace( + choices=[ + SimpleNamespace( + message=SimpleNamespace(tool_calls=calls, content="") + ) + ] + ), {"total_tokens": 1} + + @asynccontextmanager + async def context(value): + yield value + + class Run: + def __init__(self, tools): + self.tools = tools + + async def wait(self): + # Cursor invokes callbacks on provider threads; both callbacks + # must cross run_coroutine_threadsafe and open their own DB session. + contents = await asyncio.gather( + *[ + asyncio.to_thread( + self.tools[call.function.name].execute, + json.loads(call.function.arguments), + None, + ) + for call in tool_calls + ] + ) + for content in contents: + verify_observation(content) + self.tools["finish"].execute({"refs": []}, None) + return SimpleNamespace(usage=SimpleNamespace(total_tokens=1)) + + class Agent: + def __init__(self, tools): + self.tools = tools + + async def send(self, prompt): + return Run(self.tools) + + class Agents: + async def create(self, options): + return context(Agent(options.local.custom_tools)) + + class Client: + @staticmethod + async def launch_bridge(**kwargs): + return context(SimpleNamespace(agents=Agents())) + + if provider == "openai": + monkeypatch.setattr( + openai_harness, + "_resolve_client_and_model", + lambda: (CompletionClient(), "test"), + ) + harness = openai_harness.OpenAIHarness() + else: + monkeypatch.setenv("CURSOR_API_KEY", "test-provider-no-network") + monkeypatch.setattr( + cursor_harness, + "_require_cursor_sdk", + lambda: SimpleNamespace( + CustomTool=SimpleNamespace, + AgentOptions=SimpleNamespace, + LocalAgentOptions=SimpleNamespace, + AsyncClient=Client, + ), + ) + harness = cursor_harness.CursorHarness() + episode = await harness.run_episode( + db_factory=contract_db_session, + user_id="local-dev-user", + namespace=namespace, + document_scope=scope, + query="scopeprobe", + budget=EpisodeBudget(), + ) + assert len(observed) == 2 + assert episode.stop_reason == "finished" + assert all(step.error is None for step in episode.steps) diff --git a/docs/retrieval-document-scope.md b/docs/retrieval-document-scope.md new file mode 100644 index 00000000..36775430 --- /dev/null +++ b/docs/retrieval-document-scope.md @@ -0,0 +1,38 @@ +# Retrieval document scope + +Both `POST /api/v1/retrieval/query` and `POST /api/v2/retrieval/query` accept +`include_document_ids` and `exclude_document_ids`. + +| Request field | Meaning | +| --- | --- | +| `include_document_ids` omitted or `null` | All otherwise accessible documents in the requested namespace are eligible. | +| `include_document_ids: []` | No documents are eligible; retrieval returns empty results. | +| `include_document_ids: ["doc_a", "doc_b"]` | Only those documents are eligible. | +| `exclude_document_ids: ["doc_b"]` | Exclude these documents, including when they also appear in the include list. | +| `exclude_document_ids` omitted or `[]` | No additional document exclusions. | + +For example: + +```json +{ + "namespace": "default", + "query": "What are the findings?", + "include_document_ids": ["doc_a", "doc_b"], + "exclude_document_ids": ["doc_b"] +} +``` + +Only `doc_a` is eligible. IDs must be document IDs, not filenames. Unknown, +foreign-user, and other-namespace IDs do not grant access. Duplicate include +IDs do not broaden the scope. An inclusion list that leaves no eligible +documents returns empty results and does not fall back to the full corpus. + +The same boundary applies to small-corpus retrieval, classic search +(`use_agentic: false`), and agent exploration. Agent tools may select a narrower +set of documents, but cannot broaden the request boundary. Final references, +results, and connected asset hydration obey the same scope. Cache entries +distinguish unrestricted, empty, and explicitly included document sets. + +Scope restricts available evidence; it does not add an LLM routing step or +change path filtering, ranking, or threshold semantics. Omitting both fields +preserves unrestricted retrieval within the existing user/namespace boundary. diff --git a/packages/shared-python/shared/services/retrieval/agent_explore/dispatch.py b/packages/shared-python/shared/services/retrieval/agent_explore/dispatch.py index 22621db5..31f7b666 100644 --- a/packages/shared-python/shared/services/retrieval/agent_explore/dispatch.py +++ b/packages/shared-python/shared/services/retrieval/agent_explore/dispatch.py @@ -18,6 +18,8 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from collections.abc import Callable from contextlib import AbstractAsyncContextManager from typing import Any @@ -36,6 +38,7 @@ async def dispatch_tool_call( db_factory: DbFactory, user_id: str, namespace: str, + document_scope: DocumentScope = DocumentScope(), budget: ToolBudget | None = None, ) -> ToolResult: """Run one ``REGISTRY`` tool call against a fresh, call-scoped DB session. @@ -49,7 +52,8 @@ async def dispatch_tool_call( try: async with db_factory() as db: tool_ctx = ToolContext( - db=db, user_id=user_id, namespace=namespace, budget=budget or ToolBudget() + db=db, user_id=user_id, namespace=namespace, budget=budget or ToolBudget(), + document_scope=document_scope ) return await REGISTRY.dispatch(name, tool_ctx, args) except Exception as exc: # noqa: BLE001 - one broken tool must not kill the episode diff --git a/packages/shared-python/shared/services/retrieval/agent_explore/harness/base.py b/packages/shared-python/shared/services/retrieval/agent_explore/harness/base.py index 4bab470e..5972bb83 100644 --- a/packages/shared-python/shared/services/retrieval/agent_explore/harness/base.py +++ b/packages/shared-python/shared/services/retrieval/agent_explore/harness/base.py @@ -10,6 +10,8 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from typing import Protocol, runtime_checkable from shared.services.retrieval.agent_explore.budget import EpisodeBudget @@ -27,6 +29,7 @@ async def run_episode( db_factory: DbFactory, user_id: str, namespace: str, + document_scope: DocumentScope = DocumentScope(), query: str, budget: EpisodeBudget, ) -> EpisodeResult: diff --git a/packages/shared-python/shared/services/retrieval/agent_explore/harness/cursor_harness.py b/packages/shared-python/shared/services/retrieval/agent_explore/harness/cursor_harness.py index ebade0ea..351ee43f 100644 --- a/packages/shared-python/shared/services/retrieval/agent_explore/harness/cursor_harness.py +++ b/packages/shared-python/shared/services/retrieval/agent_explore/harness/cursor_harness.py @@ -62,6 +62,8 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + import asyncio import contextlib import json @@ -135,6 +137,7 @@ async def run_episode( db_factory: DbFactory, user_id: str, namespace: str, + document_scope: DocumentScope = DocumentScope(), query: str, budget: EpisodeBudget, ) -> EpisodeResult: @@ -192,6 +195,7 @@ def _dispatch_sync(tool_name: str, args: dict[str, Any]) -> str: db_factory=db_factory, user_id=user_id, namespace=namespace, + document_scope=document_scope, budget=tool_budget, ), loop, diff --git a/packages/shared-python/shared/services/retrieval/agent_explore/harness/openai_harness.py b/packages/shared-python/shared/services/retrieval/agent_explore/harness/openai_harness.py index 80fd5716..0e47cc19 100644 --- a/packages/shared-python/shared/services/retrieval/agent_explore/harness/openai_harness.py +++ b/packages/shared-python/shared/services/retrieval/agent_explore/harness/openai_harness.py @@ -51,6 +51,8 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + import asyncio import json import time @@ -175,6 +177,7 @@ async def run_episode( db_factory: DbFactory, user_id: str, namespace: str, + document_scope: DocumentScope = DocumentScope(), query: str, budget: EpisodeBudget, ) -> EpisodeResult: @@ -322,6 +325,7 @@ async def run_episode( db_factory=db_factory, user_id=user_id, namespace=namespace, + document_scope=document_scope, budget=tool_budget, ) tool_elapsed_ms = int((time.perf_counter() - tool_started) * 1000) diff --git a/packages/shared-python/shared/services/retrieval/agent_explore/ref_resolution.py b/packages/shared-python/shared/services/retrieval/agent_explore/ref_resolution.py index 5edeeac4..423e6ebf 100644 --- a/packages/shared-python/shared/services/retrieval/agent_explore/ref_resolution.py +++ b/packages/shared-python/shared/services/retrieval/agent_explore/ref_resolution.py @@ -19,6 +19,8 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from typing import Any from sqlalchemy import select @@ -38,8 +40,10 @@ async def resolve_finish_refs( user_id: str, namespace: str, refs: list[dict[str, Any]], + document_scope: DocumentScope = DocumentScope(), ) -> list[dict[str, Any]]: """Return refs with ``chunk_id`` populated; drops refs that don't resolve.""" + refs = [ref for ref in refs if document_scope.allows(str(ref.get("document_id") or "").strip())] document_ids = { str(ref.get("document_id") or "").strip() for ref in refs if ref.get("document_id") } @@ -54,6 +58,7 @@ async def resolve_finish_refs( .where(Document.user_id == user_id) .where(Document.namespace == namespace) .where(Document.status == "active") + .where(document_scope.predicate(Document.document_id)) ) ) .scalars() diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/registry.py b/packages/shared-python/shared/services/retrieval/agent_tools/registry.py index d5aa871c..5da6af86 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/registry.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/registry.py @@ -12,6 +12,8 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from collections.abc import Awaitable, Callable from dataclasses import dataclass, field from typing import Any @@ -60,6 +62,7 @@ class ToolContext: user_id: str namespace: str budget: ToolBudget = field(default_factory=ToolBudget) + document_scope: DocumentScope = DocumentScope() @dataclass diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/tools/assets.py b/packages/shared-python/shared/services/retrieval/agent_tools/tools/assets.py index 18c5a1b6..f15ecc49 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/tools/assets.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/tools/assets.py @@ -66,6 +66,7 @@ async def assets(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: host_of = [str(c).strip() for c in (args.get("host_of") or []) if str(c).strip()] scope_filters: list[Any] = [ + ctx.document_scope.predicate(Document.document_id), Document.user_id == ctx.user_id, Document.namespace == ctx.namespace, Document.status == "active", diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/tools/grep.py b/packages/shared-python/shared/services/retrieval/agent_tools/tools/grep.py index 34f86504..afad9cbd 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/tools/grep.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/tools/grep.py @@ -103,6 +103,7 @@ async def grep(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: document_ids=document_ids, chunk_types=chunk_types, ) + filters.append(ctx.document_scope.predicate(Document.document_id)) content_filter = ( DocumentChunk.content.op("~*")(pattern) if is_regex diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/tools/list_documents.py b/packages/shared-python/shared/services/retrieval/agent_tools/tools/list_documents.py index d0ddffe1..4ce879df 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/tools/list_documents.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/tools/list_documents.py @@ -45,6 +45,7 @@ async def list_documents(ctx: ToolContext, _args: dict[str, Any]) -> ToolResult: .where(Document.user_id == ctx.user_id) .where(Document.namespace == ctx.namespace) .where(Document.status == "active") + .where(ctx.document_scope.predicate(Document.document_id)) .order_by(Document.source_file_name) ) rows = (await ctx.db.execute(stmt)).all() diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/tools/neighbors.py b/packages/shared-python/shared/services/retrieval/agent_tools/tools/neighbors.py index 2da21d68..09607c5c 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/tools/neighbors.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/tools/neighbors.py @@ -47,11 +47,15 @@ async def neighbors(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: .where(Document.user_id == ctx.user_id) .where(Document.namespace == ctx.namespace) .where(Document.status == "active") + .where(ctx.document_scope.predicate(Document.document_id)) ) ).scalar_one_or_none() if document is None: return ToolResult(text="", error=f"unknown document_id: {document_id}") + allowed_nodes = select(GraphNode.node_id).where( + ctx.document_scope.predicate(GraphNode.owner_document_id) + ) node_id = f"doc:{document_id}" edges = ( ( @@ -60,6 +64,8 @@ async def neighbors(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: .where(GraphEdge.user_id == ctx.user_id) .where(GraphEdge.namespace == ctx.namespace) .where(GraphEdge.edge_kind == "related") + .where(GraphEdge.source_node_id.in_(allowed_nodes)) + .where(GraphEdge.target_node_id.in_(allowed_nodes)) .where( or_( GraphEdge.source_node_id == node_id, @@ -81,7 +87,10 @@ async def neighbors(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: peer_nodes = ( ( await ctx.db.execute( - select(GraphNode).where(GraphNode.node_id.in_(peer_node_ids)) + select(GraphNode).where( + GraphNode.node_id.in_(peer_node_ids), + ctx.document_scope.predicate(GraphNode.owner_document_id), + ) ) ) .scalars() diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/tools/node_filter.py b/packages/shared-python/shared/services/retrieval/agent_tools/tools/node_filter.py index 84c918fc..04eb6bb5 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/tools/node_filter.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/tools/node_filter.py @@ -114,6 +114,7 @@ async def node_filter(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: .where(Document.user_id == ctx.user_id) .where(Document.namespace == ctx.namespace) .where(Document.status == "active") + .where(ctx.document_scope.predicate(Document.document_id)) ) ) .scalars() diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/tools/outline.py b/packages/shared-python/shared/services/retrieval/agent_tools/tools/outline.py index 2596d468..6e24f5f8 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/tools/outline.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/tools/outline.py @@ -74,6 +74,7 @@ async def outline(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: .where(Document.user_id == ctx.user_id) .where(Document.namespace == ctx.namespace) .where(Document.status == "active") + .where(ctx.document_scope.predicate(Document.document_id)) ) ).scalar_one_or_none() if document is None or not document.current_job_result_id: diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/tools/read.py b/packages/shared-python/shared/services/retrieval/agent_tools/tools/read.py index 9d7072f7..4f8bebc3 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/tools/read.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/tools/read.py @@ -230,6 +230,7 @@ async def read(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: .where(Document.user_id == ctx.user_id) .where(Document.namespace == ctx.namespace) .where(Document.status == "active") + .where(ctx.document_scope.predicate(Document.document_id)) ) ) .scalars() @@ -388,6 +389,7 @@ async def read(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: db=ctx.db, rows=base_rows, exclude_document_ids=[], + document_scope=ctx.document_scope, exclude_sections=[], ) rows_by_chunk_id = { diff --git a/packages/shared-python/shared/services/retrieval/agent_tools/tools/recall.py b/packages/shared-python/shared/services/retrieval/agent_tools/tools/recall.py index ec46ece5..f799a171 100644 --- a/packages/shared-python/shared/services/retrieval/agent_tools/tools/recall.py +++ b/packages/shared-python/shared/services/retrieval/agent_tools/tools/recall.py @@ -27,12 +27,14 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from typing import Any from sqlalchemy import select, text from sqlalchemy.ext.asyncio import AsyncSession -from shared.models.database.document import Document, DocumentChunk +from shared.models.database.document import DocumentChunk from shared.services.retrieval.agent_tools.registry import ( ToolContext, ToolResult, @@ -65,28 +67,13 @@ """ -async def _excluded_document_ids( - db: AsyncSession, *, user_id: str, namespace: str, document_ids: list[str] -) -> list[str]: - if not document_ids: - return [] - rows = await db.execute( - select(Document.document_id) - .where(Document.user_id == user_id) - .where(Document.namespace == namespace) - .where(Document.status == "active") - .where(Document.document_id.notin_(document_ids)) - ) - return [str(r[0]) for r in rows.all()] - - async def _term_channel_rows( db: AsyncSession, *, user_id: str, namespace: str, query: str, - document_ids: list[str], + document_scope: DocumentScope, chunk_types: set[str] | None, top_k: int, ) -> list[dict[str, Any]]: @@ -100,10 +87,8 @@ async def _term_channel_rows( "needle": needle, "limit": top_k, } - doc_clause = "" - if document_ids: - doc_clause = "AND d.document_id = ANY(:doc_ids)" - params["doc_ids"] = document_ids + doc_clause, scope_params = document_scope.sql() + params.update(scope_params) statement = text(_TERM_CHANNEL_SQL.format(doc_clause=doc_clause)) unit_rows = [dict(row._mapping) for row in (await db.execute(statement, params)).all()] if not unit_rows: @@ -202,20 +187,19 @@ async def recall(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: error="no runnable channels requested (vector is reserved, not implemented)", ) + document_scope = ctx.document_scope.narrow(document_ids) if document_ids else ctx.document_scope channel_rows: list[list[dict[str, Any]]] = [] weights: list[float] = [] if "path_content" in active_channels: - exclude_document_ids = await _excluded_document_ids( - ctx.db, user_id=ctx.user_id, namespace=ctx.namespace, document_ids=document_ids - ) discovery = await map_unit_discovery( ctx.db, user_id=ctx.user_id, namespace=ctx.namespace, query=query, top_k=top_k, - exclude_document_ids=exclude_document_ids, + exclude_document_ids=[], + document_scope=document_scope, exclude_sections=[], chunk_types=chunk_types, ) @@ -228,7 +212,7 @@ async def recall(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: user_id=ctx.user_id, namespace=ctx.namespace, query=query, - document_ids=document_ids, + document_scope=document_scope, chunk_types=chunk_types, top_k=top_k, ) diff --git a/packages/shared-python/shared/services/retrieval/app_service.py b/packages/shared-python/shared/services/retrieval/app_service.py index 69727d1e..7a24d48f 100644 --- a/packages/shared-python/shared/services/retrieval/app_service.py +++ b/packages/shared-python/shared/services/retrieval/app_service.py @@ -23,6 +23,7 @@ async def run_retrieval_query( top_k: int, exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + include_document_ids: list[str] | None = None, chunk_types: set[str] | None = None, signal_paths: list[str] | None = None, filter_mode: str = "delete", @@ -43,6 +44,7 @@ async def run_retrieval_query( top_k=top_k, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, + include_document_ids=include_document_ids, chunk_types=chunk_types, signal_paths=signal_paths, filter_mode=filter_mode, diff --git a/packages/shared-python/shared/services/retrieval/cache_service.py b/packages/shared-python/shared/services/retrieval/cache_service.py index 07ea6b6a..794f608b 100644 --- a/packages/shared-python/shared/services/retrieval/cache_service.py +++ b/packages/shared-python/shared/services/retrieval/cache_service.py @@ -80,6 +80,7 @@ def _cache_shape_digest( llm_text_model: str | None = None, llm_vision_model: str | None = None, harness: str | None = None, + include_document_ids: list[str] | None = None, ) -> str: normalized_excludes = sorted(exclude_document_ids) normalized_sections = _normalize_exclude_sections(exclude_sections) @@ -101,6 +102,9 @@ def _cache_shape_digest( ] ) payload = f"{query}|{top_k}|{'|'.join(normalized_excludes)}|{'|'.join(normalized_sections)}|{extra}" + payload += "|document_scope_v1|" + repr( + None if include_document_ids is None else sorted(set(include_document_ids)) + ) return hashlib.sha256(payload.encode("utf-8")).hexdigest() diff --git a/packages/shared-python/shared/services/retrieval/document_scope.py b/packages/shared-python/shared/services/retrieval/document_scope.py new file mode 100644 index 00000000..f1861415 --- /dev/null +++ b/packages/shared-python/shared/services/retrieval/document_scope.py @@ -0,0 +1,46 @@ +"""Request-owned document boundary shared by retrieval and corpus tools.""" + +from dataclasses import dataclass +from typing import Any + +from sqlalchemy import and_, true + + +@dataclass(frozen=True) +class DocumentScope: + include: frozenset[str] | None = None + exclude: frozenset[str] = frozenset() + + def allows(self, document_id: str | None) -> bool: + return document_id not in self.exclude and ( + self.include is None or document_id in self.include + ) + + def predicate(self, column: Any) -> Any: + clauses = [] + if self.include is not None: + clauses.append(column.in_(sorted(self.include))) + if self.exclude: + clauses.append(column.notin_(sorted(self.exclude))) + return and_(*clauses) if clauses else true() + + def sql(self, column: str = "d.document_id") -> tuple[str, dict[str, Any]]: + clauses: list[str] = [] + params: dict[str, Any] = {} + if self.include is not None: + clauses.append(f"AND {column} = ANY(:scope_include)") + params["scope_include"] = sorted(self.include) + if self.exclude: + clauses.append(f"AND {column} <> ALL(:scope_exclude)") + params["scope_exclude"] = sorted(self.exclude) + return " ".join(clauses), params + + def narrow(self, document_ids: list[str]) -> "DocumentScope": + include = frozenset(document_ids) + return DocumentScope( + include if self.include is None else self.include & include, self.exclude + ) + + def excluding(self, document_ids: list[str]) -> "DocumentScope": + """Retained exclude-only callers can only further restrict a scope.""" + return DocumentScope(self.include, self.exclude | frozenset(document_ids)) diff --git a/packages/shared-python/shared/services/retrieval/execution/plan.py b/packages/shared-python/shared/services/retrieval/execution/plan.py index 8e9a9c7f..ebf5a88d 100644 --- a/packages/shared-python/shared/services/retrieval/execution/plan.py +++ b/packages/shared-python/shared/services/retrieval/execution/plan.py @@ -39,6 +39,7 @@ async def run_retrieval_query( top_k: int, exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + include_document_ids: list[str] | None = None, chunk_types: set[str] | None = None, signal_paths: list[str] | None = None, filter_mode: str = "delete", @@ -61,6 +62,7 @@ async def run_retrieval_query( top_k=top_k, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, + include_document_ids=include_document_ids, chunk_types=chunk_types, signal_paths=signal_paths, filter_mode=filter_mode, diff --git a/packages/shared-python/shared/services/retrieval/execution/query_request.py b/packages/shared-python/shared/services/retrieval/execution/query_request.py index c6d3488d..55d3690f 100644 --- a/packages/shared-python/shared/services/retrieval/execution/query_request.py +++ b/packages/shared-python/shared/services/retrieval/execution/query_request.py @@ -1,5 +1,7 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from dataclasses import dataclass from typing import Any @@ -34,6 +36,7 @@ class RetrievalQuery: use_agentic: bool | None = None conversation_id: str | None = None llm_config: LLMConfig | None = None + include_document_ids: list[str] | None = None @classmethod def from_parameters( @@ -46,6 +49,7 @@ def from_parameters( top_k: int, exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + include_document_ids: list[str] | None = None, chunk_types: set[str] | None = None, signal_paths: list[str] | None = None, filter_mode: str = "delete", @@ -66,6 +70,7 @@ def from_parameters( top_k=top_k, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, + include_document_ids=include_document_ids, chunk_types=chunk_types, signal_paths=signal_paths, filter_mode=filter_mode, @@ -88,6 +93,9 @@ def build_cache_extra(self) -> dict[str, Any]: text_model = text_provider.model if text_provider is not None else None vision_model = vision_provider.model if vision_provider is not None else None return { + "include_document_ids": ( + None if self.include_document_ids is None else sorted(set(self.include_document_ids)) + ), "chunk_types": sorted(self.chunk_types) if self.chunk_types else None, "signal_paths": self.signal_paths, "filter_mode": self.filter_mode, @@ -121,6 +129,10 @@ def build_route_context(self) -> RetrievalRouteContext: top_k=self.top_k, exclude_document_ids=self.exclude_document_ids, exclude_sections=self.exclude_sections, + document_scope=DocumentScope( + None if self.include_document_ids is None else frozenset(self.include_document_ids), + frozenset(self.exclude_document_ids), + ), allowed_chunk_types=self.resolve_allowed_chunk_types(), chunk_types=self.chunk_types, signal_paths=self.signal_paths, diff --git a/packages/shared-python/shared/services/retrieval/execution/route_types.py b/packages/shared-python/shared/services/retrieval/execution/route_types.py index 33918090..c9650494 100644 --- a/packages/shared-python/shared/services/retrieval/execution/route_types.py +++ b/packages/shared-python/shared/services/retrieval/execution/route_types.py @@ -1,5 +1,7 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from dataclasses import dataclass from typing import Any @@ -30,6 +32,13 @@ class RetrievalRouteContext: use_agentic: bool | None conversation_id: str | None = None revision_pins: RetrievalRevisionPins | None = None + document_scope: DocumentScope = DocumentScope() + + def __post_init__(self) -> None: + # Direct route callers still supply the pre-existing exclude field. + object.__setattr__( + self, "document_scope", self.document_scope.excluding(self.exclude_document_ids) + ) @dataclass(frozen=True) diff --git a/packages/shared-python/shared/services/retrieval/execution/routes.py b/packages/shared-python/shared/services/retrieval/execution/routes.py index b3a89282..039239ab 100644 --- a/packages/shared-python/shared/services/retrieval/execution/routes.py +++ b/packages/shared-python/shared/services/retrieval/execution/routes.py @@ -79,6 +79,7 @@ async def _try_run_small_corpus_route( user_id=context.user_id, namespace=context.namespace, exclude_document_ids=context.exclude_document_ids, + document_scope=context.document_scope, allowed_chunk_types=context.allowed_chunk_types, revision_pins=context.revision_pins, max_count=context.top_k + 1, @@ -97,6 +98,7 @@ async def _try_run_small_corpus_route( user_id=context.user_id, namespace=context.namespace, exclude_document_ids=context.exclude_document_ids, + document_scope=context.document_scope, exclude_sections=context.exclude_sections, allowed_chunk_types=context.allowed_chunk_types, signal_paths=context.signal_paths or [], @@ -110,6 +112,7 @@ async def _try_run_small_corpus_route( db=context.db, rows=all_rows, exclude_document_ids=context.exclude_document_ids, + document_scope=context.document_scope, exclude_sections=context.exclude_sections, allowed_chunk_types=context.allowed_chunk_types, revision_pins=context.revision_pins, @@ -142,6 +145,7 @@ async def _run_classic_topk_route( query=context.query, top_k=context.effective_recall_k, exclude_document_ids=context.exclude_document_ids, + document_scope=context.document_scope, exclude_sections=context.exclude_sections, chunk_types=context.allowed_chunk_types, signal_paths=context.signal_paths, @@ -165,6 +169,7 @@ async def _run_classic_topk_route( db=context.db, rows=ranked_rows, exclude_document_ids=context.exclude_document_ids, + document_scope=context.document_scope, exclude_sections=context.exclude_sections, allowed_chunk_types=context.allowed_chunk_types, revision_pins=context.revision_pins, @@ -217,6 +222,7 @@ async def _run_agent_explore_route( namespace=context.namespace, query=context.query, budget=EpisodeBudget(), + document_scope=context.document_scope, ) logger.info( "retrieval agent_explore stage=episode seconds={:.3f} refs={} " @@ -241,6 +247,7 @@ async def _run_agent_explore_route( user_id=context.user_id, namespace=context.namespace, refs=episode.refs, + document_scope=context.document_scope, ) resolved = await resolve_workflow_references( db=final_db, @@ -253,6 +260,7 @@ async def _run_agent_explore_route( db=final_db, rows=resolved.rows, exclude_document_ids=context.exclude_document_ids, + document_scope=context.document_scope, exclude_sections=context.exclude_sections, allowed_chunk_types=context.allowed_chunk_types, revision_pins=context.revision_pins, diff --git a/packages/shared-python/shared/services/retrieval/hydration/connected.py b/packages/shared-python/shared/services/retrieval/hydration/connected.py index 8e637957..b960946a 100644 --- a/packages/shared-python/shared/services/retrieval/hydration/connected.py +++ b/packages/shared-python/shared/services/retrieval/hydration/connected.py @@ -1,5 +1,7 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from collections.abc import Mapping from typing import Any @@ -21,11 +23,14 @@ async def hydrate_connected_target_rows( rows: list[dict[str, Any]], exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + document_scope: DocumentScope = DocumentScope(), revision_pins: Mapping[str, str] | None = None, ) -> list[dict[str, Any]]: if db is None: return [] + document_scope = document_scope.excluding(exclude_document_ids) + existing_chunk_ids = { str(row.get('chunk_id') or '').strip() for row in rows @@ -88,6 +93,7 @@ async def hydrate_connected_target_rows( ) .outerjoin(DocumentSection, DocumentSection.section_id == DocumentChunk.section_id) .join(JobResult, JobResult.id == DocumentChunk.job_result_id) + .where(document_scope.predicate(Document.document_id)) .where(or_(*revision_filters)) .order_by(DocumentChunk.sort_order) ) @@ -118,4 +124,5 @@ async def hydrate_connected_target_rows( hydrated_rows, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, + document_scope=document_scope, ) diff --git a/packages/shared-python/shared/services/retrieval/hydration/result_assembly.py b/packages/shared-python/shared/services/retrieval/hydration/result_assembly.py index 2f6aa54a..610bf99f 100644 --- a/packages/shared-python/shared/services/retrieval/hydration/result_assembly.py +++ b/packages/shared-python/shared/services/retrieval/hydration/result_assembly.py @@ -1,5 +1,7 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from collections.abc import Mapping from typing import Any @@ -24,6 +26,7 @@ async def assemble_retrieval_results( rows: list[dict[str, Any]], exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + document_scope: DocumentScope = DocumentScope(), allowed_chunk_types: set[str] | None = None, revision_pins: Mapping[str, str] | None = None, ) -> list[dict[str, Any]]: @@ -31,6 +34,7 @@ async def assemble_retrieval_results( rows, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, + document_scope=document_scope, ) if allowed_chunk_types is not None: filtered_rows = [ @@ -42,6 +46,7 @@ async def assemble_retrieval_results( rows=filtered_rows, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, + document_scope=document_scope, revision_pins=revision_pins, ) rows_by_chunk_id = { diff --git a/packages/shared-python/shared/services/retrieval/hydration/row_utils.py b/packages/shared-python/shared/services/retrieval/hydration/row_utils.py index 74fa382e..c5c81a5d 100644 --- a/packages/shared-python/shared/services/retrieval/hydration/row_utils.py +++ b/packages/shared-python/shared/services/retrieval/hydration/row_utils.py @@ -1,5 +1,7 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from typing import Any from shared.services.retrieval.search.section_filters import is_excluded_section @@ -58,12 +60,13 @@ def filter_excluded_rows( *, exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + document_scope: DocumentScope = DocumentScope(), ) -> list[dict[str, Any]]: filtered: list[dict[str, Any]] = [] - excluded_documents = set(exclude_document_ids) + document_scope = document_scope.excluding(exclude_document_ids) for row in rows: document_id = row.get('document_id') - if document_id in excluded_documents: + if not document_scope.allows(document_id): continue if is_excluded_section( document_id=document_id, diff --git a/packages/shared-python/shared/services/retrieval/search/map_unit_discovery.py b/packages/shared-python/shared/services/retrieval/search/map_unit_discovery.py index b2a72978..9a1561c9 100644 --- a/packages/shared-python/shared/services/retrieval/search/map_unit_discovery.py +++ b/packages/shared-python/shared/services/retrieval/search/map_unit_discovery.py @@ -14,6 +14,8 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + import json import time from collections.abc import Mapping @@ -68,7 +70,7 @@ AND d.namespace = :namespace AND d.status = 'active' {revision_clause} - {exclude_clause} + {document_scope_clause} {type_clause} {signal_clause} ) @@ -88,7 +90,7 @@ AND d.namespace = :namespace AND d.status = 'active' {revision_clause} - {exclude_clause} + {document_scope_clause} {type_clause} ) """ @@ -146,14 +148,6 @@ def _build_type_clause( return f"AND ({' OR '.join(clauses)})", {} -def _build_exclude_clause(exclude_document_ids: list[str]) -> tuple[str, dict[str, Any]]: - if not exclude_document_ids: - return "", {} - return "AND d.document_id <> ALL(:excluded_doc_ids)", { - "excluded_doc_ids": list(exclude_document_ids) - } - - def _build_signal_clause( signal_paths: list[str], filter_mode: str ) -> tuple[str, dict[str, Any]]: @@ -179,6 +173,7 @@ async def map_unit_discovery( top_k: int, exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + document_scope: DocumentScope = DocumentScope(), chunk_types: set[str] | None = None, signal_paths: list[str] | None = None, filter_mode: str = "delete", @@ -201,14 +196,15 @@ async def map_unit_discovery( revision_join, revision_clause, revision_params = _build_revision_scope( revision_pins ) - exclude_clause, exclude_params = _build_exclude_clause(exclude_document_ids) + document_scope = document_scope.excluding(exclude_document_ids) + document_scope_clause, document_scope_params = document_scope.sql() type_clause, type_params = _build_type_clause(chunk_types) signal_clause, signal_params = _build_signal_clause( signal_paths or [], filter_mode ) params: dict[str, Any] = {"user_id": user_id, "namespace": namespace} params.update(revision_params) - params.update(exclude_params) + params.update(document_scope_params) params.update(type_params) params.update(signal_params) @@ -218,13 +214,15 @@ async def map_unit_discovery( signal_paths, exclude_sections, exclude_document_ids, + document_scope.include is not None, + document_scope.exclude, ) ) cte = _SCOPED_UNITS_CTE.format( revision_join=revision_join, revision_clause=revision_clause, - exclude_clause=exclude_clause, + document_scope_clause=document_scope_clause, type_clause=type_clause, signal_clause=signal_clause, ) @@ -266,7 +264,7 @@ async def map_unit_discovery( frequency_scope_cte = cte if signal_paths else _SCOPED_UNIT_IDS_CTE.format( revision_join=revision_join, revision_clause=revision_clause, - exclude_clause=exclude_clause, + document_scope_clause=document_scope_clause, type_clause=type_clause, ) frequency_query = text( @@ -336,7 +334,7 @@ async def map_unit_discovery( else _SCOPED_UNIT_IDS_CTE.format( revision_join=revision_join, revision_clause=revision_clause, - exclude_clause=exclude_clause, + document_scope_clause=document_scope_clause, type_clause=type_clause, ) ) @@ -429,7 +427,7 @@ async def map_unit_discovery( _SCOPED_UNIT_IDS_CTE.format( revision_join=revision_join, revision_clause=revision_clause, - exclude_clause=exclude_clause, + document_scope_clause=document_scope_clause, type_clause=type_clause, ) + """ @@ -656,6 +654,7 @@ async def map_unit_discovery( chunk_types=chunk_types, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, + document_scope=document_scope, revision_pins=revision_pins, ) logger.info( @@ -703,6 +702,7 @@ async def _hydrate_winning_units( chunk_types: set[str] | None, exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + document_scope: DocumentScope = DocumentScope(), revision_pins: Mapping[str, str] | None, ) -> list[dict[str, Any]]: """Map each winning leaf to its one chunk; asset requests follow connect_to.""" @@ -770,6 +770,7 @@ async def _hydrate_winning_units( rows=primaries, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, + document_scope=document_scope, revision_pins=revision_pins, ) connected_by_id = { diff --git a/packages/shared-python/shared/services/retrieval/search/scoped_corpus.py b/packages/shared-python/shared/services/retrieval/search/scoped_corpus.py index 0e087549..47f27d5a 100644 --- a/packages/shared-python/shared/services/retrieval/search/scoped_corpus.py +++ b/packages/shared-python/shared/services/retrieval/search/scoped_corpus.py @@ -1,5 +1,7 @@ from __future__ import annotations +from shared.services.retrieval.document_scope import DocumentScope + from collections.abc import Mapping from typing import Any @@ -76,6 +78,7 @@ async def count_scoped_chunks( namespace: str, exclude_document_ids: list[str], allowed_chunk_types: set[str] | None, + document_scope: DocumentScope = DocumentScope(), revision_pins: Mapping[str, str] | None = None, max_count: int | None = None, ) -> int: @@ -103,8 +106,7 @@ async def count_scoped_chunks( ) ) ) - if exclude_document_ids: - stmt = stmt.where(Document.document_id.notin_(list(exclude_document_ids))) + stmt = stmt.where(document_scope.excluding(exclude_document_ids).predicate(Document.document_id)) if allowed_chunk_types is not None: stmt = stmt.where(func.lower(DocumentChunk.chunk_type).in_(list(allowed_chunk_types))) @@ -124,6 +126,7 @@ async def load_all_scoped_chunks( namespace: str, exclude_document_ids: list[str], exclude_sections: list[dict[str, str]], + document_scope: DocumentScope = DocumentScope(), allowed_chunk_types: set[str] | None, signal_paths: list[str], filter_mode: str, @@ -152,8 +155,7 @@ async def load_all_scoped_chunks( ) if revision_pins is None: stmt = stmt.where(Document.status == 'active') - if exclude_document_ids: - stmt = stmt.where(Document.document_id.notin_(list(exclude_document_ids))) + stmt = stmt.where(document_scope.excluding(exclude_document_ids).predicate(Document.document_id)) if allowed_chunk_types is not None: stmt = stmt.where(func.lower(DocumentChunk.chunk_type).in_(list(allowed_chunk_types)))