mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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),
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user