Skip to content
Closed
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
133 changes: 133 additions & 0 deletions apps/sim/lib/knowledge/__integration__/list-totals-plan.integration.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,133 @@
/**
* The counted knowledge-base list must total each base through that base's own documents, never
* through a set-based join over the whole document table. The reader's ACL tokens are shared by
* every tenant, so a join the planner drives from the ACL index instead of `knowledge_base_id`
* reads every tenant's documents to total one base. Nested loops are disabled while planning so
* the planner reaches for exactly that join wherever the query shape still allows it.
*/
import type { Principal } from '@sim/auth/principal'
import * as schema from '@sim/db/schema'
import { document, organization, user, workspace } from '@sim/db/schema'
import { readTestDatabaseUrl } from '@sim/db/testing/test-infrastructure'
import { withUtcTimestamps } from '@sim/db/timestamps'
import { generateId } from '@sim/utils/id'
import { eq, inArray } from 'drizzle-orm'
Comment thread
waleedlatif1 marked this conversation as resolved.
import { drizzle, type PostgresJsDatabase } from 'drizzle-orm/postgres-js'
import postgres from 'postgres'
import { afterAll, beforeAll, describe, expect, it, vi } from 'vitest'

const database = vi.hoisted(() => ({
current: undefined as PostgresJsDatabase<typeof import('@sim/db/schema')> | undefined,
/** Every statement the app issues, so the test can EXPLAIN the exact SQL it ran. */
statements: [] as { sql: string; params: unknown[] }[],
}))
vi.mock('@sim/db', async (importOriginal) => ({
...(await importOriginal<typeof import('@sim/db')>()),
get db() {
if (!database.current) throw new Error('Test database not initialized')
return database.current
},
}))

import {
createKnowledgeAclFixtureIds,
seedKnowledgeAclFixture,
} from '@/lib/knowledge/__integration__/seed-source-access-fixture'
import { listKnowledgeBases } from '@/lib/knowledge/application/knowledge-bases'

interface PlanNode {
'Node Type': string
'Relation Name'?: string
'Index Cond'?: string
'Recheck Cond'?: string
Filter?: string
Plans?: PlanNode[]
}

function documentScans(node: PlanNode): PlanNode[] {
const own = node['Relation Name'] === 'document' ? [node] : []
return [...own, ...(node.Plans ?? []).flatMap(documentScans)]
}

const ids = createKnowledgeAclFixtureIds()
const reader: Principal = { kind: 'session', userId: ids.bobId, sessionId: 'fixture-reader' }
const connection = postgres(
readTestDatabaseUrl(),
withUtcTimestamps({
max: 2,
prepare: false,
fetch_types: false,
connection: {},
onnotice: () => {},
})
)

describe('knowledge-base list totals plan', () => {
beforeAll(async () => {
vi.stubGlobal('fetch', async () => {
throw new Error('Unexpected provider request in knowledge-base totals plan test')
})
database.current = drizzle(connection, {
schema,
logger: {
logQuery(sql, params) {
database.statements.push({ sql, params })
},
},
})
await seedKnowledgeAclFixture(ids, { connectorType: 'google_drive' })
await database.current.insert(document).values({
id: generateId(),
knowledgeBaseId: ids.knowledgeBaseId,
filename: 'Workspace handbook',
fileUrl: 'https://fixture.test/shared',
fileSize: 10,
mimeType: 'text/plain',
tokenCount: 10,
processingStatus: 'completed',
})
})

afterAll(async () => {
await database.current?.delete(workspace).where(eq(workspace.id, ids.workspaceId))
Comment thread
waleedlatif1 marked this conversation as resolved.
Comment thread
waleedlatif1 marked this conversation as resolved.
await database.current?.delete(organization).where(eq(organization.id, ids.organizationId))
await database.current?.delete(user).where(inArray(user.id, [ids.aliceId, ids.bobId]))
await connection.end()
vi.unstubAllGlobals()
})

it('totals each base through its own documents even when the planner prefers a hash join', async () => {
database.statements.length = 0
const { knowledgeBases } = await listKnowledgeBases.execute({
principal: reader,
input: { workspaceId: ids.workspaceId },
})
expect(knowledgeBases.map(({ knowledgeBase }) => knowledgeBase)).toEqual([
expect.objectContaining({ id: ids.knowledgeBaseId, docCount: 1, tokenCount: 10 }),
])

const counted = database.statements.filter(
({ sql }) =>
/^select .* from "knowledge_base" /.test(sql) &&
sql.includes('"document"."knowledge_base_id" = "knowledge_base"."id"')
)
expect(counted).toHaveLength(1)
const plan = await connection.begin(async (tx) => {
await tx`SET LOCAL enable_nestloop = off`
const [row] = await tx.unsafe(
`EXPLAIN (FORMAT JSON) ${counted[0].sql}`,
counted[0].params as Parameters<typeof tx.unsafe>[1]
)
return (row['QUERY PLAN'] as { Plan: PlanNode }[])[0].Plan
})

const scans = documentScans(plan)
expect(scans.length).toBeGreaterThan(0)
for (const scan of scans) {
const conditions = [scan['Index Cond'], scan['Recheck Cond'], scan.Filter].join(' ')
expect(conditions, `${scan['Node Type']} on document`).toContain(
'(knowledge_base_id = knowledge_base.id)'
)
}
})
})
1 change: 0 additions & 1 deletion apps/sim/lib/knowledge/service.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -304,7 +304,6 @@ describe('knowledge base counts with live source permissions', () => {
expect(dbChainMockFns.select).not.toHaveBeenCalledWith({
connectorId: schemaMock.knowledgeConnector.id,
})
expect(dbChainMockFns.groupBy).toHaveBeenCalledOnce()
})

