diff --git a/AGENTS.md b/AGENTS.md index 7746265..dbc863b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -40,8 +40,11 @@ Server.run() ← dispatch loop | `gomcp/protocol.go` | Protocol versions, `_meta` negotiation, pagination | | `gomcp/path.go` | `SafeJoin` — path-traversal-safe filesystem helper | | `gomcp/server.go` | Server struct, Run(), all JSON-RPC method handlers | +| `gomcp/httpserver.go` | MCP Streamable HTTP transport (`Handler`, `ListenAndServe`, `Serve`) | +| `gomcp/client.go` | MCP client (`Client` interface, `NewHTTPClient`, `NewStdioClient`) | | `gomcp/server_test.go` | Unit + integration tests (pipe-based) | | `gomcp/protocol_test.go` | 2026-07-28 + security tests | +| `gomcp/httpserver_test.go`, `gomcp/client_test.go`, `gomcp/http_e2e_test.go` | HTTP transport, client, and loopback E2E tests | | `gomcp/e2e_test.go` | Subprocess E2E test | | `examples/greet/main.go` | Canonical example MCP server | diff --git a/README.md b/README.md index 2135670..5bac9db 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,8 @@ # go-mcp -Zero-dependency [Model Context Protocol](https://modelcontextprotocol.io) server framework for Go. Expose your Go code as **tools**, **resources**, and **prompts** that AI agents can call. Stdio transport. Single binary. No runtime. +> Zero-dependency Model Context Protocol (MCP) server framework for Go — stdio and Streamable HTTP transports, plus MCP clients for both. + +Zero-dependency [Model Context Protocol](https://modelcontextprotocol.io) server framework for Go. Expose your Go code as **tools**, **resources**, and **prompts** that AI agents can call. Stdio and Streamable HTTP transports, plus clients for both. Single binary. No runtime. ```bash go get github.com/BackendStack21/go-mcp @@ -304,6 +306,45 @@ func NewTextContent(text string) map[string]any Starts the server loop. Reads JSON-RPC 2.0 from `os.Stdin`, writes responses to `os.Stdout`. Blocks until stdin closes (EOF). +### `srv.Handler() http.Handler` / `srv.ListenAndServe(addr) error` / `srv.Serve(l net.Listener) error` + +Serves the same JSON-RPC dispatch over the MCP **Streamable HTTP** transport. Each `POST` carries one JSON-RPC message; responses are `application/json` or SSE (when the client's `Accept` prefers it), notifications get `202`, and anything but `POST` (with a JSON body) is `405`/`415`. Reuses `MaxRequestBytes` and `HandlerTimeout`. Stateless — no session management, no server-initiated streams. + +```go +srv := gomcp.NewServer("my-server", "1.0.0") +// ... register tools ... +srv.ListenAndServe("localhost:8080") // blocks +``` + +### Client — `gomcp.NewHTTPClient(url)` / `gomcp.NewStdioClient(path, args...)` + +Both return the same `Client` interface, safe for concurrent use: + +```go +c := gomcp.NewHTTPClient("http://localhost:8080") +// or: c, err := gomcp.NewStdioClient("./my-server") +res, _ := c.Initialize(ctx, "my-client", "1.0.0") +c.NotifyInitialized(ctx) +tools, _ := c.ListTools(ctx) +out, _ := c.CallTool(ctx, "echo", map[string]any{"message": "hi"}) +``` + +Methods: `Initialize`, `NotifyInitialized`, `ListTools`, `CallTool`, `ListResources`, `ReadResource`, `ListPrompts`, `GetPrompt`, `Close`. + +### `srv.SetAuthToken(token string)` / `srv.SetAllowedOrigins(origins []string)` + +Optional HTTP-transport hardening: + +- **Bearer auth** — when a token is set, every request must carry `Authorization: Bearer `; anything else gets `401` with a `WWW-Authenticate: Bearer` challenge. Comparison is constant-time (`crypto/subtle`). Pair with any external token issuer (OAuth resource server, API gateway, or a plain static key). +- **Origin allowlist** — when set, browser requests with an `Origin` not on the list get `403` (DNS-rebinding / CSRF defense, per 2025-11-25 spec guidance). Non-browser clients (no `Origin` header) are unaffected. + +Both default to off; stdio is never affected — its trust model is the user launching the process. + +```go +srv.SetAuthToken(os.Getenv("MCP_TOKEN")) +srv.SetAllowedOrigins([]string{"https://claude.ai"}) +``` + ### `srv.SetInstructions(text string)` Optional natural-language guidance returned by `initialize` and `server/discover`. diff --git a/docs/index.html b/docs/index.html index 7c35c12..6d3c9cf 100644 --- a/docs/index.html +++ b/docs/index.html @@ -4,7 +4,7 @@ go-mcp — MCP Server Framework for Go - + @@ -15,7 +15,7 @@ - + @@ -48,6 +48,8 @@ /* Grid */ .grid-3 { display: grid; grid-template-columns: repeat(3, 1fr); gap: 16px; } @media (max-width: 900px) { .grid-3 { grid-template-columns: 1fr; } } + .grid-2 { display: grid; grid-template-columns: repeat(2, 1fr); gap: 16px; } + @media (max-width: 900px) { .grid-2 { grid-template-columns: 1fr; } } /* Card icon blue = accent */ .card-icon.blue { background: var(--accent-subtle); color: var(--accent); } @@ -156,7 +158,7 @@
Zero dependencies. Single binary. Go-native.

Build AI tools
in Go

-

A zero-dependency MCP server framework that turns any Go program into an AI-accessible service. Tools, resources, and prompts over stdio. Compile, ship, connect.

+

A zero-dependency MCP framework that turns any Go program into an AI-accessible service. Tools, resources, and prompts over stdio or Streamable HTTP — with clients for both. Compile, ship, connect.

@@ -215,6 +217,10 @@

Build AI tools
in Go

11
MCP methods
+
+
2
+
transports — stdio & HTTP
+
<1ms
cold start
@@ -232,7 +238,7 @@

Build AI tools
in Go

Turn any Go program into an
AI-accessible service

-

Register tools, resources, and prompts. The AI client discovers them automatically. All over stdin/stdout.

+

Register tools, resources, and prompts. The AI client discovers them automatically over stdio or HTTP. Or connect to other MCP servers yourself with the built-in clients.

@@ -262,6 +268,20 @@

Kubernetes Operator

API Integration

Bridge any REST API to AI. Give agents access to GitHub, Stripe, Slack — any HTTP endpoint becomes a callable tool.

+
+
+ +
+

HTTP-Native Server

+

Serve MCP over Streamable HTTP — one handler, any port. Optional Bearer auth and Origin allowlist built in.

+
+
+
+ +
+

MCP Client Too

+

Connect to any MCP server over HTTP or spawn one over stdio — the same typed Client interface for both.

+
@@ -359,6 +379,36 @@

Prompts

Pre-defined conversation templates. The AI requests them by name with arguments via prompts/get. Returns formatted messages ready for the chat.

+ +

Two transports. One framework.

+

The same server runs over stdio or Streamable HTTP — and the built-in clients talk to either.

+
+
+
server — stdio or HTTP
+
// stdio — the default, for local AI clients
+srv.Run()
+
+// Streamable HTTP — for remote access
+srv.SetAuthToken(token)            // optional
+srv.SetAllowedOrigins(origins)     // optional
+srv.ListenAndServe(":8080")
+
+// or mount on your own mux:
+mux.Handle("/mcp", srv.Handler())
+
+
+
client — HTTP or subprocess
+
// connect to a remote server
+c := gomcp.NewHTTPClientWithToken(url, token)
+
+// or spawn a local one
+c, _ := gomcp.NewStdioClient("./my-server")
+
+c.Initialize(ctx, "client", "1.0")
+tools, _ := c.ListTools(ctx)
+out, _  := c.CallTool(ctx, "greet", args)
+
+
@@ -403,7 +453,7 @@

One command. Zero dependencies.

go get github.com/BackendStack21/go-mcp

- 15 tests. 4 example servers. View on GitHub → + 102 tests. 91.9% coverage. View on GitHub →

