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
+142
View File
@@ -2,6 +2,7 @@ package chatloop //nolint:testpackage // Uses internal symbols.
import (
"context"
"encoding/base64"
"errors"
"iter"
"strings"
@@ -9,13 +10,16 @@ import (
"sync/atomic"
"testing"
"time"
"unicode/utf8"
"charm.land/fantasy"
fantasyanthropic "charm.land/fantasy/providers/anthropic"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
"github.com/coder/coder/v2/coderd/x/chatd/chatretry"
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
@@ -2115,3 +2119,141 @@ func TestRun_PrepareMessagesOnlyFiresOnce(t *testing.T) {
// PrepareMessages is called before each of the 3 steps.
require.Equal(t, 3, int(prepareCalls.Load()))
}
func TestExecuteSingleTool_MediaBase64Encoding(t *testing.T) {
t.Parallel()
originalBytes := []byte{0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10}
metrics := NewMetrics(prometheus.NewRegistry())
logger := slog.Make()
t.Run("EncodesRawBytesToBase64", func(t *testing.T) {
t.Parallel()
tool := fantasy.NewAgentTool(
"screenshot",
"takes a screenshot",
func(_ context.Context, _ struct{}, _ fantasy.ToolCall) (fantasy.ToolResponse, error) {
return fantasy.ToolResponse{
Type: "image",
Data: originalBytes,
MediaType: "image/jpeg",
}, nil
},
)
toolMap := map[string]fantasy.AgentTool{
"screenshot": tool,
}
tc := fantasy.ToolCallContent{
ToolCallID: "call-1",
ToolName: "screenshot",
Input: "{}",
}
result := executeSingleTool(
context.Background(),
toolMap,
tc,
metrics,
logger,
"fake", "fake-model",
map[string]bool{},
[]string{"screenshot"},
map[string]struct{}{},
)
media, ok := result.Result.(fantasy.ToolResultOutputContentMedia)
require.True(t, ok, "expected ToolResultOutputContentMedia")
require.Equal(t, "image/jpeg", media.MediaType)
decoded, err := base64.StdEncoding.DecodeString(media.Data)
require.NoError(t, err, "Data should be valid base64")
require.Equal(t, originalBytes, decoded)
})
t.Run("SanitizesInvalidUTF8InContent", func(t *testing.T) {
t.Parallel()
tool := fantasy.NewAgentTool(
"screenshot",
"takes a screenshot",
func(_ context.Context, _ struct{}, _ fantasy.ToolCall) (fantasy.ToolResponse, error) {
return fantasy.ToolResponse{
Type: "image",
Data: originalBytes,
MediaType: "image/png",
Content: "hello\xffworld",
}, nil
},
)
toolMap := map[string]fantasy.AgentTool{
"screenshot": tool,
}
tc := fantasy.ToolCallContent{
ToolCallID: "call-2",
ToolName: "screenshot",
Input: "{}",
}
result := executeSingleTool(
context.Background(),
toolMap,
tc,
metrics,
logger,
"fake", "fake-model",
map[string]bool{},
[]string{"screenshot"},
map[string]struct{}{},
)
media, ok := result.Result.(fantasy.ToolResultOutputContentMedia)
require.True(t, ok, "expected ToolResultOutputContentMedia")
require.True(t, utf8.ValidString(media.Text), "Text should be valid UTF-8")
require.Contains(t, media.Text, "hello")
require.Contains(t, media.Text, "world")
})
t.Run("SanitizesInvalidUTF8InTextResult", func(t *testing.T) {
t.Parallel()
tool := fantasy.NewAgentTool(
"echo",
"echoes input",
func(_ context.Context, _ struct{}, _ fantasy.ToolCall) (fantasy.ToolResponse, error) {
return fantasy.ToolResponse{
Content: "hello\xffworld",
}, nil
},
)
toolMap := map[string]fantasy.AgentTool{
"echo": tool,
}
tc := fantasy.ToolCallContent{
ToolCallID: "call-3",
ToolName: "echo",
Input: "{}",
}
result := executeSingleTool(
context.Background(),
toolMap,
tc,
metrics,
logger,
"fake", "fake-model",
map[string]bool{},
[]string{"echo"},
map[string]struct{}{},
)
textOutput, ok := result.Result.(fantasy.ToolResultOutputContentText)
require.True(t, ok, "expected ToolResultOutputContentText, got %T", result.Result)
require.True(t, utf8.ValidString(textOutput.Text), "Text should be valid UTF-8")
require.Contains(t, textOutput.Text, "hello")
require.Contains(t, textOutput.Text, "world")
})
}