Skip to content
Open
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
129 changes: 129 additions & 0 deletions packages/gsd-agent-core/src/agent-session.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ import { describe, test } from "node:test";
import assert from "node:assert/strict";
import { parseSkillBlock } from "./agent-session.ts";
import { AgentSessionExtensionsModule } from "./session/agent-session-extensions.ts";
import { AgentSessionPromptModule } from "./session/agent-session-prompt.ts";
import { FallbackResolver } from "./fallback-resolver.ts";

describe("parseSkillBlock", () => {
test("parses a valid skill block with trailing user message", () => {
Expand Down Expand Up @@ -53,6 +55,123 @@ describe("AgentSessionExtensionsModule", () => {
});
});

describe("AgentSessionPromptModule", () => {
test("switches to a configured fallback model for retryable 502 errors", async () => {
const primary = makeModel("primary-provider", "primary-model");
const fallback = makeModel("fallback-provider", "fallback-model");
const events: unknown[] = [];
let selectedModel: unknown;
let fallbackErrorType: string | undefined;

const host = {
model: primary,
settingsManager: {
getRetrySettings: () => ({ enabled: true, maxRetries: 0, baseDelayMs: 1 }),
},
fallbackResolver: {
async findFallback(currentModel: unknown, errorType: string) {
fallbackErrorType = errorType;
assert.equal(currentModel, primary);
return {
model: fallback,
chainName: "default",
reason: "falling back to fallback-provider/fallback-model",
};
},
},
agent: {
state: {
model: primary,
messages: [
{
role: "assistant",
stopReason: "error",
errorMessage: "502 status code (no body)",
},
],
},
},
_lastAssistantMessage: {
role: "assistant",
stopReason: "error",
errorMessage: "502 status code (no body)",
},
_retryAttempt: 0,
emit: (event: unknown) => events.push(event),
setModel: async (model: unknown) => {
selectedModel = model;
host.agent.state.model = model as typeof primary;
},
checkCompaction: async () => false,
};

const shouldContinue = await new AgentSessionPromptModule(host as any).handlePostAgentRun();

assert.equal(shouldContinue, true);
assert.equal(selectedModel, fallback);
assert.equal(host.agent.state.model, fallback);
assert.equal(fallbackErrorType, "retryable");
assert.deepEqual(host.agent.state.messages, []);
assert.deepEqual(events, [
{
type: "fallback_provider_switch",
fromProvider: "primary-provider",
toProvider: "fallback-provider",
model: fallback,
reason: "falling back to fallback-provider/fallback-model",
},
]);
});
});

describe("FallbackResolver", () => {
test("allows fallback to a backup model on the same provider", async () => {
const primary = makeModel("shared-provider", "primary-model");
const backup = makeModel("shared-provider", "backup-model");
let exhaustedProvider: string | undefined;

const resolver = new FallbackResolver(
{
getFallbackSettings: () => ({
enabled: true,
chains: {
sameProvider: [
{ provider: "shared-provider", model: "primary-model" },
{ provider: "shared-provider", model: "backup-model" },
],
},
}),
} as any,
{
markProviderExhausted(provider: string) {
exhaustedProvider = provider;
},
isProviderAvailable() {
return exhaustedProvider === undefined;
},
} as any,
{
find(provider: string, model: string) {
if (provider === "shared-provider" && model === "primary-model") return primary;
if (provider === "shared-provider" && model === "backup-model") return backup;
return undefined;
},
isProviderRequestReady() {
return true;
},
} as any,
);

const result = await resolver.findFallback(primary, "retryable");

assert.deepEqual(result, {
model: backup,
chainName: "sameProvider",
reason: "falling back to shared-provider/backup-model",
});
});
});

function makeSkill(name: string) {
return {
name,
Expand All @@ -64,3 +183,13 @@ function makeSkill(name: string) {
disableModelInvocation: false,
};
}

function makeModel(provider: string, id: string) {
return {
api: "openai-responses",
provider,
id,
name: id,
contextWindow: 128000,
};
}
7 changes: 7 additions & 0 deletions packages/gsd-agent-core/src/agent-session.ts
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,7 @@ import {
type TurnLatencyVisibleKind,
} from "./turn-latency.js";
import type { AgentSessionHost } from "./session/agent-session-host.js";
import type { FallbackResolver } from "./fallback-resolver.js";
import { AgentSessionEventsModule } from "./session/agent-session-events.js";
import { AgentSessionPromptModule } from "./session/agent-session-prompt.js";
import { AgentSessionModelModule } from "./session/agent-session-model.js";
Expand Down Expand Up @@ -148,6 +149,7 @@ export class AgentSession implements AgentSessionHost {
_baseSystemPromptOptions!: BuildSystemPromptOptions;
_lastAssistantMessage: AssistantMessage | undefined = undefined;
_activeTurnLatency: TurnLatencyRecord | undefined = undefined;
_fallbackResolver: FallbackResolver | undefined = undefined;

private readonly _events = new AgentSessionEventsModule(this);
private readonly _prompt = new AgentSessionPromptModule(this);
Expand All @@ -171,6 +173,7 @@ export class AgentSession implements AgentSessionHost {
this._allowedToolNames = config.allowedToolNames ? new Set(config.allowedToolNames) : undefined;
this._baseToolsOverride = config.baseToolsOverride;
this._sessionStartEvent = config.sessionStartEvent ?? { type: "session_start", reason: "startup" };
this._fallbackResolver = config.fallbackResolver;

this._unsubscribeAgent = this.agent.subscribe(this._events.handleAgentEvent);
this._extensions.installAgentToolHooks();
Expand Down Expand Up @@ -297,6 +300,10 @@ export class AgentSession implements AgentSessionHost {
return this._extensionRunner;
}

get fallbackResolver(): FallbackResolver | undefined {
return this._fallbackResolver;
}

// AgentSessionHost cross-module surface
emit(event: AgentSessionEvent): void {
this._events.emit(event);
Expand Down
22 changes: 19 additions & 3 deletions packages/gsd-agent-core/src/fallback-resolver.ts
Original file line number Diff line number Diff line change
Expand Up @@ -53,11 +53,23 @@ export class FallbackResolver {
if (currentIndex === -1) continue;

// Try entries after the current one (already sorted by priority)
const result = await this._findAvailableInChain(chainName, entries, currentIndex + 1);
const result = await this._findAvailableInChain(
chainName,
entries,
currentIndex + 1,
undefined,
currentModel.provider,
);
if (result) return result;

// Wrap around: try entries before the current one
const wrapResult = await this._findAvailableInChain(chainName, entries, 0, currentIndex);
const wrapResult = await this._findAvailableInChain(
chainName,
entries,
0,
currentIndex,
currentModel.provider,
);
if (wrapResult) return wrapResult;
}

Expand Down Expand Up @@ -134,14 +146,18 @@ export class FallbackResolver {
entries: FallbackChainEntry[],
startIndex: number,
endIndex?: number,
allowBackedOffProvider?: string,
): Promise<FallbackResult | null> {
const end = endIndex ?? entries.length;

for (let i = startIndex; i < end; i++) {
const entry = entries[i];

// Check provider-level backoff
if (!this.authStorage.isProviderAvailable(entry.provider)) {
if (
entry.provider !== allowBackedOffProvider &&
!this.authStorage.isProviderAvailable(entry.provider)
) {
continue;
}

Expand Down
2 changes: 2 additions & 0 deletions packages/gsd-agent-core/src/sdk.ts
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import { clampThinkingLevel, type Message, type Model, streamSimple } from "@gsd
import { getAgentDir } from "@gsd/pi-coding-agent/config.js";
import { resolvePath } from "@gsd/pi-coding-agent/utils/paths.js";
import { AgentSession } from "./agent-session.js";
import { FallbackResolver } from "./fallback-resolver.js";
import { formatNoModelsAvailableMessage } from "@gsd/pi-coding-agent/core/auth-guidance.js";
import { AuthStorage } from "@gsd/pi-coding-agent/core/auth-storage.js";
import { DEFAULT_THINKING_LEVEL } from "@gsd/pi-coding-agent/core/defaults.js";
Expand Down Expand Up @@ -428,6 +429,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {}
allowedToolNames,
extensionRunnerRef,
sessionStartEvent: options.sessionStartEvent,
fallbackResolver: new FallbackResolver(settingsManager, authStorage, modelRegistry),
});
const extensionsResult = resourceLoader.getExtensions();

Expand Down
2 changes: 2 additions & 0 deletions packages/gsd-agent-core/src/session/agent-session-host.ts
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@ import type {
TurnLatencyStatus,
TurnLatencyVisibleKind,
} from "../turn-latency.js";
import type { FallbackResolver } from "../fallback-resolver.js";

/**
* Internal surface shared by AgentSession submodule classes.
Expand All @@ -56,6 +57,7 @@ export interface AgentSessionHost {
readonly settingsManager: SettingsManager;
readonly modelRegistry: ModelRegistry;
readonly resourceLoader: ResourceLoader;
readonly fallbackResolver: FallbackResolver | undefined;

// Mutable session state accessed across modules
_scopedModels: Array<{ model: Model<any>; thinkingLevel?: ThinkingLevel }>;
Expand Down
39 changes: 36 additions & 3 deletions packages/gsd-agent-core/src/session/agent-session-prompt.ts
Original file line number Diff line number Diff line change
Expand Up @@ -41,8 +41,13 @@ export class AgentSessionPromptModule {
return false;
}

if (this.isRetryableError(msg) && (await this.prepareRetry(msg))) {
return true;
if (this.isRetryableError(msg)) {
if (await this.tryFallback()) {
return true;
}
if (await this.prepareRetry(msg)) {
return true;
}
}

if (msg.stopReason === "error" && this.host._retryAttempt > 0) {
Expand Down Expand Up @@ -485,7 +490,35 @@ export class AgentSessionPromptModule {
// Match: overloaded_error, provider returned error, rate limit, 429, 500, 502, 503, 504, service unavailable, network/connection errors (including connection lost), WebSocket transport closes/errors, fetch failed, premature stream endings, HTTP/2 closed before response, terminated, retry delay exceeded
return /overloaded|provider.?returned.?error|rate.?limit|too many requests|429|500|502|503|504|service.?unavailable|server.?error|internal.?error|network.?error|connection.?error|connection.?refused|connection.?lost|websocket.?closed|websocket.?error|other side closed|fetch failed|upstream.?connect|reset before headers|socket hang up|ended without|stream ended before message_stop|http2 request did not get a response|timed? out|timeout|terminated|retry delay/i.test(
err,
);
);
}

async tryFallback(): Promise<boolean> {
const resolver = this.host.fallbackResolver;
const currentModel = this.host.model;
if (!resolver || !currentModel) {
return false;
}

const result = await resolver.findFallback(currentModel, "retryable");
if (!result) {
return false;
}

const messages = this.host.agent.state.messages;
if (messages.length > 0 && messages[messages.length - 1].role === "assistant") {
this.host.agent.state.messages = messages.slice(0, -1);
}
this.host._retryAttempt = 0;
await this.host.setModel(result.model);
this.host.emit({
type: "fallback_provider_switch",
fromProvider: currentModel.provider,
toProvider: result.model.provider,
model: result.model,
reason: result.reason,
});
return true;
}

async prepareRetry(message: AssistantMessage): Promise<boolean> {
Expand Down
3 changes: 3 additions & 0 deletions packages/gsd-agent-core/src/session/agent-session-types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ import type { SettingsManager } from "@gsd/pi-coding-agent/core/settings-manager
import type { SourceInfo } from "@gsd/pi-coding-agent/core/source-info.js";
import type { SessionStartEvent } from "@gsd/pi-coding-agent/core/extensions/index.js";
import type { CompactionResult } from "../compaction/index.js";
import type { FallbackResolver } from "../fallback-resolver.js";

// Skill Block Parsing
// ============================================================================
Expand Down Expand Up @@ -153,6 +154,8 @@ export interface AgentSessionConfig {
extensionRunnerRef?: { current?: ExtensionRunner };
/** Session start event metadata emitted when extensions bind to this runtime. */
sessionStartEvent?: SessionStartEvent;
/** Optional fallback resolver for switching models after provider errors. */
fallbackResolver?: FallbackResolver;
}

export interface ExtensionBindings {
Expand Down
52 changes: 52 additions & 0 deletions packages/pi-coding-agent/test/model-registry.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1111,6 +1111,58 @@ describe("ModelRegistry", () => {
expect(() => streamSimple(scopedModel, emptyContext)).not.toThrow("custom streamSimple override");
});

test("anthropic-compatible provider override does not hijack other providers sharing anthropic-messages", async () => {
const registry = ModelRegistry.create(authStorage, modelsJsonPath);

registry.registerProvider("custom-anthropic-wrapper", {
api: "anthropic-messages",
baseUrl: "https://wrapper.test/anthropic",
apiKey: "TEST_KEY",
models: [
{
id: "wrapper-model",
name: "Wrapper Model",
reasoning: true,
input: ["text"],
cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
contextWindow: 128000,
maxTokens: 4096,
},
],
streamSimple: () => {
throw new Error("wrapper streamSimple override");
},
});

const wrapperModel = registry.find("custom-anthropic-wrapper", "wrapper-model");
if (!wrapperModel) {
throw new Error("Expected wrapper model to be registered");
}
const anthropicModel: Model<Api> = {
id: "test-anthropic-model",
name: "Test Anthropic Model",
api: "anthropic-messages",
provider: "anthropic",
baseUrl: "https://api.anthropic.com/v1",
reasoning: true,
input: ["text"],
cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
contextWindow: 200000,
maxTokens: 4096,
};

await expect(streamSimple(wrapperModel, emptyContext).result()).resolves.toMatchObject({
provider: "custom-anthropic-wrapper",
model: "wrapper-model",
errorMessage: "No API key for provider: custom-anthropic-wrapper",
});
await expect(
getApiProvider("anthropic-messages")?.streamSimple(anthropicModel, emptyContext).result(),
).resolves.not.toMatchObject({
provider: "custom-anthropic-wrapper",
});
});

describe("dynamic provider override persistence", () => {
test("baseUrl-only override keeps built-in provider models after refresh", () => {
const registry = ModelRegistry.create(authStorage, modelsJsonPath);
Expand Down
Loading
Loading