mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix(chatd): block subagents from spawning workspaces (#22603)
## Summary Subagent (child) chats were previously given access to workspace provisioning tools (`list_templates`, `read_template`, `create_workspace`), which could lead to uncontrolled resource consumption. This PR moves those tools behind the same `!chat.ParentChatID.Valid` gate that already protects the subagent tools (`spawn_agent`, `wait_agent`, etc.). ## Changes - **`coderd/chatd/chatd.go`**: Moved `list_templates`, `read_template`, and `create_workspace` tool registration into the root-chat-only block alongside subagent tools. - **`coderd/chatd/chatd_test.go`**: Added `TestSubagentChatExcludesWorkspaceProvisioningTools` — an E2E test that spawns a subagent via a root chat and verifies the subagent's LLM call does not include workspace provisioning or subagent tools. - **`coderd/chatd/chattest/openai.go`**: Added `Tools` field to `OpenAIRequest` and supporting `OpenAITool`/`OpenAIToolFunction` types so tests can inspect which tools are sent to the model.
This commit is contained in:
+21
-18
@@ -2177,22 +2177,6 @@ func (p *Server) runChat(
|
||||
|
||||
// Here are all the tools we have for the chat.
|
||||
tools := []fantasy.AgentTool{
|
||||
chattool.ListTemplates(chattool.ListTemplatesOptions{
|
||||
DB: p.db,
|
||||
OwnerID: chat.OwnerID,
|
||||
}),
|
||||
chattool.ReadTemplate(chattool.ReadTemplateOptions{
|
||||
DB: p.db,
|
||||
OwnerID: chat.OwnerID,
|
||||
}),
|
||||
chattool.CreateWorkspace(chattool.CreateWorkspaceOptions{
|
||||
DB: p.db,
|
||||
OwnerID: chat.OwnerID,
|
||||
ChatID: chat.ID,
|
||||
CreateFn: p.createWorkspaceFn,
|
||||
AgentConnFn: chattool.AgentConnFunc(p.agentConnFn),
|
||||
WorkspaceMu: &workspaceMu,
|
||||
}),
|
||||
chattool.ReadFile(chattool.ReadFileOptions{
|
||||
GetWorkspaceConn: getWorkspaceConn,
|
||||
}),
|
||||
@@ -2216,10 +2200,29 @@ func (p *Server) runChat(
|
||||
GetWorkspaceConn: getWorkspaceConn,
|
||||
}),
|
||||
}
|
||||
// Only root chats (not delegated subagents) get subagent tools.
|
||||
// Child agents must not spawn further subagents — they should
|
||||
// Only root chats (not delegated subagents) get workspace
|
||||
// provisioning and subagent tools. Child agents must not
|
||||
// create workspaces or spawn further subagents — they should
|
||||
// focus on completing their delegated task.
|
||||
if !chat.ParentChatID.Valid {
|
||||
tools = append(tools,
|
||||
chattool.ListTemplates(chattool.ListTemplatesOptions{
|
||||
DB: p.db,
|
||||
OwnerID: chat.OwnerID,
|
||||
}),
|
||||
chattool.ReadTemplate(chattool.ReadTemplateOptions{
|
||||
DB: p.db,
|
||||
OwnerID: chat.OwnerID,
|
||||
}),
|
||||
chattool.CreateWorkspace(chattool.CreateWorkspaceOptions{
|
||||
DB: p.db,
|
||||
OwnerID: chat.OwnerID,
|
||||
ChatID: chat.ID,
|
||||
CreateFn: p.createWorkspaceFn,
|
||||
AgentConnFn: chattool.AgentConnFunc(p.agentConnFn),
|
||||
WorkspaceMu: &workspaceMu,
|
||||
}),
|
||||
)
|
||||
tools = append(tools, p.subagentTools(func() database.Chat {
|
||||
return chat
|
||||
})...)
|
||||
|
||||
@@ -84,6 +84,160 @@ func TestInterruptChatBroadcastsStatusAcrossInstances(t *testing.T) {
|
||||
}, testutil.WaitMedium, testutil.IntervalFast)
|
||||
}
|
||||
|
||||
func TestSubagentChatExcludesWorkspaceProvisioningTools(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
deploymentValues := coderdtest.DeploymentValues(t)
|
||||
deploymentValues.Experiments = []string{string(codersdk.ExperimentAgents)}
|
||||
client := coderdtest.New(t, &coderdtest.Options{
|
||||
DeploymentValues: deploymentValues,
|
||||
IncludeProvisionerDaemon: true,
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
agentToken := uuid.NewString()
|
||||
version := coderdtest.CreateTemplateVersion(t, client, user.OrganizationID, &echo.Responses{
|
||||
Parse: echo.ParseComplete,
|
||||
ProvisionPlan: echo.PlanComplete,
|
||||
ProvisionApply: echo.ApplyComplete,
|
||||
ProvisionGraph: echo.ProvisionGraphWithAgent(agentToken),
|
||||
})
|
||||
coderdtest.AwaitTemplateVersionJobCompleted(t, client, version.ID)
|
||||
coderdtest.CreateTemplate(t, client, user.OrganizationID, version.ID)
|
||||
|
||||
_ = agenttest.New(t, client.URL, agentToken)
|
||||
|
||||
// Track tools sent in LLM requests. The first call is for the
|
||||
// root chat which spawns a subagent; the second call is for the
|
||||
// subagent itself.
|
||||
var toolsMu sync.Mutex
|
||||
toolsByCall := make([][]string, 0, 2)
|
||||
|
||||
var callCount atomic.Int32
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("ok")
|
||||
}
|
||||
|
||||
names := make([]string, 0, len(req.Tools))
|
||||
for _, tool := range req.Tools {
|
||||
names = append(names, tool.Function.Name)
|
||||
}
|
||||
toolsMu.Lock()
|
||||
toolsByCall = append(toolsByCall, names)
|
||||
toolsMu.Unlock()
|
||||
|
||||
if callCount.Add(1) == 1 {
|
||||
// Root chat: model calls spawn_agent.
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAIToolCallChunk("spawn_agent", `{"prompt":"do the thing","title":"sub"}`),
|
||||
)
|
||||
}
|
||||
// Subsequent calls (including the subagent): just reply.
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("Done.")...,
|
||||
)
|
||||
})
|
||||
|
||||
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
|
||||
Provider: "openai-compat",
|
||||
APIKey: "test-api-key",
|
||||
BaseURL: openAIURL,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
contextLimit := int64(4096)
|
||||
isDefault := true
|
||||
_, err = client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai-compat",
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
IsDefault: &isDefault,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create a root chat whose first model call will spawn a subagent.
|
||||
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
Content: []codersdk.ChatInputPart{
|
||||
{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "Spawn a subagent to do the thing.",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Wait for the root chat AND the subagent to finish.
|
||||
// The root chat finishes first, then the chatd server
|
||||
// picks up and runs the child (subagent) chat.
|
||||
require.Eventually(t, func() bool {
|
||||
got, getErr := client.GetChat(ctx, chat.ID)
|
||||
if getErr != nil {
|
||||
return false
|
||||
}
|
||||
if got.Chat.Status != codersdk.ChatStatusWaiting && got.Chat.Status != codersdk.ChatStatusError {
|
||||
return false
|
||||
}
|
||||
// Also ensure the subagent LLM call has been made.
|
||||
toolsMu.Lock()
|
||||
n := len(toolsByCall)
|
||||
toolsMu.Unlock()
|
||||
// Expect at least 3 calls: root-1 (spawn_agent), child-1, root-2.
|
||||
return n >= 3
|
||||
}, testutil.WaitLong, testutil.IntervalFast)
|
||||
|
||||
// There should be at least two streamed calls: one for the root
|
||||
// chat and one for the subagent child chat.
|
||||
toolsMu.Lock()
|
||||
recorded := append([][]string(nil), toolsByCall...)
|
||||
toolsMu.Unlock()
|
||||
|
||||
require.GreaterOrEqual(t, len(recorded), 2,
|
||||
"expected at least 2 streamed LLM calls (root + subagent)")
|
||||
|
||||
workspaceTools := []string{"list_templates", "read_template", "create_workspace"}
|
||||
subagentTools := []string{"spawn_agent", "wait_agent", "message_agent", "close_agent"}
|
||||
|
||||
// Identify root and subagent calls. Root chat calls include
|
||||
// spawn_agent; the subagent call does not. Because the root chat
|
||||
// makes multiple LLM calls (before and after spawn_agent), we
|
||||
// find exactly one call that lacks spawn_agent — that's the
|
||||
// subagent.
|
||||
var rootCalls, childCalls [][]string
|
||||
for _, tools := range recorded {
|
||||
hasSpawnAgent := slice.Contains(tools, "spawn_agent")
|
||||
if hasSpawnAgent {
|
||||
rootCalls = append(rootCalls, tools)
|
||||
} else {
|
||||
childCalls = append(childCalls, tools)
|
||||
}
|
||||
}
|
||||
|
||||
require.NotEmpty(t, rootCalls, "expected at least one root chat LLM call")
|
||||
require.NotEmpty(t, childCalls, "expected at least one subagent LLM call")
|
||||
|
||||
// Root chat calls must include workspace and subagent tools.
|
||||
for _, tool := range workspaceTools {
|
||||
require.Contains(t, rootCalls[0], tool,
|
||||
"root chat should have workspace tool %q", tool)
|
||||
}
|
||||
for _, tool := range subagentTools {
|
||||
require.Contains(t, rootCalls[0], tool,
|
||||
"root chat should have subagent tool %q", tool)
|
||||
}
|
||||
|
||||
// Subagent calls must NOT include workspace or subagent tools.
|
||||
for _, tool := range workspaceTools {
|
||||
require.NotContains(t, childCalls[0], tool,
|
||||
"subagent chat should NOT have workspace tool %q", tool)
|
||||
}
|
||||
for _, tool := range subagentTools {
|
||||
require.NotContains(t, childCalls[0], tool,
|
||||
"subagent chat should NOT have subagent tool %q", tool)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInterruptChatClearsWorkerInDatabase(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -30,6 +30,7 @@ type OpenAIRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []OpenAIMessage `json:"messages"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
Tools []OpenAITool `json:"tools,omitempty"`
|
||||
Prompt []interface{} `json:"prompt,omitempty"` // For responses API
|
||||
// TODO: encoding/json ignores inline tags. Add custom UnmarshalJSON to capture unknown keys.
|
||||
Options map[string]interface{} `json:",inline"` //nolint:revive
|
||||
@@ -41,6 +42,17 @@ type OpenAIMessage struct {
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// OpenAIToolFunction represents the function definition inside a tool.
|
||||
type OpenAIToolFunction struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
// OpenAITool represents a tool definition in an OpenAI request.
|
||||
type OpenAITool struct {
|
||||
Type string `json:"type"`
|
||||
Function OpenAIToolFunction `json:"function"`
|
||||
}
|
||||
|
||||
// OpenAIToolCallFunction represents the function details in a tool call.
|
||||
type OpenAIToolCallFunction struct {
|
||||
Name string `json:"name,omitempty"`
|
||||
|
||||
Reference in New Issue
Block a user