Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 36 additions & 0 deletions cli/exp_mcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -731,6 +731,7 @@ func (s *mcpServer) startServer(ctx context.Context, inv *serpent.Invocation, in
}

// Register tools based on the allowlist. Zero length means allow everything.
registeredTools := make(map[string]bool, len(toolsdk.All))
for _, tool := range toolsdk.All {
// Skip if not allowed.
if len(allowedTools) > 0 && !slices.ContainsFunc(allowedTools, func(t string) bool {
Expand All @@ -752,6 +753,18 @@ func (s *mcpServer) startServer(ctx context.Context, inv *serpent.Invocation, in
}

mcpSrv.AddTools(mcpFromSDK(tool, toolDeps))
registeredTools[tool.Tool.Name] = true
}

// Skip prompts whose referenced tools are unavailable so clients are
// not offered workflows they cannot run.
for _, prompt := range toolsdk.AllPrompts {
if slices.ContainsFunc(prompt.RequiredTools, func(name string) bool {
return !registeredTools[name]
}) {
continue
}
mcpSrv.AddPrompts(mcpPromptFromSDK(prompt))
}

srv := server.NewStdioServer(mcpSrv)
Expand Down Expand Up @@ -1026,3 +1039,26 @@ func mcpFromSDK(sdkTool toolsdk.GenericTool, tb toolsdk.Deps) server.ServerTool
},
}
}

func mcpPromptFromSDK(sdkPrompt toolsdk.Prompt) server.ServerPrompt {
opts := []mcp.PromptOption{mcp.WithPromptDescription(sdkPrompt.Description)}
for _, arg := range sdkPrompt.Arguments {
argOpts := []mcp.ArgumentOption{mcp.ArgumentDescription(arg.Description)}
if arg.Required {
argOpts = append(argOpts, mcp.RequiredArgument())
}
opts = append(opts, mcp.WithArgument(arg.Name, argOpts...))
}
return server.ServerPrompt{
Prompt: mcp.NewPrompt(sdkPrompt.Name, opts...),
Handler: func(_ context.Context, request mcp.GetPromptRequest) (*mcp.GetPromptResult, error) {
text, err := sdkPrompt.Render(request.Params.Arguments)
if err != nil {
return nil, err
}
return mcp.NewGetPromptResult(sdkPrompt.Description, []mcp.PromptMessage{
mcp.NewPromptMessage(mcp.RoleUser, mcp.NewTextContent(text)),
}), nil
},
}
}
131 changes: 131 additions & 0 deletions cli/exp_mcp_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ import (
"github.com/coder/coder/v2/coderd/coderdtest"
"github.com/coder/coder/v2/coderd/httpapi"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/codersdk/toolsdk"
"github.com/coder/coder/v2/testutil"
"github.com/coder/coder/v2/testutil/expecter"
)
Expand Down Expand Up @@ -101,6 +102,20 @@ func TestExpMcpServer(t *testing.T) {
assert.True(t, *annotations.IdempotentHint)
assert.False(t, *annotations.OpenWorldHint)

// Prompts reference chat tools, which are excluded by this
// allowlist, so none may be advertised. With no prompts
// registered the server rejects prompts/list entirely.
stdin.WriteLine(`{"jsonrpc":"2.0","id":5,"method":"prompts/list"}`)
promptsOutput := stdout.ReadLine(ctx)
var promptsResponse struct {
Error *struct {
Code int `json:"code"`
} `json:"error"`
}
err = json.Unmarshal([]byte(promptsOutput), &promptsResponse)
require.NoError(t, err)
require.NotNil(t, promptsResponse.Error, "prompts/list should fail when no prompts are registered")

// Call the tool and ensure it works.
toolPayload := `{"jsonrpc":"2.0","id":3,"method":"tools/call", "params": {"name": "coder_get_authenticated_user", "arguments": {}}}`
stdin.WriteLine(toolPayload)
Expand All @@ -115,6 +130,122 @@ func TestExpMcpServer(t *testing.T) {
<-cmdDone
})

