From c9b9f3127719fee88026c9225253484c218de6c4 Mon Sep 17 00:00:00 2001 From: Amolith Date: Wed, 10 Jun 2026 18:23:20 -0600 Subject: [PATCH] mcp: select HTTP transport from CLI --- cmd/cooked-mcp/main.go | 40 +++++++++++----- cmd/cooked-mcp/main_test.go | 93 ++++++++++++++++++++++++++++++++++++- 2 files changed, 120 insertions(+), 13 deletions(-) diff --git a/cmd/cooked-mcp/main.go b/cmd/cooked-mcp/main.go index bb2b3229a7dd1a0ce876021f10c82a0bb14aada0..ed0e8e298edd78122c9605dcbc1656b3448f2579 100644 --- a/cmd/cooked-mcp/main.go +++ b/cmd/cooked-mcp/main.go @@ -28,6 +28,11 @@ type runOptions struct { httpAddr string } +type mcpRunner interface { + RunStdio(context.Context) error + RunHTTP(context.Context, string, string) error +} + func init() { buildVersion := builtVersion() if buildVersion != "" { @@ -74,15 +79,7 @@ func run() error { return err } - server := mcp.NewServer(client, version) - if err := server.RunStdio(ctx); err != nil { - if errors.Is(err, context.Canceled) { - return nil - } - return err - } - - return nil + return runMCPServer(ctx, mcp.NewServer(client, version), config, options.transport) } func parseRunOptions(args []string) (runOptions, error) { @@ -90,7 +87,7 @@ func parseRunOptions(args []string) (runOptions, error) { var options runOptions flagSet.StringVar(&options.configPath, "config", "", "path to TOML configuration file") - flagSet.StringVar(&options.transport, "transport", "stdio", "MCP transport: stdio") + flagSet.StringVar(&options.transport, "transport", "stdio", "MCP transport: stdio or http") flagSet.StringVar(&options.httpAddr, "http-addr", "", "MCP HTTP listen address") if err := flagSet.Parse(args); err != nil { @@ -101,11 +98,30 @@ func parseRunOptions(args []string) (runOptions, error) { } func validateTransport(transport string) error { - if transport == "stdio" { + switch transport { + case "stdio", "http": + return nil + default: + return fmt.Errorf("unsupported transport %q (supported: stdio, http)", transport) + } +} + +func runMCPServer(ctx context.Context, server mcpRunner, config appconfig.Config, transport string) error { + var err error + switch transport { + case "stdio": + err = server.RunStdio(ctx) + case "http": + err = server.RunHTTP(ctx, config.HTTPAddr, config.HTTPToken) + default: + return fmt.Errorf("unsupported transport %q (supported: stdio, http)", transport) + } + + if errors.Is(err, context.Canceled) { return nil } - return fmt.Errorf("unsupported transport %q (supported: stdio)", transport) + return err } func builtVersion() string { diff --git a/cmd/cooked-mcp/main_test.go b/cmd/cooked-mcp/main_test.go index 1983dac7068e6d796bd1aff2eaaa9dc706f8299f..846cbfdc3b646e591ad299b41f1a2bc58bb358ad 100644 --- a/cmd/cooked-mcp/main_test.go +++ b/cmd/cooked-mcp/main_test.go @@ -5,10 +5,36 @@ package main import ( + "context" + "errors" "runtime/debug" "testing" + + "git.secluded.site/cooked-mcp/internal/appconfig" ) +type fakeMCPRunner struct { + stdioCalls int + httpCalls int + httpAddr string + httpToken string + err error +} + +func (r *fakeMCPRunner) RunStdio(context.Context) error { + r.stdioCalls++ + + return r.err +} + +func (r *fakeMCPRunner) RunHTTP(_ context.Context, addr, token string) error { + r.httpCalls++ + r.httpAddr = addr + r.httpToken = token + + return r.err +} + // server.CONFIG.8 func TestParseRunOptionsACIDServerConfig8ParsesHTTPAddr(t *testing.T) { got, err := parseRunOptions([]string{"--http-addr", "127.0.0.1:9999"}) @@ -31,17 +57,82 @@ func TestValidateTransportACIDServerTransport4AcceptsStdio(t *testing.T) { } } +// server.TRANSPORT.2-2 +func TestValidateTransportACIDServerTransport2_2AcceptsHTTP(t *testing.T) { + if err := validateTransport("http"); err != nil { + t.Fatalf("validateTransport() error = %v, want nil", err) + } +} + // server.TRANSPORT.4 func TestValidateTransportACIDServerTransport4RejectsUnknownTransport(t *testing.T) { err := validateTransport("bogus") if err == nil { t.Fatal("validateTransport() error = nil, want unsupported transport error") } - if err.Error() != `unsupported transport "bogus" (supported: stdio)` { + if err.Error() != `unsupported transport "bogus" (supported: stdio, http)` { t.Fatalf("validateTransport() error = %q", err.Error()) } } +// server.TRANSPORT.2-1 +func TestRunMCPServerACIDServerTransport2_1SelectsStdio(t *testing.T) { + runner := &fakeMCPRunner{} + err := runMCPServer(context.Background(), runner, appconfig.Config{}, "stdio") + if err != nil { + t.Fatalf("runMCPServer() error = %v", err) + } + + if runner.stdioCalls != 1 { + t.Fatalf("stdio calls = %d, want 1", runner.stdioCalls) + } + if runner.httpCalls != 0 { + t.Fatalf("http calls = %d, want 0", runner.httpCalls) + } +} + +// server.TRANSPORT.2-2 server.TRANSPORT.3 +func TestRunMCPServerACIDServerTransport2_2SelectsHTTP(t *testing.T) { + runner := &fakeMCPRunner{} + err := runMCPServer(context.Background(), runner, appconfig.Config{ + HTTPAddr: "127.0.0.1:8123", + HTTPToken: "secret-token", + }, "http") + if err != nil { + t.Fatalf("runMCPServer() error = %v", err) + } + + if runner.httpCalls != 1 { + t.Fatalf("http calls = %d, want 1", runner.httpCalls) + } + if runner.httpAddr != "127.0.0.1:8123" { + t.Fatalf("http addr = %q, want configured address", runner.httpAddr) + } + if runner.httpToken != "secret-token" { + t.Fatal("HTTP token was not passed to RunHTTP") + } + if runner.stdioCalls != 0 { + t.Fatalf("stdio calls = %d, want 0", runner.stdioCalls) + } +} + +func TestRunMCPServerTreatsContextCancellationAsCleanShutdown(t *testing.T) { + runner := &fakeMCPRunner{err: context.Canceled} + err := runMCPServer(context.Background(), runner, appconfig.Config{}, "http") + if err != nil { + t.Fatalf("runMCPServer() error = %v, want nil", err) + } +} + +func TestRunMCPServerReturnsTransportErrors(t *testing.T) { + wantErr := errors.New("transport failed") + runner := &fakeMCPRunner{err: wantErr} + err := runMCPServer(context.Background(), runner, appconfig.Config{}, "stdio") + if !errors.Is(err, wantErr) { + t.Fatalf("runMCPServer() error = %v, want %v", err, wantErr) + } +} + func TestVersionFromBuildInfoPrefersModuleVersion(t *testing.T) { got := versionFromBuildInfo(&debug.BuildInfo{Main: debug.Module{Version: "1.2.3"}}) if got != "1.2.3" {