mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix(coderd): sanitize Anthropic provider tool history (#24706)
Anthropic can reject replayed chat histories when a provider-executed tool call, such as `web_search`, is present without its matching provider result block. This sanitizes unpaired Anthropic provider-executed tool calls during prompt reconstruction, before Anthropic requests, and before persistence so existing poisoned histories can continue and new malformed turns are not stored. Resolves: CODAGT-259 > Mux is acting on Mike's behalf.
This commit is contained in:
@@ -193,7 +193,7 @@ type ProviderTool struct {
|
||||
|
||||
// stepResult holds the accumulated output of a single streaming
|
||||
// step. Since we own the stream consumer, all content is tracked
|
||||
// directly here — no shadow draft state needed.
|
||||
// directly here, no shadow draft state needed.
|
||||
type stepResult struct {
|
||||
content []fantasy.Content
|
||||
usage fantasy.Usage
|
||||
@@ -391,6 +391,12 @@ 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(
|
||||
ctx, opts.Logger, "pre_request", provider, modelName, sanitizeStats,
|
||||
slog.F("step_index", step),
|
||||
slog.F("total_steps", totalSteps),
|
||||
)
|
||||
if applyAnthropicCaching {
|
||||
addAnthropicPromptCaching(prepared)
|
||||
}
|
||||
@@ -518,12 +524,18 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
})
|
||||
}
|
||||
|
||||
contextLimit := extractContextLimit(result.providerMetadata)
|
||||
if !contextLimit.Valid && opts.ContextLimitFallback > 0 {
|
||||
contextLimit = sql.NullInt64{
|
||||
Int64: opts.ContextLimitFallback,
|
||||
Valid: true,
|
||||
}
|
||||
contextLimit := extractContextLimitWithFallback(
|
||||
result.providerMetadata,
|
||||
opts.ContextLimitFallback,
|
||||
)
|
||||
|
||||
result.content = sanitizeAnthropicProviderToolStepContent(
|
||||
ctx, opts.Logger, provider, modelName,
|
||||
"dynamic_tool_persist", step, result.finishReason, result.content,
|
||||
)
|
||||
if len(result.content) == 0 && len(pending) == 0 {
|
||||
tryCompactOnExit(ctx, opts, result.usage, result.providerMetadata)
|
||||
return ErrDynamicToolCall
|
||||
}
|
||||
|
||||
if err := opts.PersistStep(ctx, PersistedStep{
|
||||
@@ -560,13 +572,21 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
}
|
||||
}
|
||||
// Extract context limit from provider metadata.
|
||||
contextLimit := extractContextLimit(result.providerMetadata)
|
||||
if !contextLimit.Valid && opts.ContextLimitFallback > 0 {
|
||||
contextLimit = sql.NullInt64{
|
||||
Int64: opts.ContextLimitFallback,
|
||||
Valid: true,
|
||||
}
|
||||
contextLimit := extractContextLimitWithFallback(
|
||||
result.providerMetadata,
|
||||
opts.ContextLimitFallback,
|
||||
)
|
||||
result.content = sanitizeAnthropicProviderToolStepContent(
|
||||
ctx, opts.Logger, provider, modelName,
|
||||
"normal_persist", step, result.finishReason, result.content,
|
||||
)
|
||||
if len(result.content) == 0 {
|
||||
lastUsage = result.usage
|
||||
lastProviderMetadata = result.providerMetadata
|
||||
stoppedByModel = true
|
||||
break
|
||||
}
|
||||
|
||||
// Persist the step. If persistence fails because
|
||||
// the chat was interrupted between the previous
|
||||
// check and here, fall back to the interrupt-safe
|
||||
@@ -714,6 +734,67 @@ 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
|
||||
@@ -1281,9 +1362,9 @@ func flushActiveState(
|
||||
}
|
||||
}
|
||||
|
||||
// persistInterruptedStep saves all accumulated content from a
|
||||
// partial stream. Since we own the stepResult directly, no shadow
|
||||
// state is needed.
|
||||
// persistInterruptedStep saves durable content from a partial stream.
|
||||
// Provider-executed calls without results are removed because their
|
||||
// result metadata cannot be synthesized safely.
|
||||
func persistInterruptedStep(
|
||||
ctx context.Context,
|
||||
opts RunOptions,
|
||||
@@ -1293,6 +1374,18 @@ func persistInterruptedStep(
|
||||
return
|
||||
}
|
||||
|
||||
provider := ""
|
||||
modelName := ""
|
||||
if opts.Model != nil {
|
||||
provider = opts.Model.Provider()
|
||||
modelName = opts.Model.Model()
|
||||
}
|
||||
var sanitizeStats chatprompt.AnthropicProviderToolSanitizationStats
|
||||
result.content, sanitizeStats = sanitizeAnthropicProviderToolContent(provider, result.content)
|
||||
chatprompt.LogAnthropicProviderToolSanitization(
|
||||
ctx, opts.Logger, "interrupted_persist", provider, modelName, sanitizeStats,
|
||||
)
|
||||
|
||||
// Track which tool calls already have results in the content.
|
||||
answeredToolCalls := make(map[string]struct{})
|
||||
for _, c := range result.content {
|
||||
@@ -1327,6 +1420,9 @@ func persistInterruptedStep(
|
||||
if _, exists := answeredToolCalls[tc.ToolCallID]; exists {
|
||||
continue
|
||||
}
|
||||
if isAnthropicProviderExecutedToolCall(provider, tc) {
|
||||
continue
|
||||
}
|
||||
content = append(content, fantasy.ToolResultContent{
|
||||
ToolCallID: tc.ToolCallID,
|
||||
ToolName: tc.ToolName,
|
||||
@@ -1344,6 +1440,10 @@ func persistInterruptedStep(
|
||||
answeredToolCalls[tc.ToolCallID] = struct{}{}
|
||||
}
|
||||
|
||||
if len(content) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
persistCtx := context.WithoutCancel(ctx)
|
||||
if err := opts.PersistStep(persistCtx, PersistedStep{
|
||||
Content: content,
|
||||
@@ -1625,6 +1725,17 @@ func extractContextLimit(metadata fantasy.ProviderMetadata) sql.NullInt64 {
|
||||
}
|
||||
}
|
||||
|
||||
func extractContextLimitWithFallback(metadata fantasy.ProviderMetadata, fallback int64) sql.NullInt64 {
|
||||
contextLimit := extractContextLimit(metadata)
|
||||
if contextLimit.Valid || fallback <= 0 {
|
||||
return contextLimit
|
||||
}
|
||||
return sql.NullInt64{
|
||||
Int64: fallback,
|
||||
Valid: true,
|
||||
}
|
||||
}
|
||||
|
||||
func findContextLimitValue(value any) (int64, bool) {
|
||||
var (
|
||||
limit int64
|
||||
|
||||
Reference in New Issue
Block a user