mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user