diff --git a/coderd/x/chatd/chatloop/chatloop.go b/coderd/x/chatd/chatloop/chatloop.go index 4e367e818b..34de2e4e74 100644 --- a/coderd/x/chatd/chatloop/chatloop.go +++ b/coderd/x/chatd/chatloop/chatloop.go @@ -128,6 +128,11 @@ type RunOptions struct { // the current step. This is used for plan turns where // propose_plan should terminate the run on success. StopAfterTools map[string]struct{} + // ExclusiveToolNames lists tool names that must be called + // alone in a batch. When any exclusive tool appears + // alongside other locally-executed tools, every tool in the + // batch receives a policy error and nothing executes. + ExclusiveToolNames map[string]bool // ModelConfig holds per-call LLM parameters (temperature, // max tokens, etc.) read from the chat model configuration. @@ -477,102 +482,10 @@ func Run(ctx context.Context, opts RunOptions) error { // blocks into separate database messages by role. var toolResults []fantasy.ToolResultContent if result.shouldContinue { - // Check for context cancellation before starting - // tool execution. If the chat was interrupted - // between stream completion and here, persist - // what we have and bail out. - if ctx.Err() != nil { - if errors.Is(context.Cause(ctx), ErrInterrupted) { - persistInterruptedStep(ctx, opts, &result) - return ErrInterrupted - } - return ctx.Err() - } - - // Partition tool calls into built-in and dynamic. - var builtinCalls, dynamicCalls []fantasy.ToolCallContent - if len(opts.DynamicToolNames) > 0 { - for _, tc := range result.toolCalls { - if opts.DynamicToolNames[tc.ToolName] { - dynamicCalls = append(dynamicCalls, tc) - } else { - builtinCalls = append(builtinCalls, tc) - } - } - } else { - builtinCalls = result.toolCalls - } - - // Execute only built-in tools. - toolResults = executeTools(ctx, opts.Tools, opts.ActiveTools, opts.ProviderTools, builtinCalls, opts.Metrics, opts.Logger, provider, modelName, opts.BuiltinToolNames, func(tr fantasy.ToolResultContent, completedAt time.Time) { - recordToolResultTimestamp(&result, tr.ToolCallID, completedAt) - publishToolAttachments(ctx, opts.Logger, tr, completedAt, publishMessagePart) - ssePart := chatprompt.PartFromContentWithLogger(ctx, opts.Logger, tr) - ssePart.CreatedAt = &completedAt - publishMessagePart(codersdk.ChatMessageRoleTool, ssePart) - }) - for _, tr := range toolResults { - result.content = append(result.content, tr) - } - - // If dynamic tools were called, persist what we - // have (assistant + built-in results) and exit so - // the caller can execute them externally. - if len(dynamicCalls) > 0 { - pending := make([]PendingToolCall, 0, len(dynamicCalls)) - for _, dc := range dynamicCalls { - pending = append(pending, PendingToolCall{ - ToolCallID: dc.ToolCallID, - ToolName: dc.ToolName, - Args: dc.Input, - }) - } - - contextLimit := extractContextLimitWithFallback( - result.providerMetadata, - opts.ContextLimitFallback, - ) - - result.content = chatsanitize.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{ - Content: result.content, - Usage: result.usage, - ContextLimit: contextLimit, - ProviderResponseID: extractOpenAIResponseIDIfStored(opts.ProviderOptions, result.providerMetadata), - Runtime: time.Since(stepStart), - PendingDynamicToolCalls: pending, - }); err != nil { - if errors.Is(err, ErrInterrupted) { - persistInterruptedStep(ctx, opts, &result) - return ErrInterrupted - } - return xerrors.Errorf("persist step: %w", err) - } - - tryCompactOnExit(ctx, opts, result.usage, result.providerMetadata) - - return ErrDynamicToolCall - } - - // Check for interruption after tool execution. - // Tools that were canceled mid-flight produce error - // results via ctx cancellation. Persist the full - // step (assistant blocks + tool results) through - // the interrupt-safe path so nothing is lost. - if ctx.Err() != nil { - if errors.Is(context.Cause(ctx), ErrInterrupted) { - persistInterruptedStep(ctx, opts, &result) - return ErrInterrupted - } - return ctx.Err() + var err error + toolResults, err = executeToolsForStep(ctx, opts, &result, provider, modelName, step, stepStart, publishMessagePart) + if err != nil { + return err } } // Extract context limit from provider metadata. @@ -1154,6 +1067,276 @@ func executeTools( return results } +// executeToolsForStep runs the tool-execution phase of a single +// chatloop step. It enforces the exclusive-tool policy, partitions +// built-in versus dynamic tool calls, dispatches built-in tools, and +// when dynamic tool calls are present persists the step and returns +// ErrDynamicToolCall so the caller can execute them externally. +// Returns the tool results to append to the step, or an error that the +// caller must propagate (ErrInterrupted, ErrDynamicToolCall, ctx.Err(), +// or a persistence failure). +func executeToolsForStep( + ctx context.Context, + opts RunOptions, + result *stepResult, + provider, modelName string, + step int, + stepStart time.Time, + publishMessagePart func(codersdk.ChatMessageRole, codersdk.ChatMessagePart), +) ([]fantasy.ToolResultContent, error) { + // Check for context cancellation before starting tool + // execution. If the chat was interrupted between stream + // completion and here, persist what we have and bail out. + if ctx.Err() != nil { + if errors.Is(context.Cause(ctx), ErrInterrupted) { + persistInterruptedStep(ctx, opts, result) + return nil, ErrInterrupted + } + return nil, ctx.Err() + } + + // Enforce exclusivity across ALL locally-executable tool + // calls (both built-in and dynamic) before partitioning. + // Checking only the built-in partition would let the model + // bypass the policy by mixing an exclusive tool with a + // dynamic tool: the exclusive tool would still run and the + // dynamic call would still be handed to the caller for + // external execution, breaking the planning-only contract. + localCandidates := make([]fantasy.ToolCallContent, 0, len(result.toolCalls)) + for _, tc := range result.toolCalls { + if !tc.ProviderExecuted { + localCandidates = append(localCandidates, tc) + } + } + policyResults, exclusiveViolation := applyExclusiveToolPolicy( + localCandidates, + opts.ExclusiveToolNames, + opts.Metrics, + provider, + modelName, + ) + if exclusiveViolation { + now := dbtime.Now() + for _, tr := range policyResults { + recordToolResultTimestamp(result, tr.ToolCallID, now) + publishToolAttachments(ctx, opts.Logger, tr, now, publishMessagePart) + ssePart := chatprompt.PartFromContentWithLogger(ctx, opts.Logger, tr) + ssePart.CreatedAt = &now + publishMessagePart(codersdk.ChatMessageRoleTool, ssePart) + } + for _, tr := range policyResults { + result.content = append(result.content, tr) + } + // Mirror the post-execution interruption check used by the + // non-policy path: if the chat was interrupted while we + // synthesized policy errors, route through + // persistInterruptedStep so the synthesized results are not + // dropped when the regular PersistStep path fails on a + // canceled context. + if ctx.Err() != nil { + if errors.Is(context.Cause(ctx), ErrInterrupted) { + persistInterruptedStep(ctx, opts, result) + return nil, ErrInterrupted + } + return nil, ctx.Err() + } + // Fall through to the normal persistence path so the loop + // continues with error results that the model can observe + // and retry. Skip partitioning, execution, and + // pending-dynamic persistence. + return policyResults, nil + } + + // Partition tool calls into built-in and dynamic. + var builtinCalls, dynamicCalls []fantasy.ToolCallContent + if len(opts.DynamicToolNames) > 0 { + for _, tc := range result.toolCalls { + if opts.DynamicToolNames[tc.ToolName] { + dynamicCalls = append(dynamicCalls, tc) + } else { + builtinCalls = append(builtinCalls, tc) + } + } + } else { + builtinCalls = result.toolCalls + } + + // Execute only built-in tools. + toolResults := executeTools(ctx, opts.Tools, opts.ActiveTools, opts.ProviderTools, builtinCalls, opts.Metrics, opts.Logger, provider, modelName, opts.BuiltinToolNames, func(tr fantasy.ToolResultContent, completedAt time.Time) { + recordToolResultTimestamp(result, tr.ToolCallID, completedAt) + publishToolAttachments(ctx, opts.Logger, tr, completedAt, publishMessagePart) + ssePart := chatprompt.PartFromContentWithLogger(ctx, opts.Logger, tr) + ssePart.CreatedAt = &completedAt + publishMessagePart(codersdk.ChatMessageRoleTool, ssePart) + }) + for _, tr := range toolResults { + result.content = append(result.content, tr) + } + + // If dynamic tools were called, persist what we have + // (assistant + built-in results) and exit so the caller can + // execute them externally. + if len(dynamicCalls) > 0 { + // Strip Anthropic provider-executed tool calls without + // matching results before persisting so the action-required + // step does not carry a malformed tool-call history into + // downstream provider requests. + result.content = chatsanitize.SanitizeAnthropicProviderToolStepContent( + ctx, opts.Logger, provider, modelName, + "dynamic_tool_persist", step, result.finishReason, result.content, + ) + if err := persistPendingDynamicStep(ctx, opts, result, stepStart, dynamicCalls); err != nil { + return nil, err + } + tryCompactOnExit(ctx, opts, result.usage, result.providerMetadata) + return nil, ErrDynamicToolCall + } + + // Check for interruption after tool execution. Tools that + // were canceled mid-flight produce error results via ctx + // cancellation. Persist the full step (assistant blocks + + // tool results) through the interrupt-safe path so nothing + // is lost. + if ctx.Err() != nil { + if errors.Is(context.Cause(ctx), ErrInterrupted) { + persistInterruptedStep(ctx, opts, result) + return nil, ErrInterrupted + } + return nil, ctx.Err() + } + + return toolResults, nil +} + +// persistPendingDynamicStep persists a step that has pending dynamic +// tool calls awaiting external execution. Returns ErrInterrupted when +// persistence fails because the chat was interrupted. +func persistPendingDynamicStep( + ctx context.Context, + opts RunOptions, + result *stepResult, + stepStart time.Time, + dynamicCalls []fantasy.ToolCallContent, +) error { + pending := make([]PendingToolCall, 0, len(dynamicCalls)) + for _, dc := range dynamicCalls { + pending = append(pending, PendingToolCall{ + ToolCallID: dc.ToolCallID, + ToolName: dc.ToolName, + Args: dc.Input, + }) + } + + contextLimit := extractContextLimitWithFallback(result.providerMetadata, opts.ContextLimitFallback) + + if err := opts.PersistStep(ctx, PersistedStep{ + Content: result.content, + Usage: result.usage, + ContextLimit: contextLimit, + ProviderResponseID: extractOpenAIResponseIDIfStored(opts.ProviderOptions, result.providerMetadata), + Runtime: time.Since(stepStart), + PendingDynamicToolCalls: pending, + }); err != nil { + if errors.Is(err, ErrInterrupted) { + persistInterruptedStep(ctx, opts, result) + return ErrInterrupted + } + return xerrors.Errorf("persist step: %w", err) + } + return nil +} + +// applyExclusiveToolPolicy checks whether toolCalls violate the +// exclusive-tool policy declared by exclusiveToolNames. When a +// violation is detected it synthesizes deterministic policy-error +// results for every tool call and records size/error metrics so the +// exclusivity failure mode is visible to operators. Returns +// (results, true) on violation; (nil, false) otherwise. +func applyExclusiveToolPolicy( + toolCalls []fantasy.ToolCallContent, + exclusiveToolNames map[string]bool, + metrics *Metrics, + provider, model string, +) ([]fantasy.ToolResultContent, bool) { + blockingToolName, ok := firstExclusiveToolName(toolCalls, exclusiveToolNames) + if !ok { + return nil, false + } + results := exclusiveToolPolicyResults(toolCalls, exclusiveToolNames, blockingToolName) + for _, tr := range results { + recordToolResultMetrics(metrics, provider, model, tr) + } + return results, true +} + +// recordToolResultMetrics observes tool result size and increments +// tool_errors_total when the result carries an error output. Mirrors +// the metric-recording defer in executeSingleTool so that synthetic +// results (e.g. exclusive-tool policy errors) contribute to operator +// visibility. +func recordToolResultMetrics(metrics *Metrics, provider, model string, tr fantasy.ToolResultContent) { + if metrics == nil { + return + } + label := tr.ToolName + if label == "" { + label = "unknown" + } + metrics.ToolResultSizeBytes.WithLabelValues(provider, model, label).Observe( + float64(ToolResultSize(tr)), + ) + if _, ok := tr.Result.(fantasy.ToolResultOutputContentError); ok { + metrics.RecordToolError(provider, model, label) + } +} + +func firstExclusiveToolName( + toolCalls []fantasy.ToolCallContent, + exclusiveToolNames map[string]bool, +) (string, bool) { + if len(toolCalls) <= 1 || len(exclusiveToolNames) == 0 { + return "", false + } + + for _, tc := range toolCalls { + if exclusiveToolNames[tc.ToolName] { + return tc.ToolName, true + } + } + + return "", false +} + +func exclusiveToolPolicyResults( + toolCalls []fantasy.ToolCallContent, + exclusiveToolNames map[string]bool, + blockingToolName string, +) []fantasy.ToolResultContent { + results := make([]fantasy.ToolResultContent, len(toolCalls)) + for i, tc := range toolCalls { + message := exclusiveToolSkippedErrorMessage(blockingToolName) + if exclusiveToolNames[tc.ToolName] { + message = exclusiveToolMustRunAloneErrorMessage(tc.ToolName) + } + results[i] = fantasy.ToolResultContent{ + ToolCallID: tc.ToolCallID, + ToolName: tc.ToolName, + Result: fantasy.ToolResultOutputContentError{ + Error: xerrors.New(message), + }, + } + } + return results +} + +func exclusiveToolMustRunAloneErrorMessage(toolName string) string { + return toolName + " must be called alone, without other tools in the same batch. Retry with only the " + toolName + " call." +} + +func exclusiveToolSkippedErrorMessage(toolName string) string { + return "this tool was skipped because " + toolName + " must run alone in its batch. Retry your tool calls without " + toolName + ", or call " + toolName + " separately first." +} + // executeSingleTool executes one tool call and converts the // response into a ToolResultContent. func executeSingleTool( diff --git a/coderd/x/chatd/chatloop/chatloop_test.go b/coderd/x/chatd/chatloop/chatloop_test.go index 08489773ce..3b3a4483ed 100644 --- a/coderd/x/chatd/chatloop/chatloop_test.go +++ b/coderd/x/chatd/chatloop/chatloop_test.go @@ -15,6 +15,7 @@ import ( "charm.land/fantasy" fantasyanthropic "charm.land/fantasy/providers/anthropic" "github.com/prometheus/client_golang/prometheus" + promtestutil "github.com/prometheus/client_golang/prometheus/testutil" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/xerrors" @@ -1079,6 +1080,19 @@ func TestRun_InterruptedStepPersistsSyntheticToolResult(t *testing.T) { "interrupted tool should have no call timestamp (never reached StreamPartTypeToolCall)") } +func requireToolResultErrorMessage( + t *testing.T, + result fantasy.ToolResultContent, + expected string, +) { + t.Helper() + + output, ok := result.Result.(fantasy.ToolResultOutputContentError) + require.Truef(t, ok, "expected error tool result, got %T", result.Result) + require.Error(t, output.Error) + require.Equal(t, expected, output.Error.Error()) +} + func streamFromParts(parts []fantasy.StreamPart) fantasy.StreamResponse { return iter.Seq[fantasy.StreamPart](func(yield func(fantasy.StreamPart) bool) { for _, part := range parts { @@ -1643,6 +1657,578 @@ func TestRun_ParallelToolExecutionTimestamps(t *testing.T) { "tc-2 tool-result timestamp must be >= tool-call timestamp") } +// TestRun_ExclusiveToolPolicyViolation exercises the full Run() -> +// executeToolsForStep() -> applyExclusiveToolPolicy() wiring. When an +// exclusive tool is called alongside other locally-executable tools, +// neither runner must fire and every call in the batch must receive a +// synthesized policy error that is both persisted and published via +// SSE. This guards against a regression where +// executeToolsForStep's policy call is accidentally removed: the +// pure-unit tests cover the policy function in isolation, but only +// this test catches a broken wiring path. +func TestRun_ExclusiveToolPolicyViolation(t *testing.T) { + t.Parallel() + + var advisorRuns atomic.Int32 + advisorTool := fantasy.NewAgentTool( + "advisor", + "returns strategic guidance", + func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) { + advisorRuns.Add(1) + return fantasy.NewTextResponse(`{"status":"ok"}`), nil + }, + ) + var readRuns atomic.Int32 + readTool := fantasy.NewAgentTool( + "read_file", + "reads a file", + func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) { + readRuns.Add(1) + return fantasy.NewTextResponse(`{"contents":"main"}`), nil + }, + ) + + var mu sync.Mutex + var streamCalls int + model := &chattest.FakeModel{ + ProviderName: "fake", + StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) { + mu.Lock() + step := streamCalls + streamCalls++ + mu.Unlock() + + if step == 0 { + // Step 0: model emits an illegal mixed batch. + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeToolInputStart, ID: "advisor-1", ToolCallName: "advisor"}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "advisor-1", Delta: `{}`}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "advisor-1"}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "advisor-1", + ToolCallName: "advisor", + ToolCallInput: `{}`, + }, + {Type: fantasy.StreamPartTypeToolInputStart, ID: "read-1", ToolCallName: "read_file"}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "read-1", Delta: `{"path":"main.go"}`}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "read-1"}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "read-1", + ToolCallName: "read_file", + ToolCallInput: `{"path":"main.go"}`, + }, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls}, + }), nil + } + // Step 1: the loop re-streams after tool results; end the run. + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeTextStart, ID: "text-1"}, + {Type: fantasy.StreamPartTypeTextDelta, ID: "text-1", Delta: "ok, retrying"}, + {Type: fantasy.StreamPartTypeTextEnd, ID: "text-1"}, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonStop}, + }), nil + }, + } + + var persistedSteps []PersistedStep + var publishedToolParts []codersdk.ChatMessagePart + err := Run(context.Background(), RunOptions{ + Model: model, + Messages: []fantasy.Message{ + textMessage(fantasy.MessageRoleUser, "please advise and read"), + }, + Tools: []fantasy.AgentTool{advisorTool, readTool}, + ExclusiveToolNames: map[string]bool{"advisor": true}, + MaxSteps: 5, + PersistStep: func(_ context.Context, step PersistedStep) error { + persistedSteps = append(persistedSteps, step) + return nil + }, + PublishMessagePart: func(role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) { + if role != codersdk.ChatMessageRoleTool { + return + } + publishedToolParts = append(publishedToolParts, part) + }, + }) + require.NoError(t, err) + + // Neither runner must have fired: the policy short-circuits + // before partitioning and execution. + require.Equal(t, int32(0), advisorRuns.Load(), + "advisor runner must not fire on mixed batches") + require.Equal(t, int32(0), readRuns.Load(), + "read_file runner must not fire on mixed batches") + + // Two steps: the mixed-batch step plus the follow-up stream. + require.Len(t, persistedSteps, 2) + firstStep := persistedSteps[0] + + advisorErr, ok := findToolResultByID(firstStep.Content, "advisor-1") + require.True(t, ok, "persisted step must contain the advisor policy result") + requireToolResultErrorMessage(t, advisorErr, + "advisor must be called alone, without other tools in the same batch. Retry with only the advisor call.") + + readErr, ok := findToolResultByID(firstStep.Content, "read-1") + require.True(t, ok, "persisted step must contain the read_file policy result") + requireToolResultErrorMessage(t, readErr, + "this tool was skipped because advisor must run alone in its batch. Retry your tool calls without advisor, or call advisor separately first.") + + // Policy-error results must be SSE-published so the client + // can render them immediately. Confirm both tool-result parts + // reached PublishMessagePart with a non-nil CreatedAt, which + // is the dbtime.Now() stamp the policy branch sets. + var sawAdvisorPart, sawReadPart bool + for _, part := range publishedToolParts { + switch part.ToolCallID { + case "advisor-1": + sawAdvisorPart = true + require.NotNil(t, part.CreatedAt, + "policy result SSE part must carry the dbtime.Now() timestamp") + case "read-1": + sawReadPart = true + require.NotNil(t, part.CreatedAt, + "policy result SSE part must carry the dbtime.Now() timestamp") + } + } + require.True(t, sawAdvisorPart, "advisor policy result must be SSE-published") + require.True(t, sawReadPart, "read_file policy result must be SSE-published") +} + +func findToolResultByID( + content []fantasy.Content, + toolCallID string, +) (fantasy.ToolResultContent, bool) { + for _, block := range content { + tr, ok := fantasy.AsContentType[fantasy.ToolResultContent](block) + if !ok { + continue + } + if tr.ToolCallID == toolCallID { + return tr, true + } + } + return fantasy.ToolResultContent{}, false +} + +func TestExclusiveToolPolicy_MixedBatchErrors(t *testing.T) { + t.Parallel() + + results, violated := applyExclusiveToolPolicy( + []fantasy.ToolCallContent{ + {ToolCallID: "advisor-1", ToolName: "advisor", Input: `{}`}, + {ToolCallID: "read-1", ToolName: "read_file", Input: `{"path":"main.go"}`}, + }, + map[string]bool{"advisor": true}, + NopMetrics(), + "fake", + "", + ) + + require.True(t, violated) + require.Len(t, results, 2) + require.Equal(t, "advisor-1", results[0].ToolCallID) + require.Equal(t, "read-1", results[1].ToolCallID) + requireToolResultErrorMessage( + t, + results[0], + "advisor must be called alone, without other tools in the same batch. Retry with only the advisor call.", + ) + requireToolResultErrorMessage( + t, + results[1], + "this tool was skipped because advisor must run alone in its batch. Retry your tool calls without advisor, or call advisor separately first.", + ) +} + +func TestApplyExclusiveToolPolicy_RecordsErrorMetrics(t *testing.T) { + t.Parallel() + + reg := prometheus.NewPedanticRegistry() + m := NewMetrics(reg) + + _, violated := applyExclusiveToolPolicy( + []fantasy.ToolCallContent{ + {ToolCallID: "advisor-1", ToolName: "advisor", Input: `{}`}, + {ToolCallID: "read-1", ToolName: "read_file", Input: `{"path":"main.go"}`}, + }, + map[string]bool{"advisor": true}, + m, + "fake", + "claude-test", + ) + require.True(t, violated) + + require.Equal(t, 1.0, promtestutil.ToFloat64( + m.ToolErrorsTotal.WithLabelValues("fake", "claude-test", "advisor"), + )) + require.Equal(t, 1.0, promtestutil.ToFloat64( + m.ToolErrorsTotal.WithLabelValues("fake", "claude-test", "read_file"), + )) +} + +func TestExclusiveToolPolicy_MultipleExclusive(t *testing.T) { + t.Parallel() + + results, violated := applyExclusiveToolPolicy( + []fantasy.ToolCallContent{ + {ToolCallID: "advisor-1", ToolName: "advisor", Input: `{}`}, + {ToolCallID: "advisor-2", ToolName: "advisor", Input: `{"mode":"second-opinion"}`}, + }, + map[string]bool{"advisor": true}, + NopMetrics(), + "fake", + "", + ) + + require.True(t, violated) + require.Len(t, results, 2) + requireToolResultErrorMessage( + t, + results[0], + "advisor must be called alone, without other tools in the same batch. Retry with only the advisor call.", + ) + requireToolResultErrorMessage( + t, + results[1], + "advisor must be called alone, without other tools in the same batch. Retry with only the advisor call.", + ) +} + +// TestRun_ExclusiveToolPolicyBlocksMixedWithDynamicTool guards the +// exclusive-over-dynamic bypass: the policy must run before the +// built-in vs dynamic partition. If a future refactor moves the +// policy check beneath the partition (so only built-in calls are +// inspected), an exclusive builtin mixed with a dynamic tool would +// still execute locally while the dynamic call is handed off via +// ErrDynamicToolCall, breaking the planning-only contract. +// +// This test has the model emit an exclusive builtin (advisor) +// alongside a dynamic tool (mcp_tool) in the same batch and asserts +// that Run does NOT exit with ErrDynamicToolCall, the advisor +// runner never fires, and both calls receive a synthesized policy +// error. +func TestRun_ExclusiveToolPolicyBlocksMixedWithDynamicTool(t *testing.T) { + t.Parallel() + + var advisorRuns atomic.Int32 + advisorTool := fantasy.NewAgentTool( + "advisor", + "returns strategic guidance", + func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) { + advisorRuns.Add(1) + return fantasy.NewTextResponse(`{"status":"ok"}`), nil + }, + ) + + var mu sync.Mutex + var streamCalls int + model := &chattest.FakeModel{ + ProviderName: "fake", + StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) { + mu.Lock() + step := streamCalls + streamCalls++ + mu.Unlock() + + if step == 0 { + // Step 0: model emits an illegal mixed batch + // combining an exclusive builtin with a + // dynamic tool. + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeToolInputStart, ID: "advisor-1", ToolCallName: "advisor"}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "advisor-1", Delta: `{}`}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "advisor-1"}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "advisor-1", + ToolCallName: "advisor", + ToolCallInput: `{}`, + }, + {Type: fantasy.StreamPartTypeToolInputStart, ID: "mcp-1", ToolCallName: "mcp_tool"}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "mcp-1", Delta: `{"q":"docs"}`}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "mcp-1"}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "mcp-1", + ToolCallName: "mcp_tool", + ToolCallInput: `{"q":"docs"}`, + }, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls}, + }), nil + } + // Step 1: after the policy error is fed back, + // terminate the run so the test assertions have a + // deterministic exit. + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeTextStart, ID: "text-1"}, + {Type: fantasy.StreamPartTypeTextDelta, ID: "text-1", Delta: "retrying"}, + {Type: fantasy.StreamPartTypeTextEnd, ID: "text-1"}, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonStop}, + }), nil + }, + } + + var persistedSteps []PersistedStep + err := Run(context.Background(), RunOptions{ + Model: model, + Messages: []fantasy.Message{ + textMessage(fantasy.MessageRoleUser, "please advise and fetch"), + }, + Tools: []fantasy.AgentTool{advisorTool}, + DynamicToolNames: map[string]bool{"mcp_tool": true}, + ExclusiveToolNames: map[string]bool{"advisor": true}, + MaxSteps: 5, + PersistStep: func(_ context.Context, step PersistedStep) error { + persistedSteps = append(persistedSteps, step) + return nil + }, + }) + // Run must NOT exit with ErrDynamicToolCall: the policy + // short-circuits before the dynamic partition so the dynamic + // call is never handed off for external execution. + require.NoError(t, err) + + // The advisor runner must not fire on mixed batches; the + // policy blocks the whole batch including the exclusive tool + // itself. + require.Equal(t, int32(0), advisorRuns.Load(), + "advisor runner must not fire on mixed batches") + + // Two steps: the mixed-batch step with synthesized policy + // errors plus the follow-up stream that ends the run. + require.Len(t, persistedSteps, 2) + firstStep := persistedSteps[0] + + // The persisted step must not record the dynamic tool as + // pending: the policy-error path returns before + // persistPendingDynamicStep runs. + require.Empty(t, firstStep.PendingDynamicToolCalls, + "policy-rejected batches must not leak dynamic tool calls to the caller") + + advisorErr, ok := findToolResultByID(firstStep.Content, "advisor-1") + require.True(t, ok, "persisted step must contain the advisor policy result") + requireToolResultErrorMessage(t, advisorErr, + "advisor must be called alone, without other tools in the same batch. Retry with only the advisor call.") + + mcpErr, ok := findToolResultByID(firstStep.Content, "mcp-1") + require.True(t, ok, "persisted step must contain the mcp_tool policy result") + requireToolResultErrorMessage(t, mcpErr, + "this tool was skipped because advisor must run alone in its batch. Retry your tool calls without advisor, or call advisor separately first.") +} + +// TestRun_ExclusiveToolAloneSucceeds is the happy-path counterpart +// to TestRun_ExclusiveToolPolicyViolation: a single exclusive tool +// emitted alone must actually execute. The `len(toolCalls) <= 1` +// guard in firstExclusiveToolName is the sole mechanism that lets +// solo exclusive-tool calls proceed. If that guard regresses to +// `< 1`, every solo exclusive-tool call would enter an infinite +// policy-error/retry loop, and every unit test on the policy +// function in isolation would still pass. Only this Run()-level +// test catches that regression. +func TestRun_ExclusiveToolAloneSucceeds(t *testing.T) { + t.Parallel() + + var advisorRuns atomic.Int32 + advisorTool := fantasy.NewAgentTool( + "advisor", + "returns strategic guidance", + func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) { + advisorRuns.Add(1) + return fantasy.NewTextResponse(`{"status":"ok"}`), nil + }, + ) + + var mu sync.Mutex + var streamCalls int + model := &chattest.FakeModel{ + ProviderName: "fake", + StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) { + mu.Lock() + step := streamCalls + streamCalls++ + mu.Unlock() + + if step == 0 { + // Step 0: model emits exactly one + // exclusive-tool call in isolation. + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeToolInputStart, ID: "advisor-1", ToolCallName: "advisor"}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "advisor-1", Delta: `{}`}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "advisor-1"}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "advisor-1", + ToolCallName: "advisor", + ToolCallInput: `{}`, + }, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls}, + }), nil + } + // Step 1: the loop re-streams after the tool + // result; end the run. + 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 + }, + } + + var persistedSteps []PersistedStep + err := Run(context.Background(), RunOptions{ + Model: model, + Messages: []fantasy.Message{ + textMessage(fantasy.MessageRoleUser, "please advise"), + }, + Tools: []fantasy.AgentTool{advisorTool}, + ExclusiveToolNames: map[string]bool{"advisor": true}, + MaxSteps: 5, + PersistStep: func(_ context.Context, step PersistedStep) error { + persistedSteps = append(persistedSteps, step) + return nil + }, + }) + require.NoError(t, err) + + // The solo exclusive tool must actually execute exactly once. + require.Equal(t, int32(1), advisorRuns.Load(), + "solo exclusive-tool call must execute") + + // The first persisted step must contain a non-error tool + // result for the advisor call, proving the policy did not + // synthesize an error and the real runner fired. + require.GreaterOrEqual(t, len(persistedSteps), 1) + result, ok := findToolResultByID(persistedSteps[0].Content, "advisor-1") + require.True(t, ok, "persisted step must contain the advisor tool result") + _, isErr := result.Result.(fantasy.ToolResultOutputContentError) + require.Falsef(t, isErr, + "solo exclusive-tool call must produce a real tool result, not a policy error: %+v", result.Result) +} + +// TestRun_ExclusiveToolWithProviderExecutedSucceeds guards the +// interaction between the ProviderExecuted filter and the +// exclusive-tool policy. executeToolsForStep builds localCandidates +// by dropping ProviderExecuted calls before passing them to +// applyExclusiveToolPolicy. That filter is the sole mechanism +// preventing a false policy violation when a solo exclusive tool +// appears in a batch where the provider also server-executed a tool +// (for example Anthropic web_search). +// +// If the filter is removed, localCandidates would contain both the +// provider-executed call and the exclusive call. firstExclusiveToolName +// would then see len > 1, find advisor, and return a violation. The +// advisor would never run and the retry loop would burn steps until +// MaxSteps. +// +// This test emits an advisor call alongside a provider-executed +// web_search call (with its provider-emitted result) and asserts the +// advisor runner actually fires. +func TestRun_ExclusiveToolWithProviderExecutedSucceeds(t *testing.T) { + t.Parallel() + + var advisorRuns atomic.Int32 + advisorTool := fantasy.NewAgentTool( + "advisor", + "returns strategic guidance", + func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) { + advisorRuns.Add(1) + return fantasy.NewTextResponse(`{"status":"ok"}`), nil + }, + ) + + var mu sync.Mutex + var streamCalls int + model := &chattest.FakeModel{ + ProviderName: "fake", + StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) { + mu.Lock() + step := streamCalls + streamCalls++ + mu.Unlock() + + if step == 0 { + // Step 0: provider server-executed web_search and + // returned its result inline, plus the model + // emitted an exclusive advisor call for local + // execution. The ProviderExecuted filter must + // drop web_search from the policy check so the + // advisor is treated as a solo exclusive call. + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeToolInputStart, ID: "ws-1", ToolCallName: "web_search", ProviderExecuted: true}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "ws-1", Delta: `{"query":"coder"}`, ProviderExecuted: true}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "ws-1"}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "ws-1", + ToolCallName: "web_search", + ToolCallInput: `{"query":"coder"}`, + ProviderExecuted: true, + }, + { + Type: fantasy.StreamPartTypeToolResult, + ID: "ws-1", + ToolCallName: "web_search", + ProviderExecuted: true, + }, + {Type: fantasy.StreamPartTypeToolInputStart, ID: "advisor-1", ToolCallName: "advisor"}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "advisor-1", Delta: `{}`}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "advisor-1"}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "advisor-1", + ToolCallName: "advisor", + ToolCallInput: `{}`, + }, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls}, + }), nil + } + // Step 1: end the run after the advisor result is + // fed back. + 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 + }, + } + + var persistedSteps []PersistedStep + err := Run(context.Background(), RunOptions{ + Model: model, + Messages: []fantasy.Message{ + textMessage(fantasy.MessageRoleUser, "search and then advise"), + }, + Tools: []fantasy.AgentTool{advisorTool}, + ExclusiveToolNames: map[string]bool{"advisor": true}, + MaxSteps: 5, + PersistStep: func(_ context.Context, step PersistedStep) error { + persistedSteps = append(persistedSteps, step) + return nil + }, + }) + require.NoError(t, err) + + // The advisor must execute exactly once: the ProviderExecuted + // filter removes web_search from the exclusivity check, so the + // advisor is treated as a solo exclusive call. + require.Equal(t, int32(1), advisorRuns.Load(), + "advisor must execute when the only other call in the batch was provider-executed") + + // The advisor result must be a real tool result, not a + // synthesized policy error. + require.GreaterOrEqual(t, len(persistedSteps), 1) + advisorResult, ok := findToolResultByID(persistedSteps[0].Content, "advisor-1") + require.True(t, ok, "persisted step must contain the advisor tool result") + _, isErr := advisorResult.Result.(fantasy.ToolResultOutputContentError) + require.Falsef(t, isErr, + "advisor must produce a real tool result, not a policy error: %+v", advisorResult.Result) +} + func TestRun_PersistStepErrorPropagates(t *testing.T) { t.Parallel()