diff --git a/packages/agent/src/types.ts b/packages/agent/src/types.ts index 7d79f4056..79e435397 100644 --- a/packages/agent/src/types.ts +++ b/packages/agent/src/types.ts @@ -169,7 +169,7 @@ export interface AgentToolResult { // Content blocks supporting text and images content: (TextContent | ImageContent)[]; // Details to be displayed in a UI or logged - details: T; + details?: T; } // Callback for streaming tool execution updates diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 09b71b1f4..a74c060e9 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,6 +1,7 @@ # Changelog ## [Unreleased] + ### Added - Added configurable fuzzy matching threshold for edit operations @@ -13,6 +14,11 @@ ### Changed +- Converted tool implementations from factory functions to class-based architecture for better modularity +- Updated tool exports to expose classes instead of factory functions +- Refactored tool instantiation patterns across all built-in tools +- Updated test files to use new class-based tool constructors +- Modified SDK exports to provide tool classes for custom usage - Enhanced patch application with better ambiguity detection and clearer error messages - Improved diff truncation algorithm to preserve context around changes more intelligently - Updated patch mode prompts to prefer larger semantic blocks over single-line edits diff --git a/packages/coding-agent/src/core/bash-executor.ts b/packages/coding-agent/src/core/bash-executor.ts index ce60105c8..88b652909 100644 --- a/packages/coding-agent/src/core/bash-executor.ts +++ b/packages/coding-agent/src/core/bash-executor.ts @@ -9,7 +9,7 @@ import type { Subprocess } from "bun"; import { getShellConfig, killProcessTree } from "../utils/shell"; import { getOrCreateSnapshot, getSnapshotSourceCommand } from "../utils/shell-snapshot"; -import { createOutputSink, pumpStream } from "./streaming-output"; +import { OutputSink, pumpStream } from "./streaming-output"; import type { BashOperations } from "./tools/bash"; import { DEFAULT_MAX_BYTES } from "./tools/truncate"; import { ScopeSignal } from "./utils"; @@ -85,7 +85,7 @@ export async function executeBash(command: string, options?: BashExecutorOptions killProcessTree(child.pid); }); - const sink = createOutputSink(DEFAULT_MAX_BYTES, DEFAULT_MAX_BYTES * 2, options?.onChunk); + const sink = new OutputSink(DEFAULT_MAX_BYTES, DEFAULT_MAX_BYTES * 2, options?.onChunk); const writer = sink.getWriter(); try { @@ -128,7 +128,7 @@ export async function executeBashWithOperations( operations: BashOperations, options?: BashExecutorOptions, ): Promise { - const sink = createOutputSink(DEFAULT_MAX_BYTES, DEFAULT_MAX_BYTES * 2, options?.onChunk); + const sink = new OutputSink(DEFAULT_MAX_BYTES, DEFAULT_MAX_BYTES * 2, options?.onChunk); const writer = sink.getWriter(); // Create a ReadableStream from the callback-based operations.exec diff --git a/packages/coding-agent/src/core/cursor/exec-bridge.ts b/packages/coding-agent/src/core/cursor/exec-bridge.ts index b32e9a427..7b57c35f7 100644 --- a/packages/coding-agent/src/core/cursor/exec-bridge.ts +++ b/packages/coding-agent/src/core/cursor/exec-bridge.ts @@ -7,7 +7,7 @@ import type { AgentToolResult, AgentToolUpdateCallback, } from "@oh-my-pi/pi-agent-core"; -import type { CursorExecHandlers, CursorMcpCall, ToolResultMessage } from "@oh-my-pi/pi-ai"; +import type { CursorMcpCall, CursorExecHandlers as ICursorExecHandlers, ToolResultMessage } from "@oh-my-pi/pi-ai"; import { resolveToCwd } from "../tools/path-utils"; interface CursorExecBridgeOptions { @@ -143,92 +143,99 @@ function formatMcpToolErrorMessage(toolName: string, availableTools: string[]): return `MCP tool "${toolName}" not found. Available tools: ${list}`; } -export function createCursorExecHandlers(options: CursorExecBridgeOptions): CursorExecHandlers { - return { - read: async (args) => { - const toolCallId = decodeToolCallId(args.toolCallId); - const toolResultMessage = await executeTool(options, "read", toolCallId, { path: args.path }); - return toolResultMessage; - }, - ls: async (args) => { - const toolCallId = decodeToolCallId(args.toolCallId); - const toolResultMessage = await executeTool(options, "ls", toolCallId, { path: args.path }); - return toolResultMessage; - }, - grep: async (args) => { - const toolCallId = decodeToolCallId(args.toolCallId); - const toolResultMessage = await executeTool(options, "grep", toolCallId, { - pattern: args.pattern, - path: args.path || undefined, - glob: args.glob || undefined, - outputMode: args.outputMode || undefined, - context: args.context ?? args.contextBefore ?? args.contextAfter ?? undefined, - ignoreCase: args.caseInsensitive || undefined, - type: args.type || undefined, - headLimit: args.headLimit ?? undefined, - multiline: args.multiline || undefined, - }); - return toolResultMessage; - }, - write: async (args) => { - const toolCallId = decodeToolCallId(args.toolCallId); - const content = args.fileText ?? new TextDecoder().decode(args.fileBytes ?? new Uint8Array()); - const toolResultMessage = await executeTool(options, "write", toolCallId, { - path: args.path, - content, - }); - return toolResultMessage; - }, - delete: async (args) => { - const toolCallId = decodeToolCallId(args.toolCallId); - const toolResultMessage = await executeDelete(options, args.path, toolCallId); - return toolResultMessage; - }, - shell: async (args) => { - const toolCallId = decodeToolCallId(args.toolCallId); - const timeoutSeconds = - args.timeout && args.timeout > 0 - ? args.timeout > 1000 - ? Math.ceil(args.timeout / 1000) - : args.timeout - : undefined; - const toolResultMessage = await executeTool(options, "bash", toolCallId, { - command: args.command, - workdir: args.workingDirectory || undefined, - timeout: timeoutSeconds, - }); - return toolResultMessage; - }, - diagnostics: async (args) => { - const toolCallId = decodeToolCallId(args.toolCallId); - const toolResultMessage = await executeTool(options, "lsp", toolCallId, { - action: "diagnostics", - file: args.path, - }); - return toolResultMessage; - }, - mcp: async (call: CursorMcpCall) => { - const toolName = call.toolName || call.name; - const toolCallId = decodeToolCallId(call.toolCallId); - const tool = options.tools.get(toolName); - if (!tool) { - const availableTools = Array.from(options.tools.keys()).filter((name) => name.startsWith("mcp_")); - const message = formatMcpToolErrorMessage(toolName, availableTools); - const toolResult: ToolResultMessage = { - role: "toolResult", - toolCallId, - toolName, - content: [{ type: "text", text: message }], - details: {}, - isError: true, - timestamp: Date.now(), - }; - return toolResult; - } +export class CursorExecHandlers implements ICursorExecHandlers { + constructor(private options: CursorExecBridgeOptions) {} - const args = Object.keys(call.args ?? {}).length > 0 ? call.args : decodeMcpArgs(call.rawArgs ?? {}); - const toolResultMessage = await executeTool(options, toolName, toolCallId, args); - return toolResultMessage; - }, - }; + async read(args: Parameters>[0]) { + const toolCallId = decodeToolCallId(args.toolCallId); + const toolResultMessage = await executeTool(this.options, "read", toolCallId, { path: args.path }); + return toolResultMessage; + } + + async ls(args: Parameters>[0]) { + const toolCallId = decodeToolCallId(args.toolCallId); + const toolResultMessage = await executeTool(this.options, "ls", toolCallId, { path: args.path }); + return toolResultMessage; + } + + async grep(args: Parameters>[0]) { + const toolCallId = decodeToolCallId(args.toolCallId); + const toolResultMessage = await executeTool(this.options, "grep", toolCallId, { + pattern: args.pattern, + path: args.path || undefined, + glob: args.glob || undefined, + outputMode: args.outputMode || undefined, + context: args.context ?? args.contextBefore ?? args.contextAfter ?? undefined, + ignoreCase: args.caseInsensitive || undefined, + type: args.type || undefined, + headLimit: args.headLimit ?? undefined, + multiline: args.multiline || undefined, + }); + return toolResultMessage; + } + + async write(args: Parameters>[0]) { + const toolCallId = decodeToolCallId(args.toolCallId); + const content = args.fileText ?? new TextDecoder().decode(args.fileBytes ?? new Uint8Array()); + const toolResultMessage = await executeTool(this.options, "write", toolCallId, { + path: args.path, + content, + }); + return toolResultMessage; + } + + async delete(args: Parameters>[0]) { + const toolCallId = decodeToolCallId(args.toolCallId); + const toolResultMessage = await executeDelete(this.options, args.path, toolCallId); + return toolResultMessage; + } + + async shell(args: Parameters>[0]) { + const toolCallId = decodeToolCallId(args.toolCallId); + const timeoutSeconds = + args.timeout && args.timeout > 0 + ? args.timeout > 1000 + ? Math.ceil(args.timeout / 1000) + : args.timeout + : undefined; + const toolResultMessage = await executeTool(this.options, "bash", toolCallId, { + command: args.command, + workdir: args.workingDirectory || undefined, + timeout: timeoutSeconds, + }); + return toolResultMessage; + } + + async diagnostics(args: Parameters>[0]) { + const toolCallId = decodeToolCallId(args.toolCallId); + const toolResultMessage = await executeTool(this.options, "lsp", toolCallId, { + action: "diagnostics", + file: args.path, + }); + return toolResultMessage; + } + + async mcp(call: CursorMcpCall) { + const toolName = call.toolName || call.name; + const toolCallId = decodeToolCallId(call.toolCallId); + const tool = this.options.tools.get(toolName); + if (!tool) { + const availableTools = Array.from(this.options.tools.keys()).filter((name) => name.startsWith("mcp_")); + const message = formatMcpToolErrorMessage(toolName, availableTools); + const toolResult: ToolResultMessage = { + role: "toolResult", + toolCallId, + toolName, + content: [{ type: "text", text: message }], + details: {}, + isError: true, + timestamp: Date.now(), + }; + return toolResult; + } + + const args = Object.keys(call.args ?? {}).length > 0 ? call.args : decodeMcpArgs(call.rawArgs ?? {}); + const toolResultMessage = await executeTool(this.options, toolName, toolCallId, args); + return toolResultMessage; + } } diff --git a/packages/coding-agent/src/core/custom-commands/bundled/review/index.ts b/packages/coding-agent/src/core/custom-commands/bundled/review/index.ts index 64e2482df..d72a009f1 100644 --- a/packages/coding-agent/src/core/custom-commands/bundled/review/index.ts +++ b/packages/coding-agent/src/core/custom-commands/bundled/review/index.ts @@ -224,154 +224,152 @@ function buildReviewPrompt(mode: string, stats: DiffStats, rawDiff: string): str }); } -export function createReviewCommand(api: CustomCommandAPI): CustomCommand { - return { - name: "review", - description: "Launch interactive code review", +export class ReviewCommand implements CustomCommand { + name = "review"; + description = "Launch interactive code review"; - async execute(_args: string[], ctx: HookCommandContext): Promise { - if (!ctx.hasUI) { - return "Use the Task tool to run the 'reviewer' agent to review recent code changes."; + constructor(private api: CustomCommandAPI) {} + + async execute(_args: string[], ctx: HookCommandContext): Promise { + if (!ctx.hasUI) { + return "Use the Task tool to run the 'reviewer' agent to review recent code changes."; + } + + const mode = await ctx.ui.select("Review Mode", [ + "1. Review against a base branch (PR Style)", + "2. Review uncommitted changes", + "3. Review a specific commit", + "4. Custom review instructions", + ]); + + if (!mode) return undefined; + + const modeNum = parseInt(mode[0], 10); + + switch (modeNum) { + case 1: { + // PR-style review against base branch + const branches = await getGitBranches(this.api); + if (branches.length === 0) { + ctx.ui.notify("No git branches found", "error"); + return undefined; + } + + const baseBranch = await ctx.ui.select("Select base branch to compare against", branches); + if (!baseBranch) return undefined; + + const currentBranch = await getCurrentBranch(this.api); + const diffResult = await this.api.exec("git", ["diff", `${baseBranch}...${currentBranch}`], { + timeout: 30000, + }); + if (diffResult.code !== 0) { + ctx.ui.notify(`Failed to get diff: ${diffResult.stderr}`, "error"); + return undefined; + } + + if (!diffResult.stdout.trim()) { + ctx.ui.notify(`No changes between ${baseBranch} and ${currentBranch}`, "warning"); + return undefined; + } + + const stats = parseDiff(diffResult.stdout); + if (stats.files.length === 0) { + ctx.ui.notify("No reviewable files (all changes filtered out)", "warning"); + return undefined; + } + + return buildReviewPrompt( + `Reviewing changes between \`${baseBranch}\` and \`${currentBranch}\` (PR-style)`, + stats, + diffResult.stdout, + ); } - const mode = await ctx.ui.select("Review Mode", [ - "1. Review against a base branch (PR Style)", - "2. Review uncommitted changes", - "3. Review a specific commit", - "4. Custom review instructions", - ]); + case 2: { + // Uncommitted changes - combine staged and unstaged + const status = await getGitStatus(this.api); + if (!status.trim()) { + ctx.ui.notify("No uncommitted changes found", "warning"); + return undefined; + } - if (!mode) return undefined; + const [unstagedResult, stagedResult] = await Promise.all([ + this.api.exec("git", ["diff"], { timeout: 30000 }), + this.api.exec("git", ["diff", "--cached"], { timeout: 30000 }), + ]); - const modeNum = parseInt(mode[0], 10); + const combinedDiff = [unstagedResult.stdout, stagedResult.stdout].filter(Boolean).join("\n"); - switch (modeNum) { - case 1: { - // PR-style review against base branch - const branches = await getGitBranches(api); - if (branches.length === 0) { - ctx.ui.notify("No git branches found", "error"); - return undefined; - } + if (!combinedDiff.trim()) { + ctx.ui.notify("No diff content found", "warning"); + return undefined; + } - const baseBranch = await ctx.ui.select("Select base branch to compare against", branches); - if (!baseBranch) return undefined; + const stats = parseDiff(combinedDiff); + if (stats.files.length === 0) { + ctx.ui.notify("No reviewable files (all changes filtered out)", "warning"); + return undefined; + } - const currentBranch = await getCurrentBranch(api); - const diffResult = await api.exec("git", ["diff", `${baseBranch}...${currentBranch}`], { - timeout: 30000, - }); - if (diffResult.code !== 0) { - ctx.ui.notify(`Failed to get diff: ${diffResult.stderr}`, "error"); - return undefined; - } + return buildReviewPrompt("Reviewing uncommitted changes (staged + unstaged)", stats, combinedDiff); + } - if (!diffResult.stdout.trim()) { - ctx.ui.notify(`No changes between ${baseBranch} and ${currentBranch}`, "warning"); - return undefined; - } + case 3: { + // Specific commit + const commits = await getRecentCommits(this.api, 20); + if (commits.length === 0) { + ctx.ui.notify("No commits found", "error"); + return undefined; + } + const selected = await ctx.ui.select("Select commit to review", commits); + if (!selected) return undefined; + + // Extract commit hash from selection (format: "abc1234 message") + const hash = selected.split(" ")[0]; + + // Get the commit diff (with timeout) + const showResult = await this.api.exec("git", ["show", "--format=", hash], { timeout: 30000 }); + if (showResult.code !== 0) { + ctx.ui.notify(`Failed to get commit: ${showResult.stderr}`, "error"); + return undefined; + } + + if (!showResult.stdout.trim()) { + ctx.ui.notify("Commit has no diff content", "warning"); + return undefined; + } + + const stats = parseDiff(showResult.stdout); + if (stats.files.length === 0) { + ctx.ui.notify("No reviewable files in commit (all changes filtered out)", "warning"); + return undefined; + } + + return buildReviewPrompt(`Reviewing commit \`${hash}\``, stats, showResult.stdout); + } + + case 4: { + // Custom instructions - still uses the old approach since user provides context + const instructions = await ctx.ui.editor("Enter custom review instructions", "Review the following:\n\n"); + if (!instructions?.trim()) return undefined; + + // For custom, we still try to get current diff for context + const diffResult = await this.api.exec("git", ["diff", "HEAD"], { timeout: 30000 }); + const hasDiff = diffResult.code === 0 && diffResult.stdout.trim(); + + if (hasDiff) { const stats = parseDiff(diffResult.stdout); - if (stats.files.length === 0) { - ctx.ui.notify("No reviewable files (all changes filtered out)", "warning"); - return undefined; - } - - return buildReviewPrompt( - `Reviewing changes between \`${baseBranch}\` and \`${currentBranch}\` (PR-style)`, + // Even if all files filtered, include the custom instructions + return `${buildReviewPrompt( + `Custom review: ${instructions.split("\n")[0].slice(0, 60)}...`, stats, diffResult.stdout, - ); + )}\n\n### Additional Instructions\n\n${instructions}`; } - case 2: { - // Uncommitted changes - combine staged and unstaged - const status = await getGitStatus(api); - if (!status.trim()) { - ctx.ui.notify("No uncommitted changes found", "warning"); - return undefined; - } - - const [unstagedResult, stagedResult] = await Promise.all([ - api.exec("git", ["diff"], { timeout: 30000 }), - api.exec("git", ["diff", "--cached"], { timeout: 30000 }), - ]); - - const combinedDiff = [unstagedResult.stdout, stagedResult.stdout].filter(Boolean).join("\n"); - - if (!combinedDiff.trim()) { - ctx.ui.notify("No diff content found", "warning"); - return undefined; - } - - const stats = parseDiff(combinedDiff); - if (stats.files.length === 0) { - ctx.ui.notify("No reviewable files (all changes filtered out)", "warning"); - return undefined; - } - - return buildReviewPrompt("Reviewing uncommitted changes (staged + unstaged)", stats, combinedDiff); - } - - case 3: { - // Specific commit - const commits = await getRecentCommits(api, 20); - if (commits.length === 0) { - ctx.ui.notify("No commits found", "error"); - return undefined; - } - - const selected = await ctx.ui.select("Select commit to review", commits); - if (!selected) return undefined; - - // Extract commit hash from selection (format: "abc1234 message") - const hash = selected.split(" ")[0]; - - // Get the commit diff (with timeout) - const showResult = await api.exec("git", ["show", "--format=", hash], { timeout: 30000 }); - if (showResult.code !== 0) { - ctx.ui.notify(`Failed to get commit: ${showResult.stderr}`, "error"); - return undefined; - } - - if (!showResult.stdout.trim()) { - ctx.ui.notify("Commit has no diff content", "warning"); - return undefined; - } - - const stats = parseDiff(showResult.stdout); - if (stats.files.length === 0) { - ctx.ui.notify("No reviewable files in commit (all changes filtered out)", "warning"); - return undefined; - } - - return buildReviewPrompt(`Reviewing commit \`${hash}\``, stats, showResult.stdout); - } - - case 4: { - // Custom instructions - still uses the old approach since user provides context - const instructions = await ctx.ui.editor( - "Enter custom review instructions", - "Review the following:\n\n", - ); - if (!instructions?.trim()) return undefined; - - // For custom, we still try to get current diff for context - const diffResult = await api.exec("git", ["diff", "HEAD"], { timeout: 30000 }); - const hasDiff = diffResult.code === 0 && diffResult.stdout.trim(); - - if (hasDiff) { - const stats = parseDiff(diffResult.stdout); - // Even if all files filtered, include the custom instructions - return `${buildReviewPrompt( - `Custom review: ${instructions.split("\n")[0].slice(0, 60)}...`, - stats, - diffResult.stdout, - )}\n\n### Additional Instructions\n\n${instructions}`; - } - - // No diff available, just pass instructions - return `## Code Review Request + // No diff available, just pass instructions + return `## Code Review Request ### Mode Custom review instructions @@ -381,13 +379,12 @@ Custom review instructions ${instructions} Use the Task tool with \`agent: "reviewer"\` to execute this review.`; - } - - default: - return undefined; } - }, - }; + + default: + return undefined; + } + } } async function getGitBranches(api: CustomCommandAPI): Promise { @@ -434,4 +431,4 @@ async function getRecentCommits(api: CustomCommandAPI, count: number): Promise { - if (args.length === 0) return formatUsage(); +export class WorktreeCommand implements CustomCommand { + name = "wt"; + description = "Git worktree management"; - const subcommand = args[0]; - const rest = args.slice(1); + // biome-ignore lint/complexity/noUselessConstructor: interface conformance - loader passes API to all commands + constructor(_api: CustomCommandAPI) {} - try { - switch (subcommand) { - case "new": { - const parsed = parseFlags(rest); - const branch = parsed.positionals[0]; - if (!branch) return formatUsage(); - const base = getFlagValue(parsed.flags, "base"); - if (parsed.flags.get("base") === true) { - return "Missing value for --base"; - } - return await handleNew({ branch, base }); + async execute(args: string[], ctx: HookCommandContext): Promise { + if (args.length === 0) return formatUsage(); + + const subcommand = args[0]; + const rest = args.slice(1); + + try { + switch (subcommand) { + case "new": { + const parsed = parseFlags(rest); + const branch = parsed.positionals[0]; + if (!branch) return formatUsage(); + const base = getFlagValue(parsed.flags, "base"); + if (parsed.flags.get("base") === true) { + return "Missing value for --base"; } - case "list": - return await handleList(ctx); - case "merge": { - const parsed = parseFlags(rest); - const source = parsed.positionals[0]; - const target = parsed.positionals[1]; - if (!source) return formatUsage(); - const strategyRaw = getFlagValue(parsed.flags, "strategy"); - if (parsed.flags.get("strategy") === true) { - return "Missing value for --strategy"; - } - const strategy = strategyRaw as CollapseStrategy | undefined; - const keep = getFlagBoolean(parsed.flags, "keep"); - return await handleMerge({ source, target, strategy, keep }); - } - case "rm": { - const parsed = parseFlags(rest); - const name = parsed.positionals[0]; - if (!name) return formatUsage(); - const force = getFlagBoolean(parsed.flags, "force"); - return await handleRm({ name, force }); - } - case "status": - return await handleStatus(); - case "spawn": { - const parsed = parseFlags(rest); - const task = parsed.positionals[0]; - if (!task) return formatUsage(); - const scope = getFlagValue(parsed.flags, "scope"); - if (parsed.flags.get("scope") === true) { - return "Missing value for --scope"; - } - const name = getFlagValue(parsed.flags, "name"); - return await handleSpawn({ task, scope, name }, ctx); - } - case "parallel": { - const tasks = parseParallelTasks(rest); - if (tasks.length === 0) return formatUsage(); - return await handleParallel(tasks, ctx); - } - default: - return formatUsage(); + return await handleNew({ branch, base }); } - } catch (err) { - return formatError(err); + case "list": + return await handleList(ctx); + case "merge": { + const parsed = parseFlags(rest); + const source = parsed.positionals[0]; + const target = parsed.positionals[1]; + if (!source) return formatUsage(); + const strategyRaw = getFlagValue(parsed.flags, "strategy"); + if (parsed.flags.get("strategy") === true) { + return "Missing value for --strategy"; + } + const strategy = strategyRaw as CollapseStrategy | undefined; + const keep = getFlagBoolean(parsed.flags, "keep"); + return await handleMerge({ source, target, strategy, keep }); + } + case "rm": { + const parsed = parseFlags(rest); + const name = parsed.positionals[0]; + if (!name) return formatUsage(); + const force = getFlagBoolean(parsed.flags, "force"); + return await handleRm({ name, force }); + } + case "status": + return await handleStatus(); + case "spawn": { + const parsed = parseFlags(rest); + const task = parsed.positionals[0]; + if (!task) return formatUsage(); + const scope = getFlagValue(parsed.flags, "scope"); + if (parsed.flags.get("scope") === true) { + return "Missing value for --scope"; + } + const name = getFlagValue(parsed.flags, "name"); + return await handleSpawn({ task, scope, name }, ctx); + } + case "parallel": { + const tasks = parseParallelTasks(rest); + if (tasks.length === 0) return formatUsage(); + return await handleParallel(tasks, ctx); + } + default: + return formatUsage(); } - }, - }; + } catch (err) { + return formatError(err); + } + } } -export default createWorktreeCommand; +export default WorktreeCommand; diff --git a/packages/coding-agent/src/core/custom-commands/loader.ts b/packages/coding-agent/src/core/custom-commands/loader.ts index c2f2504a3..267cc6da4 100644 --- a/packages/coding-agent/src/core/custom-commands/loader.ts +++ b/packages/coding-agent/src/core/custom-commands/loader.ts @@ -12,8 +12,8 @@ import { getAgentDir, getConfigDirs } from "../../config"; import * as piCodingAgent from "../../index"; import { execCommand } from "../exec"; import { logger } from "../logger"; -import { createReviewCommand } from "./bundled/review"; -import { createWorktreeCommand } from "./bundled/wt"; +import { ReviewCommand } from "./bundled/review"; +import { WorktreeCommand } from "./bundled/wt"; import type { CustomCommand, CustomCommandAPI, @@ -146,19 +146,17 @@ function loadBundledCommands(sharedApi: CustomCommandAPI): LoadedCustomCommand[] const bundled: LoadedCustomCommand[] = []; // Add bundled commands here - const reviewCommand = createReviewCommand(sharedApi); bundled.push({ path: "bundled:review", resolvedPath: "bundled:review", - command: reviewCommand, + command: new ReviewCommand(sharedApi), source: "bundled", }); - const worktreeCommand = createWorktreeCommand(sharedApi); bundled.push({ path: "bundled:wt", resolvedPath: "bundled:wt", - command: worktreeCommand, + command: new WorktreeCommand(sharedApi), source: "bundled", }); diff --git a/packages/coding-agent/src/core/custom-tools/index.ts b/packages/coding-agent/src/core/custom-tools/index.ts index cf18adfdf..62b507956 100644 --- a/packages/coding-agent/src/core/custom-tools/index.ts +++ b/packages/coding-agent/src/core/custom-tools/index.ts @@ -2,7 +2,7 @@ * Custom tools module. */ -export { discoverAndLoadCustomTools, loadCustomTools } from "./loader"; +export { CustomToolLoader, discoverAndLoadCustomTools, loadCustomTools } from "./loader"; export type { AgentToolResult, AgentToolUpdateCallback, @@ -19,4 +19,4 @@ export type { RenderResultOptions, ToolLoadError, } from "./types"; -export { wrapCustomTool, wrapCustomTools } from "./wrapper"; +export { CustomToolAdapter } from "./wrapper"; diff --git a/packages/coding-agent/src/core/custom-tools/loader.ts b/packages/coding-agent/src/core/custom-tools/loader.ts index 810606434..896ef8ae1 100644 --- a/packages/coding-agent/src/core/custom-tools/loader.ts +++ b/packages/coding-agent/src/core/custom-tools/loader.ts @@ -17,7 +17,7 @@ import { execCommand } from "../exec"; import type { HookUIContext } from "../hooks/types"; import { logger } from "../logger"; import { getAllPluginToolPaths } from "../plugins/loader"; -import type { CustomToolAPI, CustomToolFactory, CustomToolsLoadResult, LoadedCustomTool } from "./types"; +import type { CustomToolAPI, CustomToolFactory, LoadedCustomTool, ToolLoadError } from "./types"; /** * Resolve tool path. @@ -56,13 +56,6 @@ function createNoOpUIContext(): HookUIContext { }; } -/** Error with source metadata */ -interface ToolLoadError { - path: string; - error: string; - source?: { provider: string; providerName: string; level: "user" | "project" }; -} - /** * Load a single tool module using native Bun import. */ @@ -117,65 +110,81 @@ interface ToolPathWithSource { source?: { provider: string; providerName: string; level: "user" | "project" }; } +/** + * Loads custom tools from paths with conflict detection and error handling. + * + * Manages a shared API instance passed to all tool factories, providing access to + * execution context, UI, logger, and injected dependencies. The UI context can be + * updated after loading via setUIContext(). + */ +export class CustomToolLoader { + tools: LoadedCustomTool[] = []; + errors: ToolLoadError[] = []; + private sharedApi: CustomToolAPI; + private seenNames: Set; + + constructor(cwd: string, builtInToolNames: string[]) { + this.sharedApi = { + cwd, + exec: (command: string, args: string[], options?: ExecOptions) => + execCommand(command, args, options?.cwd ?? cwd, options), + ui: createNoOpUIContext(), + hasUI: false, + logger, + typebox, + pi: piCodingAgent, + }; + this.seenNames = new Set(builtInToolNames); + } + + async load(pathsWithSources: ToolPathWithSource[]): Promise { + for (const { path: toolPath, source } of pathsWithSources) { + const { tools: loadedTools, error } = await loadTool(toolPath, this.sharedApi.cwd, this.sharedApi, source); + + if (error) { + this.errors.push(error); + continue; + } + + if (loadedTools) { + for (const loadedTool of loadedTools) { + // Check for name conflicts + if (this.seenNames.has(loadedTool.tool.name)) { + this.errors.push({ + path: toolPath, + error: `Tool name "${loadedTool.tool.name}" conflicts with existing tool`, + source, + }); + continue; + } + + this.seenNames.add(loadedTool.tool.name); + this.tools.push(loadedTool); + } + } + } + } + + setUIContext(uiContext: HookUIContext, hasUI: boolean): void { + this.sharedApi.ui = uiContext; + this.sharedApi.hasUI = hasUI; + } +} + /** * Load all tools from configuration. * @param pathsWithSources - Array of tool paths with optional source metadata * @param cwd - Current working directory for resolving relative paths * @param builtInToolNames - Names of built-in tools to check for conflicts */ -export async function loadCustomTools( - pathsWithSources: ToolPathWithSource[], - cwd: string, - builtInToolNames: string[], -): Promise { - const tools: LoadedCustomTool[] = []; - const errors: ToolLoadError[] = []; - const seenNames = new Set(builtInToolNames); - - // Shared API object - all tools get the same instance - const sharedApi: CustomToolAPI = { - cwd, - exec: (command: string, args: string[], options?: ExecOptions) => - execCommand(command, args, options?.cwd ?? cwd, options), - ui: createNoOpUIContext(), - hasUI: false, - logger, - typebox, - pi: piCodingAgent, - }; - - for (const { path: toolPath, source } of pathsWithSources) { - const { tools: loadedTools, error } = await loadTool(toolPath, cwd, sharedApi, source); - - if (error) { - errors.push(error); - continue; - } - - if (loadedTools) { - for (const loadedTool of loadedTools) { - // Check for name conflicts - if (seenNames.has(loadedTool.tool.name)) { - errors.push({ - path: toolPath, - error: `Tool name "${loadedTool.tool.name}" conflicts with existing tool`, - source, - }); - continue; - } - - seenNames.add(loadedTool.tool.name); - tools.push(loadedTool); - } - } - } - +export async function loadCustomTools(pathsWithSources: ToolPathWithSource[], cwd: string, builtInToolNames: string[]) { + const loader = new CustomToolLoader(cwd, builtInToolNames); + await loader.load(pathsWithSources); return { - tools, - errors, - setUIContext(uiContext, hasUI) { - sharedApi.ui = uiContext; - sharedApi.hasUI = hasUI; + tools: loader.tools, + errors: loader.errors, + setUIContext: (uiContext: HookUIContext, hasUI: boolean) => { + loader.setUIContext(uiContext, hasUI); }, }; } @@ -190,11 +199,7 @@ export async function loadCustomTools( * @param cwd - Current working directory * @param builtInToolNames - Names of built-in tools to check for conflicts */ -export async function discoverAndLoadCustomTools( - configuredPaths: string[], - cwd: string, - builtInToolNames: string[], -): Promise { +export async function discoverAndLoadCustomTools(configuredPaths: string[], cwd: string, builtInToolNames: string[]) { const allPathsWithSources: ToolPathWithSource[] = []; const seen = new Set(); diff --git a/packages/coding-agent/src/core/custom-tools/types.ts b/packages/coding-agent/src/core/custom-tools/types.ts index fb63e4630..c1aaca6a4 100644 --- a/packages/coding-agent/src/core/custom-tools/types.ts +++ b/packages/coding-agent/src/core/custom-tools/types.ts @@ -158,13 +158,13 @@ export type CustomToolFactory = ( ) => CustomTool | CustomTool[] | Promise | CustomTool[]>; /** Loaded custom tool with metadata and wrapped AgentTool */ -export interface LoadedCustomTool { +export interface LoadedCustomTool { /** Original path (as specified) */ path: string; /** Resolved absolute path */ resolvedPath: string; /** The original custom tool instance */ - tool: CustomTool; + tool: CustomTool; /** Source metadata (provider and level) */ source?: { provider: string; providerName: string; level: "user" | "project" }; } diff --git a/packages/coding-agent/src/core/custom-tools/wrapper.ts b/packages/coding-agent/src/core/custom-tools/wrapper.ts index 90dcd97c4..4e1f81544 100644 --- a/packages/coding-agent/src/core/custom-tools/wrapper.ts +++ b/packages/coding-agent/src/core/custom-tools/wrapper.ts @@ -1,33 +1,69 @@ /** - * Wraps CustomTool instances into AgentTool for use with the agent. + * CustomToolAdapter wraps CustomTool instances into AgentTool for use with the agent. */ -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolResult, AgentToolUpdateCallback, RenderResultOptions } from "@oh-my-pi/pi-agent-core"; +import type { Component } from "@oh-my-pi/pi-tui"; +import type { Static, TSchema } from "@sinclair/typebox"; import type { Theme } from "../../modes/interactive/theme/theme"; import type { CustomTool, CustomToolContext, LoadedCustomTool } from "./types"; -/** - * Wrap a CustomTool into an AgentTool. - * The wrapper injects the ToolContext into execute calls. - */ -export function wrapCustomTool(tool: CustomTool, getContext: () => CustomToolContext): AgentTool { - return { - name: tool.name, - label: tool.label, - description: tool.description, - parameters: tool.parameters, - execute: (toolCallId, params, signal, onUpdate, context) => - tool.execute(toolCallId, params, onUpdate, context ?? getContext(), signal), - renderCall: tool.renderCall ? (args, theme) => tool.renderCall?.(args, theme as Theme) : undefined, - renderResult: tool.renderResult - ? (result, options, theme) => tool.renderResult?.(result, options, theme as Theme) - : undefined, - }; -} +export class CustomToolAdapter + implements AgentTool +{ + name: string; + label: string; + description: string; + parameters: TParams; -/** - * Wrap all loaded custom tools into AgentTools. - */ -export function wrapCustomTools(loadedTools: LoadedCustomTool[], getContext: () => CustomToolContext): AgentTool[] { - return loadedTools.map((lt) => wrapCustomTool(lt.tool, getContext)); + constructor( + private tool: CustomTool, + private getContext: () => CustomToolContext, + ) { + this.name = tool.name; + this.label = tool.label ?? ""; + this.description = tool.description; + this.parameters = tool.parameters; + } + + execute( + toolCallId: string, + params: Static, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + context?: CustomToolContext, + ) { + return this.tool.execute(toolCallId, params, onUpdate, context ?? this.getContext(), signal); + } + + /** Optional custom rendering for tool call display (returns UI component) */ + renderCall(args: Static, theme: TTheme): Component | undefined { + return this.tool.renderCall?.(args, theme); + } + + /** Optional custom rendering for tool result display (returns UI component) */ + renderResult(result: AgentToolResult, options: RenderResultOptions, theme: TTheme): Component | undefined { + return this.tool.renderResult?.(result, options, theme); + } + + /** + * Backward-compatible export of factory function for existing callers. + * Prefer CustomToolAdapter constructor directly. + */ + static wrap( + tool: CustomTool, + getContext: () => CustomToolContext, + ): AgentTool { + return new CustomToolAdapter(tool, getContext); + } + + /** + * Wrap all loaded custom tools into AgentTools. + */ + static wrapTools( + loadedTools: LoadedCustomTool[], + getContext: () => CustomToolContext, + ): AgentTool[] { + return loadedTools.map((lt) => CustomToolAdapter.wrap(lt.tool, getContext)); + } } diff --git a/packages/coding-agent/src/core/event-bus.ts b/packages/coding-agent/src/core/event-bus.ts index e3fa8c206..c84bd1b88 100644 --- a/packages/coding-agent/src/core/event-bus.ts +++ b/packages/coding-agent/src/core/event-bus.ts @@ -1,27 +1,7 @@ -export interface EventBus { - emit(channel: string, data: unknown): void; - on(channel: string, handler: (data: unknown) => void): () => void; -} +export class EventBus { + private readonly listeners = new Map void>>(); -export interface EventBusController extends EventBus { - clear(): void; -} - -class SimpleEventEmitter { - private listeners = new Map void>>(); - - on(channel: string, handler: (data: unknown) => void): void { - if (!this.listeners.has(channel)) { - this.listeners.set(channel, new Set()); - } - this.listeners.get(channel)!.add(handler); - } - - off(channel: string, handler: (data: unknown) => void): void { - this.listeners.get(channel)?.delete(handler); - } - - emit(channel: string, data: unknown): void { + public emit(channel: string, data: unknown): void { const handlers = this.listeners.get(channel); if (handlers) { for (const handler of handlers) { @@ -30,30 +10,22 @@ class SimpleEventEmitter { } } - removeAllListeners(): void { + public on(channel: string, handler: (data: unknown) => void): () => void { + if (!this.listeners.has(channel)) { + this.listeners.set(channel, new Set()); + } + const safeHandler = async (data: unknown) => { + try { + await handler(data); + } catch (err) { + console.error(`Event handler error (${channel}):`, err); + } + }; + this.listeners.get(channel)!.add(safeHandler); + return () => this.listeners.get(channel)?.delete(safeHandler); + } + + public clear(): void { this.listeners.clear(); } } - -export function createEventBus(): EventBusController { - const emitter = new SimpleEventEmitter(); - return { - emit: (channel, data) => { - emitter.emit(channel, data); - }, - on: (channel, handler) => { - const safeHandler = async (data: unknown) => { - try { - await handler(data); - } catch (err) { - console.error(`Event handler error (${channel}):`, err); - } - }; - emitter.on(channel, safeHandler); - return () => emitter.off(channel, safeHandler); - }, - clear: () => { - emitter.removeAllListeners(); - }, - }; -} diff --git a/packages/coding-agent/src/core/extensions/index.ts b/packages/coding-agent/src/core/extensions/index.ts index 365bc6909..8bf63cb18 100644 --- a/packages/coding-agent/src/core/extensions/index.ts +++ b/packages/coding-agent/src/core/extensions/index.ts @@ -2,7 +2,12 @@ * Extension system for lifecycle events and custom tools. */ -export { createExtensionRuntime, discoverAndLoadExtensions, loadExtensionFromFactory, loadExtensions } from "./loader"; +export { + discoverAndLoadExtensions, + ExtensionRuntime, + loadExtensionFromFactory, + loadExtensions, +} from "./loader"; export type { BranchHandler, ExtensionErrorListener, @@ -32,7 +37,6 @@ export type { ExecResult, Extension, ExtensionActions, - // API ExtensionAPI, ExtensionCommandContext, ExtensionCommandContextActions, @@ -45,7 +49,6 @@ export type { ExtensionFactory, ExtensionFlag, ExtensionHandler, - ExtensionRuntime, ExtensionShortcut, ExtensionUIContext, ExtensionUIDialogOptions, @@ -112,4 +115,10 @@ export { isReadToolResult, isWriteToolResult, } from "./types"; -export { wrapRegisteredTool, wrapRegisteredTools, wrapToolWithExtensions } from "./wrapper"; +export { + ExtensionToolWrapper, + RegisteredToolAdapter, + wrapRegisteredTool, + wrapRegisteredTools, + wrapToolWithExtensions, +} from "./wrapper"; diff --git a/packages/coding-agent/src/core/extensions/loader.ts b/packages/coding-agent/src/core/extensions/loader.ts index eb93893c4..e7a75ccbc 100644 --- a/packages/coding-agent/src/core/extensions/loader.ts +++ b/packages/coding-agent/src/core/extensions/loader.ts @@ -4,22 +4,26 @@ import { existsSync, readdirSync, readFileSync, statSync } from "node:fs"; import * as path from "node:path"; +import type { ThinkingLevel } from "@oh-my-pi/pi-agent-core"; +import type { ImageContent, Model, TextContent } from "@oh-my-pi/pi-ai"; import type { KeyId } from "@oh-my-pi/pi-tui"; +import type { TSchema } from "@sinclair/typebox"; import * as TypeBox from "@sinclair/typebox"; import { type ExtensionModule, extensionModuleCapability } from "../../capability/extension-module"; import { loadCapability } from "../../discovery"; import { expandPath, getExtensionNameFromPath } from "../../discovery/helpers"; import * as piCodingAgent from "../../index"; -import { createEventBus, type EventBus } from "../event-bus"; +import { EventBus } from "../event-bus"; import type { ExecOptions } from "../exec"; import { execCommand } from "../exec"; import { logger } from "../logger"; +import type { CustomMessage } from "../messages"; import type { Extension, ExtensionAPI, ExtensionContext, ExtensionFactory, - ExtensionRuntime, + ExtensionRuntime as IExtensionRuntime, LoadExtensionsResult, MessageRenderer, RegisteredCommand, @@ -36,147 +40,183 @@ function resolvePath(extPath: string, cwd: string): string { type HandlerFn = (...args: unknown[]) => Promise; -/** - * Create a runtime with throwing stubs for action methods. - * Runner.initialize() replaces these with real implementations. - */ -export function createExtensionRuntime(): ExtensionRuntime { - const notInitialized = () => { - throw new Error("Extension runtime not initialized. Action methods cannot be called during extension loading."); - }; - - return { - sendMessage: notInitialized, - sendUserMessage: notInitialized, - appendEntry: notInitialized, - setLabel: notInitialized, - getActiveTools: notInitialized, - getAllTools: notInitialized, - setActiveTools: notInitialized, - setModel: () => Promise.reject(new Error("Extension runtime not initialized")), - getThinkingLevel: notInitialized, - setThinkingLevel: notInitialized, - flagValues: new Map(), - }; +export class ExtensionRuntimeNotInitializedError extends Error { + constructor() { + super("Extension runtime not initialized. Action methods cannot be called during extension loading."); + } } /** - * Create the ExtensionAPI for an extension. + * Extension runtime with throwing stubs for action methods. + * These are replaced with real implementations during initialization. + */ +export class ExtensionRuntime implements IExtensionRuntime { + flagValues = new Map(); + + sendMessage(): void { + throw new ExtensionRuntimeNotInitializedError(); + } + + sendUserMessage(): void { + throw new ExtensionRuntimeNotInitializedError(); + } + + appendEntry(): void { + throw new ExtensionRuntimeNotInitializedError(); + } + + setLabel(): void { + throw new ExtensionRuntimeNotInitializedError(); + } + + getActiveTools(): string[] { + throw new ExtensionRuntimeNotInitializedError(); + } + + getAllTools(): string[] { + throw new ExtensionRuntimeNotInitializedError(); + } + + setActiveTools(): Promise { + throw new ExtensionRuntimeNotInitializedError(); + } + + setModel(): Promise { + throw new ExtensionRuntimeNotInitializedError(); + } + + getThinkingLevel(): ThinkingLevel { + throw new ExtensionRuntimeNotInitializedError(); + } + + setThinkingLevel(): void { + throw new ExtensionRuntimeNotInitializedError(); + } +} + +/** + * ExtensionAPI implementation for an extension. * Registration methods write to the extension object. * Action methods delegate to the shared runtime. */ -function createExtensionAPI( - extension: Extension, - runtime: ExtensionRuntime, - cwd: string, - eventBus: EventBus, -): ExtensionAPI { - const api = { - logger, - typebox: TypeBox, - pi: piCodingAgent, +class ConcreteExtensionAPI implements ExtensionAPI, IExtensionRuntime { + readonly logger = logger; + readonly typebox = TypeBox; + readonly pi = piCodingAgent; + readonly events: EventBus; + readonly flagValues = new Map(); - on(event: string, handler: HandlerFn): void { - const list = extension.handlers.get(event) ?? []; - list.push(handler); - extension.handlers.set(event, list); + constructor( + private extension: Extension, + private runtime: IExtensionRuntime, + private cwd: string, + eventBus: EventBus, + ) { + this.events = eventBus; + } + + on(event: string, handler: F): void { + const list = this.extension.handlers.get(event) ?? []; + list.push(handler); + this.extension.handlers.set(event, list); + } + + registerTool(tool: ToolDefinition): void { + this.extension.tools.set(tool.name, { + definition: tool, + extensionPath: this.extension.path, + }); + } + + registerCommand( + name: string, + options: { + description?: string; + getArgumentCompletions?: RegisteredCommand["getArgumentCompletions"]; + handler: RegisteredCommand["handler"]; }, + ): void { + this.extension.commands.set(name, { name, ...options }); + } - registerTool(tool: ToolDefinition): void { - extension.tools.set(tool.name, { - definition: tool, - extensionPath: extension.path, - }); + setLabel(label: string): void { + this.extension.label = label; + } + + registerShortcut( + shortcut: KeyId, + options: { + description?: string; + handler: (ctx: ExtensionContext) => Promise | void; }, + ): void { + this.extension.shortcuts.set(shortcut, { shortcut, extensionPath: this.extension.path, ...options }); + } - registerCommand( - name: string, - options: { - description?: string; - getArgumentCompletions?: RegisteredCommand["getArgumentCompletions"]; - handler: RegisteredCommand["handler"]; - }, - ): void { - extension.commands.set(name, { name, ...options }); - }, + registerFlag( + name: string, + options: { description?: string; type: "boolean" | "string"; default?: boolean | string }, + ): void { + this.extension.flags.set(name, { name, extensionPath: this.extension.path, ...options }); + if (options.default !== undefined) { + this.runtime.flagValues.set(name, options.default); + } + } - setLabel(label: string): void { - extension.label = label; - }, + registerMessageRenderer(customType: string, renderer: MessageRenderer): void { + this.extension.messageRenderers.set(customType, renderer as MessageRenderer); + } - registerShortcut( - shortcut: KeyId, - options: { - description?: string; - handler: (ctx: ExtensionContext) => Promise | void; - }, - ): void { - extension.shortcuts.set(shortcut, { shortcut, extensionPath: extension.path, ...options }); - }, + getFlag(name: string): boolean | string | undefined { + if (!this.extension.flags.has(name)) return undefined; + return this.runtime.flagValues.get(name); + } - registerFlag( - name: string, - options: { description?: string; type: "boolean" | "string"; default?: boolean | string }, - ): void { - extension.flags.set(name, { name, extensionPath: extension.path, ...options }); - if (options.default !== undefined) { - runtime.flagValues.set(name, options.default); - } - }, + sendMessage( + message: Pick, "customType" | "content" | "display" | "details">, + options?: { triggerTurn?: boolean; deliverAs?: "steer" | "followUp" | "nextTurn" }, + ): void { + this.runtime.sendMessage(message, options); + } - registerMessageRenderer(customType: string, renderer: MessageRenderer): void { - extension.messageRenderers.set(customType, renderer as MessageRenderer); - }, + sendUserMessage( + content: string | (TextContent | ImageContent)[], + options?: { deliverAs?: "steer" | "followUp" }, + ): void { + this.runtime.sendUserMessage(content, options); + } - getFlag(name: string): boolean | string | undefined { - if (!extension.flags.has(name)) return undefined; - return runtime.flagValues.get(name); - }, + appendEntry(customType: string, data?: unknown): void { + this.runtime.appendEntry(customType, data); + } - sendMessage(message, options): void { - runtime.sendMessage(message, options); - }, + exec(command: string, args: string[], options?: ExecOptions) { + return execCommand(command, args, options?.cwd ?? this.cwd, options); + } - sendUserMessage(content, options): void { - runtime.sendUserMessage(content, options); - }, + getActiveTools(): string[] { + return this.runtime.getActiveTools(); + } - appendEntry(customType: string, data?: unknown): void { - runtime.appendEntry(customType, data); - }, + getAllTools(): string[] { + return this.runtime.getAllTools(); + } - exec(command: string, args: string[], options?: ExecOptions) { - return execCommand(command, args, options?.cwd ?? cwd, options); - }, + setActiveTools(toolNames: string[]): Promise { + return this.runtime.setActiveTools(toolNames); + } - getActiveTools(): string[] { - return runtime.getActiveTools(); - }, + setModel(model: Model): Promise { + return this.runtime.setModel(model); + } - getAllTools(): string[] { - return runtime.getAllTools(); - }, + getThinkingLevel(): ThinkingLevel { + return this.runtime.getThinkingLevel(); + } - setActiveTools(toolNames: string[]): void { - runtime.setActiveTools(toolNames); - }, - - setModel(model) { - return runtime.setModel(model); - }, - - getThinkingLevel() { - return runtime.getThinkingLevel(); - }, - - setThinkingLevel(level) { - runtime.setThinkingLevel(level); - }, - - events: eventBus, - } as ExtensionAPI; - - return api; + setThinkingLevel(level: ThinkingLevel): void { + this.runtime.setThinkingLevel(level); + } } /** @@ -199,7 +239,7 @@ async function loadExtension( extensionPath: string, cwd: string, eventBus: EventBus, - runtime: ExtensionRuntime, + runtime: IExtensionRuntime, ): Promise<{ extension: Extension | null; error: string | null }> { const resolvedPath = resolvePath(extensionPath, cwd); @@ -215,7 +255,7 @@ async function loadExtension( } const extension = createExtension(extensionPath, resolvedPath); - const api = createExtensionAPI(extension, runtime, cwd, eventBus); + const api = new ConcreteExtensionAPI(extension, runtime, cwd, eventBus); await factory(api); return { extension, error: null }; @@ -232,11 +272,11 @@ export async function loadExtensionFromFactory( factory: ExtensionFactory, cwd: string, eventBus: EventBus, - runtime: ExtensionRuntime, + runtime: IExtensionRuntime, name = "", ): Promise { const extension = createExtension(name, name); - const api = createExtensionAPI(extension, runtime, cwd, eventBus); + const api = new ConcreteExtensionAPI(extension, runtime, cwd, eventBus); await factory(api); return extension; } @@ -247,8 +287,8 @@ export async function loadExtensionFromFactory( export async function loadExtensions(paths: string[], cwd: string, eventBus?: EventBus): Promise { const extensions: Extension[] = []; const errors: Array<{ path: string; error: string }> = []; - const resolvedEventBus = eventBus ?? createEventBus(); - const runtime = createExtensionRuntime(); + const resolvedEventBus = eventBus ?? new EventBus(); + const runtime = new ExtensionRuntime(); for (const extPath of paths) { const { extension, error } = await loadExtension(extPath, cwd, resolvedEventBus, runtime); diff --git a/packages/coding-agent/src/core/extensions/types.ts b/packages/coding-agent/src/core/extensions/types.ts index 28919a3fe..01ff8c344 100644 --- a/packages/coding-agent/src/core/extensions/types.ts +++ b/packages/coding-agent/src/core/extensions/types.ts @@ -779,8 +779,8 @@ export type ExtensionFactory = (pi: ExtensionAPI) => void | Promise; // Loaded Extension Types // ============================================================================ -export interface RegisteredTool { - definition: ToolDefinition; +export interface RegisteredTool { + definition: ToolDefinition; extensionPath: string; } @@ -877,7 +877,7 @@ export interface Extension { resolvedPath: string; label?: string; handlers: Map; - tools: Map; + tools: Map>; messageRenderers: Map; commands: Map; flags: Map; diff --git a/packages/coding-agent/src/core/extensions/wrapper.ts b/packages/coding-agent/src/core/extensions/wrapper.ts index e15eeeebd..d491e72cf 100644 --- a/packages/coding-agent/src/core/extensions/wrapper.ts +++ b/packages/coding-agent/src/core/extensions/wrapper.ts @@ -9,27 +9,53 @@ import type { ExtensionRunner } from "./runner"; import type { RegisteredTool, ToolCallEventResult, ToolResultEventResult } from "./types"; /** - * Wrap a RegisteredTool into an AgentTool. + * Adapts a RegisteredTool into an AgentTool. + */ +export class RegisteredToolAdapter implements AgentTool { + readonly name: string; + readonly label: string; + readonly description: string; + readonly parameters: any; + + constructor( + private registeredTool: RegisteredTool, + private runner: ExtensionRunner, + ) { + const { definition } = registeredTool; + this.name = definition.name; + this.label = definition.label || ""; + this.description = definition.description; + this.parameters = definition.parameters; + } + + async execute( + toolCallId: string, + params: any, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ) { + return this.registeredTool.definition.execute(toolCallId, params, onUpdate, this.runner.createContext(), signal); + } + + renderCall?(args: any, theme: any) { + return this.registeredTool.definition.renderCall?.(args, theme as Theme); + } + + renderResult?(result: any, options: any, theme: any) { + return this.registeredTool.definition.renderResult?.( + result, + { expanded: options.expanded, isPartial: options.isPartial, spinnerFrame: options.spinnerFrame }, + theme as Theme, + ); + } +} + +/** + * Backward-compatible factory function wrapper. */ export function wrapRegisteredTool(registeredTool: RegisteredTool, runner: ExtensionRunner): AgentTool { - const { definition } = registeredTool; - return { - name: definition.name, - label: definition.label, - description: definition.description, - parameters: definition.parameters, - execute: (toolCallId, params, signal, onUpdate) => - definition.execute(toolCallId, params, onUpdate, runner.createContext(), signal), - renderCall: definition.renderCall ? (args, theme) => definition.renderCall?.(args, theme as Theme) : undefined, - renderResult: definition.renderResult - ? (result, options, theme) => - definition.renderResult?.( - result, - { expanded: options.expanded, isPartial: options.isPartial, spinnerFrame: options.spinnerFrame }, - theme as Theme, - ) - : undefined, - }; + return new RegisteredToolAdapter(registeredTool, runner); } /** @@ -40,98 +66,121 @@ export function wrapRegisteredTools(registeredTools: RegisteredTool[], runner: E } /** - * Wrap a tool with extension callbacks for interception. + * Wraps a tool with extension callbacks for interception. * - Emits tool_call event before execution (can block) * - Emits tool_result event after execution (can modify result) */ -export function wrapToolWithExtensions(tool: AgentTool, runner: ExtensionRunner): AgentTool { - return { - ...tool, - execute: async ( - toolCallId: string, - params: Record, - signal?: AbortSignal, - onUpdate?: AgentToolUpdateCallback, - context?: AgentToolContext, - ) => { - // Emit tool_call event - extensions can block execution - if (runner.hasHandlers("tool_call")) { - try { - const callResult = (await runner.emitToolCall({ - type: "tool_call", - toolName: tool.name, - toolCallId, - input: params, - })) as ToolCallEventResult | undefined; +export class ExtensionToolWrapper implements AgentTool { + name: string; + label: string; + description: string; + parameters: unknown; + renderCall?: AgentTool["renderCall"]; + renderResult?: AgentTool["renderResult"]; - if (callResult?.block) { - const reason = callResult.reason || "Tool execution was blocked by an extension"; - throw new Error(reason); - } - } catch (err) { - if (err instanceof Error) { - throw err; - } - throw new Error(`Extension failed, blocking execution: ${String(err)}`); - } - } - - // Execute the actual tool - let result: { content: any; details: T }; - let executionError: Error | undefined; + constructor( + private tool: AgentTool, + private runner: ExtensionRunner, + ) { + this.name = tool.name; + this.label = tool.label ?? ""; + this.description = tool.description; + this.parameters = tool.parameters; + this.renderCall = tool.renderCall; + this.renderResult = tool.renderResult; + } + async execute( + toolCallId: string, + params: Record, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + context?: AgentToolContext, + ) { + // Emit tool_call event - extensions can block execution + if (this.runner.hasHandlers("tool_call")) { try { - result = await tool.execute(toolCallId, params, signal, onUpdate, context); - } catch (err) { - executionError = err instanceof Error ? err : new Error(String(err)); - result = { - content: [{ type: "text", text: executionError.message }], - details: undefined as T, - }; - } - - // Emit tool_result event - extensions can modify the result and error status - if (runner.hasHandlers("tool_result")) { - const resultResult = (await runner.emit({ - type: "tool_result", - toolName: tool.name, + const callResult = (await this.runner.emitToolCall({ + type: "tool_call", + toolName: this.tool.name, toolCallId, input: params, - content: result.content, - details: result.details, - isError: !!executionError, - })) as ToolResultEventResult | undefined; + })) as ToolCallEventResult | undefined; - if (resultResult) { - const modifiedContent: (TextContent | ImageContent)[] = resultResult.content ?? result.content; - const modifiedDetails = (resultResult.details ?? result.details) as T; + if (callResult?.block) { + const reason = callResult.reason || "Tool execution was blocked by an extension"; + throw new Error(reason); + } + } catch (err) { + if (err instanceof Error) { + throw err; + } + throw new Error(`Extension failed, blocking execution: ${String(err)}`); + } + } - // Extension can override error status - if (resultResult.isError === true && !executionError) { - // Extension marks a successful result as error - const textBlocks = (modifiedContent ?? []).filter((c): c is TextContent => c.type === "text"); - const errorText = - textBlocks.map((t) => t.text).join("\n") || "Tool result marked as error by extension"; - throw new Error(errorText); - } - if (resultResult.isError === false && executionError) { - // Extension clears the error - return success - return { content: modifiedContent, details: modifiedDetails }; - } + // Execute the actual tool + let result: { content: any; details?: T }; + let executionError: Error | undefined; - // Error status unchanged, but content/details may be modified - if (executionError) { - throw executionError; - } + try { + result = await this.tool.execute(toolCallId, params, signal, onUpdate, context); + } catch (err) { + executionError = err instanceof Error ? err : new Error(String(err)); + result = { + content: [{ type: "text", text: executionError.message }], + details: undefined as T, + }; + } + + // Emit tool_result event - extensions can modify the result and error status + if (this.runner.hasHandlers("tool_result")) { + const resultResult = (await this.runner.emit({ + type: "tool_result", + toolName: this.tool.name, + toolCallId, + input: params, + content: result.content, + details: result.details, + isError: !!executionError, + })) as ToolResultEventResult | undefined; + + if (resultResult) { + const modifiedContent: (TextContent | ImageContent)[] = resultResult.content ?? result.content; + const modifiedDetails = (resultResult.details ?? result.details) as T; + + // Extension can override error status + if (resultResult.isError === true && !executionError) { + // Extension marks a successful result as error + const textBlocks = (modifiedContent ?? []).filter((c): c is TextContent => c.type === "text"); + const errorText = textBlocks.map((t) => t.text).join("\n") || "Tool result marked as error by extension"; + throw new Error(errorText); + } + if (resultResult.isError === false && executionError) { + // Extension clears the error - return success return { content: modifiedContent, details: modifiedDetails }; } - } - // No extension modification - if (executionError) { - throw executionError; + // Error status unchanged, but content/details may be modified + if (executionError) { + throw executionError; + } + return { content: modifiedContent, details: modifiedDetails }; } - return result; - }, - }; + } + + // No extension modification + if (executionError) { + throw executionError; + } + return result; + } +} + +/** + * Wrap a tool with extension callbacks for interception. + * @deprecated Use `new ExtensionToolWrapper()` directly + */ +export function wrapToolWithExtensions(tool: AgentTool, runner: ExtensionRunner): AgentTool { + return new ExtensionToolWrapper(tool, runner); } diff --git a/packages/coding-agent/src/core/hooks/index.ts b/packages/coding-agent/src/core/hooks/index.ts index fff9540fc..1780603a1 100644 --- a/packages/coding-agent/src/core/hooks/index.ts +++ b/packages/coding-agent/src/core/hooks/index.ts @@ -11,6 +11,6 @@ export { type SendMessageHandler, } from "./loader"; export { execCommand, HookRunner, type HookErrorListener } from "./runner"; -export { wrapToolsWithHooks, wrapToolWithHooks } from "./tool-wrapper"; +export { HookToolWrapper, wrapToolsWithHooks, wrapToolWithHooks } from "./tool-wrapper"; export * from "./types"; export type { UsageStatistics, ReadonlySessionManager } from "../session-manager"; diff --git a/packages/coding-agent/src/core/hooks/tool-wrapper.ts b/packages/coding-agent/src/core/hooks/tool-wrapper.ts index d2afa1c69..31214fa47 100644 --- a/packages/coding-agent/src/core/hooks/tool-wrapper.ts +++ b/packages/coding-agent/src/core/hooks/tool-wrapper.ts @@ -7,93 +7,119 @@ import type { HookRunner } from "./runner"; import type { ToolCallEventResult, ToolResultEventResult } from "./types"; /** - * Wrap a tool with hook callbacks. + * Wraps an AgentTool with hook callbacks for interception. + * + * Features: * - Emits tool_call event before execution (can block) * - Emits tool_result event after execution (can modify result) * - Forwards onUpdate callback to wrapped tool for progress streaming */ -export function wrapToolWithHooks(tool: AgentTool, hookRunner: HookRunner): AgentTool { - return { - ...tool, - execute: async ( - toolCallId: string, - params: Record, - signal?: AbortSignal, - onUpdate?: AgentToolUpdateCallback, - context?: AgentToolContext, - ) => { - // Emit tool_call event - hooks can block execution - // If hook errors/times out, block by default (fail-safe) - if (hookRunner.hasHandlers("tool_call")) { - try { - const callResult = (await hookRunner.emitToolCall({ - type: "tool_call", - toolName: tool.name, - toolCallId, - input: params, - })) as ToolCallEventResult | undefined; +export class HookToolWrapper implements AgentTool { + name: string; + label: string; + description: string; + parameters: unknown; + renderCall?: AgentTool["renderCall"]; + renderResult?: AgentTool["renderResult"]; - if (callResult?.block) { - const reason = callResult.reason || "Tool execution was blocked by a hook"; - throw new Error(reason); - } - } catch (err) { - // Hook error or block - throw to mark as error - if (err instanceof Error) { - throw err; - } - throw new Error(`Hook failed, blocking execution: ${String(err)}`); - } - } + constructor( + private tool: AgentTool, + private hookRunner: HookRunner, + ) { + this.name = tool.name; + this.label = tool.label ?? ""; + this.description = tool.description; + this.parameters = tool.parameters; + this.renderCall = tool.renderCall; + this.renderResult = tool.renderResult; + } - // Execute the actual tool, forwarding onUpdate for progress streaming + async execute( + toolCallId: string, + params: Record, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + context?: AgentToolContext, + ) { + // Emit tool_call event - hooks can block execution + // If hook errors/times out, block by default (fail-safe) + if (this.hookRunner.hasHandlers("tool_call")) { try { - const result = await tool.execute(toolCallId, params, signal, onUpdate, context); + const callResult = (await this.hookRunner.emitToolCall({ + type: "tool_call", + toolName: this.tool.name, + toolCallId, + input: params, + })) as ToolCallEventResult | undefined; - // Emit tool_result event - hooks can modify the result - if (hookRunner.hasHandlers("tool_result")) { - const resultResult = (await hookRunner.emit({ - type: "tool_result", - toolName: tool.name, - toolCallId, - input: params, - content: result.content, - details: result.details, - isError: false, - })) as ToolResultEventResult | undefined; - - // Apply modifications if any - if (resultResult) { - return { - content: resultResult.content ?? result.content, - details: (resultResult.details ?? result.details) as T, - }; - } + if (callResult?.block) { + const reason = callResult.reason || "Tool execution was blocked by a hook"; + throw new Error(reason); } - - return result; } catch (err) { - // Emit tool_result event for errors so hooks can observe failures - if (hookRunner.hasHandlers("tool_result")) { - await hookRunner.emit({ - type: "tool_result", - toolName: tool.name, - toolCallId, - input: params, - content: [{ type: "text", text: err instanceof Error ? err.message : String(err) }], - details: undefined, - isError: true, - }); + // Hook error or block - throw to mark as error + if (err instanceof Error) { + throw err; } - throw err; // Re-throw original error for agent-loop + throw new Error(`Hook failed, blocking execution: ${String(err)}`); } - }, - }; + } + + // Execute the actual tool, forwarding onUpdate for progress streaming + try { + const result = await this.tool.execute(toolCallId, params, signal, onUpdate, context); + + // Emit tool_result event - hooks can modify the result + if (this.hookRunner.hasHandlers("tool_result")) { + const resultResult = (await this.hookRunner.emit({ + type: "tool_result", + toolName: this.tool.name, + toolCallId, + input: params, + content: result.content, + details: result.details, + isError: false, + })) as ToolResultEventResult | undefined; + + // Apply modifications if any + if (resultResult) { + return { + content: resultResult.content ?? result.content, + details: (resultResult.details ?? result.details) as T, + }; + } + } + + return result; + } catch (err) { + // Emit tool_result event for errors so hooks can observe failures + if (this.hookRunner.hasHandlers("tool_result")) { + await this.hookRunner.emit({ + type: "tool_result", + toolName: this.tool.name, + toolCallId, + input: params, + content: [{ type: "text", text: err instanceof Error ? err.message : String(err) }], + details: undefined, + isError: true, + }); + } + throw err; // Re-throw original error for agent-loop + } + } } /** * Wrap all tools with hook callbacks. */ export function wrapToolsWithHooks(tools: AgentTool[], hookRunner: HookRunner): AgentTool[] { - return tools.map((tool) => wrapToolWithHooks(tool, hookRunner)); + return tools.map((tool) => new HookToolWrapper(tool, hookRunner)); +} + +/** + * Backward compatibility alias - use HookToolWrapper directly. + * @deprecated Use HookToolWrapper class instead + */ +export function wrapToolWithHooks(tool: AgentTool, hookRunner: HookRunner): AgentTool { + return new HookToolWrapper(tool, hookRunner); } diff --git a/packages/coding-agent/src/core/mcp/index.ts b/packages/coding-agent/src/core/mcp/index.ts index 292cf5095..79ef5531a 100644 --- a/packages/coding-agent/src/core/mcp/index.ts +++ b/packages/coding-agent/src/core/mcp/index.ts @@ -7,7 +7,6 @@ // Client export { callTool, connectToServer, disconnectServer, listTools, serverSupportsTools } from "./client"; - // Config export type { ExaFilterResult, LoadMCPConfigsOptions, LoadMCPConfigsResult } from "./config"; export { @@ -17,6 +16,9 @@ export { loadAllMCPConfigs, validateServerConfig, } from "./config"; +// JSON-RPC (lightweight HTTP-based MCP calls) +export type { JsonRpcResponse } from "./json-rpc"; +export { callMCP, parseSSE } from "./json-rpc"; // Loader (for SDK integration) export type { MCPToolsLoadOptions, MCPToolsLoadResult } from "./loader"; export { discoverAndLoadMCPTools } from "./loader"; @@ -25,7 +27,12 @@ export type { MCPDiscoverOptions, MCPLoadResult } from "./manager"; export { createMCPManager, MCPManager } from "./manager"; // Tool bridge export type { MCPToolDetails } from "./tool-bridge"; -export { createMCPTool, createMCPToolName, createMCPTools, parseMCPToolName } from "./tool-bridge"; +export { + createMCPToolName, + DeferredMCPTool, + MCPTool, + parseMCPToolName, +} from "./tool-bridge"; // Tool cache export { MCPToolCache } from "./tool-cache"; // Transports diff --git a/packages/coding-agent/src/core/mcp/json-rpc.ts b/packages/coding-agent/src/core/mcp/json-rpc.ts new file mode 100644 index 000000000..aee1e360c --- /dev/null +++ b/packages/coding-agent/src/core/mcp/json-rpc.ts @@ -0,0 +1,88 @@ +/** + * MCP JSON-RPC 2.0 over HTTPS. + * + * Lightweight utilities for calling MCP servers directly via HTTP + * without maintaining persistent connections. + */ + +import { logger } from "../logger"; + +/** Parse SSE response format (lines starting with "data: ") */ +export function parseSSE(text: string): unknown { + const lines = text.split("\n"); + for (const line of lines) { + if (line.startsWith("data: ")) { + const data = line.slice(6).trim(); + if (data === "[DONE]") continue; + try { + return JSON.parse(data); + } catch { + // Try next line + } + } + } + // Fallback: try parsing entire response as JSON + try { + return JSON.parse(text); + } catch { + return null; + } +} + +/** JSON-RPC 2.0 response structure */ +export interface JsonRpcResponse { + jsonrpc: "2.0"; + id: string | number; + result?: T; + error?: { + code: number; + message: string; + data?: unknown; + }; +} + +/** + * Call an MCP server with JSON-RPC 2.0 over HTTPS. + * + * @param url - Full MCP server URL (including any query parameters) + * @param method - JSON-RPC method name (e.g., "tools/list", "tools/call") + * @param params - Method parameters + * @returns Parsed JSON-RPC response + */ +export async function callMCP( + url: string, + method: string, + params?: Record, +): Promise> { + const body = { + jsonrpc: "2.0", + id: Math.random().toString(36).slice(2), + method, + params: params ?? {}, + }; + + const response = await fetch(url, { + method: "POST", + headers: { + "Content-Type": "application/json", + Accept: "application/json, text/event-stream", + }, + body: JSON.stringify(body), + }); + + if (!response.ok) { + const errorMsg = `MCP request failed: ${response.status} ${response.statusText}`; + logger.error(errorMsg, { url, method, params }); + throw new Error(errorMsg); + } + + const text = await response.text(); + const result = parseSSE(text) as JsonRpcResponse | null; + + if (!result) { + logger.error("Failed to parse MCP response", { url, method, responseText: text.slice(0, 500) }); + throw new Error("Failed to parse MCP response"); + } + + return result; +} diff --git a/packages/coding-agent/src/core/mcp/manager.ts b/packages/coding-agent/src/core/mcp/manager.ts index 3c36e7618..5677aa33b 100644 --- a/packages/coding-agent/src/core/mcp/manager.ts +++ b/packages/coding-agent/src/core/mcp/manager.ts @@ -11,7 +11,7 @@ import { logger } from "../logger"; import { connectToServer, disconnectServer, listTools } from "./client"; import { loadAllMCPConfigs, validateServerConfig } from "./config"; import type { MCPToolDetails } from "./tool-bridge"; -import { createDeferredMCPTools, createMCPTools } from "./tool-bridge"; +import { DeferredMCPTool, MCPTool } from "./tool-bridge"; import type { MCPToolCache } from "./tool-cache"; import type { MCPServerConfig, MCPServerConnection, MCPToolDefinition } from "./types"; @@ -188,7 +188,7 @@ export class MCPManager { .then(({ connection, serverTools }) => { if (this.pendingToolLoads.get(name) !== toolsPromise) return; this.pendingToolLoads.delete(name); - const customTools = createMCPTools(connection, serverTools); + const customTools = MCPTool.fromTools(connection, serverTools); this.replaceServerTools(name, customTools); void this.toolCache?.set(name, config, serverTools); }) @@ -240,7 +240,7 @@ export class MCPManager { if (!value) continue; const { connection, serverTools } = value; connectedServers.add(name); - allTools.push(...createMCPTools(connection, serverTools)); + allTools.push(...MCPTool.fromTools(connection, serverTools)); } else if (task.tracked.status === "rejected") { const message = task.tracked.reason instanceof Error ? task.tracked.reason.message : String(task.tracked.reason); @@ -250,7 +250,7 @@ export class MCPManager { const cached = cachedTools.get(name); if (cached) { const source = this.sources.get(name); - allTools.push(...createDeferredMCPTools(name, cached, () => this.waitForConnection(name), source)); + allTools.push(...DeferredMCPTool.fromTools(name, cached, () => this.waitForConnection(name), source)); } } } @@ -356,7 +356,7 @@ export class MCPManager { // Reload tools const serverTools = await listTools(connection); - const customTools = createMCPTools(connection, serverTools); + const customTools = MCPTool.fromTools(connection, serverTools); void this.toolCache?.set(name, connection.config, serverTools); // Replace tools from this server diff --git a/packages/coding-agent/src/core/mcp/tool-bridge.ts b/packages/coding-agent/src/core/mcp/tool-bridge.ts index a29a6d3c8..709c094ad 100644 --- a/packages/coding-agent/src/core/mcp/tool-bridge.ts +++ b/packages/coding-agent/src/core/mcp/tool-bridge.ts @@ -4,9 +4,10 @@ * Converts MCP tool definitions to CustomTool format for the agent. */ +import type { AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { TSchema } from "@sinclair/typebox"; import type { SourceMeta } from "../../capability/types"; -import type { CustomTool, CustomToolResult } from "../custom-tools/types"; +import type { CustomTool, CustomToolContext, CustomToolResult } from "../custom-tools/types"; import { callTool } from "./client"; import type { MCPContent, MCPServerConnection, MCPToolDefinition } from "./types"; @@ -89,138 +90,155 @@ export function parseMCPToolName(name: string): { serverName: string; toolName: } /** - * Convert an MCP tool definition to a CustomTool. + * CustomTool wrapping an MCP tool with an active connection. */ -export function createMCPTool( - connection: MCPServerConnection, - tool: MCPToolDefinition, -): CustomTool { - const name = createMCPToolName(connection.name, tool.name); - const schema = convertSchema(tool.inputSchema); +export class MCPTool implements CustomTool { + public readonly name: string; + public readonly label: string; + public readonly description: string; + public readonly parameters: TSchema; - return { - name, - label: `${connection.name}/${tool.name}`, - description: tool.description ?? `MCP tool from ${connection.name}`, - parameters: schema, + /** Create MCPTool instances for all tools from an MCP server connection */ + static fromTools(connection: MCPServerConnection, tools: MCPToolDefinition[]): MCPTool[] { + return tools.map((tool) => new MCPTool(connection, tool)); + } - async execute(_toolCallId, params, _onUpdate, _ctx, _signal): Promise> { - try { - const result = await callTool(connection, tool.name, params as Record); + constructor( + private readonly connection: MCPServerConnection, + private readonly tool: MCPToolDefinition, + ) { + this.name = createMCPToolName(connection.name, tool.name); + this.label = `${connection.name}/${tool.name}`; + this.description = tool.description ?? `MCP tool from ${connection.name}`; + this.parameters = convertSchema(tool.inputSchema); + } - const text = formatMCPContent(result.content); - const details: MCPToolDetails = { - serverName: connection.name, - mcpToolName: tool.name, - isError: result.isError, - rawContent: result.content, - provider: connection._source?.provider, - providerName: connection._source?.providerName, - }; + async execute( + _toolCallId: string, + params: unknown, + _onUpdate: AgentToolUpdateCallback | undefined, + _ctx: CustomToolContext, + _signal?: AbortSignal, + ): Promise> { + try { + const result = await callTool(this.connection, this.tool.name, params as Record); - if (result.isError) { - return { - content: [{ type: "text", text: `Error: ${text}` }], - details, - }; - } + const text = formatMCPContent(result.content); + const details: MCPToolDetails = { + serverName: this.connection.name, + mcpToolName: this.tool.name, + isError: result.isError, + rawContent: result.content, + provider: this.connection._source?.provider, + providerName: this.connection._source?.providerName, + }; + if (result.isError) { return { - content: [{ type: "text", text }], + content: [{ type: "text", text: `Error: ${text}` }], details, }; - } catch (error) { - const message = error instanceof Error ? error.message : String(error); - return { - content: [{ type: "text", text: `MCP error: ${message}` }], - details: { - serverName: connection.name, - mcpToolName: tool.name, - isError: true, - provider: connection._source?.provider, - providerName: connection._source?.providerName, - }, - }; } - }, - }; -} -export function createDeferredMCPTool( - serverName: string, - tool: MCPToolDefinition, - getConnection: () => Promise, - source?: SourceMeta, -): CustomTool { - const name = createMCPToolName(serverName, tool.name); - const schema = convertSchema(tool.inputSchema); - const fallbackProvider = source?.provider; - const fallbackProviderName = source?.providerName; - - return { - name, - label: `${serverName}/${tool.name}`, - description: tool.description ?? `MCP tool from ${serverName}`, - parameters: schema, - - async execute(_toolCallId, params, _onUpdate, _ctx, _signal): Promise> { - try { - const connection = await getConnection(); - const result = await callTool(connection, tool.name, params as Record); - - const text = formatMCPContent(result.content); - const details: MCPToolDetails = { - serverName, - mcpToolName: tool.name, - isError: result.isError, - rawContent: result.content, - provider: connection._source?.provider ?? fallbackProvider, - providerName: connection._source?.providerName ?? fallbackProviderName, - }; - - if (result.isError) { - return { - content: [{ type: "text", text: `Error: ${text}` }], - details, - }; - } - - return { - content: [{ type: "text", text }], - details, - }; - } catch (error) { - const message = error instanceof Error ? error.message : String(error); - return { - content: [{ type: "text", text: `MCP error: ${message}` }], - details: { - serverName, - mcpToolName: tool.name, - isError: true, - provider: fallbackProvider, - providerName: fallbackProviderName, - }, - }; - } - }, - }; -} - -export function createDeferredMCPTools( - serverName: string, - tools: MCPToolDefinition[], - getConnection: () => Promise, - source?: SourceMeta, -): CustomTool[] { - return tools.map((tool) => createDeferredMCPTool(serverName, tool, getConnection, source)); + return { + content: [{ type: "text", text }], + details, + }; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + return { + content: [{ type: "text", text: `MCP error: ${message}` }], + details: { + serverName: this.connection.name, + mcpToolName: this.tool.name, + isError: true, + provider: this.connection._source?.provider, + providerName: this.connection._source?.providerName, + }, + }; + } + } } /** - * Convert all tools from an MCP server to CustomTools. + * CustomTool wrapping an MCP tool with deferred connection resolution. */ -export function createMCPTools( - connection: MCPServerConnection, - tools: MCPToolDefinition[], -): CustomTool[] { - return tools.map((tool) => createMCPTool(connection, tool)); +export class DeferredMCPTool implements CustomTool { + public readonly name: string; + public readonly label: string; + public readonly description: string; + public readonly parameters: TSchema; + private readonly fallbackProvider: string | undefined; + private readonly fallbackProviderName: string | undefined; + + /** Create DeferredMCPTool instances for all tools from an MCP server */ + static fromTools( + serverName: string, + tools: MCPToolDefinition[], + getConnection: () => Promise, + source?: SourceMeta, + ): DeferredMCPTool[] { + return tools.map((tool) => new DeferredMCPTool(serverName, tool, getConnection, source)); + } + + constructor( + private readonly serverName: string, + private readonly tool: MCPToolDefinition, + private readonly getConnection: () => Promise, + source?: SourceMeta, + ) { + this.name = createMCPToolName(serverName, tool.name); + this.label = `${serverName}/${tool.name}`; + this.description = tool.description ?? `MCP tool from ${serverName}`; + this.parameters = convertSchema(tool.inputSchema); + this.fallbackProvider = source?.provider; + this.fallbackProviderName = source?.providerName; + } + + async execute( + _toolCallId: string, + params: unknown, + _onUpdate: AgentToolUpdateCallback | undefined, + _ctx: CustomToolContext, + _signal?: AbortSignal, + ): Promise> { + try { + const connection = await this.getConnection(); + const result = await callTool(connection, this.tool.name, params as Record); + + const text = formatMCPContent(result.content); + const details: MCPToolDetails = { + serverName: this.serverName, + mcpToolName: this.tool.name, + isError: result.isError, + rawContent: result.content, + provider: connection._source?.provider ?? this.fallbackProvider, + providerName: connection._source?.providerName ?? this.fallbackProviderName, + }; + + if (result.isError) { + return { + content: [{ type: "text", text: `Error: ${text}` }], + details, + }; + } + + return { + content: [{ type: "text", text }], + details, + }; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + return { + content: [{ type: "text", text: `MCP error: ${message}` }], + details: { + serverName: this.serverName, + mcpToolName: this.tool.name, + isError: true, + provider: this.fallbackProvider, + providerName: this.fallbackProviderName, + }, + }; + } + } } diff --git a/packages/coding-agent/src/core/python-executor.ts b/packages/coding-agent/src/core/python-executor.ts index 83a637179..c1dc78cd7 100644 --- a/packages/coding-agent/src/core/python-executor.ts +++ b/packages/coding-agent/src/core/python-executor.ts @@ -9,7 +9,7 @@ import { type PreludeHelper, PythonKernel, } from "./python-kernel"; -import { createOutputSink } from "./streaming-output"; +import { OutputSink } from "./streaming-output"; import { DEFAULT_MAX_BYTES } from "./tools/truncate"; export type PythonKernelMode = "session" | "per-call"; @@ -218,7 +218,7 @@ async function executeWithKernel( code: string, options: PythonExecutorOptions | undefined, ): Promise { - const sink = createOutputSink(DEFAULT_MAX_BYTES, DEFAULT_MAX_BYTES * 2, options?.onChunk); + const sink = new OutputSink(DEFAULT_MAX_BYTES, DEFAULT_MAX_BYTES * 2, options?.onChunk); const writer = sink.getWriter(); const displayOutputs: KernelDisplayOutput[] = []; diff --git a/packages/coding-agent/src/core/sdk.ts b/packages/coding-agent/src/core/sdk.ts index 759091cca..b0a966ccf 100644 --- a/packages/coding-agent/src/core/sdk.ts +++ b/packages/coding-agent/src/core/sdk.ts @@ -40,13 +40,13 @@ import { initializeWithSettings } from "../discovery"; import { registerAsyncCleanup } from "../modes/cleanup"; import { AgentSession } from "./agent-session"; import { AuthStorage } from "./auth-storage"; -import { createCursorExecHandlers } from "./cursor/exec-bridge"; +import { CursorExecHandlers } from "./cursor/exec-bridge"; import { type CustomCommandsLoadResult, loadCustomCommands as loadCustomCommandsInternal, } from "./custom-commands/index"; import type { CustomTool, CustomToolContext, CustomToolSessionEvent } from "./custom-tools/types"; -import { createEventBus, type EventBus } from "./event-bus"; +import { EventBus } from "./event-bus"; import { discoverAndLoadExtensions, type ExtensionContext, @@ -79,29 +79,29 @@ import { loadProjectContextFiles as loadContextFilesInternal, } from "./system-prompt"; import { time } from "./timings"; -import { createToolContextStore } from "./tools/context"; +import { ToolContextStore } from "./tools/context"; import { getGeminiImageTools } from "./tools/gemini-image"; import { + BashTool, BUILTIN_TOOLS, - createBashTool, - createFindTool, - createGitTool, - createGrepTool, - createLsTool, - createPythonTool, - createReadTool, - createSshTool, createTools, - createWriteTool, EditTool, + FindTool, + GitTool, + GrepTool, getWebSearchTools, + LsTool, + loadSshTool, + PythonTool, + ReadTool, setPreferredImageProvider, setPreferredWebSearchProvider, type Tool, type ToolSession, + WriteTool, warmupLspServers, } from "./tools/index"; -import { createTtsrManager } from "./ttsr"; +import { TtsrManager } from "./ttsr"; // Types export interface CreateAgentSessionOptions { @@ -212,21 +212,21 @@ export type { FileSlashCommand } from "./slash-commands"; export type { Tool } from "./tools/index"; export { - // Tool factories + // Tool classes and factories BUILTIN_TOOLS, createTools, type ToolSession, - // Individual tool factories (for custom usage) - createReadTool, - createBashTool, - createPythonTool, - createSshTool, + // Individual tool classes (for custom usage) + BashTool, EditTool, - createWriteTool, - createGrepTool, - createFindTool, - createGitTool, - createLsTool, + FindTool, + GitTool, + GrepTool, + loadSshTool, + LsTool, + PythonTool, + ReadTool, + WriteTool, }; // Helper Functions @@ -551,7 +551,7 @@ function createCustomToolsExtension(tools: CustomTool[]): ExtensionFactory { export async function createAgentSession(options: CreateAgentSessionOptions = {}): Promise { const cwd = options.cwd ?? process.cwd(); const agentDir = options.agentDir ?? getDefaultAgentDir(); - const eventBus = options.eventBus ?? createEventBus(); + const eventBus = options.eventBus ?? new EventBus(); registerSshCleanup(); registerPythonCleanup(); @@ -662,7 +662,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} time("discoverSkills"); // Discover rules - const ttsrManager = createTtsrManager(settingsManager.getTtsrSettings()); + const ttsrManager = new TtsrManager(settingsManager.getTtsrSettings()); const rulesResult = await loadCapability(ruleCapability.id, { cwd }); for (const rule of rulesResult.items) { if (rule.ttsrTrigger) { @@ -850,7 +850,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} session.abort(); }, }); - const toolContextStore = createToolContextStore(getSessionContext); + const toolContextStore = new ToolContextStore(getSessionContext); const registeredTools = extensionRunner?.getAllRegisteredTools() ?? []; const allCustomTools = [ @@ -881,10 +881,10 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} time("combineTools"); let cursorEventEmitter: ((event: AgentEvent) => void) | undefined; - const cursorExecHandlers = createCursorExecHandlers({ + const cursorExecHandlers = new CursorExecHandlers({ cwd, tools: toolRegistry, - getToolContext: toolContextStore.getContext, + getToolContext: () => toolContextStore.getContext(), emitEvent: (event) => cursorEventEmitter?.(event), }); @@ -986,7 +986,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} followUpMode: settingsManager.getFollowUpMode(), interruptMode: settingsManager.getInterruptMode(), thinkingBudgets: settingsManager.getThinkingBudgets(), - getToolContext: toolContextStore.getContext, + getToolContext: (tc) => toolContextStore.getContext(tc), getApiKey: async () => { const currentModel = agent.state.model; if (!currentModel) { diff --git a/packages/coding-agent/src/core/streaming-output.ts b/packages/coding-agent/src/core/streaming-output.ts index 42ba6b337..b82562362 100644 --- a/packages/coding-agent/src/core/streaming-output.ts +++ b/packages/coding-agent/src/core/streaming-output.ts @@ -42,59 +42,68 @@ export async function pumpStream(readable: ReadableStream, writer: W } } -export function createOutputSink( - spillThreshold: number, - maxBuffer: number, - onChunk?: (text: string) => void, -): WritableStream & { - dump: (annotation?: string) => { output: string; truncated: boolean; fullOutputPath?: string }; -} { - const chunks: Array<{ text: string; bytes: number }> = []; - let chunkBytes = 0; - let totalBytes = 0; - let fullOutputPath: string | undefined; - let fullOutputStream: OutputFileSink | undefined; - - const sink = new WritableStream({ - write(text) { - const bytes = Buffer.byteLength(text, "utf-8"); - totalBytes += bytes; - - if (totalBytes > spillThreshold && !fullOutputPath) { - fullOutputPath = join(tmpdir(), `omp-${nanoid()}.buffer`); - const stream = Bun.file(fullOutputPath).writer(); - for (const chunk of chunks) { - stream.write(chunk.text); - } - fullOutputStream = stream; - } - fullOutputStream?.write(text); - - chunks.push({ text, bytes }); - chunkBytes += bytes; - while (chunkBytes > maxBuffer && chunks.length > 1) { - const removed = chunks.shift(); - if (removed) { - chunkBytes -= removed.bytes; - } - } - - onChunk?.(text); - }, - close() { - fullOutputStream?.end(); - }, - }); - - return Object.assign(sink, { - dump(annotation?: string) { - if (annotation) { - const text = `\n\n${annotation}`; - chunks.push({ text, bytes: Buffer.byteLength(text, "utf-8") }); - } - const full = chunks.map((chunk) => chunk.text).join(""); - const { content, truncated } = truncateTail(full); - return { output: truncated ? content : full, truncated, fullOutputPath }; - }, - }); +export interface OutputSinkDump { + output: string; + truncated: boolean; + fullOutputPath?: string; +} + +export class OutputSink { + private readonly stream: WritableStream; + private readonly chunks: Array<{ text: string; bytes: number }> = []; + private chunkBytes = 0; + private totalBytes = 0; + private fullOutputPath: string | undefined; + private fullOutputStream: OutputFileSink | undefined; + + constructor( + private readonly spillThreshold: number, + private readonly maxBuffer: number, + private readonly onChunk?: (text: string) => void, + ) { + this.stream = new WritableStream({ + write: (text) => { + const bytes = Buffer.byteLength(text, "utf-8"); + this.totalBytes += bytes; + + if (this.totalBytes > this.spillThreshold && !this.fullOutputPath) { + this.fullOutputPath = join(tmpdir(), `omp-${nanoid()}.buffer`); + const stream = Bun.file(this.fullOutputPath).writer(); + for (const chunk of this.chunks) { + stream.write(chunk.text); + } + this.fullOutputStream = stream; + } + this.fullOutputStream?.write(text); + + this.chunks.push({ text, bytes }); + this.chunkBytes += bytes; + while (this.chunkBytes > this.maxBuffer && this.chunks.length > 1) { + const removed = this.chunks.shift(); + if (removed) { + this.chunkBytes -= removed.bytes; + } + } + + this.onChunk?.(text); + }, + close: () => { + this.fullOutputStream?.end(); + }, + }); + } + + getWriter(): WritableStreamDefaultWriter { + return this.stream.getWriter(); + } + + dump(annotation?: string): OutputSinkDump { + if (annotation) { + const text = `\n\n${annotation}`; + this.chunks.push({ text, bytes: Buffer.byteLength(text, "utf-8") }); + } + const full = this.chunks.map((chunk) => chunk.text).join(""); + const { content, truncated } = truncateTail(full); + return { output: truncated ? content : full, truncated, fullOutputPath: this.fullOutputPath }; + } } diff --git a/packages/coding-agent/src/core/tools/ask.ts b/packages/coding-agent/src/core/tools/ask.ts index 0744efe1b..d767ec8d5 100644 --- a/packages/coding-agent/src/core/tools/ask.ts +++ b/packages/coding-agent/src/core/tools/ask.ts @@ -15,7 +15,7 @@ * and add "(Recommended)" at the end of the label */ -import type { AgentTool, AgentToolContext, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; import { Type } from "@sinclair/typebox"; @@ -188,7 +188,7 @@ function formatQuestionResult(result: QuestionResult): string { } // ============================================================================= -// Tool Implementation +// Tool Class // ============================================================================= interface AskParams { @@ -203,109 +203,108 @@ interface AskParams { }>; } -export function createAskTool(session: ToolSession): null | AgentTool { - if (!session.hasUI) { - return null; +/** + * Ask tool for interactive user prompting during execution. + * + * Allows gathering user preferences, clarifying instructions, and getting decisions + * on implementation choices as the agent works. + */ +export class AskTool implements AgentTool { + public readonly name = "ask"; + public readonly label = "Ask"; + public readonly description: string; + public readonly parameters = askSchema; + + constructor(_session: ToolSession) { + this.description = renderPromptTemplate(askDescription); } - return { - name: "ask", - label: "Ask", - description: renderPromptTemplate(askDescription), - parameters: askSchema, - async execute( - _toolCallId: string, - params: AskParams, - _signal?: AbortSignal, - _onUpdate?: AgentToolUpdateCallback, - context?: AgentToolContext, - ) { - // Headless fallback - if (!context?.hasUI || !context.ui) { - return { - content: [{ type: "text" as const, text: "Error: User prompt requires interactive mode" }], - details: {}, - }; - } + static createIf(session: ToolSession): AskTool | null { + return session.hasUI ? new AskTool(session) : null; + } - const { ui } = context; - - // Multi-part questions mode - if (params.questions && params.questions.length > 0) { - const results: QuestionResult[] = []; - - for (const q of params.questions) { - const optionLabels = q.options.map((o) => o.label); - const { selectedOptions, customInput } = await askSingleQuestion( - ui, - q.question, - optionLabels, - q.multi ?? false, - ); - - results.push({ - id: q.id, - question: q.question, - options: optionLabels, - multi: q.multi ?? false, - selectedOptions, - customInput, - }); - } - - const details: AskToolDetails = { results }; - const responseLines = results.map(formatQuestionResult); - const responseText = `User answers:\n${responseLines.join("\n")}`; - - return { content: [{ type: "text" as const, text: responseText }], details }; - } - - // Single question mode (backwards compatible) - const question = params.question ?? ""; - const options = params.options ?? []; - const multi = params.multi ?? false; - const optionLabels = options.map((o) => o.label); - - if (!question || optionLabels.length === 0) { - return { - content: [{ type: "text" as const, text: "Error: question and options are required" }], - details: {}, - }; - } - - const { selectedOptions, customInput } = await askSingleQuestion(ui, question, optionLabels, multi); - - const details: AskToolDetails = { - question, - options: optionLabels, - multi, - selectedOptions, - customInput, + public async execute( + _toolCallId: string, + params: AskParams, + _signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + context?: AgentToolContext, + ): Promise> { + // Headless fallback + if (!context?.hasUI || !context.ui) { + return { + content: [{ type: "text" as const, text: "Error: User prompt requires interactive mode" }], + details: {}, }; + } - let responseText: string; - if (customInput) { - responseText = `User provided custom input: ${customInput}`; - } else if (selectedOptions.length > 0) { - responseText = multi - ? `User selected: ${selectedOptions.join(", ")}` - : `User selected: ${selectedOptions[0]}`; - } else { - responseText = "User cancelled the selection"; + const { ui } = context; + + // Multi-part questions mode + if (params.questions && params.questions.length > 0) { + const results: QuestionResult[] = []; + + for (const q of params.questions) { + const optionLabels = q.options.map((o) => o.label); + const { selectedOptions, customInput } = await askSingleQuestion( + ui, + q.question, + optionLabels, + q.multi ?? false, + ); + + results.push({ + id: q.id, + question: q.question, + options: optionLabels, + multi: q.multi ?? false, + selectedOptions, + customInput, + }); } + const details: AskToolDetails = { results }; + const responseLines = results.map(formatQuestionResult); + const responseText = `User answers:\n${responseLines.join("\n")}`; + return { content: [{ type: "text" as const, text: responseText }], details }; - }, - }; -} + } -/** Default ask tool - returns null when no UI */ -export const askTool = createAskTool({ - cwd: process.cwd(), - hasUI: false, - getSessionFile: () => null, - getSessionSpawns: () => "*", -}); + // Single question mode (backwards compatible) + const question = params.question ?? ""; + const options = params.options ?? []; + const multi = params.multi ?? false; + const optionLabels = options.map((o) => o.label); + + if (!question || optionLabels.length === 0) { + return { + content: [{ type: "text" as const, text: "Error: question and options are required" }], + details: {}, + }; + } + + const { selectedOptions, customInput } = await askSingleQuestion(ui, question, optionLabels, multi); + + const details: AskToolDetails = { + question, + options: optionLabels, + multi, + selectedOptions, + customInput, + }; + + let responseText: string; + if (customInput) { + responseText = `User provided custom input: ${customInput}`; + } else if (selectedOptions.length > 0) { + responseText = multi ? `User selected: ${selectedOptions.join(", ")}` : `User selected: ${selectedOptions[0]}`; + } else { + responseText = "User cancelled the selection"; + } + + return { content: [{ type: "text" as const, text: responseText }], details }; + } +} // ============================================================================= // TUI Renderer diff --git a/packages/coding-agent/src/core/tools/bash.ts b/packages/coding-agent/src/core/tools/bash.ts index 445db1f5f..ae98dc34c 100644 --- a/packages/coding-agent/src/core/tools/bash.ts +++ b/packages/coding-agent/src/core/tools/bash.ts @@ -1,5 +1,5 @@ import { relative, resolve, sep } from "node:path"; -import type { AgentTool, AgentToolContext } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text, truncateToWidth } from "@oh-my-pi/pi-tui"; import { Type } from "@sinclair/typebox"; @@ -52,113 +52,122 @@ export interface BashToolOptions { operations?: BashOperations; } -export function createBashTool(session: ToolSession, options?: BashToolOptions): AgentTool { - return { - name: "bash", - label: "Bash", - description: renderPromptTemplate(bashDescription), - parameters: bashSchema, - execute: async ( - _toolCallId: string, - { command, timeout, workdir }: { command: string; timeout?: number; workdir?: string }, - signal?: AbortSignal, - onUpdate?, - ctx?: AgentToolContext, - ) => { - // Check interception if enabled and available tools are known - if (session.settings?.getBashInterceptorEnabled()) { - const rules = session.settings?.getBashInterceptorRules?.(); - const interception = checkBashInterception(command, ctx?.toolNames ?? [], rules); - if (interception.block) { - throw new Error(interception.message); - } - if (session.settings?.getBashInterceptorSimpleLsEnabled?.() !== false) { - const lsInterception = checkSimpleLsInterception(command, ctx?.toolNames ?? []); - if (lsInterception.block) { - throw new Error(lsInterception.message); - } +/** + * Bash tool implementation. + * + * Executes bash commands with optional timeout and working directory. + * Supports custom operations for remote execution. + */ +export class BashTool implements AgentTool { + public readonly name = "bash"; + public readonly label = "Bash"; + public readonly description: string; + public readonly parameters = bashSchema; + + private readonly session: ToolSession; + private readonly options?: BashToolOptions; + + constructor(session: ToolSession, options?: BashToolOptions) { + this.session = session; + this.options = options; + this.description = renderPromptTemplate(bashDescription); + } + + public async execute( + _toolCallId: string, + { command, timeout, workdir }: { command: string; timeout?: number; workdir?: string }, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + ctx?: AgentToolContext, + ): Promise> { + // Check interception if enabled and available tools are known + if (this.session.settings?.getBashInterceptorEnabled()) { + const rules = this.session.settings?.getBashInterceptorRules?.(); + const interception = checkBashInterception(command, ctx?.toolNames ?? [], rules); + if (interception.block) { + throw new Error(interception.message); + } + if (this.session.settings?.getBashInterceptorSimpleLsEnabled?.() !== false) { + const lsInterception = checkSimpleLsInterception(command, ctx?.toolNames ?? []); + if (lsInterception.block) { + throw new Error(lsInterception.message); } } + } - const commandCwd = workdir ? resolveToCwd(workdir, session.cwd) : session.cwd; - let cwdStat: Awaited>; - try { - cwdStat = await Bun.file(commandCwd).stat(); - } catch { - throw new Error(`Working directory does not exist: ${commandCwd}`); - } - if (!cwdStat.isDirectory()) { - throw new Error(`Working directory is not a directory: ${commandCwd}`); - } + const commandCwd = workdir ? resolveToCwd(workdir, this.session.cwd) : this.session.cwd; + let cwdStat: Awaited>; + try { + cwdStat = await Bun.file(commandCwd).stat(); + } catch { + throw new Error(`Working directory does not exist: ${commandCwd}`); + } + if (!cwdStat.isDirectory()) { + throw new Error(`Working directory is not a directory: ${commandCwd}`); + } - // Track output for streaming updates - let currentOutput = ""; + // Track output for streaming updates + let currentOutput = ""; - const executorOptions: BashExecutorOptions = { - cwd: commandCwd, - timeout: timeout ? timeout * 1000 : undefined, // Convert to milliseconds - signal, - onChunk: (chunk) => { - currentOutput += chunk; - if (onUpdate) { - const truncation = truncateTail(currentOutput); - onUpdate({ - content: [{ type: "text", text: truncation.content || "" }], - details: truncation.truncated - ? { - truncation, - fullOutput: currentOutput, - } - : undefined, - }); - } - }, + const executorOptions: BashExecutorOptions = { + cwd: commandCwd, + timeout: timeout ? timeout * 1000 : undefined, // Convert to milliseconds + signal, + onChunk: (chunk) => { + currentOutput += chunk; + if (onUpdate) { + const truncation = truncateTail(currentOutput); + onUpdate({ + content: [{ type: "text", text: truncation.content || "" }], + details: truncation.truncated ? { truncation, fullOutput: currentOutput } : {}, + }); + } + }, + }; + + // Use custom operations if provided, otherwise use default local executor + const result = this.options?.operations + ? await executeBashWithOperations(command, commandCwd, this.options.operations, executorOptions) + : await executeBash(command, executorOptions); + + // Handle errors + if (result.cancelled) { + throw new Error(result.output || "Command aborted"); + } + + // Apply tail truncation for final output + const truncation = truncateTail(result.output); + let outputText = truncation.content || "(no output)"; + + let details: BashToolDetails | undefined; + + if (truncation.truncated) { + details = { + truncation, + fullOutputPath: result.fullOutputPath, + fullOutput: currentOutput, }; - // Use custom operations if provided, otherwise use default local executor - const result = options?.operations - ? await executeBashWithOperations(command, commandCwd, options.operations, executorOptions) - : await executeBash(command, executorOptions); + const startLine = truncation.totalLines - truncation.outputLines + 1; + const endLine = truncation.totalLines; - // Handle errors - if (result.cancelled) { - throw new Error(result.output || "Command aborted"); + if (truncation.lastLinePartial) { + const lastLineSize = formatSize(Buffer.byteLength(result.output.split("\n").pop() || "", "utf-8")); + outputText += `\n\n[Showing last ${formatSize(truncation.outputBytes)} of line ${endLine} (line is ${lastLineSize}). Full output: ${result.fullOutputPath}]`; + } else if (truncation.truncatedBy === "lines") { + outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines}. Full output: ${result.fullOutputPath}]`; + } else { + outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines} (${formatSize(DEFAULT_MAX_BYTES)} limit). Full output: ${result.fullOutputPath}]`; } + } - // Apply tail truncation for final output - const truncation = truncateTail(result.output); - let outputText = truncation.content || "(no output)"; + if (result.exitCode !== 0 && result.exitCode !== undefined) { + outputText += `\n\nCommand exited with code ${result.exitCode}`; + throw new Error(outputText); + } - let details: BashToolDetails | undefined; - - if (truncation.truncated) { - details = { - truncation, - fullOutputPath: result.fullOutputPath, - fullOutput: currentOutput, - }; - - const startLine = truncation.totalLines - truncation.outputLines + 1; - const endLine = truncation.totalLines; - - if (truncation.lastLinePartial) { - const lastLineSize = formatSize(Buffer.byteLength(result.output.split("\n").pop() || "", "utf-8")); - outputText += `\n\n[Showing last ${formatSize(truncation.outputBytes)} of line ${endLine} (line is ${lastLineSize}). Full output: ${result.fullOutputPath}]`; - } else if (truncation.truncatedBy === "lines") { - outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines}. Full output: ${result.fullOutputPath}]`; - } else { - outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines} (${formatSize(DEFAULT_MAX_BYTES)} limit). Full output: ${result.fullOutputPath}]`; - } - } - - if (result.exitCode !== 0 && result.exitCode !== undefined) { - outputText += `\n\nCommand exited with code ${result.exitCode}`; - throw new Error(outputText); - } - - return { content: [{ type: "text", text: outputText }], details }; - }, - }; + return { content: [{ type: "text", text: outputText }], details }; + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/calculator.ts b/packages/coding-agent/src/core/tools/calculator.ts index ab7f976a9..e81d29b36 100644 --- a/packages/coding-agent/src/core/tools/calculator.ts +++ b/packages/coding-agent/src/core/tools/calculator.ts @@ -1,4 +1,4 @@ -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolResult } from "@oh-my-pi/pi-agent-core"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; import { Type } from "@sinclair/typebox"; @@ -390,32 +390,47 @@ function formatResult(value: number): string { return String(value); } -export function createCalculatorTool(_session: ToolSession): AgentTool { - return { - name: "calc", - label: "Calc", - description: renderPromptTemplate(calculatorDescription), - parameters: calculatorSchema, - execute: async ( - _toolCallId: string, - { calculations }: { calculations: Array<{ expression: string; prefix: string; suffix: string }> }, - signal?: AbortSignal, - ) => { - return untilAborted(signal, async () => { - const results = calculations.map((calc) => { - const value = evaluateExpression(calc.expression); - const output = `${calc.prefix}${formatResult(value)}${calc.suffix}`; - return { expression: calc.expression, value, output }; - }); +// ═══════════════════════════════════════════════════════════════════════════ +// Tool Class +// ═══════════════════════════════════════════════════════════════════════════ - const outputText = results.map((result) => result.output).join("\n"); - return { - content: [{ type: "text", text: outputText }], - details: { results }, - }; +type CalculatorParams = { calculations: Array<{ expression: string; prefix: string; suffix: string }> }; + +/** + * Calculator tool for evaluating mathematical expressions. + * + * Supports decimal, hex (0x), binary (0b), octal (0o) literals, + * standard arithmetic operators, and parentheses. + */ +export class CalculatorTool implements AgentTool { + public readonly name = "calc"; + public readonly label = "Calc"; + public readonly description: string; + public readonly parameters = calculatorSchema; + + constructor(_session: ToolSession) { + this.description = renderPromptTemplate(calculatorDescription); + } + + public async execute( + _toolCallId: string, + { calculations }: CalculatorParams, + signal?: AbortSignal, + ): Promise> { + return untilAborted(signal, async () => { + const results = calculations.map((calc) => { + const value = evaluateExpression(calc.expression); + const output = `${calc.prefix}${formatResult(value)}${calc.suffix}`; + return { expression: calc.expression, value, output }; }); - }, - }; + + const outputText = results.map((result) => result.output).join("\n"); + return { + content: [{ type: "text", text: outputText }], + details: { results }, + }; + }); + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/complete.ts b/packages/coding-agent/src/core/tools/complete.ts index e147c536a..476ffb686 100644 --- a/packages/coding-agent/src/core/tools/complete.ts +++ b/packages/coding-agent/src/core/tools/complete.ts @@ -4,8 +4,9 @@ * Subagents must call this tool to finish and return structured JSON output. */ -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import { StringEnum } from "@oh-my-pi/pi-ai"; +import type { Static, TObject } from "@sinclair/typebox"; import { Type } from "@sinclair/typebox"; import Ajv, { type ErrorObject, type ValidateFunction } from "ajv"; import type { ToolSession } from "./index"; @@ -52,77 +53,86 @@ function formatAjvErrors(errors: ErrorObject[] | null | undefined): string { .join("; "); } -export function createCompleteTool(session: ToolSession) { - const schemaResult = normalizeSchema(session.outputSchema); - // Convert JTD to JSON Schema if needed (auto-detected) - const normalizedSchema = - schemaResult.normalized !== undefined ? jtdToJsonSchema(schemaResult.normalized) : undefined; - let validate: ValidateFunction | undefined; - let schemaError = schemaResult.error; +export class CompleteTool implements AgentTool { + public readonly name = "complete"; + public readonly label = "Complete"; + public readonly description = + "Finish the task with structured JSON output. Call exactly once at the end of the task.\n\n" + + "If you cannot complete the task, call with status='aborted' and an error message."; + public readonly parameters: TObject; - if (normalizedSchema !== undefined && !schemaError) { - try { - validate = ajv.compile(normalizedSchema as any); - } catch (err) { - schemaError = err instanceof Error ? err.message : String(err); + private readonly validate?: ValidateFunction; + private readonly schemaError?: string; + + constructor(session: ToolSession) { + const schemaResult = normalizeSchema(session.outputSchema); + // Convert JTD to JSON Schema if needed (auto-detected) + const normalizedSchema = + schemaResult.normalized !== undefined ? jtdToJsonSchema(schemaResult.normalized) : undefined; + let schemaError = schemaResult.error; + + if (normalizedSchema !== undefined && !schemaError) { + try { + this.validate = ajv.compile(normalizedSchema as any); + } catch (err) { + schemaError = err instanceof Error ? err.message : String(err); + } } + + this.schemaError = schemaError; + + const schemaHint = formatSchema(normalizedSchema ?? session.outputSchema); + + // Use actual schema if provided, otherwise fall back to Type.Any + // Merge description into the JSON schema for better tool documentation + const dataSchema = normalizedSchema + ? Type.Unsafe({ + ...(normalizedSchema as object), + description: `Structured output matching the schema:\n${schemaHint}`, + }) + : Type.Any({ description: "Structured JSON output (no schema specified)" }); + + this.parameters = Type.Object({ + data: Type.Optional(dataSchema), + status: Type.Optional( + StringEnum(["success", "aborted"], { + description: "Use 'aborted' if the task cannot be completed, defaults to 'success'", + }), + ), + error: Type.Optional(Type.String({ description: "Error message when status is 'aborted'" })), + }); } - const schemaHint = formatSchema(normalizedSchema ?? session.outputSchema); + public async execute( + _toolCallId: string, + params: Static, + _signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const status = (params.status ?? "success") as "success" | "aborted"; - // Use actual schema if provided, otherwise fall back to Type.Any - // Merge description into the JSON schema for better tool documentation - const dataSchema = normalizedSchema - ? Type.Unsafe({ - ...(normalizedSchema as object), - description: `Structured output matching the schema:\n${schemaHint}`, - }) - : Type.Any({ description: "Structured JSON output (no schema specified)" }); - - const completeParams = Type.Object({ - data: Type.Optional(dataSchema), - status: Type.Optional( - StringEnum(["success", "aborted"], { - description: "Use 'aborted' if the task cannot be completed, defaults to 'success'", - }), - ), - error: Type.Optional(Type.String({ description: "Error message when status is 'aborted'" })), - }); - - const tool: AgentTool = { - name: "complete", - label: "Complete", - description: - "Finish the task with structured JSON output. Call exactly once at the end of the task.\n\n" + - "If you cannot complete the task, call with status='aborted' and an error message.", - parameters: completeParams, - execute: async (_toolCallId, params) => { - const status = params.status ?? "success"; - - // Skip validation when aborting - data is optional for aborts - if (status === "success") { - if (params.data === undefined) { - throw new Error("data is required when status is 'success'"); - } - if (schemaError) { - throw new Error(`Invalid output schema: ${schemaError}`); - } - if (validate && !validate(params.data)) { - throw new Error(`Output does not match schema: ${formatAjvErrors(validate.errors)}`); - } + // Skip validation when aborting - data is optional for aborts + if (status === "success") { + if (params.data === undefined) { + throw new Error("data is required when status is 'success'"); } + if (this.schemaError) { + throw new Error(`Invalid output schema: ${this.schemaError}`); + } + if (this.validate && !this.validate(params.data)) { + throw new Error(`Output does not match schema: ${formatAjvErrors(this.validate.errors)}`); + } + } - const responseText = - status === "aborted" ? `Task aborted: ${params.error || "No reason provided"}` : "Completion recorded."; + const responseText = + status === "aborted" ? `Task aborted: ${params.error || "No reason provided"}` : "Completion recorded."; - return { - content: [{ type: "text", text: responseText }], - details: { data: params.data, status, error: params.error }, - }; - }, - }; - - return tool; + return { + content: [{ type: "text", text: responseText }], + details: { data: params.data, status, error: params.error as string | undefined }, + }; + } } // Register subprocess tool handler for extraction + termination. diff --git a/packages/coding-agent/src/core/tools/context.ts b/packages/coding-agent/src/core/tools/context.ts index 4e745c24b..63a2656f6 100644 --- a/packages/coding-agent/src/core/tools/context.ts +++ b/packages/coding-agent/src/core/tools/context.ts @@ -11,31 +11,29 @@ declare module "@oh-my-pi/pi-agent-core" { } } -export interface ToolContextStore { - getContext(toolCall?: ToolCallContext): AgentToolContext; - setUIContext(uiContext: ExtensionUIContext, hasUI: boolean): void; - setToolNames(names: string[]): void; -} +export class ToolContextStore { + private uiContext: ExtensionUIContext | undefined; + private hasUI = false; + private toolNames: string[] = []; -export function createToolContextStore(getBaseContext: () => CustomToolContext): ToolContextStore { - let uiContext: ExtensionUIContext | undefined; - let hasUI = false; - let toolNames: string[] = []; + constructor(private readonly getBaseContext: () => CustomToolContext) {} - return { - getContext: (toolCall) => ({ - ...getBaseContext(), - ui: uiContext, - hasUI, - toolNames, + getContext(toolCall?: ToolCallContext): AgentToolContext { + return { + ...this.getBaseContext(), + ui: this.uiContext, + hasUI: this.hasUI, + toolNames: this.toolNames, toolCall, - }), - setUIContext: (context, uiAvailable) => { - uiContext = context; - hasUI = uiAvailable; - }, - setToolNames: (names) => { - toolNames = names; - }, - }; + }; + } + + setUIContext(uiContext: ExtensionUIContext, hasUI: boolean): void { + this.uiContext = uiContext; + this.hasUI = hasUI; + } + + setToolNames(names: string[]): void { + this.toolNames = names; + } } diff --git a/packages/coding-agent/src/core/tools/exa/index.ts b/packages/coding-agent/src/core/tools/exa/index.ts index 8c8b3c11c..238f143e2 100644 --- a/packages/coding-agent/src/core/tools/exa/index.ts +++ b/packages/coding-agent/src/core/tools/exa/index.ts @@ -48,11 +48,11 @@ export { callExaTool, callWebsetsTool, createMCPToolFromServer, - createMCPWrappedTool, fetchMCPToolSchema, findApiKey, formatSearchResults, isSearchResponse, + MCPWrappedTool, } from "./mcp-client"; export { renderExaCall, renderExaResult } from "./render"; export { researcherTools } from "./researcher"; diff --git a/packages/coding-agent/src/core/tools/exa/mcp-client.ts b/packages/coding-agent/src/core/tools/exa/mcp-client.ts index 7793c086e..1d6f3ba0e 100644 --- a/packages/coding-agent/src/core/tools/exa/mcp-client.ts +++ b/packages/coding-agent/src/core/tools/exa/mcp-client.ts @@ -1,14 +1,15 @@ /** * Exa MCP Client * - * Client for interacting with Exa MCP servers via JSON-RPC 2.0 over HTTPS. + * Client for interacting with Exa MCP servers. */ import { existsSync, readFileSync } from "node:fs"; import { homedir } from "node:os"; import type { TSchema } from "@sinclair/typebox"; -import type { CustomTool } from "../../custom-tools/types"; +import type { CustomTool, CustomToolResult } from "../../custom-tools/types"; import { logger } from "../../logger"; +import { callMCP } from "../../mcp/json-rpc"; import type { ExaRenderDetails, ExaSearchResponse, @@ -48,63 +49,6 @@ export async function findApiKey(): Promise { return null; } -/** Parse SSE response format (lines starting with "data: ") */ -function parseSSE(text: string): unknown { - const lines = text.split("\n"); - for (const line of lines) { - if (line.startsWith("data: ")) { - const data = line.slice(6).trim(); - if (data === "[DONE]") continue; - try { - return JSON.parse(data); - } catch { - // Try next line - } - } - } - // Fallback: try parsing entire response as JSON - try { - return JSON.parse(text); - } catch { - return null; - } -} - -/** Call MCP server with JSON-RPC 2.0 */ -export async function callMCP(url: string, method: string, params?: Record): Promise { - const body = { - jsonrpc: "2.0", - id: Math.random().toString(36).slice(2), - method, - params: params ?? {}, - }; - - const response = await fetch(url, { - method: "POST", - headers: { - "Content-Type": "application/json", - Accept: "application/json, text/event-stream", - }, - body: JSON.stringify(body), - }); - - if (!response.ok) { - const errorMsg = `MCP request failed: ${response.status} ${response.statusText}`; - logger.error(errorMsg, { url, method, params }); - throw new Error(errorMsg); - } - - const text = await response.text(); - const result = parseSSE(text); - - if (!result) { - logger.error("Failed to parse MCP response", { url, method, responseText: text.slice(0, 500) }); - throw new Error("Failed to parse MCP response"); - } - - return result; -} - /** Fetch available tools from Exa MCP */ export async function fetchExaTools(apiKey: string, toolNames: string[]): Promise { const url = `https://mcp.exa.ai/mcp?exaApiKey=${encodeURIComponent(apiKey)}&toolNames=${encodeURIComponent(toolNames.join(","))}`; @@ -299,56 +243,67 @@ export async function fetchMCPToolSchema( } /** - * Create a CustomTool dynamically from MCP tool metadata. + * CustomTool dynamically created from MCP tool metadata. * * This allows tools to be generated from MCP server schemas without hardcoding, * reducing drift when MCP servers add new parameters. */ -export function createMCPWrappedTool( - config: MCPToolWrapperConfig, - schema: TSchema, - description: string, -): CustomTool { - return { - name: config.name, - label: config.label, - description, - parameters: schema, - async execute(_toolCallId, params, _onUpdate, _ctx, _signal) { - try { - const apiKey = await findApiKey(); - if (!apiKey) { - return { - content: [{ type: "text" as const, text: "Error: EXA_API_KEY not found" }], - details: { error: "EXA_API_KEY not found", toolName: config.name }, - }; - } +export class MCPWrappedTool implements CustomTool { + public readonly name: string; + public readonly label: string; + public readonly description: string; + public readonly parameters: TSchema; - const response = config.isWebsetsTool - ? await callWebsetsTool(apiKey, config.mcpToolName, params as Record) - : await callExaTool(config.mcpToolName, params as Record, apiKey); + private readonly config: MCPToolWrapperConfig; - if (isSearchResponse(response)) { - const formatted = formatSearchResults(response); - return { - content: [{ type: "text" as const, text: formatted }], - details: { response, toolName: config.name }, - }; - } + constructor(config: MCPToolWrapperConfig, schema: TSchema, description: string) { + this.config = config; + this.name = config.name; + this.label = config.label; + this.description = description; + this.parameters = schema; + } + async execute( + _toolCallId: string, + params: unknown, + _onUpdate?: unknown, + _ctx?: unknown, + _signal?: AbortSignal, + ): Promise> { + try { + const apiKey = await findApiKey(); + if (!apiKey) { return { - content: [{ type: "text" as const, text: JSON.stringify(response, null, 2) }], - details: { raw: response, toolName: config.name }, - }; - } catch (error) { - const message = error instanceof Error ? error.message : String(error); - return { - content: [{ type: "text" as const, text: `Error: ${message}` }], - details: { error: message, toolName: config.name }, + content: [{ type: "text" as const, text: "Error: EXA_API_KEY not found" }], + details: { error: "EXA_API_KEY not found", toolName: this.config.name }, }; } - }, - }; + + const response = this.config.isWebsetsTool + ? await callWebsetsTool(apiKey, this.config.mcpToolName, params as Record) + : await callExaTool(this.config.mcpToolName, params as Record, apiKey); + + if (isSearchResponse(response)) { + const formatted = formatSearchResults(response); + return { + content: [{ type: "text" as const, text: formatted }], + details: { response, toolName: this.config.name }, + }; + } + + return { + content: [{ type: "text" as const, text: JSON.stringify(response, null, 2) }], + details: { raw: response, toolName: this.config.name }, + }; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + return { + content: [{ type: "text" as const, text: `Error: ${message}` }], + details: { error: message, toolName: this.config.name }, + }; + } + } } /** @@ -361,9 +316,9 @@ export async function createMCPToolFromServer( config: MCPToolWrapperConfig, fallbackSchema: TSchema, fallbackDescription: string, -): Promise> { +): Promise { const mcpTool = await fetchMCPToolSchema(apiKey, config.mcpToolName, config.isWebsetsTool); const schema = mcpTool?.inputSchema ?? fallbackSchema; const description = mcpTool?.description ?? fallbackDescription; - return createMCPWrappedTool(config, schema, description); + return new MCPWrappedTool(config, schema, description); } diff --git a/packages/coding-agent/src/core/tools/find.ts b/packages/coding-agent/src/core/tools/find.ts index 58314cda1..6a03d49aa 100644 --- a/packages/coding-agent/src/core/tools/find.ts +++ b/packages/coding-agent/src/core/tools/find.ts @@ -1,8 +1,9 @@ import path from "node:path"; -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import { StringEnum } from "@oh-my-pi/pi-ai"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; +import type { Static } from "@sinclair/typebox"; import { Type } from "@sinclair/typebox"; import { getLanguageFromPath, type Theme } from "../../modes/interactive/theme/theme"; import findDescription from "../../prompts/tools/find.md" with { type: "text" }; @@ -109,264 +110,68 @@ async function captureCommandOutput( return { stdout, stderr, exitCode, aborted: scope.aborted }; } -export function createFindTool(session: ToolSession, options?: FindToolOptions): AgentTool { - const customOps = options?.operations; +export class FindTool implements AgentTool { + public readonly name = "find"; + public readonly label = "Find"; + public readonly description: string; + public readonly parameters = findSchema; - return { - name: "find", - label: "Find", - description: renderPromptTemplate(findDescription), - parameters: findSchema, - execute: async ( - _toolCallId: string, - { - pattern, - path: searchDir, - limit, - hidden, - sortByMtime, - type, - }: { - pattern: string; - path?: string; - limit?: number; - hidden?: boolean; - sortByMtime?: boolean; - type?: "file" | "dir" | "all"; - }, - signal?: AbortSignal, - ) => { - return untilAborted(signal, async () => { - const searchPath = resolveToCwd(searchDir || ".", session.cwd); - const scopePath = (() => { - const relative = path.relative(session.cwd, searchPath).replace(/\\/g, "/"); - return relative.length === 0 ? "." : relative; - })(); - const effectiveLimit = limit ?? DEFAULT_LIMIT; - const effectiveType = type ?? "all"; - const includeHidden = hidden ?? true; - const shouldSortByMtime = sortByMtime ?? false; + private readonly session: ToolSession; + private readonly customOps?: FindOperations; - // If custom operations provided with glob, use that instead of fd - if (customOps?.glob) { - if (!(await customOps.exists(searchPath))) { - throw new Error(`Path not found: ${searchPath}`); - } + constructor(session: ToolSession, options?: FindToolOptions) { + this.session = session; + this.customOps = options?.operations; + this.description = renderPromptTemplate(findDescription); + } - const results = await customOps.glob(pattern, searchPath, { - ignore: ["**/node_modules/**", "**/.git/**"], - limit: effectiveLimit, - }); + public async execute( + _toolCallId: string, + params: Static, + signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const { pattern, path: searchDir, limit, hidden, sortByMtime, type } = params; - if (results.length === 0) { - return { - content: [{ type: "text", text: "No files found matching pattern" }], - details: { scopePath, fileCount: 0, files: [], truncated: false }, - }; - } + return untilAborted(signal, async () => { + const searchPath = resolveToCwd(searchDir || ".", this.session.cwd); + const scopePath = (() => { + const relative = path.relative(this.session.cwd, searchPath).replace(/\\/g, "/"); + return relative.length === 0 ? "." : relative; + })(); + const effectiveLimit = limit ?? DEFAULT_LIMIT; + const effectiveType = type ?? "all"; + const includeHidden = hidden ?? true; + const shouldSortByMtime = sortByMtime ?? false; - // Relativize paths - const relativized = results.map((p) => { - if (p.startsWith(searchPath)) { - return p.slice(searchPath.length + 1); - } - return path.relative(searchPath, p); - }); - - const resultLimitReached = relativized.length >= effectiveLimit; - const rawOutput = relativized.join("\n"); - const truncation = truncateHead(rawOutput, { maxLines: Number.MAX_SAFE_INTEGER }); - - let resultOutput = truncation.content; - const details: FindToolDetails = { - scopePath, - fileCount: relativized.length, - files: relativized, - truncated: resultLimitReached || truncation.truncated, - }; - const notices: string[] = []; - - if (resultLimitReached) { - notices.push( - `${effectiveLimit} results limit reached. Use limit=${effectiveLimit * 2} for more, or refine pattern`, - ); - details.resultLimitReached = effectiveLimit; - } - - if (truncation.truncated) { - notices.push(`${formatSize(DEFAULT_MAX_BYTES)} limit reached`); - details.truncation = truncation; - } - - if (notices.length > 0) { - resultOutput += `\n\n[${notices.join(". ")}]`; - } - - return { - content: [{ type: "text", text: resultOutput }], - details: Object.keys(details).length > 0 ? details : undefined, - }; + // If custom operations provided with glob, use that instead of fd + if (this.customOps?.glob) { + if (!(await this.customOps.exists(searchPath))) { + throw new Error(`Path not found: ${searchPath}`); } - // Default: use fd - const fdPath = await ensureTool("fd", true); - if (!fdPath) { - throw new Error("fd is not available and could not be downloaded"); - } + const results = await this.customOps.glob(pattern, searchPath, { + ignore: ["**/node_modules/**", "**/.git/**"], + limit: effectiveLimit, + }); - // Build fd arguments - // When pattern contains path separators (e.g. "reports/**"), use --full-path - // so fd matches against the full path, not just the filename. - // Also prepend **/ to anchor the pattern at any depth in the search path. - // Note: "**/foo.rs" is a glob construct (filename at any depth), not a path. - // Only patterns with real path components like "foo/bar" or "foo/**/bar" need --full-path. - const patternWithoutLeadingStarStar = pattern.replace(/^\*\*\//, ""); - const hasPathSeparator = - patternWithoutLeadingStarStar.includes("/") || patternWithoutLeadingStarStar.includes("\\"); - const effectivePattern = hasPathSeparator && !pattern.startsWith("**/") ? `**/${pattern}` : pattern; - const args: string[] = [ - "--glob", // Use glob pattern - ...(hasPathSeparator ? ["--full-path"] : []), - "--color=never", // No ANSI colors - "--max-results", - String(effectiveLimit), - ]; - - if (includeHidden) { - args.push("--hidden"); - } - - // Add type filter - if (effectiveType === "file") { - args.push("--type", "f"); - } else if (effectiveType === "dir") { - args.push("--type", "d"); - } - - // Include .gitignore files (root + nested) so fd respects them even outside git repos - const gitignoreFiles = new Set(); - const rootGitignore = path.join(searchPath, ".gitignore"); - if (await Bun.file(rootGitignore).exists()) { - gitignoreFiles.add(rootGitignore); - } - - try { - const gitignoreArgs = [ - "--hidden", - "--no-ignore", - "--type", - "f", - "--name", - ".gitignore", - "--exclude", - ".git", - "--exclude", - "node_modules", - "--absolute-path", - searchPath, - ]; - const { stdout: gitignoreStdout, aborted: gitignoreAborted } = await captureCommandOutput( - fdPath, - gitignoreArgs, - signal, - ); - if (gitignoreAborted) { - throw new Error("Operation aborted"); - } - for (const rawLine of gitignoreStdout.split("\n")) { - const file = rawLine.trim(); - if (!file) continue; - gitignoreFiles.add(file); - } - } catch (err) { - if (signal?.aborted) { - throw err instanceof Error ? err : new Error("Operation aborted"); - } - // Ignore lookup errors - } - - for (const gitignorePath of gitignoreFiles) { - args.push("--ignore-file", gitignorePath); - } - - // Pattern and path - args.push(effectivePattern, searchPath); - - // Run fd - const { stdout, stderr, exitCode, aborted } = await captureCommandOutput(fdPath, args, signal); - - if (aborted) { - throw new Error("Operation aborted"); - } - - const output = stdout.trim(); - - if (exitCode !== 0) { - const errorMsg = stderr.trim() || `fd exited with code ${exitCode ?? -1}`; - // fd returns non-zero for some errors but may still have partial output - if (!output) { - throw new Error(errorMsg); - } - } - - if (!output) { + if (results.length === 0) { return { content: [{ type: "text", text: "No files found matching pattern" }], details: { scopePath, fileCount: 0, files: [], truncated: false }, }; } - const lines = output.split("\n"); - const relativized: string[] = []; - const mtimes: number[] = []; - - for (const rawLine of lines) { - signal?.throwIfAborted(); - const line = rawLine.replace(/\r$/, "").trim(); - if (!line) { - continue; + // Relativize paths + const relativized = results.map((p) => { + if (p.startsWith(searchPath)) { + return p.slice(searchPath.length + 1); } + return path.relative(searchPath, p); + }); - const hadTrailingSlash = line.endsWith("/") || line.endsWith("\\"); - let relativePath = line; - if (line.startsWith(searchPath)) { - relativePath = line.slice(searchPath.length + 1); // +1 for the / - } else { - relativePath = path.relative(searchPath, line); - } - - if (hadTrailingSlash && !relativePath.endsWith("/")) { - relativePath += "/"; - } - - // When sorting by mtime, keep files that fail to stat with mtime 0 - if (shouldSortByMtime) { - try { - const fullPath = path.join(searchPath, relativePath); - const stat = await Bun.file(fullPath).stat(); - relativized.push(relativePath); - mtimes.push(stat.mtimeMs); - } catch { - relativized.push(relativePath); - mtimes.push(0); - } - } else { - relativized.push(relativePath); - } - } - - // Sort by mtime if requested (most recent first) - if (shouldSortByMtime && relativized.length > 0) { - const indexed = relativized.map((path, idx) => ({ path, mtime: mtimes[idx] })); - indexed.sort((a, b) => b.mtime - a.mtime); - relativized.length = 0; - relativized.push(...indexed.map((item) => item.path)); - } - - // Check if we hit the result limit const resultLimitReached = relativized.length >= effectiveLimit; - - // Apply byte truncation (no line limit since we already have result limit) const rawOutput = relativized.join("\n"); const truncation = truncateHead(rawOutput, { maxLines: Number.MAX_SAFE_INTEGER }); @@ -377,8 +182,6 @@ export function createFindTool(session: ToolSession, options?: FindToolOptions): files: relativized, truncated: resultLimitReached || truncation.truncated, }; - - // Build notices const notices: string[] = []; if (resultLimitReached) { @@ -399,11 +202,205 @@ export function createFindTool(session: ToolSession, options?: FindToolOptions): return { content: [{ type: "text", text: resultOutput }], - details: Object.keys(details).length > 0 ? details : undefined, + details, }; - }); - }, - }; + } + + // Default: use fd + const fdPath = await ensureTool("fd", true); + if (!fdPath) { + throw new Error("fd is not available and could not be downloaded"); + } + + // Build fd arguments + // When pattern contains path separators (e.g. "reports/**"), use --full-path + // so fd matches against the full path, not just the filename. + // Also prepend **/ to anchor the pattern at any depth in the search path. + // Note: "**/foo.rs" is a glob construct (filename at any depth), not a path. + // Only patterns with real path components like "foo/bar" or "foo/**/bar" need --full-path. + const patternWithoutLeadingStarStar = pattern.replace(/^\*\*\//, ""); + const hasPathSeparator = + patternWithoutLeadingStarStar.includes("/") || patternWithoutLeadingStarStar.includes("\\"); + const effectivePattern = hasPathSeparator && !pattern.startsWith("**/") ? `**/${pattern}` : pattern; + const args: string[] = [ + "--glob", // Use glob pattern + ...(hasPathSeparator ? ["--full-path"] : []), + "--color=never", // No ANSI colors + "--max-results", + String(effectiveLimit), + ]; + + if (includeHidden) { + args.push("--hidden"); + } + + // Add type filter + if (effectiveType === "file") { + args.push("--type", "f"); + } else if (effectiveType === "dir") { + args.push("--type", "d"); + } + + // Include .gitignore files (root + nested) so fd respects them even outside git repos + const gitignoreFiles = new Set(); + const rootGitignore = path.join(searchPath, ".gitignore"); + if (await Bun.file(rootGitignore).exists()) { + gitignoreFiles.add(rootGitignore); + } + + try { + const gitignoreArgs = [ + "--hidden", + "--no-ignore", + "--type", + "f", + "--name", + ".gitignore", + "--exclude", + ".git", + "--exclude", + "node_modules", + "--absolute-path", + searchPath, + ]; + const { stdout: gitignoreStdout, aborted: gitignoreAborted } = await captureCommandOutput( + fdPath, + gitignoreArgs, + signal, + ); + if (gitignoreAborted) { + throw new Error("Operation aborted"); + } + for (const rawLine of gitignoreStdout.split("\n")) { + const file = rawLine.trim(); + if (!file) continue; + gitignoreFiles.add(file); + } + } catch (err) { + if (signal?.aborted) { + throw err instanceof Error ? err : new Error("Operation aborted"); + } + // Ignore lookup errors + } + + for (const gitignorePath of gitignoreFiles) { + args.push("--ignore-file", gitignorePath); + } + + // Pattern and path + args.push(effectivePattern, searchPath); + + // Run fd + const { stdout, stderr, exitCode, aborted } = await captureCommandOutput(fdPath, args, signal); + + if (aborted) { + throw new Error("Operation aborted"); + } + + const output = stdout.trim(); + + if (exitCode !== 0) { + const errorMsg = stderr.trim() || `fd exited with code ${exitCode ?? -1}`; + // fd returns non-zero for some errors but may still have partial output + if (!output) { + throw new Error(errorMsg); + } + } + + if (!output) { + return { + content: [{ type: "text", text: "No files found matching pattern" }], + details: { scopePath, fileCount: 0, files: [], truncated: false }, + }; + } + + const lines = output.split("\n"); + const relativized: string[] = []; + const mtimes: number[] = []; + + for (const rawLine of lines) { + signal?.throwIfAborted(); + const line = rawLine.replace(/\r$/, "").trim(); + if (!line) { + continue; + } + + const hadTrailingSlash = line.endsWith("/") || line.endsWith("\\"); + let relativePath = line; + if (line.startsWith(searchPath)) { + relativePath = line.slice(searchPath.length + 1); // +1 for the / + } else { + relativePath = path.relative(searchPath, line); + } + + if (hadTrailingSlash && !relativePath.endsWith("/")) { + relativePath += "/"; + } + + // When sorting by mtime, keep files that fail to stat with mtime 0 + if (shouldSortByMtime) { + try { + const fullPath = path.join(searchPath, relativePath); + const stat = await Bun.file(fullPath).stat(); + relativized.push(relativePath); + mtimes.push(stat.mtimeMs); + } catch { + relativized.push(relativePath); + mtimes.push(0); + } + } else { + relativized.push(relativePath); + } + } + + // Sort by mtime if requested (most recent first) + if (shouldSortByMtime && relativized.length > 0) { + const indexed = relativized.map((path, idx) => ({ path, mtime: mtimes[idx] })); + indexed.sort((a, b) => b.mtime - a.mtime); + relativized.length = 0; + relativized.push(...indexed.map((item) => item.path)); + } + + // Check if we hit the result limit + const resultLimitReached = relativized.length >= effectiveLimit; + + // Apply byte truncation (no line limit since we already have result limit) + const rawOutput = relativized.join("\n"); + const truncation = truncateHead(rawOutput, { maxLines: Number.MAX_SAFE_INTEGER }); + + let resultOutput = truncation.content; + const details: FindToolDetails = { + scopePath, + fileCount: relativized.length, + files: relativized, + truncated: resultLimitReached || truncation.truncated, + }; + + // Build notices + const notices: string[] = []; + + if (resultLimitReached) { + notices.push( + `${effectiveLimit} results limit reached. Use limit=${effectiveLimit * 2} for more, or refine pattern`, + ); + details.resultLimitReached = effectiveLimit; + } + + if (truncation.truncated) { + notices.push(`${formatSize(DEFAULT_MAX_BYTES)} limit reached`); + details.truncation = truncation; + } + + if (notices.length > 0) { + resultOutput += `\n\n[${notices.join(". ")}]`; + } + + return { + content: [{ type: "text", text: resultOutput }], + details, + }; + }); + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/git.ts b/packages/coding-agent/src/core/tools/git.ts index 60dfc6aab..4723680e0 100644 --- a/packages/coding-agent/src/core/tools/git.ts +++ b/packages/coding-agent/src/core/tools/git.ts @@ -1,4 +1,4 @@ -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import { StringEnum } from "@oh-my-pi/pi-ai"; import { type GitParams, gitTool as gitToolCore, type ToolResponse } from "@oh-my-pi/pi-git-tool"; import { type Static, Type } from "@sinclair/typebox"; @@ -171,37 +171,43 @@ const gitSchema = Type.Object({ export type GitToolDetails = ToolResponse; -export function createGitTool(session: ToolSession): AgentTool | null { - if (session.settings?.getGitToolEnabled() === false) { - return null; +export class GitTool implements AgentTool { + public readonly name = "git"; + public readonly label = "Git"; + public readonly description: string; + public readonly parameters = gitSchema; + + private readonly session: ToolSession; + + constructor(session: ToolSession) { + this.session = session; + this.description = renderPromptTemplate(gitDescription); } - return { - name: "git", - label: "Git", - description: renderPromptTemplate(gitDescription), - parameters: gitSchema, - execute: async (_toolCallId, params: Static, _signal?: AbortSignal) => { - if (params.operation === "commit" && !params.message) { - throw new Error("Git commit requires a message to avoid an interactive editor. Provide `message`."); - } - const result = await gitToolCore(params as GitParams, session.cwd); - if ("error" in result) { - const message = result._rendered ?? result.error; - return { content: [{ type: "text", text: message }], details: result }; - } - if ("confirm" in result) { - const message = result._rendered ?? result.confirm; - return { content: [{ type: "text", text: message }], details: result }; - } - return { content: [{ type: "text", text: result._rendered }], details: result }; - }, - }; + static createIf(session: ToolSession): GitTool | null { + return session.settings?.getGitToolEnabled() === false ? null : new GitTool(session); + } + + public async execute( + _toolCallId: string, + params: Static, + _signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + if (params.operation === "commit" && !params.message) { + throw new Error("Git commit requires a message to avoid an interactive editor. Provide `message`."); + } + + const result = await gitToolCore(params as GitParams, this.session.cwd); + if ("error" in result) { + const message = result._rendered ?? result.error; + return { content: [{ type: "text", text: message }], details: result }; + } + if ("confirm" in result) { + const message = result._rendered ?? result.confirm; + return { content: [{ type: "text", text: message }], details: result }; + } + return { content: [{ type: "text", text: result._rendered }], details: result }; + } } - -export const gitTool = createGitTool({ - cwd: process.cwd(), - hasUI: false, - getSessionFile: () => null, - getSessionSpawns: () => null, -})!; diff --git a/packages/coding-agent/src/core/tools/grep.ts b/packages/coding-agent/src/core/tools/grep.ts index 7bcaa488e..248cc23b5 100644 --- a/packages/coding-agent/src/core/tools/grep.ts +++ b/packages/coding-agent/src/core/tools/grep.ts @@ -1,5 +1,5 @@ import nodePath from "node:path"; -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import { StringEnum } from "@oh-my-pi/pi-ai"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; @@ -93,416 +93,204 @@ export interface GrepToolOptions { operations?: GrepOperations; } -export function createGrepTool(session: ToolSession, options?: GrepToolOptions): AgentTool { - const ops = options?.operations ?? defaultGrepOperations; - return { - name: "grep", - label: "Grep", - description: renderPromptTemplate(grepDescription), - parameters: grepSchema, - execute: async ( - _toolCallId: string, - { - pattern, - path: searchDir, - glob, - type, - ignoreCase, - caseSensitive, - literal, - multiline, - context, - limit, - outputMode, - headLimit, - offset, - }: { - pattern: string; - path?: string; - glob?: string; - type?: string; - ignoreCase?: boolean; - caseSensitive?: boolean; - literal?: boolean; - multiline?: boolean; - context?: number; - limit?: number; - outputMode?: "content" | "files_with_matches" | "count"; - headLimit?: number; - offset?: number; - }, - signal?: AbortSignal, - ) => { - return untilAborted(signal, async () => { - const rgPath = await ensureTool("rg", true); - if (!rgPath) { - throw new Error("ripgrep (rg) is not available and could not be downloaded"); +interface GrepParams { + pattern: string; + path?: string; + glob?: string; + type?: string; + ignoreCase?: boolean; + caseSensitive?: boolean; + literal?: boolean; + multiline?: boolean; + context?: number; + limit?: number; + outputMode?: "content" | "files_with_matches" | "count"; + headLimit?: number; + offset?: number; +} + +export class GrepTool implements AgentTool { + public readonly name = "grep"; + public readonly label = "Grep"; + public readonly description: string; + public readonly parameters = grepSchema; + + private readonly session: ToolSession; + private readonly ops: GrepOperations; + + constructor(session: ToolSession, options?: GrepToolOptions) { + this.session = session; + this.ops = options?.operations ?? defaultGrepOperations; + this.description = renderPromptTemplate(grepDescription); + } + + public async execute( + _toolCallId: string, + params: GrepParams, + signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const { + pattern, + path: searchDir, + glob, + type, + ignoreCase, + caseSensitive, + literal, + multiline, + context, + limit, + outputMode, + headLimit, + offset, + } = params; + + return untilAborted(signal, async () => { + const rgPath = await ensureTool("rg", true); + if (!rgPath) { + throw new Error("ripgrep (rg) is not available and could not be downloaded"); + } + + const searchPath = resolveToCwd(searchDir || ".", this.session.cwd); + const scopePath = (() => { + const relative = nodePath.relative(this.session.cwd, searchPath).replace(/\\/g, "/"); + return relative.length === 0 ? "." : relative; + })(); + + let isDirectory: boolean; + try { + isDirectory = await this.ops.isDirectory(searchPath); + } catch { + throw new Error(`Path not found: ${searchPath}`); + } + const contextValue = context && context > 0 ? context : 0; + const effectiveLimit = Math.max(1, limit ?? DEFAULT_LIMIT); + const effectiveOutputMode = outputMode ?? "content"; + const effectiveOffset = offset && offset > 0 ? offset : 0; + const hasHeadLimit = headLimit !== undefined && headLimit > 0; + + const formatPath = (filePath: string): string => { + if (isDirectory) { + const relative = nodePath.relative(searchPath, filePath); + if (relative && !relative.startsWith("..")) { + return relative.replace(/\\/g, "/"); + } } + return nodePath.basename(filePath); + }; - const searchPath = resolveToCwd(searchDir || ".", session.cwd); - const scopePath = (() => { - const relative = nodePath.relative(session.cwd, searchPath).replace(/\\/g, "/"); - return relative.length === 0 ? "." : relative; - })(); - - let isDirectory: boolean; - try { - isDirectory = await ops.isDirectory(searchPath); - } catch { - throw new Error(`Path not found: ${searchPath}`); - } - const contextValue = context && context > 0 ? context : 0; - const effectiveLimit = Math.max(1, limit ?? DEFAULT_LIMIT); - const effectiveOutputMode = outputMode ?? "content"; - const effectiveOffset = offset && offset > 0 ? offset : 0; - const hasHeadLimit = headLimit !== undefined && headLimit > 0; - - const formatPath = (filePath: string): string => { - if (isDirectory) { - const relative = nodePath.relative(searchPath, filePath); - if (relative && !relative.startsWith("..")) { - return relative.replace(/\\/g, "/"); + const fileCache = new Map>(); + const getFileLines = async (filePath: string): Promise => { + let linesPromise = fileCache.get(filePath); + if (!linesPromise) { + linesPromise = (async () => { + try { + const content = await this.ops.readFile(filePath); + return content.replace(/\r\n/g, "\n").replace(/\r/g, "\n").split("\n"); + } catch { + return []; } - } - return nodePath.basename(filePath); - }; - - const fileCache = new Map>(); - const getFileLines = async (filePath: string): Promise => { - let linesPromise = fileCache.get(filePath); - if (!linesPromise) { - linesPromise = (async () => { - try { - const content = await ops.readFile(filePath); - return content.replace(/\r\n/g, "\n").replace(/\r/g, "\n").split("\n"); - } catch { - return []; - } - })(); - fileCache.set(filePath, linesPromise); - } - return linesPromise; - }; - - const args: string[] = []; - - // Base arguments depend on output mode - if (effectiveOutputMode === "files_with_matches") { - args.push("--files-with-matches", "--color=never", "--hidden"); - } else if (effectiveOutputMode === "count") { - args.push("--count", "--color=never", "--hidden"); - } else { - args.push("--json", "--line-number", "--color=never", "--hidden"); + })(); + fileCache.set(filePath, linesPromise); } + return linesPromise; + }; - if (caseSensitive) { - args.push("--case-sensitive"); - } else if (ignoreCase) { - args.push("--ignore-case"); - } else { - args.push("--smart-case"); + const args: string[] = []; + + // Base arguments depend on output mode + if (effectiveOutputMode === "files_with_matches") { + args.push("--files-with-matches", "--color=never", "--hidden"); + } else if (effectiveOutputMode === "count") { + args.push("--count", "--color=never", "--hidden"); + } else { + args.push("--json", "--line-number", "--color=never", "--hidden"); + } + + if (caseSensitive) { + args.push("--case-sensitive"); + } else if (ignoreCase) { + args.push("--ignore-case"); + } else { + args.push("--smart-case"); + } + + if (multiline) { + args.push("--multiline"); + } + + if (literal) { + args.push("--fixed-strings"); + } + + if (glob) { + args.push("--glob", glob); + } + + if (type) { + args.push("--type", type); + } + + args.push("--", pattern, searchPath); + + const child: Subprocess = Bun.spawn([rgPath, ...args], { + stdin: "ignore", + stdout: "pipe", + stderr: "pipe", + }); + + let stderr = ""; + let matchCount = 0; + let matchLimitReached = false; + let linesTruncated = false; + let aborted = false; + let killedDueToLimit = false; + const outputLines: string[] = []; + const files = new Set(); + const fileList: string[] = []; + const fileMatchCounts = new Map(); + + const recordFile = (filePath: string) => { + const relative = formatPath(filePath); + if (!files.has(relative)) { + files.add(relative); + fileList.push(relative); } + }; - if (multiline) { - args.push("--multiline"); - } + const recordFileMatch = (filePath: string) => { + const relative = formatPath(filePath); + fileMatchCounts.set(relative, (fileMatchCounts.get(relative) ?? 0) + 1); + }; - if (literal) { - args.push("--fixed-strings"); - } + const stopChild = (dueToLimit: boolean = false) => { + killedDueToLimit = dueToLimit; + child.kill(); + }; - if (glob) { - args.push("--glob", glob); - } + using signalScope = new ScopeSignal(signal ? { signal } : undefined); + signalScope.catch(() => { + aborted = true; + stopChild(); + }); - if (type) { - args.push("--type", type); - } - - args.push("--", pattern, searchPath); - - const child: Subprocess = Bun.spawn([rgPath, ...args], { - stdin: "ignore", - stdout: "pipe", - stderr: "pipe", - }); - - let stderr = ""; - let matchCount = 0; - let matchLimitReached = false; - let linesTruncated = false; - let aborted = false; - let killedDueToLimit = false; - const outputLines: string[] = []; - const files = new Set(); - const fileList: string[] = []; - const fileMatchCounts = new Map(); - - const recordFile = (filePath: string) => { - const relative = formatPath(filePath); - if (!files.has(relative)) { - files.add(relative); - fileList.push(relative); - } - }; - - const recordFileMatch = (filePath: string) => { - const relative = formatPath(filePath); - fileMatchCounts.set(relative, (fileMatchCounts.get(relative) ?? 0) + 1); - }; - - const stopChild = (dueToLimit: boolean = false) => { - killedDueToLimit = dueToLimit; - child.kill(); - }; - - using signalScope = new ScopeSignal(signal ? { signal } : undefined); - signalScope.catch(() => { - aborted = true; - stopChild(); - }); - - // For simple output modes (files_with_matches, count), process text directly - if (effectiveOutputMode === "files_with_matches" || effectiveOutputMode === "count") { - const stdoutReader = (child.stdout as ReadableStream).getReader(); - const stderrReader = (child.stderr as ReadableStream).getReader(); - const decoder = new TextDecoder(); - let stdout = ""; - - await Promise.all([ - (async () => { - while (true) { - const { done, value } = await stdoutReader.read(); - if (done) break; - stdout += decoder.decode(value, { stream: true }); - } - })(), - (async () => { - while (true) { - const { done, value } = await stderrReader.read(); - if (done) break; - stderr += decoder.decode(value, { stream: true }); - } - })(), - ]); - - const exitCode = await child.exited; - - if (aborted) { - throw new Error("Operation aborted"); - } - - if (exitCode !== 0 && exitCode !== 1) { - const errorMsg = stderr.trim() || `ripgrep exited with code ${exitCode}`; - throw new Error(errorMsg); - } - - const lines = stdout - .trim() - .split("\n") - .filter((line) => line.length > 0); - - if (lines.length === 0) { - return { - content: [{ type: "text", text: "No matches found" }], - details: { - scopePath, - matchCount: 0, - fileCount: 0, - files: [], - mode: effectiveOutputMode, - truncated: false, - }, - }; - } - - // Apply offset and headLimit - let processedLines = lines; - if (effectiveOffset > 0) { - processedLines = processedLines.slice(effectiveOffset); - } - if (hasHeadLimit) { - processedLines = processedLines.slice(0, headLimit); - } - - let simpleMatchCount = 0; - let fileCount = 0; - const simpleFiles = new Set(); - const simpleFileList: string[] = []; - const simpleFileMatchCounts = new Map(); - - const recordSimpleFile = (filePath: string) => { - const relative = formatPath(filePath); - if (!simpleFiles.has(relative)) { - simpleFiles.add(relative); - simpleFileList.push(relative); - } - }; - - // Count mode: ripgrep provides total count per file, so we set directly (not increment) - const setFileMatchCount = (filePath: string, count: number) => { - const relative = formatPath(filePath); - simpleFileMatchCounts.set(relative, count); - }; - - if (effectiveOutputMode === "files_with_matches") { - for (const line of lines) { - recordSimpleFile(line); - } - fileCount = simpleFiles.size; - simpleMatchCount = fileCount; - } else { - for (const line of lines) { - const separatorIndex = line.lastIndexOf(":"); - const filePart = separatorIndex === -1 ? line : line.slice(0, separatorIndex); - const countPart = separatorIndex === -1 ? "" : line.slice(separatorIndex + 1); - const count = Number.parseInt(countPart, 10); - recordSimpleFile(filePart); - if (!Number.isNaN(count)) { - simpleMatchCount += count; - setFileMatchCount(filePart, count); - } - } - fileCount = simpleFiles.size; - } - - const truncatedByHeadLimit = hasHeadLimit && processedLines.length < lines.length; - - // For count mode, format as "path:count" - if (effectiveOutputMode === "count") { - const formatted = processedLines.map((line) => { - const separatorIndex = line.lastIndexOf(":"); - const relative = formatPath(separatorIndex === -1 ? line : line.slice(0, separatorIndex)); - const count = separatorIndex === -1 ? "0" : line.slice(separatorIndex + 1); - return `${relative}:${count}`; - }); - const output = formatted.join("\n"); - return { - content: [{ type: "text", text: output }], - details: { - scopePath, - matchCount: simpleMatchCount, - fileCount, - files: simpleFileList, - fileMatches: simpleFileList.map((path) => ({ - path, - count: simpleFileMatchCounts.get(path) ?? 0, - })), - mode: effectiveOutputMode, - truncated: truncatedByHeadLimit, - headLimitReached: truncatedByHeadLimit ? headLimit : undefined, - }, - }; - } - - // For files_with_matches, format paths - const formatted = processedLines.map((line) => formatPath(line)); - const output = formatted.join("\n"); - return { - content: [{ type: "text", text: output }], - details: { - scopePath, - matchCount: simpleMatchCount, - fileCount, - files: simpleFileList, - mode: effectiveOutputMode, - truncated: truncatedByHeadLimit, - headLimitReached: truncatedByHeadLimit ? headLimit : undefined, - }, - }; - } - - // Content mode - existing JSON processing - const formatBlock = async (filePath: string, lineNumber: number): Promise => { - const relativePath = formatPath(filePath); - const lines = await getFileLines(filePath); - if (!lines.length) { - return [`${relativePath}:${lineNumber}: (unable to read file)`]; - } - - const block: string[] = []; - const start = contextValue > 0 ? Math.max(1, lineNumber - contextValue) : lineNumber; - const end = contextValue > 0 ? Math.min(lines.length, lineNumber + contextValue) : lineNumber; - - for (let current = start; current <= end; current++) { - const lineText = lines[current - 1] ?? ""; - const sanitized = lineText.replace(/\r/g, ""); - const isMatchLine = current === lineNumber; - - const { text: truncatedText, wasTruncated } = truncateLine(sanitized); - if (wasTruncated) { - linesTruncated = true; - } - - if (isMatchLine) { - block.push(`${relativePath}:${current}: ${truncatedText}`); - } else { - block.push(`${relativePath}-${current}- ${truncatedText}`); - } - } - - return block; - }; - - const processLine = async (line: string): Promise => { - if (!line.trim() || matchCount >= effectiveLimit) { - return; - } - - let event: { type: string; data?: { path?: { text?: string }; line_number?: number } }; - try { - event = JSON.parse(line); - } catch { - return; - } - - if (event.type === "match") { - matchCount++; - const filePath = event.data?.path?.text; - const lineNumber = event.data?.line_number; - - if (filePath && typeof lineNumber === "number") { - recordFile(filePath); - recordFileMatch(filePath); - const block = await formatBlock(filePath, lineNumber); - outputLines.push(...block); - } - - if (matchCount >= effectiveLimit) { - matchLimitReached = true; - stopChild(true); - } - } - }; - - // Read streams using Bun's ReadableStream API + // For simple output modes (files_with_matches, count), process text directly + if (effectiveOutputMode === "files_with_matches" || effectiveOutputMode === "count") { const stdoutReader = (child.stdout as ReadableStream).getReader(); const stderrReader = (child.stderr as ReadableStream).getReader(); const decoder = new TextDecoder(); - let stdoutBuffer = ""; + let stdout = ""; await Promise.all([ - // Process stdout line by line (async () => { while (true) { const { done, value } = await stdoutReader.read(); if (done) break; - - stdoutBuffer += decoder.decode(value, { stream: true }); - const lines = stdoutBuffer.split("\n"); - // Keep the last incomplete line in the buffer - stdoutBuffer = lines.pop() ?? ""; - - for (const line of lines) { - await processLine(line); - } - } - // Process any remaining content - if (stdoutBuffer.trim()) { - await processLine(stdoutBuffer); + stdout += decoder.decode(value, { stream: true }); } })(), - // Collect stderr (async () => { while (true) { const { done, value } = await stderrReader.read(); @@ -518,12 +306,17 @@ export function createGrepTool(session: ToolSession, options?: GrepToolOptions): throw new Error("Operation aborted"); } - if (!killedDueToLimit && exitCode !== 0 && exitCode !== 1) { + if (exitCode !== 0 && exitCode !== 1) { const errorMsg = stderr.trim() || `ripgrep exited with code ${exitCode}`; throw new Error(errorMsg); } - if (matchCount === 0) { + const lines = stdout + .trim() + .split("\n") + .filter((line) => line.length > 0); + + if (lines.length === 0) { return { content: [{ type: "text", text: "No matches found" }], details: { @@ -537,8 +330,8 @@ export function createGrepTool(session: ToolSession, options?: GrepToolOptions): }; } - // Apply offset and headLimit to output lines - let processedLines = outputLines; + // Apply offset and headLimit + let processedLines = lines; if (effectiveOffset > 0) { processedLines = processedLines.slice(effectiveOffset); } @@ -546,57 +339,277 @@ export function createGrepTool(session: ToolSession, options?: GrepToolOptions): processedLines = processedLines.slice(0, headLimit); } - // Apply byte truncation (no line limit since we already have match limit) - const rawOutput = processedLines.join("\n"); - const truncation = truncateHead(rawOutput, { maxLines: Number.MAX_SAFE_INTEGER }); + let simpleMatchCount = 0; + let fileCount = 0; + const simpleFiles = new Set(); + const simpleFileList: string[] = []; + const simpleFileMatchCounts = new Map(); - let output = truncation.content; - const truncatedByHeadLimit = hasHeadLimit && processedLines.length < outputLines.length; - const details: GrepToolDetails = { - scopePath, - matchCount, - fileCount: files.size, - files: fileList, - fileMatches: fileList.map((path) => ({ - path, - count: fileMatchCounts.get(path) ?? 0, - })), - mode: effectiveOutputMode, - truncated: matchLimitReached || truncation.truncated || truncatedByHeadLimit, - headLimitReached: truncatedByHeadLimit ? headLimit : undefined, + const recordSimpleFile = (filePath: string) => { + const relative = formatPath(filePath); + if (!simpleFiles.has(relative)) { + simpleFiles.add(relative); + simpleFileList.push(relative); + } }; - // Build notices - const notices: string[] = []; + // Count mode: ripgrep provides total count per file, so we set directly (not increment) + const setFileMatchCount = (filePath: string, count: number) => { + const relative = formatPath(filePath); + simpleFileMatchCounts.set(relative, count); + }; - if (matchLimitReached) { - notices.push( - `${effectiveLimit} matches limit reached. Use limit=${effectiveLimit * 2} for more, or refine pattern`, - ); - details.matchLimitReached = effectiveLimit; + if (effectiveOutputMode === "files_with_matches") { + for (const line of lines) { + recordSimpleFile(line); + } + fileCount = simpleFiles.size; + simpleMatchCount = fileCount; + } else { + for (const line of lines) { + const separatorIndex = line.lastIndexOf(":"); + const filePart = separatorIndex === -1 ? line : line.slice(0, separatorIndex); + const countPart = separatorIndex === -1 ? "" : line.slice(separatorIndex + 1); + const count = Number.parseInt(countPart, 10); + recordSimpleFile(filePart); + if (!Number.isNaN(count)) { + simpleMatchCount += count; + setFileMatchCount(filePart, count); + } + } + fileCount = simpleFiles.size; } - if (truncation.truncated) { - notices.push(`${formatSize(DEFAULT_MAX_BYTES)} limit reached`); - details.truncation = truncation; - } - - if (linesTruncated) { - notices.push(`Some lines truncated to ${GREP_MAX_LINE_LENGTH} chars. Use read tool to see full lines`); - details.linesTruncated = true; - } - - if (notices.length > 0) { - output += `\n\n[${notices.join(". ")}]`; + const truncatedByHeadLimit = hasHeadLimit && processedLines.length < lines.length; + + // For count mode, format as "path:count" + if (effectiveOutputMode === "count") { + const formatted = processedLines.map((line) => { + const separatorIndex = line.lastIndexOf(":"); + const relative = formatPath(separatorIndex === -1 ? line : line.slice(0, separatorIndex)); + const count = separatorIndex === -1 ? "0" : line.slice(separatorIndex + 1); + return `${relative}:${count}`; + }); + const output = formatted.join("\n"); + return { + content: [{ type: "text", text: output }], + details: { + scopePath, + matchCount: simpleMatchCount, + fileCount, + files: simpleFileList, + fileMatches: simpleFileList.map((path) => ({ + path, + count: simpleFileMatchCounts.get(path) ?? 0, + })), + mode: effectiveOutputMode, + truncated: truncatedByHeadLimit, + headLimitReached: truncatedByHeadLimit ? headLimit : undefined, + }, + }; } + // For files_with_matches, format paths + const formatted = processedLines.map((line) => formatPath(line)); + const output = formatted.join("\n"); return { content: [{ type: "text", text: output }], - details: Object.keys(details).length > 0 ? details : undefined, + details: { + scopePath, + matchCount: simpleMatchCount, + fileCount, + files: simpleFileList, + mode: effectiveOutputMode, + truncated: truncatedByHeadLimit, + headLimitReached: truncatedByHeadLimit ? headLimit : undefined, + }, }; - }); - }, - }; + } + + // Content mode - existing JSON processing + const formatBlock = async (filePath: string, lineNumber: number): Promise => { + const relativePath = formatPath(filePath); + const lines = await getFileLines(filePath); + if (!lines.length) { + return [`${relativePath}:${lineNumber}: (unable to read file)`]; + } + + const block: string[] = []; + const start = contextValue > 0 ? Math.max(1, lineNumber - contextValue) : lineNumber; + const end = contextValue > 0 ? Math.min(lines.length, lineNumber + contextValue) : lineNumber; + + for (let current = start; current <= end; current++) { + const lineText = lines[current - 1] ?? ""; + const sanitized = lineText.replace(/\r/g, ""); + const isMatchLine = current === lineNumber; + + const { text: truncatedText, wasTruncated } = truncateLine(sanitized); + if (wasTruncated) { + linesTruncated = true; + } + + if (isMatchLine) { + block.push(`${relativePath}:${current}: ${truncatedText}`); + } else { + block.push(`${relativePath}-${current}- ${truncatedText}`); + } + } + + return block; + }; + + const processLine = async (line: string): Promise => { + if (!line.trim() || matchCount >= effectiveLimit) { + return; + } + + let event: { type: string; data?: { path?: { text?: string }; line_number?: number } }; + try { + event = JSON.parse(line); + } catch { + return; + } + + if (event.type === "match") { + matchCount++; + const filePath = event.data?.path?.text; + const lineNumber = event.data?.line_number; + + if (filePath && typeof lineNumber === "number") { + recordFile(filePath); + recordFileMatch(filePath); + const block = await formatBlock(filePath, lineNumber); + outputLines.push(...block); + } + + if (matchCount >= effectiveLimit) { + matchLimitReached = true; + stopChild(true); + } + } + }; + + // Read streams using Bun's ReadableStream API + const stdoutReader = (child.stdout as ReadableStream).getReader(); + const stderrReader = (child.stderr as ReadableStream).getReader(); + const decoder = new TextDecoder(); + let stdoutBuffer = ""; + + await Promise.all([ + // Process stdout line by line + (async () => { + while (true) { + const { done, value } = await stdoutReader.read(); + if (done) break; + + stdoutBuffer += decoder.decode(value, { stream: true }); + const lines = stdoutBuffer.split("\n"); + // Keep the last incomplete line in the buffer + stdoutBuffer = lines.pop() ?? ""; + + for (const line of lines) { + await processLine(line); + } + } + // Process any remaining content + if (stdoutBuffer.trim()) { + await processLine(stdoutBuffer); + } + })(), + // Collect stderr + (async () => { + while (true) { + const { done, value } = await stderrReader.read(); + if (done) break; + stderr += decoder.decode(value, { stream: true }); + } + })(), + ]); + + const exitCode = await child.exited; + + if (aborted) { + throw new Error("Operation aborted"); + } + + if (!killedDueToLimit && exitCode !== 0 && exitCode !== 1) { + const errorMsg = stderr.trim() || `ripgrep exited with code ${exitCode}`; + throw new Error(errorMsg); + } + + if (matchCount === 0) { + return { + content: [{ type: "text", text: "No matches found" }], + details: { + scopePath, + matchCount: 0, + fileCount: 0, + files: [], + mode: effectiveOutputMode, + truncated: false, + }, + }; + } + + // Apply offset and headLimit to output lines + let processedLines = outputLines; + if (effectiveOffset > 0) { + processedLines = processedLines.slice(effectiveOffset); + } + if (hasHeadLimit) { + processedLines = processedLines.slice(0, headLimit); + } + + // Apply byte truncation (no line limit since we already have match limit) + const rawOutput = processedLines.join("\n"); + const truncation = truncateHead(rawOutput, { maxLines: Number.MAX_SAFE_INTEGER }); + + let output = truncation.content; + const truncatedByHeadLimit = hasHeadLimit && processedLines.length < outputLines.length; + const details: GrepToolDetails = { + scopePath, + matchCount, + fileCount: files.size, + files: fileList, + fileMatches: fileList.map((path) => ({ + path, + count: fileMatchCounts.get(path) ?? 0, + })), + mode: effectiveOutputMode, + truncated: matchLimitReached || truncation.truncated || truncatedByHeadLimit, + headLimitReached: truncatedByHeadLimit ? headLimit : undefined, + }; + + // Build notices + const notices: string[] = []; + + if (matchLimitReached) { + notices.push( + `${effectiveLimit} matches limit reached. Use limit=${effectiveLimit * 2} for more, or refine pattern`, + ); + details.matchLimitReached = effectiveLimit; + } + + if (truncation.truncated) { + notices.push(`${formatSize(DEFAULT_MAX_BYTES)} limit reached`); + details.truncation = truncation; + } + + if (linesTruncated) { + notices.push(`Some lines truncated to ${GREP_MAX_LINE_LENGTH} chars. Use read tool to see full lines`); + details.linesTruncated = true; + } + + if (notices.length > 0) { + output += `\n\n[${notices.join(". ")}]`; + } + + return { + content: [{ type: "text", text: output }], + details, + }; + }); + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/index.ts b/packages/coding-agent/src/core/tools/index.ts index 260085658..8b0d4ff3e 100644 --- a/packages/coding-agent/src/core/tools/index.ts +++ b/packages/coding-agent/src/core/tools/index.ts @@ -1,35 +1,35 @@ -export { type AskToolDetails, askTool, createAskTool } from "./ask"; -export { type BashOperations, type BashToolDetails, type BashToolOptions, createBashTool } from "./bash"; -export { type CalculatorToolDetails, createCalculatorTool } from "./calculator"; -export { createCompleteTool } from "./complete"; +export { AskTool, type AskToolDetails } from "./ask"; +export { type BashOperations, BashTool, type BashToolDetails, type BashToolOptions } from "./bash"; +export { CalculatorTool, type CalculatorToolDetails } from "./calculator"; +export { CompleteTool } from "./complete"; // Exa MCP tools (22 tools) export { exaTools } from "./exa/index"; export type { ExaRenderDetails, ExaSearchResponse, ExaSearchResult } from "./exa/types"; -export { createFindTool, type FindOperations, type FindToolDetails, type FindToolOptions } from "./find"; +export { type FindOperations, FindTool, type FindToolDetails, type FindToolOptions } from "./find"; export { setPreferredImageProvider } from "./gemini-image"; -export { createGitTool, type GitToolDetails, gitTool } from "./git"; -export { createGrepTool, type GrepOperations, type GrepToolDetails, type GrepToolOptions } from "./grep"; -export { createLsTool, type LsOperations, type LsToolDetails, type LsToolOptions } from "./ls"; +export { GitTool, type GitToolDetails } from "./git"; +export { type GrepOperations, GrepTool, type GrepToolDetails, type GrepToolOptions } from "./grep"; +export { type LsOperations, LsTool, type LsToolDetails, type LsToolOptions } from "./ls"; export { - createLspTool, type FileDiagnosticsResult, type FileFormatResult, getLspStatus, type LspServerStatus, + LspTool, type LspToolDetails, type LspWarmupOptions, type LspWarmupResult, warmupLspServers, } from "./lsp/index"; -export { createNotebookTool, type NotebookToolDetails } from "./notebook"; -export { createOutputTool, type OutputToolDetails } from "./output"; +export { NotebookTool, type NotebookToolDetails } from "./notebook"; +export { OutputTool, type OutputToolDetails } from "./output"; export { EditTool, type EditToolDetails } from "./patch"; -export { createPythonTool, type PythonToolDetails } from "./python"; -export { createReadTool, type ReadToolDetails } from "./read"; +export { PythonTool, type PythonToolDetails, type PythonToolOptions } from "./python"; +export { ReadTool, type ReadToolDetails } from "./read"; export { reportFindingTool, type SubmitReviewDetails } from "./review"; -export { createSshTool, type SSHToolDetails } from "./ssh"; -export { BUNDLED_AGENTS, createTaskTool, taskTool } from "./task/index"; -export { createTodoWriteTool, type TodoItem, type TodoWriteToolDetails } from "./todo-write"; +export { loadSshTool, type SSHToolDetails, SshTool } from "./ssh"; +export { BUNDLED_AGENTS, TaskTool } from "./task/index"; +export { type TodoItem, TodoWriteTool, type TodoWriteToolDetails } from "./todo-write"; export { DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, @@ -40,10 +40,9 @@ export { truncateLine, truncateTail, } from "./truncate"; -export { createWebFetchTool, type WebFetchToolDetails } from "./web-fetch"; +export { WebFetchTool, type WebFetchToolDetails } from "./web-fetch"; export { companyWebSearchTools, - createWebSearchTool, exaWebSearchTools, getWebSearchTools, hasExaWebSearch, @@ -51,6 +50,7 @@ export { setPreferredWebSearchProvider, type WebSearchProvider, type WebSearchResponse, + WebSearchTool, type WebSearchToolsOptions, webSearchCodeContextTool, webSearchCompanyTool, @@ -58,9 +58,8 @@ export { webSearchCustomTool, webSearchDeepTool, webSearchLinkedinTool, - webSearchTool, } from "./web-search/index"; -export { createWriteTool, type WriteToolDetails } from "./write"; +export { WriteTool, type WriteToolDetails } from "./write"; import type { AgentTool } from "@oh-my-pi/pi-agent-core"; import type { EventBus } from "../event-bus"; @@ -68,27 +67,27 @@ import { logger } from "../logger"; import { getPreludeDocs, warmPythonEnvironment } from "../python-executor"; import { checkPythonKernelAvailability } from "../python-kernel"; import type { BashInterceptorRule } from "../settings-manager"; -import { createAskTool } from "./ask"; -import { createBashTool } from "./bash"; -import { createCalculatorTool } from "./calculator"; -import { createCompleteTool } from "./complete"; -import { createFindTool } from "./find"; -import { createGitTool } from "./git"; -import { createGrepTool } from "./grep"; -import { createLsTool } from "./ls"; -import { createLspTool } from "./lsp/index"; -import { createNotebookTool } from "./notebook"; -import { createOutputTool } from "./output"; +import { AskTool } from "./ask"; +import { BashTool } from "./bash"; +import { CalculatorTool } from "./calculator"; +import { CompleteTool } from "./complete"; +import { FindTool } from "./find"; +import { GitTool } from "./git"; +import { GrepTool } from "./grep"; +import { LsTool } from "./ls"; +import { LspTool } from "./lsp/index"; +import { NotebookTool } from "./notebook"; +import { OutputTool } from "./output"; import { EditTool } from "./patch"; -import { createPythonTool } from "./python"; -import { createReadTool } from "./read"; +import { PythonTool } from "./python"; +import { ReadTool } from "./read"; import { reportFindingTool } from "./review"; -import { createSshTool } from "./ssh"; -import { createTaskTool } from "./task/index"; -import { createTodoWriteTool } from "./todo-write"; -import { createWebFetchTool } from "./web-fetch"; -import { createWebSearchTool } from "./web-search/index"; -import { createWriteTool } from "./write"; +import { loadSshTool } from "./ssh"; +import { TaskTool } from "./task/index"; +import { TodoWriteTool } from "./todo-write"; +import { WebFetchTool } from "./web-fetch"; +import { WebSearchTool } from "./web-search/index"; +import { WriteTool } from "./write"; /** Tool type (AgentTool from pi-ai) */ export type Tool = AgentTool; @@ -130,7 +129,7 @@ export interface ToolSession { getLspDiagnosticsOnWrite(): boolean; getLspDiagnosticsOnEdit(): boolean; getEditFuzzyMatch(): boolean; - getEditFuzzyThreshold(): number; + getEditFuzzyThreshold?(): number; getEditPatchMode?(): boolean; getGitToolEnabled(): boolean; getBashInterceptorEnabled(): boolean; @@ -145,29 +144,29 @@ export interface ToolSession { type ToolFactory = (session: ToolSession) => Tool | null | Promise; export const BUILTIN_TOOLS: Record = { - ask: createAskTool, - bash: createBashTool, - python: createPythonTool, - calc: createCalculatorTool, - ssh: createSshTool, + ask: AskTool.createIf, + bash: (s) => new BashTool(s), + python: (s) => new PythonTool(s), + calc: (s) => new CalculatorTool(s), + ssh: loadSshTool, edit: (s) => new EditTool(s), - find: createFindTool, - git: createGitTool, - grep: createGrepTool, - ls: createLsTool, - lsp: createLspTool, - notebook: createNotebookTool, - output: createOutputTool, - read: createReadTool, - task: createTaskTool, - todo_write: createTodoWriteTool, - web_fetch: createWebFetchTool, - web_search: createWebSearchTool, - write: createWriteTool, + find: (s) => new FindTool(s), + git: GitTool.createIf, + grep: (s) => new GrepTool(s), + ls: (s) => new LsTool(s), + lsp: LspTool.createIf, + notebook: (s) => new NotebookTool(s), + output: (s) => new OutputTool(s), + read: (s) => new ReadTool(s), + task: TaskTool.create, + todo_write: (s) => new TodoWriteTool(s), + web_fetch: (s) => new WebFetchTool(s), + web_search: (s) => new WebSearchTool(s), + write: (s) => new WriteTool(s), }; export const HIDDEN_TOOLS: Record = { - complete: createCompleteTool, + complete: (s) => new CompleteTool(s), report_finding: () => reportFindingTool, }; diff --git a/packages/coding-agent/src/core/tools/ls.ts b/packages/coding-agent/src/core/tools/ls.ts index 1974f6614..a03ac4051 100644 --- a/packages/coding-agent/src/core/tools/ls.ts +++ b/packages/coding-agent/src/core/tools/ls.ts @@ -1,7 +1,6 @@ import nodePath from "node:path"; -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; -import type { Component } from "@oh-my-pi/pi-tui"; -import { Text } from "@oh-my-pi/pi-tui"; +import type { AgentTool, AgentToolResult } from "@oh-my-pi/pi-agent-core"; +import { type Component, Text } from "@oh-my-pi/pi-tui"; import { Type } from "@sinclair/typebox"; import { getLanguageFromPath, type Theme } from "../../modes/interactive/theme/theme"; import type { RenderResultOptions } from "../custom-tools/types"; @@ -69,129 +68,135 @@ const defaultLsOperations: LsOperations = { }, }; -export function createLsTool(session: ToolSession, options?: LsToolOptions): AgentTool { - const ops = options?.operations ?? defaultLsOperations; +export class LsTool implements AgentTool { + public readonly name = "ls"; + public readonly label = "Ls"; + public readonly description = + 'List directory contents with modification times. Returns entries sorted alphabetically, with \'/\' suffix for directories and relative age (e.g., "2d ago", "just now"). Includes dotfiles. Output is truncated to 500 entries or 50KB (whichever is hit first).'; + public readonly parameters = lsSchema; - return { - name: "ls", - label: "Ls", - description: `List directory contents with modification times. Returns entries sorted alphabetically, with '/' suffix for directories and relative age (e.g., "2d ago", "just now"). Includes dotfiles. Output is truncated to 500 entries or 50KB (whichever is hit first).`, - parameters: lsSchema, - execute: async ( - _toolCallId: string, - { path, limit }: { path?: string; limit?: number }, - signal?: AbortSignal, - ) => { - return untilAborted(signal, async () => { - const dirPath = resolveToCwd(path || ".", session.cwd); - const effectiveLimit = limit ?? DEFAULT_LIMIT; + private readonly session: ToolSession; + private readonly ops: LsOperations; - // Check if path exists and is a directory - const dirStat = await ops.stat(dirPath); - if (!dirStat) { - throw new Error(`Path not found: ${dirPath}`); + constructor(session: ToolSession, options?: LsToolOptions) { + this.session = session; + this.ops = options?.operations ?? defaultLsOperations; + } + + public async execute( + _toolCallId: string, + { path, limit }: { path?: string; limit?: number }, + signal?: AbortSignal, + ): Promise> { + return untilAborted(signal, async () => { + const dirPath = resolveToCwd(path || ".", this.session.cwd); + const effectiveLimit = limit ?? DEFAULT_LIMIT; + + // Check if path exists and is a directory + const dirStat = await this.ops.stat(dirPath); + if (!dirStat) { + throw new Error(`Path not found: ${dirPath}`); + } + + if (!dirStat.isDirectory()) { + throw new Error(`Not a directory: ${dirPath}`); + } + + // Read directory entries + let entries: string[]; + try { + entries = await this.ops.readdir(dirPath); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + throw new Error(`Cannot read directory: ${message}`); + } + + // Sort alphabetically (case-insensitive) + entries.sort((a, b) => a.toLowerCase().localeCompare(b.toLowerCase())); + + // Format entries with directory indicators + const results: string[] = []; + let entryLimitReached = false; + let dirCount = 0; + let fileCount = 0; + + for (const entry of entries) { + signal?.throwIfAborted(); + if (results.length >= effectiveLimit) { + entryLimitReached = true; + break; } - if (!dirStat.isDirectory()) { - throw new Error(`Not a directory: ${dirPath}`); + const fullPath = nodePath.join(dirPath, entry); + let suffix = ""; + let age = ""; + + const entryStat = await this.ops.stat(fullPath); + if (!entryStat) { + // Skip entries we can't stat + continue; } - // Read directory entries - let entries: string[]; - try { - entries = await ops.readdir(dirPath); - } catch (error) { - const message = error instanceof Error ? error.message : String(error); - throw new Error(`Cannot read directory: ${message}`); + if (entryStat.isDirectory()) { + suffix = "/"; + dirCount += 1; + } else { + fileCount += 1; } + // Calculate age from mtime + const ageSeconds = Math.floor((Date.now() - entryStat.mtimeMs) / 1000); + age = formatAge(ageSeconds); - // Sort alphabetically (case-insensitive) - entries.sort((a, b) => a.toLowerCase().localeCompare(b.toLowerCase())); + // Format: "name/ (2d ago)" or "name (just now)" + const line = age ? `${entry}${suffix} (${age})` : entry + suffix; + results.push(line); + } - // Format entries with directory indicators - const results: string[] = []; - let entryLimitReached = false; - let dirCount = 0; - let fileCount = 0; + if (results.length === 0) { + return { content: [{ type: "text", text: "(empty directory)" }], details: {} }; + } - for (const entry of entries) { - signal?.throwIfAborted(); - if (results.length >= effectiveLimit) { - entryLimitReached = true; - break; - } + // Apply byte truncation (no line limit since we already have entry limit) + const rawOutput = results.join("\n"); + const truncation = truncateHead(rawOutput, { maxLines: Number.MAX_SAFE_INTEGER }); - const fullPath = nodePath.join(dirPath, entry); - let suffix = ""; - let age = ""; + let output = truncation.content; + const details: LsToolDetails = { + entries: results, + dirCount, + fileCount, + }; + const truncationReasons: Array<"entryLimit" | "byteLimit"> = []; - const entryStat = await ops.stat(fullPath); - if (!entryStat) { - // Skip entries we can't stat - continue; - } + // Build notices + const notices: string[] = []; - if (entryStat.isDirectory()) { - suffix = "/"; - dirCount += 1; - } else { - fileCount += 1; - } - // Calculate age from mtime - const ageSeconds = Math.floor((Date.now() - entryStat.mtimeMs) / 1000); - age = formatAge(ageSeconds); + if (entryLimitReached) { + notices.push(`${effectiveLimit} entries limit reached. Use limit=${effectiveLimit * 2} for more`); + details.entryLimitReached = effectiveLimit; + truncationReasons.push("entryLimit"); + } - // Format: "name/ (2d ago)" or "name (just now)" - const line = age ? `${entry}${suffix} (${age})` : entry + suffix; - results.push(line); - } + if (truncation.truncated) { + notices.push(`${formatSize(DEFAULT_MAX_BYTES)} limit reached`); + details.truncation = truncation; + truncationReasons.push("byteLimit"); + } - if (results.length === 0) { - return { content: [{ type: "text", text: "(empty directory)" }], details: undefined }; - } + if (truncationReasons.length > 0) { + details.truncationReasons = truncationReasons; + } - // Apply byte truncation (no line limit since we already have entry limit) - const rawOutput = results.join("\n"); - const truncation = truncateHead(rawOutput, { maxLines: Number.MAX_SAFE_INTEGER }); + if (notices.length > 0) { + output += `\n\n[${notices.join(". ")}]`; + } - let output = truncation.content; - const details: LsToolDetails = { - entries: results, - dirCount, - fileCount, - }; - const truncationReasons: Array<"entryLimit" | "byteLimit"> = []; - - // Build notices - const notices: string[] = []; - - if (entryLimitReached) { - notices.push(`${effectiveLimit} entries limit reached. Use limit=${effectiveLimit * 2} for more`); - details.entryLimitReached = effectiveLimit; - truncationReasons.push("entryLimit"); - } - - if (truncation.truncated) { - notices.push(`${formatSize(DEFAULT_MAX_BYTES)} limit reached`); - details.truncation = truncation; - truncationReasons.push("byteLimit"); - } - - if (truncationReasons.length > 0) { - details.truncationReasons = truncationReasons; - } - - if (notices.length > 0) { - output += `\n\n[${notices.join(". ")}]`; - } - - return { - content: [{ type: "text", text: output }], - details, - }; - }); - }, - }; + return { + content: [{ type: "text", text: output }], + details, + }; + }); + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/lsp/clients/biome-client.ts b/packages/coding-agent/src/core/tools/lsp/clients/biome-client.ts index 1ad65ae5b..e0b662b15 100644 --- a/packages/coding-agent/src/core/tools/lsp/clients/biome-client.ts +++ b/packages/coding-agent/src/core/tools/lsp/clients/biome-client.ts @@ -110,6 +110,11 @@ export class BiomeClient implements LinterClient { private config: ServerConfig; private cwd: string; + /** Factory method for creating BiomeClient instances */ + static create(config: ServerConfig, cwd: string): LinterClient { + return new BiomeClient(config, cwd); + } + constructor(config: ServerConfig, cwd: string) { this.config = config; this.cwd = cwd; @@ -198,10 +203,3 @@ export class BiomeClient implements LinterClient { // Nothing to dispose for CLI client } } - -/** - * Factory function to create a Biome client. - */ -export function createBiomeClient(config: ServerConfig, cwd: string): LinterClient { - return new BiomeClient(config, cwd); -} diff --git a/packages/coding-agent/src/core/tools/lsp/clients/index.ts b/packages/coding-agent/src/core/tools/lsp/clients/index.ts index 189c860ef..bdbc416ad 100644 --- a/packages/coding-agent/src/core/tools/lsp/clients/index.ts +++ b/packages/coding-agent/src/core/tools/lsp/clients/index.ts @@ -5,11 +5,11 @@ * Different implementations can use LSP protocol, CLI tools, or other mechanisms. */ -export { BiomeClient, createBiomeClient } from "./biome-client"; -export { createLspLinterClient, LspLinterClient } from "./lsp-linter-client"; +export { BiomeClient } from "./biome-client"; +export { LspLinterClient } from "./lsp-linter-client"; import type { LinterClient, ServerConfig } from "../types"; -import { createLspLinterClient } from "./lsp-linter-client"; +import { LspLinterClient } from "./lsp-linter-client"; // Cache of linter clients by server name + cwd const clientCache = new Map(); @@ -31,7 +31,7 @@ export function getLinterClient(serverName: string, config: ServerConfig, cwd: s client = config.createClient(config, cwd); } else { // Default to LSP - client = createLspLinterClient(config, cwd); + client = LspLinterClient.create(config, cwd); } clientCache.set(key, client); diff --git a/packages/coding-agent/src/core/tools/lsp/clients/lsp-linter-client.ts b/packages/coding-agent/src/core/tools/lsp/clients/lsp-linter-client.ts index fa54fd4f8..a55fba1f1 100644 --- a/packages/coding-agent/src/core/tools/lsp/clients/lsp-linter-client.ts +++ b/packages/coding-agent/src/core/tools/lsp/clients/lsp-linter-client.ts @@ -26,6 +26,11 @@ export class LspLinterClient implements LinterClient { private cwd: string; private client: LspClient | null = null; + /** Factory method for creating LspLinterClient instances */ + static create(config: ServerConfig, cwd: string): LinterClient { + return new LspLinterClient(config, cwd); + } + constructor(config: ServerConfig, cwd: string) { this.config = config; this.cwd = cwd; @@ -89,10 +94,3 @@ export class LspLinterClient implements LinterClient { // Client lifecycle is managed globally, nothing to dispose here } } - -/** - * Factory function to create an LSP linter client. - */ -export function createLspLinterClient(config: ServerConfig, cwd: string): LinterClient { - return new LspLinterClient(config, cwd); -} diff --git a/packages/coding-agent/src/core/tools/lsp/config.ts b/packages/coding-agent/src/core/tools/lsp/config.ts index 7d5331abe..3c12650bf 100644 --- a/packages/coding-agent/src/core/tools/lsp/config.ts +++ b/packages/coding-agent/src/core/tools/lsp/config.ts @@ -4,7 +4,7 @@ import { YAML } from "bun"; import { globSync } from "glob"; import { getConfigDirPaths } from "../../../config"; import { logger } from "../../logger"; -import { createBiomeClient } from "./clients/biome-client"; +import { BiomeClient } from "./clients/biome-client"; import DEFAULTS from "./defaults.json" with { type: "json" }; import type { ServerConfig } from "./types"; @@ -137,7 +137,7 @@ function applyRuntimeDefaults(servers: Record): Record = { ...servers }; if (updated.biome) { - updated.biome = { ...updated.biome, createClient: createBiomeClient }; + updated.biome = { ...updated.biome, createClient: BiomeClient.create }; } if (updated.omnisharp?.args) { diff --git a/packages/coding-agent/src/core/tools/lsp/index.ts b/packages/coding-agent/src/core/tools/lsp/index.ts index 3645fced4..7612f6b5a 100644 --- a/packages/coding-agent/src/core/tools/lsp/index.ts +++ b/packages/coding-agent/src/core/tools/lsp/index.ts @@ -1,7 +1,7 @@ import type { Dirent } from "node:fs"; import { existsSync, statSync } from "node:fs"; import path from "node:path"; -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { BunFile } from "bun"; import { type Theme, theme } from "../../../modes/interactive/theme/theme"; import lspDescription from "../../../prompts/tools/lsp.md" with { type: "text" }; @@ -923,692 +923,705 @@ export function createLspWritethrough(cwd: string, options?: WritethroughOptions }; } -/** Create an LSP tool */ -export function createLspTool(session: ToolSession): AgentTool | null { - if (session.enableLsp === false) { - return null; +/** + * LSP tool for language server protocol operations. + */ +export class LspTool implements AgentTool { + public readonly name = "lsp"; + public readonly label = "LSP"; + public readonly description: string; + public readonly parameters = lspSchema; + public readonly renderCall = renderCall; + public readonly renderResult = renderResult; + + private readonly session: ToolSession; + + constructor(session: ToolSession) { + this.session = session; + this.description = renderPromptTemplate(lspDescription); } - return { - name: "lsp", - label: "LSP", - description: renderPromptTemplate(lspDescription), - parameters: lspSchema, - renderCall, - renderResult, - execute: async (_toolCallId, params: LspParams, _signal) => { - const { - action, - file, - files, - line, - column, - end_line, - end_character, - query, - new_name, - replacement, - kind, - apply, - action_index, - include_declaration, - } = params; - const config = await getConfig(session.cwd); + static createIf(session: ToolSession): LspTool | null { + return session.enableLsp === false ? null : new LspTool(session); + } - // Status action doesn't need a file - if (action === "status") { - const servers = Object.keys(config.servers); - const lspmuxState = await detectLspmux(); - const lspmuxStatus = lspmuxState.available - ? lspmuxState.running - ? "lspmux: active (multiplexing enabled)" - : "lspmux: installed but server not running" - : ""; + public async execute( + _toolCallId: string, + params: LspParams, + _signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const { + action, + file, + files, + line, + column, + end_line, + end_character, + query, + new_name, + replacement, + kind, + apply, + action_index, + include_declaration, + } = params; - const serverStatus = - servers.length > 0 - ? `Active language servers: ${servers.join(", ")}` - : "No language servers configured for this project"; + const config = await getConfig(this.session.cwd); - const output = lspmuxStatus ? `${serverStatus}\n${lspmuxStatus}` : serverStatus; + // Status action doesn't need a file + if (action === "status") { + const servers = Object.keys(config.servers); + const lspmuxState = await detectLspmux(); + const lspmuxStatus = lspmuxState.available + ? lspmuxState.running + ? "lspmux: active (multiplexing enabled)" + : "lspmux: installed but server not running" + : ""; + + const serverStatus = + servers.length > 0 + ? `Active language servers: ${servers.join(", ")}` + : "No language servers configured for this project"; + + const output = lspmuxStatus ? `${serverStatus}\n${lspmuxStatus}` : serverStatus; + return { + content: [{ type: "text", text: output }], + details: { action, success: true }, + }; + } + + // Workspace diagnostics - check entire project + if (action === "workspace_diagnostics") { + const result = await runWorkspaceDiagnostics(this.session.cwd, config); + return { + content: [ + { + type: "text", + text: `Workspace diagnostics (${result.projectType.description}):\n${result.output}`, + }, + ], + details: { action, success: true }, + }; + } + + // Diagnostics can be batch or single-file - queries all applicable servers + if (action === "diagnostics") { + const targets = files?.length ? files : file ? [file] : null; + if (!targets) { return { - content: [{ type: "text", text: output }], - details: { action, success: true }, + content: [{ type: "text", text: "Error: file or files parameter required for diagnostics" }], + details: { action, success: false }, }; } - // Workspace diagnostics - check entire project - if (action === "workspace_diagnostics") { - const result = await runWorkspaceDiagnostics(session.cwd, config); - return { - content: [ - { - type: "text", - text: `Workspace diagnostics (${result.projectType.description}):\n${result.output}`, - }, - ], - details: { action, success: true }, - }; - } + const detailed = Boolean(files?.length); + const results: string[] = []; + const allServerNames = new Set(); - // Diagnostics can be batch or single-file - queries all applicable servers - if (action === "diagnostics") { - const targets = files?.length ? files : file ? [file] : null; - if (!targets) { - return { - content: [{ type: "text", text: "Error: file or files parameter required for diagnostics" }], - details: { action, success: false }, - }; + for (const target of targets) { + const resolved = resolveToCwd(target, this.session.cwd); + const servers = getServersForFile(config, resolved); + if (servers.length === 0) { + results.push(`${theme.status.error} ${target}: No language server found`); + continue; } - const detailed = Boolean(files?.length); - const results: string[] = []; - const allServerNames = new Set(); + const uri = fileToUri(resolved); + const relPath = path.relative(this.session.cwd, resolved); + const allDiagnostics: Diagnostic[] = []; - for (const target of targets) { - const resolved = resolveToCwd(target, session.cwd); - const servers = getServersForFile(config, resolved); - if (servers.length === 0) { - results.push(`${theme.status.error} ${target}: No language server found`); - continue; - } - - const uri = fileToUri(resolved); - const relPath = path.relative(session.cwd, resolved); - const allDiagnostics: Diagnostic[] = []; - - // Query all applicable servers for this file - for (const [serverName, serverConfig] of servers) { - allServerNames.add(serverName); - try { - if (serverConfig.createClient) { - const linterClient = getLinterClient(serverName, serverConfig, session.cwd); - const diagnostics = await linterClient.lint(resolved); - allDiagnostics.push(...diagnostics); - continue; - } - const client = await getOrCreateClient(serverConfig, session.cwd); - const minVersion = client.diagnosticsVersion; - await refreshFile(client, resolved); - const diagnostics = await waitForDiagnostics(client, uri, 3000, undefined, minVersion); + // Query all applicable servers for this file + for (const [serverName, serverConfig] of servers) { + allServerNames.add(serverName); + try { + if (serverConfig.createClient) { + const linterClient = getLinterClient(serverName, serverConfig, this.session.cwd); + const diagnostics = await linterClient.lint(resolved); allDiagnostics.push(...diagnostics); - } catch { - // Server failed, continue with others + continue; } + const client = await getOrCreateClient(serverConfig, this.session.cwd); + const minVersion = client.diagnosticsVersion; + await refreshFile(client, resolved); + const diagnostics = await waitForDiagnostics(client, uri, 3000, undefined, minVersion); + allDiagnostics.push(...diagnostics); + } catch { + // Server failed, continue with others } + } - // Deduplicate diagnostics - const seen = new Set(); - const uniqueDiagnostics: Diagnostic[] = []; - for (const d of allDiagnostics) { - const key = `${d.range.start.line}:${d.range.start.character}:${d.range.end.line}:${d.range.end.character}:${d.message}`; - if (!seen.has(key)) { - seen.add(key); - uniqueDiagnostics.push(d); - } + // Deduplicate diagnostics + const seen = new Set(); + const uniqueDiagnostics: Diagnostic[] = []; + for (const d of allDiagnostics) { + const key = `${d.range.start.line}:${d.range.start.character}:${d.range.end.line}:${d.range.end.character}:${d.message}`; + if (!seen.has(key)) { + seen.add(key); + uniqueDiagnostics.push(d); } + } - if (!detailed && targets.length === 1) { - if (uniqueDiagnostics.length === 0) { - return { - content: [{ type: "text", text: "No diagnostics" }], - details: { action, serverName: Array.from(allServerNames).join(", "), success: true }, - }; - } - - const summary = formatDiagnosticsSummary(uniqueDiagnostics); - const formatted = uniqueDiagnostics.map((d) => formatDiagnostic(d, relPath)); - const output = `${summary}:\n${formatted.map((f) => ` ${f}`).join("\n")}`; + if (!detailed && targets.length === 1) { + if (uniqueDiagnostics.length === 0) { return { - content: [{ type: "text", text: output }], + content: [{ type: "text", text: "No diagnostics" }], details: { action, serverName: Array.from(allServerNames).join(", "), success: true }, }; } - if (uniqueDiagnostics.length === 0) { - results.push(`${theme.status.success} ${relPath}: no issues`); - } else { - const summary = formatDiagnosticsSummary(uniqueDiagnostics); - results.push(`${theme.status.error} ${relPath}: ${summary}`); - for (const diag of uniqueDiagnostics) { - results.push(` ${formatDiagnostic(diag, relPath)}`); - } - } + const summary = formatDiagnosticsSummary(uniqueDiagnostics); + const formatted = uniqueDiagnostics.map((d) => formatDiagnostic(d, relPath)); + const output = `${summary}:\n${formatted.map((f) => ` ${f}`).join("\n")}`; + return { + content: [{ type: "text", text: output }], + details: { action, serverName: Array.from(allServerNames).join(", "), success: true }, + }; } - return { - content: [{ type: "text", text: results.join("\n") }], - details: { action, serverName: Array.from(allServerNames).join(", "), success: true }, - }; + if (uniqueDiagnostics.length === 0) { + results.push(`${theme.status.success} ${relPath}: no issues`); + } else { + const summary = formatDiagnosticsSummary(uniqueDiagnostics); + results.push(`${theme.status.error} ${relPath}: ${summary}`); + for (const diag of uniqueDiagnostics) { + results.push(` ${formatDiagnostic(diag, relPath)}`); + } + } } - const requiresFile = - !file && - action !== "workspace_symbols" && - action !== "flycheck" && - action !== "ssr" && - action !== "runnables" && - action !== "reload_workspace"; + return { + content: [{ type: "text", text: results.join("\n") }], + details: { action, serverName: Array.from(allServerNames).join(", "), success: true }, + }; + } - if (requiresFile) { - return { - content: [{ type: "text", text: "Error: file parameter required for this action" }], - details: { action, success: false }, - }; + const requiresFile = + !file && + action !== "workspace_symbols" && + action !== "flycheck" && + action !== "ssr" && + action !== "runnables" && + action !== "reload_workspace"; + + if (requiresFile) { + return { + content: [{ type: "text", text: "Error: file parameter required for this action" }], + details: { action, success: false }, + }; + } + + const resolvedFile = file ? resolveToCwd(file, this.session.cwd) : null; + const serverInfo = resolvedFile + ? getLspServerForFile(config, resolvedFile) + : getServerForWorkspaceAction(config, action); + + if (!serverInfo) { + return { + content: [{ type: "text", text: "No language server found for this action" }], + details: { action, success: false }, + }; + } + + const [serverName, serverConfig] = serverInfo; + + try { + const client = await getOrCreateClient(serverConfig, this.session.cwd); + let targetFile = resolvedFile; + if (action === "runnables" && !targetFile) { + targetFile = findFileForServer(this.session.cwd, serverConfig); + if (!targetFile) { + return { + content: [{ type: "text", text: "Error: no matching files found for runnables" }], + details: { action, serverName, success: false }, + }; + } } - const resolvedFile = file ? resolveToCwd(file, session.cwd) : null; - const serverInfo = resolvedFile - ? getLspServerForFile(config, resolvedFile) - : getServerForWorkspaceAction(config, action); - - if (!serverInfo) { - return { - content: [{ type: "text", text: "No language server found for this action" }], - details: { action, success: false }, - }; + if (targetFile) { + await ensureFileOpen(client, targetFile); } - const [serverName, serverConfig] = serverInfo; + const uri = targetFile ? fileToUri(targetFile) : ""; + const position = { line: (line || 1) - 1, character: (column || 1) - 1 }; - try { - const client = await getOrCreateClient(serverConfig, session.cwd); - let targetFile = resolvedFile; - if (action === "runnables" && !targetFile) { - targetFile = findFileForServer(session.cwd, serverConfig); - if (!targetFile) { + let output: string; + + switch (action) { + // ===================================================================== + // Standard LSP Operations + // ===================================================================== + + case "definition": { + const result = (await sendRequest(client, "textDocument/definition", { + textDocument: { uri }, + position, + })) as Location | Location[] | LocationLink | LocationLink[] | null; + + if (!result) { + output = "No definition found"; + } else { + const raw = Array.isArray(result) ? result : [result]; + const locations = raw.flatMap((loc) => { + if ("uri" in loc) { + return [loc as Location]; + } + if ("targetUri" in loc) { + // Use targetSelectionRange (the precise identifier range) with fallback to targetRange + const link = loc as LocationLink; + return [{ uri: link.targetUri, range: link.targetSelectionRange ?? link.targetRange }]; + } + return []; + }); + + if (locations.length === 0) { + output = "No definition found"; + } else { + output = `Found ${locations.length} definition(s):\n${locations + .map((loc) => ` ${formatLocation(loc, this.session.cwd)}`) + .join("\n")}`; + } + } + break; + } + + case "references": { + const result = (await sendRequest(client, "textDocument/references", { + textDocument: { uri }, + position, + context: { includeDeclaration: include_declaration ?? true }, + })) as Location[] | null; + + if (!result || result.length === 0) { + output = "No references found"; + } else { + const lines = result.map((loc) => ` ${formatLocation(loc, this.session.cwd)}`); + output = `Found ${result.length} reference(s):\n${lines.join("\n")}`; + } + break; + } + + case "hover": { + const result = (await sendRequest(client, "textDocument/hover", { + textDocument: { uri }, + position, + })) as Hover | null; + + if (!result || !result.contents) { + output = "No hover information"; + } else { + output = extractHoverText(result.contents); + } + break; + } + + case "symbols": { + const result = (await sendRequest(client, "textDocument/documentSymbol", { + textDocument: { uri }, + })) as (DocumentSymbol | SymbolInformation)[] | null; + + if (!result || result.length === 0) { + output = "No symbols found"; + } else if (!targetFile) { return { - content: [{ type: "text", text: "Error: no matching files found for runnables" }], + content: [{ type: "text", text: "Error: file parameter required for symbols" }], + details: { action, serverName, success: false }, + }; + } else { + const relPath = path.relative(this.session.cwd, targetFile); + // Check if hierarchical (DocumentSymbol) or flat (SymbolInformation) + if ("selectionRange" in result[0]) { + // Hierarchical + const lines = (result as DocumentSymbol[]).flatMap((s) => formatDocumentSymbol(s)); + output = `Symbols in ${relPath}:\n${lines.join("\n")}`; + } else { + // Flat + const lines = (result as SymbolInformation[]).map((s) => { + const line = s.location.range.start.line + 1; + const icon = symbolKindToIcon(s.kind); + return `${icon} ${s.name} @ line ${line}`; + }); + output = `Symbols in ${relPath}:\n${lines.join("\n")}`; + } + } + break; + } + + case "workspace_symbols": { + if (!query) { + return { + content: [{ type: "text", text: "Error: query parameter required for workspace_symbols" }], details: { action, serverName, success: false }, }; } + + const result = (await sendRequest(client, "workspace/symbol", { query })) as SymbolInformation[] | null; + + if (!result || result.length === 0) { + output = `No symbols matching "${query}"`; + } else { + const lines = result.map((s) => formatSymbolInformation(s, this.session.cwd)); + output = `Found ${result.length} symbol(s) matching "${query}":\n${lines.map((l) => ` ${l}`).join("\n")}`; + } + break; } - if (targetFile) { - await ensureFileOpen(client, targetFile); + case "rename": { + if (!new_name) { + return { + content: [{ type: "text", text: "Error: new_name parameter required for rename" }], + details: { action, serverName, success: false }, + }; + } + + const result = (await sendRequest(client, "textDocument/rename", { + textDocument: { uri }, + position, + newName: new_name, + })) as WorkspaceEdit | null; + + if (!result) { + output = "Rename returned no edits"; + } else { + const shouldApply = apply !== false; + if (shouldApply) { + const applied = await applyWorkspaceEdit(result, this.session.cwd); + output = `Applied rename:\n${applied.map((a) => ` ${a}`).join("\n")}`; + } else { + const preview = formatWorkspaceEdit(result, this.session.cwd); + output = `Rename preview:\n${preview.map((p) => ` ${p}`).join("\n")}`; + } + } + break; } - const uri = targetFile ? fileToUri(targetFile) : ""; - const position = { line: (line || 1) - 1, character: (column || 1) - 1 }; - - let output: string; - - switch (action) { - // ===================================================================== - // Standard LSP Operations - // ===================================================================== - - case "definition": { - const result = (await sendRequest(client, "textDocument/definition", { - textDocument: { uri }, - position, - })) as Location | Location[] | LocationLink | LocationLink[] | null; - - if (!result) { - output = "No definition found"; - } else { - const raw = Array.isArray(result) ? result : [result]; - const locations = raw.flatMap((loc) => { - if ("uri" in loc) { - return [loc as Location]; - } - if ("targetUri" in loc) { - // Use targetSelectionRange (the precise identifier range) with fallback to targetRange - const link = loc as LocationLink; - return [{ uri: link.targetUri, range: link.targetSelectionRange ?? link.targetRange }]; - } - return []; - }); - - if (locations.length === 0) { - output = "No definition found"; - } else { - output = `Found ${locations.length} definition(s):\n${locations - .map((loc) => ` ${formatLocation(loc, session.cwd)}`) - .join("\n")}`; - } - } - break; + case "actions": { + if (!targetFile) { + return { + content: [{ type: "text", text: "Error: file parameter required for actions" }], + details: { action, serverName, success: false }, + }; } - case "references": { - const result = (await sendRequest(client, "textDocument/references", { - textDocument: { uri }, - position, - context: { includeDeclaration: include_declaration ?? true }, - })) as Location[] | null; + const actionsMinVersion = client.diagnosticsVersion; + await refreshFile(client, targetFile); + const diagnostics = await waitForDiagnostics(client, uri, 3000, undefined, actionsMinVersion); + const endLine = (end_line ?? line ?? 1) - 1; + const endCharacter = (end_character ?? column ?? 1) - 1; + const range = { start: position, end: { line: endLine, character: endCharacter } }; + const relevantDiagnostics = diagnostics.filter( + (d) => d.range.start.line <= range.end.line && d.range.end.line >= range.start.line, + ); - if (!result || result.length === 0) { - output = "No references found"; - } else { - const lines = result.map((loc) => ` ${formatLocation(loc, session.cwd)}`); - output = `Found ${result.length} reference(s):\n${lines.join("\n")}`; - } - break; + const codeActionContext: { diagnostics: Diagnostic[]; only?: string[] } = { + diagnostics: relevantDiagnostics, + }; + if (kind) { + codeActionContext.only = [kind]; } - case "hover": { - const result = (await sendRequest(client, "textDocument/hover", { - textDocument: { uri }, - position, - })) as Hover | null; + const result = (await sendRequest(client, "textDocument/codeAction", { + textDocument: { uri }, + range, + context: codeActionContext, + })) as Array | null; - if (!result || !result.contents) { - output = "No hover information"; - } else { - output = extractHoverText(result.contents); - } - break; - } - - case "symbols": { - const result = (await sendRequest(client, "textDocument/documentSymbol", { - textDocument: { uri }, - })) as (DocumentSymbol | SymbolInformation)[] | null; - - if (!result || result.length === 0) { - output = "No symbols found"; - } else if (!targetFile) { + if (!result || result.length === 0) { + output = "No code actions available"; + } else if (action_index !== undefined) { + // Apply specific action + if (action_index < 0 || action_index >= result.length) { return { - content: [{ type: "text", text: "Error: file parameter required for symbols" }], - details: { action, serverName, success: false }, - }; - } else { - const relPath = path.relative(session.cwd, targetFile); - // Check if hierarchical (DocumentSymbol) or flat (SymbolInformation) - if ("selectionRange" in result[0]) { - // Hierarchical - const lines = (result as DocumentSymbol[]).flatMap((s) => formatDocumentSymbol(s)); - output = `Symbols in ${relPath}:\n${lines.join("\n")}`; - } else { - // Flat - const lines = (result as SymbolInformation[]).map((s) => { - const line = s.location.range.start.line + 1; - const icon = symbolKindToIcon(s.kind); - return `${icon} ${s.name} @ line ${line}`; - }); - output = `Symbols in ${relPath}:\n${lines.join("\n")}`; - } - } - break; - } - - case "workspace_symbols": { - if (!query) { - return { - content: [{ type: "text", text: "Error: query parameter required for workspace_symbols" }], + content: [ + { + type: "text", + text: `Error: action_index ${action_index} out of range (0-${result.length - 1})`, + }, + ], details: { action, serverName, success: false }, }; } - const result = (await sendRequest(client, "workspace/symbol", { query })) as - | SymbolInformation[] + const isCommand = (candidate: CodeAction | Command): candidate is Command => + typeof (candidate as Command).command === "string"; + const isCodeAction = (candidate: CodeAction | Command): candidate is CodeAction => + !isCommand(candidate); + const getCommandPayload = ( + candidate: CodeAction | Command, + ): { command: string; arguments?: unknown[] } | null => { + if (isCommand(candidate)) { + return { command: candidate.command, arguments: candidate.arguments }; + } + if (candidate.command) { + return { command: candidate.command.command, arguments: candidate.command.arguments }; + } + return null; + }; + + const codeAction = result[action_index]; + + // Resolve if needed + let resolvedAction = codeAction; + if ( + isCodeAction(codeAction) && + !codeAction.edit && + codeAction.data && + client.serverCapabilities?.codeActionProvider + ) { + const provider = client.serverCapabilities.codeActionProvider; + if (typeof provider === "object" && provider.resolveProvider) { + resolvedAction = (await sendRequest(client, "codeAction/resolve", codeAction)) as CodeAction; + } + } + + if (isCodeAction(resolvedAction) && resolvedAction.edit) { + const applied = await applyWorkspaceEdit(resolvedAction.edit, this.session.cwd); + output = `Applied "${codeAction.title}":\n${applied.map((a) => ` ${a}`).join("\n")}`; + } else { + const commandPayload = getCommandPayload(resolvedAction); + if (commandPayload) { + await sendRequest(client, "workspace/executeCommand", commandPayload); + output = `Executed "${codeAction.title}"`; + } else { + output = `Code action "${codeAction.title}" has no edits or command to apply`; + } + } + } else { + // List available actions + const lines = result.map((actionItem, i) => { + if ("kind" in actionItem || "isPreferred" in actionItem || "edit" in actionItem) { + const actionDetails = actionItem as CodeAction; + const preferred = actionDetails.isPreferred ? " (preferred)" : ""; + const kindInfo = actionDetails.kind ? ` [${actionDetails.kind}]` : ""; + return ` [${i}] ${actionDetails.title}${kindInfo}${preferred}`; + } + return ` [${i}] ${actionItem.title}`; + }); + output = `Available code actions:\n${lines.join("\n")}\n\nUse action_index parameter to apply a specific action.`; + } + break; + } + + case "incoming_calls": + case "outgoing_calls": { + // First, prepare the call hierarchy item at the cursor position + const prepareResult = (await sendRequest(client, "textDocument/prepareCallHierarchy", { + textDocument: { uri }, + position, + })) as CallHierarchyItem[] | null; + + if (!prepareResult || prepareResult.length === 0) { + output = "No callable symbol found at this position"; + break; + } + + const item = prepareResult[0]; + + if (action === "incoming_calls") { + const calls = (await sendRequest(client, "callHierarchy/incomingCalls", { item })) as + | CallHierarchyIncomingCall[] | null; - if (!result || result.length === 0) { - output = `No symbols matching "${query}"`; + if (!calls || calls.length === 0) { + output = `No callers found for "${item.name}"`; } else { - const lines = result.map((s) => formatSymbolInformation(s, session.cwd)); - output = `Found ${result.length} symbol(s) matching "${query}":\n${lines.map((l) => ` ${l}`).join("\n")}`; - } - break; - } - - case "rename": { - if (!new_name) { - return { - content: [{ type: "text", text: "Error: new_name parameter required for rename" }], - details: { action, serverName, success: false }, - }; - } - - const result = (await sendRequest(client, "textDocument/rename", { - textDocument: { uri }, - position, - newName: new_name, - })) as WorkspaceEdit | null; - - if (!result) { - output = "Rename returned no edits"; - } else { - const shouldApply = apply !== false; - if (shouldApply) { - const applied = await applyWorkspaceEdit(result, session.cwd); - output = `Applied rename:\n${applied.map((a) => ` ${a}`).join("\n")}`; - } else { - const preview = formatWorkspaceEdit(result, session.cwd); - output = `Rename preview:\n${preview.map((p) => ` ${p}`).join("\n")}`; - } - } - break; - } - - case "actions": { - if (!targetFile) { - return { - content: [{ type: "text", text: "Error: file parameter required for actions" }], - details: { action, serverName, success: false }, - }; - } - - const actionsMinVersion = client.diagnosticsVersion; - await refreshFile(client, targetFile); - const diagnostics = await waitForDiagnostics(client, uri, 3000, undefined, actionsMinVersion); - const endLine = (end_line ?? line ?? 1) - 1; - const endCharacter = (end_character ?? column ?? 1) - 1; - const range = { start: position, end: { line: endLine, character: endCharacter } }; - const relevantDiagnostics = diagnostics.filter( - (d) => d.range.start.line <= range.end.line && d.range.end.line >= range.start.line, - ); - - const codeActionContext: { diagnostics: Diagnostic[]; only?: string[] } = { - diagnostics: relevantDiagnostics, - }; - if (kind) { - codeActionContext.only = [kind]; - } - - const result = (await sendRequest(client, "textDocument/codeAction", { - textDocument: { uri }, - range, - context: codeActionContext, - })) as Array | null; - - if (!result || result.length === 0) { - output = "No code actions available"; - } else if (action_index !== undefined) { - // Apply specific action - if (action_index < 0 || action_index >= result.length) { - return { - content: [ - { - type: "text", - text: `Error: action_index ${action_index} out of range (0-${result.length - 1})`, - }, - ], - details: { action, serverName, success: false }, - }; - } - - const isCommand = (candidate: CodeAction | Command): candidate is Command => - typeof (candidate as Command).command === "string"; - const isCodeAction = (candidate: CodeAction | Command): candidate is CodeAction => - !isCommand(candidate); - const getCommandPayload = ( - candidate: CodeAction | Command, - ): { command: string; arguments?: unknown[] } | null => { - if (isCommand(candidate)) { - return { command: candidate.command, arguments: candidate.arguments }; - } - if (candidate.command) { - return { command: candidate.command.command, arguments: candidate.command.arguments }; - } - return null; - }; - - const codeAction = result[action_index]; - - // Resolve if needed - let resolvedAction = codeAction; - if ( - isCodeAction(codeAction) && - !codeAction.edit && - codeAction.data && - client.serverCapabilities?.codeActionProvider - ) { - const provider = client.serverCapabilities.codeActionProvider; - if (typeof provider === "object" && provider.resolveProvider) { - resolvedAction = (await sendRequest(client, "codeAction/resolve", codeAction)) as CodeAction; - } - } - - if (isCodeAction(resolvedAction) && resolvedAction.edit) { - const applied = await applyWorkspaceEdit(resolvedAction.edit, session.cwd); - output = `Applied "${codeAction.title}":\n${applied.map((a) => ` ${a}`).join("\n")}`; - } else { - const commandPayload = getCommandPayload(resolvedAction); - if (commandPayload) { - await sendRequest(client, "workspace/executeCommand", commandPayload); - output = `Executed "${codeAction.title}"`; - } else { - output = `Code action "${codeAction.title}" has no edits or command to apply`; - } - } - } else { - // List available actions - const lines = result.map((actionItem, i) => { - if ("kind" in actionItem || "isPreferred" in actionItem || "edit" in actionItem) { - const actionDetails = actionItem as CodeAction; - const preferred = actionDetails.isPreferred ? " (preferred)" : ""; - const kindInfo = actionDetails.kind ? ` [${actionDetails.kind}]` : ""; - return ` [${i}] ${actionDetails.title}${kindInfo}${preferred}`; - } - return ` [${i}] ${actionItem.title}`; + const lines = calls.map((call) => { + const loc = { uri: call.from.uri, range: call.from.selectionRange }; + const detail = call.from.detail ? ` (${call.from.detail})` : ""; + return ` ${call.from.name}${detail} @ ${formatLocation(loc, this.session.cwd)}`; }); - output = `Available code actions:\n${lines.join("\n")}\n\nUse action_index parameter to apply a specific action.`; + output = `Found ${calls.length} caller(s) of "${item.name}":\n${lines.join("\n")}`; } - break; - } + } else { + const calls = (await sendRequest(client, "callHierarchy/outgoingCalls", { item })) as + | CallHierarchyOutgoingCall[] + | null; - case "incoming_calls": - case "outgoing_calls": { - // First, prepare the call hierarchy item at the cursor position - const prepareResult = (await sendRequest(client, "textDocument/prepareCallHierarchy", { - textDocument: { uri }, - position, - })) as CallHierarchyItem[] | null; - - if (!prepareResult || prepareResult.length === 0) { - output = "No callable symbol found at this position"; - break; - } - - const item = prepareResult[0]; - - if (action === "incoming_calls") { - const calls = (await sendRequest(client, "callHierarchy/incomingCalls", { item })) as - | CallHierarchyIncomingCall[] - | null; - - if (!calls || calls.length === 0) { - output = `No callers found for "${item.name}"`; - } else { - const lines = calls.map((call) => { - const loc = { uri: call.from.uri, range: call.from.selectionRange }; - const detail = call.from.detail ? ` (${call.from.detail})` : ""; - return ` ${call.from.name}${detail} @ ${formatLocation(loc, session.cwd)}`; - }); - output = `Found ${calls.length} caller(s) of "${item.name}":\n${lines.join("\n")}`; - } + if (!calls || calls.length === 0) { + output = `"${item.name}" doesn't call any functions`; } else { - const calls = (await sendRequest(client, "callHierarchy/outgoingCalls", { item })) as - | CallHierarchyOutgoingCall[] - | null; - - if (!calls || calls.length === 0) { - output = `"${item.name}" doesn't call any functions`; - } else { - const lines = calls.map((call) => { - const loc = { uri: call.to.uri, range: call.to.selectionRange }; - const detail = call.to.detail ? ` (${call.to.detail})` : ""; - return ` ${call.to.name}${detail} @ ${formatLocation(loc, session.cwd)}`; - }); - output = `"${item.name}" calls ${calls.length} function(s):\n${lines.join("\n")}`; - } - } - break; - } - - // ===================================================================== - // Rust-Analyzer Specific Operations - // ===================================================================== - - case "flycheck": { - if (!hasCapability(serverConfig, "flycheck")) { - return { - content: [{ type: "text", text: "Error: flycheck requires rust-analyzer" }], - details: { action, serverName, success: false }, - }; - } - - await rustAnalyzer.flycheck(client, resolvedFile ?? undefined); - const collected: Array<{ filePath: string; diagnostic: Diagnostic }> = []; - for (const [diagUri, diags] of client.diagnostics.entries()) { - const relPath = path.relative(session.cwd, uriToFile(diagUri)); - for (const diag of diags) { - collected.push({ filePath: relPath, diagnostic: diag }); - } - } - - if (collected.length === 0) { - output = "Flycheck: no issues found"; - } else { - const summary = formatDiagnosticsSummary(collected.map((d) => d.diagnostic)); - const formatted = collected.slice(0, 20).map((d) => formatDiagnostic(d.diagnostic, d.filePath)); - const more = collected.length > 20 ? `\n ... and ${collected.length - 20} more` : ""; - output = `Flycheck ${summary}:\n${formatted.map((f) => ` ${f}`).join("\n")}${more}`; - } - break; - } - - case "expand_macro": { - if (!hasCapability(serverConfig, "expandMacro")) { - return { - content: [{ type: "text", text: "Error: expand_macro requires rust-analyzer" }], - details: { action, serverName, success: false }, - }; - } - - if (!targetFile) { - return { - content: [{ type: "text", text: "Error: file parameter required for expand_macro" }], - details: { action, serverName, success: false }, - }; - } - - const result = await rustAnalyzer.expandMacro(client, targetFile, line || 1, column || 1); - if (!result) { - output = "No macro expansion at this position"; - } else { - output = `Macro: ${result.name}\n\nExpansion:\n${result.expansion}`; - } - break; - } - - case "ssr": { - if (!hasCapability(serverConfig, "ssr")) { - return { - content: [{ type: "text", text: "Error: ssr requires rust-analyzer" }], - details: { action, serverName, success: false }, - }; - } - - if (!query) { - return { - content: [{ type: "text", text: "Error: query parameter (pattern) required for ssr" }], - details: { action, serverName, success: false }, - }; - } - - if (!replacement) { - return { - content: [{ type: "text", text: "Error: replacement parameter required for ssr" }], - details: { action, serverName, success: false }, - }; - } - - const shouldApply = apply === true; - const result = await rustAnalyzer.ssr(client, query, replacement, !shouldApply); - - if (shouldApply) { - const applied = await applyWorkspaceEdit(result, session.cwd); - output = - applied.length > 0 - ? `Applied SSR:\n${applied.map((a) => ` ${a}`).join("\n")}` - : "SSR: no matches found"; - } else { - const preview = formatWorkspaceEdit(result, session.cwd); - output = - preview.length > 0 - ? `SSR preview:\n${preview.map((p) => ` ${p}`).join("\n")}` - : "SSR: no matches found"; - } - break; - } - - case "runnables": { - if (!hasCapability(serverConfig, "runnables")) { - return { - content: [{ type: "text", text: "Error: runnables requires rust-analyzer" }], - details: { action, serverName, success: false }, - }; - } - - if (!targetFile) { - return { - content: [{ type: "text", text: "Error: file parameter required for runnables" }], - details: { action, serverName, success: false }, - }; - } - - const result = await rustAnalyzer.runnables(client, targetFile, line); - if (result.length === 0) { - output = "No runnables found"; - } else { - const lines = result.map((r) => { - const args = r.args?.cargoArgs?.join(" ") || ""; - return ` [${r.kind}] ${r.label}${args ? ` (cargo ${args})` : ""}`; + const lines = calls.map((call) => { + const loc = { uri: call.to.uri, range: call.to.selectionRange }; + const detail = call.to.detail ? ` (${call.to.detail})` : ""; + return ` ${call.to.name}${detail} @ ${formatLocation(loc, this.session.cwd)}`; }); - output = `Found ${result.length} runnable(s):\n${lines.join("\n")}`; + output = `"${item.name}" calls ${calls.length} function(s):\n${lines.join("\n")}`; } - break; } - - case "related_tests": { - if (!hasCapability(serverConfig, "relatedTests")) { - return { - content: [{ type: "text", text: "Error: related_tests requires rust-analyzer" }], - details: { action, serverName, success: false }, - }; - } - - if (!targetFile) { - return { - content: [{ type: "text", text: "Error: file parameter required for related_tests" }], - details: { action, serverName, success: false }, - }; - } - - const result = await rustAnalyzer.relatedTests(client, targetFile, line || 1, column || 1); - if (result.length === 0) { - output = "No related tests found"; - } else { - output = `Found ${result.length} related test(s):\n${result.map((t) => ` ${t}`).join("\n")}`; - } - break; - } - - case "reload_workspace": { - await rustAnalyzer.reloadWorkspace(client); - output = "Workspace reloaded successfully"; - break; - } - - default: - output = `Unknown action: ${action}`; + break; } - return { - content: [{ type: "text", text: output }], - details: { serverName, action, success: true }, - }; - } catch (err) { - const errorMessage = err instanceof Error ? err.message : String(err); - return { - content: [{ type: "text", text: `LSP error: ${errorMessage}` }], - details: { serverName, action, success: false }, - }; + // ===================================================================== + // Rust-Analyzer Specific Operations + // ===================================================================== + + case "flycheck": { + if (!hasCapability(serverConfig, "flycheck")) { + return { + content: [{ type: "text", text: "Error: flycheck requires rust-analyzer" }], + details: { action, serverName, success: false }, + }; + } + + await rustAnalyzer.flycheck(client, resolvedFile ?? undefined); + const collected: Array<{ filePath: string; diagnostic: Diagnostic }> = []; + for (const [diagUri, diags] of client.diagnostics.entries()) { + const relPath = path.relative(this.session.cwd, uriToFile(diagUri)); + for (const diag of diags) { + collected.push({ filePath: relPath, diagnostic: diag }); + } + } + + if (collected.length === 0) { + output = "Flycheck: no issues found"; + } else { + const summary = formatDiagnosticsSummary(collected.map((d) => d.diagnostic)); + const formatted = collected.slice(0, 20).map((d) => formatDiagnostic(d.diagnostic, d.filePath)); + const more = collected.length > 20 ? `\n ... and ${collected.length - 20} more` : ""; + output = `Flycheck ${summary}:\n${formatted.map((f) => ` ${f}`).join("\n")}${more}`; + } + break; + } + + case "expand_macro": { + if (!hasCapability(serverConfig, "expandMacro")) { + return { + content: [{ type: "text", text: "Error: expand_macro requires rust-analyzer" }], + details: { action, serverName, success: false }, + }; + } + + if (!targetFile) { + return { + content: [{ type: "text", text: "Error: file parameter required for expand_macro" }], + details: { action, serverName, success: false }, + }; + } + + const result = await rustAnalyzer.expandMacro(client, targetFile, line || 1, column || 1); + if (!result) { + output = "No macro expansion at this position"; + } else { + output = `Macro: ${result.name}\n\nExpansion:\n${result.expansion}`; + } + break; + } + + case "ssr": { + if (!hasCapability(serverConfig, "ssr")) { + return { + content: [{ type: "text", text: "Error: ssr requires rust-analyzer" }], + details: { action, serverName, success: false }, + }; + } + + if (!query) { + return { + content: [{ type: "text", text: "Error: query parameter (pattern) required for ssr" }], + details: { action, serverName, success: false }, + }; + } + + if (!replacement) { + return { + content: [{ type: "text", text: "Error: replacement parameter required for ssr" }], + details: { action, serverName, success: false }, + }; + } + + const shouldApply = apply === true; + const result = await rustAnalyzer.ssr(client, query, replacement, !shouldApply); + + if (shouldApply) { + const applied = await applyWorkspaceEdit(result, this.session.cwd); + output = + applied.length > 0 + ? `Applied SSR:\n${applied.map((a) => ` ${a}`).join("\n")}` + : "SSR: no matches found"; + } else { + const preview = formatWorkspaceEdit(result, this.session.cwd); + output = + preview.length > 0 + ? `SSR preview:\n${preview.map((p) => ` ${p}`).join("\n")}` + : "SSR: no matches found"; + } + break; + } + + case "runnables": { + if (!hasCapability(serverConfig, "runnables")) { + return { + content: [{ type: "text", text: "Error: runnables requires rust-analyzer" }], + details: { action, serverName, success: false }, + }; + } + + if (!targetFile) { + return { + content: [{ type: "text", text: "Error: file parameter required for runnables" }], + details: { action, serverName, success: false }, + }; + } + + const result = await rustAnalyzer.runnables(client, targetFile, line); + if (result.length === 0) { + output = "No runnables found"; + } else { + const lines = result.map((r) => { + const args = r.args?.cargoArgs?.join(" ") || ""; + return ` [${r.kind}] ${r.label}${args ? ` (cargo ${args})` : ""}`; + }); + output = `Found ${result.length} runnable(s):\n${lines.join("\n")}`; + } + break; + } + + case "related_tests": { + if (!hasCapability(serverConfig, "relatedTests")) { + return { + content: [{ type: "text", text: "Error: related_tests requires rust-analyzer" }], + details: { action, serverName, success: false }, + }; + } + + if (!targetFile) { + return { + content: [{ type: "text", text: "Error: file parameter required for related_tests" }], + details: { action, serverName, success: false }, + }; + } + + const result = await rustAnalyzer.relatedTests(client, targetFile, line || 1, column || 1); + if (result.length === 0) { + output = "No related tests found"; + } else { + output = `Found ${result.length} related test(s):\n${result.map((t) => ` ${t}`).join("\n")}`; + } + break; + } + + case "reload_workspace": { + await rustAnalyzer.reloadWorkspace(client); + output = "Workspace reloaded successfully"; + break; + } + + default: + output = `Unknown action: ${action}`; } - }, - }; + + return { + content: [{ type: "text", text: output }], + details: { serverName, action, success: true }, + }; + } catch (err) { + const errorMessage = err instanceof Error ? err.message : String(err); + return { + content: [{ type: "text", text: `LSP error: ${errorMessage}` }], + details: { serverName, action, success: false }, + }; + } + } } diff --git a/packages/coding-agent/src/core/tools/notebook.ts b/packages/coding-agent/src/core/tools/notebook.ts index 46b9d28d6..d4811bfb3 100644 --- a/packages/coding-agent/src/core/tools/notebook.ts +++ b/packages/coding-agent/src/core/tools/notebook.ts @@ -1,8 +1,8 @@ -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import { StringEnum } from "@oh-my-pi/pi-ai"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; -import { Type } from "@sinclair/typebox"; +import { type Static, Type } from "@sinclair/typebox"; import type { Theme } from "../../modes/interactive/theme/theme"; import type { RenderResultOptions } from "../custom-tools/types"; import type { ToolSession } from "../sdk"; @@ -63,133 +63,135 @@ function splitIntoLines(content: string): string[] { return content.split("\n").map((line, i, arr) => (i < arr.length - 1 ? `${line}\n` : line)); } -export function createNotebookTool(session: ToolSession): AgentTool { - return { - name: "notebook", - label: "Notebook", - description: - "Completely replaces the contents of a specific cell in a Jupyter notebook (.ipynb file) with new source. Jupyter notebooks are interactive documents that combine code, text, and visualizations, commonly used for data analysis and scientific computing. The notebook_path parameter must be an absolute path, not a relative path. The cell_number is 0-indexed. Use edit_mode=insert to add a new cell at the index specified by cell_number. Use edit_mode=delete to delete the cell at the index specified by cell_number.", - parameters: notebookSchema, - execute: async ( - _toolCallId: string, - { - action, - notebook_path, - cell_index, - content, - cell_type, - }: { action: string; notebook_path: string; cell_index: number; content?: string; cell_type?: string }, - signal?: AbortSignal, - ) => { - const absolutePath = resolveToCwd(notebook_path, session.cwd); +type NotebookParams = Static; - return untilAborted(signal, async () => { - // Check if file exists - const file = Bun.file(absolutePath); - if (!(await file.exists())) { - throw new Error(`Notebook not found: ${notebook_path}`); +export class NotebookTool implements AgentTool { + public readonly name = "notebook"; + public readonly label = "Notebook"; + public readonly description = + "Completely replaces the contents of a specific cell in a Jupyter notebook (.ipynb file) with new source. Jupyter notebooks are interactive documents that combine code, text, and visualizations, commonly used for data analysis and scientific computing. The notebook_path parameter must be an absolute path, not a relative path. The cell_number is 0-indexed. Use edit_mode=insert to add a new cell at the index specified by cell_number. Use edit_mode=delete to delete the cell at the index specified by cell_number."; + public readonly parameters = notebookSchema; + + private readonly session: ToolSession; + + constructor(session: ToolSession) { + this.session = session; + } + + public async execute( + _toolCallId: string, + params: NotebookParams, + signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const { action, notebook_path, cell_index, content, cell_type } = params; + const absolutePath = resolveToCwd(notebook_path, this.session.cwd); + + return untilAborted(signal, async () => { + // Check if file exists + const file = Bun.file(absolutePath); + if (!(await file.exists())) { + throw new Error(`Notebook not found: ${notebook_path}`); + } + + // Read and parse notebook + let notebook: Notebook; + try { + notebook = await file.json(); + } catch { + throw new Error(`Invalid JSON in notebook: ${notebook_path}`); + } + + // Validate notebook structure + if (!notebook.cells || !Array.isArray(notebook.cells)) { + throw new Error(`Invalid notebook structure (missing cells array): ${notebook_path}`); + } + + const cellCount = notebook.cells.length; + + // Validate cell_index based on action + if (action === "insert") { + if (cell_index < 0 || cell_index > cellCount) { + throw new Error(`Cell index ${cell_index} out of range for insert (0-${cellCount}) in ${notebook_path}`); } - - // Read and parse notebook - let notebook: Notebook; - try { - notebook = await file.json(); - } catch { - throw new Error(`Invalid JSON in notebook: ${notebook_path}`); + } else { + if (cell_index < 0 || cell_index >= cellCount) { + throw new Error(`Cell index ${cell_index} out of range (0-${cellCount - 1}) in ${notebook_path}`); } + } - // Validate notebook structure - if (!notebook.cells || !Array.isArray(notebook.cells)) { - throw new Error(`Invalid notebook structure (missing cells array): ${notebook_path}`); + // Validate content for edit/insert + if ((action === "edit" || action === "insert") && content === undefined) { + throw new Error(`Content is required for ${action} action`); + } + + // Perform the action + let resultMessage: string; + let finalCellType: string | undefined; + let cellSource: string[] | undefined; + + switch (action) { + case "edit": { + const sourceLines = splitIntoLines(content!); + notebook.cells[cell_index].source = sourceLines; + finalCellType = notebook.cells[cell_index].cell_type; + cellSource = sourceLines; + resultMessage = `Replaced cell ${cell_index} (${finalCellType})`; + break; } - - const cellCount = notebook.cells.length; - - // Validate cell_index based on action - if (action === "insert") { - if (cell_index < 0 || cell_index > cellCount) { - throw new Error( - `Cell index ${cell_index} out of range for insert (0-${cellCount}) in ${notebook_path}`, - ); - } - } else { - if (cell_index < 0 || cell_index >= cellCount) { - throw new Error(`Cell index ${cell_index} out of range (0-${cellCount - 1}) in ${notebook_path}`); + case "insert": { + const sourceLines = splitIntoLines(content!); + const newCellType = (cell_type as "code" | "markdown") || "code"; + const newCell: NotebookCell = { + cell_type: newCellType, + source: sourceLines, + metadata: {}, + }; + if (newCellType === "code") { + newCell.execution_count = null; + newCell.outputs = []; } + notebook.cells.splice(cell_index, 0, newCell); + finalCellType = newCellType; + cellSource = sourceLines; + resultMessage = `Inserted ${newCellType} cell at position ${cell_index}`; + break; } - - // Validate content for edit/insert - if ((action === "edit" || action === "insert") && content === undefined) { - throw new Error(`Content is required for ${action} action`); + case "delete": { + const removedCell = notebook.cells[cell_index]; + finalCellType = removedCell.cell_type; + cellSource = removedCell.source; + notebook.cells.splice(cell_index, 1); + resultMessage = `Deleted cell ${cell_index} (${finalCellType})`; + break; } - - // Perform the action - let resultMessage: string; - let finalCellType: string | undefined; - let cellSource: string[] | undefined; - - switch (action) { - case "edit": { - const sourceLines = splitIntoLines(content!); - notebook.cells[cell_index].source = sourceLines; - finalCellType = notebook.cells[cell_index].cell_type; - cellSource = sourceLines; - resultMessage = `Replaced cell ${cell_index} (${finalCellType})`; - break; - } - case "insert": { - const sourceLines = splitIntoLines(content!); - const newCellType = (cell_type as "code" | "markdown") || "code"; - const newCell: NotebookCell = { - cell_type: newCellType, - source: sourceLines, - metadata: {}, - }; - if (newCellType === "code") { - newCell.execution_count = null; - newCell.outputs = []; - } - notebook.cells.splice(cell_index, 0, newCell); - finalCellType = newCellType; - cellSource = sourceLines; - resultMessage = `Inserted ${newCellType} cell at position ${cell_index}`; - break; - } - case "delete": { - const removedCell = notebook.cells[cell_index]; - finalCellType = removedCell.cell_type; - cellSource = removedCell.source; - notebook.cells.splice(cell_index, 1); - resultMessage = `Deleted cell ${cell_index} (${finalCellType})`; - break; - } - default: { - throw new Error(`Invalid action: ${action}`); - } + default: { + throw new Error(`Invalid action: ${action}`); } + } - // Write back with single-space indentation - await Bun.write(absolutePath, JSON.stringify(notebook, null, 1)); + // Write back with single-space indentation + await Bun.write(absolutePath, JSON.stringify(notebook, null, 1)); - const newCellCount = notebook.cells.length; - return { - content: [ - { - type: "text", - text: `${resultMessage}. Notebook now has ${newCellCount} cells.`, - }, - ], - details: { - action: action as "edit" | "insert" | "delete", - cellIndex: cell_index, - cellType: finalCellType, - totalCells: newCellCount, - cellSource, + const newCellCount = notebook.cells.length; + return { + content: [ + { + type: "text", + text: `${resultMessage}. Notebook now has ${newCellCount} cells.`, }, - }; - }); - }, - }; + ], + details: { + action: action as "edit" | "insert" | "delete", + cellIndex: cell_index, + cellType: finalCellType, + totalCells: newCellCount, + cellSource, + }, + }; + }); + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/output.ts b/packages/coding-agent/src/core/tools/output.ts index 58f1d4706..cca071316 100644 --- a/packages/coding-agent/src/core/tools/output.ts +++ b/packages/coding-agent/src/core/tools/output.ts @@ -6,8 +6,8 @@ import * as fs from "node:fs"; import * as path from "node:path"; -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; -import { StringEnum, type TextContent } from "@oh-my-pi/pi-ai"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; +import { StringEnum } from "@oh-my-pi/pi-ai"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; import { Type } from "@sinclair/typebox"; @@ -232,166 +232,182 @@ function extractPreviewLines(content: string, maxLines: number): string[] { return preview; } -export function createOutputTool(session: ToolSession): AgentTool { - return { - name: "output", - label: "Output", - description: renderPromptTemplate(outputDescription), - parameters: outputSchema, - execute: async ( - _toolCallId: string, - params: { - ids: string[]; - format?: "raw" | "json" | "stripped"; - query?: string; - offset?: number; - limit?: number; - }, - ): Promise<{ content: TextContent[]; details: OutputToolDetails }> => { - const sessionFile = session.getSessionFile(); +type OutputParams = { + ids: string[]; + format?: "raw" | "json" | "stripped"; + query?: string; + offset?: number; + limit?: number; +}; - if (!sessionFile) { - return { - content: [{ type: "text", text: "No session - output artifacts unavailable" }], - details: { outputs: [], notFound: params.ids }, - }; +/** + * Output tool for reading agent/task outputs by ID. + * + * Resolves IDs like "reviewer_0" to artifact paths in the current session. + */ +export class OutputTool implements AgentTool { + public readonly name = "output"; + public readonly label = "Output"; + public readonly description: string; + public readonly parameters = outputSchema; + + private readonly session: ToolSession; + + constructor(session: ToolSession) { + this.session = session; + this.description = renderPromptTemplate(outputDescription); + } + + public async execute( + _toolCallId: string, + params: OutputParams, + _signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const sessionFile = this.session.getSessionFile(); + + if (!sessionFile) { + return { + content: [{ type: "text", text: "No session - output artifacts unavailable" }], + details: { outputs: [], notFound: params.ids }, + }; + } + + const artifactsDir = getArtifactsDir(sessionFile); + if (!artifactsDir || !fs.existsSync(artifactsDir)) { + return { + content: [{ type: "text", text: "No artifacts directory found" }], + details: { outputs: [], notFound: params.ids }, + }; + } + + const outputs: OutputEntry[] = []; + const notFound: string[] = []; + const outputContentById = new Map(); + const query = params.query?.trim(); + const wantsQuery = query !== undefined && query.length > 0; + const format = params.format ?? (wantsQuery ? "json" : "raw"); + + if (wantsQuery && (params.offset !== undefined || params.limit !== undefined)) { + throw new Error("query cannot be combined with offset/limit"); + } + + const queryResults: Array<{ id: string; value: unknown }> = []; + + for (const id of params.ids) { + const outputPath = path.join(artifactsDir, `${id}.out.md`); + + if (!fs.existsSync(outputPath)) { + notFound.push(id); + continue; } - const artifactsDir = getArtifactsDir(sessionFile); - if (!artifactsDir || !fs.existsSync(artifactsDir)) { - return { - content: [{ type: "text", text: "No artifacts directory found" }], - details: { outputs: [], notFound: params.ids }, - }; - } + const rawContent = fs.readFileSync(outputPath, "utf-8"); + const rawLines = rawContent.split("\n"); + const totalLines = rawLines.length; + const totalChars = rawContent.length; - const outputs: OutputEntry[] = []; - const notFound: string[] = []; - const outputContentById = new Map(); - const query = params.query?.trim(); - const wantsQuery = query !== undefined && query.length > 0; - const format = params.format ?? (wantsQuery ? "json" : "raw"); + let selectedContent = rawContent; + let range: OutputRange | undefined; - if (wantsQuery && (params.offset !== undefined || params.limit !== undefined)) { - throw new Error("query cannot be combined with offset/limit"); - } - - const queryResults: Array<{ id: string; value: unknown }> = []; - - for (const id of params.ids) { - const outputPath = path.join(artifactsDir, `${id}.out.md`); - - if (!fs.existsSync(outputPath)) { - notFound.push(id); - continue; + if (wantsQuery && query) { + let jsonValue: unknown; + try { + jsonValue = JSON.parse(rawContent); + } catch (err) { + const message = err instanceof Error ? err.message : String(err); + throw new Error(`Output ${id} is not valid JSON: ${message}`); } - - const rawContent = fs.readFileSync(outputPath, "utf-8"); - const rawLines = rawContent.split("\n"); - const totalLines = rawLines.length; - const totalChars = rawContent.length; - - let selectedContent = rawContent; - let range: OutputRange | undefined; - - if (wantsQuery && query) { - let jsonValue: unknown; - try { - jsonValue = JSON.parse(rawContent); - } catch (err) { - const message = err instanceof Error ? err.message : String(err); - throw new Error(`Output ${id} is not valid JSON: ${message}`); - } - const value = applyQuery(jsonValue, query); - queryResults.push({ id, value }); - try { - selectedContent = JSON.stringify(value, null, 2) ?? "null"; - } catch { - selectedContent = String(value); - } - } else if (params.offset !== undefined || params.limit !== undefined) { - const startLine = Math.max(1, params.offset ?? 1); - if (startLine > totalLines) { - throw new Error( - `Offset ${params.offset ?? startLine} is beyond end of output (${totalLines} lines) for ${id}`, - ); - } - const effectiveLimit = params.limit ?? totalLines - startLine + 1; - const endLine = Math.min(totalLines, startLine + effectiveLimit - 1); - const selectedLines = rawLines.slice(startLine - 1, endLine); - selectedContent = selectedLines.join("\n"); - range = { startLine, endLine, totalLines }; + const value = applyQuery(jsonValue, query); + queryResults.push({ id, value }); + try { + selectedContent = JSON.stringify(value, null, 2) ?? "null"; + } catch { + selectedContent = String(value); } - - outputContentById.set(id, selectedContent); - outputs.push({ - id, - path: outputPath, - lineCount: wantsQuery ? selectedContent.split("\n").length : totalLines, - charCount: wantsQuery ? selectedContent.length : totalChars, - provenance: parseOutputProvenance(id), - previewLines: extractPreviewLines(selectedContent, 4), - range, - query: query, - }); + } else if (params.offset !== undefined || params.limit !== undefined) { + const startLine = Math.max(1, params.offset ?? 1); + if (startLine > totalLines) { + throw new Error( + `Offset ${params.offset ?? startLine} is beyond end of output (${totalLines} lines) for ${id}`, + ); + } + const effectiveLimit = params.limit ?? totalLines - startLine + 1; + const endLine = Math.min(totalLines, startLine + effectiveLimit - 1); + const selectedLines = rawLines.slice(startLine - 1, endLine); + selectedContent = selectedLines.join("\n"); + range = { startLine, endLine, totalLines }; } - // Error case: some IDs not found - if (notFound.length > 0) { - const available = listAvailableOutputs(artifactsDir); - const errorMsg = - available.length > 0 - ? `Not found: ${notFound.join(", ")}\nAvailable: ${available.join(", ")}` - : `Not found: ${notFound.join(", ")}\nNo outputs available in current session`; + outputContentById.set(id, selectedContent); + outputs.push({ + id, + path: outputPath, + lineCount: wantsQuery ? selectedContent.split("\n").length : totalLines, + charCount: wantsQuery ? selectedContent.length : totalChars, + provenance: parseOutputProvenance(id), + previewLines: extractPreviewLines(selectedContent, 4), + range, + query: query, + }); + } - return { - content: [{ type: "text", text: errorMsg }], - details: { outputs, notFound, availableIds: available }, - }; - } - - // Success: build response based on format - let contentText: string; - - if (format === "json") { - const jsonData = wantsQuery - ? queryResults - : outputs.map((o) => ({ - id: o.id, - lineCount: o.lineCount, - charCount: o.charCount, - provenance: o.provenance, - previewLines: o.previewLines, - range: o.range, - content: outputContentById.get(o.id) ?? "", - })); - contentText = JSON.stringify(jsonData, null, 2); - } else { - // raw or stripped - const parts = outputs.map((o) => { - let content = outputContentById.get(o.id) ?? ""; - if (format === "stripped") { - content = stripAnsi(content); - } - if (o.range && o.range.endLine < o.range.totalLines) { - const nextOffset = o.range.endLine + 1; - content += `\n\n[Showing lines ${o.range.startLine}-${o.range.endLine} of ${o.range.totalLines}. Use offset=${nextOffset} to continue]`; - } - // Add header for multiple outputs - if (outputs.length > 1) { - return `=== ${o.id} (${o.lineCount} lines, ${formatBytes(o.charCount)}) ===\n${content}`; - } - return content; - }); - contentText = parts.join("\n\n"); - } + // Error case: some IDs not found + if (notFound.length > 0) { + const available = listAvailableOutputs(artifactsDir); + const errorMsg = + available.length > 0 + ? `Not found: ${notFound.join(", ")}\nAvailable: ${available.join(", ")}` + : `Not found: ${notFound.join(", ")}\nNo outputs available in current session`; return { - content: [{ type: "text", text: contentText }], - details: { outputs }, + content: [{ type: "text", text: errorMsg }], + details: { outputs, notFound, availableIds: available }, }; - }, - }; + } + + // Success: build response based on format + let contentText: string; + + if (format === "json") { + const jsonData = wantsQuery + ? queryResults + : outputs.map((o) => ({ + id: o.id, + lineCount: o.lineCount, + charCount: o.charCount, + provenance: o.provenance, + previewLines: o.previewLines, + range: o.range, + content: outputContentById.get(o.id) ?? "", + })); + contentText = JSON.stringify(jsonData, null, 2); + } else { + // raw or stripped + const parts = outputs.map((o) => { + let content = outputContentById.get(o.id) ?? ""; + if (format === "stripped") { + content = stripAnsi(content); + } + if (o.range && o.range.endLine < o.range.totalLines) { + const nextOffset = o.range.endLine + 1; + content += `\n\n[Showing lines ${o.range.startLine}-${o.range.endLine} of ${o.range.totalLines}. Use offset=${nextOffset} to continue]`; + } + // Add header for multiple outputs + if (outputs.length > 1) { + return `=== ${o.id} (${o.lineCount} lines, ${formatBytes(o.charCount)}) ===\n${content}`; + } + return content; + }); + contentText = parts.join("\n\n"); + } + + return { + content: [{ type: "text", text: contentText }], + details: { outputs }, + }; + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/patch/shared.ts b/packages/coding-agent/src/core/tools/patch/shared.ts index adef038ed..195a95ab1 100644 --- a/packages/coding-agent/src/core/tools/patch/shared.ts +++ b/packages/coding-agent/src/core/tools/patch/shared.ts @@ -8,7 +8,14 @@ import { Text } from "@oh-my-pi/pi-tui"; import { getLanguageFromPath, type Theme } from "../../../modes/interactive/theme/theme"; import type { RenderResultOptions } from "../../custom-tools/types"; import type { FileDiagnosticsResult } from "../lsp/index"; -import { createToolUIKit, formatExpandHint, getDiffStats, shortenPath, truncateDiffByHunk } from "../render-utils"; +import { + createToolUIKit, + formatExpandHint, + getDiffStats, + shortenPath, + type ToolUIKit, + truncateDiffByHunk, +} from "../render-utils"; import type { DiffError, DiffResult, Operation } from "./types"; // ═══════════════════════════════════════════════════════════════════════════ @@ -94,7 +101,7 @@ function renderDiffSection( rawPath: string, expanded: boolean, uiTheme: Theme, - ui: ReturnType, + ui: ToolUIKit, renderDiffFn: (t: string, o?: { filePath?: string }) => string, ): string { let text = ""; diff --git a/packages/coding-agent/src/core/tools/python.ts b/packages/coding-agent/src/core/tools/python.ts index bf18d9efd..9ee61f0e0 100644 --- a/packages/coding-agent/src/core/tools/python.ts +++ b/packages/coding-agent/src/core/tools/python.ts @@ -1,9 +1,9 @@ import { relative, resolve, sep } from "node:path"; -import type { AgentTool, AgentToolContext } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { ImageContent } from "@oh-my-pi/pi-ai"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text, truncateToWidth } from "@oh-my-pi/pi-tui"; -import { Type } from "@sinclair/typebox"; +import { type Static, Type } from "@sinclair/typebox"; import { truncateToVisualLines } from "../../modes/interactive/components/visual-truncate"; import type { Theme } from "../../modes/interactive/theme/theme"; import pythonDescription from "../../prompts/tools/python.md" with { type: "text" }; @@ -116,160 +116,163 @@ export function getPythonToolDescription(): string { return renderPromptTemplate(pythonDescription, { categories }); } -interface CreatePythonToolOptions { +export interface PythonToolOptions { proxyExecutor?: PythonProxyExecutor; } -export function createPythonTool( - session: ToolSession | null, - options?: CreatePythonToolOptions, -): AgentTool { - const { proxyExecutor } = options ?? {}; +export class PythonTool implements AgentTool { + public readonly name = "python"; + public readonly label = "Python"; + public readonly description: string; + public readonly parameters = pythonSchema; - return { - name: "python", - label: "Python", - description: getPythonToolDescription(), - parameters: pythonSchema, - execute: async ( - _toolCallId: string, - params: PythonToolParams, - signal?: AbortSignal, - onUpdate?, - _ctx?: AgentToolContext, - ) => { - if (proxyExecutor) { - return proxyExecutor(params, signal); + private readonly session: ToolSession | null; + private readonly proxyExecutor?: PythonProxyExecutor; + + constructor(session: ToolSession | null, options?: PythonToolOptions) { + this.session = session; + this.proxyExecutor = options?.proxyExecutor; + this.description = getPythonToolDescription(); + } + + public async execute( + _toolCallId: string, + params: Static, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + _ctx?: AgentToolContext, + ): Promise> { + if (this.proxyExecutor) { + return this.proxyExecutor(params, signal); + } + + if (!this.session) { + throw new Error("Python tool requires a session when not using proxy executor"); + } + + const { code, timeout, workdir, reset } = params; + const controller = new AbortController(); + const onAbort = () => controller.abort(); + signal?.addEventListener("abort", onAbort, { once: true }); + + try { + if (signal?.aborted) { + throw new Error("Aborted"); } - if (!session) { - throw new Error("Python tool requires a session when not using proxy executor"); - } - - const { code, timeout, workdir, reset } = params; - const controller = new AbortController(); - const onAbort = () => controller.abort(); - signal?.addEventListener("abort", onAbort, { once: true }); - + const commandCwd = workdir ? resolveToCwd(workdir, this.session.cwd) : this.session.cwd; + let cwdStat: Awaited>; try { - if (signal?.aborted) { - throw new Error("Aborted"); - } + cwdStat = await Bun.file(commandCwd).stat(); + } catch { + throw new Error(`Working directory does not exist: ${commandCwd}`); + } + if (!cwdStat.isDirectory()) { + throw new Error(`Working directory is not a directory: ${commandCwd}`); + } - const commandCwd = workdir ? resolveToCwd(workdir, session.cwd) : session.cwd; - let cwdStat: Awaited>; - try { - cwdStat = await Bun.file(commandCwd).stat(); - } catch { - throw new Error(`Working directory does not exist: ${commandCwd}`); - } - if (!cwdStat.isDirectory()) { - throw new Error(`Working directory is not a directory: ${commandCwd}`); - } + const maxTailBytes = DEFAULT_MAX_BYTES * 2; + const tailChunks: Array<{ text: string; bytes: number }> = []; + let tailBytes = 0; + const jsonOutputs: unknown[] = []; + const images: ImageContent[] = []; - const maxTailBytes = DEFAULT_MAX_BYTES * 2; - const tailChunks: Array<{ text: string; bytes: number }> = []; - let tailBytes = 0; - const jsonOutputs: unknown[] = []; - const images: ImageContent[] = []; - - const sessionFile = session.getSessionFile?.() ?? undefined; - const sessionId = sessionFile ? `session:${sessionFile}:workdir:${commandCwd}` : `cwd:${commandCwd}`; - const executorOptions: PythonExecutorOptions = { - cwd: commandCwd, - timeout: timeout ? timeout * 1000 : undefined, - signal: controller.signal, - sessionId, - kernelMode: session.settings?.getPythonKernelMode?.() ?? "session", - useSharedGateway: session.settings?.getPythonSharedGateway?.() ?? true, - reset, - onChunk: (chunk) => { - const chunkBytes = Buffer.byteLength(chunk, "utf-8"); - tailChunks.push({ text: chunk, bytes: chunkBytes }); - tailBytes += chunkBytes; - while (tailBytes > maxTailBytes && tailChunks.length > 1) { - const removed = tailChunks.shift(); - if (removed) { - tailBytes -= removed.bytes; - } + const sessionFile = this.session.getSessionFile?.() ?? undefined; + const sessionId = sessionFile ? `session:${sessionFile}:workdir:${commandCwd}` : `cwd:${commandCwd}`; + const executorOptions: PythonExecutorOptions = { + cwd: commandCwd, + timeout: timeout ? timeout * 1000 : undefined, + signal: controller.signal, + sessionId, + kernelMode: this.session.settings?.getPythonKernelMode?.() ?? "session", + useSharedGateway: this.session.settings?.getPythonSharedGateway?.() ?? true, + reset, + onChunk: (chunk) => { + const chunkBytes = Buffer.byteLength(chunk, "utf-8"); + tailChunks.push({ text: chunk, bytes: chunkBytes }); + tailBytes += chunkBytes; + while (tailBytes > maxTailBytes && tailChunks.length > 1) { + const removed = tailChunks.shift(); + if (removed) { + tailBytes -= removed.bytes; } - if (onUpdate) { - const tailText = tailChunks.map((entry) => entry.text).join(""); - const truncation = truncateTail(tailText); - onUpdate({ - content: [{ type: "text", text: truncation.content || "" }], - details: truncation.truncated ? { truncation } : undefined, - }); - } - }, + } + if (onUpdate) { + const tailText = tailChunks.map((entry) => entry.text).join(""); + const truncation = truncateTail(tailText); + onUpdate({ + content: [{ type: "text", text: truncation.content || "" }], + details: truncation.truncated ? { truncation } : undefined, + }); + } + }, + }; + + const result = await executePython(code, executorOptions); + + const statusEvents: PythonStatusEvent[] = []; + for (const output of result.displayOutputs) { + if (output.type === "json") { + jsonOutputs.push(output.data); + } + if (output.type === "image") { + images.push({ type: "image", data: output.data, mimeType: output.mimeType }); + } + if (output.type === "status") { + statusEvents.push(output.event); + } + } + + if (result.cancelled) { + throw new Error(result.output || "Command aborted"); + } + + const truncation = truncateTail(result.output); + let outputText = + truncation.content || (jsonOutputs.length > 0 || images.length > 0 ? "(no text output)" : "(no output)"); + let details: PythonToolDetails | undefined; + + if (truncation.truncated) { + const fullOutputSuffix = result.fullOutputPath ? ` Full output: ${result.fullOutputPath}` : ""; + details = { + truncation, + fullOutputPath: result.fullOutputPath, + jsonOutputs: jsonOutputs, + images, + statusEvents: statusEvents.length > 0 ? statusEvents : undefined, }; - const result = await executePython(code, executorOptions); + const startLine = truncation.totalLines - truncation.outputLines + 1; + const endLine = truncation.totalLines; - const statusEvents: PythonStatusEvent[] = []; - for (const output of result.displayOutputs) { - if (output.type === "json") { - jsonOutputs.push(output.data); - } - if (output.type === "image") { - images.push({ type: "image", data: output.data, mimeType: output.mimeType }); - } - if (output.type === "status") { - statusEvents.push(output.event); - } + if (truncation.lastLinePartial) { + const lastLineSize = formatSize(Buffer.byteLength(result.output.split("\n").pop() || "", "utf-8")); + outputText += `\n\n[Showing last ${formatSize(truncation.outputBytes)} of line ${endLine} (line is ${lastLineSize})${fullOutputSuffix}]`; + } else if (truncation.truncatedBy === "lines") { + outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines}${fullOutputSuffix}]`; + } else { + outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines} (${formatSize(DEFAULT_MAX_BYTES)} limit)${fullOutputSuffix}]`; } - - if (result.cancelled) { - throw new Error(result.output || "Command aborted"); - } - - const truncation = truncateTail(result.output); - let outputText = - truncation.content || (jsonOutputs.length > 0 || images.length > 0 ? "(no text output)" : "(no output)"); - let details: PythonToolDetails | undefined; - - if (truncation.truncated) { - const fullOutputSuffix = result.fullOutputPath ? ` Full output: ${result.fullOutputPath}` : ""; - details = { - truncation, - fullOutputPath: result.fullOutputPath, - jsonOutputs: jsonOutputs, - images, - statusEvents: statusEvents.length > 0 ? statusEvents : undefined, - }; - - const startLine = truncation.totalLines - truncation.outputLines + 1; - const endLine = truncation.totalLines; - - if (truncation.lastLinePartial) { - const lastLineSize = formatSize(Buffer.byteLength(result.output.split("\n").pop() || "", "utf-8")); - outputText += `\n\n[Showing last ${formatSize(truncation.outputBytes)} of line ${endLine} (line is ${lastLineSize})${fullOutputSuffix}]`; - } else if (truncation.truncatedBy === "lines") { - outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines}${fullOutputSuffix}]`; - } else { - outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines} (${formatSize(DEFAULT_MAX_BYTES)} limit)${fullOutputSuffix}]`; - } - } - - if (!details && (jsonOutputs.length > 0 || images.length > 0 || statusEvents.length > 0)) { - details = { - jsonOutputs: jsonOutputs.length > 0 ? jsonOutputs : undefined, - images: images.length > 0 ? images : undefined, - statusEvents: statusEvents.length > 0 ? statusEvents : undefined, - }; - } - - if (result.exitCode !== 0 && result.exitCode !== undefined) { - outputText += `\n\nCommand exited with code ${result.exitCode}`; - throw new Error(outputText); - } - - return { content: [{ type: "text", text: outputText }], details }; - } finally { - signal?.removeEventListener("abort", onAbort); } - }, - }; + + if (!details && (jsonOutputs.length > 0 || images.length > 0 || statusEvents.length > 0)) { + details = { + jsonOutputs: jsonOutputs.length > 0 ? jsonOutputs : undefined, + images: images.length > 0 ? images : undefined, + statusEvents: statusEvents.length > 0 ? statusEvents : undefined, + }; + } + + if (result.exitCode !== 0 && result.exitCode !== undefined) { + outputText += `\n\nCommand exited with code ${result.exitCode}`; + throw new Error(outputText); + } + + return { content: [{ type: "text", text: outputText }], details }; + } finally { + signal?.removeEventListener("abort", onAbort); + } + } } interface PythonRenderArgs { diff --git a/packages/coding-agent/src/core/tools/read.ts b/packages/coding-agent/src/core/tools/read.ts index e474cdbc5..d83b05f44 100644 --- a/packages/coding-agent/src/core/tools/read.ts +++ b/packages/coding-agent/src/core/tools/read.ts @@ -1,6 +1,6 @@ import { homedir } from "node:os"; import path from "node:path"; -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { ImageContent, TextContent } from "@oh-my-pi/pi-ai"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; @@ -15,7 +15,7 @@ import type { RenderResultOptions } from "../custom-tools/types"; import { renderPromptTemplate } from "../prompt-templates"; import type { ToolSession } from "../sdk"; import { ScopeSignal, untilAborted } from "../utils"; -import { createLsTool } from "./ls"; +import { LsTool } from "./ls"; import { resolveReadPath, resolveToCwd } from "./path-utils"; import { shortenPath, wrapBrackets } from "./render-utils"; import { @@ -425,74 +425,109 @@ export interface ReadToolDetails { redirectedTo?: "ls"; } -export function createReadTool(session: ToolSession): AgentTool { - const autoResizeImages = session.settings?.getImageAutoResize() ?? true; - const lsTool = createLsTool(session); - return { - name: "read", - label: "Read", - description: renderPromptTemplate(readDescription, { +type ReadParams = { path: string; offset?: number; limit?: number; lines?: boolean }; + +/** + * Read tool implementation. + * + * Reads files with support for images, documents (via markitdown), and text. + * Directories redirect to the ls tool. + */ +export class ReadTool implements AgentTool { + public readonly name = "read"; + public readonly label = "Read"; + public readonly description: string; + public readonly parameters = readSchema; + + private readonly session: ToolSession; + private readonly autoResizeImages: boolean; + private readonly lsTool: LsTool; + + constructor(session: ToolSession) { + this.session = session; + this.autoResizeImages = session.settings?.getImageAutoResize() ?? true; + this.lsTool = new LsTool(session); + this.description = renderPromptTemplate(readDescription, { DEFAULT_MAX_LINES: String(DEFAULT_MAX_LINES), - }), - parameters: readSchema, - execute: async ( - toolCallId: string, - { path: readPath, offset, limit, lines }: { path: string; offset?: number; limit?: number; lines?: boolean }, - signal?: AbortSignal, - ) => { - const absolutePath = resolveReadPath(readPath, session.cwd); + }); + } - return untilAborted(signal, async () => { - let isDirectory = false; - let fileSize = 0; - try { - const stat = await Bun.file(absolutePath).stat(); - fileSize = stat.size; - isDirectory = stat.isDirectory(); - } catch (error) { - if (isNotFoundError(error)) { - let message = `File not found: ${readPath}`; + public async execute( + toolCallId: string, + params: ReadParams, + signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const { path: readPath, offset, limit, lines } = params; + const absolutePath = resolveReadPath(readPath, this.session.cwd); - // Skip fuzzy matching for remote mounts (sshfs) to avoid hangs - if (!isRemoteMountPath(absolutePath)) { - const suggestions = await findReadPathSuggestions(readPath, session.cwd, signal); + return untilAborted(signal, async () => { + let isDirectory = false; + let fileSize = 0; + try { + const stat = await Bun.file(absolutePath).stat(); + fileSize = stat.size; + isDirectory = stat.isDirectory(); + } catch (error) { + if (isNotFoundError(error)) { + let message = `File not found: ${readPath}`; - if (suggestions?.suggestions.length) { - const scopeLabel = suggestions.scopeLabel ? ` in ${suggestions.scopeLabel}` : ""; - message += `\n\nClosest matches${scopeLabel}:\n${suggestions.suggestions.map((match) => `- ${match}`).join("\n")}`; - if (suggestions.truncated) { - message += `\n[Search truncated to first ${MAX_FUZZY_CANDIDATES} paths. Refine the path if the match isn't listed.]`; - } - } else if (suggestions?.error) { - message += `\n\nFuzzy match failed: ${suggestions.error}`; - } else if (suggestions?.scopeLabel) { - message += `\n\nNo similar paths found in ${suggestions.scopeLabel}.`; + // Skip fuzzy matching for remote mounts (sshfs) to avoid hangs + if (!isRemoteMountPath(absolutePath)) { + const suggestions = await findReadPathSuggestions(readPath, this.session.cwd, signal); + + if (suggestions?.suggestions.length) { + const scopeLabel = suggestions.scopeLabel ? ` in ${suggestions.scopeLabel}` : ""; + message += `\n\nClosest matches${scopeLabel}:\n${suggestions.suggestions.map((match) => `- ${match}`).join("\n")}`; + if (suggestions.truncated) { + message += `\n[Search truncated to first ${MAX_FUZZY_CANDIDATES} paths. Refine the path if the match isn't listed.]`; } + } else if (suggestions?.error) { + message += `\n\nFuzzy match failed: ${suggestions.error}`; + } else if (suggestions?.scopeLabel) { + message += `\n\nNo similar paths found in ${suggestions.scopeLabel}.`; } - - throw new Error(message); } - throw error; + + throw new Error(message); } + throw error; + } - if (isDirectory) { - const lsResult = await lsTool.execute(toolCallId, { path: readPath, limit }, signal); - return { - content: lsResult.content, - details: { redirectedTo: "ls", truncation: lsResult.details?.truncation }, - }; - } + if (isDirectory) { + const lsResult = await this.lsTool.execute(toolCallId, { path: readPath, limit }, signal); + return { + content: lsResult.content, + details: { redirectedTo: "ls", truncation: lsResult.details?.truncation }, + }; + } - const mimeType = await detectSupportedImageMimeTypeFromFile(absolutePath); - const ext = path.extname(absolutePath).toLowerCase(); + const mimeType = await detectSupportedImageMimeTypeFromFile(absolutePath); + const ext = path.extname(absolutePath).toLowerCase(); - // Read the file based on type - let content: (TextContent | ImageContent)[]; - let details: ReadToolDetails | undefined; + // Read the file based on type + let content: (TextContent | ImageContent)[]; + let details: ReadToolDetails | undefined; - if (mimeType) { - if (fileSize > MAX_IMAGE_SIZE) { - const sizeStr = formatSize(fileSize); + if (mimeType) { + if (fileSize > MAX_IMAGE_SIZE) { + const sizeStr = formatSize(fileSize); + const maxStr = formatSize(MAX_IMAGE_SIZE); + content = [ + { + type: "text", + text: `[Image file too large: ${sizeStr} exceeds ${maxStr} limit. Use an image viewer or resize the image.]`, + }, + ]; + } else { + // Read as image (binary) + const file = Bun.file(absolutePath); + const buffer = await file.arrayBuffer(); + + // Check actual buffer size after reading to prevent OOM during serialization + if (buffer.byteLength > MAX_IMAGE_SIZE) { + const sizeStr = formatSize(buffer.byteLength); const maxStr = formatSize(MAX_IMAGE_SIZE); content = [ { @@ -501,178 +536,164 @@ export function createReadTool(session: ToolSession): AgentTool MAX_IMAGE_SIZE) { - const sizeStr = formatSize(buffer.byteLength); - const maxStr = formatSize(MAX_IMAGE_SIZE); - content = [ - { - type: "text", - text: `[Image file too large: ${sizeStr} exceeds ${maxStr} limit. Use an image viewer or resize the image.]`, - }, - ]; - } else { - const base64 = Buffer.from(buffer).toString("base64"); + if (this.autoResizeImages) { + // Resize image if needed - catch errors from WASM + try { + const resized = await resizeImage({ type: "image", data: base64, mimeType }); + const dimensionNote = formatDimensionNote(resized); - if (autoResizeImages) { - // Resize image if needed - catch errors from WASM - try { - const resized = await resizeImage({ type: "image", data: base64, mimeType }); - const dimensionNote = formatDimensionNote(resized); - - let textNote = `Read image file [${resized.mimeType}]`; - if (dimensionNote) { - textNote += `\n${dimensionNote}`; - } - - content = [ - { type: "text", text: textNote }, - { type: "image", data: resized.data, mimeType: resized.mimeType }, - ]; - } catch { - // Fall back to original image on resize failure - content = [ - { type: "text", text: `Read image file [${mimeType}]` }, - { type: "image", data: base64, mimeType }, - ]; + let textNote = `Read image file [${resized.mimeType}]`; + if (dimensionNote) { + textNote += `\n${dimensionNote}`; } - } else { + + content = [ + { type: "text", text: textNote }, + { type: "image", data: resized.data, mimeType: resized.mimeType }, + ]; + } catch { + // Fall back to original image on resize failure content = [ { type: "text", text: `Read image file [${mimeType}]` }, { type: "image", data: base64, mimeType }, ]; } - } - } - } else if (CONVERTIBLE_EXTENSIONS.has(ext)) { - // Convert document via markitdown - const result = await convertWithMarkitdown(absolutePath, signal); - if (result.ok) { - // Apply truncation to converted content - const truncation = truncateHead(result.content); - let outputText = truncation.content; - - if (truncation.truncated) { - outputText += `\n\n[Document converted via markitdown. Output truncated to ${formatSize(DEFAULT_MAX_BYTES)}]`; - details = { truncation }; - } - - content = [{ type: "text", text: outputText }]; - } else { - // markitdown not available or failed - const errorMsg = - result.error === "markitdown not found" - ? `markitdown not installed. Install with: pip install markitdown` - : result.error || "conversion failed"; - content = [{ type: "text", text: `[Cannot read ${ext} file: ${errorMsg}]` }]; - } - } else { - // Read as text - const file = Bun.file(absolutePath); - const textContent = await file.text(); - const allLines = textContent.split("\n"); - const totalFileLines = allLines.length; - - // Apply offset if specified (1-indexed to 0-indexed) - const startLine = offset ? Math.max(0, offset - 1) : 0; - const startLineDisplay = startLine + 1; // For display (1-indexed) - - // Check if offset is out of bounds - if (startLine >= allLines.length) { - throw new Error(`Offset ${offset} is beyond end of file (${allLines.length} lines total)`); - } - - // If limit is specified by user, use it; otherwise we'll let truncateHead decide - let selectedContent: string; - let userLimitedLines: number | undefined; - if (limit !== undefined) { - const endLine = Math.min(startLine + limit, allLines.length); - selectedContent = allLines.slice(startLine, endLine).join("\n"); - userLimitedLines = endLine - startLine; - } else { - selectedContent = allLines.slice(startLine).join("\n"); - } - - // Apply truncation (respects both line and byte limits) - const truncation = truncateHead(selectedContent); - - // Add line numbers if requested (default: true) - const shouldAddLineNumbers = lines !== false; - const prependLineNumbers = (text: string, startNum: number): string => { - const lines = text.split("\n"); - const lastLineNum = startNum + lines.length - 1; - const padWidth = String(lastLineNum).length; - return lines - .map((line, i) => { - const lineNum = String(startNum + i).padStart(padWidth, " "); - return `${lineNum}\t${line}`; - }) - .join("\n"); - }; - - let outputText: string; - - if (truncation.firstLineExceedsLimit) { - const firstLine = allLines[startLine] ?? ""; - const firstLineBytes = Buffer.byteLength(firstLine, "utf-8"); - const snippet = truncateStringToBytesFromStart(firstLine, DEFAULT_MAX_BYTES); - const shownSize = formatSize(snippet.bytes); - - outputText = shouldAddLineNumbers ? prependLineNumbers(snippet.text, startLineDisplay) : snippet.text; - if (snippet.text.length > 0) { - outputText += `\n\n[Line ${startLineDisplay} is ${formatSize( - firstLineBytes, - )}, exceeds ${formatSize(DEFAULT_MAX_BYTES)} limit. Showing first ${shownSize} of the line.]`; } else { - outputText = `[Line ${startLineDisplay} is ${formatSize( - firstLineBytes, - )}, exceeds ${formatSize(DEFAULT_MAX_BYTES)} limit. Unable to display a valid UTF-8 snippet.]`; + content = [ + { type: "text", text: `Read image file [${mimeType}]` }, + { type: "image", data: base64, mimeType }, + ]; } + } + } + } else if (CONVERTIBLE_EXTENSIONS.has(ext)) { + // Convert document via markitdown + const result = await convertWithMarkitdown(absolutePath, signal); + if (result.ok) { + // Apply truncation to converted content + const truncation = truncateHead(result.content); + let outputText = truncation.content; + + if (truncation.truncated) { + outputText += `\n\n[Document converted via markitdown. Output truncated to ${formatSize(DEFAULT_MAX_BYTES)}]`; details = { truncation }; - } else if (truncation.truncated) { - // Truncation occurred - build actionable notice - const endLineDisplay = startLineDisplay + truncation.outputLines - 1; - const nextOffset = endLineDisplay + 1; - - outputText = shouldAddLineNumbers - ? prependLineNumbers(truncation.content, startLineDisplay) - : truncation.content; - - if (truncation.truncatedBy === "lines") { - outputText += `\n\n[Showing lines ${startLineDisplay}-${endLineDisplay} of ${totalFileLines}. Use offset=${nextOffset} to continue]`; - } else { - outputText += `\n\n[Showing lines ${startLineDisplay}-${endLineDisplay} of ${totalFileLines} (${formatSize( - DEFAULT_MAX_BYTES, - )} limit). Use offset=${nextOffset} to continue]`; - } - details = { truncation }; - } else if (userLimitedLines !== undefined && startLine + userLimitedLines < allLines.length) { - // User specified limit, there's more content, but no truncation - const remaining = allLines.length - (startLine + userLimitedLines); - const nextOffset = startLine + userLimitedLines + 1; - - outputText = shouldAddLineNumbers - ? prependLineNumbers(truncation.content, startLineDisplay) - : truncation.content; - outputText += `\n\n[${remaining} more lines in file. Use offset=${nextOffset} to continue]`; - } else { - // No truncation, no user limit exceeded - outputText = shouldAddLineNumbers - ? prependLineNumbers(truncation.content, startLineDisplay) - : truncation.content; } content = [{ type: "text", text: outputText }]; + } else if (result.error) { + // markitdown not available or failed + const errorMsg = + result.error === "markitdown not found" + ? `markitdown not installed. Install with: pip install markitdown` + : result.error || "conversion failed"; + content = [{ type: "text", text: `[Cannot read ${ext} file: ${errorMsg}]` }]; + } else { + content = [{ type: "text", text: `[Cannot read ${ext} file: conversion failed]` }]; + } + } else { + // Read as text + const file = Bun.file(absolutePath); + const textContent = await file.text(); + const allLines = textContent.split("\n"); + const totalFileLines = allLines.length; + + // Apply offset if specified (1-indexed to 0-indexed) + const startLine = offset ? Math.max(0, offset - 1) : 0; + const startLineDisplay = startLine + 1; // For display (1-indexed) + + // Check if offset is out of bounds + if (startLine >= allLines.length) { + throw new Error(`Offset ${offset} is beyond end of file (${allLines.length} lines total)`); } - return { content, details }; - }); - }, - }; + // If limit is specified by user, use it; otherwise we'll let truncateHead decide + let selectedContent: string; + let userLimitedLines: number | undefined; + if (limit !== undefined) { + const endLine = Math.min(startLine + limit, allLines.length); + selectedContent = allLines.slice(startLine, endLine).join("\n"); + userLimitedLines = endLine - startLine; + } else { + selectedContent = allLines.slice(startLine).join("\n"); + } + + // Apply truncation (respects both line and byte limits) + const truncation = truncateHead(selectedContent); + + // Add line numbers if requested (default: true) + const shouldAddLineNumbers = lines !== false; + const prependLineNumbers = (text: string, startNum: number): string => { + const textLines = text.split("\n"); + const lastLineNum = startNum + textLines.length - 1; + const padWidth = String(lastLineNum).length; + return textLines + .map((line, i) => { + const lineNum = String(startNum + i).padStart(padWidth, " "); + return `${lineNum}\t${line}`; + }) + .join("\n"); + }; + + let outputText: string; + + if (truncation.firstLineExceedsLimit) { + const firstLine = allLines[startLine] ?? ""; + const firstLineBytes = Buffer.byteLength(firstLine, "utf-8"); + const snippet = truncateStringToBytesFromStart(firstLine, DEFAULT_MAX_BYTES); + const shownSize = formatSize(snippet.bytes); + + outputText = shouldAddLineNumbers ? prependLineNumbers(snippet.text, startLineDisplay) : snippet.text; + if (snippet.text.length > 0) { + outputText += `\n\n[Line ${startLineDisplay} is ${formatSize( + firstLineBytes, + )}, exceeds ${formatSize(DEFAULT_MAX_BYTES)} limit. Showing first ${shownSize} of the line.]`; + } else { + outputText = `[Line ${startLineDisplay} is ${formatSize( + firstLineBytes, + )}, exceeds ${formatSize(DEFAULT_MAX_BYTES)} limit. Unable to display a valid UTF-8 snippet.]`; + } + details = { truncation }; + } else if (truncation.truncated) { + // Truncation occurred - build actionable notice + const endLineDisplay = startLineDisplay + truncation.outputLines - 1; + const nextOffset = endLineDisplay + 1; + + outputText = shouldAddLineNumbers + ? prependLineNumbers(truncation.content, startLineDisplay) + : truncation.content; + + if (truncation.truncatedBy === "lines") { + outputText += `\n\n[Showing lines ${startLineDisplay}-${endLineDisplay} of ${totalFileLines}. Use offset=${nextOffset} to continue]`; + } else { + outputText += `\n\n[Showing lines ${startLineDisplay}-${endLineDisplay} of ${totalFileLines} (${formatSize( + DEFAULT_MAX_BYTES, + )} limit). Use offset=${nextOffset} to continue]`; + } + details = { truncation }; + } else if (userLimitedLines !== undefined && startLine + userLimitedLines < allLines.length) { + // User specified limit, there's more content, but no truncation + const remaining = allLines.length - (startLine + userLimitedLines); + const nextOffset = startLine + userLimitedLines + 1; + + outputText = shouldAddLineNumbers + ? prependLineNumbers(truncation.content, startLineDisplay) + : truncation.content; + outputText += `\n\n[${remaining} more lines in file. Use offset=${nextOffset} to continue]`; + } else { + // No truncation, no user limit exceeded + outputText = shouldAddLineNumbers + ? prependLineNumbers(truncation.content, startLineDisplay) + : truncation.content; + } + + content = [{ type: "text", text: outputText }]; + } + + return { content, details }; + }); + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/render-utils.ts b/packages/coding-agent/src/core/tools/render-utils.ts index e5924040a..bb31203a6 100644 --- a/packages/coding-agent/src/core/tools/render-utils.ts +++ b/packages/coding-agent/src/core/tools/render-utils.ts @@ -231,67 +231,104 @@ export interface ToolUITitleOptions { bold?: boolean; } -export interface ToolUIKit { - theme: Theme; - title: (label: string, options?: ToolUITitleOptions) => string; - meta: (meta: string[]) => string; - count: (label: string, count: number) => string; - moreItems: (remaining: number, itemType: string) => string; - expandHint: (expanded: boolean, hasMore: boolean) => string; - scope: (scopePath?: string) => string; - truncationSuffix: (truncated: boolean) => string; - errorMessage: (message: string | undefined) => string; - emptyMessage: (message: string) => string; - badge: (label: string, color: ToolUIColor) => string; - statusIcon: (status: ToolUIStatus, spinnerFrame?: number) => string; - wrapBrackets: (text: string) => string; - truncate: (text: string, maxLen: number) => string; - previewLines: (text: string, maxLines: number, maxLineLen: number) => string[]; - formatBytes: (bytes: number) => string; - formatTokens: (tokens: number) => string; - formatDuration: (ms: number) => string; - formatAge: (ageSeconds: number | null | undefined) => string; - formatDiagnostics: ( - diag: { errored: boolean; summary: string; messages: string[] }, - expanded: boolean, - getLangIcon: (filePath: string) => string, - ) => string; - formatDiffStats: (added: number, removed: number, hunks: number) => string; -} - -export function createToolUIKit(theme: Theme): ToolUIKit { - return { - theme, - title: (label, options) => { - const content = options?.bold === false ? label : theme.bold(label); - return theme.fg("toolTitle", content); - }, - meta: (meta) => formatMeta(meta, theme), - count: (label, count) => formatCount(label, count), - moreItems: (remaining, itemType) => formatMoreItems(remaining, itemType, theme), - expandHint: (expanded, hasMore) => formatExpandHint(theme, expanded, hasMore), - scope: (scopePath) => formatScope(scopePath, theme), - truncationSuffix: (truncated) => formatTruncationSuffix(truncated, theme), - errorMessage: (message) => formatErrorMessage(message, theme), - emptyMessage: (message) => formatEmptyMessage(message, theme), - badge: (label, color) => formatBadge(label, color, theme), - statusIcon: (status, spinnerFrame) => formatStatusIcon(status, theme, spinnerFrame), - wrapBrackets: (text) => wrapBrackets(text, theme), - truncate: (text, maxLen) => truncate(text, maxLen, theme.format.ellipsis), - previewLines: (text, maxLines, maxLineLen) => getPreviewLines(text, maxLines, maxLineLen, theme.format.ellipsis), - formatBytes, - formatTokens, - formatDuration, - formatAge, - formatDiagnostics: (diag, expanded, getLangIcon) => formatDiagnostics(diag, expanded, theme, getLangIcon), - formatDiffStats: (added, removed, hunks) => formatDiffStats(added, removed, hunks, theme), - }; -} - // ============================================================================= // Diagnostic Formatting // ============================================================================= +export class ToolUIKit { + constructor(public theme: Theme) {} + + title(label: string, options?: ToolUITitleOptions): string { + const content = options?.bold === false ? label : this.theme.bold(label); + return this.theme.fg("toolTitle", content); + } + + meta(meta: string[]): string { + return formatMeta(meta, this.theme); + } + + count(label: string, count: number): string { + return formatCount(label, count); + } + + moreItems(remaining: number, itemType: string): string { + return formatMoreItems(remaining, itemType, this.theme); + } + + expandHint(expanded: boolean, hasMore: boolean): string { + return formatExpandHint(this.theme, expanded, hasMore); + } + + scope(scopePath?: string): string { + return formatScope(scopePath, this.theme); + } + + truncationSuffix(truncated: boolean): string { + return formatTruncationSuffix(truncated, this.theme); + } + + errorMessage(message: string | undefined): string { + return formatErrorMessage(message, this.theme); + } + + emptyMessage(message: string): string { + return formatEmptyMessage(message, this.theme); + } + + badge(label: string, color: ToolUIColor): string { + return formatBadge(label, color, this.theme); + } + + statusIcon(status: ToolUIStatus, spinnerFrame?: number): string { + return formatStatusIcon(status, this.theme, spinnerFrame); + } + + wrapBrackets(text: string): string { + return wrapBrackets(text, this.theme); + } + + truncate(text: string, maxLen: number): string { + return truncate(text, maxLen, this.theme.format.ellipsis); + } + + previewLines(text: string, maxLines: number, maxLineLen: number): string[] { + return getPreviewLines(text, maxLines, maxLineLen, this.theme.format.ellipsis); + } + + formatBytes(bytes: number): string { + return formatBytes(bytes); + } + + formatTokens(tokens: number): string { + return formatTokens(tokens); + } + + formatDuration(ms: number): string { + return formatDuration(ms); + } + + formatAge(ageSeconds: number | null | undefined): string { + return formatAge(ageSeconds); + } + + formatDiagnostics( + diag: { errored: boolean; summary: string; messages: string[] }, + expanded: boolean, + getLangIcon: (filePath: string) => string, + ): string { + return formatDiagnostics(diag, expanded, this.theme, getLangIcon); + } + + formatDiffStats(added: number, removed: number, hunks: number): string { + return formatDiffStats(added, removed, hunks, this.theme); + } +} + +/** @deprecated Use `new ToolUIKit(theme)` instead */ +export function createToolUIKit(theme: Theme): ToolUIKit { + return new ToolUIKit(theme); +} + interface ParsedDiagnostic { filePath: string; line: number; diff --git a/packages/coding-agent/src/core/tools/ssh.ts b/packages/coding-agent/src/core/tools/ssh.ts index 7c2801346..37fff3335 100644 --- a/packages/coding-agent/src/core/tools/ssh.ts +++ b/packages/coding-agent/src/core/tools/ssh.ts @@ -1,4 +1,4 @@ -import type { AgentTool, AgentToolContext } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; import { Type } from "@sinclair/typebox"; @@ -112,97 +112,116 @@ async function loadHosts(session: ToolSession): Promise<{ return { hostNames, hostsByName }; } -export async function createSshTool(session: ToolSession): Promise | null> { +interface SshToolParams { + host: string; + command: string; + cwd?: string; + timeout?: number; +} + +export class SshTool implements AgentTool { + public readonly name = "ssh"; + public readonly label = "SSH"; + public readonly description: string; + public readonly parameters = sshSchema; + + private readonly allowedHosts: Set; + private readonly hostsByName: Map; + private readonly hostNames: string[]; + + constructor(hostNames: string[], hostsByName: Map) { + this.hostNames = hostNames; + this.hostsByName = hostsByName; + this.allowedHosts = new Set(hostNames); + + const descriptionHosts = hostNames + .map((name) => hostsByName.get(name)) + .filter((host): host is SSHHost => host !== undefined); + + this.description = formatDescription(descriptionHosts); + } + + public async execute( + _toolCallId: string, + { host, command, cwd, timeout }: SshToolParams, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + _ctx?: AgentToolContext, + ): Promise> { + if (!this.allowedHosts.has(host)) { + throw new Error(`Unknown SSH host: ${host}. Available hosts: ${this.hostNames.join(", ")}`); + } + + const hostConfig = this.hostsByName.get(host); + if (!hostConfig) { + throw new Error(`SSH host not loaded: ${host}`); + } + + const hostInfo = await ensureHostInfo(hostConfig); + const remoteCommand = buildRemoteCommand(command, cwd, hostInfo); + let currentOutput = ""; + + const result = await executeSSH(hostConfig, remoteCommand, { + timeout: timeout ? timeout * 1000 : undefined, + signal, + compatEnabled: hostInfo.compatEnabled, + onChunk: (chunk) => { + currentOutput += chunk; + if (onUpdate) { + const truncation = truncateTail(currentOutput); + onUpdate({ + content: [{ type: "text", text: truncation.content || "" }], + details: { + truncation: truncation.truncated ? truncation : undefined, + }, + }); + } + }, + }); + + if (result.cancelled) { + throw new Error(result.output || "Command aborted"); + } + + const truncation = truncateTail(result.output); + let outputText = truncation.content || "(no output)"; + + let details: SSHToolDetails | undefined; + + if (truncation.truncated) { + details = { + truncation, + fullOutputPath: result.fullOutputPath, + }; + + const startLine = truncation.totalLines - truncation.outputLines + 1; + const endLine = truncation.totalLines; + + if (truncation.lastLinePartial) { + const lastLineSize = formatSize(Buffer.byteLength(result.output.split("\n").pop() || "", "utf-8")); + outputText += `\n\n[Showing last ${formatSize(truncation.outputBytes)} of line ${endLine} (line is ${lastLineSize}). Full output: ${result.fullOutputPath}]`; + } else if (truncation.truncatedBy === "lines") { + outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines}. Full output: ${result.fullOutputPath}]`; + } else { + outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines} (${formatSize(DEFAULT_MAX_BYTES)} limit). Full output: ${result.fullOutputPath}]`; + } + } + + if (result.exitCode !== 0 && result.exitCode !== undefined) { + outputText += `\n\nCommand exited with code ${result.exitCode}`; + throw new Error(outputText); + } + + return { content: [{ type: "text", text: outputText }], details: details ?? {} }; + } +} + +export async function loadSshTool(session: ToolSession): Promise { const { hostNames, hostsByName } = await loadHosts(session); if (hostNames.length === 0) { return null; } - - const allowedHosts = new Set(hostNames); - - const descriptionHosts = hostNames - .map((name) => hostsByName.get(name)) - .filter((host): host is SSHHost => host !== undefined); - - return { - name: "ssh", - label: "SSH", - description: formatDescription(descriptionHosts), - parameters: sshSchema, - execute: async ( - _toolCallId: string, - { host, command, cwd, timeout }: { host: string; command: string; cwd?: string; timeout?: number }, - signal?: AbortSignal, - onUpdate?, - _ctx?: AgentToolContext, - ) => { - if (!allowedHosts.has(host)) { - throw new Error(`Unknown SSH host: ${host}. Available hosts: ${hostNames.join(", ")}`); - } - - const hostConfig = hostsByName.get(host); - if (!hostConfig) { - throw new Error(`SSH host not loaded: ${host}`); - } - - const hostInfo = await ensureHostInfo(hostConfig); - const remoteCommand = buildRemoteCommand(command, cwd, hostInfo); - let currentOutput = ""; - - const result = await executeSSH(hostConfig, remoteCommand, { - timeout: timeout ? timeout * 1000 : undefined, - signal, - compatEnabled: hostInfo.compatEnabled, - onChunk: (chunk) => { - currentOutput += chunk; - if (onUpdate) { - const truncation = truncateTail(currentOutput); - onUpdate({ - content: [{ type: "text", text: truncation.content || "" }], - details: { - truncation: truncation.truncated ? truncation : undefined, - }, - }); - } - }, - }); - - if (result.cancelled) { - throw new Error(result.output || "Command aborted"); - } - - const truncation = truncateTail(result.output); - let outputText = truncation.content || "(no output)"; - - let details: SSHToolDetails | undefined; - - if (truncation.truncated) { - details = { - truncation, - fullOutputPath: result.fullOutputPath, - }; - - const startLine = truncation.totalLines - truncation.outputLines + 1; - const endLine = truncation.totalLines; - - if (truncation.lastLinePartial) { - const lastLineSize = formatSize(Buffer.byteLength(result.output.split("\n").pop() || "", "utf-8")); - outputText += `\n\n[Showing last ${formatSize(truncation.outputBytes)} of line ${endLine} (line is ${lastLineSize}). Full output: ${result.fullOutputPath}]`; - } else if (truncation.truncatedBy === "lines") { - outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines}. Full output: ${result.fullOutputPath}]`; - } else { - outputText += `\n\n[Showing lines ${startLine}-${endLine} of ${truncation.totalLines} (${formatSize(DEFAULT_MAX_BYTES)} limit). Full output: ${result.fullOutputPath}]`; - } - } - - if (result.exitCode !== 0 && result.exitCode !== undefined) { - outputText += `\n\nCommand exited with code ${result.exitCode}`; - throw new Error(outputText); - } - - return { content: [{ type: "text", text: outputText }], details }; - }, - }; + return new SshTool(hostNames, hostsByName); } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/task/executor.ts b/packages/coding-agent/src/core/tools/task/executor.ts index 83e85a483..bace27740 100644 --- a/packages/coding-agent/src/core/tools/task/executor.ts +++ b/packages/coding-agent/src/core/tools/task/executor.ts @@ -12,9 +12,9 @@ import type { MCPManager } from "../../mcp/manager"; import type { ModelRegistry } from "../../model-registry"; import { checkPythonKernelAvailability } from "../../python-kernel"; import type { ToolSession } from ".."; -import { createLspTool } from "../lsp/index"; +import { LspTool } from "../lsp/index"; import type { LspParams } from "../lsp/types"; -import { createPythonTool } from "../python"; +import { PythonTool } from "../python"; import { ensureArtifactsDir, getArtifactPaths } from "./artifacts"; import { resolveModelPattern } from "./model-resolver"; import { subprocessToolRegistry } from "./subprocess-tool-registry"; @@ -388,7 +388,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise(); const lspToolSession: ToolSession = { @@ -400,7 +400,7 @@ export async function runSubprocess(options: ExecutorOptions): Promise { }); } +// ═══════════════════════════════════════════════════════════════════════════ +// Tool Class +// ═══════════════════════════════════════════════════════════════════════════ + +type TaskParams = { + agent: string; + context?: string; + model?: string; + output?: unknown; + tasks: Array<{ id: string; task: string; description: string }>; +}; + /** - * Create the task tool configured for a specific session. + * Task tool - Delegate tasks to specialized agents. + * + * Requires async initialization to discover available agents. + * Use `TaskTool.create(session)` to instantiate. */ -export async function createTaskTool( - session: ToolSession, -): Promise> { - // Check for same-agent blocking (allows other agent types) - const blockedAgent = process.env.OMP_BLOCKED_AGENT; +export class TaskTool implements AgentTool { + public readonly name = "task"; + public readonly label = "Task"; + public readonly description: string; + public readonly parameters = taskSchema; + public readonly renderCall = renderCall; + public readonly renderResult = renderResult; - // Build description upfront - const description = await buildDescription(session.cwd); + private readonly session: ToolSession; + private readonly blockedAgent: string | undefined; - return { - name: "task", - label: "Task", - description, - parameters: taskSchema, - renderCall, - renderResult, - execute: async (_toolCallId, params, signal, onUpdate) => { - const startTime = Date.now(); - const { agents, projectAgentsDir } = await discoverAgents(session.cwd); - const { agent: agentName, context, model, output: outputSchema } = params; + private constructor(session: ToolSession, description: string) { + this.session = session; + this.description = description; + this.blockedAgent = process.env.OMP_BLOCKED_AGENT; + } - const isDefaultModelAlias = (value: string | undefined): boolean => { - if (!value) return true; - const normalized = value.trim().toLowerCase(); - return normalized === "default" || normalized === "pi/default" || normalized === "omp/default"; + /** + * Create a TaskTool instance with async agent discovery. + */ + public static async create(session: ToolSession): Promise { + const description = await buildDescription(session.cwd); + return new TaskTool(session, description); + } + + public async execute( + _toolCallId: string, + params: TaskParams, + signal?: AbortSignal, + onUpdate?: AgentToolUpdateCallback, + ): Promise> { + const startTime = Date.now(); + const { agents, projectAgentsDir } = await discoverAgents(this.session.cwd); + const { agent: agentName, context, model, output: outputSchema } = params; + + const isDefaultModelAlias = (value: string | undefined): boolean => { + if (!value) return true; + const normalized = value.trim().toLowerCase(); + return normalized === "default" || normalized === "pi/default" || normalized === "omp/default"; + }; + + // Validate agent exists + const agent = getAgent(agents, agentName); + if (!agent) { + const available = agents.map((a) => a.name).join(", ") || "none"; + return { + content: [ + { + type: "text", + text: `Unknown agent "${agentName}". Available: ${available}`, + }, + ], + details: { + projectAgentsDir, + results: [], + totalDurationMs: 0, + }, }; + } - // Validate agent exists - const agent = getAgent(agents, agentName); - if (!agent) { - const available = agents.map((a) => a.name).join(", ") || "none"; + const shouldInheritSessionModel = model === undefined && isDefaultModelAlias(agent.model); + const sessionModel = shouldInheritSessionModel ? this.session.getActiveModelString?.() : undefined; + const modelOverride = model ?? sessionModel ?? this.session.getModelString?.(); + const thinkingLevelOverride = agent.thinkingLevel; + + // Output schema priority: agent frontmatter > params > inherited from parent session + const schemaOverridden = outputSchema !== undefined && agent.output !== undefined; + const effectiveOutputSchema = agent.output ?? outputSchema ?? this.session.outputSchema; + + // Handle empty or missing tasks + if (!params.tasks || params.tasks.length === 0) { + return { + content: [ + { + type: "text", + text: `No tasks provided. Use: { agent, context, tasks: [{id, task, description}, ...] }`, + }, + ], + details: { + projectAgentsDir, + results: [], + totalDurationMs: 0, + }, + }; + } + + // Validate task count + if (params.tasks.length > MAX_PARALLEL_TASKS) { + return { + content: [ + { + type: "text", + text: `Too many tasks (${params.tasks.length}). Max is ${MAX_PARALLEL_TASKS}.`, + }, + ], + details: { + projectAgentsDir, + results: [], + totalDurationMs: 0, + }, + }; + } + + const tasks = params.tasks; + const missingTaskIndexes: number[] = []; + const idIndexes = new Map(); + + for (let i = 0; i < tasks.length; i++) { + const id = tasks[i]?.id; + if (typeof id !== "string" || id.trim() === "") { + missingTaskIndexes.push(i); + continue; + } + const normalizedId = id.toLowerCase(); + const indexes = idIndexes.get(normalizedId); + if (indexes) { + indexes.push(i); + } else { + idIndexes.set(normalizedId, [i]); + } + } + + const duplicateIds: Array<{ id: string; indexes: number[] }> = []; + for (const [normalizedId, indexes] of idIndexes.entries()) { + if (indexes.length > 1) { + duplicateIds.push({ + id: tasks[indexes[0]]?.id ?? normalizedId, + indexes, + }); + } + } + + if (missingTaskIndexes.length > 0 || duplicateIds.length > 0) { + const problems: string[] = []; + if (missingTaskIndexes.length > 0) { + problems.push(`Missing task ids at indexes: ${missingTaskIndexes.join(", ")}`); + } + if (duplicateIds.length > 0) { + const details = duplicateIds.map((entry) => `${entry.id} (indexes ${entry.indexes.join(", ")})`).join("; "); + problems.push(`Duplicate task ids detected (case-insensitive): ${details}`); + } + return { + content: [{ type: "text", text: `Invalid tasks: ${problems.join(". ")}` }], + details: { + projectAgentsDir, + results: [], + totalDurationMs: 0, + }, + }; + } + + // Derive artifacts directory + const sessionFile = this.session.getSessionFile(); + const artifactsDir = sessionFile ? getArtifactsDir(sessionFile) : null; + const tempArtifactsDir = artifactsDir ? null : createTempArtifactsDir(); + const effectiveArtifactsDir = artifactsDir || tempArtifactsDir!; + + // Initialize progress tracking + const progressMap = new Map(); + + // Update callback + const emitProgress = () => { + const progress = Array.from(progressMap.values()).sort((a, b) => a.index - b.index); + onUpdate?.({ + content: [{ type: "text", text: `Running ${params.tasks.length} agents...` }], + details: { + projectAgentsDir, + results: [], + totalDurationMs: Date.now() - startTime, + progress, + }, + }); + }; + + try { + // Check self-recursion prevention + if (this.blockedAgent && agentName === this.blockedAgent) { return { content: [ { type: "text", - text: `Unknown agent "${agentName}". Available: ${available}`, + text: `Cannot spawn ${this.blockedAgent} agent from within itself (recursion prevention). Use a different agent type.`, }, ], - details: { - projectAgentsDir, - results: [], - totalDurationMs: 0, - }, - }; - } - - const shouldInheritSessionModel = model === undefined && isDefaultModelAlias(agent.model); - const sessionModel = shouldInheritSessionModel ? session.getActiveModelString?.() : undefined; - const modelOverride = model ?? sessionModel ?? session.getModelString?.(); - const thinkingLevelOverride = agent.thinkingLevel; - - // Output schema priority: agent frontmatter > params > inherited from parent session - const schemaOverridden = outputSchema !== undefined && agent.output !== undefined; - const effectiveOutputSchema = agent.output ?? outputSchema ?? session.outputSchema; - - // Handle empty or missing tasks - if (!params.tasks || params.tasks.length === 0) { - return { - content: [ - { - type: "text", - text: `No tasks provided. Use: { agent, context, tasks: [{id, task, description}, ...] }`, - }, - ], - details: { - projectAgentsDir, - results: [], - totalDurationMs: 0, - }, - }; - } - - // Validate task count - if (params.tasks.length > MAX_PARALLEL_TASKS) { - return { - content: [ - { - type: "text", - text: `Too many tasks (${params.tasks.length}). Max is ${MAX_PARALLEL_TASKS}.`, - }, - ], - details: { - projectAgentsDir, - results: [], - totalDurationMs: 0, - }, - }; - } - - const tasks = params.tasks; - const missingTaskIndexes: number[] = []; - const idIndexes = new Map(); - - for (let i = 0; i < tasks.length; i++) { - const id = tasks[i]?.id; - if (typeof id !== "string" || id.trim() === "") { - missingTaskIndexes.push(i); - continue; - } - const normalizedId = id.toLowerCase(); - const indexes = idIndexes.get(normalizedId); - if (indexes) { - indexes.push(i); - } else { - idIndexes.set(normalizedId, [i]); - } - } - - const duplicateIds: Array<{ id: string; indexes: number[] }> = []; - for (const [normalizedId, indexes] of idIndexes.entries()) { - if (indexes.length > 1) { - duplicateIds.push({ - id: tasks[indexes[0]]?.id ?? normalizedId, - indexes, - }); - } - } - - if (missingTaskIndexes.length > 0 || duplicateIds.length > 0) { - const problems: string[] = []; - if (missingTaskIndexes.length > 0) { - problems.push(`Missing task ids at indexes: ${missingTaskIndexes.join(", ")}`); - } - if (duplicateIds.length > 0) { - const details = duplicateIds - .map((entry) => `${entry.id} (indexes ${entry.indexes.join(", ")})`) - .join("; "); - problems.push(`Duplicate task ids detected (case-insensitive): ${details}`); - } - return { - content: [{ type: "text", text: `Invalid tasks: ${problems.join(". ")}` }], - details: { - projectAgentsDir, - results: [], - totalDurationMs: 0, - }, - }; - } - - // Derive artifacts directory - const sessionFile = session.getSessionFile(); - const artifactsDir = sessionFile ? getArtifactsDir(sessionFile) : null; - const tempArtifactsDir = artifactsDir ? null : createTempArtifactsDir(); - const effectiveArtifactsDir = artifactsDir || tempArtifactsDir!; - - // Initialize progress tracking - const progressMap = new Map(); - - // Update callback - const emitProgress = () => { - const progress = Array.from(progressMap.values()).sort((a, b) => a.index - b.index); - onUpdate?.({ - content: [{ type: "text", text: `Running ${params.tasks.length} agents...` }], details: { projectAgentsDir, results: [], totalDurationMs: Date.now() - startTime, - progress, }, - }); + }; + } + + // Check spawn restrictions from parent + const parentSpawns = this.session.getSessionSpawns() ?? "*"; + const allowedSpawns = parentSpawns.split(",").map((s) => s.trim()); + const isSpawnAllowed = (): boolean => { + if (parentSpawns === "") return false; // Empty = deny all + if (parentSpawns === "*") return true; // Wildcard = allow all + return allowedSpawns.includes(agentName); }; - try { - // Check self-recursion prevention - if (blockedAgent && agentName === blockedAgent) { - return { - content: [ - { - type: "text", - text: `Cannot spawn ${blockedAgent} agent from within itself (recursion prevention). Use a different agent type.`, - }, - ], - details: { - projectAgentsDir, - results: [], - totalDurationMs: Date.now() - startTime, - }, - }; - } - - // Check spawn restrictions from parent - const parentSpawns = session.getSessionSpawns() ?? "*"; - const allowedSpawns = parentSpawns.split(",").map((s) => s.trim()); - const isSpawnAllowed = (): boolean => { - if (parentSpawns === "") return false; // Empty = deny all - if (parentSpawns === "*") return true; // Wildcard = allow all - return allowedSpawns.includes(agentName); - }; - - if (!isSpawnAllowed()) { - const allowed = parentSpawns === "" ? "none (spawns disabled for this agent)" : parentSpawns; - return { - content: [{ type: "text", text: `Cannot spawn '${agentName}'. Allowed: ${allowed}` }], - details: { - projectAgentsDir, - results: [], - totalDurationMs: Date.now() - startTime, - }, - }; - } - - // Build full prompts with context prepended - const tasksWithContext = tasks.map((t) => ({ - task: context ? `${context}\n\n${t.task}` : t.task, - description: t.description, - taskId: t.id, - })); - - // Initialize progress for all tasks - for (let i = 0; i < tasksWithContext.length; i++) { - const t = tasksWithContext[i]; - progressMap.set(i, { - index: i, - taskId: t.taskId, - agent: agentName, - agentSource: agent.source, - status: "pending", - task: t.task, - recentTools: [], - recentOutput: [], - toolCount: 0, - tokens: 0, - durationMs: 0, - modelOverride, - description: t.description, - }); - } - emitProgress(); - - // Execute in parallel with concurrency limit - const { results: partialResults, aborted } = await mapWithConcurrencyLimit( - tasksWithContext, - MAX_CONCURRENCY, - async (task, index) => { - return runSubprocess({ - cwd: session.cwd, - agent, - task: task.task, - description: task.description, - index, - taskId: task.taskId, - context: undefined, // Already prepended above - modelOverride, - thinkingLevel: thinkingLevelOverride, - outputSchema: effectiveOutputSchema, - sessionFile, - persistArtifacts: !!artifactsDir, - artifactsDir: effectiveArtifactsDir, - enableLsp: false, - signal, - eventBus: undefined, - onProgress: (progress) => { - progressMap.set(index, structuredClone(progress)); - emitProgress(); - }, - authStorage: session.authStorage, - modelRegistry: session.modelRegistry, - settingsManager: session.settingsManager, - mcpManager: session.mcpManager, - }); + if (!isSpawnAllowed()) { + const allowed = parentSpawns === "" ? "none (spawns disabled for this agent)" : parentSpawns; + return { + content: [{ type: "text", text: `Cannot spawn '${agentName}'. Allowed: ${allowed}` }], + details: { + projectAgentsDir, + results: [], + totalDurationMs: Date.now() - startTime, }, - signal, - ); + }; + } - // Fill in skipped tasks (undefined entries from abort) with placeholder results - const results: SingleResult[] = partialResults.map((result, index) => { - if (result !== undefined) return result; - const task = tasksWithContext[index]; - return { - index, - taskId: task.taskId, - agent: agentName, - agentSource: agent.source, + // Build full prompts with context prepended + const tasksWithContext = tasks.map((t) => ({ + task: context ? `${context}\n\n${t.task}` : t.task, + description: t.description, + taskId: t.id, + })); + + // Initialize progress for all tasks + for (let i = 0; i < tasksWithContext.length; i++) { + const t = tasksWithContext[i]; + progressMap.set(i, { + index: i, + taskId: t.taskId, + agent: agentName, + agentSource: agent.source, + status: "pending", + task: t.task, + recentTools: [], + recentOutput: [], + toolCount: 0, + tokens: 0, + durationMs: 0, + modelOverride, + description: t.description, + }); + } + emitProgress(); + + // Execute in parallel with concurrency limit + const { results: partialResults, aborted } = await mapWithConcurrencyLimit( + tasksWithContext, + MAX_CONCURRENCY, + async (task, index) => { + return runSubprocess({ + cwd: this.session.cwd, + agent, task: task.task, description: task.description, - exitCode: 1, - output: "", - stderr: "Skipped (cancelled before start)", - truncated: false, - durationMs: 0, - tokens: 0, + index, + taskId: task.taskId, + context: undefined, // Already prepended above modelOverride, - error: "Skipped", - aborted: true, - }; - }); - - // Aggregate usage from executor results (already accumulated incrementally) - const aggregatedUsage = createUsageTotals(); - let hasAggregatedUsage = false; - for (const result of results) { - if (result.usage) { - addUsageTotals(aggregatedUsage, result.usage); - hasAggregatedUsage = true; - } - } - - // Collect output paths (artifacts already written by executor in real-time) - const outputPaths: string[] = []; - for (const result of results) { - if (result.artifactPaths) { - outputPaths.push(result.artifactPaths.outputPath); - } - } - - // Build final output - match plugin format - const successCount = results.filter((r) => r.exitCode === 0).length; - const cancelledCount = results.filter((r) => r.aborted).length; - const totalDuration = Date.now() - startTime; - - const summaries = results.map((r) => { - const status = r.aborted ? "cancelled" : r.exitCode === 0 ? "completed" : `failed (exit ${r.exitCode})`; - const output = r.output.trim() || r.stderr.trim() || "(no output)"; - const preview = output.split("\n").slice(0, 5).join("\n"); - const meta = r.outputMeta - ? ` [${r.outputMeta.lineCount} lines, ${formatBytes(r.outputMeta.charCount)}]` - : ""; - return `[${r.agent}] ${status}${meta} ${r.taskId}\n${preview}`; - }); - - const outputIds = results.filter((r) => !r.aborted || r.output.trim()).map((r) => r.taskId); - const outputHint = - outputIds.length > 0 ? `\n\nUse output tool for full logs: output ids ${outputIds.join(", ")}` : ""; - const schemaNote = schemaOverridden - ? `\n\nNote: Agent '${agentName}' has a fixed output schema; your 'output' parameter was ignored.\nRequired schema: ${JSON.stringify(agent.output)}` - : ""; - const cancelledNote = aborted && cancelledCount > 0 ? ` (${cancelledCount} cancelled)` : ""; - const summary = `${successCount}/${results.length} succeeded${cancelledNote} [${formatDuration( - totalDuration, - )}]\n\n${summaries.join("\n\n---\n\n")}${outputHint}${schemaNote}`; - - // Cleanup temp directory if used - if (tempArtifactsDir) { - await cleanupTempDir(tempArtifactsDir); - } + thinkingLevel: thinkingLevelOverride, + outputSchema: effectiveOutputSchema, + sessionFile, + persistArtifacts: !!artifactsDir, + artifactsDir: effectiveArtifactsDir, + enableLsp: false, + signal, + eventBus: undefined, + onProgress: (progress) => { + progressMap.set(index, structuredClone(progress)); + emitProgress(); + }, + authStorage: this.session.authStorage, + modelRegistry: this.session.modelRegistry, + settingsManager: this.session.settingsManager, + mcpManager: this.session.mcpManager, + }); + }, + signal, + ); + // Fill in skipped tasks (undefined entries from abort) with placeholder results + const results: SingleResult[] = partialResults.map((result, index) => { + if (result !== undefined) return result; + const task = tasksWithContext[index]; return { - content: [{ type: "text", text: summary }], - details: { - projectAgentsDir, - results: results, - totalDurationMs: totalDuration, - usage: hasAggregatedUsage ? aggregatedUsage : undefined, - outputPaths, - }, + index, + taskId: task.taskId, + agent: agentName, + agentSource: agent.source, + task: task.task, + description: task.description, + exitCode: 1, + output: "", + stderr: "Skipped (cancelled before start)", + truncated: false, + durationMs: 0, + tokens: 0, + modelOverride, + error: "Skipped", + aborted: true, }; - } catch (err) { - // Cleanup temp directory on error - if (tempArtifactsDir) { - await cleanupTempDir(tempArtifactsDir); - } + }); - return { - content: [{ type: "text", text: `Task execution failed: ${err}` }], - details: { - projectAgentsDir, - results: [], - totalDurationMs: Date.now() - startTime, - }, - }; + // Aggregate usage from executor results (already accumulated incrementally) + const aggregatedUsage = createUsageTotals(); + let hasAggregatedUsage = false; + for (const result of results) { + if (result.usage) { + addUsageTotals(aggregatedUsage, result.usage); + hasAggregatedUsage = true; + } } - }, - }; -} -// Default task tool - returns a placeholder tool -// Real implementations should use createTaskTool(session) to initialize the tool -export const taskTool: AgentTool = { - name: "task", - label: "Task", - description: "Launch a new agent to handle complex, multi-step tasks autonomously.", - parameters: taskSchema, - execute: async () => ({ - content: [{ type: "text", text: "Task tool not properly initialized. Use createTaskTool(session) instead." }], - details: { - projectAgentsDir: null, - results: [], - totalDurationMs: 0, - }, - }), -}; + // Collect output paths (artifacts already written by executor in real-time) + const outputPaths: string[] = []; + for (const result of results) { + if (result.artifactPaths) { + outputPaths.push(result.artifactPaths.outputPath); + } + } + + // Build final output - match plugin format + const successCount = results.filter((r) => r.exitCode === 0).length; + const cancelledCount = results.filter((r) => r.aborted).length; + const totalDuration = Date.now() - startTime; + + const summaries = results.map((r) => { + const status = r.aborted ? "cancelled" : r.exitCode === 0 ? "completed" : `failed (exit ${r.exitCode})`; + const output = r.output.trim() || r.stderr.trim() || "(no output)"; + const preview = output.split("\n").slice(0, 5).join("\n"); + const meta = r.outputMeta + ? ` [${r.outputMeta.lineCount} lines, ${formatBytes(r.outputMeta.charCount)}]` + : ""; + return `[${r.agent}] ${status}${meta} ${r.taskId}\n${preview}`; + }); + + const outputIds = results.filter((r) => !r.aborted || r.output.trim()).map((r) => r.taskId); + const outputHint = + outputIds.length > 0 ? `\n\nUse output tool for full logs: output ids ${outputIds.join(", ")}` : ""; + const schemaNote = schemaOverridden + ? `\n\nNote: Agent '${agentName}' has a fixed output schema; your 'output' parameter was ignored.\nRequired schema: ${JSON.stringify(agent.output)}` + : ""; + const cancelledNote = aborted && cancelledCount > 0 ? ` (${cancelledCount} cancelled)` : ""; + const summary = `${successCount}/${results.length} succeeded${cancelledNote} [${formatDuration( + totalDuration, + )}]\n\n${summaries.join("\n\n---\n\n")}${outputHint}${schemaNote}`; + + // Cleanup temp directory if used + if (tempArtifactsDir) { + await cleanupTempDir(tempArtifactsDir); + } + + return { + content: [{ type: "text", text: summary }], + details: { + projectAgentsDir, + results: results, + totalDurationMs: totalDuration, + usage: hasAggregatedUsage ? aggregatedUsage : undefined, + outputPaths, + }, + }; + } catch (err) { + // Cleanup temp directory on error + if (tempArtifactsDir) { + await cleanupTempDir(tempArtifactsDir); + } + + return { + content: [{ type: "text", text: `Task execution failed: ${err}` }], + details: { + projectAgentsDir, + results: [], + totalDurationMs: Date.now() - startTime, + }, + }; + } + } +} diff --git a/packages/coding-agent/src/core/tools/todo-write.ts b/packages/coding-agent/src/core/tools/todo-write.ts index f4d34f58b..60d8da95f 100644 --- a/packages/coding-agent/src/core/tools/todo-write.ts +++ b/packages/coding-agent/src/core/tools/todo-write.ts @@ -1,7 +1,7 @@ import { randomUUID } from "node:crypto"; import { mkdirSync } from "node:fs"; import path from "node:path"; -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import { StringEnum } from "@oh-my-pi/pi-ai"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; @@ -49,6 +49,8 @@ export interface TodoWriteToolDetails { const TODO_FILE_NAME = "todos.json"; +type TodoWriteParams = { todos: Array<{ id?: string; content?: string; activeForm?: string; status?: string }> }; + function normalizeTodoStatus(status?: string): TodoStatus { switch (status) { case "in_progress": @@ -137,6 +139,14 @@ async function saveTodoFile(filePath: string, data: TodoFile): Promise { await Bun.write(filePath, JSON.stringify(data, null, 2)); } +function formatTodoSummary(todos: TodoItem[]): string { + if (todos.length === 0) return "Todo list cleared."; + const completed = todos.filter((t) => t.status === "completed").length; + const inProgress = todos.filter((t) => t.status === "in_progress").length; + const pending = todos.filter((t) => t.status === "pending").length; + return `Saved ${todos.length} todos (${pending} pending, ${inProgress} in progress, ${completed} completed).`; +} + function formatTodoLine(item: TodoItem, uiTheme: Theme, prefix: string): string { const checkbox = uiTheme.checkbox; const displayText = @@ -151,71 +161,76 @@ function formatTodoLine(item: TodoItem, uiTheme: Theme, prefix: string): string } } -function formatTodoSummary(todos: TodoItem[]): string { - if (todos.length === 0) return "Todo list cleared."; - const completed = todos.filter((t) => t.status === "completed").length; - const inProgress = todos.filter((t) => t.status === "in_progress").length; - const pending = todos.filter((t) => t.status === "pending").length; - return `Saved ${todos.length} todos (${pending} pending, ${inProgress} in progress, ${completed} completed).`; -} +// ============================================================================= +// Tool Class +// ============================================================================= -export function createTodoWriteTool(session: ToolSession): AgentTool { - return { - name: "todo_write", - label: "Todo Write", - description: renderPromptTemplate(todoWriteDescription), - parameters: todoWriteSchema, - execute: async ( - _toolCallId: string, - params: { todos: Array<{ id?: string; content?: string; activeForm?: string; status?: string }> }, - ) => { - const todos = normalizeTodos(params.todos ?? []); - const validation = validateSequentialTodos(todos); - if (!validation.valid) { - throw new Error(validation.error ?? "Todos must be completed sequentially."); - } - const updatedAt = Date.now(); +export class TodoWriteTool implements AgentTool { + public readonly name = "todo_write"; + public readonly label = "Todo Write"; + public readonly description: string; + public readonly parameters = todoWriteSchema; - const sessionFile = session.getSessionFile(); - if (!sessionFile) { - return { - content: [{ type: "text", text: formatTodoSummary(todos) }], - details: { todos, updatedAt, storage: "memory" }, - }; - } + private readonly session: ToolSession; - const artifactsDir = getArtifactsDir(sessionFile); - if (!artifactsDir) { - return { - content: [{ type: "text", text: formatTodoSummary(todos) }], - details: { todos, updatedAt, storage: "memory" }, - }; - } + constructor(session: ToolSession) { + this.session = session; + this.description = renderPromptTemplate(todoWriteDescription); + } - ensureArtifactsDir(artifactsDir); - const todoPath = path.join(artifactsDir, TODO_FILE_NAME); - const existing = await loadTodoFile(todoPath); - const storedTodos = existing?.todos ?? []; - const merged = todos.length > 0 ? todos : []; - const fileData: TodoFile = { updatedAt, todos: merged }; - - try { - mkdirSync(artifactsDir, { recursive: true }); - await saveTodoFile(todoPath, fileData); - } catch (error) { - logger.error("Failed to write todo file", { path: todoPath, error: String(error) }); - return { - content: [{ type: "text", text: "Failed to save todos." }], - details: { todos: storedTodos, updatedAt, storage: "session" }, - }; - } + public async execute( + _toolCallId: string, + params: TodoWriteParams, + _signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const todos = normalizeTodos(params.todos ?? []); + const validation = validateSequentialTodos(todos); + if (!validation.valid) { + throw new Error(validation.error ?? "Todos must be completed sequentially."); + } + const updatedAt = Date.now(); + const sessionFile = this.session.getSessionFile(); + if (!sessionFile) { return { - content: [{ type: "text", text: formatTodoSummary(merged) }], - details: { todos: merged, updatedAt, storage: "session" }, + content: [{ type: "text", text: formatTodoSummary(todos) }], + details: { todos, updatedAt, storage: "memory" }, }; - }, - }; + } + + const artifactsDir = getArtifactsDir(sessionFile); + if (!artifactsDir) { + return { + content: [{ type: "text", text: formatTodoSummary(todos) }], + details: { todos, updatedAt, storage: "memory" }, + }; + } + + ensureArtifactsDir(artifactsDir); + const todoPath = path.join(artifactsDir, TODO_FILE_NAME); + const existing = await loadTodoFile(todoPath); + const storedTodos = existing?.todos ?? []; + const merged = todos.length > 0 ? todos : []; + const fileData: TodoFile = { updatedAt, todos: merged }; + + try { + mkdirSync(artifactsDir, { recursive: true }); + await saveTodoFile(todoPath, fileData); + } catch (error) { + logger.error("Failed to write todo file", { path: todoPath, error: String(error) }); + return { + content: [{ type: "text", text: "Failed to save todos." }], + details: { todos: storedTodos, updatedAt, storage: "session" }, + }; + } + + return { + content: [{ type: "text", text: formatTodoSummary(merged) }], + details: { todos: merged, updatedAt, storage: "session" }, + }; + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/web-fetch.ts b/packages/coding-agent/src/core/tools/web-fetch.ts index 0852f89bc..47f6520e9 100644 --- a/packages/coding-agent/src/core/tools/web-fetch.ts +++ b/packages/coding-agent/src/core/tools/web-fetch.ts @@ -1,9 +1,9 @@ import { tmpdir } from "node:os"; import * as path from "node:path"; -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; -import { Type } from "@sinclair/typebox"; +import { type Static, Type } from "@sinclair/typebox"; import { nanoid } from "nanoid"; import { parse as parseHtml } from "node-html-parser"; import { type Theme, theme } from "../../modes/interactive/theme/theme"; @@ -848,55 +848,62 @@ export interface WebFetchToolDetails { notes: string[]; } -export function createWebFetchTool(_session: ToolSession): AgentTool { - return { - name: "web_fetch", - label: "Web Fetch", - description: renderPromptTemplate(webFetchDescription), - parameters: webFetchSchema, - execute: async ( - _toolCallId: string, - { url, timeout = DEFAULT_TIMEOUT, raw = false }: { url: string; timeout?: number; raw?: boolean }, - signal?: AbortSignal, - ) => { - if (signal?.aborted) { - throw new Error("Operation aborted"); - } +export class WebFetchTool implements AgentTool { + public readonly name = "web_fetch"; + public readonly label = "Web Fetch"; + public readonly description: string; + public readonly parameters = webFetchSchema; - // Clamp timeout - const effectiveTimeout = Math.min(Math.max(timeout, 1), 120); + constructor(_session: ToolSession) { + this.description = renderPromptTemplate(webFetchDescription); + } - const result = await renderUrl(url, effectiveTimeout, raw, signal); + public async execute( + _toolCallId: string, + params: Static, + signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + const { url, timeout = DEFAULT_TIMEOUT, raw = false } = params; - // Format output - let output = ""; - output += `URL: ${result.finalUrl}\n`; - output += `Content-Type: ${result.contentType}\n`; - output += `Method: ${result.method}\n`; - if (result.truncated) { - output += `Warning: Output was truncated\n`; - } - if (result.notes.length > 0) { - output += `Notes: ${result.notes.join("; ")}\n`; - } - output += `\n---\n\n`; - output += result.content; + if (signal?.aborted) { + throw new Error("Operation aborted"); + } - const details: WebFetchToolDetails = { - url: result.url, - finalUrl: result.finalUrl, - contentType: result.contentType, - method: result.method, - truncated: result.truncated, - notes: result.notes, - }; + // Clamp timeout + const effectiveTimeout = Math.min(Math.max(timeout, 1), 120); - return { - content: [{ type: "text", text: output }], - details, - }; - }, - }; + const result = await renderUrl(url, effectiveTimeout, raw, signal); + + // Format output + let output = ""; + output += `URL: ${result.finalUrl}\n`; + output += `Content-Type: ${result.contentType}\n`; + output += `Method: ${result.method}\n`; + if (result.truncated) { + output += `Warning: Output was truncated\n`; + } + if (result.notes.length > 0) { + output += `Notes: ${result.notes.join("; ")}\n`; + } + output += `\n---\n\n`; + output += result.content; + + const details: WebFetchToolDetails = { + url: result.url, + finalUrl: result.finalUrl, + contentType: result.contentType, + method: result.method, + truncated: result.truncated, + notes: result.notes, + }; + + return { + content: [{ type: "text", text: output }], + details, + }; + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/tools/web-search/index.ts b/packages/coding-agent/src/core/tools/web-search/index.ts index 7a1905d14..77e3e55fb 100644 --- a/packages/coding-agent/src/core/tools/web-search/index.ts +++ b/packages/coding-agent/src/core/tools/web-search/index.ts @@ -12,7 +12,7 @@ * - web_search_company: Comprehensive company research */ -import type { AgentTool } from "@oh-my-pi/pi-agent-core"; +import type { AgentTool, AgentToolContext, AgentToolResult, AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import { StringEnum } from "@oh-my-pi/pi-ai"; import { Type } from "@sinclair/typebox"; import type { Theme } from "../../../modes/interactive/theme/theme"; @@ -330,16 +330,32 @@ async function executeWebSearch( }; } -/** Web search tool as AgentTool (for allTools export) */ -export const webSearchTool: AgentTool = { - name: "web_search", - label: "Web Search", - description: renderPromptTemplate(webSearchDescription), - parameters: webSearchSchema, - execute: async (toolCallId, params) => { - return executeWebSearch(toolCallId, params as WebSearchParams); - }, -}; +/** + * Web search tool implementation. + * + * Supports Anthropic, Perplexity, and Exa providers with automatic fallback. + * Session is accepted for interface consistency but not used. + */ +export class WebSearchTool implements AgentTool { + public readonly name = "web_search"; + public readonly label = "Web Search"; + public readonly description: string; + public readonly parameters = webSearchSchema; + + constructor(_session: ToolSession) { + this.description = renderPromptTemplate(webSearchDescription); + } + + public async execute( + _toolCallId: string, + params: WebSearchParams, + _signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + _context?: AgentToolContext, + ): Promise> { + return executeWebSearch(_toolCallId, params); + } +} /** Web search tool as CustomTool (for TUI rendering support) */ export const webSearchCustomTool: CustomTool = { @@ -367,11 +383,6 @@ export const webSearchCustomTool: CustomTool { - return webSearchTool; -} - // ============================================================================ // Exa-specific tools (available when EXA_API_KEY is present) // ============================================================================ diff --git a/packages/coding-agent/src/core/tools/write.ts b/packages/coding-agent/src/core/tools/write.ts index 47a26b6a2..7f90e13f5 100644 --- a/packages/coding-agent/src/core/tools/write.ts +++ b/packages/coding-agent/src/core/tools/write.ts @@ -1,4 +1,10 @@ -import type { AgentTool, AgentToolContext, ToolCallContext } from "@oh-my-pi/pi-agent-core"; +import type { + AgentTool, + AgentToolContext, + AgentToolResult, + AgentToolUpdateCallback, + ToolCallContext, +} from "@oh-my-pi/pi-agent-core"; import type { Component } from "@oh-my-pi/pi-tui"; import { Text } from "@oh-my-pi/pi-tui"; import { Type } from "@sinclair/typebox"; @@ -8,7 +14,12 @@ import type { RenderResultOptions } from "../custom-tools/types"; import { renderPromptTemplate } from "../prompt-templates"; import type { ToolSession } from "../sdk"; import { untilAborted } from "../utils"; -import { createLspWritethrough, type FileDiagnosticsResult, writethroughNoop } from "./lsp/index"; +import { + createLspWritethrough, + type FileDiagnosticsResult, + type WritethroughCallback, + writethroughNoop, +} from "./lsp/index"; import { resolveToCwd } from "./path-utils"; import { formatDiagnostics, formatExpandHint, replaceTabs, shortenPath } from "./render-utils"; @@ -38,51 +49,69 @@ function getLspBatchRequest(toolCall: ToolCallContext | undefined): { id: string return { id: toolCall.batchId, flush: !hasLaterWrites }; } -export function createWriteTool(session: ToolSession): AgentTool { - const enableLsp = session.enableLsp ?? true; - const enableFormat = enableLsp ? (session.settings?.getLspFormatOnWrite() ?? true) : false; - const enableDiagnostics = enableLsp ? (session.settings?.getLspDiagnosticsOnWrite() ?? true) : false; - const writethrough = enableLsp - ? createLspWritethrough(session.cwd, { enableFormat, enableDiagnostics }) - : writethroughNoop; - return { - name: "write", - label: "Write", - description: renderPromptTemplate(writeDescription), - parameters: writeSchema, - execute: async ( - _toolCallId: string, - { path, content }: { path: string; content: string }, - signal?: AbortSignal, - _onUpdate?: unknown, - context?: AgentToolContext, - ) => { - return untilAborted(signal, async () => { - const absolutePath = resolveToCwd(path, session.cwd); - const batchRequest = getLspBatchRequest(context?.toolCall); +// ═══════════════════════════════════════════════════════════════════════════ +// Tool Class +// ═══════════════════════════════════════════════════════════════════════════ - const diagnostics = await writethrough(absolutePath, content, signal, undefined, batchRequest); +type WriteParams = { path: string; content: string }; - let resultText = `Successfully wrote ${content.length} bytes to ${path}`; - if (!diagnostics) { - return { - content: [{ type: "text", text: resultText }], - details: {}, - }; - } +/** + * Write tool implementation. + * + * Creates or overwrites files with optional LSP formatting and diagnostics. + */ +export class WriteTool implements AgentTool { + public readonly name = "write"; + public readonly label = "Write"; + public readonly description: string; + public readonly parameters = writeSchema; - const messages = diagnostics?.messages; - if (messages && messages.length > 0) { - resultText += `\n\nLSP Diagnostics (${diagnostics.summary}):\n`; - resultText += messages.map((d) => ` ${d}`).join("\n"); - } + private readonly session: ToolSession; + private readonly writethrough: WritethroughCallback; + + constructor(session: ToolSession) { + this.session = session; + const enableLsp = session.enableLsp ?? true; + const enableFormat = enableLsp ? (session.settings?.getLspFormatOnWrite() ?? true) : false; + const enableDiagnostics = enableLsp ? (session.settings?.getLspDiagnosticsOnWrite() ?? true) : false; + this.writethrough = enableLsp + ? createLspWritethrough(session.cwd, { enableFormat, enableDiagnostics }) + : writethroughNoop; + this.description = renderPromptTemplate(writeDescription); + } + + public async execute( + _toolCallId: string, + { path, content }: WriteParams, + signal?: AbortSignal, + _onUpdate?: AgentToolUpdateCallback, + context?: AgentToolContext, + ): Promise> { + return untilAborted(signal, async () => { + const absolutePath = resolveToCwd(path, this.session.cwd); + const batchRequest = getLspBatchRequest(context?.toolCall); + + const diagnostics = await this.writethrough(absolutePath, content, signal, undefined, batchRequest); + + let resultText = `Successfully wrote ${content.length} bytes to ${path}`; + if (!diagnostics) { return { content: [{ type: "text", text: resultText }], - details: { diagnostics }, + details: {}, }; - }); - }, - }; + } + + const messages = diagnostics?.messages; + if (messages && messages.length > 0) { + resultText += `\n\nLSP Diagnostics (${diagnostics.summary}):\n`; + resultText += messages.map((d) => ` ${d}`).join("\n"); + } + return { + content: [{ type: "text", text: resultText }], + details: { diagnostics }, + }; + }); + } } // ============================================================================= diff --git a/packages/coding-agent/src/core/ttsr.ts b/packages/coding-agent/src/core/ttsr.ts index 137f9aff9..0a05439d9 100644 --- a/packages/coding-agent/src/core/ttsr.ts +++ b/packages/coding-agent/src/core/ttsr.ts @@ -21,44 +21,6 @@ interface InjectionRecord { lastInjectedAt: number; } -export interface TtsrManager { - /** Add a TTSR rule to be monitored */ - addRule(rule: Rule): void; - - /** Check if any uninjected TTSR matches the stream buffer. Returns matching rules. */ - check(streamBuffer: string): Rule[]; - - /** Mark rules as injected (won't trigger again until conditions allow) */ - markInjected(rules: Rule[]): void; - - /** Get names of all injected rules (for persistence) */ - getInjectedRuleNames(): string[]; - - /** Restore injected state from a list of rule names */ - restoreInjected(ruleNames: string[]): void; - - /** Reset stream buffer (called on new turn) */ - resetBuffer(): void; - - /** Get current stream buffer */ - getBuffer(): string; - - /** Append to stream buffer */ - appendToBuffer(text: string): void; - - /** Check if any TTSRs are registered */ - hasRules(): boolean; - - /** Increment message counter (call after each turn) */ - incrementMessageCount(): void; - - /** Get current message count */ - getMessageCount(): number; - - /** Get settings */ - getSettings(): Required; -} - const DEFAULT_SETTINGS: Required = { enabled: true, contextMode: "discard", @@ -66,146 +28,138 @@ const DEFAULT_SETTINGS: Required = { repeatGap: 10, }; -export function createTtsrManager(settings?: TtsrSettings): TtsrManager { - /** Resolved settings with defaults */ - const resolvedSettings: Required = { - ...DEFAULT_SETTINGS, - ...settings, - }; +export class TtsrManager { + private readonly settings: Required; + private readonly rules = new Map(); + private readonly injectionRecords = new Map(); + private buffer = ""; + private messageCount = 0; - /** Map of rule name -> { rule, compiled regex } */ - const rules = new Map(); - - /** Map of rule name -> injection record */ - const injectionRecords = new Map(); - - /** Current stream buffer for pattern matching */ - let buffer = ""; - - /** Message counter for tracking gap between injections */ - let messageCount = 0; + constructor(settings?: TtsrSettings) { + this.settings = { ...DEFAULT_SETTINGS, ...settings }; + } /** Check if a rule can be triggered based on repeat settings */ - function canTrigger(ruleName: string): boolean { - const record = injectionRecords.get(ruleName); + private canTrigger(ruleName: string): boolean { + const record = this.injectionRecords.get(ruleName); if (!record) { - // Never injected, can trigger return true; } - if (resolvedSettings.repeatMode === "once") { - // Once mode: never trigger again after first injection + if (this.settings.repeatMode === "once") { return false; } - // After-gap mode: check if enough messages have passed - const gap = messageCount - record.lastInjectedAt; - return gap >= resolvedSettings.repeatGap; + const gap = this.messageCount - record.lastInjectedAt; + return gap >= this.settings.repeatGap; } - return { - addRule(rule: Rule): void { - // Only add rules that have a TTSR trigger pattern - if (!rule.ttsrTrigger) { - return; + /** Add a TTSR rule to be monitored */ + addRule(rule: Rule): void { + if (!rule.ttsrTrigger) { + return; + } + + if (this.rules.has(rule.name)) { + return; + } + + try { + const regex = new RegExp(rule.ttsrTrigger); + this.rules.set(rule.name, { rule, regex }); + logger.debug("TTSR rule registered", { + ruleName: rule.name, + pattern: rule.ttsrTrigger, + }); + } catch (err) { + logger.warn("TTSR rule has invalid regex pattern, skipping", { + ruleName: rule.name, + pattern: rule.ttsrTrigger, + error: err instanceof Error ? err.message : String(err), + }); + } + } + + /** Check if any uninjected TTSR matches the stream buffer. Returns matching rules. */ + check(streamBuffer: string): Rule[] { + const matches: Rule[] = []; + + for (const [name, entry] of this.rules) { + if (!this.canTrigger(name)) { + continue; } - // Skip if already registered - if (rules.has(rule.name)) { - return; - } - - // Compile the regex pattern - try { - const regex = new RegExp(rule.ttsrTrigger); - rules.set(rule.name, { rule, regex }); - logger.debug("TTSR rule registered", { - ruleName: rule.name, - pattern: rule.ttsrTrigger, - }); - } catch (err) { - logger.warn("TTSR rule has invalid regex pattern, skipping", { - ruleName: rule.name, - pattern: rule.ttsrTrigger, - error: err instanceof Error ? err.message : String(err), + if (entry.regex.test(streamBuffer)) { + matches.push(entry.rule); + logger.debug("TTSR pattern matched", { + ruleName: name, + pattern: entry.rule.ttsrTrigger, }); } - }, + } - check(streamBuffer: string): Rule[] { - const matches: Rule[] = []; + return matches; + } - for (const [name, entry] of rules) { - // Skip rules that can't trigger yet - if (!canTrigger(name)) { - continue; - } + /** Mark rules as injected (won't trigger again until conditions allow) */ + markInjected(rulesToMark: Rule[]): void { + for (const rule of rulesToMark) { + this.injectionRecords.set(rule.name, { lastInjectedAt: this.messageCount }); + logger.debug("TTSR rule marked as injected", { + ruleName: rule.name, + messageCount: this.messageCount, + repeatMode: this.settings.repeatMode, + }); + } + } - // Test the buffer against the rule's pattern - if (entry.regex.test(streamBuffer)) { - matches.push(entry.rule); - logger.debug("TTSR pattern matched", { - ruleName: name, - pattern: entry.rule.ttsrTrigger, - }); - } - } + /** Get names of all injected rules (for persistence) */ + getInjectedRuleNames(): string[] { + return Array.from(this.injectionRecords.keys()); + } - return matches; - }, + /** Restore injected state from a list of rule names */ + restoreInjected(ruleNames: string[]): void { + for (const name of ruleNames) { + this.injectionRecords.set(name, { lastInjectedAt: 0 }); + } + if (ruleNames.length > 0) { + logger.debug("TTSR injected state restored", { ruleNames }); + } + } - markInjected(rulesToMark: Rule[]): void { - for (const rule of rulesToMark) { - injectionRecords.set(rule.name, { lastInjectedAt: messageCount }); - logger.debug("TTSR rule marked as injected", { - ruleName: rule.name, - messageCount, - repeatMode: resolvedSettings.repeatMode, - }); - } - }, + /** Reset stream buffer (called on new turn) */ + resetBuffer(): void { + this.buffer = ""; + } - getInjectedRuleNames(): string[] { - return Array.from(injectionRecords.keys()); - }, + /** Get current stream buffer */ + getBuffer(): string { + return this.buffer; + } - restoreInjected(ruleNames: string[]): void { - // When restoring, we don't know the original message count, so use 0 - // This means in "after-gap" mode, rules can trigger again after the gap - for (const name of ruleNames) { - injectionRecords.set(name, { lastInjectedAt: 0 }); - } - if (ruleNames.length > 0) { - logger.debug("TTSR injected state restored", { ruleNames }); - } - }, + /** Append to stream buffer */ + appendToBuffer(text: string): void { + this.buffer += text; + } - resetBuffer(): void { - buffer = ""; - }, + /** Check if any TTSRs are registered */ + hasRules(): boolean { + return this.rules.size > 0; + } - getBuffer(): string { - return buffer; - }, + /** Increment message counter (call after each turn) */ + incrementMessageCount(): void { + this.messageCount++; + } - appendToBuffer(text: string): void { - buffer += text; - }, + /** Get current message count */ + getMessageCount(): number { + return this.messageCount; + } - hasRules(): boolean { - return rules.size > 0; - }, - - incrementMessageCount(): void { - messageCount++; - }, - - getMessageCount(): number { - return messageCount; - }, - - getSettings(): Required { - return resolvedSettings; - }, - }; + /** Get settings */ + getSettings(): Required { + return this.settings; + } } diff --git a/packages/coding-agent/src/index.ts b/packages/coding-agent/src/index.ts index bcb15a6e6..0109b25b0 100644 --- a/packages/coding-agent/src/index.ts +++ b/packages/coding-agent/src/index.ts @@ -63,7 +63,7 @@ export type { LoadedCustomTool, RenderResultOptions, } from "./core/custom-tools/index"; -export { discoverAndLoadCustomTools, loadCustomTools } from "./core/custom-tools/index"; +export { CustomToolLoader, discoverAndLoadCustomTools, loadCustomTools } from "./core/custom-tools/index"; // Extension types and utilities export type { AppAction, @@ -79,7 +79,6 @@ export type { ExtensionFactory, ExtensionFlag, ExtensionHandler, - ExtensionRuntime, ExtensionShortcut, ExtensionUIContext, ExtensionUIDialogOptions, @@ -97,9 +96,9 @@ export type { UserBashEventResult, } from "./core/extensions/index"; export { - createExtensionRuntime, discoverAndLoadExtensions, ExtensionRunner, + ExtensionRuntime, isBashToolResult, isEditToolResult, isFindToolResult, @@ -119,21 +118,16 @@ export { ModelRegistry } from "./core/model-registry"; export type { PromptTemplate } from "./core/prompt-templates"; // SDK for programmatic usage export { + // Factory + BashTool, // Tool factories BUILTIN_TOOLS, type BuildSystemPromptOptions, buildSystemPrompt, type CreateAgentSessionOptions, type CreateAgentSessionResult, - // Factory createAgentSession, - createBashTool, - createFindTool, - createGrepTool, - createLsTool, - createReadTool, createTools, - createWriteTool, // Discovery discoverAuthStorage, discoverContextFiles, @@ -144,8 +138,15 @@ export { discoverPromptTemplates, discoverSkills, EditTool, + FindTool, + GrepTool, + LsTool, loadSettings, + loadSshTool, + PythonTool, + ReadTool, type ToolSession, + WriteTool, } from "./core/sdk"; export { type BranchSummaryEntry, @@ -201,11 +202,11 @@ export { type FindToolDetails, type FindToolOptions, formatSize, + GitTool, type GitToolDetails, type GrepOperations, type GrepToolDetails, type GrepToolOptions, - gitTool, type LsOperations, type LsToolDetails, type LsToolOptions, diff --git a/packages/coding-agent/src/modes/interactive/components/tool-execution.ts b/packages/coding-agent/src/modes/interactive/components/tool-execution.ts index 2b14f35f4..b7114c6ef 100644 --- a/packages/coding-agent/src/modes/interactive/components/tool-execution.ts +++ b/packages/coding-agent/src/modes/interactive/components/tool-execution.ts @@ -203,11 +203,9 @@ export class ToolExecutionComponent extends Container { if (this.editDiffArgsKey === argsKey) return; this.editDiffArgsKey = argsKey; - computePatchDiff( - { path, operation, moveTo, diff }, - this.cwd, - { fuzzyThreshold: this.editFuzzyThreshold }, - ).then((result) => { + computePatchDiff({ path, operation, moveTo, diff }, this.cwd, { + fuzzyThreshold: this.editFuzzyThreshold, + }).then((result) => { if (this.editDiffArgsKey === argsKey) { this.editDiffPreview = result; this.updateDisplay(); diff --git a/packages/coding-agent/test/block-images.test.ts b/packages/coding-agent/test/block-images.test.ts index 9b5743a79..5e0a80049 100644 --- a/packages/coding-agent/test/block-images.test.ts +++ b/packages/coding-agent/test/block-images.test.ts @@ -5,7 +5,7 @@ import { join } from "node:path"; import { processFileArguments } from "../src/cli/file-processor"; import { SettingsManager } from "../src/core/settings-manager"; import type { ToolSession } from "../src/core/tools/index"; -import { createReadTool } from "../src/core/tools/read"; +import { ReadTool } from "../src/core/tools/read"; // 1x1 red PNG image as base64 (smallest valid PNG) const TINY_PNG_BASE64 = @@ -69,7 +69,7 @@ describe("blockImages setting", () => { const imagePath = join(testDir, "test.png"); writeFileSync(imagePath, Buffer.from(TINY_PNG_BASE64, "base64")); - const tool = createReadTool(createTestToolSession(testDir)); + const tool = new ReadTool(createTestToolSession(testDir)); const result = await tool.execute("test-1", { path: imagePath }); // Should have text note + image content @@ -83,7 +83,7 @@ describe("blockImages setting", () => { const textPath = join(testDir, "test.txt"); writeFileSync(textPath, "Hello, world!"); - const tool = createReadTool(createTestToolSession(testDir)); + const tool = new ReadTool(createTestToolSession(testDir)); const result = await tool.execute("test-2", { path: textPath }); expect(result.content).toHaveLength(1); diff --git a/packages/coding-agent/test/core/python-prelude.test.ts b/packages/coding-agent/test/core/python-prelude.test.ts index 4618e5a07..69e010ec7 100644 --- a/packages/coding-agent/test/core/python-prelude.test.ts +++ b/packages/coding-agent/test/core/python-prelude.test.ts @@ -2,7 +2,7 @@ import { describe, expect, it } from "bun:test"; import { existsSync } from "node:fs"; import { join } from "node:path"; import { resetPreludeDocsCache, warmPythonEnvironment } from "../../src/core/python-executor"; -import { createPythonTool, getPythonToolDescription } from "../../src/core/tools/python"; +import { getPythonToolDescription, PythonTool } from "../../src/core/tools/python"; const resolvePythonPath = (): string | null => { const venvPath = process.env.VIRTUAL_ENV; @@ -102,7 +102,7 @@ describe.skipIf(!shouldRun)("PYTHON_PRELUDE integration", () => { }, }; - const tool = createPythonTool(session); + const tool = new PythonTool(session); const code = ` helpers = ${JSON.stringify(helpers)} missing = [name for name in helpers if name not in globals() or not callable(globals()[name])] diff --git a/packages/coding-agent/test/core/streaming-output.test.ts b/packages/coding-agent/test/core/streaming-output.test.ts index 542d9ce47..8e5c80a9e 100644 --- a/packages/coding-agent/test/core/streaming-output.test.ts +++ b/packages/coding-agent/test/core/streaming-output.test.ts @@ -1,14 +1,14 @@ import { describe, expect, it } from "bun:test"; -import { createOutputSink } from "../../src/core/streaming-output"; +import { OutputSink } from "../../src/core/streaming-output"; function makeLargeOutput(size: number): string { return "x".repeat(size); } -describe("createOutputSink", () => { +describe("OutputSink", () => { it("spills to disk and truncates large output", async () => { const largeOutput = makeLargeOutput(60_000); - const sink = createOutputSink(10, 70_000); + const sink = new OutputSink(10, 70_000); const writer = sink.getWriter(); await writer.write(largeOutput); diff --git a/packages/coding-agent/test/python-tool-settings.test.ts b/packages/coding-agent/test/python-tool-settings.test.ts index df0db8151..a19cd9074 100644 --- a/packages/coding-agent/test/python-tool-settings.test.ts +++ b/packages/coding-agent/test/python-tool-settings.test.ts @@ -5,7 +5,7 @@ import { join } from "node:path"; import * as pythonExecutor from "../src/core/python-executor"; import * as pythonKernel from "../src/core/python-kernel"; import { createTools, type ToolSession } from "../src/core/tools/index"; -import { createPythonTool } from "../src/core/tools/python"; +import { PythonTool } from "../src/core/tools/python"; function createSettings(overrides?: Partial): ToolSession["settings"] { return { @@ -76,7 +76,7 @@ describe("python tool settings", () => { }); const session = createSession(testDir, { getPythonKernelMode: () => "per-call" }); - const pythonTool = createPythonTool(session); + const pythonTool = new PythonTool(session); await pythonTool.execute("tool-call", { code: "print(1)" }); diff --git a/packages/coding-agent/test/tools.test.ts b/packages/coding-agent/test/tools.test.ts index f363ab4c9..e4f22ef08 100644 --- a/packages/coding-agent/test/tools.test.ts +++ b/packages/coding-agent/test/tools.test.ts @@ -3,14 +3,14 @@ import { mkdirSync, readFileSync, rmSync, writeFileSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; import { nanoid } from "nanoid"; -import { createBashTool } from "../src/core/tools/bash"; -import { createFindTool } from "../src/core/tools/find"; -import { createGrepTool } from "../src/core/tools/grep"; +import { BashTool } from "../src/core/tools/bash"; +import { FindTool } from "../src/core/tools/find"; +import { GrepTool } from "../src/core/tools/grep"; import type { ToolSession } from "../src/core/tools/index"; -import { createLsTool } from "../src/core/tools/ls"; +import { LsTool } from "../src/core/tools/ls"; import { EditTool } from "../src/core/tools/patch"; -import { createReadTool } from "../src/core/tools/read"; -import { createWriteTool } from "../src/core/tools/write"; +import { ReadTool } from "../src/core/tools/read"; +import { WriteTool } from "../src/core/tools/write"; import * as shellModule from "../src/utils/shell"; // Helper to extract text from content blocks @@ -34,13 +34,13 @@ function createTestToolSession(cwd: string): ToolSession { describe("Coding Agent Tools", () => { let testDir: string; - let readTool: ReturnType; - let writeTool: ReturnType; + let readTool: ReadTool; + let writeTool: WriteTool; let editTool: EditTool; - let bashTool: ReturnType; - let grepTool: ReturnType; - let findTool: ReturnType; - let lsTool: ReturnType; + let bashTool: BashTool; + let grepTool: GrepTool; + let findTool: FindTool; + let lsTool: LsTool; beforeEach(() => { // Create a unique temporary directory for each test @@ -49,13 +49,13 @@ describe("Coding Agent Tools", () => { // Create tools for this test directory const session = createTestToolSession(testDir); - readTool = createReadTool(session); - writeTool = createWriteTool(session); + readTool = new ReadTool(session); + writeTool = new WriteTool(session); editTool = new EditTool(session); - bashTool = createBashTool(session); - grepTool = createGrepTool(session); - findTool = createFindTool(session); - lsTool = createLsTool(session); + bashTool = new BashTool(session); + grepTool = new GrepTool(session); + findTool = new FindTool(session); + lsTool = new LsTool(session); }); afterEach(() => { @@ -252,9 +252,9 @@ describe("Coding Agent Tools", () => { expect(getTextOutput(result)).toContain("Successfully replaced"); expect(result.details).toBeDefined(); - expect(result.details.diff).toBeDefined(); - expect(typeof result.details.diff).toBe("string"); - expect(result.details.diff).toContain("testing"); + expect(result.details!.diff).toBeDefined(); + expect(typeof result.details!.diff).toBe("string"); + expect(result.details!.diff).toContain("testing"); }); it("should fail if text not found", async () => { @@ -402,7 +402,7 @@ function b() { it("should throw error when cwd does not exist", async () => { const nonexistentCwd = "/this/directory/definitely/does/not/exist/12345"; - const bashToolWithBadCwd = createBashTool(createTestToolSession(nonexistentCwd)); + const bashToolWithBadCwd = new BashTool(createTestToolSession(nonexistentCwd)); await expect(bashToolWithBadCwd.execute("test-call-11", { command: "echo test" })).rejects.toThrow( /Working directory does not exist/, @@ -417,7 +417,7 @@ function b() { prefix: undefined, }); - const bashWithBadShell = createBashTool(createTestToolSession(testDir)); + const bashWithBadShell = new BashTool(createTestToolSession(testDir)); await expect(bashWithBadShell.execute("test-call-12", { command: "echo test" })).rejects.toThrow(/ENOENT/); diff --git a/packages/coding-agent/test/tools/python-execution.test.ts b/packages/coding-agent/test/tools/python-execution.test.ts index a2743b1db..735de5c55 100644 --- a/packages/coding-agent/test/tools/python-execution.test.ts +++ b/packages/coding-agent/test/tools/python-execution.test.ts @@ -4,7 +4,7 @@ import { tmpdir } from "node:os"; import { join } from "node:path"; import * as pythonExecutor from "../../src/core/python-executor"; import type { ToolSession } from "../../src/core/tools/index"; -import { createPythonTool } from "../../src/core/tools/python"; +import { PythonTool } from "../../src/core/tools/python"; function createSession(cwd: string): ToolSession { return { @@ -40,7 +40,7 @@ describe("python tool execution", () => { stdinRequested: false, }); - const tool = createPythonTool(createSession(tempDir)); + const tool = new PythonTool(createSession(tempDir)); const result = await tool.execute( "call-id", { code: "print('hi')", timeout: 5, workdir: tempDir, reset: true }, diff --git a/packages/coding-agent/test/tools/python.test.ts b/packages/coding-agent/test/tools/python.test.ts index ec22254ad..0b049e31b 100644 --- a/packages/coding-agent/test/tools/python.test.ts +++ b/packages/coding-agent/test/tools/python.test.ts @@ -1,7 +1,7 @@ import { afterAll, beforeAll, describe, expect, it, vi } from "bun:test"; import * as pythonExecutor from "../../src/core/python-executor"; import { createTools, type ToolSession } from "../../src/core/tools/index"; -import { createPythonTool } from "../../src/core/tools/python"; +import { PythonTool } from "../../src/core/tools/python"; let previousSkipCheck: string | undefined; @@ -46,7 +46,7 @@ function createSettings(toolMode: "ipy-only" | "bash-only" | "both") { describe("python tool schema", () => { it("exposes expected parameters", () => { - const tool = createPythonTool(createSession()); + const tool = new PythonTool(createSession()); const schema = tool.parameters as { type: string; properties: Record; @@ -74,7 +74,7 @@ describe("python tool docs template", () => { ]; const spy = vi.spyOn(pythonExecutor, "getPreludeDocs").mockReturnValue(docs); - const tool = createPythonTool(createSession()); + const tool = new PythonTool(createSession()); expect(tool.description).toContain("### File I/O"); expect(tool.description).toContain("read(path)"); @@ -86,7 +86,7 @@ describe("python tool docs template", () => { it("renders fallback when docs are unavailable", () => { const spy = vi.spyOn(pythonExecutor, "getPreludeDocs").mockReturnValue([]); - const tool = createPythonTool(createSession()); + const tool = new PythonTool(createSession()); expect(tool.description).toContain("Documentation unavailable — Python kernel failed to start");