From 832d2978b605457238ae7e6a840a027610fa0291 Mon Sep 17 00:00:00 2001 From: Roo Code Date: Thu, 24 Jul 2025 10:29:07 +0000 Subject: [PATCH] feat: improve apply_code tool with isolated API context for diff generation - Create truly isolated API handler for Stage 2 diff generation - Use hardcoded, optimized system prompt for diff generation - Ensure clean context without conversational history in Stage 2 - Update tests to reflect the new implementation This addresses the subtle but important difference pointed out by @CberYellowstone in issue #6159, where the original proposal emphasized creating a truly isolated, dedicated process for diff generation. --- .../tools/__tests__/applyCodeTool.spec.ts | 412 +++++++++++++----- src/core/tools/applyCodeTool.ts | 59 ++- 2 files changed, 348 insertions(+), 123 deletions(-) diff --git a/src/core/tools/__tests__/applyCodeTool.spec.ts b/src/core/tools/__tests__/applyCodeTool.spec.ts index a8fabd34ed..deedc8c362 100644 --- a/src/core/tools/__tests__/applyCodeTool.spec.ts +++ b/src/core/tools/__tests__/applyCodeTool.spec.ts @@ -9,9 +9,9 @@ import * as path from "path" // Mock fs/promises before any imports vi.mock("fs/promises", () => ({ default: { - readFile: vi.fn().mockResolvedValue(Buffer.from("original content")), + readFile: vi.fn(), }, - readFile: vi.fn().mockResolvedValue(Buffer.from("original content")), + readFile: vi.fn(), })) // Mock dependencies @@ -42,13 +42,19 @@ vi.mock("../../prompts/responses", () => ({ }, })) -vi.mock("../applyDiffTool", () => ({ - applyDiffToolLegacy: vi.fn().mockResolvedValue(undefined), +vi.mock("../../../api", () => ({ + buildApiHandler: vi.fn().mockReturnValue({ + createMessage: vi.fn(), + getModel: vi.fn().mockReturnValue({ + id: "claude-3", + info: { contextWindow: 200000 }, + }), + }), })) // Import after mocking to get the mocked version -import { applyDiffToolLegacy } from "../applyDiffTool" import fs from "fs/promises" +import { buildApiHandler } from "../../../api" describe("applyCodeTool", () => { // Test data @@ -63,8 +69,8 @@ describe("applyCodeTool", () => { const mockedFileExistsAtPath = fileExistsAtPath as MockedFunction const mockedGetReadablePath = getReadablePath as MockedFunction const mockedPathResolve = path.resolve as MockedFunction - const mockedApplyDiffToolLegacy = applyDiffToolLegacy as MockedFunction const mockedReadFile = fs.readFile as MockedFunction + const mockedBuildApiHandler = buildApiHandler as MockedFunction const mockCline: any = {} let mockAskApproval: ReturnType @@ -82,6 +88,8 @@ describe("applyCodeTool", () => { mockCline.cwd = "/" mockCline.consecutiveMistakeCount = 0 + mockCline.taskId = "test-task-id" + mockCline.apiConfiguration = { apiProvider: "anthropic", apiKey: "test-key" } mockCline.api = { createMessage: vi.fn(), getModel: vi.fn().mockReturnValue({ @@ -93,18 +101,40 @@ describe("applyCodeTool", () => { deref: vi.fn().mockReturnValue({ getState: vi.fn().mockResolvedValue({ applyEnabled: true, + diagnosticsEnabled: true, + writeDelayMs: 0, }), }), } mockCline.rooIgnoreController = { validateAccess: vi.fn().mockReturnValue(true), } + mockCline.rooProtectedController = { + isWriteProtected: vi.fn().mockReturnValue(false), + } + mockCline.diffStrategy = { + applyDiff: vi.fn().mockResolvedValue({ + success: true, + content: "modified content", + }), + } mockCline.diffViewProvider = { reset: vi.fn().mockResolvedValue(undefined), + editType: undefined, + open: vi.fn().mockResolvedValue(undefined), + update: vi.fn().mockResolvedValue(undefined), + scrollToFirstDiff: vi.fn(), + revertChanges: vi.fn().mockResolvedValue(undefined), + saveChanges: vi.fn().mockResolvedValue(undefined), + pushToolWriteResult: vi.fn().mockResolvedValue("File updated successfully"), + } + mockCline.fileContextTracker = { + trackFileContext: vi.fn().mockResolvedValue(undefined), } mockCline.say = vi.fn().mockResolvedValue(undefined) mockCline.ask = vi.fn().mockResolvedValue(undefined) mockCline.recordToolError = vi.fn() + mockCline.recordToolUsage = vi.fn() mockCline.sayAndCreateMissingParamError = vi.fn().mockResolvedValue("Missing param error") mockAskApproval = vi.fn().mockResolvedValue(true) @@ -178,46 +208,7 @@ describe("applyCodeTool", () => { }) }) - describe("feature flag", () => { - it("returns error when applyEnabled is false", async () => { - await executeApplyCodeTool({}, { applyEnabled: false }) - - // The actual implementation should check this flag - const provider = mockCline.providerRef.deref() - const state = await provider?.getState() - expect(state?.applyEnabled).toBe(false) - }) - - it("proceeds when applyEnabled is true", async () => { - // Mock successful API responses - const mockStream = { - [Symbol.asyncIterator]: vi.fn().mockReturnValue({ - next: vi - .fn() - .mockResolvedValueOnce({ - value: { type: "text", text: "```typescript\n" + originalContent + "\n```" }, - done: false, - }) - .mockResolvedValueOnce({ done: true }), - }), - } - mockCline.api.createMessage.mockReturnValue(mockStream) - - await executeApplyCodeTool({}, { applyEnabled: true }) - - expect(mockCline.api.createMessage).toHaveBeenCalled() - }) - }) - describe("file validation", () => { - it("returns error when file does not exist", async () => { - // For new files, the tool should still work but generate full file content - await executeApplyCodeTool({}, { fileExists: false }) - - // The tool should handle non-existent files by creating them - expect(mockCline.recordToolError).not.toHaveBeenCalled() - }) - it("validates access with rooIgnoreController", async () => { await executeApplyCodeTool({}, { accessAllowed: false }) @@ -227,30 +218,35 @@ describe("applyCodeTool", () => { }) describe("two-stage API workflow", () => { - it("makes two API calls with correct prompts", async () => { + it("makes two API calls with correct prompts and isolated context", async () => { // Mock file read - mockedReadFile.mockResolvedValue(Buffer.from(originalContent)) + mockedReadFile.mockResolvedValue(originalContent) // Mock successful API responses - const generatedCode = `function getData() { - try { - return fetch('/api/data').then(res => { - if (!res.ok) throw new Error('Failed to fetch'); - return res.json(); - }); - } catch (error) { - console.error('Error fetching data:', error); - throw error; - } + const generatedCode = `try { + return fetch('/api/data').then(res => { + if (!res.ok) throw new Error('Failed to fetch'); + return res.json(); + }); +} catch (error) { + console.error('Error fetching data:', error); + throw error; }` - // First API call response (code generation) + // First API call response (code generation) - returns JSON const mockStream1 = { [Symbol.asyncIterator]: vi.fn().mockReturnValue({ next: vi .fn() .mockResolvedValueOnce({ - value: { type: "text", text: "```typescript\n" + generatedCode + "\n```" }, + value: { + type: "text", + text: JSON.stringify({ + file: testFilePath, + type: "snippet", + code: generatedCode, + }), + }, done: false, }) .mockResolvedValueOnce({ done: true }), @@ -265,7 +261,23 @@ describe("applyCodeTool", () => { .mockResolvedValueOnce({ value: { type: "text", - text: "\ntest/file.ts\nmock diff content\n", + text: `<<<<<<< SEARCH +function getData() { + return fetch('/api/data').then(res => res.json()); +} +======= +function getData() { + try { + return fetch('/api/data').then(res => { + if (!res.ok) throw new Error('Failed to fetch'); + return res.json(); + }); + } catch (error) { + console.error('Error fetching data:', error); + throw error; + } +} +>>>>>>> REPLACE`, }, done: false, }) @@ -273,30 +285,46 @@ describe("applyCodeTool", () => { }), } - mockCline.api.createMessage.mockReturnValueOnce(mockStream1).mockReturnValueOnce(mockStream2) + // Mock the isolated API handler + const mockIsolatedApiHandler = { + createMessage: vi.fn().mockReturnValue(mockStream2), + getModel: vi.fn().mockReturnValue({ + id: "claude-3", + info: { contextWindow: 200000 }, + }), + countTokens: vi.fn().mockResolvedValue(100), + } + + mockCline.api.createMessage.mockReturnValueOnce(mockStream1) + mockedBuildApiHandler.mockReturnValue(mockIsolatedApiHandler) await executeApplyCodeTool() - // Verify two API calls were made - expect(mockCline.api.createMessage).toHaveBeenCalledTimes(2) - - // Verify first call (code generation) + // Verify first API call (code generation) uses main API + expect(mockCline.api.createMessage).toHaveBeenCalledTimes(1) const firstCall = mockCline.api.createMessage.mock.calls[0] - expect(firstCall[0]).toContain("generate code") - expect(firstCall[0]).toContain(testInstruction) + expect(firstCall[0]).toContain("code generation expert") + expect(firstCall[1][0].content[0].text).toContain(testInstruction) - // Verify second call (diff generation) - const secondCall = mockCline.api.createMessage.mock.calls[1] - expect(secondCall[0]).toContain("create a diff") - expect(secondCall[1]).toEqual([ - { role: "user", content: expect.stringContaining(originalContent) }, - { role: "assistant", content: generatedCode }, - ]) + // Verify second API call uses isolated handler + expect(mockedBuildApiHandler).toHaveBeenCalledWith(mockCline.apiConfiguration) + expect(mockIsolatedApiHandler.createMessage).toHaveBeenCalledTimes(1) + + // Verify the isolated call has the hardcoded system prompt + const secondCall = mockIsolatedApiHandler.createMessage.mock.calls[0] + expect(secondCall[0]).toContain("specialized diff generation model") + expect(secondCall[0]).toContain("Your ONLY task is to generate accurate diff patches") + + // Verify the isolated call has clean context (no conversation history) + expect(secondCall[1]).toHaveLength(1) + expect(secondCall[1][0].role).toBe("user") + expect(secondCall[1][0].content[0].text).toContain("Original file content:") + expect(secondCall[1][0].content[0].text).toContain("New code to integrate:") }) - it("delegates to applyDiffTool after generating diff", async () => { + it("applies the generated diff using diffStrategy", async () => { // Mock file read - mockedReadFile.mockResolvedValue(Buffer.from(originalContent)) + mockedReadFile.mockResolvedValue(originalContent) // Mock API responses const mockStream1 = { @@ -304,7 +332,14 @@ describe("applyCodeTool", () => { next: vi .fn() .mockResolvedValueOnce({ - value: { type: "text", text: "```typescript\ngenerated code\n```" }, + value: { + type: "text", + text: JSON.stringify({ + file: testFilePath, + type: "snippet", + code: "generated code", + }), + }, done: false, }) .mockResolvedValueOnce({ done: true }), @@ -318,7 +353,7 @@ describe("applyCodeTool", () => { .mockResolvedValueOnce({ value: { type: "text", - text: "\ntest/file.ts\ndiff content\n", + text: "<<<<<<< SEARCH\noriginal\n=======\nmodified\n>>>>>>> REPLACE", }, done: false, }) @@ -326,35 +361,79 @@ describe("applyCodeTool", () => { }), } - mockCline.api.createMessage.mockReturnValueOnce(mockStream1).mockReturnValueOnce(mockStream2) + const mockIsolatedApiHandler = { + createMessage: vi.fn().mockReturnValue(mockStream2), + getModel: vi.fn().mockReturnValue({ + id: "claude-3", + info: { contextWindow: 200000 }, + }), + countTokens: vi.fn().mockResolvedValue(100), + } + + mockCline.api.createMessage.mockReturnValueOnce(mockStream1) + mockedBuildApiHandler.mockReturnValue(mockIsolatedApiHandler) await executeApplyCodeTool() - // Verify applyDiffToolLegacy was called - expect(mockedApplyDiffToolLegacy).toHaveBeenCalledWith( - mockCline, - expect.objectContaining({ - type: "tool_use", - name: "apply_diff", - params: { - path: testFilePath, - diff: "diff content", - }, - }), - mockAskApproval, - mockHandleError, - mockPushToolResult, - mockRemoveClosingTag, + // Verify diffStrategy.applyDiff was called with the generated diff + expect(mockCline.diffStrategy.applyDiff).toHaveBeenCalledWith( + originalContent, + "<<<<<<< SEARCH\noriginal\n=======\nmodified\n>>>>>>> REPLACE", ) + + // Verify the diff view was updated + expect(mockCline.diffViewProvider.update).toHaveBeenCalledWith("modified content", true) + expect(mockPushToolResult).toHaveBeenCalledWith("File updated successfully") + }) + + it("handles new file creation", async () => { + // Mock file doesn't exist + mockedFileExistsAtPath.mockResolvedValue(false) + + // Mock API response for new file + const newFileContent = `export function newFunction() { + return "Hello, World!"; +}` + + const mockStream = { + [Symbol.asyncIterator]: vi.fn().mockReturnValue({ + next: vi + .fn() + .mockResolvedValueOnce({ + value: { + type: "text", + text: JSON.stringify({ + file: testFilePath, + type: "full_file", + code: newFileContent, + }), + }, + done: false, + }) + .mockResolvedValueOnce({ done: true }), + }), + } + + mockCline.api.createMessage.mockReturnValue(mockStream) + + await executeApplyCodeTool({}, { fileExists: false }) + + // Verify only one API call was made (no diff generation for new files) + expect(mockCline.api.createMessage).toHaveBeenCalledTimes(1) + expect(mockedBuildApiHandler).not.toHaveBeenCalled() + + // Verify the file was created + expect(mockCline.diffViewProvider.editType).toBe("create") + expect(mockCline.diffViewProvider.update).toHaveBeenCalledWith(newFileContent, true) }) }) describe("error handling", () => { it("handles API errors in first stage", async () => { // Mock file read - mockedReadFile.mockResolvedValue(Buffer.from(originalContent)) + mockedReadFile.mockResolvedValue(originalContent) - // Mock API error - need to return a proper async iterator that throws + // Mock API error const mockStream = { [Symbol.asyncIterator]: vi.fn().mockImplementation(() => ({ next: vi.fn().mockRejectedValue(new Error("API error")), @@ -376,17 +455,17 @@ describe("applyCodeTool", () => { expect(mockHandleError).toHaveBeenCalledWith("applying code", expect.any(Error)) }) - it("handles malformed API responses", async () => { + it("handles malformed JSON responses", async () => { // Mock file read - mockedReadFile.mockResolvedValue(Buffer.from(originalContent)) + mockedReadFile.mockResolvedValue(originalContent) - // Mock malformed response (no code blocks) + // Mock malformed response (invalid JSON) const mockStream = { [Symbol.asyncIterator]: vi.fn().mockReturnValue({ next: vi .fn() .mockResolvedValueOnce({ - value: { type: "text", text: "Just some text without code blocks" }, + value: { type: "text", text: "Just some text without valid JSON" }, done: false, }) .mockResolvedValueOnce({ done: true }), @@ -397,12 +476,80 @@ describe("applyCodeTool", () => { await executeApplyCodeTool() - // The error should be about parsing JSON, not "No code was generated" expect(mockCline.say).toHaveBeenCalledWith( "error", expect.stringContaining("Failed to parse code generation response"), ) }) + + it("handles diff application failures", async () => { + // Mock file read + mockedReadFile.mockResolvedValue(originalContent) + + // Mock successful code generation + const mockStream1 = { + [Symbol.asyncIterator]: vi.fn().mockReturnValue({ + next: vi + .fn() + .mockResolvedValueOnce({ + value: { + type: "text", + text: JSON.stringify({ + file: testFilePath, + type: "snippet", + code: "generated code", + }), + }, + done: false, + }) + .mockResolvedValueOnce({ done: true }), + }), + } + + // Mock successful diff generation + const mockStream2 = { + [Symbol.asyncIterator]: vi.fn().mockReturnValue({ + next: vi + .fn() + .mockResolvedValueOnce({ + value: { + type: "text", + text: "<<<<<<< SEARCH\nwrong content\n=======\nmodified\n>>>>>>> REPLACE", + }, + done: false, + }) + .mockResolvedValueOnce({ done: true }), + }), + } + + const mockIsolatedApiHandler = { + createMessage: vi.fn().mockReturnValue(mockStream2), + getModel: vi.fn().mockReturnValue({ + id: "claude-3", + info: { contextWindow: 200000 }, + }), + countTokens: vi.fn().mockResolvedValue(100), + } + + mockCline.api.createMessage.mockReturnValueOnce(mockStream1) + mockedBuildApiHandler.mockReturnValue(mockIsolatedApiHandler) + + // Mock diff application failure + mockCline.diffStrategy.applyDiff.mockResolvedValue({ + success: false, + error: "Could not find search content", + }) + + await executeApplyCodeTool() + + expect(mockCline.say).toHaveBeenCalledWith( + "error", + "Failed to apply generated diff: Could not find search content", + ) + expect(mockPushToolResult).toHaveBeenCalledWith( + "Failed to apply generated diff: Could not find search content", + ) + }) }) describe("partial execution", () => { @@ -413,4 +560,67 @@ describe("applyCodeTool", () => { expect(mockPushToolResult).not.toHaveBeenCalled() }) }) + + describe("user approval", () => { + it("reverts changes when user denies approval", async () => { + // Mock file read + mockedReadFile.mockResolvedValue(originalContent) + + // Mock successful API responses + const mockStream1 = { + [Symbol.asyncIterator]: vi.fn().mockReturnValue({ + next: vi + .fn() + .mockResolvedValueOnce({ + value: { + type: "text", + text: JSON.stringify({ + file: testFilePath, + type: "snippet", + code: "generated code", + }), + }, + done: false, + }) + .mockResolvedValueOnce({ done: true }), + }), + } + + const mockStream2 = { + [Symbol.asyncIterator]: vi.fn().mockReturnValue({ + next: vi + .fn() + .mockResolvedValueOnce({ + value: { + type: "text", + text: "<<<<<<< SEARCH\noriginal\n=======\nmodified\n>>>>>>> REPLACE", + }, + done: false, + }) + .mockResolvedValueOnce({ done: true }), + }), + } + + const mockIsolatedApiHandler = { + createMessage: vi.fn().mockReturnValue(mockStream2), + getModel: vi.fn().mockReturnValue({ + id: "claude-3", + info: { contextWindow: 200000 }, + }), + countTokens: vi.fn().mockResolvedValue(100), + } + + mockCline.api.createMessage.mockReturnValueOnce(mockStream1) + mockedBuildApiHandler.mockReturnValue(mockIsolatedApiHandler) + + // User denies approval + mockAskApproval.mockResolvedValue(false) + + await executeApplyCodeTool() + + expect(mockCline.diffViewProvider.revertChanges).toHaveBeenCalled() + expect(mockCline.diffViewProvider.saveChanges).not.toHaveBeenCalled() + expect(mockPushToolResult).not.toHaveBeenCalled() + }) + }) }) diff --git a/src/core/tools/applyCodeTool.ts b/src/core/tools/applyCodeTool.ts index 6065438528..3a01b7db73 100644 --- a/src/core/tools/applyCodeTool.ts +++ b/src/core/tools/applyCodeTool.ts @@ -95,6 +95,7 @@ IMPORTANT: - Ensure proper escaping of quotes and newlines in JSON` // Make first API call for code generation + // This uses the existing API handler with full context const codeGenMessages = [ { role: "user" as const, @@ -128,12 +129,32 @@ IMPORTANT: return } - // Stage 2: Focused Diff Generation + // Stage 2: Focused Diff Generation with ISOLATED context let diffContent = "" if (fileExists && codeGenResult.type === "snippet") { - const diffGenPrompt = `You are a diff generation expert. Given the original file content and new code, generate a standard unified diff patch to integrate the new code into the original file. + // HARDCODED OPTIMIZED PROMPT - This is the key difference + // This prompt is specifically designed for diff generation without any conversational noise + const DIFF_GENERATION_SYSTEM_PROMPT = `You are a specialized diff generation model. Your ONLY task is to generate accurate diff patches. -Original file content: +RULES: +1. You will receive EXACTLY two inputs: original file content and new code to integrate +2. You must output ONLY the diff patch in the specified format +3. Do NOT add any explanations, comments, or conversational text +4. Focus ONLY on the mechanical task of creating the diff +5. Ensure the SEARCH section matches the original content EXACTLY (including whitespace) +6. Place the new code in the most logical location within the file + +OUTPUT FORMAT: +<<<<<<< SEARCH +[exact content from original file] +======= +[integrated content with new code] +>>>>>>> REPLACE + +You may use multiple SEARCH/REPLACE blocks if needed.` + + // Simplified prompt for isolated context - no conversational instructions + const diffGenPrompt = `Original file content: \`\`\` ${originalContent} \`\`\` @@ -141,32 +162,26 @@ ${originalContent} New code to integrate: \`\`\` ${codeGenResult.code} -\`\`\` +\`\`\`` -Generate a diff in the exact format used by the apply_diff tool: -<<<<<<< SEARCH -[exact content to find including whitespace] -======= -[new content to replace with] ->>>>>>> REPLACE - -IMPORTANT: -- The SEARCH section must exactly match existing content -- Include proper indentation and whitespace -- You may use multiple SEARCH/REPLACE blocks if needed -- Focus only on integrating the new code logically` - - const diffGenMessages = [ + // Create a COMPLETELY ISOLATED API call + // This is a new, independent message array with NO conversation history + const isolatedDiffMessages = [ { role: "user" as const, content: [{ type: "text" as const, text: diffGenPrompt }], }, ] - const diffGenStream = cline.api.createMessage( - "You are a diff generation expert. Generate accurate diffs for code integration.", - diffGenMessages, - { taskId: cline.taskId, mode: "diff_generation" }, + // Create a new API handler instance to ensure complete isolation + // This prevents any context bleeding from the main conversation + const isolatedApiHandler = buildApiHandler(cline.apiConfiguration) + + // Make the isolated API call with the hardcoded system prompt + const diffGenStream = isolatedApiHandler.createMessage( + DIFF_GENERATION_SYSTEM_PROMPT, + isolatedDiffMessages, + { taskId: `${cline.taskId}-diff-gen`, mode: "diff_generation" }, ) for await (const chunk of diffGenStream) {