feat: add MCP server config ID to tool-call message parts (#23522)

This commit is contained in:
Kyle Carberry
2026-03-24 20:29:36 +00:00
committed by GitHub
parent 65a694b537
commit dda985150d
16 changed files with 310 additions and 31 deletions
+27 -1
View File
@@ -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,
+16 -1
View File
@@ -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()