feat: add plan mode with restricted tool boundary (#24236)

> This PR was authored by Mux on behalf of Mike.

## Summary
- add persistent plan mode for chats and the chat-specific plan file
flow
- add structured planning tools such as `ask_user_question` and
`propose_plan`
- keep `write_file` and `edit_files` constrained to the chat-specific
plan file during plan turns
- allow shell exploration in plan mode, including subagents, via
`execute` and `process_output`
- block implementation-oriented, provider-native, MCP, dynamic, and
computer-use tools during plan turns
- update the chat UI, tests, and docs for the new planning flow
This commit is contained in:
Michael Suchacz
2026-04-16 11:12:01 +02:00
committed by GitHub
parent e996f6d44b
commit 1cf0354f72
76 changed files with 6398 additions and 889 deletions
+54 -4
View File
@@ -43,6 +43,10 @@ const (
var (
ErrInterrupted = xerrors.New("chat interrupted")
ErrDynamicToolCall = xerrors.New("dynamic tool call")
// ErrStopAfterTool is returned when a tool listed in
// StopAfterTools produces a successful result, indicating
// the run should terminate cleanly after persistence.
ErrStopAfterTool = xerrors.New("stop after tool")
errStartupTimeout = xerrors.New(
"chat response did not start before the startup timeout",
@@ -114,6 +118,11 @@ type RunOptions struct {
// the chatloop persists partial results and exits with
// ErrDynamicToolCall instead of executing the tool.
DynamicToolNames map[string]bool
// StopAfterTools lists tool names that, when they produce a
// successful result, cause the run to stop after persisting
// the current step. This is used for plan turns where
// propose_plan should terminate the run on success.
StopAfterTools map[string]struct{}
// ModelConfig holds per-call LLM parameters (temperature,
// max tokens, etc.) read from the chat model configuration.
@@ -472,7 +481,7 @@ func Run(ctx context.Context, opts RunOptions) error {
}
// Execute only built-in tools.
toolResults = executeTools(ctx, opts.Tools, opts.ProviderTools, builtinCalls, opts.Metrics, provider, opts.BuiltinToolNames, func(tr fantasy.ToolResultContent, completedAt time.Time) {
toolResults = executeTools(ctx, opts.Tools, opts.ActiveTools, opts.ProviderTools, builtinCalls, opts.Metrics, provider, opts.BuiltinToolNames, func(tr fantasy.ToolResultContent, completedAt time.Time) {
recordToolResultTimestamp(&result, tr.ToolCallID, completedAt)
ssePart := chatprompt.PartFromContent(tr)
ssePart.CreatedAt = &completedAt
@@ -566,6 +575,12 @@ func Run(ctx context.Context, opts RunOptions) error {
lastUsage = result.usage
lastProviderMetadata = result.providerMetadata
// Check if any executed tool triggers an early stop.
if shouldStopAfterTools(opts.StopAfterTools, toolResults) {
tryCompactOnExit(ctx, opts, result.usage, result.providerMetadata)
return ErrStopAfterTool
}
// When chain mode is active (PreviousResponseID set), exit
// it after persisting the first chained step. Continuation
// steps include tool-result messages, which fantasy rejects
@@ -1022,6 +1037,7 @@ func processStepStream(
func executeTools(
ctx context.Context,
allTools []fantasy.AgentTool,
activeTools []string,
providerTools []ProviderTool,
toolCalls []fantasy.ToolCallContent,
metrics *Metrics,
@@ -1051,11 +1067,14 @@ func executeTools(
for _, t := range allTools {
toolMap[t.Info().Name] = t
}
providerRunnerNames := make(map[string]struct{}, len(providerTools))
// Include runners from provider tools so locally-executed
// provider tools (e.g. computer use) can be dispatched.
for _, pt := range providerTools {
if pt.Runner != nil {
toolMap[pt.Runner.Info().Name] = pt.Runner
name := pt.Runner.Info().Name
toolMap[name] = pt.Runner
providerRunnerNames[name] = struct{}{}
}
}
@@ -1081,7 +1100,7 @@ func executeTools(
// accurate individual completion times.
completedAt[i] = dbtime.Now()
}()
results[i] = executeSingleTool(ctx, toolMap, tc, metrics, provider, builtinToolNames)
results[i] = executeSingleTool(ctx, toolMap, tc, metrics, provider, builtinToolNames, activeTools, providerRunnerNames)
}()
}
wg.Wait()
@@ -1105,6 +1124,8 @@ func executeSingleTool(
metrics *Metrics,
provider string,
builtinToolNames map[string]bool,
activeTools []string,
providerRunnerNames map[string]struct{},
) fantasy.ToolResultContent {
result := fantasy.ToolResultContent{
ToolCallID: tc.ToolCallID,
@@ -1121,6 +1142,13 @@ func executeSingleTool(
)
}()
if _, isProviderRunner := providerRunnerNames[tc.ToolName]; !isProviderRunner && !isToolActive(tc.ToolName, activeTools) {
result.Result = fantasy.ToolResultOutputContentError{
Error: xerrors.New("Tool not active in this turn: " + tc.ToolName),
}
return result
}
tool, exists := toolMap[tc.ToolName]
if !exists {
result.Result = fantasy.ToolResultOutputContentError{
@@ -1325,6 +1353,10 @@ func tryCompactOnExit(
}
}
func isToolActive(name string, activeTools []string) bool {
return len(activeTools) == 0 || slices.Contains(activeTools, name)
}
// buildToolDefinitions converts AgentTool definitions into the
// fantasy.Tool slice expected by fantasy.Call. When activeTools
// is non-empty, only function tools whose name appears in the
@@ -1334,7 +1366,7 @@ func buildToolDefinitions(tools []fantasy.AgentTool, activeTools []string, provi
prepared := make([]fantasy.Tool, 0, len(tools)+len(providerTools))
for _, tool := range tools {
info := tool.Info()
if len(activeTools) > 0 && !slices.Contains(activeTools, info.Name) {
if !isToolActive(info.Name, activeTools) {
continue
}
@@ -1361,6 +1393,24 @@ func buildToolDefinitions(tools []fantasy.AgentTool, activeTools []string, provi
return prepared
}
// shouldStopAfterTools returns true if any tool result in the
// slice matches a name in stopTools and produced a successful
// (non-error) result.
func shouldStopAfterTools(stopTools map[string]struct{}, results []fantasy.ToolResultContent) bool {
if len(stopTools) == 0 {
return false
}
for _, tr := range results {
if _, ok := stopTools[tr.ToolName]; !ok {
continue
}
if _, isErr := tr.Result.(fantasy.ToolResultOutputContentError); !isErr {
return true
}
}
return false
}
func shouldApplyAnthropicPromptCaching(model fantasy.LanguageModel) bool {
if model == nil {
return false
+282
View File
@@ -101,6 +101,150 @@ func TestRun_ActiveToolsPrepareBehavior(t *testing.T) {
require.True(t, hasAnthropicEphemeralCacheControl(capturedCall.Prompt[4]))
}
func TestRun_ActiveToolsRejectsDisallowedExecution(t *testing.T) {
t.Parallel()
var blockedCalls atomic.Int32
blockedToolName := "write_file"
model := &chattest.FakeModel{
ProviderName: "fake",
StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) {
return streamFromParts([]fantasy.StreamPart{
{Type: fantasy.StreamPartTypeToolInputStart, ID: "tc-blocked", ToolCallName: blockedToolName},
{Type: fantasy.StreamPartTypeToolInputDelta, ID: "tc-blocked", Delta: `{"path":"/tmp/nope"}`},
{Type: fantasy.StreamPartTypeToolInputEnd, ID: "tc-blocked"},
{
Type: fantasy.StreamPartTypeToolCall,
ID: "tc-blocked",
ToolCallName: blockedToolName,
ToolCallInput: `{"path":"/tmp/nope"}`,
},
{Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls},
}), nil
},
}
blockedTool := fantasy.NewAgentTool(
blockedToolName,
"blocked tool",
func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) {
blockedCalls.Add(1)
return fantasy.NewTextResponse("should not run"), nil
},
)
var persistedStep PersistedStep
err := Run(context.Background(), RunOptions{
Model: model,
Messages: []fantasy.Message{
textMessage(fantasy.MessageRoleUser, "try the blocked tool"),
},
Tools: []fantasy.AgentTool{
newNoopTool(activeToolName),
blockedTool,
},
ActiveTools: []string{activeToolName},
MaxSteps: 1,
PersistStep: func(_ context.Context, step PersistedStep) error {
persistedStep = step
return nil
},
})
require.NoError(t, err)
require.Zero(t, blockedCalls.Load(), "disallowed tool must not execute")
var foundToolError bool
for _, block := range persistedStep.Content {
toolResult, ok := fantasy.AsContentType[fantasy.ToolResultContent](block)
if !ok || toolResult.ToolName != blockedToolName {
continue
}
errResult, ok := toolResult.Result.(fantasy.ToolResultOutputContentError)
require.True(t, ok)
assert.EqualError(t, errResult.Error, "Tool not active in this turn: "+blockedToolName)
foundToolError = true
}
require.True(t, foundToolError, "persisted step should include the rejected tool result")
}
func TestRun_ActiveToolsAllowsProviderRunnerExecution(t *testing.T) {
t.Parallel()
providerRunnerName := "computer"
var runnerCalls atomic.Int32
model := &chattest.FakeModel{
ProviderName: "fake",
StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) {
return streamFromParts([]fantasy.StreamPart{
{Type: fantasy.StreamPartTypeToolInputStart, ID: "tc-provider-runner", ToolCallName: providerRunnerName},
{Type: fantasy.StreamPartTypeToolInputDelta, ID: "tc-provider-runner", Delta: `{}`},
{Type: fantasy.StreamPartTypeToolInputEnd, ID: "tc-provider-runner"},
{
Type: fantasy.StreamPartTypeToolCall,
ID: "tc-provider-runner",
ToolCallName: providerRunnerName,
ToolCallInput: `{}`,
},
{Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls},
}), nil
},
}
runnerTool := fantasy.NewAgentTool(
providerRunnerName,
"provider runner",
func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) {
runnerCalls.Add(1)
return fantasy.NewTextResponse("ran provider runner"), nil
},
)
var persistedStep PersistedStep
err := Run(context.Background(), RunOptions{
Model: model,
Messages: []fantasy.Message{
textMessage(fantasy.MessageRoleUser, "use the computer"),
},
Tools: []fantasy.AgentTool{newNoopTool(activeToolName)},
ActiveTools: []string{activeToolName},
ProviderTools: []ProviderTool{
{
Definition: fantasy.FunctionTool{
Name: providerRunnerName,
Description: "provider runner",
InputSchema: map[string]any{
"type": "object",
"properties": map[string]any{},
},
},
Runner: runnerTool,
},
},
MaxSteps: 1,
PersistStep: func(_ context.Context, step PersistedStep) error {
persistedStep = step
return nil
},
})
require.NoError(t, err)
require.Equal(t, int32(1), runnerCalls.Load(),
"provider runner should execute even when omitted from active tools")
var foundToolResult bool
for _, block := range persistedStep.Content {
toolResult, ok := fantasy.AsContentType[fantasy.ToolResultContent](block)
if !ok || toolResult.ToolName != providerRunnerName {
continue
}
textResult, ok := toolResult.Result.(fantasy.ToolResultOutputContentText)
require.True(t, ok)
assert.Equal(t, "ran provider runner", textResult.Text)
foundToolResult = true
}
require.True(t, foundToolResult,
"persisted step should include the provider runner result")
}
func TestProcessStepStream_AnthropicUsageMatchesFinalDelta(t *testing.T) {
t.Parallel()
@@ -921,6 +1065,144 @@ func TestRun_MultiStepToolExecution(t *testing.T) {
"tool-result timestamp must be >= tool-call timestamp")
}
func TestStopAfterTool_Success(t *testing.T) {
t.Parallel()
streamCalls := 0
model := &chattest.FakeModel{
ProviderName: "fake",
StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) {
streamCalls++
return streamFromParts([]fantasy.StreamPart{
{Type: fantasy.StreamPartTypeToolInputStart, ID: "tc-plan", ToolCallName: "propose_plan"},
{Type: fantasy.StreamPartTypeToolInputDelta, ID: "tc-plan", Delta: `{"path":"/tmp/plan.md"}`},
{Type: fantasy.StreamPartTypeToolInputEnd, ID: "tc-plan"},
{
Type: fantasy.StreamPartTypeToolCall,
ID: "tc-plan",
ToolCallName: "propose_plan",
ToolCallInput: `{"path":"/tmp/plan.md"}`,
},
{Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls},
}), nil
},
}
proposePlanTool := fantasy.NewAgentTool(
"propose_plan",
"writes a plan",
func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) {
return fantasy.NewTextResponse("plan saved"), nil
},
)
var persistedSteps []PersistedStep
persistStepCalls := 0
err := Run(context.Background(), RunOptions{
Model: model,
Messages: []fantasy.Message{
textMessage(fantasy.MessageRoleUser, "propose a plan"),
},
Tools: []fantasy.AgentTool{proposePlanTool},
MaxSteps: 5,
StopAfterTools: map[string]struct{}{
"propose_plan": {},
},
PersistStep: func(_ context.Context, step PersistedStep) error {
persistStepCalls++
persistedSteps = append(persistedSteps, step)
return nil
},
})
require.ErrorIs(t, err, ErrStopAfterTool)
require.Equal(t, 1, streamCalls)
require.Equal(t, 1, persistStepCalls)
require.Len(t, persistedSteps, 1)
var foundToolResult bool
for _, block := range persistedSteps[0].Content {
toolResult, ok := fantasy.AsContentType[fantasy.ToolResultContent](block)
if !ok || toolResult.ToolName != "propose_plan" {
continue
}
foundToolResult = true
_, isErr := toolResult.Result.(fantasy.ToolResultOutputContentError)
require.False(t, isErr, "stop-after-tool should only trigger on successful tool results")
}
require.True(t, foundToolResult, "persisted step should include the successful tool result before stopping")
}
func TestStopAfterTool_IgnoresErrorResults(t *testing.T) {
t.Parallel()
streamCalls := 0
model := &chattest.FakeModel{
ProviderName: "fake",
StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) {
streamCalls++
if streamCalls == 1 {
return streamFromParts([]fantasy.StreamPart{
{Type: fantasy.StreamPartTypeToolInputStart, ID: "tc-plan", ToolCallName: "propose_plan"},
{Type: fantasy.StreamPartTypeToolInputDelta, ID: "tc-plan", Delta: `{"path":"/tmp/plan.md"}`},
{Type: fantasy.StreamPartTypeToolInputEnd, ID: "tc-plan"},
{
Type: fantasy.StreamPartTypeToolCall,
ID: "tc-plan",
ToolCallName: "propose_plan",
ToolCallInput: `{"path":"/tmp/plan.md"}`,
},
{Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls},
}), nil
}
return streamFromParts([]fantasy.StreamPart{
{Type: fantasy.StreamPartTypeTextStart, ID: "text-1"},
{Type: fantasy.StreamPartTypeTextDelta, ID: "text-1", Delta: "tool failed, continue"},
{Type: fantasy.StreamPartTypeTextEnd, ID: "text-1"},
{Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonStop},
}), nil
},
}
proposePlanTool := fantasy.NewAgentTool(
"propose_plan",
"writes a plan",
func(context.Context, struct{}, fantasy.ToolCall) (fantasy.ToolResponse, error) {
return fantasy.NewTextErrorResponse("plan failed"), nil
},
)
var persistedSteps []PersistedStep
err := Run(context.Background(), RunOptions{
Model: model,
Messages: []fantasy.Message{
textMessage(fantasy.MessageRoleUser, "propose a plan"),
},
Tools: []fantasy.AgentTool{proposePlanTool},
MaxSteps: 5,
StopAfterTools: map[string]struct{}{
"propose_plan": {},
},
PersistStep: func(_ context.Context, step PersistedStep) error {
persistedSteps = append(persistedSteps, step)
return nil
},
})
require.NoError(t, err)
require.Equal(t, 2, streamCalls)
require.Len(t, persistedSteps, 2)
var foundToolError bool
for _, block := range persistedSteps[0].Content {
toolResult, ok := fantasy.AsContentType[fantasy.ToolResultContent](block)
if !ok || toolResult.ToolName != "propose_plan" {
continue
}
_, foundToolError = toolResult.Result.(fantasy.ToolResultOutputContentError)
}
require.True(t, foundToolError, "first step should persist the failed tool result")
}
func TestRun_ParallelToolExecutionTimestamps(t *testing.T) {
t.Parallel()