mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add MCP server config ID to tool-call message parts (#23522)
This commit is contained in:
+27
-1
@@ -3017,6 +3017,16 @@ func (p *Server) runChat(
|
||||
defer mcpCleanup()
|
||||
}
|
||||
|
||||
// Build a lookup from tool name to MCP server config ID
|
||||
// so we can annotate persisted parts with the originating
|
||||
// server.
|
||||
toolNameToConfigID := make(map[string]uuid.UUID)
|
||||
for _, t := range mcpTools {
|
||||
if mcp, ok := t.(mcpclient.MCPToolIdentifier); ok {
|
||||
toolNameToConfigID[t.Info().Name] = mcp.MCPServerConfigID()
|
||||
}
|
||||
}
|
||||
|
||||
if instruction != "" {
|
||||
prompt = chatprompt.InsertSystem(prompt, instruction)
|
||||
}
|
||||
@@ -3079,7 +3089,13 @@ func (p *Server) runChat(
|
||||
if len(assistantBlocks) > 0 {
|
||||
sdkParts := make([]codersdk.ChatMessagePart, 0, len(assistantBlocks))
|
||||
for _, block := range assistantBlocks {
|
||||
sdkParts = append(sdkParts, chatprompt.PartFromContent(block))
|
||||
part := chatprompt.PartFromContent(block)
|
||||
if part.ToolName != "" {
|
||||
if configID, ok := toolNameToConfigID[part.ToolName]; ok {
|
||||
part.MCPServerConfigID = uuid.NullUUID{UUID: configID, Valid: true}
|
||||
}
|
||||
}
|
||||
sdkParts = append(sdkParts, part)
|
||||
}
|
||||
finalAssistantText = strings.TrimSpace(contentBlocksToText(sdkParts))
|
||||
var marshalErr error
|
||||
@@ -3092,6 +3108,11 @@ func (p *Server) runChat(
|
||||
toolResultContents := make([]pqtype.NullRawMessage, len(toolResults))
|
||||
for i, tr := range toolResults {
|
||||
trPart := chatprompt.PartFromContent(tr)
|
||||
if trPart.ToolName != "" {
|
||||
if configID, ok := toolNameToConfigID[trPart.ToolName]; ok {
|
||||
trPart.MCPServerConfigID = uuid.NullUUID{UUID: configID, Valid: true}
|
||||
}
|
||||
}
|
||||
var marshalErr error
|
||||
toolResultContents[i], marshalErr = chatprompt.MarshalParts([]codersdk.ChatMessagePart{trPart})
|
||||
if marshalErr != nil {
|
||||
@@ -3463,6 +3484,11 @@ func (p *Server) runChat(
|
||||
role codersdk.ChatMessageRole,
|
||||
part codersdk.ChatMessagePart,
|
||||
) {
|
||||
if part.ToolName != "" {
|
||||
if configID, ok := toolNameToConfigID[part.ToolName]; ok {
|
||||
part.MCPServerConfigID = uuid.NullUUID{UUID: configID, Valid: true}
|
||||
}
|
||||
}
|
||||
p.publishMessagePart(chat.ID, role, part)
|
||||
},
|
||||
Compaction: compactionOptions,
|
||||
|
||||
@@ -195,7 +195,7 @@ func connectOne(
|
||||
}
|
||||
|
||||
tools = append(
|
||||
tools, newMCPTool(cfg.Slug, mcpTool, mcpClient),
|
||||
tools, newMCPTool(cfg.ID, cfg.Slug, mcpTool, mcpClient),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -383,10 +383,17 @@ func redactErrorURL(err error) string {
|
||||
return err.Error()
|
||||
}
|
||||
|
||||
// MCPToolIdentifier is implemented by tools that originate from
|
||||
// an MCP server config and can report the config's database ID.
|
||||
type MCPToolIdentifier interface {
|
||||
MCPServerConfigID() uuid.UUID
|
||||
}
|
||||
|
||||
// mcpToolWrapper adapts a single MCP tool into a
|
||||
// fantasy.AgentTool. It stores the prefixed name for Info() but
|
||||
// strips the prefix when forwarding calls to the remote server.
|
||||
type mcpToolWrapper struct {
|
||||
configID uuid.UUID
|
||||
prefixedName string
|
||||
originalName string
|
||||
description string
|
||||
@@ -396,14 +403,22 @@ type mcpToolWrapper struct {
|
||||
providerOptions fantasy.ProviderOptions
|
||||
}
|
||||
|
||||
// MCPServerConfigID returns the database ID of the MCP server
|
||||
// config that this tool originates from.
|
||||
func (t *mcpToolWrapper) MCPServerConfigID() uuid.UUID {
|
||||
return t.configID
|
||||
}
|
||||
|
||||
// newMCPTool creates an mcpToolWrapper from an mcp.Tool
|
||||
// discovered on a remote server.
|
||||
func newMCPTool(
|
||||
configID uuid.UUID,
|
||||
serverSlug string,
|
||||
tool mcp.Tool,
|
||||
mcpClient *client.Client,
|
||||
) *mcpToolWrapper {
|
||||
return &mcpToolWrapper{
|
||||
configID: configID,
|
||||
prefixedName: serverSlug + toolNameSep + tool.Name,
|
||||
originalName: tool.Name,
|
||||
description: tool.Description,
|
||||
|
||||
@@ -621,6 +621,91 @@ func TestConnectAll_EmptyAccessToken(t *testing.T) {
|
||||
require.NotEmpty(t, tools)
|
||||
}
|
||||
|
||||
// TestConnectAll_MCPToolIdentifier verifies that tools returned
|
||||
// by ConnectAll implement the MCPToolIdentifier interface and
|
||||
// report the correct server config ID.
|
||||
func TestConnectAll_MCPToolIdentifier(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
|
||||
ts := newTestMCPServer(t, echoTool())
|
||||
|
||||
configID := uuid.New()
|
||||
cfg := database.MCPServerConfig{
|
||||
ID: configID,
|
||||
Slug: "id-srv",
|
||||
DisplayName: "ID Server",
|
||||
Url: ts.URL,
|
||||
Transport: "streamable_http",
|
||||
AuthType: "none",
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
tools, cleanup := mcpclient.ConnectAll(ctx, logger, []database.MCPServerConfig{cfg}, nil)
|
||||
t.Cleanup(cleanup)
|
||||
|
||||
require.Len(t, tools, 1)
|
||||
|
||||
// Assert the tool implements MCPToolIdentifier.
|
||||
identifier, ok := tools[0].(mcpclient.MCPToolIdentifier)
|
||||
require.True(t, ok, "tool should implement MCPToolIdentifier")
|
||||
assert.Equal(t, configID, identifier.MCPServerConfigID())
|
||||
}
|
||||
|
||||
// TestConnectAll_MCPToolIdentifier_MultipleServers verifies that
|
||||
// each tool from a different MCP server carries its own config ID.
|
||||
func TestConnectAll_MCPToolIdentifier_MultipleServers(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
|
||||
ts1 := newTestMCPServer(t, echoTool())
|
||||
ts2 := newTestMCPServer(t, greetTool())
|
||||
|
||||
configID1 := uuid.New()
|
||||
configID2 := uuid.New()
|
||||
cfg1 := database.MCPServerConfig{
|
||||
ID: configID1,
|
||||
Slug: "srv-a",
|
||||
DisplayName: "Server A",
|
||||
Url: ts1.URL,
|
||||
Transport: "streamable_http",
|
||||
AuthType: "none",
|
||||
Enabled: true,
|
||||
}
|
||||
cfg2 := database.MCPServerConfig{
|
||||
ID: configID2,
|
||||
Slug: "srv-b",
|
||||
DisplayName: "Server B",
|
||||
Url: ts2.URL,
|
||||
Transport: "streamable_http",
|
||||
AuthType: "none",
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
tools, cleanup := mcpclient.ConnectAll(
|
||||
ctx, logger,
|
||||
[]database.MCPServerConfig{cfg1, cfg2},
|
||||
nil,
|
||||
)
|
||||
t.Cleanup(cleanup)
|
||||
|
||||
require.Len(t, tools, 2)
|
||||
|
||||
// Map tool name to config ID via the MCPToolIdentifier
|
||||
// interface.
|
||||
idByName := make(map[string]uuid.UUID)
|
||||
for _, tool := range tools {
|
||||
identifier, ok := tool.(mcpclient.MCPToolIdentifier)
|
||||
require.True(t, ok, "tool %q should implement MCPToolIdentifier", tool.Info().Name)
|
||||
idByName[tool.Info().Name] = identifier.MCPServerConfigID()
|
||||
}
|
||||
|
||||
assert.Equal(t, configID1, idByName["srv-a__echo"])
|
||||
assert.Equal(t, configID2, idByName["srv-b__greet"])
|
||||
}
|
||||
|
||||
func TestConnectAll_CallToolError(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
|
||||
Reference in New Issue
Block a user