diff --git a/gomcp/client.go b/gomcp/client.go new file mode 100644 index 0000000..f35b6ab --- /dev/null +++ b/gomcp/client.go @@ -0,0 +1,547 @@ +package gomcp + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "os/exec" + "strings" + "sync" +) + +// Client is an MCP client speaking to a server over any transport. +// All methods are safe for concurrent use. +type Client interface { + // Initialize performs the MCP handshake and returns the server's + // capabilities. Call it once before other requests. + Initialize(ctx context.Context, clientName, clientVersion string) (*InitializeResult, error) + // NotifyInitialized sends the notifications/initialized notice. + NotifyInitialized(ctx context.Context) error + // ListTools returns the server's registered tools. + ListTools(ctx context.Context) ([]Tool, error) + // CallTool invokes a tool and returns its text output. + CallTool(ctx context.Context, name string, args map[string]any) (string, error) + // ListResources returns the server's registered resources. + ListResources(ctx context.Context) ([]Resource, error) + // ReadResource reads one resource by URI. + ReadResource(ctx context.Context, uri string) (string, error) + // ListPrompts returns the server's registered prompts. + ListPrompts(ctx context.Context) ([]Prompt, error) + // GetPrompt renders one prompt with arguments. + GetPrompt(ctx context.Context, name string, args map[string]any) ([]PromptMessage, error) + // Close releases transport resources (HTTP client: no-op; stdio: + // terminates the subprocess). + Close() error +} + +// ServerInfo identifies a server in an InitializeResult. +type ServerInfo struct { + Name string `json:"name"` + Version string `json:"version"` +} + +// InitializeResult is the server's response to initialize. +type InitializeResult struct { + ProtocolVersion string `json:"protocolVersion"` + ServerInfo ServerInfo `json:"serverInfo"` + Instructions string `json:"instructions,omitempty"` +} + +// ClientProtocolVersion is the protocol version HTTPClient and StdioClient +// request by default. +const ClientProtocolVersion = DefaultProtocolVersion + +// rpcCaller issues one JSON-RPC request and decodes the result into out. +type rpcCaller interface { + call(ctx context.Context, method string, params any, out any) error + notify(ctx context.Context, method string, params any) error +} + +// --- shared JSON-RPC plumbing --- + +type rpcWireParams struct { + Name string `json:"name,omitempty"` + Arguments map[string]any `json:"arguments,omitempty"` + URI string `json:"uri,omitempty"` +} + +func decodeRPCResponse(body []byte, out any) error { + var env struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Result json.RawMessage `json:"result"` + Error *struct { + Code int `json:"code"` + Message string `json:"message"` + } `json:"error"` + } + if err := json.Unmarshal(body, &env); err != nil { + return fmt.Errorf("malformed response: %w", err) + } + if env.Error != nil { + return fmt.Errorf("rpc error %d: %s", env.Error.Code, env.Error.Message) + } + if len(env.Result) == 0 { + return fmt.Errorf("response has no result: %s", strings.TrimSpace(string(body))) + } + if out != nil { + if err := json.Unmarshal(env.Result, out); err != nil { + return fmt.Errorf("decode result: %w", err) + } + } + return nil +} + +func clientMethods(c rpcCaller, protoVersion func() string) clientCore { + return clientCore{c: c, protoVersion: protoVersion} +} + +type clientCore struct { + c rpcCaller + protoVersion func() string +} + +func (cc clientCore) Initialize(ctx context.Context, clientName, clientVersion string) (*InitializeResult, error) { + var res InitializeResult + err := cc.c.call(ctx, "initialize", map[string]any{ + "protocolVersion": cc.protoVersion(), + "capabilities": map[string]any{}, + "clientInfo": map[string]string{"name": clientName, "version": clientVersion}, + }, &res) + if err != nil { + return nil, err + } + if !isSupportedProtocolVersion(res.ProtocolVersion) { + return nil, fmt.Errorf("server speaks unsupported protocol version %q (client requested %q)", res.ProtocolVersion, cc.protoVersion()) + } + return &res, nil +} + +func (cc clientCore) NotifyInitialized(ctx context.Context) error { + return cc.c.notify(ctx, "notifications/initialized", nil) +} + +func (cc clientCore) ListTools(ctx context.Context) ([]Tool, error) { + var res struct { + Tools []Tool `json:"tools"` + } + if err := cc.c.call(ctx, "tools/list", map[string]any{}, &res); err != nil { + return nil, err + } + return res.Tools, nil +} + +func (cc clientCore) CallTool(ctx context.Context, name string, args map[string]any) (string, error) { + var res struct { + Content []struct { + Text string `json:"text"` + } `json:"content"` + IsError bool `json:"isError"` + } + if err := cc.c.call(ctx, "tools/call", rpcWireParams{Name: name, Arguments: args}, &res); err != nil { + return "", err + } + var sb strings.Builder + for _, c := range res.Content { + sb.WriteString(c.Text) + } + if res.IsError { + return sb.String(), fmt.Errorf("tool error: %s", sb.String()) + } + return sb.String(), nil +} + +func (cc clientCore) ListResources(ctx context.Context) ([]Resource, error) { + var res struct { + Resources []Resource `json:"resources"` + } + if err := cc.c.call(ctx, "resources/list", map[string]any{}, &res); err != nil { + return nil, err + } + return res.Resources, nil +} + +func (cc clientCore) ReadResource(ctx context.Context, uri string) (string, error) { + var res struct { + Contents []struct { + Text string `json:"text"` + } `json:"contents"` + } + if err := cc.c.call(ctx, "resources/read", rpcWireParams{URI: uri}, &res); err != nil { + return "", err + } + if len(res.Contents) == 0 { + return "", fmt.Errorf("resource %q returned no contents", uri) + } + return res.Contents[0].Text, nil +} + +func (cc clientCore) ListPrompts(ctx context.Context) ([]Prompt, error) { + var res struct { + Prompts []Prompt `json:"prompts"` + } + if err := cc.c.call(ctx, "prompts/list", map[string]any{}, &res); err != nil { + return nil, err + } + return res.Prompts, nil +} + +func (cc clientCore) GetPrompt(ctx context.Context, name string, args map[string]any) ([]PromptMessage, error) { + var res struct { + Messages []PromptMessage `json:"messages"` + } + if err := cc.c.call(ctx, "prompts/get", rpcWireParams{Name: name, Arguments: args}, &res); err != nil { + return nil, err + } + return res.Messages, nil +} + +// --- HTTP transport --- + +// MaxResponseBytes caps one response body read by HTTPClient. A hostile +// or buggy server cannot OOM the client with an unbounded response. +const MaxResponseBytes int64 = 64 << 20 // 64 MiB + +// HTTPClient is a Client over the MCP Streamable HTTP transport. +type HTTPClient struct { + url string + token string + http *http.Client + core clientCore + nextID int64 + mu sync.Mutex +} + +// NewHTTPClient returns a Client posting JSON-RPC messages to url. +func NewHTTPClient(url string) *HTTPClient { + return NewHTTPClientWithToken(url, "") +} + +// NewHTTPClientWithToken returns a Client that authenticates every request +// with `Authorization: Bearer `. Prefer this over embedding +// credentials in the URL — URLs leak into logs and error messages. +func NewHTTPClientWithToken(url, token string) *HTTPClient { + hc := &HTTPClient{url: url, token: token, http: &http.Client{}} + hc.core = clientMethods(hc, func() string { return ClientProtocolVersion }) + return hc +} + +func (c *HTTPClient) setAuth(h http.Header) { + if c.token != "" { + h.Set("Authorization", "Bearer "+c.token) + } +} + +func (c *HTTPClient) call(ctx context.Context, method string, params any, out any) error { + c.mu.Lock() + c.nextID++ + id := c.nextID + c.mu.Unlock() + + payload, err := json.Marshal(map[string]any{ + "jsonrpc": "2.0", "id": id, "method": method, "params": params, + }) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.url, bytes.NewReader(payload)) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + c.setAuth(req.Header) + resp, err := c.http.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + body, err := io.ReadAll(io.LimitReader(resp.Body, MaxResponseBytes)) + if err != nil { + return err + } + if resp.StatusCode == http.StatusAccepted { + return fmt.Errorf("unexpected 202 for request %q", method) + } + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("http %d: %s", resp.StatusCode, strings.TrimSpace(string(body))) + } + // Handle both plain JSON and single-event SSE responses. + if strings.HasPrefix(resp.Header.Get("Content-Type"), "text/event-stream") { + body = sseData(body) + } + return decodeRPCResponse(body, out) +} + +func (c *HTTPClient) notify(ctx context.Context, method string, params any) error { + payload, err := json.Marshal(map[string]any{"jsonrpc": "2.0", "method": method, "params": params}) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.url, bytes.NewReader(payload)) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + c.setAuth(req.Header) + resp, err := c.http.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusAccepted && resp.StatusCode != http.StatusOK { + b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + return fmt.Errorf("notify %q: http %d: %s", method, resp.StatusCode, strings.TrimSpace(string(b))) + } + _, _ = io.Copy(io.Discard, resp.Body) + return nil +} + +// sseData extracts the JSON payload of the first data: field. The space +// after the colon is optional per the SSE spec. +func sseData(body []byte) []byte { + for _, line := range strings.Split(string(body), "\n") { + d, ok := strings.CutPrefix(line, "data:") + if !ok { + continue + } + return []byte(strings.TrimSpace(d)) + } + return body +} + +func (c *HTTPClient) Initialize(ctx context.Context, name, version string) (*InitializeResult, error) { + return c.core.Initialize(ctx, name, version) +} +func (c *HTTPClient) NotifyInitialized(ctx context.Context) error { + return c.core.NotifyInitialized(ctx) +} +func (c *HTTPClient) ListTools(ctx context.Context) ([]Tool, error) { + return c.core.ListTools(ctx) +} +func (c *HTTPClient) CallTool(ctx context.Context, name string, args map[string]any) (string, error) { + return c.core.CallTool(ctx, name, args) +} +func (c *HTTPClient) ListResources(ctx context.Context) ([]Resource, error) { + return c.core.ListResources(ctx) +} +func (c *HTTPClient) ReadResource(ctx context.Context, uri string) (string, error) { + return c.core.ReadResource(ctx, uri) +} +func (c *HTTPClient) ListPrompts(ctx context.Context) ([]Prompt, error) { + return c.core.ListPrompts(ctx) +} +func (c *HTTPClient) GetPrompt(ctx context.Context, name string, args map[string]any) ([]PromptMessage, error) { + return c.core.GetPrompt(ctx, name, args) +} +func (c *HTTPClient) Close() error { return nil } + +// --- stdio transport --- + +// StdioClient is a Client that spawns a server subprocess and speaks MCP +// stdio (newline-delimited JSON-RPC) over its pipes. +// +// A single lifetime reader goroutine demultiplexes responses by JSON-RPC +// id. A timed-out call does not desync the stream: its late response is +// delivered to the abandoned channel (buffered, then discarded) and the +// next call still receives its own response. +type StdioClient struct { + cmd *exec.Cmd + stdin io.WriteCloser + core clientCore + nextID int64 + mu sync.Mutex + + pendingMu sync.Mutex + pending map[int64]chan []byte + closed bool + readerDone chan struct{} +} + +// NewStdioClient starts the server binary at path (with optional args) and +// connects over stdio. +func NewStdioClient(path string, args ...string) (*StdioClient, error) { + cmd := exec.Command(path, args...) + stdin, err := cmd.StdinPipe() + if err != nil { + return nil, err + } + stdout, err := cmd.StdoutPipe() + if err != nil { + return nil, err + } + if err := cmd.Start(); err != nil { + return nil, err + } + sc := &StdioClient{ + cmd: cmd, + stdin: stdin, + pending: make(map[int64]chan []byte), + readerDone: make(chan struct{}), + } + sc.core = clientMethods(sc, func() string { return ClientProtocolVersion }) + go sc.readLoop(stdout) + return sc, nil +} + +// readLoop is the single lifetime reader. It matches each response to the +// pending request id; unmatched responses (e.g. late replies to timed-out +// calls) are discarded. +func (c *StdioClient) readLoop(r io.Reader) { + defer close(c.readerDone) + br := bufio.NewReader(r) + for { + line, err := br.ReadBytes('\n') + if len(bytes.TrimSpace(line)) > 0 { + var env struct { + ID *int64 `json:"id"` + } + if json.Unmarshal(line, &env) == nil && env.ID != nil { + c.pendingMu.Lock() + ch, ok := c.pending[*env.ID] + if ok { + delete(c.pending, *env.ID) + } + c.pendingMu.Unlock() + if ok { + ch <- line // buffered cap 1; never blocks + } + } + } + if err != nil { + c.failAllPending(err) + return + } + } +} + +// failAllPending unblocks every waiting call when the stream dies. +func (c *StdioClient) failAllPending(err error) { + c.pendingMu.Lock() + for id, ch := range c.pending { + delete(c.pending, id) + ch <- []byte(fmt.Sprintf(`{"error":{"code":-32000,"message":"stdio stream closed: %s"}}`, err)) + } + c.pendingMu.Unlock() +} + +func (c *StdioClient) roundTrip(ctx context.Context, payload []byte, id *int64) ([]byte, error) { + c.pendingMu.Lock() + if c.closed { + c.pendingMu.Unlock() + return nil, fmt.Errorf("client closed") + } + var ch chan []byte + if id != nil { + ch = make(chan []byte, 1) + c.pending[*id] = ch + } + c.pendingMu.Unlock() + + if _, err := c.stdin.Write(append(payload, '\n')); err != nil { + if id != nil { + c.pendingMu.Lock() + delete(c.pending, *id) + c.pendingMu.Unlock() + } + return nil, fmt.Errorf("write: %w", err) + } + if id == nil { + return nil, nil // notification: no response expected + } + + select { + case b := <-ch: + return b, nil + case <-ctx.Done(): + // Leave the pending entry; the reader deletes and drains it when + // the late response arrives, so the stream stays in sync. + return nil, ctx.Err() + case <-c.readerDone: + // Stream died; any pending entry was already failed. + select { + case b := <-ch: + return b, nil + default: + return nil, fmt.Errorf("stdio stream closed") + } + } +} + +func (c *StdioClient) call(ctx context.Context, method string, params any, out any) error { + c.mu.Lock() + c.nextID++ + id := c.nextID + c.mu.Unlock() + payload, err := json.Marshal(map[string]any{ + "jsonrpc": "2.0", "id": id, "method": method, "params": params, + }) + if err != nil { + return err + } + body, err := c.roundTrip(ctx, payload, &id) + if err != nil { + return err + } + return decodeRPCResponse(body, out) +} + +func (c *StdioClient) notify(ctx context.Context, method string, params any) error { + payload, err := json.Marshal(map[string]any{"jsonrpc": "2.0", "method": method, "params": params}) + if err != nil { + return err + } + _, err = c.roundTrip(ctx, payload, nil) + return err +} + +func (c *StdioClient) Initialize(ctx context.Context, name, version string) (*InitializeResult, error) { + return c.core.Initialize(ctx, name, version) +} +func (c *StdioClient) NotifyInitialized(ctx context.Context) error { + return c.core.NotifyInitialized(ctx) +} +func (c *StdioClient) ListTools(ctx context.Context) ([]Tool, error) { + return c.core.ListTools(ctx) +} +func (c *StdioClient) CallTool(ctx context.Context, name string, args map[string]any) (string, error) { + return c.core.CallTool(ctx, name, args) +} +func (c *StdioClient) ListResources(ctx context.Context) ([]Resource, error) { + return c.core.ListResources(ctx) +} +func (c *StdioClient) ReadResource(ctx context.Context, uri string) (string, error) { + return c.core.ReadResource(ctx, uri) +} +func (c *StdioClient) ListPrompts(ctx context.Context) ([]Prompt, error) { + return c.core.ListPrompts(ctx) +} +func (c *StdioClient) GetPrompt(ctx context.Context, name string, args map[string]any) ([]PromptMessage, error) { + return c.core.GetPrompt(ctx, name, args) +} + +func (c *StdioClient) Close() error { + c.pendingMu.Lock() + if c.closed { + c.pendingMu.Unlock() + <-c.readerDone + return nil + } + c.closed = true + c.pendingMu.Unlock() + + // Closing stdin makes the server's read loop EOF; the subprocess then + // exits and its pipes close, ending our reader goroutine. + _ = c.stdin.Close() + if c.cmd.Process != nil { + _ = c.cmd.Process.Kill() + } + err := c.cmd.Wait() + <-c.readerDone // reader observes pipe close and drains pending calls + return err +} diff --git a/gomcp/client_fixes_test.go b/gomcp/client_fixes_test.go new file mode 100644 index 0000000..2e72c82 --- /dev/null +++ b/gomcp/client_fixes_test.go @@ -0,0 +1,58 @@ +package gomcp + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +// M2: token-authenticated HTTP client works against a token-protected server. +func TestHTTPClientWithToken(t *testing.T) { + srv := NewServer("tok", "1.0.0") + srv.SetAuthToken("s3cret") + ts := httptest.NewServer(srv.Handler()) + defer ts.Close() + + // Anonymous client is rejected at handshake. + if _, err := NewHTTPClient(ts.URL).Initialize(t.Context(), "t", "0"); err == nil { + t.Fatal("anonymous client should fail against protected server") + } + + // Token client succeeds end to end. + c := NewHTTPClientWithToken(ts.URL, "s3cret") + res, err := c.Initialize(t.Context(), "t", "0") + if err != nil { + t.Fatal(err) + } + if res.ServerInfo.Name != "tok" { + t.Fatalf("serverInfo.name = %q", res.ServerInfo.Name) + } + if err := c.NotifyInitialized(t.Context()); err != nil { + t.Fatal(err) + } +} + +// M1: notify surfaces HTTP errors instead of swallowing them. +func TestHTTPClientNotifySurfacesErrors(t *testing.T) { + srv := NewServer("notify-err", "1.0.0") + srv.SetAuthToken("k") + ts := httptest.NewServer(srv.Handler()) + defer ts.Close() + + err := NewHTTPClient(ts.URL).NotifyInitialized(t.Context()) + if err == nil || !strings.Contains(err.Error(), "401") { + t.Fatalf("want 401 surfaced from notify, got %v", err) + } +} + +// M4: unsupported negotiated protocol version is rejected. +func TestClientRejectsUnsupportedProtocolVersion(t *testing.T) { + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"jsonrpc":"2.0","id":1,"result":{"protocolVersion":"1999-01-01","serverInfo":{"name":"x","version":"0"}}}`)) + })) + defer ts.Close() + if _, err := NewHTTPClient(ts.URL).Initialize(t.Context(), "t", "0"); err == nil { + t.Fatal("expected unsupported-protocol-version error") + } +} diff --git a/gomcp/client_test.go b/gomcp/client_test.go new file mode 100644 index 0000000..5eef9da --- /dev/null +++ b/gomcp/client_test.go @@ -0,0 +1,117 @@ +package gomcp + +import ( + "context" + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + "time" +) + +func TestHTTPClientInitializeAndCallTool(t *testing.T) { + _, ts := newHTTPTestServer(t) + c := NewHTTPClient(ts.URL) + + res, err := c.Initialize(context.Background(), "test-client", "1.2.3") + if err != nil { + t.Fatal(err) + } + if res.ServerInfo.Name != "http-test" { + t.Fatalf("serverInfo.name = %q", res.ServerInfo.Name) + } + if err := c.NotifyInitialized(context.Background()); err != nil { + t.Fatal(err) + } + + tools, err := c.ListTools(context.Background()) + if err != nil { + t.Fatal(err) + } + if len(tools) != 1 || tools[0].Name != "echo" { + t.Fatalf("tools = %+v", tools) + } + + out, err := c.CallTool(context.Background(), "echo", map[string]any{"message": "hello"}) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(out, "hello") { + t.Fatalf("CallTool output = %q", out) + } +} + +func TestHTTPClientRPCError(t *testing.T) { + _, ts := newHTTPTestServer(t) + c := NewHTTPClient(ts.URL) + if _, err := c.Initialize(context.Background(), "t", "0"); err != nil { + t.Fatal(err) + } + _, err := c.CallTool(context.Background(), "nope", nil) + if err == nil { + t.Fatal("expected error for unknown tool") + } +} + +func TestStdioClientAgainstSubprocess(t *testing.T) { + if testing.Short() { + t.Skip("spawns subprocess") + } + dir := t.TempDir() + src := filepath.Join(dir, "main.go") + if err := os.WriteFile(src, []byte(`package main + +import ( + "context" + "fmt" + + "github.com/BackendStack21/go-mcp/gomcp" +) + +func main() { + srv := gomcp.NewServer("stdio-echo", "1.0.0") + srv.AddTool(gomcp.Tool{ + Name: "upper", + InputSchema: gomcp.InputSchema{Type: "object"}, + Handler: func(ctx context.Context, args map[string]any) (string, error) { + return fmt.Sprintf("%v", args["text"]), nil + }, + }) + srv.Run() +} +`), 0o644); err != nil { + t.Fatal(err) + } + bin := filepath.Join(dir, "server") + // Build from the module root so the gomcp import resolves. + build := exec.Command("go", "build", "-o", bin, src) + build.Dir = "." + if out, err := build.CombinedOutput(); err != nil { + t.Fatalf("build: %v\n%s", err, out) + } + + c, err := NewStdioClient(bin) + if err != nil { + t.Fatal(err) + } + defer c.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + res, err := c.Initialize(ctx, "t", "0") + if err != nil { + t.Fatal(err) + } + if res.ServerInfo.Name != "stdio-echo" { + t.Fatalf("serverInfo.name = %q", res.ServerInfo.Name) + } + out, err := c.CallTool(ctx, "upper", map[string]any{"text": "works"}) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(out, "works") { + t.Fatalf("output = %q", out) + } +} diff --git a/gomcp/coverage_http_test.go b/gomcp/coverage_http_test.go new file mode 100644 index 0000000..55a03b9 --- /dev/null +++ b/gomcp/coverage_http_test.go @@ -0,0 +1,181 @@ +package gomcp + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "os/exec" + "path/filepath" + "strings" + "testing" +) + +// newFullTestServer registers a tool, a resource, and a prompt. +func newFullTestServer(t *testing.T) *httptest.Server { + t.Helper() + srv := NewServer("full", "1.0.0") + srv.AddTool(Tool{ + Name: "echo", + InputSchema: InputSchema{Type: "object"}, + Handler: func(ctx context.Context, args map[string]any) (string, error) { + return fmt.Sprintf("%v", args["message"]), nil + }, + }) + srv.AddResource(Resource{ + URI: "test://doc", + Name: "doc", + Handler: func(ctx context.Context) (string, error) { return "resource-body", nil }, + }) + srv.AddPrompt(Prompt{ + Name: "p1", + Handler: func(ctx context.Context, args map[string]any) ([]PromptMessage, error) { + return []PromptMessage{{Role: "user", Content: fmt.Sprintf("prompt %v", args["x"])}}, nil + }, + }) + ts := httptest.NewServer(srv.Handler()) + t.Cleanup(ts.Close) + return ts +} + +func mustInit(t *testing.T, c Client) { + t.Helper() + if _, err := c.Initialize(context.Background(), "t", "0"); err != nil { + t.Fatal(err) + } + if err := c.NotifyInitialized(context.Background()); err != nil { + t.Fatal(err) + } +} + +// buildTestServer compiles a test subprocess server from src, from the +// module root so the gomcp import resolves, and returns the binary path. +func buildTestServer(t *testing.T, dir, src string) string { + t.Helper() + bin := filepath.Join(dir, "server") + build := exec.Command("go", "build", "-o", bin, src) + build.Dir = "." + if out, err := build.CombinedOutput(); err != nil { + t.Fatalf("build: %v\n%s", err, out) + } + return bin +} + +// Exercise the full HTTPClient interface surface. +func TestHTTPClientFullSurface(t *testing.T) { + ts := newFullTestServer(t) + c := NewHTTPClient(ts.URL) + ctx := context.Background() + mustInit(t, c) + + if _, err := c.ListTools(ctx); err != nil { + t.Fatal(err) + } + if _, err := c.CallTool(ctx, "echo", map[string]any{"message": "m"}); err != nil { + t.Fatal(err) + } + + resources, err := c.ListResources(ctx) + if err != nil { + t.Fatal(err) + } + if len(resources) != 1 || resources[0].URI != "test://doc" { + t.Fatalf("resources = %+v", resources) + } + body, err := c.ReadResource(ctx, "test://doc") + if err != nil { + t.Fatal(err) + } + if body != "resource-body" { + t.Fatalf("body = %q", body) + } + // Resource that was never registered -> JSON-RPC error surfaced. + if _, err := c.ReadResource(ctx, "test://missing"); err == nil { + t.Fatal("expected error for unknown resource") + } + + prompts, err := c.ListPrompts(ctx) + if err != nil { + t.Fatal(err) + } + if len(prompts) != 1 || prompts[0].Name != "p1" { + t.Fatalf("prompts = %+v", prompts) + } + msgs, err := c.GetPrompt(ctx, "p1", map[string]any{"x": "arg"}) + if err != nil { + t.Fatal(err) + } + if len(msgs) != 1 || msgs[0].Content != "prompt arg" { + t.Fatalf("messages = %+v", msgs) + } + + if err := c.Close(); err != nil { + t.Fatal(err) + } +} + +// HTTP error statuses surface as errors from call/notify. +func TestHTTPClientErrorPaths(t *testing.T) { + var n int + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + n++ + if n == 1 { + http.Error(w, "boom", http.StatusInternalServerError) + return + } + w.WriteHeader(http.StatusAccepted) // 202 for a request with an id + })) + defer ts.Close() + c := NewHTTPClient(ts.URL) + if _, err := c.ListTools(context.Background()); err == nil || !strings.Contains(err.Error(), "500") { + t.Fatalf("want 500 surfaced, got %v", err) + } + if _, err := c.ListTools(context.Background()); err == nil { + t.Fatal("want error for 202 response to a request") + } +} + +// decodeRPCResponse error branches: broken JSON, missing result. +func TestDecodeRPCResponseErrors(t *testing.T) { + if err := decodeRPCResponse([]byte("not json"), nil); err == nil { + t.Fatal("want malformed-response error") + } + if err := decodeRPCResponse([]byte(`{"jsonrpc":"2.0","id":1}`), nil); err == nil { + t.Fatal("want missing-result error") + } + var out struct{ A int } + if err := decodeRPCResponse([]byte(`{"jsonrpc":"2.0","id":1,"result":"str-not-obj"}`), &out); err == nil { + t.Fatal("want decode error") + } + // Happy path with nil out and a plain error envelope. + if err := decodeRPCResponse([]byte(`{"jsonrpc":"2.0","id":1,"result":{}}`), nil); err != nil { + t.Fatal(err) + } + if err := decodeRPCResponse([]byte(`{"jsonrpc":"2.0","id":1,"error":{"code":-1,"message":"x"}}`), nil); err == nil { + t.Fatal("want rpc error") + } +} + +// sseData handles "data:" without space and multi-line bodies. +func TestSSEDataVariants(t *testing.T) { + if got := string(sseData([]byte("data:{\"a\":1}\n"))); got != `{"a":1}` { + t.Fatalf("no-space variant = %q", got) + } + if got := string(sseData([]byte(": ping\ndata: one\ndata: two\n"))); got != "one" { + t.Fatalf("first data line = %q", got) + } +} + +// marshalHTTP never fails for JSONRPCError; exercise it directly anyway. +func TestMarshalHTTP(t *testing.T) { + b := marshalHTTP(NewJSONRPCError(1, -32600, "x")) + if !strings.Contains(string(b), "-32600") { + t.Fatalf("marshalled = %s", b) + } +} + +// Client interface compliance checks. +func TestClientInterfaceCompliance(t *testing.T) { + var _ Client = (*HTTPClient)(nil) + var _ Client = (*StdioClient)(nil) +} diff --git a/gomcp/coverage_protocol_test.go b/gomcp/coverage_protocol_test.go new file mode 100644 index 0000000..782f118 --- /dev/null +++ b/gomcp/coverage_protocol_test.go @@ -0,0 +1,72 @@ +package gomcp + +import ( + "encoding/json" + "testing" + "time" +) + +// protocolVersionFromParams: _meta wins, then top-level, then none. +func TestProtocolVersionFromParams(t *testing.T) { + cases := []struct { + params string + want string + }{ + {`{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28"},"protocolVersion":"2024-11-05"}`, "2026-07-28"}, + {`{"protocolVersion":"2025-03-26"}`, "2025-03-26"}, + {`{}`, ""}, + {``, ""}, + {`null`, ""}, + } + for _, tc := range cases { + var raw json.RawMessage + if tc.params != "" && tc.params != "null" { + raw = json.RawMessage(tc.params) + } else if tc.params == "null" { + raw = json.RawMessage("null") + } + if got := protocolVersionFromParams(raw); got != tc.want { + t.Errorf("params %s: got %q, want %q", tc.params, got, tc.want) + } + } +} + +// negotiateProtocolVersion fallback chain. +func TestNegotiateProtocolVersionFallback(t *testing.T) { + if got := negotiateProtocolVersion("2025-11-25", "fallback-ver"); got != "2025-11-25" { + t.Errorf("supported echo = %q", got) + } + if got := negotiateProtocolVersion("1999-01-01", "fallback-ver"); got != "fallback-ver" { + t.Errorf("fallback = %q", got) + } + if got := negotiateProtocolVersion("", ""); got != DefaultProtocolVersion { + t.Errorf("default = %q", got) + } +} + +// handlerContext override semantics. +func TestHandlerContextOverride(t *testing.T) { + s := &Server{} + base := time.Second + ctx, cancel := s.handlerContext(t.Context(), base) + defer cancel() + dl, ok := ctx.Deadline() + if !ok || time.Until(dl) > base { + t.Fatalf("deadline missing or too far: %v %v", ok, dl) + } + + // Negative override disables the timeout entirely. + ctx2, cancel2 := s.handlerContext(t.Context(), -1) + defer cancel2() + if _, ok := ctx2.Deadline(); ok { + t.Fatal("negative override should disable deadline") + } +} + +// decodeRPCResponse: literal null result is present (len 4) and accepted +// with a nil out; it is not treated as missing. +func TestDecodeRPCResponseNullResult(t *testing.T) { + if err := decodeRPCResponse([]byte(`{"jsonrpc":"2.0","id":1,"result":null}`), nil); err != nil { + t.Fatalf("null result with nil out should pass, got %v", err) + } +} diff --git a/gomcp/coverage_run_test.go b/gomcp/coverage_run_test.go new file mode 100644 index 0000000..ef2e31f --- /dev/null +++ b/gomcp/coverage_run_test.go @@ -0,0 +1,72 @@ +package gomcp + +import ( + "bufio" + "context" + "os" + "strings" + "testing" + "time" +) + +// runWithRealStdio swaps os.Stdin/os.Stdout for pipes, runs fn, feeds one +// request, and returns the server's response line. Restores the originals. +func runWithRealStdio(t *testing.T, fn func() error, request string) string { + t.Helper() + oldIn, oldOut := os.Stdin, os.Stdout + inR, inW, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + outR, outW, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + os.Stdin, os.Stdout = inR, outW + defer func() { os.Stdin, os.Stdout = oldIn, oldOut }() + + done := make(chan error, 1) + go func() { done <- fn() }() + + if _, err := inW.WriteString(request + "\n"); err != nil { + t.Fatal(err) + } + outCh := make(chan string, 1) + go func() { + line, _ := bufio.NewReader(outR).ReadString('\n') + outCh <- line + }() + + var resp string + select { + case resp = <-outCh: + case <-time.After(5 * time.Second): + t.Fatal("no response within 5s") + } + // EOF the input so fn returns. + inW.Close() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("Run did not return after stdin EOF") + } + return strings.TrimSpace(resp) +} + +func TestRunUsesOsStdio(t *testing.T) { + srv := NewServer("run-stdio", "1.0.0") + resp := runWithRealStdio(t, srv.Run, `{"jsonrpc":"2.0","id":1,"method":"ping"}`) + if !strings.Contains(resp, `"result"`) { + t.Fatalf("ping response = %q", resp) + } +} + +func TestRunContextUsesOsStdio(t *testing.T) { + srv := NewServer("runctx-stdio", "1.0.0") + resp := runWithRealStdio(t, func() error { + return srv.RunContext(context.Background()) + }, `{"jsonrpc":"2.0","id":7,"method":"ping"}`) + if !strings.Contains(resp, `"id":7`) { + t.Fatalf("ping response = %q", resp) + } +} diff --git a/gomcp/coverage_stdio_test.go b/gomcp/coverage_stdio_test.go new file mode 100644 index 0000000..ee210a7 --- /dev/null +++ b/gomcp/coverage_stdio_test.go @@ -0,0 +1,171 @@ +package gomcp + +import ( + "context" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +// fullServerSrc registers a tool, resource, and prompt for stdio tests. +const fullServerSrc = `package main + +import ( + "context" + "fmt" + + "github.com/BackendStack21/go-mcp/gomcp" +) + +func main() { + srv := gomcp.NewServer("stdio-full", "1.0.0") + srv.AddTool(gomcp.Tool{ + Name: "upper", + InputSchema: gomcp.InputSchema{Type: "object"}, + Handler: func(ctx context.Context, args map[string]any) (string, error) { + return fmt.Sprintf("%v", args["text"]), nil + }, + }) + srv.AddResource(gomcp.Resource{ + URI: "test://doc", + Name: "doc", + Handler: func(ctx context.Context) (string, error) { + return "res-body", nil + }, + }) + srv.AddPrompt(gomcp.Prompt{ + Name: "p1", + Handler: func(ctx context.Context, args map[string]any) ([]gomcp.PromptMessage, error) { + return []gomcp.PromptMessage{{Role: "user", Content: "prompt-body"}}, nil + }, + }) + srv.Run() +} +` + +func buildFullServer(t *testing.T) string { + t.Helper() + dir := t.TempDir() + src := filepath.Join(dir, "main.go") + if err := os.WriteFile(src, []byte(fullServerSrc), 0o644); err != nil { + t.Fatal(err) + } + return buildTestServer(t, dir, src) +} + +// Exercise the full StdioClient interface surface. +func TestStdioClientFullSurface(t *testing.T) { + if testing.Short() { + t.Skip("spawns subprocess") + } + c, err := NewStdioClient(buildFullServer(t)) + if err != nil { + t.Fatal(err) + } + defer c.Close() + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + if _, err := c.Initialize(ctx, "t", "0"); err != nil { + t.Fatal(err) + } + if err := c.NotifyInitialized(ctx); err != nil { + t.Fatal(err) + } + + tools, err := c.ListTools(ctx) + if err != nil { + t.Fatal(err) + } + if len(tools) != 1 { + t.Fatalf("tools = %+v", tools) + } + out, err := c.CallTool(ctx, "upper", map[string]any{"text": "hi"}) + if err != nil { + t.Fatal(err) + } + if out != "hi" { + t.Fatalf("out = %q", out) + } + + resources, err := c.ListResources(ctx) + if err != nil { + t.Fatal(err) + } + if len(resources) != 1 { + t.Fatalf("resources = %+v", resources) + } + body, err := c.ReadResource(ctx, "test://doc") + if err != nil { + t.Fatal(err) + } + if body != "res-body" { + t.Fatalf("body = %q", body) + } + if _, err := c.ReadResource(ctx, "test://missing"); err == nil { + t.Fatal("expected error for unknown resource") + } + + prompts, err := c.ListPrompts(ctx) + if err != nil { + t.Fatal(err) + } + if len(prompts) != 1 { + t.Fatalf("prompts = %+v", prompts) + } + msgs, err := c.GetPrompt(ctx, "p1", nil) + if err != nil { + t.Fatal(err) + } + if len(msgs) != 1 || msgs[0].Content != "prompt-body" { + t.Fatalf("messages = %+v", msgs) + } +} + +// StdioClient call after Close must fail, not hang. +func TestStdioClientCallAfterClose(t *testing.T) { + if testing.Short() { + t.Skip("spawns subprocess") + } + c, err := NewStdioClient(buildFullServer(t)) + if err != nil { + t.Fatal(err) + } + if err := c.Close(); err != nil && !strings.Contains(err.Error(), "signal") { + t.Fatalf("Close returned unexpected error: %v", err) + } + if _, err := c.ListTools(context.Background()); err == nil { + t.Fatal("expected error after Close") + } +} + +// A call against a dead server (killed externally) fails via failAllPending. +func TestStdioClientServerDeath(t *testing.T) { + if testing.Short() { + t.Skip("spawns subprocess") + } + c, err := NewStdioClient(buildFullServer(t)) + if err != nil { + t.Fatal(err) + } + defer c.Close() + if _, err := c.Initialize(context.Background(), "t", "0"); err != nil { + t.Fatal(err) + } + if err := c.cmd.Process.Kill(); err != nil { + t.Fatal(err) + } + deadline := time.Now().Add(10 * time.Second) + for { + _, err = c.ListTools(context.Background()) + if err != nil { + break // server death surfaced as expected + } + if time.Now().After(deadline) { + t.Fatal("calls kept succeeding after server death") + } + time.Sleep(50 * time.Millisecond) + } +} diff --git a/gomcp/http_e2e_test.go b/gomcp/http_e2e_test.go new file mode 100644 index 0000000..0bc3515 --- /dev/null +++ b/gomcp/http_e2e_test.go @@ -0,0 +1,100 @@ +package gomcp + +import ( + "context" + "fmt" + "net" + "testing" + "time" +) + +// TestE2EHTTPLoopback exercises the HTTP server and HTTP client together +// over a real loopback listener: handshake, tool listing, tool call, +// resource read, and prompt rendering end to end. +func TestE2EHTTPLoopback(t *testing.T) { + if testing.Short() { + t.Skip("binds a real listener") + } + + srv := NewServer("e2e-http", "1.0.0") + srv.AddTool(Tool{ + Name: "greet", + Description: "Greet someone", + InputSchema: InputSchema{Type: "object", Properties: map[string]Property{"name": {Type: "string"}}}, + Handler: func(ctx context.Context, args map[string]any) (string, error) { + return fmt.Sprintf("hello %v", args["name"]), nil + }, + }) + srv.AddResource(Resource{ + URI: "config://app", + Name: "app-config", + Description: "App configuration", + MimeType: "application/json", + Handler: func(ctx context.Context) (string, error) { + return `{"mode":"e2e"}`, nil + }, + }) + srv.AddPrompt(Prompt{ + Name: "review", + Description: "Review code", + Handler: func(ctx context.Context, args map[string]any) ([]PromptMessage, error) { + return []PromptMessage{{Role: "user", Content: fmt.Sprintf("review %v", args["file"])}}, nil + }, + }) + + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + done := make(chan struct{}) + go func() { defer close(done); _ = srv.Serve(l) }() + defer func() { <-done }() + defer l.Close() + + c := NewHTTPClient("http://" + l.Addr().String()) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + res, err := c.Initialize(ctx, "e2e-client", "0.1.0") + if err != nil { + t.Fatal(err) + } + if res.ServerInfo.Name != "e2e-http" { + t.Fatalf("serverInfo.name = %q", res.ServerInfo.Name) + } + if err := c.NotifyInitialized(ctx); err != nil { + t.Fatal(err) + } + + tools, err := c.ListTools(ctx) + if err != nil { + t.Fatal(err) + } + if len(tools) != 1 || tools[0].Name != "greet" { + t.Fatalf("tools = %+v", tools) + } + + out, err := c.CallTool(ctx, "greet", map[string]any{"name": "e2e"}) + if err != nil { + t.Fatal(err) + } + if out != "hello e2e" { + t.Fatalf("tool output = %q", out) + } + + body, err := c.ReadResource(ctx, "config://app") + if err != nil { + t.Fatal(err) + } + if body != `{"mode":"e2e"}` { + t.Fatalf("resource body = %q", body) + } + + msgs, err := c.GetPrompt(ctx, "review", map[string]any{"file": "main.go"}) + if err != nil { + t.Fatal(err) + } + if len(msgs) != 1 || msgs[0].Content != "review main.go" { + t.Fatalf("prompt messages = %+v", msgs) + } +} diff --git a/gomcp/httpauth_test.go b/gomcp/httpauth_test.go new file mode 100644 index 0000000..8c0db7a --- /dev/null +++ b/gomcp/httpauth_test.go @@ -0,0 +1,114 @@ +package gomcp + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +const testPing = `{"jsonrpc":"2.0","id":1,"method":"ping"}` + +func TestHTTPAuthTokenRequired(t *testing.T) { + srv := NewServer("auth-test", "1.0.0") + srv.SetAuthToken("secret-token") + ts := newAuthTestServer(t, srv) + + // No credentials -> 401 with WWW-Authenticate. + resp, body := postJSON(t, ts.URL, testPing, "application/json") + if resp.StatusCode != http.StatusUnauthorized { + t.Fatalf("status = %d, want 401 (body: %s)", resp.StatusCode, body) + } + if wa := resp.Header.Get("WWW-Authenticate"); !strings.HasPrefix(wa, "Bearer") { + t.Fatalf("WWW-Authenticate = %q, want Bearer challenge", wa) + } + + // Wrong token -> 401. + req, _ := http.NewRequest(http.MethodPost, ts.URL, strings.NewReader(testPing)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer wrong") + r2, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer r2.Body.Close() + if r2.StatusCode != http.StatusUnauthorized { + t.Fatalf("wrong token: status = %d, want 401", r2.StatusCode) + } + + // Correct token -> 200 with result. + req2, _ := http.NewRequest(http.MethodPost, ts.URL, strings.NewReader(testPing)) + req2.Header.Set("Content-Type", "application/json") + req2.Header.Set("Authorization", "Bearer secret-token") + r3, err := http.DefaultClient.Do(req2) + if err != nil { + t.Fatal(err) + } + defer r3.Body.Close() + if r3.StatusCode != http.StatusOK { + t.Fatalf("correct token: status = %d, want 200", r3.StatusCode) + } +} + +func TestHTTPAuthDisabledByDefault(t *testing.T) { + srv := NewServer("noauth", "1.0.0") + ts := newAuthTestServer(t, srv) + resp, _ := postJSON(t, ts.URL, testPing, "application/json") + if resp.StatusCode != http.StatusOK { + t.Fatalf("status = %d, want 200 without auth", resp.StatusCode) + } +} + +func TestHTTPOriginAllowlist(t *testing.T) { + srv := NewServer("origin-test", "1.0.0") + srv.SetAllowedOrigins([]string{"https://claude.ai", "http://localhost:3000"}) + ts := newAuthTestServer(t, srv) + + cases := []struct { + origin string + want int + }{ + {"https://claude.ai", http.StatusOK}, + {"http://localhost:3000", http.StatusOK}, + {"https://evil.example", http.StatusForbidden}, + {"", http.StatusOK}, // non-browser clients send no Origin + } + for _, tc := range cases { + req, _ := http.NewRequest(http.MethodPost, ts.URL, strings.NewReader(testPing)) + req.Header.Set("Content-Type", "application/json") + if tc.origin != "" { + req.Header.Set("Origin", tc.origin) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != tc.want { + t.Errorf("origin %q: status = %d, want %d", tc.origin, resp.StatusCode, tc.want) + } + } +} + +func TestHTTPOriginUncheckedByDefault(t *testing.T) { + srv := NewServer("noorigin", "1.0.0") + ts := newAuthTestServer(t, srv) + req, _ := http.NewRequest(http.MethodPost, ts.URL, strings.NewReader(testPing)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Origin", "https://anything.example") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("status = %d, want 200 (no allowlist set)", resp.StatusCode) + } +} + +func newAuthTestServer(t *testing.T, srv *Server) *httptest.Server { + t.Helper() + ts := httptest.NewServer(srv.Handler()) + t.Cleanup(ts.Close) + return ts +} diff --git a/gomcp/httpserver.go b/gomcp/httpserver.go new file mode 100644 index 0000000..02cd9e4 --- /dev/null +++ b/gomcp/httpserver.go @@ -0,0 +1,193 @@ +package gomcp + +import ( + "bytes" + "crypto/subtle" + "encoding/json" + "fmt" + "io" + "mime" + "net" + "net/http" + "strings" + "time" +) + +// Handler returns an http.Handler exposing the server over the MCP +// Streamable HTTP transport (protocol versions 2025-03-26 and later). +// +// Each POST body carries one JSON-RPC 2.0 message (a JSON object — arrays, +// i.e. JSON-RPC batches, are rejected with -32600, matching the stdio +// framing of one message per request). Responses are application/json, +// or text/event-stream (a single event) when the client's Accept header +// prefers SSE. Notifications are answered with 202 and no body. +// +// The transport is stateless: there is no session state, no server-initiated +// GET stream, and DELETE is not supported. GET, DELETE, and PUT return 405. +// Content-Type must be application/json; anything else is a 415. +// +// Message size and handler behavior follow the Server's existing settings +// (MaxRequestBytes, HandlerTimeout). Bad input is answered in-band with a +// JSON-RPC error and never fails the HTTP layer. +func (s *Server) Handler() http.Handler { + return http.HandlerFunc(s.serveHTTP) +} + +// Default HTTP server timeouts. These bound header reads, full request +// reads, and idle keep-alive connections so a slow or hostile client +// cannot pin goroutines and file descriptors indefinitely (slowloris). +const ( + defaultReadHeaderTimeout = 10 * time.Second + defaultReadTimeout = 30 * time.Second + defaultIdleTimeout = 120 * time.Second +) + +// newHTTPServer builds the http.Server used by ListenAndServe and Serve. +func (s *Server) newHTTPServer() *http.Server { + return &http.Server{ + Handler: s.Handler(), + ReadHeaderTimeout: defaultReadHeaderTimeout, + ReadTimeout: defaultReadTimeout, + IdleTimeout: defaultIdleTimeout, + } +} + +// ListenAndServe starts an HTTP server on addr serving the MCP Streamable +// HTTP transport. It blocks until the listener fails or the process exits. +func (s *Server) ListenAndServe(addr string) error { + hs := s.newHTTPServer() + hs.Addr = addr + return hs.ListenAndServe() +} + +// Serve accepts connections on an existing listener, serving the MCP +// Streamable HTTP transport. It blocks until the listener fails. +func (s *Server) Serve(l net.Listener) error { + return s.newHTTPServer().Serve(l) +} + +func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) { + s.mu.RLock() + authToken := s.authToken + allowedOrigins := s.allowedOrigins + s.mu.RUnlock() + + // Origin check first: a cross-site request never reaches dispatch. + if len(allowedOrigins) > 0 && r.Header.Get("Origin") != "" && !containsString(allowedOrigins, r.Header.Get("Origin")) { + http.Error(w, "origin not allowed", http.StatusForbidden) + return + } + + if authToken != "" { + const prefix = "Bearer " + authz := r.Header.Get("Authorization") + if !strings.HasPrefix(authz, prefix) || !constantTimeEqual(strings.TrimPrefix(authz, prefix), authToken) { + w.Header().Set("WWW-Authenticate", `Bearer realm="mcp"`) + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + } + + if r.Method != http.MethodPost { + w.Header().Set("Allow", http.MethodPost) + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + + // Parse the media type strictly: "application/json" (optionally with + // charset parameters), not look-alikes such as "application/jsonx". + if mt, _, err := mime.ParseMediaType(r.Header.Get("Content-Type")); err != nil || mt != "application/json" { + http.Error(w, "Content-Type must be application/json", http.StatusUnsupportedMediaType) + return + } + + // Cap the body at maxReq+1 bytes so an oversized request is detected + // without ever buffering it fully (constant memory, like readMessage). + // A negative MaxRequestBytes disables the cap for stdio, but the HTTP + // transport always enforces a floor so a remote peer can never OOM + // the process. + s.mu.RLock() + maxReq := s.MaxRequestBytes + s.mu.RUnlock() + if maxReq == 0 { + maxReq = DefaultMaxRequestBytes + } + if maxReq <= 0 || maxReq > DefaultMaxRequestBytes { + maxReq = DefaultMaxRequestBytes + } + if maxReq > 0 { + r.Body = http.MaxBytesReader(w, r.Body, maxReq+1) + } + body, err := io.ReadAll(r.Body) + if err != nil || (maxReq > 0 && int64(len(body)) > maxReq) { + writeRawJSONHTTP(w, marshalHTTP(NewJSONRPCError(nil, ErrCodeInvalidRequest, + fmt.Sprintf("Invalid Request: message exceeds maximum size of %d bytes", maxReq)))) + return + } + + // Reuse the stdio dispatch loop: feed it one newline-terminated + // message and capture the response it writes. + var out bytes.Buffer + in := bytes.NewReader(append(body, '\n')) + _ = s.RunWithIOContext(r.Context(), in, &out) + + resp := out.Bytes() + if len(bytes.TrimSpace(resp)) == 0 { + // The loop only writes nothing for a notification (no id). + w.WriteHeader(http.StatusAccepted) + return + } + + if wantsSSE(r.Header.Get("Accept")) { + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.WriteHeader(http.StatusOK) + fmt.Fprintf(w, "event: message\ndata: %s\n\n", bytes.TrimSuffix(resp, []byte("\n"))) + return + } + writeRawJSONHTTP(w, resp) +} + +// wantsSSE reports whether the client's Accept header asks for SSE in +// preference to plain JSON. +func wantsSSE(accept string) bool { + for _, part := range strings.Split(accept, ",") { + mt := strings.TrimSpace(strings.SplitN(part, ";", 2)[0]) + if mt == "text/event-stream" { + return true + } + if mt == "application/json" { + return false + } + } + return false +} + +func marshalHTTP(v any) []byte { + b, err := json.Marshal(v) + if err != nil { + return []byte(`{"jsonrpc":"2.0","id":null,"error":{"code":-32603,"message":"internal marshal error"}}`) + } + return b +} + +func writeRawJSONHTTP(w http.ResponseWriter, body []byte) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write(body) +} + +func containsString(list []string, s string) bool { + for _, v := range list { + if v == s { + return true + } + } + return false +} + +// constantTimeEqual compares tokens without leaking value information +// through timing (crypto/subtle). +func constantTimeEqual(a, b string) bool { + return subtle.ConstantTimeCompare([]byte(a), []byte(b)) == 1 +} diff --git a/gomcp/httpserver_low_test.go b/gomcp/httpserver_low_test.go new file mode 100644 index 0000000..9b870ff --- /dev/null +++ b/gomcp/httpserver_low_test.go @@ -0,0 +1,61 @@ +package gomcp + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +// LOW sweep regression tests. + +// application/jsonx and other look-alikes must be rejected (415), +// while application/json; charset=utf-8 is accepted. +func TestHTTPStrictContentType(t *testing.T) { + srv := NewServer("ct", "1.0.0") + ts := httptest.NewServer(srv.Handler()) + defer ts.Close() + + for _, ct := range []string{"application/jsonx", "text/json", "application/json-patch"} { + req, _ := http.NewRequest(http.MethodPost, ts.URL, strings.NewReader(testPing)) + req.Header.Set("Content-Type", ct) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusUnsupportedMediaType { + t.Errorf("Content-Type %q: status = %d, want 415", ct, resp.StatusCode) + } + } + + req, _ := http.NewRequest(http.MethodPost, ts.URL, strings.NewReader(testPing)) + req.Header.Set("Content-Type", "application/json; charset=utf-8") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Errorf("charset parameter: status = %d, want 200", resp.StatusCode) + } +} + +// A negative MaxRequestBytes disables the stdio cap, but HTTP must still +// enforce the DefaultMaxRequestBytes floor. +func TestHTTPBodyCapFloor(t *testing.T) { + srv := NewServer("floor", "1.0.0") + srv.MaxRequestBytes = -1 // "uncapped" for stdio + ts := httptest.NewServer(srv.Handler()) + defer ts.Close() + + big := `{"jsonrpc":"2.0","id":1,"method":"ping","params":{"pad":"` + + strings.Repeat("x", int(DefaultMaxRequestBytes)+1024) + `"}}` + resp, body := postJSON(t, ts.URL, big, "application/json") + if resp.StatusCode != http.StatusOK { + t.Fatalf("status = %d", resp.StatusCode) + } + if !strings.Contains(body, "-32600") { + t.Fatalf("want -32600 for oversized body despite negative MaxRequestBytes, got: %.100s", body) + } +} diff --git a/gomcp/httpserver_test.go b/gomcp/httpserver_test.go new file mode 100644 index 0000000..81ff601 --- /dev/null +++ b/gomcp/httpserver_test.go @@ -0,0 +1,209 @@ +package gomcp + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" +) + +func newHTTPTestServer(t *testing.T) (*Server, *httptest.Server) { + t.Helper() + srv := NewServer("http-test", "1.0.0") + srv.AddTool(Tool{ + Name: "echo", + Description: "Echo back the message", + InputSchema: InputSchema{Type: "object", Properties: map[string]Property{"message": {Type: "string"}}}, + Handler: func(ctx context.Context, args map[string]any) (string, error) { + return fmt.Sprintf("%v", args["message"]), nil + }, + }) + ts := httptest.NewServer(srv.Handler()) + t.Cleanup(ts.Close) + return srv, ts +} + +func postJSON(t *testing.T, url string, body string, accept string) (*http.Response, string) { + t.Helper() + req, err := http.NewRequest(http.MethodPost, url, strings.NewReader(body)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-Type", "application/json") + if accept != "" { + req.Header.Set("Accept", accept) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { resp.Body.Close() }) + b, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatal(err) + } + return resp, string(b) +} + +func TestHTTPInitialize(t *testing.T) { + _, ts := newHTTPTestServer(t) + resp, body := postJSON(t, ts.URL, `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"t","version":"0"}}}`, "application/json, text/event-stream") + if resp.StatusCode != http.StatusOK { + t.Fatalf("status = %d, body = %s", resp.StatusCode, body) + } + if ct := resp.Header.Get("Content-Type"); !strings.HasPrefix(ct, "application/json") { + t.Fatalf("Content-Type = %q", ct) + } + var out struct { + Result struct { + ProtocolVersion string `json:"protocolVersion"` + ServerInfo struct{ Name string } `json:"serverInfo"` + } `json:"result"` + } + if err := json.Unmarshal([]byte(body), &out); err != nil { + t.Fatalf("bad JSON: %v (%s)", err, body) + } + if out.Result.ProtocolVersion != "2025-11-25" { + t.Fatalf("protocolVersion = %q", out.Result.ProtocolVersion) + } + if out.Result.ServerInfo.Name != "http-test" { + t.Fatalf("serverInfo.name = %q", out.Result.ServerInfo.Name) + } +} + +func TestHTTPToolsCall(t *testing.T) { + _, ts := newHTTPTestServer(t) + // initialize first, as clients do + postJSON(t, ts.URL, `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"t","version":"0"}}}`, "application/json") + _, body := postJSON(t, ts.URL, `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"echo","arguments":{"message":"hi"}}}`, "application/json") + if !strings.Contains(body, `"hi"`) { + t.Fatalf("echo result missing: %s", body) + } +} + +func TestHTTPSSEResponse(t *testing.T) { + _, ts := newHTTPTestServer(t) + req, _ := http.NewRequest(http.MethodPost, ts.URL, strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"ping"}`)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "text/event-stream") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if ct := resp.Header.Get("Content-Type"); !strings.HasPrefix(ct, "text/event-stream") { + t.Fatalf("Content-Type = %q", ct) + } + b, _ := io.ReadAll(resp.Body) + if !strings.Contains(string(b), "data: ") || !strings.Contains(string(b), `"result"`) { + t.Fatalf("SSE body missing response: %q", string(b)) + } +} + +func TestHTTPMethodNotAllowed(t *testing.T) { + _, ts := newHTTPTestServer(t) + for _, method := range []string{http.MethodGet, http.MethodDelete, http.MethodPut} { + req, _ := http.NewRequest(method, ts.URL, nil) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusMethodNotAllowed { + t.Errorf("%s: status = %d, want 405", method, resp.StatusCode) + } + if resp.Header.Get("Allow") != http.MethodPost { + t.Errorf("%s: Allow = %q, want POST", method, resp.Header.Get("Allow")) + } + } +} + +func TestHTTPBadJSON(t *testing.T) { + _, ts := newHTTPTestServer(t) + resp, body := postJSON(t, ts.URL, `{not json`, "application/json") + if resp.StatusCode != http.StatusOK { + t.Fatalf("status = %d", resp.StatusCode) + } + if !strings.Contains(body, "-32700") { + t.Fatalf("want -32700 parse error, got: %s", body) + } +} + +func TestHTTPNotificationReturns202(t *testing.T) { + _, ts := newHTTPTestServer(t) + req, _ := http.NewRequest(http.MethodPost, ts.URL, strings.NewReader(`{"jsonrpc":"2.0","method":"notifications/initialized"}`)) + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusAccepted { + t.Fatalf("status = %d, want 202", resp.StatusCode) + } +} + +func TestHTTPOversizedBody(t *testing.T) { + srv, _ := newHTTPTestServer(t) + srv.MaxRequestBytes = 128 + ts := httptest.NewServer(srv.Handler()) + defer ts.Close() + big := `{"jsonrpc":"2.0","id":1,"method":"ping","params":{"pad":"` + strings.Repeat("x", 1024) + `"}}` + resp, body := postJSON(t, ts.URL, big, "application/json") + if resp.StatusCode != http.StatusOK { + t.Fatalf("status = %d", resp.StatusCode) + } + if !strings.Contains(body, "-32600") { + t.Fatalf("want -32600, got: %s", body) + } +} + +func TestHTTPWrongContentType(t *testing.T) { + _, ts := newHTTPTestServer(t) + resp, err := http.Post(ts.URL, "text/plain", bytes.NewReader([]byte(`{}`))) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusUnsupportedMediaType { + t.Fatalf("status = %d, want 415", resp.StatusCode) + } +} + +func TestHTTPListenAndServe(t *testing.T) { + // Grab a free port, close the listener, then serve on it for real. + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + addr := l.Addr().String() + l.Close() + + srv := NewServer("listen-test", "1.0.0") + go func() { _ = srv.ListenAndServe(addr) }() + + // Retry until the listener is up. + deadline := time.Now().Add(2 * time.Second) + for { + conn, err := net.DialTimeout("tcp", addr, 100*time.Millisecond) + if err == nil { + conn.Close() + break + } + if time.Now().After(deadline) { + t.Fatal("server never started listening") + } + time.Sleep(20 * time.Millisecond) + } + _, body := postJSON(t, "http://"+addr, `{"jsonrpc":"2.0","id":1,"method":"ping"}`, "application/json") + if !strings.Contains(body, `"result"`) { + t.Fatalf("ping failed over real listener: %s", body) + } +} diff --git a/gomcp/httpserver_timeouts_test.go b/gomcp/httpserver_timeouts_test.go new file mode 100644 index 0000000..f0e8091 --- /dev/null +++ b/gomcp/httpserver_timeouts_test.go @@ -0,0 +1,19 @@ +package gomcp + +import "testing" + +// TestHTTPServerTimeoutsSet guards against slowloris exposure: every +// serve path must build an http.Server with bounded header/idle timeouts. +func TestHTTPServerTimeoutsSet(t *testing.T) { + srv := NewServer("timeout-test", "1.0.0") + hs := srv.newHTTPServer() + if hs.ReadHeaderTimeout <= 0 { + t.Errorf("ReadHeaderTimeout = %v, want > 0", hs.ReadHeaderTimeout) + } + if hs.ReadTimeout <= 0 { + t.Errorf("ReadTimeout = %v, want > 0", hs.ReadTimeout) + } + if hs.IdleTimeout <= 0 { + t.Errorf("IdleTimeout = %v, want > 0", hs.IdleTimeout) + } +} diff --git a/gomcp/server.go b/gomcp/server.go index 3154ea0..1064b73 100644 --- a/gomcp/server.go +++ b/gomcp/server.go @@ -81,6 +81,14 @@ type Server struct { // 2026-07-28 results. Empty selects "public". Use "private" when // list or read results are caller-specific. CacheScope string + + // authToken, when set, is the exact Bearer token required by the + // HTTP transport. See SetAuthToken. + authToken string + + // allowedOrigins, when non-empty, is the Origin allowlist enforced + // by the HTTP transport. See SetAllowedOrigins. + allowedOrigins []string } // NewServer creates a new MCP server with the given name and version. @@ -114,6 +122,28 @@ func (s *Server) SetInstructions(text string) { s.instructions = text } +// SetAuthToken requires HTTP transport clients to present this exact token +// as `Authorization: Bearer `. Requests without valid credentials +// get 401 with a Bearer WWW-Authenticate challenge. Empty (default) allows +// anonymous access. Stdio transport is unaffected (trust comes from the +// user launching the process). +func (s *Server) SetAuthToken(token string) { + s.mu.Lock() + defer s.mu.Unlock() + s.authToken = token +} + +// SetAllowedOrigins restricts which Origins browser-based HTTP clients may +// come from (DNS-rebinding and CSRF defense, per the 2025-11-25 spec +// guidance). Requests with an Origin not on the list get 403. Requests +// without an Origin header (non-browser clients) are always allowed. +// An empty list (default) does no Origin checking. +func (s *Server) SetAllowedOrigins(origins []string) { + s.mu.Lock() + defer s.mu.Unlock() + s.allowedOrigins = origins +} + // AddTool registers a tool with the server. Tools are callable functions // that the AI client can invoke with arguments. Safe to call concurrently // with request handling. diff --git a/gomcp/stdio_client_race_test.go b/gomcp/stdio_client_race_test.go new file mode 100644 index 0000000..43fa50a --- /dev/null +++ b/gomcp/stdio_client_race_test.go @@ -0,0 +1,134 @@ +package gomcp + +import ( + "context" + "os" + "os/exec" + "path/filepath" + "testing" + "time" +) + +// slowServer replies to tool "slow" only after a delay, controlled by the +// argument. Everything else replies instantly. +const slowServerSrc = `package main + +import ( + "context" + "fmt" + "time" + + "github.com/BackendStack21/go-mcp/gomcp" +) + +func main() { + srv := gomcp.NewServer("slow", "1.0.0") + srv.AddTool(gomcp.Tool{ + Name: "slow", + InputSchema: gomcp.InputSchema{Type: "object"}, + Handler: func(ctx context.Context, args map[string]any) (string, error) { + ms := 50 + if v, ok := args["ms"].(float64); ok { + ms = int(v) + } + time.Sleep(time.Duration(ms) * time.Millisecond) + return fmt.Sprintf("slept %dms", ms), nil + }, + }) + srv.Run() +} +` + +func buildSlowServer(t *testing.T) string { + t.Helper() + dir := t.TempDir() + src := filepath.Join(dir, "main.go") + if err := os.WriteFile(src, []byte(slowServerSrc), 0o644); err != nil { + t.Fatal(err) + } + bin := filepath.Join(dir, "slowserver") + build := exec.Command("go", "build", "-o", bin, src) + build.Dir = "." + if out, err := build.CombinedOutput(); err != nil { + t.Fatalf("build: %v\n%s", err, out) + } + return bin +} + +// TestStdioClientTimeoutDoesNotDesync: a call that times out must not +// corrupt the stream — the NEXT call must still get its OWN response, +// not the abandoned one. +func TestStdioClientTimeoutDoesNotDesync(t *testing.T) { + if testing.Short() { + t.Skip("spawns subprocess") + } + bin := buildSlowServer(t) + c, err := NewStdioClient(bin) + if err != nil { + t.Fatal(err) + } + defer c.Close() + ctx := context.Background() + + if _, err := c.Initialize(ctx, "t", "0"); err != nil { + t.Fatal(err) + } + + // Call 1: 2s server-side delay, 100ms client deadline -> timeout. + tctx, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + if _, err := c.CallTool(tctx, "slow", map[string]any{"ms": 2000}); err == nil { + t.Fatal("expected timeout error") + } + + // Call 2 must succeed and receive ITS OWN result, not the stale one. + done := make(chan string, 1) + go func() { + out, err := c.CallTool(ctx, "slow", map[string]any{"ms": 10}) + if err != nil { + done <- "err:" + err.Error() + return + } + done <- out + }() + select { + case got := <-done: + if got != "slept 10ms" { + t.Fatalf("second call got stale/desynced response: %q", got) + } + case <-time.After(10 * time.Second): + t.Fatal("second call never returned — stream desynced after timeout") + } +} + +// TestStdioClientCloseDuringInflight: Close while a call is in flight must +// return without hanging and without racing the reader. +func TestStdioClientCloseDuringInflight(t *testing.T) { + if testing.Short() { + t.Skip("spawns subprocess") + } + bin := buildSlowServer(t) + c, err := NewStdioClient(bin) + if err != nil { + t.Fatal(err) + } + if _, err := c.Initialize(context.Background(), "t", "0"); err != nil { + t.Fatal(err) + } + + done := make(chan struct{}) + go func() { + defer close(done) + _, _ = c.CallTool(context.Background(), "slow", map[string]any{"ms": 3000}) + }() + time.Sleep(100 * time.Millisecond) + + closed := make(chan error, 1) + go func() { closed <- c.Close() }() + select { + case <-closed: + case <-time.After(5 * time.Second): + t.Fatal("Close hung while a call was in flight") + } + <-done +}