Skip to content

Commit b36757f

Browse files
fix(knowledge): scope list counts to selected knowledge bases (#8575)
* fix(knowledge): scope list counts to selected knowledge bases * fix(knowledge): avoid count-query fan-out for unpaged lists
1 parent 24b4957 commit b36757f

3 files changed

Lines changed: 252 additions & 32 deletions

File tree

Lines changed: 235 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,235 @@
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 manyBases = createKnowledgeAclFixtureIds()
22+
const offPageId = generateId()
23+
const emptyId = generateId()
24+
const archivedId = generateId()
25+
const access: KnowledgeAccessScope = { kind: 'workspace', tokens: WORKSPACE_ACCESS_TOKENS }
26+
const reports: Array<Record<string, unknown>> = []
27+
28+
interface CapturedQuery {
29+
query: string
30+
parameters: NonNullable<Parameters<typeof db.$client.unsafe>[1]>
31+
}
32+
33+
interface ExplainNode {
34+
'Relation Name'?: string
35+
'Index Name'?: string
36+
'Actual Rows': number
37+
'Actual Loops': number
38+
'Rows Removed by Filter'?: number
39+
'Rows Removed by Index Recheck'?: number
40+
Plans?: ExplainNode[]
41+
}
42+
43+
function documentVisits(node: ExplainNode): number {
44+
const readsDocuments =
45+
node['Relation Name'] === 'document' || node['Index Name']?.startsWith('doc_')
46+
const own = readsDocuments
47+
? (node['Actual Rows'] +
48+
(node['Rows Removed by Filter'] ?? 0) +
49+
(node['Rows Removed by Index Recheck'] ?? 0)) *
50+
node['Actual Loops']
51+
: 0
52+
return own + (node.Plans ?? []).reduce((total, child) => total + documentVisits(child), 0)
53+
}
54+
55+
beforeAll(async () => {
56+
await seedKnowledgeAclFixture(ids, { connectorType: 'google_drive' })
57+
await seedKnowledgeAclFixture(foreign, { connectorType: 'google_drive' })
58+
await db
59+
.update(knowledgeBase)
60+
.set({ name: 'A small', createdAt: new Date('2026-01-01') })
61+
.where(eq(knowledgeBase.id, ids.knowledgeBaseId))
62+
await db.insert(knowledgeBase).values([
63+
{
64+
id: emptyId,
65+
workspaceId: ids.workspaceId,
66+
userId: ids.aliceId,
67+
name: 'B empty',
68+
createdAt: new Date('2026-01-02'),
69+
},
70+
{
71+
id: offPageId,
72+
workspaceId: ids.workspaceId,
73+
userId: ids.aliceId,
74+
name: 'C large',
75+
createdAt: new Date('2026-01-03'),
76+
},
77+
{
78+
id: archivedId,
79+
workspaceId: ids.workspaceId,
80+
userId: ids.aliceId,
81+
name: 'D archived',
82+
deletedAt: new Date(),
83+
},
84+
])
85+
await db.insert(document).values(
86+
[
87+
{ tokenCount: 7 },
88+
{ tokenCount: 11 },
89+
{ tokenCount: 100, acl: ['u:hidden@fixture.test'] },
90+
{ tokenCount: 100, archivedAt: new Date() },
91+
{ tokenCount: 100, deletedAt: new Date() },
92+
{ tokenCount: 100, userExcluded: true },
93+
].map((row) => ({
94+
id: generateId(),
95+
knowledgeBaseId: ids.knowledgeBaseId,
96+
filename: 'fixture.txt',
97+
fileUrl: 'https://fixture.invalid/document',
98+
fileSize: 1,
99+
mimeType: 'text/plain',
100+
acl: ['ws'],
101+
...row,
102+
}))
103+
)
104+
for (const baseId of [offPageId, foreign.knowledgeBaseId]) {
105+
await db.execute(sql`INSERT INTO document
106+
(id, knowledge_base_id, filename, file_url, file_size, mime_type, acl, token_count)
107+
SELECT ${baseId} || '-' || n, ${baseId}, 'bulk.txt', 'https://fixture.invalid/bulk',
108+
1, 'text/plain', ARRAY['ws'], 1 FROM generate_series(1, 10000) AS n`)
109+
}
110+
await seedKnowledgeAclFixture(manyBases, { connectorType: 'google_drive' })
111+
await db.execute(sql`INSERT INTO knowledge_base (id, workspace_id, user_id, name)
112+
SELECT ${manyBases.knowledgeBaseId} || '-' || n, ${manyBases.workspaceId},
113+
${manyBases.aliceId}, 'Scale fixture ' || n FROM generate_series(1, 10000) AS n`)
114+
await db.execute(sql`INSERT INTO document
115+
(id, knowledge_base_id, filename, file_url, file_size, mime_type, acl, token_count)
116+
VALUES (${generateId()}, ${manyBases.knowledgeBaseId}, 'first.txt',
117+
'https://fixture.invalid/first', 1, 'text/plain', ARRAY['ws'], 13),
118+
(${generateId()}, ${`${manyBases.knowledgeBaseId}-10000`}, 'last.txt',
119+
'https://fixture.invalid/last', 1, 'text/plain', ARRAY['ws'], 17)`)
120+
await db.execute(sql`ANALYZE knowledge_base`)
121+
await db.execute(sql`ANALYZE document`)
122+
}, 60_000)
123+
124+
afterAll(async () => {
125+
const reportPath = process.env.KNOWLEDGE_BASE_LIST_REPORT_PATH
126+
if (reportPath) writeFileSync(reportPath, JSON.stringify(reports, null, 2))
127+
try {
128+
for (const fixture of [ids, foreign, manyBases]) {
129+
await db.delete(workspace).where(eq(workspace.id, fixture.workspaceId))
130+
await db.delete(organization).where(eq(organization.id, fixture.organizationId))
131+
await db.delete(user).where(inArray(user.id, [fixture.aliceId, fixture.bobId]))
132+
}
133+
} finally {
134+
await db.$client.end()
135+
}
136+
})
137+
138+
describe('knowledge base list counts on real Postgres', () => {
139+
it.each(['name', 'createdAt'] as const)(
140+
'bounds document reads to a small page ordered by %s',
141+
async (sortBy) => {
142+
const captured: CapturedQuery[] = []
143+
const previousDebug = db.$client.options.debug
144+
db.$client.options.debug = (_connection, query, parameters) => {
145+
if (captured.length < 30) captured.push({ query, parameters: [...parameters] })
146+
}
147+
try {
148+
const page = await getWorkspaceKnowledgeBases(ids.workspaceId, 'active', {
149+
countsFor: access,
150+
limit: 1,
151+
sortBy,
152+
})
153+
expect(
154+
page.data.map(({ id, docCount, tokenCount }) => ({ id, docCount, tokenCount }))
155+
).toEqual([{ id: ids.knowledgeBaseId, docCount: 2, tokenCount: 18 }])
156+
expect(page.nextCursorKeys).not.toBeNull()
157+
} finally {
158+
db.$client.options.debug = previousDebug
159+
}
160+
const plans = []
161+
for (const statement of captured.filter(({ query }) => query.includes('"document"'))) {
162+
const [result] = await db.$client.unsafe<
163+
Array<{ 'QUERY PLAN': Array<{ Plan: ExplainNode }> }>
164+
>(`EXPLAIN (ANALYZE, BUFFERS, FORMAT JSON) ${statement.query}`, statement.parameters)
165+
plans.push(...result['QUERY PLAN'])
166+
}
167+
const visits = plans.reduce((total, plan) => total + documentVisits(plan.Plan), 0)
168+
reports.push({ sortBy, visits, plans })
169+
expect(plans.length).toBeGreaterThan(0)
170+
expect(visits).toBeLessThan(100)
171+
}
172+
)
173+
174+
it('keeps empty KBs and count visibility through pagination and unpaged reads', async () => {
175+
const first = await getWorkspaceKnowledgeBases(ids.workspaceId, 'active', {
176+
countsFor: access,
177+
limit: 1,
178+
sortBy: 'name',
179+
})
180+
if (!first.nextCursorKeys) throw new Error('Expected a second knowledge-base page')
181+
const second = await getWorkspaceKnowledgeBases(ids.workspaceId, 'active', {
182+
countsFor: access,
183+
limit: 1,
184+
sortBy: 'name',
185+
cursorKeys: first.nextCursorKeys,
186+
})
187+
expect(
188+
second.data.map(({ id, docCount, tokenCount }) => ({ id, docCount, tokenCount }))
189+
).toEqual([{ id: emptyId, docCount: 0, tokenCount: 0 }])
190+
const all = await getWorkspaceKnowledgeBases(ids.workspaceId, 'active', {
191+
countsFor: access,
192+
sortBy: 'name',
193+
})
194+
expect(all.data.map(({ id, docCount, tokenCount }) => ({ id, docCount, tokenCount }))).toEqual([
195+
{ id: ids.knowledgeBaseId, docCount: 2, tokenCount: 18 },
196+
{ id: emptyId, docCount: 0, tokenCount: 0 },
197+
{ id: offPageId, docCount: 10000, tokenCount: 10000 },
198+
])
199+
expect(all.nextCursorKeys).toBeNull()
200+
const archived = await getWorkspaceKnowledgeBases(ids.workspaceId, 'archived', {
201+
countsFor: access,
202+
})
203+
expect(archived.data.map(({ id }) => id)).toEqual([archivedId])
204+
})
205+
206+
it('counts a large unpaged workspace within a fixed database round-trip budget', async () => {
207+
let documentQueries = 0
208+
const previousDebug = db.$client.options.debug
209+
db.$client.options.debug = (_connection, query) => {
210+
if (query.includes('"document"')) documentQueries++
211+
}
212+
try {
213+
const all = await getWorkspaceKnowledgeBases(manyBases.workspaceId, 'active', {
214+
countsFor: access,
215+
})
216+
expect(all.data).toHaveLength(10001)
217+
expect(all.nextCursorKeys).toBeNull()
218+
expect(all.data.find((kb) => kb.id === manyBases.knowledgeBaseId)).toMatchObject({
219+
docCount: 1,
220+
tokenCount: 13,
221+
})
222+
expect(all.data.find((kb) => kb.id === `${manyBases.knowledgeBaseId}-10000`)).toMatchObject({
223+
docCount: 1,
224+
tokenCount: 17,
225+
})
226+
expect(all.data.reduce((total, kb) => total + kb.docCount, 0)).toBe(2)
227+
expect(all.data.reduce((total, kb) => total + kb.tokenCount, 0)).toBe(30)
228+
reports.push({ unpagedBases: all.data.length, documentQueries })
229+
expect(documentQueries).toBeGreaterThan(0)
230+
expect(documentQueries).toBeLessThan(10)
231+
} finally {
232+
db.$client.options.debug = previousDebug
233+
}
234+
})
235+
})

‎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: 16 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ import { resourceScopeCondition } from '@/lib/core/resource-scope.server'
3333
import { generateRestoreName } from '@/lib/core/utils/restore-name'
3434
import { findActiveFolder, resolveRestoredFolderId } from '@/lib/folders/queries'
3535
import { isKnowledgeMemberAccessAvailable } from '@/lib/knowledge/access/availability'
36-
import { knowledgeAccessCondition } from '@/lib/knowledge/access/predicate'
36+
import { knowledgeAccessCondition, textArrayLiteral } from '@/lib/knowledge/access/predicate'
3737
import type { KnowledgeAccessProvider } from '@/lib/knowledge/access/types'
3838
import { mirrorsSourceAcls } from '@/lib/knowledge/connectors/access-modes'
3939
import {
@@ -188,8 +188,8 @@ async function readKnowledgeBaseRows(
188188
}
189189

190190
/**
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.
191+
* Pages bases before counting the documents `access` admits. Explicit document base IDs keep
192+
* the count selective instead of scanning a shared ACL token across tenants before the join.
193193
*/
194194
async function readCountedKnowledgeBaseRows(
195195
where: SQL | undefined,
@@ -200,31 +200,18 @@ async function readCountedKnowledgeBaseRows(
200200
Array<ActiveKnowledgeBaseReference & Pick<KnowledgeBaseWithCounts, 'docCount' | 'tokenCount'>>
201201
> {
202202
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-
)
219-
)
220-
.where(where)
221-
.groupBy(knowledgeBase.id)
222-
.orderBy(...orderBy)
223-
224-
const rows = limit === undefined ? await query : await query.limit(limit)
203+
const rows = await readKnowledgeBaseRows(where, orderBy, limit)
204+
const totals =
205+
rows.length > 0
206+
? await countDocumentsByKnowledgeBase(
207+
sql`${document.knowledgeBaseId} = ANY(${textArrayLiteral(rows.map((kb) => kb.id))})`,
208+
knowledgeAccessCondition(scope)
209+
)
210+
: []
211+
const counts = new Map(totals.map((total) => [total.knowledgeBaseId, total]))
225212

226213
/**
227-
* The join above already counted everything the reader's stored ACL admits. Only a
214+
* The counts above already include everything the reader's stored ACL admits. Only a
228215
* provider can add documents a live source (GitHub, Confluence) authorizes beyond that,
229216
* and that supplement is resolved once for the whole list: an unpaged list is bounded by
230217
* its own filter, a page by its row IDs, so a workspace with tens of thousands of bases
@@ -243,9 +230,9 @@ async function readCountedKnowledgeBaseRows(
243230
)
244231
: undefined
245232
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),
233+
...kb,
234+
docCount: (counts.get(kb.id)?.docCount ?? 0) + (liveCounts?.get(kb.id)?.docCount ?? 0),
235+
tokenCount: (counts.get(kb.id)?.tokenCount ?? 0) + (liveCounts?.get(kb.id)?.tokenCount ?? 0),
249236
}))
250237
}
251238

0 commit comments

Comments
 (0)