diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 90043de2fe..d920e66174 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -4730,8 +4730,10 @@ func (p *Server) runChat( } } + var instructionInjected bool if instruction != "" { prompt = chatprompt.InsertSystem(prompt, instruction) + instructionInjected = true } prompt = renderPlanPathPrompt(prompt, resolvePlanPathBlock(ctx)) if skillIndex := chattool.FormatSkillIndex(skills); skillIndex != "" { @@ -5077,6 +5079,33 @@ func (p *Server) runChat( // start streaming build logs before the tool // completes. p.publishChatPubsubEvent(updatedChat, codersdk.ChatWatchEventKindStatusChange, nil) + + // When a workspace is first attached mid-turn + // (e.g. via create_workspace), fetch and persist + // instruction files immediately so the LLM has + // AGENTS.md context for the remainder of this + // turn. The persisted marker prevents redundant + // fetches on subsequent turns. + if instruction == "" && updatedChat.WorkspaceID.Valid { + newInstruction, discoveredSkills, persistErr := p.persistInstructionFiles( + ctx, + updatedChat, + modelConfig.ID, + workspaceCtx.getWorkspaceAgent, + workspaceCtx.getWorkspaceConn, + ) + if persistErr != nil { + p.logger.Warn(ctx, "failed to persist instruction files on workspace attach", + slog.F("chat_id", updatedChat.ID), + slog.Error(persistErr), + ) + } else { + instruction = newInstruction + if len(discoveredSkills) > 0 { + skills = discoveredSkills + } + } + } } tools = append(tools, chattool.ListTemplates(chat.OrganizationID, p.db, chattool.ListTemplatesOptions{ @@ -5311,11 +5340,27 @@ func (p *Server) runChat( if chat.ParentChatID.Valid { reloadedPrompt = chatprompt.InsertSystem(reloadedPrompt, defaultSubagentInstruction) } - if instruction != "" { - reloadedPrompt = chatprompt.InsertSystem(reloadedPrompt, instruction) + // Re-derive instruction and skills from the reloaded + // messages so that any context added during the + // chatloop (e.g. via persistInstructionFiles when + // the agent changes) is picked up after compaction. + // The captured instruction takes priority; fall + // back to persisted DB content otherwise. + reloadedInstruction := instruction + if reloadedInstruction == "" { + reloadedInstruction = instructionFromContextFiles(reloadedMsgs) + } + if reloadedInstruction != "" { + reloadedPrompt = chatprompt.InsertSystem(reloadedPrompt, reloadedInstruction) + instructionInjected = true } reloadedPrompt = renderPlanPathPrompt(reloadedPrompt, resolvePlanPathBlock(reloadCtx)) - if skillIndex := chattool.FormatSkillIndex(skills); skillIndex != "" { + reloadedSkills := skillsFromParts(reloadedMsgs) + if len(reloadedSkills) == 0 { + reloadedSkills = skills + } + + if skillIndex := chattool.FormatSkillIndex(reloadedSkills); skillIndex != "" { reloadedPrompt = chatprompt.InsertSystem(reloadedPrompt, skillIndex) } reloadUserPrompt := p.resolveUserPrompt(reloadCtx, chat.OwnerID) @@ -5333,7 +5378,17 @@ func (p *Server) runChat( DisableChainMode: func() { chainModeActive = false }, - + PrepareMessages: func(msgs []fantasy.Message) []fantasy.Message { + if instructionInjected || instruction == "" { + return nil + } + instructionInjected = true + result := chatprompt.InsertSystem(msgs, instruction) + if skillIndex := chattool.FormatSkillIndex(skills); skillIndex != "" { + result = chatprompt.InsertSystem(result, skillIndex) + } + return result + }, OnRetry: func( attempt int, retryErr error, @@ -5726,39 +5781,37 @@ func contextFileAgentID(messages []database.ChatMessage) (uuid.UUID, bool) { return lastID, found } -// persistInstructionFiles reads instruction files and discovers -// skills from the workspace agent, persisting both as message -// parts. This is called once when a workspace is first attached -// to a chat (or when the agent changes). Returns the formatted -// instruction string and skill index for injection into the -// current turn's prompt. -func (p *Server) persistInstructionFiles( +// fetchWorkspaceContext retrieves fresh instruction files and +// skills from the workspace agent without persisting. It handles +// agent connection, context configuration fetching, content +// sanitization, and metadata stamping. Returns the workspace +// agent, the stamped parts, discovered skills, and whether the +// workspace connection succeeded. A nil agent means the chat has +// no valid workspace or the agent lookup failed; +// workspaceConnOK is false in that case. +func (p *Server) fetchWorkspaceContext( ctx context.Context, chat database.Chat, - modelConfigID uuid.UUID, getWorkspaceAgent func(context.Context) (database.WorkspaceAgent, error), getWorkspaceConn func(context.Context) (workspacesdk.AgentConn, error), -) (instruction string, skills []chattool.SkillMeta, err error) { +) (agent *database.WorkspaceAgent, agentParts []codersdk.ChatMessagePart, discoveredSkills []chattool.SkillMeta, workspaceConnOK bool) { if !chat.WorkspaceID.Valid || getWorkspaceAgent == nil { - return "", nil, nil + return nil, nil, nil, false } - agent, err := getWorkspaceAgent(ctx) - if err != nil { - return "", nil, nil + loadedAgent, agentErr := getWorkspaceAgent(ctx) + if agentErr != nil { + return nil, nil, nil, false } - directory := agent.ExpandedDirectory + directory := loadedAgent.ExpandedDirectory if directory == "" { - directory = agent.Directory + directory = loadedAgent.Directory } // Fetch context configuration from the agent. Parts // arrive pre-populated with context-file and skill entries // so we don't need additional round-trips. - var workspaceConnOK bool - var agentParts []codersdk.ChatMessagePart - if getWorkspaceConn != nil { instructionCtx, cancel := context.WithTimeout(ctx, p.instructionLookupTimeout) defer cancel() @@ -5789,21 +5842,15 @@ func (p *Server) persistInstructionFiles( // Stamp server-side fields and sanitize content. The // agent cannot know its own UUID, OS metadata, or // directory — those are added here at the trust boundary. - var discoveredSkills []chattool.SkillMeta - var hasContent, hasContextFilePart bool - agentID := uuid.NullUUID{UUID: agent.ID, Valid: true} + agentID := uuid.NullUUID{UUID: loadedAgent.ID, Valid: true} for i := range agentParts { agentParts[i].ContextFileAgentID = agentID switch agentParts[i].Type { case codersdk.ChatMessagePartTypeContextFile: - hasContextFilePart = true agentParts[i].ContextFileContent = SanitizePromptText(agentParts[i].ContextFileContent) - agentParts[i].ContextFileOS = agent.OperatingSystem + agentParts[i].ContextFileOS = loadedAgent.OperatingSystem agentParts[i].ContextFileDirectory = directory - if agentParts[i].ContextFileContent != "" { - hasContent = true - } case codersdk.ChatMessagePartTypeSkill: discoveredSkills = append(discoveredSkills, chattool.SkillMeta{ Name: agentParts[i].SkillName, @@ -5814,6 +5861,49 @@ func (p *Server) persistInstructionFiles( } } + return &loadedAgent, agentParts, discoveredSkills, workspaceConnOK +} + +// persistInstructionFiles fetches AGENTS.md instruction files and +// skills from the workspace agent, persisting both as message +// parts. This is called once when a workspace is first attached +// to a chat (or when the agent changes). Returns the formatted +// instruction string and skill index for injection into the +// current turn's prompt. +func (p *Server) persistInstructionFiles( + ctx context.Context, + chat database.Chat, + modelConfigID uuid.UUID, + getWorkspaceAgent func(context.Context) (database.WorkspaceAgent, error), + getWorkspaceConn func(context.Context) (workspacesdk.AgentConn, error), +) (instruction string, skills []chattool.SkillMeta, err error) { + agent, agentParts, discoveredSkills, workspaceConnOK := p.fetchWorkspaceContext( + ctx, chat, getWorkspaceAgent, getWorkspaceConn, + ) + // Defensive guard: fetchWorkspaceContext returns nil when the + // chat has no valid workspace or the agent lookup fails. It's + // cheaper to guard here than push the precondition up to all + // callers. + if agent == nil { + return "", nil, nil + } + + agentID := uuid.NullUUID{UUID: agent.ID, Valid: true} + hasContent := false + hasContextFilePart := false + for _, part := range agentParts { + if part.Type == codersdk.ChatMessagePartTypeContextFile { + hasContextFilePart = true + if part.ContextFileContent != "" { + hasContent = true + } + } + } + directory := agent.ExpandedDirectory + if directory == "" { + directory = agent.Directory + } + if !hasContent { if !workspaceConnOK { return "", nil, nil diff --git a/coderd/x/chatd/chatloop/chatloop.go b/coderd/x/chatd/chatloop/chatloop.go index 1bd45d5b1f..951554bbb6 100644 --- a/coderd/x/chatd/chatloop/chatloop.go +++ b/coderd/x/chatd/chatloop/chatloop.go @@ -139,6 +139,12 @@ type RunOptions struct { Compaction *CompactionOptions ReloadMessages func(context.Context) ([]fantasy.Message, error) DisableChainMode func() + // PrepareMessages is called before each LLM step with the + // current message history. If it returns non-nil, the returned + // slice replaces messages for this and all subsequent steps. + // Used to inject system context that becomes available mid-loop + // (e.g. AGENTS.md after create_workspace). + PrepareMessages func([]fantasy.Message) []fantasy.Message // OnRetry is called before each retry attempt when the LLM // stream fails with a retryable error. It provides the attempt @@ -363,6 +369,11 @@ func Run(ctx context.Context, opts RunOptions) error { // copy copies Message structs by value, so field // reassignments in addAnthropicPromptCaching only // affect the prepared slice. + if opts.PrepareMessages != nil { + if updated := opts.PrepareMessages(messages); updated != nil { + messages = updated + } + } prepared := make([]fantasy.Message, len(messages)) copy(prepared, messages) if applyAnthropicCaching { @@ -370,6 +381,7 @@ func Run(ctx context.Context, opts RunOptions) error { } opts.Metrics.MessageCount.WithLabelValues(provider).Observe(float64(len(prepared))) opts.Metrics.PromptSizeBytes.WithLabelValues(provider).Observe(float64(EstimatePromptSize(prepared))) + call := fantasy.Call{ Prompt: prepared, Tools: tools, diff --git a/coderd/x/chatd/chatloop/chatloop_test.go b/coderd/x/chatd/chatloop/chatloop_test.go index 47d08b75db..b33f8ac796 100644 --- a/coderd/x/chatd/chatloop/chatloop_test.go +++ b/coderd/x/chatd/chatloop/chatloop_test.go @@ -1552,3 +1552,205 @@ func TestRun_PersistStepInterruptedFallback(t *testing.T) { } require.True(t, foundText, "fallback should persist the text content") } + +func TestRun_PrepareMessagesInjectsSystemContextMidLoop(t *testing.T) { + t.Parallel() + + const injectedInstruction = "You are working in /home/coder/project. Follow AGENTS.md guidelines." + + var mu sync.Mutex + var streamCalls int + var secondCallPrompt []fantasy.Message + + // Step 0 calls a tool. Step 1 sees the injected system message. + model := &chattest.FakeModel{ + ProviderName: "fake", + StreamFn: func(_ context.Context, call fantasy.Call) (fantasy.StreamResponse, error) { + mu.Lock() + step := streamCalls + streamCalls++ + mu.Unlock() + + switch step { + case 0: + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeToolInputStart, ID: "tc-1", ToolCallName: "create_workspace"}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "tc-1", Delta: `{}`}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "tc-1"}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "tc-1", + ToolCallName: "create_workspace", + ToolCallInput: `{}`, + }, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls}, + }), nil + default: + mu.Lock() + secondCallPrompt = append([]fantasy.Message(nil), call.Prompt...) + mu.Unlock() + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeTextStart, ID: "text-1"}, + {Type: fantasy.StreamPartTypeTextDelta, ID: "text-1", Delta: "done"}, + {Type: fantasy.StreamPartTypeTextEnd, ID: "text-1"}, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonStop}, + }), nil + } + }, + } + + // Simulate: after the tool executes (step 0), instruction + // becomes available. PrepareMessages injects it before step 1. + instructionInjected := make(chan struct{}) + var instructionAvailable atomic.Value + // The tool sets instruction after execution. + tool := fantasy.NewAgentTool( + "create_workspace", + "create a workspace", + func(_ context.Context, _ struct{}, _ fantasy.ToolCall) (fantasy.ToolResponse, error) { + instructionAvailable.Store(injectedInstruction) + return fantasy.ToolResponse{}, nil + }, + ) + + err := Run(context.Background(), RunOptions{ + Model: model, + Messages: []fantasy.Message{ + textMessage(fantasy.MessageRoleUser, "create a workspace and open a PR"), + }, + Tools: []fantasy.AgentTool{tool}, + MaxSteps: 5, + PersistStep: func(_ context.Context, _ PersistedStep) error { + return nil + }, + PrepareMessages: func(msgs []fantasy.Message) []fantasy.Message { + select { + case <-instructionInjected: + return nil + default: + } + instr, ok := instructionAvailable.Load().(string) + if !ok || instr == "" { + return nil + } + close(instructionInjected) + // Insert a system message after existing system messages. + result := make([]fantasy.Message, 0, len(msgs)+1) + inserted := false + for i, msg := range msgs { + result = append(result, msg) + if !inserted && msg.Role == fantasy.MessageRoleSystem { + // Insert after the last system message. + if i+1 >= len(msgs) || msgs[i+1].Role != fantasy.MessageRoleSystem { + result = append(result, fantasy.Message{ + Role: fantasy.MessageRoleSystem, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: instr}, + }, + }) + inserted = true + } + } + } + if !inserted { + // No system messages — prepend. + result = append([]fantasy.Message{{ + Role: fantasy.MessageRoleSystem, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: instr}, + }, + }}, result...) + } + return result + }, + }) + require.NoError(t, err) + require.Equal(t, 2, streamCalls) + + // The second LLM call should contain the injected instruction. + require.NotEmpty(t, secondCallPrompt) + var foundInstruction bool + for _, msg := range secondCallPrompt { + if msg.Role != fantasy.MessageRoleSystem { + continue + } + for _, part := range msg.Content { + if tp, ok := fantasy.AsMessagePart[fantasy.TextPart](part); ok { + if strings.Contains(tp.Text, "AGENTS.md") { + foundInstruction = true + } + } + } + } + require.True(t, foundInstruction, + "step 1 prompt should contain the injected system instruction") +} + +func TestRun_PrepareMessagesOnlyFiresOnce(t *testing.T) { + t.Parallel() + + var mu sync.Mutex + var streamCalls int + + // Three steps: tool call, tool call, text. PrepareMessages + // should inject on step 1 and return nil on step 2. + model := &chattest.FakeModel{ + ProviderName: "fake", + StreamFn: func(_ context.Context, _ fantasy.Call) (fantasy.StreamResponse, error) { + mu.Lock() + step := streamCalls + streamCalls++ + mu.Unlock() + + if step < 2 { + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeToolInputStart, ID: "tc-" + strings.Repeat("x", step+1), ToolCallName: "noop"}, + {Type: fantasy.StreamPartTypeToolInputDelta, ID: "tc-" + strings.Repeat("x", step+1), Delta: `{}`}, + {Type: fantasy.StreamPartTypeToolInputEnd, ID: "tc-" + strings.Repeat("x", step+1)}, + { + Type: fantasy.StreamPartTypeToolCall, + ID: "tc-" + strings.Repeat("x", step+1), + ToolCallName: "noop", + ToolCallInput: `{}`, + }, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonToolCalls}, + }), nil + } + return streamFromParts([]fantasy.StreamPart{ + {Type: fantasy.StreamPartTypeTextStart, ID: "text-1"}, + {Type: fantasy.StreamPartTypeTextDelta, ID: "text-1", Delta: "done"}, + {Type: fantasy.StreamPartTypeTextEnd, ID: "text-1"}, + {Type: fantasy.StreamPartTypeFinish, FinishReason: fantasy.FinishReasonStop}, + }), nil + }, + } + + var prepareCalls atomic.Int32 + err := Run(context.Background(), RunOptions{ + Model: model, + Messages: []fantasy.Message{ + textMessage(fantasy.MessageRoleUser, "do something"), + }, + Tools: []fantasy.AgentTool{newNoopTool("noop")}, + MaxSteps: 5, + PersistStep: func(_ context.Context, _ PersistedStep) error { + return nil + }, + PrepareMessages: func(msgs []fantasy.Message) []fantasy.Message { + call := prepareCalls.Add(1) + if call == 1 { + // First call: inject a message. + return append(msgs, fantasy.Message{ + Role: fantasy.MessageRoleSystem, + Content: []fantasy.MessagePart{fantasy.TextPart{Text: "injected"}}, + }) + } + // Subsequent calls: no changes. + return nil + }, + }) + require.NoError(t, err) + require.Equal(t, 3, streamCalls) + // PrepareMessages is called before each of the 3 steps. + require.Equal(t, 3, int(prepareCalls.Load())) +} diff --git a/coderd/x/chatd/chattool/createworkspace.go b/coderd/x/chatd/chattool/createworkspace.go index 0419a7ed54..b57f0e031c 100644 --- a/coderd/x/chatd/chattool/createworkspace.go +++ b/coderd/x/chatd/chattool/createworkspace.go @@ -250,6 +250,15 @@ func CreateWorkspace(organizationID uuid.UUID, db database.Store, options Create )), nil } } + + // The agent is now online — re-fire so callers can + // load instruction files from the running agent. + if options.OnChatUpdated != nil { + if latest, err := db.GetChatByID(ctx, options.ChatID); err == nil { + options.OnChatUpdated(latest) + } + } + result := map[string]any{ "created": true, "workspace_name": workspace.FullName(), diff --git a/coderd/x/chatd/chattool/createworkspace_test.go b/coderd/x/chatd/chattool/createworkspace_test.go index c043d4bec8..94ee88d700 100644 --- a/coderd/x/chatd/chattool/createworkspace_test.go +++ b/coderd/x/chatd/chattool/createworkspace_test.go @@ -1129,6 +1129,136 @@ func expectExistingWorkspaceLookup( }, nil) } +func TestCreateWorkspace_OnChatUpdatedFiresAfterBuild(t *testing.T) { + t.Parallel() + + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + + ownerID := uuid.New() + templateID := uuid.New() + workspaceID := uuid.New() + chatID := uuid.New() + jobID := uuid.New() + buildID := uuid.New() + + // checkExistingWorkspace calls GetChatByID first. Return a chat + // with no workspace so the tool proceeds to creation. + db.EXPECT(). + GetChatByID(gomock.Any(), chatID). + Return(database.Chat{ + ID: chatID, + }, nil) + + db.EXPECT(). + GetAuthorizationUserRoles(gomock.Any(), ownerID). + Return(database.GetAuthorizationUserRolesRow{ + ID: ownerID, + Roles: []string{}, + Groups: []string{}, + Status: database.UserStatusActive, + }, nil) + + // Org check: GetTemplateByID returns a template in the + // same org (uuid.Nil matches our organizationID param). + db.EXPECT(). + GetTemplateByID(gomock.Any(), templateID). + Return(database.Template{ + ID: templateID, + OrganizationID: uuid.Nil, + }, nil) + + db.EXPECT(). + GetChatWorkspaceTTL(gomock.Any()). + Return("0s", nil) + + // UpdateChatWorkspaceBinding — triggers first OnChatUpdated. + db.EXPECT(). + UpdateChatWorkspaceBinding(gomock.Any(), gomock.Any()). + Return(database.Chat{ + ID: chatID, + WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}, + }, nil) + + // waitForBuild: fetch build, then poll job as completed. + db.EXPECT(). + GetWorkspaceBuildByID(gomock.Any(), buildID). + Return(database.WorkspaceBuild{ + ID: buildID, + WorkspaceID: workspaceID, + JobID: jobID, + }, nil) + db.EXPECT(). + GetProvisionerJobByID(gomock.Any(), jobID). + Return(database.ProvisionerJob{ + ID: jobID, + JobStatus: database.ProvisionerJobStatusSucceeded, + CompletedAt: validNullTime(time.Now()), + }, nil) + + // GetChatByID — called after waitForBuild for second OnChatUpdated. + // GetChatByID — called after waitForBuild for second OnChatUpdated. + db.EXPECT(). + GetChatByID(gomock.Any(), chatID). + Return(database.Chat{ + ID: chatID, + WorkspaceID: uuid.NullUUID{UUID: workspaceID, Valid: true}, + }, nil) + + // Agent lookup after build completes — return empty so we skip + // agent selection and waitForAgentReady. + db.EXPECT(). + GetWorkspaceAgentsInLatestBuildByWorkspaceID(gomock.Any(), workspaceID). + Return([]database.WorkspaceAgent{}, nil) + + var mu sync.Mutex + var callbackChats []database.Chat + + createFn := func(_ context.Context, _ uuid.UUID, req codersdk.CreateWorkspaceRequest) (codersdk.Workspace, error) { + return codersdk.Workspace{ + ID: workspaceID, + Name: req.Name, + OwnerName: "testuser", + LatestBuild: codersdk.WorkspaceBuild{ + ID: buildID, + }, + }, nil + } + + tool := CreateWorkspace(uuid.Nil, db, CreateWorkspaceOptions{ + OwnerID: ownerID, + + ChatID: chatID, + CreateFn: createFn, + WorkspaceMu: &sync.Mutex{}, + Logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}), + OnChatUpdated: func(chat database.Chat) { + mu.Lock() + callbackChats = append(callbackChats, chat) + mu.Unlock() + }, + }) + + input := fmt.Sprintf(`{"template_id":%q,"name":"test-callback"}`, templateID.String()) + resp, err := tool.Run(context.Background(), fantasy.ToolCall{ + ID: "call-1", + Name: "create_workspace", + Input: input, + }) + require.NoError(t, err) + require.False(t, resp.IsError) + + mu.Lock() + defer mu.Unlock() + require.Len(t, callbackChats, 2, + "OnChatUpdated should fire twice: once on binding, once after build completes") + // Both callbacks should carry the workspace ID. + for i, chat := range callbackChats { + require.True(t, chat.WorkspaceID.Valid, "callback %d should have workspace ID", i) + require.Equal(t, workspaceID, chat.WorkspaceID.UUID) + } +} + func validNullTime(t time.Time) sql.NullTime { return sql.NullTime{Time: t, Valid: true} } diff --git a/coderd/x/chatd/chattool/startworkspace.go b/coderd/x/chatd/chattool/startworkspace.go index e9291d2599..c703cab077 100644 --- a/coderd/x/chatd/chattool/startworkspace.go +++ b/coderd/x/chatd/chattool/startworkspace.go @@ -134,6 +134,13 @@ func StartWorkspace(options StartWorkspaceOptions) fantasy.AgentTool { build.ID, )), nil } + // The agent is now online — re-fire so callers can + // load instruction files from the running agent. + if options.OnChatUpdated != nil { + if latest, err := options.DB.GetChatByID(ctx, options.ChatID); err == nil { + options.OnChatUpdated(latest) + } + } return waitForAgentAndRespond(ctx, options.DB, options.AgentConnFn, ws, build.ID) case database.ProvisionerJobStatusSucceeded: // If the latest successful build is a start @@ -190,6 +197,14 @@ func StartWorkspace(options StartWorkspaceOptions) fantasy.AgentTool { )), nil } + // The agent is now online — re-fire so callers can + // load instruction files from the running agent. + if options.OnChatUpdated != nil { + if latest, err := options.DB.GetChatByID(ctx, options.ChatID); err == nil { + options.OnChatUpdated(latest) + } + } + return waitForAgentAndRespond(ctx, options.DB, options.AgentConnFn, ws, startBuild.ID) }) } diff --git a/coderd/x/chatd/instruction_test.go b/coderd/x/chatd/instruction_test.go index 9b8f3dfc10..514a8ff4cb 100644 --- a/coderd/x/chatd/instruction_test.go +++ b/coderd/x/chatd/instruction_test.go @@ -1,12 +1,15 @@ package chatd //nolint:testpackage // Uses internal symbols. import ( + "encoding/json" "strings" "testing" "charm.land/fantasy" + "github.com/sqlc-dev/pqtype" "github.com/stretchr/testify/require" + "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" "github.com/coder/coder/v2/coderd/x/chatd/chattool" "github.com/coder/coder/v2/codersdk" @@ -232,3 +235,55 @@ func TestFormatSystemInstructions(t *testing.T) { require.Contains(t, got, "Source: /real/AGENTS.md") }) } + +func TestInstructionFromContextFiles(t *testing.T) { + t.Parallel() + + makeMsg := func(parts []codersdk.ChatMessagePart) database.ChatMessage { + raw, _ := json.Marshal(parts) + return database.ChatMessage{ + Content: pqtype.NullRawMessage{RawMessage: raw, Valid: true}, + } + } + + t.Run("EmptyMessages", func(t *testing.T) { + t.Parallel() + got := instructionFromContextFiles(nil) + require.Empty(t, got) + }) + + t.Run("NoContextFileParts", func(t *testing.T) { + t.Parallel() + msgs := []database.ChatMessage{ + makeMsg([]codersdk.ChatMessagePart{ + { + Type: codersdk.ChatMessagePartTypeSkill, + SkillName: "test", + SkillDescription: "test skill", + }, + }), + } + got := instructionFromContextFiles(msgs) + require.Empty(t, got) + }) + + t.Run("ReconstructsFromContextFileParts", func(t *testing.T) { + t.Parallel() + msgs := []database.ChatMessage{ + makeMsg([]codersdk.ChatMessagePart{ + { + Type: codersdk.ChatMessagePartTypeContextFile, + ContextFileOS: "linux", + ContextFileDirectory: "/home/coder/project", + ContextFileContent: "project rules", + ContextFilePath: "/home/coder/project/AGENTS.md", + }, + }), + } + got := instructionFromContextFiles(msgs) + require.Contains(t, got, "Operating System: linux") + require.Contains(t, got, "Working Directory: /home/coder/project") + require.Contains(t, got, "Source: /home/coder/project/AGENTS.md") + require.Contains(t, got, "project rules") + }) +}