mirror of
https://github.com/simstudioai/sim.git
synced 2026-09-24 15:45:35 +08:00
feat(embeddings): add OpenRouter support (#6396)
* feat(knowledge): add OpenRouter embedding fallback * fix(knowledge): preserve successful embedding batches * feat(embeddings): add OpenRouter provider * fix(knowledge): bill only platform embedding tokens * test(embeddings): include OpenRouter provider * feat(embeddings): load OpenRouter model catalog * fix(embeddings): preserve legacy provider default * fix(embeddings): batch OpenRouter requests * fix(embeddings): reset stale OpenRouter model
This commit is contained in:
@@ -25,7 +25,7 @@ Sim's knowledge bases embed separately, at a fixed vector width and from a small
|
||||
|
||||
## Usage Instructions
|
||||
|
||||
Turn text into embedding vectors for semantic search, clustering, and similarity. Supports OpenAI, Google Gemini, Cohere, and Mistral embedding models.
|
||||
Turn text into embedding vectors for semantic search, clustering, and similarity. Supports OpenAI, OpenRouter, Google Gemini, Cohere, and Mistral embedding models.
|
||||
|
||||
|
||||
|
||||
@@ -55,6 +55,30 @@ Generate embeddings from text using OpenAI's embedding models
|
||||
| `dimensions` | number | Dimensionality of each vector |
|
||||
| `usage` | json | Token usage |
|
||||
|
||||
### OpenRouter Embeddings
|
||||
|
||||
Generate embeddings through OpenRouter
|
||||
|
||||
#### Input
|
||||
|
||||
| Parameter | Type | Required | Description |
|
||||
| --------- | ---- | -------- | ----------- |
|
||||
| `input` | string | Yes | Text to embed, or an array of texts to embed in one call |
|
||||
| `model` | string | No | Embedding model to use |
|
||||
| `taskType` | string | No | What the embedding is for, when the model supports task conditioning: document, query, similarity, classification, or clustering |
|
||||
| `dimensions` | number | No | Output dimensions, when the model supports truncation. Defaults to native. |
|
||||
| `apiKey` | string | Yes | API key for the selected embedding provider |
|
||||
|
||||
#### Output
|
||||
|
||||
| Parameter | Type | Description |
|
||||
| --------- | ---- | ----------- |
|
||||
| `embeddings` | json | Generated embeddings |
|
||||
| `model` | string | Model used |
|
||||
| `provider` | string | Provider used |
|
||||
| `dimensions` | number | Dimensionality of each vector |
|
||||
| `usage` | json | Token usage |
|
||||
|
||||
### Gemini Embeddings
|
||||
|
||||
Generate embeddings from text using Google's Gemini embedding models
|
||||
|
||||
@@ -92,6 +92,7 @@ CRON_SECRET=your_cron_secret # Use `openssl rand -hex 32` to generate. Authentic
|
||||
# VLLM_API_KEY= # Optional bearer token if your vLLM instance requires auth
|
||||
# LITELLM_BASE_URL=http://localhost:4000 # Base URL for your LiteLLM proxy (OpenAI-compatible)
|
||||
# LITELLM_API_KEY= # Optional bearer token if your LiteLLM proxy requires auth
|
||||
# OPENROUTER_API_KEY= # Optional self-hosted fallback for OpenAI knowledge-base embeddings
|
||||
# NEXT_PUBLIC_FORCE_HOSTED=true # Dev only: treat this instance as hosted Sim (sim-auto pool, platform keys); ignored in production builds
|
||||
# FIREWORKS_API_KEY= # Optional Fireworks AI API key for model listing and inference
|
||||
# FIREWORKS_API_KEY_1= # Optional Fireworks API key for rotation (hosted deployments)
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
/**
|
||||
* @vitest-environment node
|
||||
*/
|
||||
import { createMockRequest } from '@sim/testing'
|
||||
import { afterAll, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
|
||||
const { mockFetch, mockFilterBlacklistedModels, mockIsProviderBlacklisted } = vi.hoisted(() => ({
|
||||
mockFetch: vi.fn(),
|
||||
mockFilterBlacklistedModels: vi.fn(),
|
||||
mockIsProviderBlacklisted: vi.fn(),
|
||||
}))
|
||||
|
||||
vi.mock('@/providers/utils', () => ({
|
||||
filterBlacklistedModels: mockFilterBlacklistedModels,
|
||||
isProviderBlacklisted: mockIsProviderBlacklisted,
|
||||
}))
|
||||
|
||||
import { GET } from '@/app/api/providers/openrouter/embeddings/models/route'
|
||||
|
||||
const request = () => createMockRequest('GET')
|
||||
|
||||
describe('GET /api/providers/openrouter/embeddings/models', () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
vi.stubGlobal('fetch', mockFetch)
|
||||
mockIsProviderBlacklisted.mockReturnValue(false)
|
||||
mockFilterBlacklistedModels.mockImplementation((models: string[]) => models)
|
||||
})
|
||||
|
||||
afterAll(() => {
|
||||
vi.unstubAllGlobals()
|
||||
})
|
||||
|
||||
it('returns every unique embedding model with the OpenRouter prefix', async () => {
|
||||
mockFetch.mockResolvedValue({
|
||||
ok: true,
|
||||
status: 200,
|
||||
statusText: 'OK',
|
||||
json: async () => ({
|
||||
data: [
|
||||
{ id: 'qwen/qwen3-embedding-8b', context_length: 32768 },
|
||||
{ id: 'openai/text-embedding-3-small', context_length: 8192 },
|
||||
{ id: 'qwen/qwen3-embedding-8b', context_length: 32768 },
|
||||
],
|
||||
}),
|
||||
})
|
||||
|
||||
const response = await GET(request(), undefined as never)
|
||||
|
||||
expect(response.status).toBe(200)
|
||||
await expect(response.json()).resolves.toEqual({
|
||||
models: ['openrouter/qwen/qwen3-embedding-8b', 'openrouter/openai/text-embedding-3-small'],
|
||||
})
|
||||
expect(mockFetch).toHaveBeenCalledWith(
|
||||
'https://openrouter.ai/api/v1/embeddings/models',
|
||||
expect.objectContaining({ next: { revalidate: 300 } })
|
||||
)
|
||||
})
|
||||
|
||||
it('does not fetch when OpenRouter is blacklisted', async () => {
|
||||
mockIsProviderBlacklisted.mockReturnValue(true)
|
||||
|
||||
const response = await GET(request(), undefined as never)
|
||||
|
||||
expect(response.status).toBe(200)
|
||||
await expect(response.json()).resolves.toEqual({ models: [] })
|
||||
expect(mockFetch).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('fails fast when OpenRouter rejects the model-list request', async () => {
|
||||
mockFetch.mockResolvedValue({
|
||||
ok: false,
|
||||
status: 503,
|
||||
statusText: 'Service Unavailable',
|
||||
})
|
||||
|
||||
const response = await GET(request(), undefined as never)
|
||||
|
||||
expect(response.status).toBe(500)
|
||||
expect(mockFilterBlacklistedModels).not.toHaveBeenCalled()
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,24 @@
|
||||
import { createLogger } from '@sim/logger'
|
||||
import { type NextRequest, NextResponse } from 'next/server'
|
||||
import { providerModelsResponseSchema } from '@/lib/api/contracts/providers'
|
||||
import { withRouteHandler } from '@/lib/core/utils/with-route-handler'
|
||||
import { fetchOpenRouterEmbeddingModelCatalog } from '@/lib/embeddings/openrouter-model-catalog.server'
|
||||
import { filterBlacklistedModels, isProviderBlacklisted } from '@/providers/utils'
|
||||
|
||||
const logger = createLogger('OpenRouterEmbeddingModelsAPI')
|
||||
|
||||
export const GET = withRouteHandler(async (_request: NextRequest) => {
|
||||
if (isProviderBlacklisted('openrouter')) {
|
||||
logger.info('OpenRouter provider is blacklisted, returning empty embedding models')
|
||||
return NextResponse.json({ models: [] })
|
||||
}
|
||||
|
||||
const uniqueModels = (await fetchOpenRouterEmbeddingModelCatalog()).map((model) => model.id)
|
||||
const models = filterBlacklistedModels(uniqueModels)
|
||||
|
||||
logger.info('Successfully fetched OpenRouter embedding models', {
|
||||
count: models.length,
|
||||
filtered: uniqueModels.length - models.length,
|
||||
})
|
||||
return NextResponse.json(providerModelsResponseSchema.parse({ models }))
|
||||
})
|
||||
@@ -4,20 +4,37 @@
|
||||
import { createMockRequest, hybridAuthMockFns } from '@sim/testing'
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
|
||||
const { mockEmbed } = vi.hoisted(() => ({
|
||||
mockEmbed: vi.fn(),
|
||||
const { mockEmbed, mockEmbedOpenRouter, mockGetOpenRouterEmbeddingModelMetadata } = vi.hoisted(
|
||||
() => ({
|
||||
mockEmbed: vi.fn(),
|
||||
mockEmbedOpenRouter: vi.fn(),
|
||||
mockGetOpenRouterEmbeddingModelMetadata: vi.fn(),
|
||||
})
|
||||
)
|
||||
|
||||
vi.mock('@/lib/embeddings/openrouter-model-catalog.server', () => ({
|
||||
getOpenRouterEmbeddingModelMetadata: mockGetOpenRouterEmbeddingModelMetadata,
|
||||
OpenRouterEmbeddingModelNotFoundError: class OpenRouterEmbeddingModelNotFoundError extends Error {
|
||||
constructor(model: string) {
|
||||
super(`Unsupported OpenRouter embedding model: ${model}`)
|
||||
this.name = 'OpenRouterEmbeddingModelNotFoundError'
|
||||
}
|
||||
},
|
||||
}))
|
||||
|
||||
vi.mock('@/lib/embeddings', async () => {
|
||||
const catalog = await import('@/lib/embeddings/catalog')
|
||||
return {
|
||||
embed: mockEmbed,
|
||||
embedOpenRouter: mockEmbedOpenRouter,
|
||||
DEFAULT_OPENROUTER_EMBEDDING_MODEL: 'openrouter/openai/text-embedding-3-small',
|
||||
findEmbeddingModelInfo: catalog.findEmbeddingModelInfo,
|
||||
getModelsForProvider: catalog.getModelsForProvider,
|
||||
resolveDimensions: catalog.resolveDimensions,
|
||||
}
|
||||
})
|
||||
|
||||
import { OpenRouterEmbeddingModelNotFoundError } from '@/lib/embeddings/openrouter-model-catalog.server'
|
||||
import { POST } from '@/app/api/tools/embeddings/route'
|
||||
|
||||
const baseBody = {
|
||||
@@ -34,6 +51,10 @@ function post(body: Record<string, unknown>) {
|
||||
describe('POST /api/tools/embeddings', () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
mockGetOpenRouterEmbeddingModelMetadata.mockResolvedValue({
|
||||
id: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
maxInputTokens: 32768,
|
||||
})
|
||||
hybridAuthMockFns.mockCheckInternalAuth.mockResolvedValue({
|
||||
success: true,
|
||||
userId: 'user-1',
|
||||
@@ -47,6 +68,15 @@ describe('POST /api/tools/embeddings', () => {
|
||||
pricingId: 'text-embedding-3-small',
|
||||
dimensions: 1536,
|
||||
})
|
||||
mockEmbedOpenRouter.mockResolvedValue({
|
||||
embeddings: [[0.1, 0.2]],
|
||||
totalTokens: 3,
|
||||
billableTokens: 0,
|
||||
isBYOK: true,
|
||||
modelName: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
pricingId: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
dimensions: 2,
|
||||
})
|
||||
})
|
||||
|
||||
it('rejects an unauthenticated caller', async () => {
|
||||
@@ -117,6 +147,82 @@ describe('POST /api/tools/embeddings', () => {
|
||||
)
|
||||
})
|
||||
|
||||
it('routes OpenRouter through its transport with an explicit key', async () => {
|
||||
const response = await post({
|
||||
provider: 'openrouter',
|
||||
model: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
input: 'hello world',
|
||||
apiKey: 'or-test',
|
||||
})
|
||||
|
||||
expect(response.status).toBe(200)
|
||||
expect(mockEmbedOpenRouter).toHaveBeenCalledWith(
|
||||
['hello world'],
|
||||
expect.objectContaining({
|
||||
apiKey: 'or-test',
|
||||
maxInputTokens: 32768,
|
||||
model: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
})
|
||||
)
|
||||
expect(mockEmbed).not.toHaveBeenCalled()
|
||||
expect((await response.json()).provider).toBe('openrouter')
|
||||
})
|
||||
|
||||
it('rejects OpenRouter without an explicit key', async () => {
|
||||
const response = await post({
|
||||
provider: 'openrouter',
|
||||
model: 'openrouter/openai/text-embedding-3-small',
|
||||
input: 'hello world',
|
||||
})
|
||||
|
||||
expect(response.status).toBe(400)
|
||||
expect((await response.json()).error).toContain('apiKey')
|
||||
expect(mockEmbed).not.toHaveBeenCalled()
|
||||
expect(mockEmbedOpenRouter).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('rejects an invalid OpenRouter model id', async () => {
|
||||
const response = await post({
|
||||
provider: 'openrouter',
|
||||
model: 'openrouter/not-qualified',
|
||||
input: 'hello world',
|
||||
apiKey: 'or-test',
|
||||
})
|
||||
|
||||
expect(response.status).toBe(400)
|
||||
expect((await response.json()).error).toContain('Invalid OpenRouter embedding model')
|
||||
expect(mockEmbedOpenRouter).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('rejects a qualified model that is absent from OpenRouter', async () => {
|
||||
mockGetOpenRouterEmbeddingModelMetadata.mockRejectedValue(
|
||||
new OpenRouterEmbeddingModelNotFoundError('openrouter/example/missing')
|
||||
)
|
||||
|
||||
const response = await post({
|
||||
provider: 'openrouter',
|
||||
model: 'openrouter/example/missing',
|
||||
input: 'hello world',
|
||||
apiKey: 'or-test',
|
||||
})
|
||||
|
||||
expect(response.status).toBe(400)
|
||||
expect((await response.json()).error).toContain('Unsupported OpenRouter embedding model')
|
||||
expect(mockEmbedOpenRouter).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('keeps API keys required for non-OpenRouter providers', async () => {
|
||||
const response = await post({
|
||||
provider: 'openai',
|
||||
model: 'text-embedding-3-small',
|
||||
input: 'hello world',
|
||||
})
|
||||
|
||||
expect(response.status).toBe(400)
|
||||
expect((await response.json()).error).toContain('apiKey')
|
||||
expect(mockEmbed).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('surfaces a provider failure as 502', async () => {
|
||||
mockEmbed.mockRejectedValue(new Error('Embedding API failed: 429 Too Many Requests'))
|
||||
const response = await post(baseBody)
|
||||
|
||||
@@ -11,10 +11,18 @@ import { checkInternalAuth } from '@/lib/auth/hybrid'
|
||||
import { withRouteHandler } from '@/lib/core/utils/with-route-handler'
|
||||
import {
|
||||
DEFAULT_MODEL_BY_PROVIDER,
|
||||
DEFAULT_OPENROUTER_EMBEDDING_MODEL,
|
||||
embed,
|
||||
embedOpenRouter,
|
||||
findEmbeddingModelInfo,
|
||||
resolveDimensions,
|
||||
} from '@/lib/embeddings'
|
||||
import {
|
||||
getOpenRouterEmbeddingModelMetadata,
|
||||
type OpenRouterEmbeddingModelMetadata,
|
||||
OpenRouterEmbeddingModelNotFoundError,
|
||||
} from '@/lib/embeddings/openrouter-model-catalog.server'
|
||||
import { normalizeOpenRouterEmbeddingModelId } from '@/lib/embeddings/openrouter-models'
|
||||
|
||||
const logger = createLogger('EmbeddingsToolAPI')
|
||||
|
||||
@@ -63,7 +71,6 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
|
||||
if (!parsed.success) return parsed.response
|
||||
|
||||
const { provider, apiKey, model, input, taskType, dimensions } = parsed.data.body
|
||||
|
||||
const texts = normalizeInput(input)
|
||||
|
||||
/**
|
||||
@@ -109,53 +116,102 @@ export const POST = withRouteHandler(async (request: NextRequest) => {
|
||||
)
|
||||
}
|
||||
|
||||
const resolvedModel = model || DEFAULT_MODEL_BY_PROVIDER[provider]
|
||||
const info = findEmbeddingModelInfo(resolvedModel)
|
||||
if (!info) {
|
||||
return NextResponse.json(
|
||||
{ success: false, error: `Unsupported embedding model: ${resolvedModel}` },
|
||||
{ status: 400 }
|
||||
)
|
||||
}
|
||||
if (info.provider !== provider) {
|
||||
return NextResponse.json(
|
||||
{
|
||||
success: false,
|
||||
error: `Model ${resolvedModel} belongs to ${info.provider}, not ${provider}`,
|
||||
},
|
||||
{ status: 400 }
|
||||
)
|
||||
let resolvedModel: string
|
||||
let openRouterModelMetadata: OpenRouterEmbeddingModelMetadata | undefined
|
||||
if (provider === 'openrouter') {
|
||||
try {
|
||||
resolvedModel = normalizeOpenRouterEmbeddingModelId(
|
||||
model || DEFAULT_OPENROUTER_EMBEDDING_MODEL
|
||||
)
|
||||
} catch (error) {
|
||||
return NextResponse.json(
|
||||
{ success: false, error: getErrorMessage(error, 'Invalid OpenRouter embedding model') },
|
||||
{ status: 400 }
|
||||
)
|
||||
}
|
||||
try {
|
||||
openRouterModelMetadata = await getOpenRouterEmbeddingModelMetadata(resolvedModel)
|
||||
} catch (error) {
|
||||
const modelError = error instanceof OpenRouterEmbeddingModelNotFoundError
|
||||
return NextResponse.json(
|
||||
{
|
||||
success: false,
|
||||
error: getErrorMessage(
|
||||
error,
|
||||
modelError
|
||||
? 'Unsupported OpenRouter embedding model'
|
||||
: 'Failed to load OpenRouter embedding model metadata'
|
||||
),
|
||||
},
|
||||
{ status: modelError ? 400 : 502 }
|
||||
)
|
||||
}
|
||||
} else {
|
||||
resolvedModel = model || DEFAULT_MODEL_BY_PROVIDER[provider]
|
||||
}
|
||||
|
||||
/**
|
||||
* Resolved here as well as inside `embed()` so an unsupported `dimensions`
|
||||
* is reported as the client error it is. The block's dropdown constrains the
|
||||
* field, but a reference expression can put any value on the wire.
|
||||
*/
|
||||
try {
|
||||
resolveDimensions(info, dimensions)
|
||||
} catch (error) {
|
||||
return NextResponse.json(
|
||||
{ success: false, error: getErrorMessage(error, 'Invalid dimensions') },
|
||||
{ status: 400 }
|
||||
)
|
||||
if (provider !== 'openrouter') {
|
||||
const info = findEmbeddingModelInfo(resolvedModel)
|
||||
if (!info) {
|
||||
return NextResponse.json(
|
||||
{ success: false, error: `Unsupported embedding model: ${resolvedModel}` },
|
||||
{ status: 400 }
|
||||
)
|
||||
}
|
||||
if (info.provider !== provider) {
|
||||
return NextResponse.json(
|
||||
{
|
||||
success: false,
|
||||
error: `Model ${resolvedModel} belongs to ${info.provider}, not ${provider}`,
|
||||
},
|
||||
{ status: 400 }
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* Resolved here as well as inside `embed()` so an unsupported `dimensions`
|
||||
* is reported as the client error it is. The block's dropdown constrains the
|
||||
* field, but a reference expression can put any value on the wire.
|
||||
*/
|
||||
try {
|
||||
resolveDimensions(info, dimensions)
|
||||
} catch (error) {
|
||||
return NextResponse.json(
|
||||
{ success: false, error: getErrorMessage(error, 'Invalid dimensions') },
|
||||
{ status: 400 }
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
logger.info(`Embedding ${texts.length} input(s) with ${provider}/${resolvedModel}`)
|
||||
|
||||
try {
|
||||
const result = await embed(texts, {
|
||||
model: resolvedModel,
|
||||
taskType,
|
||||
dimensions,
|
||||
apiKey,
|
||||
/**
|
||||
* Callers reach this route through a tool whose `request.modelInput`
|
||||
* already projected `input` at the HTTP hop, so projecting again here
|
||||
* would run the substitution over already-projected content.
|
||||
*/
|
||||
projectInputs: null,
|
||||
})
|
||||
let result
|
||||
if (provider === 'openrouter') {
|
||||
if (!openRouterModelMetadata) {
|
||||
throw new Error('OpenRouter embedding model metadata was not resolved')
|
||||
}
|
||||
result = await embedOpenRouter(texts, {
|
||||
model: resolvedModel,
|
||||
dimensions,
|
||||
apiKey,
|
||||
maxInputTokens: openRouterModelMetadata.maxInputTokens,
|
||||
projectInputs: null,
|
||||
})
|
||||
} else {
|
||||
result = await embed(texts, {
|
||||
model: resolvedModel,
|
||||
taskType,
|
||||
dimensions,
|
||||
apiKey,
|
||||
/**
|
||||
* Callers reach this route through a tool whose `request.modelInput`
|
||||
* already projected `input` at the HTTP hop, so projecting again here
|
||||
* would run the substitution over already-projected content.
|
||||
*/
|
||||
projectInputs: null,
|
||||
})
|
||||
}
|
||||
|
||||
return NextResponse.json({
|
||||
success: true,
|
||||
|
||||
@@ -893,19 +893,23 @@ describe.concurrent('Blocks Module', () => {
|
||||
|
||||
expect(providerSubBlock?.commandSearchable).toBe(true)
|
||||
expect(providerSubBlock?.value?.()).toBe('openai')
|
||||
expect(providerIds).toEqual(['openai', 'gemini', 'cohere', 'mistral'])
|
||||
expect(providerIds).toEqual(['openai', 'gemini', 'cohere', 'mistral', 'openrouter'])
|
||||
|
||||
for (const provider of providerIds) {
|
||||
// Each provider routes to its own registered tool...
|
||||
const toolId = block?.tools.config?.tool?.({ provider })
|
||||
expect(block?.tools.access).toContain(toolId)
|
||||
// ...and has a model dropdown with at least one option.
|
||||
// ...and has either a static model list or a dynamic model loader.
|
||||
const modelSubBlock = block?.subBlocks.find(
|
||||
(sb) => sb.id === 'model' && sb.condition?.value === provider
|
||||
)
|
||||
expect(
|
||||
Array.isArray(modelSubBlock?.options) ? modelSubBlock.options.length : 0
|
||||
).toBeGreaterThan(0)
|
||||
if (provider === 'openrouter') {
|
||||
expect(modelSubBlock?.fetchOptions).toBeTypeOf('function')
|
||||
} else {
|
||||
expect(
|
||||
Array.isArray(modelSubBlock?.options) ? modelSubBlock.options.length : 0
|
||||
).toBeGreaterThan(0)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
@@ -1,9 +1,29 @@
|
||||
/**
|
||||
* @vitest-environment node
|
||||
*/
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
|
||||
const { mockFetchQuery } = vi.hoisted(() => ({
|
||||
mockFetchQuery: vi.fn(),
|
||||
}))
|
||||
|
||||
vi.mock('@/app/_shell/providers/get-query-client', () => ({
|
||||
getQueryClient: () => ({ fetchQuery: mockFetchQuery }),
|
||||
}))
|
||||
|
||||
import { DEFAULT_MODEL_BY_PROVIDER, EMBEDDING_MODELS } from '@/lib/embeddings/catalog'
|
||||
import { EmbeddingsBlock, TOOL_ID_BY_PROVIDER } from '@/blocks/blocks/embeddings'
|
||||
import { DEFAULT_OPENROUTER_EMBEDDING_MODEL } from '@/lib/embeddings/openrouter-models'
|
||||
import {
|
||||
EMBEDDING_BLOCK_PROVIDERS,
|
||||
EmbeddingsBlock,
|
||||
TOOL_ID_BY_PROVIDER,
|
||||
} from '@/blocks/blocks/embeddings'
|
||||
|
||||
const OPENROUTER_MODELS = [
|
||||
'openrouter/openai/text-embedding-3-small',
|
||||
'openrouter/qwen/qwen3-embedding-8b',
|
||||
'openrouter/google/gemini-embedding-001',
|
||||
]
|
||||
|
||||
/**
|
||||
* The block derives its model, task-type, and dimension options from the
|
||||
@@ -17,7 +37,8 @@ function subBlocksById(id: string) {
|
||||
}
|
||||
|
||||
function optionIds(options: unknown): string[] {
|
||||
return Array.isArray(options) ? options.map((option) => (option as { id: string }).id) : []
|
||||
const resolved = typeof options === 'function' ? options() : options
|
||||
return Array.isArray(resolved) ? resolved.map((option) => (option as { id: string }).id) : []
|
||||
}
|
||||
|
||||
/** The provider a `{ field: 'provider', value: X }` condition selects. */
|
||||
@@ -34,19 +55,31 @@ function conditionModel(subBlock: { condition?: unknown }): string | undefined {
|
||||
}
|
||||
|
||||
describe('Embeddings block', () => {
|
||||
it('offers exactly the catalog models for each provider', () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
mockFetchQuery.mockResolvedValue({ models: OPENROUTER_MODELS })
|
||||
})
|
||||
|
||||
it('offers static catalog models for direct providers', () => {
|
||||
const modelSubBlocks = subBlocksById('model')
|
||||
const offered = new Map<string, string[]>()
|
||||
|
||||
for (const subBlock of modelSubBlocks) {
|
||||
const provider = conditionProvider(subBlock)
|
||||
expect(provider).toBeDefined()
|
||||
if (provider === 'openrouter') continue
|
||||
offered.set(provider as string, optionIds(subBlock.options))
|
||||
}
|
||||
|
||||
const expected = new Map<string, string[]>()
|
||||
for (const [modelId, info] of Object.entries(EMBEDDING_MODELS)) {
|
||||
expected.set(info.provider, [...(expected.get(info.provider) ?? []), modelId])
|
||||
for (const provider of EMBEDDING_BLOCK_PROVIDERS) {
|
||||
if (provider === 'openrouter') continue
|
||||
expected.set(
|
||||
provider,
|
||||
Object.entries(EMBEDDING_MODELS).flatMap(([modelId, info]) =>
|
||||
info.provider === provider ? [modelId] : []
|
||||
)
|
||||
)
|
||||
}
|
||||
|
||||
expect([...offered.keys()].sort()).toEqual([...expected.keys()].sort())
|
||||
@@ -63,12 +96,17 @@ describe('Embeddings block', () => {
|
||||
})
|
||||
|
||||
it('shows a task-type dropdown for exactly the models that support one', () => {
|
||||
const withTaskTypes = Object.entries(EMBEDDING_MODELS)
|
||||
.filter(([, info]) => info.supportedTaskTypes)
|
||||
.map(([id]) => id)
|
||||
const withTaskTypes = Object.entries(EMBEDDING_MODELS).flatMap(([id, info]) => {
|
||||
if (!info.supportedTaskTypes) return []
|
||||
return [`${info.provider}:${id}`]
|
||||
})
|
||||
|
||||
const subBlocks = subBlocksById('taskType')
|
||||
expect(subBlocks.map(conditionModel).sort()).toEqual(withTaskTypes.slice().sort())
|
||||
expect(
|
||||
subBlocks
|
||||
.map((subBlock) => `${conditionProvider(subBlock)}:${conditionModel(subBlock)}`)
|
||||
.sort()
|
||||
).toEqual(withTaskTypes.slice().sort())
|
||||
|
||||
for (const subBlock of subBlocks) {
|
||||
const model = conditionModel(subBlock) as string
|
||||
@@ -79,12 +117,17 @@ describe('Embeddings block', () => {
|
||||
})
|
||||
|
||||
it('shows a dimensions dropdown for exactly the models that support reduction', () => {
|
||||
const withDimensions = Object.entries(EMBEDDING_MODELS)
|
||||
.filter(([, info]) => info.supportedDimensions)
|
||||
.map(([id]) => id)
|
||||
const withDimensions = Object.entries(EMBEDDING_MODELS).flatMap(([id, info]) => {
|
||||
if (!info.supportedDimensions) return []
|
||||
return [`${info.provider}:${id}`]
|
||||
})
|
||||
|
||||
const subBlocks = subBlocksById('dimensions')
|
||||
expect(subBlocks.map(conditionModel).sort()).toEqual(withDimensions.slice().sort())
|
||||
expect(
|
||||
subBlocks
|
||||
.map((subBlock) => `${conditionProvider(subBlock)}:${conditionModel(subBlock)}`)
|
||||
.sort()
|
||||
).toEqual(withDimensions.slice().sort())
|
||||
|
||||
for (const subBlock of subBlocks) {
|
||||
const model = conditionModel(subBlock) as string
|
||||
@@ -103,6 +146,51 @@ describe('Embeddings block', () => {
|
||||
expect(EmbeddingsBlock.tools.config?.tool?.({ provider })).toBe(toolId)
|
||||
}
|
||||
expect(EmbeddingsBlock.tools.access).toHaveLength(Object.keys(TOOL_ID_BY_PROVIDER).length)
|
||||
expect(EmbeddingsBlock.tools.config?.tool?.({})).toBe('embeddings_openai')
|
||||
expect(() => EmbeddingsBlock.tools.config?.tool?.({ provider: 'unknown' })).toThrow(
|
||||
'Unsupported embedding provider: unknown'
|
||||
)
|
||||
})
|
||||
|
||||
it('loads every OpenRouter embedding model and maps its dedicated key', async () => {
|
||||
const openRouterModels = subBlocksById('model').find(
|
||||
(subBlock) => conditionProvider(subBlock) === 'openrouter'
|
||||
)
|
||||
expect(openRouterModels?.type).toBe('combobox')
|
||||
expect(openRouterModels?.options).toEqual([])
|
||||
expect(optionIds(await openRouterModels?.fetchOptions?.('block-1'))).toEqual(OPENROUTER_MODELS)
|
||||
expect(mockFetchQuery).toHaveBeenCalledOnce()
|
||||
|
||||
expect(
|
||||
EmbeddingsBlock.tools.config?.params?.({
|
||||
provider: 'openrouter',
|
||||
model: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
input: 'hello',
|
||||
apiKey: 'stale-openai-key',
|
||||
openRouterApiKey: 'or-test',
|
||||
dimensions: '1024',
|
||||
})
|
||||
).toEqual({
|
||||
apiKey: 'or-test',
|
||||
input: 'hello',
|
||||
model: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
taskType: undefined,
|
||||
dimensions: undefined,
|
||||
})
|
||||
})
|
||||
|
||||
it('requires the OpenRouter key on hosted and self-hosted deployments', () => {
|
||||
const keySubBlock = subBlocksById('openRouterApiKey')[0]
|
||||
expect(keySubBlock.required).toBe(true)
|
||||
expect(keySubBlock.hideWhenHosted).toBeUndefined()
|
||||
expect(keySubBlock.placeholder).not.toContain('OPENROUTER_API_KEY')
|
||||
expect(() =>
|
||||
EmbeddingsBlock.tools.config?.params?.({
|
||||
provider: 'openrouter',
|
||||
model: DEFAULT_OPENROUTER_EMBEDDING_MODEL,
|
||||
input: 'hello',
|
||||
})
|
||||
).toThrow('OpenRouter API key is required')
|
||||
})
|
||||
|
||||
it('only forwards capabilities the selected model declares', () => {
|
||||
@@ -194,6 +282,28 @@ describe('Embeddings block', () => {
|
||||
expect(merged.model).toBe('embed-v4.0')
|
||||
})
|
||||
|
||||
it('replaces a catalog model left over when switching to OpenRouter', () => {
|
||||
const merged = mergeLikeExecutor({
|
||||
provider: 'openrouter',
|
||||
model: 'gemini-embedding-001',
|
||||
input: 'hello',
|
||||
openRouterApiKey: 'or-test',
|
||||
})
|
||||
|
||||
expect(merged.model).toBe(DEFAULT_OPENROUTER_EMBEDDING_MODEL)
|
||||
})
|
||||
|
||||
it('rejects an invalid model entered for OpenRouter', () => {
|
||||
expect(() =>
|
||||
mergeLikeExecutor({
|
||||
provider: 'openrouter',
|
||||
model: 'not-an-openrouter-model',
|
||||
input: 'hello',
|
||||
openRouterApiKey: 'or-test',
|
||||
})
|
||||
).toThrow('Invalid OpenRouter embedding model: not-an-openrouter-model')
|
||||
})
|
||||
|
||||
it('keeps a model that does belong to the selected provider', () => {
|
||||
const merged = mergeLikeExecutor({
|
||||
provider: 'openai',
|
||||
|
||||
@@ -11,20 +11,39 @@ import {
|
||||
EMBEDDING_MODELS,
|
||||
getModelsForProvider,
|
||||
} from '@/lib/embeddings/catalog'
|
||||
import type { EmbeddingCatalogProvider, EmbeddingTaskType } from '@/lib/embeddings/types'
|
||||
import {
|
||||
DEFAULT_OPENROUTER_EMBEDDING_MODEL,
|
||||
normalizeOpenRouterEmbeddingModelId,
|
||||
} from '@/lib/embeddings/openrouter-models'
|
||||
import type { EmbeddingTaskType } from '@/lib/embeddings/types'
|
||||
import { getQueryClient } from '@/app/_shell/providers/get-query-client'
|
||||
import type { BlockConfig, BlockMeta, SubBlockConfig } from '@/blocks/types'
|
||||
import { AuthMode, IntegrationType } from '@/blocks/types'
|
||||
import { providerModelsQueryOptions } from '@/hooks/queries/providers'
|
||||
import type { EmbeddingsResponse } from '@/tools/embeddings/types'
|
||||
|
||||
const TOOL_ID_BY_PROVIDER: Record<EmbeddingCatalogProvider, string> = {
|
||||
export const EMBEDDING_BLOCK_PROVIDERS = [...EMBEDDING_CATALOG_PROVIDERS, 'openrouter'] as const
|
||||
|
||||
type EmbeddingBlockProvider = (typeof EMBEDDING_BLOCK_PROVIDERS)[number]
|
||||
|
||||
async function fetchOpenRouterEmbeddingModelOptions() {
|
||||
const { models } = await getQueryClient().fetchQuery(
|
||||
providerModelsQueryOptions('openrouter-embeddings')
|
||||
)
|
||||
return models.map((model) => ({ label: model, id: model }))
|
||||
}
|
||||
|
||||
const TOOL_ID_BY_PROVIDER: Record<EmbeddingBlockProvider, string> = {
|
||||
openai: 'embeddings_openai',
|
||||
openrouter: 'embeddings_openrouter',
|
||||
gemini: 'embeddings_gemini',
|
||||
cohere: 'embeddings_cohere',
|
||||
mistral: 'embeddings_mistral',
|
||||
}
|
||||
|
||||
const PROVIDER_LABELS: Record<EmbeddingCatalogProvider, string> = {
|
||||
const PROVIDER_LABELS: Record<EmbeddingBlockProvider, string> = {
|
||||
openai: 'OpenAI',
|
||||
openrouter: 'OpenRouter',
|
||||
gemini: 'Google Gemini',
|
||||
cohere: 'Cohere',
|
||||
mistral: 'Mistral',
|
||||
@@ -43,15 +62,31 @@ const TASK_TYPE_LABELS: Record<EmbeddingTaskType, string> = {
|
||||
* catalog model cannot leave this block stale. Every variant shares one
|
||||
* sub-block id, so each is scoped by a `condition` naming the provider.
|
||||
*/
|
||||
const MODEL_SUB_BLOCKS: SubBlockConfig[] = EMBEDDING_CATALOG_PROVIDERS.map((provider) => ({
|
||||
const MODEL_SUB_BLOCKS: SubBlockConfig[] = EMBEDDING_CATALOG_PROVIDERS.map((provider) => {
|
||||
return {
|
||||
id: 'model',
|
||||
title: 'Model',
|
||||
type: 'dropdown',
|
||||
options: getModelsForProvider(provider).map((id) => ({
|
||||
label: EMBEDDING_MODELS[id].label,
|
||||
id,
|
||||
})),
|
||||
value: () => DEFAULT_MODEL_BY_PROVIDER[provider],
|
||||
condition: { field: 'provider', value: provider },
|
||||
dependsOn: ['provider'],
|
||||
}
|
||||
})
|
||||
|
||||
MODEL_SUB_BLOCKS.push({
|
||||
id: 'model',
|
||||
title: 'Model',
|
||||
type: 'dropdown',
|
||||
options: getModelsForProvider(provider).map((id) => ({ label: EMBEDDING_MODELS[id].label, id })),
|
||||
value: () => DEFAULT_MODEL_BY_PROVIDER[provider],
|
||||
condition: { field: 'provider', value: provider },
|
||||
type: 'combobox',
|
||||
options: [],
|
||||
fetchOptions: fetchOpenRouterEmbeddingModelOptions,
|
||||
value: () => DEFAULT_OPENROUTER_EMBEDDING_MODEL,
|
||||
condition: { field: 'provider', value: 'openrouter' },
|
||||
dependsOn: ['provider'],
|
||||
}))
|
||||
})
|
||||
|
||||
/**
|
||||
* Task-type and dimension dropdowns, which are per-model rather than
|
||||
@@ -60,40 +95,42 @@ const MODEL_SUB_BLOCKS: SubBlockConfig[] = EMBEDDING_CATALOG_PROVIDERS.map((prov
|
||||
*/
|
||||
const CAPABILITY_SUB_BLOCKS: SubBlockConfig[] = Object.entries(EMBEDDING_MODELS).flatMap(
|
||||
([model, info]) => {
|
||||
const scope = { field: 'provider', value: info.provider, and: { field: 'model', value: model } }
|
||||
const subBlocks: SubBlockConfig[] = []
|
||||
return [info.provider].flatMap((provider) => {
|
||||
const scope = { field: 'provider', value: provider, and: { field: 'model', value: model } }
|
||||
const subBlocks: SubBlockConfig[] = []
|
||||
|
||||
if (info.supportedTaskTypes) {
|
||||
subBlocks.push({
|
||||
id: 'taskType',
|
||||
title: 'Task Type',
|
||||
type: 'dropdown',
|
||||
options: info.supportedTaskTypes.map((task) => ({
|
||||
label: TASK_TYPE_LABELS[task],
|
||||
id: task,
|
||||
})),
|
||||
value: () => 'document',
|
||||
condition: scope,
|
||||
dependsOn: ['provider', 'model'],
|
||||
})
|
||||
}
|
||||
if (info.supportedTaskTypes) {
|
||||
subBlocks.push({
|
||||
id: 'taskType',
|
||||
title: 'Task Type',
|
||||
type: 'dropdown',
|
||||
options: info.supportedTaskTypes.map((task) => ({
|
||||
label: TASK_TYPE_LABELS[task],
|
||||
id: task,
|
||||
})),
|
||||
value: () => 'document',
|
||||
condition: scope,
|
||||
dependsOn: ['provider', 'model'],
|
||||
})
|
||||
}
|
||||
|
||||
if (info.supportedDimensions) {
|
||||
subBlocks.push({
|
||||
id: 'dimensions',
|
||||
title: 'Dimensions',
|
||||
type: 'dropdown',
|
||||
options: info.supportedDimensions.map((size) => ({
|
||||
label: size === info.nativeDimensions ? `${size} (default)` : String(size),
|
||||
id: String(size),
|
||||
})),
|
||||
value: () => String(info.nativeDimensions),
|
||||
condition: scope,
|
||||
dependsOn: ['provider', 'model'],
|
||||
})
|
||||
}
|
||||
if (info.supportedDimensions) {
|
||||
subBlocks.push({
|
||||
id: 'dimensions',
|
||||
title: 'Dimensions',
|
||||
type: 'dropdown',
|
||||
options: info.supportedDimensions.map((size) => ({
|
||||
label: size === info.nativeDimensions ? `${size} (default)` : String(size),
|
||||
id: String(size),
|
||||
})),
|
||||
value: () => String(info.nativeDimensions),
|
||||
condition: scope,
|
||||
dependsOn: ['provider', 'model'],
|
||||
})
|
||||
}
|
||||
|
||||
return subBlocks
|
||||
return subBlocks
|
||||
})
|
||||
}
|
||||
)
|
||||
|
||||
@@ -103,7 +140,7 @@ export const EmbeddingsBlock: BlockConfig<EmbeddingsResponse> = {
|
||||
description: 'Generate embeddings',
|
||||
authMode: AuthMode.ApiKey,
|
||||
longDescription:
|
||||
'Turn text into embedding vectors for semantic search, clustering, and similarity. Supports OpenAI, Google Gemini, Cohere, and Mistral embedding models.',
|
||||
'Turn text into embedding vectors for semantic search, clustering, and similarity. Supports OpenAI, OpenRouter, Google Gemini, Cohere, and Mistral embedding models.',
|
||||
category: 'tools',
|
||||
integrationType: IntegrationType.AI,
|
||||
docsLink: 'https://docs.sim.ai/integrations/embeddings',
|
||||
@@ -121,7 +158,7 @@ export const EmbeddingsBlock: BlockConfig<EmbeddingsResponse> = {
|
||||
id: 'provider',
|
||||
title: 'Provider',
|
||||
type: 'dropdown',
|
||||
options: EMBEDDING_CATALOG_PROVIDERS.map((provider) => ({
|
||||
options: EMBEDDING_BLOCK_PROVIDERS.map((provider) => ({
|
||||
label: PROVIDER_LABELS[provider],
|
||||
id: provider,
|
||||
})),
|
||||
@@ -131,9 +168,9 @@ export const EmbeddingsBlock: BlockConfig<EmbeddingsResponse> = {
|
||||
...MODEL_SUB_BLOCKS,
|
||||
...CAPABILITY_SUB_BLOCKS,
|
||||
/**
|
||||
* One field for every provider. Sim stocks a hosted key for all four
|
||||
* (`OPENAI_API_KEY`, `GEMINI_API_KEY`, `COHERE_API_KEY`, `MISTRAL_API_KEY`),
|
||||
* so none of them needs the user to supply one on hosted Sim.
|
||||
* Sim stocks a hosted key for each catalog provider, so none of those
|
||||
* fields needs the user to supply one on hosted Sim. OpenRouter is always
|
||||
* explicit BYOK for this block.
|
||||
*/
|
||||
{
|
||||
id: 'apiKey',
|
||||
@@ -142,12 +179,29 @@ export const EmbeddingsBlock: BlockConfig<EmbeddingsResponse> = {
|
||||
placeholder: 'Enter your provider API key',
|
||||
password: true,
|
||||
required: true,
|
||||
condition: { field: 'provider', value: 'openrouter', not: true },
|
||||
connectionDroppable: false,
|
||||
hideWhenHosted: true,
|
||||
},
|
||||
{
|
||||
id: 'openRouterApiKey',
|
||||
title: 'OpenRouter API Key',
|
||||
type: 'short-input',
|
||||
placeholder: 'Enter your OpenRouter API key',
|
||||
password: true,
|
||||
required: true,
|
||||
condition: { field: 'provider', value: 'openrouter' },
|
||||
connectionDroppable: false,
|
||||
},
|
||||
],
|
||||
tools: {
|
||||
access: ['embeddings_openai', 'embeddings_gemini', 'embeddings_cohere', 'embeddings_mistral'],
|
||||
access: [
|
||||
'embeddings_openai',
|
||||
'embeddings_openrouter',
|
||||
'embeddings_gemini',
|
||||
'embeddings_cohere',
|
||||
'embeddings_mistral',
|
||||
],
|
||||
config: {
|
||||
/**
|
||||
* Runs at serialization, before variable resolution, so this only ever
|
||||
@@ -155,8 +209,10 @@ export const EmbeddingsBlock: BlockConfig<EmbeddingsResponse> = {
|
||||
* `<Block.output>` references.
|
||||
*/
|
||||
tool: (params) => {
|
||||
const provider = params.provider as EmbeddingCatalogProvider
|
||||
return TOOL_ID_BY_PROVIDER[provider] ?? TOOL_ID_BY_PROVIDER.openai
|
||||
const provider = (params.provider as EmbeddingBlockProvider | undefined) ?? 'openai'
|
||||
const toolId = TOOL_ID_BY_PROVIDER[provider]
|
||||
if (!toolId) throw new Error(`Unsupported embedding provider: ${String(params.provider)}`)
|
||||
return toolId
|
||||
},
|
||||
/**
|
||||
* Every per-provider dropdown shares one subblock id (`model`,
|
||||
@@ -171,17 +227,42 @@ export const EmbeddingsBlock: BlockConfig<EmbeddingsResponse> = {
|
||||
* overrides it.
|
||||
*/
|
||||
params: (params) => {
|
||||
const provider = (params.provider as EmbeddingCatalogProvider) || 'openai'
|
||||
const provider = (params.provider as EmbeddingBlockProvider) || 'openai'
|
||||
if (!params.input) {
|
||||
throw new Error('Input text is required')
|
||||
}
|
||||
|
||||
if (provider === 'openrouter') {
|
||||
if (typeof params.openRouterApiKey !== 'string' || !params.openRouterApiKey.trim()) {
|
||||
throw new Error('OpenRouter API key is required')
|
||||
}
|
||||
const savedModel =
|
||||
typeof params.model === 'string' && params.model ? params.model : undefined
|
||||
const savedCatalogProvider = savedModel
|
||||
? EMBEDDING_MODELS[savedModel]?.provider
|
||||
: undefined
|
||||
const model = normalizeOpenRouterEmbeddingModelId(
|
||||
savedCatalogProvider && savedCatalogProvider !== 'openai'
|
||||
? DEFAULT_OPENROUTER_EMBEDDING_MODEL
|
||||
: (savedModel ?? DEFAULT_OPENROUTER_EMBEDDING_MODEL)
|
||||
)
|
||||
return {
|
||||
apiKey: params.openRouterApiKey,
|
||||
input: params.input,
|
||||
model,
|
||||
taskType: undefined,
|
||||
dimensions: undefined,
|
||||
}
|
||||
}
|
||||
|
||||
const catalogProvider = provider
|
||||
|
||||
/** A model saved under a previous provider must not survive the switch. */
|
||||
const savedModel = params.model as string | undefined
|
||||
const model =
|
||||
savedModel && EMBEDDING_MODELS[savedModel]?.provider === provider
|
||||
savedModel && EMBEDDING_MODELS[savedModel]?.provider === catalogProvider
|
||||
? savedModel
|
||||
: DEFAULT_MODEL_BY_PROVIDER[provider]
|
||||
: DEFAULT_MODEL_BY_PROVIDER[catalogProvider]
|
||||
|
||||
const info = EMBEDDING_MODELS[model]
|
||||
const requested =
|
||||
@@ -217,6 +298,10 @@ export const EmbeddingsBlock: BlockConfig<EmbeddingsResponse> = {
|
||||
taskType: { type: 'string', description: 'What the embedding will be used for' },
|
||||
dimensions: { type: 'number', description: 'Output vector dimensions' },
|
||||
apiKey: { type: 'string', description: 'Provider API key' },
|
||||
openRouterApiKey: {
|
||||
type: 'string',
|
||||
description: 'OpenRouter API key',
|
||||
},
|
||||
},
|
||||
outputs: {
|
||||
embeddings: { type: 'json', description: 'Generated embeddings' },
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import { createLogger } from '@sim/logger'
|
||||
import { getErrorMessage } from '@sim/utils/errors'
|
||||
import { useQuery } from '@tanstack/react-query'
|
||||
import { queryOptions, useQuery } from '@tanstack/react-query'
|
||||
import { requestJson } from '@/lib/api/client/request'
|
||||
import {
|
||||
getBaseProviderModelsContract,
|
||||
@@ -9,6 +9,7 @@ import {
|
||||
getLitellmProviderModelsContract,
|
||||
getOllamaCloudProviderModelsContract,
|
||||
getOllamaProviderModelsContract,
|
||||
getOpenRouterEmbeddingModelsContract,
|
||||
getOpenRouterProviderModelsContract,
|
||||
getTogetherProviderModelsContract,
|
||||
getVllmProviderModelsContract,
|
||||
@@ -16,6 +17,8 @@ import {
|
||||
} from '@/lib/api/contracts/providers'
|
||||
import type { ProviderName } from '@/stores/providers'
|
||||
|
||||
type ProviderModelSource = ProviderName | 'openrouter-embeddings'
|
||||
|
||||
const logger = createLogger('ProviderModelsQuery')
|
||||
|
||||
export const PROVIDER_MODELS_STALE_TIME = 5 * 60 * 1000
|
||||
@@ -28,14 +31,14 @@ export const providerKeys = {
|
||||
}
|
||||
|
||||
async function fetchProviderModels(
|
||||
provider: ProviderName,
|
||||
provider: ProviderModelSource,
|
||||
signal?: AbortSignal,
|
||||
workspaceId?: string
|
||||
): Promise<ProviderModelsResponse> {
|
||||
try {
|
||||
const data = await requestProviderModels(provider, signal, workspaceId)
|
||||
const models: string[] = Array.isArray(data.models) ? data.models : []
|
||||
const uniqueModels = provider === 'openrouter' ? Array.from(new Set(models)) : models
|
||||
const uniqueModels = provider.startsWith('openrouter') ? Array.from(new Set(models)) : models
|
||||
|
||||
return {
|
||||
models: uniqueModels,
|
||||
@@ -50,7 +53,7 @@ async function fetchProviderModels(
|
||||
}
|
||||
|
||||
async function requestProviderModels(
|
||||
provider: ProviderName,
|
||||
provider: ProviderModelSource,
|
||||
signal?: AbortSignal,
|
||||
workspaceId?: string
|
||||
): Promise<ProviderModelsResponse> {
|
||||
@@ -70,6 +73,8 @@ async function requestProviderModels(
|
||||
return requestJson(getLitellmProviderModelsContract, { signal })
|
||||
case 'openrouter':
|
||||
return requestJson(getOpenRouterProviderModelsContract, { signal })
|
||||
case 'openrouter-embeddings':
|
||||
return requestJson(getOpenRouterEmbeddingModelsContract, { signal })
|
||||
case 'fireworks':
|
||||
return requestJson(getFireworksProviderModelsContract, {
|
||||
query: { workspaceId },
|
||||
@@ -88,10 +93,14 @@ async function requestProviderModels(
|
||||
}
|
||||
}
|
||||
|
||||
export function useProviderModels(provider: ProviderName, workspaceId?: string) {
|
||||
return useQuery({
|
||||
export function providerModelsQueryOptions(provider: ProviderModelSource, workspaceId?: string) {
|
||||
return queryOptions({
|
||||
queryKey: providerKeys.list(provider, workspaceId),
|
||||
queryFn: ({ signal }) => fetchProviderModels(provider, signal, workspaceId),
|
||||
staleTime: PROVIDER_MODELS_STALE_TIME,
|
||||
})
|
||||
}
|
||||
|
||||
export function useProviderModels(provider: ProviderModelSource, workspaceId?: string) {
|
||||
return useQuery(providerModelsQueryOptions(provider, workspaceId))
|
||||
}
|
||||
|
||||
@@ -59,6 +59,20 @@ export const openRouterUpstreamResponseSchema = z.object({
|
||||
.default([]),
|
||||
})
|
||||
|
||||
export const openRouterEmbeddingModelsUpstreamResponseSchema = z.object({
|
||||
data: z.array(
|
||||
z
|
||||
.object({
|
||||
id: z.string().min(1, 'OpenRouter embedding model id cannot be empty'),
|
||||
context_length: z
|
||||
.number()
|
||||
.int('OpenRouter embedding context length must be an integer')
|
||||
.positive('OpenRouter embedding context length must be positive'),
|
||||
})
|
||||
.passthrough()
|
||||
),
|
||||
})
|
||||
|
||||
export const vllmUpstreamResponseSchema = z.object({
|
||||
data: z
|
||||
.array(
|
||||
@@ -254,6 +268,15 @@ export const getOpenRouterProviderModelsContract = defineRouteContract({
|
||||
},
|
||||
})
|
||||
|
||||
export const getOpenRouterEmbeddingModelsContract = defineRouteContract({
|
||||
method: 'GET',
|
||||
path: '/api/providers/openrouter/embeddings/models',
|
||||
response: {
|
||||
mode: 'json',
|
||||
schema: providerModelsResponseSchema,
|
||||
},
|
||||
})
|
||||
|
||||
export const getLitellmProviderModelsContract = defineRouteContract({
|
||||
method: 'GET',
|
||||
path: '/api/providers/litellm/models',
|
||||
|
||||
@@ -8,12 +8,15 @@ import type { EmbeddingCatalogProvider, EmbeddingTaskType } from '@/lib/embeddin
|
||||
* addition — a new catalog provider or task type stays absent from the wire
|
||||
* enum until it is added below.
|
||||
*/
|
||||
type EmbeddingToolProvider = EmbeddingCatalogProvider | 'openrouter'
|
||||
|
||||
export const embeddingProviders = [
|
||||
'openai',
|
||||
'openrouter',
|
||||
'gemini',
|
||||
'cohere',
|
||||
'mistral',
|
||||
] as const satisfies readonly EmbeddingCatalogProvider[]
|
||||
] as const satisfies readonly EmbeddingToolProvider[]
|
||||
|
||||
export const embeddingTaskTypes = [
|
||||
'document',
|
||||
@@ -28,15 +31,15 @@ export const MAX_EMBEDDING_INPUTS = 1000
|
||||
/** Caps total payload size independently of the input count. */
|
||||
export const MAX_EMBEDDING_TOTAL_CHARS = 1_000_000
|
||||
|
||||
const MISSING_EMBEDDING_FIELDS_ERROR = 'Missing required fields: provider, apiKey, and input'
|
||||
const MISSING_EMBEDDING_INPUT_ERROR = 'Missing required field: input'
|
||||
const embeddingCatalogProviders = [
|
||||
'openai',
|
||||
'gemini',
|
||||
'cohere',
|
||||
'mistral',
|
||||
] as const satisfies readonly EmbeddingCatalogProvider[]
|
||||
|
||||
export const embeddingsToolBodySchema = z.object({
|
||||
provider: z.enum(embeddingProviders, {
|
||||
error: `Invalid provider. Must be one of: ${embeddingProviders.join(', ')}`,
|
||||
}),
|
||||
apiKey: z
|
||||
.string({ error: MISSING_EMBEDDING_FIELDS_ERROR })
|
||||
.min(1, MISSING_EMBEDDING_FIELDS_ERROR),
|
||||
const embeddingToolCommonShape = {
|
||||
model: z.string().min(1, 'model cannot be empty').optional(),
|
||||
/** A single text, or an array of texts embedded in one call. */
|
||||
input: z.union(
|
||||
@@ -47,7 +50,7 @@ export const embeddingsToolBodySchema = z.object({
|
||||
.min(1, 'input must contain at least one text')
|
||||
.max(MAX_EMBEDDING_INPUTS, `input cannot exceed ${MAX_EMBEDDING_INPUTS} texts`),
|
||||
],
|
||||
{ error: MISSING_EMBEDDING_FIELDS_ERROR }
|
||||
{ error: MISSING_EMBEDDING_INPUT_ERROR }
|
||||
),
|
||||
taskType: z.enum(embeddingTaskTypes).optional(),
|
||||
/** Matryoshka output size. Omitted means the model's native dimensionality. */
|
||||
@@ -57,7 +60,20 @@ export const embeddingsToolBodySchema = z.object({
|
||||
.min(1, 'dimensions must be at least 1')
|
||||
.max(4096, 'dimensions cannot exceed 4096')
|
||||
.optional(),
|
||||
})
|
||||
}
|
||||
|
||||
export const embeddingsToolBodySchema = z.discriminatedUnion('provider', [
|
||||
z.object({
|
||||
...embeddingToolCommonShape,
|
||||
provider: z.enum(embeddingCatalogProviders),
|
||||
apiKey: z.string({ error: 'apiKey is required' }).min(1, 'apiKey cannot be empty'),
|
||||
}),
|
||||
z.object({
|
||||
...embeddingToolCommonShape,
|
||||
provider: z.literal('openrouter'),
|
||||
apiKey: z.string({ error: 'apiKey is required' }).min(1, 'apiKey cannot be empty'),
|
||||
}),
|
||||
])
|
||||
|
||||
const embeddingsUsageSchema = z.object({
|
||||
prompt_tokens: z.number(),
|
||||
|
||||
@@ -9,6 +9,7 @@ import {
|
||||
envField,
|
||||
inspectCapability,
|
||||
inspectOAuthClientCapability,
|
||||
KNOWLEDGE_EMBEDDINGS_CAPABILITY,
|
||||
LLM_KEY_POOLS,
|
||||
OCR_CAPABILITY,
|
||||
requireCapability,
|
||||
@@ -111,6 +112,43 @@ describe('env capabilities', () => {
|
||||
})
|
||||
|
||||
describe('fallback capabilities', () => {
|
||||
it('resolves configured knowledge embedding transports in fallback order', () => {
|
||||
expect(
|
||||
inspectCapability(KNOWLEDGE_EMBEDDINGS_CAPABILITY, {
|
||||
OPENROUTER_API_KEY: 'openrouter-key',
|
||||
}).providerIds
|
||||
).toEqual(['openrouter'])
|
||||
|
||||
expect(
|
||||
inspectCapability(KNOWLEDGE_EMBEDDINGS_CAPABILITY, {
|
||||
AZURE_OPENAI_API_KEY: 'azure-key',
|
||||
AZURE_OPENAI_ENDPOINT: 'https://azure.example.com',
|
||||
AZURE_OPENAI_API_VERSION: '2024-10-21',
|
||||
OPENAI_API_KEY_1: 'openai-key',
|
||||
OPENROUTER_API_KEY: 'openrouter-key',
|
||||
}).providerIds
|
||||
).toEqual(['azure-openai', 'openai', 'openrouter'])
|
||||
})
|
||||
|
||||
it('reports partially configured Azure knowledge embeddings', () => {
|
||||
const inspection = inspectCapability(KNOWLEDGE_EMBEDDINGS_CAPABILITY, {
|
||||
AZURE_OPENAI_API_KEY: 'azure-key',
|
||||
})
|
||||
|
||||
expect(inspection).toMatchObject({
|
||||
configured: false,
|
||||
providerIds: [],
|
||||
error: expect.any(EnvCapabilityConfigurationError),
|
||||
})
|
||||
expect(inspection.providers[0]).toMatchObject({
|
||||
state: 'partial',
|
||||
missingFields: expect.arrayContaining([
|
||||
'AZURE_OPENAI_ENDPOINT',
|
||||
'AZURE_OPENAI_API_VERSION',
|
||||
]),
|
||||
})
|
||||
})
|
||||
|
||||
it('resolves every ready email provider subset in declaration order', () => {
|
||||
for (let mask = 0; mask < 1 << EMAIL_PROVIDER_ORDER.length; mask += 1) {
|
||||
const expected = EMAIL_PROVIDER_ORDER.filter((_, index) => (mask & (1 << index)) !== 0)
|
||||
@@ -251,6 +289,27 @@ describe('env capabilities', () => {
|
||||
expect(onFailure.mock.calls.map(([providerId]) => providerId)).toEqual(['resend', 'ses'])
|
||||
})
|
||||
|
||||
it('stops fallback immediately when the error predicate rejects an error', async () => {
|
||||
const fatal = new Error('invalid credentials')
|
||||
const first = { send: vi.fn().mockRejectedValue(fatal) }
|
||||
const second = { send: vi.fn().mockResolvedValue('should not run') }
|
||||
const fallback = wireFallback({
|
||||
definition: EMAIL_CAPABILITY,
|
||||
values: { RESEND_API_KEY: 're_test', AWS_SES_REGION: 'us-east-1' },
|
||||
factories: {
|
||||
resend: () => first,
|
||||
ses: () => second,
|
||||
smtp: () => null,
|
||||
azure: () => null,
|
||||
gmail: () => null,
|
||||
},
|
||||
shouldFallback: () => false,
|
||||
})
|
||||
|
||||
await expect(fallback.execute((provider) => provider.send())).rejects.toBe(fatal)
|
||||
expect(second.send).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('fails immediately when a ready provider has no runtime implementation', () => {
|
||||
expect(() =>
|
||||
wireFallback({
|
||||
|
||||
@@ -845,6 +845,7 @@ export interface WireFallbackOptions<TDefinition extends FallbackCapabilityDefin
|
||||
definition: TDefinition
|
||||
values: EnvCapabilityValues
|
||||
factories: FallbackFactories<TDefinition, TProvider>
|
||||
shouldFallback?: (error: unknown, providerId: DeclaredProviderId<TDefinition>) => boolean
|
||||
onFailure?: (providerId: DeclaredProviderId<TDefinition>, error: unknown) => void
|
||||
}
|
||||
|
||||
@@ -852,6 +853,7 @@ export function wireFallback<const TDefinition extends FallbackCapabilityDefinit
|
||||
definition,
|
||||
values,
|
||||
factories,
|
||||
shouldFallback,
|
||||
onFailure,
|
||||
}: WireFallbackOptions<TDefinition, TProvider>) {
|
||||
const resolution = inspectCapability(definition, values)
|
||||
@@ -889,6 +891,7 @@ export function wireFallback<const TDefinition extends FallbackCapabilityDefinit
|
||||
try {
|
||||
return await operation(provider, providerId)
|
||||
} catch (error) {
|
||||
if (shouldFallback && !shouldFallback(error, providerId)) throw error
|
||||
failures.push(error)
|
||||
onFailure?.(providerId, error)
|
||||
}
|
||||
@@ -1207,6 +1210,54 @@ export const OCR_CAPABILITY = defineCapability({
|
||||
],
|
||||
} as const)
|
||||
|
||||
export const KNOWLEDGE_EMBEDDINGS_CAPABILITY = defineCapability({
|
||||
strategy: 'fallback',
|
||||
id: 'knowledge-embeddings',
|
||||
label: 'Knowledge embeddings',
|
||||
providers: [
|
||||
{
|
||||
id: 'azure-openai',
|
||||
label: 'Azure OpenAI',
|
||||
activation: {
|
||||
mode: 'any-present',
|
||||
keys: ['AZURE_OPENAI_API_KEY', 'AZURE_OPENAI_ENDPOINT', 'AZURE_OPENAI_API_VERSION'],
|
||||
},
|
||||
requires: allOf(
|
||||
envField('AZURE_OPENAI_API_KEY'),
|
||||
envField('AZURE_OPENAI_ENDPOINT', {
|
||||
validation: {
|
||||
kind: 'url',
|
||||
protocols: ['http:', 'https:'],
|
||||
message: 'must be a valid HTTP(S) URL',
|
||||
},
|
||||
}),
|
||||
envField('AZURE_OPENAI_API_VERSION')
|
||||
),
|
||||
optionalFields: [envField('KB_OPENAI_MODEL_NAME')],
|
||||
},
|
||||
{
|
||||
id: 'openai',
|
||||
label: 'OpenAI',
|
||||
activation: {
|
||||
mode: 'any-present',
|
||||
keys: ['OPENAI_API_KEY', 'OPENAI_API_KEY_1', 'OPENAI_API_KEY_2', 'OPENAI_API_KEY_3'],
|
||||
},
|
||||
requires: anyOf(
|
||||
envField('OPENAI_API_KEY'),
|
||||
envField('OPENAI_API_KEY_1'),
|
||||
envField('OPENAI_API_KEY_2'),
|
||||
envField('OPENAI_API_KEY_3')
|
||||
),
|
||||
},
|
||||
{
|
||||
id: 'openrouter',
|
||||
label: 'OpenRouter',
|
||||
activation: { mode: 'any-present', keys: ['OPENROUTER_API_KEY'] },
|
||||
requires: envField('OPENROUTER_API_KEY'),
|
||||
},
|
||||
],
|
||||
} as const)
|
||||
|
||||
export const OAUTH_CLIENT_CAPABILITIES = {
|
||||
google: ['GOOGLE_CLIENT_ID', 'GOOGLE_CLIENT_SECRET'],
|
||||
x: ['X_CLIENT_ID', 'X_CLIENT_SECRET'],
|
||||
@@ -1250,6 +1301,7 @@ export const ENV_CAPABILITIES = [
|
||||
ASYNC_JOBS_CAPABILITY,
|
||||
CACHE_CAPABILITY,
|
||||
OCR_CAPABILITY,
|
||||
KNOWLEDGE_EMBEDDINGS_CAPABILITY,
|
||||
] as const
|
||||
|
||||
export const LLM_KEY_POOLS = {
|
||||
|
||||
@@ -183,6 +183,7 @@ export const env = createEnv({
|
||||
OPENAI_API_KEY_1: z.string().min(1).optional(), // Additional OpenAI API key for load balancing
|
||||
OPENAI_API_KEY_2: z.string().min(1).optional(), // Additional OpenAI API key for load balancing
|
||||
OPENAI_API_KEY_3: z.string().min(1).optional(), // Additional OpenAI API key for load balancing
|
||||
OPENROUTER_API_KEY: z.string().min(1).optional(), // OpenRouter API key; self-hosted fallback for OpenAI knowledge-base embeddings
|
||||
MISTRAL_API_KEY: z.string().min(1).optional(), // Mistral AI API key
|
||||
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
|
||||
|
||||
@@ -1,8 +1,23 @@
|
||||
/**
|
||||
* @vitest-environment node
|
||||
*/
|
||||
import { resetEnvMock, setEnv } from '@sim/testing'
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import { embed } from '@/lib/embeddings/client'
|
||||
import {
|
||||
EmbeddingAPIError,
|
||||
embed,
|
||||
embedKnowledgeForDeployment,
|
||||
embedOpenRouter,
|
||||
isTransientEmbeddingError,
|
||||
} from '@/lib/embeddings/client'
|
||||
|
||||
const { mockGetBYOKKey } = vi.hoisted(() => ({
|
||||
mockGetBYOKKey: vi.fn(),
|
||||
}))
|
||||
|
||||
vi.mock('@/lib/api-key/byok', () => ({
|
||||
getBYOKKey: mockGetBYOKKey,
|
||||
}))
|
||||
|
||||
/**
|
||||
* Exercises the orchestrator end-to-end against a mocked transport: batching,
|
||||
@@ -35,11 +50,25 @@ let fetchMock: ReturnType<typeof vi.fn>
|
||||
beforeEach(() => {
|
||||
fetchMock = vi.fn()
|
||||
global.fetch = fetchMock as unknown as typeof fetch
|
||||
mockGetBYOKKey.mockResolvedValue(null)
|
||||
setEnv({
|
||||
AZURE_OPENAI_API_KEY: undefined,
|
||||
AZURE_OPENAI_ENDPOINT: undefined,
|
||||
AZURE_OPENAI_API_VERSION: undefined,
|
||||
GEMINI_API_KEY: undefined,
|
||||
OPENAI_API_KEY: undefined,
|
||||
OPENAI_API_KEY_1: undefined,
|
||||
OPENAI_API_KEY_2: undefined,
|
||||
OPENAI_API_KEY_3: undefined,
|
||||
OPENROUTER_API_KEY: undefined,
|
||||
})
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
global.fetch = originalFetch
|
||||
vi.useRealTimers()
|
||||
vi.restoreAllMocks()
|
||||
resetEnvMock()
|
||||
})
|
||||
|
||||
describe('embed', () => {
|
||||
@@ -258,6 +287,26 @@ describe('embed', () => {
|
||||
})
|
||||
|
||||
expect(result.isBYOK).toBe(true)
|
||||
expect(result.billableTokens).toBe(0)
|
||||
})
|
||||
|
||||
it('uses OpenRouter as an explicit transport for an OpenAI catalog model', async () => {
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2]])))
|
||||
|
||||
await embed(['hello'], {
|
||||
model: 'text-embedding-3-large',
|
||||
transport: 'openrouter',
|
||||
apiKey: 'or-test',
|
||||
dimensions: 1024,
|
||||
projectInputs: null,
|
||||
})
|
||||
|
||||
const [url, init] = fetchMock.mock.calls[0]
|
||||
expect(url).toBe('https://openrouter.ai/api/v1/embeddings')
|
||||
expect(JSON.parse((init as RequestInit).body as string)).toMatchObject({
|
||||
model: 'openai/text-embedding-3-large',
|
||||
dimensions: 1024,
|
||||
})
|
||||
})
|
||||
|
||||
/**
|
||||
@@ -384,3 +433,297 @@ describe('embed', () => {
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe('embedOpenRouter', () => {
|
||||
it('uses a dynamic model and reports the returned native dimensions', async () => {
|
||||
fetchMock.mockResolvedValue(
|
||||
jsonResponse(
|
||||
openAIBody(
|
||||
[
|
||||
[1, 2, 3],
|
||||
[4, 5, 6],
|
||||
],
|
||||
7
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
const result = await embedOpenRouter(['alpha', 'beta'], {
|
||||
model: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
apiKey: 'or-test',
|
||||
maxInputTokens: 32768,
|
||||
projectInputs: null,
|
||||
})
|
||||
|
||||
const [url, init] = fetchMock.mock.calls[0]
|
||||
expect(url).toBe('https://openrouter.ai/api/v1/embeddings')
|
||||
expect(JSON.parse((init as RequestInit).body as string)).toMatchObject({
|
||||
input: ['alpha', 'beta'],
|
||||
model: 'qwen/qwen3-embedding-8b',
|
||||
})
|
||||
expect(result).toMatchObject({
|
||||
embeddings: [
|
||||
[1, 2, 3],
|
||||
[4, 5, 6],
|
||||
],
|
||||
dimensions: 3,
|
||||
totalTokens: 7,
|
||||
billableTokens: 0,
|
||||
isBYOK: true,
|
||||
modelName: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
})
|
||||
})
|
||||
|
||||
it('fails when OpenRouter returns the wrong number of vectors', async () => {
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2]])))
|
||||
|
||||
await expect(
|
||||
embedOpenRouter(['alpha', 'beta'], {
|
||||
model: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
apiKey: 'or-test',
|
||||
maxInputTokens: 32768,
|
||||
projectInputs: null,
|
||||
})
|
||||
).rejects.toThrow('returned 1 embeddings for 2 inputs')
|
||||
})
|
||||
|
||||
it('fails when OpenRouter returns inconsistent vector dimensions', async () => {
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2], [3]])))
|
||||
|
||||
await expect(
|
||||
embedOpenRouter(['alpha', 'beta'], {
|
||||
model: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
apiKey: 'or-test',
|
||||
maxInputTokens: 32768,
|
||||
projectInputs: null,
|
||||
})
|
||||
).rejects.toThrow('inconsistent dimensions')
|
||||
})
|
||||
|
||||
it('truncates inputs to the selected model context length', async () => {
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2]])))
|
||||
|
||||
await embedOpenRouter(['alpha beta gamma'], {
|
||||
model: 'openrouter/thenlper/gte-base',
|
||||
apiKey: 'or-test',
|
||||
maxInputTokens: 1,
|
||||
projectInputs: null,
|
||||
})
|
||||
|
||||
const [, init] = fetchMock.mock.calls[0]
|
||||
const body = JSON.parse((init as RequestInit).body as string)
|
||||
expect(body.input).toHaveLength(1)
|
||||
expect(body.input[0]).not.toBe('alpha beta gamma')
|
||||
})
|
||||
|
||||
it('splits dynamic models at the provider item limit and recombines in order', async () => {
|
||||
fetchMock.mockImplementation(async (_url, init) => {
|
||||
const body = JSON.parse((init as RequestInit).body as string)
|
||||
const inputs = body.input as string[]
|
||||
return jsonResponse(
|
||||
openAIBody(
|
||||
inputs.map((input) => [Number(input.slice(1))]),
|
||||
inputs.length
|
||||
)
|
||||
)
|
||||
})
|
||||
const inputs = Array.from({ length: 2049 }, (_, index) => `i${index}`)
|
||||
|
||||
const result = await embedOpenRouter(inputs, {
|
||||
model: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
apiKey: 'or-test',
|
||||
maxInputTokens: 32768,
|
||||
projectInputs: null,
|
||||
})
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledTimes(2)
|
||||
expect(
|
||||
fetchMock.mock.calls
|
||||
.map(([, init]) => JSON.parse((init as RequestInit).body as string).input.length)
|
||||
.sort((a, b) => a - b)
|
||||
).toEqual([1, 2048])
|
||||
expect(result.embeddings).toHaveLength(2049)
|
||||
expect(result.embeddings[0]).toEqual([0])
|
||||
expect(result.embeddings[2048]).toEqual([2048])
|
||||
})
|
||||
})
|
||||
|
||||
describe('knowledge embedding transport fallback', () => {
|
||||
const options = {
|
||||
model: 'text-embedding-3-small',
|
||||
taskType: 'document' as const,
|
||||
dimensions: 1536,
|
||||
projectInputs: null,
|
||||
}
|
||||
|
||||
it('uses OpenRouter when it is the only configured self-hosted transport', async () => {
|
||||
setEnv({ OPENROUTER_API_KEY: 'or-test' })
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2]], 3)))
|
||||
|
||||
const result = await embedKnowledgeForDeployment(['hello'], options, false)
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledOnce()
|
||||
const [url, init] = fetchMock.mock.calls[0]
|
||||
expect(url).toBe('https://openrouter.ai/api/v1/embeddings')
|
||||
expect(JSON.parse((init as RequestInit).body as string)).toMatchObject({
|
||||
model: 'openai/text-embedding-3-small',
|
||||
dimensions: 1536,
|
||||
})
|
||||
expect(result).toMatchObject({
|
||||
embeddings: [[1, 2]],
|
||||
billableTokens: 3,
|
||||
isBYOK: false,
|
||||
modelName: 'text-embedding-3-small',
|
||||
dimensions: 1536,
|
||||
})
|
||||
})
|
||||
|
||||
it('keeps the original OpenAI path when OpenRouter is not configured', async () => {
|
||||
setEnv({ OPENAI_API_KEY: 'openai-test' })
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2]])))
|
||||
|
||||
await embedKnowledgeForDeployment(['hello'], options, false)
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledOnce()
|
||||
expect(fetchMock.mock.calls[0][0]).toBe('https://api.openai.com/v1/embeddings')
|
||||
})
|
||||
|
||||
it('uses Azure before OpenAI and OpenRouter when all are configured', async () => {
|
||||
setEnv({
|
||||
AZURE_OPENAI_API_KEY: 'azure-test',
|
||||
AZURE_OPENAI_ENDPOINT: 'https://example.openai.azure.com',
|
||||
AZURE_OPENAI_API_VERSION: '2024-10-21',
|
||||
KB_OPENAI_MODEL_NAME: 'kb-embedding-deployment',
|
||||
OPENAI_API_KEY: 'openai-test',
|
||||
OPENROUTER_API_KEY: 'or-test',
|
||||
})
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2]])))
|
||||
|
||||
const result = await embedKnowledgeForDeployment(['hello'], options, false)
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledOnce()
|
||||
expect(fetchMock.mock.calls[0][0]).toBe(
|
||||
'https://example.openai.azure.com/openai/deployments/kb-embedding-deployment/embeddings?api-version=2024-10-21'
|
||||
)
|
||||
expect(result.modelName).toBe('kb-embedding-deployment')
|
||||
})
|
||||
|
||||
it('uses a workspace OpenAI key before OpenRouter', async () => {
|
||||
setEnv({ OPENROUTER_API_KEY: 'or-test' })
|
||||
mockGetBYOKKey.mockResolvedValue({ apiKey: 'workspace-openai-test', isBYOK: true })
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2]])))
|
||||
|
||||
const result = await embedKnowledgeForDeployment(
|
||||
['hello'],
|
||||
{ ...options, workspaceId: 'workspace-1' },
|
||||
false
|
||||
)
|
||||
|
||||
expect(mockGetBYOKKey).toHaveBeenCalledWith('workspace-1', 'openai')
|
||||
expect(fetchMock).toHaveBeenCalledOnce()
|
||||
const [url, init] = fetchMock.mock.calls[0]
|
||||
expect(url).toBe('https://api.openai.com/v1/embeddings')
|
||||
expect((init as RequestInit).headers).toMatchObject({
|
||||
Authorization: 'Bearer workspace-openai-test',
|
||||
})
|
||||
expect(result.isBYOK).toBe(true)
|
||||
})
|
||||
|
||||
it('does not use OpenRouter for non-OpenAI knowledge models', async () => {
|
||||
setEnv({ GEMINI_API_KEY: 'gemini-test', OPENROUTER_API_KEY: 'or-test' })
|
||||
fetchMock.mockResolvedValue(jsonResponse({ embeddings: [{ values: [1, 2] }] }))
|
||||
|
||||
await embedKnowledgeForDeployment(
|
||||
['hello'],
|
||||
{ ...options, model: 'gemini-embedding-001' },
|
||||
false
|
||||
)
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledOnce()
|
||||
expect(fetchMock.mock.calls[0][0]).toContain('generativelanguage.googleapis.com')
|
||||
})
|
||||
|
||||
it('ignores OpenRouter on hosted deployments', async () => {
|
||||
setEnv({ OPENAI_API_KEY: 'openai-test', OPENROUTER_API_KEY: 'or-test' })
|
||||
fetchMock.mockResolvedValue(jsonResponse(openAIBody([[1, 2]])))
|
||||
|
||||
await embedKnowledgeForDeployment(['hello'], options, true)
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledOnce()
|
||||
expect(fetchMock.mock.calls[0][0]).toBe('https://api.openai.com/v1/embeddings')
|
||||
})
|
||||
|
||||
it('does not fall back after a fatal provider error', async () => {
|
||||
setEnv({ OPENAI_API_KEY: 'openai-test', OPENROUTER_API_KEY: 'or-test' })
|
||||
fetchMock.mockResolvedValue(jsonResponse({ error: 'invalid key' }, 401))
|
||||
|
||||
await expect(embedKnowledgeForDeployment(['hello'], options, false)).rejects.toThrow(
|
||||
/Embedding API failed: 401/
|
||||
)
|
||||
expect(fetchMock).toHaveBeenCalledOnce()
|
||||
})
|
||||
|
||||
it('falls back after transient retries and projects inputs only once', async () => {
|
||||
vi.useFakeTimers()
|
||||
setEnv({ OPENAI_API_KEY: 'openai-test', OPENROUTER_API_KEY: 'or-test' })
|
||||
const projectInputs = vi.fn(() => ['projected'])
|
||||
fetchMock.mockImplementation(async (url) =>
|
||||
url === 'https://api.openai.com/v1/embeddings'
|
||||
? jsonResponse({ error: 'unavailable' }, 503)
|
||||
: jsonResponse(openAIBody([[7, 8]], 2))
|
||||
)
|
||||
|
||||
const pending = embedKnowledgeForDeployment(['secret'], { ...options, projectInputs }, false)
|
||||
await vi.runAllTimersAsync()
|
||||
const result = await pending
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledTimes(5)
|
||||
expect(fetchMock.mock.calls.slice(0, 4).every(([url]) => url.includes('api.openai.com'))).toBe(
|
||||
true
|
||||
)
|
||||
expect(fetchMock.mock.calls[4][0]).toBe('https://openrouter.ai/api/v1/embeddings')
|
||||
expect(projectInputs).toHaveBeenCalledOnce()
|
||||
expect(result.embeddings).toEqual([[7, 8]])
|
||||
})
|
||||
|
||||
it('falls back only the failed batch and retains successful provider work', async () => {
|
||||
vi.useFakeTimers()
|
||||
setEnv({ OPENAI_API_KEY: 'openai-test', OPENROUTER_API_KEY: 'or-test' })
|
||||
mockGetBYOKKey.mockResolvedValue({ apiKey: 'workspace-openai-test', isBYOK: true })
|
||||
const firstInput = `first ${'word '.repeat(5000)}`
|
||||
const secondInput = `second ${'word '.repeat(5000)}`
|
||||
fetchMock.mockImplementation(async (url, init) => {
|
||||
const body = JSON.parse((init as RequestInit).body as string)
|
||||
const input = body.input[0] as string
|
||||
if (url === 'https://api.openai.com/v1/embeddings' && input.startsWith('second')) {
|
||||
return jsonResponse({ error: 'unavailable' }, 503)
|
||||
}
|
||||
return jsonResponse(openAIBody([[input.startsWith('first') ? 1 : 2]], 3))
|
||||
})
|
||||
|
||||
const pending = embedKnowledgeForDeployment(
|
||||
[firstInput, secondInput],
|
||||
{ ...options, workspaceId: 'workspace-1' },
|
||||
false
|
||||
)
|
||||
await vi.runAllTimersAsync()
|
||||
const result = await pending
|
||||
|
||||
const openRouterInputs = fetchMock.mock.calls
|
||||
.filter(([url]) => url === 'https://openrouter.ai/api/v1/embeddings')
|
||||
.flatMap(([, init]) => JSON.parse((init as RequestInit).body as string).input as string[])
|
||||
expect(openRouterInputs).toEqual([secondInput])
|
||||
expect(fetchMock).toHaveBeenCalledTimes(6)
|
||||
expect(result.embeddings).toEqual([[1], [2]])
|
||||
expect(result.totalTokens).toBe(6)
|
||||
expect(result.billableTokens).toBe(3)
|
||||
expect(result.isBYOK).toBe(false)
|
||||
})
|
||||
|
||||
it('classifies only transient embedding failures for failover', () => {
|
||||
expect(isTransientEmbeddingError(new EmbeddingAPIError('unavailable', 503))).toBe(true)
|
||||
expect(isTransientEmbeddingError(new EmbeddingAPIError('rate limited', 429))).toBe(true)
|
||||
expect(isTransientEmbeddingError(new EmbeddingAPIError('invalid key', 401))).toBe(false)
|
||||
expect(isTransientEmbeddingError(new DOMException('timed out', 'AbortError'))).toBe(true)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,6 +1,14 @@
|
||||
import { createLogger } from '@sim/logger'
|
||||
import { chunkArray } from '@sim/utils/helpers'
|
||||
import { getBYOKKey } from '@/lib/api-key/byok'
|
||||
import { getRotatingApiKey } from '@/lib/core/config/api-keys'
|
||||
import { env, envNumber } from '@/lib/core/config/env'
|
||||
import {
|
||||
type FallbackFactories,
|
||||
KNOWLEDGE_EMBEDDINGS_CAPABILITY,
|
||||
wireFallback,
|
||||
} from '@/lib/core/config/env-capabilities'
|
||||
import { isHosted } from '@/lib/core/config/env-flags'
|
||||
import { mapWithConcurrency } from '@/lib/core/utils/concurrency'
|
||||
import {
|
||||
DEFAULT_EMBEDDING_MODEL,
|
||||
@@ -10,12 +18,14 @@ import {
|
||||
resolveDimensions,
|
||||
} from '@/lib/embeddings/catalog'
|
||||
import { resolveProviderKey } from '@/lib/embeddings/keys'
|
||||
import { DEFAULT_OPENROUTER_EMBEDDING_MODEL } from '@/lib/embeddings/openrouter-models'
|
||||
import { getAdapterFactory } from '@/lib/embeddings/providers'
|
||||
import type {
|
||||
EmbeddingProviderAdapter,
|
||||
EmbeddingTaskType,
|
||||
EmbedOptions,
|
||||
EmbedResult,
|
||||
OpenRouterEmbedOptions,
|
||||
} from '@/lib/embeddings/types'
|
||||
import { isRetryableError, retryWithExponentialBackoff } from '@/lib/knowledge/documents/utils'
|
||||
import { batchByTokenLimit, estimateTokenCount, truncateToTokenLimit } from '@/lib/tokenization'
|
||||
@@ -47,6 +57,14 @@ export class EmbeddingAPIError extends Error {
|
||||
}
|
||||
}
|
||||
|
||||
export function isTransientEmbeddingError(error: unknown): boolean {
|
||||
if (error instanceof EmbeddingAPIError) {
|
||||
return error.status === 429 || error.status >= 500
|
||||
}
|
||||
if (error instanceof Error && error.name === 'AbortError') return true
|
||||
return isRetryableError(error)
|
||||
}
|
||||
|
||||
interface ResolvedProvider {
|
||||
adapter: EmbeddingProviderAdapter
|
||||
info: EmbeddingModelInfo
|
||||
@@ -80,6 +98,26 @@ async function resolveProvider(model: string, options: EmbedOptions): Promise<Re
|
||||
const info = getEmbeddingModelInfo(model)
|
||||
const dimensions = resolveDimensions(info, options.dimensions)
|
||||
|
||||
if (options.transport === 'openrouter') {
|
||||
if (info.provider !== 'openai') {
|
||||
throw new Error(`OpenRouter transport does not support catalog provider: ${info.provider}`)
|
||||
}
|
||||
if (!options.apiKey) {
|
||||
throw new Error('OPENROUTER_API_KEY is not configured')
|
||||
}
|
||||
return {
|
||||
adapter: getAdapterFactory('openrouter')({
|
||||
modelName: model,
|
||||
apiKey: options.apiKey,
|
||||
nativeDimensions: info.nativeDimensions,
|
||||
}),
|
||||
info,
|
||||
modelName: model,
|
||||
dimensions,
|
||||
isBYOK: true,
|
||||
}
|
||||
}
|
||||
|
||||
if (!options.apiKey) {
|
||||
const azure = resolveAzureOverride(info, model)
|
||||
if (azure) {
|
||||
@@ -116,10 +154,11 @@ async function resolveProvider(model: string, options: EmbedOptions): Promise<Re
|
||||
}
|
||||
}
|
||||
|
||||
/** `inputs` are already projected and batched by {@link embed}. */
|
||||
/** `inputs` are already projected and batched by the embedding orchestrator. */
|
||||
async function callEmbeddingAPI(
|
||||
inputs: string[],
|
||||
provider: ResolvedProvider,
|
||||
adapter: EmbeddingProviderAdapter,
|
||||
tokenizerProvider: string,
|
||||
taskType: EmbeddingTaskType,
|
||||
/**
|
||||
* The caller's explicit reduction, or undefined when none was requested. Kept
|
||||
@@ -131,7 +170,7 @@ async function callEmbeddingAPI(
|
||||
): Promise<{ embeddings: number[][]; totalTokens: number }> {
|
||||
return retryWithExponentialBackoff(
|
||||
async () => {
|
||||
const request = provider.adapter.buildRequest({
|
||||
const request = adapter.buildRequest({
|
||||
inputs,
|
||||
taskType,
|
||||
dimensions: requestedDimensions,
|
||||
@@ -164,10 +203,7 @@ async function callEmbeddingAPI(
|
||||
*/
|
||||
const totalTokens =
|
||||
request.parseTokens?.(json) ??
|
||||
inputs.reduce(
|
||||
(sum, text) => sum + estimateTokenCount(text, provider.info.tokenizerProvider).count,
|
||||
0
|
||||
)
|
||||
inputs.reduce((sum, text) => sum + estimateTokenCount(text, tokenizerProvider).count, 0)
|
||||
|
||||
return { embeddings, totalTokens }
|
||||
},
|
||||
@@ -175,25 +211,33 @@ async function callEmbeddingAPI(
|
||||
maxRetries: 3,
|
||||
initialDelayMs: 1000,
|
||||
maxDelayMs: 10000,
|
||||
retryCondition: (error: unknown) => {
|
||||
if (error instanceof EmbeddingAPIError) {
|
||||
return error.status === 429 || error.status >= 500
|
||||
}
|
||||
return isRetryableError(error)
|
||||
},
|
||||
retryCondition: isTransientEmbeddingError,
|
||||
}
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* Generates embeddings for a batch of texts with token-aware batching,
|
||||
* per-provider item caps, bounded concurrency, and retry on transient failures.
|
||||
*/
|
||||
export async function embed(texts: string[], options: EmbedOptions): Promise<EmbedResult> {
|
||||
const model = options.model ?? DEFAULT_EMBEDDING_MODEL
|
||||
const taskType = options.taskType ?? 'document'
|
||||
const provider = await resolveProvider(model, options)
|
||||
interface EmbeddingInputLimits {
|
||||
maxInputTokens: number
|
||||
maxTokensPerRequest?: number
|
||||
tokenizerProvider: string
|
||||
approximateTokenCount: boolean
|
||||
}
|
||||
|
||||
function getEmbeddingInputLimits(info: EmbeddingModelInfo): EmbeddingInputLimits {
|
||||
return {
|
||||
maxInputTokens: info.maxInputTokens,
|
||||
maxTokensPerRequest: info.maxTokensPerRequest,
|
||||
tokenizerProvider: info.tokenizerProvider,
|
||||
approximateTokenCount: hasApproximateTokenCount(info),
|
||||
}
|
||||
}
|
||||
|
||||
function prepareEmbeddingInputs(
|
||||
texts: string[],
|
||||
model: string,
|
||||
limits: EmbeddingInputLimits,
|
||||
projectInputs: EmbedOptions['projectInputs']
|
||||
): string[] {
|
||||
/**
|
||||
* Projected before batching, not after. The projector rewrites resolved-secret
|
||||
* plaintext to placeholders, which changes length, and `batchByTokenLimit`
|
||||
@@ -205,7 +249,7 @@ export async function embed(texts: string[], options: EmbedOptions): Promise<Emb
|
||||
* Doing it here also keeps projection to exactly once per call, so no retry
|
||||
* can re-project already-projected content.
|
||||
*/
|
||||
const modelInputs = options.projectInputs ? options.projectInputs(texts) : texts
|
||||
const modelInputs = projectInputs ? projectInputs(texts) : texts
|
||||
|
||||
/**
|
||||
* Each input is held to the model's own per-input ceiling, exactly as declared.
|
||||
@@ -218,18 +262,75 @@ export async function embed(texts: string[], options: EmbedOptions): Promise<Emb
|
||||
* embedding input is otherwise indistinguishable from a good one, both to the
|
||||
* caller and in the vector it produces.
|
||||
*/
|
||||
const ceiling = provider.info.maxInputTokens
|
||||
const ceiling = limits.maxInputTokens
|
||||
const boundedInputs = modelInputs.map((text) => {
|
||||
if (estimateTokenCount(text, provider.info.tokenizerProvider).count <= ceiling) return text
|
||||
if (estimateTokenCount(text, limits.tokenizerProvider).count <= ceiling) return text
|
||||
logger.warn('Embedding input exceeds the model token limit and will be truncated', {
|
||||
model,
|
||||
maxInputTokens: ceiling,
|
||||
chars: text.length,
|
||||
approximateTokenCount: hasApproximateTokenCount(provider.info),
|
||||
approximateTokenCount: limits.approximateTokenCount,
|
||||
})
|
||||
return truncateToTokenLimit(text, ceiling, model)
|
||||
})
|
||||
|
||||
return boundedInputs
|
||||
}
|
||||
|
||||
async function embedWithProvider(
|
||||
boundedInputs: string[],
|
||||
model: string,
|
||||
taskType: EmbeddingTaskType,
|
||||
requestedDimensions: number | undefined,
|
||||
provider: ResolvedProvider
|
||||
): Promise<EmbedResult> {
|
||||
const batches = createEmbeddingBatches(
|
||||
boundedInputs,
|
||||
model,
|
||||
getEmbeddingInputLimits(provider.info),
|
||||
provider.adapter.maxItemsPerRequest
|
||||
)
|
||||
|
||||
const batchResults = await mapWithConcurrency(
|
||||
batches,
|
||||
MAX_CONCURRENT_BATCHES,
|
||||
async (batch, i) => {
|
||||
try {
|
||||
return await callEmbeddingAPI(
|
||||
batch,
|
||||
provider.adapter,
|
||||
provider.info.tokenizerProvider,
|
||||
taskType,
|
||||
requestedDimensions
|
||||
)
|
||||
} catch (error) {
|
||||
logger.error(`Failed to generate embeddings for batch ${i + 1}/${batches.length}:`, error)
|
||||
throw error
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
const { embeddings, totalTokens } = combineEmbeddingBatches(batchResults)
|
||||
|
||||
return {
|
||||
embeddings,
|
||||
totalTokens,
|
||||
billableTokens: provider.isBYOK ? 0 : totalTokens,
|
||||
isBYOK: provider.isBYOK,
|
||||
modelName: provider.modelName,
|
||||
pricingId: provider.info.pricingId,
|
||||
dimensions: provider.dimensions,
|
||||
}
|
||||
}
|
||||
|
||||
function createEmbeddingBatches(
|
||||
boundedInputs: string[],
|
||||
model: string,
|
||||
limits: Pick<EmbeddingInputLimits, 'maxInputTokens' | 'maxTokensPerRequest'>,
|
||||
itemLimit: number | undefined
|
||||
): string[][] {
|
||||
const ceiling = limits.maxInputTokens
|
||||
|
||||
/**
|
||||
* How many tokens may share one request — a different limit from the per-input
|
||||
* ceiling above, and the one that decides how many inputs go in a batch.
|
||||
@@ -245,29 +346,17 @@ export async function embed(texts: string[], options: EmbedOptions): Promise<Emb
|
||||
* Cohere takes 128k tokens in one text, far above the target.
|
||||
*/
|
||||
const requestBudget = Math.max(
|
||||
Math.min(provider.info.maxTokensPerRequest ?? BATCH_TOKEN_TARGET, BATCH_TOKEN_TARGET),
|
||||
Math.min(limits.maxTokensPerRequest ?? BATCH_TOKEN_TARGET, BATCH_TOKEN_TARGET),
|
||||
ceiling
|
||||
)
|
||||
|
||||
const tokenBatches = batchByTokenLimit(boundedInputs, requestBudget, model)
|
||||
const itemLimit = provider.adapter.maxItemsPerRequest
|
||||
const batches = itemLimit
|
||||
? tokenBatches.flatMap((batch) => chunkArray(batch, itemLimit))
|
||||
: tokenBatches
|
||||
|
||||
const batchResults = await mapWithConcurrency(
|
||||
batches,
|
||||
MAX_CONCURRENT_BATCHES,
|
||||
async (batch, i) => {
|
||||
try {
|
||||
return await callEmbeddingAPI(batch, provider, taskType, options.dimensions)
|
||||
} catch (error) {
|
||||
logger.error(`Failed to generate embeddings for batch ${i + 1}/${batches.length}:`, error)
|
||||
throw error
|
||||
}
|
||||
}
|
||||
)
|
||||
return itemLimit ? tokenBatches.flatMap((batch) => chunkArray(batch, itemLimit)) : tokenBatches
|
||||
}
|
||||
|
||||
function combineEmbeddingBatches(
|
||||
batchResults: readonly { embeddings: number[][]; totalTokens: number }[]
|
||||
): { embeddings: number[][]; totalTokens: number } {
|
||||
const embeddings: number[][] = []
|
||||
let totalTokens = 0
|
||||
for (const batch of batchResults) {
|
||||
@@ -276,13 +365,228 @@ export async function embed(texts: string[], options: EmbedOptions): Promise<Emb
|
||||
}
|
||||
totalTokens += batch.totalTokens
|
||||
}
|
||||
return { embeddings, totalTokens }
|
||||
}
|
||||
|
||||
/**
|
||||
* Generates embeddings for a batch of texts with token-aware batching,
|
||||
* per-provider item caps, bounded concurrency, and retry on transient failures.
|
||||
*/
|
||||
export async function embed(texts: string[], options: EmbedOptions): Promise<EmbedResult> {
|
||||
const model = options.model ?? DEFAULT_EMBEDDING_MODEL
|
||||
const taskType = options.taskType ?? 'document'
|
||||
const provider = await resolveProvider(model, options)
|
||||
const boundedInputs = prepareEmbeddingInputs(
|
||||
texts,
|
||||
model,
|
||||
getEmbeddingInputLimits(provider.info),
|
||||
options.projectInputs
|
||||
)
|
||||
return embedWithProvider(boundedInputs, model, taskType, options.dimensions, provider)
|
||||
}
|
||||
|
||||
/** Generates embeddings for any model returned by OpenRouter's embedding catalog. */
|
||||
export async function embedOpenRouter(
|
||||
texts: string[],
|
||||
options: OpenRouterEmbedOptions
|
||||
): Promise<EmbedResult> {
|
||||
if (texts.length === 0) throw new Error('At least one embedding input is required')
|
||||
if (!options.apiKey) throw new Error('OpenRouter API key is required')
|
||||
if (!Number.isInteger(options.maxInputTokens) || options.maxInputTokens <= 0) {
|
||||
throw new Error('OpenRouter max input tokens must be a positive integer')
|
||||
}
|
||||
|
||||
const model = options.model ?? DEFAULT_OPENROUTER_EMBEDDING_MODEL
|
||||
const limits: EmbeddingInputLimits = {
|
||||
maxInputTokens: options.maxInputTokens,
|
||||
tokenizerProvider: 'openrouter',
|
||||
approximateTokenCount: true,
|
||||
}
|
||||
const boundedInputs = prepareEmbeddingInputs(texts, model, limits, options.projectInputs)
|
||||
const adapter = getAdapterFactory('openrouter')({
|
||||
modelName: model,
|
||||
apiKey: options.apiKey,
|
||||
nativeDimensions: options.dimensions ?? 0,
|
||||
})
|
||||
const batches = createEmbeddingBatches(boundedInputs, model, limits, adapter.maxItemsPerRequest)
|
||||
const batchResults = await mapWithConcurrency(batches, MAX_CONCURRENT_BATCHES, async (batch) =>
|
||||
callEmbeddingAPI(batch, adapter, limits.tokenizerProvider, 'document', options.dimensions)
|
||||
)
|
||||
const result = combineEmbeddingBatches(batchResults)
|
||||
|
||||
if (result.embeddings.length !== boundedInputs.length) {
|
||||
throw new Error(
|
||||
`OpenRouter returned ${result.embeddings.length} embeddings for ${boundedInputs.length} inputs`
|
||||
)
|
||||
}
|
||||
const dimensions = result.embeddings[0]?.length
|
||||
if (!dimensions) throw new Error('OpenRouter returned an empty embedding vector')
|
||||
if (result.embeddings.some((embedding) => embedding.length !== dimensions)) {
|
||||
throw new Error('OpenRouter returned embedding vectors with inconsistent dimensions')
|
||||
}
|
||||
if (options.dimensions !== undefined && dimensions !== options.dimensions) {
|
||||
throw new Error(
|
||||
`OpenRouter returned ${dimensions} dimensions instead of the requested ${options.dimensions}`
|
||||
)
|
||||
}
|
||||
|
||||
return {
|
||||
embeddings: result.embeddings,
|
||||
totalTokens: result.totalTokens,
|
||||
billableTokens: 0,
|
||||
isBYOK: true,
|
||||
modelName: model,
|
||||
pricingId: model,
|
||||
dimensions,
|
||||
}
|
||||
}
|
||||
|
||||
type KnowledgeEmbedOptions = Omit<EmbedOptions, 'apiKey' | 'transport'>
|
||||
|
||||
function resolveEnvironmentOpenAIKey(): string {
|
||||
if (env.OPENAI_API_KEY) return env.OPENAI_API_KEY
|
||||
return getRotatingApiKey('openai')
|
||||
}
|
||||
|
||||
/** @internal Exported for deterministic hosted/self-hosted routing tests. */
|
||||
export async function embedKnowledgeForDeployment(
|
||||
texts: string[],
|
||||
options: KnowledgeEmbedOptions,
|
||||
hosted: boolean
|
||||
): Promise<EmbedResult> {
|
||||
const model = options.model ?? DEFAULT_EMBEDDING_MODEL
|
||||
const info = getEmbeddingModelInfo(model)
|
||||
if (hosted || !env.OPENROUTER_API_KEY || info.provider !== 'openai') {
|
||||
return embed(texts, options)
|
||||
}
|
||||
|
||||
const dimensions = resolveDimensions(info, options.dimensions)
|
||||
const taskType = options.taskType ?? 'document'
|
||||
const boundedInputs = prepareEmbeddingInputs(
|
||||
texts,
|
||||
model,
|
||||
getEmbeddingInputLimits(info),
|
||||
options.projectInputs
|
||||
)
|
||||
const workspaceKey = options.workspaceId ? await getBYOKKey(options.workspaceId, 'openai') : null
|
||||
const capabilityValues = workspaceKey ? { ...env, OPENAI_API_KEY: workspaceKey.apiKey } : env
|
||||
|
||||
const factories = {
|
||||
'azure-openai': () => {
|
||||
const azure = resolveAzureOverride(info, model)
|
||||
if (!azure) return null
|
||||
return {
|
||||
adapter: getAdapterFactory('azure-openai')({
|
||||
modelName: azure.deployment,
|
||||
apiKey: azure.apiKey,
|
||||
nativeDimensions: info.nativeDimensions,
|
||||
endpoint: azure.endpoint,
|
||||
apiVersion: azure.apiVersion,
|
||||
}),
|
||||
info,
|
||||
modelName: azure.deployment,
|
||||
dimensions,
|
||||
isBYOK: false,
|
||||
}
|
||||
},
|
||||
openai: () => {
|
||||
const apiKey = workspaceKey?.apiKey ?? resolveEnvironmentOpenAIKey()
|
||||
return {
|
||||
adapter: getAdapterFactory('openai')({
|
||||
modelName: model,
|
||||
apiKey,
|
||||
nativeDimensions: info.nativeDimensions,
|
||||
}),
|
||||
info,
|
||||
modelName: model,
|
||||
dimensions,
|
||||
isBYOK: Boolean(workspaceKey),
|
||||
}
|
||||
},
|
||||
openrouter: () => {
|
||||
if (!env.OPENROUTER_API_KEY) return null
|
||||
return {
|
||||
adapter: getAdapterFactory('openrouter')({
|
||||
modelName: model,
|
||||
apiKey: env.OPENROUTER_API_KEY,
|
||||
nativeDimensions: info.nativeDimensions,
|
||||
}),
|
||||
info,
|
||||
modelName: model,
|
||||
dimensions,
|
||||
isBYOK: false,
|
||||
}
|
||||
},
|
||||
} satisfies FallbackFactories<typeof KNOWLEDGE_EMBEDDINGS_CAPABILITY, ResolvedProvider>
|
||||
|
||||
const fallback = wireFallback<typeof KNOWLEDGE_EMBEDDINGS_CAPABILITY, ResolvedProvider>({
|
||||
definition: KNOWLEDGE_EMBEDDINGS_CAPABILITY,
|
||||
values: capabilityValues,
|
||||
factories,
|
||||
shouldFallback: isTransientEmbeddingError,
|
||||
onFailure(providerId, error) {
|
||||
logger.warn('Knowledge embedding provider failed; continuing fallback chain', {
|
||||
providerId,
|
||||
error,
|
||||
})
|
||||
},
|
||||
})
|
||||
|
||||
const itemLimits = fallback.providers.flatMap((provider) =>
|
||||
provider.adapter.maxItemsPerRequest ? [provider.adapter.maxItemsPerRequest] : []
|
||||
)
|
||||
const batches = createEmbeddingBatches(
|
||||
boundedInputs,
|
||||
model,
|
||||
info,
|
||||
itemLimits.length > 0 ? Math.min(...itemLimits) : undefined
|
||||
)
|
||||
const batchResults = await mapWithConcurrency(
|
||||
batches,
|
||||
MAX_CONCURRENT_BATCHES,
|
||||
async (batch, i) => {
|
||||
try {
|
||||
return await fallback.execute(async (provider) => ({
|
||||
...(await callEmbeddingAPI(
|
||||
batch,
|
||||
provider.adapter,
|
||||
provider.info.tokenizerProvider,
|
||||
taskType,
|
||||
options.dimensions
|
||||
)),
|
||||
provider,
|
||||
}))
|
||||
} catch (error) {
|
||||
logger.error(`Failed to generate embeddings for batch ${i + 1}/${batches.length}:`, error)
|
||||
throw error
|
||||
}
|
||||
}
|
||||
)
|
||||
const { embeddings, totalTokens } = combineEmbeddingBatches(batchResults)
|
||||
const defaultProvider = fallback.providers[0]
|
||||
const usedProviders = batchResults.map((batch) => batch.provider)
|
||||
const metadataProvider = usedProviders[0] ?? defaultProvider
|
||||
const modelNames = new Set(usedProviders.map((provider) => provider.modelName))
|
||||
const billableTokens = batchResults.reduce(
|
||||
(sum, batch) => sum + (batch.provider.isBYOK ? 0 : batch.totalTokens),
|
||||
0
|
||||
)
|
||||
|
||||
return {
|
||||
embeddings,
|
||||
totalTokens,
|
||||
isBYOK: provider.isBYOK,
|
||||
modelName: provider.modelName,
|
||||
pricingId: provider.info.pricingId,
|
||||
dimensions: provider.dimensions,
|
||||
billableTokens,
|
||||
isBYOK: usedProviders.length > 0 ? billableTokens === 0 : metadataProvider.isBYOK,
|
||||
modelName: modelNames.size > 1 ? model : metadataProvider.modelName,
|
||||
pricingId: info.pricingId,
|
||||
dimensions,
|
||||
}
|
||||
}
|
||||
|
||||
/** Generates KB document/query embeddings with opt-in self-hosted OpenRouter fallback. */
|
||||
export async function embedKnowledge(
|
||||
texts: string[],
|
||||
options: KnowledgeEmbedOptions
|
||||
): Promise<EmbedResult> {
|
||||
return embedKnowledgeForDeployment(texts, options, isHosted)
|
||||
}
|
||||
|
||||
@@ -9,5 +9,11 @@ export {
|
||||
findEmbeddingModelInfo,
|
||||
resolveDimensions,
|
||||
} from '@/lib/embeddings/catalog'
|
||||
export { embed } from '@/lib/embeddings/client'
|
||||
export type { EmbeddingTaskType, EmbedOptions, EmbedResult } from '@/lib/embeddings/types'
|
||||
export { embed, embedKnowledge, embedOpenRouter } from '@/lib/embeddings/client'
|
||||
export { DEFAULT_OPENROUTER_EMBEDDING_MODEL } from '@/lib/embeddings/openrouter-models'
|
||||
export type {
|
||||
EmbeddingTaskType,
|
||||
EmbedOptions,
|
||||
EmbedResult,
|
||||
OpenRouterEmbedOptions,
|
||||
} from '@/lib/embeddings/types'
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
/**
|
||||
* @vitest-environment node
|
||||
*/
|
||||
import { afterAll, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import {
|
||||
getOpenRouterEmbeddingModelMetadata,
|
||||
OpenRouterEmbeddingModelNotFoundError,
|
||||
} from '@/lib/embeddings/openrouter-model-catalog.server'
|
||||
|
||||
const fetchMock = vi.fn()
|
||||
|
||||
describe('OpenRouter embedding model catalog', () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
vi.stubGlobal('fetch', fetchMock)
|
||||
})
|
||||
|
||||
afterAll(() => {
|
||||
vi.unstubAllGlobals()
|
||||
})
|
||||
|
||||
it('resolves a prefixed model with its live input ceiling', async () => {
|
||||
fetchMock.mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => ({
|
||||
data: [{ id: 'qwen/qwen3-embedding-8b', context_length: 32768 }],
|
||||
}),
|
||||
})
|
||||
|
||||
await expect(
|
||||
getOpenRouterEmbeddingModelMetadata('openrouter/qwen/qwen3-embedding-8b')
|
||||
).resolves.toEqual({
|
||||
id: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
maxInputTokens: 32768,
|
||||
})
|
||||
})
|
||||
|
||||
it('rejects a model absent from the live embedding catalog', async () => {
|
||||
fetchMock.mockResolvedValue({ ok: true, json: async () => ({ data: [] }) })
|
||||
|
||||
await expect(
|
||||
getOpenRouterEmbeddingModelMetadata('openrouter/example/missing')
|
||||
).rejects.toBeInstanceOf(OpenRouterEmbeddingModelNotFoundError)
|
||||
})
|
||||
|
||||
it('fails fast when OpenRouter omits a model context length', async () => {
|
||||
fetchMock.mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => ({ data: [{ id: 'example/missing-context' }] }),
|
||||
})
|
||||
|
||||
await expect(
|
||||
getOpenRouterEmbeddingModelMetadata('openrouter/example/missing-context')
|
||||
).rejects.toThrow('Invalid input')
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,54 @@
|
||||
import { openRouterEmbeddingModelsUpstreamResponseSchema } from '@/lib/api/contracts/providers'
|
||||
import {
|
||||
toOpenRouterEmbeddingModelId,
|
||||
toOpenRouterWireEmbeddingModelId,
|
||||
} from '@/lib/embeddings/openrouter-models'
|
||||
|
||||
const OPENROUTER_EMBEDDING_MODELS_URL = 'https://openrouter.ai/api/v1/embeddings/models'
|
||||
|
||||
export interface OpenRouterEmbeddingModelMetadata {
|
||||
id: string
|
||||
maxInputTokens: number
|
||||
}
|
||||
|
||||
export class OpenRouterEmbeddingModelNotFoundError extends Error {
|
||||
constructor(model: string) {
|
||||
super(`Unsupported OpenRouter embedding model: ${model}`)
|
||||
this.name = 'OpenRouterEmbeddingModelNotFoundError'
|
||||
}
|
||||
}
|
||||
|
||||
/** Loads OpenRouter's current embedding-only catalog with its input ceilings. */
|
||||
export async function fetchOpenRouterEmbeddingModelCatalog(): Promise<
|
||||
OpenRouterEmbeddingModelMetadata[]
|
||||
> {
|
||||
const response = await fetch(OPENROUTER_EMBEDDING_MODELS_URL, {
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
next: { revalidate: 300 },
|
||||
})
|
||||
if (!response.ok) {
|
||||
throw new Error(
|
||||
`Failed to fetch OpenRouter embedding models: ${response.status} ${response.statusText}`
|
||||
)
|
||||
}
|
||||
|
||||
const data = openRouterEmbeddingModelsUpstreamResponseSchema.parse(await response.json())
|
||||
const models = new Map<string, OpenRouterEmbeddingModelMetadata>()
|
||||
for (const model of data.data) {
|
||||
const id = toOpenRouterEmbeddingModelId(model.id)
|
||||
models.set(id, { id, maxInputTokens: model.context_length })
|
||||
}
|
||||
return Array.from(models.values())
|
||||
}
|
||||
|
||||
/** Resolves and validates one selected model against OpenRouter's live catalog. */
|
||||
export async function getOpenRouterEmbeddingModelMetadata(
|
||||
model: string
|
||||
): Promise<OpenRouterEmbeddingModelMetadata> {
|
||||
const normalizedId = toOpenRouterEmbeddingModelId(toOpenRouterWireEmbeddingModelId(model))
|
||||
const metadata = (await fetchOpenRouterEmbeddingModelCatalog()).find(
|
||||
(candidate) => candidate.id === normalizedId
|
||||
)
|
||||
if (!metadata) throw new OpenRouterEmbeddingModelNotFoundError(model)
|
||||
return metadata
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
const OPENROUTER_MODEL_PREFIX = 'openrouter/'
|
||||
|
||||
export const DEFAULT_OPENROUTER_EMBEDDING_MODEL = 'openrouter/openai/text-embedding-3-small'
|
||||
|
||||
const LEGACY_OPENAI_EMBEDDING_MODELS = new Set([
|
||||
'text-embedding-3-small',
|
||||
'text-embedding-3-large',
|
||||
'text-embedding-ada-002',
|
||||
])
|
||||
|
||||
function assertUpstreamModelId(model: string): string {
|
||||
if (!/^[^/\s]+\/[^\s]+$/.test(model)) {
|
||||
throw new Error(`Invalid OpenRouter embedding model: ${model}`)
|
||||
}
|
||||
return model
|
||||
}
|
||||
|
||||
/** Converts an OpenRouter upstream model id into Sim's provider-prefixed id. */
|
||||
export function toOpenRouterEmbeddingModelId(upstreamModel: string): string {
|
||||
return `${OPENROUTER_MODEL_PREFIX}${assertUpstreamModelId(upstreamModel)}`
|
||||
}
|
||||
|
||||
/** Converts Sim's model id into the provider-qualified id OpenRouter accepts. */
|
||||
export function toOpenRouterWireEmbeddingModelId(model: string): string {
|
||||
if (model.startsWith(OPENROUTER_MODEL_PREFIX)) {
|
||||
return assertUpstreamModelId(model.slice(OPENROUTER_MODEL_PREFIX.length))
|
||||
}
|
||||
if (LEGACY_OPENAI_EMBEDDING_MODELS.has(model)) {
|
||||
return `openai/${model}`
|
||||
}
|
||||
throw new Error(`Invalid OpenRouter embedding model: ${model}`)
|
||||
}
|
||||
|
||||
/** Normalizes persisted OpenAI-only values to the dynamic OpenRouter id format. */
|
||||
export function normalizeOpenRouterEmbeddingModelId(model: string): string {
|
||||
if (model.startsWith(OPENROUTER_MODEL_PREFIX)) {
|
||||
toOpenRouterWireEmbeddingModelId(model)
|
||||
return model
|
||||
}
|
||||
if (LEGACY_OPENAI_EMBEDDING_MODELS.has(model)) {
|
||||
return `${OPENROUTER_MODEL_PREFIX}openai/${model}`
|
||||
}
|
||||
throw new Error(`Invalid OpenRouter embedding model: ${model}`)
|
||||
}
|
||||
@@ -3,6 +3,7 @@ import { createCohereAdapter } from '@/lib/embeddings/providers/cohere'
|
||||
import { createGeminiAdapter } from '@/lib/embeddings/providers/gemini'
|
||||
import { createMistralAdapter } from '@/lib/embeddings/providers/mistral'
|
||||
import { createOpenAIAdapter } from '@/lib/embeddings/providers/openai'
|
||||
import { createOpenRouterAdapter } from '@/lib/embeddings/providers/openrouter'
|
||||
import type {
|
||||
AzureEmbeddingAdapterContext,
|
||||
EmbeddingAdapterFactory,
|
||||
@@ -17,6 +18,7 @@ type AdapterFactoryFor<K extends EmbeddingProviderKind> = K extends 'azure-opena
|
||||
const ADAPTER_FACTORIES: { [K in EmbeddingProviderKind]: AdapterFactoryFor<K> } = {
|
||||
openai: createOpenAIAdapter,
|
||||
'azure-openai': createAzureOpenAIAdapter,
|
||||
openrouter: createOpenRouterAdapter,
|
||||
gemini: createGeminiAdapter,
|
||||
cohere: createCohereAdapter,
|
||||
mistral: createMistralAdapter,
|
||||
@@ -34,4 +36,5 @@ export {
|
||||
createGeminiAdapter,
|
||||
createMistralAdapter,
|
||||
createOpenAIAdapter,
|
||||
createOpenRouterAdapter,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
import { toOpenRouterWireEmbeddingModelId } from '@/lib/embeddings/openrouter-models'
|
||||
import {
|
||||
OPENAI_MAX_ITEMS_PER_REQUEST,
|
||||
type OpenAIEmbeddingResponse,
|
||||
} from '@/lib/embeddings/providers/openai'
|
||||
import type { EmbeddingAdapterFactory } from '@/lib/embeddings/types'
|
||||
|
||||
/** OpenRouter exposes embedding models under provider-qualified model ids. */
|
||||
export const createOpenRouterAdapter: EmbeddingAdapterFactory = ({ modelName, apiKey }) => ({
|
||||
maxItemsPerRequest: OPENAI_MAX_ITEMS_PER_REQUEST,
|
||||
buildRequest: ({ inputs, dimensions }) => ({
|
||||
apiUrl: 'https://openrouter.ai/api/v1/embeddings',
|
||||
headers: {
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
'Content-Type': 'application/json',
|
||||
},
|
||||
body: {
|
||||
input: inputs,
|
||||
model: toOpenRouterWireEmbeddingModelId(modelName),
|
||||
encoding_format: 'float',
|
||||
...(dimensions !== undefined && { dimensions }),
|
||||
},
|
||||
parse: (json) => (json as OpenAIEmbeddingResponse).data.map((item) => item.embedding),
|
||||
parseTokens: (json) => (json as OpenAIEmbeddingResponse).usage?.total_tokens,
|
||||
}),
|
||||
})
|
||||
@@ -8,6 +8,7 @@ import {
|
||||
createGeminiAdapter,
|
||||
createMistralAdapter,
|
||||
createOpenAIAdapter,
|
||||
createOpenRouterAdapter,
|
||||
} from '@/lib/embeddings/providers'
|
||||
|
||||
const INPUTS = ['alpha', 'beta']
|
||||
@@ -59,6 +60,56 @@ describe('OpenAI adapter', () => {
|
||||
})
|
||||
})
|
||||
|
||||
describe('OpenRouter adapter', () => {
|
||||
const adapter = createOpenRouterAdapter({
|
||||
modelName: 'text-embedding-3-small',
|
||||
apiKey: 'or-test',
|
||||
nativeDimensions: 1536,
|
||||
})
|
||||
|
||||
it('uses the OpenRouter endpoint and provider-qualified OpenAI model id', () => {
|
||||
const request = adapter.buildRequest({
|
||||
inputs: INPUTS,
|
||||
taskType: 'document',
|
||||
dimensions: 1536,
|
||||
})
|
||||
|
||||
expect(request.apiUrl).toBe('https://openrouter.ai/api/v1/embeddings')
|
||||
expect(request.headers.Authorization).toBe('Bearer or-test')
|
||||
expect(request.body).toMatchObject({
|
||||
input: INPUTS,
|
||||
model: 'openai/text-embedding-3-small',
|
||||
dimensions: 1536,
|
||||
encoding_format: 'float',
|
||||
})
|
||||
})
|
||||
|
||||
it('parses OpenAI-compatible vectors and usage', () => {
|
||||
const request = adapter.buildRequest({ inputs: INPUTS, taskType: 'query' })
|
||||
const response = {
|
||||
data: [{ embedding: [1, 2] }, { embedding: [3, 4] }],
|
||||
usage: { total_tokens: 9 },
|
||||
}
|
||||
|
||||
expect(request.parse(response)).toEqual([
|
||||
[1, 2],
|
||||
[3, 4],
|
||||
])
|
||||
expect(request.parseTokens?.(response)).toBe(9)
|
||||
})
|
||||
|
||||
it('passes a dynamic OpenRouter model id through without rewriting its provider', () => {
|
||||
const dynamicAdapter = createOpenRouterAdapter({
|
||||
modelName: 'openrouter/qwen/qwen3-embedding-8b',
|
||||
apiKey: 'or-test',
|
||||
nativeDimensions: 0,
|
||||
})
|
||||
const request = dynamicAdapter.buildRequest({ inputs: INPUTS, taskType: 'document' })
|
||||
|
||||
expect(request.body).toMatchObject({ model: 'qwen/qwen3-embedding-8b' })
|
||||
})
|
||||
})
|
||||
|
||||
describe('Gemini adapter', () => {
|
||||
const adapter = createGeminiAdapter({
|
||||
modelName: 'gemini-embedding-001',
|
||||
|
||||
@@ -4,14 +4,19 @@
|
||||
* `@/lib/embeddings/providers`.
|
||||
*/
|
||||
|
||||
export type EmbeddingProviderKind = 'openai' | 'azure-openai' | 'gemini' | 'cohere' | 'mistral'
|
||||
export type EmbeddingProviderKind =
|
||||
| 'openai'
|
||||
| 'azure-openai'
|
||||
| 'openrouter'
|
||||
| 'gemini'
|
||||
| 'cohere'
|
||||
| 'mistral'
|
||||
|
||||
/**
|
||||
* Providers a catalog model can belong to. Azure OpenAI is excluded because it
|
||||
* is a transport override for OpenAI models rather than a provider users pick:
|
||||
* no model is ever catalogued under it, and it resolves its own credentials.
|
||||
* Providers a catalog model can belong to. Azure OpenAI and OpenRouter are
|
||||
* transports for OpenAI models, so no model is catalogued under either one.
|
||||
*/
|
||||
export type EmbeddingCatalogProvider = Exclude<EmbeddingProviderKind, 'azure-openai'>
|
||||
export type EmbeddingCatalogProvider = Exclude<EmbeddingProviderKind, 'azure-openai' | 'openrouter'>
|
||||
|
||||
/** Provider id for `estimateTokenCount` so token counts match the embedding provider's tokenization. */
|
||||
export type TokenizerProviderId = 'openai' | 'google' | 'cohere' | 'mistral'
|
||||
@@ -75,6 +80,8 @@ export type EmbeddingAdapterFactory<Ctx extends EmbeddingAdapterContext = Embedd
|
||||
export interface EmbedOptions {
|
||||
/** Catalog model id. Defaults to the platform default when omitted. */
|
||||
model?: string
|
||||
/** Transport override for catalog models exposed through another provider. */
|
||||
transport?: 'openrouter'
|
||||
/** Workspace used to look up a BYOK key before falling back to platform keys. */
|
||||
workspaceId?: string | null
|
||||
taskType?: EmbeddingTaskType
|
||||
@@ -101,7 +108,9 @@ export interface EmbedOptions {
|
||||
export interface EmbedResult {
|
||||
embeddings: number[][]
|
||||
totalTokens: number
|
||||
/** True when a workspace-owned key was used, meaning Sim does not bill for it. */
|
||||
/** Tokens processed with a Sim-funded key and therefore eligible for billing. */
|
||||
billableTokens: number
|
||||
/** True when every successful embedding used a caller- or workspace-owned key. */
|
||||
isBYOK: boolean
|
||||
/** Model name as sent to the provider. */
|
||||
modelName: string
|
||||
@@ -110,3 +119,13 @@ export interface EmbedResult {
|
||||
/** Dimensionality of the returned vectors. */
|
||||
dimensions: number
|
||||
}
|
||||
|
||||
export interface OpenRouterEmbedOptions {
|
||||
apiKey: string
|
||||
model?: string
|
||||
/** Per-input ceiling reported by OpenRouter's embedding model catalog. */
|
||||
maxInputTokens: number
|
||||
/** Forwarded when a caller explicitly requests a provider-supported reduction. */
|
||||
dimensions?: number
|
||||
projectInputs: ((values: readonly string[]) => string[]) | null
|
||||
}
|
||||
|
||||
@@ -5926,7 +5926,7 @@
|
||||
"slug": "embeddings",
|
||||
"name": "Embeddings",
|
||||
"description": "Generate embeddings",
|
||||
"longDescription": "Turn text into embedding vectors for semantic search, clustering, and similarity. Supports OpenAI, Google Gemini, Cohere, and Mistral embedding models.",
|
||||
"longDescription": "Turn text into embedding vectors for semantic search, clustering, and similarity. Supports OpenAI, OpenRouter, Google Gemini, Cohere, and Mistral embedding models.",
|
||||
"bgColor": "#7B4DFF",
|
||||
"iconName": "EmbeddingsIcon",
|
||||
"docsUrl": "https://docs.sim.ai/integrations/embeddings",
|
||||
|
||||
@@ -924,8 +924,7 @@ export async function processDocumentAsync(
|
||||
.where(eq(document.id, documentId))
|
||||
return
|
||||
}
|
||||
let totalEmbeddingTokens = 0
|
||||
let embeddingIsBYOK = false
|
||||
let billableEmbeddingTokens = 0
|
||||
let embeddingModelName = kbEmbeddingModel
|
||||
let embeddingPricingId = kbEmbeddingModel
|
||||
|
||||
@@ -992,17 +991,15 @@ export async function processDocumentAsync(
|
||||
logger.info(`[${documentId}] Processing embedding batch ${batchNum}/${totalBatches}`)
|
||||
const {
|
||||
embeddings: batchEmbeddings,
|
||||
totalTokens: batchTokens,
|
||||
isBYOK,
|
||||
billableTokens: batchBillableTokens,
|
||||
modelName,
|
||||
pricingId,
|
||||
} = await generateEmbeddings(batch, kbEmbeddingModel, ctx.workspaceId)
|
||||
for (const emb of batchEmbeddings) {
|
||||
embeddings.push(emb)
|
||||
}
|
||||
totalEmbeddingTokens += batchTokens
|
||||
billableEmbeddingTokens += batchBillableTokens
|
||||
if (i === 0) {
|
||||
embeddingIsBYOK = isBYOK
|
||||
embeddingModelName = modelName
|
||||
embeddingPricingId = pricingId
|
||||
}
|
||||
@@ -1136,12 +1133,12 @@ export async function processDocumentAsync(
|
||||
const processingTime = Date.now() - startTime
|
||||
logger.info(`[${documentId}] Successfully processed document in ${processingTime}ms`)
|
||||
|
||||
if (!embeddingIsBYOK && totalEmbeddingTokens > 0) {
|
||||
if (billableEmbeddingTokens > 0) {
|
||||
try {
|
||||
const costMultiplier = getCostMultiplier()
|
||||
const { total: cost } = calculateCost(
|
||||
embeddingPricingId,
|
||||
totalEmbeddingTokens,
|
||||
billableEmbeddingTokens,
|
||||
0,
|
||||
false,
|
||||
costMultiplier
|
||||
@@ -1158,7 +1155,7 @@ export async function processDocumentAsync(
|
||||
description: embeddingModelName,
|
||||
cost,
|
||||
sourceReference: `knowledge-document:${documentId}:${startTime}`,
|
||||
metadata: { inputTokens: totalEmbeddingTokens, outputTokens: 0 },
|
||||
metadata: { inputTokens: billableEmbeddingTokens, outputTokens: 0 },
|
||||
},
|
||||
],
|
||||
})
|
||||
@@ -1170,7 +1167,7 @@ export async function processDocumentAsync(
|
||||
} else {
|
||||
logger.warn(
|
||||
`[${documentId}] Embedding model "${embeddingModelName}" has no pricing entry — billing skipped`,
|
||||
{ totalEmbeddingTokens, embeddingModelName }
|
||||
{ billableEmbeddingTokens, embeddingModelName }
|
||||
)
|
||||
}
|
||||
} catch (billingError) {
|
||||
|
||||
@@ -7,7 +7,7 @@ import {
|
||||
import { recordUsage } from '@/lib/billing/core/usage-log'
|
||||
import { checkAndBillPayerOverageThreshold } from '@/lib/billing/threshold-billing'
|
||||
import { env } from '@/lib/core/config/env'
|
||||
import { embed } from '@/lib/embeddings'
|
||||
import { embedKnowledge } from '@/lib/embeddings'
|
||||
import {
|
||||
assertKbEmbeddingModel,
|
||||
DEFAULT_EMBEDDING_MODEL,
|
||||
@@ -46,6 +46,7 @@ export function getConfiguredEmbeddingModel(): string {
|
||||
export interface GenerateEmbeddingsResult {
|
||||
embeddings: number[][]
|
||||
totalTokens: number
|
||||
billableTokens: number
|
||||
isBYOK: boolean
|
||||
modelName: string
|
||||
/** Pricing identifier for use with calculateCost / EMBEDDING_MODEL_PRICING. */
|
||||
@@ -65,7 +66,7 @@ export async function generateEmbeddings(
|
||||
): Promise<GenerateEmbeddingsResult> {
|
||||
assertKbEmbeddingModel(embeddingModel)
|
||||
|
||||
const result = await embed(texts, {
|
||||
const result = await embedKnowledge(texts, {
|
||||
model: embeddingModel,
|
||||
workspaceId,
|
||||
taskType: 'document',
|
||||
@@ -76,6 +77,7 @@ export async function generateEmbeddings(
|
||||
return {
|
||||
embeddings: result.embeddings,
|
||||
totalTokens: result.totalTokens,
|
||||
billableTokens: result.billableTokens,
|
||||
isBYOK: result.isBYOK,
|
||||
modelName: result.modelName,
|
||||
pricingId: result.pricingId,
|
||||
@@ -89,7 +91,7 @@ export async function generateSearchEmbedding(
|
||||
): Promise<{ embedding: number[]; isBYOK: boolean }> {
|
||||
assertKbEmbeddingModel(embeddingModel)
|
||||
|
||||
const result = await embed([query], {
|
||||
const result = await embedKnowledge([query], {
|
||||
model: embeddingModel,
|
||||
workspaceId,
|
||||
taskType: 'query',
|
||||
|
||||
@@ -6,6 +6,7 @@ import { embeddingsCohereTool } from '@/tools/embeddings/cohere'
|
||||
import { embeddingsGeminiTool } from '@/tools/embeddings/gemini'
|
||||
import { embeddingsMistralTool } from '@/tools/embeddings/mistral'
|
||||
import { embeddingsOpenAITool } from '@/tools/embeddings/openai'
|
||||
import { embeddingsOpenRouterTool } from '@/tools/embeddings/openrouter'
|
||||
import { embeddingsTool as legacyOpenAIEmbeddingsTool } from '@/tools/openai/embeddings'
|
||||
|
||||
const ALL_TOOLS = [
|
||||
@@ -13,6 +14,7 @@ const ALL_TOOLS = [
|
||||
embeddingsGeminiTool,
|
||||
embeddingsCohereTool,
|
||||
embeddingsMistralTool,
|
||||
embeddingsOpenRouterTool,
|
||||
legacyOpenAIEmbeddingsTool,
|
||||
]
|
||||
|
||||
@@ -76,4 +78,13 @@ describe('embeddings tools model-input projection', () => {
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
it('requires an explicit OpenRouter key without hosted-key injection', () => {
|
||||
expect(embeddingsOpenRouterTool.params.apiKey.required).toBe(true)
|
||||
expect(embeddingsOpenRouterTool.hosting).toBeUndefined()
|
||||
expect(embeddingsOpenRouterTool.request.body({ input: 'hello' } as never)).toMatchObject({
|
||||
provider: 'openrouter',
|
||||
model: 'openrouter/openai/text-embedding-3-small',
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import type { EmbeddingProvider } from '@/lib/api/contracts/tools/embeddings'
|
||||
import { BYOK_PROVIDER_IDS, DEFAULT_MODEL_BY_PROVIDER } from '@/lib/embeddings/catalog'
|
||||
import type { EmbeddingCatalogProvider } from '@/lib/embeddings/types'
|
||||
import { getEmbeddingModelPricing } from '@/providers/models'
|
||||
import type { EmbeddingsParams, EmbeddingsResponse } from '@/tools/embeddings/types'
|
||||
import type { ToolConfig } from '@/tools/types'
|
||||
@@ -11,28 +12,68 @@ const HOSTED_KEY_RATE_LIMIT = {
|
||||
burstMultiplier: 1,
|
||||
} as const
|
||||
|
||||
interface CreateEmbeddingToolOptions {
|
||||
interface EmbeddingToolBaseOptions {
|
||||
id: string
|
||||
name: string
|
||||
provider: EmbeddingProvider
|
||||
description: string
|
||||
/** Env var prefix for the hosted key pool. */
|
||||
envKeyPrefix: string
|
||||
}
|
||||
|
||||
interface HostedEmbeddingToolOptions extends EmbeddingToolBaseOptions {
|
||||
provider: EmbeddingCatalogProvider
|
||||
envKeyPrefix: string
|
||||
defaultModel?: never
|
||||
}
|
||||
|
||||
interface ExplicitKeyEmbeddingToolOptions extends EmbeddingToolBaseOptions {
|
||||
provider: Extract<EmbeddingProvider, 'openrouter'>
|
||||
envKeyPrefix?: never
|
||||
defaultModel: string
|
||||
}
|
||||
|
||||
type CreateEmbeddingToolOptions = HostedEmbeddingToolOptions | ExplicitKeyEmbeddingToolOptions
|
||||
|
||||
/**
|
||||
* Builds a provider-specific embeddings tool. Every provider shares the same
|
||||
* params, transport, and output shape; only key resolution and the default
|
||||
* model differ, so they are produced from one definition rather than copied.
|
||||
*/
|
||||
export function createEmbeddingTool({
|
||||
id,
|
||||
name,
|
||||
provider,
|
||||
description,
|
||||
envKeyPrefix,
|
||||
}: CreateEmbeddingToolOptions): ToolConfig<EmbeddingsParams, EmbeddingsResponse> {
|
||||
const defaultModel = DEFAULT_MODEL_BY_PROVIDER[provider]
|
||||
export function createEmbeddingTool(
|
||||
options: CreateEmbeddingToolOptions
|
||||
): ToolConfig<EmbeddingsParams, EmbeddingsResponse> {
|
||||
const { id, name, provider, description } = options
|
||||
const defaultModel =
|
||||
provider === 'openrouter' ? options.defaultModel : DEFAULT_MODEL_BY_PROVIDER[provider]
|
||||
/**
|
||||
* Sim-hosted catalog providers are billed per input token with no markup.
|
||||
* OpenRouter requires an explicit user key and therefore has no hosting config.
|
||||
*/
|
||||
const hostingConfig: ToolConfig<EmbeddingsParams, EmbeddingsResponse>['hosting'] =
|
||||
provider === 'openrouter'
|
||||
? undefined
|
||||
: {
|
||||
envKeyPrefix: options.envKeyPrefix,
|
||||
apiKeyParam: 'apiKey',
|
||||
byokProviderId: BYOK_PROVIDER_IDS[provider],
|
||||
pricing: {
|
||||
type: 'custom',
|
||||
getCost: (_params, output) => {
|
||||
const tokens = output.__embeddingTokens
|
||||
if (typeof tokens !== 'number' || Number.isNaN(tokens)) {
|
||||
throw new Error('Embedding response missing token usage')
|
||||
}
|
||||
const model = typeof output.model === 'string' ? output.model : defaultModel
|
||||
const pricing = getEmbeddingModelPricing(model)
|
||||
if (!pricing) {
|
||||
throw new Error(`No pricing configured for embedding model: ${model}`)
|
||||
}
|
||||
return {
|
||||
cost: (tokens * pricing.input) / 1_000_000,
|
||||
metadata: { model, totalTokens: tokens, inputPricePerMillion: pricing.input },
|
||||
}
|
||||
},
|
||||
},
|
||||
rateLimit: HOSTED_KEY_RATE_LIMIT,
|
||||
}
|
||||
|
||||
return {
|
||||
id,
|
||||
@@ -75,34 +116,7 @@ export function createEmbeddingTool({
|
||||
},
|
||||
},
|
||||
|
||||
hosting: {
|
||||
envKeyPrefix,
|
||||
apiKeyParam: 'apiKey',
|
||||
byokProviderId: BYOK_PROVIDER_IDS[provider],
|
||||
/**
|
||||
* Billed per input token with no markup, matching how the knowledge-base
|
||||
* path bills the same models.
|
||||
*/
|
||||
pricing: {
|
||||
type: 'custom',
|
||||
getCost: (_params, output) => {
|
||||
const tokens = output.__embeddingTokens
|
||||
if (typeof tokens !== 'number' || Number.isNaN(tokens)) {
|
||||
throw new Error('Embedding response missing token usage')
|
||||
}
|
||||
const model = typeof output.model === 'string' ? output.model : defaultModel
|
||||
const pricing = getEmbeddingModelPricing(model)
|
||||
if (!pricing) {
|
||||
throw new Error(`No pricing configured for embedding model: ${model}`)
|
||||
}
|
||||
return {
|
||||
cost: (tokens * pricing.input) / 1_000_000,
|
||||
metadata: { model, totalTokens: tokens, inputPricePerMillion: pricing.input },
|
||||
}
|
||||
},
|
||||
},
|
||||
rateLimit: HOSTED_KEY_RATE_LIMIT,
|
||||
},
|
||||
hosting: hostingConfig,
|
||||
|
||||
request: {
|
||||
url: '/api/tools/embeddings',
|
||||
|
||||
@@ -3,4 +3,5 @@ export { createEmbeddingTool } from '@/tools/embeddings/factory'
|
||||
export { embeddingsGeminiTool } from '@/tools/embeddings/gemini'
|
||||
export { embeddingsMistralTool } from '@/tools/embeddings/mistral'
|
||||
export { embeddingsOpenAITool } from '@/tools/embeddings/openai'
|
||||
export { embeddingsOpenRouterTool } from '@/tools/embeddings/openrouter'
|
||||
export type { EmbeddingsParams, EmbeddingsResponse } from '@/tools/embeddings/types'
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
import { DEFAULT_OPENROUTER_EMBEDDING_MODEL } from '@/lib/embeddings/openrouter-models'
|
||||
import { createEmbeddingTool } from '@/tools/embeddings/factory'
|
||||
|
||||
export const embeddingsOpenRouterTool = createEmbeddingTool({
|
||||
id: 'embeddings_openrouter',
|
||||
name: 'OpenRouter Embeddings',
|
||||
provider: 'openrouter',
|
||||
description: 'Generate embeddings through OpenRouter',
|
||||
defaultModel: DEFAULT_OPENROUTER_EMBEDDING_MODEL,
|
||||
})
|
||||
@@ -2,7 +2,7 @@ import type { EmbeddingTaskTypeName } from '@/lib/api/contracts/tools/embeddings
|
||||
import type { ToolResponse } from '@/tools/types'
|
||||
|
||||
export interface EmbeddingsParams {
|
||||
apiKey: string
|
||||
apiKey?: string
|
||||
input: string | string[]
|
||||
model?: string
|
||||
taskType?: EmbeddingTaskTypeName
|
||||
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -1011,6 +1011,7 @@ import {
|
||||
embeddingsGeminiTool,
|
||||
embeddingsMistralTool,
|
||||
embeddingsOpenAITool,
|
||||
embeddingsOpenRouterTool,
|
||||
} from '@/tools/embeddings'
|
||||
import {
|
||||
enrichCheckCreditsTool,
|
||||
@@ -6839,6 +6840,7 @@ export const tools: Record<string, ToolConfig> = {
|
||||
embeddings_gemini: embeddingsGeminiTool,
|
||||
embeddings_cohere: embeddingsCohereTool,
|
||||
embeddings_mistral: embeddingsMistralTool,
|
||||
embeddings_openrouter: embeddingsOpenRouterTool,
|
||||
evernote_copy_note: evernoteCopyNoteTool,
|
||||
evernote_create_note: evernoteCreateNoteTool,
|
||||
evernote_create_notebook: evernoteCreateNotebookTool,
|
||||
|
||||
@@ -9,8 +9,8 @@ const QUERY_HOOKS_DIR = path.join(ROOT, 'apps/sim/hooks/queries')
|
||||
const SELECTOR_HOOKS_DIR = path.join(ROOT, 'apps/sim/hooks/selectors')
|
||||
|
||||
const BASELINE = {
|
||||
totalRoutes: 1009,
|
||||
zodRoutes: 1009,
|
||||
totalRoutes: 1010,
|
||||
zodRoutes: 1010,
|
||||
nonZodRoutes: 0,
|
||||
} as const
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ import {
|
||||
type EnvCapabilityValues,
|
||||
getCapabilityFields,
|
||||
hasEnvCapabilityValue,
|
||||
KNOWLEDGE_EMBEDDINGS_CAPABILITY,
|
||||
OAUTH_CLIENT_CAPABILITIES,
|
||||
type OAuthClientCapabilityField,
|
||||
type OAuthClientCapabilityId,
|
||||
@@ -730,6 +731,101 @@ export const KNOWLEDGE_SETUP = defineCapabilitySetup(OCR_CAPABILITY, {
|
||||
optionOrder: ['local', 'mistral', 'azure-mistral'],
|
||||
})
|
||||
|
||||
export const KNOWLEDGE_EMBEDDINGS_SETUP = defineCapabilitySetup(KNOWLEDGE_EMBEDDINGS_CAPABILITY, {
|
||||
label: 'Knowledge embeddings',
|
||||
message: 'Knowledge embedding provider?',
|
||||
actions: {},
|
||||
providers: {
|
||||
'azure-openai': {
|
||||
hint: 'preferred when configured; falls back to other configured providers',
|
||||
prompts: [
|
||||
{
|
||||
type: 'field',
|
||||
key: 'AZURE_OPENAI_ENDPOINT',
|
||||
input: 'text',
|
||||
required: true,
|
||||
validate: true,
|
||||
},
|
||||
{
|
||||
type: 'field',
|
||||
key: 'AZURE_OPENAI_API_VERSION',
|
||||
input: 'text',
|
||||
required: true,
|
||||
},
|
||||
{
|
||||
type: 'field',
|
||||
key: 'AZURE_OPENAI_API_KEY',
|
||||
input: 'secret',
|
||||
required: true,
|
||||
},
|
||||
{
|
||||
type: 'field',
|
||||
key: 'KB_OPENAI_MODEL_NAME',
|
||||
input: 'text',
|
||||
hint: 'optional Azure deployment name; defaults to the embedding model id',
|
||||
},
|
||||
],
|
||||
},
|
||||
openai: {
|
||||
hint: 'direct OpenAI with optional key rotation',
|
||||
prompts: [
|
||||
{
|
||||
type: 'choice',
|
||||
id: 'openai-credentials',
|
||||
message: 'OpenAI credentials?',
|
||||
options: [
|
||||
{
|
||||
id: 'single',
|
||||
label: 'Single API key',
|
||||
currentWhen: { kind: 'present', key: 'OPENAI_API_KEY' },
|
||||
prompts: [{ type: 'field', key: 'OPENAI_API_KEY', input: 'secret', required: true }],
|
||||
},
|
||||
{
|
||||
id: 'rotating',
|
||||
label: 'Rotating key pool',
|
||||
currentWhen: {
|
||||
kind: 'any',
|
||||
conditions: [
|
||||
{ kind: 'present', key: 'OPENAI_API_KEY_1' },
|
||||
{ kind: 'present', key: 'OPENAI_API_KEY_2' },
|
||||
{ kind: 'present', key: 'OPENAI_API_KEY_3' },
|
||||
],
|
||||
},
|
||||
prompts: [
|
||||
{ type: 'field', key: 'OPENAI_API_KEY_1', input: 'secret', required: true },
|
||||
{
|
||||
type: 'field',
|
||||
key: 'OPENAI_API_KEY_2',
|
||||
input: 'secret',
|
||||
hint: 'optional rotating key',
|
||||
},
|
||||
{
|
||||
type: 'field',
|
||||
key: 'OPENAI_API_KEY_3',
|
||||
input: 'secret',
|
||||
hint: 'optional rotating key',
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
},
|
||||
openrouter: {
|
||||
hint: 'fallback for OpenAI knowledge embedding models',
|
||||
prompts: [
|
||||
{
|
||||
type: 'field',
|
||||
key: 'OPENROUTER_API_KEY',
|
||||
input: 'secret',
|
||||
required: true,
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
optionOrder: ['azure-openai', 'openai', 'openrouter'],
|
||||
})
|
||||
|
||||
export const CAPABILITY_SETUPS = [
|
||||
EMAIL_SETUP,
|
||||
STORAGE_SETUP,
|
||||
@@ -737,6 +833,7 @@ export const CAPABILITY_SETUPS = [
|
||||
JOBS_SETUP,
|
||||
CACHE_SETUP,
|
||||
KNOWLEDGE_SETUP,
|
||||
KNOWLEDGE_EMBEDDINGS_SETUP,
|
||||
] as const
|
||||
|
||||
const configuredCapabilityIds = new Set(CAPABILITY_SETUPS.map((setup) => setup.definition.id))
|
||||
|
||||
@@ -14,6 +14,10 @@ describe('env capability status', () => {
|
||||
expect(status.features.jobs).toMatchObject({ state: 'default', providerId: 'database' })
|
||||
expect(status.features.cache).toMatchObject({ state: 'default', providerId: 'database' })
|
||||
expect(status.features.knowledge).toMatchObject({ state: 'default', providerId: 'local' })
|
||||
expect(status.features['knowledge-embeddings']).toMatchObject({
|
||||
state: 'missing',
|
||||
providerIds: [],
|
||||
})
|
||||
expect(status.features.llm).toMatchObject({
|
||||
state: 'missing',
|
||||
configuredPoolCount: 0,
|
||||
@@ -60,6 +64,20 @@ describe('env capability status', () => {
|
||||
expect(JSON.stringify(status)).not.toContain('secret-private-key')
|
||||
})
|
||||
|
||||
it('reports configured knowledge embedding transports without exposing keys', () => {
|
||||
const status = buildEnvCapabilityStatus({
|
||||
OPENAI_API_KEY: 'secret-openai-key',
|
||||
OPENROUTER_API_KEY: 'secret-openrouter-key',
|
||||
})
|
||||
|
||||
expect(status.features['knowledge-embeddings']).toMatchObject({
|
||||
state: 'configured',
|
||||
providerIds: ['openai', 'openrouter'],
|
||||
})
|
||||
expect(JSON.stringify(status)).not.toContain('secret-openai-key')
|
||||
expect(JSON.stringify(status)).not.toContain('secret-openrouter-key')
|
||||
})
|
||||
|
||||
it('reports selected Daytona and cloud storage providers', () => {
|
||||
const status = buildEnvCapabilityStatus({
|
||||
SANDBOX_PROVIDER: 'daytona',
|
||||
|
||||
@@ -15,6 +15,7 @@ import {
|
||||
inspectCapability,
|
||||
inspectOAuthClientCapability,
|
||||
isTruthyEnvCapabilityValue,
|
||||
KNOWLEDGE_EMBEDDINGS_CAPABILITY,
|
||||
LLM_KEY_POOLS,
|
||||
OAUTH_CLIENT_CAPABILITIES,
|
||||
type OAuthClientCapabilityId,
|
||||
@@ -43,6 +44,8 @@ interface FeatureStatusBase<TId extends SetupStatusFeatureId> {
|
||||
}
|
||||
|
||||
type EmailProviderId = (typeof EMAIL_CAPABILITY.providers)[number]['id']
|
||||
type KnowledgeEmbeddingsProviderId =
|
||||
(typeof KNOWLEDGE_EMBEDDINGS_CAPABILITY.providers)[number]['id']
|
||||
type StorageProviderId =
|
||||
| (typeof STORAGE_CAPABILITY)['defaultProvider']['id']
|
||||
| (typeof STORAGE_CAPABILITY.providers)[number]['id']
|
||||
@@ -77,6 +80,13 @@ export interface KnowledgeCapabilityStatus extends FeatureStatusBase<'knowledge'
|
||||
providerId: 'local' | 'mistral' | 'azure-mistral' | null
|
||||
}
|
||||
|
||||
export interface KnowledgeEmbeddingsCapabilityStatus
|
||||
extends FeatureStatusBase<'knowledge-embeddings'> {
|
||||
strategy: 'fallback'
|
||||
providerIds: readonly KnowledgeEmbeddingsProviderId[]
|
||||
providers: readonly ProviderInspection<KnowledgeEmbeddingsProviderId>[]
|
||||
}
|
||||
|
||||
export interface LlmKeyPoolStatus {
|
||||
id: LlmKeyPoolId
|
||||
state: 'configured' | 'missing'
|
||||
@@ -99,6 +109,7 @@ interface FeatureStatusById {
|
||||
jobs: JobsCapabilityStatus
|
||||
cache: CacheCapabilityStatus
|
||||
knowledge: KnowledgeCapabilityStatus
|
||||
'knowledge-embeddings': KnowledgeEmbeddingsCapabilityStatus
|
||||
llm: LlmCapabilityStatus
|
||||
}
|
||||
|
||||
@@ -391,6 +402,26 @@ function inspectKnowledge(values: EnvCapabilityValues): KnowledgeCapabilityStatu
|
||||
}
|
||||
}
|
||||
|
||||
function inspectKnowledgeEmbeddings(
|
||||
values: EnvCapabilityValues
|
||||
): KnowledgeEmbeddingsCapabilityStatus {
|
||||
const inspection = inspectCapability(KNOWLEDGE_EMBEDDINGS_CAPABILITY, values)
|
||||
const brokenState = brokenProviderState(inspection.providers)
|
||||
const state = inspection.configured ? 'configured' : (brokenState ?? 'missing')
|
||||
const configurationError =
|
||||
inspection.error ??
|
||||
getCapabilityConfigurationError(KNOWLEDGE_EMBEDDINGS_CAPABILITY, inspection.providers)
|
||||
|
||||
return {
|
||||
...featureMetadata('knowledge-embeddings'),
|
||||
strategy: 'fallback',
|
||||
state,
|
||||
providerIds: inspection.providerIds,
|
||||
providers: inspection.providers,
|
||||
...(configurationError ? { issue: issue(brokenState ?? 'invalid', configurationError) } : {}),
|
||||
}
|
||||
}
|
||||
|
||||
function inspectLlm(values: EnvCapabilityValues): LlmCapabilityStatus {
|
||||
const pools = {} as Record<LlmKeyPoolId, LlmKeyPoolStatus>
|
||||
let configuredPoolCount = 0
|
||||
@@ -463,6 +494,7 @@ const FEATURE_STATUS_BUILDERS = {
|
||||
jobs: inspectJobs,
|
||||
cache: inspectCache,
|
||||
knowledge: inspectKnowledge,
|
||||
'knowledge-embeddings': inspectKnowledgeEmbeddings,
|
||||
llm: inspectLlm,
|
||||
} satisfies Record<SetupStatusFeatureId, (values: EnvCapabilityValues) => unknown>
|
||||
|
||||
@@ -476,6 +508,7 @@ export function buildEnvCapabilityStatus(values: EnvCapabilityValues): EnvCapabi
|
||||
jobs: FEATURE_STATUS_BUILDERS.jobs(values),
|
||||
cache: FEATURE_STATUS_BUILDERS.cache(values),
|
||||
knowledge: FEATURE_STATUS_BUILDERS.knowledge(values),
|
||||
'knowledge-embeddings': FEATURE_STATUS_BUILDERS['knowledge-embeddings'](values),
|
||||
llm: FEATURE_STATUS_BUILDERS.llm(values),
|
||||
},
|
||||
oauthClients: inspectOAuthClients(values),
|
||||
|
||||
@@ -104,6 +104,15 @@ function featureDetail(feature: FeatureStatus): string {
|
||||
if (feature.providerId === 'local') return 'Local parser (default)'
|
||||
if (feature.providerId === 'azure-mistral') return 'Azure Mistral OCR'
|
||||
return feature.providerId === 'mistral' ? 'Mistral OCR' : 'Not configured'
|
||||
case 'knowledge-embeddings':
|
||||
return feature.providerIds.length > 0
|
||||
? feature.providerIds
|
||||
.map(
|
||||
(id) =>
|
||||
feature.providers.find((provider) => provider.id === id)?.label ?? titleCase(id)
|
||||
)
|
||||
.join(' → ')
|
||||
: 'Not configured'
|
||||
case 'llm': {
|
||||
const configured = Object.values(feature.pools)
|
||||
.filter((pool) => pool.state === 'configured')
|
||||
|
||||
Reference in New Issue
Block a user