diff --git a/.gitignore b/.gitignore index 3ea1f95..7bf519e 100644 --- a/.gitignore +++ b/.gitignore @@ -41,6 +41,8 @@ config.local.json # Build output /bin/ /dist/ +/server +/agent # Temporary files /tmp/ diff --git a/cmd/server/main.go b/cmd/server/main.go index 82d59c7..ec3c215 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -1,13 +1,250 @@ package main import ( + "context" "fmt" + "log" + "log/slog" + "net/http" "os" + "os/signal" + "syscall" + "time" + "github.com/codeready-toolchain/cli-mcp-server/pkg/server" + "github.com/codeready-toolchain/cli-mcp-server/pkg/session" + "github.com/codeready-toolchain/cli-mcp-server/pkg/tools" "github.com/codeready-toolchain/cli-mcp-server/pkg/version" + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/spf13/cobra" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/client-go/kubernetes" + "k8s.io/client-go/rest" + "k8s.io/client-go/tools/clientcmd" ) +const shutdownTimeout = 310 * time.Second + func main() { fmt.Fprintf(os.Stderr, "cli-mcp-server %s (built %s)\n", version.Commit, version.BuildTime) - // TODO: Cobra root command, MCP server setup (SANDBOX-1814) + + var ( + address string + transport string + stateless bool + namespace string + sandboxImage string + kubeconfig string + hmacKeyFile string + idleTimeout time.Duration + warmPoolSize int + ) + + rootCmd := &cobra.Command{ + Use: "cli-mcp-server", + Short: "Sandboxed exec environment MCP server for LLM investigation", + RunE: func(_ *cobra.Command, _ []string) error { + return runServer(runConfig{ + address: address, + transport: transport, + stateless: stateless, + namespace: namespace, + sandboxImage: sandboxImage, + kubeconfig: kubeconfig, + hmacKeyFile: hmacKeyFile, + idleTimeout: idleTimeout, + warmPoolSize: warmPoolSize, + }) + }, + } + + rootCmd.Flags().StringVarP(&address, "address", "a", "localhost:8080", "Server address (host:port)") + rootCmd.Flags().StringVarP(&transport, "transport", "t", "stdio", "Transport (stdio, http)") + rootCmd.Flags().BoolVar(&stateless, "stateless", false, "Enable stateless mode (required for HTTP)") + rootCmd.Flags().StringVar(&namespace, "namespace", "tarsy", "Namespace for sandbox pods") + rootCmd.Flags().StringVar(&sandboxImage, "sandbox-image", "", "Container image for sandbox pods (required)") + rootCmd.Flags().StringVar(&kubeconfig, "kubeconfig", "", "Path to kubeconfig for sandbox pods") + rootCmd.Flags().StringVar(&hmacKeyFile, "hmac-key-file", "", "Path to HMAC shared secret file (required)") + rootCmd.Flags().DurationVar(&idleTimeout, "idle-timeout", 30*time.Minute, "Idle timeout for sandbox pods") + rootCmd.Flags().IntVar(&warmPoolSize, "warm-pool-size", 0, "Pre-warmed sandbox pods (0 = disabled)") + + if err := rootCmd.Execute(); err != nil { + os.Exit(1) + } +} + +type runConfig struct { + address string + transport string + stateless bool + namespace string + sandboxImage string + kubeconfig string + hmacKeyFile string + idleTimeout time.Duration + warmPoolSize int +} + +func runServer(cfg runConfig) error { + if err := server.ValidateTransportFlags(cfg.transport, cfg.stateless, cfg.address); err != nil { + return err + } + if cfg.sandboxImage == "" { + return fmt.Errorf("--sandbox-image is required") + } + if cfg.idleTimeout <= 0 { + return fmt.Errorf("--idle-timeout must be greater than zero") + } + if cfg.warmPoolSize < 0 { + return fmt.Errorf("--warm-pool-size must not be negative") + } + hmacKey, err := loadHMACKey(cfg.hmacKeyFile) + if err != nil { + return err + } + + logger := slog.New(slog.NewJSONHandler(os.Stderr, &slog.HandlerOptions{ + Level: slog.LevelInfo, + AddSource: true, + })).With("server", "cli-mcp-server") + + clientset, err := buildClientset(cfg.kubeconfig) + if err != nil { + return fmt.Errorf("failed to create kubernetes client: %w", err) + } + + sandboxCfg := session.DefaultConfig() + sandboxCfg.Image = cfg.sandboxImage + sandboxCfg.HMACKey = hmacKey + sandboxCfg.Namespace = cfg.namespace + sandboxCfg.IdleTimeout = cfg.idleTimeout + sandboxCfg.WarmPoolSize = cfg.warmPoolSize + + mgr, err := session.NewSessionManager(clientset, sandboxCfg, logger) + if err != nil { + return fmt.Errorf("failed to create session manager: %w", err) + } + + mcpServer := server.NewMCPServer("cli-mcp-server", cfg.stateless, logger) + tools.NewBashTool(mgr).RegisterWith(mcpServer) + + // Shared context cancelled on SIGTERM/SIGINT — stops background workers for both transports. + ctx, cancel := signal.NotifyContext(context.Background(), syscall.SIGTERM, syscall.SIGINT) + defer cancel() + + startCleanupLoop(ctx, mgr, logger) + if cfg.warmPoolSize > 0 { + mgr.StartPool(ctx) + } + + switch cfg.transport { + case "http": + return serveHTTP(ctx, cancel, cfg.address, mcpServer, mgr, clientset, cfg.namespace, logger) + case "stdio": + return serveStdio(ctx, mcpServer) + default: + return fmt.Errorf("unsupported transport: %s", cfg.transport) + } +} + +func serveHTTP(ctx context.Context, cancel context.CancelFunc, address string, mcpServer *mcp.Server, mgr *session.SessionManager, clientset kubernetes.Interface, namespace string, logger *slog.Logger) error { + checker := &k8sHealthChecker{clientset: clientset, namespace: namespace} + mux := server.NewMux(mcpServer, mgr, checker, logger) + + srv := &http.Server{ + Addr: address, + Handler: mux, + ReadHeaderTimeout: 10 * time.Second, // mitigate slowloris; leave ReadTimeout unset for MCP streams + } + serverErrChan := make(chan error, 1) + go func() { + log.Printf("listening on %s (HTTP, stateless)", address) + if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed { + serverErrChan <- err + } + }() + + select { + case <-ctx.Done(): + log.Println("received shutdown signal, draining...") + case err := <-serverErrChan: + cancel() + return fmt.Errorf("HTTP serve failed: %w", err) + } + + cancel() + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), shutdownTimeout) + defer shutdownCancel() + if err := srv.Shutdown(shutdownCtx); err != nil { + log.Printf("shutdown error: %v", err) + } + log.Println("shutdown complete") + return nil +} + +func serveStdio(ctx context.Context, mcpServer *mcp.Server) error { + log.Println("serving on stdio") + if err := mcpServer.Run(ctx, &mcp.StdioTransport{}); err != nil && ctx.Err() == nil { + return fmt.Errorf("stdio transport error: %w", err) + } + log.Println("shutdown complete") + return nil +} + +func startCleanupLoop(ctx context.Context, mgr *session.SessionManager, logger *slog.Logger) { + go func() { + ticker := time.NewTicker(5 * time.Minute) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + cleaned, err := mgr.CleanupStale(ctx) + if err != nil { + logger.Error("stale cleanup failed", "error", err) + } else if cleaned > 0 { + logger.Info("cleaned stale sessions", "count", cleaned) + } + } + } + }() +} + +func loadHMACKey(path string) (string, error) { + if path == "" { + return "", fmt.Errorf("--hmac-key-file is required") + } + data, err := os.ReadFile(path) + if err != nil { + return "", fmt.Errorf("failed to read HMAC key file: %w", err) + } + if len(data) == 0 { + return "", fmt.Errorf("HMAC key file is empty (zero bytes): %s", path) + } + return string(data), nil +} + +func buildClientset(kubeconfigPath string) (kubernetes.Interface, error) { + var config *rest.Config + var err error + if kubeconfigPath != "" { + config, err = clientcmd.BuildConfigFromFlags("", kubeconfigPath) + } else { + config, err = rest.InClusterConfig() + } + if err != nil { + return nil, err + } + return kubernetes.NewForConfig(config) +} + +type k8sHealthChecker struct { + clientset kubernetes.Interface + namespace string +} + +func (c *k8sHealthChecker) CheckHealth(ctx context.Context) error { + _, err := c.clientset.CoreV1().Namespaces().Get(ctx, c.namespace, metav1.GetOptions{}) + return err } diff --git a/go.mod b/go.mod index d171646..c1f5d48 100644 --- a/go.mod +++ b/go.mod @@ -3,8 +3,11 @@ module github.com/codeready-toolchain/cli-mcp-server go 1.25.12 require ( + github.com/codeready-toolchain/mcp-common v0.0.0-20260311065550-14ceedd27660 github.com/google/uuid v1.6.0 github.com/modelcontextprotocol/go-sdk v1.4.1 + github.com/prometheus/client_golang v1.22.0 + github.com/spf13/cobra v1.9.1 github.com/stretchr/testify v1.11.1 k8s.io/api v0.33.4 k8s.io/apimachinery v0.33.4 @@ -12,6 +15,8 @@ require ( ) require ( + github.com/beorn7/perks v1.0.1 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect github.com/emicklei/go-restful/v3 v3.13.0 // indirect github.com/fxamacker/cbor/v2 v2.9.0 // indirect @@ -23,6 +28,7 @@ require ( github.com/google/gnostic-models v0.6.9 // indirect github.com/google/go-cmp v0.7.0 // indirect github.com/google/jsonschema-go v0.4.2 // indirect + github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/josharian/intern v1.0.0 // indirect github.com/json-iterator/go v1.1.12 // indirect github.com/mailru/easyjson v0.7.7 // indirect @@ -30,15 +36,19 @@ require ( github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 // indirect + github.com/prometheus/client_model v0.6.1 // indirect + github.com/prometheus/common v0.62.0 // indirect + github.com/prometheus/procfs v0.15.1 // indirect github.com/segmentio/asm v1.1.3 // indirect github.com/segmentio/encoding v0.5.4 // indirect + github.com/spf13/pflag v1.0.6 // indirect github.com/x448/float16 v0.8.4 // indirect github.com/yosida95/uritemplate/v3 v3.0.2 // indirect - golang.org/x/net v0.49.0 // indirect + golang.org/x/net v0.55.0 // indirect golang.org/x/oauth2 v0.34.0 // indirect - golang.org/x/sys v0.40.0 // indirect - golang.org/x/term v0.39.0 // indirect - golang.org/x/text v0.33.0 // indirect + golang.org/x/sys v0.45.0 // indirect + golang.org/x/term v0.43.0 // indirect + golang.org/x/text v0.37.0 // indirect golang.org/x/time v0.14.0 // indirect google.golang.org/protobuf v1.36.12-0.20260120151049-f2248ac996af // indirect gopkg.in/evanphx/json-patch.v4 v4.13.0 // indirect diff --git a/go.sum b/go.sum index 7f9acd3..320c222 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,10 @@ +github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= +github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/codeready-toolchain/mcp-common v0.0.0-20260311065550-14ceedd27660 h1:j5vQQF19SBjzaFW3k4FPqLhL0EXmKJf8i+tEM3bSKuY= +github.com/codeready-toolchain/mcp-common v0.0.0-20260311065550-14ceedd27660/go.mod h1:kMt0I6nCr2uUERVj8UNmp+JjJ2rXiCREedsarmLzMQI= +github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= @@ -35,12 +42,16 @@ github.com/google/pprof v0.0.0-20241029153458-d1b30febd7db h1:097atOisP2aRj7vFgY github.com/google/pprof v0.0.0-20241029153458-d1b30febd7db/go.mod h1:vavhavw2zAxS5dIdcRluK6cSGGPlZynqzFM8NdvU144= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= +github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/josharian/intern v1.0.0 h1:vlS4z54oSdjm0bgjRigI+G1HpF+tI+9rE5LLzOg8HmY= github.com/josharian/intern v1.0.0/go.mod h1:5DoeVV0s6jJacbCEi61lwdGj/aVlrQvzHFFd8Hwg//Y= github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= +github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= +github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ= github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI= github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= @@ -48,6 +59,8 @@ github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= +github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= github.com/mailru/easyjson v0.7.7 h1:UGYAvKxe3sBsEDzO8ZeWOSlIQfWFlxbzLZe7hwFURr0= github.com/mailru/easyjson v0.7.7/go.mod h1:xzfreul335JAWq5oZzymOObrkdz5UnU4kGfJJLY9Nlc= github.com/modelcontextprotocol/go-sdk v1.4.1 h1:M4x9GyIPj+HoIlHNGpK2hq5o3BFhC+78PkEaldQRphc= @@ -67,14 +80,25 @@ github.com/onsi/gomega v1.35.1/go.mod h1:PvZbdDc8J6XJEpDK4HCuRBm8a6Fzp9/DmhC9C7y github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U= github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/prometheus/client_golang v1.22.0 h1:rb93p9lokFEsctTys46VnV1kLCDpVZ0a/Y92Vm0Zc6Q= +github.com/prometheus/client_golang v1.22.0/go.mod h1:R7ljNsLXhuQXYZYtw6GAE9AZg8Y7vEW5scdCXrWRXC0= +github.com/prometheus/client_model v0.6.1 h1:ZKSh/rekM+n3CeS952MLRAdFwIKqeY8b62p8ais2e9E= +github.com/prometheus/client_model v0.6.1/go.mod h1:OrxVMOVHjw3lKMa8+x6HeMGkHMQyHDk9E3jmP2AmGiY= +github.com/prometheus/common v0.62.0 h1:xasJaQlnWAeyHdUBeGjXmutelfJHWMRr+Fg4QszZ2Io= +github.com/prometheus/common v0.62.0/go.mod h1:vyBcEuLSvWos9B1+CyL7JZ2up+uFzXhkqml0W5zIY1I= +github.com/prometheus/procfs v0.15.1 h1:YagwOFzUgYfKKHX6Dr+sHT7km/hxC76UB0learggepc= +github.com/prometheus/procfs v0.15.1/go.mod h1:fB45yRUv8NstnjriLhBQLuOUt+WW4BsoGhij/e3PBqk= github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII= github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o= +github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= github.com/segmentio/asm v1.1.3 h1:WM03sfUOENvvKexOLp+pCqgb/WDjsi7EK8gIsICtzhc= github.com/segmentio/asm v1.1.3/go.mod h1:Ld3L4ZXGNcSLRg4JBsZ3//1+f/TjYl0Mzen/DQy1EJg= github.com/segmentio/encoding v0.5.4 h1:OW1VRern8Nw6ITAtwSZ7Idrl3MXCFwXHPgqESYfvNt0= github.com/segmentio/encoding v0.5.4/go.mod h1:HS1ZKa3kSN32ZHVZ7ZLPLXWvOVIiZtyJnO1gPH1sKt0= -github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA= -github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= +github.com/spf13/cobra v1.9.1 h1:CXSaggrXdbHK9CF+8ywj8Amf7PBRmPCOJugH954Nnlo= +github.com/spf13/cobra v1.9.1/go.mod h1:nDyEzZ8ogv936Cinf6g1RU9MRY64Ir93oCnqb9wxYW0= +github.com/spf13/pflag v1.0.6 h1:jFzHGLGAlb3ruxLB8MhbI6A8+AQX/2eW4qeyNZXNp2o= +github.com/spf13/pflag v1.0.6/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= @@ -101,8 +125,8 @@ golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= -golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o= -golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8= +golang.org/x/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= +golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= golang.org/x/oauth2 v0.34.0 h1:hqK/t4AKgbqWkdkcAeI8XLmbK+4m4G5YeQRrmiotGlw= golang.org/x/oauth2 v0.34.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -111,22 +135,22 @@ golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJ golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.40.0 h1:DBZZqJ2Rkml6QMQsZywtnjnnGvHza6BTfYFWY9kjEWQ= -golang.org/x/sys v0.40.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= -golang.org/x/term v0.39.0 h1:RclSuaJf32jOqZz74CkPA9qFuVTX7vhLlpfj/IGWlqY= -golang.org/x/term v0.39.0/go.mod h1:yxzUCTP/U+FzoxfdKmLaA0RV1WgE0VY7hXBwKtY/4ww= +golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= +golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4= +golang.org/x/term v0.43.0/go.mod h1:lrhlHNdQJHO+1qVYiHfFKVuVioJIheAc3fBSMFYEIsk= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= -golang.org/x/text v0.33.0 h1:B3njUFyqtHDUI5jMn1YIr5B0IE2U0qck04r6d4KPAxE= -golang.org/x/text v0.33.0/go.mod h1:LuMebE6+rBincTi9+xWTY8TztLzKHc/9C1uBCG27+q8= +golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= +golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI= golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE= golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA= -golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc= -golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg= +golang.org/x/tools v0.44.0 h1:UP4ajHPIcuMjT1GqzDWRlalUEoY+uzoZKnhOjbIPD2c= +golang.org/x/tools v0.44.0/go.mod h1:KA0AfVErSdxRZIsOVipbv3rQhVXTnlU6UhKxHd1seDI= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= diff --git a/pkg/server/server.go b/pkg/server/server.go new file mode 100644 index 0000000..ed78831 --- /dev/null +++ b/pkg/server/server.go @@ -0,0 +1,138 @@ +package server + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "net" + "net/http" + "time" + + "github.com/codeready-toolchain/cli-mcp-server/pkg/session" + "github.com/codeready-toolchain/cli-mcp-server/pkg/version" + "github.com/codeready-toolchain/mcp-common/pkg/middleware" + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/prometheus/client_golang/prometheus/promhttp" +) + +// SessionCleaner abstracts session cleanup for testability. +type SessionCleaner interface { + CleanupSession(ctx context.Context, sessionID string) error +} + +// HealthChecker abstracts the Kubernetes API reachability check. +type HealthChecker interface { + CheckHealth(ctx context.Context) error +} + +// NewMCPServer constructs an mcp.Server with mcp-common middleware. +func NewMCPServer(name string, stateless bool, logger *slog.Logger) *mcp.Server { + srv := mcp.NewServer( + &mcp.Implementation{Name: name, Version: version.Commit}, + &mcp.ServerOptions{ + Capabilities: &mcp.ServerCapabilities{ + Tools: &mcp.ToolCapabilities{ListChanged: !stateless}, + }, + Logger: logger, + }) + srv.AddReceivingMiddleware(middleware.NewMetricsMiddleware(name, logger)) + srv.AddReceivingMiddleware(middleware.NewLoggingMiddleware(logger)) + return srv +} + +// NewMux builds the HTTP mux with all operational endpoints. +func NewMux(mcpServer *mcp.Server, cleaner SessionCleaner, checker HealthChecker, logger *slog.Logger) *http.ServeMux { + mux := http.NewServeMux() + + mux.Handle("/mcp", mcp.NewStreamableHTTPHandler(func(*http.Request) *mcp.Server { + return mcpServer + }, &mcp.StreamableHTTPOptions{ + Stateless: true, + DisableLocalhostProtection: true, + })) + + mux.Handle("/metrics", promhttp.Handler()) + + mux.HandleFunc("/live", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]string{"status": "alive"}) + }) + + mux.HandleFunc("/health", func(w http.ResponseWriter, r *http.Request) { + checkCtx, cancel := context.WithTimeout(r.Context(), 5*time.Second) + defer cancel() + + w.Header().Set("Content-Type", "application/json") + if err := checker.CheckHealth(checkCtx); err != nil { + w.WriteHeader(http.StatusServiceUnavailable) + _ = json.NewEncoder(w).Encode(map[string]string{ + "status": "unhealthy", + "error": err.Error(), + "time": time.Now().Format(time.RFC3339), + }) + return + } + _ = json.NewEncoder(w).Encode(map[string]string{ + "status": "healthy", + "time": time.Now().Format(time.RFC3339), + }) + }) + + mux.HandleFunc("DELETE /sessions/{id}", func(w http.ResponseWriter, r *http.Request) { + sessionID := r.PathValue("id") + if sessionID == "" { + http.Error(w, "missing session ID", http.StatusBadRequest) + return + } + if err := session.ValidateSessionID(sessionID); err != nil { + http.Error(w, "invalid session ID format", http.StatusBadRequest) + return + } + if err := cleaner.CleanupSession(r.Context(), sessionID); err != nil { + logger.Error("session cleanup failed", "session_id", sessionID, "error", err) + http.Error(w, "failed to end session", http.StatusInternalServerError) + return + } + w.WriteHeader(http.StatusNoContent) + }) + + mux.HandleFunc("DELETE /sessions/", func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "missing session ID", http.StatusBadRequest) + }) + + return mux +} + +// IsLoopback returns true if the host part of addr is a loopback address. +func IsLoopback(addr string) bool { + host, _, err := net.SplitHostPort(addr) + if err != nil { + host = addr + } + if host == "localhost" { + return true + } + ip := net.ParseIP(host) + return ip != nil && ip.IsLoopback() +} + +// ValidateTransportFlags checks stateless/transport/address consistency. +func ValidateTransportFlags(transport string, stateless bool, address string) error { + switch transport { + case "http": + if !stateless { + return fmt.Errorf("--stateless is required for HTTP transport") + } + if !IsLoopback(address) { + return fmt.Errorf("HTTP transport requires a loopback --address (got %q); non-loopback addresses are not allowed without an authentication boundary", address) + } + case "stdio": + if stateless { + return fmt.Errorf("--stateless is not supported for stdio transport") + } + default: + return fmt.Errorf("unsupported transport %q (use \"stdio\" or \"http\")", transport) + } + return nil +} diff --git a/pkg/server/server_test.go b/pkg/server/server_test.go new file mode 100644 index 0000000..d1adff8 --- /dev/null +++ b/pkg/server/server_test.go @@ -0,0 +1,260 @@ +package server + +import ( + "context" + "encoding/json" + "errors" + "log/slog" + "net/http" + "net/http/httptest" + "os" + "testing" + + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type mockCleaner struct { + calledWith string + err error +} + +func (m *mockCleaner) CleanupSession(_ context.Context, sessionID string) error { + m.calledWith = sessionID + return m.err +} + +type mockHealthChecker struct { + err error +} + +func (m *mockHealthChecker) CheckHealth(_ context.Context) error { + return m.err +} + +func newTestLogger() *slog.Logger { + return slog.New(slog.NewTextHandler(os.Stderr, nil)) +} + +func TestSessionDelete(t *testing.T) { + t.Run("returns 204 and calls cleanup with session id", func(t *testing.T) { + // given + cleaner := &mockCleaner{} + checker := &mockHealthChecker{} + logger := newTestLogger() + mcpSrv := NewMCPServer("test", true, logger) + mux := NewMux(mcpSrv, cleaner, checker, logger) + + // when + req := httptest.NewRequestWithContext(context.Background(), http.MethodDelete, "/sessions/inv-abc", nil) + rr := httptest.NewRecorder() + mux.ServeHTTP(rr, req) + + // then + assert.Equal(t, http.StatusNoContent, rr.Code) + assert.Equal(t, "inv-abc", cleaner.calledWith) + }) + + t.Run("returns 400 for empty session id", func(t *testing.T) { + // given + cleaner := &mockCleaner{} + checker := &mockHealthChecker{} + logger := newTestLogger() + mcpSrv := NewMCPServer("test", true, logger) + mux := NewMux(mcpSrv, cleaner, checker, logger) + + // when + req := httptest.NewRequestWithContext(context.Background(), http.MethodDelete, "/sessions/", nil) + rr := httptest.NewRecorder() + mux.ServeHTTP(rr, req) + + // then + assert.Equal(t, http.StatusBadRequest, rr.Code) + assert.Equal(t, "", cleaner.calledWith) + }) + + t.Run("returns 500 with generic body on cleanup error", func(t *testing.T) { + // given + cleaner := &mockCleaner{err: errors.New("k8s pod delete failed: namespace xyz")} + checker := &mockHealthChecker{} + logger := newTestLogger() + mcpSrv := NewMCPServer("test", true, logger) + mux := NewMux(mcpSrv, cleaner, checker, logger) + + // when + req := httptest.NewRequestWithContext(context.Background(), http.MethodDelete, "/sessions/inv-fail", nil) + rr := httptest.NewRecorder() + mux.ServeHTTP(rr, req) + + // then + assert.Equal(t, http.StatusInternalServerError, rr.Code) + assert.Contains(t, rr.Body.String(), "failed to end session") + assert.NotContains(t, rr.Body.String(), "k8s pod delete failed") + assert.Equal(t, "inv-fail", cleaner.calledWith) + }) + + t.Run("returns 400 for invalid session id format", func(t *testing.T) { + // given + cleaner := &mockCleaner{} + checker := &mockHealthChecker{} + logger := newTestLogger() + mcpSrv := NewMCPServer("test", true, logger) + mux := NewMux(mcpSrv, cleaner, checker, logger) + + // when + req := httptest.NewRequestWithContext(context.Background(), http.MethodDelete, "/sessions/INVALID_ID!", nil) + rr := httptest.NewRecorder() + mux.ServeHTTP(rr, req) + + // then + assert.Equal(t, http.StatusBadRequest, rr.Code) + assert.Contains(t, rr.Body.String(), "invalid session ID format") + assert.Equal(t, "", cleaner.calledWith) + }) +} + +func TestLive(t *testing.T) { + // given + cleaner := &mockCleaner{} + checker := &mockHealthChecker{} + logger := newTestLogger() + mcpSrv := NewMCPServer("test", true, logger) + mux := NewMux(mcpSrv, cleaner, checker, logger) + + // when + req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/live", nil) + rr := httptest.NewRecorder() + mux.ServeHTTP(rr, req) + + // then + assert.Equal(t, http.StatusOK, rr.Code) + var body map[string]string + require.NoError(t, json.Unmarshal(rr.Body.Bytes(), &body)) + assert.Equal(t, "alive", body["status"]) +} + +func TestHealth(t *testing.T) { + t.Run("returns 200 when healthy", func(t *testing.T) { + // given + cleaner := &mockCleaner{} + checker := &mockHealthChecker{} + logger := newTestLogger() + mcpSrv := NewMCPServer("test", true, logger) + mux := NewMux(mcpSrv, cleaner, checker, logger) + + // when + req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/health", nil) + rr := httptest.NewRecorder() + mux.ServeHTTP(rr, req) + + // then + assert.Equal(t, http.StatusOK, rr.Code) + var body map[string]string + require.NoError(t, json.Unmarshal(rr.Body.Bytes(), &body)) + assert.Equal(t, "healthy", body["status"]) + }) + + t.Run("returns 503 when unhealthy", func(t *testing.T) { + // given + cleaner := &mockCleaner{} + checker := &mockHealthChecker{err: errors.New("connection refused")} + logger := newTestLogger() + mcpSrv := NewMCPServer("test", true, logger) + mux := NewMux(mcpSrv, cleaner, checker, logger) + + // when + req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/health", nil) + rr := httptest.NewRecorder() + mux.ServeHTTP(rr, req) + + // then + assert.Equal(t, http.StatusServiceUnavailable, rr.Code) + var body map[string]string + require.NoError(t, json.Unmarshal(rr.Body.Bytes(), &body)) + assert.Equal(t, "unhealthy", body["status"]) + assert.NotEmpty(t, body["error"]) + }) +} + +func TestBashToolRegistered(t *testing.T) { + // given + logger := newTestLogger() + mcpSrv := NewMCPServer("test", true, logger) + + ct, st := mcp.NewInMemoryTransports() + ss, err := mcpSrv.Connect(context.Background(), st, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = ss.Close() }) + + client := mcp.NewClient(&mcp.Implementation{Name: "test-client", Version: "0.1"}, nil) + cs, err := client.Connect(context.Background(), ct, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = cs.Close() }) + + // when — server has no tools registered yet (bash tool registration happens in cmd/server) + result, err := cs.ListTools(context.Background(), nil) + + // then + require.NoError(t, err) + assert.Empty(t, result.Tools) +} + +func TestIsLoopback(t *testing.T) { + tests := []struct { + name string + addr string + want bool + }{ + {"localhost with port", "localhost:8080", true}, + {"127.0.0.1 with port", "127.0.0.1:8080", true}, + {"::1 with port", "[::1]:8080", true}, + {"bare localhost", "localhost", true}, + {"bare 127.0.0.1", "127.0.0.1", true}, + {"non-loopback IP", "10.0.0.1:8080", false}, + {"non-loopback hostname", "myhost.example.com:8080", false}, + {"0.0.0.0", "0.0.0.0:8080", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // when + got := IsLoopback(tt.addr) + + // then + assert.Equal(t, tt.want, got) + }) + } +} + +func TestValidateTransportFlags(t *testing.T) { + tests := []struct { + name string + transport string + stateless bool + address string + wantErr string + }{ + {"http valid", "http", true, "localhost:8080", ""}, + {"http missing stateless", "http", false, "localhost:8080", "--stateless is required"}, + {"http non-loopback", "http", true, "10.0.0.1:8080", "loopback"}, + {"stdio valid", "stdio", false, "localhost:8080", ""}, + {"stdio with stateless", "stdio", true, "localhost:8080", "not supported for stdio"}, + {"invalid transport", "websocket", false, "", "unsupported transport"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // when + err := ValidateTransportFlags(tt.transport, tt.stateless, tt.address) + + // then + if tt.wantErr == "" { + assert.NoError(t, err) + } else { + require.Error(t, err) + assert.Contains(t, err.Error(), tt.wantErr) + } + }) + } +} diff --git a/pkg/session/manager.go b/pkg/session/manager.go index da0623a..d6d849d 100644 --- a/pkg/session/manager.go +++ b/pkg/session/manager.go @@ -102,7 +102,8 @@ func (m *SessionManager) SetCache(c *PodCache) { m.cache = c } -func validateSessionID(sessionID string) error { +// ValidateSessionID checks that sessionID matches RFC 1123 DNS label format. +func ValidateSessionID(sessionID string) error { if !sessionIDRegex.MatchString(sessionID) { return fmt.Errorf("invalid session ID %q: must match RFC 1123 DNS label format", sessionID) } @@ -112,7 +113,7 @@ func validateSessionID(sessionID string) error { // GetOrCreatePod resolves or creates a sandbox pod for the session. // Lookup order: cache → label-based K8s API discovery → idempotent create. func (m *SessionManager) GetOrCreatePod(ctx context.Context, sessionID string) (podIP string, err error) { - if err := validateSessionID(sessionID); err != nil { + if err := ValidateSessionID(sessionID); err != nil { return "", err } diff --git a/pkg/session/manager_test.go b/pkg/session/manager_test.go index 82245f3..5f49876 100644 --- a/pkg/session/manager_test.go +++ b/pkg/session/manager_test.go @@ -101,7 +101,7 @@ func TestValidateSessionID(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { // when - err := validateSessionID(tt.id) + err := ValidateSessionID(tt.id) // then if tt.wantErr {