Files
oh-my-pi/packages/coding-agent/src/session/stream-guards.ts
T
can1357 7eeaba0471 refactor(coding-agent/session): restructured monolithic agent session
- Extracted internal handlers and logic from AgentSession into dedicated runner, guard, and coordinator modules.
- Created standalone modules for bash execution, evaluation runners, IRC bridging, and prewalk coordination.
- Established dedicated session components for tracking stats, todos, streams, and retry fallback chains.
- Preserved existing session behavior while significantly reducing monolithic class size and complexity.
2026-07-24 01:24:44 +02:00

418 lines
16 KiB
TypeScript

import * as fs from "node:fs";
import type { Agent, AgentEvent, AgentMessage, AgentTurnEndContext } from "@oh-my-pi/pi-agent-core";
import type { AssistantMessage, AssistantMessageEvent, Model, ToolCall } from "@oh-my-pi/pi-ai";
import { GeminiHeaderRunDetector, isGeminiThinkingModel } from "@oh-my-pi/pi-ai/utils/thinking-loop";
import { type RepeatedToolCallDetection, ToolCallLoopGuard } from "@oh-my-pi/pi-ai/utils/tool-call-loop-guard";
import { isEnoent, logger, prompt } from "@oh-my-pi/pi-utils";
import type { Settings } from "../config/settings";
import { normalizeDiff, normalizeToLF, ParseError, previewPatch, stripBom } from "../edit";
import { type LocalProtocolOptions, resolveLocalUrlToPath } from "../internal-urls";
import geminiToolReminderTemplate from "../prompts/system/gemini-tool-call-reminder.md" with { type: "text" };
import toolCallLoopRedirectTemplate from "../prompts/system/tool-call-loop-redirect.md" with { type: "text" };
import type { SecretObfuscator } from "../secrets/obfuscator";
import { assertEditableFile } from "../tools/auto-generated-guard";
import { isInternalUrlPath, normalizeLocalScheme, resolveToCwd } from "../tools/path-utils";
import { ToolError } from "../tools/tool-errors";
import type { CustomMessage } from "./messages";
import type { SessionManager } from "./session-manager";
const GEMINI_HEADER_INTERRUPT_REASON = "Interrupted: emit a tool call instead of more planning";
const GEMINI_TOOL_REMINDER_TYPE = "gemini-tool-call-reminder";
const TOOL_CALL_LOOP_REDIRECT_TYPE = "tool-call-loop-redirect";
/** Capabilities borrowed by the session's streaming and loop guards. */
export interface StreamGuardsHost {
agent: Agent;
settings: Settings;
sessionManager: SessionManager;
obfuscator: SecretObfuscator | undefined;
model(): Model | undefined;
isDisposed(): boolean;
promptGeneration(): number;
localProtocolOptions(): LocalProtocolOptions;
emitNotice(level: "info" | "warning" | "error", message: string, source?: string): void;
schedulePostPromptTask(task: (signal: AbortSignal) => Promise<void>): void;
discardAssistantTurn(message: AssistantMessage): void;
}
/** Guards streamed edit calls against generated files and invalid patch previews. */
export class StreamingEditGuard {
readonly #host: StreamGuardsHost;
#abortTriggered = false;
#checkedLineCounts = new Map<string, number>();
#precheckedToolCallIds = new Set<string>();
#fileCache = new Map<string, string>();
#lastToolCallId: string | undefined;
constructor(host: StreamGuardsHost) {
this.#host = host;
}
/** Whether the current turn was aborted by streaming edit validation. */
get abortTriggered(): boolean {
return this.#abortTriggered;
}
/** Clears all turn-scoped streaming edit state. */
reset(): void {
this.#abortTriggered = false;
this.#checkedLineCounts.clear();
this.#precheckedToolCallIds.clear();
this.#fileCache.clear();
}
/** Pre-caches and validates a streamed edit as its arguments arrive. */
preCache(event: AgentEvent): void {
if (this.#abortTriggered || 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.#getToolCall(event);
if (!streamingEdit) return;
// The auto-generated guard runs unconditionally: editing a generated file
// is never the user's intent, and the cost of a false-positive abort is one
// wasted turn vs. silently corrupting a regenerated source.
const shouldCheckAutoGenerated =
!streamingEdit.toolCall.id || !this.#precheckedToolCallIds.has(streamingEdit.toolCall.id);
if (shouldCheckAutoGenerated) {
if (streamingEdit.toolCall.id) this.#precheckedToolCallIds.add(streamingEdit.toolCall.id);
this.#abortForAutoGeneratedPath(streamingEdit.toolCall, streamingEdit.path, streamingEdit.resolvedPath);
}
// File-cache priming feeds maybeAbort's removed-lines check, which is the
// optional patch-preview verification gated by edit.streamingAbort.
if (this.#host.settings.get("edit.streamingAbort")) this.#ensureFileCache(streamingEdit.resolvedPath);
}
/** Invalidates cached source text after an edit tool result lands. */
invalidate(filePath: string): void {
const resolvedPath = this.#resolveSessionFsPath(filePath);
if (resolvedPath !== undefined) this.#fileCache.delete(resolvedPath);
}
/** Aborts a streamed edit whose completed patch preview cannot apply. */
maybeAbort(event: AgentEvent): void {
if (!this.#host.settings.get("edit.streamingAbort") || this.#abortTriggered || event.type !== "message_update") {
return;
}
const assistantEvent = event.assistantMessageEvent;
if (assistantEvent.type !== "toolcall_end" && assistantEvent.type !== "toolcall_delta") return;
const streamingEdit = this.#getToolCall(event);
if (!streamingEdit?.toolCall.id) return;
const { toolCall, path, resolvedPath, diff, op, rename } = streamingEdit;
if (!diff || (op && op !== "update") || !diff.includes("\n")) return;
const lastNewlineIndex = diff.lastIndexOf("\n");
if (lastNewlineIndex < 0) return;
const diffForCheck = diff.endsWith("\n") ? diff : diff.slice(0, lastNewlineIndex + 1);
if (diffForCheck.trim().length === 0) return;
let normalizedDiff = normalizeDiff(diffForCheck.replace(/\r/g, ""));
if (!normalizedDiff) return;
if (this.#host.obfuscator) normalizedDiff = this.#host.obfuscator.deobfuscate(normalizedDiff);
if (!normalizedDiff) return;
const lines = normalizedDiff.split("\n");
if (!lines.some(line => line.startsWith("+") || line.startsWith("-"))) return;
const lineCount = lines.length;
const lastChecked = this.#checkedLineCounts.get(toolCall.id);
if (lastChecked !== undefined && lineCount <= lastChecked) return;
this.#checkedLineCounts.set(toolCall.id, lineCount);
const removedLines = lines
.filter(line => line.startsWith("-") && !line.startsWith("--- "))
.map(line => line.slice(1));
if (removedLines.length > 0) {
let cachedContent = this.#fileCache.get(resolvedPath);
if (cachedContent === undefined) {
this.#ensureFileCache(resolvedPath);
cachedContent = this.#fileCache.get(resolvedPath);
}
if (cachedContent !== undefined) {
const missing = removedLines.find(line => !cachedContent.includes(normalizeToLF(line)));
if (missing) this.#abortPatch(toolCall.id, path, `Failed to find expected lines in ${path}:\n${missing}`);
return;
}
if (assistantEvent.type === "toolcall_delta") return;
void this.#checkRemovedLines(toolCall.id, path, resolvedPath, removedLines);
return;
}
if (assistantEvent.type === "toolcall_delta") return;
void this.#checkPreviewPatch(toolCall.id, path, rename, normalizedDiff);
}
#getToolCall(event: AgentEvent):
| {
toolCall: ToolCall;
path: string;
resolvedPath: string;
diff?: string;
op?: string;
rename?: string;
}
| undefined {
if (event.type !== "message_update" || event.message.role !== "assistant") return undefined;
const contentIndex = event.assistantMessageEvent.contentIndex ?? 0;
const messageContent = event.message.content;
if (!Array.isArray(messageContent) || contentIndex < 0 || contentIndex >= messageContent.length) return undefined;
const toolCall = messageContent[contentIndex] as ToolCall;
if (toolCall.name !== "edit") return undefined;
const args = toolCall.arguments;
if (!args || typeof args !== "object" || Array.isArray(args) || "old_text" in args || "new_text" in args) {
return undefined;
}
const filePath = typeof args.path === "string" ? args.path : undefined;
if (!filePath) return undefined;
// local:// URLs resolve to artifacts; other internal URLs have no local path.
const resolvedPath = this.#resolveSessionFsPath(filePath);
if (resolvedPath === undefined) return undefined;
return {
toolCall,
path: filePath,
resolvedPath,
diff: typeof args.diff === "string" ? args.diff : undefined,
op: typeof args.op === "string" ? args.op : undefined,
rename: typeof args.rename === "string" ? args.rename : undefined,
};
}
#abortForAutoGeneratedPath(toolCall: ToolCall, filePath: string, resolvedPath: string): void {
if (this.#lastToolCallId === toolCall.id) return;
this.#lastToolCallId = toolCall.id;
void assertEditableFile(resolvedPath, filePath).catch(error => {
if (!(error instanceof ToolError) || this.#lastToolCallId !== toolCall.id) return;
if (!this.#abortTriggered) {
this.#abortTriggered = true;
logger.warn("Streaming edit aborted due to auto-generated file guard", {
toolCallId: toolCall.id,
path: filePath,
});
this.#host.agent.abort();
}
});
}
#ensureFileCache(resolvedPath: string): void {
if (this.#fileCache.has(resolvedPath)) return;
try {
const rawText = fs.readFileSync(resolvedPath, "utf-8");
const { text } = stripBom(rawText);
this.#fileCache.set(resolvedPath, normalizeToLF(text));
} catch {
// Read errors are handled by the edit tool itself.
}
}
#resolveSessionFsPath(filePath: string): string | undefined {
const normalized = normalizeLocalScheme(filePath);
if (normalized.startsWith("local:")) {
return resolveLocalUrlToPath(normalized, this.#host.localProtocolOptions());
}
if (isInternalUrlPath(normalized)) return undefined;
return resolveToCwd(normalized, this.#host.sessionManager.getCwd());
}
async #checkRemovedLines(
toolCallId: string,
filePath: string,
resolvedPath: string,
removedLines: string[],
): Promise<void> {
if (this.#abortTriggered) return;
try {
const { text } = stripBom(await Bun.file(resolvedPath).text());
const normalizedContent = normalizeToLF(text);
const missing = removedLines.find(line => !normalizedContent.includes(normalizeToLF(line)));
if (missing)
this.#abortPatch(toolCallId, filePath, `Failed to find expected lines in ${filePath}:\n${missing}`);
} catch (error) {
if (!isEnoent(error)) {
// Unexpected fallback read errors remain non-fatal.
}
}
}
async #checkPreviewPatch(
toolCallId: string,
filePath: string,
rename: string | undefined,
normalizedDiff: string,
): Promise<void> {
if (this.#abortTriggered) return;
try {
await previewPatch(
{ path: filePath, op: "update", rename, diff: normalizedDiff },
{
cwd: this.#host.sessionManager.getCwd(),
allowFuzzy: this.#host.settings.get("edit.fuzzyMatch"),
fuzzyThreshold: this.#host.settings.get("edit.fuzzyThreshold"),
},
);
} catch (error) {
if (error instanceof ParseError) return;
this.#abortPatch(toolCallId, filePath, error instanceof Error ? error.message : String(error));
}
}
#abortPatch(toolCallId: string, filePath: string, error: string): void {
this.#abortTriggered = true;
logger.warn("Streaming edit aborted due to patch preview failure", { toolCallId, path: filePath, error });
this.#host.agent.abort();
}
}
/** Detects cross-turn tool loops and Gemini reasoning-header runaways. */
export class LoopGuards {
readonly #host: StreamGuardsHost;
#geminiHeaderDetector: GeminiHeaderRunDetector | undefined;
#toolCallLoopGuard: ToolCallLoopGuard | undefined;
#toolCallLoopGuardSettingsKey: string | undefined;
constructor(host: StreamGuardsHost) {
this.#host = host;
}
/** Records a completed turn and injects a redirect when calls repeat. */
recordTurn(messages: AgentMessage[], context: AgentTurnEndContext | undefined): void {
if (context?.message.role !== "assistant") return;
const detection = this.#activeToolCallLoopGuard()?.recordTurn({
message: context.message,
toolResults: context.toolResults,
});
if (detection) this.#injectToolCallLoopRedirect(messages, detection);
}
/** Feeds a streamed assistant event to the Gemini header-runaway detector. */
onAssistantEvent(message: AssistantMessage, event: AssistantMessageEvent): void {
if (event.type === "thinking_start") {
this.#geminiHeaderDetector = this.#geminiHeaderGuardActive() ? new GeminiHeaderRunDetector() : undefined;
return;
}
const detector = this.#geminiHeaderDetector;
if (!detector) return;
if (event.type === "thinking_delta") {
if (detector.push(event.delta)) this.#interruptGeminiHeaderRunaway(detector.count, message.timestamp);
return;
}
if (event.type === "text_start" || event.type === "toolcall_start") detector.reset();
}
#activeToolCallLoopGuard(): ToolCallLoopGuard | undefined {
if (this.#host.settings.get("model.toolCallLoopGuard.enabled") !== true) {
this.#toolCallLoopGuard = undefined;
this.#toolCallLoopGuardSettingsKey = undefined;
return undefined;
}
const threshold = this.#host.settings.get("model.toolCallLoopGuard.threshold");
const exemptTools = this.#host.settings
.get("model.toolCallLoopGuard.exemptTools")
.filter((tool): tool is string => typeof tool === "string" && tool.length > 0);
const settingsKey = `${threshold}:${JSON.stringify(exemptTools)}`;
if (!this.#toolCallLoopGuard || this.#toolCallLoopGuardSettingsKey !== settingsKey) {
this.#toolCallLoopGuard = new ToolCallLoopGuard({ threshold, exemptTools });
this.#toolCallLoopGuardSettingsKey = settingsKey;
}
return this.#toolCallLoopGuard;
}
#injectToolCallLoopRedirect(messages: AgentMessage[], detection: RepeatedToolCallDetection): void {
const content = prompt.render(toolCallLoopRedirectTemplate, {
tool_name: detection.toolName,
count: detection.count,
arguments_summary: detection.argumentsSummary,
result_summary: detection.resultSummary || "(no text result)",
});
const details = {
toolName: detection.toolName,
count: detection.count,
argumentsSummary: detection.argumentsSummary,
resultSummary: detection.resultSummary,
};
logger.warn("cross-turn tool-call loop detected", { toolName: detection.toolName, count: detection.count });
const redirectMessage: CustomMessage = {
role: "custom",
customType: TOOL_CALL_LOOP_REDIRECT_TYPE,
content,
display: false,
details,
attribution: "agent",
timestamp: Date.now(),
};
messages.push(redirectMessage);
if (this.#host.agent.state.messages !== messages) this.#host.agent.appendMessage(redirectMessage);
this.#host.sessionManager.appendCustomMessageEntry(
TOOL_CALL_LOOP_REDIRECT_TYPE,
content,
false,
details,
"agent",
);
}
#geminiHeaderGuardActive(): boolean {
const model = this.#host.model();
return (
process.env.PI_NO_THINKING_LOOP_GUARD !== "1" &&
this.#host.settings.get("model.loopGuard.enabled") === true &&
this.#host.settings.get("model.loopGuard.toolCallReminder") === true &&
model !== undefined &&
isGeminiThinkingModel(model)
);
}
#interruptGeminiHeaderRunaway(headerCount: number, targetTimestamp: number): void {
const model = this.#host.model();
logger.warn("Gemini reasoning-header runaway; interrupting to require a tool call", {
model: model?.id,
provider: model?.provider,
headers: headerCount,
});
this.#host.emitNotice(
"warning",
`Interrupted ${headerCount} planning headers with no tool call; reminded the model to issue one.`,
"loop-guard",
);
this.#host.agent.abort(GEMINI_HEADER_INTERRUPT_REASON);
const generation = this.#host.promptGeneration();
this.#host.schedulePostPromptTask(async signal => {
if (signal.aborted || this.#host.isDisposed() || this.#host.promptGeneration() !== generation) return;
await this.#host.agent.waitForIdle();
if (signal.aborted || this.#host.isDisposed() || this.#host.promptGeneration() !== generation) return;
const aborted = this.#host.agent.state.messages.findLast(
(message): message is AssistantMessage =>
message.role === "assistant" && message.timestamp === targetTimestamp,
);
if (aborted) this.#host.discardAssistantTurn(aborted);
const content = prompt.render(geminiToolReminderTemplate, { count: headerCount });
const details = { headers: headerCount };
this.#host.agent.appendMessage({
role: "custom",
customType: GEMINI_TOOL_REMINDER_TYPE,
content,
display: false,
details,
attribution: "agent",
timestamp: Date.now(),
});
this.#host.sessionManager.appendCustomMessageEntry(
GEMINI_TOOL_REMINDER_TYPE,
content,
false,
details,
"agent",
);
try {
await this.#host.agent.continue();
} catch (error) {
logger.warn("gemini tool-call reminder continue failed", { error: String(error) });
}
});
}
}