diff --git a/.changeset/retrieval-include-document-ids.md b/.changeset/retrieval-include-document-ids.md new file mode 100644 index 0000000..84204e9 --- /dev/null +++ b/.changeset/retrieval-include-document-ids.md @@ -0,0 +1,5 @@ +--- +"@ontos-ai/knowhere-sdk": minor +--- + +Add `includeDocumentIds` to retrieval queries so callers can restrict one request to an explicit document set. Empty arrays are preserved on the wire; omitted means unrestricted; exclusions still win. diff --git a/README.md b/README.md index b30a13a..b981054 100644 --- a/README.md +++ b/README.md @@ -515,12 +515,17 @@ without creating SDK local-disk cache state. Use `syncParsedDocument(...)` to explicitly resume or retry parsed-storage sync for an existing `documentId`, `jobId`, or local parsed result. -Follow-up queries can exclude documents or sections for one request: +Retrieval queries can limit documents and exclude documents or sections for one request. +`includeDocumentIds?: string[]` restricts retrieval to the supplied document IDs. +Omitting it leaves documents unrestricted by inclusion; passing `[]` matches no +documents. Exclusions take precedence over inclusions, including when a document +ID appears in both `includeDocumentIds` and `excludeDocumentIds`. ```typescript const followUp = await client.retrieval.query({ namespace: 'support-center', query: 'battery charging', + includeDocumentIds: ['doc_123', 'doc_old'], excludeDocumentIds: ['doc_old'], excludeSections: [{ documentId: 'doc_123', sectionPath: 'Appendix / Legal' }], }); diff --git a/src/resources/__tests__/retrieval-wire.test.ts b/src/resources/__tests__/retrieval-wire.test.ts new file mode 100644 index 0000000..54aa7b6 --- /dev/null +++ b/src/resources/__tests__/retrieval-wire.test.ts @@ -0,0 +1,83 @@ +import { createServer } from 'node:http'; +import { once } from 'node:events'; +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { Knowhere, type RetrievalQueryParams } from '../../index.js'; + +describe.each(['apiKey', 'authTokenProvider'] as const)('Retrieval wire payload (%s)', (auth) => { + beforeEach(() => { + vi.stubEnv('KNOWHERE_API_KEY', ''); + }); + + afterEach(() => { + vi.unstubAllEnvs(); + }); + + it.each<{ name: string; params: RetrievalQueryParams; body: string }>([ + { + name: 'includes overlapping include/exclude IDs and section exclusions', + params: { + query: 'battery charging', + includeDocumentIds: ['doc_123', 'doc_old'], + excludeDocumentIds: ['doc_old'], + excludeSections: [{ documentId: 'doc_123', sectionPath: 'Appendix / Legal' }], + }, + body: '{"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"}]}', + }, + { + name: 'preserves empty inclusion and exclusion arrays', + params: { query: 'battery charging', includeDocumentIds: [], excludeDocumentIds: [] }, + body: '{"query":"battery charging","include_document_ids":[],"exclude_document_ids":[]}', + }, + { + name: 'omits document filters when not supplied', + params: { query: 'battery charging' }, + body: '{"query":"battery charging"}', + }, + ])('$name', async ({ params, body }) => { + let receivedBody = ''; + let receivedMethod: string | undefined; + let receivedUrl: string | undefined; + let receivedAuthorization: string | undefined; + const server = createServer((request, response) => { + receivedMethod = request.method; + receivedUrl = request.url; + receivedAuthorization = request.headers.authorization; + request.setEncoding('utf8'); + request.on('data', (chunk: string) => { + receivedBody += chunk; + }); + request.on('end', () => { + response.writeHead(200, { 'Content-Type': 'application/json' }); + response.end('{"results":[]}'); + }); + }); + + try { + server.listen(0, '127.0.0.1'); + await once(server, 'listening'); + const address = server.address(); + if (!address || typeof address === 'string') throw new Error('Expected TCP server address'); + const client = new Knowhere({ + baseURL: `http://127.0.0.1:${address.port}`, + ...(auth === 'apiKey' + ? { apiKey: 'test-key' } + : { authTokenProvider: (): string => 'test-token' }), + maxRetries: 0, + }); + + await client.retrieval.query(params); + + expect(receivedMethod).toBe('POST'); + expect(receivedUrl).toBe('/v2/retrieval/query'); + expect(receivedAuthorization).toBe( + auth === 'apiKey' ? 'Bearer test-key' : 'Bearer test-token', + ); + expect(receivedBody).toBe(body); + } finally { + await new Promise((resolve, reject) => { + server.close((error) => (error ? reject(error) : resolve())); + server.closeAllConnections(); + }); + } + }); +}); diff --git a/src/resources/__tests__/retrieval.test.ts b/src/resources/__tests__/retrieval.test.ts index cf10fbb..60ce4c1 100644 --- a/src/resources/__tests__/retrieval.test.ts +++ b/src/resources/__tests__/retrieval.test.ts @@ -50,6 +50,7 @@ describe('Retrieval Resource', () => { rerank: true, threshold: 0.2, internalRecallK: 25, + includeDocumentIds: ['doc-123', 'doc-old'], excludeDocumentIds: ['doc-old'], excludeSections: [ { @@ -72,6 +73,7 @@ describe('Retrieval Resource', () => { rerank: true, threshold: 0.2, internalRecallK: 25, + includeDocumentIds: ['doc-123', 'doc-old'], excludeDocumentIds: ['doc-old'], excludeSections: [ { diff --git a/src/types/retrieval.ts b/src/types/retrieval.ts index 6555254..2038211 100644 --- a/src/types/retrieval.ts +++ b/src/types/retrieval.ts @@ -58,7 +58,13 @@ export interface RetrievalQueryParams { threshold?: number; /** Override the internal per-channel recall count */ internalRecallK?: number; - /** Documents to exclude for this request only */ + /** + * Limit retrieval to these documents for this request only. + * Omitted means no document inclusion restriction; [] matches no documents. + * Exclusions take precedence over inclusions. + */ + includeDocumentIds?: string[]; + /** Documents to exclude for this request only, even when included in includeDocumentIds. */ excludeDocumentIds?: string[]; /** Document sections to exclude for this request only */ excludeSections?: RetrievalSectionExclusion[];