it('does not retain stale totals when a live source no longer authorizes its documents', async () => {
Expand Down
33 changes: 22 additions & 11 deletions apps/sim/lib/knowledge/service.ts
Original file line number Diff line number Diff line change
Expand Up @@ -200,25 +200,36 @@ async function readCountedKnowledgeBaseRows(
Array<ActiveKnowledgeBaseReference & Pick<KnowledgeBaseWithCounts, 'docCount' | 'tokenCount'>>
> {
const scope = 'get' in access ? await access.get() : access
const query = db
/**
* A lateral aggregate, so each base is counted through its own `knowledge_base_id` index and
* never through a bitmap of every tenant's documents sharing the `ws` token. An aggregate always
* yields one row, so an empty base stays at zero.
*/
const totals = db
.select({
...ACTIVE_KNOWLEDGE_BASE_REFERENCE_FIELDS,
tokenCount: sql<number>`COALESCE(SUM(${document.tokenCount}), 0)`.mapWith(Number),
docCount: count(document.knowledgeBaseId),
tokenCount: sql<number>`COALESCE(SUM(${document.tokenCount}), 0)`
.mapWith(Number)
.as('readable_token_count'),
docCount: count().as('readable_doc_count'),
})
.from(knowledgeBase)
.leftJoin(
document,
.from(document)
.where(
and(
eq(document.knowledgeBaseId, knowledgeBase.id),
eq(document.userExcluded, false),
isNull(document.archivedAt),
isNull(document.deletedAt),
...ACTIVE_DOCUMENT_CONDITIONS,
knowledgeAccessCondition(scope)
)
)
.as('totals')
const query = db
.select({
...ACTIVE_KNOWLEDGE_BASE_REFERENCE_FIELDS,
tokenCount: totals.tokenCount,
docCount: totals.docCount,
})
.from(knowledgeBase)
.innerJoinLateral(totals, sql`true`)
Comment thread
waleedlatif1 marked this conversation as resolved.
.where(where)
.groupBy(knowledgeBase.id)
.orderBy(...orderBy)

const rows = limit === undefined ? await query : await query.limit(limit)
Expand Down
21 changes: 20 additions & 1 deletion packages/testing/src/mocks/database.mock.ts
Original file line number Diff line number Diff line change
Expand Up @@ -91,7 +91,16 @@ function createMockSqlOperators() {
gte: vi.fn((a, b) => ({ type: 'gte', left: a, right: b })),
lt: vi.fn((a, b) => ({ type: 'lt', left: a, right: b })),
lte: vi.fn((a, b) => ({ type: 'lte', left: a, right: b })),
count: vi.fn((column) => ({ type: 'count', column })),
/** Drizzle's `count()` is an `SQL` expression, so it aliases and decodes like one. */
count: vi.fn((column) => {
const expression = {
type: 'count',
column,
as: (alias: string) => ({ ...expression, alias }),
mapWith: (decoder: unknown) => ({ ...expression, decoder }),
}
return expression
}),
avg: vi.fn((column) => ({ type: 'avg', column })),
sum: vi.fn((column) => ({ type: 'sum', column })),
min: vi.fn((column) => ({ type: 'min', column })),
Expand Down Expand Up @@ -235,6 +244,8 @@ const asAlias = chainSpy()
const forClause = chainSpy()
const innerJoin = chainSpy()
const leftJoin = chainSpy()
const innerJoinLateral = chainSpy()
const leftJoinLateral = chainSpy()
const insert = chainSpy()
const update = chainSpy()
const set = chainSpy()
Expand Down Expand Up @@ -334,6 +345,12 @@ const joinBuilder = (tables: unknown[], fields: SelectedFields = {}): any => {
builder.leftJoin = spyOrDefault(leftJoin, (table: unknown) =>
joinBuilder([...tables, table], fields)
)
builder.innerJoinLateral = spyOrDefault(innerJoinLateral, (subquery: unknown) =>
joinBuilder([...tables, subquery], fields)
)
builder.leftJoinLateral = spyOrDefault(leftJoinLateral, (subquery: unknown) =>
joinBuilder([...tables, subquery], fields)
)
return builder
}

Expand All @@ -359,6 +376,8 @@ export const dbChainMockFns = {
returning,
innerJoin,
leftJoin,
innerJoinLateral,
leftJoinLateral,
groupBy,
having,
as: asAlias,
Expand Down
Loading