perf(coderd/chatd): reuse workspace context within a turn (#23145)

## Summary
- reuse workspace agent context within a single `runChat()` turn
- remove duplicate latest-build agent lookups between
`resolveInstructions()` and `getWorkspaceConn()`
- avoid the extra `GetWorkspaceAgentByID` fetch when the selected
`WorkspaceAgent` already has the needed metadata
- add focused internal tests for reuse and refresh-on-dial-failure

## Why
This came out of a 5000-chat / 10-turn scaletest on bravo against a
single workspace.

The run completed successfully, but coderd stayed DB-pool bound, and one
workspace-backed hot path stood out:
- `GetWorkspaceAgentsInLatestBuildByWorkspaceID ≈ 46.7k`
- `GetWorkspaceByID ≈ 48.0k`
- `GetWorkspaceAgentByID ≈ 2.2k`

Within one `runChat()` turn, chatd was rediscovering the same workspace
agent multiple times just to resolve instructions and open the workspace
connection.

## What this changes
This PR introduces a **turn-local** workspace context helper so a single
acquired turn can:
- resolve the selected workspace agent once
- reuse that agent for instruction resolution
- reuse the same `AgentConn` for workspace tools and reload/compaction

This stays turn-local only, so a later turn on another replica still
rebuilds fresh context from the DB.

## Expected impact
This is an incremental improvement, not a full fix.

It should reduce duplicated workspace-agent lookups and shave some DB
pressure from a hot path for workspace-backed chats, while preserving
multi-replica correctness.

## Testing
- `go test ./coderd/chatd/...`
- `golangci-lint run ./coderd/chatd/...`
This commit is contained in:
Ethan
2026-03-18 00:33:44 +11:00
committed by GitHub
parent 3c430a67fa
commit a33605df58
2 changed files with 331 additions and 110 deletions
+192 -110
View File
@@ -106,6 +106,167 @@ type cachedInstruction struct {
fetchedAt time.Time
}
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()
}
func (c *turnWorkspaceContext) close() {
c.mu.Lock()
releaseConn := c.releaseConn
c.conn = nil
c.releaseConn = nil
c.mu.Unlock()
if releaseConn != nil {
releaseConn()
}
}
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 {
c.chatStateMu.Lock()
chatSnapshot := *c.currentChat
c.chatStateMu.Unlock()
return chatSnapshot, c.agent, nil
}
return c.loadWorkspaceAgentLocked(ctx)
}
func (c *turnWorkspaceContext) refreshWorkspaceAgent(
ctx context.Context,
) (database.Chat, database.WorkspaceAgent, error) {
c.mu.Lock()
defer c.mu.Unlock()
c.agent = database.WorkspaceAgent{}
c.agentLoaded = false
return c.loadWorkspaceAgentLocked(ctx)
}
func (c *turnWorkspaceContext) loadWorkspaceAgentLocked(
ctx context.Context,
) (database.Chat, database.WorkspaceAgent, error) {
c.chatStateMu.Lock()
chatSnapshot := *c.currentChat
c.chatStateMu.Unlock()
if !chatSnapshot.WorkspaceID.Valid {
refreshedChat, refreshErr := refreshChatWorkspaceSnapshot(
ctx,
chatSnapshot,
c.loadChatSnapshot,
)
if refreshErr != nil {
return chatSnapshot, database.WorkspaceAgent{}, refreshErr
}
if refreshedChat.WorkspaceID.Valid {
c.chatStateMu.Lock()
*c.currentChat = refreshedChat
c.chatStateMu.Unlock()
chatSnapshot = refreshedChat
}
}
if !chatSnapshot.WorkspaceID.Valid {
return chatSnapshot, database.WorkspaceAgent{}, xerrors.New("chat has no workspace")
}
agents, err := c.server.db.GetWorkspaceAgentsInLatestBuildByWorkspaceID(
ctx,
chatSnapshot.WorkspaceID.UUID,
)
if err != nil || len(agents) == 0 {
return chatSnapshot, database.WorkspaceAgent{}, xerrors.New("chat has no workspace agent")
}
c.agent = agents[0]
c.agentLoaded = true
return chatSnapshot, c.agent, nil
}
func (c *turnWorkspaceContext) getWorkspaceConn(ctx context.Context) (workspacesdk.AgentConn, error) {
c.mu.Lock()
if c.conn != nil {
currentConn := c.conn
c.mu.Unlock()
return currentConn, nil
}
c.mu.Unlock()
if c.server.agentConnFn == nil {
return nil, xerrors.New("workspace agent connector is not configured")
}
chatSnapshot, agent, err := c.ensureWorkspaceAgent(ctx)
if err != nil {
return nil, err
}
agentConn, agentRelease, err := c.server.agentConnFn(ctx, agent.ID)
if err != nil {
refreshedChat, refreshedAgent, refreshErr := c.refreshWorkspaceAgent(ctx)
if refreshErr != nil {
return nil, xerrors.Errorf("connect to workspace agent: %w", err)
}
retryConn, retryRelease, retryErr := c.server.agentConnFn(ctx, refreshedAgent.ID)
if retryErr != nil {
return nil, xerrors.Errorf("connect to workspace agent after refresh: %w", retryErr)
}
chatSnapshot = refreshedChat
agentConn = retryConn
agentRelease = retryRelease
}
c.mu.Lock()
if c.conn == nil {
c.conn = agentConn
c.releaseConn = agentRelease
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()
return agentConn, nil
}
currentConn := c.conn
c.mu.Unlock()
agentRelease()
return currentConn, nil
}
// AgentConnFunc provides access to workspace agent connections.
type AgentConnFunc func(ctx context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error)
@@ -2293,98 +2454,24 @@ func (p *Server) runChat(
var (
chatStateMu sync.Mutex
workspaceMu sync.Mutex
conn workspacesdk.AgentConn
releaseConn func()
)
closeConn := func() {
if releaseConn != nil {
releaseConn()
releaseConn = nil
}
}
defer closeConn()
getWorkspaceConn := func(ctx context.Context) (workspacesdk.AgentConn, error) {
chatStateMu.Lock()
if conn != nil {
currentConn := conn
chatStateMu.Unlock()
return currentConn, nil
}
chatSnapshot := currentChat
chatStateMu.Unlock()
if p.agentConnFn == nil {
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")
}
agents, err := p.db.GetWorkspaceAgentsInLatestBuildByWorkspaceID(
ctx,
chatSnapshot.WorkspaceID.UUID,
)
if err != nil || len(agents) == 0 {
return nil, xerrors.New("chat has no workspace agent")
}
agentConn, agentRelease, err := p.agentConnFn(ctx, agents[0].ID)
if err != nil {
return nil, xerrors.Errorf("connect to workspace agent: %w", err)
}
chatStateMu.Lock()
if conn == nil {
conn = agentConn
releaseConn = agentRelease
var ancestorIDs []string
if chatSnapshot.ParentChatID.Valid {
ancestorIDs = append(ancestorIDs, chatSnapshot.ParentChatID.UUID.String())
}
ancestorJSON, err := json.Marshal(ancestorIDs)
if err != nil {
logger.Warn(ctx, "failed to marshal ancestor chat IDs", slog.Error(err))
ancestorJSON = []byte("[]")
}
agentConn.SetExtraHeaders(http.Header{
workspacesdk.CoderChatIDHeader: {chatSnapshot.ID.String()},
workspacesdk.CoderAncestorChatIDsHeader: {string(ancestorJSON)},
})
chatStateMu.Unlock()
return agentConn, nil
}
currentConn := conn
chatStateMu.Unlock()
agentRelease()
return currentConn, nil
workspaceCtx := turnWorkspaceContext{
server: p,
chatStateMu: &chatStateMu,
currentChat: &currentChat,
loadChatSnapshot: loadChatSnapshot,
}
defer workspaceCtx.close()
var instruction, resolvedUserPrompt string
var g2 errgroup.Group
g2.Go(func() error {
instruction = p.resolveInstructions(ctx, chat, getWorkspaceConn)
instruction = p.resolveInstructions(
ctx,
chat,
workspaceCtx.getWorkspaceAgent,
workspaceCtx.getWorkspaceConn,
)
return nil
})
g2.Go(func() error {
@@ -2652,25 +2739,25 @@ func (p *Server) runChat(
// Here are all the tools we have for the chat.
tools := []fantasy.AgentTool{
chattool.ReadFile(chattool.ReadFileOptions{
GetWorkspaceConn: getWorkspaceConn,
GetWorkspaceConn: workspaceCtx.getWorkspaceConn,
}),
chattool.WriteFile(chattool.WriteFileOptions{
GetWorkspaceConn: getWorkspaceConn,
GetWorkspaceConn: workspaceCtx.getWorkspaceConn,
}),
chattool.EditFiles(chattool.EditFilesOptions{
GetWorkspaceConn: getWorkspaceConn,
GetWorkspaceConn: workspaceCtx.getWorkspaceConn,
}),
chattool.Execute(chattool.ExecuteOptions{
GetWorkspaceConn: getWorkspaceConn,
GetWorkspaceConn: workspaceCtx.getWorkspaceConn,
}),
chattool.ProcessOutput(chattool.ProcessToolOptions{
GetWorkspaceConn: getWorkspaceConn,
GetWorkspaceConn: workspaceCtx.getWorkspaceConn,
}),
chattool.ProcessList(chattool.ProcessToolOptions{
GetWorkspaceConn: getWorkspaceConn,
GetWorkspaceConn: workspaceCtx.getWorkspaceConn,
}),
chattool.ProcessSignal(chattool.ProcessToolOptions{
GetWorkspaceConn: getWorkspaceConn,
GetWorkspaceConn: workspaceCtx.getWorkspaceConn,
}),
}
// Only root chats (not delegated subagents) get workspace
@@ -2725,7 +2812,7 @@ func (p *Server) runChat(
Runner: chattool.NewComputerUseTool(
workspacesdk.DesktopDisplayWidth,
workspacesdk.DesktopDisplayHeight,
getWorkspaceConn, quartz.NewReal(),
workspaceCtx.getWorkspaceConn, quartz.NewReal(),
),
})
}
@@ -2763,7 +2850,12 @@ func (p *Server) runChat(
var reloadInstruction, reloadUserPrompt string
var rg errgroup.Group
rg.Go(func() error {
reloadInstruction = p.resolveInstructions(reloadCtx, chat, getWorkspaceConn)
reloadInstruction = p.resolveInstructions(
reloadCtx,
chat,
workspaceCtx.getWorkspaceAgent,
workspaceCtx.getWorkspaceConn,
)
return nil
})
rg.Go(func() error {
@@ -3144,20 +3236,18 @@ func refreshChatWorkspaceSnapshot(
func (p *Server) resolveInstructions(
ctx context.Context,
chat database.Chat,
getWorkspaceAgent func(context.Context) (database.WorkspaceAgent, error),
getWorkspaceConn func(context.Context) (workspacesdk.AgentConn, error),
) string {
if !chat.WorkspaceID.Valid {
if !chat.WorkspaceID.Valid || getWorkspaceAgent == nil {
return ""
}
agents, agentsErr := p.db.GetWorkspaceAgentsInLatestBuildByWorkspaceID(
ctx,
chat.WorkspaceID.UUID,
)
if agentsErr != nil || len(agents) == 0 {
agent, agentErr := getWorkspaceAgent(ctx)
if agentErr != nil {
return ""
}
agentID := agents[0].ID
agentID := agent.ID
p.instructionCacheMu.Lock()
cached, ok := p.instructionCache[agentID]
@@ -3167,14 +3257,6 @@ func (p *Server) resolveInstructions(
return cached.instruction
}
// Look up the agent's OS and working directory.
agent, err := p.db.GetWorkspaceAgentByID(ctx, agentID)
if err != nil {
p.logger.Debug(ctx, "failed to look up workspace agent for instruction context",
slog.F("agent_id", agentID),
slog.Error(err),
)
}
directory := agent.ExpandedDirectory
if directory == "" {
directory = agent.Directory
+139
View File
@@ -2,13 +2,20 @@ package chatd
import (
"context"
"sync"
"testing"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"golang.org/x/xerrors"
"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbmock"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/codersdk/workspacesdk"
"github.com/coder/coder/v2/codersdk/workspacesdk/agentconnmock"
)
func TestRefreshChatWorkspaceSnapshot_NoReloadWhenWorkspacePresent(t *testing.T) {
@@ -84,3 +91,135 @@ func TestRefreshChatWorkspaceSnapshot_ReturnsReloadError(t *testing.T) {
require.ErrorContains(t, err, loadErr.Error())
require.Equal(t, chat, refreshed)
}
func TestResolveInstructionsReusesTurnLocalWorkspaceAgent(t *testing.T) {
t.Parallel()
ctx := context.Background()
ctrl := gomock.NewController(t)
db := dbmock.NewMockStore(ctrl)
workspaceID := uuid.New()
chat := database.Chat{
ID: uuid.New(),
WorkspaceID: uuid.NullUUID{
UUID: workspaceID,
Valid: true,
},
}
workspaceAgent := database.WorkspaceAgent{
ID: uuid.New(),
OperatingSystem: "linux",
Directory: "/home/coder/project",
ExpandedDirectory: "/home/coder/project",
}
db.EXPECT().GetWorkspaceAgentsInLatestBuildByWorkspaceID(
gomock.Any(),
workspaceID,
).Return([]database.WorkspaceAgent{workspaceAgent}, nil).Times(1)
conn := agentconnmock.NewMockAgentConn(ctrl)
conn.EXPECT().SetExtraHeaders(gomock.Any()).Times(1)
conn.EXPECT().LS(gomock.Any(), "", gomock.Any()).Return(
workspacesdk.LSResponse{},
codersdk.NewTestError(404, "POST", "/api/v0/list-directory"),
).Times(1)
conn.EXPECT().ReadFile(
gomock.Any(),
"/home/coder/project/AGENTS.md",
int64(0),
int64(maxInstructionFileBytes+1),
).Return(
nil,
"",
codersdk.NewTestError(404, "GET", "/api/v0/read-file"),
).Times(1)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
server := &Server{
db: db,
logger: logger,
instructionCache: make(map[uuid.UUID]cachedInstruction),
agentConnFn: func(context.Context, uuid.UUID) (workspacesdk.AgentConn, func(), error) {
return conn, func() {}, nil
},
}
chatStateMu := &sync.Mutex{}
currentChat := chat
workspaceCtx := turnWorkspaceContext{
server: server,
chatStateMu: chatStateMu,
currentChat: &currentChat,
loadChatSnapshot: func(context.Context, uuid.UUID) (database.Chat, error) { return database.Chat{}, nil },
}
t.Cleanup(workspaceCtx.close)
instruction := server.resolveInstructions(
ctx,
chat,
workspaceCtx.getWorkspaceAgent,
workspaceCtx.getWorkspaceConn,
)
require.Contains(t, instruction, "Operating System: linux")
require.Contains(t, instruction, "Working Directory: /home/coder/project")
}
func TestTurnWorkspaceContextGetWorkspaceConnRefreshesWorkspaceAgent(t *testing.T) {
t.Parallel()
ctx := context.Background()
ctrl := gomock.NewController(t)
db := dbmock.NewMockStore(ctrl)
workspaceID := uuid.New()
chat := database.Chat{
ID: uuid.New(),
WorkspaceID: uuid.NullUUID{
UUID: workspaceID,
Valid: true,
},
}
initialAgent := database.WorkspaceAgent{ID: uuid.New()}
refreshedAgent := database.WorkspaceAgent{ID: uuid.New()}
gomock.InOrder(
db.EXPECT().GetWorkspaceAgentsInLatestBuildByWorkspaceID(
gomock.Any(),
workspaceID,
).Return([]database.WorkspaceAgent{initialAgent}, nil),
db.EXPECT().GetWorkspaceAgentsInLatestBuildByWorkspaceID(
gomock.Any(),
workspaceID,
).Return([]database.WorkspaceAgent{refreshedAgent}, nil),
)
conn := agentconnmock.NewMockAgentConn(ctrl)
conn.EXPECT().SetExtraHeaders(gomock.Any()).Times(1)
var dialed []uuid.UUID
server := &Server{db: db}
server.agentConnFn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
dialed = append(dialed, agentID)
if agentID == initialAgent.ID {
return nil, nil, xerrors.New("dial failed")
}
return conn, func() {}, nil
}
chatStateMu := &sync.Mutex{}
currentChat := chat
workspaceCtx := turnWorkspaceContext{
server: server,
chatStateMu: chatStateMu,
currentChat: &currentChat,
loadChatSnapshot: func(context.Context, uuid.UUID) (database.Chat, error) { return database.Chat{}, nil },
}
t.Cleanup(workspaceCtx.close)
gotConn, err := workspaceCtx.getWorkspaceConn(ctx)
require.NoError(t, err)
require.Same(t, conn, gotConn)
require.Equal(t, []uuid.UUID{initialAgent.ID, refreshedAgent.ID}, dialed)
}