diff --git a/README.md b/README.md index c741687..453f6c7 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,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. 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. 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/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 962a07f..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 @@ -48,6 +48,12 @@ 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. + +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 @@ -78,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 @@ -117,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: @@ -137,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/go.mod b/go.mod index 6153a0b..2e0c08e 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,12 +26,16 @@ 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 + 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 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 @@ -43,6 +48,8 @@ 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/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 @@ -51,7 +58,9 @@ 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/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 @@ -64,8 +73,11 @@ 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 + 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..e35d114 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,10 @@ 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/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= @@ -126,6 +136,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 +197,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= @@ -202,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= @@ -219,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/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/in/tui/active_provider_test.go b/internal/adapters/in/tui/active_provider_test.go new file mode 100644 index 0000000..820b1f4 --- /dev/null +++ b/internal/adapters/in/tui/active_provider_test.go @@ -0,0 +1,226 @@ +package tui + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + + "ero/internal/core" + "ero/internal/ports" + portmocks "ero/internal/ports/mocks" +) + +type mockActiveProviderController struct{ mock.Mock } + +func (m *mockActiveProviderController) Catalog(ctx context.Context) ([]ports.ReviewProviderDescriptor, error) { + args := m.Called(ctx) + return args.Get(0).([]ports.ReviewProviderDescriptor), args.Error(1) +} +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 (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 (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 (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 (m *mockActiveProviderController) Generation() int64 { + args := m.Called() + return args.Get(0).(int64) +} +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 (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) + + 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"}} + 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() + 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) + updated, refreshCmd := m.Update(cmd()) + m = updated.(Model) + + require.NotNil(t, refreshCmd) + 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 := &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()()) + 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) + 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 := &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"}} + + msg := m.refreshActiveProviderCmd(true)() + updated, _ := m.Update(msg) + m = updated.(Model) + + 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 := &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"}} + + msg := m.switchActiveProviderCmd("other")() + updated, _ := m.Update(msg) + m = updated.(Model) + + 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 := &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) + + 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 := &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" + 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) + + 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 := &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" + 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) + controller.AssertExpectations(t) +} diff --git a/internal/adapters/in/tui/component/help_pane.go b/internal/adapters/in/tui/component/help_pane.go index bd9eb9c..1c3887c 100644 --- a/internal/adapters/in/tui/component/help_pane.go +++ b/internal/adapters/in/tui/component/help_pane.go @@ -19,7 +19,16 @@ func RenderHelpPane(width, height int, enterKeyLabel, commentSubmitKeyLabel stri renderHelpShortcut("f", "find file", contentWidth), renderHelpShortcut("/", "grep references", contentWidth), renderHelpShortcut("d", "switch diff mode", contentWidth), - renderHelpShortcut("h/l", "previous/next file", contentWidth), + renderHelpShortcut("p", "switch provider", contentWidth), + renderHelpShortcut("alt+p", "cycle 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 chunk", contentWidth), renderHelpShortcut("a", "expand all context", contentWidth), renderHelpShortcut(enterKeyLabel, "expand more context", contentWidth), renderHelpShortcut("s/space", "select lines", contentWidth), @@ -29,11 +38,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..e728fa9 100644 --- a/internal/adapters/in/tui/component/statusbar.go +++ b/internal/adapters/in/tui/component/statusbar.go @@ -2,22 +2,32 @@ package component import ( "fmt" + "image/color" "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 + ProviderSwitch bool + CurrentFile string + Message string + ScrollPercent float64 + ActiveProviderLabel string + ActiveRuntimeName string + ProviderSync core.ProviderSyncState + DraftCommentCount int + ShowNoProvider bool + NerdFont bool } type StatusBar struct { @@ -30,7 +40,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{ @@ -38,9 +48,15 @@ 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 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)}) + } prefix := renderStatusSegments(leftWidth, segments...) percent := renderStatusSegments(leftWidth-lipgloss.Width(prefix), statusSegment{style: theme.StatusInfoStyle, label: fmt.Sprintf("%3.0f%%", model.ScrollPercent*100)}) @@ -60,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 { @@ -69,10 +86,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: "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 { @@ -81,6 +101,9 @@ func renderStatusHint(width, providerCount int) string { fallback := "? help" if providerCount > 0 { fallback = "P publish" + if providerSwitch { + fallback = "p provider" + } } return theme.StatusInfoStyle.Render(TruncateRunes(fallback, max(width-theme.StatusInfoStyle.GetHorizontalPadding(), 0))) } @@ -93,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 { @@ -125,6 +152,161 @@ 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 == "" { + 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 != "" { + 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 draftCommentCountLabel(count int) string { + if count == 1 { + return "1 draft comment" + } + return fmt.Sprintf("%d draft comments", count) +} + +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)) + } + 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 nerdFontGitHubLarge + } + return providerAbbreviation(provider) +} + +func providerStatusDotColor(status core.ProviderSyncStatus) color.Color { + switch status { + case core.ProviderSyncStatusSynced: + return lipgloss.Color("#3fb950") + case core.ProviderSyncStatusFailed: + return lipgloss.Color("#ff7b72") + case core.ProviderSyncStatusBackingOff: + return lipgloss.Color("#ffa657") + case core.ProviderSyncStatusLoadingCache, core.ProviderSyncStatusSyncing: + return lipgloss.Color("#58a6ff") + default: + return lipgloss.Color("81") + } +} + +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: + 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 formatStatusTime(value time.Time) string { + return value.Local().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..c2a7557 --- /dev/null +++ b/internal/adapters/in/tui/component/statusbar_test.go @@ -0,0 +1,160 @@ +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) + baseLabel := baseTime.Local().Format("15:04") + nextLabel := nextTime.Local().Format("15:04") + + 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 " + 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 " + 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 " + baseLabel, "next " + nextLabel}, + }, + } + + 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 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, nerdFontGitHubLarge+" "+nerdFontSyncDot) + } +} + +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, nerdFontGitHubLarge+" "+nerdFontSyncDot) + 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.ProviderSwitch = true + model.NerdFont = true + + raw := NewStatusBar(120).Render(model) + view := stripANSIForStatusbarTest(raw) + + 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") +} + +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) + 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/cursor_navigation_test.go b/internal/adapters/in/tui/cursor_navigation_test.go index cb3192b..5927b2f 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() @@ -102,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/help_pane_test.go b/internal/adapters/in/tui/help_pane_test.go index b10c0ca..d1b98c4 100644 --- a/internal/adapters/in/tui/help_pane_test.go +++ b/internal/adapters/in/tui/help_pane_test.go @@ -24,6 +24,11 @@ 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, "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") updated, _ = model.Update(tea.KeyPressMsg{Code: tea.KeyEsc}) 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/keymap/action.go b/internal/adapters/in/tui/keymap/action.go index 16028e0..25c3b10 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 { @@ -53,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 @@ -67,15 +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 "?": + case "alt+p": + return ActionCycleProvider + case "p": + return ActionOpenProviderPicker + case "r": + return ActionRefreshProvider + case "o": + return ActionTogglePRSheet + 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 bff27e0..d71ec55 100644 --- a/internal/adapters/in/tui/keymap/action_test.go +++ b/internal/adapters/in/tui/keymap/action_test.go @@ -30,21 +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: "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/markdown_renderer.go b/internal/adapters/in/tui/markdown_renderer.go new file mode 100644 index 0000000..7d66c2e --- /dev/null +++ b/internal/adapters/in/tui/markdown_renderer.go @@ -0,0 +1,171 @@ +package tui + +import ( + "crypto/sha256" + "encoding/hex" + "regexp" + "strings" + + "charm.land/glamour/v2" + glamouransi "charm.land/glamour/v2/ansi" + + "ero/internal/adapters/in/tui/theme" +) + +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) + } + rendered = sanitizeRenderedMarkdown(rendered) + 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.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;?]*[ -/]*[@-~]`) + oscEscapePattern = regexp.MustCompile(`\x1b\][^\x1b\x07]*(?:\x07|\x1b\\)`) +) + +func sanitizeRenderedMarkdown(input string) string { + return oscEscapePattern.ReplaceAllString(input, "") +} + +func safeMarkdownFallback(input string) string { + 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 new file mode 100644 index 0000000..4f1c156 --- /dev/null +++ b/internal/adapters/in/tui/markdown_renderer_test.go @@ -0,0 +1,109 @@ +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 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() + + 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 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) { + return "", errors.New("boom") + }}, nil + }) + + 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") || strings.Contains(got, "]8;") { + 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 aca7cca..50a1c95 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,20 @@ 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 + markdownRenderer *MarkdownRenderer ctx context.Context reviewLineCache *render.ReviewLineCache cachedEditorWidth int @@ -120,6 +164,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,9 +190,11 @@ 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, + markdownRenderer: NewMarkdownRenderer(), ctx: ctx, reviewLineCache: render.NewReviewLineCache(), } @@ -157,6 +207,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 } @@ -180,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() @@ -213,6 +276,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 @@ -231,6 +330,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() { @@ -249,10 +350,65 @@ 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) } - return m.updateReviewAction(keymap.ReviewAction(msg.String())) + if m.prSheet.open { + return m.updatePRSheetAction(msg) + } + return m.updateReviewAction(keymap.ReviewAction(msg.Keystroke())) + default: + return m, nil + } +} + +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: + 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 + } + 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.ActionPreviousFile: + m.moveChunk(-1) + return m, nil + case keymap.ActionNextFile: + m.moveChunk(1) + return m, 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 } @@ -297,13 +453,27 @@ 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: 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) + case keymap.ActionTogglePRSheet: + m = m.TogglePRSheet() case keymap.ActionOpenHelp: m.helpActive = true case keymap.ActionNone: @@ -313,6 +483,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 { @@ -323,13 +506,20 @@ 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(), + ProviderSwitch: m.canSwitchProvider(), + CurrentFile: m.activeLocation(), + Message: m.copyFeedback, + ScrollPercent: m.reviewViewport.ScrollPercent(), + 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, }), ) if m.search.active() { @@ -338,11 +528,18 @@ 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) } 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_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_test.go b/internal/adapters/in/tui/model_test.go index a914c30..ad4b0e2 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,11 +27,28 @@ 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) }) } } +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/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.go b/internal/adapters/in/tui/pr_sheet.go new file mode 100644 index 0000000..060b63f --- /dev/null +++ b/internal/adapters/in/tui/pr_sheet.go @@ -0,0 +1,265 @@ +package tui + +import ( + "strconv" + "strings" + "time" + + "charm.land/lipgloss/v2" + "github.com/charmbracelet/x/ansi" + + "ero/internal/core" +) + +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) + paneWidth := prSheetWidth(width) + pane := m.renderPRSheet(width, height) + + 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 := "│ " + ansi.Truncate(text, contentWidth, "") + rows[i] = padRightANSI(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) + } + + lines := []string{ + "Pull request", + "", + "Provider: " + provider, + } + 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 := sanitizeRenderedMarkdown(renderer.Render(markdown, width, MarkdownThemeDark)) + if strings.TrimSpace(safeMarkdownFallback(rendered)) == "" { + return []string{"(empty)"} + } + 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 { + 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 + } + return max(totalWidth/2, 1) +} + +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 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 new file mode 100644 index 0000000..944a289 --- /dev/null +++ b/internal/adapters/in/tui/pr_sheet_test.go @@ -0,0 +1,249 @@ +package tui + +import ( + "strings" + "testing" + "time" + + 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) + 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 { + 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 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() + + 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 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, model.reviewAnchors.LineRows[ReviewLineAnchor{FileIndex: 0, SectionIndex: 2, LineIndex: 0}], model.cursorRow) + + updated, _ = model.Update(keyPress("h")) + model = updated.(Model) + 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, 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, 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) { + 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() + + 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") +} + +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 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() + + 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/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/provider_picker.go b/internal/adapters/in/tui/provider_picker.go new file mode 100644 index 0000000..5861ce1 --- /dev/null +++ b/internal/adapters/in/tui/provider_picker.go @@ -0,0 +1,193 @@ +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.Keystroke() { + 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) + case "alt+p": + m = m.closeProviderPicker() + return m, m.cycleProviderCmd() + 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 { + // 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("Provider"), theme.MutedStyle.Render("Active publish destination"), ""} + 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 = "› " + } + state := "available" + marker := "○" + if row.Active { + state = "active" + marker = "●" + } + meta := strings.TrimSpace(strings.Join([]string{row.PluginName, row.PluginSource}, " ")) + line := fmt.Sprintf("%s%s %-12s %s", cursor, marker, row.Label, state) + 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 • alt+p cycle • 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..cc055da --- /dev/null +++ b/internal/adapters/in/tui/provider_picker_test.go @@ -0,0 +1,129 @@ +package tui + +import ( + "context" + "testing" + + tea "charm.land/bubbletea/v2" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + + "ero/internal/core" + "ero/internal/ports" +) + +func TestProviderPickerDisplaysDescriptorRowsWithoutStartingInactiveProviders(t *testing.T) { + 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) + + 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, "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") + require.Contains(t, view, "gl-plugin local") + controller.AssertExpectations(t) +} + +func TestProviderPickerSelectEmitsSwitchCommandWithStableKey(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) + 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) + controller.AssertCalled(t, "Switch", mock.Anything, mock.Anything, "gitlab") + require.Equal(t, "gitlab", m.activeProviderKey) + 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() + 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) + + updated, cmd := m.Update(tea.KeyPressMsg{Text: "p", Code: 'p', Mod: tea.ModAlt}) + m = updated.(Model) + require.NotNil(t, cmd) + updated, _ = m.Update(cmd()) + m = updated.(Model) + controller.AssertCalled(t, "Switch", mock.Anything, mock.Anything, "gitlab") + + _, cmd = m.Update(keyPress("r")) + require.NotNil(t, cmd) + _ = cmd() + controller.AssertCalled(t, "Refresh", mock.Anything, mock.Anything, true) + controller.AssertExpectations(t) +} 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 9a57589..313d360 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,120 @@ func (m Model) closeReviewProvidersCmd() tea.Cmd { } } +func (m Model) startActiveProviderCmd() tea.Cmd { + activeProvider := m.activeProvider + if activeProvider == nil || !m.activeProviderSyncEnabled() { + return nil + } + ctx := m.ctx + reviewContext := m.reviewContext + 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 + } + return activeProviderStartedMsg{catalog: catalog, state: state, err: err} + } +} + +func (m Model) refreshActiveProviderCmd(manual bool) tea.Cmd { + activeProvider := m.activeProvider + if activeProvider == nil || !m.activeProviderSyncEnabled() { + 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 || !m.activeProviderSyncEnabled() { + 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.activeProviderSyncEnabled() || 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 || !m.activeProviderSyncEnabled() { + 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) canSwitchProvider() bool { + 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) { + 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.activeProviderKey = "" + 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..b286d8e 100644 --- a/internal/adapters/in/tui/review_publish.go +++ b/internal/adapters/in/tui/review_publish.go @@ -254,6 +254,12 @@ 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 { + result = append(result, providerClientWithInfo{info: info, client: m.activeProvider}) + delete(selected, m.activeRuntimeInfo.ID) + } + } for _, client := range m.reviewProviders { providerInfo, ok := m.providerInfoByClient[client] if !ok { diff --git a/internal/adapters/in/tui/review_publish_test.go b/internal/adapters/in/tui/review_publish_test.go index cbea3b7..bfc53d3 100644 --- a/internal/adapters/in/tui/review_publish_test.go +++ b/internal/adapters/in/tui/review_publish_test.go @@ -1,25 +1,25 @@ 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")) + updated, cmd := m.Update(tea.KeyPressMsg{Text: "P", Code: 'p', Mod: tea.ModShift}) m = updated.(Model) require.Nil(t, cmd) @@ -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/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") } diff --git a/internal/adapters/in/tui/theme/styles.go b/internal/adapters/in/tui/theme/styles.go index 70a3933..4a9857c 100644 --- a/internal/adapters/in/tui/theme/styles.go +++ b/internal/adapters/in/tui/theme/styles.go @@ -2,33 +2,48 @@ package theme import "charm.land/lipgloss/v2" +const ( + ColorText = "#c9d1d9" + ColorMutedText = "#8b949e" + ColorAccent = "#58a6ff" + ColorWarning = "#ffa657" + ColorKeyword = "#ff7b72" + ColorFunction = "#d2a8ff" + ColorType = "#ffa657" + 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")) 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/client.go b/internal/adapters/out/plugin/client.go index 2248238..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{ @@ -277,7 +300,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 +337,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 { @@ -325,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, @@ -479,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 dc4aa7d..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}, }, @@ -90,6 +91,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{ { @@ -104,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{ @@ -240,6 +275,56 @@ 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) + + _, 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/plugin/provider_loader.go b/internal/adapters/out/plugin/provider_loader.go index d2ab7cd..de68bf1 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,73 +20,196 @@ 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") - continue - } - command, args := splitRuntimeCommand(manifest.Runtime.Command) - if command == "" { - log.Warn().Str("plugin_path", descriptor.Path).Msg("plugin runtime command is empty") + log.Warn().Err(err).Str("plugin_path", descriptor.PluginPath).Str("contribution_id", descriptor.ContributionID).Msg("create plugin review provider client failed") continue } - 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") + 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 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; using existing runtime") } - 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 } - 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) + } + 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) } } - return providers, nil + if descriptor.Path != "" { + if abs, err := filepath.Abs(descriptor.Path); err == nil { + return "path:" + filepath.Clean(abs) + } + return "path:" + filepath.Clean(descriptor.Path) + } + return strings.Join([]string{"manifest", descriptor.Name, descriptor.Version, descriptor.Source}, ":") } -func runtimeCommandAvailable(command, pluginDir string) bool { +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 8331744..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" @@ -12,6 +13,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() @@ -53,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() diff --git a/internal/adapters/out/providercache/cache.go b/internal/adapters/out/providercache/cache.go new file mode 100644 index 0000000..03e263e --- /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 func() { _ = 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..024d9f8 --- /dev/null +++ b/internal/adapters/out/providercache/cache_test.go @@ -0,0 +1,53 @@ +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 { + 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)) + } +} + +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..50100c5 --- /dev/null +++ b/internal/app/active_provider_service.go @@ -0,0 +1,497 @@ +package app + +import ( + "context" + "strings" + "sync" + "time" + + "github.com/bnema/zerowrap" + + "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 + RuntimeInfo core.ReviewProviderInfo + 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) { + 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 + } + 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() + if err := s.closeLocked(); err != nil { + log.Warn().Err(err).Msg("active provider close before start failed") + } + s.stableKey = "" + s.runtimeID = "" + s.generation++ + startGen := s.generation + s.state = ActiveProviderState{} + 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() + 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") + } + 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 && (!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") + 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) + if !s.setState(startGen, failed) { + return s.State(), nil + } + 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() + 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++ + switchGen := s.generation + s.state = ActiveProviderState{} + 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) + if !s.setState(switchGen, failed) { + return s.State(), nil + } + 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") + } + 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 { + 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") + if !s.setState(gen, st) { + return s.State(), nil + } + return st, nil + } + } + log.Warn().Str("provider_key", stableKey).Msg("active provider switch descriptor not found") + 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) { + 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) { + log := zerowrap.FromCtx(ctx) + 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 { + 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.(ports.ReviewProviderSnapshotClient); 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 + } + if err != nil { + if st, stale := s.stateIfGenerationChanged(gen); stale { + return st, 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 + 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") + } + 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 + } + 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 { + 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} + if !s.setState(gen, st) { + return s.State(), nil + } + 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() + 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 { + out := make([]ports.ReviewProviderDescriptor, 0, len(descs)) + used := map[string]bool{} + preferenceFound := false + if s.prefs != nil { + if key, ok := s.preferredProviderKey(ctx, review); 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, 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 { + contextKey := core.NewReviewContextKey(key, review) + 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 + 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 +} +func (s *ActiveProviderService) stateIfGenerationChanged(gen int64) (ActiveProviderState, bool) { + s.mu.Lock() + defer s.mu.Unlock() + if gen == s.generation { + 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() + 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 (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 { + 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..4797df7 --- /dev/null +++ b/internal/app/active_provider_service_test.go @@ -0,0 +1,385 @@ +package app + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/mock" + + "ero/internal/core" + "ero/internal/ports" + "ero/internal/ports/mocks" +) + +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 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 := 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) + } + if st.StableProviderKey != "github" { + t.Fatalf("got %q", st.StableProviderKey) + } +} + +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{}) + + 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) + cached := core.ProviderSnapshot{StableProviderKey: "github", ContextKey: key, Threads: []core.RemoteReviewThread{{ExternalID: "old"}}} + cache := &memCache{snap: cached, ok: true} + 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) + } + 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 := 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" + if st.StableProviderKey == "b" { + t.Fatal("expected deterministic initial provider a") + } + 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 !aClosed { + t.Fatalf("switch should close old client") + } +} + +func TestActiveProviderServiceSwitchClosesCurrentBeforeStartingTarget(t *testing.T) { + review := testReviewContext() + 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) + } +} + +func TestActiveProviderServiceFailedSwitchClearsOldProviderState(t *testing.T) { + review := testReviewContext() + 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) + } + 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) + } +} + +func TestActiveProviderServiceUsesStableCatalogOrder(t *testing.T) { + review := testReviewContext() + 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 := made; len(got) != 1 || got[0] != "github" { + t.Fatalf("github fallback should be selected in catalog order, got %v", got) + } + + 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 := 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 598e96a..ea14f02 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,15 +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) - providers, err := buildReviewProviders(ctx, 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 @@ -117,7 +114,8 @@ 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)) + 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/app_test.go b/internal/app/app_test.go index 71d7c54..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"}, } @@ -246,27 +246,32 @@ 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 := 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, loader) + providers, err := buildReviewProviders(ctx, catalog, factory) require.NoError(t, err) require.Equal(t, []ports.ReviewProviderClient{provider}, providers) } 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", @@ -298,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) @@ -315,29 +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 fakeStartupPrompt struct { mode core.DiffMode err error diff --git a/internal/app/provider_polling_config.go b/internal/app/provider_polling_config.go new file mode 100644 index 0000000..6853411 --- /dev/null +++ b/internal/app/provider_polling_config.go @@ -0,0 +1,33 @@ +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 + } + if poll.MinBackoff > poll.MaxBackoff { + poll.MinBackoff = poll.MaxBackoff + } + return poll +} diff --git a/internal/app/review_providers.go b/internal/app/review_providers.go index f9f85b3..3d0d173 100644 --- a/internal/app/review_providers.go +++ b/internal/app/review_providers.go @@ -2,13 +2,35 @@ package app import ( "context" + "fmt" + + "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)) + 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 new file mode 100644 index 0000000..34978ab --- /dev/null +++ b/internal/app/tui_active_provider.go @@ -0,0 +1,86 @@ +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, core.NewProviderError(core.ProviderErrorNotApplicable, "no active provider catalog", nil) + } + return c.catalog.ListReviewProviderDescriptors(ctx) +} + +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 +} + +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) { + 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 +} + +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, + } +} 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..a9d7a97 --- /dev/null +++ b/internal/core/provider_snapshot.go @@ -0,0 +1,149 @@ +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 { + 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 + +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..f5d36ba --- /dev/null +++ b/internal/core/provider_snapshot_test.go @@ -0,0 +1,73 @@ +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/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/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/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..e837d77 100644 --- a/internal/ports/plugin.go +++ b/internal/ports/plugin.go @@ -52,6 +52,28 @@ 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 { LoadReviewProviders(ctx context.Context) ([]ReviewProviderClient, error) @@ -82,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/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 +} 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..4821b28 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,66 @@ 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 { + 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) + } +} + +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() 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..3bb2dcf --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/graphql.go @@ -0,0 +1,321 @@ +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:"diffSide"` + StartLine int `json:"startLine"` + StartSide string `json:"startDiffSide"` + 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 + diffSide + startLine + startDiffSide + 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 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 { + lastErr = err + continue + } + 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") + } + if len(matches) > 1 { + return githubRemote{}, githubPRCandidate{}, plugin.NewErrorf(plugin.ErrorNotApplicable, "ambiguous GitHub pull request match across %d remotes", len(matches)) + } + return matches[0].remote, matches[0].pr, nil +} + +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 { + 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) + } + } + 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..0efb13d --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/graphql_test.go @@ -0,0 +1,209 @@ +package main + +import ( + "context" + "maps" + "os" + "strconv" + "strings" + "testing" + "time" + + "ero/pkg/plugin" +) + +type fakeGraphQLClient struct { + listPages []ghPRListResponse + listErrs []error + snapshotPages []ghPRSnapshotResponse + snapshotCalls int + listCalls 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: + 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: + *r = f.snapshotPages[f.snapshotCalls] + f.snapshotCalls++ + default: + panic("unexpected GraphQL response type") + } + _ = query + 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"} + 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 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"} + 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 TestLoadRemoteSnapshotKeepsPartialThreadWhenNestedCommentsArePaginated(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 }} + 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) + } +} + +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..871d0af 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,58 @@ 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) { + 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 + } + client, err := p.graphQLClient() + if err != nil { + return plugin.DetectContextResult{}, plugin.NewErrorf(plugin.ErrorAuthRequired, "create GitHub GraphQL client: %v", err) + } + _, match, err := matchGitHubPRAcrossRemotes(ctx, client, remotes, req.Context) + if err != 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 +} + +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") } - return plugin.DetectContextResult{Result: plugin.DetectionResult{Applicable: false, Reason: "no GitHub remote detected"}}, nil + client, err := p.graphQLClient() + if err != nil { + return plugin.LoadRemoteSnapshotResult{}, plugin.NewErrorf(plugin.ErrorAuthRequired, "create GitHub GraphQL client: %v", err) + } + 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 + } + 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) { + snapshotReq := plugin.LoadRemoteSnapshotRequest(req) + snapshot, err := p.LoadRemoteSnapshot(ctx, snapshotReq) + 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 +129,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 +143,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 +160,31 @@ 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 remotes := githubRemotes(reviewCtx.Repository.Remotes); len(remotes) > 0 { + if client, err := p.graphQLClient(); err == nil { + _, match, err := matchGitHubPRAcrossRemotes(ctx, client, remotes, reviewCtx) + if err == nil { + return ghPR{Number: match.Number, URL: match.URL}, nil + } + if pe := plugin.AsError(err); pe == nil || pe.Code != plugin.ErrorNotApplicable { + return ghPR{}, err + } + } + } 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 +196,49 @@ 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 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") { + 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 +333,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..14709f8 100644 --- a/plugins/github/cmd/ero-plugin-github/main_test.go +++ b/plugins/github/cmd/ero-plugin-github/main_test.go @@ -9,14 +9,194 @@ 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 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"}} + 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: "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: "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"}}, + 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 +235,96 @@ 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 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"}} + 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"}} + 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..4f21361 --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/map.go @@ -0,0 +1,114 @@ +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 { + // 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 { + thread.Range.Start = lineRef(t.StartLine, firstNonEmpty(t.StartSide, t.Side)) + } 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 + } + 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..079d857 --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/match.go @@ -0,0 +1,119 @@ +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 nonBranchHeadRef(headRef) { + 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 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 == "" { + 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) +} + +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..62bd70e --- /dev/null +++ b/plugins/github/cmd/ero-plugin-github/remote.go @@ -0,0 +1,69 @@ +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 githubRemotes(remotes []plugin.GitRemote) []githubRemote { + out := make([]githubRemote, 0, len(remotes)) + seen := map[string]bool{} + for _, remote := range remotes { + 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 out +}