diff --git a/packages/ai/src/providers/anthropic.ts b/packages/ai/src/providers/anthropic.ts index 11c17701..8535448a 100644 --- a/packages/ai/src/providers/anthropic.ts +++ b/packages/ai/src/providers/anthropic.ts @@ -4,6 +4,7 @@ import type { ContentBlockParam, MessageCreateParamsStreaming, MessageParam, + RawMessageStreamEvent, } from "@anthropic-ai/sdk/resources/messages.js"; import { getEnvApiKey } from "../env-api-keys.js"; import { calculateCost } from "../models.js"; @@ -27,7 +28,7 @@ import type { } from "../types.js"; import { AssistantMessageEventStream } from "../utils/event-stream.js"; import { headersToRecord } from "../utils/headers.js"; -import { parseStreamingJson } from "../utils/json-parse.js"; +import { parseJsonWithRepair, parseStreamingJson } from "../utils/json-parse.js"; import { sanitizeSurrogates } from "../utils/sanitize-unicode.js"; import { buildCopilotDynamicHeaders, hasCopilotVisionInput } from "./github-copilot-headers.js"; @@ -213,6 +214,176 @@ function mergeHeaders(...headerSources: (Record | undefined)[]): return merged; } +interface ServerSentEvent { + event: string | null; + data: string; + raw: string[]; +} + +interface SseDecoderState { + event: string | null; + data: string[]; + raw: string[]; +} + +function flushSseEvent(state: SseDecoderState): ServerSentEvent | null { + if (!state.event && state.data.length === 0) { + return null; + } + + const event: ServerSentEvent = { + event: state.event, + data: state.data.join("\n"), + raw: [...state.raw], + }; + state.event = null; + state.data = []; + state.raw = []; + return event; +} + +function decodeSseLine(line: string, state: SseDecoderState): ServerSentEvent | null { + if (line === "") { + return flushSseEvent(state); + } + + state.raw.push(line); + if (line.startsWith(":")) { + return null; + } + + const delimiterIndex = line.indexOf(":"); + const fieldName = delimiterIndex === -1 ? line : line.slice(0, delimiterIndex); + let value = delimiterIndex === -1 ? "" : line.slice(delimiterIndex + 1); + if (value.startsWith(" ")) { + value = value.slice(1); + } + + if (fieldName === "event") { + state.event = value; + } else if (fieldName === "data") { + state.data.push(value); + } + + return null; +} + +function nextLineBreakIndex(text: string): number { + const carriageReturnIndex = text.indexOf("\r"); + const newlineIndex = text.indexOf("\n"); + if (carriageReturnIndex === -1) { + return newlineIndex; + } + if (newlineIndex === -1) { + return carriageReturnIndex; + } + return Math.min(carriageReturnIndex, newlineIndex); +} + +function consumeLine(text: string): { line: string; rest: string } | null { + const lineBreakIndex = nextLineBreakIndex(text); + if (lineBreakIndex === -1) { + return null; + } + + let nextIndex = lineBreakIndex + 1; + if (text[lineBreakIndex] === "\r" && text[nextIndex] === "\n") { + nextIndex += 1; + } + + return { + line: text.slice(0, lineBreakIndex), + rest: text.slice(nextIndex), + }; +} + +async function* iterateSseMessages( + body: ReadableStream, + signal?: AbortSignal, +): AsyncGenerator { + const reader = body.getReader(); + const decoder = new TextDecoder(); + const state: SseDecoderState = { event: null, data: [], raw: [] }; + let buffer = ""; + + try { + while (true) { + if (signal?.aborted) { + throw new Error("Request was aborted"); + } + + const { value, done } = await reader.read(); + if (done) { + break; + } + + buffer += decoder.decode(value, { stream: true }); + let consumed = consumeLine(buffer); + while (consumed) { + buffer = consumed.rest; + const event = decodeSseLine(consumed.line, state); + if (event) { + yield event; + } + consumed = consumeLine(buffer); + } + } + + buffer += decoder.decode(); + let consumed = consumeLine(buffer); + while (consumed) { + buffer = consumed.rest; + const event = decodeSseLine(consumed.line, state); + if (event) { + yield event; + } + consumed = consumeLine(buffer); + } + + if (buffer.length > 0) { + const event = decodeSseLine(buffer, state); + if (event) { + yield event; + } + } + + const trailingEvent = flushSseEvent(state); + if (trailingEvent) { + yield trailingEvent; + } + } finally { + reader.releaseLock(); + } +} + +async function* iterateAnthropicEvents( + response: Response, + signal?: AbortSignal, +): AsyncGenerator { + if (!response.body) { + throw new Error("Attempted to iterate over an Anthropic response with no body"); + } + + for await (const sse of iterateSseMessages(response.body, signal)) { + if (!sse.event || sse.event === "ping") { + continue; + } + + if (sse.event === "error") { + throw new Error(sse.data); + } + + try { + yield parseJsonWithRepair(sse.data); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + throw new Error( + `Could not parse Anthropic SSE event ${sse.event}: ${message}; data=${sse.data}; raw=${sse.raw.join("\\n")}`, + ); + } + } +} + export const streamAnthropic: StreamFunction<"anthropic-messages", AnthropicOptions> = ( model: Model<"anthropic-messages">, context: Context, @@ -273,16 +444,16 @@ export const streamAnthropic: StreamFunction<"anthropic-messages", AnthropicOpti if (nextParams !== undefined) { params = nextParams as MessageCreateParamsStreaming; } - const { data: anthropicStream, response } = await client.messages - .stream({ ...params, stream: true }, { signal: options?.signal }) - .withResponse(); + const response = await client.messages + .create({ ...params, stream: true }, { signal: options?.signal }) + .asResponse(); await options?.onResponse?.({ status: response.status, headers: headersToRecord(response.headers) }, model); stream.push({ type: "start", partial: output }); type Block = (ThinkingContent | TextContent | (ToolCall & { partialJson: string })) & { index: number }; const blocks = output.content as Block[]; - for await (const event of anthropicStream) { + for await (const event of iterateAnthropicEvents(response, options?.signal)) { if (event.type === "message_start") { output.responseId = event.message.id; // Capture initial token usage from message_start event diff --git a/packages/ai/src/utils/json-parse.ts b/packages/ai/src/utils/json-parse.ts index feeb32ad..f39f9f09 100644 --- a/packages/ai/src/utils/json-parse.ts +++ b/packages/ai/src/utils/json-parse.ts @@ -1,5 +1,99 @@ import { parse as partialParse } from "partial-json"; +const VALID_JSON_ESCAPES = new Set(['"', "\\", "/", "b", "f", "n", "r", "t", "u"]); + +function isControlCharacter(char: string): boolean { + const codePoint = char.codePointAt(0); + return codePoint !== undefined && codePoint >= 0x00 && codePoint <= 0x1f; +} + +function escapeControlCharacter(char: string): string { + switch (char) { + case "\b": + return "\\b"; + case "\f": + return "\\f"; + case "\n": + return "\\n"; + case "\r": + return "\\r"; + case "\t": + return "\\t"; + default: + return `\\u${char.codePointAt(0)?.toString(16).padStart(4, "0") ?? "0000"}`; + } +} + +/** + * Repairs malformed JSON string literals by: + * - escaping raw control characters inside strings + * - doubling backslashes before invalid escape characters + */ +export function repairJson(json: string): string { + let repaired = ""; + let inString = false; + + for (let index = 0; index < json.length; index++) { + const char = json[index]; + + if (!inString) { + repaired += char; + if (char === '"') { + inString = true; + } + continue; + } + + if (char === '"') { + repaired += char; + inString = false; + continue; + } + + if (char === "\\") { + const nextChar = json[index + 1]; + if (nextChar === undefined) { + repaired += "\\\\"; + continue; + } + + if (nextChar === "u") { + const unicodeDigits = json.slice(index + 2, index + 6); + if (/^[0-9a-fA-F]{4}$/.test(unicodeDigits)) { + repaired += `\\u${unicodeDigits}`; + index += 5; + continue; + } + } + + if (VALID_JSON_ESCAPES.has(nextChar)) { + repaired += `\\${nextChar}`; + index += 1; + continue; + } + + repaired += "\\\\"; + continue; + } + + repaired += isControlCharacter(char) ? escapeControlCharacter(char) : char; + } + + return repaired; +} + +export function parseJsonWithRepair(json: string): T { + try { + return JSON.parse(json) as T; + } catch (error) { + const repairedJson = repairJson(json); + if (repairedJson !== json) { + return JSON.parse(repairedJson) as T; + } + throw error; + } +} + /** * Attempts to parse potentially incomplete JSON during streaming. * Always returns a valid object, even if the JSON is incomplete. @@ -7,22 +101,24 @@ import { parse as partialParse } from "partial-json"; * @param partialJson The partial JSON string from streaming * @returns Parsed object or empty object if parsing fails */ -export function parseStreamingJson(partialJson: string | undefined): T { +export function parseStreamingJson>(partialJson: string | undefined): T { if (!partialJson || partialJson.trim() === "") { return {} as T; } - // Try standard parsing first (fastest for complete JSON) try { - return JSON.parse(partialJson) as T; + return parseJsonWithRepair(partialJson); } catch { - // Try partial-json for incomplete JSON try { const result = partialParse(partialJson); return (result ?? {}) as T; } catch { - // If all parsing fails, return empty object - return {} as T; + try { + const result = partialParse(repairJson(partialJson)); + return (result ?? {}) as T; + } catch { + return {} as T; + } } } } diff --git a/packages/ai/test/anthropic-sse-parsing.test.ts b/packages/ai/test/anthropic-sse-parsing.test.ts new file mode 100644 index 00000000..55af15af --- /dev/null +++ b/packages/ai/test/anthropic-sse-parsing.test.ts @@ -0,0 +1,113 @@ +import type Anthropic from "@anthropic-ai/sdk"; +import { Type } from "@sinclair/typebox"; +import { describe, expect, it } from "vitest"; +import { getModel } from "../src/models.js"; +import { streamAnthropic } from "../src/providers/anthropic.js"; +import type { Context, ToolCall } from "../src/types.js"; + +function createSseResponse(events: Array<{ event: string; data: string }>): Response { + const body = events.map(({ event, data }) => `event: ${event}\ndata: ${data}\n`).join("\n"); + return new Response(body, { + status: 200, + headers: { "content-type": "text/event-stream" }, + }); +} + +function createFakeAnthropicClient(response: Response): Anthropic { + return { + messages: { + create: () => ({ + asResponse: async () => response, + }), + }, + } as unknown as Anthropic; +} + +describe("Anthropic raw SSE parsing", () => { + it("repairs malformed SSE JSON and malformed streamed tool JSON", async () => { + const model = getModel("anthropic", "claude-haiku-4-5"); + const context: Context = { + messages: [{ role: "user", content: "Use the edit tool.", timestamp: Date.now() }], + tools: [ + { + name: "edit", + description: "Edit a file.", + parameters: Type.Object({ + path: Type.String(), + text: Type.String(), + }), + }, + ], + }; + + const malformedToolJsonDelta = String.raw`{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{\"path\":\"A\H\",\"text\":\"col1 col2\"}"}}`; + + const response = createSseResponse([ + { + event: "message_start", + data: JSON.stringify({ + type: "message_start", + message: { + id: "msg_test", + usage: { + input_tokens: 12, + output_tokens: 0, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }, + }, + }), + }, + { + event: "content_block_start", + data: JSON.stringify({ + type: "content_block_start", + index: 0, + content_block: { + type: "tool_use", + id: "toolu_test", + name: "edit", + input: {}, + }, + }), + }, + { event: "content_block_delta", data: malformedToolJsonDelta }, + { + event: "content_block_stop", + data: JSON.stringify({ type: "content_block_stop", index: 0 }), + }, + { + event: "message_delta", + data: JSON.stringify({ + type: "message_delta", + delta: { stop_reason: "tool_use" }, + usage: { + input_tokens: 12, + output_tokens: 5, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }, + }), + }, + { + event: "message_stop", + data: JSON.stringify({ type: "message_stop" }), + }, + ]); + + const stream = streamAnthropic(model, context, { + client: createFakeAnthropicClient(response), + }); + const result = await stream.result(); + + expect(result.stopReason).toBe("toolUse"); + expect(result.errorMessage).toBeUndefined(); + + const toolCall = result.content.find((block): block is ToolCall => block.type === "toolCall"); + expect(toolCall).toBeDefined(); + expect(toolCall?.arguments).toEqual({ + path: "A\\H", + text: "col1\tcol2", + }); + }); +}); diff --git a/packages/ai/test/cache-retention.test.ts b/packages/ai/test/cache-retention.test.ts index 530679c6..b517d0ba 100644 --- a/packages/ai/test/cache-retention.test.ts +++ b/packages/ai/test/cache-retention.test.ts @@ -27,7 +27,7 @@ describe("Cache Retention (PI_CACHE_RETENTION)", () => { it.skipIf(!process.env.ANTHROPIC_API_KEY)( "should use default cache TTL (no ttl field) when PI_CACHE_RETENTION is not set", async () => { - const model = getModel("anthropic", "claude-3-5-haiku-20241022"); + const model = getModel("anthropic", "claude-haiku-4-5"); let capturedPayload: any = null; const s = stream(model, context, { @@ -50,7 +50,7 @@ describe("Cache Retention (PI_CACHE_RETENTION)", () => { it.skipIf(!process.env.ANTHROPIC_API_KEY)("should use 1h cache TTL when PI_CACHE_RETENTION=long", async () => { process.env.PI_CACHE_RETENTION = "long"; - const model = getModel("anthropic", "claude-3-5-haiku-20241022"); + const model = getModel("anthropic", "claude-haiku-4-5"); let capturedPayload: any = null; const s = stream(model, context, { @@ -74,7 +74,7 @@ describe("Cache Retention (PI_CACHE_RETENTION)", () => { process.env.PI_CACHE_RETENTION = "long"; // Create a model with a different baseUrl (simulating a proxy) - const baseModel = getModel("anthropic", "claude-3-5-haiku-20241022"); + const baseModel = getModel("anthropic", "claude-haiku-4-5"); const proxyModel = { ...baseModel, baseUrl: "https://my-proxy.example.com/v1", @@ -114,7 +114,7 @@ describe("Cache Retention (PI_CACHE_RETENTION)", () => { }); it("should omit cache_control when cacheRetention is none", async () => { - const baseModel = getModel("anthropic", "claude-3-5-haiku-20241022"); + const baseModel = getModel("anthropic", "claude-haiku-4-5"); let capturedPayload: any = null; const { streamAnthropic } = await import("../src/providers/anthropic.js"); @@ -140,7 +140,7 @@ describe("Cache Retention (PI_CACHE_RETENTION)", () => { }); it("should add cache_control to string user messages", async () => { - const baseModel = getModel("anthropic", "claude-3-5-haiku-20241022"); + const baseModel = getModel("anthropic", "claude-haiku-4-5"); let capturedPayload: any = null; const { streamAnthropic } = await import("../src/providers/anthropic.js"); @@ -168,7 +168,7 @@ describe("Cache Retention (PI_CACHE_RETENTION)", () => { }); it("should set 1h cache TTL when cacheRetention is long", async () => { - const baseModel = getModel("anthropic", "claude-3-5-haiku-20241022"); + const baseModel = getModel("anthropic", "claude-haiku-4-5"); let capturedPayload: any = null; const { streamAnthropic } = await import("../src/providers/anthropic.js"); diff --git a/packages/ai/test/context-overflow.test.ts b/packages/ai/test/context-overflow.test.ts index 4a3e876b..5b1e08ef 100644 --- a/packages/ai/test/context-overflow.test.ts +++ b/packages/ai/test/context-overflow.test.ts @@ -100,8 +100,8 @@ function logResult(result: OverflowResult) { describe("Context overflow error handling", () => { describe.skipIf(!process.env.ANTHROPIC_API_KEY)("Anthropic (API Key)", () => { - it("claude-3-5-haiku - should detect overflow via isContextOverflow", async () => { - const model = getModel("anthropic", "claude-3-5-haiku-20241022"); + it("claude-haiku-4-5 - should detect overflow via isContextOverflow", async () => { + const model = getModel("anthropic", "claude-haiku-4-5"); const result = await testContextOverflow(model, process.env.ANTHROPIC_API_KEY!); logResult(result); diff --git a/packages/ai/test/empty.test.ts b/packages/ai/test/empty.test.ts index 088c3529..dd967975 100644 --- a/packages/ai/test/empty.test.ts +++ b/packages/ai/test/empty.test.ts @@ -229,7 +229,7 @@ describe("AI Providers Empty Message Tests", () => { }); describe.skipIf(!process.env.ANTHROPIC_API_KEY)("Anthropic Provider Empty Messages", () => { - const llm = getModel("anthropic", "claude-3-5-haiku-20241022"); + const llm = getModel("anthropic", "claude-haiku-4-5"); it("should handle empty content array", { retry: 3, timeout: 30000 }, async () => { await testEmptyMessage(llm); @@ -453,7 +453,7 @@ describe("AI Providers Empty Message Tests", () => { // ========================================================================= describe("Anthropic OAuth Provider Empty Messages", () => { - const llm = getModel("anthropic", "claude-3-5-haiku-20241022"); + const llm = getModel("anthropic", "claude-haiku-4-5"); it.skipIf(!anthropicOAuthToken)("should handle empty content array", { retry: 3, timeout: 30000 }, async () => { await testEmptyMessage(llm, { apiKey: anthropicOAuthToken }); diff --git a/packages/ai/test/github-copilot-anthropic.test.ts b/packages/ai/test/github-copilot-anthropic.test.ts index b0000beb..2fb44f23 100644 --- a/packages/ai/test/github-copilot-anthropic.test.ts +++ b/packages/ai/test/github-copilot-anthropic.test.ts @@ -4,37 +4,42 @@ import type { Context } from "../src/types.js"; const mockState = vi.hoisted(() => ({ constructorOpts: undefined as Record | undefined, - streamParams: undefined as Record | undefined, + createParams: undefined as Record | undefined, })); vi.mock("@anthropic-ai/sdk", () => { - const fakeStream = { - async *[Symbol.asyncIterator]() { - yield { + function createSseResponse(): Response { + const body = [ + `event: message_start\ndata: ${JSON.stringify({ type: "message_start", message: { + id: "msg_test", usage: { input_tokens: 10, output_tokens: 0 }, }, - }; - yield { + })}\n`, + `event: message_delta\ndata: ${JSON.stringify({ type: "message_delta", delta: { stop_reason: "end_turn" }, usage: { output_tokens: 5 }, - }; - }, - finalMessage: async () => ({ - usage: { input_tokens: 10, output_tokens: 5, cache_creation_input_tokens: 0, cache_read_input_tokens: 0 }, - }), - }; + })}\n`, + ].join("\n"); + + return new Response(body, { + status: 200, + headers: { "content-type": "text/event-stream" }, + }); + } class FakeAnthropic { constructor(opts: Record) { mockState.constructorOpts = opts; } messages = { - stream: (params: Record) => { - mockState.streamParams = params; - return fakeStream; + create: (params: Record) => { + mockState.createParams = params; + return { + asResponse: async () => createSseResponse(), + }; }, }; } @@ -79,7 +84,7 @@ describe("Copilot Claude via Anthropic Messages", () => { expect(beta).not.toContain("fine-grained-tool-streaming"); // Payload is valid Anthropic Messages format - const params = mockState.streamParams!; + const params = mockState.createParams!; expect(params.model).toBe("claude-sonnet-4"); expect(params.stream).toBe(true); expect(params.max_tokens).toBeGreaterThan(0); diff --git a/packages/ai/test/stream.test.ts b/packages/ai/test/stream.test.ts index a9f579f4..35029c0a 100644 --- a/packages/ai/test/stream.test.ts +++ b/packages/ai/test/stream.test.ts @@ -475,8 +475,8 @@ describe("Generate E2E Tests", () => { }); }); - describe.skipIf(!process.env.ANTHROPIC_API_KEY)("Anthropic Provider (claude-3-5-haiku-20241022)", () => { - const model = getModel("anthropic", "claude-3-5-haiku-20241022"); + describe.skipIf(!process.env.ANTHROPIC_API_KEY)("Anthropic Provider (claude-haiku-4-5)", () => { + const model = getModel("anthropic", "claude-haiku-4-5"); it("should complete basic text generation", { retry: 3 }, async () => { await basicTextGeneration(model, { thinkingEnabled: true }); diff --git a/packages/ai/test/tool-call-without-result.test.ts b/packages/ai/test/tool-call-without-result.test.ts index 7a21597e..7732eda4 100644 --- a/packages/ai/test/tool-call-without-result.test.ts +++ b/packages/ai/test/tool-call-without-result.test.ts @@ -137,7 +137,7 @@ describe("Tool Call Without Result Tests", () => { }); describe.skipIf(!process.env.ANTHROPIC_API_KEY)("Anthropic Provider", () => { - const model = getModel("anthropic", "claude-3-5-haiku-20241022"); + const model = getModel("anthropic", "claude-haiku-4-5"); it("should filter out tool calls without corresponding tool results", { retry: 3, timeout: 30000 }, async () => { await testToolCallWithoutResult(model); @@ -229,7 +229,7 @@ describe("Tool Call Without Result Tests", () => { // ========================================================================= describe("Anthropic OAuth Provider", () => { - const model = getModel("anthropic", "claude-3-5-haiku-20241022"); + const model = getModel("anthropic", "claude-haiku-4-5"); it.skipIf(!anthropicOAuthToken)( "should filter out tool calls without corresponding tool results", diff --git a/packages/ai/test/total-tokens.test.ts b/packages/ai/test/total-tokens.test.ts index 732a522b..e81cc125 100644 --- a/packages/ai/test/total-tokens.test.ts +++ b/packages/ai/test/total-tokens.test.ts @@ -106,10 +106,10 @@ describe("totalTokens field", () => { describe.skipIf(!process.env.ANTHROPIC_API_KEY)("Anthropic (API Key)", () => { it( - "claude-3-5-haiku - should return totalTokens equal to sum of components", + "claude-sonnet-4-5 - should return totalTokens equal to sum of components", { retry: 3, timeout: 60000 }, async () => { - const llm = getModel("anthropic", "claude-3-5-haiku-20241022"); + const llm = getModel("anthropic", "claude-sonnet-4-5"); console.log(`\nAnthropic / ${llm.id}:`); const { first, second } = await testTotalTokensWithCache(llm, { apiKey: process.env.ANTHROPIC_API_KEY }); diff --git a/packages/ai/test/unicode-surrogate.test.ts b/packages/ai/test/unicode-surrogate.test.ts index 56a81a95..915b35e6 100644 --- a/packages/ai/test/unicode-surrogate.test.ts +++ b/packages/ai/test/unicode-surrogate.test.ts @@ -352,7 +352,7 @@ describe("AI Providers Unicode Surrogate Pair Tests", () => { }); describe.skipIf(!process.env.ANTHROPIC_API_KEY)("Anthropic Provider Unicode Handling", () => { - const llm = getModel("anthropic", "claude-3-5-haiku-20241022"); + const llm = getModel("anthropic", "claude-haiku-4-5"); it("should handle emoji in tool results", { retry: 3, timeout: 30000 }, async () => { await testEmojiInToolResults(llm); @@ -372,7 +372,7 @@ describe("AI Providers Unicode Surrogate Pair Tests", () => { // ========================================================================= describe("Anthropic OAuth Provider Unicode Handling", () => { - const llm = getModel("anthropic", "claude-3-5-haiku-20241022"); + const llm = getModel("anthropic", "claude-haiku-4-5"); it.skipIf(!anthropicOAuthToken)("should handle emoji in tool results", { retry: 3, timeout: 30000 }, async () => { await testEmojiInToolResults(llm, { apiKey: anthropicOAuthToken }); diff --git a/packages/coding-agent/test/agent-session-model-switch-thinking.test.ts b/packages/coding-agent/test/agent-session-model-switch-thinking.test.ts deleted file mode 100644 index 28ff93eb..00000000 --- a/packages/coding-agent/test/agent-session-model-switch-thinking.test.ts +++ /dev/null @@ -1,79 +0,0 @@ -import { Agent, type ThinkingLevel } from "@mariozechner/pi-agent-core"; -import { getModel } from "@mariozechner/pi-ai"; -import { describe, expect, it } from "vitest"; -import { AgentSession } from "../src/core/agent-session.js"; -import { AuthStorage } from "../src/core/auth-storage.js"; -import { ModelRegistry } from "../src/core/model-registry.js"; -import { SessionManager } from "../src/core/session-manager.js"; -import { SettingsManager } from "../src/core/settings-manager.js"; -import { createTestResourceLoader } from "./utilities.js"; - -const reasoningModel = getModel("anthropic", "claude-sonnet-4-5")!; -const nonReasoningModel = getModel("anthropic", "claude-3-5-haiku-latest")!; - -function createSession({ - thinkingLevel = "high", - defaultThinkingLevel = thinkingLevel, - scopedModels, -}: { - thinkingLevel?: ThinkingLevel; - defaultThinkingLevel?: ThinkingLevel; - scopedModels?: Array<{ model: typeof reasoningModel; thinkingLevel?: ThinkingLevel }>; -} = {}) { - const settingsManager = SettingsManager.inMemory({ defaultThinkingLevel }); - const sessionManager = SessionManager.inMemory(); - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("anthropic", "test-key"); - const session = new AgentSession({ - agent: new Agent({ - getApiKey: () => "test-key", - initialState: { - model: reasoningModel, - systemPrompt: "You are a helpful assistant.", - tools: [], - thinkingLevel, - }, - }), - sessionManager, - settingsManager, - cwd: process.cwd(), - modelRegistry: ModelRegistry.inMemory(authStorage), - resourceLoader: createTestResourceLoader(), - scopedModels, - }); - - return { session, sessionManager, settingsManager }; -} - -describe("AgentSession model switching", () => { - it("preserves the saved thinking preference through non-reasoning models", async () => { - const { session, sessionManager, settingsManager } = createSession({ - scopedModels: [{ model: reasoningModel }, { model: nonReasoningModel }], - }); - - try { - await session.setModel(nonReasoningModel); - expect(session.thinkingLevel).toBe("off"); - expect(settingsManager.getDefaultThinkingLevel()).toBe("high"); - - await session.setModel(reasoningModel); - expect(session.thinkingLevel).toBe("high"); - - await session.cycleModel(); - expect(session.thinkingLevel).toBe("off"); - expect(settingsManager.getDefaultThinkingLevel()).toBe("high"); - - await session.cycleModel(); - expect(session.thinkingLevel).toBe("high"); - expect(settingsManager.getDefaultThinkingLevel()).toBe("high"); - expect( - sessionManager - .getEntries() - .filter((entry) => entry.type === "thinking_level_change") - .map((entry) => entry.thinkingLevel), - ).toEqual(["off", "high", "off", "high"]); - } finally { - session.dispose(); - } - }); -});