From 733a46c2260d17dcd0f3ab763b588bb7a7a77d11 Mon Sep 17 00:00:00 2001 From: can1357 Date: Fri, 13 Feb 2026 12:54:45 +0100 Subject: [PATCH] fix(coding-agent): discovered ollama models at runtime --- packages/ai/test/google-tool-schema.test.ts | 8 +- packages/coding-agent/src/cli/setup-cli.ts | 5 +- .../coding-agent/src/commit/agentic/index.ts | 1 + packages/coding-agent/src/commit/pipeline.ts | 1 + .../coding-agent/src/config/model-registry.ts | 203 +++++++++++++++++- packages/coding-agent/src/ipy/runtime.ts | 3 +- packages/coding-agent/src/main.ts | 1 + packages/coding-agent/src/sdk.ts | 1 + packages/coding-agent/src/task/executor.ts | 1 + .../test/auth-storage-rotation.test.ts | 2 +- .../coding-agent/test/model-registry.test.ts | 84 +++++++- 11 files changed, 283 insertions(+), 27 deletions(-) diff --git a/packages/ai/test/google-tool-schema.test.ts b/packages/ai/test/google-tool-schema.test.ts index 7be128879..e05ff788b 100644 --- a/packages/ai/test/google-tool-schema.test.ts +++ b/packages/ai/test/google-tool-schema.test.ts @@ -1,7 +1,11 @@ -import type { TSchema } from "@sinclair/typebox"; import { describe, expect, it } from "bun:test"; -import { convertTools, sanitizeSchemaForCloudCodeAssistClaude, sanitizeSchemaForGoogle } from "@oh-my-pi/pi-ai/providers/google-shared"; +import { + convertTools, + sanitizeSchemaForCloudCodeAssistClaude, + sanitizeSchemaForGoogle, +} from "@oh-my-pi/pi-ai/providers/google-shared"; import type { Model, Tool } from "@oh-my-pi/pi-ai/types"; +import type { TSchema } from "@sinclair/typebox"; function createModel(id: string): Model<"google-gemini-cli"> { return { diff --git a/packages/coding-agent/src/cli/setup-cli.ts b/packages/coding-agent/src/cli/setup-cli.ts index abd7bfffe..a6b12f9c6 100644 --- a/packages/coding-agent/src/cli/setup-cli.ts +++ b/packages/coding-agent/src/cli/setup-cli.ts @@ -267,7 +267,9 @@ async function handlePythonSetup(flags: { json?: boolean; check?: boolean }): Pr if (install.usedManagedEnv) { if (check.uvPath) { console.error(chalk.dim(` uv venv ${MANAGED_PYTHON_ENV}`)); - console.error(chalk.dim(` uv pip install --python ${MANAGED_PYTHON_ENV} ${check.missingPackages.join(" ")}`)); + console.error( + chalk.dim(` uv pip install --python ${MANAGED_PYTHON_ENV} ${check.missingPackages.join(" ")}`), + ); } else { console.error(chalk.dim(` ${check.pythonPath} -m venv ${MANAGED_PYTHON_ENV}`)); console.error(chalk.dim(` ${managedPythonPath()} -m pip install ${check.missingPackages.join(" ")}`)); @@ -284,7 +286,6 @@ async function handlePythonSetup(flags: { json?: boolean; check?: boolean }): Pr if (recheck.usingManagedEnv) { console.log(chalk.dim(`Managed Python environment: ${recheck.managedEnvPath}`)); } - } else { console.error(chalk.red(`\n${theme.status.error} Setup incomplete`)); console.error(chalk.dim(`Still missing: ${recheck.missingPackages.join(", ")}`)); diff --git a/packages/coding-agent/src/commit/agentic/index.ts b/packages/coding-agent/src/commit/agentic/index.ts index b9f70fdbb..a7e281776 100644 --- a/packages/coding-agent/src/commit/agentic/index.ts +++ b/packages/coding-agent/src/commit/agentic/index.ts @@ -32,6 +32,7 @@ export async function runAgenticCommit(args: CommitCommandArgs): Promise { writeStdout("● Resolving model..."); const modelRegistry = new ModelRegistry(authStorage); + await modelRegistry.refresh(); const stagedFilesPromise = (async () => { let stagedFiles = await git.getStagedFiles(); if (stagedFiles.length === 0) { diff --git a/packages/coding-agent/src/commit/pipeline.ts b/packages/coding-agent/src/commit/pipeline.ts index 809b8bc34..ce75b783e 100644 --- a/packages/coding-agent/src/commit/pipeline.ts +++ b/packages/coding-agent/src/commit/pipeline.ts @@ -43,6 +43,7 @@ async function runLegacyCommitCommand(args: CommitCommandArgs): Promise { const commitSettings = settings.getGroup("commit"); const authStorage = await discoverAuthStorage(); const modelRegistry = new ModelRegistry(authStorage); + await modelRegistry.refresh(); const { model: primaryModel, apiKey: primaryApiKey } = await resolvePrimaryModel( args.model, diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index 34fe9d7ff..237a6667d 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -6,6 +6,7 @@ import { type Model, normalizeDomain, } from "@oh-my-pi/pi-ai"; +import { logger } from "@oh-my-pi/pi-utils"; import { type Static, Type } from "@sinclair/typebox"; import AjvModule from "ajv"; import { type ConfigError, ConfigFile } from "../config"; @@ -106,6 +107,12 @@ const ModelOverrideSchema = Type.Object({ type ModelOverride = Static; +const ProviderDiscoverySchema = Type.Object({ + type: Type.Union([Type.Literal("ollama")]), +}); + +const ProviderAuthSchema = Type.Union([Type.Literal("apiKey"), Type.Literal("none")]); + const ProviderConfigSchema = Type.Object({ baseUrl: Type.Optional(Type.String({ minLength: 1 })), apiKey: Type.Optional(Type.String({ minLength: 1 })), @@ -122,6 +129,8 @@ const ProviderConfigSchema = Type.Object({ ), headers: Type.Optional(Type.Record(Type.String(), Type.String())), authHeader: Type.Optional(Type.Boolean()), + auth: Type.Optional(ProviderAuthSchema), + discovery: Type.Optional(ProviderDiscoverySchema), models: Type.Optional(Type.Array(ModelDefinitionSchema)), modelOverrides: Type.Optional(Type.Record(Type.String(), ModelOverrideSchema)), }); @@ -132,6 +141,9 @@ const ModelsConfigSchema = Type.Object({ type ModelsConfig = Static; +type ProviderAuthMode = Static; +type ProviderDiscovery = Static; + export const ModelsConfigFile = new ConfigFile("models", ModelsConfigSchema).withValidation( "models", config => { @@ -140,22 +152,30 @@ export const ModelsConfigFile = new ConfigFile("models", ModelsCon const models = providerConfig.models ?? []; if (models.length === 0) { - // Override-only config: needs baseUrl or modelOverrides + // Override-only config: needs baseUrl, modelOverrides, or discovery const hasModelOverrides = providerConfig.modelOverrides && Object.keys(providerConfig.modelOverrides).length > 0; - if (!providerConfig.baseUrl && !hasModelOverrides) { - throw new Error(`Provider ${providerName}: must specify "baseUrl", "modelOverrides", or "models".`); + if (!providerConfig.baseUrl && !hasModelOverrides && !providerConfig.discovery) { + throw new Error( + `Provider ${providerName}: must specify "baseUrl", "modelOverrides", "discovery", or "models".`, + ); } } else { - // Full replacement: needs baseUrl and apiKey + // Full replacement: needs baseUrl and apiKey unless auth is disabled if (!providerConfig.baseUrl) { throw new Error(`Provider ${providerName}: "baseUrl" is required when defining custom models.`); } - if (!providerConfig.apiKey) { - throw new Error(`Provider ${providerName}: "apiKey" is required when defining custom models.`); + if (!providerConfig.apiKey && providerConfig.auth !== "none") { + throw new Error( + `Provider ${providerName}: "apiKey" is required when defining custom models unless auth is "none".`, + ); } } + if (providerConfig.discovery && !providerConfig.api) { + throw new Error(`Provider ${providerName}: "api" is required when discovery is enabled at provider level.`); + } + for (const modelDef of models) { const hasModelApi = !!modelDef.api; @@ -183,6 +203,14 @@ interface ProviderOverride { apiKey?: string; } +interface DiscoveryProviderConfig { + provider: string; + api: Api; + baseUrl?: string; + headers?: Record; + discovery: ProviderDiscovery; +} + /** * Serialized representation of ModelRegistry for passing to subagent workers. */ @@ -196,6 +224,8 @@ interface CustomModelsResult { models?: Model[]; overrides?: Map; modelOverrides?: Map>; + keylessProviders?: Set; + discoverableProviders?: DiscoveryProviderConfig[]; error?: ConfigError; found: boolean; } @@ -255,6 +285,9 @@ function applyModelOverride(model: Model, override: ModelOverride): Model[] = []; #customProviderApiKeys: Map = new Map(); + #keylessProviders: Set = new Set(); + #discoverableProviders: DiscoveryProviderConfig[] = []; + #modelOverrides: Map> = new Map(); #configError: ConfigError | undefined = undefined; #modelsConfigFile: ConfigFile; @@ -281,11 +314,15 @@ export class ModelRegistry { /** * Reload models from disk (built-in + custom from models.json). */ - refresh(): void { + async refresh(): Promise { this.#modelsConfigFile.invalidate(); this.#customProviderApiKeys.clear(); + this.#keylessProviders.clear(); + this.#discoverableProviders = []; + this.#modelOverrides.clear(); this.#configError = undefined; this.#loadModels(); + await this.#refreshRuntimeDiscoveries(); } /** @@ -301,9 +338,14 @@ export class ModelRegistry { models: customModels = [], overrides = new Map(), modelOverrides = new Map(), + keylessProviders = new Set(), + discoverableProviders = [], error: configError, } = this.#loadCustomModels(); this.#configError = configError; + this.#keylessProviders = keylessProviders; + this.#discoverableProviders = discoverableProviders; + this.#modelOverrides = modelOverrides; const builtInModels = this.#loadBuiltInModels(overrides, modelOverrides); const combined = this.#mergeCustomModels(builtInModels, customModels); @@ -367,13 +409,30 @@ export class ModelRegistry { const { value, error, status } = this.#modelsConfigFile.tryLoad(); if (status === "error") { - return { models: [], overrides: new Map(), modelOverrides: new Map(), error, found: true }; + return { + models: [], + overrides: new Map(), + modelOverrides: new Map(), + keylessProviders: new Set(), + discoverableProviders: [], + error, + found: true, + }; } else if (status === "not-found") { - return { models: [], overrides: new Map(), modelOverrides: new Map(), found: false }; + return { + models: [], + overrides: new Map(), + modelOverrides: new Map(), + keylessProviders: new Set(), + discoverableProviders: [], + found: false, + }; } const overrides = new Map(); const allModelOverrides = new Map>(); + const keylessProviders = new Set(); + const discoverableProviders: DiscoveryProviderConfig[] = []; for (const [providerName, providerConfig] of Object.entries(value.providers)) { // Always set overrides when baseUrl/headers present @@ -385,6 +444,21 @@ export class ModelRegistry { }); } + const authMode = (providerConfig.auth ?? "apiKey") as ProviderAuthMode; + if (authMode === "none") { + keylessProviders.add(providerName); + } + + if (providerConfig.discovery && providerConfig.api) { + discoverableProviders.push({ + provider: providerName, + api: providerConfig.api as Api, + baseUrl: providerConfig.baseUrl, + headers: providerConfig.headers, + discovery: providerConfig.discovery, + }); + } + // Always store API key for fallback resolver if (providerConfig.apiKey) { this.#customProviderApiKeys.set(providerName, providerConfig.apiKey); @@ -400,7 +474,108 @@ export class ModelRegistry { } } - return { models: this.#parseModels(value), overrides, modelOverrides: allModelOverrides, found: true }; + return { + models: this.#parseModels(value), + overrides, + modelOverrides: allModelOverrides, + keylessProviders, + discoverableProviders, + found: true, + }; + } + + async #refreshRuntimeDiscoveries(): Promise { + if (this.#discoverableProviders.length === 0) return; + const discovered = await Promise.all( + this.#discoverableProviders.map(provider => this.#discoverProviderModels(provider)), + ); + const merged = this.#mergeCustomModels(this.#models, discovered.flat()); + this.#models = this.#applyModelOverrides(merged, this.#modelOverrides); + } + + async #discoverProviderModels(providerConfig: DiscoveryProviderConfig): Promise[]> { + switch (providerConfig.discovery.type) { + case "ollama": + return this.#discoverOllamaModels(providerConfig); + } + } + + async #discoverOllamaModels(providerConfig: DiscoveryProviderConfig): Promise[]> { + const endpoint = this.#normalizeOllamaBaseUrl(providerConfig.baseUrl); + const tagsUrl = `${endpoint}/api/tags`; + try { + const response = await fetch(tagsUrl, { + headers: { ...(providerConfig.headers ?? {}) }, + signal: AbortSignal.timeout(3000), + }); + if (!response.ok) { + logger.warn("model discovery failed for provider", { + provider: providerConfig.provider, + status: response.status, + url: tagsUrl, + }); + return []; + } + const payload = (await response.json()) as { models?: Array<{ name?: string; model?: string }> }; + const models = payload.models ?? []; + const discovered: Model[] = []; + for (const item of models) { + const id = item.model || item.name; + if (!id) continue; + discovered.push({ + id, + name: item.name || id, + api: providerConfig.api, + provider: providerConfig.provider, + baseUrl: `${endpoint}/v1`, + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128000, + maxTokens: 8192, + headers: providerConfig.headers, + }); + } + return this.#applyProviderModelOverrides(providerConfig.provider, discovered); + } catch (error) { + logger.warn("model discovery failed for provider", { + provider: providerConfig.provider, + url: tagsUrl, + error: error instanceof Error ? error.message : String(error), + }); + return []; + } + } + + #normalizeOllamaBaseUrl(baseUrl?: string): string { + const raw = baseUrl || "http://127.0.0.1:11434"; + try { + const parsed = new URL(raw); + return `${parsed.protocol}//${parsed.host}`; + } catch { + return "http://127.0.0.1:11434"; + } + } + + #applyProviderModelOverrides(provider: string, models: Model[]): Model[] { + const overrides = this.#modelOverrides.get(provider); + if (!overrides || overrides.size === 0) return models; + return models.map(model => { + const override = overrides.get(model.id); + if (!override) return model; + return applyModelOverride(model, override); + }); + } + + #applyModelOverrides(models: Model[], overrides: Map>): Model[] { + if (overrides.size === 0) return models; + return models.map(model => { + const providerOverrides = overrides.get(model.provider); + if (!providerOverrides) return model; + const override = providerOverrides.get(model.id); + if (!override) return model; + return applyModelOverride(model, override); + }); } #parseModels(config: ModelsConfig): Model[] { @@ -469,7 +644,7 @@ export class ModelRegistry { * This is a fast check that doesn't refresh OAuth tokens. */ getAvailable(): Model[] { - return this.#models.filter(m => this.authStorage.hasAuth(m.provider)); + return this.#models.filter(m => this.#keylessProviders.has(m.provider) || this.authStorage.hasAuth(m.provider)); } /** @@ -490,6 +665,9 @@ export class ModelRegistry { * Get API key for a model. */ async getApiKey(model: Model, sessionId?: string): Promise { + if (this.#keylessProviders.has(model.provider)) { + return ""; + } return this.authStorage.getApiKey(model.provider, sessionId, { baseUrl: model.baseUrl }); } @@ -497,6 +675,9 @@ export class ModelRegistry { * Get API key for a provider (e.g., "openai"). */ async getApiKeyForProvider(provider: string, sessionId?: string, baseUrl?: string): Promise { + if (this.#keylessProviders.has(provider)) { + return ""; + } return this.authStorage.getApiKey(provider, sessionId, { baseUrl }); } diff --git a/packages/coding-agent/src/ipy/runtime.ts b/packages/coding-agent/src/ipy/runtime.ts index 1a9d46671..ef9ccd25d 100644 --- a/packages/coding-agent/src/ipy/runtime.ts +++ b/packages/coding-agent/src/ipy/runtime.ts @@ -202,7 +202,8 @@ export function resolvePythonRuntime(cwd: string, baseEnv: Record { expect(googleModels.some(m => m.id === "gemini-custom")).toBe(true); }); - test("refresh() picks up baseUrl override changes", () => { + test("refresh() picks up baseUrl override changes", async () => { writeRawModelsJson({ anthropic: overrideConfig("https://first-proxy.example.com/v1"), }); @@ -173,7 +173,7 @@ describe("ModelRegistry", () => { writeRawModelsJson({ anthropic: overrideConfig("https://second-proxy.example.com/v1"), }); - registry.refresh(); + await registry.refresh(); expect(getModelsForProvider(registry, "anthropic")[0].baseUrl).toBe("https://second-proxy.example.com/v1"); }); @@ -270,7 +270,7 @@ describe("ModelRegistry", () => { ); }); - test("refresh() reloads merged custom models from disk", () => { + test("refresh() reloads merged custom models from disk", async () => { writeModelsJson({ anthropic: providerConfig("https://first-proxy.example.com/v1", [{ id: "claude-custom" }]), }); @@ -281,7 +281,7 @@ describe("ModelRegistry", () => { writeModelsJson({ anthropic: providerConfig("https://second-proxy.example.com/v1", [{ id: "claude-custom-2" }]), }); - registry.refresh(); + await registry.refresh(); const anthropicModels = getModelsForProvider(registry, "anthropic"); expect(anthropicModels.some(m => m.id === "claude-custom")).toBe(false); @@ -289,7 +289,7 @@ describe("ModelRegistry", () => { expect(anthropicModels.some(m => m.id.includes("claude"))).toBe(true); }); - test("removing custom models from models.json keeps built-in provider models", () => { + test("removing custom models from models.json keeps built-in provider models", async () => { writeModelsJson({ anthropic: providerConfig("https://proxy.example.com/v1", [{ id: "claude-custom" }]), }); @@ -299,7 +299,7 @@ describe("ModelRegistry", () => { // Remove custom models and refresh writeModelsJson({}); - registry.refresh(); + await registry.refresh(); const anthropicModels = getModelsForProvider(registry, "anthropic"); expect(anthropicModels.length).toBeGreaterThan(1); @@ -484,7 +484,7 @@ describe("ModelRegistry", () => { expect(sonnet?.headers?.["X-Custom-Model-Header"]).toBe("value"); }); - test("refresh() picks up model override changes", () => { + test("refresh() picks up model override changes", async () => { writeRawModelsJson({ openrouter: { modelOverrides: { @@ -510,14 +510,14 @@ describe("ModelRegistry", () => { }, }, }); - registry.refresh(); + await registry.refresh(); expect( getModelsForProvider(registry, "openrouter").find(m => m.id === "anthropic/claude-sonnet-4")?.name, ).toBe("Second Name"); }); - test("removing model override restores built-in values", () => { + test("removing model override restores built-in values", async () => { writeRawModelsJson({ openrouter: { modelOverrides: { @@ -536,7 +536,7 @@ describe("ModelRegistry", () => { // Remove override and refresh writeRawModelsJson({}); - registry.refresh(); + await registry.refresh(); const restoredName = getModelsForProvider(registry, "openrouter").find( m => m.id === "anthropic/claude-sonnet-4", @@ -544,4 +544,68 @@ describe("ModelRegistry", () => { expect(restoredName).not.toBe("Custom Name"); }); }); + + describe("runtime discovery", () => { + test("discovers ollama models at runtime and treats auth:none providers as available", async () => { + writeRawModelsJson({ + ollama: { + baseUrl: "http://127.0.0.1:11434/v1", + api: "openai-completions", + auth: "none", + discovery: { type: "ollama" }, + }, + }); + + const originalFetch = globalThis.fetch; + globalThis.fetch = (async (input: string | URL | Request) => { + expect(String(input)).toBe("http://127.0.0.1:11434/api/tags"); + return new Response( + JSON.stringify({ + models: [{ name: "qwen2.5-coder:7b" }, { model: "llama3.2:3b", name: "llama3.2:3b" }], + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); + }) as unknown as typeof fetch; + + try { + const registry = new ModelRegistry(authStorage, modelsJsonPath); + await registry.refresh(); + + const ollamaModels = getModelsForProvider(registry, "ollama"); + expect(ollamaModels.some(m => m.id === "qwen2.5-coder:7b")).toBe(true); + expect(ollamaModels.some(m => m.id === "llama3.2:3b")).toBe(true); + + const available = registry.getAvailable().filter(m => m.provider === "ollama"); + expect(available.length).toBe(2); + expect(await registry.getApiKey(available[0])).toBe(""); + } finally { + globalThis.fetch = originalFetch; + } + }); + + test("discovery failure does not fail model registry refresh", async () => { + writeRawModelsJson({ + ollama: { + baseUrl: "http://127.0.0.1:11434", + api: "openai-completions", + auth: "none", + discovery: { type: "ollama" }, + }, + }); + + const originalFetch = globalThis.fetch; + globalThis.fetch = (async () => { + throw new Error("connection refused"); + }) as unknown as typeof fetch; + + try { + const registry = new ModelRegistry(authStorage, modelsJsonPath); + await registry.refresh(); + expect(getModelsForProvider(registry, "ollama")).toHaveLength(0); + expect(registry.getError()).toBeUndefined(); + } finally { + globalThis.fetch = originalFetch; + } + }); + }); });