From 49d022bffaf2fac5f5f3880fdfecb6f46082aa4d Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 14:48:50 +0200 Subject: [PATCH 01/22] refactor(plugins): separate provider discovery from clients --- docs/plugins.md | 2 + .../adapters/out/plugin/provider_loader.go | 121 +++++++++++----- .../out/plugin/provider_loader_test.go | 71 ++++++++++ internal/app/app.go | 2 +- internal/app/app_test.go | 30 +++- internal/app/review_providers.go | 22 ++- .../mocks/review_provider_catalog_mock.go | 101 ++++++++++++++ .../review_provider_client_factory_mock.go | 107 ++++++++++++++ .../mocks/review_provider_loader_mock.go | 130 ++++++++++++++++++ internal/ports/plugin.go | 24 ++++ 10 files changed, 571 insertions(+), 39 deletions(-) create mode 100644 internal/ports/mocks/review_provider_catalog_mock.go create mode 100644 internal/ports/mocks/review_provider_client_factory_mock.go diff --git a/docs/plugins.md b/docs/plugins.md index 962a07f..f93c93b 100644 --- a/docs/plugins.md +++ b/docs/plugins.md @@ -48,6 +48,8 @@ label = "Example" Required fields are `name`, `version`, `manifest_version = "1"`, `protocol = "ero.plugin.v1"`, `runtime.command`, and at least one contribution with `type` and `id`. Contribution type strings are lower snake_case; the currently implemented public contribution type is `review_provider`. +Ero discovers available review providers from installed plugin manifests before starting plugin subprocesses. Each discovered provider has a host-owned stable key derived from a canonical installed-plugin identity plus the contribution `id`; runtime provider IDs returned by `initialize` remain provider-owned metadata and are not used as the host selection key. + `runtime.command` is executed with the plugin root as the working directory. Keep it stable for installed users; use the optional `build.command` for local development or release packaging. ## Protocol diff --git a/internal/adapters/out/plugin/provider_loader.go b/internal/adapters/out/plugin/provider_loader.go index d2ab7cd..c5e7ab0 100644 --- a/internal/adapters/out/plugin/provider_loader.go +++ b/internal/adapters/out/plugin/provider_loader.go @@ -2,6 +2,8 @@ package pluginadapter import ( "context" + "crypto/sha256" + "encoding/hex" "fmt" "os" "os/exec" @@ -18,57 +20,114 @@ import ( // ReviewProviderLoader builds review provider clients from installed plugin manifests. type ReviewProviderLoader struct { - registry ports.PluginRegistry - timeout time.Duration + registry ports.PluginRegistry + timeout time.Duration + clientFactory func(context.Context, ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) } // NewReviewProviderLoader creates a loader backed by an installed plugin registry. func NewReviewProviderLoader(registry ports.PluginRegistry) *ReviewProviderLoader { - return &ReviewProviderLoader{registry: registry, timeout: DefaultPluginTimeout} + loader := &ReviewProviderLoader{registry: registry, timeout: DefaultPluginTimeout} + loader.clientFactory = loader.createReviewProviderClient + return loader } -// LoadReviewProviders implements ports.ReviewProviderLoader. -func (l *ReviewProviderLoader) LoadReviewProviders(ctx context.Context) ([]ports.ReviewProviderClient, error) { +// ListReviewProviderDescriptors implements ports.ReviewProviderCatalog. +func (l *ReviewProviderLoader) ListReviewProviderDescriptors(ctx context.Context) ([]ports.ReviewProviderDescriptor, error) { descriptors, err := l.registry.InstalledPlugins(ctx) if err != nil { return nil, err } + providers := make([]ports.ReviewProviderDescriptor, 0) + for _, descriptor := range descriptors { + for _, contribution := range descriptor.Contributions { + if contribution.Type != pluginsdk.ContributionReviewProvider { + continue + } + providers = append(providers, ports.ReviewProviderDescriptor{ + Key: stableReviewProviderKey(descriptor, contribution), + PluginName: descriptor.Name, + PluginVersion: descriptor.Version, + PluginSource: descriptor.Source, + PluginPath: descriptor.Path, + ContributionID: contribution.ID, + Label: contribution.Label, + Type: contribution.Type, + }) + } + } + return providers, nil +} + +// CreateReviewProviderClient implements ports.ReviewProviderClientFactory. +func (l *ReviewProviderLoader) CreateReviewProviderClient(ctx context.Context, descriptor ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { + return l.clientFactory(ctx, descriptor) +} + +// LoadReviewProviders implements ports.ReviewProviderLoader as a temporary compatibility shim. +func (l *ReviewProviderLoader) LoadReviewProviders(ctx context.Context) ([]ports.ReviewProviderClient, error) { + descriptors, err := l.ListReviewProviderDescriptors(ctx) + if err != nil { + return nil, err + } log := zerowrap.FromCtx(ctx) - providers := make([]ports.ReviewProviderClient, 0) + providers := make([]ports.ReviewProviderClient, 0, len(descriptors)) for _, descriptor := range descriptors { - manifest, err := LoadManifest(descriptor.Path) + client, err := l.CreateReviewProviderClient(ctx, descriptor) if err != nil { - log.Warn().Err(err).Str("plugin_path", descriptor.Path).Msg("load plugin manifest failed") + log.Warn().Err(err).Str("plugin_path", descriptor.PluginPath).Str("contribution_id", descriptor.ContributionID).Msg("create plugin review provider client failed") continue } - command, args := splitRuntimeCommand(manifest.Runtime.Command) - if command == "" { - log.Warn().Str("plugin_path", descriptor.Path).Msg("plugin runtime command is empty") - continue + providers = append(providers, client) + } + return providers, nil +} + +func (l *ReviewProviderLoader) createReviewProviderClient(ctx context.Context, descriptor ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { + manifest, err := LoadManifest(descriptor.PluginPath) + if err != nil { + return nil, err + } + command, args := splitRuntimeCommand(manifest.Runtime.Command) + if command == "" { + return nil, fmt.Errorf("plugin runtime command is empty") + } + if !runtimeCommandAvailable(command, descriptor.PluginPath) && strings.TrimSpace(manifest.Build.Command) != "" { + if err := runPluginBuildCommand(ctx, descriptor.PluginPath, manifest.Build.Command, l.timeout); err != nil { + log := zerowrap.FromCtx(ctx) + log.Warn().Err(err).Str("plugin_path", descriptor.PluginPath).Msg("build plugin runtime failed") } - if !runtimeCommandAvailable(command, descriptor.Path) && strings.TrimSpace(manifest.Build.Command) != "" { - if err := runPluginBuildCommand(ctx, descriptor.Path, manifest.Build.Command, l.timeout); err != nil { - log.Warn().Err(err).Str("plugin_path", descriptor.Path).Msg("build plugin runtime failed") - } + } + if !strings.Contains(command, "/") { + if resolved, err := exec.LookPath(command); err == nil { + command = resolved } - if !strings.Contains(command, "/") { - if resolved, err := exec.LookPath(command); err == nil { - command = resolved - } + } + return NewClientForContribution(command, args, descriptor.PluginPath, descriptor.ContributionID, l.timeout) +} + +func stableReviewProviderKey(descriptor ports.PluginDescriptor, contribution ports.PluginContribution) string { + identity := canonicalInstalledPluginIdentity(descriptor) + digest := sha256.Sum256([]byte(identity)) + return "plugin:" + hex.EncodeToString(digest[:8]) + "#review_provider:" + contribution.ID +} + +func canonicalInstalledPluginIdentity(descriptor ports.PluginDescriptor) string { + if source, err := ParseSource(descriptor.Source); err == nil { + switch source.Type { + case SourceTypeGit: + return strings.Join([]string{"git", strings.ToLower(source.Host), strings.TrimSuffix(source.Path, ".git"), source.Ref}, ":") + case SourceTypeLocal: + return "local:" + filepath.Clean(source.LocalPath) } - for _, contribution := range descriptor.Contributions { - if contribution.Type != pluginsdk.ContributionReviewProvider { - continue - } - client, err := NewClientForContribution(command, args, descriptor.Path, contribution.ID, l.timeout) - if err != nil { - log.Warn().Err(err).Str("plugin_path", descriptor.Path).Str("contribution_id", contribution.ID).Msg("create plugin review provider client failed") - continue - } - providers = append(providers, client) + } + if descriptor.Path != "" { + if abs, err := filepath.Abs(descriptor.Path); err == nil { + return "path:" + filepath.Clean(abs) } + return "path:" + filepath.Clean(descriptor.Path) } - return providers, nil + return strings.Join([]string{"manifest", descriptor.Name, descriptor.Version, descriptor.Source}, ":") } func runtimeCommandAvailable(command, pluginDir string) bool { diff --git a/internal/adapters/out/plugin/provider_loader_test.go b/internal/adapters/out/plugin/provider_loader_test.go index 8331744..96084b5 100644 --- a/internal/adapters/out/plugin/provider_loader_test.go +++ b/internal/adapters/out/plugin/provider_loader_test.go @@ -12,6 +12,77 @@ import ( "ero/internal/ports/mocks" ) +func TestReviewProviderLoaderListsReviewProviderDescriptorsWithoutStartingClients(t *testing.T) { + t.Parallel() + + clientFactoryCalls := 0 + multiDescriptor := ports.PluginDescriptor{ + Name: "multi", + Version: "0.1.0", + Source: "git:example.com/owner/multi@v0.1.0", + Path: "plugins/multi", + Contributions: []ports.PluginContribution{ + {Type: "review_provider", ID: "github", Label: "GitHub"}, + {Type: "theme", ID: "dark", Label: "Dark"}, + {Type: "review_provider", ID: "gitlab", Label: "GitLab"}, + }, + } + registry := mocks.NewMockPluginRegistry(t) + registry.EXPECT().InstalledPlugins(context.Background()).Return([]ports.PluginDescriptor{multiDescriptor, { + Name: "other", + Version: "2.0.0", + Source: "git:example.com/owner/other@v2.0.0", + Path: "plugins/other", + Contributions: []ports.PluginContribution{ + {Type: "workflow", ID: "triage", Label: "Triage"}, + }, + }}, nil) + + loader := NewReviewProviderLoader(registry) + loader.clientFactory = func(context.Context, ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { + clientFactoryCalls++ + return nil, nil + } + + descriptors, err := loader.ListReviewProviderDescriptors(context.Background()) + require.NoError(t, err) + require.Equal(t, []ports.ReviewProviderDescriptor{{ + Key: stableReviewProviderKey(multiDescriptor, multiDescriptor.Contributions[0]), + PluginName: "multi", + PluginVersion: "0.1.0", + PluginSource: "git:example.com/owner/multi@v0.1.0", + PluginPath: "plugins/multi", + ContributionID: "github", + Label: "GitHub", + Type: "review_provider", + }, { + Key: stableReviewProviderKey(multiDescriptor, multiDescriptor.Contributions[2]), + PluginName: "multi", + PluginVersion: "0.1.0", + PluginSource: "git:example.com/owner/multi@v0.1.0", + PluginPath: "plugins/multi", + ContributionID: "gitlab", + Label: "GitLab", + Type: "review_provider", + }}, descriptors) + require.Zero(t, clientFactoryCalls) +} + +func TestStableReviewProviderKeyUsesCanonicalInstalledPluginIdentity(t *testing.T) { + t.Parallel() + + first := ports.PluginDescriptor{Name: "same", Version: "1.0.0", Source: "git:example.com/owner/one@v1"} + second := ports.PluginDescriptor{Name: "same", Version: "1.0.0", Source: "git:example.com/owner/two@v1"} + contribution := ports.PluginContribution{Type: "review_provider", ID: "github", Label: "GitHub"} + + firstKey := stableReviewProviderKey(first, contribution) + secondKey := stableReviewProviderKey(second, contribution) + + require.NotEqual(t, firstKey, secondKey) + require.NotContains(t, firstKey, "same@1.0.0") + require.Contains(t, firstKey, "#review_provider:github") +} + func TestReviewProviderLoaderBuildsMissingRuntimeBeforeStartingProvider(t *testing.T) { t.Parallel() diff --git a/internal/app/app.go b/internal/app/app.go index 598e96a..52d084a 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -104,7 +104,7 @@ func newAppWithClipboard(cfg *viper.Viper, loader reviewLoader, runner tuiRunner var reviewProviders []ports.ReviewProviderClient pluginManager := pluginadapter.NewManager() providerLoader := pluginadapter.NewReviewProviderLoader(pluginManager) - providers, err := buildReviewProviders(ctx, providerLoader) + providers, err := buildReviewProviders(ctx, providerLoader, providerLoader) if err != nil { log.Warn().Err(err).Msg("load review providers failed") } else { diff --git a/internal/app/app_test.go b/internal/app/app_test.go index 71d7c54..4f8d21d 100644 --- a/internal/app/app_test.go +++ b/internal/app/app_test.go @@ -246,15 +246,16 @@ func TestRunLoadsReviewAndRunsTUIWithConfig(t *testing.T) { } } -func TestBuildReviewProvidersDelegatesToLoader(t *testing.T) { +func TestBuildReviewProvidersUsesDescriptorsAndFactory(t *testing.T) { t.Parallel() ctx := context.Background() + descriptor := ports.ReviewProviderDescriptor{Key: "plugin@1.0.0/github", ContributionID: "github"} provider := mocks.NewMockReviewProviderClient(t) - loader := mocks.NewMockReviewProviderLoader(t) - loader.EXPECT().LoadReviewProviders(ctx).Return([]ports.ReviewProviderClient{provider}, nil) + catalog := &fakeReviewProviderCatalog{descriptors: []ports.ReviewProviderDescriptor{descriptor}} + factory := &fakeReviewProviderClientFactory{clients: map[string]ports.ReviewProviderClient{descriptor.Key: provider}} - providers, err := buildReviewProviders(ctx, loader) + providers, err := buildReviewProviders(ctx, catalog, factory) require.NoError(t, err) require.Equal(t, []ports.ReviewProviderClient{provider}, providers) } @@ -338,6 +339,27 @@ func (f *fakeGitMetadataReader) ResolveRevision(_ string, revision string) (stri } func (f *fakeGitMetadataReader) DefaultBranch(string) (string, error) { return f.defaultBranch, f.err } +type fakeReviewProviderCatalog struct { + descriptors []ports.ReviewProviderDescriptor + err error +} + +func (f *fakeReviewProviderCatalog) ListReviewProviderDescriptors(context.Context) ([]ports.ReviewProviderDescriptor, error) { + return f.descriptors, f.err +} + +type fakeReviewProviderClientFactory struct { + clients map[string]ports.ReviewProviderClient + err error +} + +func (f *fakeReviewProviderClientFactory) CreateReviewProviderClient(_ context.Context, descriptor ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { + if f.err != nil { + return nil, f.err + } + return f.clients[descriptor.Key], nil +} + type fakeStartupPrompt struct { mode core.DiffMode err error diff --git a/internal/app/review_providers.go b/internal/app/review_providers.go index f9f85b3..2871aa7 100644 --- a/internal/app/review_providers.go +++ b/internal/app/review_providers.go @@ -3,12 +3,28 @@ package app import ( "context" + "github.com/bnema/zerowrap" + "ero/internal/ports" ) -func buildReviewProviders(ctx context.Context, loader ports.ReviewProviderLoader) ([]ports.ReviewProviderClient, error) { - if loader == nil { +func buildReviewProviders(ctx context.Context, catalog ports.ReviewProviderCatalog, factory ports.ReviewProviderClientFactory) ([]ports.ReviewProviderClient, error) { + if catalog == nil || factory == nil { return nil, nil } - return loader.LoadReviewProviders(ctx) + descriptors, err := catalog.ListReviewProviderDescriptors(ctx) + if err != nil { + return nil, err + } + log := zerowrap.FromCtx(ctx) + providers := make([]ports.ReviewProviderClient, 0, len(descriptors)) + for _, descriptor := range descriptors { + provider, err := factory.CreateReviewProviderClient(ctx, descriptor) + if err != nil { + log.Warn().Err(err).Str("provider_key", descriptor.Key).Str("contribution_id", descriptor.ContributionID).Msg("create plugin review provider client failed") + continue + } + providers = append(providers, provider) + } + return providers, nil } diff --git a/internal/ports/mocks/review_provider_catalog_mock.go b/internal/ports/mocks/review_provider_catalog_mock.go new file mode 100644 index 0000000..c0e94a5 --- /dev/null +++ b/internal/ports/mocks/review_provider_catalog_mock.go @@ -0,0 +1,101 @@ +// Code generated by mockery; DO NOT EDIT. +// github.com/vektra/mockery +// template: testify + +package mocks + +import ( + "context" + "ero/internal/ports" + + mock "github.com/stretchr/testify/mock" +) + +// NewMockReviewProviderCatalog creates a new instance of MockReviewProviderCatalog. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockReviewProviderCatalog(t interface { + mock.TestingT + Cleanup(func()) +}) *MockReviewProviderCatalog { + mock := &MockReviewProviderCatalog{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} + +// MockReviewProviderCatalog is an autogenerated mock type for the ReviewProviderCatalog type +type MockReviewProviderCatalog struct { + mock.Mock +} + +type MockReviewProviderCatalog_Expecter struct { + mock *mock.Mock +} + +func (_m *MockReviewProviderCatalog) EXPECT() *MockReviewProviderCatalog_Expecter { + return &MockReviewProviderCatalog_Expecter{mock: &_m.Mock} +} + +// ListReviewProviderDescriptors provides a mock function for the type MockReviewProviderCatalog +func (_mock *MockReviewProviderCatalog) ListReviewProviderDescriptors(ctx context.Context) ([]ports.ReviewProviderDescriptor, error) { + ret := _mock.Called(ctx) + + if len(ret) == 0 { + panic("no return value specified for ListReviewProviderDescriptors") + } + + var r0 []ports.ReviewProviderDescriptor + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context) ([]ports.ReviewProviderDescriptor, error)); ok { + return returnFunc(ctx) + } + if returnFunc, ok := ret.Get(0).(func(context.Context) []ports.ReviewProviderDescriptor); ok { + r0 = returnFunc(ctx) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]ports.ReviewProviderDescriptor) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context) error); ok { + r1 = returnFunc(ctx) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockReviewProviderCatalog_ListReviewProviderDescriptors_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListReviewProviderDescriptors' +type MockReviewProviderCatalog_ListReviewProviderDescriptors_Call struct { + *mock.Call +} + +// ListReviewProviderDescriptors is a helper method to define mock.On call +// - ctx context.Context +func (_e *MockReviewProviderCatalog_Expecter) ListReviewProviderDescriptors(ctx interface{}) *MockReviewProviderCatalog_ListReviewProviderDescriptors_Call { + return &MockReviewProviderCatalog_ListReviewProviderDescriptors_Call{Call: _e.mock.On("ListReviewProviderDescriptors", ctx)} +} + +func (_c *MockReviewProviderCatalog_ListReviewProviderDescriptors_Call) Run(run func(ctx context.Context)) *MockReviewProviderCatalog_ListReviewProviderDescriptors_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + run( + arg0, + ) + }) + return _c +} + +func (_c *MockReviewProviderCatalog_ListReviewProviderDescriptors_Call) Return(reviewProviderDescriptors []ports.ReviewProviderDescriptor, err error) *MockReviewProviderCatalog_ListReviewProviderDescriptors_Call { + _c.Call.Return(reviewProviderDescriptors, err) + return _c +} + +func (_c *MockReviewProviderCatalog_ListReviewProviderDescriptors_Call) RunAndReturn(run func(ctx context.Context) ([]ports.ReviewProviderDescriptor, error)) *MockReviewProviderCatalog_ListReviewProviderDescriptors_Call { + _c.Call.Return(run) + return _c +} diff --git a/internal/ports/mocks/review_provider_client_factory_mock.go b/internal/ports/mocks/review_provider_client_factory_mock.go new file mode 100644 index 0000000..dc368fa --- /dev/null +++ b/internal/ports/mocks/review_provider_client_factory_mock.go @@ -0,0 +1,107 @@ +// Code generated by mockery; DO NOT EDIT. +// github.com/vektra/mockery +// template: testify + +package mocks + +import ( + "context" + "ero/internal/ports" + + mock "github.com/stretchr/testify/mock" +) + +// NewMockReviewProviderClientFactory creates a new instance of MockReviewProviderClientFactory. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockReviewProviderClientFactory(t interface { + mock.TestingT + Cleanup(func()) +}) *MockReviewProviderClientFactory { + mock := &MockReviewProviderClientFactory{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} + +// MockReviewProviderClientFactory is an autogenerated mock type for the ReviewProviderClientFactory type +type MockReviewProviderClientFactory struct { + mock.Mock +} + +type MockReviewProviderClientFactory_Expecter struct { + mock *mock.Mock +} + +func (_m *MockReviewProviderClientFactory) EXPECT() *MockReviewProviderClientFactory_Expecter { + return &MockReviewProviderClientFactory_Expecter{mock: &_m.Mock} +} + +// CreateReviewProviderClient provides a mock function for the type MockReviewProviderClientFactory +func (_mock *MockReviewProviderClientFactory) CreateReviewProviderClient(ctx context.Context, descriptor ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { + ret := _mock.Called(ctx, descriptor) + + if len(ret) == 0 { + panic("no return value specified for CreateReviewProviderClient") + } + + var r0 ports.ReviewProviderClient + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error)); ok { + return returnFunc(ctx, descriptor) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, ports.ReviewProviderDescriptor) ports.ReviewProviderClient); ok { + r0 = returnFunc(ctx, descriptor) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(ports.ReviewProviderClient) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, ports.ReviewProviderDescriptor) error); ok { + r1 = returnFunc(ctx, descriptor) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockReviewProviderClientFactory_CreateReviewProviderClient_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateReviewProviderClient' +type MockReviewProviderClientFactory_CreateReviewProviderClient_Call struct { + *mock.Call +} + +// CreateReviewProviderClient is a helper method to define mock.On call +// - ctx context.Context +// - descriptor ports.ReviewProviderDescriptor +func (_e *MockReviewProviderClientFactory_Expecter) CreateReviewProviderClient(ctx interface{}, descriptor interface{}) *MockReviewProviderClientFactory_CreateReviewProviderClient_Call { + return &MockReviewProviderClientFactory_CreateReviewProviderClient_Call{Call: _e.mock.On("CreateReviewProviderClient", ctx, descriptor)} +} + +func (_c *MockReviewProviderClientFactory_CreateReviewProviderClient_Call) Run(run func(ctx context.Context, descriptor ports.ReviewProviderDescriptor)) *MockReviewProviderClientFactory_CreateReviewProviderClient_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 ports.ReviewProviderDescriptor + if args[1] != nil { + arg1 = args[1].(ports.ReviewProviderDescriptor) + } + run( + arg0, + arg1, + ) + }) + return _c +} + +func (_c *MockReviewProviderClientFactory_CreateReviewProviderClient_Call) Return(reviewProviderClient ports.ReviewProviderClient, err error) *MockReviewProviderClientFactory_CreateReviewProviderClient_Call { + _c.Call.Return(reviewProviderClient, err) + return _c +} + +func (_c *MockReviewProviderClientFactory_CreateReviewProviderClient_Call) RunAndReturn(run func(ctx context.Context, descriptor ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error)) *MockReviewProviderClientFactory_CreateReviewProviderClient_Call { + _c.Call.Return(run) + return _c +} diff --git a/internal/ports/mocks/review_provider_loader_mock.go b/internal/ports/mocks/review_provider_loader_mock.go index c61bfeb..c9211a3 100644 --- a/internal/ports/mocks/review_provider_loader_mock.go +++ b/internal/ports/mocks/review_provider_loader_mock.go @@ -38,6 +38,136 @@ func (_m *MockReviewProviderLoader) EXPECT() *MockReviewProviderLoader_Expecter return &MockReviewProviderLoader_Expecter{mock: &_m.Mock} } +// CreateReviewProviderClient provides a mock function for the type MockReviewProviderLoader +func (_mock *MockReviewProviderLoader) CreateReviewProviderClient(ctx context.Context, descriptor ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { + ret := _mock.Called(ctx, descriptor) + + if len(ret) == 0 { + panic("no return value specified for CreateReviewProviderClient") + } + + var r0 ports.ReviewProviderClient + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error)); ok { + return returnFunc(ctx, descriptor) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, ports.ReviewProviderDescriptor) ports.ReviewProviderClient); ok { + r0 = returnFunc(ctx, descriptor) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(ports.ReviewProviderClient) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, ports.ReviewProviderDescriptor) error); ok { + r1 = returnFunc(ctx, descriptor) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockReviewProviderLoader_CreateReviewProviderClient_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateReviewProviderClient' +type MockReviewProviderLoader_CreateReviewProviderClient_Call struct { + *mock.Call +} + +// CreateReviewProviderClient is a helper method to define mock.On call +// - ctx context.Context +// - descriptor ports.ReviewProviderDescriptor +func (_e *MockReviewProviderLoader_Expecter) CreateReviewProviderClient(ctx interface{}, descriptor interface{}) *MockReviewProviderLoader_CreateReviewProviderClient_Call { + return &MockReviewProviderLoader_CreateReviewProviderClient_Call{Call: _e.mock.On("CreateReviewProviderClient", ctx, descriptor)} +} + +func (_c *MockReviewProviderLoader_CreateReviewProviderClient_Call) Run(run func(ctx context.Context, descriptor ports.ReviewProviderDescriptor)) *MockReviewProviderLoader_CreateReviewProviderClient_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 ports.ReviewProviderDescriptor + if args[1] != nil { + arg1 = args[1].(ports.ReviewProviderDescriptor) + } + run( + arg0, + arg1, + ) + }) + return _c +} + +func (_c *MockReviewProviderLoader_CreateReviewProviderClient_Call) Return(reviewProviderClient ports.ReviewProviderClient, err error) *MockReviewProviderLoader_CreateReviewProviderClient_Call { + _c.Call.Return(reviewProviderClient, err) + return _c +} + +func (_c *MockReviewProviderLoader_CreateReviewProviderClient_Call) RunAndReturn(run func(ctx context.Context, descriptor ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error)) *MockReviewProviderLoader_CreateReviewProviderClient_Call { + _c.Call.Return(run) + return _c +} + +// ListReviewProviderDescriptors provides a mock function for the type MockReviewProviderLoader +func (_mock *MockReviewProviderLoader) ListReviewProviderDescriptors(ctx context.Context) ([]ports.ReviewProviderDescriptor, error) { + ret := _mock.Called(ctx) + + if len(ret) == 0 { + panic("no return value specified for ListReviewProviderDescriptors") + } + + var r0 []ports.ReviewProviderDescriptor + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context) ([]ports.ReviewProviderDescriptor, error)); ok { + return returnFunc(ctx) + } + if returnFunc, ok := ret.Get(0).(func(context.Context) []ports.ReviewProviderDescriptor); ok { + r0 = returnFunc(ctx) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]ports.ReviewProviderDescriptor) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context) error); ok { + r1 = returnFunc(ctx) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockReviewProviderLoader_ListReviewProviderDescriptors_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListReviewProviderDescriptors' +type MockReviewProviderLoader_ListReviewProviderDescriptors_Call struct { + *mock.Call +} + +// ListReviewProviderDescriptors is a helper method to define mock.On call +// - ctx context.Context +func (_e *MockReviewProviderLoader_Expecter) ListReviewProviderDescriptors(ctx interface{}) *MockReviewProviderLoader_ListReviewProviderDescriptors_Call { + return &MockReviewProviderLoader_ListReviewProviderDescriptors_Call{Call: _e.mock.On("ListReviewProviderDescriptors", ctx)} +} + +func (_c *MockReviewProviderLoader_ListReviewProviderDescriptors_Call) Run(run func(ctx context.Context)) *MockReviewProviderLoader_ListReviewProviderDescriptors_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + run( + arg0, + ) + }) + return _c +} + +func (_c *MockReviewProviderLoader_ListReviewProviderDescriptors_Call) Return(reviewProviderDescriptors []ports.ReviewProviderDescriptor, err error) *MockReviewProviderLoader_ListReviewProviderDescriptors_Call { + _c.Call.Return(reviewProviderDescriptors, err) + return _c +} + +func (_c *MockReviewProviderLoader_ListReviewProviderDescriptors_Call) RunAndReturn(run func(ctx context.Context) ([]ports.ReviewProviderDescriptor, error)) *MockReviewProviderLoader_ListReviewProviderDescriptors_Call { + _c.Call.Return(run) + return _c +} + // LoadReviewProviders provides a mock function for the type MockReviewProviderLoader func (_mock *MockReviewProviderLoader) LoadReviewProviders(ctx context.Context) ([]ports.ReviewProviderClient, error) { ret := _mock.Called(ctx) diff --git a/internal/ports/plugin.go b/internal/ports/plugin.go index c363354..836f1dd 100644 --- a/internal/ports/plugin.go +++ b/internal/ports/plugin.go @@ -52,8 +52,32 @@ type PluginRemoveResult struct { RemovedRepo bool `json:"removed_repo"` } +// ReviewProviderDescriptor describes a review_provider contribution without starting its runtime. +type ReviewProviderDescriptor struct { + Key string + PluginName string + PluginVersion string + PluginSource string + PluginPath string + ContributionID string + Label string + Type string +} + +// ReviewProviderCatalog discovers review_provider contribution descriptors. +type ReviewProviderCatalog interface { + ListReviewProviderDescriptors(ctx context.Context) ([]ReviewProviderDescriptor, error) +} + +// ReviewProviderClientFactory creates live provider clients for selected descriptors. +type ReviewProviderClientFactory interface { + CreateReviewProviderClient(ctx context.Context, descriptor ReviewProviderDescriptor) (ReviewProviderClient, error) +} + // ReviewProviderLoader builds provider clients from installed plugin sources. type ReviewProviderLoader interface { + ReviewProviderCatalog + ReviewProviderClientFactory LoadReviewProviders(ctx context.Context) ([]ReviewProviderClient, error) } From 9a1bf8a0cca7635d4ae0e7c4bae57091040e0e47 Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 15:01:33 +0200 Subject: [PATCH 02/22] feat(providers): add active provider sync service --- internal/adapters/in/cli/root.go | 6 + internal/adapters/out/plugin/client.go | 26 +- internal/adapters/out/plugin/client_test.go | 12 + internal/adapters/out/providercache/cache.go | 130 +++++++ .../adapters/out/providercache/cache_test.go | 33 ++ internal/app/active_provider_service.go | 332 ++++++++++++++++++ internal/app/active_provider_service_test.go | 247 +++++++++++++ internal/app/app.go | 1 + internal/app/provider_polling_config.go | 30 ++ internal/core/provider_error.go | 74 ++++ internal/core/provider_error_test.go | 38 ++ internal/core/provider_snapshot.go | 122 +++++++ internal/core/provider_snapshot_test.go | 69 ++++ internal/ports/provider_sync.go | 19 + 14 files changed, 1138 insertions(+), 1 deletion(-) create mode 100644 internal/adapters/out/providercache/cache.go create mode 100644 internal/adapters/out/providercache/cache_test.go create mode 100644 internal/app/active_provider_service.go create mode 100644 internal/app/active_provider_service_test.go create mode 100644 internal/app/provider_polling_config.go create mode 100644 internal/core/provider_error.go create mode 100644 internal/core/provider_error_test.go create mode 100644 internal/core/provider_snapshot.go create mode 100644 internal/core/provider_snapshot_test.go create mode 100644 internal/ports/provider_sync.go diff --git a/internal/adapters/in/cli/root.go b/internal/adapters/in/cli/root.go index ba25a44..5991d9b 100644 --- a/internal/adapters/in/cli/root.go +++ b/internal/adapters/in/cli/root.go @@ -3,6 +3,7 @@ package cli import ( "fmt" "strings" + "time" "github.com/spf13/cobra" "github.com/spf13/viper" @@ -43,11 +44,13 @@ func NewRootCommand(cfg *viper.Viper, run RunFunc) (*cobra.Command, error) { flags.Int("context-lines", 3, "Number of unchanged context lines to keep around changes") flags.String("log-level", "info", "Log level (trace, debug, info, warn, error, disabled)") flags.String("log-file", "", "Write logs to this file instead of the default XDG state log") + flags.Duration("provider-sync-interval", 2*time.Minute, "Interval for active review provider background sync") cfg.SetDefault("repo-path", ".") cfg.SetDefault("context-lines", 3) cfg.SetDefault("diff-mode", string(core.DiffModeBranch)) cfg.SetDefault("startup-detect", true) cfg.SetDefault("log-level", "info") + cfg.SetDefault("provider-sync-interval", 2*time.Minute) if err := cfg.BindPFlag("repo-path", flags.Lookup("repo-path")); err != nil { return nil, fmt.Errorf("bind repo-path flag: %w", err) } @@ -60,6 +63,9 @@ func NewRootCommand(cfg *viper.Viper, run RunFunc) (*cobra.Command, error) { if err := cfg.BindPFlag("log-file", flags.Lookup("log-file")); err != nil { return nil, fmt.Errorf("bind log-file flag: %w", err) } + if err := cfg.BindPFlag("provider-sync-interval", flags.Lookup("provider-sync-interval")); err != nil { + return nil, fmt.Errorf("bind provider-sync-interval flag: %w", err) + } cfg.SetEnvPrefix("ERO") cfg.SetEnvKeyReplacer(strings.NewReplacer("-", "_")) cfg.AutomaticEnv() diff --git a/internal/adapters/out/plugin/client.go b/internal/adapters/out/plugin/client.go index 2248238..817c314 100644 --- a/internal/adapters/out/plugin/client.go +++ b/internal/adapters/out/plugin/client.go @@ -277,7 +277,7 @@ func (c *Client) call(ctx context.Context, method string, params any, result any // Check for protocol error. if resp.Error != nil { - return resp.Error + return toCoreProviderError(resp.Error) } // Decode the result into the caller's type. @@ -314,6 +314,30 @@ var _ ports.ReviewProviderClient = (*Client)(nil) // ---------- conversion helpers ---------- +func toCoreProviderError(err *protocol.Error) error { + if err == nil { + return nil + } + kind := core.ProviderErrorInternal + switch err.Code { + case protocol.ErrorAuthRequired: + kind = core.ProviderErrorAuthenticationRequired + case protocol.ErrorNotApplicable: + kind = core.ProviderErrorNotApplicable + case protocol.ErrorUnsupportedCapability: + kind = core.ProviderErrorUnsupportedCapability + case protocol.ErrorRemoteRateLimited: + kind = core.ProviderErrorRateLimited + case protocol.ErrorNetwork: + kind = core.ProviderErrorTransientNetwork + case protocol.ErrorRemoteValidationFailed, protocol.ErrorInvalidRequest: + kind = core.ProviderErrorRemoteValidation + case protocol.ErrorInternal, protocol.ErrorPartialPublishUnknown: + kind = core.ProviderErrorInternal + } + return core.NewProviderError(kind, err.Error(), err) +} + func toCoreProviderInfo(info protocol.ReviewProviderInfo) core.ReviewProviderInfo { decisions := make([]core.ReviewDecision, len(info.Capabilities.Decisions)) for i, d := range info.Capabilities.Decisions { diff --git a/internal/adapters/out/plugin/client_test.go b/internal/adapters/out/plugin/client_test.go index dc4aa7d..0fa402c 100644 --- a/internal/adapters/out/plugin/client_test.go +++ b/internal/adapters/out/plugin/client_test.go @@ -90,6 +90,9 @@ func (s *fakePluginServer) DetectContext(_ context.Context, req pluginsdk.Detect } func (s *fakePluginServer) LoadRemoteThreads(_ context.Context, req pluginsdk.LoadRemoteThreadsRequest) (pluginsdk.LoadRemoteThreadsResult, error) { + if code := os.Getenv("FAKE_PLUGIN_LOAD_ERROR_CODE"); code != "" { + return pluginsdk.LoadRemoteThreadsResult{}, protocol.NewError(code, "remote load failed") + } return pluginsdk.LoadRemoteThreadsResult{ Threads: []pluginsdk.RemoteReviewThread{ { @@ -240,6 +243,15 @@ func TestClientLoadRemoteThreads(t *testing.T) { assert.Equal(t, "LGTM", threads[0].Comments[0].Body) } +func TestClientMapsProtocolErrorsToProviderErrors(t *testing.T) { + client := setupFakeClientWithEnv(t, DefaultPluginTimeout, "FAKE_PLUGIN_LOAD_ERROR_CODE="+protocol.ErrorRemoteRateLimited) + + _, err := client.LoadRemoteThreads(context.Background(), core.ReviewContext{}) + require.Error(t, err) + assert.Equal(t, core.ProviderErrorRateLimited, core.ClassifyProviderError(err)) + assert.True(t, core.IsRetryableProviderError(err)) +} + func TestClientPublishReview(t *testing.T) { client := setupFakeClient(t) diff --git a/internal/adapters/out/providercache/cache.go b/internal/adapters/out/providercache/cache.go new file mode 100644 index 0000000..c19baea --- /dev/null +++ b/internal/adapters/out/providercache/cache.go @@ -0,0 +1,130 @@ +package providercache + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "os" + "path/filepath" + "strings" + + "ero/internal/core" +) + +// Store persists normalized provider snapshots under cache and preferences under config. +type Store struct { + cacheDir string + configDir string +} + +func NewStore(cacheDir, configDir string) *Store { + return &Store{cacheDir: cacheDir, configDir: configDir} +} + +func NewXDGStore() *Store { + return NewStore(xdgDir("XDG_CACHE_HOME", ".cache", "ero"), xdgDir("XDG_CONFIG_HOME", ".config", "ero")) +} + +func (s *Store) LoadProviderSnapshot(_ context.Context, key core.ReviewContextKey) (core.ProviderSnapshot, bool, error) { + path := s.snapshotPath(key) + data, err := os.ReadFile(path) + if errors.Is(err, os.ErrNotExist) { + return core.ProviderSnapshot{}, false, nil + } + if err != nil { + return core.ProviderSnapshot{}, false, err + } + var snapshot core.ProviderSnapshot + if err := json.Unmarshal(data, &snapshot); err != nil { + return core.ProviderSnapshot{}, false, err + } + return snapshot, true, nil +} + +func (s *Store) SaveProviderSnapshot(_ context.Context, snapshot core.ProviderSnapshot) error { + return writeJSONAtomic(s.snapshotPath(snapshot.ContextKey), snapshot) +} + +func (s *Store) LoadActiveProviderKey(_ context.Context, repositoryIdentity string) (string, bool, error) { + data, err := os.ReadFile(s.preferencePath(repositoryIdentity)) + if errors.Is(err, os.ErrNotExist) { + return "", false, nil + } + if err != nil { + return "", false, err + } + var pref struct { + StableProviderKey string `json:"stable_provider_key"` + } + if err := json.Unmarshal(data, &pref); err != nil { + return "", false, err + } + if pref.StableProviderKey == "" { + return "", false, nil + } + return pref.StableProviderKey, true, nil +} + +func (s *Store) SaveActiveProviderKey(_ context.Context, repositoryIdentity string, stableProviderKey string) error { + pref := struct { + RepositoryIdentity string `json:"repository_identity"` + StableProviderKey string `json:"stable_provider_key"` + }{repositoryIdentity, stableProviderKey} + return writeJSONAtomic(s.preferencePath(repositoryIdentity), pref) +} + +func (s *Store) snapshotPath(key core.ReviewContextKey) string { + return filepath.Join(s.cacheDir, "provider-snapshots", safeName(key.StableProviderKey), key.Digest()+".json") +} + +func (s *Store) preferencePath(repositoryIdentity string) string { + return filepath.Join(s.configDir, "provider-preferences", safeName(repositoryIdentity)+".json") +} + +func writeJSONAtomic(path string, v any) error { + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + return err + } + data, err := json.MarshalIndent(v, "", " ") + if err != nil { + return err + } + data = append(data, '\n') + tmp, err := os.CreateTemp(filepath.Dir(path), ".tmp-*") + if err != nil { + return err + } + tmpName := tmp.Name() + defer os.Remove(tmpName) + if _, err := tmp.Write(data); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Close(); err != nil { + return err + } + return os.Rename(tmpName, path) +} + +func safeName(raw string) string { + replacer := strings.NewReplacer("/", "_", "\\", "_", ":", "_", "#", "_", " ", "_") + name := replacer.Replace(raw) + if len(name) <= 80 { + return name + } + sum := sha256.Sum256([]byte(raw)) + return name[:40] + "-" + hex.EncodeToString(sum[:8]) +} + +func xdgDir(envVar, defaultBase, appName string) string { + if dir := os.Getenv(envVar); dir != "" { + return filepath.Join(dir, appName) + } + home, err := os.UserHomeDir() + if err != nil { + return filepath.Join(".", appName) + } + return filepath.Join(home, defaultBase, appName) +} diff --git a/internal/adapters/out/providercache/cache_test.go b/internal/adapters/out/providercache/cache_test.go new file mode 100644 index 0000000..fbf61d6 --- /dev/null +++ b/internal/adapters/out/providercache/cache_test.go @@ -0,0 +1,33 @@ +package providercache + +import ( + "context" + "testing" + "time" + + "ero/internal/core" +) + +func TestCacheRoundTripNormalizedSnapshot(t *testing.T) { + store := NewStore(t.TempDir(), t.TempDir()) + key := core.ReviewContextKey{StableProviderKey: "plugin:github#review_provider:github", RepositoryIdentity: "remotes:github.com/o/r", TargetMode: core.DiffModeBranch, BaseRef: "main", HeadRef: "feature", BaseSHA: "b", HeadSHA: "h"} + snapshot := core.ProviderSnapshot{StableProviderKey: key.StableProviderKey, RuntimeProviderID: "github", ContextKey: key, Threads: []core.RemoteReviewThread{{ProviderID: "github", ExternalID: "thread-1"}}, FetchedAt: time.Unix(10, 0).UTC()} + + if err := store.SaveProviderSnapshot(context.Background(), snapshot); err != nil { t.Fatal(err) } + got, ok, err := store.LoadProviderSnapshot(context.Background(), key) + if err != nil { t.Fatal(err) } + if !ok { t.Fatal("expected cached snapshot") } + if got.StableProviderKey != snapshot.StableProviderKey || got.RuntimeProviderID != "github" || len(got.Threads) != 1 { + t.Fatalf("unexpected snapshot: %#v", got) + } +} + +func TestPreferenceRoundTrip(t *testing.T) { + store := NewStore(t.TempDir(), t.TempDir()) + repoID := "remotes:github.com/o/r" + if _, ok, err := store.LoadActiveProviderKey(context.Background(), repoID); err != nil || ok { t.Fatalf("empty load = ok %v err %v", ok, err) } + if err := store.SaveActiveProviderKey(context.Background(), repoID, "provider-key"); err != nil { t.Fatal(err) } + got, ok, err := store.LoadActiveProviderKey(context.Background(), repoID) + if err != nil { t.Fatal(err) } + if !ok || got != "provider-key" { t.Fatalf("got %q ok %v", got, ok) } +} diff --git a/internal/app/active_provider_service.go b/internal/app/active_provider_service.go new file mode 100644 index 0000000..9cc0fea --- /dev/null +++ b/internal/app/active_provider_service.go @@ -0,0 +1,332 @@ +package app + +import ( + "context" + "strings" + "sync" + "time" + + "ero/internal/core" + "ero/internal/ports" +) + +type ProviderPollingConfig struct{ Interval, MinBackoff, MaxBackoff time.Duration } + +func DefaultProviderPollingConfig() ProviderPollingConfig { + return ProviderPollingConfig{Interval: 2 * time.Minute, MinBackoff: 5 * time.Second, MaxBackoff: time.Minute} +} + +type ActiveProviderState struct { + StableProviderKey string + RuntimeProviderID string + Snapshot core.ProviderSnapshot + FromCache bool + Syncing bool + LastError error + NextSyncAt time.Time +} + +type ActiveProviderService struct { + catalog ports.ReviewProviderCatalog + factory ports.ReviewProviderClientFactory + cache ports.ProviderSnapshotCache + prefs ports.ActiveProviderPreferenceStore + poll ProviderPollingConfig + + mu sync.Mutex + client ports.ReviewProviderClient + stableKey string + runtimeID string + generation int64 + backoff time.Duration + state ActiveProviderState +} + +func NewActiveProviderService(catalog ports.ReviewProviderCatalog, factory ports.ReviewProviderClientFactory, cache ports.ProviderSnapshotCache, prefs ports.ActiveProviderPreferenceStore, poll ProviderPollingConfig) *ActiveProviderService { + if poll.Interval == 0 { + poll = DefaultProviderPollingConfig() + } + if poll.MinBackoff == 0 { + poll.MinBackoff = 5 * time.Second + } + if poll.MaxBackoff == 0 { + poll.MaxBackoff = time.Minute + } + return &ActiveProviderService{catalog: catalog, factory: factory, cache: cache, prefs: prefs, poll: poll} +} + +func (s *ActiveProviderService) Start(ctx context.Context, review core.ReviewContext) (ActiveProviderState, error) { + descs, err := s.catalog.ListReviewProviderDescriptors(ctx) + if err != nil { + return ActiveProviderState{}, err + } + ordered := s.orderCandidates(ctx, descs, review) + s.mu.Lock() + s.closeLocked() + s.stableKey = "" + s.runtimeID = "" + s.generation++ + startGen := s.generation + s.state = ActiveProviderState{} + s.mu.Unlock() + var lastErr error + for _, d := range ordered { + client, info, err := s.probe(ctx, d, review) + if err != nil { + lastErr = err + continue + } + s.mu.Lock() + s.closeLocked() + s.client = client + s.stableKey = d.Key + s.runtimeID = info.ID + s.generation++ + gen := s.generation + s.backoff = 0 + s.mu.Unlock() + if s.prefs != nil { + _ = s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key) + } + st := s.loadCachedState(ctx, review, d.Key, info.ID) + s.setState(gen, st) + return st, nil + } + failed := failedProviderState(lastErr) + s.setState(startGen, failed) + return failed, lastErr +} + +func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewContext, stableKey string) (ActiveProviderState, error) { + descs, err := s.catalog.ListReviewProviderDescriptors(ctx) + if err != nil { + return ActiveProviderState{}, err + } + for _, d := range descs { + if d.Key == stableKey { + s.mu.Lock() + s.closeLocked() + s.stableKey = "" + s.runtimeID = "" + s.generation++ + switchGen := s.generation + s.state = ActiveProviderState{} + s.mu.Unlock() + client, info, err := s.probe(ctx, d, review) + if err != nil { + failed := failedProviderState(err) + s.setState(switchGen, failed) + return failed, err + } + s.mu.Lock() + s.closeLocked() + s.client = client + s.stableKey = d.Key + s.runtimeID = info.ID + s.generation++ + gen := s.generation + s.backoff = 0 + s.mu.Unlock() + if s.prefs != nil { + _ = s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key) + } + st := s.loadCachedState(ctx, review, d.Key, info.ID) + s.setState(gen, st) + return st, nil + } + } + return ActiveProviderState{}, core.NewProviderError(core.ProviderErrorNotApplicable, "provider descriptor not found", nil) +} + +func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewContext, manual bool) (ActiveProviderState, error) { + s.mu.Lock() + client := s.client + key := s.stableKey + runtimeID := s.runtimeID + s.generation++ + gen := s.generation + prev := s.state + if manual { + s.backoff = 0 + } + s.mu.Unlock() + if client == nil { + return prev, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) + } + threads, err := client.LoadRemoteThreads(ctx, review) + if err != nil { + st := prev + st.LastError = err + st.Syncing = false + st.Snapshot.Sync.LastError = err.Error() + st.NextSyncAt = s.nextBackoff(err) + if st.NextSyncAt.IsZero() { + st.Snapshot.Sync.Status = core.ProviderSyncStatusFailed + st.Snapshot.Sync.NextSyncAt = nil + } else { + st.Snapshot.Sync.Status = core.ProviderSyncStatusBackingOff + st.Snapshot.Sync.NextSyncAt = new(st.NextSyncAt) + } + s.setState(gen, st) + return st, err + } + now := time.Now().UTC() + next := now.Add(s.poll.Interval) + snap := core.ProviderSnapshot{StableProviderKey: key, RuntimeProviderID: runtimeID, ContextKey: core.NewReviewContextKey(key, review), Threads: threads, FetchedAt: now, Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusSynced, LastSyncAt: new(now), NextSyncAt: new(next)}} + if s.cache != nil { + _ = s.cache.SaveProviderSnapshot(ctx, snap) + } + st := ActiveProviderState{StableProviderKey: key, RuntimeProviderID: runtimeID, Snapshot: snap, NextSyncAt: next} + s.setState(gen, st) + return st, nil +} + +func (s *ActiveProviderService) CompleteTimer(ctx context.Context, review core.ReviewContext, generation int64) (ActiveProviderState, error) { + s.mu.Lock() + cur := s.generation + s.mu.Unlock() + if generation != cur { + return s.State(), nil + } + return s.Refresh(ctx, review, false) +} +func (s *ActiveProviderService) Generation() int64 { + s.mu.Lock() + defer s.mu.Unlock() + return s.generation +} +func (s *ActiveProviderService) State() ActiveProviderState { + s.mu.Lock() + defer s.mu.Unlock() + return s.state +} +func (s *ActiveProviderService) Close() error { + s.mu.Lock() + defer s.mu.Unlock() + return s.closeLocked() +} + +func (s *ActiveProviderService) orderCandidates(ctx context.Context, descs []ports.ReviewProviderDescriptor, review core.ReviewContext) []ports.ReviewProviderDescriptor { + out := make([]ports.ReviewProviderDescriptor, 0, len(descs)) + used := map[string]bool{} + preferenceFound := false + if s.prefs != nil { + if key, ok, _ := s.prefs.LoadActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository)); ok { + for _, d := range descs { + if d.Key == key { + out = append(out, d) + used[d.Key] = true + preferenceFound = true + break + } + } + } + } + if !preferenceFound { + for _, d := range descs { + if plausibleGitHub(review, d) { + out = append(out, d) + used[d.Key] = true + break + } + } + } + for _, d := range descs { + if !used[d.Key] { + out = append(out, d) + } + } + return out +} + +func plausibleGitHub(review core.ReviewContext, d ports.ReviewProviderDescriptor) bool { + descriptorText := strings.ToLower(strings.Join([]string{d.Key, d.PluginName, d.PluginSource, d.ContributionID, d.Label, d.Type}, " ")) + if !strings.Contains(descriptorText, "github") { + return false + } + if strings.Contains(d.PluginSource, "github") || strings.Contains(strings.ToLower(d.ContributionID), "github") || strings.Contains(strings.ToLower(d.Label), "github") || strings.Contains(strings.ToLower(d.PluginName), "github") { + return true + } + for _, r := range review.Repository.Remotes { + if strings.Contains(strings.ToLower(r.URL), "github.com") { + return true + } + } + return false +} + +func (s *ActiveProviderService) probe(ctx context.Context, d ports.ReviewProviderDescriptor, review core.ReviewContext) (ports.ReviewProviderClient, core.ReviewProviderInfo, error) { + client, err := s.factory.CreateReviewProviderClient(ctx, d) + if err != nil { + return nil, core.ReviewProviderInfo{}, err + } + info, err := client.Initialize(ctx) + if err != nil { + _ = client.Close() + return nil, info, err + } + det, err := client.DetectContext(ctx, review) + if err != nil { + _ = client.Close() + return nil, info, err + } + if !det.Applicable { + _ = client.Close() + return nil, info, core.NewProviderError(core.ProviderErrorNotApplicable, det.Reason, nil) + } + return client, info, nil +} +func (s *ActiveProviderService) loadCachedState(ctx context.Context, review core.ReviewContext, key, runtimeID string) ActiveProviderState { + st := ActiveProviderState{StableProviderKey: key, RuntimeProviderID: runtimeID} + if s.cache != nil { + if snap, ok, _ := s.cache.LoadProviderSnapshot(ctx, core.NewReviewContextKey(key, review)); ok { + snap.Cached = true + st.Snapshot = snap + st.FromCache = true + if snap.Sync.NextSyncAt != nil { + st.NextSyncAt = *snap.Sync.NextSyncAt + } + } + } + return st +} +func (s *ActiveProviderService) setState(gen int64, st ActiveProviderState) { + s.mu.Lock() + defer s.mu.Unlock() + if gen == s.generation { + s.state = st + } +} +func (s *ActiveProviderService) nextBackoff(err error) time.Time { + s.mu.Lock() + defer s.mu.Unlock() + if !core.IsRetryableProviderError(err) { + return time.Time{} + } + if s.backoff == 0 { + s.backoff = s.poll.MinBackoff + } else { + s.backoff *= 2 + if s.backoff > s.poll.MaxBackoff { + s.backoff = s.poll.MaxBackoff + } + } + return time.Now().UTC().Add(s.backoff) +} +func (s *ActiveProviderService) closeLocked() error { + if s.client == nil { + return nil + } + err := s.client.Close() + s.client = nil + return err +} + +func failedProviderState(err error) ActiveProviderState { + st := ActiveProviderState{LastError: err} + if err != nil { + st.Snapshot.Sync.Status = core.ProviderSyncStatusFailed + st.Snapshot.Sync.LastError = err.Error() + } + return st +} diff --git a/internal/app/active_provider_service_test.go b/internal/app/active_provider_service_test.go new file mode 100644 index 0000000..85804a5 --- /dev/null +++ b/internal/app/active_provider_service_test.go @@ -0,0 +1,247 @@ +package app + +import ( + "context" + "errors" + "testing" + "time" + + "ero/internal/core" + "ero/internal/ports" +) + +type memCatalog []ports.ReviewProviderDescriptor + +func (m memCatalog) ListReviewProviderDescriptors(context.Context) ([]ports.ReviewProviderDescriptor, error) { + return []ports.ReviewProviderDescriptor(m), nil +} + +type memFactory struct { + clients map[string]*fakeProvider + made []string + beforeCreate func(string) +} + +func (f *memFactory) CreateReviewProviderClient(_ context.Context, d ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { + if f.beforeCreate != nil { + f.beforeCreate(d.Key) + } + f.made = append(f.made, d.Key) + return f.clients[d.Key], nil +} + +type fakeProvider struct { + id string + applicable bool + detectErr, errorLoad error + closed int + threads []core.RemoteReviewThread +} + +func (f *fakeProvider) Initialize(context.Context) (core.ReviewProviderInfo, error) { + return core.ReviewProviderInfo{ID: f.id, Capabilities: core.ReviewProviderCapabilities{LoadRemoteComments: true}}, nil +} +func (f *fakeProvider) DetectContext(context.Context, core.ReviewContext) (core.DetectionResult, error) { + if f.detectErr != nil { + return core.DetectionResult{}, f.detectErr + } + return core.DetectionResult{Applicable: f.applicable, Reason: "nope"}, nil +} +func (f *fakeProvider) LoadRemoteThreads(context.Context, core.ReviewContext) ([]core.RemoteReviewThread, error) { + if f.errorLoad != nil { + return nil, f.errorLoad + } + return f.threads, nil +} +func (f *fakeProvider) PublishReview(context.Context, core.PublishReviewRequest) (core.PublishReviewResult, error) { + return core.PublishReviewResult{}, nil +} +func (f *fakeProvider) Close() error { f.closed++; return nil } + +type memCache struct { + snap core.ProviderSnapshot + ok bool +} + +func (m *memCache) LoadProviderSnapshot(context.Context, core.ReviewContextKey) (core.ProviderSnapshot, bool, error) { + return m.snap, m.ok, nil +} +func (m *memCache) SaveProviderSnapshot(_ context.Context, s core.ProviderSnapshot) error { + m.snap = s + m.ok = true + return nil +} + +type memPrefs struct { + key string + ok bool +} + +func (m *memPrefs) LoadActiveProviderKey(context.Context, string) (string, bool, error) { + return m.key, m.ok, nil +} +func (m *memPrefs) SaveActiveProviderKey(_ context.Context, _ string, k string) error { + m.key = k + m.ok = true + return nil +} + +func testReviewContext() core.ReviewContext { + return core.ReviewContext{Repository: core.RepositoryMetadata{Remotes: []core.GitRemote{{URL: "https://github.com/acme/repo.git"}}}, Target: core.ReviewTargetMetadata{Mode: core.DiffModeWorking}} +} + +func TestActiveProviderServicePreferenceFallbackAndClosesFailedClients(t *testing.T) { + bad := &fakeProvider{id: "bad", applicable: false} + good := &fakeProvider{id: "good", applicable: true} + fac := &memFactory{clients: map[string]*fakeProvider{"preferred": bad, "github": good}} + svc := NewActiveProviderService(memCatalog{{Key: "preferred"}, {Key: "github", Type: "github"}}, fac, nil, &memPrefs{key: "preferred", ok: true}, ProviderPollingConfig{}) + st, err := svc.Start(context.Background(), testReviewContext()) + if err != nil { + t.Fatal(err) + } + if st.StableProviderKey != "github" { + t.Fatalf("got %q", st.StableProviderKey) + } + if bad.closed != 1 { + t.Fatalf("failed client not closed") + } + if good.closed != 0 { + t.Fatalf("active client was closed") + } +} + +func TestActiveProviderServiceCacheFirstRefreshPreservesCacheOnRetryableFailure(t *testing.T) { + review := testReviewContext() + key := core.NewReviewContextKey("github", review) + cached := core.ProviderSnapshot{StableProviderKey: "github", ContextKey: key, Threads: []core.RemoteReviewThread{{ExternalID: "old"}}} + cache := &memCache{snap: cached, ok: true} + p := &fakeProvider{id: "rt", applicable: true, errorLoad: core.NewProviderError(core.ProviderErrorTransientNetwork, "offline", errors.New("dial"))} + svc := NewActiveProviderService(memCatalog{{Key: "github", Type: "github"}}, &memFactory{clients: map[string]*fakeProvider{"github": p}}, cache, nil, ProviderPollingConfig{Interval: time.Minute, MinBackoff: time.Second, MaxBackoff: time.Second}) + st, err := svc.Start(context.Background(), review) + if err != nil { + t.Fatal(err) + } + if !st.FromCache || len(st.Snapshot.Threads) != 1 { + t.Fatalf("expected cached snapshot first") + } + st, err = svc.Refresh(context.Background(), review, false) + if err == nil { + t.Fatal("expected refresh error") + } + if len(st.Snapshot.Threads) != 1 || st.Snapshot.Threads[0].ExternalID != "old" { + t.Fatalf("cache not preserved") + } + if st.NextSyncAt.IsZero() { + t.Fatalf("retryable error should set next sync") + } + if st.Snapshot.Sync.Status != core.ProviderSyncStatusBackingOff { + t.Fatalf("retryable error status = %q, want %q", st.Snapshot.Sync.Status, core.ProviderSyncStatusBackingOff) + } +} + +func TestActiveProviderServiceSwitchGenerationIgnoresStaleTimer(t *testing.T) { + review := testReviewContext() + a := &fakeProvider{id: "a", applicable: true} + b := &fakeProvider{id: "b", applicable: true} + svc := NewActiveProviderService(memCatalog{{Key: "a"}, {Key: "b"}}, &memFactory{clients: map[string]*fakeProvider{"a": a, "b": b}}, nil, nil, ProviderPollingConfig{}) + st, err := svc.Start(context.Background(), review) + if err != nil { + t.Fatal(err) + } + target := "b" + oldProvider := a + if st.StableProviderKey == "b" { + target = "a" + oldProvider = b + } + old := svc.Generation() + if _, err := svc.Switch(context.Background(), review, target); err != nil { + t.Fatal(err) + } + if _, err := svc.CompleteTimer(context.Background(), review, old); err != nil { + t.Fatal(err) + } + if got := svc.State().StableProviderKey; got != target { + t.Fatalf("stale timer changed state to %q", got) + } + if oldProvider.closed != 1 { + t.Fatalf("switch should close old client") + } +} + +func TestActiveProviderServiceSwitchClosesCurrentBeforeStartingTarget(t *testing.T) { + review := testReviewContext() + a := &fakeProvider{id: "a", applicable: true} + b := &fakeProvider{id: "b", applicable: true} + factory := &memFactory{clients: map[string]*fakeProvider{"a": a, "b": b}} + svc := NewActiveProviderService(memCatalog{{Key: "a"}, {Key: "b"}}, factory, nil, nil, ProviderPollingConfig{}) + if _, err := svc.Start(context.Background(), review); err != nil { + t.Fatal(err) + } + factory.beforeCreate = func(key string) { + if key == "b" && a.closed != 1 { + t.Fatalf("current provider was still live when target provider started") + } + } + if _, err := svc.Switch(context.Background(), review, "b"); err != nil { + t.Fatal(err) + } +} + +func TestActiveProviderServiceFailedSwitchClearsOldProviderState(t *testing.T) { + review := testReviewContext() + a := &fakeProvider{id: "a", applicable: true, threads: []core.RemoteReviewThread{{ExternalID: "old"}}} + b := &fakeProvider{id: "b", applicable: false} + svc := NewActiveProviderService(memCatalog{{Key: "a"}, {Key: "b"}}, &memFactory{clients: map[string]*fakeProvider{"a": a, "b": b}}, nil, nil, ProviderPollingConfig{}) + if _, err := svc.Start(context.Background(), review); err != nil { + t.Fatal(err) + } + if _, err := svc.Refresh(context.Background(), review, false); err != nil { + t.Fatal(err) + } + if state := svc.State(); state.StableProviderKey != "a" || len(state.Snapshot.Threads) != 1 { + t.Fatalf("expected old active provider state before switch, got %#v", state) + } + + st, err := svc.Switch(context.Background(), review, "b") + if err == nil { + t.Fatal("expected switch failure") + } + if st.StableProviderKey != "" || len(st.Snapshot.Threads) != 0 { + t.Fatalf("failed switch returned old provider state: %#v", st) + } + if state := svc.State(); state.StableProviderKey != "" || len(state.Snapshot.Threads) != 0 || state.LastError == nil { + t.Fatalf("failed switch left incoherent service state: %#v", state) + } + if a.closed != 1 { + t.Fatalf("old provider should be closed, got %d", a.closed) + } +} + +func TestActiveProviderServiceUsesStableCatalogOrder(t *testing.T) { + review := testReviewContext() + fac := &memFactory{clients: map[string]*fakeProvider{ + "first": {id: "first", applicable: true}, + "github": {id: "github", applicable: true}, + "second": {id: "second", applicable: true}, + }} + svc := NewActiveProviderService(memCatalog{{Key: "first"}, {Key: "github", ContributionID: "github"}, {Key: "second"}}, fac, nil, nil, ProviderPollingConfig{}) + if _, err := svc.Start(context.Background(), review); err != nil { + t.Fatal(err) + } + if got := fac.made; len(got) != 1 || got[0] != "github" { + t.Fatalf("github fallback should be selected in catalog order, got %v", got) + } + + fac = &memFactory{clients: map[string]*fakeProvider{ + "z": {id: "z", applicable: false}, + "a": {id: "a", applicable: true}, + }} + svc = NewActiveProviderService(memCatalog{{Key: "z"}, {Key: "a"}}, fac, nil, nil, ProviderPollingConfig{}) + if _, err := svc.Start(context.Background(), review); err != nil { + t.Fatal(err) + } + if got := fac.made; len(got) != 2 || got[0] != "z" || got[1] != "a" { + t.Fatalf("remaining candidates should preserve catalog order, got %v", got) + } +} diff --git a/internal/app/app.go b/internal/app/app.go index 52d084a..7f78e39 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -104,6 +104,7 @@ func newAppWithClipboard(cfg *viper.Viper, loader reviewLoader, runner tuiRunner var reviewProviders []ports.ReviewProviderClient pluginManager := pluginadapter.NewManager() providerLoader := pluginadapter.NewReviewProviderLoader(pluginManager) + _ = providerPollingConfigFromConfig(cfg) providers, err := buildReviewProviders(ctx, providerLoader, providerLoader) if err != nil { log.Warn().Err(err).Msg("load review providers failed") diff --git a/internal/app/provider_polling_config.go b/internal/app/provider_polling_config.go new file mode 100644 index 0000000..ad8f9e3 --- /dev/null +++ b/internal/app/provider_polling_config.go @@ -0,0 +1,30 @@ +package app + +import ( + "time" + + "github.com/spf13/viper" +) + +func providerPollingConfigFromConfig(cfg *viper.Viper) ProviderPollingConfig { + poll := DefaultProviderPollingConfig() + if cfg == nil { + return poll + } + if interval := cfg.GetDuration("provider-sync-interval"); interval > 0 { + poll.Interval = interval + } + if min := cfg.GetDuration("provider-sync-min-backoff"); min > 0 { + poll.MinBackoff = min + } + if max := cfg.GetDuration("provider-sync-max-backoff"); max > 0 { + poll.MaxBackoff = max + } + if poll.MinBackoff == 0 { + poll.MinBackoff = 5 * time.Second + } + if poll.MaxBackoff == 0 { + poll.MaxBackoff = time.Minute + } + return poll +} diff --git a/internal/core/provider_error.go b/internal/core/provider_error.go new file mode 100644 index 0000000..f8cb9bd --- /dev/null +++ b/internal/core/provider_error.go @@ -0,0 +1,74 @@ +package core + +import ( + "errors" + "fmt" +) + +// ProviderErrorKind classifies provider failures across app/port/adapter boundaries. +type ProviderErrorKind string + +const ( + ProviderErrorAuthenticationRequired ProviderErrorKind = "authentication_required" + ProviderErrorNotApplicable ProviderErrorKind = "not_applicable" + ProviderErrorUnsupportedCapability ProviderErrorKind = "unsupported_capability" + ProviderErrorRateLimited ProviderErrorKind = "rate_limited" + ProviderErrorTransientNetwork ProviderErrorKind = "transient_network" + ProviderErrorRemoteValidation ProviderErrorKind = "remote_validation" + ProviderErrorInternal ProviderErrorKind = "internal" +) + +// ProviderError is the common classified provider error type returned by ports/adapters. +type ProviderError struct { + Kind ProviderErrorKind + Message string + Err error +} + +func (e *ProviderError) Error() string { + if e == nil { + return "" + } + if e.Message != "" { + return e.Message + } + if e.Err != nil { + return e.Err.Error() + } + return string(e.Kind) +} + +func (e *ProviderError) Unwrap() error { + if e == nil { + return nil + } + return e.Err +} + +func NewProviderError(kind ProviderErrorKind, message string, err error) *ProviderError { + return &ProviderError{Kind: kind, Message: message, Err: err} +} + +func ClassifyProviderError(err error) ProviderErrorKind { + if err == nil { + return "" + } + var providerErr *ProviderError + if errors.As(err, &providerErr) && providerErr.Kind != "" { + return providerErr.Kind + } + return ProviderErrorInternal +} + +func IsRetryableProviderError(err error) bool { + switch ClassifyProviderError(err) { + case ProviderErrorRateLimited, ProviderErrorTransientNetwork: + return true + default: + return false + } +} + +func FormatProviderError(kind ProviderErrorKind, format string, args ...any) *ProviderError { + return NewProviderError(kind, fmt.Sprintf(format, args...), nil) +} diff --git a/internal/core/provider_error_test.go b/internal/core/provider_error_test.go new file mode 100644 index 0000000..70c9f9b --- /dev/null +++ b/internal/core/provider_error_test.go @@ -0,0 +1,38 @@ +package core + +import ( + "fmt" + "testing" +) + +func TestClassifyProviderError(t *testing.T) { + tests := []ProviderErrorKind{ + ProviderErrorAuthenticationRequired, + ProviderErrorNotApplicable, + ProviderErrorUnsupportedCapability, + ProviderErrorRateLimited, + ProviderErrorTransientNetwork, + ProviderErrorRemoteValidation, + ProviderErrorInternal, + } + for _, kind := range tests { + t.Run(string(kind), func(t *testing.T) { + err := fmt.Errorf("client boundary: %w", NewProviderError(kind, "boom", nil)) + if got := ClassifyProviderError(err); got != kind { + t.Fatalf("got %q want %q", got, kind) + } + }) + } +} + +func TestIsRetryableProviderError(t *testing.T) { + if !IsRetryableProviderError(NewProviderError(ProviderErrorRateLimited, "", nil)) { + t.Fatal("rate limited should retry") + } + if !IsRetryableProviderError(NewProviderError(ProviderErrorTransientNetwork, "", nil)) { + t.Fatal("network should retry") + } + if IsRetryableProviderError(NewProviderError(ProviderErrorAuthenticationRequired, "", nil)) { + t.Fatal("auth should not retry") + } +} diff --git a/internal/core/provider_snapshot.go b/internal/core/provider_snapshot.go new file mode 100644 index 0000000..cea186e --- /dev/null +++ b/internal/core/provider_snapshot.go @@ -0,0 +1,122 @@ +package core + +import ( + "crypto/sha256" + "encoding/hex" + "net/url" + "path/filepath" + "sort" + "strings" + "time" +) + +// ReviewContextKey identifies a provider snapshot for a stable provider and review target. +type ReviewContextKey struct { + StableProviderKey string `json:"stable_provider_key"` + RepositoryIdentity string `json:"repository_identity"` + TargetMode DiffMode `json:"target_mode"` + BaseRef string `json:"base_ref,omitempty"` + HeadRef string `json:"head_ref,omitempty"` + BaseSHA string `json:"base_sha,omitempty"` + HeadSHA string `json:"head_sha,omitempty"` + MergeBaseSHA string `json:"merge_base_sha,omitempty"` +} + +// NewReviewContextKey builds a cache/preference identity from stable review inputs only. +func NewReviewContextKey(stableProviderKey string, ctx ReviewContext) ReviewContextKey { + return ReviewContextKey{ + StableProviderKey: stableProviderKey, + RepositoryIdentity: RepositoryIdentity(ctx.Repository), + TargetMode: ctx.Target.Mode, + BaseRef: ctx.Target.BaseRef, + HeadRef: ctx.Target.HeadRef, + BaseSHA: ctx.Target.BaseSHA, + HeadSHA: ctx.Target.HeadSHA, + MergeBaseSHA: ctx.Target.MergeBaseSHA, + } +} + +// RepositoryIdentity returns a stable repository identity, preferring normalized remotes. +func RepositoryIdentity(repo RepositoryMetadata) string { + remotes := make([]string, 0, len(repo.Remotes)) + for _, remote := range repo.Remotes { + if normalized := normalizeRemoteURL(remote.URL); normalized != "" { + remotes = append(remotes, normalized) + } + } + if len(remotes) > 0 { + sort.Strings(remotes) + return "remotes:" + strings.Join(remotes, ",") + } + if repo.WorktreeRoot != "" { + return "path:" + filepath.Clean(repo.WorktreeRoot) + } + if repo.RepoPath != "" { + return "path:" + filepath.Clean(repo.RepoPath) + } + return "unknown" +} + +// Digest returns a filesystem-safe digest for the context key. +func (k ReviewContextKey) Digest() string { + parts := []string{k.StableProviderKey, k.RepositoryIdentity, string(k.TargetMode), k.BaseRef, k.HeadRef, k.BaseSHA, k.HeadSHA, k.MergeBaseSHA} + sum := sha256.Sum256([]byte(strings.Join(parts, "\x00"))) + return hex.EncodeToString(sum[:]) +} + +// ProviderSnapshot is Ero's normalized cached view of remote review data. +type ProviderSnapshot struct { + StableProviderKey string `json:"stable_provider_key"` + RuntimeProviderID string `json:"runtime_provider_id,omitempty"` + ContextKey ReviewContextKey `json:"context_key"` + Threads []RemoteReviewThread `json:"threads"` + Overview *ProviderOverview `json:"overview,omitempty"` + Metadata map[string]string `json:"metadata,omitempty"` + FetchedAt time.Time `json:"fetched_at"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` + Cached bool `json:"cached"` + Stale bool `json:"stale"` + Sync ProviderSyncState `json:"sync"` +} + +// ProviderOverview is a placeholder for richer provider overview data. +type ProviderOverview struct { + Title string `json:"title,omitempty"` + ExternalURL string `json:"external_url,omitempty"` + Body string `json:"body,omitempty"` +} + +type ProviderSyncStatus string + +const ( + ProviderSyncStatusIdle ProviderSyncStatus = "idle" + ProviderSyncStatusLoadingCache ProviderSyncStatus = "loading_cache" + ProviderSyncStatusSyncing ProviderSyncStatus = "syncing" + ProviderSyncStatusSynced ProviderSyncStatus = "synced" + ProviderSyncStatusFailed ProviderSyncStatus = "failed" + ProviderSyncStatusBackingOff ProviderSyncStatus = "backing_off" +) + +type ProviderSyncState struct { + Status ProviderSyncStatus `json:"status"` + LastSyncAt *time.Time `json:"last_sync_at,omitempty"` + NextSyncAt *time.Time `json:"next_sync_at,omitempty"` + LastError string `json:"last_error,omitempty"` +} + +func normalizeRemoteURL(raw string) string { + raw = strings.TrimSpace(raw) + if raw == "" { + return "" + } + if strings.HasPrefix(raw, "git@") && strings.Contains(raw, ":") { + trimmed := strings.TrimPrefix(raw, "git@") + parts := strings.SplitN(trimmed, ":", 2) + return strings.ToLower(parts[0]) + "/" + strings.TrimSuffix(parts[1], ".git") + } + if u, err := url.Parse(raw); err == nil && u.Host != "" { + path := strings.TrimPrefix(strings.TrimSuffix(u.Path, ".git"), "/") + return strings.ToLower(u.Host) + "/" + path + } + return strings.TrimSuffix(raw, ".git") +} diff --git a/internal/core/provider_snapshot_test.go b/internal/core/provider_snapshot_test.go new file mode 100644 index 0000000..90c478d --- /dev/null +++ b/internal/core/provider_snapshot_test.go @@ -0,0 +1,69 @@ +package core + +import ( + "testing" + "time" +) + +func TestReviewContextKeyStableAcrossSessionFields(t *testing.T) { + ctx := sampleProviderReviewContext() + key := NewReviewContextKey("plugin:github#review_provider:github", ctx) + + ctx.Session.LocalReviewID = "different" + ctx.Session.IdempotencyKey = "different-key" + ctx.Session.CreatedAt = time.Now().Add(24 * time.Hour) + + if got := NewReviewContextKey("plugin:github#review_provider:github", ctx); got != key { + t.Fatalf("key changed for runtime session fields\nwant: %#v\n got: %#v", key, got) + } +} + +func TestReviewContextKeyIncludesIdentityInputs(t *testing.T) { + base := sampleProviderReviewContext() + baseKey := NewReviewContextKey("provider-a", base) + + cases := map[string]func(*ReviewContext){ + "provider": nil, + "remote": func(c *ReviewContext) { c.Repository.Remotes[0].URL = "git@github.com:owner/other.git" }, + "mode": func(c *ReviewContext) { c.Target.Mode = DiffModeRange }, + "base ref": func(c *ReviewContext) { c.Target.BaseRef = "main" }, + "head ref": func(c *ReviewContext) { c.Target.HeadRef = "feature-2" }, + "base sha": func(c *ReviewContext) { c.Target.BaseSHA = "base2" }, + "head sha": func(c *ReviewContext) { c.Target.HeadSHA = "head2" }, + "merge base": func(c *ReviewContext) { c.Target.MergeBaseSHA = "merge2" }, + } + + for name, mutate := range cases { + ctx := base + provider := "provider-a" + if mutate == nil { provider = "provider-b" } else { mutate(&ctx) } + if got := NewReviewContextKey(provider, ctx); got == baseKey { + t.Fatalf("%s did not change key", name) + } + } +} + +func TestRepositoryIdentityPrefersRemotesAndFallsBackToPath(t *testing.T) { + ctx := sampleProviderReviewContext() + remoteID := NewReviewContextKey("provider", ctx).RepositoryIdentity + + ctx.Repository.RepoPath = "/different/path" + ctx.Repository.WorktreeRoot = "/different/worktree" + if got := NewReviewContextKey("provider", ctx).RepositoryIdentity; got != remoteID { + t.Fatalf("path changed remote-backed identity: %q != %q", got, remoteID) + } + + ctx.Repository.Remotes = nil + got := NewReviewContextKey("provider", ctx).RepositoryIdentity + if got == "" || got == remoteID { + t.Fatalf("expected path fallback identity, got %q", got) + } +} + +func sampleProviderReviewContext() ReviewContext { + return ReviewContext{ + Repository: RepositoryMetadata{RepoPath: "/repo", WorktreeRoot: "/repo", Remotes: []GitRemote{{Name: "origin", URL: "git@github.com:owner/repo.git"}}}, + Target: ReviewTargetMetadata{Mode: DiffModeBranch, BaseRef: "origin/main", HeadRef: "feature", BaseSHA: "base1", HeadSHA: "head1", MergeBaseSHA: "merge1"}, + Session: ReviewSessionMetadata{LocalReviewID: "local", IdempotencyKey: "idem", CreatedAt: time.Unix(1, 0)}, + } +} diff --git a/internal/ports/provider_sync.go b/internal/ports/provider_sync.go new file mode 100644 index 0000000..3cf7f7d --- /dev/null +++ b/internal/ports/provider_sync.go @@ -0,0 +1,19 @@ +package ports + +import ( + "context" + + "ero/internal/core" +) + +// ProviderSnapshotCache stores normalized Ero provider snapshots, not raw provider payloads. +type ProviderSnapshotCache interface { + LoadProviderSnapshot(ctx context.Context, key core.ReviewContextKey) (core.ProviderSnapshot, bool, error) + SaveProviderSnapshot(ctx context.Context, snapshot core.ProviderSnapshot) error +} + +// ActiveProviderPreferenceStore stores the user's active stable provider key per repository identity. +type ActiveProviderPreferenceStore interface { + LoadActiveProviderKey(ctx context.Context, repositoryIdentity string) (string, bool, error) + SaveActiveProviderKey(ctx context.Context, repositoryIdentity string, stableProviderKey string) error +} From 1ec6d958e3804e165210dc973a920276a10e95d1 Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 15:24:03 +0200 Subject: [PATCH 03/22] feat(tui): add active provider controls and PR sheet --- .../adapters/in/tui/active_provider_test.go | 181 +++++++++++++++++ .../adapters/in/tui/component/help_pane.go | 13 +- .../adapters/in/tui/component/statusbar.go | 94 ++++++++- .../in/tui/component/statusbar_test.go | 122 ++++++++++++ internal/adapters/in/tui/help_pane_test.go | 3 + internal/adapters/in/tui/keymap/action.go | 60 +++--- .../adapters/in/tui/keymap/action_test.go | 4 + internal/adapters/in/tui/model.go | 133 +++++++++++-- internal/adapters/in/tui/pr_sheet.go | 123 ++++++++++++ internal/adapters/in/tui/pr_sheet_test.go | 95 +++++++++ internal/adapters/in/tui/provider_picker.go | 186 ++++++++++++++++++ .../adapters/in/tui/provider_picker_test.go | 89 +++++++++ internal/adapters/in/tui/review_providers.go | 109 ++++++++++ internal/adapters/in/tui/review_publish.go | 5 + internal/app/active_provider_service.go | 26 ++- internal/app/app.go | 14 +- internal/app/tui_active_provider.go | 77 ++++++++ 17 files changed, 1273 insertions(+), 61 deletions(-) create mode 100644 internal/adapters/in/tui/active_provider_test.go create mode 100644 internal/adapters/in/tui/component/statusbar_test.go create mode 100644 internal/adapters/in/tui/pr_sheet.go create mode 100644 internal/adapters/in/tui/pr_sheet_test.go create mode 100644 internal/adapters/in/tui/provider_picker.go create mode 100644 internal/adapters/in/tui/provider_picker_test.go create mode 100644 internal/app/tui_active_provider.go diff --git a/internal/adapters/in/tui/active_provider_test.go b/internal/adapters/in/tui/active_provider_test.go new file mode 100644 index 0000000..6e3b43b --- /dev/null +++ b/internal/adapters/in/tui/active_provider_test.go @@ -0,0 +1,181 @@ +package tui + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/require" + + "ero/internal/core" + "ero/internal/ports" +) + +type fakeActiveProviderController struct { + catalog []ports.ReviewProviderDescriptor + startState ActiveProviderState + refreshState ActiveProviderState + switchStates map[string]ActiveProviderState + publishResult core.PublishReviewResult + startErr error + refreshErr error + switchErrs map[string]error + publishErr error + startCalls int + refreshManual []bool + switchKeys []string + publishRequests []core.PublishReviewRequest + closed bool +} + +func (f *fakeActiveProviderController) Catalog(context.Context) ([]ports.ReviewProviderDescriptor, error) { + return append([]ports.ReviewProviderDescriptor(nil), f.catalog...), nil +} +func (f *fakeActiveProviderController) Start(context.Context, core.ReviewContext) (ActiveProviderState, error) { + f.startCalls++ + return f.startState, f.startErr +} +func (f *fakeActiveProviderController) Refresh(_ context.Context, _ core.ReviewContext, manual bool) (ActiveProviderState, error) { + f.refreshManual = append(f.refreshManual, manual) + return f.refreshState, f.refreshErr +} +func (f *fakeActiveProviderController) Switch(_ context.Context, _ core.ReviewContext, key string) (ActiveProviderState, error) { + f.switchKeys = append(f.switchKeys, key) + if err := f.switchErrs[key]; err != nil { + return ActiveProviderState{}, err + } + return f.switchStates[key], nil +} +func (f *fakeActiveProviderController) PublishReview(_ context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) { + f.publishRequests = append(f.publishRequests, request) + return f.publishResult, f.publishErr +} +func (f *fakeActiveProviderController) Generation() int64 { + return int64(len(f.refreshManual) + len(f.switchKeys) + f.startCalls) +} +func (f *fakeActiveProviderController) CompleteTimer(ctx context.Context, review core.ReviewContext, _ int64) (ActiveProviderState, error) { + return f.Refresh(ctx, review, false) +} +func (f *fakeActiveProviderController) Close() error { f.closed = true; return nil } + +func TestActiveProviderStartupLoadsOnlyActiveProviderState(t *testing.T) { + controller := &fakeActiveProviderController{ + catalog: []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "other", Label: "Other"}}, + startState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github-runtime", RuntimeInfo: core.ReviewProviderInfo{ID: "github-runtime", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ProviderID: "github-runtime", ExternalID: "t1"}}, Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusSynced}}}, + } + m := NewModelWithActiveProviderContext(context.Background(), []core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, []ports.ReviewProviderClient{fakeProviderAsPort{&fakeReviewProvider{}}}) + + cmd := m.Init() + require.NotNil(t, cmd) + updated, refreshCmd := m.Update(cmd()) + m = updated.(Model) + + require.NotNil(t, refreshCmd) + require.Equal(t, 1, controller.startCalls) + require.Len(t, m.providerCatalog, 2) + require.Equal(t, "github", m.activeProviderKey) + require.Equal(t, "github-runtime", m.activeRuntimeID) + require.Equal(t, core.ProviderSyncStatusSynced, m.providerSyncState.Status) + require.Len(t, m.remoteThreads, 1) + require.Empty(t, m.providerInfoByClient) +} + +func TestProviderStartupRefreshesRemoteThreadsAfterCacheState(t *testing.T) { + controller := &fakeActiveProviderController{ + catalog: []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}}, + startState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "cached"}}}}, + refreshState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "fresh"}}}}, + } + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + + started, refreshCmd := m.Update(m.Init()()) + m = started.(Model) + require.Len(t, m.remoteThreads, 1) + require.Equal(t, "cached", m.remoteThreads[0].ExternalID) + require.NotNil(t, refreshCmd) + + refreshed, _ := m.Update(refreshCmd()) + m = refreshed.(Model) + require.Equal(t, []bool{false}, controller.refreshManual) + require.Len(t, m.remoteThreads, 1) + require.Equal(t, "fresh", m.remoteThreads[0].ExternalID) +} + +func TestProviderRefreshManualReplacesRemoteThreads(t *testing.T) { + controller := &fakeActiveProviderController{refreshState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "new"}}}}} + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + m.remoteThreads = []core.RemoteReviewThread{{ExternalID: "old"}} + + msg := m.refreshActiveProviderCmd(true)() + updated, _ := m.Update(msg) + m = updated.(Model) + + require.Equal(t, []bool{true}, controller.refreshManual) + require.Len(t, m.remoteThreads, 1) + require.Equal(t, "new", m.remoteThreads[0].ExternalID) +} + +func TestProviderSwitchReplacesRemoteData(t *testing.T) { + controller := &fakeActiveProviderController{switchStates: map[string]ActiveProviderState{"other": {StableProviderKey: "other", RuntimeProviderID: "other-runtime", RuntimeInfo: core.ReviewProviderInfo{ID: "other-runtime"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "other-thread"}}}}}, switchErrs: map[string]error{}} + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + m.remoteThreads = []core.RemoteReviewThread{{ExternalID: "old"}} + + msg := m.switchActiveProviderCmd("other")() + updated, _ := m.Update(msg) + m = updated.(Model) + + require.Equal(t, []string{"other"}, controller.switchKeys) + require.Equal(t, "other", m.activeProviderKey) + require.Len(t, m.remoteThreads, 1) + require.Equal(t, "other-thread", m.remoteThreads[0].ExternalID) +} + +func TestActiveProviderPollTimerRefreshesWithGeneration(t *testing.T) { + controller := &fakeActiveProviderController{refreshState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "polled"}}}}} + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + + msg := m.completeActiveProviderTimerCmd(42)() + updated, _ := m.Update(msg) + m = updated.(Model) + + require.Equal(t, []bool{false}, controller.refreshManual) + require.Len(t, m.remoteThreads, 1) + require.Equal(t, "polled", m.remoteThreads[0].ExternalID) +} + +func TestActiveProviderPublishUsesActiveProviderClient(t *testing.T) { + controller := &fakeActiveProviderController{publishResult: core.PublishReviewResult{ProviderID: "github", ExternalReviewID: "review-1"}} + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + m.activeRuntimeInfo = core.ReviewProviderInfo{ID: "github", Label: "GitHub", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}} + m.activeRuntimeID = "github" + m.providerInfos = []core.ReviewProviderInfo{m.activeRuntimeInfo} + + m, _ = m.openPublishReview() + updated, cmd := m.publishSelectedProviders() + m = updated + require.NotNil(t, cmd) + msg := cmd().(publishReviewCompletedMsg) + m, _ = m.handlePublishReviewCompleted(msg) + + require.Len(t, controller.publishRequests, 1) + require.Equal(t, "github", controller.publishRequests[0].ProviderID) + require.False(t, m.publish.active) +} + +func TestProviderSwitchFailureClearsRemoteThreads(t *testing.T) { + controller := &fakeActiveProviderController{switchStates: map[string]ActiveProviderState{}, switchErrs: map[string]error{"other": errors.New("auth")}} + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + m.activeProviderKey = "github" + m.activeRuntimeID = "github" + m.providerInfos = []core.ReviewProviderInfo{{ID: "github"}} + m.remoteThreads = []core.RemoteReviewThread{{ExternalID: "old"}} + + msg := m.switchActiveProviderCmd("other")() + updated, _ := m.Update(msg) + m = updated.(Model) + + require.Equal(t, "other", m.activeProviderKey) + require.Empty(t, m.activeRuntimeID) + require.Empty(t, m.providerInfos) + require.Empty(t, m.remoteThreads) +} diff --git a/internal/adapters/in/tui/component/help_pane.go b/internal/adapters/in/tui/component/help_pane.go index bd9eb9c..d2e7f16 100644 --- a/internal/adapters/in/tui/component/help_pane.go +++ b/internal/adapters/in/tui/component/help_pane.go @@ -19,6 +19,14 @@ func RenderHelpPane(width, height int, enterKeyLabel, commentSubmitKeyLabel stri renderHelpShortcut("f", "find file", contentWidth), renderHelpShortcut("/", "grep references", contentWidth), renderHelpShortcut("d", "switch diff mode", contentWidth), + renderHelpShortcut("g/G", "cycle/pick provider", contentWidth), + renderHelpShortcut("r/o", "refresh provider/PR sheet", contentWidth), + "", + theme.HelpSectionStyle.Render("Search"), + renderHelpShortcut("↑↓", "select result", contentWidth), + renderHelpShortcut(enterKeyLabel, "jump to result", contentWidth), + renderHelpShortcut("esc", "cancel search", contentWidth), + "", renderHelpShortcut("h/l", "previous/next file", contentWidth), renderHelpShortcut("a", "expand all context", contentWidth), renderHelpShortcut(enterKeyLabel, "expand more context", contentWidth), @@ -29,11 +37,6 @@ func RenderHelpPane(width, height int, enterKeyLabel, commentSubmitKeyLabel stri renderHelpShortcut("y/Y", "copy plain/rich", contentWidth), renderHelpShortcut("q", "quit", contentWidth), "", - theme.HelpSectionStyle.Render("Search"), - renderHelpShortcut("↑↓", "select result", contentWidth), - renderHelpShortcut(enterKeyLabel, "jump to result", contentWidth), - renderHelpShortcut("esc", "cancel search", contentWidth), - "", theme.HelpSectionStyle.Render("Comment editor"), renderHelpShortcut(enterKeyLabel, "new line", contentWidth), renderHelpShortcut(commentSubmitKeyLabel, "submit comment", contentWidth), diff --git a/internal/adapters/in/tui/component/statusbar.go b/internal/adapters/in/tui/component/statusbar.go index e01570e..aa903d8 100644 --- a/internal/adapters/in/tui/component/statusbar.go +++ b/internal/adapters/in/tui/component/statusbar.go @@ -3,21 +3,28 @@ package component import ( "fmt" "strings" + "time" "charm.land/lipgloss/v2" "github.com/charmbracelet/x/ansi" "ero/internal/adapters/in/tui/theme" + "ero/internal/core" ) type StatusModel struct { - AppName string - Mode string - FileCount int - ProviderCount int - CurrentFile string - Message string - ScrollPercent float64 + AppName string + Mode string + FileCount int + ProviderCount int + CurrentFile string + Message string + ScrollPercent float64 + ActiveProviderLabel string + ActiveRuntimeName string + ProviderSync core.ProviderSyncState + ShowNoProvider bool + NerdFont bool } type StatusBar struct { @@ -41,6 +48,9 @@ func (c StatusBar) Render(model StatusModel) string { if model.ProviderCount > 0 { segments = append(segments, statusSegment{style: theme.StatusInfoStyle, label: providerCountLabel(model.ProviderCount)}) } + if syncLabel := providerSyncLabel(model); syncLabel != "" { + segments = append(segments, statusSegment{style: theme.StatusInfoStyle, label: syncLabel}) + } prefix := renderStatusSegments(leftWidth, segments...) percent := renderStatusSegments(leftWidth-lipgloss.Width(prefix), statusSegment{style: theme.StatusInfoStyle, label: fmt.Sprintf("%3.0f%%", model.ScrollPercent*100)}) @@ -125,6 +135,76 @@ func providerCountLabel(count int) string { return fmt.Sprintf("%d providers", count) } +func providerSyncLabel(model StatusModel) string { + provider := strings.TrimSpace(model.ActiveProviderLabel) + if provider == "" { + if model.ShowNoProvider { + return "no provider" + } + return "" + } + if runtimeName := strings.TrimSpace(model.ActiveRuntimeName); runtimeName != "" && runtimeName != provider { + provider += "/" + runtimeName + } + + parts := []string{provider} + status := providerSyncStatusLabel(model.ProviderSync.Status) + if status != "" { + if model.NerdFont { + status = providerSyncStatusSymbol(model.ProviderSync.Status) + " " + status + } + parts = append(parts, status) + } + if model.ProviderSync.LastError != "" { + parts = append(parts, TruncateRunes(model.ProviderSync.LastError, 24)) + } + if model.ProviderSync.LastSyncAt != nil { + parts = append(parts, "last "+formatStatusTime(*model.ProviderSync.LastSyncAt)) + } + if model.ProviderSync.NextSyncAt != nil { + parts = append(parts, "next "+formatStatusTime(*model.ProviderSync.NextSyncAt)) + } + return strings.Join(parts, " ") +} + +func providerSyncStatusLabel(status core.ProviderSyncStatus) string { + switch status { + case core.ProviderSyncStatusLoadingCache: + return "cache" + case core.ProviderSyncStatusSyncing: + return "syncing" + case core.ProviderSyncStatusSynced: + return "synced" + case core.ProviderSyncStatusFailed: + return "failed" + case core.ProviderSyncStatusBackingOff: + return "backoff" + default: + return "" + } +} + +func providerSyncStatusSymbol(status core.ProviderSyncStatus) string { + switch status { + case core.ProviderSyncStatusLoadingCache: + return "󰃨" + case core.ProviderSyncStatusSyncing: + return "󰑓" + case core.ProviderSyncStatusSynced: + return "" + case core.ProviderSyncStatusFailed: + return "" + case core.ProviderSyncStatusBackingOff: + return "󰌾" + default: + return "󰓦" + } +} + +func formatStatusTime(value time.Time) string { + return value.UTC().Format("15:04") +} + func TruncateRunes(value string, width int) string { if width <= 0 { return "" diff --git a/internal/adapters/in/tui/component/statusbar_test.go b/internal/adapters/in/tui/component/statusbar_test.go new file mode 100644 index 0000000..b714c2c --- /dev/null +++ b/internal/adapters/in/tui/component/statusbar_test.go @@ -0,0 +1,122 @@ +package component + +import ( + "regexp" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "ero/internal/core" +) + +func TestStatusbarProviderSync(t *testing.T) { + baseTime := time.Date(2026, 6, 5, 12, 34, 0, 0, time.UTC) + nextTime := baseTime.Add(5 * time.Minute) + + tests := []struct { + name string + model StatusModel + want []string + wantNot []string + }{ + { + name: "no provider", + model: noProviderStatusModel(), + want: []string{"ero", "branch", "1 file", "no provider"}, + wantNot: []string{"GitHub", "synced", "syncing", "cache", "failed", "backoff"}, + }, + { + name: "cache", + model: syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: core.ProviderSyncStatusLoadingCache}), + want: []string{"GitHub/gh-runtime", "cache"}, + }, + { + name: "syncing", + model: syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: core.ProviderSyncStatusSyncing}), + want: []string{"GitHub/gh-runtime", "syncing"}, + }, + { + name: "synced", + model: syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: core.ProviderSyncStatusSynced, LastSyncAt: &baseTime, NextSyncAt: &nextTime}), + want: []string{"GitHub/gh-runtime", "synced", "last 12:34", "next 12:39"}, + }, + { + name: "failed", + model: syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: core.ProviderSyncStatusFailed, LastSyncAt: &baseTime, LastError: "boom"}), + want: []string{"GitHub/gh-runtime", "failed", "boom", "last 12:34"}, + }, + { + name: "backing-off", + model: syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: core.ProviderSyncStatusBackingOff, LastSyncAt: &baseTime, NextSyncAt: &nextTime}), + want: []string{"GitHub/gh-runtime", "backoff", "last 12:34", "next 12:39"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + view := stripANSIForStatusbarTest(NewStatusBar(120).Render(tt.model)) + for _, want := range tt.want { + require.Contains(t, view, want) + } + for _, wantNot := range tt.wantNot { + require.NotContains(t, view, wantNot) + } + }) + } +} + +func TestStatusbarProviderSyncUsesNerdFontSymbolWhenSupported(t *testing.T) { + for status, symbol := range map[core.ProviderSyncStatus]string{ + core.ProviderSyncStatusLoadingCache: "󰃨", + core.ProviderSyncStatusSyncing: "󰑓", + core.ProviderSyncStatusSynced: "", + core.ProviderSyncStatusFailed: "", + core.ProviderSyncStatusBackingOff: "󰌾", + } { + model := syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: status}) + model.NerdFont = true + + view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) + + require.Contains(t, view, symbol) + } +} + +func TestStatusbarProviderSyncNarrowWidthDegradesGracefully(t *testing.T) { + last := time.Date(2026, 6, 5, 12, 34, 0, 0, time.UTC) + next := last.Add(5 * time.Minute) + model := syncStatusModel("VeryLongProvider", "very-long-runtime-name", core.ProviderSyncState{Status: core.ProviderSyncStatusBackingOff, LastSyncAt: &last, NextSyncAt: &next}) + + view := stripANSIForStatusbarTest(NewStatusBar(32).Render(model)) + + require.NotContains(t, view, "\n") + require.LessOrEqual(t, len([]rune(view)), 32) + require.True(t, strings.Contains(view, "Very") || strings.Contains(view, "back") || strings.Contains(view, "? help")) +} + +func baseStatusModel() StatusModel { + return StatusModel{AppName: "ero", Mode: "branch", FileCount: 1, CurrentFile: "demo.go"} +} + +func noProviderStatusModel() StatusModel { + model := baseStatusModel() + model.ShowNoProvider = true + return model +} + +func syncStatusModel(label, runtime string, sync core.ProviderSyncState) StatusModel { + model := baseStatusModel() + model.ProviderCount = 1 + model.ActiveProviderLabel = label + model.ActiveRuntimeName = runtime + model.ProviderSync = sync + return model +} + +var statusbarANSIPattern = regexp.MustCompile(`\x1b\[[0-9;:]*[A-Za-z]`) + +func stripANSIForStatusbarTest(s string) string { + return statusbarANSIPattern.ReplaceAllString(s, "") +} diff --git a/internal/adapters/in/tui/help_pane_test.go b/internal/adapters/in/tui/help_pane_test.go index b10c0ca..e35d0ce 100644 --- a/internal/adapters/in/tui/help_pane_test.go +++ b/internal/adapters/in/tui/help_pane_test.go @@ -24,6 +24,9 @@ func TestModelHelpModalShowsShortcutsAndCloses(t *testing.T) { assert.Contains(t, view, "select result") assert.Contains(t, view, "a expand all") assert.Contains(t, view, "expand more") + assert.Contains(t, view, "cycle/pick provider") + assert.Contains(t, view, "refresh provider") + assert.Contains(t, view, "PR sheet") assert.NotContains(t, view, "a/b") updated, _ = model.Update(tea.KeyPressMsg{Code: tea.KeyEsc}) diff --git a/internal/adapters/in/tui/keymap/action.go b/internal/adapters/in/tui/keymap/action.go index 16028e0..403138d 100644 --- a/internal/adapters/in/tui/keymap/action.go +++ b/internal/adapters/in/tui/keymap/action.go @@ -3,30 +3,34 @@ package keymap type Action string const ( - ActionNone Action = "" - ActionQuit Action = "quit" - ActionMoveUp Action = "move_up" - ActionMoveDown Action = "move_down" - ActionPageUp Action = "page_up" - ActionPageDown Action = "page_down" - ActionMoveStart Action = "move_start" - ActionMoveEnd Action = "move_end" - ActionToggleSelection Action = "toggle_selection" - ActionClearSelection Action = "clear_selection" - ActionOpenComment Action = "open_comment" - ActionClearReview Action = "clear_review" - ActionPublishReview Action = "publish_review" - ActionCopyReviewJSON Action = "copy_review_json" - ActionCopyPlain Action = "copy_plain" - ActionCopyWithMetadata Action = "copy_with_metadata" - ActionOpenFileSearch Action = "open_file_search" - ActionOpenGrepSearch Action = "open_grep_search" - ActionOpenDiffMode Action = "open_diff_mode" - ActionPreviousFile Action = "previous_file" - ActionNextFile Action = "next_file" - ActionExpandAllContext Action = "expand_all_context" - ActionExpandMoreContext Action = "expand_more_context" - ActionOpenHelp Action = "open_help" + ActionNone Action = "" + ActionQuit Action = "quit" + ActionMoveUp Action = "move_up" + ActionMoveDown Action = "move_down" + ActionPageUp Action = "page_up" + ActionPageDown Action = "page_down" + ActionMoveStart Action = "move_start" + ActionMoveEnd Action = "move_end" + ActionToggleSelection Action = "toggle_selection" + ActionClearSelection Action = "clear_selection" + ActionOpenComment Action = "open_comment" + ActionClearReview Action = "clear_review" + ActionPublishReview Action = "publish_review" + ActionCopyReviewJSON Action = "copy_review_json" + ActionCopyPlain Action = "copy_plain" + ActionCopyWithMetadata Action = "copy_with_metadata" + ActionOpenFileSearch Action = "open_file_search" + ActionOpenGrepSearch Action = "open_grep_search" + ActionOpenDiffMode Action = "open_diff_mode" + ActionPreviousFile Action = "previous_file" + ActionNextFile Action = "next_file" + ActionExpandAllContext Action = "expand_all_context" + ActionExpandMoreContext Action = "expand_more_context" + ActionCycleProvider Action = "cycle_provider" + ActionOpenProviderPicker Action = "open_provider_picker" + ActionRefreshProvider Action = "refresh_provider" + ActionTogglePRSheet Action = "toggle_pr_sheet" + ActionOpenHelp Action = "open_help" ) func ReviewAction(key string) Action { @@ -75,6 +79,14 @@ func ReviewAction(key string) Action { return ActionExpandAllContext case "enter": return ActionExpandMoreContext + case "g": + return ActionCycleProvider + case "G": + return ActionOpenProviderPicker + case "r": + return ActionRefreshProvider + case "o": + return ActionTogglePRSheet case "?": return ActionOpenHelp default: diff --git a/internal/adapters/in/tui/keymap/action_test.go b/internal/adapters/in/tui/keymap/action_test.go index bff27e0..243a612 100644 --- a/internal/adapters/in/tui/keymap/action_test.go +++ b/internal/adapters/in/tui/keymap/action_test.go @@ -44,6 +44,10 @@ func TestReviewAction(t *testing.T) { {name: "next file n alias", key: "n", want: ActionNextFile}, {name: "expand all context", key: "a", want: ActionExpandAllContext}, {name: "expand more context enter binding", key: "enter", want: ActionExpandMoreContext}, + {name: "cycle provider", key: "g", want: ActionCycleProvider}, + {name: "open provider picker", key: "G", want: ActionOpenProviderPicker}, + {name: "refresh provider", key: "r", want: ActionRefreshProvider}, + {name: "toggle pr sheet", key: "o", want: ActionTogglePRSheet}, {name: "open help", key: "?", want: ActionOpenHelp}, {name: "unknown", key: "R", want: ActionNone}, } diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index aca7cca..d7c8e6f 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -51,13 +51,47 @@ type clipboardCopyFailedMsg struct { err error } -type reviewProvidersLoadedMsg struct { - infos []core.ReviewProviderInfo - threads []core.RemoteReviewThread - clients map[ports.ReviewProviderClient]core.ReviewProviderInfo - errs []string +type activeProviderController interface { + Catalog(context.Context) ([]ports.ReviewProviderDescriptor, error) + Start(context.Context, core.ReviewContext) (ActiveProviderState, error) + Refresh(context.Context, core.ReviewContext, bool) (ActiveProviderState, error) + Switch(context.Context, core.ReviewContext, string) (ActiveProviderState, error) + PublishReview(context.Context, core.PublishReviewRequest) (core.PublishReviewResult, error) + Generation() int64 + CompleteTimer(context.Context, core.ReviewContext, int64) (ActiveProviderState, error) + Close() error } +// ActiveProviderState is the TUI-facing active provider snapshot. +type ActiveProviderState struct { + StableProviderKey string + RuntimeProviderID string + RuntimeInfo core.ReviewProviderInfo + Snapshot core.ProviderSnapshot + FromCache bool + Syncing bool + LastError error +} + +type activeProviderStartedMsg struct { + catalog []ports.ReviewProviderDescriptor + state ActiveProviderState + err error +} + +type activeProviderRefreshedMsg struct { + state ActiveProviderState + err error +} + +type activeProviderSwitchedMsg struct { + stableKey string + state ActiveProviderState + err error +} + +type activeProviderPollDueMsg struct{ generation int64 } + type Model struct { title string files []core.ReviewFile @@ -89,10 +123,19 @@ type Model struct { commentEditor *InlineCommentEditor reviewContext core.ReviewContext reviewProviders []ports.ReviewProviderClient + activeProvider activeProviderController + providerCatalog []ports.ReviewProviderDescriptor + activeProviderKey string + activeRuntimeID string + activeRuntimeInfo core.ReviewProviderInfo + providerSyncState core.ProviderSyncState + providerOverview *core.ProviderOverview remoteThreads []core.RemoteReviewThread providerInfos []core.ReviewProviderInfo providerInfoByClient map[ports.ReviewProviderClient]core.ReviewProviderInfo + providerPicker providerPickerState publish publishState + prSheet prSheetState ctx context.Context reviewLineCache *render.ReviewLineCache cachedEditorWidth int @@ -120,6 +163,10 @@ func NewModelWithReviewProviders(files []core.ReviewFile, terminal ports.Termina } func NewModelWithReviewProvidersContext(ctx context.Context, files []core.ReviewFile, terminal ports.Terminal, loader reviewLoader, request core.ReviewRequest, clipboardWriter ports.ClipboardWriter, reviewContext core.ReviewContext, providers []ports.ReviewProviderClient) Model { + return NewModelWithActiveProviderContext(ctx, files, terminal, loader, request, clipboardWriter, reviewContext, nil, providers) +} + +func NewModelWithActiveProviderContext(ctx context.Context, files []core.ReviewFile, terminal ports.Terminal, loader reviewLoader, request core.ReviewRequest, clipboardWriter ports.ClipboardWriter, reviewContext core.ReviewContext, activeProvider activeProviderController, providers []ports.ReviewProviderClient) Model { if ctx == nil { ctx = context.Background() } @@ -142,6 +189,7 @@ func NewModelWithReviewProvidersContext(ctx context.Context, files []core.Review reviewDraft: core.NewReviewDraft(), reviewContext: reviewContext, reviewProviders: append([]ports.ReviewProviderClient(nil), providers...), + activeProvider: activeProvider, providerInfos: nil, providerInfoByClient: map[ports.ReviewProviderClient]core.ReviewProviderInfo{}, remoteThreads: nil, @@ -157,6 +205,9 @@ func NewModelWithReviewProvidersContext(ctx context.Context, files []core.Review } func (m Model) Init() tea.Cmd { + if m.activeProvider != nil { + return m.startActiveProviderCmd() + } if len(m.reviewProviders) == 0 { return nil } @@ -213,6 +264,42 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case clipboardCopyFailedMsg: m.setCopyFeedback("Copy failed: " + msg.err.Error()) return m, m.expireCopyFeedbackCmd() + case prSheetToggledMsg: + return m.TogglePRSheet(), nil + case prSheetScrolledMsg: + return m.ScrollPRSheet(msg.delta), nil + case activeProviderStartedMsg: + m.providerCatalog = msg.catalog + m.applyActiveProviderState(msg.state) + if msg.err != nil { + m.setCopyFeedback("Provider unavailable: " + msg.err.Error()) + m.syncReviewViewport() + return m, m.expireCopyFeedbackCmd() + } + m.syncReviewViewport() + return m, m.refreshActiveProviderCmd(false) + case activeProviderRefreshedMsg: + m.applyActiveProviderState(msg.state) + if msg.err != nil { + m.setCopyFeedback("Provider refresh failed: " + msg.err.Error()) + m.syncReviewViewport() + return m, tea.Batch(m.expireCopyFeedbackCmd(), m.scheduleActiveProviderPollCmd()) + } + m.syncReviewViewport() + return m, m.scheduleActiveProviderPollCmd() + case activeProviderSwitchedMsg: + if msg.err != nil { + m.clearActiveProviderRemoteData() + m.activeProviderKey = msg.stableKey + m.setCopyFeedback("Provider switch failed: " + msg.err.Error()) + m.syncReviewViewport() + return m, m.expireCopyFeedbackCmd() + } + m.applyActiveProviderState(msg.state) + m.syncReviewViewport() + return m, m.refreshActiveProviderCmd(false) + case activeProviderPollDueMsg: + return m, m.completeActiveProviderTimerCmd(msg.generation) case reviewProvidersLoadedMsg: m.providerInfos = msg.infos m.remoteThreads = msg.threads @@ -249,6 +336,9 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { if m.search.active() { return m.updateSearch(msg) } + if m.providerPicker.open { + return m.updateProviderPicker(msg) + } if m.publish.active { return m.updatePublishReview(msg) } @@ -304,6 +394,14 @@ func (m Model) updateReviewAction(action keymap.Action) (tea.Model, tea.Cmd) { m.showAllContext() case keymap.ActionExpandMoreContext: m.showMoreContext(contextStep) + case keymap.ActionCycleProvider: + return m, m.cycleProviderCmd() + case keymap.ActionOpenProviderPicker: + m = m.openProviderPicker() + case keymap.ActionRefreshProvider: + return m, m.refreshActiveProviderCmd(true) + case keymap.ActionTogglePRSheet: + m = m.TogglePRSheet() case keymap.ActionOpenHelp: m.helpActive = true case keymap.ActionNone: @@ -323,13 +421,18 @@ func (m Model) View() tea.View { content := lipgloss.JoinVertical(lipgloss.Left, review, NewStatusBar(m.width).Render(StatusModel{ - AppName: m.title, - Mode: diffModeLabel(m.diffMode, m.nerdFont), - FileCount: len(m.files), - ProviderCount: len(m.providerInfos), - CurrentFile: m.activeLocation(), - Message: m.copyFeedback, - ScrollPercent: m.reviewViewport.ScrollPercent(), + AppName: m.title, + Mode: diffModeLabel(m.diffMode, m.nerdFont), + FileCount: len(m.files), + ProviderCount: m.statusProviderCount(), + CurrentFile: m.activeLocation(), + Message: m.copyFeedback, + ScrollPercent: m.reviewViewport.ScrollPercent(), + ActiveProviderLabel: m.activeRuntimeInfo.Label, + ActiveRuntimeName: m.activeRuntimeID, + ProviderSync: m.providerSyncState, + ShowNoProvider: m.activeProvider != nil && m.activeProviderKey == "" && m.activeRuntimeInfo.ID == "", + NerdFont: m.nerdFont, }), ) if m.search.active() { @@ -338,6 +441,12 @@ func (m Model) View() tea.View { if m.publish.active { content = m.renderPublishOverlay(content) } + if m.providerPicker.open { + content = m.renderProviderPickerOverlay(content) + } + if m.prSheet.open { + content = m.renderPRSheetOverlay(content) + } if m.helpActive { content = m.renderHelpOverlay(content) } diff --git a/internal/adapters/in/tui/pr_sheet.go b/internal/adapters/in/tui/pr_sheet.go new file mode 100644 index 0000000..1ecc196 --- /dev/null +++ b/internal/adapters/in/tui/pr_sheet.go @@ -0,0 +1,123 @@ +package tui + +import ( + "strconv" + "strings" + + tea "charm.land/bubbletea/v2" + "charm.land/lipgloss/v2" +) + +type prSheetToggledMsg struct{} + +type prSheetScrolledMsg struct { + delta int +} + +type prSheetState struct { + open bool + yOffset int +} + +func (m Model) TogglePRSheet() Model { + m.prSheet.open = !m.prSheet.open + return m +} + +func (m Model) ScrollPRSheet(delta int) Model { + m.prSheet.yOffset = clampPRSheetOffset(m.prSheet.yOffset+delta, m.prSheetLineCount()) + return m +} + +func (m Model) renderPRSheetOverlay(content string) string { + width := max(m.width, 1) + height := max(m.height, 1) + pane := m.renderPRSheet(width, height) + paneWidth := lipgloss.Width(pane) + + canvas := lipgloss.NewCanvas(width, height) + compositor := lipgloss.NewCompositor( + lipgloss.NewLayer(content), + lipgloss.NewLayer(pane).X(max(width-paneWidth, 0)).Y(0).Z(1), + ) + canvas.Compose(compositor) + return canvas.Render() +} + +func (m Model) renderPRSheet(width, height int) string { + sheetWidth := prSheetWidth(width) + contentWidth := max(sheetWidth-2, 1) + lines := m.prSheetLines() + m.prSheet.yOffset = clampPRSheetOffset(m.prSheet.yOffset, len(lines)) + visibleLines := visiblePRSheetLines(lines, m.prSheet.yOffset, height) + + rows := make([]string, height) + for i := range height { + text := "" + if i < len(visibleLines) { + text = visibleLines[i] + } + row := "│ " + truncatePlainRow(text, contentWidth) + rows[i] = padRight(row, sheetWidth) + } + return strings.Join(rows, "\n") +} + +func (m Model) prSheetLines() []string { + provider := "No active provider" + if m.activeRuntimeInfo.ID != "" || m.activeRuntimeInfo.Label != "" || m.activeRuntimeInfo.Name != "" { + provider = providerDisplayLabel(m.activeRuntimeInfo) + } + return []string{ + "Pull request", + "", + "Provider: " + provider, + "", + "Overview placeholder", + "Phase 4 will render PR markdown and metadata here.", + "", + "Remote threads: " + pluralCount(len(m.remoteThreads), "thread"), + } +} + +func (m Model) prSheetLineCount() int { + return len(m.prSheetLines()) +} + +func prSheetWidth(totalWidth int) int { + if totalWidth <= 1 { + return 1 + } + return min(max(totalWidth/3, 32), totalWidth) +} + +func visiblePRSheetLines(lines []string, offset, height int) []string { + if height <= 0 || len(lines) == 0 { + return nil + } + offset = clampPRSheetOffset(offset, len(lines)) + end := min(offset+height, len(lines)) + return lines[offset:end] +} + +func clampPRSheetOffset(offset, lineCount int) int { + if lineCount <= 0 { + return 0 + } + return min(max(offset, 0), lineCount-1) +} + +func pluralCount(count int, singular string) string { + return strconv.Itoa(count) + " " + pluralize(singular, count) +} + +func padRight(s string, width int) string { + if lipgloss.Width(s) >= width { + return s + } + return s + strings.Repeat(" ", width-lipgloss.Width(s)) +} + +func togglePRSheetCmd() tea.Cmd { + return func() tea.Msg { return prSheetToggledMsg{} } +} diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go new file mode 100644 index 0000000..6666639 --- /dev/null +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -0,0 +1,95 @@ +package tui + +import ( + "strings" + "testing" + + tea "charm.land/bubbletea/v2" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "ero/internal/core" +) + +func TestPRSheetOverlaysRightSideFullHeightWithLeftSeparatorOnly(t *testing.T) { + t.Parallel() + + model := NewModel([]core.ReviewFile{reviewFile("demo.go", "package demo")}) + updated, _ := model.Update(tea.WindowSizeMsg{Width: 60, Height: 8}) + model = updated.(Model).TogglePRSheet() + + view := stripANSI(model.View().Content) + lines := strings.Split(view, "\n") + require.Len(t, lines, 8) + + sheetWidth := prSheetWidth(model.width) + separatorColumn := model.width - sheetWidth + separatorRows := 0 + for _, line := range lines { + runes := []rune(line) + if len(runes) <= separatorColumn { + continue + } + if runes[separatorColumn] == '│' { + separatorRows++ + } + if len(runes) > separatorColumn+1 { + assert.NotEqual(t, '│', runes[len(runes)-1], "right edge should not be bordered") + } + } + assert.GreaterOrEqual(t, separatorRows, 6) + assert.Contains(t, view, "Pull request") + assert.NotContains(t, view, "┌") + assert.NotContains(t, view, "┐") + assert.NotContains(t, view, "└") + assert.NotContains(t, view, "┘") +} + +func TestPRSheetOverlayDoesNotReflowUnderlyingDiffContent(t *testing.T) { + t.Parallel() + + model := NewModel([]core.ReviewFile{reviewFile("demo.go", "package demo")}) + updated, _ := model.Update(tea.WindowSizeMsg{Width: 80, Height: 10}) + model = updated.(Model) + closedReviewWidth := model.reviewViewport.Width() + closedFirstLine := strings.Split(stripANSI(model.View().Content), "\n")[0] + + model = model.TogglePRSheet() + openReviewWidth := model.reviewViewport.Width() + openFirstLine := strings.Split(stripANSI(model.View().Content), "\n")[0] + + assert.Equal(t, closedReviewWidth, openReviewWidth) + sheetStart := model.width - prSheetWidth(model.width) + assert.Equal(t, string([]rune(closedFirstLine)[:sheetStart]), string([]rune(openFirstLine)[:sheetStart])) +} + +func TestPRSheetCanToggleByMethodAndMessage(t *testing.T) { + t.Parallel() + + model := NewModel(nil) + assert.False(t, model.prSheet.open) + model = model.TogglePRSheet() + assert.True(t, model.prSheet.open) + + updated, _ := model.Update(prSheetToggledMsg{}) + model = updated.(Model) + assert.False(t, model.prSheet.open) +} + +func TestPRSheetHasIndependentScrollState(t *testing.T) { + t.Parallel() + + model := NewModel([]core.ReviewFile{reviewFileWithLines("demo.go", 20)}) + updated, _ := model.Update(tea.WindowSizeMsg{Width: 80, Height: 4}) + model = updated.(Model).TogglePRSheet() + model.moveCursor(5) + reviewOffset := model.reviewViewport.YOffset() + + updated, _ = model.Update(prSheetScrolledMsg{delta: 2}) + model = updated.(Model) + + assert.Equal(t, 2, model.prSheet.yOffset) + assert.Equal(t, reviewOffset, model.reviewViewport.YOffset()) + assert.Contains(t, stripANSI(model.renderPRSheet(model.width, model.height)), "Provider:") + assert.NotContains(t, stripANSI(model.renderPRSheet(model.width, model.height)), "Pull request") +} diff --git a/internal/adapters/in/tui/provider_picker.go b/internal/adapters/in/tui/provider_picker.go new file mode 100644 index 0000000..042d3d4 --- /dev/null +++ b/internal/adapters/in/tui/provider_picker.go @@ -0,0 +1,186 @@ +package tui + +import ( + "fmt" + "strings" + + tea "charm.land/bubbletea/v2" + "charm.land/lipgloss/v2" + + "ero/internal/adapters/in/tui/theme" + "ero/internal/ports" +) + +type providerPickerState struct { + open bool + selected int + rows []providerPickerRow +} + +type providerPickerRow struct { + Key string + Label string + PluginName string + PluginSource string + Reason string + Active bool +} + +func (m Model) openProviderPicker() Model { + m.providerPicker.open = true + m.providerPicker.rows = m.providerPickerRows() + m.providerPicker.selected = clampProviderPickerSelection(m.providerPicker.selected, len(m.providerPicker.rows)) + for i, row := range m.providerPicker.rows { + if row.Active { + m.providerPicker.selected = i + break + } + } + return m +} + +func (m Model) closeProviderPicker() Model { + m.providerPicker.open = false + return m +} + +func (m Model) updateProviderPicker(msg tea.KeyPressMsg) (tea.Model, tea.Cmd) { + switch msg.String() { + case "esc": + return m.closeProviderPicker(), nil + case "up", "k": + m.providerPicker.selected = clampProviderPickerSelection(m.providerPicker.selected-1, len(m.providerPicker.rows)) + return m, nil + case "down", "j": + m.providerPicker.selected = clampProviderPickerSelection(m.providerPicker.selected+1, len(m.providerPicker.rows)) + return m, nil + case "enter": + if len(m.providerPicker.rows) == 0 { + return m.closeProviderPicker(), nil + } + key := m.providerPicker.rows[m.providerPicker.selected].Key + m = m.closeProviderPicker() + return m, m.switchActiveProviderCmd(key) + default: + return m, nil + } +} + +func (m Model) cycleProviderCmd() tea.Cmd { + rows := m.providerPickerRows() + if len(rows) == 0 { + return nil + } + next := 0 + for i, row := range rows { + if row.Active { + next = (i + 1) % len(rows) + break + } + } + return m.switchActiveProviderCmd(rows[next].Key) +} + +func (m Model) providerPickerRows() []providerPickerRow { + rows := make([]providerPickerRow, 0, len(m.providerCatalog)) + for _, descriptor := range m.providerCatalog { + rows = append(rows, m.providerPickerRow(descriptor)) + } + return rows +} + +func (m Model) providerPickerRow(descriptor ports.ReviewProviderDescriptor) providerPickerRow { + label := descriptor.Label + if label == "" { + label = descriptor.ContributionID + } + if label == "" { + label = descriptor.Key + } + reason := "" + if descriptor.Key == m.activeProviderKey && m.providerSyncState.LastError != "" { + reason = m.providerSyncState.LastError + } + return providerPickerRow{ + Key: descriptor.Key, + Label: label, + PluginName: descriptor.PluginName, + PluginSource: descriptor.PluginSource, + Reason: reason, + Active: descriptor.Key == m.activeProviderKey, + } +} + +func clampProviderPickerSelection(selected, rowCount int) int { + if rowCount <= 0 || selected < 0 { + return 0 + } + if selected >= rowCount { + return rowCount - 1 + } + return selected +} + +func (m Model) renderProviderPickerOverlay(content string) string { + width := max(m.width, 1) + height := max(m.height, 1) + pane := m.renderProviderPicker(width, height) + return renderCenteredOverlay(content, pane, width, height, max((height-lipgloss.Height(pane))/2, 0)) +} + +func (m Model) renderProviderPicker(width, height int) string { + paneWidth := min(max(width-8, 36), 76) + contentWidth := max(paneWidth-6, 1) + lines := []string{theme.HelpPaneTitleStyle.Render("Review providers"), ""} + rows := m.providerPicker.rows + if len(rows) == 0 { + lines = append(lines, theme.MutedStyle.Render("No providers discovered")) + } else { + for i, row := range rows { + cursor := " " + if i == m.providerPicker.selected { + cursor = "> " + } + active := " " + if row.Active { + active = "*" + } + meta := strings.TrimSpace(strings.Join([]string{row.PluginName, row.PluginSource}, " ")) + line := fmt.Sprintf("%s%s %s", cursor, active, row.Label) + if meta != "" { + line += " — " + meta + } + if row.Reason != "" { + line += " (" + row.Reason + ")" + } + lines = append(lines, theme.HelpLabelStyle.Render(componentTruncate(line, contentWidth))) + } + } + lines = append(lines, "", theme.HelpLabelStyle.Render("enter switch • esc close")) + lines = fitProviderPickerLines(lines, max(height-2, 1)) + return theme.HelpPaneStyle.Width(paneWidth).Render(strings.Join(lines, "\n")) +} + +func fitProviderPickerLines(lines []string, maxLines int) []string { + if len(lines) <= maxLines { + return lines + } + if maxLines <= 1 { + return lines[:maxLines] + } + result := append([]string(nil), lines[:maxLines-1]...) + result = append(result, lines[len(lines)-1]) + return result +} + +func componentTruncate(s string, width int) string { + trimmed := strings.TrimRight(s, " ") + runes := []rune(trimmed) + if width < 0 { + width = 0 + } + if len(runes) <= width { + return trimmed + } + return string(runes[:width]) +} diff --git a/internal/adapters/in/tui/provider_picker_test.go b/internal/adapters/in/tui/provider_picker_test.go new file mode 100644 index 0000000..71af212 --- /dev/null +++ b/internal/adapters/in/tui/provider_picker_test.go @@ -0,0 +1,89 @@ +package tui + +import ( + "context" + "testing" + + tea "charm.land/bubbletea/v2" + "github.com/stretchr/testify/require" + + "ero/internal/core" + "ero/internal/ports" +) + +func TestProviderPickerDisplaysDescriptorRowsWithoutStartingInactiveProviders(t *testing.T) { + controller := &fakeActiveProviderController{ + catalog: []ports.ReviewProviderDescriptor{ + {Key: "github", Label: "GitHub", PluginName: "gh-plugin", PluginSource: "builtin"}, + {Key: "gitlab", Label: "GitLab", PluginName: "gl-plugin", PluginSource: "local"}, + }, + startState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusFailed, LastError: "missing token"}}}, + } + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + updated, _ := m.Update(m.Init()()) + m = updated.(Model) + + updated, _ = m.Update(keyPress("G")) + m = updated.(Model) + + require.True(t, m.providerPicker.open) + require.Equal(t, 1, controller.startCalls) + view := stripANSI(m.View().Content) + require.Contains(t, view, "Review providers") + require.Contains(t, view, "* GitHub") + require.Contains(t, view, "gh-plugin builtin") + require.Contains(t, view, "missing token") + require.Contains(t, view, "GitLab") + require.Contains(t, view, "gl-plugin local") +} + +func TestProviderPickerSelectEmitsSwitchCommandWithStableKey(t *testing.T) { + controller := &fakeActiveProviderController{ + catalog: []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "gitlab", Label: "GitLab"}}, + startState: ActiveProviderState{StableProviderKey: "github"}, + switchStates: map[string]ActiveProviderState{"gitlab": {StableProviderKey: "gitlab"}}, + switchErrs: map[string]error{}, + } + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + updated, _ := m.Update(m.Init()()) + m = updated.(Model) + updated, _ = m.Update(keyPress("G")) + m = updated.(Model) + updated, _ = m.Update(tea.KeyPressMsg{Code: tea.KeyDown}) + m = updated.(Model) + + updated, cmd := m.Update(tea.KeyPressMsg{Code: tea.KeyEnter}) + m = updated.(Model) + require.NotNil(t, cmd) + updated, _ = m.Update(cmd()) + m = updated.(Model) + + require.False(t, m.providerPicker.open) + require.Equal(t, []string{"gitlab"}, controller.switchKeys) + require.Equal(t, "gitlab", m.activeProviderKey) +} + +func TestProviderCycleAndRefreshShortcuts(t *testing.T) { + controller := &fakeActiveProviderController{ + catalog: []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "gitlab", Label: "GitLab"}}, + startState: ActiveProviderState{StableProviderKey: "github"}, + refreshState: ActiveProviderState{StableProviderKey: "github"}, + switchStates: map[string]ActiveProviderState{"gitlab": {StableProviderKey: "gitlab"}}, + switchErrs: map[string]error{}, + } + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + updated, _ := m.Update(m.Init()()) + m = updated.(Model) + + updated, cmd := m.Update(keyPress("g")) + m = updated.(Model) + require.NotNil(t, cmd) + updated, _ = m.Update(cmd()) + m = updated.(Model) + require.Equal(t, []string{"gitlab"}, controller.switchKeys) + + _, cmd = m.Update(keyPress("r")) + require.NotNil(t, cmd) + _ = cmd() + require.Equal(t, []bool{true}, controller.refreshManual) +} diff --git a/internal/adapters/in/tui/review_providers.go b/internal/adapters/in/tui/review_providers.go index 9a57589..3721c31 100644 --- a/internal/adapters/in/tui/review_providers.go +++ b/internal/adapters/in/tui/review_providers.go @@ -2,6 +2,7 @@ package tui import ( "fmt" + "time" tea "charm.land/bubbletea/v2" @@ -11,7 +12,18 @@ import ( "ero/internal/ports" ) +type reviewProvidersLoadedMsg struct { + infos []core.ReviewProviderInfo + threads []core.RemoteReviewThread + clients map[ports.ReviewProviderClient]core.ReviewProviderInfo + errs []string +} + func (m Model) closeReviewProvidersCmd() tea.Cmd { + if m.activeProvider != nil { + activeProvider := m.activeProvider + return func() tea.Msg { _ = activeProvider.Close(); return nil } + } providers := append([]ports.ReviewProviderClient(nil), m.reviewProviders...) return func() tea.Msg { for _, provider := range providers { @@ -21,6 +33,103 @@ func (m Model) closeReviewProvidersCmd() tea.Cmd { } } +func (m Model) startActiveProviderCmd() tea.Cmd { + activeProvider := m.activeProvider + if activeProvider == nil { + return nil + } + ctx := m.ctx + reviewContext := m.reviewContext + return func() tea.Msg { + catalog, catalogErr := activeProvider.Catalog(ctx) + state, err := activeProvider.Start(ctx, reviewContext) + if err == nil { + err = catalogErr + } + return activeProviderStartedMsg{catalog: catalog, state: state, err: err} + } +} + +func (m Model) refreshActiveProviderCmd(manual bool) tea.Cmd { + activeProvider := m.activeProvider + if activeProvider == nil { + return nil + } + ctx := m.ctx + reviewContext := m.reviewContext + return func() tea.Msg { + state, err := activeProvider.Refresh(ctx, reviewContext, manual) + return activeProviderRefreshedMsg{state: state, err: err} + } +} + +func (m Model) switchActiveProviderCmd(stableKey string) tea.Cmd { + activeProvider := m.activeProvider + if activeProvider == nil { + return nil + } + ctx := m.ctx + reviewContext := m.reviewContext + return func() tea.Msg { + state, err := activeProvider.Switch(ctx, reviewContext, stableKey) + return activeProviderSwitchedMsg{stableKey: stableKey, state: state, err: err} + } +} + +func (m Model) scheduleActiveProviderPollCmd() tea.Cmd { + if m.activeProvider == nil || m.providerSyncState.NextSyncAt == nil { + return nil + } + delay := max(time.Until(*m.providerSyncState.NextSyncAt), 0) + generation := m.activeProvider.Generation() + return tea.Tick(delay, func(time.Time) tea.Msg { return activeProviderPollDueMsg{generation: generation} }) +} + +func (m Model) completeActiveProviderTimerCmd(generation int64) tea.Cmd { + activeProvider := m.activeProvider + if activeProvider == nil { + return nil + } + ctx := m.ctx + reviewContext := m.reviewContext + return func() tea.Msg { + state, err := activeProvider.CompleteTimer(ctx, reviewContext, generation) + return activeProviderRefreshedMsg{state: state, err: err} + } +} + +func (m Model) statusProviderCount() int { + if m.activeProvider != nil { + return len(m.providerCatalog) + } + return len(m.providerInfos) +} + +func (m *Model) applyActiveProviderState(state ActiveProviderState) { + m.activeProviderKey = state.StableProviderKey + m.activeRuntimeID = state.RuntimeProviderID + m.activeRuntimeInfo = state.RuntimeInfo + m.providerSyncState = state.Snapshot.Sync + m.providerOverview = state.Snapshot.Overview + m.remoteThreads = append([]core.RemoteReviewThread(nil), state.Snapshot.Threads...) + if state.RuntimeInfo.ID != "" { + m.providerInfos = []core.ReviewProviderInfo{state.RuntimeInfo} + } else { + m.providerInfos = nil + } + m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{} +} + +func (m *Model) clearActiveProviderRemoteData() { + m.activeRuntimeID = "" + m.activeRuntimeInfo = core.ReviewProviderInfo{} + m.providerSyncState = core.ProviderSyncState{} + m.providerOverview = nil + m.remoteThreads = nil + m.providerInfos = nil + m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{} +} + func (m Model) loadReviewProvidersCmd() tea.Cmd { providers := make([]ports.ReviewProviderClient, len(m.reviewProviders)) copy(providers, m.reviewProviders) diff --git a/internal/adapters/in/tui/review_publish.go b/internal/adapters/in/tui/review_publish.go index f977881..fce8020 100644 --- a/internal/adapters/in/tui/review_publish.go +++ b/internal/adapters/in/tui/review_publish.go @@ -254,6 +254,11 @@ func (m Model) providerClientsFor(infos []core.ReviewProviderInfo) []providerCli selected[info.ID] = info } result := make([]providerClientWithInfo, 0, len(infos)) + if m.activeProvider != nil && m.activeRuntimeInfo.ID != "" { + if info, ok := selected[m.activeRuntimeInfo.ID]; ok { + return append(result, providerClientWithInfo{info: info, client: m.activeProvider}) + } + } for _, client := range m.reviewProviders { providerInfo, ok := m.providerInfoByClient[client] if !ok { diff --git a/internal/app/active_provider_service.go b/internal/app/active_provider_service.go index 9cc0fea..00f5380 100644 --- a/internal/app/active_provider_service.go +++ b/internal/app/active_provider_service.go @@ -19,6 +19,7 @@ func DefaultProviderPollingConfig() ProviderPollingConfig { type ActiveProviderState struct { StableProviderKey string RuntimeProviderID string + RuntimeInfo core.ReviewProviderInfo Snapshot core.ProviderSnapshot FromCache bool Syncing bool @@ -88,7 +89,7 @@ func (s *ActiveProviderService) Start(ctx context.Context, review core.ReviewCon if s.prefs != nil { _ = s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key) } - st := s.loadCachedState(ctx, review, d.Key, info.ID) + st := s.loadCachedState(ctx, review, d.Key, info.ID, info) s.setState(gen, st) return st, nil } @@ -130,7 +131,7 @@ func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewCo if s.prefs != nil { _ = s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key) } - st := s.loadCachedState(ctx, review, d.Key, info.ID) + st := s.loadCachedState(ctx, review, d.Key, info.ID, info) s.setState(gen, st) return st, nil } @@ -138,6 +139,20 @@ func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewCo return ActiveProviderState{}, core.NewProviderError(core.ProviderErrorNotApplicable, "provider descriptor not found", nil) } +func (s *ActiveProviderService) PublishReview(ctx context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) { + s.mu.Lock() + client := s.client + runtimeID := s.runtimeID + s.mu.Unlock() + if client == nil { + return core.PublishReviewResult{}, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) + } + if request.ProviderID == "" { + request.ProviderID = runtimeID + } + return client.PublishReview(ctx, request) +} + func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewContext, manual bool) (ActiveProviderState, error) { s.mu.Lock() client := s.client @@ -176,7 +191,7 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC if s.cache != nil { _ = s.cache.SaveProviderSnapshot(ctx, snap) } - st := ActiveProviderState{StableProviderKey: key, RuntimeProviderID: runtimeID, Snapshot: snap, NextSyncAt: next} + st := ActiveProviderState{StableProviderKey: key, RuntimeProviderID: runtimeID, RuntimeInfo: prev.RuntimeInfo, Snapshot: snap, NextSyncAt: next} s.setState(gen, st) return st, nil } @@ -276,8 +291,11 @@ func (s *ActiveProviderService) probe(ctx context.Context, d ports.ReviewProvide } return client, info, nil } -func (s *ActiveProviderService) loadCachedState(ctx context.Context, review core.ReviewContext, key, runtimeID string) ActiveProviderState { +func (s *ActiveProviderService) loadCachedState(ctx context.Context, review core.ReviewContext, key, runtimeID string, info ...core.ReviewProviderInfo) ActiveProviderState { st := ActiveProviderState{StableProviderKey: key, RuntimeProviderID: runtimeID} + if len(info) > 0 { + st.RuntimeInfo = info[0] + } if s.cache != nil { if snap, ok, _ := s.cache.LoadProviderSnapshot(ctx, core.NewReviewContextKey(key, review)); ok { snap.Cached = true diff --git a/internal/app/app.go b/internal/app/app.go index 7f78e39..7fb3d5a 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -17,6 +17,7 @@ import ( clipboardadapter "ero/internal/adapters/out/clipboard" gitadapter "ero/internal/adapters/out/git" pluginadapter "ero/internal/adapters/out/plugin" + providercache "ero/internal/adapters/out/providercache" chromatokenizer "ero/internal/adapters/out/syntax/chroma" "ero/internal/core" "ero/internal/logging" @@ -101,16 +102,11 @@ func newAppWithClipboard(cfg *viper.Viper, loader reviewLoader, runner tuiRunner return err } log.Info().Int("files", len(files)).Msg("review loaded") - var reviewProviders []ports.ReviewProviderClient pluginManager := pluginadapter.NewManager() providerLoader := pluginadapter.NewReviewProviderLoader(pluginManager) - _ = providerPollingConfigFromConfig(cfg) - providers, err := buildReviewProviders(ctx, providerLoader, providerLoader) - if err != nil { - log.Warn().Err(err).Msg("load review providers failed") - } else { - reviewProviders = providers - } + providerStore := providercache.NewXDGStore() + activeProviderService := NewActiveProviderService(providerLoader, providerLoader, providerStore, providerStore, providerPollingConfigFromConfig(cfg)) + activeProvider := &tuiActiveProviderController{catalog: providerLoader, service: activeProviderService} var metadata ports.GitMetadataReader if reader, ok := loader.(ports.GitMetadataReader); ok { metadata = reader @@ -118,7 +114,7 @@ func newAppWithClipboard(cfg *viper.Viper, loader reviewLoader, runner tuiRunner metadata = reader } reviewContext := buildReviewContext(initialRequest, files, metadata, version) - err = runner.Run(tui.NewModelWithReviewProvidersContext(ctx, files, terminal.NewCapabilities(), loader, initialRequest, clipboardWriter, reviewContext, reviewProviders)) + err = runner.Run(tui.NewModelWithActiveProviderContext(ctx, files, terminal.NewCapabilities(), loader, initialRequest, clipboardWriter, reviewContext, activeProvider, nil)) if err != nil { log.Error().Err(err).Msg("tui exited with error") return err diff --git a/internal/app/tui_active_provider.go b/internal/app/tui_active_provider.go new file mode 100644 index 0000000..e3ba56a --- /dev/null +++ b/internal/app/tui_active_provider.go @@ -0,0 +1,77 @@ +package app + +import ( + "context" + + "ero/internal/adapters/in/tui" + "ero/internal/core" + "ero/internal/ports" +) + +type tuiActiveProviderController struct { + catalog ports.ReviewProviderCatalog + service *ActiveProviderService +} + +func (c *tuiActiveProviderController) Catalog(ctx context.Context) ([]ports.ReviewProviderDescriptor, error) { + if c == nil || c.catalog == nil { + return nil, nil + } + return c.catalog.ListReviewProviderDescriptors(ctx) +} + +func (c *tuiActiveProviderController) Start(ctx context.Context, review core.ReviewContext) (tui.ActiveProviderState, error) { + state, err := c.service.Start(ctx, review) + return c.toTUIState(state), err +} + +func (c *tuiActiveProviderController) Refresh(ctx context.Context, review core.ReviewContext, manual bool) (tui.ActiveProviderState, error) { + state, err := c.service.Refresh(ctx, review, manual) + return c.toTUIState(state), err +} + +func (c *tuiActiveProviderController) PublishReview(ctx context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) { + if c == nil || c.service == nil { + return core.PublishReviewResult{}, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) + } + return c.service.PublishReview(ctx, request) +} + +func (c *tuiActiveProviderController) Generation() int64 { + if c == nil || c.service == nil { + return 0 + } + return c.service.Generation() +} + +func (c *tuiActiveProviderController) CompleteTimer(ctx context.Context, review core.ReviewContext, generation int64) (tui.ActiveProviderState, error) { + if c == nil || c.service == nil { + return tui.ActiveProviderState{}, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) + } + state, err := c.service.CompleteTimer(ctx, review, generation) + return c.toTUIState(state), err +} + +func (c *tuiActiveProviderController) Switch(ctx context.Context, review core.ReviewContext, stableKey string) (tui.ActiveProviderState, error) { + state, err := c.service.Switch(ctx, review, stableKey) + return c.toTUIState(state), err +} + +func (c *tuiActiveProviderController) Close() error { + if c == nil || c.service == nil { + return nil + } + return c.service.Close() +} + +func (c *tuiActiveProviderController) toTUIState(state ActiveProviderState) tui.ActiveProviderState { + return tui.ActiveProviderState{ + StableProviderKey: state.StableProviderKey, + RuntimeProviderID: state.RuntimeProviderID, + RuntimeInfo: state.RuntimeInfo, + Snapshot: state.Snapshot, + FromCache: state.FromCache, + Syncing: state.Syncing, + LastError: state.LastError, + } +} From e7817dca7c8e1e12a54bfbe10f9871ee87eb3fa8 Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 15:45:40 +0200 Subject: [PATCH 04/22] feat(providers): support PR overview snapshots --- go.mod | 7 + go.sum | 15 ++ internal/adapters/in/tui/markdown_renderer.go | 122 +++++++++++++++ .../adapters/in/tui/markdown_renderer_test.go | 75 +++++++++ internal/adapters/in/tui/model.go | 2 + internal/adapters/in/tui/pr_sheet.go | 148 +++++++++++++++++- internal/adapters/in/tui/pr_sheet_test.go | 60 +++++++ internal/adapters/out/plugin/client.go | 71 ++++++++- internal/adapters/out/plugin/client_test.go | 73 +++++++++ internal/app/active_provider_service.go | 24 ++- internal/core/provider_snapshot.go | 33 +++- internal/core/review_provider.go | 1 + pkg/plugin/plugin.go | 15 +- pkg/plugin/protocol/types.go | 49 ++++++ pkg/plugin/server.go | 22 +++ pkg/plugin/server_test.go | 65 ++++++++ 16 files changed, 764 insertions(+), 18 deletions(-) create mode 100644 internal/adapters/in/tui/markdown_renderer.go create mode 100644 internal/adapters/in/tui/markdown_renderer_test.go diff --git a/go.mod b/go.mod index 6153a0b..a6da7a8 100644 --- a/go.mod +++ b/go.mod @@ -5,6 +5,7 @@ go 1.26.3 require ( charm.land/bubbles/v2 v2.1.0 charm.land/bubbletea/v2 v2.0.6 + charm.land/glamour/v2 v2.0.0 charm.land/lipgloss/v2 v2.0.3 github.com/alecthomas/chroma/v2 v2.24.1 github.com/bnema/zerowrap v1.4.0 @@ -25,8 +26,10 @@ require ( github.com/Microsoft/go-winio v0.6.2 // indirect github.com/ProtonMail/go-crypto v1.1.6 // indirect github.com/atotto/clipboard v0.1.4 // indirect + github.com/aymerick/douceur v0.2.0 // indirect github.com/charmbracelet/colorprofile v0.4.3 // indirect github.com/charmbracelet/ultraviolet v0.0.0-20260416155717-489999b90468 // indirect + github.com/charmbracelet/x/exp/slice v0.0.0-20250327172914-2fdc97757edf // indirect github.com/charmbracelet/x/term v0.2.2 // indirect github.com/charmbracelet/x/termios v0.1.1 // indirect github.com/charmbracelet/x/windows v0.2.2 // indirect @@ -43,6 +46,7 @@ require ( github.com/go-git/go-billy/v5 v5.9.0 // indirect github.com/go-viper/mapstructure/v2 v2.4.0 // indirect github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8 // indirect + github.com/gorilla/css v1.0.1 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/jbenet/go-context v0.0.0-20150711004518-d14ea06fba99 // indirect github.com/kevinburke/ssh_config v1.2.0 // indirect @@ -51,6 +55,7 @@ require ( github.com/mattn/go-colorable v0.1.13 // indirect github.com/mattn/go-isatty v0.0.20 // indirect github.com/mattn/go-runewidth v0.0.23 // indirect + github.com/microcosm-cc/bluemonday v1.0.27 // indirect github.com/muesli/cancelreader v0.2.2 // indirect github.com/pjbgf/sha1cd v0.6.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect @@ -66,6 +71,8 @@ require ( github.com/subosito/gotenv v1.6.0 // indirect github.com/xanzy/ssh-agent v0.3.3 // indirect github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect + github.com/yuin/goldmark v1.7.8 // indirect + github.com/yuin/goldmark-emoji v1.0.5 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/crypto v0.50.0 // indirect golang.org/x/net v0.53.0 // indirect diff --git a/go.sum b/go.sum index c9b4605..56d9dc5 100644 --- a/go.sum +++ b/go.sum @@ -2,6 +2,8 @@ charm.land/bubbles/v2 v2.1.0 h1:YSnNh5cPYlYjPxRrzs5VEn3vwhtEn3jVGRBT3M7/I0g= charm.land/bubbles/v2 v2.1.0/go.mod h1:l97h4hym2hvWBVfmJDtrEHHCtkIKeTEb3TTJ4ZOB3wY= charm.land/bubbletea/v2 v2.0.6 h1:UHN/91OyuhaOFGSrBXQ/hMZD8IO1Uc4BvHlgHXL2WJo= charm.land/bubbletea/v2 v2.0.6/go.mod h1:MH/D8ZLlN3op37vQvijKuU29g3rqTp+aQapURFonF9g= +charm.land/glamour/v2 v2.0.0 h1:IDBoqLEy7Hdpb9VOXN+khLP/XSxtJy1VsHuW/yF87+U= +charm.land/glamour/v2 v2.0.0/go.mod h1:kjq9WB0s8vuUYZNYey2jp4Lgd9f4cKdzAw88FZtpj/w= charm.land/lipgloss/v2 v2.0.3 h1:yM2zJ4Cf5Y51b7RHIwioil4ApI/aypFXXVHSwlM6RzU= charm.land/lipgloss/v2 v2.0.3/go.mod h1:7myLU9iG/3xluAWzpY/fSxYYHCgoKTie7laxk6ATwXA= dario.cat/mergo v1.0.1 h1:Ra4+bf83h2ztPIQYNP99R6m+Y7KfnARDfID+a+vLl4s= @@ -29,6 +31,8 @@ github.com/aymanbagabas/go-osc52/v2 v2.0.1 h1:HwpRHbFMcZLEVr42D4p7XBqjyuxQH5SMiE github.com/aymanbagabas/go-osc52/v2 v2.0.1/go.mod h1:uYgXzlJ7ZpABp8OJ+exZzJJhRNQ2ASbcXHWsFqH8hp8= github.com/aymanbagabas/go-udiff v0.4.1 h1:OEIrQ8maEeDBXQDoGCbbTTXYJMYRCRO1fnodZ12Gv5o= github.com/aymanbagabas/go-udiff v0.4.1/go.mod h1:0L9PGwj20lrtmEMeyw4WKJ/TMyDtvAoK9bf2u/mNo3w= +github.com/aymerick/douceur v0.2.0 h1:Mv+mAeH1Q+n9Fr+oyamOlAkUNPWPlA8PPGR0QAaYuPk= +github.com/aymerick/douceur v0.2.0/go.mod h1:wlT5vV2O3h55X9m7iVYN0TBM0NH/MmbLnd30/FjWUq4= github.com/bnema/zerowrap v1.4.0 h1:QYb+/dLS4PPNc4HN0q6C0pIWO5lmqAjdauAWwaeX2LM= github.com/bnema/zerowrap v1.4.0/go.mod h1:30FCqzwS7FNTNH1GC+ScGC99GITOzOaCYW266L4QF5s= github.com/charmbracelet/colorprofile v0.4.3 h1:QPa1IWkYI+AOB+fE+mg/5/4HRMZcaXex9t5KX76i20Q= @@ -43,6 +47,8 @@ github.com/charmbracelet/x/cellbuf v0.0.13 h1:/KBBKHuVRbq1lYx5BzEHBAFBP8VcQzJejZ github.com/charmbracelet/x/cellbuf v0.0.13/go.mod h1:xe0nKWGd3eJgtqZRaN9RjMtK7xUYchjzPr7q6kcvCCs= github.com/charmbracelet/x/exp/golden v0.0.0-20250806222409-83e3a29d542f h1:pk6gmGpCE7F3FcjaOEKYriCvpmIN4+6OS/RD0vm4uIA= github.com/charmbracelet/x/exp/golden v0.0.0-20250806222409-83e3a29d542f/go.mod h1:IfZAMTHB6XkZSeXUqriemErjAWCCzT0LwjKFYCZyw0I= +github.com/charmbracelet/x/exp/slice v0.0.0-20250327172914-2fdc97757edf h1:rLG0Yb6MQSDKdB52aGX55JT1oi0P0Kuaj7wi1bLUpnI= +github.com/charmbracelet/x/exp/slice v0.0.0-20250327172914-2fdc97757edf/go.mod h1:B3UgsnsBZS/eX42BlaNiJkD1pPOUa+oF1IYC6Yd2CEU= github.com/charmbracelet/x/term v0.2.2 h1:xVRT/S2ZcKdhhOuSP4t5cLi5o+JxklsoEObBSgfgZRk= github.com/charmbracelet/x/term v0.2.2/go.mod h1:kF8CY5RddLWrsgVwpw4kAa6TESp6EB5y3uxGLeCqzAI= github.com/charmbracelet/x/termios v0.1.1 h1:o3Q2bT8eqzGnGPOYheoYS8eEleT5ZVNYNy8JawjaNZY= @@ -97,6 +103,8 @@ github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/shlex v0.0.0-20191202100458-e7afc7fbc510 h1:El6M4kTTCOh6aBiKaUGG7oYTSPP8MxqL4YI3kZKwcP4= github.com/google/shlex v0.0.0-20191202100458-e7afc7fbc510/go.mod h1:pupxD2MaaD3pAXIBCelhxNneeOaAeabZDe5s4K6zSpQ= +github.com/gorilla/css v1.0.1 h1:ntNaBIghp6JmvWnxbZKANoLyuXTPZ4cAMlo6RyhlbO8= +github.com/gorilla/css v1.0.1/go.mod h1:BvnYkspnSzMmwRK+b8/xgNPLiIuNZr6vbZBTPQ2A3b0= github.com/henvic/httpretty v0.0.6 h1:JdzGzKZBajBfnvlMALXXMVQWxWMF/ofTy8C3/OSUTxs= github.com/henvic/httpretty v0.0.6/go.mod h1:X38wLjWXHkXT7r2+uK8LjCMne9rsuNaBLJ+5cU2/Pmo= github.com/hexops/gotextdiff v1.0.3 h1:gitA9+qJrrTCsiCl7+kh75nPqQt1cx4ZkudSTLoUqJM= @@ -126,6 +134,8 @@ github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWE github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/mattn/go-runewidth v0.0.23 h1:7ykA0T0jkPpzSvMS5i9uoNn2Xy3R383f9HDx3RybWcw= github.com/mattn/go-runewidth v0.0.23/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs= +github.com/microcosm-cc/bluemonday v1.0.27 h1:MpEUotklkwCSLeH+Qdx1VJgNqLlpY2KXwXFM08ygZfk= +github.com/microcosm-cc/bluemonday v1.0.27/go.mod h1:jFi9vgW+H7c3V0lb6nR74Ib/DIB5OBs92Dimizgw2cA= github.com/muesli/cancelreader v0.2.2 h1:3I4Kt4BQjOR54NavqnDogx/MIoWBFa0StPA8ELUXHmA= github.com/muesli/cancelreader v0.2.2/go.mod h1:3XuTXfFS2VjM+HTLZY9Ak0l6eUKfijIfMUZ4EgX0QYo= github.com/muesli/reflow v0.3.0 h1:IFsN6K9NfGtjeggFP+68I4chLZV2yIKsXJFNZ+eWh6s= @@ -185,6 +195,11 @@ github.com/xanzy/ssh-agent v0.3.3 h1:+/15pJfg/RsTxqYcX6fHqOXZwwMP+2VyYWJeWM2qQFM github.com/xanzy/ssh-agent v0.3.3/go.mod h1:6dzNDKs0J9rVPHPhaGCukekBHKqfl+L3KghI1Bc68Uw= github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e h1:JVG44RsyaB9T2KIHavMF/ppJZNG9ZpyihvCd0w101no= github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e/go.mod h1:RbqR21r5mrJuqunuUZ/Dhy/avygyECGrLceyNeo4LiM= +github.com/yuin/goldmark v1.7.1/go.mod h1:uzxRWxtg69N339t3louHJ7+O03ezfj6PlliRlaOzY1E= +github.com/yuin/goldmark v1.7.8 h1:iERMLn0/QJeHFhxSt3p6PeN9mGnvIKSpG9YYorDMnic= +github.com/yuin/goldmark v1.7.8/go.mod h1:uzxRWxtg69N339t3louHJ7+O03ezfj6PlliRlaOzY1E= +github.com/yuin/goldmark-emoji v1.0.5 h1:EMVWyCGPlXJfUXBXpuMu+ii3TIaxbVBnEX9uaDC4cIk= +github.com/yuin/goldmark-emoji v1.0.5/go.mod h1:tTkZEbwu5wkPmgTcitqddVxY9osFZiavD+r4AzQrh1U= go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= golang.org/x/crypto v0.0.0-20220622213112-05595931fe9d/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4= diff --git a/internal/adapters/in/tui/markdown_renderer.go b/internal/adapters/in/tui/markdown_renderer.go new file mode 100644 index 0000000..ed7c0d3 --- /dev/null +++ b/internal/adapters/in/tui/markdown_renderer.go @@ -0,0 +1,122 @@ +package tui + +import ( + "crypto/sha256" + "encoding/hex" + "regexp" + "strings" + + "charm.land/glamour/v2" +) + +type MarkdownTheme string + +const ( + MarkdownThemeDark MarkdownTheme = "dark" + MarkdownThemeLight MarkdownTheme = "light" +) + +type markdownTermRenderer interface { + Render(markdown string) (string, error) +} + +type markdownRendererFactory func(width int, theme MarkdownTheme) (markdownTermRenderer, error) + +type MarkdownRenderer struct { + factory markdownRendererFactory + renderers map[markdownRendererConfig]markdownTermRenderer + entries map[markdownRendererCacheKey]string +} + +type markdownRendererConfig struct { + Width int + Theme MarkdownTheme +} + +type markdownRendererCacheKey struct { + InputHash string + Width int + Theme MarkdownTheme +} + +func NewMarkdownRenderer() *MarkdownRenderer { + return NewMarkdownRendererWithFactory(newGlamourTermRenderer) +} + +func NewMarkdownRendererWithFactory(factory markdownRendererFactory) *MarkdownRenderer { + if factory == nil { + factory = newGlamourTermRenderer + } + return &MarkdownRenderer{ + factory: factory, + renderers: map[markdownRendererConfig]markdownTermRenderer{}, + entries: map[markdownRendererCacheKey]string{}, + } +} + +func (r *MarkdownRenderer) Render(markdown string, width int, theme MarkdownTheme) string { + if r == nil { + return safeMarkdownFallback(markdown) + } + if width < 1 { + width = 1 + } + if theme != MarkdownThemeLight { + theme = MarkdownThemeDark + } + if r.renderers == nil { + r.renderers = map[markdownRendererConfig]markdownTermRenderer{} + } + if r.entries == nil { + r.entries = map[markdownRendererCacheKey]string{} + } + + key := markdownRendererCacheKey{InputHash: hashMarkdownInput(markdown), Width: width, Theme: theme} + if rendered, ok := r.entries[key]; ok { + return rendered + } + + renderer, err := r.renderer(width, theme) + if err != nil { + return safeMarkdownFallback(markdown) + } + rendered, err := renderer.Render(markdown) + if err != nil { + return safeMarkdownFallback(markdown) + } + r.entries[key] = rendered + return rendered +} + +func (r *MarkdownRenderer) renderer(width int, theme MarkdownTheme) (markdownTermRenderer, error) { + config := markdownRendererConfig{Width: width, Theme: theme} + if renderer, ok := r.renderers[config]; ok { + return renderer, nil + } + renderer, err := r.factory(width, theme) + if err != nil { + return nil, err + } + r.renderers[config] = renderer + return renderer, nil +} + +func newGlamourTermRenderer(width int, theme MarkdownTheme) (markdownTermRenderer, error) { + return glamour.NewTermRenderer( + glamour.WithStandardStyle(string(theme)), + glamour.WithWordWrap(width), + ) +} + +func hashMarkdownInput(input string) string { + sum := sha256.Sum256([]byte(input)) + return hex.EncodeToString(sum[:]) +} + +var ansiEscapePattern = regexp.MustCompile(`\x1b\[[0-9;?]*[ -/]*[@-~]`) + +func safeMarkdownFallback(input string) string { + withoutEscapes := ansiEscapePattern.ReplaceAllString(input, "") + withoutEscapes = strings.ReplaceAll(withoutEscapes, "\x1b", "") + return strings.TrimSpace(withoutEscapes) +} diff --git a/internal/adapters/in/tui/markdown_renderer_test.go b/internal/adapters/in/tui/markdown_renderer_test.go new file mode 100644 index 0000000..bd91522 --- /dev/null +++ b/internal/adapters/in/tui/markdown_renderer_test.go @@ -0,0 +1,75 @@ +package tui + +import ( + "errors" + "regexp" + "strings" + "testing" +) + +type fakeMarkdownTermRenderer struct { + render func(string) (string, error) +} + +func (r fakeMarkdownTermRenderer) Render(markdown string) (string, error) { + return r.render(markdown) +} + +func TestMarkdownRendererCachesByInputWidthAndTheme(t *testing.T) { + calls := 0 + renderer := NewMarkdownRendererWithFactory(func(width int, theme MarkdownTheme) (markdownTermRenderer, error) { + return fakeMarkdownTermRenderer{render: func(markdown string) (string, error) { + calls++ + return string(theme) + ":" + markdown + ":" + string(rune(width)), nil + }}, nil + }) + + if got := renderer.Render("**hello**", 80, MarkdownThemeDark); got == "**hello**" { + t.Fatalf("expected rendered output, got fallback %q", got) + } + if got := renderer.Render("**hello**", 80, MarkdownThemeDark); got == "**hello**" { + t.Fatalf("expected cached rendered output, got fallback %q", got) + } + if calls != 1 { + t.Fatalf("expected identical input/width/theme to use cache, got %d render calls", calls) + } + + renderer.Render("**hello**", 100, MarkdownThemeDark) + renderer.Render("**hello**", 100, MarkdownThemeLight) + renderer.Render("_hello_", 100, MarkdownThemeLight) + + if calls != 4 { + t.Fatalf("expected width, theme, and input changes to invalidate cache, got %d render calls", calls) + } +} + +func TestMarkdownRendererRendersFencedCodeBlocks(t *testing.T) { + renderer := NewMarkdownRenderer() + + got := renderer.Render("```go\nfmt.Println(\"hi\")\n```", 80, MarkdownThemeDark) + plain := regexp.MustCompile(`\x1b\[[0-9;?]*[ -/]*[@-~]`).ReplaceAllString(got, "") + + if !strings.Contains(plain, "fmt.Println") { + t.Fatalf("expected rendered fenced code block to include code, got %q", got) + } + if strings.Contains(plain, "```") { + t.Fatalf("expected rendered fenced code block to omit markdown fences, got %q", got) + } +} + +func TestMarkdownRendererReturnsSafeFallbackOnRenderError(t *testing.T) { + renderer := NewMarkdownRendererWithFactory(func(width int, theme MarkdownTheme) (markdownTermRenderer, error) { + return fakeMarkdownTermRenderer{render: func(markdown string) (string, error) { + return "", errors.New("boom") + }}, nil + }) + + got := renderer.Render("hello\x1b[31m **world**", 80, MarkdownThemeDark) + + if strings.Contains(got, "\x1b") { + t.Fatalf("expected fallback to strip escape characters, got %q", got) + } + if !strings.Contains(got, "hello") || !strings.Contains(got, "world") { + t.Fatalf("expected fallback to preserve readable text, got %q", got) + } +} diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index d7c8e6f..7ae2cc4 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -136,6 +136,7 @@ type Model struct { providerPicker providerPickerState publish publishState prSheet prSheetState + markdownRenderer *MarkdownRenderer ctx context.Context reviewLineCache *render.ReviewLineCache cachedEditorWidth int @@ -193,6 +194,7 @@ func NewModelWithActiveProviderContext(ctx context.Context, files []core.ReviewF providerInfos: nil, providerInfoByClient: map[ports.ReviewProviderClient]core.ReviewProviderInfo{}, remoteThreads: nil, + markdownRenderer: NewMarkdownRenderer(), ctx: ctx, reviewLineCache: render.NewReviewLineCache(), } diff --git a/internal/adapters/in/tui/pr_sheet.go b/internal/adapters/in/tui/pr_sheet.go index 1ecc196..b89197a 100644 --- a/internal/adapters/in/tui/pr_sheet.go +++ b/internal/adapters/in/tui/pr_sheet.go @@ -3,9 +3,12 @@ package tui import ( "strconv" "strings" + "time" tea "charm.land/bubbletea/v2" "charm.land/lipgloss/v2" + + "ero/internal/core" ) type prSheetToggledMsg struct{} @@ -15,7 +18,7 @@ type prSheetScrolledMsg struct { } type prSheetState struct { - open bool + open bool yOffset int } @@ -68,22 +71,153 @@ func (m Model) prSheetLines() []string { if m.activeRuntimeInfo.ID != "" || m.activeRuntimeInfo.Label != "" || m.activeRuntimeInfo.Name != "" { provider = providerDisplayLabel(m.activeRuntimeInfo) } - return []string{ + + lines := []string{ "Pull request", "", "Provider: " + provider, - "", - "Overview placeholder", - "Phase 4 will render PR markdown and metadata here.", - "", - "Remote threads: " + pluralCount(len(m.remoteThreads), "thread"), } + if m.providerSyncState.Status != "" { + lines = append(lines, "Sync: "+string(m.providerSyncState.Status)) + } + lines = append(lines, "Remote threads: "+pluralCount(len(m.remoteThreads), "thread"), "") + + if m.providerOverview == nil { + return append(lines, + "No provider overview loaded.", + "Open an active provider that supports PR snapshots to show PR metadata, body, comments, and reviews here.", + ) + } + return append(lines, m.providerOverviewLines(m.providerOverview)...) +} + +func (m Model) providerOverviewLines(overview *core.ProviderOverview) []string { + contentWidth := max(prSheetWidth(m.width)-2, 1) + lines := []string{} + if strings.TrimSpace(overview.Title) != "" { + lines = append(lines, overview.Title) + } else { + lines = append(lines, "Untitled pull request") + } + metadata := providerOverviewMetadata(overview) + if len(metadata) > 0 { + lines = append(lines, metadata...) + } + if strings.TrimSpace(overview.ExternalURL) != "" { + lines = append(lines, overview.ExternalURL) + } + lines = append(lines, "") + + if strings.TrimSpace(overview.Body) != "" { + lines = append(lines, "Body", "") + lines = append(lines, renderPRSheetMarkdown(m.markdownRenderer, overview.Body, contentWidth)...) + lines = append(lines, "") + } + + lines = append(lines, "Issue comments: "+strconv.Itoa(len(overview.Comments))) + for _, comment := range overview.Comments { + lines = append(lines, "", commentHeader(comment.Author, comment.CreatedAt)) + lines = append(lines, renderPRSheetMarkdown(m.markdownRenderer, comment.Body, contentWidth)...) + } + + lines = append(lines, "", "Review summaries: "+strconv.Itoa(len(overview.Reviews))) + for _, review := range overview.Reviews { + lines = append(lines, "", reviewSummaryHeader(review)) + if strings.TrimSpace(review.Body) != "" { + lines = append(lines, renderPRSheetMarkdown(m.markdownRenderer, review.Body, contentWidth)...) + } + } + return trimTrailingBlankLines(lines) } func (m Model) prSheetLineCount() int { return len(m.prSheetLines()) } +func renderPRSheetMarkdown(renderer *MarkdownRenderer, markdown string, width int) []string { + rendered := renderer.Render(markdown, width, MarkdownThemeDark) + plain := safeMarkdownFallback(rendered) + if strings.TrimSpace(plain) == "" { + return []string{"(empty)"} + } + return strings.Split(plain, "\n") +} + +func providerOverviewMetadata(overview *core.ProviderOverview) []string { + if overview == nil { + return nil + } + parts := make([]string, 0, 4) + if overview.Number > 0 { + parts = append(parts, "#"+strconv.Itoa(overview.Number)) + } + if strings.TrimSpace(overview.State) != "" { + parts = append(parts, strings.TrimSpace(overview.State)) + } + if strings.TrimSpace(overview.Author) != "" { + parts = append(parts, "by "+strings.TrimSpace(overview.Author)) + } + if strings.TrimSpace(overview.BaseRef) != "" || strings.TrimSpace(overview.HeadRef) != "" { + parts = append(parts, strings.TrimSpace(overview.BaseRef)+" ← "+strings.TrimSpace(overview.HeadRef)) + } + lines := []string{} + if len(parts) > 0 { + lines = append(lines, strings.Join(parts, " · ")) + } + if overview.UpdatedAt != nil && !overview.UpdatedAt.IsZero() { + lines = append(lines, "Updated "+overview.UpdatedAt.Local().Format("2006-01-02 15:04")) + } + return lines +} + +func commentHeader(author string, createdAt time.Time) string { + parts := []string{"• Comment"} + if strings.TrimSpace(author) != "" { + parts = append(parts, "by "+strings.TrimSpace(author)) + } + if !createdAt.IsZero() { + parts = append(parts, createdAt.Local().Format("2006-01-02 15:04")) + } + return strings.Join(parts, " ") +} + +func reviewSummaryHeader(review core.ProviderReviewSummary) string { + marker := reviewStateMarker(review.State) + parts := []string{marker} + if strings.TrimSpace(review.State) != "" { + parts = append(parts, strings.ToUpper(strings.TrimSpace(review.State))) + } + if strings.TrimSpace(review.Author) != "" { + parts = append(parts, "by "+strings.TrimSpace(review.Author)) + } + if !review.SubmittedAt.IsZero() { + parts = append(parts, review.SubmittedAt.Local().Format("2006-01-02 15:04")) + } + return strings.Join(parts, " ") +} + +func reviewStateMarker(state string) string { + switch strings.ToLower(strings.TrimSpace(state)) { + case "approved", "approve": + return "✓" + case "changes_requested", "request_changes", "changes requested": + return "!" + case "commented", "comment": + return "•" + case "dismissed": + return "×" + default: + return "-" + } +} + +func trimTrailingBlankLines(lines []string) []string { + for len(lines) > 0 && strings.TrimSpace(lines[len(lines)-1]) == "" { + lines = lines[:len(lines)-1] + } + return lines +} + func prSheetWidth(totalWidth int) int { if totalWidth <= 1 { return 1 diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go index 6666639..647a8c2 100644 --- a/internal/adapters/in/tui/pr_sheet_test.go +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -3,6 +3,7 @@ package tui import ( "strings" "testing" + "time" tea "charm.land/bubbletea/v2" "github.com/stretchr/testify/assert" @@ -93,3 +94,62 @@ func TestPRSheetHasIndependentScrollState(t *testing.T) { assert.Contains(t, stripANSI(model.renderPRSheet(model.width, model.height)), "Provider:") assert.NotContains(t, stripANSI(model.renderPRSheet(model.width, model.height)), "Pull request") } + +func TestPRSheetRendersOverviewMarkdownCommentsAndReviews(t *testing.T) { + t.Parallel() + + createdAt := time.Date(2026, 1, 2, 3, 4, 0, 0, time.UTC) + model := NewModel(nil) + model.width = 180 + model.height = 40 + model.activeRuntimeInfo = core.ReviewProviderInfo{ID: "github", Label: "GitHub"} + model.providerOverview = &core.ProviderOverview{ + Title: "Add provider snapshots", + Number: 42, + State: "OPEN", + ExternalURL: "https://example.invalid/pull/1", + Author: "author", + Body: "This **adds** snapshots.\n\n```go\nfmt.Println(\"hi\")\n```", + BaseRef: "main", + HeadRef: "feature", + UpdatedAt: &createdAt, + Comments: []core.ProviderIssueComment{{ + Author: "reviewer", + Body: "Please update the **docs**.", + CreatedAt: createdAt, + }}, + Reviews: []core.ProviderReviewSummary{{ + Author: "maintainer", + State: "approved", + Body: "Looks **good** to me.", + SubmittedAt: createdAt, + }}, + } + + plain := stripANSI(model.renderPRSheet(model.width, model.height)) + + assert.Contains(t, plain, "Add provider snapshots") + assert.Contains(t, plain, "#42 · OPEN · by author · main ← feature") + assert.Contains(t, plain, "Updated "+createdAt.Local().Format("2006-01-02 15:04")) + assert.Contains(t, plain, "https://example.invalid/pull/1") + assert.Contains(t, plain, "This adds snapshots.") + assert.Contains(t, plain, "fmt.Println") + assert.Contains(t, plain, "Issue comments: 1") + assert.Contains(t, plain, "Comment by reviewer") + assert.Contains(t, plain, "Please update the docs.") + assert.Contains(t, plain, "Review summaries: 1") + assert.Contains(t, plain, "✓ APPROVED by maintainer") + assert.Contains(t, plain, "Looks good to me.") +} + +func TestPRSheetRendersNilOverviewFallback(t *testing.T) { + t.Parallel() + + model := NewModel(nil) + model.width = 100 + model.height = 12 + plain := stripANSI(model.renderPRSheet(model.width, model.height)) + + assert.Contains(t, plain, "No provider overview loaded.") + assert.Contains(t, plain, "Remote threads: 0 threads") +} diff --git a/internal/adapters/out/plugin/client.go b/internal/adapters/out/plugin/client.go index 817c314..3eba898 100644 --- a/internal/adapters/out/plugin/client.go +++ b/internal/adapters/out/plugin/client.go @@ -33,6 +33,7 @@ type Client struct { nextID int timeout time.Duration contributionID string + providerInfo core.ReviewProviderInfo closed bool } @@ -93,7 +94,11 @@ func (c *Client) Initialize(ctx context.Context) (core.ReviewProviderInfo, error return core.ReviewProviderInfo{}, fmt.Errorf("plugin protocol mismatch: expected %q, got %q", protocol.ProtocolVersion, result.Protocol) } - return toCoreProviderInfo(result.Provider), nil + info := toCoreProviderInfo(result.Provider) + c.mu.Lock() + c.providerInfo = info + c.mu.Unlock() + return info, nil } // DetectContext asks the plugin whether it considers the review context @@ -128,6 +133,24 @@ func (c *Client) LoadRemoteThreads(ctx context.Context, review core.ReviewContex return threads, nil } +func (c *Client) LoadRemoteSnapshot(ctx context.Context, review core.ReviewContext) (core.ProviderSnapshot, error) { + c.mu.Lock() + info := c.providerInfo + c.mu.Unlock() + if !info.Capabilities.LoadRemoteSnapshot { + threads, err := c.LoadRemoteThreads(ctx, review) + if err != nil { + return core.ProviderSnapshot{}, err + } + return core.ProviderSnapshot{RuntimeProviderID: info.ID, Threads: threads}, nil + } + var result protocol.LoadRemoteSnapshotResult + if err := c.call(ctx, "load_remote_snapshot", protocol.LoadRemoteSnapshotRequest{Context: toProtocolReviewContext(review)}, &result); err != nil { + return core.ProviderSnapshot{}, err + } + return toCoreProviderSnapshot(result, info.ID), nil +} + // PublishReview sends the draft to the plugin for publication. func (c *Client) PublishReview(ctx context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) { params := protocol.PublishReviewParams{ @@ -349,6 +372,7 @@ func toCoreProviderInfo(info protocol.ReviewProviderInfo) core.ReviewProviderInf Name: info.Name, Capabilities: core.ReviewProviderCapabilities{ LoadRemoteComments: info.Capabilities.LoadRemoteComments, + LoadRemoteSnapshot: info.Capabilities.LoadRemoteSnapshot, PublishReview: info.Capabilities.PublishReview, Decisions: decisions, IdempotentPublish: info.Capabilities.IdempotentPublish, @@ -503,6 +527,51 @@ func toProtocolLineRange(r core.ReviewLineRange) protocol.ReviewLineRange { } } +func toCoreProviderSnapshot(s protocol.LoadRemoteSnapshotResult, fallbackProviderID string) core.ProviderSnapshot { + threads := make([]core.RemoteReviewThread, len(s.Threads)) + for i, t := range s.Threads { + threads[i] = toCoreRemoteThread(t) + } + runtimeID := s.RuntimeProviderID + if runtimeID == "" { + runtimeID = fallbackProviderID + } + snap := core.ProviderSnapshot{RuntimeProviderID: runtimeID, Threads: threads, Overview: toCoreProviderOverview(s.Overview), Metadata: s.Metadata} + if s.FetchedAt != nil { + snap.FetchedAt = *s.FetchedAt + } + snap.ExpiresAt = s.ExpiresAt + return snap +} + +func toCoreProviderOverview(o *protocol.ProviderOverview) *core.ProviderOverview { + if o == nil { + return nil + } + comments := make([]core.ProviderIssueComment, len(o.Comments)) + for i, c := range o.Comments { + comments[i] = core.ProviderIssueComment{ExternalID: c.ExternalID, Author: c.Author, Body: c.Body, CreatedAt: c.CreatedAt, UpdatedAt: c.UpdatedAt, ExternalURL: c.ExternalURL} + } + reviews := make([]core.ProviderReviewSummary, len(o.Reviews)) + for i, r := range o.Reviews { + reviews[i] = core.ProviderReviewSummary{ExternalID: r.ExternalID, Author: r.Author, State: r.State, Body: r.Body, SubmittedAt: r.SubmittedAt, ExternalURL: r.ExternalURL} + } + return &core.ProviderOverview{ + RuntimeProviderID: o.RuntimeProviderID, + Title: o.Title, + Number: o.Number, + State: o.State, + ExternalURL: o.ExternalURL, + Author: o.Author, + Body: o.Body, + BaseRef: o.BaseRef, + HeadRef: o.HeadRef, + UpdatedAt: o.UpdatedAt, + Comments: comments, + Reviews: reviews, + } +} + func toCoreRemoteThread(t protocol.RemoteReviewThread) core.RemoteReviewThread { comments := make([]core.RemoteReviewComment, len(t.Comments)) for i, c := range t.Comments { diff --git a/internal/adapters/out/plugin/client_test.go b/internal/adapters/out/plugin/client_test.go index 0fa402c..72fc022 100644 --- a/internal/adapters/out/plugin/client_test.go +++ b/internal/adapters/out/plugin/client_test.go @@ -76,6 +76,7 @@ func (s *fakePluginServer) Initialize(_ context.Context, req pluginsdk.Initializ Name: "fake-plugin", Capabilities: pluginsdk.ReviewProviderCapabilities{ LoadRemoteComments: true, + LoadRemoteSnapshot: os.Getenv("FAKE_PLUGIN_SNAPSHOT") == "1", PublishReview: true, Decisions: []pluginsdk.ReviewDecision{pluginsdk.ReviewDecisionComment, pluginsdk.ReviewDecisionApprove}, }, @@ -107,6 +108,37 @@ func (s *fakePluginServer) LoadRemoteThreads(_ context.Context, req pluginsdk.Lo }, nil } +func (s *fakePluginServer) LoadRemoteSnapshot(_ context.Context, req pluginsdk.LoadRemoteSnapshotRequest) (pluginsdk.LoadRemoteSnapshotResult, error) { + if os.Getenv("FAKE_PLUGIN_SNAPSHOT") != "1" { + return pluginsdk.LoadRemoteSnapshotResult{}, protocol.NewError(protocol.ErrorUnsupportedCapability, "snapshot unsupported") + } + updated := time.Date(2026, 6, 5, 12, 30, 0, 0, time.UTC) + return pluginsdk.LoadRemoteSnapshotResult{ + RuntimeProviderID: "snapshot-runtime", + Threads: []pluginsdk.RemoteReviewThread{{ + ProviderID: "snapshot-runtime", + ExternalID: "snapshot-thread", + FilePath: "snapshot.go", + }}, + Overview: &pluginsdk.ProviderOverview{ + RuntimeProviderID: "snapshot-runtime", + Title: "Snapshot PR", + Number: 42, + State: "OPEN", + ExternalURL: "https://example.com/pr/42", + Author: "alice", + Body: "## Body", + BaseRef: "main", + HeadRef: "feature", + UpdatedAt: &updated, + Comments: []pluginsdk.ProviderIssueComment{{ExternalID: "ic1", Author: "bob", Body: "hello", CreatedAt: updated, UpdatedAt: updated, ExternalURL: "https://example.com/c/1"}}, + Reviews: []pluginsdk.ProviderReviewSummary{{ExternalID: "rv1", Author: "carol", State: "APPROVED", Body: "approved", SubmittedAt: updated, ExternalURL: "https://example.com/r/1"}}, + }, + Metadata: map[string]string{"source": "snapshot"}, + FetchedAt: &updated, + }, nil +} + func (s *fakePluginServer) PublishReview(_ context.Context, req pluginsdk.PublishReviewParams) (pluginsdk.PublishReviewResultData, error) { return pluginsdk.PublishReviewResultData{ Result: pluginsdk.ReviewPublishResult{ @@ -243,6 +275,47 @@ func TestClientLoadRemoteThreads(t *testing.T) { assert.Equal(t, "LGTM", threads[0].Comments[0].Body) } +func TestClientLoadRemoteSnapshotFallsBackToThreadsWhenCapabilityMissing(t *testing.T) { + client := setupFakeClient(t) + _, err := client.Initialize(context.Background()) + require.NoError(t, err) + + snapshot, err := client.LoadRemoteSnapshot(context.Background(), core.ReviewContext{}) + require.NoError(t, err) + require.Len(t, snapshot.Threads, 1) + assert.Equal(t, "thread-1", snapshot.Threads[0].ExternalID) + assert.Nil(t, snapshot.Overview) +} + +func TestClientLoadRemoteSnapshotUsesAdvertisedSnapshotAndMapsOverview(t *testing.T) { + client := setupFakeClientWithEnv(t, DefaultPluginTimeout, "FAKE_PLUGIN_SNAPSHOT=1") + info, err := client.Initialize(context.Background()) + require.NoError(t, err) + require.True(t, info.Capabilities.LoadRemoteSnapshot) + + snapshot, err := client.LoadRemoteSnapshot(context.Background(), core.ReviewContext{}) + require.NoError(t, err) + require.Equal(t, "snapshot-runtime", snapshot.RuntimeProviderID) + require.Len(t, snapshot.Threads, 1) + assert.Equal(t, "snapshot-thread", snapshot.Threads[0].ExternalID) + require.NotNil(t, snapshot.Overview) + assert.Equal(t, "snapshot-runtime", snapshot.Overview.RuntimeProviderID) + assert.Equal(t, "Snapshot PR", snapshot.Overview.Title) + assert.Equal(t, 42, snapshot.Overview.Number) + assert.Equal(t, "OPEN", snapshot.Overview.State) + assert.Equal(t, "alice", snapshot.Overview.Author) + assert.Equal(t, "main", snapshot.Overview.BaseRef) + assert.Equal(t, "feature", snapshot.Overview.HeadRef) + assert.Equal(t, "https://example.com/pr/42", snapshot.Overview.ExternalURL) + require.NotNil(t, snapshot.Overview.UpdatedAt) + assert.False(t, snapshot.Overview.UpdatedAt.IsZero()) + require.Len(t, snapshot.Overview.Comments, 1) + assert.Equal(t, "bob", snapshot.Overview.Comments[0].Author) + require.Len(t, snapshot.Overview.Reviews, 1) + assert.Equal(t, "APPROVED", snapshot.Overview.Reviews[0].State) + assert.Equal(t, "snapshot", snapshot.Metadata["source"]) +} + func TestClientMapsProtocolErrorsToProviderErrors(t *testing.T) { client := setupFakeClientWithEnv(t, DefaultPluginTimeout, "FAKE_PLUGIN_LOAD_ERROR_CODE="+protocol.ErrorRemoteRateLimited) diff --git a/internal/app/active_provider_service.go b/internal/app/active_provider_service.go index 00f5380..fe47241 100644 --- a/internal/app/active_provider_service.go +++ b/internal/app/active_provider_service.go @@ -27,6 +27,10 @@ type ActiveProviderState struct { NextSyncAt time.Time } +type remoteSnapshotLoader interface { + LoadRemoteSnapshot(ctx context.Context, review core.ReviewContext) (core.ProviderSnapshot, error) +} + type ActiveProviderService struct { catalog ports.ReviewProviderCatalog factory ports.ReviewProviderClientFactory @@ -168,7 +172,15 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC if client == nil { return prev, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) } - threads, err := client.LoadRemoteThreads(ctx, review) + var snap core.ProviderSnapshot + var err error + if loader, ok := client.(remoteSnapshotLoader); ok { + snap, err = loader.LoadRemoteSnapshot(ctx, review) + } else { + var threads []core.RemoteReviewThread + threads, err = client.LoadRemoteThreads(ctx, review) + snap.Threads = threads + } if err != nil { st := prev st.LastError = err @@ -187,7 +199,15 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC } now := time.Now().UTC() next := now.Add(s.poll.Interval) - snap := core.ProviderSnapshot{StableProviderKey: key, RuntimeProviderID: runtimeID, ContextKey: core.NewReviewContextKey(key, review), Threads: threads, FetchedAt: now, Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusSynced, LastSyncAt: new(now), NextSyncAt: new(next)}} + if snap.RuntimeProviderID == "" { + snap.RuntimeProviderID = runtimeID + } + snap.StableProviderKey = key + snap.ContextKey = core.NewReviewContextKey(key, review) + if snap.FetchedAt.IsZero() { + snap.FetchedAt = now + } + snap.Sync = core.ProviderSyncState{Status: core.ProviderSyncStatusSynced, LastSyncAt: new(now), NextSyncAt: new(next)} if s.cache != nil { _ = s.cache.SaveProviderSnapshot(ctx, snap) } diff --git a/internal/core/provider_snapshot.go b/internal/core/provider_snapshot.go index cea186e..a9d7a97 100644 --- a/internal/core/provider_snapshot.go +++ b/internal/core/provider_snapshot.go @@ -81,9 +81,36 @@ type ProviderSnapshot struct { // ProviderOverview is a placeholder for richer provider overview data. type ProviderOverview struct { - Title string `json:"title,omitempty"` - ExternalURL string `json:"external_url,omitempty"` - Body string `json:"body,omitempty"` + RuntimeProviderID string `json:"runtime_provider_id,omitempty"` + Title string `json:"title,omitempty"` + Number int `json:"number,omitempty"` + State string `json:"state,omitempty"` + ExternalURL string `json:"external_url,omitempty"` + Author string `json:"author,omitempty"` + Body string `json:"body,omitempty"` + BaseRef string `json:"base_ref,omitempty"` + HeadRef string `json:"head_ref,omitempty"` + UpdatedAt *time.Time `json:"updated_at,omitempty"` + Comments []ProviderIssueComment `json:"comments,omitempty"` + Reviews []ProviderReviewSummary `json:"reviews,omitempty"` +} + +type ProviderIssueComment struct { + ExternalID string `json:"external_id"` + Author string `json:"author,omitempty"` + Body string `json:"body"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + ExternalURL string `json:"external_url,omitempty"` +} + +type ProviderReviewSummary struct { + ExternalID string `json:"external_id"` + Author string `json:"author,omitempty"` + State string `json:"state,omitempty"` + Body string `json:"body,omitempty"` + SubmittedAt time.Time `json:"submitted_at"` + ExternalURL string `json:"external_url,omitempty"` } type ProviderSyncStatus string diff --git a/internal/core/review_provider.go b/internal/core/review_provider.go index 4bce51b..7a49e3f 100644 --- a/internal/core/review_provider.go +++ b/internal/core/review_provider.go @@ -27,6 +27,7 @@ const ( // ReviewProviderCapabilities declares what a provider can do. type ReviewProviderCapabilities struct { LoadRemoteComments bool `json:"load_remote_comments"` + LoadRemoteSnapshot bool `json:"load_remote_snapshot,omitempty"` PublishReview bool `json:"publish_review"` Decisions []ReviewDecision `json:"decisions"` IdempotentPublish bool `json:"idempotent_publish"` diff --git a/pkg/plugin/plugin.go b/pkg/plugin/plugin.go index 2951459..739c0d0 100644 --- a/pkg/plugin/plugin.go +++ b/pkg/plugin/plugin.go @@ -46,12 +46,17 @@ type ( DetectionResult = protocol.DetectionResult ) -// Remote thread types +// Remote thread and snapshot types type ( - LoadRemoteThreadsRequest = protocol.LoadRemoteThreadsRequest - LoadRemoteThreadsResult = protocol.LoadRemoteThreadsResult - RemoteReviewThread = protocol.RemoteReviewThread - RemoteReviewComment = protocol.RemoteReviewComment + LoadRemoteThreadsRequest = protocol.LoadRemoteThreadsRequest + LoadRemoteThreadsResult = protocol.LoadRemoteThreadsResult + LoadRemoteSnapshotRequest = protocol.LoadRemoteSnapshotRequest + LoadRemoteSnapshotResult = protocol.LoadRemoteSnapshotResult + ProviderOverview = protocol.ProviderOverview + ProviderIssueComment = protocol.ProviderIssueComment + ProviderReviewSummary = protocol.ProviderReviewSummary + RemoteReviewThread = protocol.RemoteReviewThread + RemoteReviewComment = protocol.RemoteReviewComment ) // Publish types diff --git a/pkg/plugin/protocol/types.go b/pkg/plugin/protocol/types.go index 26f1494..8498a45 100644 --- a/pkg/plugin/protocol/types.go +++ b/pkg/plugin/protocol/types.go @@ -58,6 +58,7 @@ type ReviewProviderInfo struct { // ReviewProviderCapabilities declares what a provider can do. type ReviewProviderCapabilities struct { LoadRemoteComments bool `json:"load_remote_comments"` + LoadRemoteSnapshot bool `json:"load_remote_snapshot,omitempty"` PublishReview bool `json:"publish_review"` Decisions []ReviewDecision `json:"decisions"` IdempotentPublish bool `json:"idempotent_publish"` @@ -102,6 +103,54 @@ type LoadRemoteThreadsResult struct { Threads []RemoteReviewThread `json:"threads"` } +// ---- load_remote_snapshot ---- + +type LoadRemoteSnapshotRequest struct { + Context ReviewContext `json:"context"` +} + +type LoadRemoteSnapshotResult struct { + RuntimeProviderID string `json:"runtime_provider_id,omitempty"` + Threads []RemoteReviewThread `json:"threads"` + Overview *ProviderOverview `json:"overview,omitempty"` + Metadata map[string]string `json:"metadata,omitempty"` + FetchedAt *time.Time `json:"fetched_at,omitempty"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` +} + +type ProviderOverview struct { + RuntimeProviderID string `json:"runtime_provider_id,omitempty"` + Title string `json:"title,omitempty"` + Number int `json:"number,omitempty"` + State string `json:"state,omitempty"` + ExternalURL string `json:"external_url,omitempty"` + Author string `json:"author,omitempty"` + Body string `json:"body,omitempty"` + BaseRef string `json:"base_ref,omitempty"` + HeadRef string `json:"head_ref,omitempty"` + UpdatedAt *time.Time `json:"updated_at,omitempty"` + Comments []ProviderIssueComment `json:"comments,omitempty"` + Reviews []ProviderReviewSummary `json:"reviews,omitempty"` +} + +type ProviderIssueComment struct { + ExternalID string `json:"external_id"` + Author string `json:"author,omitempty"` + Body string `json:"body"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + ExternalURL string `json:"external_url,omitempty"` +} + +type ProviderReviewSummary struct { + ExternalID string `json:"external_id"` + Author string `json:"author,omitempty"` + State string `json:"state,omitempty"` + Body string `json:"body,omitempty"` + SubmittedAt time.Time `json:"submitted_at"` + ExternalURL string `json:"external_url,omitempty"` +} + // ---- publish_review ---- // PublishReviewParams holds the payload for a publish request. diff --git a/pkg/plugin/server.go b/pkg/plugin/server.go index 046ece1..c1b7fe8 100644 --- a/pkg/plugin/server.go +++ b/pkg/plugin/server.go @@ -20,6 +20,10 @@ type ReviewProvider interface { PublishReview(ctx context.Context, req PublishReviewParams) (PublishReviewResultData, error) } +type RemoteSnapshotProvider interface { + LoadRemoteSnapshot(ctx context.Context, req LoadRemoteSnapshotRequest) (LoadRemoteSnapshotResult, error) +} + // ServeReviewProvider runs the JSON-lines protocol server loop. It reads // requests from stdin, dispatches to the provider, and writes responses to // stdout. The server returns after EOF on stdin or when ctx is cancelled. @@ -68,6 +72,8 @@ func dispatch(ctx context.Context, provider ReviewProvider, req protocol.Request return handleDetectContext(ctx, provider, req) case "load_remote_threads": return handleLoadRemoteThreads(ctx, provider, req) + case "load_remote_snapshot": + return handleLoadRemoteSnapshot(ctx, provider, req) case "publish_review": return handlePublishReview(ctx, provider, req) default: @@ -129,6 +135,22 @@ func handleLoadRemoteThreads(ctx context.Context, provider ReviewProvider, req p return protocol.Response{ID: req.ID, Result: result} } +func handleLoadRemoteSnapshot(ctx context.Context, provider ReviewProvider, req protocol.Request) protocol.Response { + snapshotProvider, ok := provider.(RemoteSnapshotProvider) + if !ok { + return protocol.Response{ID: req.ID, Error: protocol.NewError(ErrorUnsupportedCapability, "method not supported: load_remote_snapshot")} + } + var params protocol.LoadRemoteSnapshotRequest + if err := json.Unmarshal(req.Params, ¶ms); err != nil { + return protocol.Response{ID: req.ID, Error: protocol.NewError(ErrorInvalidRequest, "invalid load_remote_snapshot params: "+err.Error())} + } + result, err := snapshotProvider.LoadRemoteSnapshot(ctx, params) + if err != nil { + return protocol.Response{ID: req.ID, Error: toProtocolError(err)} + } + return protocol.Response{ID: req.ID, Result: result} +} + func handlePublishReview(ctx context.Context, provider ReviewProvider, req protocol.Request) protocol.Response { var params protocol.PublishReviewParams if err := json.Unmarshal(req.Params, ¶ms); err != nil { diff --git a/pkg/plugin/server_test.go b/pkg/plugin/server_test.go index d72ddf2..bb6fea0 100644 --- a/pkg/plugin/server_test.go +++ b/pkg/plugin/server_test.go @@ -25,6 +25,17 @@ type fakeProvider struct { methodsCalled []string } +type fakeSnapshotProvider struct { + *fakeProvider + snapshotResult plugin.LoadRemoteSnapshotResult + snapshotErr error +} + +func (f *fakeSnapshotProvider) LoadRemoteSnapshot(_ context.Context, _ plugin.LoadRemoteSnapshotRequest) (plugin.LoadRemoteSnapshotResult, error) { + f.methodsCalled = append(f.methodsCalled, "load_remote_snapshot") + return f.snapshotResult, f.snapshotErr +} + func (f *fakeProvider) Initialize(_ context.Context, _ plugin.InitializeRequest) (plugin.InitializeResult, error) { f.methodsCalled = append(f.methodsCalled, "initialize") return f.initResult, f.initErr @@ -273,6 +284,60 @@ func TestServeReviewProviderLargeRequestLine(t *testing.T) { } } +func TestLoadRemoteSnapshotDispatchesOptionalMethod(t *testing.T) { + provider := &fakeSnapshotProvider{fakeProvider: &fakeProvider{}, snapshotResult: plugin.LoadRemoteSnapshotResult{RuntimeProviderID: "github", Overview: &plugin.ProviderOverview{Title: "PR"}}} + input := bytes.NewBufferString(`{"id":"s1","method":"load_remote_snapshot","params":{"context":{"repository":{"repo_path":"."}}}}` + "\n") + var output bytes.Buffer + if err := plugin.ServeReviewProvider(context.Background(), provider, input, &output); err != nil { + t.Fatalf("ServeReviewProvider returned error: %v", err) + } + var response struct { + ID string `json:"id"` + Result plugin.LoadRemoteSnapshotResult `json:"result"` + Error *plugin.Error `json:"error,omitempty"` + } + if err := json.Unmarshal(output.Bytes(), &response); err != nil { + t.Fatalf("invalid json: %v", err) + } + if response.Error != nil || response.Result.Overview == nil || response.Result.Overview.Title != "PR" { + t.Fatalf("unexpected response: %#v raw=%s", response, output.String()) + } + if !reflect.DeepEqual(provider.methodsCalled, []string{"load_remote_snapshot"}) { + t.Fatalf("expected load_remote_snapshot call, got %v", provider.methodsCalled) + } +} + +func TestLoadRemoteSnapshotUnsupportedKeepsOldProvidersValid(t *testing.T) { + provider := &fakeProvider{threadsResult: plugin.LoadRemoteThreadsResult{Threads: []plugin.RemoteReviewThread{{ProviderID: "old", ExternalID: "t1"}}}} + input := bytes.NewBufferString(strings.Join([]string{ + `{"id":"s1","method":"load_remote_snapshot","params":{"context":{"repository":{"repo_path":"."}}}}`, + `{"id":"t1","method":"load_remote_threads","params":{"context":{"repository":{"repo_path":"."}}}}`, + }, "\n") + "\n") + var output bytes.Buffer + if err := plugin.ServeReviewProvider(context.Background(), provider, input, &output); err != nil { + t.Fatalf("ServeReviewProvider returned error: %v", err) + } + dec := json.NewDecoder(&output) + var unsupported plugin.Response + if err := dec.Decode(&unsupported); err != nil { + t.Fatalf("decode unsupported: %v", err) + } + if unsupported.Error == nil || unsupported.Error.Code != plugin.ErrorUnsupportedCapability { + t.Fatalf("expected unsupported capability, got %#v", unsupported.Error) + } + var fallback struct { + ID string `json:"id"` + Result plugin.LoadRemoteThreadsResult `json:"result"` + Error *plugin.Error `json:"error,omitempty"` + } + if err := dec.Decode(&fallback); err != nil { + t.Fatalf("decode fallback: %v", err) + } + if fallback.Error != nil || len(fallback.Result.Threads) != 1 { + t.Fatalf("expected fallback threads, got %#v", fallback) + } +} + func TestServeReviewProviderRoundTripSuccess(t *testing.T) { t.Parallel() From bba937dde3bf08d09d5e638e3035adb4441185aa Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 16:07:12 +0200 Subject: [PATCH 05/22] feat(github): sync pull request review context --- go.mod | 5 + go.sum | 5 + .../github/cmd/ero-plugin-github/graphql.go | 237 ++++++++++++++++++ .../cmd/ero-plugin-github/graphql_test.go | 163 ++++++++++++ plugins/github/cmd/ero-plugin-github/main.go | 123 +++++++-- .../github/cmd/ero-plugin-github/main_test.go | 160 +++++++++++- plugins/github/cmd/ero-plugin-github/map.go | 112 +++++++++ plugins/github/cmd/ero-plugin-github/match.go | 79 ++++++ .../github/cmd/ero-plugin-github/remote.go | 65 +++++ 9 files changed, 923 insertions(+), 26 deletions(-) create mode 100644 plugins/github/cmd/ero-plugin-github/graphql.go create mode 100644 plugins/github/cmd/ero-plugin-github/graphql_test.go create mode 100644 plugins/github/cmd/ero-plugin-github/map.go create mode 100644 plugins/github/cmd/ero-plugin-github/match.go create mode 100644 plugins/github/cmd/ero-plugin-github/remote.go diff --git a/go.mod b/go.mod index a6da7a8..2e0c08e 100644 --- a/go.mod +++ b/go.mod @@ -26,6 +26,7 @@ require ( github.com/Microsoft/go-winio v0.6.2 // indirect github.com/ProtonMail/go-crypto v1.1.6 // indirect github.com/atotto/clipboard v0.1.4 // indirect + github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect github.com/aymerick/douceur v0.2.0 // indirect github.com/charmbracelet/colorprofile v0.4.3 // indirect github.com/charmbracelet/ultraviolet v0.0.0-20260416155717-489999b90468 // indirect @@ -34,6 +35,7 @@ require ( github.com/charmbracelet/x/termios v0.1.1 // indirect github.com/charmbracelet/x/windows v0.2.2 // indirect github.com/cli/safeexec v1.0.0 // indirect + github.com/cli/shurcooL-graphql v0.0.4 // indirect github.com/clipperhouse/displaywidth v0.11.0 // indirect github.com/clipperhouse/uax29/v2 v2.7.0 // indirect github.com/cloudflare/circl v1.6.3 // indirect @@ -47,6 +49,7 @@ require ( github.com/go-viper/mapstructure/v2 v2.4.0 // indirect github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8 // indirect github.com/gorilla/css v1.0.1 // indirect + github.com/henvic/httpretty v0.0.6 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/jbenet/go-context v0.0.0-20150711004518-d14ea06fba99 // indirect github.com/kevinburke/ssh_config v1.2.0 // indirect @@ -57,6 +60,7 @@ require ( github.com/mattn/go-runewidth v0.0.23 // indirect github.com/microcosm-cc/bluemonday v1.0.27 // indirect github.com/muesli/cancelreader v0.2.2 // indirect + github.com/muesli/termenv v0.16.0 // indirect github.com/pjbgf/sha1cd v0.6.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect github.com/rivo/uniseg v0.4.7 // indirect @@ -69,6 +73,7 @@ require ( github.com/spf13/cast v1.10.0 // indirect github.com/stretchr/objx v0.5.2 // indirect github.com/subosito/gotenv v1.6.0 // indirect + github.com/thlib/go-timezone-local v0.0.0-20210907160436-ef149e42d28e // indirect github.com/xanzy/ssh-agent v0.3.3 // indirect github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect github.com/yuin/goldmark v1.7.8 // indirect diff --git a/go.sum b/go.sum index 56d9dc5..e35d114 100644 --- a/go.sum +++ b/go.sum @@ -105,6 +105,8 @@ github.com/google/shlex v0.0.0-20191202100458-e7afc7fbc510 h1:El6M4kTTCOh6aBiKaU github.com/google/shlex v0.0.0-20191202100458-e7afc7fbc510/go.mod h1:pupxD2MaaD3pAXIBCelhxNneeOaAeabZDe5s4K6zSpQ= github.com/gorilla/css v1.0.1 h1:ntNaBIghp6JmvWnxbZKANoLyuXTPZ4cAMlo6RyhlbO8= github.com/gorilla/css v1.0.1/go.mod h1:BvnYkspnSzMmwRK+b8/xgNPLiIuNZr6vbZBTPQ2A3b0= +github.com/h2non/parth v0.0.0-20190131123155-b4df798d6542 h1:2VTzZjLZBgl62/EtslCrtky5vbi9dd7HrQPQIx6wqiw= +github.com/h2non/parth v0.0.0-20190131123155-b4df798d6542/go.mod h1:Ow0tF8D4Kplbc8s8sSb3V2oUCygFHVp8gC3Dn6U4MNI= github.com/henvic/httpretty v0.0.6 h1:JdzGzKZBajBfnvlMALXXMVQWxWMF/ofTy8C3/OSUTxs= github.com/henvic/httpretty v0.0.6/go.mod h1:X38wLjWXHkXT7r2+uK8LjCMne9rsuNaBLJ+5cU2/Pmo= github.com/hexops/gotextdiff v1.0.3 h1:gitA9+qJrrTCsiCl7+kh75nPqQt1cx4ZkudSTLoUqJM= @@ -217,6 +219,7 @@ golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7w golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20210831042530-f4d43177bf5e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= @@ -234,6 +237,8 @@ gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8 gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= +gopkg.in/h2non/gock.v1 v1.1.2 h1:jBbHXgGBK/AoPVfJh5x4r/WxIrElvbLel8TCZkkZJoY= +gopkg.in/h2non/gock.v1 v1.1.2/go.mod h1:n7UGz/ckNChHiK05rDoiC4MYSunEC/lyaUm2WWaDva0= gopkg.in/natefinch/lumberjack.v2 v2.2.1 h1:bBRl1b0OH9s/DuPhuXpNl+VtCaJXFZ5/uEFST95x9zc= gopkg.in/natefinch/lumberjack.v2 v2.2.1/go.mod h1:YD8tP3GAjkrDg1eZH7EGmyESg/lsYskCTPBJVb9jqSc= gopkg.in/warnings.v0 v0.1.2 h1:wFXVbFY8DY5/xOe1ECiWdKCzZlxgshcYVNkBHstARME= diff --git a/plugins/github/cmd/ero-plugin-github/graphql.go b/plugins/github/cmd/ero-plugin-github/graphql.go new file mode 100644 index 0000000..3c6e12b --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/graphql.go @@ -0,0 +1,237 @@ +package main + +import ( + "context" + "time" + + "github.com/cli/go-gh/v2/pkg/api" + + "ero/pkg/plugin" +) + +type graphQLDoer interface { + DoWithContext(ctx context.Context, query string, variables map[string]any, response any) error +} + +func defaultGraphQLClient() (graphQLDoer, error) { return api.DefaultGraphQLClient() } + +type graphQLClientFactory func() (graphQLDoer, error) + +type githubPRSnapshot struct { + Number int + URL string + Title string + State string + Body string + Author string + BaseRef string + HeadRef string + HeadRepoOwner string + HeadRepoName string + HeadSHA string + UpdatedAt *time.Time + IssueComments []plugin.ProviderIssueComment + Reviews []plugin.ProviderReviewSummary + Threads []plugin.RemoteReviewThread +} + +type ghPageInfo struct { + HasNextPage bool `json:"hasNextPage"` + EndCursor string `json:"endCursor"` +} + +type ghActor struct { + Login string `json:"login"` +} + +type ghPRListResponse struct { + Repository struct { + PullRequests struct { + Nodes []ghPRNode `json:"nodes"` + PageInfo ghPageInfo `json:"pageInfo"` + } `json:"pullRequests"` + } `json:"repository"` +} + +type ghPRSnapshotResponse struct { + Repository struct { + PullRequest ghPRNode `json:"pullRequest"` + } `json:"repository"` +} + +type ghPRNode struct { + Number int `json:"number"` + URL string `json:"url"` + Title string `json:"title"` + State string `json:"state"` + Body string `json:"body"` + Author *ghActor `json:"author"` + BaseRefName string `json:"baseRefName"` + HeadRefName string `json:"headRefName"` + HeadRefOid string `json:"headRefOid"` + UpdatedAt *time.Time `json:"updatedAt"` + HeadRepositoryOwner *ghActor `json:"headRepositoryOwner"` + HeadRepository *struct { + Name string `json:"name"` + } `json:"headRepository"` + Comments struct { + Nodes []ghIssueComment `json:"nodes"` + PageInfo ghPageInfo `json:"pageInfo"` + } `json:"comments"` + Reviews struct { + Nodes []ghReview `json:"nodes"` + PageInfo ghPageInfo `json:"pageInfo"` + } `json:"reviews"` + ReviewThreads struct { + Nodes []ghReviewThread `json:"nodes"` + PageInfo ghPageInfo `json:"pageInfo"` + } `json:"reviewThreads"` +} + +type ghIssueComment struct { + ID string `json:"id"` + URL string `json:"url"` + Body string `json:"body"` + Author *ghActor `json:"author"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` +} +type ghReview struct { + ID string `json:"id"` + URL string `json:"url"` + State string `json:"state"` + Body string `json:"body"` + Author *ghActor `json:"author"` + SubmittedAt time.Time `json:"submittedAt"` +} +type ghReviewThread struct { + ID string `json:"id"` + Path string `json:"path"` + Line int `json:"line"` + Side string `json:"side"` + StartLine int `json:"startLine"` + StartSide string `json:"startSide"` + IsOutdated bool `json:"isOutdated"` + Comments struct { + Nodes []ghReviewThreadComment `json:"nodes"` + PageInfo ghPageInfo `json:"pageInfo"` + } `json:"comments"` +} +type ghReviewThreadComment struct { + ID string `json:"id"` + URL string `json:"url"` + Body string `json:"body"` + Author *ghActor `json:"author"` + CreatedAt time.Time `json:"createdAt"` +} + +const githubPRListQuery = `query EroPRList($owner:String!, $name:String!, $after:String) { repository(owner:$owner, name:$name) { pullRequests(first:50, after:$after, states:[OPEN]) { nodes { number url title state baseRefName headRefName headRefOid headRepositoryOwner { login } headRepository { name } } pageInfo { hasNextPage endCursor } } } } }` +const githubPRSnapshotQuery = `query EroPRSnapshot($owner:String!, $name:String!, $number:Int!, $commentsAfter:String, $reviewsAfter:String, $threadsAfter:String) { repository(owner:$owner, name:$name) { pullRequest(number:$number) { number url title state body updatedAt author { login } baseRefName headRefName headRefOid headRepositoryOwner { login } headRepository { name } comments(first:100, after:$commentsAfter) { nodes { id url body createdAt updatedAt author { login } } pageInfo { hasNextPage endCursor } } reviews(first:100, after:$reviewsAfter) { nodes { id url state body submittedAt author { login } } pageInfo { hasNextPage endCursor } } reviewThreads(first:100, after:$threadsAfter) { nodes { id path line side startLine startSide isOutdated comments(first:100) { nodes { id url body createdAt author { login } } pageInfo { hasNextPage endCursor } } } pageInfo { hasNextPage endCursor } } } } }` + +func (p githubProvider) graphQLClient() (graphQLDoer, error) { + if p.newGraphQLClient != nil { + return p.newGraphQLClient() + } + return defaultGraphQLClient() +} + +func fetchGitHubSnapshot(ctx context.Context, client graphQLDoer, remote githubRemote, reviewCtx plugin.ReviewContext) (githubPRSnapshot, error) { + candidates, err := fetchGitHubPRCandidates(ctx, client, remote) + if err != nil { + return githubPRSnapshot{}, err + } + match, err := matchGitHubPR(reviewCtx, candidates) + if err != nil { + return githubPRSnapshot{}, err + } + return fetchGitHubPRSnapshot(ctx, client, remote, match.Number) +} + +func fetchGitHubPRCandidates(ctx context.Context, client graphQLDoer, remote githubRemote) ([]githubPRCandidate, error) { + var out []githubPRCandidate + after := "" + for { + vars := map[string]any{"owner": remote.Owner, "name": remote.Name, "after": cursorValue(after)} + var resp ghPRListResponse + if err := client.DoWithContext(ctx, githubPRListQuery, vars, &resp); err != nil { + return nil, classifyGitHubRemoteError("fetch GitHub pull requests", err) + } + for _, n := range resp.Repository.PullRequests.Nodes { + out = append(out, candidateFromNode(n)) + } + pi := resp.Repository.PullRequests.PageInfo + if !pi.HasNextPage { + return out, nil + } + if pi.EndCursor == "" { + return nil, plugin.NewError(plugin.ErrorRemoteValidationFailed, "GitHub pull request pagination missing end cursor") + } + after = pi.EndCursor + } +} + +func fetchGitHubPRSnapshot(ctx context.Context, client graphQLDoer, remote githubRemote, number int) (githubPRSnapshot, error) { + var commentsAfter, reviewsAfter, threadsAfter string + commentsDone, reviewsDone, threadsDone := false, false, false + var accum githubPRSnapshot + for { + vars := map[string]any{"owner": remote.Owner, "name": remote.Name, "number": number, "commentsAfter": cursorValue(commentsAfter), "reviewsAfter": cursorValue(reviewsAfter), "threadsAfter": cursorValue(threadsAfter)} + var resp ghPRSnapshotResponse + if err := client.DoWithContext(ctx, githubPRSnapshotQuery, vars, &resp); err != nil { + return githubPRSnapshot{}, classifyGitHubRemoteError("fetch GitHub pull request snapshot", err) + } + pr := resp.Repository.PullRequest + page := mapGitHubPRMetadata(pr) + if accum.Number == 0 { + accum = page + } + if !commentsDone { + accum.IssueComments = append(accum.IssueComments, mapGitHubIssueComments(pr.Comments.Nodes)...) + } + if !reviewsDone { + accum.Reviews = append(accum.Reviews, mapGitHubReviews(pr.Reviews.Nodes)...) + } + if !threadsDone { + for _, thread := range pr.ReviewThreads.Nodes { + if thread.Comments.PageInfo.HasNextPage { + return githubPRSnapshot{}, plugin.NewError(plugin.ErrorRemoteValidationFailed, "GitHub review thread comments pagination beyond first page is not supported") + } + accum.Threads = append(accum.Threads, mapGitHubThread(thread)) + } + } + if !commentsDone && pr.Comments.PageInfo.HasNextPage { + if pr.Comments.PageInfo.EndCursor == "" { + return githubPRSnapshot{}, plugin.NewError(plugin.ErrorRemoteValidationFailed, "GitHub comments pagination missing end cursor") + } + commentsAfter = pr.Comments.PageInfo.EndCursor + } else { + commentsDone = true + } + if !reviewsDone && pr.Reviews.PageInfo.HasNextPage { + if pr.Reviews.PageInfo.EndCursor == "" { + return githubPRSnapshot{}, plugin.NewError(plugin.ErrorRemoteValidationFailed, "GitHub reviews pagination missing end cursor") + } + reviewsAfter = pr.Reviews.PageInfo.EndCursor + } else { + reviewsDone = true + } + if !threadsDone && pr.ReviewThreads.PageInfo.HasNextPage { + if pr.ReviewThreads.PageInfo.EndCursor == "" { + return githubPRSnapshot{}, plugin.NewError(plugin.ErrorRemoteValidationFailed, "GitHub reviewThreads pagination missing end cursor") + } + threadsAfter = pr.ReviewThreads.PageInfo.EndCursor + } else { + threadsDone = true + } + if commentsDone && reviewsDone && threadsDone { + return accum, nil + } + } +} + +func cursorValue(cursor string) any { + if cursor == "" { + return nil + } + return cursor +} diff --git a/plugins/github/cmd/ero-plugin-github/graphql_test.go b/plugins/github/cmd/ero-plugin-github/graphql_test.go new file mode 100644 index 0000000..2900683 --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/graphql_test.go @@ -0,0 +1,163 @@ +package main + +import ( + "context" + "maps" + "testing" + "time" + + "ero/pkg/plugin" +) + +type fakeGraphQLClient struct { + listPages []ghPRListResponse + snapshotPages []ghPRSnapshotResponse + listCalls int + snapshotCalls int + vars []map[string]any +} + +func (f *fakeGraphQLClient) DoWithContext(_ context.Context, query string, variables map[string]any, response any) error { + copied := make(map[string]any, len(variables)) + maps.Copy(copied, variables) + f.vars = append(f.vars, copied) + switch r := response.(type) { + case *ghPRListResponse: + *r = f.listPages[f.listCalls] + f.listCalls++ + case *ghPRSnapshotResponse: + *r = f.snapshotPages[f.snapshotCalls] + f.snapshotCalls++ + default: + panic("unexpected GraphQL response type") + } + _ = query + return nil +} + +func TestGitHubGraphQLMapping(t *testing.T) { + now := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) + actor := &ghActor{Login: "octo"} + pr := ghPRNode{Number: 12, URL: "https://github.com/owner/repo/pull/12", Title: "Add feature", State: "OPEN", Body: "**markdown**", Author: actor, BaseRefName: "main", HeadRefName: "feature", HeadRefOid: "abc", UpdatedAt: &now, HeadRepositoryOwner: actor, HeadRepository: &struct { + Name string `json:"name"` + }{Name: "repo"}} + pr.Comments.Nodes = []ghIssueComment{{ID: "ic1", URL: "https://c", Body: "issue body", Author: actor, CreatedAt: now, UpdatedAt: now}} + pr.Reviews.Nodes = []ghReview{{ID: "rv1", URL: "https://r", State: "COMMENTED", Body: "summary", Author: actor, SubmittedAt: now}} + pr.ReviewThreads.Nodes = []ghReviewThread{ + thread("t1", "file.go", 10, "RIGHT", 0, "", false, "c1", "single"), + thread("t2", "file.go", 20, "RIGHT", 18, "RIGHT", false, "c2", "multi"), + thread("t3", "old.go", 7, "LEFT", 0, "", false, "c3", "deleted"), + thread("t4", "gone.go", 0, "RIGHT", 0, "", true, "c4", "outdated"), + } + s := mapGitHubPRPage(pr) + if s.Number != 12 || s.Body != "**markdown**" || s.Author != "octo" || s.BaseRef != "main" || s.HeadRef != "feature" || s.UpdatedAt != &now { + t.Fatalf("metadata not mapped: %#v", s) + } + if len(s.IssueComments) != 1 || s.IssueComments[0].Body != "issue body" || len(s.Reviews) != 1 || s.Reviews[0].Body != "summary" { + t.Fatalf("overview comments/reviews not mapped: %#v", s) + } + if len(s.Threads) != 4 || s.Threads[0].Range.End.NewLineNumber != 10 || s.Threads[1].Range.Start.NewLineNumber != 18 || s.Threads[2].Range.End.OldLineNumber != 7 || !s.Threads[3].Unmapped { + t.Fatalf("threads not mapped/classified: %#v", s.Threads) + } +} + +func TestLoadRemoteSnapshotPaginatesAndMaps(t *testing.T) { + now := time.Now().UTC() + actor := &ghActor{Login: "octo"} + owner := &ghActor{Login: "owner"} + list1 := ghPRListResponse{} + list1.Repository.PullRequests.Nodes = []ghPRNode{{Number: 12, BaseRefName: "main", HeadRefName: "other"}} + list1.Repository.PullRequests.PageInfo = ghPageInfo{HasNextPage: true, EndCursor: "p2"} + list2 := ghPRListResponse{} + list2.Repository.PullRequests.Nodes = []ghPRNode{{Number: 13, BaseRefName: "main", HeadRefName: "feature", HeadRepositoryOwner: owner, HeadRepository: &struct { + Name string `json:"name"` + }{Name: "repo"}}} + page1 := ghPRSnapshotResponse{} + page1.Repository.PullRequest = ghPRNode{Number: 13, URL: "https://pr", Title: "PR", State: "OPEN", Author: actor, BaseRefName: "main", HeadRefName: "feature", UpdatedAt: &now} + page1.Repository.PullRequest.Comments.Nodes = []ghIssueComment{{ID: "i1", Body: "first", CreatedAt: now, UpdatedAt: now}} + page1.Repository.PullRequest.Comments.PageInfo = ghPageInfo{HasNextPage: true, EndCursor: "c2"} + page1.Repository.PullRequest.Reviews.PageInfo = ghPageInfo{HasNextPage: true, EndCursor: "r2"} + page1.Repository.PullRequest.ReviewThreads.PageInfo = ghPageInfo{HasNextPage: true, EndCursor: "t2"} + page2 := ghPRSnapshotResponse{} + page2.Repository.PullRequest = page1.Repository.PullRequest + page2.Repository.PullRequest.Comments.PageInfo = ghPageInfo{} + page2.Repository.PullRequest.Reviews.PageInfo = ghPageInfo{} + page2.Repository.PullRequest.ReviewThreads.PageInfo = ghPageInfo{} + page2.Repository.PullRequest.Comments.Nodes = []ghIssueComment{{ID: "i2", Body: "second", CreatedAt: now, UpdatedAt: now}} + page2.Repository.PullRequest.Reviews.Nodes = []ghReview{{ID: "r1", Body: "review", SubmittedAt: now}} + page2.Repository.PullRequest.ReviewThreads.Nodes = []ghReviewThread{thread("t", "file.go", 3, "RIGHT", 0, "", false, "tc", "body")} + fake := &fakeGraphQLClient{listPages: []ghPRListResponse{list1, list2}, snapshotPages: []ghPRSnapshotResponse{page1, page2}} + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return fake, nil }} + got, err := provider.LoadRemoteSnapshot(context.Background(), plugin.LoadRemoteSnapshotRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "git@github.com:owner/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}}) + if err != nil { + t.Fatalf("LoadRemoteSnapshot returned error: %v", err) + } + if fake.listCalls != 2 || fake.snapshotCalls != 2 || got.Overview.Number != 13 || len(got.Overview.Comments) != 2 || len(got.Overview.Reviews) != 1 || len(got.Threads) != 1 || got.Threads[0].Comments[0].Body != "body" { + t.Fatalf("unexpected snapshot: calls %d/%d %#v", fake.listCalls, fake.snapshotCalls, got) + } +} + +func TestLoadRemoteSnapshotDoesNotRefetchCompletedCollections(t *testing.T) { + now := time.Now().UTC() + list := ghPRListResponse{} + list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 1, BaseRefName: "main", HeadRefName: "feature"}} + page1 := ghPRSnapshotResponse{} + page1.Repository.PullRequest = ghPRNode{Number: 1, Title: "PR"} + page1.Repository.PullRequest.Comments.Nodes = []ghIssueComment{{ID: "i1", Body: "first", CreatedAt: now, UpdatedAt: now}} + page1.Repository.PullRequest.Comments.PageInfo = ghPageInfo{} + page1.Repository.PullRequest.ReviewThreads.Nodes = []ghReviewThread{thread("t1", "file.go", 1, "RIGHT", 0, "", false, "c1", "first thread")} + page1.Repository.PullRequest.ReviewThreads.PageInfo = ghPageInfo{HasNextPage: true, EndCursor: "t2"} + page2 := ghPRSnapshotResponse{} + page2.Repository.PullRequest = ghPRNode{Number: 1, Title: "PR"} + page2.Repository.PullRequest.Comments.Nodes = []ghIssueComment{{ID: "i1-duplicate-if-refetched", Body: "duplicate", CreatedAt: now, UpdatedAt: now}} + page2.Repository.PullRequest.Comments.PageInfo = ghPageInfo{} + page2.Repository.PullRequest.ReviewThreads.Nodes = []ghReviewThread{thread("t2", "file.go", 2, "RIGHT", 0, "", false, "c2", "second thread")} + fake := &fakeGraphQLClient{listPages: []ghPRListResponse{list}, snapshotPages: []ghPRSnapshotResponse{page1, page2}} + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return fake, nil }} + got, err := provider.LoadRemoteSnapshot(context.Background(), plugin.LoadRemoteSnapshotRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "https://github.com/owner/repo"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}}) + if err != nil { + t.Fatalf("LoadRemoteSnapshot returned error: %v", err) + } + if len(got.Overview.Comments) != 1 || got.Overview.Comments[0].ExternalID != "i1" || len(got.Threads) != 2 { + t.Fatalf("completed comments should not be re-appended while threads page: %#v", got) + } + if len(fake.vars) < 3 || fake.vars[1]["commentsAfter"] != nil || fake.vars[2]["commentsAfter"] != nil || fake.vars[2]["threadsAfter"] != "t2" { + t.Fatalf("unexpected pagination variables: %#v", fake.vars) + } +} + +func TestLoadRemoteSnapshotRejectsUnpaginatedThreadComments(t *testing.T) { + list := ghPRListResponse{} + list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 1, BaseRefName: "main", HeadRefName: "feature"}} + page := ghPRSnapshotResponse{} + page.Repository.PullRequest = ghPRNode{Number: 1} + thread := thread("t", "file.go", 1, "RIGHT", 0, "", false, "c", "body") + thread.Comments.PageInfo = ghPageInfo{HasNextPage: true, EndCursor: "more"} + page.Repository.PullRequest.ReviewThreads.Nodes = []ghReviewThread{thread} + fake := &fakeGraphQLClient{listPages: []ghPRListResponse{list}, snapshotPages: []ghPRSnapshotResponse{page}} + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return fake, nil }} + _, err := provider.LoadRemoteSnapshot(context.Background(), plugin.LoadRemoteSnapshotRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "https://github.com/owner/repo"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}}) + if plugin.AsError(err) == nil || plugin.AsError(err).Code != plugin.ErrorRemoteValidationFailed { + t.Fatalf("expected remote_validation_failed for nested comment pagination, got %v", err) + } +} + +func TestLoadRemoteThreadsUsesSnapshotThreads(t *testing.T) { + list := ghPRListResponse{} + list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 1, BaseRefName: "main", HeadRefName: "feature"}} + page := ghPRSnapshotResponse{} + page.Repository.PullRequest = ghPRNode{Number: 1} + page.Repository.PullRequest.ReviewThreads.Nodes = []ghReviewThread{thread("t", "f", 1, "RIGHT", 0, "", false, "c", "b")} + fake := &fakeGraphQLClient{listPages: []ghPRListResponse{list}, snapshotPages: []ghPRSnapshotResponse{page}} + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return fake, nil }} + got, err := provider.LoadRemoteThreads(context.Background(), plugin.LoadRemoteThreadsRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "https://github.com/owner/repo"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}}) + if err != nil || len(got.Threads) != 1 { + t.Fatalf("LoadRemoteThreads = %#v, %v", got, err) + } +} + +func thread(id, path string, line int, side string, start int, startSide string, outdated bool, cid, body string) ghReviewThread { + t := ghReviewThread{ID: id, Path: path, Line: line, Side: side, StartLine: start, StartSide: startSide, IsOutdated: outdated} + t.Comments.Nodes = []ghReviewThreadComment{{ID: cid, URL: "https://thread", Body: body, Author: &ghActor{Login: "reviewer"}, CreatedAt: time.Now().UTC()}} + return t +} diff --git a/plugins/github/cmd/ero-plugin-github/main.go b/plugins/github/cmd/ero-plugin-github/main.go index eac361f..c25d1a4 100644 --- a/plugins/github/cmd/ero-plugin-github/main.go +++ b/plugins/github/cmd/ero-plugin-github/main.go @@ -17,8 +17,9 @@ import ( const providerID = "github" type githubProvider struct { - getenv func(string) string - execGH func(context.Context, ...string) (string, string, error) + getenv func(string) string + execGH func(context.Context, ...string) (string, string, error) + newGraphQLClient graphQLClientFactory } type ghPR struct { @@ -53,7 +54,8 @@ func (p githubProvider) Initialize(_ context.Context, req plugin.InitializeReque Label: "GitHub", Name: "ero-plugin-github", Capabilities: plugin.ReviewProviderCapabilities{ - LoadRemoteComments: false, + LoadRemoteComments: true, + LoadRemoteSnapshot: true, PublishReview: true, Decisions: []plugin.ReviewDecision{ plugin.ReviewDecisionComment, @@ -66,17 +68,48 @@ func (p githubProvider) Initialize(_ context.Context, req plugin.InitializeReque }, nil } -func (p githubProvider) DetectContext(_ context.Context, req plugin.DetectContextRequest) (plugin.DetectContextResult, error) { - for _, remote := range req.Context.Repository.Remotes { - if isGitHubRemote(remote.URL) { - return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: true, Reason: "GitHub remote detected"}}, nil - } +func (p githubProvider) DetectContext(ctx context.Context, req plugin.DetectContextRequest) (plugin.DetectContextResult, error) { + remote, ok := firstGitHubRemote(req.Context.Repository.Remotes) + if !ok { + return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: "no GitHub remote detected"}}, nil + } + client, err := p.graphQLClient() + if err != nil { + return plugin.DetectContextResult{}, plugin.NewErrorf(plugin.ErrorAuthRequired, "create GitHub GraphQL client: %v", err) + } + candidates, err := fetchGitHubPRCandidates(ctx, client, remote) + if err != nil { + return plugin.DetectContextResult{}, err + } + match, err := matchGitHubPR(req.Context, candidates) + if err != nil { + return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: err.Error()}}, nil + } + return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: true, Reason: "matched GitHub pull request " + githubPRSummary(match)}}, nil +} + +func (p githubProvider) LoadRemoteSnapshot(ctx context.Context, req plugin.LoadRemoteSnapshotRequest) (plugin.LoadRemoteSnapshotResult, error) { + remote, ok := firstGitHubRemote(req.Context.Repository.Remotes) + if !ok { + return plugin.LoadRemoteSnapshotResult{}, plugin.NewError(plugin.ErrorNotApplicable, "no GitHub remote detected") + } + client, err := p.graphQLClient() + if err != nil { + return plugin.LoadRemoteSnapshotResult{}, plugin.NewErrorf(plugin.ErrorAuthRequired, "create GitHub GraphQL client: %v", err) } - return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: "no GitHub remote detected"}}, nil + snapshot, err := fetchGitHubSnapshot(ctx, client, remote, req.Context) + if err != nil { + return plugin.LoadRemoteSnapshotResult{}, err + } + return snapshotResultFromGitHub(snapshot), nil } -func (p githubProvider) LoadRemoteThreads(_ context.Context, _ plugin.LoadRemoteThreadsRequest) (plugin.LoadRemoteThreadsResult, error) { - return plugin.LoadRemoteThreadsResult{}, plugin.NewError(plugin.ErrorUnsupportedCapability, "GitHub remote comment loading is not implemented yet") +func (p githubProvider) LoadRemoteThreads(ctx context.Context, req plugin.LoadRemoteThreadsRequest) (plugin.LoadRemoteThreadsResult, error) { + snapshot, err := p.LoadRemoteSnapshot(ctx, plugin.LoadRemoteSnapshotRequest{Context: req.Context}) + if err != nil { + return plugin.LoadRemoteThreadsResult{}, err + } + return plugin.LoadRemoteThreadsResult{Threads: snapshot.Threads}, nil } func (p githubProvider) PublishReview(ctx context.Context, req plugin.PublishReviewParams) (plugin.PublishReviewResultData, error) { @@ -86,7 +119,7 @@ func (p githubProvider) PublishReview(ctx context.Context, req plugin.PublishRev ctx, cancel := context.WithTimeout(ctx, 15*time.Second) defer cancel() - pr, err := p.currentPullRequest(ctx) + pr, err := p.currentPullRequest(ctx, req.Payload.Context) if err != nil { return plugin.PublishReviewResultData{}, err } @@ -100,7 +133,7 @@ func (p githubProvider) PublishReview(ctx context.Context, req plugin.PublishRev if message == "" { message = err.Error() } - return plugin.PublishReviewResultData{}, plugin.NewErrorf(plugin.ErrorNetwork, "publish GitHub review: %s", message) + return plugin.PublishReviewResultData{}, classifyGitHubRemoteMessage("publish GitHub review", message) } var response ghReviewResponse if err := json.Unmarshal([]byte(stdout), &response); err != nil { @@ -117,20 +150,33 @@ func (p githubProvider) PublishReview(ctx context.Context, req plugin.PublishRev }}, nil } -func (p githubProvider) currentPullRequest(ctx context.Context) (ghPR, error) { +func (p githubProvider) currentPullRequest(ctx context.Context, reviewCtx plugin.ReviewContext) (ghPR, error) { + if remote, ok := firstGitHubRemote(reviewCtx.Repository.Remotes); ok { + if client, err := p.graphQLClient(); err == nil { + candidates, err := fetchGitHubPRCandidates(ctx, client, remote) + if err != nil { + return ghPR{}, err + } + match, err := matchGitHubPR(reviewCtx, candidates) + if err != nil { + return ghPR{}, err + } + return ghPR{Number: match.Number, URL: match.URL}, nil + } + } if p.execGH == nil { p.execGH = execGH } - stdout, stderr, err := p.execGH(ctx, "pr", "view", "--json", "number,url") + stdout, stderr, err := p.execGH(ctx, ghPRViewArgs(reviewCtx)...) if err != nil { message := strings.TrimSpace(stderr) if message == "" { message = err.Error() } if strings.Contains(strings.ToLower(message), "no pull request") || strings.Contains(strings.ToLower(message), "no pull requests") { - return ghPR{}, plugin.NewErrorf(plugin.ErrorNotApplicable, "no pull request found for current branch: %s", message) + return ghPR{}, plugin.NewErrorf(plugin.ErrorNotApplicable, "no pull request found for review context: %s", message) } - return ghPR{}, plugin.NewErrorf(plugin.ErrorAuthRequired, "GitHub CLI PR lookup failed; ensure gh is installed and authenticated: %s", message) + return ghPR{}, classifyGitHubRemoteMessage("GitHub CLI PR lookup failed", message) } var pr ghPR if err := json.Unmarshal([]byte(stdout), &pr); err != nil { @@ -142,6 +188,44 @@ func (p githubProvider) currentPullRequest(ctx context.Context) (ghPR, error) { return pr, nil } +func ghPRViewArgs(reviewCtx plugin.ReviewContext) []string { + args := []string{"pr", "view", "--json", "number,url"} + branch := publishPRLookupBranch(reviewCtx) + if branch != "" { + args = append(args, branch) + } + return args +} + +func publishPRLookupBranch(reviewCtx plugin.ReviewContext) string { + headRef := strings.TrimSpace(reviewCtx.Target.HeadRef) + if headRef == "" && strings.EqualFold(reviewCtx.Target.Mode, "branch") { + headRef = strings.TrimSpace(reviewCtx.Repository.CurrentBranch) + } + if headRef == "" { + return "" + } + return normalizeRef(headRef) +} + +func classifyGitHubRemoteError(action string, err error) error { + if pe := plugin.AsError(err); pe != nil { + return pe + } + return classifyGitHubRemoteMessage(action, err.Error()) +} + +func classifyGitHubRemoteMessage(action, message string) error { + lower := strings.ToLower(message) + code := plugin.ErrorNetwork + if strings.Contains(lower, "rate limit") || strings.Contains(lower, "secondary rate") || strings.Contains(lower, "api rate limit exceeded") { + code = plugin.ErrorRemoteRateLimited + } else if strings.Contains(lower, "auth") || strings.Contains(lower, "authentication") || strings.Contains(lower, "credential") || strings.Contains(lower, "401") || strings.Contains(lower, "403") { + code = plugin.ErrorAuthRequired + } + return plugin.NewErrorf(code, "%s: %s", action, message) +} + func buildReviewArgs(prNumber int, payload plugin.ReviewPublishPayload) ([]string, error) { args := []string{"api", "-X", "POST", fmt.Sprintf("repos/{owner}/{repo}/pulls/%d/reviews", prNumber)} if payload.Context.Repository.HeadSHA != "" { @@ -236,8 +320,3 @@ func execGH(ctx context.Context, args ...string) (string, string, error) { stdout, stderr, err := gh.ExecContext(ctx, args...) return stdout.String(), stderr.String(), err } - -func isGitHubRemote(url string) bool { - url = strings.ToLower(url) - return strings.Contains(url, "github.com:") || strings.Contains(url, "github.com/") -} diff --git a/plugins/github/cmd/ero-plugin-github/main_test.go b/plugins/github/cmd/ero-plugin-github/main_test.go index eaf7b57..7f82079 100644 --- a/plugins/github/cmd/ero-plugin-github/main_test.go +++ b/plugins/github/cmd/ero-plugin-github/main_test.go @@ -9,14 +9,139 @@ import ( "ero/pkg/plugin" ) -func TestDetectContextRequiresGitHubRemote(t *testing.T) { - provider := githubProvider{} - result, err := provider.DetectContext(context.Background(), plugin.DetectContextRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{Name: "origin", URL: "git@github.com:owner/repo.git"}}}}}) +func TestGitHubRemoteParsing(t *testing.T) { + valid := []string{ + "git@github.com:owner/repo.git", + "git@github.com:owner/repo.git/", + "https://github.com/owner/repo.git", + "https://github.com/owner/repo", + "https://github.com/owner/repo/", + "ssh://git@github.com/owner/repo.git", + } + for _, raw := range valid { + remote, ok := parseGitHubRemote(raw) + if !ok || remote.Owner != "owner" || remote.Name != "repo" { + t.Fatalf("parseGitHubRemote(%q) = %#v, %v", raw, remote, ok) + } + } + + invalid := []string{ + "git@example.com:owner/repo.git", + "http://github.com/owner/repo", + "git://github.com/owner/repo", + "https://notgithub.com/owner/repo", + "https://github.com/owner", + "https://github.com/owner/repo/extra", + "github.com/owner/repo", + "ssh://github.com/owner/repo.git", + "ssh://user@github.com/owner/repo.git", + } + for _, raw := range invalid { + if remote, ok := parseGitHubRemote(raw); ok { + t.Fatalf("parseGitHubRemote(%q) unexpectedly matched %#v", raw, remote) + } + } +} + +func TestDetectContextRequiresMatchingGitHubPullRequest(t *testing.T) { + list := ghPRListResponse{} + list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 12, URL: "https://github.com/owner/repo/pull/12", BaseRefName: "main", HeadRefName: "feature"}} + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return &fakeGraphQLClient{listPages: []ghPRListResponse{list}}, nil }} + review := plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{Name: "origin", URL: "git@github.com:owner/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}} + result, err := provider.DetectContext(context.Background(), plugin.DetectContextRequest{Context: review}) if err != nil { t.Fatalf("DetectContext returned error: %v", err) } if !result.Result.Applicable { - t.Fatalf("expected GitHub remote to be applicable: %#v", result) + t.Fatalf("expected matching GitHub PR to be applicable: %#v", result) + } + + list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 13, BaseRefName: "main", HeadRefName: "other"}} + provider = githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return &fakeGraphQLClient{listPages: []ghPRListResponse{list}}, nil }} + result, err = provider.DetectContext(context.Background(), plugin.DetectContextRequest{Context: review}) + if err != nil { + t.Fatalf("DetectContext returned error: %v", err) + } + if result.Result.Applicable || !strings.Contains(result.Result.Reason, "no matching") { + t.Fatalf("expected no matching PR to be unavailable: %#v", result) + } + + result, err = provider.DetectContext(context.Background(), plugin.DetectContextRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{Name: "origin", URL: "https://example.com/github.com/owner/repo"}}}}}) + if err != nil { + t.Fatalf("DetectContext returned error: %v", err) + } + if result.Result.Applicable { + t.Fatalf("expected malformed/non-GitHub remote to be rejected: %#v", result) + } +} + +func TestGitHubPRMatching(t *testing.T) { + tests := []struct { + name string + ctx plugin.ReviewContext + prs []githubPRCandidate + wantNumber int + wantErr string + }{ + { + name: "branch mode matches current head branch and default base fallback", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}, + prs: []githubPRCandidate{{Number: 1, BaseRef: "main", HeadRef: "feature", HeadRepoOwner: "owner", HeadRepoName: "repo"}}, + wantNumber: 1, + }, + { + name: "range mode uses target base and head refs", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "range", BaseRef: "release", HeadRef: "topic"}}, + prs: []githubPRCandidate{{Number: 2, BaseRef: "release", HeadRef: "topic", HeadRepoOwner: "owner", HeadRepoName: "repo"}}, + wantNumber: 2, + }, + { + name: "fork PR matches branch even when head repository differs from base remote", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}, + prs: []githubPRCandidate{{Number: 3, BaseRef: "main", HeadRef: "feature", HeadRepoOwner: "forker", HeadRepoName: "repo"}}, + wantNumber: 3, + }, + { + name: "detached range-only SHA matches exact head SHA", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "range", HeadSHA: "abc123"}}, + prs: []githubPRCandidate{{Number: 4, BaseRef: "main", HeadSHA: "abc123"}}, + wantNumber: 4, + }, + { + name: "ambiguous multiple matches returns not applicable", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}, + prs: []githubPRCandidate{{Number: 5, BaseRef: "main", HeadRef: "feature"}, {Number: 6, BaseRef: "main", HeadRef: "feature"}}, + wantErr: plugin.ErrorNotApplicable, + }, + { + name: "default branch fallback does not override explicit base ref", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch", BaseRef: "release"}}, + prs: []githubPRCandidate{{Number: 7, BaseRef: "main", HeadRef: "feature"}}, + wantErr: plugin.ErrorNotApplicable, + }, + { + name: "no match returns not applicable", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}, + prs: []githubPRCandidate{{Number: 8, BaseRef: "main", HeadRef: "other"}}, + wantErr: plugin.ErrorNotApplicable, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := matchGitHubPR(tt.ctx, tt.prs) + if tt.wantErr != "" { + if plugin.AsError(err) == nil || plugin.AsError(err).Code != tt.wantErr { + t.Fatalf("expected error %q, got %v", tt.wantErr, err) + } + return + } + if err != nil { + t.Fatalf("matchGitHubPR returned error: %v", err) + } + if got.Number != tt.wantNumber { + t.Fatalf("expected PR %d, got %#v", tt.wantNumber, got) + } + }) } } @@ -55,6 +180,33 @@ func TestPublishReviewRejectsMalformedGitHubReviewResponse(t *testing.T) { } } +func TestPublishReviewUsesGraphQLMatchedPullRequestWhenContextHasRemote(t *testing.T) { + list := ghPRListResponse{} + list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 77, URL: "https://github.com/owner/repo/pull/77", BaseRefName: "release", HeadRefName: "topic"}} + var calls [][]string + provider := githubProvider{ + newGraphQLClient: func() (graphQLDoer, error) { return &fakeGraphQLClient{listPages: []ghPRListResponse{list}}, nil }, + execGH: func(_ context.Context, args ...string) (string, string, error) { + calls = append(calls, slices.Clone(args)) + return `{"id": 77, "html_url": "https://github.com/owner/repo/pull/77#pullrequestreview-77"}`, "", nil + }, + } + _, err := provider.PublishReview(context.Background(), plugin.PublishReviewParams{Payload: plugin.ReviewPublishPayload{ + Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "git@github.com:owner/repo.git"}}}, Target: plugin.ReviewTargetMetadata{Mode: "range", BaseRef: "release", HeadRef: "topic"}}, + Draft: plugin.ReviewDraftSnapshot{Summary: "summary"}, + }}) + if err != nil { + t.Fatalf("PublishReview returned error: %v", err) + } + if len(calls) != 1 { + t.Fatalf("expected only publish gh call, got %#v", calls) + } + joined := strings.Join(calls[0], "\x00") + if strings.Contains(joined, "pr\x00view") || !strings.Contains(joined, "pulls/77/reviews") { + t.Fatalf("publish did not use GraphQL-matched PR: %#v", calls[0]) + } +} + func TestPublishReviewSubmitsGitHubReview(t *testing.T) { var calls [][]string provider := githubProvider{execGH: func(_ context.Context, args ...string) (string, string, error) { diff --git a/plugins/github/cmd/ero-plugin-github/map.go b/plugins/github/cmd/ero-plugin-github/map.go new file mode 100644 index 0000000..1ac8c01 --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/map.go @@ -0,0 +1,112 @@ +package main + +import ( + "time" + + "ero/pkg/plugin" +) + +func candidateFromNode(n ghPRNode) githubPRCandidate { + owner := "" + if n.HeadRepositoryOwner != nil { + owner = n.HeadRepositoryOwner.Login + } + repo := "" + if n.HeadRepository != nil { + repo = n.HeadRepository.Name + } + return githubPRCandidate{Number: n.Number, URL: n.URL, Title: n.Title, State: n.State, BaseRef: n.BaseRefName, HeadRef: n.HeadRefName, HeadRepoOwner: owner, HeadRepoName: repo, HeadSHA: n.HeadRefOid} +} + +func mapGitHubPRPage(n ghPRNode) githubPRSnapshot { + s := mapGitHubPRMetadata(n) + s.IssueComments = mapGitHubIssueComments(n.Comments.Nodes) + s.Reviews = mapGitHubReviews(n.Reviews.Nodes) + for _, t := range n.ReviewThreads.Nodes { + s.Threads = append(s.Threads, mapGitHubThread(t)) + } + return s +} + +func mapGitHubPRMetadata(n ghPRNode) githubPRSnapshot { + s := githubPRSnapshot{Number: n.Number, URL: n.URL, Title: n.Title, State: n.State, Body: n.Body, BaseRef: n.BaseRefName, HeadRef: n.HeadRefName, HeadSHA: n.HeadRefOid, UpdatedAt: n.UpdatedAt} + if n.Author != nil { + s.Author = n.Author.Login + } + if n.HeadRepositoryOwner != nil { + s.HeadRepoOwner = n.HeadRepositoryOwner.Login + } + if n.HeadRepository != nil { + s.HeadRepoName = n.HeadRepository.Name + } + return s +} + +func mapGitHubIssueComments(comments []ghIssueComment) []plugin.ProviderIssueComment { + out := make([]plugin.ProviderIssueComment, 0, len(comments)) + for _, c := range comments { + author := "" + if c.Author != nil { + author = c.Author.Login + } + out = append(out, plugin.ProviderIssueComment{ExternalID: c.ID, Author: author, Body: c.Body, CreatedAt: c.CreatedAt, UpdatedAt: c.UpdatedAt, ExternalURL: c.URL}) + } + return out +} + +func mapGitHubReviews(reviews []ghReview) []plugin.ProviderReviewSummary { + out := make([]plugin.ProviderReviewSummary, 0, len(reviews)) + for _, r := range reviews { + author := "" + if r.Author != nil { + author = r.Author.Login + } + out = append(out, plugin.ProviderReviewSummary{ExternalID: r.ID, Author: author, State: r.State, Body: r.Body, SubmittedAt: r.SubmittedAt, ExternalURL: r.URL}) + } + return out +} + +func mapGitHubThread(t ghReviewThread) plugin.RemoteReviewThread { + thread := plugin.RemoteReviewThread{ProviderID: providerID, ExternalID: t.ID, FilePath: t.Path, ExternalURL: firstThreadURL(t), Unmapped: t.IsOutdated || t.Path == "" || t.Line <= 0} + thread.Range = plugin.ReviewLineRange{End: lineRef(t.Line, t.Side)} + if t.StartLine > 0 { + thread.Range.Start = lineRef(t.StartLine, firstNonEmpty(t.StartSide, t.Side)) + } else { + thread.Range.Start = thread.Range.End + } + if thread.Range.End.NewLineNumber == 0 && thread.Range.End.OldLineNumber == 0 { + thread.Unmapped = true + } + for _, c := range t.Comments.Nodes { + author := "" + if c.Author != nil { + author = c.Author.Login + } + thread.Comments = append(thread.Comments, plugin.RemoteReviewComment{ExternalID: c.ID, Author: author, Body: c.Body, CreatedAt: c.CreatedAt}) + } + return thread +} + +func lineRef(line int, side string) plugin.ReviewLineRef { + if line <= 0 { + return plugin.ReviewLineRef{} + } + if side == "LEFT" { + return plugin.ReviewLineRef{OldLineNumber: line, Kind: "LEFT"} + } + return plugin.ReviewLineRef{NewLineNumber: line, Kind: "RIGHT"} +} + +func firstThreadURL(t ghReviewThread) string { + if len(t.Comments.Nodes) == 0 { + return "" + } + return t.Comments.Nodes[0].URL +} + +func snapshotResultFromGitHub(s githubPRSnapshot) plugin.LoadRemoteSnapshotResult { + now := pluginNow() + return plugin.LoadRemoteSnapshotResult{RuntimeProviderID: providerID, Threads: s.Threads, FetchedAt: &now, Overview: &plugin.ProviderOverview{RuntimeProviderID: providerID, Title: s.Title, Number: s.Number, State: s.State, ExternalURL: s.URL, Author: s.Author, Body: s.Body, BaseRef: s.BaseRef, HeadRef: s.HeadRef, UpdatedAt: s.UpdatedAt, Comments: s.IssueComments, Reviews: s.Reviews}, Metadata: map[string]string{"provider": "github"}} +} + +var pluginNow = time.Now diff --git a/plugins/github/cmd/ero-plugin-github/match.go b/plugins/github/cmd/ero-plugin-github/match.go new file mode 100644 index 0000000..32f5f20 --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/match.go @@ -0,0 +1,79 @@ +package main + +import ( + "fmt" + "strings" + + "ero/pkg/plugin" +) + +type githubPRCandidate struct { + Number int + URL string + Title string + State string + BaseRef string + HeadRef string + HeadRepoOwner string + HeadRepoName string + HeadSHA string +} + +func matchGitHubPR(ctx plugin.ReviewContext, candidates []githubPRCandidate) (githubPRCandidate, error) { + matches := make([]githubPRCandidate, 0, 1) + for _, pr := range candidates { + if githubPRMatches(ctx, pr) { + matches = append(matches, pr) + } + } + if len(matches) == 0 { + return githubPRCandidate{}, plugin.NewError(plugin.ErrorNotApplicable, "no matching GitHub pull request found") + } + if len(matches) > 1 { + return githubPRCandidate{}, plugin.NewErrorf(plugin.ErrorNotApplicable, "ambiguous GitHub pull request match: %d candidates matched", len(matches)) + } + return matches[0], nil +} + +func githubPRMatches(ctx plugin.ReviewContext, pr githubPRCandidate) bool { + base := strings.TrimSpace(ctx.Target.BaseRef) + if base == "" { + base = strings.TrimSpace(ctx.Repository.DefaultBranch) + } + if base != "" && !refEqual(pr.BaseRef, base) { + return false + } + + headSHA := firstNonEmpty(ctx.Target.HeadSHA, ctx.Repository.HeadSHA) + headRef := strings.TrimSpace(ctx.Target.HeadRef) + if headRef == "" && strings.EqualFold(ctx.Target.Mode, "branch") { + headRef = strings.TrimSpace(ctx.Repository.CurrentBranch) + } + if headRef != "" { + if !refEqual(pr.HeadRef, headRef) { + return false + } + if headSHA != "" && pr.HeadSHA != "" && !strings.EqualFold(pr.HeadSHA, headSHA) { + return false + } + return true + } + + return headSHA != "" && pr.HeadSHA != "" && strings.EqualFold(pr.HeadSHA, headSHA) +} + +func refEqual(a, b string) bool { + return normalizeRef(a) == normalizeRef(b) +} + +func normalizeRef(ref string) string { + ref = strings.TrimSpace(ref) + for _, prefix := range []string{"refs/heads/", "origin/"} { + ref = strings.TrimPrefix(ref, prefix) + } + return strings.ToLower(ref) +} + +func githubPRSummary(pr githubPRCandidate) string { + return fmt.Sprintf("#%d %s", pr.Number, pr.URL) +} diff --git a/plugins/github/cmd/ero-plugin-github/remote.go b/plugins/github/cmd/ero-plugin-github/remote.go new file mode 100644 index 0000000..f258c25 --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/remote.go @@ -0,0 +1,65 @@ +package main + +import ( + "net/url" + "regexp" + "strings" + + "ero/pkg/plugin" +) + +type githubRemote struct { + Owner string + Name string +} + +var scpGitHubRemotePattern = regexp.MustCompile(`^git@github\.com:([^/]+)/(.+)$`) + +func parseGitHubRemote(raw string) (githubRemote, bool) { + raw = strings.TrimSpace(raw) + if raw == "" { + return githubRemote{}, false + } + if m := scpGitHubRemotePattern.FindStringSubmatch(raw); m != nil { + return cleanGitHubRepo(m[1], m[2]) + } + u, err := url.Parse(raw) + if err != nil || !strings.EqualFold(u.Hostname(), "github.com") { + return githubRemote{}, false + } + if u.Scheme != "https" && u.Scheme != "ssh" { + return githubRemote{}, false + } + parts := strings.Split(strings.Trim(u.Path, "/"), "/") + if len(parts) != 2 { + return githubRemote{}, false + } + if u.Scheme == "ssh" && (u.User == nil || u.User.Username() != "git") { + return githubRemote{}, false + } + return cleanGitHubRepo(parts[0], parts[1]) +} + +func cleanGitHubRepo(owner, repo string) (githubRemote, bool) { + owner = strings.TrimSpace(owner) + repo = strings.TrimSuffix(strings.TrimSpace(repo), "/") + repo = strings.TrimSuffix(repo, ".git") + if owner == "" || repo == "" || strings.Contains(owner, "/") || strings.Contains(repo, "/") || strings.Contains(owner, " ") || strings.Contains(repo, " ") { + return githubRemote{}, false + } + return githubRemote{Owner: owner, Name: repo}, true +} + +func isGitHubRemote(raw string) bool { + _, ok := parseGitHubRemote(raw) + return ok +} + +func firstGitHubRemote(remotes []plugin.GitRemote) (githubRemote, bool) { + for _, remote := range remotes { + if parsed, ok := parseGitHubRemote(remote.URL); ok { + return parsed, true + } + } + return githubRemote{}, false +} From e4c2e14f32879c3ad9cbb5d280fbe8a835c1a8df Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 16:34:23 +0200 Subject: [PATCH 06/22] chore(providers): finalize active provider sync --- README.md | 4 +- docs/architecture.md | 5 ++ docs/plugins.md | 17 +++-- internal/adapters/in/tui/pr_sheet.go | 5 -- internal/adapters/in/tui/provider_picker.go | 2 + internal/adapters/in/tui/review_publish.go | 3 +- .../adapters/out/providercache/cache_test.go | 28 ++++++--- internal/app/app.go | 3 +- internal/app/tui_active_provider.go | 9 +++ internal/core/provider_snapshot_test.go | 24 ++++--- .../github/cmd/ero-plugin-github/graphql.go | 63 ++++++++++++++++++- plugins/github/cmd/ero-plugin-github/main.go | 3 +- plugins/github/cmd/ero-plugin-github/map.go | 2 + .../github/cmd/ero-plugin-github/remote.go | 5 -- 14 files changed, 134 insertions(+), 39 deletions(-) diff --git a/README.md b/README.md index c741687..9aff151 100644 --- a/README.md +++ b/README.md @@ -13,6 +13,8 @@ A terminal UI for reviewing Git diffs file by file. - expandable unchanged context around diff hunks - syntax-aware diff rendering - selection copy support +- active review-provider sync with inline remote review threads +- GitHub PR overview sheet with Markdown-rendered PR body, issue comments, and review summaries ## Install @@ -44,7 +46,7 @@ ero --context-lines 5 ## Plugins -Ero supports a general local subprocess plugin system, managed with `ero plugin install`, `ero plugin list`, `ero plugin update`, and `ero plugin remove`. The first shipped contribution type is `review_provider`, used by the maintained GitHub and pi-coding-agent plugins. The GitHub plugin requires the GitHub CLI (`gh`) installed and authenticated with `gh auth login`. See [docs/plugins.md](docs/plugins.md) for authoring details. +Ero supports a general local subprocess plugin system, managed with `ero plugin install`, `ero plugin list`, `ero plugin update`, and `ero plugin remove`. The first shipped contribution type is `review_provider`, used by the maintained GitHub and pi-coding-agent plugins. Ero discovers all provider contributions but activates one review provider at a time, with provider switching, manual refresh, cache-first sync, and provider sync status in the TUI. The GitHub plugin uses GitHub CLI-compatible authentication through `go-gh` and requires `gh auth login`. See [docs/plugins.md](docs/plugins.md) for authoring details. ## Development diff --git a/docs/architecture.md b/docs/architecture.md index 256e523..608f166 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -72,6 +72,7 @@ It composes: - syntax adapter - core business/application services - startup mode detection and mixed-change prompting +- active review-provider selection, cache-first sync, polling/backoff, and provider preference storage - CLI adapter - TUI adapter factories @@ -121,6 +122,8 @@ Initial outbound ports: - `StartupStateReader` for smart default launch - `FileContentReader` - `SyntaxTokenizer` +- plugin contribution catalogs and selected review-provider client factories +- normalized provider snapshot cache and active-provider preference storage - `ReviewCallbackPublisher` (reserved for later integration) For v1, the concrete git adapter should be built on top of `github.com/go-git/go-git/v5`. @@ -132,6 +135,8 @@ If a core service needs external data or side effects, it depends on a port. The Bubble Tea app should be an orchestrator, not a god object. +Review-provider subprocess lifecycle and sync policy belong to `internal/app`, not the TUI. The TUI consumes active provider state, sends switch/refresh/publish intents, renders sync status and provider overview data, and keeps provider picker rows descriptor/cache-based so inactive providers are not started just for display. + ### Thin app shell `internal/adapters/in/tui/` should mainly: diff --git a/docs/plugins.md b/docs/plugins.md index f93c93b..b182470 100644 --- a/docs/plugins.md +++ b/docs/plugins.md @@ -2,7 +2,7 @@ Ero plugins are a general extension mechanism based on local subprocesses. A plugin declares one or more contributions in its manifest; future contribution types can extend other parts of Ero, such as themes or additional workflows. -This first release ships the `review_provider` contribution type, which lets plugins publish reviews and load remote review comments. +The `review_provider` contribution type lets plugins publish reviews, load remote review threads, and provide provider-specific review context such as pull request metadata. ## Install and manage plugins @@ -50,6 +50,10 @@ Required fields are `name`, `version`, `manifest_version = "1"`, `protocol = "er Ero discovers available review providers from installed plugin manifests before starting plugin subprocesses. Each discovered provider has a host-owned stable key derived from a canonical installed-plugin identity plus the contribution `id`; runtime provider IDs returned by `initialize` remain provider-owned metadata and are not used as the host selection key. +Ero keeps the plugin system global while activating only one review provider at a time. The TUI can switch providers, manually refresh the active provider, show cache/sync status, and display provider overview data in the PR sheet. Inactive providers remain descriptors plus cached/previously observed status; Ero does not start inactive provider subprocesses just to populate the picker. + +Provider snapshots are normalized Ero data stored under the XDG cache directory. Ero loads cached provider data first, refreshes in the background, and keeps good cached data when refresh fails. + `runtime.command` is executed with the plugin root as the working directory. Keep it stable for installed users; use the optional `build.command` for local development or release packaging. ## Protocol @@ -80,10 +84,11 @@ Review provider methods: - `initialize`: negotiate `ero.plugin.v1`, bind to the requested `contribution_id`, and return provider metadata/capabilities. - `detect_context`: decide whether the current repository/review context applies. -- `load_remote_threads`: return remote review comments when `load_remote_comments` is supported. +- `load_remote_threads`: return remote review threads when `load_remote_comments` is supported. +- `load_remote_snapshot`: return remote review threads plus provider overview data when `load_remote_snapshot` is supported. Hosts prefer this method when advertised and fall back to `load_remote_threads` for older providers. - `publish_review`: publish a draft review when `publish_review` is supported. -Capabilities include `load_remote_comments`, `publish_review`, supported `decisions` (`comment`, `request_changes`, `approve`), and `idempotent_publish`. +Capabilities include `load_remote_comments`, `load_remote_snapshot`, `publish_review`, supported `decisions` (`comment`, `request_changes`, `approve`), and `idempotent_publish`. ## Go SDK @@ -119,7 +124,7 @@ Do not put secrets in `ero-plugin.toml`, command-line arguments, or stdout. Read Ero ships maintained plugin implementations under `plugins/`: -- `plugins/github`: GitHub review provider. It requires the GitHub CLI (`gh`) installed and authenticated with `gh auth login`; the plugin uses `go-gh`/`gh` for GitHub auth, current-branch PR lookup, and PR review submission. Publishing returns a fast error when the current branch has no associated pull request. +- `plugins/github`: GitHub review provider. It uses GitHub CLI-compatible authentication through `go-gh`, so `gh auth login` must be configured. The provider parses GitHub remotes, detects the matching pull request for the current branch/range context, fetches PR metadata, issue comments, review summaries, and review threads through GraphQL, and publishes reviews to the matched pull request. Publishing returns a fast error when no matching pull request is available. - `plugins/pi-coding-agent`: pi-coding-agent destination. Load its Pi extension, then Ero can publish a review into the matching Pi session as a user message. Build them with: @@ -139,6 +144,6 @@ For a one-off development session, `pi -e ./plugins/pi-coding-agent` also works, The bridge records active sessions in an owner-only runtime registry and uses per-session Unix sockets. Ero selects a session by `PI_CODING_AGENT_SESSION_ID` when set, otherwise by repository path plus branch/SHA when available. -## First-release limitations +## Current limitations -The first plugin release focuses on review providers launched as local subprocesses. Ero does not provide a sandbox, plugin marketplace, background daemon, automatic secret storage, or full forge implementations. Remote APIs, authentication flows, and provider-specific publish semantics belong in individual plugins. +Ero review providers run as local subprocesses. Ero does not provide a sandbox, plugin marketplace, background daemon, automatic secret storage, or full forge implementations. Remote APIs, authentication flows, and provider-specific publish semantics belong in individual plugins. diff --git a/internal/adapters/in/tui/pr_sheet.go b/internal/adapters/in/tui/pr_sheet.go index b89197a..b70eb5a 100644 --- a/internal/adapters/in/tui/pr_sheet.go +++ b/internal/adapters/in/tui/pr_sheet.go @@ -5,7 +5,6 @@ import ( "strings" "time" - tea "charm.land/bubbletea/v2" "charm.land/lipgloss/v2" "ero/internal/core" @@ -251,7 +250,3 @@ func padRight(s string, width int) string { } return s + strings.Repeat(" ", width-lipgloss.Width(s)) } - -func togglePRSheetCmd() tea.Cmd { - return func() tea.Msg { return prSheetToggledMsg{} } -} diff --git a/internal/adapters/in/tui/provider_picker.go b/internal/adapters/in/tui/provider_picker.go index 042d3d4..7ca1a40 100644 --- a/internal/adapters/in/tui/provider_picker.go +++ b/internal/adapters/in/tui/provider_picker.go @@ -129,7 +129,9 @@ func (m Model) renderProviderPickerOverlay(content string) string { } func (m Model) renderProviderPicker(width, height int) string { + // Keep an eight-column outer margin, a readable 36-column minimum, and a 76-column maximum. paneWidth := min(max(width-8, 36), 76) + // Account for pane padding/chrome while preserving at least one content column. contentWidth := max(paneWidth-6, 1) lines := []string{theme.HelpPaneTitleStyle.Render("Review providers"), ""} rows := m.providerPicker.rows diff --git a/internal/adapters/in/tui/review_publish.go b/internal/adapters/in/tui/review_publish.go index fce8020..b286d8e 100644 --- a/internal/adapters/in/tui/review_publish.go +++ b/internal/adapters/in/tui/review_publish.go @@ -256,7 +256,8 @@ func (m Model) providerClientsFor(infos []core.ReviewProviderInfo) []providerCli result := make([]providerClientWithInfo, 0, len(infos)) if m.activeProvider != nil && m.activeRuntimeInfo.ID != "" { if info, ok := selected[m.activeRuntimeInfo.ID]; ok { - return append(result, providerClientWithInfo{info: info, client: m.activeProvider}) + result = append(result, providerClientWithInfo{info: info, client: m.activeProvider}) + delete(selected, m.activeRuntimeInfo.ID) } } for _, client := range m.reviewProviders { diff --git a/internal/adapters/out/providercache/cache_test.go b/internal/adapters/out/providercache/cache_test.go index fbf61d6..b17b3b8 100644 --- a/internal/adapters/out/providercache/cache_test.go +++ b/internal/adapters/out/providercache/cache_test.go @@ -13,10 +13,16 @@ func TestCacheRoundTripNormalizedSnapshot(t *testing.T) { key := core.ReviewContextKey{StableProviderKey: "plugin:github#review_provider:github", RepositoryIdentity: "remotes:github.com/o/r", TargetMode: core.DiffModeBranch, BaseRef: "main", HeadRef: "feature", BaseSHA: "b", HeadSHA: "h"} snapshot := core.ProviderSnapshot{StableProviderKey: key.StableProviderKey, RuntimeProviderID: "github", ContextKey: key, Threads: []core.RemoteReviewThread{{ProviderID: "github", ExternalID: "thread-1"}}, FetchedAt: time.Unix(10, 0).UTC()} - if err := store.SaveProviderSnapshot(context.Background(), snapshot); err != nil { t.Fatal(err) } + if err := store.SaveProviderSnapshot(context.Background(), snapshot); err != nil { + t.Fatal(err) + } got, ok, err := store.LoadProviderSnapshot(context.Background(), key) - if err != nil { t.Fatal(err) } - if !ok { t.Fatal("expected cached snapshot") } + if err != nil { + t.Fatal(err) + } + if !ok { + t.Fatal("expected cached snapshot") + } if got.StableProviderKey != snapshot.StableProviderKey || got.RuntimeProviderID != "github" || len(got.Threads) != 1 { t.Fatalf("unexpected snapshot: %#v", got) } @@ -25,9 +31,17 @@ func TestCacheRoundTripNormalizedSnapshot(t *testing.T) { func TestPreferenceRoundTrip(t *testing.T) { store := NewStore(t.TempDir(), t.TempDir()) repoID := "remotes:github.com/o/r" - if _, ok, err := store.LoadActiveProviderKey(context.Background(), repoID); err != nil || ok { t.Fatalf("empty load = ok %v err %v", ok, err) } - if err := store.SaveActiveProviderKey(context.Background(), repoID, "provider-key"); err != nil { t.Fatal(err) } + if _, ok, err := store.LoadActiveProviderKey(context.Background(), repoID); err != nil || ok { + t.Fatalf("empty load = ok %v err %v", ok, err) + } + if err := store.SaveActiveProviderKey(context.Background(), repoID, "provider-key"); err != nil { + t.Fatal(err) + } got, ok, err := store.LoadActiveProviderKey(context.Background(), repoID) - if err != nil { t.Fatal(err) } - if !ok || got != "provider-key" { t.Fatalf("got %q ok %v", got, ok) } + if err != nil { + t.Fatal(err) + } + if !ok || got != "provider-key" { + t.Fatalf("got %q ok %v", got, ok) + } } diff --git a/internal/app/app.go b/internal/app/app.go index 7fb3d5a..ea14f02 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -114,7 +114,8 @@ func newAppWithClipboard(cfg *viper.Viper, loader reviewLoader, runner tuiRunner metadata = reader } reviewContext := buildReviewContext(initialRequest, files, metadata, version) - err = runner.Run(tui.NewModelWithActiveProviderContext(ctx, files, terminal.NewCapabilities(), loader, initialRequest, clipboardWriter, reviewContext, activeProvider, nil)) + var compatibilityProviders []ports.ReviewProviderClient + err = runner.Run(tui.NewModelWithActiveProviderContext(ctx, files, terminal.NewCapabilities(), loader, initialRequest, clipboardWriter, reviewContext, activeProvider, compatibilityProviders)) if err != nil { log.Error().Err(err).Msg("tui exited with error") return err diff --git a/internal/app/tui_active_provider.go b/internal/app/tui_active_provider.go index e3ba56a..0eda535 100644 --- a/internal/app/tui_active_provider.go +++ b/internal/app/tui_active_provider.go @@ -21,11 +21,17 @@ func (c *tuiActiveProviderController) Catalog(ctx context.Context) ([]ports.Revi } func (c *tuiActiveProviderController) Start(ctx context.Context, review core.ReviewContext) (tui.ActiveProviderState, error) { + if c == nil || c.service == nil { + return tui.ActiveProviderState{}, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) + } state, err := c.service.Start(ctx, review) return c.toTUIState(state), err } func (c *tuiActiveProviderController) Refresh(ctx context.Context, review core.ReviewContext, manual bool) (tui.ActiveProviderState, error) { + if c == nil || c.service == nil { + return tui.ActiveProviderState{}, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) + } state, err := c.service.Refresh(ctx, review, manual) return c.toTUIState(state), err } @@ -53,6 +59,9 @@ func (c *tuiActiveProviderController) CompleteTimer(ctx context.Context, review } func (c *tuiActiveProviderController) Switch(ctx context.Context, review core.ReviewContext, stableKey string) (tui.ActiveProviderState, error) { + if c == nil || c.service == nil { + return tui.ActiveProviderState{}, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) + } state, err := c.service.Switch(ctx, review, stableKey) return c.toTUIState(state), err } diff --git a/internal/core/provider_snapshot_test.go b/internal/core/provider_snapshot_test.go index 90c478d..f5d36ba 100644 --- a/internal/core/provider_snapshot_test.go +++ b/internal/core/provider_snapshot_test.go @@ -23,20 +23,24 @@ func TestReviewContextKeyIncludesIdentityInputs(t *testing.T) { baseKey := NewReviewContextKey("provider-a", base) cases := map[string]func(*ReviewContext){ - "provider": nil, - "remote": func(c *ReviewContext) { c.Repository.Remotes[0].URL = "git@github.com:owner/other.git" }, - "mode": func(c *ReviewContext) { c.Target.Mode = DiffModeRange }, - "base ref": func(c *ReviewContext) { c.Target.BaseRef = "main" }, - "head ref": func(c *ReviewContext) { c.Target.HeadRef = "feature-2" }, - "base sha": func(c *ReviewContext) { c.Target.BaseSHA = "base2" }, - "head sha": func(c *ReviewContext) { c.Target.HeadSHA = "head2" }, + "provider": nil, + "remote": func(c *ReviewContext) { c.Repository.Remotes[0].URL = "git@github.com:owner/other.git" }, + "mode": func(c *ReviewContext) { c.Target.Mode = DiffModeRange }, + "base ref": func(c *ReviewContext) { c.Target.BaseRef = "main" }, + "head ref": func(c *ReviewContext) { c.Target.HeadRef = "feature-2" }, + "base sha": func(c *ReviewContext) { c.Target.BaseSHA = "base2" }, + "head sha": func(c *ReviewContext) { c.Target.HeadSHA = "head2" }, "merge base": func(c *ReviewContext) { c.Target.MergeBaseSHA = "merge2" }, } for name, mutate := range cases { ctx := base provider := "provider-a" - if mutate == nil { provider = "provider-b" } else { mutate(&ctx) } + if mutate == nil { + provider = "provider-b" + } else { + mutate(&ctx) + } if got := NewReviewContextKey(provider, ctx); got == baseKey { t.Fatalf("%s did not change key", name) } @@ -63,7 +67,7 @@ func TestRepositoryIdentityPrefersRemotesAndFallsBackToPath(t *testing.T) { func sampleProviderReviewContext() ReviewContext { return ReviewContext{ Repository: RepositoryMetadata{RepoPath: "/repo", WorktreeRoot: "/repo", Remotes: []GitRemote{{Name: "origin", URL: "git@github.com:owner/repo.git"}}}, - Target: ReviewTargetMetadata{Mode: DiffModeBranch, BaseRef: "origin/main", HeadRef: "feature", BaseSHA: "base1", HeadSHA: "head1", MergeBaseSHA: "merge1"}, - Session: ReviewSessionMetadata{LocalReviewID: "local", IdempotencyKey: "idem", CreatedAt: time.Unix(1, 0)}, + Target: ReviewTargetMetadata{Mode: DiffModeBranch, BaseRef: "origin/main", HeadRef: "feature", BaseSHA: "base1", HeadSHA: "head1", MergeBaseSHA: "merge1"}, + Session: ReviewSessionMetadata{LocalReviewID: "local", IdempotencyKey: "idem", CreatedAt: time.Unix(1, 0)}, } } diff --git a/plugins/github/cmd/ero-plugin-github/graphql.go b/plugins/github/cmd/ero-plugin-github/graphql.go index 3c6e12b..2a037e8 100644 --- a/plugins/github/cmd/ero-plugin-github/graphql.go +++ b/plugins/github/cmd/ero-plugin-github/graphql.go @@ -125,8 +125,67 @@ type ghReviewThreadComment struct { CreatedAt time.Time `json:"createdAt"` } -const githubPRListQuery = `query EroPRList($owner:String!, $name:String!, $after:String) { repository(owner:$owner, name:$name) { pullRequests(first:50, after:$after, states:[OPEN]) { nodes { number url title state baseRefName headRefName headRefOid headRepositoryOwner { login } headRepository { name } } pageInfo { hasNextPage endCursor } } } } }` -const githubPRSnapshotQuery = `query EroPRSnapshot($owner:String!, $name:String!, $number:Int!, $commentsAfter:String, $reviewsAfter:String, $threadsAfter:String) { repository(owner:$owner, name:$name) { pullRequest(number:$number) { number url title state body updatedAt author { login } baseRefName headRefName headRefOid headRepositoryOwner { login } headRepository { name } comments(first:100, after:$commentsAfter) { nodes { id url body createdAt updatedAt author { login } } pageInfo { hasNextPage endCursor } } reviews(first:100, after:$reviewsAfter) { nodes { id url state body submittedAt author { login } } pageInfo { hasNextPage endCursor } } reviewThreads(first:100, after:$threadsAfter) { nodes { id path line side startLine startSide isOutdated comments(first:100) { nodes { id url body createdAt author { login } } pageInfo { hasNextPage endCursor } } } pageInfo { hasNextPage endCursor } } } } }` +const githubPRListQuery = `query EroPRList($owner:String!, $name:String!, $after:String) { + repository(owner:$owner, name:$name) { + pullRequests(first:50, after:$after, states:[OPEN]) { + nodes { + number + url + title + state + baseRefName + headRefName + headRefOid + headRepositoryOwner { login } + headRepository { name } + } + pageInfo { hasNextPage endCursor } + } + } +}` + +const githubPRSnapshotQuery = `query EroPRSnapshot($owner:String!, $name:String!, $number:Int!, $commentsAfter:String, $reviewsAfter:String, $threadsAfter:String) { + repository(owner:$owner, name:$name) { + pullRequest(number:$number) { + number + url + title + state + body + updatedAt + author { login } + baseRefName + headRefName + headRefOid + headRepositoryOwner { login } + headRepository { name } + comments(first:100, after:$commentsAfter) { + nodes { id url body createdAt updatedAt author { login } } + pageInfo { hasNextPage endCursor } + } + reviews(first:100, after:$reviewsAfter) { + nodes { id url state body submittedAt author { login } } + pageInfo { hasNextPage endCursor } + } + reviewThreads(first:100, after:$threadsAfter) { + nodes { + id + path + line + side + startLine + startSide + isOutdated + comments(first:100) { + nodes { id url body createdAt author { login } } + pageInfo { hasNextPage endCursor } + } + } + pageInfo { hasNextPage endCursor } + } + } + } +}` func (p githubProvider) graphQLClient() (graphQLDoer, error) { if p.newGraphQLClient != nil { diff --git a/plugins/github/cmd/ero-plugin-github/main.go b/plugins/github/cmd/ero-plugin-github/main.go index c25d1a4..55ea6e6 100644 --- a/plugins/github/cmd/ero-plugin-github/main.go +++ b/plugins/github/cmd/ero-plugin-github/main.go @@ -105,7 +105,8 @@ func (p githubProvider) LoadRemoteSnapshot(ctx context.Context, req plugin.LoadR } func (p githubProvider) LoadRemoteThreads(ctx context.Context, req plugin.LoadRemoteThreadsRequest) (plugin.LoadRemoteThreadsResult, error) { - snapshot, err := p.LoadRemoteSnapshot(ctx, plugin.LoadRemoteSnapshotRequest{Context: req.Context}) + snapshotReq := plugin.LoadRemoteSnapshotRequest(req) + snapshot, err := p.LoadRemoteSnapshot(ctx, snapshotReq) if err != nil { return plugin.LoadRemoteThreadsResult{}, err } diff --git a/plugins/github/cmd/ero-plugin-github/map.go b/plugins/github/cmd/ero-plugin-github/map.go index 1ac8c01..4f21361 100644 --- a/plugins/github/cmd/ero-plugin-github/map.go +++ b/plugins/github/cmd/ero-plugin-github/map.go @@ -67,6 +67,7 @@ func mapGitHubReviews(reviews []ghReview) []plugin.ProviderReviewSummary { } func mapGitHubThread(t ghReviewThread) plugin.RemoteReviewThread { + // Mark clearly stale or unanchored threads as unmapped before line conversion. thread := plugin.RemoteReviewThread{ProviderID: providerID, ExternalID: t.ID, FilePath: t.Path, ExternalURL: firstThreadURL(t), Unmapped: t.IsOutdated || t.Path == "" || t.Line <= 0} thread.Range = plugin.ReviewLineRange{End: lineRef(t.Line, t.Side)} if t.StartLine > 0 { @@ -74,6 +75,7 @@ func mapGitHubThread(t ghReviewThread) plugin.RemoteReviewThread { } else { thread.Range.Start = thread.Range.End } + // Also mark unmapped when GitHub supplied a line but no supported side mapping resolved. if thread.Range.End.NewLineNumber == 0 && thread.Range.End.OldLineNumber == 0 { thread.Unmapped = true } diff --git a/plugins/github/cmd/ero-plugin-github/remote.go b/plugins/github/cmd/ero-plugin-github/remote.go index f258c25..cea3fff 100644 --- a/plugins/github/cmd/ero-plugin-github/remote.go +++ b/plugins/github/cmd/ero-plugin-github/remote.go @@ -50,11 +50,6 @@ func cleanGitHubRepo(owner, repo string) (githubRemote, bool) { return githubRemote{Owner: owner, Name: repo}, true } -func isGitHubRemote(raw string) bool { - _, ok := parseGitHubRemote(raw) - return ok -} - func firstGitHubRemote(remotes []plugin.GitRemote) (githubRemote, bool) { for _, remote := range remotes { if parsed, ok := parseGitHubRemote(remote.URL); ok { From 32e1a1c5d58cadb78b0f69d77530fa17ea330c6b Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 20:32:55 +0200 Subject: [PATCH 07/22] chore(providers): trace active provider sync --- internal/app/active_provider_service.go | 46 ++++++++++++++++++++++--- 1 file changed, 41 insertions(+), 5 deletions(-) diff --git a/internal/app/active_provider_service.go b/internal/app/active_provider_service.go index fe47241..e7051e7 100644 --- a/internal/app/active_provider_service.go +++ b/internal/app/active_provider_service.go @@ -6,6 +6,8 @@ import ( "sync" "time" + "github.com/bnema/zerowrap" + "ero/internal/core" "ero/internal/ports" ) @@ -61,13 +63,18 @@ func NewActiveProviderService(catalog ports.ReviewProviderCatalog, factory ports } func (s *ActiveProviderService) Start(ctx context.Context, review core.ReviewContext) (ActiveProviderState, error) { + log := zerowrap.FromCtx(ctx) descs, err := s.catalog.ListReviewProviderDescriptors(ctx) if err != nil { + log.Warn().Err(err).Msg("active provider catalog load failed") return ActiveProviderState{}, err } ordered := s.orderCandidates(ctx, descs, review) + log.Info().Int("descriptor_count", len(descs)).Int("candidate_count", len(ordered)).Str("repo", core.RepositoryIdentity(review.Repository)).Msg("active provider start") s.mu.Lock() - s.closeLocked() + if err := s.closeLocked(); err != nil { + log.Warn().Err(err).Msg("active provider close before start failed") + } s.stableKey = "" s.runtimeID = "" s.generation++ @@ -76,13 +83,17 @@ func (s *ActiveProviderService) Start(ctx context.Context, review core.ReviewCon s.mu.Unlock() var lastErr error for _, d := range ordered { + log.Debug().Str("provider_key", d.Key).Str("contribution_id", d.ContributionID).Str("plugin", d.PluginName).Msg("active provider probe candidate") client, info, err := s.probe(ctx, d, review) if err != nil { + log.Debug().Err(err).Str("provider_key", d.Key).Str("contribution_id", d.ContributionID).Msg("active provider probe failed") lastErr = err continue } s.mu.Lock() - s.closeLocked() + if err := s.closeLocked(); err != nil { + log.Warn().Err(err).Msg("active provider close before activate failed") + } s.client = client s.stableKey = d.Key s.runtimeID = info.ID @@ -94,23 +105,30 @@ func (s *ActiveProviderService) Start(ctx context.Context, review core.ReviewCon _ = s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key) } st := s.loadCachedState(ctx, review, d.Key, info.ID, info) + log.Info().Str("provider_key", d.Key).Str("runtime_provider_id", info.ID).Bool("from_cache", st.FromCache).Int("remote_thread_count", len(st.Snapshot.Threads)).Msg("active provider selected") s.setState(gen, st) return st, nil } + log.Warn().Err(lastErr).Msg("active provider start found no applicable provider") failed := failedProviderState(lastErr) s.setState(startGen, failed) return failed, lastErr } func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewContext, stableKey string) (ActiveProviderState, error) { + log := zerowrap.FromCtx(ctx) descs, err := s.catalog.ListReviewProviderDescriptors(ctx) if err != nil { + log.Warn().Err(err).Str("provider_key", stableKey).Msg("active provider switch catalog load failed") return ActiveProviderState{}, err } for _, d := range descs { if d.Key == stableKey { + log.Info().Str("provider_key", stableKey).Str("contribution_id", d.ContributionID).Msg("active provider switch start") s.mu.Lock() - s.closeLocked() + if err := s.closeLocked(); err != nil { + log.Warn().Err(err).Str("provider_key", stableKey).Msg("active provider close before switch failed") + } s.stableKey = "" s.runtimeID = "" s.generation++ @@ -119,12 +137,15 @@ func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewCo s.mu.Unlock() client, info, err := s.probe(ctx, d, review) if err != nil { + log.Warn().Err(err).Str("provider_key", stableKey).Msg("active provider switch probe failed") failed := failedProviderState(err) s.setState(switchGen, failed) return failed, err } s.mu.Lock() - s.closeLocked() + if err := s.closeLocked(); err != nil { + log.Warn().Err(err).Str("provider_key", stableKey).Msg("active provider close before switched activate failed") + } s.client = client s.stableKey = d.Key s.runtimeID = info.ID @@ -136,10 +157,12 @@ func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewCo _ = s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key) } st := s.loadCachedState(ctx, review, d.Key, info.ID, info) + log.Info().Str("provider_key", d.Key).Str("runtime_provider_id", info.ID).Bool("from_cache", st.FromCache).Int("remote_thread_count", len(st.Snapshot.Threads)).Msg("active provider switch complete") s.setState(gen, st) return st, nil } } + log.Warn().Str("provider_key", stableKey).Msg("active provider switch descriptor not found") return ActiveProviderState{}, core.NewProviderError(core.ProviderErrorNotApplicable, "provider descriptor not found", nil) } @@ -158,6 +181,7 @@ func (s *ActiveProviderService) PublishReview(ctx context.Context, request core. } func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewContext, manual bool) (ActiveProviderState, error) { + log := zerowrap.FromCtx(ctx) s.mu.Lock() client := s.client key := s.stableKey @@ -170,13 +194,17 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC } s.mu.Unlock() if client == nil { + log.Warn().Bool("manual", manual).Msg("active provider refresh requested with no provider") return prev, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider", nil) } + log.Info().Str("provider_key", key).Str("runtime_provider_id", runtimeID).Bool("manual", manual).Msg("active provider refresh start") var snap core.ProviderSnapshot var err error if loader, ok := client.(remoteSnapshotLoader); ok { + log.Debug().Str("provider_key", key).Msg("active provider loading remote snapshot") snap, err = loader.LoadRemoteSnapshot(ctx, review) } else { + log.Debug().Str("provider_key", key).Msg("active provider loading remote threads") var threads []core.RemoteReviewThread threads, err = client.LoadRemoteThreads(ctx, review) snap.Threads = threads @@ -190,9 +218,11 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC if st.NextSyncAt.IsZero() { st.Snapshot.Sync.Status = core.ProviderSyncStatusFailed st.Snapshot.Sync.NextSyncAt = nil + log.Warn().Err(err).Str("provider_key", key).Str("runtime_provider_id", runtimeID).Bool("manual", manual).Msg("active provider refresh failed") } else { st.Snapshot.Sync.Status = core.ProviderSyncStatusBackingOff st.Snapshot.Sync.NextSyncAt = new(st.NextSyncAt) + log.Warn().Err(err).Str("provider_key", key).Str("runtime_provider_id", runtimeID).Bool("manual", manual).Time("next_sync_at", st.NextSyncAt).Msg("active provider refresh backing off") } s.setState(gen, st) return st, err @@ -211,6 +241,7 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC if s.cache != nil { _ = s.cache.SaveProviderSnapshot(ctx, snap) } + log.Info().Str("provider_key", key).Str("runtime_provider_id", snap.RuntimeProviderID).Bool("manual", manual).Int("remote_thread_count", len(snap.Threads)).Bool("has_overview", snap.Overview != nil).Time("next_sync_at", next).Msg("active provider refresh synced") st := ActiveProviderState{StableProviderKey: key, RuntimeProviderID: runtimeID, RuntimeInfo: prev.RuntimeInfo, Snapshot: snap, NextSyncAt: next} s.setState(gen, st) return st, nil @@ -312,18 +343,23 @@ func (s *ActiveProviderService) probe(ctx context.Context, d ports.ReviewProvide return client, info, nil } func (s *ActiveProviderService) loadCachedState(ctx context.Context, review core.ReviewContext, key, runtimeID string, info ...core.ReviewProviderInfo) ActiveProviderState { + log := zerowrap.FromCtx(ctx) st := ActiveProviderState{StableProviderKey: key, RuntimeProviderID: runtimeID} if len(info) > 0 { st.RuntimeInfo = info[0] } if s.cache != nil { - if snap, ok, _ := s.cache.LoadProviderSnapshot(ctx, core.NewReviewContextKey(key, review)); ok { + contextKey := core.NewReviewContextKey(key, review) + if snap, ok, _ := s.cache.LoadProviderSnapshot(ctx, contextKey); ok { snap.Cached = true st.Snapshot = snap st.FromCache = true if snap.Sync.NextSyncAt != nil { st.NextSyncAt = *snap.Sync.NextSyncAt } + log.Debug().Str("provider_key", key).Any("context_key", contextKey).Int("remote_thread_count", len(snap.Threads)).Bool("has_overview", snap.Overview != nil).Msg("active provider cache hit") + } else { + log.Debug().Str("provider_key", key).Any("context_key", contextKey).Msg("active provider cache miss") } } return st From cf17c6cb59a054f4c0843e75ef4a9556e3b14eb5 Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 20:39:58 +0200 Subject: [PATCH 08/22] fix(plugins): rebuild stale local provider runtimes --- .../adapters/out/plugin/provider_loader.go | 78 +++++++++++++++++-- .../out/plugin/provider_loader_test.go | 50 ++++++++++++ 2 files changed, 123 insertions(+), 5 deletions(-) diff --git a/internal/adapters/out/plugin/provider_loader.go b/internal/adapters/out/plugin/provider_loader.go index c5e7ab0..e0df002 100644 --- a/internal/adapters/out/plugin/provider_loader.go +++ b/internal/adapters/out/plugin/provider_loader.go @@ -92,7 +92,7 @@ func (l *ReviewProviderLoader) createReviewProviderClient(ctx context.Context, d if command == "" { return nil, fmt.Errorf("plugin runtime command is empty") } - if !runtimeCommandAvailable(command, descriptor.PluginPath) && strings.TrimSpace(manifest.Build.Command) != "" { + if shouldBuildRuntime(command, descriptor.PluginPath, manifest.Build.Command) { if err := runPluginBuildCommand(ctx, descriptor.PluginPath, manifest.Build.Command, l.timeout); err != nil { log := zerowrap.FromCtx(ctx) log.Warn().Err(err).Str("plugin_path", descriptor.PluginPath).Msg("build plugin runtime failed") @@ -131,19 +131,87 @@ func canonicalInstalledPluginIdentity(descriptor ports.PluginDescriptor) string } func runtimeCommandAvailable(command, pluginDir string) bool { + _, ok := runtimeCommandInfo(command, pluginDir) + return ok +} + +func runtimeCommandInfo(command, pluginDir string) (os.FileInfo, bool) { if command == "" { - return false + return nil, false } if !strings.Contains(command, "/") { _, err := exec.LookPath(command) - return err == nil + return nil, err == nil } + path := runtimeCommandPath(command, pluginDir) + info, err := os.Stat(path) + return info, err == nil && !info.IsDir() +} + +func runtimeCommandPath(command, pluginDir string) string { path := command if !filepath.IsAbs(path) { path = filepath.Join(pluginDir, path) } - info, err := os.Stat(path) - return err == nil && !info.IsDir() + return filepath.Clean(path) +} + +func shouldBuildRuntime(command, pluginDir, buildCommand string) bool { + buildCommand = strings.TrimSpace(buildCommand) + if buildCommand == "" { + return false + } + runtimeInfo, available := runtimeCommandInfo(command, pluginDir) + if !available { + return true + } + if !strings.Contains(command, "/") { + return false + } + return pluginSourceNewerThanRuntime(pluginDir, runtimeCommandPath(command, pluginDir), runtimeInfo.ModTime(), buildCommand) +} + +func pluginSourceNewerThanRuntime(pluginDir, runtimePath string, runtimeModTime time.Time, buildCommand string) bool { + runtimePath = filepath.Clean(runtimePath) + newer := false + _ = filepath.WalkDir(pluginDir, func(path string, d os.DirEntry, err error) error { + if err != nil || newer { + return nil + } + if d.IsDir() { + if d.Name() == ".git" { + return filepath.SkipDir + } + return nil + } + path = filepath.Clean(path) + if path == runtimePath || !isPluginSourcePath(path) { + return nil + } + info, err := d.Info() + if err == nil && info.ModTime().After(runtimeModTime) { + newer = true + } + return nil + }) + if newer { + return true + } + buildCommandName, _ := splitRuntimeCommand(buildCommand) + if buildCommandName == "" || !strings.Contains(buildCommandName, "/") { + return false + } + buildPath := runtimeCommandPath(buildCommandName, pluginDir) + info, err := os.Stat(buildPath) + return err == nil && !info.IsDir() && info.ModTime().After(runtimeModTime) +} + +func isPluginSourcePath(path string) bool { + switch filepath.Base(path) { + case "ero-plugin.toml", "go.mod", "go.sum": + return true + } + return filepath.Ext(path) == ".go" } func runPluginBuildCommand(ctx context.Context, pluginDir, buildCommand string, timeout time.Duration) error { diff --git a/internal/adapters/out/plugin/provider_loader_test.go b/internal/adapters/out/plugin/provider_loader_test.go index 96084b5..45d0393 100644 --- a/internal/adapters/out/plugin/provider_loader_test.go +++ b/internal/adapters/out/plugin/provider_loader_test.go @@ -5,6 +5,7 @@ import ( "os" "path/filepath" "testing" + "time" "github.com/stretchr/testify/require" @@ -124,6 +125,55 @@ label = "pi-coding-agent" require.NoError(t, providers[0].Close()) } +func TestReviewProviderLoaderRebuildsStaleLocalRuntimeBeforeStartingProvider(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + runtimePath := filepath.Join(dir, "runtime-plugin") + buildMarker := filepath.Join(dir, "build-ran") + buildScript := filepath.Join(dir, "build-runtime.sh") + err := os.WriteFile(buildScript, []byte("#!/bin/sh\necho rebuilt > ./build-ran\ncat > ./runtime-plugin <<'EOF'\n#!/bin/sh\ncat\nEOF\nchmod +x ./runtime-plugin\n"), 0o755) + require.NoError(t, err) + err = os.WriteFile(runtimePath, []byte("#!/bin/sh\necho stale-runtime\n"), 0o755) + require.NoError(t, err) + oldTime := time.Now().Add(-2 * time.Hour) + require.NoError(t, os.Chtimes(runtimePath, oldTime, oldTime)) + err = os.WriteFile(filepath.Join(dir, "cmd-source.go"), []byte("package main\n"), 0o644) + require.NoError(t, err) + err = os.WriteFile(filepath.Join(dir, "ero-plugin.toml"), []byte(`name = "buildable" +version = "0.1.0" +manifest_version = "1" +protocol = "ero.plugin.v1" + +[runtime] +command = "./runtime-plugin" + +[build] +command = "`+buildScript+`" + +[[contributions]] +type = "review_provider" +id = "github" +label = "GitHub" +`), 0o644) + require.NoError(t, err) + + registry := mocks.NewMockPluginRegistry(t) + registry.EXPECT().InstalledPlugins(context.Background()).Return([]ports.PluginDescriptor{{ + Name: "buildable", + Path: dir, + Contributions: []ports.PluginContribution{ + {Type: "review_provider", ID: "github", Label: "GitHub"}, + }, + }}, nil) + + providers, err := NewReviewProviderLoader(registry).LoadReviewProviders(context.Background()) + require.NoError(t, err) + require.Len(t, providers, 1) + require.FileExists(t, buildMarker) + require.NoError(t, providers[0].Close()) +} + func TestReviewProviderLoaderStartsOneClientPerReviewProviderContribution(t *testing.T) { t.Parallel() From b1c441fa6b5f72b6fe5be9b9041d886fa39bfd66 Mon Sep 17 00:00:00 2001 From: brice Date: Fri, 5 Jun 2026 20:53:45 +0200 Subject: [PATCH 09/22] fix(github): use review thread diff side fields --- .../github/cmd/ero-plugin-github/graphql.go | 8 ++-- .../cmd/ero-plugin-github/graphql_test.go | 37 +++++++++++++++++++ 2 files changed, 41 insertions(+), 4 deletions(-) diff --git a/plugins/github/cmd/ero-plugin-github/graphql.go b/plugins/github/cmd/ero-plugin-github/graphql.go index 2a037e8..b329db0 100644 --- a/plugins/github/cmd/ero-plugin-github/graphql.go +++ b/plugins/github/cmd/ero-plugin-github/graphql.go @@ -108,9 +108,9 @@ type ghReviewThread struct { ID string `json:"id"` Path string `json:"path"` Line int `json:"line"` - Side string `json:"side"` + Side string `json:"diffSide"` StartLine int `json:"startLine"` - StartSide string `json:"startSide"` + StartSide string `json:"startDiffSide"` IsOutdated bool `json:"isOutdated"` Comments struct { Nodes []ghReviewThreadComment `json:"nodes"` @@ -172,9 +172,9 @@ const githubPRSnapshotQuery = `query EroPRSnapshot($owner:String!, $name:String! id path line - side + diffSide startLine - startSide + startDiffSide isOutdated comments(first:100) { nodes { id url body createdAt author { login } } diff --git a/plugins/github/cmd/ero-plugin-github/graphql_test.go b/plugins/github/cmd/ero-plugin-github/graphql_test.go index 2900683..5fc1351 100644 --- a/plugins/github/cmd/ero-plugin-github/graphql_test.go +++ b/plugins/github/cmd/ero-plugin-github/graphql_test.go @@ -3,6 +3,9 @@ package main import ( "context" "maps" + "os" + "strconv" + "strings" "testing" "time" @@ -35,6 +38,15 @@ func (f *fakeGraphQLClient) DoWithContext(_ context.Context, query string, varia return nil } +func TestGitHubPRSnapshotQueryUsesReviewThreadSchemaFieldNames(t *testing.T) { + if strings.Contains(githubPRSnapshotQuery, "\n side\n") || strings.Contains(githubPRSnapshotQuery, "\n startSide\n") { + t.Fatalf("review thread query must use GitHub schema fields diffSide/startDiffSide, not side/startSide:\n%s", githubPRSnapshotQuery) + } + if !strings.Contains(githubPRSnapshotQuery, "\n diffSide\n") || !strings.Contains(githubPRSnapshotQuery, "\n startDiffSide\n") { + t.Fatalf("review thread query is missing diffSide/startDiffSide:\n%s", githubPRSnapshotQuery) + } +} + func TestGitHubGraphQLMapping(t *testing.T) { now := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) actor := &ghActor{Login: "octo"} @@ -61,6 +73,31 @@ func TestGitHubGraphQLMapping(t *testing.T) { } } +func TestLoadRemoteSnapshotAgainstGitHubWhenEnabled(t *testing.T) { + prNumber := strings.TrimSpace(os.Getenv("ERO_GITHUB_INTEGRATION_PR")) + if prNumber == "" { + t.Skip("set ERO_GITHUB_INTEGRATION_PR to run GitHub integration snapshot fetch") + } + number, err := strconv.Atoi(prNumber) + if err != nil { + t.Fatalf("invalid ERO_GITHUB_INTEGRATION_PR: %v", err) + } + client, err := defaultGraphQLClient() + if err != nil { + t.Fatalf("defaultGraphQLClient: %v", err) + } + got, err := fetchGitHubPRSnapshot(context.Background(), client, githubRemote{Owner: "bnema", Name: "ero"}, number) + if err != nil { + t.Fatalf("fetchGitHubPRSnapshot(%d): %v", number, err) + } + if got.Number != number || got.Title == "" { + t.Fatalf("unexpected snapshot metadata: %#v", got) + } + if len(got.Threads) == 0 { + t.Fatalf("expected at least one review thread on PR %d", number) + } +} + func TestLoadRemoteSnapshotPaginatesAndMaps(t *testing.T) { now := time.Now().UTC() actor := &ghActor{Login: "octo"} From e58d9b132fb6dd57981a9b72eb7b2a3f49020eaa Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 06:02:06 +0200 Subject: [PATCH 10/22] feat(tui): support mouse wheel scrolling --- .../adapters/in/tui/cursor_navigation_test.go | 22 +++++++++++++++++++ internal/adapters/in/tui/model.go | 18 +++++++++++++++ internal/adapters/in/tui/model_test.go | 6 +++-- 3 files changed, 44 insertions(+), 2 deletions(-) diff --git a/internal/adapters/in/tui/cursor_navigation_test.go b/internal/adapters/in/tui/cursor_navigation_test.go index cb3192b..4b9a41a 100644 --- a/internal/adapters/in/tui/cursor_navigation_test.go +++ b/internal/adapters/in/tui/cursor_navigation_test.go @@ -64,6 +64,28 @@ func TestModelCursorNavigationKeepsCursorVisibleWithoutRebuildingDocument(t *tes } } +func TestModelMouseWheelScrollsReviewDocument(t *testing.T) { + t.Parallel() + + model := NewModel([]core.ReviewFile{reviewFileWithLines("demo.go", 80)}) + updated, _ := model.Update(tea.WindowSizeMsg{Width: 80, Height: 10}) + model = updated.(Model) + + for range 10 { + updated, _ = model.Update(tea.MouseWheelMsg(tea.Mouse{Button: tea.MouseWheelDown})) + model = updated.(Model) + } + + assert.Equal(t, 12, model.cursorRow) + assert.Equal(t, 4, model.reviewViewport.YOffset()) + + updated, _ = model.Update(tea.MouseWheelMsg(tea.Mouse{Button: tea.MouseWheelUp})) + model = updated.(Model) + + assert.Equal(t, 11, model.cursorRow) + assert.Equal(t, 4, model.reviewViewport.YOffset()) +} + func TestModelCursorNavigationPreservesAbsoluteAndPageViewportSemantics(t *testing.T) { t.Parallel() diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index 7ae2cc4..7ee854d 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -320,6 +320,8 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.height = max(msg.Height, 0) m.syncReviewViewport() return m, nil + case tea.MouseWheelMsg: + return m.updateMouseWheel(msg) case tea.KeyPressMsg: if m.helpActive { switch msg.String() { @@ -350,6 +352,21 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { } } +func (m Model) updateMouseWheel(msg tea.MouseWheelMsg) (tea.Model, tea.Cmd) { + if m.helpActive || m.commentEditor != nil || m.search.active() || m.providerPicker.open || m.publish.active { + return m, nil + } + switch msg.Mouse().Button { + case tea.MouseWheelUp: + m.moveCursor(-1) + case tea.MouseWheelDown: + m.moveCursor(1) + default: + return m, nil + } + return m, nil +} + func (m Model) updateReviewAction(action keymap.Action) (tea.Model, tea.Cmd) { switch action { case keymap.ActionQuit: @@ -454,6 +471,7 @@ func (m Model) View() tea.View { } view := tea.NewView(content) view.AltScreen = true + view.MouseMode = tea.MouseModeCellMotion if m.commentEditor != nil { view.KeyboardEnhancements.ReportAllKeysAsEscapeCodes = true view.KeyboardEnhancements.ReportAssociatedText = true diff --git a/internal/adapters/in/tui/model_test.go b/internal/adapters/in/tui/model_test.go index a914c30..6229fe7 100644 --- a/internal/adapters/in/tui/model_test.go +++ b/internal/adapters/in/tui/model_test.go @@ -13,7 +13,7 @@ import ( "ero/internal/core" ) -func TestModelViewUsesAlternateScreen(t *testing.T) { +func TestModelViewUsesAlternateScreenAndMouseWheelEvents(t *testing.T) { t.Parallel() tests := []struct { @@ -27,7 +27,9 @@ func TestModelViewUsesAlternateScreen(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() model := NewModel(tt.files) - assert.True(t, model.View().AltScreen) + view := model.View() + assert.True(t, view.AltScreen) + assert.Equal(t, tea.MouseModeCellMotion, view.MouseMode) }) } } From 830eb60fc48c462d02cca8fe4925749ea8159698 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 06:15:20 +0200 Subject: [PATCH 11/22] feat(tui): show draft comment count --- .../adapters/in/tui/active_provider_test.go | 125 +++++++++--------- .../adapters/in/tui/component/statusbar.go | 13 +- .../in/tui/component/statusbar_test.go | 20 +++ internal/adapters/in/tui/model.go | 14 ++ internal/adapters/in/tui/model_test.go | 15 +++ .../adapters/in/tui/provider_picker_test.go | 47 +++---- 6 files changed, 151 insertions(+), 83 deletions(-) diff --git a/internal/adapters/in/tui/active_provider_test.go b/internal/adapters/in/tui/active_provider_test.go index 6e3b43b..569e783 100644 --- a/internal/adapters/in/tui/active_provider_test.go +++ b/internal/adapters/in/tui/active_provider_test.go @@ -5,64 +5,51 @@ import ( "errors" "testing" + "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" "ero/internal/core" "ero/internal/ports" ) -type fakeActiveProviderController struct { - catalog []ports.ReviewProviderDescriptor - startState ActiveProviderState - refreshState ActiveProviderState - switchStates map[string]ActiveProviderState - publishResult core.PublishReviewResult - startErr error - refreshErr error - switchErrs map[string]error - publishErr error - startCalls int - refreshManual []bool - switchKeys []string - publishRequests []core.PublishReviewRequest - closed bool -} +type mockActiveProviderController struct{ mock.Mock } -func (f *fakeActiveProviderController) Catalog(context.Context) ([]ports.ReviewProviderDescriptor, error) { - return append([]ports.ReviewProviderDescriptor(nil), f.catalog...), nil +func (m *mockActiveProviderController) Catalog(ctx context.Context) ([]ports.ReviewProviderDescriptor, error) { + args := m.Called(ctx) + return args.Get(0).([]ports.ReviewProviderDescriptor), args.Error(1) } -func (f *fakeActiveProviderController) Start(context.Context, core.ReviewContext) (ActiveProviderState, error) { - f.startCalls++ - return f.startState, f.startErr +func (m *mockActiveProviderController) Start(ctx context.Context, review core.ReviewContext) (ActiveProviderState, error) { + args := m.Called(ctx, review) + return args.Get(0).(ActiveProviderState), args.Error(1) } -func (f *fakeActiveProviderController) Refresh(_ context.Context, _ core.ReviewContext, manual bool) (ActiveProviderState, error) { - f.refreshManual = append(f.refreshManual, manual) - return f.refreshState, f.refreshErr +func (m *mockActiveProviderController) Refresh(ctx context.Context, review core.ReviewContext, manual bool) (ActiveProviderState, error) { + args := m.Called(ctx, review, manual) + return args.Get(0).(ActiveProviderState), args.Error(1) } -func (f *fakeActiveProviderController) Switch(_ context.Context, _ core.ReviewContext, key string) (ActiveProviderState, error) { - f.switchKeys = append(f.switchKeys, key) - if err := f.switchErrs[key]; err != nil { - return ActiveProviderState{}, err - } - return f.switchStates[key], nil +func (m *mockActiveProviderController) Switch(ctx context.Context, review core.ReviewContext, stableKey string) (ActiveProviderState, error) { + args := m.Called(ctx, review, stableKey) + return args.Get(0).(ActiveProviderState), args.Error(1) } -func (f *fakeActiveProviderController) PublishReview(_ context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) { - f.publishRequests = append(f.publishRequests, request) - return f.publishResult, f.publishErr +func (m *mockActiveProviderController) PublishReview(ctx context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) { + args := m.Called(ctx, request) + return args.Get(0).(core.PublishReviewResult), args.Error(1) } -func (f *fakeActiveProviderController) Generation() int64 { - return int64(len(f.refreshManual) + len(f.switchKeys) + f.startCalls) +func (m *mockActiveProviderController) Generation() int64 { + args := m.Called() + return args.Get(0).(int64) } -func (f *fakeActiveProviderController) CompleteTimer(ctx context.Context, review core.ReviewContext, _ int64) (ActiveProviderState, error) { - return f.Refresh(ctx, review, false) +func (m *mockActiveProviderController) CompleteTimer(ctx context.Context, review core.ReviewContext, generation int64) (ActiveProviderState, error) { + args := m.Called(ctx, review, generation) + return args.Get(0).(ActiveProviderState), args.Error(1) } -func (f *fakeActiveProviderController) Close() error { f.closed = true; return nil } +func (m *mockActiveProviderController) Close() error { return m.Called().Error(0) } func TestActiveProviderStartupLoadsOnlyActiveProviderState(t *testing.T) { - controller := &fakeActiveProviderController{ - catalog: []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "other", Label: "Other"}}, - startState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github-runtime", RuntimeInfo: core.ReviewProviderInfo{ID: "github-runtime", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ProviderID: "github-runtime", ExternalID: "t1"}}, Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusSynced}}}, - } + controller := &mockActiveProviderController{} + catalog := []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "other", Label: "Other"}} + startState := ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github-runtime", RuntimeInfo: core.ReviewProviderInfo{ID: "github-runtime", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ProviderID: "github-runtime", ExternalID: "t1"}}, Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusSynced}}} + controller.On("Catalog", mock.Anything).Return(catalog, nil).Once() + controller.On("Start", mock.Anything, mock.Anything).Return(startState, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), []core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, []ports.ReviewProviderClient{fakeProviderAsPort{&fakeReviewProvider{}}}) cmd := m.Init() @@ -71,21 +58,23 @@ func TestActiveProviderStartupLoadsOnlyActiveProviderState(t *testing.T) { m = updated.(Model) require.NotNil(t, refreshCmd) - require.Equal(t, 1, controller.startCalls) + controller.AssertNumberOfCalls(t, "Start", 1) require.Len(t, m.providerCatalog, 2) require.Equal(t, "github", m.activeProviderKey) require.Equal(t, "github-runtime", m.activeRuntimeID) require.Equal(t, core.ProviderSyncStatusSynced, m.providerSyncState.Status) require.Len(t, m.remoteThreads, 1) require.Empty(t, m.providerInfoByClient) + controller.AssertExpectations(t) } func TestProviderStartupRefreshesRemoteThreadsAfterCacheState(t *testing.T) { - controller := &fakeActiveProviderController{ - catalog: []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}}, - startState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "cached"}}}}, - refreshState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "fresh"}}}}, - } + controller := &mockActiveProviderController{} + startState := ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "cached"}}}} + refreshState := ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "fresh"}}}} + controller.On("Catalog", mock.Anything).Return([]ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}}, nil).Once() + controller.On("Start", mock.Anything, mock.Anything).Return(startState, nil).Once() + controller.On("Refresh", mock.Anything, mock.Anything, false).Return(refreshState, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) started, refreshCmd := m.Update(m.Init()()) @@ -96,13 +85,16 @@ func TestProviderStartupRefreshesRemoteThreadsAfterCacheState(t *testing.T) { refreshed, _ := m.Update(refreshCmd()) m = refreshed.(Model) - require.Equal(t, []bool{false}, controller.refreshManual) + controller.AssertCalled(t, "Refresh", mock.Anything, mock.Anything, false) require.Len(t, m.remoteThreads, 1) require.Equal(t, "fresh", m.remoteThreads[0].ExternalID) + controller.AssertExpectations(t) } func TestProviderRefreshManualReplacesRemoteThreads(t *testing.T) { - controller := &fakeActiveProviderController{refreshState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "new"}}}}} + controller := &mockActiveProviderController{} + refreshState := ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "new"}}}} + controller.On("Refresh", mock.Anything, mock.Anything, true).Return(refreshState, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) m.remoteThreads = []core.RemoteReviewThread{{ExternalID: "old"}} @@ -110,13 +102,16 @@ func TestProviderRefreshManualReplacesRemoteThreads(t *testing.T) { updated, _ := m.Update(msg) m = updated.(Model) - require.Equal(t, []bool{true}, controller.refreshManual) + controller.AssertCalled(t, "Refresh", mock.Anything, mock.Anything, true) require.Len(t, m.remoteThreads, 1) require.Equal(t, "new", m.remoteThreads[0].ExternalID) + controller.AssertExpectations(t) } func TestProviderSwitchReplacesRemoteData(t *testing.T) { - controller := &fakeActiveProviderController{switchStates: map[string]ActiveProviderState{"other": {StableProviderKey: "other", RuntimeProviderID: "other-runtime", RuntimeInfo: core.ReviewProviderInfo{ID: "other-runtime"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "other-thread"}}}}}, switchErrs: map[string]error{}} + controller := &mockActiveProviderController{} + switchState := ActiveProviderState{StableProviderKey: "other", RuntimeProviderID: "other-runtime", RuntimeInfo: core.ReviewProviderInfo{ID: "other-runtime"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "other-thread"}}}} + controller.On("Switch", mock.Anything, mock.Anything, "other").Return(switchState, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) m.remoteThreads = []core.RemoteReviewThread{{ExternalID: "old"}} @@ -124,27 +119,36 @@ func TestProviderSwitchReplacesRemoteData(t *testing.T) { updated, _ := m.Update(msg) m = updated.(Model) - require.Equal(t, []string{"other"}, controller.switchKeys) + controller.AssertCalled(t, "Switch", mock.Anything, mock.Anything, "other") require.Equal(t, "other", m.activeProviderKey) require.Len(t, m.remoteThreads, 1) require.Equal(t, "other-thread", m.remoteThreads[0].ExternalID) + controller.AssertExpectations(t) } func TestActiveProviderPollTimerRefreshesWithGeneration(t *testing.T) { - controller := &fakeActiveProviderController{refreshState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "polled"}}}}} + controller := &mockActiveProviderController{} + refreshed := ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ExternalID: "polled"}}}} + controller.On("CompleteTimer", mock.Anything, mock.Anything, int64(42)).Return(refreshed, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) msg := m.completeActiveProviderTimerCmd(42)() updated, _ := m.Update(msg) m = updated.(Model) - require.Equal(t, []bool{false}, controller.refreshManual) + controller.AssertCalled(t, "CompleteTimer", mock.Anything, mock.Anything, int64(42)) require.Len(t, m.remoteThreads, 1) require.Equal(t, "polled", m.remoteThreads[0].ExternalID) + controller.AssertExpectations(t) } func TestActiveProviderPublishUsesActiveProviderClient(t *testing.T) { - controller := &fakeActiveProviderController{publishResult: core.PublishReviewResult{ProviderID: "github", ExternalReviewID: "review-1"}} + controller := &mockActiveProviderController{} + publishResult := core.PublishReviewResult{ProviderID: "github", ExternalReviewID: "review-1"} + var publishedRequest core.PublishReviewRequest + controller.On("PublishReview", mock.Anything, mock.Anything).Run(func(args mock.Arguments) { + publishedRequest = args.Get(1).(core.PublishReviewRequest) + }).Return(publishResult, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) m.activeRuntimeInfo = core.ReviewProviderInfo{ID: "github", Label: "GitHub", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}} m.activeRuntimeID = "github" @@ -157,13 +161,15 @@ func TestActiveProviderPublishUsesActiveProviderClient(t *testing.T) { msg := cmd().(publishReviewCompletedMsg) m, _ = m.handlePublishReviewCompleted(msg) - require.Len(t, controller.publishRequests, 1) - require.Equal(t, "github", controller.publishRequests[0].ProviderID) + controller.AssertNumberOfCalls(t, "PublishReview", 1) + require.Equal(t, "github", publishedRequest.ProviderID) require.False(t, m.publish.active) + controller.AssertExpectations(t) } func TestProviderSwitchFailureClearsRemoteThreads(t *testing.T) { - controller := &fakeActiveProviderController{switchStates: map[string]ActiveProviderState{}, switchErrs: map[string]error{"other": errors.New("auth")}} + controller := &mockActiveProviderController{} + controller.On("Switch", mock.Anything, mock.Anything, "other").Return(ActiveProviderState{}, errors.New("auth")).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) m.activeProviderKey = "github" m.activeRuntimeID = "github" @@ -178,4 +184,5 @@ func TestProviderSwitchFailureClearsRemoteThreads(t *testing.T) { require.Empty(t, m.activeRuntimeID) require.Empty(t, m.providerInfos) require.Empty(t, m.remoteThreads) + controller.AssertExpectations(t) } diff --git a/internal/adapters/in/tui/component/statusbar.go b/internal/adapters/in/tui/component/statusbar.go index aa903d8..7b5aaf7 100644 --- a/internal/adapters/in/tui/component/statusbar.go +++ b/internal/adapters/in/tui/component/statusbar.go @@ -23,6 +23,7 @@ type StatusModel struct { ActiveProviderLabel string ActiveRuntimeName string ProviderSync core.ProviderSyncState + DraftCommentCount int ShowNoProvider bool NerdFont bool } @@ -51,6 +52,9 @@ func (c StatusBar) Render(model StatusModel) string { if syncLabel := providerSyncLabel(model); syncLabel != "" { segments = append(segments, statusSegment{style: theme.StatusInfoStyle, label: syncLabel}) } + if model.DraftCommentCount > 0 { + segments = append(segments, statusSegment{style: theme.StatusInfoStyle, label: draftCommentCountLabel(model.DraftCommentCount)}) + } prefix := renderStatusSegments(leftWidth, segments...) percent := renderStatusSegments(leftWidth-lipgloss.Width(prefix), statusSegment{style: theme.StatusInfoStyle, label: fmt.Sprintf("%3.0f%%", model.ScrollPercent*100)}) @@ -151,7 +155,7 @@ func providerSyncLabel(model StatusModel) string { status := providerSyncStatusLabel(model.ProviderSync.Status) if status != "" { if model.NerdFont { - status = providerSyncStatusSymbol(model.ProviderSync.Status) + " " + status + status = providerSyncStatusSymbol(model.ProviderSync.Status) } parts = append(parts, status) } @@ -167,6 +171,13 @@ func providerSyncLabel(model StatusModel) string { return strings.Join(parts, " ") } +func draftCommentCountLabel(count int) string { + if count == 1 { + return "1 draft comment" + } + return fmt.Sprintf("%d draft comments", count) +} + func providerSyncStatusLabel(status core.ProviderSyncStatus) string { switch status { case core.ProviderSyncStatusLoadingCache: diff --git a/internal/adapters/in/tui/component/statusbar_test.go b/internal/adapters/in/tui/component/statusbar_test.go index b714c2c..cfaa9f5 100644 --- a/internal/adapters/in/tui/component/statusbar_test.go +++ b/internal/adapters/in/tui/component/statusbar_test.go @@ -84,6 +84,26 @@ func TestStatusbarProviderSyncUsesNerdFontSymbolWhenSupported(t *testing.T) { } } +func TestStatusbarProviderSyncOmitsStatusWordWhenNerdFontSymbolIsShown(t *testing.T) { + model := syncStatusModel("GitHub", "github", core.ProviderSyncState{Status: core.ProviderSyncStatusSynced}) + model.NerdFont = true + + view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) + + require.Contains(t, view, "GitHub/github") + require.Contains(t, view, "") + require.NotContains(t, view, "synced") +} + +func TestStatusbarShowsDraftCommentCount(t *testing.T) { + model := baseStatusModel() + model.DraftCommentCount = 2 + + view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) + + require.Contains(t, view, "2 draft comments") +} + func TestStatusbarProviderSyncNarrowWidthDegradesGracefully(t *testing.T) { last := time.Date(2026, 6, 5, 12, 34, 0, 0, time.UTC) next := last.Add(5 * time.Minute) diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index 7ee854d..b28c93d 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -430,6 +430,19 @@ func (m Model) updateReviewAction(action keymap.Action) (tea.Model, tea.Cmd) { return m, nil } +func (m Model) unpublishedDraftCommentCount() int { + if m.reviewDraft == nil { + return 0 + } + count := 0 + for _, comment := range m.reviewDraft.Comments() { + if comment.State != core.ReviewCommentStatePublished { + count++ + } + } + return count +} + func (m Model) View() tea.View { review := m.reviewViewport.View(m.reviewVisualState()) if m.loading { @@ -450,6 +463,7 @@ func (m Model) View() tea.View { ActiveProviderLabel: m.activeRuntimeInfo.Label, ActiveRuntimeName: m.activeRuntimeID, ProviderSync: m.providerSyncState, + DraftCommentCount: m.unpublishedDraftCommentCount(), ShowNoProvider: m.activeProvider != nil && m.activeProviderKey == "" && m.activeRuntimeInfo.ID == "", NerdFont: m.nerdFont, }), diff --git a/internal/adapters/in/tui/model_test.go b/internal/adapters/in/tui/model_test.go index 6229fe7..ad4b0e2 100644 --- a/internal/adapters/in/tui/model_test.go +++ b/internal/adapters/in/tui/model_test.go @@ -34,6 +34,21 @@ func TestModelViewUsesAlternateScreenAndMouseWheelEvents(t *testing.T) { } } +func TestModelViewShowsUnpublishedDraftCommentCount(t *testing.T) { + t.Parallel() + + model := NewModel([]core.ReviewFile{reviewFile("demo.go", "package main")}) + _, err := model.reviewDraft.AddComment(core.ReviewCommentInput{FilePath: "demo.go", Range: core.ReviewLineRange{Start: core.ReviewLineRef{NewLineNumber: 1}, End: core.ReviewLineRef{NewLineNumber: 1}}, Body: "first"}) + require.NoError(t, err) + _, err = model.reviewDraft.AddComment(core.ReviewCommentInput{FilePath: "demo.go", Range: core.ReviewLineRange{Start: core.ReviewLineRef{NewLineNumber: 1}, End: core.ReviewLineRef{NewLineNumber: 1}}, Body: "second"}) + require.NoError(t, err) + model.reviewDraft.ApplyPublishedRefs("github", []core.PublishedReviewCommentRef{{LocalCommentID: "comment-1", ExternalID: "remote-1"}}) + + view := stripANSI(model.View().Content) + + require.Contains(t, view, "1 draft comment") +} + func TestModelViewRendersSequentialReviewDocumentWithoutFileExplorer(t *testing.T) { t.Parallel() diff --git a/internal/adapters/in/tui/provider_picker_test.go b/internal/adapters/in/tui/provider_picker_test.go index 71af212..0637863 100644 --- a/internal/adapters/in/tui/provider_picker_test.go +++ b/internal/adapters/in/tui/provider_picker_test.go @@ -5,6 +5,7 @@ import ( "testing" tea "charm.land/bubbletea/v2" + "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" "ero/internal/core" @@ -12,13 +13,14 @@ import ( ) func TestProviderPickerDisplaysDescriptorRowsWithoutStartingInactiveProviders(t *testing.T) { - controller := &fakeActiveProviderController{ - catalog: []ports.ReviewProviderDescriptor{ - {Key: "github", Label: "GitHub", PluginName: "gh-plugin", PluginSource: "builtin"}, - {Key: "gitlab", Label: "GitLab", PluginName: "gl-plugin", PluginSource: "local"}, - }, - startState: ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusFailed, LastError: "missing token"}}}, + controller := &mockActiveProviderController{} + catalog := []ports.ReviewProviderDescriptor{ + {Key: "github", Label: "GitHub", PluginName: "gh-plugin", PluginSource: "builtin"}, + {Key: "gitlab", Label: "GitLab", PluginName: "gl-plugin", PluginSource: "local"}, } + startState := ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github", RuntimeInfo: core.ReviewProviderInfo{ID: "github", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusFailed, LastError: "missing token"}}} + controller.On("Catalog", mock.Anything).Return(catalog, nil).Once() + controller.On("Start", mock.Anything, mock.Anything).Return(startState, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) updated, _ := m.Update(m.Init()()) m = updated.(Model) @@ -27,7 +29,7 @@ func TestProviderPickerDisplaysDescriptorRowsWithoutStartingInactiveProviders(t m = updated.(Model) require.True(t, m.providerPicker.open) - require.Equal(t, 1, controller.startCalls) + controller.AssertNumberOfCalls(t, "Start", 1) view := stripANSI(m.View().Content) require.Contains(t, view, "Review providers") require.Contains(t, view, "* GitHub") @@ -35,15 +37,14 @@ func TestProviderPickerDisplaysDescriptorRowsWithoutStartingInactiveProviders(t require.Contains(t, view, "missing token") require.Contains(t, view, "GitLab") require.Contains(t, view, "gl-plugin local") + controller.AssertExpectations(t) } func TestProviderPickerSelectEmitsSwitchCommandWithStableKey(t *testing.T) { - controller := &fakeActiveProviderController{ - catalog: []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "gitlab", Label: "GitLab"}}, - startState: ActiveProviderState{StableProviderKey: "github"}, - switchStates: map[string]ActiveProviderState{"gitlab": {StableProviderKey: "gitlab"}}, - switchErrs: map[string]error{}, - } + controller := &mockActiveProviderController{} + controller.On("Catalog", mock.Anything).Return([]ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "gitlab", Label: "GitLab"}}, nil).Once() + controller.On("Start", mock.Anything, mock.Anything).Return(ActiveProviderState{StableProviderKey: "github"}, nil).Once() + controller.On("Switch", mock.Anything, mock.Anything, "gitlab").Return(ActiveProviderState{StableProviderKey: "gitlab"}, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) updated, _ := m.Update(m.Init()()) m = updated.(Model) @@ -59,18 +60,17 @@ func TestProviderPickerSelectEmitsSwitchCommandWithStableKey(t *testing.T) { m = updated.(Model) require.False(t, m.providerPicker.open) - require.Equal(t, []string{"gitlab"}, controller.switchKeys) + controller.AssertCalled(t, "Switch", mock.Anything, mock.Anything, "gitlab") require.Equal(t, "gitlab", m.activeProviderKey) + controller.AssertExpectations(t) } func TestProviderCycleAndRefreshShortcuts(t *testing.T) { - controller := &fakeActiveProviderController{ - catalog: []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "gitlab", Label: "GitLab"}}, - startState: ActiveProviderState{StableProviderKey: "github"}, - refreshState: ActiveProviderState{StableProviderKey: "github"}, - switchStates: map[string]ActiveProviderState{"gitlab": {StableProviderKey: "gitlab"}}, - switchErrs: map[string]error{}, - } + controller := &mockActiveProviderController{} + controller.On("Catalog", mock.Anything).Return([]ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "gitlab", Label: "GitLab"}}, nil).Once() + controller.On("Start", mock.Anything, mock.Anything).Return(ActiveProviderState{StableProviderKey: "github"}, nil).Once() + controller.On("Switch", mock.Anything, mock.Anything, "gitlab").Return(ActiveProviderState{StableProviderKey: "gitlab"}, nil).Once() + controller.On("Refresh", mock.Anything, mock.Anything, true).Return(ActiveProviderState{StableProviderKey: "github"}, nil).Once() m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) updated, _ := m.Update(m.Init()()) m = updated.(Model) @@ -80,10 +80,11 @@ func TestProviderCycleAndRefreshShortcuts(t *testing.T) { require.NotNil(t, cmd) updated, _ = m.Update(cmd()) m = updated.(Model) - require.Equal(t, []string{"gitlab"}, controller.switchKeys) + controller.AssertCalled(t, "Switch", mock.Anything, mock.Anything, "gitlab") _, cmd = m.Update(keyPress("r")) require.NotNil(t, cmd) _ = cmd() - require.Equal(t, []bool{true}, controller.refreshManual) + controller.AssertCalled(t, "Refresh", mock.Anything, mock.Anything, true) + controller.AssertExpectations(t) } From 7b5870c5f56a3de0702d53d03484cfb00f17ffaf Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 06:36:59 +0200 Subject: [PATCH 12/22] test: replace port fakes with generated mocks --- .../adapters/in/tui/active_provider_test.go | 4 +- .../in/tui/inline_review_interaction_test.go | 19 +- .../adapters/in/tui/review_publish_test.go | 112 +++++------ internal/app/active_provider_service_test.go | 183 +++++++++--------- internal/app/app_test.go | 74 ++----- 5 files changed, 163 insertions(+), 229 deletions(-) diff --git a/internal/adapters/in/tui/active_provider_test.go b/internal/adapters/in/tui/active_provider_test.go index 569e783..cd0c34d 100644 --- a/internal/adapters/in/tui/active_provider_test.go +++ b/internal/adapters/in/tui/active_provider_test.go @@ -10,6 +10,7 @@ import ( "ero/internal/core" "ero/internal/ports" + portmocks "ero/internal/ports/mocks" ) type mockActiveProviderController struct{ mock.Mock } @@ -50,7 +51,8 @@ func TestActiveProviderStartupLoadsOnlyActiveProviderState(t *testing.T) { startState := ActiveProviderState{StableProviderKey: "github", RuntimeProviderID: "github-runtime", RuntimeInfo: core.ReviewProviderInfo{ID: "github-runtime", Label: "GitHub"}, Snapshot: core.ProviderSnapshot{Threads: []core.RemoteReviewThread{{ProviderID: "github-runtime", ExternalID: "t1"}}, Sync: core.ProviderSyncState{Status: core.ProviderSyncStatusSynced}}} controller.On("Catalog", mock.Anything).Return(catalog, nil).Once() controller.On("Start", mock.Anything, mock.Anything).Return(startState, nil).Once() - m := NewModelWithActiveProviderContext(context.Background(), []core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, []ports.ReviewProviderClient{fakeProviderAsPort{&fakeReviewProvider{}}}) + legacyProvider := portmocks.NewMockReviewProviderClient(t) + m := NewModelWithActiveProviderContext(context.Background(), []core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, []ports.ReviewProviderClient{legacyProvider}) cmd := m.Init() require.NotNil(t, cmd) diff --git a/internal/adapters/in/tui/inline_review_interaction_test.go b/internal/adapters/in/tui/inline_review_interaction_test.go index 2946c87..b044d9f 100644 --- a/internal/adapters/in/tui/inline_review_interaction_test.go +++ b/internal/adapters/in/tui/inline_review_interaction_test.go @@ -14,15 +14,6 @@ import ( "ero/internal/ports/mocks" ) -type deadlineRecordingClipboardWriter struct { - hadDeadline bool -} - -func (w *deadlineRecordingClipboardWriter) WriteClipboard(ctx context.Context, _ string) error { - _, w.hadDeadline = ctx.Deadline() - return nil -} - func TestModelOpenCommentEditorUsesSelectedLineRange(t *testing.T) { model := NewModel([]core.ReviewFile{reviewFileWithLines("demo.go", 3)}) updated, _ := model.Update(keyPress("s")) @@ -78,8 +69,8 @@ func TestModelInlineCommentSubmitCopiesReviewJSON(t *testing.T) { } func TestModelInlineCommentSubmitCopiesReviewJSONWithClipboardDeadline(t *testing.T) { - writer := &deadlineRecordingClipboardWriter{} - model := NewModelWithClipboardWriter([]core.ReviewFile{reviewFileWithLines("demo.go", 1)}, nil, nil, core.ReviewRequest{DiffMode: core.DiffModeBranch}, writer) + clipboard := mocks.NewMockClipboardWriter(t) + model := NewModelWithClipboardWriter([]core.ReviewFile{reviewFileWithLines("demo.go", 1)}, nil, nil, core.ReviewRequest{DiffMode: core.DiffModeBranch}, clipboard) updated, _ := model.Update(keyPress("c")) model = updated.(Model) @@ -88,13 +79,17 @@ func TestModelInlineCommentSubmitCopiesReviewJSONWithClipboardDeadline(t *testin model = updated.(Model) } + hadDeadline := false + clipboard.EXPECT().WriteClipboard(mock.Anything, mock.Anything).Run(func(ctx context.Context, _ string) { + _, hadDeadline = ctx.Deadline() + }).Return(nil).Once() updated, cmd := model.Update(tea.KeyPressMsg{Code: tea.KeyEnter, Mod: tea.ModCtrl}) model = updated.(Model) require.NotNil(t, cmd) updated, _ = model.Update(cmd()) model = updated.(Model) - assert.True(t, writer.hadDeadline, "clipboard writes should use a deadline so status feedback cannot remain stuck on Copying review JSON… forever") + assert.True(t, hadDeadline, "clipboard writes should use a deadline so status feedback cannot remain stuck on Copying review JSON… forever") assert.Contains(t, model.copyFeedback, "Review JSON copied") } diff --git a/internal/adapters/in/tui/review_publish_test.go b/internal/adapters/in/tui/review_publish_test.go index cbea3b7..656e03d 100644 --- a/internal/adapters/in/tui/review_publish_test.go +++ b/internal/adapters/in/tui/review_publish_test.go @@ -1,23 +1,23 @@ package tui import ( - "context" "errors" "testing" tea "charm.land/bubbletea/v2" + "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" "ero/internal/core" "ero/internal/ports" + "ero/internal/ports/mocks" ) func TestUppercasePOpensPublishOverlayWithProvider(t *testing.T) { - provider := &fakeReviewProvider{info: core.ReviewProviderInfo{ID: "pi-coding-agent", Label: "pi-coding-agent", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}}} + info := core.ReviewProviderInfo{ID: "pi-coding-agent", Label: "pi-coding-agent", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}} + provider := mocks.NewMockReviewProviderClient(t) m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, nil) - m.reviewProviders = []ports.ReviewProviderClient{fakeProviderAsPort{provider}} - m.providerInfos = []core.ReviewProviderInfo{provider.info} - m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{m.reviewProviders[0]: provider.info} + attachProviderClient(&m, provider, info) updated, cmd := m.Update(keyPress("P")) m = updated.(Model) @@ -60,11 +60,13 @@ func TestPublishOverlaySupportsKeyboardFocusAndToggle(t *testing.T) { } func TestPublishReviewSuccess(t *testing.T) { - provider := &fakeReviewProvider{info: core.ReviewProviderInfo{ID: "github", Label: "GitHub", Capabilities: core.ReviewProviderCapabilities{PublishReview: true, Decisions: []core.ReviewDecision{core.ReviewDecisionComment}}}} + info := core.ReviewProviderInfo{ID: "github", Label: "GitHub", Capabilities: core.ReviewProviderCapabilities{PublishReview: true, Decisions: []core.ReviewDecision{core.ReviewDecisionComment}}} + provider := mocks.NewMockReviewProviderClient(t) + provider.EXPECT().PublishReview(mock.Anything, mock.MatchedBy(func(req core.PublishReviewRequest) bool { + return req.ProviderID == "github" && req.Draft.Decision == core.ReviewDecisionComment + })).Return(core.PublishReviewResult{ProviderID: "github", ExternalReviewID: "review-1"}, nil).Once() m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, nil) - m.reviewProviders = []ports.ReviewProviderClient{fakeProviderAsPort{provider}} - m.providerInfos = []core.ReviewProviderInfo{provider.info} - m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{m.reviewProviders[0]: provider.info} + attachProviderClient(&m, provider, info) m.reviewDraft.SetDecision(core.ReviewDecisionComment) m, _ = m.openPublishReview() updated, cmd := m.publishSelectedProviders() @@ -78,8 +80,7 @@ func TestPublishReviewSuccess(t *testing.T) { model, _ := m.Update(msg) m = model.(Model) require.False(t, m.publish.active) - require.Len(t, provider.requests, 1) - require.Equal(t, core.ReviewDecisionComment, provider.requests[0].Draft.Decision) + provider.AssertNumberOfCalls(t, "PublishReview", 1) comments := m.reviewDraft.Comments() require.Len(t, comments, 1) require.Len(t, comments[0].ProviderRefs, 1) @@ -87,11 +88,11 @@ func TestPublishReviewSuccess(t *testing.T) { } func TestPublishReviewFailedProvider(t *testing.T) { - provider := &fakeReviewProvider{info: core.ReviewProviderInfo{ID: "github", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}}, publishErr: errors.New("auth required")} + info := core.ReviewProviderInfo{ID: "github", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}} + provider := mocks.NewMockReviewProviderClient(t) + provider.EXPECT().PublishReview(mock.Anything, mock.MatchedBy(func(req core.PublishReviewRequest) bool { return req.ProviderID == "github" })).Return(core.PublishReviewResult{}, errors.New("auth required")).Once() m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, nil) - m.reviewProviders = []ports.ReviewProviderClient{fakeProviderAsPort{provider}} - m.providerInfos = []core.ReviewProviderInfo{provider.info} - m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{m.reviewProviders[0]: provider.info} + attachProviderClient(&m, provider, info) m, _ = m.openPublishReview() updated, cmd := m.publishSelectedProviders() m = updated @@ -102,9 +103,9 @@ func TestPublishReviewFailedProvider(t *testing.T) { } func TestStatusBarShowsProviderPublishHint(t *testing.T) { - provider := &fakeReviewProvider{info: core.ReviewProviderInfo{ID: "pi-coding-agent", Label: "pi-coding-agent", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}}} + info := core.ReviewProviderInfo{ID: "pi-coding-agent", Label: "pi-coding-agent", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}} m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, nil) - m.providerInfos = []core.ReviewProviderInfo{provider.info} + m.providerInfos = []core.ReviewProviderInfo{info} view := stripANSI(m.View().Content) require.Contains(t, view, "1 provider") @@ -112,11 +113,11 @@ func TestStatusBarShowsProviderPublishHint(t *testing.T) { } func TestProviderUnavailableReasonAppearsInStatusBar(t *testing.T) { - provider := &fakeReviewProvider{ - info: core.ReviewProviderInfo{ID: "pi-coding-agent", Label: "pi-coding-agent", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}}, - detection: &core.DetectionResult{Applicable: false, Reason: "no active bridge session"}, - } - m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, []ports.ReviewProviderClient{fakeProviderAsPort{provider}}) + info := core.ReviewProviderInfo{ID: "pi-coding-agent", Label: "pi-coding-agent", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}} + provider := mocks.NewMockReviewProviderClient(t) + provider.EXPECT().Initialize(mock.Anything).Return(info, nil).Once() + provider.EXPECT().DetectContext(mock.Anything, mock.Anything).Return(core.DetectionResult{Applicable: false, Reason: "no active bridge session"}, nil).Once() + m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, []ports.ReviewProviderClient{provider}) cmd := m.Init() require.NotNil(t, cmd) @@ -133,72 +134,49 @@ func publishTestRange() core.ReviewLineRange { } func TestPublishReviewUsesClientMatchedByProviderID(t *testing.T) { - firstClient := &fakeReviewProvider{info: core.ReviewProviderInfo{ID: "first", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}}} - selectedClient := &fakeReviewProvider{info: core.ReviewProviderInfo{ID: "selected", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}}} + firstInfo := core.ReviewProviderInfo{ID: "first", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}} + selectedInfo := core.ReviewProviderInfo{ID: "selected", Capabilities: core.ReviewProviderCapabilities{PublishReview: true}} + firstClient := mocks.NewMockReviewProviderClient(t) + selectedClient := mocks.NewMockReviewProviderClient(t) + selectedClient.EXPECT().PublishReview(mock.Anything, mock.MatchedBy(func(req core.PublishReviewRequest) bool { return req.ProviderID == "selected" })).Return(core.PublishReviewResult{ProviderID: "selected"}, nil).Once() m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, nil) - m.reviewProviders = []ports.ReviewProviderClient{fakeProviderAsPort{firstClient}, fakeProviderAsPort{selectedClient}} - m.providerInfos = []core.ReviewProviderInfo{selectedClient.info} - m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{m.reviewProviders[0]: firstClient.info, m.reviewProviders[1]: selectedClient.info} + m.reviewProviders = []ports.ReviewProviderClient{firstClient, selectedClient} + m.providerInfos = []core.ReviewProviderInfo{selectedInfo} + m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{m.reviewProviders[0]: firstInfo, m.reviewProviders[1]: selectedInfo} m, _ = m.openPublishReview() updated, cmd := m.publishSelectedProviders() m = updated require.NotNil(t, cmd) _ = cmd().(publishReviewCompletedMsg) - require.Empty(t, firstClient.requests) - require.Len(t, selectedClient.requests, 1) + firstClient.AssertNotCalled(t, "PublishReview", mock.Anything, mock.Anything) + selectedClient.AssertNumberOfCalls(t, "PublishReview", 1) } func TestPublishReviewUnsupportedDecisionWarning(t *testing.T) { - provider := &fakeReviewProvider{info: core.ReviewProviderInfo{ID: "pi-coding-agent", Capabilities: core.ReviewProviderCapabilities{PublishReview: true, Decisions: []core.ReviewDecision{core.ReviewDecisionComment}}}} + info := core.ReviewProviderInfo{ID: "pi-coding-agent", Capabilities: core.ReviewProviderCapabilities{PublishReview: true, Decisions: []core.ReviewDecision{core.ReviewDecisionComment}}} + provider := mocks.NewMockReviewProviderClient(t) + provider.EXPECT().PublishReview(mock.Anything, mock.MatchedBy(func(req core.PublishReviewRequest) bool { + return req.ProviderID == "pi-coding-agent" && req.Draft.Decision == "" + })).Return(core.PublishReviewResult{ProviderID: "pi-coding-agent"}, nil).Once() m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, nil) - m.reviewProviders = []ports.ReviewProviderClient{fakeProviderAsPort{provider}} - m.providerInfos = []core.ReviewProviderInfo{provider.info} - m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{m.reviewProviders[0]: provider.info} + attachProviderClient(&m, provider, info) m.reviewDraft.SetDecision(core.ReviewDecisionApprove) m, _ = m.openPublishReview() updated, cmd := m.publishSelectedProviders() m = updated require.Nil(t, cmd) require.Contains(t, m.publish.message, "Decision unsupported") - require.Empty(t, provider.requests) + provider.AssertNotCalled(t, "PublishReview", mock.Anything, mock.Anything) updated, cmd = m.publishSelectedProviders() m = updated require.NotNil(t, cmd) _ = cmd().(publishReviewCompletedMsg) - require.Len(t, provider.requests, 1) - require.Empty(t, provider.requests[0].Draft.Decision) + provider.AssertNumberOfCalls(t, "PublishReview", 1) } -type fakeReviewProvider struct { - info core.ReviewProviderInfo - detection *core.DetectionResult - threads []core.RemoteReviewThread - publishErr error - requests []core.PublishReviewRequest +func attachProviderClient(m *Model, provider ports.ReviewProviderClient, info core.ReviewProviderInfo) { + m.reviewProviders = []ports.ReviewProviderClient{provider} + m.providerInfos = []core.ReviewProviderInfo{info} + m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{provider: info} } - -type fakeProviderAsPort struct{ *fakeReviewProvider } - -func (f fakeProviderAsPort) Initialize(context.Context) (core.ReviewProviderInfo, error) { - return f.info, nil -} -func (f fakeProviderAsPort) DetectContext(context.Context, core.ReviewContext) (core.DetectionResult, error) { - if f.detection != nil { - return *f.detection, nil - } - return core.DetectionResult{Applicable: true}, nil -} -func (f fakeProviderAsPort) LoadRemoteThreads(context.Context, core.ReviewContext) ([]core.RemoteReviewThread, error) { - return f.threads, nil -} -func (f fakeProviderAsPort) PublishReview(_ context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) { - f.requests = append(f.requests, request) - if f.publishErr != nil { - return core.PublishReviewResult{}, f.publishErr - } - return core.PublishReviewResult{ProviderID: request.ProviderID}, nil -} -func (f fakeProviderAsPort) Close() error { return nil } - -var _ tea.Cmd diff --git a/internal/app/active_provider_service_test.go b/internal/app/active_provider_service_test.go index 85804a5..9f08f35 100644 --- a/internal/app/active_provider_service_test.go +++ b/internal/app/active_provider_service_test.go @@ -6,58 +6,13 @@ import ( "testing" "time" + "github.com/stretchr/testify/mock" + "ero/internal/core" "ero/internal/ports" + "ero/internal/ports/mocks" ) -type memCatalog []ports.ReviewProviderDescriptor - -func (m memCatalog) ListReviewProviderDescriptors(context.Context) ([]ports.ReviewProviderDescriptor, error) { - return []ports.ReviewProviderDescriptor(m), nil -} - -type memFactory struct { - clients map[string]*fakeProvider - made []string - beforeCreate func(string) -} - -func (f *memFactory) CreateReviewProviderClient(_ context.Context, d ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { - if f.beforeCreate != nil { - f.beforeCreate(d.Key) - } - f.made = append(f.made, d.Key) - return f.clients[d.Key], nil -} - -type fakeProvider struct { - id string - applicable bool - detectErr, errorLoad error - closed int - threads []core.RemoteReviewThread -} - -func (f *fakeProvider) Initialize(context.Context) (core.ReviewProviderInfo, error) { - return core.ReviewProviderInfo{ID: f.id, Capabilities: core.ReviewProviderCapabilities{LoadRemoteComments: true}}, nil -} -func (f *fakeProvider) DetectContext(context.Context, core.ReviewContext) (core.DetectionResult, error) { - if f.detectErr != nil { - return core.DetectionResult{}, f.detectErr - } - return core.DetectionResult{Applicable: f.applicable, Reason: "nope"}, nil -} -func (f *fakeProvider) LoadRemoteThreads(context.Context, core.ReviewContext) ([]core.RemoteReviewThread, error) { - if f.errorLoad != nil { - return nil, f.errorLoad - } - return f.threads, nil -} -func (f *fakeProvider) PublishReview(context.Context, core.PublishReviewRequest) (core.PublishReviewResult, error) { - return core.PublishReviewResult{}, nil -} -func (f *fakeProvider) Close() error { f.closed++; return nil } - type memCache struct { snap core.ProviderSnapshot ok bool @@ -90,11 +45,31 @@ func testReviewContext() core.ReviewContext { return core.ReviewContext{Repository: core.RepositoryMetadata{Remotes: []core.GitRemote{{URL: "https://github.com/acme/repo.git"}}}, Target: core.ReviewTargetMetadata{Mode: core.DiffModeWorking}} } +func mockCatalog(t *testing.T, descriptors ...ports.ReviewProviderDescriptor) *mocks.MockReviewProviderCatalog { + catalog := mocks.NewMockReviewProviderCatalog(t) + catalog.EXPECT().ListReviewProviderDescriptors(mock.Anything).Return(descriptors, nil) + return catalog +} + +func expectFactoryClient(factory *mocks.MockReviewProviderClientFactory, key string, client ports.ReviewProviderClient) { + factory.EXPECT().CreateReviewProviderClient(mock.Anything, mock.MatchedBy(func(d ports.ReviewProviderDescriptor) bool { return d.Key == key })).Return(client, nil).Once() +} + +func expectProbe(provider *mocks.MockReviewProviderClient, id string, applicable bool) { + provider.EXPECT().Initialize(mock.Anything).Return(core.ReviewProviderInfo{ID: id, Capabilities: core.ReviewProviderCapabilities{LoadRemoteComments: true}}, nil).Once() + provider.EXPECT().DetectContext(mock.Anything, mock.Anything).Return(core.DetectionResult{Applicable: applicable, Reason: "nope"}, nil).Once() +} + func TestActiveProviderServicePreferenceFallbackAndClosesFailedClients(t *testing.T) { - bad := &fakeProvider{id: "bad", applicable: false} - good := &fakeProvider{id: "good", applicable: true} - fac := &memFactory{clients: map[string]*fakeProvider{"preferred": bad, "github": good}} - svc := NewActiveProviderService(memCatalog{{Key: "preferred"}, {Key: "github", Type: "github"}}, fac, nil, &memPrefs{key: "preferred", ok: true}, ProviderPollingConfig{}) + bad := mocks.NewMockReviewProviderClient(t) + expectProbe(bad, "bad", false) + bad.EXPECT().Close().Return(nil).Once() + good := mocks.NewMockReviewProviderClient(t) + expectProbe(good, "good", true) + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "preferred", bad) + expectFactoryClient(factory, "github", good) + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "preferred"}, ports.ReviewProviderDescriptor{Key: "github", Type: "github"}), factory, nil, &memPrefs{key: "preferred", ok: true}, ProviderPollingConfig{}) st, err := svc.Start(context.Background(), testReviewContext()) if err != nil { t.Fatal(err) @@ -102,12 +77,6 @@ func TestActiveProviderServicePreferenceFallbackAndClosesFailedClients(t *testin if st.StableProviderKey != "github" { t.Fatalf("got %q", st.StableProviderKey) } - if bad.closed != 1 { - t.Fatalf("failed client not closed") - } - if good.closed != 0 { - t.Fatalf("active client was closed") - } } func TestActiveProviderServiceCacheFirstRefreshPreservesCacheOnRetryableFailure(t *testing.T) { @@ -115,8 +84,12 @@ func TestActiveProviderServiceCacheFirstRefreshPreservesCacheOnRetryableFailure( key := core.NewReviewContextKey("github", review) cached := core.ProviderSnapshot{StableProviderKey: "github", ContextKey: key, Threads: []core.RemoteReviewThread{{ExternalID: "old"}}} cache := &memCache{snap: cached, ok: true} - p := &fakeProvider{id: "rt", applicable: true, errorLoad: core.NewProviderError(core.ProviderErrorTransientNetwork, "offline", errors.New("dial"))} - svc := NewActiveProviderService(memCatalog{{Key: "github", Type: "github"}}, &memFactory{clients: map[string]*fakeProvider{"github": p}}, cache, nil, ProviderPollingConfig{Interval: time.Minute, MinBackoff: time.Second, MaxBackoff: time.Second}) + provider := mocks.NewMockReviewProviderClient(t) + expectProbe(provider, "rt", true) + provider.EXPECT().LoadRemoteThreads(mock.Anything, mock.Anything).Return(nil, core.NewProviderError(core.ProviderErrorTransientNetwork, "offline", errors.New("dial"))).Once() + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "github", provider) + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "github", Type: "github"}), factory, cache, nil, ProviderPollingConfig{Interval: time.Minute, MinBackoff: time.Second, MaxBackoff: time.Second}) st, err := svc.Start(context.Background(), review) if err != nil { t.Fatal(err) @@ -141,18 +114,23 @@ func TestActiveProviderServiceCacheFirstRefreshPreservesCacheOnRetryableFailure( func TestActiveProviderServiceSwitchGenerationIgnoresStaleTimer(t *testing.T) { review := testReviewContext() - a := &fakeProvider{id: "a", applicable: true} - b := &fakeProvider{id: "b", applicable: true} - svc := NewActiveProviderService(memCatalog{{Key: "a"}, {Key: "b"}}, &memFactory{clients: map[string]*fakeProvider{"a": a, "b": b}}, nil, nil, ProviderPollingConfig{}) + a := mocks.NewMockReviewProviderClient(t) + expectProbe(a, "a", true) + aClosed := false + a.EXPECT().Close().Run(func() { aClosed = true }).Return(nil).Once() + b := mocks.NewMockReviewProviderClient(t) + expectProbe(b, "b", true) + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "a", a) + expectFactoryClient(factory, "b", b) + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "a"}, ports.ReviewProviderDescriptor{Key: "b"}), factory, nil, nil, ProviderPollingConfig{}) st, err := svc.Start(context.Background(), review) if err != nil { t.Fatal(err) } target := "b" - oldProvider := a if st.StableProviderKey == "b" { - target = "a" - oldProvider = b + t.Fatal("expected deterministic initial provider a") } old := svc.Generation() if _, err := svc.Switch(context.Background(), review, target); err != nil { @@ -164,24 +142,29 @@ func TestActiveProviderServiceSwitchGenerationIgnoresStaleTimer(t *testing.T) { if got := svc.State().StableProviderKey; got != target { t.Fatalf("stale timer changed state to %q", got) } - if oldProvider.closed != 1 { + if !aClosed { t.Fatalf("switch should close old client") } } func TestActiveProviderServiceSwitchClosesCurrentBeforeStartingTarget(t *testing.T) { review := testReviewContext() - a := &fakeProvider{id: "a", applicable: true} - b := &fakeProvider{id: "b", applicable: true} - factory := &memFactory{clients: map[string]*fakeProvider{"a": a, "b": b}} - svc := NewActiveProviderService(memCatalog{{Key: "a"}, {Key: "b"}}, factory, nil, nil, ProviderPollingConfig{}) - if _, err := svc.Start(context.Background(), review); err != nil { - t.Fatal(err) - } - factory.beforeCreate = func(key string) { - if key == "b" && a.closed != 1 { + a := mocks.NewMockReviewProviderClient(t) + expectProbe(a, "a", true) + aClosed := false + a.EXPECT().Close().Run(func() { aClosed = true }).Return(nil).Once() + b := mocks.NewMockReviewProviderClient(t) + expectProbe(b, "b", true) + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "a", a) + factory.EXPECT().CreateReviewProviderClient(mock.Anything, mock.MatchedBy(func(d ports.ReviewProviderDescriptor) bool { return d.Key == "b" })).Run(func(_ context.Context, _ ports.ReviewProviderDescriptor) { + if !aClosed { t.Fatalf("current provider was still live when target provider started") } + }).Return(b, nil).Once() + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "a"}, ports.ReviewProviderDescriptor{Key: "b"}), factory, nil, nil, ProviderPollingConfig{}) + if _, err := svc.Start(context.Background(), review); err != nil { + t.Fatal(err) } if _, err := svc.Switch(context.Background(), review, "b"); err != nil { t.Fatal(err) @@ -190,9 +173,17 @@ func TestActiveProviderServiceSwitchClosesCurrentBeforeStartingTarget(t *testing func TestActiveProviderServiceFailedSwitchClearsOldProviderState(t *testing.T) { review := testReviewContext() - a := &fakeProvider{id: "a", applicable: true, threads: []core.RemoteReviewThread{{ExternalID: "old"}}} - b := &fakeProvider{id: "b", applicable: false} - svc := NewActiveProviderService(memCatalog{{Key: "a"}, {Key: "b"}}, &memFactory{clients: map[string]*fakeProvider{"a": a, "b": b}}, nil, nil, ProviderPollingConfig{}) + a := mocks.NewMockReviewProviderClient(t) + expectProbe(a, "a", true) + a.EXPECT().LoadRemoteThreads(mock.Anything, mock.Anything).Return([]core.RemoteReviewThread{{ExternalID: "old"}}, nil).Once() + a.EXPECT().Close().Return(nil).Once() + b := mocks.NewMockReviewProviderClient(t) + expectProbe(b, "b", false) + b.EXPECT().Close().Return(nil).Once() + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "a", a) + expectFactoryClient(factory, "b", b) + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "a"}, ports.ReviewProviderDescriptor{Key: "b"}), factory, nil, nil, ProviderPollingConfig{}) if _, err := svc.Start(context.Background(), review); err != nil { t.Fatal(err) } @@ -213,35 +204,37 @@ func TestActiveProviderServiceFailedSwitchClearsOldProviderState(t *testing.T) { if state := svc.State(); state.StableProviderKey != "" || len(state.Snapshot.Threads) != 0 || state.LastError == nil { t.Fatalf("failed switch left incoherent service state: %#v", state) } - if a.closed != 1 { - t.Fatalf("old provider should be closed, got %d", a.closed) - } } func TestActiveProviderServiceUsesStableCatalogOrder(t *testing.T) { review := testReviewContext() - fac := &memFactory{clients: map[string]*fakeProvider{ - "first": {id: "first", applicable: true}, - "github": {id: "github", applicable: true}, - "second": {id: "second", applicable: true}, - }} - svc := NewActiveProviderService(memCatalog{{Key: "first"}, {Key: "github", ContributionID: "github"}, {Key: "second"}}, fac, nil, nil, ProviderPollingConfig{}) + github := mocks.NewMockReviewProviderClient(t) + expectProbe(github, "github", true) + factory := mocks.NewMockReviewProviderClientFactory(t) + made := []string{} + factory.EXPECT().CreateReviewProviderClient(mock.Anything, mock.MatchedBy(func(d ports.ReviewProviderDescriptor) bool { return d.Key == "github" })).Run(func(_ context.Context, d ports.ReviewProviderDescriptor) { made = append(made, d.Key) }).Return(github, nil).Once() + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "first"}, ports.ReviewProviderDescriptor{Key: "github", ContributionID: "github"}, ports.ReviewProviderDescriptor{Key: "second"}), factory, nil, nil, ProviderPollingConfig{}) if _, err := svc.Start(context.Background(), review); err != nil { t.Fatal(err) } - if got := fac.made; len(got) != 1 || got[0] != "github" { + if got := made; len(got) != 1 || got[0] != "github" { t.Fatalf("github fallback should be selected in catalog order, got %v", got) } - fac = &memFactory{clients: map[string]*fakeProvider{ - "z": {id: "z", applicable: false}, - "a": {id: "a", applicable: true}, - }} - svc = NewActiveProviderService(memCatalog{{Key: "z"}, {Key: "a"}}, fac, nil, nil, ProviderPollingConfig{}) + z := mocks.NewMockReviewProviderClient(t) + expectProbe(z, "z", false) + z.EXPECT().Close().Return(nil).Once() + a := mocks.NewMockReviewProviderClient(t) + expectProbe(a, "a", true) + factory = mocks.NewMockReviewProviderClientFactory(t) + made = []string{} + factory.EXPECT().CreateReviewProviderClient(mock.Anything, mock.MatchedBy(func(d ports.ReviewProviderDescriptor) bool { return d.Key == "z" })).Run(func(_ context.Context, d ports.ReviewProviderDescriptor) { made = append(made, d.Key) }).Return(z, nil).Once() + factory.EXPECT().CreateReviewProviderClient(mock.Anything, mock.MatchedBy(func(d ports.ReviewProviderDescriptor) bool { return d.Key == "a" })).Run(func(_ context.Context, d ports.ReviewProviderDescriptor) { made = append(made, d.Key) }).Return(a, nil).Once() + svc = NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "z"}, ports.ReviewProviderDescriptor{Key: "a"}), factory, nil, nil, ProviderPollingConfig{}) if _, err := svc.Start(context.Background(), review); err != nil { t.Fatal(err) } - if got := fac.made; len(got) != 2 || got[0] != "z" || got[1] != "a" { + if got := made; len(got) != 2 || got[0] != "z" || got[1] != "a" { t.Fatalf("remaining candidates should preserve catalog order, got %v", got) } } diff --git a/internal/app/app_test.go b/internal/app/app_test.go index 4f8d21d..8d2ec49 100644 --- a/internal/app/app_test.go +++ b/internal/app/app_test.go @@ -252,8 +252,10 @@ func TestBuildReviewProvidersUsesDescriptorsAndFactory(t *testing.T) { ctx := context.Background() descriptor := ports.ReviewProviderDescriptor{Key: "plugin@1.0.0/github", ContributionID: "github"} provider := mocks.NewMockReviewProviderClient(t) - catalog := &fakeReviewProviderCatalog{descriptors: []ports.ReviewProviderDescriptor{descriptor}} - factory := &fakeReviewProviderClientFactory{clients: map[string]ports.ReviewProviderClient{descriptor.Key: provider}} + catalog := mocks.NewMockReviewProviderCatalog(t) + catalog.EXPECT().ListReviewProviderDescriptors(ctx).Return([]ports.ReviewProviderDescriptor{descriptor}, nil) + factory := mocks.NewMockReviewProviderClientFactory(t) + factory.EXPECT().CreateReviewProviderClient(ctx, descriptor).Return(provider, nil) providers, err := buildReviewProviders(ctx, catalog, factory) require.NoError(t, err) @@ -261,13 +263,15 @@ func TestBuildReviewProvidersUsesDescriptorsAndFactory(t *testing.T) { } func TestBuildReviewContext(t *testing.T) { - metadata := &fakeGitMetadataReader{ - worktreeRoot: "/repo", - currentBranch: "feature", - defaultBranch: "main", - headSHA: "headsha", - remotes: []ports.GitRemoteInfo{{Name: "origin", URL: "git@example.com:repo.git"}}, - } + metadata := mocks.NewMockGitMetadataReader(t) + metadata.EXPECT().WorktreeRoot("/repo").Return("/repo", nil) + metadata.EXPECT().CurrentBranch("/repo").Return("feature", nil) + metadata.EXPECT().DefaultBranch("/repo").Return("main", nil) + metadata.EXPECT().HeadSHA("/repo").Return("headsha", nil) + metadata.EXPECT().Remotes("/repo").Return([]ports.GitRemoteInfo{{Name: "origin", URL: "git@example.com:repo.git"}}, nil) + metadata.EXPECT().ResolveRevision("/repo", "main").Return("mainsha", nil) + metadata.EXPECT().ResolveRevision("/repo", "feature").Return("featuresha", nil) + metadata.EXPECT().MergeBase("/repo", "main", "feature").Return("mergebase", nil) ctx := buildReviewContext(core.ReviewRequest{RepoPath: "/repo", DiffMode: core.DiffModeRange, BaseRevision: "main", HeadRevision: "feature"}, []core.ReviewFile{{ Path: "demo.go", OldPath: "old_demo.go", @@ -299,7 +303,13 @@ func TestBuildReviewContext(t *testing.T) { } func TestBuildReviewContextMetadataBestEffort(t *testing.T) { - ctx := buildReviewContext(core.ReviewRequest{RepoPath: "/repo", DiffMode: core.DiffModeBranch}, minimalReviewFiles(), &fakeGitMetadataReader{err: errors.New("git unavailable")}, "dev") + metadata := mocks.NewMockGitMetadataReader(t) + metadata.EXPECT().WorktreeRoot("/repo").Return("", errors.New("git unavailable")) + metadata.EXPECT().CurrentBranch("/repo").Return("", errors.New("git unavailable")) + metadata.EXPECT().DefaultBranch("/repo").Return("", errors.New("git unavailable")) + metadata.EXPECT().HeadSHA("/repo").Return("", errors.New("git unavailable")) + metadata.EXPECT().Remotes("/repo").Return(nil, errors.New("git unavailable")) + ctx := buildReviewContext(core.ReviewRequest{RepoPath: "/repo", DiffMode: core.DiffModeBranch}, minimalReviewFiles(), metadata, "dev") require.Equal(t, "/repo", ctx.Repository.RepoPath) require.Empty(t, ctx.Repository.WorktreeRoot) require.NotEmpty(t, ctx.Session.IdempotencyKey) @@ -316,50 +326,6 @@ func minimalReviewFiles() []core.ReviewFile { }} } -type fakeGitMetadataReader struct { - worktreeRoot string - currentBranch string - defaultBranch string - headSHA string - remotes []ports.GitRemoteInfo - err error -} - -func (f *fakeGitMetadataReader) WorktreeRoot(string) (string, error) { return f.worktreeRoot, f.err } -func (f *fakeGitMetadataReader) CurrentBranch(string) (string, error) { return f.currentBranch, f.err } -func (f *fakeGitMetadataReader) HeadSHA(string) (string, error) { return f.headSHA, f.err } -func (f *fakeGitMetadataReader) Remotes(string) ([]ports.GitRemoteInfo, error) { - return f.remotes, f.err -} -func (f *fakeGitMetadataReader) MergeBase(string, string, string) (string, error) { - return "mergebase", f.err -} -func (f *fakeGitMetadataReader) ResolveRevision(_ string, revision string) (string, error) { - return revision + "sha", f.err -} -func (f *fakeGitMetadataReader) DefaultBranch(string) (string, error) { return f.defaultBranch, f.err } - -type fakeReviewProviderCatalog struct { - descriptors []ports.ReviewProviderDescriptor - err error -} - -func (f *fakeReviewProviderCatalog) ListReviewProviderDescriptors(context.Context) ([]ports.ReviewProviderDescriptor, error) { - return f.descriptors, f.err -} - -type fakeReviewProviderClientFactory struct { - clients map[string]ports.ReviewProviderClient - err error -} - -func (f *fakeReviewProviderClientFactory) CreateReviewProviderClient(_ context.Context, descriptor ports.ReviewProviderDescriptor) (ports.ReviewProviderClient, error) { - if f.err != nil { - return nil, f.err - } - return f.clients[descriptor.Key], nil -} - type fakeStartupPrompt struct { mode core.DiffMode err error From 4c87d1418877e94b561b25c1a8bdfb8b022a27f0 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 06:54:49 +0200 Subject: [PATCH 13/22] feat(tui): improve provider switching UX --- .../adapters/in/tui/component/help_pane.go | 3 +- .../adapters/in/tui/component/statusbar.go | 69 +++++++++++++++++-- .../in/tui/component/statusbar_test.go | 33 ++++++--- internal/adapters/in/tui/help_pane_test.go | 4 +- internal/adapters/in/tui/keymap/action.go | 16 ++--- .../adapters/in/tui/keymap/action_test.go | 13 ++-- internal/adapters/in/tui/model.go | 2 +- internal/adapters/in/tui/provider_picker.go | 19 +++-- .../adapters/in/tui/provider_picker_test.go | 49 +++++++++++-- .../adapters/in/tui/review_publish_test.go | 2 +- 10 files changed, 165 insertions(+), 45 deletions(-) diff --git a/internal/adapters/in/tui/component/help_pane.go b/internal/adapters/in/tui/component/help_pane.go index d2e7f16..6a2b6b3 100644 --- a/internal/adapters/in/tui/component/help_pane.go +++ b/internal/adapters/in/tui/component/help_pane.go @@ -19,7 +19,8 @@ func RenderHelpPane(width, height int, enterKeyLabel, commentSubmitKeyLabel stri renderHelpShortcut("f", "find file", contentWidth), renderHelpShortcut("/", "grep references", contentWidth), renderHelpShortcut("d", "switch diff mode", contentWidth), - renderHelpShortcut("g/G", "cycle/pick provider", contentWidth), + renderHelpShortcut("p", "switch provider", contentWidth), + renderHelpShortcut("alt+p", "cycle provider", contentWidth), renderHelpShortcut("r/o", "refresh provider/PR sheet", contentWidth), "", theme.HelpSectionStyle.Render("Search"), diff --git a/internal/adapters/in/tui/component/statusbar.go b/internal/adapters/in/tui/component/statusbar.go index 7b5aaf7..75c6690 100644 --- a/internal/adapters/in/tui/component/statusbar.go +++ b/internal/adapters/in/tui/component/statusbar.go @@ -46,7 +46,7 @@ func (c StatusBar) Render(model StatusModel) string { {style: theme.StatusModeStyle, label: model.Mode}, {style: theme.StatusInfoStyle, label: fileCountLabel(model.FileCount)}, } - if model.ProviderCount > 0 { + if model.ProviderCount > 0 && (strings.TrimSpace(model.ActiveProviderLabel) == "" || !model.NerdFont) { segments = append(segments, statusSegment{style: theme.StatusInfoStyle, label: providerCountLabel(model.ProviderCount)}) } if syncLabel := providerSyncLabel(model); syncLabel != "" { @@ -86,7 +86,7 @@ type KeyHint struct { func renderStatusHint(width, providerCount int) string { hints := []KeyHint{{Key: "?", Label: "help"}} if providerCount > 0 { - hints = []KeyHint{{Key: "P", Label: "publish"}, {Key: "?", Label: "help"}} + hints = []KeyHint{{Key: "p", Label: "provider"}, {Key: "P", Label: "publish"}, {Key: "?", Label: "help"}} } full := RenderKeyHints(hints) if lipgloss.Width(full) <= width { @@ -94,7 +94,7 @@ func renderStatusHint(width, providerCount int) string { } fallback := "? help" if providerCount > 0 { - fallback = "P publish" + fallback = "p provider" } return theme.StatusInfoStyle.Render(TruncateRunes(fallback, max(width-theme.StatusInfoStyle.GetHorizontalPadding(), 0))) } @@ -152,11 +152,11 @@ func providerSyncLabel(model StatusModel) string { } parts := []string{provider} + if model.NerdFont { + parts = []string{compactProviderLabel(provider, model.ProviderCount, model.ProviderSync.Status)} + } status := providerSyncStatusLabel(model.ProviderSync.Status) - if status != "" { - if model.NerdFont { - status = providerSyncStatusSymbol(model.ProviderSync.Status) - } + if status != "" && !model.NerdFont { parts = append(parts, status) } if model.ProviderSync.LastError != "" { @@ -178,6 +178,61 @@ func draftCommentCountLabel(count int) string { return fmt.Sprintf("%d draft comments", count) } +func compactProviderLabel(provider string, providerCount int, status core.ProviderSyncStatus) string { + label := providerGlyph(provider) + providerStatusDot(status) + if providerCount > 1 { + label += fmt.Sprintf(" +%d", providerCount-1) + } + return label +} + +func providerGlyph(provider string) string { + provider = strings.ToLower(strings.TrimSpace(provider)) + if strings.Contains(provider, "github") { + return "" + } + return providerAbbreviation(provider) +} + +func providerStatusDot(status core.ProviderSyncStatus) string { + color := lipgloss.Color("81") + switch status { + case core.ProviderSyncStatusSynced: + color = lipgloss.Color("#3fb950") + case core.ProviderSyncStatusFailed: + color = lipgloss.Color("#ff7b72") + case core.ProviderSyncStatusBackingOff: + color = lipgloss.Color("#ffa657") + case core.ProviderSyncStatusLoadingCache, core.ProviderSyncStatusSyncing: + color = lipgloss.Color("#58a6ff") + } + return theme.StatusBaseStyle.Foreground(color).Render("●") +} + +func providerAbbreviation(provider string) string { + provider = strings.ToLower(strings.TrimSpace(provider)) + if provider == "" { + return "?" + } + switch { + case strings.Contains(provider, "github"): + return "gh" + case strings.Contains(provider, "gitlab"): + return "gl" + case strings.Contains(provider, "bitbucket"): + return "bb" + case strings.Contains(provider, "forgejo"): + return "fj" + case strings.Contains(provider, "gitea"): + return "gt" + } + runes := []rune(provider) + if len(runes) > 2 { + runes = runes[:2] + } + return string(runes) +} + func providerSyncStatusLabel(status core.ProviderSyncStatus) string { switch status { case core.ProviderSyncStatusLoadingCache: diff --git a/internal/adapters/in/tui/component/statusbar_test.go b/internal/adapters/in/tui/component/statusbar_test.go index cfaa9f5..913edfb 100644 --- a/internal/adapters/in/tui/component/statusbar_test.go +++ b/internal/adapters/in/tui/component/statusbar_test.go @@ -67,20 +67,20 @@ func TestStatusbarProviderSync(t *testing.T) { } } -func TestStatusbarProviderSyncUsesNerdFontSymbolWhenSupported(t *testing.T) { - for status, symbol := range map[core.ProviderSyncStatus]string{ - core.ProviderSyncStatusLoadingCache: "󰃨", - core.ProviderSyncStatusSyncing: "󰑓", - core.ProviderSyncStatusSynced: "", - core.ProviderSyncStatusFailed: "", - core.ProviderSyncStatusBackingOff: "󰌾", +func TestStatusbarProviderSyncUsesNerdFontProviderGlyphAndStatusDot(t *testing.T) { + for _, status := range []core.ProviderSyncStatus{ + core.ProviderSyncStatusLoadingCache, + core.ProviderSyncStatusSyncing, + core.ProviderSyncStatusSynced, + core.ProviderSyncStatusFailed, + core.ProviderSyncStatusBackingOff, } { model := syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: status}) model.NerdFont = true view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) - require.Contains(t, view, symbol) + require.Contains(t, view, "●") } } @@ -90,11 +90,24 @@ func TestStatusbarProviderSyncOmitsStatusWordWhenNerdFontSymbolIsShown(t *testin view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) - require.Contains(t, view, "GitHub/github") - require.Contains(t, view, "") + require.Contains(t, view, "●") + require.NotContains(t, view, "") require.NotContains(t, view, "synced") } +func TestStatusbarShowsCompactActiveProviderAndAdditionalCount(t *testing.T) { + model := syncStatusModel("GitHub", "github", core.ProviderSyncState{Status: core.ProviderSyncStatusSynced}) + model.ProviderCount = 2 + model.NerdFont = true + + view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) + + require.Contains(t, view, "● +1") + require.NotContains(t, view, "2 providers") + require.Contains(t, view, "p provider") + require.Contains(t, view, "P publish") +} + func TestStatusbarShowsDraftCommentCount(t *testing.T) { model := baseStatusModel() model.DraftCommentCount = 2 diff --git a/internal/adapters/in/tui/help_pane_test.go b/internal/adapters/in/tui/help_pane_test.go index e35d0ce..d1b98c4 100644 --- a/internal/adapters/in/tui/help_pane_test.go +++ b/internal/adapters/in/tui/help_pane_test.go @@ -24,7 +24,9 @@ func TestModelHelpModalShowsShortcutsAndCloses(t *testing.T) { assert.Contains(t, view, "select result") assert.Contains(t, view, "a expand all") assert.Contains(t, view, "expand more") - assert.Contains(t, view, "cycle/pick provider") + assert.Contains(t, view, "switch provider") + assert.Contains(t, view, "cycle provider") + assert.NotContains(t, view, "g/G") assert.Contains(t, view, "refresh provider") assert.Contains(t, view, "PR sheet") assert.NotContains(t, view, "a/b") diff --git a/internal/adapters/in/tui/keymap/action.go b/internal/adapters/in/tui/keymap/action.go index 403138d..25c3b10 100644 --- a/internal/adapters/in/tui/keymap/action.go +++ b/internal/adapters/in/tui/keymap/action.go @@ -57,13 +57,13 @@ func ReviewAction(key string) Action { return ActionOpenComment case "x": return ActionClearReview - case "P": + case "P", "shift+p": return ActionPublishReview - case "C": + case "C", "shift+c": return ActionCopyReviewJSON case "y": return ActionCopyPlain - case "Y": + case "Y", "shift+y": return ActionCopyWithMetadata case "f": return ActionOpenFileSearch @@ -71,23 +71,23 @@ func ReviewAction(key string) Action { return ActionOpenGrepSearch case "d": return ActionOpenDiffMode - case "left", "h", "p": + case "left", "h": return ActionPreviousFile - case "right", "l", "n": + case "right", "l": return ActionNextFile case "a": return ActionExpandAllContext case "enter": return ActionExpandMoreContext - case "g": + case "alt+p": return ActionCycleProvider - case "G": + case "p": return ActionOpenProviderPicker case "r": return ActionRefreshProvider case "o": return ActionTogglePRSheet - case "?": + case "?", "shift+/": return ActionOpenHelp default: return ActionNone diff --git a/internal/adapters/in/tui/keymap/action_test.go b/internal/adapters/in/tui/keymap/action_test.go index 243a612..d71ec55 100644 --- a/internal/adapters/in/tui/keymap/action_test.go +++ b/internal/adapters/in/tui/keymap/action_test.go @@ -30,25 +30,30 @@ func TestReviewAction(t *testing.T) { {name: "open comment", key: "c", want: ActionOpenComment}, {name: "clear review binding", key: "x", want: ActionClearReview}, {name: "publish review shifted binding", key: "P", want: ActionPublishReview}, + {name: "publish review explicit shift binding", key: "shift+p", want: ActionPublishReview}, {name: "copy review json shifted binding", key: "C", want: ActionCopyReviewJSON}, + {name: "copy review json explicit shift binding", key: "shift+c", want: ActionCopyReviewJSON}, {name: "copy plain", key: "y", want: ActionCopyPlain}, {name: "copy with metadata shifted binding", key: "Y", want: ActionCopyWithMetadata}, + {name: "copy with metadata explicit shift binding", key: "shift+y", want: ActionCopyWithMetadata}, {name: "open file search", key: "f", want: ActionOpenFileSearch}, {name: "open grep search", key: "/", want: ActionOpenGrepSearch}, {name: "open diff mode", key: "d", want: ActionOpenDiffMode}, {name: "previous file left arrow", key: "left", want: ActionPreviousFile}, {name: "previous file vim alias", key: "h", want: ActionPreviousFile}, - {name: "previous file p alias", key: "p", want: ActionPreviousFile}, + {name: "open provider picker", key: "p", want: ActionOpenProviderPicker}, {name: "next file right arrow", key: "right", want: ActionNextFile}, {name: "next file vim alias", key: "l", want: ActionNextFile}, - {name: "next file n alias", key: "n", want: ActionNextFile}, + {name: "next file n removed", key: "n", want: ActionNone}, {name: "expand all context", key: "a", want: ActionExpandAllContext}, {name: "expand more context enter binding", key: "enter", want: ActionExpandMoreContext}, - {name: "cycle provider", key: "g", want: ActionCycleProvider}, - {name: "open provider picker", key: "G", want: ActionOpenProviderPicker}, + {name: "old provider cycle binding removed", key: "g", want: ActionNone}, + {name: "old provider picker binding removed", key: "G", want: ActionNone}, + {name: "cycle provider alt p", key: "alt+p", want: ActionCycleProvider}, {name: "refresh provider", key: "r", want: ActionRefreshProvider}, {name: "toggle pr sheet", key: "o", want: ActionTogglePRSheet}, {name: "open help", key: "?", want: ActionOpenHelp}, + {name: "open help explicit shift binding", key: "shift+/", want: ActionOpenHelp}, {name: "unknown", key: "R", want: ActionNone}, } diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index b28c93d..81bd3bf 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -346,7 +346,7 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { if m.publish.active { return m.updatePublishReview(msg) } - return m.updateReviewAction(keymap.ReviewAction(msg.String())) + return m.updateReviewAction(keymap.ReviewAction(msg.Keystroke())) default: return m, nil } diff --git a/internal/adapters/in/tui/provider_picker.go b/internal/adapters/in/tui/provider_picker.go index 7ca1a40..5861ce1 100644 --- a/internal/adapters/in/tui/provider_picker.go +++ b/internal/adapters/in/tui/provider_picker.go @@ -45,7 +45,7 @@ func (m Model) closeProviderPicker() Model { } func (m Model) updateProviderPicker(msg tea.KeyPressMsg) (tea.Model, tea.Cmd) { - switch msg.String() { + switch msg.Keystroke() { case "esc": return m.closeProviderPicker(), nil case "up", "k": @@ -61,6 +61,9 @@ func (m Model) updateProviderPicker(msg tea.KeyPressMsg) (tea.Model, tea.Cmd) { key := m.providerPicker.rows[m.providerPicker.selected].Key m = m.closeProviderPicker() return m, m.switchActiveProviderCmd(key) + case "alt+p": + m = m.closeProviderPicker() + return m, m.cycleProviderCmd() default: return m, nil } @@ -133,7 +136,7 @@ func (m Model) renderProviderPicker(width, height int) string { paneWidth := min(max(width-8, 36), 76) // Account for pane padding/chrome while preserving at least one content column. contentWidth := max(paneWidth-6, 1) - lines := []string{theme.HelpPaneTitleStyle.Render("Review providers"), ""} + lines := []string{theme.HelpPaneTitleStyle.Render("Provider"), theme.MutedStyle.Render("Active publish destination"), ""} rows := m.providerPicker.rows if len(rows) == 0 { lines = append(lines, theme.MutedStyle.Render("No providers discovered")) @@ -141,14 +144,16 @@ func (m Model) renderProviderPicker(width, height int) string { for i, row := range rows { cursor := " " if i == m.providerPicker.selected { - cursor = "> " + cursor = "› " } - active := " " + state := "available" + marker := "○" if row.Active { - active = "*" + state = "active" + marker = "●" } meta := strings.TrimSpace(strings.Join([]string{row.PluginName, row.PluginSource}, " ")) - line := fmt.Sprintf("%s%s %s", cursor, active, row.Label) + line := fmt.Sprintf("%s%s %-12s %s", cursor, marker, row.Label, state) if meta != "" { line += " — " + meta } @@ -158,7 +163,7 @@ func (m Model) renderProviderPicker(width, height int) string { lines = append(lines, theme.HelpLabelStyle.Render(componentTruncate(line, contentWidth))) } } - lines = append(lines, "", theme.HelpLabelStyle.Render("enter switch • esc close")) + lines = append(lines, "", theme.HelpLabelStyle.Render("enter switch • alt+p cycle • esc close")) lines = fitProviderPickerLines(lines, max(height-2, 1)) return theme.HelpPaneStyle.Width(paneWidth).Render(strings.Join(lines, "\n")) } diff --git a/internal/adapters/in/tui/provider_picker_test.go b/internal/adapters/in/tui/provider_picker_test.go index 0637863..cc055da 100644 --- a/internal/adapters/in/tui/provider_picker_test.go +++ b/internal/adapters/in/tui/provider_picker_test.go @@ -25,14 +25,15 @@ func TestProviderPickerDisplaysDescriptorRowsWithoutStartingInactiveProviders(t updated, _ := m.Update(m.Init()()) m = updated.(Model) - updated, _ = m.Update(keyPress("G")) + updated, _ = m.Update(keyPress("p")) m = updated.(Model) require.True(t, m.providerPicker.open) controller.AssertNumberOfCalls(t, "Start", 1) view := stripANSI(m.View().Content) - require.Contains(t, view, "Review providers") - require.Contains(t, view, "* GitHub") + require.Contains(t, view, "Provider") + require.Contains(t, view, "● GitHub") + require.Contains(t, view, "active") require.Contains(t, view, "gh-plugin builtin") require.Contains(t, view, "missing token") require.Contains(t, view, "GitLab") @@ -48,7 +49,7 @@ func TestProviderPickerSelectEmitsSwitchCommandWithStableKey(t *testing.T) { m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) updated, _ := m.Update(m.Init()()) m = updated.(Model) - updated, _ = m.Update(keyPress("G")) + updated, _ = m.Update(keyPress("p")) m = updated.(Model) updated, _ = m.Update(tea.KeyPressMsg{Code: tea.KeyDown}) m = updated.(Model) @@ -65,6 +66,44 @@ func TestProviderPickerSelectEmitsSwitchCommandWithStableKey(t *testing.T) { controller.AssertExpectations(t) } +func TestProviderPickerRendersActiveAndAvailableStates(t *testing.T) { + model := NewModel(nil) + model.providerCatalog = []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub", PluginName: "gh-plugin", PluginSource: "builtin"}, {Key: "gitlab", Label: "GitLab", PluginName: "gl-plugin", PluginSource: "local"}} + model.activeProviderKey = "github" + model.providerPicker = model.openProviderPicker().providerPicker + + view := stripANSI(model.renderProviderPicker(80, 20)) + + require.Contains(t, view, "● GitHub") + require.Contains(t, view, "active") + require.Contains(t, view, "○ GitLab") + require.Contains(t, view, "available") + require.Contains(t, view, "alt+p cycle") +} + +func TestProviderPickerAltPCyclesAndClosesPicker(t *testing.T) { + controller := &mockActiveProviderController{} + controller.On("Catalog", mock.Anything).Return([]ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "gitlab", Label: "GitLab"}}, nil).Once() + controller.On("Start", mock.Anything, mock.Anything).Return(ActiveProviderState{StableProviderKey: "github"}, nil).Once() + controller.On("Switch", mock.Anything, mock.Anything, "gitlab").Return(ActiveProviderState{StableProviderKey: "gitlab"}, nil).Once() + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, controller, nil) + updated, _ := m.Update(m.Init()()) + m = updated.(Model) + updated, _ = m.Update(keyPress("p")) + m = updated.(Model) + require.True(t, m.providerPicker.open) + + updated, cmd := m.Update(tea.KeyPressMsg{Text: "p", Code: 'p', Mod: tea.ModAlt}) + m = updated.(Model) + require.NotNil(t, cmd) + require.False(t, m.providerPicker.open) + updated, _ = m.Update(cmd()) + m = updated.(Model) + + controller.AssertCalled(t, "Switch", mock.Anything, mock.Anything, "gitlab") + controller.AssertExpectations(t) +} + func TestProviderCycleAndRefreshShortcuts(t *testing.T) { controller := &mockActiveProviderController{} controller.On("Catalog", mock.Anything).Return([]ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "gitlab", Label: "GitLab"}}, nil).Once() @@ -75,7 +114,7 @@ func TestProviderCycleAndRefreshShortcuts(t *testing.T) { updated, _ := m.Update(m.Init()()) m = updated.(Model) - updated, cmd := m.Update(keyPress("g")) + updated, cmd := m.Update(tea.KeyPressMsg{Text: "p", Code: 'p', Mod: tea.ModAlt}) m = updated.(Model) require.NotNil(t, cmd) updated, _ = m.Update(cmd()) diff --git a/internal/adapters/in/tui/review_publish_test.go b/internal/adapters/in/tui/review_publish_test.go index 656e03d..bfc53d3 100644 --- a/internal/adapters/in/tui/review_publish_test.go +++ b/internal/adapters/in/tui/review_publish_test.go @@ -19,7 +19,7 @@ func TestUppercasePOpensPublishOverlayWithProvider(t *testing.T) { m := NewModelWithReviewProviders([]core.ReviewFile{reviewFile("demo.go", "package main")}, nil, nil, core.ReviewRequest{}, nil, core.ReviewContext{}, nil) attachProviderClient(&m, provider, info) - updated, cmd := m.Update(keyPress("P")) + updated, cmd := m.Update(tea.KeyPressMsg{Text: "P", Code: 'p', Mod: tea.ModShift}) m = updated.(Model) require.Nil(t, cmd) From b592bd4a8d4d1c662914465bac784bc68b84c1b6 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 07:21:43 +0200 Subject: [PATCH 14/22] fix(providers): harden sync review follow-ups --- .../adapters/in/tui/component/statusbar.go | 32 ++---- .../in/tui/component/statusbar_test.go | 1 + internal/adapters/in/tui/model.go | 38 +++++++ internal/adapters/in/tui/pr_sheet_test.go | 15 +++ internal/adapters/in/tui/review_providers.go | 5 + .../adapters/out/plugin/provider_loader.go | 5 - internal/adapters/out/providercache/cache.go | 2 +- internal/app/active_provider_service.go | 100 ++++++++++++++---- internal/app/active_provider_service_test.go | 79 ++++++++++++++ internal/app/review_providers.go | 6 ++ internal/app/tui_active_provider.go | 2 +- internal/ports/plugin.go | 8 +- .../github/cmd/ero-plugin-github/graphql.go | 42 ++++++-- .../cmd/ero-plugin-github/graphql_test.go | 11 +- plugins/github/cmd/ero-plugin-github/main.go | 32 +++--- .../github/cmd/ero-plugin-github/main_test.go | 49 +++++++++ plugins/github/cmd/ero-plugin-github/match.go | 23 ++++ .../github/cmd/ero-plugin-github/remote.go | 17 ++- 18 files changed, 382 insertions(+), 85 deletions(-) diff --git a/internal/adapters/in/tui/component/statusbar.go b/internal/adapters/in/tui/component/statusbar.go index 75c6690..a1cc7ca 100644 --- a/internal/adapters/in/tui/component/statusbar.go +++ b/internal/adapters/in/tui/component/statusbar.go @@ -17,6 +17,7 @@ type StatusModel struct { Mode string FileCount int ProviderCount int + ProviderSwitch bool CurrentFile string Message string ScrollPercent float64 @@ -38,7 +39,7 @@ func NewStatusBar(width int) StatusBar { func (c StatusBar) Render(model StatusModel) string { width := max(c.width, 1) - right := renderStatusHint(width, model.ProviderCount) + right := renderStatusHint(width, model.ProviderCount, model.ProviderSwitch) leftWidth := max(width-lipgloss.Width(right)-1, 0) segments := []statusSegment{ @@ -83,10 +84,13 @@ type KeyHint struct { Label string } -func renderStatusHint(width, providerCount int) string { +func renderStatusHint(width, providerCount int, providerSwitch bool) string { hints := []KeyHint{{Key: "?", Label: "help"}} if providerCount > 0 { - hints = []KeyHint{{Key: "p", Label: "provider"}, {Key: "P", Label: "publish"}, {Key: "?", Label: "help"}} + hints = []KeyHint{{Key: "P", Label: "publish"}, {Key: "?", Label: "help"}} + if providerSwitch { + hints = []KeyHint{{Key: "p", Label: "provider"}, {Key: "P", Label: "publish"}, {Key: "?", Label: "help"}} + } } full := RenderKeyHints(hints) if lipgloss.Width(full) <= width { @@ -94,7 +98,10 @@ func renderStatusHint(width, providerCount int) string { } fallback := "? help" if providerCount > 0 { - fallback = "p provider" + fallback = "P publish" + if providerSwitch { + fallback = "p provider" + } } return theme.StatusInfoStyle.Render(TruncateRunes(fallback, max(width-theme.StatusInfoStyle.GetHorizontalPadding(), 0))) } @@ -250,23 +257,6 @@ func providerSyncStatusLabel(status core.ProviderSyncStatus) string { } } -func providerSyncStatusSymbol(status core.ProviderSyncStatus) string { - switch status { - case core.ProviderSyncStatusLoadingCache: - return "󰃨" - case core.ProviderSyncStatusSyncing: - return "󰑓" - case core.ProviderSyncStatusSynced: - return "" - case core.ProviderSyncStatusFailed: - return "" - case core.ProviderSyncStatusBackingOff: - return "󰌾" - default: - return "󰓦" - } -} - func formatStatusTime(value time.Time) string { return value.UTC().Format("15:04") } diff --git a/internal/adapters/in/tui/component/statusbar_test.go b/internal/adapters/in/tui/component/statusbar_test.go index 913edfb..b16843e 100644 --- a/internal/adapters/in/tui/component/statusbar_test.go +++ b/internal/adapters/in/tui/component/statusbar_test.go @@ -98,6 +98,7 @@ func TestStatusbarProviderSyncOmitsStatusWordWhenNerdFontSymbolIsShown(t *testin func TestStatusbarShowsCompactActiveProviderAndAdditionalCount(t *testing.T) { model := syncStatusModel("GitHub", "github", core.ProviderSyncState{Status: core.ProviderSyncStatusSynced}) model.ProviderCount = 2 + model.ProviderSwitch = true model.NerdFont = true view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index 81bd3bf..c59f9bb 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -346,6 +346,9 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { if m.publish.active { return m.updatePublishReview(msg) } + if m.prSheet.open { + return m.updatePRSheetAction(msg) + } return m.updateReviewAction(keymap.ReviewAction(msg.Keystroke())) default: return m, nil @@ -358,8 +361,14 @@ func (m Model) updateMouseWheel(msg tea.MouseWheelMsg) (tea.Model, tea.Cmd) { } switch msg.Mouse().Button { case tea.MouseWheelUp: + if m.prSheet.open { + return m.ScrollPRSheet(-1), nil + } m.moveCursor(-1) case tea.MouseWheelDown: + if m.prSheet.open { + return m.ScrollPRSheet(1), nil + } m.moveCursor(1) default: return m, nil @@ -367,6 +376,28 @@ func (m Model) updateMouseWheel(msg tea.MouseWheelMsg) (tea.Model, tea.Cmd) { return m, nil } +func (m Model) updatePRSheetAction(msg tea.KeyPressMsg) (tea.Model, tea.Cmd) { + switch keymap.ReviewAction(msg.Keystroke()) { + case keymap.ActionMoveUp: + return m.ScrollPRSheet(-1), nil + case keymap.ActionMoveDown: + return m.ScrollPRSheet(1), nil + case keymap.ActionPageUp: + return m.ScrollPRSheet(-max(m.height-1, 1)), nil + case keymap.ActionPageDown: + return m.ScrollPRSheet(max(m.height-1, 1)), nil + case keymap.ActionTogglePRSheet: + return m.TogglePRSheet(), nil + case keymap.ActionOpenHelp: + m.helpActive = true + return m, nil + case keymap.ActionQuit: + return m, tea.Batch(m.closeReviewProvidersCmd(), tea.Quit) + default: + return m, nil + } +} + func (m Model) updateReviewAction(action keymap.Action) (tea.Model, tea.Cmd) { switch action { case keymap.ActionQuit: @@ -414,8 +445,14 @@ func (m Model) updateReviewAction(action keymap.Action) (tea.Model, tea.Cmd) { case keymap.ActionExpandMoreContext: m.showMoreContext(contextStep) case keymap.ActionCycleProvider: + if !m.canSwitchProvider() { + return m, nil + } return m, m.cycleProviderCmd() case keymap.ActionOpenProviderPicker: + if !m.canSwitchProvider() { + return m, nil + } m = m.openProviderPicker() case keymap.ActionRefreshProvider: return m, m.refreshActiveProviderCmd(true) @@ -457,6 +494,7 @@ func (m Model) View() tea.View { Mode: diffModeLabel(m.diffMode, m.nerdFont), FileCount: len(m.files), ProviderCount: m.statusProviderCount(), + ProviderSwitch: m.canSwitchProvider(), CurrentFile: m.activeLocation(), Message: m.copyFeedback, ScrollPercent: m.reviewViewport.ScrollPercent(), diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go index 647a8c2..f407f57 100644 --- a/internal/adapters/in/tui/pr_sheet_test.go +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -77,6 +77,21 @@ func TestPRSheetCanToggleByMethodAndMessage(t *testing.T) { assert.False(t, model.prSheet.open) } +func TestPRSheetKeyboardAndWheelScrollSheetNotReview(t *testing.T) { + model := NewModel([]core.ReviewFile{reviewFileWithLines("demo.go", 30)}) + updated, _ := model.Update(tea.WindowSizeMsg{Width: 80, Height: 4}) + model = updated.(Model).TogglePRSheet() + reviewOffset := model.reviewViewport.YOffset() + + updated, _ = model.Update(keyPress("j")) + model = updated.(Model) + updated, _ = model.Update(tea.MouseWheelMsg(tea.Mouse{Button: tea.MouseWheelDown})) + model = updated.(Model) + + assert.Equal(t, 2, model.prSheet.yOffset) + assert.Equal(t, reviewOffset, model.reviewViewport.YOffset()) +} + func TestPRSheetHasIndependentScrollState(t *testing.T) { t.Parallel() diff --git a/internal/adapters/in/tui/review_providers.go b/internal/adapters/in/tui/review_providers.go index 3721c31..646f8ee 100644 --- a/internal/adapters/in/tui/review_providers.go +++ b/internal/adapters/in/tui/review_providers.go @@ -43,6 +43,7 @@ func (m Model) startActiveProviderCmd() tea.Cmd { return func() tea.Msg { catalog, catalogErr := activeProvider.Catalog(ctx) state, err := activeProvider.Start(ctx, reviewContext) + // Provider switching needs the catalog, so surface catalog failure even when start succeeds. if err == nil { err = catalogErr } @@ -105,6 +106,10 @@ func (m Model) statusProviderCount() int { return len(m.providerInfos) } +func (m Model) canSwitchProvider() bool { + return m.activeProvider != nil && len(m.providerCatalog) > 0 +} + func (m *Model) applyActiveProviderState(state ActiveProviderState) { m.activeProviderKey = state.StableProviderKey m.activeRuntimeID = state.RuntimeProviderID diff --git a/internal/adapters/out/plugin/provider_loader.go b/internal/adapters/out/plugin/provider_loader.go index e0df002..f8a87b8 100644 --- a/internal/adapters/out/plugin/provider_loader.go +++ b/internal/adapters/out/plugin/provider_loader.go @@ -130,11 +130,6 @@ func canonicalInstalledPluginIdentity(descriptor ports.PluginDescriptor) string return strings.Join([]string{"manifest", descriptor.Name, descriptor.Version, descriptor.Source}, ":") } -func runtimeCommandAvailable(command, pluginDir string) bool { - _, ok := runtimeCommandInfo(command, pluginDir) - return ok -} - func runtimeCommandInfo(command, pluginDir string) (os.FileInfo, bool) { if command == "" { return nil, false diff --git a/internal/adapters/out/providercache/cache.go b/internal/adapters/out/providercache/cache.go index c19baea..03e263e 100644 --- a/internal/adapters/out/providercache/cache.go +++ b/internal/adapters/out/providercache/cache.go @@ -97,7 +97,7 @@ func writeJSONAtomic(path string, v any) error { return err } tmpName := tmp.Name() - defer os.Remove(tmpName) + defer func() { _ = os.Remove(tmpName) }() if _, err := tmp.Write(data); err != nil { _ = tmp.Close() return err diff --git a/internal/app/active_provider_service.go b/internal/app/active_provider_service.go index e7051e7..1b9538c 100644 --- a/internal/app/active_provider_service.go +++ b/internal/app/active_provider_service.go @@ -29,10 +29,6 @@ type ActiveProviderState struct { NextSyncAt time.Time } -type remoteSnapshotLoader interface { - LoadRemoteSnapshot(ctx context.Context, review core.ReviewContext) (core.ProviderSnapshot, error) -} - type ActiveProviderService struct { catalog ports.ReviewProviderCatalog factory ports.ReviewProviderClientFactory @@ -69,6 +65,7 @@ func (s *ActiveProviderService) Start(ctx context.Context, review core.ReviewCon log.Warn().Err(err).Msg("active provider catalog load failed") return ActiveProviderState{}, err } + preferredKey, hasPreference := s.preferredProviderKey(ctx, review) ordered := s.orderCandidates(ctx, descs, review) log.Info().Int("descriptor_count", len(descs)).Int("candidate_count", len(ordered)).Str("repo", core.RepositoryIdentity(review.Repository)).Msg("active provider start") s.mu.Lock() @@ -101,17 +98,26 @@ func (s *ActiveProviderService) Start(ctx context.Context, review core.ReviewCon gen := s.generation s.backoff = 0 s.mu.Unlock() - if s.prefs != nil { - _ = s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key) + if s.prefs != nil && (!hasPreference || preferredKey == d.Key) { + if err := s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key); err != nil { + log.Warn().Err(err).Str("provider_key", d.Key).Msg("active provider preference save failed") + } } st := s.loadCachedState(ctx, review, d.Key, info.ID, info) log.Info().Str("provider_key", d.Key).Str("runtime_provider_id", info.ID).Bool("from_cache", st.FromCache).Int("remote_thread_count", len(st.Snapshot.Threads)).Msg("active provider selected") - s.setState(gen, st) + if !s.setState(gen, st) { + return s.State(), nil + } return st, nil } + if lastErr == nil { + lastErr = core.NewProviderError(core.ProviderErrorNotApplicable, "no applicable provider found", nil) + } log.Warn().Err(lastErr).Msg("active provider start found no applicable provider") failed := failedProviderState(lastErr) - s.setState(startGen, failed) + if !s.setState(startGen, failed) { + return s.State(), nil + } return failed, lastErr } @@ -139,7 +145,9 @@ func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewCo if err != nil { log.Warn().Err(err).Str("provider_key", stableKey).Msg("active provider switch probe failed") failed := failedProviderState(err) - s.setState(switchGen, failed) + if !s.setState(switchGen, failed) { + return s.State(), nil + } return failed, err } s.mu.Lock() @@ -154,11 +162,15 @@ func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewCo s.backoff = 0 s.mu.Unlock() if s.prefs != nil { - _ = s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key) + if err := s.prefs.SaveActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository), d.Key); err != nil { + log.Warn().Err(err).Str("provider_key", d.Key).Msg("active provider preference save failed") + } } st := s.loadCachedState(ctx, review, d.Key, info.ID, info) log.Info().Str("provider_key", d.Key).Str("runtime_provider_id", info.ID).Bool("from_cache", st.FromCache).Int("remote_thread_count", len(st.Snapshot.Threads)).Msg("active provider switch complete") - s.setState(gen, st) + if !s.setState(gen, st) { + return s.State(), nil + } return st, nil } } @@ -200,7 +212,7 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC log.Info().Str("provider_key", key).Str("runtime_provider_id", runtimeID).Bool("manual", manual).Msg("active provider refresh start") var snap core.ProviderSnapshot var err error - if loader, ok := client.(remoteSnapshotLoader); ok { + if loader, ok := client.(ports.ReviewProviderSnapshotClient); ok { log.Debug().Str("provider_key", key).Msg("active provider loading remote snapshot") snap, err = loader.LoadRemoteSnapshot(ctx, review) } else { @@ -210,6 +222,9 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC snap.Threads = threads } if err != nil { + if st, stale := s.stateIfGenerationChanged(gen); stale { + return st, nil + } st := prev st.LastError = err st.Syncing = false @@ -224,11 +239,16 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC st.Snapshot.Sync.NextSyncAt = new(st.NextSyncAt) log.Warn().Err(err).Str("provider_key", key).Str("runtime_provider_id", runtimeID).Bool("manual", manual).Time("next_sync_at", st.NextSyncAt).Msg("active provider refresh backing off") } - s.setState(gen, st) + if !s.setState(gen, st) { + return s.State(), nil + } return st, err } now := time.Now().UTC() next := now.Add(s.poll.Interval) + if st, stale := s.stateIfGenerationChanged(gen); stale { + return st, nil + } if snap.RuntimeProviderID == "" { snap.RuntimeProviderID = runtimeID } @@ -239,11 +259,15 @@ func (s *ActiveProviderService) Refresh(ctx context.Context, review core.ReviewC } snap.Sync = core.ProviderSyncState{Status: core.ProviderSyncStatusSynced, LastSyncAt: new(now), NextSyncAt: new(next)} if s.cache != nil { - _ = s.cache.SaveProviderSnapshot(ctx, snap) + if err := s.cache.SaveProviderSnapshot(ctx, snap); err != nil { + log.Warn().Err(err).Str("provider_key", key).Msg("active provider cache save failed") + } } log.Info().Str("provider_key", key).Str("runtime_provider_id", snap.RuntimeProviderID).Bool("manual", manual).Int("remote_thread_count", len(snap.Threads)).Bool("has_overview", snap.Overview != nil).Time("next_sync_at", next).Msg("active provider refresh synced") st := ActiveProviderState{StableProviderKey: key, RuntimeProviderID: runtimeID, RuntimeInfo: prev.RuntimeInfo, Snapshot: snap, NextSyncAt: next} - s.setState(gen, st) + if !s.setState(gen, st) { + return s.State(), nil + } return st, nil } @@ -269,7 +293,22 @@ func (s *ActiveProviderService) State() ActiveProviderState { func (s *ActiveProviderService) Close() error { s.mu.Lock() defer s.mu.Unlock() - return s.closeLocked() + err := s.closeLocked() + s.clearActiveLocked() + return err +} + +func (s *ActiveProviderService) preferredProviderKey(ctx context.Context, review core.ReviewContext) (string, bool) { + if s.prefs == nil { + return "", false + } + key, ok, err := s.prefs.LoadActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository)) + if err != nil { + log := zerowrap.FromCtx(ctx) + log.Warn().Err(err).Msg("active provider preference load failed") + return "", false + } + return key, ok } func (s *ActiveProviderService) orderCandidates(ctx context.Context, descs []ports.ReviewProviderDescriptor, review core.ReviewContext) []ports.ReviewProviderDescriptor { @@ -277,7 +316,7 @@ func (s *ActiveProviderService) orderCandidates(ctx context.Context, descs []por used := map[string]bool{} preferenceFound := false if s.prefs != nil { - if key, ok, _ := s.prefs.LoadActiveProviderKey(ctx, core.RepositoryIdentity(review.Repository)); ok { + if key, ok := s.preferredProviderKey(ctx, review); ok { for _, d := range descs { if d.Key == key { out = append(out, d) @@ -350,7 +389,9 @@ func (s *ActiveProviderService) loadCachedState(ctx context.Context, review core } if s.cache != nil { contextKey := core.NewReviewContextKey(key, review) - if snap, ok, _ := s.cache.LoadProviderSnapshot(ctx, contextKey); ok { + if snap, ok, err := s.cache.LoadProviderSnapshot(ctx, contextKey); err != nil { + log.Warn().Err(err).Str("provider_key", key).Any("context_key", contextKey).Msg("active provider cache load failed") + } else if ok { snap.Cached = true st.Snapshot = snap st.FromCache = true @@ -364,12 +405,23 @@ func (s *ActiveProviderService) loadCachedState(ctx context.Context, review core } return st } -func (s *ActiveProviderService) setState(gen int64, st ActiveProviderState) { +func (s *ActiveProviderService) stateIfGenerationChanged(gen int64) (ActiveProviderState, bool) { s.mu.Lock() defer s.mu.Unlock() if gen == s.generation { - s.state = st + return ActiveProviderState{}, false + } + return s.state, true +} + +func (s *ActiveProviderService) setState(gen int64, st ActiveProviderState) bool { + s.mu.Lock() + defer s.mu.Unlock() + if gen != s.generation { + return false } + s.state = st + return true } func (s *ActiveProviderService) nextBackoff(err error) time.Time { s.mu.Lock() @@ -396,6 +448,14 @@ func (s *ActiveProviderService) closeLocked() error { return err } +func (s *ActiveProviderService) clearActiveLocked() { + s.stableKey = "" + s.runtimeID = "" + s.generation++ + s.backoff = 0 + s.state = ActiveProviderState{} +} + func failedProviderState(err error) ActiveProviderState { st := ActiveProviderState{LastError: err} if err != nil { diff --git a/internal/app/active_provider_service_test.go b/internal/app/active_provider_service_test.go index 9f08f35..70b78b6 100644 --- a/internal/app/active_provider_service_test.go +++ b/internal/app/active_provider_service_test.go @@ -3,6 +3,7 @@ package app import ( "context" "errors" + "sync" "testing" "time" @@ -79,6 +80,84 @@ func TestActiveProviderServicePreferenceFallbackAndClosesFailedClients(t *testin } } +func TestActiveProviderServiceStartWithoutCandidatesReturnsNotApplicable(t *testing.T) { + svc := NewActiveProviderService(mockCatalog(t), mocks.NewMockReviewProviderClientFactory(t), nil, nil, ProviderPollingConfig{}) + + st, err := svc.Start(context.Background(), testReviewContext()) + + if got := core.ClassifyProviderError(err); got != core.ProviderErrorNotApplicable { + t.Fatalf("expected not applicable error, got %q state %#v error %v", got, st, err) + } +} + +func TestActiveProviderServiceCloseInvalidatesInFlightRefresh(t *testing.T) { + review := testReviewContext() + started := make(chan struct{}) + release := make(chan struct{}) + provider := mocks.NewMockReviewProviderClient(t) + expectProbe(provider, "github", true) + provider.EXPECT().LoadRemoteThreads(mock.Anything, mock.Anything).Run(func(context.Context, core.ReviewContext) { + close(started) + <-release + }).Return([]core.RemoteReviewThread{{ExternalID: "stale"}}, nil).Once() + provider.EXPECT().Close().Return(nil).Once() + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "github", provider) + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "github", Type: "github"}), factory, nil, nil, ProviderPollingConfig{}) + if _, err := svc.Start(context.Background(), review); err != nil { + t.Fatal(err) + } + + var wg sync.WaitGroup + var refreshState ActiveProviderState + var refreshErr error + wg.Add(1) + go func() { + defer wg.Done() + refreshState, refreshErr = svc.Refresh(context.Background(), review, false) + }() + <-started + if err := svc.Close(); err != nil { + t.Fatal(err) + } + close(release) + wg.Wait() + + if refreshErr != nil { + t.Fatalf("stale refresh should be discarded without surfacing old error: %v", refreshErr) + } + if refreshState.StableProviderKey != "" || len(refreshState.Snapshot.Threads) != 0 { + t.Fatalf("stale refresh returned closed provider state: %#v", refreshState) + } + if state := svc.State(); state.StableProviderKey != "" || len(state.Snapshot.Threads) != 0 { + t.Fatalf("close should clear service state, got %#v", state) + } +} + +func TestActiveProviderServiceAutomaticFallbackDoesNotOverwriteStoredPreference(t *testing.T) { + bad := mocks.NewMockReviewProviderClient(t) + expectProbe(bad, "bad", false) + bad.EXPECT().Close().Return(nil).Once() + good := mocks.NewMockReviewProviderClient(t) + expectProbe(good, "good", true) + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "preferred", bad) + expectFactoryClient(factory, "github", good) + prefs := &memPrefs{key: "preferred", ok: true} + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "preferred"}, ports.ReviewProviderDescriptor{Key: "github", Type: "github"}), factory, nil, prefs, ProviderPollingConfig{}) + + st, err := svc.Start(context.Background(), testReviewContext()) + if err != nil { + t.Fatal(err) + } + if st.StableProviderKey != "github" { + t.Fatalf("got %q", st.StableProviderKey) + } + if prefs.key != "preferred" { + t.Fatalf("automatic fallback overwrote explicit preference with %q", prefs.key) + } +} + func TestActiveProviderServiceCacheFirstRefreshPreservesCacheOnRetryableFailure(t *testing.T) { review := testReviewContext() key := core.NewReviewContextKey("github", review) diff --git a/internal/app/review_providers.go b/internal/app/review_providers.go index 2871aa7..3d0d173 100644 --- a/internal/app/review_providers.go +++ b/internal/app/review_providers.go @@ -2,6 +2,7 @@ package app import ( "context" + "fmt" "github.com/bnema/zerowrap" @@ -18,13 +19,18 @@ func buildReviewProviders(ctx context.Context, catalog ports.ReviewProviderCatal } log := zerowrap.FromCtx(ctx) providers := make([]ports.ReviewProviderClient, 0, len(descriptors)) + failed := 0 for _, descriptor := range descriptors { provider, err := factory.CreateReviewProviderClient(ctx, descriptor) if err != nil { + failed++ log.Warn().Err(err).Str("provider_key", descriptor.Key).Str("contribution_id", descriptor.ContributionID).Msg("create plugin review provider client failed") continue } providers = append(providers, provider) } + if len(descriptors) > 0 && len(providers) == 0 { + return nil, fmt.Errorf("create review providers: all %d configured provider(s) failed to load", failed) + } return providers, nil } diff --git a/internal/app/tui_active_provider.go b/internal/app/tui_active_provider.go index 0eda535..34978ab 100644 --- a/internal/app/tui_active_provider.go +++ b/internal/app/tui_active_provider.go @@ -15,7 +15,7 @@ type tuiActiveProviderController struct { func (c *tuiActiveProviderController) Catalog(ctx context.Context) ([]ports.ReviewProviderDescriptor, error) { if c == nil || c.catalog == nil { - return nil, nil + return nil, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider catalog", nil) } return c.catalog.ListReviewProviderDescriptors(ctx) } diff --git a/internal/ports/plugin.go b/internal/ports/plugin.go index 836f1dd..e837d77 100644 --- a/internal/ports/plugin.go +++ b/internal/ports/plugin.go @@ -76,8 +76,6 @@ type ReviewProviderClientFactory interface { // ReviewProviderLoader builds provider clients from installed plugin sources. type ReviewProviderLoader interface { - ReviewProviderCatalog - ReviewProviderClientFactory LoadReviewProviders(ctx context.Context) ([]ReviewProviderClient, error) } @@ -106,3 +104,9 @@ type ReviewProviderClient interface { PublishReview(ctx context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) Close() error } + +// ReviewProviderSnapshotClient is implemented by providers that can load a +// normalized PR/provider snapshot in one call instead of remote threads only. +type ReviewProviderSnapshotClient interface { + LoadRemoteSnapshot(ctx context.Context, review core.ReviewContext) (core.ProviderSnapshot, error) +} diff --git a/plugins/github/cmd/ero-plugin-github/graphql.go b/plugins/github/cmd/ero-plugin-github/graphql.go index b329db0..7d7a23a 100644 --- a/plugins/github/cmd/ero-plugin-github/graphql.go +++ b/plugins/github/cmd/ero-plugin-github/graphql.go @@ -194,16 +194,37 @@ func (p githubProvider) graphQLClient() (graphQLDoer, error) { return defaultGraphQLClient() } -func fetchGitHubSnapshot(ctx context.Context, client graphQLDoer, remote githubRemote, reviewCtx plugin.ReviewContext) (githubPRSnapshot, error) { - candidates, err := fetchGitHubPRCandidates(ctx, client, remote) - if err != nil { - return githubPRSnapshot{}, err +func matchGitHubPRAcrossRemotes(ctx context.Context, client graphQLDoer, remotes []githubRemote, reviewCtx plugin.ReviewContext) (githubRemote, githubPRCandidate, error) { + matches := make([]struct { + remote githubRemote + pr githubPRCandidate + }, 0, 1) + var lastErr error + for _, remote := range remotes { + candidates, err := fetchGitHubPRCandidates(ctx, client, remote) + if err != nil { + return githubRemote{}, githubPRCandidate{}, err + } + match, err := matchGitHubPR(reviewCtx, candidates) + if err != nil { + lastErr = err + continue + } + matches = append(matches, struct { + remote githubRemote + pr githubPRCandidate + }{remote: remote, pr: match}) + } + if len(matches) == 0 { + if lastErr != nil { + return githubRemote{}, githubPRCandidate{}, lastErr + } + return githubRemote{}, githubPRCandidate{}, plugin.NewError(plugin.ErrorNotApplicable, "no matching GitHub pull request found") } - match, err := matchGitHubPR(reviewCtx, candidates) - if err != nil { - return githubPRSnapshot{}, err + if len(matches) > 1 { + return githubRemote{}, githubPRCandidate{}, plugin.NewErrorf(plugin.ErrorNotApplicable, "ambiguous GitHub pull request match across %d remotes", len(matches)) } - return fetchGitHubPRSnapshot(ctx, client, remote, match.Number) + return matches[0].remote, matches[0].pr, nil } func fetchGitHubPRCandidates(ctx context.Context, client graphQLDoer, remote githubRemote) ([]githubPRCandidate, error) { @@ -252,10 +273,11 @@ func fetchGitHubPRSnapshot(ctx context.Context, client graphQLDoer, remote githu } if !threadsDone { for _, thread := range pr.ReviewThreads.Nodes { + mapped := mapGitHubThread(thread) if thread.Comments.PageInfo.HasNextPage { - return githubPRSnapshot{}, plugin.NewError(plugin.ErrorRemoteValidationFailed, "GitHub review thread comments pagination beyond first page is not supported") + mapped.Unmapped = true } - accum.Threads = append(accum.Threads, mapGitHubThread(thread)) + accum.Threads = append(accum.Threads, mapped) } } if !commentsDone && pr.Comments.PageInfo.HasNextPage { diff --git a/plugins/github/cmd/ero-plugin-github/graphql_test.go b/plugins/github/cmd/ero-plugin-github/graphql_test.go index 5fc1351..093f19e 100644 --- a/plugins/github/cmd/ero-plugin-github/graphql_test.go +++ b/plugins/github/cmd/ero-plugin-github/graphql_test.go @@ -163,7 +163,7 @@ func TestLoadRemoteSnapshotDoesNotRefetchCompletedCollections(t *testing.T) { } } -func TestLoadRemoteSnapshotRejectsUnpaginatedThreadComments(t *testing.T) { +func TestLoadRemoteSnapshotKeepsPartialThreadWhenNestedCommentsArePaginated(t *testing.T) { list := ghPRListResponse{} list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 1, BaseRefName: "main", HeadRefName: "feature"}} page := ghPRSnapshotResponse{} @@ -173,9 +173,12 @@ func TestLoadRemoteSnapshotRejectsUnpaginatedThreadComments(t *testing.T) { page.Repository.PullRequest.ReviewThreads.Nodes = []ghReviewThread{thread} fake := &fakeGraphQLClient{listPages: []ghPRListResponse{list}, snapshotPages: []ghPRSnapshotResponse{page}} provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return fake, nil }} - _, err := provider.LoadRemoteSnapshot(context.Background(), plugin.LoadRemoteSnapshotRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "https://github.com/owner/repo"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}}) - if plugin.AsError(err) == nil || plugin.AsError(err).Code != plugin.ErrorRemoteValidationFailed { - t.Fatalf("expected remote_validation_failed for nested comment pagination, got %v", err) + got, err := provider.LoadRemoteSnapshot(context.Background(), plugin.LoadRemoteSnapshotRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "https://github.com/owner/repo"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}}) + if err != nil { + t.Fatalf("LoadRemoteSnapshot returned error: %v", err) + } + if len(got.Threads) != 1 || !got.Threads[0].Unmapped || len(got.Threads[0].Comments) != 1 { + t.Fatalf("expected partial thread first page marked unmapped, got %#v", got.Threads) } } diff --git a/plugins/github/cmd/ero-plugin-github/main.go b/plugins/github/cmd/ero-plugin-github/main.go index 55ea6e6..681fcc6 100644 --- a/plugins/github/cmd/ero-plugin-github/main.go +++ b/plugins/github/cmd/ero-plugin-github/main.go @@ -69,19 +69,15 @@ func (p githubProvider) Initialize(_ context.Context, req plugin.InitializeReque } func (p githubProvider) DetectContext(ctx context.Context, req plugin.DetectContextRequest) (plugin.DetectContextResult, error) { - remote, ok := firstGitHubRemote(req.Context.Repository.Remotes) - if !ok { + remotes := githubRemotes(req.Context.Repository.Remotes) + if len(remotes) == 0 { return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: "no GitHub remote detected"}}, nil } client, err := p.graphQLClient() if err != nil { return plugin.DetectContextResult{}, plugin.NewErrorf(plugin.ErrorAuthRequired, "create GitHub GraphQL client: %v", err) } - candidates, err := fetchGitHubPRCandidates(ctx, client, remote) - if err != nil { - return plugin.DetectContextResult{}, err - } - match, err := matchGitHubPR(req.Context, candidates) + _, match, err := matchGitHubPRAcrossRemotes(ctx, client, remotes, req.Context) if err != nil { return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: err.Error()}}, nil } @@ -89,15 +85,19 @@ func (p githubProvider) DetectContext(ctx context.Context, req plugin.DetectCont } func (p githubProvider) LoadRemoteSnapshot(ctx context.Context, req plugin.LoadRemoteSnapshotRequest) (plugin.LoadRemoteSnapshotResult, error) { - remote, ok := firstGitHubRemote(req.Context.Repository.Remotes) - if !ok { + remotes := githubRemotes(req.Context.Repository.Remotes) + if len(remotes) == 0 { return plugin.LoadRemoteSnapshotResult{}, plugin.NewError(plugin.ErrorNotApplicable, "no GitHub remote detected") } client, err := p.graphQLClient() if err != nil { return plugin.LoadRemoteSnapshotResult{}, plugin.NewErrorf(plugin.ErrorAuthRequired, "create GitHub GraphQL client: %v", err) } - snapshot, err := fetchGitHubSnapshot(ctx, client, remote, req.Context) + remote, match, err := matchGitHubPRAcrossRemotes(ctx, client, remotes, req.Context) + if err != nil { + return plugin.LoadRemoteSnapshotResult{}, err + } + snapshot, err := fetchGitHubPRSnapshot(ctx, client, remote, match.Number) if err != nil { return plugin.LoadRemoteSnapshotResult{}, err } @@ -152,17 +152,15 @@ func (p githubProvider) PublishReview(ctx context.Context, req plugin.PublishRev } func (p githubProvider) currentPullRequest(ctx context.Context, reviewCtx plugin.ReviewContext) (ghPR, error) { - if remote, ok := firstGitHubRemote(reviewCtx.Repository.Remotes); ok { + if remotes := githubRemotes(reviewCtx.Repository.Remotes); len(remotes) > 0 { if client, err := p.graphQLClient(); err == nil { - candidates, err := fetchGitHubPRCandidates(ctx, client, remote) - if err != nil { - return ghPR{}, err + _, match, err := matchGitHubPRAcrossRemotes(ctx, client, remotes, reviewCtx) + if err == nil { + return ghPR{Number: match.Number, URL: match.URL}, nil } - match, err := matchGitHubPR(reviewCtx, candidates) - if err != nil { + if pe := plugin.AsError(err); pe == nil || pe.Code != plugin.ErrorNotApplicable { return ghPR{}, err } - return ghPR{Number: match.Number, URL: match.URL}, nil } } if p.execGH == nil { diff --git a/plugins/github/cmd/ero-plugin-github/main_test.go b/plugins/github/cmd/ero-plugin-github/main_test.go index 7f82079..0d9e130 100644 --- a/plugins/github/cmd/ero-plugin-github/main_test.go +++ b/plugins/github/cmd/ero-plugin-github/main_test.go @@ -107,6 +107,12 @@ func TestGitHubPRMatching(t *testing.T) { prs: []githubPRCandidate{{Number: 4, BaseRef: "main", HeadSHA: "abc123"}}, wantNumber: 4, }, + { + name: "HEAD pseudo ref falls through to exact head SHA", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "range", HeadRef: "HEAD", HeadSHA: "def456"}}, + prs: []githubPRCandidate{{Number: 9, BaseRef: "main", HeadRef: "feature", HeadSHA: "def456"}}, + wantNumber: 9, + }, { name: "ambiguous multiple matches returns not applicable", ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}, @@ -180,6 +186,49 @@ func TestPublishReviewRejectsMalformedGitHubReviewResponse(t *testing.T) { } } +func TestPublishReviewFallsBackToGHCLIWhenGraphQLHasNoMatch(t *testing.T) { + list := ghPRListResponse{} + list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 88, URL: "https://github.com/owner/repo/pull/88", BaseRefName: "main", HeadRefName: "other"}} + var calls [][]string + provider := githubProvider{ + newGraphQLClient: func() (graphQLDoer, error) { return &fakeGraphQLClient{listPages: []ghPRListResponse{list}}, nil }, + execGH: func(_ context.Context, args ...string) (string, string, error) { + calls = append(calls, slices.Clone(args)) + if len(calls) == 1 { + return `{"number": 12, "url": "https://github.com/owner/repo/pull/12"}`, "", nil + } + return `{"id": 99, "html_url": "https://github.com/owner/repo/pull/12#pullrequestreview-99"}`, "", nil + }, + } + + _, err := provider.PublishReview(context.Background(), plugin.PublishReviewParams{Payload: plugin.ReviewPublishPayload{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "git@github.com:owner/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}, Draft: plugin.ReviewDraftSnapshot{Summary: "summary"}}}) + if err != nil { + t.Fatalf("PublishReview returned error: %v", err) + } + if len(calls) != 2 || !strings.Contains(strings.Join(calls[0], "\x00"), "pr\x00view") { + t.Fatalf("expected gh pr view fallback then publish, got %#v", calls) + } +} + +func TestLoadRemoteSnapshotSearchesAllGitHubRemotes(t *testing.T) { + forkList := ghPRListResponse{} + forkList.Repository.PullRequests.Nodes = []ghPRNode{{Number: 1, BaseRefName: "main", HeadRefName: "other"}} + upstreamList := ghPRListResponse{} + upstreamList.Repository.PullRequests.Nodes = []ghPRNode{{Number: 2, BaseRefName: "main", HeadRefName: "feature"}} + snapshot := ghPRSnapshotResponse{} + snapshot.Repository.PullRequest = ghPRNode{Number: 2, URL: "https://github.com/upstream/repo/pull/2", Title: "PR", BaseRefName: "main", HeadRefName: "feature"} + fake := &fakeGraphQLClient{listPages: []ghPRListResponse{forkList, upstreamList}, snapshotPages: []ghPRSnapshotResponse{snapshot}} + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return fake, nil }} + + got, err := provider.LoadRemoteSnapshot(context.Background(), plugin.LoadRemoteSnapshotRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{Name: "origin", URL: "git@github.com:fork/repo.git"}, {Name: "upstream", URL: "git@github.com:upstream/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}}) + if err != nil { + t.Fatalf("LoadRemoteSnapshot returned error: %v", err) + } + if fake.listCalls != 2 || got.Overview == nil || got.Overview.Number != 2 { + t.Fatalf("expected upstream PR match, calls=%d snapshot=%#v", fake.listCalls, got.Overview) + } +} + func TestPublishReviewUsesGraphQLMatchedPullRequestWhenContextHasRemote(t *testing.T) { list := ghPRListResponse{} list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 77, URL: "https://github.com/owner/repo/pull/77", BaseRefName: "release", HeadRefName: "topic"}} diff --git a/plugins/github/cmd/ero-plugin-github/match.go b/plugins/github/cmd/ero-plugin-github/match.go index 32f5f20..66ef1ef 100644 --- a/plugins/github/cmd/ero-plugin-github/match.go +++ b/plugins/github/cmd/ero-plugin-github/match.go @@ -46,6 +46,9 @@ func githubPRMatches(ctx plugin.ReviewContext, pr githubPRCandidate) bool { headSHA := firstNonEmpty(ctx.Target.HeadSHA, ctx.Repository.HeadSHA) headRef := strings.TrimSpace(ctx.Target.HeadRef) + if nonBranchHeadRef(headRef) { + headRef = "" + } if headRef == "" && strings.EqualFold(ctx.Target.Mode, "branch") { headRef = strings.TrimSpace(ctx.Repository.CurrentBranch) } @@ -62,6 +65,26 @@ func githubPRMatches(ctx plugin.ReviewContext, pr githubPRCandidate) bool { return headSHA != "" && pr.HeadSHA != "" && strings.EqualFold(pr.HeadSHA, headSHA) } +func nonBranchHeadRef(ref string) bool { + ref = strings.TrimSpace(ref) + if ref == "" { + return false + } + if strings.EqualFold(ref, "HEAD") || strings.Contains(ref, "..") { + return true + } + trimmed := strings.TrimPrefix(strings.ToLower(ref), "refs/heads/") + if len(trimmed) < 7 || len(trimmed) > 40 { + return false + } + for _, r := range trimmed { + if (r < '0' || r > '9') && (r < 'a' || r > 'f') { + return false + } + } + return true +} + func refEqual(a, b string) bool { return normalizeRef(a) == normalizeRef(b) } diff --git a/plugins/github/cmd/ero-plugin-github/remote.go b/plugins/github/cmd/ero-plugin-github/remote.go index cea3fff..62bd70e 100644 --- a/plugins/github/cmd/ero-plugin-github/remote.go +++ b/plugins/github/cmd/ero-plugin-github/remote.go @@ -50,11 +50,20 @@ func cleanGitHubRepo(owner, repo string) (githubRemote, bool) { return githubRemote{Owner: owner, Name: repo}, true } -func firstGitHubRemote(remotes []plugin.GitRemote) (githubRemote, bool) { +func githubRemotes(remotes []plugin.GitRemote) []githubRemote { + out := make([]githubRemote, 0, len(remotes)) + seen := map[string]bool{} for _, remote := range remotes { - if parsed, ok := parseGitHubRemote(remote.URL); ok { - return parsed, true + parsed, ok := parseGitHubRemote(remote.URL) + if !ok { + continue } + key := strings.ToLower(parsed.Owner + "/" + parsed.Name) + if seen[key] { + continue + } + seen[key] = true + out = append(out, parsed) } - return githubRemote{}, false + return out } From 3d1e8cf7b63b6d713f471a13b6cc75ce76330cc5 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 08:33:09 +0200 Subject: [PATCH 15/22] fix(tui): polish provider status and PR sheet --- .../adapters/in/tui/component/statusbar.go | 88 ++++++++++++++----- .../in/tui/component/statusbar_test.go | 10 ++- internal/adapters/in/tui/pr_sheet.go | 4 +- internal/adapters/in/tui/pr_sheet_test.go | 10 +++ 4 files changed, 85 insertions(+), 27 deletions(-) diff --git a/internal/adapters/in/tui/component/statusbar.go b/internal/adapters/in/tui/component/statusbar.go index a1cc7ca..a07310e 100644 --- a/internal/adapters/in/tui/component/statusbar.go +++ b/internal/adapters/in/tui/component/statusbar.go @@ -2,6 +2,7 @@ package component import ( "fmt" + "image/color" "strings" "time" @@ -50,8 +51,8 @@ func (c StatusBar) Render(model StatusModel) string { if model.ProviderCount > 0 && (strings.TrimSpace(model.ActiveProviderLabel) == "" || !model.NerdFont) { segments = append(segments, statusSegment{style: theme.StatusInfoStyle, label: providerCountLabel(model.ProviderCount)}) } - if syncLabel := providerSyncLabel(model); syncLabel != "" { - segments = append(segments, statusSegment{style: theme.StatusInfoStyle, label: syncLabel}) + if syncSegment := providerSyncSegment(model); syncSegment.label != "" || syncSegment.rendered != "" { + segments = append(segments, syncSegment) } if model.DraftCommentCount > 0 { segments = append(segments, statusSegment{style: theme.StatusInfoStyle, label: draftCommentCountLabel(model.DraftCommentCount)}) @@ -75,8 +76,9 @@ func (c StatusBar) Render(model StatusModel) string { } type statusSegment struct { - style lipgloss.Style - label string + style lipgloss.Style + label string + rendered string } type KeyHint struct { @@ -114,6 +116,10 @@ func renderStatusSegments(width int, segments ...statusSegment) string { if remaining <= 0 { break } + if segment.rendered != "" { + rendered.WriteString(ansi.Truncate(segment.rendered, remaining, "…")) + continue + } padding := segment.style.GetHorizontalPadding() labelWidth := remaining - padding if labelWidth <= 0 { @@ -146,6 +152,22 @@ func providerCountLabel(count int) string { return fmt.Sprintf("%d providers", count) } +const ( + nerdFontGitHubLarge = "\uf113" // nf-fa-github_alt + nerdFontSyncDot = "\u25cf" +) + +func providerSyncSegment(model StatusModel) statusSegment { + if model.NerdFont && strings.TrimSpace(model.ActiveProviderLabel) != "" { + return statusSegment{rendered: renderNerdFontProviderSync(model)} + } + label := providerSyncLabel(model) + if label == "" { + return statusSegment{} + } + return statusSegment{style: theme.StatusInfoStyle, label: label} +} + func providerSyncLabel(model StatusModel) string { provider := strings.TrimSpace(model.ActiveProviderLabel) if provider == "" { @@ -159,11 +181,8 @@ func providerSyncLabel(model StatusModel) string { } parts := []string{provider} - if model.NerdFont { - parts = []string{compactProviderLabel(provider, model.ProviderCount, model.ProviderSync.Status)} - } status := providerSyncStatusLabel(model.ProviderSync.Status) - if status != "" && !model.NerdFont { + if status != "" { parts = append(parts, status) } if model.ProviderSync.LastError != "" { @@ -185,35 +204,62 @@ func draftCommentCountLabel(count int) string { return fmt.Sprintf("%d draft comments", count) } -func compactProviderLabel(provider string, providerCount int, status core.ProviderSyncStatus) string { - label := providerGlyph(provider) + providerStatusDot(status) - if providerCount > 1 { - label += fmt.Sprintf(" +%d", providerCount-1) +func renderNerdFontProviderSync(model StatusModel) string { + provider := strings.TrimSpace(model.ActiveProviderLabel) + if runtimeName := strings.TrimSpace(model.ActiveRuntimeName); runtimeName != "" && runtimeName != provider { + provider += "/" + runtimeName + } + var b strings.Builder + b.WriteString(theme.StatusBaseStyle.Render(" ")) + b.WriteString(theme.StatusBaseStyle.Foreground(lipgloss.Color("248")).Render(providerGlyph(provider))) + b.WriteString(theme.StatusBaseStyle.Render(" ")) + b.WriteString(theme.StatusBaseStyle.Foreground(providerStatusDotColor(model.ProviderSync.Status)).Render(nerdFontSyncDot)) + for _, part := range nerdFontProviderSyncTextParts(model) { + b.WriteString(theme.StatusBaseStyle.Render(" ")) + b.WriteString(theme.StatusBaseStyle.Foreground(lipgloss.Color("248")).Render(part)) + } + b.WriteString(theme.StatusBaseStyle.Render(" ")) + return b.String() +} + +func nerdFontProviderSyncTextParts(model StatusModel) []string { + parts := []string{} + if model.ProviderCount > 1 { + parts = append(parts, fmt.Sprintf("+%d", model.ProviderCount-1)) } - return label + if model.ProviderSync.LastError != "" { + parts = append(parts, TruncateRunes(model.ProviderSync.LastError, 24)) + } + if model.ProviderSync.LastSyncAt != nil { + parts = append(parts, "last "+formatStatusTime(*model.ProviderSync.LastSyncAt)) + } + if model.ProviderSync.NextSyncAt != nil { + parts = append(parts, "next "+formatStatusTime(*model.ProviderSync.NextSyncAt)) + } + return parts } func providerGlyph(provider string) string { provider = strings.ToLower(strings.TrimSpace(provider)) if strings.Contains(provider, "github") { - return "" + return nerdFontGitHubLarge } return providerAbbreviation(provider) } -func providerStatusDot(status core.ProviderSyncStatus) string { - color := lipgloss.Color("81") +func providerStatusDotColor(status core.ProviderSyncStatus) color.Color { switch status { case core.ProviderSyncStatusSynced: - color = lipgloss.Color("#3fb950") + return lipgloss.Color("#3fb950") case core.ProviderSyncStatusFailed: - color = lipgloss.Color("#ff7b72") + return lipgloss.Color("#ff7b72") case core.ProviderSyncStatusBackingOff: - color = lipgloss.Color("#ffa657") + return lipgloss.Color("#ffa657") case core.ProviderSyncStatusLoadingCache, core.ProviderSyncStatusSyncing: - color = lipgloss.Color("#58a6ff") + return lipgloss.Color("#58a6ff") + default: + return lipgloss.Color("81") } - return theme.StatusBaseStyle.Foreground(color).Render("●") } func providerAbbreviation(provider string) string { diff --git a/internal/adapters/in/tui/component/statusbar_test.go b/internal/adapters/in/tui/component/statusbar_test.go index b16843e..c54e303 100644 --- a/internal/adapters/in/tui/component/statusbar_test.go +++ b/internal/adapters/in/tui/component/statusbar_test.go @@ -80,7 +80,7 @@ func TestStatusbarProviderSyncUsesNerdFontProviderGlyphAndStatusDot(t *testing.T view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) - require.Contains(t, view, "●") + require.Contains(t, view, nerdFontGitHubLarge+" "+nerdFontSyncDot) } } @@ -90,7 +90,7 @@ func TestStatusbarProviderSyncOmitsStatusWordWhenNerdFontSymbolIsShown(t *testin view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) - require.Contains(t, view, "●") + require.Contains(t, view, nerdFontGitHubLarge+" "+nerdFontSyncDot) require.NotContains(t, view, "") require.NotContains(t, view, "synced") } @@ -101,10 +101,12 @@ func TestStatusbarShowsCompactActiveProviderAndAdditionalCount(t *testing.T) { model.ProviderSwitch = true model.NerdFont = true - view := stripANSIForStatusbarTest(NewStatusBar(120).Render(model)) + raw := NewStatusBar(120).Render(model) + view := stripANSIForStatusbarTest(raw) - require.Contains(t, view, "● +1") + require.Contains(t, view, nerdFontGitHubLarge+" "+nerdFontSyncDot+" +1") require.NotContains(t, view, "2 providers") + require.Contains(t, raw, "48;5;236") require.Contains(t, view, "p provider") require.Contains(t, view, "P publish") } diff --git a/internal/adapters/in/tui/pr_sheet.go b/internal/adapters/in/tui/pr_sheet.go index b70eb5a..b3b2db1 100644 --- a/internal/adapters/in/tui/pr_sheet.go +++ b/internal/adapters/in/tui/pr_sheet.go @@ -34,8 +34,8 @@ func (m Model) ScrollPRSheet(delta int) Model { func (m Model) renderPRSheetOverlay(content string) string { width := max(m.width, 1) height := max(m.height, 1) + paneWidth := prSheetWidth(width) pane := m.renderPRSheet(width, height) - paneWidth := lipgloss.Width(pane) canvas := lipgloss.NewCanvas(width, height) compositor := lipgloss.NewCompositor( @@ -221,7 +221,7 @@ func prSheetWidth(totalWidth int) int { if totalWidth <= 1 { return 1 } - return min(max(totalWidth/3, 32), totalWidth) + return max(totalWidth/2, 1) } func visiblePRSheetLines(lines []string, offset, height int) []string { diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go index f407f57..839926c 100644 --- a/internal/adapters/in/tui/pr_sheet_test.go +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -24,8 +24,12 @@ func TestPRSheetOverlaysRightSideFullHeightWithLeftSeparatorOnly(t *testing.T) { require.Len(t, lines, 8) sheetWidth := prSheetWidth(model.width) + assert.Equal(t, model.width/2, sheetWidth) separatorColumn := model.width - sheetWidth separatorRows := 0 + firstLineRunes := []rune(lines[0]) + require.Greater(t, len(firstLineRunes), separatorColumn) + assert.Equal(t, '│', firstLineRunes[separatorColumn]) for _, line := range lines { runes := []rune(line) if len(runes) <= separatorColumn { @@ -46,6 +50,12 @@ func TestPRSheetOverlaysRightSideFullHeightWithLeftSeparatorOnly(t *testing.T) { assert.NotContains(t, view, "┘") } +func TestPRSheetWidthIsHalfOfView(t *testing.T) { + assert.Equal(t, 40, prSheetWidth(80)) + assert.Equal(t, 30, prSheetWidth(60)) + assert.Equal(t, 1, prSheetWidth(1)) +} + func TestPRSheetOverlayDoesNotReflowUnderlyingDiffContent(t *testing.T) { t.Parallel() From cd667000901f5fee23ab567cd1defdfdeb02b0c8 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 11:21:14 +0200 Subject: [PATCH 16/22] fix(tui): align PR sheet overlay rows --- internal/adapters/in/tui/pr_sheet.go | 35 ++++++++++++++++++----- internal/adapters/in/tui/pr_sheet_test.go | 12 ++++++++ 2 files changed, 40 insertions(+), 7 deletions(-) diff --git a/internal/adapters/in/tui/pr_sheet.go b/internal/adapters/in/tui/pr_sheet.go index b3b2db1..cbed8d3 100644 --- a/internal/adapters/in/tui/pr_sheet.go +++ b/internal/adapters/in/tui/pr_sheet.go @@ -6,6 +6,7 @@ import ( "time" "charm.land/lipgloss/v2" + "github.com/charmbracelet/x/ansi" "ero/internal/core" ) @@ -36,14 +37,27 @@ func (m Model) renderPRSheetOverlay(content string) string { height := max(m.height, 1) paneWidth := prSheetWidth(width) pane := m.renderPRSheet(width, height) + return composeRightOverlay(content, pane, width, height, paneWidth) +} - canvas := lipgloss.NewCanvas(width, height) - compositor := lipgloss.NewCompositor( - lipgloss.NewLayer(content), - lipgloss.NewLayer(pane).X(max(width-paneWidth, 0)).Y(0).Z(1), - ) - canvas.Compose(compositor) - return canvas.Render() +func composeRightOverlay(content, overlay string, width, height, overlayWidth int) string { + leftWidth := max(width-overlayWidth, 0) + contentLines := strings.Split(content, "\n") + overlayLines := strings.Split(overlay, "\n") + rows := make([]string, height) + for i := range height { + left := "" + if i < len(contentLines) { + left = ansi.Truncate(contentLines[i], leftWidth, "") + } + left = padRightANSI(left, leftWidth) + right := strings.Repeat(" ", overlayWidth) + if i < len(overlayLines) { + right = overlayLines[i] + } + rows[i] = left + right + } + return strings.Join(rows, "\n") } func (m Model) renderPRSheet(width, height int) string { @@ -250,3 +264,10 @@ func padRight(s string, width int) string { } return s + strings.Repeat(" ", width-lipgloss.Width(s)) } + +func padRightANSI(s string, width int) string { + if ansi.StringWidth(s) >= width { + return s + } + return s + strings.Repeat(" ", width-ansi.StringWidth(s)) +} diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go index 839926c..1a156cb 100644 --- a/internal/adapters/in/tui/pr_sheet_test.go +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -50,6 +50,18 @@ func TestPRSheetOverlaysRightSideFullHeightWithLeftSeparatorOnly(t *testing.T) { assert.NotContains(t, view, "┘") } +func TestComposeRightOverlayAlignsFirstRowWithRemainingRows(t *testing.T) { + content := "abcdef\nghijkl\nmnopqr" + overlay := "│ first\n│ second\n│ third" + plain := stripANSI(composeRightOverlay(content, overlay, 12, 3, 7)) + + for i, line := range strings.Split(plain, "\n") { + runes := []rune(line) + require.Greater(t, len(runes), 5) + assert.Equal(t, '│', runes[5], "line %d separator column", i) + } +} + func TestPRSheetWidthIsHalfOfView(t *testing.T) { assert.Equal(t, 40, prSheetWidth(80)) assert.Equal(t, 30, prSheetWidth(60)) From dc9c0616ca63c0ecbf125d0fb7799e534fc34c09 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 12:54:27 +0200 Subject: [PATCH 17/22] fix(tui): hide unanchored remote threads inline --- internal/adapters/in/tui/pr_sheet.go | 35 +++--------- internal/adapters/in/tui/pr_sheet_test.go | 12 ---- .../adapters/in/tui/presenter/document.go | 10 ---- .../in/tui/presenter/document_test.go | 56 +++++++++++++------ .../adapters/in/tui/review_remote_test.go | 6 +- 5 files changed, 50 insertions(+), 69 deletions(-) diff --git a/internal/adapters/in/tui/pr_sheet.go b/internal/adapters/in/tui/pr_sheet.go index cbed8d3..b3b2db1 100644 --- a/internal/adapters/in/tui/pr_sheet.go +++ b/internal/adapters/in/tui/pr_sheet.go @@ -6,7 +6,6 @@ import ( "time" "charm.land/lipgloss/v2" - "github.com/charmbracelet/x/ansi" "ero/internal/core" ) @@ -37,27 +36,14 @@ func (m Model) renderPRSheetOverlay(content string) string { height := max(m.height, 1) paneWidth := prSheetWidth(width) pane := m.renderPRSheet(width, height) - return composeRightOverlay(content, pane, width, height, paneWidth) -} -func composeRightOverlay(content, overlay string, width, height, overlayWidth int) string { - leftWidth := max(width-overlayWidth, 0) - contentLines := strings.Split(content, "\n") - overlayLines := strings.Split(overlay, "\n") - rows := make([]string, height) - for i := range height { - left := "" - if i < len(contentLines) { - left = ansi.Truncate(contentLines[i], leftWidth, "") - } - left = padRightANSI(left, leftWidth) - right := strings.Repeat(" ", overlayWidth) - if i < len(overlayLines) { - right = overlayLines[i] - } - rows[i] = left + right - } - return strings.Join(rows, "\n") + canvas := lipgloss.NewCanvas(width, height) + compositor := lipgloss.NewCompositor( + lipgloss.NewLayer(content), + lipgloss.NewLayer(pane).X(max(width-paneWidth, 0)).Y(0).Z(1), + ) + canvas.Compose(compositor) + return canvas.Render() } func (m Model) renderPRSheet(width, height int) string { @@ -264,10 +250,3 @@ func padRight(s string, width int) string { } return s + strings.Repeat(" ", width-lipgloss.Width(s)) } - -func padRightANSI(s string, width int) string { - if ansi.StringWidth(s) >= width { - return s - } - return s + strings.Repeat(" ", width-ansi.StringWidth(s)) -} diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go index 1a156cb..839926c 100644 --- a/internal/adapters/in/tui/pr_sheet_test.go +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -50,18 +50,6 @@ func TestPRSheetOverlaysRightSideFullHeightWithLeftSeparatorOnly(t *testing.T) { assert.NotContains(t, view, "┘") } -func TestComposeRightOverlayAlignsFirstRowWithRemainingRows(t *testing.T) { - content := "abcdef\nghijkl\nmnopqr" - overlay := "│ first\n│ second\n│ third" - plain := stripANSI(composeRightOverlay(content, overlay, 12, 3, 7)) - - for i, line := range strings.Split(plain, "\n") { - runes := []rune(line) - require.Greater(t, len(runes), 5) - assert.Equal(t, '│', runes[5], "line %d separator column", i) - } -} - func TestPRSheetWidthIsHalfOfView(t *testing.T) { assert.Equal(t, 40, prSheetWidth(80)) assert.Equal(t, 30, prSheetWidth(60)) diff --git a/internal/adapters/in/tui/presenter/document.go b/internal/adapters/in/tui/presenter/document.go index 4160417..d3a8ee2 100644 --- a/internal/adapters/in/tui/presenter/document.go +++ b/internal/adapters/in/tui/presenter/document.go @@ -43,7 +43,6 @@ type reviewDocumentBuilder struct { } func (b *reviewDocumentBuilder) build() ReviewDocument { - b.appendUnmappedRemoteThreads() if len(b.input.Files) == 0 { b.rows = append(b.rows, ReviewRow{Kind: ReviewRowKindMessage, Message: "Review"}, @@ -58,15 +57,6 @@ func (b *reviewDocumentBuilder) build() ReviewDocument { return b.document() } -func (b *reviewDocumentBuilder) appendUnmappedRemoteThreads() { - for _, thread := range b.input.Annotations.RemoteThreads { - if !thread.Unmapped && thread.FilePath != "" { - continue - } - b.rows = append(b.rows, ReviewRow{Kind: ReviewRowKindRemoteThread, FileIndex: -1, SectionIndex: -1, LineIndex: -1, Annotation: ReviewAnnotation{RemoteThread: thread}}) - } -} - func (b *reviewDocumentBuilder) appendFile(fileIndex int, file core.ReviewFile) { if fileIndex > 0 { b.rows = append(b.rows, ReviewRow{Kind: ReviewRowKindBlank, FileIndex: fileIndex, FilePath: file.Path}) diff --git a/internal/adapters/in/tui/presenter/document_test.go b/internal/adapters/in/tui/presenter/document_test.go index 715e8bc..363d343 100644 --- a/internal/adapters/in/tui/presenter/document_test.go +++ b/internal/adapters/in/tui/presenter/document_test.go @@ -91,7 +91,6 @@ func TestBuildReviewDocumentProjectsAnnotationRowsAndRebuildsAnchors(t *testing. }) assert.Equal(t, []ReviewRowKind{ - ReviewRowKindRemoteThread, ReviewRowKindFile, ReviewRowKindRule, ReviewRowKindLine, @@ -104,22 +103,45 @@ func TestBuildReviewDocumentProjectsAnnotationRowsAndRebuildsAnchors(t *testing. ReviewRowKindEditor, ReviewRowKindEditor, }, rowKinds(doc.Rows)) - assert.Equal(t, 1, doc.Anchors.FileRows[0]) - assert.Equal(t, 3, doc.Anchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 0}]) - assert.Equal(t, 6, doc.Anchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 1}]) - assert.Equal(t, 9, doc.Anchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 2}]) - assert.Equal(t, "comment-1", doc.Rows[4].Annotation.Comment.ID) - assert.Equal(t, 0, doc.Rows[4].Annotation.LineIndex) - assert.Equal(t, "local note", doc.Rows[5].Annotation.Body) - assert.Equal(t, "github", doc.Rows[0].Annotation.RemoteThread.ProviderID) - assert.True(t, doc.Rows[0].Annotation.RemoteThread.Unmapped) - assert.Equal(t, "github", doc.Rows[7].Annotation.RemoteThread.ProviderID) - assert.Equal(t, "octocat", doc.Rows[8].Annotation.Author) - assert.Equal(t, "remote note", doc.Rows[8].Annotation.Body) - assert.Equal(t, "demo.go", doc.Rows[10].Annotation.Editor.FilePath) - assert.Equal(t, 0, doc.Rows[10].Annotation.LineIndex) - assert.Equal(t, 1, doc.Rows[11].Annotation.LineIndex) - assert.False(t, doc.Rows[4].Selectable) + assert.Equal(t, 0, doc.Anchors.FileRows[0]) + assert.Equal(t, 2, doc.Anchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 0}]) + assert.Equal(t, 5, doc.Anchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 1}]) + assert.Equal(t, 8, doc.Anchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 2}]) + assert.Equal(t, "comment-1", doc.Rows[3].Annotation.Comment.ID) + assert.Equal(t, 0, doc.Rows[3].Annotation.LineIndex) + assert.Equal(t, "local note", doc.Rows[4].Annotation.Body) + assert.Equal(t, "github", doc.Rows[6].Annotation.RemoteThread.ProviderID) + assert.Equal(t, "octocat", doc.Rows[7].Annotation.Author) + assert.Equal(t, "remote note", doc.Rows[7].Annotation.Body) + assert.Equal(t, "demo.go", doc.Rows[9].Annotation.Editor.FilePath) + assert.Equal(t, 0, doc.Rows[9].Annotation.LineIndex) + assert.Equal(t, 1, doc.Rows[10].Annotation.LineIndex) + assert.False(t, doc.Rows[3].Selectable) +} + +func TestBuildReviewDocumentDoesNotProjectUnmappedOrFilelessRemoteThreads(t *testing.T) { + t.Parallel() + + doc := BuildReviewDocument(ReviewDocumentInput{ + Files: []core.ReviewFile{{ + Path: "demo.go", + Sections: []core.ReviewSection{{Kind: core.SectionKindChanged, Lines: []core.ReviewLine{{NewLineNumber: 1, Content: "one", Kind: core.LineKindAdded}}}}, + }}, + Annotations: ReviewAnnotations{RemoteThreads: []core.RemoteReviewThread{ + {ProviderID: "github", ExternalID: "unmapped", Unmapped: true, Comments: []core.RemoteReviewComment{{Body: "orphaned"}}}, + {ProviderID: "github", ExternalID: "fileless", Comments: []core.RemoteReviewComment{{Body: "no file"}}}, + {ProviderID: "github", ExternalID: "mapped", FilePath: "demo.go", Range: core.ReviewLineRange{Start: core.ReviewLineRef{NewLineNumber: 1, Kind: core.LineKindAdded}, End: core.ReviewLineRef{NewLineNumber: 1, Kind: core.LineKindAdded}}, Comments: []core.RemoteReviewComment{{Author: "octocat", Body: "mapped note"}}}, + }}, + }) + + require.Equal(t, ReviewRowKindFile, doc.Rows[0].Kind) + var remoteIDs []string + for _, row := range doc.Rows { + if row.Kind == ReviewRowKindRemoteThread { + remoteIDs = append(remoteIDs, row.Annotation.RemoteThread.ExternalID) + } + } + require.Equal(t, []string{"mapped", "mapped"}, remoteIDs) } func TestBuildReviewDocumentAnnotationRowsMatchRenderedLineCountsBeforeLaterAnchors(t *testing.T) { diff --git a/internal/adapters/in/tui/review_remote_test.go b/internal/adapters/in/tui/review_remote_test.go index 7a77b9f..8bdd5d2 100644 --- a/internal/adapters/in/tui/review_remote_test.go +++ b/internal/adapters/in/tui/review_remote_test.go @@ -24,10 +24,12 @@ func TestRemoteReviewAnnotations(t *testing.T) { require.Contains(t, view, "octocat: remote note") } -func TestRemoteReviewAnnotationsUnmapped(t *testing.T) { +func TestRemoteReviewAnnotationsUnmappedAreHiddenFromDiff(t *testing.T) { rendered := renderReviewForTest([]core.ReviewFile{reviewFile("demo.go", "package main")}, 80, -1, -1, presenter.ReviewAnnotations{ RemoteThreads: []core.RemoteReviewThread{{ProviderID: "github", Unmapped: true, Comments: []core.RemoteReviewComment{{Body: "orphaned"}}}}, }) view := stripANSI(rendered.Content) - require.Contains(t, view, "[github] unmapped: orphaned") + require.NotContains(t, view, "[github] unmapped") + require.NotContains(t, view, "orphaned") + require.Contains(t, view, "demo.go") } From 2e53167ad0f119d7cd34cb5d17c11c968bcbb294 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 13:05:10 +0200 Subject: [PATCH 18/22] feat(tui): color PR sheet markdown --- internal/adapters/in/tui/markdown_renderer.go | 55 ++++++++++++++++++- .../adapters/in/tui/markdown_renderer_test.go | 38 ++++++++++++- internal/adapters/in/tui/pr_sheet.go | 31 ++++++++--- internal/adapters/in/tui/pr_sheet_test.go | 35 ++++++++++++ internal/adapters/in/tui/theme/styles.go | 14 +++++ 5 files changed, 159 insertions(+), 14 deletions(-) diff --git a/internal/adapters/in/tui/markdown_renderer.go b/internal/adapters/in/tui/markdown_renderer.go index ed7c0d3..7d66c2e 100644 --- a/internal/adapters/in/tui/markdown_renderer.go +++ b/internal/adapters/in/tui/markdown_renderer.go @@ -7,6 +7,9 @@ import ( "strings" "charm.land/glamour/v2" + glamouransi "charm.land/glamour/v2/ansi" + + "ero/internal/adapters/in/tui/theme" ) type MarkdownTheme string @@ -84,6 +87,7 @@ func (r *MarkdownRenderer) Render(markdown string, width int, theme MarkdownThem if err != nil { return safeMarkdownFallback(markdown) } + rendered = sanitizeRenderedMarkdown(rendered) r.entries[key] = rendered return rendered } @@ -103,20 +107,65 @@ func (r *MarkdownRenderer) renderer(width int, theme MarkdownTheme) (markdownTer func newGlamourTermRenderer(width int, theme MarkdownTheme) (markdownTermRenderer, error) { return glamour.NewTermRenderer( - glamour.WithStandardStyle(string(theme)), + glamour.WithStyles(eroMarkdownStyle(theme)), glamour.WithWordWrap(width), ) } +func eroMarkdownStyle(markdownTheme MarkdownTheme) glamouransi.StyleConfig { + text := theme.ColorStatusInfo + muted := theme.ColorMutedText + heading := theme.ColorAccent + section := theme.ColorWarning + codeBackground := theme.ColorCodeBg + if markdownTheme == MarkdownThemeLight { + text = theme.ColorStatusBase + muted = "244" + codeBackground = "#f6f8fa" + } + bold := true + italic := true + underline := true + zeroIndent := uint(0) + quoteIndent := uint(1) + quoteToken := "│ " + return glamouransi.StyleConfig{ + Document: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: &text}}, + Text: glamouransi.StylePrimitive{Color: &text}, + Paragraph: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: &text}}, + Heading: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: &heading, Bold: &bold}}, + H1: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: &heading, Bold: &bold}}, + H2: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: §ion, Bold: &bold}}, + H3: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: §ion, Bold: &bold}}, + Strong: glamouransi.StylePrimitive{Color: §ion, Bold: &bold}, + Emph: glamouransi.StylePrimitive{Color: &text, Italic: &italic}, + Item: glamouransi.StylePrimitive{Color: &text}, + Enumeration: glamouransi.StylePrimitive{Color: &muted}, + Link: glamouransi.StylePrimitive{Color: &heading, Underline: &underline}, + LinkText: glamouransi.StylePrimitive{Color: &heading, Underline: &underline}, + Code: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: §ion, BackgroundColor: &codeBackground}}, + CodeBlock: glamouransi.StyleCodeBlock{StyleBlock: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: &text, BackgroundColor: &codeBackground}, Margin: &zeroIndent}, Theme: "github-dark"}, + BlockQuote: glamouransi.StyleBlock{StylePrimitive: glamouransi.StylePrimitive{Color: &muted}, Indent: "eIndent, IndentToken: "eToken}, + } +} + func hashMarkdownInput(input string) string { sum := sha256.Sum256([]byte(input)) return hex.EncodeToString(sum[:]) } -var ansiEscapePattern = regexp.MustCompile(`\x1b\[[0-9;?]*[ -/]*[@-~]`) +var ( + ansiEscapePattern = regexp.MustCompile(`\x1b\[[0-9;?]*[ -/]*[@-~]`) + oscEscapePattern = regexp.MustCompile(`\x1b\][^\x1b\x07]*(?:\x07|\x1b\\)`) +) + +func sanitizeRenderedMarkdown(input string) string { + return oscEscapePattern.ReplaceAllString(input, "") +} func safeMarkdownFallback(input string) string { - withoutEscapes := ansiEscapePattern.ReplaceAllString(input, "") + withoutEscapes := sanitizeRenderedMarkdown(input) + withoutEscapes = ansiEscapePattern.ReplaceAllString(withoutEscapes, "") withoutEscapes = strings.ReplaceAll(withoutEscapes, "\x1b", "") return strings.TrimSpace(withoutEscapes) } diff --git a/internal/adapters/in/tui/markdown_renderer_test.go b/internal/adapters/in/tui/markdown_renderer_test.go index bd91522..4f1c156 100644 --- a/internal/adapters/in/tui/markdown_renderer_test.go +++ b/internal/adapters/in/tui/markdown_renderer_test.go @@ -43,6 +43,23 @@ func TestMarkdownRendererCachesByInputWidthAndTheme(t *testing.T) { } } +func TestMarkdownRendererColorsHeadingsAndFencedCodeBlocks(t *testing.T) { + renderer := NewMarkdownRenderer() + + got := renderer.Render("## Validation\n\n```go\nfmt.Println(\"hi\")\n```", 80, MarkdownThemeDark) + plain := regexp.MustCompile(`\x1b\[[0-9;?]*[ -/]*[@-~]`).ReplaceAllString(got, "") + + if !strings.Contains(got, "\x1b[") { + t.Fatalf("expected rendered markdown to include ANSI color/style escapes, got %q", got) + } + if !strings.Contains(plain, "Validation") || !strings.Contains(plain, "fmt.Println") { + t.Fatalf("expected heading and code block content, got %q", got) + } + if strings.Contains(plain, "```") { + t.Fatalf("expected rendered fenced code block to omit markdown fences, got %q", got) + } +} + func TestMarkdownRendererRendersFencedCodeBlocks(t *testing.T) { renderer := NewMarkdownRenderer() @@ -57,6 +74,23 @@ func TestMarkdownRendererRendersFencedCodeBlocks(t *testing.T) { } } +func TestMarkdownRendererStripsOSC8Hyperlinks(t *testing.T) { + renderer := NewMarkdownRendererWithFactory(func(width int, theme MarkdownTheme) (markdownTermRenderer, error) { + return fakeMarkdownTermRenderer{render: func(markdown string) (string, error) { + return "\x1b]8;id=abc;https://github.com/bnema/ero\x1b\\GitHub\x1b]8;;\x1b\\", nil + }}, nil + }) + + got := renderer.Render("[GitHub](https://github.com/bnema/ero)", 80, MarkdownThemeDark) + + if strings.Contains(got, "]8;") || strings.Contains(got, "id=abc") || strings.Contains(got, "\x1b]") { + t.Fatalf("expected OSC-8 hyperlink sequences to be stripped, got %q", got) + } + if !strings.Contains(got, "GitHub") { + t.Fatalf("expected link text to remain, got %q", got) + } +} + func TestMarkdownRendererReturnsSafeFallbackOnRenderError(t *testing.T) { renderer := NewMarkdownRendererWithFactory(func(width int, theme MarkdownTheme) (markdownTermRenderer, error) { return fakeMarkdownTermRenderer{render: func(markdown string) (string, error) { @@ -64,9 +98,9 @@ func TestMarkdownRendererReturnsSafeFallbackOnRenderError(t *testing.T) { }}, nil }) - got := renderer.Render("hello\x1b[31m **world**", 80, MarkdownThemeDark) + got := renderer.Render("hello\x1b[31m \x1b]8;id=x;https://example.invalid\x1b\\**world**\x1b]8;;\x1b\\", 80, MarkdownThemeDark) - if strings.Contains(got, "\x1b") { + if strings.Contains(got, "\x1b") || strings.Contains(got, "]8;") { t.Fatalf("expected fallback to strip escape characters, got %q", got) } if !strings.Contains(got, "hello") || !strings.Contains(got, "world") { diff --git a/internal/adapters/in/tui/pr_sheet.go b/internal/adapters/in/tui/pr_sheet.go index b3b2db1..060b63f 100644 --- a/internal/adapters/in/tui/pr_sheet.go +++ b/internal/adapters/in/tui/pr_sheet.go @@ -6,6 +6,7 @@ import ( "time" "charm.land/lipgloss/v2" + "github.com/charmbracelet/x/ansi" "ero/internal/core" ) @@ -59,8 +60,8 @@ func (m Model) renderPRSheet(width, height int) string { if i < len(visibleLines) { text = visibleLines[i] } - row := "│ " + truncatePlainRow(text, contentWidth) - rows[i] = padRight(row, sheetWidth) + row := "│ " + ansi.Truncate(text, contentWidth, "") + rows[i] = padRightANSI(row, sheetWidth) } return strings.Join(rows, "\n") } @@ -134,12 +135,24 @@ func (m Model) prSheetLineCount() int { } func renderPRSheetMarkdown(renderer *MarkdownRenderer, markdown string, width int) []string { - rendered := renderer.Render(markdown, width, MarkdownThemeDark) - plain := safeMarkdownFallback(rendered) - if strings.TrimSpace(plain) == "" { + rendered := sanitizeRenderedMarkdown(renderer.Render(markdown, width, MarkdownThemeDark)) + if strings.TrimSpace(safeMarkdownFallback(rendered)) == "" { return []string{"(empty)"} } - return strings.Split(plain, "\n") + return trimRenderedMarkdownBlankLines(strings.Split(rendered, "\n")) +} + +func trimRenderedMarkdownBlankLines(lines []string) []string { + for len(lines) > 0 && strings.TrimSpace(safeMarkdownFallback(lines[0])) == "" { + lines = lines[1:] + } + for len(lines) > 0 && strings.TrimSpace(safeMarkdownFallback(lines[len(lines)-1])) == "" { + lines = lines[:len(lines)-1] + } + if len(lines) == 0 { + return []string{"(empty)"} + } + return lines } func providerOverviewMetadata(overview *core.ProviderOverview) []string { @@ -244,9 +257,9 @@ func pluralCount(count int, singular string) string { return strconv.Itoa(count) + " " + pluralize(singular, count) } -func padRight(s string, width int) string { - if lipgloss.Width(s) >= width { +func padRightANSI(s string, width int) string { + if ansi.StringWidth(s) >= width { return s } - return s + strings.Repeat(" ", width-lipgloss.Width(s)) + return s + strings.Repeat(" ", width-ansi.StringWidth(s)) } diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go index 839926c..d353690 100644 --- a/internal/adapters/in/tui/pr_sheet_test.go +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -167,6 +167,41 @@ func TestPRSheetRendersOverviewMarkdownCommentsAndReviews(t *testing.T) { assert.Contains(t, plain, "Looks good to me.") } +func TestPRSheetRendersMarkdownHeadingsAndCodeBlocksWithColor(t *testing.T) { + model := NewModel(nil) + model.width = 120 + model.height = 20 + model.providerOverview = &core.ProviderOverview{ + Title: "PR", + Body: "## Details\n\n```go\nfunc main() {}\n```", + } + + raw := model.renderPRSheet(model.width, model.height) + plain := stripANSI(raw) + + require.Contains(t, plain, "Details") + require.Contains(t, plain, "func main") + require.Contains(t, raw, "\x1b[") +} + +func TestPRSheetDoesNotLeakOSCHyperlinkSequences(t *testing.T) { + model := NewModel(nil) + model.width = 120 + model.height = 20 + model.providerOverview = &core.ProviderOverview{ + Title: "PR", + Body: "See https://github.com/example/repo/pull/1", + } + + raw := model.renderPRSheet(model.width, model.height) + plain := stripANSI(raw) + + require.Contains(t, plain, "github.com/example/repo") + require.NotContains(t, raw, "\x1b]8;") + require.NotContains(t, raw, "]8;id=") + require.NotContains(t, plain, "]8;id=") +} + func TestPRSheetRendersNilOverviewFallback(t *testing.T) { t.Parallel() diff --git a/internal/adapters/in/tui/theme/styles.go b/internal/adapters/in/tui/theme/styles.go index 70a3933..8def904 100644 --- a/internal/adapters/in/tui/theme/styles.go +++ b/internal/adapters/in/tui/theme/styles.go @@ -2,6 +2,20 @@ package theme import "charm.land/lipgloss/v2" +const ( + ColorText = "#c9d1d9" + ColorMutedText = "#8b949e" + ColorAccent = "#58a6ff" + ColorWarning = "#ffa657" + ColorKeyword = "#ff7b72" + ColorFunction = "#d2a8ff" + ColorString = "#a5d6ff" + ColorNumber = "#79c0ff" + ColorStatusBase = "236" + ColorStatusInfo = "248" + ColorCodeBg = "#1f2a44" +) + var ( FileHeaderStyle = lipgloss.NewStyle().Bold(true).Foreground(lipgloss.Color("15")) FileRuleStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("8")) From ac9e235f524f4092d270be22df3d2be0244eadc1 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 13:16:43 +0200 Subject: [PATCH 19/22] fix(tui): keep file navigation in PR sheet --- internal/adapters/in/tui/model.go | 6 +++++ internal/adapters/in/tui/pr_sheet_test.go | 22 +++++++++++++++++++ .../github/cmd/ero-plugin-github/main_test.go | 6 +++++ plugins/github/cmd/ero-plugin-github/match.go | 8 +------ 4 files changed, 35 insertions(+), 7 deletions(-) diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index c59f9bb..57b721b 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -386,6 +386,12 @@ func (m Model) updatePRSheetAction(msg tea.KeyPressMsg) (tea.Model, tea.Cmd) { return m.ScrollPRSheet(-max(m.height-1, 1)), nil case keymap.ActionPageDown: return m.ScrollPRSheet(max(m.height-1, 1)), nil + case keymap.ActionPreviousFile: + m.moveFile(-1) + return m, nil + case keymap.ActionNextFile: + m.moveFile(1) + return m, nil case keymap.ActionTogglePRSheet: return m.TogglePRSheet(), nil case keymap.ActionOpenHelp: diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go index d353690..fcb1424 100644 --- a/internal/adapters/in/tui/pr_sheet_test.go +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -87,6 +87,28 @@ func TestPRSheetCanToggleByMethodAndMessage(t *testing.T) { assert.False(t, model.prSheet.open) } +func TestPRSheetAllowsFileNavigationShortcuts(t *testing.T) { + model := NewModel([]core.ReviewFile{reviewFile("a.go", "package a"), reviewFile("b.go", "package b")}) + updated, _ := model.Update(tea.WindowSizeMsg{Width: 80, Height: 8}) + model = updated.(Model).TogglePRSheet() + + updated, _ = model.Update(keyPress("l")) + model = updated.(Model) + assert.Equal(t, 1, model.selectedFile) + + updated, _ = model.Update(keyPress("h")) + model = updated.(Model) + assert.Equal(t, 0, model.selectedFile) + + updated, _ = model.Update(tea.KeyPressMsg{Code: tea.KeyRight}) + model = updated.(Model) + assert.Equal(t, 1, model.selectedFile) + + updated, _ = model.Update(tea.KeyPressMsg{Code: tea.KeyLeft}) + model = updated.(Model) + assert.Equal(t, 0, model.selectedFile) +} + func TestPRSheetKeyboardAndWheelScrollSheetNotReview(t *testing.T) { model := NewModel([]core.ReviewFile{reviewFileWithLines("demo.go", 30)}) updated, _ := model.Update(tea.WindowSizeMsg{Width: 80, Height: 4}) diff --git a/plugins/github/cmd/ero-plugin-github/main_test.go b/plugins/github/cmd/ero-plugin-github/main_test.go index 0d9e130..8948584 100644 --- a/plugins/github/cmd/ero-plugin-github/main_test.go +++ b/plugins/github/cmd/ero-plugin-github/main_test.go @@ -113,6 +113,12 @@ func TestGitHubPRMatching(t *testing.T) { prs: []githubPRCandidate{{Number: 9, BaseRef: "main", HeadRef: "feature", HeadSHA: "def456"}}, wantNumber: 9, }, + { + name: "matching branch accepts local unpushed head SHA", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch", HeadSHA: "local-unpushed"}}, + prs: []githubPRCandidate{{Number: 10, BaseRef: "main", HeadRef: "feature", HeadSHA: "remote-pr-head"}}, + wantNumber: 10, + }, { name: "ambiguous multiple matches returns not applicable", ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}, diff --git a/plugins/github/cmd/ero-plugin-github/match.go b/plugins/github/cmd/ero-plugin-github/match.go index 66ef1ef..6c18c69 100644 --- a/plugins/github/cmd/ero-plugin-github/match.go +++ b/plugins/github/cmd/ero-plugin-github/match.go @@ -53,13 +53,7 @@ func githubPRMatches(ctx plugin.ReviewContext, pr githubPRCandidate) bool { headRef = strings.TrimSpace(ctx.Repository.CurrentBranch) } if headRef != "" { - if !refEqual(pr.HeadRef, headRef) { - return false - } - if headSHA != "" && pr.HeadSHA != "" && !strings.EqualFold(pr.HeadSHA, headSHA) { - return false - } - return true + return refEqual(pr.HeadRef, headRef) } return headSHA != "" && pr.HeadSHA != "" && strings.EqualFold(pr.HeadSHA, headSHA) From 0d781d5fff38ea0d49f8bbbb608110236c5e1eba Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 13:25:55 +0200 Subject: [PATCH 20/22] fix(providers): skip PR sync outside branch mode --- .../adapters/in/tui/active_provider_test.go | 11 +++++++ internal/adapters/in/tui/review_providers.go | 23 ++++++++++---- internal/app/app_test.go | 2 +- internal/core/startup.go | 4 +-- internal/core/startup_test.go | 12 +++---- plugins/github/cmd/ero-plugin-github/main.go | 11 +++++++ .../github/cmd/ero-plugin-github/main_test.go | 31 +++++++++++++++++++ 7 files changed, 79 insertions(+), 15 deletions(-) diff --git a/internal/adapters/in/tui/active_provider_test.go b/internal/adapters/in/tui/active_provider_test.go index cd0c34d..971ac89 100644 --- a/internal/adapters/in/tui/active_provider_test.go +++ b/internal/adapters/in/tui/active_provider_test.go @@ -45,6 +45,17 @@ func (m *mockActiveProviderController) CompleteTimer(ctx context.Context, review } func (m *mockActiveProviderController) Close() error { return m.Called().Error(0) } +func TestActiveProviderDoesNotStartOutsideBranchMode(t *testing.T) { + controller := &mockActiveProviderController{} + m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{DiffMode: core.DiffModeUpstream}, nil, core.ReviewContext{Target: core.ReviewTargetMetadata{Mode: core.DiffModeUpstream}}, controller, nil) + + cmd := m.Init() + + require.Nil(t, cmd) + controller.AssertNotCalled(t, "Catalog", mock.Anything) + controller.AssertNotCalled(t, "Start", mock.Anything, mock.Anything) +} + func TestActiveProviderStartupLoadsOnlyActiveProviderState(t *testing.T) { controller := &mockActiveProviderController{} catalog := []ports.ReviewProviderDescriptor{{Key: "github", Label: "GitHub"}, {Key: "other", Label: "Other"}} diff --git a/internal/adapters/in/tui/review_providers.go b/internal/adapters/in/tui/review_providers.go index 646f8ee..bcc1c56 100644 --- a/internal/adapters/in/tui/review_providers.go +++ b/internal/adapters/in/tui/review_providers.go @@ -35,7 +35,7 @@ func (m Model) closeReviewProvidersCmd() tea.Cmd { func (m Model) startActiveProviderCmd() tea.Cmd { activeProvider := m.activeProvider - if activeProvider == nil { + if activeProvider == nil || !m.activeProviderSyncEnabled() { return nil } ctx := m.ctx @@ -53,7 +53,7 @@ func (m Model) startActiveProviderCmd() tea.Cmd { func (m Model) refreshActiveProviderCmd(manual bool) tea.Cmd { activeProvider := m.activeProvider - if activeProvider == nil { + if activeProvider == nil || !m.activeProviderSyncEnabled() { return nil } ctx := m.ctx @@ -66,7 +66,7 @@ func (m Model) refreshActiveProviderCmd(manual bool) tea.Cmd { func (m Model) switchActiveProviderCmd(stableKey string) tea.Cmd { activeProvider := m.activeProvider - if activeProvider == nil { + if activeProvider == nil || !m.activeProviderSyncEnabled() { return nil } ctx := m.ctx @@ -78,7 +78,7 @@ func (m Model) switchActiveProviderCmd(stableKey string) tea.Cmd { } func (m Model) scheduleActiveProviderPollCmd() tea.Cmd { - if m.activeProvider == nil || m.providerSyncState.NextSyncAt == nil { + if m.activeProvider == nil || !m.activeProviderSyncEnabled() || m.providerSyncState.NextSyncAt == nil { return nil } delay := max(time.Until(*m.providerSyncState.NextSyncAt), 0) @@ -88,7 +88,7 @@ func (m Model) scheduleActiveProviderPollCmd() tea.Cmd { func (m Model) completeActiveProviderTimerCmd(generation int64) tea.Cmd { activeProvider := m.activeProvider - if activeProvider == nil { + if activeProvider == nil || !m.activeProviderSyncEnabled() { return nil } ctx := m.ctx @@ -107,7 +107,18 @@ func (m Model) statusProviderCount() int { } func (m Model) canSwitchProvider() bool { - return m.activeProvider != nil && len(m.providerCatalog) > 0 + return m.activeProvider != nil && m.activeProviderSyncEnabled() && len(m.providerCatalog) > 0 +} + +func (m Model) activeProviderSyncEnabled() bool { + mode := m.reviewContext.Target.Mode + if mode == "" { + mode = m.request.DiffMode + } + if mode == "" { + mode = core.DiffModeBranch + } + return mode == core.DiffModeBranch } func (m *Model) applyActiveProviderState(state ActiveProviderState) { diff --git a/internal/app/app_test.go b/internal/app/app_test.go index 8d2ec49..ad5a3a5 100644 --- a/internal/app/app_test.go +++ b/internal/app/app_test.go @@ -133,7 +133,7 @@ func TestRunDetectsStartupModeWhenNoExplicitCommand(t *testing.T) { }{ {name: "working changes", state: core.StartupState{HasUnstagedChanges: true}, wantMode: core.DiffModeWorking}, {name: "staged changes", state: core.StartupState{HasStagedChanges: true}, wantMode: core.DiffModeStaged}, - {name: "ahead of upstream", state: core.StartupState{HasUpstream: true, Ahead: 1}, wantMode: core.DiffModeUpstream}, + {name: "ahead of upstream with default branch", state: core.StartupState{HasUpstream: true, Ahead: 1, HasDefaultBranch: true}, wantMode: core.DiffModeBranch}, {name: "mixed prompts and uses selected mode", state: core.StartupState{HasStagedChanges: true, HasUnstagedChanges: true}, promptMode: core.DiffModeLocal, wantMode: core.DiffModeLocal}, {name: "mixed non-interactive errors", state: core.StartupState{HasStagedChanges: true, HasUnstagedChanges: true}, wantErr: "choose explicitly"}, } diff --git a/internal/core/startup.go b/internal/core/startup.go index 6ebdd89..173ac7c 100644 --- a/internal/core/startup.go +++ b/internal/core/startup.go @@ -46,12 +46,12 @@ func ResolveStartupDecision(state StartupState) StartupDecision { return StartupDecision{Kind: StartupDecisionUseMode, DiffMode: DiffModeStaged} case state.DetachedHead: return StartupDecision{Kind: StartupDecisionNoReviewableChanges, Message: "detached HEAD has no safe default diff; choose an explicit diff mode"} + case state.HasDefaultBranch: + return StartupDecision{Kind: StartupDecisionUseMode, DiffMode: DiffModeBranch} case state.HasUpstream && state.Ahead > 0: return StartupDecision{Kind: StartupDecisionUseMode, DiffMode: DiffModeUpstream} case state.HasUpstream && state.Behind > 0: return StartupDecision{Kind: StartupDecisionNoReviewableChanges, Message: "branch is behind upstream; pull first or choose an explicit diff mode"} - case state.HasDefaultBranch: - return StartupDecision{Kind: StartupDecisionUseMode, DiffMode: DiffModeBranch} default: return StartupDecision{Kind: StartupDecisionNoReviewableChanges, Message: "no local changes or upstream/default branch diff detected"} } diff --git a/internal/core/startup_test.go b/internal/core/startup_test.go index 1167a69..56feaff 100644 --- a/internal/core/startup_test.go +++ b/internal/core/startup_test.go @@ -59,14 +59,14 @@ func TestResolveStartupDecisionChoosesBranchStateModes(t *testing.T) { want StartupDecision }{ { - name: "ahead of upstream uses upstream diff", - state: StartupState{HasUpstream: true, Ahead: 2}, - want: StartupDecision{Kind: StartupDecisionUseMode, DiffMode: DiffModeUpstream}, + name: "ahead of upstream uses branch diff when default branch exists", + state: StartupState{HasUpstream: true, Ahead: 2, HasDefaultBranch: true}, + want: StartupDecision{Kind: StartupDecisionUseMode, DiffMode: DiffModeBranch}, }, { - name: "diverged from upstream uses upstream diff from merge base", - state: StartupState{HasUpstream: true, Ahead: 1, Behind: 1}, - want: StartupDecision{Kind: StartupDecisionUseMode, DiffMode: DiffModeUpstream}, + name: "diverged from upstream uses branch diff when default branch exists", + state: StartupState{HasUpstream: true, Ahead: 1, Behind: 1, HasDefaultBranch: true}, + want: StartupDecision{Kind: StartupDecisionUseMode, DiffMode: DiffModeBranch}, }, { name: "behind upstream only asks user to pull instead of reviewing remote-only changes", diff --git a/plugins/github/cmd/ero-plugin-github/main.go b/plugins/github/cmd/ero-plugin-github/main.go index 681fcc6..3da18cc 100644 --- a/plugins/github/cmd/ero-plugin-github/main.go +++ b/plugins/github/cmd/ero-plugin-github/main.go @@ -69,6 +69,9 @@ func (p githubProvider) Initialize(_ context.Context, req plugin.InitializeReque } func (p githubProvider) DetectContext(ctx context.Context, req plugin.DetectContextRequest) (plugin.DetectContextResult, error) { + if !isBranchReviewMode(req.Context) { + return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: "GitHub PR sync is available in branch mode only"}}, nil + } remotes := githubRemotes(req.Context.Repository.Remotes) if len(remotes) == 0 { return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: "no GitHub remote detected"}}, nil @@ -85,6 +88,9 @@ func (p githubProvider) DetectContext(ctx context.Context, req plugin.DetectCont } func (p githubProvider) LoadRemoteSnapshot(ctx context.Context, req plugin.LoadRemoteSnapshotRequest) (plugin.LoadRemoteSnapshotResult, error) { + if !isBranchReviewMode(req.Context) { + return plugin.LoadRemoteSnapshotResult{}, plugin.NewError(plugin.ErrorNotApplicable, "GitHub PR sync is available in branch mode only") + } remotes := githubRemotes(req.Context.Repository.Remotes) if len(remotes) == 0 { return plugin.LoadRemoteSnapshotResult{}, plugin.NewError(plugin.ErrorNotApplicable, "no GitHub remote detected") @@ -196,6 +202,11 @@ func ghPRViewArgs(reviewCtx plugin.ReviewContext) []string { return args } +func isBranchReviewMode(reviewCtx plugin.ReviewContext) bool { + mode := reviewCtx.Target.Mode + return mode == "" || strings.EqualFold(mode, "branch") +} + func publishPRLookupBranch(reviewCtx plugin.ReviewContext) string { headRef := strings.TrimSpace(reviewCtx.Target.HeadRef) if headRef == "" && strings.EqualFold(reviewCtx.Target.Mode, "branch") { diff --git a/plugins/github/cmd/ero-plugin-github/main_test.go b/plugins/github/cmd/ero-plugin-github/main_test.go index 8948584..416f2fc 100644 --- a/plugins/github/cmd/ero-plugin-github/main_test.go +++ b/plugins/github/cmd/ero-plugin-github/main_test.go @@ -43,6 +43,37 @@ func TestGitHubRemoteParsing(t *testing.T) { } } +func TestDetectContextSkipsNonBranchReviewModes(t *testing.T) { + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { + t.Fatal("non-branch detection should not create a GraphQL client") + return nil, nil + }} + review := plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{Name: "origin", URL: "git@github.com:owner/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "upstream", HeadRef: "HEAD"}} + + result, err := provider.DetectContext(context.Background(), plugin.DetectContextRequest{Context: review}) + + if err != nil { + t.Fatalf("DetectContext returned error: %v", err) + } + if result.Result.Applicable || !strings.Contains(result.Result.Reason, "branch mode") { + t.Fatalf("expected non-branch mode to be not applicable without sync failure: %#v", result) + } +} + +func TestLoadRemoteSnapshotSkipsNonBranchReviewModes(t *testing.T) { + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { + t.Fatal("non-branch snapshot should not create a GraphQL client") + return nil, nil + }} + review := plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{Name: "origin", URL: "git@github.com:owner/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "range", BaseRef: "main", HeadRef: "HEAD"}} + + _, err := provider.LoadRemoteSnapshot(context.Background(), plugin.LoadRemoteSnapshotRequest{Context: review}) + + if plugin.AsError(err) == nil || plugin.AsError(err).Code != plugin.ErrorNotApplicable || !strings.Contains(err.Error(), "branch mode") { + t.Fatalf("expected non-branch snapshot to be not applicable, got %v", err) + } +} + func TestDetectContextRequiresMatchingGitHubPullRequest(t *testing.T) { list := ghPRListResponse{} list.Repository.PullRequests.Nodes = []ghPRNode{{Number: 12, URL: "https://github.com/owner/repo/pull/12", BaseRefName: "main", HeadRefName: "feature"}} From 893d1310cef8e3def8368bf681e18e793ab669b4 Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 13:25:44 +0200 Subject: [PATCH 21/22] fix(tui): navigate horizontally by diff chunk --- .../adapters/in/tui/component/help_pane.go | 2 +- .../adapters/in/tui/cursor_navigation_test.go | 59 +++++++++++++++++++ internal/adapters/in/tui/model.go | 8 +-- internal/adapters/in/tui/model_context.go | 11 ---- internal/adapters/in/tui/model_viewport.go | 41 +++++++++++++ internal/adapters/in/tui/pr_sheet_test.go | 24 ++++++-- 6 files changed, 123 insertions(+), 22 deletions(-) diff --git a/internal/adapters/in/tui/component/help_pane.go b/internal/adapters/in/tui/component/help_pane.go index 6a2b6b3..1c3887c 100644 --- a/internal/adapters/in/tui/component/help_pane.go +++ b/internal/adapters/in/tui/component/help_pane.go @@ -28,7 +28,7 @@ func RenderHelpPane(width, height int, enterKeyLabel, commentSubmitKeyLabel stri renderHelpShortcut(enterKeyLabel, "jump to result", contentWidth), renderHelpShortcut("esc", "cancel search", contentWidth), "", - renderHelpShortcut("h/l", "previous/next file", contentWidth), + renderHelpShortcut("h/l", "previous/next chunk", contentWidth), renderHelpShortcut("a", "expand all context", contentWidth), renderHelpShortcut(enterKeyLabel, "expand more context", contentWidth), renderHelpShortcut("s/space", "select lines", contentWidth), diff --git a/internal/adapters/in/tui/cursor_navigation_test.go b/internal/adapters/in/tui/cursor_navigation_test.go index 4b9a41a..5927b2f 100644 --- a/internal/adapters/in/tui/cursor_navigation_test.go +++ b/internal/adapters/in/tui/cursor_navigation_test.go @@ -124,6 +124,65 @@ func TestModelCursorNavigationPreservesAbsoluteAndPageViewportSemantics(t *testi } } +func TestModelHorizontalNavigationMovesBetweenChangedChunks(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + keys []tea.KeyPressMsg + want ReviewLineAnchor + }{ + {name: "l moves to next chunk in next file", keys: []tea.KeyPressMsg{keyPress("l"), keyPress("l")}, want: ReviewLineAnchor{FileIndex: 1, SectionIndex: 1, LineIndex: 0}}, + {name: "right moves to next chunk in next file", keys: []tea.KeyPressMsg{{Code: tea.KeyRight}, {Code: tea.KeyRight}}, want: ReviewLineAnchor{FileIndex: 1, SectionIndex: 1, LineIndex: 0}}, + {name: "h moves to previous chunk in previous file", keys: []tea.KeyPressMsg{keyPress("l"), keyPress("l"), keyPress("h")}, want: ReviewLineAnchor{FileIndex: 0, SectionIndex: 2, LineIndex: 0}}, + {name: "left moves to previous chunk in previous file", keys: []tea.KeyPressMsg{{Code: tea.KeyRight}, {Code: tea.KeyRight}, {Code: tea.KeyLeft}}, want: ReviewLineAnchor{FileIndex: 0, SectionIndex: 2, LineIndex: 0}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + model := NewModel(chunkNavigationReviewFiles()) + updated, _ := model.Update(tea.WindowSizeMsg{Width: 80, Height: 5}) + model = updated.(Model) + model.cursorRow = model.reviewAnchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 0}] + model.updateAfterCursorMove() + selectionAnchorRow := model.cursorRow + model.selectionAnchorRow = &selectionAnchorRow + model.selectedFile = 0 + model.selectedContext = 0 + + for _, key := range tt.keys { + updated, _ = model.Update(key) + model = updated.(Model) + } + + wantCursor := model.reviewAnchors.LineRows[tt.want] + assert.Equal(t, wantCursor, model.cursorRow) + assert.LessOrEqual(t, model.reviewViewport.YOffset(), model.cursorRow) + assert.Less(t, model.cursorRow, model.reviewViewport.YOffset()+model.reviewViewport.Height()) + assert.NotEqual(t, model.reviewAnchors.FileRows[tt.want.FileIndex], model.cursorRow, "horizontal navigation should target changed chunk lines, not file headers") + assert.NotNil(t, model.selectionAnchorRow, "horizontal navigation should not toggle line selection") + assert.Equal(t, model.reviewAnchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 0}], *model.selectionAnchorRow) + assert.Equal(t, -1, model.selectedContext, "horizontal navigation should not highlight context expanders") + }) + } +} + +func chunkNavigationReviewFiles() []core.ReviewFile { + return []core.ReviewFile{ + {Path: "a.go", Sections: []core.ReviewSection{ + {ID: "a-1", Kind: core.SectionKindChanged, Lines: []core.ReviewLine{{NewLineNumber: 1, Content: "a first", Kind: core.LineKindAdded}}}, + {ID: "a-context", Kind: core.SectionKindContext, Lines: []core.ReviewLine{{OldLineNumber: 2, NewLineNumber: 2, Content: "hidden a", Kind: core.LineKindUnchanged}}}, + {ID: "a-2", Kind: core.SectionKindChanged, Lines: []core.ReviewLine{{NewLineNumber: 3, Content: "a second", Kind: core.LineKindAdded}}}, + }}, + {Path: "b.go", Sections: []core.ReviewSection{ + {ID: "b-context", Kind: core.SectionKindContext, Lines: []core.ReviewLine{{OldLineNumber: 1, NewLineNumber: 1, Content: "hidden b", Kind: core.LineKindUnchanged}}}, + {ID: "b-1", Kind: core.SectionKindChanged, Lines: []core.ReviewLine{{NewLineNumber: 2, Content: "b first", Kind: core.LineKindAdded}}}, + }}, + } +} + func repeatKey(key tea.KeyPressMsg, count int) []tea.KeyPressMsg { keys := make([]tea.KeyPressMsg, count) for i := range keys { diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index 57b721b..b5d0fde 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -387,10 +387,10 @@ func (m Model) updatePRSheetAction(msg tea.KeyPressMsg) (tea.Model, tea.Cmd) { case keymap.ActionPageDown: return m.ScrollPRSheet(max(m.height-1, 1)), nil case keymap.ActionPreviousFile: - m.moveFile(-1) + m.moveChunk(-1) return m, nil case keymap.ActionNextFile: - m.moveFile(1) + m.moveChunk(1) return m, nil case keymap.ActionTogglePRSheet: return m.TogglePRSheet(), nil @@ -443,9 +443,9 @@ func (m Model) updateReviewAction(action keymap.Action) (tea.Model, tea.Cmd) { case keymap.ActionOpenDiffMode: return m.openSearch(searchModeDiff) case keymap.ActionPreviousFile: - m.moveFile(-1) + m.moveChunk(-1) case keymap.ActionNextFile: - m.moveFile(1) + m.moveChunk(1) case keymap.ActionExpandAllContext: m.showAllContext() case keymap.ActionExpandMoreContext: diff --git a/internal/adapters/in/tui/model_context.go b/internal/adapters/in/tui/model_context.go index fbe4313..bfbc64e 100644 --- a/internal/adapters/in/tui/model_context.go +++ b/internal/adapters/in/tui/model_context.go @@ -2,17 +2,6 @@ package tui import "ero/internal/core" -func (m *Model) moveFile(delta int) { - if len(m.files) == 0 { - return - } - - m.selectedFile = min(max(m.selectedFile+delta, 0), len(m.files)-1) - m.clearSelection() - m.resetContextSelection() - m.syncReviewViewport() -} - func (m *Model) showMoreContext(count int) { fileIndex, sectionIndex, ok := m.contextSectionLocationForExpansion() if !ok || !m.contextBarActionAllowed(fileIndex, sectionIndex, ContextBarActionShowMore) { diff --git a/internal/adapters/in/tui/model_viewport.go b/internal/adapters/in/tui/model_viewport.go index 7c567aa..66d50b1 100644 --- a/internal/adapters/in/tui/model_viewport.go +++ b/internal/adapters/in/tui/model_viewport.go @@ -120,6 +120,47 @@ func (m *Model) moveCursor(delta int) { m.updateAfterCursorMove() } +func (m *Model) moveChunk(delta int) { + if delta == 0 { + return + } + chunks := m.changedChunkRows() + if len(chunks) == 0 { + return + } + index := sort.Search(len(chunks), func(i int) bool { return chunks[i] >= m.cursorRow }) + if delta > 0 { + if index < len(chunks) && chunks[index] == m.cursorRow { + index++ + } + if index >= len(chunks) { + index = len(chunks) - 1 + } + } else { + if index >= len(chunks) || chunks[index] >= m.cursorRow { + index-- + } + if index < 0 { + index = 0 + } + } + m.cursorRow = chunks[index] + m.selectedContext = -1 + m.keepCursorVisible() + m.updateActiveFileFromCursor() + m.syncReviewVisualState() +} + +func (m Model) changedChunkRows() []int { + rows := make([]int, 0) + for rowIndex, row := range m.reviewRows { + if row.Kind == ReviewRowKindLine && row.LineIndex == 0 && row.FileIndex >= 0 && row.FileIndex < len(m.files) && row.SectionIndex >= 0 && row.SectionIndex < len(m.files[row.FileIndex].Sections) && m.files[row.FileIndex].Sections[row.SectionIndex].Kind == core.SectionKindChanged { + rows = append(rows, rowIndex) + } + } + return rows +} + func (m *Model) moveCursorToStart() { m.cursorRow = m.firstSelectableRow() m.updateAfterCursorMoveWithOffset(0) diff --git a/internal/adapters/in/tui/pr_sheet_test.go b/internal/adapters/in/tui/pr_sheet_test.go index fcb1424..944a289 100644 --- a/internal/adapters/in/tui/pr_sheet_test.go +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -87,26 +87,38 @@ func TestPRSheetCanToggleByMethodAndMessage(t *testing.T) { assert.False(t, model.prSheet.open) } -func TestPRSheetAllowsFileNavigationShortcuts(t *testing.T) { - model := NewModel([]core.ReviewFile{reviewFile("a.go", "package a"), reviewFile("b.go", "package b")}) +func TestPRSheetAllowsChunkNavigationShortcuts(t *testing.T) { + model := NewModel(prSheetChunkNavigationReviewFiles()) updated, _ := model.Update(tea.WindowSizeMsg{Width: 80, Height: 8}) model = updated.(Model).TogglePRSheet() + model.cursorRow = model.reviewAnchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 0}] + model.updateAfterCursorMove() updated, _ = model.Update(keyPress("l")) model = updated.(Model) - assert.Equal(t, 1, model.selectedFile) + assert.Equal(t, model.reviewAnchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 2, LineIndex: 0}], model.cursorRow) updated, _ = model.Update(keyPress("h")) model = updated.(Model) - assert.Equal(t, 0, model.selectedFile) + assert.Equal(t, model.reviewAnchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 0}], model.cursorRow) updated, _ = model.Update(tea.KeyPressMsg{Code: tea.KeyRight}) model = updated.(Model) - assert.Equal(t, 1, model.selectedFile) + assert.Equal(t, model.reviewAnchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 2, LineIndex: 0}], model.cursorRow) updated, _ = model.Update(tea.KeyPressMsg{Code: tea.KeyLeft}) model = updated.(Model) - assert.Equal(t, 0, model.selectedFile) + assert.Equal(t, model.reviewAnchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 0, LineIndex: 0}], model.cursorRow) +} + +func prSheetChunkNavigationReviewFiles() []core.ReviewFile { + return []core.ReviewFile{ + {Path: "a.go", Sections: []core.ReviewSection{ + {ID: "a-1", Kind: core.SectionKindChanged, Lines: []core.ReviewLine{{NewLineNumber: 1, Content: "a first", Kind: core.LineKindAdded}}}, + {ID: "a-context", Kind: core.SectionKindContext, Lines: []core.ReviewLine{{OldLineNumber: 2, NewLineNumber: 2, Content: "hidden a", Kind: core.LineKindUnchanged}}}, + {ID: "a-2", Kind: core.SectionKindChanged, Lines: []core.ReviewLine{{NewLineNumber: 3, Content: "a second", Kind: core.LineKindAdded}}}, + }}, + } } func TestPRSheetKeyboardAndWheelScrollSheetNotReview(t *testing.T) { From d9f07ffcb7f9556f7811324f392d977b9dc23aea Mon Sep 17 00:00:00 2001 From: brice Date: Sat, 6 Jun 2026 15:28:31 +0200 Subject: [PATCH 22/22] fix(providers): address final review findings --- README.md | 4 +- .../adapters/in/tui/active_provider_test.go | 25 +++++ .../adapters/in/tui/component/statusbar.go | 2 +- .../in/tui/component/statusbar_test.go | 8 +- internal/adapters/in/tui/model.go | 14 ++- internal/adapters/in/tui/review_context.go | 92 +++++++++++++++++++ internal/adapters/in/tui/review_providers.go | 1 + internal/adapters/in/tui/theme/styles.go | 35 +++---- .../adapters/out/plugin/provider_loader.go | 5 +- .../adapters/out/providercache/cache_test.go | 10 +- internal/app/active_provider_service.go | 33 ++++++- internal/app/active_provider_service_test.go | 66 +++++++++++++ internal/app/provider_polling_config.go | 3 + pkg/plugin/server_test.go | 10 +- .../github/cmd/ero-plugin-github/graphql.go | 5 +- .../cmd/ero-plugin-github/graphql_test.go | 8 +- plugins/github/cmd/ero-plugin-github/main.go | 5 +- .../github/cmd/ero-plugin-github/main_test.go | 38 +++++++- plugins/github/cmd/ero-plugin-github/match.go | 25 ++++- 19 files changed, 352 insertions(+), 37 deletions(-) create mode 100644 internal/adapters/in/tui/review_context.go diff --git a/README.md b/README.md index 9aff151..453f6c7 100644 --- a/README.md +++ b/README.md @@ -46,7 +46,9 @@ ero --context-lines 5 ## Plugins -Ero supports a general local subprocess plugin system, managed with `ero plugin install`, `ero plugin list`, `ero plugin update`, and `ero plugin remove`. The first shipped contribution type is `review_provider`, used by the maintained GitHub and pi-coding-agent plugins. Ero discovers all provider contributions but activates one review provider at a time, with provider switching, manual refresh, cache-first sync, and provider sync status in the TUI. The GitHub plugin uses GitHub CLI-compatible authentication through `go-gh` and requires `gh auth login`. See [docs/plugins.md](docs/plugins.md) for authoring details. +Ero supports a general local subprocess plugin system, managed with `ero plugin install`, `ero plugin list`, `ero plugin update`, and `ero plugin remove`. The first shipped contribution type is `review_provider`, used by the maintained GitHub and pi-coding-agent plugins. + +Ero discovers all provider contributions but activates one review provider at a time. The TUI supports provider switching, manual refresh, cache-first sync, and provider sync status. The GitHub plugin uses GitHub CLI-compatible authentication through `go-gh` and requires `gh auth login`. See [docs/plugins.md](docs/plugins.md) for authoring details. ## Development diff --git a/internal/adapters/in/tui/active_provider_test.go b/internal/adapters/in/tui/active_provider_test.go index 971ac89..820b1f4 100644 --- a/internal/adapters/in/tui/active_provider_test.go +++ b/internal/adapters/in/tui/active_provider_test.go @@ -45,6 +45,31 @@ func (m *mockActiveProviderController) CompleteTimer(ctx context.Context, review } func (m *mockActiveProviderController) Close() error { return m.Called().Error(0) } +func TestActiveProviderReloadToNonBranchModeClearsAndClosesProvider(t *testing.T) { + controller := &mockActiveProviderController{} + controller.On("Close").Return(nil).Once() + m := NewModelWithActiveProviderContext(context.Background(), []core.ReviewFile{reviewFile("branch.go", "package branch")}, nil, nil, core.ReviewRequest{DiffMode: core.DiffModeBranch}, nil, core.ReviewContext{Target: core.ReviewTargetMetadata{Mode: core.DiffModeBranch}}, controller, nil) + m.activeProviderKey = "github" + m.activeRuntimeID = "github" + m.activeRuntimeInfo = core.ReviewProviderInfo{ID: "github"} + m.providerOverview = &core.ProviderOverview{Title: "PR"} + m.providerSyncState = core.ProviderSyncState{Status: core.ProviderSyncStatusSynced} + m.remoteThreads = []core.RemoteReviewThread{{ExternalID: "old"}} + + updated, cmd := m.Update(reviewLoadedMsg{mode: core.DiffModeWorking, files: []core.ReviewFile{reviewFile("working.go", "package working")}}) + m = updated.(Model) + if cmd != nil { + _ = cmd() + } + + require.Empty(t, m.activeProviderKey) + require.Empty(t, m.activeRuntimeID) + require.Nil(t, m.providerOverview) + require.Empty(t, m.remoteThreads) + require.Equal(t, core.DiffModeWorking, m.reviewContext.Target.Mode) + controller.AssertExpectations(t) +} + func TestActiveProviderDoesNotStartOutsideBranchMode(t *testing.T) { controller := &mockActiveProviderController{} m := NewModelWithActiveProviderContext(context.Background(), nil, nil, nil, core.ReviewRequest{DiffMode: core.DiffModeUpstream}, nil, core.ReviewContext{Target: core.ReviewTargetMetadata{Mode: core.DiffModeUpstream}}, controller, nil) diff --git a/internal/adapters/in/tui/component/statusbar.go b/internal/adapters/in/tui/component/statusbar.go index a07310e..e728fa9 100644 --- a/internal/adapters/in/tui/component/statusbar.go +++ b/internal/adapters/in/tui/component/statusbar.go @@ -304,7 +304,7 @@ func providerSyncStatusLabel(status core.ProviderSyncStatus) string { } func formatStatusTime(value time.Time) string { - return value.UTC().Format("15:04") + return value.Local().Format("15:04") } func TruncateRunes(value string, width int) string { diff --git a/internal/adapters/in/tui/component/statusbar_test.go b/internal/adapters/in/tui/component/statusbar_test.go index c54e303..c2a7557 100644 --- a/internal/adapters/in/tui/component/statusbar_test.go +++ b/internal/adapters/in/tui/component/statusbar_test.go @@ -14,6 +14,8 @@ import ( func TestStatusbarProviderSync(t *testing.T) { baseTime := time.Date(2026, 6, 5, 12, 34, 0, 0, time.UTC) nextTime := baseTime.Add(5 * time.Minute) + baseLabel := baseTime.Local().Format("15:04") + nextLabel := nextTime.Local().Format("15:04") tests := []struct { name string @@ -40,17 +42,17 @@ func TestStatusbarProviderSync(t *testing.T) { { name: "synced", model: syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: core.ProviderSyncStatusSynced, LastSyncAt: &baseTime, NextSyncAt: &nextTime}), - want: []string{"GitHub/gh-runtime", "synced", "last 12:34", "next 12:39"}, + want: []string{"GitHub/gh-runtime", "synced", "last " + baseLabel, "next " + nextLabel}, }, { name: "failed", model: syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: core.ProviderSyncStatusFailed, LastSyncAt: &baseTime, LastError: "boom"}), - want: []string{"GitHub/gh-runtime", "failed", "boom", "last 12:34"}, + want: []string{"GitHub/gh-runtime", "failed", "boom", "last " + baseLabel}, }, { name: "backing-off", model: syncStatusModel("GitHub", "gh-runtime", core.ProviderSyncState{Status: core.ProviderSyncStatusBackingOff, LastSyncAt: &baseTime, NextSyncAt: &nextTime}), - want: []string{"GitHub/gh-runtime", "backoff", "last 12:34", "next 12:39"}, + want: []string{"GitHub/gh-runtime", "backoff", "last " + baseLabel, "next " + nextLabel}, }, } diff --git a/internal/adapters/in/tui/model.go b/internal/adapters/in/tui/model.go index b5d0fde..50a1c95 100644 --- a/internal/adapters/in/tui/model.go +++ b/internal/adapters/in/tui/model.go @@ -233,13 +233,23 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.remoteThreads = nil m.providerInfoByClient = map[ports.ReviewProviderClient]core.ReviewProviderInfo{} m.publish = publishState{} + m.reviewContext = m.reviewContextForLoadedReview(msg.mode) + m.clearActiveProviderRemoteData() m.resetContextSelection() m.reviewViewport.GotoTop() m.syncReviewViewport() + cmds := []tea.Cmd{} + if m.activeProvider != nil { + if m.activeProviderSyncEnabled() { + cmds = append(cmds, m.startActiveProviderCmd()) + } else { + cmds = append(cmds, m.closeReviewProvidersCmd()) + } + } if len(m.reviewProviders) > 0 { - return m, m.loadReviewProvidersCmd() + cmds = append(cmds, m.loadReviewProvidersCmd()) } - return m, nil + return m, tea.Batch(cmds...) case reviewLoadFailedMsg: m.loading = false m.loadError = msg.err.Error() diff --git a/internal/adapters/in/tui/review_context.go b/internal/adapters/in/tui/review_context.go new file mode 100644 index 0000000..5e33603 --- /dev/null +++ b/internal/adapters/in/tui/review_context.go @@ -0,0 +1,92 @@ +package tui + +import ( + "path/filepath" + + "ero/internal/core" +) + +func (m Model) reviewContextForLoadedReview(mode core.DiffMode) core.ReviewContext { + ctx := m.reviewContext + ctx.Target = reviewTargetForRequest(m.request, mode, ctx.Repository) + ctx.Diff = core.DiffMetadata{FilesChanged: len(m.files)} + ctx.Files = reviewFileMetadataFromFiles(m.files, &ctx.Diff) + return ctx +} + +func reviewTargetForRequest(request core.ReviewRequest, mode core.DiffMode, repo core.RepositoryMetadata) core.ReviewTargetMetadata { + target := core.ReviewTargetMetadata{Mode: mode, BaseRef: request.BaseRevision, HeadRef: request.HeadRevision} + switch mode { + case core.DiffModeBranch: + target.BaseRef = request.BaseRevision + if target.BaseRef == "" { + target.BaseRef = repo.DefaultBranch + } + target.HeadRef = repo.CurrentBranch + target.HeadSHA = repo.HeadSHA + case core.DiffModeCommit: + target.HeadRef = request.Revision + if target.HeadRef == "" { + target.HeadRef = "HEAD" + } + target.HeadSHA = repo.HeadSHA + case core.DiffModeRange: + target.BaseRef = request.BaseRevision + target.HeadRef = request.HeadRevision + case core.DiffModeUpstream: + target.BaseRef = request.UpstreamRef + if target.BaseRef == "" { + target.BaseRef = "@{upstream}" + } + target.HeadRef = "HEAD" + target.HeadSHA = repo.HeadSHA + } + return target +} + +func reviewFileMetadataFromFiles(files []core.ReviewFile, diff *core.DiffMetadata) []core.ReviewFileMetadata { + metadata := make([]core.ReviewFileMetadata, 0, len(files)) + for _, file := range files { + status := file.Status + if status == "" { + status = core.ReviewFileStatusModified + } + meta := core.ReviewFileMetadata{Path: file.Path, OldPath: file.OldPath, Status: status, Language: languageFromReviewPath(file.Path)} + for _, section := range file.Sections { + if section.Kind == core.SectionKindChanged { + anchor := core.ReviewHunkAnchor{SectionID: section.ID} + for _, line := range section.VisibleLines() { + if anchor.OldStartLine == 0 && line.OldLineNumber > 0 { + anchor.OldStartLine = line.OldLineNumber + } + if anchor.NewStartLine == 0 && line.NewLineNumber > 0 { + anchor.NewStartLine = line.NewLineNumber + } + if anchor.OldStartLine > 0 && anchor.NewStartLine > 0 { + break + } + } + meta.Hunks = append(meta.Hunks, anchor) + } + for _, line := range section.VisibleLines() { + if line.Kind == core.LineKindAdded { + diff.Additions++ + } + if line.Kind == core.LineKindDeleted { + diff.Deletions++ + } + meta.LineAnchors = append(meta.LineAnchors, core.NewReviewLineAnchor(file.Path, line)) + } + } + metadata = append(metadata, meta) + } + return metadata +} + +func languageFromReviewPath(path string) string { + ext := filepath.Ext(path) + if len(ext) > 1 { + return ext[1:] + } + return "" +} diff --git a/internal/adapters/in/tui/review_providers.go b/internal/adapters/in/tui/review_providers.go index bcc1c56..313d360 100644 --- a/internal/adapters/in/tui/review_providers.go +++ b/internal/adapters/in/tui/review_providers.go @@ -137,6 +137,7 @@ func (m *Model) applyActiveProviderState(state ActiveProviderState) { } func (m *Model) clearActiveProviderRemoteData() { + m.activeProviderKey = "" m.activeRuntimeID = "" m.activeRuntimeInfo = core.ReviewProviderInfo{} m.providerSyncState = core.ProviderSyncState{} diff --git a/internal/adapters/in/tui/theme/styles.go b/internal/adapters/in/tui/theme/styles.go index 8def904..4a9857c 100644 --- a/internal/adapters/in/tui/theme/styles.go +++ b/internal/adapters/in/tui/theme/styles.go @@ -9,6 +9,7 @@ const ( ColorWarning = "#ffa657" ColorKeyword = "#ff7b72" ColorFunction = "#d2a8ff" + ColorType = "#ffa657" ColorString = "#a5d6ff" ColorNumber = "#79c0ff" ColorStatusBase = "236" @@ -21,28 +22,28 @@ var ( FileRuleStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("8")) PanelTitleStyle = lipgloss.NewStyle().Bold(true).Underline(true) MutedStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("8")) - AddedLineStyle = lipgloss.NewStyle().Background(lipgloss.Color("#011209")).Foreground(lipgloss.Color("#c9d1d9")) - DeletedLineStyle = lipgloss.NewStyle().Background(lipgloss.Color("#1f0101")).Foreground(lipgloss.Color("#c9d1d9")) + AddedLineStyle = lipgloss.NewStyle().Background(lipgloss.Color("#011209")).Foreground(lipgloss.Color(ColorText)) + DeletedLineStyle = lipgloss.NewStyle().Background(lipgloss.Color("#1f0101")).Foreground(lipgloss.Color(ColorText)) AddedMarkerStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#3fb950")).Bold(true) - DeletedMarkerStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#ff7b72")).Bold(true) - LineNumberStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#8b949e")) - SelectedExpander = lipgloss.NewStyle().Bold(true).Foreground(lipgloss.Color("#58a6ff")) - CursorRowStyle = lipgloss.NewStyle().Background(lipgloss.Color("#1f2a44")) + DeletedMarkerStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorKeyword)).Bold(true) + LineNumberStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorMutedText)) + SelectedExpander = lipgloss.NewStyle().Bold(true).Foreground(lipgloss.Color(ColorAccent)) + CursorRowStyle = lipgloss.NewStyle().Background(lipgloss.Color(ColorCodeBg)) SelectedRowStyle = lipgloss.NewStyle().Background(lipgloss.Color("#25351f")) CommentRangeRowStyle = lipgloss.NewStyle().Background(lipgloss.Color("#201a35")) - KeywordStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#ff7b72")) - FunctionStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#d2a8ff")) - TypeStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#ffa657")) - NameStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#c9d1d9")) - StringStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#a5d6ff")) - NumberStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#79c0ff")) - CommentStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#8b949e")).Italic(true) - OperatorStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#ff7b72")) - PunctuationStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#c9d1d9")) - StatusBaseStyle = lipgloss.NewStyle().Background(lipgloss.Color("236")).Foreground(lipgloss.Color("252")) + KeywordStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorKeyword)) + FunctionStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorFunction)) + TypeStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorType)) + NameStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorText)) + StringStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorString)) + NumberStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorNumber)) + CommentStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorMutedText)).Italic(true) + OperatorStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorKeyword)) + PunctuationStyle = lipgloss.NewStyle().Foreground(lipgloss.Color(ColorText)) + StatusBaseStyle = lipgloss.NewStyle().Background(lipgloss.Color(ColorStatusBase)).Foreground(lipgloss.Color("252")) StatusAppStyle = StatusBaseStyle.Bold(true).Background(lipgloss.Color("62")).Foreground(lipgloss.Color("230")).Padding(0, 1) StatusModeStyle = StatusBaseStyle.Foreground(lipgloss.Color("229")).Padding(0, 1) - StatusInfoStyle = StatusBaseStyle.Foreground(lipgloss.Color("248")).Padding(0, 1) + StatusInfoStyle = StatusBaseStyle.Foreground(lipgloss.Color(ColorStatusInfo)).Padding(0, 1) StatusKeyStyle = StatusBaseStyle.Foreground(lipgloss.Color("81")).Bold(true) StatusHintTextStyle = StatusBaseStyle.Foreground(lipgloss.Color("244")) diff --git a/internal/adapters/out/plugin/provider_loader.go b/internal/adapters/out/plugin/provider_loader.go index f8a87b8..de68bf1 100644 --- a/internal/adapters/out/plugin/provider_loader.go +++ b/internal/adapters/out/plugin/provider_loader.go @@ -94,8 +94,11 @@ func (l *ReviewProviderLoader) createReviewProviderClient(ctx context.Context, d } if shouldBuildRuntime(command, descriptor.PluginPath, manifest.Build.Command) { if err := runPluginBuildCommand(ctx, descriptor.PluginPath, manifest.Build.Command, l.timeout); err != nil { + if _, available := runtimeCommandInfo(command, descriptor.PluginPath); !available { + return nil, fmt.Errorf("build plugin runtime: %w", err) + } log := zerowrap.FromCtx(ctx) - log.Warn().Err(err).Str("plugin_path", descriptor.PluginPath).Msg("build plugin runtime failed") + log.Warn().Err(err).Str("plugin_path", descriptor.PluginPath).Msg("build plugin runtime failed; using existing runtime") } } if !strings.Contains(command, "/") { diff --git a/internal/adapters/out/providercache/cache_test.go b/internal/adapters/out/providercache/cache_test.go index b17b3b8..024d9f8 100644 --- a/internal/adapters/out/providercache/cache_test.go +++ b/internal/adapters/out/providercache/cache_test.go @@ -23,8 +23,14 @@ func TestCacheRoundTripNormalizedSnapshot(t *testing.T) { if !ok { t.Fatal("expected cached snapshot") } - if got.StableProviderKey != snapshot.StableProviderKey || got.RuntimeProviderID != "github" || len(got.Threads) != 1 { - t.Fatalf("unexpected snapshot: %#v", got) + if got.StableProviderKey != snapshot.StableProviderKey { + t.Fatalf("StableProviderKey mismatch: got %q, want %q", got.StableProviderKey, snapshot.StableProviderKey) + } + if got.RuntimeProviderID != "github" { + t.Fatalf("RuntimeProviderID mismatch: got %q, want %q", got.RuntimeProviderID, "github") + } + if len(got.Threads) != 1 { + t.Fatalf("unexpected thread count: got %d, want 1", len(got.Threads)) } } diff --git a/internal/app/active_provider_service.go b/internal/app/active_provider_service.go index 1b9538c..50100c5 100644 --- a/internal/app/active_provider_service.go +++ b/internal/app/active_provider_service.go @@ -88,6 +88,14 @@ func (s *ActiveProviderService) Start(ctx context.Context, review core.ReviewCon continue } s.mu.Lock() + if s.generation != startGen { + staleState := s.state + s.mu.Unlock() + if err := client.Close(); err != nil { + log.Warn().Err(err).Str("provider_key", d.Key).Msg("active provider stale start close failed") + } + return staleState, nil + } if err := s.closeLocked(); err != nil { log.Warn().Err(err).Msg("active provider close before activate failed") } @@ -151,6 +159,14 @@ func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewCo return failed, err } s.mu.Lock() + if s.generation != switchGen { + staleState := s.state + s.mu.Unlock() + if err := client.Close(); err != nil { + log.Warn().Err(err).Str("provider_key", stableKey).Msg("active provider stale switch close failed") + } + return staleState, nil + } if err := s.closeLocked(); err != nil { log.Warn().Err(err).Str("provider_key", stableKey).Msg("active provider close before switched activate failed") } @@ -175,7 +191,22 @@ func (s *ActiveProviderService) Switch(ctx context.Context, review core.ReviewCo } } log.Warn().Str("provider_key", stableKey).Msg("active provider switch descriptor not found") - return ActiveProviderState{}, core.NewProviderError(core.ProviderErrorNotApplicable, "provider descriptor not found", nil) + err = core.NewProviderError(core.ProviderErrorNotApplicable, "provider descriptor not found", nil) + failed := failedProviderState(err) + s.mu.Lock() + if closeErr := s.closeLocked(); closeErr != nil { + log.Warn().Err(closeErr).Str("provider_key", stableKey).Msg("active provider close after missing switch descriptor failed") + } + s.stableKey = "" + s.runtimeID = "" + s.generation++ + gen := s.generation + s.state = ActiveProviderState{} + s.mu.Unlock() + if !s.setState(gen, failed) { + return s.State(), nil + } + return failed, err } func (s *ActiveProviderService) PublishReview(ctx context.Context, request core.PublishReviewRequest) (core.PublishReviewResult, error) { diff --git a/internal/app/active_provider_service_test.go b/internal/app/active_provider_service_test.go index 70b78b6..4797df7 100644 --- a/internal/app/active_provider_service_test.go +++ b/internal/app/active_provider_service_test.go @@ -80,6 +80,72 @@ func TestActiveProviderServicePreferenceFallbackAndClosesFailedClients(t *testin } } +func TestActiveProviderServiceCloseInvalidatesInFlightStart(t *testing.T) { + review := testReviewContext() + started := make(chan struct{}) + release := make(chan struct{}) + provider := mocks.NewMockReviewProviderClient(t) + provider.EXPECT().Initialize(mock.Anything).Run(func(context.Context) { + close(started) + <-release + }).Return(core.ReviewProviderInfo{ID: "github", Capabilities: core.ReviewProviderCapabilities{LoadRemoteComments: true}}, nil).Once() + provider.EXPECT().DetectContext(mock.Anything, mock.Anything).Return(core.DetectionResult{Applicable: true}, nil).Once() + provider.EXPECT().Close().Return(nil).Once() + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "github", provider) + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "github", Type: "github"}), factory, nil, nil, ProviderPollingConfig{}) + + var wg sync.WaitGroup + var startState ActiveProviderState + var startErr error + wg.Add(1) + go func() { + defer wg.Done() + startState, startErr = svc.Start(context.Background(), review) + }() + <-started + if err := svc.Close(); err != nil { + t.Fatal(err) + } + close(release) + wg.Wait() + + if startErr != nil { + t.Fatalf("stale start should be discarded without surfacing old error: %v", startErr) + } + if startState.StableProviderKey != "" { + t.Fatalf("stale start returned active provider state: %#v", startState) + } + if state := svc.State(); state.StableProviderKey != "" { + t.Fatalf("close should win over stale start, got %#v", state) + } +} + +func TestActiveProviderServiceMissingSwitchDescriptorClearsOldProviderState(t *testing.T) { + review := testReviewContext() + provider := mocks.NewMockReviewProviderClient(t) + expectProbe(provider, "github", true) + provider.EXPECT().Close().Return(nil).Once() + factory := mocks.NewMockReviewProviderClientFactory(t) + expectFactoryClient(factory, "github", provider) + svc := NewActiveProviderService(mockCatalog(t, ports.ReviewProviderDescriptor{Key: "github", Type: "github"}), factory, nil, nil, ProviderPollingConfig{}) + if _, err := svc.Start(context.Background(), review); err != nil { + t.Fatal(err) + } + + st, err := svc.Switch(context.Background(), review, "missing") + + if got := core.ClassifyProviderError(err); got != core.ProviderErrorNotApplicable { + t.Fatalf("expected not applicable, got %q (%v)", got, err) + } + if st.StableProviderKey != "" || st.LastError == nil { + t.Fatalf("missing descriptor should return failed empty state, got %#v", st) + } + if state := svc.State(); state.StableProviderKey != "" || state.LastError == nil { + t.Fatalf("missing descriptor should clear service state, got %#v", state) + } +} + func TestActiveProviderServiceStartWithoutCandidatesReturnsNotApplicable(t *testing.T) { svc := NewActiveProviderService(mockCatalog(t), mocks.NewMockReviewProviderClientFactory(t), nil, nil, ProviderPollingConfig{}) diff --git a/internal/app/provider_polling_config.go b/internal/app/provider_polling_config.go index ad8f9e3..6853411 100644 --- a/internal/app/provider_polling_config.go +++ b/internal/app/provider_polling_config.go @@ -26,5 +26,8 @@ func providerPollingConfigFromConfig(cfg *viper.Viper) ProviderPollingConfig { if poll.MaxBackoff == 0 { poll.MaxBackoff = time.Minute } + if poll.MinBackoff > poll.MaxBackoff { + poll.MinBackoff = poll.MaxBackoff + } return poll } diff --git a/pkg/plugin/server_test.go b/pkg/plugin/server_test.go index bb6fea0..4821b28 100644 --- a/pkg/plugin/server_test.go +++ b/pkg/plugin/server_test.go @@ -299,8 +299,14 @@ func TestLoadRemoteSnapshotDispatchesOptionalMethod(t *testing.T) { if err := json.Unmarshal(output.Bytes(), &response); err != nil { t.Fatalf("invalid json: %v", err) } - if response.Error != nil || response.Result.Overview == nil || response.Result.Overview.Title != "PR" { - t.Fatalf("unexpected response: %#v raw=%s", response, output.String()) + if response.Error != nil { + t.Fatalf("unexpected error: %v raw=%s", response.Error, output.String()) + } + if response.Result.Overview == nil { + t.Fatalf("missing overview in result: %#v raw=%s", response.Result, output.String()) + } + if response.Result.Overview.Title != "PR" { + t.Fatalf("unexpected overview title: got %q, want %q", response.Result.Overview.Title, "PR") } if !reflect.DeepEqual(provider.methodsCalled, []string{"load_remote_snapshot"}) { t.Fatalf("expected load_remote_snapshot call, got %v", provider.methodsCalled) diff --git a/plugins/github/cmd/ero-plugin-github/graphql.go b/plugins/github/cmd/ero-plugin-github/graphql.go index 7d7a23a..3bb2dcf 100644 --- a/plugins/github/cmd/ero-plugin-github/graphql.go +++ b/plugins/github/cmd/ero-plugin-github/graphql.go @@ -203,7 +203,8 @@ func matchGitHubPRAcrossRemotes(ctx context.Context, client graphQLDoer, remotes for _, remote := range remotes { candidates, err := fetchGitHubPRCandidates(ctx, client, remote) if err != nil { - return githubRemote{}, githubPRCandidate{}, err + lastErr = err + continue } match, err := matchGitHubPR(reviewCtx, candidates) if err != nil { @@ -275,6 +276,8 @@ func fetchGitHubPRSnapshot(ctx context.Context, client graphQLDoer, remote githu for _, thread := range pr.ReviewThreads.Nodes { mapped := mapGitHubThread(thread) if thread.Comments.PageInfo.HasNextPage { + // Keep nested pagination bounded: threads with more than 100 comments may be incomplete, + // so mark them unmapped rather than anchoring partial discussion inline. mapped.Unmapped = true } accum.Threads = append(accum.Threads, mapped) diff --git a/plugins/github/cmd/ero-plugin-github/graphql_test.go b/plugins/github/cmd/ero-plugin-github/graphql_test.go index 093f19e..0efb13d 100644 --- a/plugins/github/cmd/ero-plugin-github/graphql_test.go +++ b/plugins/github/cmd/ero-plugin-github/graphql_test.go @@ -14,9 +14,10 @@ import ( type fakeGraphQLClient struct { listPages []ghPRListResponse + listErrs []error snapshotPages []ghPRSnapshotResponse - listCalls int snapshotCalls int + listCalls int vars []map[string]any } @@ -26,6 +27,11 @@ func (f *fakeGraphQLClient) DoWithContext(_ context.Context, query string, varia f.vars = append(f.vars, copied) switch r := response.(type) { case *ghPRListResponse: + if f.listCalls < len(f.listErrs) && f.listErrs[f.listCalls] != nil { + err := f.listErrs[f.listCalls] + f.listCalls++ + return err + } *r = f.listPages[f.listCalls] f.listCalls++ case *ghPRSnapshotResponse: diff --git a/plugins/github/cmd/ero-plugin-github/main.go b/plugins/github/cmd/ero-plugin-github/main.go index 3da18cc..871d0af 100644 --- a/plugins/github/cmd/ero-plugin-github/main.go +++ b/plugins/github/cmd/ero-plugin-github/main.go @@ -82,7 +82,10 @@ func (p githubProvider) DetectContext(ctx context.Context, req plugin.DetectCont } _, match, err := matchGitHubPRAcrossRemotes(ctx, client, remotes, req.Context) if err != nil { - return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: err.Error()}}, nil + if pe := plugin.AsError(err); pe != nil && pe.Code == plugin.ErrorNotApplicable { + return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: err.Error()}}, nil + } + return plugin.DetectContextResult{}, err } return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: true, Reason: "matched GitHub pull request " + githubPRSummary(match)}}, nil } diff --git a/plugins/github/cmd/ero-plugin-github/main_test.go b/plugins/github/cmd/ero-plugin-github/main_test.go index 416f2fc..14709f8 100644 --- a/plugins/github/cmd/ero-plugin-github/main_test.go +++ b/plugins/github/cmd/ero-plugin-github/main_test.go @@ -145,11 +145,23 @@ func TestGitHubPRMatching(t *testing.T) { wantNumber: 9, }, { - name: "matching branch accepts local unpushed head SHA", - ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch", HeadSHA: "local-unpushed"}}, - prs: []githubPRCandidate{{Number: 10, BaseRef: "main", HeadRef: "feature", HeadSHA: "remote-pr-head"}}, + name: "matching branch accepts local unpushed head SHA from local remote", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "git@github.com:owner/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch", HeadSHA: "local-unpushed"}}, + prs: []githubPRCandidate{{Number: 10, BaseRef: "main", HeadRef: "feature", HeadRepoOwner: "owner", HeadRepoName: "repo", HeadSHA: "remote-pr-head"}}, wantNumber: 10, }, + { + name: "matching branch rejects another fork with different SHA", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "git@github.com:owner/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch", HeadSHA: "local-head"}}, + prs: []githubPRCandidate{{Number: 11, BaseRef: "main", HeadRef: "feature", HeadRepoOwner: "other", HeadRepoName: "repo", HeadSHA: "other-head"}}, + wantErr: plugin.ErrorNotApplicable, + }, + { + name: "matching branch accepts another fork when SHA matches", + ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{URL: "git@github.com:owner/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch", HeadSHA: "same-head"}}, + prs: []githubPRCandidate{{Number: 12, BaseRef: "main", HeadRef: "feature", HeadRepoOwner: "other", HeadRepoName: "repo", HeadSHA: "same-head"}}, + wantNumber: 12, + }, { name: "ambiguous multiple matches returns not applicable", ctx: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}, @@ -247,6 +259,26 @@ func TestPublishReviewFallsBackToGHCLIWhenGraphQLHasNoMatch(t *testing.T) { } } +func TestLoadRemoteSnapshotContinuesAfterRemoteListError(t *testing.T) { + upstreamList := ghPRListResponse{} + upstreamList.Repository.PullRequests.Nodes = []ghPRNode{{Number: 2, BaseRefName: "main", HeadRefName: "feature", HeadRepositoryOwner: &ghActor{Login: "upstream"}, HeadRepository: &struct { + Name string `json:"name"` + }{Name: "repo"}}} + snapshot := ghPRSnapshotResponse{} + snapshot.Repository.PullRequest = ghPRNode{Number: 2, URL: "https://github.com/upstream/repo/pull/2", Title: "PR", BaseRefName: "main", HeadRefName: "feature"} + fake := &fakeGraphQLClient{listErrs: []error{plugin.NewError(plugin.ErrorNetwork, "remote unavailable")}, listPages: []ghPRListResponse{{}, upstreamList}, snapshotPages: []ghPRSnapshotResponse{snapshot}} + provider := githubProvider{newGraphQLClient: func() (graphQLDoer, error) { return fake, nil }} + + got, err := provider.LoadRemoteSnapshot(context.Background(), plugin.LoadRemoteSnapshotRequest{Context: plugin.ReviewContext{Repository: plugin.RepositoryMetadata{Remotes: []plugin.GitRemote{{Name: "origin", URL: "git@github.com:fork/repo.git"}, {Name: "upstream", URL: "git@github.com:upstream/repo.git"}}, CurrentBranch: "feature", DefaultBranch: "main"}, Target: plugin.ReviewTargetMetadata{Mode: "branch"}}}) + + if err != nil { + t.Fatalf("LoadRemoteSnapshot returned error: %v", err) + } + if fake.listCalls != 2 || got.Overview == nil || got.Overview.Number != 2 { + t.Fatalf("expected second remote PR match after first remote error, calls=%d snapshot=%#v", fake.listCalls, got.Overview) + } +} + func TestLoadRemoteSnapshotSearchesAllGitHubRemotes(t *testing.T) { forkList := ghPRListResponse{} forkList.Repository.PullRequests.Nodes = []ghPRNode{{Number: 1, BaseRefName: "main", HeadRefName: "other"}} diff --git a/plugins/github/cmd/ero-plugin-github/match.go b/plugins/github/cmd/ero-plugin-github/match.go index 6c18c69..079d857 100644 --- a/plugins/github/cmd/ero-plugin-github/match.go +++ b/plugins/github/cmd/ero-plugin-github/match.go @@ -53,12 +53,35 @@ func githubPRMatches(ctx plugin.ReviewContext, pr githubPRCandidate) bool { headRef = strings.TrimSpace(ctx.Repository.CurrentBranch) } if headRef != "" { - return refEqual(pr.HeadRef, headRef) + if !refEqual(pr.HeadRef, headRef) { + return false + } + if localGitHubRemoteKnown(ctx) && prHeadRepositoryKnown(pr) && !prHeadRepositoryMatchesLocalRemote(ctx, pr) { + return headSHA != "" && pr.HeadSHA != "" && strings.EqualFold(pr.HeadSHA, headSHA) + } + return true } return headSHA != "" && pr.HeadSHA != "" && strings.EqualFold(pr.HeadSHA, headSHA) } +func localGitHubRemoteKnown(ctx plugin.ReviewContext) bool { + return len(githubRemotes(ctx.Repository.Remotes)) > 0 +} + +func prHeadRepositoryKnown(pr githubPRCandidate) bool { + return strings.TrimSpace(pr.HeadRepoOwner) != "" && strings.TrimSpace(pr.HeadRepoName) != "" +} + +func prHeadRepositoryMatchesLocalRemote(ctx plugin.ReviewContext, pr githubPRCandidate) bool { + for _, remote := range githubRemotes(ctx.Repository.Remotes) { + if strings.EqualFold(remote.Owner, pr.HeadRepoOwner) && strings.EqualFold(remote.Name, pr.HeadRepoName) { + return true + } + } + return false +} + func nonBranchHeadRef(ref string) bool { ref = strings.TrimSpace(ref) if ref == "" {