diff options
Diffstat (limited to 'cmd')
| -rw-r--r-- | cmd/hexai-mcp-server/main.go | 28 | ||||
| -rw-r--r-- | cmd/hexai-mcp-server/main_test.go | 71 |
2 files changed, 51 insertions, 48 deletions
diff --git a/cmd/hexai-mcp-server/main.go b/cmd/hexai-mcp-server/main.go index 32817f3..d2495e6 100644 --- a/cmd/hexai-mcp-server/main.go +++ b/cmd/hexai-mcp-server/main.go @@ -12,7 +12,17 @@ import ( "codeberg.org/snonux/hexai/internal/hexaimcp" ) +// buildOverrides constructs MCPOverrides from parsed CLI options. +func buildOverrides(opts mcpOptions) hexaimcp.MCPOverrides { + return hexaimcp.MCPOverrides{ + PromptsDir: opts.promptsDir, + SlashCommandSync: opts.slashCommandSync, + SlashCommandDir: opts.slashCommandDir, + } +} + // 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 @@ -86,29 +96,21 @@ func main() { } // run executes the MCP server logic with the given options and I/O streams. +// CLI flag values are passed via MCPOverrides instead of environment variables. func run(opts mcpOptions, stdin io.Reader, stdout, stderr io.Writer) error { - // Set environment variables for RunWithFactory based on flag values - if opts.promptsDir != "" { - os.Setenv("HEXAI_MCP_PROMPTS_DIR", opts.promptsDir) - } - if opts.slashCommandSync { - os.Setenv("HEXAI_MCP_SLASHCOMMAND_SYNC", "true") - } - if opts.slashCommandDir != "" { - os.Setenv("HEXAI_MCP_SLASHCOMMAND_DIR", opts.slashCommandDir) - } - if opts.showVersion { fmt.Fprintln(stdout, internal.Version) return nil } + overrides := buildOverrides(opts) + // Handle backfill operation if opts.syncAll { - return runBackfill(opts.logPath, opts.configPath) + return runBackfill(opts.logPath, opts.configPath, overrides) } - return runMCP(opts.logPath, opts.configPath, stdin, stdout, stderr) + return runMCP(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 76d0fca..3d85fbb 100644 --- a/cmd/hexai-mcp-server/main_test.go +++ b/cmd/hexai-mcp-server/main_test.go @@ -9,6 +9,7 @@ import ( "testing" "codeberg.org/snonux/hexai/internal" + "codeberg.org/snonux/hexai/internal/hexaimcp" ) func TestPrintDeprecationWarning(t *testing.T) { @@ -65,39 +66,21 @@ func TestRun_ShowVersion(t *testing.T) { } } -func TestRun_SetsPromptsDir(t *testing.T) { - t.Setenv("HEXAI_MCP_PROMPTS_DIR", "") - opts := mcpOptions{showVersion: true, promptsDir: "/tmp/test-prompts"} - var stdout bytes.Buffer - if err := run(opts, nil, &stdout, nil); err != nil { - t.Fatalf("run: %v", err) - } - if got := os.Getenv("HEXAI_MCP_PROMPTS_DIR"); got != "/tmp/test-prompts" { - t.Fatalf("expected HEXAI_MCP_PROMPTS_DIR=/tmp/test-prompts, got %q", got) - } -} - -func TestRun_SetsSlashCommandSync(t *testing.T) { - t.Setenv("HEXAI_MCP_SLASHCOMMAND_SYNC", "") - opts := mcpOptions{showVersion: true, slashCommandSync: true} - var stdout bytes.Buffer - if err := run(opts, nil, &stdout, nil); err != nil { - t.Fatalf("run: %v", err) +func TestBuildOverrides(t *testing.T) { + opts := mcpOptions{ + promptsDir: "/tmp/test-prompts", + slashCommandSync: true, + slashCommandDir: "/tmp/test-cmds", } - if got := os.Getenv("HEXAI_MCP_SLASHCOMMAND_SYNC"); got != "true" { - t.Fatalf("expected HEXAI_MCP_SLASHCOMMAND_SYNC=true, got %q", got) + overrides := buildOverrides(opts) + if overrides.PromptsDir != "/tmp/test-prompts" { + t.Fatalf("expected PromptsDir=/tmp/test-prompts, got %q", overrides.PromptsDir) } -} - -func TestRun_SetsSlashCommandDir(t *testing.T) { - t.Setenv("HEXAI_MCP_SLASHCOMMAND_DIR", "") - opts := mcpOptions{showVersion: true, slashCommandDir: "/tmp/test-cmds"} - var stdout bytes.Buffer - if err := run(opts, nil, &stdout, nil); err != nil { - t.Fatalf("run: %v", err) + if !overrides.SlashCommandSync { + t.Fatal("expected SlashCommandSync=true") } - if got := os.Getenv("HEXAI_MCP_SLASHCOMMAND_DIR"); got != "/tmp/test-cmds" { - t.Fatalf("expected HEXAI_MCP_SLASHCOMMAND_DIR=/tmp/test-cmds, got %q", got) + if overrides.SlashCommandDir != "/tmp/test-cmds" { + t.Fatalf("expected SlashCommandDir=/tmp/test-cmds, got %q", overrides.SlashCommandDir) } } @@ -106,13 +89,22 @@ func TestRun_SyncAll(t *testing.T) { t.Cleanup(func() { runBackfill = old }) var gotLog, gotConfig string - runBackfill = func(logPath, configPath string) error { + var gotOverrides hexaimcp.MCPOverrides + runBackfill = func(logPath, configPath string, overrides hexaimcp.MCPOverrides) error { gotLog = logPath gotConfig = configPath + gotOverrides = overrides return nil } - opts := mcpOptions{syncAll: true, logPath: "/tmp/test.log", configPath: "/tmp/cfg.toml"} + opts := mcpOptions{ + syncAll: true, + logPath: "/tmp/test.log", + configPath: "/tmp/cfg.toml", + promptsDir: "/tmp/prompts", + slashCommandSync: true, + slashCommandDir: "/tmp/cmds", + } if err := run(opts, nil, nil, nil); err != nil { t.Fatalf("run syncAll: %v", err) } @@ -122,6 +114,15 @@ func TestRun_SyncAll(t *testing.T) { if gotConfig != "/tmp/cfg.toml" { t.Fatalf("expected configPath=/tmp/cfg.toml, got %q", gotConfig) } + if gotOverrides.PromptsDir != "/tmp/prompts" { + t.Fatalf("expected overrides.PromptsDir=/tmp/prompts, got %q", gotOverrides.PromptsDir) + } + if !gotOverrides.SlashCommandSync { + t.Fatal("expected overrides.SlashCommandSync=true") + } + if gotOverrides.SlashCommandDir != "/tmp/cmds" { + t.Fatalf("expected overrides.SlashCommandDir=/tmp/cmds, got %q", gotOverrides.SlashCommandDir) + } } func TestRun_SyncAllError(t *testing.T) { @@ -129,7 +130,7 @@ func TestRun_SyncAllError(t *testing.T) { t.Cleanup(func() { runBackfill = old }) wantErr := errors.New("backfill failed") - runBackfill = func(_, _ string) error { return wantErr } + runBackfill = func(_, _ string, _ hexaimcp.MCPOverrides) error { return wantErr } opts := mcpOptions{syncAll: true} if err := run(opts, nil, nil, nil); !errors.Is(err, wantErr) { @@ -142,7 +143,7 @@ func TestRun_MCPServer(t *testing.T) { t.Cleanup(func() { runMCP = old }) called := false - runMCP = func(logPath, configPath string, stdin io.Reader, stdout, stderr io.Writer) error { + runMCP = func(logPath, configPath string, overrides hexaimcp.MCPOverrides, stdin io.Reader, stdout, stderr io.Writer) error { called = true return nil } @@ -161,7 +162,7 @@ func TestRun_MCPServerError(t *testing.T) { t.Cleanup(func() { runMCP = old }) wantErr := errors.New("server failed") - runMCP = func(_, _ string, _ io.Reader, _, _ io.Writer) error { return wantErr } + runMCP = func(_, _ string, _ hexaimcp.MCPOverrides, _ io.Reader, _, _ io.Writer) error { return wantErr } if err := run(mcpOptions{}, nil, nil, nil); !errors.Is(err, wantErr) { t.Fatalf("expected server error, got: %v", err) |
