fix: accept null tool arguments and bound HTTP resource use

Review follow-ups on the SDK migration.

An "arguments": null is what clients send for parameterless tools like get_me,
and what mcp-go accepted by returning a nil map. The new adapter rejected it
with InvalidParams, which broke those calls outright.

The /mcp endpoint took unlimited request bodies and never expired idle
sessions, so a peer that goes away without DELETE kept its session for the
process lifetime. Both are reachable before any token check, so neither can
stay unbounded; the body cap sits above the SDK default to leave room for the
base64 content create_or_update_file accepts.

Required() smuggled a bool through the property schema map and deleted it
again, colliding with the JSON Schema keyword of the same name. It now sets a
field on Property, so an object property can carry its own required list.

The tool contract fixture cost a manual regeneration step and four
hand-maintained counts on every tool change, and a snapshot freezes defects
rather than reporting them. Property assertions cover the same surface and
reject a duplicate tool name, a readOnlyHint that disagrees with the register
call, and a default that contradicts its own type or enum.

Co-Authored-By: Claude (Opus 5) <noreply@anthropic.com>
This commit is contained in:
silverwind
2026-08-02 19:34:40 +02:00
parent 80c8b25d6e
commit 0dc9868e2e
11 changed files with 316 additions and 3225 deletions
+44 -14
View File
@@ -1,20 +1,50 @@
package annotation
import "testing"
import (
"encoding/json"
"maps"
"testing"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
// The hints are what clients use to decide whether a tool needs confirmation, so
// assert the encoded form: an omitted readOnlyHint reads as false either way, but
// only the explicit form survives a client that checks for the key.
func TestAnnotations(t *testing.T) {
readOnly := ReadOnly("Read")
if readOnly.Title != "Read" || !readOnly.ReadOnlyHint || readOnly.DestructiveHint != nil {
t.Errorf("ReadOnly() = %#v", readOnly)
}
write := Write("Write")
if write.Title != "Write" || write.ReadOnlyHint || write.DestructiveHint != nil {
t.Errorf("Write() = %#v", write)
}
destructive := Destructive("Delete")
if destructive.Title != "Delete" || destructive.ReadOnlyHint || destructive.DestructiveHint == nil || !*destructive.DestructiveHint {
t.Errorf("Destructive() = %#v", destructive)
for _, test := range []struct {
name string
annotations *mcp.ToolAnnotations
want map[string]any
}{
{
name: "ReadOnly",
annotations: ReadOnly("Read"),
want: map[string]any{"title": "Read", "readOnlyHint": true, "idempotentHint": false},
},
{
name: "Write",
annotations: Write("Write"),
want: map[string]any{"title": "Write", "readOnlyHint": false, "idempotentHint": false},
},
{
name: "Destructive",
annotations: Destructive("Delete"),
want: map[string]any{"title": "Delete", "readOnlyHint": false, "idempotentHint": false, "destructiveHint": true},
},
} {
t.Run(test.name, func(t *testing.T) {
encoded, err := json.Marshal(test.annotations)
if err != nil {
t.Fatalf("json.Marshal() error = %v", err)
}
var got map[string]any
if err := json.Unmarshal(encoded, &got); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
if !maps.Equal(got, test.want) {
t.Errorf("annotations = %s, want %v", encoded, test.want)
}
})
}
}
+24 -28
View File
@@ -10,7 +10,7 @@ type Property struct {
}
// PropertyOption configures one property in a tool's input schema.
type PropertyOption func(map[string]any)
type PropertyOption func(*Property)
// NewDefinition builds a tool definition without enabling SDK-side validation.
func NewDefinition(name, description string, annotations *mcp.ToolAnnotations, properties ...Property) *mcp.Tool {
@@ -40,71 +40,67 @@ func NewDefinition(name, description string, annotations *mcp.ToolAnnotations, p
}
func String(name string, options ...PropertyOption) Property {
return newProperty(name, "string", false, options...)
return newProperty(name, map[string]any{"type": "string"}, options...)
}
func Number(name string, options ...PropertyOption) Property {
return newProperty(name, "number", false, options...)
return newProperty(name, map[string]any{"type": "number"}, options...)
}
func Boolean(name string, options ...PropertyOption) Property {
return newProperty(name, "boolean", false, options...)
return newProperty(name, map[string]any{"type": "boolean"}, options...)
}
func Array(name string, options ...PropertyOption) Property {
return newProperty(name, "array", false, options...)
return newProperty(name, map[string]any{"type": "array"}, options...)
}
func Object(name string, options ...PropertyOption) Property {
return newProperty(name, "object", true, options...)
return newProperty(name, map[string]any{"type": "object", "properties": map[string]any{}}, options...)
}
func newProperty(name, propertyType string, object bool, options ...PropertyOption) Property {
schema := map[string]any{"type": propertyType}
if object {
schema["properties"] = map[string]any{}
}
func newProperty(name string, schema map[string]any, options ...PropertyOption) Property {
property := Property{name: name, schema: schema}
for _, option := range options {
option(schema)
option(&property)
}
required, _ := schema["required"].(bool)
delete(schema, "required")
return Property{name: name, schema: schema, required: required}
return property
}
// Required marks the property as required on the parent schema. It is not a
// property-level keyword, so it never touches the emitted property schema.
func Required() PropertyOption {
return func(schema map[string]any) {
schema["required"] = true
return func(property *Property) {
property.required = true
}
}
func Description(description string) PropertyOption {
return func(schema map[string]any) {
schema["description"] = description
return func(property *Property) {
property.schema["description"] = description
}
}
func Enum(values ...string) PropertyOption {
return func(schema map[string]any) {
schema["enum"] = values
return func(property *Property) {
property.schema["enum"] = values
}
}
func Default(value any) PropertyOption {
return func(schema map[string]any) {
schema["default"] = value
return func(property *Property) {
property.schema["default"] = value
}
}
func Minimum(value float64) PropertyOption {
return func(schema map[string]any) {
schema["minimum"] = value
return func(property *Property) {
property.schema["minimum"] = value
}
}
func Items(schema any) PropertyOption {
return func(propertySchema map[string]any) {
propertySchema["items"] = schema
return func(property *Property) {
property.schema["items"] = schema
}
}
-9
View File
@@ -1,7 +1,6 @@
package tool
import (
"encoding/json"
"reflect"
"testing"
@@ -58,14 +57,6 @@ func TestNewDefinition(t *testing.T) {
if !reflect.DeepEqual(definition.InputSchema, want) {
t.Errorf("InputSchema = %#v, want %#v", definition.InputSchema, want)
}
data, err := json.Marshal(definition)
if err != nil {
t.Fatalf("json.Marshal() error = %v", err)
}
if !json.Valid(data) {
t.Fatalf("json.Marshal() returned invalid JSON: %s", data)
}
}
func TestNewDefinitionWithoutRequiredProperties(t *testing.T) {
+53 -48
View File
@@ -10,19 +10,23 @@ import (
"github.com/modelcontextprotocol/go-sdk/mcp"
)
func callTool(handler Handler, arguments json.RawMessage) (*mcp.CallToolResult, error) {
serverTool := ServerTool{Tool: &mcp.Tool{Name: "example"}, Handler: handler}
return serverTool.MCPHandler()(context.Background(), &mcp.CallToolRequest{
Params: &mcp.CallToolParamsRaw{Arguments: arguments},
})
}
func captureArguments(into *map[string]any) Handler {
return func(_ context.Context, arguments map[string]any) (*mcp.CallToolResult, error) {
*into = arguments
return &mcp.CallToolResult{}, nil
}
}
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}}`)},
})
result, err := callTool(captureArguments(&got), json.RawMessage(`{"count":2,"nested":{"enabled":true}}`))
if err != nil {
t.Fatalf("MCPHandler() error = %v", err)
}
@@ -36,18 +40,13 @@ func TestMCPHandler(t *testing.T) {
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
},
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},
})
for _, arguments := range []json.RawMessage{json.RawMessage(`[]`), json.RawMessage(`"text"`), json.RawMessage(`{"broken"`)} {
_, err := callTool(handler, arguments)
assertProtocolErrorCode(t, err, jsonrpc.CodeInvalidParams)
}
if called {
@@ -55,37 +54,43 @@ func TestMCPHandlerRejectsInvalidArguments(t *testing.T) {
}
}
// Tools without parameters are callable with an omitted or null "arguments",
// which is what clients send and what mcp-go accepted before the SDK migration.
func TestMCPHandlerAcceptsAbsentArguments(t *testing.T) {
for _, arguments := range []json.RawMessage{nil, json.RawMessage(`null`)} {
var got map[string]any
if _, err := callTool(captureArguments(&got), arguments); err != nil {
t.Fatalf("MCPHandler() with arguments %s error = %v", arguments, err)
}
if got == nil || len(got) != 0 {
t.Errorf("arguments = %#v, want an empty map", got)
}
}
}
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) {
for _, test := range []struct {
name string
handler Handler
}{
{
name: "handler error",
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
},
{
name: "panic",
handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
panic("failed")
},
}
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)
}
})
},
} {
t.Run(test.name, func(t *testing.T) {
_, err := callTool(test.handler, nil)
assertProtocolErrorCode(t, err, jsonrpc.CodeInternalError)
})
}
}
func assertProtocolErrorCode(t *testing.T, err error, want int64) {
+9 -26
View File
@@ -1,7 +1,6 @@
package tool
import (
"bytes"
"context"
"encoding/json"
"errors"
@@ -89,30 +88,18 @@ func (t *Tool) Tools() []ServerTool {
// MCPHandler adapts a project handler to the official SDK's low-level handler.
func (s ServerTool) MCPHandler() mcp.ToolHandler {
return func(ctx context.Context, req *mcp.CallToolRequest) (result *mcp.CallToolResult, err error) {
name := ""
if s.Tool != nil {
name = s.Tool.Name
}
defer func() {
if recovered := recover(); recovered != nil {
panicErr := fmt.Errorf("panic recovered in %s tool handler: %v", name, recovered)
panicErr := fmt.Errorf("panic recovered in %s tool handler: %v", s.Tool.Name, recovered)
log.Errorf("%s", panicErr)
result = nil
err = &jsonrpc.Error{Code: jsonrpc.CodeInternalError, Message: panicErr.Error()}
err = internalError(panicErr)
}
}()
if req == nil || req.Params == nil {
return nil, invalidParamsError("missing tool call parameters")
}
arguments, err := decodeArguments(req.Params.Arguments)
if err != nil {
return nil, err
}
if s.Handler == nil {
return nil, internalError(fmt.Errorf("tool %q has no handler", name))
}
result, err = s.Handler(ctx, arguments)
if err != nil {
@@ -127,25 +114,21 @@ func (s ServerTool) MCPHandler() mcp.ToolHandler {
}
func decodeArguments(raw json.RawMessage) (map[string]any, error) {
trimmed := bytes.TrimSpace(raw)
if len(trimmed) == 0 {
// An omitted and a null "arguments" both mean the tool was called without any.
if len(raw) == 0 || string(raw) == "null" {
return map[string]any{}, nil
}
if bytes.Equal(trimmed, []byte("null")) {
return nil, invalidParamsError("tool arguments must be an object")
}
var arguments map[string]any
if err := json.Unmarshal(trimmed, &arguments); err != nil {
return nil, invalidParamsError(fmt.Sprintf("invalid tool arguments: %v", err))
if err := json.Unmarshal(raw, &arguments); err != nil {
return nil, &jsonrpc.Error{
Code: jsonrpc.CodeInvalidParams,
Message: fmt.Sprintf("invalid tool arguments: %v", err),
}
}
return arguments, nil
}
func invalidParamsError(message string) error {
return &jsonrpc.Error{Code: jsonrpc.CodeInvalidParams, Message: message}
}
func internalError(err error) error {
return &jsonrpc.Error{Code: jsonrpc.CodeInternalError, Message: err.Error()}
}