import type { ExtensionAPI, ExtensionCommandContext } from "@earendil-works/pi-coding-agent"; import { Runtime } from "pi-ccs/extensions/runtime.ts"; import { resolveSqlitePath } from "pi-ccs/src/sqlite-path.ts"; import { isSwitchable } from "pi-ccs/src/parse/index.ts"; import { fetchRemoteModels, mergeModelLists } from "pi-ccs/src/models-fetch.ts"; import { registerProvider, type PiRegisterApi } from "pi-ccs/src/register.ts"; import { summarizeModelMeta } from "pi-ccs/src/model-meta.ts"; import type { CcProvider, PiSwitchSelection } from "pi-ccs/src/types.ts"; function registerCurrentProvider(pi: ExtensionAPI, rt: Runtime, provider: CcProvider, modelId: string): boolean { return registerProvider(pi as unknown as PiRegisterApi, provider, [modelId], { rules: rt.headerRules, ...rt.headerOverrideOpts(provider), vars: rt.headerVars(), debug: rt.config.debug, onReject: rt.rejectSink(), modelMeta: rt.modelMetaFor(provider), }); } async function activate( pi: ExtensionAPI, rt: Runtime, ctx: ExtensionCommandContext, provider: CcProvider, modelId: string, ): Promise { if (!registerCurrentProvider(pi, rt, provider, modelId)) { ctx.ui.notify(`切换失败:${provider.parseError ?? "无法注册当前渠道"}`, "error"); return false; } const model = ctx.modelRegistry.find(provider.piName, modelId); if (!model) { ctx.ui.notify(`切换失败:未找到模型 ${provider.piName} / ${modelId}`, "error"); return false; } if (!(await pi.setModel(model))) { ctx.ui.notify("切换失败:Pi 未能激活该模型", "error"); return false; } const selection: PiSwitchSelection = { dbId: provider.id, model: modelId, tab: provider.appType, appType: provider.appType, provider: provider.piName, }; const saved = rt.state.saveSelection(selection); const recent = rt.state.recordRecent({ dbId: provider.id, model: modelId }); if (!recent.ok && rt.config.debug) { console.warn("[pi-ccs-current-model] could not save recent:", recent.error); } const metaHint = summarizeModelMeta(rt.modelMetaFor(provider)); ctx.ui.notify( saved.ok ? `已切换到 ${provider.displayName} · ${modelId}(${metaHint})` : `已切换到 ${provider.displayName} · ${modelId},但未保存:${saved.error ?? "unknown"}`, saved.ok ? "info" : "warning", ); ctx.ui.setStatus?.("pi-switch", `${modelId} @ ${provider.appType}/${provider.displayName}`); return true; } async function pickCurrentModel( pi: ExtensionAPI, rt: Runtime, ctx: ExtensionCommandContext, args: string, ): Promise { rt.reloadConfig(); rt.reloadHeaderRules(); const { providers, error } = rt.refreshSnapshot(); if (error) ctx.ui.notify(error, "warning"); const selection = rt.state.readOrMigrateSelection(providers); if (!selection?.dbId) { ctx.ui.notify("还没有当前渠道,请先使用 /ccs 选择一次 Provider", "warning"); return; } const provider = providers.find((item) => item.id === selection.dbId); if (!provider) { ctx.ui.notify("当前渠道在 cc-switch 数据库中不存在,请重新用 /ccs 选择", "warning"); return; } if (!isSwitchable(provider)) { ctx.ui.notify(`当前渠道不可切换:${provider.parseError ?? "unknown"}`, "warning"); return; } const directModel = args.trim(); if (directModel) { await activate(pi, rt, ctx, provider, directModel); return; } let remoteModels: string[] = []; try { const userAgent = rt.overridesFor(provider)?.headers?.["User-Agent"]; const result = await fetchRemoteModels({ baseUrl: provider.baseUrl, apiKey: provider.apiKey, modelsUrl: provider.modelsUrl, isFullUrl: provider.isFullUrl, userAgent, }); if (result.error) { ctx.ui.notify(`远程模型刷新失败,使用本地列表:${result.error}`, "warning"); } else { remoteModels = result.models; } } catch (error) { const message = error instanceof Error ? error.message : String(error); ctx.ui.notify(`远程模型刷新失败,使用本地列表:${message}`, "warning"); } const models = mergeModelLists(provider.configModels, remoteModels); if (!models.length) { const manual = await ctx.ui.input(`手动输入模型 ID · ${provider.displayName}`, selection.model ?? ""); if (manual?.trim()) await activate(pi, rt, ctx, provider, manual.trim()); return; } const ordered = [ ...models.filter((id) => id === selection.model), ...models.filter((id) => id !== selection.model), ]; const labels = ordered.map((id) => (id === selection.model ? `★ ${id}` : id)); const manualLabel = "手动输入模型 ID"; const picked = await ctx.ui.select( `当前渠道模型 · ${provider.appType}/${provider.displayName}`, [...labels, manualLabel], ); if (!picked) return; if (picked === manualLabel) { const manual = await ctx.ui.input("手动输入模型 ID", selection.model ?? ""); if (manual?.trim()) await activate(pi, rt, ctx, provider, manual.trim()); return; } const modelId = ordered[labels.indexOf(picked)]; if (modelId) await activate(pi, rt, ctx, provider, modelId); } export default async function currentModelExtension(pi: ExtensionAPI): Promise { const [cp, fs, os] = await Promise.all([ import("node:child_process"), import("node:fs"), import("node:os"), ]); const rt = new Runtime({ execFileSync: cp.execFileSync, existsSync: fs.existsSync, readFileSync: fs.readFileSync, writeFileSync: fs.writeFileSync, renameSync: fs.renameSync, release: os.release(), home: os.homedir(), }); rt.reloadConfig(); rt.reloadHeaderRules(); const sqlite = resolveSqlitePath({ configPath: rt.config.sqlitePath, exists: rt.io.existsSync }); rt.sqlite3Path = sqlite.path ?? ""; rt.sqlite3Tried = sqlite.tried ?? []; const handler = async (args: string, ctx: ExtensionCommandContext) => { if (!rt.sqlite3Path) { ctx.ui.notify("未找到 sqlite3;请配置 SQLITE3_PATH 或 pi-switch.json.sqlitePath", "error"); return; } await pickCurrentModel(pi, rt, ctx, args); }; pi.registerCommand("ccs-model", { description: "只选择当前 cc-switch Provider 的模型;可附模型 ID 直接切换", handler, }); pi.registerCommand("ps-model", { description: "ccs-model 别名:只切换当前 Provider 的模型", handler, }); }