Skip to content

Commit 885d8e0

Browse files
fix(router-provider): address review feedback on fetch failures and double-fetch
Make fetchModel single-flight and short-circuit once models are loaded so auth-scoped providers do not hit the models endpoint twice per request. Catch ensureModelFetched failures in Task via safeEnsureModelFetched so a metadata fetch error falls back to defaults instead of ending the task. Add reject-then-recover coverage and Task tests for the new call sites. Co-authored-by: Cursor <cursoragent@cursor.com>
1 parent 47e0732 commit 885d8e0

4 files changed

Lines changed: 306 additions & 17 deletions

File tree

src/api/providers/__tests__/zoo-gateway.spec.ts

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -659,6 +659,17 @@ describe("ZooGatewayHandler", () => {
659659
expect(getModels).not.toHaveBeenCalled()
660660
})
661661

662+
it("short-circuits a subsequent fetchModel call after models are populated", async () => {
663+
const handler = new ZooGatewayHandler(mockOptions)
664+
const { getModels } = await import("../fetchers/modelCache")
665+
666+
await handler.ensureModelFetched()
667+
vitest.mocked(getModels).mockClear()
668+
669+
await handler.fetchModel()
670+
expect(getModels).not.toHaveBeenCalled()
671+
})
672+
662673
it("deduplicates concurrent calls into a single fetch", async () => {
663674
const handler = new ZooGatewayHandler(mockOptions)
664675
const { getModels } = await import("../fetchers/modelCache")
@@ -669,6 +680,26 @@ describe("ZooGatewayHandler", () => {
669680
expect(getModels).toHaveBeenCalledTimes(1)
670681
})
671682

683+
it("recovers after a rejected fetch so later calls are not poisoned", async () => {
684+
const handler = new ZooGatewayHandler(mockOptions)
685+
const { getModels } = await import("../fetchers/modelCache")
686+
687+
vitest.mocked(getModels).mockRejectedValueOnce(new Error("network down"))
688+
await expect(handler.ensureModelFetched()).rejects.toThrow("network down")
689+
690+
vitest.mocked(getModels).mockResolvedValueOnce({
691+
"anthropic/claude-sonnet-4": {
692+
maxTokens: 64000,
693+
contextWindow: 1000000,
694+
supportsImages: true,
695+
supportsPromptCache: true,
696+
},
697+
})
698+
await handler.ensureModelFetched()
699+
700+
expect(handler.getModel().info.contextWindow).toBe(1000000)
701+
})
702+
672703
it("makes getModel return the fetched context window instead of the default", async () => {
673704
const { getModels } = await import("../fetchers/modelCache")
674705
vitest.mocked(getModels).mockResolvedValueOnce({

src/api/providers/router-provider.ts

Lines changed: 23 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -56,24 +56,33 @@ export abstract class RouterProvider extends BaseProvider {
5656
})
5757
}
5858

59-
public async fetchModel() {
60-
this.models = await getModels({ provider: this.name, apiKey: this.client.apiKey, baseUrl: this.client.baseURL })
61-
return this.getModel()
62-
}
59+
private modelFetchPromise?: Promise<{ id: string; info: ModelInfo }>
6360

64-
private modelFetchPromise?: Promise<void>
61+
public async fetchModel() {
62+
if (Object.keys(this.models).length > 0) {
63+
return this.getModel()
64+
}
6565

66-
async ensureModelFetched(): Promise<void> {
67-
if (Object.keys(this.models).length === 0) {
68-
const fetchPromise = (this.modelFetchPromise ??= this.fetchModel().then(() => undefined))
69-
try {
70-
await fetchPromise
71-
} finally {
72-
if (this.modelFetchPromise === fetchPromise) {
66+
if (!this.modelFetchPromise) {
67+
this.modelFetchPromise = getModels({
68+
provider: this.name,
69+
apiKey: this.client.apiKey,
70+
baseUrl: this.client.baseURL,
71+
})
72+
.then((models) => {
73+
this.models = models
74+
return this.getModel()
75+
})
76+
.finally(() => {
7377
this.modelFetchPromise = undefined
74-
}
75-
}
78+
})
7679
}
80+
81+
return this.modelFetchPromise
82+
}
83+
84+
async ensureModelFetched(): Promise<void> {
85+
await this.fetchModel()
7786
}
7887

7988
override getModel(): { id: string; info: ModelInfo } {

src/core/task/Task.ts

Lines changed: 19 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2762,7 +2762,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
27622762

27632763
await this.diffViewProvider.reset()
27642764

2765-
await this.api.ensureModelFetched?.()
2765+
await this.safeEnsureModelFetched()
27662766

27672767
// Cache model info once per API request to avoid repeated calls during streaming
27682768
// This is especially important for tools and background usage collection
@@ -3839,12 +3839,28 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
38393839
)
38403840
}
38413841

3842+
/**
3843+
* Ensures router-provider model metadata is loaded before getModel() is used for
3844+
* context management or streaming. Failures fall back to hardcoded defaults rather
3845+
* than aborting the task.
3846+
*/
3847+
private async safeEnsureModelFetched(): Promise<void> {
3848+
try {
3849+
await this.api.ensureModelFetched?.()
3850+
} catch (error) {
3851+
console.error(
3852+
`[Task#${this.taskId}] Failed to fetch model metadata:`,
3853+
error instanceof Error ? error.message : error,
3854+
)
3855+
}
3856+
}
3857+
38423858
private async handleContextWindowExceededError(): Promise<void> {
38433859
const state = await this.providerRef.deref()?.getState()
38443860
const { profileThresholds = {}, mode, apiConfiguration } = state ?? {}
38453861

38463862
const { contextTokens } = this.getTokenUsage()
3847-
await this.api.ensureModelFetched?.()
3863+
await this.safeEnsureModelFetched()
38483864
const modelInfo = this.api.getModel().info
38493865

38503866
const maxTokens = getModelMaxOutputTokens({
@@ -4045,7 +4061,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
40454061
const { contextTokens } = this.getTokenUsage()
40464062

40474063
if (contextTokens) {
4048-
await this.api.ensureModelFetched?.()
4064+
await this.safeEnsureModelFetched()
40494065
const modelInfo = this.api.getModel().info
40504066

40514067
const maxTokens = getModelMaxOutputTokens({

src/core/task/__tests__/Task.spec.ts

Lines changed: 233 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ type TaskTestAccess = {
3232
presentAssistantMessageSafe: () => void
3333
updateClineMessage: (message: import("@roo-code/types").ClineMessage) => Promise<void>
3434
saveClineMessages: () => Promise<boolean>
35+
safeEnsureModelFetched: () => Promise<void>
3536
}
3637

3738
function getTaskTestAccess(task: Task): TaskTestAccess {
@@ -185,6 +186,14 @@ vi.mock("../../environment/getEnvironmentDetails", () => ({
185186
getEnvironmentDetails: vi.fn().mockResolvedValue(""),
186187
}))
187188

189+
vi.mock("../../mentions/processUserContentMentions", async (importOriginal) => {
190+
const actual = await importOriginal<typeof import("../../mentions/processUserContentMentions")>()
191+
return {
192+
...actual,
193+
processUserContentMentions: vi.fn().mockImplementation(actual.processUserContentMentions),
194+
}
195+
})
196+
188197
vi.mock("../../ignore/RooIgnoreController")
189198

190199
vi.mock("../../../i18n", () => {
@@ -2507,6 +2516,230 @@ describe("Cline", () => {
25072516
})
25082517
})
25092518

2519+
describe("safeEnsureModelFetched", () => {
2520+
it("loads model metadata before getModel is used", async () => {
2521+
const task = new Task({
2522+
provider: mockProvider,
2523+
apiConfiguration: mockApiConfig,
2524+
task: "test task",
2525+
startTask: false,
2526+
})
2527+
2528+
const ensureModelFetched = vi.fn().mockResolvedValue(undefined)
2529+
Object.assign(task.api, { ensureModelFetched })
2530+
2531+
await getTaskTestAccess(task).safeEnsureModelFetched()
2532+
2533+
expect(ensureModelFetched).toHaveBeenCalledTimes(1)
2534+
})
2535+
2536+
it("swallows fetch failures so callers can fall back to defaults", async () => {
2537+
const task = new Task({
2538+
provider: mockProvider,
2539+
apiConfiguration: mockApiConfig,
2540+
task: "test task",
2541+
startTask: false,
2542+
})
2543+
2544+
const ensureModelFetched = vi.fn().mockRejectedValue(new Error("network down"))
2545+
Object.assign(task.api, { ensureModelFetched })
2546+
const errorSpy = vi.spyOn(console, "error").mockImplementation(() => {})
2547+
2548+
await expect(getTaskTestAccess(task).safeEnsureModelFetched()).resolves.toBeUndefined()
2549+
2550+
expect(errorSpy).toHaveBeenCalledWith(
2551+
expect.stringContaining("Failed to fetch model metadata"),
2552+
"network down",
2553+
)
2554+
errorSpy.mockRestore()
2555+
})
2556+
2557+
it("is a no-op when the api handler does not implement ensureModelFetched", async () => {
2558+
const task = new Task({
2559+
provider: mockProvider,
2560+
apiConfiguration: mockApiConfig,
2561+
task: "test task",
2562+
startTask: false,
2563+
})
2564+
2565+
await expect(getTaskTestAccess(task).safeEnsureModelFetched()).resolves.toBeUndefined()
2566+
})
2567+
2568+
it("calls safeEnsureModelFetched from attemptApiRequest when context tokens are present", async () => {
2569+
const task = new Task({
2570+
provider: mockProvider,
2571+
apiConfiguration: mockApiConfig,
2572+
task: "test task",
2573+
startTask: false,
2574+
})
2575+
2576+
vi.spyOn(getTaskTestAccess(task), "getSystemPrompt").mockResolvedValue("mock system prompt")
2577+
vi.spyOn(task, "getTokenUsage").mockReturnValue({
2578+
totalCost: 0,
2579+
totalTokensIn: 0,
2580+
totalTokensOut: 0,
2581+
contextTokens: 50_000,
2582+
})
2583+
const safeSpy = vi.spyOn(getTaskTestAccess(task), "safeEnsureModelFetched").mockResolvedValue(undefined)
2584+
vi.spyOn(task.api, "getModel").mockReturnValue({
2585+
id: mockApiConfig.apiModelId!,
2586+
info: {
2587+
supportsImages: false,
2588+
supportsPromptCache: true,
2589+
contextWindow: 200_000,
2590+
maxTokens: 4096,
2591+
} as ModelInfo,
2592+
})
2593+
vi.spyOn(task.api, "createMessage").mockReturnValue({
2594+
async *[Symbol.asyncIterator]() {
2595+
yield { type: "text", text: "ok" }
2596+
},
2597+
async next() {
2598+
return { done: true, value: undefined }
2599+
},
2600+
async return() {
2601+
return { done: true, value: undefined }
2602+
},
2603+
async throw(error: unknown) {
2604+
throw error
2605+
},
2606+
async [Symbol.asyncDispose]() {},
2607+
} as AsyncGenerator<ApiStreamChunk>)
2608+
2609+
task.apiConversationHistory = [
2610+
{
2611+
role: "user" as const,
2612+
content: [{ type: "text" as const, text: "test message" }],
2613+
ts: Date.now(),
2614+
},
2615+
]
2616+
2617+
const iterator = task.attemptApiRequest(0)
2618+
await iterator.next()
2619+
2620+
expect(safeSpy).toHaveBeenCalled()
2621+
})
2622+
2623+
it("continues attemptApiRequest when model metadata fetch fails", async () => {
2624+
const task = new Task({
2625+
provider: mockProvider,
2626+
apiConfiguration: mockApiConfig,
2627+
task: "test task",
2628+
startTask: false,
2629+
})
2630+
2631+
vi.spyOn(getTaskTestAccess(task), "getSystemPrompt").mockResolvedValue("mock system prompt")
2632+
vi.spyOn(task, "getTokenUsage").mockReturnValue({
2633+
totalCost: 0,
2634+
totalTokensIn: 0,
2635+
totalTokensOut: 0,
2636+
contextTokens: 50_000,
2637+
})
2638+
const ensureModelFetched = vi.fn().mockRejectedValue(new Error("fetch failed"))
2639+
Object.assign(task.api, { ensureModelFetched })
2640+
vi.spyOn(task.api, "getModel").mockReturnValue({
2641+
id: mockApiConfig.apiModelId!,
2642+
info: {
2643+
supportsImages: false,
2644+
supportsPromptCache: true,
2645+
contextWindow: 200_000,
2646+
maxTokens: 4096,
2647+
} as ModelInfo,
2648+
})
2649+
vi.spyOn(task.api, "createMessage").mockReturnValue({
2650+
async *[Symbol.asyncIterator]() {
2651+
yield { type: "text", text: "ok" }
2652+
},
2653+
async next() {
2654+
return { done: false, value: { type: "text", text: "ok" } }
2655+
},
2656+
async return() {
2657+
return { done: true, value: undefined }
2658+
},
2659+
async throw(error: unknown) {
2660+
throw error
2661+
},
2662+
async [Symbol.asyncDispose]() {},
2663+
} as AsyncGenerator<ApiStreamChunk>)
2664+
const errorSpy = vi.spyOn(console, "error").mockImplementation(() => {})
2665+
2666+
task.apiConversationHistory = [
2667+
{
2668+
role: "user" as const,
2669+
content: [{ type: "text" as const, text: "test message" }],
2670+
ts: Date.now(),
2671+
},
2672+
]
2673+
2674+
const iterator = task.attemptApiRequest(0)
2675+
await expect(iterator.next()).resolves.toMatchObject({
2676+
done: false,
2677+
value: { type: "text", text: "ok" },
2678+
})
2679+
expect(errorSpy).toHaveBeenCalled()
2680+
errorSpy.mockRestore()
2681+
})
2682+
2683+
it("fetches model metadata before caching the streaming model", async () => {
2684+
const task = new Task({
2685+
provider: mockProvider,
2686+
apiConfiguration: mockApiConfig,
2687+
task: "test task",
2688+
startTask: false,
2689+
})
2690+
2691+
const ensureModelFetched = vi.fn().mockResolvedValue(undefined)
2692+
Object.assign(task.api, { ensureModelFetched })
2693+
vi.spyOn(task.api, "getModel").mockReturnValue({
2694+
id: mockApiConfig.apiModelId!,
2695+
info: {
2696+
supportsImages: false,
2697+
supportsPromptCache: true,
2698+
contextWindow: 200_000,
2699+
maxTokens: 4096,
2700+
} as ModelInfo,
2701+
})
2702+
vi.mocked(processUserContentMentions).mockResolvedValueOnce({
2703+
content: [{ type: "text", text: "hello" }],
2704+
mode: undefined,
2705+
})
2706+
const safeSpy = vi.spyOn(getTaskTestAccess(task), "safeEnsureModelFetched")
2707+
vi.spyOn(task, "attemptApiRequest").mockImplementation(() => {
2708+
throw new Error("stop after model metadata fetch")
2709+
})
2710+
vi.spyOn(getTaskTestAccess(task), "saveClineMessages").mockResolvedValue(true)
2711+
vi.spyOn(task.diffViewProvider, "reset").mockResolvedValue(undefined as never)
2712+
vi.spyOn(task, "addToApiConversationHistory").mockResolvedValue(undefined as never)
2713+
2714+
task.clineMessages = [
2715+
{
2716+
ts: Date.now(),
2717+
type: "say",
2718+
say: "api_req_started",
2719+
text: "{}",
2720+
},
2721+
]
2722+
vi.spyOn(task, "say").mockImplementation(async (type) => {
2723+
if (type === "api_req_started") {
2724+
task.clineMessages.push({
2725+
ts: Date.now(),
2726+
type: "say",
2727+
say: "api_req_started",
2728+
text: "{}",
2729+
})
2730+
}
2731+
return undefined as never
2732+
})
2733+
2734+
const result = await task.recursivelyMakeClineRequests([{ type: "text", text: "hello" }], false)
2735+
2736+
expect(result).toBe(true)
2737+
expect(safeSpy).toHaveBeenCalled()
2738+
expect(ensureModelFetched).toHaveBeenCalled()
2739+
expect(task.cachedStreamingModel?.id).toBe(mockApiConfig.apiModelId)
2740+
})
2741+
})
2742+
25102743
describe("start()", () => {
25112744
it("should be a no-op if the task was already started in the constructor", () => {
25122745
const task = new Task({

0 commit comments

Comments
 (0)