feat: added assistant message event interception for streaming edit abort control

- Added `onAssistantMessageEvent` callback to Agent API for inspecting and aborting assistant streaming events.
- Added `setAssistantMessageEventInterceptor()` method to dynamically update assistant message event handlers.
- Converted `checkAutoGeneratedFileContent()` from async to synchronous for improved streaming edit abort detection performance.
- Implemented LRU caching in auto-generated file detection with early path-based checks to prevent unnecessary edits.
- Refactored streaming edit pre-caching to use assistant message event interception for real-time abort capability.
- Extracted `peekFile()` utility for efficient file prefix reading with pooled buffer reuse strategy.
This commit is contained in:
can1357
2026-04-08 14:44:06 +02:00
parent 211e5369a0
commit 757bdca7a1
16 changed files with 337 additions and 124 deletions
+8 -12
View File
@@ -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!(
+5 -1
View File
@@ -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).
+4
View File
@@ -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,
+15 -1
View File
@@ -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) {
+10 -2
View File
@@ -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<string, unknown>, toolName: string) => Record<string, unknown>;
/**
* 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`.
+8
View File
@@ -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
@@ -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 {
@@ -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);
@@ -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);
@@ -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<string, number>();
#streamingEditPrecheckedToolCallIds = new Set<string>();
#streamingEditFileCache = new Map<string, string>();
#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<void> {
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);
@@ -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<string> {
return Bun.file(path).slice(0, maxBytes).text();
return peekFile(path, maxBytes, header => utf8Decoder.decode(header));
}
async writeText(path: string, content: string): Promise<void> {
@@ -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<string | null> {
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<string, { marker: string | undefined }>({ max: 10 });
async function getAutoGeneratedMarker(filePath: string): Promise<string | undefined> {
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<void> {
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<void> {
export function assertEditableFileContent(content: string, displayPath: string): void {
if (!settings.get("edit.blockAutoGenerated")) {
return;
}
+2 -2
View File
@@ -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<typeof writeSchema, WriteToolDetails
// Check if file exists and is auto-generated before overwriting
if (await fs.exists(absolutePath)) {
await checkAutoGeneratedFile(absolutePath, path);
await assertEditableFile(absolutePath, path);
}
const diagnostics = await this.#writethrough(absolutePath, cleanContent, signal, undefined, batchRequest);
@@ -2,7 +2,7 @@
* Streaming edit abort tests.
*/
import { afterEach, beforeEach, describe, expect, it } from "bun:test";
import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test";
import * as fs from "node:fs";
import * as os from "node:os";
import * as path from "node:path";
@@ -14,6 +14,8 @@ import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings";
import { AgentSession } from "@oh-my-pi/pi-coding-agent/session/agent-session";
import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage";
import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager";
import * as autoGeneratedGuard from "@oh-my-pi/pi-coding-agent/tools/auto-generated-guard";
import { ToolError } from "@oh-my-pi/pi-coding-agent/tools/tool-errors";
import { Snowflake } from "@oh-my-pi/pi-utils";
import { Type } from "@sinclair/typebox";
@@ -83,7 +85,7 @@ async function createSession(
tempDir: string,
streamFn: Agent["streamFn"],
tool: AgentTool,
): Promise<{ session: AgentSession; authStorage: AuthStorage }> {
): 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();
}
}
});
});
@@ -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();
});
});
+18
View File
@@ -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> | number {
const index = availableAsyncPoolIndexes.pop();
if (index !== undefined) {
@@ -63,6 +74,10 @@ function withSyncPoolBuffer<T>(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<T>(filePath: string, maxBytes: number, op: (header: Uint8Array) => T): T {
if (maxBytes <= 0) {
return op(EMPTY_BUFFER);
@@ -79,6 +94,9 @@ export function peekFileSync<T>(filePath: string, maxBytes: number, op: (header:
}
}
/**
* Like {@link peekFileSync} but uses async I/O.
*/
export async function peekFile<T>(filePath: string, maxBytes: number, op: (header: Uint8Array) => T): Promise<T> {
if (maxBytes <= 0) {
return op(EMPTY_BUFFER);