package tool import ( "context" "encoding/json" "errors" "testing" "github.com/modelcontextprotocol/go-sdk/jsonrpc" "github.com/modelcontextprotocol/go-sdk/mcp" ) func TestMCPHandler(t *testing.T) { var got map[string]any serverTool := ServerTool{ Tool: &mcp.Tool{Name: "example"}, Handler: func(_ context.Context, arguments map[string]any) (*mcp.CallToolResult, error) { got = arguments return &mcp.CallToolResult{}, nil }, } result, err := serverTool.MCPHandler()(context.Background(), &mcp.CallToolRequest{ Params: &mcp.CallToolParamsRaw{Arguments: json.RawMessage(`{"count":2,"nested":{"enabled":true}}`)}, }) if err != nil { t.Fatalf("MCPHandler() error = %v", err) } if result == nil { t.Fatal("MCPHandler() result is nil") } if got["count"] != float64(2) { t.Errorf("count type/value = %T(%v), want float64(2)", got["count"], got["count"]) } } func TestMCPHandlerRejectsInvalidArguments(t *testing.T) { called := false serverTool := ServerTool{ Tool: &mcp.Tool{Name: "example"}, Handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) { called = true return &mcp.CallToolResult{}, nil }, } for _, arguments := range []json.RawMessage{json.RawMessage(`[]`), json.RawMessage(`null`), json.RawMessage(`{"broken"`)} { _, err := serverTool.MCPHandler()(context.Background(), &mcp.CallToolRequest{ Params: &mcp.CallToolParamsRaw{Arguments: arguments}, }) assertProtocolErrorCode(t, err, jsonrpc.CodeInvalidParams) } if called { t.Fatal("handler was called with invalid arguments") } } func TestMCPHandlerConvertsErrorsAndRecoversPanics(t *testing.T) { t.Run("handler error", func(t *testing.T) { serverTool := ServerTool{ Tool: &mcp.Tool{Name: "example"}, Handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) { return nil, errors.New("failed") }, } _, err := serverTool.MCPHandler()(context.Background(), &mcp.CallToolRequest{Params: &mcp.CallToolParamsRaw{}}) assertProtocolErrorCode(t, err, jsonrpc.CodeInternalError) }) t.Run("panic", func(t *testing.T) { calls := 0 serverTool := ServerTool{ Tool: &mcp.Tool{Name: "example"}, Handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) { calls++ if calls == 1 { panic("failed") } return &mcp.CallToolResult{}, nil }, } handler := serverTool.MCPHandler() _, err := handler(context.Background(), &mcp.CallToolRequest{Params: &mcp.CallToolParamsRaw{}}) assertProtocolErrorCode(t, err, jsonrpc.CodeInternalError) if _, err := handler(context.Background(), &mcp.CallToolRequest{Params: &mcp.CallToolParamsRaw{}}); err != nil { t.Fatalf("second handler call after panic error = %v", err) } }) } func assertProtocolErrorCode(t *testing.T, err error, want int64) { t.Helper() var protocolErr *jsonrpc.Error if !errors.As(err, &protocolErr) { t.Fatalf("error = %v, want *jsonrpc.Error", err) } if protocolErr.Code != want { t.Errorf("error code = %d, want %d", protocolErr.Code, want) } }