diff options
| author | Paul Buetow <paul@buetow.org> | 2026-06-19 08:57:03 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-06-19 08:57:03 +0300 |
| commit | 3b0816f11504ac9ef8f55b91b99371fbb55bb040 (patch) | |
| tree | 38a6e4239fc2ac0f23977a69065468880ec70b07 | |
| parent | 426068b0ea1fff4e2bfc708d8a3914379f084ad8 (diff) | |
fix mcp protocol test cleanup (7k0)
| -rw-r--r-- | internal/hexaimcp/run.go | 20 | ||||
| -rw-r--r-- | internal/hexaimcp/run_test.go | 154 |
2 files changed, 144 insertions, 30 deletions
diff --git a/internal/hexaimcp/run.go b/internal/hexaimcp/run.go index 7043ed2..25897c6 100644 --- a/internal/hexaimcp/run.go +++ b/internal/hexaimcp/run.go @@ -67,6 +67,20 @@ func RunWithFactory( stderr io.Writer, factory ServerFactory, ) error { + loadOpts := appconfig.LoadOptions{ConfigPath: configPath} + return runWithFactoryLoadOptions(ctx, logPath, loadOpts, overrides, stdin, stdout, stderr, factory) +} + +func runWithFactoryLoadOptions( + ctx context.Context, + logPath string, + loadOpts appconfig.LoadOptions, + overrides MCPOverrides, + stdin io.Reader, + stdout io.Writer, + stderr io.Writer, + factory ServerFactory, +) error { // Setup logger logger, err := setupLogger(logPath) if err != nil { @@ -82,7 +96,7 @@ func RunWithFactory( logger.Printf("WARNING: hexai-mcp-server is DEPRECATED and experimental - not actively maintained") // Load configuration and apply CLI overrides - cfg := loadConfig(ctx, logger, configPath) + cfg := loadConfigWithOptions(ctx, logger, loadOpts) applyOverrides(&cfg, overrides) return runServer(ctx, cfg, logger, stdin, stdout, factory) @@ -148,6 +162,10 @@ func loadConfig(ctx context.Context, logger *log.Logger, configPath string) appc ConfigPath: configPath, IgnoreEnv: false, } + return loadConfigWithOptions(ctx, logger, opts) +} + +func loadConfigWithOptions(ctx context.Context, logger *log.Logger, opts appconfig.LoadOptions) appconfig.App { return appconfig.LoadWithOptions(ctx, logger, opts) } diff --git a/internal/hexaimcp/run_test.go b/internal/hexaimcp/run_test.go index d567f2e..f4e130b 100644 --- a/internal/hexaimcp/run_test.go +++ b/internal/hexaimcp/run_test.go @@ -2,6 +2,7 @@ package hexaimcp import ( + "bufio" "bytes" "context" "encoding/json" @@ -12,6 +13,7 @@ import ( "path/filepath" "strings" "testing" + "time" "codeberg.org/snonux/hexai/internal/appconfig" "codeberg.org/snonux/hexai/internal/mcp" @@ -23,6 +25,13 @@ type mockServerRunner struct { runFunc func() error } +type fullProtocolServer struct { + stdin *io.PipeWriter + stdout io.Reader + done <-chan error + cancel context.CancelFunc +} + func (m *mockServerRunner) Run(context.Context) error { if m.runFunc != nil { return m.runFunc() @@ -32,20 +41,79 @@ func (m *mockServerRunner) Run(context.Context) error { // TestFullProtocolFlow tests the complete MCP protocol interaction func TestFullProtocolFlow(t *testing.T) { + tmpDir, promptsDir := setupFullProtocolConfig(t) + server := startFullProtocolServer(t, tmpDir, promptsDir) + defer server.cancel() + + writeJSONRPCLine(t, server.stdin, initializeRequest()) + assertInitializeResponse(t, readJSONRPCLine(t, server.stdout)) + assertServerStillRunning(t, server.done) + + if err := server.stdin.Close(); err != nil { + t.Fatalf("close stdin writer: %v", err) + } + if err := waitForServer(t, server.done); err != nil { + t.Fatalf("server returned error: %v", err) + } + + if _, err := os.Stat(filepath.Join(promptsDir, "backups")); err != nil { + t.Fatalf("prompts store was not initialized in test dir: %v", err) + } + if err := os.RemoveAll(tmpDir); err != nil { + t.Fatalf("remove temp dir after server exit: %v", err) + } +} + +func setupFullProtocolConfig(t *testing.T) (string, string) { + t.Helper() tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, "config.toml") + if err := os.WriteFile(configPath, []byte(""), 0o644); err != nil { + t.Fatalf("write config: %v", err) + } + + t.Setenv("HEXAI_MCP_SLASHCOMMAND_SYNC", "true") + t.Setenv("HEXAI_MCP_SLASHCOMMAND_DIR", "/dev/null/impossible") - // Create test server factory + return tmpDir, filepath.Join(tmpDir, "prompts") +} + +func startFullProtocolServer(t *testing.T, tmpDir, promptsDir string) fullProtocolServer { + t.Helper() + stdinReader, stdinWriter := io.Pipe() + stdoutReader, stdoutWriter := io.Pipe() + done := make(chan error, 1) + ctx, cancel := context.WithCancel(context.Background()) + + loadOpts := appconfig.LoadOptions{ + ConfigPath: filepath.Join(tmpDir, "config.toml"), + IgnoreEnv: true, + ProjectRoot: tmpDir, + } + overrides := MCPOverrides{PromptsDir: promptsDir} serverFactory := func(r io.Reader, w io.Writer, logger *log.Logger, store promptstore.PromptStore, syncer mcp.SlashCommandSyncer) ServerRunner { return mcp.NewServer(r, w, logger, store, syncer) } - // Setup I/O pipes - inBuf := &bytes.Buffer{} - outBuf := &bytes.Buffer{} - errBuf := &bytes.Buffer{} + go func() { + defer func() { + _ = stdinReader.Close() + _ = stdoutWriter.Close() + }() + logPath := filepath.Join(tmpDir, "mcp.log") + done <- runWithFactoryLoadOptions(ctx, logPath, loadOpts, overrides, stdinReader, stdoutWriter, &bytes.Buffer{}, serverFactory) + }() - // Send initialize request - initReq := map[string]any{ + return fullProtocolServer{ + stdin: stdinWriter, + stdout: stdoutReader, + done: done, + cancel: cancel, + } +} + +func initializeRequest() map[string]any { + return map[string]any{ "jsonrpc": "2.0", "id": 1, "method": "initialize", @@ -58,38 +126,66 @@ func TestFullProtocolFlow(t *testing.T) { }, }, } - - writeJSONRPC(t, inBuf, initReq) - - // Run server in background (it will read from inBuf and write to outBuf) - go func() { - // Pass prompts dir via overrides instead of environment variable - overrides := MCPOverrides{PromptsDir: tmpDir} - - // Note: This will hang waiting for more input, which is expected - _ = RunWithFactory(context.Background(), "", "", overrides, inBuf, outBuf, errBuf, serverFactory) - }() - - // Give server time to process - // Note: In a real test, you'd use proper synchronization - - // For now, just verify the server starts and creates the prompts directory - // A full integration test would require more sophisticated I/O handling } -func writeJSONRPC(t *testing.T, w io.Writer, req map[string]any) { +func writeJSONRPCLine(t *testing.T, w io.Writer, req map[string]any) { t.Helper() data, err := json.Marshal(req) if err != nil { t.Fatalf("marshal request: %v", err) } - header := fmt.Sprintf("Content-Length: %d\r\n\r\n", len(data)) - if _, err := io.WriteString(w, header); err != nil { - t.Fatalf("write header: %v", err) - } if _, err := w.Write(data); err != nil { t.Fatalf("write body: %v", err) } + if _, err := io.WriteString(w, "\n"); err != nil { + t.Fatalf("write newline: %v", err) + } +} + +func readJSONRPCLine(t *testing.T, r io.Reader) mcp.Response { + t.Helper() + line, err := bufio.NewReader(r).ReadBytes('\n') + if err != nil { + t.Fatalf("read response: %v", err) + } + var resp mcp.Response + if err := json.Unmarshal(line, &resp); err != nil { + t.Fatalf("unmarshal response %q: %v", line, err) + } + return resp +} + +func assertInitializeResponse(t *testing.T, resp mcp.Response) { + t.Helper() + if resp.JSONRPC != "2.0" { + t.Fatalf("response JSONRPC = %q, want 2.0", resp.JSONRPC) + } + if resp.ID != float64(1) { + t.Fatalf("response ID = %#v, want 1", resp.ID) + } + if resp.Error != nil { + t.Fatalf("initialize response error = %+v", resp.Error) + } +} + +func assertServerStillRunning(t *testing.T, done <-chan error) { + t.Helper() + select { + case err := <-done: + t.Fatalf("server returned before stdin was closed: %v", err) + default: + } +} + +func waitForServer(t *testing.T, done <-chan error) error { + t.Helper() + select { + case err := <-done: + return err + case <-time.After(2 * time.Second): + t.Fatal("server did not exit after stdin closed") + } + return nil } func TestGetPromptsDir(t *testing.T) { |