t.Run("PromptsPartialAllowlist", func(t *testing.T) {
t.Parallel()

ctx := testutil.Context(t, testutil.WaitShort)
logger := testutil.Logger(t)
cancelCtx, cancel := context.WithCancel(ctx)
t.Cleanup(cancel)

client := coderdtest.New(t, nil)
_ = coderdtest.CreateFirstUser(t, client)
// The model-list tool is an optional suggestion in the delegate
// workflow, so its absence must not suppress the prompt.
inv, root := clitest.New(t, "exp", "mcp", "server",
"--allowed-tools=coder_create_chat,coder_get_chat,coder_get_chat_messages,coder_send_chat_message")
inv = inv.WithContext(cancelCtx)

var stdout *expecter.Expecter
stdout, inv.Stdout = expecter.NewPiped(t)
stdin := testutil.NewWriterAttachedToInvocation(t, logger.Named("stdin"), inv)
clitest.SetupConfig(t, client, root)

cmdDone := make(chan struct{})
go func() {
defer close(cmdDone)
err := inv.Run()
assert.NoError(t, err)
}()

stdin.WriteLine(`{"jsonrpc":"2.0","id":1,"method":"prompts/list"}`)
output := stdout.ReadLine(ctx)
cancel()
<-cmdDone

var listResponse struct {
Result struct {
Prompts []struct {
Name string `json:"name"`
} `json:"prompts"`
} `json:"result"`
}
err := json.Unmarshal([]byte(output), &listResponse)
require.NoError(t, err)
foundPrompts := make([]string, 0, len(listResponse.Result.Prompts))
for _, prompt := range listResponse.Result.Prompts {
foundPrompts = append(foundPrompts, prompt.Name)
}
require.Contains(t, foundPrompts, toolsdk.PromptNameAgentsDelegate)
require.Contains(t, foundPrompts, toolsdk.PromptNameAgentsCheck)
})

t.Run("Prompts", func(t *testing.T) {
t.Parallel()

ctx := testutil.Context(t, testutil.WaitShort)
logger := testutil.Logger(t)
cancelCtx, cancel := context.WithCancel(ctx)
t.Cleanup(cancel)

client := coderdtest.New(t, nil)
_ = coderdtest.CreateFirstUser(t, client)
inv, root := clitest.New(t, "exp", "mcp", "server")
inv = inv.WithContext(cancelCtx)

var stdout *expecter.Expecter
stdout, inv.Stdout = expecter.NewPiped(t)
stdin := testutil.NewWriterAttachedToInvocation(t, logger.Named("stdin"), inv)
clitest.SetupConfig(t, client, root)

cmdDone := make(chan struct{})
go func() {
defer close(cmdDone)
err := inv.Run()
assert.NoError(t, err)
}()

stdin.WriteLine(`{"jsonrpc":"2.0","id":1,"method":"prompts/list"}`)
output := stdout.ReadLine(ctx)
var listResponse struct {
Result struct {
Prompts []struct {
Name string `json:"name"`
} `json:"prompts"`
} `json:"result"`
}
err := json.Unmarshal([]byte(output), &listResponse)
require.NoError(t, err)
foundPrompts := make([]string, 0, len(listResponse.Result.Prompts))
for _, prompt := range listResponse.Result.Prompts {
foundPrompts = append(foundPrompts, prompt.Name)
}
for _, prompt := range toolsdk.AllPrompts {
require.Contains(t, foundPrompts, prompt.Name)
}

stdin.WriteLine(`{"jsonrpc":"2.0","id":2,"method":"prompts/get","params":{"name":"coder_agents_delegate","arguments":{"task":"Fix the flaky test."}}}`)
output = stdout.ReadLine(ctx)
cancel()
<-cmdDone

var getResponse struct {
Result struct {
Messages []struct {
Role string `json:"role"`
Content struct {
Text string `json:"text"`
} `json:"content"`
} `json:"messages"`
} `json:"result"`
}
err = json.Unmarshal([]byte(output), &getResponse)
require.NoError(t, err)
require.Len(t, getResponse.Result.Messages, 1)
require.Equal(t, "user", getResponse.Result.Messages[0].Role)
require.Contains(t, getResponse.Result.Messages[0].Content.Text, "Fix the flaky test.")
})

