diff --git a/src/api/index.ts b/src/api/index.ts index c9992c9c39..f48ab50c0e 100644 --- a/src/api/index.ts +++ b/src/api/index.ts @@ -125,6 +125,13 @@ export interface ApiHandler { getModel(): { id: string; info: ModelInfo } + /** + * Ensures model metadata has been fetched from the remote API so that getModel() + * returns accurate info (context window, pricing, etc.) instead of hardcoded defaults. + * Only router providers that discover models over the network implement this. + */ + ensureModelFetched?(): Promise + /** * Optional context window for context-management / auto-condense when it must differ from * getModel().info.contextWindow. Only VS Code LM overrides it (static `maxInputTokens` vs its diff --git a/src/api/providers/__tests__/zoo-gateway.spec.ts b/src/api/providers/__tests__/zoo-gateway.spec.ts index e0c060db3b..e797dc9745 100644 --- a/src/api/providers/__tests__/zoo-gateway.spec.ts +++ b/src/api/providers/__tests__/zoo-gateway.spec.ts @@ -635,4 +635,92 @@ describe("ZooGatewayHandler", () => { ) }) }) + + describe("ensureModelFetched", () => { + it("fetches models when instance models are empty", async () => { + const handler = new ZooGatewayHandler(mockOptions) + const { getModels } = await import("../fetchers/modelCache") + + expect(handler.getModel().info.contextWindow).toBe(200000) + + await handler.ensureModelFetched() + + expect(getModels).toHaveBeenCalled() + }) + + it("skips the fetch when models are already populated", async () => { + const handler = new ZooGatewayHandler(mockOptions) + const { getModels } = await import("../fetchers/modelCache") + + await handler.ensureModelFetched() + vitest.mocked(getModels).mockClear() + + await handler.ensureModelFetched() + expect(getModels).not.toHaveBeenCalled() + }) + + it("short-circuits a subsequent fetchModel call after models are populated", async () => { + const handler = new ZooGatewayHandler(mockOptions) + const { getModels } = await import("../fetchers/modelCache") + + await handler.ensureModelFetched() + vitest.mocked(getModels).mockClear() + + await handler.fetchModel() + expect(getModels).not.toHaveBeenCalled() + }) + + it("deduplicates concurrent calls into a single fetch", async () => { + const handler = new ZooGatewayHandler(mockOptions) + const { getModels } = await import("../fetchers/modelCache") + vitest.mocked(getModels).mockClear() + + await Promise.all([handler.ensureModelFetched(), handler.ensureModelFetched()]) + + expect(getModels).toHaveBeenCalledTimes(1) + }) + + it("recovers after a rejected fetch so later calls are not poisoned", async () => { + const handler = new ZooGatewayHandler(mockOptions) + const { getModels } = await import("../fetchers/modelCache") + + vitest.mocked(getModels).mockRejectedValueOnce(new Error("network down")) + await expect(handler.ensureModelFetched()).rejects.toThrow("network down") + + vitest.mocked(getModels).mockResolvedValueOnce({ + "anthropic/claude-sonnet-4": { + maxTokens: 64000, + contextWindow: 1000000, + supportsImages: true, + supportsPromptCache: true, + }, + }) + await handler.ensureModelFetched() + + expect(handler.getModel().info.contextWindow).toBe(1000000) + }) + + it("makes getModel return the fetched context window instead of the default", async () => { + const { getModels } = await import("../fetchers/modelCache") + vitest.mocked(getModels).mockResolvedValueOnce({ + "google/gemini-2.5-pro": { + maxTokens: 65536, + contextWindow: 1048576, + supportsImages: true, + supportsPromptCache: false, + }, + }) + + const handler = new ZooGatewayHandler({ + ...mockOptions, + zooGatewayModelId: "google/gemini-2.5-pro", + }) + + expect(handler.getModel().info.contextWindow).toBe(200000) + + await handler.ensureModelFetched() + + expect(handler.getModel().info.contextWindow).toBe(1048576) + }) + }) }) diff --git a/src/api/providers/router-provider.ts b/src/api/providers/router-provider.ts index 7a983ed12e..cbdd49e58b 100644 --- a/src/api/providers/router-provider.ts +++ b/src/api/providers/router-provider.ts @@ -56,9 +56,33 @@ export abstract class RouterProvider extends BaseProvider { }) } + private modelFetchPromise?: Promise<{ id: string; info: ModelInfo }> + public async fetchModel() { - this.models = await getModels({ provider: this.name, apiKey: this.client.apiKey, baseUrl: this.client.baseURL }) - return this.getModel() + if (Object.keys(this.models).length > 0) { + return this.getModel() + } + + if (!this.modelFetchPromise) { + this.modelFetchPromise = getModels({ + provider: this.name, + apiKey: this.client.apiKey, + baseUrl: this.client.baseURL, + }) + .then((models) => { + this.models = models + return this.getModel() + }) + .finally(() => { + this.modelFetchPromise = undefined + }) + } + + return this.modelFetchPromise + } + + async ensureModelFetched(): Promise { + await this.fetchModel() } override getModel(): { id: string; info: ModelInfo } { diff --git a/src/core/task/Task.ts b/src/core/task/Task.ts index 8f69a3a0d4..4ba2996c91 100644 --- a/src/core/task/Task.ts +++ b/src/core/task/Task.ts @@ -2762,6 +2762,8 @@ export class Task extends EventEmitter implements TaskLike { await this.diffViewProvider.reset() + await this.safeEnsureModelFetched() + // Cache model info once per API request to avoid repeated calls during streaming // This is especially important for tools and background usage collection this.cachedStreamingModel = this.api.getModel() @@ -3837,11 +3839,28 @@ export class Task extends EventEmitter implements TaskLike { ) } + /** + * Ensures router-provider model metadata is loaded before getModel() is used for + * context management or streaming. Failures fall back to hardcoded defaults rather + * than aborting the task. + */ + private async safeEnsureModelFetched(): Promise { + try { + await this.api.ensureModelFetched?.() + } catch (error) { + console.error( + `[Task#${this.taskId}] Failed to fetch model metadata:`, + error instanceof Error ? error.message : error, + ) + } + } + private async handleContextWindowExceededError(): Promise { const state = await this.providerRef.deref()?.getState() const { profileThresholds = {}, mode, apiConfiguration } = state ?? {} const { contextTokens } = this.getTokenUsage() + await this.safeEnsureModelFetched() const modelInfo = this.api.getModel().info const maxTokens = getModelMaxOutputTokens({ @@ -4042,6 +4061,7 @@ export class Task extends EventEmitter implements TaskLike { const { contextTokens } = this.getTokenUsage() if (contextTokens) { + await this.safeEnsureModelFetched() const modelInfo = this.api.getModel().info const maxTokens = getModelMaxOutputTokens({ diff --git a/src/core/task/__tests__/Task.spec.ts b/src/core/task/__tests__/Task.spec.ts index d9f7240c5c..330241e221 100644 --- a/src/core/task/__tests__/Task.spec.ts +++ b/src/core/task/__tests__/Task.spec.ts @@ -32,6 +32,8 @@ type TaskTestAccess = { presentAssistantMessageSafe: () => void updateClineMessage: (message: import("@roo-code/types").ClineMessage) => Promise saveClineMessages: () => Promise + safeEnsureModelFetched: () => Promise + addToApiConversationHistory: (message: unknown, reasoning?: string) => Promise } function getTaskTestAccess(task: Task): TaskTestAccess { @@ -185,6 +187,14 @@ vi.mock("../../environment/getEnvironmentDetails", () => ({ getEnvironmentDetails: vi.fn().mockResolvedValue(""), })) +vi.mock("../../mentions/processUserContentMentions", async (importOriginal) => { + const actual = await importOriginal() + return { + ...actual, + processUserContentMentions: vi.fn().mockImplementation(actual.processUserContentMentions), + } +}) + vi.mock("../../ignore/RooIgnoreController") vi.mock("../../../i18n", () => { @@ -2507,6 +2517,230 @@ describe("Cline", () => { }) }) + describe("safeEnsureModelFetched", () => { + it("loads model metadata before getModel is used", async () => { + const task = new Task({ + provider: mockProvider, + apiConfiguration: mockApiConfig, + task: "test task", + startTask: false, + }) + + const ensureModelFetched = vi.fn().mockResolvedValue(undefined) + Object.assign(task.api, { ensureModelFetched }) + + await getTaskTestAccess(task).safeEnsureModelFetched() + + expect(ensureModelFetched).toHaveBeenCalledTimes(1) + }) + + it("swallows fetch failures so callers can fall back to defaults", async () => { + const task = new Task({ + provider: mockProvider, + apiConfiguration: mockApiConfig, + task: "test task", + startTask: false, + }) + + const ensureModelFetched = vi.fn().mockRejectedValue(new Error("network down")) + Object.assign(task.api, { ensureModelFetched }) + const errorSpy = vi.spyOn(console, "error").mockImplementation(() => {}) + + await expect(getTaskTestAccess(task).safeEnsureModelFetched()).resolves.toBeUndefined() + + expect(errorSpy).toHaveBeenCalledWith( + expect.stringContaining("Failed to fetch model metadata"), + "network down", + ) + errorSpy.mockRestore() + }) + + it("is a no-op when the api handler does not implement ensureModelFetched", async () => { + const task = new Task({ + provider: mockProvider, + apiConfiguration: mockApiConfig, + task: "test task", + startTask: false, + }) + + await expect(getTaskTestAccess(task).safeEnsureModelFetched()).resolves.toBeUndefined() + }) + + it("calls safeEnsureModelFetched from attemptApiRequest when context tokens are present", async () => { + const task = new Task({ + provider: mockProvider, + apiConfiguration: mockApiConfig, + task: "test task", + startTask: false, + }) + + vi.spyOn(getTaskTestAccess(task), "getSystemPrompt").mockResolvedValue("mock system prompt") + vi.spyOn(task, "getTokenUsage").mockReturnValue({ + totalCost: 0, + totalTokensIn: 0, + totalTokensOut: 0, + contextTokens: 50_000, + }) + const safeSpy = vi.spyOn(getTaskTestAccess(task), "safeEnsureModelFetched").mockResolvedValue(undefined) + vi.spyOn(task.api, "getModel").mockReturnValue({ + id: mockApiConfig.apiModelId!, + info: { + supportsImages: false, + supportsPromptCache: true, + contextWindow: 200_000, + maxTokens: 4096, + } as ModelInfo, + }) + vi.spyOn(task.api, "createMessage").mockReturnValue({ + async *[Symbol.asyncIterator]() { + yield { type: "text", text: "ok" } + }, + async next() { + return { done: true, value: undefined } + }, + async return() { + return { done: true, value: undefined } + }, + async throw(error: unknown) { + throw error + }, + async [Symbol.asyncDispose]() {}, + } as AsyncGenerator) + + task.apiConversationHistory = [ + { + role: "user" as const, + content: [{ type: "text" as const, text: "test message" }], + ts: Date.now(), + }, + ] + + const iterator = task.attemptApiRequest(0) + await iterator.next() + + expect(safeSpy).toHaveBeenCalled() + }) + + it("continues attemptApiRequest when model metadata fetch fails", async () => { + const task = new Task({ + provider: mockProvider, + apiConfiguration: mockApiConfig, + task: "test task", + startTask: false, + }) + + vi.spyOn(getTaskTestAccess(task), "getSystemPrompt").mockResolvedValue("mock system prompt") + vi.spyOn(task, "getTokenUsage").mockReturnValue({ + totalCost: 0, + totalTokensIn: 0, + totalTokensOut: 0, + contextTokens: 50_000, + }) + const ensureModelFetched = vi.fn().mockRejectedValue(new Error("fetch failed")) + Object.assign(task.api, { ensureModelFetched }) + vi.spyOn(task.api, "getModel").mockReturnValue({ + id: mockApiConfig.apiModelId!, + info: { + supportsImages: false, + supportsPromptCache: true, + contextWindow: 200_000, + maxTokens: 4096, + } as ModelInfo, + }) + vi.spyOn(task.api, "createMessage").mockReturnValue({ + async *[Symbol.asyncIterator]() { + yield { type: "text", text: "ok" } + }, + async next() { + return { done: false, value: { type: "text", text: "ok" } } + }, + async return() { + return { done: true, value: undefined } + }, + async throw(error: unknown) { + throw error + }, + async [Symbol.asyncDispose]() {}, + } as AsyncGenerator) + const errorSpy = vi.spyOn(console, "error").mockImplementation(() => {}) + + task.apiConversationHistory = [ + { + role: "user" as const, + content: [{ type: "text" as const, text: "test message" }], + ts: Date.now(), + }, + ] + + const iterator = task.attemptApiRequest(0) + await expect(iterator.next()).resolves.toMatchObject({ + done: false, + value: { type: "text", text: "ok" }, + }) + expect(errorSpy).toHaveBeenCalled() + errorSpy.mockRestore() + }) + + it("fetches model metadata before caching the streaming model", async () => { + const task = new Task({ + provider: mockProvider, + apiConfiguration: mockApiConfig, + task: "test task", + startTask: false, + }) + + const ensureModelFetched = vi.fn().mockResolvedValue(undefined) + Object.assign(task.api, { ensureModelFetched }) + vi.spyOn(task.api, "getModel").mockReturnValue({ + id: mockApiConfig.apiModelId!, + info: { + supportsImages: false, + supportsPromptCache: true, + contextWindow: 200_000, + maxTokens: 4096, + } as ModelInfo, + }) + vi.mocked(processUserContentMentions).mockResolvedValueOnce({ + content: [{ type: "text", text: "hello" }], + mode: undefined, + }) + const safeSpy = vi.spyOn(getTaskTestAccess(task), "safeEnsureModelFetched") + vi.spyOn(task, "attemptApiRequest").mockImplementation(() => { + throw new Error("stop after model metadata fetch") + }) + vi.spyOn(getTaskTestAccess(task), "saveClineMessages").mockResolvedValue(true) + vi.spyOn(task.diffViewProvider, "reset").mockResolvedValue(undefined as never) + vi.spyOn(getTaskTestAccess(task), "addToApiConversationHistory").mockResolvedValue(undefined) + + task.clineMessages = [ + { + ts: Date.now(), + type: "say", + say: "api_req_started", + text: "{}", + }, + ] + vi.spyOn(task, "say").mockImplementation(async (type) => { + if (type === "api_req_started") { + task.clineMessages.push({ + ts: Date.now(), + type: "say", + say: "api_req_started", + text: "{}", + }) + } + return undefined as never + }) + + const result = await task.recursivelyMakeClineRequests([{ type: "text", text: "hello" }], false) + + expect(result).toBe(true) + expect(safeSpy).toHaveBeenCalled() + expect(ensureModelFetched).toHaveBeenCalled() + expect(task.cachedStreamingModel?.id).toBe(mockApiConfig.apiModelId) + }) + }) + describe("start()", () => { it("should be a no-op if the task was already started in the constructor", () => { const task = new Task({