summaryrefslogtreecommitdiff
path: root/cmd
diff options
context:
space:
mode:
Diffstat (limited to 'cmd')
-rw-r--r--cmd/hexai-mcp-server/main.go33
-rw-r--r--cmd/hexai-mcp-server/main_test.go54
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)
}