From 5a25344a5569290f753f2cf64a505f86f56b89a3 Mon Sep 17 00:00:00 2001 From: Vikhyath Mondreti Date: Wed, 16 Sep 2026 17:38:29 -0700 Subject: [PATCH 1/2] fix(knowledge): bound KB block vector retrieval --- .github/workflows/test-build.yml | 1 + .../kb-block-search.integration.ts | 135 ++++++++++ apps/sim/lib/knowledge/search/queries.test.ts | 235 ++++++++++++++++-- apps/sim/lib/knowledge/search/queries.ts | 117 ++++++--- 4 files changed, 430 insertions(+), 58 deletions(-) create mode 100644 apps/sim/lib/knowledge/__integration__/kb-block-search.integration.ts diff --git a/.github/workflows/test-build.yml b/.github/workflows/test-build.yml index bb504e22fe2..c2ac8153946 100644 --- a/.github/workflows/test-build.yml +++ b/.github/workflows/test-build.yml @@ -221,6 +221,7 @@ jobs: lib/knowledge/__integration__/search-source-progress.integration.ts lib/knowledge/__integration__/search-source-pagination.integration.ts lib/knowledge/__integration__/search-reference-batching.integration.ts + lib/knowledge/__integration__/kb-block-search.integration.ts lib/core/outbox/service.integration.ts lib/knowledge/__integration__/connector-upload.integration.ts lib/uploads/contexts/organization-logo/application.integration.ts diff --git a/apps/sim/lib/knowledge/__integration__/kb-block-search.integration.ts b/apps/sim/lib/knowledge/__integration__/kb-block-search.integration.ts new file mode 100644 index 00000000000..dba85e74b3d --- /dev/null +++ b/apps/sim/lib/knowledge/__integration__/kb-block-search.integration.ts @@ -0,0 +1,135 @@ +/** KB block retrieval against disposable PostgreSQL, using a workspace API-key identity. */ +import type { Principal } from '@sim/auth/principal' +import { db } from '@sim/db' +import { document, embedding, knowledgeBase, organization, user, workspace } from '@sim/db/schema' +import { generateId } from '@sim/utils/id' +import { eq, inArray } from 'drizzle-orm' +import { afterAll, beforeAll, describe, expect, it } from 'vitest' +import { + createKnowledgeAclFixtureIds, + seedKnowledgeAclFixture, +} from '@/lib/knowledge/__integration__/seed-source-access-fixture' +import { createKnowledgeAccessProvider } from '@/lib/knowledge/access/scope' +import { retrieveKnowledgeSearch } from '@/lib/knowledge/search/queries' +import { embeddingVectorValues } from '@/lib/knowledge/vector-columns' + +describe('API-key KB block fan-out', () => { + const ids = createKnowledgeAclFixtureIds() + const bases = Array.from({ length: 18 }, () => ({ + id: generateId(), + visible: generateId(), + denied: generateId(), + excluded: generateId(), + })) + const principal: Principal = { + kind: 'workspace_api_key', + workspaceId: ids.workspaceId, + keyId: 'fixture-key', + } + const vector = [1, ...Array(1535).fill(0)] + const queryVector = { + vector: JSON.stringify(vector), + dimensions: 1536 as const, + model: 'text-embedding-3-small', + } + + beforeAll(async () => { + await seedKnowledgeAclFixture(ids, { connectorType: 'google_drive' }) + await db.insert(knowledgeBase).values( + bases.map((base, index) => ({ + id: base.id, + userId: ids.aliceId, + workspaceId: ids.workspaceId, + name: `KB block ${index}`, + })) + ) + await db.insert(document).values( + bases.flatMap((base) => + (['visible', 'denied', 'excluded'] as const).map((kind) => ({ + id: base[kind], + knowledgeBaseId: base.id, + filename: kind, + fileUrl: `https://fixture.invalid/${base[kind]}`, + fileSize: 12, + mimeType: 'text/plain', + processingStatus: 'completed', + acl: kind === 'denied' ? [`u:${ids.aliceId}@fixture.test`] : ['ws'], + userExcluded: kind === 'excluded', + })) + ) + ) + await db.insert(embedding).values( + bases.flatMap((base) => + (['visible', 'denied', 'excluded'] as const).map((kind) => ({ + id: generateId(), + documentId: base[kind], + knowledgeBaseId: base.id, + chunkIndex: 0, + chunkHash: base[kind], + content: `Fixture policy ${kind}`, + contentLength: 24, + tokenCount: 5, + startOffset: 0, + endOffset: 24, + tag1: 'policy', + ...embeddingVectorValues(1536, vector), + })) + ) + ) + }) + + afterAll(async () => { + await db.delete(workspace).where(eq(workspace.id, ids.workspaceId)) + await db.delete(organization).where(eq(organization.id, ids.organizationId)) + await db.delete(user).where(inArray(user.id, [ids.aliceId, ids.bobId])) + await db.$client.end() + }) + + it.each([false, true])( + 'completes 18 concurrent KB searches with access checks intact (tag filter: %s)', + async (withTags) => { + const previousDebug = db.$client.options.debug + const statements: string[] = [] + db.$client.options.debug = (_connection, query) => { + if (statements.length < 250) statements.push(query) + } + try { + const results = await Promise.all( + bases.map(async (base) => { + const accessProvider = createKnowledgeAccessProvider(principal, { + workspaceId: ids.workspaceId, + knowledgeBaseIds: [base.id], + }) + const access = await accessProvider.get() + expect(access.kind).toBe('workspace') + return retrieveKnowledgeSearch({ + knowledgeBaseIds: [base.id], + topK: 2, + access, + accessProvider, + searchMode: 'vector', + query: 'Find the fixture policy', + queryVector, + ...(withTags && { + structuredFilters: [ + { tagSlot: 'tag1', fieldType: 'text', operator: 'eq', value: 'policy' }, + ], + }), + }) + }) + ) + for (const [index, result] of results.entries()) { + expect(result.retrieval).toEqual({ status: 'complete', timedOutLegs: [] }) + expect(result.rows.map((row) => row.documentId)).toEqual([bases[index].visible]) + expect(result.rows[0].knowledgeBaseId).toBe(bases[index].id) + expect(result.rows[0].distance).toBeCloseTo(0) + } + expect(statements.filter((query) => query.includes('statement_timeout'))).toHaveLength(36) + expect(statements.filter((query) => query.includes('+ 0'))).toHaveLength(18) + expect(statements.some((query) => query.includes('hnsw.iterative_scan'))).toBe(false) + } finally { + db.$client.options.debug = previousDebug + } + } + ) +}) diff --git a/apps/sim/lib/knowledge/search/queries.test.ts b/apps/sim/lib/knowledge/search/queries.test.ts index a5f20dc7f20..b7a5b9ebabe 100644 --- a/apps/sim/lib/knowledge/search/queries.test.ts +++ b/apps/sim/lib/knowledge/search/queries.test.ts @@ -1,6 +1,7 @@ /** * @vitest-environment node */ +import { db } from '@sim/db' import { dbChainMockFns, hasMockCondition, @@ -15,7 +16,12 @@ import { WORKSPACE_ACCESS_TOKENS, } from '@/lib/knowledge/access/types' import { buildTagFilterCondition } from '@/lib/knowledge/documents/tag-filter' -import { SearchBudget } from '@/lib/knowledge/search/budget' +import { + SearchBudget, + SearchDeadlineError, + type SearchExecutor, +} from '@/lib/knowledge/search/budget' +import type { SearchStage } from '@/lib/knowledge/search/diagnostics' import { executeKeywordSearch, getStructuredTagFilters, @@ -310,7 +316,174 @@ describe('getStructuredTagFilters', () => { }) }) +describe('KB block vector retrieval', () => { + const params: SearchParams = { + knowledgeBaseIds: ['kb-small'], + topK: 2, + access: { kind: 'workspace', tokens: WORKSPACE_ACCESS_TOKENS }, + queryVector: { vector: '[0.1,0.2]', dimensions: 1536, model: 'text-embedding-3-small' }, + distanceThreshold: 1, + } + + beforeEach(() => resetDbChainMock()) + afterEach(() => { + vi.restoreAllMocks() + vi.useRealTimers() + }) + + it.each([handleVectorOnlySearch, handleTagAndVectorSearch])( + 'does not acquire a connection or start SQL after the KB retrieval deadline', + async (search) => { + const budget = new SearchBudget('vector', performance.now() - 1) + expect( + await search({ + ...params, + budget, + structuredFilters: [ + { tagSlot: 'tag1', fieldType: 'text', operator: 'eq', value: 'release' }, + ], + }) + ).toEqual([]) + expect(budget.timedOut).toBe(true) + expect(dbChainMockFns.transaction).not.toHaveBeenCalled() + expect(dbChainMockFns.select).not.toHaveBeenCalled() + } + ) + + it('ranks all candidates in a small KB exactly instead of traversing the shared vector index', async () => { + queueTableRows(schemaMock.embedding, [{ id: 'near' }, { id: 'far' }]) + queueTableRows(schemaMock.embedding, [ + { id: 'far', distance: 0.2 }, + { id: 'near', distance: 0.1 }, + ]) + const rows = await handleVectorOnlySearch(params) + expect(rows.map((row) => row.id)).toEqual(['near', 'far']) + expect(dbChainMockFns.execute).not.toHaveBeenCalled() + expect(Object.keys(dbChainMockFns.select.mock.calls[0][0])).toEqual(['id']) + expect(dbChainMockFns.limit).toHaveBeenNthCalledWith(1, 201) + expect(render(dbChainMockFns.orderBy.mock.calls[0][0]).sql).toContain('+ 0') + expect( + hasMockCondition( + dbChainMockFns.where.mock.calls[1][0], + (node) => + node.type === 'inArray' && + node.column === schemaMock.embedding.id && + JSON.stringify(node.values) === JSON.stringify(['near', 'far']) + ) + ).toBe(true) + }) + + it.each([1, 200, 201])( + 'reports a %i-candidate SQL timeout as partial, not an empty complete result', + async (count) => { + queueTableRows( + schemaMock.embedding, + Array.from({ length: count }, (_, index) => ({ id: `candidate-${index}` })) + ) + dbChainMockFns.orderBy + .mockImplementationOnce(dbChainMockFns.orderBy.getMockImplementation()!) + .mockRejectedValueOnce(new Error('Statement canceled', { cause: { code: '57014' } })) + const result = await retrieveKnowledgeSearch({ + ...params, + query: 'fixture policy', + searchMode: 'vector', + }) + expect(result).toEqual({ + rows: [], + retrieval: { status: 'partial', timedOutLegs: ['vector'] }, + }) + const usesAnn = dbChainMockFns.execute.mock.calls.some(([statement]) => + render(statement).sql.includes('hnsw.iterative_scan') + ) + expect(usesAnn).toBe(count > 200) + } + ) + + it('does not convert an unexpected ranking error into partial retrieval', async () => { + queueTableRows(schemaMock.embedding, [{ id: 'candidate' }]) + const failure = new Error('Connection lost', { cause: { code: '08006' } }) + dbChainMockFns.orderBy + .mockImplementationOnce(dbChainMockFns.orderBy.getMockImplementation()!) + .mockRejectedValueOnce(failure) + await expect( + retrieveKnowledgeSearch({ ...params, query: 'fixture policy', searchMode: 'vector' }) + ).rejects.toBe(failure) + }) + + it('reports incomplete retrieval for 18 expired pool waiters without starting their SQL later', async () => { + vi.useFakeTimers() + const release: Array<() => void> = [] + const transactions: Array> = [] + vi.spyOn(db, 'transaction').mockImplementation((callback) => { + const transaction = new Promise((resolve) => release.push(resolve)).then(() => + callback(db as never) + ) + transactions.push(transaction) + return transaction as ReturnType + }) + const pending = Promise.all( + Array.from({ length: 18 }, (_, index) => + retrieveKnowledgeSearch({ + ...params, + knowledgeBaseIds: [`kb-${index}`], + query: 'fixture policy', + searchMode: 'vector', + vectorBudgetMs: 50, + }) + ) + ) + await vi.advanceTimersByTimeAsync(60) + const results = await pending + expect(results).toHaveLength(18) + for (const result of results) { + expect(result).toEqual({ + rows: [], + retrieval: { status: 'partial', timedOutLegs: ['vector'] }, + }) + } + for (const resume of release) resume() + const settled = await Promise.allSettled(transactions) + expect(settled).toHaveLength(18) + for (const transaction of settled) { + expect(transaction.status).toBe('rejected') + if (transaction.status === 'rejected') + expect(transaction.reason).toBeInstanceOf(SearchDeadlineError) + } + expect(dbChainMockFns.select).not.toHaveBeenCalled() + expect(dbChainMockFns.execute).not.toHaveBeenCalled() + }) + + it.each([1, 201])( + 'shares the remaining SQL budget between the probe and %i-candidate ranking', + async (count) => { + vi.spyOn(performance, 'now').mockReturnValue(0) + queueTableRows( + schemaMock.embedding, + Array.from({ length: count }, (_, index) => ({ id: `candidate-${index}` })) + ) + const query = SearchBudget.prototype.query + vi.spyOn(SearchBudget.prototype, 'query').mockImplementation(async function ( + this: SearchBudget, + stage: SearchStage, + run: (executor: SearchExecutor) => PromiseLike + ) { + const runQuery: SearchBudget['query'] = query.bind(this) + const result = await runQuery(stage, run) + if (stage === 'vector.probe') vi.spyOn(performance, 'now').mockReturnValue(30) + return result + }) + await handleVectorOnlySearch({ ...params, budget: new SearchBudget('vector', 100) }) + const timeouts = dbChainMockFns.execute.mock.calls + .map(([statement]) => render(statement)) + .filter((statement) => statement.sql.includes('statement_timeout')) + .map((statement) => statement.params[0]) + expect(timeouts).toEqual(count === 1 ? ['100', '70'] : ['100', '70', '70']) + } + ) +}) + describe('vector scan settings', () => { + const largeProbe = Array.from({ length: 201 }, (_, index) => ({ id: `probe-${index}` })) const params: SearchParams = { knowledgeBaseIds: ['kb-small'], topK: 2, @@ -321,13 +494,14 @@ describe('vector scan settings', () => { beforeEach(() => { resetDbChainMock() + queueTableRows(schemaMock.embedding, largeProbe) }) afterEach(() => { vi.useRealTimers() }) - it('tunes a small workspace search before querying its KB scope, preserving distance ordering', async () => { + it('tunes an overflowing KB scope without limiting ANN to the probe prefix', async () => { queueTableRows(schemaMock.embedding, [ { id: 'far', distance: 0.2 }, { id: 'near', distance: 0.1 }, @@ -341,11 +515,11 @@ describe('vector scan settings', () => { params: ['20000'], }) expect(dbChainMockFns.execute.mock.invocationCallOrder[0]).toBeLessThan( - dbChainMockFns.select.mock.invocationCallOrder[0] + dbChainMockFns.select.mock.invocationCallOrder[1] ) expect( hasMockCondition( - dbChainMockFns.where.mock.calls[0][0], + dbChainMockFns.where.mock.calls[1][0], (node) => node.type === 'inArray' && node.column === schemaMock.embedding.knowledgeBaseId && @@ -354,10 +528,16 @@ describe('vector scan settings', () => { ) ).toBe(true) expect(dbChainMockFns.limit).toHaveBeenCalledWith(2) - expect(dbChainMockFns.limit).toHaveBeenCalledOnce() - expect(Object.keys(dbChainMockFns.select.mock.calls[0][0])).toEqual(['id', 'distance']) - const ranked = dbChainMockFns.from.mock.calls[1][0] - expect(dbChainMockFns.select.mock.calls[1][0].distance).toBe(ranked.distance) + expect(dbChainMockFns.limit).toHaveBeenCalledTimes(2) + expect( + hasMockCondition( + dbChainMockFns.where.mock.calls[1][0], + (node) => node.type === 'inArray' && node.column === schemaMock.embedding.id + ) + ).toBe(false) + expect(Object.keys(dbChainMockFns.select.mock.calls[1][0])).toEqual(['id', 'distance']) + const ranked = dbChainMockFns.from.mock.calls[2][0] + expect(dbChainMockFns.select.mock.calls[2][0].distance).toBe(ranked.distance) expect(dbChainMockFns.innerJoin).toHaveBeenCalledWith( schemaMock.embedding, expect.objectContaining({ @@ -367,20 +547,22 @@ describe('vector scan settings', () => { }) ) expect(dbChainMockFns.orderBy).toHaveBeenLastCalledWith(ranked.distance) - expect(dbChainMockFns.limit.mock.invocationCallOrder[0]).toBeLessThan( - dbChainMockFns.select.mock.invocationCallOrder[1] + expect(dbChainMockFns.limit.mock.invocationCallOrder[1]).toBeLessThan( + dbChainMockFns.select.mock.invocationCallOrder[2] ) }) - it('shares one local configuration across all KB vector legs and trims their sorted merge', async () => { + it('tunes each KB leg and trims their sorted merge', async () => { const knowledgeBaseIds = ['kb-1', 'kb-2', 'kb-3', 'kb-4', 'kb-5'] - for (let index = 0; index < knowledgeBaseIds.length; index++) + for (let index = 0; index < knowledgeBaseIds.length; index++) { + if (index > 0) queueTableRows(schemaMock.embedding, largeProbe) queueTableRows(schemaMock.embedding, [{ id: `row-${index}`, distance: (5 - index) / 10 }]) + } const rows = await handleVectorOnlySearch({ ...params, knowledgeBaseIds }) expect(rows.map((row) => row.id)).toEqual(['row-4', 'row-3']) - expect(dbChainMockFns.transaction).toHaveBeenCalledOnce() - expect(dbChainMockFns.execute).toHaveBeenCalledOnce() - expect(dbChainMockFns.select).toHaveBeenCalledTimes(10) + expect(dbChainMockFns.transaction).toHaveBeenCalledTimes(5) + expect(dbChainMockFns.execute).toHaveBeenCalledTimes(5) + expect(dbChainMockFns.select).toHaveBeenCalledTimes(15) for (const kbId of knowledgeBaseIds) expect( dbChainMockFns.where.mock.calls.some(([condition]) => @@ -432,20 +614,23 @@ describe('vector scan settings', () => { queueTableRows(schemaMock.embedding, [{ id: 'fallback', distance: 0.1 }]) expect((await handleVectorOnlySearch(params)).map((row) => row.id)).toEqual(['fallback']) expect(dbChainMockFns.execute).toHaveBeenCalledOnce() - expect(dbChainMockFns.select).toHaveBeenCalledTimes(2) + expect(dbChainMockFns.select).toHaveBeenCalledTimes(3) await handleVectorOnlySearch(params) expect(dbChainMockFns.transaction).toHaveBeenCalledOnce() await vi.advanceTimersByTimeAsync(10 * 60 * 1000 + 1) + queueTableRows(schemaMock.embedding, largeProbe) await handleVectorOnlySearch(params) expect(dbChainMockFns.transaction).toHaveBeenCalledTimes(2) expect(dbChainMockFns.execute).toHaveBeenCalledTimes(2) }) - it('propagates an unrelated settings failure without issuing a query or disabling later tuning', async () => { + it('propagates an unrelated settings failure without ranking or disabling later tuning', async () => { const failure = { code: '08006', message: 'Connection lost' } dbChainMockFns.execute.mockRejectedValueOnce(failure) await expect(handleVectorOnlySearch(params)).rejects.toBe(failure) - expect(dbChainMockFns.select).not.toHaveBeenCalled() + expect(dbChainMockFns.select).toHaveBeenCalledOnce() + expect(dbChainMockFns.orderBy).not.toHaveBeenCalled() + queueTableRows(schemaMock.embedding, largeProbe) await handleVectorOnlySearch(params) expect(dbChainMockFns.execute).toHaveBeenCalledTimes(2) }) @@ -456,7 +641,8 @@ describe('vector scan settings', () => { .mockImplementationOnce(dbChainMockFns.orderBy.getMockImplementation()!) .mockRejectedValueOnce(failure) await expect(handleVectorOnlySearch(params)).rejects.toBe(failure) - expect(dbChainMockFns.select).toHaveBeenCalledTimes(2) + expect(dbChainMockFns.select).toHaveBeenCalledTimes(3) + queueTableRows(schemaMock.embedding, largeProbe) await handleVectorOnlySearch(params) expect(dbChainMockFns.execute).toHaveBeenCalledTimes(2) }) @@ -477,9 +663,10 @@ describe('workspace search filters before ranking', () => { } beforeEach(() => resetDbChainMock()) - function expectScopeOnEveryQuery() { - expect(dbChainMockFns.where).toHaveBeenCalled() - for (const [condition] of dbChainMockFns.where.mock.calls) { + function expectScopeOnEveryQuery(skipIdentityProbe = false) { + const queries = dbChainMockFns.where.mock.calls.slice(skipIdentityProbe ? 1 : 0) + expect(queries.length).toBeGreaterThan(0) + for (const [condition] of queries) { expect( hasMockCondition( condition, @@ -512,13 +699,15 @@ describe('workspace search filters before ranking', () => { it.each([handleVectorOnlySearch, handleTagOnlySearch, handleTagAndVectorSearch])( 'applies the full document scope to vector and tag searches', async (search) => { + const hasIdentityProbe = search !== handleTagOnlySearch + if (hasIdentityProbe) queueTableRows(schemaMock.embedding, [{ id: 'candidate' }]) await search({ ...params, structuredFilters: [ { tagSlot: 'tag1', fieldType: 'text', operator: 'eq', value: 'launch' }, ], }) - expectScopeOnEveryQuery() + expectScopeOnEveryQuery(hasIdentityProbe) } ) diff --git a/apps/sim/lib/knowledge/search/queries.ts b/apps/sim/lib/knowledge/search/queries.ts index 84ace644039..e67a126fa5e 100644 --- a/apps/sim/lib/knowledge/search/queries.ts +++ b/apps/sim/lib/knowledge/search/queries.ts @@ -56,6 +56,7 @@ const CANDIDATE_HNSW_SCAN_MEM_MULTIPLIER = '2' const MIN_VECTOR_RERANK_CANDIDATES = 400 const MAX_VECTOR_RERANK_CANDIDATES = 1600 const VECTOR_RERANK_OVERSAMPLING = 8 +const MAX_EXACT_KB_VECTOR_CANDIDATES = 200 /** How long to stop trying the iterative-scan settings after the server rejected them. */ const HNSW_SETTINGS_UNSUPPORTED_RETRY_MS = 10 * 60 * 1000 @@ -777,40 +778,91 @@ export async function handleVectorOnlySearch(params: SearchParams): Promise - selectRankedVectorResults( - executor, - distance, - [ - kbScope, - ...getVisibilityConditions(access, params.filters), - sql`${distance} < ${distanceThreshold}`, - ], - limit - ) - /** * A relaxed-order iterative scan may hand rows back slightly out of distance * order, so both paths re-sort in memory before trimming to `topK`. */ if (strategy.useParallel) { const parallelLimit = Math.ceil(topK / knowledgeBaseIds.length) + 5 - const allResults = await withVectorScanSettings(async (executor) => { - const parallelResults = await Promise.all( - knowledgeBaseIds.map((kbId) => - vectorLeg(executor, eq(embedding.knowledgeBaseId, kbId), parallelLimit) - ) + const allResults: SearchResult[] = [] + /** Keep one active KB leg per request so multi-base searches cannot monopolize the pool. */ + for (const kbId of knowledgeBaseIds) { + allResults.push( + ...(await selectScopedVectorResults( + params, + distance, + eq(embedding.knowledgeBaseId, kbId), + parallelLimit + )) ) - return parallelResults.flat() - }) + if (params.budget?.timedOut) break + } return allResults.sort((a, b) => a.distance - b.distance).slice(0, topK) } - const rows = await withVectorScanSettings((executor) => - vectorLeg(executor, inArray(embedding.knowledgeBaseId, knowledgeBaseIds), topK) + const rows = await selectScopedVectorResults( + params, + distance, + inArray(embedding.knowledgeBaseId, knowledgeBaseIds), + topK ) return rows.sort((a, b) => a.distance - b.distance) } +/** + * KB runs without a human subject still need bounded small-scope ranking. Probe only chunk + * identities, then reapply every access and visibility predicate before ranking and hydration. + * An overflowing probe selects ANN over the whole scope, never a truncated candidate prefix. + */ +async function selectScopedVectorResults( + params: SearchParams, + distance: SQL, + kbScope: SQL | undefined, + limit: number, + tagConditions: (SQL | undefined)[] = [] +): Promise { + try { + const probe = await runSearchQuery(params.budget, 'vector.probe', (executor) => + executor + .select({ id: embedding.id }) + .from(embedding) + .where(and(kbScope, eq(embedding.enabled, true), ...tagConditions)) + .limit(MAX_EXACT_KB_VECTOR_CANDIDATES + 1) + ) + if (probe.length === 0) return [] + const conditions = [ + kbScope, + ...getVisibilityConditions(params.access, params.filters), + ...tagConditions, + sql`${distance} < ${params.distanceThreshold}`, + ] + if (probe.length <= MAX_EXACT_KB_VECTOR_CANDIDATES) { + annotateSearchDiagnostics({ vectorRanking: 'exact' }) + return await runSearchQuery(params.budget, 'vector.exact', (executor) => + selectRankedVectorResults( + executor, + distance, + [ + ...conditions, + inArray( + embedding.id, + probe.map((candidate) => candidate.id) + ), + ], + limit, + true + ) + ) + } + return await withVectorScanSettings( + (executor) => selectRankedVectorResults(executor, distance, conditions, limit), + params.budget + ) + } catch (error) { + if (!params.budget?.isTimeout(error)) throw error + return [] + } +} + /** * Bound ANN traversal and rerank a small candidate pool against the original vectors. * Nearest-neighbor traversal drives document visibility lookups, avoiding a sort of @@ -1023,14 +1075,15 @@ function selectRankedVectorResults( executor: SearchExecutor, distance: SQL, conditions: (SQL | undefined)[], - limit: number + limit: number, + exact = false ) { const ranked = executor .select({ id: embedding.id, distance: distance.as('distance') }) .from(embedding) .innerJoin(document, eq(embedding.documentId, document.id)) .where(and(...conditions)) - .orderBy(distance) + .orderBy(exact ? sql`(${distance}) + 0` : distance) .limit(limit) .as('ranked_embeddings') @@ -1317,18 +1370,12 @@ export async function handleTagAndVectorSearch(params: SearchParams): Promise - selectRankedVectorResults( - executor, - distance, - [ - inArray(embedding.knowledgeBaseId, knowledgeBaseIds), - ...getVisibilityConditions(access, params.filters), - ...tagFilterConditions, - sql`${distance} < ${distanceThreshold}`, - ], - topK - ) + const rows = await selectScopedVectorResults( + params, + distance, + inArray(embedding.knowledgeBaseId, knowledgeBaseIds), + topK, + tagFilterConditions ) return rows.sort((a, b) => a.distance - b.distance) } From f6dd97f7eefd1c617f094ea7bafdaaa1a6e1f63d Mon Sep 17 00:00:00 2001 From: Vikhyath Mondreti Date: Wed, 16 Sep 2026 17:49:40 -0700 Subject: [PATCH 2/2] fix(knowledge): align search fixtures with scoped probing --- .../app/api/knowledge/search/utils.test.ts | 64 +++++++++---------- 1 file changed, 31 insertions(+), 33 deletions(-) diff --git a/apps/sim/app/api/knowledge/search/utils.test.ts b/apps/sim/app/api/knowledge/search/utils.test.ts index 2c0a4abaf59..fc19016b170 100644 --- a/apps/sim/app/api/knowledge/search/utils.test.ts +++ b/apps/sim/app/api/knowledge/search/utils.test.ts @@ -212,6 +212,10 @@ describe('Knowledge Search Utils', () => { describe('handleTagAndVectorSearch', () => { it('returns only bounded ranked rows without first materializing every matching tag ID', async () => { resetDbChainMock() + queueTableRows( + schemaMock.embedding, + Array.from({ length: 201 }, (_, index) => ({ id: `candidate-${index}` })) + ) queueTableRows(schemaMock.embedding, [makeResult('second', 0.2), makeResult('first', 0.1)]) const results = await handleTagAndVectorSearch({ @@ -226,9 +230,11 @@ describe('Knowledge Search Utils', () => { }) expect(results.map((row) => row.id)).toEqual(['first', 'second']) - expect(dbChainMockFns.select).toHaveBeenCalledTimes(2) + expect(dbChainMockFns.select).toHaveBeenCalledTimes(3) expect(dbChainMockFns.as).toHaveBeenCalledWith('ranked_embeddings') - expect(dbChainMockFns.select.mock.calls[0][0]).toHaveProperty('distance') + expect(Object.keys(dbChainMockFns.select.mock.calls[0][0])).toEqual(['id']) + expect(dbChainMockFns.limit).toHaveBeenNthCalledWith(1, 201) + expect(dbChainMockFns.select.mock.calls[1][0]).toHaveProperty('distance') expect(dbChainMockFns.limit).toHaveBeenCalledWith(2) }) @@ -536,6 +542,7 @@ describe('Knowledge Search Utils', () => { }) it('runs a single retrieval leg in vector mode', async () => { + queueTableRows(schemaMock.embedding, [{ id: 'vector-hit' }]) queueTableRows(schemaMock.embedding, [makeResult('vector-hit')]) const results = await executeKnowledgeSearch({ @@ -548,20 +555,19 @@ describe('Knowledge Search Utils', () => { }) expect(results.map((r) => r.id)).toEqual(['vector-hit']) - expect(dbChainMockFns.select).toHaveBeenCalledTimes(2) + expect(dbChainMockFns.select).toHaveBeenCalledTimes(3) expect(dbChainMockFns.as).toHaveBeenCalledWith('ranked_embeddings') }) it('runs both legs and fuses them in hybrid mode', async () => { /** - * Chains dequeue in creation order. Hybrid legs over-fetch past the - * plain scan's candidate pool, so the vector leg opens its transaction - * and applies the scan settings before selecting: the keyword ranking - * pass is built first, then the vector select, then hydration. + * Chains dequeue in creation order: keyword ranking, the budgeted vector + * probe, keyword hydration, then vector ranking and hydration in one query. */ queueTableRows(schemaMock.embedding, [{ id: 'keyword-hit', keywordRank: 0.9 }]) - queueTableRows(schemaMock.embedding, [makeResult('vector-hit')]) + queueTableRows(schemaMock.embedding, [{ id: 'vector-hit' }]) queueTableRows(schemaMock.embedding, [makeResult('keyword-hit')]) + queueTableRows(schemaMock.embedding, [makeResult('vector-hit')]) const results = await executeKnowledgeSearch({ knowledgeBaseIds: ['kb-123'], @@ -573,39 +579,31 @@ describe('Knowledge Search Utils', () => { }) expect(results.map((r) => r.id).sort()).toEqual(['keyword-hit', 'vector-hit']) - expect(dbChainMockFns.select).toHaveBeenCalledTimes(4) + expect(dbChainMockFns.select).toHaveBeenCalledTimes(5) }) - it('falls back to vector results when the keyword leg fails', async () => { + it('propagates unexpected keyword errors after the vector leg finishes', async () => { /** The failing ranking chain is still built first and takes the first queued set. */ queueTableRows(schemaMock.embedding, [{ id: 'never-ranked', keywordRank: 0 }]) + queueTableRows(schemaMock.embedding, [{ id: 'vector-hit' }]) queueTableRows(schemaMock.embedding, [makeResult('vector-hit')]) - /** - * Both legs share one `orderBy` spy, so target the keyword leg by its - * ranking expression. Calling the untouched spy first captures the - * sentinel that tells the mock to build its normal chain, which the - * vector leg still needs. - */ - const chainDefault = dbChainMockFns.orderBy() - dbChainMockFns.orderBy.mockImplementation((fragment: unknown) => { - const text = (fragment as { strings?: string[] })?.strings?.join('') ?? '' - if (text.includes('ts_rank_cd')) { - throw new Error('tsquery blew up') - } - return chainDefault - }) - - const results = await executeKnowledgeSearch({ - knowledgeBaseIds: ['kb-123'], - access: WORKSPACE_ACCESS_SCOPE, - topK: 10, - searchMode: 'hybrid', - query: 'PROJ-1234', - queryVector: JSON.stringify([0.1, 0.2, 0.3]), + const failure = new Error('tsquery failed') + dbChainMockFns.orderBy.mockImplementationOnce(() => { + throw failure }) - expect(results.map((r) => r.id)).toEqual(['vector-hit']) + await expect( + executeKnowledgeSearch({ + knowledgeBaseIds: ['kb-123'], + access: WORKSPACE_ACCESS_SCOPE, + topK: 10, + searchMode: 'hybrid', + query: 'PROJ-1234', + queryVector: JSON.stringify([0.1, 0.2, 0.3]), + }) + ).rejects.toBe(failure) + expect(dbChainMockFns.as).toHaveBeenCalledWith('ranked_embeddings') }) it('skips both query legs when only tag filters are provided', async () => {