mcp: select HTTP transport from CLI

Amolith created

Change summary

cmd/cooked-mcp/main.go      | 40 +++++++++++-----
cmd/cooked-mcp/main_test.go | 93 ++++++++++++++++++++++++++++++++++++++
2 files changed, 120 insertions(+), 13 deletions(-)

Detailed changes

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 {

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" {