feat(knowledge): add embedding model selection and Cohere reranker (#4349)

* feat(knowledge): add embedding model selection and Cohere reranker

* fix(knowledge): split reranker model constants into client-safe module

* fix(knowledge): bill rerank on every successful API call and fix MDX docs literal

* test(knowledge): align embedding tests with provider abstraction changes

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): require explicit Azure deployment per OpenAI embedding model

Greptile P1: when AZURE_OPENAI_* was set, every OpenAI embedding model was
routed to the single KB_OPENAI_MODEL_NAME deployment. A KB created with
text-embedding-3-large would be embedded by whatever model that deployment
serves while billing tracked 3-large pricing — and chunks ingested via Azure
versus queried via real OpenAI would land in mismatched vector spaces.

Now require AZURE_OPENAI_DEPLOYMENT_TEXT_EMBEDDING_3_(SMALL|LARGE) per model.
Falls back to KB_OPENAI_MODEL_NAME only for text-embedding-3-small (legacy).
If no deployment is configured for the chosen model, route to direct OpenAI
instead of silently routing to the wrong deployment.

Also fix type predicate in search/route.ts to use KnowledgeBaseAccessResult
so the build passes.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): skip platform reranker billing for BYOK Cohere keys

Cursor bugbot found that resolveCohereKey discarded BYOK status, so the
search route always added platform rerankerCost even when the workspace
supplied its own Cohere key.

Now resolveCohereKey returns { apiKey, isBYOK } and rerank() returns
{ results, isBYOK }. The search route checks rerankIsBYOK before adding
rerankerCost or emitting the rerankerCost/rerankerSearchUnits fields,
mirroring how generateEmbeddings handles BYOK billing.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): match search tokenizer to embedding provider; remove dead var

Cursor bugbot:
- Token estimation was hardcoded to 'openai' for every embedding model.
  For gemini-embedding-001 the cost was computed against an OpenAI-tokenized
  count, producing wrong input.tokens.prompt and (slightly) wrong cost.
  Now derive the tokenizer provider from the embedding model's provider.
- rerankApplied was set but never read. Removed.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): match chunk tokenizer to KB embedding provider

Cursor bugbot: createChunk and updateChunk hardcoded the 'openai' tokenizer
when computing the stored tokenCount. For KBs using gemini-embedding-001 the
count was estimated with the wrong heuristic, leading to inaccurate stored
counts (and any billing derived from them). Now derive the tokenizer from
the KB's embedding model provider, matching the search route.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* refactor(knowledge): centralize tokenizer mapping on EmbeddingModelInfo

Add tokenizerProvider directly to EmbeddingModelInfo so callers read it
from the registry instead of reimplementing the gemini→google / openai→openai
map at each call site. Removes the local helper in chunks/service.ts and
the inline ternary in search/route.ts.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* refactor(knowledge): lock embedding model to KB_EMBEDDING_MODEL env var

Remove the user-facing model picker from the KB create modal and the
embeddingModel field from the create/update API schemas. The active model
is now selected server-side via KB_EMBEDDING_MODEL, which collapses Azure
routing to a single deployment (KB_OPENAI_MODEL_NAME) and drops the
per-model AZURE_OPENAI_DEPLOYMENT_TEXT_EMBEDDING_3_* env vars and
SUPPORTED_EMBEDDING_MODEL_IDS / UI-only label+description registry fields.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): use provider tokenizer for chunks and bound rerank indices

- documents/service.ts: replace ceil(len/4) heuristic with estimateTokenCount using the embedding model's tokenizerProvider so token counts match billing
- reranker.ts: filter Cohere rerank results to valid indices before mapping to defend against malformed responses
- utils.test.ts: add embeddingModel to kb fixture so getEmbeddingModelInfo resolves

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): use .count from estimateTokenCount return value

estimateTokenCount returns a TokenEstimate object, not a number — access
.count so the integer token count is stored instead of an object.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): only enforce single embedding model when query is present

Tag-only searches don't generate a query embedding, so two KBs with
different embedding models can be filtered together. Gate the guard on
hasQuery so cross-model tag-only queries no longer 400.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): use getConfiguredEmbeddingModel in copilot KB creation

Copilot-created KBs were hardcoded to text-embedding-3-small, ignoring
KB_EMBEDDING_MODEL. This caused cross-KB searches mixing copilot- and
API-created KBs to hit the embedding-model-mismatch guard.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): make EMBEDDING_DIMENSIONS a literal type

CreateKnowledgeBaseData.embeddingDimension is typed as the literal 1536,
so EMBEDDING_DIMENSIONS needs `as const` to satisfy it after the copilot
path switched to passing the constant.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): use per-KB embedding model in v1 search route

The v1 search endpoint was passing undefined to generateSearchEmbedding,
which silently fell back to text-embedding-3-small. KBs created while
KB_EMBEDDING_MODEL=gemini-embedding-001 (or any non-default) would have
their queries embedded with the wrong model. Now resolves the model from
the KB rows like the internal route, with the same multi-model guard.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* chore(knowledge): polish embedding/reranker implementation

- Drop unused supportsCustomDimensions from EmbeddingModelInfo (every
  registered model supports it; OpenAI/Azure paths now always send
  dimensions: 1536).
- Type SUPPORTED_EMBEDDING_MODELS as Partial<Record<...>> so index lookups
  surface as possibly-undefined in the type system instead of relying on
  runtime null checks alone.
- Require AZURE_OPENAI_API_VERSION in the Azure routing gate. Missing
  api-version no longer slips through as ?api-version=undefined; it now
  falls back to direct OpenAI.
- Use the embedding provider's tokenizer (estimateTokenCount) for the
  Gemini fallback token estimate instead of len/4, so billing matches
  the model's tokenization.
- Drop unreachable 'text-embedding-3-small' fallback in the manual chunk
  upload route — accessCheck.knowledgeBase is non-null after the access
  guard.
- docs-chunker now reads getConfiguredEmbeddingModel() so Sim's docs
  ingestion respects KB_EMBEDDING_MODEL like the user-facing paths.
- Add v1 search route test covering per-KB model resolution and the
  cross-KB mixed-model rejection.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): resolve type errors and unhandled rejection in search routes

- Use accessCheck.knowledgeBase.embeddingModel directly in chunks response
- Narrow access-check predicate to KnowledgeBaseAccessResult in v1 search
- Move inaccessible-KB 404 check before query embedding promise creation

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): pass Gemini API key via x-goog-api-key header

URLs end up in server access logs, proxy logs, and APM tools, so embedding
the key as a query param risks accidental exposure. Google explicitly
recommends the header form for the Gemini REST API.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): default Azure deployment name to embedding model name

Restore the prior fallback so existing Azure deployments — which conventionally
name the deployment after the model — continue to route through Azure when
KB_OPENAI_MODEL_NAME is unset. Before this fix, those deployments silently fell
through to direct OpenAI.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): cap Gemini batches at 100 items, add singular GEMINI_API_KEY fallback

- Gemini's batchEmbedContents API rejects requests with more than 100
  items. The token-based batcher could pack hundreds of short chunks
  into a single request, causing 400s. Add maxItemsPerRequest on
  ResolvedProvider and split token batches further when set.
- Mirror resolveOpenAIKey by accepting GEMINI_API_KEY (singular) as a
  fallback before requiring the rotating GEMINI_API_KEY_1/2/3 keys.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

* fix(knowledge): prefer singular Cohere key before rotation

