mirror of
https://gitea.com/gitea/gitea-mcp.git
synced 2026-08-03 15:49:23 +02:00
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:
@@ -6,6 +6,8 @@ LDFLAGS := -X "main.Version=$(VERSION)"
|
|||||||
GOLANGCI_LINT_PACKAGE ?= github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.12.2 # renovate: datasource=go
|
GOLANGCI_LINT_PACKAGE ?= github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.12.2 # renovate: datasource=go
|
||||||
GOVULNCHECK_PACKAGE ?= golang.org/x/vuln/cmd/govulncheck@v1.3.0 # renovate: datasource=go
|
GOVULNCHECK_PACKAGE ?= golang.org/x/vuln/cmd/govulncheck@v1.3.0 # renovate: datasource=go
|
||||||
|
|
||||||
|
GOTEST_FLAGS ?= -race -timeout 20m
|
||||||
|
|
||||||
.PHONY: help
|
.PHONY: help
|
||||||
help: ## print this help message
|
help: ## print this help message
|
||||||
@echo "Usage: make [target]"
|
@echo "Usage: make [target]"
|
||||||
@@ -40,7 +42,7 @@ build: ## build the application
|
|||||||
|
|
||||||
.PHONY: test
|
.PHONY: test
|
||||||
test: ## run Go tests
|
test: ## run Go tests
|
||||||
$(GO) test ./...
|
$(GO) test $(GOTEST_FLAGS) ./...
|
||||||
|
|
||||||
.PHONY: air
|
.PHONY: air
|
||||||
air: ## install air for hot reload
|
air: ## install air for hot reload
|
||||||
|
|||||||
+17
-14
@@ -32,6 +32,15 @@ import (
|
|||||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// maxRequestBodyBytes raises the SDK's 4 MiB default, which is too tight for the
|
||||||
|
// base64 file content create_or_update_file accepts.
|
||||||
|
const maxRequestBodyBytes = 32 << 20
|
||||||
|
|
||||||
|
// sessionTimeout expires idle sessions, which the SDK otherwise keeps for the
|
||||||
|
// process lifetime: a client that goes away without DELETE /mcp leaks its
|
||||||
|
// session, and initialize takes no token. Clients re-initialize on the 404.
|
||||||
|
const sessionTimeout = 30 * time.Minute
|
||||||
|
|
||||||
var (
|
var (
|
||||||
mcpServer *mcp.Server
|
mcpServer *mcp.Server
|
||||||
|
|
||||||
@@ -95,24 +104,18 @@ func authTokenMiddleware(next mcp.MethodHandler) mcp.MethodHandler {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func newStreamableHTTPHandler(s *mcp.Server) http.Handler {
|
func newHTTPServer(addr string, s *mcp.Server) *http.Server {
|
||||||
return mcp.NewStreamableHTTPHandler(
|
mux := http.NewServeMux()
|
||||||
|
mux.Handle("/mcp", mcp.NewStreamableHTTPHandler(
|
||||||
func(*http.Request) *mcp.Server { return s },
|
func(*http.Request) *mcp.Server { return s },
|
||||||
&mcp.StreamableHTTPOptions{
|
&mcp.StreamableHTTPOptions{
|
||||||
Logger: log.Slog(),
|
Logger: log.Slog(),
|
||||||
MaxRequestBodyBytes: -1,
|
MaxRequestBodyBytes: maxRequestBodyBytes,
|
||||||
Stateless: false,
|
Stateless: false, // PR 2 switches this on
|
||||||
|
SessionTimeout: sessionTimeout,
|
||||||
},
|
},
|
||||||
)
|
))
|
||||||
}
|
return &http.Server{Addr: addr, Handler: mux}
|
||||||
|
|
||||||
func newHTTPServer(addr string, s *mcp.Server) *http.Server {
|
|
||||||
mux := http.NewServeMux()
|
|
||||||
mux.Handle("/mcp", newStreamableHTTPHandler(s))
|
|
||||||
return &http.Server{
|
|
||||||
Addr: addr,
|
|
||||||
Handler: mux,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func Run() error {
|
func Run() error {
|
||||||
|
|||||||
@@ -1,53 +1,6 @@
|
|||||||
package operation
|
package operation
|
||||||
|
|
||||||
import (
|
import "testing"
|
||||||
"testing"
|
|
||||||
|
|
||||||
"gitea.com/gitea/gitea-mcp/pkg/flag"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestAllToolsHaveDescriptions ensures every registered tool sets a non-empty
|
|
||||||
// Tool.Description, as strict MCP clients reject tools without one.
|
|
||||||
func TestAllToolsHaveDescriptions(t *testing.T) {
|
|
||||||
origRO, origAllow := flag.ReadOnly, flag.AllowedTools
|
|
||||||
t.Cleanup(func() {
|
|
||||||
flag.ReadOnly, flag.AllowedTools = origRO, origAllow
|
|
||||||
})
|
|
||||||
flag.ReadOnly = false
|
|
||||||
flag.AllowedTools = nil
|
|
||||||
|
|
||||||
var missing []string
|
|
||||||
for _, d := range domainTools {
|
|
||||||
for _, st := range d.Tools() {
|
|
||||||
if st.Tool.Description == "" {
|
|
||||||
missing = append(missing, st.Tool.Name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if len(missing) > 0 {
|
|
||||||
t.Errorf("tools missing a description: %v", missing)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestDomainToolsScopesAreUniqueAndNonEmpty ensures every entry registered in
|
|
||||||
// domainTools has a canonical, non-empty scope name and that no two domains
|
|
||||||
// share the same scope (each domain.Tools() call is filtered by exactly one
|
|
||||||
// scope name via flag.AllowedScopes).
|
|
||||||
func TestDomainToolsScopesAreUniqueAndNonEmpty(t *testing.T) {
|
|
||||||
seen := map[string]struct{}{}
|
|
||||||
for _, d := range domainTools {
|
|
||||||
scope := d.Scope()
|
|
||||||
if scope == "" {
|
|
||||||
t.Errorf("domainTools contains a domain with an empty scope")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if _, ok := seen[scope]; ok {
|
|
||||||
t.Errorf("domainTools contains a duplicate scope %q", scope)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
seen[scope] = struct{}{}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseAuthToken(t *testing.T) {
|
func TestParseAuthToken(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"os"
|
"os"
|
||||||
@@ -41,28 +41,39 @@ func exposeAllTools(t *testing.T) {
|
|||||||
flag.Version = testServerVersion
|
flag.Version = testServerVersion
|
||||||
}
|
}
|
||||||
|
|
||||||
func assertVersionToolResult(t *testing.T, result *mcp.CallToolResult) {
|
// registeredToolCount is what the registry exposes under the current flags, so
|
||||||
|
// the transport assertions track tool additions without being edited.
|
||||||
|
func registeredToolCount() int {
|
||||||
|
count := 0
|
||||||
|
for _, domain := range domainTools {
|
||||||
|
count += len(domain.Tools())
|
||||||
|
}
|
||||||
|
return count
|
||||||
|
}
|
||||||
|
|
||||||
|
func textContent(t *testing.T, result *mcp.CallToolResult) string {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
if len(result.Content) != 1 {
|
if len(result.Content) != 1 {
|
||||||
t.Fatalf("version tool content count = %d, want 1", len(result.Content))
|
t.Fatalf("content count = %d, want 1", len(result.Content))
|
||||||
}
|
}
|
||||||
content, ok := result.Content[0].(*mcp.TextContent)
|
content, ok := result.Content[0].(*mcp.TextContent)
|
||||||
if !ok {
|
if !ok {
|
||||||
t.Fatalf("version tool content type = %T, want *mcp.TextContent", result.Content[0])
|
t.Fatalf("content type = %T, want *mcp.TextContent", result.Content[0])
|
||||||
}
|
|
||||||
if !strings.Contains(content.Text, testServerVersion) {
|
|
||||||
t.Errorf("version tool result = %q, want it to contain %q", content.Text, testServerVersion)
|
|
||||||
}
|
}
|
||||||
|
return content.Text
|
||||||
}
|
}
|
||||||
|
|
||||||
func listAndCallVersion(ctx context.Context, t *testing.T, session *mcp.ClientSession) {
|
// listAndCallVersion is the round trip every transport must support. wantText
|
||||||
|
// differs per transport: the stdio subprocess resolves its version from the VCS
|
||||||
|
// build info (main.go:14), so only the in-process servers have a known one.
|
||||||
|
func listAndCallVersion(ctx context.Context, t *testing.T, session *mcp.ClientSession, wantText string) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
result, err := session.ListTools(ctx, nil)
|
result, err := session.ListTools(ctx, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("ListTools() error = %v", err)
|
t.Fatalf("ListTools() error = %v", err)
|
||||||
}
|
}
|
||||||
if len(result.Tools) != 54 {
|
if want := registeredToolCount(); len(result.Tools) != want {
|
||||||
t.Fatalf("ListTools() count = %d, want 54", len(result.Tools))
|
t.Fatalf("ListTools() count = %d, want %d", len(result.Tools), want)
|
||||||
}
|
}
|
||||||
callResult, err := session.CallTool(ctx, &mcp.CallToolParams{
|
callResult, err := session.CallTool(ctx, &mcp.CallToolParams{
|
||||||
Name: "get_gitea_mcp_server_version",
|
Name: "get_gitea_mcp_server_version",
|
||||||
@@ -70,7 +81,9 @@ func listAndCallVersion(ctx context.Context, t *testing.T, session *mcp.ClientSe
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("CallTool() error = %v", err)
|
t.Fatalf("CallTool() error = %v", err)
|
||||||
}
|
}
|
||||||
assertVersionToolResult(t, callResult)
|
if got := textContent(t, callResult); !strings.Contains(got, wantText) {
|
||||||
|
t.Errorf("version tool result = %q, want it to contain %q", got, wantText)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestOfficialSDKInMemory(t *testing.T) {
|
func TestOfficialSDKInMemory(t *testing.T) {
|
||||||
@@ -94,7 +107,7 @@ func TestOfficialSDKInMemory(t *testing.T) {
|
|||||||
if got := session.InitializeResult().ProtocolVersion; got != "2026-07-28" {
|
if got := session.InitializeResult().ProtocolVersion; got != "2026-07-28" {
|
||||||
t.Errorf("protocol version = %q, want %q", got, "2026-07-28")
|
t.Errorf("protocol version = %q, want %q", got, "2026-07-28")
|
||||||
}
|
}
|
||||||
listAndCallVersion(ctx, t, session)
|
listAndCallVersion(ctx, t, session, testServerVersion)
|
||||||
if err := session.Close(); err != nil {
|
if err := session.Close(); err != nil {
|
||||||
t.Fatalf("Close() error = %v", err)
|
t.Fatalf("Close() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -132,7 +145,7 @@ func TestStreamableHTTPStateful(t *testing.T) {
|
|||||||
if got := session.InitializeResult().ProtocolVersion; got != "2025-11-25" {
|
if got := session.InitializeResult().ProtocolVersion; got != "2025-11-25" {
|
||||||
t.Errorf("protocol version = %q, want %q", got, "2025-11-25")
|
t.Errorf("protocol version = %q, want %q", got, "2025-11-25")
|
||||||
}
|
}
|
||||||
listAndCallVersion(ctx, t, session)
|
listAndCallVersion(ctx, t, session, testServerVersion)
|
||||||
|
|
||||||
response, err := httpTestServer.Client().Get(httpTestServer.URL + "/not-mcp")
|
response, err := httpTestServer.Client().Get(httpTestServer.URL + "/not-mcp")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -144,25 +157,47 @@ func TestStreamableHTTPStateful(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestStreamableHTTPAllowsLegacyLargeBodies(t *testing.T) {
|
// spaceReader yields an endless run of spaces, so oversized bodies can be sent
|
||||||
|
// without allocating them.
|
||||||
|
type spaceReader struct{}
|
||||||
|
|
||||||
|
func (spaceReader) Read(p []byte) (int, error) {
|
||||||
|
for index := range p {
|
||||||
|
p[index] = ' '
|
||||||
|
}
|
||||||
|
return len(p), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStreamableHTTPRequestBodyLimit(t *testing.T) {
|
||||||
server := newMCPServer(testServerVersion)
|
server := newMCPServer(testServerVersion)
|
||||||
httpTestServer := httptest.NewServer(newHTTPServer("", server).Handler)
|
httpTestServer := httptest.NewServer(newHTTPServer("", server).Handler)
|
||||||
defer httpTestServer.Close()
|
defer httpTestServer.Close()
|
||||||
|
|
||||||
body := strings.NewReader(strings.Repeat(" ", mcp.DefaultMaxRequestBodyBytes+1))
|
for _, test := range []struct {
|
||||||
request, err := http.NewRequest(http.MethodPost, httpTestServer.URL+"/mcp", body)
|
name string
|
||||||
|
size int64
|
||||||
|
tooLarge bool
|
||||||
|
}{
|
||||||
|
{name: "above the SDK default", size: mcp.DefaultMaxRequestBodyBytes + 1},
|
||||||
|
{name: "above our own limit", size: maxRequestBodyBytes + 1, tooLarge: true},
|
||||||
|
} {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
request, err := http.NewRequest(http.MethodPost, httpTestServer.URL+"/mcp", io.LimitReader(spaceReader{}, test.size))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewRequest() error = %v", err)
|
t.Fatalf("NewRequest() error = %v", err)
|
||||||
}
|
}
|
||||||
|
request.ContentLength = test.size
|
||||||
request.Header.Set("Content-Type", "application/json")
|
request.Header.Set("Content-Type", "application/json")
|
||||||
request.Header.Set("Accept", "application/json, text/event-stream")
|
request.Header.Set("Accept", "application/json, text/event-stream")
|
||||||
response, err := httpTestServer.Client().Do(request)
|
response, err := httpTestServer.Client().Do(request)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("POST large body error = %v", err)
|
t.Fatalf("POST %d bytes error = %v", test.size, err)
|
||||||
}
|
}
|
||||||
defer response.Body.Close()
|
defer response.Body.Close()
|
||||||
if response.StatusCode == http.StatusRequestEntityTooLarge {
|
if gotTooLarge := response.StatusCode == http.StatusRequestEntityTooLarge; gotTooLarge != test.tooLarge {
|
||||||
t.Errorf("POST large body status = %d; PR 1 must preserve the previous unlimited body behavior", response.StatusCode)
|
t.Errorf("POST %d bytes status = %d, want %d = %v", test.size, response.StatusCode, http.StatusRequestEntityTooLarge, test.tooLarge)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -179,8 +214,7 @@ func (t *authorizationTransport) set(value string) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *authorizationTransport) RoundTrip(request *http.Request) (*http.Response, error) {
|
func (t *authorizationTransport) RoundTrip(request *http.Request) (*http.Response, error) {
|
||||||
clone := request.Clone(request.Context())
|
clone := request.Clone(request.Context()) // Clone already copies the header
|
||||||
clone.Header = request.Header.Clone()
|
|
||||||
t.mu.RLock()
|
t.mu.RLock()
|
||||||
value := t.value
|
value := t.value
|
||||||
t.mu.RUnlock()
|
t.mu.RUnlock()
|
||||||
@@ -313,6 +347,7 @@ func TestStdioCommandTransport(t *testing.T) {
|
|||||||
if testing.Short() {
|
if testing.Short() {
|
||||||
t.Skip("skipping subprocess build in short mode")
|
t.Skip("skipping subprocess build in short mode")
|
||||||
}
|
}
|
||||||
|
exposeAllTools(t) // the subprocess runs with default flags, so match them here
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
@@ -336,58 +371,5 @@ func TestStdioCommandTransport(t *testing.T) {
|
|||||||
if got := session.InitializeResult().ProtocolVersion; got != "2026-07-28" {
|
if got := session.InitializeResult().ProtocolVersion; got != "2026-07-28" {
|
||||||
t.Errorf("protocol version = %q, want %q", got, "2026-07-28")
|
t.Errorf("protocol version = %q, want %q", got, "2026-07-28")
|
||||||
}
|
}
|
||||||
result, err := session.ListTools(ctx, nil)
|
listAndCallVersion(ctx, t, session, "Gitea MCP Server version:")
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("ListTools() error = %v", err)
|
|
||||||
}
|
|
||||||
if len(result.Tools) != 54 {
|
|
||||||
t.Errorf("ListTools() count = %d, want 54", len(result.Tools))
|
|
||||||
}
|
|
||||||
callResult, err := session.CallTool(ctx, &mcp.CallToolParams{Name: "get_gitea_mcp_server_version"})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("CallTool() error = %v", err)
|
|
||||||
}
|
|
||||||
content, ok := callResult.Content[0].(*mcp.TextContent)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("CallTool() content type = %T, want *mcp.TextContent", callResult.Content[0])
|
|
||||||
}
|
|
||||||
if !strings.Contains(content.Text, "Gitea MCP Server version:") {
|
|
||||||
t.Errorf("CallTool() result = %q, want server version", content.Text)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewHTTPServerAddress(t *testing.T) {
|
|
||||||
server := newHTTPServer(":12345", newMCPServer(testServerVersion))
|
|
||||||
if server.Addr != ":12345" {
|
|
||||||
t.Errorf("server address = %q, want %q", server.Addr, ":12345")
|
|
||||||
}
|
|
||||||
if server.Handler == nil {
|
|
||||||
t.Error("server handler is nil")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHTTPServerGracefulShutdown(t *testing.T) {
|
|
||||||
server := newHTTPServer("127.0.0.1:0", newMCPServer(testServerVersion))
|
|
||||||
listener, err := net.Listen("tcp", server.Addr)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Listen() error = %v", err)
|
|
||||||
}
|
|
||||||
serveDone := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
serveDone <- server.Serve(listener)
|
|
||||||
}()
|
|
||||||
|
|
||||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
||||||
defer cancel()
|
|
||||||
if err := server.Shutdown(shutdownCtx); err != nil {
|
|
||||||
t.Fatalf("Shutdown() error = %v", err)
|
|
||||||
}
|
|
||||||
select {
|
|
||||||
case err := <-serveDone:
|
|
||||||
if !errors.Is(err, http.ErrServerClosed) {
|
|
||||||
t.Errorf("Serve() error = %v, want http.ErrServerClosed", err)
|
|
||||||
}
|
|
||||||
case <-shutdownCtx.Done():
|
|
||||||
t.Fatal("server did not stop after Shutdown()")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
-2817
File diff suppressed because it is too large
Load Diff
+97
-134
@@ -1,179 +1,142 @@
|
|||||||
package operation
|
package operation
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"os"
|
"slices"
|
||||||
"path/filepath"
|
|
||||||
"sort"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||||
)
|
)
|
||||||
|
|
||||||
const updateToolContractEnv = "UPDATE_TOOL_CONTRACT"
|
// TestToolContract checks the properties every exposed tool must hold, rather
|
||||||
|
// than a snapshot of the current surface, so adding a tool needs no fixture
|
||||||
type toolContract struct {
|
// update and a malformed schema fails here instead of panicking in AddTool.
|
||||||
Scope string `json:"scope"`
|
|
||||||
Access string `json:"access"`
|
|
||||||
Name string `json:"name"`
|
|
||||||
Description string `json:"description"`
|
|
||||||
InputSchema any `json:"inputSchema"`
|
|
||||||
Annotations contractAnnotations `json:"annotations"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type contractAnnotations struct {
|
|
||||||
Title string `json:"title"`
|
|
||||||
ReadOnlyHint bool `json:"readOnlyHint"`
|
|
||||||
DestructiveHint bool `json:"destructiveHint"`
|
|
||||||
IdempotentHint bool `json:"idempotentHint"`
|
|
||||||
OpenWorldHint bool `json:"openWorldHint"`
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestToolContract(t *testing.T) {
|
func TestToolContract(t *testing.T) {
|
||||||
const (
|
scopeByName := map[string]string{}
|
||||||
wantDomains = 18
|
seenScopes := map[string]struct{}{}
|
||||||
wantTools = 54
|
|
||||||
wantRead = 33
|
|
||||||
wantWrite = 21
|
|
||||||
)
|
|
||||||
|
|
||||||
contracts := make([]toolContract, 0, wantTools)
|
|
||||||
seenScopes := make(map[string]struct{}, wantDomains)
|
|
||||||
seenNames := make(map[string]struct{}, wantTools)
|
|
||||||
readCount, writeCount := 0, 0
|
|
||||||
|
|
||||||
for _, domain := range domainTools {
|
for _, domain := range domainTools {
|
||||||
scope := domain.Scope()
|
scope := domain.Scope()
|
||||||
if scope == "" {
|
if scope == "" {
|
||||||
t.Fatal("registered tool domain has an empty scope")
|
t.Error("domainTools contains a domain with an empty scope")
|
||||||
}
|
}
|
||||||
|
// Tools() filters one domain by exactly one scope name, so a shared
|
||||||
|
// scope would make --scope select more than the caller asked for.
|
||||||
if _, duplicate := seenScopes[scope]; duplicate {
|
if _, duplicate := seenScopes[scope]; duplicate {
|
||||||
t.Fatalf("duplicate tool domain scope %q", scope)
|
t.Errorf("domainTools contains a duplicate scope %q", scope)
|
||||||
}
|
}
|
||||||
seenScopes[scope] = struct{}{}
|
seenScopes[scope] = struct{}{}
|
||||||
|
|
||||||
for _, registered := range domain.ReadTools() {
|
for _, registered := range domain.ReadTools() {
|
||||||
contracts = append(contracts, decodeToolContract(t, scope, "read", registered.Tool))
|
assertToolContract(t, scope, registered.Tool, true, scopeByName)
|
||||||
readCount++
|
|
||||||
}
|
}
|
||||||
for _, registered := range domain.WriteTools() {
|
for _, registered := range domain.WriteTools() {
|
||||||
contracts = append(contracts, decodeToolContract(t, scope, "write", registered.Tool))
|
assertToolContract(t, scope, registered.Tool, false, scopeByName)
|
||||||
writeCount++
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if len(scopeByName) == 0 {
|
||||||
if len(seenScopes) != wantDomains {
|
t.Fatal("no tools are registered")
|
||||||
t.Errorf("domain count = %d, want %d", len(seenScopes), wantDomains)
|
|
||||||
}
|
|
||||||
if len(contracts) != wantTools {
|
|
||||||
t.Errorf("tool count = %d, want %d", len(contracts), wantTools)
|
|
||||||
}
|
|
||||||
if readCount != wantRead {
|
|
||||||
t.Errorf("read tool count = %d, want %d", readCount, wantRead)
|
|
||||||
}
|
|
||||||
if writeCount != wantWrite {
|
|
||||||
t.Errorf("write tool count = %d, want %d", writeCount, wantWrite)
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, contract := range contracts {
|
|
||||||
if _, duplicate := seenNames[contract.Name]; duplicate {
|
|
||||||
t.Errorf("duplicate tool name %q", contract.Name)
|
|
||||||
}
|
|
||||||
seenNames[contract.Name] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
sort.Slice(contracts, func(i, j int) bool {
|
|
||||||
if contracts[i].Scope != contracts[j].Scope {
|
|
||||||
return contracts[i].Scope < contracts[j].Scope
|
|
||||||
}
|
|
||||||
if contracts[i].Access != contracts[j].Access {
|
|
||||||
return contracts[i].Access < contracts[j].Access
|
|
||||||
}
|
|
||||||
return contracts[i].Name < contracts[j].Name
|
|
||||||
})
|
|
||||||
|
|
||||||
got, err := json.MarshalIndent(contracts, "", " ")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("marshal tool contract: %v", err)
|
|
||||||
}
|
|
||||||
got = append(got, '\n')
|
|
||||||
|
|
||||||
goldenPath := filepath.Join("testdata", "tools.golden.json")
|
|
||||||
if os.Getenv(updateToolContractEnv) == "1" {
|
|
||||||
if err := os.WriteFile(goldenPath, got, 0o644); err != nil {
|
|
||||||
t.Fatalf("update tool contract: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
want, err := os.ReadFile(goldenPath)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read tool contract: %v", err)
|
|
||||||
}
|
|
||||||
if !bytes.Equal(got, want) {
|
|
||||||
t.Errorf("tool contract changed; inspect the semantic diff before running %s=1 go test -run '^TestToolContract$' ./operation/", updateToolContractEnv)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func decodeToolContract(t *testing.T, scope, access string, toolDefinition any) toolContract {
|
func assertToolContract(t *testing.T, scope string, definition *mcp.Tool, readOnly bool, scopeByName map[string]string) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
data, err := json.Marshal(toolDefinition)
|
t.Run(definition.Name, func(t *testing.T) {
|
||||||
if err != nil {
|
if previous, duplicate := scopeByName[definition.Name]; duplicate {
|
||||||
t.Fatalf("marshal %s tool in scope %q: %v", access, scope, err)
|
t.Errorf("tool name is already registered in scope %q; AddTool would silently replace it", previous)
|
||||||
}
|
}
|
||||||
var definition map[string]any
|
scopeByName[definition.Name] = scope
|
||||||
if err := json.Unmarshal(data, &definition); err != nil {
|
|
||||||
t.Fatalf("decode %s tool in scope %q: %v", access, scope, err)
|
// Strict MCP clients reject a tools/list entry without a description.
|
||||||
|
if definition.Description == "" {
|
||||||
|
t.Error("tool has no description")
|
||||||
}
|
}
|
||||||
|
|
||||||
name := requiredString(t, definition, "name", scope)
|
// A write tool registered as read stays exposed under --read-only.
|
||||||
description := requiredString(t, definition, "description", name)
|
if definition.Annotations == nil || definition.Annotations.ReadOnlyHint != readOnly {
|
||||||
inputSchema, ok := definition["inputSchema"].(map[string]any)
|
t.Errorf("annotations = %+v, want readOnlyHint %v", definition.Annotations, readOnly)
|
||||||
|
}
|
||||||
|
|
||||||
|
schema := decodeJSON(t, definition.InputSchema)
|
||||||
|
if schema["type"] != "object" {
|
||||||
|
t.Fatalf("input schema type = %v, want object", schema["type"])
|
||||||
|
}
|
||||||
|
properties, ok := schema["properties"].(map[string]any)
|
||||||
if !ok {
|
if !ok {
|
||||||
t.Fatalf("tool %q has inputSchema of type %T, want JSON object", name, definition["inputSchema"])
|
t.Fatalf("input schema properties = %T, want a JSON object", schema["properties"])
|
||||||
}
|
}
|
||||||
// An omitted required keyword and an empty array have the same JSON Schema meaning.
|
|
||||||
if _, ok := inputSchema["required"]; !ok {
|
|
||||||
inputSchema["required"] = []any{}
|
|
||||||
}
|
|
||||||
annotations, _ := definition["annotations"].(map[string]any)
|
|
||||||
|
|
||||||
// Normalize protocol defaults independently of SDK omitempty behavior.
|
for name, raw := range properties {
|
||||||
return toolContract{
|
property, ok := raw.(map[string]any)
|
||||||
Scope: scope,
|
if !ok {
|
||||||
Access: access,
|
t.Errorf("property %q = %T, want a JSON object", name, raw)
|
||||||
Name: name,
|
continue
|
||||||
Description: description,
|
|
||||||
InputSchema: inputSchema,
|
|
||||||
Annotations: contractAnnotations{
|
|
||||||
Title: stringField(annotations, "title", ""),
|
|
||||||
ReadOnlyHint: boolField(annotations, "readOnlyHint", false),
|
|
||||||
DestructiveHint: boolField(annotations, "destructiveHint", true),
|
|
||||||
IdempotentHint: boolField(annotations, "idempotentHint", false),
|
|
||||||
OpenWorldHint: boolField(annotations, "openWorldHint", true),
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
assertPropertyContract(t, name, property)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func requiredString(t *testing.T, object map[string]any, key, owner string) string {
|
func assertPropertyContract(t *testing.T, name string, property map[string]any) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
value, ok := object[key].(string)
|
propertyType, ok := property["type"].(string)
|
||||||
if !ok || value == "" {
|
if !ok {
|
||||||
t.Fatalf("%s has missing or empty %q", owner, key)
|
t.Errorf("property %q has no type", name)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
enum, hasEnum := property["enum"].([]any)
|
||||||
|
if _, declared := property["enum"]; declared && len(enum) == 0 {
|
||||||
|
t.Errorf("property %q has an empty enum", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
defaultValue, hasDefault := property["default"]
|
||||||
|
if !hasDefault {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !matchesJSONType(defaultValue, propertyType) {
|
||||||
|
t.Errorf("property %q default %#v is not a %s", name, defaultValue, propertyType)
|
||||||
|
}
|
||||||
|
if hasEnum && !slices.Contains(enum, defaultValue) {
|
||||||
|
t.Errorf("property %q default %#v is not one of its enum values %#v", name, defaultValue, enum)
|
||||||
}
|
}
|
||||||
return value
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func stringField(object map[string]any, key, fallback string) string {
|
func matchesJSONType(value any, propertyType string) bool {
|
||||||
if value, ok := object[key].(string); ok {
|
switch propertyType {
|
||||||
return value
|
case "string":
|
||||||
|
_, ok := value.(string)
|
||||||
|
return ok
|
||||||
|
case "number":
|
||||||
|
_, ok := value.(float64)
|
||||||
|
return ok
|
||||||
|
case "boolean":
|
||||||
|
_, ok := value.(bool)
|
||||||
|
return ok
|
||||||
|
case "array":
|
||||||
|
_, ok := value.([]any)
|
||||||
|
return ok
|
||||||
|
case "object":
|
||||||
|
_, ok := value.(map[string]any)
|
||||||
|
return ok
|
||||||
|
default:
|
||||||
|
return false
|
||||||
}
|
}
|
||||||
return fallback
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func boolField(object map[string]any, key string, fallback bool) bool {
|
// decodeJSON round-trips through JSON so the assertions see what an MCP client
|
||||||
if value, ok := object[key].(bool); ok {
|
// receives rather than the Go values behind it.
|
||||||
return value
|
func decodeJSON(t *testing.T, value any) map[string]any {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
encoded, err := json.Marshal(value)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal: %v", err)
|
||||||
}
|
}
|
||||||
return fallback
|
var decoded map[string]any
|
||||||
|
if err := json.Unmarshal(encoded, &decoded); err != nil {
|
||||||
|
t.Fatalf("decode: %v", err)
|
||||||
|
}
|
||||||
|
return decoded
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,20 +1,50 @@
|
|||||||
package annotation
|
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) {
|
func TestAnnotations(t *testing.T) {
|
||||||
readOnly := ReadOnly("Read")
|
for _, test := range []struct {
|
||||||
if readOnly.Title != "Read" || !readOnly.ReadOnlyHint || readOnly.DestructiveHint != nil {
|
name string
|
||||||
t.Errorf("ReadOnly() = %#v", readOnly)
|
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
|
||||||
write := Write("Write")
|
if err := json.Unmarshal(encoded, &got); err != nil {
|
||||||
if write.Title != "Write" || write.ReadOnlyHint || write.DestructiveHint != nil {
|
t.Fatalf("json.Unmarshal() error = %v", err)
|
||||||
t.Errorf("Write() = %#v", write)
|
|
||||||
}
|
}
|
||||||
|
if !maps.Equal(got, test.want) {
|
||||||
destructive := Destructive("Delete")
|
t.Errorf("annotations = %s, want %v", encoded, test.want)
|
||||||
if destructive.Title != "Delete" || destructive.ReadOnlyHint || destructive.DestructiveHint == nil || !*destructive.DestructiveHint {
|
}
|
||||||
t.Errorf("Destructive() = %#v", destructive)
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+24
-28
@@ -10,7 +10,7 @@ type Property struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// PropertyOption configures one property in a tool's input schema.
|
// 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.
|
// NewDefinition builds a tool definition without enabling SDK-side validation.
|
||||||
func NewDefinition(name, description string, annotations *mcp.ToolAnnotations, properties ...Property) *mcp.Tool {
|
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 {
|
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 {
|
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 {
|
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 {
|
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 {
|
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 {
|
func newProperty(name string, schema map[string]any, options ...PropertyOption) Property {
|
||||||
schema := map[string]any{"type": propertyType}
|
property := Property{name: name, schema: schema}
|
||||||
if object {
|
|
||||||
schema["properties"] = map[string]any{}
|
|
||||||
}
|
|
||||||
for _, option := range options {
|
for _, option := range options {
|
||||||
option(schema)
|
option(&property)
|
||||||
}
|
}
|
||||||
|
return property
|
||||||
required, _ := schema["required"].(bool)
|
|
||||||
delete(schema, "required")
|
|
||||||
return Property{name: name, schema: schema, required: required}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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 {
|
func Required() PropertyOption {
|
||||||
return func(schema map[string]any) {
|
return func(property *Property) {
|
||||||
schema["required"] = true
|
property.required = true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func Description(description string) PropertyOption {
|
func Description(description string) PropertyOption {
|
||||||
return func(schema map[string]any) {
|
return func(property *Property) {
|
||||||
schema["description"] = description
|
property.schema["description"] = description
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func Enum(values ...string) PropertyOption {
|
func Enum(values ...string) PropertyOption {
|
||||||
return func(schema map[string]any) {
|
return func(property *Property) {
|
||||||
schema["enum"] = values
|
property.schema["enum"] = values
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func Default(value any) PropertyOption {
|
func Default(value any) PropertyOption {
|
||||||
return func(schema map[string]any) {
|
return func(property *Property) {
|
||||||
schema["default"] = value
|
property.schema["default"] = value
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func Minimum(value float64) PropertyOption {
|
func Minimum(value float64) PropertyOption {
|
||||||
return func(schema map[string]any) {
|
return func(property *Property) {
|
||||||
schema["minimum"] = value
|
property.schema["minimum"] = value
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func Items(schema any) PropertyOption {
|
func Items(schema any) PropertyOption {
|
||||||
return func(propertySchema map[string]any) {
|
return func(property *Property) {
|
||||||
propertySchema["items"] = schema
|
property.schema["items"] = schema
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
package tool
|
package tool
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
|
||||||
"reflect"
|
"reflect"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -58,14 +57,6 @@ func TestNewDefinition(t *testing.T) {
|
|||||||
if !reflect.DeepEqual(definition.InputSchema, want) {
|
if !reflect.DeepEqual(definition.InputSchema, want) {
|
||||||
t.Errorf("InputSchema = %#v, want %#v", 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) {
|
func TestNewDefinitionWithoutRequiredProperties(t *testing.T) {
|
||||||
|
|||||||
+49
-44
@@ -10,19 +10,23 @@ import (
|
|||||||
"github.com/modelcontextprotocol/go-sdk/mcp"
|
"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) {
|
func TestMCPHandler(t *testing.T) {
|
||||||
var got map[string]any
|
var got map[string]any
|
||||||
serverTool := ServerTool{
|
result, err := callTool(captureArguments(&got), json.RawMessage(`{"count":2,"nested":{"enabled":true}}`))
|
||||||
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 {
|
if err != nil {
|
||||||
t.Fatalf("MCPHandler() error = %v", err)
|
t.Fatalf("MCPHandler() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -36,18 +40,13 @@ func TestMCPHandler(t *testing.T) {
|
|||||||
|
|
||||||
func TestMCPHandlerRejectsInvalidArguments(t *testing.T) {
|
func TestMCPHandlerRejectsInvalidArguments(t *testing.T) {
|
||||||
called := false
|
called := false
|
||||||
serverTool := ServerTool{
|
handler := func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
|
||||||
Tool: &mcp.Tool{Name: "example"},
|
|
||||||
Handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
|
|
||||||
called = true
|
called = true
|
||||||
return &mcp.CallToolResult{}, nil
|
return &mcp.CallToolResult{}, nil
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, arguments := range []json.RawMessage{json.RawMessage(`[]`), json.RawMessage(`null`), json.RawMessage(`{"broken"`)} {
|
for _, arguments := range []json.RawMessage{json.RawMessage(`[]`), json.RawMessage(`"text"`), json.RawMessage(`{"broken"`)} {
|
||||||
_, err := serverTool.MCPHandler()(context.Background(), &mcp.CallToolRequest{
|
_, err := callTool(handler, arguments)
|
||||||
Params: &mcp.CallToolParamsRaw{Arguments: arguments},
|
|
||||||
})
|
|
||||||
assertProtocolErrorCode(t, err, jsonrpc.CodeInvalidParams)
|
assertProtocolErrorCode(t, err, jsonrpc.CodeInvalidParams)
|
||||||
}
|
}
|
||||||
if called {
|
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) {
|
func TestMCPHandlerConvertsErrorsAndRecoversPanics(t *testing.T) {
|
||||||
t.Run("handler error", func(t *testing.T) {
|
for _, test := range []struct {
|
||||||
serverTool := ServerTool{
|
name string
|
||||||
Tool: &mcp.Tool{Name: "example"},
|
handler Handler
|
||||||
Handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
|
}{
|
||||||
|
{
|
||||||
|
name: "handler error",
|
||||||
|
handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
|
||||||
return nil, errors.New("failed")
|
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()
|
name: "panic",
|
||||||
_, err := handler(context.Background(), &mcp.CallToolRequest{Params: &mcp.CallToolParamsRaw{}})
|
handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
|
||||||
|
panic("failed")
|
||||||
|
},
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
_, err := callTool(test.handler, nil)
|
||||||
assertProtocolErrorCode(t, err, jsonrpc.CodeInternalError)
|
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) {
|
func assertProtocolErrorCode(t *testing.T, err error, want int64) {
|
||||||
|
|||||||
+9
-26
@@ -1,7 +1,6 @@
|
|||||||
package tool
|
package tool
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
@@ -89,30 +88,18 @@ func (t *Tool) Tools() []ServerTool {
|
|||||||
// MCPHandler adapts a project handler to the official SDK's low-level handler.
|
// MCPHandler adapts a project handler to the official SDK's low-level handler.
|
||||||
func (s ServerTool) MCPHandler() mcp.ToolHandler {
|
func (s ServerTool) MCPHandler() mcp.ToolHandler {
|
||||||
return func(ctx context.Context, req *mcp.CallToolRequest) (result *mcp.CallToolResult, err error) {
|
return func(ctx context.Context, req *mcp.CallToolRequest) (result *mcp.CallToolResult, err error) {
|
||||||
name := ""
|
|
||||||
if s.Tool != nil {
|
|
||||||
name = s.Tool.Name
|
|
||||||
}
|
|
||||||
defer func() {
|
defer func() {
|
||||||
if recovered := recover(); recovered != nil {
|
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)
|
log.Errorf("%s", panicErr)
|
||||||
result = nil
|
err = internalError(panicErr)
|
||||||
err = &jsonrpc.Error{Code: jsonrpc.CodeInternalError, Message: panicErr.Error()}
|
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
if req == nil || req.Params == nil {
|
|
||||||
return nil, invalidParamsError("missing tool call parameters")
|
|
||||||
}
|
|
||||||
|
|
||||||
arguments, err := decodeArguments(req.Params.Arguments)
|
arguments, err := decodeArguments(req.Params.Arguments)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if s.Handler == nil {
|
|
||||||
return nil, internalError(fmt.Errorf("tool %q has no handler", name))
|
|
||||||
}
|
|
||||||
|
|
||||||
result, err = s.Handler(ctx, arguments)
|
result, err = s.Handler(ctx, arguments)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -127,25 +114,21 @@ func (s ServerTool) MCPHandler() mcp.ToolHandler {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func decodeArguments(raw json.RawMessage) (map[string]any, error) {
|
func decodeArguments(raw json.RawMessage) (map[string]any, error) {
|
||||||
trimmed := bytes.TrimSpace(raw)
|
// An omitted and a null "arguments" both mean the tool was called without any.
|
||||||
if len(trimmed) == 0 {
|
if len(raw) == 0 || string(raw) == "null" {
|
||||||
return map[string]any{}, nil
|
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
|
var arguments map[string]any
|
||||||
if err := json.Unmarshal(trimmed, &arguments); err != nil {
|
if err := json.Unmarshal(raw, &arguments); err != nil {
|
||||||
return nil, invalidParamsError(fmt.Sprintf("invalid tool arguments: %v", err))
|
return nil, &jsonrpc.Error{
|
||||||
|
Code: jsonrpc.CodeInvalidParams,
|
||||||
|
Message: fmt.Sprintf("invalid tool arguments: %v", err),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return arguments, nil
|
return arguments, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func invalidParamsError(message string) error {
|
|
||||||
return &jsonrpc.Error{Code: jsonrpc.CodeInvalidParams, Message: message}
|
|
||||||
}
|
|
||||||
|
|
||||||
func internalError(err error) error {
|
func internalError(err error) error {
|
||||||
return &jsonrpc.Error{Code: jsonrpc.CodeInternalError, Message: err.Error()}
|
return &jsonrpc.Error{Code: jsonrpc.CodeInternalError, Message: err.Error()}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user