From 83de37b579005ce4799b71c97fa527ecf85e1293 Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Sat, 29 Aug 2026 09:23:49 +0000 Subject: [PATCH 1/4] refactor(coderd/x/chatd): deepen turn environment preparation --- coderd/x/chatd/chatd.go | 1199 --------- coderd/x/chatd/generation.go | 171 +- coderd/x/chatd/generation_preparer.go | 926 ------- .../generation_preparer_internal_test.go | 34 +- coderd/x/chatd/toolinput.go | 10 +- coderd/x/chatd/toolinput_internal_test.go | 18 +- coderd/x/chatd/turn_environment.go | 2226 +++++++++++++++++ 7 files changed, 2324 insertions(+), 2260 deletions(-) delete mode 100644 coderd/x/chatd/generation_preparer.go create mode 100644 coderd/x/chatd/turn_environment.go diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 456031add59..b1bc24137cb 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -7,7 +7,6 @@ import ( "encoding/json" "errors" "fmt" - "net/http" "slices" "strconv" "strings" @@ -16,7 +15,6 @@ import ( "time" "charm.land/fantasy" - "charm.land/fantasy/providers/anthropic" "github.com/dustin/go-humanize" "github.com/google/uuid" "github.com/prometheus/client_golang/prometheus" @@ -39,20 +37,17 @@ import ( "github.com/coder/coder/v2/coderd/webpush" "github.com/coder/coder/v2/coderd/workspacestats" "github.com/coder/coder/v2/coderd/x/agenthooks/dispatch" - "github.com/coder/coder/v2/coderd/x/chatd/agentselect" "github.com/coder/coder/v2/coderd/x/chatd/chatadvisor" "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" "github.com/coder/coder/v2/coderd/x/chatd/chaterror" "github.com/coder/coder/v2/coderd/x/chatd/chathooks" "github.com/coder/coder/v2/coderd/x/chatd/chatloop" - "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" "github.com/coder/coder/v2/coderd/x/chatd/chatstate" "github.com/coder/coder/v2/coderd/x/chatd/chattool" "github.com/coder/coder/v2/coderd/x/chatd/mcpclient" "github.com/coder/coder/v2/coderd/x/chatd/messagepartbuffer" - skillspkg "github.com/coder/coder/v2/coderd/x/skills" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/codersdk/workspacesdk" "github.com/coder/coder/v2/codersdk/x/agenthooks" @@ -460,632 +455,6 @@ func (p *Server) pinnedWorkspaceMCPTools( return chattool.NewWorkspaceMCPTools(infos, getConn, nil), nil } -type turnWorkspaceContext struct { - server *Server - chatStateMu *sync.Mutex - currentChat *database.Chat - loadChatSnapshot func(context.Context, uuid.UUID) (database.Chat, error) - - mu sync.Mutex - agent database.WorkspaceAgent - agentLoaded bool - conn workspacesdk.AgentConn - releaseConn func() - cachedWorkspaceID uuid.NullUUID -} - -func (c *turnWorkspaceContext) close() { - c.clearCachedWorkspaceState() -} - -func (c *turnWorkspaceContext) clearCachedWorkspaceState() { - c.mu.Lock() - releaseConn := c.releaseConn - c.agent = database.WorkspaceAgent{} - c.agentLoaded = false - c.conn = nil - c.releaseConn = nil - c.cachedWorkspaceID = uuid.NullUUID{} - c.mu.Unlock() - - if releaseConn != nil { - releaseConn() - } -} - -func (c *turnWorkspaceContext) setCurrentChat(chat database.Chat) { - c.chatStateMu.Lock() - *c.currentChat = chat - c.chatStateMu.Unlock() -} - -func (c *turnWorkspaceContext) currentChatSnapshot() database.Chat { - c.chatStateMu.Lock() - chatSnapshot := *c.currentChat - c.chatStateMu.Unlock() - return chatSnapshot -} - -func (c *turnWorkspaceContext) selectWorkspace(chat database.Chat) { - c.setCurrentChat(chat) - c.clearCachedWorkspaceState() -} - -func (c *turnWorkspaceContext) currentWorkspaceMatches(expected uuid.NullUUID) (database.Chat, bool) { - chatSnapshot := c.currentChatSnapshot() - return chatSnapshot, nullUUIDEqual(chatSnapshot.WorkspaceID, expected) -} - -func (c *turnWorkspaceContext) trackWorkspaceUsage(ctx context.Context, chatSnapshot database.Chat) { - if c.server == nil || !chatSnapshot.WorkspaceID.Valid { - return - } - logger := c.server.logger.With( - slog.F("chat_id", chatSnapshot.ID), - slog.F("owner_id", chatSnapshot.OwnerID), - ) - c.server.trackWorkspaceUsage(ctx, chatSnapshot.ID, chatSnapshot.WorkspaceID, logger) -} - -func nullUUIDEqual(left, right uuid.NullUUID) bool { - if left.Valid != right.Valid { - return false - } - if !left.Valid { - return true - } - return left.UUID == right.UUID -} - -func (c *turnWorkspaceContext) persistBuildAgentBinding( - ctx context.Context, - chatSnapshot database.Chat, - buildID uuid.UUID, - agentID uuid.UUID, -) (database.Chat, error) { - updatedChat, err := c.server.db.UpdateChatBuildAgentBinding( - ctx, - database.UpdateChatBuildAgentBindingParams{ - ID: chatSnapshot.ID, - BuildID: uuid.NullUUID{ - UUID: buildID, - Valid: true, - }, - AgentID: uuid.NullUUID{ - UUID: agentID, - Valid: true, - }, - }, - ) - if err != nil { - return chatSnapshot, xerrors.Errorf( - "update chat build/agent binding: %w", err, - ) - } - - // If the chat was rebound to a different agent (e.g. a workspace rebuild - // produced a new agent), re-pin its context to the new agent so it stops - // injecting the previous agent's resources. Workspace lifecycle tools clear - // the agent binding while preserving the pin, so a missing prior agent also - // requires a re-pin when pinned context exists. Best-effort: a context error - // must never fail the binding. The pinned context fields on updatedChat are - // background state, reloaded on the next snapshot fetch. - hasStaleUnboundContext := !chatSnapshot.AgentID.Valid && chatSnapshot.ContextAggregateHash != nil - if hasStaleUnboundContext || (chatSnapshot.AgentID.Valid && chatSnapshot.AgentID.UUID != agentID) { - //nolint:gocritic // Chatd re-pins chats it does not own as the daemon subject. - repinCtx := dbauthz.AsChatd(ctx) - if repinErr := database.ReadModifyUpdate(c.server.db, func(tx database.Store) error { - return repinChatContext(repinCtx, tx, chatSnapshot.ID, uuid.NullUUID{UUID: agentID, Valid: true}) - }); repinErr != nil { - c.server.logger.Warn(ctx, "re-pin chat context after agent rebind", - slog.F("chat_id", chatSnapshot.ID), - slog.F("agent_id", agentID), - slog.Error(repinErr)) - } - } - - c.setCurrentChat(updatedChat) - return updatedChat, nil -} - -func (c *turnWorkspaceContext) getWorkspaceAgent(ctx context.Context) (database.WorkspaceAgent, error) { - _, agent, err := c.ensureWorkspaceAgent(ctx) - return agent, err -} - -func (c *turnWorkspaceContext) ensureWorkspaceAgent( - ctx context.Context, -) (database.Chat, database.WorkspaceAgent, error) { - c.mu.Lock() - defer c.mu.Unlock() - - if c.agentLoaded { - chatSnapshot := c.currentChatSnapshot() - if nullUUIDEqual(c.cachedWorkspaceID, chatSnapshot.WorkspaceID) { - return chatSnapshot, c.agent, nil - } - c.agent = database.WorkspaceAgent{} - c.agentLoaded = false - } - - return c.loadWorkspaceAgentLocked(ctx) -} - -func (c *turnWorkspaceContext) loadWorkspaceAgentLocked( - ctx context.Context, -) (database.Chat, database.WorkspaceAgent, error) { - chatSnapshot := c.currentChatSnapshot() - - for attempt := 0; attempt < 2; attempt++ { - if !chatSnapshot.WorkspaceID.Valid { - refreshedChat, refreshErr := refreshChatWorkspaceSnapshot( - ctx, - chatSnapshot, - c.loadChatSnapshot, - ) - if refreshErr != nil { - return chatSnapshot, database.WorkspaceAgent{}, refreshErr - } - if refreshedChat.WorkspaceID.Valid { - c.setCurrentChat(refreshedChat) - chatSnapshot = refreshedChat - } - } - - if !chatSnapshot.WorkspaceID.Valid { - return chatSnapshot, database.WorkspaceAgent{}, xerrors.New("no workspace is associated with this chat. Use the create_workspace tool to create one") - } - - if chatSnapshot.AgentID.Valid { - agent, err := c.server.db.GetWorkspaceAgentByID(ctx, chatSnapshot.AgentID.UUID) - if err == nil { - latestChat, workspaceMatches := c.currentWorkspaceMatches(chatSnapshot.WorkspaceID) - if !workspaceMatches { - chatSnapshot = latestChat - continue - } - c.agent = agent - c.agentLoaded = true - c.cachedWorkspaceID = chatSnapshot.WorkspaceID - return chatSnapshot, c.agent, nil - } - if !xerrors.Is(err, sql.ErrNoRows) { - c.server.logger.Warn(ctx, "agent binding lookup failed, re-resolving", - slog.F("agent_id", chatSnapshot.AgentID.UUID), - slog.Error(err), - ) - } - } - - agents, err := c.server.db.GetWorkspaceAgentsInLatestBuildByWorkspaceID( - ctx, - chatSnapshot.WorkspaceID.UUID, - ) - if err != nil { - return chatSnapshot, database.WorkspaceAgent{}, xerrors.Errorf( - "get workspace agents in latest build: %w", - err, - ) - } - if len(agents) == 0 { - return chatSnapshot, database.WorkspaceAgent{}, errChatHasNoWorkspaceAgent - } - selected, err := agentselect.FindChatAgent(agents) - if err != nil { - return chatSnapshot, database.WorkspaceAgent{}, xerrors.Errorf( - "find chat agent: %w", - err, - ) - } - - build, err := c.server.db.GetLatestWorkspaceBuildByWorkspaceID(ctx, chatSnapshot.WorkspaceID.UUID) - if err != nil { - return chatSnapshot, database.WorkspaceAgent{}, xerrors.Errorf("get latest workspace build: %w", err) - } - - updatedChat, err := c.persistBuildAgentBinding( - ctx, - chatSnapshot, - build.ID, - selected.ID, - ) - if err != nil { - return chatSnapshot, database.WorkspaceAgent{}, err - } - - chatSnapshot = updatedChat - latestChat, workspaceMatches := c.currentWorkspaceMatches(chatSnapshot.WorkspaceID) - if !workspaceMatches { - chatSnapshot = latestChat - continue - } - c.agent = selected - c.agentLoaded = true - c.cachedWorkspaceID = chatSnapshot.WorkspaceID - return chatSnapshot, c.agent, nil - } - - return chatSnapshot, database.WorkspaceAgent{}, xerrors.New( - "chat workspace changed while resolving agent", - ) -} - -func (c *turnWorkspaceContext) latestWorkspaceAgentID( - ctx context.Context, - workspaceID uuid.UUID, -) (uuid.UUID, error) { - agents, err := c.server.db.GetWorkspaceAgentsInLatestBuildByWorkspaceID( - ctx, - workspaceID, - ) - if err != nil { - return uuid.Nil, xerrors.Errorf( - "get workspace agents in latest build: %w", - err, - ) - } - if len(agents) == 0 { - return uuid.Nil, errChatHasNoWorkspaceAgent - } - selected, err := agentselect.FindChatAgent(agents) - if err != nil { - return uuid.Nil, xerrors.Errorf( - "find chat agent: %w", - err, - ) - } - return selected.ID, nil -} - -func (c *turnWorkspaceContext) workspaceAgentIDForConn( - ctx context.Context, -) (database.Chat, uuid.UUID, error) { - for attempt := 0; attempt < 2; attempt++ { - chatSnapshot := c.currentChatSnapshot() - if !chatSnapshot.WorkspaceID.Valid || !chatSnapshot.AgentID.Valid { - updatedChat, agent, err := c.ensureWorkspaceAgent(ctx) - if err != nil { - return updatedChat, uuid.Nil, err - } - return updatedChat, agent.ID, nil - } - - currentAgentID, err := c.latestWorkspaceAgentID( - ctx, - chatSnapshot.WorkspaceID.UUID, - ) - if err != nil { - if xerrors.Is(err, errChatHasNoWorkspaceAgent) { - c.clearCachedWorkspaceState() - } - return chatSnapshot, uuid.Nil, err - } - - latestChat, workspaceMatches := c.currentWorkspaceMatches( - chatSnapshot.WorkspaceID, - ) - if !workspaceMatches { - continue - } - return latestChat, currentAgentID, nil - } - - chatSnapshot := c.currentChatSnapshot() - return chatSnapshot, uuid.Nil, xerrors.New( - "chat workspace changed while resolving agent", - ) -} - -// getWorkspaceConnLocked returns the cached connection when it still matches -// the current workspace. When the workspace changed, it clears the stale -// cached state and returns the release func for the caller to run after -// unlocking. -func (c *turnWorkspaceContext) getWorkspaceConnLocked() (workspacesdk.AgentConn, func()) { - if c.conn == nil { - return nil, nil - } - - chatSnapshot := c.currentChatSnapshot() - if nullUUIDEqual(c.cachedWorkspaceID, chatSnapshot.WorkspaceID) { - return c.conn, nil - } - - agentRelease := c.releaseConn - c.agent = database.WorkspaceAgent{} - c.agentLoaded = false - c.conn = nil - c.releaseConn = nil - c.cachedWorkspaceID = uuid.NullUUID{} - return nil, agentRelease -} - -// isAgentUnreachable reports whether the given agent row's -// status is disconnected or timed out. It uses timestamp -// arithmetic on the row. The "connecting" state is allowed -// through because it is normal after a fresh workspace build. -func isAgentUnreachable(now time.Time, agent database.WorkspaceAgent, inactiveTimeout time.Duration) bool { - status := agent.Status(now, inactiveTimeout) - return status.Status == database.WorkspaceAgentStatusDisconnected || - status.Status == database.WorkspaceAgentStatusTimeout -} - -func agentDisconnectedFor(now time.Time, agent database.WorkspaceAgent, inactiveTimeout time.Duration) (time.Duration, bool) { - status := agent.Status(now, inactiveTimeout) - if status.Status != database.WorkspaceAgentStatusDisconnected || status.DisconnectedAt == nil { - return 0, false - } - - disconnectedFor := now.Sub(*status.DisconnectedAt) - if disconnectedFor < 0 { - disconnectedFor = 0 - } - return disconnectedFor, true -} - -func (c *turnWorkspaceContext) latestWorkspaceAgentRecoveryError( - ctx context.Context, - workspaceID uuid.UUID, -) error { - agentID, err := c.latestWorkspaceAgentID(ctx, workspaceID) - if err != nil { - if xerrors.Is(err, errChatHasNoWorkspaceAgent) { - return err - } - c.server.logger.Warn(ctx, "failed to resolve latest agent for timeout classification", slog.Error(err)) - return errChatDialTimeout - } - - agent, err := c.server.db.GetWorkspaceAgentByID(ctx, agentID) - if err != nil { - c.server.logger.Warn(ctx, "failed to load latest agent for timeout classification", - slog.F("agent_id", agentID), - slog.Error(err), - ) - return errChatDialTimeout - } - - now := c.server.clock.Now() - status := agent.Status(now, c.server.agentInactiveDisconnectTimeout) - recoveryErr := errChatDialTimeout - if status.Status == database.WorkspaceAgentStatusTimeout { - recoveryErr = errChatAgentNeverConnected - } else if status.Status == database.WorkspaceAgentStatusDisconnected && status.DisconnectedAt != nil { - disconnectedFor := now.Sub(*status.DisconnectedAt) - if disconnectedFor < 0 { - disconnectedFor = 0 - } - if disconnectedFor >= agentDisconnectedRecoveryThreshold { - recoveryErr = errChatAgentDisconnected - } - } - return c.externalAgentError(ctx, agent, recoveryErr) -} - -func (c *turnWorkspaceContext) externalAgentError( - ctx context.Context, - agent database.WorkspaceAgent, - fallback error, -) error { - isExternal, err := chattool.IsExternalWorkspaceAgent(ctx, c.server.db, agent) - if err != nil || !isExternal { - return fallback - } - return newChatExternalAgentUnavailableError(agent) -} - -func (c *turnWorkspaceContext) externalAgentPreflightError( - ctx context.Context, - chatSnapshot database.Chat, - agent database.WorkspaceAgent, -) error { - // Mirror the cache-hit gate: only short-circuit on clearly offline - // states (Disconnected/Timeout). Connecting is allowed through so - // an external agent the user just started can still connect inside - // the normal dial window. - if !isAgentUnreachable(c.server.clock.Now(), agent, c.server.agentInactiveDisconnectTimeout) { - return nil - } - - isExternal, err := chattool.IsExternalWorkspaceAgent(ctx, c.server.db, agent) - if err != nil || !isExternal || !chatSnapshot.WorkspaceID.Valid { - return nil - } - - // Stale agent bindings rely on dialWithLazyValidation to discover - // replacement agents, so only skip the dial when this agent is still - // the latest selected chat agent for the workspace. - latestAgentID, err := c.latestWorkspaceAgentID(ctx, chatSnapshot.WorkspaceID.UUID) - if err != nil || latestAgentID != agent.ID { - return nil - } - return newChatExternalAgentUnavailableError(agent) -} - -func (c *turnWorkspaceContext) getWorkspaceConn(ctx context.Context) (workspacesdk.AgentConn, error) { - if c.server.agentConnFn == nil { - return nil, xerrors.New("workspace agent connector is not configured") - } - - for attempt := 0; attempt < 2; attempt++ { - c.mu.Lock() - currentConn, staleRelease := c.getWorkspaceConnLocked() - // Capture agentID in the same lock section as - // currentConn to prevent a TOCTOU race with - // concurrent clearCachedWorkspaceState calls. - agentID := c.agent.ID - c.mu.Unlock() - - // Status check on cache hit: re-fetch the agent - // row so we see the latest heartbeat rather than - // a potentially stale cached copy. - if currentConn != nil { - chatSnapshot := c.currentChatSnapshot() - if agentID != uuid.Nil { - freshAgent, err := c.server.db.GetWorkspaceAgentByID(ctx, agentID) - if err != nil { - c.server.logger.Warn(ctx, "failed to re-fetch agent for status check", - slog.F("agent_id", agentID), - slog.Error(err), - ) - // On DB error the check re-runs on the - // next tool call. - } else if _, disconnected := agentDisconnectedFor( - c.server.clock.Now(), - freshAgent, - c.server.agentInactiveDisconnectTimeout, - ); disconnected { - c.clearCachedWorkspaceState() - continue - } - } - c.trackWorkspaceUsage(ctx, chatSnapshot) - return currentConn, nil - } - if staleRelease != nil { - staleRelease() - } - - chatSnapshot, agent, err := c.ensureWorkspaceAgent(ctx) - if err != nil { - return nil, err - } - if err := c.externalAgentPreflightError(ctx, chatSnapshot, agent); err != nil { - return nil, err - } - - // Wrap the dial in a timeout to bound the time spent - // waiting for an unreachable agent. The timeout scopes - // only dialWithLazyValidation, not ensureWorkspaceAgent - // or the post-dial binding steps. - dialCtx, dialCancelCause := context.WithCancelCause(ctx) - dialTimer := c.server.clock.AfterFunc( - c.server.dialTimeout, - func() { dialCancelCause(errChatDialTimeout) }, - "chatd", - dialTimeoutTimerTag, - ) - dialCancel := func() { - dialTimer.Stop() - dialCancelCause(nil) - } - dialResult, err := dialWithLazyValidation( - dialCtx, - c.server.clock, - agent.ID, - chatSnapshot.WorkspaceID.UUID, - DialFunc(c.server.agentConnFn), - func(ctx context.Context, workspaceID uuid.UUID) (uuid.UUID, error) { - return c.latestWorkspaceAgentID(ctx, workspaceID) - }, - workspaceDialValidationDelay, - ) - dialCancel() - if err != nil { - if xerrors.Is(err, errChatHasNoWorkspaceAgent) { - c.clearCachedWorkspaceState() - return nil, err - } - // Surface the dial timeout sentinel only when the - // parent context is still alive. If the parent was - // canceled (e.g. ErrInterrupted), its error must - // propagate unchanged so the chatloop can detect it. - if ctx.Err() == nil && errors.Is(context.Cause(dialCtx), errChatDialTimeout) { - c.clearCachedWorkspaceState() - return nil, c.latestWorkspaceAgentRecoveryError(ctx, chatSnapshot.WorkspaceID.UUID) - } - return nil, err - } - agentConn := dialResult.Conn - agentRelease := dialResult.Release - if dialResult.WasSwitched { - build, err := c.server.db.GetLatestWorkspaceBuildByWorkspaceID(ctx, chatSnapshot.WorkspaceID.UUID) - if err != nil { - if agentRelease != nil { - agentRelease() - } - return nil, xerrors.Errorf("get latest workspace build: %w", err) - } - - switchedAgent, err := c.server.db.GetWorkspaceAgentByID(ctx, dialResult.AgentID) - if err != nil { - if agentRelease != nil { - agentRelease() - } - return nil, xerrors.Errorf("get workspace agent by id: %w", err) - } - - updatedChat, err := c.persistBuildAgentBinding( - ctx, - chatSnapshot, - build.ID, - switchedAgent.ID, - ) - if err != nil { - if agentRelease != nil { - agentRelease() - } - return nil, err - } - chatSnapshot = updatedChat - - c.mu.Lock() - c.agent = switchedAgent - c.agentLoaded = true - c.cachedWorkspaceID = chatSnapshot.WorkspaceID - c.mu.Unlock() - } - - if _, workspaceMatches := c.currentWorkspaceMatches(chatSnapshot.WorkspaceID); !workspaceMatches { - if agentRelease != nil { - agentRelease() - } - c.clearCachedWorkspaceState() - continue - } - - c.mu.Lock() - if c.conn == nil { - c.conn = agentConn - c.releaseConn = agentRelease - c.cachedWorkspaceID = chatSnapshot.WorkspaceID - - var ancestorIDs []string - if chatSnapshot.ParentChatID.Valid { - ancestorIDs = append(ancestorIDs, chatSnapshot.ParentChatID.UUID.String()) - } - ancestorJSON, marshalErr := json.Marshal(ancestorIDs) - if marshalErr != nil { - ancestorJSON = []byte("[]") - } - agentConn.SetExtraHeaders(http.Header{ - workspacesdk.CoderChatIDHeader: {chatSnapshot.ID.String()}, - workspacesdk.CoderAncestorChatIDsHeader: {string(ancestorJSON)}, - }) - - c.mu.Unlock() - c.server.logger.Debug(ctx, "set chat headers on agent conn", - slog.F("chat_id", chatSnapshot.ID), - slog.F("ancestor_chat_ids", ancestorIDs), - slog.F("workspace_id", chatSnapshot.WorkspaceID.UUID), - slog.F("agent_id", dialResult.AgentID), - ) - c.trackWorkspaceUsage(ctx, chatSnapshot) - return agentConn, nil - } - currentConn = c.conn - c.mu.Unlock() - - if agentRelease != nil { - agentRelease() - } - c.trackWorkspaceUsage(ctx, chatSnapshot) - return currentConn, nil - } - - return nil, xerrors.New("chat workspace changed while connecting") -} - -// AgentConnFunc provides access to workspace agent connections. type AgentConnFunc func(ctx context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) var ( @@ -3498,574 +2867,6 @@ type runChatResult struct { HistoryTipMessageID int64 } -func allToolNames(allTools []fantasy.AgentTool) []string { - toolNames := make([]string, 0, len(allTools)) - for _, tool := range allTools { - toolNames = append(toolNames, tool.Info().Name) - } - return toolNames -} - -func isExploreSubagentMode(mode database.NullChatMode) bool { - return mode.Valid && mode.ChatMode == database.ChatModeExplore -} - -// filterExternalMCPConfigsForTurn returns the external MCP server configs -// visible on the current turn. Explore children snapshot this filtered set at -// spawn time so later model overrides cannot widen the external-tool boundary. -func filterExternalMCPConfigsForTurn( - configs []database.MCPServerConfig, - mode database.NullChatPlanMode, - parentChatID uuid.NullUUID, -) ([]database.MCPServerConfig, map[uuid.UUID]struct{}) { - if !mode.Valid || mode.ChatPlanMode != database.ChatPlanModePlan { - return configs, nil - } - if parentChatID.Valid { - // Plan-mode subagents do not receive external MCP tools because - // their trust boundary is narrower than the root chat's. - return nil, map[uuid.UUID]struct{}{} - } - - filtered := make([]database.MCPServerConfig, 0, len(configs)) - approvedIDs := make(map[uuid.UUID]struct{}) - for _, cfg := range configs { - if !cfg.AllowInPlanMode { - continue - } - filtered = append(filtered, cfg) - approvedIDs[cfg.ID] = struct{}{} - } - return filtered, approvedIDs -} - -func builtinPlanToolAllowed(name string, isRootChat bool) bool { - switch name { - case "read_file", "execute", "process_output", "read_skill", "read_skill_file": - return true - case "write_file", "edit_files", "list_templates", "read_template", - "create_workspace", "start_workspace", "stop_workspace", "propose_plan", "spawn_agent", - "spawn_explore_agent", "wait_agent", "list_agents", "list_subagent_models", - "ask_user_question", "attach_file": - return isRootChat - case "process_list", "process_signal", "message_agent", "interrupt_agent", "close_agent", - "spawn_computer_use_agent": - return false - default: - return false - } -} - -func toolAllowedForTurn( - tool fantasy.AgentTool, - mode database.NullChatPlanMode, - parentChatID uuid.NullUUID, - approvedMCPConfigIDs map[uuid.UUID]struct{}, -) bool { - if !mode.Valid || mode.ChatPlanMode != database.ChatPlanModePlan { - return true - } - if builtinPlanToolAllowed(tool.Info().Name, !parentChatID.Valid) { - return true - } - mcpTool, ok := tool.(mcpclient.MCPToolIdentifier) - if !ok { - return false - } - _, approved := approvedMCPConfigIDs[mcpTool.MCPServerConfigID()] - return approved -} - -func filterToolsForTurn( - allTools []fantasy.AgentTool, - mode database.NullChatPlanMode, - parentChatID uuid.NullUUID, - approvedMCPConfigIDs map[uuid.UUID]struct{}, -) []fantasy.AgentTool { - if !mode.Valid || mode.ChatPlanMode != database.ChatPlanModePlan { - return allTools - } - - filtered := make([]fantasy.AgentTool, 0, len(allTools)) - for _, tool := range allTools { - if toolAllowedForTurn(tool, mode, parentChatID, approvedMCPConfigIDs) { - filtered = append(filtered, tool) - } - } - return filtered -} - -// activeToolNamesForTurn extends the built-in plan allowlist with approved -// external MCP tools for root plan-mode chats. -func activeToolNamesForTurn( - allTools []fantasy.AgentTool, - mode database.NullChatPlanMode, - parentChatID uuid.NullUUID, - approvedMCPConfigIDs map[uuid.UUID]struct{}, -) []string { - toolNames := make([]string, 0, len(allTools)) - for _, tool := range allTools { - if toolAllowedForTurn(tool, mode, parentChatID, approvedMCPConfigIDs) { - toolNames = append(toolNames, tool.Info().Name) - } - } - return toolNames -} - -func allowedExploreToolNames(allTools []fantasy.AgentTool) []string { - builtinExplorePolicy := map[string]bool{ - "read_file": true, - "write_file": false, - "edit_files": false, - "execute": true, - "process_output": true, - "process_list": false, - "process_signal": false, - "list_templates": false, - "read_template": false, - "create_workspace": false, - "start_workspace": false, - "stop_workspace": false, - "propose_plan": false, - "spawn_agent": false, - "wait_agent": false, - "message_agent": false, - "interrupt_agent": false, - "close_agent": false, - "list_agents": false, - "list_subagent_models": false, - "read_skill": true, - "read_skill_file": true, - "ask_user_question": false, - } - - toolNames := make([]string, 0, len(allTools)) - for _, tool := range allTools { - name := tool.Info().Name - if builtinExplorePolicy[name] { - toolNames = append(toolNames, name) - continue - } - // External MCP tools pass through here. They were snapshot-filtered - // at spawn time on chat.MCPServerIDs. WorkspaceMCPTool does not - // implement MCPToolIdentifier, so workspace tools are excluded - // here too, in addition to the structural exclusion in runChat - // tool assembly. - if _, ok := tool.(mcpclient.MCPToolIdentifier); ok { - toolNames = append(toolNames, name) - } - } - return toolNames -} - -// allowedBehaviorToolNames runs only on non-plan turns because -// appendDynamicTools returns early for plan mode. Within that boundary, -// Explore mode wins over the default behavior that allows all tools. -func allowedBehaviorToolNames( - allTools []fantasy.AgentTool, - chatMode database.NullChatMode, -) []string { - if isExploreSubagentMode(chatMode) { - return allowedExploreToolNames(allTools) - } - return allToolNames(allTools) -} - -func stopAfterPlanTools( - planMode database.NullChatPlanMode, - parentChatID uuid.NullUUID, -) map[string]struct{} { - if !planMode.Valid || planMode.ChatPlanMode != database.ChatPlanModePlan { - return nil - } - stopTools := map[string]struct{}{ - "propose_plan": {}, - } - if !parentChatID.Valid { - stopTools["ask_user_question"] = struct{}{} - } - return stopTools -} - -func stopAfterBehaviorTools( - planMode database.NullChatPlanMode, - chatMode database.NullChatMode, - parentChatID uuid.NullUUID, -) map[string]struct{} { - if isExploreSubagentMode(chatMode) { - return nil - } - return stopAfterPlanTools(planMode, parentChatID) -} - -type systemPromptBehaviorContext struct { - planMode database.NullChatPlanMode - chatMode database.NullChatMode - planModeInstructions string - isRootChat bool -} - -func workspaceSkillsForResolution(workspaceSkills []chattool.SkillMeta) []skillspkg.Skill { - if len(workspaceSkills) == 0 { - return nil - } - resolved := make([]skillspkg.Skill, 0, len(workspaceSkills)) - for _, skill := range workspaceSkills { - resolved = append(resolved, skillspkg.Skill{ - Name: skill.Name, - Description: skill.Description, - Source: skillspkg.SourceWorkspace, - }) - } - return resolved -} - -func mergeTurnSkills( - personalSkills []skillspkg.Skill, - workspaceSkills []chattool.SkillMeta, -) []skillspkg.ResolvedSkill { - return skillspkg.MergeSkills( - personalSkills, - workspaceSkillsForResolution(workspaceSkills), - ) -} - -// buildSystemPrompt applies system-level prompt injections in a fixed -// order: subagent instruction, chat instruction, skill index, user prompt, -// then mode overlay prompts. -func buildSystemPrompt( - prompt []fantasy.Message, - subagentInstruction string, - instruction string, - resolvedSkills []skillspkg.ResolvedSkill, - userPrompt string, - behaviorContext systemPromptBehaviorContext, -) []fantasy.Message { - if subagentInstruction != "" { - prompt = chatprompt.InsertSystem(prompt, subagentInstruction) - } - if instruction != "" { - prompt = chatprompt.InsertSystem(prompt, instruction) - } - if skillIndex := chattool.FormatResolvedSkillIndex(resolvedSkills); skillIndex != "" { - prompt = chatprompt.InsertSystem(prompt, skillIndex) - } - if userPrompt != "" { - prompt = chatprompt.InsertSystem(prompt, userPrompt) - } - if isExploreSubagentMode(behaviorContext.chatMode) { - prompt = chatprompt.InsertSystem(prompt, ExploreSubagentOverlayPrompt) - return prompt - } - isPlanModeTurn := behaviorContext.planMode.Valid && behaviorContext.planMode.ChatPlanMode == database.ChatPlanModePlan - if isPlanModeTurn { - if behaviorContext.isRootChat { - prompt = chatprompt.InsertSystem(prompt, PlanningOverlayPrompt()) - if behaviorContext.planModeInstructions != "" { - prompt = chatprompt.InsertSystem(prompt, behaviorContext.planModeInstructions) - } - } else { - prompt = chatprompt.InsertSystem(prompt, PlanningSubagentOverlayPrompt) - } - } - return prompt -} - -func removeSkillIndexMessages(prompt []fantasy.Message) []fantasy.Message { - out := make([]fantasy.Message, 0, len(prompt)) - removed := false - for _, message := range prompt { - if isSkillIndexMessage(message) { - removed = true - continue - } - out = append(out, message) - } - if !removed { - return prompt - } - return out -} - -func isSkillIndexMessage(message fantasy.Message) bool { - if message.Role != fantasy.MessageRoleSystem || len(message.Content) != 1 { - return false - } - textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](message.Content[0]) - if !ok { - return false - } - text := strings.TrimSpace(textPart.Text) - return strings.HasPrefix(text, chattool.AvailableSkillsOpenTag+"\n") && strings.HasSuffix(text, chattool.AvailableSkillsCloseTag) -} - -type rootChatToolsOptions struct { - chat database.Chat - modelConfigID uuid.UUID - workspaceCtx *turnWorkspaceContext - workspaceMu *sync.Mutex - resolvePlanPath func(context.Context) (string, string, error) - storeFile chattool.StoreFileFunc - isPlanModeTurn bool -} - -func (p *Server) loadPlanModeInstructions( - ctx context.Context, - mode database.NullChatPlanMode, - logger slog.Logger, -) string { - if !mode.Valid || mode.ChatPlanMode != database.ChatPlanModePlan { - return "" - } - - // Plan-mode instructions live in deployment config, but chat workers do - // not carry a deployment-config actor during background execution. - //nolint:gocritic // Required to read deployment config during background chat processing. - systemCtx := dbauthz.AsSystemRestricted(ctx) - fetched, err := p.db.GetChatPlanModeInstructions(systemCtx) - if err != nil { - logger.Warn(ctx, - "failed to fetch plan mode instructions", - slog.Error(err), - ) - return "" - } - - return fetched -} - -func userSkillContext(ctx context.Context, userID uuid.UUID) context.Context { - actor := rbac.Subject{ - Type: rbac.SubjectTypeUser, - ID: userID.String(), - Roles: rbac.RoleIdentifiers{rbac.RoleMember()}, - Scope: rbac.ScopeAll, - }.WithCachedASTValue() - // Chat turns run asynchronously after admission, so the original request - // actor may no longer be available when a worker loads personal skills. - // We synthesize the chat owner as a member instead of reusing that actor. - // Hardcoding RoleMember is safe because dbauthz enforces - // ResourceUserSkill.WithOwner(userID), so this actor cannot read any other - // user's skills regardless of role. Org scoping is not needed because - // personal skills are user-scoped, not org-scoped. - //nolint:gocritic // The synthetic actor is intentional for the reasons above. - return dbauthz.As(ctx, actor) -} - -func (p *Server) fetchPersonalSkillMetadata( - ctx context.Context, - userID uuid.UUID, - logger slog.Logger, -) []skillspkg.Skill { - rows, err := p.db.ListUserSkillMetadataByUserID(userSkillContext(ctx, userID), userID) - // See package coderd/x/skills (doc.go) for why metadata fetch failures - // intentionally degrade to an empty personal-skill list instead of - // failing the chat turn. - if err != nil { - logger.Warn(ctx, "failed to load personal skill metadata", - slog.F("owner_id", userID), - slog.Error(err), - ) - return nil - } - - personalSkills := make([]skillspkg.Skill, 0, len(rows)) - for _, row := range rows { - personalSkills = append(personalSkills, skillspkg.Skill{ - Name: row.Name, - Description: row.Description, - Source: skillspkg.SourcePersonal, - }) - } - return personalSkills -} - -func (p *Server) loadPersonalSkillBody( - ctx context.Context, - userID uuid.UUID, - name string, -) (skillspkg.ParsedSkill, error) { - row, err := p.db.GetUserSkillByUserIDAndName( - userSkillContext(ctx, userID), - database.GetUserSkillByUserIDAndNameParams{ - UserID: userID, - Name: name, - }, - ) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return skillspkg.ParsedSkill{}, skillspkg.ErrSkillNotFound - } - p.logger.Error(ctx, "load personal skill body failed", - slog.F("user_id", userID), - slog.F("name", name), - slog.Error(err), - ) - return skillspkg.ParsedSkill{}, xerrors.Errorf("load personal skill body: %w", err) - } - - parsed, err := skillspkg.ParsePersonalSkillMarkdown([]byte(row.Content)) - if err != nil { - p.logger.Error(ctx, "parse personal skill body failed", - slog.F("user_id", userID), - slog.F("name", name), - slog.Error(err), - ) - return skillspkg.ParsedSkill{}, xerrors.Errorf("parse personal skill body: %w", err) - } - return parsed, nil -} - -func (p *Server) appendRootChatTools( - ctx context.Context, - tools []fantasy.AgentTool, - opts rootChatToolsOptions, -) []fantasy.AgentTool { - onChatUpdated := func(updatedChat database.Chat) { - opts.workspaceCtx.selectWorkspace(updatedChat) - // Notify the frontend immediately so it can start streaming - // build logs before the tool completes. - p.publishChatPubsubEvent(updatedChat, codersdk.ChatWatchEventKindStatusChange, nil) - } - - tools = append(tools, - chattool.ListTemplates(p.db, opts.chat.OrganizationID, chattool.ListTemplatesOptions{ - OwnerID: opts.chat.OwnerID, - Logger: p.logger, - Clock: p.clock, - }), - chattool.ReadTemplate(p.db, opts.chat.OrganizationID, chattool.ReadTemplateOptions{ - OwnerID: opts.chat.OwnerID, - }), - chattool.CreateWorkspace(p.db, opts.chat.OrganizationID, opts.chat.ID, chattool.CreateWorkspaceOptions{ - OwnerID: opts.chat.OwnerID, - CreateFn: p.createWorkspaceFn, - AgentConnFn: chattool.AgentConnFunc(p.agentConnFn), - AgentInactiveDisconnectTimeout: p.agentInactiveDisconnectTimeout, - WorkspaceMu: opts.workspaceMu, - OnChatUpdated: onChatUpdated, - Logger: p.logger, - }), - chattool.StartWorkspace(p.db, opts.chat.ID, chattool.StartWorkspaceOptions{ - OwnerID: opts.chat.OwnerID, - StartFn: p.startWorkspaceFn, - AgentConnFn: chattool.AgentConnFunc(p.agentConnFn), - WorkspaceMu: opts.workspaceMu, - OnChatUpdated: onChatUpdated, - Logger: p.logger, - }), - chattool.StopWorkspace(p.db, opts.chat.ID, chattool.StopWorkspaceOptions{ - OwnerID: opts.chat.OwnerID, - StopFn: p.stopWorkspaceFn, - WorkspaceMu: opts.workspaceMu, - OnChatUpdated: onChatUpdated, - Logger: p.logger, - }), - ) - if opts.isPlanModeTurn { - tools = append(tools, chattool.ProposePlan(chattool.ProposePlanOptions{ - GetWorkspaceConn: opts.workspaceCtx.getWorkspaceConn, - ResolvePlanPath: opts.resolvePlanPath, - IsPlanTurn: opts.isPlanModeTurn, - StoreFile: opts.storeFile, - })) - } - - return append(tools, p.subagentTools(ctx, func() database.Chat { - return opts.chat - }, opts.modelConfigID)...) -} - -func appendDynamicTools( - ctx context.Context, - logger slog.Logger, - tools []fantasy.AgentTool, - raw pqtype.NullRawMessage, - planMode database.NullChatPlanMode, - chatMode database.NullChatMode, -) ([]fantasy.AgentTool, map[string]bool, error) { - if isExploreSubagentMode(chatMode) || (planMode.Valid && planMode.ChatPlanMode == database.ChatPlanModePlan) { - return tools, nil, nil - } - - dynamicToolNames, err := parseDynamicToolNames(raw) - if err != nil { - return nil, nil, xerrors.Errorf("parse dynamic tool names: %w", err) - } - if len(dynamicToolNames) == 0 { - return tools, dynamicToolNames, nil - } - - var dynamicToolDefs []codersdk.DynamicTool - if raw.Valid { - if err := json.Unmarshal(raw.RawMessage, &dynamicToolDefs); err != nil { - return nil, nil, xerrors.Errorf("unmarshal dynamic tools: %w", err) - } - } - - activeToolNames := make(map[string]struct{}, len(tools)) - for _, name := range allowedBehaviorToolNames(tools, chatMode) { - activeToolNames[name] = struct{}{} - } - for _, t := range tools { - info := t.Info() - if _, active := activeToolNames[info.Name]; !active { - continue - } - if dynamicToolNames[info.Name] { - logger.Warn(ctx, "dynamic tool name collides with built-in tool, built-in takes precedence", - slog.F("tool_name", info.Name)) - delete(dynamicToolNames, info.Name) - } - } - - var filteredDefs []codersdk.DynamicTool - for _, dt := range dynamicToolDefs { - if dynamicToolNames[dt.Name] { - filteredDefs = append(filteredDefs, dt) - } - } - - return append(tools, dynamicToolsFromSDK(logger, filteredDefs)...), dynamicToolNames, nil -} - -// buildProviderTools creates provider-native tool definitions -// (like web search) based on the model configuration. These -// tools are executed server-side by the LLM provider. -func buildProviderTools(options *codersdk.ChatModelProviderOptions) []chatloop.ProviderTool { - var tools []chatloop.ProviderTool - - if options == nil { - return nil - } - - if options.Anthropic != nil && options.Anthropic.WebSearchEnabled != nil && *options.Anthropic.WebSearchEnabled { - tools = append(tools, chatloop.ProviderTool{ - Definition: anthropic.WebSearchTool(&anthropic.WebSearchToolOptions{ - AllowedDomains: options.Anthropic.AllowedDomains, - BlockedDomains: options.Anthropic.BlockedDomains, - }), - }) - } - - if tool, ok := chatopenai.WebSearchTool(options.OpenAI); ok { - tools = append(tools, chatloop.ProviderTool{ - Definition: tool, - }) - } - - if options.Google != nil && options.Google.WebSearchEnabled != nil && *options.Google.WebSearchEnabled { - tools = append(tools, chatloop.ProviderTool{ - Definition: fantasy.ProviderDefinedTool{ - ID: "web_search", - Name: "web_search", - }, - }) - } - - return tools -} - func (p *Server) aiProviderConfig(ctx context.Context, provider database.AIProvider) (chatprovider.ConfiguredProvider, error) { keys, err := p.db.GetAIProviderKeysByProviderID(ctx, provider.ID) if err != nil { diff --git a/coderd/x/chatd/generation.go b/coderd/x/chatd/generation.go index 5956d93eba1..6f4e48c41f5 100644 --- a/coderd/x/chatd/generation.go +++ b/coderd/x/chatd/generation.go @@ -20,7 +20,6 @@ import ( "github.com/coder/coder/v2/coderd/x/chatd/chathooks" "github.com/coder/coder/v2/coderd/x/chatd/chatloop" "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" - "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" "github.com/coder/coder/v2/coderd/x/chatd/chatretry" "github.com/coder/coder/v2/coderd/x/chatd/chatstate" "github.com/coder/coder/v2/coderd/x/chatd/chattool" @@ -50,42 +49,6 @@ type generationPrepareInput struct { ) } -// generationPrepared contains the side-effect inputs for a generation task. -type generationPrepared struct { - Chat database.Chat - Messages []database.ChatMessage - - Model chatprovider.Model - Prompt []fantasy.Message - Tools []fantasy.AgentTool - ActiveTools []string - AllowInactiveTools map[string]bool - ProviderTools []chatloop.ProviderTool - ModelRoute aiGatewayModelRoute - ModelBuildOptions modelBuildOptions - - // ResolvedProvider is the configured provider identity used to label - // user-facing errors. See chatloop.GenerateAssistantOptions.ErrorProvider. - ResolvedProvider string - - ModelConfigID uuid.UUID - CallTemplate fantasy.Call - ContextLimitFallback int64 - - DynamicToolNames map[string]bool - StopAfterTools map[string]struct{} - ExclusiveToolNames map[string]bool - BuiltinToolNames map[string]bool - ToolNameToConfigID map[string]uuid.UUID - - MaxSteps int - Compaction *generationCompaction - // Cleanup is always non-nil when prepareGeneration succeeds. - Cleanup func() - - Debug *generationDebug -} - // generationCompaction contains compaction inputs prepared for generation. type generationCompaction struct { // Override, when non-nil, is the compaction model override resolved at @@ -459,8 +422,8 @@ func (s *taskStarter) StartGeneration(ctx context.Context, input chatWorkerTaskS Messages: messages, RecordMCPConnectSummaries: input.DebugTurn.RecordMCPConnectSummaries, } - prepared, err := retryGenerationPhase(ctx, s, "prepare", func() (generationPrepared, error) { - return s.server.prepareGeneration(ctx, prepareInput) + prepared, err := retryGenerationPhase(ctx, s, "prepare", func() (turnEnvironment, error) { + return s.server.buildTurnEnvironment(ctx, prepareInput) }) if err != nil { if errors.Is(err, errTaskExpectedExit) || errors.Is(err, errTaskRetryable) { @@ -468,23 +431,23 @@ func (s *taskStarter) StartGeneration(ctx context.Context, input chatWorkerTaskS } return s.finishGenerationError(ctx, machine, input, err, generationAttemptNotRequired) } - cleanup := prepared.Cleanup + cleanup := prepared.Close var decision generationDecision - if input.StopNudges.consume(stopNudgeKey(prepared.Messages)) { + if input.StopNudges.consume(stopNudgeKey(prepared.Turn().messages)) { decision = generationDecision{kind: generationActionGenerateAssistant} } else { decision, err = retryGenerationPhase(ctx, s, "decide", func() (generationDecision, error) { return decideGenerationAction(generationDecisionInput{ - chat: prepared.Chat, - messages: prepared.Messages, - dynamicToolNames: prepared.DynamicToolNames, - exclusiveToolNames: prepared.ExclusiveToolNames, - stopAfterTools: prepared.StopAfterTools, - maxSteps: prepared.MaxSteps, - compactionEnabled: prepared.Compaction != nil, - compactionNeeded: prepared.Compaction != nil && prepared.Compaction.Required, - compactionThresholdPercent: generationCompactionThreshold(prepared.Compaction), - compactionContextLimit: generationCompactionContextLimit(prepared.Compaction), + chat: prepared.Turn().chat, + messages: prepared.Turn().messages, + dynamicToolNames: prepared.Toolset().dynamicToolNames, + exclusiveToolNames: prepared.Toolset().exclusiveToolNames, + stopAfterTools: prepared.Toolset().stopAfterTools, + maxSteps: prepared.Turn().maxSteps, + compactionEnabled: prepared.CompactionConfig() != nil, + compactionNeeded: prepared.CompactionConfig() != nil && prepared.CompactionConfig().Required, + compactionThresholdPercent: generationCompactionThreshold(prepared.CompactionConfig()), + compactionContextLimit: generationCompactionContextLimit(prepared.CompactionConfig()), }) }) } @@ -493,8 +456,8 @@ func (s *taskStarter) StartGeneration(ctx context.Context, input chatWorkerTaskS if errors.Is(err, errTaskExpectedExit) || errors.Is(err, errTaskRetryable) { return xerrors.Errorf("decide generation: %w", err) } - if errors.Is(err, errCompactionStillOverLimit) && prepared.Compaction != nil { - metricProvider, metricModel := compactionMetricIdentity(prepared.Compaction) + if errors.Is(err, errCompactionStillOverLimit) && prepared.CompactionConfig() != nil { + metricProvider, metricModel := compactionMetricIdentity(prepared.CompactionConfig()) s.server.metrics.RecordCompaction( metricProvider, metricModel, @@ -731,23 +694,23 @@ func (s *taskStarter) generateAssistant( ctx context.Context, machine *chatstate.ChatMachine, input chatWorkerTaskStartInput, - prepared generationPrepared, + prepared turnEnvironment, ) error { attempt, err := s.beginGenerationAttempt(ctx, machine, input) if err != nil { return xerrors.Errorf("begin generation attempt: %w", err) } defer attempt.closeEpisode() - runCtx := input.DebugTurn.Ensure(ctx, prepared.Chat, prepared.Debug) + runCtx := input.DebugTurn.Ensure(ctx, prepared.Turn().chat, prepared.Turn().debug) outcome, err := chatloop.GenerateAssistant(runCtx, chatloop.GenerateAssistantOptions{ - Model: prepared.Model.LanguageModel(), - ErrorProvider: prepared.ResolvedProvider, - Messages: prepared.Prompt, - Tools: prepared.Tools, - ActiveTools: prepared.ActiveTools, - ProviderTools: prepared.ProviderTools, - ContextLimitFallback: prepared.ContextLimitFallback, - CallTemplate: prepared.CallTemplate, + Model: prepared.ModelConfig().model.LanguageModel(), + ErrorProvider: prepared.ModelConfig().resolvedProvider, + Messages: prepared.Prompt(), + Tools: prepared.Toolset().tools, + ActiveTools: prepared.Toolset().activeTools, + ProviderTools: prepared.Toolset().providerTools, + ContextLimitFallback: prepared.ModelConfig().contextLimitFallback, + CallTemplate: prepared.ModelConfig().callTemplate, PublishMessagePart: attempt.publish, OnModelStreamStart: attempt.startModelInvocation, Logger: s.opts.Logger, @@ -766,9 +729,9 @@ func (s *taskStarter) generateAssistant( } outcome.Step.Content = chathooks.ApplyAdmittedToolCalls(outcome.Step.Content, preflight) messages, err := buildCommitStepMessages(buildCommitStepMessagesInput{ - modelConfigID: prepared.ModelConfigID, + modelConfigID: prepared.ModelConfig().configID, step: stepDataFromPersisted(outcome.Step), - toolNameToConfigID: prepared.ToolNameToConfigID, + toolNameToConfigID: prepared.Toolset().toolNameToConfigID, logger: s.opts.Logger, contentVersion: chatprompt.CurrentContentVersion, hookRewrittenToolCalls: preflight.Overrides, @@ -776,7 +739,7 @@ func (s *taskStarter) generateAssistant( if err != nil { return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) } - messages, err = appendHookResultMessages(messages, preflight.Results, prepared.ModelConfigID) + messages, err = appendHookResultMessages(messages, preflight.Results, prepared.ModelConfig().configID) if err != nil { return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) } @@ -786,21 +749,21 @@ func (s *taskStarter) generateAssistant( func (s *taskStarter) admitStepToolCalls( ctx context.Context, input chatWorkerTaskStartInput, - prepared generationPrepared, + prepared turnEnvironment, content []fantasy.Content, ) (chathooks.PreToolUseExecutionResult, error) { if !s.server.hooks.Enabled() { return chathooks.PreToolUseExecutionResult{}, nil } toolCalls := chathooks.PendingToolCalls(content) - if len(toolCalls) == 0 || exclusiveBatchRejected(toolCalls, prepared.ExclusiveToolNames) { + if len(toolCalls) == 0 || exclusiveBatchRejected(toolCalls, prepared.Toolset().exclusiveToolNames) { return chathooks.PreToolUseExecutionResult{}, nil } // An admission error discards the whole batch before it can be // committed, so its find_tools calls would otherwise never reach // the executeLocalTools counter; count them at each error exit. countBatch := func() { - if !prepared.BuiltinToolNames[chattool.FindToolsName] { + if !prepared.Toolset().builtinToolNames[chattool.FindToolsName] { return } for _, toolCall := range toolCalls { @@ -816,7 +779,7 @@ func (s *taskStarter) admitStepToolCalls( return chathooks.PreToolUseExecutionResult{}, chathooks.GenerationDispatchError(agenthooks.EventPreToolUse, err) } unambiguous, _, ambiguous := partitionAmbiguousToolCalls(prepared, toolCalls) - preflight, err := s.server.hooks.PreflightPendingToolCalls(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), unambiguous) + preflight, err := s.server.hooks.PreflightPendingToolCalls(ctx, chathooks.ChatFor(prepared.Turn().chat, input.hookTurnID()), unambiguous) if err != nil { countBatch() return chathooks.PreToolUseExecutionResult{}, chathooks.GenerationDispatchError(agenthooks.EventPreToolUse, err) @@ -829,7 +792,7 @@ func (s *taskStarter) admitStepToolCalls( // Calls denied at admission persist synthetic results with the // assistant step, so they never surface as unresolved calls where // executeLocalTools counts find_tools invocations; count them here. - if prepared.BuiltinToolNames[chattool.FindToolsName] { + if prepared.Toolset().builtinToolNames[chattool.FindToolsName] { for _, result := range preflight.Denied { if result.ToolName == chattool.FindToolsName { s.server.metrics.FindToolsCallsTotal.Inc() @@ -865,10 +828,10 @@ func (s *taskStarter) executeLocalTools( ctx context.Context, machine *chatstate.ChatMachine, input chatWorkerTaskStartInput, - prepared generationPrepared, + prepared turnEnvironment, decision generationDecision, ) error { - exclusiveRejected := exclusiveBatchRejected(decision.localToolCalls, prepared.ExclusiveToolNames) + exclusiveRejected := exclusiveBatchRejected(decision.localToolCalls, prepared.Toolset().exclusiveToolNames) allowed := decision.localToolCalls var allowedIndexes []int var denied []fantasy.ToolResultContent @@ -879,7 +842,7 @@ func (s *taskStarter) executeLocalTools( // model-emitted call passes through, because rejections upstream of // the tool (partition denials, hook denials, exclusive-policy // batches) never reach its handler or OnCall. - if prepared.BuiltinToolNames[chattool.FindToolsName] { + if prepared.Toolset().builtinToolNames[chattool.FindToolsName] { for _, toolCall := range decision.localToolCalls { if toolCall.ToolName == chattool.FindToolsName { s.server.metrics.FindToolsCallsTotal.Inc() @@ -893,9 +856,9 @@ func (s *taskStarter) executeLocalTools( defer attempt.closeEpisode() provider := "" modelName := "" - if prepared.Model.Valid() { - provider = prepared.Model.Provider() - modelName = prepared.Model.ModelID() + if prepared.ModelConfig().model.Valid() { + provider = prepared.ModelConfig().model.Provider() + modelName = prepared.ModelConfig().model.ModelID() } var outcome chatloop.PersistedStep var spawnDispatchErr error @@ -909,17 +872,17 @@ func (s *taskStarter) executeLocalTools( } } outcome, err = chatloop.ExecuteLocalTools(ctx, chatloop.ExecuteLocalToolsOptions{ - Tools: prepared.Tools, - ActiveTools: prepared.ActiveTools, - AllowInactiveTools: prepared.AllowInactiveTools, - ProviderTools: prepared.ProviderTools, + Tools: prepared.Toolset().tools, + ActiveTools: prepared.Toolset().activeTools, + AllowInactiveTools: prepared.Toolset().allowInactiveTools, + ProviderTools: prepared.Toolset().providerTools, ToolCalls: allowed, ObservedToolCalls: decision.localToolCalls, - ExclusiveToolNames: prepared.ExclusiveToolNames, - BuiltinToolNames: prepared.BuiltinToolNames, + ExclusiveToolNames: prepared.Toolset().exclusiveToolNames, + BuiltinToolNames: prepared.Toolset().builtinToolNames, ModelProvider: provider, ModelName: modelName, - ContextLimit: prepared.ContextLimitFallback, + ContextLimit: prepared.ModelConfig().contextLimitFallback, ToolNameAliases: subagentToolNameAliases, UnbilledToolNames: unbilledSubagentToolNames, BillingRecorder: billingRecorder, @@ -939,23 +902,23 @@ func (s *taskStarter) executeLocalTools( spawnDispatchErr = chathooks.GenerationDispatchError(agenthooks.EventUserPromptSubmit, hookErr) } } - postResults, postDispatchErr := s.server.hooks.PostToolUseResults(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), outcome.Content) + postResults, postDispatchErr := s.server.hooks.PostToolUseResults(ctx, chathooks.ChatFor(prepared.Turn().chat, input.hookTurnID()), outcome.Content) for _, result := range denied { outcome.Content = append(outcome.Content, result) } chathooks.RestoreToolCallOrder(outcome.Content, decision.localToolCalls) step := stepDataFromPersisted(outcome) messages, err := buildCommitStepMessages(buildCommitStepMessagesInput{ - modelConfigID: prepared.ModelConfigID, + modelConfigID: prepared.ModelConfig().configID, step: step, - toolNameToConfigID: prepared.ToolNameToConfigID, + toolNameToConfigID: prepared.Toolset().toolNameToConfigID, logger: s.opts.Logger, contentVersion: chatprompt.CurrentContentVersion, }) if err != nil { return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) } - messages, err = appendHookResultMessages(messages, postResults, prepared.ModelConfigID) + messages, err = appendHookResultMessages(messages, postResults, prepared.ModelConfig().configID) if err != nil { return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) } @@ -994,7 +957,7 @@ func (s *taskStarter) generateCompaction( ctx context.Context, machine *chatstate.ChatMachine, input chatWorkerTaskStartInput, - prepared generationPrepared, + prepared turnEnvironment, source chatloop.CompactionSource, ) error { attempt, err := s.beginGenerationAttempt(ctx, machine, input) @@ -1002,27 +965,27 @@ func (s *taskStarter) generateCompaction( return xerrors.Errorf("beginGenerationAttempt: %w", err) } defer attempt.closeEpisode() - if prepared.Compaction == nil { + if prepared.CompactionConfig() == nil { return s.finishGenerationError(ctx, machine, input, xerrors.New("compaction action missing options"), requireGenerationAttempt(attempt.number)) } - compactionOpts := prepared.Compaction.Options - metricProvider, metricModel := compactionMetricIdentity(prepared.Compaction) - if override := prepared.Compaction.Override; override != nil { + compactionOpts := prepared.CompactionConfig().Options + metricProvider, metricModel := compactionMetricIdentity(prepared.CompactionConfig()) + if override := prepared.CompactionConfig().Override; override != nil { // A usable override that fails to build is a hard generation failure. overrideModel, err := s.server.resolveModelCall(ctx, modelCallSpec{ purpose: "compaction", - chat: prepared.Chat, + chat: prepared.Turn().chat, explicitConfig: &override.Config, requestedEffort: override.ReasoningEffort, chatdScopedRoute: true, - buildOptions: prepared.ModelBuildOptions, + buildOptions: prepared.ModelConfig().buildOptions, }) if err != nil { return xerrors.Errorf("build compaction model override: %w", err) } logger := s.server.logger.With( - slog.F("chat_id", prepared.Chat.ID), - slog.F("owner_id", prepared.Chat.OwnerID), + slog.F("chat_id", prepared.Turn().chat.ID), + slog.F("owner_id", prepared.Turn().chat.OwnerID), ) compactionOpts.Model = overrideModel.model.LanguageModel() compactionOpts.ResolvedProvider = overrideModel.resolvedProvider @@ -1034,11 +997,11 @@ func (s *taskStarter) generateCompaction( logger, compactionOpts.Messages, overrideModel.model, - prepared.Compaction.ChatModelConfig, + prepared.CompactionConfig().ChatModelConfig, overrideModel.dbConfig, ) } - preResult, err := s.server.hooks.Trigger(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), chathooks.Message{}, agenthooks.EventPreCompact, dispatch.CapacityClassGeneration) + preResult, err := s.server.hooks.Trigger(ctx, chathooks.ChatFor(prepared.Turn().chat, input.hookTurnID()), chathooks.Message{}, agenthooks.EventPreCompact, dispatch.CapacityClassGeneration) if err != nil { return chathooks.GenerationDispatchError(agenthooks.EventPreCompact, err) } @@ -1051,7 +1014,7 @@ func (s *taskStarter) generateCompaction( // Attach the turn debug run so the compaction call records a child // debug run; without it startCompactionDebugRun finds no parent and // skips debug instrumentation entirely. - runCtx := input.DebugTurn.Ensure(ctx, prepared.Chat, prepared.Debug) + runCtx := input.DebugTurn.Ensure(ctx, prepared.Turn().chat, prepared.Turn().debug) outcome, err := chatloop.GenerateCompaction(runCtx, compactionOpts) if err != nil { s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err) @@ -1063,7 +1026,7 @@ func (s *taskStarter) generateCompaction( return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) } messages, err := buildCompactionMessages(buildCompactionMessagesInput{ - modelConfigID: prepared.ModelConfigID, + modelConfigID: prepared.ModelConfig().configID, toolCallID: compactionOpts.ToolCallID, toolName: compactionOpts.ToolName, compaction: compactionOutcome(outcome), @@ -1079,7 +1042,7 @@ func (s *taskStarter) generateCompaction( Messages: messages.Messages, VisibleIndexes: visibleMessageIndexes(messages.Messages), ConsumeCompactionRequest: true, - }, []*chathooks.Result{persistedPreResult}, prepared.ModelConfigID) + }, []*chathooks.Result{persistedPreResult}, prepared.ModelConfig().configID) if err != nil { s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err) return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) @@ -1087,12 +1050,12 @@ func (s *taskStarter) generateCompaction( // Hook effects and fail-closed errors must commit atomically with // compaction; a separate commit races the runner and can be dropped // on crash. - postResult, postDispatchErr := s.server.hooks.Trigger(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), chathooks.Message{}, agenthooks.EventPostCompact, dispatch.CapacityClassGeneration) + postResult, postDispatchErr := s.server.hooks.Trigger(ctx, chathooks.ChatFor(prepared.Turn().chat, input.hookTurnID()), chathooks.Message{}, agenthooks.EventPostCompact, dispatch.CapacityClassGeneration) var postCommitErr error if postDispatchErr != nil { postCommitErr = chathooks.GenerationDispatchError(agenthooks.EventPostCompact, postDispatchErr) } else { - commitMessages, err = appendHookResultMessages(commitMessages, []*chathooks.Result{postResult}, prepared.ModelConfigID) + commitMessages, err = appendHookResultMessages(commitMessages, []*chathooks.Result{postResult}, prepared.ModelConfig().configID) if err != nil { s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err) return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) diff --git a/coderd/x/chatd/generation_preparer.go b/coderd/x/chatd/generation_preparer.go deleted file mode 100644 index ff49843deb6..00000000000 --- a/coderd/x/chatd/generation_preparer.go +++ /dev/null @@ -1,926 +0,0 @@ -package chatd - -import ( - "context" - "slices" - "strings" - "sync" - - "charm.land/fantasy" - "github.com/google/uuid" - "golang.org/x/sync/errgroup" - "golang.org/x/xerrors" - - "cdr.dev/slog/v3" - "github.com/coder/coder/v2/coderd/database" - "github.com/coder/coder/v2/coderd/x/chatd/chatadvisor" - "github.com/coder/coder/v2/coderd/x/chatd/chatloop" - "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" - "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" - "github.com/coder/coder/v2/coderd/x/chatd/chatsanitize" - "github.com/coder/coder/v2/coderd/x/chatd/chattool" - "github.com/coder/coder/v2/coderd/x/chatd/mcpclient" - skillspkg "github.com/coder/coder/v2/coderd/x/skills" - "github.com/coder/coder/v2/codersdk" -) - -// effectiveMCPServerConfigs loads the chat's stored selection plus -// owner-readable Force On configs at generation time, so stored lists -// predating enforcement cannot dodge the policy (Cure53 CDM-02-010). -// Explore chats keep their immutable spawn-time snapshot instead. -func (server *Server) effectiveMCPServerConfigs( - ctx context.Context, - logger slog.Logger, - chat database.Chat, -) ([]database.MCPServerConfig, error) { - var configs []database.MCPServerConfig - if len(chat.MCPServerIDs) > 0 { - var err error - configs, err = enabledMCPServerConfigsForChatOrg(ctx, server.db, chat.OrganizationID, chat.MCPServerIDs) - if err != nil { - // Best-effort for the user-selected set, matching prior - // behavior: a load failure degrades the turn rather than - // failing it. - logger.Warn(ctx, "failed to load MCP server configs", slog.Error(err)) - configs = nil - } - } - if isExploreSubagentMode(chat.Mode) { - return configs, nil - } - forced, err := forcedMCPServerConfigsForOwner(ctx, server.db, chat.OrganizationID, chat.OwnerID) - if err != nil { - // Fail closed: running the turn without the forced set would - // silently bypass a security policy. - return nil, err - } - seen := make(map[uuid.UUID]struct{}, len(configs)) - for _, cfg := range configs { - seen[cfg.ID] = struct{}{} - } - for _, cfg := range forced { - if _, ok := seen[cfg.ID]; !ok { - configs = append(configs, cfg) - } - } - return configs, nil -} - -func (server *Server) prepareGeneration( - ctx context.Context, - input generationPrepareInput, -) (generationPrepared, error) { - chat := input.Chat - logger := server.logger.With( - slog.F("chat_id", chat.ID), - slog.F("owner_id", chat.OwnerID), - ) - - prepStart := server.clock.Now() - defer func() { - if prepDuration := server.clock.Since(prepStart); prepDuration >= slowPrepareThreshold { - logger.Warn(ctx, "slow generation preparation", - slog.F("duration", prepDuration), - ) - } - }() - - var ( - promptRows []database.ChatMessage - mcpConfigs []database.MCPServerConfig - mcpTokens []database.MCPServerUserToken - ) - - var g errgroup.Group - g.Go(func() error { - var err error - promptRows, err = server.db.GetChatMessagesForPromptByChatID(ctx, chat.ID) - if err != nil { - return xerrors.Errorf("get chat messages for prompt: %w", err) - } - return nil - }) - g.Go(func() error { - var err error - mcpConfigs, err = server.effectiveMCPServerConfigs(ctx, logger, chat) - return err - }) - if err := g.Wait(); err != nil { - return generationPrepared{}, err - } - - apiKeyID, err := server.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) - if err != nil { - return generationPrepared{}, xerrors.Errorf("ensure synthetic API key: %w", err) - } - modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} - - requestedEffort := chatRequestedEffort(chat) - resolved, err := server.resolveModelCall(ctx, modelCallSpec{ - purpose: "standard_turn", - chat: chat, - requestedEffort: requestedEffort, - buildOptions: modelOpts, - }) - if err != nil { - return generationPrepared{}, err - } - // The chat config keeps driving compaction, sanitization, and debug - // attribution even when computer use swaps the resolved call below. - modelConfig := resolved.dbConfig - - // Computer-use turns swap in a specialized model, so the substitution - // must happen before anything model-sensitive runs: file-part - // classification, history sanitization, and provider option preparation - // must all agree with the client actually used for the turn. - isComputerUse := chat.Mode.Valid && chat.Mode.ChatMode == database.ChatModeComputerUse - var computerUseProvider codersdk.ChatComputerUseProvider - if isComputerUse { - var cuModelProvider, cuModelName string - computerUseProvider, cuModelProvider, cuModelName, err = server.computerUseProviderAndModelFromConfig(ctx) - if err != nil { - return generationPrepared{}, xerrors.Errorf("resolve computer use provider and model: %w", err) - } - cuResolved, cuErr := server.resolveModelCall(ctx, modelCallSpec{ - purpose: "computer_use", - chat: chat, - fixedModel: &fixedModelCall{ - providerType: cuModelProvider, - modelName: cuModelName, - callConfig: resolved.callConfig, - }, - requestedEffort: requestedEffort, - buildOptions: modelOpts, - }) - if cuErr != nil { - return generationPrepared{}, xerrors.Errorf( - "resolve computer use model for provider %q model %q: %w", - computerUseProvider, - cuModelName, - cuErr, - ) - } - resolved = cuResolved - } - model := resolved.model - callConfig := resolved.callConfig - modelRoute := resolved.route - - currentPlanMode := chat.PlanMode - isPlanModeTurn := currentPlanMode.Valid && currentPlanMode.ChatPlanMode == database.ChatPlanModePlan - isExploreSubagent := isExploreSubagentMode(chat.Mode) - isRootChat := !chat.ParentChatID.Valid - - mcpConnectConfigs, approvedPlanMCPConfigIDs := filterExternalMCPConfigsForTurn( - mcpConfigs, - currentPlanMode, - chat.ParentChatID, - ) - if isExploreSubagent && isRootChat { - mcpConnectConfigs = nil - approvedPlanMCPConfigIDs = map[uuid.UUID]struct{}{} - } - - planModeInstructions := server.loadPlanModeInstructions(ctx, currentPlanMode, logger) - advisorCfg := server.loadAdvisorConfig(ctx, logger) - // Force Enabled from the experiment; the stored DB value is ignored. - advisorCfg.Enabled = server.experiments.Enabled(codersdk.ExperimentChatAdvisor) - - var advisorRuntime *chatadvisor.Runtime - if advisorCfg.Enabled && isRootChat && !isPlanModeTurn && !isExploreSubagent { - var advisorErr error - advisorRuntime, advisorErr = server.newAdvisorRuntime( - ctx, - chat, - advisorCfg, - modelOpts, - logger, - ) - if advisorErr != nil { - return generationPrepared{}, advisorErr - } - } - - var advisorPromptSnapshot []fantasy.Message - setAdvisorPromptSnapshot := func(msgs []fantasy.Message) { - if advisorRuntime == nil { - return - } - advisorPromptSnapshot = slices.Clone(msgs) - } - - currentChat := chat - loadChatSnapshot := func(loadCtx context.Context, chatID uuid.UUID) (database.Chat, error) { - return server.db.GetChatByID(loadCtx, chatID) - } - var chatStateMu sync.Mutex - var workspaceMu sync.Mutex - workspaceCtx := turnWorkspaceContext{ - server: server, - chatStateMu: &chatStateMu, - currentChat: ¤tChat, - loadChatSnapshot: loadChatSnapshot, - } - cleanup := func() { - workspaceCtx.close() - } - - planPathFn := func(ctx context.Context) (string, string, error) { - conn, err := workspaceCtx.getWorkspaceConn(ctx) - if err != nil { - return "", "", err - } - home, err := chattool.ResolveWorkspaceHome(ctx, conn) - if err != nil { - return "", "", err - } - return chattool.PlanPathForChat(home, chat.ID), home, nil - } - resolvePlanPathForTools := func(ctx context.Context) (string, string, error) { - planCtx, cancel := context.WithTimeout(ctx, planPathLookupTimeout) - defer cancel() - return planPathFn(planCtx) - } - resolvePlanPathBlock := func(resolveCtx context.Context) string { - if chat.ParentChatID.Valid { - return "" - } - - planCtx, cancel := context.WithTimeout(resolveCtx, planPathLookupTimeout) - defer cancel() - - if _, _, err := workspaceCtx.workspaceAgentIDForConn(planCtx); err != nil { - logger.Debug(resolveCtx, "plan path instruction: agent not reachable", - slog.Error(err), - slog.F("chat_id", chat.ID), - ) - return "" - } - - planPath, home, err := planPathFn(planCtx) - if err != nil { - logger.Debug(resolveCtx, "plan path instruction: failed to resolve plan path", - slog.Error(err), - slog.F("chat_id", chat.ID), - ) - return "" - } - return formatPlanPathBlock(planPath, home) - } - - var ( - prompt []fantasy.Message - instruction string - mcpTools []fantasy.AgentTool - mcpSummaries []mcpclient.ConnectSummary - mcpCleanup func() - workspaceMCPTools []fantasy.AgentTool - workspaceSkills []chattool.SkillMeta - personalSkills []skillspkg.Skill - resolvedUserPrompt string - planPathBlock string - ) - - // Drop provider-executed tool history produced by a different provider - // before building the prompt. A provider that shares another's wire format - // (e.g. Bedrock and Anthropic) can still reject the other's - // provider-executed blocks, so a mid-chat provider switch must not replay - // them. - promptRows = server.sanitizeForeignProviderExecutedToolRows(ctx, logger, promptRows, chat.OwnerID, modelConfig.ID) - - if chat.WorkspaceID.Valid { - // Resolve the workspace agent so the chat row's AgentID and - // BuildID bindings are up to date before the chatworker - // decision helper inspects them. ensureWorkspaceAgent does a - // DB lookup and lazily calls persistBuildAgentBinding when - // the bound agent has changed, so this is a cheap metadata - // refresh, not a workspace dial. It must not insert chat - // history; only metadata is mutated here. - agent, _ := workspaceCtx.getWorkspaceAgent(ctx) - - // API-created chats bind their agent lazily here, after - // hydrateChatContextOnCreate ran with no agent. Pin the chat to the - // bound agent's pushed snapshot now if it is still unpinned, so the - // first turn reads workspace context instead of waiting for the - // agent's next push. Idempotent and snapshot-gated; runs before the - // pinned context is read below. - server.ensureChatContextPinnedOnFirstTurn(ctx, workspaceCtx.currentChatSnapshot()) - - var resolveErr error - instruction, workspaceSkills, resolveErr = server.resolveTurnWorkspaceContext(ctx, chat, agent) - if resolveErr != nil { - cleanup() - return generationPrepared{}, resolveErr - } - } - - // Build the debug context before the connect phase so its - // outcomes can be recorded with run-creation context even when a - // later preparation step fails. - triggerMessageID, historyTipMessageID, triggerLabel := deriveChatDebugSeed(promptRows) - debugSvc := server.existingDebugService() - var debug *generationDebug - if resolved.debugEnabled { - if debugSvc == nil { - cleanup() - return generationPrepared{}, xerrors.New("chat debug service missing after enablement check") - } - debug = &generationDebug{ - Enabled: true, - Service: debugSvc, - Provider: resolved.resolvedProvider, - Model: resolved.resolvedModel, - TriggerMessageID: triggerMessageID, - HistoryTipMessageID: historyTipMessageID, - TriggerLabel: triggerLabel, - ModelConfig: modelConfig, - } - } - - var g2 errgroup.Group - g2.Go(func() error { - var err error - // Key the file-part acceptance on model.Provider() (the fantasy - // transport identity), not the configured provider, because - // aibridge routing rewrites the provider (e.g. Bedrock to the - // Anthropic transport). The conversion that actually drops or - // accepts a file part is the one for model.Provider(). - acceptsFilePart := model.AcceptsFilePartMediaType - providerType := string(modelRoute.Provider.Type) - prompt, err = chatprompt.ConvertMessagesWithFiles(ctx, promptRows, server.chatFileResolver(providerType), logger, acceptsFilePart) - if err != nil { - return xerrors.Errorf("build chat prompt: %w", err) - } - return nil - }) - g2.Go(func() error { - personalSkills = server.fetchPersonalSkillMetadata(ctx, chat.OwnerID, logger) - return nil - }) - g2.Go(func() error { - resolvedUserPrompt = server.resolveUserPrompt(ctx, chat.OwnerID) - return nil - }) - if len(mcpConnectConfigs) > 0 { - g2.Go(func() error { - var tokenErr error - mcpTokens, tokenErr = server.db.GetMCPServerUserTokensByUserID(ctx, chat.OwnerID) - if tokenErr != nil { - logger.Warn(ctx, "failed to load MCP user tokens", slog.Error(tokenErr)) - } - mcpTokens = server.refreshExpiredMCPTokens(ctx, logger, mcpConnectConfigs, mcpTokens) - mcpTools, mcpSummaries, mcpCleanup = mcpclient.ConnectAll( - ctx, - logger, - mcpConnectConfigs, - mcpTokens, - chat.OwnerID, - server.oidcTokenSource, - chatprovider.CoderHeaders(chat), - ) - return nil - }) - } - if chat.WorkspaceID.Valid && !isPlanModeTurn && !isExploreSubagent { - g2.Go(func() error { - workspaceMCPTools = server.resolveWorkspaceMCPTools(ctx, logger, chat, &workspaceCtx) - return nil - }) - } - // Resolve the per-chat plan path block in the parallel phase. It dials - // the workspace agent to read the home directory, so running it here lets - // the cold dial overlap with the rest of turn preparation instead of - // blocking system prompt assembly on a sequential dial. Best-effort: - // resolvePlanPathBlock logs and returns an empty block on failure. - if chat.WorkspaceID.Valid && !chat.ParentChatID.Valid { - g2.Go(func() error { - planPathBlock = resolvePlanPathBlock(ctx) - return nil - }) - } - g2Err := g2.Wait() - // Record connect outcomes before acting on any preparation error: - // ConnectAll has already run, so a failure below (or in g2 itself) - // would otherwise discard this attempt's outcomes. - if debug != nil && input.RecordMCPConnectSummaries != nil && len(mcpSummaries) > 0 { - input.RecordMCPConnectSummaries(ctx, chat, debug, mcpSummaries) - } - if g2Err != nil { - cleanup() - return generationPrepared{}, g2Err - } - - if mcpCleanup != nil { - previousCleanup := cleanup - cleanup = func() { - mcpCleanup() - previousCleanup() - } - } - - prompt, sanitizeStats := chatsanitize.SanitizeAnthropicProviderToolHistory(model.Provider(), prompt) - chatsanitize.LogAnthropicProviderToolSanitization( - ctx, - logger, - "persisted_history_replay", - model.Provider(), - model.ModelID(), - sanitizeStats, - ) - - subagentInstruction := "" - if !isRootChat { - subagentInstruction = defaultSubagentInstruction - } - resolvedSkillsFor := func(workspaceSkills []chattool.SkillMeta) []skillspkg.ResolvedSkill { - return mergeTurnSkills(personalSkills, workspaceSkills) - } - resolveSkillAlias := func(alias string) (skillspkg.ResolvedSkill, error) { - return skillspkg.Lookup(resolvedSkillsFor(workspaceSkills), alias) - } - initialResolvedSkills := resolvedSkillsFor(workspaceSkills) - - prompt = buildSystemPrompt( - prompt, - subagentInstruction, - instruction, - initialResolvedSkills, - resolvedUserPrompt, - systemPromptBehaviorContext{ - planMode: currentPlanMode, - chatMode: chat.Mode, - planModeInstructions: planModeInstructions, - isRootChat: isRootChat, - }, - ) - if advisorRuntime != nil { - prompt = chatprompt.InsertSystem(prompt, chatadvisor.ParentGuidanceBlock) - } - prompt = renderPlanPathPrompt(prompt, planPathBlock) - setAdvisorPromptSnapshot(prompt) - - storeChatAttachment := server.newStoreChatAttachmentFunc(&workspaceCtx) - tools := []fantasy.AgentTool{ - chattool.ReadFile(chattool.ReadFileOptions{GetWorkspaceConn: workspaceCtx.getWorkspaceConn}), - chattool.WriteFile(chattool.WriteFileOptions{ - GetWorkspaceConn: workspaceCtx.getWorkspaceConn, - ResolvePlanPath: resolvePlanPathForTools, - IsPlanTurn: isPlanModeTurn, - }), - chattool.EditFiles(chattool.EditFilesOptions{ - GetWorkspaceConn: workspaceCtx.getWorkspaceConn, - ResolvePlanPath: resolvePlanPathForTools, - IsPlanTurn: isPlanModeTurn, - }), - chattool.AttachFile(chattool.AttachFileOptions{ - GetWorkspaceConn: workspaceCtx.getWorkspaceConn, - StoreFile: storeChatAttachment, - }), - chattool.Execute(chattool.ExecuteOptions{ - GetWorkspaceConn: workspaceCtx.getWorkspaceConn, - AgentBrowserSession: chat.ID.String(), - }), - chattool.ProcessOutput(chattool.ProcessToolOptions{GetWorkspaceConn: workspaceCtx.getWorkspaceConn}), - chattool.ProcessList(chattool.ProcessToolOptions{GetWorkspaceConn: workspaceCtx.getWorkspaceConn}), - chattool.ProcessSignal(chattool.ProcessToolOptions{GetWorkspaceConn: workspaceCtx.getWorkspaceConn}), - } - if isPlanModeTurn && isRootChat { - tools = append(tools, chattool.NewAskUserQuestionTool()) - } - if isRootChat { - tools = server.appendRootChatTools(ctx, tools, rootChatToolsOptions{ - chat: chat, - modelConfigID: modelConfig.ID, - workspaceCtx: &workspaceCtx, - workspaceMu: &workspaceMu, - resolvePlanPath: resolvePlanPathForTools, - storeFile: storeChatAttachment, - isPlanModeTurn: isPlanModeTurn, - }) - } - - skillOpts := chattool.ReadSkillOptions{ - GetWorkspaceConn: workspaceCtx.getWorkspaceConn, - GetSkills: func() []chattool.SkillMeta { - return workspaceSkills - }, - ResolveAlias: resolveSkillAlias, - LoadPersonalSkillBody: func(ctx context.Context, name string) (skillspkg.ParsedSkill, error) { - return server.loadPersonalSkillBody(ctx, chat.OwnerID, name) - }, - } - appendCurrentSkillTools := func(current []fantasy.AgentTool) ([]fantasy.AgentTool, bool) { - if len(personalSkills) == 0 && len(workspaceSkills) == 0 { - return current, false - } - updated := current - changed := false - appendTool := func(tool fantasy.AgentTool) { - name := tool.Info().Name - if slices.ContainsFunc(current, func(existing fantasy.AgentTool) bool { - return existing.Info().Name == name - }) { - return - } - if !changed { - updated = slices.Clone(current) - changed = true - } - updated = append(updated, tool) - } - appendTool(chattool.ReadSkill(skillOpts)) - if len(workspaceSkills) > 0 { - appendTool(chattool.ReadSkillFile(skillOpts)) - } - return updated, changed - } - tools, _ = appendCurrentSkillTools(tools) - if advisorRuntime != nil { - tools = append(tools, chatadvisor.Tool(chatadvisor.ToolOptions{ - Runtime: advisorRuntime, - GetConversationSnapshot: func() []fantasy.Message { - return stripAdvisorGuidanceBlock(slices.Clone(advisorPromptSnapshot)) - }, - })) - } - - var exclusiveToolNames map[string]bool - if advisorRuntime != nil { - exclusiveToolNames = map[string]bool{chatadvisor.ToolName: true} - } - - builtinToolNames := make(map[string]bool, len(tools)) - for _, t := range tools { - builtinToolNames[t.Info().Name] = true - } - - mcpConfigByID := make(map[uuid.UUID]database.MCPServerConfig, len(mcpConnectConfigs)) - for _, config := range mcpConnectConfigs { - mcpConfigByID[config.ID] = config - } - deferredCandidates := collectDeferredMCPCandidates(deferredMCPCandidateInput{ - mcpTools: mcpTools, - workspaceMCPTools: workspaceMCPTools, - mcpConfigByID: mcpConfigByID, - planMode: currentPlanMode, - parentChatID: chat.ParentChatID, - approvedMCPConfigIDs: approvedPlanMCPConfigIDs, - includeWorkspaceTools: !isExploreSubagent, - }) - tools = append(tools, mcpTools...) - if !isExploreSubagent { - tools = append(tools, workspaceMCPTools...) - } - tools = filterToolsForTurn(tools, currentPlanMode, chat.ParentChatID, approvedPlanMCPConfigIDs) - - tools, dynamicToolNames, err := appendDynamicTools(ctx, logger, tools, chat.DynamicTools, currentPlanMode, chat.Mode) - if err != nil { - cleanup() - return generationPrepared{}, err - } - - var providerTools []chatloop.ProviderTool - if !isPlanModeTurn && callConfig.ProviderOptions != nil { - providerTools = buildProviderTools(callConfig.ProviderOptions) - if isExploreSubagent { - if !chat.ParentChatID.Valid { - providerTools = nil - } else { - providerTools = slices.DeleteFunc(providerTools, func(tool chatloop.ProviderTool) bool { - return tool.Definition.GetName() != "web_search" - }) - } - } - } - - if isComputerUse { - providerTools, err = appendComputerUseProviderTool(providerTools, computerUseProviderToolOptions{ - provider: computerUseProvider, - isPlanModeTurn: isPlanModeTurn, - isComputerUse: isComputerUse, - getWorkspaceConn: workspaceCtx.getWorkspaceConn, - storeFile: storeChatAttachment, - clock: server.clock, - logger: server.logger.Named("computer_use"), - }) - if err != nil { - cleanup() - return generationPrepared{}, xerrors.Errorf("register computer use provider tool for provider %q: %w", computerUseProvider, err) - } - } else { - providerTools, err = appendComputerUseProviderTool(providerTools, computerUseProviderToolOptions{ - isPlanModeTurn: isPlanModeTurn, - isComputerUse: false, - }) - if err != nil { - cleanup() - return generationPrepared{}, err - } - } - - activeToolNames := activeToolNamesForTurn(tools, currentPlanMode, chat.ParentChatID, approvedPlanMCPConfigIDs) - if isExploreSubagent { - activeToolNames = allowedExploreToolNames(tools) - } - var allowInactiveTools map[string]bool - if decideMCPToolSearch(mcpToolSearchInput{ - experimentEnabled: server.experiments.Enabled(codersdk.ExperimentMCPToolSearch), - candidates: deferredCandidates, - dynamicToolNames: dynamicToolNames, - }) { - activationTokenBudget := float64(modelConfig.ContextLimit) / mcpToolSearchBudgetDivisor - findTools := chattool.FindTools(chattool.FindToolsOptions{ - Entries: deferredMCPToolEntries(deferredCandidates), - SchemaTokenBudget: activationTokenBudget, - CatalogTokenBudget: activationTokenBudget, - // Calls total is counted in executeLocalTools, which also - // sees calls rejected before the tool runs; OnCall covers - // only calls that reach the handler or its decode. - OnCall: func(callCtx context.Context, call chattool.FindToolsCall) { - if call.Rejection == "" { - server.metrics.FindToolsMatchCount.Observe(float64(call.MatchCount)) - server.metrics.FindToolsActivationsTotal.Add(float64(len(call.Activated))) - if call.MatchCount == 0 { - server.metrics.FindToolsEmptyTotal.Inc() - } - } - // Queries and names are model output that can echo - // prompt content, so standard logs carry only - // aggregate fields; raw values are visible through - // the opt-in chat debug logging path. - logger.Info(callCtx, "deferred MCP tool search", - slog.F("query_count", len(call.Queries)), - slog.F("name_count", len(call.Names)), - slog.F("match_count", call.MatchCount), - slog.F("activated_count", len(call.Activated)), - slog.F("total_deferred", call.TotalDeferred), - slog.F("rejection", call.Rejection), - ) - }, - }) - tools, activeToolNames, allowInactiveTools = configureDeferredMCPToolSearch( - tools, - activeToolNames, - deferredCandidates, - findTools, - deriveDeferredMCPActivations(promptRows, deferredCandidates, activationTokenBudget), - ) - builtinToolNames[chattool.FindToolsName] = true - } - - toolNameToConfigID := make(map[string]uuid.UUID) - for _, t := range tools { - if mcpTool, ok := t.(mcpclient.MCPToolIdentifier); ok { - toolNameToConfigID[t.Info().Name] = mcpTool.MCPServerConfigID() - } - } - - compactionToolCallID := "chat_summarized_" + uuid.NewString() - effectiveThreshold := modelConfig.CompressionThreshold - if override, ok := server.resolveUserCompactionThreshold(ctx, chat.OwnerID, modelConfig.ID); ok { - effectiveThreshold = override - } - // The compaction trigger uses the stricter of the chat and override - // models' context limits: the history must also fit the summarizer's - // window. - compactionContextLimit := modelConfig.ContextLimit - compactionOverride, err := server.resolveCompactionOverrideConfig(ctx, chat) - if err != nil { - cleanup() - return generationPrepared{}, err - } - if compactionOverride != nil { - if overrideLimit := compactionOverride.Config.ContextLimit; overrideLimit > 0 && - (compactionContextLimit <= 0 || overrideLimit < compactionContextLimit) { - compactionContextLimit = overrideLimit - } - } - compactionStepUsage := latestPromptUsage(promptRows) - compactionNeeded := shouldCompactPromptUsage(compactionStepUsage, compactionContextLimit, effectiveThreshold) - // The options carry the chat model; generateCompaction swaps in the - // override client when one is configured. - compactionOptions := chatloop.GenerateCompactionOptions{ - Model: model.LanguageModel(), - Messages: prompt, - ThresholdPercent: effectiveThreshold, - ContextLimit: compactionContextLimit, - ContextLimitFallback: compactionContextLimit, - ToolCallID: compactionToolCallID, - ToolName: "chat_summarized", - DebugSvc: debugSvc, - ChatID: chat.ID, - HistoryTipMessageID: historyTipMessageID, - ResolvedProvider: resolved.resolvedProvider, - ResolvedModel: resolved.resolvedModel, - ModelConfigID: modelConfig.ID, - StepUsage: compactionStepUsage, - SummaryCall: compactionSummaryCall(resolved), - } - - // workspaceCtx.currentChatSnapshot may carry a freshly persisted - // AgentID/BuildID binding from the getWorkspaceAgent call above. - // Return that snapshot so downstream consumers see the up-to-date - // metadata. - refreshedChat := workspaceCtx.currentChatSnapshot() - if refreshedChat.ID == uuid.Nil { - refreshedChat = chat - } - - return generationPrepared{ - Chat: refreshedChat, - Messages: input.Messages, - Model: model, - Prompt: prompt, - Tools: tools, - ActiveTools: activeToolNames, - AllowInactiveTools: allowInactiveTools, - ProviderTools: providerTools, - ModelRoute: modelRoute, - ModelBuildOptions: modelOpts, - ResolvedProvider: resolved.resolvedProvider, - ModelConfigID: modelConfig.ID, - CallTemplate: resolved.newCall(), - ContextLimitFallback: modelConfig.ContextLimit, - DynamicToolNames: dynamicToolNames, - StopAfterTools: stopAfterBehaviorTools(currentPlanMode, chat.Mode, chat.ParentChatID), - ExclusiveToolNames: exclusiveToolNames, - BuiltinToolNames: builtinToolNames, - ToolNameToConfigID: toolNameToConfigID, - MaxSteps: maxChatSteps, - Compaction: &generationCompaction{ - Override: compactionOverride, - ChatModelConfig: modelConfig, - Required: compactionNeeded, - Options: compactionOptions, - }, - Cleanup: cleanup, - Debug: debug, - }, nil -} - -func latestPromptUsage(messages []database.ChatMessage) fantasy.Usage { - for i := len(messages) - 1; i >= 0; i-- { - usage := usageFromMessage(messages[i]) - if usage != (fantasy.Usage{}) { - return usage - } - } - return fantasy.Usage{} -} - -func shouldCompactPromptUsage(usage fantasy.Usage, contextLimit int64, thresholdPercent int32) bool { - if thresholdPercent >= 100 || contextLimit <= 0 { - return false - } - contextTokens := contextTokensFromUsage(usage) - if contextTokens <= 0 { - return false - } - usagePercent := (float64(contextTokens) / float64(contextLimit)) * 100 - return usagePercent >= float64(thresholdPercent) -} - -func contextTokensFromUsage(usage fantasy.Usage) int64 { - total := int64(0) - hasContextTokens := false - if usage.InputTokens > 0 { - total += usage.InputTokens - hasContextTokens = true - } - if usage.CacheReadTokens > 0 { - total += usage.CacheReadTokens - hasContextTokens = true - } - if usage.CacheCreationTokens > 0 { - total += usage.CacheCreationTokens - hasContextTokens = true - } - if !hasContextTokens && usage.TotalTokens > 0 { - total = usage.TotalTokens - } - return total -} - -func (server *Server) afterInterruptionOutcome( - ctx context.Context, - outcome interruptionOutcome, -) error { - chat := outcome.Chat - logger := server.logger.With(slog.F("chat_id", chat.ID), slog.F("owner_id", chat.OwnerID)) - - if outcome.Kind == runnerActionKindFinishInterruption && !chat.ParentChatID.Valid { - server.clearLastTurnSummaryAsync(context.WithoutCancel(ctx), chat, logger) - } - return nil -} - -func (server *Server) afterGenerationOutcome( - ctx context.Context, - outcome generationOutcome, -) error { - chat := outcome.Chat - logger := server.logger.With(slog.F("chat_id", chat.ID), slog.F("owner_id", chat.OwnerID)) - - switch outcome.Kind { - case runnerActionKindFinishTurn: - finalizeCtx := context.WithoutCancel(ctx) - runResult := server.deriveFinalTurnRunResult(finalizeCtx, chat, logger) - server.maybeFinalizeTurnStatusLabelAndPush(finalizeCtx, chat, chat.Status, "", runResult, logger) - case runnerActionKindFinishError: - server.maybeFinalizeTurnStatusLabelAndPush(context.WithoutCancel(ctx), chat, chat.Status, outcome.LastError, runChatResult{}, logger) - case runnerActionKindEnterRequiresAction: - server.maybeFinalizeTurnStatusLabelAndPush(context.WithoutCancel(ctx), chat, chat.Status, "", runChatResult{}, logger) - } - return nil -} - -// deriveFinalTurnRunResult rebuilds the inputs needed to generate the -// end-of-turn status label directly from persisted state. -func (server *Server) deriveFinalTurnRunResult( - ctx context.Context, - chat database.Chat, - logger slog.Logger, -) runChatResult { - // generateFinalTurnStatusLabel only produces a model-generated label for - // the Waiting status, so skip the model resolution and history read - // otherwise. - if chat.Status != database.ChatStatusWaiting { - return runChatResult{} - } - - promptRows, err := server.db.GetChatMessagesForPromptByChatID(ctx, chat.ID) - if err != nil { - logger.Warn(ctx, "derive final turn status label: load prompt rows", slog.Error(err)) - return runChatResult{} - } - triggerMessageID, historyTipMessageID, _ := deriveChatDebugSeed(promptRows) - finalAssistantText := latestAssistantText(promptRows) - if finalAssistantText == "" { - return runChatResult{} - } - - apiKeyID, err := server.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) - if err != nil { - logger.Warn(ctx, "derive final turn status label: ensure synthetic API key", slog.Error(err)) - return runChatResult{FinalAssistantText: finalAssistantText, TriggerMessageID: triggerMessageID, HistoryTipMessageID: historyTipMessageID} - } - modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} - resolved, err := server.resolveModelCall(ctx, modelCallSpec{ - purpose: "turn_status_label", - chat: chat, - buildOptions: modelOpts, - }) - if err != nil { - // Preserve the text and IDs for the generic-label fallback. - logger.Warn(ctx, "derive final turn status label: resolve model", slog.Error(err)) - return runChatResult{ - FinalAssistantText: finalAssistantText, - TriggerMessageID: triggerMessageID, - HistoryTipMessageID: historyTipMessageID, - } - } - - return runChatResult{ - FinalAssistantText: finalAssistantText, - StatusLabelCall: &resolved, - TriggerMessageID: triggerMessageID, - HistoryTipMessageID: historyTipMessageID, - } -} - -// latestAssistantText returns the trimmed text of the most recent assistant -// message. It mirrors the FinalAssistantText that buildCommitStepMessages -// produced from the freshly generated step, making persisted history the -// single source of truth for the turn status label input. -func latestAssistantText(messages []database.ChatMessage) string { - for i := len(messages) - 1; i >= 0; i-- { - if messages[i].Role != database.ChatMessageRoleAssistant { - continue - } - parts, err := chatprompt.ParseContent(messages[i]) - if err != nil { - return "" - } - return strings.TrimSpace(textFromParts(parts)) - } - return "" -} - -// ACLs are deliberately not re-checked: revocation blocks new selection but -// leaves already-selected servers usable, like template ACLs for running -// workspaces. Disabling or deleting the config cuts off existing chats. -func enabledMCPServerConfigsForChatOrg( - ctx context.Context, - db database.Store, - organizationID uuid.UUID, - ids []uuid.UUID, -) ([]database.MCPServerConfig, error) { - configs, err := db.GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx, database.GetEnabledMCPServerConfigsByOrganizationAndIDsParams{ - OrganizationID: organizationID, - IDs: ids, - }) - if err != nil { - return nil, xerrors.Errorf("get enabled MCP server configs for organization: %w", err) - } - return configs, nil -} diff --git a/coderd/x/chatd/generation_preparer_internal_test.go b/coderd/x/chatd/generation_preparer_internal_test.go index e5b9edfcad0..a90abacef08 100644 --- a/coderd/x/chatd/generation_preparer_internal_test.go +++ b/coderd/x/chatd/generation_preparer_internal_test.go @@ -154,26 +154,26 @@ func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) { chatprovider.ProviderAPIKeys{}, withInternalTestServerTransportFactory(&aibridgeTestFactory{}), ) - prepared, err := server.prepareGeneration(ctx, generationPrepareInput{ + prepared, err := server.buildTurnEnvironment(ctx, generationPrepareInput{ Chat: created.Chat, Messages: created.InitialMessages, }) require.NoError(t, err) - t.Cleanup(prepared.Cleanup) + t.Cleanup(prepared.Close) - providerOptions, ok := prepared.CallTemplate.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) - require.True(t, ok, "%T", prepared.CallTemplate.ProviderOptions[fantasyopenai.Name]) + providerOptions, ok := prepared.ModelConfig().callTemplate.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) + require.True(t, ok, "%T", prepared.ModelConfig().callTemplate.ProviderOptions[fantasyopenai.Name]) require.NotNil(t, providerOptions.ReasoningEffort) require.Equal(t, fantasyopenai.ReasoningEffortMedium, *providerOptions.ReasoningEffort) require.NotNil(t, providerOptions.User) require.Equal(t, "turn-options-sentinel", *providerOptions.User) - require.NotNil(t, prepared.CallTemplate.MaxOutputTokens) - require.Equal(t, defaultChatMaxOutputTokens, *prepared.CallTemplate.MaxOutputTokens) + require.NotNil(t, prepared.ModelConfig().callTemplate.MaxOutputTokens) + require.Equal(t, defaultChatMaxOutputTokens, *prepared.ModelConfig().callTemplate.MaxOutputTokens) - require.NotNil(t, prepared.Compaction) - summaryCall := prepared.Compaction.Options.SummaryCall - require.Equal(t, prepared.CallTemplate.ProviderOptions, summaryCall.ProviderOptions) + require.NotNil(t, prepared.CompactionConfig()) + summaryCall := prepared.CompactionConfig().Options.SummaryCall + require.Equal(t, prepared.ModelConfig().callTemplate.ProviderOptions, summaryCall.ProviderOptions) require.NotNil(t, summaryCall.ToolChoice) require.Equal(t, fantasy.ToolChoiceNone, *summaryCall.ToolChoice) // Non-streaming summaries must not inherit the default output cap the @@ -259,24 +259,24 @@ func TestPrepareGenerationComputerUseIgnoresChatTransportOverride(t *testing.T) chatprovider.ProviderAPIKeys{}, withInternalTestServerTransportFactory(&aibridgeTestFactory{}), ) - prepared, err := server.prepareGeneration(ctx, generationPrepareInput{ + prepared, err := server.buildTurnEnvironment(ctx, generationPrepareInput{ Chat: created.Chat, Messages: created.InitialMessages, }) require.NoError(t, err) - t.Cleanup(prepared.Cleanup) + t.Cleanup(prepared.Close) // The computer-use model is Responses-selected by the SDK and its client // ignores the config's forced Chat Completions, so the options must be the // Responses type or the SDK discards them. - _, ok := prepared.CallTemplate.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) - require.True(t, ok, "%T", prepared.CallTemplate.ProviderOptions[fantasyopenai.Name]) + _, ok := prepared.ModelConfig().callTemplate.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) + require.True(t, ok, "%T", prepared.ModelConfig().callTemplate.ProviderOptions[fantasyopenai.Name]) // File classification must also key on the substituted model: the // Responses transport drops native text file parts, so the attachment // must be inlined as text rather than kept as a FilePart. var sawInlinedText bool - for _, message := range prepared.Prompt { + for _, message := range prepared.Prompt() { for _, part := range message.Content { if filePart, isFile := part.(fantasy.FilePart); isFile { t.Fatalf("text attachment survived as FilePart %q", filePart.Filename) @@ -344,19 +344,19 @@ func TestPrepareGenerationSubagentUsesOwnerSyntheticAPIKey(t *testing.T) { chatprovider.ProviderAPIKeys{}, withInternalTestServerTransportFactory(&aibridgeTestFactory{}), ) - prepared, err := server.prepareGeneration(ctx, generationPrepareInput{ + prepared, err := server.buildTurnEnvironment(ctx, generationPrepareInput{ Chat: created.Chat, Messages: created.InitialMessages, }) require.NoError(t, err) - t.Cleanup(prepared.Cleanup) + t.Cleanup(prepared.Close) gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ UserID: user.ID, TokenName: GatewayTokenName(user.ID), }) require.NoError(t, err) - require.Equal(t, gatewayKey.ID, prepared.ModelBuildOptions.ActiveAPIKeyID) + require.Equal(t, gatewayKey.ID, prepared.ModelConfig().buildOptions.ActiveAPIKeyID) } // TestDeriveFinalTurnRunResult exercises the re-derivation path that replaces diff --git a/coderd/x/chatd/toolinput.go b/coderd/x/chatd/toolinput.go index 6e961f72219..6cbe49fe0b1 100644 --- a/coderd/x/chatd/toolinput.go +++ b/coderd/x/chatd/toolinput.go @@ -18,7 +18,7 @@ import ( // of a dispatch failure. allowedIndexes maps allowed calls back to the input // order without relying on duplicate-prone IDs. func partitionAmbiguousToolCalls( - prepared generationPrepared, + prepared turnEnvironment, toolCalls []fantasy.ToolCallContent, ) (allowed []fantasy.ToolCallContent, allowedIndexes []int, rejected []fantasy.ToolResultContent) { for i, toolCall := range toolCalls { @@ -39,7 +39,7 @@ func partitionAmbiguousToolCalls( // validateOverriddenToolInputs rechecks the inputs a pre_tool_use consumer // replaced. The model cannot fix an ambiguous override, so the turn fails // closed instead of executing it. -func validateOverriddenToolInputs(prepared generationPrepared, preflight chathooks.PreToolUseExecutionResult) error { +func validateOverriddenToolInputs(prepared turnEnvironment, preflight chathooks.PreToolUseExecutionResult) error { for _, toolCall := range preflight.Allowed { if _, overridden := preflight.Overrides[toolCall.ToolCallID]; !overridden { continue @@ -54,16 +54,16 @@ func validateOverriddenToolInputs(prepared generationPrepared, preflight chathoo // validateBuiltinToolInput only guards builtin tools, whose input coderd // decodes itself. Dynamic calls are executed by the client and MCP calls by // their own server, and a dynamic tool cannot shadow a builtin name. -func validateBuiltinToolInput(prepared generationPrepared, toolName string, input []byte) error { +func validateBuiltinToolInput(prepared turnEnvironment, toolName string, input []byte) error { // Execution resolves a deprecated alias to its canonical tool, so // skipping that here would let the old name bypass validation. if canonical, aliased := subagentToolNameAliases[toolName]; aliased { toolName = canonical } - if !prepared.BuiltinToolNames[toolName] { + if !prepared.Toolset().builtinToolNames[toolName] { return nil } - for _, tool := range prepared.Tools { + for _, tool := range prepared.Toolset().tools { info := tool.Info() if info.Name != toolName { continue diff --git a/coderd/x/chatd/toolinput_internal_test.go b/coderd/x/chatd/toolinput_internal_test.go index 56b720b87bf..fef0b29ca1c 100644 --- a/coderd/x/chatd/toolinput_internal_test.go +++ b/coderd/x/chatd/toolinput_internal_test.go @@ -36,7 +36,7 @@ func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) { t.Run("builtin", func(t *testing.T) { t.Parallel() - prepared := generationPrepared{ + prepared := turnEnvironmentState{ Tools: []fantasy.AgentTool{fetch}, BuiltinToolNames: map[string]bool{"fetch": true}, } @@ -51,7 +51,7 @@ func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) { t.Run("non-builtin", func(t *testing.T) { t.Parallel() - prepared := generationPrepared{Tools: []fantasy.AgentTool{fetch}} + prepared := turnEnvironmentState{Tools: []fantasy.AgentTool{fetch}} allowed, allowedIndexes, rejected := partitionAmbiguousToolCalls(prepared, []fantasy.ToolCallContent{ambiguous, clean}) require.Empty(t, rejected) require.Len(t, allowed, 2) @@ -83,7 +83,7 @@ func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) { Input: `{"chat_id":"a","CHAT_ID":"b"}`, } - prepared := generationPrepared{ + prepared := turnEnvironmentState{ Tools: []fantasy.AgentTool{tool}, BuiltinToolNames: map[string]bool{canonical: true}, } @@ -95,7 +95,7 @@ func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) { func TestValidateOverriddenToolInputs(t *testing.T) { t.Parallel() - prepared := generationPrepared{ + prepared := turnEnvironmentState{ Tools: []fantasy.AgentTool{fetchToolStub()}, BuiltinToolNames: map[string]bool{"fetch": true}, } @@ -165,12 +165,12 @@ func TestBuiltinToolSchemasDescribeTheirInputs(t *testing.T) { chatprovider.ProviderAPIKeys{}, withInternalTestServerTransportFactory(&aibridgeTestFactory{}), ) - prepared, err := server.prepareGeneration(ctx, generationPrepareInput{ + prepared, err := server.buildTurnEnvironment(ctx, generationPrepareInput{ Chat: created.Chat, Messages: created.InitialMessages, }) require.NoError(t, err) - t.Cleanup(prepared.Cleanup) + t.Cleanup(prepared.Close) // These take an empty struct, so they carry no keys to validate. noInput := map[string]bool{ @@ -179,10 +179,10 @@ func TestBuiltinToolSchemasDescribeTheirInputs(t *testing.T) { "list_subagent_models": true, } var unvalidated []string - require.NotEmpty(t, prepared.BuiltinToolNames) - for _, tool := range prepared.Tools { + require.NotEmpty(t, prepared.Toolset().builtinToolNames) + for _, tool := range prepared.Toolset().tools { info := tool.Info() - if !prepared.BuiltinToolNames[info.Name] || len(info.Parameters) > 0 || noInput[info.Name] { + if !prepared.Toolset().builtinToolNames[info.Name] || len(info.Parameters) > 0 || noInput[info.Name] { continue } unvalidated = append(unvalidated, info.Name) diff --git a/coderd/x/chatd/turn_environment.go b/coderd/x/chatd/turn_environment.go new file mode 100644 index 00000000000..acbc00457e0 --- /dev/null +++ b/coderd/x/chatd/turn_environment.go @@ -0,0 +1,2226 @@ +package chatd + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "net/http" + "slices" + "strings" + "sync" + "time" + + "charm.land/fantasy" + "charm.land/fantasy/providers/anthropic" + "github.com/google/uuid" + "github.com/sqlc-dev/pqtype" + "golang.org/x/sync/errgroup" + "golang.org/x/xerrors" + + "cdr.dev/slog/v3" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" + "github.com/coder/coder/v2/coderd/rbac" + "github.com/coder/coder/v2/coderd/x/chatd/agentselect" + "github.com/coder/coder/v2/coderd/x/chatd/chatadvisor" + "github.com/coder/coder/v2/coderd/x/chatd/chatloop" + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" + "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" + "github.com/coder/coder/v2/coderd/x/chatd/chatsanitize" + "github.com/coder/coder/v2/coderd/x/chatd/chattool" + "github.com/coder/coder/v2/coderd/x/chatd/mcpclient" + skillspkg "github.com/coder/coder/v2/coderd/x/skills" + "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/codersdk/workspacesdk" +) + +// effectiveMCPServerConfigs loads the chat's stored selection plus +// owner-readable Force On configs at generation time, so stored lists +// predating enforcement cannot dodge the policy (Cure53 CDM-02-010). +// Explore chats keep their immutable spawn-time snapshot instead. +type turnEnvironment interface { + Turn() turnState + ModelConfig() turnModelConfig + Prompt() []fantasy.Message + Toolset() turnToolset + CompactionConfig() *generationCompaction + Close() +} + +type turnState struct { + chat database.Chat + messages []database.ChatMessage + maxSteps int + debug *generationDebug +} + +type turnModelConfig struct { + model chatprovider.Model + route aiGatewayModelRoute + buildOptions modelBuildOptions + resolvedProvider string + configID uuid.UUID + callTemplate fantasy.Call + contextLimitFallback int64 +} + +type turnToolset struct { + tools []fantasy.AgentTool + activeTools []string + allowInactiveTools map[string]bool + providerTools []chatloop.ProviderTool + dynamicToolNames map[string]bool + stopAfterTools map[string]struct{} + exclusiveToolNames map[string]bool + builtinToolNames map[string]bool + toolNameToConfigID map[string]uuid.UUID +} + +type turnEnvironmentState struct { + Chat database.Chat + Messages []database.ChatMessage + + Model chatprovider.Model + PromptMessages []fantasy.Message + Tools []fantasy.AgentTool + ActiveTools []string + AllowInactiveTools map[string]bool + ProviderTools []chatloop.ProviderTool + ModelRoute aiGatewayModelRoute + ModelBuildOptions modelBuildOptions + ResolvedProvider string + ModelConfigID uuid.UUID + CallTemplate fantasy.Call + ContextLimitFallback int64 + + DynamicToolNames map[string]bool + StopAfterTools map[string]struct{} + ExclusiveToolNames map[string]bool + BuiltinToolNames map[string]bool + ToolNameToConfigID map[string]uuid.UUID + + MaxSteps int + Compaction *generationCompaction + Cleanup func() + Debug *generationDebug +} + +func (e turnEnvironmentState) Turn() turnState { + return turnState{chat: e.Chat, messages: e.Messages, maxSteps: e.MaxSteps, debug: e.Debug} +} + +func (e turnEnvironmentState) ModelConfig() turnModelConfig { + return turnModelConfig{ + model: e.Model, route: e.ModelRoute, buildOptions: e.ModelBuildOptions, + resolvedProvider: e.ResolvedProvider, configID: e.ModelConfigID, + callTemplate: e.CallTemplate, contextLimitFallback: e.ContextLimitFallback, + } +} + +func (e turnEnvironmentState) Prompt() []fantasy.Message { return e.PromptMessages } + +func (e turnEnvironmentState) Toolset() turnToolset { + return turnToolset{ + tools: e.Tools, activeTools: e.ActiveTools, allowInactiveTools: e.AllowInactiveTools, + providerTools: e.ProviderTools, dynamicToolNames: e.DynamicToolNames, + stopAfterTools: e.StopAfterTools, exclusiveToolNames: e.ExclusiveToolNames, + builtinToolNames: e.BuiltinToolNames, toolNameToConfigID: e.ToolNameToConfigID, + } +} + +func (e turnEnvironmentState) CompactionConfig() *generationCompaction { return e.Compaction } +func (e turnEnvironmentState) Close() { e.Cleanup() } + +func (server *Server) effectiveMCPServerConfigs( + ctx context.Context, + logger slog.Logger, + chat database.Chat, +) ([]database.MCPServerConfig, error) { + var configs []database.MCPServerConfig + if len(chat.MCPServerIDs) > 0 { + var err error + configs, err = enabledMCPServerConfigsForChatOrg(ctx, server.db, chat.OrganizationID, chat.MCPServerIDs) + if err != nil { + // Best-effort for the user-selected set, matching prior + // behavior: a load failure degrades the turn rather than + // failing it. + logger.Warn(ctx, "failed to load MCP server configs", slog.Error(err)) + configs = nil + } + } + if isExploreSubagentMode(chat.Mode) { + return configs, nil + } + forced, err := forcedMCPServerConfigsForOwner(ctx, server.db, chat.OrganizationID, chat.OwnerID) + if err != nil { + // Fail closed: running the turn without the forced set would + // silently bypass a security policy. + return nil, err + } + seen := make(map[uuid.UUID]struct{}, len(configs)) + for _, cfg := range configs { + seen[cfg.ID] = struct{}{} + } + for _, cfg := range forced { + if _, ok := seen[cfg.ID]; !ok { + configs = append(configs, cfg) + } + } + return configs, nil +} + +func (server *Server) buildTurnEnvironment( + ctx context.Context, + input generationPrepareInput, +) (turnEnvironment, error) { + chat := input.Chat + logger := server.logger.With( + slog.F("chat_id", chat.ID), + slog.F("owner_id", chat.OwnerID), + ) + + prepStart := server.clock.Now() + defer func() { + if prepDuration := server.clock.Since(prepStart); prepDuration >= slowPrepareThreshold { + logger.Warn(ctx, "slow generation preparation", + slog.F("duration", prepDuration), + ) + } + }() + + var ( + promptRows []database.ChatMessage + mcpConfigs []database.MCPServerConfig + mcpTokens []database.MCPServerUserToken + ) + + var g errgroup.Group + g.Go(func() error { + var err error + promptRows, err = server.db.GetChatMessagesForPromptByChatID(ctx, chat.ID) + if err != nil { + return xerrors.Errorf("get chat messages for prompt: %w", err) + } + return nil + }) + g.Go(func() error { + var err error + mcpConfigs, err = server.effectiveMCPServerConfigs(ctx, logger, chat) + return err + }) + if err := g.Wait(); err != nil { + return nil, err + } + + apiKeyID, err := server.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) + if err != nil { + return nil, xerrors.Errorf("ensure synthetic API key: %w", err) + } + modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} + + requestedEffort := chatRequestedEffort(chat) + resolved, err := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "standard_turn", + chat: chat, + requestedEffort: requestedEffort, + buildOptions: modelOpts, + }) + if err != nil { + return nil, err + } + // The chat config keeps driving compaction, sanitization, and debug + // attribution even when computer use swaps the resolved call below. + modelConfig := resolved.dbConfig + + // Computer-use turns swap in a specialized model, so the substitution + // must happen before anything model-sensitive runs: file-part + // classification, history sanitization, and provider option preparation + // must all agree with the client actually used for the turn. + isComputerUse := chat.Mode.Valid && chat.Mode.ChatMode == database.ChatModeComputerUse + var computerUseProvider codersdk.ChatComputerUseProvider + if isComputerUse { + var cuModelProvider, cuModelName string + computerUseProvider, cuModelProvider, cuModelName, err = server.computerUseProviderAndModelFromConfig(ctx) + if err != nil { + return nil, xerrors.Errorf("resolve computer use provider and model: %w", err) + } + cuResolved, cuErr := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "computer_use", + chat: chat, + fixedModel: &fixedModelCall{ + providerType: cuModelProvider, + modelName: cuModelName, + callConfig: resolved.callConfig, + }, + requestedEffort: requestedEffort, + buildOptions: modelOpts, + }) + if cuErr != nil { + return nil, xerrors.Errorf( + "resolve computer use model for provider %q model %q: %w", + computerUseProvider, + cuModelName, + cuErr, + ) + } + resolved = cuResolved + } + model := resolved.model + callConfig := resolved.callConfig + modelRoute := resolved.route + + currentPlanMode := chat.PlanMode + isPlanModeTurn := currentPlanMode.Valid && currentPlanMode.ChatPlanMode == database.ChatPlanModePlan + isExploreSubagent := isExploreSubagentMode(chat.Mode) + isRootChat := !chat.ParentChatID.Valid + + mcpConnectConfigs, approvedPlanMCPConfigIDs := filterExternalMCPConfigsForTurn( + mcpConfigs, + currentPlanMode, + chat.ParentChatID, + ) + if isExploreSubagent && isRootChat { + mcpConnectConfigs = nil + approvedPlanMCPConfigIDs = map[uuid.UUID]struct{}{} + } + + planModeInstructions := server.loadPlanModeInstructions(ctx, currentPlanMode, logger) + advisorCfg := server.loadAdvisorConfig(ctx, logger) + // Force Enabled from the experiment; the stored DB value is ignored. + advisorCfg.Enabled = server.experiments.Enabled(codersdk.ExperimentChatAdvisor) + + var advisorRuntime *chatadvisor.Runtime + if advisorCfg.Enabled && isRootChat && !isPlanModeTurn && !isExploreSubagent { + var advisorErr error + advisorRuntime, advisorErr = server.newAdvisorRuntime( + ctx, + chat, + advisorCfg, + modelOpts, + logger, + ) + if advisorErr != nil { + return nil, advisorErr + } + } + + var advisorPromptSnapshot []fantasy.Message + setAdvisorPromptSnapshot := func(msgs []fantasy.Message) { + if advisorRuntime == nil { + return + } + advisorPromptSnapshot = slices.Clone(msgs) + } + + currentChat := chat + loadChatSnapshot := func(loadCtx context.Context, chatID uuid.UUID) (database.Chat, error) { + return server.db.GetChatByID(loadCtx, chatID) + } + var chatStateMu sync.Mutex + var workspaceMu sync.Mutex + workspaceCtx := turnWorkspaceContext{ + server: server, + chatStateMu: &chatStateMu, + currentChat: ¤tChat, + loadChatSnapshot: loadChatSnapshot, + } + cleanup := func() { + workspaceCtx.close() + } + + planPathFn := func(ctx context.Context) (string, string, error) { + conn, err := workspaceCtx.getWorkspaceConn(ctx) + if err != nil { + return "", "", err + } + home, err := chattool.ResolveWorkspaceHome(ctx, conn) + if err != nil { + return "", "", err + } + return chattool.PlanPathForChat(home, chat.ID), home, nil + } + resolvePlanPathForTools := func(ctx context.Context) (string, string, error) { + planCtx, cancel := context.WithTimeout(ctx, planPathLookupTimeout) + defer cancel() + return planPathFn(planCtx) + } + resolvePlanPathBlock := func(resolveCtx context.Context) string { + if chat.ParentChatID.Valid { + return "" + } + + planCtx, cancel := context.WithTimeout(resolveCtx, planPathLookupTimeout) + defer cancel() + + if _, _, err := workspaceCtx.workspaceAgentIDForConn(planCtx); err != nil { + logger.Debug(resolveCtx, "plan path instruction: agent not reachable", + slog.Error(err), + slog.F("chat_id", chat.ID), + ) + return "" + } + + planPath, home, err := planPathFn(planCtx) + if err != nil { + logger.Debug(resolveCtx, "plan path instruction: failed to resolve plan path", + slog.Error(err), + slog.F("chat_id", chat.ID), + ) + return "" + } + return formatPlanPathBlock(planPath, home) + } + + var ( + prompt []fantasy.Message + instruction string + mcpTools []fantasy.AgentTool + mcpSummaries []mcpclient.ConnectSummary + mcpCleanup func() + workspaceMCPTools []fantasy.AgentTool + workspaceSkills []chattool.SkillMeta + personalSkills []skillspkg.Skill + resolvedUserPrompt string + planPathBlock string + ) + + // Drop provider-executed tool history produced by a different provider + // before building the prompt. A provider that shares another's wire format + // (e.g. Bedrock and Anthropic) can still reject the other's + // provider-executed blocks, so a mid-chat provider switch must not replay + // them. + promptRows = server.sanitizeForeignProviderExecutedToolRows(ctx, logger, promptRows, chat.OwnerID, modelConfig.ID) + + if chat.WorkspaceID.Valid { + // Resolve the workspace agent so the chat row's AgentID and + // BuildID bindings are up to date before the chatworker + // decision helper inspects them. ensureWorkspaceAgent does a + // DB lookup and lazily calls persistBuildAgentBinding when + // the bound agent has changed, so this is a cheap metadata + // refresh, not a workspace dial. It must not insert chat + // history; only metadata is mutated here. + agent, _ := workspaceCtx.getWorkspaceAgent(ctx) + + // API-created chats bind their agent lazily here, after + // hydrateChatContextOnCreate ran with no agent. Pin the chat to the + // bound agent's pushed snapshot now if it is still unpinned, so the + // first turn reads workspace context instead of waiting for the + // agent's next push. Idempotent and snapshot-gated; runs before the + // pinned context is read below. + server.ensureChatContextPinnedOnFirstTurn(ctx, workspaceCtx.currentChatSnapshot()) + + var resolveErr error + instruction, workspaceSkills, resolveErr = server.resolveTurnWorkspaceContext(ctx, chat, agent) + if resolveErr != nil { + cleanup() + return nil, resolveErr + } + } + + // Build the debug context before the connect phase so its + // outcomes can be recorded with run-creation context even when a + // later preparation step fails. + triggerMessageID, historyTipMessageID, triggerLabel := deriveChatDebugSeed(promptRows) + debugSvc := server.existingDebugService() + var debug *generationDebug + if resolved.debugEnabled { + if debugSvc == nil { + cleanup() + return nil, xerrors.New("chat debug service missing after enablement check") + } + debug = &generationDebug{ + Enabled: true, + Service: debugSvc, + Provider: resolved.resolvedProvider, + Model: resolved.resolvedModel, + TriggerMessageID: triggerMessageID, + HistoryTipMessageID: historyTipMessageID, + TriggerLabel: triggerLabel, + ModelConfig: modelConfig, + } + } + + var g2 errgroup.Group + g2.Go(func() error { + var err error + // Key the file-part acceptance on model.Provider() (the fantasy + // transport identity), not the configured provider, because + // aibridge routing rewrites the provider (e.g. Bedrock to the + // Anthropic transport). The conversion that actually drops or + // accepts a file part is the one for model.Provider(). + acceptsFilePart := model.AcceptsFilePartMediaType + providerType := string(modelRoute.Provider.Type) + prompt, err = chatprompt.ConvertMessagesWithFiles(ctx, promptRows, server.chatFileResolver(providerType), logger, acceptsFilePart) + if err != nil { + return xerrors.Errorf("build chat prompt: %w", err) + } + return nil + }) + g2.Go(func() error { + personalSkills = server.fetchPersonalSkillMetadata(ctx, chat.OwnerID, logger) + return nil + }) + g2.Go(func() error { + resolvedUserPrompt = server.resolveUserPrompt(ctx, chat.OwnerID) + return nil + }) + if len(mcpConnectConfigs) > 0 { + g2.Go(func() error { + var tokenErr error + mcpTokens, tokenErr = server.db.GetMCPServerUserTokensByUserID(ctx, chat.OwnerID) + if tokenErr != nil { + logger.Warn(ctx, "failed to load MCP user tokens", slog.Error(tokenErr)) + } + mcpTokens = server.refreshExpiredMCPTokens(ctx, logger, mcpConnectConfigs, mcpTokens) + mcpTools, mcpSummaries, mcpCleanup = mcpclient.ConnectAll( + ctx, + logger, + mcpConnectConfigs, + mcpTokens, + chat.OwnerID, + server.oidcTokenSource, + chatprovider.CoderHeaders(chat), + ) + return nil + }) + } + if chat.WorkspaceID.Valid && !isPlanModeTurn && !isExploreSubagent { + g2.Go(func() error { + workspaceMCPTools = server.resolveWorkspaceMCPTools(ctx, logger, chat, &workspaceCtx) + return nil + }) + } + // Resolve the per-chat plan path block in the parallel phase. It dials + // the workspace agent to read the home directory, so running it here lets + // the cold dial overlap with the rest of turn preparation instead of + // blocking system prompt assembly on a sequential dial. Best-effort: + // resolvePlanPathBlock logs and returns an empty block on failure. + if chat.WorkspaceID.Valid && !chat.ParentChatID.Valid { + g2.Go(func() error { + planPathBlock = resolvePlanPathBlock(ctx) + return nil + }) + } + g2Err := g2.Wait() + // Record connect outcomes before acting on any preparation error: + // ConnectAll has already run, so a failure below (or in g2 itself) + // would otherwise discard this attempt's outcomes. + if debug != nil && input.RecordMCPConnectSummaries != nil && len(mcpSummaries) > 0 { + input.RecordMCPConnectSummaries(ctx, chat, debug, mcpSummaries) + } + if g2Err != nil { + cleanup() + return nil, g2Err + } + + if mcpCleanup != nil { + previousCleanup := cleanup + cleanup = func() { + mcpCleanup() + previousCleanup() + } + } + + prompt, sanitizeStats := chatsanitize.SanitizeAnthropicProviderToolHistory(model.Provider(), prompt) + chatsanitize.LogAnthropicProviderToolSanitization( + ctx, + logger, + "persisted_history_replay", + model.Provider(), + model.ModelID(), + sanitizeStats, + ) + + subagentInstruction := "" + if !isRootChat { + subagentInstruction = defaultSubagentInstruction + } + resolvedSkillsFor := func(workspaceSkills []chattool.SkillMeta) []skillspkg.ResolvedSkill { + return mergeTurnSkills(personalSkills, workspaceSkills) + } + resolveSkillAlias := func(alias string) (skillspkg.ResolvedSkill, error) { + return skillspkg.Lookup(resolvedSkillsFor(workspaceSkills), alias) + } + initialResolvedSkills := resolvedSkillsFor(workspaceSkills) + + prompt = buildSystemPrompt( + prompt, + subagentInstruction, + instruction, + initialResolvedSkills, + resolvedUserPrompt, + systemPromptBehaviorContext{ + planMode: currentPlanMode, + chatMode: chat.Mode, + planModeInstructions: planModeInstructions, + isRootChat: isRootChat, + }, + ) + if advisorRuntime != nil { + prompt = chatprompt.InsertSystem(prompt, chatadvisor.ParentGuidanceBlock) + } + prompt = renderPlanPathPrompt(prompt, planPathBlock) + setAdvisorPromptSnapshot(prompt) + + storeChatAttachment := server.newStoreChatAttachmentFunc(&workspaceCtx) + tools := []fantasy.AgentTool{ + chattool.ReadFile(chattool.ReadFileOptions{GetWorkspaceConn: workspaceCtx.getWorkspaceConn}), + chattool.WriteFile(chattool.WriteFileOptions{ + GetWorkspaceConn: workspaceCtx.getWorkspaceConn, + ResolvePlanPath: resolvePlanPathForTools, + IsPlanTurn: isPlanModeTurn, + }), + chattool.EditFiles(chattool.EditFilesOptions{ + GetWorkspaceConn: workspaceCtx.getWorkspaceConn, + ResolvePlanPath: resolvePlanPathForTools, + IsPlanTurn: isPlanModeTurn, + }), + chattool.AttachFile(chattool.AttachFileOptions{ + GetWorkspaceConn: workspaceCtx.getWorkspaceConn, + StoreFile: storeChatAttachment, + }), + chattool.Execute(chattool.ExecuteOptions{ + GetWorkspaceConn: workspaceCtx.getWorkspaceConn, + AgentBrowserSession: chat.ID.String(), + }), + chattool.ProcessOutput(chattool.ProcessToolOptions{GetWorkspaceConn: workspaceCtx.getWorkspaceConn}), + chattool.ProcessList(chattool.ProcessToolOptions{GetWorkspaceConn: workspaceCtx.getWorkspaceConn}), + chattool.ProcessSignal(chattool.ProcessToolOptions{GetWorkspaceConn: workspaceCtx.getWorkspaceConn}), + } + if isPlanModeTurn && isRootChat { + tools = append(tools, chattool.NewAskUserQuestionTool()) + } + if isRootChat { + tools = server.appendRootChatTools(ctx, tools, rootChatToolsOptions{ + chat: chat, + modelConfigID: modelConfig.ID, + workspaceCtx: &workspaceCtx, + workspaceMu: &workspaceMu, + resolvePlanPath: resolvePlanPathForTools, + storeFile: storeChatAttachment, + isPlanModeTurn: isPlanModeTurn, + }) + } + + skillOpts := chattool.ReadSkillOptions{ + GetWorkspaceConn: workspaceCtx.getWorkspaceConn, + GetSkills: func() []chattool.SkillMeta { + return workspaceSkills + }, + ResolveAlias: resolveSkillAlias, + LoadPersonalSkillBody: func(ctx context.Context, name string) (skillspkg.ParsedSkill, error) { + return server.loadPersonalSkillBody(ctx, chat.OwnerID, name) + }, + } + appendCurrentSkillTools := func(current []fantasy.AgentTool) ([]fantasy.AgentTool, bool) { + if len(personalSkills) == 0 && len(workspaceSkills) == 0 { + return current, false + } + updated := current + changed := false + appendTool := func(tool fantasy.AgentTool) { + name := tool.Info().Name + if slices.ContainsFunc(current, func(existing fantasy.AgentTool) bool { + return existing.Info().Name == name + }) { + return + } + if !changed { + updated = slices.Clone(current) + changed = true + } + updated = append(updated, tool) + } + appendTool(chattool.ReadSkill(skillOpts)) + if len(workspaceSkills) > 0 { + appendTool(chattool.ReadSkillFile(skillOpts)) + } + return updated, changed + } + tools, _ = appendCurrentSkillTools(tools) + if advisorRuntime != nil { + tools = append(tools, chatadvisor.Tool(chatadvisor.ToolOptions{ + Runtime: advisorRuntime, + GetConversationSnapshot: func() []fantasy.Message { + return stripAdvisorGuidanceBlock(slices.Clone(advisorPromptSnapshot)) + }, + })) + } + + var exclusiveToolNames map[string]bool + if advisorRuntime != nil { + exclusiveToolNames = map[string]bool{chatadvisor.ToolName: true} + } + + builtinToolNames := make(map[string]bool, len(tools)) + for _, t := range tools { + builtinToolNames[t.Info().Name] = true + } + + mcpConfigByID := make(map[uuid.UUID]database.MCPServerConfig, len(mcpConnectConfigs)) + for _, config := range mcpConnectConfigs { + mcpConfigByID[config.ID] = config + } + deferredCandidates := collectDeferredMCPCandidates(deferredMCPCandidateInput{ + mcpTools: mcpTools, + workspaceMCPTools: workspaceMCPTools, + mcpConfigByID: mcpConfigByID, + planMode: currentPlanMode, + parentChatID: chat.ParentChatID, + approvedMCPConfigIDs: approvedPlanMCPConfigIDs, + includeWorkspaceTools: !isExploreSubagent, + }) + tools = append(tools, mcpTools...) + if !isExploreSubagent { + tools = append(tools, workspaceMCPTools...) + } + tools = filterToolsForTurn(tools, currentPlanMode, chat.ParentChatID, approvedPlanMCPConfigIDs) + + tools, dynamicToolNames, err := appendDynamicTools(ctx, logger, tools, chat.DynamicTools, currentPlanMode, chat.Mode) + if err != nil { + cleanup() + return nil, err + } + + var providerTools []chatloop.ProviderTool + if !isPlanModeTurn && callConfig.ProviderOptions != nil { + providerTools = buildProviderTools(callConfig.ProviderOptions) + if isExploreSubagent { + if !chat.ParentChatID.Valid { + providerTools = nil + } else { + providerTools = slices.DeleteFunc(providerTools, func(tool chatloop.ProviderTool) bool { + return tool.Definition.GetName() != "web_search" + }) + } + } + } + + if isComputerUse { + providerTools, err = appendComputerUseProviderTool(providerTools, computerUseProviderToolOptions{ + provider: computerUseProvider, + isPlanModeTurn: isPlanModeTurn, + isComputerUse: isComputerUse, + getWorkspaceConn: workspaceCtx.getWorkspaceConn, + storeFile: storeChatAttachment, + clock: server.clock, + logger: server.logger.Named("computer_use"), + }) + if err != nil { + cleanup() + return nil, xerrors.Errorf("register computer use provider tool for provider %q: %w", computerUseProvider, err) + } + } else { + providerTools, err = appendComputerUseProviderTool(providerTools, computerUseProviderToolOptions{ + isPlanModeTurn: isPlanModeTurn, + isComputerUse: false, + }) + if err != nil { + cleanup() + return nil, err + } + } + + activeToolNames := activeToolNamesForTurn(tools, currentPlanMode, chat.ParentChatID, approvedPlanMCPConfigIDs) + if isExploreSubagent { + activeToolNames = allowedExploreToolNames(tools) + } + var allowInactiveTools map[string]bool + if decideMCPToolSearch(mcpToolSearchInput{ + experimentEnabled: server.experiments.Enabled(codersdk.ExperimentMCPToolSearch), + candidates: deferredCandidates, + dynamicToolNames: dynamicToolNames, + }) { + activationTokenBudget := float64(modelConfig.ContextLimit) / mcpToolSearchBudgetDivisor + findTools := chattool.FindTools(chattool.FindToolsOptions{ + Entries: deferredMCPToolEntries(deferredCandidates), + SchemaTokenBudget: activationTokenBudget, + CatalogTokenBudget: activationTokenBudget, + // Calls total is counted in executeLocalTools, which also + // sees calls rejected before the tool runs; OnCall covers + // only calls that reach the handler or its decode. + OnCall: func(callCtx context.Context, call chattool.FindToolsCall) { + if call.Rejection == "" { + server.metrics.FindToolsMatchCount.Observe(float64(call.MatchCount)) + server.metrics.FindToolsActivationsTotal.Add(float64(len(call.Activated))) + if call.MatchCount == 0 { + server.metrics.FindToolsEmptyTotal.Inc() + } + } + // Queries and names are model output that can echo + // prompt content, so standard logs carry only + // aggregate fields; raw values are visible through + // the opt-in chat debug logging path. + logger.Info(callCtx, "deferred MCP tool search", + slog.F("query_count", len(call.Queries)), + slog.F("name_count", len(call.Names)), + slog.F("match_count", call.MatchCount), + slog.F("activated_count", len(call.Activated)), + slog.F("total_deferred", call.TotalDeferred), + slog.F("rejection", call.Rejection), + ) + }, + }) + tools, activeToolNames, allowInactiveTools = configureDeferredMCPToolSearch( + tools, + activeToolNames, + deferredCandidates, + findTools, + deriveDeferredMCPActivations(promptRows, deferredCandidates, activationTokenBudget), + ) + builtinToolNames[chattool.FindToolsName] = true + } + + toolNameToConfigID := make(map[string]uuid.UUID) + for _, t := range tools { + if mcpTool, ok := t.(mcpclient.MCPToolIdentifier); ok { + toolNameToConfigID[t.Info().Name] = mcpTool.MCPServerConfigID() + } + } + + compactionToolCallID := "chat_summarized_" + uuid.NewString() + effectiveThreshold := modelConfig.CompressionThreshold + if override, ok := server.resolveUserCompactionThreshold(ctx, chat.OwnerID, modelConfig.ID); ok { + effectiveThreshold = override + } + // The compaction trigger uses the stricter of the chat and override + // models' context limits: the history must also fit the summarizer's + // window. + compactionContextLimit := modelConfig.ContextLimit + compactionOverride, err := server.resolveCompactionOverrideConfig(ctx, chat) + if err != nil { + cleanup() + return nil, err + } + if compactionOverride != nil { + if overrideLimit := compactionOverride.Config.ContextLimit; overrideLimit > 0 && + (compactionContextLimit <= 0 || overrideLimit < compactionContextLimit) { + compactionContextLimit = overrideLimit + } + } + compactionStepUsage := latestPromptUsage(promptRows) + compactionNeeded := shouldCompactPromptUsage(compactionStepUsage, compactionContextLimit, effectiveThreshold) + // The options carry the chat model; generateCompaction swaps in the + // override client when one is configured. + compactionOptions := chatloop.GenerateCompactionOptions{ + Model: model.LanguageModel(), + Messages: prompt, + ThresholdPercent: effectiveThreshold, + ContextLimit: compactionContextLimit, + ContextLimitFallback: compactionContextLimit, + ToolCallID: compactionToolCallID, + ToolName: "chat_summarized", + DebugSvc: debugSvc, + ChatID: chat.ID, + HistoryTipMessageID: historyTipMessageID, + ResolvedProvider: resolved.resolvedProvider, + ResolvedModel: resolved.resolvedModel, + ModelConfigID: modelConfig.ID, + StepUsage: compactionStepUsage, + SummaryCall: compactionSummaryCall(resolved), + } + + // workspaceCtx.currentChatSnapshot may carry a freshly persisted + // AgentID/BuildID binding from the getWorkspaceAgent call above. + // Return that snapshot so downstream consumers see the up-to-date + // metadata. + refreshedChat := workspaceCtx.currentChatSnapshot() + if refreshedChat.ID == uuid.Nil { + refreshedChat = chat + } + + return turnEnvironmentState{ + Chat: refreshedChat, + Messages: input.Messages, + Model: model, + PromptMessages: prompt, + Tools: tools, + ActiveTools: activeToolNames, + AllowInactiveTools: allowInactiveTools, + ProviderTools: providerTools, + ModelRoute: modelRoute, + ModelBuildOptions: modelOpts, + ResolvedProvider: resolved.resolvedProvider, + ModelConfigID: modelConfig.ID, + CallTemplate: resolved.newCall(), + ContextLimitFallback: modelConfig.ContextLimit, + DynamicToolNames: dynamicToolNames, + StopAfterTools: stopAfterBehaviorTools(currentPlanMode, chat.Mode, chat.ParentChatID), + ExclusiveToolNames: exclusiveToolNames, + BuiltinToolNames: builtinToolNames, + ToolNameToConfigID: toolNameToConfigID, + MaxSteps: maxChatSteps, + Compaction: &generationCompaction{ + Override: compactionOverride, + ChatModelConfig: modelConfig, + Required: compactionNeeded, + Options: compactionOptions, + }, + Cleanup: cleanup, + Debug: debug, + }, nil +} + +func latestPromptUsage(messages []database.ChatMessage) fantasy.Usage { + for i := len(messages) - 1; i >= 0; i-- { + usage := usageFromMessage(messages[i]) + if usage != (fantasy.Usage{}) { + return usage + } + } + return fantasy.Usage{} +} + +func shouldCompactPromptUsage(usage fantasy.Usage, contextLimit int64, thresholdPercent int32) bool { + if thresholdPercent >= 100 || contextLimit <= 0 { + return false + } + contextTokens := contextTokensFromUsage(usage) + if contextTokens <= 0 { + return false + } + usagePercent := (float64(contextTokens) / float64(contextLimit)) * 100 + return usagePercent >= float64(thresholdPercent) +} + +func contextTokensFromUsage(usage fantasy.Usage) int64 { + total := int64(0) + hasContextTokens := false + if usage.InputTokens > 0 { + total += usage.InputTokens + hasContextTokens = true + } + if usage.CacheReadTokens > 0 { + total += usage.CacheReadTokens + hasContextTokens = true + } + if usage.CacheCreationTokens > 0 { + total += usage.CacheCreationTokens + hasContextTokens = true + } + if !hasContextTokens && usage.TotalTokens > 0 { + total = usage.TotalTokens + } + return total +} + +func (server *Server) afterInterruptionOutcome( + ctx context.Context, + outcome interruptionOutcome, +) error { + chat := outcome.Chat + logger := server.logger.With(slog.F("chat_id", chat.ID), slog.F("owner_id", chat.OwnerID)) + + if outcome.Kind == runnerActionKindFinishInterruption && !chat.ParentChatID.Valid { + server.clearLastTurnSummaryAsync(context.WithoutCancel(ctx), chat, logger) + } + return nil +} + +func (server *Server) afterGenerationOutcome( + ctx context.Context, + outcome generationOutcome, +) error { + chat := outcome.Chat + logger := server.logger.With(slog.F("chat_id", chat.ID), slog.F("owner_id", chat.OwnerID)) + + switch outcome.Kind { + case runnerActionKindFinishTurn: + finalizeCtx := context.WithoutCancel(ctx) + runResult := server.deriveFinalTurnRunResult(finalizeCtx, chat, logger) + server.maybeFinalizeTurnStatusLabelAndPush(finalizeCtx, chat, chat.Status, "", runResult, logger) + case runnerActionKindFinishError: + server.maybeFinalizeTurnStatusLabelAndPush(context.WithoutCancel(ctx), chat, chat.Status, outcome.LastError, runChatResult{}, logger) + case runnerActionKindEnterRequiresAction: + server.maybeFinalizeTurnStatusLabelAndPush(context.WithoutCancel(ctx), chat, chat.Status, "", runChatResult{}, logger) + } + return nil +} + +// deriveFinalTurnRunResult rebuilds the inputs needed to generate the +// end-of-turn status label directly from persisted state. +func (server *Server) deriveFinalTurnRunResult( + ctx context.Context, + chat database.Chat, + logger slog.Logger, +) runChatResult { + // generateFinalTurnStatusLabel only produces a model-generated label for + // the Waiting status, so skip the model resolution and history read + // otherwise. + if chat.Status != database.ChatStatusWaiting { + return runChatResult{} + } + + promptRows, err := server.db.GetChatMessagesForPromptByChatID(ctx, chat.ID) + if err != nil { + logger.Warn(ctx, "derive final turn status label: load prompt rows", slog.Error(err)) + return runChatResult{} + } + triggerMessageID, historyTipMessageID, _ := deriveChatDebugSeed(promptRows) + finalAssistantText := latestAssistantText(promptRows) + if finalAssistantText == "" { + return runChatResult{} + } + + apiKeyID, err := server.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) + if err != nil { + logger.Warn(ctx, "derive final turn status label: ensure synthetic API key", slog.Error(err)) + return runChatResult{FinalAssistantText: finalAssistantText, TriggerMessageID: triggerMessageID, HistoryTipMessageID: historyTipMessageID} + } + modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} + resolved, err := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "turn_status_label", + chat: chat, + buildOptions: modelOpts, + }) + if err != nil { + // Preserve the text and IDs for the generic-label fallback. + logger.Warn(ctx, "derive final turn status label: resolve model", slog.Error(err)) + return runChatResult{ + FinalAssistantText: finalAssistantText, + TriggerMessageID: triggerMessageID, + HistoryTipMessageID: historyTipMessageID, + } + } + + return runChatResult{ + FinalAssistantText: finalAssistantText, + StatusLabelCall: &resolved, + TriggerMessageID: triggerMessageID, + HistoryTipMessageID: historyTipMessageID, + } +} + +// latestAssistantText returns the trimmed text of the most recent assistant +// message. It mirrors the FinalAssistantText that buildCommitStepMessages +// produced from the freshly generated step, making persisted history the +// single source of truth for the turn status label input. +func latestAssistantText(messages []database.ChatMessage) string { + for i := len(messages) - 1; i >= 0; i-- { + if messages[i].Role != database.ChatMessageRoleAssistant { + continue + } + parts, err := chatprompt.ParseContent(messages[i]) + if err != nil { + return "" + } + return strings.TrimSpace(textFromParts(parts)) + } + return "" +} + +// ACLs are deliberately not re-checked: revocation blocks new selection but +// leaves already-selected servers usable, like template ACLs for running +// workspaces. Disabling or deleting the config cuts off existing chats. +func enabledMCPServerConfigsForChatOrg( + ctx context.Context, + db database.Store, + organizationID uuid.UUID, + ids []uuid.UUID, +) ([]database.MCPServerConfig, error) { + configs, err := db.GetEnabledMCPServerConfigsByOrganizationAndIDs(ctx, database.GetEnabledMCPServerConfigsByOrganizationAndIDsParams{ + OrganizationID: organizationID, + IDs: ids, + }) + if err != nil { + return nil, xerrors.Errorf("get enabled MCP server configs for organization: %w", err) + } + return configs, nil +} + +type turnWorkspaceContext struct { + server *Server + chatStateMu *sync.Mutex + currentChat *database.Chat + loadChatSnapshot func(context.Context, uuid.UUID) (database.Chat, error) + + mu sync.Mutex + agent database.WorkspaceAgent + agentLoaded bool + conn workspacesdk.AgentConn + releaseConn func() + cachedWorkspaceID uuid.NullUUID +} + +func (c *turnWorkspaceContext) close() { + c.clearCachedWorkspaceState() +} + +func (c *turnWorkspaceContext) clearCachedWorkspaceState() { + c.mu.Lock() + releaseConn := c.releaseConn + c.agent = database.WorkspaceAgent{} + c.agentLoaded = false + c.conn = nil + c.releaseConn = nil + c.cachedWorkspaceID = uuid.NullUUID{} + c.mu.Unlock() + + if releaseConn != nil { + releaseConn() + } +} + +func (c *turnWorkspaceContext) setCurrentChat(chat database.Chat) { + c.chatStateMu.Lock() + *c.currentChat = chat + c.chatStateMu.Unlock() +} + +func (c *turnWorkspaceContext) currentChatSnapshot() database.Chat { + c.chatStateMu.Lock() + chatSnapshot := *c.currentChat + c.chatStateMu.Unlock() + return chatSnapshot +} + +func (c *turnWorkspaceContext) selectWorkspace(chat database.Chat) { + c.setCurrentChat(chat) + c.clearCachedWorkspaceState() +} + +func (c *turnWorkspaceContext) currentWorkspaceMatches(expected uuid.NullUUID) (database.Chat, bool) { + chatSnapshot := c.currentChatSnapshot() + return chatSnapshot, nullUUIDEqual(chatSnapshot.WorkspaceID, expected) +} + +func (c *turnWorkspaceContext) trackWorkspaceUsage(ctx context.Context, chatSnapshot database.Chat) { + if c.server == nil || !chatSnapshot.WorkspaceID.Valid { + return + } + logger := c.server.logger.With( + slog.F("chat_id", chatSnapshot.ID), + slog.F("owner_id", chatSnapshot.OwnerID), + ) + c.server.trackWorkspaceUsage(ctx, chatSnapshot.ID, chatSnapshot.WorkspaceID, logger) +} + +func nullUUIDEqual(left, right uuid.NullUUID) bool { + if left.Valid != right.Valid { + return false + } + if !left.Valid { + return true + } + return left.UUID == right.UUID +} + +func (c *turnWorkspaceContext) persistBuildAgentBinding( + ctx context.Context, + chatSnapshot database.Chat, + buildID uuid.UUID, + agentID uuid.UUID, +) (database.Chat, error) { + updatedChat, err := c.server.db.UpdateChatBuildAgentBinding( + ctx, + database.UpdateChatBuildAgentBindingParams{ + ID: chatSnapshot.ID, + BuildID: uuid.NullUUID{ + UUID: buildID, + Valid: true, + }, + AgentID: uuid.NullUUID{ + UUID: agentID, + Valid: true, + }, + }, + ) + if err != nil { + return chatSnapshot, xerrors.Errorf( + "update chat build/agent binding: %w", err, + ) + } + + // If the chat was rebound to a different agent (e.g. a workspace rebuild + // produced a new agent), re-pin its context to the new agent so it stops + // injecting the previous agent's resources. Workspace lifecycle tools clear + // the agent binding while preserving the pin, so a missing prior agent also + // requires a re-pin when pinned context exists. Best-effort: a context error + // must never fail the binding. The pinned context fields on updatedChat are + // background state, reloaded on the next snapshot fetch. + hasStaleUnboundContext := !chatSnapshot.AgentID.Valid && chatSnapshot.ContextAggregateHash != nil + if hasStaleUnboundContext || (chatSnapshot.AgentID.Valid && chatSnapshot.AgentID.UUID != agentID) { + //nolint:gocritic // Chatd re-pins chats it does not own as the daemon subject. + repinCtx := dbauthz.AsChatd(ctx) + if repinErr := database.ReadModifyUpdate(c.server.db, func(tx database.Store) error { + return repinChatContext(repinCtx, tx, chatSnapshot.ID, uuid.NullUUID{UUID: agentID, Valid: true}) + }); repinErr != nil { + c.server.logger.Warn(ctx, "re-pin chat context after agent rebind", + slog.F("chat_id", chatSnapshot.ID), + slog.F("agent_id", agentID), + slog.Error(repinErr)) + } + } + + c.setCurrentChat(updatedChat) + return updatedChat, nil +} + +func (c *turnWorkspaceContext) getWorkspaceAgent(ctx context.Context) (database.WorkspaceAgent, error) { + _, agent, err := c.ensureWorkspaceAgent(ctx) + return agent, err +} + +func (c *turnWorkspaceContext) ensureWorkspaceAgent( + ctx context.Context, +) (database.Chat, database.WorkspaceAgent, error) { + c.mu.Lock() + defer c.mu.Unlock() + + if c.agentLoaded { + chatSnapshot := c.currentChatSnapshot() + if nullUUIDEqual(c.cachedWorkspaceID, chatSnapshot.WorkspaceID) { + return chatSnapshot, c.agent, nil + } + c.agent = database.WorkspaceAgent{} + c.agentLoaded = false + } + + return c.loadWorkspaceAgentLocked(ctx) +} + +func (c *turnWorkspaceContext) loadWorkspaceAgentLocked( + ctx context.Context, +) (database.Chat, database.WorkspaceAgent, error) { + chatSnapshot := c.currentChatSnapshot() + + for attempt := 0; attempt < 2; attempt++ { + if !chatSnapshot.WorkspaceID.Valid { + refreshedChat, refreshErr := refreshChatWorkspaceSnapshot( + ctx, + chatSnapshot, + c.loadChatSnapshot, + ) + if refreshErr != nil { + return chatSnapshot, database.WorkspaceAgent{}, refreshErr + } + if refreshedChat.WorkspaceID.Valid { + c.setCurrentChat(refreshedChat) + chatSnapshot = refreshedChat + } + } + + if !chatSnapshot.WorkspaceID.Valid { + return chatSnapshot, database.WorkspaceAgent{}, xerrors.New("no workspace is associated with this chat. Use the create_workspace tool to create one") + } + + if chatSnapshot.AgentID.Valid { + agent, err := c.server.db.GetWorkspaceAgentByID(ctx, chatSnapshot.AgentID.UUID) + if err == nil { + latestChat, workspaceMatches := c.currentWorkspaceMatches(chatSnapshot.WorkspaceID) + if !workspaceMatches { + chatSnapshot = latestChat + continue + } + c.agent = agent + c.agentLoaded = true + c.cachedWorkspaceID = chatSnapshot.WorkspaceID + return chatSnapshot, c.agent, nil + } + if !xerrors.Is(err, sql.ErrNoRows) { + c.server.logger.Warn(ctx, "agent binding lookup failed, re-resolving", + slog.F("agent_id", chatSnapshot.AgentID.UUID), + slog.Error(err), + ) + } + } + + agents, err := c.server.db.GetWorkspaceAgentsInLatestBuildByWorkspaceID( + ctx, + chatSnapshot.WorkspaceID.UUID, + ) + if err != nil { + return chatSnapshot, database.WorkspaceAgent{}, xerrors.Errorf( + "get workspace agents in latest build: %w", + err, + ) + } + if len(agents) == 0 { + return chatSnapshot, database.WorkspaceAgent{}, errChatHasNoWorkspaceAgent + } + selected, err := agentselect.FindChatAgent(agents) + if err != nil { + return chatSnapshot, database.WorkspaceAgent{}, xerrors.Errorf( + "find chat agent: %w", + err, + ) + } + + build, err := c.server.db.GetLatestWorkspaceBuildByWorkspaceID(ctx, chatSnapshot.WorkspaceID.UUID) + if err != nil { + return chatSnapshot, database.WorkspaceAgent{}, xerrors.Errorf("get latest workspace build: %w", err) + } + + updatedChat, err := c.persistBuildAgentBinding( + ctx, + chatSnapshot, + build.ID, + selected.ID, + ) + if err != nil { + return chatSnapshot, database.WorkspaceAgent{}, err + } + + chatSnapshot = updatedChat + latestChat, workspaceMatches := c.currentWorkspaceMatches(chatSnapshot.WorkspaceID) + if !workspaceMatches { + chatSnapshot = latestChat + continue + } + c.agent = selected + c.agentLoaded = true + c.cachedWorkspaceID = chatSnapshot.WorkspaceID + return chatSnapshot, c.agent, nil + } + + return chatSnapshot, database.WorkspaceAgent{}, xerrors.New( + "chat workspace changed while resolving agent", + ) +} + +func (c *turnWorkspaceContext) latestWorkspaceAgentID( + ctx context.Context, + workspaceID uuid.UUID, +) (uuid.UUID, error) { + agents, err := c.server.db.GetWorkspaceAgentsInLatestBuildByWorkspaceID( + ctx, + workspaceID, + ) + if err != nil { + return uuid.Nil, xerrors.Errorf( + "get workspace agents in latest build: %w", + err, + ) + } + if len(agents) == 0 { + return uuid.Nil, errChatHasNoWorkspaceAgent + } + selected, err := agentselect.FindChatAgent(agents) + if err != nil { + return uuid.Nil, xerrors.Errorf( + "find chat agent: %w", + err, + ) + } + return selected.ID, nil +} + +func (c *turnWorkspaceContext) workspaceAgentIDForConn( + ctx context.Context, +) (database.Chat, uuid.UUID, error) { + for attempt := 0; attempt < 2; attempt++ { + chatSnapshot := c.currentChatSnapshot() + if !chatSnapshot.WorkspaceID.Valid || !chatSnapshot.AgentID.Valid { + updatedChat, agent, err := c.ensureWorkspaceAgent(ctx) + if err != nil { + return updatedChat, uuid.Nil, err + } + return updatedChat, agent.ID, nil + } + + currentAgentID, err := c.latestWorkspaceAgentID( + ctx, + chatSnapshot.WorkspaceID.UUID, + ) + if err != nil { + if xerrors.Is(err, errChatHasNoWorkspaceAgent) { + c.clearCachedWorkspaceState() + } + return chatSnapshot, uuid.Nil, err + } + + latestChat, workspaceMatches := c.currentWorkspaceMatches( + chatSnapshot.WorkspaceID, + ) + if !workspaceMatches { + continue + } + return latestChat, currentAgentID, nil + } + + chatSnapshot := c.currentChatSnapshot() + return chatSnapshot, uuid.Nil, xerrors.New( + "chat workspace changed while resolving agent", + ) +} + +// getWorkspaceConnLocked returns the cached connection when it still matches +// the current workspace. When the workspace changed, it clears the stale +// cached state and returns the release func for the caller to run after +// unlocking. +func (c *turnWorkspaceContext) getWorkspaceConnLocked() (workspacesdk.AgentConn, func()) { + if c.conn == nil { + return nil, nil + } + + chatSnapshot := c.currentChatSnapshot() + if nullUUIDEqual(c.cachedWorkspaceID, chatSnapshot.WorkspaceID) { + return c.conn, nil + } + + agentRelease := c.releaseConn + c.agent = database.WorkspaceAgent{} + c.agentLoaded = false + c.conn = nil + c.releaseConn = nil + c.cachedWorkspaceID = uuid.NullUUID{} + return nil, agentRelease +} + +// isAgentUnreachable reports whether the given agent row's +// status is disconnected or timed out. It uses timestamp +// arithmetic on the row. The "connecting" state is allowed +// through because it is normal after a fresh workspace build. +func isAgentUnreachable(now time.Time, agent database.WorkspaceAgent, inactiveTimeout time.Duration) bool { + status := agent.Status(now, inactiveTimeout) + return status.Status == database.WorkspaceAgentStatusDisconnected || + status.Status == database.WorkspaceAgentStatusTimeout +} + +func agentDisconnectedFor(now time.Time, agent database.WorkspaceAgent, inactiveTimeout time.Duration) (time.Duration, bool) { + status := agent.Status(now, inactiveTimeout) + if status.Status != database.WorkspaceAgentStatusDisconnected || status.DisconnectedAt == nil { + return 0, false + } + + disconnectedFor := now.Sub(*status.DisconnectedAt) + if disconnectedFor < 0 { + disconnectedFor = 0 + } + return disconnectedFor, true +} + +func (c *turnWorkspaceContext) latestWorkspaceAgentRecoveryError( + ctx context.Context, + workspaceID uuid.UUID, +) error { + agentID, err := c.latestWorkspaceAgentID(ctx, workspaceID) + if err != nil { + if xerrors.Is(err, errChatHasNoWorkspaceAgent) { + return err + } + c.server.logger.Warn(ctx, "failed to resolve latest agent for timeout classification", slog.Error(err)) + return errChatDialTimeout + } + + agent, err := c.server.db.GetWorkspaceAgentByID(ctx, agentID) + if err != nil { + c.server.logger.Warn(ctx, "failed to load latest agent for timeout classification", + slog.F("agent_id", agentID), + slog.Error(err), + ) + return errChatDialTimeout + } + + now := c.server.clock.Now() + status := agent.Status(now, c.server.agentInactiveDisconnectTimeout) + recoveryErr := errChatDialTimeout + if status.Status == database.WorkspaceAgentStatusTimeout { + recoveryErr = errChatAgentNeverConnected + } else if status.Status == database.WorkspaceAgentStatusDisconnected && status.DisconnectedAt != nil { + disconnectedFor := now.Sub(*status.DisconnectedAt) + if disconnectedFor < 0 { + disconnectedFor = 0 + } + if disconnectedFor >= agentDisconnectedRecoveryThreshold { + recoveryErr = errChatAgentDisconnected + } + } + return c.externalAgentError(ctx, agent, recoveryErr) +} + +func (c *turnWorkspaceContext) externalAgentError( + ctx context.Context, + agent database.WorkspaceAgent, + fallback error, +) error { + isExternal, err := chattool.IsExternalWorkspaceAgent(ctx, c.server.db, agent) + if err != nil || !isExternal { + return fallback + } + return newChatExternalAgentUnavailableError(agent) +} + +func (c *turnWorkspaceContext) externalAgentPreflightError( + ctx context.Context, + chatSnapshot database.Chat, + agent database.WorkspaceAgent, +) error { + // Mirror the cache-hit gate: only short-circuit on clearly offline + // states (Disconnected/Timeout). Connecting is allowed through so + // an external agent the user just started can still connect inside + // the normal dial window. + if !isAgentUnreachable(c.server.clock.Now(), agent, c.server.agentInactiveDisconnectTimeout) { + return nil + } + + isExternal, err := chattool.IsExternalWorkspaceAgent(ctx, c.server.db, agent) + if err != nil || !isExternal || !chatSnapshot.WorkspaceID.Valid { + return nil + } + + // Stale agent bindings rely on dialWithLazyValidation to discover + // replacement agents, so only skip the dial when this agent is still + // the latest selected chat agent for the workspace. + latestAgentID, err := c.latestWorkspaceAgentID(ctx, chatSnapshot.WorkspaceID.UUID) + if err != nil || latestAgentID != agent.ID { + return nil + } + return newChatExternalAgentUnavailableError(agent) +} + +func (c *turnWorkspaceContext) getWorkspaceConn(ctx context.Context) (workspacesdk.AgentConn, error) { + if c.server.agentConnFn == nil { + return nil, xerrors.New("workspace agent connector is not configured") + } + + for attempt := 0; attempt < 2; attempt++ { + c.mu.Lock() + currentConn, staleRelease := c.getWorkspaceConnLocked() + // Capture agentID in the same lock section as + // currentConn to prevent a TOCTOU race with + // concurrent clearCachedWorkspaceState calls. + agentID := c.agent.ID + c.mu.Unlock() + + // Status check on cache hit: re-fetch the agent + // row so we see the latest heartbeat rather than + // a potentially stale cached copy. + if currentConn != nil { + chatSnapshot := c.currentChatSnapshot() + if agentID != uuid.Nil { + freshAgent, err := c.server.db.GetWorkspaceAgentByID(ctx, agentID) + if err != nil { + c.server.logger.Warn(ctx, "failed to re-fetch agent for status check", + slog.F("agent_id", agentID), + slog.Error(err), + ) + // On DB error the check re-runs on the + // next tool call. + } else if _, disconnected := agentDisconnectedFor( + c.server.clock.Now(), + freshAgent, + c.server.agentInactiveDisconnectTimeout, + ); disconnected { + c.clearCachedWorkspaceState() + continue + } + } + c.trackWorkspaceUsage(ctx, chatSnapshot) + return currentConn, nil + } + if staleRelease != nil { + staleRelease() + } + + chatSnapshot, agent, err := c.ensureWorkspaceAgent(ctx) + if err != nil { + return nil, err + } + if err := c.externalAgentPreflightError(ctx, chatSnapshot, agent); err != nil { + return nil, err + } + + // Wrap the dial in a timeout to bound the time spent + // waiting for an unreachable agent. The timeout scopes + // only dialWithLazyValidation, not ensureWorkspaceAgent + // or the post-dial binding steps. + dialCtx, dialCancelCause := context.WithCancelCause(ctx) + dialTimer := c.server.clock.AfterFunc( + c.server.dialTimeout, + func() { dialCancelCause(errChatDialTimeout) }, + "chatd", + dialTimeoutTimerTag, + ) + dialCancel := func() { + dialTimer.Stop() + dialCancelCause(nil) + } + dialResult, err := dialWithLazyValidation( + dialCtx, + c.server.clock, + agent.ID, + chatSnapshot.WorkspaceID.UUID, + DialFunc(c.server.agentConnFn), + func(ctx context.Context, workspaceID uuid.UUID) (uuid.UUID, error) { + return c.latestWorkspaceAgentID(ctx, workspaceID) + }, + workspaceDialValidationDelay, + ) + dialCancel() + if err != nil { + if xerrors.Is(err, errChatHasNoWorkspaceAgent) { + c.clearCachedWorkspaceState() + return nil, err + } + // Surface the dial timeout sentinel only when the + // parent context is still alive. If the parent was + // canceled (e.g. ErrInterrupted), its error must + // propagate unchanged so the chatloop can detect it. + if ctx.Err() == nil && errors.Is(context.Cause(dialCtx), errChatDialTimeout) { + c.clearCachedWorkspaceState() + return nil, c.latestWorkspaceAgentRecoveryError(ctx, chatSnapshot.WorkspaceID.UUID) + } + return nil, err + } + agentConn := dialResult.Conn + agentRelease := dialResult.Release + if dialResult.WasSwitched { + build, err := c.server.db.GetLatestWorkspaceBuildByWorkspaceID(ctx, chatSnapshot.WorkspaceID.UUID) + if err != nil { + if agentRelease != nil { + agentRelease() + } + return nil, xerrors.Errorf("get latest workspace build: %w", err) + } + + switchedAgent, err := c.server.db.GetWorkspaceAgentByID(ctx, dialResult.AgentID) + if err != nil { + if agentRelease != nil { + agentRelease() + } + return nil, xerrors.Errorf("get workspace agent by id: %w", err) + } + + updatedChat, err := c.persistBuildAgentBinding( + ctx, + chatSnapshot, + build.ID, + switchedAgent.ID, + ) + if err != nil { + if agentRelease != nil { + agentRelease() + } + return nil, err + } + chatSnapshot = updatedChat + + c.mu.Lock() + c.agent = switchedAgent + c.agentLoaded = true + c.cachedWorkspaceID = chatSnapshot.WorkspaceID + c.mu.Unlock() + } + + if _, workspaceMatches := c.currentWorkspaceMatches(chatSnapshot.WorkspaceID); !workspaceMatches { + if agentRelease != nil { + agentRelease() + } + c.clearCachedWorkspaceState() + continue + } + + c.mu.Lock() + if c.conn == nil { + c.conn = agentConn + c.releaseConn = agentRelease + c.cachedWorkspaceID = chatSnapshot.WorkspaceID + + var ancestorIDs []string + if chatSnapshot.ParentChatID.Valid { + ancestorIDs = append(ancestorIDs, chatSnapshot.ParentChatID.UUID.String()) + } + ancestorJSON, marshalErr := json.Marshal(ancestorIDs) + if marshalErr != nil { + ancestorJSON = []byte("[]") + } + agentConn.SetExtraHeaders(http.Header{ + workspacesdk.CoderChatIDHeader: {chatSnapshot.ID.String()}, + workspacesdk.CoderAncestorChatIDsHeader: {string(ancestorJSON)}, + }) + + c.mu.Unlock() + c.server.logger.Debug(ctx, "set chat headers on agent conn", + slog.F("chat_id", chatSnapshot.ID), + slog.F("ancestor_chat_ids", ancestorIDs), + slog.F("workspace_id", chatSnapshot.WorkspaceID.UUID), + slog.F("agent_id", dialResult.AgentID), + ) + c.trackWorkspaceUsage(ctx, chatSnapshot) + return agentConn, nil + } + currentConn = c.conn + c.mu.Unlock() + + if agentRelease != nil { + agentRelease() + } + c.trackWorkspaceUsage(ctx, chatSnapshot) + return currentConn, nil + } + + return nil, xerrors.New("chat workspace changed while connecting") +} + +// AgentConnFunc provides access to workspace agent connections. + +func allToolNames(allTools []fantasy.AgentTool) []string { + toolNames := make([]string, 0, len(allTools)) + for _, tool := range allTools { + toolNames = append(toolNames, tool.Info().Name) + } + return toolNames +} + +func isExploreSubagentMode(mode database.NullChatMode) bool { + return mode.Valid && mode.ChatMode == database.ChatModeExplore +} + +// filterExternalMCPConfigsForTurn returns the external MCP server configs +// visible on the current turn. Explore children snapshot this filtered set at +// spawn time so later model overrides cannot widen the external-tool boundary. +func filterExternalMCPConfigsForTurn( + configs []database.MCPServerConfig, + mode database.NullChatPlanMode, + parentChatID uuid.NullUUID, +) ([]database.MCPServerConfig, map[uuid.UUID]struct{}) { + if !mode.Valid || mode.ChatPlanMode != database.ChatPlanModePlan { + return configs, nil + } + if parentChatID.Valid { + // Plan-mode subagents do not receive external MCP tools because + // their trust boundary is narrower than the root chat's. + return nil, map[uuid.UUID]struct{}{} + } + + filtered := make([]database.MCPServerConfig, 0, len(configs)) + approvedIDs := make(map[uuid.UUID]struct{}) + for _, cfg := range configs { + if !cfg.AllowInPlanMode { + continue + } + filtered = append(filtered, cfg) + approvedIDs[cfg.ID] = struct{}{} + } + return filtered, approvedIDs +} + +func builtinPlanToolAllowed(name string, isRootChat bool) bool { + switch name { + case "read_file", "execute", "process_output", "read_skill", "read_skill_file": + return true + case "write_file", "edit_files", "list_templates", "read_template", + "create_workspace", "start_workspace", "stop_workspace", "propose_plan", "spawn_agent", + "spawn_explore_agent", "wait_agent", "list_agents", "list_subagent_models", + "ask_user_question", "attach_file": + return isRootChat + case "process_list", "process_signal", "message_agent", "interrupt_agent", "close_agent", + "spawn_computer_use_agent": + return false + default: + return false + } +} + +func toolAllowedForTurn( + tool fantasy.AgentTool, + mode database.NullChatPlanMode, + parentChatID uuid.NullUUID, + approvedMCPConfigIDs map[uuid.UUID]struct{}, +) bool { + if !mode.Valid || mode.ChatPlanMode != database.ChatPlanModePlan { + return true + } + if builtinPlanToolAllowed(tool.Info().Name, !parentChatID.Valid) { + return true + } + mcpTool, ok := tool.(mcpclient.MCPToolIdentifier) + if !ok { + return false + } + _, approved := approvedMCPConfigIDs[mcpTool.MCPServerConfigID()] + return approved +} + +func filterToolsForTurn( + allTools []fantasy.AgentTool, + mode database.NullChatPlanMode, + parentChatID uuid.NullUUID, + approvedMCPConfigIDs map[uuid.UUID]struct{}, +) []fantasy.AgentTool { + if !mode.Valid || mode.ChatPlanMode != database.ChatPlanModePlan { + return allTools + } + + filtered := make([]fantasy.AgentTool, 0, len(allTools)) + for _, tool := range allTools { + if toolAllowedForTurn(tool, mode, parentChatID, approvedMCPConfigIDs) { + filtered = append(filtered, tool) + } + } + return filtered +} + +// activeToolNamesForTurn extends the built-in plan allowlist with approved +// external MCP tools for root plan-mode chats. +func activeToolNamesForTurn( + allTools []fantasy.AgentTool, + mode database.NullChatPlanMode, + parentChatID uuid.NullUUID, + approvedMCPConfigIDs map[uuid.UUID]struct{}, +) []string { + toolNames := make([]string, 0, len(allTools)) + for _, tool := range allTools { + if toolAllowedForTurn(tool, mode, parentChatID, approvedMCPConfigIDs) { + toolNames = append(toolNames, tool.Info().Name) + } + } + return toolNames +} + +func allowedExploreToolNames(allTools []fantasy.AgentTool) []string { + builtinExplorePolicy := map[string]bool{ + "read_file": true, + "write_file": false, + "edit_files": false, + "execute": true, + "process_output": true, + "process_list": false, + "process_signal": false, + "list_templates": false, + "read_template": false, + "create_workspace": false, + "start_workspace": false, + "stop_workspace": false, + "propose_plan": false, + "spawn_agent": false, + "wait_agent": false, + "message_agent": false, + "interrupt_agent": false, + "close_agent": false, + "list_agents": false, + "list_subagent_models": false, + "read_skill": true, + "read_skill_file": true, + "ask_user_question": false, + } + + toolNames := make([]string, 0, len(allTools)) + for _, tool := range allTools { + name := tool.Info().Name + if builtinExplorePolicy[name] { + toolNames = append(toolNames, name) + continue + } + // External MCP tools pass through here. They were snapshot-filtered + // at spawn time on chat.MCPServerIDs. WorkspaceMCPTool does not + // implement MCPToolIdentifier, so workspace tools are excluded + // here too, in addition to the structural exclusion in runChat + // tool assembly. + if _, ok := tool.(mcpclient.MCPToolIdentifier); ok { + toolNames = append(toolNames, name) + } + } + return toolNames +} + +// allowedBehaviorToolNames runs only on non-plan turns because +// appendDynamicTools returns early for plan mode. Within that boundary, +// Explore mode wins over the default behavior that allows all tools. +func allowedBehaviorToolNames( + allTools []fantasy.AgentTool, + chatMode database.NullChatMode, +) []string { + if isExploreSubagentMode(chatMode) { + return allowedExploreToolNames(allTools) + } + return allToolNames(allTools) +} + +func stopAfterPlanTools( + planMode database.NullChatPlanMode, + parentChatID uuid.NullUUID, +) map[string]struct{} { + if !planMode.Valid || planMode.ChatPlanMode != database.ChatPlanModePlan { + return nil + } + stopTools := map[string]struct{}{ + "propose_plan": {}, + } + if !parentChatID.Valid { + stopTools["ask_user_question"] = struct{}{} + } + return stopTools +} + +func stopAfterBehaviorTools( + planMode database.NullChatPlanMode, + chatMode database.NullChatMode, + parentChatID uuid.NullUUID, +) map[string]struct{} { + if isExploreSubagentMode(chatMode) { + return nil + } + return stopAfterPlanTools(planMode, parentChatID) +} + +type systemPromptBehaviorContext struct { + planMode database.NullChatPlanMode + chatMode database.NullChatMode + planModeInstructions string + isRootChat bool +} + +func workspaceSkillsForResolution(workspaceSkills []chattool.SkillMeta) []skillspkg.Skill { + if len(workspaceSkills) == 0 { + return nil + } + resolved := make([]skillspkg.Skill, 0, len(workspaceSkills)) + for _, skill := range workspaceSkills { + resolved = append(resolved, skillspkg.Skill{ + Name: skill.Name, + Description: skill.Description, + Source: skillspkg.SourceWorkspace, + }) + } + return resolved +} + +func mergeTurnSkills( + personalSkills []skillspkg.Skill, + workspaceSkills []chattool.SkillMeta, +) []skillspkg.ResolvedSkill { + return skillspkg.MergeSkills( + personalSkills, + workspaceSkillsForResolution(workspaceSkills), + ) +} + +// buildSystemPrompt applies system-level prompt injections in a fixed +// order: subagent instruction, chat instruction, skill index, user prompt, +// then mode overlay prompts. +func buildSystemPrompt( + prompt []fantasy.Message, + subagentInstruction string, + instruction string, + resolvedSkills []skillspkg.ResolvedSkill, + userPrompt string, + behaviorContext systemPromptBehaviorContext, +) []fantasy.Message { + if subagentInstruction != "" { + prompt = chatprompt.InsertSystem(prompt, subagentInstruction) + } + if instruction != "" { + prompt = chatprompt.InsertSystem(prompt, instruction) + } + if skillIndex := chattool.FormatResolvedSkillIndex(resolvedSkills); skillIndex != "" { + prompt = chatprompt.InsertSystem(prompt, skillIndex) + } + if userPrompt != "" { + prompt = chatprompt.InsertSystem(prompt, userPrompt) + } + if isExploreSubagentMode(behaviorContext.chatMode) { + prompt = chatprompt.InsertSystem(prompt, ExploreSubagentOverlayPrompt) + return prompt + } + isPlanModeTurn := behaviorContext.planMode.Valid && behaviorContext.planMode.ChatPlanMode == database.ChatPlanModePlan + if isPlanModeTurn { + if behaviorContext.isRootChat { + prompt = chatprompt.InsertSystem(prompt, PlanningOverlayPrompt()) + if behaviorContext.planModeInstructions != "" { + prompt = chatprompt.InsertSystem(prompt, behaviorContext.planModeInstructions) + } + } else { + prompt = chatprompt.InsertSystem(prompt, PlanningSubagentOverlayPrompt) + } + } + return prompt +} + +func removeSkillIndexMessages(prompt []fantasy.Message) []fantasy.Message { + out := make([]fantasy.Message, 0, len(prompt)) + removed := false + for _, message := range prompt { + if isSkillIndexMessage(message) { + removed = true + continue + } + out = append(out, message) + } + if !removed { + return prompt + } + return out +} + +func isSkillIndexMessage(message fantasy.Message) bool { + if message.Role != fantasy.MessageRoleSystem || len(message.Content) != 1 { + return false + } + textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](message.Content[0]) + if !ok { + return false + } + text := strings.TrimSpace(textPart.Text) + return strings.HasPrefix(text, chattool.AvailableSkillsOpenTag+"\n") && strings.HasSuffix(text, chattool.AvailableSkillsCloseTag) +} + +type rootChatToolsOptions struct { + chat database.Chat + modelConfigID uuid.UUID + workspaceCtx *turnWorkspaceContext + workspaceMu *sync.Mutex + resolvePlanPath func(context.Context) (string, string, error) + storeFile chattool.StoreFileFunc + isPlanModeTurn bool +} + +func (server *Server) loadPlanModeInstructions( + ctx context.Context, + mode database.NullChatPlanMode, + logger slog.Logger, +) string { + if !mode.Valid || mode.ChatPlanMode != database.ChatPlanModePlan { + return "" + } + + // Plan-mode instructions live in deployment config, but chat workers do + // not carry a deployment-config actor during background execution. + //nolint:gocritic // Required to read deployment config during background chat processing. + systemCtx := dbauthz.AsSystemRestricted(ctx) + fetched, err := server.db.GetChatPlanModeInstructions(systemCtx) + if err != nil { + logger.Warn(ctx, + "failed to fetch plan mode instructions", + slog.Error(err), + ) + return "" + } + + return fetched +} + +func userSkillContext(ctx context.Context, userID uuid.UUID) context.Context { + actor := rbac.Subject{ + Type: rbac.SubjectTypeUser, + ID: userID.String(), + Roles: rbac.RoleIdentifiers{rbac.RoleMember()}, + Scope: rbac.ScopeAll, + }.WithCachedASTValue() + // Chat turns run asynchronously after admission, so the original request + // actor may no longer be available when a worker loads personal skills. + // We synthesize the chat owner as a member instead of reusing that actor. + // Hardcoding RoleMember is safe because dbauthz enforces + // ResourceUserSkill.WithOwner(userID), so this actor cannot read any other + // user's skills regardless of role. Org scoping is not needed because + // personal skills are user-scoped, not org-scoped. + //nolint:gocritic // The synthetic actor is intentional for the reasons above. + return dbauthz.As(ctx, actor) +} + +func (server *Server) fetchPersonalSkillMetadata( + ctx context.Context, + userID uuid.UUID, + logger slog.Logger, +) []skillspkg.Skill { + rows, err := server.db.ListUserSkillMetadataByUserID(userSkillContext(ctx, userID), userID) + // See package coderd/x/skills (doc.go) for why metadata fetch failures + // intentionally degrade to an empty personal-skill list instead of + // failing the chat turn. + if err != nil { + logger.Warn(ctx, "failed to load personal skill metadata", + slog.F("owner_id", userID), + slog.Error(err), + ) + return nil + } + + personalSkills := make([]skillspkg.Skill, 0, len(rows)) + for _, row := range rows { + personalSkills = append(personalSkills, skillspkg.Skill{ + Name: row.Name, + Description: row.Description, + Source: skillspkg.SourcePersonal, + }) + } + return personalSkills +} + +func (server *Server) loadPersonalSkillBody( + ctx context.Context, + userID uuid.UUID, + name string, +) (skillspkg.ParsedSkill, error) { + row, err := server.db.GetUserSkillByUserIDAndName( + userSkillContext(ctx, userID), + database.GetUserSkillByUserIDAndNameParams{ + UserID: userID, + Name: name, + }, + ) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return skillspkg.ParsedSkill{}, skillspkg.ErrSkillNotFound + } + server.logger.Error(ctx, "load personal skill body failed", + slog.F("user_id", userID), + slog.F("name", name), + slog.Error(err), + ) + return skillspkg.ParsedSkill{}, xerrors.Errorf("load personal skill body: %w", err) + } + + parsed, err := skillspkg.ParsePersonalSkillMarkdown([]byte(row.Content)) + if err != nil { + server.logger.Error(ctx, "parse personal skill body failed", + slog.F("user_id", userID), + slog.F("name", name), + slog.Error(err), + ) + return skillspkg.ParsedSkill{}, xerrors.Errorf("parse personal skill body: %w", err) + } + return parsed, nil +} + +func (server *Server) appendRootChatTools( + ctx context.Context, + tools []fantasy.AgentTool, + opts rootChatToolsOptions, +) []fantasy.AgentTool { + onChatUpdated := func(updatedChat database.Chat) { + opts.workspaceCtx.selectWorkspace(updatedChat) + // Notify the frontend immediately so it can start streaming + // build logs before the tool completes. + server.publishChatPubsubEvent(updatedChat, codersdk.ChatWatchEventKindStatusChange, nil) + } + + tools = append(tools, + chattool.ListTemplates(server.db, opts.chat.OrganizationID, chattool.ListTemplatesOptions{ + OwnerID: opts.chat.OwnerID, + Logger: server.logger, + Clock: server.clock, + }), + chattool.ReadTemplate(server.db, opts.chat.OrganizationID, chattool.ReadTemplateOptions{ + OwnerID: opts.chat.OwnerID, + }), + chattool.CreateWorkspace(server.db, opts.chat.OrganizationID, opts.chat.ID, chattool.CreateWorkspaceOptions{ + OwnerID: opts.chat.OwnerID, + CreateFn: server.createWorkspaceFn, + AgentConnFn: chattool.AgentConnFunc(server.agentConnFn), + AgentInactiveDisconnectTimeout: server.agentInactiveDisconnectTimeout, + WorkspaceMu: opts.workspaceMu, + OnChatUpdated: onChatUpdated, + Logger: server.logger, + }), + chattool.StartWorkspace(server.db, opts.chat.ID, chattool.StartWorkspaceOptions{ + OwnerID: opts.chat.OwnerID, + StartFn: server.startWorkspaceFn, + AgentConnFn: chattool.AgentConnFunc(server.agentConnFn), + WorkspaceMu: opts.workspaceMu, + OnChatUpdated: onChatUpdated, + Logger: server.logger, + }), + chattool.StopWorkspace(server.db, opts.chat.ID, chattool.StopWorkspaceOptions{ + OwnerID: opts.chat.OwnerID, + StopFn: server.stopWorkspaceFn, + WorkspaceMu: opts.workspaceMu, + OnChatUpdated: onChatUpdated, + Logger: server.logger, + }), + ) + if opts.isPlanModeTurn { + tools = append(tools, chattool.ProposePlan(chattool.ProposePlanOptions{ + GetWorkspaceConn: opts.workspaceCtx.getWorkspaceConn, + ResolvePlanPath: opts.resolvePlanPath, + IsPlanTurn: opts.isPlanModeTurn, + StoreFile: opts.storeFile, + })) + } + + return append(tools, server.subagentTools(ctx, func() database.Chat { + return opts.chat + }, opts.modelConfigID)...) +} + +func appendDynamicTools( + ctx context.Context, + logger slog.Logger, + tools []fantasy.AgentTool, + raw pqtype.NullRawMessage, + planMode database.NullChatPlanMode, + chatMode database.NullChatMode, +) ([]fantasy.AgentTool, map[string]bool, error) { + if isExploreSubagentMode(chatMode) || (planMode.Valid && planMode.ChatPlanMode == database.ChatPlanModePlan) { + return tools, nil, nil + } + + dynamicToolNames, err := parseDynamicToolNames(raw) + if err != nil { + return nil, nil, xerrors.Errorf("parse dynamic tool names: %w", err) + } + if len(dynamicToolNames) == 0 { + return tools, dynamicToolNames, nil + } + + var dynamicToolDefs []codersdk.DynamicTool + if raw.Valid { + if err := json.Unmarshal(raw.RawMessage, &dynamicToolDefs); err != nil { + return nil, nil, xerrors.Errorf("unmarshal dynamic tools: %w", err) + } + } + + activeToolNames := make(map[string]struct{}, len(tools)) + for _, name := range allowedBehaviorToolNames(tools, chatMode) { + activeToolNames[name] = struct{}{} + } + for _, t := range tools { + info := t.Info() + if _, active := activeToolNames[info.Name]; !active { + continue + } + if dynamicToolNames[info.Name] { + logger.Warn(ctx, "dynamic tool name collides with built-in tool, built-in takes precedence", + slog.F("tool_name", info.Name)) + delete(dynamicToolNames, info.Name) + } + } + + var filteredDefs []codersdk.DynamicTool + for _, dt := range dynamicToolDefs { + if dynamicToolNames[dt.Name] { + filteredDefs = append(filteredDefs, dt) + } + } + + return append(tools, dynamicToolsFromSDK(logger, filteredDefs)...), dynamicToolNames, nil +} + +// buildProviderTools creates provider-native tool definitions +// (like web search) based on the model configuration. These +// tools are executed server-side by the LLM provider. +func buildProviderTools(options *codersdk.ChatModelProviderOptions) []chatloop.ProviderTool { + var tools []chatloop.ProviderTool + + if options == nil { + return nil + } + + if options.Anthropic != nil && options.Anthropic.WebSearchEnabled != nil && *options.Anthropic.WebSearchEnabled { + tools = append(tools, chatloop.ProviderTool{ + Definition: anthropic.WebSearchTool(&anthropic.WebSearchToolOptions{ + AllowedDomains: options.Anthropic.AllowedDomains, + BlockedDomains: options.Anthropic.BlockedDomains, + }), + }) + } + + if tool, ok := chatopenai.WebSearchTool(options.OpenAI); ok { + tools = append(tools, chatloop.ProviderTool{ + Definition: tool, + }) + } + + if options.Google != nil && options.Google.WebSearchEnabled != nil && *options.Google.WebSearchEnabled { + tools = append(tools, chatloop.ProviderTool{ + Definition: fantasy.ProviderDefinedTool{ + ID: "web_search", + Name: "web_search", + }, + }) + } + + return tools +} From c1d022cbbe0fc6c22101d2b94d72103e873a0118 Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Sat, 29 Aug 2026 09:49:31 +0000 Subject: [PATCH 2/4] test(coderd/x/chatd): replace turn policy fixtures --- coderd/x/chatd/ARCHITECTURE.md | 2 + coderd/x/chatd/chatd.go | 1 + coderd/x/chatd/chatd_internal_test.go | 12 +- coderd/x/chatd/chatd_test.go | 953 ------------------ coderd/x/chatd/generation.go | 14 +- coderd/x/chatd/toolinput.go | 14 +- coderd/x/chatd/toolinput_internal_test.go | 32 +- coderd/x/chatd/turn_environment.go | 210 ++-- ...t.go => turn_environment_internal_test.go} | 43 +- 9 files changed, 171 insertions(+), 1110 deletions(-) rename coderd/x/chatd/{generation_preparer_internal_test.go => turn_environment_internal_test.go} (95%) diff --git a/coderd/x/chatd/ARCHITECTURE.md b/coderd/x/chatd/ARCHITECTURE.md index 42b7caf7c93..22b3979a587 100644 --- a/coderd/x/chatd/ARCHITECTURE.md +++ b/coderd/x/chatd/ARCHITECTURE.md @@ -855,6 +855,8 @@ Retriable conditions include, but are not limited to: #### Generation goroutine +TODO(human): Document `turn_environment.go` as the preparation boundary for prompts, models, workspace connections, and tool policy. + The generation goroutine is responsible for calling the LLM API and executing tools. It is spawned when the event indicates the core state machine is in `R0` or `R1` (status is `running`). It inspects the chat's message history, and decides what's the next step to take. The result of that step is the application of one of the following core state machine transitions: diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index b1bc24137cb..8a35b323ab8 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -455,6 +455,7 @@ func (p *Server) pinnedWorkspaceMCPTools( return chattool.NewWorkspaceMCPTools(infos, getConn, nil), nil } +// AgentConnFunc provides access to workspace agent connections. type AgentConnFunc func(ctx context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) var ( diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index de66681f1de..3848bc48211 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -1899,7 +1899,7 @@ func TestFetchPersonalSkillMetadata(t *testing.T) { }, ) - got := server.fetchPersonalSkillMetadata(context.Background(), userID, logger) + got := (turnEnvironmentBuilder{server: server}).fetchPersonalSkillMetadata(context.Background(), userID, logger) require.Equal(t, []skillspkg.Skill{{ Name: "personal-review", Description: "Personal review process", @@ -1919,7 +1919,7 @@ func TestFetchPersonalSkillMetadata(t *testing.T) { db.EXPECT().ListUserSkillMetadataByUserID(gomock.Any(), userID).Return(nil, xerrors.New("boom")) - got := server.fetchPersonalSkillMetadata(context.Background(), userID, logger) + got := (turnEnvironmentBuilder{server: server}).fetchPersonalSkillMetadata(context.Background(), userID, logger) require.Empty(t, got) warns := sink.Entries(func(e slog.SinkEntry) bool { return e.Level == slog.LevelWarn && strings.Contains(e.Message, "personal skill metadata") @@ -1955,7 +1955,7 @@ func TestLoadPersonalSkillBody(t *testing.T) { }, ) - got, err := server.loadPersonalSkillBody(context.Background(), userID, "personal-review") + got, err := (turnEnvironmentBuilder{server: server}).loadPersonalSkillBody(context.Background(), userID, "personal-review") require.NoError(t, err) require.Equal(t, "personal-review", got.Name) require.Equal(t, "Personal review process", got.Description) @@ -1983,7 +1983,7 @@ func TestLoadPersonalSkillBody(t *testing.T) { }, ) - _, err := server.loadPersonalSkillBody(context.Background(), userID, "missing-skill") + _, err := (turnEnvironmentBuilder{server: server}).loadPersonalSkillBody(context.Background(), userID, "missing-skill") require.ErrorIs(t, err, skillspkg.ErrSkillNotFound) }) @@ -2009,7 +2009,7 @@ func TestLoadPersonalSkillBody(t *testing.T) { }, ) - _, err := server.loadPersonalSkillBody(context.Background(), userID, "error-skill") + _, err := (turnEnvironmentBuilder{server: server}).loadPersonalSkillBody(context.Background(), userID, "error-skill") require.ErrorContains(t, err, "load personal skill body") require.ErrorIs(t, err, dbErr) @@ -2045,7 +2045,7 @@ func TestLoadPersonalSkillBody(t *testing.T) { }, ) - _, err := server.loadPersonalSkillBody(context.Background(), userID, "broken-skill") + _, err := (turnEnvironmentBuilder{server: server}).loadPersonalSkillBody(context.Background(), userID, "broken-skill") require.ErrorContains(t, err, "parse personal skill body") require.ErrorIs(t, err, skillspkg.ErrSkillBodyRequired) diff --git a/coderd/x/chatd/chatd_test.go b/coderd/x/chatd/chatd_test.go index c2f022081d2..22e49572ba7 100644 --- a/coderd/x/chatd/chatd_test.go +++ b/coderd/x/chatd/chatd_test.go @@ -373,959 +373,6 @@ func TestSubagentChatExcludesWorkspaceProvisioningTools(t *testing.T) { "subagent chat should NOT have ask_user_question") } -func TestPlanModeSubagentChatExcludesAskUserQuestion(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitLong) - client, _, api := coderdtest.NewWithAPI(t, &coderdtest.Options{ - DeploymentValues: coderdtest.DeploymentValues(t), - IncludeProvisionerDaemon: true, - }) - aibridgedtest.StartTestAIBridgeDaemon(t.Context(), t, api, nil) - user := coderdtest.CreateFirstUser(t, client) - expClient := codersdk.NewExperimentalClient(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) - - // Start an external MCP server whose tools should remain available to the - // root plan-mode chat but stay hidden from plan-mode subagents. - mcpSrv := newTestMCPServer("plan-root-mcp") - addTestMCPTextTool(mcpSrv, "echo", "Echoes the input", "echo: ") - mcpTS := httptest.NewServer(testMCPHTTPHandler(mcpSrv)) - t.Cleanup(mcpTS.Close) - - mcpConfig, err := client.CreateMCPServerConfig(ctx, user.OrganizationID, codersdk.CreateMCPServerConfigRequest{ - DisplayName: "Plan Root MCP", - Slug: "plan-root-mcp", - Transport: "streamable_http", - URL: mcpTS.URL, - AuthType: "none", - Availability: "default_off", - Enabled: true, - AllowInPlanMode: true, - }) - require.NoError(t, err) - - var toolsMu sync.Mutex - toolsByCall := make([][]string, 0, 2) - requestsByCall := make([]recordedOpenAIRequest, 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) - requestsByCall = append(requestsByCall, recordOpenAIRequest(req)) - toolsMu.Unlock() - - if callCount.Add(1) == 1 { - return chattest.OpenAIStreamingResponse( - chattest.OpenAIToolCallChunk("spawn_agent", `{"type":"general","prompt":"inspect the codebase","title":"sub"}`), - ) - } - return chattest.OpenAIStreamingResponse( - chattest.OpenAITextChunks("done")..., - ) - }) - - coderdtest.CreateOpenAICompatChatModel(t, expClient, openAIURL) - - chat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{ - OrganizationID: user.OrganizationID, - PlanMode: codersdk.ChatPlanModePlan, - MCPServerIDs: []uuid.UUID{mcpConfig.ID}, - Content: []codersdk.ChatInputPart{ - { - Type: codersdk.ChatInputPartTypeText, - Text: "Spawn a subagent to inspect the codebase.", - }, - }, - }) - require.NoError(t, err) - - require.Eventually(t, func() bool { - got, getErr := expClient.GetChat(ctx, chat.ID) - if getErr != nil { - return false - } - if got.Status != codersdk.ChatStatusWaiting && got.Status != codersdk.ChatStatusError { - return false - } - toolsMu.Lock() - n := len(toolsByCall) - toolsMu.Unlock() - return n >= 3 - }, testutil.WaitLong, testutil.IntervalFast) - - toolsMu.Lock() - recorded := append([][]string(nil), toolsByCall...) - recordedRequests := append([]recordedOpenAIRequest(nil), requestsByCall...) - toolsMu.Unlock() - - require.GreaterOrEqual(t, len(recorded), 2, - "expected at least 2 streamed LLM calls (root + subagent)") - require.Len(t, recordedRequests, len(recorded)) - - var rootCalls, childCalls [][]string - var rootRequests, childRequests []recordedOpenAIRequest - for i, tools := range recorded { - if slice.Contains(tools, "spawn_agent") { - rootCalls = append(rootCalls, tools) - rootRequests = append(rootRequests, recordedRequests[i]) - continue - } - childCalls = append(childCalls, tools) - childRequests = append(childRequests, recordedRequests[i]) - } - - require.NotEmpty(t, rootCalls, "expected at least one root chat LLM call") - require.NotEmpty(t, childCalls, "expected at least one subagent LLM call") - require.NotEmpty(t, rootRequests, "expected at least one root prompt") - require.NotEmpty(t, childRequests, "expected at least one subagent prompt") - require.Contains(t, rootCalls[0], "ask_user_question", - "root plan-mode chat should have ask_user_question") - require.Contains(t, rootCalls[0], "write_file", - "root plan-mode chat should have write_file") - require.Contains(t, rootCalls[0], "edit_files", - "root plan-mode chat should have edit_files") - require.Contains(t, rootCalls[0], "execute", - "root plan-mode chat should have execute") - require.Contains(t, rootCalls[0], "process_output", - "root plan-mode chat should have process_output") - require.Contains(t, rootCalls[0], "plan-root-mcp__echo", - "root plan-mode chat should have approved external MCP tools") - require.NotContains(t, childCalls[0], "ask_user_question", - "plan-mode subagent should NOT have ask_user_question") - require.NotContains(t, childCalls[0], "write_file", - "plan-mode subagent should NOT have write_file") - require.NotContains(t, childCalls[0], "edit_files", - "plan-mode subagent should NOT have edit_files") - require.Contains(t, childCalls[0], "execute", - "plan-mode subagent should have execute") - require.Contains(t, childCalls[0], "process_output", - "plan-mode subagent should have process_output") - require.NotContains(t, childCalls[0], "plan-root-mcp__echo", - "plan-mode subagent should NOT have external MCP tools") - require.True(t, requestHasSystemSubstring(rootRequests[0], "You are in Plan Mode.")) - require.True(t, requestHasSystemSubstring(childRequests[0], "You are in Plan Mode as a delegated sub-agent.")) - require.False(t, requestHasSystemSubstring(childRequests[0], "When the plan is ready, call propose_plan")) -} - -func TestExploreSubagentIsReadOnly(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitLong) - client, _, api := coderdtest.NewWithAPI(t, &coderdtest.Options{ - DeploymentValues: coderdtest.DeploymentValues(t), - IncludeProvisionerDaemon: true, - }) - db := api.Database - aibridgedtest.StartTestAIBridgeDaemon(t.Context(), t, api, nil) - user := coderdtest.CreateFirstUser(t, client) - expClient := codersdk.NewExperimentalClient(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) - template := coderdtest.CreateTemplate(t, client, user.OrganizationID, version.ID) - workspace := coderdtest.CreateWorkspace(t, client, template.ID, func(cwr *codersdk.CreateWorkspaceRequest) { - cwr.AutomaticUpdates = codersdk.AutomaticUpdatesNever - }) - coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID) - _ = agenttest.New(t, client.URL, agentToken) - coderdtest.NewWorkspaceAgentWaiter(t, client, workspace.ID).Wait() - - var toolsMu sync.Mutex - toolsByCall := make([][]string, 0, 2) - requestsByCall := make([]recordedOpenAIRequest, 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) - requestsByCall = append(requestsByCall, recordOpenAIRequest(req)) - toolsMu.Unlock() - - if callCount.Add(1) == 1 { - return chattest.OpenAIStreamingResponse( - chattest.OpenAIToolCallChunk("spawn_agent", `{"type":"explore","prompt":"investigate the codebase","title":"sub"}`), - ) - } - return chattest.OpenAIStreamingResponse( - chattest.OpenAITextChunks("done")..., - ) - }) - - coderdtest.CreateOpenAICompatChatModel(t, expClient, openAIURL) - - _, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{ - OrganizationID: user.OrganizationID, - WorkspaceID: &workspace.ID, - Content: []codersdk.ChatInputPart{ - { - Type: codersdk.ChatInputPartTypeText, - Text: "Spawn an Explore subagent to inspect the codebase.", - }, - }, - }) - require.NoError(t, err) - - require.Eventually(t, func() bool { - toolsMu.Lock() - defer toolsMu.Unlock() - - sawRoot := false - sawChild := false - for _, tools := range toolsByCall { - if slice.Contains(tools, "spawn_agent") { - sawRoot = true - continue - } - sawChild = true - } - return sawRoot && sawChild - }, testutil.WaitLong, testutil.IntervalFast) - - toolsMu.Lock() - recorded := append([][]string(nil), toolsByCall...) - recordedRequests := append([]recordedOpenAIRequest(nil), requestsByCall...) - toolsMu.Unlock() - - require.GreaterOrEqual(t, len(recorded), 2, - "expected at least 2 streamed LLM calls (root + subagent)") - require.Len(t, recordedRequests, len(recorded)) - - var rootCalls, childCalls [][]string - var rootRequests, childRequests []recordedOpenAIRequest - for i, tools := range recorded { - if slice.Contains(tools, "spawn_agent") { - rootCalls = append(rootCalls, tools) - rootRequests = append(rootRequests, recordedRequests[i]) - continue - } - childCalls = append(childCalls, tools) - childRequests = append(childRequests, recordedRequests[i]) - } - - require.NotEmpty(t, rootCalls, "expected at least one root chat LLM call") - require.NotEmpty(t, childCalls, "expected at least one subagent LLM call") - require.NotEmpty(t, rootRequests, "expected at least one root prompt") - require.NotEmpty(t, childRequests, "expected at least one subagent prompt") - require.Contains(t, rootCalls[0], "spawn_agent") - require.Contains(t, rootCalls[0], "write_file") - require.Contains(t, rootCalls[0], "edit_files") - require.NotContains(t, childCalls[0], "write_file") - require.NotContains(t, childCalls[0], "edit_files") - require.NotContains(t, childCalls[0], "spawn_agent") - require.NotContains(t, childCalls[0], "wait_agent") - require.Contains(t, childCalls[0], "read_file") - require.Contains(t, childCalls[0], "execute") - require.Contains(t, childCalls[0], "process_output") - require.True(t, requestHasSystemSubstring(childRequests[0], "You are in Explore Mode as a delegated sub-agent.")) - require.False(t, requestHasSystemSubstring(rootRequests[0], "You are in Explore Mode as a delegated sub-agent.")) - - rootChats, err := db.GetChats(dbauthz.AsChatd(ctx), database.GetChatsParams{ - OwnedOnly: true, - ViewerID: user.UserID, - }) - require.NoError(t, err) - rootIDs := make([]uuid.UUID, 0, len(rootChats)) - for _, root := range rootChats { - rootIDs = append(rootIDs, root.Chat.ID) - } - childRows, err := db.GetChildChatsByParentIDs(dbauthz.AsChatd(ctx), database.GetChildChatsByParentIDsParams{ - ParentIds: rootIDs, - }) - require.NoError(t, err) - var exploreChildren []database.Chat - for _, candidate := range childRows { - if candidate.Chat.Mode.Valid && candidate.Chat.Mode.ChatMode == database.ChatModeExplore { - exploreChildren = append(exploreChildren, candidate.Chat) - } - } - require.Len(t, exploreChildren, 1) -} - -func TestExploreChatUsesPersistedMCPSnapshot(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - ctx := testutil.Context(t, testutil.WaitLong) - - externalMCP := newTestMCPServer("external-snapshot-mcp") - addTestMCPTextTool(externalMCP, "echo", "Echoes the input", "echo: ") - externalMCPServer := httptest.NewServer(testMCPHTTPHandler(externalMCP)) - defer externalMCPServer.Close() - - secondMCP := newTestMCPServer("second-mcp") - addTestMCPTextTool(secondMCP, "echo", "Echoes the input", "echo: ") - secondMCPServer := httptest.NewServer(testMCPHTTPHandler(secondMCP)) - defer secondMCPServer.Close() - - var ( - requestsMu sync.Mutex - requests []recordedOpenAIRequest - ) - openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { - if !req.Stream { - return chattest.OpenAINonStreamingResponse("ok") - } - - requestsMu.Lock() - requests = append(requests, recordOpenAIRequest(req)) - requestsMu.Unlock() - - return chattest.OpenAIStreamingResponse( - chattest.OpenAITextChunks("done")..., - ) - }) - - user, org, _ := seedChatDependenciesWithProvider(t, db, "openai", openAIURL) - webSearchEnabled := true - storeEnabled := true - // OpenAI only serializes web_search through the Responses API. - // Store=true routes there only for supported Responses models. - webSearchModel := insertChatModelConfigWithCallConfig( - t, - db, - user.ID, - "openai", - "gpt-4o", - codersdk.ChatModelCallConfig{ - ProviderOptions: &codersdk.ChatModelProviderOptions{ - OpenAI: &codersdk.ChatModelOpenAIProviderOptions{ - Store: &storeEnabled, - WebSearchEnabled: &webSearchEnabled, - }, - }, - }, - ) - mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{ - OrganizationID: org.ID, - DisplayName: "External Snapshot MCP", - Slug: "external-snapshot-mcp", - Url: externalMCPServer.URL, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - }) - dbgen.MCPServerConfig(t, db, database.MCPServerConfig{ - OrganizationID: org.ID, - DisplayName: "Second MCP", - Slug: "second-mcp", - Url: secondMCPServer.URL, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - }) - - ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID) - rootChat := dbgen.Chat(t, db, database.Chat{ - OrganizationID: org.ID, - OwnerID: user.ID, - WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, - AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, - LastModelConfigID: webSearchModel.ID, - Title: "root", - ClientType: database.ChatClientTypeApi, - }) - - userContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{ - codersdk.ChatMessageText("inspect the codebase"), - }) - require.NoError(t, err) - createdExplore, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{ - OrganizationID: org.ID, - OwnerID: user.ID, - WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, - AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, - ParentChatID: uuid.NullUUID{UUID: rootChat.ID, Valid: true}, - RootChatID: uuid.NullUUID{UUID: rootChat.ID, Valid: true}, - LastModelConfigID: webSearchModel.ID, - Title: "explore", - Mode: database.NullChatMode{ - ChatMode: database.ChatModeExplore, - Valid: true, - }, - MCPServerIDs: []uuid.UUID{mcpConfig.ID}, - ClientType: database.ChatClientTypeApi, - InitialMessages: []chatstate.Message{ - { - Role: database.ChatMessageRoleUser, - Content: userContent, - Visibility: database.ChatMessageVisibilityBoth, - ContentVersion: chatprompt.CurrentContentVersion, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: webSearchModel.ID, Valid: true}, - }, - }, - }) - require.NoError(t, err) - exploreChat := createdExplore.Chat - - ctrl := gomock.NewController(t) - mockConn := agentconnmock.NewMockAgentConn(ctrl) - mockConn.EXPECT().SetExtraHeaders(gomock.Any()).AnyTimes() - mockConn.EXPECT().ContextConfig(gomock.Any()). - Return(workspacesdk.ContextConfigResponse{}, xerrors.New("not supported")).AnyTimes() - workspaceToolName := "workspace-snapshot-mcp__echo" - mockConn.EXPECT().LS(gomock.Any(), gomock.Any(), gomock.Any()). - Return(workspacesdk.LSResponse{AbsolutePathString: "/home/coder"}, nil).AnyTimes() - mockConn.EXPECT().ReadFile(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). - Return(io.NopCloser(strings.NewReader("")), "", nil).AnyTimes() - - factory := chattest.NewMockAIBridgeTransport(t, openAIURL) - server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) { - withoutMCPToolSearch(cfg) - cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory) - cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) { - require.Equal(t, dbAgent.ID, agentID) - return mockConn, func() {}, nil - } - }) - _ = server - - chatResult := waitForTerminalChat(ctx, t, db, exploreChat.ID) - if chatResult.Status == database.ChatStatusError { - require.FailNowf(t, "explore chat failed", "last_error=%q", chatLastErrorMessage(chatResult.LastError)) - } - - requestsMu.Lock() - recorded := append([]recordedOpenAIRequest(nil), requests...) - requestsMu.Unlock() - require.Len(t, recorded, 1) - - tools := recorded[0].Tools - require.Contains(t, tools, "read_file") - require.Contains(t, tools, "execute") - require.Contains(t, tools, "process_output") - require.Contains(t, tools, "external-snapshot-mcp__echo") - require.Contains(t, tools, "web_search", "Explore provider tool filter should let web_search through when the current model supports it") - require.NotContains(t, tools, "second-mcp__echo") - require.NotContains(t, tools, workspaceToolName) - require.NotContains(t, tools, "write_file") - require.NotContains(t, tools, "edit_files") - require.NotContains(t, tools, "spawn_agent") -} - -func TestRootExploreChatStaysBuiltinOnlyAtRuntime(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - ctx := testutil.Context(t, testutil.WaitLong) - - externalMCP := newTestMCPServer("root-explore-runtime-mcp") - addTestMCPTextTool(externalMCP, "echo", "Echoes the input", "echo: ") - externalMCPServer := httptest.NewServer(testMCPHTTPHandler(externalMCP)) - defer externalMCPServer.Close() - - var ( - requestsMu sync.Mutex - requests []recordedOpenAIRequest - ) - openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { - if !req.Stream { - return chattest.OpenAINonStreamingResponse("ok") - } - - requestsMu.Lock() - requests = append(requests, recordOpenAIRequest(req)) - requestsMu.Unlock() - - return chattest.OpenAIStreamingResponse( - chattest.OpenAITextChunks("done")..., - ) - }) - - user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL) - mcpConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{ - OrganizationID: org.ID, - DisplayName: "Root Explore Runtime MCP", - Slug: "root-explore-runtime-mcp", - Url: externalMCPServer.URL, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - }) - - factory := chattest.NewMockAIBridgeTransport(t, openAIURL) - server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) { - cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory) - }) - - exploreChat, err := server.CreateChat(ctx, chatd.CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - Title: "root-explore-builtin-only", - ModelConfigID: model.ID, - ChatMode: database.NullChatMode{ - ChatMode: database.ChatModeExplore, - Valid: true, - }, - MCPServerIDs: []uuid.UUID{mcpConfig.ID}, - InitialUserContent: []codersdk.ChatMessagePart{ - codersdk.ChatMessageText("Inspect the codebase."), - }, - }) - require.NoError(t, err) - waitForChatProcessed(ctx, t, db, exploreChat.ID, server) - - storedChat, err := db.GetChatByID(ctx, exploreChat.ID) - require.NoError(t, err) - if storedChat.Status == database.ChatStatusError { - require.FailNowf(t, "explore chat failed", "last_error=%q", chatLastErrorMessage(storedChat.LastError)) - } - require.Equal(t, database.ChatStatusWaiting, storedChat.Status) - require.ElementsMatch(t, []uuid.UUID{mcpConfig.ID}, storedChat.MCPServerIDs) - - requestsMu.Lock() - recorded := append([]recordedOpenAIRequest(nil), requests...) - requestsMu.Unlock() - require.Len(t, recorded, 1) - - tools := recorded[0].Tools - require.Contains(t, tools, "read_file") - require.Contains(t, tools, "execute") - require.NotContains(t, tools, "write_file") - require.NotContains(t, tools, "root-explore-runtime-mcp__echo", - "root Explore chats should strip persisted external MCP tools at runtime") -} - -func TestRootExploreChatExcludesWebSearchProviderToolAtRuntime(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - ctx := testutil.Context(t, testutil.WaitLong) - - var ( - requestsMu sync.Mutex - requests []recordedOpenAIRequest - ) - openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { - if !req.Stream { - return chattest.OpenAINonStreamingResponse("ok") - } - - requestsMu.Lock() - requests = append(requests, recordOpenAIRequest(req)) - requestsMu.Unlock() - - return chattest.OpenAIStreamingResponse( - chattest.OpenAITextChunks("done")..., - ) - }) - - user, org, _ := seedChatDependenciesWithProvider(t, db, "openai", openAIURL) - webSearchEnabled := true - storeEnabled := true - // OpenAI only serializes web_search through the Responses API. - // Store=true routes there only for supported Responses models. - webSearchModel := insertChatModelConfigWithCallConfig( - t, - db, - user.ID, - "openai", - "gpt-4o", - codersdk.ChatModelCallConfig{ - ProviderOptions: &codersdk.ChatModelProviderOptions{ - OpenAI: &codersdk.ChatModelOpenAIProviderOptions{ - Store: &storeEnabled, - WebSearchEnabled: &webSearchEnabled, - }, - }, - }, - ) - - factory := chattest.NewMockAIBridgeTransport(t, openAIURL) - server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) { - cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory) - }) - - exploreChat, err := server.CreateChat(ctx, chatd.CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - Title: "root-explore-no-provider-web-search", - ModelConfigID: webSearchModel.ID, - ChatMode: database.NullChatMode{ - ChatMode: database.ChatModeExplore, - Valid: true, - }, - InitialUserContent: []codersdk.ChatMessagePart{ - codersdk.ChatMessageText("Inspect the codebase."), - }, - }) - require.NoError(t, err) - waitForChatProcessed(ctx, t, db, exploreChat.ID, server) - - storedChat, err := db.GetChatByID(ctx, exploreChat.ID) - require.NoError(t, err) - if storedChat.Status == database.ChatStatusError { - require.FailNowf(t, "explore chat failed", "last_error=%q", chatLastErrorMessage(storedChat.LastError)) - } - require.Equal(t, database.ChatStatusWaiting, storedChat.Status) - - requestsMu.Lock() - recorded := append([]recordedOpenAIRequest(nil), requests...) - requestsMu.Unlock() - require.Len(t, recorded, 1) - - tools := recorded[0].Tools - require.Contains(t, tools, "read_file") - require.Contains(t, tools, "execute") - require.NotContains(t, tools, "web_search", - "root Explore chats should stay builtin-only and must not inherit provider-native web_search at runtime") - require.NotContains(t, tools, "write_file") -} - -func TestExploreChatSendMessageCannotMutateMCPSnapshot(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - ctx := testutil.Context(t, testutil.WaitLong) - - newEchoMCPServer := func(name string) *httptest.Server { - t.Helper() - - mcpSrv := newTestMCPServer(name) - addTestMCPTextTool(mcpSrv, "echo", "Echoes the input", "echo: ") - mcpTS := httptest.NewServer(testMCPHTTPHandler(mcpSrv)) - t.Cleanup(mcpTS.Close) - return mcpTS - } - - parentTS := newEchoMCPServer("runtime-parent-mcp") - injectedTS := newEchoMCPServer("runtime-injected-mcp") - - var ( - requestsMu sync.Mutex - requests []recordedOpenAIRequest - ) - childRequests := func() []recordedOpenAIRequest { - requestsMu.Lock() - defer requestsMu.Unlock() - - filtered := make([]recordedOpenAIRequest, 0, len(requests)) - for _, req := range requests { - if requestHasSystemSubstring(req, "You are in Explore Mode as a delegated sub-agent.") { - filtered = append(filtered, req) - } - } - return filtered - } - - var streamCallCount atomic.Int32 - openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { - if !req.Stream { - return chattest.OpenAINonStreamingResponse("ok") - } - - requestsMu.Lock() - requests = append(requests, recordOpenAIRequest(req)) - requestsMu.Unlock() - - if streamCallCount.Add(1) == 1 { - return chattest.OpenAIStreamingResponse( - chattest.OpenAIToolCallChunk("spawn_agent", `{"type":"explore","prompt":"inspect the codebase","title":"sub"}`), - ) - } - - return chattest.OpenAIStreamingResponse( - chattest.OpenAITextChunks("done")..., - ) - }) - - user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL) - parentConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{ - OrganizationID: org.ID, - DisplayName: "Runtime Parent MCP", - Slug: "runtime-parent-mcp", - Url: parentTS.URL, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - }) - injectedConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{ - OrganizationID: org.ID, - DisplayName: "Runtime Injected MCP", - Slug: "runtime-injected-mcp", - Url: injectedTS.URL, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - }) - - factory := chattest.NewMockAIBridgeTransport(t, openAIURL) - server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) { - withoutMCPToolSearch(cfg) - cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory) - }) - - rootChat, err := server.CreateChat(ctx, chatd.CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - Title: "runtime-parent", - ModelConfigID: model.ID, - MCPServerIDs: []uuid.UUID{parentConfig.ID}, - InitialUserContent: []codersdk.ChatMessagePart{ - codersdk.ChatMessageText("Spawn an Explore subagent to inspect the codebase."), - }, - }) - require.NoError(t, err) - - var exploreChat database.Chat - testutil.Eventually(ctx, t, func(ctx context.Context) bool { - childRows, err := db.GetChildChatsByParentIDs(dbauthz.AsChatd(ctx), database.GetChildChatsByParentIDsParams{ - ParentIds: []uuid.UUID{rootChat.ID}, - }) - if err != nil { - return false - } - for _, candidate := range childRows { - if candidate.Chat.Mode.Valid && candidate.Chat.Mode.ChatMode == database.ChatModeExplore { - exploreChat = candidate.Chat - return true - } - } - return false - }, testutil.IntervalFast) - - chatResult := waitForTerminalChat(ctx, t, db, exploreChat.ID) - if chatResult.Status == database.ChatStatusError { - require.FailNowf(t, "explore chat failed", "last_error=%q", chatLastErrorMessage(chatResult.LastError)) - } - - exploreChat, err = db.GetChatByID(ctx, exploreChat.ID) - require.NoError(t, err) - require.ElementsMatch(t, []uuid.UUID{parentConfig.ID}, exploreChat.MCPServerIDs) - - initialChildRequestCount := len(childRequests()) - require.GreaterOrEqual(t, initialChildRequestCount, 1) - - updatedMCPServerIDs := []uuid.UUID{injectedConfig.ID} - _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ - ChatID: exploreChat.ID, - CreatedBy: user.ID, - Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("inspect the codebase again")}, - MCPServerIDs: &updatedMCPServerIDs, - }) - require.NoError(t, err) - - storedExploreChat, err := db.GetChatByID(ctx, exploreChat.ID) - require.NoError(t, err) - require.ElementsMatch(t, []uuid.UUID{parentConfig.ID}, storedExploreChat.MCPServerIDs) - - testutil.Eventually(ctx, t, func(ctx context.Context) bool { - return len(childRequests()) > initialChildRequestCount - }, testutil.IntervalFast) - - chatResult = waitForTerminalChat(ctx, t, db, exploreChat.ID) - if chatResult.Status == database.ChatStatusError { - require.FailNowf(t, "explore chat failed", "last_error=%q", chatLastErrorMessage(chatResult.LastError)) - } - - recordedChildRequests := childRequests() - require.GreaterOrEqual(t, len(recordedChildRequests), initialChildRequestCount+1) - - tools := recordedChildRequests[len(recordedChildRequests)-1].Tools - require.Contains(t, tools, "runtime-parent-mcp__echo") - require.NotContains(t, tools, "runtime-injected-mcp__echo", - "Explore child runtime should keep the spawn-time MCP snapshot after SendMessage") -} - -func TestPlanModeRootChatAllowsApprovedExternalMCPTools(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - ctx := testutil.Context(t, testutil.WaitLong) - - echoMCP := newTestMCPServer("plan-visibility-echo") - addTestMCPTextTool(echoMCP, "echo", "Echoes the input", "echo: ") - echoTS := httptest.NewServer(testMCPHTTPHandler(echoMCP)) - t.Cleanup(echoTS.Close) - - filteredMCP := newTestMCPServer("plan-visibility-filtered") - addTestMCPTextTool(filteredMCP, "visible", "Visible tool", "visible: ") - addTestMCPTextTool(filteredMCP, "hidden", "Hidden tool", "hidden: ") - filteredTS := httptest.NewServer(testMCPHTTPHandler(filteredMCP)) - t.Cleanup(filteredTS.Close) - - var ( - requests []recordedOpenAIRequest - requestsMu sync.Mutex - ) - openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { - if !req.Stream { - return chattest.OpenAINonStreamingResponse("title") - } - - requestsMu.Lock() - requests = append(requests, recordOpenAIRequest(req)) - requestsMu.Unlock() - - return chattest.OpenAIStreamingResponse( - chattest.OpenAITextChunks("Done.")..., - ) - }) - - user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL) - - approvedConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{ - OrganizationID: org.ID, - DisplayName: "Plan Approved MCP", - Slug: "plan-approved-mcp", - Url: echoTS.URL, - AllowInPlanMode: true, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - }) - - blockedConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{ - OrganizationID: org.ID, - DisplayName: "Plan Blocked MCP", - Slug: "plan-blocked-mcp", - Url: echoTS.URL, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - }) - - filteredConfig := dbgen.MCPServerConfig(t, db, database.MCPServerConfig{ - OrganizationID: org.ID, - DisplayName: "Plan Filtered MCP", - Slug: "plan-filtered-mcp", - Url: filteredTS.URL, - AllowInPlanMode: true, - ToolAllowList: []string{"visible"}, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - }) - - ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID) - // Workspace MCP tools now come from the agent's pinned snapshot, not live - // discovery. Seed the workspace MCP server so chats bound to the agent - // hydrate the "workspace-plan-mcp__echo" tool. - seedAgentMCPToolContext(ctx, t, db, agentMCPToolContext{ - AgentID: dbAgent.ID, - ServerName: "workspace-plan-mcp", - ToolName: "echo", - ToolDescription: "Workspace echo tool", - }) - ctrl := gomock.NewController(t) - mockConn := agentconnmock.NewMockAgentConn(ctrl) - mockConn.EXPECT().SetExtraHeaders(gomock.Any()).AnyTimes() - mockConn.EXPECT().ContextConfig(gomock.Any()). - Return(workspacesdk.ContextConfigResponse{}, xerrors.New("not supported")).AnyTimes() - workspaceToolName := "workspace-plan-mcp__echo" - mockConn.EXPECT().LS(gomock.Any(), gomock.Any(), gomock.Any()). - Return(workspacesdk.LSResponse{AbsolutePathString: "/home/coder"}, nil).AnyTimes() - mockConn.EXPECT().ReadFile(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). - Return(io.NopCloser(strings.NewReader("")), "", nil).AnyTimes() - - factory := chattest.NewMockAIBridgeTransport(t, openAIURL) - server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) { - withoutMCPToolSearch(cfg) - cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory) - cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) { - require.Equal(t, dbAgent.ID, agentID) - return mockConn, func() {}, nil - } - }) - - planChat, err := server.CreateChat(ctx, chatd.CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - Title: "plan-mode-root-mcp-visibility", - ModelConfigID: model.ID, - WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, - PlanMode: database.NullChatPlanMode{ChatPlanMode: database.ChatPlanModePlan, Valid: true}, - MCPServerIDs: []uuid.UUID{approvedConfig.ID, blockedConfig.ID, filteredConfig.ID}, - InitialUserContent: []codersdk.ChatMessagePart{ - codersdk.ChatMessageText("List the available tools in plan mode."), - }, - }) - require.NoError(t, err) - waitForChatProcessed(ctx, t, db, planChat.ID, server) - - planChatResult, err := db.GetChatByID(ctx, planChat.ID) - require.NoError(t, err) - require.Equal(t, database.ChatStatusWaiting, planChatResult.Status) - - askChat, err := server.CreateChat(ctx, chatd.CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - Title: "ask-mode-root-mcp-visibility", - ModelConfigID: model.ID, - WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, - MCPServerIDs: []uuid.UUID{approvedConfig.ID, blockedConfig.ID, filteredConfig.ID}, - InitialUserContent: []codersdk.ChatMessagePart{ - codersdk.ChatMessageText("List the available tools outside plan mode."), - }, - }) - require.NoError(t, err) - waitForChatProcessed(ctx, t, db, askChat.ID, server) - - askChatResult, err := db.GetChatByID(ctx, askChat.ID) - require.NoError(t, err) - require.Equal(t, database.ChatStatusWaiting, askChatResult.Status) - - requestsMu.Lock() - recorded := append([]recordedOpenAIRequest(nil), requests...) - requestsMu.Unlock() - require.Len(t, recorded, 2, "expected exactly one streamed model call per chat") - - planTools := recorded[0].Tools - askTools := recorded[1].Tools - - require.Contains(t, planTools, "plan-approved-mcp__echo", - "root plan mode should expose approved external MCP tools") - require.NotContains(t, planTools, "plan-blocked-mcp__echo", - "root plan mode should hide unapproved external MCP tools") - require.Contains(t, planTools, "plan-filtered-mcp__visible", - "root plan mode should keep allowlisted tools from approved MCP servers") - require.NotContains(t, planTools, "plan-filtered-mcp__hidden", - "root plan mode should still respect MCP tool allowlists") - require.NotContains(t, planTools, workspaceToolName, - "root plan mode should exclude workspace MCP tools") - - require.Contains(t, askTools, "plan-approved-mcp__echo", - "ask mode should keep approved external MCP tools") - require.Contains(t, askTools, "plan-blocked-mcp__echo", - "ask mode should keep unapproved-for-plan external MCP tools") - require.Contains(t, askTools, "plan-filtered-mcp__visible", - "ask mode should keep allowlisted tools from external MCP servers") - require.NotContains(t, askTools, "plan-filtered-mcp__hidden", - "ask mode should continue respecting MCP tool allowlists") - require.Contains(t, askTools, workspaceToolName, - "ask mode should continue exposing workspace MCP tools") -} - -// TestUnarchiveChildChat covers the deterministic branches of the -// Server.UnarchiveChat child path: every child unarchive attempt is -// rejected with chatd.ErrArchiveRequiresRootChat. func TestUnarchiveChildChat(t *testing.T) { t.Parallel() diff --git a/coderd/x/chatd/generation.go b/coderd/x/chatd/generation.go index 6f4e48c41f5..435f833731c 100644 --- a/coderd/x/chatd/generation.go +++ b/coderd/x/chatd/generation.go @@ -423,7 +423,7 @@ func (s *taskStarter) StartGeneration(ctx context.Context, input chatWorkerTaskS RecordMCPConnectSummaries: input.DebugTurn.RecordMCPConnectSummaries, } prepared, err := retryGenerationPhase(ctx, s, "prepare", func() (turnEnvironment, error) { - return s.server.buildTurnEnvironment(ctx, prepareInput) + return buildTurnEnvironment(ctx, s.server, prepareInput) }) if err != nil { if errors.Is(err, errTaskExpectedExit) || errors.Is(err, errTaskRetryable) { @@ -763,7 +763,7 @@ func (s *taskStarter) admitStepToolCalls( // committed, so its find_tools calls would otherwise never reach // the executeLocalTools counter; count them at each error exit. countBatch := func() { - if !prepared.Toolset().builtinToolNames[chattool.FindToolsName] { + if !prepared.Toolset().IsBuiltin(chattool.FindToolsName) { return } for _, toolCall := range toolCalls { @@ -778,13 +778,13 @@ func (s *taskStarter) admitStepToolCalls( countBatch() return chathooks.PreToolUseExecutionResult{}, chathooks.GenerationDispatchError(agenthooks.EventPreToolUse, err) } - unambiguous, _, ambiguous := partitionAmbiguousToolCalls(prepared, toolCalls) + unambiguous, _, ambiguous := partitionAmbiguousToolCalls(prepared.Toolset(), toolCalls) preflight, err := s.server.hooks.PreflightPendingToolCalls(ctx, chathooks.ChatFor(prepared.Turn().chat, input.hookTurnID()), unambiguous) if err != nil { countBatch() return chathooks.PreToolUseExecutionResult{}, chathooks.GenerationDispatchError(agenthooks.EventPreToolUse, err) } - if err := validateOverriddenToolInputs(prepared, preflight); err != nil { + if err := validateOverriddenToolInputs(prepared.Toolset(), preflight); err != nil { countBatch() return chathooks.PreToolUseExecutionResult{}, chathooks.GenerationDispatchError(agenthooks.EventPreToolUse, err) } @@ -792,7 +792,7 @@ func (s *taskStarter) admitStepToolCalls( // Calls denied at admission persist synthetic results with the // assistant step, so they never surface as unresolved calls where // executeLocalTools counts find_tools invocations; count them here. - if prepared.Toolset().builtinToolNames[chattool.FindToolsName] { + if prepared.Toolset().IsBuiltin(chattool.FindToolsName) { for _, result := range preflight.Denied { if result.ToolName == chattool.FindToolsName { s.server.metrics.FindToolsCallsTotal.Inc() @@ -836,13 +836,13 @@ func (s *taskStarter) executeLocalTools( var allowedIndexes []int var denied []fantasy.ToolResultContent if !exclusiveRejected { - allowed, allowedIndexes, denied = partitionAmbiguousToolCalls(prepared, decision.localToolCalls) + allowed, allowedIndexes, denied = partitionAmbiguousToolCalls(prepared.Toolset(), decision.localToolCalls) } // find_tools calls are counted here, at the single point every // model-emitted call passes through, because rejections upstream of // the tool (partition denials, hook denials, exclusive-policy // batches) never reach its handler or OnCall. - if prepared.Toolset().builtinToolNames[chattool.FindToolsName] { + if prepared.Toolset().IsBuiltin(chattool.FindToolsName) { for _, toolCall := range decision.localToolCalls { if toolCall.ToolName == chattool.FindToolsName { s.server.metrics.FindToolsCallsTotal.Inc() diff --git a/coderd/x/chatd/toolinput.go b/coderd/x/chatd/toolinput.go index 6cbe49fe0b1..9a2018c140d 100644 --- a/coderd/x/chatd/toolinput.go +++ b/coderd/x/chatd/toolinput.go @@ -18,7 +18,7 @@ import ( // of a dispatch failure. allowedIndexes maps allowed calls back to the input // order without relying on duplicate-prone IDs. func partitionAmbiguousToolCalls( - prepared turnEnvironment, + toolset *turnToolset, toolCalls []fantasy.ToolCallContent, ) (allowed []fantasy.ToolCallContent, allowedIndexes []int, rejected []fantasy.ToolResultContent) { for i, toolCall := range toolCalls { @@ -26,7 +26,7 @@ func partitionAmbiguousToolCalls( rejected = append(rejected, malformedToolResult(toolCall)) continue } - if err := validateBuiltinToolInput(prepared, toolCall.ToolName, []byte(toolCall.Input)); err != nil { + if err := validateBuiltinToolInput(toolset, toolCall.ToolName, []byte(toolCall.Input)); err != nil { rejected = append(rejected, ambiguousToolResult(toolCall, err)) continue } @@ -39,12 +39,12 @@ func partitionAmbiguousToolCalls( // validateOverriddenToolInputs rechecks the inputs a pre_tool_use consumer // replaced. The model cannot fix an ambiguous override, so the turn fails // closed instead of executing it. -func validateOverriddenToolInputs(prepared turnEnvironment, preflight chathooks.PreToolUseExecutionResult) error { +func validateOverriddenToolInputs(toolset *turnToolset, preflight chathooks.PreToolUseExecutionResult) error { for _, toolCall := range preflight.Allowed { if _, overridden := preflight.Overrides[toolCall.ToolCallID]; !overridden { continue } - if err := validateBuiltinToolInput(prepared, toolCall.ToolName, []byte(toolCall.Input)); err != nil { + if err := validateBuiltinToolInput(toolset, toolCall.ToolName, []byte(toolCall.Input)); err != nil { return xerrors.Errorf("hook input override for tool %s: %w", toolCall.ToolName, err) } } @@ -54,16 +54,16 @@ func validateOverriddenToolInputs(prepared turnEnvironment, preflight chathooks. // validateBuiltinToolInput only guards builtin tools, whose input coderd // decodes itself. Dynamic calls are executed by the client and MCP calls by // their own server, and a dynamic tool cannot shadow a builtin name. -func validateBuiltinToolInput(prepared turnEnvironment, toolName string, input []byte) error { +func validateBuiltinToolInput(toolset *turnToolset, toolName string, input []byte) error { // Execution resolves a deprecated alias to its canonical tool, so // skipping that here would let the old name bypass validation. if canonical, aliased := subagentToolNameAliases[toolName]; aliased { toolName = canonical } - if !prepared.Toolset().builtinToolNames[toolName] { + if !toolset.IsBuiltin(toolName) { return nil } - for _, tool := range prepared.Toolset().tools { + for _, tool := range toolset.tools { info := tool.Info() if info.Name != toolName { continue diff --git a/coderd/x/chatd/toolinput_internal_test.go b/coderd/x/chatd/toolinput_internal_test.go index fef0b29ca1c..efa77125f3e 100644 --- a/coderd/x/chatd/toolinput_internal_test.go +++ b/coderd/x/chatd/toolinput_internal_test.go @@ -36,11 +36,11 @@ func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) { t.Run("builtin", func(t *testing.T) { t.Parallel() - prepared := turnEnvironmentState{ - Tools: []fantasy.AgentTool{fetch}, - BuiltinToolNames: map[string]bool{"fetch": true}, + prepared := turnToolset{ + tools: []fantasy.AgentTool{fetch}, + builtinToolNames: map[string]bool{"fetch": true}, } - allowed, allowedIndexes, rejected := partitionAmbiguousToolCalls(prepared, []fantasy.ToolCallContent{ambiguous, clean}) + allowed, allowedIndexes, rejected := partitionAmbiguousToolCalls(&prepared, []fantasy.ToolCallContent{ambiguous, clean}) require.Len(t, rejected, 1) require.Equal(t, "call_ambiguous", rejected[0].ToolCallID) require.Len(t, allowed, 1) @@ -51,8 +51,8 @@ func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) { t.Run("non-builtin", func(t *testing.T) { t.Parallel() - prepared := turnEnvironmentState{Tools: []fantasy.AgentTool{fetch}} - allowed, allowedIndexes, rejected := partitionAmbiguousToolCalls(prepared, []fantasy.ToolCallContent{ambiguous, clean}) + prepared := turnToolset{tools: []fantasy.AgentTool{fetch}} + allowed, allowedIndexes, rejected := partitionAmbiguousToolCalls(&prepared, []fantasy.ToolCallContent{ambiguous, clean}) require.Empty(t, rejected) require.Len(t, allowed, 2) require.Equal(t, []int{0, 1}, allowedIndexes) @@ -83,11 +83,11 @@ func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) { Input: `{"chat_id":"a","CHAT_ID":"b"}`, } - prepared := turnEnvironmentState{ - Tools: []fantasy.AgentTool{tool}, - BuiltinToolNames: map[string]bool{canonical: true}, + prepared := turnToolset{ + tools: []fantasy.AgentTool{tool}, + builtinToolNames: map[string]bool{canonical: true}, } - _, _, rejected := partitionAmbiguousToolCalls(prepared, []fantasy.ToolCallContent{aliased}) + _, _, rejected := partitionAmbiguousToolCalls(&prepared, []fantasy.ToolCallContent{aliased}) require.Len(t, rejected, 1) }) } @@ -95,9 +95,9 @@ func TestPartitionAmbiguousToolCallsGatesOnBuiltins(t *testing.T) { func TestValidateOverriddenToolInputs(t *testing.T) { t.Parallel() - prepared := turnEnvironmentState{ - Tools: []fantasy.AgentTool{fetchToolStub()}, - BuiltinToolNames: map[string]bool{"fetch": true}, + prepared := turnToolset{ + tools: []fantasy.AgentTool{fetchToolStub()}, + builtinToolNames: map[string]bool{"fetch": true}, } overridden := chathooks.PreToolUseExecutionResult{ Allowed: []fantasy.ToolCallContent{{ @@ -109,13 +109,13 @@ func TestValidateOverriddenToolInputs(t *testing.T) { "call_overridden": json.RawMessage(`{"URL":"https://other.test"}`), }, } - require.ErrorContains(t, validateOverriddenToolInputs(prepared, overridden), + require.ErrorContains(t, validateOverriddenToolInputs(&prepared, overridden), `hook input override for tool fetch: input key "URL" differs from schema property "url" only by case`) // The same input is left alone when no consumer replaced it, because // the model-authored batch is checked before the dispatch instead. untouched := chathooks.PreToolUseExecutionResult{Allowed: overridden.Allowed} - require.NoError(t, validateOverriddenToolInputs(prepared, untouched)) + require.NoError(t, validateOverriddenToolInputs(&prepared, untouched)) } // TestBuiltinToolSchemasDescribeTheirInputs guards the validator's reach: it @@ -165,7 +165,7 @@ func TestBuiltinToolSchemasDescribeTheirInputs(t *testing.T) { chatprovider.ProviderAPIKeys{}, withInternalTestServerTransportFactory(&aibridgeTestFactory{}), ) - prepared, err := server.buildTurnEnvironment(ctx, generationPrepareInput{ + prepared, err := buildTurnEnvironment(ctx, server, generationPrepareInput{ Chat: created.Chat, Messages: created.InitialMessages, }) diff --git a/coderd/x/chatd/turn_environment.go b/coderd/x/chatd/turn_environment.go index acbc00457e0..ed5de674711 100644 --- a/coderd/x/chatd/turn_environment.go +++ b/coderd/x/chatd/turn_environment.go @@ -36,15 +36,11 @@ import ( "github.com/coder/coder/v2/codersdk/workspacesdk" ) -// effectiveMCPServerConfigs loads the chat's stored selection plus -// owner-readable Force On configs at generation time, so stored lists -// predating enforcement cannot dodge the policy (Cure53 CDM-02-010). -// Explore chats keep their immutable spawn-time snapshot instead. type turnEnvironment interface { - Turn() turnState - ModelConfig() turnModelConfig + Turn() *turnState + ModelConfig() *turnModelConfig Prompt() []fantasy.Message - Toolset() turnToolset + Toolset() *turnToolset CompactionConfig() *generationCompaction Close() } @@ -58,7 +54,6 @@ type turnState struct { type turnModelConfig struct { model chatprovider.Model - route aiGatewayModelRoute buildOptions modelBuildOptions resolvedProvider string configID uuid.UUID @@ -78,61 +73,40 @@ type turnToolset struct { toolNameToConfigID map[string]uuid.UUID } -type turnEnvironmentState struct { - Chat database.Chat - Messages []database.ChatMessage - - Model chatprovider.Model - PromptMessages []fantasy.Message - Tools []fantasy.AgentTool - ActiveTools []string - AllowInactiveTools map[string]bool - ProviderTools []chatloop.ProviderTool - ModelRoute aiGatewayModelRoute - ModelBuildOptions modelBuildOptions - ResolvedProvider string - ModelConfigID uuid.UUID - CallTemplate fantasy.Call - ContextLimitFallback int64 - - DynamicToolNames map[string]bool - StopAfterTools map[string]struct{} - ExclusiveToolNames map[string]bool - BuiltinToolNames map[string]bool - ToolNameToConfigID map[string]uuid.UUID - - MaxSteps int - Compaction *generationCompaction - Cleanup func() - Debug *generationDebug -} - -func (e turnEnvironmentState) Turn() turnState { - return turnState{chat: e.Chat, messages: e.Messages, maxSteps: e.MaxSteps, debug: e.Debug} +func (t turnToolset) IsExclusive(name string) bool { return t.exclusiveToolNames[name] } +func (t turnToolset) IsDynamic(name string) bool { return t.dynamicToolNames[name] } +func (t turnToolset) IsBuiltin(name string) bool { return t.builtinToolNames[name] } +func (t turnToolset) AllowsInactive(name string) bool { return t.allowInactiveTools[name] } +func (t turnToolset) StopsAfter(name string) bool { _, ok := t.stopAfterTools[name]; return ok } +func (t turnToolset) ConfigID(name string) (uuid.UUID, bool) { + id, ok := t.toolNameToConfigID[name] + return id, ok } -func (e turnEnvironmentState) ModelConfig() turnModelConfig { - return turnModelConfig{ - model: e.Model, route: e.ModelRoute, buildOptions: e.ModelBuildOptions, - resolvedProvider: e.ResolvedProvider, configID: e.ModelConfigID, - callTemplate: e.CallTemplate, contextLimitFallback: e.ContextLimitFallback, - } +type turnEnvironmentState struct { + turn turnState + model turnModelConfig + prompt []fantasy.Message + toolset turnToolset + compaction *generationCompaction + cleanup func() } -func (e turnEnvironmentState) Prompt() []fantasy.Message { return e.PromptMessages } - -func (e turnEnvironmentState) Toolset() turnToolset { - return turnToolset{ - tools: e.Tools, activeTools: e.ActiveTools, allowInactiveTools: e.AllowInactiveTools, - providerTools: e.ProviderTools, dynamicToolNames: e.DynamicToolNames, - stopAfterTools: e.StopAfterTools, exclusiveToolNames: e.ExclusiveToolNames, - builtinToolNames: e.BuiltinToolNames, toolNameToConfigID: e.ToolNameToConfigID, +func (e *turnEnvironmentState) Turn() *turnState { return &e.turn } +func (e *turnEnvironmentState) ModelConfig() *turnModelConfig { return &e.model } +func (e *turnEnvironmentState) Prompt() []fantasy.Message { return e.prompt } +func (e *turnEnvironmentState) Toolset() *turnToolset { return &e.toolset } +func (e *turnEnvironmentState) CompactionConfig() *generationCompaction { return e.compaction } +func (e *turnEnvironmentState) Close() { + if e.cleanup != nil { + e.cleanup() } } -func (e turnEnvironmentState) CompactionConfig() *generationCompaction { return e.Compaction } -func (e turnEnvironmentState) Close() { e.Cleanup() } - +// effectiveMCPServerConfigs loads the chat's stored selection plus +// owner-readable Force On configs at generation time, so stored lists +// predating enforcement cannot dodge the policy (Cure53 CDM-02-010). +// Explore chats keep their immutable spawn-time snapshot instead. func (server *Server) effectiveMCPServerConfigs( ctx context.Context, logger slog.Logger, @@ -171,10 +145,12 @@ func (server *Server) effectiveMCPServerConfigs( return configs, nil } -func (server *Server) buildTurnEnvironment( +func buildTurnEnvironment( ctx context.Context, + server *Server, input generationPrepareInput, ) (turnEnvironment, error) { + builder := turnEnvironmentBuilder{server: server} chat := input.Chat logger := server.logger.With( slog.F("chat_id", chat.ID), @@ -286,7 +262,7 @@ func (server *Server) buildTurnEnvironment( approvedPlanMCPConfigIDs = map[uuid.UUID]struct{}{} } - planModeInstructions := server.loadPlanModeInstructions(ctx, currentPlanMode, logger) + planModeInstructions := builder.loadPlanModeInstructions(ctx, currentPlanMode, logger) advisorCfg := server.loadAdvisorConfig(ctx, logger) // Force Enabled from the experiment; the stored DB value is ignored. advisorCfg.Enabled = server.experiments.Enabled(codersdk.ExperimentChatAdvisor) @@ -459,7 +435,7 @@ func (server *Server) buildTurnEnvironment( return nil }) g2.Go(func() error { - personalSkills = server.fetchPersonalSkillMetadata(ctx, chat.OwnerID, logger) + personalSkills = builder.fetchPersonalSkillMetadata(ctx, chat.OwnerID, logger) return nil }) g2.Go(func() error { @@ -593,7 +569,7 @@ func (server *Server) buildTurnEnvironment( tools = append(tools, chattool.NewAskUserQuestionTool()) } if isRootChat { - tools = server.appendRootChatTools(ctx, tools, rootChatToolsOptions{ + tools = builder.appendRootChatTools(ctx, tools, rootChatToolsOptions{ chat: chat, modelConfigID: modelConfig.ID, workspaceCtx: &workspaceCtx, @@ -611,7 +587,7 @@ func (server *Server) buildTurnEnvironment( }, ResolveAlias: resolveSkillAlias, LoadPersonalSkillBody: func(ctx context.Context, name string) (skillspkg.ParsedSkill, error) { - return server.loadPersonalSkillBody(ctx, chat.OwnerID, name) + return builder.loadPersonalSkillBody(ctx, chat.OwnerID, name) }, } appendCurrentSkillTools := func(current []fantasy.AgentTool) ([]fantasy.AgentTool, bool) { @@ -831,35 +807,31 @@ func (server *Server) buildTurnEnvironment( refreshedChat = chat } - return turnEnvironmentState{ - Chat: refreshedChat, - Messages: input.Messages, - Model: model, - PromptMessages: prompt, - Tools: tools, - ActiveTools: activeToolNames, - AllowInactiveTools: allowInactiveTools, - ProviderTools: providerTools, - ModelRoute: modelRoute, - ModelBuildOptions: modelOpts, - ResolvedProvider: resolved.resolvedProvider, - ModelConfigID: modelConfig.ID, - CallTemplate: resolved.newCall(), - ContextLimitFallback: modelConfig.ContextLimit, - DynamicToolNames: dynamicToolNames, - StopAfterTools: stopAfterBehaviorTools(currentPlanMode, chat.Mode, chat.ParentChatID), - ExclusiveToolNames: exclusiveToolNames, - BuiltinToolNames: builtinToolNames, - ToolNameToConfigID: toolNameToConfigID, - MaxSteps: maxChatSteps, - Compaction: &generationCompaction{ - Override: compactionOverride, - ChatModelConfig: modelConfig, - Required: compactionNeeded, - Options: compactionOptions, + return &turnEnvironmentState{ + turn: turnState{ + chat: refreshedChat, messages: input.Messages, + maxSteps: maxChatSteps, debug: debug, }, - Cleanup: cleanup, - Debug: debug, + model: turnModelConfig{ + model: model, buildOptions: modelOpts, + resolvedProvider: resolved.resolvedProvider, + configID: modelConfig.ID, callTemplate: resolved.newCall(), + contextLimitFallback: modelConfig.ContextLimit, + }, + prompt: prompt, + toolset: turnToolset{ + tools: tools, activeTools: activeToolNames, + allowInactiveTools: allowInactiveTools, providerTools: providerTools, + dynamicToolNames: dynamicToolNames, + stopAfterTools: stopAfterBehaviorTools(currentPlanMode, chat.Mode, chat.ParentChatID), + exclusiveToolNames: exclusiveToolNames, builtinToolNames: builtinToolNames, + toolNameToConfigID: toolNameToConfigID, + }, + compaction: &generationCompaction{ + Override: compactionOverride, ChatModelConfig: modelConfig, + Required: compactionNeeded, Options: compactionOptions, + }, + cleanup: cleanup, }, nil } @@ -1655,8 +1627,6 @@ func (c *turnWorkspaceContext) getWorkspaceConn(ctx context.Context) (workspaces return nil, xerrors.New("chat workspace changed while connecting") } -// AgentConnFunc provides access to workspace agent connections. - func allToolNames(allTools []fantasy.AgentTool) []string { toolNames := make([]string, 0, len(allTools)) for _, tool := range allTools { @@ -1958,6 +1928,10 @@ func isSkillIndexMessage(message fantasy.Message) bool { return strings.HasPrefix(text, chattool.AvailableSkillsOpenTag+"\n") && strings.HasSuffix(text, chattool.AvailableSkillsCloseTag) } +type turnEnvironmentBuilder struct { + server *Server +} + type rootChatToolsOptions struct { chat database.Chat modelConfigID uuid.UUID @@ -1968,7 +1942,7 @@ type rootChatToolsOptions struct { isPlanModeTurn bool } -func (server *Server) loadPlanModeInstructions( +func (builder turnEnvironmentBuilder) loadPlanModeInstructions( ctx context.Context, mode database.NullChatPlanMode, logger slog.Logger, @@ -1981,7 +1955,7 @@ func (server *Server) loadPlanModeInstructions( // not carry a deployment-config actor during background execution. //nolint:gocritic // Required to read deployment config during background chat processing. systemCtx := dbauthz.AsSystemRestricted(ctx) - fetched, err := server.db.GetChatPlanModeInstructions(systemCtx) + fetched, err := builder.server.db.GetChatPlanModeInstructions(systemCtx) if err != nil { logger.Warn(ctx, "failed to fetch plan mode instructions", @@ -2011,12 +1985,12 @@ func userSkillContext(ctx context.Context, userID uuid.UUID) context.Context { return dbauthz.As(ctx, actor) } -func (server *Server) fetchPersonalSkillMetadata( +func (builder turnEnvironmentBuilder) fetchPersonalSkillMetadata( ctx context.Context, userID uuid.UUID, logger slog.Logger, ) []skillspkg.Skill { - rows, err := server.db.ListUserSkillMetadataByUserID(userSkillContext(ctx, userID), userID) + rows, err := builder.server.db.ListUserSkillMetadataByUserID(userSkillContext(ctx, userID), userID) // See package coderd/x/skills (doc.go) for why metadata fetch failures // intentionally degrade to an empty personal-skill list instead of // failing the chat turn. @@ -2039,12 +2013,12 @@ func (server *Server) fetchPersonalSkillMetadata( return personalSkills } -func (server *Server) loadPersonalSkillBody( +func (builder turnEnvironmentBuilder) loadPersonalSkillBody( ctx context.Context, userID uuid.UUID, name string, ) (skillspkg.ParsedSkill, error) { - row, err := server.db.GetUserSkillByUserIDAndName( + row, err := builder.server.db.GetUserSkillByUserIDAndName( userSkillContext(ctx, userID), database.GetUserSkillByUserIDAndNameParams{ UserID: userID, @@ -2055,7 +2029,7 @@ func (server *Server) loadPersonalSkillBody( if errors.Is(err, sql.ErrNoRows) { return skillspkg.ParsedSkill{}, skillspkg.ErrSkillNotFound } - server.logger.Error(ctx, "load personal skill body failed", + builder.server.logger.Error(ctx, "load personal skill body failed", slog.F("user_id", userID), slog.F("name", name), slog.Error(err), @@ -2065,7 +2039,7 @@ func (server *Server) loadPersonalSkillBody( parsed, err := skillspkg.ParsePersonalSkillMarkdown([]byte(row.Content)) if err != nil { - server.logger.Error(ctx, "parse personal skill body failed", + builder.server.logger.Error(ctx, "parse personal skill body failed", slog.F("user_id", userID), slog.F("name", name), slog.Error(err), @@ -2075,7 +2049,7 @@ func (server *Server) loadPersonalSkillBody( return parsed, nil } -func (server *Server) appendRootChatTools( +func (builder turnEnvironmentBuilder) appendRootChatTools( ctx context.Context, tools []fantasy.AgentTool, opts rootChatToolsOptions, @@ -2084,41 +2058,41 @@ func (server *Server) appendRootChatTools( opts.workspaceCtx.selectWorkspace(updatedChat) // Notify the frontend immediately so it can start streaming // build logs before the tool completes. - server.publishChatPubsubEvent(updatedChat, codersdk.ChatWatchEventKindStatusChange, nil) + builder.server.publishChatPubsubEvent(updatedChat, codersdk.ChatWatchEventKindStatusChange, nil) } tools = append(tools, - chattool.ListTemplates(server.db, opts.chat.OrganizationID, chattool.ListTemplatesOptions{ + chattool.ListTemplates(builder.server.db, opts.chat.OrganizationID, chattool.ListTemplatesOptions{ OwnerID: opts.chat.OwnerID, - Logger: server.logger, - Clock: server.clock, + Logger: builder.server.logger, + Clock: builder.server.clock, }), - chattool.ReadTemplate(server.db, opts.chat.OrganizationID, chattool.ReadTemplateOptions{ + chattool.ReadTemplate(builder.server.db, opts.chat.OrganizationID, chattool.ReadTemplateOptions{ OwnerID: opts.chat.OwnerID, }), - chattool.CreateWorkspace(server.db, opts.chat.OrganizationID, opts.chat.ID, chattool.CreateWorkspaceOptions{ + chattool.CreateWorkspace(builder.server.db, opts.chat.OrganizationID, opts.chat.ID, chattool.CreateWorkspaceOptions{ OwnerID: opts.chat.OwnerID, - CreateFn: server.createWorkspaceFn, - AgentConnFn: chattool.AgentConnFunc(server.agentConnFn), - AgentInactiveDisconnectTimeout: server.agentInactiveDisconnectTimeout, + CreateFn: builder.server.createWorkspaceFn, + AgentConnFn: chattool.AgentConnFunc(builder.server.agentConnFn), + AgentInactiveDisconnectTimeout: builder.server.agentInactiveDisconnectTimeout, WorkspaceMu: opts.workspaceMu, OnChatUpdated: onChatUpdated, - Logger: server.logger, + Logger: builder.server.logger, }), - chattool.StartWorkspace(server.db, opts.chat.ID, chattool.StartWorkspaceOptions{ + chattool.StartWorkspace(builder.server.db, opts.chat.ID, chattool.StartWorkspaceOptions{ OwnerID: opts.chat.OwnerID, - StartFn: server.startWorkspaceFn, - AgentConnFn: chattool.AgentConnFunc(server.agentConnFn), + StartFn: builder.server.startWorkspaceFn, + AgentConnFn: chattool.AgentConnFunc(builder.server.agentConnFn), WorkspaceMu: opts.workspaceMu, OnChatUpdated: onChatUpdated, - Logger: server.logger, + Logger: builder.server.logger, }), - chattool.StopWorkspace(server.db, opts.chat.ID, chattool.StopWorkspaceOptions{ + chattool.StopWorkspace(builder.server.db, opts.chat.ID, chattool.StopWorkspaceOptions{ OwnerID: opts.chat.OwnerID, - StopFn: server.stopWorkspaceFn, + StopFn: builder.server.stopWorkspaceFn, WorkspaceMu: opts.workspaceMu, OnChatUpdated: onChatUpdated, - Logger: server.logger, + Logger: builder.server.logger, }), ) if opts.isPlanModeTurn { @@ -2130,7 +2104,7 @@ func (server *Server) appendRootChatTools( })) } - return append(tools, server.subagentTools(ctx, func() database.Chat { + return append(tools, builder.server.subagentTools(ctx, func() database.Chat { return opts.chat }, opts.modelConfigID)...) } diff --git a/coderd/x/chatd/generation_preparer_internal_test.go b/coderd/x/chatd/turn_environment_internal_test.go similarity index 95% rename from coderd/x/chatd/generation_preparer_internal_test.go rename to coderd/x/chatd/turn_environment_internal_test.go index a90abacef08..860eb9a7ca4 100644 --- a/coderd/x/chatd/generation_preparer_internal_test.go +++ b/coderd/x/chatd/turn_environment_internal_test.go @@ -154,7 +154,7 @@ func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) { chatprovider.ProviderAPIKeys{}, withInternalTestServerTransportFactory(&aibridgeTestFactory{}), ) - prepared, err := server.buildTurnEnvironment(ctx, generationPrepareInput{ + prepared, err := buildTurnEnvironment(ctx, server, generationPrepareInput{ Chat: created.Chat, Messages: created.InitialMessages, }) @@ -259,7 +259,7 @@ func TestPrepareGenerationComputerUseIgnoresChatTransportOverride(t *testing.T) chatprovider.ProviderAPIKeys{}, withInternalTestServerTransportFactory(&aibridgeTestFactory{}), ) - prepared, err := server.buildTurnEnvironment(ctx, generationPrepareInput{ + prepared, err := buildTurnEnvironment(ctx, server, generationPrepareInput{ Chat: created.Chat, Messages: created.InitialMessages, }) @@ -344,7 +344,7 @@ func TestPrepareGenerationSubagentUsesOwnerSyntheticAPIKey(t *testing.T) { chatprovider.ProviderAPIKeys{}, withInternalTestServerTransportFactory(&aibridgeTestFactory{}), ) - prepared, err := server.buildTurnEnvironment(ctx, generationPrepareInput{ + prepared, err := buildTurnEnvironment(ctx, server, generationPrepareInput{ Chat: created.Chat, Messages: created.InitialMessages, }) @@ -769,3 +769,40 @@ func TestEnabledMCPServerConfigsForChatOrg(t *testing.T) { require.Empty(t, configs) }) } + +func TestTurnToolsetMetadata(t *testing.T) { + t.Parallel() + + configID := uuid.New() + toolset := turnToolset{ + exclusiveToolNames: map[string]bool{"exclusive": true}, + dynamicToolNames: map[string]bool{"dynamic": true}, + builtinToolNames: map[string]bool{"builtin": true}, + allowInactiveTools: map[string]bool{"inactive": true}, + stopAfterTools: map[string]struct{}{"stop": {}}, + toolNameToConfigID: map[string]uuid.UUID{"configured": configID}, + } + + tests := []struct { + name string + got bool + }{ + {name: "exclusive", got: toolset.IsExclusive("exclusive")}, + {name: "dynamic", got: toolset.IsDynamic("dynamic")}, + {name: "builtin", got: toolset.IsBuiltin("builtin")}, + {name: "inactive", got: toolset.AllowsInactive("inactive")}, + {name: "stop", got: toolset.StopsAfter("stop")}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.True(t, tt.got) + }) + } + + gotConfigID, ok := toolset.ConfigID("configured") + require.True(t, ok) + require.Equal(t, configID, gotConfigID) + _, ok = toolset.ConfigID("missing") + require.False(t, ok) +} From cd7871fde18ed5c5873c9e2552f7449db638167f Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Sat, 29 Aug 2026 10:57:22 +0000 Subject: [PATCH 3/4] refactor(coderd/x/chatd): simplify turn policy helpers --- coderd/x/chatd/chatd_internal_test.go | 81 ------------------- coderd/x/chatd/turn_environment.go | 66 +-------------- .../x/chatd/turn_environment_internal_test.go | 37 --------- 3 files changed, 1 insertion(+), 183 deletions(-) diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 3848bc48211..a0c4a765225 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -810,40 +810,6 @@ func TestAllowedExploreToolNames(t *testing.T) { require.NotContains(t, got, chattool.FindToolsName) } -func TestAllowedBehaviorToolNames(t *testing.T) { - t.Parallel() - - makeTools := func(names ...string) []fantasy.AgentTool { - tools := make([]fantasy.AgentTool, 0, len(names)) - for _, name := range names { - tools = append(tools, newTestAgentTool(name)) - } - return tools - } - - allTools := makeTools("read_file", "custom_tool", "spawn_agent") - exploreMode := database.NullChatMode{ - ChatMode: database.ChatModeExplore, - Valid: true, - } - - t.Run("DefaultModeReturnsAllTools", func(t *testing.T) { - t.Parallel() - require.Equal(t, []string{"read_file", "custom_tool", "spawn_agent"}, allowedBehaviorToolNames( - allTools, - database.NullChatMode{}, - )) - }) - - t.Run("ExploreModeUsesExploreAllowlist", func(t *testing.T) { - t.Parallel() - require.Equal(t, []string{"read_file"}, allowedBehaviorToolNames( - allTools, - exploreMode, - )) - }) -} - func TestStopAfterPlanTools(t *testing.T) { t.Parallel() @@ -1819,53 +1785,6 @@ func TestPersonalAndWorkspaceSkillCollisionInSystemPrompt(t *testing.T) { require.ErrorContains(t, err, "workspace/deploy") } -func TestSkillIndexRefreshReplacesStaleAliases(t *testing.T) { - t.Parallel() - - initialResolved := mergeTurnSkills( - []skillspkg.Skill{{ - Name: "deploy", - Description: "Personal deployment process", - Source: skillspkg.SourcePersonal, - }}, - nil, - ) - prompt := buildSystemPrompt( - []fantasy.Message{{ - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "Create a workspace."}, - }, - }}, - "", - "", - initialResolved, - "", - systemPromptBehaviorContext{}, - ) - - mergedIndex := chattool.FormatResolvedSkillIndex(mergeTurnSkills( - []skillspkg.Skill{{ - Name: "deploy", - Description: "Personal deployment process", - Source: skillspkg.SourcePersonal, - }}, - []chattool.SkillMeta{{ - Name: "deploy", - Description: "Workspace deployment process", - Dir: "/skills/deploy", - }}, - )) - prompt = removeSkillIndexMessages(prompt) - prompt = chatprompt.InsertSystem(prompt, mergedIndex) - - text := systemPromptText(t, prompt) - require.Equal(t, 1, strings.Count(text, "")) - require.NotContains(t, text, "\n- deploy: Personal deployment process") - require.Contains(t, text, "- personal/deploy: Personal deployment process") - require.Contains(t, text, "- workspace/deploy: Workspace deployment process") -} - func requireUserSkillContextActor(ctx context.Context, t *testing.T, userID uuid.UUID) { t.Helper() actor, ok := dbauthz.ActorFromContext(ctx) diff --git a/coderd/x/chatd/turn_environment.go b/coderd/x/chatd/turn_environment.go index ed5de674711..4314dbc5358 100644 --- a/coderd/x/chatd/turn_environment.go +++ b/coderd/x/chatd/turn_environment.go @@ -73,15 +73,7 @@ type turnToolset struct { toolNameToConfigID map[string]uuid.UUID } -func (t turnToolset) IsExclusive(name string) bool { return t.exclusiveToolNames[name] } -func (t turnToolset) IsDynamic(name string) bool { return t.dynamicToolNames[name] } -func (t turnToolset) IsBuiltin(name string) bool { return t.builtinToolNames[name] } -func (t turnToolset) AllowsInactive(name string) bool { return t.allowInactiveTools[name] } -func (t turnToolset) StopsAfter(name string) bool { _, ok := t.stopAfterTools[name]; return ok } -func (t turnToolset) ConfigID(name string) (uuid.UUID, bool) { - id, ok := t.toolNameToConfigID[name] - return id, ok -} +func (t turnToolset) IsBuiltin(name string) bool { return t.builtinToolNames[name] } type turnEnvironmentState struct { turn turnState @@ -1627,14 +1619,6 @@ func (c *turnWorkspaceContext) getWorkspaceConn(ctx context.Context) (workspaces return nil, xerrors.New("chat workspace changed while connecting") } -func allToolNames(allTools []fantasy.AgentTool) []string { - toolNames := make([]string, 0, len(allTools)) - for _, tool := range allTools { - toolNames = append(toolNames, tool.Info().Name) - } - return toolNames -} - func isExploreSubagentMode(mode database.NullChatMode) bool { return mode.Valid && mode.ChatMode == database.ChatModeExplore } @@ -1787,19 +1771,6 @@ func allowedExploreToolNames(allTools []fantasy.AgentTool) []string { return toolNames } -// allowedBehaviorToolNames runs only on non-plan turns because -// appendDynamicTools returns early for plan mode. Within that boundary, -// Explore mode wins over the default behavior that allows all tools. -func allowedBehaviorToolNames( - allTools []fantasy.AgentTool, - chatMode database.NullChatMode, -) []string { - if isExploreSubagentMode(chatMode) { - return allowedExploreToolNames(allTools) - } - return allToolNames(allTools) -} - func stopAfterPlanTools( planMode database.NullChatPlanMode, parentChatID uuid.NullUUID, @@ -1900,34 +1871,6 @@ func buildSystemPrompt( return prompt } -func removeSkillIndexMessages(prompt []fantasy.Message) []fantasy.Message { - out := make([]fantasy.Message, 0, len(prompt)) - removed := false - for _, message := range prompt { - if isSkillIndexMessage(message) { - removed = true - continue - } - out = append(out, message) - } - if !removed { - return prompt - } - return out -} - -func isSkillIndexMessage(message fantasy.Message) bool { - if message.Role != fantasy.MessageRoleSystem || len(message.Content) != 1 { - return false - } - textPart, ok := fantasy.AsMessagePart[fantasy.TextPart](message.Content[0]) - if !ok { - return false - } - text := strings.TrimSpace(textPart.Text) - return strings.HasPrefix(text, chattool.AvailableSkillsOpenTag+"\n") && strings.HasSuffix(text, chattool.AvailableSkillsCloseTag) -} - type turnEnvironmentBuilder struct { server *Server } @@ -2136,15 +2079,8 @@ func appendDynamicTools( } } - activeToolNames := make(map[string]struct{}, len(tools)) - for _, name := range allowedBehaviorToolNames(tools, chatMode) { - activeToolNames[name] = struct{}{} - } for _, t := range tools { info := t.Info() - if _, active := activeToolNames[info.Name]; !active { - continue - } if dynamicToolNames[info.Name] { logger.Warn(ctx, "dynamic tool name collides with built-in tool, built-in takes precedence", slog.F("tool_name", info.Name)) diff --git a/coderd/x/chatd/turn_environment_internal_test.go b/coderd/x/chatd/turn_environment_internal_test.go index 860eb9a7ca4..61817f306b0 100644 --- a/coderd/x/chatd/turn_environment_internal_test.go +++ b/coderd/x/chatd/turn_environment_internal_test.go @@ -769,40 +769,3 @@ func TestEnabledMCPServerConfigsForChatOrg(t *testing.T) { require.Empty(t, configs) }) } - -func TestTurnToolsetMetadata(t *testing.T) { - t.Parallel() - - configID := uuid.New() - toolset := turnToolset{ - exclusiveToolNames: map[string]bool{"exclusive": true}, - dynamicToolNames: map[string]bool{"dynamic": true}, - builtinToolNames: map[string]bool{"builtin": true}, - allowInactiveTools: map[string]bool{"inactive": true}, - stopAfterTools: map[string]struct{}{"stop": {}}, - toolNameToConfigID: map[string]uuid.UUID{"configured": configID}, - } - - tests := []struct { - name string - got bool - }{ - {name: "exclusive", got: toolset.IsExclusive("exclusive")}, - {name: "dynamic", got: toolset.IsDynamic("dynamic")}, - {name: "builtin", got: toolset.IsBuiltin("builtin")}, - {name: "inactive", got: toolset.AllowsInactive("inactive")}, - {name: "stop", got: toolset.StopsAfter("stop")}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - require.True(t, tt.got) - }) - } - - gotConfigID, ok := toolset.ConfigID("configured") - require.True(t, ok) - require.Equal(t, configID, gotConfigID) - _, ok = toolset.ConfigID("missing") - require.False(t, ok) -} From 1add6c35c62c17193bc2de5287a2f71990b7e910 Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Sat, 29 Aug 2026 12:32:58 +0000 Subject: [PATCH 4/4] refactor(coderd/x/chatd): remove turn environment builder wrapper --- coderd/x/chatd/chatd_internal_test.go | 18 +- coderd/x/chatd/context_prompt.go | 2 +- .../x/chatd/context_prompt_internal_test.go | 2 +- coderd/x/chatd/generation.go | 4 +- coderd/x/chatd/turn_environment.go | 74 +++---- .../x/chatd/turn_environment_internal_test.go | 197 ++++++++++-------- 6 files changed, 155 insertions(+), 142 deletions(-) diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index a0c4a765225..34c3f5d7934 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -1803,7 +1803,6 @@ func TestFetchPersonalSkillMetadata(t *testing.T) { ctrl := gomock.NewController(t) db := dbmock.NewMockStore(ctrl) logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug) - server := &Server{db: db} userID := uuid.New() db.EXPECT().ListUserSkillMetadataByUserID(gomock.Any(), userID).DoAndReturn( @@ -1818,7 +1817,7 @@ func TestFetchPersonalSkillMetadata(t *testing.T) { }, ) - got := (turnEnvironmentBuilder{server: server}).fetchPersonalSkillMetadata(context.Background(), userID, logger) + got := fetchPersonalSkillMetadata(context.Background(), db, userID, logger) require.Equal(t, []skillspkg.Skill{{ Name: "personal-review", Description: "Personal review process", @@ -1833,12 +1832,11 @@ func TestFetchPersonalSkillMetadata(t *testing.T) { db := dbmock.NewMockStore(ctrl) sink := testutil.NewFakeSink(t) logger := sink.Logger().Leveled(slog.LevelDebug) - server := &Server{db: db} userID := uuid.New() db.EXPECT().ListUserSkillMetadataByUserID(gomock.Any(), userID).Return(nil, xerrors.New("boom")) - got := (turnEnvironmentBuilder{server: server}).fetchPersonalSkillMetadata(context.Background(), userID, logger) + got := fetchPersonalSkillMetadata(context.Background(), db, userID, logger) require.Empty(t, got) warns := sink.Entries(func(e slog.SinkEntry) bool { return e.Level == slog.LevelWarn && strings.Contains(e.Message, "personal skill metadata") @@ -1855,7 +1853,6 @@ func TestLoadPersonalSkillBody(t *testing.T) { ctrl := gomock.NewController(t) db := dbmock.NewMockStore(ctrl) - server := &Server{db: db} userID := uuid.New() params := database.GetUserSkillByUserIDAndNameParams{ UserID: userID, @@ -1874,7 +1871,7 @@ func TestLoadPersonalSkillBody(t *testing.T) { }, ) - got, err := (turnEnvironmentBuilder{server: server}).loadPersonalSkillBody(context.Background(), userID, "personal-review") + got, err := loadPersonalSkillBody(context.Background(), db, slogtest.Make(t, nil), userID, "personal-review") require.NoError(t, err) require.Equal(t, "personal-review", got.Name) require.Equal(t, "Personal review process", got.Description) @@ -1887,7 +1884,6 @@ func TestLoadPersonalSkillBody(t *testing.T) { ctrl := gomock.NewController(t) db := dbmock.NewMockStore(ctrl) - server := &Server{db: db} userID := uuid.New() params := database.GetUserSkillByUserIDAndNameParams{ UserID: userID, @@ -1902,7 +1898,7 @@ func TestLoadPersonalSkillBody(t *testing.T) { }, ) - _, err := (turnEnvironmentBuilder{server: server}).loadPersonalSkillBody(context.Background(), userID, "missing-skill") + _, err := loadPersonalSkillBody(context.Background(), db, slogtest.Make(t, nil), userID, "missing-skill") require.ErrorIs(t, err, skillspkg.ErrSkillNotFound) }) @@ -1912,7 +1908,6 @@ func TestLoadPersonalSkillBody(t *testing.T) { ctrl := gomock.NewController(t) db := dbmock.NewMockStore(ctrl) sink := testutil.NewFakeSink(t) - server := &Server{db: db, logger: sink.Logger()} userID := uuid.New() params := database.GetUserSkillByUserIDAndNameParams{ UserID: userID, @@ -1928,7 +1923,7 @@ func TestLoadPersonalSkillBody(t *testing.T) { }, ) - _, err := (turnEnvironmentBuilder{server: server}).loadPersonalSkillBody(context.Background(), userID, "error-skill") + _, err := loadPersonalSkillBody(context.Background(), db, sink.Logger(), userID, "error-skill") require.ErrorContains(t, err, "load personal skill body") require.ErrorIs(t, err, dbErr) @@ -1945,7 +1940,6 @@ func TestLoadPersonalSkillBody(t *testing.T) { ctrl := gomock.NewController(t) db := dbmock.NewMockStore(ctrl) sink := testutil.NewFakeSink(t) - server := &Server{db: db, logger: sink.Logger()} userID := uuid.New() params := database.GetUserSkillByUserIDAndNameParams{ UserID: userID, @@ -1964,7 +1958,7 @@ func TestLoadPersonalSkillBody(t *testing.T) { }, ) - _, err := (turnEnvironmentBuilder{server: server}).loadPersonalSkillBody(context.Background(), userID, "broken-skill") + _, err := loadPersonalSkillBody(context.Background(), db, sink.Logger(), userID, "broken-skill") require.ErrorContains(t, err, "parse personal skill body") require.ErrorIs(t, err, skillspkg.ErrSkillBodyRequired) diff --git a/coderd/x/chatd/context_prompt.go b/coderd/x/chatd/context_prompt.go index 0646109e544..14d2f73d750 100644 --- a/coderd/x/chatd/context_prompt.go +++ b/coderd/x/chatd/context_prompt.go @@ -173,7 +173,7 @@ func decodeSkillIdentity(body json.RawMessage) (name, description string, decode // workspace skills from the chat's pinned context resources // (chat_context_resources), populated at hydrate and refresh time. A chat // with no pinned rows yields no context. A read error is returned rather than -// swallowed, matching the other prompt-input reads in prepareGeneration. +// swallowed, matching the other prompt-input reads in buildTurnEnvironment. // // agent only decorates the instruction header with its OS and directory; an // unresolved (zero-value) agent does not blank the context, so the pin keeps diff --git a/coderd/x/chatd/context_prompt_internal_test.go b/coderd/x/chatd/context_prompt_internal_test.go index 96738f67831..fba425e5f54 100644 --- a/coderd/x/chatd/context_prompt_internal_test.go +++ b/coderd/x/chatd/context_prompt_internal_test.go @@ -438,7 +438,7 @@ func TestPinnedWorkspaceContextFromHydratedPin(t *testing.T) { require.Empty(t, emptySkills) } -// TestResolveTurnWorkspaceContext covers the dispatch that prepareGeneration +// TestResolveTurnWorkspaceContext covers the dispatch that buildTurnEnvironment // wires up: the pinned copy when the chat has pinned rows, and nothing for a // non-workspace chat or a chat without pinned rows. func TestResolveTurnWorkspaceContext(t *testing.T) { diff --git a/coderd/x/chatd/generation.go b/coderd/x/chatd/generation.go index 435f833731c..2e7900fec09 100644 --- a/coderd/x/chatd/generation.go +++ b/coderd/x/chatd/generation.go @@ -622,9 +622,9 @@ func (s *taskStarter) waitGenerationRetry(ctx context.Context, delay time.Durati } const ( - // generationPhaseMaxAttempts bounds how many times prepareGeneration + // generationPhaseMaxAttempts bounds how many times buildTurnEnvironment // and decideGenerationAction run before the turn finishes with an - // error. Both phases are retried because prepareGeneration performs + // error. Both phases are retried because buildTurnEnvironment performs // I/O (DB reads, MCP connects, workspace dials) that can fail // transiently. generationPhaseMaxAttempts = 3 diff --git a/coderd/x/chatd/turn_environment.go b/coderd/x/chatd/turn_environment.go index 4314dbc5358..422f7f0932b 100644 --- a/coderd/x/chatd/turn_environment.go +++ b/coderd/x/chatd/turn_environment.go @@ -73,7 +73,7 @@ type turnToolset struct { toolNameToConfigID map[string]uuid.UUID } -func (t turnToolset) IsBuiltin(name string) bool { return t.builtinToolNames[name] } +func (t *turnToolset) IsBuiltin(name string) bool { return t.builtinToolNames[name] } type turnEnvironmentState struct { turn turnState @@ -142,7 +142,6 @@ func buildTurnEnvironment( server *Server, input generationPrepareInput, ) (turnEnvironment, error) { - builder := turnEnvironmentBuilder{server: server} chat := input.Chat logger := server.logger.With( slog.F("chat_id", chat.ID), @@ -254,7 +253,7 @@ func buildTurnEnvironment( approvedPlanMCPConfigIDs = map[uuid.UUID]struct{}{} } - planModeInstructions := builder.loadPlanModeInstructions(ctx, currentPlanMode, logger) + planModeInstructions := loadPlanModeInstructions(ctx, server.db, currentPlanMode, logger) advisorCfg := server.loadAdvisorConfig(ctx, logger) // Force Enabled from the experiment; the stored DB value is ignored. advisorCfg.Enabled = server.experiments.Enabled(codersdk.ExperimentChatAdvisor) @@ -427,7 +426,7 @@ func buildTurnEnvironment( return nil }) g2.Go(func() error { - personalSkills = builder.fetchPersonalSkillMetadata(ctx, chat.OwnerID, logger) + personalSkills = fetchPersonalSkillMetadata(ctx, server.db, chat.OwnerID, logger) return nil }) g2.Go(func() error { @@ -561,7 +560,7 @@ func buildTurnEnvironment( tools = append(tools, chattool.NewAskUserQuestionTool()) } if isRootChat { - tools = builder.appendRootChatTools(ctx, tools, rootChatToolsOptions{ + tools = appendRootChatTools(ctx, server, tools, rootChatToolsOptions{ chat: chat, modelConfigID: modelConfig.ID, workspaceCtx: &workspaceCtx, @@ -579,7 +578,7 @@ func buildTurnEnvironment( }, ResolveAlias: resolveSkillAlias, LoadPersonalSkillBody: func(ctx context.Context, name string) (skillspkg.ParsedSkill, error) { - return builder.loadPersonalSkillBody(ctx, chat.OwnerID, name) + return loadPersonalSkillBody(ctx, server.db, server.logger, chat.OwnerID, name) }, } appendCurrentSkillTools := func(current []fantasy.AgentTool) ([]fantasy.AgentTool, bool) { @@ -1871,10 +1870,6 @@ func buildSystemPrompt( return prompt } -type turnEnvironmentBuilder struct { - server *Server -} - type rootChatToolsOptions struct { chat database.Chat modelConfigID uuid.UUID @@ -1885,8 +1880,9 @@ type rootChatToolsOptions struct { isPlanModeTurn bool } -func (builder turnEnvironmentBuilder) loadPlanModeInstructions( +func loadPlanModeInstructions( ctx context.Context, + db database.Store, mode database.NullChatPlanMode, logger slog.Logger, ) string { @@ -1898,7 +1894,7 @@ func (builder turnEnvironmentBuilder) loadPlanModeInstructions( // not carry a deployment-config actor during background execution. //nolint:gocritic // Required to read deployment config during background chat processing. systemCtx := dbauthz.AsSystemRestricted(ctx) - fetched, err := builder.server.db.GetChatPlanModeInstructions(systemCtx) + fetched, err := db.GetChatPlanModeInstructions(systemCtx) if err != nil { logger.Warn(ctx, "failed to fetch plan mode instructions", @@ -1928,12 +1924,13 @@ func userSkillContext(ctx context.Context, userID uuid.UUID) context.Context { return dbauthz.As(ctx, actor) } -func (builder turnEnvironmentBuilder) fetchPersonalSkillMetadata( +func fetchPersonalSkillMetadata( ctx context.Context, + db database.Store, userID uuid.UUID, logger slog.Logger, ) []skillspkg.Skill { - rows, err := builder.server.db.ListUserSkillMetadataByUserID(userSkillContext(ctx, userID), userID) + rows, err := db.ListUserSkillMetadataByUserID(userSkillContext(ctx, userID), userID) // See package coderd/x/skills (doc.go) for why metadata fetch failures // intentionally degrade to an empty personal-skill list instead of // failing the chat turn. @@ -1956,12 +1953,14 @@ func (builder turnEnvironmentBuilder) fetchPersonalSkillMetadata( return personalSkills } -func (builder turnEnvironmentBuilder) loadPersonalSkillBody( +func loadPersonalSkillBody( ctx context.Context, + db database.Store, + logger slog.Logger, userID uuid.UUID, name string, ) (skillspkg.ParsedSkill, error) { - row, err := builder.server.db.GetUserSkillByUserIDAndName( + row, err := db.GetUserSkillByUserIDAndName( userSkillContext(ctx, userID), database.GetUserSkillByUserIDAndNameParams{ UserID: userID, @@ -1972,7 +1971,7 @@ func (builder turnEnvironmentBuilder) loadPersonalSkillBody( if errors.Is(err, sql.ErrNoRows) { return skillspkg.ParsedSkill{}, skillspkg.ErrSkillNotFound } - builder.server.logger.Error(ctx, "load personal skill body failed", + logger.Error(ctx, "load personal skill body failed", slog.F("user_id", userID), slog.F("name", name), slog.Error(err), @@ -1982,7 +1981,7 @@ func (builder turnEnvironmentBuilder) loadPersonalSkillBody( parsed, err := skillspkg.ParsePersonalSkillMarkdown([]byte(row.Content)) if err != nil { - builder.server.logger.Error(ctx, "parse personal skill body failed", + logger.Error(ctx, "parse personal skill body failed", slog.F("user_id", userID), slog.F("name", name), slog.Error(err), @@ -1992,8 +1991,9 @@ func (builder turnEnvironmentBuilder) loadPersonalSkillBody( return parsed, nil } -func (builder turnEnvironmentBuilder) appendRootChatTools( +func appendRootChatTools( ctx context.Context, + server *Server, tools []fantasy.AgentTool, opts rootChatToolsOptions, ) []fantasy.AgentTool { @@ -2001,41 +2001,41 @@ func (builder turnEnvironmentBuilder) appendRootChatTools( opts.workspaceCtx.selectWorkspace(updatedChat) // Notify the frontend immediately so it can start streaming // build logs before the tool completes. - builder.server.publishChatPubsubEvent(updatedChat, codersdk.ChatWatchEventKindStatusChange, nil) + server.publishChatPubsubEvent(updatedChat, codersdk.ChatWatchEventKindStatusChange, nil) } tools = append(tools, - chattool.ListTemplates(builder.server.db, opts.chat.OrganizationID, chattool.ListTemplatesOptions{ + chattool.ListTemplates(server.db, opts.chat.OrganizationID, chattool.ListTemplatesOptions{ OwnerID: opts.chat.OwnerID, - Logger: builder.server.logger, - Clock: builder.server.clock, + Logger: server.logger, + Clock: server.clock, }), - chattool.ReadTemplate(builder.server.db, opts.chat.OrganizationID, chattool.ReadTemplateOptions{ + chattool.ReadTemplate(server.db, opts.chat.OrganizationID, chattool.ReadTemplateOptions{ OwnerID: opts.chat.OwnerID, }), - chattool.CreateWorkspace(builder.server.db, opts.chat.OrganizationID, opts.chat.ID, chattool.CreateWorkspaceOptions{ + chattool.CreateWorkspace(server.db, opts.chat.OrganizationID, opts.chat.ID, chattool.CreateWorkspaceOptions{ OwnerID: opts.chat.OwnerID, - CreateFn: builder.server.createWorkspaceFn, - AgentConnFn: chattool.AgentConnFunc(builder.server.agentConnFn), - AgentInactiveDisconnectTimeout: builder.server.agentInactiveDisconnectTimeout, + CreateFn: server.createWorkspaceFn, + AgentConnFn: chattool.AgentConnFunc(server.agentConnFn), + AgentInactiveDisconnectTimeout: server.agentInactiveDisconnectTimeout, WorkspaceMu: opts.workspaceMu, OnChatUpdated: onChatUpdated, - Logger: builder.server.logger, + Logger: server.logger, }), - chattool.StartWorkspace(builder.server.db, opts.chat.ID, chattool.StartWorkspaceOptions{ + chattool.StartWorkspace(server.db, opts.chat.ID, chattool.StartWorkspaceOptions{ OwnerID: opts.chat.OwnerID, - StartFn: builder.server.startWorkspaceFn, - AgentConnFn: chattool.AgentConnFunc(builder.server.agentConnFn), + StartFn: server.startWorkspaceFn, + AgentConnFn: chattool.AgentConnFunc(server.agentConnFn), WorkspaceMu: opts.workspaceMu, OnChatUpdated: onChatUpdated, - Logger: builder.server.logger, + Logger: server.logger, }), - chattool.StopWorkspace(builder.server.db, opts.chat.ID, chattool.StopWorkspaceOptions{ + chattool.StopWorkspace(server.db, opts.chat.ID, chattool.StopWorkspaceOptions{ OwnerID: opts.chat.OwnerID, - StopFn: builder.server.stopWorkspaceFn, + StopFn: server.stopWorkspaceFn, WorkspaceMu: opts.workspaceMu, OnChatUpdated: onChatUpdated, - Logger: builder.server.logger, + Logger: server.logger, }), ) if opts.isPlanModeTurn { @@ -2047,7 +2047,7 @@ func (builder turnEnvironmentBuilder) appendRootChatTools( })) } - return append(tools, builder.server.subagentTools(ctx, func() database.Chat { + return append(tools, server.subagentTools(ctx, func() database.Chat { return opts.chat }, opts.modelConfigID)...) } diff --git a/coderd/x/chatd/turn_environment_internal_test.go b/coderd/x/chatd/turn_environment_internal_test.go index 61817f306b0..05c78721012 100644 --- a/coderd/x/chatd/turn_environment_internal_test.go +++ b/coderd/x/chatd/turn_environment_internal_test.go @@ -47,49 +47,54 @@ func textMessage(t *testing.T, id int64, role database.ChatMessageRole, parts .. func TestLatestAssistantText(t *testing.T) { t.Parallel() - t.Run("ReturnsMostRecentAssistantMessage", func(t *testing.T) { - t.Parallel() - messages := []database.ChatMessage{ - textMessage(t, 1, database.ChatMessageRoleUser, "hi"), - textMessage(t, 2, database.ChatMessageRoleAssistant, "first answer"), - textMessage(t, 3, database.ChatMessageRoleTool, "tool result"), - textMessage(t, 4, database.ChatMessageRoleAssistant, " final answer "), - } - require.Equal(t, "final answer", latestAssistantText(messages)) - }) - - t.Run("ConcatenatesTextParts", func(t *testing.T) { - t.Parallel() - messages := []database.ChatMessage{ - textMessage(t, 1, database.ChatMessageRoleAssistant, "foo", "bar"), - } - require.Equal(t, "foobar", latestAssistantText(messages)) - }) - - t.Run("NoAssistantMessage", func(t *testing.T) { - t.Parallel() - messages := []database.ChatMessage{ - textMessage(t, 1, database.ChatMessageRoleUser, "hi"), - textMessage(t, 2, database.ChatMessageRoleTool, "tool result"), - } - require.Empty(t, latestAssistantText(messages)) - }) - - t.Run("EmptyAssistantText", func(t *testing.T) { - t.Parallel() - messages := []database.ChatMessage{ - textMessage(t, 1, database.ChatMessageRoleAssistant, " "), - } - require.Empty(t, latestAssistantText(messages)) - }) - - t.Run("EmptyHistory", func(t *testing.T) { - t.Parallel() - require.Empty(t, latestAssistantText(nil)) - }) + tests := []struct { + name string + messages []database.ChatMessage + want string + }{ + { + name: "ReturnsMostRecentAssistantMessage", + messages: []database.ChatMessage{ + textMessage(t, 1, database.ChatMessageRoleUser, "hi"), + textMessage(t, 2, database.ChatMessageRoleAssistant, "first answer"), + textMessage(t, 3, database.ChatMessageRoleTool, "tool result"), + textMessage(t, 4, database.ChatMessageRoleAssistant, " final answer "), + }, + want: "final answer", + }, + { + name: "ConcatenatesTextParts", + messages: []database.ChatMessage{ + textMessage(t, 1, database.ChatMessageRoleAssistant, "foo", "bar"), + }, + want: "foobar", + }, + { + name: "NoAssistantMessage", + messages: []database.ChatMessage{ + textMessage(t, 1, database.ChatMessageRoleUser, "hi"), + textMessage(t, 2, database.ChatMessageRoleTool, "tool result"), + }, + }, + { + name: "EmptyAssistantText", + messages: []database.ChatMessage{ + textMessage(t, 1, database.ChatMessageRoleAssistant, " "), + }, + }, + { + name: "EmptyHistory", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.Equal(t, tt.want, latestAssistantText(tt.messages)) + }) + } } -func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) { +func TestBuildTurnEnvironmentClampsRequestedReasoningEffortToMax(t *testing.T) { t.Parallel() db, ps := dbtestutil.NewDB(t) @@ -181,7 +186,7 @@ func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) { require.Nil(t, summaryCall.MaxOutputTokens) } -func TestPrepareGenerationComputerUseIgnoresChatTransportOverride(t *testing.T) { +func TestBuildTurnEnvironmentComputerUseIgnoresChatTransportOverride(t *testing.T) { t.Parallel() db, ps := dbtestutil.NewDB(t) @@ -290,7 +295,7 @@ func TestPrepareGenerationComputerUseIgnoresChatTransportOverride(t *testing.T) require.True(t, sawInlinedText, "attachment was not inlined as text") } -func TestPrepareGenerationSubagentUsesOwnerSyntheticAPIKey(t *testing.T) { +func TestBuildTurnEnvironmentSubagentUsesOwnerSyntheticAPIKey(t *testing.T) { t.Parallel() db, ps := dbtestutil.NewDB(t) @@ -361,7 +366,7 @@ func TestPrepareGenerationSubagentUsesOwnerSyntheticAPIKey(t *testing.T) { // TestDeriveFinalTurnRunResult exercises the re-derivation path that replaces // the old in-memory generationSideEffects stash. The server here never ran -// prepareGeneration, so a passing test proves the finish-turn inputs are +// buildTurnEnvironment, so a passing test proves the finish-turn inputs are // rebuilt purely from persisted state. func TestDeriveFinalTurnRunResult(t *testing.T) { t.Parallel() @@ -596,51 +601,65 @@ func TestShouldCompactPromptUsage(t *testing.T) { const contextLimit = int64(262144) // 256K, as in the poolside report - t.Run("inflated cumulative usage triggers compaction", func(t *testing.T) { - t.Parallel() - // 417,012 tokens: what the aibridge cross-chunk sum bug - // produced for a ~6,000-token conversation. - assert.True(t, shouldCompactPromptUsage( - fantasy.Usage{InputTokens: 417012, TotalTokens: 418846}, - contextLimit, 80)) - }) - - t.Run("correct per-step usage does not trigger", func(t *testing.T) { - t.Parallel() - assert.False(t, shouldCompactPromptUsage( - fantasy.Usage{InputTokens: 6000, TotalTokens: 6030}, - contextLimit, 80)) - }) - - t.Run("threshold 100 disables compaction", func(t *testing.T) { - t.Parallel() - assert.False(t, shouldCompactPromptUsage( - fantasy.Usage{InputTokens: 500000}, contextLimit, 100)) - }) - - t.Run("zero context limit disables compaction", func(t *testing.T) { - t.Parallel() - assert.False(t, shouldCompactPromptUsage( - fantasy.Usage{InputTokens: 6000}, 0, 80)) - }) - - t.Run("counts cache read and creation tokens", func(t *testing.T) { - t.Parallel() - usage := fantasy.Usage{ - InputTokens: 6000, - CacheReadTokens: 200000, - CacheCreationTokens: 5000, - } - // 211,000 / 262,144 = ~80.5% - assert.True(t, shouldCompactPromptUsage(usage, contextLimit, 80)) - }) - - t.Run("falls back to TotalTokens when granular fields are missing", func(t *testing.T) { - t.Parallel() - assert.True(t, shouldCompactPromptUsage( - fantasy.Usage{TotalTokens: 211000}, - contextLimit, 80)) - }) + tests := []struct { + name string + usage fantasy.Usage + limit int64 + threshold int32 + want bool + }{ + { + // 417,012 tokens: what the aibridge cross-chunk sum bug + // produced for a ~6,000-token conversation. + name: "inflated cumulative usage triggers compaction", + usage: fantasy.Usage{InputTokens: 417012, TotalTokens: 418846}, + limit: contextLimit, + threshold: 80, + want: true, + }, + { + name: "correct per-step usage does not trigger", + usage: fantasy.Usage{InputTokens: 6000, TotalTokens: 6030}, + limit: contextLimit, + threshold: 80, + }, + { + name: "threshold 100 disables compaction", + usage: fantasy.Usage{InputTokens: 500000}, + limit: contextLimit, + threshold: 100, + }, + { + name: "zero context limit disables compaction", + usage: fantasy.Usage{InputTokens: 6000}, + threshold: 80, + }, + { + // 211,000 / 262,144 = ~80.5% + name: "counts cache read and creation tokens", + usage: fantasy.Usage{ + InputTokens: 6000, + CacheReadTokens: 200000, + CacheCreationTokens: 5000, + }, + limit: contextLimit, + threshold: 80, + want: true, + }, + { + name: "falls back to TotalTokens when granular fields are missing", + usage: fantasy.Usage{TotalTokens: 211000}, + limit: contextLimit, + threshold: 80, + want: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + assert.Equal(t, tt.want, shouldCompactPromptUsage(tt.usage, tt.limit, tt.threshold)) + }) + } } func TestEnabledMCPServerConfigsForChatOrg(t *testing.T) {