Match resolveOpenAIKey/resolveGeminiKey order: check the singular
COHERE_API_KEY before falling back to rotating keys.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
Waleed
2026-04-30 01:45:56 -07:00
committed by GitHub
co-authored by Claude Opus 4.7
parent cb8ea3a870
commit d94f4c9943
29 changed files with 1041 additions and 190 deletions
@@ -47,6 +47,8 @@ Search for similar content in a knowledge base using vector similarity
| `properties` | string | No | No description |
| `tagName` | string | No | No description |
| `tagValue` | string | No | No description |
| `rerankerEnabled` | boolean | No | Whether to apply Cohere reranking to vector search results |
| `rerankerModel` | string | No | Cohere rerank model to use \(one of: rerank-v4.0-pro, rerank-v4.0-fast, rerank-v3.5\) |
| `tagFilters` | string | No | No description |
#### Output
@@ -215,7 +215,12 @@ export const POST = withRouteHandler(
let cost = null
try {
cost = calculateCost('text-embedding-3-small', newChunk.tokenCount, 0, false)
cost = calculateCost(
accessCheck.knowledgeBase.embeddingModel,
newChunk.tokenCount,
0,
false
)
} catch (error) {
logger.warn(`[${requestId}] Failed to calculate cost for chunk upload`, {
error: error instanceof Error ? error.message : 'Unknown error',
@@ -240,7 +245,7 @@ export const POST = withRouteHandler(
completion: 0,
total: newChunk.tokenCount,
},
model: 'text-embedding-3-small',
model: accessCheck.knowledgeBase.embeddingModel,
pricing: cost.pricing,
},
}
-2
View File
@@ -27,8 +27,6 @@ const logger = createLogger('KnowledgeBaseByIdAPI')
const UpdateKnowledgeBaseSchema = z.object({
name: z.string().min(1, 'Name is required').optional(),
description: z.string().optional(),
embeddingModel: z.literal('text-embedding-3-small').optional(),
embeddingDimension: z.literal(1536).optional(),
workspaceId: z.string().nullable().optional(),
chunkingConfig: z
.object({
+7 -4
View File
@@ -6,6 +6,7 @@ import { getSession } from '@/lib/auth'
import { PlatformEvents } from '@/lib/core/telemetry'
import { generateRequestId } from '@/lib/core/utils/request'
import { withRouteHandler } from '@/lib/core/utils/with-route-handler'
import { EMBEDDING_DIMENSIONS, getConfiguredEmbeddingModel } from '@/lib/knowledge/embeddings'
import {
createKnowledgeBase,
getKnowledgeBases,
@@ -20,8 +21,6 @@ const CreateKnowledgeBaseSchema = z.object({
name: z.string().min(1, 'Name is required'),
description: z.string().optional(),
workspaceId: z.string().min(1, 'Workspace ID is required'),
embeddingModel: z.literal('text-embedding-3-small').default('text-embedding-3-small'),
embeddingDimension: z.literal(1536).default(1536),
chunkingConfig: z
.object({
maxSize: z.number().min(100).max(4000).default(1024),
@@ -118,9 +117,13 @@ export const POST = withRouteHandler(async (req: NextRequest) => {
try {
const validatedData = CreateKnowledgeBaseSchema.parse(body)
const embeddingModel = getConfiguredEmbeddingModel()
const createData = {
...validatedData,
userId: session.user.id,
embeddingModel,
embeddingDimension: EMBEDDING_DIMENSIONS,
}
const newKnowledgeBase = await createKnowledgeBase(createData, requestId)
@@ -166,8 +169,8 @@ export const POST = withRouteHandler(async (req: NextRequest) => {
metadata: {
name: validatedData.name,
description: validatedData.description,
embeddingModel: validatedData.embeddingModel,
embeddingDimension: validatedData.embeddingDimension,
embeddingModel,
embeddingDimension: EMBEDDING_DIMENSIONS,
chunkingStrategy: validatedData.chunkingConfig.strategy,
chunkMaxSize: validatedData.chunkingConfig.maxSize,
chunkMinSize: validatedData.chunkingConfig.minSize,
@@ -432,6 +432,7 @@ describe('Knowledge Search API Route', () => {
userId: 'user-123',
name: 'Test KB',
deletedAt: null,
embeddingModel: 'text-embedding-3-small',
},
})
@@ -524,6 +525,7 @@ describe('Knowledge Search API Route', () => {
userId: 'user-123',
name: 'Test KB',
deletedAt: null,
embeddingModel: 'text-embedding-3-small',
},
})
@@ -571,6 +573,7 @@ describe('Knowledge Search API Route', () => {
userId: 'user-123',
name: 'Test KB',
deletedAt: null,
embeddingModel: 'text-embedding-3-small',
},
})
@@ -625,6 +628,7 @@ describe('Knowledge Search API Route', () => {
userId: 'user-123',
name: 'Test KB',
deletedAt: null,
embeddingModel: 'text-embedding-3-small',
},
})
@@ -694,6 +698,7 @@ describe('Knowledge Search API Route', () => {
userId: 'user-123',
name: 'Test KB',
deletedAt: null,
embeddingModel: 'text-embedding-3-small',
},
})
@@ -739,6 +744,7 @@ describe('Knowledge Search API Route', () => {
userId: 'user-123',
name: 'Test KB',
deletedAt: null,
embeddingModel: 'text-embedding-3-small',
},
})
@@ -877,6 +883,7 @@ describe('Knowledge Search API Route', () => {
userId: 'user-123',
name: 'Test KB',
deletedAt: null,
embeddingModel: 'text-embedding-3-small',
},
})
@@ -921,11 +928,17 @@ describe('Knowledge Search API Route', () => {
userId: 'user-123',
name: 'Test KB',
deletedAt: null,
embeddingModel: 'text-embedding-3-small',
},
})
.mockResolvedValueOnce({
hasAccess: true,
knowledgeBase: { id: 'kb-456', userId: 'user-123', name: 'Test KB 2' },
knowledgeBase: {
id: 'kb-456',
userId: 'user-123',
name: 'Test KB 2',
embeddingModel: 'text-embedding-3-small',
},
})
mockGetDocumentTagDefinitions.mockResolvedValue(mockTagDefinitions)
+123 -15
View File
@@ -7,6 +7,8 @@ import { PlatformEvents } from '@/lib/core/telemetry'
import { generateRequestId } from '@/lib/core/utils/request'
import { withRouteHandler } from '@/lib/core/utils/with-route-handler'
import { ALL_TAG_SLOTS } from '@/lib/knowledge/constants'
import { getEmbeddingModelInfo } from '@/lib/knowledge/embedding-models'
import { DEFAULT_RERANKER_MODEL, rerank, SUPPORTED_RERANKER_MODELS } from '@/lib/knowledge/reranker'
import { getDocumentTagDefinitions } from '@/lib/knowledge/tags/service'
import { buildUndefinedTagsError, validateTagValue } from '@/lib/knowledge/tags/utils'
import type { StructuredFilter } from '@/lib/knowledge/types'
@@ -20,7 +22,8 @@ import {
handleVectorOnlySearch,
type SearchResult,
} from '@/app/api/knowledge/search/utils'
import { checkKnowledgeBaseAccess } from '@/app/api/knowledge/utils'
import { checkKnowledgeBaseAccess, type KnowledgeBaseAccessResult } from '@/app/api/knowledge/utils'
import { getRerankModelPricing } from '@/providers/models'
import { calculateCost } from '@/providers/utils'
const logger = createLogger('VectorSearchAPI')
@@ -59,6 +62,11 @@ const VectorSearchSchema = z
.optional()
.nullable()
.transform((val) => val || undefined),
rerankerEnabled: z.boolean().optional().default(false),
rerankerModel: z
.enum(SUPPORTED_RERANKER_MODELS as unknown as [string, ...string[]])
.optional()
.default(DEFAULT_RERANKER_MODEL),
})
.refine(
(data) => {
@@ -235,12 +243,26 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
)
}
const workspaceId = accessChecks.find((ac) => ac?.hasAccess)?.knowledgeBase?.workspaceId
const accessibleKbs = accessChecks
.filter((ac): ac is KnowledgeBaseAccessResult => Boolean(ac?.hasAccess))
.map((ac) => ac.knowledgeBase)
const workspaceId = accessibleKbs[0]?.workspaceId
const useReranker = validatedData.rerankerEnabled && Boolean(validatedData.query?.trim())
const rerankerModel = useReranker ? validatedData.rerankerModel : null
const hasQuery = validatedData.query && validatedData.query.trim().length > 0
const queryEmbeddingPromise = hasQuery
? generateSearchEmbedding(validatedData.query!, undefined, workspaceId)
: Promise.resolve(null)
const embeddingModels = Array.from(new Set(accessibleKbs.map((kb) => kb.embeddingModel)))
if (hasQuery && embeddingModels.length > 1) {
return NextResponse.json(
{
error:
'Selected knowledge bases use different embedding models and cannot be searched together. Search them separately.',
},
{ status: 400 }
)
}
const queryEmbeddingModel = embeddingModels[0]
// Check if any requested knowledge bases were not accessible
const inaccessibleKbIds = knowledgeBaseIds.filter((id) => !accessibleKbIds.includes(id))
@@ -252,6 +274,10 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
)
}
const queryEmbeddingPromise = hasQuery
? generateSearchEmbedding(validatedData.query!, queryEmbeddingModel, workspaceId)
: Promise.resolve(null)
if (workflowId) {
const authorization = await authorizeWorkflowByWorkspacePermission({
workflowId,
@@ -278,6 +304,10 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
const hasFilters = structuredFilters && structuredFilters.length > 0
// Oversample candidates when reranking so the reranker has more to choose from.
// Cap at 100 to bound Cohere request cost (1 search unit = ≤100 docs).
const candidateTopK = useReranker ? Math.min(100, validatedData.topK * 4) : validatedData.topK
if (!hasQuery && hasFilters) {
// Tag-only search without vector similarity
results = await handleTagOnlySearch({
@@ -291,24 +321,24 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
`[${requestId}] Executing tag + vector search with filters:`,
structuredFilters
)
const strategy = getQueryStrategy(accessibleKbIds.length, validatedData.topK)
const strategy = getQueryStrategy(accessibleKbIds.length, candidateTopK)
const queryVector = JSON.stringify(await queryEmbeddingPromise)
results = await handleTagAndVectorSearch({
knowledgeBaseIds: accessibleKbIds,
topK: validatedData.topK,
topK: candidateTopK,
structuredFilters,
queryVector,
distanceThreshold: strategy.distanceThreshold,
})
} else if (hasQuery && !hasFilters) {
// Vector-only search
const strategy = getQueryStrategy(accessibleKbIds.length, validatedData.topK)
const strategy = getQueryStrategy(accessibleKbIds.length, candidateTopK)
const queryVector = JSON.stringify(await queryEmbeddingPromise)
results = await handleVectorOnlySearch({
knowledgeBaseIds: accessibleKbIds,
topK: validatedData.topK,
topK: candidateTopK,
queryVector,
distanceThreshold: strategy.distanceThreshold,
})
@@ -323,13 +353,60 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
)
}
// Optional Cohere rerank pass on top of vector results.
const rerankedScores = new Map<string, number>()
// `rerankBilled` = Cohere was successfully called (even with 0 results) and we owe the search unit.
let rerankBilled = false
let rerankIsBYOK = false
if (useReranker && rerankerModel && results.length > 0) {
const candidateCount = results.length
try {
const { results: ranked, isBYOK } = await rerank(
validatedData.query!,
results.map((r) => ({ id: r.id, text: r.content })),
{ model: rerankerModel, topN: validatedData.topK, workspaceId }
)
rerankBilled = true
rerankIsBYOK = isBYOK
if (ranked.length === 0) {
logger.warn(
`[${requestId}] Reranker returned 0 results; falling back to vector ordering`,
{ model: rerankerModel, candidateCount }
)
results = results.slice(0, validatedData.topK)
} else {
const idToResult = new Map(results.map((r) => [r.id, r]))
results = ranked
.map((r) => idToResult.get(r.item.id))
.filter((r): r is SearchResult => Boolean(r))
for (const r of ranked) rerankedScores.set(r.item.id, r.relevanceScore)
logger.info(`[${requestId}] Reranked ${candidateCount} → ${results.length} results`, {
model: rerankerModel,
})
}
} catch (error) {
logger.warn(`[${requestId}] Reranker failed; falling back to vector ordering`, {
error: error instanceof Error ? error.message : 'Unknown error',
model: rerankerModel,
candidateCount,
workspaceId,
})
results = results.slice(0, validatedData.topK)
}
} else if (useReranker) {
results = results.slice(0, validatedData.topK)
}
// Calculate cost for the embedding (with fallback if calculation fails)
let cost = null
let tokenCount = null
if (hasQuery) {
try {
tokenCount = estimateTokenCount(validatedData.query!, 'openai')
cost = calculateCost('text-embedding-3-small', tokenCount.count, 0, false)
tokenCount = estimateTokenCount(
validatedData.query!,
getEmbeddingModelInfo(queryEmbeddingModel).tokenizerProvider
)
cost = calculateCost(queryEmbeddingModel, tokenCount.count, 0, false)
} catch (error) {
logger.warn(`[${requestId}] Failed to calculate cost for search query`, {
error: error instanceof Error ? error.message : 'Unknown error',
@@ -338,6 +415,32 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
}
}
// Add Cohere rerank cost (1 search unit per successful call, since we cap candidates ≤100).
// Bill on every successful API response — Cohere charges even when 0 results are returned.
let rerankerCost = 0
if (rerankBilled && rerankerModel && !rerankIsBYOK) {
const pricing = getRerankModelPricing(rerankerModel)
if (pricing) {
rerankerCost = pricing.perSearchUnit
if (cost) {
cost = {
...cost,
input: cost.input + rerankerCost,
total: cost.total + rerankerCost,
}
} else {
cost = {
input: rerankerCost,
output: 0,
total: rerankerCost,
pricing: { input: 0, output: 0, updatedAt: pricing.updatedAt },
}
}
} else {
logger.warn(`[${requestId}] No pricing entry for rerank model ${rerankerModel}`)
}
}
// Fetch tag definitions for display name mapping (reuse the same fetch from filtering)
const tagDefsResults = await Promise.all(
accessibleKbIds.map(async (kbId) => {
@@ -400,6 +503,7 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
}
})
const rerankerScore = rerankedScores.get(result.id)
return {
documentId: result.documentId,
documentName: documentNameMap[result.documentId] || undefined,
@@ -407,6 +511,7 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
chunkIndex: result.chunkIndex,
metadata: tags, // Clean display name mapped tags
similarity: hasQuery ? 1 - result.distance : 1, // Perfect similarity for tag-only searches
...(rerankerScore !== undefined && { rerankerScore }),
}
}),
query: validatedData.query || '',
@@ -414,19 +519,22 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
knowledgeBaseId: accessibleKbIds[0],
topK: validatedData.topK,
totalResults: results.length,
...(cost && tokenCount
...(cost
? {
cost: {
input: cost.input,
output: cost.output,
total: cost.total,
tokens: {
prompt: tokenCount.count,
prompt: tokenCount?.count ?? 0,
completion: 0,
total: tokenCount.count,
total: tokenCount?.count ?? 0,
},
model: 'text-embedding-3-small',
model: queryEmbeddingModel,
pricing: cost.pricing,
...(rerankBilled && !rerankIsBYOK
? { rerankerCost, rerankerModel, rerankerSearchUnits: 1 }
: {}),
},
}
: {}),
@@ -220,7 +220,7 @@ describe('Knowledge Search Utils', () => {
Object.keys(env).forEach((key) => delete (env as any)[key])
})
it('should use default API version when not provided in Azure config', async () => {
it('falls back to OpenAI when AZURE_OPENAI_API_VERSION is not set', async () => {
const { env } = await import('@/lib/core/config/env')
Object.keys(env).forEach((key) => delete (env as any)[key])
Object.assign(env, {
@@ -240,7 +240,7 @@ describe('Knowledge Search Utils', () => {
await generateSearchEmbedding('test query')
expect(vi.mocked(fetch)).toHaveBeenCalledWith(
expect.stringContaining('api-version='),
'https://api.openai.com/v1/embeddings',
expect.any(Object)
)
@@ -282,7 +282,7 @@ describe('Knowledge Search Utils', () => {
Object.keys(env).forEach((key) => delete (env as any)[key])
await expect(generateSearchEmbedding('test query')).rejects.toThrow(
'Either OPENAI_API_KEY or Azure OpenAI configuration (AZURE_OPENAI_API_KEY + AZURE_OPENAI_ENDPOINT) must be configured'
'OPENAI_API_KEY is not configured'
)
})
@@ -354,6 +354,7 @@ describe('Knowledge Search Utils', () => {
body: JSON.stringify({
input: ['test query'],
encoding_format: 'float',
dimensions: 1536,
}),
})
)
+2 -1
View File
@@ -212,6 +212,7 @@ describe('Knowledge Utils', () => {
id: 'kb1',
userId: 'user1',
workspaceId: null,
embeddingModel: 'text-embedding-3-small',
chunkingConfig: { maxSize: 1024, minSize: 1, overlap: 200 },
})
docRows.push({ id: 'doc1', knowledgeBaseId: 'kb1' })
@@ -370,7 +371,7 @@ describe('Knowledge Utils', () => {
Object.keys(env).forEach((key) => delete (env as any)[key])
await expect(generateEmbeddings(['test text'])).rejects.toThrow(
'Either OPENAI_API_KEY or Azure OpenAI configuration (AZURE_OPENAI_API_KEY + AZURE_OPENAI_ENDPOINT) must be configured'
'OPENAI_API_KEY is not configured'
)
})
})
+14 -3
View File
@@ -103,7 +103,10 @@ export interface EmbeddingData {
export interface KnowledgeBaseAccessResult {
hasAccess: true
knowledgeBase: Pick<KnowledgeBaseData, 'id' | 'userId' | 'workspaceId' | 'name'>
knowledgeBase: Pick<
KnowledgeBaseData,
'id' | 'userId' | 'workspaceId' | 'name' | 'embeddingModel'
>
}
export interface KnowledgeBaseAccessDenied {
@@ -117,7 +120,10 @@ export type KnowledgeBaseAccessCheck = KnowledgeBaseAccessResult | KnowledgeBase
export interface DocumentAccessResult {
hasAccess: true
document: DocumentData
knowledgeBase: Pick<KnowledgeBaseData, 'id' | 'userId' | 'workspaceId' | 'name'>
knowledgeBase: Pick<
KnowledgeBaseData,
'id' | 'userId' | 'workspaceId' | 'name' | 'embeddingModel'
>
}
export interface DocumentAccessDenied {
@@ -132,7 +138,10 @@ export interface ChunkAccessResult {
hasAccess: true
chunk: EmbeddingData
document: DocumentData
knowledgeBase: Pick<KnowledgeBaseData, 'id' | 'userId' | 'workspaceId' | 'name'>
knowledgeBase: Pick<
KnowledgeBaseData,
'id' | 'userId' | 'workspaceId' | 'name' | 'embeddingModel'
>
}
export interface ChunkAccessDenied {
@@ -156,6 +165,7 @@ export async function checkKnowledgeBaseAccess(
userId: knowledgeBase.userId,
workspaceId: knowledgeBase.workspaceId,
name: knowledgeBase.name,
embeddingModel: knowledgeBase.embeddingModel,
})
.from(knowledgeBase)
.where(and(eq(knowledgeBase.id, knowledgeBaseId), isNull(knowledgeBase.deletedAt)))
@@ -200,6 +210,7 @@ export async function checkKnowledgeBaseWriteAccess(
userId: knowledgeBase.userId,
workspaceId: knowledgeBase.workspaceId,
name: knowledgeBase.name,
embeddingModel: knowledgeBase.embeddingModel,
})
.from(knowledgeBase)
.where(and(eq(knowledgeBase.id, knowledgeBaseId), isNull(knowledgeBase.deletedAt)))
+3 -2
View File
@@ -2,6 +2,7 @@ import { AuditAction, AuditResourceType, recordAudit } from '@sim/audit'
import { type NextRequest, NextResponse } from 'next/server'
import { z } from 'zod'
import { withRouteHandler } from '@/lib/core/utils/with-route-handler'
import { EMBEDDING_DIMENSIONS, getConfiguredEmbeddingModel } from '@/lib/knowledge/embeddings'
import { createKnowledgeBase, getKnowledgeBases } from '@/lib/knowledge/service'
import {
authenticateRequest,
@@ -92,8 +93,8 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
description,
workspaceId,
userId,
embeddingModel: 'text-embedding-3-small',
embeddingDimension: 1536,
embeddingModel: getConfiguredEmbeddingModel(),
embeddingDimension: EMBEDDING_DIMENSIONS,
chunkingConfig: chunkingConfig ?? { maxSize: 1024, minSize: 100, overlap: 200 },
},
requestId
@@ -0,0 +1,179 @@
/**
* Tests for v1 knowledge search API route.
* Specifically guards the per-KB embedding model resolution and the
* multi-model rejection so the v1 endpoint stays in lockstep with the
* internal route.
*
* @vitest-environment node
*/
import { createMockRequest, knowledgeApiUtilsMock, knowledgeApiUtilsMockFns } from '@sim/testing'
import { beforeEach, describe, expect, it, vi } from 'vitest'
const {
mockHandleVectorOnlySearch,
mockHandleTagOnlySearch,
mockHandleTagAndVectorSearch,
mockGetQueryStrategy,
mockGenerateSearchEmbedding,
mockGetDocumentNamesByIds,
mockAuthenticateRequest,
mockValidateWorkspaceAccess,
} = vi.hoisted(() => ({
mockHandleVectorOnlySearch: vi.fn(),
mockHandleTagOnlySearch: vi.fn(),
mockHandleTagAndVectorSearch: vi.fn(),
mockGetQueryStrategy: vi.fn(),
mockGenerateSearchEmbedding: vi.fn(),
mockGetDocumentNamesByIds: vi.fn(),
mockAuthenticateRequest: vi.fn(),
mockValidateWorkspaceAccess: vi.fn(),
}))
vi.mock('@/app/api/knowledge/search/utils', () => ({
handleVectorOnlySearch: mockHandleVectorOnlySearch,
handleTagOnlySearch: mockHandleTagOnlySearch,
handleTagAndVectorSearch: mockHandleTagAndVectorSearch,
getQueryStrategy: mockGetQueryStrategy,
generateSearchEmbedding: mockGenerateSearchEmbedding,
getDocumentNamesByIds: mockGetDocumentNamesByIds,
}))
vi.mock('@/app/api/knowledge/utils', () => knowledgeApiUtilsMock)
vi.mock('@/app/api/v1/knowledge/utils', () => ({
authenticateRequest: mockAuthenticateRequest,
validateWorkspaceAccess: mockValidateWorkspaceAccess,
parseJsonBody: async (req: Request) => {
try {
return { success: true, data: await req.json() }
} catch {
return {
success: false,
response: new Response(JSON.stringify({ error: 'Invalid JSON' }), { status: 400 }),
}
}
},
validateSchema: <T>(
schema: {
safeParse: (v: unknown) => {
success: boolean
data?: T
error?: { issues: { message: string }[] }
}
},
data: unknown
) => {
const result = schema.safeParse(data)
if (!result.success) {
return {
success: false,
response: new Response(
JSON.stringify({ error: result.error?.issues.map((i) => i.message).join(', ') }),
{ status: 400 }
),
}
}
return { success: true, data: result.data }
},
handleError: (e: unknown) =>
new Response(JSON.stringify({ error: e instanceof Error ? e.message : 'error' }), {
status: 500,
}),
}))
vi.mock('@/lib/knowledge/tags/service', () => ({
getDocumentTagDefinitions: vi.fn().mockResolvedValue([]),
}))
import { POST } from '@/app/api/v1/knowledge/search/route'
const mockCheckKnowledgeBaseAccess = knowledgeApiUtilsMockFns.mockCheckKnowledgeBaseAccess
const baseKb = (id: string, embeddingModel: string) => ({
id,
userId: 'user-1',
name: `KB ${id}`,
workspaceId: 'ws-1',
embeddingModel,
deletedAt: null,
})
describe('v1 knowledge search route — per-KB embedding model', () => {
beforeEach(() => {
vi.clearAllMocks()
mockAuthenticateRequest.mockResolvedValue({
requestId: 'req-1',
userId: 'user-1',
rateLimit: {},
})
mockValidateWorkspaceAccess.mockResolvedValue(null)
mockGetQueryStrategy.mockReturnValue({ distanceThreshold: 0.5 })
mockGenerateSearchEmbedding.mockResolvedValue([0.1, 0.2, 0.3])
mockHandleVectorOnlySearch.mockResolvedValue([])
mockGetDocumentNamesByIds.mockResolvedValue({})
})
it('passes the KB embedding model into generateSearchEmbedding', async () => {
mockCheckKnowledgeBaseAccess.mockResolvedValueOnce({
hasAccess: true,
knowledgeBase: baseKb('kb-gemini', 'gemini-embedding-001'),
})
const req = createMockRequest('POST', {
workspaceId: 'ws-1',
knowledgeBaseIds: 'kb-gemini',
query: 'hello',
})
const res = await POST(req)
expect(res.status).toBe(200)
expect(mockGenerateSearchEmbedding).toHaveBeenCalledWith(
'hello',
'gemini-embedding-001',
'ws-1'
)
})
it('rejects cross-KB queries with mixed embedding models', async () => {
mockCheckKnowledgeBaseAccess
.mockResolvedValueOnce({
hasAccess: true,
knowledgeBase: baseKb('kb-openai', 'text-embedding-3-small'),
})
.mockResolvedValueOnce({
hasAccess: true,
knowledgeBase: baseKb('kb-gemini', 'gemini-embedding-001'),
})
const req = createMockRequest('POST', {
workspaceId: 'ws-1',
knowledgeBaseIds: ['kb-openai', 'kb-gemini'],
query: 'hello',
})
const res = await POST(req)
expect(res.status).toBe(400)
expect(mockGenerateSearchEmbedding).not.toHaveBeenCalled()
})
it('allows tag-only search across mixed embedding models', async () => {
mockHandleTagOnlySearch.mockResolvedValue([])
mockCheckKnowledgeBaseAccess.mockResolvedValueOnce({
hasAccess: true,
knowledgeBase: baseKb('kb-mixed', 'text-embedding-3-small'),
})
const req = createMockRequest('POST', {
workspaceId: 'ws-1',
knowledgeBaseIds: 'kb-mixed',
tagFilters: [{ tagName: 'category', operator: 'eq', value: 'docs' }],
})
const res = await POST(req)
expect(res.status).toBe(400)
// tagName "category" is undefined in our empty getDocumentTagDefinitions mock,
// so the route returns 400 before reaching the search handlers — but crucially
// it never tries to generate an embedding.
expect(mockGenerateSearchEmbedding).not.toHaveBeenCalled()
})
})
+22 -8
View File
@@ -14,7 +14,7 @@ import {
handleVectorOnlySearch,
type SearchResult,
} from '@/app/api/knowledge/search/utils'
import { checkKnowledgeBaseAccess } from '@/app/api/knowledge/utils'
import { checkKnowledgeBaseAccess, type KnowledgeBaseAccessResult } from '@/app/api/knowledge/utils'
import {
authenticateRequest,
handleError,
@@ -84,11 +84,13 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
const accessChecks = await Promise.all(
knowledgeBaseIds.map((kbId) => checkKnowledgeBaseAccess(kbId, userId))
)
const accessibleKbIds = knowledgeBaseIds.filter(
(_, idx) =>
accessChecks[idx]?.hasAccess &&
accessChecks[idx]?.knowledgeBase?.workspaceId === workspaceId
)
const accessibleKbs = accessChecks
.filter(
(ac): ac is KnowledgeBaseAccessResult =>
ac.hasAccess === true && ac.knowledgeBase.workspaceId === workspaceId
)
.map((ac) => ac.knowledgeBase)
const accessibleKbIds = accessibleKbs.map((kb) => kb.id)
if (accessibleKbIds.length === 0) {
return NextResponse.json(
@@ -173,6 +175,18 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
const hasQuery = query && query.trim().length > 0
const hasFilters = structuredFilters.length > 0
const embeddingModels = Array.from(new Set(accessibleKbs.map((kb) => kb.embeddingModel)))
if (hasQuery && embeddingModels.length > 1) {
return NextResponse.json(
{
error:
'Selected knowledge bases use different embedding models and cannot be searched together. Search them separately.',
},
{ status: 400 }
)
}
const queryEmbeddingModel = embeddingModels[0]
let results: SearchResult[]
if (!hasQuery && hasFilters) {
@@ -184,7 +198,7 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
} else if (hasQuery && hasFilters) {
const strategy = getQueryStrategy(accessibleKbIds.length, topK)
const queryVector = JSON.stringify(
await generateSearchEmbedding(query!, undefined, workspaceId)
await generateSearchEmbedding(query!, queryEmbeddingModel, workspaceId)
)
results = await handleTagAndVectorSearch({
knowledgeBaseIds: accessibleKbIds,
@@ -196,7 +210,7 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
} else if (hasQuery) {
const strategy = getQueryStrategy(accessibleKbIds.length, topK)
const queryVector = JSON.stringify(
await generateSearchEmbedding(query!, undefined, workspaceId)
await generateSearchEmbedding(query!, queryEmbeddingModel, workspaceId)
)
results = await handleVectorOnlySearch({
knowledgeBaseIds: accessibleKbIds,
+21
View File
@@ -1,4 +1,5 @@
import { PackageSearchIcon } from '@/components/icons'
import { DEFAULT_RERANKER_MODEL, SUPPORTED_RERANKER_MODELS } from '@/lib/knowledge/reranker-models'
import type { BlockConfig } from '@/blocks/types'
export const KnowledgeBlock: BlockConfig = {
@@ -86,6 +87,24 @@ export const KnowledgeBlock: BlockConfig = {
dependsOn: ['knowledgeBaseSelector'],
condition: { field: 'operation', value: 'search' },
},
{
id: 'rerankerEnabled',
title: 'Rerank Results',
type: 'switch',
condition: { field: 'operation', value: 'search' },
},
{
id: 'rerankerModel',
title: 'Rerank Model',
type: 'dropdown',
options: SUPPORTED_RERANKER_MODELS.map((id) => ({ label: id, id })),
value: () => DEFAULT_RERANKER_MODEL,
condition: {
field: 'operation',
value: 'search',
and: { field: 'rerankerEnabled', value: true },
},
},
// --- List Documents ---
{
@@ -397,6 +416,8 @@ export const KnowledgeBlock: BlockConfig = {
limit: { type: 'number', description: 'Max items to return' },
offset: { type: 'number', description: 'Pagination offset' },
tagFilters: { type: 'string', description: 'Tag filter criteria' },
rerankerEnabled: { type: 'boolean', description: 'Apply Cohere reranking to search results' },
rerankerModel: { type: 'string', description: 'Cohere rerank model identifier' },
documentTags: { type: 'string', description: 'Document tags' },
chunkSearch: { type: 'string', description: 'Search filter for chunks' },
chunkEnabledFilter: { type: 'string', description: 'Filter chunks by enabled status' },
+3 -3
View File
@@ -4,7 +4,7 @@ import { createLogger } from '@sim/logger'
import { TextChunker } from '@/lib/chunkers/text-chunker'
import type { DocChunk, DocsChunkerOptions } from '@/lib/chunkers/types'
import { estimateTokens } from '@/lib/chunkers/utils'
import { generateEmbeddings } from '@/lib/knowledge/embeddings'
import { generateEmbeddings, getConfiguredEmbeddingModel } from '@/lib/knowledge/embeddings'
interface HeaderInfo {
level: number
@@ -74,9 +74,9 @@ export class DocsChunker {
const headers = this.extractHeaders(cleanedContent)
logger.info(`Generating embeddings for ${textChunks.length} chunks in ${relativePath}`)
const embeddingModel = getConfiguredEmbeddingModel()
const embeddings: number[][] =
textChunks.length > 0 ? (await generateEmbeddings(textChunks)).embeddings : []
const embeddingModel = 'text-embedding-3-small'
textChunks.length > 0 ? (await generateEmbeddings(textChunks, embeddingModel)).embeddings : []
const chunks: DocChunk[] = []
let currentPosition = 0
@@ -18,7 +18,11 @@ import {
processDocumentAsync,
updateDocument,
} from '@/lib/knowledge/documents/service'
import { generateSearchEmbedding } from '@/lib/knowledge/embeddings'
import {
EMBEDDING_DIMENSIONS,
generateSearchEmbedding,
getConfiguredEmbeddingModel,
} from '@/lib/knowledge/embeddings'
import {
createKnowledgeBase,
deleteKnowledgeBase,
@@ -107,8 +111,8 @@ export const knowledgeBaseServerTool: BaseServerTool<KnowledgeBaseArgs, Knowledg
description: args.description,
workspaceId,
userId: context.userId,
embeddingModel: 'text-embedding-3-small',
embeddingDimension: 1536,
embeddingModel: getConfiguredEmbeddingModel(),
embeddingDimension: EMBEDDING_DIMENSIONS,
chunkingConfig: args.chunkingConfig || {
maxSize: 1024,
minSize: 1,
@@ -220,7 +224,7 @@ export const knowledgeBaseServerTool: BaseServerTool<KnowledgeBaseArgs, Knowledg
const queryEmbedding = await generateSearchEmbedding(
args.query,
undefined,
kb.embeddingModel,
kb.workspaceId
)
const queryVector = JSON.stringify(queryEmbedding)
+10 -1
View File
@@ -7,7 +7,12 @@ import { env } from '@/lib/core/config/env'
* @throws Error if no API keys are configured for rotation
*/
export function getRotatingApiKey(provider: string): string {
if (provider !== 'openai' && provider !== 'anthropic' && provider !== 'gemini') {
if (
provider !== 'openai' &&
provider !== 'anthropic' &&
provider !== 'gemini' &&
provider !== 'cohere'
) {
throw new Error(`No rotation implemented for provider: ${provider}`)
}
@@ -25,6 +30,10 @@ export function getRotatingApiKey(provider: string): string {
if (env.GEMINI_API_KEY_1) keys.push(env.GEMINI_API_KEY_1)
if (env.GEMINI_API_KEY_2) keys.push(env.GEMINI_API_KEY_2)
if (env.GEMINI_API_KEY_3) keys.push(env.GEMINI_API_KEY_3)
} else if (provider === 'cohere') {
if (env.COHERE_API_KEY_1) keys.push(env.COHERE_API_KEY_1)
if (env.COHERE_API_KEY_2) keys.push(env.COHERE_API_KEY_2)
if (env.COHERE_API_KEY_3) keys.push(env.COHERE_API_KEY_3)
}
if (keys.length === 0) {
+7 -1
View File
@@ -96,6 +96,7 @@ export const env = createEnv({
ANTHROPIC_API_KEY_1: z.string().min(1).optional(), // Primary Anthropic Claude API key
ANTHROPIC_API_KEY_2: z.string().min(1).optional(), // Additional Anthropic API key for load balancing
ANTHROPIC_API_KEY_3: z.string().min(1).optional(), // Additional Anthropic API key for load balancing
GEMINI_API_KEY: z.string().min(1).optional(), // Singular Gemini API key (used as fallback when rotation keys are unset)
GEMINI_API_KEY_1: z.string().min(1).optional(), // Primary Gemini API key
GEMINI_API_KEY_2: z.string().min(1).optional(), // Additional Gemini API key for load balancing
GEMINI_API_KEY_3: z.string().min(1).optional(), // Additional Gemini API key for load balancing
@@ -103,6 +104,10 @@ export const env = createEnv({
VLLM_BASE_URL: z.string().url().optional(), // vLLM self-hosted base URL (OpenAI-compatible)
VLLM_API_KEY: z.string().optional(), // Optional bearer token for vLLM
FIREWORKS_API_KEY: z.string().optional(), // Optional Fireworks AI API key for model listing
COHERE_API_KEY: z.string().min(1).optional(), // Cohere API key for reranker (rerank-v4.0-pro, rerank-v4.0-fast, rerank-v3.5)
COHERE_API_KEY_1: z.string().min(1).optional(), // Primary Cohere API key for rotation
COHERE_API_KEY_2: z.string().min(1).optional(), // Additional Cohere API key for load balancing
COHERE_API_KEY_3: z.string().min(1).optional(), // Additional Cohere API key for load balancing
ELEVENLABS_API_KEY: z.string().min(1).optional(), // ElevenLabs API key for text-to-speech in deployed chat
SERPER_API_KEY: z.string().min(1).optional(), // Serper API key for online search
EXA_API_KEY: z.string().min(1).optional(), // Exa AI API key for enhanced online search
@@ -118,7 +123,8 @@ export const env = createEnv({
AZURE_ANTHROPIC_ENDPOINT: z.string().url().optional(), // Azure Anthropic service endpoint
AZURE_ANTHROPIC_API_KEY: z.string().min(1).optional(), // Azure Anthropic API key
AZURE_ANTHROPIC_API_VERSION: z.string().min(1).optional(), // Azure Anthropic API version (e.g. 2023-06-01)
KB_OPENAI_MODEL_NAME: z.string().optional(), // Knowledge base OpenAI model name (works with both regular OpenAI and Azure OpenAI)
KB_OPENAI_MODEL_NAME: z.string().optional(), // Azure deployment name serving the configured KB embedding model (used only when AZURE_OPENAI_* credentials are set).
KB_EMBEDDING_MODEL: z.string().optional(), // Embedding model used for all new knowledge bases. Must be one of the supported model ids; defaults to text-embedding-3-small.
WAND_OPENAI_MODEL_NAME: z.string().optional(), // Wand generation OpenAI model name (works with both regular OpenAI and Azure OpenAI)
OCR_AZURE_ENDPOINT: z.string().url().optional(), // Azure Mistral OCR service endpoint
OCR_AZURE_MODEL_NAME: z.string().optional(), // Azure Mistral OCR model name for document processing
+35 -7
View File
@@ -11,6 +11,7 @@ import type {
ChunkQueryResult,
CreateChunkData,
} from '@/lib/knowledge/chunks/types'
import { getEmbeddingModelInfo } from '@/lib/knowledge/embedding-models'
import { generateEmbeddings } from '@/lib/knowledge/embeddings'
import { estimateTokenCount } from '@/lib/tokenization/estimators'
@@ -111,10 +112,25 @@ export async function createChunk(
workspaceId?: string | null
): Promise<ChunkData> {
logger.info(`[${requestId}] Generating embedding for manual chunk`)
const { embeddings } = await generateEmbeddings([chunkData.content], undefined, workspaceId)
const kbRow = await db
.select({ embeddingModel: knowledgeBase.embeddingModel })
.from(knowledgeBase)
.where(and(eq(knowledgeBase.id, knowledgeBaseId), isNull(knowledgeBase.deletedAt)))
.limit(1)
if (kbRow.length === 0) {
throw new Error('Knowledge base not found')
}
const kbEmbeddingModel = kbRow[0].embeddingModel
const { embeddings } = await generateEmbeddings(
[chunkData.content],
kbEmbeddingModel,
workspaceId
)
// Calculate accurate token count
const tokenCount = estimateTokenCount(chunkData.content, 'openai')
const tokenCount = estimateTokenCount(
chunkData.content,
getEmbeddingModelInfo(kbEmbeddingModel).tokenizerProvider
)
const chunkId = generateId()
const now = new Date()
@@ -160,7 +176,7 @@ export async function createChunk(
contentLength: chunkData.content.length,
tokenCount: tokenCount.count,
embedding: embeddings[0],
embeddingModel: 'text-embedding-3-small',
embeddingModel: kbEmbeddingModel,
startOffset: 0, // Manual chunks don't have document offsets
endOffset: chunkData.content.length,
// Inherit text tags from parent document
@@ -360,10 +376,22 @@ export async function updateChunk(
if (content !== currentChunk[0].content) {
logger.info(`[${requestId}] Content changed, regenerating embedding for chunk ${chunkId}`)
const { embeddings } = await generateEmbeddings([content], undefined, workspaceId)
const kbRow = await tx
.select({ embeddingModel: knowledgeBase.embeddingModel })
.from(knowledgeBase)
.innerJoin(document, eq(document.knowledgeBaseId, knowledgeBase.id))
.where(eq(document.id, currentChunk[0].documentId))
.limit(1)
const chunkEmbeddingModel = kbRow[0]?.embeddingModel
if (!chunkEmbeddingModel) {
throw new Error('Knowledge base for chunk not found')
}
const { embeddings } = await generateEmbeddings([content], chunkEmbeddingModel, workspaceId)
// Calculate accurate token count
const tokenCount = estimateTokenCount(content, 'openai')
const tokenCount = estimateTokenCount(
content,
getEmbeddingModelInfo(chunkEmbeddingModel).tokenizerProvider
)
dbUpdateData.content = content
dbUpdateData.contentLength = newContentLength
+14 -5
View File
@@ -34,6 +34,7 @@ import { env } from '@/lib/core/config/env'
import { getCostMultiplier, isTriggerDevEnabled } from '@/lib/core/config/feature-flags'
import { processDocument } from '@/lib/knowledge/documents/document-processor'
import type { DocumentSortField, SortOrder } from '@/lib/knowledge/documents/types'
import { getEmbeddingModelInfo } from '@/lib/knowledge/embedding-models'
import { generateEmbeddings } from '@/lib/knowledge/embeddings'
import {
buildUndefinedTagsError,
@@ -43,6 +44,7 @@ import {
validateTagValue,
} from '@/lib/knowledge/tags/utils'
import type { ProcessedDocumentTags } from '@/lib/knowledge/types'
import { estimateTokenCount } from '@/lib/tokenization/estimators'
import { deleteFile } from '@/lib/uploads/core/storage-service'
import { extractStorageKey } from '@/lib/uploads/utils/file-utils'
import type { DocumentProcessingPayload } from '@/background/knowledge-processing'
@@ -380,6 +382,7 @@ export async function processDocumentAsync(
userId: knowledgeBase.userId,
workspaceId: knowledgeBase.workspaceId,
chunkingConfig: knowledgeBase.chunkingConfig,
embeddingModel: knowledgeBase.embeddingModel,
})
.from(knowledgeBase)
.where(and(eq(knowledgeBase.id, knowledgeBaseId), isNull(knowledgeBase.deletedAt)))
@@ -429,9 +432,11 @@ export async function processDocumentAsync(
overlap: rawConfig?.overlap ?? 200,
}
const kbEmbeddingModel = kb[0].embeddingModel
let totalEmbeddingTokens = 0
let embeddingIsBYOK = false
let embeddingModelName = 'text-embedding-3-small'
let embeddingModelName = kbEmbeddingModel
let embeddingPricingId = kbEmbeddingModel
await withTimeout(
(async () => {
@@ -480,7 +485,8 @@ export async function processDocumentAsync(
totalTokens: batchTokens,
isBYOK,
modelName,
} = await generateEmbeddings(batch, undefined, kb[0].workspaceId)
pricingId,
} = await generateEmbeddings(batch, kbEmbeddingModel, kb[0].workspaceId)
for (const emb of batchEmbeddings) {
embeddings.push(emb)
}
@@ -488,6 +494,7 @@ export async function processDocumentAsync(
if (i === 0) {
embeddingIsBYOK = isBYOK
embeddingModelName = modelName
embeddingPricingId = pricingId
}
}
}
@@ -528,6 +535,8 @@ export async function processDocumentAsync(
logger.info(`[${documentId}] Creating embedding records with tags`)
const tokenizerProvider = getEmbeddingModelInfo(kbEmbeddingModel).tokenizerProvider
const embeddingRecords = processed.chunks.map((chunk, chunkIndex) => ({
id: generateId(),
knowledgeBaseId,
@@ -536,9 +545,9 @@ export async function processDocumentAsync(
chunkHash: sha256Hex(chunk.text),
content: chunk.text,
contentLength: chunk.text.length,
tokenCount: Math.ceil(chunk.text.length / 4),
tokenCount: estimateTokenCount(chunk.text, tokenizerProvider).count,
embedding: embeddings[chunkIndex] || null,
embeddingModel: 'text-embedding-3-small',
embeddingModel: kbEmbeddingModel,
startOffset: chunk.metadata.startIndex,
endOffset: chunk.metadata.endIndex,
tag1: documentTags.tag1,
@@ -620,7 +629,7 @@ export async function processDocumentAsync(
try {
const costMultiplier = getCostMultiplier()
const { total: cost } = calculateCost(
embeddingModelName,
embeddingPricingId,
totalEmbeddingTokens,
0,
false,
@@ -0,0 +1,48 @@
/**
* Registry of embedding models supported by the platform.
* Selection happens server-side via the `KB_EMBEDDING_MODEL` env var; this
* registry exists to resolve provider, tokenizer, and pricing metadata at
* runtime for any model recorded on a knowledge base row.
*/
export const EMBEDDING_DIMENSIONS = 1536 as const
export const DEFAULT_EMBEDDING_MODEL = 'text-embedding-3-small'
export type EmbeddingProviderKind = 'openai' | 'azure-openai' | 'gemini'
export type TokenizerProviderId = 'openai' | 'google'
export interface EmbeddingModelInfo {
provider: EmbeddingProviderKind
/** Pricing/billing label — must match an entry in EMBEDDING_MODEL_PRICING when billed. */
pricingId: string
/** Provider id for `estimateTokenCount` so token counts match the embedding provider's tokenization. */
tokenizerProvider: TokenizerProviderId
}
export const SUPPORTED_EMBEDDING_MODELS: Partial<Record<string, EmbeddingModelInfo>> = {
'text-embedding-3-small': {
provider: 'openai',
pricingId: 'text-embedding-3-small',
tokenizerProvider: 'openai',
},
'text-embedding-3-large': {
provider: 'openai',
pricingId: 'text-embedding-3-large',
tokenizerProvider: 'openai',
},
'gemini-embedding-001': {
provider: 'gemini',
pricingId: 'gemini-embedding-001',
tokenizerProvider: 'google',
},
}
export function getEmbeddingModelInfo(model: string): EmbeddingModelInfo {
const info = SUPPORTED_EMBEDDING_MODELS[model]
if (!info) {
throw new Error(`Unsupported embedding model: ${model}`)
}
return info
}
+266 -112
View File
@@ -1,24 +1,30 @@
import { createLogger } from '@sim/logger'
import { getBYOKKey } from '@/lib/api-key/byok'
import { getRotatingApiKey } from '@/lib/core/config/api-keys'
import { env } from '@/lib/core/config/env'
import { isRetryableError, retryWithExponentialBackoff } from '@/lib/knowledge/documents/utils'
import { batchByTokenLimit } from '@/lib/tokenization'
import {
DEFAULT_EMBEDDING_MODEL,
EMBEDDING_DIMENSIONS,
getEmbeddingModelInfo,
SUPPORTED_EMBEDDING_MODELS,
type TokenizerProviderId,
} from '@/lib/knowledge/embedding-models'
import { batchByTokenLimit, estimateTokenCount } from '@/lib/tokenization'
const logger = createLogger('EmbeddingUtils')
const MAX_TOKENS_PER_REQUEST = 8000
const MAX_CONCURRENT_BATCHES = env.KB_CONFIG_CONCURRENCY_LIMIT || 50
const EMBEDDING_DIMENSIONS = 1536
const EMBEDDING_REQUEST_TIMEOUT_MS = 60_000
/**
* Check if the model supports custom dimensions.
* text-embedding-3-* models support the dimensions parameter.
* Checks for 'embedding-3' to handle Azure deployments with custom naming conventions.
*/
function supportsCustomDimensions(modelName: string): boolean {
const name = modelName.toLowerCase()
return name.includes('embedding-3') && !name.includes('ada')
}
export type { EmbeddingModelInfo } from '@/lib/knowledge/embedding-models'
export {
DEFAULT_EMBEDDING_MODEL,
EMBEDDING_DIMENSIONS,
getEmbeddingModelInfo,
SUPPORTED_EMBEDDING_MODELS,
} from '@/lib/knowledge/embedding-models'
export class EmbeddingAPIError extends Error {
public status: number
@@ -30,112 +36,245 @@ export class EmbeddingAPIError extends Error {
}
}
interface EmbeddingConfig {
useAzure: boolean
export type EmbeddingInputType = 'document' | 'query'
interface ProviderRequest {
apiUrl: string
headers: Record<string, string>
body: unknown
parse: (json: unknown) => number[][]
}
interface ResolvedProvider {
modelName: string
pricingId: string
isBYOK: boolean
/** Tokenizer used to estimate tokens when the API does not return a usage field. */
tokenizerProvider: TokenizerProviderId
/** Hard per-request item cap enforced by the provider (e.g. Gemini caps at 100). */
maxItemsPerRequest?: number
buildRequest: (inputs: string[], inputType: EmbeddingInputType) => ProviderRequest
}
interface EmbeddingResponseItem {
embedding: number[]
index: number
}
interface EmbeddingAPIResponse {
data: EmbeddingResponseItem[]
model: string
usage: {
prompt_tokens: number
total_tokens: number
}
}
async function getEmbeddingConfig(
embeddingModel = 'text-embedding-3-small',
workspaceId?: string | null
): Promise<EmbeddingConfig> {
const azureApiKey = env.AZURE_OPENAI_API_KEY
const azureEndpoint = env.AZURE_OPENAI_ENDPOINT
const azureApiVersion = env.AZURE_OPENAI_API_VERSION
const kbModelName = env.KB_OPENAI_MODEL_NAME || embeddingModel
const useAzure = !!(azureApiKey && azureEndpoint)
if (useAzure) {
return {
useAzure: true,
apiUrl: `${azureEndpoint}/openai/deployments/${kbModelName}/embeddings?api-version=${azureApiVersion}`,
headers: {
'api-key': azureApiKey!,
'Content-Type': 'application/json',
},
modelName: kbModelName,
isBYOK: false,
}
}
let openaiApiKey = env.OPENAI_API_KEY
let isBYOK = false
/** Gemini's `batchEmbedContents` rejects requests with more than 100 items. */
const GEMINI_MAX_ITEMS_PER_REQUEST = 100
async function resolveOpenAIKey(workspaceId?: string | null): Promise<{
apiKey: string
isBYOK: boolean
}> {
if (workspaceId) {
const byokResult = await getBYOKKey(workspaceId, 'openai')
if (byokResult) {
logger.info('Using workspace BYOK key for OpenAI embeddings')
openaiApiKey = byokResult.apiKey
isBYOK = true
return { apiKey: byokResult.apiKey, isBYOK: true }
}
}
if (!openaiApiKey) {
throw new Error(
'Either OPENAI_API_KEY or Azure OpenAI configuration (AZURE_OPENAI_API_KEY + AZURE_OPENAI_ENDPOINT) must be configured'
)
if (env.OPENAI_API_KEY) {
return { apiKey: env.OPENAI_API_KEY, isBYOK: false }
}
return {
useAzure: false,
apiUrl: 'https://api.openai.com/v1/embeddings',
headers: {
Authorization: `Bearer ${openaiApiKey}`,
'Content-Type': 'application/json',
},
modelName: embeddingModel,
isBYOK,
try {
return { apiKey: getRotatingApiKey('openai'), isBYOK: false }
} catch {
throw new Error('OPENAI_API_KEY is not configured')
}
}
const EMBEDDING_REQUEST_TIMEOUT_MS = 60_000
async function resolveGeminiKey(workspaceId?: string | null): Promise<{
apiKey: string
isBYOK: boolean
}> {
if (workspaceId) {
const byokResult = await getBYOKKey(workspaceId, 'google')
if (byokResult) {
logger.info('Using workspace BYOK key for Gemini embeddings')
return { apiKey: byokResult.apiKey, isBYOK: true }
}
}
if (env.GEMINI_API_KEY) {
return { apiKey: env.GEMINI_API_KEY, isBYOK: false }
}
try {
return { apiKey: getRotatingApiKey('gemini'), isBYOK: false }
} catch {
throw new Error(
'GEMINI_API_KEY (or GEMINI_API_KEY_1/2/3 for rotation) must be configured for Gemini embeddings'
)
}
}
function buildOpenAIProvider(modelName: string, apiKey: string): ResolvedProvider['buildRequest'] {
return (inputs) => ({
apiUrl: 'https://api.openai.com/v1/embeddings',
headers: {
Authorization: `Bearer ${apiKey}`,
'Content-Type': 'application/json',
},
body: {
input: inputs,
model: modelName,
encoding_format: 'float',
dimensions: EMBEDDING_DIMENSIONS,
},
parse: (json) => {
const data = json as { data: Array<{ embedding: number[] }> }
return data.data.map((item) => item.embedding)
},
})
}
function buildAzureOpenAIProvider(
deployment: string,
apiKey: string,
endpoint: string,
apiVersion: string
): ResolvedProvider['buildRequest'] {
return (inputs) => ({
apiUrl: `${endpoint}/openai/deployments/${deployment}/embeddings?api-version=${apiVersion}`,
headers: {
'api-key': apiKey,
'Content-Type': 'application/json',
},
body: {
input: inputs,
encoding_format: 'float',
dimensions: EMBEDDING_DIMENSIONS,
},
parse: (json) => {
const data = json as { data: Array<{ embedding: number[] }> }
return data.data.map((item) => item.embedding)
},
})
}
/**
* Gemini does NOT auto-normalize embeddings when `outputDimensionality` is set below the
* native 3072 dimension on `gemini-embedding-001`. Manually L2-normalize so cosine and
* inner-product similarity work correctly.
*/
function l2Normalize(vector: number[]): number[] {
let sumSquares = 0
for (const v of vector) sumSquares += v * v
const norm = Math.sqrt(sumSquares)
if (norm === 0) return vector
return vector.map((v) => v / norm)
}
function buildGeminiProvider(modelName: string, apiKey: string): ResolvedProvider['buildRequest'] {
return (inputs, inputType) => ({
apiUrl: `https://generativelanguage.googleapis.com/v1beta/models/${modelName}:batchEmbedContents`,
headers: {
'Content-Type': 'application/json',
'x-goog-api-key': apiKey,
},
body: {
requests: inputs.map((text) => ({
model: `models/${modelName}`,
content: { parts: [{ text }] },
taskType: inputType === 'query' ? 'RETRIEVAL_QUERY' : 'RETRIEVAL_DOCUMENT',
outputDimensionality: EMBEDDING_DIMENSIONS,
})),
},
parse: (json) => {
const data = json as { embeddings: Array<{ values: number[] }> }
return data.embeddings.map((item) => l2Normalize(item.values))
},
})
}
/**
* Returns the embedding model to use for new knowledge bases.
* Sourced from the `KB_EMBEDDING_MODEL` env var; falls back to the default if
* unset or set to an unsupported model.
*/
export function getConfiguredEmbeddingModel(): string {
const configured = env.KB_EMBEDDING_MODEL
if (configured && SUPPORTED_EMBEDDING_MODELS[configured]) {
return configured
}
if (configured) {
logger.warn(
`KB_EMBEDDING_MODEL="${configured}" is not a supported embedding model — falling back to ${DEFAULT_EMBEDDING_MODEL}`
)
}
return DEFAULT_EMBEDDING_MODEL
}
async function resolveProvider(
embeddingModel: string,
workspaceId?: string | null
): Promise<ResolvedProvider> {
const azureApiKey = env.AZURE_OPENAI_API_KEY
const azureEndpoint = env.AZURE_OPENAI_ENDPOINT
const azureApiVersion = env.AZURE_OPENAI_API_VERSION
const isOpenAIModel = SUPPORTED_EMBEDDING_MODELS[embeddingModel]?.provider === 'openai'
/**
* Azure deployment names default to the embedding model name when
* `KB_OPENAI_MODEL_NAME` is unset — this matches the pre-existing
* convention where deployments are named after the model they host.
*/
const azureDeploymentName = env.KB_OPENAI_MODEL_NAME || embeddingModel
const useAzure = Boolean(isOpenAIModel && azureApiKey && azureEndpoint && azureApiVersion)
const info = getEmbeddingModelInfo(embeddingModel)
if (useAzure) {
return {
modelName: azureDeploymentName,
pricingId: info.pricingId,
isBYOK: false,
tokenizerProvider: info.tokenizerProvider,
buildRequest: buildAzureOpenAIProvider(
azureDeploymentName,
azureApiKey!,
azureEndpoint!,
azureApiVersion!
),
}
}
if (info.provider === 'openai') {
const { apiKey, isBYOK } = await resolveOpenAIKey(workspaceId)
return {
modelName: embeddingModel,
pricingId: info.pricingId,
isBYOK,
tokenizerProvider: info.tokenizerProvider,
buildRequest: buildOpenAIProvider(embeddingModel, apiKey),
}
}
if (info.provider === 'gemini') {
const { apiKey, isBYOK } = await resolveGeminiKey(workspaceId)
return {
modelName: embeddingModel,
pricingId: info.pricingId,
isBYOK,
tokenizerProvider: info.tokenizerProvider,
maxItemsPerRequest: GEMINI_MAX_ITEMS_PER_REQUEST,
buildRequest: buildGeminiProvider(embeddingModel, apiKey),
}
}
throw new Error(`Unknown embedding provider for model ${embeddingModel}`)
}
async function callEmbeddingAPI(
inputs: string[],
config: EmbeddingConfig
provider: ResolvedProvider,
inputType: EmbeddingInputType
): Promise<{ embeddings: number[][]; totalTokens: number }> {
return retryWithExponentialBackoff(
async () => {
const useDimensions = supportsCustomDimensions(config.modelName)
const requestBody = config.useAzure
? {
input: inputs,
encoding_format: 'float',
...(useDimensions && { dimensions: EMBEDDING_DIMENSIONS }),
}
: {
input: inputs,
model: config.modelName,
encoding_format: 'float',
...(useDimensions && { dimensions: EMBEDDING_DIMENSIONS }),
}
const request = provider.buildRequest(inputs, inputType)
const controller = new AbortController()
const timeout = setTimeout(() => controller.abort(), EMBEDDING_REQUEST_TIMEOUT_MS)
const response = await fetch(config.apiUrl, {
const response = await fetch(request.apiUrl, {
method: 'POST',
headers: config.headers,
body: JSON.stringify(requestBody),
headers: request.headers,
body: JSON.stringify(request.body),
signal: controller.signal,
}).finally(() => clearTimeout(timeout))
@@ -147,11 +286,18 @@ async function callEmbeddingAPI(
)
}
const data: EmbeddingAPIResponse = await response.json()
return {
embeddings: data.data.map((item) => item.embedding),
totalTokens: data.usage.total_tokens,
}
const json = await response.json()
const embeddings = request.parse(json)
const usage = (json as { usage?: { total_tokens?: number } }).usage
const totalTokens =
usage?.total_tokens ??
// Gemini does not return usage.total_tokens — estimate with the provider's tokenizer
inputs.reduce(
(sum, text) => sum + estimateTokenCount(text, provider.tokenizerProvider).count,
0
)
return { embeddings, totalTokens }
},
{
maxRetries: 3,
@@ -167,9 +313,15 @@ async function callEmbeddingAPI(
)
}
/**
* Process batches with controlled concurrency
*/
function splitByItemLimit<T>(items: T[], limit: number): T[][] {
if (items.length <= limit) return [items]
const result: T[][] = []
for (let i = 0; i < items.length; i += limit) {
result.push(items.slice(i, i + limit))
}
return result
}
async function processWithConcurrency<T, R>(
items: T[],
concurrency: number,
@@ -194,28 +346,31 @@ export interface GenerateEmbeddingsResult {
totalTokens: number
isBYOK: boolean
modelName: string
/** Pricing identifier for use with calculateCost / EMBEDDING_MODEL_PRICING. */
pricingId: string
}
/**
* Generate embeddings for multiple texts with token-aware batching and parallel processing.
* Returns embeddings alongside actual token count, model name, and whether a workspace BYOK key
* was used (vs. the platform's shared key) — enabling callers to make correct billing decisions.
*/
export async function generateEmbeddings(
texts: string[],
embeddingModel = 'text-embedding-3-small',
embeddingModel: string = DEFAULT_EMBEDDING_MODEL,
workspaceId?: string | null
): Promise<GenerateEmbeddingsResult> {
const config = await getEmbeddingConfig(embeddingModel, workspaceId)
const provider = await resolveProvider(embeddingModel, workspaceId)
const batches = batchByTokenLimit(texts, MAX_TOKENS_PER_REQUEST, embeddingModel)
const tokenBatches = batchByTokenLimit(texts, MAX_TOKENS_PER_REQUEST, embeddingModel)
const batches = provider.maxItemsPerRequest
? tokenBatches.flatMap((batch) => splitByItemLimit(batch, provider.maxItemsPerRequest!))
: tokenBatches
const batchResults = await processWithConcurrency(
batches,
MAX_CONCURRENT_BATCHES,
async (batch, i) => {
try {
return await callEmbeddingAPI(batch, config)
return await callEmbeddingAPI(batch, provider, 'document')
} catch (error) {
logger.error(`Failed to generate embeddings for batch ${i + 1}/${batches.length}:`, error)
throw error
@@ -235,25 +390,24 @@ export async function generateEmbeddings(
return {
embeddings: allEmbeddings,
totalTokens,
isBYOK: config.isBYOK,
modelName: config.modelName,
isBYOK: provider.isBYOK,
modelName: provider.modelName,
pricingId: provider.pricingId,
}
}
/**
* Generate embedding for a single search query
* Generate embedding for a single search query.
*/
export async function generateSearchEmbedding(
query: string,
embeddingModel = 'text-embedding-3-small',
embeddingModel: string = DEFAULT_EMBEDDING_MODEL,
workspaceId?: string | null
): Promise<number[]> {
const config = await getEmbeddingConfig(embeddingModel, workspaceId)
const provider = await resolveProvider(embeddingModel, workspaceId)
logger.info(
`Using ${config.useAzure ? 'Azure OpenAI' : 'OpenAI'} for search embedding generation`
)
logger.info(`Using ${provider.modelName} for search embedding generation`)
const { embeddings } = await callEmbeddingAPI([query], config)
const { embeddings } = await callEmbeddingAPI([query], provider, 'query')
return embeddings[0]
}
+18
View File
@@ -0,0 +1,18 @@
/**
* Client-safe registry of Cohere rerank models supported by the platform.
* Kept free of server imports so it can be imported into UI / block code.
*/
/** Cohere rerank model identifiers we accept. Must match Cohere's model ids exactly. */
export const SUPPORTED_RERANKER_MODELS = [
'rerank-v4.0-pro',
'rerank-v4.0-fast',
'rerank-v3.5',
] as const
export type RerankerModelId = (typeof SUPPORTED_RERANKER_MODELS)[number]
export const DEFAULT_RERANKER_MODEL: RerankerModelId = 'rerank-v4.0-fast'
export function isSupportedRerankerModel(model: string): model is RerankerModelId {
return (SUPPORTED_RERANKER_MODELS as readonly string[]).includes(model)
}
+163
View File
@@ -0,0 +1,163 @@
import { createLogger } from '@sim/logger'
import { getBYOKKey } from '@/lib/api-key/byok'
import { getRotatingApiKey } from '@/lib/core/config/api-keys'
import { env } from '@/lib/core/config/env'
import { isRetryableError, retryWithExponentialBackoff } from '@/lib/knowledge/documents/utils'
import {
DEFAULT_RERANKER_MODEL,
isSupportedRerankerModel,
type RerankerModelId,
SUPPORTED_RERANKER_MODELS,
} from '@/lib/knowledge/reranker-models'
export {
DEFAULT_RERANKER_MODEL,
isSupportedRerankerModel,
type RerankerModelId,
SUPPORTED_RERANKER_MODELS,
}
const logger = createLogger('Reranker')
const RERANK_REQUEST_TIMEOUT_MS = 30_000
/**
* Cohere bills per "search unit" = one query with up to 100 documents.
* We cap at 100 so each rerank call costs exactly 1 unit and matches
* `RERANK_MODEL_PRICING` in `providers/models.ts`. The search route also
* caps `candidateTopK` at 100, so this is a defensive ceiling.
*/
const MAX_DOCUMENTS_PER_RERANK = 100
export interface RerankItem {
/** Stable identifier so callers can correlate ranked results back to source rows. */
id: string
text: string
}
export interface RerankedResult<T extends RerankItem> {
item: T
relevanceScore: number
}
export interface RerankResponse<T extends RerankItem> {
results: RerankedResult<T>[]
/** True when a workspace-supplied (BYOK) Cohere key was used. Callers should skip platform billing in that case. */
isBYOK: boolean
}
class RerankAPIError extends Error {
public status: number
constructor(message: string, status: number) {
super(message)
this.name = 'RerankAPIError'
this.status = status
}
}
async function resolveCohereKey(
workspaceId?: string | null
): Promise<{ apiKey: string; isBYOK: boolean }> {
if (workspaceId) {
const byokResult = await getBYOKKey(workspaceId, 'cohere')
if (byokResult) {
logger.info('Using workspace BYOK key for Cohere reranker')
return { apiKey: byokResult.apiKey, isBYOK: true }
}
}
if (env.COHERE_API_KEY) {
return { apiKey: env.COHERE_API_KEY, isBYOK: false }
}
try {
return { apiKey: getRotatingApiKey('cohere'), isBYOK: false }
} catch {
throw new Error(
'No Cohere API key configured. Set COHERE_API_KEY_1/2/3 (rotation) or COHERE_API_KEY.'
)
}
}
interface CohereRerankResponse {
results: Array<{ index: number; relevance_score: number }>
}
/**
* Rerank documents against a query using Cohere's `/v2/rerank` endpoint.
* Returns the items in descending order of relevance, capped at `topN`.
*/
export async function rerank<T extends RerankItem>(
query: string,
items: T[],
options: {
model: string
topN?: number
workspaceId?: string | null
}
): Promise<RerankResponse<T>> {
if (items.length === 0) return { results: [], isBYOK: false }
if (!isSupportedRerankerModel(options.model)) {
throw new Error(`Unsupported reranker model: ${options.model}`)
}
const { apiKey, isBYOK } = await resolveCohereKey(options.workspaceId)
const cappedItems =
items.length > MAX_DOCUMENTS_PER_RERANK ? items.slice(0, MAX_DOCUMENTS_PER_RERANK) : items
if (items.length > MAX_DOCUMENTS_PER_RERANK) {
logger.warn(`Rerank input capped from ${items.length} to ${MAX_DOCUMENTS_PER_RERANK} documents`)
}
const documents = cappedItems.map((it) => it.text)
const response = await retryWithExponentialBackoff(
async () => {
const controller = new AbortController()
const timeout = setTimeout(() => controller.abort(), RERANK_REQUEST_TIMEOUT_MS)
const res = await fetch('https://api.cohere.com/v2/rerank', {
method: 'POST',
headers: {
Authorization: `Bearer ${apiKey}`,
'Content-Type': 'application/json',
},
body: JSON.stringify({
model: options.model,
query,
documents,
top_n: options.topN ?? cappedItems.length,
}),
signal: controller.signal,
}).finally(() => clearTimeout(timeout))
if (!res.ok) {
const errorText = await res.text()
throw new RerankAPIError(
`Cohere rerank failed: ${res.status} ${res.statusText} - ${errorText}`,
res.status
)
}
return (await res.json()) as CohereRerankResponse
},
{
maxRetries: 3,
initialDelayMs: 500,
maxDelayMs: 5000,
retryCondition: (error: unknown) => {
if (error instanceof RerankAPIError) {
return error.status === 429 || error.status >= 500
}
return isRetryableError(error)
},
}
)
return {
results: response.results
.filter((r) => r.index >= 0 && r.index < cappedItems.length)
.map((r) => ({
item: cappedItems[r.index],
relevanceScore: r.relevance_score,
})),
isBYOK,
}
}
-2
View File
@@ -250,8 +250,6 @@ export async function updateKnowledgeBase(
if (updates.workspaceId !== undefined) updateData.workspaceId = updates.workspaceId
if (updates.chunkingConfig !== undefined) {
updateData.chunkingConfig = updates.chunkingConfig
updateData.embeddingModel = 'text-embedding-3-small'
updateData.embeddingDimension = 1536
}
if (updates.name !== undefined) {
+1 -1
View File
@@ -34,7 +34,7 @@ export interface CreateKnowledgeBaseData {
name: string
description?: string
workspaceId: string
embeddingModel: 'text-embedding-3-small'
embeddingModel: string
embeddingDimension: 1536
chunkingConfig: ChunkingConfig
userId: string
+21
View File
@@ -3023,12 +3023,33 @@ export const EMBEDDING_MODEL_PRICING: Record<string, ModelPricing> = {
output: 0.0,
updatedAt: '2026-04-01',
},
'gemini-embedding-001': {
input: 0.15, // $0.15 per 1M tokens
output: 0.0,
updatedAt: '2026-04-29',
},
}
export function getEmbeddingModelPricing(modelId: string): ModelPricing | null {
return EMBEDDING_MODEL_PRICING[modelId] || null
}
/**
* Cohere rerank pricing in USD per single search unit (one query × ≤100 docs).
* Sim caps every rerank request to ≤100 documents, so each call = 1 unit.
*/
export const RERANK_MODEL_PRICING: Record<string, { perSearchUnit: number; updatedAt: string }> = {
'rerank-v4.0-pro': { perSearchUnit: 0.0025, updatedAt: '2026-04-29' },
'rerank-v4.0-fast': { perSearchUnit: 0.002, updatedAt: '2026-04-29' },
'rerank-v3.5': { perSearchUnit: 0.002, updatedAt: '2026-04-29' },
}
export function getRerankModelPricing(
modelId: string
): { perSearchUnit: number; updatedAt: string } | null {
return RERANK_MODEL_PRICING[modelId] || null
}
export function getModelsWithReasoningEffort(): string[] {
const models: string[] = []
for (const provider of Object.values(PROVIDER_DEFINITIONS)) {
+38 -2
View File
@@ -1,3 +1,4 @@
import { DEFAULT_RERANKER_MODEL, SUPPORTED_RERANKER_MODELS } from '@/lib/knowledge/reranker-models'
import type { KnowledgeSearchResponse } from '@/tools/knowledge/types'
import { enrichKBTagFiltersSchema } from '@/tools/schema-enrichers'
import { parseTagFilters } from '@/tools/shared/tags'
@@ -41,6 +42,18 @@ export const knowledgeSearchTool: ToolConfig<any, KnowledgeSearchResponse> = {
},
},
},
rerankerEnabled: {
type: 'boolean',
required: false,
visibility: 'user-only',
description: 'Whether to apply Cohere reranking to vector search results',
},
rerankerModel: {
type: 'string',
required: false,
visibility: 'user-only',
description: `Cohere rerank model to use (one of: ${SUPPORTED_RERANKER_MODELS.join(', ')})`,
},
},
schemaEnrichment: {
@@ -65,11 +78,18 @@ export const knowledgeSearchTool: ToolConfig<any, KnowledgeSearchResponse> = {
// Parse tag filters from various formats (array, JSON string)
const structuredFilters = parseTagFilters(params.tagFilters)
const rerankerEnabled = params.rerankerEnabled === true || params.rerankerEnabled === 'true'
const rerankerModel =
typeof params.rerankerModel === 'string' && params.rerankerModel.length > 0
? params.rerankerModel
: DEFAULT_RERANKER_MODEL
const requestBody = {
knowledgeBaseIds,
query: params.query,
topK: params.topK ? Math.max(1, Math.min(100, Number(params.topK))) : 10,
...(structuredFilters.length > 0 && { tagFilters: structuredFilters }),
...(rerankerEnabled && { rerankerEnabled: true, rerankerModel }),
...(workflowId && { workflowId }),
}
@@ -83,9 +103,25 @@ export const knowledgeSearchTool: ToolConfig<any, KnowledgeSearchResponse> = {
// Restructure cost: extract tokens/model to top level for logging
let costFields: Record<string, unknown> = {}
if (data.cost && typeof data.cost === 'object') {
const { tokens, model, input, output: outputCost, total } = data.cost
const {
tokens,
model,
input,
output: outputCost,
total,
rerankerCost,
rerankerModel,
rerankerSearchUnits,
} = data.cost
costFields = {
cost: { input, output: outputCost, total },
cost: {
input,
output: outputCost,
total,
...(typeof rerankerCost === 'number' && { rerankerCost }),
...(typeof rerankerModel === 'string' && { rerankerModel }),
...(typeof rerankerSearchUnits === 'number' && { rerankerSearchUnits }),
},
...(tokens && { tokens }),
...(model && { model }),
}
+10 -11
View File
@@ -40,6 +40,7 @@ export interface KnowledgeSearchResult {
chunkIndex: number
metadata: Record<string, any>
similarity: number
rerankerScore?: number
}
export interface KnowledgeSearchResponse {
@@ -52,18 +53,16 @@ export interface KnowledgeSearchResponse {
input: number
output: number
total: number
tokens: {
prompt: number
completion: number
total: number
}
model: string
pricing: {
input: number
output: number
updatedAt: string
}
rerankerCost?: number
rerankerModel?: string
rerankerSearchUnits?: number
}
tokens?: {
prompt: number
completion: number
total: number
}
model?: string
}
error?: string
}
+1
View File
@@ -17,6 +17,7 @@ export type BYOKProviderId =
| 'linkup'
| 'brandfetch'
| 'parallel_ai'
| 'cohere'
export type HttpMethod = 'GET' | 'POST' | 'PUT' | 'DELETE' | 'PATCH' | 'HEAD'