mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-09-05 08:10:14 +00:00
feat: implement apply_code tool for two-stage code generation
- Add new apply_code tool that uses a two-stage approach: 1. Creative code generation based on instructions 2. Focused diff generation with clean context - Add applyEnabled configuration option to enable/disable the feature - Update tool descriptions and type definitions - Add comprehensive tests for the new tool - Integrate with existing diff application infrastructure This addresses issue #6159 by providing a more reliable alternative to apply_diff that separates creative code generation from technical diff creation, reducing errors from noisy conversational contexts.
This commit is contained in:
parent
1e17b3b3bb
commit
d875f222be
9 changed files with 756 additions and 1 deletions
|
|
@ -63,6 +63,7 @@ export const DEFAULT_CONSECUTIVE_MISTAKE_LIMIT = 3
|
|||
const baseProviderSettingsSchema = z.object({
|
||||
includeMaxTokens: z.boolean().optional(),
|
||||
diffEnabled: z.boolean().optional(),
|
||||
applyEnabled: z.boolean().optional(),
|
||||
todoListEnabled: z.boolean().optional(),
|
||||
fuzzyMatchThreshold: z.number().optional(),
|
||||
modelTemperature: z.number().nullish(),
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ export const toolNames = [
|
|||
"read_file",
|
||||
"write_to_file",
|
||||
"apply_diff",
|
||||
"apply_code",
|
||||
"insert_content",
|
||||
"search_and_replace",
|
||||
"search_files",
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ import { Task } from "../task/Task"
|
|||
import { codebaseSearchTool } from "../tools/codebaseSearchTool"
|
||||
import { experiments, EXPERIMENT_IDS } from "../../shared/experiments"
|
||||
import { applyDiffToolLegacy } from "../tools/applyDiffTool"
|
||||
import { applyCodeTool } from "../tools/applyCodeTool"
|
||||
|
||||
/**
|
||||
* Processes and presents assistant message content to the user interface.
|
||||
|
|
@ -214,6 +215,8 @@ export async function presentAssistantMessage(cline: Task) {
|
|||
const modeName = getModeBySlug(mode, customModes)?.name ?? mode
|
||||
return `[${block.name} in ${modeName} mode: '${message}']`
|
||||
}
|
||||
case "apply_code":
|
||||
return `[${block.name} for '${block.params.path}' with instruction: '${block.params.instruction}']`
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -522,6 +525,9 @@ export async function presentAssistantMessage(cline: Task) {
|
|||
askFinishSubTaskApproval,
|
||||
)
|
||||
break
|
||||
case "apply_code":
|
||||
await applyCodeTool(cline, block, askApproval, handleError, pushToolResult, removeClosingTag)
|
||||
break
|
||||
}
|
||||
|
||||
break
|
||||
|
|
|
|||
34
src/core/prompts/tools/apply-code.ts
Normal file
34
src/core/prompts/tools/apply-code.ts
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
import { ToolArgs } from "./types"
|
||||
|
||||
export function getApplyCodeDescription(args: ToolArgs): string {
|
||||
return `## apply_code
|
||||
Description: Request to apply code changes using a two-stage approach for improved reliability. This tool first generates code based on your instruction, then creates an accurate diff to integrate it into the existing file. This approach separates creative code generation from technical diff creation, resulting in more reliable code modifications.
|
||||
|
||||
Parameters:
|
||||
- path: (required) The path of the file to modify (relative to the current workspace directory ${args.cwd})
|
||||
- instruction: (required) Clear instruction describing what code changes to make
|
||||
|
||||
Usage:
|
||||
<apply_code>
|
||||
<path>File path here</path>
|
||||
<instruction>Your instruction for code changes</instruction>
|
||||
</apply_code>
|
||||
|
||||
Example: Adding a new function to an existing file
|
||||
<apply_code>
|
||||
<path>src/utils.ts</path>
|
||||
<instruction>Add a function called calculateAverage that takes an array of numbers and returns their average</instruction>
|
||||
</apply_code>
|
||||
|
||||
Example: Modifying existing code
|
||||
<apply_code>
|
||||
<path>src/api/handler.ts</path>
|
||||
<instruction>Update the error handling in the fetchData function to include retry logic with exponential backoff</instruction>
|
||||
</apply_code>
|
||||
|
||||
Benefits over apply_diff:
|
||||
- More reliable: Separates code generation from diff creation
|
||||
- Cleaner context: Each stage has focused, minimal context
|
||||
- Better success rate: Reduces failures due to inaccurate diffs
|
||||
- Natural instructions: Use plain language instead of crafting diffs`
|
||||
}
|
||||
|
|
@ -23,6 +23,7 @@ import { getSwitchModeDescription } from "./switch-mode"
|
|||
import { getNewTaskDescription } from "./new-task"
|
||||
import { getCodebaseSearchDescription } from "./codebase-search"
|
||||
import { getUpdateTodoListDescription } from "./update-todo-list"
|
||||
import { getApplyCodeDescription } from "./apply-code"
|
||||
import { CodeIndexManager } from "../../../services/code-index/manager"
|
||||
|
||||
// Map of tool names to their description functions
|
||||
|
|
@ -46,6 +47,7 @@ const toolDescriptionMap: Record<string, (args: ToolArgs) => string | undefined>
|
|||
search_and_replace: (args) => getSearchAndReplaceDescription(args),
|
||||
apply_diff: (args) =>
|
||||
args.diffStrategy ? args.diffStrategy.getToolDescription({ cwd: args.cwd, toolOptions: args.toolOptions }) : "",
|
||||
apply_code: (args) => getApplyCodeDescription(args),
|
||||
update_todo_list: (args) => getUpdateTodoListDescription(args),
|
||||
}
|
||||
|
||||
|
|
@ -114,6 +116,11 @@ export function getToolDescriptionsForMode(
|
|||
tools.delete("update_todo_list")
|
||||
}
|
||||
|
||||
// Conditionally exclude apply_code if disabled in settings
|
||||
if (settings?.applyEnabled === false) {
|
||||
tools.delete("apply_code")
|
||||
}
|
||||
|
||||
// Map tool descriptions for allowed tools
|
||||
const descriptions = Array.from(tools).map((toolName) => {
|
||||
const descriptionFn = toolDescriptionMap[toolName]
|
||||
|
|
@ -148,4 +155,5 @@ export {
|
|||
getInsertContentDescription,
|
||||
getSearchAndReplaceDescription,
|
||||
getCodebaseSearchDescription,
|
||||
getApplyCodeDescription,
|
||||
}
|
||||
|
|
|
|||
416
src/core/tools/__tests__/applyCodeTool.spec.ts
Normal file
416
src/core/tools/__tests__/applyCodeTool.spec.ts
Normal file
|
|
@ -0,0 +1,416 @@
|
|||
import { vi, describe, it, expect, beforeEach } from "vitest"
|
||||
import type { MockedFunction } from "vitest"
|
||||
import { applyCodeTool } from "../applyCodeTool"
|
||||
import { ToolUse, ToolResponse } from "../../../shared/tools"
|
||||
import { fileExistsAtPath } from "../../../utils/fs"
|
||||
import { getReadablePath } from "../../../utils/path"
|
||||
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().mockResolvedValue(Buffer.from("original content")),
|
||||
}))
|
||||
|
||||
// Mock dependencies
|
||||
vi.mock("path", async () => {
|
||||
const originalPath = await vi.importActual("path")
|
||||
return {
|
||||
...originalPath,
|
||||
resolve: vi.fn().mockImplementation((...args) => {
|
||||
const separator = process.platform === "win32" ? "\\" : "/"
|
||||
return args.join(separator)
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
vi.mock("../../../utils/fs", () => ({
|
||||
fileExistsAtPath: vi.fn().mockResolvedValue(true),
|
||||
}))
|
||||
|
||||
vi.mock("../../../utils/path", () => ({
|
||||
getReadablePath: vi.fn().mockReturnValue("test/file.ts"),
|
||||
}))
|
||||
|
||||
vi.mock("../../prompts/responses", () => ({
|
||||
formatResponse: {
|
||||
toolError: vi.fn((msg) => `Error: ${msg}`),
|
||||
rooIgnoreError: vi.fn((path) => `Access denied: ${path}`),
|
||||
createPrettyPatch: vi.fn(() => "mock-diff"),
|
||||
},
|
||||
}))
|
||||
|
||||
vi.mock("../applyDiffTool", () => ({
|
||||
applyDiffToolLegacy: vi.fn().mockResolvedValue(undefined),
|
||||
}))
|
||||
|
||||
// Import after mocking to get the mocked version
|
||||
import { applyDiffToolLegacy } from "../applyDiffTool"
|
||||
import fs from "fs/promises"
|
||||
|
||||
describe("applyCodeTool", () => {
|
||||
// Test data
|
||||
const testFilePath = "test/file.ts"
|
||||
const absoluteFilePath = process.platform === "win32" ? "C:\\test\\file.ts" : "/test/file.ts"
|
||||
const testInstruction = "Add error handling to the function"
|
||||
const originalContent = `function getData() {
|
||||
return fetch('/api/data').then(res => res.json());
|
||||
}`
|
||||
|
||||
// Mocked functions
|
||||
const mockedFileExistsAtPath = fileExistsAtPath as MockedFunction<typeof fileExistsAtPath>
|
||||
const mockedGetReadablePath = getReadablePath as MockedFunction<typeof getReadablePath>
|
||||
const mockedPathResolve = path.resolve as MockedFunction<typeof path.resolve>
|
||||
const mockedApplyDiffToolLegacy = applyDiffToolLegacy as MockedFunction<typeof applyDiffToolLegacy>
|
||||
const mockedReadFile = fs.readFile as MockedFunction<typeof fs.readFile>
|
||||
|
||||
const mockCline: any = {}
|
||||
let mockAskApproval: ReturnType<typeof vi.fn>
|
||||
let mockHandleError: ReturnType<typeof vi.fn>
|
||||
let mockPushToolResult: ReturnType<typeof vi.fn>
|
||||
let mockRemoveClosingTag: ReturnType<typeof vi.fn>
|
||||
let toolResult: ToolResponse | undefined
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
|
||||
mockedPathResolve.mockReturnValue(absoluteFilePath)
|
||||
mockedFileExistsAtPath.mockResolvedValue(true)
|
||||
mockedGetReadablePath.mockReturnValue(testFilePath)
|
||||
|
||||
mockCline.cwd = "/"
|
||||
mockCline.consecutiveMistakeCount = 0
|
||||
mockCline.api = {
|
||||
createMessage: vi.fn(),
|
||||
getModel: vi.fn().mockReturnValue({
|
||||
id: "claude-3",
|
||||
info: { contextWindow: 200000 },
|
||||
}),
|
||||
}
|
||||
mockCline.providerRef = {
|
||||
deref: vi.fn().mockReturnValue({
|
||||
getState: vi.fn().mockResolvedValue({
|
||||
applyEnabled: true,
|
||||
}),
|
||||
}),
|
||||
}
|
||||
mockCline.rooIgnoreController = {
|
||||
validateAccess: vi.fn().mockReturnValue(true),
|
||||
}
|
||||
mockCline.diffViewProvider = {
|
||||
reset: vi.fn().mockResolvedValue(undefined),
|
||||
}
|
||||
mockCline.say = vi.fn().mockResolvedValue(undefined)
|
||||
mockCline.ask = vi.fn().mockResolvedValue(undefined)
|
||||
mockCline.recordToolError = vi.fn()
|
||||
mockCline.sayAndCreateMissingParamError = vi.fn().mockResolvedValue("Missing param error")
|
||||
|
||||
mockAskApproval = vi.fn().mockResolvedValue(true)
|
||||
mockHandleError = vi.fn().mockResolvedValue(undefined)
|
||||
mockRemoveClosingTag = vi.fn((tag, content) => content)
|
||||
mockPushToolResult = vi.fn()
|
||||
|
||||
toolResult = undefined
|
||||
})
|
||||
|
||||
/**
|
||||
* Helper function to execute the apply code tool
|
||||
*/
|
||||
async function executeApplyCodeTool(
|
||||
params: Partial<ToolUse["params"]> = {},
|
||||
options: {
|
||||
fileExists?: boolean
|
||||
isPartial?: boolean
|
||||
accessAllowed?: boolean
|
||||
applyEnabled?: boolean
|
||||
} = {},
|
||||
): Promise<ToolResponse | undefined> {
|
||||
const fileExists = options.fileExists ?? true
|
||||
const isPartial = options.isPartial ?? false
|
||||
const accessAllowed = options.accessAllowed ?? true
|
||||
const applyEnabled = options.applyEnabled ?? true
|
||||
|
||||
mockedFileExistsAtPath.mockResolvedValue(fileExists)
|
||||
mockCline.rooIgnoreController.validateAccess.mockReturnValue(accessAllowed)
|
||||
mockCline.providerRef.deref().getState.mockResolvedValue({ applyEnabled })
|
||||
|
||||
const toolUse: ToolUse = {
|
||||
type: "tool_use",
|
||||
name: "apply_code",
|
||||
params: {
|
||||
path: testFilePath,
|
||||
instruction: testInstruction,
|
||||
...params,
|
||||
},
|
||||
partial: isPartial,
|
||||
}
|
||||
|
||||
await applyCodeTool(
|
||||
mockCline,
|
||||
toolUse,
|
||||
mockAskApproval,
|
||||
mockHandleError,
|
||||
(result: ToolResponse) => {
|
||||
toolResult = result
|
||||
mockPushToolResult(result)
|
||||
},
|
||||
mockRemoveClosingTag,
|
||||
)
|
||||
|
||||
return toolResult
|
||||
}
|
||||
|
||||
describe("parameter validation", () => {
|
||||
it("handles missing path parameter", async () => {
|
||||
await executeApplyCodeTool({ path: undefined })
|
||||
|
||||
expect(mockCline.sayAndCreateMissingParamError).toHaveBeenCalledWith("apply_code", "path")
|
||||
expect(mockPushToolResult).toHaveBeenCalledWith("Missing param error")
|
||||
})
|
||||
|
||||
it("handles missing instruction parameter", async () => {
|
||||
await executeApplyCodeTool({ instruction: undefined })
|
||||
|
||||
expect(mockCline.sayAndCreateMissingParamError).toHaveBeenCalledWith("apply_code", "instruction")
|
||||
expect(mockPushToolResult).toHaveBeenCalledWith("Missing param error")
|
||||
})
|
||||
})
|
||||
|
||||
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 })
|
||||
|
||||
expect(mockCline.rooIgnoreController.validateAccess).toHaveBeenCalledWith(testFilePath)
|
||||
expect(mockPushToolResult).toHaveBeenCalledWith(expect.stringContaining("Access denied"))
|
||||
})
|
||||
})
|
||||
|
||||
describe("two-stage API workflow", () => {
|
||||
it("makes two API calls with correct prompts", async () => {
|
||||
// Mock file read
|
||||
mockedReadFile.mockResolvedValue(Buffer.from(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;
|
||||
}
|
||||
}`
|
||||
|
||||
// First API call response (code generation)
|
||||
const mockStream1 = {
|
||||
[Symbol.asyncIterator]: vi.fn().mockReturnValue({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
value: { type: "text", text: "```typescript\n" + generatedCode + "\n```" },
|
||||
done: false,
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
|
||||
// Second API call response (diff generation)
|
||||
const mockStream2 = {
|
||||
[Symbol.asyncIterator]: vi.fn().mockReturnValue({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
value: {
|
||||
type: "text",
|
||||
text: "<apply_diff>\n<path>test/file.ts</path>\n<diff>mock diff content</diff>\n</apply_diff>",
|
||||
},
|
||||
done: false,
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
|
||||
mockCline.api.createMessage.mockReturnValueOnce(mockStream1).mockReturnValueOnce(mockStream2)
|
||||
|
||||
await executeApplyCodeTool()
|
||||
|
||||
// Verify two API calls were made
|
||||
expect(mockCline.api.createMessage).toHaveBeenCalledTimes(2)
|
||||
|
||||
// Verify first call (code generation)
|
||||
const firstCall = mockCline.api.createMessage.mock.calls[0]
|
||||
expect(firstCall[0]).toContain("generate code")
|
||||
expect(firstCall[0]).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 },
|
||||
])
|
||||
})
|
||||
|
||||
it("delegates to applyDiffTool after generating diff", async () => {
|
||||
// Mock file read
|
||||
mockedReadFile.mockResolvedValue(Buffer.from(originalContent))
|
||||
|
||||
// Mock API responses
|
||||
const mockStream1 = {
|
||||
[Symbol.asyncIterator]: vi.fn().mockReturnValue({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
value: { type: "text", text: "```typescript\ngenerated code\n```" },
|
||||
done: false,
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
|
||||
const mockStream2 = {
|
||||
[Symbol.asyncIterator]: vi.fn().mockReturnValue({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
value: {
|
||||
type: "text",
|
||||
text: "<apply_diff>\n<path>test/file.ts</path>\n<diff>diff content</diff>\n</apply_diff>",
|
||||
},
|
||||
done: false,
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
|
||||
mockCline.api.createMessage.mockReturnValueOnce(mockStream1).mockReturnValueOnce(mockStream2)
|
||||
|
||||
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,
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
describe("error handling", () => {
|
||||
it("handles API errors in first stage", async () => {
|
||||
// Mock file read
|
||||
mockedReadFile.mockResolvedValue(Buffer.from(originalContent))
|
||||
|
||||
// Mock API error - need to return a proper async iterator that throws
|
||||
const mockStream = {
|
||||
[Symbol.asyncIterator]: vi.fn().mockImplementation(() => ({
|
||||
next: vi.fn().mockRejectedValue(new Error("API error")),
|
||||
})),
|
||||
}
|
||||
mockCline.api.createMessage.mockReturnValue(mockStream)
|
||||
|
||||
await executeApplyCodeTool()
|
||||
|
||||
expect(mockHandleError).toHaveBeenCalledWith("applying code", expect.any(Error))
|
||||
})
|
||||
|
||||
it("handles file read errors", async () => {
|
||||
// Mock file read error
|
||||
mockedReadFile.mockRejectedValue(new Error("File read error"))
|
||||
|
||||
await executeApplyCodeTool()
|
||||
|
||||
expect(mockHandleError).toHaveBeenCalledWith("applying code", expect.any(Error))
|
||||
})
|
||||
|
||||
it("handles malformed API responses", async () => {
|
||||
// Mock file read
|
||||
mockedReadFile.mockResolvedValue(Buffer.from(originalContent))
|
||||
|
||||
// Mock malformed response (no code blocks)
|
||||
const mockStream = {
|
||||
[Symbol.asyncIterator]: vi.fn().mockReturnValue({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
value: { type: "text", text: "Just some text without code blocks" },
|
||||
done: false,
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
|
||||
mockCline.api.createMessage.mockReturnValue(mockStream)
|
||||
|
||||
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"),
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
describe("partial execution", () => {
|
||||
it("returns early for partial blocks", async () => {
|
||||
await executeApplyCodeTool({}, { isPartial: true })
|
||||
|
||||
expect(mockCline.api.createMessage).not.toHaveBeenCalled()
|
||||
expect(mockPushToolResult).not.toHaveBeenCalled()
|
||||
})
|
||||
})
|
||||
})
|
||||
280
src/core/tools/applyCodeTool.ts
Normal file
280
src/core/tools/applyCodeTool.ts
Normal file
|
|
@ -0,0 +1,280 @@
|
|||
import path from "path"
|
||||
import fs from "fs/promises"
|
||||
|
||||
import { TelemetryService } from "@roo-code/telemetry"
|
||||
import { DEFAULT_WRITE_DELAY_MS } from "@roo-code/types"
|
||||
|
||||
import { ClineSayTool } from "../../shared/ExtensionMessage"
|
||||
import { getReadablePath } from "../../utils/path"
|
||||
import { Task } from "../task/Task"
|
||||
import { ToolUse, RemoveClosingTag, AskApproval, HandleError, PushToolResult } from "../../shared/tools"
|
||||
import { formatResponse } from "../prompts/responses"
|
||||
import { fileExistsAtPath } from "../../utils/fs"
|
||||
import { RecordSource } from "../context-tracking/FileContextTrackerTypes"
|
||||
import { buildApiHandler } from "../../api"
|
||||
|
||||
interface CodeGenerationResult {
|
||||
file: string
|
||||
type: "snippet" | "full_file"
|
||||
code: string
|
||||
}
|
||||
|
||||
export async function applyCodeTool(
|
||||
cline: Task,
|
||||
block: ToolUse,
|
||||
askApproval: AskApproval,
|
||||
handleError: HandleError,
|
||||
pushToolResult: PushToolResult,
|
||||
removeClosingTag: RemoveClosingTag,
|
||||
) {
|
||||
const relPath: string | undefined = block.params.path
|
||||
const instruction: string | undefined = block.params.instruction
|
||||
|
||||
const sharedMessageProps: ClineSayTool = {
|
||||
tool: "applyCode",
|
||||
path: getReadablePath(cline.cwd, removeClosingTag("path", relPath)),
|
||||
instruction: removeClosingTag("instruction", instruction),
|
||||
}
|
||||
|
||||
try {
|
||||
if (block.partial) {
|
||||
// Update GUI message
|
||||
await cline.ask("tool", JSON.stringify(sharedMessageProps), block.partial).catch(() => {})
|
||||
return
|
||||
} else {
|
||||
if (!relPath) {
|
||||
cline.consecutiveMistakeCount++
|
||||
cline.recordToolError("apply_code")
|
||||
pushToolResult(await cline.sayAndCreateMissingParamError("apply_code", "path"))
|
||||
return
|
||||
}
|
||||
|
||||
if (!instruction) {
|
||||
cline.consecutiveMistakeCount++
|
||||
cline.recordToolError("apply_code")
|
||||
pushToolResult(await cline.sayAndCreateMissingParamError("apply_code", "instruction"))
|
||||
return
|
||||
}
|
||||
|
||||
const accessAllowed = cline.rooIgnoreController?.validateAccess(relPath)
|
||||
|
||||
if (!accessAllowed) {
|
||||
await cline.say("rooignore_error", relPath)
|
||||
pushToolResult(formatResponse.toolError(formatResponse.rooIgnoreError(relPath)))
|
||||
return
|
||||
}
|
||||
|
||||
const absolutePath = path.resolve(cline.cwd, relPath)
|
||||
const fileExists = await fileExistsAtPath(absolutePath)
|
||||
|
||||
// Read the original file content if it exists
|
||||
let originalContent = ""
|
||||
if (fileExists) {
|
||||
originalContent = await fs.readFile(absolutePath, "utf-8")
|
||||
}
|
||||
|
||||
// Stage 1: Creative Code Generation
|
||||
const codeGenPrompt = `You are a code generation expert. Generate code based on the following instruction.
|
||||
|
||||
File: ${relPath}
|
||||
${fileExists ? `Current content:\n\`\`\`\n${originalContent}\n\`\`\`` : "File does not exist yet."}
|
||||
|
||||
Instruction: ${instruction}
|
||||
|
||||
Respond with a JSON object in this exact format:
|
||||
{
|
||||
"file": "${relPath}",
|
||||
"type": "${fileExists ? "snippet" : "full_file"}",
|
||||
"code": "your generated code here"
|
||||
}
|
||||
|
||||
IMPORTANT:
|
||||
- For existing files, generate only the new/modified code snippet
|
||||
- For new files, generate the complete file content
|
||||
- Do not include any markdown code blocks in the "code" field
|
||||
- Ensure proper escaping of quotes and newlines in JSON`
|
||||
|
||||
// Make first API call for code generation
|
||||
const codeGenMessages = [
|
||||
{
|
||||
role: "user" as const,
|
||||
content: [{ type: "text" as const, text: codeGenPrompt }],
|
||||
},
|
||||
]
|
||||
|
||||
const codeGenStream = cline.api.createMessage(
|
||||
"You are a code generation expert. Generate code exactly as requested.",
|
||||
codeGenMessages,
|
||||
{ taskId: cline.taskId, mode: "code_generation" },
|
||||
)
|
||||
|
||||
let codeGenResponse = ""
|
||||
for await (const chunk of codeGenStream) {
|
||||
if (chunk.type === "text") {
|
||||
codeGenResponse += chunk.text
|
||||
}
|
||||
}
|
||||
|
||||
// Parse the code generation result
|
||||
let codeGenResult: CodeGenerationResult
|
||||
try {
|
||||
codeGenResult = JSON.parse(codeGenResponse)
|
||||
} catch (error) {
|
||||
cline.consecutiveMistakeCount++
|
||||
cline.recordToolError("apply_code")
|
||||
const formattedError = `Failed to parse code generation response: ${error.message}\n\nResponse: ${codeGenResponse}`
|
||||
await cline.say("error", formattedError)
|
||||
pushToolResult(formattedError)
|
||||
return
|
||||
}
|
||||
|
||||
// Stage 2: Focused Diff Generation
|
||||
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.
|
||||
|
||||
Original file content:
|
||||
\`\`\`
|
||||
${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 = [
|
||||
{
|
||||
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" },
|
||||
)
|
||||
|
||||
for await (const chunk of diffGenStream) {
|
||||
if (chunk.type === "text") {
|
||||
diffContent += chunk.text
|
||||
}
|
||||
}
|
||||
|
||||
// Apply the diff using the existing diff strategy
|
||||
const diffResult = (await cline.diffStrategy?.applyDiff(originalContent, diffContent)) ?? {
|
||||
success: false,
|
||||
error: "No diff strategy available",
|
||||
}
|
||||
|
||||
if (!diffResult.success) {
|
||||
cline.consecutiveMistakeCount++
|
||||
cline.recordToolError("apply_code")
|
||||
const formattedError = `Failed to apply generated diff: ${diffResult.error}`
|
||||
await cline.say("error", formattedError)
|
||||
pushToolResult(formattedError)
|
||||
return
|
||||
}
|
||||
|
||||
// Show diff view before asking for approval
|
||||
cline.diffViewProvider.editType = "modify"
|
||||
await cline.diffViewProvider.open(relPath)
|
||||
await cline.diffViewProvider.update(diffResult.content, true)
|
||||
cline.diffViewProvider.scrollToFirstDiff()
|
||||
|
||||
// Check if file is write-protected
|
||||
const isWriteProtected = cline.rooProtectedController?.isWriteProtected(relPath) || false
|
||||
|
||||
const completeMessage = JSON.stringify({
|
||||
...sharedMessageProps,
|
||||
diff: formatResponse.createPrettyPatch(relPath, originalContent, diffResult.content),
|
||||
isProtected: isWriteProtected,
|
||||
} satisfies ClineSayTool)
|
||||
|
||||
const didApprove = await askApproval("tool", completeMessage, undefined, isWriteProtected)
|
||||
|
||||
if (!didApprove) {
|
||||
await cline.diffViewProvider.revertChanges()
|
||||
return
|
||||
}
|
||||
|
||||
// Save the changes
|
||||
const provider = cline.providerRef.deref()
|
||||
const state = await provider?.getState()
|
||||
const diagnosticsEnabled = state?.diagnosticsEnabled ?? true
|
||||
const writeDelayMs = state?.writeDelayMs ?? DEFAULT_WRITE_DELAY_MS
|
||||
await cline.diffViewProvider.saveChanges(diagnosticsEnabled, writeDelayMs)
|
||||
} else {
|
||||
// For new files or full file replacements, use the generated code directly
|
||||
cline.diffViewProvider.editType = fileExists ? "modify" : "create"
|
||||
await cline.diffViewProvider.open(relPath)
|
||||
await cline.diffViewProvider.update(codeGenResult.code, true)
|
||||
cline.diffViewProvider.scrollToFirstDiff()
|
||||
|
||||
// Check if file is write-protected
|
||||
const isWriteProtected = cline.rooProtectedController?.isWriteProtected(relPath) || false
|
||||
|
||||
const completeMessage = JSON.stringify({
|
||||
...sharedMessageProps,
|
||||
content: fileExists ? undefined : codeGenResult.code,
|
||||
diff: fileExists
|
||||
? formatResponse.createPrettyPatch(relPath, originalContent, codeGenResult.code)
|
||||
: undefined,
|
||||
isProtected: isWriteProtected,
|
||||
} satisfies ClineSayTool)
|
||||
|
||||
const didApprove = await askApproval("tool", completeMessage, undefined, isWriteProtected)
|
||||
|
||||
if (!didApprove) {
|
||||
await cline.diffViewProvider.revertChanges()
|
||||
return
|
||||
}
|
||||
|
||||
// Save the changes
|
||||
const provider = cline.providerRef.deref()
|
||||
const state = await provider?.getState()
|
||||
const diagnosticsEnabled = state?.diagnosticsEnabled ?? true
|
||||
const writeDelayMs = state?.writeDelayMs ?? DEFAULT_WRITE_DELAY_MS
|
||||
await cline.diffViewProvider.saveChanges(diagnosticsEnabled, writeDelayMs)
|
||||
}
|
||||
|
||||
// Track file edit operation
|
||||
if (relPath) {
|
||||
await cline.fileContextTracker.trackFileContext(relPath, "roo_edited" as RecordSource)
|
||||
}
|
||||
|
||||
// Used to determine if we should wait for busy terminal to update before sending api request
|
||||
cline.didEditFile = true
|
||||
|
||||
// Get the formatted response message
|
||||
const message = await cline.diffViewProvider.pushToolWriteResult(cline, cline.cwd, !fileExists)
|
||||
|
||||
pushToolResult(message)
|
||||
|
||||
await cline.diffViewProvider.reset()
|
||||
|
||||
cline.consecutiveMistakeCount = 0
|
||||
cline.recordToolUsage("apply_code")
|
||||
|
||||
return
|
||||
}
|
||||
} catch (error) {
|
||||
await handleError("applying code", error)
|
||||
await cline.diffViewProvider.reset()
|
||||
return
|
||||
}
|
||||
}
|
||||
|
|
@ -315,6 +315,7 @@ export interface ClineSayTool {
|
|||
tool:
|
||||
| "editedExistingFile"
|
||||
| "appliedDiff"
|
||||
| "applyCode"
|
||||
| "newFileCreated"
|
||||
| "codebaseSearch"
|
||||
| "readFile"
|
||||
|
|
@ -346,6 +347,7 @@ export interface ClineSayTool {
|
|||
endLine?: number
|
||||
lineNumber?: number
|
||||
query?: string
|
||||
instruction?: string
|
||||
batchFiles?: Array<{
|
||||
path: string
|
||||
lineSnippet: string
|
||||
|
|
|
|||
|
|
@ -65,6 +65,7 @@ export const toolParamNames = [
|
|||
"query",
|
||||
"args",
|
||||
"todos",
|
||||
"instruction",
|
||||
] as const
|
||||
|
||||
export type ToolParamName = (typeof toolParamNames)[number]
|
||||
|
|
@ -164,6 +165,11 @@ export interface SearchAndReplaceToolUse extends ToolUse {
|
|||
Partial<Pick<Record<ToolParamName, string>, "use_regex" | "ignore_case" | "start_line" | "end_line">>
|
||||
}
|
||||
|
||||
export interface ApplyCodeToolUse extends ToolUse {
|
||||
name: "apply_code"
|
||||
params: Partial<Pick<Record<ToolParamName, string>, "path" | "instruction">>
|
||||
}
|
||||
|
||||
// Define tool group configuration
|
||||
export type ToolGroupConfig = {
|
||||
tools: readonly string[]
|
||||
|
|
@ -176,6 +182,7 @@ export const TOOL_DISPLAY_NAMES: Record<ToolName, string> = {
|
|||
fetch_instructions: "fetch instructions",
|
||||
write_to_file: "write files",
|
||||
apply_diff: "apply changes",
|
||||
apply_code: "apply code",
|
||||
search_files: "search files",
|
||||
list_files: "list files",
|
||||
list_code_definition_names: "list definitions",
|
||||
|
|
@ -205,7 +212,7 @@ export const TOOL_GROUPS: Record<ToolGroup, ToolGroupConfig> = {
|
|||
],
|
||||
},
|
||||
edit: {
|
||||
tools: ["apply_diff", "write_to_file", "insert_content", "search_and_replace"],
|
||||
tools: ["apply_diff", "apply_code", "write_to_file", "insert_content", "search_and_replace"],
|
||||
},
|
||||
browser: {
|
||||
tools: ["browser_action"],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue