From 8112ac4c49af4e1405d4ff05aa2cbc4cb6bd2c54 Mon Sep 17 00:00:00 2001 From: awoaCrim Date: Sun, 9 Aug 2026 17:43:56 +0800 Subject: [PATCH 1/2] feat: add structured custom model editor --- frontend/package.json | 1 + frontend/src/actions/provider-actions.js | 212 +++++++++----- frontend/src/actions/provider-actions.test.js | 208 ++++++++++++++ frontend/src/app/main.js | 22 +- frontend/src/components/modals.js | 231 ++++++++++++--- frontend/src/components/provider-drawer.js | 42 ++- frontend/src/config/model-editor.js | 201 ++++++++++++++ frontend/src/config/model-editor.test.js | 161 +++++++++++ frontend/src/config/model-fields.js | 47 ++++ frontend/src/styles/drawer.css | 35 ++- frontend/src/styles/main.css | 1 + frontend/src/styles/modal.css | 16 +- frontend/src/styles/model-editor.css | 262 ++++++++++++++++++ frontend/wailsjs/go/models.ts | 18 ++ internal/configsync/coordinator.go | 71 ++++- internal/configsync/coordinator_test.go | 131 +++++++++ internal/pi/models.go | 66 ++++- internal/pi/models_test.go | 158 +++++++++++ internal/provider/types.go | 98 ++++++- 19 files changed, 1830 insertions(+), 151 deletions(-) create mode 100644 frontend/src/actions/provider-actions.test.js create mode 100644 frontend/src/config/model-editor.js create mode 100644 frontend/src/config/model-editor.test.js create mode 100644 frontend/src/config/model-fields.js create mode 100644 frontend/src/styles/model-editor.css create mode 100644 internal/configsync/coordinator_test.go create mode 100644 internal/pi/models_test.go diff --git a/frontend/package.json b/frontend/package.json index 04f80bb..2ceba89 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -6,6 +6,7 @@ "scripts": { "dev": "vite", "build": "vite build", + "test": "node --test src/config/model-editor.test.js src/actions/provider-actions.test.js", "preview": "vite preview" }, "devDependencies": { diff --git a/frontend/src/actions/provider-actions.js b/frontend/src/actions/provider-actions.js index ef3604b..29c430a 100644 --- a/frontend/src/actions/provider-actions.js +++ b/frontend/src/actions/provider-actions.js @@ -1,5 +1,6 @@ import { createProviderFromPreset } from "../config/presets.js"; import { customHeadersForProvider, headersForMode } from "../config/header-presets.js"; +import { modelExtraFields, readModelDraft } from "../config/model-editor.js"; import { withTimeout } from "../utils/async.js"; const FETCH_MODELS_TIMEOUT_MS = 10_000; @@ -12,6 +13,15 @@ function getModelCheckboxes(root) { return Array.from(root.querySelectorAll("[data-model-id]")); } +function definedFields(value) { + return Object.fromEntries(Object.entries(value ?? {}).filter(([, fieldValue]) => fieldValue !== undefined)); +} + +function modelContextInput(root, modelId) { + return Array.from(root.querySelectorAll("[data-cw-model]")) + .find((input) => input.dataset.cwModel === modelId); +} + function syncToggleAllButton(root) { const button = root.querySelector("[data-toggle-model-selection-all]"); if (!button) return; @@ -45,6 +55,18 @@ function mergeModels(existing, incoming) { return merged; } +function syncDefaultModelState(state, providerId, models, { oldId = "", newId = "" } = {}) { + if (state.defaultProviderId !== providerId) return state; + + let defaultModelId = state.defaultModelId; + if (oldId && defaultModelId === oldId) defaultModelId = newId; + if (defaultModelId && !models.some((model) => model.id === defaultModelId)) { + defaultModelId = ""; + } + + return { ...state, defaultModelId }; +} + export function createProviderActions({ root, api, store, providerForm, feedback }) { let headerSaveTimer = null; @@ -163,9 +185,8 @@ export function createProviderActions({ root, api, store, providerForm, feedback ...state, providers: rest, selectedProviderId: fallback?.id ?? "", - defaultProviderId: state.defaultProviderId === provider.id ? fallback?.id ?? "" : state.defaultProviderId, - defaultModelId: - state.defaultProviderId === provider.id ? fallback?.selectedModelId ?? "" : state.defaultModelId, + defaultProviderId: state.defaultProviderId === provider.id ? "" : state.defaultProviderId, + defaultModelId: state.defaultProviderId === provider.id ? "" : state.defaultModelId, drawer: null, modal: null }; @@ -219,94 +240,145 @@ export function createProviderActions({ root, api, store, providerForm, feedback const modal = store.getState().modal; if (!modal || modal.kind !== "fetch-models") return; + const provider = store.getState().providers.find((item) => item.id === modal.payload.providerId); + if (!provider) return; const selected = Array.from(root.querySelectorAll("[data-model-id]:checked")) .map((checkbox) => { const model = modal.payload.models.find((m) => m.id === checkbox.dataset.modelId); if (!model) return null; - const cwInput = root.querySelector(`[data-cw-model="${model.id}"]`); + const existingModel = provider.models?.find((item) => item.id === model.id) ?? {}; + const cwInput = modelContextInput(root, model.id); const cwK = parseInt(cwInput?.value, 10) || 256; - return { ...model, contextWindow: cwK * 1000 }; + const { selected: _selected, extraFields: _extraFields, ...modelDraft } = model; + const extraFields = { ...modelExtraFields(existingModel), ...modelExtraFields(model) }; + return { + ...existingModel, + ...definedFields(modelDraft), + ...(Object.keys(extraFields).length ? { extraFields } : {}), + contextWindow: cwK * 1000 + }; }) .filter(Boolean); try { await api.replaceModels(modal.payload.providerId, selected); - store.setState((state) => ({ - ...state, - providers: state.providers.map((provider) => { - if (provider.id !== modal.payload.providerId) { - return provider; - } - const nextModels = selected; - const selectedModelId = - nextModels.some((model) => model.id === provider.selectedModelId) - ? provider.selectedModelId - : nextModels[0]?.id ?? ""; - return { ...provider, models: nextModels, selectedModelId }; - }), - modal: null - })); + store.setState((state) => { + const nextState = { + ...state, + providers: state.providers.map((provider) => { + if (provider.id !== modal.payload.providerId) return provider; + const selectedModelId = + selected.some((model) => model.id === provider.selectedModelId) + ? provider.selectedModelId + : selected[0]?.id ?? ""; + return { ...provider, models: selected, selectedModelId }; + }), + modal: null + }; + return syncDefaultModelState(nextState, modal.payload.providerId, selected); + }); } catch (error) { feedback.showError("导入模型失败", error); } } - function openManualModel() { - const modal = store.getState().modal; - const providerId = - modal?.payload?.providerId || currentProvider(store.getState())?.id || ""; + function openModelEditor(modelId = "") { + const state = store.getState(); + const provider = currentProvider(state); + const operation = state.modal; + const providerId = operation?.payload?.providerId || provider?.id || ""; if (!providerId) return; - store.setState((state) => ({ - ...state, + + const model = provider?.id === providerId + ? provider.models?.find((item) => item.id === modelId) + : undefined; + store.setState((nextState) => ({ + ...nextState, modal: { - kind: "manual-model", + kind: "model-editor", payload: { providerId, - modelId: "", - contextWindowK: 256 + mode: model ? "edit" : "add", + originalModelId: model?.id || "", + model: model || { + id: "", + name: "", + reasoning: false, + contextWindow: 128000, + maxTokens: 16384 + } } } })); } - async function importManualModel() { + async function saveModelEditor() { const modal = store.getState().modal; - if (!modal || modal.kind !== "manual-model") return; + if (!modal || modal.kind !== "model-editor") return; - const modelId = root.querySelector('input[name="manualModelId"]')?.value?.trim(); - const contextWindowK = parseInt(root.querySelector('input[name="manualContextWindow"]')?.value, 10) || 256; - if (!modelId) { - feedback.showError("导入模型失败", new Error("Model ID 不能为空")); + const readInput = (name) => root.querySelector(`[name="${name}"]`); + let draft; + try { + draft = readModelDraft({ + originalModel: modal.payload.model || {}, + readValue: (name) => readInput(name)?.value ?? "", + readChecked: (name) => !!readInput(name)?.checked, + readCheckedValues: (name) => Array.from(root.querySelectorAll(`input[name="${name}"]:checked`)) + .map((input) => input.value) + }); + } catch (error) { + feedback.showError("保存模型失败", error); return; } - const manualModel = { - id: modelId, - name: modelId, - contextWindow: contextWindowK * 1000, - reasoning: false - }; + const provider = store.getState().providers.find((item) => item.id === modal.payload.providerId); + if (!provider) return; + const oldId = modal.payload.originalModelId; + const duplicate = (provider.models ?? []).find((model) => model.id === draft.id && model.id !== oldId); + if (duplicate) { + feedback.showError("保存模型失败", new Error(`模型 ID 已存在:${draft.id}`)); + return; + } + const nextModels = mergeModels( + (provider.models ?? []).filter((model) => model.id !== oldId && model.id !== draft.id), + [draft] + ); + const stateModels = nextModels.map((model) => { + const { + __piSwitchReplaceDocument: _replaceDocument, + __piSwitchOriginalId: _originalId, + ...stateModel + } = model; + return stateModel; + }); try { - await api.importModels(modal.payload.providerId, [manualModel]); - store.setState((state) => ({ - ...state, - providers: state.providers.map((provider) => { - if (provider.id !== modal.payload.providerId) { - return provider; - } - const mergedModels = mergeModels(provider.models, [manualModel]); - const selectedModelId = - provider.selectedModelId?.trim() || mergedModels[0]?.id || manualModel.id; - return { ...provider, models: mergedModels, selectedModelId }; - }), - modal: null - })); + await api.replaceModels(provider.id, nextModels); + store.setState((state) => { + const nextState = { + ...state, + providers: state.providers.map((item) => { + if (item.id !== provider.id) return item; + const selectedModelId = item.selectedModelId === oldId + ? draft.id + : stateModels.some((model) => model.id === item.selectedModelId) + ? item.selectedModelId + : stateModels[0]?.id ?? ""; + return { ...item, models: stateModels, selectedModelId }; + }), + modal: null + }; + return syncDefaultModelState(nextState, provider.id, stateModels, { oldId, newId: draft.id }); + }); } catch (error) { - feedback.showError("导入模型失败", error); + feedback.showError("保存模型失败", error); } } + function openManualModel() { + openModelEditor(); + } + function toggleAllModelSelections() { const checkboxes = getModelCheckboxes(root); const nextChecked = checkboxes.some((checkbox) => !checkbox.checked); @@ -322,17 +394,20 @@ export function createProviderActions({ root, api, store, providerForm, feedback const nextModels = (provider.models ?? []).filter((m) => m.id !== modelId); try { await api.replaceModels(provider.id, nextModels); - store.setState((state) => ({ - ...state, - providers: state.providers.map((item) => { - if (item.id !== provider.id) return item; - const selectedModelId = - nextModels.some((m) => m.id === item.selectedModelId) - ? item.selectedModelId - : nextModels[0]?.id ?? ""; - return { ...item, models: nextModels, selectedModelId }; - }) - })); + store.setState((state) => { + const nextState = { + ...state, + providers: state.providers.map((item) => { + if (item.id !== provider.id) return item; + const selectedModelId = + nextModels.some((m) => m.id === item.selectedModelId) + ? item.selectedModelId + : nextModels[0]?.id ?? ""; + return { ...item, models: nextModels, selectedModelId }; + }) + }; + return syncDefaultModelState(nextState, provider.id, nextModels); + }); } catch (error) { feedback.showError("移除模型失败", error); } @@ -343,8 +418,9 @@ export function createProviderActions({ root, api, store, providerForm, feedback remove, fetchModels, importModels, + openModelEditor, + saveModelEditor, openManualModel, - importManualModel, toggleAllModelSelections, deleteModel, setHeaderMode, diff --git a/frontend/src/actions/provider-actions.test.js b/frontend/src/actions/provider-actions.test.js new file mode 100644 index 0000000..baba9ce --- /dev/null +++ b/frontend/src/actions/provider-actions.test.js @@ -0,0 +1,208 @@ +import assert from "node:assert/strict"; +import test from "node:test"; + +import { createProviderActions } from "./provider-actions.js"; +import { createStore } from "../state/store.js"; + +function modelFormRoot(overrides = {}) { + const fields = { + modelId: { value: "demo" }, + modelName: { value: "Demo" }, + modelContextWindow: { value: "128000" }, + modelMaxTokens: { value: "16384" }, + modelReasoning: { checked: false }, + modelCostTiers: { value: "" }, + modelSamplingParams: { value: "" }, + modelHeaders: { value: "" }, + modelExtraJson: { value: "" }, + ...overrides.fields + }; + return { + querySelector(selector) { + const match = selector.match(/^\[name="([^"]+)"\]$/); + return match ? fields[match[1]] ?? null : null; + }, + querySelectorAll(selector) { + if (selector === 'input[name="modelInput"]:checked') { + return (overrides.inputTypes ?? ["text"]).map((value) => ({ value })); + } + return []; + } + }; +} + +function actionsFor({ root, state, api = {} }) { + const errors = []; + const store = createStore(state); + const actions = createProviderActions({ + root, + api, + store, + providerForm: { commit: async () => true }, + feedback: { + showError(title, error) { + errors.push({ title, error }); + }, + showLoading() {} + } + }); + return { actions, store, errors }; +} + +test("renames selected and default model without a second default call", async () => { + let replaced; + const root = modelFormRoot({ + fields: { + modelId: { value: "new-id" }, + modelName: { value: "Renamed" } + } + }); + const { actions, store, errors } = actionsFor({ + root, + state: { + selectedProviderId: "p", + defaultProviderId: "p", + defaultModelId: "old-id", + modal: { + kind: "model-editor", + payload: { + providerId: "p", + originalModelId: "old-id", + model: { id: "old-id", name: "Old", reasoning: false } + } + }, + providers: [{ id: "p", selectedModelId: "old-id", models: [{ id: "old-id", name: "Old" }] }] + }, + api: { + async replaceModels(providerId, models) { + assert.equal(providerId, "p"); + replaced = models; + }, + async setDefaultModel() { + assert.fail("saveModelEditor must not issue a second default-model write"); + } + } + }); + + await actions.saveModelEditor(); + + assert.equal(errors.length, 0); + assert.equal(replaced[0].__piSwitchReplaceDocument, true); + assert.equal(replaced[0].__piSwitchOriginalId, "old-id"); + assert.equal(store.getState().defaultModelId, "new-id"); + assert.equal(store.getState().providers[0].selectedModelId, "new-id"); + assert.equal(store.getState().providers[0].models[0].__piSwitchReplaceDocument, undefined); + assert.equal(store.getState().providers[0].models[0].__piSwitchOriginalId, undefined); +}); + +test("rejects a duplicate model id before writing", async () => { + let writes = 0; + const { actions, errors } = actionsFor({ + root: modelFormRoot({ fields: { modelId: { value: "taken" } } }), + state: { + selectedProviderId: "p", + defaultProviderId: "", + defaultModelId: "", + modal: { + kind: "model-editor", + payload: { + providerId: "p", + originalModelId: "old-id", + model: { id: "old-id", name: "Old" } + } + }, + providers: [{ + id: "p", + selectedModelId: "old-id", + models: [{ id: "old-id", name: "Old" }, { id: "taken", name: "Taken" }] + }] + }, + api: { async replaceModels() { writes += 1; } } + }); + + await actions.saveModelEditor(); + + assert.equal(writes, 0); + assert.equal(errors[0].error.message, "模型 ID 已存在:taken"); +}); + +test("deleting a default model clears the frontend default state", async () => { + const { actions, store } = actionsFor({ + root: modelFormRoot(), + state: { + selectedProviderId: "p", + defaultProviderId: "p", + defaultModelId: "remove", + modal: null, + providers: [{ + id: "p", + selectedModelId: "remove", + models: [{ id: "remove", name: "Remove" }, { id: "keep", name: "Keep" }] + }] + }, + api: { async replaceModels() {} } + }); + + await actions.deleteModel("remove"); + + assert.equal(store.getState().defaultModelId, ""); + assert.equal(store.getState().providers[0].selectedModelId, "keep"); +}); + +test("bulk import keeps existing and fetched unknown fields in the transport envelope", async () => { + let replaced; + const checkboxes = [{ dataset: { modelId: "demo" } }]; + const root = { + querySelector() { + return null; + }, + querySelectorAll(selector) { + if (selector === "[data-model-id]:checked") return checkboxes; + if (selector === "[data-cw-model]") return [{ dataset: { cwModel: "demo" }, value: "256" }]; + return []; + } + }; + const { actions, errors } = actionsFor({ + root, + state: { + selectedProviderId: "p", + defaultProviderId: "", + defaultModelId: "", + modal: { + kind: "fetch-models", + payload: { + providerId: "p", + models: [{ + id: "demo", + name: "Fetched", + api: undefined, + baseUrl: undefined, + selected: true, + extraFields: { fetchedFlag: 2 } + }] + } + }, + providers: [{ + id: "p", + selectedModelId: "demo", + models: [{ + id: "demo", + name: "Existing", + api: "openai-completions", + baseUrl: "https://model.example/v1", + extraFields: { vendorFlag: 1 } + }] + }] + }, + api: { async replaceModels(_providerId, models) { replaced = models; } } + }); + + await actions.importModels(); + + assert.equal(errors.length, 0); + assert.deepEqual(replaced[0].extraFields, { vendorFlag: 1, fetchedFlag: 2 }); + assert.equal(replaced[0].api, "openai-completions"); + assert.equal(replaced[0].baseUrl, "https://model.example/v1"); + assert.equal(replaced[0].selected, undefined); + assert.equal(replaced[0].contextWindow, 256000); +}); diff --git a/frontend/src/app/main.js b/frontend/src/app/main.js index aa07232..d58613d 100644 --- a/frontend/src/app/main.js +++ b/frontend/src/app/main.js @@ -151,7 +151,7 @@ const clickActions = { "data-import-models": providerActions.importModels, "data-toggle-model-selection-all": providerActions.toggleAllModelSelections, "data-open-manual-model": providerActions.openManualModel, - "data-import-manual-model": providerActions.importManualModel, + "data-save-model": providerActions.saveModelEditor, "data-delete-provider": confirmProviderDeletion, "data-set-default": appActions.setDefault, "data-launch-pi": appActions.directLaunch, @@ -181,6 +181,8 @@ function bindClickEvents() { if (target.dataset.providerId) return selectProvider(target.dataset.providerId); if (target.dataset.selectModel) return selectModel(target.dataset.selectModel); + if (target.dataset.editModel) return providerActions.openModelEditor(target.dataset.editModel); + if (target.hasAttribute("data-add-model")) return providerActions.openModelEditor(); if (target.dataset.deleteModel) return providerActions.deleteModel(target.dataset.deleteModel); if (target.dataset.setHeaderMode) return providerActions.setHeaderMode(target.dataset.setHeaderMode); if (target.dataset.openProviderSettings) return openProvider(target.dataset.openProviderSettings); @@ -188,13 +190,29 @@ function bindClickEvents() { if (target.dataset.confirmDeleteProvider) return providerActions.remove(target.dataset.confirmDeleteProvider); const attribute = findActionAttribute(target); - if (attribute) await clickActions[attribute](); + if (attribute) { + try { + await clickActions[attribute](); + } catch (error) { + feedback.showError("操作失败", error); + } + } }); } function bindFormEvents() { root.addEventListener("change", async (event) => { const target = event.target; + if (target.matches("[data-model-thinking-mode]")) { + const level = target.name.replace("modelThinkingMode_", ""); + const valueInput = root.querySelector(`[name="modelThinkingValue_${level}"]`); + if (valueInput) { + valueInput.disabled = target.value !== "custom"; + if (target.value !== "custom") valueInput.value = ""; + } + return; + } + if (target.matches("[data-model-id]")) { providerActions.syncToggleAllButton(); return; diff --git a/frontend/src/components/modals.js b/frontend/src/components/modals.js index e3cfe38..926acc9 100644 --- a/frontend/src/components/modals.js +++ b/frontend/src/components/modals.js @@ -1,9 +1,11 @@ import { escapeHtml } from "./view-utils.js"; +import { modelExtraFields } from "../config/model-editor.js"; +import { COMPAT_BOOLEAN_FIELDS, COMPAT_ENUM_FIELDS, THINKING_LEVELS } from "../config/model-fields.js"; -function modalFrame({ tone = "", eyebrow = "", title, description = "", body = "", actions = "", wide = false }) { +function modalFrame({ tone = "", eyebrow = "", title, description = "", body = "", actions = "", wide = false, className = "" }) { return ` - + `, diff --git a/frontend/src/config/model-editor.js b/frontend/src/config/model-editor.js index d4a11d2..dd81ff1 100644 --- a/frontend/src/config/model-editor.js +++ b/frontend/src/config/model-editor.js @@ -1,4 +1,4 @@ -import { COMPAT_BOOLEAN_FIELDS, COMPAT_ENUM_FIELDS, MODEL_EDITOR_FIELDS, THINKING_LEVELS } from "./model-fields.js"; +import { COMPAT_BOOLEAN_FIELDS, COMPAT_ENUM_FIELDS, THINKING_LEVELS } from "./model-fields.js"; const COST_FIELDS = ["input", "output", "cacheRead", "cacheWrite"]; const COST_TIER_FIELDS = ["inputTokensAbove", ...COST_FIELDS]; @@ -34,12 +34,15 @@ function parseObjectJson(rawValue, label) { return value; } -function readNumber(rawValue, label, { integer = false } = {}) { +function readNumber(rawValue, label, { integer = false, positive = false } = {}) { const value = String(rawValue ?? "").trim(); if (!value) return undefined; const number = Number(value); - if (!Number.isFinite(number) || number < 0 || (integer && !Number.isInteger(number))) { - throw new Error(`${label} 必须是非负${integer ? "整数" : "数字"}`); + const invalid = !Number.isFinite(number) + || (positive ? number <= 0 : number < 0) + || (integer && !Number.isInteger(number)); + if (invalid) { + throw new Error(`${label} 必须是${positive ? "正" : "非负"}${integer ? "整数" : "数字"}`); } return number; } @@ -106,14 +109,7 @@ function readCost(originalModel, value) { } function readCompat(originalModel, value) { - const originalCompat = objectEntries(originalModel.compat); - const compat = Object.fromEntries( - originalCompat.filter(([field]) => - !COMPAT_BOOLEAN_FIELDS.some(([knownField]) => knownField === field) - && !COMPAT_ENUM_FIELDS.some(([knownField]) => knownField === field) - && !COMPAT_JSON_FIELDS.includes(field) - ) - ); + const compat = {}; for (const [name] of COMPAT_BOOLEAN_FIELDS) { const fieldValue = value(`modelCompat_${name}`); @@ -125,7 +121,7 @@ function readCompat(originalModel, value) { compat[name] = fieldValue; } if (fieldValue === "__custom__") { - const originalValue = originalCompat.find(([field]) => field === name)?.[1]; + const originalValue = originalModel.compat?.[name]; if (originalValue !== undefined) compat[name] = originalValue; } } @@ -151,24 +147,23 @@ export function readModelDraft({ const id = value("modelId"); if (!id) throw new Error("模型 id 不能为空"); - const contextWindow = readNumber(value("modelContextWindow"), "上下文窗口", { integer: true }); - const maxTokens = readNumber(value("modelMaxTokens"), "最大输出", { integer: true }); + const contextWindow = readNumber(value("modelContextWindow"), "上下文窗口", { integer: true, positive: true }); + const maxTokens = readNumber(value("modelMaxTokens"), "最大输出", { integer: true, positive: true }); const thinkingLevelMap = readThinkingLevelMap(originalModel, value); const cost = readCost(originalModel, value); const compat = readCompat(originalModel, value); const samplingParams = parseObjectJson(value("modelSamplingParams"), "samplingParams"); const headers = validateHeaders(parseObjectJson(value("modelHeaders"), "headers"), "headers"); - const extra = parseJson(value("modelExtraJson"), "其他字段", {}); + const compatExtraFieldsJson = value("modelCompatExtraJson"); + parseObjectJson(compatExtraFieldsJson, "其他 compat 字段"); + const extraFieldsJson = value("modelExtraJson"); + const extra = parseJson(extraFieldsJson, "其他字段", {}); if (!extra || Array.isArray(extra) || typeof extra !== "object") { throw new Error("其他字段必须是 JSON 对象"); } const inputs = readCheckedValues("modelInput"); - const draft = { ...extra }; - delete draft.selected; - delete draft.extraFields; - delete draft.__piSwitchReplaceDocument; - delete draft.__piSwitchOriginalId; + const draft = {}; Object.assign(draft, { id, ...(Object.prototype.hasOwnProperty.call(originalModel, "api") ? { api: originalModel.api } : {}), @@ -183,19 +178,12 @@ export function readModelDraft({ ...(samplingParams === undefined ? {} : { samplingParams }), ...(headers === undefined ? {} : { headers }), ...(compat === undefined ? {} : { compat }), - __piSwitchReplaceDocument: true, - ...(originalModel.id && originalModel.id !== id ? { __piSwitchOriginalId: originalModel.id } : {}) + ...(compatExtraFieldsJson ? { compatExtraFieldsJson } : {}), + ...(extraFieldsJson ? { extraFieldsJson } : {}), + ...(originalModel.revision ? { revision: originalModel.revision } : {}), + replaceDocument: true, + ...(originalModel.id && originalModel.id !== id ? { originalId: originalModel.id } : {}) }); return draft; } - -export function modelExtraFields(model) { - const transportedExtra = model?.extraFields && typeof model.extraFields === "object" && !Array.isArray(model.extraFields) - ? model.extraFields - : {}; - return { - ...transportedExtra, - ...Object.fromEntries(objectEntries(model).filter(([field]) => !MODEL_EDITOR_FIELDS.has(field))) - }; -} diff --git a/frontend/src/config/model-editor.test.js b/frontend/src/config/model-editor.test.js index af3e0c5..d97296f 100644 --- a/frontend/src/config/model-editor.test.js +++ b/frontend/src/config/model-editor.test.js @@ -1,7 +1,7 @@ import assert from "node:assert/strict"; import test from "node:test"; -import { modelExtraFields, readModelDraft } from "./model-editor.js"; +import { readModelDraft } from "./model-editor.js"; const levels = ["off", "minimal", "low", "medium", "high", "xhigh", "max"]; const compatEnums = ["maxTokensField", "thinkingFormat", "cacheControlFormat", "deferredToolsMode", "sessionAffinityFormat"]; @@ -15,6 +15,7 @@ function form(overrides = {}) { modelCostTiers: "", modelSamplingParams: "", modelHeaders: "", + modelCompatExtraJson: "", modelExtraJson: "", ...Object.fromEntries(levels.map((level) => [`modelThinkingMode_${level}`, "inherit"])), ...Object.fromEntries(compatEnums.map((field) => [`modelCompat_${field}`, "inherit"])), @@ -63,7 +64,6 @@ test("builds a complete model draft and keeps hidden API fields", () => { })); assert.deepEqual(draft, { - vendorOption: { fast: true }, id: "renamed", api: "openai-completions", baseUrl: "https://model.example/v1", @@ -76,8 +76,9 @@ test("builds a complete model draft and keeps hidden API fields", () => { cost: { input: 1.25 }, samplingParams: { temperature: 0.7 }, headers: { "x-model-route": "fast" }, - __piSwitchReplaceDocument: true, - __piSwitchOriginalId: "old-id" + extraFieldsJson: '{"vendorOption":{"fast":true}}', + replaceDocument: true, + originalId: "old-id" }); }); @@ -86,17 +87,53 @@ test("preserves unknown nested fields and unknown compat enum values", () => { originalModel: { thinkingLevelMap: { future: "ultra" }, cost: { input: 1, vendorRate: 2 }, - compat: { thinkingFormat: "future-format", vendorCompat: true } + compat: { thinkingFormat: "future-format" }, + compatExtraFieldsJson: '{"vendorCompat":true}' }, values: { modelCost_input: "1", - modelCompat_thinkingFormat: "__custom__" + modelCompat_thinkingFormat: "__custom__", + modelCompatExtraJson: '{"vendorCompat":true}' } })); assert.deepEqual(draft.thinkingLevelMap, { future: "ultra" }); assert.deepEqual(draft.cost, { vendorRate: 2, input: 1 }); - assert.deepEqual(draft.compat, { vendorCompat: true, thinkingFormat: "future-format" }); + assert.deepEqual(draft.compat, { thinkingFormat: "future-format" }); + assert.equal(draft.compatExtraFieldsJson, '{"vendorCompat":true}'); +}); + +test("preserves editable future compat fields alongside structured fields", () => { + const draft = readModelDraft(form({ + originalModel: { + compat: { + supportsDeveloperRole: false, + supportsThinkingTokenBudget: true + }, + compatExtraFieldsJson: '{"vendorCompat":{"mode":"future"}}' + }, + values: { + modelCompat_supportsDeveloperRole: "false", + modelCompat_supportsThinkingTokenBudget: "true", + modelCompatExtraJson: '{"vendorCompat":{"mode":"updated"}}' + } + })); + + assert.deepEqual(draft.compat, { + supportsDeveloperRole: false, + supportsThinkingTokenBudget: true + }); + assert.equal(draft.compatExtraFieldsJson, '{"vendorCompat":{"mode":"updated"}}'); +}); + +test("clearing future compat JSON removes those fields", () => { + const draft = readModelDraft(form({ + originalModel: { compatExtraFieldsJson: '{"vendorCompat":true}' }, + values: { modelCompatExtraJson: "" } + })); + + assert.equal(draft.compat, undefined); + assert.equal(draft.compatExtraFieldsJson, undefined); }); test("does not expand a partial cost object to four zero fields", () => { @@ -125,7 +162,15 @@ test("validates JSON object fields and header values", () => { test("validates integers, JSON syntax, and complete cost tiers", () => { assert.equal( errorMessage(() => readModelDraft(form({ values: { modelContextWindow: "12.5" } }))), - "上下文窗口 必须是非负整数" + "上下文窗口 必须是正整数" + ); + assert.equal( + errorMessage(() => readModelDraft(form({ values: { modelContextWindow: "0" } }))), + "上下文窗口 必须是正整数" + ); + assert.equal( + errorMessage(() => readModelDraft(form({ values: { modelMaxTokens: "0" } }))), + "最大输出 必须是正整数" ); assert.equal( errorMessage(() => readModelDraft(form({ values: { modelExtraJson: "{" } }))), @@ -146,16 +191,12 @@ test("requires a provider value for custom thinking levels", () => { ); }); -test("filters known and UI-only fields from the extra JSON editor", () => { - assert.deepEqual(modelExtraFields({ - id: "demo", - name: "Demo", - compat: {}, - selected: true, - extraFields: { legacy: true }, - vendorOption: 1 - }), { - legacy: true, - vendorOption: 1 - }); +test("keeps raw large integers as text in unknown field envelopes", () => { + const raw = '{"vendorId":9007199254740993}'; + const draft = readModelDraft(form({ + originalModel: { extraFieldsJson: raw }, + values: { modelExtraJson: raw } + })); + + assert.equal(draft.extraFieldsJson, raw); }); diff --git a/frontend/src/config/model-fields.js b/frontend/src/config/model-fields.js index b1ffff3..3413a4d 100644 --- a/frontend/src/config/model-fields.js +++ b/frontend/src/config/model-fields.js @@ -1,9 +1,3 @@ -export const MODEL_EDITOR_FIELDS = new Set([ - "id", "name", "api", "baseUrl", "reasoning", "thinkingLevelMap", "input", - "cost", "contextWindow", "maxTokens", "samplingParams", "headers", "compat", - "__piSwitchReplaceDocument", "__piSwitchOriginalId", "extraFields", "selected" -]); - export const THINKING_LEVELS = [ ["off", "关闭"], ["minimal", "Minimal"], @@ -35,7 +29,11 @@ export const COMPAT_BOOLEAN_FIELDS = [ ["allowEmptySignature", "允许空 thinking signature", "仅适用于会返回空思考签名的 Anthropic 兼容代理。"], ["supportsStrictTools", "Anthropic 严格工具", "Anthropic 接口是否接受严格 JSON Schema 工具。"], ["supportsToolReferences", "支持工具引用", "Anthropic 兼容接口是否支持工具引用。"], - ["supportsToolSearch", "支持工具搜索", "Responses API 是否支持工具搜索能力。"] + ["supportsToolSearch", "支持工具搜索", "Responses API 是否支持工具搜索能力。"], + ["zaiToolStream", "支持 Z.ai 工具流", "是否发送 tool_stream 以流式接收 Z.ai 工具调用。"], + ["supportsThinkingTokenBudget", "支持思考 token 预算", "是否支持 vLLM 风格的 thinking_token_budget。"], + ["supportsAdditionalTools", "支持 additional_tools", "Responses API 是否支持按消息挂载 additional_tools。"], + ["supportsExplicitPromptCacheMode", "支持显式提示缓存", "是否支持 prompt_cache_options 的显式缓存模式。"] ]; export const COMPAT_ENUM_FIELDS = [ diff --git a/frontend/src/services/wails-api.js b/frontend/src/services/wails-api.js index 2f8ee29..e54485b 100644 --- a/frontend/src/services/wails-api.js +++ b/frontend/src/services/wails-api.js @@ -51,8 +51,8 @@ export class WailsApi { return ImportModels(providerId, models); } - async replaceModels(providerId, models) { - return ReplaceModels(providerId, models); + async replaceModels(providerId, models, expectedRevision = "") { + return ReplaceModels(providerId, models, expectedRevision); } async setDefaultModel(providerId, modelId) { diff --git a/frontend/wailsjs/go/main/App.d.ts b/frontend/wailsjs/go/main/App.d.ts index 9cb684b..47f5928 100644 --- a/frontend/wailsjs/go/main/App.d.ts +++ b/frontend/wailsjs/go/main/App.d.ts @@ -10,34 +10,34 @@ export function CheckEnvVar(arg1:string):Promise; export function CheckForUpdate():Promise; -export function CreateProvider(arg1:provider.Config):Promise; +export function CreateProvider(arg1:provider.ConfigTransport):Promise; export function DeleteProvider(arg1:string):Promise; export function ExecuteLaunchPi(arg1:string,arg2:string):Promise; -export function FetchModels(arg1:string):Promise>; +export function FetchModels(arg1:string):Promise>; export function GetAppState():Promise; -export function ImportModels(arg1:string,arg2:Array):Promise; +export function ImportModels(arg1:string,arg2:Array):Promise; export function InstallUpdate():Promise; export function LaunchPi(arg1:string,arg2:string):Promise; -export function ListProviders():Promise>; +export function ListProviders():Promise>; export function MarkUpdateChecked():Promise; export function OpenConfigFolder():Promise; -export function ReplaceModels(arg1:string,arg2:Array):Promise; +export function ReplaceModels(arg1:string,arg2:Array,arg3:string):Promise; export function SetDefaultModel(arg1:string,arg2:string):Promise; export function TestConnection(arg1:string):Promise; -export function UpdateProvider(arg1:string,arg2:provider.Config):Promise; +export function UpdateProvider(arg1:string,arg2:provider.ConfigTransport):Promise; export function UpdateSettings(arg1:config.AppSettings):Promise; diff --git a/frontend/wailsjs/go/main/App.js b/frontend/wailsjs/go/main/App.js index 2ce0748..e051b8c 100644 --- a/frontend/wailsjs/go/main/App.js +++ b/frontend/wailsjs/go/main/App.js @@ -54,8 +54,8 @@ export function OpenConfigFolder() { return window['go']['main']['App']['OpenConfigFolder'](); } -export function ReplaceModels(arg1, arg2) { - return window['go']['main']['App']['ReplaceModels'](arg1, arg2); +export function ReplaceModels(arg1, arg2, arg3) { + return window['go']['main']['App']['ReplaceModels'](arg1, arg2, arg3); } export function SetDefaultModel(arg1, arg2) { diff --git a/frontend/wailsjs/go/models.ts b/frontend/wailsjs/go/models.ts index 145385d..e9606ea 100644 --- a/frontend/wailsjs/go/models.ts +++ b/frontend/wailsjs/go/models.ts @@ -30,7 +30,7 @@ export namespace config { } export class AppState { version: string; - providers: provider.Config[]; + providers: provider.ConfigTransport[]; selectedProviderId: string; defaultProviderId: string; defaultModelId: string; @@ -44,7 +44,7 @@ export namespace config { constructor(source: any = {}) { if ('string' === typeof source) source = JSON.parse(source); this.version = source["version"]; - this.providers = this.convertValues(source["providers"], provider.Config); + this.providers = this.convertValues(source["providers"], provider.ConfigTransport); this.selectedProviderId = source["selectedProviderId"]; this.defaultProviderId = source["defaultProviderId"]; this.defaultModelId = source["defaultModelId"]; @@ -94,7 +94,7 @@ export namespace pi { export namespace provider { - export class ModelInfo { + export class ModelTransport { id: string; name: string; api?: string; @@ -108,10 +108,14 @@ export namespace provider { samplingParams?: Record; headers?: Record; compat?: Record; - extraFields?: Record; + compatExtraFieldsJson?: string; + extraFieldsJson?: string; + revision?: string; + replaceDocument?: boolean; + originalId?: string; static createFrom(source: any = {}) { - return new ModelInfo(source); + return new ModelTransport(source); } constructor(source: any = {}) { @@ -129,10 +133,14 @@ export namespace provider { this.samplingParams = source["samplingParams"]; this.headers = source["headers"]; this.compat = source["compat"]; - this.extraFields = source["extraFields"]; + this.compatExtraFieldsJson = source["compatExtraFieldsJson"]; + this.extraFieldsJson = source["extraFieldsJson"]; + this.revision = source["revision"]; + this.replaceDocument = source["replaceDocument"]; + this.originalId = source["originalId"]; } } - export class Config { + export class ConfigTransport { id: string; name: string; type: string; @@ -144,12 +152,14 @@ export namespace provider { headerMode: string; headers: Record; customHeaders?: Record; - models: ModelInfo[]; + models: ModelTransport[]; host: string; selectedModelId: string; + extraFieldsJson?: string; + modelsRevision?: string; static createFrom(source: any = {}) { - return new Config(source); + return new ConfigTransport(source); } constructor(source: any = {}) { @@ -165,9 +175,11 @@ export namespace provider { this.headerMode = source["headerMode"]; this.headers = source["headers"]; this.customHeaders = source["customHeaders"]; - this.models = this.convertValues(source["models"], ModelInfo); + this.models = this.convertValues(source["models"], ModelTransport); this.host = source["host"]; this.selectedModelId = source["selectedModelId"]; + this.extraFieldsJson = source["extraFieldsJson"]; + this.modelsRevision = source["modelsRevision"]; } convertValues(a: any, classs: any, asMap: boolean = false): any { @@ -204,6 +216,38 @@ export namespace provider { this.lines = source["lines"]; } } + export class ModelListTransport { + models: ModelTransport[]; + revision: string; + + static createFrom(source: any = {}) { + return new ModelListTransport(source); + } + + constructor(source: any = {}) { + if ('string' === typeof source) source = JSON.parse(source); + this.models = this.convertValues(source["models"], ModelTransport); + this.revision = source["revision"]; + } + + convertValues(a: any, classs: any, asMap: boolean = false): any { + if (!a) { + return a; + } + if (a.slice && a.map) { + return (a as any[]).map(elem => this.convertValues(elem, classs)); + } else if ("object" === typeof a) { + if (asMap) { + for (const key of Object.keys(a)) { + a[key] = new classs(a[key]); + } + return a; + } + return new classs(a); + } + return a; + } + } } diff --git a/internal/config/switch_config.go b/internal/config/switch_config.go index b5d4c57..d86a470 100644 --- a/internal/config/switch_config.go +++ b/internal/config/switch_config.go @@ -31,13 +31,13 @@ type SwitchConfig struct { } type AppState struct { - Version string `json:"version"` - Providers []provider.Config `json:"providers"` - SelectedProviderID string `json:"selectedProviderId"` - DefaultProviderID string `json:"defaultProviderId"` - DefaultModelID string `json:"defaultModelId"` - Settings AppSettings `json:"settings"` - Logs []string `json:"logs"` + Version string `json:"version"` + Providers []provider.ConfigTransport `json:"providers"` + SelectedProviderID string `json:"selectedProviderId"` + DefaultProviderID string `json:"defaultProviderId"` + DefaultModelID string `json:"defaultModelId"` + Settings AppSettings `json:"settings"` + Logs []string `json:"logs"` } type Service struct { diff --git a/internal/configsync/coordinator.go b/internal/configsync/coordinator.go index 5a22efe..6913507 100644 --- a/internal/configsync/coordinator.go +++ b/internal/configsync/coordinator.go @@ -1,7 +1,6 @@ package configsync import ( - "encoding/json" "errors" "sync" "time" @@ -139,7 +138,7 @@ func (c *Coordinator) MergeModels(providerID string, models []provider.ModelInfo return c.saveAppConfig(cfg) } -func (c *Coordinator) ReplaceModels(providerID string, models []provider.ModelInfo) error { +func (c *Coordinator) ReplaceModels(providerID string, models []provider.ModelInfo, expectedRevisions ...string) ([]provider.ModelInfo, error) { c.mu.Lock() defer c.mu.Unlock() @@ -147,16 +146,16 @@ func (c *Coordinator) ReplaceModels(providerID string, models []provider.ModelIn cfg, current, err := c.providerState(providerID) if err != nil { - return err + return nil, err } defaults, err := pi.ReadDefaults(cfg.Settings.PiSettingsPath) if err != nil { - return err + return nil, err } - replaced, err := pi.ReplaceModels(cfg.Settings.PiModelsPath, providerID, models) + replaced, err := pi.ReplaceModels(cfg.Settings.PiModelsPath, providerID, models, expectedRevisions...) if err != nil { - return err + return nil, err } c.wrote(cfg.Settings.PiModelsPath) current.Models = replaced @@ -178,7 +177,7 @@ func (c *Coordinator) ReplaceModels(providerID string, models []provider.ModelIn if renamedID := renames[defaultModelID]; renamedID != "" { defaultModelID = renamedID if err := pi.PatchDefaults(cfg.Settings.PiSettingsPath, pi.DefaultPatch{DefaultModel: stringPointer(defaultModelID)}); err != nil { - return err + return nil, err } c.wrote(cfg.Settings.PiSettingsPath) } @@ -187,7 +186,7 @@ func (c *Coordinator) ReplaceModels(providerID string, models []provider.ModelIn if defaultRemoved { empty := "" if err := pi.PatchDefaults(cfg.Settings.PiSettingsPath, pi.DefaultPatch{DefaultModel: &empty}); err != nil { - return err + return nil, err } c.wrote(cfg.Settings.PiSettingsPath) } @@ -199,7 +198,10 @@ func (c *Coordinator) ReplaceModels(providerID string, models []provider.ModelIn cfg.Settings.LastDefaultModelID = "" } } - return c.saveAppConfig(cfg) + if err := c.saveAppConfig(cfg); err != nil { + return nil, err + } + return replaced, nil } func (c *Coordinator) SetDefault(providerID string, modelID string) error { @@ -231,24 +233,10 @@ func (c *Coordinator) SetDefault(providerID string, modelID string) error { func modelRenames(models []provider.ModelInfo) map[string]string { renames := map[string]string{} for index := range models { - value, exists := models[index].ExtraFields[provider.ModelOriginalIDField] - if !exists { - if rawValue, rawExists := models[index].Extra[provider.ModelOriginalIDField]; rawExists { - var oldID string - if json.Unmarshal(rawValue, &oldID) == nil { - value = oldID - exists = true - } - } - } - if exists { - oldID, ok := value.(string) - if ok && oldID != "" && oldID != models[index].ID { - renames[oldID] = models[index].ID - } + oldID := models[index].OriginalID + if oldID != "" && oldID != models[index].ID { + renames[oldID] = models[index].ID } - delete(models[index].ExtraFields, provider.ModelOriginalIDField) - delete(models[index].Extra, provider.ModelOriginalIDField) } return renames } diff --git a/internal/configsync/coordinator_test.go b/internal/configsync/coordinator_test.go index 11e42a8..00a81ab 100644 --- a/internal/configsync/coordinator_test.go +++ b/internal/configsync/coordinator_test.go @@ -1,7 +1,6 @@ package configsync import ( - "encoding/json" "os" "path/filepath" "strings" @@ -46,18 +45,25 @@ func TestReplaceModelsRenamesSelectedAndDefaultModel(t *testing.T) { } initialConfig.Settings.LastDefaultProviderID = "p" initialConfig.Settings.LastDefaultModelID = "old-id" + oldModel, err := initialConfig.ProviderByID("p") + if err != nil { + t.Fatal(err) + } + oldRevision := oldModel.Models[0].Revision + if oldRevision == "" { + t.Fatal("loaded model has no revision") + } if err := service.Save(initialConfig); err != nil { t.Fatal(err) } coordinator := New(service, nil) - if err := coordinator.ReplaceModels("p", []provider.ModelInfo{{ - ID: "new-id", - Name: "New", - ExtraFields: map[string]any{ - provider.ModelReplaceDocumentField: true, - provider.ModelOriginalIDField: "old-id", - }, + if _, err := coordinator.ReplaceModels("p", []provider.ModelInfo{{ + ID: "new-id", + Name: "New", + ReplaceDocument: true, + OriginalID: "old-id", + Revision: oldRevision, }}); err != nil { t.Fatal(err) } @@ -93,7 +99,7 @@ func TestReplaceModelsRenamesSelectedAndDefaultModel(t *testing.T) { t.Fatal("models document is empty") } encoded := string(data) - if containsAny(encoded, provider.ModelReplaceDocumentField, provider.ModelOriginalIDField, "extraFields") { + if containsAny(encoded, "replaceDocument", "originalId") { t.Fatalf("internal transport metadata leaked into models.json: %s", encoded) } } @@ -107,25 +113,17 @@ func containsAny(value string, needles ...string) bool { return false } -func TestModelRenamesExtractsAndRemovesInternalMetadata(t *testing.T) { +func TestModelRenamesReadsMetadataWithoutMutatingModels(t *testing.T) { models := []provider.ModelInfo{{ - ID: "new-id", - ExtraFields: map[string]any{ - provider.ModelOriginalIDField: "old-id", - }, - Extra: map[string]json.RawMessage{ - provider.ModelOriginalIDField: json.RawMessage(`"old-id"`), - }, + ID: "new-id", + OriginalID: "old-id", }} renames := modelRenames(models) if renames["old-id"] != "new-id" { t.Fatalf("renames = %#v, want old-id -> new-id", renames) } - if _, exists := models[0].ExtraFields[provider.ModelOriginalIDField]; exists { - t.Fatal("original model id metadata remained in ExtraFields") - } - if _, exists := models[0].Extra[provider.ModelOriginalIDField]; exists { - t.Fatal("original model id metadata remained in Extra") + if models[0].OriginalID != "old-id" { + t.Fatal("rename metadata was mutated before the persistence layer used it") } } diff --git a/internal/pi/models.go b/internal/pi/models.go index 026fa3b..6f49cfb 100644 --- a/internal/pi/models.go +++ b/internal/pi/models.go @@ -1,8 +1,11 @@ package pi import ( + "crypto/sha256" + "encoding/hex" "encoding/json" "errors" + "fmt" "os" "sort" "strings" @@ -119,30 +122,35 @@ func DeleteProvider(path string, providerID string) error { } func MergeModels(path string, providerID string, incoming []provider.ModelInfo) ([]provider.ModelInfo, error) { - var result []provider.ModelInfo - err := mutateProviderModels(path, providerID, func(existing []provider.ModelInfo) []provider.ModelInfo { - result = provider.MergeModels(existing, incoming) - return result - }, true) - return result, err + return mutateProviderModels(path, providerID, func(existing []provider.ModelInfo) ([]provider.ModelInfo, error) { + return provider.MergeModels(existing, incoming), nil + }, modelDocumentMergeIncremental) } -func ReplaceModels(path string, providerID string, models []provider.ModelInfo) ([]provider.ModelInfo, error) { +func ReplaceModels(path string, providerID string, models []provider.ModelInfo, expectedRevisions ...string) ([]provider.ModelInfo, error) { result := provider.NormalizeModels(models) - err := mutateProviderModels(path, providerID, func([]provider.ModelInfo) []provider.ModelInfo { - return result - }, true) - for index := range result { - delete(result[index].Extra, replaceModelDocumentField) - delete(result[index].Extra, provider.ModelOriginalIDField) - delete(result[index].ExtraFields, replaceModelDocumentField) - delete(result[index].ExtraFields, provider.ModelOriginalIDField) + expectedRevision := "" + if len(expectedRevisions) > 0 { + expectedRevision = expectedRevisions[0] } - return result, err + return mutateProviderModels(path, providerID, func(existing []provider.ModelInfo) ([]provider.ModelInfo, error) { + if expectedRevision != "" && provider.ModelListRevision(existing) != expectedRevision { + return nil, errors.New("模型列表已被外部修改,请重新打开编辑器") + } + return result, nil + }, modelDocumentReplaceList) } -func mutateProviderModels(path string, providerID string, mutate func([]provider.ModelInfo) []provider.ModelInfo, preserveExistingFields bool) error { - return mutateJSONDocument(path, func(payload map[string]json.RawMessage) error { +type modelDocumentMergeMode int + +const ( + modelDocumentMergeIncremental modelDocumentMergeMode = iota + modelDocumentReplaceList +) + +func mutateProviderModels(path string, providerID string, mutate func([]provider.ModelInfo) ([]provider.ModelInfo, error), mode modelDocumentMergeMode) ([]provider.ModelInfo, error) { + var persisted []provider.ModelInfo + err := mutateJSONDocument(path, func(payload map[string]json.RawMessage) error { providers, err := decodeProviders(payload) if err != nil { return err @@ -160,18 +168,27 @@ func mutateProviderModels(path string, providerID string, mutate func([]provider if err != nil { return err } - nextModels := provider.NormalizeModels(mutate(existingModels)) - mergedModels, err := mergeModelDocuments(fields["models"], nextModels, preserveExistingFields) + next, err := mutate(existingModels) + if err != nil { + return err + } + nextModels := provider.NormalizeModels(next) + mergedModels, err := mergeModelDocuments(fields["models"], nextModels, mode) if err != nil { return err } fields["models"] = mergedModels + persisted, err = decodeModels(mergedModels) + if err != nil { + return err + } providers[providerID], err = json.Marshal(fields) if err != nil { return err } return encodeProviders(payload, providers) }) + return persisted, err } func decodeProviders(payload map[string]json.RawMessage) (map[string]json.RawMessage, error) { @@ -218,7 +235,7 @@ func mergeProviderDocument(existing json.RawMessage, cfg provider.Config, includ return nil, err } if includeModels { - models, err := mergeModelDocuments(fields["models"], cfg.Models, true) + models, err := mergeModelDocuments(fields["models"], cfg.Models, modelDocumentReplaceList) if err != nil { return nil, err } @@ -232,9 +249,7 @@ var knownModelDocumentFields = map[string]struct{}{ "cost": {}, "contextWindow": {}, "maxTokens": {}, "samplingParams": {}, "headers": {}, "compat": {}, } -const replaceModelDocumentField = provider.ModelReplaceDocumentField - -func mergeModelDocuments(existing json.RawMessage, models []provider.ModelInfo, preserveExistingFields bool) (json.RawMessage, error) { +func mergeModelDocuments(existing json.RawMessage, models []provider.ModelInfo, mode modelDocumentMergeMode) (json.RawMessage, error) { existingByID := map[string]map[string]json.RawMessage{} if len(existing) > 0 { var rawModels []json.RawMessage @@ -264,45 +279,35 @@ func mergeModelDocuments(existing json.RawMessage, models []provider.ModelInfo, if err := json.Unmarshal(encodedModel, &nextFields); err != nil { return nil, err } - if rawExtraFields, exists := nextFields["extraFields"]; exists { - extraFields := map[string]json.RawMessage{} - if err := json.Unmarshal(rawExtraFields, &extraFields); err != nil { - return nil, err + delete(nextFields, "selected") + replaceDocument := model.ReplaceDocument + originalID := strings.TrimSpace(model.OriginalID) + oldFields := existingByID[model.ID] + if oldFields == nil && originalID != "" { + oldFields = existingByID[originalID] + } + if model.Revision != "" { + if oldFields == nil || modelDocumentRevision(oldFields) != model.Revision { + return nil, fmt.Errorf("模型 %s 已被外部修改,请重新打开编辑器", model.ID) } - for key, value := range extraFields { - if key == "selected" { - continue + } + if !replaceDocument && oldFields != nil { + if rawCompat, ok := oldFields["compat"]; ok { + if err := mergeCompatFields(nextFields, rawCompat, mode == modelDocumentMergeIncremental); err != nil { + return nil, err } - if key == replaceModelDocumentField || key == provider.ModelOriginalIDField { - nextFields[key] = value + } + for key, value := range oldFields { + if key == "selected" { continue } - if _, known := knownModelDocumentFields[key]; !known { - nextFields[key] = value - } - } - delete(nextFields, "extraFields") - } - delete(nextFields, provider.ModelOriginalIDField) - replaceDocument := false - if rawReplace, exists := nextFields[replaceModelDocumentField]; exists { - _ = json.Unmarshal(rawReplace, &replaceDocument) - delete(nextFields, replaceModelDocumentField) - } - if preserveExistingFields && !replaceDocument { - if oldFields := existingByID[model.ID]; oldFields != nil { - for key, value := range oldFields { - // selected is a UI-only flag used by the import dialog, - // never a Pi model parameter. - if key == "selected" || key == replaceModelDocumentField || key == provider.ModelOriginalIDField { - continue - } + if mode == modelDocumentReplaceList { if _, known := knownModelDocumentFields[key]; known { continue } - if _, exists := nextFields[key]; !exists { - nextFields[key] = value - } + } + if _, exists := nextFields[key]; !exists { + nextFields[key] = value } } } @@ -315,6 +320,46 @@ func mergeModelDocuments(existing json.RawMessage, models []provider.ModelInfo, return json.Marshal(merged) } +func mergeCompatFields(modelFields map[string]json.RawMessage, existingCompat json.RawMessage, preserveKnown bool) error { + oldFields := map[string]json.RawMessage{} + if err := json.Unmarshal(existingCompat, &oldFields); err != nil { + return err + } + compatFields := map[string]json.RawMessage{} + if rawCompat, ok := modelFields["compat"]; ok { + if err := json.Unmarshal(rawCompat, &compatFields); err != nil { + return err + } + } + for field, value := range oldFields { + if !preserveKnown && provider.IsTransportCompatField(field) { + continue + } + if _, exists := compatFields[field]; !exists { + compatFields[field] = value + } + } + if len(compatFields) == 0 { + delete(modelFields, "compat") + return nil + } + encoded, err := json.Marshal(compatFields) + if err != nil { + return err + } + modelFields["compat"] = encoded + return nil +} + +func modelDocumentRevision(fields map[string]json.RawMessage) string { + encoded, err := json.Marshal(fields) + if err != nil { + return "" + } + hash := sha256.Sum256(encoded) + return hex.EncodeToString(hash[:]) +} + func decodeModels(raw json.RawMessage) ([]provider.ModelInfo, error) { if len(raw) == 0 { return nil, nil diff --git a/internal/pi/models_test.go b/internal/pi/models_test.go index 3a55565..d7ec5e0 100644 --- a/internal/pi/models_test.go +++ b/internal/pi/models_test.go @@ -20,32 +20,284 @@ func TestModelInfoExposesUnknownFieldsThroughTransportEnvelope(t *testing.T) { }`), &model); err != nil { t.Fatal(err) } - encoded, err := json.Marshal(model) + transport, err := provider.NewModelTransport(model) if err != nil { t.Fatal(err) } - var fields map[string]json.RawMessage - if err := json.Unmarshal(encoded, &fields); err != nil { + if transport.ExtraFieldsJSON != `{ + "vendorOption": { + "fast": true + } +}` { + t.Fatalf("extraFieldsJson = %s", transport.ExtraFieldsJSON) + } + if transport.API != "openai-completions" { + t.Fatalf("api = %q", transport.API) + } + if transport.ThinkingLevelMap["high"] != "high" { + t.Fatalf("thinkingLevelMap = %#v", transport.ThinkingLevelMap) + } +} + +func TestMergeModelsPreservesMissingKnownFields(t *testing.T) { + path := filepath.Join(t.TempDir(), "models.json") + initial := []byte(`{ + "providers": { + "demo": { + "models": [{ + "id": "keep", + "name": "Old name", + "contextWindow": 200000, + "maxTokens": 32000, + "compat": {"supportsDeveloperRole": false} + }] + } + } +}`) + if err := os.WriteFile(path, initial, 0o644); err != nil { + t.Fatal(err) + } + + if _, err := MergeModels(path, "demo", []provider.ModelInfo{{ID: "keep", Name: "New name"}}); err != nil { + t.Fatal(err) + } + + model := readModelDocument(t, path, "demo", "keep") + for _, field := range []string{"contextWindow", "maxTokens", "compat"} { + if _, ok := model[field]; !ok { + t.Errorf("incremental merge removed existing known field %q", field) + } + } +} + +func TestMergeModelsPreservesMissingNestedCompatFields(t *testing.T) { + path := filepath.Join(t.TempDir(), "models.json") + initial := []byte(`{ + "providers": { + "demo": { + "models": [{ + "id":"keep", + "name":"Keep", + "compat":{"supportsDeveloperRole":false,"supportsTemperature":true,"vendorCompat":1} + }] + } + } +}`) + if err := os.WriteFile(path, initial, 0o644); err != nil { + t.Fatal(err) + } + + if _, err := MergeModels(path, "demo", []provider.ModelInfo{{ + ID: "keep", + Name: "Updated", + Compat: map[string]any{"supportsDeveloperRole": true}, + }}); err != nil { + t.Fatal(err) + } + + model := readModelDocument(t, path, "demo", "keep") + var compat map[string]json.RawMessage + if err := json.Unmarshal(model["compat"], &compat); err != nil { + t.Fatal(err) + } + for _, field := range []string{"supportsTemperature", "vendorCompat"} { + if _, ok := compat[field]; !ok { + t.Errorf("incremental merge removed compat field %q", field) + } + } + if got := string(compat["supportsDeveloperRole"]); got != "true" { + t.Fatalf("supportsDeveloperRole = %s, want true", got) + } +} +func TestReplaceModelsPreservesUnknownNestedCompatFields(t *testing.T) { + path := filepath.Join(t.TempDir(), "models.json") + initial := []byte(`{ + "providers": { + "demo": { + "models": [{ + "id":"keep", + "name":"Keep", + "compat":{"supportsDeveloperRole":false,"vendorCompat":{"mode":"future"}} + }] + } + } +}`) + if err := os.WriteFile(path, initial, 0o644); err != nil { t.Fatal(err) } - if _, ok := fields["vendorOption"]; ok { - t.Fatal("unknown fields must use the Wails transport envelope") + + if _, err := ReplaceModels(path, "demo", []provider.ModelInfo{{ + ID: "keep", + Name: "Updated", + Compat: map[string]any{"supportsDeveloperRole": true}, + }}); err != nil { + t.Fatal(err) } - var extra map[string]json.RawMessage - if err := json.Unmarshal(fields["extraFields"], &extra); err != nil { + + model := readModelDocument(t, path, "demo", "keep") + var compat map[string]json.RawMessage + if err := json.Unmarshal(model["compat"], &compat); err != nil { t.Fatal(err) } - if _, ok := extra["vendorOption"]; !ok { - t.Fatal("transport envelope lost vendorOption") + if _, ok := compat["vendorCompat"]; !ok { + t.Fatal("replace list removed unknown nested compat field") } - if _, ok := fields["api"]; !ok { - t.Fatal("typed model api field missing") + if got := string(compat["supportsDeveloperRole"]); got != "true" { + t.Fatalf("supportsDeveloperRole = %s, want true", got) } - if _, ok := fields["thinkingLevelMap"]; !ok { - t.Fatal("typed thinkingLevelMap field missing") +} + +func TestReplaceModelsPreservesUnknownTopLevelExtraFieldsName(t *testing.T) { + path := filepath.Join(t.TempDir(), "models.json") + initial := []byte(`{ + "providers": { + "demo": { + "models": [{"id":"keep","name":"Keep","extraFields":{"vendorFlag":true}}] + } + } +}`) + if err := os.WriteFile(path, initial, 0o644); err != nil { + t.Fatal(err) + } + + providers, err := ReadAllModels(path) + if err != nil { + t.Fatal(err) + } + if _, err := ReplaceModels(path, "demo", providers[0].Models); err != nil { + t.Fatal(err) + } + + model := readModelDocument(t, path, "demo", "keep") + if _, ok := model["extraFields"]; !ok { + t.Fatal("unknown top-level extraFields field was not preserved") + } + if _, ok := model["vendorFlag"]; ok { + t.Fatal("nested extraFields value leaked into the model top level") } } +func TestReplaceModelsPreservesUnknownLargeInteger(t *testing.T) { + path := filepath.Join(t.TempDir(), "models.json") + initial := []byte(`{ + "providers": { + "demo": { + "models": [{"id":"keep","name":"Keep","vendorId":9007199254740993}] + } + } +}`) + if err := os.WriteFile(path, initial, 0o644); err != nil { + t.Fatal(err) + } + + providers, err := ReadAllModels(path) + if err != nil { + t.Fatal(err) + } + if _, err := ReplaceModels(path, "demo", providers[0].Models); err != nil { + t.Fatal(err) + } + + model := readModelDocument(t, path, "demo", "keep") + if got := string(model["vendorId"]); got != "9007199254740993" { + t.Fatalf("vendorId = %s, want exact 9007199254740993", got) + } +} + +func TestReplaceModelsRejectsStaleListRevision(t *testing.T) { + path := filepath.Join(t.TempDir(), "models.json") + initial := []byte(`{"providers":{"demo":{"models":[{"id":"keep","name":"Keep"}]}}}`) + if err := os.WriteFile(path, initial, 0o644); err != nil { + t.Fatal(err) + } + providers, err := ReadAllModels(path) + if err != nil { + t.Fatal(err) + } + expectedRevision := provider.ModelListRevision(providers[0].Models) + if expectedRevision == "" { + t.Fatal("model list has no revision") + } + + external := []byte(`{"providers":{"demo":{"models":[{"id":"keep","name":"Keep"},{"id":"external","name":"External"}]}}}`) + if err := os.WriteFile(path, external, 0o644); err != nil { + t.Fatal(err) + } + if _, err := ReplaceModels(path, "demo", providers[0].Models, expectedRevision); err == nil { + t.Fatal("stale list replacement unexpectedly removed an external model") + } + if model := readModelDocument(t, path, "demo", "external"); model == nil { + t.Fatal("external model was removed") + } +} +func TestReplaceModelsRejectsStaleRevision(t *testing.T) { + path := filepath.Join(t.TempDir(), "models.json") + initial := []byte(`{ + "providers": { + "demo": { + "models": [{"id":"keep","name":"Before","contextWindow":128000}] + } + } +}`) + if err := os.WriteFile(path, initial, 0o644); err != nil { + t.Fatal(err) + } + + providers, err := ReadAllModels(path) + if err != nil { + t.Fatal(err) + } + stale := providers[0].Models[0] + stale.Name = "Edited in Pi Switch" + stale.ReplaceDocument = true + + external := []byte(`{ + "providers": { + "demo": { + "models": [{"id":"keep","name":"Externally changed","contextWindow":200000}] + } + } +}`) + if err := os.WriteFile(path, external, 0o644); err != nil { + t.Fatal(err) + } + + if _, err := ReplaceModels(path, "demo", []provider.ModelInfo{stale}); err == nil { + t.Fatal("stale model edit unexpectedly overwrote an external change") + } + model := readModelDocument(t, path, "demo", "keep") + var name string + if err := json.Unmarshal(model["name"], &name); err != nil { + t.Fatal(err) + } + if name != "Externally changed" { + t.Fatalf("name = %q, want external value", name) + } +} + +func TestReplaceModelsReturnsFreshRevision(t *testing.T) { + path := filepath.Join(t.TempDir(), "models.json") + initial := []byte(`{"providers":{"demo":{"models":[{"id":"keep","name":"Before"}]}}}`) + if err := os.WriteFile(path, initial, 0o644); err != nil { + t.Fatal(err) + } + providers, err := ReadAllModels(path) + if err != nil { + t.Fatal(err) + } + model := providers[0].Models[0] + oldRevision := model.Revision + model.Name = "After" + model.ReplaceDocument = true + + replaced, err := ReplaceModels(path, "demo", []provider.ModelInfo{model}) + if err != nil { + t.Fatal(err) + } + if len(replaced) != 1 || replaced[0].Revision == "" || replaced[0].Revision == oldRevision { + t.Fatalf("revisions = old %q, new %#v", oldRevision, replaced) + } +} func TestReplaceModelsUsesSubmittedDocumentForExplicitEdits(t *testing.T) { path := filepath.Join(t.TempDir(), "models.json") initial := []byte(`{ @@ -66,33 +318,19 @@ func TestReplaceModelsUsesSubmittedDocumentForExplicitEdits(t *testing.T) { } _, err := ReplaceModels(path, "demo", []provider.ModelInfo{{ - ID: "keep", - Name: "New name", - Extra: map[string]json.RawMessage{ - "__piSwitchReplaceDocument": json.RawMessage(`true`), - }, + ID: "keep", + Name: "New name", + ReplaceDocument: true, }}) if err != nil { t.Fatal(err) } - data, err := os.ReadFile(path) - if err != nil { - t.Fatal(err) - } - var document struct { - Providers map[string]struct { - Models []map[string]json.RawMessage `json:"models"` - } `json:"providers"` - } - if err := json.Unmarshal(data, &document); err != nil { - t.Fatal(err) - } - model := document.Providers["demo"].Models[0] + model := readModelDocument(t, path, "demo", "keep") if _, ok := model["vendorOption"]; ok { t.Fatal("explicit editor save must allow removal of unknown fields") } - if _, ok := model["__piSwitchReplaceDocument"]; ok { + if _, ok := model["replaceDocument"]; ok { t.Fatal("internal replacement marker must not be persisted") } } @@ -156,3 +394,27 @@ func TestReplaceModelsPreservesUnknownFieldsForRetainedModels(t *testing.T) { t.Fatalf("name = %q, want New name", name) } } + +func readModelDocument(t *testing.T, path string, providerID string, modelID string) map[string]json.RawMessage { + t.Helper() + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + var document struct { + Providers map[string]struct { + Models []map[string]json.RawMessage `json:"models"` + } `json:"providers"` + } + if err := json.Unmarshal(data, &document); err != nil { + t.Fatal(err) + } + for _, model := range document.Providers[providerID].Models { + var id string + if err := json.Unmarshal(model["id"], &id); err == nil && id == modelID { + return model + } + } + t.Fatalf("model %s/%s not found", providerID, modelID) + return nil +} diff --git a/internal/provider/types.go b/internal/provider/types.go index 6d493f6..754c1fe 100644 --- a/internal/provider/types.go +++ b/internal/provider/types.go @@ -1,6 +1,8 @@ package provider import ( + "crypto/sha256" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -68,11 +70,6 @@ func (cfg Config) MarshalJSON() ([]byte, error) { return json.Marshal(fields) } -const ( - ModelReplaceDocumentField = "__piSwitchReplaceDocument" - ModelOriginalIDField = "__piSwitchOriginalId" -) - type ModelInfo struct { ID string `json:"id"` Name string `json:"name"` @@ -87,18 +84,34 @@ type ModelInfo struct { SamplingParams map[string]any `json:"samplingParams,omitempty"` Headers map[string]string `json:"headers,omitempty"` Compat map[string]any `json:"compat,omitempty"` - ExtraFields map[string]any `json:"extraFields,omitempty"` Extra map[string]json.RawMessage `json:"-"` + CompatRaw map[string]json.RawMessage `json:"-"` + Revision string `json:"-"` + ReplaceDocument bool `json:"-"` + OriginalID string `json:"-"` } var modelFields = map[string]struct{}{ "id": {}, "name": {}, "api": {}, "baseUrl": {}, "reasoning": {}, "thinkingLevelMap": {}, "input": {}, - "cost": {}, "contextWindow": {}, "maxTokens": {}, "samplingParams": {}, "headers": {}, "compat": {}, "extraFields": {}, + "cost": {}, "contextWindow": {}, "maxTokens": {}, "samplingParams": {}, "headers": {}, "compat": {}, +} + +var transportCompatFields = map[string]struct{}{ + "supportsStore": {}, "supportsDeveloperRole": {}, "supportsReasoningEffort": {}, "supportsUsageInStreaming": {}, + "supportsFinishReason": {}, "requiresToolResultName": {}, "requiresAssistantAfterToolResult": {}, + "requiresThinkingAsText": {}, "requiresReasoningContentOnAssistantMessages": {}, "supportsOpenAIGrammarTools": {}, + "supportsStrictMode": {}, "sendSessionAffinityHeaders": {}, "supportsLongCacheRetention": {}, + "supportsEagerToolInputStreaming": {}, "supportsCacheControlOnTools": {}, "supportsTemperature": {}, + "forceAdaptiveThinking": {}, "allowEmptySignature": {}, "supportsStrictTools": {}, "supportsToolReferences": {}, + "supportsToolSearch": {}, "zaiToolStream": {}, "supportsThinkingTokenBudget": {}, "supportsAdditionalTools": {}, + "supportsExplicitPromptCacheMode": {}, "maxTokensField": {}, "thinkingFormat": {}, "cacheControlFormat": {}, + "deferredToolsMode": {}, "sessionAffinityFormat": {}, "chatTemplateKwargs": {}, "chatTemplateArgs": {}, + "openRouterRouting": {}, "vercelGatewayRouting": {}, } -// UnmarshalJSON keeps model-specific fields that Pi or a provider may support. -// This lets the editor expose the complete model object without losing fields -// unknown to Pi Switch. +// UnmarshalJSON keeps model-specific fields in their original JSON form. The +// raw values are the persistence source of truth; transport conversion is kept +// separate so a real model field can never collide with Pi Switch metadata. func (model *ModelInfo) UnmarshalJSON(data []byte) error { type modelAlias ModelInfo var decoded modelAlias @@ -109,25 +122,40 @@ func (model *ModelInfo) UnmarshalJSON(data []byte) error { if err := json.Unmarshal(data, &raw); err != nil { return err } - for field := range modelFields { - delete(raw, field) + canonical, err := json.Marshal(raw) + if err != nil { + return err } - *model = ModelInfo(decoded) - model.Extra = raw - if len(raw) > 0 { - if model.ExtraFields == nil { - model.ExtraFields = make(map[string]any, len(raw)) + compatRaw := map[string]json.RawMessage{} + structuredCompat := map[string]any{} + if value, ok := raw["compat"]; ok && len(value) > 0 { + if err := json.Unmarshal(value, &compatRaw); err != nil { + return err } - for field, value := range raw { + for field, rawValue := range compatRaw { + if _, known := transportCompatFields[field]; !known { + continue + } var decodedValue any - if err := json.Unmarshal(value, &decodedValue); err != nil { + if err := json.Unmarshal(rawValue, &decodedValue); err != nil { return err } - if _, exists := model.ExtraFields[field]; !exists { - model.ExtraFields[field] = decodedValue - } + structuredCompat[field] = decodedValue } } + for field := range modelFields { + delete(raw, field) + } + *model = ModelInfo(decoded) + if len(structuredCompat) == 0 { + model.Compat = nil + } else { + model.Compat = structuredCompat + } + model.Extra = raw + model.CompatRaw = compatRaw + hash := sha256.Sum256(canonical) + model.Revision = hex.EncodeToString(hash[:]) return nil } @@ -141,29 +169,299 @@ func (model ModelInfo) MarshalJSON() ([]byte, error) { if err := json.Unmarshal(data, &fields); err != nil { return nil, err } - extraFields := make(map[string]any, len(model.Extra)+len(model.ExtraFields)) - for field, value := range model.Extra { - var decodedValue any - if err := json.Unmarshal(value, &decodedValue); err != nil { + + compatFields := cloneRawFields(model.CompatRaw) + if len(model.Compat) > 0 { + encodedCompat, err := json.Marshal(model.Compat) + if err != nil { return nil, err } - extraFields[field] = decodedValue - } - for field, value := range model.ExtraFields { - extraFields[field] = value + structuredCompat := map[string]json.RawMessage{} + if err := json.Unmarshal(encodedCompat, &structuredCompat); err != nil { + return nil, err + } + for field, value := range structuredCompat { + compatFields[field] = value + } } - if len(extraFields) == 0 { - delete(fields, "extraFields") + if len(compatFields) == 0 { + delete(fields, "compat") } else { - encodedExtraFields, err := json.Marshal(extraFields) + fields["compat"], err = json.Marshal(compatFields) if err != nil { return nil, err } - fields["extraFields"] = encodedExtraFields + } + + for field, value := range model.Extra { + if _, known := modelFields[field]; !known { + fields[field] = value + } } return json.Marshal(fields) } +// ModelTransport is the Wails-facing model DTO. Unknown model fields are sent +// as raw JSON text so JavaScript never coerces large integers, and transport +// metadata stays outside the models.json namespace. +type ModelTransport struct { + ID string `json:"id"` + Name string `json:"name"` + API string `json:"api,omitempty"` + BaseURL string `json:"baseUrl,omitempty"` + Reasoning bool `json:"reasoning"` + ThinkingLevelMap map[string]any `json:"thinkingLevelMap,omitempty"` + Input []string `json:"input,omitempty"` + Cost map[string]any `json:"cost,omitempty"` + ContextWindow int `json:"contextWindow,omitempty"` + MaxTokens int `json:"maxTokens,omitempty"` + SamplingParams map[string]any `json:"samplingParams,omitempty"` + Headers map[string]string `json:"headers,omitempty"` + Compat map[string]any `json:"compat,omitempty"` + CompatExtraFieldsJSON string `json:"compatExtraFieldsJson,omitempty"` + ExtraFieldsJSON string `json:"extraFieldsJson,omitempty"` + Revision string `json:"revision,omitempty"` + ReplaceDocument bool `json:"replaceDocument,omitempty"` + OriginalID string `json:"originalId,omitempty"` +} + +func IsTransportCompatField(field string) bool { + _, known := transportCompatFields[field] + return known +} + +func NewModelTransport(model ModelInfo) (ModelTransport, error) { + compat, compatExtra, err := splitCompatForTransport(model) + if err != nil { + return ModelTransport{}, err + } + extraFields := cloneRawFields(model.Extra) + delete(extraFields, "selected") + extraJSON, err := marshalRawObject(extraFields) + if err != nil { + return ModelTransport{}, err + } + return ModelTransport{ + ID: model.ID, Name: model.Name, API: model.API, BaseURL: model.BaseURL, Reasoning: model.Reasoning, + ThinkingLevelMap: model.ThinkingLevelMap, Input: model.Input, Cost: model.Cost, + ContextWindow: model.ContextWindow, MaxTokens: model.MaxTokens, SamplingParams: model.SamplingParams, + Headers: model.Headers, Compat: compat, CompatExtraFieldsJSON: compatExtra, + ExtraFieldsJSON: extraJSON, Revision: model.Revision, + }, nil +} + +func (transport ModelTransport) ModelInfo() (ModelInfo, error) { + extra, err := parseRawObject(transport.ExtraFieldsJSON) + if err != nil { + return ModelInfo{}, fmt.Errorf("其他模型字段无效:%w", err) + } + for field := range modelFields { + delete(extra, field) + } + delete(extra, "selected") + + compatRaw, err := parseRawObject(transport.CompatExtraFieldsJSON) + if err != nil { + return ModelInfo{}, fmt.Errorf("其他 compat 字段无效:%w", err) + } + for field := range transportCompatFields { + delete(compatRaw, field) + } + return ModelInfo{ + ID: transport.ID, Name: transport.Name, API: transport.API, BaseURL: transport.BaseURL, + Reasoning: transport.Reasoning, ThinkingLevelMap: transport.ThinkingLevelMap, Input: transport.Input, + Cost: transport.Cost, ContextWindow: transport.ContextWindow, MaxTokens: transport.MaxTokens, + SamplingParams: transport.SamplingParams, Headers: transport.Headers, Compat: transport.Compat, + Extra: extra, CompatRaw: compatRaw, Revision: transport.Revision, + ReplaceDocument: transport.ReplaceDocument, OriginalID: transport.OriginalID, + }, nil +} + +func ModelTransports(models []ModelInfo) ([]ModelTransport, error) { + result := make([]ModelTransport, 0, len(models)) + for _, model := range models { + converted, err := NewModelTransport(model) + if err != nil { + return nil, err + } + result = append(result, converted) + } + return result, nil +} + +func ModelsFromTransport(models []ModelTransport) ([]ModelInfo, error) { + result := make([]ModelInfo, 0, len(models)) + for _, model := range models { + converted, err := model.ModelInfo() + if err != nil { + return nil, err + } + result = append(result, converted) + } + return result, nil +} + +type ModelListTransport struct { + Models []ModelTransport `json:"models"` + Revision string `json:"revision"` +} + +func NewModelListTransport(models []ModelInfo) (ModelListTransport, error) { + converted, err := ModelTransports(models) + if err != nil { + return ModelListTransport{}, err + } + return ModelListTransport{Models: converted, Revision: ModelListRevision(models)}, nil +} + +func ModelListRevision(models []ModelInfo) string { + entries := make([][2]string, 0, len(models)) + for _, model := range models { + if model.Revision == "" { + return "" + } + entries = append(entries, [2]string{model.ID, model.Revision}) + } + encoded, err := json.Marshal(entries) + if err != nil { + return "" + } + hash := sha256.Sum256(encoded) + return hex.EncodeToString(hash[:]) +} + +func splitCompatForTransport(model ModelInfo) (map[string]any, string, error) { + structured := map[string]any{} + extra := cloneRawFields(model.CompatRaw) + for field, value := range model.CompatRaw { + if _, known := transportCompatFields[field]; !known { + continue + } + var decoded any + if err := json.Unmarshal(value, &decoded); err != nil { + return nil, "", err + } + structured[field] = decoded + delete(extra, field) + } + for field, value := range model.Compat { + if _, known := transportCompatFields[field]; known { + structured[field] = value + delete(extra, field) + continue + } + if _, exists := extra[field]; !exists { + encoded, err := json.Marshal(value) + if err != nil { + return nil, "", err + } + extra[field] = encoded + } + } + extraJSON, err := marshalRawObject(extra) + if len(structured) == 0 { + structured = nil + } + return structured, extraJSON, err +} + +func parseRawObject(value string) (map[string]json.RawMessage, error) { + if strings.TrimSpace(value) == "" { + return map[string]json.RawMessage{}, nil + } + fields := map[string]json.RawMessage{} + if err := json.Unmarshal([]byte(value), &fields); err != nil { + return nil, err + } + return fields, nil +} + +func marshalRawObject(fields map[string]json.RawMessage) (string, error) { + if len(fields) == 0 { + return "", nil + } + encoded, err := json.MarshalIndent(fields, "", " ") + return string(encoded), err +} + +func cloneRawFields(input map[string]json.RawMessage) map[string]json.RawMessage { + cloned := make(map[string]json.RawMessage, len(input)) + for field, value := range input { + cloned[field] = value + } + return cloned +} + +type ConfigTransport struct { + ID string `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + BaseURL string `json:"baseUrl"` + APIKeyEnv string `json:"apiKeyEnv"` + APIKeyLiteral string `json:"apiKeyLiteral"` + API string `json:"api"` + Proxy string `json:"proxy"` + HeaderMode string `json:"headerMode"` + Headers map[string]string `json:"headers"` + CustomHeaders map[string]string `json:"customHeaders,omitempty"` + Models []ModelTransport `json:"models"` + Host string `json:"host"` + SelectedModelID string `json:"selectedModelId"` + ExtraFieldsJSON string `json:"extraFieldsJson,omitempty"` + ModelsRevision string `json:"modelsRevision,omitempty"` +} + +func NewConfigTransport(cfg Config) (ConfigTransport, error) { + models, err := ModelTransports(cfg.Models) + if err != nil { + return ConfigTransport{}, err + } + extraJSON, err := marshalRawObject(cfg.Extra) + if err != nil { + return ConfigTransport{}, err + } + return ConfigTransport{ + ID: cfg.ID, Name: cfg.Name, Type: cfg.Type, BaseURL: cfg.BaseURL, + APIKeyEnv: cfg.APIKeyEnv, APIKeyLiteral: cfg.APIKeyLiteral, API: cfg.API, Proxy: cfg.Proxy, + HeaderMode: cfg.HeaderMode, Headers: cfg.Headers, CustomHeaders: cfg.CustomHeaders, + Models: models, Host: cfg.Host, SelectedModelID: cfg.SelectedModelID, ExtraFieldsJSON: extraJSON, + ModelsRevision: ModelListRevision(cfg.Models), + }, nil +} + +func (transport ConfigTransport) Config() (Config, error) { + models, err := ModelsFromTransport(transport.Models) + if err != nil { + return Config{}, err + } + extra, err := parseRawObject(transport.ExtraFieldsJSON) + if err != nil { + return Config{}, fmt.Errorf("其他 Provider 字段无效:%w", err) + } + for field := range configFields { + delete(extra, field) + } + return Config{ + ID: transport.ID, Name: transport.Name, Type: transport.Type, BaseURL: transport.BaseURL, + APIKeyEnv: transport.APIKeyEnv, APIKeyLiteral: transport.APIKeyLiteral, API: transport.API, + Proxy: transport.Proxy, HeaderMode: transport.HeaderMode, Headers: transport.Headers, + CustomHeaders: transport.CustomHeaders, Models: models, Host: transport.Host, + SelectedModelID: transport.SelectedModelID, Extra: extra, + }, nil +} + +func ConfigTransports(configs []Config) ([]ConfigTransport, error) { + result := make([]ConfigTransport, 0, len(configs)) + for _, cfg := range configs { + converted, err := NewConfigTransport(cfg) + if err != nil { + return nil, err + } + result = append(result, converted) + } + return result, nil +} + type ConnectionTestResult struct { OK bool `json:"ok"` Title string `json:"title"` diff --git a/internal/provider/types_test.go b/internal/provider/types_test.go new file mode 100644 index 0000000..b4f7a30 --- /dev/null +++ b/internal/provider/types_test.go @@ -0,0 +1,180 @@ +package provider + +import ( + "encoding/json" + "testing" +) + +func TestModelTransportPreservesUnknownRawJSON(t *testing.T) { + var model ModelInfo + if err := json.Unmarshal([]byte(`{ + "id":"demo", + "name":"Demo", + "reasoning":false, + "extraFields":{"vendorFlag":true}, + "vendorId":9007199254740993, + "compat":{ + "supportsDeveloperRole":false, + "futureCounter":9007199254740993 + } +}`), &model); err != nil { + t.Fatal(err) + } + + transport, err := NewModelTransport(model) + if err != nil { + t.Fatal(err) + } + if transport.ExtraFieldsJSON == "" || transport.CompatExtraFieldsJSON == "" { + t.Fatalf("transport lost raw fields: %#v", transport) + } + + roundTripped, err := transport.ModelInfo() + if err != nil { + t.Fatal(err) + } + encoded, err := json.Marshal(roundTripped) + if err != nil { + t.Fatal(err) + } + var fields map[string]json.RawMessage + if err := json.Unmarshal(encoded, &fields); err != nil { + t.Fatal(err) + } + if got := string(fields["vendorId"]); got != "9007199254740993" { + t.Fatalf("vendorId = %s", got) + } + if _, ok := fields["extraFields"]; !ok { + t.Fatal("real top-level extraFields field was not preserved") + } + var compat map[string]json.RawMessage + if err := json.Unmarshal(fields["compat"], &compat); err != nil { + t.Fatal(err) + } + if got := string(compat["futureCounter"]); got != "9007199254740993" { + t.Fatalf("futureCounter = %s", got) + } + if got := string(compat["supportsDeveloperRole"]); got != "false" { + t.Fatalf("supportsDeveloperRole = %s", got) + } +} + +func TestModelTransportDoesNotInventPersistenceRevision(t *testing.T) { + transport, err := NewModelTransport(ModelInfo{ID: "fetched", Name: "Fetched"}) + if err != nil { + t.Fatal(err) + } + if transport.Revision != "" { + t.Fatalf("transient fetched model received persistence revision %q", transport.Revision) + } +} +func TestModelTransportFiltersSelectedFromUnknownFields(t *testing.T) { + model := ModelInfo{ + ID: "demo", + Name: "Demo", + Extra: map[string]json.RawMessage{ + "selected": json.RawMessage(`true`), + "vendorFlag": json.RawMessage(`1`), + }, + } + transport, err := NewModelTransport(model) + if err != nil { + t.Fatal(err) + } + fields, err := parseRawObject(transport.ExtraFieldsJSON) + if err != nil { + t.Fatal(err) + } + if _, ok := fields["selected"]; ok { + t.Fatal("UI-only selected field leaked into the transport envelope") + } + if _, ok := fields["vendorFlag"]; !ok { + t.Fatal("real unknown field was removed with selected") + } +} + +func TestModelTransportMetadataDoesNotConsumeSameNamedUnknownFields(t *testing.T) { + transport := ModelTransport{ + ID: "demo", + Name: "Demo", + Revision: "transport-revision", + ReplaceDocument: true, + OriginalID: "old-demo", + ExtraFieldsJSON: `{ + "revision":"provider-revision", + "replaceDocument":"provider-value", + "originalId":"provider-original" +}`, + } + model, err := transport.ModelInfo() + if err != nil { + t.Fatal(err) + } + if model.Revision != "transport-revision" || !model.ReplaceDocument || model.OriginalID != "old-demo" { + t.Fatalf("transport metadata = %#v", model) + } + encoded, err := json.Marshal(model) + if err != nil { + t.Fatal(err) + } + var fields map[string]json.RawMessage + if err := json.Unmarshal(encoded, &fields); err != nil { + t.Fatal(err) + } + if string(fields["revision"]) != `"provider-revision"` { + t.Fatalf("real revision field = %s", fields["revision"]) + } + if string(fields["replaceDocument"]) != `"provider-value"` { + t.Fatalf("real replaceDocument field = %s", fields["replaceDocument"]) + } + if string(fields["originalId"]) != `"provider-original"` { + t.Fatalf("real originalId field = %s", fields["originalId"]) + } +} + +func TestConfigTransportIncludesStableModelListRevision(t *testing.T) { + var persisted ModelInfo + if err := json.Unmarshal([]byte(`{"id":"model","name":"Model"}`), &persisted); err != nil { + t.Fatal(err) + } + cfg := Config{ + ID: "demo", + Name: "Demo", + Models: []ModelInfo{persisted}, + } + first, err := NewConfigTransport(cfg) + if err != nil { + t.Fatal(err) + } + second, err := NewConfigTransport(cfg) + if err != nil { + t.Fatal(err) + } + if first.ModelsRevision == "" || first.ModelsRevision != second.ModelsRevision { + t.Fatalf("model revisions = %q and %q", first.ModelsRevision, second.ModelsRevision) + } + if len(first.Models) != 1 || first.Models[0].Revision == "" { + t.Fatalf("model transport revision missing: %#v", first.Models) + } +} +func TestConfigTransportPreservesProviderUnknownFields(t *testing.T) { + cfg := Config{ + ID: "demo", + Name: "Demo", + Models: []ModelInfo{{ID: "model", Name: "Model"}}, + Extra: map[string]json.RawMessage{ + "providerVendorId": json.RawMessage(`9007199254740993`), + }, + } + transport, err := NewConfigTransport(cfg) + if err != nil { + t.Fatal(err) + } + roundTripped, err := transport.Config() + if err != nil { + t.Fatal(err) + } + if got := string(roundTripped.Extra["providerVendorId"]); got != "9007199254740993" { + t.Fatalf("providerVendorId = %s", got) + } +}