Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .github/workflows/test-build.yml
Original file line number Diff line number Diff line change
Expand Up @@ -187,6 +187,7 @@ jobs:
bunx vitest run --mode integration
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/core/outbox/service.integration.ts
lib/knowledge/__integration__/connector-upload.integration.ts

Expand Down
6 changes: 4 additions & 2 deletions apps/sim/app/api/v2/knowledge/search/route.provenance.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,8 @@ vi.mock('@/lib/knowledge/application/contexts', () => ({
}))

vi.mock('@/lib/knowledge/service', () => ({
getActiveKnowledgeBaseReference: mocks.getKnowledgeBase,
getActiveKnowledgeBaseReferences: (ids: string[]) =>
Promise.all(ids.map((id) => mocks.getKnowledgeBase(id))),
}))

vi.mock('@/lib/knowledge/embeddings', () => ({
Expand All @@ -63,7 +64,8 @@ vi.mock('@/lib/knowledge/search/queries', () => ({
}))

vi.mock('@/lib/knowledge/tags/service', () => ({
getDocumentTagDefinitions: mocks.getTagDefinitions,
getDocumentTagDefinitionsByKnowledgeBaseIds: async (ids: string[]) =>
new Map(await Promise.all(ids.map(async (id) => [id, await mocks.getTagDefinitions(id)]))),
}))

vi.mock('@/lib/knowledge/tags/utils', () => ({
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,123 @@
import { db } from '@sim/db'
import {
knowledgeBase,
knowledgeBaseTagDefinitions,
organization,
user,
workspace,
} from '@sim/db/schema'
import { generateId } from '@sim/utils/id'
import { and, eq, inArray } from 'drizzle-orm'
import { afterAll, beforeAll, describe, expect, it, vi } from 'vitest'
import {
createKnowledgeAclFixtureIds,
seedKnowledgeAclFixture,
} from '@/lib/knowledge/__integration__/seed-source-access-fixture'
import {
getActiveKnowledgeBaseReference,
getActiveKnowledgeBaseReferences,
} from '@/lib/knowledge/service'
import {
getDocumentTagDefinitions,
getDocumentTagDefinitionsByKnowledgeBaseIds,
} from '@/lib/knowledge/tags/service'

describe('batched search reference reads', () => {
const fixture = createKnowledgeAclFixtureIds()
const baseIds = [fixture.knowledgeBaseId, ...Array.from({ length: 19 }, () => generateId())]
const deletedBaseId = generateId()

beforeAll(async () => {
await seedKnowledgeAclFixture(fixture)
await db.insert(knowledgeBase).values(
[...baseIds.slice(1), deletedBaseId].map((id, index) => ({
id,
userId: fixture.aliceId,
workspaceId: fixture.workspaceId,
name: `Batched search fixture ${index}`,
deletedAt: id === deletedBaseId ? new Date() : null,
}))
)
await db.insert(knowledgeBaseTagDefinitions).values(
baseIds.slice(0, -1).flatMap((knowledgeBaseId) =>
(['tag3', 'tag2'] as const).map((tagSlot) => ({
id: generateId(),
knowledgeBaseId,
tagSlot,
displayName: tagSlot,
fieldType: 'text' as const,
}))
)
)
})

afterAll(async () => {
vi.restoreAllMocks()
await db.delete(workspace).where(eq(workspace.id, fixture.workspaceId))
await db.delete(organization).where(eq(organization.id, fixture.organizationId))
await db.delete(user).where(inArray(user.id, [fixture.aliceId, fixture.bobId]))
await db.$client.end()
})

it('returns identical active references with one query instead of twenty', async () => {
const select = vi.spyOn(db, 'select')
try {
const expected = await Promise.all(baseIds.map(getActiveKnowledgeBaseReference))
expect(select).toHaveBeenCalledTimes(20)
select.mockClear()
expect(await getActiveKnowledgeBaseReferences(baseIds)).toEqual(expected)
expect(select).toHaveBeenCalledOnce()
} finally {
select.mockRestore()
}
})

it('preserves input order, duplicates, missing identities, and soft deletion', async () => {
const ids = [baseIds[8], generateId(), baseIds[0], deletedBaseId, baseIds[8]]
expect(await getActiveKnowledgeBaseReferences(ids)).toEqual(
await Promise.all(ids.map(getActiveKnowledgeBaseReference))
)
})

it('returns identical ordered tag definitions with one query instead of twenty', async () => {
const select = vi.spyOn(db, 'select')
try {
const expected = new Map(
await Promise.all(
baseIds.map(async (id) => [id, await getDocumentTagDefinitions(id)] as const)
)
)
expect(select).toHaveBeenCalledTimes(20)
select.mockClear()
expect(await getDocumentTagDefinitionsByKnowledgeBaseIds(baseIds)).toEqual(expected)
expect(select).toHaveBeenCalledOnce()
expect(expected.get(baseIds.at(-1)!)).toEqual([])
} finally {
select.mockRestore()
}
})

it('does not cache updated references or tag definitions across reads', async () => {
await getActiveKnowledgeBaseReferences(baseIds)
await getDocumentTagDefinitionsByKnowledgeBaseIds(baseIds)
await db
.update(knowledgeBase)
.set({ name: 'Updated reference' })
.where(eq(knowledgeBase.id, baseIds[0]))
await db
.update(knowledgeBaseTagDefinitions)
.set({ displayName: 'Updated definition' })
.where(
and(
eq(knowledgeBaseTagDefinitions.knowledgeBaseId, baseIds[0]),
eq(knowledgeBaseTagDefinitions.tagSlot, 'tag1')
)
)
const references = await getActiveKnowledgeBaseReferences(baseIds)
const tags = await getDocumentTagDefinitionsByKnowledgeBaseIds(baseIds)
expect(references[0]?.name).toBe('Updated reference')
expect(
tags.get(baseIds[0])?.find((definition) => definition.tagSlot === 'tag1')?.displayName
).toBe('Updated definition')
})
})
4 changes: 4 additions & 0 deletions apps/sim/lib/knowledge/application/documents.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,10 @@ vi.mock('@/lib/knowledge/documents/service', () => ({

vi.mock('@/lib/knowledge/tags/service', () => ({
getDocumentTagDefinitions: mocks.getDocumentTagDefinitions,
getDocumentTagDefinitionsByKnowledgeBaseIds: async (ids: string[]) =>
new Map(
await Promise.all(ids.map(async (id) => [id, await mocks.getDocumentTagDefinitions(id)]))
),
}))

vi.mock('@/lib/knowledge/orchestration/documents', () => ({
Expand Down
125 changes: 123 additions & 2 deletions apps/sim/lib/knowledge/application/search.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -13,13 +13,15 @@ const mocks = vi.hoisted(() => ({
requireOrganizationSearch: vi.fn(),
resolvePermission: vi.fn(),
getKnowledgeBase: vi.fn(),
getKnowledgeBases: vi.fn(),
resolveBilling: vi.fn(),
checkUsage: vi.fn(),
checkActorUsage: vi.fn(),
generateEmbedding: vi.fn(),
executeSearch: vi.fn(),
getDocumentMetadata: vi.fn(),
getTagDefinitions: vi.fn(),
getTagDefinitionsBatch: vi.fn(),
recordEmbeddingUsage: vi.fn(),
importProvenance: vi.fn(),
rerank: vi.fn(),
Expand Down Expand Up @@ -72,7 +74,7 @@ vi.mock('@/lib/permission-groups/resolve.server', () => ({
}))

vi.mock('@/lib/knowledge/service', () => ({
getActiveKnowledgeBaseReference: mocks.getKnowledgeBase,
getActiveKnowledgeBaseReferences: mocks.getKnowledgeBases,
}))

vi.mock('@/lib/knowledge/embeddings', () => ({
Expand All @@ -87,7 +89,7 @@ vi.mock('@/lib/knowledge/search/queries', () => ({
}))

vi.mock('@/lib/knowledge/tags/service', () => ({
getDocumentTagDefinitions: mocks.getTagDefinitions,
getDocumentTagDefinitionsByKnowledgeBaseIds: mocks.getTagDefinitionsBatch,
}))

vi.mock('@/lib/knowledge/tags/utils', () => ({
Expand Down Expand Up @@ -130,6 +132,13 @@ describe('knowledge search application use case', () => {
mocks.resolveWorkspace.mockResolvedValue(workspace)
mocks.resolvePermission.mockResolvedValue('read')
mocks.getKnowledgeBase.mockResolvedValue(knowledgeBase)
mocks.getKnowledgeBases.mockImplementation((ids: string[]) =>
Promise.all(ids.map((id) => mocks.getKnowledgeBase(id)))
)
mocks.getTagDefinitionsBatch.mockImplementation(
async (ids: string[]) =>
new Map(await Promise.all(ids.map(async (id) => [id, await mocks.getTagDefinitions(id)])))
)
mocks.resolveBilling.mockResolvedValue({
actorUserId: 'user-1',
workspaceId: 'workspace-1',
Expand Down Expand Up @@ -554,6 +563,118 @@ describe('knowledge search application use case', () => {

expect(mocks.resolveWorkspace).not.toHaveBeenCalled()
expect(mocks.getKnowledgeBase).not.toHaveBeenCalled()
expect(mocks.getKnowledgeBases).not.toHaveBeenCalled()
expect(mocks.getTagDefinitionsBatch).not.toHaveBeenCalled()
})

it('loads references and tags once for twenty bases while preserving requested order', async () => {
const ids = Array.from({ length: 20 }, (_, index) => `knowledge-${20 - index}`)
mocks.getKnowledgeBases.mockResolvedValue(ids.map((id) => ({ ...knowledgeBase, id })))
mocks.getTagDefinitionsBatch.mockResolvedValue(new Map(ids.map((id) => [id, []])))

const result = await searchKnowledge.execute({
principal: { kind: 'session', userId: 'user-1', sessionId: 'session-1' },
input: { knowledgeBaseIds: ids, query: 'answer', topK: 5 },
})

expect(mocks.getKnowledgeBases).toHaveBeenCalledExactlyOnceWith(ids)
expect(mocks.getTagDefinitionsBatch).toHaveBeenCalledExactlyOnceWith(ids)
expect(result.knowledgeBaseIds).toEqual(ids)
expect(result.knowledgeBaseId).toBe(ids[0])
expect(result.knowledgeBases.map((base) => base.id)).toEqual(ids)
expect(mocks.executeSearch).toHaveBeenCalledWith(
expect.objectContaining({ knowledgeBaseIds: ids })
)
})

it('preserves duplicate requested bases in retrieval and the response', async () => {
const ids = ['knowledge-2', 'knowledge-1', 'knowledge-2']
mocks.getKnowledgeBases.mockResolvedValue(ids.map((id) => ({ ...knowledgeBase, id })))

const result = await searchKnowledge.execute({
principal: { kind: 'session', userId: 'user-1', sessionId: 'session-1' },
input: { knowledgeBaseIds: ids, query: 'answer', topK: 5 },
})

expect(result.knowledgeBaseIds).toEqual(ids)
expect(mocks.executeSearch).toHaveBeenCalledWith(
expect.objectContaining({ knowledgeBaseIds: ids })
)
})

it('preserves missing-id order and duplicates in the concealed error before authorization or billing', async () => {
mocks.getKnowledgeBases.mockResolvedValue([null, knowledgeBase, null, null])

await expect(
searchKnowledge.execute({
principal: { kind: 'session', userId: 'user-1', sessionId: 'session-1' },
input: {
knowledgeBaseIds: ['missing-2', 'knowledge-1', 'missing-1', 'missing-2'],
query: 'answer',
topK: 5,
},
})
).rejects.toMatchObject({
code: 'not_found',
message: 'Knowledge bases not found or access denied: missing-2, missing-1, missing-2',
})
expect(mocks.resolvePermission).not.toHaveBeenCalled()
expect(mocks.resolveBilling).not.toHaveBeenCalled()
expect(mocks.executeSearch).not.toHaveBeenCalled()
})

it('rejects a batch spanning different canonical workspaces before billing', async () => {
mocks.getKnowledgeBases.mockResolvedValue([
knowledgeBase,
{ ...knowledgeBase, id: 'knowledge-2', workspaceId: 'workspace-2' },
])

await expect(
searchKnowledge.execute({
principal: { kind: 'session', userId: 'user-1', sessionId: 'session-1' },
input: { knowledgeBaseIds: ['knowledge-1', 'knowledge-2'], query: 'answer', topK: 5 },
})
).rejects.toMatchObject({
code: 'validation',
message: 'Selected knowledge bases must belong to the same workspace',
})
expect(mocks.resolveBilling).not.toHaveBeenCalled()
expect(mocks.executeSearch).not.toHaveBeenCalled()
})

it('reuses the tag filter batch when naming result metadata', async () => {
const ids = ['knowledge-1', 'knowledge-2']
mocks.getKnowledgeBases.mockResolvedValue(ids.map((id) => ({ ...knowledgeBase, id })))
mocks.getTagDefinitionsBatch.mockResolvedValue(
new Map(
ids.map((id) => [
id,
[{ knowledgeBaseId: id, tagSlot: 'tag1', displayName: 'team', fieldType: 'text' }],
])
)
)
mocks.executeSearch.mockResolvedValue([
{
id: 'chunk-1',
documentId: 'document-1',
knowledgeBaseId: ids[0],
content: 'answer',
tag1: 'docs',
},
])

const result = await searchKnowledge.execute({
principal: { kind: 'session', userId: 'user-1', sessionId: 'session-1' },
input: {
knowledgeBaseIds: ids,
topK: 5,
tagFilters: [{ tagName: 'team', operator: 'eq', value: 'docs' }],
},
})

expect(mocks.getTagDefinitionsBatch).toHaveBeenCalledExactlyOnceWith(ids)
expect(result.results[0].metadata).toEqual({ team: 'docs' })
expect(mocks.generateEmbedding).not.toHaveBeenCalled()
})

it('rejects multi-knowledge-base tag filters without embedding spend', async () => {
Expand Down
28 changes: 12 additions & 16 deletions apps/sim/lib/knowledge/application/search.ts
Original file line number Diff line number Diff line change
Expand Up @@ -49,13 +49,13 @@ import {
import { importKnowledgeSearchResultSecretProvenance } from '@/lib/knowledge/secret-provenance'
import {
type ActiveKnowledgeBaseReference,
getActiveKnowledgeBaseReference,
getActiveKnowledgeBaseReferences,
} from '@/lib/knowledge/service'
import {
type KnowledgeTagNameFilter,
resolveKnowledgeTagFilters,
} from '@/lib/knowledge/tags/filter-resolution'
import { getDocumentTagDefinitions } from '@/lib/knowledge/tags/service'
import { getDocumentTagDefinitionsByKnowledgeBaseIds } from '@/lib/knowledge/tags/service'
import type { DocumentTagDefinition } from '@/lib/knowledge/tags/types'
import type { StructuredFilter } from '@/lib/knowledge/types'
import { estimateTokenCount } from '@/lib/tokenization/estimators'
Expand Down Expand Up @@ -189,9 +189,7 @@ async function resolveKnowledgeSearchContext(
`topK must be an integer between 1 and ${KNOWLEDGE_SEARCH_COST_POLICY.maxTopK}`
)
}
const knowledgeBases = await Promise.all(
input.knowledgeBaseIds.map(getActiveKnowledgeBaseReference)
)
const knowledgeBases = await getActiveKnowledgeBaseReferences(input.knowledgeBaseIds)
const missingIds = input.knowledgeBaseIds.filter((_, index) => {
const knowledgeBase = knowledgeBases[index]
return !knowledgeBase || (!knowledgeBase.workspaceId && !knowledgeBase.organizationId)
Expand Down Expand Up @@ -558,18 +556,16 @@ export const searchKnowledge = defineAuthorizedKnowledgeUseCase({
}
}

const tagDefinitionEntries = await Promise.all(
knowledgeBaseIds.map(async (knowledgeBaseId) => {
const definitions =
definitionsByKnowledgeBase.get(knowledgeBaseId) ??
(await getDocumentTagDefinitions(knowledgeBaseId))
return [
knowledgeBaseId,
new Map(definitions.map((definition) => [definition.tagSlot, definition.displayName])),
] as const
})
if (filters.length === 0) {
definitionsByKnowledgeBase =
await getDocumentTagDefinitionsByKnowledgeBaseIds(knowledgeBaseIds)
}
const tagMaps = new Map(
[...definitionsByKnowledgeBase].map(([knowledgeBaseId, definitions]) => [
knowledgeBaseId,
new Map(definitions.map((definition) => [definition.tagSlot, definition.displayName])),
])
)
const tagMaps = new Map(tagDefinitionEntries)
/**
* Always read: the provenance snapshot vouches for the name, URL, and tags
* a model may see, but the source card's modified time and connector type
Expand Down
Loading
Loading