fix(coderd/x/chatd): repair Anthropic provider tool history (#24744)

## Problem

Anthropic returns HTTP 400 when an assistant message contains a
`web_search_tool_result` block whose `tool_use_id` has no matching
earlier `server_tool_use` block in the same assistant message. A
previous fix (#24706) sanitized provider-executed tool calls without
matching results, but the opposite direction, orphaned or misordered
provider-executed results, could still slip through both the prompt
sanitizer and the persistence path.

## Fix

Tighten Anthropic provider-executed tool history handling while
preserving the useful result payload as normal assistant text when the
provider-tool metadata is unsafe.

1. Extract Anthropic provider-tool sanitization into
`coderd/x/chatd/chatsanitize` so provider-specific repair logic is no
longer spread through `chatprompt` and `chatloop`.

2. `chatsanitize.SanitizeAnthropicProviderToolHistory` removes invalid
provider-executed tool structure for Anthropic prompts: orphans in
either direction, result-before-call, duplicate IDs, invalid JSON
inputs, empty IDs and tool names, unsupported tool names, mismatched
`ProviderExecuted` flags, provider-executed blocks outside assistant
messages, and web-search results without serializable Anthropic result
metadata. Provider-executed result payloads are textified instead of
being discarded when there is text to preserve.

3. `chatsanitize.SanitizeAnthropicProviderToolContent` mirrors the same
rule at the streamed step content level. Persisted history no longer
carries invalid provider-tool blocks forward, but it keeps the result
text for future turns.

4. `chatsanitize.ApplyAnthropicProviderToolGuard` only repairs
structurally invalid Anthropic provider-tool history. It no longer
strips otherwise-valid historical `web_search` blocks just because web
search is disabled for the current request. The fail-closed fallback
also textifies provider results before removing provider-tool metadata.

Tests cover prompt sanitization, validation reason strings, result
payload textification, content-level persistence sanitization, disabled
web-search history preservation, direct pre-request guard behavior, and
the fallback strip path.

> Mux is acting on Mike's behalf.
This commit is contained in:
Michael Suchacz
2026-04-28 12:45:23 +02:00
committed by GitHub
parent 68c8499c9a
commit 99eb46dac1
8 changed files with 3467 additions and 516 deletions
+5 -4
View File
@@ -43,6 +43,7 @@ import (
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
"github.com/coder/coder/v2/coderd/x/chatd/chatretry"
"github.com/coder/coder/v2/coderd/x/chatd/chatsanitize"
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
"github.com/coder/coder/v2/coderd/x/chatd/internal/agentselect"
"github.com/coder/coder/v2/coderd/x/chatd/mcpclient"
@@ -6243,8 +6244,8 @@ func (p *Server) runChat(
if err := g2.Wait(); err != nil {
return result, err
}
prompt, sanitizeStats := chatprompt.SanitizeAnthropicProviderToolCalls(model.Provider(), prompt)
chatprompt.LogAnthropicProviderToolSanitization(
prompt, sanitizeStats := chatsanitize.SanitizeAnthropicProviderToolHistory(model.Provider(), prompt)
chatsanitize.LogAnthropicProviderToolSanitization(
ctx, logger, "persisted_history_replay", model.Provider(), model.Model(), sanitizeStats,
)
subagentInstruction := ""
@@ -6871,8 +6872,8 @@ func (p *Server) runChat(
if err != nil {
return nil, xerrors.Errorf("convert reloaded messages: %w", err)
}
reloadedPrompt, sanitizeStats := chatprompt.SanitizeAnthropicProviderToolCalls(model.Provider(), reloadedPrompt)
chatprompt.LogAnthropicProviderToolSanitization(
reloadedPrompt, sanitizeStats := chatsanitize.SanitizeAnthropicProviderToolHistory(model.Provider(), reloadedPrompt)
chatsanitize.LogAnthropicProviderToolSanitization(
reloadCtx, logger, "reload_messages", model.Provider(), model.Model(), sanitizeStats,
)
// Re-derive instruction and skills from the reloaded
+12 -69
View File
@@ -26,6 +26,7 @@ import (
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
"github.com/coder/coder/v2/coderd/x/chatd/chatretry"
"github.com/coder/coder/v2/coderd/x/chatd/chatsanitize"
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/quartz"
@@ -391,12 +392,15 @@ func Run(ctx context.Context, opts RunOptions) error {
}
prepared := make([]fantasy.Message, len(messages))
copy(prepared, messages)
prepared, sanitizeStats := chatprompt.SanitizeAnthropicProviderToolCalls(provider, prepared)
chatprompt.LogAnthropicProviderToolSanitization(
prepared, sanitizeStats := chatsanitize.SanitizeAnthropicProviderToolHistory(provider, prepared)
chatsanitize.LogAnthropicProviderToolSanitization(
ctx, opts.Logger, "pre_request", provider, modelName, sanitizeStats,
slog.F("step_index", step),
slog.F("total_steps", totalSteps),
)
prepared = chatsanitize.ApplyAnthropicProviderToolGuard(
ctx, opts.Logger, provider, modelName, prepared,
)
if applyAnthropicCaching {
addAnthropicPromptCaching(prepared)
}
@@ -529,7 +533,7 @@ func Run(ctx context.Context, opts RunOptions) error {
opts.ContextLimitFallback,
)
result.content = sanitizeAnthropicProviderToolStepContent(
result.content = chatsanitize.SanitizeAnthropicProviderToolStepContent(
ctx, opts.Logger, provider, modelName,
"dynamic_tool_persist", step, result.finishReason, result.content,
)
@@ -576,7 +580,7 @@ func Run(ctx context.Context, opts RunOptions) error {
result.providerMetadata,
opts.ContextLimitFallback,
)
result.content = sanitizeAnthropicProviderToolStepContent(
result.content = chatsanitize.SanitizeAnthropicProviderToolStepContent(
ctx, opts.Logger, provider, modelName,
"normal_persist", step, result.finishReason, result.content,
)
@@ -734,67 +738,6 @@ func Run(ctx context.Context, opts RunOptions) error {
return nil
}
func sanitizeAnthropicProviderToolStepContent(
ctx context.Context,
logger slog.Logger,
provider string,
modelName string,
phase string,
step int,
finishReason fantasy.FinishReason,
content []fantasy.Content,
) []fantasy.Content {
sanitized, stats := sanitizeAnthropicProviderToolContent(provider, content)
chatprompt.LogAnthropicProviderToolSanitization(
ctx, logger, phase, provider, modelName, stats,
slog.F("step_index", step),
slog.F("finish_reason", finishReason),
)
return sanitized
}
func sanitizeAnthropicProviderToolContent(
provider string,
content []fantasy.Content,
) ([]fantasy.Content, chatprompt.AnthropicProviderToolSanitizationStats) {
var stats chatprompt.AnthropicProviderToolSanitizationStats
if provider != fantasyanthropic.Name || len(content) == 0 {
return content, stats
}
matchedResultIDs := make(map[string]struct{})
for _, block := range content {
result, ok := fantasy.AsContentType[fantasy.ToolResultContent](block)
if !ok || !result.ProviderExecuted || result.ToolCallID == "" {
continue
}
matchedResultIDs[result.ToolCallID] = struct{}{}
}
out := make([]fantasy.Content, 0, len(content))
for _, block := range content {
toolCall, ok := fantasy.AsContentType[fantasy.ToolCallContent](block)
if ok && isAnthropicProviderExecutedToolCall(provider, toolCall) {
if _, hasResult := matchedResultIDs[toolCall.ToolCallID]; !hasResult {
stats.RemovedToolCalls++
continue
}
}
out = append(out, block)
}
if stats.RemovedToolCalls == 0 {
return content, stats
}
return out, stats
}
func isAnthropicProviderExecutedToolCall(
provider string,
toolCall fantasy.ToolCallContent,
) bool {
return provider == fantasyanthropic.Name && toolCall.ProviderExecuted
}
// guardedAttempt owns an attempt-scoped context and startup guard
// around a provider stream. release is idempotent and frees the
// attempt-scoped timer/context. finish canonicalizes startup timeout
@@ -1380,9 +1323,9 @@ func persistInterruptedStep(
provider = opts.Model.Provider()
modelName = opts.Model.Model()
}
var sanitizeStats chatprompt.AnthropicProviderToolSanitizationStats
result.content, sanitizeStats = sanitizeAnthropicProviderToolContent(provider, result.content)
chatprompt.LogAnthropicProviderToolSanitization(
var sanitizeStats chatsanitize.AnthropicProviderToolSanitizationStats
result.content, sanitizeStats = chatsanitize.SanitizeAnthropicProviderToolContent(provider, result.content)
chatsanitize.LogAnthropicProviderToolSanitization(
ctx, opts.Logger, "interrupted_persist", provider, modelName, sanitizeStats,
)
@@ -1420,7 +1363,7 @@ func persistInterruptedStep(
if _, exists := answeredToolCalls[tc.ToolCallID]; exists {
continue
}
if isAnthropicProviderExecutedToolCall(provider, tc) {
if chatsanitize.IsAnthropicProviderExecutedToolCall(provider, tc) {
continue
}
content = append(content, fantasy.ToolResultContent{
+703
View File
@@ -20,8 +20,10 @@ import (
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"cdr.dev/slog/v3/sloggers/slogtest"
"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/chatsanitize"
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
@@ -30,6 +32,93 @@ import (
const activeToolName = "read_file"
func validWebSearchProviderMetadataForTest() fantasy.ProviderMetadata {
return fantasy.ProviderMetadata{
fantasyanthropic.Name: &fantasyanthropic.WebSearchResultMetadata{
Results: []fantasyanthropic.WebSearchResultItem{
{
URL: "https://example.com",
Title: "Example",
EncryptedContent: "encrypted",
},
},
},
}
}
func safeToolCallContent(block fantasy.Content) (fantasy.ToolCallContent, bool) {
var zero fantasy.ToolCallContent
switch value := block.(type) {
case fantasy.ToolCallContent:
return value, true
case *fantasy.ToolCallContent:
if value == nil {
return zero, false
}
return *value, true
default:
return zero, false
}
}
func safeToolResultContent(block fantasy.Content) (fantasy.ToolResultContent, bool) {
var zero fantasy.ToolResultContent
switch value := block.(type) {
case fantasy.ToolResultContent:
return value, true
case *fantasy.ToolResultContent:
if value == nil {
return zero, false
}
return *value, true
default:
return zero, false
}
}
func safeToolCallPart(part fantasy.MessagePart) (fantasy.ToolCallPart, bool) {
var zero fantasy.ToolCallPart
if part == nil {
return zero, false
}
if value, ok := part.(*fantasy.ToolCallPart); ok && value == nil {
return zero, false
}
type toolCallPart = fantasy.ToolCallPart
return fantasy.AsMessagePart[toolCallPart](part)
}
func safeToolResultPart(part fantasy.MessagePart) (fantasy.ToolResultPart, bool) {
var zero fantasy.ToolResultPart
if part == nil {
return zero, false
}
if value, ok := part.(*fantasy.ToolResultPart); ok && value == nil {
return zero, false
}
type toolResultPart = fantasy.ToolResultPart
return fantasy.AsMessagePart[toolResultPart](part)
}
func toolCallContentToPart(toolCall fantasy.ToolCallContent) fantasy.ToolCallPart {
return fantasy.ToolCallPart{
ToolCallID: toolCall.ToolCallID,
ToolName: toolCall.ToolName,
Input: toolCall.Input,
ProviderExecuted: toolCall.ProviderExecuted,
ProviderOptions: fantasy.ProviderOptions(toolCall.ProviderMetadata),
}
}
func toolResultContentToPart(toolResult fantasy.ToolResultContent) fantasy.ToolResultPart {
return fantasy.ToolResultPart{
ToolCallID: toolResult.ToolCallID,
Output: toolResult.Result,
ProviderExecuted: toolResult.ProviderExecuted,
ProviderOptions: fantasy.ProviderOptions(toolResult.ProviderMetadata),
}
}
func awaitRunResult(ctx context.Context, t *testing.T, done <-chan error) error {
t.Helper()
@@ -1041,6 +1130,21 @@ func requireNoProviderExecutedToolResultContent(t *testing.T, content []fantasy.
}
}
func requireTextPrompt(t *testing.T, prompt []fantasy.Message, text string) fantasy.TextPart {
t.Helper()
for _, message := range prompt {
for _, part := range message.Content {
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](part)
if ok && textPart.Text == text {
return textPart
}
}
}
t.Fatalf("missing prompt text %q", text)
return fantasy.TextPart{}
}
func requireNoProviderExecutedToolCallPrompt(t *testing.T, prompt []fantasy.Message) {
t.Helper()
@@ -1108,6 +1212,75 @@ func requireToolResultPrompt(t *testing.T, prompt []fantasy.Message, id string)
return fantasy.ToolResultPart{}
}
func requireNoProviderExecutedToolResultPrompt(t *testing.T, prompt []fantasy.Message) {
t.Helper()
for i, message := range prompt {
for j, part := range message.Content {
toolResult, ok := safeToolResultPart(part)
if ok && toolResult.ProviderExecuted {
t.Fatalf("prompt[%d].content[%d]: unexpected provider-executed result", i, j)
}
}
}
}
func requireProviderExecutedToolCallPrompt(
t *testing.T,
prompt []fantasy.Message,
id string,
) fantasy.ToolCallPart {
t.Helper()
for _, message := range prompt {
for _, part := range message.Content {
toolCall, ok := safeToolCallPart(part)
if ok && toolCall.ProviderExecuted && toolCall.ToolCallID == id {
return toolCall
}
}
}
t.Fatalf("missing provider-executed prompt tool call %q", id)
return fantasy.ToolCallPart{}
}
func requireProviderExecutedToolResultPrompt(
t *testing.T,
prompt []fantasy.Message,
id string,
) fantasy.ToolResultPart {
t.Helper()
for _, message := range prompt {
for _, part := range message.Content {
toolResult, ok := safeToolResultPart(part)
if ok && toolResult.ProviderExecuted && toolResult.ToolCallID == id {
return toolResult
}
}
}
t.Fatalf("missing provider-executed prompt tool result %q", id)
return fantasy.ToolResultPart{}
}
func requireAnthropicProviderToolPromptSafe(t *testing.T, prompt []fantasy.Message) {
t.Helper()
require.Empty(t, chatsanitize.ValidateAnthropicProviderToolHistory(prompt))
}
func requireLogField(t *testing.T, entry slog.SinkEntry, name string) any {
t.Helper()
for _, field := range entry.Fields {
if field.Name == name {
return field.Value
}
}
t.Fatalf("missing log field %q", name)
return nil
}
func containsPromptSentinel(prompt []fantasy.Message) bool {
for _, message := range prompt {
if message.Role != fantasy.MessageRoleUser || len(message.Content) != 1 {
@@ -2027,6 +2200,7 @@ func TestRun_AnthropicKeepsPairedWebSearchBeforePersist(t *testing.T) {
ID: "ws-1",
ToolCallName: "web_search",
ProviderExecuted: true,
ProviderMetadata: validWebSearchProviderMetadataForTest(),
},
{Type: fantasy.StreamPartTypeTextStart, ID: "text-1"},
{Type: fantasy.StreamPartTypeTextDelta, ID: "text-1", Delta: "search done"},
@@ -2322,6 +2496,535 @@ func TestRun_AnthropicSanitizesWebSearchBeforeContinuation(t *testing.T) {
require.False(t, promptResult.ProviderExecuted)
}
func TestSanitizeAnthropicProviderToolContent(t *testing.T) {
t.Parallel()
providerCall := func(id, name, input string) fantasy.ToolCallContent {
return fantasy.ToolCallContent{
ToolCallID: id,
ToolName: name,
Input: input,
ProviderExecuted: true,
}
}
providerResult := func(id, name string) fantasy.ToolResultContent {
return fantasy.ToolResultContent{
ToolCallID: id,
ToolName: name,
ProviderExecuted: true,
ProviderMetadata: validWebSearchProviderMetadataForTest(),
Result: fantasy.ToolResultOutputContentText{Text: "ok"},
}
}
localCall := func(id, name string) fantasy.ToolCallContent {
return fantasy.ToolCallContent{
ToolCallID: id,
ToolName: name,
Input: `{}`,
}
}
localResult := func(id, name string) fantasy.ToolResultContent {
return fantasy.ToolResultContent{
ToolCallID: id,
ToolName: name,
Result: fantasy.ToolResultOutputContentText{Text: "ok"},
}
}
type contentSummary struct {
providerCalls []string
providerResults []string
localCalls []string
localResults []string
}
summarizeContent := func(content []fantasy.Content) contentSummary {
var summary contentSummary
for _, block := range content {
if toolCall, ok := safeToolCallContent(block); ok {
if toolCall.ProviderExecuted {
summary.providerCalls = append(summary.providerCalls, toolCall.ToolCallID)
} else {
summary.localCalls = append(summary.localCalls, toolCall.ToolCallID)
}
continue
}
if toolResult, ok := safeToolResultContent(block); ok {
if toolResult.ProviderExecuted {
summary.providerResults = append(summary.providerResults, toolResult.ToolCallID)
} else {
summary.localResults = append(summary.localResults, toolResult.ToolCallID)
}
}
}
return summary
}
assertProviderHistoryValid := func(t *testing.T, content []fantasy.Content) {
t.Helper()
parts := make([]fantasy.MessagePart, 0)
for _, block := range content {
if toolCall, ok := safeToolCallContent(block); ok && toolCall.ProviderExecuted {
parts = append(parts, toolCallContentToPart(toolCall))
continue
}
if toolResult, ok := safeToolResultContent(block); ok && toolResult.ProviderExecuted {
parts = append(parts, toolResultContentToPart(toolResult))
}
}
if len(parts) == 0 {
return
}
require.Empty(t, chatsanitize.ValidateAnthropicProviderToolHistory([]fantasy.Message{
{
Role: fantasy.MessageRoleAssistant,
Content: parts,
},
}))
}
metadataCall := providerCall("ws-meta", "web_search", `{"query":"coder"}`)
metadataCall.ProviderMetadata = fantasy.ProviderMetadata{fantasyanthropic.Name: nil}
metadataResult := providerResult("ws-meta", "web_search")
metadataResult.ProviderMetadata = fantasy.ProviderMetadata{fantasyanthropic.Name: nil}
pointerCall := providerCall("ws-pointer", "web_search", `{"query":"coder"}`)
var nilToolCall *fantasy.ToolCallContent
testCases := []struct {
name string
provider string
content []fantasy.Content
wantSummary contentSummary
wantRemovedCalls int
wantRemovedResults int
wantTexts []string
validateAnthropic bool
}{
{
name: "orphan provider result textified",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
fantasy.TextContent{Text: "keep"},
providerResult("ws-1", "web_search"),
},
wantRemovedResults: 1,
wantTexts: []string{"keep", "ok"},
validateAnthropic: true,
},
{
name: "result before call removes both provider blocks",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerResult("ws-1", "web_search"),
providerCall("ws-1", "web_search", `{"query":"coder"}`),
},
wantRemovedCalls: 1,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "valid web search pair preserved",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerCall("ws-1", "web_search", `{"query":"coder"}`),
providerResult("ws-1", "web_search"),
fantasy.TextContent{Text: "search done"},
},
wantSummary: contentSummary{
providerCalls: []string{"ws-1"},
providerResults: []string{"ws-1"},
},
wantTexts: []string{"search done"},
validateAnthropic: true,
},
{
name: "invalid JSON provider call drops pair",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerCall("ws-1", "web_search", `{`),
providerResult("ws-1", "web_search"),
},
wantRemovedCalls: 1,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "empty ID provider call drops pair",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerCall("", "web_search", `{"query":"coder"}`),
providerResult("", "web_search"),
},
wantRemovedCalls: 1,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "empty tool name provider call drops pair",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerCall("ws-empty", "", `{"query":"coder"}`),
providerResult("ws-empty", ""),
},
wantRemovedCalls: 1,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "non web search provider pair drops through serializable helper",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerCall("code-1", "code_execution", `{"code":"print(1)"}`),
providerResult("code-1", "code_execution"),
},
wantRemovedCalls: 1,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "mismatched provider result tool name drops pair",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerCall("ws-mismatch", "web_search", `{"query":"coder"}`),
providerResult("ws-mismatch", "code_execution"),
},
wantRemovedCalls: 1,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "duplicate provider IDs drop all provider content for ID",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerCall("dup-1", "web_search", `{"query":"coder"}`),
providerResult("dup-1", "web_search"),
providerCall("dup-1", "web_search", `{"query":"coder"}`),
},
wantRemovedCalls: 2,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "mismatched provider flags remove only provider side",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
providerCall("mix-1", "web_search", `{"query":"coder"}`),
localResult("mix-1", "web_search"),
localCall("mix-2", "read_file"),
providerResult("mix-2", "web_search"),
},
wantSummary: contentSummary{
localCalls: []string{"mix-2"},
localResults: []string{"mix-1"},
},
wantRemovedCalls: 1,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "malformed provider metadata textifies result",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
metadataCall,
metadataResult,
},
wantRemovedCalls: 1,
wantRemovedResults: 1,
wantTexts: []string{"ok"},
validateAnthropic: true,
},
{
name: "pointer and nil pointer variants are handled safely",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
nilToolCall,
&pointerCall,
providerResult("ws-pointer", "web_search"),
},
wantSummary: contentSummary{
providerCalls: []string{"ws-pointer"},
providerResults: []string{"ws-pointer"},
},
validateAnthropic: true,
},
{
name: "local tool content is unchanged",
provider: fantasyanthropic.Name,
content: []fantasy.Content{
localCall("tc-1", "read_file"),
localResult("tc-1", "read_file"),
},
wantSummary: contentSummary{
localCalls: []string{"tc-1"},
localResults: []string{"tc-1"},
},
validateAnthropic: true,
},
{
name: "non Anthropic provider content is unchanged",
provider: "fake",
content: []fantasy.Content{
providerCall("ws-1", "web_search", `{"query":"coder"}`),
},
wantSummary: contentSummary{
providerCalls: []string{"ws-1"},
},
},
}
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
sanitized, stats := chatsanitize.SanitizeAnthropicProviderToolContent(tc.provider, tc.content)
require.Equal(t, tc.wantRemovedCalls, stats.RemovedToolCalls)
require.Equal(t, tc.wantRemovedResults, stats.RemovedToolResults)
require.Zero(t, stats.DroppedMessages)
summary := summarizeContent(sanitized)
assert.ElementsMatch(t, tc.wantSummary.providerCalls, summary.providerCalls)
assert.ElementsMatch(t, tc.wantSummary.providerResults, summary.providerResults)
assert.ElementsMatch(t, tc.wantSummary.localCalls, summary.localCalls)
assert.ElementsMatch(t, tc.wantSummary.localResults, summary.localResults)
for _, text := range tc.wantTexts {
requireTextContent(t, sanitized, text)
}
if tc.validateAnthropic {
assertProviderHistoryValid(t, sanitized)
}
})
}
}
func TestRun_AnthropicProviderToolPreRequestGuard(t *testing.T) {
t.Parallel()
webSearchTool := ProviderTool{
Definition: fantasy.ProviderDefinedTool{
ID: "anthropic.web_search",
Name: "web_search",
},
}
providerPair := func(id string) []fantasy.MessagePart {
return []fantasy.MessagePart{
fantasy.ToolCallPart{
ToolCallID: id,
ToolName: "web_search",
Input: `{"query":"coder"}`,
ProviderExecuted: true,
},
fantasy.ToolResultPart{
ToolCallID: id,
Output: fantasy.ToolResultOutputContentText{Text: "ok"},
ProviderExecuted: true,
ProviderOptions: fantasy.ProviderOptions(validWebSearchProviderMetadataForTest()),
},
}
}
completionModel := func(capturedPrompt *[]fantasy.Message) *chattest.FakeModel {
return &chattest.FakeModel{
ProviderName: fantasyanthropic.Name,
ModelName: "claude-test",
StreamFn: func(_ context.Context, call fantasy.Call) (fantasy.StreamResponse, error) {
*capturedPrompt = append([]fantasy.Message(nil), call.Prompt...)
return streamFromParts([]fantasy.StreamPart{
{Type: fantasy.StreamPartTypeTextStart, ID: "text-1"},
{Type: fantasy.StreamPartTypeTextDelta, ID: "text-1", Delta: "done"},
{Type: fantasy.StreamPartTypeTextEnd, ID: "text-1"},
{Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonStop},
}), nil
},
}
}
t.Run("allowed web search survives when provider tool is enabled", func(t *testing.T) {
t.Parallel()
var capturedPrompt []fantasy.Message
err := Run(context.Background(), RunOptions{
Model: completionModel(&capturedPrompt),
Messages: []fantasy.Message{
textMessage(fantasy.MessageRoleUser, "search"),
{
Role: fantasy.MessageRoleAssistant,
Content: providerPair("ws-allowed"),
},
textMessage(fantasy.MessageRoleUser, "continue"),
},
ProviderTools: []ProviderTool{webSearchTool},
MaxSteps: 1,
PersistStep: func(_ context.Context, _ PersistedStep) error {
return nil
},
})
require.NoError(t, err)
toolCall := requireProviderExecutedToolCallPrompt(t, capturedPrompt, "ws-allowed")
require.Equal(t, "web_search", toolCall.ToolName)
requireProviderExecutedToolResultPrompt(t, capturedPrompt, "ws-allowed")
requireAnthropicProviderToolPromptSafe(t, capturedPrompt)
})
t.Run("web search history survives when provider tool is disabled", func(t *testing.T) {
t.Parallel()
var capturedPrompt []fantasy.Message
err := Run(context.Background(), RunOptions{
Model: completionModel(&capturedPrompt),
Messages: []fantasy.Message{
textMessage(fantasy.MessageRoleUser, "search and read"),
{
Role: fantasy.MessageRoleAssistant,
Content: append(providerPair("ws-disabled"), fantasy.ToolCallPart{
ToolCallID: "tc-1",
ToolName: "read_file",
Input: `{"path":"main.go"}`,
}),
},
{
Role: fantasy.MessageRoleTool,
Content: []fantasy.MessagePart{
fantasy.ToolResultPart{
ToolCallID: "tc-1",
Output: fantasy.ToolResultOutputContentText{Text: "file"},
},
},
},
textMessage(fantasy.MessageRoleUser, "continue"),
},
MaxSteps: 1,
PersistStep: func(_ context.Context, _ PersistedStep) error {
return nil
},
})
require.NoError(t, err)
requireProviderExecutedToolCallPrompt(t, capturedPrompt, "ws-disabled")
requireProviderExecutedToolResultPrompt(t, capturedPrompt, "ws-disabled")
promptResult := requireToolResultPrompt(t, capturedPrompt, "tc-1")
require.False(t, promptResult.ProviderExecuted)
requireAnthropicProviderToolPromptSafe(t, capturedPrompt)
})
t.Run("direct guard textifies orphaned provider result", func(t *testing.T) {
t.Parallel()
guarded := chatsanitize.ApplyAnthropicProviderToolGuard(
context.Background(),
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
fantasyanthropic.Name,
"claude-test",
[]fantasy.Message{
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: "keep"},
fantasy.ToolResultPart{
ToolCallID: "ws-orphan",
Output: fantasy.ToolResultOutputContentText{Text: "search result"},
ProviderExecuted: true,
},
},
},
},
)
requireNoProviderExecutedToolResultPrompt(t, guarded)
requireAnthropicProviderToolPromptSafe(t, guarded)
require.Len(t, guarded, 1)
require.Len(t, guarded[0].Content, 2)
textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](guarded[0].Content[0])
require.True(t, ok)
require.Equal(t, "keep", textPart.Text)
textPart, ok = fantasy.AsMessagePart[fantasy.TextPart](guarded[0].Content[1])
require.True(t, ok)
require.Equal(t, "search result", textPart.Text)
})
t.Run("direct guard leaves valid provider history unchanged", func(t *testing.T) {
t.Parallel()
content := []fantasy.MessagePart{fantasy.TextPart{Text: "keep"}}
content = append(content, providerPair("ws-one")...)
content = append(content, providerPair("ws-two")...)
guarded := chatsanitize.ApplyAnthropicProviderToolGuard(
context.Background(),
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
fantasyanthropic.Name,
"claude-test",
[]fantasy.Message{{Role: fantasy.MessageRoleAssistant, Content: content}},
)
requireAnthropicProviderToolPromptSafe(t, guarded)
require.Len(t, guarded, 1)
require.Len(t, guarded[0].Content, len(content))
requireProviderExecutedToolCallPrompt(t, guarded, "ws-one")
requireProviderExecutedToolResultPrompt(t, guarded, "ws-one")
requireProviderExecutedToolCallPrompt(t, guarded, "ws-two")
requireProviderExecutedToolResultPrompt(t, guarded, "ws-two")
})
t.Run("direct guard leaves non Anthropic providers unchanged", func(t *testing.T) {
t.Parallel()
prompt := []fantasy.Message{
{
Role: fantasy.MessageRoleAssistant,
Content: providerPair("ws-other-provider"),
},
}
guarded := chatsanitize.ApplyAnthropicProviderToolGuard(
context.Background(),
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
"fake",
"fake-model",
prompt,
)
require.Equal(t, prompt, guarded)
})
t.Run("guard logs removals", func(t *testing.T) {
t.Parallel()
logSink := testutil.NewFakeSink(t)
logger := logSink.Logger()
logPair := providerPair("ws-log")
guarded := chatsanitize.ApplyAnthropicProviderToolGuard(
context.Background(),
logger,
fantasyanthropic.Name,
"claude-test",
[]fantasy.Message{
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
logPair[1],
logPair[0],
},
},
},
)
requireNoProviderExecutedToolCallPrompt(t, guarded)
requireNoProviderExecutedToolResultPrompt(t, guarded)
requireTextPrompt(t, guarded, "ok")
entries := logSink.Entries(func(e slog.SinkEntry) bool {
return e.Level == slog.LevelWarn &&
e.Message == "removed provider-executed tool history"
})
require.Len(t, entries, 1)
require.Equal(t, "pre_request_guard", requireLogField(t, entries[0], "phase"))
require.Equal(t, 1, requireLogField(t, entries[0], "removed_tool_calls"))
require.Equal(t, 1, requireLogField(t, entries[0], "removed_tool_results"))
})
}
// TestRun_PersistStepInterruptedFallback verifies that when the normal
// PersistStep call returns ErrInterrupted (e.g., context canceled in a
// race), the step is retried via the interrupt-safe path.
+22 -187
View File
@@ -11,7 +11,6 @@ import (
"strings"
"charm.land/fantasy"
fantasyanthropic "charm.land/fantasy/providers/anthropic"
"github.com/google/uuid"
"github.com/sqlc-dev/pqtype"
"golang.org/x/xerrors"
@@ -35,192 +34,28 @@ var toolCallIDSanitizer = regexp.MustCompile(`[^a-zA-Z0-9_-]`)
var syntheticPasteFileNamePattern = regexp.MustCompile(`^pasted-text-\d{4}-\d{2}-\d{2}-\d{2}-\d{2}-\d{2}\.txt$`)
// AnthropicProviderToolSanitizationStats describes prompt changes made
// while removing unpaired Anthropic provider-executed tool calls.
type AnthropicProviderToolSanitizationStats struct {
RemovedToolCalls int
DroppedMessages int
func safeAsToolCallPart(part fantasy.MessagePart) (fantasy.ToolCallPart, bool) {
var zero fantasy.ToolCallPart
if part == nil {
return zero, false
}
if value, ok := part.(*fantasy.ToolCallPart); ok && value == nil {
return zero, false
}
type toolCallPart = fantasy.ToolCallPart
return fantasy.AsMessagePart[toolCallPart](part)
}
// LogAnthropicProviderToolSanitization logs prompt changes made while removing
// unpaired Anthropic provider-executed tool calls.
func LogAnthropicProviderToolSanitization(
ctx context.Context,
logger slog.Logger,
phase string,
provider string,
modelName string,
stats AnthropicProviderToolSanitizationStats,
extra ...slog.Field,
) {
if stats.RemovedToolCalls == 0 {
return
func safeAsToolResultPart(part fantasy.MessagePart) (fantasy.ToolResultPart, bool) {
var zero fantasy.ToolResultPart
if part == nil {
return zero, false
}
fields := []slog.Field{
slog.F("phase", phase),
slog.F("tool_type", "provider_executed"),
slog.F("provider", provider),
slog.F("model", modelName),
slog.F("removed_tool_calls", stats.RemovedToolCalls),
slog.F("dropped_messages", stats.DroppedMessages),
if value, ok := part.(*fantasy.ToolResultPart); ok && value == nil {
return zero, false
}
fields = append(fields, extra...)
logger.Warn(ctx, "removed unpaired provider-executed tool calls", fields...)
}
// SanitizeAnthropicProviderToolCalls removes Anthropic provider-executed
// calls that do not have a same-message provider result.
func SanitizeAnthropicProviderToolCalls(
provider string,
messages []fantasy.Message,
) ([]fantasy.Message, AnthropicProviderToolSanitizationStats) {
var stats AnthropicProviderToolSanitizationStats
if provider != fantasyanthropic.Name || len(messages) == 0 {
return messages, stats
}
out := make([]fantasy.Message, 0, len(messages))
changed := false
for _, msg := range messages {
if msg.Role != fantasy.MessageRoleAssistant {
out = appendSanitizedMessage(out, msg)
continue
}
matchedResultIDs := make(map[string]struct{})
for _, part := range msg.Content {
result, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](part)
if !ok || !result.ProviderExecuted || result.ToolCallID == "" {
continue
}
matchedResultIDs[result.ToolCallID] = struct{}{}
}
parts := make([]fantasy.MessagePart, 0, len(msg.Content))
removedFromMessage := 0
for _, part := range msg.Content {
toolCall, ok := fantasy.AsMessagePart[fantasy.ToolCallPart](part)
if ok && toolCall.ProviderExecuted {
if _, hasResult := matchedResultIDs[toolCall.ToolCallID]; !hasResult {
stats.RemovedToolCalls++
removedFromMessage++
changed = true
continue
}
}
parts = append(parts, part)
}
if removedFromMessage > 0 {
if len(parts) == 0 {
stats.DroppedMessages++
continue
}
msg.Content = parts
}
out = appendSanitizedMessage(out, msg)
}
if !changed {
return messages, stats
}
return out, stats
}
func appendSanitizedMessage(out []fantasy.Message, msg fantasy.Message) []fantasy.Message {
if len(out) == 0 || out[len(out)-1].Role != msg.Role {
return append(out, msg)
}
last := &out[len(out)-1]
lastContent := applyMessageProviderOptionsToLastPart(last.Content, last.ProviderOptions)
msgContent := applyMessageProviderOptionsToLastPart(msg.Content, msg.ProviderOptions)
content := make([]fantasy.MessagePart, 0, len(lastContent)+len(msgContent))
content = append(content, lastContent...)
content = append(content, msgContent...)
last.Content = content
last.ProviderOptions = nil
return out
}
func applyMessageProviderOptionsToLastPart(
parts []fantasy.MessagePart,
options fantasy.ProviderOptions,
) []fantasy.MessagePart {
if len(options) == 0 || len(parts) == 0 {
return parts
}
out := make([]fantasy.MessagePart, len(parts))
copy(out, parts)
lastIndex := len(out) - 1
switch part := out[lastIndex].(type) {
case fantasy.TextPart:
part.ProviderOptions = mergeProviderOptions(part.ProviderOptions, options)
out[lastIndex] = part
case *fantasy.TextPart:
if part != nil {
clone := *part
clone.ProviderOptions = mergeProviderOptions(clone.ProviderOptions, options)
out[lastIndex] = &clone
}
case fantasy.ReasoningPart:
part.ProviderOptions = mergeProviderOptions(part.ProviderOptions, options)
out[lastIndex] = part
case *fantasy.ReasoningPart:
if part != nil {
clone := *part
clone.ProviderOptions = mergeProviderOptions(clone.ProviderOptions, options)
out[lastIndex] = &clone
}
case fantasy.FilePart:
part.ProviderOptions = mergeProviderOptions(part.ProviderOptions, options)
out[lastIndex] = part
case *fantasy.FilePart:
if part != nil {
clone := *part
clone.ProviderOptions = mergeProviderOptions(clone.ProviderOptions, options)
out[lastIndex] = &clone
}
case fantasy.ToolCallPart:
part.ProviderOptions = mergeProviderOptions(part.ProviderOptions, options)
out[lastIndex] = part
case *fantasy.ToolCallPart:
if part != nil {
clone := *part
clone.ProviderOptions = mergeProviderOptions(clone.ProviderOptions, options)
out[lastIndex] = &clone
}
case fantasy.ToolResultPart:
part.ProviderOptions = mergeProviderOptions(part.ProviderOptions, options)
out[lastIndex] = part
case *fantasy.ToolResultPart:
if part != nil {
clone := *part
clone.ProviderOptions = mergeProviderOptions(clone.ProviderOptions, options)
out[lastIndex] = &clone
}
}
return out
}
func mergeProviderOptions(first, second fantasy.ProviderOptions) fantasy.ProviderOptions {
if len(first) == 0 {
return second
}
if len(second) == 0 {
return first
}
merged := make(fantasy.ProviderOptions, len(first)+len(second))
for provider, options := range first {
merged[provider] = options
}
for provider, options := range second {
if options != nil {
merged[provider] = options
}
}
return merged
type toolResultPart = fantasy.ToolResultPart
return fantasy.AsMessagePart[toolResultPart](part)
}
// FileData holds resolved file content for LLM prompt building.
@@ -764,7 +599,7 @@ func normalizeAssistantToolCallInputs(
) []fantasy.MessagePart {
normalized := make([]fantasy.MessagePart, 0, len(parts))
for _, part := range parts {
toolCall, ok := fantasy.AsMessagePart[fantasy.ToolCallPart](part)
toolCall, ok := safeAsToolCallPart(part)
if !ok {
normalized = append(normalized, part)
continue
@@ -797,7 +632,7 @@ func normalizeToolCallInput(input string) string {
func ExtractToolCalls(parts []fantasy.MessagePart) []fantasy.ToolCallContent {
toolCalls := make([]fantasy.ToolCallContent, 0, len(parts))
for _, part := range parts {
toolCall, ok := fantasy.AsMessagePart[fantasy.ToolCallPart](part)
toolCall, ok := safeAsToolCallPart(part)
if !ok {
continue
}
@@ -1148,7 +983,7 @@ func injectMissingToolResults(prompt []fantasy.Message) []fantasy.Message {
break
}
for _, part := range prompt[j].Content {
tr, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](part)
tr, ok := safeAsToolResultPart(part)
if !ok {
continue
}
@@ -1205,7 +1040,7 @@ func injectMissingToolUses(
allToolResults := make([]fantasy.ToolResultPart, 0, len(msg.Content))
for _, part := range msg.Content {
toolResult, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](part)
toolResult, ok := safeAsToolResultPart(part)
if !ok {
continue
}
+32 -256
View File
@@ -22,6 +22,7 @@ import (
"github.com/coder/coder/v2/coderd/database/dbgen"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
"github.com/coder/coder/v2/coderd/x/chatd/chatsanitize"
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
@@ -59,243 +60,16 @@ func convertMessagesWithoutFiles(t *testing.T, messages []database.ChatMessage)
return prompt
}
func TestSanitizeAnthropicProviderToolCalls(t *testing.T) {
t.Parallel()
type testToolCallPart = fantasy.ToolCallPart
textPart := fantasy.TextPart{Text: "Here is a summary."}
webSearchCall := fantasy.ToolCallPart{
ToolCallID: "srvtoolu_search",
ToolName: "web_search",
Input: `{"query":"coder"}`,
ProviderExecuted: true,
}
matchedResult := fantasy.ToolResultPart{
ToolCallID: "srvtoolu_search",
Output: fantasy.ToolResultOutputContentText{Text: `{"ok":true}`},
ProviderExecuted: true,
}
codeExecutionCall := fantasy.ToolCallPart{
ToolCallID: "srvtoolu_code",
ToolName: "code_execution",
Input: `{"code":"print(1)"}`,
ProviderExecuted: true,
}
localCall := fantasy.ToolCallPart{
ToolCallID: "toolu_local",
ToolName: "read_file",
Input: `{"path":"main.go"}`,
}
unpairedWebSearchCall := webSearchCall
unpairedWebSearchCall.ToolCallID = "srvtoolu_unpaired"
disableParallelToolUse := true
providerOptions := fantasy.ProviderOptions{
fantasyanthropic.Name: &fantasyanthropic.ProviderOptions{
DisableParallelToolUse: &disableParallelToolUse,
},
}
enableParallelToolUse := false
providerOptionsAllowParallel := fantasy.ProviderOptions{
fantasyanthropic.Name: &fantasyanthropic.ProviderOptions{
DisableParallelToolUse: &enableParallelToolUse,
},
}
type testToolResultPart = fantasy.ToolResultPart
testCases := []struct {
name string
provider string
messages []fantasy.Message
want []fantasy.Message
wantRemoved int
wantDropped int
}{
{
name: "removes unpaired call and keeps text",
provider: fantasyanthropic.Name,
messages: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
textPart,
webSearchCall,
},
}},
want: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{textPart},
}},
wantRemoved: 1,
},
{
name: "drops assistant message when only part is removed",
provider: fantasyanthropic.Name,
messages: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{webSearchCall},
}},
want: []fantasy.Message{},
wantRemoved: 1,
wantDropped: 1,
},
{
name: "coalesces adjacent roles after dropping empty message",
provider: fantasyanthropic.Name,
messages: []fantasy.Message{
{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: "search for coder"},
},
},
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{webSearchCall},
},
{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: "now summarize"},
},
ProviderOptions: providerOptions,
},
},
want: []fantasy.Message{{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: "search for coder"},
fantasy.TextPart{
Text: "now summarize",
ProviderOptions: providerOptions,
},
},
}},
wantRemoved: 1,
wantDropped: 1,
},
{
name: "coalesces adjacent provider options without flattening boundaries",
provider: fantasyanthropic.Name,
messages: []fantasy.Message{
{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: "search for coder"},
},
ProviderOptions: providerOptionsAllowParallel,
},
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{webSearchCall},
},
{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: "now summarize"},
},
ProviderOptions: providerOptions,
},
},
want: []fantasy.Message{{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{
fantasy.TextPart{
Text: "search for coder",
ProviderOptions: providerOptionsAllowParallel,
},
fantasy.TextPart{
Text: "now summarize",
ProviderOptions: providerOptions,
},
},
}},
wantRemoved: 1,
wantDropped: 1,
},
{
name: "keeps matched call and result",
provider: fantasyanthropic.Name,
messages: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
webSearchCall,
matchedResult,
},
}},
want: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
webSearchCall,
matchedResult,
},
}},
},
{
name: "removes only unpaired call from mixed message",
provider: fantasyanthropic.Name,
messages: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
textPart,
webSearchCall,
matchedResult,
unpairedWebSearchCall,
},
}},
want: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
textPart,
webSearchCall,
matchedResult,
},
}},
wantRemoved: 1,
},
{
name: "removes unpaired provider call and keeps local call",
provider: fantasyanthropic.Name,
messages: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
textPart,
codeExecutionCall,
localCall,
},
}},
want: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
textPart,
localCall,
},
}},
wantRemoved: 1,
},
{
name: "leaves other providers unchanged",
provider: "fake",
messages: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{webSearchCall},
}},
want: []fantasy.Message{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{webSearchCall},
}},
},
}
func asToolCallPartForTest(part fantasy.MessagePart) (fantasy.ToolCallPart, bool) {
return fantasy.AsMessagePart[testToolCallPart](part)
}
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
sanitized, stats := chatprompt.SanitizeAnthropicProviderToolCalls(
tc.provider,
tc.messages,
)
require.Equal(t, tc.wantRemoved, stats.RemovedToolCalls)
require.Equal(t, tc.wantDropped, stats.DroppedMessages)
require.Equal(t, tc.want, sanitized)
})
}
func asToolResultPartForTest(part fantasy.MessagePart) (fantasy.ToolResultPart, bool) {
return fantasy.AsMessagePart[testToolResultPart](part)
}
func TestConvertMessagesWithFiles_NormalizesAssistantToolCallInput(t *testing.T) {
@@ -742,18 +516,20 @@ func TestInjectMissingToolResults_SkipsProviderExecuted(t *testing.T) {
// The tool message should have exactly one result (the local one).
var resultIDs []string
for _, part := range prompt[1].Content {
tr, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](part)
tr, ok := asToolResultPartForTest(part)
if ok {
resultIDs = append(resultIDs, tr.ToolCallID)
}
}
require.Equal(t, []string{"toolu_local"}, resultIDs)
sanitized, sanitizeStats := chatprompt.SanitizeAnthropicProviderToolCalls(
sanitized, sanitizeStats := chatsanitize.SanitizeAnthropicProviderToolHistory(
fantasyanthropic.Name,
prompt,
)
require.Equal(t, 1, sanitizeStats.RemovedToolCalls)
require.Equal(t, 0, sanitizeStats.RemovedToolResults)
require.Len(t, sanitized, 2)
require.Empty(t, chatsanitize.ValidateAnthropicProviderToolHistory(sanitized))
remainingToolCalls := chatprompt.ExtractToolCalls(sanitized[0].Content)
require.Len(t, remainingToolCalls, 1)
require.Equal(t, "toolu_local", remainingToolCalls[0].ToolCallID)
@@ -799,7 +575,7 @@ func TestInjectMissingToolResults_SkipsProviderExecutedAndInjectsLocal(t *testin
require.Equal(t, []string{"toolu_read"}, extractToolResultIDs(t, prompt[1]))
require.Len(t, prompt[1].Content, 1)
toolResult, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
toolResult, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok, "expected synthetic ToolResultPart")
require.Equal(t, "toolu_read", toolResult.ToolCallID)
require.False(t, toolResult.ProviderExecuted)
@@ -992,7 +768,7 @@ func TestInjectMissingToolUses_DropsProviderExecutedOrphans(t *testing.T) {
for i, msg := range prompt {
if msg.Role == fantasy.MessageRoleAssistant {
for _, part := range msg.Content {
tc, ok := fantasy.AsMessagePart[fantasy.ToolCallPart](part)
tc, ok := asToolCallPartForTest(part)
if ok && tc.Input == "{}" && tc.ToolCallID == "srvtoolu_C" {
t.Errorf("message[%d]: unexpected synthetic tool_use for srvtoolu_C", i)
}
@@ -1092,12 +868,12 @@ func TestProviderExecutedResultInAssistantContent(t *testing.T) {
// The assistant message must contain 3 parts: tool_call, tool_result, text.
var foundToolCall, foundToolResult, foundText bool
for _, part := range prompt[0].Content {
if tc, ok := fantasy.AsMessagePart[fantasy.ToolCallPart](part); ok {
if tc, ok := asToolCallPartForTest(part); ok {
require.Equal(t, "srvtoolu_WS", tc.ToolCallID)
require.True(t, tc.ProviderExecuted, "ToolCallPart.ProviderExecuted must be true")
foundToolCall = true
}
if tr, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](part); ok {
if tr, ok := asToolResultPartForTest(part); ok {
require.Equal(t, "srvtoolu_WS", tr.ToolCallID)
require.True(t, tr.ProviderExecuted, "ToolResultPart.ProviderExecuted must be true")
foundToolResult = true
@@ -1844,7 +1620,7 @@ func TestMixedFormatConversation(t *testing.T) {
// 4. Old tool: result paired with call_1.
require.Equal(t, fantasy.MessageRoleTool, prompt[3].Role)
require.Len(t, prompt[3].Content, 1)
toolResult, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[3].Content[0])
toolResult, ok := asToolResultPartForTest(prompt[3].Content[0])
require.True(t, ok)
assert.Equal(t, "call_1", toolResult.ToolCallID)
@@ -2026,7 +1802,7 @@ func extractToolResultIDs(t *testing.T, msgs ...fantasy.Message) []string {
var ids []string
for _, msg := range msgs {
for _, part := range msg.Content {
tr, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](part)
tr, ok := asToolResultPartForTest(part)
if ok {
ids = append(ids, tr.ToolCallID)
}
@@ -2304,11 +2080,11 @@ func TestConvertMessagesWithFiles_FiltersEmptyTextAndReasoningParts(t *testing.T
require.True(t, ok, "expected ReasoningPart at index 2")
require.Equal(t, "thinking deeply", reasoningPart.Text)
toolCallPart, ok := fantasy.AsMessagePart[fantasy.ToolCallPart](resultParts[3])
toolCallPart, ok := asToolCallPartForTest(resultParts[3])
require.True(t, ok, "expected ToolCallPart at index 3")
require.Equal(t, "call-1", toolCallPart.ToolCallID)
toolResultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](resultParts[4])
toolResultPart, ok := asToolResultPartForTest(resultParts[4])
require.True(t, ok, "expected ToolResultPart at index 4")
require.Equal(t, "call-1", toolResultPart.ToolCallID)
})
@@ -2342,7 +2118,7 @@ func TestConvertMessagesWithFiles_FiltersEmptyTextAndReasoningParts(t *testing.T
require.True(t, ok, "expected TextPart")
require.Equal(t, " reply ", textPart.Text)
tcPart, ok := fantasy.AsMessagePart[fantasy.ToolCallPart](resultParts[1])
tcPart, ok := asToolCallPartForTest(resultParts[1])
require.True(t, ok, "expected ToolCallPart")
require.Equal(t, "tc-1", tcPart.ToolCallID)
})
@@ -2739,7 +2515,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
require.Equal(t, fantasy.MessageRoleTool, toolMsg.Role)
require.Len(t, toolMsg.Content, 1)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](toolMsg.Content[0])
resultPart, ok := asToolResultPartForTest(toolMsg.Content[0])
require.True(t, ok, "expected ToolResultPart")
require.Equal(t, callID, resultPart.ToolCallID)
require.False(t, resultPart.ProviderExecuted)
@@ -2796,7 +2572,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok, "expected ToolResultPart")
mediaOutput, ok := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
@@ -2869,7 +2645,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
require.False(t, resultPart.ProviderExecuted)
@@ -2895,7 +2671,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
_, isMedia := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
@@ -2921,7 +2697,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
_, isMedia := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
@@ -2943,7 +2719,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
_, isMedia := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
@@ -2971,7 +2747,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
errOutput, isError := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentError](resultPart.Output)
@@ -3003,7 +2779,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
_, isMedia := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
@@ -3032,7 +2808,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
_, isMedia := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
@@ -3060,7 +2836,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
_, isMedia := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
@@ -3091,7 +2867,7 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
prompt := loadPrompt(t, chat)
require.Len(t, prompt, 2)
resultPart, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
resultPart, ok := asToolResultPartForTest(prompt[1].Content[0])
require.True(t, ok)
_, isMedia := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentMedia](resultPart.Output)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,146 @@
package chatsanitize
import (
"testing"
"charm.land/fantasy"
fantasyanthropic "charm.land/fantasy/providers/anthropic"
"github.com/stretchr/testify/require"
)
func textMessageForTest(role fantasy.MessageRole, text string) fantasy.Message {
return fantasy.Message{
Role: role,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: text},
},
}
}
func TestProviderExecutedToolMessageIndexes(t *testing.T) {
t.Parallel()
messages := []fantasy.Message{
textMessageForTest(fantasy.MessageRoleUser, "plain"),
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
fantasy.ToolResultPart{
ToolCallID: "ws-result-only",
ProviderExecuted: true,
},
},
},
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
fantasy.ToolCallPart{
ToolCallID: "ws-call",
ToolName: "web_search",
Input: `{"query":"coder"}`,
ProviderExecuted: true,
},
},
},
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
fantasy.ToolCallPart{
ToolCallID: "local-call",
ToolName: "read_file",
Input: `{"path":"main.go"}`,
},
},
},
}
require.Equal(t, map[int]struct{}{1: {}, 2: {}}, providerExecutedToolMessageIndexes(messages))
}
func TestAnthropicProviderToolFallbackStripHelpers(t *testing.T) {
t.Parallel()
providerCall := fantasy.ToolCallPart{
ToolCallID: "ws-strip",
ToolName: "web_search",
Input: `{"query":"coder"}`,
ProviderExecuted: true,
}
providerResult := fantasy.ToolResultPart{
ToolCallID: "ws-strip",
Output: fantasy.ToolResultOutputContentText{Text: "ok"},
ProviderExecuted: true,
}
messages := []fantasy.Message{
textMessageForTest(fantasy.MessageRoleAssistant, "first"),
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
providerCall,
providerResult,
},
},
textMessageForTest(fantasy.MessageRoleAssistant, "second"),
{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: "keep"},
fantasy.ToolResultPart{
ToolCallID: "ws-user",
ProviderExecuted: true,
},
},
},
}
stripped, stats := stripAnthropicProviderToolHistoryFromMessages(
messages,
map[int]struct{}{1: {}, 3: {}},
)
require.Equal(t, 1, stats.RemovedToolCalls)
require.Equal(t, 2, stats.RemovedToolResults)
require.Zero(t, stats.DroppedMessages)
sanitized, sanitizeStats := SanitizeAnthropicProviderToolHistory(
fantasyanthropic.Name,
stripped,
)
require.Zero(t, sanitizeStats.RemovedToolCalls)
require.Zero(t, sanitizeStats.RemovedToolResults)
require.Empty(t, ValidateAnthropicProviderToolHistory(sanitized))
require.Len(t, sanitized, 2)
require.Equal(t, fantasy.MessageRoleAssistant, sanitized[0].Role)
require.Len(t, sanitized[0].Content, 3)
firstText, ok := fantasy.AsMessagePart[fantasy.TextPart](sanitized[0].Content[0])
require.True(t, ok)
require.Equal(t, "first", firstText.Text)
stripText, ok := fantasy.AsMessagePart[fantasy.TextPart](sanitized[0].Content[1])
require.True(t, ok)
require.Equal(t, "ok", stripText.Text)
secondText, ok := fantasy.AsMessagePart[fantasy.TextPart](sanitized[0].Content[2])
require.True(t, ok)
require.Equal(t, "second", secondText.Text)
require.Equal(t, fantasy.MessageRoleUser, sanitized[1].Role)
require.Len(t, sanitized[1].Content, 1)
keepText, ok := fantasy.AsMessagePart[fantasy.TextPart](sanitized[1].Content[0])
require.True(t, ok)
require.Equal(t, "keep", keepText.Text)
violations := make([]AnthropicProviderToolHistoryViolation, 33)
for i := range violations {
violations[i] = AnthropicProviderToolHistoryViolation{
MessageIndex: i,
PartIndex: i + 1,
ID: "ws-detail",
Reason: "test_reason",
}
}
details, truncated := anthropicProviderToolViolationLogDetails(violations)
require.True(t, truncated)
require.Len(t, details, maxAnthropicProviderToolViolationLogDetails)
require.Len(t, details[0], 4)
require.Equal(t, 0, details[0]["message_index"])
require.Equal(t, 1, details[0]["part_index"])
require.Equal(t, "ws-detail", details[0]["id"])
require.Equal(t, "test_reason", details[0]["reason"])
}
File diff suppressed because it is too large Load Diff