diff --git a/README.md b/README.md index 440bd83..3885e6d 100644 --- a/README.md +++ b/README.md @@ -142,13 +142,15 @@ for document in documents.documents: print(documents.pagination.total_pages) ``` -Retrieval supports exclusions when clients want follow-up results that avoid -previously used documents or sections: +Retrieval can limit documents for one request and exclude documents or sections. +Omitting `include_document_ids` leaves documents unrestricted by inclusion; +passing `[]` matches no documents. Exclusions take precedence over inclusions. ```python response = client.retrieval.query( namespace="support-center", query="battery charging", + include_document_ids=["doc_123", "doc_old"], exclude_document_ids=["doc_old"], exclude_sections=[ {"document_id": "doc_123", "section_path": "Appendix / Legal"} diff --git a/docs/usage.md b/docs/usage.md index 8b60424..d78e8a2 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -585,15 +585,18 @@ result.source.source_file_name result.source.section_path ``` -### Exclude documents or sections +### Limit or exclude documents or sections -Use exclusions for follow-up queries that should avoid already-used context. +Use `include_document_ids` to search only the supplied documents for this +request. Omit it to leave documents unrestricted by inclusion; pass `[]` to +match no documents. Exclusions take precedence over inclusions. ```python response = client.retrieval.query( namespace="support-center", query="battery charging", top_k=10, + include_document_ids=["doc_123", "doc_old"], exclude_document_ids=["doc_old"], exclude_sections=[ {"document_id": "doc_123", "section_path": "Appendix / Legal"} diff --git a/src/knowhere/resources/retrieval.py b/src/knowhere/resources/retrieval.py index 6061edf..2ce427e 100644 --- a/src/knowhere/resources/retrieval.py +++ b/src/knowhere/resources/retrieval.py @@ -34,6 +34,7 @@ def query( rerank: Optional[bool] = None, threshold: Optional[float] = None, internal_recall_k: Optional[int] = None, + include_document_ids: Optional[list[str]] = None, exclude_document_ids: Optional[list[str]] = None, exclude_sections: Optional[list[RetrievalSectionExclusion]] = None, llm_config: Optional[LLMConfig] = None, @@ -64,6 +65,8 @@ def query( body["threshold"] = threshold if internal_recall_k is not None: body["internal_recall_k"] = internal_recall_k + if include_document_ids is not None: + body["include_document_ids"] = include_document_ids if exclude_document_ids is not None: body["exclude_document_ids"] = exclude_document_ids if exclude_sections is not None: @@ -98,6 +101,7 @@ async def query( rerank: Optional[bool] = None, threshold: Optional[float] = None, internal_recall_k: Optional[int] = None, + include_document_ids: Optional[list[str]] = None, exclude_document_ids: Optional[list[str]] = None, exclude_sections: Optional[list[RetrievalSectionExclusion]] = None, llm_config: Optional[LLMConfig] = None, @@ -128,6 +132,8 @@ async def query( body["threshold"] = threshold if internal_recall_k is not None: body["internal_recall_k"] = internal_recall_k + if include_document_ids is not None: + body["include_document_ids"] = include_document_ids if exclude_document_ids is not None: body["exclude_document_ids"] = exclude_document_ids if exclude_sections is not None: diff --git a/tests/test_retrieval.py b/tests/test_retrieval.py index 6529626..515fb3b 100644 --- a/tests/test_retrieval.py +++ b/tests/test_retrieval.py @@ -105,6 +105,7 @@ def test_query_sends_request_and_returns_results(self, sync_client: Any) -> None rerank=True, threshold=0.2, internal_recall_k=25, + include_document_ids=["doc_123", "doc_old"], exclude_document_ids=["doc_old"], exclude_sections=[ { @@ -128,6 +129,7 @@ def test_query_sends_request_and_returns_results(self, sync_client: Any) -> None "rerank": True, "threshold": 0.2, "internal_recall_k": 25, + "include_document_ids": ["doc_123", "doc_old"], "exclude_document_ids": ["doc_old"], "exclude_sections": [ { @@ -239,6 +241,27 @@ def test_query_omits_defaulted_optional_fields(self, sync_client: Any) -> None: request_body: Dict[str, Any] = json.loads(route.calls[0].request.read()) assert request_body == {"query": "refund policy"} + assert "include_document_ids" not in request_body + assert "exclude_document_ids" not in request_body + + @respx.mock + def test_query_preserves_empty_document_scope_arrays(self, sync_client: Any) -> None: + route = respx.post(RETRIEVAL_QUERY_URL).mock( + return_value=httpx.Response(200, json=_make_retrieval_response()) + ) + + sync_client.retrieval.query( + query="refund policy", + include_document_ids=[], + exclude_document_ids=[], + ) + + request_body: Dict[str, Any] = json.loads(route.calls[0].request.read()) + assert request_body == { + "query": "refund policy", + "include_document_ids": [], + "exclude_document_ids": [], + } @respx.mock @pytest.mark.asyncio