diff --git a/cmd/entire/cli/agent/external/discovery.go b/cmd/entire/cli/agent/external/discovery.go index 2fb5fcb1f7..7a7ca7ff18 100644 --- a/cmd/entire/cli/agent/external/discovery.go +++ b/cmd/entire/cli/agent/external/discovery.go @@ -2,8 +2,11 @@ package external import ( "context" + "errors" + "fmt" "log/slog" "os" + "os/exec" "path/filepath" "runtime" "strings" @@ -23,6 +26,11 @@ const ( // discoveryTimeout caps the total time spent scanning $PATH for external agents. const discoveryTimeout = 10 * time.Second +var ( + statExternalAgent = os.Stat //nolint:gochecknoglobals // narrow test seam for stat failures + lookPathExternalAgent = exec.LookPath //nolint:gochecknoglobals // narrow test seam for lookup failures +) + // DiscoverAndRegister scans $PATH for executables matching "entire-agent-", // calls their "info" subcommand, and registers them in the agent registry. // Binaries whose name conflicts with an already-registered agent are skipped. @@ -43,6 +51,51 @@ func DiscoverAndRegisterAlways(ctx context.Context) { discoverAndRegister(ctx) } +// DiscoverAndRegisterNamedAlways discovers and registers only the external +// agent binary matching name. It bypasses the external_agents setting for +// explicit, one-invocation selections without executing unrelated plugins. +func DiscoverAndRegisterNamedAlways(ctx context.Context, name types.AgentName) error { + return discoverAndRegisterNamed(ctx, name, discoveryTimeout) +} + +func discoverAndRegisterNamed(ctx context.Context, name types.AgentName, timeout time.Duration) error { + if name == "" { + return nil + } + if strings.ContainsAny(string(name), `/\`) { + return fmt.Errorf("invalid external agent name %q: contains path separators", name) + } + if _, err := agent.Get(name); err == nil { + return nil + } + + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + if err := ctx.Err(); err != nil { + return fmt.Errorf("discovering external agent %q: %w", name, err) + } + + binName := binaryPrefix + string(name) + binPath, err := lookPathExternalAgent(binName) + if ctxErr := ctx.Err(); ctxErr != nil { + return fmt.Errorf("looking up external agent %q binary %q: %w", name, binName, ctxErr) + } + if err != nil { + if errors.Is(err, exec.ErrNotFound) { + return nil + } + return fmt.Errorf("looking up external agent %q binary %q: %w", name, binName, err) + } + registered, err := registerExternalAgent(ctx, binPath, name) + if err != nil { + return err + } + if !registered { + return fmt.Errorf("external agent %q binary %q was found but could not be registered", name, binPath) + } + return nil +} + // discoverAndRegister contains the shared scanning logic for external agent discovery. func discoverAndRegister(ctx context.Context) { ctx, cancel := context.WithTimeout(ctx, discoveryTimeout) @@ -93,44 +146,60 @@ func discoverAndRegister(ctx context.Context) { continue } - finfo, err := os.Stat(binPath) //nolint:gosec // PATH entries are trusted - if err != nil || finfo.IsDir() { - continue - } - // Check executable bit (on Unix; Windows doesn't set execute bits) - if runtime.GOOS != osWindows && finfo.Mode()&0o111 == 0 { - continue - } - - ea, err := New(ctx, binPath) + registeredAgent, err := registerExternalAgent(ctx, binPath, agentName) if err != nil { - logging.Debug(ctx, "skipping external agent (info failed)", + logging.Debug(ctx, "skipping external agent (registration failed)", slog.String("binary", binPath), + slog.String("agent", string(agentName)), slog.String("error", err.Error())) continue } - - // Wrap with capability interfaces and register - wrapped, err := Wrap(ea) - if err != nil { - logging.Debug(ctx, "skipping external agent (wrap failed)", - slog.String("binary", binPath), - slog.String("error", err.Error())) - continue + if registeredAgent { + registered[agentName] = true } - agent.Register(agentName, func() agent.Agent { - return wrapped - }) - registered[agentName] = true - - logging.Debug(ctx, "registered external agent", - slog.String("name", string(agentName)), - slog.String("type", string(ea.Type())), - slog.String("binary", binPath)) } } } +func registerExternalAgent(ctx context.Context, binPath string, name types.AgentName) (bool, error) { + finfo, err := statExternalAgent(binPath) + if err != nil { + if errors.Is(err, os.ErrNotExist) { + return false, nil + } + return false, fmt.Errorf("inspecting external agent %q binary %q: %w", name, binPath, err) + } + if finfo.IsDir() { + return false, nil + } + // Check executable bit (on Unix; Windows doesn't set execute bits). + if runtime.GOOS != osWindows && finfo.Mode()&0o111 == 0 { + return false, nil + } + + ea, err := New(ctx, binPath) + if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return false, fmt.Errorf("loading info for external agent %q from binary %q: %w: %w", name, binPath, ctxErr, err) + } + return false, fmt.Errorf("loading info for external agent %q from binary %q: %w", name, binPath, err) + } + + wrapped, err := Wrap(ea) + if err != nil { + return false, fmt.Errorf("wrapping external agent %q from binary %q: %w", name, binPath, err) + } + agent.Register(name, func() agent.Agent { + return wrapped + }) + + logging.Debug(ctx, "registered external agent", + slog.String("name", string(name)), + slog.String("type", string(ea.Type())), + slog.String("binary", binPath)) + return true, nil +} + // StripExeExt removes Windows executable extensions (.exe, .bat, .cmd, .com) // from a file name so that the derived name matches on all platforms. On Unix // this is effectively a no-op because binaries have no extension. diff --git a/cmd/entire/cli/agent/external/discovery_test.go b/cmd/entire/cli/agent/external/discovery_test.go index 60057b094d..0358f8c294 100644 --- a/cmd/entire/cli/agent/external/discovery_test.go +++ b/cmd/entire/cli/agent/external/discovery_test.go @@ -2,11 +2,15 @@ package external import ( "context" + "errors" + "fmt" "os" "os/exec" "path/filepath" "runtime" + "strings" "testing" + "time" "github.com/entireio/cli/cmd/entire/cli/agent" "github.com/entireio/cli/cmd/entire/cli/agent/types" @@ -290,6 +294,210 @@ func TestDiscoverAndRegisterAlways_FindsAgentWithoutSettings(t *testing.T) { } } +func TestDiscoverAndRegisterNamedAlways_TimesOutStalledInfo(t *testing.T) { + sleepPath, err := exec.LookPath("sleep") + if err != nil { + t.Skip("sleep not available") + } + + name := types.AgentName("disc-named-timeout") + dir := t.TempDir() + binPath := filepath.Join(dir, binaryPrefix+string(name)) + script := fmt.Sprintf("#!/bin/sh\nexec %q 60\n", sleepPath) + if err := os.WriteFile(binPath, []byte(script), 0o755); err != nil { + t.Fatalf("write stalled mock binary: %v", err) + } + t.Setenv("PATH", dir) + + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + + started := time.Now() + err = DiscoverAndRegisterNamedAlways(ctx, name) + if elapsed := time.Since(started); elapsed > 2*time.Second { + t.Fatalf("named discovery took %v, want cancellation near context deadline", elapsed) + } + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %v, want context deadline exceeded", err) + } + if _, err := agent.Get(name); err == nil { + t.Fatal("stalled external agent was registered") + } +} + +func TestDiscoverAndRegisterNamedAlways_CanceledContext(t *testing.T) { + if _, err := exec.LookPath("sh"); err != nil { + t.Skip("sh not available") + } + + name := types.AgentName("disc-named-canceled") + dir := setupDiscoveryDir(t, string(name), makeInfoJSON(string(name))) + t.Setenv("PATH", dir) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + err := DiscoverAndRegisterNamedAlways(ctx, name) + if !errors.Is(err, context.Canceled) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %v, want context canceled", err) + } +} + +func TestDiscoverAndRegisterNamedAlways_InvalidInfo(t *testing.T) { + if _, err := exec.LookPath("sh"); err != nil { + t.Skip("sh not available") + } + + name := types.AgentName("disc-named-invalid-info") + dir := setupDiscoveryDir(t, string(name), "not json") + t.Setenv("PATH", dir) + + err := DiscoverAndRegisterNamedAlways(context.Background(), name) + if err == nil { + t.Fatal("DiscoverAndRegisterNamedAlways() error = nil, want invalid info error") + } + if !strings.Contains(err.Error(), string(name)) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %q, want agent name %q", err, name) + } + if !strings.Contains(err.Error(), "info: invalid JSON") { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %q, want invalid info context", err) + } +} + +func TestDiscoverAndRegisterNamedAlways_MissingHelper(t *testing.T) { + name := types.AgentName("disc-named-missing") + t.Setenv("PATH", t.TempDir()) + + if err := DiscoverAndRegisterNamedAlways(context.Background(), name); err != nil { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %v, want nil for missing helper", err) + } +} + +func TestDiscoverAndRegisterNamedAlways_RejectsPathSeparators(t *testing.T) { + originalLookPath := lookPathExternalAgent + t.Cleanup(func() { lookPathExternalAgent = originalLookPath }) + + for _, name := range []types.AgentName{"foo/../../agent", `foo\bar`} { + lookedUp := false + lookPathExternalAgent = func(string) (string, error) { + lookedUp = true + return "", exec.ErrNotFound + } + + err := DiscoverAndRegisterNamedAlways(context.Background(), name) + if err == nil || !strings.Contains(err.Error(), "path separators") { + t.Errorf("DiscoverAndRegisterNamedAlways(%q) error = %v, want path separator error", name, err) + } + if lookedUp { + t.Errorf("DiscoverAndRegisterNamedAlways(%q) called exec.LookPath for an invalid name", name) + } + } +} + +func TestDiscoverAndRegisterNamedAlways_DeadlineWhileLookingUpMissingHelper(t *testing.T) { + name := types.AgentName("disc-named-lookup-deadline") + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + + originalLookPath := lookPathExternalAgent + t.Cleanup(func() { lookPathExternalAgent = originalLookPath }) + lookPathExternalAgent = func(string) (string, error) { + <-ctx.Done() + return "", exec.ErrNotFound + } + + err := DiscoverAndRegisterNamedAlways(ctx, name) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %v, want context deadline exceeded", err) + } +} + +func TestDiscoverAndRegisterNamedAlways_HelperDisappearsAfterLookup(t *testing.T) { + name := types.AgentName("disc-named-helper-disappeared") + binPath := filepath.Join(t.TempDir(), binaryPrefix+string(name)) + + originalLookPath := lookPathExternalAgent + originalStat := statExternalAgent + t.Cleanup(func() { + lookPathExternalAgent = originalLookPath + statExternalAgent = originalStat + }) + lookPathExternalAgent = func(string) (string, error) { return binPath, nil } + statExternalAgent = func(string) (os.FileInfo, error) { return nil, os.ErrNotExist } + + err := DiscoverAndRegisterNamedAlways(context.Background(), name) + if err == nil { + t.Fatal("DiscoverAndRegisterNamedAlways() error = nil, want helper-disappeared error") + } + if !strings.Contains(err.Error(), string(name)) || !strings.Contains(err.Error(), binPath) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %q, want agent and binary context", err) + } + if !strings.Contains(err.Error(), "was found but could not be registered") { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %q, want actionable registration context", err) + } +} + +func TestDiscoverAndRegisterNamedAlways_StatError(t *testing.T) { + name := types.AgentName("disc-named-stat-error") + binPath := filepath.Join(t.TempDir(), binaryPrefix+string(name)) + wantErr := errors.New("stat failed") + + originalLookPath := lookPathExternalAgent + originalStat := statExternalAgent + t.Cleanup(func() { + lookPathExternalAgent = originalLookPath + statExternalAgent = originalStat + }) + lookPathExternalAgent = func(string) (string, error) { return binPath, nil } + statExternalAgent = func(string) (os.FileInfo, error) { return nil, wantErr } + + err := DiscoverAndRegisterNamedAlways(context.Background(), name) + if !errors.Is(err, wantErr) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %v, want stat error", err) + } + if !strings.Contains(err.Error(), string(name)) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %q, want agent name %q", err, name) + } +} + +func TestDiscoverAndRegisterNamedAlways_LookPathError(t *testing.T) { + name := types.AgentName("disc-named-lookpath-error") + wantErr := errors.New("lookup failed") + + originalLookPath := lookPathExternalAgent + t.Cleanup(func() { lookPathExternalAgent = originalLookPath }) + lookPathExternalAgent = func(string) (string, error) { return "", wantErr } + + err := DiscoverAndRegisterNamedAlways(context.Background(), name) + if !errors.Is(err, wantErr) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %v, want lookup error", err) + } + if !strings.Contains(err.Error(), string(name)) { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %q, want agent name %q", err, name) + } +} + +func TestDiscoverAndRegister_ContinuesAfterRegistrationError(t *testing.T) { + if _, err := exec.LookPath("sh"); err != nil { + t.Skip("sh not available") + } + + badName := "disc-scan-a-invalid" + badDir := setupDiscoveryDir(t, badName, "not json") + goodName := "disc-scan-z-valid" + goodDir := setupDiscoveryDir(t, goodName, makeInfoJSON(goodName)) + t.Setenv("PATH", badDir+string(os.PathListSeparator)+goodDir) + + DiscoverAndRegisterAlways(context.Background()) + + if _, err := agent.Get(types.AgentName(badName)); err == nil { + t.Fatalf("invalid external agent %q was registered", badName) + } + if _, err := agent.Get(types.AgentName(goodName)); err != nil { + t.Fatalf("valid external agent %q was not registered after earlier failure: %v", goodName, err) + } +} + func TestIsExternal_WrappedAgent(t *testing.T) { if _, err := exec.LookPath("sh"); err != nil { t.Skip("sh not available") @@ -411,3 +619,35 @@ func TestDiscoverAndRegister_RegistersBatOnWindows(t *testing.T) { t.Errorf("agent Name() = %q, want %q", ag.Name(), name) } } + +// TestDiscoverAndRegisterNamedAlways_RegistersBatOnWindows covers the explicit +// named-discovery path, which uses exec.LookPath and therefore depends on +// Windows PATHEXT handling rather than the scan-all filepath.Glob path above. +func TestDiscoverAndRegisterNamedAlways_RegistersBatOnWindows(t *testing.T) { + if runtime.GOOS != osWindows { + t.Skip("this test only applies on Windows") + } + + name := types.AgentName("disc-named-bat") + infoJSON := `{"protocol_version":1,"name":"` + string(name) + `","type":"` + string(name) + ` Agent","description":"Named Windows agent","is_preview":false,"protected_dirs":[],"hook_names":[],"capabilities":{}}` + script := "@echo off\r\nif not \"%1\"==\"info\" goto :notinfo\r\necho " + infoJSON + "\r\ngoto :eof\r\n:notinfo\r\necho unknown subcommand: %1 1>&2\r\nexit /b 1\r\n" + + dir := t.TempDir() + binPath := filepath.Join(dir, binaryPrefix+string(name)+".bat") + if err := os.WriteFile(binPath, []byte(script), 0o755); err != nil { + t.Fatalf("write mock binary: %v", err) + } + t.Setenv("PATH", dir) + t.Setenv("PATHEXT", ".COM;.EXE;.BAT;.CMD") + + if err := DiscoverAndRegisterNamedAlways(context.Background(), name); err != nil { + t.Fatalf("DiscoverAndRegisterNamedAlways() error = %v", err) + } + ag, err := agent.Get(name) + if err != nil { + t.Fatalf("expected named .bat agent %q to be registered: %v", name, err) + } + if ag.Name() != name { + t.Fatalf("agent Name() = %q, want %q", ag.Name(), name) + } +} diff --git a/cmd/entire/cli/dispatch.go b/cmd/entire/cli/dispatch.go index 52804b859a..569bf706bb 100644 --- a/cmd/entire/cli/dispatch.go +++ b/cmd/entire/cli/dispatch.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "os" + "strings" dispatchpkg "github.com/entireio/cli/cmd/entire/cli/dispatch" "github.com/entireio/cli/cmd/entire/cli/interactive" @@ -18,6 +19,10 @@ var renderDispatchMarkdown = dispatchpkg.RenderMarkdown var dispatchTerminalMode = interactive.IsTerminalWriter var runInteractiveDispatch = defaultRunInteractiveDispatch var renderTerminalMarkdown = defaultRenderTerminalMarkdown +var shouldRunDispatchWizardForCommand = shouldRunDispatchWizard +var runDispatchWizardForCommand = runDispatchWizard +var prepareLocalDispatch = dispatchpkg.PrepareLocal +var resolveDispatchProvider = resolveDispatchSummaryProvider func newDispatchCmd() *cobra.Command { var ( @@ -27,6 +32,7 @@ func newDispatchCmd() *cobra.Command { flagAllBranches bool flagRepos []string flagVoice string + flagAgent string flagInsecureHTTPAuth bool ) @@ -38,16 +44,26 @@ func newDispatchCmd() *cobra.Command { Examples: entire dispatch entire dispatch --local --all-branches + entire dispatch --local --agent codex entire dispatch --repos entireio/cli entire dispatch --voice neutral`, RunE: func(cmd *cobra.Command, _ []string) error { + agentOverride := strings.TrimSpace(flagAgent) + agentFlagSet := cmd.Flags().Changed("agent") + if agentFlagSet && !flagLocal { + return errors.New("--agent only applies to --local (cloud dispatch uses Entire's server-side generator)") + } + if agentFlagSet && agentOverride == "" { + return errors.New("--agent requires a non-empty value") + } + var ( opts dispatchpkg.Options err error ) - if shouldRunDispatchWizard(cmd.Flags().NFlag(), isTerminalStdin(os.Stdin), interactive.IsTerminalWriter(cmd.OutOrStdout())) { - opts, err = runDispatchWizard(cmd) + if shouldRunDispatchWizardForCommand(cmd.Flags().NFlag(), isTerminalStdin(os.Stdin), interactive.IsTerminalWriter(cmd.OutOrStdout())) { + opts, err = runDispatchWizardForCommand(cmd) } else { opts, err = parseDispatchFlags(cmd, flagLocal, flagSince, flagUntil, flagAllBranches, flagRepos, flagVoice, flagInsecureHTTPAuth) } @@ -57,6 +73,18 @@ Examples: } return err } + if opts.Mode == dispatchpkg.ModeLocal { + opts, err = prepareLocalDispatch(cmd.Context(), opts) + if err != nil { + return err + } + provider, err := resolveDispatchProvider(cmd.Context(), cmd.ErrOrStderr(), agentOverride) + if err != nil { + return err + } + opts.TextGenerator = provider.TextGenerator + opts.Model = provider.Model + } if err := runDispatchCommand(cmd.Context(), cmd.OutOrStdout(), opts); err != nil { if errors.Is(err, errDispatchCancelled) { @@ -74,6 +102,7 @@ Examples: cmd.Flags().BoolVar(&flagAllBranches, "all-branches", false, "include every existing local branch (--local only; renamed or deleted branches are skipped)") cmd.Flags().StringSliceVar(&flagRepos, "repos", nil, fmt.Sprintf("cloud repo slugs, up to %d (for example entireio/cli)", dispatchpkg.CloudRepoLimit)) cmd.Flags().StringVar(&flagVoice, "voice", "", "voice preset name or literal description") + cmd.Flags().StringVar(&flagAgent, "agent", "", "local text-generation agent (requires --local)") cmd.Flags().BoolVar(&flagInsecureHTTPAuth, "insecure-http-auth", false, "Allow authentication over plain HTTP (insecure, for local development only)") if err := cmd.Flags().MarkHidden("insecure-http-auth"); err != nil { panic(fmt.Sprintf("hide insecure-http-auth flag: %v", err)) diff --git a/cmd/entire/cli/dispatch/dispatch.go b/cmd/entire/cli/dispatch/dispatch.go index 73dcd23c2f..8dedd4e5cf 100644 --- a/cmd/entire/cli/dispatch/dispatch.go +++ b/cmd/entire/cli/dispatch/dispatch.go @@ -12,6 +12,10 @@ const ( ModeLocal ) +type TextGenerator interface { + GenerateText(ctx context.Context, prompt string, model string) (string, error) +} + func (m Mode) String() string { switch m { case ModeServer: @@ -33,6 +37,9 @@ type Options struct { ImplicitCurrentBranch bool Voice string InsecureHTTPAuth bool + TextGenerator TextGenerator + Model string + localPreflight *localPreflight } // CloudRepoLimit caps how many repos the cloud mode may query in one request. diff --git a/cmd/entire/cli/dispatch/generate.go b/cmd/entire/cli/dispatch/generate.go index 339a295b07..6767e0cf58 100644 --- a/cmd/entire/cli/dispatch/generate.go +++ b/cmd/entire/cli/dispatch/generate.go @@ -8,28 +8,18 @@ import ( "strings" "time" - "github.com/entireio/cli/cmd/entire/cli/agent" - "github.com/entireio/cli/cmd/entire/cli/agent/claudecode" "github.com/entireio/cli/cmd/entire/cli/jsonutil" - "github.com/entireio/cli/cmd/entire/cli/summarize" ) -type dispatchTextGenerator interface { - GenerateText(ctx context.Context, prompt string, model string) (string, error) -} - -var dispatchTextGeneratorFactory = func() (dispatchTextGenerator, error) { - textGenerator, ok := agent.AsTextGenerator(claudecode.NewClaudeCodeAgent()) - if !ok { - return nil, errors.New("default dispatch generator does not support text generation") - } - return textGenerator, nil -} - -func generateLocalDispatch(ctx context.Context, dispatch *Dispatch, voice string) (string, error) { - textGenerator, err := dispatchTextGeneratorFactory() - if err != nil { - return "", err +func generateLocalDispatch( + ctx context.Context, + dispatch *Dispatch, + voice string, + textGenerator TextGenerator, + model string, +) (string, error) { + if textGenerator == nil { + return "", errors.New("local dispatch text generator is not configured") } prompt, err := buildDispatchPrompt(dispatch, voice) @@ -37,7 +27,7 @@ func generateLocalDispatch(ctx context.Context, dispatch *Dispatch, voice string return "", err } - text, err := textGenerator.GenerateText(ctx, prompt, summarize.DefaultModel) + text, err := textGenerator.GenerateText(ctx, prompt, model) if err != nil { return "", fmt.Errorf("generate dispatch text: %w", err) } diff --git a/cmd/entire/cli/dispatch/generate_test.go b/cmd/entire/cli/dispatch/generate_test.go index a89eb0271c..6202b5ccf9 100644 --- a/cmd/entire/cli/dispatch/generate_test.go +++ b/cmd/entire/cli/dispatch/generate_test.go @@ -10,10 +10,10 @@ import ( ) func TestGenerateLocalDispatch_UsesVoiceAndBullets(t *testing.T) { - mock := &stubTextGenerator{text: "generated dispatch"} - oldFactory := dispatchTextGeneratorFactory - dispatchTextGeneratorFactory = func() (dispatchTextGenerator, error) { return mock, nil } - t.Cleanup(func() { dispatchTextGeneratorFactory = oldFactory }) + t.Parallel() + + mock := &stubTextGenerator{text: " generated dispatch\n"} + var generator TextGenerator = mock dispatch := &Dispatch{ Repos: []RepoGroup{{ @@ -27,13 +27,20 @@ func TestGenerateLocalDispatch_UsesVoiceAndBullets(t *testing.T) { }}, } - got, err := generateLocalDispatch(context.Background(), dispatch, "marvin") + expectedPrompt, err := buildDispatchPrompt(dispatch, "marvin") + if err != nil { + t.Fatal(err) + } + got, err := generateLocalDispatch(context.Background(), dispatch, "marvin", generator, "test-model") if err != nil { t.Fatal(err) } if got != "generated dispatch" { t.Fatalf("unexpected text: %q", got) } + if mock.prompt != expectedPrompt { + t.Fatalf("unexpected prompt:\n%s\nwant:\n%s", mock.prompt, expectedPrompt) + } if !strings.Contains(mock.prompt, "You write concise markdown engineering dispatches.") { t.Fatalf("missing server instruction block in prompt: %s", mock.prompt) } @@ -52,9 +59,14 @@ func TestGenerateLocalDispatch_UsesVoiceAndBullets(t *testing.T) { if !strings.Contains(mock.prompt, "Write the final dispatch in markdown.") { t.Fatalf("missing final dispatch instruction in prompt: %s", mock.prompt) } + if mock.model != "test-model" { + t.Fatalf("unexpected model: %q", mock.model) + } } func TestBuildDispatchPrompt_SanitizesVoiceAndEscapesPromptTags(t *testing.T) { + t.Parallel() + dispatch := &Dispatch{ CoveredRepos: []string{"entireio/cli"}, Repos: []RepoGroup{{ @@ -257,26 +269,43 @@ func TestMarshalDispatchPromptPayload_OmitsRepoURLWhenFullNameSanitized(t *testi } func TestGenerateLocalDispatch_PropagatesGeneratorError(t *testing.T) { - oldFactory := dispatchTextGeneratorFactory - dispatchTextGeneratorFactory = func() (dispatchTextGenerator, error) { - return &stubTextGenerator{err: errors.New("boom")}, nil + t.Parallel() + + providerErr := errors.New("boom") + _, err := generateLocalDispatch( + context.Background(), + &Dispatch{}, + "", + &stubTextGenerator{err: providerErr}, + "test-model", + ) + if !errors.Is(err, providerErr) { + t.Fatalf("expected wrapped provider error, got %v", err) + } + if err.Error() != "generate dispatch text: boom" { + t.Fatalf("unexpected error: %v", err) } - t.Cleanup(func() { dispatchTextGeneratorFactory = oldFactory }) +} + +func TestGenerateLocalDispatch_RejectsNilGenerator(t *testing.T) { + t.Parallel() - _, err := generateLocalDispatch(context.Background(), &Dispatch{}, "") - if err == nil || !strings.Contains(err.Error(), "boom") { - t.Fatalf("expected generator error, got %v", err) + _, err := generateLocalDispatch(context.Background(), &Dispatch{}, "", nil, "test-model") + if err == nil || err.Error() != "local dispatch text generator is not configured" { + t.Fatalf("unexpected error: %v", err) } } type stubTextGenerator struct { prompt string + model string text string err error } -func (s *stubTextGenerator) GenerateText(_ context.Context, prompt string, _ string) (string, error) { +func (s *stubTextGenerator) GenerateText(_ context.Context, prompt string, model string) (string, error) { s.prompt = prompt + s.model = model if s.err != nil { return "", s.err } diff --git a/cmd/entire/cli/dispatch/mode_local.go b/cmd/entire/cli/dispatch/mode_local.go index 9659e0a8f5..3e262a357a 100644 --- a/cmd/entire/cli/dispatch/mode_local.go +++ b/cmd/entire/cli/dispatch/mode_local.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "os/exec" + "slices" "sort" "strings" "sync" @@ -35,7 +36,30 @@ var ( nowUTC = func() time.Time { return time.Now().UTC() } ) -func runLocal(ctx context.Context, opts Options) (*Dispatch, error) { +type localPreflight struct { + normalizedSince time.Time + normalizedUntil time.Time + repoRoots []string + sinceInput string + untilInput string + repoPathsInput []string +} + +func (p *localPreflight) matches(opts Options) bool { + return p != nil && + p.sinceInput == opts.Since && + p.untilInput == opts.Until && + slices.Equal(p.repoPathsInput, opts.RepoPaths) +} + +// PrepareLocal validates and resolves the inputs needed before local dispatch +// generation can begin. The returned options can be passed to Run without +// repeating time-window parsing or repository-root discovery. +func PrepareLocal(ctx context.Context, opts Options) (Options, error) { + if opts.Mode != ModeLocal { + return Options{}, errors.New("local dispatch preflight requires local mode") + } + now := nowUTC() sinceInput := strings.TrimSpace(opts.Since) if sinceInput == "" { @@ -43,28 +67,49 @@ func runLocal(ctx context.Context, opts Options) (*Dispatch, error) { } since, err := ParseSinceAtNow(sinceInput, now) if err != nil { - return nil, err + return Options{}, err } until, err := ParseUntilAtNow(opts.Until, now) if err != nil { - return nil, err + return Options{}, err } normalizedSince, normalizedUntil := NormalizeWindow(since, until) if !normalizedSince.Before(normalizedUntil) { - return nil, errors.New("--since must be before --until") + return Options{}, errors.New("--since must be before --until") } repoRoots, err := resolveRepoRoots(ctx, opts.RepoPaths) if err != nil { - return nil, err + return Options{}, err + } + + opts.localPreflight = &localPreflight{ + normalizedSince: normalizedSince, + normalizedUntil: normalizedUntil, + repoRoots: repoRoots, + sinceInput: opts.Since, + untilInput: opts.Until, + repoPathsInput: slices.Clone(opts.RepoPaths), + } + return opts, nil +} + +func runLocal(ctx context.Context, opts Options) (*Dispatch, error) { + if !opts.localPreflight.matches(opts) { + prepared, err := PrepareLocal(ctx, opts) + if err != nil { + return nil, err + } + opts = prepared } + preflight := opts.localPreflight allCandidates := make([]candidate, 0) var candidatesMu sync.Mutex group, groupCtx := errgroup.WithContext(ctx) - for _, repoRoot := range repoRoots { + for _, repoRoot := range preflight.repoRoots { group.Go(func() error { - candidates, err := enumerateRepoCandidates(groupCtx, repoRoot, opts, normalizedSince, normalizedUntil) + candidates, err := enumerateRepoCandidates(groupCtx, repoRoot, opts, preflight.normalizedSince, preflight.normalizedUntil) if err != nil { return err } @@ -83,14 +128,14 @@ func runLocal(ctx context.Context, opts Options) (*Dispatch, error) { CoveredRepos: coveredRepos(allCandidates), Repos: groupBulletsByRepo(fallback.Used), Window: Window{ - NormalizedSince: normalizedSince, - NormalizedUntil: normalizedUntil, + NormalizedSince: preflight.normalizedSince, + NormalizedUntil: preflight.normalizedUntil, FirstCheckpointAt: firstAt(fallback.Used), LastCheckpointAt: lastAt(fallback.Used), }, } - text, err := generateLocalDispatch(ctx, dispatch, opts.Voice) + text, err := generateLocalDispatch(ctx, dispatch, opts.Voice, opts.TextGenerator, opts.Model) if err != nil { return nil, err } diff --git a/cmd/entire/cli/dispatch/mode_local_test.go b/cmd/entire/cli/dispatch/mode_local_test.go index 24ff20edee..4241a87d7e 100644 --- a/cmd/entire/cli/dispatch/mode_local_test.go +++ b/cmd/entire/cli/dispatch/mode_local_test.go @@ -20,9 +20,201 @@ import ( "github.com/go-git/go-git/v6/plumbing/object" ) +func TestPrepareLocal_RejectsServerMode(t *testing.T) { + t.Parallel() + + _, err := PrepareLocal(context.Background(), Options{Mode: ModeServer}) + if err == nil || !strings.Contains(err.Error(), "local") { + t.Fatalf("expected local-mode error, got %v", err) + } +} + +func TestPrepareLocal_RejectsInvalidSince(t *testing.T) { + t.Parallel() + + _, err := PrepareLocal(context.Background(), Options{ + Mode: ModeLocal, + Since: "definitely-not-a-time", + }) + if err == nil || !strings.Contains(err.Error(), "unparseable time") { + t.Fatalf("expected invalid --since error, got %v", err) + } +} + +func TestPrepareLocal_RejectsInvalidUntil(t *testing.T) { + t.Parallel() + + _, err := PrepareLocal(context.Background(), Options{ + Mode: ModeLocal, + Since: "2026-07-16T12:00:00Z", + Until: "definitely-not-a-time", + }) + if err == nil || !strings.Contains(err.Error(), "unparseable time") { + t.Fatalf("expected invalid --until error, got %v", err) + } +} + +func TestPrepareLocal_RejectsReversedNormalizedWindow(t *testing.T) { + t.Parallel() + + _, err := PrepareLocal(context.Background(), Options{ + Mode: ModeLocal, + Since: "2026-07-17T12:01:00Z", + Until: "2026-07-17T12:00:00Z", + }) + if err == nil || err.Error() != "--since must be before --until" { + t.Fatalf("expected reversed-window error, got %v", err) + } +} + +func TestPrepareLocal_RejectsEqualNormalizedWindow(t *testing.T) { + t.Parallel() + + _, err := PrepareLocal(context.Background(), Options{ + Mode: ModeLocal, + Since: "2026-07-17T12:00:00Z", + Until: "2026-07-17T12:00:00Z", + }) + if err == nil || err.Error() != "--since must be before --until" { + t.Fatalf("expected equal-window error, got %v", err) + } +} + +func TestPrepareLocal_ValidWindowOutsideGitFailsRepoResolution(t *testing.T) { + t.Chdir(t.TempDir()) + + _, err := PrepareLocal(context.Background(), Options{ + Mode: ModeLocal, + Since: "2026-07-16T12:00:00Z", + Until: "2026-07-17T12:00:00Z", + }) + if err == nil || !strings.Contains(err.Error(), "not in a git repository") { + t.Fatalf("expected repository-root error, got %v", err) + } +} + +func TestPrepareLocal_RunAutoPreparesDirectCall(t *testing.T) { + t.Chdir(t.TempDir()) + + _, err := Run(context.Background(), Options{ + Mode: ModeLocal, + Since: "2026-07-16T12:00:00Z", + Until: "2026-07-17T12:00:00Z", + AllBranches: true, + TextGenerator: stubGeneratedLocalDispatch(), + }) + if err == nil || !strings.Contains(err.Error(), "not in a git repository") { + t.Fatalf("expected direct Run to perform repository preflight, got %v", err) + } +} + +func TestRunLocal_UsesPreparedWindowAndRepoRoots(t *testing.T) { + repoDir := t.TempDir() + testutil.InitRepo(t, repoDir) + testutil.WriteFile(t, repoDir, "a.txt", "x") + testutil.GitAdd(t, repoDir, "a.txt") + testutil.GitCommit(t, repoDir, "initial") + addOriginRemote(t, repoDir) + + preparedAt := time.Date(2026, 7, 17, 12, 34, 45, 0, time.UTC) + oldNow := nowUTC + nowUTC = func() time.Time { return preparedAt } + t.Cleanup(func() { nowUTC = oldNow }) + + t.Chdir(repoDir) + generator := stubGeneratedLocalDispatch() + prepared, err := PrepareLocal(context.Background(), Options{ + Mode: ModeLocal, + Since: "1h", + AllBranches: true, + TextGenerator: generator, + Model: "prepared-model", + }) + if err != nil { + t.Fatal(err) + } + if prepared.TextGenerator != generator || prepared.Model != "prepared-model" { + t.Fatal("preflight must preserve injected generation options") + } + + nowUTC = func() time.Time { return preparedAt.Add(24 * time.Hour) } + t.Chdir(t.TempDir()) + got, err := Run(context.Background(), prepared) + if err != nil { + t.Fatal(err) + } + + wantSince := time.Date(2026, 7, 17, 11, 34, 0, 0, time.UTC) + wantUntil := time.Date(2026, 7, 17, 12, 35, 0, 0, time.UTC) + if !got.Window.NormalizedSince.Equal(wantSince) || !got.Window.NormalizedUntil.Equal(wantUntil) { + t.Fatalf("prepared window was recomputed: got [%s, %s), want [%s, %s)", + got.Window.NormalizedSince, got.Window.NormalizedUntil, wantSince, wantUntil) + } + if got.GeneratedText != "generated dispatch" { + t.Fatalf("unexpected generated text: %q", got.GeneratedText) + } +} + +func TestRunLocal_RepreparesWhenPreflightInputsChange(t *testing.T) { + repoDir := t.TempDir() + testutil.InitRepo(t, repoDir) + testutil.WriteFile(t, repoDir, "a.txt", "x") + testutil.GitAdd(t, repoDir, "a.txt") + testutil.GitCommit(t, repoDir, "initial") + addOriginRemote(t, repoDir) + t.Chdir(repoDir) + + tests := []struct { + name string + mutate func(*Options) + wantError string + }{ + { + name: "since", + mutate: func(opts *Options) { + opts.Since = "definitely-not-a-time" + }, + wantError: "unparseable time", + }, + { + name: "until", + mutate: func(opts *Options) { + opts.Until = "definitely-not-a-time" + }, + wantError: "unparseable time", + }, + { + name: "repo paths", + mutate: func(opts *Options) { + opts.RepoPaths = []string{filepath.Join(t.TempDir(), "missing")} + }, + wantError: "resolve repo root", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + prepared, err := PrepareLocal(context.Background(), Options{ + Mode: ModeLocal, + Since: "7d", + AllBranches: true, + TextGenerator: stubGeneratedLocalDispatch(), + }) + if err != nil { + t.Fatal(err) + } + + tt.mutate(&prepared) + _, err = Run(context.Background(), prepared) + if err == nil || !strings.Contains(err.Error(), tt.wantError) { + t.Fatalf("Run() error = %v, want error containing %q", err, tt.wantError) + } + }) + } +} + func TestLocalMode_EnumeratesCheckpoints(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) testutil.WriteFile(t, dir, "a.txt", "x") testutil.GitAdd(t, dir, "a.txt") @@ -47,9 +239,10 @@ func TestLocalMode_EnumeratesCheckpoints(t *testing.T) { t.Chdir(dir) got, err := Run(context.Background(), Options{ - Mode: ModeLocal, - Since: "7d", - Branches: []string{"main"}, + Mode: ModeLocal, + Since: "7d", + Branches: []string{"main"}, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -74,7 +267,6 @@ func TestLocalMode_EnumeratesCheckpoints(t *testing.T) { func TestLocalMode_ExplicitRepoUsesTargetRepoCheckpointSettings(t *testing.T) { cwdDir := t.TempDir() targetDir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, cwdDir) if err := os.MkdirAll(filepath.Join(cwdDir, ".entire"), 0o755); err != nil { @@ -112,10 +304,11 @@ func TestLocalMode_ExplicitRepoUsesTargetRepoCheckpointSettings(t *testing.T) { t.Chdir(cwdDir) got, err := Run(context.Background(), Options{ - Mode: ModeLocal, - RepoPaths: []string{targetDir}, - Since: "7d", - Branches: []string{"main"}, + Mode: ModeLocal, + RepoPaths: []string{targetDir}, + Since: "7d", + Branches: []string{"main"}, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -130,7 +323,6 @@ func TestLocalMode_ExplicitRepoUsesTargetRepoCheckpointSettings(t *testing.T) { func TestLocalMode_UsesUntilWindow(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) testutil.WriteFile(t, dir, "a.txt", "x") testutil.GitAdd(t, dir, "a.txt") @@ -155,10 +347,11 @@ func TestLocalMode_UsesUntilWindow(t *testing.T) { t.Chdir(dir) got, err := Run(context.Background(), Options{ - Mode: ModeLocal, - Since: "7d", - Until: now.Add(-time.Hour).Format(time.RFC3339), - Branches: []string{"main"}, + Mode: ModeLocal, + Since: "7d", + Until: now.Add(-time.Hour).Format(time.RFC3339), + Branches: []string{"main"}, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -170,7 +363,6 @@ func TestLocalMode_UsesUntilWindow(t *testing.T) { func TestLocalMode_FallsBackToCommitSubjectWhenSummaryMissing(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) testutil.WriteFile(t, dir, "a.txt", "x") testutil.GitAdd(t, dir, "a.txt") @@ -199,9 +391,10 @@ func TestLocalMode_FallsBackToCommitSubjectWhenSummaryMissing(t *testing.T) { t.Chdir(dir) got, err := Run(context.Background(), Options{ - Mode: ModeLocal, - Since: "7d", - Branches: []string{"main"}, + Mode: ModeLocal, + Since: "7d", + Branches: []string{"main"}, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -232,23 +425,20 @@ func TestLocalMode_GenerateProducesInlineText(t *testing.T) { }) oldNow := nowUTC - oldFactory := dispatchTextGeneratorFactory nowUTC = func() time.Time { return createdAt.Add(2 * time.Hour) } mock := &stubTextGenerator{text: "generated inline dispatch"} - dispatchTextGeneratorFactory = func() (dispatchTextGenerator, error) { - return mock, nil - } t.Cleanup(func() { nowUTC = oldNow - dispatchTextGeneratorFactory = oldFactory }) t.Chdir(dir) got, err := Run(context.Background(), Options{ - Mode: ModeLocal, - Since: "7d", - Branches: []string{"main"}, + Mode: ModeLocal, + Since: "7d", + Branches: []string{"main"}, + TextGenerator: mock, + Model: "test-model", }) if err != nil { t.Fatal(err) @@ -256,6 +446,9 @@ func TestLocalMode_GenerateProducesInlineText(t *testing.T) { if got.GeneratedText != "generated inline dispatch" { t.Fatalf("expected generated text, got %q", got.GeneratedText) } + if mock.model != "test-model" { + t.Fatalf("unexpected model: %q", mock.model) + } } func TestLocalMode_FailsWhenGeneratedMarkdownIsEmpty(t *testing.T) { @@ -276,22 +469,18 @@ func TestLocalMode_FailsWhenGeneratedMarkdownIsEmpty(t *testing.T) { }) oldNow := nowUTC - oldFactory := dispatchTextGeneratorFactory nowUTC = func() time.Time { return createdAt.Add(2 * time.Hour) } - dispatchTextGeneratorFactory = func() (dispatchTextGenerator, error) { - return &stubTextGenerator{text: " \n\t "}, nil - } t.Cleanup(func() { nowUTC = oldNow - dispatchTextGeneratorFactory = oldFactory }) t.Chdir(dir) _, err := Run(context.Background(), Options{ - Mode: ModeLocal, - Since: "7d", - Branches: []string{"main"}, + Mode: ModeLocal, + Since: "7d", + Branches: []string{"main"}, + TextGenerator: &stubTextGenerator{text: " \n\t "}, }) if err == nil { t.Fatal("expected error when local generation returns empty markdown") @@ -303,7 +492,6 @@ func TestLocalMode_FailsWhenGeneratedMarkdownIsEmpty(t *testing.T) { func TestLocalMode_ImplicitCurrentBranchUsesHEADReachability(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) testutil.WriteFile(t, dir, "a.txt", "x") testutil.GitAdd(t, dir, "a.txt") @@ -360,6 +548,7 @@ func TestLocalMode_ImplicitCurrentBranchUsesHEADReachability(t *testing.T) { Since: "7d", Branches: []string{"entire-dispatch-codex"}, ImplicitCurrentBranch: true, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -371,7 +560,6 @@ func TestLocalMode_ImplicitCurrentBranchUsesHEADReachability(t *testing.T) { func TestLocalMode_ExplicitBranchesRemainExact(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) testutil.WriteFile(t, dir, "a.txt", "x") testutil.GitAdd(t, dir, "a.txt") @@ -423,9 +611,10 @@ func TestLocalMode_ExplicitBranchesRemainExact(t *testing.T) { t.Chdir(dir) got, err := Run(context.Background(), Options{ - Mode: ModeLocal, - Since: "7d", - Branches: []string{"entire-dispatch-codex"}, + Mode: ModeLocal, + Since: "7d", + Branches: []string{"entire-dispatch-codex"}, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -437,7 +626,6 @@ func TestLocalMode_ExplicitBranchesRemainExact(t *testing.T) { func TestLocalMode_ImplicitCurrentBranchUsesCheckpointBranchWithoutTrailerReachability(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) testutil.WriteFile(t, dir, "a.txt", "x") testutil.GitAdd(t, dir, "a.txt") @@ -468,6 +656,7 @@ func TestLocalMode_ImplicitCurrentBranchUsesCheckpointBranchWithoutTrailerReacha Since: "7d", Branches: []string{"entire-dispatch-codex"}, ImplicitCurrentBranch: true, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -479,7 +668,6 @@ func TestLocalMode_ImplicitCurrentBranchUsesCheckpointBranchWithoutTrailerReacha func TestLocalMode_ImplicitCurrentBranchExcludesDefaultBranchHistory(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) addOriginRemote(t, dir) @@ -522,6 +710,7 @@ func TestLocalMode_ImplicitCurrentBranchExcludesDefaultBranchHistory(t *testing. Since: "7d", Branches: []string{"my-feature"}, ImplicitCurrentBranch: true, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -548,7 +737,6 @@ func TestLocalMode_ImplicitCurrentBranchExcludesDefaultBranchHistory(t *testing. // the dispatch came back empty. func TestLocalMode_ImplicitCurrentBranchOnDefaultBranchIncludesMergedWork(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) addOriginRemote(t, dir) @@ -590,6 +778,7 @@ func TestLocalMode_ImplicitCurrentBranchOnDefaultBranchIncludesMergedWork(t *tes Since: "7d", Branches: []string{defaultBranch}, ImplicitCurrentBranch: true, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -610,7 +799,6 @@ func TestLocalMode_ImplicitCurrentBranchOnDefaultBranchIncludesMergedWork(t *tes // store.List, which only sees local checkpoints, so this work was invisible. func TestLocalMode_IncludesReachableCheckpointMissingFromLocalStore(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) addOriginRemote(t, dir) @@ -642,6 +830,7 @@ func TestLocalMode_IncludesReachableCheckpointMissingFromLocalStore(t *testing.T Since: "7d", Branches: []string{defaultBranch}, ImplicitCurrentBranch: true, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -656,7 +845,6 @@ func TestLocalMode_IncludesReachableCheckpointMissingFromLocalStore(t *testing.T func TestLocalMode_AllBranchesRestrictsToLocalBranches(t *testing.T) { dir := t.TempDir() - stubGeneratedLocalDispatch(t) testutil.InitRepo(t, dir) testutil.WriteFile(t, dir, "a.txt", "x") testutil.GitAdd(t, dir, "a.txt") @@ -688,9 +876,10 @@ func TestLocalMode_AllBranchesRestrictsToLocalBranches(t *testing.T) { t.Chdir(dir) got, err := Run(context.Background(), Options{ - Mode: ModeLocal, - Since: "7d", - AllBranches: true, + Mode: ModeLocal, + Since: "7d", + AllBranches: true, + TextGenerator: stubGeneratedLocalDispatch(), }) if err != nil { t.Fatal(err) @@ -994,16 +1183,8 @@ type seededCheckpoint struct { outcome string } -func stubGeneratedLocalDispatch(t *testing.T) { - t.Helper() - - oldFactory := dispatchTextGeneratorFactory - dispatchTextGeneratorFactory = func() (dispatchTextGenerator, error) { - return &stubTextGenerator{text: "generated dispatch"}, nil - } - t.Cleanup(func() { - dispatchTextGeneratorFactory = oldFactory - }) +func stubGeneratedLocalDispatch() TextGenerator { + return &stubTextGenerator{text: "generated dispatch"} } func seedCommittedCheckpoint(t *testing.T, repoDir string, cp seededCheckpoint) { diff --git a/cmd/entire/cli/dispatch_test.go b/cmd/entire/cli/dispatch_test.go index 5b16ba988a..2e5dd29dd2 100644 --- a/cmd/entire/cli/dispatch_test.go +++ b/cmd/entire/cli/dispatch_test.go @@ -3,11 +3,16 @@ package cli import ( "bytes" "context" + "errors" "io" "strings" "testing" + "github.com/entireio/cli/cmd/entire/cli/agent" + "github.com/entireio/cli/cmd/entire/cli/agent/types" dispatchpkg "github.com/entireio/cli/cmd/entire/cli/dispatch" + "github.com/entireio/cli/cmd/entire/cli/settings" + "github.com/entireio/cli/cmd/entire/cli/testutil" "github.com/spf13/cobra" ) @@ -201,6 +206,556 @@ func TestNewDispatchCmd_LocalHelpText(t *testing.T) { } } +func TestNewDispatchCmd_AgentFlagHelpText(t *testing.T) { + t.Parallel() + + cmd := newDispatchCmd() + flag := cmd.Flags().Lookup("agent") + if flag == nil { + t.Fatal("expected --agent flag to be registered") + } + want := "local text-generation agent (requires --local)" + if flag.Usage != want { + t.Fatalf("unexpected --agent help text: %q", flag.Usage) + } + if modelFlag := cmd.Flags().Lookup("model"); modelFlag != nil { + t.Fatal("did not expect --model flag to be registered") + } +} + +func TestNewDispatchCmd_LongHelpIncludesLocalAgentExample(t *testing.T) { + t.Parallel() + + cmd := newDispatchCmd() + if !strings.Contains(cmd.Long, "entire dispatch --local --agent codex") { + t.Fatalf("long help missing local-agent example:\n%s", cmd.Long) + } +} + +func TestDispatchPreflight_InvalidTimeBeforeProvider(t *testing.T) { + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + providerCalled := false + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + providerCalled = true + return nil, errors.New("provider must not run") + } + runDispatch = func(context.Context, dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + t.Fatal("dispatch must not run after preflight fails") + return nil, errors.New("dispatch must not run") + } + t.Cleanup(func() { + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + }) + + cmd := newDispatchCmd() + cmd.SilenceErrors = true + cmd.SilenceUsage = true + cmd.SetArgs([]string{"--local", "--all-branches", "--since", "not-a-time"}) + + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "unparseable time") { + t.Fatalf("expected invalid time preflight error, got %v", err) + } + if providerCalled { + t.Fatal("provider resolution ran before invalid time preflight returned") + } +} + +func TestDispatchPreflight_RepositoryFailureBeforeProvider(t *testing.T) { + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + providerCalled := false + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + providerCalled = true + return nil, errors.New("provider must not run") + } + runDispatch = func(context.Context, dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + t.Fatal("dispatch must not run after preflight fails") + return nil, errors.New("dispatch must not run") + } + t.Cleanup(func() { + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + }) + t.Chdir(t.TempDir()) + + cmd := newDispatchCmd() + cmd.SilenceErrors = true + cmd.SilenceUsage = true + cmd.SetArgs([]string{ + "--local", "--all-branches", + "--since", "2026-07-16T12:00:00Z", + "--until", "2026-07-17T12:00:00Z", + }) + + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "not in a git repository") { + t.Fatalf("expected repository preflight error, got %v", err) + } + if providerCalled { + t.Fatal("provider resolution ran before repository preflight returned") + } +} + +func TestDispatchProvider_LocalRunsAfterPreflightAndInjectsOptions(t *testing.T) { + oldPrepare := prepareLocalDispatch + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + oldTerminalMode := dispatchTerminalMode + oldMarkdown := renderDispatchMarkdown + var calls []string + generator := &stubTextAgent{} + prepareLocalDispatch = func(_ context.Context, opts dispatchpkg.Options) (dispatchpkg.Options, error) { + calls = append(calls, "prepare") + return opts, nil + } + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + calls = append(calls, "provider") + return &checkpointSummaryProvider{TextGenerator: generator, Model: "ordered-model"}, nil + } + runDispatch = func(_ context.Context, opts dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + calls = append(calls, "dispatch") + if opts.TextGenerator != generator || opts.Model != "ordered-model" { + t.Fatalf("provider options not passed to dispatch: generator=%T model=%q", opts.TextGenerator, opts.Model) + } + return &dispatchpkg.Dispatch{}, nil + } + dispatchTerminalMode = func(io.Writer) bool { return false } + renderDispatchMarkdown = func(*dispatchpkg.Dispatch) string { return "" } + t.Cleanup(func() { + prepareLocalDispatch = oldPrepare + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + dispatchTerminalMode = oldTerminalMode + renderDispatchMarkdown = oldMarkdown + }) + + cmd := newDispatchCmd() + cmd.SetArgs([]string{"--local", "--all-branches"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if got := strings.Join(calls, ","); got != "prepare,provider,dispatch" { + t.Fatalf("call order = %q, want prepare,provider,dispatch", got) + } +} + +func TestDispatchPreflight_CloudSkipsLocalPreparationAndProvider(t *testing.T) { + oldPrepare := prepareLocalDispatch + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + oldTerminalMode := dispatchTerminalMode + prepareLocalDispatch = func(context.Context, dispatchpkg.Options) (dispatchpkg.Options, error) { + t.Fatal("cloud dispatch must not run local preflight") + return dispatchpkg.Options{}, errors.New("local preflight must not run") + } + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + t.Fatal("cloud dispatch must not resolve a local provider") + return nil, errors.New("provider must not run") + } + runDispatch = func(_ context.Context, opts dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + if opts.Mode != dispatchpkg.ModeServer { + t.Fatalf("mode = %v, want server", opts.Mode) + } + return &dispatchpkg.Dispatch{}, nil + } + dispatchTerminalMode = func(io.Writer) bool { return false } + t.Cleanup(func() { + prepareLocalDispatch = oldPrepare + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + dispatchTerminalMode = oldTerminalMode + }) + + cmd := newDispatchCmd() + cmd.SetArgs([]string{"--repos", "entireio/cli"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } +} + +func TestNewDispatchCmd_CloudAgentFailsBeforeProviderOrDispatch(t *testing.T) { + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + unexpectedCallErr := errors.New("unexpected command dependency call") + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + t.Fatal("local provider resolution must not run for cloud --agent validation") + return nil, unexpectedCallErr + } + runDispatch = func(context.Context, dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + t.Fatal("dispatch must not run after invalid cloud --agent validation") + return nil, unexpectedCallErr + } + t.Cleanup(func() { + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + }) + + cmd := newDispatchCmd() + cmd.SilenceErrors = true + cmd.SilenceUsage = true + cmd.SetArgs([]string{"--agent", string(agent.AgentNameCodex)}) + + err := cmd.Execute() + if err == nil { + t.Fatal("expected cloud --agent validation error") + } + want := "--agent only applies to --local (cloud dispatch uses Entire's server-side generator)" + if err.Error() != want { + t.Fatalf("unexpected error: %q", err) + } +} + +func TestNewDispatchCmd_CloudExplicitEmptyAgentUsesLocalOnlyErrorPrecedence(t *testing.T) { + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + providerCalled := false + dispatchCalled := false + unexpectedCallErr := errors.New("unexpected command dependency call") + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + providerCalled = true + return nil, unexpectedCallErr + } + runDispatch = func(context.Context, dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + dispatchCalled = true + return nil, unexpectedCallErr + } + t.Cleanup(func() { + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + }) + + cmd := newDispatchCmd() + cmd.SilenceErrors = true + cmd.SilenceUsage = true + cmd.SetArgs([]string{"--agent="}) + + err := cmd.Execute() + want := "--agent only applies to --local (cloud dispatch uses Entire's server-side generator)" + if err == nil || err.Error() != want { + t.Fatalf("unexpected error: %v", err) + } + if providerCalled { + t.Fatal("local provider resolution must not run for cloud --agent validation") + } + if dispatchCalled { + t.Fatal("dispatch must not run after invalid cloud --agent validation") + } +} + +func TestNewDispatchCmd_LocalExplicitEmptyAgentFailsBeforeProviderOrDispatch(t *testing.T) { + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + providerCalled := false + dispatchCalled := false + unexpectedCallErr := errors.New("unexpected command dependency call") + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + providerCalled = true + return nil, unexpectedCallErr + } + runDispatch = func(context.Context, dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + dispatchCalled = true + return nil, unexpectedCallErr + } + t.Cleanup(func() { + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + }) + + for _, args := range [][]string{ + {"--local", "--all-branches", "--agent="}, + {"--local", "--all-branches", "--agent", " "}, + } { + providerCalled = false + dispatchCalled = false + cmd := newDispatchCmd() + cmd.SilenceErrors = true + cmd.SilenceUsage = true + cmd.SetArgs(args) + + err := cmd.Execute() + if err == nil || err.Error() != "--agent requires a non-empty value" { + t.Fatalf("args %q: unexpected error: %v", args, err) + } + if providerCalled { + t.Fatalf("args %q: provider resolution must not run", args) + } + if dispatchCalled { + t.Fatalf("args %q: dispatch must not run", args) + } + } +} + +func TestNewDispatchCmd_LocalAgentInjectsProviderAndKeepsOutputSeparated(t *testing.T) { + oldPrepare := prepareLocalDispatch + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + oldTerminalMode := dispatchTerminalMode + oldMarkdown := renderDispatchMarkdown + generator := &stubTextAgent{} + prepareLocalDispatch = func(_ context.Context, opts dispatchpkg.Options) (dispatchpkg.Options, error) { + return opts, nil + } + resolveDispatchProvider = func(_ context.Context, w io.Writer, override string) (*checkpointSummaryProvider, error) { + if override != string(agent.AgentNameCodex) { + t.Fatalf("provider override = %q, want codex", override) + } + if _, err := io.WriteString(w, "provider notice\n"); err != nil { + t.Fatal(err) + } + return &checkpointSummaryProvider{TextGenerator: generator, Model: "exact-model"}, nil + } + runDispatch = func(_ context.Context, opts dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + if opts.Mode != dispatchpkg.ModeLocal { + t.Fatalf("mode = %v, want local", opts.Mode) + } + if opts.TextGenerator != generator { + t.Fatalf("TextGenerator = %T, want raw provider generator", opts.TextGenerator) + } + if opts.Model != "exact-model" { + t.Fatalf("Model = %q, want exact-model", opts.Model) + } + return &dispatchpkg.Dispatch{GeneratedText: "generated dispatch"}, nil + } + dispatchTerminalMode = func(io.Writer) bool { return false } + renderDispatchMarkdown = func(*dispatchpkg.Dispatch) string { return testDispatchGeneratedMarkdown } + t.Cleanup(func() { + prepareLocalDispatch = oldPrepare + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + dispatchTerminalMode = oldTerminalMode + renderDispatchMarkdown = oldMarkdown + }) + + cmd := newDispatchCmd() + cmd.SilenceErrors = true + cmd.SilenceUsage = true + var stdout bytes.Buffer + var stderr bytes.Buffer + cmd.SetOut(&stdout) + cmd.SetErr(&stderr) + cmd.SetArgs([]string{"--local", "--all-branches", "--agent", " " + string(agent.AgentNameCodex) + " "}) + + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if got := stdout.String(); got != testDispatchGeneratedMarkdown { + t.Fatalf("unexpected stdout: %q", got) + } + if got := stderr.String(); got != "provider notice\n" { + t.Fatalf("unexpected stderr: %q", got) + } +} + +func TestNewDispatchCmd_LocalWithoutAgentResolvesConfiguredProvider(t *testing.T) { + oldPrepare := prepareLocalDispatch + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + oldTerminalMode := dispatchTerminalMode + prepareLocalDispatch = func(_ context.Context, opts dispatchpkg.Options) (dispatchpkg.Options, error) { + return opts, nil + } + resolveDispatchProvider = func(_ context.Context, _ io.Writer, override string) (*checkpointSummaryProvider, error) { + if override != "" { + t.Fatalf("provider override = %q, want empty", override) + } + return &checkpointSummaryProvider{TextGenerator: &stubTextAgent{}}, nil + } + runDispatch = func(context.Context, dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + return &dispatchpkg.Dispatch{}, nil + } + dispatchTerminalMode = func(io.Writer) bool { return false } + t.Cleanup(func() { + prepareLocalDispatch = oldPrepare + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + dispatchTerminalMode = oldTerminalMode + }) + + cmd := newDispatchCmd() + cmd.SetArgs([]string{"--local", "--all-branches"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } +} + +func TestDispatchWizard_LocalWithoutConfiguredAgentPromptsAndPersistsSelection(t *testing.T) { + repoDir := t.TempDir() + testutil.InitRepo(t, repoDir) + t.Chdir(repoDir) + + oldShouldRunWizard := shouldRunDispatchWizardForCommand + oldRunWizard := runDispatchWizardForCommand + oldLoad := loadSummarySettings + oldLoadFile := loadSummarySettingsFromFile + oldSave := saveLocalSummarySettings + oldDiscover := discoverSummaryProvidersAlways + oldList := listRegisteredAgents + oldGet := getSummaryAgent + oldAvailable := isSummaryCLIAvailable + oldCanPrompt := canPromptForSummaryProvider + oldPrompt := promptSummaryProvider + oldRunDispatch := runDispatch + oldTerminalMode := dispatchTerminalMode + oldMarkdown := renderDispatchMarkdown + t.Cleanup(func() { + shouldRunDispatchWizardForCommand = oldShouldRunWizard + runDispatchWizardForCommand = oldRunWizard + loadSummarySettings = oldLoad + loadSummarySettingsFromFile = oldLoadFile + saveLocalSummarySettings = oldSave + discoverSummaryProvidersAlways = oldDiscover + listRegisteredAgents = oldList + getSummaryAgent = oldGet + isSummaryCLIAvailable = oldAvailable + canPromptForSummaryProvider = oldCanPrompt + promptSummaryProvider = oldPrompt + runDispatch = oldRunDispatch + dispatchTerminalMode = oldTerminalMode + renderDispatchMarkdown = oldMarkdown + }) + + var calls []string + shouldRunDispatchWizardForCommand = func(int, bool, bool) bool { return true } + runDispatchWizardForCommand = func(*cobra.Command) (dispatchpkg.Options, error) { + calls = append(calls, "wizard") + return dispatchpkg.Options{Mode: dispatchpkg.ModeLocal, Since: "7d", AllBranches: true}, nil + } + loadSummarySettings = func(context.Context) (*settings.EntireSettings, error) { + return &settings.EntireSettings{Enabled: true}, nil + } + loadSummarySettingsFromFile = func(string) (*settings.EntireSettings, error) { + return &settings.EntireSettings{}, nil + } + discoverSummaryProvidersAlways = func(context.Context) {} + listRegisteredAgents = func() []types.AgentName { + return []types.AgentName{agent.AgentNameCodex, agent.AgentNameGemini} + } + getSummaryAgent = func(name types.AgentName) (agent.Agent, error) { + kind := agent.AgentTypeCodex + if name == agent.AgentNameGemini { + kind = agent.AgentTypeGemini + } + return &stubTextAgent{name: name, kind: kind}, nil + } + isSummaryCLIAvailable = func(types.AgentName) bool { return true } + canPromptForSummaryProvider = func() bool { return true } + promptSummaryProvider = func(providers []checkpointSummaryProvider) (types.AgentName, error) { + calls = append(calls, "picker") + if len(providers) != 2 || providers[0].Name != agent.AgentNameCodex || providers[1].Name != agent.AgentNameGemini { + t.Fatalf("picker providers = %+v, want enabled codex and gemini", providers) + } + return agent.AgentNameGemini, nil + } + var persistedProvider string + saveLocalSummarySettings = func(_ context.Context, s *settings.EntireSettings) error { + if s.SummaryGeneration != nil { + persistedProvider = s.SummaryGeneration.Provider + } + return nil + } + runDispatch = func(_ context.Context, opts dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + calls = append(calls, "dispatch") + selected, ok := opts.TextGenerator.(*stubTextAgent) + if !ok || selected.name != agent.AgentNameGemini { + t.Fatalf("dispatch generator = %#v, want selected gemini agent", opts.TextGenerator) + } + return &dispatchpkg.Dispatch{}, nil + } + dispatchTerminalMode = func(io.Writer) bool { return false } + renderDispatchMarkdown = func(*dispatchpkg.Dispatch) string { return "" } + + cmd := newDispatchCmd() + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if got := strings.Join(calls, ","); got != "wizard,picker,dispatch" { + t.Fatalf("call order = %q, want wizard,picker,dispatch", got) + } + if persistedProvider != string(agent.AgentNameGemini) { + t.Fatalf("persisted provider = %q, want %q", persistedProvider, agent.AgentNameGemini) + } +} + +func TestNewDispatchCmd_CloudDispatchDoesNotResolveLocalProvider(t *testing.T) { + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + oldTerminalMode := dispatchTerminalMode + unexpectedCallErr := errors.New("unexpected local provider resolution") + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + t.Fatal("normal cloud dispatch must not resolve a local provider") + return nil, unexpectedCallErr + } + runDispatch = func(_ context.Context, opts dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + if opts.Mode != dispatchpkg.ModeServer { + t.Fatalf("mode = %v, want server", opts.Mode) + } + return &dispatchpkg.Dispatch{}, nil + } + dispatchTerminalMode = func(io.Writer) bool { return false } + t.Cleanup(func() { + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + dispatchTerminalMode = oldTerminalMode + }) + + cmd := newDispatchCmd() + cmd.SetArgs([]string{"--repos", "entireio/cli"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } +} + +func TestNewDispatchCmd_ProviderErrorUsesStderrAndSkipsDispatch(t *testing.T) { + oldPrepare := prepareLocalDispatch + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + unexpectedCallErr := errors.New("unexpected dispatch call") + prepareLocalDispatch = func(_ context.Context, opts dispatchpkg.Options) (dispatchpkg.Options, error) { + return opts, nil + } + resolveDispatchProvider = func(_ context.Context, w io.Writer, override string) (*checkpointSummaryProvider, error) { + if override != string(agent.AgentNameCodex) { + t.Fatalf("provider override = %q, want codex", override) + } + if _, err := io.WriteString(w, "provider warning\n"); err != nil { + t.Fatal(err) + } + return nil, errors.New("provider failed") + } + runDispatch = func(context.Context, dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + t.Fatal("dispatch must not run after provider resolution fails") + return nil, unexpectedCallErr + } + t.Cleanup(func() { + prepareLocalDispatch = oldPrepare + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + }) + + cmd := newDispatchCmd() + cmd.SilenceErrors = true + cmd.SilenceUsage = true + var stdout bytes.Buffer + var stderr bytes.Buffer + cmd.SetOut(&stdout) + cmd.SetErr(&stderr) + cmd.SetArgs([]string{"--local", "--all-branches", "--agent", string(agent.AgentNameCodex)}) + + err := cmd.Execute() + if err == nil || err.Error() != "provider failed" { + t.Fatalf("unexpected error: %v", err) + } + if got := stdout.String(); got != "" { + t.Fatalf("unexpected stdout: %q", got) + } + if got := stderr.String(); got != "provider warning\n" { + t.Fatalf("unexpected stderr: %q", got) + } +} + func TestShouldRunDispatchWizard(t *testing.T) { t.Parallel() diff --git a/cmd/entire/cli/dispatch_tui_test.go b/cmd/entire/cli/dispatch_tui_test.go index b866d4103b..779a1526b7 100644 --- a/cmd/entire/cli/dispatch_tui_test.go +++ b/cmd/entire/cli/dispatch_tui_test.go @@ -23,6 +23,80 @@ func (p fakeDispatchProgram) Run() (tea.Model, error) { return model, nil } +type dispatchProgramFunc func() (tea.Model, error) + +func (f dispatchProgramFunc) Run() (tea.Model, error) { + return f() +} + +func TestDispatchTerminal_ResolvesProviderBeforeProgramRun(t *testing.T) { + oldPrepare := prepareLocalDispatch + oldProvider := resolveDispatchProvider + oldRunDispatch := runDispatch + oldTerminalMode := dispatchTerminalMode + oldInteractiveDispatch := runInteractiveDispatch + oldRenderTerminal := renderTerminalMarkdown + oldProgramFactory := newDispatchProgram + providerResolved := false + programRunning := false + generator := &stubTextAgent{} + prepareLocalDispatch = func(_ context.Context, opts dispatchpkg.Options) (dispatchpkg.Options, error) { + return opts, nil + } + resolveDispatchProvider = func(context.Context, io.Writer, string) (*checkpointSummaryProvider, error) { + if programRunning { + t.Fatal("provider resolution ran after Bubble Tea took terminal ownership") + } + providerResolved = true + return &checkpointSummaryProvider{TextGenerator: generator, Model: "terminal-model"}, nil + } + runDispatch = func(_ context.Context, opts dispatchpkg.Options) (*dispatchpkg.Dispatch, error) { + if !programRunning { + t.Fatal("interactive dispatch callback ran before program Run") + } + if opts.TextGenerator != generator || opts.Model != "terminal-model" { + t.Fatalf("provider options not passed to interactive dispatch: generator=%T model=%q", opts.TextGenerator, opts.Model) + } + return &dispatchpkg.Dispatch{GeneratedText: "# terminal dispatch\n"}, nil + } + dispatchTerminalMode = func(io.Writer) bool { return true } + runInteractiveDispatch = defaultRunInteractiveDispatch + renderTerminalMarkdown = func(_ io.Writer, markdown string) (string, error) { return markdown, nil } + newDispatchProgram = func(model tea.Model, _ io.Writer, _ bool) dispatchProgram { + if !providerResolved { + t.Fatal("Bubble Tea program was created before provider resolution completed") + } + return dispatchProgramFunc(func() (tea.Model, error) { + programRunning = true + status, ok := model.(dispatchStatusModel) + if !ok { + t.Fatalf("unexpected model type %T", model) + } + markdown, err := status.run(context.Background()) + status.result = dispatchRenderResult{markdown: markdown, err: err} + return status, nil + }) + } + t.Cleanup(func() { + prepareLocalDispatch = oldPrepare + resolveDispatchProvider = oldProvider + runDispatch = oldRunDispatch + dispatchTerminalMode = oldTerminalMode + runInteractiveDispatch = oldInteractiveDispatch + renderTerminalMarkdown = oldRenderTerminal + newDispatchProgram = oldProgramFactory + }) + + cmd := newDispatchCmd() + cmd.SetArgs([]string{"--local", "--all-branches"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if !providerResolved { + t.Fatal("provider was not resolved") + } +} + func TestDefaultRunInteractiveDispatch_DoesNotUseAltScreen(t *testing.T) { // Cannot run in parallel: mutates package-level newDispatchProgram, which // races with TestDefaultRunInteractiveDispatch_ClearsLoadingCardBeforeReturn. diff --git a/cmd/entire/cli/explain_summary_provider.go b/cmd/entire/cli/explain_summary_provider.go index 64cf0205e2..fd74bd6bfe 100644 --- a/cmd/entire/cli/explain_summary_provider.go +++ b/cmd/entire/cli/explain_summary_provider.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "io" + "strings" "github.com/entireio/cli/cmd/entire/cli/agent" "github.com/entireio/cli/cmd/entire/cli/agent/external" @@ -19,21 +20,25 @@ import ( ) var ( - loadSummarySettings = LoadEntireSettings - loadSummarySettingsFromFile = settings.LoadFromFile - saveLocalSummarySettings = SaveEntireSettingsLocal - getSummaryAgent = agent.Get - listRegisteredAgents = agent.List - isSummaryCLIAvailable = agent.IsSummaryCLIAvailable - discoverSummaryProviders = external.DiscoverAndRegister - discoverSummaryProvidersAlways = external.DiscoverAndRegisterAlways + loadSummarySettings = LoadEntireSettings + loadSummarySettingsFromFile = settings.LoadFromFile + saveLocalSummarySettings = SaveEntireSettingsLocal + getSummaryAgent = agent.Get + listRegisteredAgents = agent.List + isSummaryCLIAvailable = agent.IsSummaryCLIAvailable + discoverSummaryProviders = external.DiscoverAndRegister + discoverSummaryProvidersAlways = external.DiscoverAndRegisterAlways + discoverDispatchSummaryProvider = external.DiscoverAndRegisterNamedAlways + canPromptForSummaryProvider = interactive.CanPromptInteractively + promptSummaryProvider = promptForSummaryProvider ) type checkpointSummaryProvider struct { - Name types.AgentName - DisplayName string - Model string - Generator summarize.Generator + Name types.AgentName + DisplayName string + Model string + TextGenerator agent.TextGenerator + Generator summarize.Generator // Streaming reports whether the underlying text generator supports the // streaming path (the same predicate TextGeneratorAdapter dispatches on), // so the explain layer can attribute timeouts to the streaming diagnostic @@ -41,6 +46,24 @@ type checkpointSummaryProvider struct { Streaming bool } +func resolveDispatchSummaryProvider(ctx context.Context, w io.Writer, override string) (*checkpointSummaryProvider, error) { + override = strings.TrimSpace(override) + if override == "" { + return resolveCheckpointSummaryProvider(ctx, w) + } + + providerName := types.AgentName(override) + if _, err := getSummaryAgent(providerName); err != nil { + if err := discoverDispatchSummaryProvider(ctx, providerName); err != nil { + return nil, err + } + } + if err := validateSummaryProvider(override); err != nil { + return nil, err + } + return buildCheckpointSummaryProvider(providerName, "") +} + func resolveCheckpointSummaryProvider(ctx context.Context, w io.Writer) (*checkpointSummaryProvider, error) { s, err := loadSummarySettings(ctx) if err != nil { @@ -70,11 +93,11 @@ func resolveCheckpointSummaryProvider(ctx context.Context, w io.Writer) (*checkp case 1: return autoSelectSummaryProvider(ctx, w, candidates[0].Name, "non-interactive auto-select: single installed provider") default: - if !interactive.CanPromptInteractively() { + if !canPromptForSummaryProvider() { return autoSelectSummaryProvider(ctx, w, candidates[0].Name, "non-interactive auto-select: first detected of multiple") } - selected, err := promptForSummaryProvider(candidates) + selected, err := promptSummaryProvider(candidates) if err != nil { return nil, err } @@ -172,6 +195,10 @@ func promptForSummaryProvider(providers []checkpointSummaryProvider) (types.Agen } func buildCheckpointSummaryProvider(name types.AgentName, model string) (*checkpointSummaryProvider, error) { + return buildCheckpointSummaryProviderWithEffectiveModel(name, summarize.ResolveModel(name, model)) +} + +func buildCheckpointSummaryProviderWithEffectiveModel(name types.AgentName, effectiveModel string) (*checkpointSummaryProvider, error) { ag, err := getSummaryAgent(name) if err != nil { return nil, fmt.Errorf("loading summary provider %s: %w", name, err) @@ -182,15 +209,14 @@ func buildCheckpointSummaryProvider(name types.AgentName, model string) (*checkp return nil, fmt.Errorf("agent %s does not support summary generation", name) } - effectiveModel := summarize.ResolveModel(name, model) - _, streaming := agent.AsStreamingTextGenerator(textGenerator) return &checkpointSummaryProvider{ - Name: name, - DisplayName: string(ag.Type()), - Model: effectiveModel, - Streaming: streaming, + Name: name, + DisplayName: string(ag.Type()), + Model: effectiveModel, + TextGenerator: textGenerator, + Streaming: streaming, Generator: &summarize.TextGeneratorAdapter{ TextGenerator: textGenerator, Model: effectiveModel, @@ -227,7 +253,7 @@ func validateSummaryProvider(provider string) error { return fmt.Errorf("agent %q does not support summary generation", provider) } if !isSummaryProviderAvailable(name, ag) { - return fmt.Errorf("summary provider %q is configured but its CLI binary is not on PATH; install it or choose another provider", provider) + return fmt.Errorf("summary provider %q CLI binary is not on PATH; install it or choose another provider", provider) } return nil } diff --git a/cmd/entire/cli/explain_summary_provider_test.go b/cmd/entire/cli/explain_summary_provider_test.go index cf37380d8c..21523e5e73 100644 --- a/cmd/entire/cli/explain_summary_provider_test.go +++ b/cmd/entire/cli/explain_summary_provider_test.go @@ -3,6 +3,8 @@ package cli import ( "bytes" "context" + "errors" + "fmt" "os" "os/exec" "path/filepath" @@ -81,6 +83,25 @@ func (s *stubTextAgent) GenerateText(context.Context, string, string) (string, e return `{"intent":"Intent","outcome":"Outcome","learnings":{"repo":[],"code":[],"workflow":[]},"friction":[],"open_items":[]}`, nil } +type stubNonTextAgent struct { + agent.Agent +} + +func writeInfoSentinelExternalAgentBinary(t *testing.T, dir, name string) { + t.Helper() + + script := `#!/bin/sh +if [ "$1" = "info" ]; then + : > "$ENTIRE_TEST_UNRELATED_INFO_SENTINEL" + exit 1 +fi +echo '{}' +` + if err := os.WriteFile(filepath.Join(dir, "entire-agent-"+name), []byte(script), 0o755); err != nil { + t.Fatalf("write unrelated external agent binary: %v", err) + } +} + func TestResolveCheckpointSummaryProvider_UsesConfiguredProvider(t *testing.T) { // Cannot use t.Parallel() because we use t.Chdir and package-level var stubs ctx := context.Background() @@ -130,6 +151,467 @@ func TestResolveCheckpointSummaryProvider_UsesConfiguredProvider(t *testing.T) { if provider.Model != "haiku" { t.Fatalf("provider.Model = %q, want %q", provider.Model, "haiku") } + if provider.TextGenerator == nil { + t.Fatal("provider.TextGenerator = nil, want configured provider's raw text generator") + } +} + +func TestResolveDispatchSummaryProvider_ExplicitCodexUsesDefaultModelWithoutPersistence(t *testing.T) { + // Cannot use t.Parallel(): mutates package-level resolution seams. + ctx := context.Background() + codex := &stubTextAgent{name: agent.AgentNameCodex, kind: agent.AgentTypeCodex} + + originalLoad := loadSummarySettings + originalLoadFile := loadSummarySettingsFromFile + originalSave := saveLocalSummarySettings + originalGet := getSummaryAgent + originalCLI := isSummaryCLIAvailable + originalDiscover := discoverDispatchSummaryProvider + t.Cleanup(func() { + loadSummarySettings = originalLoad + loadSummarySettingsFromFile = originalLoadFile + saveLocalSummarySettings = originalSave + getSummaryAgent = originalGet + isSummaryCLIAvailable = originalCLI + discoverDispatchSummaryProvider = originalDiscover + }) + + loadSummarySettings = func(context.Context) (*settings.EntireSettings, error) { + t.Fatal("explicit dispatch provider must not load summary settings") + return nil, errors.New("unexpected settings load") + } + loadSummarySettingsFromFile = func(string) (*settings.EntireSettings, error) { + t.Fatal("explicit dispatch provider must not load settings for persistence") + return nil, errors.New("unexpected settings load for persistence") + } + saveLocalSummarySettings = func(context.Context, *settings.EntireSettings) error { + t.Fatal("explicit dispatch provider must not persist settings") + return nil + } + getSummaryAgent = func(name types.AgentName) (agent.Agent, error) { + if name != agent.AgentNameCodex { + t.Fatalf("getSummaryAgent(%q), want %q", name, agent.AgentNameCodex) + } + return codex, nil + } + isSummaryCLIAvailable = func(name types.AgentName) bool { + return name == agent.AgentNameCodex + } + discoverDispatchSummaryProvider = func(context.Context, types.AgentName) error { + t.Fatal("registered explicit provider should not trigger external discovery") + return nil + } + + provider, err := resolveDispatchSummaryProvider(ctx, &bytes.Buffer{}, " codex ") + if err != nil { + t.Fatalf("resolveDispatchSummaryProvider() error = %v", err) + } + if provider.Name != agent.AgentNameCodex { + t.Fatalf("provider.Name = %q, want %q", provider.Name, agent.AgentNameCodex) + } + if provider.Model != "" { + t.Fatalf("provider.Model = %q, want provider CLI default", provider.Model) + } + if provider.TextGenerator != codex { + t.Fatalf("provider.TextGenerator = %T %p, want raw generator %T %p", provider.TextGenerator, provider.TextGenerator, codex, codex) + } +} + +func TestResolveDispatchSummaryProvider_EmptyOverrideUsesConfiguredProviderAndModel(t *testing.T) { + // Cannot use t.Parallel(): mutates package-level resolution seams. + ctx := context.Background() + configured := &stubTextAgent{name: agent.AgentNameGemini, kind: agent.AgentTypeGemini} + + originalLoad := loadSummarySettings + originalGet := getSummaryAgent + originalCLI := isSummaryCLIAvailable + originalDiscover := discoverSummaryProvidersAlways + t.Cleanup(func() { + loadSummarySettings = originalLoad + getSummaryAgent = originalGet + isSummaryCLIAvailable = originalCLI + discoverSummaryProvidersAlways = originalDiscover + }) + + loadSummarySettings = func(context.Context) (*settings.EntireSettings, error) { + return &settings.EntireSettings{SummaryGeneration: &settings.SummaryGenerationSettings{ + Provider: string(agent.AgentNameGemini), + Model: "gemini-saved-model", + }}, nil + } + getSummaryAgent = func(name types.AgentName) (agent.Agent, error) { + if name != agent.AgentNameGemini { + t.Fatalf("getSummaryAgent(%q), want %q", name, agent.AgentNameGemini) + } + return configured, nil + } + isSummaryCLIAvailable = func(name types.AgentName) bool { + return name == agent.AgentNameGemini + } + discoverSummaryProvidersAlways = func(context.Context) { + t.Fatal("configured registered provider should not trigger external discovery") + } + + provider, err := resolveDispatchSummaryProvider(ctx, &bytes.Buffer{}, " \t\n") + if err != nil { + t.Fatalf("resolveDispatchSummaryProvider() error = %v", err) + } + if provider.Name != agent.AgentNameGemini { + t.Fatalf("provider.Name = %q, want %q", provider.Name, agent.AgentNameGemini) + } + if provider.Model != "gemini-saved-model" { + t.Fatalf("provider.Model = %q, want configured model", provider.Model) + } + if provider.TextGenerator != configured { + t.Fatalf("provider.TextGenerator = %T, want configured raw generator", provider.TextGenerator) + } +} + +func TestResolveDispatchSummaryProvider_ExplicitProviderIgnoresSavedProviderAndModel(t *testing.T) { + // Cannot use t.Parallel(): mutates package-level resolution seams. + ctx := context.Background() + codex := &stubTextAgent{name: agent.AgentNameCodex, kind: agent.AgentTypeCodex} + loadCalls := 0 + + originalLoad := loadSummarySettings + originalGet := getSummaryAgent + originalCLI := isSummaryCLIAvailable + t.Cleanup(func() { + loadSummarySettings = originalLoad + getSummaryAgent = originalGet + isSummaryCLIAvailable = originalCLI + }) + + loadSummarySettings = func(context.Context) (*settings.EntireSettings, error) { + loadCalls++ + return &settings.EntireSettings{SummaryGeneration: &settings.SummaryGenerationSettings{ + Provider: string(agent.AgentNameClaudeCode), + Model: "sonnet", + }}, nil + } + getSummaryAgent = func(types.AgentName) (agent.Agent, error) { return codex, nil } + isSummaryCLIAvailable = func(types.AgentName) bool { return true } + + provider, err := resolveDispatchSummaryProvider(ctx, &bytes.Buffer{}, string(agent.AgentNameCodex)) + if err != nil { + t.Fatalf("resolveDispatchSummaryProvider() error = %v", err) + } + if loadCalls != 0 { + t.Fatalf("loadSummarySettings calls = %d, want 0 for explicit override", loadCalls) + } + if provider.Name != agent.AgentNameCodex || provider.Model != "" { + t.Fatalf("provider = %+v, want explicit Codex with provider-default model", provider) + } +} + +func TestResolveDispatchSummaryProvider_ExplicitClaudeUsesSummaryDefaultModel(t *testing.T) { + // Cannot use t.Parallel(): mutates package-level resolution seams. + ctx := context.Background() + claude := &stubTextAgent{name: agent.AgentNameClaudeCode, kind: agent.AgentTypeClaudeCode} + loadCalls := 0 + + originalLoad := loadSummarySettings + originalGet := getSummaryAgent + originalCLI := isSummaryCLIAvailable + t.Cleanup(func() { + loadSummarySettings = originalLoad + getSummaryAgent = originalGet + isSummaryCLIAvailable = originalCLI + }) + + loadSummarySettings = func(context.Context) (*settings.EntireSettings, error) { + loadCalls++ + return &settings.EntireSettings{SummaryGeneration: &settings.SummaryGenerationSettings{ + Provider: string(agent.AgentNameClaudeCode), + Model: "opus", + }}, nil + } + getSummaryAgent = func(types.AgentName) (agent.Agent, error) { return claude, nil } + isSummaryCLIAvailable = func(types.AgentName) bool { return true } + + provider, err := resolveDispatchSummaryProvider(ctx, &bytes.Buffer{}, string(agent.AgentNameClaudeCode)) + if err != nil { + t.Fatalf("resolveDispatchSummaryProvider() error = %v", err) + } + if loadCalls != 0 { + t.Fatalf("loadSummarySettings calls = %d, want 0 for explicit override", loadCalls) + } + if provider.Model != summarize.DefaultModel { + t.Fatalf("provider.Model = %q, want summary default %q", provider.Model, summarize.DefaultModel) + } +} + +func TestResolveDispatchSummaryProvider_PropagatesDiscoveryDeadline(t *testing.T) { + // Cannot use t.Parallel(): mutates package-level resolution seams. + providerName := types.AgentName("external-discovery-deadline") + + originalGet := getSummaryAgent + originalDiscover := discoverDispatchSummaryProvider + t.Cleanup(func() { + getSummaryAgent = originalGet + discoverDispatchSummaryProvider = originalDiscover + }) + + getSummaryAgent = func(types.AgentName) (agent.Agent, error) { + return nil, errors.New("not registered") + } + discoverDispatchSummaryProvider = func(context.Context, types.AgentName) error { + return fmt.Errorf("discovering external agent %q: %w", providerName, context.DeadlineExceeded) + } + + _, err := resolveDispatchSummaryProvider(context.Background(), &bytes.Buffer{}, string(providerName)) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("resolveDispatchSummaryProvider() error = %v, want context deadline exceeded", err) + } + if strings.Contains(err.Error(), "unknown summary provider") { + t.Fatalf("resolveDispatchSummaryProvider() error = %q, do not want unknown-provider rewrite", err) + } +} + +func TestResolveDispatchSummaryProvider_PropagatesDiscoveryCancellation(t *testing.T) { + // Cannot use t.Parallel(): mutates package-level resolution seams. + providerName := types.AgentName("external-discovery-canceled") + + originalGet := getSummaryAgent + originalDiscover := discoverDispatchSummaryProvider + t.Cleanup(func() { + getSummaryAgent = originalGet + discoverDispatchSummaryProvider = originalDiscover + }) + + getSummaryAgent = func(types.AgentName) (agent.Agent, error) { + return nil, errors.New("not registered") + } + discoverDispatchSummaryProvider = func(context.Context, types.AgentName) error { + return fmt.Errorf("discovering external agent %q: %w", providerName, context.Canceled) + } + + _, err := resolveDispatchSummaryProvider(context.Background(), &bytes.Buffer{}, string(providerName)) + if !errors.Is(err, context.Canceled) { + t.Fatalf("resolveDispatchSummaryProvider() error = %v, want context canceled", err) + } + if strings.Contains(err.Error(), "unknown summary provider") { + t.Fatalf("resolveDispatchSummaryProvider() error = %q, do not want unknown-provider rewrite", err) + } +} + +func TestResolveDispatchSummaryProvider_PropagatesInvalidExternalInfo(t *testing.T) { + // Cannot use t.Parallel(): mutates package-level resolution seams. + providerName := types.AgentName("external-discovery-invalid-info") + infoErr := errors.New("invalid helper info") + + originalGet := getSummaryAgent + originalDiscover := discoverDispatchSummaryProvider + t.Cleanup(func() { + getSummaryAgent = originalGet + discoverDispatchSummaryProvider = originalDiscover + }) + + getSummaryAgent = func(types.AgentName) (agent.Agent, error) { + return nil, errors.New("not registered") + } + discoverDispatchSummaryProvider = func(context.Context, types.AgentName) error { + return fmt.Errorf("loading info for external agent %q: info: invalid JSON: %w", providerName, infoErr) + } + + _, err := resolveDispatchSummaryProvider(context.Background(), &bytes.Buffer{}, string(providerName)) + if !errors.Is(err, infoErr) { + t.Fatalf("resolveDispatchSummaryProvider() error = %v, want invalid-info cause", err) + } + if !strings.Contains(err.Error(), string(providerName)) || !strings.Contains(err.Error(), "info: invalid JSON") { + t.Fatalf("resolveDispatchSummaryProvider() error = %q, want provider and invalid-info context", err) + } + if strings.Contains(err.Error(), "unknown summary provider") { + t.Fatalf("resolveDispatchSummaryProvider() error = %q, do not want unknown-provider rewrite", err) + } +} + +func TestResolveDispatchSummaryProvider_MissingExternalKeepsUnknownProviderError(t *testing.T) { + // Cannot use t.Parallel(): mutates package-level resolution seams. + providerName := types.AgentName("external-discovery-missing") + + originalGet := getSummaryAgent + originalDiscover := discoverDispatchSummaryProvider + t.Cleanup(func() { + getSummaryAgent = originalGet + discoverDispatchSummaryProvider = originalDiscover + }) + + getSummaryAgent = func(types.AgentName) (agent.Agent, error) { + return nil, errors.New("not registered") + } + discoverDispatchSummaryProvider = func(context.Context, types.AgentName) error { return nil } + + _, err := resolveDispatchSummaryProvider(context.Background(), &bytes.Buffer{}, string(providerName)) + if err == nil || !strings.Contains(err.Error(), "unknown summary provider") { + t.Fatalf("resolveDispatchSummaryProvider() error = %v, want existing unknown-provider error", err) + } +} + +func TestResolveDispatchSummaryProvider_ExplicitValidationErrors(t *testing.T) { + // Cannot use t.Parallel(): subtests mutate package-level resolution seams. + tests := []struct { + name string + override string + agent agent.Agent + getErr error + available bool + wantError string + unwantedError string + }{ + { + name: "unknown provider", + override: "missing-provider", + getErr: errors.New("not registered"), + available: true, + wantError: "unknown summary provider", + }, + { + name: "no text generator capability", + override: "no-text", + agent: &stubNonTextAgent{Agent: &stubTextAgent{ + name: "no-text", + kind: agent.AgentTypeUnknown, + }}, + available: true, + wantError: "does not support summary generation", + }, + { + name: "CLI unavailable", + override: string(agent.AgentNameCodex), + agent: &stubTextAgent{name: agent.AgentNameCodex, kind: agent.AgentTypeCodex}, + available: false, + wantError: "install it or choose another provider", + unwantedError: "configured", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + originalGet := getSummaryAgent + originalCLI := isSummaryCLIAvailable + originalDiscover := discoverDispatchSummaryProvider + t.Cleanup(func() { + getSummaryAgent = originalGet + isSummaryCLIAvailable = originalCLI + discoverDispatchSummaryProvider = originalDiscover + }) + + getSummaryAgent = func(types.AgentName) (agent.Agent, error) { + if tt.getErr != nil { + return nil, tt.getErr + } + return tt.agent, nil + } + isSummaryCLIAvailable = func(types.AgentName) bool { return tt.available } + discoverDispatchSummaryProvider = func(context.Context, types.AgentName) error { return nil } + + _, err := resolveDispatchSummaryProvider(context.Background(), &bytes.Buffer{}, tt.override) + if err == nil { + t.Fatalf("resolveDispatchSummaryProvider(%q) error = nil, want %q", tt.override, tt.wantError) + } + if !strings.Contains(err.Error(), tt.wantError) { + t.Fatalf("resolveDispatchSummaryProvider(%q) error = %q, want substring %q", tt.override, err, tt.wantError) + } + if tt.unwantedError != "" && strings.Contains(err.Error(), tt.unwantedError) { + t.Fatalf("resolveDispatchSummaryProvider(%q) error = %q, do not want substring %q", tt.override, err, tt.unwantedError) + } + }) + } +} + +func TestResolveDispatchSummaryProvider_ExplicitExternalProviderDoesNotWriteLocalSettings(t *testing.T) { + // Cannot use t.Parallel(): subtests mutate cwd, PATH, and the agent registry. + if _, err := exec.LookPath("sh"); err != nil { + t.Skip("sh not available") + } + + tests := []struct { + name string + providerName string + localContent string + }{ + {name: "does not create settings.local.json", providerName: "external-dispatch-no-create"}, + { + name: "does not update settings.local.json", + providerName: "external-dispatch-no-update", + localContent: `{"external_agents":false,"summary_generation":{"provider":"codex","model":"saved-model"}}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + testutil.InitRepo(t, tmpDir) + t.Chdir(tmpDir) + + if err := os.MkdirAll(filepath.Join(tmpDir, ".entire"), 0o755); err != nil { + t.Fatalf("mkdir .entire: %v", err) + } + if err := os.WriteFile(filepath.Join(tmpDir, ".entire", "settings.json"), []byte(`{"enabled":true,"external_agents":false}`), 0o644); err != nil { + t.Fatalf("write settings.json: %v", err) + } + + localPath := filepath.Join(tmpDir, ".entire", "settings.local.json") + if tt.localContent != "" { + if err := os.WriteFile(localPath, []byte(tt.localContent), 0o644); err != nil { + t.Fatalf("write settings.local.json: %v", err) + } + } + + externalDir := t.TempDir() + writeExternalSummaryAgentBinary(t, externalDir, tt.providerName) + writeInfoSentinelExternalAgentBinary(t, externalDir, tt.providerName+"-unrelated") + t.Setenv("PATH", externalDir+string(os.PathListSeparator)+os.Getenv("PATH")) + unrelatedInfoSentinel := filepath.Join(t.TempDir(), "unrelated-info-called") + t.Setenv("ENTIRE_TEST_UNRELATED_INFO_SENTINEL", unrelatedInfoSentinel) + modelRecord := filepath.Join(t.TempDir(), "model-args") + t.Setenv("ENTIRE_TEST_EXTERNAL_MODEL_RECORD", modelRecord) + + provider, err := resolveDispatchSummaryProvider(ctx, &bytes.Buffer{}, tt.providerName) + if err != nil { + t.Fatalf("resolveDispatchSummaryProvider() error = %v", err) + } + if provider.Name != types.AgentName(tt.providerName) { + t.Fatalf("provider.Name = %q, want %q", provider.Name, tt.providerName) + } + if provider.Model != "" { + t.Fatalf("provider.Model = %q, want external CLI default", provider.Model) + } + if provider.TextGenerator == nil { + t.Fatal("provider.TextGenerator = nil, want external raw generator") + } + if _, err := os.Stat(unrelatedInfoSentinel); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("unrelated plugin info was invoked: stat error = %v", err) + } + + generated, err := provider.TextGenerator.GenerateText(ctx, "generate a summary", provider.Model) + if err != nil { + t.Fatalf("provider.TextGenerator.GenerateText() error = %v", err) + } + if !strings.Contains(generated, `"intent":"Intent"`) { + t.Fatalf("provider.TextGenerator.GenerateText() = %q, want generated summary", generated) + } + modelArgs, err := os.ReadFile(modelRecord) + if err != nil { + t.Fatalf("read external model args: %v", err) + } + if string(modelArgs) != "--model\n\n" { + t.Fatalf("external generate-text args = %q, want empty model argument", modelArgs) + } + + got, err := os.ReadFile(localPath) + switch { + case tt.localContent == "" && !errors.Is(err, os.ErrNotExist): + t.Fatalf("settings.local.json read error = %v, want file to remain absent (content %q)", err, got) + case tt.localContent != "" && err != nil: + t.Fatalf("read settings.local.json: %v", err) + case tt.localContent != "" && string(got) != tt.localContent: + t.Fatalf("settings.local.json changed:\n got: %s\nwant: %s", got, tt.localContent) + } + }) + } } func TestResolveCheckpointSummaryProvider_SavesSingleInstalledProvider(t *testing.T) { diff --git a/cmd/entire/cli/setup_test.go b/cmd/entire/cli/setup_test.go index a4cefd0f12..d3f821ae25 100644 --- a/cmd/entire/cli/setup_test.go +++ b/cmd/entire/cli/setup_test.go @@ -177,6 +177,9 @@ case "$1" in echo '{"present": true}' ;; generate-text) + if [ -n "$ENTIRE_TEST_EXTERNAL_MODEL_RECORD" ]; then + printf '%s\n%s\n' "$2" "$3" > "$ENTIRE_TEST_EXTERNAL_MODEL_RECORD" + fi echo '{"text":"{\"intent\":\"Intent\",\"outcome\":\"Outcome\",\"learnings\":{\"repo\":[],\"code\":[],\"workflow\":[]},\"friction\":[],\"open_items\":[]}"}' ;; *)