From cc91f80e99af5385937b48e7e94fcd75f9e06b94 Mon Sep 17 00:00:00 2001 From: can1357 Date: Wed, 28 Jan 2026 02:39:28 +0100 Subject: [PATCH] feat: added toolChoice support and optimized TUI - Added ToolChoice type and toolChoice parameter support across all AI providers (OpenAI, Azure OpenAI, Anthropic, Google) enabling fine-grained control over tool/function selection during LLM calls. - Added toolChoice override capability to Agent.prompt() method and session prompt options allowing callers to control tool selection behavior per request. - Added provider-specific tool choice mapping functions (mapAnthropicToolChoice, mapGoogleToolChoice, mapOpenAiToolChoice) to normalize tool choice formats across different LLM APIs. - Removed kernel heartbeat/ping mechanism from PythonKernel, simplifying health monitoring by relying on direct isAlive() checks instead of periodic HTTP requests. --- .gitignore | 2 + packages/agent/src/agent.ts | 30 ++- packages/ai/scripts/generate-models.ts | 4 + packages/ai/src/models.generated.ts | 47 ++++- .../src/providers/azure-openai-responses.ts | 5 + .../src/providers/openai-codex-responses.ts | 27 ++- .../openai-codex/request-transformer.ts | 1 + .../ai/src/providers/openai-completions.ts | 7 +- packages/ai/src/providers/openai-responses.ts | 5 + packages/ai/src/stream.ts | 102 +++++++++- packages/ai/src/types.ts | 13 ++ ...nai-completions-tool-result-images.test.ts | 1 + packages/coding-agent/src/ipy/executor.ts | 5 +- packages/coding-agent/src/ipy/kernel.ts | 42 ----- .../src/prompts/agents/reviewer.md | 2 + .../coding-agent/src/session/agent-session.ts | 6 +- packages/coding-agent/src/task/executor.ts | 175 ++++++++++++++---- packages/coding-agent/src/task/render.ts | 23 +-- .../test/core/python-kernel.test.ts | 45 ----- packages/tui/src/tui.ts | 9 +- packages/tui/src/utils.ts | 8 +- 21 files changed, 392 insertions(+), 167 deletions(-) diff --git a/.gitignore b/.gitignore index 0d567c228..b96c70299 100644 --- a/.gitignore +++ b/.gitignore @@ -42,3 +42,5 @@ changes/ __pycache__/ *.b64.js + +*.cpuprofile diff --git a/packages/agent/src/agent.ts b/packages/agent/src/agent.ts index 891f7b6e7..928176753 100644 --- a/packages/agent/src/agent.ts +++ b/packages/agent/src/agent.ts @@ -13,6 +13,7 @@ import { streamSimple, type TextContent, type ThinkingBudgets, + type ToolChoice, type ToolResultMessage, } from "@oh-my-pi/pi-ai"; import { agentLoop, agentLoopContinue } from "./agent-loop"; @@ -107,6 +108,10 @@ export interface AgentOptions { cursorOnToolResult?: CursorToolResultHandler; } +export interface AgentPromptOptions { + toolChoice?: ToolChoice; +} + /** Buffered Cursor tool result with text position at time of call */ interface CursorToolResultEntry { toolResult: ToolResultMessage; @@ -358,9 +363,13 @@ export class Agent { } /** Send a prompt with an AgentMessage */ - async prompt(message: AgentMessage | AgentMessage[]): Promise; - async prompt(input: string, images?: ImageContent[]): Promise; - async prompt(input: string | AgentMessage | AgentMessage[], images?: ImageContent[]) { + async prompt(message: AgentMessage | AgentMessage[], options?: AgentPromptOptions): Promise; + async prompt(input: string, images?: ImageContent[], options?: AgentPromptOptions): Promise; + async prompt( + input: string | AgentMessage | AgentMessage[], + imagesOrOptions?: ImageContent[] | AgentPromptOptions, + options?: AgentPromptOptions, + ) { if (this._state.isStreaming) { throw new Error( "Agent is already processing a prompt. Use steer() or followUp() to queue messages, or wait for completion.", @@ -371,10 +380,19 @@ export class Agent { if (!model) throw new Error("No model configured"); let msgs: AgentMessage[]; + let promptOptions: AgentPromptOptions | undefined; + let images: ImageContent[] | undefined; if (Array.isArray(input)) { msgs = input; + promptOptions = imagesOrOptions as AgentPromptOptions | undefined; } else if (typeof input === "string") { + if (Array.isArray(imagesOrOptions)) { + images = imagesOrOptions; + promptOptions = options; + } else { + promptOptions = imagesOrOptions; + } const content: Array = [{ type: "text", text: input }]; if (images && images.length > 0) { content.push(...images); @@ -388,9 +406,10 @@ export class Agent { ]; } else { msgs = [input]; + promptOptions = imagesOrOptions as AgentPromptOptions | undefined; } - await this._runLoop(msgs); + await this._runLoop(msgs, promptOptions); } /** Continue from current context (for retry after overflow) */ @@ -415,7 +434,7 @@ export class Agent { * If messages are provided, starts a new conversation turn with those messages. * Otherwise, continues from existing context. */ - private async _runLoop(messages?: AgentMessage[]) { + private async _runLoop(messages?: AgentMessage[], options?: AgentPromptOptions) { const model = this._state.model; if (!model) throw new Error("No model configured"); @@ -467,6 +486,7 @@ export class Agent { interruptMode: this.interruptMode, sessionId: this._sessionId, thinkingBudgets: this._thinkingBudgets, + toolChoice: options?.toolChoice, convertToLlm: this.convertToLlm, transformContext: this.transformContext, getApiKey: this.getApiKey, diff --git a/packages/ai/scripts/generate-models.ts b/packages/ai/scripts/generate-models.ts index ff74be31e..3aa5a21d6 100644 --- a/packages/ai/scripts/generate-models.ts +++ b/packages/ai/scripts/generate-models.ts @@ -81,6 +81,8 @@ async function fetchOpenRouterModels(): Promise[]> { const outputCost = parseFloat(model.pricing?.completion || "0") * 1_000_000; const cacheReadCost = parseFloat(model.pricing?.input_cache_read || "0") * 1_000_000; const cacheWriteCost = parseFloat(model.pricing?.input_cache_write || "0") * 1_000_000; + // Check if model supports tool_choice parameter + const supportsToolChoice = model.supported_parameters?.includes("tool_choice") ?? false; const normalizedModel: Model = { id: modelKey, @@ -98,6 +100,8 @@ async function fetchOpenRouterModels(): Promise[]> { }, contextWindow: model.context_length || 4096, maxTokens: model.top_provider?.max_completion_tokens || 4096, + // Only add compat if tool_choice is not supported (default is true) + ...(supportsToolChoice ? {} : { compat: { supportsToolChoice: false } }), }; models.push(normalizedModel); } diff --git a/packages/ai/src/models.generated.ts b/packages/ai/src/models.generated.ts index 3a0b93c8d..ab94dae13 100644 --- a/packages/ai/src/models.generated.ts +++ b/packages/ai/src/models.generated.ts @@ -4610,6 +4610,7 @@ export const MODELS = { api: "openai-completions", provider: "openrouter", baseUrl: "https://openrouter.ai/api/v1", + compat: {"supportsToolChoice":false}, reasoning: false, input: ["text", "image"], cost: { @@ -4627,6 +4628,7 @@ export const MODELS = { api: "openai-completions", provider: "openrouter", baseUrl: "https://openrouter.ai/api/v1", + compat: {"supportsToolChoice":false}, reasoning: false, input: ["text"], cost: { @@ -4644,6 +4646,7 @@ export const MODELS = { api: "openai-completions", provider: "openrouter", baseUrl: "https://openrouter.ai/api/v1", + compat: {"supportsToolChoice":false}, reasoning: false, input: ["text", "image"], cost: { @@ -4661,6 +4664,7 @@ export const MODELS = { api: "openai-completions", provider: "openrouter", baseUrl: "https://openrouter.ai/api/v1", + compat: {"supportsToolChoice":false}, reasoning: false, input: ["text", "image"], cost: { @@ -4859,6 +4863,24 @@ export const MODELS = { contextWindow: 1000000, maxTokens: 64000, } satisfies Model<"openai-completions">, + "arcee-ai/trinity-large-preview:free": { + id: "arcee-ai/trinity-large-preview:free", + name: "Arcee AI: Trinity Large Preview (free)", + api: "openai-completions", + provider: "openrouter", + baseUrl: "https://openrouter.ai/api/v1", + compat: {"supportsToolChoice":false}, + reasoning: false, + input: ["text"], + cost: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + }, + contextWindow: 131000, + maxTokens: 4096, + } satisfies Model<"openai-completions">, "arcee-ai/trinity-mini": { id: "arcee-ai/trinity-mini", name: "Arcee AI: Trinity Mini", @@ -5227,7 +5249,7 @@ export const MODELS = { cost: { input: 0.21, output: 0.32, - cacheRead: 0, + cacheRead: 0.21, cacheWrite: 0, }, contextWindow: 163840, @@ -6264,11 +6286,11 @@ export const MODELS = { cost: { input: 0.6, output: 3, - cacheRead: 0.09999999999999999, + cacheRead: 0, cacheWrite: 0, }, contextWindow: 262144, - maxTokens: 4096, + maxTokens: 262144, } satisfies Model<"openai-completions">, "nex-agi/deepseek-v3.1-nex-n1": { id: "nex-agi/deepseek-v3.1-nex-n1", @@ -7624,7 +7646,7 @@ export const MODELS = { cost: { input: 0.049999999999999996, output: 0.25, - cacheRead: 0, + cacheRead: 0.049999999999999996, cacheWrite: 0, }, contextWindow: 32000, @@ -8788,6 +8810,23 @@ export const MODELS = { contextWindow: 1000000, maxTokens: 64000, } satisfies Model<"anthropic-messages">, + "arcee-ai/trinity-large-preview": { + id: "arcee-ai/trinity-large-preview", + name: "Trinity Large Preview", + api: "anthropic-messages", + provider: "vercel-ai-gateway", + baseUrl: "https://ai-gateway.vercel.sh", + reasoning: false, + input: ["text"], + cost: { + input: 0.25, + output: 1, + cacheRead: 0, + cacheWrite: 0, + }, + contextWindow: 131000, + maxTokens: 131000, + } satisfies Model<"anthropic-messages">, "bytedance/seed-1.6": { id: "bytedance/seed-1.6", name: "Seed 1.6", diff --git a/packages/ai/src/providers/azure-openai-responses.ts b/packages/ai/src/providers/azure-openai-responses.ts index 98a7c1487..625140dfe 100644 --- a/packages/ai/src/providers/azure-openai-responses.ts +++ b/packages/ai/src/providers/azure-openai-responses.ts @@ -26,6 +26,7 @@ import type { ThinkingContent, Tool, ToolCall, + ToolChoice, } from "../types"; import { AssistantMessageEventStream } from "../utils/event-stream"; import { parseStreamingJson } from "../utils/json-parse"; @@ -64,6 +65,7 @@ export interface AzureOpenAIResponsesOptions extends StreamOptions { azureResourceName?: string; azureBaseUrl?: string; azureDeploymentName?: string; + toolChoice?: ToolChoice; } /** @@ -449,6 +451,9 @@ function buildParams( if (context.tools) { params.tools = convertTools(context.tools); + if (options?.toolChoice) { + params.tool_choice = options.toolChoice; + } } if (model.reasoning) { diff --git a/packages/ai/src/providers/openai-codex-responses.ts b/packages/ai/src/providers/openai-codex-responses.ts index c1bfdadf5..0cbfb990e 100644 --- a/packages/ai/src/providers/openai-codex-responses.ts +++ b/packages/ai/src/providers/openai-codex-responses.ts @@ -24,6 +24,7 @@ import type { ThinkingContent, Tool, ToolCall, + ToolChoice, } from "../types"; import { AssistantMessageEventStream } from "../utils/event-stream"; import { parseStreamingJson } from "../utils/json-parse"; @@ -46,6 +47,7 @@ export interface OpenAICodexResponsesOptions extends StreamOptions { textVerbosity?: "low" | "medium" | "high"; include?: string[]; codexMode?: boolean; + toolChoice?: ToolChoice; } export const CODEX_INSTRUCTIONS = `You are an expert coding assistant operating inside pi, a coding agent harness.`; @@ -83,6 +85,23 @@ function normalizeResponsesToolCallId(id: string): { callId: string; itemId: str return { callId: `call_${hash}`, itemId: `item_${hash}` }; } +function normalizeCodexToolChoice(choice: ToolChoice | undefined): string | Record | undefined { + if (!choice) return undefined; + if (typeof choice === "string") return choice; + if (choice.type === "function") { + if ("function" in choice && choice.function?.name) { + return { type: "function", name: choice.function.name }; + } + if ("name" in choice && choice.name) { + return { type: "function", name: choice.name }; + } + } + if (choice.type === "tool" && choice.name) { + return { type: "function", name: choice.name }; + } + return undefined; +} + export const streamOpenAICodexResponses: StreamFunction<"openai-codex-responses"> = ( model: Model<"openai-codex-responses">, context: Context, @@ -139,8 +158,14 @@ export const streamOpenAICodexResponses: StreamFunction<"openai-codex-responses" params.temperature = options.temperature; } - if (context.tools) { + if (context.tools && context.tools.length > 0) { params.tools = convertTools(context.tools); + if (options?.toolChoice) { + const toolChoice = normalizeCodexToolChoice(options.toolChoice); + if (toolChoice) { + params.tool_choice = toolChoice; + } + } } const systemPrompt = buildCodexSystemPrompt({ diff --git a/packages/ai/src/providers/openai-codex/request-transformer.ts b/packages/ai/src/providers/openai-codex/request-transformer.ts index b5c847e52..69b46b03c 100644 --- a/packages/ai/src/providers/openai-codex/request-transformer.ts +++ b/packages/ai/src/providers/openai-codex/request-transformer.ts @@ -28,6 +28,7 @@ export interface RequestBody { instructions?: string; input?: InputItem[]; tools?: unknown; + tool_choice?: unknown; temperature?: number; reasoning?: Partial; text?: { diff --git a/packages/ai/src/providers/openai-completions.ts b/packages/ai/src/providers/openai-completions.ts index 81725aa0d..877f666f0 100644 --- a/packages/ai/src/providers/openai-completions.ts +++ b/packages/ai/src/providers/openai-completions.ts @@ -23,6 +23,7 @@ import type { ThinkingContent, Tool, ToolCall, + ToolChoice, ToolResultMessage, } from "../types"; import { AssistantMessageEventStream } from "../utils/event-stream"; @@ -74,7 +75,7 @@ function hasToolHistory(messages: Message[]): boolean { } export interface OpenAICompletionsOptions extends StreamOptions { - toolChoice?: "auto" | "none" | "required" | { type: "function"; function: { name: string } }; + toolChoice?: ToolChoice; reasoningEffort?: "minimal" | "low" | "medium" | "high" | "xhigh"; } @@ -420,7 +421,7 @@ function buildParams(model: Model<"openai-completions">, context: Context, optio params.tools = []; } - if (options?.toolChoice) { + if (options?.toolChoice && compat.supportsToolChoice) { params.tool_choice = options.toolChoice; } @@ -794,6 +795,7 @@ function detectCompat(model: Model<"openai-completions">): ResolvedOpenAICompat supportsDeveloperRole: !isNonStandard, supportsReasoningEffort: !isGrok && !isZai, supportsUsageInStreaming: true, + supportsToolChoice: true, maxTokensField: useMaxTokens ? "max_tokens" : "max_completion_tokens", requiresToolResultName: isMistral, requiresAssistantAfterToolResult: false, // Mistral no longer requires this as of Dec 2024 @@ -820,6 +822,7 @@ function getCompat(model: Model<"openai-completions">): ResolvedOpenAICompat { supportsDeveloperRole: model.compat.supportsDeveloperRole ?? detected.supportsDeveloperRole, supportsReasoningEffort: model.compat.supportsReasoningEffort ?? detected.supportsReasoningEffort, supportsUsageInStreaming: model.compat.supportsUsageInStreaming ?? detected.supportsUsageInStreaming, + supportsToolChoice: model.compat.supportsToolChoice ?? detected.supportsToolChoice, maxTokensField: model.compat.maxTokensField ?? detected.maxTokensField, requiresToolResultName: model.compat.requiresToolResultName ?? detected.requiresToolResultName, requiresAssistantAfterToolResult: diff --git a/packages/ai/src/providers/openai-responses.ts b/packages/ai/src/providers/openai-responses.ts index f0d6f34ad..abf647900 100644 --- a/packages/ai/src/providers/openai-responses.ts +++ b/packages/ai/src/providers/openai-responses.ts @@ -24,6 +24,7 @@ import type { ThinkingContent, Tool, ToolCall, + ToolChoice, } from "../types"; import { AssistantMessageEventStream } from "../utils/event-stream"; import { parseStreamingJson } from "../utils/json-parse"; @@ -36,6 +37,7 @@ export interface OpenAIResponsesOptions extends StreamOptions { reasoningEffort?: "minimal" | "low" | "medium" | "high" | "xhigh"; reasoningSummary?: "auto" | "detailed" | "concise" | null; serviceTier?: ResponseCreateParamsStreaming["service_tier"]; + toolChoice?: ToolChoice; /** * Enforce strict tool call/result pairing when building Responses API inputs. * Azure OpenAI Responses API requires tool results to have a matching tool call. @@ -408,6 +410,9 @@ function buildParams(model: Model<"openai-responses">, context: Context, options if (context.tools) { params.tools = convertTools(context.tools); + if (options?.toolChoice) { + params.tool_choice = options.toolChoice; + } } if (model.reasoning) { diff --git a/packages/ai/src/stream.ts b/packages/ai/src/stream.ts index 7eb6e3a76..e28a30719 100644 --- a/packages/ai/src/stream.ts +++ b/packages/ai/src/stream.ts @@ -27,6 +27,7 @@ import type { SimpleStreamOptions, ThinkingBudgets, ThinkingLevel, + ToolChoice, } from "./types"; // Set up http proxy according to env variables for `fetch` based SDKs in Node.js. @@ -263,6 +264,52 @@ function resolveBedrockThinkingBudget( return { budget, level }; } +function mapAnthropicToolChoice(choice?: ToolChoice): AnthropicOptions["toolChoice"] { + if (!choice) return undefined; + if (typeof choice === "string") { + if (choice === "required") return "any"; + if (choice === "auto" || choice === "none" || choice === "any") return choice; + return undefined; + } + if (choice.type === "tool") { + return choice.name ? { type: "tool", name: choice.name } : undefined; + } + if (choice.type === "function") { + const name = "function" in choice ? choice.function?.name : choice.name; + return name ? { type: "tool", name } : undefined; + } + return undefined; +} + +function mapGoogleToolChoice( + choice?: ToolChoice, +): GoogleOptions["toolChoice"] | GoogleGeminiCliOptions["toolChoice"] | GoogleVertexOptions["toolChoice"] { + if (!choice) return undefined; + if (typeof choice === "string") { + if (choice === "required") return "any"; + if (choice === "auto" || choice === "none" || choice === "any") return choice; + return undefined; + } + return "any"; +} + +function mapOpenAiToolChoice(choice?: ToolChoice): OpenAICompletionsOptions["toolChoice"] { + if (!choice) return undefined; + if (typeof choice === "string") { + if (choice === "any") return "required"; + if (choice === "auto" || choice === "none" || choice === "required") return choice; + return undefined; + } + if (choice.type === "tool") { + return choice.name ? { type: "function", function: { name: choice.name } } : undefined; + } + if (choice.type === "function") { + const name = "function" in choice ? choice.function?.name : choice.name; + return name ? { type: "function", function: { name } } : undefined; + } + return undefined; +} + function mapOptionsForApi( model: Model, options?: SimpleStreamOptions, @@ -287,12 +334,20 @@ function mapOptionsForApi( // Explicitly disable thinking when reasoning is not specified const reasoning = options?.reasoning; if (!reasoning) { - return { ...base, thinkingEnabled: false } satisfies AnthropicOptions; + return { + ...base, + thinkingEnabled: false, + toolChoice: mapAnthropicToolChoice(options?.toolChoice), + } satisfies AnthropicOptions; } let thinkingBudget = options.thinkingBudgets?.[reasoning] ?? ANTHROPIC_THINKING[reasoning]; if (thinkingBudget <= 0) { - return { ...base, thinkingEnabled: false } satisfies AnthropicOptions; + return { + ...base, + thinkingEnabled: false, + toolChoice: mapAnthropicToolChoice(options?.toolChoice), + } satisfies AnthropicOptions; } if (ANTHROPIC_USE_INTERLEAVED_THINKING) { @@ -300,6 +355,7 @@ function mapOptionsForApi( ...base, thinkingEnabled: true, thinkingBudgetTokens: thinkingBudget, + toolChoice: mapAnthropicToolChoice(options?.toolChoice), } satisfies AnthropicOptions; } @@ -313,13 +369,18 @@ function mapOptionsForApi( // If thinking budget is too low, disable thinking if (thinkingBudget <= 0) { - return { ...base, thinkingEnabled: false } satisfies AnthropicOptions; + return { + ...base, + thinkingEnabled: false, + toolChoice: mapAnthropicToolChoice(options?.toolChoice), + } satisfies AnthropicOptions; } else { return { ...base, maxTokens, thinkingEnabled: true, thinkingBudgetTokens: thinkingBudget, + toolChoice: mapAnthropicToolChoice(options?.toolChoice), } satisfies AnthropicOptions; } } @@ -329,6 +390,7 @@ function mapOptionsForApi( ...base, reasoning: options?.reasoning, thinkingBudgets: options?.thinkingBudgets, + toolChoice: mapAnthropicToolChoice(options?.toolChoice), }; const budgetInfo = resolveBedrockThinkingBudget(model as Model<"bedrock-converse-stream">, options); if (!budgetInfo) return bedrockBase as OptionsForApi; @@ -351,31 +413,39 @@ function mapOptionsForApi( return { ...base, reasoningEffort: supportsXhigh(model) ? options?.reasoning : clampReasoning(options?.reasoning), + toolChoice: mapOpenAiToolChoice(options?.toolChoice), } satisfies OpenAICompletionsOptions; case "openai-responses": return { ...base, reasoningEffort: supportsXhigh(model) ? options?.reasoning : clampReasoning(options?.reasoning), + toolChoice: mapOpenAiToolChoice(options?.toolChoice), } satisfies OpenAIResponsesOptions; case "azure-openai-responses": return { ...base, reasoningEffort: supportsXhigh(model) ? options?.reasoning : clampReasoning(options?.reasoning), + toolChoice: mapOpenAiToolChoice(options?.toolChoice), } satisfies AzureOpenAIResponsesOptions; case "openai-codex-responses": return { ...base, reasoningEffort: supportsXhigh(model) ? options?.reasoning : clampReasoning(options?.reasoning), + toolChoice: mapOpenAiToolChoice(options?.toolChoice), } satisfies OpenAICodexResponsesOptions; case "google-generative-ai": { // Explicitly disable thinking when reasoning is not specified // This is needed because Gemini has "dynamic thinking" enabled by default if (!options?.reasoning) { - return { ...base, thinking: { enabled: false } } satisfies GoogleOptions; + return { + ...base, + thinking: { enabled: false }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), + } satisfies GoogleOptions; } const googleModel = model as Model<"google-generative-ai">; @@ -390,6 +460,7 @@ function mapOptionsForApi( enabled: true, level: getGemini3ThinkingLevel(effort, googleModel), }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), } satisfies GoogleOptions; } @@ -399,12 +470,17 @@ function mapOptionsForApi( enabled: true, budgetTokens: getGoogleBudget(googleModel, effort, options?.thinkingBudgets), }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), } satisfies GoogleOptions; } case "google-gemini-cli": { if (!options?.reasoning) { - return { ...base, thinking: { enabled: false } } satisfies GoogleGeminiCliOptions; + return { + ...base, + thinking: { enabled: false }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), + } satisfies GoogleGeminiCliOptions; } const effort = clampReasoning(options.reasoning)!; @@ -417,6 +493,7 @@ function mapOptionsForApi( enabled: true, level: getGeminiCliThinkingLevel(effort, model.id), }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), } satisfies GoogleGeminiCliOptions; } @@ -432,12 +509,17 @@ function mapOptionsForApi( // If thinking budget is too low, disable thinking if (thinkingBudget <= 0) { - return { ...base, thinking: { enabled: false } } satisfies GoogleGeminiCliOptions; + return { + ...base, + thinking: { enabled: false }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), + } satisfies GoogleGeminiCliOptions; } else { return { ...base, maxTokens, thinking: { enabled: true, budgetTokens: thinkingBudget }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), } satisfies GoogleGeminiCliOptions; } } @@ -445,7 +527,11 @@ function mapOptionsForApi( case "google-vertex": { // Explicitly disable thinking when reasoning is not specified if (!options?.reasoning) { - return { ...base, thinking: { enabled: false } } satisfies GoogleVertexOptions; + return { + ...base, + thinking: { enabled: false }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), + } satisfies GoogleVertexOptions; } const vertexModel = model as Model<"google-vertex">; @@ -459,6 +545,7 @@ function mapOptionsForApi( enabled: true, level: getGemini3ThinkingLevel(effort, geminiModel), }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), } satisfies GoogleVertexOptions; } @@ -468,6 +555,7 @@ function mapOptionsForApi( enabled: true, budgetTokens: getGoogleBudget(geminiModel, effort, options?.thinkingBudgets), }, + toolChoice: mapGoogleToolChoice(options?.toolChoice), } satisfies GoogleVertexOptions; } diff --git a/packages/ai/src/types.ts b/packages/ai/src/types.ts index 85dfe12eb..204336b3c 100644 --- a/packages/ai/src/types.ts +++ b/packages/ai/src/types.ts @@ -94,6 +94,15 @@ export type ThinkingLevel = "minimal" | "low" | "medium" | "high" | "xhigh"; /** Token budgets for each thinking level (token-based providers only) */ export type ThinkingBudgets = { [key in ThinkingLevel]?: number }; +export type ToolChoice = + | "auto" + | "none" + | "any" + | "required" + | { type: "function"; name: string } + | { type: "function"; function: { name: string } } + | { type: "tool"; name: string }; + // Base options all providers share export interface StreamOptions { temperature?: number; @@ -129,6 +138,8 @@ export interface SimpleStreamOptions extends StreamOptions { cursorExecHandlers?: CursorExecHandlers; /** Hook to handle tool results from Cursor exec */ cursorOnToolResult?: CursorToolResultHandler; + /** Optional tool choice override for compatible providers */ + toolChoice?: ToolChoice; } // Generic StreamFunction with typed options @@ -298,6 +309,8 @@ export interface OpenAICompat { requiresReasoningContentForToolCalls?: boolean; /** Whether assistant tool-call messages must include non-empty content. Default: false. */ requiresAssistantContentForToolCalls?: boolean; + /** Whether the provider supports the `tool_choice` parameter. Default: true. */ + supportsToolChoice?: boolean; /** OpenRouter-specific routing preferences. Only used when baseUrl points to OpenRouter. */ openRouterRouting?: OpenRouterRouting; } diff --git a/packages/ai/test/openai-completions-tool-result-images.test.ts b/packages/ai/test/openai-completions-tool-result-images.test.ts index 85bed2deb..b4cf7e825 100644 --- a/packages/ai/test/openai-completions-tool-result-images.test.ts +++ b/packages/ai/test/openai-completions-tool-result-images.test.ts @@ -17,6 +17,7 @@ const compat: Required = { supportsDeveloperRole: true, supportsReasoningEffort: true, supportsUsageInStreaming: true, + supportsToolChoice: true, maxTokensField: "max_completion_tokens", requiresToolResultName: false, requiresAssistantAfterToolResult: false, diff --git a/packages/coding-agent/src/ipy/executor.ts b/packages/coding-agent/src/ipy/executor.ts index caa52b827..3378f33b8 100644 --- a/packages/coding-agent/src/ipy/executor.ts +++ b/packages/coding-agent/src/ipy/executor.ts @@ -255,10 +255,9 @@ async function createKernelSession( lastUsedAt: Date.now(), }; - session.heartbeatTimer = setInterval(async () => { + session.heartbeatTimer = setInterval(() => { if (session.dead) return; - const ok = await session.kernel.ping().catch(() => false); - if (!ok) { + if (!session.kernel.isAlive()) { session.dead = true; } }, 5000); diff --git a/packages/coding-agent/src/ipy/kernel.ts b/packages/coding-agent/src/ipy/kernel.ts index 9f18a3b49..a10c799d9 100644 --- a/packages/coding-agent/src/ipy/kernel.ts +++ b/packages/coding-agent/src/ipy/kernel.ts @@ -13,9 +13,6 @@ import { PYTHON_PRELUDE } from "./prelude"; const TEXT_ENCODER = new TextEncoder(); const TEXT_DECODER = new TextDecoder(); -const HEARTBEAT_INTERVAL_MS = 5000; -const HEARTBEAT_TIMEOUT_MS = 2000; -const HEARTBEAT_FAILURE_LIMIT = 1; const GATEWAY_STARTUP_TIMEOUT_MS = 30000; const GATEWAY_STARTUP_ATTEMPTS = 3; const TRACE_IPC = process.env.OMP_PYTHON_IPC_TRACE === "1"; @@ -465,8 +462,6 @@ export class PythonKernel { #ws: WebSocket | null = null; #disposed = false; #alive = true; - #heartbeatTimer?: NodeJS.Timeout; - #heartbeatFailures = 0; #messageHandlers = new Map void>(); #channelHandlers = new Map void>>(); #pendingExecutions = new Map void>(); @@ -562,7 +557,6 @@ export class PythonKernel { try { await kernel.connectWebSocket(); await kernel.initializeKernelEnvironment(cwd, env); - kernel.startHeartbeat(); const preludeResult = await kernel.execute(PYTHON_PRELUDE, { silent: true, storeHistory: false }); if (preludeResult.cancelled || preludeResult.status === "error") { throw new Error("Failed to initialize Python kernel prelude"); @@ -602,7 +596,6 @@ export class PythonKernel { time("startWithSharedGateway:connectWS"); await kernel.initializeKernelEnvironment(cwd, env); time("startWithSharedGateway:initEnv"); - kernel.startHeartbeat(); const preludeResult = await kernel.execute(PYTHON_PRELUDE, { silent: true, storeHistory: false }); time("startWithSharedGateway:prelude"); if (preludeResult.cancelled || preludeResult.status === "error") { @@ -721,7 +714,6 @@ export class PythonKernel { try { await kernel.connectWebSocket(); await kernel.initializeKernelEnvironment(options.cwd, options.env); - kernel.startHeartbeat(); const preludeResult = await kernel.execute(PYTHON_PRELUDE, { silent: true, storeHistory: false }); if (preludeResult.cancelled || preludeResult.status === "error") { throw new Error("Failed to initialize Python kernel prelude"); @@ -1108,11 +1100,6 @@ export class PythonKernel { this.#alive = false; this.abortPendingExecutions("Kernel shutdown"); - if (this.#heartbeatTimer) { - clearInterval(this.#heartbeatTimer); - this.#heartbeatTimer = undefined; - } - try { await fetch(`${this.gatewayUrl}/api/kernels/${this.kernelId}`, { method: "DELETE", @@ -1140,35 +1127,6 @@ export class PythonKernel { } } - async ping(timeoutMs: number = HEARTBEAT_TIMEOUT_MS): Promise { - if (!this.isAlive()) return false; - try { - const response = await fetch(`${this.gatewayUrl}/api/kernels/${this.kernelId}`, { - signal: AbortSignal.timeout(timeoutMs), - headers: this.#authHeaders(), - }); - if (response.ok) { - this.#heartbeatFailures = 0; - return true; - } - throw new Error(`Kernel status check failed: ${response.status}`); - } catch (err: unknown) { - this.#heartbeatFailures += 1; - if (this.#heartbeatFailures > HEARTBEAT_FAILURE_LIMIT) { - this.#alive = false; - logger.warn("Kernel heartbeat failed", { error: err instanceof Error ? err.message : String(err) }); - } - return false; - } - } - - private startHeartbeat(): void { - if (this.#heartbeatTimer) return; - this.#heartbeatTimer = setInterval(() => { - void this.ping(); - }, HEARTBEAT_INTERVAL_MS); - } - private renderDisplay(content: Record): { text: string; outputs: KernelDisplayOutput[] } { const data = content.data as Record | undefined; if (!data) return { text: "", outputs: [] }; diff --git a/packages/coding-agent/src/prompts/agents/reviewer.md b/packages/coding-agent/src/prompts/agents/reviewer.md index e9e49b9d7..32a24a3e9 100644 --- a/packages/coding-agent/src/prompts/agents/reviewer.md +++ b/packages/coding-agent/src/prompts/agents/reviewer.md @@ -115,6 +115,8 @@ Final `complete` call (payload goes under `data`): - `data.confidence`: 0.0-1.0 - `data.findings`: Optional; MUST omit (it is populated from `report_finding` calls) +Do not output JSON or code blocks. You must call the `complete` tool. + Correctness judgment ignores non-blocking issues (style, docs, nits). diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index d69076aa3..af1a4c569 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -23,6 +23,7 @@ import type { Model, TextContent, ToolCall, + ToolChoice, Usage, UsageReport, } from "@oh-my-pi/pi-ai"; @@ -157,6 +158,8 @@ export interface PromptOptions { images?: ImageContent[]; /** When streaming, how to queue the message: "steer" (interrupt) or "followUp" (wait). */ streamingBehavior?: "steer" | "followUp"; + /** Optional tool choice override for the next LLM call. */ + toolChoice?: ToolChoice; } /** Result from cycleModel() */ @@ -1252,7 +1255,8 @@ export class AgentSession { } } - await this.agent.prompt(messages); + const agentPromptOptions = options?.toolChoice ? { toolChoice: options.toolChoice } : undefined; + await this.agent.prompt(messages, agentPromptOptions); await this.waitForRetry(); } diff --git a/packages/coding-agent/src/task/executor.ts b/packages/coding-agent/src/task/executor.ts index 830c03ffb..73abad597 100644 --- a/packages/coding-agent/src/task/executor.ts +++ b/packages/coding-agent/src/task/executor.ts @@ -5,7 +5,7 @@ */ import path from "node:path"; import type { AgentEvent, ThinkingLevel } from "@oh-my-pi/pi-agent-core"; -import type { Api, Model } from "@oh-my-pi/pi-ai"; +import type { Api, Model, ToolChoice } from "@oh-my-pi/pi-ai"; import type { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; import { parseModelPattern } from "@oh-my-pi/pi-coding-agent/config/model-resolver"; import type { PromptTemplate } from "@oh-my-pi/pi-coding-agent/config/prompt-templates"; @@ -19,10 +19,12 @@ import type { AgentSession, AgentSessionEvent } from "@oh-my-pi/pi-coding-agent/ import type { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; import type { ContextFileEntry } from "@oh-my-pi/pi-coding-agent/tools"; +import { jtdToJsonSchema } from "@oh-my-pi/pi-coding-agent/tools/jtd-to-json-schema"; import { ToolAbortError } from "@oh-my-pi/pi-coding-agent/tools/tool-errors"; import type { EventBus } from "@oh-my-pi/pi-coding-agent/utils/event-bus"; import { logger, untilAborted } from "@oh-my-pi/pi-utils"; import type { TSchema } from "@sinclair/typebox"; +import Ajv, { type ValidateFunction } from "ajv"; import { subprocessToolRegistry } from "./subprocess-tool-registry"; import { type AgentDefinition, @@ -37,6 +39,7 @@ import { const DEFAULT_MODEL_ALIASES = new Set(["default", "pi/default", "omp/default"]); const MCP_CALL_TIMEOUT_MS = 60_000; +const ajv = new Ajv({ allErrors: true, strict: false }); /** Agent event types to forward for progress tracking. */ const agentEventTypes = new Set([ @@ -152,6 +155,22 @@ function resolveModelOverride( return {}; } +function buildCompleteToolChoice(model?: Model): ToolChoice | undefined { + if (!model) return undefined; + if ( + model.api === "openai-codex-responses" || + model.api === "openai-responses" || + model.api === "openai-completions" || + model.api === "azure-openai-responses" + ) { + return { type: "function", name: "complete" }; + } + if (model.api === "anthropic-messages" || model.api === "bedrock-converse-stream") { + return { type: "tool", name: "complete" }; + } + return undefined; +} + /** Options for subagent execution */ export interface ExecutorOptions { cwd: string; @@ -227,6 +246,89 @@ function truncateOutput(output: string): { text: string; truncated: boolean } { return { text: output, truncated }; } +function parseStringifiedJson(value: unknown): unknown { + if (typeof value !== "string") return value; + const trimmed = value.trim(); + if (!trimmed) return value; + if (!(trimmed.startsWith("{") || trimmed.startsWith("["))) return value; + try { + return JSON.parse(trimmed); + } catch { + return value; + } +} + +function normalizeOutputSchema(schema: unknown): { normalized?: unknown; error?: string } { + if (schema === undefined || schema === null) return {}; + if (typeof schema === "string") { + try { + return { normalized: JSON.parse(schema) }; + } catch (err) { + return { error: err instanceof Error ? err.message : String(err) }; + } + } + return { normalized: schema }; +} + +function buildOutputValidator(schema: unknown): { validate?: ValidateFunction; error?: string } { + const { normalized, error } = normalizeOutputSchema(schema); + if (error) return { error }; + if (normalized === undefined) return {}; + const jsonSchema = jtdToJsonSchema(normalized); + try { + return { validate: ajv.compile(jsonSchema as any) }; + } catch (err) { + return { error: err instanceof Error ? err.message : String(err) }; + } +} + +function tryParseJsonOutput(text: string): unknown | undefined { + const trimmed = text.trim(); + if (!trimmed) return undefined; + try { + return JSON.parse(trimmed); + } catch { + return undefined; + } +} + +function extractCompletionData(parsed: unknown): unknown { + if (!parsed || typeof parsed !== "object") return parsed; + const record = parsed as Record; + if ("data" in record) { + return record.data; + } + return parsed; +} + +function normalizeCompleteData(data: unknown, reportFindings?: ReviewFinding[]): unknown { + let normalized = parseStringifiedJson(data ?? null); + if ( + Array.isArray(reportFindings) && + reportFindings.length > 0 && + normalized && + typeof normalized === "object" && + !Array.isArray(normalized) + ) { + const record = normalized as Record; + if (!("findings" in record)) { + normalized = { ...record, findings: reportFindings }; + } + } + return normalized; +} + +function resolveFallbackCompletion(rawOutput: string, outputSchema: unknown): { data: unknown } | null { + const parsed = tryParseJsonOutput(rawOutput); + if (parsed === undefined) return null; + const candidate = parseStringifiedJson(extractCompletionData(parsed)); + if (candidate === undefined) return null; + const { validate, error } = buildOutputValidator(outputSchema); + if (error) return null; + if (validate && !validate(candidate)) return null; + return { data: candidate }; +} + /** * Extract a short preview from tool args for display. */ @@ -819,7 +921,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise + let previousTools: string[] | null = null; + try { + while (!completeCalled && retryCount < MAX_COMPLETE_RETRIES && !abortSignal.aborted) { + retryCount++; + if (!previousTools) { + previousTools = session.getActiveToolNames(); + await session.setActiveToolsByName(["complete"]); + } + const reminder = ` CRITICAL: You stopped without calling the complete tool. This is reminder ${retryCount} of ${MAX_COMPLETE_RETRIES}. You MUST call the complete tool to finish your task. Options: @@ -959,7 +1069,12 @@ Failure to call complete after ${MAX_COMPLETE_RETRIES} reminders will result in Call complete now.`; - await session.prompt(reminder); + await session.prompt(reminder, reminderToolChoice ? { toolChoice: reminderToolChoice } : undefined); + } + } finally { + if (previousTools) { + await session.setActiveToolsByName(previousTools); + } } const lastMessage = session.state.messages[session.state.messages.length - 1]; @@ -1026,6 +1141,7 @@ Call complete now.`; const completeItems = progress.extractedToolData?.complete as | Array<{ data?: unknown; status?: "success" | "aborted"; error?: string }> | undefined; + const reportFindings = progress.extractedToolData?.report_finding as ReviewFinding[] | undefined; const hasComplete = Array.isArray(completeItems) && completeItems.length > 0; if (hasComplete) { const lastComplete = completeItems[completeItems.length - 1]; @@ -1041,29 +1157,7 @@ Call complete now.`; } } else { // Normal successful completion - let completeData = lastComplete?.data ?? null; - // Handle double-stringified JSON (subagent returned JSON string instead of object) - if (typeof completeData === "string" && (completeData.startsWith("{") || completeData.startsWith("["))) { - try { - completeData = JSON.parse(completeData); - } catch { - // Not valid JSON, keep as string - } - } - // Special case: merge report_finding data into review output for parent visibility - const reportFindings = progress.extractedToolData?.report_finding as ReviewFinding[] | undefined; - if ( - Array.isArray(reportFindings) && - reportFindings.length > 0 && - completeData && - typeof completeData === "object" && - !Array.isArray(completeData) - ) { - const record = completeData as Record; - if (!("findings" in record)) { - completeData = { ...record, findings: reportFindings }; - } - } + const completeData = normalizeCompleteData(lastComplete?.data ?? null, reportFindings); try { rawOutput = JSON.stringify(completeData, null, 2) ?? "null"; } catch (err) { @@ -1074,8 +1168,27 @@ Call complete now.`; stderr = ""; } } else { - const warning = "SYSTEM WARNING: Subagent exited without calling complete tool after 3 reminders."; - rawOutput = rawOutput ? `${warning}\n\n${rawOutput}` : warning; + const allowFallback = !done.aborted && !signal?.aborted; + const { normalized: normalizedSchema, error: schemaError } = normalizeOutputSchema(outputSchema); + const hasOutputSchema = normalizedSchema !== undefined && !schemaError; + const fallback = allowFallback ? resolveFallbackCompletion(rawOutput, outputSchema) : null; + if (fallback) { + const completeData = normalizeCompleteData(fallback.data, reportFindings); + try { + rawOutput = JSON.stringify(completeData, null, 2) ?? "null"; + } catch (err) { + const errorMessage = err instanceof Error ? err.message : String(err); + rawOutput = `{"error":"Failed to serialize fallback completion: ${errorMessage}"}`; + } + exitCode = 0; + stderr = ""; + } else if (!hasOutputSchema && allowFallback && rawOutput.trim().length > 0) { + exitCode = 0; + stderr = ""; + } else { + const warning = "SYSTEM WARNING: Subagent exited without calling complete tool after 3 reminders."; + rawOutput = rawOutput ? `${warning}\n\n${rawOutput}` : warning; + } } const { text: truncatedOutput, truncated } = truncateOutput(rawOutput); diff --git a/packages/coding-agent/src/task/render.ts b/packages/coding-agent/src/task/render.ts index a40b2d5e0..43510530a 100644 --- a/packages/coding-agent/src/task/render.ts +++ b/packages/coding-agent/src/task/render.ts @@ -77,8 +77,7 @@ function formatJsonScalar(value: unknown, theme: Theme): string { return ""; } -const MISSING_COMPLETE_WARNING_PREFIX = - "SYSTEM WARNING: Subagent exited without calling complete tool"; +const MISSING_COMPLETE_WARNING_PREFIX = "SYSTEM WARNING: Subagent exited without calling complete tool"; function extractMissingCompleteWarning(output: string): { warning?: string; rest: string } { const lines = output.split("\n"); @@ -89,7 +88,7 @@ function extractMissingCompleteWarning(output: string): { warning?: string; rest const rest = lines .slice(1) .join("\n") - .replace(/^\s*\n+/, "") + .replace(/^\s*\n+/, ""); return { warning: firstLine, rest }; } @@ -294,10 +293,7 @@ function renderOutputSection( if (outputLines.length > previewCount) { lines.push( - `${continuePrefix} ${theme.fg( - "dim", - formatMoreItems(outputLines.length - previewCount, "line", theme), - )}`, + `${continuePrefix} ${theme.fg("dim", formatMoreItems(outputLines.length - previewCount, "line", theme))}`, ); } @@ -736,8 +732,7 @@ function renderAgentResult(result: SingleResult, isLast: boolean, expanded: bool const prefix = isLast ? theme.fg("dim", theme.tree.last) : theme.fg("dim", theme.tree.branch); const continuePrefix = isLast ? " " : `${theme.fg("dim", theme.tree.vertical)} `; - const { warning: missingCompleteWarning, rest: outputWithoutWarning } = - extractMissingCompleteWarning(result.output); + const { warning: missingCompleteWarning, rest: outputWithoutWarning } = extractMissingCompleteWarning(result.output); const aborted = result.aborted ?? false; const success = !aborted && result.exitCode === 0; const needsWarning = Boolean(missingCompleteWarning) && success; @@ -841,15 +836,7 @@ function renderAgentResult(result: SingleResult, isLast: boolean, expanded: bool // Fallback to output preview if no custom rendering if (!hasCustomRendering) { lines.push( - ...renderOutputSection( - outputWithoutWarning, - continuePrefix, - expanded, - theme, - 3, - 12, - missingCompleteWarning, - ), + ...renderOutputSection(outputWithoutWarning, continuePrefix, expanded, theme, 3, 12, missingCompleteWarning), ); } diff --git a/packages/coding-agent/test/core/python-kernel.test.ts b/packages/coding-agent/test/core/python-kernel.test.ts index f029d1df6..acad1a857 100644 --- a/packages/coding-agent/test/core/python-kernel.test.ts +++ b/packages/coding-agent/test/core/python-kernel.test.ts @@ -302,51 +302,6 @@ describe("PythonKernel (external gateway)", () => { }); }); - it("marks kernel dead after repeated ping failures", async () => { - const fetchMock = vi.fn(async (url: string, init?: RequestInit) => { - if (url.endsWith("/api/kernels") && init?.method === "POST") { - return new Response(JSON.stringify({ id: "kernel-2" }), { status: 201 }); - } - if (url.includes("/api/kernels/kernel-2") && !init?.method) { - throw new Error("ping failed"); - } - return new Response("", { status: 200 }); - }); - globalThis.fetch = fetchMock as unknown as typeof fetch; - - let initSeen = false; - let preludeSeen = false; - - const kernelPromise = PythonKernel.start({ cwd: "/" }); - await Bun.sleep(10); - const ws = FakeWebSocket.lastInstance; - if (!ws) throw new Error("WebSocket not initialized"); - ws.setSendHandler(data => { - const msg = typeof data === "string" ? (JSON.parse(data) as JupyterMessage) : decodeMessage(data); - const code = String(msg.content.code ?? ""); - if (!initSeen) { - initSeen = true; - sendOkExecution(ws, msg.header.msg_id); - return; - } - if (!preludeSeen) { - expect(code).toBe(PYTHON_PRELUDE); - preludeSeen = true; - } - sendOkExecution(ws, msg.header.msg_id); - }); - - const kernel = await kernelPromise; - const firstPing = await kernel.ping(1); - const secondPing = await kernel.ping(1); - - expect(firstPing).toBe(false); - expect(secondPing).toBe(false); - expect(kernel.isAlive()).toBe(false); - - await kernel.shutdown(); - }); - it("initializes the IPython prelude", async () => { const fetchMock = vi.fn(async (url: string, init?: RequestInit) => { if (url.endsWith("/api/kernels") && init?.method === "POST") { diff --git a/packages/tui/src/tui.ts b/packages/tui/src/tui.ts index 64475cf26..85856dc98 100644 --- a/packages/tui/src/tui.ts +++ b/packages/tui/src/tui.ts @@ -770,10 +770,13 @@ export class TUI extends Container { /** * Find and extract cursor position from rendered lines. * Searches for CURSOR_MARKER, calculates its position, and strips it from the output. + * @param lines - Rendered lines to search + * @param height - Terminal height (to calculate viewport) * @returns Cursor position { row, col } or null if no marker found */ - private extractCursorPosition(lines: string[]): { row: number; col: number } | null { - for (let row = 0; row < lines.length; row++) { + private extractCursorPosition(lines: string[], height: number): { row: number; col: number } | null { + const viewportTop = Math.max(0, lines.length - height); + for (let row = lines.length - 1; row >= viewportTop; row--) { const line = lines[row]; const markerIndex = line.indexOf(CURSOR_MARKER); if (markerIndex !== -1) { @@ -913,7 +916,7 @@ export class TUI extends Container { } // Extract cursor position before applying line resets (marker must be found first) - const cursorPos = this.extractCursorPosition(newLines); + const cursorPos = this.extractCursorPosition(newLines, height); newLines = this.applyLineResets(newLines); diff --git a/packages/tui/src/utils.ts b/packages/tui/src/utils.ts index b732904d9..755beed52 100644 --- a/packages/tui/src/utils.ts +++ b/packages/tui/src/utils.ts @@ -303,6 +303,8 @@ class AnsiCodeTracker { } } +const WRAP_OPTIONS = { wordWrap: true, hard: true, trim: false } as const; + /** * Wrap text with ANSI codes preserved. * @@ -315,11 +317,7 @@ class AnsiCodeTracker { * @returns Array of wrapped lines (NOT padded to width) */ export function wrapTextWithAnsi(text: string, width: number): string[] { - if (!text) { - return [""]; - } - - return Bun.wrapAnsi(text, width, { wordWrap: true, hard: true, trim: false }).split("\n"); + return Bun.wrapAnsi(text, width, WRAP_OPTIONS).split("\n"); } const PUNCTUATION_REGEX = /[(){}[\]<>.,;:'"!?+\-=*/\\|&%^$#@~`]/;