Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions src/agent.ts
Original file line number Diff line number Diff line change
Expand Up @@ -489,6 +489,10 @@ export class Agent<zO = unknown, zI = unknown, const Tools extends AnyTool[] = [
let currentModelIndex = 0;
let sameModelRetries = 0;

let modelCalls = 0;
const maxModelCalls = (1 + this.#retryStrategy.sameModelRetries) * (1 + this.#maxRecoveryAttempts) *
this.#models.length * this.#retryStrategy.modelCycles;

for (let cycle = 0; cycle < this.#retryStrategy.modelCycles; cycle++) {
currentModelIndex = 0;
sameModelRetries = 0;
Expand Down Expand Up @@ -523,10 +527,18 @@ export class Agent<zO = unknown, zI = unknown, const Tools extends AnyTool[] = [
trace: agentTrace.id,
};
sameModelRetries = 0;
// Otherwise the next pass on this model counts as another switch.
previousModel = null;
}

// attempt 0 = initial call, 1..N = recovery retries via handleModelError
for (let attempt = 0; attempt <= this.#maxRecoveryAttempts; attempt++) {
modelCalls++;
if (modelCalls > maxModelCalls) {
agentTrace.log(`Exceeded maximum model calls (${maxModelCalls}) in one turn`);
throw lastError ?? new Error(`Exceeded maximum model calls (${maxModelCalls}) in one turn`);
}

using modelTrace = newTrace({
type: "model",
parent: agentTrace,
Expand Down
43 changes: 43 additions & 0 deletions tests/simple/agent.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -651,3 +651,46 @@ Deno.test("a deterministic client error retires a model for the rest of the run"
assertEquals(primaryCalls, 1);
assertObjectMatch(run.history.at(-1)!, { type: "output_text", content: "search done" });
});

function serverErrorModel(name: string, calls: { count: number }): Adapter<unknown, unknown> {
return {
provider: name,
model: name,
stream() {
calls.count += 1;
throw new Error("500 Internal Server Error");
},
};
}

Deno.test("a failing fallback keeps its own retry budget and switches models once", async () => {
const primaryCalls = { count: 0 };
const fallbackCalls = { count: 0 };
const sameModelRetries = 2;
const maxRecoveryAttempts = 0;

const agent = new Agent({
model: [serverErrorModel("primary", primaryCalls), serverErrorModel("fallback", fallbackCalls)],
instructions: "You are a friendly assistant",
maxRecoveryAttempts,
retryStrategy: { modelCycles: 1, sameModelRetries },
});

const streamItems: StreamItem[] = [];
await assertRejects(
async () => {
for await (const item of agent.stream("Hello!")) {
streamItems.push(item);
}
},
Error,
"500 Internal Server Error",
);

assertEquals(primaryCalls.count, 1 + sameModelRetries);
assertEquals(fallbackCalls.count, 1 + sameModelRetries);
assertEquals(streamItems.filter((item) => item.type === "model_switched").length, 1);

const maxModelCalls = (1 + sameModelRetries) * (1 + maxRecoveryAttempts) * 2 * 1;
assertEquals(primaryCalls.count + fallbackCalls.count, maxModelCalls);
});
Loading