[refactoring][mcp] sdk migration part 1: replace method constants (#61679)

This commit is contained in:
STeve (Xin) Huang
2025-11-22 02:53:56 +00:00
committed by GitHub
parent 2428299b47
commit 86b0474e4d
13 changed files with 114 additions and 61 deletions
+1 -1
View File
@@ -405,7 +405,7 @@ func (r *serverConnWithAutoReconnect) cacheMessageLocked(ctx context.Context, ms
switch m := msg.(type) {
case *mcputils.JSONRPCRequest:
if r.initRequest == nil && m.Method == mcp.MethodInitialize {
if r.initRequest == nil && m.Method == mcputils.MethodInitialize {
r.initRequest = m
r.Logger.DebugContext(ctx, "Cached initialize", "request", m)
}
+8 -9
View File
@@ -25,7 +25,6 @@ import (
"strings"
"github.com/gravitational/trace"
"github.com/mark3labs/mcp-go/mcp"
"github.com/gravitational/teleport"
apidefaults "github.com/gravitational/teleport/api/defaults"
@@ -113,14 +112,14 @@ func eventWithHeader(header http.Header) eventOptionFunc {
}
}
func (a *sessionAuditor) shouldEmitEvent(method mcp.MCPMethod) bool {
func (a *sessionAuditor) shouldEmitEvent(method string) bool {
// Do not record discovery, ping calls.
switch method {
case mcp.MethodPing,
mcp.MethodResourcesList,
mcp.MethodResourcesTemplatesList,
mcp.MethodPromptsList,
mcp.MethodToolsList:
case mcputils.MethodPing,
mcputils.MethodResourcesList,
mcputils.MethodResourcesTemplatesList,
mcputils.MethodPromptsList,
mcputils.MethodToolsList:
return false
default:
return true
@@ -187,7 +186,7 @@ func (a *sessionAuditor) emitNotificationEvent(ctx context.Context, msg *mcputil
AppMetadata: a.makeAppMetadata(),
Message: apievents.MCPJSONRPCMessage{
JSONRPC: msg.JSONRPC,
Method: string(msg.Method),
Method: msg.Method,
Params: msg.Params.GetEventParams(),
},
Status: apievents.Status{
@@ -221,7 +220,7 @@ func (a *sessionAuditor) emitRequestEvent(ctx context.Context, msg *mcputils.JSO
},
Message: apievents.MCPJSONRPCMessage{
JSONRPC: msg.JSONRPC,
Method: string(msg.Method),
Method: msg.Method,
ID: msg.ID.String(),
Params: msg.Params.GetEventParams(),
},
+2 -2
View File
@@ -228,7 +228,7 @@ func (c *requestBuilder) makeToolsCallRequest(toolName string) *mcputils.JSONRPC
return &mcputils.JSONRPCRequest{
JSONRPC: mcp.JSONRPC_VERSION,
ID: c.makeRequestID(),
Method: mcp.MethodToolsCall,
Method: mcputils.MethodToolsCall,
Params: mcputils.JSONRPCParams{
"name": toolName,
},
@@ -239,7 +239,7 @@ func (c *requestBuilder) makeToolsListRequest() *mcputils.JSONRPCRequest {
return &mcputils.JSONRPCRequest{
JSONRPC: mcp.JSONRPC_VERSION,
ID: c.makeRequestID(),
Method: mcp.MethodToolsList,
Method: mcputils.MethodToolsList,
}
}
+1 -1
View File
@@ -241,7 +241,7 @@ func (t *streamableHTTPTransport) handleMCPMessage(r *http.Request) (*http.Respo
case baseMessage.IsRequest():
mcpRequest := baseMessage.MakeRequest()
// Only emit session start if "initialize" succeeded.
if mcpRequest.Method == mcp.MethodInitialize && respErrForAudit == nil {
if mcpRequest.Method == mcputils.MethodInitialize && respErrForAudit == nil {
t.emitStartEvent(t.parentCtx, eventWithHeader(r.Header))
}
t.emitRequestEvent(t.parentCtx, mcpRequest, eventWithError(respErrForAudit), eventWithHeader(r.Header))
+25 -25
View File
@@ -21,10 +21,10 @@ package mcp
import (
"slices"
"github.com/mark3labs/mcp-go/mcp"
"github.com/prometheus/client_golang/prometheus"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/lib/utils/mcputils"
)
var (
@@ -89,17 +89,17 @@ var (
// The list is obtained by searching these in addition to mcp-go:
// - https://github.com/modelcontextprotocol/modelcontextprotocol
// - https://github.com/modelcontextprotocol/typescript-sdk/blob/main/src/server/index.ts
knownNotificationMethods = []mcp.MCPMethod{
knownNotificationMethods = []string{
//nolint:misspell // "cancelled" is "UK" spelling but our linter is set to use US locale
"notifications/cancelled",
"notifications/initialized",
"notifications/message",
"notifications/progress",
mcp.MethodNotificationPromptsListChanged, // notifications/prompts/list_changed
mcp.MethodNotificationResourcesListChanged, // notifications/resources/list_changed
mcp.MethodNotificationResourceUpdated, // notifications/resources/updated
mcp.MethodNotificationToolsListChanged, // notifications/tools/list_changed
"notifications/roots/list_changed",
mcputils.MethodNotificationPromptsListChanged, // notifications/prompts/list_changed
mcputils.MethodNotificationResourcesListChanged, // notifications/resources/list_changed
mcputils.MethodNotificationResourceUpdated, // notifications/resources/updated
mcputils.MethodNotificationToolsListChanged, // notifications/tools/list_changed
mcputils.MethodNotificationRootsListChanged, // notifications/roots/list_changed
}
// knownRequestMethods is a list of known method names for requests.
@@ -107,33 +107,33 @@ var (
// The list is obtained by searching these in addition to mcp-go:
// - https://github.com/modelcontextprotocol/modelcontextprotocol
// - https://github.com/modelcontextprotocol/typescript-sdk/blob/main/src/server/index.ts
knownRequestMethods = []mcp.MCPMethod{
mcp.MethodInitialize, // initialize
mcp.MethodPing, // ping
mcp.MethodResourcesList, // resources/list
mcp.MethodResourcesTemplatesList, // resources/templates/list
mcp.MethodResourcesRead, // resources/read
mcp.MethodPromptsList, // prompts/list
mcp.MethodPromptsGet, // prompts/get
mcp.MethodToolsList, // tools/list
mcp.MethodToolsCall, // tools/call
mcp.MethodSetLogLevel, // logging/setLevel
mcp.MethodElicitationCreate, // elicitation/create
"roots/list",
"sampling/createMessage",
knownRequestMethods = []string{
mcputils.MethodInitialize, // initialize
mcputils.MethodPing, // ping
mcputils.MethodResourcesList, // resources/list
mcputils.MethodResourcesTemplatesList, // resources/templates/list
mcputils.MethodResourcesRead, // resources/read
mcputils.MethodPromptsList, // prompts/list
mcputils.MethodPromptsGet, // prompts/get
mcputils.MethodToolsList, // tools/list
mcputils.MethodToolsCall, // tools/call
mcputils.MethodSetLogLevel, // logging/setLevel
mcputils.MethodElicitationCreate, // elicitation/create
mcputils.MethodListRoots, // roots/list
mcputils.MethodSamplingCreateMessage, // sampling/createMessage
}
)
func reportNotificationMethod(method mcp.MCPMethod) string {
func reportNotificationMethod(method string) string {
if slices.Contains(knownNotificationMethods, method) {
return string(method)
return method
}
return "unknown"
}
func reportRequestMethod(method mcp.MCPMethod) string {
func reportRequestMethod(method string) string {
if slices.Contains(knownRequestMethods, method) {
return string(method)
return method
}
return "unknown"
}
+2 -2
View File
@@ -263,7 +263,7 @@ func (s *sessionHandler) processClientRequestNoAudit(ctx context.Context, req *m
s.idTracker.PushRequest(req)
switch req.Method {
case mcp.MethodToolsCall:
case mcputils.MethodToolsCall:
methodName, _ := req.Params.GetName()
if authErr := s.checkAccessToTool(ctx, methodName); authErr != nil {
return makeToolAccessDeniedResponse(req, authErr), trace.Wrap(authErr)
@@ -277,7 +277,7 @@ func (s *sessionHandler) processServerResponse(ctx context.Context, response *mc
messagesFromServer.WithLabelValues(s.transport, "response", reportRequestMethod(method)).Inc()
switch method {
case mcp.MethodToolsList:
case mcputils.MethodToolsList:
return s.makeToolsCallResponse(ctx, response)
}
return response
+2 -2
View File
@@ -105,7 +105,7 @@ func Test_sessionHandler(t *testing.T) {
requestEvent, ok := event.(*apievents.MCPSessionRequest)
require.True(t, ok)
require.True(t, requestEvent.Success)
require.Equal(t, string(mcp.MethodToolsCall), requestEvent.Message.Method)
require.Equal(t, mcputils.MethodToolsCall, requestEvent.Message.Method)
checkParamsHaveNameField(t, requestEvent.Message.Params, allowedTool)
})
}
@@ -124,7 +124,7 @@ func Test_sessionHandler(t *testing.T) {
requestEvent, ok := event.(*apievents.MCPSessionRequest)
require.True(t, ok)
require.False(t, requestEvent.Success)
require.Equal(t, string(mcp.MethodToolsCall), requestEvent.Message.Method)
require.Equal(t, mcputils.MethodToolsCall, requestEvent.Message.Method)
checkParamsHaveNameField(t, requestEvent.Message.Params, deniedTool)
})
}
+2 -2
View File
@@ -213,7 +213,7 @@ func (h *HTTPReaderWriter) WriteMessage(ctx context.Context, msg mcp.JSONRPCMess
resp, err := h.targetClient.SendRequest(ctx, mcpclienttransport.JSONRPCRequest{
JSONRPC: v.JSONRPC,
ID: v.ID,
Method: string(v.Method),
Method: v.Method,
Params: v.Params,
})
if err != nil {
@@ -225,7 +225,7 @@ func (h *HTTPReaderWriter) WriteMessage(ctx context.Context, msg mcp.JSONRPCMess
return trace.Wrap(h.targetClient.SendNotification(ctx, mcp.JSONRPCNotification{
JSONRPC: v.JSONRPC,
Notification: mcp.Notification{
Method: string(v.Method),
Method: v.Method,
Params: mcp.NotificationParams{
AdditionalFields: v.Params,
},
+3 -3
View File
@@ -31,12 +31,12 @@ import (
// growing infinitely. IDTracker is safe for concurrent use.
type IDTracker struct {
mu sync.Mutex
lruCache *simplelru.LRU[mcp.RequestId, mcp.MCPMethod]
lruCache *simplelru.LRU[mcp.RequestId, string]
}
// NewIDTracker creates a new IDTracker with provided maximum size.
func NewIDTracker(size int) (*IDTracker, error) {
lruCache, err := simplelru.NewLRU[mcp.RequestId, mcp.MCPMethod](size, nil)
lruCache, err := simplelru.NewLRU[mcp.RequestId, string](size, nil)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -58,7 +58,7 @@ func (t *IDTracker) PushRequest(msg *JSONRPCRequest) bool {
}
// PopByID retrieves the tracked information and remove it from the tracker.
func (t *IDTracker) PopByID(id mcp.RequestId) (mcp.MCPMethod, bool) {
func (t *IDTracker) PopByID(id mcp.RequestId) (string, bool) {
if id.IsNil() {
return "", false
}
+6 -6
View File
@@ -42,7 +42,7 @@ func TestIDTracker(t *testing.T) {
t.Run("request tracked", func(t *testing.T) {
require.True(t, tracker.PushRequest(&JSONRPCRequest{
ID: mcp.NewRequestId(0),
Method: mcp.MethodToolsList,
Method: MethodToolsList,
}))
require.Equal(t, 1, tracker.Len())
})
@@ -65,7 +65,7 @@ func TestIDTracker(t *testing.T) {
t.Run("pop tracked id", func(t *testing.T) {
method, ok := tracker.PopByID(mcp.NewRequestId(0))
require.True(t, ok)
require.Equal(t, mcp.MethodToolsList, method)
require.Equal(t, MethodToolsList, method)
require.Empty(t, tracker.Len())
})
@@ -73,14 +73,14 @@ func TestIDTracker(t *testing.T) {
for i := range 20 {
tracker.PushRequest(&JSONRPCRequest{
ID: mcp.NewRequestId(i + 1),
Method: mcp.MethodToolsCall,
Method: MethodToolsCall,
})
require.LessOrEqual(t, tracker.Len(), 10)
}
for i := range 5 {
method, ok := tracker.PopByID(mcp.NewRequestId(20 - i))
require.True(t, ok)
require.Equal(t, mcp.MethodToolsCall, method)
require.Equal(t, MethodToolsCall, method)
}
require.Empty(t, tracker.Len())
})
@@ -93,7 +93,7 @@ func BenchmarkIDTracker(b *testing.B) {
for i := range 100 {
idTracker.PushRequest(&JSONRPCRequest{
ID: mcp.NewRequestId(i),
Method: mcp.MethodToolsList,
Method: MethodToolsList,
})
}
@@ -102,7 +102,7 @@ func BenchmarkIDTracker(b *testing.B) {
for b.Loop() {
idTracker.PushRequest(&JSONRPCRequest{
ID: mcp.NewRequestId(2000),
Method: mcp.MethodToolsList,
Method: MethodToolsList,
})
idTracker.PopByID(mcp.NewRequestId(2000))
}
+59 -5
View File
@@ -68,7 +68,7 @@ type BaseJSONRPCMessage struct {
// ID is the ID for request and response. ID is nil for notification.
ID mcp.RequestId `json:"id"`
// Method is the request or notification method. Method is empty for response.
Method mcp.MCPMethod `json:"method,omitempty"`
Method string `json:"method,omitempty"`
// Params is the params for request and notification.
Params JSONRPCParams `json:"params,omitempty"`
// Result is the response result.
@@ -126,7 +126,7 @@ func (m *BaseJSONRPCMessage) MakeResponse() *JSONRPCResponse {
// https://modelcontextprotocol.io/specification/2025-03-26/basic#notifications
type JSONRPCNotification struct {
JSONRPC string `json:"jsonrpc"`
Method mcp.MCPMethod `json:"method"`
Method string `json:"method"`
Params JSONRPCParams `json:"params,omitempty"`
}
@@ -135,7 +135,7 @@ type JSONRPCNotification struct {
// https://modelcontextprotocol.io/specification/2025-03-26/basic#requests
type JSONRPCRequest struct {
JSONRPC string `json:"jsonrpc"`
Method mcp.MCPMethod `json:"method"`
Method string `json:"method"`
ID mcp.RequestId `json:"id"`
Params JSONRPCParams `json:"params,omitempty"`
}
@@ -154,7 +154,7 @@ type JSONRPCResponse struct {
Error json.RawMessage `json:"error,omitempty"`
}
// GetListToolResult assumes the result is for mcp.MethodToolsList and returns
// GetListToolResult assumes the result is for MethodToolsList and returns
// the corresponding go object.
func (r *JSONRPCResponse) GetListToolResult() (*mcp.ListToolsResult, error) {
var listResult mcp.ListToolsResult
@@ -164,7 +164,7 @@ func (r *JSONRPCResponse) GetListToolResult() (*mcp.ListToolsResult, error) {
return &listResult, nil
}
// GetInitializeResult assumes the result is for mcp.MethodInitialize and
// GetInitializeResult assumes the result is for MethodInitialize and
// returns the corresponding go object.
func (r *JSONRPCResponse) GetInitializeResult() (*mcp.InitializeResult, error) {
var result mcp.InitializeResult
@@ -188,6 +188,60 @@ func unmarshalResponse(rawMessage string) (*JSONRPCResponse, error) {
}
const (
// MethodInitialize initiates connection and negotiates protocol capabilities.
MethodInitialize = "initialize"
// MethodPing verifies connection liveness between client and server.
MethodPing = "ping"
// MethodResourcesList lists all available server resources.
MethodResourcesList = "resources/list"
// MethodResourcesTemplatesList provides URI templates for constructing resource URIs.
MethodResourcesTemplatesList = "resources/templates/list"
// MethodResourcesRead retrieves content of a specific resource by URI.
MethodResourcesRead = "resources/read"
// MethodPromptsList lists all available prompt templates.
MethodPromptsList = "prompts/list"
// MethodPromptsGet retrieves a specific prompt template with filled parameters.
MethodPromptsGet = "prompts/get"
// MethodToolsList lists all available executable tools.
MethodToolsList = "tools/list"
// MethodToolsCall invokes a specific tool with provided parameters.
MethodToolsCall = "tools/call"
// MethodSetLogLevel configures the minimum log level for client
MethodSetLogLevel = "logging/setLevel"
// MethodElicitationCreate requests additional information from the user during interactions.
MethodElicitationCreate = "elicitation/create"
// MethodListRoots requests roots list from the client during interactions.
MethodListRoots = "roots/list"
// MethodSamplingCreateMessage is sent by server to request client to sample messages from LLM.
MethodSamplingCreateMessage = "sampling/createMessage"
// MethodNotificationResourcesListChanged notifies when the list of available resources changes.
MethodNotificationResourcesListChanged = "notifications/resources/list_changed"
// MethodNotificationResourceUpdated notifies when a resource changes.
MethodNotificationResourceUpdated = "notifications/resources/updated"
// MethodNotificationPromptsListChanged notifies when the list of available prompt templates changes.
MethodNotificationPromptsListChanged = "notifications/prompts/list_changed"
// MethodNotificationToolsListChanged notifies when the list of available tools changes.
MethodNotificationToolsListChanged = "notifications/tools/list_changed"
// MethodNotificationRootsListChanged notifies when the list of available roots changes.
MethodNotificationRootsListChanged = "notifications/roots/list_changed"
// MethodNotificationInitialized defines the method used for "initialized"
// notification. This notification is sent by the client after it receives
// the initialize response.
+2 -2
View File
@@ -89,7 +89,7 @@ func TestJSONRPCNotification(t *testing.T) {
m := base.MakeNotification()
require.NotNil(t, m)
assert.Equal(t, mcp.MCPMethod("notifications/message"), m.Method)
assert.Equal(t, "notifications/message", m.Method)
assert.Len(t, base.Params, 3)
outputJSON, err := json.MarshalIndent(m, "", " ")
@@ -106,7 +106,7 @@ func TestJSONRPCRequest(t *testing.T) {
m := base.MakeRequest()
require.NotNil(t, m)
assert.Equal(t, mcp.MethodToolsCall, m.Method)
assert.Equal(t, MethodToolsCall, m.Method)
assert.Equal(t, "int64:2", m.ID.String())
name, ok := m.Params.GetName()
assert.True(t, ok)
+1 -1
View File
@@ -47,7 +47,7 @@ func TestConnectSSEServer(t *testing.T) {
initReq := mcpclienttransport.JSONRPCRequest{
JSONRPC: mcp.JSONRPC_VERSION,
ID: mcp.NewRequestId(int64(1)),
Method: string(mcp.MethodInitialize),
Method: MethodInitialize,
Params: mcp.InitializeParams{
ProtocolVersion: mcp.LATEST_PROTOCOL_VERSION,
ClientInfo: mcp.Implementation{