diff options
Diffstat (limited to 'cmd')
| -rw-r--r-- | cmd/hexai-mcp-server/main.go | 33 | ||||
| -rw-r--r-- | cmd/hexai-mcp-server/main_test.go | 54 |
2 files changed, 44 insertions, 43 deletions
diff --git a/cmd/hexai-mcp-server/main.go b/cmd/hexai-mcp-server/main.go index 03150bb..741e310 100644 --- a/cmd/hexai-mcp-server/main.go +++ b/cmd/hexai-mcp-server/main.go @@ -25,13 +25,6 @@ func buildOverrides(opts mcpOptions) hexaimcp.MCPOverrides { } } -// Seams for testing: override in tests to avoid launching real MCP server. -// Signatures match hexaimcp.Run and hexaimcp.RunBackfill respectively. -var ( - runMCP = hexaimcp.Run - runBackfill = hexaimcp.RunBackfill -) - // deprecationWarning is the notice runMain emits on every startup so users // see this binary is experimental. Kept as a constant (not printf'd) so // tests can assert on its contents directly. @@ -65,6 +58,18 @@ type mcpOptions struct { showVersion bool } +type mcpDeps struct { + runMCP func(context.Context, string, string, hexaimcp.MCPOverrides, io.Reader, io.Writer, io.Writer) error + runBackfill func(context.Context, string, string, hexaimcp.MCPOverrides) error +} + +func defaultMCPDeps() mcpDeps { + return mcpDeps{ + runMCP: hexaimcp.Run, + runBackfill: hexaimcp.RunBackfill, + } +} + func main() { os.Exit(runMain(os.Args[1:], os.Stdin, os.Stdout, os.Stderr)) } // runMain prints the deprecation warning, parses flags, and delegates to @@ -73,6 +78,10 @@ func main() { os.Exit(runMain(os.Args[1:], os.Stdin, os.Stdout, os.Stderr)) } // Pulling this out of main keeps it testable without touching package-level // flag state. func runMain(args []string, stdin io.Reader, stdout, stderr io.Writer) int { + return runMainWithDeps(args, stdin, stdout, stderr, defaultMCPDeps()) +} + +func runMainWithDeps(args []string, stdin io.Reader, stdout, stderr io.Writer, deps mcpDeps) int { fmt.Fprint(stderr, deprecationWarning) defaultLog, err := defaultLogPath() @@ -107,7 +116,7 @@ func runMain(args []string, stdin io.Reader, stdout, stderr io.Writer) int { ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) defer stop() - if err := run(ctx, opts, stdin, stdout, stderr); err != nil { + if err := runWithDeps(ctx, opts, stdin, stdout, stderr, deps); err != nil { fmt.Fprintf(stderr, "error: %v\n", err) return 1 } @@ -118,6 +127,10 @@ func runMain(args []string, stdin io.Reader, stdout, stderr io.Writer) int { // CLI flag values are passed via MCPOverrides instead of environment variables. // ctx is threaded into the server/backfill so they stop on shutdown signals. func run(ctx context.Context, opts mcpOptions, stdin io.Reader, stdout, stderr io.Writer) error { + return runWithDeps(ctx, opts, stdin, stdout, stderr, defaultMCPDeps()) +} + +func runWithDeps(ctx context.Context, opts mcpOptions, stdin io.Reader, stdout, stderr io.Writer, deps mcpDeps) error { if opts.showVersion { fmt.Fprintln(stdout, internal.Version) return nil @@ -127,10 +140,10 @@ func run(ctx context.Context, opts mcpOptions, stdin io.Reader, stdout, stderr i // Handle backfill operation if opts.syncAll { - return runBackfill(ctx, opts.logPath, opts.configPath, overrides) + return deps.runBackfill(ctx, opts.logPath, opts.configPath, overrides) } - return runMCP(ctx, opts.logPath, opts.configPath, overrides, stdin, stdout, stderr) + return deps.runMCP(ctx, opts.logPath, opts.configPath, overrides, stdin, stdout, stderr) } // defaultLogPath returns the default MCP log file path in the state directory. diff --git a/cmd/hexai-mcp-server/main_test.go b/cmd/hexai-mcp-server/main_test.go index b2a3895..ded0be7 100644 --- a/cmd/hexai-mcp-server/main_test.go +++ b/cmd/hexai-mcp-server/main_test.go @@ -66,12 +66,10 @@ func TestBuildOverrides(t *testing.T) { } func TestRun_SyncAll(t *testing.T) { - old := runBackfill - t.Cleanup(func() { runBackfill = old }) - var gotLog, gotConfig string var gotOverrides hexaimcp.MCPOverrides - runBackfill = func(_ context.Context, logPath, configPath string, overrides hexaimcp.MCPOverrides) error { + deps := defaultMCPDeps() + deps.runBackfill = func(_ context.Context, logPath, configPath string, overrides hexaimcp.MCPOverrides) error { gotLog = logPath gotConfig = configPath gotOverrides = overrides @@ -86,7 +84,7 @@ func TestRun_SyncAll(t *testing.T) { slashCommandSync: true, slashCommandDir: "/tmp/cmds", } - if err := run(context.Background(), opts, nil, nil, nil); err != nil { + if err := runWithDeps(context.Background(), opts, nil, nil, nil, deps); err != nil { t.Fatalf("run syncAll: %v", err) } if gotLog != "/tmp/test.log" { @@ -107,30 +105,26 @@ func TestRun_SyncAll(t *testing.T) { } func TestRun_SyncAllError(t *testing.T) { - old := runBackfill - t.Cleanup(func() { runBackfill = old }) - wantErr := errors.New("backfill failed") - runBackfill = func(_ context.Context, _, _ string, _ hexaimcp.MCPOverrides) error { return wantErr } + deps := defaultMCPDeps() + deps.runBackfill = func(_ context.Context, _, _ string, _ hexaimcp.MCPOverrides) error { return wantErr } opts := mcpOptions{syncAll: true} - if err := run(context.Background(), opts, nil, nil, nil); !errors.Is(err, wantErr) { + if err := runWithDeps(context.Background(), opts, nil, nil, nil, deps); !errors.Is(err, wantErr) { t.Fatalf("expected backfill error, got: %v", err) } } func TestRun_MCPServer(t *testing.T) { - old := runMCP - t.Cleanup(func() { runMCP = old }) - called := false - runMCP = func(_ context.Context, logPath, configPath string, overrides hexaimcp.MCPOverrides, stdin io.Reader, stdout, stderr io.Writer) error { + deps := defaultMCPDeps() + deps.runMCP = func(_ context.Context, logPath, configPath string, overrides hexaimcp.MCPOverrides, stdin io.Reader, stdout, stderr io.Writer) error { called = true return nil } opts := mcpOptions{logPath: "/tmp/mcp.log"} - if err := run(context.Background(), opts, nil, nil, nil); err != nil { + if err := runWithDeps(context.Background(), opts, nil, nil, nil, deps); err != nil { t.Fatalf("run MCP: %v", err) } if !called { @@ -139,15 +133,13 @@ func TestRun_MCPServer(t *testing.T) { } func TestRun_MCPServerError(t *testing.T) { - old := runMCP - t.Cleanup(func() { runMCP = old }) - wantErr := errors.New("server failed") - runMCP = func(_ context.Context, _, _ string, _ hexaimcp.MCPOverrides, _ io.Reader, _, _ io.Writer) error { + deps := defaultMCPDeps() + deps.runMCP = func(_ context.Context, _, _ string, _ hexaimcp.MCPOverrides, _ io.Reader, _, _ io.Writer) error { return wantErr } - if err := run(context.Background(), mcpOptions{}, nil, nil, nil); !errors.Is(err, wantErr) { + if err := runWithDeps(context.Background(), mcpOptions{}, nil, nil, nil, deps); !errors.Is(err, wantErr) { t.Fatalf("expected server error, got: %v", err) } } @@ -172,17 +164,15 @@ func TestRunMain_VersionFlag(t *testing.T) { // runMain --sync-all path: forwards parsed options to runBackfill and // returns 0 on success. func TestRunMain_SyncAllSuccess(t *testing.T) { - old := runBackfill - t.Cleanup(func() { runBackfill = old }) - var gotLog string - runBackfill = func(_ context.Context, logPath string, _ string, _ hexaimcp.MCPOverrides) error { + deps := defaultMCPDeps() + deps.runBackfill = func(_ context.Context, logPath string, _ string, _ hexaimcp.MCPOverrides) error { gotLog = logPath return nil } var stdout, stderr bytes.Buffer - code := runMain([]string{"-sync-all", "-log", "/tmp/sync.log"}, nil, &stdout, &stderr) + code := runMainWithDeps([]string{"-sync-all", "-log", "/tmp/sync.log"}, nil, &stdout, &stderr, deps) if code != 0 { t.Fatalf("runMain code = %d, want 0; stderr=%q", code, stderr.String()) } @@ -194,14 +184,13 @@ func TestRunMain_SyncAllSuccess(t *testing.T) { // runMain run-error path: when the underlying server fails, runMain must // return 1 (the production exit code) and write the error to stderr. func TestRunMain_ServerErrorReturnsOne(t *testing.T) { - old := runMCP - t.Cleanup(func() { runMCP = old }) - runMCP = func(context.Context, string, string, hexaimcp.MCPOverrides, io.Reader, io.Writer, io.Writer) error { + deps := defaultMCPDeps() + deps.runMCP = func(context.Context, string, string, hexaimcp.MCPOverrides, io.Reader, io.Writer, io.Writer) error { return errors.New("mcp boom") } var stdout, stderr bytes.Buffer - code := runMain(nil, nil, &stdout, &stderr) + code := runMainWithDeps(nil, nil, &stdout, &stderr, deps) if code != 1 { t.Fatalf("runMain code = %d, want 1", code) } @@ -212,15 +201,14 @@ func TestRunMain_ServerErrorReturnsOne(t *testing.T) { // Bad flag must yield exit 2 without ever invoking the server stub. func TestRunMain_BadFlagReturnsTwo(t *testing.T) { - old := runMCP - t.Cleanup(func() { runMCP = old }) called := false - runMCP = func(context.Context, string, string, hexaimcp.MCPOverrides, io.Reader, io.Writer, io.Writer) error { + deps := defaultMCPDeps() + deps.runMCP = func(context.Context, string, string, hexaimcp.MCPOverrides, io.Reader, io.Writer, io.Writer) error { called = true return nil } var stdout, stderr bytes.Buffer - code := runMain([]string{"--bogus"}, nil, &stdout, &stderr) + code := runMainWithDeps([]string{"--bogus"}, nil, &stdout, &stderr, deps) if code != 2 { t.Fatalf("runMain code = %d, want 2", code) } |
