diff --git a/open-sse/config/embeddingRegistry.ts b/open-sse/config/embeddingRegistry.ts index 9ed26527fe..a96327e8f0 100644 --- a/open-sse/config/embeddingRegistry.ts +++ b/open-sse/config/embeddingRegistry.ts @@ -8,7 +8,43 @@ * keyed by provider ID (e.g. "nebius", "openai"). */ -export const EMBEDDING_PROVIDERS = { +export interface EmbeddingProvider { + id: string; + baseUrl: string; + authType: string; + authHeader: string; + models: { id: string; name: string; dimensions?: number }[]; +} + +export interface EmbeddingProviderNodeRow { + prefix: string; + name: string; + baseUrl: string; + apiType?: string; +} + +/** + * Build a dynamic EmbeddingProvider from a local provider_node. + * Only used for local providers (localhost) — caller must filter by hostname. + */ +export function buildDynamicEmbeddingProvider(node: EmbeddingProviderNodeRow): EmbeddingProvider { + if (!node.prefix || !node.baseUrl) { + throw new Error(`Invalid provider_node: missing prefix or baseUrl`); + } + if (node.prefix.includes("/") || node.prefix.includes(" ")) { + throw new Error(`Invalid provider_node prefix "${node.prefix}": must not contain / or spaces`); + } + const baseUrl = node.baseUrl.replace(/\/+$/, ""); + return { + id: node.prefix, + baseUrl: `${baseUrl}/embeddings`, + authType: "none", + authHeader: "none", + models: [], + }; +} + +export const EMBEDDING_PROVIDERS: Record = { nebius: { id: "nebius", baseUrl: "https://api.tokenfactory.nebius.com/v1/embeddings", @@ -70,7 +106,7 @@ export const EMBEDDING_PROVIDERS = { /** * Get embedding provider config by ID */ -export function getEmbeddingProvider(providerId) { +export function getEmbeddingProvider(providerId: string): EmbeddingProvider | null { return EMBEDDING_PROVIDERS[providerId] || null; } @@ -78,26 +114,36 @@ export function getEmbeddingProvider(providerId) { * Parse embedding model string (format: "provider/model" or just "model") * Returns { provider, model } */ -export function parseEmbeddingModel(modelStr) { +export function parseEmbeddingModel( + modelStr: string | null, + dynamicProviders?: EmbeddingProvider[] +): { provider: string | null; model: string | null } { if (!modelStr) return { provider: null, model: null }; // Check for "provider/model" format const slashIdx = modelStr.indexOf("/"); if (slashIdx > 0) { - // Handle nested model IDs like "nebius/Qwen/Qwen3-Embedding-8B" - // Try each provider prefix - for (const [providerId, config] of Object.entries(EMBEDDING_PROVIDERS)) { + // Phase 1: Try each hardcoded provider prefix + for (const [providerId] of Object.entries(EMBEDDING_PROVIDERS)) { if (modelStr.startsWith(providerId + "/")) { return { provider: providerId, model: modelStr.slice(providerId.length + 1) }; } } - // Fallback: first segment is provider + // Phase 2: Try dynamic provider_nodes prefix + if (dynamicProviders) { + for (const dp of dynamicProviders) { + if (modelStr.startsWith(dp.id + "/")) { + return { provider: dp.id, model: modelStr.slice(dp.id.length + 1) }; + } + } + } + // Phase 3: Fallback — first segment is provider const provider = modelStr.slice(0, slashIdx); const model = modelStr.slice(slashIdx + 1); return { provider, model }; } - // No provider prefix — search all providers for the model + // No provider prefix — search hardcoded providers for the model for (const [providerId, config] of Object.entries(EMBEDDING_PROVIDERS)) { if (config.models.some((m) => m.id === modelStr)) { return { provider: providerId, model: modelStr }; diff --git a/open-sse/handlers/embeddings.ts b/open-sse/handlers/embeddings.ts index 9dbf84dd28..0d20ab1dbb 100644 --- a/open-sse/handlers/embeddings.ts +++ b/open-sse/handlers/embeddings.ts @@ -13,18 +13,48 @@ * } */ -import { getEmbeddingProvider, parseEmbeddingModel } from "../config/embeddingRegistry.ts"; +import { + getEmbeddingProvider, + parseEmbeddingModel, + type EmbeddingProvider, +} from "../config/embeddingRegistry.ts"; import { saveCallLog } from "@/lib/usageDb"; /** - * Handle embedding request - * @param {object} options - * @param {object} options.body - Request body - * @param {object} options.credentials - Provider credentials { apiKey, accessToken } - * @param {object} options.log - Logger + * Handle embedding request. + * Supports both hardcoded cloud providers and dynamic local provider_nodes. + * When resolvedProvider is passed, uses it directly (injection pattern from route handler). + * Falls back to hardcoded registry lookup for backward compatibility. */ -export async function handleEmbedding({ body, credentials, log }) { - const { provider, model } = parseEmbeddingModel(body.model); +export async function handleEmbedding({ + body, + credentials, + log, + resolvedProvider = null, + resolvedModel = null, +}: { + body: Record; + credentials: { apiKey?: string; accessToken?: string } | null; + log?: { info: (...args: unknown[]) => void; error: (...args: unknown[]) => void }; + resolvedProvider?: EmbeddingProvider | null; + resolvedModel?: string | null; +}) { + // Use pre-resolved provider/model from route handler if available (supports dynamic provider_nodes). + let provider: string | null; + let model: string | null; + let providerConfig: EmbeddingProvider | null; + + if (resolvedProvider) { + provider = resolvedProvider.id; + model = resolvedModel; + providerConfig = resolvedProvider; + } else { + const parsed = parseEmbeddingModel(body.model as string); + provider = parsed.provider; + model = parsed.model; + providerConfig = provider ? getEmbeddingProvider(provider) : null; + } + const startTime = Date.now(); // Summarized request body for call log (avoid storing large embedding input arrays) @@ -42,7 +72,6 @@ export async function handleEmbedding({ body, credentials, log }) { }; } - const providerConfig = getEmbeddingProvider(provider); if (!providerConfig) { return { success: false, @@ -66,11 +95,15 @@ export async function handleEmbedding({ body, credentials, log }) { "Content-Type": "application/json", }; - const token = credentials.apiKey || credentials.accessToken; - if (providerConfig.authHeader === "bearer") { - headers["Authorization"] = `Bearer ${token}`; - } else if (providerConfig.authHeader === "x-api-key") { - headers["x-api-key"] = token; + // Skip credential injection for local providers (authType: "none") + const token = + providerConfig.authType === "none" ? null : credentials?.apiKey || credentials?.accessToken; + if (token) { + if (providerConfig.authHeader === "bearer") { + headers["Authorization"] = `Bearer ${token}`; + } else if (providerConfig.authHeader === "x-api-key") { + headers["x-api-key"] = token; + } } if (log) { diff --git a/src/app/api/v1/embeddings/route.ts b/src/app/api/v1/embeddings/route.ts index bd4c5a96ba..515cf2e28c 100644 --- a/src/app/api/v1/embeddings/route.ts +++ b/src/app/api/v1/embeddings/route.ts @@ -9,6 +9,9 @@ import { import { parseEmbeddingModel, getAllEmbeddingModels, + getEmbeddingProvider, + buildDynamicEmbeddingProvider, + type EmbeddingProviderNodeRow, } from "@omniroute/open-sse/config/embeddingRegistry.ts"; import { errorResponse } from "@omniroute/open-sse/utils/error.ts"; import { HTTP_STATUS } from "@omniroute/open-sse/config/constants.ts"; @@ -18,7 +21,7 @@ import { enforceApiKeyPolicy } from "@/shared/utils/apiKeyPolicy"; import { v1EmbeddingsSchema } from "@/shared/validation/schemas"; import { isValidationFailure, validateBody } from "@/shared/validation/helpers"; -import { getAllCustomModels } from "@/lib/localDb"; +import { getAllCustomModels, getProviderNodes } from "@/lib/localDb"; /** * Handle CORS preflight @@ -110,8 +113,42 @@ export async function POST(request) { const policy = await enforceApiKeyPolicy(request, body.model); if (policy.rejection) return policy.rejection; + // Load local provider_nodes for embedding routing (only localhost — prevents auth bypass/SSRF) + let dynamicProviders: ReturnType[] = []; + try { + const nodes = await getProviderNodes(); + dynamicProviders = (Array.isArray(nodes) ? nodes : []) + .filter((n: EmbeddingProviderNodeRow) => { + // provider_nodes apiType is "chat" or "responses" (not "embeddings") — local OpenAI-compatible + // backends expose /embeddings under the same base URL as chat, so we build the URL as baseUrl + /embeddings. + if (n.apiType !== "chat" && n.apiType !== "responses") return false; + try { + const hostname = new URL(n.baseUrl).hostname; + return ( + hostname === "localhost" || + hostname === "127.0.0.1" || + hostname === "::1" || + hostname === "[::1]" + ); + } catch { + return false; + } + }) + .map((n) => { + try { + return buildDynamicEmbeddingProvider(n); + } catch (err) { + log.error("EMBED", `Skipping invalid provider_node ${n.prefix}: ${err}`); + return null; + } + }) + .filter((p): p is NonNullable => p !== null); + } catch (err) { + log.error("EMBED", `Failed to load provider_nodes for embeddings: ${err}`); + } + // Parse model to get provider - const { provider } = parseEmbeddingModel(body.model); + const { provider, model: resolvedModel } = parseEmbeddingModel(body.model, dynamicProviders); if (!provider) { return errorResponse( HTTP_STATUS.BAD_REQUEST, @@ -119,19 +156,39 @@ export async function POST(request) { ); } - // Get credentials for the embedding provider - const credentials = await getProviderCredentials(provider); - if (!credentials) { + // Resolve provider config — dynamic first (local override), then hardcoded + const providerConfig = + dynamicProviders.find((dp) => dp.id === provider) || getEmbeddingProvider(provider) || null; + + if (!providerConfig) { return errorResponse( HTTP_STATUS.BAD_REQUEST, - `No credentials for embedding provider: ${provider}` + `Unknown embedding provider: ${provider}. No matching hardcoded or local provider found.` ); } - const result = await handleEmbedding({ body, credentials, log }); + // Get credentials — skip for local providers (authType: "none") + let credentials = null; + if (providerConfig && providerConfig.authType !== "none") { + credentials = await getProviderCredentials(provider); + if (!credentials) { + return errorResponse( + HTTP_STATUS.BAD_REQUEST, + `No credentials for embedding provider: ${provider}` + ); + } + } + + const result = await handleEmbedding({ + body, + credentials, + log, + resolvedProvider: providerConfig, + resolvedModel, + }); if (result.success) { - await clearRecoveredProviderState(credentials); + if (credentials) await clearRecoveredProviderState(credentials); return new Response(JSON.stringify(result.data), { status: 200, headers: { "Content-Type": "application/json" }, diff --git a/typescript b/typescript new file mode 100644 index 0000000000..e69de29bb2