fix: reuse shared tailnet for coderd-hosted MCP workspace tools (#24460)

## Problem

Coderd can expose an MCP server at `/api/experimental/mcp/http` (we have
this enabled on dogfood). Its workspace tools dialed agents through a
per-call client-side tailnet stack. Every tool call re-created a
WireGuard device, netstack, magicsock + UDP sockets, DERP connection,
coordinator websocket, and their goroutines — in a process that already
runs a long-lived shared tailnet. The duplicate stacks drove up resource
usage under load.

## Fix

Route this server's tool calls through the existing shared tailnet, so
none of those transports are reconstructed per call. Closing an
`AgentConn` now releases a tunnel reference instead of tearing down a
transport.

## Potential follow-up

`coder exp mcp server` still builds a fresh tailnet per call. It pays
per-call latency and causes coordinator/DERP churn. A shared CLI tailnet
is more involved — unlike coderd, the CLI has no existing shared tailnet
to reuse, so it would need a new long-lived client-side tailnet with
reconnect, sleep/wake, and idle-destination handling. There's less
motivation to optimize this, given the client-side MCP does not compete
for resources with coderd.

Closes CODAGT-199

> Generated by mux, but reviewed by a human
This commit is contained in:
Ethan
2026-04-21 11:37:10 +10:00
committed by GitHub
parent 1203f625b7
commit 181e103201
8 changed files with 283 additions and 99 deletions
+1 -1
View File
@@ -101,7 +101,7 @@ Examples:
ctx, cancel := context.WithTimeoutCause(ctx, 5*time.Minute, xerrors.New("MCP handler timeout after 5 min"))
defer cancel()
conn, err := newAgentConn(ctx, deps.coderClient, args.Workspace)
conn, err := openAgentConn(ctx, deps, args.Workspace)
if err != nil {
return WorkspaceBashResult{}, err
}
+65 -40
View File
@@ -65,6 +65,16 @@ func NewDeps(client *codersdk.Client, opts ...func(*Deps)) (Deps, error) {
for _, opt := range opts {
opt(&d)
}
if d.agentConnFn == nil && d.coderClient != nil {
workspaceClient := workspacesdk.New(d.coderClient)
d.agentConnFn = func(ctx context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
conn, err := workspaceClient.DialAgent(ctx, agentID, nil)
if err != nil {
return nil, nil, err
}
return conn, nil, nil
}
}
// Allow nil client for unauthenticated operation
// This enables tools that don't require user authentication to function
return d, nil
@@ -74,6 +84,7 @@ func NewDeps(client *codersdk.Client, opts ...func(*Deps)) (Deps, error) {
type Deps struct {
coderClient *codersdk.Client
report func(ReportTaskArgs) error
agentConnFn workspacesdk.AgentConnFunc
}
func (d Deps) ServerURL() string {
@@ -89,6 +100,55 @@ func WithTaskReporter(fn func(ReportTaskArgs) error) func(*Deps) {
}
}
// WithAgentConnFunc overrides how workspace tools open logical connections to
// workspace agents.
func WithAgentConnFunc(agentConnFn workspacesdk.AgentConnFunc) func(*Deps) {
return func(d *Deps) {
d.agentConnFn = agentConnFn
}
}
// openAgentConn opens a ready workspace agent session for workspace inputs in
// [owner/]workspace[.agent] format.
func openAgentConn(ctx context.Context, deps Deps, workspace string) (workspacesdk.AgentConn, error) {
if deps.coderClient == nil {
return nil, xerrors.New("workspace tools require an authenticated client")
}
workspaceName := NormalizeWorkspaceInput(workspace)
_, workspaceAgent, err := findWorkspaceAndAgent(ctx, deps.coderClient, workspaceName)
if err != nil {
return nil, xerrors.Errorf("failed to find workspace: %w", err)
}
if err := cliui.Agent(ctx, io.Discard, workspaceAgent.ID, cliui.AgentOptions{
FetchInterval: 0,
Fetch: deps.coderClient.WorkspaceAgent,
FetchLogs: deps.coderClient.WorkspaceAgentLogsAfter,
// Always wait for startup scripts.
Wait: true,
}); err != nil {
return nil, xerrors.Errorf("agent not ready: %w", err)
}
conn, release, err := deps.agentConnFn(ctx, workspaceAgent.ID)
if err != nil {
return nil, xerrors.Errorf("failed to dial agent: %w", err)
}
wrappedConn := workspacesdk.WrapAgentConn(conn, func() error {
if release != nil {
release()
}
return nil
})
if wrappedConn == nil {
return nil, xerrors.New("agent connection function returned nil connection")
}
return wrappedConn, nil
}
// HandlerFunc is a typed function that handles a tool call.
type HandlerFunc[Arg, Ret any] func(context.Context, Deps, Arg) (Ret, error)
@@ -1501,7 +1561,7 @@ var WorkspaceLS = Tool[WorkspaceLSArgs, WorkspaceLSResponse]{
MCPAnnotations: mcpReadOnlyAnnotations,
UserClientOptional: true,
Handler: func(ctx context.Context, deps Deps, args WorkspaceLSArgs) (WorkspaceLSResponse, error) {
conn, err := newAgentConn(ctx, deps.coderClient, args.Workspace)
conn, err := openAgentConn(ctx, deps, args.Workspace)
if err != nil {
return WorkspaceLSResponse{}, err
}
@@ -1567,7 +1627,7 @@ var WorkspaceReadFile = Tool[WorkspaceReadFileArgs, WorkspaceReadFileResponse]{
MCPAnnotations: mcpReadOnlyAnnotations,
UserClientOptional: true,
Handler: func(ctx context.Context, deps Deps, args WorkspaceReadFileArgs) (WorkspaceReadFileResponse, error) {
conn, err := newAgentConn(ctx, deps.coderClient, args.Workspace)
conn, err := openAgentConn(ctx, deps, args.Workspace)
if err != nil {
return WorkspaceReadFileResponse{}, err
}
@@ -1641,7 +1701,7 @@ content you are trying to write, then re-encode it properly.
MCPAnnotations: mcpDestructiveAnnotations,
UserClientOptional: true,
Handler: func(ctx context.Context, deps Deps, args WorkspaceWriteFileArgs) (codersdk.Response, error) {
conn, err := newAgentConn(ctx, deps.coderClient, args.Workspace)
conn, err := openAgentConn(ctx, deps, args.Workspace)
if err != nil {
return codersdk.Response{}, err
}
@@ -1716,7 +1776,7 @@ var WorkspaceEditFile = Tool[WorkspaceEditFileArgs, WorkspaceEditFilesResponse]{
MCPAnnotations: mcpDestructiveAnnotations,
UserClientOptional: true,
Handler: func(ctx context.Context, deps Deps, args WorkspaceEditFileArgs) (WorkspaceEditFilesResponse, error) {
conn, err := newAgentConn(ctx, deps.coderClient, args.Workspace)
conn, err := openAgentConn(ctx, deps, args.Workspace)
if err != nil {
return WorkspaceEditFilesResponse{}, err
}
@@ -1800,7 +1860,7 @@ var WorkspaceEditFiles = Tool[WorkspaceEditFilesArgs, WorkspaceEditFilesResponse
MCPAnnotations: mcpDestructiveAnnotations,
UserClientOptional: true,
Handler: func(ctx context.Context, deps Deps, args WorkspaceEditFilesArgs) (WorkspaceEditFilesResponse, error) {
conn, err := newAgentConn(ctx, deps.coderClient, args.Workspace)
conn, err := openAgentConn(ctx, deps, args.Workspace)
if err != nil {
return WorkspaceEditFilesResponse{}, err
}
@@ -2245,41 +2305,6 @@ func NormalizeWorkspaceInput(input string) string {
return normalized
}
// newAgentConn returns a connection to the agent specified by the workspace,
// which must be in the format [owner/]workspace[.agent].
func newAgentConn(ctx context.Context, client *codersdk.Client, workspace string) (workspacesdk.AgentConn, error) {
workspaceName := NormalizeWorkspaceInput(workspace)
_, workspaceAgent, err := findWorkspaceAndAgent(ctx, client, workspaceName)
if err != nil {
return nil, xerrors.Errorf("failed to find workspace: %w", err)
}
// Wait for agent to be ready.
if err := cliui.Agent(ctx, io.Discard, workspaceAgent.ID, cliui.AgentOptions{
FetchInterval: 0,
Fetch: client.WorkspaceAgent,
FetchLogs: client.WorkspaceAgentLogsAfter,
Wait: true, // Always wait for startup scripts
}); err != nil {
return nil, xerrors.Errorf("agent not ready: %w", err)
}
wsClient := workspacesdk.New(client)
conn, err := wsClient.DialAgent(ctx, workspaceAgent.ID, &workspacesdk.DialAgentOptions{
BlockEndpoints: false,
})
if err != nil {
return nil, xerrors.Errorf("failed to dial agent: %w", err)
}
if !conn.AwaitReachable(ctx) {
conn.Close()
return nil, xerrors.New("agent connection not reachable")
}
return conn, nil
}
const workspaceDescription = "The workspace ID or name in the format [owner/]workspace. If an owner is not specified, the authenticated user is used."
const workspaceAgentDescription = "The workspace name in the format [owner/]workspace[.agent]. If an owner is not specified, the authenticated user is used."
+126
View File
@@ -20,6 +20,7 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/goleak"
"golang.org/x/xerrors"
agentapi "github.com/coder/agentapi-sdk-go"
"github.com/coder/aisdk-go"
@@ -61,6 +62,22 @@ func setupWorkspaceForAgent(t *testing.T, opts *coderdtest.Options) (*codersdk.C
return userClient, r.Workspace, r.AgentToken
}
type recordingAgentConnFunc struct {
conn workspacesdk.AgentConn
err error
agentID uuid.UUID
calls int
}
func (d *recordingAgentConnFunc) AgentConn(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
d.calls++
d.agentID = agentID
if d.err != nil {
return nil, nil, d.err
}
return d.conn, nil, nil
}
// These tests are dependent on the state of the coder server.
// Running them in parallel is prone to racy behavior.
// nolint:tparallel,paralleltest
@@ -597,6 +614,115 @@ func TestTools(t *testing.T) {
}, res.Contents)
})
t.Run("WorkspaceToolsUseInjectedAgentConnFunc", func(t *testing.T) {
t.Parallel()
client, workspace, agentToken := setupWorkspaceForAgent(t, nil)
_ = agenttest.New(t, client.URL, agentToken)
coderdtest.NewWorkspaceAgentWaiter(t, client, workspace.ID).Wait()
ws, err := client.Workspace(t.Context(), workspace.ID)
require.NoError(t, err)
require.NotEmpty(t, ws.LatestBuild.Resources)
require.NotEmpty(t, ws.LatestBuild.Resources[0].Agents)
agentID := ws.LatestBuild.Resources[0].Agents[0].ID
sentinelErr := xerrors.New("injected agent connection function used")
tests := []struct {
name string
run func(t *testing.T, tb toolsdk.Deps) error
}{
{
name: "WorkspaceLS",
run: func(t *testing.T, tb toolsdk.Deps) error {
_, err := testTool(t, toolsdk.WorkspaceLS, tb, toolsdk.WorkspaceLSArgs{
Workspace: workspace.Name,
Path: "/tmp",
})
return err
},
},
{
name: "WorkspaceReadFile",
run: func(t *testing.T, tb toolsdk.Deps) error {
_, err := testTool(t, toolsdk.WorkspaceReadFile, tb, toolsdk.WorkspaceReadFileArgs{
Workspace: workspace.Name,
Path: "/tmp/file",
})
return err
},
},
{
name: "WorkspaceWriteFile",
run: func(t *testing.T, tb toolsdk.Deps) error {
_, err := testTool(t, toolsdk.WorkspaceWriteFile, tb, toolsdk.WorkspaceWriteFileArgs{
Workspace: workspace.Name,
Path: "/tmp/file",
Content: []byte("hello from agent connection function"),
})
return err
},
},
{
name: "WorkspaceEditFile",
run: func(t *testing.T, tb toolsdk.Deps) error {
_, err := testTool(t, toolsdk.WorkspaceEditFile, tb, toolsdk.WorkspaceEditFileArgs{
Workspace: workspace.Name,
Path: "/tmp/file",
Edits: []workspacesdk.FileEdit{{
Search: "hello",
Replace: "goodbye",
}},
})
return err
},
},
{
name: "WorkspaceEditFiles",
run: func(t *testing.T, tb toolsdk.Deps) error {
_, err := testTool(t, toolsdk.WorkspaceEditFiles, tb, toolsdk.WorkspaceEditFilesArgs{
Workspace: workspace.Name,
Files: []workspacesdk.FileEdits{{
Path: "/tmp/file",
Edits: []workspacesdk.FileEdit{{
Search: "hello",
Replace: "goodbye",
}},
}},
})
return err
},
},
{
name: "WorkspaceBash",
run: func(t *testing.T, tb toolsdk.Deps) error {
_, err := testTool(t, toolsdk.WorkspaceBash, tb, toolsdk.WorkspaceBashArgs{
Workspace: workspace.Name,
Command: "echo hello",
})
return err
},
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
agentConnFn := &recordingAgentConnFunc{err: sentinelErr}
tb, err := toolsdk.NewDeps(client, toolsdk.WithAgentConnFunc(agentConnFn.AgentConn))
require.NoError(t, err)
err = tt.run(t, tb)
require.ErrorIs(t, err, sentinelErr)
require.ErrorContains(t, err, "failed to dial agent")
require.Equal(t, 1, agentConnFn.calls)
require.Equal(t, agentID, agentConnFn.agentID)
})
}
})
t.Run("WorkspaceReadFile", func(t *testing.T) {
t.Parallel()