mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: re-fetch context files and skills from workspace on each turn (#24360)
Context files (AGENTS.md) and skills were only fetched from the
workspace on the first turn or when the agent changed. On subsequent
turns, stale content from persisted messages was used. This meant that
if AGENTS.md or skills were modified on the workspace between turns, the
agent wouldn't see the changes until the user created a new chat.
## Changes
- Extract `fetchWorkspaceContext` from `persistInstructionFiles` to
allow fetching workspace context without persisting
- On subsequent turns, re-fetch fresh context from the workspace instead
of reading stale persisted content; falls back to persisted messages if
the workspace dial fails
- Update `ReloadMessages` callback to re-derive instruction and skills
from reloaded database messages after compaction, instead of using
captured closure variables
- Add `formatSystemInstructionsFromParts` helper to build system
instructions directly from agent parts without requiring separate
OS/directory params
- Add tests for the new helper
<details><summary>Implementation Notes</summary>
### Root cause
In `runChat`, the `else if hasContextFiles` branch (subsequent turns)
called `instructionFromContextFiles(messages)` which read stale content
from persisted DB messages. The `ReloadMessages` callback
(post-compaction) also used captured `instruction`/`skills` closure
variables from the start of the turn, never re-deriving them.
### Approach
1. **Extract `fetchWorkspaceContext`** — Pure refactor of the fetch-only
part of `persistInstructionFiles` (agent connection, context config
retrieval, content sanitization, metadata stamping). Returns parts +
skills without persisting.
2. **Subsequent turns**: Instead of reading from persisted messages,
launch a `g2` goroutine that calls `fetchWorkspaceContext` to get fresh
context from the workspace. Falls back gracefully to persisted messages
if the workspace is unreachable.
3. **ReloadMessages**: Re-derive `instruction` from
`instructionFromContextFiles(reloadedMsgs)` and `skills` from
`skillsFromParts(reloadedMsgs)` using the freshly loaded messages, with
fallback to captured values if the reloaded messages don't contain
context (e.g. compacted away).
</details>
> 🤖 Generated by Coder Agents
This commit is contained in:
+120
-30
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()))
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user