diff --git a/src/api/providers/__tests__/vscode-lm.spec.ts b/src/api/providers/__tests__/vscode-lm.spec.ts index 86a9a30c38..780494cf54 100644 --- a/src/api/providers/__tests__/vscode-lm.spec.ts +++ b/src/api/providers/__tests__/vscode-lm.spec.ts @@ -293,6 +293,398 @@ describe("VsCodeLmHandler", () => { }) }) + describe("leaked tool-call recovery during streaming", () => { + const salvageTools = [ + { + type: "function" as const, + function: { + name: "calculator", + description: "A simple calculator", + parameters: { type: "object", properties: { operation: { type: "string" } } }, + }, + }, + ] + + const streamTextParts = (parts: string[]) => { + mockLanguageModelChat.sendRequest.mockResolvedValueOnce({ + stream: (async function* () { + for (const part of parts) { + yield new vscode.LanguageModelTextPart(part) + } + return + })(), + text: (async function* () { + yield parts.join("") + return + })(), + }) + } + + const streamMixedParts = (parts: Array) => { + mockLanguageModelChat.sendRequest.mockResolvedValueOnce({ + stream: (async function* () { + for (const part of parts) { + yield typeof part === "string" + ? new vscode.LanguageModelTextPart(part) + : new vscode.LanguageModelToolCallPart("native-1", part.name, part.input) + } + return + })(), + text: (async function* () { + yield "" + return + })(), + }) + } + + const drain = async () => { + const stream = handler.createMessage("system", [{ role: "user" as const, content: "hi" }], { + taskId: "test-task", + tools: salvageTools, + }) + const chunks = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + return chunks + } + + const collect = async (parts: string[]) => { + streamTextParts(parts) + return drain() + } + + it("recovers a tool call the model streamed as raw invoke XML", async () => { + const chunks = await collect([ + "Thinking. ", + 'add', + ]) + + expect(chunks.filter((chunk) => chunk.type === "text")).toEqual([{ type: "text", text: "Thinking. " }]) + expect(chunks.filter((chunk) => chunk.type === "tool_call")).toEqual([ + { + type: "tool_call", + id: expect.stringContaining("vscodelm-salvaged-"), + name: "calculator", + arguments: JSON.stringify({ operation: "add" }), + }, + ]) + }) + + it("detects a marker split across stream chunks", async () => { + const chunks = await collect([ + "abc sub', + ]) + + expect(chunks.filter((chunk) => chunk.type === "text")).toEqual([{ type: "text", text: "abc " }]) + expect(chunks.filter((chunk) => chunk.type === "tool_call")).toMatchObject([ + { name: "calculator", arguments: JSON.stringify({ operation: "sub" }) }, + ]) + }) + + it("emits a carried tail as plain text when it never becomes a marker", async () => { + const chunks = await collect(["hello chunk.type === "text")).toEqual([ + { type: "text", text: "hello " }, + { type: "text", text: " chunk.type === "tool_call")).toBe(false) + }) + + it("buffers across chunks that arrive after the marker", async () => { + const chunks = await collect([ + 'prose ', + '', + "mul", + "", + ]) + + expect(chunks.filter((chunk) => chunk.type === "text")).toEqual([{ type: "text", text: "prose " }]) + expect(chunks.filter((chunk) => chunk.type === "tool_call")).toMatchObject([ + { name: "calculator", arguments: JSON.stringify({ operation: "mul" }) }, + ]) + }) + + it("recovers a null-only declared parameter as JSON null through createMessage", async () => { + mockLanguageModelChat.sendRequest.mockResolvedValueOnce({ + stream: (async function* () { + yield new vscode.LanguageModelTextPart( + 'null', + ) + return + })(), + text: (async function* () { + yield "" + return + })(), + }) + + const stream = handler.createMessage("system", [{ role: "user" as const, content: "hi" }], { + taskId: "test-task", + tools: [ + { + type: "function" as const, + function: { + name: "nuller", + description: "", + parameters: { type: "object", properties: { cursor: { type: "null" } } }, + }, + }, + ], + }) + const chunks = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + expect(chunks.filter((chunk) => chunk.type === "tool_call")).toMatchObject([ + { name: "nuller", arguments: JSON.stringify({ cursor: null }) }, + ]) + }) + + it("keeps an invoke block for an unknown tool as literal text", async () => { + const block = '1' + const chunks = await collect([block]) + + expect(chunks.filter((chunk) => chunk.type === "text")).toEqual([{ type: "text", text: block }]) + expect(chunks.some((chunk) => chunk.type === "tool_call")).toBe(false) + }) + + it("emits prose before the recovered tool call", async () => { + const chunks = await collect([ + "Thinking. ", + 'add', + ]) + + expect(chunks.map((chunk) => chunk.type)).toEqual(["text", "tool_call", "usage"]) + }) + + it("flushes buffered text before a native tool call so no text follows a tool_use", async () => { + streamMixedParts([ + 'partial ', + { name: "calculator", input: { operation: "div" } }, + ]) + const chunks = await drain() + + // The ordering comparison is only meaningful once both kinds of chunk exist: a + // silently broken flush emits no text at all, and -1 < firstToolCall would still hold. + expect(chunks.filter((chunk) => chunk.type === "text")).toEqual([ + { type: "text", text: "partial " }, + { type: "text", text: '' }, + ]) + const types = chunks.map((chunk) => chunk.type) + const lastText = types.lastIndexOf("text") + const firstToolCall = types.indexOf("tool_call") + expect(lastText).toBeGreaterThanOrEqual(0) + expect(firstToolCall).toBeGreaterThanOrEqual(0) + expect(firstToolCall).toBeGreaterThan(lastText) + }) + + it("does not latch buffering on prose that merely mentions the tag", async () => { + const chunks = await collect(["never emit markup as text. ", "Streaming continues."]) + + expect(chunks.filter((chunk) => chunk.type === "text")).toEqual([ + { type: "text", text: "never emit markup as text. " }, + { type: "text", text: "Streaming continues." }, + ]) + expect(chunks.some((chunk) => chunk.type === "tool_call")).toBe(false) + }) + + it("does not recover an invoke block quoted inside a fenced code block", async () => { + const block = 'add' + // Open the wrapper outside the fence so only the quoting guard prevents recovery. + const chunks = await collect(["\n```\n" + block + "\n```\n"]) + + expect(chunks.some((chunk) => chunk.type === "tool_call")).toBe(false) + }) + + it.each([ + { name: "write_to_file", parameter: "content" }, + { name: "update_todo_list", parameter: "todos" }, + ])( + "recovers a large $name call split across chunks without leaking markup", + async ({ name, parameter }) => { + const payload = "x".repeat(5000) + streamTextParts([ + ``, + payload, + "", + ]) + const chunks = [] + for await (const chunk of handler.createMessage("system", [{ role: "user", content: "hi" }], { + taskId: "test-task", + tools: [ + { + type: "function", + function: { + name, + parameters: { type: "object", properties: { [parameter]: { type: "string" } } }, + }, + }, + ], + })) { + chunks.push(chunk) + } + + expect(chunks.filter((chunk) => chunk.type !== "usage")).toEqual([ + { + type: "tool_call", + id: expect.stringContaining("vscodelm-salvaged-"), + name, + arguments: JSON.stringify({ [parameter]: payload }), + }, + ]) + }, + ) + + it("flushes an over-long never-closing invoke as plain text before the stream ends", async () => { + // End-of-stream flushing produces identical text, so track production timing to + // prove the buffer releases text before the stream ends. + const filler = "x".repeat(256 * 1024) + const parts = ['', filler, filler, filler, filler] + let partsProduced = 0 + + mockLanguageModelChat.sendRequest.mockResolvedValueOnce({ + stream: (async function* () { + for (const part of parts) { + partsProduced++ + yield new vscode.LanguageModelTextPart(part) + } + return + })(), + text: (async function* () { + yield "" + return + })(), + }) + + const stream = handler.createMessage("system", [{ role: "user" as const, content: "hi" }], { + taskId: "test-task", + tools: salvageTools, + }) + + let sawTextBeforeStreamEnd = false + let streamedText = "" + for await (const chunk of stream) { + if (chunk.type === "text") { + streamedText += chunk.text + if (partsProduced < parts.length) { + sawTextBeforeStreamEnd = true + } + } + } + + expect(sawTextBeforeStreamEnd).toBe(true) + expect(streamedText).toContain('') + expect(streamedText).toContain(filler) + }) + + it("recovers a completed call and releases text when the buffer passes the cap", async () => { + // The cap used to be bypassed whenever a complete block sat in the buffer, so both the + // call and 512 KB of trailing prose were withheld until the stream ended. + const filler = "y".repeat(256 * 1024) + const parts = [ + 'add', + filler, + filler, + ] + let partsProduced = 0 + + mockLanguageModelChat.sendRequest.mockResolvedValueOnce({ + stream: (async function* () { + for (const part of parts) { + partsProduced++ + yield new vscode.LanguageModelTextPart(part) + } + return + })(), + text: (async function* () { + yield "" + return + })(), + }) + + const stream = handler.createMessage("system", [{ role: "user" as const, content: "hi" }], { + taskId: "test-task", + tools: salvageTools, + }) + + let sawTextBeforeStreamEnd = false + const toolCalls = [] + for await (const chunk of stream) { + if (chunk.type === "text" && partsProduced < parts.length) { + sawTextBeforeStreamEnd = true + } + if (chunk.type === "tool_call") { + toolCalls.push(chunk) + } + } + + expect(sawTextBeforeStreamEnd).toBe(true) + expect(toolCalls).toMatchObject([ + { name: "calculator", arguments: JSON.stringify({ operation: "add" }) }, + ]) + }) + + it("keeps an invoke still in flight buffered when the cap is passed", async () => { + // Only the decided prefix may drain; emitting the open block as text would strand the + // call whose closing tag arrives in a later chunk. Every complete block in the + // over-cap buffer must drain now, so timing is asserted rather than final order. + const bulky = "z".repeat(256 * 1024) + const parts = [ + `\n${bulky}\n` + + 'mid\n', + 'sub', + "", + ] + let partsProduced = 0 + + mockLanguageModelChat.sendRequest.mockResolvedValueOnce({ + stream: (async function* () { + for (const part of parts) { + partsProduced++ + yield new vscode.LanguageModelTextPart(part) + } + return + })(), + text: (async function* () { + yield "" + return + })(), + }) + + const stream = handler.createMessage("system", [{ role: "user" as const, content: "hi" }], { + taskId: "test-task", + tools: salvageTools, + }) + + const toolCalls = [] + const drainedEarly = [] + for await (const chunk of stream) { + if (chunk.type === "tool_call") { + toolCalls.push(chunk) + if (partsProduced < parts.length) { + drainedEarly.push(chunk) + } + } + } + + // Both complete blocks sit in the buffer when the cap is crossed, so both must drain + // before the stream ends; the still-open third only resolves at the final chunk. + expect(drainedEarly).toMatchObject([ + { arguments: JSON.stringify({ operation: bulky }) }, + { arguments: JSON.stringify({ operation: "mid" }) }, + ]) + expect(toolCalls).toMatchObject([ + { name: "calculator", arguments: JSON.stringify({ operation: bulky }) }, + { name: "calculator", arguments: JSON.stringify({ operation: "mid" }) }, + { name: "calculator", arguments: JSON.stringify({ operation: "sub" }) }, + ]) + }) + }) + it("returns the original registry name for a tool declared with an encoded name", async () => { const systemPrompt = "You are a helpful assistant" const originalName = `read\uD800file` diff --git a/src/api/providers/vscode-lm.ts b/src/api/providers/vscode-lm.ts index 2771227680..c839312d6a 100644 --- a/src/api/providers/vscode-lm.ts +++ b/src/api/providers/vscode-lm.ts @@ -96,9 +96,38 @@ function convertToVsCodeLmTools(tools: OpenAI.Chat.ChatCompletionTool[]): vscode * contain bare markup, so passing it through is the safer default rather than a security boundary. * Widening to the bare case needs a reproduction first. */ +// Latching on a bare `` block in `text`, or 0 when none is closed. + * Mirrors `extractLeakedToolCalls`' forward scan so both agree on what is decidable. + */ +function lastCompleteInvokeBlockEnd(text: string): number { + const openPattern = /<(?:antml:)?invoke\s+name="([^"]+)"\s*>/gi + const closePattern = /<\/(?:antml:)?invoke\s*>/gi + let end = 0 + for (let open = openPattern.exec(text); open !== null; open = openPattern.exec(text)) { + closePattern.lastIndex = open.index + open[0].length + const close = closePattern.exec(text) + // With no closing tag after this open there is none after any later open either. + if (!close) { + break + } + end = close.index + close[0].length + openPattern.lastIndex = end + } + return end +} + /** * Left-to-right scan state behind the quoting heuristics: open code fence, current line start, * backticks seen on the current line, and whether a `` wrapper is open. @@ -1103,6 +1132,71 @@ export class VsCodeLmHandler extends BaseProvider implements SingleCompletionHan // Accumulate the text and count at the end of the stream to reduce token counting overhead. let accumulatedText: string = "" + // Only offered tools may be recovered; their schemas keep recovered parameters typed + // rather than passing every value through as a string. + const providedToolSchemas: LeakedToolSchemas = new Map( + (metadata?.tools ?? []) + .filter((tool) => tool.type === "function") + .filter((tool) => tool.function.name.length > 0) + .map((tool) => [tool.function.name, tool.function.parameters as Record | undefined]), + ) + const salvageLeakedToolCalls = providedToolSchemas.size > 0 + let salvageBuffering = false + let salvageBuffer = "" + let salvageCarry = "" + let salvageEmittedText = "" + let salvagedToolCallIndex = 0 + + // Parses one fully-decided span and returns its ordered chunks: prose first, then recovered + // calls. Advancing `salvageEmittedText` by exactly the consumed span is what preserves the + // quoting and wrapper context for whatever tail stays buffered. + const drainSalvagePrefix = (span: string): ApiStreamChunk[] => { + const drained: ApiStreamChunk[] = [] + const { calls, leftoverText } = extractLeakedToolCalls(span, providedToolSchemas, salvageEmittedText) + salvageEmittedText += span + if (leftoverText) { + drained.push({ type: "text", text: leftoverText }) + } + for (const call of calls) { + console.warn( + "Zoo Code : Recovered a tool call the model emitted as text instead of a structured tool call:", + { name: call.name, params: Object.keys(call.input) }, + ) + drained.push({ + type: "tool_call", + id: `vscodelm-salvaged-${Date.now()}-${salvagedToolCallIndex++}`, + name: call.name, + arguments: JSON.stringify(call.input), + }) + } + return drained + } + + // Drains the salvage state into ordered chunks: prose first, then any recovered calls. Must + // run before a native tool_call is yielded — text after a tool_use block is rejected by + // Anthropic once the turn is serialized back into history. + const flushSalvage = (): ApiStreamChunk[] => { + if (!salvageLeakedToolCalls) { + return [] + } + + if (!salvageBuffering) { + if (salvageCarry) { + const carried = salvageCarry + salvageCarry = "" + return [{ type: "text", text: carried }] + } + return [] + } + + // Buffering is only entered with the marker already in the buffer, so `buffered` is + // always non-empty here; an emptiness guard would be unreachable code. + const buffered = salvageBuffer + salvageBuffering = false + salvageBuffer = "" + return drainSalvagePrefix(buffered) + } + try { // Create the response stream with required options const requestOptions: vscode.LanguageModelChatRequestOptions = { @@ -1126,11 +1220,64 @@ export class VsCodeLmHandler extends BaseProvider implements SingleCompletionHan } accumulatedText += chunk.value - yield { - type: "text", - text: chunk.value, + + // Fast path: when we didn't offer any tools there is nothing to salvage, so + // stream the text straight through exactly as before. + if (!salvageLeakedToolCalls) { + yield { type: "text", text: chunk.value } + continue + } + + // Once we've seen the start of a leaked tool-call block, buffer the rest of the + // stream so the full markup can be parsed and replayed as a structured call. + if (salvageBuffering) { + salvageBuffer += chunk.value + if (salvageBuffer.length > MAX_SALVAGE_BUFFER_CHARS) { + // Past the cap, drain only the decided prefix: emitting the undecided tail as + // text would strand a call whose closing tag is still in flight. + const decidedEnd = lastCompleteInvokeBlockEnd(salvageBuffer) + if (decidedEnd > 0) { + const decided = salvageBuffer.slice(0, decidedEnd) + salvageBuffer = salvageBuffer.slice(decidedEnd) + yield* drainSalvagePrefix(decided) + } + // Still over the cap with nothing decidable: markup that never closes must not + // withhold the stream, so release it and let a later marker re-arm recovery. + if (salvageBuffer.length > MAX_SALVAGE_BUFFER_CHARS) { + const overflowed = salvageBuffer + salvageBuffering = false + salvageBuffer = "" + salvageEmittedText += overflowed + yield { type: "text", text: overflowed } + } + } + continue + } + + // Watch for the start of a leaked tool-call block, carrying a small tail across + // chunks so a marker split across chunk boundaries is still detected. + const combined = salvageCarry + chunk.value + const markerMatch = combined.match(LEAKED_TOOL_CALL_START) + if (markerMatch) { + const before = combined.slice(0, markerMatch.index) + if (before) { + salvageEmittedText += before + yield { type: "text", text: before } + } + salvageBuffering = true + salvageBuffer = combined.slice(markerMatch.index) + salvageCarry = "" + } else { + const carryLength = trailingPartialToolMarkerLength(combined) + const emit = carryLength > 0 ? combined.slice(0, combined.length - carryLength) : combined + salvageCarry = carryLength > 0 ? combined.slice(combined.length - carryLength) : "" + if (emit) { + salvageEmittedText += emit + yield { type: "text", text: emit } + } } } else if (chunk instanceof vscode.LanguageModelToolCallPart) { + yield* flushSalvage() try { // Validate tool call parameters if (!chunk.name || typeof chunk.name !== "string") { @@ -1178,6 +1325,9 @@ export class VsCodeLmHandler extends BaseProvider implements SingleCompletionHan } } + // Flush any leaked tool-call recovery state accumulated during streaming. + yield* flushSalvage() + // Count tokens in the accumulated text after stream completion const totalOutputTokens: number = await this.internalCountTokens(accumulatedText)