fix(coderd/x/chatd): prevent invalid tool results from poisoning chat history (#24663)

- **computeruse.go**: Decode base64 screenshot data before storing in
`ToolResponse.Data` (was casting base64 string to bytes without
decoding)
- **chatloop.go**: Re-encode `ToolResponse.Data` to base64 via
`base64.StdEncoding.EncodeToString` instead of `string()` cast
- **mcpclient.go**: UTF-8 validate all text from MCP responses in
`convertCallResult()` using `strings.ToValidUTF8`
- **chatprompt.go (persist)**: Defense-in-depth UTF-8 sanitization of
text and media Text fields before database storage
- **chatprompt.go (replay)**: Antivenom layer that validates base64 and
UTF-8 at read time, auto-healing already-poisoned chats without
requiring a migration
- `TestToolResultAntivenom`: 4 subtests covering poisoned text, poisoned
media, valid media round-trip, and media with invalid UTF-8 text
-  Adds `TestConvertCallResult_UTF8Sanitization`: 4 subtests covering invalid
UTF-8 in TextContent, EmbeddedResource, valid passthrough, and
multi-part
- Adds `TestComputerUseTool_Run_ScreenshotDataIsDecodedBinary`: Verifies no
double-encode in the computer-use path
- Updated existing computer-use tests for the new decoded-binary
contract

> 🤖
This commit is contained in:
Cian Johnston
2026-04-23 19:58:38 +01:00
committed by GitHub
parent c602a31856
commit a02339c66a
12 changed files with 638 additions and 49 deletions
+25 -4
View File
@@ -2,6 +2,7 @@ package chattool
import (
"context"
"encoding/base64"
"fmt"
"time"
@@ -168,7 +169,7 @@ func (t *computerUseTool) Run(ctx context.Context, call fantasy.ToolCall) (fanta
return t.captureScreenshot(ctx, conn, declaredWidth, declaredHeight)
}
func (*computerUseTool) captureScreenshot(
func (t *computerUseTool) captureScreenshot(
ctx context.Context,
conn workspacesdk.AgentConn,
declaredWidth, declaredHeight int,
@@ -179,7 +180,16 @@ func (*computerUseTool) captureScreenshot(
fmt.Sprintf("screenshot failed: %v", err),
), nil
}
return fantasy.NewImageResponse([]byte(screenResp.ScreenshotData), "image/png"), nil
screenData, err := base64.StdEncoding.DecodeString(screenResp.ScreenshotData)
if err != nil {
t.logger.Error(ctx, "failed to decode screenshot base64 in captureScreenshot",
slog.Error(err),
)
return fantasy.NewTextErrorResponse(
fmt.Sprintf("failed to decode screenshot data: %v", err),
), nil
}
return fantasy.NewImageResponse(screenData, "image/png"), nil
}
func (t *computerUseTool) captureSharedScreenshot(
@@ -194,22 +204,33 @@ func (t *computerUseTool) captureSharedScreenshot(
), nil
}
screenData, err := base64.StdEncoding.DecodeString(screenResp.ScreenshotData)
if err != nil {
t.logger.Error(ctx, "failed to decode screenshot base64 in captureSharedScreenshot",
slog.Error(err),
)
return fantasy.NewTextErrorResponse(
fmt.Sprintf("failed to decode screenshot data: %v", err),
), nil
}
attachmentName := fmt.Sprintf(
"screenshot-%s.png",
t.clock.Now().UTC().Format("2006-01-02T15-04-05Z"),
)
if t.storeFile == nil {
t.logger.Warn(ctx, "screenshot attachment storage is not configured")
return fantasy.NewImageResponse([]byte(screenResp.ScreenshotData), "image/png"), nil
return fantasy.NewImageResponse(screenData, "image/png"), nil
}
response := fantasy.NewImageResponse(screenData, "image/png")
attachment, err := storeScreenshotAttachment(
ctx,
t.storeFile,
attachmentName,
screenResp.ScreenshotData,
)
response := fantasy.NewImageResponse([]byte(screenResp.ScreenshotData), "image/png")
if err != nil {
t.logger.Warn(ctx, "failed to persist screenshot attachment",
slog.F("attachment_name", attachmentName),
+73 -5
View File
@@ -70,7 +70,9 @@ func TestComputerUseTool_Run_Screenshot(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, "image", resp.Type)
assert.Equal(t, "image/png", resp.MediaType)
assert.Equal(t, []byte("iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR4nGP4n539HwAHFwLVF8kc1wAAAABJRU5ErkJggg=="), resp.Data)
expectedBinary, decErr := base64.StdEncoding.DecodeString("iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR4nGP4n539HwAHFwLVF8kc1wAAAABJRU5ErkJggg==")
require.NoError(t, decErr)
assert.Equal(t, expectedBinary, resp.Data)
assert.False(t, resp.IsError)
}
@@ -118,7 +120,9 @@ func TestComputerUseTool_Run_Screenshot_PersistsAttachment(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, "image", resp.Type)
assert.Equal(t, "image/png", resp.MediaType)
assert.Equal(t, []byte(screenshotPNG), resp.Data)
expectedBinary, decErr := base64.StdEncoding.DecodeString(screenshotPNG)
require.NoError(t, decErr)
assert.Equal(t, expectedBinary, resp.Data)
assert.Contains(t, storedName, "screenshot-")
assert.Equal(t, "image/png", storedType)
expectedPNG, decodeErr := base64.StdEncoding.DecodeString(screenshotPNG)
@@ -200,7 +204,9 @@ func TestComputerUseTool_Run_Screenshot_OversizedAttachmentFallsBackToImage(t *t
assert.Equal(t, "image", resp.Type)
assert.Equal(t, "image/png", resp.MediaType)
assert.False(t, resp.IsError)
require.Len(t, resp.Data, len(oversizedScreenshot))
expectedOversized, decErr := base64.StdEncoding.DecodeString(oversizedScreenshot)
require.NoError(t, decErr)
require.Len(t, resp.Data, len(expectedOversized))
attachments, err := chattool.AttachmentsFromMetadata(resp.Metadata)
require.NoError(t, err)
assert.Empty(t, attachments)
@@ -260,7 +266,9 @@ func TestComputerUseTool_Run_LeftClick(t *testing.T) {
resp, err := tool.Run(context.Background(), call)
require.NoError(t, err)
assert.Equal(t, "image", resp.Type)
assert.Equal(t, []byte(followUpScreenshot), resp.Data)
expectedBinary, decErr := base64.StdEncoding.DecodeString(followUpScreenshot)
require.NoError(t, decErr)
assert.Equal(t, expectedBinary, resp.Data)
attachments, err := chattool.AttachmentsFromMetadata(resp.Metadata)
require.NoError(t, err)
assert.Empty(t, attachments)
@@ -307,13 +315,73 @@ func TestComputerUseTool_Run_Wait(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, "image", resp.Type)
assert.Equal(t, "image/png", resp.MediaType)
assert.Equal(t, []byte(followUpScreenshot), resp.Data)
expectedBinary, decErr := base64.StdEncoding.DecodeString(followUpScreenshot)
require.NoError(t, decErr)
assert.Equal(t, expectedBinary, resp.Data)
assert.False(t, resp.IsError)
attachments, err := chattool.AttachmentsFromMetadata(resp.Metadata)
require.NoError(t, err)
assert.Empty(t, attachments)
}
func TestComputerUseTool_Run_ScreenshotDataIsDecodedBinary(t *testing.T) {
t.Parallel()
ctrl := gomock.NewController(t)
mockConn := agentconnmock.NewMockAgentConn(ctrl)
geometry := workspacesdk.DefaultDesktopGeometry()
// A known base64 string (1x1 red PNG).
const screenshotBase64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR4nGP4z8BQDwAEgAF/pooBPQAAAABJRU5ErkJggg=="
mockConn.EXPECT().ExecuteDesktopAction(
gomock.Any(),
gomock.AssignableToTypeOf(workspacesdk.DesktopAction{}),
).Return(workspacesdk.DesktopActionResponse{
Output: "screenshot",
ScreenshotData: screenshotBase64,
ScreenshotWidth: geometry.DeclaredWidth,
ScreenshotHeight: geometry.DeclaredHeight,
}, nil)
tool := chattool.NewComputerUseTool(
geometry.DeclaredWidth,
geometry.DeclaredHeight,
func(_ context.Context) (workspacesdk.AgentConn, error) {
return mockConn, nil
},
nil,
quartz.NewReal(),
slogtest.Make(t, nil),
)
call := fantasy.ToolCall{
ID: "test-decode-1",
Name: "computer",
Input: `{"action":"screenshot"}`,
}
resp, err := tool.Run(context.Background(), call)
require.NoError(t, err)
assert.Equal(t, "image", resp.Type)
assert.Equal(t, "image/png", resp.MediaType)
// Data must contain decoded binary, not the base64 string
// reinterpreted as bytes.
expectedBinary, err := base64.StdEncoding.DecodeString(screenshotBase64)
require.NoError(t, err)
assert.Equal(t, expectedBinary, resp.Data,
"ToolResponse.Data should contain decoded binary, not base64-as-bytes")
// Verify that re-encoding produces the original base64 string.
// This is the round-trip that the chat loop performs when
// building the API response.
reEncoded := base64.StdEncoding.EncodeToString(resp.Data)
assert.Equal(t, screenshotBase64, reEncoded,
"re-encoding Data should produce the original base64 string (no double-encode)")
}
func TestComputerUseTool_Run_ConnError(t *testing.T) {
t.Parallel()
+2 -2
View File
@@ -107,7 +107,7 @@ func convertMCPToolResponse(
for _, c := range resp.Content {
switch c.Type {
case "text":
textParts = append(textParts, c.Text)
textParts = append(textParts, strings.ToValidUTF8(c.Text, "\uFFFD"))
case "image", "audio":
if c.Data == "" {
continue
@@ -129,7 +129,7 @@ func convertMCPToolResponse(
binaryResult = &r
}
default:
textParts = append(textParts, c.Text)
textParts = append(textParts, strings.ToValidUTF8(c.Text, "\uFFFD"))
}
}