diff --git a/src/core/tools/ExecuteCommandTool.ts b/src/core/tools/ExecuteCommandTool.ts index 8fcb917b13..b356ebab1b 100644 --- a/src/core/tools/ExecuteCommandTool.ts +++ b/src/core/tools/ExecuteCommandTool.ts @@ -11,7 +11,7 @@ import { Task } from "../task/Task" import { ToolUse, ToolResponse } from "../../shared/tools" import { formatResponse } from "../prompts/responses" -import { unescapeHtmlEntities } from "../../utils/text-normalization" +import { unescapeHtmlEntities, sanitizeForPromptInjection } from "../../utils/text-normalization" import { ExitCodeDetails, RooTerminalCallbacks, RooTerminalProcess } from "../../integrations/terminal/types" import { TerminalRegistry } from "../../integrations/terminal/TerminalRegistry" import { Terminal } from "../../integrations/terminal/Terminal" @@ -459,6 +459,8 @@ export async function executeCommandInTerminal( await onCompletedPromise } + const safeResult = sanitizeForPromptInjection(result) + if (message) { const { text, images } = message await task.say("user_feedback", text, images) @@ -468,7 +470,7 @@ export async function executeCommandInTerminal( formatResponse.toolResult( [ `Command is still running in terminal from '${terminal.getCurrentWorkingDirectory().toPosix()}'.`, - result.length > 0 ? `Here's the output so far:\n${result}\n` : "\n", + safeResult.length > 0 ? `Here's the output so far:\n${safeResult}\n` : "\n", `\n${text}\n`, ].join("\n"), images, @@ -509,14 +511,14 @@ export async function executeCommandInTerminal( return [ false, - `Command executed in terminal within working directory '${currentWorkingDir}'. ${exitStatus}\nOutput:\n${result}`, + `Command executed in terminal within working directory '${currentWorkingDir}'. ${exitStatus}\nOutput:\n${safeResult}`, ] } else { return [ false, [ `Command is still running in terminal ${workingDir ? ` from '${workingDir.toPosix()}'` : ""}.`, - result.length > 0 ? `Here's the output so far:\n${result}\n` : "\n", + safeResult.length > 0 ? `Here's the output so far:\n${safeResult}\n` : "\n", "You will be updated on the terminal status and new output in the future.", ].join("\n"), ] @@ -569,7 +571,7 @@ function formatPersistedOutput( `Output (${sizeStr}) persisted. Artifact ID: ${artifactId}`, "", "Preview:", - result.preview, + sanitizeForPromptInjection(result.preview), "", "Use read_command_output tool to view full output if needed.", ].join("\n") diff --git a/src/core/tools/ReadFileTool.ts b/src/core/tools/ReadFileTool.ts index 8ad6a3b33d..200419b04d 100644 --- a/src/core/tools/ReadFileTool.ts +++ b/src/core/tools/ReadFileTool.ts @@ -21,6 +21,7 @@ import { RecordSource } from "../context-tracking/FileContextTrackerTypes" import { isPathOutsideWorkspace } from "../../utils/pathUtils" import { getReadablePath } from "../../utils/path" import { extractTextFromFile, addLineNumbers, getSupportedBinaryFormats } from "../../integrations/misc/extract-text" +import { sanitizeForPromptInjection } from "../../utils/text-normalization" import { readWithIndentation, readWithSlice } from "../../integrations/misc/indentation-reader" import { DEFAULT_LINE_LIMIT } from "../prompts/tools/native-tools/read_file" import type { ToolUse, PushToolResult } from "../../shared/tools" @@ -221,7 +222,7 @@ export class ReadFileTool extends BaseTool<"read_file"> { await task.fileContextTracker.trackFileContext(relPath, "read_tool" as RecordSource) updateFileResult(relPath, { - nativeContent: `File: ${relPath}\n${result}`, + nativeContent: `File: ${relPath}\n${sanitizeForPromptInjection(result)}`, }) } catch (error) { const errorMsg = error instanceof Error ? error.message : String(error) @@ -397,7 +398,7 @@ export class ReadFileTool extends BaseTool<"read_file"> { updateFileResult(relPath, { nativeContent: lineCount > 0 - ? `File: ${relPath}\nLines 1-${lineCount}:\n${numberedContent}` + ? `File: ${relPath}\nLines 1-${lineCount}:\n${sanitizeForPromptInjection(numberedContent)}` : `File: ${relPath}\nNote: File is empty`, }) return @@ -794,7 +795,7 @@ export class ReadFileTool extends BaseTool<"read_file"> { } } - results.push(`File: ${relPath}\n${content}`) + results.push(`File: ${relPath}\n${sanitizeForPromptInjection(content)}`) // Track file in context await task.fileContextTracker.trackFileContext(relPath, "read_tool") diff --git a/src/integrations/misc/extract-text.ts b/src/integrations/misc/extract-text.ts index f29fa915d1..fc4eca812a 100644 --- a/src/integrations/misc/extract-text.ts +++ b/src/integrations/misc/extract-text.ts @@ -7,6 +7,7 @@ import { isBinaryFile } from "isbinaryfile" import { extractTextFromXLSX } from "./extract-text-from-xlsx" import { readWithSlice } from "./indentation-reader" import { DEFAULT_LINE_LIMIT } from "../../core/prompts/tools/native-tools/read_file" +import { sanitizeForPromptInjection } from "../../utils/text-normalization" async function extractTextFromPDF(filePath: string): Promise { const dataBuffer = await fs.readFile(filePath) @@ -91,7 +92,7 @@ export async function extractTextFromFileWithMetadata( const extractor = SUPPORTED_BINARY_FORMATS[fileExtension as keyof typeof SUPPORTED_BINARY_FORMATS] if (extractor) { // For binary formats, extract and count lines - const content = await extractor(filePath) + const content = sanitizeForPromptInjection(await extractor(filePath)) const lines = content.split("\n") return { content, @@ -130,7 +131,7 @@ export async function extractTextFromFileWithMetadata( */ export async function extractTextFromFile(filePath: string): Promise { const result = await extractTextFromFileWithMetadata(filePath) - return result.content + return sanitizeForPromptInjection(result.content) } export function addLineNumbers(content: string, startLine: number = 1): string { diff --git a/src/utils/__tests__/text-normalization.spec.ts b/src/utils/__tests__/text-normalization.spec.ts index e672617d18..b5f3470997 100644 --- a/src/utils/__tests__/text-normalization.spec.ts +++ b/src/utils/__tests__/text-normalization.spec.ts @@ -1,4 +1,4 @@ -import { normalizeString, unescapeHtmlEntities } from "../text-normalization" +import { normalizeString, unescapeHtmlEntities, sanitizeForPromptInjection } from "../text-normalization" describe("Text normalization utilities", () => { describe("normalizeString", () => { @@ -100,5 +100,26 @@ describe("Text normalization utilities", () => { const expected = "array[0] and [1]" expect(unescapeHtmlEntities(input)).toBe(expected) }) + + describe("sanitizeForPromptInjection", () => { + it("escapes XML-like tags", () => { + expect(sanitizeForPromptInjection("inject")).toBe( + "\\inject\\", + ) + }) + + it("escapes HTML comment-like sequences", () => { + expect(sanitizeForPromptInjection("")).toBe("\\") + }) + + it("does not escape standalone less-than signs", () => { + expect(sanitizeForPromptInjection("a < b")).toBe("a < b") + }) + + it("returns original string when no tags are present", () => { + const original = "Plain text without any markup" + expect(sanitizeForPromptInjection(original)).toBe(original) + }) + }) }) }) diff --git a/src/utils/text-normalization.ts b/src/utils/text-normalization.ts index 9e25d140c4..88df034a35 100644 --- a/src/utils/text-normalization.ts +++ b/src/utils/text-normalization.ts @@ -76,6 +76,17 @@ export function normalizeString(str: string, options: NormalizeOptions = DEFAULT return normalized } +/** + * Escapes potential XML/HTML-like tags to prevent indirect prompt injection + * via tool outputs (command output, file contents, etc.). + * + * @param content The untrusted content to sanitize + * @returns The sanitized content with tag-like sequences escaped + */ +export function sanitizeForPromptInjection(content: string): string { + return content.replace(/<(\/?[a-zA-Z!?])/g, "\\<$1") +} + /** * Unescapes common HTML entities in a string *