Skip to content

Commit 0bc7f98

Browse files
committed
fix(knowledge): scope list counts to selected knowledge bases
1 parent 8a82103 commit 0bc7f98

3 files changed

Lines changed: 214 additions & 30 deletions

File tree

Lines changed: 194 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,194 @@
1+
import { writeFileSync } from 'node:fs'
2+
import { db } from '@sim/db'
3+
import { document, knowledgeBase, organization, user, workspace } from '@sim/db/schema'
4+
import { generateId } from '@sim/utils/id'
5+
import { eq, inArray, sql } from 'drizzle-orm'
6+
import { afterAll, beforeAll, describe, expect, it } from 'vitest'
7+
import {
8+
createKnowledgeAclFixtureIds,
9+
seedKnowledgeAclFixture,
10+
} from '@/lib/knowledge/__integration__/seed-source-access-fixture'
11+
import { type KnowledgeAccessScope, WORKSPACE_ACCESS_TOKENS } from '@/lib/knowledge/access/types'
12+
import { getWorkspaceKnowledgeBases } from '@/lib/knowledge/service'
13+
14+
/**
15+
* A small KB page must not count the rest of its workspace or other tenants before applying
16+
* its limit. Real query plans catch this even when warm caches hide it from a timing test.
17+
* The same read must retain ACL/lifecycle filtering, empty bases, and keyset continuity.
18+
*/
19+
const ids = createKnowledgeAclFixtureIds()
20+
const foreign = createKnowledgeAclFixtureIds()
21+
const offPageId = generateId()
22+
const emptyId = generateId()
23+
const archivedId = generateId()
24+
const access: KnowledgeAccessScope = { kind: 'workspace', tokens: WORKSPACE_ACCESS_TOKENS }
25+
const reports: Array<Record<string, unknown>> = []
26+
27+
interface CapturedQuery {
28+
query: string
29+
parameters: NonNullable<Parameters<typeof db.$client.unsafe>[1]>
30+
}
31+
32+
interface ExplainNode {
33+
'Relation Name'?: string
34+
'Index Name'?: string
35+
'Actual Rows': number
36+
'Actual Loops': number
37+
'Rows Removed by Filter'?: number
38+
'Rows Removed by Index Recheck'?: number
39+
Plans?: ExplainNode[]
40+
}
41+
42+
function documentVisits(node: ExplainNode): number {
43+
const readsDocuments =
44+
node['Relation Name'] === 'document' || node['Index Name']?.startsWith('doc_')
45+
const own = readsDocuments
46+
? (node['Actual Rows'] +
47+
(node['Rows Removed by Filter'] ?? 0) +
48+
(node['Rows Removed by Index Recheck'] ?? 0)) *
49+
node['Actual Loops']
50+
: 0
51+
return own + (node.Plans ?? []).reduce((total, child) => total + documentVisits(child), 0)
52+
}
53+
54+
beforeAll(async () => {
55+
await seedKnowledgeAclFixture(ids, { connectorType: 'google_drive' })
56+
await seedKnowledgeAclFixture(foreign, { connectorType: 'google_drive' })
57+
await db
58+
.update(knowledgeBase)
59+
.set({ name: 'A small', createdAt: new Date('2026-01-01') })
60+
.where(eq(knowledgeBase.id, ids.knowledgeBaseId))
61+
await db.insert(knowledgeBase).values([
62+
{
63+
id: emptyId,
64+
workspaceId: ids.workspaceId,
65+
userId: ids.aliceId,
66+
name: 'B empty',
67+
createdAt: new Date('2026-01-02'),
68+
},
69+
{
70+
id: offPageId,
71+
workspaceId: ids.workspaceId,
72+
userId: ids.aliceId,
73+
name: 'C large',
74+
createdAt: new Date('2026-01-03'),
75+
},
76+
{
77+
id: archivedId,
78+
workspaceId: ids.workspaceId,
79+
userId: ids.aliceId,
80+
name: 'D archived',
81+
deletedAt: new Date(),
82+
},
83+
])
84+
await db.insert(document).values(
85+
[
86+
{ tokenCount: 7 },
87+
{ tokenCount: 11 },
88+
{ tokenCount: 100, acl: ['u:hidden@fixture.test'] },
89+
{ tokenCount: 100, archivedAt: new Date() },
90+
{ tokenCount: 100, deletedAt: new Date() },
91+
{ tokenCount: 100, userExcluded: true },
92+
].map((row) => ({
93+
id: generateId(),
94+
knowledgeBaseId: ids.knowledgeBaseId,
95+
filename: 'fixture.txt',
96+
fileUrl: 'https://fixture.invalid/document',
97+
fileSize: 1,
98+
mimeType: 'text/plain',
99+
acl: ['ws'],
100+
...row,
101+
}))
102+
)
103+
for (const baseId of [offPageId, foreign.knowledgeBaseId]) {
104+
await db.execute(sql`INSERT INTO document
105+
(id, knowledge_base_id, filename, file_url, file_size, mime_type, acl, token_count)
106+
SELECT ${baseId} || '-' || n, ${baseId}, 'bulk.txt', 'https://fixture.invalid/bulk',
107+
1, 'text/plain', ARRAY['ws'], 1 FROM generate_series(1, 10000) AS n`)
108+
}
109+
await db.execute(sql`ANALYZE knowledge_base`)
110+
await db.execute(sql`ANALYZE document`)
111+
}, 60_000)
112+
113+
afterAll(async () => {
114+
const reportPath = process.env.KNOWLEDGE_BASE_LIST_REPORT_PATH
115+
if (reportPath) writeFileSync(reportPath, JSON.stringify(reports, null, 2))
116+
try {
117+
for (const fixture of [ids, foreign]) {
118+
await db.delete(workspace).where(eq(workspace.id, fixture.workspaceId))
119+
await db.delete(organization).where(eq(organization.id, fixture.organizationId))
120+
await db.delete(user).where(inArray(user.id, [fixture.aliceId, fixture.bobId]))
121+
}
122+
} finally {
123+
await db.$client.end()
124+
}
125+
})
126+
127+
describe('knowledge base list counts on real Postgres', () => {
128+
it.each(['name', 'createdAt'] as const)(
129+
'bounds document reads to a small page ordered by %s',
130+
async (sortBy) => {
131+
const captured: CapturedQuery[] = []
132+
const previousDebug = db.$client.options.debug
133+
db.$client.options.debug = (_connection, query, parameters) => {
134+
if (captured.length < 30) captured.push({ query, parameters: [...parameters] })
135+
}
136+
try {
137+
const page = await getWorkspaceKnowledgeBases(ids.workspaceId, 'active', {
138+
countsFor: access,
139+
limit: 1,
140+
sortBy,
141+
})
142+
expect(
143+
page.data.map(({ id, docCount, tokenCount }) => ({ id, docCount, tokenCount }))
144+
).toEqual([{ id: ids.knowledgeBaseId, docCount: 2, tokenCount: 18 }])
145+
expect(page.nextCursorKeys).not.toBeNull()
146+
} finally {
147+
db.$client.options.debug = previousDebug
148+
}
149+
const plans = []
150+
for (const statement of captured.filter(({ query }) => query.includes('"document"'))) {
151+
const [result] = await db.$client.unsafe<
152+
Array<{ 'QUERY PLAN': Array<{ Plan: ExplainNode }> }>
153+
>(`EXPLAIN (ANALYZE, BUFFERS, FORMAT JSON) ${statement.query}`, statement.parameters)
154+
plans.push(...result['QUERY PLAN'])
155+
}
156+
const visits = plans.reduce((total, plan) => total + documentVisits(plan.Plan), 0)
157+
reports.push({ sortBy, visits, plans })
158+
expect(plans.length).toBeGreaterThan(0)
159+
expect(visits).toBeLessThan(100)
160+
}
161+
)
162+
163+
it('keeps empty KBs and count visibility through pagination and unpaged reads', async () => {
164+
const first = await getWorkspaceKnowledgeBases(ids.workspaceId, 'active', {
165+
countsFor: access,
166+
limit: 1,
167+
sortBy: 'name',
168+
})
169+
if (!first.nextCursorKeys) throw new Error('Expected a second knowledge-base page')
170+
const second = await getWorkspaceKnowledgeBases(ids.workspaceId, 'active', {
171+
countsFor: access,
172+
limit: 1,
173+
sortBy: 'name',
174+
cursorKeys: first.nextCursorKeys,
175+
})
176+
expect(
177+
second.data.map(({ id, docCount, tokenCount }) => ({ id, docCount, tokenCount }))
178+
).toEqual([{ id: emptyId, docCount: 0, tokenCount: 0 }])
179+
const all = await getWorkspaceKnowledgeBases(ids.workspaceId, 'active', {
180+
countsFor: access,
181+
sortBy: 'name',
182+
})
183+
expect(all.data.map(({ id, docCount, tokenCount }) => ({ id, docCount, tokenCount }))).toEqual([
184+
{ id: ids.knowledgeBaseId, docCount: 2, tokenCount: 18 },
185+
{ id: emptyId, docCount: 0, tokenCount: 0 },
186+
{ id: offPageId, docCount: 10000, tokenCount: 10000 },
187+
])
188+
expect(all.nextCursorKeys).toBeNull()
189+
const archived = await getWorkspaceKnowledgeBases(ids.workspaceId, 'archived', {
190+
countsFor: access,
191+
})
192+
expect(archived.data.map(({ id }) => id)).toEqual([archivedId])
193+
})
194+
})

‎apps/sim/lib/knowledge/service.test.ts‎

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -293,18 +293,16 @@ describe('knowledge base counts with live source permissions', () => {
293293
id: 'kb-1',
294294
workspaceId: 'ws-1',
295295
chunkingConfig: {},
296-
docCount: 2,
297-
tokenCount: 10,
298296
createdAt: new Date('2026-01-01'),
299297
},
300298
])
299+
queueTableRows(schemaMock.document, [{ knowledgeBaseId: 'kb-1', docCount: 2, tokenCount: 10 }])
301300
const result = await getWorkspaceKnowledgeBases('ws-1', 'archived', { countsFor: access })
302301
expect(result.data[0]).toMatchObject({ docCount: 2, tokenCount: 10 })
303302
expect(getForConnectors).not.toHaveBeenCalled()
304303
expect(dbChainMockFns.select).not.toHaveBeenCalledWith({
305304
connectorId: schemaMock.knowledgeConnector.id,
306305
})
307-
expect(dbChainMockFns.groupBy).toHaveBeenCalledOnce()
308306
})
309307

310308
it('does not retain stale totals when a live source no longer authorizes its documents', async () => {

‎apps/sim/lib/knowledge/service.ts‎

Lines changed: 19 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ import { db } from '@sim/db'
22
import { document, knowledgeBase, knowledgeConnector, workspaceFiles } from '@sim/db/schema'
33
import { createLogger } from '@sim/logger'
44
import { getPostgresConstraintName, getPostgresErrorCode } from '@sim/utils/errors'
5+
import { chunkArray } from '@sim/utils/helpers'
56
import { generateId } from '@sim/utils/id'
67
import { filterUndefined } from '@sim/utils/object'
78
import type { SQL } from 'drizzle-orm'
@@ -52,6 +53,7 @@ import type {
5253
import { getUserEntityPermissions } from '@/lib/workspaces/permissions/utils'
5354

5455
const logger = createLogger('KnowledgeBaseService')
56+
const KNOWLEDGE_BASE_COUNT_BATCH_SIZE = 100
5557

5658
/**
5759
* Every caller-fixable knowledge-base failure is an {@link OrchestrationError},
@@ -188,8 +190,8 @@ async function readKnowledgeBaseRows(
188190
}
189191

190192
/**
191-
* {@link readKnowledgeBaseRows} plus the live totals of the documents `access` admits. Only the
192-
* surfaces that display totals pay for the document join, and they always count as a reader.
193+
* Pages bases before counting the documents `access` admits. Explicit document base IDs keep
194+
* the count selective instead of scanning a shared ACL token across tenants before the join.
193195
*/
194196
async function readCountedKnowledgeBaseRows(
195197
where: SQL | undefined,
@@ -200,31 +202,21 @@ async function readCountedKnowledgeBaseRows(
200202
Array<ActiveKnowledgeBaseReference & Pick<KnowledgeBaseWithCounts, 'docCount' | 'tokenCount'>>
201203
> {
202204
const scope = 'get' in access ? await access.get() : access
203-
const query = db
204-
.select({
205-
...ACTIVE_KNOWLEDGE_BASE_REFERENCE_FIELDS,
206-
tokenCount: sql<number>`COALESCE(SUM(${document.tokenCount}), 0)`.mapWith(Number),
207-
docCount: count(document.knowledgeBaseId),
208-
})
209-
.from(knowledgeBase)
210-
.leftJoin(
211-
document,
212-
and(
213-
eq(document.knowledgeBaseId, knowledgeBase.id),
214-
eq(document.userExcluded, false),
215-
isNull(document.archivedAt),
216-
isNull(document.deletedAt),
217-
knowledgeAccessCondition(scope)
218-
)
205+
const rows = await readKnowledgeBaseRows(where, orderBy, limit)
206+
const counts = new Map<string, { docCount: number; tokenCount: number }>()
207+
for (const batch of chunkArray(rows, KNOWLEDGE_BASE_COUNT_BATCH_SIZE)) {
208+
const totals = await countDocumentsByKnowledgeBase(
209+
inArray(
210+
document.knowledgeBaseId,
211+
batch.map((kb) => kb.id)
212+
),
213+
knowledgeAccessCondition(scope)
219214
)
220-
.where(where)
221-
.groupBy(knowledgeBase.id)
222-
.orderBy(...orderBy)
223-
224-
const rows = limit === undefined ? await query : await query.limit(limit)
215+
for (const total of totals) counts.set(total.knowledgeBaseId, total)
216+
}
225217

226218
/**
227-
* The join above already counted everything the reader's stored ACL admits. Only a
219+
* The counts above already include everything the reader's stored ACL admits. Only a
228220
* provider can add documents a live source (GitHub, Confluence) authorizes beyond that,
229221
* and that supplement is resolved once for the whole list: an unpaged list is bounded by
230222
* its own filter, a page by its row IDs, so a workspace with tens of thousands of bases
@@ -243,9 +235,9 @@ async function readCountedKnowledgeBaseRows(
243235
)
244236
: undefined
245237
return rows.map((kb) => ({
246-
...toActiveKnowledgeBaseReference(kb),
247-
docCount: Number(kb.docCount) + (liveCounts?.get(kb.id)?.docCount ?? 0),
248-
tokenCount: kb.tokenCount + (liveCounts?.get(kb.id)?.tokenCount ?? 0),
238+
...kb,
239+
docCount: (counts.get(kb.id)?.docCount ?? 0) + (liveCounts?.get(kb.id)?.docCount ?? 0),
240+
tokenCount: (counts.get(kb.id)?.tokenCount ?? 0) + (liveCounts?.get(kb.id)?.tokenCount ?? 0),
249241
}))
250242
}
251243

0 commit comments

Comments
 (0)