feat: added strict-mode fallback for OpenAI tool calls with all_strict

- Added `toolStrictMode` support with `all_strict`/`none`/`mixed` options to OpenAI compatibility.
- Fixed OpenAI-completion strict-mode flows by capturing failed HTTP responses and retrying once as non-strict.
- Fixed completion error reporting by surfacing captured status, headers, and JSON `type`/`param`/`code` details.
- Improved strict-schema enforcement with WeakMap memoization and circular-schema detection in sanitization.
- Fixed OpenRouter provider lookup by resolving fallback model IDs for suffix and date variants in registry resolution.
- Refactored benchmark tooling and added async RPC error-window tracking for scheduled run execution.
This commit is contained in:
can1357
2026-04-13 15:30:34 +02:00
parent b69dea4823
commit 212d56bc11
38 changed files with 1492 additions and 548 deletions
+12
View File
@@ -1,10 +1,22 @@
# Changelog
## [Unreleased]
### Added
- Added `toolStrictMode` compatibility option (`"all_strict"` or `"none"`) to OpenAI-compatible model config to force tool schemas to be sent uniformly strict, uniformly non-strict, or keep mixed per-tool behavior
### Changed
- Changed Cerebras OpenAI-compatible providers to default `toolStrictMode` to `"all_strict"` unless explicitly overridden
### Fixed
- Fixed OpenAI Completions handling for providers that reject mixed `strict` flags by automatically retrying with non-strict tool schemas when an initial all-strict tool request fails with strict-format 400/422 errors
- Fixed OpenAI-completions error reporting by including captured JSON error body details such as type, param, and code when a request fails without a body in the thrown SDK error
- Fixed shell execution failure responses to preserve all result fields when sanitizing, preventing truncated metadata in stream results
- Fixed context overflow detection to recognize `model_context_window_exceeded` from z.ai / GLM providers, preventing infinite retry loops when context window is exceeded ([#638](https://github.com/can1357/oh-my-pi/issues/638))
- Fixed strict tool schema enforcement to preserve `additionalProperties: false` and required keys for reused nested object schemas, preventing invalid `todo_write` function schemas in Codex/OpenAI requests
## [14.1.0] - 2026-04-11
### Added
@@ -1,13 +1,15 @@
import type { Model, OpenAICompat } from "../types";
type OpenAIReasoningEffort = "minimal" | "low" | "medium" | "high" | "xhigh";
type ResolvedToolStrictMode = NonNullable<OpenAICompat["toolStrictMode"]> | "mixed";
export type ResolvedOpenAICompat = Required<
Omit<OpenAICompat, "openRouterRouting" | "vercelGatewayRouting" | "extraBody">
Omit<OpenAICompat, "openRouterRouting" | "vercelGatewayRouting" | "extraBody" | "toolStrictMode">
> & {
openRouterRouting?: OpenAICompat["openRouterRouting"];
vercelGatewayRouting?: OpenAICompat["vercelGatewayRouting"];
extraBody?: OpenAICompat["extraBody"];
toolStrictMode: ResolvedToolStrictMode;
};
function detectStrictModeSupport(provider: string, baseUrl: string): boolean {
@@ -109,6 +111,7 @@ export function detectOpenAICompat(model: Model<"openai-completions">, resolvedB
vercelGatewayRouting: undefined,
supportsStrictMode: detectStrictModeSupport(provider, baseUrl),
extraBody: undefined,
toolStrictMode: isCerebras ? "all_strict" : "mixed",
};
}
@@ -151,5 +154,6 @@ export function resolveOpenAICompat(
vercelGatewayRouting: model.compat.vercelGatewayRouting ?? detected.vercelGatewayRouting,
supportsStrictMode: model.compat.supportsStrictMode ?? detected.supportsStrictMode,
extraBody: model.compat.extraBody,
toolStrictMode: model.compat.toolStrictMode ?? detected.toolStrictMode,
};
}
+162 -33
View File
@@ -31,7 +31,12 @@ import {
} from "../types";
import { createAbortSourceTracker } from "../utils/abort";
import { AssistantMessageEventStream } from "../utils/event-stream";
import { finalizeErrorMessage, type RawHttpRequestDump, rewriteCopilotAuthError } from "../utils/http-inspector";
import {
type CapturedHttpErrorResponse,
finalizeErrorMessage,
type RawHttpRequestDump,
rewriteCopilotAuthError,
} from "../utils/http-inspector";
import {
createFirstEventWatchdog,
getOpenAIStreamIdleTimeoutMs,
@@ -42,6 +47,7 @@ import {
import { parseStreamingJson } from "../utils/json-parse";
import { parseGitHubCopilotApiKey } from "../utils/oauth/github-copilot";
import { getKimiCommonHeaders } from "../utils/oauth/kimi";
import { extractHttpStatusFromError } from "../utils/retry";
import { adaptSchemaForStrict, NO_STRICT } from "../utils/schema";
import { mapToOpenAICompletionsToolChoice } from "../utils/tool-choice";
import {
@@ -126,6 +132,14 @@ type OpenAICompletionsSamplingParams = OpenAI.Chat.Completions.ChatCompletionCre
repetition_penalty?: number;
};
type AppliedToolStrictMode = "mixed" | "all_strict" | "none";
type ToolStrictModeOverride = Exclude<ResolvedOpenAICompat["toolStrictMode"], "mixed"> | undefined;
type BuiltOpenAICompletionTools = {
tools: OpenAI.Chat.Completions.ChatCompletionTool[];
toolStrictMode: AppliedToolStrictMode;
};
// LIMITATION: The think tag parser uses naive string matching for <think>/<thinking> tags.
// If MiniMax models output these literal strings in code blocks, XML examples, or explanations,
// they will be incorrectly consumed as thinking delimiters, truncating visible output.
@@ -177,6 +191,7 @@ export const streamOpenAICompletions: StreamFunction<"openai-completions"> = (
(async () => {
const startTime = Date.now();
let firstTokenTime: number | undefined;
let getCapturedErrorResponse: (() => CapturedHttpErrorResponse | undefined) | undefined;
const output: AssistantMessage = {
role: "assistant",
@@ -203,24 +218,42 @@ export const streamOpenAICompletions: StreamFunction<"openai-completions"> = (
try {
const apiKey = options?.apiKey || getEnvApiKey(model.provider) || "";
const idleTimeoutMs = getOpenAIStreamIdleTimeoutMs();
const { client, copilotPremiumRequests, baseUrl } = await createClient(
model,
context,
apiKey,
options?.headers,
options?.initiatorOverride,
);
const params = buildParams(model, context, options, baseUrl);
options?.onPayload?.(params);
rawRequestDump = {
provider: model.provider,
api: output.api,
model: model.id,
method: "POST",
url: `${baseUrl}/chat/completions`,
body: params,
const {
client,
copilotPremiumRequests,
baseUrl,
requestHeaders,
getCapturedErrorResponse: captureErrorResponse,
clearCapturedErrorResponse,
} = await createClient(model, context, apiKey, options?.headers, options?.initiatorOverride);
getCapturedErrorResponse = captureErrorResponse;
let appliedToolStrictMode: AppliedToolStrictMode = "mixed";
const createCompletionsStream = async (toolStrictModeOverride?: ToolStrictModeOverride) => {
clearCapturedErrorResponse();
const { params, toolStrictMode } = buildParams(model, context, options, baseUrl, toolStrictModeOverride);
appliedToolStrictMode = toolStrictMode;
options?.onPayload?.(params);
rawRequestDump = {
provider: model.provider,
api: output.api,
model: model.id,
method: "POST",
url: `${baseUrl}/chat/completions`,
headers: requestHeaders,
body: params,
};
return client.chat.completions.create(params, { signal: requestSignal });
};
const openaiStream = await client.chat.completions.create(params, { signal: requestSignal });
let openaiStream: AsyncIterable<ChatCompletionChunk>;
try {
openaiStream = await createCompletionsStream();
} catch (error) {
const capturedErrorResponse = getCapturedErrorResponse();
if (!shouldRetryWithoutStrictTools(error, capturedErrorResponse, appliedToolStrictMode, context.tools)) {
throw error;
}
openaiStream = await createCompletionsStream("none");
}
const firstEventWatchdog = createFirstEventWatchdog(
options?.streamFirstEventTimeoutMs ?? getStreamFirstEventTimeoutMs(idleTimeoutMs),
() => abortTracker.abortLocally(firstEventTimeoutAbortError),
@@ -513,7 +546,9 @@ export const streamOpenAICompletions: StreamFunction<"openai-completions"> = (
for (const block of output.content) delete (block as any).index;
const firstEventTimeoutError = abortTracker.getLocalAbortReason();
output.stopReason = abortTracker.wasCallerAbort() ? "aborted" : "error";
output.errorMessage = firstEventTimeoutError?.message ?? (await finalizeErrorMessage(error, rawRequestDump));
output.errorMessage =
firstEventTimeoutError?.message ??
(await finalizeErrorMessage(error, rawRequestDump, getCapturedErrorResponse?.()));
// Some providers via OpenRouter include extra details here.
const rawMetadata = (error as { error?: { metadata?: { raw?: string } } })?.error?.metadata?.raw;
if (rawMetadata) output.errorMessage += `\n${rawMetadata}`;
@@ -538,6 +573,9 @@ async function createClient(
client: OpenAI;
copilotPremiumRequests: number | undefined;
baseUrl: string | undefined;
requestHeaders: Record<string, string>;
getCapturedErrorResponse: () => CapturedHttpErrorResponse | undefined;
clearCapturedErrorResponse: () => void;
}> {
if (!apiKey) {
if (!$env.OPENAI_API_KEY) {
@@ -573,6 +611,34 @@ async function createClient(
copilotPremiumRequests = copilot.premiumRequests;
baseUrl = resolveGitHubCopilotBaseUrl(model.baseUrl, rawApiKey) ?? model.baseUrl;
}
let capturedErrorResponse: CapturedHttpErrorResponse | undefined;
const wrappedFetch = Object.assign(
async (input: string | URL | Request, init?: RequestInit): Promise<Response> => {
const response = await fetch(input, init);
if (response.ok) {
capturedErrorResponse = undefined;
return response;
}
let bodyText: string | undefined;
let bodyJson: unknown;
try {
bodyText = await response.clone().text();
if (bodyText.trim().length > 0) {
try {
bodyJson = JSON.parse(bodyText);
} catch {}
}
} catch {}
capturedErrorResponse = {
status: response.status,
headers: response.headers,
bodyText,
bodyJson,
};
return response;
},
{ preconnect: fetch.preconnect },
);
return {
client: new OpenAI({
apiKey,
@@ -580,9 +646,15 @@ async function createClient(
dangerouslyAllowBrowser: true,
maxRetries: 5,
defaultHeaders: headers,
fetch: wrappedFetch,
}),
copilotPremiumRequests,
baseUrl,
requestHeaders: headers,
getCapturedErrorResponse: () => capturedErrorResponse,
clearCapturedErrorResponse: () => {
capturedErrorResponse = undefined;
},
};
}
@@ -591,7 +663,8 @@ function buildParams(
context: Context,
options: OpenAICompletionsOptions | undefined,
resolvedBaseUrl?: string,
) {
toolStrictModeOverride?: ToolStrictModeOverride,
): { params: OpenAICompletionsSamplingParams; toolStrictMode: AppliedToolStrictMode } {
const compat = getCompat(model, resolvedBaseUrl);
const messages = convertMessages(model, context, compat);
maybeAddOpenRouterAnthropicCacheControl(model, messages);
@@ -607,6 +680,7 @@ function buildParams(
messages,
stream: true,
};
let toolStrictMode: AppliedToolStrictMode = "none";
if (compat.supportsUsageInStreaming !== false) {
(params as { stream_options?: { include_usage: boolean } }).stream_options = { include_usage: true };
@@ -647,7 +721,9 @@ function buildParams(
}
if (context.tools) {
params.tools = convertTools(context.tools, compat);
const builtTools = convertTools(context.tools, compat, toolStrictModeOverride);
params.tools = builtTools.tools;
toolStrictMode = builtTools.toolStrictMode;
} else if (hasToolHistory(context.messages)) {
// Anthropic (via LiteLLM/proxy) requires tools param when conversation has tool_calls/tool_results
params.tools = [];
@@ -697,7 +773,14 @@ function buildParams(
Object.assign(params, compat.extraBody);
}
return params;
return buildParamsResult(params, toolStrictMode);
}
function buildParamsResult(
params: OpenAICompletionsSamplingParams,
toolStrictMode: AppliedToolStrictMode,
): { params: OpenAICompletionsSamplingParams; toolStrictMode: AppliedToolStrictMode } {
return { params, toolStrictMode };
}
function getOptionalNumberProperty(value: object, key: string): number | undefined {
@@ -1102,22 +1185,68 @@ export function convertMessages(
return params;
}
function convertTools(tools: Tool[], compat: ResolvedOpenAICompat): OpenAI.Chat.Completions.ChatCompletionTool[] {
return tools.map(tool => {
function convertTools(
tools: Tool[],
compat: ResolvedOpenAICompat,
toolStrictModeOverride?: ToolStrictModeOverride,
): BuiltOpenAICompletionTools {
const adaptedTools = tools.map(tool => {
const strict = !NO_STRICT && compat.supportsStrictMode !== false && tool.strict !== false;
const baseParameters = tool.parameters as unknown as Record<string, unknown>;
const { schema: parameters, strict: effectiveStrict } = adaptSchemaForStrict(baseParameters, strict);
const adapted = adaptSchemaForStrict(baseParameters, strict);
return {
type: "function",
function: {
name: tool.name,
description: tool.description || "",
parameters,
// Only include strict if provider supports it. Some reject unknown fields.
...(effectiveStrict && { strict: true }),
},
tool,
baseParameters,
parameters: adapted.schema,
strict: adapted.strict,
};
});
const requestedStrictMode = toolStrictModeOverride ?? compat.toolStrictMode;
const toolStrictMode =
requestedStrictMode === "none"
? "none"
: requestedStrictMode === "all_strict"
? adaptedTools.every(tool => tool.strict)
? "all_strict"
: "none"
: "mixed";
return {
tools: adaptedTools.map(({ tool, baseParameters, parameters, strict }) => {
const includeStrict = toolStrictMode === "all_strict" || (toolStrictMode === "mixed" && strict);
return {
type: "function",
function: {
name: tool.name,
description: tool.description || "",
parameters: includeStrict ? parameters : baseParameters,
// Only include strict if provider supports it. Some reject unknown fields.
...(includeStrict && { strict: true }),
},
};
}),
toolStrictMode,
};
}
function shouldRetryWithoutStrictTools(
error: unknown,
capturedErrorResponse: CapturedHttpErrorResponse | undefined,
toolStrictMode: AppliedToolStrictMode,
tools: Tool[] | undefined,
): boolean {
if (!tools || tools.length === 0 || toolStrictMode !== "all_strict") {
return false;
}
const status = extractHttpStatusFromError(error) ?? capturedErrorResponse?.status;
if (status !== 400 && status !== 422) {
return false;
}
const messageParts = [error instanceof Error ? error.message : undefined, capturedErrorResponse?.bodyText]
.filter((value): value is string => typeof value === "string" && value.trim().length > 0)
.join("\n");
return /wrong_api_format|mixed values for 'strict'|tool[s]?\b.*strict|\bstrict\b.*tool/i.test(messageParts);
}
function mapStopReason(reason: ChatCompletionChunk.Choice["finish_reason"] | string): {
+2
View File
@@ -483,6 +483,8 @@ export interface OpenAICompat {
extraBody?: Record<string, unknown>;
/** Whether the provider supports the `strict` field in tool definitions. Default: auto-detected per provider/baseUrl (conservative for unknown providers). */
supportsStrictMode?: boolean;
/** Whether tool schemas must be sent either all strict or all non-strict. Undefined keeps the existing per-tool mixed behavior. */
toolStrictMode?: "all_strict" | "none";
}
/**
+68 -1
View File
@@ -13,6 +13,13 @@ export type RawHttpRequestDump = {
body?: unknown;
};
export type CapturedHttpErrorResponse = {
status: number;
headers?: Headers;
bodyText?: string;
bodyJson?: unknown;
};
type ErrorWithStatus = {
status?: unknown;
};
@@ -44,8 +51,18 @@ export async function appendRawHttpRequestDumpFor400(
export async function finalizeErrorMessage(
error: unknown,
rawRequestDump: RawHttpRequestDump | undefined,
capturedErrorResponse?: CapturedHttpErrorResponse,
): Promise<string> {
return appendRawHttpRequestDumpFor400(formatErrorMessageWithRetryAfter(error), error, rawRequestDump);
let message = formatErrorMessageWithRetryAfter(error, capturedErrorResponse?.headers);
const capturedMessage = formatCapturedHttpError(capturedErrorResponse);
if (capturedMessage) {
if (/\bstatus code\s*\(no body\)/i.test(message)) {
message = `${capturedErrorResponse?.status ?? "HTTP"} status code: ${capturedMessage}`;
} else if (!message.includes(capturedMessage)) {
message = `${message}\n${capturedMessage}`;
}
}
return appendRawHttpRequestDumpFor400(message, error, rawRequestDump);
}
export function withHttpStatus(error: unknown, status: number): Error {
@@ -96,3 +113,53 @@ function redactHeaders(headers: Record<string, string> | undefined): Record<stri
}
return redacted;
}
function formatCapturedHttpError(captured: CapturedHttpErrorResponse | undefined): string | undefined {
if (!captured) return undefined;
const bodyText = captured.bodyText?.trim();
if (!bodyText) return undefined;
const payload = parseCapturedErrorPayload(captured);
if (!payload) return bodyText;
const errorPayload = getObjectProperty(payload, "error") ?? payload;
const message = getStringProperty(errorPayload, "message") ?? getStringProperty(payload, "message") ?? bodyText;
const extras = [
getStringProperty(errorPayload, "type") ?? getStringProperty(payload, "type"),
getStringProperty(errorPayload, "param") ?? getStringProperty(payload, "param"),
getStringProperty(errorPayload, "code") ?? getStringProperty(payload, "code"),
]
.filter(Boolean)
.map((value, index) => {
if (index === 0) return `type=${value}`;
if (index === 1) return `param=${value}`;
return `code=${value}`;
});
return extras.length > 0 ? `${message} (${extras.join(" ")})` : message;
}
function parseCapturedErrorPayload(captured: CapturedHttpErrorResponse): Record<string, unknown> | undefined {
if (isObject(captured.bodyJson)) {
return captured.bodyJson;
}
if (!captured.bodyText) return undefined;
try {
const parsed = JSON.parse(captured.bodyText);
return isObject(parsed) ? parsed : undefined;
} catch {
return undefined;
}
}
function getObjectProperty(value: Record<string, unknown>, key: string): Record<string, unknown> | undefined {
const property = value[key];
return isObject(property) ? property : undefined;
}
function getStringProperty(value: Record<string, unknown>, key: string): string | undefined {
const property = value[key];
return typeof property === "string" && property.trim().length > 0 ? property : undefined;
}
function isObject(value: unknown): value is Record<string, unknown> {
return typeof value === "object" && value !== null && !Array.isArray(value);
}
+46 -15
View File
@@ -88,8 +88,12 @@ function hasUnrepresentableStrictObjectMap(schema: Record<string, unknown>, seen
export function sanitizeSchemaForStrictMode(
schema: Record<string, unknown>,
seen?: WeakSet<object>,
cache?: WeakMap<Record<string, unknown>, Record<string, unknown>>,
): Record<string, unknown> {
if (!seen) seen = new WeakSet();
if (!cache) cache = new WeakMap();
const cached = cache.get(schema);
if (cached) return cached;
if (seen.has(schema)) return {};
seen.add(schema);
const typeValue = schema.type;
@@ -98,8 +102,10 @@ export function sanitizeSchemaForStrictMode(
const schemaWithoutType = { ...schema };
delete schemaWithoutType.type;
const sanitizedWithoutType = sanitizeSchemaForStrictMode(schemaWithoutType, seen);
const sanitizedWithoutType = sanitizeSchemaForStrictMode(schemaWithoutType, seen, cache);
if (typeVariants.length === 0) {
cache.set(schema, sanitizedWithoutType);
seen.delete(schema);
return sanitizedWithoutType;
}
@@ -113,19 +119,25 @@ export function sanitizeSchemaForStrictMode(
if (variantType !== "array") {
delete variantSchema.items;
}
return sanitizeSchemaForStrictMode(variantSchema, seen);
return sanitizeSchemaForStrictMode(variantSchema, seen, cache);
});
if (variants.length === 1) {
cache.set(schema, variants[0] as Record<string, unknown>);
seen.delete(schema);
return variants[0] as Record<string, unknown>;
}
return {
const result = {
anyOf: variants,
};
cache.set(schema, result);
seen.delete(schema);
return result;
}
const sanitized: Record<string, unknown> = {};
cache.set(schema, sanitized);
for (const [key, value] of Object.entries(schema)) {
if (NON_STRUCTURAL_SCHEMA_KEYS.has(key) || key === "type" || key === "const" || key === "nullable") {
continue;
@@ -135,7 +147,7 @@ export function sanitizeSchemaForStrictMode(
const properties = Object.fromEntries(
Object.entries(value).map(([propertyName, propertySchema]) => [
propertyName,
isJsonObject(propertySchema) ? sanitizeSchemaForStrictMode(propertySchema, seen) : propertySchema,
isJsonObject(propertySchema) ? sanitizeSchemaForStrictMode(propertySchema, seen, cache) : propertySchema,
]),
);
sanitized.properties = properties;
@@ -144,10 +156,10 @@ export function sanitizeSchemaForStrictMode(
if (key === "items") {
if (isJsonObject(value)) {
sanitized.items = sanitizeSchemaForStrictMode(value, seen);
sanitized.items = sanitizeSchemaForStrictMode(value, seen, cache);
} else if (Array.isArray(value)) {
sanitized.items = value.map(entry =>
isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, seen) : entry,
isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, seen, cache) : entry,
);
} else {
sanitized.items = value;
@@ -156,7 +168,9 @@ export function sanitizeSchemaForStrictMode(
}
if (COMBINATOR_KEYS.includes(key as (typeof COMBINATOR_KEYS)[number]) && Array.isArray(value)) {
sanitized[key] = value.map(entry => (isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, seen) : entry));
sanitized[key] = value.map(entry =>
isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, seen, cache) : entry,
);
continue;
}
@@ -164,7 +178,9 @@ export function sanitizeSchemaForStrictMode(
sanitized[key] = Object.fromEntries(
Object.entries(value).map(([definitionName, definitionSchema]) => [
definitionName,
isJsonObject(definitionSchema) ? sanitizeSchemaForStrictMode(definitionSchema, seen) : definitionSchema,
isJsonObject(definitionSchema)
? sanitizeSchemaForStrictMode(definitionSchema, seen, cache)
: definitionSchema,
]),
);
continue;
@@ -221,9 +237,11 @@ export function sanitizeSchemaForStrictMode(
if (schema.nullable === true) {
const { nullable: _, ...withoutNullable } = sanitized;
seen.delete(schema);
return { anyOf: [withoutNullable, { type: "null" }] };
}
seen.delete(schema);
return sanitized;
}
@@ -241,11 +259,23 @@ export function sanitizeSchemaForStrictMode(
* i.e. the node is not representable in strict mode. Prefer
* {@link tryEnforceStrictSchema} which catches this and degrades gracefully.
*/
export function enforceStrictSchema(schema: Record<string, unknown>, seen?: WeakSet<object>): Record<string, unknown> {
export function enforceStrictSchema(
schema: Record<string, unknown>,
seen?: WeakSet<object>,
cache?: WeakMap<Record<string, unknown>, Record<string, unknown>>,
): Record<string, unknown> {
if (!seen) seen = new WeakSet();
if (seen.has(schema)) return schema;
if (!cache) cache = new WeakMap();
if (seen.has(schema)) {
throw new Error("Schema contains a circular object graph — cannot enforce strict mode");
}
const cached = cache.get(schema);
if (cached) {
return cached;
}
seen.add(schema);
const result = { ...schema };
cache.set(schema, result);
const isObjectType = result.type === "object";
if (isObjectType) {
result.additionalProperties = false;
@@ -263,7 +293,7 @@ export function enforceStrictSchema(schema: Record<string, unknown>, seen?: Weak
Object.entries(props).map(([key, value]) => {
const processed =
value != null && typeof value === "object" && !Array.isArray(value)
? enforceStrictSchema(value as Record<string, unknown>, seen)
? enforceStrictSchema(value as Record<string, unknown>, seen, cache)
: value;
// Optional property — wrap as nullable so strict mode accepts it
if (!originalRequired.has(key)) {
@@ -287,18 +317,18 @@ export function enforceStrictSchema(schema: Record<string, unknown>, seen?: Weak
if (Array.isArray(result.items)) {
result.items = result.items.map(entry =>
entry != null && typeof entry === "object" && !Array.isArray(entry)
? enforceStrictSchema(entry as Record<string, unknown>, seen)
? enforceStrictSchema(entry as Record<string, unknown>, seen, cache)
: entry,
);
} else {
result.items = enforceStrictSchema(result.items as Record<string, unknown>, seen);
result.items = enforceStrictSchema(result.items as Record<string, unknown>, seen, cache);
}
}
for (const key of COMBINATOR_KEYS) {
if (Array.isArray(result[key])) {
result[key] = (result[key] as unknown[]).map(entry =>
entry != null && typeof entry === "object" && !Array.isArray(entry)
? enforceStrictSchema(entry as Record<string, unknown>, seen)
? enforceStrictSchema(entry as Record<string, unknown>, seen, cache)
: entry,
);
}
@@ -310,7 +340,7 @@ export function enforceStrictSchema(schema: Record<string, unknown>, seen?: Weak
Object.entries(defs).map(([name, def]) => [
name,
def != null && typeof def === "object" && !Array.isArray(def)
? enforceStrictSchema(def as Record<string, unknown>, seen)
? enforceStrictSchema(def as Record<string, unknown>, seen, cache)
: def,
]),
);
@@ -326,6 +356,7 @@ export function enforceStrictSchema(schema: Record<string, unknown>, seen?: Weak
) {
throw new Error("Schema node has no type, combinator, or $ref — cannot enforce strict mode");
}
seen.delete(schema);
return result;
}
@@ -86,6 +86,7 @@ describe("openai-completions compatibility", () => {
vercelGatewayRouting: {},
extraBody: {},
supportsStrictMode: true,
toolStrictMode: "none",
} satisfies Required<OpenAICompat>;
const assistantMessage: AssistantMessage = {
role: "assistant",
@@ -32,6 +32,7 @@ const compat: Required<OpenAICompat> = {
vercelGatewayRouting: {},
extraBody: {},
supportsStrictMode: true,
toolStrictMode: "none",
};
function buildToolResult(toolCallId: string, timestamp: number): ToolResultMessage {
@@ -37,9 +37,20 @@ function createAbortedSignal(): AbortSignal {
return controller.signal;
}
function captureCompletionsPayload(model: Model<"openai-completions">): Promise<unknown> {
function createSseResponse(events: unknown[]): Response {
const payload = `${events.map(event => `data: ${typeof event === "string" ? event : JSON.stringify(event)}`).join("\n\n")}\n\n`;
return new Response(payload, {
status: 200,
headers: { "content-type": "text/event-stream" },
});
}
function captureCompletionsPayload(
model: Model<"openai-completions">,
context: Context = testContext,
): Promise<unknown> {
const { promise, resolve } = Promise.withResolvers<unknown>();
streamOpenAICompletions(model, testContext, {
streamOpenAICompletions(model, context, {
apiKey: "test-key",
signal: createAbortedSignal(),
onPayload: payload => resolve(payload),
@@ -110,6 +121,117 @@ describe("OpenAI tool strict mode", () => {
expect(payload.stream_options).toBeUndefined();
});
it("uses uniformly non-strict tool schemas when provider requires all-or-none strictness", async () => {
const model: Model<"openai-completions"> = {
...getBundledModel("openai", "gpt-4o-mini"),
api: "openai-completions",
compat: { toolStrictMode: "all_strict" } satisfies OpenAICompat,
};
const context: Context = {
...testContext,
tools: [
testTool,
{
name: "dynamic_map",
description: "Dynamic object map",
parameters: Type.Object({
values: Type.Optional(Type.Record(Type.String(), Type.String())),
}),
},
],
};
const payload = (await captureCompletionsPayload(model, context)) as {
tools?: Array<{ function?: { strict?: boolean } }>;
};
expect(payload.tools).toHaveLength(2);
expect(payload.tools?.every(tool => tool.function?.strict === undefined)).toBe(true);
});
it("surfaces captured JSON error bodies when the SDK reports no body", async () => {
const model: Model<"openai-completions"> = {
...getBundledModel("openai", "gpt-4o-mini"),
api: "openai-completions",
};
global.fetch = Object.assign(
async (_input: string | URL | Request, _init?: RequestInit): Promise<Response> =>
new Response(
JSON.stringify({
message: "Tools with mixed values for 'strict' are not allowed.",
type: "invalid_request_error",
param: "tools",
code: "wrong_api_format",
}),
{
status: 422,
headers: { "content-type": "application/json" },
},
),
{ preconnect: originalFetch.preconnect },
);
const result = await streamOpenAICompletions(model, testContext, { apiKey: "test-key" }).result();
expect(result.stopReason).toBe("error");
expect(result.errorMessage).toContain("Tools with mixed values for 'strict' are not allowed.");
expect(result.errorMessage).toContain("param=tools");
expect(result.errorMessage).toContain("code=wrong_api_format");
});
it("retries with non-strict tool schemas after strict-mode request errors", async () => {
const model: Model<"openai-completions"> = {
...getBundledModel("openai", "gpt-4o-mini"),
api: "openai-completions",
compat: { toolStrictMode: "all_strict" } satisfies OpenAICompat,
};
const strictFlags: boolean[][] = [];
global.fetch = Object.assign(
async (_input: string | URL | Request, init?: RequestInit): Promise<Response> => {
const bodyText = typeof init?.body === "string" ? init.body : "";
const payload = JSON.parse(bodyText) as {
tools?: Array<{ function?: { strict?: boolean } }>;
};
strictFlags.push((payload.tools ?? []).map(tool => tool.function?.strict === true));
if (strictFlags.length === 1) {
return new Response(
JSON.stringify({
message: "Strict tool schema validation failed.",
type: "invalid_request_error",
param: "tools",
code: "wrong_api_format",
}),
{
status: 422,
headers: { "content-type": "application/json" },
},
);
}
return createSseResponse([
{
id: "chatcmpl-retry",
object: "chat.completion.chunk",
created: 0,
model: model.id,
choices: [{ index: 0, delta: { content: "Hello" } }],
},
{
id: "chatcmpl-retry",
object: "chat.completion.chunk",
created: 0,
model: model.id,
choices: [{ index: 0, delta: {}, finish_reason: "stop" }],
},
"[DONE]",
]);
},
{ preconnect: originalFetch.preconnect },
);
const result = await streamOpenAICompletions(model, testContext, { apiKey: "test-key" }).result();
expect(result.stopReason).toBe("stop");
expect(result.content).toContainEqual({ type: "text", text: "Hello" });
expect(strictFlags).toEqual([[true], [false]]);
});
it("sends strict=true for openai-responses tool schemas on OpenAI", async () => {
const model = getBundledModel("openai", "gpt-5-mini") as Model<"openai-responses">;
@@ -178,6 +178,42 @@ describe("enforceStrictSchema", () => {
expect(validBranch.additionalProperties).toBe(false);
});
it("reuses enforced object schemas across shared branches", () => {
const sharedTaskSchema = {
type: "object",
properties: {
content: { type: "string" },
notes: { type: "string" },
},
required: ["content"],
} as Record<string, unknown>;
const schema = {
type: "object",
properties: {
primary: {
type: "array",
items: sharedTaskSchema,
},
secondary: {
anyOf: [{ type: "array", items: sharedTaskSchema }, { type: "null" }],
},
},
required: ["primary", "secondary"],
} as Record<string, unknown>;
const strict = enforceStrictSchema(schema);
const rootProperties = strict.properties as Record<string, Record<string, unknown>>;
const primaryItems = rootProperties.primary.items as Record<string, unknown>;
const secondaryBranches = rootProperties.secondary.anyOf as Array<Record<string, unknown>>;
const secondaryItems = secondaryBranches[0]?.items as Record<string, unknown>;
expect(primaryItems.additionalProperties).toBe(false);
expect(primaryItems.required).toEqual(["content", "notes"]);
expect(secondaryItems.additionalProperties).toBe(false);
expect(secondaryItems.required).toEqual(["content", "notes"]);
expect(secondaryItems.properties).toEqual(primaryItems.properties);
});
it("treats type arrays containing object as object schemas via tryEnforceStrictSchema", () => {
const schema = {
type: ["object", "null"],
@@ -287,4 +323,45 @@ describe("tryEnforceStrictSchema", () => {
expect(result.strict).toBe(false);
expect(result.schema).toBe(schema);
});
it("keeps shared object schemas strict-compatible after adaptation", () => {
const sharedTaskSchema = Type.Object({
content: Type.String(),
status: Type.Optional(Type.String()),
notes: Type.Optional(Type.String()),
});
const schema = Type.Object({
ops: Type.Array(
Type.Union([
Type.Object({
op: Type.Literal("replace"),
tasks: Type.Array(sharedTaskSchema),
}),
Type.Object({
op: Type.Literal("update"),
tasks: Type.Optional(Type.Array(sharedTaskSchema)),
}),
]),
),
});
const result = tryEnforceStrictSchema(schema as unknown as Record<string, unknown>);
const rootProperties = result.schema.properties as Record<string, Record<string, unknown>>;
const opBranches = ((rootProperties.ops.items as Record<string, unknown>).anyOf ?? []) as Array<
Record<string, unknown>
>;
const replaceTasks = ((opBranches[0]?.properties as Record<string, Record<string, unknown>>)?.tasks?.items ??
{}) as Record<string, unknown>;
const updateTasks = (
((opBranches[1]?.properties as Record<string, Record<string, unknown>>)?.tasks?.anyOf ?? []) as Array<
Record<string, unknown>
>
)[0]?.items as Record<string, unknown>;
expect(result.strict).toBe(true);
expect(replaceTasks.additionalProperties).toBe(false);
expect(replaceTasks.required).toEqual(["content", "status", "notes"]);
expect(updateTasks.additionalProperties).toBe(false);
expect(updateTasks.required).toEqual(["content", "status", "notes"]);
});
});
+3 -1
View File
@@ -1,9 +1,9 @@
# Changelog
## [Unreleased]
### Breaking Changes
- Changed the `vim` tool API to require either `open: "path"` or `kbd: [...]` per call and removed direct `line`/`col` cursor parameters from `open`, so callers must position the cursor via key sequences after opening
- Changed the `edit` schemas for patch, replace, hashline, and chunk modes from top-level request fields to `edits` array entries, requiring path/mode details on each edit and breaking callers that send legacy top-level `path`, `old_text`, `new_text`, `op`, `move`, or `delete` payloads
### Added
@@ -16,6 +16,7 @@
### Changed
- Changed the `task` tool `schema` field to require JSON-encoded JTD schema text instead of a schema object, matching prompt guidance and task-subagent invocation
- Changed chunk edit payloads to encode selectors as `path: "file:selector"` and updated chunk tool guidance and examples to match
- Updated `edit` call/result rendering to show per-file diff sections and append a `(+N more)` hint when edits target multiple files
- Grouped chunk-mode `grep` results by directory, file, and chunk so directory searches now render as hierarchical sections (`#`/`##`) with per-chunk anchor lines
@@ -24,6 +25,7 @@
### Fixed
- Fixed OpenRouter model resolution to accept dated routed selectors such as `openrouter/z-ai/glm-4.7-20251222:nitro`, inheriting metadata from the base catalog model when the exact variant is not listed yet
- Fixed pre-execution edit preview routing so replace/patch/hashline mode diffs are computed from the new structured edit entries
- Adjusted chunk/hashline/prompt guidance and validation to align with the refactored per-entry schema
- Fixed chunk streaming output detection to verify chunk edits with `chunkToolEditSchema`, preventing non-chunk edit payloads from being rendered as chunk diffs
@@ -1,4 +1,4 @@
import { logger, Snowflake } from "@oh-my-pi/pi-utils";
import { logger } from "@oh-my-pi/pi-utils";
const DELIVERY_RETRY_BASE_MS = 500;
const DELIVERY_RETRY_MAX_MS = 30_000;
@@ -79,7 +79,7 @@ export function createAnalyzeFileTool(options: {
});
const taskParams: TaskParams = {
agent: "quick_task",
schema: analyzeFileOutputSchema,
schema: JSON.stringify(analyzeFileOutputSchema),
tasks,
};
return taskTool.execute(toolCallId, taskParams, signal, onUpdate);
@@ -51,6 +51,7 @@ const TRAILING_CANONICAL_MARKERS = [
"xhigh",
"free",
"exacto",
"nitro",
"original",
"optimized",
"nvfp4",
@@ -28,7 +28,7 @@ import {
import { isRecord, logger } from "@oh-my-pi/pi-utils";
import { type Static, Type } from "@sinclair/typebox";
import { type ConfigError, ConfigFile } from "../config";
import { parseModelString } from "../config/model-resolver";
import { parseModelString, resolveProviderModelReference } from "../config/model-resolver";
import { isValidThemeColor, type ThemeColor } from "../modes/theme/theme";
import type { AuthStorage, OAuthCredential } from "../session/auth-storage";
import {
@@ -160,6 +160,7 @@ const OpenAICompatSchema = Type.Object({
vercelGatewayRouting: Type.Optional(VercelGatewayRoutingSchema),
extraBody: Type.Optional(Type.Record(Type.String(), Type.Unknown())),
supportsStrictMode: Type.Optional(Type.Boolean()),
toolStrictMode: Type.Optional(Type.Union([Type.Literal("all_strict"), Type.Literal("none")])),
});
const EffortSchema = Type.Union([
@@ -1871,7 +1872,7 @@ export class ModelRegistry {
* Find a model by provider and ID.
*/
find(provider: string, modelId: string): Model<Api> | undefined {
return this.#models.find(m => m.provider === provider && m.id === modelId);
return resolveProviderModelReference(provider, modelId, this.#models);
}
/**
@@ -62,6 +62,101 @@ export function formatModelSelectorValue(selector: string, thinkingLevel: Thinki
return thinkingLevel && thinkingLevel !== ThinkingLevel.Inherit ? `${selector}:${thinkingLevel}` : selector;
}
function getOpenRouterRouteSuffix(modelId: string): { baseId: string; suffix: string } | undefined {
const colonIdx = modelId.lastIndexOf(":");
if (colonIdx === -1) {
return undefined;
}
const suffix = modelId.slice(colonIdx + 1).trim();
if (!suffix || parseThinkingLevel(suffix)) {
return undefined;
}
return { baseId: modelId.slice(0, colonIdx), suffix };
}
function stripOpenRouterDateSuffix(modelId: string): string | undefined {
const stripped = modelId.replace(/-\d{8}(?=$|:)/i, "");
return stripped !== modelId ? stripped : undefined;
}
function getOpenRouterFallbackModelIds(modelId: string): string[] {
const orderedCandidates: string[] = [];
const queue = [modelId];
const seen = new Set<string>();
while (queue.length > 0) {
const candidate = queue.shift();
if (!candidate || seen.has(candidate)) {
continue;
}
seen.add(candidate);
orderedCandidates.push(candidate);
const routedSuffix = getOpenRouterRouteSuffix(candidate);
if (routedSuffix) {
queue.push(routedSuffix.baseId);
}
const strippedDate = stripOpenRouterDateSuffix(candidate);
if (strippedDate) {
queue.push(strippedDate);
}
}
return orderedCandidates;
}
function cloneModelWithRequestedId(model: Model<Api>, requestedId: string): Model<Api> {
return {
...model,
id: requestedId,
...(model.name === model.id ? { name: requestedId } : {}),
};
}
export function resolveProviderModelReference(
provider: string,
modelId: string,
availableModels: readonly Model<Api>[],
): Model<Api> | undefined {
const normalizedProvider = provider.trim().toLowerCase();
const normalizedModelId = modelId.trim().toLowerCase();
if (!normalizedProvider || !normalizedModelId) {
return undefined;
}
const exactMatches = availableModels.filter(
model => model.provider.toLowerCase() === normalizedProvider && model.id.toLowerCase() === normalizedModelId,
);
if (exactMatches.length === 1) {
return exactMatches[0];
}
if (exactMatches.length > 1) {
return undefined;
}
if (normalizedProvider !== "openrouter") {
return undefined;
}
for (const fallbackId of getOpenRouterFallbackModelIds(modelId).slice(1)) {
const baseMatches = availableModels.filter(
model =>
model.provider.toLowerCase() === normalizedProvider && model.id.toLowerCase() === fallbackId.toLowerCase(),
);
if (baseMatches.length === 1) {
return cloneModelWithRequestedId(baseMatches[0], modelId);
}
if (baseMatches.length > 1) {
return undefined;
}
}
return undefined;
}
export interface ModelMatchPreferences {
/** Most-recently-used model keys (provider/modelId) to prefer when ambiguous. */
usageOrder?: string[];
@@ -171,17 +266,7 @@ export function findExactModelReferenceMatch(
const provider = trimmedReference.substring(0, slashIndex).trim();
const modelId = trimmedReference.substring(slashIndex + 1).trim();
if (provider && modelId) {
const providerMatches = availableModels.filter(
model =>
model.provider.toLowerCase() === provider.toLowerCase() &&
model.id.toLowerCase() === modelId.toLowerCase(),
);
if (providerMatches.length === 1) {
return providerMatches[0];
}
if (providerMatches.length > 1) {
return undefined;
}
return resolveProviderModelReference(provider, modelId, availableModels);
}
}
return undefined;
@@ -853,10 +938,8 @@ export function resolveCliModel(options: {
let exact: (typeof availableModels)[number] | undefined;
if (slashIdx !== -1) {
const prefix = lower.substring(0, slashIdx);
const suffix = lower.substring(slashIdx + 1);
exact = availableModels.find(
model => model.provider.toLowerCase() === prefix && model.id.toLowerCase() === suffix,
);
const suffix = trimmedModel.substring(slashIdx + 1);
exact = resolveProviderModelReference(prefix, suffix, availableModels);
}
if (!exact && !trimmedModel.includes(":")) {
const canonicalMatch = modelRegistry.resolveCanonicalModel?.(trimmedModel, { availableOnly: false });
@@ -905,6 +988,19 @@ export function resolveCliModel(options: {
}
}
if (provider) {
const exactProviderMatch = resolveProviderModelReference(provider, pattern, availableModels);
if (exactProviderMatch) {
return {
model: exactProviderMatch,
selector: formatModelString(exactProviderMatch),
warning: undefined,
thinkingLevel: undefined,
error: undefined,
};
}
}
const candidates = provider ? availableModels.filter(model => model.provider === provider) : availableModels;
const { model, thinkingLevel, warning } = parseModelPattern(pattern, candidates, preferences, {
allowInvalidThinkingSelectorFallback: false,
+3 -1
View File
@@ -602,7 +602,9 @@ export async function waitForProjectLoaded(client: LspClient, signal?: AbortSign
if (signal?.aborted) return;
await Promise.race([
client.projectLoaded,
...(signal ? [new Promise<void>(resolve => signal.addEventListener("abort", () => resolve(), { once: true }))] : []),
...(signal
? [new Promise<void>(resolve => signal.addEventListener("abort", () => resolve(), { once: true }))]
: []),
]);
}
+1 -1
View File
@@ -19,8 +19,8 @@ import {
sendRequest,
setIdleTimeout,
syncContent,
waitForProjectLoaded,
WARMUP_TIMEOUT_MS,
waitForProjectLoaded,
} from "./client";
import { getLinterClient } from "./clients";
import { getServersForFile, type LspConfig, loadConfig } from "./config";
@@ -315,7 +315,8 @@ export class ToolExecutionComponent extends Container {
*/
#updateSpinnerAnimation(): void {
// Spinner for: task tool with partial result, or edit/write while args streaming
const isStreamingArgs = !this.#argsComplete && (this.#toolName === "edit" || this.#toolName === "write" || this.#toolName === "vim");
const isStreamingArgs =
!this.#argsComplete && (this.#toolName === "edit" || this.#toolName === "write" || this.#toolName === "vim");
const isBackgroundAsyncTask =
this.#toolName === "task" &&
(this.#result?.details as { async?: { state?: string } } | undefined)?.async?.state === "running";
@@ -1,5 +1,4 @@
Edits files via syntax-aware chunks. Run `read(path="file.ts")` first.
- `write` rewrites the entire targeted region — best for most edits.
- `replace` does surgical find-and-replace within a chunk — use when making small changes to a large chunk, or batching multiple substitutions.
@@ -13,7 +13,7 @@ Subagents lack your conversation history. Every decision, file content, and user
- `.description`: UI display only — subagent never sees it
- `.assignment`: Complete self-contained instructions. One-liners PROHIBITED; missing acceptance criteria = too vague.
- `context`: Shared background prepended to every assignment. Session-specific info only.
- `schema`: JTD schema for expected output. Format lives here — **MUST NOT** be duplicated in assignments.
- `schema`: JSON-encoded JTD schema for expected output. Format lives here — **MUST NOT** be duplicated in assignments.
- `tasks`: Tasks to execute in parallel.
- `isolated`: Run in isolated environment; returns patches. Use when tasks edit overlapping files.
</parameters>
+22 -22
View File
@@ -1,30 +1,30 @@
Stateful single-buffer Vim-style editor.
Stateful Vim-style editor with multi-buffer support.
Use this for surgical text edits when motions and compact viewport feedback are more efficient than rewriting full regions.
Every call requires `file` — the path to edit. The buffer is loaded automatically on first use.
Actions:
- `open`: load a file into the buffer (auto-saves any previous buffer)
- `kbd`: run Vim key sequences, then optionally insert literal text
- `{"file": "path/to/file.py"}` — view file (loads buffer if needed)
- `{"file": "path/to/file.py", "kbd": ["…"], "insert": "…"}` — edit file
How `kbd` + `insert` work together:
- `kbd` runs Vim key sequences (motions, commands, operators)
- `insert` is **raw text** (with real `\n` newlines in JSON) that gets typed into the buffer
- For `insert` to work, the last `kbd` entry **MUST** leave the buffer in INSERT mode (via `i`, `o`, `O`, `a`, `A`, `cc`, `C`, `s`, `S`, etc.)
- After the call, the tool auto-exits INSERT mode and auto-saves to disk
- Set `pause: true` to skip auto-save and stay in the current mode
Rules:
- One active buffer at a time; `open` replaces it
- Each `kbd` call auto-saves to disk (unless `pause: true`)
- `kbd` array entries run in order; every non-final entry must leave NORMAL mode — if an entry enters INSERT mode, end it with `<Esc>` or merge into one string
- `insert` is **raw text** (newlines = real `\n` in JSON), NOT Vim key syntax; the buffer must already be in INSERT mode (via `i`, `o`, `O`, `a`, `A`, `cc`, etc.)
- After `insert`, the tool exits INSERT mode and saves automatically (unless `pause: true`)
- `pause: true` keeps the current mode active and skips auto-save; use it for multi-step edits
- Use `:e!` to reload from disk and discard unsaved changes
- Each non-final `kbd` entry must end in NORMAL mode — use `<Esc>` or merge into one string
- To recover from mistakes: `{"file": "f.py", "kbd": ["u"]}` to undo, or `{"file": "f.py", "kbd": [":e!<CR>"]}` to reload from disk
Supported Vim subset: motions (`h/j/k/l`, `w/b/e`, `0/$`, `gg/G`, `{/}`, `f/t`), counts, `.` repeat, insert commands (`i/a/o/O/I/A/cc/C/s/S`), visual mode (`v/V`), operators (`d/c/y/p`), text objects (`iw/aw/i"/a"/i(/a(`), undo/redo (`u`/`<C-r>`), search (`/pattern<CR>`, `n/N`), ex commands (`:s`, `:%s`, `:e`, `:e!`, ranged `:d`).
Supported: motions (`h/j/k/l`, `w/b/e`, `0/$`, `gg/G`, `{/}`, `f/t`), counts, `.` repeat, insert (`i/a/o/O/I/A/cc/C/s/S`), visual (`v/V`), operators (`d/c/y/p`), text objects (`iw/aw/i"/a"/i(/a(`), undo/redo (`u`/`<C-r>`), search (`/pattern<CR>`, `n/N`), ex (`:s`, `:%s`, `:e`, `:e!`, ranged `:d`).
Special keys: `<Esc>` or `<Escape>`, `<CR>` or `<Enter>`, `<BS>`, `<Tab>`, `<C-d>`, `<C-u>`, `<C-r>`, `<C-w>`, `<C-o>`.
Special keys: `<Esc>`, `<Escape>`, `<CR>`, `<Enter>`, `<BS>`, `<Tab>`, `<C-d>`, `<C-u>`, `<C-r>`, `<C-w>`, `<C-o>`.
Examples:
- Open file: `{"open":"src/app.ts"}`
- Open at line: `{"open":"src/app.ts", "line":42}`
- Rename word: `{"kbd":["42G", "ciwnewName<Esc>"]}`
- Replace line with multi-line text: `{"kbd":["5G", "cc"], "insert":" if b == 0:\n return None"}`
- Add lines below: `{"kbd":["3G", "o"], "insert":"def multiply(a, b):\n return a * b"}`
- Global substitution: `{"kbd":[":%s/oldName/newName/g<CR>"]}`
- Search and delete: `{"kbd":["/TODO<CR>", "dd"]}`
- Delete range of lines: `{"kbd":[":3,5d<CR>"]}`
- `{"file": "src/app.ts"}` — view file
- `{"file": "src/app.ts", "kbd": ["3G", "ciwnewName<Esc>"]}` — rename word on line 3
- `{"file": "src/app.ts", "kbd": ["5G", "cc"], "insert": " if b == 0:\n return None"}` — replace line 5
- `{"file": "src/app.ts", "kbd": ["3G", "o"], "insert": "def multiply(a, b):\n return a * b"}` — insert after line 3
- `{"file": "src/app.ts", "kbd": [":%s/oldName/newName/g<CR>"]}` — find and replace
- `{"file": "src/app.ts", "kbd": ["/TODO<CR>", "dd"]}` — search and delete
- `{"file": "src/app.ts", "kbd": [":3,5d<CR>"]}` — delete line range
+2 -2
View File
@@ -82,9 +82,9 @@ const createTaskSchema = (options: { isolationEnabled: boolean }) => {
}),
),
schema: Type.Optional(
Type.Record(Type.String(), Type.Unknown(), {
Type.String({
description:
"JTD schema defining expected response structure. Use typed properties. Output format belongs here — never in context or assignment.",
"JSON-encoded JTD schema defining expected response structure. Output format belongs here — never in context or assignment.",
}),
),
tasks: Type.Array(taskItemSchema, {
+1 -3
View File
@@ -162,9 +162,7 @@ export class GrepTool implements AgentTool<typeof grepSchema, GrepToolDetails> {
const stat = await Bun.file(searchPath).stat();
isDirectory = stat.isDirectory();
} catch {
const hint = scopePath.includes(",")
? ` (comma-separated paths must each exist relative to cwd)`
: "";
const hint = scopePath.includes(",") ? ` (comma-separated paths must each exist relative to cwd)` : "";
throw new ToolError(`Path not found: ${scopePath}${hint}`);
}
+2 -3
View File
@@ -23,10 +23,8 @@ import { SearchTool } from "../web/search";
import { AskTool } from "./ask";
import { AstEditTool } from "./ast-edit";
import { AstGrepTool } from "./ast-grep";
import { PollTool } from "./poll-tool";
import { BashTool } from "./bash";
import { BrowserTool } from "./browser";
import { CalculatorTool } from "./calculator";
import { CancelJobTool } from "./cancel-job";
import { type CheckpointState, CheckpointTool, RewindTool } from "./checkpoint";
@@ -48,6 +46,7 @@ import { GrepTool } from "./grep";
import { InspectImageTool } from "./inspect-image";
import { NotebookTool } from "./notebook";
import { wrapToolWithMetaNotice } from "./output-meta";
import { PollTool } from "./poll-tool";
import { PythonTool } from "./python";
import { ReadTool } from "./read";
import { RenderMermaidTool } from "./render-mermaid";
@@ -73,7 +72,6 @@ export * from "../web/search";
export * from "./ask";
export * from "./ast-edit";
export * from "./ast-grep";
export * from "./poll-tool";
export * from "./bash";
export * from "./browser";
export * from "./calculator";
@@ -87,6 +85,7 @@ export * from "./gh";
export * from "./grep";
export * from "./inspect-image";
export * from "./notebook";
export * from "./poll-tool";
export * from "./python";
export * from "./read";
export * from "./render-mermaid";
@@ -57,7 +57,7 @@ export class SubmitResultTool implements AgentTool<TSchema, SubmitResultDetails>
readonly label = "Submit Result";
readonly description =
"Finish the task with structured JSON output. Call exactly once at the end of the task.\n\n" +
"Pass `result: { data: <your output> }` for success, or `result: { error: \"message\" }` for failure.\n" +
'Pass `result: { data: <your output> }` for success, or `result: { error: "message" }` for failure.\n' +
"The `data`/`error` wrapper is required — do not put your output directly in `result`.";
readonly parameters: TSchema;
strict = true;
@@ -173,7 +173,7 @@ export class SubmitResultTool implements AgentTool<TSchema, SubmitResultDetails>
}
if (errorMessage === undefined && data === undefined) {
throw new Error(
"result must contain either `data` or `error`. Use `{result: {data: <your output>}}` for success or `{result: {error: \"message\"}}` for failure.",
'result must contain either `data` or `error`. Use `{result: {data: <your output>}}` for success or `{result: {error: "message"}}` for failure.',
);
}
+111 -116
View File
@@ -37,24 +37,27 @@ const INTERNAL_URL_PREFIX = /^(agent|artifact|skill|rule|local|mcp):\/\//;
const utf8Decoder = new TextDecoder("utf-8", { fatal: true });
const vimSchema = Type.Object({
open: Type.Optional(Type.String({ description: "File path to open" })),
line: Type.Optional(Type.Number({ description: "1-indexed line to place cursor on open" })),
col: Type.Optional(Type.Number({ description: "1-indexed column to place cursor on open" })),
kbd: Type.Optional(Type.Array(Type.String(), { description: "Vim key sequences to execute" })),
insert: Type.Optional(Type.String({ description: "Raw text to insert literally while in INSERT mode" })),
pause: Type.Optional(Type.Boolean({ description: "Return an intermediate snapshot without forcing a mode exit" })),
file: Type.String({ description: "File path to edit." }),
kbd: Type.Optional(
Type.Array(Type.String(), {
description: "Vim key sequences to execute against the buffer. Null when just viewing the file.",
}),
),
insert: Type.Optional(
Type.String({
description:
"Raw text to type into the buffer. kbd must leave INSERT mode active first (e.g. via o, O, i, cc). Null when not inserting.",
}),
),
pause: Type.Optional(
Type.Boolean({
description: "If true, skip auto-save and keep current mode. Null or false for normal auto-save.",
}),
),
});
type VimParams = Static<typeof vimSchema>;
function isOpenParams(params: VimParams): boolean {
return params.open !== undefined;
}
function isKbdParams(params: VimParams): boolean {
return params.kbd! !== undefined;
}
function fingerprintEqual(left: VimFingerprint | null, right: VimFingerprint | null): boolean {
if (left === null || right === null) {
return left === right;
@@ -216,7 +219,7 @@ export class VimTool implements AgentTool<typeof vimSchema, VimToolDetails> {
readonly parameters = vimSchema;
readonly concurrency = "exclusive";
#engine: VimEngine | null = null;
#engines = new Map<string, VimEngine>();
#writethrough: WritethroughCallback;
constructor(private readonly session: ToolSession) {
@@ -329,125 +332,117 @@ export class VimTool implements AgentTool<typeof vimSchema, VimToolDetails> {
_context?: AgentToolContext,
): Promise<AgentToolResult<VimToolDetails>> {
return untilAborted(signal, async () => {
if (isOpenParams(params)) {
// Auto-save previous buffer before opening new file
if (this.#engine?.buffer.modified) {
await this.#saveBuffer(this.#engine.buffer);
}
this.#engine = null;
const loaded = await this.#loadBuffer(params.open!);
const engine = new VimEngine(new VimBuffer(loaded), {
// Resolve file path and get-or-create engine for this buffer
const { absolutePath } = normalizeTargetPath(params.file, this.session.cwd);
let engine = this.#engines.get(absolutePath);
let isNewBuffer = false;
if (!engine) {
const loaded = await this.#loadBuffer(params.file);
engine = new VimEngine(new VimBuffer(loaded), {
beforeMutate: buffer => this.#beforeMutate(buffer),
loadBuffer: path => this.#loadBuffer(path),
saveBuffer: (buffer, options) => this.#saveBuffer(buffer, options),
});
if (params.line || params.col) {
engine.setCursor(Math.max(0, (params.line ?? 1) - 1), Math.max(0, (params.col ?? 1) - 1));
}
engine.viewportStart = params.line ? Math.max(1, params.line - 20) : 1;
engine.statusMessage = `Opened ${engine.buffer.displayPath}`;
this.#engine = engine;
return this.#renderFromEngine(
engine,
VIM_OPEN_VIEWPORT_LINES,
params.line ? Math.max(1, params.line - 20) : 1,
);
engine.viewportStart = 1;
this.#engines.set(absolutePath, engine);
isNewBuffer = true;
}
if (isKbdParams(params)) {
if (!this.#engine) {
throw new ToolError("No active vim buffer. Open a file first.");
const sequences = Array.isArray(params.kbd) ? params.kbd : undefined;
if (!sequences) {
// No kbd — just show the file viewport
if (isNewBuffer) {
engine.statusMessage = `Opened ${engine.buffer.displayPath}`;
}
const engine = this.#engine;
const sequences = params.kbd!;
const commandText = sequences.join(" ");
const tokenGroups = splitTokensBySequence(sequences);
const beforeText = serializeBufferText(engine.buffer);
return this.#renderFromEngine(engine, VIM_OPEN_VIEWPORT_LINES, engine.viewportStart);
}
if (this.session.getPlanModeState?.()?.enabled) {
if (params.insert !== undefined) {
throw new ToolError("Plan mode: vim is read-only; insert payloads are not allowed.");
}
const preview = engine.clone({
beforeMutate: async () => {
throw new VimInputError(
"Plan mode: vim is read-only; only navigation, search, open, and close are allowed.",
);
},
saveBuffer: async () => {
throw new VimInputError("Plan mode: :w is not allowed.");
},
});
await executeKeySequences(preview, tokenGroups, commandText);
// Execute kbd sequences
const commandText = sequences.join(" ");
const tokenGroups = splitTokensBySequence(sequences);
const beforeText = serializeBufferText(engine.buffer);
if (this.session.getPlanModeState?.()?.enabled) {
if (params.insert !== undefined) {
throw new ToolError("Plan mode: vim is read-only; insert payloads are not allowed.");
}
const preview = engine.clone({
beforeMutate: async () => {
throw new VimInputError(
"Plan mode: vim is read-only; only navigation, search, open, and close are allowed.",
);
},
saveBuffer: async () => {
throw new VimInputError("Plan mode: :w is not allowed.");
},
});
await executeKeySequences(preview, tokenGroups, commandText);
}
try {
const FRAME_INTERVAL_MS = 16; // ~60fps
let lastUpdateTime = 0;
try {
const FRAME_INTERVAL_MS = 16; // ~60fps
let lastUpdateTime = 0;
const emitUpdate = onUpdate
? async () => {
const now = Date.now();
if (now - lastUpdateTime < FRAME_INTERVAL_MS) {
return; // throttle: skip if too soon
}
onUpdate(this.#renderFromEngine(engine, VIM_DEFAULT_VIEWPORT_LINES, engine.viewportStart));
lastUpdateTime = Date.now();
await Bun.sleep(FRAME_INTERVAL_MS); // real delay for terminal to render
const emitUpdate = onUpdate
? async () => {
const now = Date.now();
if (now - lastUpdateTime < FRAME_INTERVAL_MS) {
return; // throttle: skip if too soon
}
: undefined;
onUpdate(this.#renderFromEngine(engine, VIM_DEFAULT_VIEWPORT_LINES, engine.viewportStart));
lastUpdateTime = Date.now();
await Bun.sleep(FRAME_INTERVAL_MS); // real delay for terminal to render
}
: undefined;
await executeKeySequences(engine, tokenGroups, commandText, emitUpdate);
await executeKeySequences(engine, tokenGroups, commandText, emitUpdate);
if (!engine.closed && params.insert !== undefined) {
await engine.applyLiteralInsert(params.insert, params.pause !== true);
await emitUpdate?.();
}
if (!engine.closed && params.insert !== undefined) {
await engine.applyLiteralInsert(params.insert, params.pause !== true);
await emitUpdate?.();
}
if (params.pause === true && !engine.closed && engine.getPendingInput()) {
engine.statusMessage = engine.statusMessage ?? `Paused in ${engine.getPublicMode()} mode`;
if (params.pause === true && !engine.closed && engine.getPendingInput()) {
engine.statusMessage = engine.statusMessage ?? `Paused in ${engine.getPublicMode()} mode`;
}
} catch (error) {
this.#throwWithSnapshot(engine, error);
}
if (beforeText !== serializeBufferText(engine.buffer)) {
engine.centerViewportOnCursor();
}
// Auto-save when buffer was modified
if (!engine.closed && engine.buffer.modified && params.pause !== true) {
try {
const result = await this.#saveBuffer(engine.buffer);
engine.buffer.markSaved(result.loaded);
engine.diagnostics = result.diagnostics;
if (beforeText !== serializeBufferText(engine.buffer)) {
engine.centerViewportOnCursor();
}
} catch (error) {
this.#throwWithSnapshot(engine, error);
}
if (beforeText !== serializeBufferText(engine.buffer)) {
engine.centerViewportOnCursor();
}
// Auto-save when buffer was modified
if (!engine.closed && engine.buffer.modified && params.pause !== true) {
try {
const result = await this.#saveBuffer(engine.buffer);
engine.buffer.markSaved(result.loaded);
engine.diagnostics = result.diagnostics;
if (beforeText !== serializeBufferText(engine.buffer)) {
engine.centerViewportOnCursor();
}
} catch (error) {
this.#throwWithSnapshot(engine, error);
}
}
const afterText = serializeBufferText(engine.buffer);
const modelDiff = buildModelDiff(beforeText, afterText);
const result = this.#renderFromEngine(
engine,
VIM_DEFAULT_VIEWPORT_LINES,
engine.viewportStart,
engine.closed,
undefined,
undefined,
modelDiff,
);
if (engine.closed) {
this.#engine = null;
}
return result;
}
throw new ToolError("Invalid vim parameters");
const afterText = serializeBufferText(engine.buffer);
const modelDiff = buildModelDiff(beforeText, afterText);
const result = this.#renderFromEngine(
engine,
VIM_DEFAULT_VIEWPORT_LINES,
engine.viewportStart,
engine.closed,
undefined,
undefined,
modelDiff,
);
if (engine.closed) {
this.#engines.delete(absolutePath);
}
return result;
});
}
}
@@ -535,7 +530,7 @@ function getInsertForDisplay(args: VimRenderArgs): string | undefined {
}
interface VimRenderArgs {
open?: string;
file?: string;
kbd?: string[];
insert?: string;
pause?: boolean;
@@ -544,8 +539,8 @@ interface VimRenderArgs {
export const vimToolRenderer = {
renderCall(args: VimRenderArgs, options: RenderResultOptions, uiTheme: Theme): Component {
if (args.open) {
return renderText(`${uiTheme.bold("Vim")} open ${args.open}`);
if (args.file && !args.kbd) {
return renderText(`${uiTheme.bold("Vim")} open ${args.file}`);
}
// Build a description of the streaming args for the header
@@ -116,7 +116,6 @@ describe("executeBash", () => {
expect(() => process.kill(pid, "SIGKILL")).not.toThrow();
});
it("times out commands", async () => {
if (process.platform === "win32") {
return;
@@ -189,6 +189,17 @@ describe("ModelRegistry", () => {
expect(sonnetVariants.some(variant => variant.selector === "demo/claude-4.5-sonnet")).toBe(true);
});
test("collapses nitro-suffixed OpenRouter variants under the upstream canonical id", () => {
writeRawModelsJson({
openrouter: providerConfig("https://openrouter.ai/api/v1", [{ id: "z-ai/glm-4.7-20251222:nitro" }]),
});
const registry = new ModelRegistry(authStorage, modelsJsonPath);
const variants = registry.getCanonicalVariants("glm-4.7");
expect(variants.some(variant => variant.selector === "openrouter/z-ai/glm-4.7-20251222:nitro")).toBe(true);
});
test("collapses anthropic latest aliases into the best upstream claude family id", () => {
writeRawModelsJson({
demo: providerConfig("https://demo.example.com/v1", [
@@ -336,6 +347,21 @@ describe("ModelRegistry", () => {
});
});
describe("OpenRouter routed suffix fallback", () => {
test("find synthesizes a routed model id from the base OpenRouter metadata", () => {
writeRawModelsJson({
openrouter: providerConfig("https://openrouter.ai/api/v1", [{ id: "z-ai/glm-4.7" }]),
});
const registry = new ModelRegistry(authStorage, modelsJsonPath);
const model = registry.find("openrouter", "z-ai/glm-4.7-20251222:nitro");
expect(model?.provider).toBe("openrouter");
expect(model?.id).toBe("z-ai/glm-4.7-20251222:nitro");
expect(model?.name).toBe("z-ai/glm-4.7-20251222:nitro");
});
});
describe("baseUrl override (no custom models)", () => {
test("overriding baseUrl keeps all built-in models", () => {
writeRawModelsJson({
@@ -77,6 +77,23 @@ const mockOpenRouterModels: Model<"anthropic-messages">[] = [
contextWindow: 128000,
maxTokens: 4096,
},
{
id: "z-ai/glm-4.7",
name: "GLM 4.7",
api: "anthropic-messages",
provider: "openrouter",
baseUrl: "https://openrouter.ai/api/v1",
reasoning: true,
thinking: {
mode: "budget",
minLevel: Effort.Minimal,
maxLevel: Effort.High,
},
input: ["text"],
cost: { input: 1, output: 2, cacheRead: 0.1, cacheWrite: 1 },
contextWindow: 128000,
maxTokens: 8192,
},
];
const mockProviderOverlapModels: Model<"anthropic-messages">[] = [
@@ -319,6 +336,24 @@ describe("parseModelPattern", () => {
expect(result.explicitThinkingLevel).toBe(false);
expect(result.warning).toBeUndefined();
});
test("supports OpenRouter route suffixes that are not present in the catalog", () => {
const result = parseModelPattern("openrouter/z-ai/glm-4.7-20251222:nitro", allModels);
expect(result.model?.provider).toBe("openrouter");
expect(result.model?.id).toBe("z-ai/glm-4.7-20251222:nitro");
expect(result.thinkingLevel).toBeUndefined();
expect(result.explicitThinkingLevel).toBe(false);
expect(result.warning).toBeUndefined();
});
test("supports OpenRouter route suffixes with an appended thinking level", () => {
const result = parseModelPattern("openrouter/z-ai/glm-4.7-20251222:nitro:high", allModels);
expect(result.model?.provider).toBe("openrouter");
expect(result.model?.id).toBe("z-ai/glm-4.7-20251222:nitro");
expect(result.thinkingLevel).toBe(Effort.High);
expect(result.explicitThinkingLevel).toBe(true);
expect(result.warning).toBeUndefined();
});
});
describe("invalid thinking levels with OpenRouter models", () => {
@@ -646,6 +681,37 @@ describe("resolveCliModel", () => {
expect(result.error).toContain("not found");
});
test("supports provider-prefixed OpenRouter route suffixes even when the base model is cataloged without them", () => {
const registry = {
getAll: () => allModels,
} as unknown as Parameters<typeof resolveCliModel>[0]["modelRegistry"];
const result = resolveCliModel({
cliModel: "openrouter/z-ai/glm-4.7-20251222:nitro",
modelRegistry: registry,
});
expect(result.error).toBeUndefined();
expect(result.model?.provider).toBe("openrouter");
expect(result.model?.id).toBe("z-ai/glm-4.7-20251222:nitro");
});
test("supports explicit OpenRouter provider with route suffixes that are not in the catalog", () => {
const registry = {
getAll: () => allModels,
} as unknown as Parameters<typeof resolveCliModel>[0]["modelRegistry"];
const result = resolveCliModel({
cliProvider: "openrouter",
cliModel: "z-ai/glm-4.7-20251222:nitro",
modelRegistry: registry,
});
expect(result.error).toBeUndefined();
expect(result.model?.provider).toBe("openrouter");
expect(result.model?.id).toBe("z-ai/glm-4.7-20251222:nitro");
});
test("returns a clear error when there are no models", () => {
const registry = {
getAll: () => [],
+1 -1
View File
@@ -10,12 +10,12 @@ import { DEFAULT_BASH_INTERCEPTOR_RULES, Settings } from "@oh-my-pi/pi-coding-ag
import { EditTool } from "@oh-my-pi/pi-coding-agent/edit";
import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager";
import type { ToolSession } from "@oh-my-pi/pi-coding-agent/tools";
import { PollTool } from "@oh-my-pi/pi-coding-agent/tools/poll-tool";
import { BashTool } from "@oh-my-pi/pi-coding-agent/tools/bash";
import { CancelJobTool } from "@oh-my-pi/pi-coding-agent/tools/cancel-job";
import { FindTool } from "@oh-my-pi/pi-coding-agent/tools/find";
import { GrepTool } from "@oh-my-pi/pi-coding-agent/tools/grep";
import { wrapToolWithMetaNotice } from "@oh-my-pi/pi-coding-agent/tools/output-meta";
import { PollTool } from "@oh-my-pi/pi-coding-agent/tools/poll-tool";
import { ReadTool } from "@oh-my-pi/pi-coding-agent/tools/read";
import { WriteTool } from "@oh-my-pi/pi-coding-agent/tools/write";
import * as markitUtils from "@oh-my-pi/pi-coding-agent/utils/markit";
@@ -77,6 +77,19 @@ function formatCompatibilityIssues(
}
describe("builtin tool schemas provider compatibility", () => {
it("keeps task and todo_write strict-compatible for OpenAI-style providers", async () => {
const toolSchemas = await collectToolSchemas();
for (const toolName of ["task", "todo_write"]) {
const entry = toolSchemas.find(tool => tool.name === toolName);
expect(entry).toBeDefined();
if (!entry) {
continue;
}
const strictResult = adaptSchemaForStrict(entry.schema, true);
expect(strictResult.strict).toBe(true);
}
});
it("keeps all builtin and hidden tool schemas valid after provider enforcement", async () => {
const toolSchemas = await collectToolSchemas();
const failures: string[] = [];
+31 -31
View File
@@ -159,9 +159,9 @@ describe("vim tool", () => {
await Bun.write(filePath, "foo = 1;\nfoo = foo + 1;\n");
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "sample.ts" });
await tool.execute("edit", { kbd: ["ciwbar<Esc>", "j", "."] });
await tool.execute("save", { kbd: [":w<CR>"] });
await tool.execute("open", { file: "sample.ts" });
await tool.execute("edit", { file: "sample.ts", kbd: ["ciwbar<Esc>", "j", "."] });
await tool.execute("save", { file: "sample.ts", kbd: [":w<CR>"] });
const saved = await Bun.file(filePath).text();
expect(saved).toContain("bar = 1;");
@@ -173,8 +173,8 @@ describe("vim tool", () => {
await Bun.write(filePath, Array.from({ length: 1100 }, (_, index) => `line ${index + 1};`).join("\n"));
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "long.ts" });
const moved = await tool.execute("jump", { kbd: ["1014G"] });
await tool.execute("open", { file: "long.ts" });
const moved = await tool.execute("jump", { file: "long.ts", kbd: ["1014G"] });
const text = textResult(moved);
expect(text).toContain(">1014│line 1014;");
expect(moved.details?.cursor.line).toBe(1014);
@@ -185,8 +185,8 @@ describe("vim tool", () => {
await Bun.write(filePath, Array.from({ length: 500 }, (_, index) => `line ${index + 1};`).join("\n"));
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "center.ts" });
const edited = await tool.execute("edit", { kbd: ["386Go"], insert: "inserted", pause: true });
await tool.execute("open", { file: "center.ts" });
const edited = await tool.execute("edit", { file: "center.ts", kbd: ["386Go"], insert: "inserted", pause: true });
expect(edited.details?.cursor.line).toBe(387);
expect(edited.details?.viewport.start).toBe(367);
expect(edited.details?.viewport.end).toBe(406);
@@ -199,8 +199,8 @@ describe("vim tool", () => {
await Bun.write(filePath, Array.from({ length: 1100 }, (_, index) => `line ${index + 1};`).join("\n"));
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "long-edit.ts" });
const edited = await tool.execute("edit", { kbd: ["1014G", "o"], insert: "inserted" });
await tool.execute("open", { file: "long-edit.ts" });
const edited = await tool.execute("edit", { file: "long-edit.ts", kbd: ["1014G", "o"], insert: "inserted" });
const text = textResult(edited);
expect(edited.details?.cursor.line).toBe(1015);
expect(edited.details?.viewport.start).toBe(995);
@@ -213,9 +213,9 @@ describe("vim tool", () => {
await Bun.write(filePath, "first\nsecond\n");
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "replace.ts" });
const replaced = await tool.execute("replace", { kbd: ["cc"], insert: "alpha\nbeta" });
await tool.execute("save", { kbd: [":w<CR>"] });
await tool.execute("open", { file: "replace.ts" });
const replaced = await tool.execute("replace", { file: "replace.ts", kbd: ["cc"], insert: "alpha\nbeta" });
await tool.execute("save", { file: "replace.ts", kbd: [":w<CR>"] });
const saved = await Bun.file(filePath).text();
expect(saved).toBe("alpha\nbeta\nsecond\n");
@@ -228,8 +228,8 @@ describe("vim tool", () => {
await Bun.write(filePath, "first\n");
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "ambiguous.ts" });
await expect(tool.execute("bad", { kbd: ["o", "o"] })).rejects.toThrow(/left Vim in INSERT mode/i);
await tool.execute("open", { file: "ambiguous.ts" });
await expect(tool.execute("bad", { file: "ambiguous.ts", kbd: ["o", "o"] })).rejects.toThrow(/left Vim in INSERT mode/i);
});
it("rejects additional kbd entries after entering insert mode", async () => {
@@ -237,8 +237,8 @@ describe("vim tool", () => {
await Bun.write(filePath, "alpha\nbeta\n");
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "insert-boundary.ts" });
await expect(tool.execute("edit", { kbd: ["2G", "o", "o"] })).rejects.toThrow(/insert field|<Esc>/i);
await tool.execute("open", { file: "insert-boundary.ts" });
await expect(tool.execute("edit", { file: "insert-boundary.ts", kbd: ["2G", "o", "o"] })).rejects.toThrow(/insert field|<Esc>/i);
const saved = await Bun.file(filePath).text();
expect(saved).toBe("alpha\nbeta\n");
});
@@ -248,13 +248,13 @@ describe("vim tool", () => {
await Bun.write(filePath, "first\n");
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "pause.ts" });
const paused = await tool.execute("pause", { kbd: ["cc"], pause: true });
await tool.execute("open", { file: "pause.ts" });
const paused = await tool.execute("pause", { file: "pause.ts", kbd: ["cc"], pause: true });
expect(paused.details?.mode).toBe("INSERT");
expect(textResult(paused)).toContain("Pending: INSERT mode");
await tool.execute("resume", { kbd: [], insert: "replacement" });
await tool.execute("save", { kbd: [":w<CR>"] });
await tool.execute("resume", { file: "pause.ts", kbd: [], insert: "replacement" });
await tool.execute("save", { file: "pause.ts", kbd: [":w<CR>"] });
const saved = await Bun.file(filePath).text();
expect(saved).toBe("replacement\n");
});
@@ -264,8 +264,8 @@ describe("vim tool", () => {
await Bun.write(filePath, "first\n");
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "bad-insert.ts" });
await expect(tool.execute("bad", { kbd: [], insert: "nope" })).rejects.toThrow(
await tool.execute("open", { file: "bad-insert.ts" });
await expect(tool.execute("bad", { file: "bad-insert.ts", kbd: [], insert: "nope" })).rejects.toThrow(
/Insert payload requires INSERT mode/i,
);
});
@@ -275,7 +275,7 @@ describe("vim tool", () => {
await Bun.write(filePath, "\treturn value;\n");
const tool = new VimTool(createSession(tmpDir));
const opened = await tool.execute("open", { open: "tabs.ts" });
const opened = await tool.execute("open", { file: "tabs.ts" });
const text = textResult(opened);
expect(text).toContain("Focus:");
expect(text).toContain(" → return value;");
@@ -287,8 +287,8 @@ describe("vim tool", () => {
await Bun.write(filePath, "alpha\nbeta\n");
const tool = new VimTool(createSession(tmpDir));
await tool.execute("open", { open: "search.ts" });
const paused = await tool.execute("search", { kbd: ["/be"], pause: true });
await tool.execute("open", { file: "search.ts" });
const paused = await tool.execute("search", { file: "search.ts", kbd: ["/be"], pause: true });
expect(paused.details?.pendingInput?.kind).toBe("search-forward");
expect(textResult(paused)).toContain("Pending: /be");
});
@@ -299,8 +299,8 @@ describe("vim tool", () => {
const tool = new VimTool(createSession(tmpDir));
const pendingInputs: string[] = [];
await tool.execute("open", { open: "command.ts" });
const result = await tool.execute("command", { kbd: [":%s/foo/bar/g<CR>"] }, undefined, update => {
await tool.execute("open", { file: "command.ts" });
const result = await tool.execute("command", { file: "command.ts", kbd: [":%s/foo/bar/g<CR>"] }, undefined, update => {
const pending = update.details?.pendingInput;
if (pending?.kind === "command") {
pendingInputs.push(pending.text);
@@ -325,11 +325,11 @@ describe("vim tool", () => {
}),
);
await tool.execute("open", { open: "plan.ts" });
const moved = await tool.execute("move", { kbd: ["2G"] });
await tool.execute("open", { file: "plan.ts" });
const moved = await tool.execute("move", { file: "plan.ts", kbd: ["2G"] });
expect(textResult(moved)).toContain("L2:1");
await expect(tool.execute("edit", { kbd: ["dd"] })).rejects.toThrow(/Plan mode/i);
await expect(tool.execute("insert", { kbd: ["cc"], insert: "blocked" })).rejects.toThrow(/Plan mode/i);
await expect(tool.execute("edit", { file: "plan.ts", kbd: ["dd"] })).rejects.toThrow(/Plan mode/i);
await expect(tool.execute("insert", { file: "plan.ts", kbd: ["cc"], insert: "blocked" })).rejects.toThrow(/Plan mode/i);
});
});
@@ -1132,7 +1132,6 @@ async function runSingleTask(
}
}
// Retry if the model didn't attempt any edit/write (read-only or no tool calls)
const madeEditAttempt = toolStats.edit > 0 || toolStats.write > 0;
if (!madeEditAttempt && zeroToolRetries < noOpRetryLimit) {
@@ -1413,7 +1412,7 @@ async function _runRpcBenchmarkRun(
} else if (toolName === "write") {
toolStats.write++;
}
if (e.args) {
toolStats.totalInputChars += JSON.stringify(e.args).length;
}
@@ -1460,7 +1459,7 @@ async function _runRpcBenchmarkRun(
}
}
}
// Retry if the model didn't attempt any edit/write (read-only or no tool calls)
const madeEditAttempt = toolStats.edit > 0 || toolStats.write > 0;
if (!madeEditAttempt && zeroToolRetries < noOpRetryLimit) {
@@ -1470,15 +1469,15 @@ async function _runRpcBenchmarkRun(
attempt--; // Don't consume a regular attempt slot
continue;
}
patchApplied = toolStats.edit > 0;
const filesToVerify = task.files.length > 0 ? task.files : undefined;
const verification = await verifyExpectedFileSubset(expectedDir, cwd, filesToVerify);
if (config.autoFormat) {
await formatDirectory(cwd);
}
verificationPassed = verification.success;
indentScore = verification.indentScore;
formattedEquivalent = verification.formattedEquivalent;
@@ -1488,11 +1487,11 @@ async function _runRpcBenchmarkRun(
if (!verification.success && verification.error) {
error = verification.error;
}
if (verification.success) {
break;
}
const mutationIntentSuffix = mutationIntentValidation
? `\n\nMutation intent: ${mutationIntentValidation.matched ? "matched" : "not matched"} (${mutationIntentValidation.reason})`
: "";
+10 -1
View File
@@ -325,6 +325,7 @@ class RpcClient:
self._async_errors = _BoundedHistory[BaseException](_DEFAULT_ERROR_HISTORY_LIMIT)
self._scheduled_agent_runs = 0
self._completed_agent_runs = 0
self._last_schedule_async_error_index = 0
self._ui_requests: queue.Queue[ExtensionUiRequest] = queue.Queue()
self._stderr_chunks = _BoundedHistory[str](self._max_stderr_chunks)
self._closed_error: BaseException | None = None
@@ -381,6 +382,7 @@ class RpcClient:
self._async_errors.clear()
self._scheduled_agent_runs = 0
self._completed_agent_runs = 0
self._last_schedule_async_error_index = 0
self._ui_requests = queue.Queue()
with self._state_lock:
self._stderr_chunks.clear()
@@ -810,6 +812,7 @@ class RpcClient:
self._prompt_lifecycle.acquire(operation)
try:
if self._is_agent_idle():
self._check_async_errors()
return
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
@@ -841,7 +844,7 @@ class RpcClient:
def _mark_agent_run_scheduled(self) -> None:
with self._event_condition:
self._scheduled_agent_runs += 1
self._last_schedule_async_error_index = self._async_errors.current_index()
def _mark_agent_run_completed(self) -> None:
with self._event_condition:
self._completed_agent_runs += 1
@@ -851,6 +854,12 @@ class RpcClient:
with self._event_condition:
return self._scheduled_agent_runs == self._completed_agent_runs
def _check_async_errors(self) -> None:
with self._event_condition:
errors = self._async_errors.snapshot_from(self._last_schedule_async_error_index)
if errors:
raise errors[0]
def _build_prompt_turn(self, events: tuple[RpcAgentEvent, ...]) -> PromptTurn:
final_messages: tuple[AgentMessage, ...] = ()
for event in reversed(events):
+36
View File
@@ -0,0 +1,36 @@
#!/usr/bin/env python3
"""
Chunk edit benchmark: tests chunk-mode edit tool usage across models with a simple edit task.
"""
from __future__ import annotations
from edit_benchmark_common import BenchmarkSpec, EDIT_DIFF, EXPECTED_CONTENT, run_benchmark_main
EDIT_PROMPT = f"""\
Use the `read` tool to inspect `test.py`, then use the `edit` tool in chunk mode to make `test.py` exactly match the requested change.
Apply this diff:
```diff
{EDIT_DIFF}```
Final expected file content:
```python
{EXPECTED_CONTENT}```
"""
CHUNK_BENCHMARK = BenchmarkSpec(
description="Benchmark chunk-mode edit tool across models with simple edit tasks.",
workspace_prefix="chunk-benchmark",
tools=("edit", "read"),
env={"PI_EDIT_VARIANT": "chunk", "PI_STRICT_EDIT_MODE": "1"},
initial_prompt=EDIT_PROMPT,
retry_instruction='Use `read(path="test.py")` to refresh chunk selectors if needed, then try again using the edit tool.',
)
def main() -> int:
return run_benchmark_main(CHUNK_BENCHMARK)
if __name__ == "__main__":
raise SystemExit(main())
+526
View File
@@ -0,0 +1,526 @@
#!/usr/bin/env python3
"""
Shared helpers for edit benchmark scripts.
"""
from __future__ import annotations
import argparse
import asyncio
import json
import os
import shutil
import sys
import tempfile
import threading
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable
REPO_ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(REPO_ROOT / "python/omp-rpc/src"))
from omp_rpc import MessageEndEvent, MessageStartEvent, MessageUpdateEvent, RpcClient, ToolExecutionStartEvent # noqa: E402
MODELS = [
"openrouter/moonshotai/kimi-k2.5",
"openrouter/anthropic/claude-haiku-4.5",
"openrouter/google/gemini-3.1-flash-lite-preview",
"openrouter/z-ai/glm-4.7-20251222:nitro"
# "openrouter/anthropic/claude-sonnet-4.6",
# "openrouter/google/gemini-3-flash-preview",
# "openrouter/z-ai/glm-5-turbo",
# "openrouter/minimax/minimax-m2.7",
]
INITIAL_CONTENT = """\
def divide(a, b):
return a / b
def greet(name):
return f"Hello, {name}!"
def main():
print(divide(10, 2))
print(greet("World"))
"""
EXPECTED_CONTENT = """\
def divide(a, b):
if b == 0:
return None
return a / b
def multiply(a, b):
return a * b
def greet(name):
return f"Hello, {name}!"
def main():
print(divide(10, 2))
print(multiply(3, 4))
print(greet("World"))
"""
EDIT_DIFF = """\
@@ -1,9 +1,14 @@
def divide(a, b):
+ if b == 0:
+ return None
return a / b
+def multiply(a, b):
+ return a * b
+
def greet(name):
return f"Hello, {name}!"
def main():
print(divide(10, 2))
+ print(multiply(3, 4))
print(greet("World"))
"""
FEEDBACK_PROMPT = """\
STOP. The editing task is complete. Do NOT make any more edits or tool calls.
This is a survey. Answer these 6 questions about your experience using the editing tool (2-3 sentences each):
1. Tool input schema: Was the input schema intuitive? What confused you?
2. Tool description: Was the description clear enough? What was missing?
3. Tool behaviour: What would make the tool easier to use?
4. Tool results & errors: Were error messages helpful? What could improve?
5. Bugs: Did anything behave unexpectedly?
6. Other thoughts: Anything else?
"""
DEFAULT_MAX_TURNS = 20
_PRINT_LOCK = threading.Lock()
@dataclass(frozen=True)
class BenchmarkSpec:
description: str
workspace_prefix: str
tools: tuple[str, ...]
env: dict[str, str]
initial_prompt: str
retry_instruction: str
@dataclass
class BenchmarkResult:
model: str
success: bool
turns_used: int
prompt_attempts: int
edit_calls: int
token_input: int
token_output: int
feedback: str
error: str | None = None
class VerbosePrinter:
def __init__(self, model: str):
self._label = model.removeprefix("openrouter/")
self._open_kind: str | None = None
self._seen_block_lengths: dict[tuple[str, int], int] = {}
def _prefix(self, kind: str) -> str:
return f"[{self._label}] {kind}> "
def flush(self) -> None:
with _PRINT_LOCK:
if self._open_kind is None:
return
sys.stderr.write("\n")
sys.stderr.flush()
self._open_kind = None
def emit_delta(self, kind: str, delta: str, content_index: int | None = None) -> None:
if not delta:
return
if content_index is not None:
key = (kind, content_index)
self._seen_block_lengths[key] = self._seen_block_lengths.get(key, 0) + len(delta)
with _PRINT_LOCK:
if self._open_kind != kind:
if self._open_kind is not None:
sys.stderr.write("\n")
sys.stderr.write(self._prefix(kind))
self._open_kind = kind
parts = delta.splitlines(keepends=True)
for index, part in enumerate(parts):
if index > 0:
sys.stderr.write(self._prefix(kind))
sys.stderr.write(part)
if delta.endswith("\n"):
self._open_kind = None
sys.stderr.flush()
def emit_tool_call(self, tool_name: str, args: Any) -> None:
rendered_args = json.dumps(args, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
with _PRINT_LOCK:
if self._open_kind is not None:
sys.stderr.write("\n")
self._open_kind = None
sys.stderr.write(f"{self._prefix('tool')}{tool_name} {rendered_args}\n")
sys.stderr.flush()
def reset_message(self) -> None:
self.flush()
self._seen_block_lengths.clear()
def emit_missing_from_message(self, message: dict[str, Any]) -> None:
content = message.get("content")
if not isinstance(content, list):
return
for content_index, block in enumerate(content):
if not isinstance(block, dict):
continue
block_type = block.get("type")
if block_type == "text":
text = block.get("text")
kind = "text"
elif block_type == "thinking":
text = block.get("thinking")
kind = "thinking"
else:
continue
if not isinstance(text, str) or not text:
continue
key = (kind, content_index)
seen = self._seen_block_lengths.get(key, 0)
if seen < len(text):
self.emit_delta(kind, text[seen:], content_index)
def emit_redacted_thinking_notice(self, message: dict[str, Any]) -> None:
content = message.get("content")
if not isinstance(content, list):
return
has_redacted = any(isinstance(block, dict) and block.get("type") == "redactedThinking" for block in content)
if not has_redacted:
return
with _PRINT_LOCK:
if self._open_kind is not None:
sys.stderr.write("\n")
self._open_kind = None
sys.stderr.write(f"{self._prefix('thinking')}[redacted by provider]\n")
sys.stderr.flush()
def resolve_repo_omp_bin() -> str | None:
cli_path = REPO_ROOT / "packages/coding-agent" / "src/cli.ts"
if not cli_path.exists():
return None
return str(cli_path)
def resolve_omp_bin(raw: str | None) -> str:
if raw:
return raw
repo_bin = resolve_repo_omp_bin()
if repo_bin:
return repo_bin
found = shutil.which("omp")
if not found:
raise SystemExit("Could not find `omp` on PATH and could not resolve the repo CLI. Set --omp-bin or OMP_BIN.")
return found
def build_retry_prompt(spec: BenchmarkSpec, current_content: str) -> str:
return (
"The file doesn't match the expected result yet.\n\n"
f"Current content:\n```\n{current_content}```\n\n"
f"Expected:\n```\n{EXPECTED_CONTENT}```\n\n"
f"{spec.retry_instruction}"
)
def install_verbose_logging(
client: RpcClient,
model: str,
mode: str | None,
thinking: str | None,
) -> Callable[[], None] | None:
if mode is None:
return None
printer = VerbosePrinter(model)
include_messages = mode == "verbose"
if include_messages and thinking is None:
with _PRINT_LOCK:
sys.stderr.write(
f"[{model.removeprefix('openrouter/')}] verbose> "
"no thinking level requested; pass --thinking low|medium|high|xhigh if the provider exposes reasoning.\n"
)
sys.stderr.flush()
def handle_message_start(event: MessageStartEvent) -> None:
if not include_messages:
return
if event.message.get("role") == "assistant":
printer.reset_message()
def handle_message_update(event: MessageUpdateEvent) -> None:
if not include_messages:
return
if event.message.get("role") != "assistant":
return
message_event = event.assistant_message_event
event_type = message_event["type"]
if event_type == "text_delta":
printer.emit_delta("text", message_event["delta"], message_event["contentIndex"])
elif event_type == "thinking_delta":
printer.emit_delta("thinking", message_event["delta"], message_event["contentIndex"])
def handle_message_end(event: MessageEndEvent) -> None:
if not include_messages:
return
if event.message.get("role") == "assistant":
printer.emit_missing_from_message(event.message)
printer.emit_redacted_thinking_notice(event.message)
printer.flush()
printer.reset_message()
def handle_tool_start(event: ToolExecutionStartEvent) -> None:
printer.emit_tool_call(event.tool_name, event.args)
removers = [
client.on_message_start(handle_message_start),
client.on_message_update(handle_message_update),
client.on_message_end(handle_message_end),
client.on_tool_execution_start(handle_tool_start),
]
def cleanup() -> None:
for remove in reversed(removers):
remove()
printer.flush()
return cleanup
def run_benchmark_for_model(
*,
spec: BenchmarkSpec,
model: str,
omp_bin: str,
workspace: Path,
timeout: float,
log_mode: str | None,
thinking: str | None,
max_turns: int,
) -> BenchmarkResult:
"""Run a single edit benchmark for one model."""
test_file = workspace / "test.py"
test_file.write_text(INITIAL_CONTENT)
prompt_attempts = 0
token_input = 0
token_output = 0
turns_used = 0
edit_vim_tool_calls = 0
success = False
feedback = ""
error_msg: str | None = None
counting_edit_turns = True
try:
with RpcClient(
executable=omp_bin,
model=model,
cwd=workspace,
env={**spec.env},
thinking=thinking,
tools=spec.tools,
no_skills=True,
no_rules=True,
no_session=True,
startup_timeout=30.0,
request_timeout=120.0,
) as client:
client.install_headless_ui()
verbose_cleanup = install_verbose_logging(client, model, log_mode, thinking)
def handle_tool_count(event: ToolExecutionStartEvent) -> None:
nonlocal edit_vim_tool_calls, turns_used
if counting_edit_turns:
turns_used += 1
if event.tool_name in {"edit", "vim"}:
edit_vim_tool_calls += 1
tool_count_remover = client.on_tool_execution_start(handle_tool_count)
try:
for turn in range(1, max_turns + 1):
prompt_attempts = turn
if turn == 1:
client.prompt(spec.initial_prompt)
else:
client.prompt(build_retry_prompt(spec, test_file.read_text()))
client.wait_for_idle(timeout=timeout)
current_content = test_file.read_text()
if current_content.strip() == EXPECTED_CONTENT.strip():
success = True
break
stats = client.get_session_stats()
token_input = stats.tokens.input
token_output = stats.tokens.output
counting_edit_turns = False
client.prompt(FEEDBACK_PROMPT)
client.wait_for_idle(timeout=timeout)
feedback = client.get_last_assistant_text() or ""
stats = client.get_session_stats()
token_input = stats.tokens.input
token_output = stats.tokens.output
finally:
tool_count_remover()
if verbose_cleanup is not None:
verbose_cleanup()
except Exception as exc:
error_msg = f"{type(exc).__name__}: {exc}"
return BenchmarkResult(
model=model,
success=success,
turns_used=turns_used,
prompt_attempts=prompt_attempts,
edit_calls=edit_vim_tool_calls,
token_input=token_input,
token_output=token_output,
feedback=feedback.strip(),
error=error_msg,
)
async def run_all(spec: BenchmarkSpec, args: argparse.Namespace) -> dict[str, dict[str, Any]]:
omp_bin = resolve_omp_bin(args.omp_bin)
timestamp = time.strftime("%Y%m%d-%H%M%S")
workspace_root = Path(tempfile.gettempdir()) / f"{spec.workspace_prefix}-{timestamp}"
workspace_root.mkdir(parents=True, exist_ok=True)
selected_models = args.models or MODELS
tasks = []
for model in selected_models:
model_slug = model.replace("/", "_")
workspace = workspace_root / model_slug
workspace.mkdir(parents=True, exist_ok=True)
print(f"Starting benchmark for {model}...", file=sys.stderr)
tasks.append(
asyncio.to_thread(
run_benchmark_for_model,
spec=spec,
model=model,
omp_bin=omp_bin,
workspace=workspace,
timeout=args.timeout,
log_mode="verbose" if args.verbose else ("print" if args.print else None),
thinking=args.thinking,
max_turns=args.max_turns,
)
)
benchmark_results = await asyncio.gather(*tasks, return_exceptions=True)
results: dict[str, dict[str, Any]] = {}
for model, result in zip(selected_models, benchmark_results):
if isinstance(result, Exception):
results[model] = {
"tokens_in": 0,
"tokens_out": 0,
"model_feedback": "",
"success": False,
"turns_used": 0,
"prompt_attempts": 0,
"edit_calls": 0,
"error": f"{type(result).__name__}: {result}",
}
print(f" {model}: error - {result}", file=sys.stderr)
continue
results[model] = {
"tokens_in": result.token_input,
"tokens_out": result.token_output,
"model_feedback": result.feedback,
"success": result.success,
"turns_used": result.turns_used,
"edit_calls": result.edit_calls,
"prompt_attempts": result.prompt_attempts,
"error": result.error,
}
status = "success" if result.success else "failed"
print(f" {model}: {status} in {result.turns_used} turns", file=sys.stderr)
return results
def parse_args(description: str) -> argparse.Namespace:
parser = argparse.ArgumentParser(description=description)
parser.add_argument(
"--omp-bin",
default=os.environ.get("OMP_BIN"),
help="Executable to launch. Defaults to the repo checkout CLI, then falls back to `omp` on PATH.",
)
parser.add_argument(
"--timeout", type=float, default=300.0, help="Per-turn timeout in seconds."
)
parser.add_argument(
"--max-turns",
type=int,
default=DEFAULT_MAX_TURNS,
help=f"Maximum edit/retry turns before the benchmark gives up (default: {DEFAULT_MAX_TURNS}).",
)
parser.add_argument(
"--model",
dest="models",
action="append",
help="Repeat to limit execution to specific models.",
)
logging_group = parser.add_mutually_exclusive_group()
logging_group.add_argument(
"--print",
action="store_true",
help="Print tool calls to stderr while the benchmark runs.",
)
logging_group.add_argument(
"--verbose",
action="store_true",
help="Print assistant text, thinking, and tool calls to stderr while the benchmark runs.",
)
parser.add_argument(
"--thinking",
choices=["off", "minimal", "low", "medium", "high", "xhigh"],
default="medium",
help="Request a specific thinking level for models that support reasoning (default: medium).",
)
return parser.parse_args()
def run_benchmark_main(spec: BenchmarkSpec) -> int:
args = parse_args(spec.description)
results = asyncio.run(run_all(spec, args))
print(json.dumps(results, indent=2))
return 0
+11 -281
View File
@@ -1,85 +1,10 @@
#!/usr/bin/env python3
"""
Vim edit benchmark: Tests vim tool across 3 models with simple edit tasks.
Retries up to 10 turns until file matches expected, then asks for feedback.
Outputs JSON results with tokens, feedback, and success status.
Vim edit benchmark: tests the vim tool across models with a simple edit task.
"""
from __future__ import annotations
import argparse
import asyncio
import json
import os
import shutil
import sys
import tempfile
import time
from dataclasses import dataclass
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(REPO_ROOT / "python/omp-rpc/src"))
from omp_rpc import RpcClient, RpcError # noqa: E402
MODELS = [
"openrouter/moonshotai/kimi-k2.5",
"openrouter/anthropic/claude-haiku-4.5",
"openrouter/anthropic/claude-sonnet-4.6",
"openrouter/google/gemini-3-flash-preview",
"openrouter/z-ai/glm-5-turbo",
"openrouter/minimax/minimax-m2.7",
]
# Edit task: add error handling and a new method
INITIAL_CONTENT = """\
def divide(a, b):
return a / b
def greet(name):
return f"Hello, {name}!"
def main():
print(divide(10, 2))
print(greet("World"))
"""
EXPECTED_CONTENT = """\
def divide(a, b):
if b == 0:
return None
return a / b
def multiply(a, b):
return a * b
def greet(name):
return f"Hello, {name}!"
def main():
print(divide(10, 2))
print(multiply(3, 4))
print(greet("World"))
"""
EDIT_DIFF = """\
@@ -1,9 +1,14 @@
def divide(a, b):
+ if b == 0:
+ return None
return a / b
+def multiply(a, b):
+ return a * b
+
def greet(name):
return f"Hello, {name}!"
def main():
print(divide(10, 2))
+ print(multiply(3, 4))
print(greet("World"))
"""
from edit_benchmark_common import BenchmarkSpec, EDIT_DIFF, run_benchmark_main
EDIT_PROMPT = f"""\
Apply the following diff to the file `test.py` using the vim tool with the minimum amount of "moves":
@@ -87,213 +12,18 @@ Apply the following diff to the file `test.py` using the vim tool with the minim
{EDIT_DIFF}```
"""
FEEDBACK_PROMPT = """\
You just used the edit tool in vim mode to make edits. Please share your honest feedback on each point below (2-3 sentences each):
1. **Tool input schema**: Was the input schema intuitive? What could be better?
2. **Tool description**: Was the tool description helpful enough to use it correctly? How could it be improved?
3. **Tool behaviour**: Any improvements or changes to how the tool works that would lead to smoother outcomes?
4. **Tool results & errors**: What could be improved about the tool results or error messages?
5. **Bugs**: Did you encounter any bugs or unexpected behaviour?
6. **Other thoughts**: Anything else worth mentioning?
"""
MAX_TURNS = 10
@dataclass
class BenchmarkResult:
model: str
success: bool
turns_used: int
token_input: int
token_output: int
feedback: str
error: str | None = None
def require_openrouter_key() -> str:
key = os.environ.get("OPENROUTER_API_KEY")
if not key:
raise SystemExit("OPENROUTER_API_KEY is not set")
return key
def resolve_omp_bin(raw: str | None) -> str:
if raw:
return raw
found = shutil.which("omp")
if not found:
raise SystemExit("Could not find `omp` on PATH. Set --omp-bin or OMP_BIN.")
return found
def run_benchmark_for_model(
*,
model: str,
omp_bin: str,
workspace: Path,
timeout: float,
openrouter_key: str,
) -> BenchmarkResult:
"""Run the vim edit benchmark for a single model."""
test_file = workspace / "test.py"
test_file.write_text(INITIAL_CONTENT)
token_input = 0
token_output = 0
turns_used = 0
success = False
feedback = ""
error_msg: str | None = None
try:
with RpcClient(
executable=omp_bin,
model=model,
cwd=workspace,
env={"OPENROUTER_API_KEY": openrouter_key, "PI_EDIT_VARIANT": "vim", "PI_STRICT_EDIT_MODE": "1"},
tools=("edit", "read"),
no_skills=True,
no_rules=True,
no_session=True,
startup_timeout=30.0,
request_timeout=120.0,
) as client:
client.install_headless_ui()
# Edit loop: keep prompting until file matches or max turns
for turn in range(1, MAX_TURNS + 1):
turns_used = turn
if turn == 1:
client.prompt(EDIT_PROMPT)
else:
current = test_file.read_text()
client.prompt(
f"The file doesn't match the expected result yet.\n\n"
f"Current content:\n```\n{current}```\n\n"
f"Expected:\n```\n{EXPECTED_CONTENT}```\n\n"
f"Please try again using the edit tool."
)
client.wait_for_idle(timeout=timeout)
# Check if file matches expected
current_content = test_file.read_text()
if current_content.strip() == EXPECTED_CONTENT.strip():
success = True
break
# Get token usage from session stats
stats = client.get_session_stats()
token_input = stats.tokens.input
token_output = stats.tokens.output
# Ask for feedback
client.prompt(FEEDBACK_PROMPT)
client.wait_for_idle(timeout=timeout)
feedback = client.get_last_assistant_text() or ""
# Update final token counts
stats = client.get_session_stats()
token_input = stats.tokens.input
token_output = stats.tokens.output
except Exception as e:
error_msg = f"{type(e).__name__}: {e}"
return BenchmarkResult(
model=model,
success=success,
turns_used=turns_used,
token_input=token_input,
token_output=token_output,
feedback=feedback.strip(),
error=error_msg,
)
async def run_all(args: argparse.Namespace) -> dict:
openrouter_key = require_openrouter_key()
omp_bin = resolve_omp_bin(args.omp_bin)
timestamp = time.strftime("%Y%m%d-%H%M%S")
workspace_root = Path(tempfile.gettempdir()) / f"vim-benchmark-{timestamp}"
workspace_root.mkdir(parents=True, exist_ok=True)
selected_models = args.models or MODELS
# Create workspaces and tasks
tasks = []
for model in selected_models:
model_slug = model.replace("/", "_")
workspace = workspace_root / model_slug
workspace.mkdir(parents=True, exist_ok=True)
print(f"Starting benchmark for {model}...", file=sys.stderr)
tasks.append(
asyncio.to_thread(
run_benchmark_for_model,
model=model,
omp_bin=omp_bin,
workspace=workspace,
timeout=args.timeout,
openrouter_key=openrouter_key,
)
)
# Run all in parallel
benchmark_results = await asyncio.gather(*tasks, return_exceptions=True)
results: dict[str, dict] = {}
for model, result in zip(selected_models, benchmark_results):
if isinstance(result, Exception):
results[model] = {
"tokens_in": 0,
"tokens_out": 0,
"model_feedback": "",
"success": False,
"turns_used": 0,
"error": f"{type(result).__name__}: {result}",
}
print(f" {model}: error - {result}", file=sys.stderr)
else:
results[model] = {
"tokens_in": result.token_input,
"tokens_out": result.token_output,
"model_feedback": result.feedback,
"success": result.success,
"turns_used": result.turns_used,
"error": result.error,
}
status = "success" if result.success else "failed"
print(f" {model}: {status} in {result.turns_used} turns", file=sys.stderr)
return results
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Benchmark vim tool across models with simple edit tasks."
)
parser.add_argument("--omp-bin", default=os.environ.get("OMP_BIN"))
parser.add_argument(
"--timeout", type=float, default=300.0, help="Per-turn timeout in seconds."
)
parser.add_argument(
"--model",
dest="models",
action="append",
help="Repeat to limit execution to specific models.",
)
return parser.parse_args()
VIM_BENCHMARK = BenchmarkSpec(
description="Benchmark vim tool across models with simple edit tasks.",
workspace_prefix="vim-benchmark",
tools=("edit", "read"),
env={"PI_EDIT_VARIANT": "vim", "PI_STRICT_EDIT_MODE": "1"},
initial_prompt=EDIT_PROMPT,
retry_instruction="Please try again using the vim tool.",
)
def main() -> int:
args = parse_args()
results = asyncio.run(run_all(args))
print(json.dumps(results, indent=2))
return 0
return run_benchmark_main(VIM_BENCHMARK)
if __name__ == "__main__":