feat(coding-agent): added ToolCallContext interface for batch execution metadata
- Added ToolCallContext interface with batch execution metadata including batchId, index, and total. - Enhanced getToolContext() to accept optional ToolCallContext parameter for batch-aware tool execution. - Implemented LSP batching support to coalesce formatting and diagnostics operations on parallel edits. - Added comprehensive tests for tool call batch context and LSP writethrough batching functionality.
This commit is contained in:
@@ -1,6 +1,9 @@
|
||||
# Changelog
|
||||
|
||||
## [Unreleased]
|
||||
### Added
|
||||
|
||||
- Enhanced getToolContext to receive tool call batch information including batchId, index, total count, and tool call details
|
||||
|
||||
## [5.6.7] - 2026-01-18
|
||||
### Fixed
|
||||
|
||||
@@ -377,6 +377,8 @@ async function executeToolCalls(
|
||||
const results: ToolResultMessage[] = [];
|
||||
let steeringMessages: AgentMessage[] | undefined;
|
||||
const shouldInterruptImmediately = interruptMode !== "wait";
|
||||
const toolCallInfos = toolCalls.map((call) => ({ id: call.id, name: call.name }));
|
||||
const batchId = `${assistantMessage.timestamp ?? Date.now()}_${toolCalls[0]?.id ?? "batch"}`;
|
||||
|
||||
for (let index = 0; index < toolCalls.length; index++) {
|
||||
const toolCall = toolCalls[index];
|
||||
@@ -397,7 +399,14 @@ async function executeToolCalls(
|
||||
|
||||
const validatedArgs = validateToolArguments(tool, toolCall);
|
||||
|
||||
const toolContext = getToolContext ? getToolContext() : undefined;
|
||||
const toolContext = getToolContext
|
||||
? getToolContext({
|
||||
batchId,
|
||||
index,
|
||||
total: toolCalls.length,
|
||||
toolCalls: toolCallInfos,
|
||||
})
|
||||
: undefined;
|
||||
result = await tool.execute(
|
||||
toolCall.id,
|
||||
validatedArgs,
|
||||
|
||||
@@ -27,6 +27,7 @@ import type {
|
||||
AgentToolContext,
|
||||
StreamFn,
|
||||
ThinkingLevel,
|
||||
ToolCallContext,
|
||||
} from "./types";
|
||||
|
||||
/**
|
||||
@@ -94,7 +95,7 @@ export interface AgentOptions {
|
||||
* Provides tool execution context, resolved per tool call.
|
||||
* Use for late-bound UI or session state access.
|
||||
*/
|
||||
getToolContext?: () => AgentToolContext | undefined;
|
||||
getToolContext?: (toolCall?: ToolCallContext) => AgentToolContext | undefined;
|
||||
|
||||
/**
|
||||
* Cursor exec handlers for local tool execution.
|
||||
@@ -139,7 +140,7 @@ export class Agent {
|
||||
private _sessionId?: string;
|
||||
private _thinkingBudgets?: ThinkingBudgets;
|
||||
public getApiKey?: (provider: string) => Promise<string | undefined> | string | undefined;
|
||||
private getToolContext?: () => AgentToolContext | undefined;
|
||||
private getToolContext?: (toolCall?: ToolCallContext) => AgentToolContext | undefined;
|
||||
private cursorExecHandlers?: CursorExecHandlers;
|
||||
private cursorOnToolResult?: CursorToolResultHandler;
|
||||
private runningPrompt?: Promise<void>;
|
||||
|
||||
@@ -109,7 +109,14 @@ export interface AgentLoopConfig extends SimpleStreamOptions {
|
||||
* Provides tool execution context, resolved per tool call.
|
||||
* Use for late-bound UI or session state access.
|
||||
*/
|
||||
getToolContext?: () => AgentToolContext | undefined;
|
||||
getToolContext?: (toolCall?: ToolCallContext) => AgentToolContext | undefined;
|
||||
}
|
||||
|
||||
export interface ToolCallContext {
|
||||
batchId: string;
|
||||
index: number;
|
||||
total: number;
|
||||
toolCalls: Array<{ id: string; name: string }>;
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -9,7 +9,15 @@ import {
|
||||
import { Type } from "@sinclair/typebox";
|
||||
import { describe, expect, it } from "vitest";
|
||||
import { agentLoop, agentLoopContinue } from "../src/agent-loop";
|
||||
import type { AgentContext, AgentEvent, AgentLoopConfig, AgentMessage, AgentTool } from "../src/types";
|
||||
import type {
|
||||
AgentContext,
|
||||
AgentEvent,
|
||||
AgentLoopConfig,
|
||||
AgentMessage,
|
||||
AgentTool,
|
||||
AgentToolContext,
|
||||
ToolCallContext,
|
||||
} from "../src/types";
|
||||
|
||||
// Mock stream for testing - mimics MockAssistantStream
|
||||
class MockAssistantStream extends EventStream<AssistantMessageEvent, AssistantMessage> {
|
||||
@@ -236,6 +244,79 @@ describe("agentLoop with AgentMessage", () => {
|
||||
expect(convertedMessages.length).toBe(2);
|
||||
});
|
||||
|
||||
it("provides tool call batch context", async () => {
|
||||
const toolSchema = Type.Object({ value: Type.String() });
|
||||
const contexts: ToolCallContext[] = [];
|
||||
const tool: AgentTool<typeof toolSchema, { value: string }> = {
|
||||
name: "echo",
|
||||
label: "Echo",
|
||||
description: "Echo tool",
|
||||
parameters: toolSchema,
|
||||
async execute(_toolCallId, params, _signal, _onUpdate, ctx) {
|
||||
const toolCall = (ctx as { toolCall?: ToolCallContext })?.toolCall;
|
||||
if (toolCall) {
|
||||
contexts.push(toolCall);
|
||||
}
|
||||
return {
|
||||
content: [{ type: "text", text: `echoed: ${params.value}` }],
|
||||
details: { value: params.value },
|
||||
};
|
||||
},
|
||||
};
|
||||
|
||||
const context: AgentContext = {
|
||||
systemPrompt: "",
|
||||
messages: [],
|
||||
tools: [tool],
|
||||
};
|
||||
|
||||
const userPrompt: AgentMessage = createUserMessage("echo something");
|
||||
|
||||
const config: AgentLoopConfig = {
|
||||
model: createModel(),
|
||||
convertToLlm: identityConverter,
|
||||
getToolContext: (toolCall) => ({ toolCall }) as AgentToolContext,
|
||||
};
|
||||
|
||||
let callIndex = 0;
|
||||
const streamFn = () => {
|
||||
const stream = new MockAssistantStream();
|
||||
queueMicrotask(() => {
|
||||
if (callIndex === 0) {
|
||||
const message = createAssistantMessage(
|
||||
[
|
||||
{ type: "toolCall", id: "tool-1", name: "echo", arguments: { value: "hello" } },
|
||||
{ type: "toolCall", id: "tool-2", name: "echo", arguments: { value: "world" } },
|
||||
],
|
||||
"toolUse",
|
||||
);
|
||||
stream.push({ type: "done", reason: "toolUse", message });
|
||||
} else {
|
||||
const message = createAssistantMessage([{ type: "text", text: "done" }]);
|
||||
stream.push({ type: "done", reason: "stop", message });
|
||||
}
|
||||
callIndex++;
|
||||
});
|
||||
return stream;
|
||||
};
|
||||
|
||||
const stream = agentLoop([userPrompt], context, config, undefined, streamFn);
|
||||
|
||||
for await (const _ of stream) {
|
||||
// consume
|
||||
}
|
||||
|
||||
expect(contexts).toHaveLength(2);
|
||||
expect(contexts[0]?.batchId).toBe(contexts[1]?.batchId);
|
||||
expect(contexts[0]?.total).toBe(2);
|
||||
expect(contexts[0]?.toolCalls).toEqual([
|
||||
{ id: "tool-1", name: "echo" },
|
||||
{ id: "tool-2", name: "echo" },
|
||||
]);
|
||||
expect(contexts[0]?.index).toBe(0);
|
||||
expect(contexts[1]?.index).toBe(1);
|
||||
});
|
||||
|
||||
it("should handle tool calls and results", async () => {
|
||||
const toolSchema = Type.Object({ value: Type.String() });
|
||||
const executed: string[] = [];
|
||||
|
||||
@@ -1,6 +1,14 @@
|
||||
# Changelog
|
||||
|
||||
## [Unreleased]
|
||||
### Changed
|
||||
|
||||
- Improved LSP batching to coalesce formatting and diagnostics for parallel edits
|
||||
- Updated edit and write tools to support batched LSP operations
|
||||
|
||||
### Fixed
|
||||
|
||||
- Coalesced LSP formatting/diagnostics for parallel edits so only the final write triggers LSP across touched files
|
||||
|
||||
## [6.1.0] - 2026-01-19
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import type { AgentToolContext } from "@oh-my-pi/pi-agent-core";
|
||||
import type { AgentToolContext, ToolCallContext } from "@oh-my-pi/pi-agent-core";
|
||||
import type { CustomToolContext } from "../custom-tools/types";
|
||||
import type { ExtensionUIContext } from "../extensions/types";
|
||||
|
||||
@@ -7,11 +7,12 @@ declare module "@oh-my-pi/pi-agent-core" {
|
||||
ui?: ExtensionUIContext;
|
||||
hasUI?: boolean;
|
||||
toolNames?: string[];
|
||||
toolCall?: ToolCallContext;
|
||||
}
|
||||
}
|
||||
|
||||
export interface ToolContextStore {
|
||||
getContext(): AgentToolContext;
|
||||
getContext(toolCall?: ToolCallContext): AgentToolContext;
|
||||
setUIContext(uiContext: ExtensionUIContext, hasUI: boolean): void;
|
||||
setToolNames(names: string[]): void;
|
||||
}
|
||||
@@ -22,11 +23,12 @@ export function createToolContextStore(getBaseContext: () => CustomToolContext):
|
||||
let toolNames: string[] = [];
|
||||
|
||||
return {
|
||||
getContext: () => ({
|
||||
getContext: (toolCall) => ({
|
||||
...getBaseContext(),
|
||||
ui: uiContext,
|
||||
hasUI,
|
||||
toolNames,
|
||||
toolCall,
|
||||
}),
|
||||
setUIContext: (context, uiAvailable) => {
|
||||
uiContext = context;
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import type { AgentTool } from "@oh-my-pi/pi-agent-core";
|
||||
import type { AgentTool, AgentToolContext, ToolCallContext } from "@oh-my-pi/pi-agent-core";
|
||||
import type { Component } from "@oh-my-pi/pi-tui";
|
||||
import { Text } from "@oh-my-pi/pi-tui";
|
||||
import { Type } from "@sinclair/typebox";
|
||||
@@ -41,6 +41,22 @@ export interface EditToolDetails {
|
||||
diagnostics?: FileDiagnosticsResult;
|
||||
}
|
||||
|
||||
const LSP_BATCH_TOOLS = new Set(["edit", "write"]);
|
||||
|
||||
function getLspBatchRequest(toolCall: ToolCallContext | undefined): { id: string; flush: boolean } | undefined {
|
||||
if (!toolCall) {
|
||||
return undefined;
|
||||
}
|
||||
const hasOtherWrites = toolCall.toolCalls.some(
|
||||
(call, index) => index !== toolCall.index && LSP_BATCH_TOOLS.has(call.name),
|
||||
);
|
||||
if (!hasOtherWrites) {
|
||||
return undefined;
|
||||
}
|
||||
const hasLaterWrites = toolCall.toolCalls.slice(toolCall.index + 1).some((call) => LSP_BATCH_TOOLS.has(call.name));
|
||||
return { id: toolCall.batchId, flush: !hasLaterWrites };
|
||||
}
|
||||
|
||||
export function createEditTool(session: ToolSession): AgentTool<typeof editSchema> {
|
||||
const allowFuzzy = session.settings?.getEditFuzzyMatch() ?? true;
|
||||
const enableLsp = session.enableLsp ?? true;
|
||||
@@ -58,6 +74,8 @@ export function createEditTool(session: ToolSession): AgentTool<typeof editSchem
|
||||
_toolCallId: string,
|
||||
{ path, oldText, newText, all }: { path: string; oldText: string; newText: string; all?: boolean },
|
||||
signal?: AbortSignal,
|
||||
_onUpdate?: unknown,
|
||||
context?: AgentToolContext,
|
||||
) => {
|
||||
// Reject .ipynb files - use NotebookEdit tool instead
|
||||
if (path.endsWith(".ipynb")) {
|
||||
@@ -163,7 +181,8 @@ export function createEditTool(session: ToolSession): AgentTool<typeof editSchem
|
||||
}
|
||||
|
||||
const finalContent = bom + restoreLineEndings(normalizedNewContent, originalEnding);
|
||||
const diagnostics = await writethrough(absolutePath, finalContent, signal, file);
|
||||
const batchRequest = getLspBatchRequest(context?.toolCall);
|
||||
const diagnostics = await writethrough(absolutePath, finalContent, signal, file, batchRequest);
|
||||
|
||||
const diffResult = generateDiffString(normalizedContent, normalizedNewContent);
|
||||
|
||||
|
||||
@@ -667,6 +667,7 @@ export type WritethroughCallback = (
|
||||
content: string,
|
||||
signal?: AbortSignal,
|
||||
file?: BunFile,
|
||||
batch?: LspWritethroughBatchRequest,
|
||||
) => Promise<FileDiagnosticsResult | undefined>;
|
||||
|
||||
/** No-op writethrough callback */
|
||||
@@ -684,83 +685,241 @@ export async function writethroughNoop(
|
||||
return undefined;
|
||||
}
|
||||
|
||||
interface PendingWritethrough {
|
||||
dst: string;
|
||||
content: string;
|
||||
file?: BunFile;
|
||||
}
|
||||
|
||||
interface LspWritethroughBatchRequest {
|
||||
id: string;
|
||||
flush: boolean;
|
||||
}
|
||||
|
||||
interface LspWritethroughBatchState {
|
||||
entries: Map<string, PendingWritethrough>;
|
||||
options: Required<WritethroughOptions>;
|
||||
}
|
||||
|
||||
const writethroughBatches = new Map<string, LspWritethroughBatchState>();
|
||||
|
||||
function getOrCreateWritethroughBatch(id: string, options: Required<WritethroughOptions>): LspWritethroughBatchState {
|
||||
const existing = writethroughBatches.get(id);
|
||||
if (existing) {
|
||||
existing.options.enableFormat ||= options.enableFormat;
|
||||
existing.options.enableDiagnostics ||= options.enableDiagnostics;
|
||||
return existing;
|
||||
}
|
||||
const batch: LspWritethroughBatchState = {
|
||||
entries: new Map<string, PendingWritethrough>(),
|
||||
options: { ...options },
|
||||
};
|
||||
writethroughBatches.set(id, batch);
|
||||
return batch;
|
||||
}
|
||||
|
||||
function summarizeDiagnosticMessages(messages: string[]): { summary: string; errored: boolean } {
|
||||
const counts = { error: 0, warning: 0, info: 0, hint: 0 };
|
||||
for (const message of messages) {
|
||||
const match = message.match(/\[(error|warning|info|hint)\]/i);
|
||||
if (!match) continue;
|
||||
const key = match[1].toLowerCase() as keyof typeof counts;
|
||||
counts[key] += 1;
|
||||
}
|
||||
|
||||
const parts: string[] = [];
|
||||
if (counts.error > 0) parts.push(`${counts.error} error(s)`);
|
||||
if (counts.warning > 0) parts.push(`${counts.warning} warning(s)`);
|
||||
if (counts.info > 0) parts.push(`${counts.info} info(s)`);
|
||||
if (counts.hint > 0) parts.push(`${counts.hint} hint(s)`);
|
||||
|
||||
return {
|
||||
summary: parts.length > 0 ? parts.join(", ") : "no issues",
|
||||
errored: counts.error > 0,
|
||||
};
|
||||
}
|
||||
|
||||
function mergeDiagnostics(
|
||||
results: Array<FileDiagnosticsResult | undefined>,
|
||||
options: Required<WritethroughOptions>,
|
||||
): FileDiagnosticsResult | undefined {
|
||||
const messages: string[] = [];
|
||||
const servers = new Set<string>();
|
||||
let hasResults = false;
|
||||
let hasFormatter = false;
|
||||
let formatted = false;
|
||||
|
||||
for (const result of results) {
|
||||
if (!result) continue;
|
||||
hasResults = true;
|
||||
if (result.server) {
|
||||
for (const server of result.server.split(",")) {
|
||||
const trimmed = server.trim();
|
||||
if (trimmed) {
|
||||
servers.add(trimmed);
|
||||
}
|
||||
}
|
||||
}
|
||||
if (result.messages.length > 0) {
|
||||
messages.push(...result.messages);
|
||||
}
|
||||
if (result.formatter !== undefined) {
|
||||
hasFormatter = true;
|
||||
if (result.formatter === FileFormatResult.FORMATTED) {
|
||||
formatted = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (!hasResults && !hasFormatter) {
|
||||
return undefined;
|
||||
}
|
||||
|
||||
let summary = options.enableDiagnostics ? "no issues" : "OK";
|
||||
let errored = false;
|
||||
if (messages.length > 0) {
|
||||
const summaryInfo = summarizeDiagnosticMessages(messages);
|
||||
summary = summaryInfo.summary;
|
||||
errored = summaryInfo.errored;
|
||||
}
|
||||
const formatter = hasFormatter ? (formatted ? FileFormatResult.FORMATTED : FileFormatResult.UNCHANGED) : undefined;
|
||||
|
||||
return {
|
||||
server: servers.size > 0 ? Array.from(servers).join(", ") : undefined,
|
||||
messages,
|
||||
summary,
|
||||
errored,
|
||||
formatter,
|
||||
};
|
||||
}
|
||||
|
||||
async function runLspWritethrough(
|
||||
dst: string,
|
||||
content: string,
|
||||
cwd: string,
|
||||
options: Required<WritethroughOptions>,
|
||||
signal?: AbortSignal,
|
||||
file?: BunFile,
|
||||
): Promise<FileDiagnosticsResult | undefined> {
|
||||
const { enableFormat, enableDiagnostics } = options;
|
||||
const config = await getConfig(cwd);
|
||||
const servers = getServersForFile(config, dst);
|
||||
if (servers.length === 0) {
|
||||
return writethroughNoop(dst, content, signal, file);
|
||||
}
|
||||
const { lspServers, customLinterServers } = splitServers(servers);
|
||||
|
||||
let finalContent = content;
|
||||
const writeContent = async (value: string) => (file ? file.write(value) : Bun.write(dst, value));
|
||||
const getWritePromise = once(() => writeContent(finalContent));
|
||||
const useCustomFormatter = enableFormat && customLinterServers.length > 0;
|
||||
|
||||
// Capture diagnostic versions BEFORE syncing to detect stale diagnostics
|
||||
const minVersions = enableDiagnostics ? await captureDiagnosticVersions(cwd, servers) : undefined;
|
||||
|
||||
let formatter: FileFormatResult | undefined;
|
||||
let diagnostics: FileDiagnosticsResult | undefined;
|
||||
try {
|
||||
const timeoutSignal = AbortSignal.timeout(10_000);
|
||||
const operationSignal = signal ? AbortSignal.any([signal, timeoutSignal]) : timeoutSignal;
|
||||
await untilAborted(operationSignal, async () => {
|
||||
if (useCustomFormatter) {
|
||||
// Custom linters (e.g. Biome CLI) require on-disk input.
|
||||
await writeContent(content);
|
||||
finalContent = await formatContent(dst, content, cwd, customLinterServers, operationSignal);
|
||||
formatter = finalContent !== content ? FileFormatResult.FORMATTED : FileFormatResult.UNCHANGED;
|
||||
await writeContent(finalContent);
|
||||
await syncFileContent(dst, finalContent, cwd, lspServers, operationSignal);
|
||||
} else {
|
||||
// 1. Sync original content to LSP servers
|
||||
await syncFileContent(dst, content, cwd, lspServers, operationSignal);
|
||||
|
||||
// 2. Format in-memory via LSP
|
||||
if (enableFormat) {
|
||||
finalContent = await formatContent(dst, content, cwd, lspServers, operationSignal);
|
||||
formatter = finalContent !== content ? FileFormatResult.FORMATTED : FileFormatResult.UNCHANGED;
|
||||
}
|
||||
|
||||
// 3. If formatted, sync formatted content to LSP servers
|
||||
if (finalContent !== content) {
|
||||
await syncFileContent(dst, finalContent, cwd, lspServers, operationSignal);
|
||||
}
|
||||
|
||||
// 4. Write to disk
|
||||
await getWritePromise();
|
||||
}
|
||||
|
||||
// 5. Notify saved to LSP servers
|
||||
await notifyFileSaved(dst, cwd, lspServers, operationSignal);
|
||||
|
||||
// 6. Get diagnostics from all servers (wait for fresh results)
|
||||
if (enableDiagnostics) {
|
||||
diagnostics = await getDiagnosticsForFile(dst, cwd, servers, operationSignal, minVersions);
|
||||
}
|
||||
});
|
||||
} catch {
|
||||
await getWritePromise();
|
||||
}
|
||||
|
||||
if (formatter !== undefined) {
|
||||
diagnostics ??= {
|
||||
server: servers.map(([name]) => name).join(", "),
|
||||
messages: [],
|
||||
summary: "OK",
|
||||
errored: false,
|
||||
};
|
||||
diagnostics.formatter = formatter;
|
||||
}
|
||||
|
||||
return diagnostics;
|
||||
}
|
||||
|
||||
async function flushWritethroughBatch(
|
||||
batch: PendingWritethrough[],
|
||||
cwd: string,
|
||||
options: Required<WritethroughOptions>,
|
||||
signal?: AbortSignal,
|
||||
): Promise<FileDiagnosticsResult | undefined> {
|
||||
if (batch.length === 0) {
|
||||
return undefined;
|
||||
}
|
||||
const results: Array<FileDiagnosticsResult | undefined> = [];
|
||||
for (const entry of batch) {
|
||||
results.push(await runLspWritethrough(entry.dst, entry.content, cwd, options, signal, entry.file));
|
||||
}
|
||||
return mergeDiagnostics(results, options);
|
||||
}
|
||||
|
||||
/** Create a writethrough callback for LSP aware write operations */
|
||||
export function createLspWritethrough(cwd: string, options?: WritethroughOptions): WritethroughCallback {
|
||||
const { enableFormat = false, enableDiagnostics = false } = options ?? {};
|
||||
if (!enableFormat && !enableDiagnostics) {
|
||||
const resolvedOptions: Required<WritethroughOptions> = {
|
||||
enableFormat: options?.enableFormat ?? false,
|
||||
enableDiagnostics: options?.enableDiagnostics ?? false,
|
||||
};
|
||||
if (!resolvedOptions.enableFormat && !resolvedOptions.enableDiagnostics) {
|
||||
return writethroughNoop;
|
||||
}
|
||||
return async (dst: string, content: string, signal?: AbortSignal, file?: BunFile) => {
|
||||
const config = await getConfig(cwd);
|
||||
const servers = getServersForFile(config, dst);
|
||||
if (servers.length === 0) {
|
||||
return writethroughNoop(dst, content, signal, file);
|
||||
}
|
||||
const { lspServers, customLinterServers } = splitServers(servers);
|
||||
|
||||
let finalContent = content;
|
||||
const writeContent = async (value: string) => (file ? file.write(value) : Bun.write(dst, value));
|
||||
const getWritePromise = once(() => writeContent(finalContent));
|
||||
const useCustomFormatter = enableFormat && customLinterServers.length > 0;
|
||||
|
||||
// Capture diagnostic versions BEFORE syncing to detect stale diagnostics
|
||||
const minVersions = enableDiagnostics ? await captureDiagnosticVersions(cwd, servers) : undefined;
|
||||
|
||||
let formatter: FileFormatResult | undefined;
|
||||
let diagnostics: FileDiagnosticsResult | undefined;
|
||||
try {
|
||||
const timeoutSignal = AbortSignal.timeout(10_000);
|
||||
const operationSignal = signal ? AbortSignal.any([signal, timeoutSignal]) : timeoutSignal;
|
||||
await untilAborted(operationSignal, async () => {
|
||||
if (useCustomFormatter) {
|
||||
// Custom linters (e.g. Biome CLI) require on-disk input.
|
||||
await writeContent(content);
|
||||
finalContent = await formatContent(dst, content, cwd, customLinterServers, operationSignal);
|
||||
formatter = finalContent !== content ? FileFormatResult.FORMATTED : FileFormatResult.UNCHANGED;
|
||||
await writeContent(finalContent);
|
||||
await syncFileContent(dst, finalContent, cwd, lspServers, operationSignal);
|
||||
} else {
|
||||
// 1. Sync original content to LSP servers
|
||||
await syncFileContent(dst, content, cwd, lspServers, operationSignal);
|
||||
|
||||
// 2. Format in-memory via LSP
|
||||
if (enableFormat) {
|
||||
finalContent = await formatContent(dst, content, cwd, lspServers, operationSignal);
|
||||
formatter = finalContent !== content ? FileFormatResult.FORMATTED : FileFormatResult.UNCHANGED;
|
||||
}
|
||||
|
||||
// 3. If formatted, sync formatted content to LSP servers
|
||||
if (finalContent !== content) {
|
||||
await syncFileContent(dst, finalContent, cwd, lspServers, operationSignal);
|
||||
}
|
||||
|
||||
// 4. Write to disk
|
||||
await getWritePromise();
|
||||
}
|
||||
|
||||
// 5. Notify saved to LSP servers
|
||||
await notifyFileSaved(dst, cwd, lspServers, operationSignal);
|
||||
|
||||
// 6. Get diagnostics from all servers (wait for fresh results)
|
||||
if (enableDiagnostics) {
|
||||
diagnostics = await getDiagnosticsForFile(dst, cwd, servers, operationSignal, minVersions);
|
||||
}
|
||||
});
|
||||
} catch {
|
||||
await getWritePromise();
|
||||
return async (
|
||||
dst: string,
|
||||
content: string,
|
||||
signal?: AbortSignal,
|
||||
file?: BunFile,
|
||||
batch?: LspWritethroughBatchRequest,
|
||||
) => {
|
||||
if (!batch) {
|
||||
return runLspWritethrough(dst, content, cwd, resolvedOptions, signal, file);
|
||||
}
|
||||
|
||||
if (formatter !== undefined) {
|
||||
diagnostics ??= {
|
||||
server: servers.map(([name]) => name).join(", "),
|
||||
messages: [],
|
||||
summary: "OK",
|
||||
errored: false,
|
||||
};
|
||||
diagnostics.formatter = formatter;
|
||||
const state = getOrCreateWritethroughBatch(batch.id, resolvedOptions);
|
||||
state.entries.set(dst, { dst, content, file });
|
||||
|
||||
if (!batch.flush) {
|
||||
await writethroughNoop(dst, content, signal, file);
|
||||
return undefined;
|
||||
}
|
||||
|
||||
return diagnostics;
|
||||
writethroughBatches.delete(batch.id);
|
||||
return flushWritethroughBatch(Array.from(state.entries.values()), cwd, state.options, signal);
|
||||
};
|
||||
}
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import type { AgentTool } from "@oh-my-pi/pi-agent-core";
|
||||
import type { AgentTool, AgentToolContext, ToolCallContext } from "@oh-my-pi/pi-agent-core";
|
||||
import type { Component } from "@oh-my-pi/pi-tui";
|
||||
import { Text } from "@oh-my-pi/pi-tui";
|
||||
import { Type } from "@sinclair/typebox";
|
||||
@@ -22,6 +22,22 @@ export interface WriteToolDetails {
|
||||
diagnostics?: FileDiagnosticsResult;
|
||||
}
|
||||
|
||||
const LSP_BATCH_TOOLS = new Set(["edit", "write"]);
|
||||
|
||||
function getLspBatchRequest(toolCall: ToolCallContext | undefined): { id: string; flush: boolean } | undefined {
|
||||
if (!toolCall) {
|
||||
return undefined;
|
||||
}
|
||||
const hasOtherWrites = toolCall.toolCalls.some(
|
||||
(call, index) => index !== toolCall.index && LSP_BATCH_TOOLS.has(call.name),
|
||||
);
|
||||
if (!hasOtherWrites) {
|
||||
return undefined;
|
||||
}
|
||||
const hasLaterWrites = toolCall.toolCalls.slice(toolCall.index + 1).some((call) => LSP_BATCH_TOOLS.has(call.name));
|
||||
return { id: toolCall.batchId, flush: !hasLaterWrites };
|
||||
}
|
||||
|
||||
export function createWriteTool(session: ToolSession): AgentTool<typeof writeSchema, WriteToolDetails> {
|
||||
const enableLsp = session.enableLsp ?? true;
|
||||
const enableFormat = enableLsp ? (session.settings?.getLspFormatOnWrite() ?? true) : false;
|
||||
@@ -38,11 +54,14 @@ export function createWriteTool(session: ToolSession): AgentTool<typeof writeSch
|
||||
_toolCallId: string,
|
||||
{ path, content }: { path: string; content: string },
|
||||
signal?: AbortSignal,
|
||||
_onUpdate?: unknown,
|
||||
context?: AgentToolContext,
|
||||
) => {
|
||||
return untilAborted(signal, async () => {
|
||||
const absolutePath = resolveToCwd(path, session.cwd);
|
||||
const batchRequest = getLspBatchRequest(context?.toolCall);
|
||||
|
||||
const diagnostics = await writethrough(absolutePath, content, signal);
|
||||
const diagnostics = await writethrough(absolutePath, content, signal, undefined, batchRequest);
|
||||
|
||||
let resultText = `Successfully wrote ${content.length} bytes to ${path}`;
|
||||
if (!diagnostics) {
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test";
|
||||
import { mkdtempSync, rmSync } from "node:fs";
|
||||
import { tmpdir } from "node:os";
|
||||
import { join } from "node:path";
|
||||
import * as lspConfig from "../../src/core/tools/lsp/config";
|
||||
import { createLspWritethrough } from "../../src/core/tools/lsp/index";
|
||||
|
||||
describe("createLspWritethrough batching", () => {
|
||||
let tempDir: string;
|
||||
|
||||
beforeEach(() => {
|
||||
tempDir = mkdtempSync(join(tmpdir(), "omp-lsp-batch-"));
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
rmSync(tempDir, { recursive: true, force: true });
|
||||
});
|
||||
|
||||
it("defers LSP work until the batch flush", async () => {
|
||||
const loadConfigSpy = vi
|
||||
.spyOn(lspConfig, "loadConfig")
|
||||
.mockResolvedValue({ servers: {}, idleTimeoutMs: undefined });
|
||||
const getServersSpy = vi.spyOn(lspConfig, "getServersForFile").mockReturnValue([]);
|
||||
const writethrough = createLspWritethrough(tempDir, { enableFormat: true, enableDiagnostics: true });
|
||||
|
||||
const fileA = join(tempDir, "a.ts");
|
||||
const fileB = join(tempDir, "b.ts");
|
||||
const batchId = `batch-${Date.now()}`;
|
||||
|
||||
const firstResult = await writethrough(fileA, "const a = 1;\n", undefined, undefined, {
|
||||
id: batchId,
|
||||
flush: false,
|
||||
});
|
||||
|
||||
expect(firstResult).toBeUndefined();
|
||||
expect(getServersSpy).toHaveBeenCalledTimes(0);
|
||||
expect(loadConfigSpy).toHaveBeenCalledTimes(0);
|
||||
expect(await Bun.file(fileA).text()).toBe("const a = 1;\n");
|
||||
|
||||
const secondResult = await writethrough(fileB, "const b = 2;\n", undefined, undefined, {
|
||||
id: batchId,
|
||||
flush: true,
|
||||
});
|
||||
|
||||
expect(secondResult).toBeUndefined();
|
||||
expect(getServersSpy).toHaveBeenCalledTimes(2);
|
||||
expect(loadConfigSpy).toHaveBeenCalledTimes(1);
|
||||
expect(await Bun.file(fileA).text()).toBe("const a = 1;\n");
|
||||
expect(await Bun.file(fileB).text()).toBe("const b = 2;\n");
|
||||
});
|
||||
|
||||
it("runs LSP immediately when no batch is provided", async () => {
|
||||
const loadConfigSpy = vi
|
||||
.spyOn(lspConfig, "loadConfig")
|
||||
.mockResolvedValue({ servers: {}, idleTimeoutMs: undefined });
|
||||
const getServersSpy = vi.spyOn(lspConfig, "getServersForFile").mockReturnValue([]);
|
||||
const writethrough = createLspWritethrough(tempDir, { enableFormat: true, enableDiagnostics: true });
|
||||
|
||||
const filePath = join(tempDir, "single.ts");
|
||||
const result = await writethrough(filePath, "const single = true;\n");
|
||||
|
||||
expect(result).toBeUndefined();
|
||||
expect(getServersSpy).toHaveBeenCalledTimes(1);
|
||||
expect(loadConfigSpy).toHaveBeenCalledTimes(1);
|
||||
expect(await Bun.file(filePath).text()).toBe("const single = true;\n");
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user