From 60d0b4d00dd056eeef11154547eb9effac0383e0 Mon Sep 17 00:00:00 2001 From: Apoorv Raj Saxena Date: Fri, 23 Jan 2026 13:46:27 +0400 Subject: [PATCH] Add stdio transport for subprocess MCP communication --- cmd/proxy/main.go | 46 ++- internal/policy/compiler/capability.go | 6 +- internal/policy/compiler/templates.go | 1 + internal/policy/compiler/validation.go | 12 +- internal/transport/sse/handler.go | 6 +- internal/transport/stdio/reader.go | 61 ++++ internal/transport/stdio/server.go | 234 +++++++++++++ internal/transport/stdio/server_test.go | 418 ++++++++++++++++++++++++ internal/transport/stdio/writer.go | 43 +++ internal/transport/transport.go | 13 +- 10 files changed, 811 insertions(+), 29 deletions(-) create mode 100644 internal/transport/stdio/reader.go create mode 100644 internal/transport/stdio/server.go create mode 100644 internal/transport/stdio/server_test.go create mode 100644 internal/transport/stdio/writer.go diff --git a/cmd/proxy/main.go b/cmd/proxy/main.go index ea69899..62894a1 100644 --- a/cmd/proxy/main.go +++ b/cmd/proxy/main.go @@ -5,6 +5,7 @@ import ( "encoding/json" "flag" "fmt" + "io" "os" "os/signal" "strings" @@ -17,7 +18,9 @@ import ( "github.com/agentfacts/mcp-proxy/internal/policy" "github.com/agentfacts/mcp-proxy/internal/router" "github.com/agentfacts/mcp-proxy/internal/session" + "github.com/agentfacts/mcp-proxy/internal/transport" "github.com/agentfacts/mcp-proxy/internal/transport/sse" + "github.com/agentfacts/mcp-proxy/internal/transport/stdio" "github.com/agentfacts/mcp-proxy/internal/upstream" "github.com/rs/zerolog" "github.com/rs/zerolog/log" @@ -34,7 +37,7 @@ type Application struct { cfg *config.Config sessionManager *session.Manager router *router.Router - sseServer *sse.Server + transport transport.Transport upstreamClient *upstream.Client policyEngine *policy.Engine auditStore *audit.Store @@ -270,11 +273,19 @@ func newApplication(cfg *config.Config) (*Application, error) { }, nil }) - // Initialize SSE server - app.sseServer = sse.NewServer(cfg.Server, cfg.Agent, app.sessionManager) + // Initialize transport based on config + switch cfg.Server.Transport { + case "sse": + app.transport = sse.NewServer(cfg.Server, cfg.Agent, app.sessionManager) + case "stdio": + stdioServer := stdio.NewServer(cfg.Agent, app.sessionManager) + app.transport = stdioServer + default: + return nil, fmt.Errorf("unknown transport: %s", cfg.Server.Transport) + } // Set up message handler to use router - app.sseServer.SetMessageHandler(app.handleMessage) + app.transport.SetMessageHandler(app.handleMessage) // Initialize observability app.metrics = observability.NewMetrics("mcp_proxy") @@ -347,9 +358,9 @@ func (app *Application) Start(ctx context.Context) error { } } - // Start SSE server - if err := app.sseServer.Start(ctx); err != nil { - return fmt.Errorf("failed to start SSE server: %w", err) + // Start transport server + if err := app.transport.Start(ctx); err != nil { + return fmt.Errorf("failed to start %s server: %w", app.transport.Name(), err) } // Start observability server @@ -375,9 +386,9 @@ func (app *Application) Stop(ctx context.Context) error { log.Error().Err(err).Msg("Error stopping observability server") } - // Stop SSE server first (stop accepting new connections) - if err := app.sseServer.Stop(ctx); err != nil { - log.Error().Err(err).Msg("Error stopping SSE server") + // Stop transport server first (stop accepting new connections) + if err := app.transport.Stop(ctx); err != nil { + log.Error().Err(err).Msg("Error stopping transport server") } // Disconnect from upstream @@ -417,16 +428,27 @@ func initLogger(cfg config.LoggingConfig) { } zerolog.SetGlobalLevel(level) + // Determine output destination + var output io.Writer = os.Stdout + switch cfg.Output { + case "stderr": + output = os.Stderr + case "stdout", "": + output = os.Stdout + // File output could be added here if needed + } + // Configure output format if cfg.Format == "text" { log.Logger = log.Output(zerolog.ConsoleWriter{ - Out: os.Stdout, + Out: output, TimeFormat: time.RFC3339, }) } else { // JSON format (default) zerolog.TimeFieldFormat = time.RFC3339Nano + log.Logger = log.Output(output) } - log.Debug().Str("level", cfg.Level).Str("format", cfg.Format).Msg("Logger initialized") + log.Debug().Str("level", cfg.Level).Str("format", cfg.Format).Str("output", cfg.Output).Msg("Logger initialized") } diff --git a/internal/policy/compiler/capability.go b/internal/policy/compiler/capability.go index 410dd1c..7155cc4 100644 --- a/internal/policy/compiler/capability.go +++ b/internal/policy/compiler/capability.go @@ -33,9 +33,9 @@ func CompileCapabilityRules(rules []RuleDefinition, policyName string) (string, // Replace placeholders in message message = replacePlaceholders(message, map[string]string{ - "agent.id": "' + input.agent.id + '", - "tool": tool, - "required": capability, + "agent.id": "' + input.agent.id + '", + "tool": tool, + "required": capability, }) data := CapabilityData{ diff --git a/internal/policy/compiler/templates.go b/internal/policy/compiler/templates.go index 2d85fd8..6765ef8 100644 --- a/internal/policy/compiler/templates.go +++ b/internal/policy/compiler/templates.go @@ -9,6 +9,7 @@ import ( // Templates for Rego code generation. var templates *template.Template +//nolint:gochecknoinits // template initialization is idiomatic with init func init() { templates = template.New("rego").Funcs(template.FuncMap{ "quote": quoteString, diff --git a/internal/policy/compiler/validation.go b/internal/policy/compiler/validation.go index 1824a7b..593d392 100644 --- a/internal/policy/compiler/validation.go +++ b/internal/policy/compiler/validation.go @@ -157,13 +157,11 @@ func (v *Validator) validateRateLimitRule(rule *RuleDefinition) error { return fmt.Errorf("'limit' must be a number") } - // Either agent_id or agent_pattern should be specified - _, hasID := rule.Conditions["agent_id"] - _, hasPattern := rule.Conditions["agent_pattern"] - - if !hasID && !hasPattern { - // Allow default rate limit without agent specification - } + // Either agent_id or agent_pattern can be specified (both optional for default rate limits) + // No validation needed - all combinations are valid: + // - No agent specification = default rate limit for all agents + // - agent_id = rate limit for specific agent + // - agent_pattern = rate limit for agents matching pattern if window, ok := rule.Conditions["window"]; ok { w, ok := window.(string) diff --git a/internal/transport/sse/handler.go b/internal/transport/sse/handler.go index ee11721..4358f36 100644 --- a/internal/transport/sse/handler.go +++ b/internal/transport/sse/handler.go @@ -1,7 +1,6 @@ package sse import ( - "context" "encoding/json" "fmt" "io" @@ -11,11 +10,12 @@ import ( "github.com/agentfacts/mcp-proxy/internal/config" "github.com/agentfacts/mcp-proxy/internal/session" + "github.com/agentfacts/mcp-proxy/internal/transport" "github.com/rs/zerolog/log" ) -// MessageHandler is the callback for processing incoming MCP messages. -type MessageHandler func(ctx context.Context, sess *session.Session, message []byte) ([]byte, error) +// MessageHandler is an alias for the transport.MessageHandler type. +type MessageHandler = transport.MessageHandler // Handler handles SSE connections and messages. type Handler struct { diff --git a/internal/transport/stdio/reader.go b/internal/transport/stdio/reader.go new file mode 100644 index 0000000..98925c9 --- /dev/null +++ b/internal/transport/stdio/reader.go @@ -0,0 +1,61 @@ +package stdio + +import ( + "bufio" + "encoding/json" + "fmt" + "io" +) + +// DefaultMaxMessageSize is the default maximum size of a single JSON message (1MB). +const DefaultMaxMessageSize = 1024 * 1024 + +// Reader handles reading newline-delimited JSON messages from stdin. +type Reader struct { + scanner *bufio.Scanner + maxMessageSize int +} + +// NewReader creates a new Reader for the given input stream. +func NewReader(in io.Reader) *Reader { + return NewReaderWithMaxSize(in, DefaultMaxMessageSize) +} + +// NewReaderWithMaxSize creates a new Reader with a custom max message size. +func NewReaderWithMaxSize(in io.Reader, maxSize int) *Reader { + scanner := bufio.NewScanner(in) + scanner.Buffer(make([]byte, 0, 64*1024), maxSize) + + return &Reader{ + scanner: scanner, + maxMessageSize: maxSize, + } +} + +// ReadMessage reads the next JSON message from the input. +// Returns io.EOF when there are no more messages. +func (r *Reader) ReadMessage() ([]byte, error) { + if !r.scanner.Scan() { + if err := r.scanner.Err(); err != nil { + return nil, fmt.Errorf("reading input: %w", err) + } + return nil, io.EOF + } + + line := r.scanner.Bytes() + if len(line) == 0 { + // Skip empty lines + return r.ReadMessage() + } + + // Make a copy since scanner reuses the buffer + msg := make([]byte, len(line)) + copy(msg, line) + + // Validate JSON + if !json.Valid(msg) { + return nil, fmt.Errorf("invalid JSON message") + } + + return msg, nil +} diff --git a/internal/transport/stdio/server.go b/internal/transport/stdio/server.go new file mode 100644 index 0000000..2ca2f18 --- /dev/null +++ b/internal/transport/stdio/server.go @@ -0,0 +1,234 @@ +package stdio + +import ( + "context" + "encoding/json" + "fmt" + "io" + "os" + "sync" + + "github.com/agentfacts/mcp-proxy/internal/config" + "github.com/agentfacts/mcp-proxy/internal/session" + "github.com/agentfacts/mcp-proxy/internal/transport" + "github.com/rs/zerolog/log" +) + +// MessageHandler is an alias for the transport.MessageHandler type. +type MessageHandler = transport.MessageHandler + +// Server implements the stdio transport for MCP. +// It reads JSON-RPC messages from stdin and writes responses to stdout. +type Server struct { + agentCfg config.AgentConfig + sessionManager *session.Manager + messageHandler MessageHandler + session *session.Session // Single session for stdio + + // I/O streams (configurable for testing) + stdin io.Reader + stdout io.Writer + + // Lifecycle + mu sync.RWMutex + started bool + done chan struct{} + wg sync.WaitGroup +} + +// NewServer creates a new stdio transport server. +func NewServer(agentCfg config.AgentConfig, sessionMgr *session.Manager) *Server { + return &Server{ + agentCfg: agentCfg, + sessionManager: sessionMgr, + stdin: os.Stdin, + stdout: os.Stdout, + done: make(chan struct{}), + } +} + +// NewServerWithIO creates a new stdio transport server with custom I/O streams. +// This is primarily useful for testing. +func NewServerWithIO(agentCfg config.AgentConfig, sessionMgr *session.Manager, stdin io.Reader, stdout io.Writer) *Server { + return &Server{ + agentCfg: agentCfg, + sessionManager: sessionMgr, + stdin: stdin, + stdout: stdout, + done: make(chan struct{}), + } +} + +// SetMessageHandler sets the callback for processing incoming messages. +func (s *Server) SetMessageHandler(h MessageHandler) { + s.messageHandler = h +} + +// Start begins reading from stdin and processing messages. +func (s *Server) Start(ctx context.Context) error { + s.mu.Lock() + if s.started { + s.mu.Unlock() + return fmt.Errorf("server already started") + } + s.started = true + s.mu.Unlock() + + // Create a single session for the entire process lifetime + sess, err := s.sessionManager.Create(ctx) + if err != nil { + return fmt.Errorf("failed to create session: %w", err) + } + s.session = sess + + // Set default agent info from config + s.session.SetAgent(s.agentCfg.ID, s.agentCfg.Name, s.agentCfg.Capabilities) + s.session.SetClientInfo("stdio", "stdio-client") + + log.Info(). + Str("session_id", s.session.ID). + Str("transport", "stdio"). + Msg("Stdio server started") + + // Start the read loop in a goroutine + s.wg.Add(1) + go s.readLoop(ctx) + + return nil +} + +// Stop gracefully shuts down the server. +func (s *Server) Stop(ctx context.Context) error { + s.mu.Lock() + if !s.started { + s.mu.Unlock() + return nil + } + s.started = false + s.mu.Unlock() + + // Signal shutdown + close(s.done) + + // Close session + if s.session != nil { + s.session.Close() + s.sessionManager.Delete(s.session.ID) + } + + // Wait for read loop to finish with timeout + done := make(chan struct{}) + go func() { + s.wg.Wait() + close(done) + }() + + select { + case <-done: + log.Info().Msg("Stdio server stopped") + case <-ctx.Done(): + log.Warn().Msg("Stdio server stop timed out") + } + + return nil +} + +// Name returns the transport name. +func (s *Server) Name() string { + return "stdio" +} + +// readLoop continuously reads messages from stdin and processes them. +func (s *Server) readLoop(ctx context.Context) { + defer s.wg.Done() + + reader := NewReader(s.stdin) + writer := NewWriter(s.stdout) + + for { + select { + case <-s.done: + return + case <-ctx.Done(): + return + default: + } + + // Read next message + msg, err := reader.ReadMessage() + if err != nil { + if err == io.EOF { + log.Info().Msg("Stdin closed (EOF), shutting down") + return + } + log.Error().Err(err).Msg("Error reading message") + s.writeError(writer, nil, -32700, "Parse error") + continue + } + + // Increment request count + s.session.IncrementRequestCount() + + log.Debug(). + Str("session_id", s.session.ID). + Int("body_size", len(msg)). + Int("request_count", s.session.GetRequestCount()). + Msg("Received MCP message") + + // Process message through handler + var response []byte + if s.messageHandler != nil { + response, err = s.messageHandler(ctx, s.session, msg) + if err != nil { + log.Error().Err(err).Str("session_id", s.session.ID).Msg("Message handler error") + // Try to extract request ID for error response + id := extractRequestID(msg) + s.writeError(writer, id, -32603, "Internal error") + continue + } + } else { + // No handler configured - echo back for testing + response = msg + } + + // Write response + if response != nil { + if err := writer.Write(response); err != nil { + log.Error().Err(err).Msg("Error writing response") + } + } + } +} + +// writeError writes a JSON-RPC error response to stdout. +func (s *Server) writeError(writer *Writer, id interface{}, code int, message string) { + errResp := map[string]interface{}{ + "jsonrpc": "2.0", + "id": id, + "error": map[string]interface{}{ + "code": code, + "message": message, + }, + } + + data, err := json.Marshal(errResp) + if err != nil { + log.Error().Err(err).Msg("Failed to marshal error response") + return + } + + if err := writer.Write(data); err != nil { + log.Error().Err(err).Msg("Failed to write error response") + } +} + +// extractRequestID attempts to extract the request ID from a JSON-RPC message. +func extractRequestID(msg []byte) interface{} { + var req struct { + ID interface{} `json:"id"` + } + if err := json.Unmarshal(msg, &req); err != nil { + return nil + } + return req.ID +} diff --git a/internal/transport/stdio/server_test.go b/internal/transport/stdio/server_test.go new file mode 100644 index 0000000..e430318 --- /dev/null +++ b/internal/transport/stdio/server_test.go @@ -0,0 +1,418 @@ +package stdio + +import ( + "bytes" + "context" + "encoding/json" + "io" + "strings" + "testing" + "time" + + "github.com/agentfacts/mcp-proxy/internal/config" + "github.com/agentfacts/mcp-proxy/internal/session" +) + +func newTestSessionManager() *session.Manager { + return session.NewManager(session.ManagerConfig{ + SessionTTL: time.Hour, + CleanupInterval: time.Minute, + MaxSessions: 100, + }) +} + +func TestServerStartStop(t *testing.T) { + sessionMgr := newTestSessionManager() + agentCfg := config.AgentConfig{ + ID: "test-agent", + Name: "Test Agent", + Capabilities: []string{"tools"}, + } + + stdin := strings.NewReader("") + stdout := &bytes.Buffer{} + + server := NewServerWithIO(agentCfg, sessionMgr, stdin, stdout) + + ctx := context.Background() + + // Start should succeed + if err := server.Start(ctx); err != nil { + t.Fatalf("Start failed: %v", err) + } + + // Starting again should fail + if err := server.Start(ctx); err == nil { + t.Fatal("Expected error on second Start, got nil") + } + + // Stop should succeed + stopCtx, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + if err := server.Stop(stopCtx); err != nil { + t.Fatalf("Stop failed: %v", err) + } + + // Stopping again should be a no-op + if err := server.Stop(stopCtx); err != nil { + t.Fatalf("Second Stop failed: %v", err) + } +} + +func TestServerName(t *testing.T) { + sessionMgr := newTestSessionManager() + server := NewServer(config.AgentConfig{}, sessionMgr) + + if got := server.Name(); got != "stdio" { + t.Errorf("Name() = %q, want %q", got, "stdio") + } +} + +func TestServerMessageProcessing(t *testing.T) { + sessionMgr := newTestSessionManager() + agentCfg := config.AgentConfig{ + ID: "test-agent", + Name: "Test Agent", + } + + // Create a pipe to simulate stdin + stdinReader, stdinWriter := io.Pipe() + stdout := &bytes.Buffer{} + + server := NewServerWithIO(agentCfg, sessionMgr, stdinReader, stdout) + + // Set up a message handler that returns a response + server.SetMessageHandler(func(ctx context.Context, sess *session.Session, msg []byte) ([]byte, error) { + // Parse request and create response + var req map[string]interface{} + if err := json.Unmarshal(msg, &req); err != nil { + return nil, err + } + + response := map[string]interface{}{ + "jsonrpc": "2.0", + "id": req["id"], + "result": map[string]interface{}{"status": "ok"}, + } + return json.Marshal(response) + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + if err := server.Start(ctx); err != nil { + t.Fatalf("Start failed: %v", err) + } + + // Send a message + request := `{"jsonrpc":"2.0","method":"test","id":1}` + go func() { + stdinWriter.Write([]byte(request + "\n")) + // Give time for processing then close + time.Sleep(100 * time.Millisecond) + stdinWriter.Close() + }() + + // Wait for response + time.Sleep(200 * time.Millisecond) + + stopCtx, stopCancel := context.WithTimeout(ctx, time.Second) + defer stopCancel() + server.Stop(stopCtx) + + // Check response + output := stdout.String() + if output == "" { + t.Fatal("Expected output, got empty string") + } + + var response map[string]interface{} + if err := json.Unmarshal([]byte(strings.TrimSpace(output)), &response); err != nil { + t.Fatalf("Failed to parse response: %v, output was: %s", err, output) + } + + if response["jsonrpc"] != "2.0" { + t.Errorf("Expected jsonrpc 2.0, got %v", response["jsonrpc"]) + } + + if response["id"].(float64) != 1 { + t.Errorf("Expected id 1, got %v", response["id"]) + } +} + +func TestServerEchoWithoutHandler(t *testing.T) { + sessionMgr := newTestSessionManager() + agentCfg := config.AgentConfig{ + ID: "test-agent", + Name: "Test Agent", + } + + stdinReader, stdinWriter := io.Pipe() + stdout := &bytes.Buffer{} + + server := NewServerWithIO(agentCfg, sessionMgr, stdinReader, stdout) + // Don't set a message handler - should echo + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + if err := server.Start(ctx); err != nil { + t.Fatalf("Start failed: %v", err) + } + + request := `{"jsonrpc":"2.0","method":"test","id":1}` + go func() { + stdinWriter.Write([]byte(request + "\n")) + time.Sleep(100 * time.Millisecond) + stdinWriter.Close() + }() + + time.Sleep(200 * time.Millisecond) + + stopCtx, stopCancel := context.WithTimeout(ctx, time.Second) + defer stopCancel() + server.Stop(stopCtx) + + output := strings.TrimSpace(stdout.String()) + if output != request { + t.Errorf("Expected echo of request, got %s", output) + } +} + +func TestServerInvalidJSON(t *testing.T) { + sessionMgr := newTestSessionManager() + agentCfg := config.AgentConfig{ + ID: "test-agent", + Name: "Test Agent", + } + + stdinReader, stdinWriter := io.Pipe() + stdout := &bytes.Buffer{} + + server := NewServerWithIO(agentCfg, sessionMgr, stdinReader, stdout) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + if err := server.Start(ctx); err != nil { + t.Fatalf("Start failed: %v", err) + } + + // Send invalid JSON + go func() { + stdinWriter.Write([]byte("not valid json\n")) + time.Sleep(100 * time.Millisecond) + stdinWriter.Close() + }() + + time.Sleep(200 * time.Millisecond) + + stopCtx, stopCancel := context.WithTimeout(ctx, time.Second) + defer stopCancel() + server.Stop(stopCtx) + + output := stdout.String() + if !strings.Contains(output, "error") { + t.Errorf("Expected error response for invalid JSON, got: %s", output) + } + + var response map[string]interface{} + if err := json.Unmarshal([]byte(strings.TrimSpace(output)), &response); err != nil { + t.Fatalf("Failed to parse error response: %v", err) + } + + errObj, ok := response["error"].(map[string]interface{}) + if !ok { + t.Fatalf("Expected error object in response") + } + + if errObj["code"].(float64) != -32700 { + t.Errorf("Expected error code -32700, got %v", errObj["code"]) + } +} + +func TestServerEOFShutdown(t *testing.T) { + sessionMgr := newTestSessionManager() + agentCfg := config.AgentConfig{ + ID: "test-agent", + Name: "Test Agent", + } + + // Use a simple reader that immediately returns EOF + stdin := strings.NewReader("") + stdout := &bytes.Buffer{} + + server := NewServerWithIO(agentCfg, sessionMgr, stdin, stdout) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + if err := server.Start(ctx); err != nil { + t.Fatalf("Start failed: %v", err) + } + + // Server should detect EOF and be ready to stop + time.Sleep(100 * time.Millisecond) + + stopCtx, stopCancel := context.WithTimeout(ctx, time.Second) + defer stopCancel() + if err := server.Stop(stopCtx); err != nil { + t.Fatalf("Stop failed: %v", err) + } +} + +func TestReaderBasic(t *testing.T) { + input := `{"jsonrpc":"2.0","method":"test","id":1} +{"jsonrpc":"2.0","method":"test2","id":2} +` + reader := NewReader(strings.NewReader(input)) + + msg1, err := reader.ReadMessage() + if err != nil { + t.Fatalf("ReadMessage() error: %v", err) + } + + var req1 map[string]interface{} + if err := json.Unmarshal(msg1, &req1); err != nil { + t.Fatalf("Failed to parse message 1: %v", err) + } + if req1["id"].(float64) != 1 { + t.Errorf("Expected id 1, got %v", req1["id"]) + } + + msg2, err := reader.ReadMessage() + if err != nil { + t.Fatalf("ReadMessage() error: %v", err) + } + + var req2 map[string]interface{} + if err := json.Unmarshal(msg2, &req2); err != nil { + t.Fatalf("Failed to parse message 2: %v", err) + } + if req2["id"].(float64) != 2 { + t.Errorf("Expected id 2, got %v", req2["id"]) + } + + // Third read should return EOF + _, err = reader.ReadMessage() + if err != io.EOF { + t.Errorf("Expected io.EOF, got %v", err) + } +} + +func TestReaderSkipsEmptyLines(t *testing.T) { + input := ` +{"jsonrpc":"2.0","id":1} + +{"jsonrpc":"2.0","id":2} +` + reader := NewReader(strings.NewReader(input)) + + msg1, err := reader.ReadMessage() + if err != nil { + t.Fatalf("ReadMessage() error: %v", err) + } + var req1 map[string]interface{} + json.Unmarshal(msg1, &req1) + if req1["id"].(float64) != 1 { + t.Errorf("Expected id 1, got %v", req1["id"]) + } + + msg2, err := reader.ReadMessage() + if err != nil { + t.Fatalf("ReadMessage() error: %v", err) + } + var req2 map[string]interface{} + json.Unmarshal(msg2, &req2) + if req2["id"].(float64) != 2 { + t.Errorf("Expected id 2, got %v", req2["id"]) + } +} + +func TestReaderInvalidJSON(t *testing.T) { + input := "not valid json\n" + reader := NewReader(strings.NewReader(input)) + + _, err := reader.ReadMessage() + if err == nil { + t.Fatal("Expected error for invalid JSON") + } + if !strings.Contains(err.Error(), "invalid JSON") { + t.Errorf("Expected 'invalid JSON' error, got: %v", err) + } +} + +func TestWriterBasic(t *testing.T) { + buf := &bytes.Buffer{} + writer := NewWriter(buf) + + data := []byte(`{"jsonrpc":"2.0","result":"ok","id":1}`) + if err := writer.Write(data); err != nil { + t.Fatalf("Write() error: %v", err) + } + + expected := string(data) + "\n" + if got := buf.String(); got != expected { + t.Errorf("Write() output = %q, want %q", got, expected) + } +} + +func TestWriterMultiple(t *testing.T) { + buf := &bytes.Buffer{} + writer := NewWriter(buf) + + msg1 := []byte(`{"id":1}`) + msg2 := []byte(`{"id":2}`) + + writer.Write(msg1) + writer.Write(msg2) + + expected := string(msg1) + "\n" + string(msg2) + "\n" + if got := buf.String(); got != expected { + t.Errorf("Output = %q, want %q", got, expected) + } +} + +func TestServerSessionInfo(t *testing.T) { + sessionMgr := newTestSessionManager() + agentCfg := config.AgentConfig{ + ID: "custom-agent-id", + Name: "Custom Agent", + Capabilities: []string{"tools", "resources"}, + } + + stdin := strings.NewReader("") + stdout := &bytes.Buffer{} + + server := NewServerWithIO(agentCfg, sessionMgr, stdin, stdout) + + ctx := context.Background() + if err := server.Start(ctx); err != nil { + t.Fatalf("Start failed: %v", err) + } + + // Verify session was created with correct info + if server.session == nil { + t.Fatal("Session not created") + } + + if server.session.AgentID != "custom-agent-id" { + t.Errorf("AgentID = %q, want %q", server.session.AgentID, "custom-agent-id") + } + + if server.session.AgentName != "Custom Agent" { + t.Errorf("AgentName = %q, want %q", server.session.AgentName, "Custom Agent") + } + + if len(server.session.Capabilities) != 2 { + t.Errorf("Capabilities length = %d, want 2", len(server.session.Capabilities)) + } + + if server.session.SourceIP != "stdio" { + t.Errorf("SourceIP = %q, want %q", server.session.SourceIP, "stdio") + } + + stopCtx, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + server.Stop(stopCtx) +} diff --git a/internal/transport/stdio/writer.go b/internal/transport/stdio/writer.go new file mode 100644 index 0000000..08d1541 --- /dev/null +++ b/internal/transport/stdio/writer.go @@ -0,0 +1,43 @@ +package stdio + +import ( + "io" + "sync" +) + +// Writer handles thread-safe writes to stdout with newline framing. +type Writer struct { + out io.Writer + mu sync.Mutex +} + +// NewWriter creates a new Writer for the given output stream. +func NewWriter(out io.Writer) *Writer { + return &Writer{ + out: out, + } +} + +// Write writes a JSON message followed by a newline. +// It is safe for concurrent use. +func (w *Writer) Write(data []byte) error { + w.mu.Lock() + defer w.mu.Unlock() + + // Write the message + if _, err := w.out.Write(data); err != nil { + return err + } + + // Write newline delimiter + if _, err := w.out.Write([]byte{'\n'}); err != nil { + return err + } + + // Flush if the writer supports it + if f, ok := w.out.(interface{ Flush() error }); ok { + return f.Flush() + } + + return nil +} diff --git a/internal/transport/transport.go b/internal/transport/transport.go index 2766540..183821e 100644 --- a/internal/transport/transport.go +++ b/internal/transport/transport.go @@ -2,8 +2,14 @@ package transport import ( "context" + + "github.com/agentfacts/mcp-proxy/internal/session" ) +// MessageHandler is called when a message is received from a client. +// It receives the session and raw message, and returns a response. +type MessageHandler func(ctx context.Context, sess *session.Session, message []byte) ([]byte, error) + // Transport defines the interface for MCP transport implementations. // Supported transports: SSE, stdio, HTTP type Transport interface { @@ -15,11 +21,10 @@ type Transport interface { // Name returns the transport type name (e.g., "sse", "stdio", "http") Name() string -} -// MessageHandler is called when a message is received from a client. -// It should process the message and return a response. -type MessageHandler func(ctx context.Context, sessionID string, message []byte) ([]byte, error) + // SetMessageHandler sets the callback for processing incoming messages. + SetMessageHandler(handler MessageHandler) +} // ConnectionHandler is called when a new connection is established or closed. type ConnectionHandler interface {