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:
@@ -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,
|
||||
};
|
||||
}
|
||||
|
||||
@@ -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): {
|
||||
|
||||
@@ -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";
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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"]);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 }))]
|
||||
: []),
|
||||
]);
|
||||
}
|
||||
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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, {
|
||||
|
||||
@@ -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}`);
|
||||
}
|
||||
|
||||
|
||||
@@ -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.',
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -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: () => [],
|
||||
|
||||
@@ -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[] = [];
|
||||
|
||||
@@ -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})`
|
||||
: "";
|
||||
|
||||
@@ -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):
|
||||
|
||||
Executable
+36
@@ -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())
|
||||
@@ -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
@@ -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__":
|
||||
|
||||
Reference in New Issue
Block a user