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
28 changes: 18 additions & 10 deletions apps/sim/app/api/knowledge/search/utils.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ import { ResolvedSecretTraceRegistry } from '@/executor/utils/resolved-secret-tr
vi.mock('@/lib/core/rate-limiter/provider-admission', () => ({
PROVIDER_QUOTA_COOLDOWN_MS: 300_000,
ProviderQuotaExhaustedError: class ProviderQuotaExhaustedError extends Error {},
ProviderAdmissionTimeoutError: class ProviderAdmissionTimeoutError extends Error {},
isProviderQuotaExhausted: vi.fn().mockResolvedValue(false),
recordProviderCooldown: vi.fn().mockResolvedValue(undefined),
waitForProviderAdmission: vi.fn().mockResolvedValue(undefined),
Expand Down Expand Up @@ -105,18 +106,23 @@ function makeResult(id: string, distance = 0.1): SearchResult {
}
}

const TEST_EMBEDDING = [0.1, 0.2, 0.3, ...Array.from({ length: 1533 }, () => 0)]
const TEST_EMBEDDING = [0.1, 0.2, 0.3, ...Array.from({ length: 1533 }, () => 0)].map(Math.fround)

function mockNextEmbeddingResponse(): void {
vi.mocked(fetch).mockResolvedValueOnce(
new Response(
vi.mocked(fetch).mockImplementationOnce(async (_url, init) => {
const request = JSON.parse(String(init?.body))
const embedding =
request.encoding_format === 'base64'
? Buffer.from(new Float32Array(TEST_EMBEDDING).buffer).toString('base64')
: TEST_EMBEDDING
return new Response(
JSON.stringify({
data: [{ embedding: TEST_EMBEDDING, index: 0 }],
data: [{ embedding, index: 0 }],
usage: { prompt_tokens: 1, total_tokens: 1 },
}),
{ status: 200, headers: { 'Content-Type': 'application/json' } }
)
)
})
}

describe('Knowledge Search Utils', () => {
Expand Down Expand Up @@ -220,7 +226,8 @@ describe('Knowledge Search Utils', () => {
})

expect(results.map((row) => row.id)).toEqual(['first', 'second'])
expect(dbChainMockFns.select).toHaveBeenCalledTimes(1)
expect(dbChainMockFns.select).toHaveBeenCalledTimes(2)
expect(dbChainMockFns.as).toHaveBeenCalledWith('ranked_embeddings')
expect(dbChainMockFns.select.mock.calls[0][0]).toHaveProperty('distance')
expect(dbChainMockFns.limit).toHaveBeenCalledWith(2)
})
Expand Down Expand Up @@ -541,7 +548,8 @@ describe('Knowledge Search Utils', () => {
})

expect(results.map((r) => r.id)).toEqual(['vector-hit'])
expect(dbChainMockFns.select).toHaveBeenCalledTimes(1)
expect(dbChainMockFns.select).toHaveBeenCalledTimes(2)
expect(dbChainMockFns.as).toHaveBeenCalledWith('ranked_embeddings')
})

it('runs both legs and fuses them in hybrid mode', async () => {
Expand All @@ -565,7 +573,7 @@ describe('Knowledge Search Utils', () => {
})

expect(results.map((r) => r.id).sort()).toEqual(['keyword-hit', 'vector-hit'])
expect(dbChainMockFns.select).toHaveBeenCalledTimes(3)
expect(dbChainMockFns.select).toHaveBeenCalledTimes(4)
})

it('falls back to vector results when the keyword leg fails', async () => {
Expand Down Expand Up @@ -830,7 +838,7 @@ describe('Knowledge Search Utils', () => {
body: JSON.stringify({
input: ['test query'],
model: 'text-embedding-3-small',
encoding_format: 'float',
encoding_format: 'base64',
dimensions: 1536,
}),
})
Expand Down Expand Up @@ -860,7 +868,7 @@ describe('Knowledge Search Utils', () => {
body: JSON.stringify({
input: ['prefix {{TOKEN}} suffix'],
model: 'text-embedding-3-small',
encoding_format: 'float',
encoding_format: 'base64',
dimensions: 1536,
}),
})
Expand Down
17 changes: 12 additions & 5 deletions apps/sim/app/api/knowledge/utils.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ import * as workspacesUtilsModule from '@/lib/workspaces/utils'
vi.mock('@/lib/core/rate-limiter/provider-admission', () => ({
PROVIDER_QUOTA_COOLDOWN_MS: 300_000,
ProviderQuotaExhaustedError: class ProviderQuotaExhaustedError extends Error {},
ProviderAdmissionTimeoutError: class ProviderAdmissionTimeoutError extends Error {},
isProviderQuotaExhausted: vi.fn().mockResolvedValue(false),
recordProviderCooldown: vi.fn().mockResolvedValue(undefined),
waitForProviderAdmission: vi.fn().mockResolvedValue(undefined),
Expand Down Expand Up @@ -125,18 +126,24 @@ function createTestEmbedding(value: number): number[] {
return Array.from({ length: TEST_EMBEDDING_DIMENSION }, () => value)
}

function createEmbeddingResponse(values: number[]): Response {
function createEmbeddingResponse(values: number[], encoding: 'base64' | 'float'): Response {
return new Response(
JSON.stringify({
data: values.map((value, index) => ({ embedding: createTestEmbedding(value), index })),
data: values.map((value, index) => ({
embedding:
encoding === 'base64'
? Buffer.from(new Float32Array(createTestEmbedding(value)).buffer).toString('base64')
: createTestEmbedding(value),
index,
})),
usage: { prompt_tokens: values.length, total_tokens: values.length },
}),
{ status: 200, headers: { 'Content-Type': 'application/json' } }
)
}

function createEmbeddingFetchMock() {
return vi.fn().mockResolvedValue(createEmbeddingResponse([0.1, 0.3]))
return vi.fn().mockResolvedValue(createEmbeddingResponse([0.1, 0.3], 'base64'))
}

vi.stubGlobal('fetch', createEmbeddingFetchMock())
Expand Down Expand Up @@ -286,7 +293,7 @@ describe('Knowledge Utils', () => {
})

const fetchSpy = vi.mocked(fetch)
fetchSpy.mockResolvedValueOnce(createEmbeddingResponse([0.1]))
fetchSpy.mockResolvedValueOnce(createEmbeddingResponse([0.1], 'float'))

await generateEmbeddings(['test text'], DEFAULT_EMBEDDING_TARGET)

Expand All @@ -310,7 +317,7 @@ describe('Knowledge Utils', () => {
})

const fetchSpy = vi.mocked(fetch)
fetchSpy.mockResolvedValueOnce(createEmbeddingResponse([0.1]))
fetchSpy.mockResolvedValueOnce(createEmbeddingResponse([0.1], 'base64'))

await generateEmbeddings(['test text'], DEFAULT_EMBEDDING_TARGET)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,7 @@ vi.mock('@/lib/knowledge/application/contexts', () => ({
}))

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

vi.mock('@/lib/knowledge/embeddings', () => ({
Expand Down
48 changes: 48 additions & 0 deletions apps/sim/lib/core/rate-limiter/storage/db-token-bucket.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -47,4 +47,52 @@ describe('PostgreSQL token bucket', () => {
})
expect(dbChainMockFns.set).toHaveBeenCalledWith(expect.objectContaining({ tokens: '0' }))
})

it('initializes request and cooldown buckets in a consistent order without duplicate keys', async () => {
const now = new Date()
dbChainMockFns.limit.mockResolvedValue([
{ key: 'cooldown', tokens: '0', lastRefillAt: now, blockedUntil: null },
{ key: 'requests', tokens: '10', lastRefillAt: now, blockedUntil: null },
{ key: 'tokens', tokens: '10', lastRefillAt: now, blockedUntil: null },
])

expect(
await new DbTokenBucket().consumeTokensAtomically(
[
{ key: 'tokens', cost: 2, config: CONFIG },
{ key: 'requests', cost: 1, config: CONFIG },
],
{ cooldownKeys: ['cooldown', 'cooldown'], deadlineAt: now.getTime() + 1000 }
)
).toEqual({ allowed: true, retryAfterMs: 0 })

expect(dbChainMockFns.values).toHaveBeenCalledExactlyOnceWith([
{ key: 'cooldown', tokens: '0', lastRefillAt: now, updatedAt: now },
{ key: 'requests', tokens: '10', lastRefillAt: now, updatedAt: now },
{ key: 'tokens', tokens: '10', lastRefillAt: now, updatedAt: now },
])
expect(dbChainMockFns.set).toHaveBeenCalledWith({
tokens: '8',
lastRefillAt: now,
updatedAt: now,
})
expect(dbChainMockFns.set).toHaveBeenCalledWith({
tokens: '9',
lastRefillAt: now,
updatedAt: now,
})
})

it('accepts an empty reservation without attempting an empty insert', async () => {
dbChainMockFns.limit.mockResolvedValue([])

expect(
await new DbTokenBucket().consumeTokensAtomically([], {
cooldownKeys: [],
deadlineAt: Date.now() + 1000,
})
).toEqual({ allowed: true, retryAfterMs: 0 })
expect(dbChainMockFns.insert).not.toHaveBeenCalled()
expect(dbChainMockFns.update).not.toHaveBeenCalled()
})
})
20 changes: 12 additions & 8 deletions apps/sim/lib/core/rate-limiter/storage/db-token-bucket.ts
Original file line number Diff line number Diff line change
Expand Up @@ -94,16 +94,20 @@ export class DbTokenBucket implements RateLimitStorageAdapter {
...new Set([...options.cooldownKeys, ...reservations.map((item) => item.key)]),
].sort()
const createdAt = new Date()
for (const key of keys) {
const reservation = reservations.find((item) => item.key === key)
if (keys.length > 0) {
await tx
.insert(rateLimitBucket)
.values({
key,
tokens: String(reservation?.config.maxTokens ?? 0),
lastRefillAt: createdAt,
updatedAt: createdAt,
})
.values(
keys.map((key) => {
const reservation = reservations.find((item) => item.key === key)
return {
key,
tokens: String(reservation?.config.maxTokens ?? 0),
lastRefillAt: createdAt,
updatedAt: createdAt,
}
})
)
.onConflictDoNothing()
}
const rows = await tx
Expand Down
Loading
Loading