diff --git a/crates/pi-natives/src/chunk/edit.rs b/crates/pi-natives/src/chunk/edit.rs index 5cb82d3d9..8b95deaec 100644 --- a/crates/pi-natives/src/chunk/edit.rs +++ b/crates/pi-natives/src/chunk/edit.rs @@ -1608,18 +1608,14 @@ mod tests { // The chunk range should include the #[napi] attribute. assert_eq!(chunk.start_line, 1, "chunk should start at the attribute line"); - let result = apply_single_edit( - &state, - "test.rs", - EditOperation { - op: ChunkEditOp::Replace, - sel: Some("fn_close".to_owned()), - crc: Some(chunk.checksum.clone()), - region: None, - content: Some("/// doc\n#[napi]\nfn close() {\n new();\n}".to_owned()), - find: None, - }, - ); + let result = apply_single_edit(&state, "test.rs", EditOperation { + op: ChunkEditOp::Replace, + sel: Some("fn_close".to_owned()), + crc: Some(chunk.checksum.clone()), + region: None, + content: Some("/// doc\n#[napi]\nfn close() {\n new();\n}".to_owned()), + find: None, + }); let occurrences = result.diff_after.matches("#[napi]").count(); assert_eq!( diff --git a/packages/agent/CHANGELOG.md b/packages/agent/CHANGELOG.md index ae5971110..db933da62 100644 --- a/packages/agent/CHANGELOG.md +++ b/packages/agent/CHANGELOG.md @@ -1,6 +1,10 @@ # Changelog ## [Unreleased] +### Added + +- Added `onAssistantMessageEvent` callback option to inspect assistant streaming events before they are emitted, enabling abort decisions before buffered events continue flowing +- Added `setAssistantMessageEventInterceptor()` method to dynamically set or update the assistant message event interceptor ## [13.13.0] - 2026-03-18 @@ -310,4 +314,4 @@ Initial release under @oh-my-pi scope. See previous releases at [badlogic/pi-mon - `Agent` constructor now has all options optional (empty options use defaults). -- `queueMessage()` is now synchronous (no longer returns a Promise). +- `queueMessage()` is now synchronous (no longer returns a Promise). \ No newline at end of file diff --git a/packages/agent/src/agent-loop.ts b/packages/agent/src/agent-loop.ts index e32feaea3..18bbc7fc9 100644 --- a/packages/agent/src/agent-loop.ts +++ b/packages/agent/src/agent-loop.ts @@ -391,6 +391,10 @@ async function streamAssistantResponse( if (partialMessage) { partialMessage = event.partial; context.messages[context.messages.length - 1] = partialMessage; + config.onAssistantMessageEvent?.(partialMessage, event); + if (signal?.aborted) { + continue; + } stream.push({ type: "message_update", assistantMessageEvent: event, diff --git a/packages/agent/src/agent.ts b/packages/agent/src/agent.ts index 38dbc6f74..17c18ec4e 100644 --- a/packages/agent/src/agent.ts +++ b/packages/agent/src/agent.ts @@ -3,6 +3,7 @@ */ import { type AssistantMessage, + type AssistantMessageEvent, type CursorExecHandlers, type CursorToolResultHandler, type Effort, @@ -130,7 +131,11 @@ export interface AgentOptions { * Inspect or replace provider payloads before they are sent. */ onPayload?: SimpleStreamOptions["onPayload"]; - + /** + * Inspect assistant streaming events before they are emitted to subscribers. + * Use this when abort decisions must happen before buffered events continue flowing. + */ + onAssistantMessageEvent?: (message: AssistantMessage, event: AssistantMessageEvent) => void; /** * Custom token budgets for thinking levels (token-based providers only). */ @@ -239,6 +244,7 @@ export class Agent { #intentTracing: boolean; #getToolChoice?: () => ToolChoice | undefined; #onPayload?: SimpleStreamOptions["onPayload"]; + #onAssistantMessageEvent?: (message: AssistantMessage, event: AssistantMessageEvent) => void; /** Buffered Cursor tool results with text length at time of call (for correct ordering) */ #cursorToolResultBuffer: CursorToolResultEntry[] = []; @@ -275,6 +281,7 @@ export class Agent { this.#transformToolCallArguments = opts.transformToolCallArguments; this.#intentTracing = opts.intentTracing === true; this.#getToolChoice = opts.getToolChoice; + this.#onAssistantMessageEvent = opts.onAssistantMessageEvent; } /** @@ -406,6 +413,12 @@ export class Agent { return () => this.#listeners.delete(fn); } + setAssistantMessageEventInterceptor( + fn: ((message: AssistantMessage, event: AssistantMessageEvent) => void) | undefined, + ): void { + this.#onAssistantMessageEvent = fn; + } + emitExternalEvent(event: AgentEvent) { switch (event.type) { case "message_start": @@ -762,6 +775,7 @@ export class Agent { cursorOnToolResult, transformToolCallArguments: this.#transformToolCallArguments, intentTracing: this.#intentTracing, + onAssistantMessageEvent: this.#onAssistantMessageEvent, getToolChoice, getSteeringMessages: async () => { if (skipInitialSteeringPoll) { diff --git a/packages/agent/src/types.ts b/packages/agent/src/types.ts index 43aa1d275..9ab64c8aa 100644 --- a/packages/agent/src/types.ts +++ b/packages/agent/src/types.ts @@ -1,4 +1,5 @@ import type { + AssistantMessage, AssistantMessageEvent, AssistantMessageEventStream, Effort, @@ -27,8 +28,8 @@ export interface AgentLoopConfig extends SimpleStreamOptions { /** * When to interrupt tool execution for steering messages. - * - "immediate": check after each tool call (default) - * - "wait": defer steering until the current turn completes + * - "immediate" = check after each tool call (default) + * - "wait" = defer steering until the current turn completes */ interruptMode?: "immediate" | "wait"; @@ -125,6 +126,7 @@ export interface AgentLoopConfig extends SimpleStreamOptions { * Use for deobfuscating secrets or rewriting arguments. */ transformToolCallArguments?: (args: Record, toolName: string) => Record; + /** * Enable intent tracing for tool calls. * When enabled, the harness injects an `_i: string` field into tool schemas sent to the model, @@ -132,6 +134,12 @@ export interface AgentLoopConfig extends SimpleStreamOptions { */ intentTracing?: boolean; + /** + * Inspect assistant streaming events before they are published to the outer agent event stream. + * Callers may abort synchronously to stop consuming buffered provider events. + */ + onAssistantMessageEvent?: (message: AssistantMessage, event: AssistantMessageEvent) => void; + /** * Dynamic tool choice override, resolved per LLM call. * When set and returns a value, overrides the static `toolChoice`. diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 036c7b7f3..ca1294f51 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,8 +1,12 @@ # Changelog ## [Unreleased] + ### Changed +- Made `assertEditableFileContent` synchronous instead of async for improved performance in streaming edit checks +- Enhanced streaming edit abort detection to check for auto-generated files as soon as the file path is available, rather than waiting for the full diff +- Improved file prefix reading in session storage to use `peekFile` utility from @oh-my-pi/pi-utils for better efficiency - Moved image metadata detection to @oh-my-pi/pi-utils package for shared use across projects - Simplified image loading API by removing redundant metadata parameters and consolidating image utilities - Updated imports to use readImageMetadata and parseImageMetadata from @oh-my-pi/pi-utils instead of local implementations @@ -13,6 +17,10 @@ - Removed mime.ts utility module; MIME detection moved to @oh-my-pi/pi-utils - Removed ImageMetadata interface and ReadImageMetadataOptions from local codebase +### Fixed + +- Fixed streaming edit abort for auto-generated files by adding LRU caching and early path-based detection to prevent unnecessary edits + ## [14.0.0] - 2026-04-08 ### Breaking Changes diff --git a/packages/coding-agent/src/edit/modes/chunk.ts b/packages/coding-agent/src/edit/modes/chunk.ts index 536c8d30a..bf8bafaa8 100644 --- a/packages/coding-agent/src/edit/modes/chunk.ts +++ b/packages/coding-agent/src/edit/modes/chunk.ts @@ -18,7 +18,7 @@ import type { Settings } from "../../config/settings"; import type { WritethroughCallback, WritethroughDeferredHandle } from "../../lsp"; import { getLanguageFromPath } from "../../modes/theme/theme"; import type { ToolSession } from "../../tools"; -import { checkAutoGeneratedFileContent } from "../../tools/auto-generated-guard"; +import { assertEditableFileContent } from "../../tools/auto-generated-guard"; import { invalidateFsScanAfterWrite } from "../../tools/fs-cache-invalidation"; import { outputMeta } from "../../tools/output-meta"; import { enforcePlanModeWrite, resolvePlanPath } from "../../tools/plan-mode-guard"; @@ -125,7 +125,7 @@ async function resolveChunkSourceContext(session: ToolSession, path: string): Pr let rawContent = ""; if (sourceExists) { rawContent = await sourceFile.text(); - await checkAutoGeneratedFileContent(rawContent, path); + assertEditableFileContent(rawContent, path); } return { diff --git a/packages/coding-agent/src/edit/modes/hashline.ts b/packages/coding-agent/src/edit/modes/hashline.ts index 1ac834868..76dce4016 100644 --- a/packages/coding-agent/src/edit/modes/hashline.ts +++ b/packages/coding-agent/src/edit/modes/hashline.ts @@ -20,7 +20,7 @@ import { type Static, Type } from "@sinclair/typebox"; import type { BunFile } from "bun"; import type { WritethroughCallback, WritethroughDeferredHandle } from "../../lsp"; import type { ToolSession } from "../../tools"; -import { checkAutoGeneratedFileContent } from "../../tools/auto-generated-guard"; +import { assertEditableFileContent } from "../../tools/auto-generated-guard"; import { invalidateFsScanAfterDelete, invalidateFsScanAfterRename, @@ -1263,7 +1263,7 @@ export async function executeHashlineMode( const anchorEdits = resolveEditAnchors(edits); const rawContent = await sourceFile.text(); - await checkAutoGeneratedFileContent(rawContent, path); + assertEditableFileContent(rawContent, path); const { bom, text } = stripBom(rawContent); const originalEnding = detectLineEnding(text); diff --git a/packages/coding-agent/src/edit/modes/patch.ts b/packages/coding-agent/src/edit/modes/patch.ts index 43224616c..e9e27f4ba 100644 --- a/packages/coding-agent/src/edit/modes/patch.ts +++ b/packages/coding-agent/src/edit/modes/patch.ts @@ -18,7 +18,7 @@ import { type WritethroughDeferredHandle, } from "../../lsp"; import type { ToolSession } from "../../tools"; -import { checkAutoGeneratedFile } from "../../tools/auto-generated-guard"; +import { assertEditableFile } from "../../tools/auto-generated-guard"; import { invalidateFsScanAfterDelete, invalidateFsScanAfterRename, @@ -1718,7 +1718,7 @@ export async function executePatchMode( throw new Error("Cannot edit Jupyter notebooks with the Edit tool. Use the NotebookEdit tool instead."); } - await checkAutoGeneratedFile(resolvedPath, path); + await assertEditableFile(resolvedPath, path); const input: PatchInput = { path: resolvedPath, op, rename: resolvedRename, diff }; const patchFileSystem = new LspFileSystem(writethrough, signal, batchRequest, beginDeferredDiagnosticsForPath); diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index 2140e00e9..f433f71de 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -116,12 +116,14 @@ import planModeToolDecisionReminderPrompt from "../prompts/system/plan-mode-tool import ttsrInterruptTemplate from "../prompts/system/ttsr-interrupt.md" with { type: "text" }; import { deobfuscateSessionContext, type SecretObfuscator } from "../secrets/obfuscator"; import { resolveThinkingLevelForModel, toReasoningEffort } from "../thinking"; +import { assertEditableFile } from "../tools/auto-generated-guard"; import type { CheckpointState } from "../tools/checkpoint"; import { outputMeta } from "../tools/output-meta"; import { resolveToCwd } from "../tools/path-utils"; import type { PendingActionStore } from "../tools/pending-action"; import { isAutoQaEnabled } from "../tools/report-tool-issue"; import { getLatestTodoPhasesFromEntries, type TodoItem, type TodoPhase } from "../tools/todo-write"; +import { ToolError } from "../tools/tool-errors"; import { clampTimeout } from "../tools/tool-timeouts"; import { parseCommandArgs } from "../utils/command-args"; import { resolveFileDisplayMode } from "../utils/file-display-mode"; @@ -503,6 +505,9 @@ export class AgentSession { #streamingEditAbortTriggered = false; #streamingEditCheckedLineCounts = new Map(); + + #streamingEditPrecheckedToolCallIds = new Set(); + #streamingEditFileCache = new Map(); #promptInFlightCount = 0; #obfuscator: SecretObfuscator | undefined; @@ -585,6 +590,15 @@ export class AgentSession { ); this.#ttsrManager = config.ttsrManager; this.#obfuscator = config.obfuscator; + this.agent.setAssistantMessageEventInterceptor((message, assistantMessageEvent) => { + const event: AgentEvent = { + type: "message_update", + message, + assistantMessageEvent, + }; + this.#preCacheStreamingEditFile(event); + this.#maybeAbortStreamingEdit(event); + }); this.agent.providerSessionState = this.#providerSessionState; this.#pendingActionStore = config.pendingActionStore; this.#unsubscribePendingActionPush = this.#pendingActionStore?.subscribePush(action => { @@ -822,8 +836,13 @@ export class AgentSession { } } - if (event.type === "message_update" && event.assistantMessageEvent.type === "toolcall_start") { - this.#preCacheStreamingEditFile(event); + if ( + event.type === "message_update" && + (event.assistantMessageEvent.type === "toolcall_start" || + event.assistantMessageEvent.type === "toolcall_delta" || + event.assistantMessageEvent.type === "toolcall_end") + ) { + void this.#preCacheStreamingEditFile(event); } if ( @@ -1354,31 +1373,101 @@ export class AgentSession { #resetStreamingEditState(): void { this.#streamingEditAbortTriggered = false; this.#streamingEditCheckedLineCounts.clear(); + this.#streamingEditPrecheckedToolCallIds.clear(); this.#streamingEditFileCache.clear(); } - async #preCacheStreamingEditFile(event: AgentEvent): Promise { - if (!this.settings.get("edit.streamingAbort")) return; - if (event.type !== "message_update") return; - const assistantEvent = event.assistantMessageEvent; - if (assistantEvent.type !== "toolcall_start") return; - if (event.message.role !== "assistant") return; + #getStreamingEditToolCall(event: AgentEvent): + | { + toolCall: ToolCall; + path: string; + resolvedPath: string; + diff?: string; + op?: string; + rename?: string; + } + | undefined { + if (event.type !== "message_update") return undefined; + if (event.message.role !== "assistant") return undefined; - const contentIndex = assistantEvent.contentIndex; + const contentIndex = event.assistantMessageEvent.contentIndex ?? 0; const messageContent = event.message.content; - if (!Array.isArray(messageContent) || contentIndex >= messageContent.length) return; + if (!Array.isArray(messageContent) || contentIndex < 0 || contentIndex >= messageContent.length) { + return undefined; + } + const toolCall = messageContent[contentIndex] as ToolCall; - if (toolCall.name !== "edit") return; + if (toolCall.name !== "edit") return undefined; const args = toolCall.arguments; - if (!args || typeof args !== "object" || Array.isArray(args)) return; - if ("old_text" in args || "new_text" in args) return; + if (!args || typeof args !== "object" || Array.isArray(args)) return undefined; + if ("old_text" in args || "new_text" in args) return undefined; const path = typeof args.path === "string" ? args.path : undefined; - if (!path) return; + if (!path) return undefined; - const resolvedPath = resolveToCwd(path, this.sessionManager.getCwd()); - this.#ensureFileCache(resolvedPath); + return { + toolCall, + path, + resolvedPath: resolveToCwd(path, this.sessionManager.getCwd()), + diff: typeof args.diff === "string" ? args.diff : undefined, + op: typeof args.op === "string" ? args.op : undefined, + rename: typeof args.rename === "string" ? args.rename : undefined, + }; + } + + #lastStreamingEditToolCallId: string | undefined; + #abortStreamingEditForAutoGeneratedPath(toolCall: ToolCall, path: string, resolvedPath: string): void { + if (this.#lastStreamingEditToolCallId === toolCall.id) return; + this.#lastStreamingEditToolCallId = toolCall.id; + void assertEditableFile(resolvedPath, path).catch(err => { + // peekFile and other I/O can reject with ENOENT, etc. Only ToolError means + // auto-generated detection; other failures are left for the edit tool. + if (!(err instanceof ToolError)) return; + if (this.#lastStreamingEditToolCallId !== toolCall.id) return; + + if (!this.#streamingEditAbortTriggered) { + this.#streamingEditAbortTriggered = true; + logger.warn("Streaming edit aborted due to auto-generated file guard", { + toolCallId: toolCall.id, + path, + }); + this.agent.abort(); + } + }); + } + + #preCacheStreamingEditFile(event: AgentEvent): void { + if (!this.settings.get("edit.streamingAbort")) return; + if (this.#streamingEditAbortTriggered) return; + if (event.type !== "message_update") return; + + const assistantEvent = event.assistantMessageEvent; + if ( + assistantEvent.type !== "toolcall_start" && + assistantEvent.type !== "toolcall_delta" && + assistantEvent.type !== "toolcall_end" + ) { + return; + } + + const streamingEdit = this.#getStreamingEditToolCall(event); + if (!streamingEdit) return; + + const shouldCheckAutoGenerated = + !streamingEdit.toolCall.id || !this.#streamingEditPrecheckedToolCallIds.has(streamingEdit.toolCall.id); + if (shouldCheckAutoGenerated) { + if (streamingEdit.toolCall.id) { + this.#streamingEditPrecheckedToolCallIds.add(streamingEdit.toolCall.id); + } + this.#abortStreamingEditForAutoGeneratedPath( + streamingEdit.toolCall, + streamingEdit.path, + streamingEdit.resolvedPath, + ); + } + + this.#ensureFileCache(streamingEdit.resolvedPath); } #ensureFileCache(resolvedPath: string): void { @@ -1403,24 +1492,15 @@ export class AgentSession { if (!this.settings.get("edit.streamingAbort")) return; if (this.#streamingEditAbortTriggered) return; if (event.type !== "message_update") return; + const assistantEvent = event.assistantMessageEvent; if (assistantEvent.type !== "toolcall_end" && assistantEvent.type !== "toolcall_delta") return; - if (event.message.role !== "assistant") return; - const contentIndex = assistantEvent.contentIndex; - const messageContent = event.message.content; - if (!Array.isArray(messageContent) || contentIndex >= messageContent.length) return; - const toolCall = messageContent[contentIndex] as ToolCall; - if (toolCall.name !== "edit" || !toolCall.id) return; + const streamingEdit = this.#getStreamingEditToolCall(event); + if (!streamingEdit?.toolCall.id) return; - const args = toolCall.arguments; - if (!args || typeof args !== "object" || Array.isArray(args)) return; - if ("old_text" in args || "new_text" in args) return; - - const path = typeof args.path === "string" ? args.path : undefined; - const diff = typeof args.diff === "string" ? args.diff : undefined; - const op = typeof args.op === "string" ? args.op : undefined; - if (!path || !diff) return; + const { toolCall, path, resolvedPath, diff, op, rename } = streamingEdit; + if (!diff) return; if (op && op !== "update") return; if (!diff.includes("\n")) return; @@ -1443,13 +1523,10 @@ export class AgentSession { if (lastChecked !== undefined && lineCount <= lastChecked) return; this.#streamingEditCheckedLineCounts.set(toolCall.id, lineCount); - const rename = typeof args.rename === "string" ? args.rename : undefined; - const removedLines = lines .filter(line => line.startsWith("-") && !line.startsWith("--- ")) .map(line => line.slice(1)); if (removedLines.length > 0) { - const resolvedPath = resolveToCwd(path, this.sessionManager.getCwd()); let cachedContent = this.#streamingEditFileCache.get(resolvedPath); if (cachedContent === undefined) { this.#ensureFileCache(resolvedPath); diff --git a/packages/coding-agent/src/session/session-storage.ts b/packages/coding-agent/src/session/session-storage.ts index d33984088..68dd53534 100644 --- a/packages/coding-agent/src/session/session-storage.ts +++ b/packages/coding-agent/src/session/session-storage.ts @@ -1,7 +1,9 @@ import * as fs from "node:fs"; import * as fsp from "node:fs/promises"; import * as path from "node:path"; -import { isEnoent, toError } from "@oh-my-pi/pi-utils"; +import { isEnoent, peekFile, toError } from "@oh-my-pi/pi-utils"; + +const utf8Decoder = new TextDecoder("utf-8"); export interface SessionStorageStat { size: number; @@ -163,7 +165,7 @@ export class FileSessionStorage implements SessionStorage { } async readTextPrefix(path: string, maxBytes: number): Promise { - return Bun.file(path).slice(0, maxBytes).text(); + return peekFile(path, maxBytes, header => utf8Decoder.decode(header)); } async writeText(path: string, content: string): Promise { diff --git a/packages/coding-agent/src/tools/auto-generated-guard.ts b/packages/coding-agent/src/tools/auto-generated-guard.ts index 648415ef3..d917e1ea4 100644 --- a/packages/coding-agent/src/tools/auto-generated-guard.ts +++ b/packages/coding-agent/src/tools/auto-generated-guard.ts @@ -4,8 +4,8 @@ * Prevents editing of files that appear to be automatically generated * by code generation tools (protoc, sqlc, buf, swagger, etc.). */ -import * as fs from "node:fs/promises"; import * as path from "node:path"; +import { LRUCache } from "lru-cache/raw"; import { settings } from "../config/settings"; import { ToolError } from "./tool-errors"; @@ -209,40 +209,19 @@ function isAutoGeneratedFileName(filePath: string): boolean { /** * Check if leading header comments contain auto-generated markers. - * Returns the matched marker text if found, null otherwise. + * Returns the matched marker text if found, undefined otherwise. */ -function detectAutoGeneratedMarker(content: string, filePath: string): string | null { +function detectAutoGeneratedMarker(content: string, filePath: string): string | undefined { const commentStyles = getCommentStylesForPath(filePath); const headerCommentText = extractLeadingHeaderCommentText(content, commentStyles); - if (!headerCommentText) return null; + if (!headerCommentText) return undefined; for (const markerPattern of AUTO_GENERATED_HEADER_MARKERS) { const match = markerPattern.exec(headerCommentText); if (match?.[0]) return match[0]; } - return null; -} - -/** - * Read the first N bytes of a file as a UTF-8 string. - * More efficient than reading the entire file. - */ -async function readFilePrefix(filePath: string, bytes: number): Promise { - const handle = await fs.open(filePath, "r").catch(() => null); - if (!handle) { - return null; - } - - try { - const buffer = Buffer.allocUnsafe(bytes); - const { bytesRead } = await handle.read(buffer, 0, bytes, 0); - return buffer.toString("utf-8", 0, bytesRead); - } catch { - return null; - } finally { - await handle.close(); - } + return undefined; } /** @@ -259,6 +238,25 @@ function buildAutoGeneratedError(displayPath: string, detected: string): ToolErr ); } +import { peekFile } from "@oh-my-pi/pi-utils"; + +const decoder = new TextDecoder("utf-8"); + +const autoGeneratedMap = new LRUCache({ max: 10 }); + +async function getAutoGeneratedMarker(filePath: string): Promise { + if (isAutoGeneratedFileName(filePath)) { + return filePath.split("/").pop() ?? ""; + } + const cached = autoGeneratedMap.get(filePath); + if (cached) return cached.marker; + + const content = await peekFile(filePath, CHECK_BYTE_COUNT, header => decoder.decode(header)); + const marker = detectAutoGeneratedMarker(content, filePath); + autoGeneratedMap.set(filePath, { marker }); + return marker; +} + /** * Check if a file is auto-generated by examining its content. * Throws ToolError if the file appears to be auto-generated. @@ -266,24 +264,12 @@ function buildAutoGeneratedError(displayPath: string, detected: string): ToolErr * @param absolutePath - Absolute path to the file * @param displayPath - Path to show in error messages (relative or as provided) */ -export async function checkAutoGeneratedFile(absolutePath: string, displayPath?: string): Promise { +export async function assertEditableFile(absolutePath: string, displayPath?: string) { if (!settings.get("edit.blockAutoGenerated")) { return; } - const pathForDisplay = displayPath ?? absolutePath; - - if (isAutoGeneratedFileName(absolutePath)) { - const fileName = absolutePath.split("/").pop() ?? ""; - throw buildAutoGeneratedError(pathForDisplay, fileName); - } - - const content = await readFilePrefix(absolutePath, CHECK_BYTE_COUNT); - if (content === null) { - return; - } - - const marker = detectAutoGeneratedMarker(content, absolutePath); + const marker = await getAutoGeneratedMarker(absolutePath); if (marker) { throw buildAutoGeneratedError(pathForDisplay, marker); } @@ -297,7 +283,7 @@ export async function checkAutoGeneratedFile(absolutePath: string, displayPath?: * @param content - File content to check (can be full content or prefix) * @param displayPath - Path to show in error messages */ -export async function checkAutoGeneratedFileContent(content: string, displayPath: string): Promise { +export function assertEditableFileContent(content: string, displayPath: string): void { if (!settings.get("edit.blockAutoGenerated")) { return; } diff --git a/packages/coding-agent/src/tools/write.ts b/packages/coding-agent/src/tools/write.ts index 1c6e181fa..71246d765 100644 --- a/packages/coding-agent/src/tools/write.ts +++ b/packages/coding-agent/src/tools/write.ts @@ -21,7 +21,7 @@ import type { ToolSession } from "../sdk"; import { Ellipsis, Hasher, type RenderCache, renderStatusLine, truncateToWidth } from "../tui"; import { resolveFileDisplayMode } from "../utils/file-display-mode"; import { parseArchivePathCandidates } from "./archive-reader"; -import { checkAutoGeneratedFile } from "./auto-generated-guard"; +import { assertEditableFile } from "./auto-generated-guard"; import { invalidateFsScanAfterWrite } from "./fs-cache-invalidation"; import { type OutputMeta, outputMeta } from "./output-meta"; import { enforcePlanModeWrite, resolvePlanPath } from "./plan-mode-guard"; @@ -297,7 +297,7 @@ export class WriteTool implements AgentTool { +): Promise<{ agent: Agent; session: AgentSession; authStorage: AuthStorage }> { const model = getBundledModel("anthropic", "claude-sonnet-4-5")!; const agent = new Agent({ getApiKey: () => "test-key", @@ -102,6 +104,7 @@ async function createSession( const modelRegistry = new ModelRegistry(authStorage, path.join(tempDir, "models.yml")); return { + agent, session: new AgentSession({ agent, sessionManager, @@ -135,6 +138,7 @@ function createStreamForDiff( path: string, chunks: string[], abortSignalRef: { current?: AbortSignal }, + streamStateRef?: { deltaCount: number }, ): Agent["streamFn"] { let callIndex = 0; return (_model, _context, options) => { @@ -175,6 +179,9 @@ function createStreamForDiff( for (const chunk of chunks) { if (aborted) return; diffSoFar += chunk; + if (streamStateRef) { + streamStateRef.deltaCount += 1; + } const partialCall = createToolCall(toolCallId, { path, diff: diffSoFar }); stream.push({ type: "toolcall_delta", @@ -265,4 +272,96 @@ describe("streaming edit abort", () => { } } }); + + it("does not abort when auto-generated peek fails with ENOENT (non-ToolError)", async () => { + const checkSpy = vi + .spyOn(autoGeneratedGuard, "checkAutoGeneratedFile") + .mockRejectedValue(Object.assign(new Error("ENOENT"), { code: "ENOENT" })); + + await Bun.write(path.join(tempDir, "sample.txt"), "alpha\nbeta\ngamma\n"); + const diff = "@@\n-beta\n+beta2\n"; + const abortSignalRef: { current?: AbortSignal } = {}; + const streamFn = createStreamForDiff("sample.txt", chunkStringRandomly(diff, 7), abortSignalRef); + const { agent, session, authStorage } = await createSession(tempDir, streamFn, editTool); + const abortSpy = vi.spyOn(agent, "abort"); + + try { + await session.prompt("apply patch"); + + expect(abortSpy).not.toHaveBeenCalled(); + expect(abortSignalRef.current?.aborted ?? false).toBe(false); + const lastAssistant = lastAssistantMessage(session.state.messages); + expect(lastAssistant?.stopReason).not.toBe("aborted"); + } finally { + checkSpy.mockRestore(); + abortSpy.mockRestore(); + try { + await session.dispose(); + } finally { + authStorage.close(); + } + } + }); + + it("aborts when auto-generated check rejects with ToolError", async () => { + const checkSpy = vi + .spyOn(autoGeneratedGuard, "checkAutoGeneratedFile") + .mockRejectedValue(new ToolError("Cannot modify auto-generated file")); + + await Bun.write(path.join(tempDir, "sample.txt"), "alpha\nbeta\ngamma\n"); + const diff = "@@\n-beta\n+beta2\n"; + const abortSignalRef: { current?: AbortSignal } = {}; + const streamFn = createStreamForDiff("sample.txt", chunkStringRandomly(diff, 7), abortSignalRef); + const { agent, session, authStorage } = await createSession(tempDir, streamFn, editTool); + const abortSpy = vi.spyOn(agent, "abort"); + + try { + await session.prompt("apply patch"); + + expect(abortSpy).toHaveBeenCalled(); + expect(abortSignalRef.current?.aborted ?? false).toBe(true); + const lastAssistant = lastAssistantMessage(session.state.messages); + expect(lastAssistant?.stopReason).toBe("aborted"); + } finally { + checkSpy.mockRestore(); + abortSpy.mockRestore(); + try { + await session.dispose(); + } finally { + authStorage.close(); + } + } + }); + + it("aborts auto-generated file edits as soon as the path is available", async () => { + const generatedPath = path.join(tempDir, "generated.ts"); + await Bun.write(generatedPath, "// Code generated by sqlc. DO NOT EDIT.\nexport const foo = 1;\n"); + const abortSignalRef: { current?: AbortSignal } = {}; + const diff = "@@\n-export const foo = 1;\n+export const foo = 2;\n"; + const streamFn = createStreamForDiff("generated.ts", chunkStringRandomly(diff, 7), abortSignalRef); + const { agent, session, authStorage } = await createSession(tempDir, streamFn, editTool); + let nonEmptyToolcallDeltaUpdates = 0; + const unsubscribe = agent.subscribe(event => { + if ( + event.type === "message_update" && + event.assistantMessageEvent.type === "toolcall_delta" && + event.assistantMessageEvent.delta.length > 0 + ) { + nonEmptyToolcallDeltaUpdates += 1; + } + }); + + try { + await session.prompt("apply patch"); + + expect(nonEmptyToolcallDeltaUpdates).toBe(2); + } finally { + unsubscribe(); + try { + await session.dispose(); + } finally { + authStorage.close(); + } + } + }); }); diff --git a/packages/coding-agent/test/tools/auto-generated-guard.test.ts b/packages/coding-agent/test/tools/auto-generated-guard.test.ts index fbe11d192..8c3acbf85 100644 --- a/packages/coding-agent/test/tools/auto-generated-guard.test.ts +++ b/packages/coding-agent/test/tools/auto-generated-guard.test.ts @@ -3,10 +3,7 @@ import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; import { _resetSettingsForTest, Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; -import { - checkAutoGeneratedFile, - checkAutoGeneratedFileContent, -} from "@oh-my-pi/pi-coding-agent/tools/auto-generated-guard"; +import { assertEditableFile, assertEditableFileContent } from "@oh-my-pi/pi-coding-agent/tools/auto-generated-guard"; import { ToolError } from "@oh-my-pi/pi-coding-agent/tools/tool-errors"; let tempDir: string; @@ -17,92 +14,92 @@ beforeAll(async () => { await Settings.init({ inMemory: true, cwd: tempDir }); }); -describe("checkAutoGeneratedFileContent", () => { +describe("assertEditableFileContent", () => { it("detects canonical TypeScript generated header", async () => { const content = "// Code generated by sqlc. DO NOT EDIT.\n\nexport const foo = 1;"; - await expect(checkAutoGeneratedFileContent(content, "test.ts")).rejects.toBeInstanceOf(ToolError); + await expect(assertEditableFileContent(content, "test.ts")).rejects.toBeInstanceOf(ToolError); }); it("detects @generated marker", async () => { const content = "// @generated\n\nexport const foo = 1;"; - await expect(checkAutoGeneratedFileContent(content, "test.ts")).rejects.toBeInstanceOf(ToolError); + await expect(assertEditableFileContent(content, "test.ts")).rejects.toBeInstanceOf(ToolError); }); it("detects generated-by marker for Python files", async () => { const content = "# Generated by buf\n\nvalue = 1"; - await expect(checkAutoGeneratedFileContent(content, "test.py")).rejects.toBeInstanceOf(ToolError); + await expect(assertEditableFileContent(content, "test.py")).rejects.toBeInstanceOf(ToolError); }); it("detects generated-by marker for SQL files", async () => { const content = "-- generated by sqlc\n\nselect 1;"; - await expect(checkAutoGeneratedFileContent(content, "query.sql")).rejects.toBeInstanceOf(ToolError); + await expect(assertEditableFileContent(content, "query.sql")).rejects.toBeInstanceOf(ToolError); }); it("detects generated markers in leading block comments", async () => { const content = "/*\n * Code generated by mockery. DO NOT EDIT.\n */\nexport const foo = 1;"; - await expect(checkAutoGeneratedFileContent(content, "test.ts")).rejects.toBeInstanceOf(ToolError); + await expect(assertEditableFileContent(content, "test.ts")).rejects.toBeInstanceOf(ToolError); }); it("detects kysely-codegen generated header", async () => { const content = "/**\n * This file was generated by kysely-codegen.\n * Please do not edit it manually.\n */\n\nexport interface Database {}"; - await expect(checkAutoGeneratedFileContent(content, "db.ts")).rejects.toBeInstanceOf(ToolError); + await expect(assertEditableFileContent(content, "db.ts")).rejects.toBeInstanceOf(ToolError); }); it("does not block broad prose comment markers", async () => { const content = "// auto generated dont edit bla bla\n// this is a hand-written file note\nexport const foo = 1;"; - await expect(checkAutoGeneratedFileContent(content, "test.ts")).resolves.toBeUndefined(); + await expect(assertEditableFileContent(content, "test.ts")).resolves.toBeUndefined(); }); it("does not match generated markers after code starts", async () => { const content = "export const foo = 1;\n\n// Code generated by sqlc. DO NOT EDIT."; - await expect(checkAutoGeneratedFileContent(content, "test.ts")).resolves.toBeUndefined(); + await expect(assertEditableFileContent(content, "test.ts")).resolves.toBeUndefined(); }); it("uses language-specific comment styles", async () => { const tsContent = "# Code generated by sqlc. DO NOT EDIT.\nexport const foo = 1;"; - await expect(checkAutoGeneratedFileContent(tsContent, "test.ts")).resolves.toBeUndefined(); + await expect(assertEditableFileContent(tsContent, "test.ts")).resolves.toBeUndefined(); const pyContent = "// Code generated by sqlc. DO NOT EDIT.\nvalue = 1"; - await expect(checkAutoGeneratedFileContent(pyContent, "test.py")).resolves.toBeUndefined(); + await expect(assertEditableFileContent(pyContent, "test.py")).resolves.toBeUndefined(); }); it("does not block editing the guard file itself", async () => { const guardPath = path.join(import.meta.dir, "../../src/tools/auto-generated-guard.ts"); const content = await Bun.file(guardPath).text(); await expect( - checkAutoGeneratedFileContent(content, "packages/coding-agent/src/tools/auto-generated-guard.ts"), + assertEditableFileContent(content, "packages/coding-agent/src/tools/auto-generated-guard.ts"), ).resolves.toBeUndefined(); }); it("checks only first 1024 bytes of content", async () => { const prefix = "A".repeat(1024); const content = `${prefix}\n// Code generated by sqlc. DO NOT EDIT.`; - await expect(checkAutoGeneratedFileContent(content, "test.ts")).resolves.toBeUndefined(); + await expect(assertEditableFileContent(content, "test.ts")).resolves.toBeUndefined(); }); }); -describe("checkAutoGeneratedFile", () => { +describe("assertEditableFile", () => { it("detects auto-generated filename patterns", async () => { const filePath = path.join(tempDir, "zz_generated.deepcopy.go"); await Bun.write(filePath, "package generated"); - await expect(checkAutoGeneratedFile(filePath)).rejects.toBeInstanceOf(ToolError); + await expect(assertEditableFile(filePath)).rejects.toBeInstanceOf(ToolError); }); it("detects content marker from file prefix", async () => { const filePath = path.join(tempDir, "service.ts"); await Bun.write(filePath, "// Code generated by sqlc. DO NOT EDIT.\nexport const foo = 1;"); - await expect(checkAutoGeneratedFile(filePath)).rejects.toBeInstanceOf(ToolError); + await expect(assertEditableFile(filePath)).rejects.toBeInstanceOf(ToolError); }); it("allows normal files", async () => { const filePath = path.join(tempDir, "normal.ts"); await Bun.write(filePath, "// Regular source file\nexport const foo = 1;"); - await expect(checkAutoGeneratedFile(filePath)).resolves.toBeUndefined(); + await expect(assertEditableFile(filePath)).resolves.toBeUndefined(); }); it("handles missing files gracefully", async () => { const filePath = path.join(tempDir, "does-not-exist.ts"); - await expect(checkAutoGeneratedFile(filePath)).resolves.toBeUndefined(); + await expect(assertEditableFile(filePath)).resolves.toBeUndefined(); }); }); diff --git a/packages/utils/src/peek-file.ts b/packages/utils/src/peek-file.ts index c5cbcd9c2..ec647dc20 100644 --- a/packages/utils/src/peek-file.ts +++ b/packages/utils/src/peek-file.ts @@ -1,7 +1,17 @@ +/** + * Read the first `maxBytes` of a file (offset 0) and pass that slice to `op`. + * + * Buffers are reused to avoid allocating on every peek: sync uses one growable + * `Uint8Array`; async uses a small fixed pool of `Buffer`s with a bounded wait + * queue, falling back to a fresh allocation when the pool and queue are saturated + * or when `maxBytes` exceeds the pool slot size. + */ import * as fs from "node:fs"; +/** Async pool slot size; larger peeks allocate ad hoc. */ const POOLED_BUFFER_SIZE = 512; const ASYNC_POOL_SIZE = 10; +/** Cap waiter queue so heavy concurrency does not queue unbounded; overflow uses alloc. */ const MAX_ASYNC_WAITERS = 4; const INITIAL_SYNC_BUFFER_SIZE = 1024; const EMPTY_BUFFER = Buffer.alloc(0); @@ -11,6 +21,7 @@ const availableAsyncPoolIndexes = Array.from({ length: ASYNC_POOL_SIZE }, (_, in const asyncPoolWaiters: Array<(index: number) => void> = []; let syncPool = new Uint8Array(INITIAL_SYNC_BUFFER_SIZE); +/** Returns a pool slot index, or `-1` when the caller should use a standalone buffer. */ function acquireAsyncPoolIndex(): Promise | number { const index = availableAsyncPoolIndexes.pop(); if (index !== undefined) { @@ -63,6 +74,10 @@ function withSyncPoolBuffer(maxBytes: number, op: (buffer: Uint8Array) => T): return op(syncPool.subarray(0, maxBytes)); } +/** + * Synchronously reads up to `maxBytes` from the start of `filePath` and returns `op(header)`. + * If the file is shorter, `header` is only the bytes actually read. + */ export function peekFileSync(filePath: string, maxBytes: number, op: (header: Uint8Array) => T): T { if (maxBytes <= 0) { return op(EMPTY_BUFFER); @@ -79,6 +94,9 @@ export function peekFileSync(filePath: string, maxBytes: number, op: (header: } } +/** + * Like {@link peekFileSync} but uses async I/O. + */ export async function peekFile(filePath: string, maxBytes: number, op: (header: Uint8Array) => T): Promise { if (maxBytes <= 0) { return op(EMPTY_BUFFER);