mirror of
https://github.com/coder/coder.git
synced 2026-09-21 20:51:01 +08:00
fix: workspace auto-refresh during the chat flow (#22447)
This commit is contained in:
@@ -1844,6 +1844,13 @@ func (p *Server) runChat(
|
||||
}()
|
||||
|
||||
currentChat := chat
|
||||
loadChatSnapshot := func(
|
||||
loadCtx context.Context,
|
||||
chatID uuid.UUID,
|
||||
) (database.Chat, error) {
|
||||
//nolint:gocritic // System context required to load chat snapshots for the stream.
|
||||
return p.db.GetChatByID(dbauthz.AsSystemRestricted(loadCtx), chatID)
|
||||
}
|
||||
var (
|
||||
chatStateMu sync.Mutex
|
||||
workspaceMu sync.Mutex
|
||||
@@ -1872,6 +1879,23 @@ func (p *Server) runChat(
|
||||
return nil, xerrors.New("workspace agent connector is not configured")
|
||||
}
|
||||
|
||||
if !chatSnapshot.WorkspaceID.Valid {
|
||||
refreshedChat, refreshErr := refreshChatWorkspaceSnapshot(
|
||||
ctx,
|
||||
chatSnapshot,
|
||||
loadChatSnapshot,
|
||||
)
|
||||
if refreshErr != nil {
|
||||
return nil, refreshErr
|
||||
}
|
||||
if refreshedChat.WorkspaceID.Valid {
|
||||
chatStateMu.Lock()
|
||||
currentChat = refreshedChat
|
||||
chatSnapshot = refreshedChat
|
||||
chatStateMu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
if !chatSnapshot.WorkspaceID.Valid {
|
||||
return nil, xerrors.New("chat has no workspace")
|
||||
}
|
||||
@@ -2390,6 +2414,23 @@ func usageNullInt64(value int64, valid bool) sql.NullInt64 {
|
||||
}
|
||||
}
|
||||
|
||||
func refreshChatWorkspaceSnapshot(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
loadChat func(context.Context, uuid.UUID) (database.Chat, error),
|
||||
) (database.Chat, error) {
|
||||
if chat.WorkspaceID.Valid || loadChat == nil {
|
||||
return chat, nil
|
||||
}
|
||||
|
||||
refreshedChat, err := loadChat(ctx, chat.ID)
|
||||
if err != nil {
|
||||
return chat, xerrors.Errorf("reload chat workspace state: %w", err)
|
||||
}
|
||||
|
||||
return refreshedChat, nil
|
||||
}
|
||||
|
||||
// resolveInstructions returns the combined system instructions for the
|
||||
// workspace agent. It reads the home-level (~/.coder/AGENTS.md) and
|
||||
// working-directory-level (<pwd>/AGENTS.md) instruction files, combines
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
)
|
||||
|
||||
func TestRefreshChatWorkspaceSnapshot_NoReloadWhenWorkspacePresent(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
workspaceID := uuid.New()
|
||||
chat := database.Chat{
|
||||
ID: uuid.New(),
|
||||
WorkspaceID: uuid.NullUUID{
|
||||
UUID: workspaceID,
|
||||
Valid: true,
|
||||
},
|
||||
}
|
||||
|
||||
calls := 0
|
||||
refreshed, err := refreshChatWorkspaceSnapshot(
|
||||
context.Background(),
|
||||
chat,
|
||||
func(context.Context, uuid.UUID) (database.Chat, error) {
|
||||
calls++
|
||||
return database.Chat{}, nil
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, chat, refreshed)
|
||||
require.Equal(t, 0, calls)
|
||||
}
|
||||
|
||||
func TestRefreshChatWorkspaceSnapshot_ReloadsWhenWorkspaceMissing(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
chatID := uuid.New()
|
||||
workspaceID := uuid.New()
|
||||
chat := database.Chat{ID: chatID}
|
||||
reloaded := database.Chat{
|
||||
ID: chatID,
|
||||
WorkspaceID: uuid.NullUUID{
|
||||
UUID: workspaceID,
|
||||
Valid: true,
|
||||
},
|
||||
}
|
||||
|
||||
calls := 0
|
||||
refreshed, err := refreshChatWorkspaceSnapshot(
|
||||
context.Background(),
|
||||
chat,
|
||||
func(_ context.Context, id uuid.UUID) (database.Chat, error) {
|
||||
calls++
|
||||
require.Equal(t, chatID, id)
|
||||
return reloaded, nil
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, reloaded, refreshed)
|
||||
require.Equal(t, 1, calls)
|
||||
}
|
||||
|
||||
func TestRefreshChatWorkspaceSnapshot_ReturnsReloadError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
chat := database.Chat{ID: uuid.New()}
|
||||
loadErr := xerrors.New("boom")
|
||||
|
||||
refreshed, err := refreshChatWorkspaceSnapshot(
|
||||
context.Background(),
|
||||
chat,
|
||||
func(context.Context, uuid.UUID) (database.Chat, error) {
|
||||
return database.Chat{}, loadErr
|
||||
},
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.ErrorContains(t, err, "reload chat workspace state")
|
||||
require.ErrorContains(t, err, loadErr.Error())
|
||||
require.Equal(t, chat, refreshed)
|
||||
}
|
||||
@@ -5,6 +5,9 @@ import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -15,14 +18,17 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/agent/agenttest"
|
||||
"github.com/coder/coder/v2/coderd/chatd"
|
||||
"github.com/coder/coder/v2/coderd/chatd/chattest"
|
||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/provisioner/echo"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
@@ -702,6 +708,159 @@ func TestSubscribeNoPubsubNoDuplicateMessageParts(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateWorkspaceTool_EndToEnd(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
deploymentValues := coderdtest.DeploymentValues(t)
|
||||
deploymentValues.Experiments = []string{string(codersdk.ExperimentAgents)}
|
||||
client := coderdtest.New(t, &coderdtest.Options{
|
||||
DeploymentValues: deploymentValues,
|
||||
IncludeProvisionerDaemon: true,
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
agentToken := uuid.NewString()
|
||||
version := coderdtest.CreateTemplateVersion(t, client, user.OrganizationID, &echo.Responses{
|
||||
Parse: echo.ParseComplete,
|
||||
ProvisionPlan: echo.PlanComplete,
|
||||
ProvisionApply: echo.ApplyComplete,
|
||||
ProvisionGraph: echo.ProvisionGraphWithAgent(agentToken),
|
||||
})
|
||||
coderdtest.AwaitTemplateVersionJobCompleted(t, client, version.ID)
|
||||
template := coderdtest.CreateTemplate(t, client, user.OrganizationID, version.ID)
|
||||
|
||||
// Start the test workspace agent so create_workspace can wait for
|
||||
// the agent to become reachable before returning.
|
||||
_ = agenttest.New(t, client.URL, agentToken)
|
||||
|
||||
workspaceName := "chat-ws-" + strings.ReplaceAll(uuid.NewString(), "-", "")[:8]
|
||||
createWorkspaceArgs := fmt.Sprintf(
|
||||
`{"template_id":%q,"name":%q}`,
|
||||
template.ID.String(),
|
||||
workspaceName,
|
||||
)
|
||||
|
||||
var streamedCallCount atomic.Int32
|
||||
var streamedCallsMu sync.Mutex
|
||||
streamedCalls := make([][]chattest.OpenAIMessage, 0, 2)
|
||||
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("Create workspace test")
|
||||
}
|
||||
|
||||
streamedCallsMu.Lock()
|
||||
streamedCalls = append(streamedCalls, append([]chattest.OpenAIMessage(nil), req.Messages...))
|
||||
streamedCallsMu.Unlock()
|
||||
|
||||
if streamedCallCount.Add(1) == 1 {
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAIToolCallChunk("create_workspace", createWorkspaceArgs),
|
||||
)
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("Workspace created and ready.")...,
|
||||
)
|
||||
})
|
||||
|
||||
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
|
||||
Provider: "openai-compat",
|
||||
APIKey: "test-api-key",
|
||||
BaseURL: openAIURL,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
contextLimit := int64(4096)
|
||||
isDefault := true
|
||||
_, err = client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai-compat",
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
IsDefault: &isDefault,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
Content: []codersdk.ChatInputPart{
|
||||
{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: "Create a workspace from the template and continue.",
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
var chatWithMessages codersdk.ChatWithMessages
|
||||
require.Eventually(t, func() bool {
|
||||
got, getErr := client.GetChat(ctx, chat.ID)
|
||||
if getErr != nil {
|
||||
return false
|
||||
}
|
||||
chatWithMessages = got
|
||||
return got.Chat.Status == codersdk.ChatStatusWaiting || got.Chat.Status == codersdk.ChatStatusError
|
||||
}, testutil.WaitLong, testutil.IntervalFast)
|
||||
|
||||
if chatWithMessages.Chat.Status == codersdk.ChatStatusError {
|
||||
lastError := ""
|
||||
if chatWithMessages.Chat.LastError != nil {
|
||||
lastError = *chatWithMessages.Chat.LastError
|
||||
}
|
||||
require.FailNowf(t, "chat run failed", "last_error=%q", lastError)
|
||||
}
|
||||
|
||||
require.NotNil(t, chatWithMessages.Chat.WorkspaceID)
|
||||
workspaceID := *chatWithMessages.Chat.WorkspaceID
|
||||
workspace, err := client.Workspace(ctx, workspaceID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, workspaceName, workspace.Name)
|
||||
|
||||
var foundCreateWorkspaceResult bool
|
||||
for _, message := range chatWithMessages.Messages {
|
||||
if message.Role != "tool" {
|
||||
continue
|
||||
}
|
||||
for _, part := range message.Content {
|
||||
if part.Type != codersdk.ChatMessagePartTypeToolResult || part.ToolName != "create_workspace" {
|
||||
continue
|
||||
}
|
||||
var result map[string]any
|
||||
require.NoError(t, json.Unmarshal(part.Result, &result))
|
||||
created, ok := result["created"].(bool)
|
||||
require.True(t, ok)
|
||||
require.True(t, created)
|
||||
foundCreateWorkspaceResult = true
|
||||
}
|
||||
}
|
||||
require.True(t, foundCreateWorkspaceResult, "expected create_workspace tool result message")
|
||||
|
||||
require.GreaterOrEqual(t, streamedCallCount.Load(), int32(2))
|
||||
streamedCallsMu.Lock()
|
||||
recordedStreamCalls := append([][]chattest.OpenAIMessage(nil), streamedCalls...)
|
||||
streamedCallsMu.Unlock()
|
||||
require.GreaterOrEqual(t, len(recordedStreamCalls), 2)
|
||||
|
||||
var foundToolResultInSecondCall bool
|
||||
for _, message := range recordedStreamCalls[1] {
|
||||
if message.Role != "tool" {
|
||||
continue
|
||||
}
|
||||
if !json.Valid([]byte(message.Content)) {
|
||||
continue
|
||||
}
|
||||
var result map[string]any
|
||||
if err := json.Unmarshal([]byte(message.Content), &result); err != nil {
|
||||
continue
|
||||
}
|
||||
created, ok := result["created"].(bool)
|
||||
if ok && created {
|
||||
foundToolResultInSecondCall = true
|
||||
break
|
||||
}
|
||||
}
|
||||
require.True(t, foundToolResultInSecondCall, "expected second streamed model call to include create_workspace tool output")
|
||||
}
|
||||
|
||||
func newTestServer(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
|
||||
Reference in New Issue
Block a user