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
*