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:
@@ -0,0 +1,5 @@
|
||||
package mcpclient
|
||||
|
||||
// ConvertCallResultForTest exposes convertCallResult for external
|
||||
// tests.
|
||||
var ConvertCallResultForTest = convertCallResult
|
||||
@@ -607,7 +607,7 @@ func convertCallResult(
|
||||
for _, item := range result.Content {
|
||||
switch c := item.(type) {
|
||||
case mcp.TextContent:
|
||||
textParts = append(textParts, c.Text)
|
||||
textParts = append(textParts, strings.ToValidUTF8(c.Text, "\uFFFD"))
|
||||
case mcp.ImageContent:
|
||||
data, err := base64.StdEncoding.DecodeString(
|
||||
c.Data,
|
||||
@@ -653,7 +653,7 @@ func convertCallResult(
|
||||
// regardless of form.
|
||||
switch r := c.Resource.(type) {
|
||||
case mcp.TextResourceContents:
|
||||
textParts = append(textParts, r.Text)
|
||||
textParts = append(textParts, strings.ToValidUTF8(r.Text, "\uFFFD"))
|
||||
case mcp.BlobResourceContents:
|
||||
data, err := base64.StdEncoding.DecodeString(
|
||||
r.Blob,
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
@@ -1271,3 +1272,77 @@ func TestModelIntent_Run_FallbackOnBadJSON(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
assert.True(t, resp.IsError, "malformed input should produce an error response")
|
||||
}
|
||||
|
||||
func TestConvertCallResult_UTF8Sanitization(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
result *mcp.CallToolResult
|
||||
wantContains []string
|
||||
}{
|
||||
{
|
||||
name: "InvalidUTF8InTextContent",
|
||||
result: &mcp.CallToolResult{
|
||||
Content: []mcp.Content{
|
||||
mcp.TextContent{
|
||||
Text: "Hello" + string([]byte{0xFF, 0xFE, 0x80}) + "World",
|
||||
},
|
||||
},
|
||||
},
|
||||
wantContains: []string{"Hello", "World", "\uFFFD"},
|
||||
},
|
||||
{
|
||||
name: "InvalidUTF8InEmbeddedResourceText",
|
||||
result: &mcp.CallToolResult{
|
||||
Content: []mcp.Content{
|
||||
mcp.EmbeddedResource{
|
||||
Resource: mcp.TextResourceContents{
|
||||
Text: "Content" + string([]byte{0x80, 0x81, 0x82}),
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
wantContains: []string{"Content"},
|
||||
},
|
||||
{
|
||||
name: "ValidUTF8PassesThrough",
|
||||
result: &mcp.CallToolResult{
|
||||
Content: []mcp.Content{
|
||||
mcp.TextContent{
|
||||
Text: "Hello, 世界! 🌍",
|
||||
},
|
||||
},
|
||||
},
|
||||
wantContains: []string{"Hello, 世界! 🌍"},
|
||||
},
|
||||
{
|
||||
name: "MultipleTextPartsAllSanitized",
|
||||
result: &mcp.CallToolResult{
|
||||
Content: []mcp.Content{
|
||||
mcp.TextContent{
|
||||
Text: "Part1" + string([]byte{0xFF}),
|
||||
},
|
||||
mcp.TextContent{
|
||||
Text: "Part2" + string([]byte{0xFE}),
|
||||
},
|
||||
},
|
||||
},
|
||||
wantContains: []string{"Part1", "Part2"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
resp := mcpclient.ConvertCallResultForTest(tt.result)
|
||||
|
||||
require.True(t, utf8.ValidString(resp.Content),
|
||||
"response content must be valid UTF-8")
|
||||
for _, want := range tt.wantContains {
|
||||
require.Contains(t, resp.Content, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user