From cc7c50c9594432614b1ecac023c66358999f8d22 Mon Sep 17 00:00:00 2001 From: Hannes Rudolph Date: Fri, 7 Nov 2025 20:05:04 -0700 Subject: [PATCH] fix(openai-native): consolidate loaders + tolerant validation; surface custom models in OpenAI/OpenAI Native UI; fix lint types --- packages/types/src/providers/openai.ts | 37 ++++++++++++++--------- src/api/providers/openai-native.ts | 5 +-- src/core/webview/webviewMessageHandler.ts | 9 +++--- 3 files changed, 30 insertions(+), 21 deletions(-) diff --git a/packages/types/src/providers/openai.ts b/packages/types/src/providers/openai.ts index f30701f62b..bd5430aced 100644 --- a/packages/types/src/providers/openai.ts +++ b/packages/types/src/providers/openai.ts @@ -1,4 +1,4 @@ -import type { ModelInfo } from "../model.js" +import { modelInfoSchema, type ModelInfo } from "../model.js" // https://openai.com/api/pricing/ export type OpenAiNativeModelId = keyof typeof openAiNativeModels @@ -328,12 +328,7 @@ function loadUserOpenAiNativeModels(): Record { try { const parsed = JSON.parse(inlineJson) if (parsed && typeof parsed === "object" && !Array.isArray(parsed)) { - const envResult: Record = {} - for (const [modelId, info] of Object.entries(parsed)) { - if (info && typeof info === "object" && !Array.isArray(info)) { - envResult[modelId] = info as ModelInfo - } - } + const envResult = validateModelInfoRecord(parsed) return envResult } } catch { @@ -363,20 +358,32 @@ function loadUserOpenAiNativeModels(): Record { if (!parsed || typeof parsed !== "object" || Array.isArray(parsed)) return {} - // Best-effort shallow validation; only keep object entries - const result: Record = {} - for (const [modelId, info] of Object.entries(parsed)) { - if (info && typeof info === "object" && !Array.isArray(info)) { - result[modelId] = info as ModelInfo - } - } - return result + // Use shared validator (merges sane defaults and accepts partials) + return validateModelInfoRecord(parsed) } catch { // On any error (missing file, invalid JSON, restricted env), fall back to empty extras return {} } } +/** + * Validate an arbitrary record of model info objects, keeping only entries that conform. + */ +export function validateModelInfoRecord(input: unknown): Record { + if (!input || typeof input !== "object" || Array.isArray(input)) return {} + const out: Record = {} + for (const [id, info] of Object.entries(input as Record)) { + if (info && typeof info === "object" && !Array.isArray(info)) { + const parsedInfo = modelInfoSchema.partial().safeParse(info) + const merged = parsedInfo.success + ? { ...openAiModelInfoSaneDefaults, ...parsedInfo.data } + : { ...openAiModelInfoSaneDefaults, ...(info as Record) } + out[id] = merged as ModelInfo + } + } + return out +} + /** * Returns built-in OpenAI native models merged with user-defined additions. * User models override built-ins on key collision. diff --git a/src/api/providers/openai-native.ts b/src/api/providers/openai-native.ts index 4da50aed64..bab2673443 100644 --- a/src/api/providers/openai-native.ts +++ b/src/api/providers/openai-native.ts @@ -9,6 +9,7 @@ import { openAiNativeDefaultModelId, OpenAiNativeModelId, openAiNativeModels, + validateModelInfoRecord, OPENAI_NATIVE_DEFAULT_TEMPERATURE, GPT5_DEFAULT_TEMPERATURE, type ReasoningEffort, @@ -47,7 +48,7 @@ function loadMergedOpenAiNativeModelsOnHostSync(): Record { if (inline) { const parsed = JSON.parse(inline) if (parsed && typeof parsed === "object" && !Array.isArray(parsed)) { - extras = parsed as Record + extras = validateModelInfoRecord(parsed) } } } catch { @@ -64,7 +65,7 @@ function loadMergedOpenAiNativeModelsOnHostSync(): Record { const raw = fsSync.readFileSync(customPath, "utf8") const parsed = JSON.parse(raw) if (parsed && typeof parsed === "object" && !Array.isArray(parsed)) { - extras = parsed as Record + extras = validateModelInfoRecord(parsed) } } } catch { diff --git a/src/core/webview/webviewMessageHandler.ts b/src/core/webview/webviewMessageHandler.ts index 040e99b4ce..29f9f8a6a9 100644 --- a/src/core/webview/webviewMessageHandler.ts +++ b/src/core/webview/webviewMessageHandler.ts @@ -11,6 +11,7 @@ import { type ClineMessage, type TelemetrySetting, type ModelInfo, + validateModelInfoRecord, TelemetryEventName, UserSettingsConfig, DEFAULT_CHECKPOINT_TIMEOUT_SECONDS, @@ -114,7 +115,7 @@ export const webviewMessageHandler = async ( try { const parsed = JSON.parse(inline) if (parsed && typeof parsed === "object" && !Array.isArray(parsed)) { - extras = parsed + extras = validateModelInfoRecord(parsed) } } catch { // ignore malformed env json @@ -130,7 +131,7 @@ export const webviewMessageHandler = async ( const raw = await fs.readFile(customPath, "utf8") const parsed = JSON.parse(raw) if (parsed && typeof parsed === "object" && !Array.isArray(parsed)) { - extras = parsed + extras = validateModelInfoRecord(parsed) } } catch { // missing/invalid file: ignore @@ -1030,8 +1031,8 @@ export const webviewMessageHandler = async ( case "requestOpenAiNativeModels": { // Return merged built-ins + user-defined from ~/.roo/models/openai-native.json try { - const openAiNativeModels = await getMergedOpenAiNativeModelsOnHost() - provider.postMessageToWebview({ type: "openAiNativeModels", openAiNativeModels }) + const mergedModels = await getMergedOpenAiNativeModelsOnHost() + provider.postMessageToWebview({ type: "openAiNativeModels", openAiNativeModels: mergedModels }) } catch (error) { console.error("Failed to load OpenAI Native models:", error) provider.postMessageToWebview({ type: "openAiNativeModels", openAiNativeModels: {} })