diff --git a/src/app/api/provider-models/route.ts b/src/app/api/provider-models/route.ts index 7c0b3fbd8e..6b286b5b19 100644 --- a/src/app/api/provider-models/route.ts +++ b/src/app/api/provider-models/route.ts @@ -264,18 +264,26 @@ export async function PUT(request) { [ "provider", "modelId", + "modelName", + "source", "normalizeToolCallId", "preserveOpenAIDeveloperRole", "upstreamHeaders", "compatByProtocol", "contextWindowOverride", + "apiFormat", + "targetFormat", + "supportsVision", ].includes(k) ) && ("normalizeToolCallId" in raw || "preserveOpenAIDeveloperRole" in raw || "upstreamHeaders" in raw || "compatByProtocol" in raw || - "contextWindowOverride" in raw); + "contextWindowOverride" in raw || + "apiFormat" in raw || + "targetFormat" in raw || + "supportsVision" in raw); if (compatOnly) { const knownProvider = !!provider && @@ -310,6 +318,18 @@ export async function PUT(request) { ? upstreamHeaders : undefined; } + if ("apiFormat" in raw) { + patch.apiFormat = typeof apiFormat === "string" ? apiFormat : null; + } + if ("targetFormat" in raw) { + patch.targetFormat = typeof targetFormat === "string" ? targetFormat : null; + } + if ("supportsVision" in raw) { + patch.supportsVision = + supportsVision === null || typeof supportsVision === "boolean" + ? supportsVision + : undefined; + } if (Object.keys(patch).length > 0) { mergeModelCompatOverride(provider, modelId, patch); } diff --git a/src/lib/db/models/compat.ts b/src/lib/db/models/compat.ts index aa964fd2ef..640f87b264 100644 --- a/src/lib/db/models/compat.ts +++ b/src/lib/db/models/compat.ts @@ -1,6 +1,7 @@ /** db/models/compat.ts — model-compat overrides (normalizeToolCallId, per-protocol flags, upstream headers). */ import { getDbInstance } from "../core"; +import { resolveProviderAlias } from "@omniroute/open-sse/services/model.ts"; import { MODEL_COMPAT_PROTOCOL_KEYS, type ModelCompatProtocolKey, @@ -114,13 +115,17 @@ export type ModelCompatOverride = { compatByProtocol?: CompatByProtocolMap; upstreamHeaders?: Record; isHidden?: boolean; + apiFormat?: string; + targetFormat?: string; + supportsVision?: boolean; }; export function readCompatList(providerId: string): ModelCompatOverride[] { + const canonicalId = resolveProviderAlias(providerId) || providerId; const db = getDbInstance(); const row = db .prepare("SELECT value FROM key_value WHERE namespace = ? AND key = ?") - .get(MODEL_COMPAT_NAMESPACE, providerId); + .get(MODEL_COMPAT_NAMESPACE, canonicalId); const value = getKeyValue(row).value; if (!value) return []; try { @@ -140,16 +145,17 @@ export function readCompatList(providerId: string): ModelCompatOverride[] { } export function writeCompatList(providerId: string, list: ModelCompatOverride[]) { + const canonicalId = resolveProviderAlias(providerId) || providerId; const db = getDbInstance(); if (list.length === 0) { db.prepare("DELETE FROM key_value WHERE namespace = ? AND key = ?").run( MODEL_COMPAT_NAMESPACE, - providerId + canonicalId ); } else { db.prepare("INSERT OR REPLACE INTO key_value (namespace, key, value) VALUES (?, ?, ?)").run( MODEL_COMPAT_NAMESPACE, - providerId, + canonicalId, JSON.stringify(list) ); } @@ -168,6 +174,9 @@ export type ModelCompatPatch = { /** Replace top-level extra headers for override-only rows; omit to leave unchanged. */ upstreamHeaders?: Record | null; isHidden?: boolean | null; + apiFormat?: string | null; + targetFormat?: string | null; + supportsVision?: boolean | null; }; export function compatByProtocolHasEntries(map: CompatByProtocolMap | undefined): boolean { @@ -230,12 +239,39 @@ export function mergeModelCompatOverride( next.isHidden = Boolean(patch.isHidden); } } + if ("apiFormat" in patch) { + if (!patch.apiFormat) { + delete next.apiFormat; + } else { + next.apiFormat = patch.apiFormat; + } + } + if ("targetFormat" in patch) { + if (!patch.targetFormat) { + delete next.targetFormat; + } else { + next.targetFormat = patch.targetFormat; + } + } + if ("supportsVision" in patch) { + if (patch.supportsVision === null) { + delete next.supportsVision; + } else { + next.supportsVision = Boolean(patch.supportsVision); + } + } const hasHiddenFlag = Object.prototype.hasOwnProperty.call(next, "isHidden"); + const hasApiFormat = Object.prototype.hasOwnProperty.call(next, "apiFormat"); + const hasTargetFormat = Object.prototype.hasOwnProperty.call(next, "targetFormat"); + const hasVisionFlag = Object.prototype.hasOwnProperty.call(next, "supportsVision"); if ( next.normalizeToolCallId || hasPreserveFlag || hasVideoUrlFlag || hasHiddenFlag || + hasApiFormat || + hasTargetFormat || + hasVisionFlag || compatByProtocolHasEntries(next.compatByProtocol) || hasTopUpstream ) { diff --git a/src/sse/services/model.ts b/src/sse/services/model.ts index 7a84606cd1..12179740e3 100644 --- a/src/sse/services/model.ts +++ b/src/sse/services/model.ts @@ -9,6 +9,7 @@ import { } from "@/lib/localDb"; import { getCachedSettings } from "@/lib/localDb"; import { getActiveSyncedCatalog } from "@/lib/db/models/activeSyncedCatalog"; +import { getModelCompatOverrides } from "@/lib/db/models/compat"; import { parseModel, getModelInfoCore, @@ -224,18 +225,36 @@ function findLiveCatalogModelMeta( ); } -function resolveRuntimeFormats(customMatch: any, syncedMatch: any): RuntimeModelMeta { +function resolveRuntimeFormats( + customMatch: any, + syncedMatch: any, + compatOverrideMatch: any +): RuntimeModelMeta { const apiFormat = - customMatch?.apiFormat === "responses" || syncedMatch?.apiFormat === "responses" - ? "responses" - : undefined; + (typeof customMatch?.apiFormat === "string" ? customMatch.apiFormat : undefined) || + (typeof compatOverrideMatch?.apiFormat === "string" ? compatOverrideMatch.apiFormat : undefined) || + (syncedMatch?.apiFormat === "responses" ? "responses" : undefined); const targetFormat = typeof customMatch?.targetFormat === "string" ? customMatch.targetFormat - : typeof syncedMatch?.targetFormat === "string" - ? syncedMatch.targetFormat - : undefined; - return { ...(apiFormat ? { apiFormat } : {}), ...(targetFormat ? { targetFormat } : {}) }; + : typeof compatOverrideMatch?.targetFormat === "string" + ? compatOverrideMatch.targetFormat + : typeof syncedMatch?.targetFormat === "string" + ? syncedMatch.targetFormat + : undefined; + const supportsVision = + typeof customMatch?.supportsVision === "boolean" + ? customMatch.supportsVision + : typeof compatOverrideMatch?.supportsVision === "boolean" + ? compatOverrideMatch.supportsVision + : typeof syncedMatch?.supportsVision === "boolean" + ? syncedMatch.supportsVision + : undefined; + return { + ...(apiFormat ? { apiFormat } : {}), + ...(targetFormat ? { targetFormat } : {}), + ...(supportsVision !== undefined ? { supportsVision } : {}), + }; } function copySyncedThinkingMetadata(metadata: RuntimeModelMeta, syncedMatch: any): void { @@ -269,9 +288,10 @@ function copyRegistryThinkingMetadata(metadata: RuntimeModelMeta, registryMatch: function buildRuntimeModelMeta( customMatch: any, syncedMatch: any, - registryMatch: any + registryMatch: any, + compatOverrideMatch: any ): RuntimeModelMeta { - const metadata = resolveRuntimeFormats(customMatch, syncedMatch); + const metadata = resolveRuntimeFormats(customMatch, syncedMatch, compatOverrideMatch); copyRegistryThinkingMetadata(metadata, registryMatch); copySyncedThinkingMetadata(metadata, syncedMatch); return metadata; @@ -286,9 +306,10 @@ async function lookupModelMeta( available: boolean; }> { try { - const [customModels, liveCatalog] = await Promise.all([ + const [customModels, liveCatalog, compatOverrides] = await Promise.all([ getCustomModels(providerId), getActiveSyncedCatalog(providerId), + Promise.resolve(getModelCompatOverrides(providerId)), ]); const syncedModels = liveCatalog.models; @@ -327,6 +348,9 @@ async function lookupModelMeta( syncedModels ); const registryMatch = findRegistryModel(providerId, resolvedModelId); + const compatOverrideMatch = Array.isArray(compatOverrides) + ? compatOverrides.find((m) => m.id === resolvedModelId || m.id === modelId) + : undefined; const effortBaseModelId = getRegisteredProviderEffortBaseModelId(providerId, modelId); const liveBackedEffortVariant = @@ -335,7 +359,7 @@ async function lookupModelMeta( const available = !liveCatalog.authoritative || Boolean(customMatch || syncedMatch || liveBackedEffortVariant); - const metadata = buildRuntimeModelMeta(customMatch, syncedMatch, registryMatch); + const metadata = buildRuntimeModelMeta(customMatch, syncedMatch, registryMatch, compatOverrideMatch); if (effort) metadata.resolvedThinkingEffort = effort; return { modelId: resolvedModelId, metadata, available }; diff --git a/tests/unit/model-protocol-persistence.test.ts b/tests/unit/model-protocol-persistence.test.ts new file mode 100644 index 0000000000..df34fbaf96 --- /dev/null +++ b/tests/unit/model-protocol-persistence.test.ts @@ -0,0 +1,44 @@ +import test from "node:test"; +import assert from "node:assert/strict"; +import fs from "node:fs"; +import os from "node:os"; +import path from "node:path"; + +const TEST_DATA_DIR = fs.mkdtempSync(path.join(os.tmpdir(), "omniroute-model-protocol-test-")); +process.env.DATA_DIR = TEST_DATA_DIR; +process.env.API_KEY_SECRET = process.env.API_KEY_SECRET || "model-protocol-test-secret"; + +const { mergeModelCompatOverride, getModelCompatOverrides } = + await import("../../src/lib/db/models/compat.ts"); +const { getModelInfo } = await import("../../src/sse/services/model.ts"); + +test("mergeModelCompatOverride persists apiFormat, targetFormat, and supportsVision overrides for built-in models", () => { + mergeModelCompatOverride("opencode", "claude-opus-5", { + apiFormat: "responses", + targetFormat: "claude", + supportsVision: true, + }); + + const overrides = getModelCompatOverrides("opencode"); + const modelOverride = overrides.find((m) => m.id === "claude-opus-5"); + + assert.ok(modelOverride, "Model override row should exist"); + assert.equal(modelOverride.apiFormat, "responses"); + assert.equal(modelOverride.targetFormat, "claude"); + assert.equal(modelOverride.supportsVision, true); +}); + +test("getModelInfo resolves apiFormat, targetFormat, and supportsVision from modelCompatOverrides", async () => { + mergeModelCompatOverride("opencode", "claude-opus-5", { + apiFormat: "responses", + targetFormat: "claude", + supportsVision: true, + }); + + const info = await getModelInfo("opencode/claude-opus-5"); + + assert.equal(info.provider, "opencode-zen"); + assert.equal(info.apiFormat, "responses"); + assert.equal(info.targetFormat, "claude"); + assert.equal(info.supportsVision, true); +});