diff --git a/open-sse/config/petals.ts b/open-sse/config/petals.ts new file mode 100644 index 0000000000..569c65ba1d --- /dev/null +++ b/open-sse/config/petals.ts @@ -0,0 +1,19 @@ +export const PETALS_DEFAULT_BASE_URL = "https://chat.petals.dev/api/v1/generate"; +export const PETALS_DEFAULT_MODEL = "stabilityai/StableBeluga2"; + +export function normalizePetalsBaseUrl(baseUrl: string | null | undefined): string { + const normalized = String(baseUrl || PETALS_DEFAULT_BASE_URL) + .trim() + .replace(/\/+$/, ""); + + if (normalized.endsWith("/api/v1/generate")) { + return normalized; + } + if (normalized.endsWith("/api/v1")) { + return `${normalized}/generate`; + } + if (normalized.endsWith("/api")) { + return `${normalized}/v1/generate`; + } + return `${normalized}/api/v1/generate`; +} diff --git a/open-sse/config/providerRegistry.ts b/open-sse/config/providerRegistry.ts index 73e8fbe4a6..445b04c5d0 100644 --- a/open-sse/config/providerRegistry.ts +++ b/open-sse/config/providerRegistry.ts @@ -30,6 +30,7 @@ import { BEDROCK_DEFAULT_BASE_URL } from "./bedrock.ts"; import { WATSONX_DEFAULT_BASE_URL } from "./watsonx.ts"; import { OCI_DEFAULT_BASE_URL } from "./oci.ts"; import { SAP_DEFAULT_BASE_URL } from "./sap.ts"; +import { PETALS_DEFAULT_BASE_URL } from "./petals.ts"; import { CURSOR_REGISTRY_VERSION, getAntigravityProviderHeaders, @@ -205,6 +206,26 @@ const CHAT_OPENAI_COMPAT_MODELS: Record = { "meta-llama/Meta-Llama-3-8B-Instruct", "meta-llama/Meta-Llama-3-70B-Instruct", ]), + "nous-research": [ + { id: "nousresearch/hermes-4-70b", name: "Nous: Hermes 4 70B", contextLength: 131072 }, + { id: "nousresearch/hermes-4-405b", name: "Nous: Hermes 4 405B", contextLength: 131072 }, + { + id: "nousresearch/hermes-3-llama-3.1-70b", + name: "Nous: Hermes 3 70B Instruct", + contextLength: 131072, + }, + { + id: "nousresearch/hermes-3-llama-3.1-405b", + name: "Nous: Hermes 3 405B Instruct", + contextLength: 131072, + }, + { + id: "nousresearch/hermes-2-pro-llama-3-8b", + name: "NousResearch: Hermes 2 Pro - Llama-3 8B", + contextLength: 8192, + }, + ], + petals: [{ id: "stabilityai/StableBeluga2", name: "Stable Beluga 2 (70B)", contextLength: 8192 }], poe: buildModels(["Claude-Sonnet-4.5", "GPT-5-Pro", "GPT-5-Codex", "Gemini-2.5-Pro"]), gitlab: [{ id: "gitlab-duo-code-suggestions", name: "GitLab Duo Code Suggestions" }], "gitlab-duo": [{ id: "gitlab-duo-code-suggestions", name: "GitLab Duo Code Suggestions" }], @@ -2075,6 +2096,30 @@ export const REGISTRY: Record = { passthroughModels: true, }, + "nous-research": { + id: "nous-research", + alias: "nous", + format: "openai", + executor: "default", + baseUrl: "https://inference-api.nousresearch.com/v1", + modelsUrl: "https://inference-api.nousresearch.com/v1/models", + authType: "apikey", + authHeader: "bearer", + models: CHAT_OPENAI_COMPAT_MODELS["nous-research"], + passthroughModels: true, + }, + + petals: { + id: "petals", + alias: "petals", + format: "openai", + executor: "petals", + baseUrl: PETALS_DEFAULT_BASE_URL, + authType: "apikey", + authHeader: "bearer", + models: CHAT_OPENAI_COMPAT_MODELS.petals, + }, + poe: { id: "poe", alias: "poe", diff --git a/open-sse/executors/index.ts b/open-sse/executors/index.ts index d941229962..f270d4e2ed 100644 --- a/open-sse/executors/index.ts +++ b/open-sse/executors/index.ts @@ -19,6 +19,7 @@ import { MuseSparkWebExecutor } from "./muse-spark-web.ts"; import { AzureOpenAIExecutor } from "./azure-openai.ts"; import { GitlabExecutor } from "./gitlab.ts"; import { NlpCloudExecutor } from "./nlpcloud.ts"; +import { PetalsExecutor } from "./petals.ts"; const executors = { antigravity: new AntigravityExecutor(), @@ -34,6 +35,7 @@ const executors = { gitlab: new GitlabExecutor(), "gitlab-duo": new GitlabExecutor("gitlab-duo"), nlpcloud: new NlpCloudExecutor(), + petals: new PetalsExecutor(), pollinations: new PollinationsExecutor(), pol: new PollinationsExecutor(), // Alias "cloudflare-ai": new CloudflareAIExecutor(), @@ -89,3 +91,4 @@ export { MuseSparkWebExecutor } from "./muse-spark-web.ts"; export { AzureOpenAIExecutor } from "./azure-openai.ts"; export { GitlabExecutor } from "./gitlab.ts"; export { NlpCloudExecutor } from "./nlpcloud.ts"; +export { PetalsExecutor } from "./petals.ts"; diff --git a/open-sse/executors/petals.ts b/open-sse/executors/petals.ts new file mode 100644 index 0000000000..fb0dfd7b21 --- /dev/null +++ b/open-sse/executors/petals.ts @@ -0,0 +1,385 @@ +import { randomUUID } from "node:crypto"; + +import { + BaseExecutor, + mergeUpstreamExtraHeaders, + type ExecuteInput, + type ProviderCredentials, +} from "./base.ts"; +import { + PETALS_DEFAULT_BASE_URL, + PETALS_DEFAULT_MODEL, + normalizePetalsBaseUrl, +} from "../config/petals.ts"; +import { PROVIDERS } from "../config/constants.ts"; + +type JsonRecord = Record; +type OpenAIMessage = { + role?: string; + content?: unknown; +}; + +function asRecord(value: unknown): JsonRecord { + return value && typeof value === "object" && !Array.isArray(value) ? (value as JsonRecord) : {}; +} + +function extractTextContent(content: unknown): string { + if (typeof content === "string") { + return content.trim(); + } + + if (!Array.isArray(content)) { + return ""; + } + + return content + .map((part) => { + if (!part || typeof part !== "object") return ""; + const item = part as Record; + if (item.type === "text" && typeof item.text === "string") { + return item.text; + } + if (item.type === "input_text" && typeof item.text === "string") { + return item.text; + } + return ""; + }) + .filter((text) => text.trim().length > 0) + .join("\n") + .trim(); +} + +function resolvePrompt(body: unknown): string { + const payload = asRecord(body); + + const directPrompt = extractTextContent(payload.prompt); + if (directPrompt) { + return directPrompt; + } + + const directInput = extractTextContent(payload.input); + if (directInput) { + return directInput; + } + + const messages = Array.isArray(payload.messages) ? (payload.messages as OpenAIMessage[]) : []; + if (messages.length === 0) return ""; + + const systemParts: string[] = []; + const transcript: string[] = []; + let lastRole = ""; + + for (const message of messages) { + const role = String(message?.role || "user").toLowerCase(); + const text = extractTextContent(message?.content); + if (!text) continue; + + if (role === "system" || role === "developer") { + systemParts.push(text); + continue; + } + + if (role === "assistant") { + transcript.push(`Assistant: ${text}`); + lastRole = "assistant"; + continue; + } + + transcript.push(`User: ${text}`); + lastRole = "user"; + } + + if (transcript.length === 0) { + return systemParts.join("\n\n").trim(); + } + + const parts: string[] = []; + if (systemParts.length > 0) { + parts.push(`System:\n${systemParts.join("\n\n")}`); + } + parts.push(transcript.join("\n\n")); + if (lastRole !== "assistant") { + parts.push("Assistant:"); + } + + return parts.join("\n\n").trim(); +} + +function resolveMaxNewTokens(body: unknown): number { + const payload = asRecord(body); + const candidates = [ + payload.max_new_tokens, + payload.max_completion_tokens, + payload.max_output_tokens, + payload.max_tokens, + ]; + + for (const value of candidates) { + if (typeof value === "number" && Number.isFinite(value) && value > 0) { + return Math.max(1, Math.min(4096, Math.floor(value))); + } + } + + return 256; +} + +function buildRequestPayload(model: string, body: unknown): URLSearchParams | null { + const payload = asRecord(body); + const prompt = resolvePrompt(payload); + if (!prompt) return null; + + const form = new URLSearchParams(); + form.set("model", model || PETALS_DEFAULT_MODEL); + form.set("inputs", prompt); + form.set("max_new_tokens", String(resolveMaxNewTokens(payload))); + + const hasSampling = + typeof payload.temperature === "number" || + typeof payload.top_k === "number" || + typeof payload.top_p === "number"; + + if (hasSampling) { + form.set("do_sample", "1"); + } + + if (typeof payload.temperature === "number") { + form.set("temperature", String(payload.temperature)); + } + if (typeof payload.top_k === "number") { + form.set("top_k", String(Math.max(1, Math.floor(payload.top_k)))); + } + if (typeof payload.top_p === "number") { + form.set("top_p", String(payload.top_p)); + } + if (typeof payload.repetition_penalty === "number") { + form.set("repetition_penalty", String(payload.repetition_penalty)); + } + + return form; +} + +function estimateTokens(text: string): number { + return Math.max(1, Math.ceil(text.length / 4)); +} + +function buildSseChunk(data: unknown): string { + return `data: ${JSON.stringify(data)}\n\n`; +} + +function buildOpenAiJsonCompletion( + content: string, + model: string, + id: string, + created: number +): Response { + const completionTokens = estimateTokens(content); + + return new Response( + JSON.stringify({ + id, + object: "chat.completion", + created, + model, + choices: [ + { + index: 0, + message: { role: "assistant", content }, + finish_reason: "stop", + }, + ], + usage: { + prompt_tokens: completionTokens, + completion_tokens: completionTokens, + total_tokens: completionTokens * 2, + }, + }), + { + status: 200, + headers: { "Content-Type": "application/json" }, + } + ); +} + +function buildSynthesizedStream( + content: string, + model: string, + id: string, + created: number +): Response { + const encoder = new TextEncoder(); + + const body = new ReadableStream({ + start(controller) { + controller.enqueue( + encoder.encode( + buildSseChunk({ + id, + object: "chat.completion.chunk", + created, + model, + choices: [{ index: 0, delta: { role: "assistant" }, finish_reason: null }], + }) + ) + ); + + if (content) { + controller.enqueue( + encoder.encode( + buildSseChunk({ + id, + object: "chat.completion.chunk", + created, + model, + choices: [{ index: 0, delta: { content }, finish_reason: null }], + }) + ) + ); + } + + controller.enqueue( + encoder.encode( + buildSseChunk({ + id, + object: "chat.completion.chunk", + created, + model, + choices: [{ index: 0, delta: {}, finish_reason: "stop" }], + }) + ) + ); + controller.enqueue(encoder.encode("data: [DONE]\n\n")); + controller.close(); + }, + }); + + return new Response(body, { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + }); +} + +function toOpenAiError(status: number, message: string): Response { + return new Response( + JSON.stringify({ + error: { + message, + type: + status === 401 || status === 403 + ? "authentication_error" + : status === 429 + ? "rate_limit_error" + : "api_error", + }, + }), + { + status, + headers: { "Content-Type": "application/json" }, + } + ); +} + +export class PetalsExecutor extends BaseExecutor { + constructor() { + super("petals", PROVIDERS.petals || { format: "openai", baseUrl: PETALS_DEFAULT_BASE_URL }); + } + + buildUrl( + _model: string, + _stream: boolean, + _urlIndex = 0, + credentials: ProviderCredentials | null = null + ): string { + const rawBaseUrl = + typeof credentials?.providerSpecificData?.baseUrl === "string" + ? credentials.providerSpecificData.baseUrl + : this.config.baseUrl; + return normalizePetalsBaseUrl(rawBaseUrl); + } + + buildHeaders(credentials: ProviderCredentials | null): Record { + const token = credentials?.apiKey || credentials?.accessToken; + return { + "Content-Type": "application/x-www-form-urlencoded", + Accept: "application/json", + ...(token ? { Authorization: `Bearer ${token}` } : {}), + }; + } + + async execute({ model, body, stream, credentials, signal, upstreamExtraHeaders }: ExecuteInput) { + const resolvedModel = model || PETALS_DEFAULT_MODEL; + const payload = buildRequestPayload(resolvedModel, body); + const url = this.buildUrl(resolvedModel, stream, 0, credentials); + const headers = this.buildHeaders(credentials); + mergeUpstreamExtraHeaders(headers, upstreamExtraHeaders); + + if (!payload) { + return { + response: toOpenAiError(400, "Petals requests require at least one user prompt."), + url, + headers, + transformedBody: body, + }; + } + + const transformedBody = Object.fromEntries(payload.entries()); + + try { + const response = await fetch(url, { + method: "POST", + headers, + body: payload.toString(), + signal, + }); + + if (!response.ok) { + const errorText = await response.text(); + return { + response: toOpenAiError( + response.status, + `Petals API failed with status ${response.status}: ${errorText || "Unknown error"}` + ), + url, + headers, + transformedBody, + }; + } + + const json = asRecord(await response.json()); + if (json.ok === false) { + const traceback = + typeof json.traceback === "string" && json.traceback.trim() + ? json.traceback.trim() + : "Unknown Petals upstream error"; + return { + response: toOpenAiError(502, `Petals API error: ${traceback}`), + url, + headers, + transformedBody, + }; + } + + const content = typeof json.outputs === "string" ? json.outputs : ""; + const id = `chatcmpl-petals-${randomUUID()}`; + const created = Math.floor(Date.now() / 1000); + + return { + response: stream + ? buildSynthesizedStream(content, resolvedModel, id, created) + : buildOpenAiJsonCompletion(content, resolvedModel, id, created), + url, + headers, + transformedBody, + }; + } catch (error) { + const message = error instanceof Error ? error.message : String(error || "Unknown error"); + return { + response: toOpenAiError(502, `Petals fetch error: ${message}`), + url, + headers, + transformedBody, + }; + } + } +} + +export default PetalsExecutor; diff --git a/open-sse/handlers/chatCore.ts b/open-sse/handlers/chatCore.ts index 3a521536b0..e7fbeba6fd 100644 --- a/open-sse/handlers/chatCore.ts +++ b/open-sse/handlers/chatCore.ts @@ -1500,7 +1500,11 @@ export async function handleChatCore({ content?: unknown; }; - const normalizeClaudeUpstreamMessages = (payload: Record) => { + const normalizeClaudeUpstreamMessages = ( + payload: Record, + options?: { preserveToolResultBlocks?: boolean } + ) => { + const preserveToolResultBlocks = options?.preserveToolResultBlocks === true; if (!Array.isArray(payload.messages)) return; const messages = payload.messages as ClaudeMessage[]; @@ -1549,6 +1553,9 @@ export async function handleChatCore({ } if (block.type === "tool_result") { + if (preserveToolResultBlocks) { + return [block]; + } const toolId = block.tool_use_id ?? block.id ?? "unknown"; const resultContent = block.content ?? block.text ?? block.output ?? ""; const resultText = @@ -1626,7 +1633,7 @@ export async function handleChatCore({ // regardless of combo strategy or cache_control settings. translatedBody = { ...body }; translatedBody._disableToolPrefix = true; - normalizeClaudeUpstreamMessages(translatedBody); + normalizeClaudeUpstreamMessages(translatedBody, { preserveToolResultBlocks: true }); log?.debug?.("FORMAT", `claude passthrough (preserveCache=${preserveCacheControl})`); } else { diff --git a/open-sse/services/batchProcessor.ts b/open-sse/services/batchProcessor.ts index c0047a5ad0..a3105ae6c1 100644 --- a/open-sse/services/batchProcessor.ts +++ b/open-sse/services/batchProcessor.ts @@ -422,8 +422,13 @@ async function processBatchItems(batch: BatchRecord, items: BatchRequestItem[]) async function finalizeBatch(batchId: string, results: any[], itemsWithErrors: any[]) { const current = getBatch(batchId); - if (current?.status === "cancelling") { - updateBatch(batchId, { status: "cancelled", cancelledAt: Math.floor(Date.now() / 1000) }); + if (current?.status === "cancelling" || current?.status === "cancelled") { + if (current?.inputFileId) { + updateFileStatus(current.inputFileId, "processed"); + } + if (current?.status === "cancelling") { + updateBatch(batchId, { status: "cancelled", cancelledAt: Math.floor(Date.now() / 1000) }); + } return; } diff --git a/open-sse/services/claudeCodeCompatible.ts b/open-sse/services/claudeCodeCompatible.ts index 4e114f0803..3da884ff9e 100644 --- a/open-sse/services/claudeCodeCompatible.ts +++ b/open-sse/services/claudeCodeCompatible.ts @@ -231,19 +231,25 @@ export function buildClaudeCodeCompatibleRequest({ const preparedClaudeBody = claudeBody ? prepareClaudeCodeCompatibleBody(claudeBody, preserveCacheControl) : null; - const messages = preparedClaudeBody + const normalizedMessages = Array.isArray(normalized.messages) + ? (normalized.messages as MessageLike[]) + : []; + const extractedClaudeBody = + !preparedClaudeBody && sourceBody + ? extractClaudeBodyFromSource(sourceBody, preserveCacheControl) + : null; + const effectiveClaudeBody = preparedClaudeBody || extractedClaudeBody; + const messages = effectiveClaudeBody ? buildClaudeCodeCompatibleMessagesFromClaude( - preparedClaudeBody.messages as MessageLike[], + effectiveClaudeBody.messages as MessageLike[], preserveCacheControl ) - : Array.isArray(normalized.messages) - ? buildClaudeCodeCompatibleMessages(normalized.messages as MessageLike[]) - : []; + : buildClaudeCodeCompatibleMessages(normalizedMessages); const system = buildClaudeCodeCompatibleSystemBlocks({ - messages: normalized.messages as MessageLike[], - systemBlocks: preparedClaudeBody?.system as Record[] | undefined, + messages: normalizedMessages, + systemBlocks: effectiveClaudeBody?.system as Record[] | undefined, preserveCacheControl, - injectDefaultSkeleton: !preparedClaudeBody, + injectDefaultSkeleton: !effectiveClaudeBody, }); const resolvedSessionId = sessionId || randomUUID(); const effort = resolveClaudeCodeCompatibleEffort(sourceBody, normalizedBody, model); @@ -483,17 +489,34 @@ function buildClaudeCodeCompatibleMessagesFromClaude( : []; const merged: Array<{ role: "user" | "assistant"; content: Array> }> = []; + let previousAssistantHadToolUse = false; for (const message of converted) { + const hasToolUse = message.content.some((block) => block.type === "tool_use"); + const hasToolResult = message.content.some((block) => block.type === "tool_result"); const last = merged[merged.length - 1]; - if (last && last.role === message.role) { + const shouldKeepSeparate = + hasToolUse || + hasToolResult || + previousAssistantHadToolUse || + last?.content?.some((block) => block.type === "tool_use") || + last?.content?.some((block) => block.type === "tool_result"); + + if (last && last.role === message.role && !shouldKeepSeparate) { last.content.push(...message.content); - continue; + } else { + merged.push({ role: message.role, content: [...message.content] }); } - merged.push({ role: message.role, content: [...message.content] }); + + previousAssistantHadToolUse = message.role === "assistant" && hasToolUse; } - while (merged.length > 0 && merged[merged.length - 1].role === "assistant") { + while (merged.length > 0) { + const last = merged[merged.length - 1]; + const hasToolUse = last.content.some((block) => block.type === "tool_use"); + if (last.role !== "assistant" || hasToolUse) { + break; + } merged.pop(); } @@ -688,6 +711,40 @@ function prepareClaudeCodeCompatibleBody( return readRecord(prepared); } +function extractClaudeBodyFromSource( + sourceBody: Record, + preserveCacheControl: boolean +): Record | null { + const rawMessages = Array.isArray(sourceBody.messages) + ? (sourceBody.messages as MessageLike[]) + : []; + const hasSystemRoleMessages = rawMessages.some((message) => { + const role = String(message?.role || "").toLowerCase(); + return role === "system" || role === "developer"; + }); + const hasClaudeSystem = + typeof sourceBody.system === "string" || + (Array.isArray(sourceBody.system) && sourceBody.system.length > 0); + + if (!hasClaudeSystem && !hasSystemRoleMessages) { + return null; + } + + const normalizedMessages = rawMessages.filter((message) => { + const role = String(message?.role || "").toLowerCase(); + return role !== "system" && role !== "developer"; + }); + + return prepareClaudeCodeCompatibleBody( + { + ...sourceBody, + ...(hasClaudeSystem ? {} : { system: extractCustomSystemBlocks(rawMessages) }), + messages: normalizedMessages, + }, + preserveCacheControl + ); +} + function normalizeClaudeSystemInput(system: unknown) { if (typeof system === "string") { const text = system.trim(); diff --git a/open-sse/services/rateLimitManager.ts b/open-sse/services/rateLimitManager.ts index b4d83961e0..e1ec39dd4c 100644 --- a/open-sse/services/rateLimitManager.ts +++ b/open-sse/services/rateLimitManager.ts @@ -80,6 +80,78 @@ function buildLimiterDefaults() { }; } +function updateAllLimiterSettings() { + for (const limiter of limiters.values()) { + limiter.updateSettings({ + maxConcurrent: currentRequestQueueSettings.concurrentRequests, + minTime: currentRequestQueueSettings.minTimeBetweenRequestsMs, + maxWait: currentRequestQueueSettings.maxWaitMs, + reservoir: currentRequestQueueSettings.requestsPerMinute, + reservoirRefreshAmount: currentRequestQueueSettings.requestsPerMinute, + reservoirRefreshInterval: 60 * 1000, + }); + } +} + +function reconcileEnabledConnections( + connectionsRaw: unknown[], + requestQueueSettings: RequestQueueSettings +) { + const nextEnabledConnections = new Set(); + let explicitCount = 0; + let autoCount = 0; + + for (const connRaw of connectionsRaw) { + const conn = toRecord(connRaw); + const connectionId = typeof conn.id === "string" ? conn.id : ""; + const provider = typeof conn.provider === "string" ? conn.provider : ""; + const isActive = conn.isActive === true; + const rateLimitProtection = conn.rateLimitProtection === true; + if (!connectionId || !provider) continue; + + if (rateLimitProtection) { + nextEnabledConnections.add(connectionId); + explicitCount++; + continue; + } + + if ( + requestQueueSettings.autoEnableApiKeyProviders && + getProviderCategory(provider) === "apikey" && + isActive + ) { + nextEnabledConnections.add(connectionId); + autoCount++; + + const key = `${provider}:${connectionId}`; + if (!limiters.has(key)) { + limiters.set( + key, + new Bottleneck({ + ...buildLimiterDefaults(), + id: key, + }) + ); + } + } + } + + for (const connectionId of Array.from(enabledConnections)) { + if (!nextEnabledConnections.has(connectionId)) { + disableRateLimitProtection(connectionId); + } + } + + for (const connectionId of nextEnabledConnections) { + enabledConnections.add(connectionId); + } + + return { + explicitCount, + autoCount, + }; +} + function trackAsyncOperation(promise: Promise): Promise { pendingAsyncOperations.add(promise); promise.finally(() => { @@ -100,44 +172,12 @@ export async function initializeRateLimits() { const { getProviderConnections, getSettings } = await import("@/lib/localDb"); const [connections, settings] = await Promise.all([getProviderConnections(), getSettings()]); const resilience = resolveResilienceSettings(settings); - applyRequestQueueSettings(resilience.requestQueue); - let explicitCount = 0; - let autoCount = 0; - - for (const connRaw of connections as unknown[]) { - const conn = toRecord(connRaw); - const connectionId = typeof conn.id === "string" ? conn.id : ""; - const provider = typeof conn.provider === "string" ? conn.provider : ""; - const isActive = conn.isActive === true; - const rateLimitProtection = conn.rateLimitProtection === true; - if (!connectionId || !provider) continue; - - if (rateLimitProtection) { - // Explicitly enabled by user - enabledConnections.add(connectionId); - explicitCount++; - } else if ( - resilience.requestQueue.autoEnableApiKeyProviders && - getProviderCategory(provider) === "apikey" && - isActive - ) { - // Auto-enable for API key providers (safety net) - enabledConnections.add(connectionId); - autoCount++; - - // Create a pre-configured limiter with conservative defaults - const key = `${provider}:${connectionId}`; - if (!limiters.has(key)) { - limiters.set( - key, - new Bottleneck({ - ...buildLimiterDefaults(), - id: key, - }) - ); - } - } - } + currentRequestQueueSettings = { ...resilience.requestQueue }; + const { explicitCount, autoCount } = reconcileEnabledConnections( + connections as unknown[], + currentRequestQueueSettings + ); + updateAllLimiterSettings(); if (explicitCount > 0 || autoCount > 0) { console.log( @@ -152,19 +192,12 @@ export async function initializeRateLimits() { } } -export function applyRequestQueueSettings(nextSettings: RequestQueueSettings) { +export async function applyRequestQueueSettings(nextSettings: RequestQueueSettings) { currentRequestQueueSettings = { ...nextSettings }; - - for (const limiter of limiters.values()) { - limiter.updateSettings({ - maxConcurrent: currentRequestQueueSettings.concurrentRequests, - minTime: currentRequestQueueSettings.minTimeBetweenRequestsMs, - maxWait: currentRequestQueueSettings.maxWaitMs, - reservoir: currentRequestQueueSettings.requestsPerMinute, - reservoirRefreshAmount: currentRequestQueueSettings.requestsPerMinute, - reservoirRefreshInterval: 60 * 1000, - }); - } + const { getProviderConnections } = await import("@/lib/localDb"); + const connections = await getProviderConnections(); + reconcileEnabledConnections(connections as unknown[], currentRequestQueueSettings); + updateAllLimiterSettings(); } /** diff --git a/src/app/(dashboard)/dashboard/providers/[id]/page.tsx b/src/app/(dashboard)/dashboard/providers/[id]/page.tsx index e6454826f4..9d2c5ee37d 100644 --- a/src/app/(dashboard)/dashboard/providers/[id]/page.tsx +++ b/src/app/(dashboard)/dashboard/providers/[id]/page.tsx @@ -5336,6 +5336,7 @@ const CONFIGURABLE_BASE_URL_PROVIDERS = new Set([ "databricks", "snowflake", "searxng-search", + "petals", ]); const DEFAULT_PROVIDER_BASE_URLS: Record = { @@ -5343,6 +5344,7 @@ const DEFAULT_PROVIDER_BASE_URLS: Record = { "bailian-coding-plan": "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic/v1", "xiaomi-mimo": "https://token-plan-ams.xiaomimimo.com/v1", "searxng-search": "http://localhost:8888/search", + petals: "https://chat.petals.dev/api/v1/generate", }; function getLocalProviderMetadata(providerId?: string | null) { @@ -5486,7 +5488,8 @@ function AddApiKeyModal({ const isBlackboxWeb = provider === "blackbox-web"; const isMuseSparkWeb = provider === "muse-spark-web"; const isWebSessionProvider = isGrokWeb || isPerplexityWeb || isBlackboxWeb || isMuseSparkWeb; - const apiKeyOptional = isSearxng || isLocalSelfHostedProvider; + const isPetals = provider === "petals"; + const apiKeyOptional = isSearxng || isPetals || isLocalSelfHostedProvider; const [formData, setFormData] = useState({ name: "", @@ -5546,7 +5549,7 @@ function AddApiKeyModal({ ? t("localProviderApiKeyOptionalHint", { provider: localProviderMetadata?.name || providerName || provider || "", }) - : isSearxng + : isSearxng || isPetals ? t("apiKeyOptionalHint") : undefined; @@ -5975,14 +5978,15 @@ function EditConnectionModal({ isOpen, connection, onSave, onClose }: EditConnec const isLocalSelfHostedProvider = !!localProviderMetadata; const isSearxng = connection?.provider === "searxng-search"; const isGooglePse = connection?.provider === "google-pse-search"; - const apiKeyOptional = isSearxng || isLocalSelfHostedProvider; + const isPetals = connection?.provider === "petals"; + const apiKeyOptional = isSearxng || isPetals || isLocalSelfHostedProvider; const isCcCompatible = isClaudeCodeCompatibleProvider(connection?.provider); const defaultRegion = "us-central1"; const apiCredentialHint = isLocalSelfHostedProvider ? t("localProviderApiKeyOptionalHint", { provider: localProviderMetadata?.name || connection?.provider || "", }) - : isSearxng + : isSearxng || isPetals ? t("apiKeyOptionalHint") : t("leaveBlankKeepCurrentApiKey"); diff --git a/src/app/api/providers/[id]/models/route.ts b/src/app/api/providers/[id]/models/route.ts index ac24b454b8..25fa0b413e 100755 --- a/src/app/api/providers/[id]/models/route.ts +++ b/src/app/api/providers/[id]/models/route.ts @@ -80,7 +80,14 @@ function isLocalOpenAIStyleProvider(provider: string): boolean { return isSelfHostedChatProvider(provider); } -const NAMED_OPENAI_STYLE_PROVIDERS = new Set(["bedrock", "modal", "reka", "empower", "poe"]); +const NAMED_OPENAI_STYLE_PROVIDERS = new Set([ + "bedrock", + "modal", + "reka", + "empower", + "nous-research", + "poe", +]); function isNamedOpenAIStyleProvider(provider: string): boolean { return NAMED_OPENAI_STYLE_PROVIDERS.has(provider); @@ -792,8 +799,8 @@ export async function GET( }; const buildApiDiscoveryResponse = async (models: any[]) => { - if (Array.isArray(models) && models.length > 0) { - await persistDiscoveredModels(provider, connectionId, models); + const discoveredModels = await persistDiscoveredModels(provider, connectionId, models); + if (discoveredModels.length > 0) { return buildResponse({ provider, connectionId, @@ -802,10 +809,9 @@ export async function GET( }); } - const fallback = buildDiscoveryFallbackResponse({ - cacheWarning: "No remote models discovered — using cached catalog", - localWarning: "No remote models discovered — using local catalog", - }); + const fallback = buildLocalCatalogResponse( + "No remote models discovered — using local catalog" + ); if (fallback) return fallback; return buildResponse({ diff --git a/src/app/api/resilience/route.ts b/src/app/api/resilience/route.ts index 9436aae3cf..08fe0aeb05 100644 --- a/src/app/api/resilience/route.ts +++ b/src/app/api/resilience/route.ts @@ -111,7 +111,7 @@ function normalizeLegacyPatch(body: JsonRecord): ResilienceSettingsPatch { async function syncRuntimeSettings(resilienceSettings: ResilienceSettings) { const { applyRequestQueueSettings } = await import("@omniroute/open-sse/services/rateLimitManager"); - applyRequestQueueSettings(resilienceSettings.requestQueue); + await applyRequestQueueSettings(resilienceSettings.requestQueue); } /** diff --git a/src/lib/providerModels/modelDiscovery.ts b/src/lib/providerModels/modelDiscovery.ts index d14546ef4d..0ed9dbcca5 100644 --- a/src/lib/providerModels/modelDiscovery.ts +++ b/src/lib/providerModels/modelDiscovery.ts @@ -77,7 +77,6 @@ export async function persistDiscoveredModels( models: unknown ): Promise { const normalized = normalizeDiscoveredModels(models); - if (normalized.length === 0) return []; await replaceSyncedAvailableModelsForConnection(providerId, connectionId, normalized); return normalized; } diff --git a/src/lib/providers/validation.ts b/src/lib/providers/validation.ts index 328526ae98..7ca8cd6cd4 100644 --- a/src/lib/providers/validation.ts +++ b/src/lib/providers/validation.ts @@ -66,6 +66,7 @@ import { buildRunwayHeaders, normalizeRunwayBaseUrl, } from "@omniroute/open-sse/config/runway.ts"; +import { PETALS_DEFAULT_MODEL, normalizePetalsBaseUrl } from "@omniroute/open-sse/config/petals.ts"; const OPENAI_LIKE_FORMATS = new Set(["openai", "openai-responses"]); const GEMINI_LIKE_FORMATS = new Set(["gemini", "gemini-cli"]); @@ -1579,6 +1580,119 @@ async function validateRunwayProvider({ apiKey, providerSpecificData = {} }: any return { valid: false, error: "Connection failed while testing Runway" }; } +async function validatePetalsProvider({ apiKey, providerSpecificData = {} }: any) { + const url = normalizePetalsBaseUrl(providerSpecificData.baseUrl); + const modelId = + typeof providerSpecificData.validationModelId === "string" && + providerSpecificData.validationModelId.trim() + ? providerSpecificData.validationModelId.trim() + : PETALS_DEFAULT_MODEL; + const headers: Record = { + "Content-Type": "application/x-www-form-urlencoded", + }; + if (apiKey) { + headers.Authorization = `Bearer ${apiKey}`; + } + + const body = new URLSearchParams({ + model: modelId, + inputs: "test", + max_new_tokens: "1", + }); + + try { + const response = await validationWrite(url, { + method: "POST", + headers, + body: body.toString(), + }); + + if (response.ok) { + const payload = (await response.json().catch(() => ({}))) as Record; + if (payload.ok === false) { + return { + valid: false, + error: "Petals API rejected validation request", + }; + } + return { valid: true, error: null, method: "petals_generate" }; + } + + if (response.status === 401 || response.status === 403) { + return { valid: false, error: "Invalid API key" }; + } + + if (response.status === 429) { + return { + valid: true, + error: null, + method: "petals_generate", + warning: "Rate limited, but endpoint is reachable", + }; + } + + if (response.status >= 500) { + return { valid: false, error: `Provider unavailable (${response.status})` }; + } + } catch (error: any) { + return toValidationErrorResult(error); + } + + return { valid: false, error: "Connection failed while testing Petals" }; +} + +async function validateNousResearchProvider({ apiKey, providerSpecificData = {} }: any) { + const baseUrl = + normalizeBaseUrl(providerSpecificData.baseUrl) || "https://inference-api.nousresearch.com/v1"; + const chatUrl = `${baseUrl}/chat/completions`; + const modelId = + typeof providerSpecificData.validationModelId === "string" && + providerSpecificData.validationModelId.trim() + ? providerSpecificData.validationModelId.trim() + : "nousresearch/hermes-4-70b"; + + try { + const response = await validationWrite(chatUrl, { + method: "POST", + headers: buildBearerHeaders(apiKey, providerSpecificData), + body: JSON.stringify({ + model: modelId, + messages: [{ role: "user", content: "test" }], + max_tokens: 1, + }), + }); + + if (response.ok) { + return { valid: true, error: null, method: "nous_chat_completions" }; + } + + if (response.status === 401 || response.status === 403) { + return { valid: false, error: "Invalid API key" }; + } + + if (response.status === 429) { + return { + valid: true, + error: null, + method: "nous_chat_completions", + warning: "Rate limited, but credentials are valid", + }; + } + + if (response.status === 402) { + return { valid: false, error: "Payment required or API key missing" }; + } + + if (response.status >= 500) { + return { valid: false, error: `Provider unavailable (${response.status})` }; + } + } catch (error: any) { + return toValidationErrorResult(error); + } + + return { valid: false, error: "Connection failed while testing Nous Research" }; +} + async function validatePoeProvider({ apiKey, providerSpecificData = {} }: any) { const baseUrl = normalizeBaseUrl(providerSpecificData.baseUrl) || "https://api.poe.com/v1"; const balanceUrl = new URL("/usage/current_balance", baseUrl).toString(); @@ -2487,7 +2601,8 @@ async function validateMuseSparkWebProvider({ apiKey, providerSpecificData = {} } export async function validateProviderApiKey({ provider, apiKey, providerSpecificData = {} }: any) { - const requiresApiKey = provider !== "searxng-search" && !isSelfHostedChatProvider(provider); + const requiresApiKey = + provider !== "searxng-search" && provider !== "petals" && !isSelfHostedChatProvider(provider); if (!provider || (requiresApiKey && !apiKey)) { return { valid: false, error: "Provider and API key required", unsupported: false }; } @@ -2548,6 +2663,8 @@ export async function validateProviderApiKey({ provider, apiKey, providerSpecifi baseUrl: normalizeBaseUrl(providerSpecificData?.baseUrl || ""), modelId: "Qwen/Qwen3-4B-Thinking-2507-FP8", }), + "nous-research": validateNousResearchProvider, + petals: validatePetalsProvider, poe: validatePoeProvider, clarifai: validateClarifaiProvider, reka: validateRekaProvider, diff --git a/src/shared/constants/providers.ts b/src/shared/constants/providers.ts index 9df284e822..eaed4b33ee 100644 --- a/src/shared/constants/providers.ts +++ b/src/shared/constants/providers.ts @@ -1140,6 +1140,33 @@ export const APIKEY_PROVIDERS = { "Empower exposes OpenAI-compatible chat on https://app.empower.dev/api/v1 with tool-calling support on empower-functions.", passthroughModels: true, }, + "nous-research": { + id: "nous-research", + alias: "nous", + name: "Nous Research", + icon: "hub", + color: "#2563EB", + textIcon: "NO", + website: "https://portal.nousresearch.com/help", + authHint: + "Use your Nous Portal API key. OmniRoute targets the official OpenAI-compatible inference endpoint at https://inference-api.nousresearch.com/v1.", + apiHint: + "Nous exposes an OpenAI-compatible /v1 surface with a large remote /models catalog. The /chat/completions endpoint requires a valid API key for programmatic inference.", + passthroughModels: true, + }, + petals: { + id: "petals", + alias: "petals", + name: "Petals", + icon: "hub", + color: "#10B981", + textIcon: "PT", + website: "https://chat.petals.dev", + authHint: + "No API key is required for the public research endpoint. Leave the field blank, or provide a bearer token if your self-hosted Petals gateway uses auth.", + apiHint: + "Petals exposes a public HTTP API at https://chat.petals.dev/api/v1/generate and a WebSocket API at /api/v2/generate. OmniRoute targets the HTTP generate endpoint and supports self-hosted base URLs.", + }, poe: { id: "poe", alias: "poe", diff --git a/src/shared/validation/schemas.ts b/src/shared/validation/schemas.ts index 0adc1f3398..ba31860acf 100644 --- a/src/shared/validation/schemas.ts +++ b/src/shared/validation/schemas.ts @@ -245,7 +245,10 @@ export const createProviderSchema = z }) .superRefine((data, ctx) => { const apiKey = typeof data.apiKey === "string" ? data.apiKey.trim() : ""; - const apiKeyOptional = data.provider === "searxng-search" || isLocalProvider(data.provider); + const apiKeyOptional = + data.provider === "searxng-search" || + data.provider === "petals" || + isLocalProvider(data.provider); if (!apiKeyOptional && apiKey.length === 0) { ctx.addIssue({ code: z.ZodIssueCode.custom, diff --git a/tests/unit/batch_api.test.ts b/tests/unit/batch_api.test.ts index 5a4a323e84..0f7b49bb0a 100644 --- a/tests/unit/batch_api.test.ts +++ b/tests/unit/batch_api.test.ts @@ -844,6 +844,95 @@ test("Batch Cancel API", async () => { assert.strictEqual(canCancel, false); }); +test("Batch processor keeps cancelled status for in-flight batches", async () => { + const originalFetch = globalThis.fetch; + const apiKey = await createApiKey("In Flight Cancel Key", "test-machine"); + + await createProviderConnection({ + provider: "openai", + authType: "apikey", + name: "Cancelable OpenAI", + apiKey: "sk-cancel-batch", + isActive: true, + }); + + const batchItems = [ + JSON.stringify({ + custom_id: "cancel-mid-flight", + method: "POST", + url: "/v1/chat/completions", + body: { + model: "gpt-4o-mini", + messages: [{ role: "user", content: "cancel me" }], + }, + }), + ].join("\n"); + + const file = createFile({ + bytes: Buffer.byteLength(batchItems), + filename: "cancel_mid_flight.jsonl", + purpose: "batch", + content: Buffer.from(batchItems), + apiKeyId: apiKey.id, + }); + + const batch = createBatch({ + endpoint: "/v1/chat/completions", + completionWindow: "24h", + inputFileId: file.id, + apiKeyId: apiKey.id, + }); + + globalThis.fetch = async () => { + updateBatch(batch.id, { + status: "cancelled", + cancelledAt: Math.floor(Date.now() / 1000), + }); + + return Response.json({ + id: "chatcmpl-batch-cancelled", + object: "chat.completion", + model: "gpt-4o-mini", + choices: [ + { + index: 0, + message: { role: "assistant", content: "ok" }, + finish_reason: "stop", + }, + ], + usage: { + prompt_tokens: 1, + completion_tokens: 1, + total_tokens: 2, + completion_tokens_details: { reasoning_tokens: 0 }, + }, + }); + }; + + try { + await processPendingBatches(); + + let currentBatch = getBatch(batch.id); + let remainingAttempts = 40; + while ( + remainingAttempts > 0 && + currentBatch && + !["cancelled", "completed", "failed", "expired"].includes(currentBatch.status) + ) { + await new Promise((resolve) => setTimeout(resolve, 50)); + currentBatch = getBatch(batch.id); + remainingAttempts--; + } + + assert.strictEqual(currentBatch?.status, "cancelled"); + assert.ok(!currentBatch?.outputFileId, "Cancelled batch must not emit an output file"); + assert.ok(!currentBatch?.errorFileId, "Cancelled batch must not emit an error file"); + assert.strictEqual(getFile(file.id)?.status, "processed"); + } finally { + globalThis.fetch = originalFetch; + } +}); + test("List files pagination and response format", async () => { const apiKey = await createApiKey("File List Test Key", "test-machine"); diff --git a/tests/unit/chat-openai-compat-providers.test.ts b/tests/unit/chat-openai-compat-providers.test.ts index f4a35f2654..7e2c96c206 100644 --- a/tests/unit/chat-openai-compat-providers.test.ts +++ b/tests/unit/chat-openai-compat-providers.test.ts @@ -27,6 +27,7 @@ const CHAT_OPENAI_COMPAT_PROVIDER_IDS = [ "databricks", "datarobot", "clarifai", + "nous-research", "poe", "azure-ai", "bedrock", diff --git a/tests/unit/chatcore-sanitization.test.ts b/tests/unit/chatcore-sanitization.test.ts index 5c44835993..be5cc01da2 100644 --- a/tests/unit/chatcore-sanitization.test.ts +++ b/tests/unit/chatcore-sanitization.test.ts @@ -445,6 +445,45 @@ test("chatCore sanitization normalizes mixed content blocks and removes unsuppor ); }); +test("chatCore preserves Claude passthrough tool_result blocks instead of converting them to plain text", async () => { + const { call } = await invokeChatCore({ + endpoint: "/v1/messages", + provider: "claude", + model: "claude-opus-4-7", + userAgent: "claude-cli/2.1.114", + body: { + model: "claude-opus-4-7", + max_tokens: 64, + system: [{ type: "text", text: "sys" }], + messages: [ + { + role: "assistant", + content: [ + { type: "tool_use", id: "toolu_keep", name: "Bash", input: { command: "pwd" } }, + ], + }, + { + role: "user", + content: [{ type: "tool_result", tool_use_id: "toolu_keep", content: "done" }], + }, + ], + tools: [{ name: "Bash", input_schema: { type: "object", properties: {} } }], + }, + }); + + assert.equal(call.body.messages[0].role, "assistant"); + assert.equal(call.body.messages[0].content[0].type, "tool_use"); + assert.equal(call.body.messages[1].role, "user"); + assert.equal(call.body.messages[1].content[0].type, "tool_result"); + assert.equal(call.body.messages[1].content[0].tool_use_id, "toolu_keep"); + assert.equal( + call.body.messages[1].content.some( + (block) => block.type === "text" && /\[Tool Result:/.test(block.text) + ), + false + ); +}); + test("chatCore resolves stream mode from body.stream and Accept header", async () => { const explicitTrue = await invokeChatCore({ accept: "application/json", diff --git a/tests/unit/claude-code-compatible-request.test.ts b/tests/unit/claude-code-compatible-request.test.ts index 1279d384b5..8df02fed1a 100644 --- a/tests/unit/claude-code-compatible-request.test.ts +++ b/tests/unit/claude-code-compatible-request.test.ts @@ -75,9 +75,15 @@ test("Claude Code compatible effort and max token helpers cover priority fallbac ); }); -test("buildClaudeCodeCompatibleRequest covers normalized OpenAI-style messages, source tools and fallback text", () => { +test("buildClaudeCodeCompatibleRequest promotes source system/developer messages into top-level Claude system blocks", () => { const payload = buildClaudeCodeCompatibleRequest({ sourceBody: { + messages: [ + { role: "system", content: "system note" }, + { role: "developer", content: [{ type: "text", text: "developer note" }] }, + { role: "assistant", content: [{ type: "text", text: "draft answer" }] }, + { role: "tool", content: "ignored" }, + ], tools: [ null, { @@ -111,13 +117,12 @@ test("buildClaudeCodeCompatibleRequest covers normalized OpenAI-style messages, assert.deepEqual(payload.messages, [ { role: "user", - content: [{ type: "text", text: "draft answer\nalternate answer" }], + content: [{ type: "text", text: "draft answer" }], }, ]); - assert.equal(payload.system.length, 3); - assert.match((payload as any).system[0].text, /Claude Agent SDK/); - assert.equal(payload.system[1].text, "system note"); - assert.equal(payload.system[2].text, "developer note"); + assert.equal(payload.system.length, 2); + assert.equal(payload.system[0].text, "system note"); + assert.equal(payload.system[1].text, "developer note"); assert.equal(payload.tools.length, 1); assert.deepEqual(payload.tools[0], { name: "lookup_account", @@ -130,6 +135,33 @@ test("buildClaudeCodeCompatibleRequest covers normalized OpenAI-style messages, assert.equal(payload.max_tokens, 19); }); +test("buildClaudeCodeCompatibleRequest prefers existing Claude top-level system over extracted source messages", () => { + const payload = buildClaudeCodeCompatibleRequest({ + sourceBody: { + system: [{ type: "text", text: "top-level system" }], + messages: [ + { role: "system", content: "stale system message" }, + { role: "user", content: "hello" }, + ], + }, + normalizedBody: { + messages: [ + { role: "system", content: "stale system message" }, + { role: "user", content: "hello" }, + ], + }, + model: "claude-sonnet-4-6", + cwd: "/tmp/claude-code-compatible", + now: new Date("2026-01-02T12:00:00.000Z"), + }); + + assert.equal(payload.system.length, 1); + assert.equal((payload.system[0] as any).text, "top-level system"); + assert.deepEqual(payload.messages, [ + { role: "user", content: [{ type: "text", text: "hello" }] }, + ]); +}); + test("buildClaudeCodeCompatibleRequest covers Claude-native bodies and cache-control stripping", () => { const stripped = buildClaudeCodeCompatibleRequest({ claudeBody: { @@ -207,6 +239,77 @@ test("buildClaudeCodeCompatibleRequest covers Claude-native bodies and cache-con assert.equal((preserved.tools[0].cache_control as any).type, "ephemeral"); }); +test("buildClaudeCodeCompatibleRequest keeps the next user message anchored by matching tool_result after assistant tool_use", () => { + const payload = buildClaudeCodeCompatibleRequest({ + claudeBody: { + messages: [ + { + role: "assistant", + content: [ + { type: "text", text: "Calling tool" }, + { type: "tool_use", id: "toolu_123", name: "Bash", input: { command: "pwd" } }, + ], + }, + { + role: "user", + content: [ + { + type: "tool_result", + tool_use_id: "toolu_123", + content: [{ type: "text", text: "/tmp" }], + }, + ], + }, + { + role: "user", + content: [{ type: "text", text: "Continue" }], + }, + ], + tools: [{ name: "Bash", input_schema: { type: "object", properties: {} } }], + }, + model: "claude-sonnet-4-6", + cwd: "/tmp/claude-code-compatible", + now: new Date("2026-01-02T12:00:00.000Z"), + }); + + assert.equal(payload.messages.length, 2); + assert.equal(payload.messages[0].role, "assistant"); + assert.equal((payload.messages[0].content[1] as any).type, "tool_use"); + assert.equal(payload.messages[1].role, "user"); + assert.equal((payload.messages[1].content[0] as any).type, "tool_result"); + assert.equal((payload.messages[1].content[0] as any).tool_use_id, "toolu_123"); + assert.equal((payload.messages[1].content[1] as any).type, "text"); + assert.equal((payload.messages[1].content[1] as any).text, "Continue"); +}); + +test("buildClaudeCodeCompatibleRequest preserves trailing assistant tool_use message awaiting next tool_result", () => { + const payload = buildClaudeCodeCompatibleRequest({ + claudeBody: { + messages: [ + { + role: "user", + content: [{ type: "text", text: "Run pwd" }], + }, + { + role: "assistant", + content: [ + { type: "tool_use", id: "toolu_pending", name: "Bash", input: { command: "pwd" } }, + ], + }, + ], + tools: [{ name: "Bash", input_schema: { type: "object", properties: {} } }], + }, + model: "claude-sonnet-4-6", + cwd: "/tmp/claude-code-compatible", + now: new Date("2026-01-02T12:00:00.000Z"), + }); + + assert.equal(payload.messages.length, 2); + assert.equal(payload.messages[1].role, "assistant"); + assert.equal((payload.messages[1].content[0] as any).type, "tool_use"); + assert.equal((payload.messages[1].content[0] as any).id, "toolu_pending"); +}); + test("buildClaudeCodeCompatibleRequest omits tool choice when there are no tools", () => { const payload = buildClaudeCodeCompatibleRequest({ normalizedBody: { diff --git a/tests/unit/executor-petals.test.ts b/tests/unit/executor-petals.test.ts new file mode 100644 index 0000000000..d56306ea76 --- /dev/null +++ b/tests/unit/executor-petals.test.ts @@ -0,0 +1,148 @@ +import test from "node:test"; +import assert from "node:assert/strict"; + +import { getExecutor, hasSpecializedExecutor } from "../../open-sse/executors/index.ts"; +import { PetalsExecutor } from "../../open-sse/executors/petals.ts"; + +function jsonResponse(body: unknown, status = 200) { + return new Response(JSON.stringify(body), { + status, + headers: { "Content-Type": "application/json" }, + }); +} + +test("PetalsExecutor is registered in the executor index", () => { + assert.equal(hasSpecializedExecutor("petals"), true); + assert.ok(getExecutor("petals") instanceof PetalsExecutor); +}); + +test("PetalsExecutor converts OpenAI messages into form data and wraps JSON responses", async () => { + const executor = new PetalsExecutor(); + const originalFetch = globalThis.fetch; + const calls: Array<{ + url: string; + body: URLSearchParams; + headers: Record; + }> = []; + + globalThis.fetch = async (url, init = {}) => { + calls.push({ + url: String(url), + body: new URLSearchParams(String(init.body || "")), + headers: init.headers as Record, + }); + + return jsonResponse({ + ok: true, + outputs: "Hi back from Petals.", + }); + }; + + try { + const result = await executor.execute({ + model: "stabilityai/StableBeluga2", + body: { + messages: [ + { role: "system", content: "You are concise." }, + { role: "user", content: "Hello" }, + { role: "assistant", content: "Hi there!" }, + { role: "user", content: "How are you?" }, + ], + max_tokens: 32, + temperature: 0.7, + top_p: 0.9, + }, + stream: false, + credentials: { apiKey: "" }, + signal: AbortSignal.timeout(10_000), + log: null, + }); + + assert.equal(calls.length, 1); + assert.equal(calls[0].url, "https://chat.petals.dev/api/v1/generate"); + assert.equal(calls[0].headers.Authorization, undefined); + assert.equal(calls[0].headers["Content-Type"], "application/x-www-form-urlencoded"); + assert.equal(calls[0].body.get("model"), "stabilityai/StableBeluga2"); + assert.equal(calls[0].body.get("max_new_tokens"), "32"); + assert.equal(calls[0].body.get("temperature"), "0.7"); + assert.equal(calls[0].body.get("top_p"), "0.9"); + assert.match( + calls[0].body.get("inputs") || "", + /System:\nYou are concise\.\n\nUser: Hello\n\nAssistant: Hi there!\n\nUser: How are you\?\n\nAssistant:/ + ); + + assert.deepEqual(result.transformedBody, { + model: "stabilityai/StableBeluga2", + inputs: + "System:\nYou are concise.\n\nUser: Hello\n\nAssistant: Hi there!\n\nUser: How are you?\n\nAssistant:", + max_new_tokens: "32", + do_sample: "1", + temperature: "0.7", + top_p: "0.9", + }); + + const body = (await result.response.json()) as any; + assert.equal(body.object, "chat.completion"); + assert.equal(body.choices[0].message.role, "assistant"); + assert.equal(body.choices[0].message.content, "Hi back from Petals."); + assert.equal(body.model, "stabilityai/StableBeluga2"); + } finally { + globalThis.fetch = originalFetch; + } +}); + +test("PetalsExecutor synthesizes OpenAI-compatible SSE responses for streaming requests", async () => { + const executor = new PetalsExecutor(); + const originalFetch = globalThis.fetch; + + globalThis.fetch = async () => + jsonResponse({ + ok: true, + outputs: "Petals stream output", + }); + + try { + const result = await executor.execute({ + model: "stabilityai/StableBeluga2", + body: { + messages: [{ role: "user", content: "Say hello" }], + }, + stream: true, + credentials: { apiKey: "" }, + signal: AbortSignal.timeout(10_000), + log: null, + }); + + assert.equal(result.response.headers.get("Content-Type"), "text/event-stream"); + const text = await result.response.text(); + assert.match(text, /data: \{\"id\":\"chatcmpl-petals-/); + assert.match(text, /Petals stream output/); + assert.match(text, /data: \[DONE\]/); + } finally { + globalThis.fetch = originalFetch; + } +}); + +test("PetalsExecutor maps upstream failures to OpenAI-style errors", async () => { + const executor = new PetalsExecutor(); + const originalFetch = globalThis.fetch; + + globalThis.fetch = async () => jsonResponse({ ok: false, traceback: "petals exploded" }, 200); + + try { + const result = await executor.execute({ + model: "stabilityai/StableBeluga2", + body: { messages: [{ role: "user", content: "hi" }] }, + stream: false, + credentials: { apiKey: "" }, + signal: AbortSignal.timeout(10_000), + log: null, + }); + + assert.equal(result.response.status, 502); + const body = (await result.response.json()) as any; + assert.match(body.error.message, /Petals API error: petals exploded/); + } finally { + globalThis.fetch = originalFetch; + } +}); diff --git a/tests/unit/provider-models-config.test.ts b/tests/unit/provider-models-config.test.ts index 6e7e106f67..b69663a820 100644 --- a/tests/unit/provider-models-config.test.ts +++ b/tests/unit/provider-models-config.test.ts @@ -117,6 +117,23 @@ test("Poe registry exposes current OpenAI-compatible examples", () => { assert.ok(ids.has("Gemini-2.5-Pro")); }); +test("Nous Research registry exposes current official Hermes examples", () => { + const nousModels = getProviderModels("nous"); + const ids = new Set(nousModels.map((model) => model.id)); + + assert.ok(ids.has("nousresearch/hermes-4-70b")); + assert.ok(ids.has("nousresearch/hermes-4-405b")); + assert.ok(ids.has("nousresearch/hermes-3-llama-3.1-70b")); + assert.ok(ids.has("nousresearch/hermes-3-llama-3.1-405b")); +}); + +test("Petals registry exposes the current public HTTP API fallback model", () => { + const petalsModels = getProviderModels("petals"); + const ids = new Set(petalsModels.map((model) => model.id)); + + assert.ok(ids.has("stabilityai/StableBeluga2")); +}); + test("Azure AI Foundry registry exposes fallback marketplace examples", () => { const azureAiModels = getProviderModels("azure-ai"); const ids = new Set(azureAiModels.map((model) => model.id)); diff --git a/tests/unit/provider-models-route.test.ts b/tests/unit/provider-models-route.test.ts index 10061fab55..1e82627550 100644 --- a/tests/unit/provider-models-route.test.ts +++ b/tests/unit/provider-models-route.test.ts @@ -341,6 +341,16 @@ test("provider models route fetches remote catalogs for new OpenAI-compatible ga expectedUrl: "https://app.empower.dev/api/v1/models", model: { id: "empower-functions", name: "Empower Functions", owned_by: "empower" }, }, + { + provider: "nous-research", + apiKey: "nous-key", + expectedUrl: "https://inference-api.nousresearch.com/v1/models", + model: { + id: "nousresearch/hermes-4-70b", + name: "Nous: Hermes 4 70B", + owned_by: "nous-research", + }, + }, { provider: "poe", apiKey: "poe-key", @@ -425,6 +435,20 @@ test("provider models route returns the local catalog for NLP Cloud", async () = assert.ok(body.models.some((model) => model.id === "dolphin-mixtral-8x7b")); }); +test("provider models route returns the local catalog for Petals", async () => { + const connection = await seedConnection("petals", { + apiKey: null, + }); + + const response = await callRoute(connection.id); + const body = (await response.json()) as any; + + assert.equal(response.status, 200); + assert.equal(body.provider, "petals"); + assert.equal(body.source, "local_catalog"); + assert.ok(body.models.some((model) => model.id === "stabilityai/StableBeluga2")); +}); + test("provider models route returns the local catalog for Runway video models", async () => { const connection = await seedConnection("runwayml", { apiKey: "runway-key", @@ -583,6 +607,32 @@ test("provider models route falls back to cached models when a refresh fails", a assert.equal(fetchCalls, 1); }); +test("provider models route clears cached discovery when a refresh returns no remote models", async () => { + const connection = await seedConnection("opencode-go", { + apiKey: "opencode-go-key", + }); + await modelsDb.replaceSyncedAvailableModelsForConnection("opencode-go", connection.id, [ + { id: "cached-go", name: "Cached Go", source: "api-sync" }, + ]); + + globalThis.fetch = async () => { + return Response.json({ data: [] }); + }; + + const response = await callRoute(connection.id, "?refresh=true"); + const body = (await response.json()) as any; + const cachedModels = await modelsDb.getSyncedAvailableModelsForConnection( + "opencode-go", + connection.id + ); + + assert.equal(response.status, 200); + assert.equal(body.source, "local_catalog"); + assert.match(body.warning, /no remote models discovered/i); + assert.ok(body.models.every((model) => model.id !== "cached-go")); + assert.deepEqual(cachedModels, []); +}); + test("provider models route honors autoFetchModels=false and skips remote discovery", async () => { const connection = await seedConnection("opencode-go", { apiKey: "opencode-go-key", diff --git a/tests/unit/provider-validation-specialty.test.ts b/tests/unit/provider-validation-specialty.test.ts index e88a53bfc9..6e6ec8274e 100644 --- a/tests/unit/provider-validation-specialty.test.ts +++ b/tests/unit/provider-validation-specialty.test.ts @@ -1211,6 +1211,105 @@ test("specialty validator accepts Poe credentials on the current balance endpoin assert.equal(poe.method, "poe_current_balance"); }); +test("specialty validator accepts Nous Research credentials on chat completions", async () => { + globalThis.fetch = async (url, init = {}) => { + const target = String(url); + + if (target === "https://inference-api.nousresearch.com/v1/chat/completions") { + const headers = init.headers as Record; + const body = JSON.parse(String(init.body)); + assert.equal(headers.Authorization, "Bearer nous-key"); + assert.equal(body.model, "nousresearch/hermes-4-70b"); + return new Response( + JSON.stringify({ + id: "chatcmpl-nous", + choices: [{ message: { role: "assistant", content: "ok" } }], + }), + { status: 200 } + ); + } + + throw new Error(`unexpected fetch: ${target}`); + }; + + const nous = await validateProviderApiKey({ + provider: "nous-research", + apiKey: "nous-key", + }); + + assert.equal(nous.valid, true); + assert.equal(nous.method, "nous_chat_completions"); +}); + +test("specialty validator rejects invalid Nous Research credentials", async () => { + globalThis.fetch = async (url, init = {}) => { + const target = String(url); + + if (target === "https://inference-api.nousresearch.com/v1/chat/completions") { + const headers = init.headers as Record; + assert.equal(headers.Authorization, "Bearer nous-bad"); + return new Response(JSON.stringify({ message: "invalid" }), { status: 401 }); + } + + throw new Error(`unexpected fetch: ${target}`); + }; + + const nous = await validateProviderApiKey({ + provider: "nous-research", + apiKey: "nous-bad", + }); + + assert.equal(nous.error, "Invalid API key"); +}); + +test("specialty validator accepts the public Petals generate endpoint without an API key", async () => { + globalThis.fetch = async (url, init = {}) => { + const target = String(url); + + if (target === "https://chat.petals.dev/api/v1/generate") { + const headers = init.headers as Record; + const body = new URLSearchParams(String(init.body)); + assert.equal(headers.Authorization, undefined); + assert.equal(headers["Content-Type"], "application/x-www-form-urlencoded"); + assert.equal(body.get("model"), "stabilityai/StableBeluga2"); + assert.equal(body.get("inputs"), "test"); + assert.equal(body.get("max_new_tokens"), "1"); + return new Response(JSON.stringify({ ok: true, outputs: "hi" }), { status: 200 }); + } + + throw new Error(`unexpected fetch: ${target}`); + }; + + const petals = await validateProviderApiKey({ + provider: "petals", + apiKey: "", + }); + + assert.equal(petals.valid, true); + assert.equal(petals.method, "petals_generate"); +}); + +test("specialty validator surfaces Petals upstream unavailability", async () => { + globalThis.fetch = async (url, init = {}) => { + const target = String(url); + + if (target === "https://chat.petals.dev/api/v1/generate") { + const headers = init.headers as Record; + assert.equal(headers.Authorization, undefined); + return new Response(JSON.stringify({ error: "unavailable" }), { status: 503 }); + } + + throw new Error(`unexpected fetch: ${target}`); + }; + + const petals = await validateProviderApiKey({ + provider: "petals", + apiKey: "", + }); + + assert.equal(petals.error, "Provider unavailable (503)"); +}); + test("specialty validator rejects invalid Poe credentials", async () => { globalThis.fetch = async (url, init = {}) => { const target = String(url); diff --git a/tests/unit/providers-page-utils.test.ts b/tests/unit/providers-page-utils.test.ts index 2a7476b97c..72b2880c6e 100644 --- a/tests/unit/providers-page-utils.test.ts +++ b/tests/unit/providers-page-utils.test.ts @@ -229,6 +229,8 @@ test("static catalog entries resolve local, search, audio, web-cookie and upstre const datarobotProvider = providerPageUtils.resolveDashboardProviderInfo("datarobot"); const clarifaiProvider = providerPageUtils.resolveDashboardProviderInfo("clarifai"); const empowerProvider = providerPageUtils.resolveDashboardProviderInfo("empower"); + const nousProvider = providerPageUtils.resolveDashboardProviderInfo("nous-research"); + const petalsProvider = providerPageUtils.resolveDashboardProviderInfo("petals"); const poeProvider = providerPageUtils.resolveDashboardProviderInfo("poe"); const azureAiProvider = providerPageUtils.resolveDashboardProviderInfo("azure-ai"); const watsonxProvider = providerPageUtils.resolveDashboardProviderInfo("watsonx"); @@ -275,6 +277,10 @@ test("static catalog entries resolve local, search, audio, web-cookie and upstre assert.equal(clarifaiProvider?.name, providers.APIKEY_PROVIDERS.clarifai.name); assert.equal(empowerProvider?.category, "apikey"); assert.equal(empowerProvider?.name, providers.APIKEY_PROVIDERS.empower.name); + assert.equal(nousProvider?.category, "apikey"); + assert.equal(nousProvider?.name, providers.APIKEY_PROVIDERS["nous-research"].name); + assert.equal(petalsProvider?.category, "apikey"); + assert.equal(petalsProvider?.name, providers.APIKEY_PROVIDERS.petals.name); assert.equal(poeProvider?.category, "apikey"); assert.equal(poeProvider?.name, providers.APIKEY_PROVIDERS.poe.name); assert.equal(azureAiProvider?.category, "apikey"); @@ -330,6 +336,8 @@ test("managed provider connection ids include supported static categories and ex assert.equal(providerCatalog.isManagedProviderConnectionId("datarobot"), true); assert.equal(providerCatalog.isManagedProviderConnectionId("clarifai"), true); assert.equal(providerCatalog.isManagedProviderConnectionId("empower"), true); + assert.equal(providerCatalog.isManagedProviderConnectionId("nous-research"), true); + assert.equal(providerCatalog.isManagedProviderConnectionId("petals"), true); assert.equal(providerCatalog.isManagedProviderConnectionId("poe"), true); assert.equal(providerCatalog.isManagedProviderConnectionId("azure-ai"), true); assert.equal(providerCatalog.isManagedProviderConnectionId("bedrock"), true); @@ -381,6 +389,8 @@ test("grok-web taxonomy stays web-cookie only and does not leak into api-key ent assert.equal("datarobot" in providers.APIKEY_PROVIDERS, true); assert.equal("clarifai" in providers.APIKEY_PROVIDERS, true); assert.equal("empower" in providers.APIKEY_PROVIDERS, true); + assert.equal("nous-research" in providers.APIKEY_PROVIDERS, true); + assert.equal("petals" in providers.APIKEY_PROVIDERS, true); assert.equal("poe" in providers.APIKEY_PROVIDERS, true); assert.equal("azure-ai" in providers.APIKEY_PROVIDERS, true); assert.equal("bedrock" in providers.APIKEY_PROVIDERS, true); @@ -460,6 +470,14 @@ test("grok-web taxonomy stays web-cookie only and does not leak into api-key ent apiKeyEntries.some((entry) => entry.providerId === "empower"), true ); + assert.equal( + apiKeyEntries.some((entry) => entry.providerId === "nous-research"), + true + ); + assert.equal( + apiKeyEntries.some((entry) => entry.providerId === "petals"), + true + ); assert.equal( apiKeyEntries.some((entry) => entry.providerId === "poe"), true diff --git a/tests/unit/providers-route-managed-catalog.test.ts b/tests/unit/providers-route-managed-catalog.test.ts index 2bf013c437..654c700225 100644 --- a/tests/unit/providers-route-managed-catalog.test.ts +++ b/tests/unit/providers-route-managed-catalog.test.ts @@ -102,6 +102,21 @@ test("providers route accepts managed local, audio, web-cookie and search provid name: "Empower Primary", }, }, + { + provider: "nous-research", + body: { + provider: "nous-research", + apiKey: "nous-key", + name: "Nous Research Primary", + }, + }, + { + provider: "petals", + body: { + provider: "petals", + name: "Petals Public Endpoint", + }, + }, { provider: "poe", body: { diff --git a/tests/unit/rate-limit-manager.test.ts b/tests/unit/rate-limit-manager.test.ts index e61d523103..be11a34fce 100644 --- a/tests/unit/rate-limit-manager.test.ts +++ b/tests/unit/rate-limit-manager.test.ts @@ -8,6 +8,8 @@ const TEST_DATA_DIR = fs.mkdtempSync(path.join(os.tmpdir(), "omniroute-rate-limi process.env.DATA_DIR = TEST_DATA_DIR; const core = await import("../../src/lib/db/core.ts"); +const providersDb = await import("../../src/lib/db/providers.ts"); +const resilienceSettings = await import("../../src/lib/resilience/settings.ts"); const rateLimitManager = await import("../../open-sse/services/rateLimitManager.ts"); const accountFallback = await import("../../open-sse/services/accountFallback.ts"); @@ -175,3 +177,45 @@ test("rate limit manager parses retry hints from response bodies and locks model ); assert.equal(rateLimitManager.getRateLimitStatus("openai", "conn-body").active, true); }); + +test("rate limit manager recomputes auto-enabled API key connections when queue settings change", async () => { + const autoConnection = await providersDb.createProviderConnection({ + provider: "openai", + authType: "apikey", + name: "Auto OpenAI", + apiKey: "sk-auto", + isActive: true, + }); + const explicitConnection = await providersDb.createProviderConnection({ + provider: "openai", + authType: "apikey", + name: "Explicit OpenAI", + apiKey: "sk-explicit", + isActive: true, + rateLimitProtection: true, + }); + + await rateLimitManager.initializeRateLimits(); + + assert.equal(rateLimitManager.isRateLimitEnabled(autoConnection.id), true); + assert.equal(rateLimitManager.isRateLimitEnabled(explicitConnection.id), true); + assert.ok(rateLimitManager.getAllRateLimitStatus()[`openai:${autoConnection.id}`]); + + await rateLimitManager.applyRequestQueueSettings({ + ...resilienceSettings.DEFAULT_RESILIENCE_SETTINGS.requestQueue, + autoEnableApiKeyProviders: false, + }); + + assert.equal(rateLimitManager.isRateLimitEnabled(autoConnection.id), false); + assert.equal(rateLimitManager.isRateLimitEnabled(explicitConnection.id), true); + assert.equal(rateLimitManager.getAllRateLimitStatus()[`openai:${autoConnection.id}`], undefined); + + await rateLimitManager.applyRequestQueueSettings({ + ...resilienceSettings.DEFAULT_RESILIENCE_SETTINGS.requestQueue, + autoEnableApiKeyProviders: true, + }); + + assert.equal(rateLimitManager.isRateLimitEnabled(autoConnection.id), true); + assert.equal(rateLimitManager.isRateLimitEnabled(explicitConnection.id), true); + assert.ok(rateLimitManager.getAllRateLimitStatus()[`openai:${autoConnection.id}`]); +});