t.Run("OK", func(t *testing.T) {
t.Parallel()

Expand Down
30 changes: 30 additions & 0 deletions coderd/mcp/mcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,13 @@ func (s *Server) RegisterTools(client *codersdk.Client, opts ...func(*toolsdk.De
return nil
}

// RegisterPrompts registers all MCP prompt templates with the server.
func (s *Server) RegisterPrompts() {
for _, prompt := range toolsdk.AllPrompts {
s.mcpServer.AddPrompts(mcpPromptFromSDK(prompt))
}
}

// ChatGPT tools are the search and fetch tools as defined in https://platform.openai.com/docs/mcp.
// We do not expose any extra ones because ChatGPT has an undocumented "Safety Scan" feature.
// In my experiments, if I included extra tools in the MCP server, ChatGPT would often - but not always -
Expand Down Expand Up @@ -161,6 +168,29 @@ func mcpFromSDK(sdkTool toolsdk.GenericTool, tb toolsdk.Deps) server.ServerTool
}
}

func mcpPromptFromSDK(sdkPrompt toolsdk.Prompt) server.ServerPrompt {
opts := []mcp.PromptOption{mcp.WithPromptDescription(sdkPrompt.Description)}
for _, arg := range sdkPrompt.Arguments {
argOpts := []mcp.ArgumentOption{mcp.ArgumentDescription(arg.Description)}
if arg.Required {
argOpts = append(argOpts, mcp.RequiredArgument())
}
opts = append(opts, mcp.WithArgument(arg.Name, argOpts...))
}
return server.ServerPrompt{
Prompt: mcp.NewPrompt(sdkPrompt.Name, opts...),
Handler: func(_ context.Context, request mcp.GetPromptRequest) (*mcp.GetPromptResult, error) {
text, err := sdkPrompt.Render(request.Params.Arguments)
if err != nil {
return nil, err
}
return mcp.NewGetPromptResult(sdkPrompt.Description, []mcp.PromptMessage{
mcp.NewPromptMessage(mcp.RoleUser, mcp.NewTextContent(text)),
}), nil
},
}
}

// mcpLoggerAdapter adapts slog.Logger to the mcp-go util.Logger interface
type mcpLoggerAdapter struct {
logger slog.Logger
Expand Down
29 changes: 29 additions & 0 deletions coderd/mcp/mcp_e2e_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -115,6 +115,35 @@ func TestMCPHTTP_E2E_ClientIntegration(t *testing.T) {

// Check for some basic tools that should be available
assert.Contains(t, foundTools, toolsdk.ToolNameGetAuthenticatedUser, "Should have authenticated user tool")

prompts, err := mcpClient.ListPrompts(ctx, mcp.ListPromptsRequest{})
require.NoError(t, err)
var foundPrompts []string
for _, prompt := range prompts.Prompts {
foundPrompts = append(foundPrompts, prompt.Name)
}
for _, prompt := range toolsdk.AllPrompts {
require.Contains(t, foundPrompts, prompt.Name)
}

promptResult, err := mcpClient.GetPrompt(ctx, mcp.GetPromptRequest{
Params: mcp.GetPromptParams{
Name: toolsdk.PromptNameAgentsDelegate,
Arguments: map[string]string{"task": "Fix the flaky test."},
},
})
require.NoError(t, err)
require.Len(t, promptResult.Messages, 1)
require.Equal(t, mcp.RoleUser, promptResult.Messages[0].Role)
promptText, ok := promptResult.Messages[0].Content.(mcp.TextContent)
require.True(t, ok)
require.Contains(t, promptText.Text, "Fix the flaky test.")
require.Contains(t, promptText.Text, toolsdk.ToolNameCreateChat)

_, err = mcpClient.GetPrompt(ctx, mcp.GetPromptRequest{
Params: mcp.GetPromptParams{Name: toolsdk.PromptNameAgentsDelegate},
})
require.ErrorContains(t, err, "missing required prompt argument: task")
require.NotNil(t, userTool)
require.NotNil(t, writeFileTool)
require.NotNil(t, userTool.Annotations.ReadOnlyHint)
Expand Down
1 change: 1 addition & 0 deletions coderd/mcp_http.go
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,7 @@ func (api *API) mcpHTTPHandler() http.Handler {
if err := mcpServer.RegisterTools(authenticatedClient, toolOpt); err != nil {
api.Logger.Warn(r.Context(), "failed to register MCP tools", slog.Error(err))
}
mcpServer.RegisterPrompts()
case MCPToolsetChatGPT:
if err := mcpServer.RegisterChatGPTTools(authenticatedClient, toolOpt); err != nil {
api.Logger.Warn(r.Context(), "failed to register MCP tools", slog.Error(err))
Expand Down
Loading
Loading