diff --git a/.gitignore b/.gitignore index 6e94bfb9..0f4ce00b 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,5 @@ # Binaries -/chief +/bin/ *.exe # Log files @@ -11,3 +11,6 @@ node_modules/ # VitePress docs/.vitepress/cache/ docs/.vitepress/dist/ + +# Contract fixtures (synced from chief-uplink) +contract/fixtures/ diff --git a/.goreleaser.yaml b/.goreleaser.yaml index a28af181..280cc570 100644 --- a/.goreleaser.yaml +++ b/.goreleaser.yaml @@ -87,7 +87,7 @@ brews: name: homebrew-chief token: "{{ .Env.HOMEBREW_TAP_GITHUB_TOKEN }}" directory: Formula - homepage: "https://chiefloop.com" + homepage: "https://minicodemonkey.github.io/chief/" description: "Autonomous agent loop for working through PRDs with Claude Code" license: "MIT" # Custom install script diff --git a/Makefile b/Makefile index 9c005d48..d5467623 100644 --- a/Makefile +++ b/Makefile @@ -1,7 +1,7 @@ # Chief - Autonomous PRD Agent # https://github.com/minicodemonkey/chief -BINARY_NAME := chief +BINARY_NAME := bin/chief VERSION := $(shell git describe --tags --always --dirty 2>/dev/null || echo "dev") BIN_DIR := ./bin BUILD_DIR := ./build @@ -10,14 +10,14 @@ MAIN_PKG := ./cmd/chief # Go build flags LDFLAGS := -ldflags "-X main.Version=$(VERSION)" -.PHONY: all build install test lint clean release snapshot help +.PHONY: all build install test lint clean release snapshot help sync-fixtures test-contract all: build ## build: Build the binary build: @mkdir -p $(BIN_DIR) - go build $(LDFLAGS) -o $(BIN_DIR)/$(BINARY_NAME) $(MAIN_PKG) + go build $(LDFLAGS) -o $(BINARY_NAME) $(MAIN_PKG) ## install: Install to $GOPATH/bin install: @@ -64,7 +64,34 @@ release: ## run: Build and run the TUI run: build - $(BIN_DIR)/$(BINARY_NAME) + $(BINARY_NAME) + +## Contract fixtures — chief-uplink is the source of truth. +## Override FIXTURES_REPO for local dev: make sync-fixtures FIXTURES_REPO=../chief-uplink/contract/fixtures +FIXTURES_REPO ?= https://raw.githubusercontent.com/MiniCodeMonkey/chief-uplink/main/contract/fixtures +FIXTURES_DIR := contract/fixtures + +## sync-fixtures: Download contract fixtures from chief-uplink +sync-fixtures: + @mkdir -p $(FIXTURES_DIR)/cli-to-server $(FIXTURES_DIR)/server-to-cli + @for f in cli-to-server/connect_request.json cli-to-server/state_snapshot.json \ + cli-to-server/messages_batch.json cli-to-server/prds_response.json \ + cli-to-server/settings_response.json cli-to-server/diffs_response.json \ + server-to-cli/welcome_response.json server-to-cli/command_create_project.json \ + server-to-cli/command_list_projects.json server-to-cli/command_start_run.json \ + server-to-cli/command_get_prds.json server-to-cli/command_get_settings.json \ + server-to-cli/command_get_diffs.json server-to-cli/command_new_prd.json \ + server-to-cli/command_refine_prd.json server-to-cli/command_prd_message.json; do \ + if echo "$(FIXTURES_REPO)" | grep -q "^http"; then \ + curl -sf "$(FIXTURES_REPO)/$$f" -o "$(FIXTURES_DIR)/$$f" || echo "WARN: failed to fetch $$f"; \ + else \ + cp "$(FIXTURES_REPO)/$$f" "$(FIXTURES_DIR)/$$f" || echo "WARN: failed to copy $$f"; \ + fi; \ + done + +## test-contract: Run contract tests (syncs fixtures first) +test-contract: sync-fixtures + go test ./internal/contract/ -v ## help: Show this help help: diff --git a/README.md b/README.md index 0d6b8d59..61c474bd 100644 --- a/README.md +++ b/README.md @@ -1,9 +1,5 @@ # Chief -

- Chief -

- Build big projects with Claude. Chief breaks your work into tasks and runs Claude Code in a loop until they're done. **[Documentation](https://minicodemonkey.github.io/chief/)** · **[Quick Start](https://minicodemonkey.github.io/chief/guide/quick-start)** @@ -44,17 +40,7 @@ See the [documentation](https://minicodemonkey.github.io/chief/concepts/how-it-w ## Requirements -- **[Claude Code CLI](https://docs.anthropic.com/en/docs/claude-code)**, **[Codex CLI](https://developers.openai.com/codex/cli/reference)**, or **[OpenCode CLI](https://opencode.ai)** installed and authenticated - -Use Claude by default, or configure Codex or OpenCode in `.chief/config.yaml`: - -```yaml -agent: - provider: opencode - cliPath: /usr/local/bin/opencode # optional -``` - -Or run with `chief --agent opencode` or set `CHIEF_AGENT=opencode`. +- [Claude Code CLI](https://docs.anthropic.com/en/docs/claude-code) installed and authenticated ## License @@ -62,8 +48,6 @@ MIT ## Acknowledgments -- [@Simon-BEE](https://github.com/Simon-BEE) — Multi-agent architecture and Codex CLI integration -- [@tpaulshippy](https://github.com/tpaulshippy) — OpenCode CLI support and NDJSON parser - [snarktank/ralph](https://github.com/snarktank/ralph) — The original Ralph implementation that inspired this project - [Geoffrey Huntley](https://ghuntley.com/ralph/) — For coining the "Ralph Wiggum loop" pattern - [Bubble Tea](https://github.com/charmbracelet/bubbletea) — TUI framework diff --git a/WEBSOCKET_REFACTOR.md b/WEBSOCKET_REFACTOR.md new file mode 100644 index 00000000..e0d78a03 --- /dev/null +++ b/WEBSOCKET_REFACTOR.md @@ -0,0 +1,511 @@ +# WebSocket Refactor: Managed Reverb Compatibility + +## Problem + +Laravel Cloud's managed Reverb is a separate service that only runs the standard Pusher protocol routes (`/app/{appKey}`, `/up`, etc.). Our custom `/ws/server` endpoint (registered via `ChiefReverbFactory`) never gets loaded because it lives in our application code, not in the managed Reverb cluster. This means the chief CLI can't connect. + +## Solution + +Keep managed Reverb for browser broadcasting (it already works). Adapt the chief CLI to: + +1. **Send data to the server via HTTP POST** (replaces WebSocket sends) +2. **Receive commands via Reverb's Pusher protocol** (subscribes to a private channel as a standard Pusher client) + +No infrastructure changes required — managed Reverb stays as-is. + +## Architecture: Before vs After + +### Before (custom WebSocket) + +``` +Chief CLI ──WebSocket /ws/server──→ ChiefServerController ──broadcast──→ Browser + ↑ │ + └──── WebSocket (sendToDevice) ←── CommandRelayController ←── HTTP POST─┘ +``` + +- Single persistent WebSocket for all bidirectional communication +- Custom hello/welcome handshake for auth +- In-memory connection tracking (ServerConnectionManager) +- ChiefReverbFactory adds custom route to Reverb server + +### After (HTTP + Pusher channel) + +``` +Chief CLI ──HTTP POST /api/device/messages──→ MessageIngestionController ──broadcast──→ Browser + ↑ │ + └──── Reverb private channel (Pusher protocol) ←── broadcast ←── CommandRelayController ←─┘ +``` + +- CLI sends data via HTTP POST (batched) +- CLI receives commands by subscribing to `private-chief-server.{deviceId}` on managed Reverb +- Auth via OAuth access token (existing) on both HTTP and channel subscription +- No custom Reverb routes, no in-memory connection tracking + +--- + +## Detailed Changes + +### Phase 1: New HTTP ingestion endpoint (chief-uplink) + +Create a new HTTP API endpoint that accepts messages from chief CLIs — the replacement for the WebSocket receive path. + +#### 1.1 New controller: `MessageIngestionController` + +**File:** `app/Http/Controllers/Api/MessageIngestionController.php` + +Accepts batched messages from the CLI via HTTP POST. Replaces `ChiefServerController::handleMessage()`. + +``` +POST /api/device/messages +Authorization: Bearer {access_token} +Content-Type: application/json + +{ + "messages": [ + {"type": "state_snapshot", "id": "...", "timestamp": "...", ...}, + {"type": "claude_output", "id": "...", "timestamp": "...", ...}, + ... + ] +} +``` + +Responsibilities: +- Validate OAuth access token (reuse `DeviceOAuthController::validateAccessToken()`) +- Check device not revoked +- For each message: + - If `project_state` → update `CachedProjectState` (same as current `handleProjectState()`) + - If bufferable → buffer via `WebSocketMessageBuffer` + - Broadcast to browser via `ChiefMessageReceived` event +- Return acknowledgment with any pending server-side state (e.g., new session_id) + +Response: +```json +{ + "accepted": 5, + "session_id": "uuid" +} +``` + +#### 1.2 New controller: `DevicePresenceController` + +**File:** `app/Http/Controllers/Api/DevicePresenceController.php` + +Handles explicit connect/disconnect lifecycle — replaces the WebSocket open/close events. + +``` +POST /api/device/connect +Authorization: Bearer {access_token} +Content-Type: application/json + +{ + "chief_version": "0.5.0", + "device_name": "sierra", + "os": "darwin", + "arch": "arm64", + "protocol_version": 1 +} +``` + +Responsibilities: +- Validate access token, check device not revoked +- Update device metadata (chief_version, os, arch, device_name, is_online, last_connected_at) +- Generate session_id for message buffering +- Mark device reconnected in buffer +- Dispatch `DeviceConnected` event +- Return welcome response (same fields as current WebSocket welcome) + +```json +{ + "type": "welcome", + "protocol_version": 1, + "device_id": 42, + "session_id": "uuid", + "reverb": { + "key": "app-key", + "host": "ws-xxx-reverb.laravel.cloud", + "port": 443, + "scheme": "https" + } +} +``` + +The `reverb` block tells the CLI where to connect as a Pusher client. + +``` +POST /api/device/disconnect +Authorization: Bearer {access_token} +``` + +Responsibilities: +- Mark device offline, dispatch `DeviceDisconnected` +- Start buffer grace period + +#### 1.3 New middleware: `AuthenticateDevice` + +**File:** `app/Http/Middleware/AuthenticateDevice.php` + +Extracts and validates the OAuth access token from the `Authorization: Bearer` header. Sets `$request->attributes->set('device_id', ...)` and `$request->attributes->set('user_id', ...)` for downstream controllers. Reuses `DeviceOAuthController::validateAccessToken()`. + +Apply to all `/api/device/*` routes. + +#### 1.4 Channel auth for CLI devices + +**File:** `routes/channels.php` + +Add authorization for the new CLI channel: + +```php +Broadcast::channel('chief-server.{deviceId}', function ($user, $deviceId) { + // CLI authenticates channel subscription using its OAuth token. + // The token is passed as the auth token in the Pusher subscription. + // We need a custom auth endpoint for this — see 1.5. + return $user->deviceAuthorizations() + ->where('id', $deviceId) + ->whereNull('revoked_at') + ->exists(); +}); +``` + +#### 1.5 Custom Pusher auth endpoint for CLI + +The chief CLI authenticates via OAuth access tokens, not Laravel sessions. We need a custom broadcasting auth endpoint that accepts Bearer tokens. + +**File:** `app/Http/Controllers/Api/DeviceBroadcastAuthController.php` + +``` +POST /api/device/broadcasting/auth +Authorization: Bearer {access_token} +Content-Type: application/json + +{"socket_id": "...", "channel_name": "private-chief-server.42"} +``` + +This endpoint: +- Validates the access token +- Checks the device owns the requested channel +- Returns the Pusher auth signature (same format as Laravel's standard broadcast auth) + +#### 1.6 Refactor `CommandRelayController` + +**File:** `app/Http/Controllers/Api/CommandRelayController.php` + +Currently calls `$this->connectionManager->sendToDevice()` which sends via in-memory WebSocket. Change to broadcast the command via Reverb to the CLI's channel: + +```php +// Before: +$sent = $this->connectionManager->sendToDevice($deviceId, $message); + +// After: +broadcast(new ChiefCommandDispatched($deviceId, $userId, $message)); +``` + +New event: `ChiefCommandDispatched` broadcasts on `private-chief-server.{deviceId}` with event name `chief.command`. + +The `isDeviceOnline` check changes from in-memory lookup to checking `DeviceAuthorization.is_online` in the database. + +#### 1.7 Refactor `ServerConnectionManager` + +The in-memory connection tracking (`$connections`, `$deviceToConnection`, `$connectionObjects`) is no longer needed. The class simplifies to a stateless service that: + +- Delegates to `WebSocketMessageBuffer` for buffering +- Checks device online status via database +- No longer stores Connection objects or session IDs in memory (session IDs stored in Redis or on the DeviceAuthorization model) + +Alternatively, this class can be removed entirely and its responsibilities distributed to the new controllers. + +#### 1.8 New event: `ChiefCommandDispatched` + +**File:** `app/Events/ChiefCommandDispatched.php` + +```php +class ChiefCommandDispatched implements ShouldBroadcast +{ + public function __construct( + public readonly int $deviceId, + public readonly int $userId, + public readonly array $command, + ) {} + + public function broadcastOn(): array + { + return [new Channel("private-chief-server.{$this->deviceId}")]; + } + + public function broadcastAs(): string + { + return 'chief.command'; + } + + public function broadcastWith(): array + { + return $this->command; + } +} +``` + +#### 1.9 New route: `DeviceHeartbeatController` + +The CLI needs to periodically confirm it's still alive, since we no longer have a persistent WebSocket connection to detect disconnects. + +``` +POST /api/device/heartbeat +Authorization: Bearer {access_token} +``` + +Called every 30-60 seconds by the CLI. Updates `last_heartbeat_at` on the device. A scheduled job marks devices as offline if no heartbeat received within 2 minutes. + +--- + +### Phase 2: CLI changes (chief — Go) + +#### 2.1 New HTTP client: `internal/uplink/client.go` + +Replaces the WebSocket client for sending data to the server. Handles: + +- `POST /api/device/connect` — on startup +- `POST /api/device/messages` — batched message sending +- `POST /api/device/heartbeat` — periodic keepalive +- `POST /api/device/disconnect` — on shutdown +- OAuth token refresh (existing logic, but applied to HTTP headers) +- Retry with exponential backoff on failure + +#### 2.2 Message batcher: `internal/uplink/batcher.go` + +Batches outgoing messages to reduce HTTP request volume. Key design: + +- Collects messages in a buffer +- Flushes when: buffer reaches N messages (e.g., 20), OR time threshold elapsed (e.g., 200ms), OR a priority message arrives (e.g., `run_complete`) +- Each flush sends a single `POST /api/device/messages` with the batch +- Priority messages (user-visible state changes) flush immediately +- Streaming messages (`claude_output`) batch on the 200ms timer + +Categories: +- **Immediate flush:** `run_complete`, `run_paused`, `error`, `clone_complete`, `session_expired`, `quota_exhausted` +- **Batched (200ms):** `claude_output`, `prd_output`, `run_progress`, `clone_progress` +- **Low priority (1s):** `state_snapshot`, `project_state`, `project_list`, `settings`, `log_lines` + +#### 2.3 Pusher client: `internal/uplink/pusher.go` + +Subscribes to `private-chief-server.{deviceId}` on managed Reverb to receive commands from the browser. Uses a Go Pusher client library (e.g., `pusher/pusher-websocket-go` or a lightweight implementation). + +Responsibilities: +- Connect to Reverb as a Pusher client using the app key and host from the connect response +- Authenticate the private channel via `POST /api/device/broadcasting/auth` +- Listen for `chief.command` events +- Route received commands to the existing message dispatcher + +The existing `Dispatcher` pattern (type-based routing) stays the same — commands arrive through a different transport but are dispatched identically. + +#### 2.4 Refactor `internal/cmd/serve.go` + +Replace `ws.Client` usage with the new uplink client: + +```go +// Before: +client = ws.New(wsURL, ws.WithOnReconnect(func() { ... })) +client.Connect(ctx) +client.Handshake(creds.AccessToken, version, deviceName) +// ... main loop reads from client.Receive() + +// After: +uplink := uplink.New(baseURL, creds.AccessToken, uplink.WithOnReconnect(func() { ... })) +uplink.Connect(ctx) // POST /api/device/connect + subscribe to Pusher channel +// ... main loop reads from uplink.Receive() (same interface, different transport) +``` + +The `Send()` method now enqueues into the batcher instead of writing to WebSocket. The `Receive()` channel is fed by the Pusher client instead of WebSocket reads. + +#### 2.5 Handshake changes + +The current WebSocket handshake (hello → welcome) becomes an HTTP request: + +```go +// Before: +client.Handshake(accessToken, version, deviceName) + +// After: +welcome, err := httpClient.Connect(ConnectRequest{ + ChiefVersion: version, + DeviceName: deviceName, + OS: runtime.GOOS, + Arch: runtime.GOARCH, +}) +// welcome contains device_id, session_id, reverb config +``` + +#### 2.6 Reconnection logic + +Two reconnection paths: + +1. **HTTP failures:** Retry with exponential backoff (same as current WebSocket retry). If the server is unreachable, buffer messages locally and flush on reconnect. +2. **Pusher disconnection:** The Pusher client library handles reconnection automatically. On reconnect, re-authenticate the channel. + +On any reconnection, re-send state snapshot via HTTP POST (same as current `onRecon` callback). + +#### 2.7 Heartbeat + +Add a periodic heartbeat goroutine that calls `POST /api/device/heartbeat` every 30 seconds. If the heartbeat fails, trigger reconnection logic. + +#### 2.8 Graceful shutdown + +On SIGTERM/SIGINT: +1. Stop accepting new commands from Pusher +2. Flush message batcher +3. Call `POST /api/device/disconnect` +4. Disconnect Pusher client +5. Kill Claude sessions and runs (existing logic) + +--- + +### Phase 3: Remove old WebSocket infrastructure (chief-uplink) + +After the new system is working: + +#### 3.1 Delete files +- `app/WebSocket/ChiefReverbFactory.php` +- `app/WebSocket/ChiefServerController.php` +- `app/Console/Commands/StartReverbServer.php` + +#### 3.2 Simplify `WebSocketServiceProvider` +- Remove the `StartServer::class` → `StartReverbServer::class` container binding +- Remove `ServerConnectionManager` singleton if fully replaced +- Keep `PrdSessionManager` singleton + +#### 3.3 Remove Reverb dependency from `ServerConnectionManager` +- Remove `use Laravel\Reverb\Servers\Reverb\Connection;` +- Remove `$connectionObjects` array and all methods that reference it +- Or delete the class entirely if all responsibilities moved to new controllers + +#### 3.4 Clean up tests +- Update `tests/Feature/WebSocket/MessageRelayTest.php` → test HTTP endpoints +- Add tests for `MessageIngestionController`, `DevicePresenceController`, `DeviceBroadcastAuthController` + +--- + +### Phase 4: Frontend changes (chief-uplink — minimal) + +The browser-side code requires almost no changes. + +#### 4.1 `CommandRelayController` response + +The `isDeviceOnline` check changes from in-memory to database, but the HTTP response contract stays the same. Frontend `useCommandRelay.ts` is unchanged. + +#### 4.2 `useChiefMessages.ts` — unchanged + +Still subscribes to `private-device.{deviceId}` and listens for `chief.message` events. The messages arrive via the same Reverb channel — only the server-side path that triggers the broadcast changes (HTTP controller instead of WebSocket controller). + +#### 4.3 `useEcho.ts` — unchanged + +No changes to Echo setup or connection management. + +#### 4.4 `echo.ts` — unchanged + +Still uses `broadcaster: 'reverb'` with managed Reverb config. + +--- + +## Message Flow Comparison + +### CLI → Browser (e.g., `claude_output` streaming) + +**Before:** +1. CLI sends `claude_output` JSON over WebSocket +2. `ChiefServerController::handleMessage()` receives it +3. Buffers message via `WebSocketMessageBuffer` +4. Dispatches `ChiefMessageReceived` broadcast event +5. Reverb delivers to browser on `private-device.{deviceId}` + +**After:** +1. CLI enqueues `claude_output` into message batcher +2. Batcher flushes batch via `POST /api/device/messages` +3. `MessageIngestionController` receives batch +4. For each message: buffer + dispatch `ChiefMessageReceived` +5. Reverb delivers to browser on `private-device.{deviceId}` (same as before) + +### Browser → CLI (e.g., `start_run` command) + +**Before:** +1. Browser calls `POST /ws/command/{deviceId}` with `{type: "start_run", payload: {...}}` +2. `CommandRelayController::send()` validates request +3. Calls `ServerConnectionManager::sendToDevice()` → writes to in-memory WebSocket connection +4. CLI receives `start_run` in `readLoop()` → dispatches to handler + +**After:** +1. Browser calls `POST /ws/command/{deviceId}` with `{type: "start_run", payload: {...}}` (same) +2. `CommandRelayController::send()` validates request (same) +3. Dispatches `ChiefCommandDispatched` broadcast event on `private-chief-server.{deviceId}` +4. Reverb delivers to CLI's Pusher subscription → dispatches to handler + +--- + +## Device Lifecycle + +### Connect + +1. CLI calls `POST /api/device/connect` with metadata + access token +2. Server validates token, updates device record, generates session_id +3. Server dispatches `DeviceConnected` event to browser +4. Server returns welcome response with Reverb config +5. CLI connects to Reverb as Pusher client, subscribes to `private-chief-server.{deviceId}` +6. CLI sends initial `state_snapshot` via `POST /api/device/messages` + +### Steady state + +- CLI sends messages in batches via `POST /api/device/messages` (every 200ms or on priority) +- CLI sends heartbeat via `POST /api/device/heartbeat` (every 30s) +- CLI receives commands via Pusher channel subscription +- Browser sends commands via `POST /ws/command/{deviceId}` (unchanged) +- Browser receives messages via `private-device.{deviceId}` channel (unchanged) + +### Disconnect + +**Graceful (CLI shutdown):** +1. CLI calls `POST /api/device/disconnect` +2. Server marks device offline, starts buffer grace period +3. Server dispatches `DeviceDisconnected` event + +**Ungraceful (network failure, crash):** +1. Heartbeat stops arriving +2. Scheduled job detects stale heartbeat (>2 min) +3. Marks device offline, starts buffer grace period +4. Dispatches `DeviceDisconnected` event + +--- + +## Risks and Mitigations + +### Latency from HTTP batching +**Risk:** 200ms batch window adds latency to `claude_output` streaming. +**Mitigation:** 200ms is imperceptible for terminal-like output. Priority messages flush immediately. Can tune batch window down to 100ms if needed. + +### HTTP overhead vs WebSocket +**Risk:** More HTTP requests than a single persistent connection. +**Mitigation:** Batching reduces request count significantly. A typical streaming session generates ~5 HTTP requests/second (vs thousands of individual WebSocket frames). HTTP/2 connection reuse minimizes TCP overhead. + +### Pusher message size limits +**Risk:** Managed Reverb (Pusher protocol) may have message size limits for channel events. +**Mitigation:** Commands from browser → CLI are small (typically <1KB). The large data flow (CLI → server) goes via HTTP, not Pusher channels. Reverb's default max message size is 10KB, and commands never approach this. + +### Heartbeat-based disconnect detection +**Risk:** Up to 2 minutes to detect a crashed CLI (vs instant WebSocket close detection). +**Mitigation:** Acceptable for the use case — the browser already debounces disconnect events by 2 seconds. The "offline" indicator updates within 2 minutes, which is fine. Can reduce heartbeat interval to 15s and detection to 45s if needed. + +### Authentication on Pusher channel +**Risk:** The CLI uses OAuth tokens, not Laravel sessions, for auth. Pusher channel auth requires a custom endpoint. +**Mitigation:** `DeviceBroadcastAuthController` provides a standard Pusher auth response using the CLI's OAuth token. The Pusher client library supports custom auth endpoints. + +--- + +## Implementation Order + +1. **Phase 1.1–1.3:** New HTTP endpoints + middleware (can deploy independently, no breaking changes) +2. **Phase 1.4–1.5:** Channel auth for CLI (deploy with Phase 1) +3. **Phase 1.8–1.9:** New event + heartbeat (deploy with Phase 1) +4. **Phase 2.1–2.3:** New Go uplink client, batcher, Pusher client +5. **Phase 2.4–2.8:** Refactor serve command to use new client +6. **Phase 1.6–1.7:** Refactor CommandRelayController to broadcast instead of direct WebSocket send +7. **Test end-to-end with both old and new CLI versions** +8. **Phase 3:** Remove old WebSocket infrastructure after confirming new system works +9. **Phase 4:** Any minor frontend adjustments + +Total estimate: ~15 new/modified files across both repos. diff --git a/cmd/chief/main.go b/cmd/chief/main.go index 016c188d..73b12ed1 100644 --- a/cmd/chief/main.go +++ b/cmd/chief/main.go @@ -4,17 +4,15 @@ import ( "fmt" "os" "path/filepath" - "strconv" "strings" tea "github.com/charmbracelet/bubbletea" - "github.com/minicodemonkey/chief/internal/agent" "github.com/minicodemonkey/chief/internal/cmd" "github.com/minicodemonkey/chief/internal/config" "github.com/minicodemonkey/chief/internal/git" - "github.com/minicodemonkey/chief/internal/loop" "github.com/minicodemonkey/chief/internal/prd" "github.com/minicodemonkey/chief/internal/tui" + "github.com/spf13/cobra" ) // Version is set at build time via ldflags @@ -28,340 +26,319 @@ type TUIOptions struct { Merge bool Force bool NoRetry bool - Agent string // --agent claude|codex|opencode|cursor - AgentPath string // --agent-path } func main() { - // Handle subcommands first - if len(os.Args) > 1 { - switch os.Args[1] { - case "new": - runNew() - return - case "edit": - runEdit() - return - case "status": - runStatus() - return - case "list": - runList() - return - case "help": - printHelp() - return - case "--help", "-h": - printHelp() - return - case "--version", "-v": - fmt.Printf("chief version %s\n", Version) - return - case "update": - runUpdate() - return - case "wiggum": - printWiggum() - return - } - } - - // Non-blocking version check on startup (for interactive TUI sessions) - cmd.CheckVersionOnStartup(Version) - - // Parse flags for TUI mode - opts := parseTUIFlags() - - // Handle special flags that were parsed - if opts == nil { - // Already handled (--help or --version) - return + rootCmd := buildRootCmd() + if err := rootCmd.Execute(); err != nil { + fmt.Fprintf(os.Stderr, "Error: %v\n", err) + os.Exit(1) } - - // Run the TUI - runTUIWithOptions(opts) } -// findAvailablePRD looks for any available PRD in .chief/prds/ -// Returns the path to the first PRD found, or empty string if none exist. -func findAvailablePRD() string { - prdsDir := ".chief/prds" - entries, err := os.ReadDir(prdsDir) - if err != nil { - return "" - } +func buildRootCmd() *cobra.Command { + opts := &TUIOptions{} - for _, entry := range entries { - if entry.IsDir() { - prdPath := filepath.Join(prdsDir, entry.Name(), "prd.md") - if _, err := os.Stat(prdPath); err == nil { - return prdPath + rootCmd := &cobra.Command{ + Use: "chief [name|path/to/prd.json]", + Short: "Chief - Autonomous PRD Agent", + Long: "Chief breaks down PRDs into user stories and uses Claude Code to implement them autonomously.", + // Accept arbitrary args so positional PRD name/path works + Args: cobra.ArbitraryArgs, + Version: Version, + // Silence Cobra's default error/usage printing so we control output + SilenceErrors: true, + SilenceUsage: true, + PersistentPreRun: func(c *cobra.Command, args []string) { + // Non-blocking version check on startup for all interactive commands + // Skip for update command itself and serve (which has its own check) + name := c.Name() + if name != "update" && name != "serve" && name != "version" { + cmd.CheckVersionOnStartup(Version) } - } - } - return "" -} - -// listAvailablePRDs returns all PRD names in .chief/prds/ -func listAvailablePRDs() []string { - prdsDir := ".chief/prds" - entries, err := os.ReadDir(prdsDir) - if err != nil { - return nil + }, + RunE: func(c *cobra.Command, args []string) error { + // Resolve positional argument as PRD name or path + if len(args) > 0 { + arg := args[0] + if strings.HasSuffix(arg, ".json") || strings.HasSuffix(arg, "/") { + opts.PRDPath = arg + } else { + opts.PRDPath = fmt.Sprintf(".chief/prds/%s/prd.json", arg) + } + } + runTUIWithOptions(opts) + return nil + }, } - var names []string - for _, entry := range entries { - if entry.IsDir() { - prdPath := filepath.Join(prdsDir, entry.Name(), "prd.md") - if _, err := os.Stat(prdPath); err == nil { - names = append(names, entry.Name()) - } + // Set custom version template to match previous output format + rootCmd.SetVersionTemplate("chief version {{.Version}}\n") + + // Root flags (TUI mode) + rootCmd.Flags().IntVarP(&opts.MaxIterations, "max-iterations", "n", 0, "Set maximum iterations (default: dynamic)") + rootCmd.Flags().BoolVar(&opts.NoRetry, "no-retry", false, "Disable auto-retry on Claude crashes") + rootCmd.Flags().BoolVar(&opts.Verbose, "verbose", false, "Show raw Claude output in log") + rootCmd.Flags().BoolVar(&opts.Merge, "merge", false, "Auto-merge progress on conversion conflicts") + rootCmd.Flags().BoolVar(&opts.Force, "force", false, "Auto-overwrite on conversion conflicts") + + // Subcommands + rootCmd.AddCommand(newNewCmd()) + rootCmd.AddCommand(newEditCmd()) + rootCmd.AddCommand(newStatusCmd()) + rootCmd.AddCommand(newListCmd()) + rootCmd.AddCommand(newUpdateCmd()) + rootCmd.AddCommand(newLoginCmd()) + rootCmd.AddCommand(newLogoutCmd()) + rootCmd.AddCommand(newServeCmd()) + rootCmd.AddCommand(newWiggumCmd()) + + // Custom help for root command only (subcommands use default Cobra help) + defaultHelp := rootCmd.HelpFunc() + rootCmd.SetHelpFunc(func(c *cobra.Command, args []string) { + if c != rootCmd { + defaultHelp(c, args) + return } - } - return names + fmt.Print(`Chief - Autonomous PRD Agent + +Usage: + chief [options] [|] + chief [arguments] + +Commands: + new [name] [context] Create a new PRD interactively + edit [name] [options] Edit an existing PRD interactively + status [name] Show progress for a PRD (default: main) + list List all PRDs with progress + update Update Chief to the latest version + login Authenticate with uplink.chiefloop.com + logout Log out and deauthorize this device + serve Start headless daemon for web app + update Update Chief to the latest version + +Options: + --max-iterations N, -n N Set maximum iterations (default: dynamic) + --no-retry Disable auto-retry on Claude crashes + --verbose Show raw Claude output in log + --merge Auto-merge progress on conversion conflicts + --force Auto-overwrite on conversion conflicts + -h, --help Show this help message + -v, --version Show version number + +Examples: + chief Launch TUI with default PRD (.chief/prds/main/) + chief auth Launch TUI with named PRD (.chief/prds/auth/) + chief ./my-prd.json Launch TUI with specific PRD file + chief -n 20 Launch with 20 max iterations + chief --max-iterations=5 auth + Launch auth PRD with 5 max iterations + chief --verbose Launch with raw Claude output visible + chief new Create PRD in .chief/prds/main/ + chief new auth Create PRD in .chief/prds/auth/ + chief new auth "JWT authentication for REST API" + Create PRD with context hint + chief edit Edit PRD in .chief/prds/main/ + chief edit auth Edit PRD in .chief/prds/auth/ + chief edit auth --merge Edit and auto-merge progress + chief status Show progress for default PRD + chief status auth Show progress for auth PRD + chief list List all PRDs with progress + chief update Update to the latest version + chief --version Show version number +`) + }) + + return rootCmd } -// parseAgentFlags extracts --agent and --agent-path from args[startIdx:], -// returning the agent name, agent path, remaining args (with agent flags removed), -// and the updated index offsets. It exits on missing values. -func parseAgentFlags(args []string, startIdx int) (agentName, agentPath string, remaining []string) { - for i := startIdx; i < len(args); i++ { - arg := args[i] - switch { - case arg == "--agent": - if i+1 < len(args) { - i++ - agentName = args[i] - } else { - fmt.Fprintf(os.Stderr, "Error: --agent requires a value (claude, codex, opencode, or cursor)\n") - os.Exit(1) +func newNewCmd() *cobra.Command { + return &cobra.Command{ + Use: "new [name] [context...]", + Short: "Create a new PRD interactively", + Args: cobra.ArbitraryArgs, + RunE: func(c *cobra.Command, args []string) error { + opts := cmd.NewOptions{} + if len(args) > 0 { + opts.Name = args[0] } - case strings.HasPrefix(arg, "--agent="): - agentName = strings.TrimPrefix(arg, "--agent=") - case arg == "--agent-path": - if i+1 < len(args) { - i++ - agentPath = args[i] - } else { - fmt.Fprintf(os.Stderr, "Error: --agent-path requires a value\n") - os.Exit(1) + if len(args) > 1 { + opts.Context = strings.Join(args[1:], " ") } - case strings.HasPrefix(arg, "--agent-path="): - agentPath = strings.TrimPrefix(arg, "--agent-path=") - default: - remaining = append(remaining, arg) - } + return cmd.RunNew(opts) + }, } - return } -// parseTUIFlags parses command-line flags for TUI mode -func parseTUIFlags() *TUIOptions { - opts := &TUIOptions{ - PRDPath: "", // Will be resolved later - MaxIterations: 0, // 0 signals dynamic calculation (remaining stories + 5) - Verbose: false, - Merge: false, - Force: false, - NoRetry: false, +func newEditCmd() *cobra.Command { + editOpts := &cmd.EditOptions{} + + editCmd := &cobra.Command{ + Use: "edit [name]", + Short: "Edit an existing PRD interactively", + Args: cobra.MaximumNArgs(1), + RunE: func(c *cobra.Command, args []string) error { + if len(args) > 0 { + editOpts.Name = args[0] + } + return cmd.RunEdit(*editOpts) + }, } - // Pre-extract agent flags so they don't interfere with positional arg parsing - opts.Agent, opts.AgentPath, _ = parseAgentFlags(os.Args, 1) + editCmd.Flags().BoolVar(&editOpts.Merge, "merge", false, "Auto-merge progress on conversion conflicts") + editCmd.Flags().BoolVar(&editOpts.Force, "force", false, "Auto-overwrite on conversion conflicts") - for i := 1; i < len(os.Args); i++ { - arg := os.Args[i] + return editCmd +} - switch { - case arg == "--help" || arg == "-h": - printHelp() - return nil - case arg == "--version" || arg == "-v": - fmt.Printf("chief version %s\n", Version) - return nil - case arg == "--verbose": - opts.Verbose = true - case arg == "--merge": - opts.Merge = true - case arg == "--force": - opts.Force = true - case arg == "--no-retry": - opts.NoRetry = true - case arg == "--agent" || arg == "--agent-path": - i++ // skip value (already parsed by parseAgentFlags) - case strings.HasPrefix(arg, "--agent=") || strings.HasPrefix(arg, "--agent-path="): - // already parsed by parseAgentFlags - case arg == "--max-iterations" || arg == "-n": - // Next argument should be the number - if i+1 < len(os.Args) { - i++ - n, err := strconv.Atoi(os.Args[i]) - if err != nil { - fmt.Fprintf(os.Stderr, "Error: invalid value for %s: %s\n", arg, os.Args[i]) - os.Exit(1) - } - if n < 1 { - fmt.Fprintf(os.Stderr, "Error: --max-iterations must be at least 1\n") - os.Exit(1) - } - opts.MaxIterations = n - } else { - fmt.Fprintf(os.Stderr, "Error: %s requires a value\n", arg) - os.Exit(1) - } - case strings.HasPrefix(arg, "--max-iterations="): - val := strings.TrimPrefix(arg, "--max-iterations=") - n, err := strconv.Atoi(val) - if err != nil { - fmt.Fprintf(os.Stderr, "Error: invalid value for --max-iterations: %s\n", val) - os.Exit(1) - } - if n < 1 { - fmt.Fprintf(os.Stderr, "Error: --max-iterations must be at least 1\n") - os.Exit(1) - } - opts.MaxIterations = n - case strings.HasPrefix(arg, "-n="): - val := strings.TrimPrefix(arg, "-n=") - n, err := strconv.Atoi(val) - if err != nil { - fmt.Fprintf(os.Stderr, "Error: invalid value for -n: %s\n", val) - os.Exit(1) - } - if n < 1 { - fmt.Fprintf(os.Stderr, "Error: -n must be at least 1\n") - os.Exit(1) +func newStatusCmd() *cobra.Command { + return &cobra.Command{ + Use: "status [name]", + Short: "Show progress for a PRD (default: main)", + Args: cobra.MaximumNArgs(1), + RunE: func(c *cobra.Command, args []string) error { + opts := cmd.StatusOptions{} + if len(args) > 0 { + opts.Name = args[0] } - opts.MaxIterations = n - case strings.HasPrefix(arg, "-"): - // Unknown flag - fmt.Fprintf(os.Stderr, "Error: unknown flag: %s\n", arg) - fmt.Fprintf(os.Stderr, "Run 'chief --help' for usage.\n") - os.Exit(1) - default: - // Positional argument: PRD name or path - if strings.HasSuffix(arg, ".md") || strings.HasSuffix(arg, ".json") || strings.HasSuffix(arg, "/") { - opts.PRDPath = arg - } else { - // Treat as PRD name - opts.PRDPath = fmt.Sprintf(".chief/prds/%s/prd.md", arg) - } - } + return cmd.RunStatus(opts) + }, } - - return opts } -func runNew() { - opts := cmd.NewOptions{} - - // Parse arguments: chief new [name] [context...] [--agent X] [--agent-path X] - flagAgent, flagPath, positional := parseAgentFlags(os.Args, 2) - // Filter out remaining flags, keep only positional args - var args []string - for _, a := range positional { - if !strings.HasPrefix(a, "-") { - args = append(args, a) - } - } - if len(args) > 0 { - opts.Name = args[0] - } - if len(args) > 1 { - opts.Context = strings.Join(args[1:], " ") +func newListCmd() *cobra.Command { + return &cobra.Command{ + Use: "list", + Short: "List all PRDs with progress", + Args: cobra.NoArgs, + RunE: func(c *cobra.Command, args []string) error { + return cmd.RunList(cmd.ListOptions{}) + }, } +} - opts.Provider = resolveProvider(flagAgent, flagPath) - if err := cmd.RunNew(opts); err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) +func newUpdateCmd() *cobra.Command { + return &cobra.Command{ + Use: "update", + Short: "Update Chief to the latest version", + Args: cobra.NoArgs, + RunE: func(c *cobra.Command, args []string) error { + return cmd.RunUpdate(cmd.UpdateOptions{Version: Version}) + }, } } -func runEdit() { - opts := cmd.EditOptions{} +func newLoginCmd() *cobra.Command { + loginOpts := &cmd.LoginOptions{} - // Parse arguments: chief edit [name] [--agent X] [--agent-path X] - flagAgent, flagPath, remaining := parseAgentFlags(os.Args, 2) - for _, arg := range remaining { - if opts.Name == "" && !strings.HasPrefix(arg, "-") { - opts.Name = arg - } - } - - opts.Provider = resolveProvider(flagAgent, flagPath) - if err := cmd.RunEdit(opts); err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) + loginCmd := &cobra.Command{ + Use: "login", + Short: "Authenticate with uplink.chiefloop.com", + Args: cobra.NoArgs, + RunE: func(c *cobra.Command, args []string) error { + return cmd.RunLogin(*loginOpts) + }, } -} -func runStatus() { - opts := cmd.StatusOptions{} + loginCmd.Flags().StringVar(&loginOpts.DeviceName, "name", "", "Override device name (default: hostname)") + loginCmd.Flags().StringVar(&loginOpts.SetupToken, "setup-token", "", "One-time setup token for automated auth") + loginCmd.Flags().StringVar(&loginOpts.BaseURL, "server-url", "", "Override server URL (default: https://uplink.chiefloop.com)") - // Parse arguments: chief status [name] - if len(os.Args) > 2 && !strings.HasPrefix(os.Args[2], "-") { - opts.Name = os.Args[2] - } + return loginCmd +} - if err := cmd.RunStatus(opts); err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) +func newLogoutCmd() *cobra.Command { + return &cobra.Command{ + Use: "logout", + Short: "Log out and deauthorize this device", + Args: cobra.NoArgs, + RunE: func(c *cobra.Command, args []string) error { + return cmd.RunLogout(cmd.LogoutOptions{}) + }, } } -func runUpdate() { - if err := cmd.RunUpdate(cmd.UpdateOptions{ - Version: Version, - }); err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) +func newServeCmd() *cobra.Command { + serveOpts := &cmd.ServeOptions{} + + serveCmd := &cobra.Command{ + Use: "serve", + Short: "Start headless daemon for web app", + Long: "Starts a headless daemon that connects to uplink.chiefloop.com via WebSocket and accepts commands from the web app.", + Args: cobra.NoArgs, + RunE: func(c *cobra.Command, args []string) error { + serveOpts.Version = Version + return cmd.RunServe(*serveOpts) + }, } -} -func runList() { - opts := cmd.ListOptions{} + serveCmd.Flags().StringVar(&serveOpts.Workspace, "workspace", ".", "Path to workspace directory (default: current directory)") + serveCmd.Flags().StringVar(&serveOpts.DeviceName, "name", "", "Override device name for this session") + serveCmd.Flags().StringVar(&serveOpts.LogFile, "log-file", "", "Path to log file (default: stdout)") + serveCmd.Flags().StringVar(&serveOpts.ServerURL, "server-url", "", "Override server URL (default: https://uplink.chiefloop.com)") - if err := cmd.RunList(opts); err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) + return serveCmd +} + +func newWiggumCmd() *cobra.Command { + return &cobra.Command{ + Use: "wiggum", + Short: "Bake 'em away, toys!", + Hidden: true, + Args: cobra.NoArgs, + Run: func(c *cobra.Command, args []string) { + printWiggum() + }, } } -// resolveProvider loads config and resolves the agent provider, exiting on error. -func resolveProvider(flagAgent, flagPath string) loop.Provider { - cwd, err := os.Getwd() +// findAvailablePRD looks for any available PRD in .chief/prds/ +// Returns the path to the first PRD found, or empty string if none exist. +func findAvailablePRD() string { + prdsDir := ".chief/prds" + entries, err := os.ReadDir(prdsDir) if err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) + return "" } - cfg, err := config.Load(cwd) - if err != nil { - fmt.Fprintf(os.Stderr, "Error: failed to load .chief/config.yaml: %v\n", err) - os.Exit(1) + + for _, entry := range entries { + if entry.IsDir() { + prdPath := filepath.Join(prdsDir, entry.Name(), "prd.json") + if _, err := os.Stat(prdPath); err == nil { + return prdPath + } + } } - provider, err := agent.Resolve(flagAgent, flagPath, cfg) + return "" +} + +// listAvailablePRDs returns all PRD names in .chief/prds/ +func listAvailablePRDs() []string { + prdsDir := ".chief/prds" + entries, err := os.ReadDir(prdsDir) if err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) + return nil } - if err := agent.CheckInstalled(provider); err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) + + var names []string + for _, entry := range entries { + if entry.IsDir() { + prdPath := filepath.Join(prdsDir, entry.Name(), "prd.json") + if _, err := os.Stat(prdPath); err == nil { + names = append(names, entry.Name()) + } + } } - return provider + return names } func runTUIWithOptions(opts *TUIOptions) { - provider := resolveProvider(opts.Agent, opts.AgentPath) - prdPath := opts.PRDPath // If no PRD specified, try to find one if prdPath == "" { // Try "main" first - mainPath := ".chief/prds/main/prd.md" + mainPath := ".chief/prds/main/prd.json" if _, err := os.Stat(mainPath); err == nil { prdPath = mainPath } else { @@ -395,8 +372,7 @@ func runTUIWithOptions(opts *TUIOptions) { // Create the PRD newOpts := cmd.NewOptions{ - Name: result.PRDName, - Provider: provider, + Name: result.PRDName, } if err := cmd.RunNew(newOpts); err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) @@ -404,7 +380,7 @@ func runTUIWithOptions(opts *TUIOptions) { } // Restart TUI with the new PRD - opts.PRDPath = fmt.Sprintf(".chief/prds/%s/prd.md", result.PRDName) + opts.PRDPath = fmt.Sprintf(".chief/prds/%s/prd.json", result.PRDName) runTUIWithOptions(opts) return } @@ -412,18 +388,25 @@ func runTUIWithOptions(opts *TUIOptions) { prdDir := filepath.Dir(prdPath) - // Auto-migrate: if prd.json exists alongside prd.md, migrate status - jsonPath := filepath.Join(prdDir, "prd.json") - if _, err := os.Stat(jsonPath); err == nil { - fmt.Println("Migrating status from prd.json to prd.md...") - if err := prd.MigrateFromJSON(prdDir); err != nil { - fmt.Printf("Warning: migration failed: %v\n", err) - } else { - fmt.Println("Migration complete (prd.json renamed to prd.json.bak).") + // Check if prd.md is newer than prd.json and run conversion if needed + needsConvert, err := prd.NeedsConversion(prdDir) + if err != nil { + fmt.Printf("Warning: failed to check conversion status: %v\n", err) + } else if needsConvert { + fmt.Println("prd.md is newer than prd.json, running conversion...") + convertOpts := prd.ConvertOptions{ + PRDDir: prdDir, + Merge: opts.Merge, + Force: opts.Force, } + if err := prd.Convert(convertOpts); err != nil { + fmt.Printf("Error converting PRD: %v\n", err) + os.Exit(1) + } + fmt.Println("Conversion complete.") } - app, err := tui.NewAppWithOptions(prdPath, opts.MaxIterations, provider) + app, err := tui.NewAppWithOptions(prdPath, opts.MaxIterations) if err != nil { // Check if this is a missing PRD file error if os.IsNotExist(err) || strings.Contains(err.Error(), "no such file") { @@ -470,91 +453,34 @@ func runTUIWithOptions(opts *TUIOptions) { case tui.PostExitInit: // Run new command then restart TUI newOpts := cmd.NewOptions{ - Name: finalApp.PostExitPRD, - Provider: provider, + Name: finalApp.PostExitPRD, } if err := cmd.RunNew(newOpts); err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) } // Restart TUI with the new PRD - opts.PRDPath = fmt.Sprintf(".chief/prds/%s/prd.md", finalApp.PostExitPRD) + opts.PRDPath = fmt.Sprintf(".chief/prds/%s/prd.json", finalApp.PostExitPRD) runTUIWithOptions(opts) case tui.PostExitEdit: // Run edit command then restart TUI editOpts := cmd.EditOptions{ - Name: finalApp.PostExitPRD, - Provider: provider, + Name: finalApp.PostExitPRD, + Merge: opts.Merge, + Force: opts.Force, } if err := cmd.RunEdit(editOpts); err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) } // Restart TUI with the edited PRD - opts.PRDPath = fmt.Sprintf(".chief/prds/%s/prd.md", finalApp.PostExitPRD) + opts.PRDPath = fmt.Sprintf(".chief/prds/%s/prd.json", finalApp.PostExitPRD) runTUIWithOptions(opts) } } } -func printHelp() { - fmt.Println(`Chief - Autonomous PRD Agent - -Usage: - chief [options] [|] - chief [arguments] - -Commands: - new [name] [context] Create a new PRD interactively - edit [name] [options] Edit an existing PRD interactively - status [name] Show progress for a PRD (default: main) - list List all PRDs with progress - update Update Chief to the latest version - help Show this help message - -Global Options: - --agent Agent CLI to use: claude (default), codex, opencode, or cursor - --agent-path Custom path to agent CLI binary - --max-iterations N, -n N Set maximum iterations (default: dynamic) - --no-retry Disable auto-retry on agent crashes - --verbose Show raw agent output in log - --merge Auto-merge progress on conversion conflicts - --force Auto-overwrite on conversion conflicts - --help, -h Show this help message - --version, -v Show version number - -Edit Options: - --merge Auto-merge progress on conversion conflicts - --force Auto-overwrite on conversion conflicts - -Positional Arguments: - PRD name (loads .chief/prds//prd.md) - Direct path to a prd.md file - -Examples: - chief Launch TUI with default PRD (.chief/prds/main/) - chief auth Launch TUI with named PRD (.chief/prds/auth/) - chief ./my-prd.md Launch TUI with specific PRD file - chief -n 20 Launch with 20 max iterations - chief --max-iterations=5 auth - Launch auth PRD with 5 max iterations - chief --verbose Launch with raw agent output visible - chief --agent codex Use Codex CLI instead of Claude - chief --agent cursor Use Cursor CLI as agent - chief new Create PRD in .chief/prds/main/ - chief new auth Create PRD in .chief/prds/auth/ - chief new auth "JWT authentication for REST API" - Create PRD with context hint - chief edit Edit PRD in .chief/prds/main/ - chief edit auth Edit PRD in .chief/prds/auth/ - chief edit auth --merge Edit and auto-merge progress - chief status Show progress for default PRD - chief status auth Show progress for auth PRD - chief list List all PRDs with progress - chief --version Show version number`) -} - func printWiggum() { // ANSI color codes blue := "\033[34m" diff --git a/deploy/chief.service b/deploy/chief.service new file mode 100644 index 00000000..80af505b --- /dev/null +++ b/deploy/chief.service @@ -0,0 +1,25 @@ +[Unit] +Description=Chief - Autonomous PRD Agent +Documentation=https://github.com/MiniCodeMonkey/chief +After=network-online.target +Wants=network-online.target +ConditionPathExists=/home/chief/.chief/credentials.yaml + +[Service] +Type=simple +User=chief +Group=chief +WorkingDirectory=/home/chief +ExecStart=/usr/local/bin/chief serve --workspace /home/chief/projects --log-file /home/chief/.chief/serve.log +Restart=always +RestartSec=5 +Environment=HOME=/home/chief + +# Security hardening +NoNewPrivileges=true +ProtectSystem=strict +ProtectHome=false +ReadWritePaths=/home/chief + +[Install] +WantedBy=multi-user.target diff --git a/deploy/cloud-init.sh b/deploy/cloud-init.sh new file mode 100755 index 00000000..abc07803 --- /dev/null +++ b/deploy/cloud-init.sh @@ -0,0 +1,263 @@ +#!/bin/bash +# Chief Cloud-Init Setup Script +# https://github.com/MiniCodeMonkey/chief +# +# This script sets up a VPS to run Chief as a systemd service. +# It is designed to be run via cloud-init during VPS provisioning. +# +# Usage (cloud-init user-data): +# #!/bin/bash +# curl -fsSL https://raw.githubusercontent.com/MiniCodeMonkey/chief/main/deploy/cloud-init.sh | bash +# +# With setup token (automated auth): +# #!/bin/bash +# curl -fsSL https://raw.githubusercontent.com/MiniCodeMonkey/chief/main/deploy/cloud-init.sh | CHIEF_SETUP_TOKEN= bash +# +# What this script does: +# 1. Creates a 'chief' user +# 2. Installs the Chief binary +# 3. Installs Claude Code CLI (via npm) +# 4. Creates the workspace directory +# 5. Writes and enables the systemd unit file +# +# After this script runs, you must: +# 1. SSH into the server +# 2. Run: sudo -u chief chief login (skipped if CHIEF_SETUP_TOKEN is set) +# 3. Authenticate Claude Code: sudo -u chief claude +# 4. Start the service: sudo systemctl start chief +# +# This script is idempotent (safe to run multiple times). + +set -euo pipefail + +GITHUB_REPO="MiniCodeMonkey/chief" +CHIEF_USER="chief" +CHIEF_HOME="/home/${CHIEF_USER}" +WORKSPACE_DIR="${CHIEF_HOME}/projects" +BINARY_PATH="/usr/local/bin/chief" +SERVICE_FILE="/etc/systemd/system/chief.service" + +info() { + echo "==> $1" +} + +warn() { + echo "WARNING: $1" +} + +error() { + echo "ERROR: $1" >&2 + exit 1 +} + +# Create chief user if it doesn't exist +create_user() { + if id "${CHIEF_USER}" &>/dev/null; then + info "User '${CHIEF_USER}' already exists" + else + info "Creating user '${CHIEF_USER}'..." + useradd --create-home --shell /bin/bash "${CHIEF_USER}" + fi +} + +# Install Chief binary +install_chief() { + info "Installing Chief binary..." + curl -fsSL "https://raw.githubusercontent.com/${GITHUB_REPO}/main/install.sh" | CHIEF_INSTALL_DIR=/usr/local/bin sh +} + +# Install Node.js and Claude Code CLI +install_claude_code() { + if command -v claude &>/dev/null; then + info "Claude Code CLI already installed" + return 0 + fi + + # Install Node.js if not present + if ! command -v node &>/dev/null; then + info "Installing Node.js..." + if command -v apt-get &>/dev/null; then + curl -fsSL https://deb.nodesource.com/setup_lts.x | bash - + apt-get install -y nodejs + elif command -v dnf &>/dev/null; then + curl -fsSL https://rpm.nodesource.com/setup_lts.x | bash - + dnf install -y nodejs + elif command -v yum &>/dev/null; then + curl -fsSL https://rpm.nodesource.com/setup_lts.x | bash - + yum install -y nodejs + else + warn "Could not install Node.js automatically. Please install it manually." + return 1 + fi + fi + + info "Installing Claude Code CLI..." + npm install -g @anthropic-ai/claude-code +} + +# Create workspace directory +create_workspace() { + if [ -d "${WORKSPACE_DIR}" ]; then + info "Workspace directory already exists: ${WORKSPACE_DIR}" + else + info "Creating workspace directory: ${WORKSPACE_DIR}" + mkdir -p "${WORKSPACE_DIR}" + fi + chown -R "${CHIEF_USER}:${CHIEF_USER}" "${WORKSPACE_DIR}" +} + +# Create .chief config directory +create_config_dir() { + local config_dir="${CHIEF_HOME}/.chief" + if [ -d "${config_dir}" ]; then + info "Config directory already exists: ${config_dir}" + else + info "Creating config directory: ${config_dir}" + mkdir -p "${config_dir}" + fi + chown -R "${CHIEF_USER}:${CHIEF_USER}" "${config_dir}" +} + +# Install and enable systemd service +install_service() { + info "Installing systemd service..." + + cat > "${SERVICE_FILE}" <<'UNIT' +[Unit] +Description=Chief - Autonomous PRD Agent +Documentation=https://github.com/MiniCodeMonkey/chief +After=network-online.target +Wants=network-online.target +ConditionPathExists=/home/chief/.chief/credentials.yaml + +[Service] +Type=simple +User=chief +Group=chief +WorkingDirectory=/home/chief +ExecStart=/usr/local/bin/chief serve --workspace /home/chief/projects --log-file /home/chief/.chief/serve.log +Restart=always +RestartSec=5 +Environment=HOME=/home/chief + +# Security hardening +NoNewPrivileges=true +ProtectSystem=strict +ProtectHome=false +ReadWritePaths=/home/chief + +[Install] +WantedBy=multi-user.target +UNIT + + systemctl daemon-reload + systemctl enable chief.service + info "Service enabled (but NOT started — authentication required first)" +} + +# Handle setup token if provided +handle_setup_token() { + if [ -z "${CHIEF_SETUP_TOKEN:-}" ]; then + return 0 + fi + + info "Setup token provided, configuring automated authentication..." + + # Write the setup token to a temporary file readable only by the chief user + local token_file="/tmp/chief-setup-token" + echo "${CHIEF_SETUP_TOKEN}" > "${token_file}" + chown "${CHIEF_USER}:${CHIEF_USER}" "${token_file}" + chmod 600 "${token_file}" + + # Create a one-shot systemd service that exchanges the token + cat > /etc/systemd/system/chief-setup.service < +{{PRD_CONTENT}} + + +Do NOT use any tools. Do NOT write any files. Output ONLY the raw JSON to stdout — no markdown fences, no explanation, no preamble, no commentary. The JSON must follow this exact structure: + +{ + "project": "Project Name", + "description": "Brief project description", + "userStories": [ + { + "id": "US-001", + "title": "Story Title", + "description": "Full description of what the user story accomplishes", + "acceptanceCriteria": [ + "First acceptance criterion", + "Second acceptance criterion" + ], + "priority": 1, + "passes": false + } + ] +} + +Rules: +1. Extract the project name from the main heading (# heading) +2. Extract the description from the introductory paragraph +3. For each user story: + - Generate sequential IDs: US-001, US-002, etc. + - Extract title from story heading + - Extract description from story body + - Extract acceptance criteria as an array of strings + - Assign priority based on order (first story = 1, second = 2, etc.) + - Set "passes" to false for all stories (progress tracking happens later) +4. Do NOT include "inProgress" field for new stories +5. CRITICAL - JSON string escaping: All double quotes inside JSON string values MUST be escaped with a backslash. For example: + - WRONG: "description": "Click the "Submit" button" + - RIGHT: "description": "Click the \"Submit\" button" + This applies to ALL string fields: title, description, and every entry in acceptanceCriteria. +6. Ensure the JSON is valid and properly formatted with 2-space indentation diff --git a/embed/edit_prompt.txt b/embed/edit_prompt.txt index f67082d2..08bd16f7 100644 --- a/embed/edit_prompt.txt +++ b/embed/edit_prompt.txt @@ -127,10 +127,12 @@ Before saving the edited PRD, verify: - [ ] Functional requirements are numbered and unambiguous - [ ] Non-goals updated if scope changed - [ ] Success metrics reviewed and updated if needed +- [ ] The file is ready for conversion to prd.json + --- ## Final Step -Once the edits are complete, tell the user to type `exit` to finish. +Once the edits are complete, tell the user to type `/exit` to finish. Chief will automatically convert the updated PRD. Start by reading the existing PRD and understanding what changes the user wants to make. diff --git a/embed/embed.go b/embed/embed.go index f4734207..d7b1e73c 100644 --- a/embed/embed.go +++ b/embed/embed.go @@ -16,18 +16,15 @@ var initPromptTemplate string //go:embed edit_prompt.txt var editPromptTemplate string +//go:embed convert_prompt.txt +var convertPromptTemplate string + //go:embed detect_setup_prompt.txt var detectSetupPromptTemplate string -// GetPrompt returns the agent prompt with the progress path and -// current story context substituted. The storyContext is the JSON of the -// current story to work on, inlined directly into the prompt so that the -// agent does not need to read the entire prd.md file. -func GetPrompt(progressPath, storyContext, storyID, storyTitle string) string { - result := strings.ReplaceAll(promptTemplate, "{{PROGRESS_PATH}}", progressPath) - result = strings.ReplaceAll(result, "{{STORY_CONTEXT}}", storyContext) - result = strings.ReplaceAll(result, "{{STORY_ID}}", storyID) - return strings.ReplaceAll(result, "{{STORY_TITLE}}", storyTitle) +// GetPrompt returns the agent prompt with the PRD path substituted. +func GetPrompt(prdPath string) string { + return strings.ReplaceAll(promptTemplate, "{{PRD_PATH}}", prdPath) } // GetInitPrompt returns the PRD generator prompt with the PRD directory and optional context substituted. @@ -44,6 +41,11 @@ func GetEditPrompt(prdDir string) string { return strings.ReplaceAll(editPromptTemplate, "{{PRD_DIR}}", prdDir) } +// GetConvertPrompt returns the PRD converter prompt with the PRD content inlined. +func GetConvertPrompt(prdContent string) string { + return strings.ReplaceAll(convertPromptTemplate, "{{PRD_CONTENT}}", prdContent) +} + // GetDetectSetupPrompt returns the prompt for detecting project setup commands. func GetDetectSetupPrompt() string { return detectSetupPromptTemplate diff --git a/embed/embed_test.go b/embed/embed_test.go index aff9f436..ac3fc1b3 100644 --- a/embed/embed_test.go +++ b/embed/embed_test.go @@ -6,51 +6,30 @@ import ( ) func TestGetPrompt(t *testing.T) { - progressPath := "/path/to/progress.md" - storyContext := `{"id":"US-001","title":"Test Story"}` - prompt := GetPrompt(progressPath, storyContext, "US-001", "Test Story") + prdPath := "/path/to/prd.json" + prompt := GetPrompt(prdPath) - // Verify all placeholders were substituted - if strings.Contains(prompt, "{{PROGRESS_PATH}}") { - t.Error("Expected {{PROGRESS_PATH}} to be substituted") - } - if strings.Contains(prompt, "{{STORY_CONTEXT}}") { - t.Error("Expected {{STORY_CONTEXT}} to be substituted") - } - if strings.Contains(prompt, "{{STORY_ID}}") { - t.Error("Expected {{STORY_ID}} to be substituted") - } - if strings.Contains(prompt, "{{STORY_TITLE}}") { - t.Error("Expected {{STORY_TITLE}} to be substituted") - } - - // Verify the commit message contains the exact story ID and title - if !strings.Contains(prompt, "feat: US-001 - Test Story") { - t.Error("Expected prompt to contain exact commit message 'feat: US-001 - Test Story'") + // Verify the PRD path placeholder was substituted + if strings.Contains(prompt, "{{PRD_PATH}}") { + t.Error("Expected {{PRD_PATH}} to be substituted") } - // Verify the progress path appears in the prompt - if !strings.Contains(prompt, progressPath) { - t.Errorf("Expected prompt to contain progress path %q", progressPath) + // Verify the PRD path appears in the prompt + if !strings.Contains(prompt, prdPath) { + t.Errorf("Expected prompt to contain PRD path %q", prdPath) } - // Verify the story context is inlined in the prompt - if !strings.Contains(prompt, storyContext) { - t.Error("Expected prompt to contain inlined story context") + // Verify the prompt contains key instructions + if !strings.Contains(prompt, "chief-complete") { + t.Error("Expected prompt to contain chief-complete instruction") } - // Verify the prompt contains chief-done stop condition - if !strings.Contains(prompt, "chief-done") { - t.Error("Expected prompt to contain chief-done instruction") + if !strings.Contains(prompt, "ralph-status") { + t.Error("Expected prompt to contain ralph-status instruction") } -} -func TestGetPrompt_NoFileReadInstruction(t *testing.T) { - prompt := GetPrompt("/path/progress.md", `{"id":"US-001"}`, "US-001", "Test Story") - - // The prompt should NOT instruct Claude to read the PRD file - if strings.Contains(prompt, "Read the PRD") { - t.Error("Expected prompt to NOT contain 'Read the PRD' file-read instruction") + if !strings.Contains(prompt, "passes: true") { + t.Error("Expected prompt to contain passes: true instruction") } } @@ -60,15 +39,34 @@ func TestPromptTemplateNotEmpty(t *testing.T) { } } -func TestGetPrompt_ChiefExclusion(t *testing.T) { - prompt := GetPrompt("/path/progress.md", `{"id":"US-001"}`, "US-001", "Test Story") +func TestGetConvertPrompt(t *testing.T) { + prdContent := "# My Feature\n\nA cool feature PRD." + prompt := GetConvertPrompt(prdContent) + + // Verify the prompt is not empty + if prompt == "" { + t.Error("Expected GetConvertPrompt() to return non-empty prompt") + } + + // Verify PRD content is inlined + if !strings.Contains(prompt, prdContent) { + t.Error("Expected prompt to contain the inlined PRD content") + } + if strings.Contains(prompt, "{{PRD_CONTENT}}") { + t.Error("Expected {{PRD_CONTENT}} to be substituted") + } + + // Verify key instructions are present + if !strings.Contains(prompt, "JSON") { + t.Error("Expected prompt to mention JSON") + } - // The prompt must instruct Claude to never stage or commit .chief/ files - if !strings.Contains(prompt, ".chief/") { - t.Error("Expected prompt to contain .chief/ exclusion instruction") + if !strings.Contains(prompt, "userStories") { + t.Error("Expected prompt to describe userStories structure") } - if !strings.Contains(prompt, "NEVER stage or commit") { - t.Error("Expected prompt to explicitly say NEVER stage or commit .chief/ files") + + if !strings.Contains(prompt, `"passes": false`) { + t.Error("Expected prompt to specify passes: false default") } } diff --git a/embed/init_prompt.txt b/embed/init_prompt.txt index 19c3fe73..cab267d7 100644 --- a/embed/init_prompt.txt +++ b/embed/init_prompt.txt @@ -230,10 +230,12 @@ Before saving the PRD, verify: - [ ] Functional requirements are numbered and unambiguous - [ ] Non-goals section defines clear boundaries - [ ] Success metrics are specific and measurable +- [ ] The file is ready for conversion to prd.json + --- ## Final Step -Once the PRD file is written, tell the user to type `exit` to finish. +Once the PRD file is written, tell the user to type `/exit` to finish. Chief will automatically convert it to the format needed for implementation. Start by understanding what the user wants to build. Ask your clarifying questions first, then create the PRD. diff --git a/embed/prompt.txt b/embed/prompt.txt index 8b9dee1f..e8359f30 100644 --- a/embed/prompt.txt +++ b/embed/prompt.txt @@ -4,22 +4,18 @@ You are an autonomous coding agent working on a software project. ## Your Task -Your current story: - -{{STORY_CONTEXT}} - - -1. Read `{{PROGRESS_PATH}}` if it exists (check Codebase Patterns section first) -2. Implement the user story above -3. Run quality checks (e.g., typecheck, lint, test - use whatever your project requires) -4. If checks pass, commit changes with message: `feat: {{STORY_ID}} - {{STORY_TITLE}}` - - **NEVER stage or commit `.chief/` files** — these are local working files and must stay out of version control - - Stage only the files you changed for the story (do NOT use `git add -A` or `git add .`) -5. Append your progress to `{{PROGRESS_PATH}}` +1. Read the PRD at `{{PRD_PATH}}` +2. Read `progress.md` if it exists (check Codebase Patterns section first) +3. Pick the **highest priority** user story where `passes: false` -- After determining which story to work on, output exact story id, e.g.: US-056 +4. Implement that single user story +5. Run quality checks (e.g., typecheck, lint, test - use whatever your project requires) +6. If checks pass, commit ALL changes with message: `feat: [Story ID] - [Story Title]` +7. Update the PRD to set `passes: true` for the completed story +8. Append your progress to `progress.md` ## Progress Report Format -APPEND to `{{PROGRESS_PATH}}` (never replace, always append): +APPEND to progress.md (never replace, always append): ``` ## [Date/Time] - [Story ID] - What was implemented @@ -35,7 +31,7 @@ The learnings section is critical - it helps future iterations avoid repeating m ## Consolidate Patterns -If you discover a **reusable pattern** that future iterations should know, add it to the `## Codebase Patterns` section at the TOP of `{{PROGRESS_PATH}}` (create it if it doesn't exist). This section should consolidate the most important learnings: +If you discover a **reusable pattern** that future iterations should know, add it to the `## Codebase Patterns` section at the TOP of progress.md (create it if it doesn't exist). This section should consolidate the most important learnings: ``` ## Codebase Patterns @@ -55,14 +51,16 @@ Only add patterns that are **general and reusable**, not story-specific details. ## Stop Condition -After implementing the story: -1. Review EACH acceptance criterion one by one and verify it is met -2. Only if ALL criteria pass: output -3. If any criterion is NOT met: end your response WITHOUT +After completing a user story, check if ALL stories have `passes: true`. + +If ALL stories are complete and passing, reply with: + + +If there are still stories with `passes: false`, end your response normally (another iteration will pick up the next story). ## Important - Work on ONE story per iteration - Commit frequently - Keep CI green -- Read the Codebase Patterns section in `{{PROGRESS_PATH}}` before starting +- Read the Codebase Patterns section in progress.md before starting diff --git a/go.mod b/go.mod index 63b53d00..8e1851e1 100644 --- a/go.mod +++ b/go.mod @@ -5,9 +5,12 @@ go 1.24.0 require ( github.com/alecthomas/chroma/v2 v2.23.1 github.com/charmbracelet/bubbletea v1.3.10 + github.com/charmbracelet/glamour v0.10.0 github.com/charmbracelet/lipgloss v1.1.1-0.20250404203927-76690c660834 github.com/charmbracelet/x/term v0.2.1 github.com/fsnotify/fsnotify v1.9.0 + github.com/gorilla/websocket v1.5.3 + github.com/spf13/cobra v1.10.2 gopkg.in/yaml.v3 v3.0.1 ) @@ -15,13 +18,14 @@ require ( github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect github.com/aymerick/douceur v0.2.0 // indirect github.com/charmbracelet/colorprofile v0.2.3-0.20250311203215-f60798e515dc // indirect - github.com/charmbracelet/glamour v0.10.0 // indirect github.com/charmbracelet/x/ansi v0.10.1 // indirect github.com/charmbracelet/x/cellbuf v0.0.13 // indirect github.com/charmbracelet/x/exp/slice v0.0.0-20250327172914-2fdc97757edf // indirect + github.com/creack/pty v1.1.24 // indirect github.com/dlclark/regexp2 v1.11.5 // indirect github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f // indirect github.com/gorilla/css v1.0.1 // indirect + github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/lucasb-eyer/go-colorful v1.2.0 // indirect github.com/mattn/go-isatty v0.0.20 // indirect github.com/mattn/go-localereader v0.0.1 // indirect @@ -32,6 +36,7 @@ require ( github.com/muesli/reflow v0.3.0 // indirect github.com/muesli/termenv v0.16.0 // indirect github.com/rivo/uniseg v0.4.7 // indirect + github.com/spf13/pflag v1.0.9 // 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 diff --git a/go.sum b/go.sum index a5fd2d79..36e24cde 100644 --- a/go.sum +++ b/go.sum @@ -6,6 +6,8 @@ github.com/alecthomas/repr v0.5.2 h1:SU73FTI9D1P5UNtvseffFSGmdNci/O6RsqzeXJtP0Qs github.com/alecthomas/repr v0.5.2/go.mod h1:Fr0507jx4eOXV7AlPV6AVZLYrLIuIeSOWtW57eE/O/4= github.com/aymanbagabas/go-osc52/v2 v2.0.1 h1:HwpRHbFMcZLEVr42D4p7XBqjyuxQH5SMiErDT4WkJ2k= github.com/aymanbagabas/go-osc52/v2 v2.0.1/go.mod h1:uYgXzlJ7ZpABp8OJ+exZzJJhRNQ2ASbcXHWsFqH8hp8= +github.com/aymanbagabas/go-udiff v0.2.0 h1:TK0fH4MteXUDspT88n8CKzvK0X9O2xu9yQjWpi6yML8= +github.com/aymanbagabas/go-udiff v0.2.0/go.mod h1:RE4Ex0qsGkTAJoQdQQCA0uG+nAzJO/pI/QwceO5fgrA= 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/charmbracelet/bubbletea v1.3.10 h1:otUDHWMMzQSB0Pkc87rm691KZ3SWa4KUlvF9nRvCICw= @@ -14,20 +16,21 @@ github.com/charmbracelet/colorprofile v0.2.3-0.20250311203215-f60798e515dc h1:4p github.com/charmbracelet/colorprofile v0.2.3-0.20250311203215-f60798e515dc/go.mod h1:X4/0JoqgTIPSFcRA/P6INZzIuyqdFY5rm8tb41s9okk= github.com/charmbracelet/glamour v0.10.0 h1:MtZvfwsYCx8jEPFJm3rIBFIMZUfUJ765oX8V6kXldcY= github.com/charmbracelet/glamour v0.10.0/go.mod h1:f+uf+I/ChNmqo087elLnVdCiVgjSKWuXa/l6NU2ndYk= -github.com/charmbracelet/lipgloss v1.1.0 h1:vYXsiLHVkK7fp74RkV7b2kq9+zDLoEU4MZoFqR/noCY= -github.com/charmbracelet/lipgloss v1.1.0/go.mod h1:/6Q8FR2o+kj8rz4Dq0zQc3vYf7X+B0binUUBwA0aL30= github.com/charmbracelet/lipgloss v1.1.1-0.20250404203927-76690c660834 h1:ZR7e0ro+SZZiIZD7msJyA+NjkCNNavuiPBLgerbOziE= github.com/charmbracelet/lipgloss v1.1.1-0.20250404203927-76690c660834/go.mod h1:aKC/t2arECF6rNOnaKaVU6y4t4ZeHQzqfxedE/VkVhA= github.com/charmbracelet/x/ansi v0.10.1 h1:rL3Koar5XvX0pHGfovN03f5cxLbCF2YvLeyz7D2jVDQ= github.com/charmbracelet/x/ansi v0.10.1/go.mod h1:3RQDQ6lDnROptfpWuUVIUG64bD2g2BgntdxH0Ya5TeE= -github.com/charmbracelet/x/cellbuf v0.0.13-0.20250311204145-2c3ea96c31dd h1:vy0GVL4jeHEwG5YOXDmi86oYw2yuYUGqz6a8sLwg0X8= -github.com/charmbracelet/x/cellbuf v0.0.13-0.20250311204145-2c3ea96c31dd/go.mod h1:xe0nKWGd3eJgtqZRaN9RjMtK7xUYchjzPr7q6kcvCCs= github.com/charmbracelet/x/cellbuf v0.0.13 h1:/KBBKHuVRbq1lYx5BzEHBAFBP8VcQzJejZ/IA3iR28k= github.com/charmbracelet/x/cellbuf v0.0.13/go.mod h1:xe0nKWGd3eJgtqZRaN9RjMtK7xUYchjzPr7q6kcvCCs= +github.com/charmbracelet/x/exp/golden v0.0.0-20240806155701-69247e0abc2a h1:G99klV19u0QnhiizODirwVksQB91TJKV/UaTnACcG30= +github.com/charmbracelet/x/exp/golden v0.0.0-20240806155701-69247e0abc2a/go.mod h1:wDlXFlCrmJ8J+swcL/MnGUuYnqgQdW9rhSD61oNMb6U= 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.1 h1:AQeHeLZ1OqSXhrAWpYUtZyX1T3zVxfpZuEQMIQaGIAQ= github.com/charmbracelet/x/term v0.2.1/go.mod h1:oQ4enTYFV7QN4m0i9mzHrViD7TQKvNEEkHUMCmsxdUg= +github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= +github.com/creack/pty v1.1.24 h1:bJrF4RRfyJnbTJqzRLHzcGaZK1NeM5kTC9jGgovnR1s= +github.com/creack/pty v1.1.24/go.mod h1:08sCNb52WyoAwi2QDyzUCTgcvVFhUzewun7wtTfvcwE= github.com/dlclark/regexp2 v1.11.5 h1:Q/sSnsKerHeCkc/jSTNq1oCm7KiVgUMZRDUoRu0JQZQ= github.com/dlclark/regexp2 v1.11.5/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8= github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f h1:Y/CXytFA4m6baUTXGLOoWe4PQhGxaX0KpnayAqC48p4= @@ -36,8 +39,12 @@ github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= github.com/gorilla/css v1.0.1 h1:ntNaBIghp6JmvWnxbZKANoLyuXTPZ4cAMlo6RyhlbO8= github.com/gorilla/css v1.0.1/go.mod h1:BvnYkspnSzMmwRK+b8/xgNPLiIuNZr6vbZBTPQ2A3b0= +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/hexops/gotextdiff v1.0.3 h1:gitA9+qJrrTCsiCl7+kh75nPqQt1cx4ZkudSTLoUqJM= github.com/hexops/gotextdiff v1.0.3/go.mod h1:pSWU5MAI3yDq+fZBTazCSJysOMbxWL1BSow5/V2vxeg= +github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= +github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/lucasb-eyer/go-colorful v1.2.0 h1:1nnpGOrhyZZuNyfu1QjKiUICQ74+3FNCN69Aj6K7nkY= github.com/lucasb-eyer/go-colorful v1.2.0/go.mod h1:R4dSotOR9KMtayYi1e77YzuveK+i7ruzyGqttikkLy0= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= @@ -61,6 +68,11 @@ github.com/rivo/uniseg v0.1.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJ github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= +github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= +github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU= +github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4= +github.com/spf13/pflag v1.0.9 h1:9exaQaMOCwffKiiiYk6/BndUBv+iRViNW+4lEMi0PvY= +github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= 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= @@ -68,6 +80,7 @@ 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/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= golang.org/x/exp v0.0.0-20220909182711-5c715a9e8561 h1:MDc5xs78ZrZr3HMQugiXOAkSZtfTpbJLDr/lwfgO53E= golang.org/x/exp v0.0.0-20220909182711-5c715a9e8561/go.mod h1:cyybsKvd6eL0RnXn6p/Grxp8F5bW7iYuBgsNCOHpMYE= golang.org/x/net v0.33.0 h1:74SYHlV8BIgHIFC/LrYkOGIwL19eTYXQ5wc6TBuO36I= @@ -78,8 +91,6 @@ golang.org/x/sys v0.36.0 h1:KVRy2GtZBrk1cBYA7MKu5bEZFxQk4NIDV6RLVcC8o0k= golang.org/x/sys v0.36.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= golang.org/x/term v0.31.0 h1:erwDkOK1Msy6offm1mOgvspSkslFnIGsFnxOKoufg3o= golang.org/x/term v0.31.0/go.mod h1:R4BeIy7D95HzImkxGkTW1UQTtP54tio2RyHz7PwK0aw= -golang.org/x/text v0.3.8 h1:nAL+RVCQ9uMn3vJZbV+MRnydTJFPf8qqY42YiA6MrqY= -golang.org/x/text v0.3.8/go.mod h1:E6s5w1FMmriuDzIBO73fBruAKo1PCIq6d2Q6DHfQ8WQ= golang.org/x/text v0.24.0 h1:dd5Bzh4yt5KYA8f9CJHCP4FB4D51c2c6JvN37xJJkJ0= golang.org/x/text v0.24.0/go.mod h1:L8rBsPeo2pSS+xqN0d5u2ikmjtmoJbDBT1b7nHvFCdU= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= diff --git a/internal/agent/claude.go b/internal/agent/claude.go deleted file mode 100644 index e6f5999e..00000000 --- a/internal/agent/claude.go +++ /dev/null @@ -1,58 +0,0 @@ -package agent - -import ( - "context" - "os/exec" - - "github.com/minicodemonkey/chief/internal/loop" -) - -// ClaudeProvider implements loop.Provider for the Claude Code CLI. -type ClaudeProvider struct { - cliPath string -} - -// NewClaudeProvider returns a Provider for the Claude CLI. -// If cliPath is empty, "claude" is used. -func NewClaudeProvider(cliPath string) *ClaudeProvider { - if cliPath == "" { - cliPath = "claude" - } - return &ClaudeProvider{cliPath: cliPath} -} - -// Name implements loop.Provider. -func (p *ClaudeProvider) Name() string { return "Claude" } - -// CLIPath implements loop.Provider. -func (p *ClaudeProvider) CLIPath() string { return p.cliPath } - -// LoopCommand implements loop.Provider. -func (p *ClaudeProvider) LoopCommand(ctx context.Context, prompt, workDir string) *exec.Cmd { - cmd := exec.CommandContext(ctx, p.cliPath, - "--dangerously-skip-permissions", - "-p", prompt, - "--output-format", "stream-json", - "--verbose", - ) - cmd.Dir = workDir - return cmd -} - -// InteractiveCommand implements loop.Provider. -func (p *ClaudeProvider) InteractiveCommand(workDir, prompt string) *exec.Cmd { - cmd := exec.Command(p.cliPath, prompt) - cmd.Dir = workDir - return cmd -} - -// ParseLine implements loop.Provider. -func (p *ClaudeProvider) ParseLine(line string) *loop.Event { - return loop.ParseLine(line) -} - -// LogFileName implements loop.Provider. -func (p *ClaudeProvider) LogFileName() string { return "claude.log" } - -// CleanOutput implements loop.Provider - Claude doesn't use a special format. -func (p *ClaudeProvider) CleanOutput(output string) string { return output } diff --git a/internal/agent/claude_test.go b/internal/agent/claude_test.go deleted file mode 100644 index 3ddad551..00000000 --- a/internal/agent/claude_test.go +++ /dev/null @@ -1,87 +0,0 @@ -package agent - -import ( - "context" - "testing" - - "github.com/minicodemonkey/chief/internal/loop" -) - -func TestClaudeProvider_Name(t *testing.T) { - p := NewClaudeProvider("") - if p.Name() != "Claude" { - t.Errorf("Name() = %q, want Claude", p.Name()) - } -} - -func TestClaudeProvider_CLIPath(t *testing.T) { - p := NewClaudeProvider("") - if p.CLIPath() != "claude" { - t.Errorf("CLIPath() empty arg = %q, want claude", p.CLIPath()) - } - p2 := NewClaudeProvider("/usr/local/bin/claude") - if p2.CLIPath() != "/usr/local/bin/claude" { - t.Errorf("CLIPath() custom = %q, want /usr/local/bin/claude", p2.CLIPath()) - } -} - -func TestClaudeProvider_LogFileName(t *testing.T) { - p := NewClaudeProvider("") - if p.LogFileName() != "claude.log" { - t.Errorf("LogFileName() = %q, want claude.log", p.LogFileName()) - } -} - -func TestClaudeProvider_LoopCommand(t *testing.T) { - ctx := context.Background() - p := NewClaudeProvider("/bin/claude") - cmd := p.LoopCommand(ctx, "hello world", "/work/dir") - - if cmd.Path != "/bin/claude" { - t.Errorf("LoopCommand Path = %q, want /bin/claude", cmd.Path) - } - wantArgs := []string{"/bin/claude", "--dangerously-skip-permissions", "-p", "hello world", "--output-format", "stream-json", "--verbose"} - if len(cmd.Args) != len(wantArgs) { - t.Fatalf("LoopCommand Args len = %d, want %d: %v", len(cmd.Args), len(wantArgs), cmd.Args) - } - for i, w := range wantArgs { - if cmd.Args[i] != w { - t.Errorf("LoopCommand Args[%d] = %q, want %q", i, cmd.Args[i], w) - } - } - if cmd.Dir != "/work/dir" { - t.Errorf("LoopCommand Dir = %q, want /work/dir", cmd.Dir) - } -} - -func TestClaudeProvider_InteractiveCommand(t *testing.T) { - p := NewClaudeProvider("/bin/claude") - cmd := p.InteractiveCommand("/work", "my prompt") - if cmd.Dir != "/work" { - t.Errorf("InteractiveCommand Dir = %q, want /work", cmd.Dir) - } - if len(cmd.Args) != 2 || cmd.Args[0] != "/bin/claude" || cmd.Args[1] != "my prompt" { - t.Errorf("InteractiveCommand Args = %v, want [/bin/claude my prompt]", cmd.Args) - } -} - -func TestClaudeProvider_ParseLine(t *testing.T) { - p := NewClaudeProvider("") - // Valid assistant text event - line := `{"type":"assistant","message":{"type":"assistant","content":[{"type":"text","text":"hello"}]}}` - e := p.ParseLine(line) - if e == nil { - t.Fatal("ParseLine(assistant text) returned nil") - } - if e.Type != loop.EventAssistantText { - t.Errorf("ParseLine(assistant text) Type = %v, want EventAssistantText", e.Type) - } -} - -func TestClaudeProvider_CleanOutput(t *testing.T) { - p := NewClaudeProvider("") - input := "some output" - if p.CleanOutput(input) != input { - t.Errorf("CleanOutput should return input unchanged") - } -} diff --git a/internal/agent/codex.go b/internal/agent/codex.go deleted file mode 100644 index 480b1ebf..00000000 --- a/internal/agent/codex.go +++ /dev/null @@ -1,55 +0,0 @@ -package agent - -import ( - "context" - "os/exec" - "strings" - - "github.com/minicodemonkey/chief/internal/loop" -) - -// CodexProvider implements loop.Provider for the Codex CLI. -type CodexProvider struct { - cliPath string -} - -// NewCodexProvider returns a Provider for the Codex CLI. -// If cliPath is empty, "codex" is used. -func NewCodexProvider(cliPath string) *CodexProvider { - if cliPath == "" { - cliPath = "codex" - } - return &CodexProvider{cliPath: cliPath} -} - -// Name implements loop.Provider. -func (p *CodexProvider) Name() string { return "Codex" } - -// CLIPath implements loop.Provider. -func (p *CodexProvider) CLIPath() string { return p.cliPath } - -// LoopCommand implements loop.Provider. -func (p *CodexProvider) LoopCommand(ctx context.Context, prompt, workDir string) *exec.Cmd { - cmd := exec.CommandContext(ctx, p.cliPath, "exec", "--json", "--yolo", "--skip-git-repo-check", "-C", workDir, "-") - cmd.Dir = workDir - cmd.Stdin = strings.NewReader(prompt) - return cmd -} - -// InteractiveCommand implements loop.Provider. -func (p *CodexProvider) InteractiveCommand(workDir, prompt string) *exec.Cmd { - cmd := exec.Command(p.cliPath, prompt) - cmd.Dir = workDir - return cmd -} - -// ParseLine implements loop.Provider. -func (p *CodexProvider) ParseLine(line string) *loop.Event { - return loop.ParseLineCodex(line) -} - -// LogFileName implements loop.Provider. -func (p *CodexProvider) LogFileName() string { return "codex.log" } - -// CleanOutput implements loop.Provider - Codex doesn't use a special format. -func (p *CodexProvider) CleanOutput(output string) string { return output } diff --git a/internal/agent/codex_test.go b/internal/agent/codex_test.go deleted file mode 100644 index 77cc2d2f..00000000 --- a/internal/agent/codex_test.go +++ /dev/null @@ -1,89 +0,0 @@ -package agent - -import ( - "context" - "testing" - - "github.com/minicodemonkey/chief/internal/loop" -) - -func TestCodexProvider_Name(t *testing.T) { - p := NewCodexProvider("") - if p.Name() != "Codex" { - t.Errorf("Name() = %q, want Codex", p.Name()) - } -} - -func TestCodexProvider_CLIPath(t *testing.T) { - p := NewCodexProvider("") - if p.CLIPath() != "codex" { - t.Errorf("CLIPath() empty arg = %q, want codex", p.CLIPath()) - } - p2 := NewCodexProvider("/usr/local/bin/codex") - if p2.CLIPath() != "/usr/local/bin/codex" { - t.Errorf("CLIPath() custom = %q, want /usr/local/bin/codex", p2.CLIPath()) - } -} - -func TestCodexProvider_LogFileName(t *testing.T) { - p := NewCodexProvider("") - if p.LogFileName() != "codex.log" { - t.Errorf("LogFileName() = %q, want codex.log", p.LogFileName()) - } -} - -func TestCodexProvider_LoopCommand(t *testing.T) { - ctx := context.Background() - p := NewCodexProvider("/bin/codex") - cmd := p.LoopCommand(ctx, "hello world", "/work/dir") - - if cmd.Path != "/bin/codex" { - t.Errorf("LoopCommand Path = %q, want /bin/codex", cmd.Path) - } - wantArgs := []string{"/bin/codex", "exec", "--json", "--yolo", "--skip-git-repo-check", "-C", "/work/dir", "-"} - if len(cmd.Args) != len(wantArgs) { - t.Fatalf("LoopCommand Args len = %d, want %d: %v", len(cmd.Args), len(wantArgs), cmd.Args) - } - for i, w := range wantArgs { - if cmd.Args[i] != w { - t.Errorf("LoopCommand Args[%d] = %q, want %q", i, cmd.Args[i], w) - } - } - if cmd.Dir != "/work/dir" { - t.Errorf("LoopCommand Dir = %q, want /work/dir", cmd.Dir) - } - if cmd.Stdin == nil { - t.Error("LoopCommand Stdin must be set (prompt on stdin)") - } - // Stdin should contain the prompt - // We can't easily read cmd.Stdin without running; just check it's non-nil (done above) -} - -func TestCodexProvider_InteractiveCommand(t *testing.T) { - p := NewCodexProvider("codex") - cmd := p.InteractiveCommand("/work", "my prompt") - if cmd.Dir != "/work" { - t.Errorf("InteractiveCommand Dir = %q, want /work", cmd.Dir) - } - wantInteractiveArgs := []string{"codex", "my prompt"} - if len(cmd.Args) != len(wantInteractiveArgs) { - t.Fatalf("InteractiveCommand Args len = %d, want %d: %v", len(cmd.Args), len(wantInteractiveArgs), cmd.Args) - } - for i, w := range wantInteractiveArgs { - if cmd.Args[i] != w { - t.Errorf("InteractiveCommand Args[%d] = %q, want %q", i, cmd.Args[i], w) - } - } -} - -func TestCodexProvider_ParseLine(t *testing.T) { - p := NewCodexProvider("") - // thread.started -> EventIterationStart - e := p.ParseLine(`{"type":"thread.started"}`) - if e == nil { - t.Fatal("ParseLine(thread.started) returned nil") - } - if e.Type != loop.EventIterationStart { - t.Errorf("ParseLine(thread.started) Type = %v, want EventIterationStart", e.Type) - } -} diff --git a/internal/agent/cursor.go b/internal/agent/cursor.go deleted file mode 100644 index cc6c430b..00000000 --- a/internal/agent/cursor.go +++ /dev/null @@ -1,95 +0,0 @@ -package agent - -import ( - "context" - "encoding/json" - "os/exec" - "strings" - - "github.com/minicodemonkey/chief/internal/loop" -) - -// CursorProvider implements loop.Provider for the Cursor CLI (agent). -type CursorProvider struct { - cliPath string -} - -// NewCursorProvider returns a Provider for the Cursor CLI. -// If cliPath is empty, "agent" is used. -func NewCursorProvider(cliPath string) *CursorProvider { - if cliPath == "" { - cliPath = "agent" - } - return &CursorProvider{cliPath: cliPath} -} - -// Name implements loop.Provider. -func (p *CursorProvider) Name() string { return "Cursor" } - -// CLIPath implements loop.Provider. -func (p *CursorProvider) CLIPath() string { return p.cliPath } - -// LoopCommand implements loop.Provider. -// Prompt is supplied via stdin; Cursor CLI reads it when -p has no argument. -func (p *CursorProvider) LoopCommand(ctx context.Context, prompt, workDir string) *exec.Cmd { - cmd := exec.CommandContext(ctx, p.cliPath, - "-p", - "--output-format", "stream-json", - "--force", - "--workspace", workDir, - "--trust", - ) - cmd.Dir = workDir - cmd.Stdin = strings.NewReader(prompt) - return cmd -} - -// InteractiveCommand implements loop.Provider. -func (p *CursorProvider) InteractiveCommand(workDir, prompt string) *exec.Cmd { - cmd := exec.Command(p.cliPath, prompt) - cmd.Dir = workDir - return cmd -} - -// ParseLine implements loop.Provider. -func (p *CursorProvider) ParseLine(line string) *loop.Event { - return loop.ParseLineCursor(line) -} - -// LogFileName implements loop.Provider. -func (p *CursorProvider) LogFileName() string { return "cursor.log" } - -// cursorResultLine is the structure for Cursor's result/success JSON lines. -type cursorResultLine struct { - Type string `json:"type"` - Subtype string `json:"subtype,omitempty"` - Result string `json:"result,omitempty"` -} - -// CleanOutput extracts the result from Cursor's json or stream-json output. -// For stream-json, finds the last type "result", subtype "success" and returns its result field. -// For single-line json, parses and returns result. -func (p *CursorProvider) CleanOutput(output string) string { - output = strings.TrimSpace(output) - if output == "" { - return output - } - // Try single JSON object (json output format) - var single cursorResultLine - if json.Unmarshal([]byte(output), &single) == nil && single.Type == "result" && single.Subtype == "success" && single.Result != "" { - return single.Result - } - // NDJSON: find last result/success line - lines := strings.Split(output, "\n") - for i := len(lines) - 1; i >= 0; i-- { - line := strings.TrimSpace(lines[i]) - if line == "" { - continue - } - var ev cursorResultLine - if json.Unmarshal([]byte(line), &ev) == nil && ev.Type == "result" && ev.Subtype == "success" && ev.Result != "" { - return ev.Result - } - } - return output -} diff --git a/internal/agent/cursor_test.go b/internal/agent/cursor_test.go deleted file mode 100644 index 5ba57a70..00000000 --- a/internal/agent/cursor_test.go +++ /dev/null @@ -1,101 +0,0 @@ -package agent - -import ( - "context" - "testing" - - "github.com/minicodemonkey/chief/internal/loop" -) - -func TestCursorProvider_Name(t *testing.T) { - p := NewCursorProvider("") - if p.Name() != "Cursor" { - t.Errorf("Name() = %q, want Cursor", p.Name()) - } -} - -func TestCursorProvider_CLIPath(t *testing.T) { - p := NewCursorProvider("") - if p.CLIPath() != "agent" { - t.Errorf("CLIPath() empty arg = %q, want agent", p.CLIPath()) - } - p2 := NewCursorProvider("/usr/local/bin/agent") - if p2.CLIPath() != "/usr/local/bin/agent" { - t.Errorf("CLIPath() custom = %q, want /usr/local/bin/agent", p2.CLIPath()) - } -} - -func TestCursorProvider_LogFileName(t *testing.T) { - p := NewCursorProvider("") - if p.LogFileName() != "cursor.log" { - t.Errorf("LogFileName() = %q, want cursor.log", p.LogFileName()) - } -} - -func TestCursorProvider_LoopCommand(t *testing.T) { - ctx := context.Background() - p := NewCursorProvider("/bin/agent") - cmd := p.LoopCommand(ctx, "hello world", "/work/dir") - - if cmd.Path != "/bin/agent" { - t.Errorf("LoopCommand Path = %q, want /bin/agent", cmd.Path) - } - wantArgs := []string{"/bin/agent", "-p", "--output-format", "stream-json", "--force", "--workspace", "/work/dir", "--trust"} - if len(cmd.Args) != len(wantArgs) { - t.Fatalf("LoopCommand Args len = %d, want %d: %v", len(cmd.Args), len(wantArgs), cmd.Args) - } - for i, w := range wantArgs { - if cmd.Args[i] != w { - t.Errorf("LoopCommand Args[%d] = %q, want %q", i, cmd.Args[i], w) - } - } - if cmd.Dir != "/work/dir" { - t.Errorf("LoopCommand Dir = %q, want /work/dir", cmd.Dir) - } - if cmd.Stdin == nil { - t.Error("LoopCommand Stdin must be set (prompt via stdin)") - } -} - -func TestCursorProvider_InteractiveCommand(t *testing.T) { - p := NewCursorProvider("/bin/agent") - cmd := p.InteractiveCommand("/work", "my prompt") - if cmd.Dir != "/work" { - t.Errorf("InteractiveCommand Dir = %q, want /work", cmd.Dir) - } - if len(cmd.Args) != 2 || cmd.Args[0] != "/bin/agent" || cmd.Args[1] != "my prompt" { - t.Errorf("InteractiveCommand Args = %v, want [/bin/agent my prompt]", cmd.Args) - } -} - -func TestCursorProvider_ParseLine(t *testing.T) { - p := NewCursorProvider("") - line := `{"type":"system","subtype":"init","session_id":"x"}` - e := p.ParseLine(line) - if e == nil { - t.Fatal("ParseLine(system init) returned nil") - } - if e.Type != loop.EventIterationStart { - t.Errorf("ParseLine(system init) Type = %v, want EventIterationStart", e.Type) - } -} - -func TestCursorProvider_CleanOutput(t *testing.T) { - p := NewCursorProvider("") - // NDJSON: last result/success - ndjson := `{"type":"system","subtype":"init"} -{"type":"result","subtype":"success","result":"final answer","session_id":"x"}` - if got := p.CleanOutput(ndjson); got != "final answer" { - t.Errorf("CleanOutput(NDJSON) = %q, want final answer", got) - } - // Single JSON result - single := `{"type":"result","subtype":"success","result":"single result","session_id":"x"}` - if got := p.CleanOutput(single); got != "single result" { - t.Errorf("CleanOutput(single JSON) = %q, want single result", got) - } - // No result: return as-is - plain := "plain text" - if got := p.CleanOutput(plain); got != plain { - t.Errorf("CleanOutput(plain) = %q, want %q", got, plain) - } -} diff --git a/internal/agent/opencode.go b/internal/agent/opencode.go deleted file mode 100644 index 25f2a10c..00000000 --- a/internal/agent/opencode.go +++ /dev/null @@ -1,73 +0,0 @@ -package agent - -import ( - "context" - "encoding/json" - "os/exec" - "strings" - - "github.com/minicodemonkey/chief/internal/loop" -) - -type OpenCodeProvider struct { - cliPath string -} - -func NewOpenCodeProvider(cliPath string) *OpenCodeProvider { - if cliPath == "" { - cliPath = "opencode" - } - return &OpenCodeProvider{cliPath: cliPath} -} - -func (p *OpenCodeProvider) Name() string { return "OpenCode" } - -func (p *OpenCodeProvider) CLIPath() string { return p.cliPath } - -func (p *OpenCodeProvider) LoopCommand(ctx context.Context, prompt, workDir string) *exec.Cmd { - cmd := exec.CommandContext(ctx, p.cliPath, "run", "--format", "json", prompt) - cmd.Dir = workDir - return cmd -} - -func (p *OpenCodeProvider) InteractiveCommand(workDir, prompt string) *exec.Cmd { - cmd := exec.Command(p.cliPath, "--prompt", prompt) - cmd.Dir = workDir - return cmd -} - -func (p *OpenCodeProvider) ParseLine(line string) *loop.Event { - return loop.ParseLineOpenCode(line) -} - -func (p *OpenCodeProvider) LogFileName() string { return "opencode.log" } - -// CleanOutput extracts JSON from opencode's NDJSON output format. -// It looks for the last "text" event line and returns its part.text content. -func (p *OpenCodeProvider) CleanOutput(output string) string { - output = strings.TrimSpace(output) - if !strings.Contains(output, "\n") { - return output - } - - var lastText string - for _, line := range strings.Split(output, "\n") { - line = strings.TrimSpace(line) - if line == "" { - continue - } - var ev struct { - Type string `json:"type"` - Part struct { - Text string `json:"text"` - } `json:"part"` - } - if json.Unmarshal([]byte(line), &ev) == nil && ev.Type == "text" && ev.Part.Text != "" { - lastText = ev.Part.Text - } - } - if lastText != "" { - return lastText - } - return output -} diff --git a/internal/agent/opencode_test.go b/internal/agent/opencode_test.go deleted file mode 100644 index c7f2ea68..00000000 --- a/internal/agent/opencode_test.go +++ /dev/null @@ -1,118 +0,0 @@ -package agent - -import ( - "context" - "testing" - - "github.com/minicodemonkey/chief/internal/loop" -) - -func TestOpenCodeProvider_Name(t *testing.T) { - p := NewOpenCodeProvider("") - if p.Name() != "OpenCode" { - t.Errorf("Name() = %q, want OpenCode", p.Name()) - } -} - -func TestOpenCodeProvider_CLIPath(t *testing.T) { - p := NewOpenCodeProvider("") - if p.CLIPath() != "opencode" { - t.Errorf("CLIPath() empty arg = %q, want opencode", p.CLIPath()) - } - p2 := NewOpenCodeProvider("/usr/local/bin/opencode") - if p2.CLIPath() != "/usr/local/bin/opencode" { - t.Errorf("CLIPath() custom = %q, want /usr/local/bin/opencode", p2.CLIPath()) - } -} - -func TestOpenCodeProvider_LogFileName(t *testing.T) { - p := NewOpenCodeProvider("") - if p.LogFileName() != "opencode.log" { - t.Errorf("LogFileName() = %q, want opencode.log", p.LogFileName()) - } -} - -func TestOpenCodeProvider_LoopCommand(t *testing.T) { - ctx := context.Background() - p := NewOpenCodeProvider("/bin/opencode") - cmd := p.LoopCommand(ctx, "hello world", "/work/dir") - - if cmd.Path != "/bin/opencode" { - t.Errorf("LoopCommand Path = %q, want /bin/opencode", cmd.Path) - } - wantArgs := []string{"/bin/opencode", "run", "--format", "json", "hello world"} - if len(cmd.Args) != len(wantArgs) { - t.Fatalf("LoopCommand Args len = %d, want %d: %v", len(cmd.Args), len(wantArgs), cmd.Args) - } - for i, w := range wantArgs { - if cmd.Args[i] != w { - t.Errorf("LoopCommand Args[%d] = %q, want %q", i, cmd.Args[i], w) - } - } - if cmd.Dir != "/work/dir" { - t.Errorf("LoopCommand Dir = %q, want /work/dir", cmd.Dir) - } -} - -func TestOpenCodeProvider_CleanOutput_PlainText(t *testing.T) { - p := NewOpenCodeProvider("") - // Non-NDJSON input should be returned as-is - input := `{"project": "test"}` - got := p.CleanOutput(input) - if got != input { - t.Errorf("CleanOutput(plain) = %q, want %q", got, input) - } -} - -func TestOpenCodeProvider_CleanOutput_NDJSON(t *testing.T) { - p := NewOpenCodeProvider("") - input := `{"type":"step_start","timestamp":1234,"sessionID":"ses_1"} -{"type":"text","timestamp":1235,"sessionID":"ses_1","part":{"id":"prt_1","type":"text","text":"hello world"}} -{"type":"step_finish","timestamp":1236,"sessionID":"ses_1","part":{"id":"prt_2","reason":"stop"}}` - got := p.CleanOutput(input) - if got != "hello world" { - t.Errorf("CleanOutput(ndjson) = %q, want %q", got, "hello world") - } -} - -func TestOpenCodeProvider_CleanOutput_LastTextWins(t *testing.T) { - p := NewOpenCodeProvider("") - input := `{"type":"text","timestamp":1,"sessionID":"s","part":{"id":"a","type":"text","text":"first"}} -{"type":"text","timestamp":2,"sessionID":"s","part":{"id":"b","type":"text","text":"second"}}` - got := p.CleanOutput(input) - if got != "second" { - t.Errorf("CleanOutput(multi-text) = %q, want %q", got, "second") - } -} - -func TestOpenCodeProvider_CleanOutput_NoTextEvent(t *testing.T) { - p := NewOpenCodeProvider("") - input := `{"type":"step_start","timestamp":1,"sessionID":"s"} -{"type":"step_finish","timestamp":2,"sessionID":"s","part":{"id":"a","reason":"stop"}}` - got := p.CleanOutput(input) - if got != input { - t.Errorf("CleanOutput(no-text) should return original, got %q", got) - } -} - -func TestOpenCodeProvider_InteractiveCommand(t *testing.T) { - p := NewOpenCodeProvider("opencode") - cmd := p.InteractiveCommand("/work", "my prompt") - if cmd.Dir != "/work" { - t.Errorf("InteractiveCommand Dir = %q, want /work", cmd.Dir) - } - if len(cmd.Args) != 3 || cmd.Args[0] != "opencode" || cmd.Args[1] != "--prompt" || cmd.Args[2] != "my prompt" { - t.Errorf("InteractiveCommand Args = %v, want [opencode --prompt 'my prompt']", cmd.Args) - } -} - -func TestOpenCodeProvider_ParseLine(t *testing.T) { - p := NewOpenCodeProvider("") - e := p.ParseLine(`{"type":"step_start","timestamp":1234567890,"sessionID":"ses_test123"}`) - if e == nil { - t.Fatal("ParseLine(step_start) returned nil") - } - if e.Type != loop.EventIterationStart { - t.Errorf("ParseLine(step_start) Type = %v, want EventIterationStart", e.Type) - } -} diff --git a/internal/agent/resolve.go b/internal/agent/resolve.go deleted file mode 100644 index 2458740c..00000000 --- a/internal/agent/resolve.go +++ /dev/null @@ -1,56 +0,0 @@ -package agent - -import ( - "fmt" - "os" - "os/exec" - "strings" - - "github.com/minicodemonkey/chief/internal/config" - "github.com/minicodemonkey/chief/internal/loop" -) - -// Resolve returns the agent Provider using priority: flagAgent > CHIEF_AGENT env > config > "claude". -// flagPath overrides the CLI path when non-empty (flag > CHIEF_AGENT_PATH > config agent.cliPath). -// Returns an error if the resolved provider name is not recognised. -func Resolve(flagAgent, flagPath string, cfg *config.Config) (loop.Provider, error) { - providerName := "claude" - if flagAgent != "" { - providerName = strings.ToLower(strings.TrimSpace(flagAgent)) - } else if v := os.Getenv("CHIEF_AGENT"); v != "" { - providerName = strings.ToLower(strings.TrimSpace(v)) - } else if cfg != nil && cfg.Agent.Provider != "" { - providerName = strings.ToLower(strings.TrimSpace(cfg.Agent.Provider)) - } - - cliPath := "" - if flagPath != "" { - cliPath = flagPath - } else if v := os.Getenv("CHIEF_AGENT_PATH"); v != "" { - cliPath = strings.TrimSpace(v) - } else if cfg != nil && cfg.Agent.CLIPath != "" { - cliPath = strings.TrimSpace(cfg.Agent.CLIPath) - } - - switch providerName { - case "claude": - return NewClaudeProvider(cliPath), nil - case "codex": - return NewCodexProvider(cliPath), nil - case "opencode": - return NewOpenCodeProvider(cliPath), nil - case "cursor": - return NewCursorProvider(cliPath), nil - default: - return nil, fmt.Errorf("unknown agent provider %q: expected \"claude\", \"codex\", \"opencode\", or \"cursor\"", providerName) - } -} - -// CheckInstalled verifies that the provider's CLI binary is found in PATH (or at cliPath). -func CheckInstalled(p loop.Provider) error { - _, err := exec.LookPath(p.CLIPath()) - if err != nil { - return fmt.Errorf("%s CLI not found in PATH. Install it or set agent.cliPath in .chief/config.yaml", p.Name()) - } - return nil -} diff --git a/internal/agent/resolve_test.go b/internal/agent/resolve_test.go deleted file mode 100644 index 32796f19..00000000 --- a/internal/agent/resolve_test.go +++ /dev/null @@ -1,210 +0,0 @@ -package agent - -import ( - "os" - "os/exec" - "path/filepath" - "strings" - "testing" - - "github.com/minicodemonkey/chief/internal/config" - "github.com/minicodemonkey/chief/internal/loop" -) - -func mustResolve(t *testing.T, flagAgent, flagPath string, cfg *config.Config) loop.Provider { - t.Helper() - p, err := Resolve(flagAgent, flagPath, cfg) - if err != nil { - t.Fatalf("Resolve(%q, %q, cfg) unexpected error: %v", flagAgent, flagPath, err) - } - return p -} - -func TestResolve_priority(t *testing.T) { - // Default: no flag, no env, nil config -> Claude - got := mustResolve(t, "", "", nil) - if got.Name() != "Claude" { - t.Errorf("Resolve(_, _, nil) name = %q, want Claude", got.Name()) - } - if got.CLIPath() != "claude" { - t.Errorf("Resolve(_, _, nil) CLIPath = %q, want claude", got.CLIPath()) - } - - // Flag overrides everything - got = mustResolve(t, "codex", "", nil) - if got.Name() != "Codex" { - t.Errorf("Resolve(codex, _, nil) name = %q, want Codex", got.Name()) - } - - // Config only (no flag, no env) - cfg := &config.Config{} - cfg.Agent.Provider = "codex" - cfg.Agent.CLIPath = "/usr/local/bin/codex" - got = mustResolve(t, "", "", cfg) - if got.Name() != "Codex" { - t.Errorf("Resolve(_, _, config codex) name = %q, want Codex", got.Name()) - } - if got.CLIPath() != "/usr/local/bin/codex" { - t.Errorf("Resolve(_, _, config) CLIPath = %q, want /usr/local/bin/codex", got.CLIPath()) - } - - // Flag overrides config - got = mustResolve(t, "claude", "", cfg) - if got.Name() != "Claude" { - t.Errorf("Resolve(claude, _, config codex) name = %q, want Claude", got.Name()) - } - // flag path overrides config path - got = mustResolve(t, "codex", "/opt/codex", cfg) - if got.CLIPath() != "/opt/codex" { - t.Errorf("Resolve(codex, /opt/codex, cfg) CLIPath = %q, want /opt/codex", got.CLIPath()) - } -} - -func TestResolve_env(t *testing.T) { - const keyAgent = "CHIEF_AGENT" - const keyPath = "CHIEF_AGENT_PATH" - saveAgent := os.Getenv(keyAgent) - savePath := os.Getenv(keyPath) - defer func() { - if saveAgent != "" { - os.Setenv(keyAgent, saveAgent) - } else { - os.Unsetenv(keyAgent) - } - if savePath != "" { - os.Setenv(keyPath, savePath) - } else { - os.Unsetenv(keyPath) - } - }() - - os.Unsetenv(keyAgent) - os.Unsetenv(keyPath) - - // Env provider when no flag - os.Setenv(keyAgent, "codex") - got := mustResolve(t, "", "", nil) - if got.Name() != "Codex" { - t.Errorf("with CHIEF_AGENT=codex, name = %q, want Codex", got.Name()) - } - os.Unsetenv(keyAgent) - - // Env path when no flag path - os.Setenv(keyAgent, "codex") - os.Setenv(keyPath, "/env/codex") - got = mustResolve(t, "", "", nil) - if got.CLIPath() != "/env/codex" { - t.Errorf("with CHIEF_AGENT_PATH, CLIPath = %q, want /env/codex", got.CLIPath()) - } - os.Unsetenv(keyPath) - os.Unsetenv(keyAgent) -} - -func TestResolve_normalize(t *testing.T) { - got := mustResolve(t, " CODEX ", "", nil) - if got.Name() != "Codex" { - t.Errorf("Resolve(' CODEX ') name = %q, want Codex", got.Name()) - } -} - -func TestResolve_opencode(t *testing.T) { - // Test OpenCode provider resolution - got := mustResolve(t, "opencode", "", nil) - if got.Name() != "OpenCode" { - t.Errorf("Resolve(opencode) name = %q, want OpenCode", got.Name()) - } - if got.CLIPath() != "opencode" { - t.Errorf("Resolve(opencode) CLIPath = %q, want opencode", got.CLIPath()) - } - - // Test OpenCode with custom path - got = mustResolve(t, "opencode", "/usr/local/bin/opencode", nil) - if got.CLIPath() != "/usr/local/bin/opencode" { - t.Errorf("Resolve(opencode, /usr/local/bin/opencode) CLIPath = %q, want /usr/local/bin/opencode", got.CLIPath()) - } - - // Test from config - cfg := &config.Config{} - cfg.Agent.Provider = "opencode" - cfg.Agent.CLIPath = "/opt/opencode" - got = mustResolve(t, "", "", cfg) - if got.Name() != "OpenCode" { - t.Errorf("Resolve(_, _, config opencode) name = %q, want OpenCode", got.Name()) - } - if got.CLIPath() != "/opt/opencode" { - t.Errorf("Resolve(_, _, config opencode) CLIPath = %q, want /opt/opencode", got.CLIPath()) - } -} - -func TestResolve_cursor(t *testing.T) { - got := mustResolve(t, "cursor", "", nil) - if got.Name() != "Cursor" { - t.Errorf("Resolve(cursor) name = %q, want Cursor", got.Name()) - } - if got.CLIPath() != "agent" { - t.Errorf("Resolve(cursor) CLIPath = %q, want agent", got.CLIPath()) - } - got = mustResolve(t, "cursor", "/usr/local/bin/agent", nil) - if got.CLIPath() != "/usr/local/bin/agent" { - t.Errorf("Resolve(cursor, path) CLIPath = %q, want /usr/local/bin/agent", got.CLIPath()) - } -} - -func TestResolve_unknownProvider(t *testing.T) { - _, err := Resolve("typo", "", nil) - if err == nil { - t.Fatal("Resolve(typo) expected error, got nil") - } - if !strings.Contains(err.Error(), "typo") { - t.Errorf("error should mention the bad provider name: %v", err) - } -} - -func TestCheckInstalled_notFound(t *testing.T) { - // Use a path that does not exist - p := NewCodexProvider("/nonexistent/codex-binary-that-does-not-exist") - err := CheckInstalled(p) - if err == nil { - t.Error("CheckInstalled(nonexistent) expected error, got nil") - } - if err != nil && !strings.Contains(err.Error(), "Codex") { - t.Errorf("CheckInstalled error should mention Codex: %v", err) - } -} - -func TestCheckInstalled_found(t *testing.T) { - // Go test binary is in PATH - goPath, err := exec.LookPath("go") - if err != nil { - t.Skip("go not in PATH, skipping CheckInstalled found test") - } - p := NewClaudeProvider(goPath) // abuse: use "go" as cli path to get a binary that exists - err = CheckInstalled(p) - if err != nil { - t.Errorf("CheckInstalled(existing binary) err = %v", err) - } -} - -func TestResolve_configFile(t *testing.T) { - dir := t.TempDir() - cfgPath := filepath.Join(dir, ".chief", "config.yaml") - if err := os.MkdirAll(filepath.Dir(cfgPath), 0o755); err != nil { - t.Fatal(err) - } - const yamlContent = ` -agent: - provider: codex - cliPath: /usr/local/bin/codex -` - if err := os.WriteFile(cfgPath, []byte(yamlContent), 0o644); err != nil { - t.Fatal(err) - } - cfg, err := config.Load(dir) - if err != nil { - t.Fatal(err) - } - got := mustResolve(t, "", "", cfg) - if got.Name() != "Codex" || got.CLIPath() != "/usr/local/bin/codex" { - t.Errorf("Resolve from config: name=%q path=%q", got.Name(), got.CLIPath()) - } -} diff --git a/internal/auth/auth.go b/internal/auth/auth.go new file mode 100644 index 00000000..81b1ba4c --- /dev/null +++ b/internal/auth/auth.go @@ -0,0 +1,255 @@ +package auth + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "net/http" + "os" + "path/filepath" + "sync" + "time" + + "gopkg.in/yaml.v3" +) + +const credentialsFile = "credentials.yaml" + +const defaultBaseURL = "https://uplink.chiefloop.com" + +// ErrNotLoggedIn is returned when no credentials file exists. +var ErrNotLoggedIn = errors.New("not logged in — run 'chief login' first") + +// ErrSessionExpired is returned when the refresh token is revoked or expired. +var ErrSessionExpired = errors.New("session expired — run 'chief login' again") + +// refreshMu protects concurrent token refresh operations. +var refreshMu sync.Mutex + +// Credentials holds authentication token data for uplink.chiefloop.com. +type Credentials struct { + AccessToken string `yaml:"access_token"` + RefreshToken string `yaml:"refresh_token"` + ExpiresAt time.Time `yaml:"expires_at"` + DeviceName string `yaml:"device_name"` + User string `yaml:"user"` + WSURL string `yaml:"ws_url,omitempty"` +} + +// IsExpired returns true if the access token has expired. +func (c *Credentials) IsExpired() bool { + return time.Now().After(c.ExpiresAt) +} + +// IsNearExpiry returns true if the access token will expire within the given duration. +func (c *Credentials) IsNearExpiry(d time.Duration) bool { + return time.Now().Add(d).After(c.ExpiresAt) +} + +// credentialsDir returns the path to the ~/.chief directory. +func credentialsDir() (string, error) { + home, err := os.UserHomeDir() + if err != nil { + return "", fmt.Errorf("determining home directory: %w", err) + } + return filepath.Join(home, ".chief"), nil +} + +// credentialsPath returns the full path to ~/.chief/credentials.yaml. +func credentialsPath() (string, error) { + dir, err := credentialsDir() + if err != nil { + return "", err + } + return filepath.Join(dir, credentialsFile), nil +} + +// LoadCredentials reads credentials from ~/.chief/credentials.yaml. +// Returns ErrNotLoggedIn when the file does not exist. +func LoadCredentials() (*Credentials, error) { + path, err := credentialsPath() + if err != nil { + return nil, err + } + + data, err := os.ReadFile(path) + if err != nil { + if os.IsNotExist(err) { + return nil, ErrNotLoggedIn + } + return nil, fmt.Errorf("reading credentials: %w", err) + } + + var creds Credentials + if err := yaml.Unmarshal(data, &creds); err != nil { + return nil, fmt.Errorf("parsing credentials: %w", err) + } + + return &creds, nil +} + +// SaveCredentials writes credentials to ~/.chief/credentials.yaml atomically. +// It writes to a temporary file first, then renames it into place. +// The file is created with 0600 permissions (owner read/write only). +func SaveCredentials(creds *Credentials) error { + path, err := credentialsPath() + if err != nil { + return err + } + + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0o755); err != nil { + return fmt.Errorf("creating credentials directory: %w", err) + } + + data, err := yaml.Marshal(creds) + if err != nil { + return fmt.Errorf("marshaling credentials: %w", err) + } + + // Write to temp file in the same directory for atomic rename. + tmp, err := os.CreateTemp(dir, "credentials-*.yaml") + if err != nil { + return fmt.Errorf("creating temp file: %w", err) + } + tmpPath := tmp.Name() + + if err := os.Chmod(tmpPath, 0o600); err != nil { + tmp.Close() + os.Remove(tmpPath) + return fmt.Errorf("setting temp file permissions: %w", err) + } + + if _, err := tmp.Write(data); err != nil { + tmp.Close() + os.Remove(tmpPath) + return fmt.Errorf("writing temp file: %w", err) + } + + if err := tmp.Close(); err != nil { + os.Remove(tmpPath) + return fmt.Errorf("closing temp file: %w", err) + } + + if err := os.Rename(tmpPath, path); err != nil { + os.Remove(tmpPath) + return fmt.Errorf("renaming temp file: %w", err) + } + + return nil +} + +// DeleteCredentials removes the credentials file. +// Returns nil if the file does not exist. +func DeleteCredentials() error { + path, err := credentialsPath() + if err != nil { + return err + } + + if err := os.Remove(path); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("removing credentials: %w", err) + } + + return nil +} + +// refreshResponse is the response from the token refresh endpoint. +type refreshResponse struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ExpiresIn int `json:"expires_in"` + WSURL string `json:"ws_url"` + Error string `json:"error"` +} + +// RefreshToken refreshes the access token using the refresh token. +// It is thread-safe (mutex-protected for concurrent use by serve). +// baseURL can be empty to use the default (https://uplink.chiefloop.com). +func RefreshToken(baseURL string) (*Credentials, error) { + refreshMu.Lock() + defer refreshMu.Unlock() + + creds, err := LoadCredentials() + if err != nil { + return nil, err + } + + // If token was already refreshed by another goroutine, return it. + if !creds.IsNearExpiry(5 * time.Minute) { + return creds, nil + } + + if baseURL == "" { + baseURL = defaultBaseURL + } + + reqBody, _ := json.Marshal(map[string]string{ + "grant_type": "refresh_token", + "refresh_token": creds.RefreshToken, + }) + + client := &http.Client{Timeout: 10 * time.Second} + req, err := http.NewRequest(http.MethodPost, baseURL+"/api/oauth/token", bytes.NewReader(reqBody)) + if err != nil { + return nil, fmt.Errorf("creating refresh request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + + resp, err := client.Do(req) + if err != nil { + return nil, fmt.Errorf("refreshing token: %w", err) + } + defer resp.Body.Close() + + var tokenResp refreshResponse + if err := json.NewDecoder(resp.Body).Decode(&tokenResp); err != nil { + return nil, fmt.Errorf("parsing refresh response: %w", err) + } + + if tokenResp.Error != "" || resp.StatusCode != http.StatusOK { + return nil, ErrSessionExpired + } + + creds.AccessToken = tokenResp.AccessToken + if tokenResp.RefreshToken != "" { + creds.RefreshToken = tokenResp.RefreshToken + } + if tokenResp.WSURL != "" { + creds.WSURL = tokenResp.WSURL + } + creds.ExpiresAt = time.Now().Add(time.Duration(tokenResp.ExpiresIn) * time.Second) + + if err := SaveCredentials(creds); err != nil { + return nil, fmt.Errorf("saving refreshed credentials: %w", err) + } + + return creds, nil +} + +// RevokeDevice calls the revocation endpoint to deauthorize the device server-side. +// baseURL can be empty to use the default (https://uplink.chiefloop.com). +func RevokeDevice(accessToken, baseURL string) error { + if baseURL == "" { + baseURL = defaultBaseURL + } + + reqBody, _ := json.Marshal(map[string]string{ + "access_token": accessToken, + }) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Post(baseURL+"/api/oauth/revoke", "application/json", bytes.NewReader(reqBody)) + if err != nil { + return fmt.Errorf("revoking device: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("revocation failed: server returned %s", resp.Status) + } + + return nil +} diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go new file mode 100644 index 00000000..15694304 --- /dev/null +++ b/internal/auth/auth_test.go @@ -0,0 +1,610 @@ +package auth + +import ( + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + "testing" + "time" +) + +// setTestHome overrides HOME so credentials are read/written inside t.TempDir(). +// It returns a cleanup function that restores the original HOME. +func setTestHome(t *testing.T, dir string) { + t.Helper() + orig := os.Getenv("HOME") + t.Setenv("HOME", dir) + t.Cleanup(func() { + os.Setenv("HOME", orig) + }) +} + +func TestLoadCredentials_NotLoggedIn(t *testing.T) { + setTestHome(t, t.TempDir()) + + _, err := LoadCredentials() + if !errors.Is(err, ErrNotLoggedIn) { + t.Fatalf("expected ErrNotLoggedIn, got %v", err) + } +} + +func TestSaveAndLoadCredentials(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + expires := time.Date(2026, 6, 15, 12, 0, 0, 0, time.UTC) + creds := &Credentials{ + AccessToken: "access-abc", + RefreshToken: "refresh-xyz", + ExpiresAt: expires, + DeviceName: "my-laptop", + User: "user@example.com", + } + + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + loaded, err := LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials failed: %v", err) + } + + if loaded.AccessToken != "access-abc" { + t.Errorf("expected access_token %q, got %q", "access-abc", loaded.AccessToken) + } + if loaded.RefreshToken != "refresh-xyz" { + t.Errorf("expected refresh_token %q, got %q", "refresh-xyz", loaded.RefreshToken) + } + if !loaded.ExpiresAt.Equal(expires) { + t.Errorf("expected expires_at %v, got %v", expires, loaded.ExpiresAt) + } + if loaded.DeviceName != "my-laptop" { + t.Errorf("expected device_name %q, got %q", "my-laptop", loaded.DeviceName) + } + if loaded.User != "user@example.com" { + t.Errorf("expected user %q, got %q", "user@example.com", loaded.User) + } +} + +func TestSaveCredentials_FilePermissions(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + creds := &Credentials{ + AccessToken: "token", + } + + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + path := filepath.Join(home, ".chief", "credentials.yaml") + info, err := os.Stat(path) + if err != nil { + t.Fatalf("Stat failed: %v", err) + } + + perm := info.Mode().Perm() + if perm != 0o600 { + t.Errorf("expected permissions 0600, got %04o", perm) + } +} + +func TestSaveCredentials_Atomic(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Save initial credentials. + initial := &Credentials{ + AccessToken: "first", + User: "user1", + } + if err := SaveCredentials(initial); err != nil { + t.Fatalf("SaveCredentials (initial) failed: %v", err) + } + + // Save updated credentials (should atomically replace). + updated := &Credentials{ + AccessToken: "second", + User: "user2", + } + if err := SaveCredentials(updated); err != nil { + t.Fatalf("SaveCredentials (updated) failed: %v", err) + } + + loaded, err := LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials failed: %v", err) + } + if loaded.AccessToken != "second" { + t.Errorf("expected access_token %q, got %q", "second", loaded.AccessToken) + } + if loaded.User != "user2" { + t.Errorf("expected user %q, got %q", "user2", loaded.User) + } +} + +func TestDeleteCredentials(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + creds := &Credentials{AccessToken: "to-delete"} + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + if err := DeleteCredentials(); err != nil { + t.Fatalf("DeleteCredentials failed: %v", err) + } + + _, err := LoadCredentials() + if !errors.Is(err, ErrNotLoggedIn) { + t.Fatalf("expected ErrNotLoggedIn after delete, got %v", err) + } +} + +func TestDeleteCredentials_NonExistent(t *testing.T) { + setTestHome(t, t.TempDir()) + + // Deleting when file doesn't exist should not error. + if err := DeleteCredentials(); err != nil { + t.Fatalf("DeleteCredentials on non-existent file failed: %v", err) + } +} + +func TestSaveLoadDeleteCycle(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // 1. Not logged in initially. + _, err := LoadCredentials() + if !errors.Is(err, ErrNotLoggedIn) { + t.Fatalf("expected ErrNotLoggedIn initially, got %v", err) + } + + // 2. Save credentials. + creds := &Credentials{ + AccessToken: "cycle-token", + RefreshToken: "cycle-refresh", + ExpiresAt: time.Now().Add(time.Hour), + DeviceName: "test-device", + User: "cycle-user", + } + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + // 3. Load and verify. + loaded, err := LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials failed: %v", err) + } + if loaded.AccessToken != "cycle-token" { + t.Errorf("expected access_token %q, got %q", "cycle-token", loaded.AccessToken) + } + + // 4. Delete. + if err := DeleteCredentials(); err != nil { + t.Fatalf("DeleteCredentials failed: %v", err) + } + + // 5. Not logged in again. + _, err = LoadCredentials() + if !errors.Is(err, ErrNotLoggedIn) { + t.Fatalf("expected ErrNotLoggedIn after delete, got %v", err) + } +} + +func TestIsExpired(t *testing.T) { + // Expired token. + expired := &Credentials{ + ExpiresAt: time.Now().Add(-time.Hour), + } + if !expired.IsExpired() { + t.Error("expected token to be expired") + } + + // Valid token. + valid := &Credentials{ + ExpiresAt: time.Now().Add(time.Hour), + } + if valid.IsExpired() { + t.Error("expected token to not be expired") + } +} + +func TestIsNearExpiry(t *testing.T) { + // Token expires in 3 minutes — should be near expiry within 5 minutes. + creds := &Credentials{ + ExpiresAt: time.Now().Add(3 * time.Minute), + } + + if !creds.IsNearExpiry(5 * time.Minute) { + t.Error("expected token to be near expiry within 5 minutes") + } + + if creds.IsNearExpiry(1 * time.Minute) { + t.Error("expected token to NOT be near expiry within 1 minute") + } +} + +func TestIsNearExpiry_AlreadyExpired(t *testing.T) { + creds := &Credentials{ + ExpiresAt: time.Now().Add(-time.Hour), + } + + if !creds.IsNearExpiry(5 * time.Minute) { + t.Error("expected already-expired token to be near expiry") + } +} + +func TestSaveCredentials_CreatesDirectory(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + chiefDir := filepath.Join(home, ".chief") + if _, err := os.Stat(chiefDir); !os.IsNotExist(err) { + t.Fatal("expected .chief directory to not exist initially") + } + + creds := &Credentials{AccessToken: "create-dir"} + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + info, err := os.Stat(chiefDir) + if err != nil { + t.Fatalf("expected .chief directory to exist after save, got: %v", err) + } + if !info.IsDir() { + t.Error("expected .chief to be a directory") + } +} + +func TestRefreshToken_Success(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Save credentials that are near expiry (within 5 minutes) + creds := &Credentials{ + AccessToken: "old-access-token", + RefreshToken: "test-refresh-token", + ExpiresAt: time.Now().Add(2 * time.Minute), // near expiry + DeviceName: "test-device", + User: "user@example.com", + } + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/oauth/token" { + if accept := r.Header.Get("Accept"); accept != "application/json" { + t.Errorf("expected Accept header %q, got %q", "application/json", accept) + } + var body map[string]string + json.NewDecoder(r.Body).Decode(&body) + if body["grant_type"] != "refresh_token" { + t.Errorf("expected grant_type %q, got %q", "refresh_token", body["grant_type"]) + } + if body["refresh_token"] != "test-refresh-token" { + t.Errorf("expected refresh_token %q, got %q", "test-refresh-token", body["refresh_token"]) + } + json.NewEncoder(w).Encode(refreshResponse{ + AccessToken: "new-access-token", + RefreshToken: "new-refresh-token", + ExpiresIn: 3600, + }) + return + } + http.NotFound(w, r) + })) + defer server.Close() + + refreshed, err := RefreshToken(server.URL) + if err != nil { + t.Fatalf("RefreshToken failed: %v", err) + } + + if refreshed.AccessToken != "new-access-token" { + t.Errorf("expected access_token %q, got %q", "new-access-token", refreshed.AccessToken) + } + if refreshed.RefreshToken != "new-refresh-token" { + t.Errorf("expected refresh_token %q, got %q", "new-refresh-token", refreshed.RefreshToken) + } + + // Verify credentials were persisted + loaded, err := LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials failed: %v", err) + } + if loaded.AccessToken != "new-access-token" { + t.Errorf("expected persisted access_token %q, got %q", "new-access-token", loaded.AccessToken) + } +} + +func TestRefreshToken_NotNearExpiry(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Save credentials that are NOT near expiry + creds := &Credentials{ + AccessToken: "valid-token", + RefreshToken: "refresh-token", + ExpiresAt: time.Now().Add(30 * time.Minute), + DeviceName: "test-device", + User: "user@example.com", + } + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + // Server should not be called since token is still valid + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + t.Error("server should not be called when token is not near expiry") + })) + defer server.Close() + + refreshed, err := RefreshToken(server.URL) + if err != nil { + t.Fatalf("RefreshToken failed: %v", err) + } + + if refreshed.AccessToken != "valid-token" { + t.Errorf("expected access_token %q, got %q", "valid-token", refreshed.AccessToken) + } +} + +func TestRefreshToken_SessionExpired(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + creds := &Credentials{ + AccessToken: "old-token", + RefreshToken: "revoked-refresh-token", + ExpiresAt: time.Now().Add(2 * time.Minute), + DeviceName: "test-device", + User: "user@example.com", + } + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/oauth/token" { + json.NewEncoder(w).Encode(refreshResponse{ + Error: "invalid_grant", + }) + return + } + http.NotFound(w, r) + })) + defer server.Close() + + _, err := RefreshToken(server.URL) + if !errors.Is(err, ErrSessionExpired) { + t.Fatalf("expected ErrSessionExpired, got %v", err) + } +} + +func TestRefreshToken_NotLoggedIn(t *testing.T) { + setTestHome(t, t.TempDir()) + + _, err := RefreshToken("") + if !errors.Is(err, ErrNotLoggedIn) { + t.Fatalf("expected ErrNotLoggedIn, got %v", err) + } +} + +func TestRefreshToken_ThreadSafe(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + creds := &Credentials{ + AccessToken: "old-token", + RefreshToken: "refresh-token", + ExpiresAt: time.Now().Add(2 * time.Minute), + DeviceName: "test-device", + User: "user@example.com", + } + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + var callCount int + var mu sync.Mutex + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/oauth/token" { + mu.Lock() + callCount++ + mu.Unlock() + json.NewEncoder(w).Encode(refreshResponse{ + AccessToken: "new-token", + RefreshToken: "new-refresh", + ExpiresIn: 3600, + }) + return + } + http.NotFound(w, r) + })) + defer server.Close() + + // Run multiple concurrent refreshes + var wg sync.WaitGroup + for i := 0; i < 5; i++ { + wg.Add(1) + go func() { + defer wg.Done() + _, err := RefreshToken(server.URL) + if err != nil { + t.Errorf("RefreshToken failed: %v", err) + } + }() + } + wg.Wait() + + // Only one actual refresh should have hit the server + // (the first one refreshes, subsequent ones see it's no longer near expiry) + mu.Lock() + count := callCount + mu.Unlock() + if count != 1 { + t.Errorf("expected exactly 1 server call, got %d", count) + } +} + +func TestSaveAndLoadCredentials_WithWSURL(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + creds := &Credentials{ + AccessToken: "access-abc", + RefreshToken: "refresh-xyz", + ExpiresAt: time.Date(2026, 6, 15, 12, 0, 0, 0, time.UTC), + DeviceName: "my-laptop", + User: "user@example.com", + WSURL: "wss://ws-abc123-reverb.laravel.cloud/ws/server", + } + + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + loaded, err := LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials failed: %v", err) + } + + if loaded.WSURL != "wss://ws-abc123-reverb.laravel.cloud/ws/server" { + t.Errorf("expected ws_url %q, got %q", "wss://ws-abc123-reverb.laravel.cloud/ws/server", loaded.WSURL) + } +} + +func TestRefreshToken_PreservesWSURL(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + creds := &Credentials{ + AccessToken: "old-access-token", + RefreshToken: "test-refresh-token", + ExpiresAt: time.Now().Add(2 * time.Minute), + DeviceName: "test-device", + User: "user@example.com", + WSURL: "wss://old-host/ws/server", + } + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/oauth/token" { + json.NewEncoder(w).Encode(refreshResponse{ + AccessToken: "new-access-token", + RefreshToken: "new-refresh-token", + ExpiresIn: 3600, + WSURL: "wss://new-host/ws/server", + }) + return + } + http.NotFound(w, r) + })) + defer server.Close() + + refreshed, err := RefreshToken(server.URL) + if err != nil { + t.Fatalf("RefreshToken failed: %v", err) + } + + if refreshed.WSURL != "wss://new-host/ws/server" { + t.Errorf("expected ws_url %q, got %q", "wss://new-host/ws/server", refreshed.WSURL) + } + + // Verify persisted + loaded, err := LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials failed: %v", err) + } + if loaded.WSURL != "wss://new-host/ws/server" { + t.Errorf("expected persisted ws_url %q, got %q", "wss://new-host/ws/server", loaded.WSURL) + } +} + +func TestRefreshToken_WSURLNotReturned_PreservesExisting(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + creds := &Credentials{ + AccessToken: "old-access-token", + RefreshToken: "test-refresh-token", + ExpiresAt: time.Now().Add(2 * time.Minute), + DeviceName: "test-device", + User: "user@example.com", + WSURL: "wss://existing-host/ws/server", + } + if err := SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/oauth/token" { + json.NewEncoder(w).Encode(refreshResponse{ + AccessToken: "new-access-token", + RefreshToken: "new-refresh-token", + ExpiresIn: 3600, + // WSURL intentionally omitted + }) + return + } + http.NotFound(w, r) + })) + defer server.Close() + + refreshed, err := RefreshToken(server.URL) + if err != nil { + t.Fatalf("RefreshToken failed: %v", err) + } + + if refreshed.WSURL != "wss://existing-host/ws/server" { + t.Errorf("expected existing ws_url to be preserved %q, got %q", "wss://existing-host/ws/server", refreshed.WSURL) + } +} + +func TestRevokeDevice_Success(t *testing.T) { + var receivedToken string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/oauth/revoke" { + var body map[string]string + json.NewDecoder(r.Body).Decode(&body) + receivedToken = body["access_token"] + w.WriteHeader(http.StatusOK) + return + } + http.NotFound(w, r) + })) + defer server.Close() + + err := RevokeDevice("my-token", server.URL) + if err != nil { + t.Fatalf("RevokeDevice failed: %v", err) + } + if receivedToken != "my-token" { + t.Errorf("expected token %q, got %q", "my-token", receivedToken) + } +} + +func TestRevokeDevice_ServerError(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + err := RevokeDevice("my-token", server.URL) + if err == nil { + t.Fatal("expected error for server error response") + } +} diff --git a/internal/cmd/clone.go b/internal/cmd/clone.go new file mode 100644 index 00000000..333ad397 --- /dev/null +++ b/internal/cmd/clone.go @@ -0,0 +1,251 @@ +package cmd + +import ( + "bufio" + "encoding/json" + "fmt" + "log" + "os" + "os/exec" + "path/filepath" + "regexp" + "strconv" + "strings" + + "github.com/minicodemonkey/chief/internal/workspace" + "github.com/minicodemonkey/chief/internal/ws" +) + +// handleCloneRepo handles a clone_repo request. +func handleCloneRepo(sender messageSender, scanner *workspace.Scanner, msg ws.Message) { + var req ws.CloneRepoMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing clone_repo message: %v", err) + return + } + + if req.URL == "" { + sendError(sender, ws.ErrCodeCloneFailed, "URL is required", msg.ID) + return + } + + workspaceDir := scanner.WorkspacePath() + + // Determine target directory name + dirName := req.DirectoryName + if dirName == "" { + dirName = inferDirName(req.URL) + } + + targetDir := filepath.Join(workspaceDir, dirName) + + // Check if target already exists + if _, err := os.Stat(targetDir); err == nil { + sendError(sender, ws.ErrCodeCloneFailed, + fmt.Sprintf("Directory %q already exists in workspace", dirName), msg.ID) + return + } + + // Run clone in a goroutine so we don't block the message loop + go runClone(sender, scanner, req.URL, dirName, workspaceDir) +} + +// inferDirName extracts a directory name from a git URL. +func inferDirName(url string) string { + // Remove trailing .git + url = strings.TrimSuffix(url, ".git") + // Remove trailing slash + url = strings.TrimRight(url, "/") + // Get the last path component + parts := strings.Split(url, "/") + if len(parts) > 0 { + name := parts[len(parts)-1] + // Also handle ssh-style urls like git@github.com:user/repo + if idx := strings.LastIndex(name, ":"); idx >= 0 { + name = name[idx+1:] + } + if name != "" { + return name + } + } + return "cloned-repo" +} + +// percentPattern matches git clone progress percentages. +var percentPattern = regexp.MustCompile(`(\d+)%`) + +// runClone executes the git clone and streams progress messages. +func runClone(sender messageSender, scanner *workspace.Scanner, url, dirName, workspaceDir string) { + cmd := exec.Command("git", "clone", "--progress", url, dirName) + cmd.Dir = workspaceDir + + // Git clone writes progress to stderr + stderr, err := cmd.StderrPipe() + if err != nil { + sendCloneComplete(sender, url, "", false, fmt.Sprintf("Failed to set up clone: %v", err)) + return + } + + if err := cmd.Start(); err != nil { + sendCloneComplete(sender, url, "", false, fmt.Sprintf("Failed to start clone: %v", err)) + return + } + + // Stream progress from stderr + stderrScanner := bufio.NewScanner(stderr) + stderrScanner.Split(scanGitProgress) + for stderrScanner.Scan() { + line := strings.TrimSpace(stderrScanner.Text()) + if line == "" { + continue + } + + percent := 0 + if matches := percentPattern.FindStringSubmatch(line); len(matches) > 1 { + percent, _ = strconv.Atoi(matches[1]) + } + + sendCloneProgress(sender, url, line, percent) + } + + if err := cmd.Wait(); err != nil { + sendCloneComplete(sender, url, "", false, fmt.Sprintf("Clone failed: %v", err)) + return + } + + // Trigger a rescan so the new project appears immediately + scanner.ScanAndUpdate() + + sendCloneComplete(sender, url, dirName, true, "") +} + +// scanGitProgress is a bufio.SplitFunc that splits on \r or \n, +// since git clone uses \r for progress updates. +func scanGitProgress(data []byte, atEOF bool) (advance int, token []byte, err error) { + if atEOF && len(data) == 0 { + return 0, nil, nil + } + // Find the first \r or \n + for i, b := range data { + if b == '\r' || b == '\n' { + return i + 1, data[:i], nil + } + } + if atEOF { + return len(data), data, nil + } + return 0, nil, nil +} + +// sendCloneProgress sends a clone_progress message. +func sendCloneProgress(sender messageSender, url, progressText string, percent int) { + if sender == nil { + return + } + envelope := ws.NewMessage(ws.TypeCloneProgress) + msg := ws.CloneProgressMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + URL: url, + ProgressText: progressText, + Percent: percent, + } + if err := sender.Send(msg); err != nil { + log.Printf("Error sending clone_progress: %v", err) + } +} + +// sendCloneComplete sends a clone_complete message. +func sendCloneComplete(sender messageSender, url, project string, success bool, errMsg string) { + if sender == nil { + return + } + envelope := ws.NewMessage(ws.TypeCloneComplete) + msg := ws.CloneCompleteMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + URL: url, + Success: success, + Error: errMsg, + Project: project, + } + if err := sender.Send(msg); err != nil { + log.Printf("Error sending clone_complete: %v", err) + } +} + +// handleCreateProject handles a create_project request. +func handleCreateProject(sender messageSender, scanner *workspace.Scanner, msg ws.Message) { + var req ws.CreateProjectMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing create_project message: %v", err) + return + } + + if req.Name == "" { + sendError(sender, ws.ErrCodeFilesystemError, "Project name is required", msg.ID) + return + } + + workspaceDir := scanner.WorkspacePath() + projectDir := filepath.Join(workspaceDir, req.Name) + + // Check if directory already exists + if _, err := os.Stat(projectDir); err == nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Directory %q already exists", req.Name), msg.ID) + return + } + + // Create the directory + if err := os.MkdirAll(projectDir, 0o755); err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Failed to create directory: %v", err), msg.ID) + return + } + + // Optionally run git init + if req.GitInit { + cmd := exec.Command("git", "init", projectDir) + if out, err := cmd.CombinedOutput(); err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("git init failed: %v\n%s", err, strings.TrimSpace(string(out))), msg.ID) + return + } + } + + // Trigger rescan so new project appears immediately + scanner.ScanAndUpdate() + + // Send updated project_state if git init was done (it's a discoverable project) + if req.GitInit { + project, found := scanner.FindProject(req.Name) + if found { + envelope := ws.NewMessage(ws.TypeProjectState) + psMsg := ws.ProjectStateMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: project, + } + if err := sender.Send(psMsg); err != nil { + log.Printf("Error sending project_state: %v", err) + } + return + } + } + + // Send a simple project_list update for non-git projects + envelope := ws.NewMessage(ws.TypeProjectList) + plMsg := ws.ProjectListMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Projects: scanner.Projects(), + } + if err := sender.Send(plMsg); err != nil { + log.Printf("Error sending project_list: %v", err) + } +} diff --git a/internal/cmd/clone_test.go b/internal/cmd/clone_test.go new file mode 100644 index 00000000..3a6e2602 --- /dev/null +++ b/internal/cmd/clone_test.go @@ -0,0 +1,531 @@ +package cmd + +import ( + "encoding/json" + "os" + "os/exec" + "path/filepath" + "strconv" + "strings" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/ws" +) + +func TestInferDirName(t *testing.T) { + tests := []struct { + url string + expected string + }{ + {"https://github.com/user/repo.git", "repo"}, + {"https://github.com/user/repo", "repo"}, + {"git@github.com:user/repo.git", "repo"}, + {"https://github.com/user/repo/", "repo"}, + {"https://github.com/user/my-project.git", "my-project"}, + {"git@github.com:org/my-lib.git", "my-lib"}, + } + + for _, tt := range tests { + t.Run(tt.url, func(t *testing.T) { + got := inferDirName(tt.url) + if got != tt.expected { + t.Errorf("inferDirName(%q) = %q, want %q", tt.url, got, tt.expected) + } + }) + } +} + +func TestHandleCloneRepo_Success(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a bare git repo to clone from + bareRepo := filepath.Join(home, "bare-repo.git") + cmd := exec.Command("git", "init", "--bare", bareRepo) + cmd.Env = append(os.Environ(), "GIT_CONFIG_GLOBAL=/dev/null", "GIT_CONFIG_SYSTEM=/dev/null") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git init --bare failed: %v\n%s", err, out) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send clone_repo request + cloneReq := map[string]interface{}{ + "type": "clone_repo", + "id": "req-clone-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "url": bareRepo, + } + if err := ms.sendCommand(cloneReq); err != nil { + t.Fatalf("failed to send clone command: %v", err) + } + + // Wait for clone_complete message + raw, err := ms.waitForMessageType("clone_complete", 5*time.Second) + if err != nil { + t.Fatalf("failed to receive clone_complete: %v", err) + } + + var cloneComplete map[string]interface{} + if err := json.Unmarshal(raw, &cloneComplete); err != nil { + t.Fatalf("failed to unmarshal clone_complete: %v", err) + } + + if cloneComplete["success"] != true { + t.Errorf("expected success=true, got %v (error: %v)", cloneComplete["success"], cloneComplete["error"]) + } + if cloneComplete["project"] != "bare-repo" { + t.Errorf("expected project 'bare-repo', got %v", cloneComplete["project"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify the directory was created + clonedDir := filepath.Join(workspaceDir, "bare-repo") + if _, err := os.Stat(filepath.Join(clonedDir, ".git")); os.IsNotExist(err) { + t.Error("cloned repository directory does not have .git") + } +} + +func TestHandleCloneRepo_CustomDirectoryName(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a bare git repo to clone from + bareRepo := filepath.Join(home, "source.git") + cmd := exec.Command("git", "init", "--bare", bareRepo) + cmd.Env = append(os.Environ(), "GIT_CONFIG_GLOBAL=/dev/null", "GIT_CONFIG_SYSTEM=/dev/null") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git init --bare failed: %v\n%s", err, out) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + cloneReq := map[string]interface{}{ + "type": "clone_repo", + "id": "req-clone-2", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "url": bareRepo, + "directory_name": "my-custom-name", + } + if err := ms.sendCommand(cloneReq); err != nil { + t.Fatalf("failed to send clone command: %v", err) + } + + // Wait for clone_complete message + raw, err := ms.waitForMessageType("clone_complete", 5*time.Second) + if err != nil { + t.Fatalf("failed to receive clone_complete: %v", err) + } + + var cloneComplete map[string]interface{} + if err := json.Unmarshal(raw, &cloneComplete); err != nil { + t.Fatalf("failed to unmarshal clone_complete: %v", err) + } + + if cloneComplete["success"] != true { + t.Errorf("expected success=true, got %v", cloneComplete["success"]) + } + if cloneComplete["project"] != "my-custom-name" { + t.Errorf("expected project 'my-custom-name', got %v", cloneComplete["project"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify directory exists under custom name + if _, err := os.Stat(filepath.Join(workspaceDir, "my-custom-name", ".git")); os.IsNotExist(err) { + t.Error("cloned repo not found at custom directory name") + } +} + +func TestHandleCloneRepo_DirectoryExists(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create the target directory ahead of time + if err := os.MkdirAll(filepath.Join(workspaceDir, "existing-repo"), 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + cloneReq := map[string]interface{}{ + "type": "clone_repo", + "id": "req-clone-3", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "url": "https://github.com/user/existing-repo.git", + } + if err := ms.sendCommand(cloneReq); err != nil { + t.Fatalf("failed to send clone command: %v", err) + } + + // Wait for error message + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err != nil { + t.Fatalf("failed to receive error message: %v", err) + } + + var errorReceived map[string]interface{} + if err := json.Unmarshal(raw, &errorReceived); err != nil { + t.Fatalf("failed to unmarshal error: %v", err) + } + + if errorReceived["code"] != "CLONE_FAILED" { + t.Errorf("expected code CLONE_FAILED, got %v", errorReceived["code"]) + } + if !strings.Contains(errorReceived["message"].(string), "already exists") { + t.Errorf("expected 'already exists' in message, got %v", errorReceived["message"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestHandleCloneRepo_InvalidURL(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + cloneReq := map[string]interface{}{ + "type": "clone_repo", + "id": "req-clone-4", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "url": "/nonexistent/invalid-repo", + } + if err := ms.sendCommand(cloneReq); err != nil { + t.Fatalf("failed to send clone command: %v", err) + } + + // Wait for clone_complete message + raw, err := ms.waitForMessageType("clone_complete", 5*time.Second) + if err != nil { + t.Fatalf("failed to receive clone_complete: %v", err) + } + + var cloneComplete map[string]interface{} + if err := json.Unmarshal(raw, &cloneComplete); err != nil { + t.Fatalf("failed to unmarshal clone_complete: %v", err) + } + + if cloneComplete["success"] != false { + t.Errorf("expected success=false, got %v", cloneComplete["success"]) + } + errMsg, ok := cloneComplete["error"].(string) + if !ok || errMsg == "" { + t.Error("expected non-empty error message for failed clone") + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestHandleCreateProject_Success(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + createReq := map[string]interface{}{ + "type": "create_project", + "id": "req-create-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "name": "new-project", + "git_init": false, + } + if err := ms.sendCommand(createReq); err != nil { + t.Fatalf("failed to send create command: %v", err) + } + + // Wait for project_list message (without git_init, project won't show up in scanner) + raw, err := ms.waitForMessageType("project_list", 2*time.Second) + if err != nil { + t.Fatalf("failed to receive project_list: %v", err) + } + + var response map[string]interface{} + if err := json.Unmarshal(raw, &response); err != nil { + t.Fatalf("failed to unmarshal project_list: %v", err) + } + + if response["type"] != "project_list" { + t.Errorf("expected type 'project_list', got %v", response["type"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify directory was created + projectDir := filepath.Join(workspaceDir, "new-project") + info, err := os.Stat(projectDir) + if err != nil { + t.Fatalf("project directory not created: %v", err) + } + if !info.IsDir() { + t.Error("expected project path to be a directory") + } +} + +func TestHandleCreateProject_WithGitInit(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + createReq := map[string]interface{}{ + "type": "create_project", + "id": "req-create-2", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "name": "git-project", + "git_init": true, + } + if err := ms.sendCommand(createReq); err != nil { + t.Fatalf("failed to send create command: %v", err) + } + + // Wait for project_state message (with git_init, scanner finds the project) + raw, err := ms.waitForMessageType("project_state", 2*time.Second) + if err != nil { + t.Fatalf("failed to receive project_state: %v", err) + } + + var response map[string]interface{} + if err := json.Unmarshal(raw, &response); err != nil { + t.Fatalf("failed to unmarshal project_state: %v", err) + } + + if response["type"] != "project_state" { + t.Errorf("expected type 'project_state', got %v", response["type"]) + } + project, ok := response["project"].(map[string]interface{}) + if !ok { + t.Fatal("expected project object in response") + } + if project["name"] != "git-project" { + t.Errorf("expected project name 'git-project', got %v", project["name"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify git repo was initialized + projectDir := filepath.Join(workspaceDir, "git-project") + gitDir := filepath.Join(projectDir, ".git") + if _, err := os.Stat(gitDir); os.IsNotExist(err) { + t.Error("expected .git directory to be created") + } +} + +func TestHandleCreateProject_AlreadyExists(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create directory ahead of time + if err := os.MkdirAll(filepath.Join(workspaceDir, "existing"), 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + createReq := map[string]interface{}{ + "type": "create_project", + "id": "req-create-3", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "name": "existing", + "git_init": false, + } + if err := ms.sendCommand(createReq); err != nil { + t.Fatalf("failed to send create command: %v", err) + } + + // Wait for error message + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err != nil { + t.Fatalf("failed to receive error message: %v", err) + } + + var errorReceived map[string]interface{} + if err := json.Unmarshal(raw, &errorReceived); err != nil { + t.Fatalf("failed to unmarshal error: %v", err) + } + + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "FILESYSTEM_ERROR" { + t.Errorf("expected code FILESYSTEM_ERROR, got %v", errorReceived["code"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestHandleCreateProject_EmptyName(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + createReq := map[string]interface{}{ + "type": "create_project", + "id": "req-create-4", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "name": "", + "git_init": false, + } + if err := ms.sendCommand(createReq); err != nil { + t.Fatalf("failed to send create command: %v", err) + } + + // Wait for error message + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err != nil { + t.Fatalf("failed to receive error message: %v", err) + } + + var errorReceived map[string]interface{} + if err := json.Unmarshal(raw, &errorReceived); err != nil { + t.Fatalf("failed to unmarshal error: %v", err) + } + + if errorReceived["code"] != "FILESYSTEM_ERROR" { + t.Errorf("expected code FILESYSTEM_ERROR, got %v", errorReceived["code"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestScanGitProgress(t *testing.T) { + input := "Cloning into 'repo'...\rReceiving objects: 50%\rReceiving objects: 100%\nDone.\n" + var tokens []string + data := []byte(input) + for len(data) > 0 { + advance, token, err := scanGitProgress(data, false) + if err != nil { + t.Fatal(err) + } + if advance == 0 { + // Process remaining at EOF + _, token, _ = scanGitProgress(data, true) + if token != nil { + tokens = append(tokens, string(token)) + } + break + } + if token != nil { + tokens = append(tokens, string(token)) + } + data = data[advance:] + } + + expected := []string{"Cloning into 'repo'...", "Receiving objects: 50%", "Receiving objects: 100%", "Done."} + if len(tokens) != len(expected) { + t.Fatalf("expected %d tokens, got %d: %v", len(expected), len(tokens), tokens) + } + for i, tok := range tokens { + if tok != expected[i] { + t.Errorf("token[%d] = %q, want %q", i, tok, expected[i]) + } + } +} + +// Unit tests for clone/create functions with mock projectFinder + +type mockScanner struct { + workspacePath string + projects []ws.ProjectSummary +} + +func (m *mockScanner) FindProject(name string) (ws.ProjectSummary, bool) { + for _, p := range m.projects { + if p.Name == name { + return p, true + } + } + return ws.ProjectSummary{}, false +} + +func TestCloneProgressParsing(t *testing.T) { + tests := []struct { + input string + percent int + }{ + {"Receiving objects: 50% (1/2)", 50}, + {"Resolving deltas: 100% (10/10)", 100}, + {"Cloning into 'repo'...", 0}, + {"Receiving objects: 3% (1/33)", 3}, + } + + for _, tt := range tests { + matches := percentPattern.FindStringSubmatch(tt.input) + got := 0 + if len(matches) > 1 { + got, _ = strconv.Atoi(matches[1]) + } + if got != tt.percent { + t.Errorf("percent for %q: got %d, want %d", tt.input, got, tt.percent) + } + } +} + +func TestSendCloneComplete_NilClient(t *testing.T) { + // Should not panic + sendCloneComplete(nil, "https://example.com/repo.git", "repo", true, "") +} + +func TestSendCloneProgress_NilClient(t *testing.T) { + // Should not panic + sendCloneProgress(nil, "https://example.com/repo.git", "progress", 50) +} diff --git a/internal/cmd/diffs.go b/internal/cmd/diffs.go new file mode 100644 index 00000000..e9a62866 --- /dev/null +++ b/internal/cmd/diffs.go @@ -0,0 +1,232 @@ +package cmd + +import ( + "encoding/json" + "fmt" + "log" + "os" + "os/exec" + "path/filepath" + "strings" + + "github.com/minicodemonkey/chief/internal/ws" +) + +// handleGetDiffs handles a get_diffs request from the browser. +// Unlike get_diff, this does not require prd_id and returns parsed per-file diffs. +func handleGetDiffs(sender messageSender, finder projectFinder, msg ws.Message) { + var req ws.GetDiffsMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing get_diffs message: %v", err) + return + } + + project, found := finder.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + diffText, _, err := getStoryDiff(project.Path, req.StoryID) + if err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Failed to get diff for story %q: %v", req.StoryID, err), msg.ID) + return + } + + files := parseDiffFiles(diffText) + + resp := ws.DiffsResponseMessage{ + Type: ws.TypeDiffsResponse, + Payload: ws.DiffsResponsePayload{ + Project: req.Project, + StoryID: req.StoryID, + Files: files, + }, + } + if err := sender.Send(resp); err != nil { + log.Printf("Error sending diffs_response: %v", err) + } +} + +// parseDiffFiles splits a unified diff into per-file details. +func parseDiffFiles(diffText string) []ws.DiffFileDetail { + if diffText == "" { + return []ws.DiffFileDetail{} + } + + // Split on "diff --git" boundaries + chunks := strings.Split(diffText, "diff --git ") + var files []ws.DiffFileDetail + + for _, chunk := range chunks { + chunk = strings.TrimSpace(chunk) + if chunk == "" { + continue + } + + // Extract filename from first line: "a/path b/path" + firstLine := chunk + if idx := strings.IndexByte(chunk, '\n'); idx != -1 { + firstLine = chunk[:idx] + } + + filename := "" + if parts := strings.SplitN(firstLine, " b/", 2); len(parts) == 2 { + filename = parts[1] + } + + // Count additions and deletions + additions, deletions := 0, 0 + for _, line := range strings.Split(chunk, "\n") { + if strings.HasPrefix(line, "+") && !strings.HasPrefix(line, "+++") { + additions++ + } else if strings.HasPrefix(line, "-") && !strings.HasPrefix(line, "---") { + deletions++ + } + } + + files = append(files, ws.DiffFileDetail{ + Filename: filename, + Additions: additions, + Deletions: deletions, + Patch: "diff --git " + chunk, + }) + } + + if files == nil { + files = []ws.DiffFileDetail{} + } + return files +} + +// handleGetDiff handles a get_diff request. +func handleGetDiff(sender messageSender, finder projectFinder, msg ws.Message) { + var req ws.GetDiffMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing get_diff message: %v", err) + return + } + + project, found := finder.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + prdDir := filepath.Join(project.Path, ".chief", "prds", req.PRDID) + if _, err := os.Stat(prdDir); os.IsNotExist(err) { + sendError(sender, ws.ErrCodePRDNotFound, + fmt.Sprintf("PRD %q not found in project %q", req.PRDID, req.Project), msg.ID) + return + } + + diffText, files, err := getStoryDiff(project.Path, req.StoryID) + if err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Failed to get diff for story %q: %v", req.StoryID, err), msg.ID) + return + } + + sendDiffMessage(sender, req.Project, req.PRDID, req.StoryID, files, diffText) +} + +// getStoryDiff returns the diff and list of changed files for a story's commit(s). +// It finds commits matching the story ID pattern "feat: -" in the commit message. +func getStoryDiff(repoDir, storyID string) (string, []string, error) { + // Find commit hash(es) for this story by searching commit messages + commitHash, err := findStoryCommit(repoDir, storyID) + if err != nil { + return "", nil, err + } + if commitHash == "" { + return "", nil, fmt.Errorf("no commit found for story %s", storyID) + } + + // Get the unified diff for the commit + diffText, err := getCommitDiff(repoDir, commitHash) + if err != nil { + return "", nil, fmt.Errorf("getting diff: %w", err) + } + + // Get the list of changed files + files, err := getCommitFiles(repoDir, commitHash) + if err != nil { + return "", nil, fmt.Errorf("getting changed files: %w", err) + } + + return diffText, files, nil +} + +// findStoryCommit finds the most recent commit hash matching a story ID. +// It searches for commits with messages matching "feat: -" or +// containing the story ID. +func findStoryCommit(repoDir, storyID string) (string, error) { + // Search for commits with messages containing the story ID + cmd := exec.Command("git", "log", "--format=%H", "--grep", storyID, "-1") + cmd.Dir = repoDir + output, err := cmd.Output() + if err != nil { + return "", fmt.Errorf("searching git log: %w", err) + } + + hash := strings.TrimSpace(string(output)) + return hash, nil +} + +// getCommitDiff returns the unified diff for a specific commit. +func getCommitDiff(repoDir, commitHash string) (string, error) { + cmd := exec.Command("git", "show", "--format=", "--patch", commitHash) + cmd.Dir = repoDir + output, err := cmd.Output() + if err != nil { + return "", err + } + return string(output), nil +} + +// getCommitFiles returns the list of files changed in a specific commit. +func getCommitFiles(repoDir, commitHash string) ([]string, error) { + cmd := exec.Command("git", "show", "--format=", "--name-only", commitHash) + cmd.Dir = repoDir + output, err := cmd.Output() + if err != nil { + return nil, err + } + + raw := strings.TrimSpace(string(output)) + if raw == "" { + return []string{}, nil + } + + files := strings.Split(raw, "\n") + return files, nil +} + +// sendDiffMessage sends a diff message. +func sendDiffMessage(sender messageSender, project, prdID, storyID string, files []string, diffText string) { + if sender == nil { + return + } + + if files == nil { + files = []string{} + } + + envelope := ws.NewMessage(ws.TypeDiff) + msg := ws.DiffMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: project, + PRDID: prdID, + StoryID: storyID, + Files: files, + DiffText: diffText, + } + if err := sender.Send(msg); err != nil { + log.Printf("Error sending diff: %v", err) + } +} diff --git a/internal/cmd/diffs_test.go b/internal/cmd/diffs_test.go new file mode 100644 index 00000000..4f5a1a45 --- /dev/null +++ b/internal/cmd/diffs_test.go @@ -0,0 +1,636 @@ +package cmd + +import ( + "encoding/json" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/engine" +) + +// gitCmd runs a git command in the given directory with test-safe env. +func gitCmd(t *testing.T, dir string, args ...string) string { + t.Helper() + cmd := exec.Command("git", args...) + cmd.Dir = dir + cmd.Env = append(os.Environ(), "GIT_CONFIG_GLOBAL=/dev/null", "GIT_CONFIG_SYSTEM=/dev/null") + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("git %s failed: %v\n%s", strings.Join(args, " "), err, out) + } + return strings.TrimSpace(string(out)) +} + +// createGitRepoWithStoryCommit creates a git repo with an initial commit +// and a story commit matching the "feat: - " pattern. +func createGitRepoWithStoryCommit(t *testing.T, dir, storyID, title string) { + t.Helper() + createGitRepo(t, dir) + + // Create a file and commit it with the story commit message + filePath := filepath.Join(dir, "feature.go") + if err := os.WriteFile(filePath, []byte("package main\n\nfunc feature() {}\n"), 0o644); err != nil { + t.Fatal(err) + } + gitCmd(t, dir, "add", "feature.go") + gitCmd(t, dir, "commit", "-m", "feat: "+storyID+" - "+title) +} + +func TestGetStoryDiff_Success(t *testing.T) { + dir := t.TempDir() + createGitRepoWithStoryCommit(t, dir, "US-001", "Add feature") + + diffText, files, err := getStoryDiff(dir, "US-001") + if err != nil { + t.Fatalf("getStoryDiff failed: %v", err) + } + + if diffText == "" { + t.Error("expected non-empty diff text") + } + if !strings.Contains(diffText, "feature.go") { + t.Errorf("expected diff to contain 'feature.go', got: %s", diffText) + } + + if len(files) != 1 { + t.Errorf("expected 1 changed file, got %d: %v", len(files), files) + } + if len(files) > 0 && files[0] != "feature.go" { + t.Errorf("expected file 'feature.go', got %q", files[0]) + } +} + +func TestGetStoryDiff_NoCommitFound(t *testing.T) { + dir := t.TempDir() + createGitRepo(t, dir) + + _, _, err := getStoryDiff(dir, "US-999") + if err == nil { + t.Fatal("expected error for missing story commit") + } + if !strings.Contains(err.Error(), "no commit found") { + t.Errorf("expected 'no commit found' error, got: %v", err) + } +} + +func TestGetStoryDiff_MultipleFiles(t *testing.T) { + dir := t.TempDir() + createGitRepo(t, dir) + + // Create multiple files and commit + for _, name := range []string{"a.go", "b.go", "c.go"} { + if err := os.WriteFile(filepath.Join(dir, name), []byte("package main\n"), 0o644); err != nil { + t.Fatal(err) + } + } + gitCmd(t, dir, "add", ".") + gitCmd(t, dir, "commit", "-m", "feat: US-002 - Add multiple files") + + diffText, files, err := getStoryDiff(dir, "US-002") + if err != nil { + t.Fatalf("getStoryDiff failed: %v", err) + } + + if len(files) != 3 { + t.Errorf("expected 3 changed files, got %d: %v", len(files), files) + } + + if !strings.Contains(diffText, "a.go") || !strings.Contains(diffText, "b.go") || !strings.Contains(diffText, "c.go") { + t.Errorf("expected diff to contain all files, got: %s", diffText) + } +} + +func TestFindStoryCommit_FindsMostRecent(t *testing.T) { + dir := t.TempDir() + createGitRepo(t, dir) + + // Create first commit for the story + if err := os.WriteFile(filepath.Join(dir, "v1.go"), []byte("package v1\n"), 0o644); err != nil { + t.Fatal(err) + } + gitCmd(t, dir, "add", ".") + gitCmd(t, dir, "commit", "-m", "feat: US-003 - Initial attempt") + + // Create second commit for the same story (more recent) + if err := os.WriteFile(filepath.Join(dir, "v2.go"), []byte("package v2\n"), 0o644); err != nil { + t.Fatal(err) + } + gitCmd(t, dir, "add", ".") + gitCmd(t, dir, "commit", "-m", "feat: US-003 - Fixed version") + + hash, err := findStoryCommit(dir, "US-003") + if err != nil { + t.Fatalf("findStoryCommit failed: %v", err) + } + if hash == "" { + t.Fatal("expected non-empty commit hash") + } + + // The most recent commit should be the "Fixed version" one + // Verify by checking the commit message + cmd := exec.Command("git", "log", "--format=%s", "-1", hash) + cmd.Dir = dir + out, err := cmd.Output() + if err != nil { + t.Fatalf("git log failed: %v", err) + } + msg := strings.TrimSpace(string(out)) + if msg != "feat: US-003 - Fixed version" { + t.Errorf("expected most recent commit, got: %q", msg) + } +} + +func TestSendDiffMessage(t *testing.T) { + // sendDiffMessage with nil client should not panic + sendDiffMessage(nil, "project", "prd", "US-001", []string{"a.go"}, "diff text") +} + +func TestSendDiffMessage_NilFiles(t *testing.T) { + // sendDiffMessage with nil files should not panic + sendDiffMessage(nil, "project", "prd", "US-001", nil, "diff text") +} + +func TestRunManager_SendStoryDiff(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) // nil client — just verify no panic + + // Create a temp project with a git repo and story commit + projectDir := t.TempDir() + createGitRepoWithStoryCommit(t, projectDir, "US-001", "Test Story") + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(`{}`), 0o644); err != nil { + t.Fatal(err) + } + + rm.mu.Lock() + rm.runs["myproject/feature"] = &runInfo{ + project: "myproject", + prdID: "feature", + prdPath: prdPath, + startTime: time.Now(), + storyID: "US-001", + } + rm.mu.Unlock() + + // Call sendStoryDiff with nil client — should not panic, just log + info := rm.runs["myproject/feature"] + rm.sendStoryDiff(info, engine.ManagerEvent{}.Event) +} + +func TestRunServe_GetDiff(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo with a story commit + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepoWithStoryCommit(t, projectDir, "US-001", "Add feature") + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Add feature", "passes": true}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + getDiffReq := map[string]interface{}{ + "type": "get_diff", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + "story_id": "US-001", + } + if err := ms.sendCommand(getDiffReq); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("diff", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "diff" { + t.Errorf("expected type 'diff', got %v", response["type"]) + } + if response["project"] != "myproject" { + t.Errorf("expected project 'myproject', got %v", response["project"]) + } + if response["prd_id"] != "feature" { + t.Errorf("expected prd_id 'feature', got %v", response["prd_id"]) + } + if response["story_id"] != "US-001" { + t.Errorf("expected story_id 'US-001', got %v", response["story_id"]) + } + + // Verify files array + files, ok := response["files"].([]interface{}) + if !ok { + t.Fatal("expected files to be an array") + } + if len(files) != 1 { + t.Errorf("expected 1 file, got %d", len(files)) + } + if len(files) > 0 && files[0] != "feature.go" { + t.Errorf("expected file 'feature.go', got %v", files[0]) + } + + // Verify diff_text is non-empty and contains the file + diffText, ok := response["diff_text"].(string) + if !ok || diffText == "" { + t.Error("expected non-empty diff_text") + } + if !strings.Contains(diffText, "feature.go") { + t.Errorf("expected diff_text to contain 'feature.go'") + } +} + +func TestRunServe_GetDiffProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + getDiffReq := map[string]interface{}{ + "type": "get_diff", + "id": "req-2", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "prd_id": "feature", + "story_id": "US-001", + } + if err := ms.sendCommand(getDiffReq); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "error" { + t.Errorf("expected type 'error', got %v", response["type"]) + } + if response["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", response["code"]) + } +} + +func TestRunServe_GetDiffPRDNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + getDiffReq := map[string]interface{}{ + "type": "get_diff", + "id": "req-3", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "nonexistent", + "story_id": "US-001", + } + if err := ms.sendCommand(getDiffReq); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "error" { + t.Errorf("expected type 'error', got %v", response["type"]) + } + if response["code"] != "PRD_NOT_FOUND" { + t.Errorf("expected code 'PRD_NOT_FOUND', got %v", response["code"]) + } +} + +func TestRunServe_GetDiffNoCommit(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test", "passes": false}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + getDiffReq := map[string]interface{}{ + "type": "get_diff", + "id": "req-4", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + "story_id": "US-001", + } + if err := ms.sendCommand(getDiffReq); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "error" { + t.Errorf("expected type 'error', got %v", response["type"]) + } + if response["code"] != "FILESYSTEM_ERROR" { + t.Errorf("expected code 'FILESYSTEM_ERROR', got %v", response["code"]) + } +} + +func TestRunServe_GetDiffs(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo with a story commit + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepoWithStoryCommit(t, projectDir, "US-001", "Add feature") + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // get_diffs does not require prd_id (unlike get_diff) + req := map[string]interface{}{ + "type": "get_diffs", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "story_id": "US-001", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("diffs_response", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("diffs_response was not received") + } + if response["type"] != "diffs_response" { + t.Errorf("expected type 'diffs_response', got %v", response["type"]) + } + + payload, ok := response["payload"].(map[string]interface{}) + if !ok { + t.Fatal("expected payload to be an object") + } + if payload["project"] != "myproject" { + t.Errorf("expected project 'myproject', got %v", payload["project"]) + } + if payload["story_id"] != "US-001" { + t.Errorf("expected story_id 'US-001', got %v", payload["story_id"]) + } + + files, ok := payload["files"].([]interface{}) + if !ok { + t.Fatal("expected files to be an array") + } + if len(files) != 1 { + t.Fatalf("expected 1 file, got %d", len(files)) + } + + file := files[0].(map[string]interface{}) + if file["filename"] != "feature.go" { + t.Errorf("expected filename 'feature.go', got %v", file["filename"]) + } + if int(file["additions"].(float64)) == 0 { + t.Error("expected additions > 0") + } + if _, ok := file["patch"].(string); !ok || file["patch"] == "" { + t.Error("expected non-empty patch string") + } +} + +func TestRunServe_GetDiffs_ProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]interface{}{ + "type": "get_diffs", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "story_id": "US-001", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("error message was not received") + } + if response["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", response["code"]) + } +} + +func TestParseDiffFiles(t *testing.T) { + diffText := `diff --git a/main.go b/main.go +index abc..def 100644 +--- a/main.go ++++ b/main.go +@@ -1,3 +1,5 @@ + package main ++import "fmt" ++func hello() { fmt.Println("hi") } + func main() {} +diff --git a/util.go b/util.go +new file mode 100644 +--- /dev/null ++++ b/util.go +@@ -0,0 +1,3 @@ ++package main ++func helper() {} ++func other() {} +` + + files := parseDiffFiles(diffText) + if len(files) != 2 { + t.Fatalf("expected 2 files, got %d", len(files)) + } + + if files[0].Filename != "main.go" { + t.Errorf("files[0].filename = %q, want %q", files[0].Filename, "main.go") + } + if files[0].Additions != 2 { + t.Errorf("files[0].additions = %d, want 2", files[0].Additions) + } + if files[0].Deletions != 0 { + t.Errorf("files[0].deletions = %d, want 0", files[0].Deletions) + } + + if files[1].Filename != "util.go" { + t.Errorf("files[1].filename = %q, want %q", files[1].Filename, "util.go") + } + if files[1].Additions != 3 { + t.Errorf("files[1].additions = %d, want 3", files[1].Additions) + } +} + +func TestParseDiffFiles_Empty(t *testing.T) { + files := parseDiffFiles("") + if len(files) != 0 { + t.Errorf("expected 0 files for empty diff, got %d", len(files)) + } +} diff --git a/internal/cmd/e2e_test.go b/internal/cmd/e2e_test.go new file mode 100644 index 00000000..6a85e32f --- /dev/null +++ b/internal/cmd/e2e_test.go @@ -0,0 +1,400 @@ +package cmd + +import ( + "context" + "encoding/json" + "fmt" + "os" + "path/filepath" + "sync" + "testing" + "time" +) + +// End-to-end integration tests verifying the complete CLI ↔ server message flow +// via the uplink HTTP+Pusher transport. These tests complement the unit-level +// uplink tests (internal/uplink/*_test.go) and the existing serve_test.go tests. + +func TestE2E_DeviceLifecycle_ConnectMessagesHeartbeatDisconnect(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + // Wait for full connection (Pusher subscribe). + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for state_snapshot (CLI sends on connect). + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send a command via Pusher → CLI should respond. + listReq := map[string]string{ + "type": "list_projects", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + } + if err := ms.sendCommand(listReq); err != nil { + t.Logf("sendCommand: %v", err) + cancel() + return + } + + // Wait for project_list response via HTTP messages endpoint. + if _, err := ms.waitForMessageType("project_list", 5*time.Second); err != nil { + t.Logf("waitForMessageType(project_list): %v", err) + } + + // Cancel to trigger graceful shutdown (disconnect). + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspaceDir, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify the full lifecycle happened: + // 1. Connect was called + if got := ms.connectCalls.Load(); got < 1 { + t.Errorf("connect calls = %d, want >= 1", got) + } + + // 2. Messages were sent (state_snapshot + project_list at minimum) + if got := ms.messagesCalls.Load(); got < 1 { + t.Errorf("messages calls = %d, want >= 1", got) + } + + // 3. State snapshot was received + if _, err := ms.waitForMessageType("state_snapshot", time.Second); err != nil { + t.Error("state_snapshot not received") + } + + // 4. Project list response was received + raw, err := ms.waitForMessageType("project_list", time.Second) + if err != nil { + t.Error("project_list not received") + } else { + var resp map[string]interface{} + json.Unmarshal(raw, &resp) + projects := resp["projects"].([]interface{}) + if len(projects) != 1 { + t.Errorf("expected 1 project, got %d", len(projects)) + } + } + + // 5. Disconnect was called during shutdown + if got := ms.disconnectCalls.Load(); got != 1 { + t.Errorf("disconnect calls = %d, want 1", got) + } +} + +func TestE2E_BidirectionalMessageFlow(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "alpha")) + createGitRepo(t, filepath.Join(workspaceDir, "beta")) + + var responses []map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send multiple commands rapidly — verify CLI processes and responds to each. + + // Command 1: list_projects + ms.sendCommand(map[string]string{ + "type": "list_projects", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + }) + + // Command 2: get_project + ms.sendCommand(map[string]string{ + "type": "get_project", + "id": "req-2", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "alpha", + }) + + // Command 3: ping + ms.sendCommand(map[string]string{ + "type": "ping", + "id": "req-3", + "timestamp": time.Now().UTC().Format(time.RFC3339), + }) + + // Wait for all three response types. + types := []string{"project_list", "project_state", "pong"} + for _, typ := range types { + raw, err := ms.waitForMessageType(typ, 5*time.Second) + if err == nil { + var resp map[string]interface{} + json.Unmarshal(raw, &resp) + mu.Lock() + responses = append(responses, resp) + mu.Unlock() + } else { + t.Logf("waitForMessageType(%s): %v", typ, err) + } + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if len(responses) != 3 { + t.Fatalf("expected 3 responses, got %d", len(responses)) + } + + // Verify we received all expected types. + typeSet := make(map[string]bool) + for _, r := range responses { + typeSet[r["type"].(string)] = true + } + for _, expected := range []string{"project_list", "project_state", "pong"} { + if !typeSet[expected] { + t.Errorf("missing response type %q in %v", expected, typeSet) + } + } +} + +func TestE2E_HeartbeatSentDuringSession(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for state_snapshot. + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Wait for at least one heartbeat call (up to 40s since default interval is 30s). + deadline := time.After(40 * time.Second) + for { + if ms.heartbeatCalls.Load() > 0 { + break + } + select { + case <-deadline: + t.Logf("timeout waiting for heartbeat") + cancel() + return + case <-time.After(100 * time.Millisecond): + } + } + + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + if got := ms.heartbeatCalls.Load(); got < 1 { + t.Errorf("heartbeat calls = %d, want >= 1", got) + } +} + +func TestE2E_MultipleCommandsBatchedIntoHTTPPosts(t *testing.T) { + // Verifies that multiple CLI responses are batched into HTTP POST calls + // (not one per message) via the batcher. + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var finalPongCount int + var finalTotalMsgs int + var finalHTTPCalls int + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send several rapid commands that produce responses. + for i := 0; i < 5; i++ { + ms.sendCommand(map[string]string{ + "type": "ping", + "id": fmt.Sprintf("ping-%d", i), + "timestamp": time.Now().UTC().Format(time.RFC3339), + }) + } + + // Wait for all pongs. + deadline := time.After(5 * time.Second) + pongCount := 0 + for pongCount < 5 { + msgs := ms.getMessages() + pongCount = 0 + for _, raw := range msgs { + var msg map[string]interface{} + json.Unmarshal(raw, &msg) + if msg["type"] == "pong" { + pongCount++ + } + } + if pongCount >= 5 { + break + } + select { + case <-deadline: + t.Logf("timeout: got %d pongs", pongCount) + mu.Lock() + finalPongCount = pongCount + finalTotalMsgs = len(ms.getMessages()) + finalHTTPCalls = int(ms.messagesCalls.Load()) + mu.Unlock() + return + case <-time.After(50 * time.Millisecond): + } + } + + mu.Lock() + finalPongCount = pongCount + finalTotalMsgs = len(ms.getMessages()) + finalHTTPCalls = int(ms.messagesCalls.Load()) + mu.Unlock() + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if finalPongCount != 5 { + t.Errorf("expected 5 pong messages, got %d", finalPongCount) + } + + t.Logf("total messages: %d, HTTP POST calls: %d", finalTotalMsgs, finalHTTPCalls) + if finalHTTPCalls > finalTotalMsgs { + t.Errorf("HTTP POST calls (%d) > total messages (%d), batching may not be working", finalHTTPCalls, finalTotalMsgs) + } +} + +func TestE2E_GracefulShutdownFlushesMessagesAndDisconnects(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspace, "myproject")) + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for state_snapshot. + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send a ping command. + ms.sendCommand(map[string]string{ + "type": "ping", + "id": "ping-before-shutdown", + "timestamp": time.Now().UTC().Format(time.RFC3339), + }) + + // Wait for the pong to actually be delivered to the server. + // This confirms the full send path (enqueue → batcher flush → HTTP POST) + // completed before we trigger shutdown. + if _, err := ms.waitForMessageType("pong", 5*time.Second); err != nil { + t.Logf("waitForMessageType(pong): %v", err) + } + + // Trigger shutdown after messages have been flushed. + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify pong was received by the server (already confirmed in goroutine). + if _, err := ms.waitForMessageType("pong", time.Second); err != nil { + t.Error("pong not received by server") + } + + // Verify disconnect was called during shutdown. + if got := ms.disconnectCalls.Load(); got != 1 { + t.Errorf("disconnect calls = %d, want 1", got) + } +} diff --git a/internal/cmd/edit.go b/internal/cmd/edit.go index 878fa743..87c2c96f 100644 --- a/internal/cmd/edit.go +++ b/internal/cmd/edit.go @@ -6,15 +6,14 @@ import ( "path/filepath" "github.com/minicodemonkey/chief/embed" - "github.com/minicodemonkey/chief/internal/loop" - "github.com/minicodemonkey/chief/internal/prd" ) // EditOptions contains configuration for the edit command. type EditOptions struct { - Name string // PRD name (default: "main") - BaseDir string // Base directory for .chief/prds/ (default: current directory) - Provider loop.Provider // Agent CLI provider (Claude or Codex) + Name string // PRD name (default: "main") + BaseDir string // Base directory for .chief/prds/ (default: current directory) + Merge bool // Auto-merge without prompting on conversion conflicts + Force bool // Auto-overwrite without prompting on conversion conflicts } // RunEdit edits an existing PRD by launching an interactive Claude session. @@ -47,24 +46,26 @@ func RunEdit(opts EditOptions) error { // Get the edit prompt with the PRD directory path prompt := embed.GetEditPrompt(prdDir) - if opts.Provider == nil { - return fmt.Errorf("edit command requires Provider to be set") - } - // Launch interactive agent session + // Launch interactive Claude session fmt.Printf("Editing PRD at %s...\n", prdDir) - fmt.Printf("Launching %s to help you edit your PRD...\n", opts.Provider.Name()) + fmt.Println("Launching Claude to help you edit your PRD...") fmt.Println() - if err := runInteractiveAgent(opts.Provider, opts.BaseDir, prompt); err != nil { - return fmt.Errorf("%s session failed: %w", opts.Provider.Name(), err) + if err := runInteractiveClaude(opts.BaseDir, prompt); err != nil { + return fmt.Errorf("Claude session failed: %w", err) } fmt.Println("\nPRD editing complete!") - // Validate the edited prd.md can be parsed - if _, err := prd.ParseMarkdownPRD(prdMdPath); err != nil { - fmt.Printf("Warning: prd.md could not be parsed: %v\n", err) + // Run conversion from prd.md to prd.json with progress protection + convertOpts := ConvertOptions{ + PRDDir: prdDir, + Merge: opts.Merge, + Force: opts.Force, + } + if err := RunConvertWithOptions(convertOpts); err != nil { + return fmt.Errorf("conversion failed: %w", err) } fmt.Printf("\nYour PRD is updated! Run 'chief' or 'chief %s' to continue working on it.\n", opts.Name) diff --git a/internal/cmd/edit_test.go b/internal/cmd/edit_test.go index 35ba0589..21d12be9 100644 --- a/internal/cmd/edit_test.go +++ b/internal/cmd/edit_test.go @@ -73,39 +73,50 @@ func TestRunEditDefaultsToMain(t *testing.T) { } } -func TestEditOptionsDefaults(t *testing.T) { - opts := EditOptions{} +func TestRunEditWithMergeFlag(t *testing.T) { + opts := EditOptions{ + Name: "test", + Merge: true, + Force: false, + } - if opts.Name != "" { - t.Error("Name should default to empty (filled later)") + if !opts.Merge { + t.Error("Merge flag should be true") } - if opts.BaseDir != "" { - t.Error("BaseDir should default to empty (filled later)") + if opts.Force { + t.Error("Force flag should be false") } } -func TestRunEditRequiresProvider(t *testing.T) { - tmpDir := t.TempDir() - prdDir := filepath.Join(tmpDir, ".chief", "prds", "main") - if err := os.MkdirAll(prdDir, 0755); err != nil { - t.Fatalf("Failed to create directory: %v", err) - } - prdMdPath := filepath.Join(prdDir, "prd.md") - if err := os.WriteFile(prdMdPath, []byte("# Main PRD"), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) +func TestRunEditWithForceFlag(t *testing.T) { + opts := EditOptions{ + Name: "test", + Merge: false, + Force: true, } - opts := EditOptions{ - Name: "main", - BaseDir: tmpDir, + if opts.Merge { + t.Error("Merge flag should be false") + } + if !opts.Force { + t.Error("Force flag should be true") } +} - err := RunEdit(opts) - if err == nil { - t.Fatal("expected provider validation error") +func TestEditOptionsDefaults(t *testing.T) { + opts := EditOptions{} + + if opts.Name != "" { + t.Error("Name should default to empty (filled later)") } - if !contains(err.Error(), "Provider") { - t.Fatalf("expected error to mention Provider, got: %v", err) + if opts.Merge { + t.Error("Merge should default to false") + } + if opts.Force { + t.Error("Force should default to false") + } + if opts.BaseDir != "" { + t.Error("BaseDir should default to empty (filled later)") } } diff --git a/internal/cmd/login.go b/internal/cmd/login.go new file mode 100644 index 00000000..8f4846e6 --- /dev/null +++ b/internal/cmd/login.go @@ -0,0 +1,260 @@ +package cmd + +import ( + "bufio" + "bytes" + "encoding/json" + "errors" + "fmt" + "net/http" + "os" + "os/exec" + "runtime" + "strings" + "time" + + "github.com/minicodemonkey/chief/internal/auth" +) + +const ( + defaultBaseURL = "https://uplink.chiefloop.com" + pollInterval = 5 * time.Second + loginTimeout = 5 * time.Minute +) + +// LoginOptions contains configuration for the login command. +type LoginOptions struct { + DeviceName string // Override device name (default: hostname) + BaseURL string // Override base URL (for testing) + SetupToken string // One-time setup token for automated auth +} + +// deviceCodeResponse is the response from the device code endpoint. +type deviceCodeResponse struct { + DeviceCode string `json:"device_code"` + UserCode string `json:"user_code"` +} + +// tokenResponse is the response from the token polling endpoint. +type tokenResponse struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ExpiresIn int `json:"expires_in"` + User string `json:"user"` + WSURL string `json:"ws_url"` + Error string `json:"error"` +} + +// RunLogin performs the device OAuth login flow. +func RunLogin(opts LoginOptions) error { + baseURL := opts.BaseURL + if baseURL == "" { + baseURL = os.Getenv("CHIEF_SERVER_URL") + } + if baseURL == "" { + baseURL = defaultBaseURL + } + + deviceName := opts.DeviceName + if deviceName == "" { + hostname, err := os.Hostname() + if err != nil { + deviceName = "unknown" + } else { + deviceName = hostname + } + } + + // Setup token mode: exchange token for credentials directly + if opts.SetupToken != "" { + return exchangeSetupToken(baseURL, opts.SetupToken, deviceName) + } + + // Check if already logged in + existing, err := auth.LoadCredentials() + if err == nil && existing != nil { + fmt.Printf("Already logged in as %s (%s).\n", existing.User, existing.DeviceName) + fmt.Print("Do you want to log in again? This will replace your existing credentials. [y/N] ") + reader := bufio.NewReader(os.Stdin) + answer, _ := reader.ReadString('\n') + answer = strings.TrimSpace(strings.ToLower(answer)) + if answer != "y" && answer != "yes" { + fmt.Println("Login cancelled.") + return nil + } + } + + // Request device code + codeReqBody, _ := json.Marshal(map[string]string{ + "device_name": deviceName, + }) + + resp, err := http.Post(baseURL+"/api/oauth/device/code", "application/json", bytes.NewReader(codeReqBody)) + if err != nil { + return fmt.Errorf("requesting device code: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("requesting device code: server returned %s", resp.Status) + } + + var codeResp deviceCodeResponse + if err := json.NewDecoder(resp.Body).Decode(&codeResp); err != nil { + return fmt.Errorf("parsing device code response: %w", err) + } + + // Display the user code and URL + deviceURL := baseURL + "/oauth/device" + fmt.Println() + fmt.Println("To authenticate, open this URL in your browser:") + fmt.Printf("\n %s\n\n", deviceURL) + fmt.Printf("And enter this code: %s\n\n", codeResp.UserCode) + + // Try to open browser automatically + openBrowserFunc(deviceURL) + + fmt.Println("Waiting for authorization...") + + // Poll for token + creds, err := pollForToken(baseURL, codeResp.DeviceCode, deviceName) + if err != nil { + return err + } + + // Save credentials + if err := auth.SaveCredentials(creds); err != nil { + return fmt.Errorf("saving credentials: %w", err) + } + + fmt.Printf("\nLogged in as %s (%s)\n", creds.User, creds.DeviceName) + return nil +} + +// exchangeSetupToken exchanges a one-time setup token for credentials. +// This is used during automated VPS provisioning via cloud-init. +func exchangeSetupToken(baseURL, token, deviceName string) error { + reqBody, _ := json.Marshal(map[string]string{ + "setup_token": token, + "device_name": deviceName, + }) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Post(baseURL+"/api/oauth/device/exchange", "application/json", bytes.NewReader(reqBody)) + if err != nil { + return fmt.Errorf("exchanging setup token: %w", err) + } + defer resp.Body.Close() + + var tokenResp tokenResponse + if err := json.NewDecoder(resp.Body).Decode(&tokenResp); err != nil { + return fmt.Errorf("parsing setup token response: %w", err) + } + + if resp.StatusCode != http.StatusOK || tokenResp.Error != "" { + errMsg := tokenResp.Error + if errMsg == "" { + errMsg = resp.Status + } + fmt.Fprintf(os.Stderr, "Setup token exchange failed: %s\n", errMsg) + fmt.Fprintf(os.Stderr, "Please authenticate manually by running: chief login\n") + return fmt.Errorf("setup token exchange failed: %s", errMsg) + } + + creds := &auth.Credentials{ + AccessToken: tokenResp.AccessToken, + RefreshToken: tokenResp.RefreshToken, + ExpiresAt: time.Now().Add(time.Duration(tokenResp.ExpiresIn) * time.Second), + DeviceName: deviceName, + User: tokenResp.User, + WSURL: tokenResp.WSURL, + } + + if err := auth.SaveCredentials(creds); err != nil { + return fmt.Errorf("saving credentials: %w", err) + } + + fmt.Printf("Logged in as %s (%s)\n", creds.User, creds.DeviceName) + return nil +} + +// pollForToken polls the token endpoint until authorization is granted or timeout. +func pollForToken(baseURL, deviceCode, deviceName string) (*auth.Credentials, error) { + deadline := time.Now().Add(loginTimeout) + client := &http.Client{Timeout: 10 * time.Second} + + for { + if time.Now().After(deadline) { + return nil, errors.New("login timed out — you did not authorize the device within 5 minutes") + } + + time.Sleep(pollInterval) + + reqBody, _ := json.Marshal(map[string]string{ + "device_code": deviceCode, + }) + + resp, err := client.Post(baseURL+"/api/oauth/device/token", "application/json", bytes.NewReader(reqBody)) + if err != nil { + fmt.Fprintf(os.Stderr, "Network error while polling (will retry): %v\n", err) + continue + } + + var tokenResp tokenResponse + if err := json.NewDecoder(resp.Body).Decode(&tokenResp); err != nil { + resp.Body.Close() + fmt.Fprintf(os.Stderr, "Error parsing token response (will retry): %v\n", err) + continue + } + resp.Body.Close() + + // Check for pending authorization + if tokenResp.Error == "authorization_pending" { + continue + } + + // Check for other errors + if tokenResp.Error != "" { + return nil, fmt.Errorf("authorization failed: %s", tokenResp.Error) + } + + // Check for successful token response + if resp.StatusCode == http.StatusOK && tokenResp.AccessToken != "" { + return &auth.Credentials{ + AccessToken: tokenResp.AccessToken, + RefreshToken: tokenResp.RefreshToken, + ExpiresAt: time.Now().Add(time.Duration(tokenResp.ExpiresIn) * time.Second), + DeviceName: deviceName, + User: tokenResp.User, + WSURL: tokenResp.WSURL, + }, nil + } + + // Non-200 status without a recognized error + if resp.StatusCode != http.StatusOK { + fmt.Fprintf(os.Stderr, "Unexpected status %s (will retry)\n", resp.Status) + continue + } + } +} + +// openBrowserFunc is the function used to open URLs in the browser. +// It can be replaced in tests to prevent actual browser launches. +var openBrowserFunc = openBrowserDefault + +// openBrowserDefault attempts to open the given URL in the default browser. +func openBrowserDefault(url string) { + var cmd *exec.Cmd + switch runtime.GOOS { + case "darwin": + cmd = exec.Command("open", url) + case "linux": + cmd = exec.Command("xdg-open", url) + case "windows": + cmd = exec.Command("rundll32", "url.dll,FileProtocolHandler", url) + default: + return + } + // Ignore errors — browser open is best-effort + cmd.Start() +} diff --git a/internal/cmd/login_test.go b/internal/cmd/login_test.go new file mode 100644 index 00000000..92f1dd7e --- /dev/null +++ b/internal/cmd/login_test.go @@ -0,0 +1,476 @@ +package cmd + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "strings" + "sync/atomic" + "testing" + + "github.com/minicodemonkey/chief/internal/auth" +) + +func init() { + // Prevent tests from opening actual browser windows + openBrowserFunc = func(url string) {} +} + +func setTestHome(t *testing.T, dir string) { + t.Helper() + t.Setenv("HOME", dir) +} + +func TestRunLogin_Success(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + var pollCount atomic.Int32 + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/oauth/device/code": + json.NewEncoder(w).Encode(deviceCodeResponse{ + DeviceCode: "test-device-code", + UserCode: "ABCD-1234", + }) + case "/api/oauth/device/token": + count := pollCount.Add(1) + if count < 2 { + // First poll: authorization pending + json.NewEncoder(w).Encode(tokenResponse{ + Error: "authorization_pending", + }) + return + } + // Second poll: success + json.NewEncoder(w).Encode(tokenResponse{ + AccessToken: "test-access-token", + RefreshToken: "test-refresh-token", + ExpiresIn: 3600, + User: "testuser@example.com", + WSURL: "wss://ws-test-reverb.laravel.cloud/ws/server", + }) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + // Override stdin to avoid blocking on "already logged in" prompt + oldStdin := os.Stdin + defer func() { os.Stdin = oldStdin }() + r, w, _ := os.Pipe() + os.Stdin = r + w.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "test-device", + BaseURL: server.URL, + }) + if err != nil { + t.Fatalf("RunLogin failed: %v", err) + } + + // Verify credentials were saved + creds, err := auth.LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials after login failed: %v", err) + } + if creds.AccessToken != "test-access-token" { + t.Errorf("expected access_token %q, got %q", "test-access-token", creds.AccessToken) + } + if creds.RefreshToken != "test-refresh-token" { + t.Errorf("expected refresh_token %q, got %q", "test-refresh-token", creds.RefreshToken) + } + if creds.User != "testuser@example.com" { + t.Errorf("expected user %q, got %q", "testuser@example.com", creds.User) + } + if creds.DeviceName != "test-device" { + t.Errorf("expected device_name %q, got %q", "test-device", creds.DeviceName) + } + if creds.WSURL != "wss://ws-test-reverb.laravel.cloud/ws/server" { + t.Errorf("expected ws_url %q, got %q", "wss://ws-test-reverb.laravel.cloud/ws/server", creds.WSURL) + } +} + +func TestRunLogin_DeviceCodeError(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "test-device", + BaseURL: server.URL, + }) + if err == nil { + t.Fatal("expected error for server error response") + } +} + +func TestRunLogin_AuthorizationDenied(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/oauth/device/code": + json.NewEncoder(w).Encode(deviceCodeResponse{ + DeviceCode: "test-device-code", + UserCode: "ABCD-1234", + }) + case "/api/oauth/device/token": + json.NewEncoder(w).Encode(tokenResponse{ + Error: "access_denied", + }) + } + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "test-device", + BaseURL: server.URL, + }) + if err == nil { + t.Fatal("expected error for denied authorization") + } +} + +func TestRunLogin_DefaultDeviceName(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + var receivedDeviceName string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/oauth/device/code": + var body map[string]string + json.NewDecoder(r.Body).Decode(&body) + receivedDeviceName = body["device_name"] + json.NewEncoder(w).Encode(deviceCodeResponse{ + DeviceCode: "test-device-code", + UserCode: "TEST-CODE", + }) + case "/api/oauth/device/token": + json.NewEncoder(w).Encode(tokenResponse{ + AccessToken: "token", + RefreshToken: "refresh", + ExpiresIn: 3600, + User: "user", + }) + } + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + BaseURL: server.URL, + // DeviceName left empty — should default to hostname + }) + if err != nil { + t.Fatalf("RunLogin failed: %v", err) + } + + hostname, _ := os.Hostname() + if receivedDeviceName != hostname { + t.Errorf("expected device name %q (hostname), got %q", hostname, receivedDeviceName) + } +} + +func TestRunLogin_AlreadyLoggedIn_Decline(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Save existing credentials + existing := &auth.Credentials{ + AccessToken: "existing-token", + User: "existing-user", + DeviceName: "existing-device", + } + if err := auth.SaveCredentials(existing); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + // Pipe "n\n" to stdin to decline + oldStdin := os.Stdin + defer func() { os.Stdin = oldStdin }() + r, w, _ := os.Pipe() + w.Write([]byte("n\n")) + w.Close() + os.Stdin = r + + // Server should not be called at all when declining + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + t.Error("server should not be called when login is declined") + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "new-device", + BaseURL: server.URL, + }) + if err != nil { + t.Fatalf("RunLogin should not error when declining: %v", err) + } + + // Credentials should remain unchanged + creds, err := auth.LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials failed: %v", err) + } + if creds.AccessToken != "existing-token" { + t.Errorf("credentials should not have changed, got access_token %q", creds.AccessToken) + } +} + +func TestRunLogin_SetupToken_Success(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + var receivedToken string + var receivedDeviceName string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/oauth/device/exchange" { + http.NotFound(w, r) + return + } + var body map[string]string + json.NewDecoder(r.Body).Decode(&body) + receivedToken = body["setup_token"] + receivedDeviceName = body["device_name"] + json.NewEncoder(w).Encode(tokenResponse{ + AccessToken: "setup-access-token", + RefreshToken: "setup-refresh-token", + ExpiresIn: 3600, + User: "setupuser@example.com", + WSURL: "wss://ws-setup-reverb.laravel.cloud/ws/server", + }) + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "setup-device", + BaseURL: server.URL, + SetupToken: "test-setup-token-abc123", + }) + if err != nil { + t.Fatalf("RunLogin with setup token failed: %v", err) + } + + if receivedToken != "test-setup-token-abc123" { + t.Errorf("expected setup_token %q, got %q", "test-setup-token-abc123", receivedToken) + } + if receivedDeviceName != "setup-device" { + t.Errorf("expected device_name %q, got %q", "setup-device", receivedDeviceName) + } + + // Verify credentials were saved + creds, err := auth.LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials after setup token login failed: %v", err) + } + if creds.AccessToken != "setup-access-token" { + t.Errorf("expected access_token %q, got %q", "setup-access-token", creds.AccessToken) + } + if creds.RefreshToken != "setup-refresh-token" { + t.Errorf("expected refresh_token %q, got %q", "setup-refresh-token", creds.RefreshToken) + } + if creds.User != "setupuser@example.com" { + t.Errorf("expected user %q, got %q", "setupuser@example.com", creds.User) + } + if creds.DeviceName != "setup-device" { + t.Errorf("expected device_name %q, got %q", "setup-device", creds.DeviceName) + } + if creds.WSURL != "wss://ws-setup-reverb.laravel.cloud/ws/server" { + t.Errorf("expected ws_url %q, got %q", "wss://ws-setup-reverb.laravel.cloud/ws/server", creds.WSURL) + } +} + +func TestRunLogin_WSURLNotReturned(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + var pollCount atomic.Int32 + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/oauth/device/code": + json.NewEncoder(w).Encode(deviceCodeResponse{ + DeviceCode: "test-device-code", + UserCode: "ABCD-1234", + }) + case "/api/oauth/device/token": + count := pollCount.Add(1) + if count < 2 { + json.NewEncoder(w).Encode(tokenResponse{ + Error: "authorization_pending", + }) + return + } + // No ws_url in response + json.NewEncoder(w).Encode(tokenResponse{ + AccessToken: "test-access-token", + RefreshToken: "test-refresh-token", + ExpiresIn: 3600, + User: "testuser@example.com", + }) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + oldStdin := os.Stdin + defer func() { os.Stdin = oldStdin }() + r, w, _ := os.Pipe() + os.Stdin = r + w.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "test-device", + BaseURL: server.URL, + }) + if err != nil { + t.Fatalf("RunLogin failed: %v", err) + } + + creds, err := auth.LoadCredentials() + if err != nil { + t.Fatalf("LoadCredentials after login failed: %v", err) + } + if creds.WSURL != "" { + t.Errorf("expected empty ws_url, got %q", creds.WSURL) + } +} + +func TestRunLogin_SetupToken_InvalidToken(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/oauth/device/exchange" { + http.NotFound(w, r) + return + } + w.WriteHeader(http.StatusUnauthorized) + json.NewEncoder(w).Encode(tokenResponse{ + Error: "invalid_token", + }) + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "setup-device", + BaseURL: server.URL, + SetupToken: "expired-token", + }) + if err == nil { + t.Fatal("expected error for invalid setup token") + } + if !strings.Contains(err.Error(), "invalid_token") { + t.Errorf("expected error to mention invalid_token, got: %v", err) + } + + // Verify no credentials were saved + _, err = auth.LoadCredentials() + if err == nil { + t.Error("credentials should not have been saved for invalid token") + } +} + +func TestRunLogin_SetupToken_ExpiredToken(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/oauth/device/exchange" { + http.NotFound(w, r) + return + } + w.WriteHeader(http.StatusGone) + json.NewEncoder(w).Encode(tokenResponse{ + Error: "token_expired", + }) + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "setup-device", + BaseURL: server.URL, + SetupToken: "expired-token", + }) + if err == nil { + t.Fatal("expected error for expired setup token") + } + if !strings.Contains(err.Error(), "token_expired") { + t.Errorf("expected error to mention token_expired, got: %v", err) + } +} + +func TestRunLogin_SetupToken_ServerError(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + DeviceName: "setup-device", + BaseURL: server.URL, + SetupToken: "some-token", + }) + if err == nil { + t.Fatal("expected error for server error") + } +} + +func TestRunLogin_SetupToken_DefaultDeviceName(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + var receivedDeviceName string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/oauth/device/exchange" { + http.NotFound(w, r) + return + } + var body map[string]string + json.NewDecoder(r.Body).Decode(&body) + receivedDeviceName = body["device_name"] + json.NewEncoder(w).Encode(tokenResponse{ + AccessToken: "token", + RefreshToken: "refresh", + ExpiresIn: 3600, + User: "user", + }) + })) + defer server.Close() + + err := RunLogin(LoginOptions{ + BaseURL: server.URL, + SetupToken: "test-token", + // DeviceName left empty — should default to hostname + }) + if err != nil { + t.Fatalf("RunLogin failed: %v", err) + } + + hostname, _ := os.Hostname() + if receivedDeviceName != hostname { + t.Errorf("expected device name %q (hostname), got %q", hostname, receivedDeviceName) + } +} + +func TestOpenBrowser(t *testing.T) { + // Just verifying it doesn't panic — browser open is best-effort + openBrowserFunc("https://example.com") +} diff --git a/internal/cmd/logout.go b/internal/cmd/logout.go new file mode 100644 index 00000000..4174d45a --- /dev/null +++ b/internal/cmd/logout.go @@ -0,0 +1,42 @@ +package cmd + +import ( + "errors" + "fmt" + + "github.com/minicodemonkey/chief/internal/auth" +) + +// LogoutOptions contains configuration for the logout command. +type LogoutOptions struct { + BaseURL string // Override base URL (for testing) +} + +// RunLogout performs the device logout flow. +func RunLogout(opts LogoutOptions) error { + // Load credentials to get device name and access token + creds, err := auth.LoadCredentials() + if err != nil { + if errors.Is(err, auth.ErrNotLoggedIn) { + fmt.Println("Not logged in.") + return nil + } + return err + } + + deviceName := creds.DeviceName + + // Call revocation endpoint + if err := auth.RevokeDevice(creds.AccessToken, opts.BaseURL); err != nil { + fmt.Printf("Warning: could not deauthorize device server-side: %v\n", err) + fmt.Println("Local credentials will still be removed.") + } + + // Delete local credentials + if err := auth.DeleteCredentials(); err != nil { + return fmt.Errorf("removing local credentials: %w", err) + } + + fmt.Printf("Logged out. Device %q has been deauthorized.\n", deviceName) + return nil +} diff --git a/internal/cmd/logout_test.go b/internal/cmd/logout_test.go new file mode 100644 index 00000000..fa8c2bdb --- /dev/null +++ b/internal/cmd/logout_test.go @@ -0,0 +1,100 @@ +package cmd + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/auth" +) + +func TestRunLogout_Success(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Save credentials + creds := &auth.Credentials{ + AccessToken: "test-token", + RefreshToken: "test-refresh", + ExpiresAt: time.Now().Add(time.Hour), + DeviceName: "test-device", + User: "user@example.com", + } + if err := auth.SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + var revokedToken string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/oauth/revoke" { + var body map[string]string + json.NewDecoder(r.Body).Decode(&body) + revokedToken = body["access_token"] + w.WriteHeader(http.StatusOK) + return + } + http.NotFound(w, r) + })) + defer server.Close() + + err := RunLogout(LogoutOptions{BaseURL: server.URL}) + if err != nil { + t.Fatalf("RunLogout failed: %v", err) + } + + // Verify revocation was called with correct token + if revokedToken != "test-token" { + t.Errorf("expected revoked token %q, got %q", "test-token", revokedToken) + } + + // Verify credentials are deleted + _, err = auth.LoadCredentials() + if err != auth.ErrNotLoggedIn { + t.Fatalf("expected ErrNotLoggedIn after logout, got %v", err) + } +} + +func TestRunLogout_NotLoggedIn(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + err := RunLogout(LogoutOptions{}) + if err != nil { + t.Fatalf("RunLogout should not error when not logged in: %v", err) + } +} + +func TestRunLogout_RevocationFails_StillDeletesCredentials(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + creds := &auth.Credentials{ + AccessToken: "test-token", + RefreshToken: "test-refresh", + ExpiresAt: time.Now().Add(time.Hour), + DeviceName: "test-device", + User: "user@example.com", + } + if err := auth.SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + // Server returns error on revocation + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + err := RunLogout(LogoutOptions{BaseURL: server.URL}) + if err != nil { + t.Fatalf("RunLogout should not error when revocation fails: %v", err) + } + + // Credentials should still be deleted + _, err = auth.LoadCredentials() + if err != auth.ErrNotLoggedIn { + t.Fatalf("expected ErrNotLoggedIn after logout, got %v", err) + } +} diff --git a/internal/cmd/logs.go b/internal/cmd/logs.go new file mode 100644 index 00000000..761de5d5 --- /dev/null +++ b/internal/cmd/logs.go @@ -0,0 +1,210 @@ +package cmd + +import ( + "bufio" + "encoding/json" + "fmt" + "log" + "os" + "path/filepath" + "strings" + "sync" + + "github.com/minicodemonkey/chief/internal/ws" +) + +// storyLogger manages per-story log files during Ralph loop runs. +type storyLogger struct { + mu sync.Mutex + logDir string // .chief/prds/<id>/logs/ + files map[string]*os.File // story_id -> open file +} + +// newStoryLogger creates a story logger for a given PRD. +// It creates the logs directory and removes any previous log files (V1 simplicity). +func newStoryLogger(prdPath string) (*storyLogger, error) { + prdDir := filepath.Dir(prdPath) + logDir := filepath.Join(prdDir, "logs") + + // Remove previous logs (overwrite on new run) + os.RemoveAll(logDir) + + if err := os.MkdirAll(logDir, 0o755); err != nil { + return nil, fmt.Errorf("creating log directory: %w", err) + } + + return &storyLogger{ + logDir: logDir, + files: make(map[string]*os.File), + }, nil +} + +// WriteLog writes a line to the log file for the given story. +func (sl *storyLogger) WriteLog(storyID, line string) { + if storyID == "" { + return + } + + sl.mu.Lock() + defer sl.mu.Unlock() + + f, ok := sl.files[storyID] + if !ok { + var err error + logPath := filepath.Join(sl.logDir, storyID+".log") + f, err = os.OpenFile(logPath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644) + if err != nil { + log.Printf("Error opening story log file %s: %v", logPath, err) + return + } + sl.files[storyID] = f + } + + f.WriteString(line + "\n") +} + +// Close closes all open log files. +func (sl *storyLogger) Close() { + sl.mu.Lock() + defer sl.mu.Unlock() + + for _, f := range sl.files { + f.Close() + } + sl.files = make(map[string]*os.File) +} + +// handleGetLogs handles a get_logs request. +func handleGetLogs(sender messageSender, finder projectFinder, msg ws.Message) { + var req ws.GetLogsMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing get_logs message: %v", err) + return + } + + project, found := finder.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + prdDir := filepath.Join(project.Path, ".chief", "prds", req.PRDID) + if _, err := os.Stat(prdDir); os.IsNotExist(err) { + sendError(sender, ws.ErrCodePRDNotFound, + fmt.Sprintf("PRD %q not found in project %q", req.PRDID, req.Project), msg.ID) + return + } + + logDir := filepath.Join(prdDir, "logs") + + // If story_id is provided, return that specific story's logs + if req.StoryID != "" { + lines, err := readLogFile(filepath.Join(logDir, req.StoryID+".log"), req.Lines) + if err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Failed to read logs for story %q: %v", req.StoryID, err), msg.ID) + return + } + + sendLogLines(sender, req.Project, req.PRDID, req.StoryID, lines) + return + } + + // If story_id is omitted, return the most recent log activity for the PRD + storyID, lines, err := readMostRecentLog(logDir, req.Lines) + if err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Failed to read logs: %v", err), msg.ID) + return + } + + sendLogLines(sender, req.Project, req.PRDID, storyID, lines) +} + +// readLogFile reads lines from a log file. If maxLines is 0, reads all lines. +func readLogFile(path string, maxLines int) ([]string, error) { + f, err := os.Open(path) + if err != nil { + if os.IsNotExist(err) { + return []string{}, nil + } + return nil, err + } + defer f.Close() + + var lines []string + scanner := bufio.NewScanner(f) + // Increase buffer size for long lines + buf := make([]byte, 0, 64*1024) + scanner.Buffer(buf, 1024*1024) + + for scanner.Scan() { + lines = append(lines, scanner.Text()) + } + + if maxLines > 0 && len(lines) > maxLines { + lines = lines[len(lines)-maxLines:] + } + + return lines, scanner.Err() +} + +// readMostRecentLog finds the most recently modified log file and reads it. +func readMostRecentLog(logDir string, maxLines int) (string, []string, error) { + entries, err := os.ReadDir(logDir) + if err != nil { + if os.IsNotExist(err) { + return "", []string{}, nil + } + return "", nil, err + } + + // Find the most recently modified .log file + var mostRecent string + var mostRecentTime int64 + + for _, entry := range entries { + if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".log") { + continue + } + info, err := entry.Info() + if err != nil { + continue + } + if info.ModTime().UnixNano() > mostRecentTime { + mostRecentTime = info.ModTime().UnixNano() + mostRecent = entry.Name() + } + } + + if mostRecent == "" { + return "", []string{}, nil + } + + storyID := strings.TrimSuffix(mostRecent, ".log") + lines, err := readLogFile(filepath.Join(logDir, mostRecent), maxLines) + return storyID, lines, err +} + +// sendLogLines sends a log_lines message over WebSocket. +func sendLogLines(sender messageSender, project, prdID, storyID string, lines []string) { + if lines == nil { + lines = []string{} + } + + envelope := ws.NewMessage(ws.TypeLogLines) + msg := ws.LogLinesMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: project, + PRDID: prdID, + StoryID: storyID, + Lines: lines, + Level: "info", + } + if err := sender.Send(msg); err != nil { + log.Printf("Error sending log_lines: %v", err) + } +} diff --git a/internal/cmd/logs_test.go b/internal/cmd/logs_test.go new file mode 100644 index 00000000..bb1bdf8c --- /dev/null +++ b/internal/cmd/logs_test.go @@ -0,0 +1,757 @@ +package cmd + +import ( + "encoding/json" + "os" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/engine" +) + +func TestStoryLogger_WriteAndRead(t *testing.T) { + tmpDir := t.TempDir() + prdDir := filepath.Join(tmpDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(`{}`), 0o644); err != nil { + t.Fatal(err) + } + + sl, err := newStoryLogger(prdPath) + if err != nil { + t.Fatalf("newStoryLogger failed: %v", err) + } + defer sl.Close() + + // Write some log lines + sl.WriteLog("US-001", "Starting story US-001") + sl.WriteLog("US-001", "Working on implementation") + sl.WriteLog("US-001", "Story complete") + + sl.WriteLog("US-002", "Starting story US-002") + sl.WriteLog("US-002", "Done") + + // Close to flush + sl.Close() + + // Read the log files + logDir := filepath.Join(prdDir, "logs") + + lines, err := readLogFile(filepath.Join(logDir, "US-001.log"), 0) + if err != nil { + t.Fatalf("readLogFile failed: %v", err) + } + if len(lines) != 3 { + t.Errorf("expected 3 lines for US-001, got %d", len(lines)) + } + if lines[0] != "Starting story US-001" { + t.Errorf("expected first line 'Starting story US-001', got %q", lines[0]) + } + + lines, err = readLogFile(filepath.Join(logDir, "US-002.log"), 0) + if err != nil { + t.Fatalf("readLogFile failed: %v", err) + } + if len(lines) != 2 { + t.Errorf("expected 2 lines for US-002, got %d", len(lines)) + } +} + +func TestStoryLogger_WriteEmptyStoryID(t *testing.T) { + tmpDir := t.TempDir() + prdDir := filepath.Join(tmpDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(`{}`), 0o644); err != nil { + t.Fatal(err) + } + + sl, err := newStoryLogger(prdPath) + if err != nil { + t.Fatalf("newStoryLogger failed: %v", err) + } + defer sl.Close() + + // Writing with empty story ID should be a no-op + sl.WriteLog("", "This should not be written") + sl.Close() + + // Verify no files were created + logDir := filepath.Join(prdDir, "logs") + entries, err := os.ReadDir(logDir) + if err != nil { + t.Fatalf("ReadDir failed: %v", err) + } + if len(entries) != 0 { + t.Errorf("expected no log files, got %d", len(entries)) + } +} + +func TestStoryLogger_OverwriteOnNewRun(t *testing.T) { + tmpDir := t.TempDir() + prdDir := filepath.Join(tmpDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(`{}`), 0o644); err != nil { + t.Fatal(err) + } + + // Create first logger and write some logs + sl1, err := newStoryLogger(prdPath) + if err != nil { + t.Fatal(err) + } + sl1.WriteLog("US-001", "First run output") + sl1.Close() + + // Create second logger — should overwrite previous logs + sl2, err := newStoryLogger(prdPath) + if err != nil { + t.Fatal(err) + } + sl2.WriteLog("US-001", "Second run output") + sl2.Close() + + // Read the log — should only have second run's content + logDir := filepath.Join(prdDir, "logs") + lines, err := readLogFile(filepath.Join(logDir, "US-001.log"), 0) + if err != nil { + t.Fatalf("readLogFile failed: %v", err) + } + if len(lines) != 1 { + t.Errorf("expected 1 line (overwritten), got %d: %v", len(lines), lines) + } + if len(lines) > 0 && lines[0] != "Second run output" { + t.Errorf("expected 'Second run output', got %q", lines[0]) + } +} + +func TestReadLogFile_WithLineLimit(t *testing.T) { + tmpDir := t.TempDir() + logPath := filepath.Join(tmpDir, "test.log") + + // Write 10 lines + var content string + for i := 1; i <= 10; i++ { + content += "Line " + string(rune('0'+i)) + "\n" + } + if err := os.WriteFile(logPath, []byte(content), 0o644); err != nil { + t.Fatal(err) + } + + // Read with limit of 3 — should get last 3 lines + lines, err := readLogFile(logPath, 3) + if err != nil { + t.Fatalf("readLogFile failed: %v", err) + } + if len(lines) != 3 { + t.Errorf("expected 3 lines, got %d", len(lines)) + } +} + +func TestReadLogFile_Nonexistent(t *testing.T) { + lines, err := readLogFile("/nonexistent/path/test.log", 0) + if err != nil { + t.Fatalf("expected no error for nonexistent file, got: %v", err) + } + if len(lines) != 0 { + t.Errorf("expected empty lines for nonexistent file, got %d", len(lines)) + } +} + +func TestReadMostRecentLog(t *testing.T) { + tmpDir := t.TempDir() + logDir := filepath.Join(tmpDir, "logs") + if err := os.MkdirAll(logDir, 0o755); err != nil { + t.Fatal(err) + } + + // Write two log files with different mod times + if err := os.WriteFile(filepath.Join(logDir, "US-001.log"), []byte("old log\n"), 0o644); err != nil { + t.Fatal(err) + } + + // Ensure US-002 is newer + time.Sleep(10 * time.Millisecond) + if err := os.WriteFile(filepath.Join(logDir, "US-002.log"), []byte("new log line 1\nnew log line 2\n"), 0o644); err != nil { + t.Fatal(err) + } + + storyID, lines, err := readMostRecentLog(logDir, 0) + if err != nil { + t.Fatalf("readMostRecentLog failed: %v", err) + } + if storyID != "US-002" { + t.Errorf("expected most recent story 'US-002', got %q", storyID) + } + if len(lines) != 2 { + t.Errorf("expected 2 lines, got %d", len(lines)) + } +} + +func TestReadMostRecentLog_EmptyDir(t *testing.T) { + tmpDir := t.TempDir() + logDir := filepath.Join(tmpDir, "logs") + if err := os.MkdirAll(logDir, 0o755); err != nil { + t.Fatal(err) + } + + storyID, lines, err := readMostRecentLog(logDir, 0) + if err != nil { + t.Fatalf("readMostRecentLog failed: %v", err) + } + if storyID != "" { + t.Errorf("expected empty story ID, got %q", storyID) + } + if len(lines) != 0 { + t.Errorf("expected no lines, got %d", len(lines)) + } +} + +func TestReadMostRecentLog_NonexistentDir(t *testing.T) { + storyID, lines, err := readMostRecentLog("/nonexistent/logs", 0) + if err != nil { + t.Fatalf("expected no error, got: %v", err) + } + if storyID != "" || len(lines) != 0 { + t.Errorf("expected empty results, got storyID=%q lines=%v", storyID, lines) + } +} + +func TestRunManager_StoryLogWriting(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Create a temp project with a PRD + projectDir := t.TempDir() + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "Test", "userStories": [{"id": "US-001", "title": "Story", "passes": false}]}` + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Start a run (creates logger) + err := rm.startRun("myproject", "feature", projectDir) + if err != nil { + t.Fatalf("startRun failed: %v", err) + } + + // Verify logger was created + rm.mu.RLock() + _, hasLogger := rm.loggers["myproject/feature"] + rm.mu.RUnlock() + if !hasLogger { + t.Fatal("expected logger to be created for the run") + } + + // Write some story logs + rm.writeStoryLog("myproject/feature", "US-001", "Hello from story log") + rm.writeStoryLog("myproject/feature", "US-001", "Another line") + + // Stop and cleanup + rm.stopAll() + + // Read the log file + logPath := filepath.Join(prdDir, "logs", "US-001.log") + lines, err := readLogFile(logPath, 0) + if err != nil { + t.Fatalf("readLogFile failed: %v", err) + } + if len(lines) != 2 { + t.Errorf("expected 2 lines, got %d", len(lines)) + } + if len(lines) > 0 && lines[0] != "Hello from story log" { + t.Errorf("expected 'Hello from story log', got %q", lines[0]) + } +} + +func TestRunServe_GetLogs(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo with a PRD that has log files + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + logDir := filepath.Join(prdDir, "logs") + if err := os.MkdirAll(logDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test Story", "passes": true}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Write a log file for US-001 + logContent := "Starting story US-001\nWorking on implementation\nStory complete\n" + if err := os.WriteFile(filepath.Join(logDir, "US-001.log"), []byte(logContent), 0o644); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send get_logs request with story_id + getLogsReq := map[string]interface{}{ + "type": "get_logs", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + "story_id": "US-001", + } + ms.sendCommand(getLogsReq) + + raw, err := ms.waitForMessageType("log_lines", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "log_lines" { + t.Errorf("expected type 'log_lines', got %v", response["type"]) + } + if response["project"] != "myproject" { + t.Errorf("expected project 'myproject', got %v", response["project"]) + } + if response["prd_id"] != "feature" { + t.Errorf("expected prd_id 'feature', got %v", response["prd_id"]) + } + if response["story_id"] != "US-001" { + t.Errorf("expected story_id 'US-001', got %v", response["story_id"]) + } + + lines, ok := response["lines"].([]interface{}) + if !ok { + t.Fatal("expected lines to be an array") + } + if len(lines) != 3 { + t.Errorf("expected 3 lines, got %d", len(lines)) + } +} + +func TestRunServe_GetLogsNoStoryID(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + logDir := filepath.Join(prdDir, "logs") + if err := os.MkdirAll(logDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test Story", "passes": true}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Write a log file — when no story_id provided, should return most recent + if err := os.WriteFile(filepath.Join(logDir, "US-001.log"), []byte("recent log line\n"), 0o644); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + getLogsReq := map[string]interface{}{ + "type": "get_logs", + "id": "req-2", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(getLogsReq) + + raw, err := ms.waitForMessageType("log_lines", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "log_lines" { + t.Errorf("expected type 'log_lines', got %v", response["type"]) + } + if response["story_id"] != "US-001" { + t.Errorf("expected story_id 'US-001', got %v", response["story_id"]) + } +} + +func TestRunServe_GetLogsProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + getLogsReq := map[string]interface{}{ + "type": "get_logs", + "id": "req-3", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "prd_id": "feature", + "story_id": "US-001", + } + ms.sendCommand(getLogsReq) + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "error" { + t.Errorf("expected type 'error', got %v", response["type"]) + } + if response["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", response["code"]) + } +} + +func TestRunServe_GetLogsPRDNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + getLogsReq := map[string]interface{}{ + "type": "get_logs", + "id": "req-4", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "nonexistent", + "story_id": "US-001", + } + ms.sendCommand(getLogsReq) + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "error" { + t.Errorf("expected type 'error', got %v", response["type"]) + } + if response["code"] != "PRD_NOT_FOUND" { + t.Errorf("expected code 'PRD_NOT_FOUND', got %v", response["code"]) + } +} + +func TestRunServe_GetLogsWithLineLimit(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + logDir := filepath.Join(prdDir, "logs") + if err := os.MkdirAll(logDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test Story", "passes": true}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Write 5 log lines + logContent := "Line 1\nLine 2\nLine 3\nLine 4\nLine 5\n" + if err := os.WriteFile(filepath.Join(logDir, "US-001.log"), []byte(logContent), 0o644); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Request only 2 lines + getLogsReq := map[string]interface{}{ + "type": "get_logs", + "id": "req-5", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + "story_id": "US-001", + "lines": 2, + } + ms.sendCommand(getLogsReq) + + raw, err := ms.waitForMessageType("log_lines", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("response was not received") + } + if response["type"] != "log_lines" { + t.Errorf("expected type 'log_lines', got %v", response["type"]) + } + + lines, ok := response["lines"].([]interface{}) + if !ok { + t.Fatal("expected lines to be an array") + } + if len(lines) != 2 { + t.Errorf("expected 2 lines (limited), got %d", len(lines)) + } + // Should return the last 2 lines + if len(lines) >= 2 { + if lines[0] != "Line 4" { + t.Errorf("expected 'Line 4', got %v", lines[0]) + } + if lines[1] != "Line 5" { + t.Errorf("expected 'Line 5', got %v", lines[1]) + } + } +} + +func TestRunServe_LoggingIntegration(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test Story", "passes": false}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Create a mock claude that outputs stream-json + mockDir := t.TempDir() + mockScript := `#!/bin/sh +echo '{"type":"system","subtype":"init"}' +echo '{"type":"assistant","message":{"content":[{"type":"text","text":"Working on <ralph-status>US-001</ralph-status>"}]}}' +echo '{"type":"assistant","message":{"content":[{"type":"text","text":"Implementing feature"}]}}' +echo '{"type":"result"}' +exit 0 +` + if err := os.WriteFile(filepath.Join(mockDir, "claude"), []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", mockDir+":"+origPath) + + var messages []map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send start_run request + startReq := map[string]string{ + "type": "start_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(startReq) + + // Wait for multiple messages (expecting at least some stream messages) + rawMessages, err := ms.waitForMessages(15, 5*time.Second) + if err == nil { + mu.Lock() + for _, raw := range rawMessages { + var msg map[string]interface{} + if json.Unmarshal(raw, &msg) == nil { + messages = append(messages, msg) + } + } + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify that log files were created in .chief/prds/feature/logs/ + logDir := filepath.Join(prdDir, "logs") + if _, err := os.Stat(logDir); os.IsNotExist(err) { + t.Error("expected logs directory to be created") + } + + // The log file should exist for US-001 (the story that was started) + logFile := filepath.Join(logDir, "US-001.log") + if _, err := os.Stat(logFile); os.IsNotExist(err) { + t.Error("expected US-001.log to be created") + } else { + lines, err := readLogFile(logFile, 0) + if err != nil { + t.Fatalf("readLogFile failed: %v", err) + } + // Should have at least some log content + if len(lines) == 0 { + t.Error("expected log file to have content") + } + } +} + +func TestRunManager_CleanupClosesLogger(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + tmpDir := t.TempDir() + prdDir := filepath.Join(tmpDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdPath := filepath.Join(prdDir, "prd.json") + prdState := `{"project": "Test", "userStories": [{"id": "US-001", "title": "Story", "passes": false}]}` + if err := os.WriteFile(prdPath, []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + if err := rm.startRun("myproject", "feature", tmpDir); err != nil { + t.Fatalf("startRun failed: %v", err) + } + + key := "myproject/feature" + + // Verify logger exists + rm.mu.RLock() + _, hasLogger := rm.loggers[key] + rm.mu.RUnlock() + if !hasLogger { + t.Fatal("expected logger to be created") + } + + // Cleanup should remove the logger + rm.cleanup(key) + + rm.mu.RLock() + _, hasLogger = rm.loggers[key] + rm.mu.RUnlock() + if hasLogger { + t.Error("expected logger to be removed after cleanup") + } +} diff --git a/internal/cmd/new.go b/internal/cmd/new.go index 04922b7d..3d7b6dbf 100644 --- a/internal/cmd/new.go +++ b/internal/cmd/new.go @@ -6,22 +6,21 @@ package cmd import ( "fmt" "os" + "os/exec" "path/filepath" "github.com/minicodemonkey/chief/embed" - "github.com/minicodemonkey/chief/internal/loop" "github.com/minicodemonkey/chief/internal/prd" ) // NewOptions contains configuration for the new command. type NewOptions struct { - Name string // PRD name (default: "main") - Context string // Optional context to pass to the agent - BaseDir string // Base directory for .chief/prds/ (default: current directory) - Provider loop.Provider // Agent CLI provider (Claude or Codex) + Name string // PRD name (default: "main") + Context string // Optional context to pass to Claude + BaseDir string // Base directory for .chief/prds/ (default: current directory) } -// RunNew creates a new PRD by launching an interactive agent session. +// RunNew creates a new PRD by launching an interactive Claude session. func RunNew(opts NewOptions) error { // Set defaults if opts.Name == "" { @@ -54,51 +53,67 @@ func RunNew(opts NewOptions) error { // Get the init prompt with the PRD directory path prompt := embed.GetInitPrompt(prdDir, opts.Context) - if opts.Provider == nil { - return fmt.Errorf("new command requires Provider to be set") - } - // Launch interactive agent session + // Launch interactive Claude session fmt.Printf("Creating PRD in %s...\n", prdDir) - fmt.Printf("Launching %s to help you create your PRD...\n", opts.Provider.Name()) + fmt.Println("Launching Claude to help you create your PRD...") fmt.Println() - if err := runInteractiveAgent(opts.Provider, opts.BaseDir, prompt); err != nil { - return fmt.Errorf("%s session failed: %w", opts.Provider.Name(), err) + if err := runInteractiveClaude(opts.BaseDir, prompt); err != nil { + return fmt.Errorf("Claude session failed: %w", err) } // Check if prd.md was created if _, err := os.Stat(prdMdPath); os.IsNotExist(err) { - // Clean up empty directory to prevent broken picker entries - os.Remove(prdDir) fmt.Println("\nNo prd.md was created. Run 'chief new' again to try again.") return nil } - // Validate the created prd.md can be parsed - if _, err := prd.ParseMarkdownPRD(prdMdPath); err != nil { - fmt.Printf("\nWarning: prd.md was created but could not be parsed: %v\n", err) - fmt.Println("You may need to edit it to match the expected format.") - } else { - fmt.Println("\nPRD created successfully!") + fmt.Println("\nPRD created successfully!") + + // Run conversion from prd.md to prd.json + if err := RunConvert(prdDir); err != nil { + return fmt.Errorf("conversion failed: %w", err) } fmt.Printf("\nYour PRD is ready! Run 'chief' or 'chief %s' to start working on it.\n", opts.Name) return nil } -// runInteractiveAgent launches an interactive agent session in the specified directory. -func runInteractiveAgent(provider loop.Provider, workDir, prompt string) error { - if provider == nil { - return fmt.Errorf("interactive agent requires Provider to be set") - } - cmd := provider.InteractiveCommand(workDir, prompt) +// runInteractiveClaude launches an interactive Claude session in the specified directory. +func runInteractiveClaude(workDir, prompt string) error { + // Pass prompt as argument (not -p which is print mode / non-interactive) + cmd := exec.Command("claude", prompt) + cmd.Dir = workDir cmd.Stdin = os.Stdin cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr + return cmd.Run() } +// ConvertOptions contains configuration for the conversion command. +type ConvertOptions struct { + PRDDir string // PRD directory containing prd.md + Merge bool // Auto-merge without prompting on conversion conflicts + Force bool // Auto-overwrite without prompting on conversion conflicts +} + +// RunConvert converts prd.md to prd.json using Claude. +func RunConvert(prdDir string) error { + return RunConvertWithOptions(ConvertOptions{PRDDir: prdDir}) +} + +// RunConvertWithOptions converts prd.md to prd.json using Claude with options. +// The Merge and Force flags will be fully implemented in US-019. +func RunConvertWithOptions(opts ConvertOptions) error { + return prd.Convert(prd.ConvertOptions{ + PRDDir: opts.PRDDir, + Merge: opts.Merge, + Force: opts.Force, + }) +} + // isValidPRDName checks if the name contains only valid characters. func isValidPRDName(name string) bool { if name == "" { diff --git a/internal/cmd/new_test.go b/internal/cmd/new_test.go deleted file mode 100644 index 02f672bf..00000000 --- a/internal/cmd/new_test.go +++ /dev/null @@ -1,169 +0,0 @@ -package cmd - -import ( - "os" - "path/filepath" - "strings" - "testing" -) - -func TestIsValidPRDName(t *testing.T) { - tests := []struct { - name string - input string - expected bool - }{ - {"valid lowercase", "main", true}, - {"valid with numbers", "feature1", true}, - {"valid with hyphen", "my-feature", true}, - {"valid with underscore", "my_feature", true}, - {"valid mixed case", "MyFeature", true}, - {"valid complex", "auth-v2_final", true}, - {"empty string", "", false}, - {"with space", "my feature", false}, - {"with dot", "my.feature", false}, - {"with slash", "my/feature", false}, - {"with special char", "my@feature", false}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - result := isValidPRDName(tt.input) - if result != tt.expected { - t.Errorf("isValidPRDName(%q) = %v, want %v", tt.input, result, tt.expected) - } - }) - } -} - -func TestRunNewCreatesDirectory(t *testing.T) { - // Create a temporary directory for testing - tmpDir := t.TempDir() - - // Test that directory structure is created correctly - // We can't fully test RunNew without Claude, but we can verify directory creation logic - name := "test-prd" - prdDir := filepath.Join(tmpDir, ".chief", "prds", name) - - // Simulate what RunNew does for directory creation - if err := os.MkdirAll(prdDir, 0755); err != nil { - t.Fatalf("Failed to create test directory: %v", err) - } - - // Verify directory was created at expected path - if _, err := os.Stat(prdDir); os.IsNotExist(err) { - t.Error("Expected directory to be created") - } - - // Verify parent directories also exist - chiefDir := filepath.Join(tmpDir, ".chief") - if _, err := os.Stat(chiefDir); os.IsNotExist(err) { - t.Error("Expected .chief directory to be created") - } - - prdsDir := filepath.Join(chiefDir, "prds") - if _, err := os.Stat(prdsDir); os.IsNotExist(err) { - t.Error("Expected .chief/prds directory to be created") - } -} - -func TestRunNewRejectsInvalidName(t *testing.T) { - tmpDir := t.TempDir() - - opts := NewOptions{ - Name: "invalid name with space", - BaseDir: tmpDir, - } - - err := RunNew(opts) - if err == nil { - t.Error("Expected error for invalid name") - } -} - -func TestRunNewCleansUpEmptyDirOnCancel(t *testing.T) { - tmpDir := t.TempDir() - - // Simulate what RunNew does: create directory, then check prd.md doesn't exist - name := "cancelled" - prdDir := filepath.Join(tmpDir, ".chief", "prds", name) - if err := os.MkdirAll(prdDir, 0755); err != nil { - t.Fatalf("Failed to create directory: %v", err) - } - - // prd.md was never created (user cancelled) — simulate the cleanup - prdMdPath := filepath.Join(prdDir, "prd.md") - if _, err := os.Stat(prdMdPath); os.IsNotExist(err) { - os.Remove(prdDir) - } - - // Directory should be removed - if _, err := os.Stat(prdDir); !os.IsNotExist(err) { - t.Error("Expected empty directory to be cleaned up after cancellation") - } -} - -func TestRunNewKeepsDirWhenPrdMdExists(t *testing.T) { - tmpDir := t.TempDir() - - name := "has-prd" - prdDir := filepath.Join(tmpDir, ".chief", "prds", name) - if err := os.MkdirAll(prdDir, 0755); err != nil { - t.Fatalf("Failed to create directory: %v", err) - } - - // Create prd.md (simulates successful Claude session) - prdMdPath := filepath.Join(prdDir, "prd.md") - if err := os.WriteFile(prdMdPath, []byte("# My PRD"), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) - } - - // Cleanup should NOT trigger since prd.md exists - if _, err := os.Stat(prdMdPath); os.IsNotExist(err) { - os.Remove(prdDir) - } - - // Directory should still exist - if _, err := os.Stat(prdDir); os.IsNotExist(err) { - t.Error("Expected directory to be kept when prd.md exists") - } -} - -func TestRunNewRejectsExistingPRD(t *testing.T) { - tmpDir := t.TempDir() - - // Create existing prd.md - prdDir := filepath.Join(tmpDir, ".chief", "prds", "existing") - if err := os.MkdirAll(prdDir, 0755); err != nil { - t.Fatalf("Failed to create directory: %v", err) - } - prdMdPath := filepath.Join(prdDir, "prd.md") - if err := os.WriteFile(prdMdPath, []byte("# Existing PRD"), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) - } - - opts := NewOptions{ - Name: "existing", - BaseDir: tmpDir, - } - - err := RunNew(opts) - if err == nil { - t.Error("Expected error for existing PRD") - } -} - -func TestRunNewRequiresProvider(t *testing.T) { - opts := NewOptions{ - Name: "main", - BaseDir: t.TempDir(), - } - - err := RunNew(opts) - if err == nil { - t.Fatal("expected provider validation error") - } - if !strings.Contains(err.Error(), "Provider") { - t.Fatalf("expected error to mention Provider, got: %v", err) - } -} diff --git a/internal/cmd/prds.go b/internal/cmd/prds.go new file mode 100644 index 00000000..f9345d01 --- /dev/null +++ b/internal/cmd/prds.go @@ -0,0 +1,64 @@ +package cmd + +import ( + "encoding/json" + "fmt" + "log" + "strings" + + "github.com/minicodemonkey/chief/internal/ws" +) + +// handleGetPRDs handles a get_prds request. +func handleGetPRDs(sender messageSender, finder projectFinder, msg ws.Message) { + var req ws.GetPRDsMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing get_prds message: %v", err) + return + } + + project, found := finder.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + items := make([]ws.PRDItem, 0, len(project.PRDs)) + for _, prd := range project.PRDs { + items = append(items, ws.PRDItem{ + ID: prd.ID, + Name: prd.Name, + StoryCount: prd.StoryCount, + Status: mapCompletionStatus(prd.CompletionStatus), + }) + } + + resp := ws.PRDsResponseMessage{ + Type: ws.TypePRDsResponse, + Payload: ws.PRDsResponsePayload{ + Project: req.Project, + PRDs: items, + }, + } + if err := sender.Send(resp); err != nil { + log.Printf("Error sending prds_response: %v", err) + } +} + +// mapCompletionStatus converts a "passed/total" completion status to a +// browser-friendly status string: "draft", "active", or "done". +func mapCompletionStatus(status string) string { + parts := strings.SplitN(status, "/", 2) + if len(parts) != 2 { + return "draft" + } + passed, total := parts[0], parts[1] + if total == "0" { + return "draft" + } + if passed == total { + return "done" + } + return "active" +} diff --git a/internal/cmd/prds_test.go b/internal/cmd/prds_test.go new file mode 100644 index 00000000..14d41261 --- /dev/null +++ b/internal/cmd/prds_test.go @@ -0,0 +1,260 @@ +package cmd + +import ( + "encoding/json" + "os" + "path/filepath" + "sync" + "testing" + "time" +) + +func TestRunServe_GetPRDs(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + // Create .chief/prds with two PRDs + prd1Dir := filepath.Join(projectDir, ".chief", "prds", "feature-auth") + prd2Dir := filepath.Join(projectDir, ".chief", "prds", "feature-dashboard") + if err := os.MkdirAll(prd1Dir, 0o755); err != nil { + t.Fatal(err) + } + if err := os.MkdirAll(prd2Dir, 0o755); err != nil { + t.Fatal(err) + } + + // PRD with 2/3 stories passing → "active" + prd1JSON := `{"project": "Auth System", "userStories": [ + {"id": "US-001", "passes": true}, + {"id": "US-002", "passes": true}, + {"id": "US-003", "passes": false} + ]}` + if err := os.WriteFile(filepath.Join(prd1Dir, "prd.json"), []byte(prd1JSON), 0o644); err != nil { + t.Fatal(err) + } + + // PRD with 2/2 stories passing → "done" + prd2JSON := `{"project": "Dashboard", "userStories": [ + {"id": "US-010", "passes": true}, + {"id": "US-011", "passes": true} + ]}` + if err := os.WriteFile(filepath.Join(prd2Dir, "prd.json"), []byte(prd2JSON), 0o644); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]interface{}{ + "type": "get_prds", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("prds_response", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("prds_response was not received") + } + if response["type"] != "prds_response" { + t.Errorf("expected type 'prds_response', got %v", response["type"]) + } + + payload, ok := response["payload"].(map[string]interface{}) + if !ok { + t.Fatal("expected payload to be an object") + } + if payload["project"] != "myproject" { + t.Errorf("expected project 'myproject', got %v", payload["project"]) + } + + prds, ok := payload["prds"].([]interface{}) + if !ok { + t.Fatal("expected prds to be an array") + } + if len(prds) != 2 { + t.Fatalf("expected 2 PRDs, got %d", len(prds)) + } + + // Build a map by ID for easier assertions + prdMap := make(map[string]map[string]interface{}) + for _, p := range prds { + prd := p.(map[string]interface{}) + prdMap[prd["id"].(string)] = prd + } + + // feature-auth: 2/3 passing → "active" + auth := prdMap["feature-auth"] + if auth == nil { + t.Fatal("expected feature-auth PRD") + } + if auth["name"] != "Auth System" { + t.Errorf("expected name 'Auth System', got %v", auth["name"]) + } + if int(auth["story_count"].(float64)) != 3 { + t.Errorf("expected story_count 3, got %v", auth["story_count"]) + } + if auth["status"] != "active" { + t.Errorf("expected status 'active', got %v", auth["status"]) + } + + // feature-dashboard: 2/2 passing → "done" + dash := prdMap["feature-dashboard"] + if dash == nil { + t.Fatal("expected feature-dashboard PRD") + } + if dash["status"] != "done" { + t.Errorf("expected status 'done', got %v", dash["status"]) + } +} + +func TestRunServe_GetPRDs_ProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]interface{}{ + "type": "get_prds", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("error message was not received") + } + if response["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", response["code"]) + } +} + +func TestRunServe_GetPRDs_EmptyProject(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Project with no .chief directory → empty PRD list + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var response map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]interface{}{ + "type": "get_prds", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("prds_response", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &response) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if response == nil { + t.Fatal("prds_response was not received") + } + + payload := response["payload"].(map[string]interface{}) + prds, ok := payload["prds"].([]interface{}) + if !ok { + t.Fatal("expected prds to be an array") + } + if len(prds) != 0 { + t.Errorf("expected 0 PRDs for project without .chief, got %d", len(prds)) + } +} + +func TestMapCompletionStatus(t *testing.T) { + tests := []struct { + input string + expected string + }{ + {"0/0", "draft"}, + {"0/5", "active"}, + {"3/5", "active"}, + {"5/5", "done"}, + {"", "draft"}, + {"invalid", "draft"}, + } + + for _, tc := range tests { + result := mapCompletionStatus(tc.input) + if result != tc.expected { + t.Errorf("mapCompletionStatus(%q) = %q, want %q", tc.input, result, tc.expected) + } + } +} diff --git a/internal/cmd/remote_update.go b/internal/cmd/remote_update.go new file mode 100644 index 00000000..d577e901 --- /dev/null +++ b/internal/cmd/remote_update.go @@ -0,0 +1,77 @@ +package cmd + +import ( + "fmt" + "log" + "strings" + + "github.com/minicodemonkey/chief/internal/update" + "github.com/minicodemonkey/chief/internal/ws" +) + +// handleTriggerUpdate handles a trigger_update request from the web app. +// It downloads and installs the latest binary, sends confirmation, +// and returns true if the process should exit (so systemd Restart=always picks up the new binary). +func handleTriggerUpdate(sender messageSender, msg ws.Message, version, releasesURL string) bool { + log.Println("Received trigger_update request") + + // Check for update + result, err := update.CheckForUpdate(version, update.Options{ + ReleasesURL: releasesURL, + }) + if err != nil { + sendError(sender, ws.ErrCodeUpdateFailed, + fmt.Sprintf("checking for updates: %v", err), msg.ID) + return false + } + + if !result.UpdateAvailable { + // Already on latest — send informational message (not an error) + envelope := ws.NewMessage(ws.TypeUpdateAvailable) + infoMsg := ws.UpdateAvailableMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + CurrentVersion: result.CurrentVersion, + LatestVersion: result.LatestVersion, + } + if err := sender.Send(infoMsg); err != nil { + log.Printf("Error sending update_available: %v", err) + } + log.Printf("Already on latest version (v%s)", result.CurrentVersion) + return false + } + + // Perform the update + log.Printf("Downloading v%s (current: v%s)...", result.LatestVersion, result.CurrentVersion) + _, err = update.PerformUpdate(version, update.Options{ + ReleasesURL: releasesURL, + }) + if err != nil { + errMsg := err.Error() + if strings.Contains(errMsg, "Permission denied") { + sendError(sender, ws.ErrCodeUpdateFailed, + "Permission denied. The chief binary is not writable. Ensure the service user has write permissions to the binary path.", msg.ID) + } else { + sendError(sender, ws.ErrCodeUpdateFailed, + fmt.Sprintf("update failed: %v", err), msg.ID) + } + return false + } + + // Send confirmation before exiting + log.Printf("Updated to v%s. Exiting for restart.", result.LatestVersion) + envelope := ws.NewMessage(ws.TypeUpdateAvailable) + confirmMsg := ws.UpdateAvailableMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + CurrentVersion: result.CurrentVersion, + LatestVersion: result.LatestVersion, + } + if err := sender.Send(confirmMsg); err != nil { + log.Printf("Error sending update confirmation: %v", err) + } + + return true +} diff --git a/internal/cmd/remote_update_test.go b/internal/cmd/remote_update_test.go new file mode 100644 index 00000000..126aa809 --- /dev/null +++ b/internal/cmd/remote_update_test.go @@ -0,0 +1,266 @@ +package cmd + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/update" + "github.com/minicodemonkey/chief/internal/ws" +) + +func TestHandleTriggerUpdate_AlreadyLatest(t *testing.T) { + // Mock GitHub releases API — same version + release := update.Release{TagName: "v1.0.0"} + releaseSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + json.NewEncoder(w).Encode(release) + })) + defer releaseSrv.Close() + + sender := &captureSender{} + + msg := ws.Message{ + Type: ws.TypeTriggerUpdate, + ID: "req-1", + } + + shouldExit := handleTriggerUpdate(sender, msg, "1.0.0", releaseSrv.URL) + + if shouldExit { + t.Error("should not exit when already on latest version") + } + + msgs := sender.getMessages() + if len(msgs) == 0 { + t.Fatal("expected update_available message to be sent") + } + + var receivedMsg map[string]interface{} + for _, m := range msgs { + if m["type"] == "update_available" { + receivedMsg = m + break + } + } + + if receivedMsg == nil { + t.Fatal("expected update_available message to be sent") + } + if receivedMsg["current_version"] != "1.0.0" { + t.Errorf("expected current_version '1.0.0', got %v", receivedMsg["current_version"]) + } + if receivedMsg["latest_version"] != "1.0.0" { + t.Errorf("expected latest_version '1.0.0', got %v", receivedMsg["latest_version"]) + } +} + +func TestHandleTriggerUpdate_APIError(t *testing.T) { + // Mock GitHub releases API — error + releaseSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer releaseSrv.Close() + + sender := &captureSender{} + + msg := ws.Message{ + Type: ws.TypeTriggerUpdate, + ID: "req-1", + } + + shouldExit := handleTriggerUpdate(sender, msg, "1.0.0", releaseSrv.URL) + + if shouldExit { + t.Error("should not exit on API error") + } + + msgs := sender.getMessages() + if len(msgs) == 0 { + t.Fatal("expected error message to be sent") + } + + var receivedMsg map[string]interface{} + for _, m := range msgs { + if m["type"] == "error" { + receivedMsg = m + break + } + } + + if receivedMsg == nil { + t.Fatal("expected error message to be sent") + } + if receivedMsg["code"] != "UPDATE_FAILED" { + t.Errorf("expected code 'UPDATE_FAILED', got %v", receivedMsg["code"]) + } + if receivedMsg["request_id"] != "req-1" { + t.Errorf("expected request_id 'req-1', got %v", receivedMsg["request_id"]) + } +} + +func TestRunServe_TriggerUpdateAlreadyLatest(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Mock releases API — same version + release := update.Release{TagName: "v1.0.0"} + releaseSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + json.NewEncoder(w).Encode(release) + })) + defer releaseSrv.Close() + + var responseReceived map[string]interface{} + var mu sync.Mutex + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for initial state_snapshot + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send trigger_update command via Pusher + triggerReq := map[string]string{ + "type": "trigger_update", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + } + ms.sendCommand(triggerReq) + + // Wait for update_available response + raw, err := ms.waitForMessageType("update_available", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &responseReceived) + mu.Unlock() + } + + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspaceDir, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + ReleasesURL: releaseSrv.URL, + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if responseReceived == nil { + t.Fatal("expected response message") + } + if responseReceived["type"] != "update_available" { + t.Errorf("expected type 'update_available', got %v", responseReceived["type"]) + } + if responseReceived["current_version"] != "1.0.0" { + t.Errorf("expected current_version '1.0.0', got %v", responseReceived["current_version"]) + } +} + +func TestRunServe_TriggerUpdateAPIError(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Mock releases API — error + releaseSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer releaseSrv.Close() + + var errorReceived map[string]interface{} + var mu sync.Mutex + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for initial state_snapshot + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send trigger_update command via Pusher + triggerReq := map[string]string{ + "type": "trigger_update", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + } + ms.sendCommand(triggerReq) + + // Wait for error response + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspaceDir, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + ReleasesURL: releaseSrv.URL, + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("expected error message") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "UPDATE_FAILED" { + t.Errorf("expected code 'UPDATE_FAILED', got %v", errorReceived["code"]) + } +} diff --git a/internal/cmd/runs.go b/internal/cmd/runs.go new file mode 100644 index 00000000..08594cc3 --- /dev/null +++ b/internal/cmd/runs.go @@ -0,0 +1,632 @@ +package cmd + +import ( + "context" + "encoding/json" + "fmt" + "log" + "path/filepath" + "sync" + "time" + + "github.com/minicodemonkey/chief/internal/engine" + "github.com/minicodemonkey/chief/internal/loop" + "github.com/minicodemonkey/chief/internal/prd" + "github.com/minicodemonkey/chief/internal/ws" +) + +// runManager manages Ralph loop runs driven by WebSocket commands. +type runManager struct { + mu sync.RWMutex + eng *engine.Engine + sender messageSender + // tracks which engine registration key maps to which project/prd + runs map[string]*runInfo + loggers map[string]*storyLogger +} + +// runInfo tracks metadata about a registered run. +type runInfo struct { + project string + prdID string + prdPath string // absolute path to prd.json + startTime time.Time + storyID string // currently active story ID +} + +// runKey returns the engine registration key for a project/PRD combination. +func runKey(project, prdID string) string { + return project + "/" + prdID +} + +// newRunManager creates a new run manager. +func newRunManager(eng *engine.Engine, sender messageSender) *runManager { + return &runManager{ + eng: eng, + sender: sender, + runs: make(map[string]*runInfo), + loggers: make(map[string]*storyLogger), + } +} + +// startEventMonitor subscribes to engine events and handles progress streaming, +// run completion, Claude output streaming, and quota exhaustion. +// It runs until the context is cancelled. +func (rm *runManager) startEventMonitor(ctx context.Context) { + eventCh, unsub := rm.eng.Subscribe() + go func() { + defer unsub() + for { + select { + case <-ctx.Done(): + return + case event, ok := <-eventCh: + if !ok { + return + } + rm.handleEvent(event) + } + } + }() +} + +// handleEvent routes an engine event to the appropriate handler. +func (rm *runManager) handleEvent(event engine.ManagerEvent) { + rm.mu.RLock() + info, exists := rm.runs[event.PRDName] + rm.mu.RUnlock() + + if !exists { + // Events from runs we don't track (e.g., TUI-driven runs) + return + } + + switch event.Event.Type { + case loop.EventQuotaExhausted: + rm.handleQuotaExhausted(event.PRDName) + + case loop.EventIterationStart: + rm.sendRunProgress(info, "iteration_started", event.Event) + + case loop.EventStoryStarted: + // Track the current story ID + rm.mu.Lock() + info.storyID = event.Event.StoryID + rm.mu.Unlock() + rm.sendRunProgress(info, "story_started", event.Event) + + case loop.EventStoryCompleted: + rm.sendRunProgress(info, "story_completed", event.Event) + rm.sendStoryDiff(info, event.Event) + + case loop.EventComplete: + rm.sendRunProgress(info, "complete", event.Event) + rm.sendRunComplete(info, event.PRDName) + + case loop.EventMaxIterationsReached: + rm.sendRunProgress(info, "max_iterations_reached", event.Event) + rm.sendRunComplete(info, event.PRDName) + + case loop.EventRetrying: + rm.sendRunProgress(info, "retrying", event.Event) + + case loop.EventAssistantText: + rm.writeStoryLog(event.PRDName, info.storyID, event.Event.Text) + rm.sendClaudeOutput(info, event.Event.Text, false) + + case loop.EventToolStart: + text := fmt.Sprintf("[tool_use] %s", event.Event.Tool) + rm.writeStoryLog(event.PRDName, info.storyID, text) + rm.sendClaudeOutput(info, text, false) + + case loop.EventToolResult: + rm.writeStoryLog(event.PRDName, info.storyID, event.Event.Text) + rm.sendClaudeOutput(info, event.Event.Text, false) + + case loop.EventError: + errText := "" + if event.Event.Err != nil { + errText = event.Event.Err.Error() + } + rm.writeStoryLog(event.PRDName, info.storyID, errText) + rm.sendClaudeOutput(info, errText, true) + } +} + +// sendRunProgress sends a run_progress message over WebSocket. +func (rm *runManager) sendRunProgress(info *runInfo, status string, event loop.Event) { + if rm.sender == nil { + return + } + + rm.mu.RLock() + storyID := info.storyID + rm.mu.RUnlock() + + // Use the event's story ID if available, otherwise use tracked story ID + if event.StoryID != "" { + storyID = event.StoryID + } + + envelope := ws.NewMessage(ws.TypeRunProgress) + msg := ws.RunProgressMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: info.project, + PRDID: info.prdID, + StoryID: storyID, + Status: status, + Iteration: event.Iteration, + Attempt: event.RetryCount, + } + if err := rm.sender.Send(msg); err != nil { + log.Printf("Error sending run_progress: %v", err) + } +} + +// sendRunComplete sends a run_complete message over WebSocket. +func (rm *runManager) sendRunComplete(info *runInfo, prdName string) { + if rm.sender == nil { + return + } + + // Calculate duration + rm.mu.RLock() + duration := time.Since(info.startTime) + rm.mu.RUnlock() + + // Load PRD to get pass/fail counts + var passCount, failCount, storiesCompleted int + p, err := prd.LoadPRD(info.prdPath) + if err == nil { + for _, s := range p.UserStories { + if s.Passes { + passCount++ + storiesCompleted++ + } else { + failCount++ + } + } + } + + envelope := ws.NewMessage(ws.TypeRunComplete) + msg := ws.RunCompleteMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: info.project, + PRDID: info.prdID, + StoriesCompleted: storiesCompleted, + Duration: duration.Round(time.Second).String(), + PassCount: passCount, + FailCount: failCount, + } + if err := rm.sender.Send(msg); err != nil { + log.Printf("Error sending run_complete: %v", err) + } +} + +// sendClaudeOutput sends a claude_output message for an active run over WebSocket. +func (rm *runManager) sendClaudeOutput(info *runInfo, data string, done bool) { + if rm.sender == nil { + return + } + + rm.mu.RLock() + storyID := info.storyID + rm.mu.RUnlock() + + envelope := ws.NewMessage(ws.TypeClaudeOutput) + msg := ws.ClaudeOutputMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: info.project, + PRDID: info.prdID, + StoryID: storyID, + Data: data, + Done: done, + } + if err := rm.sender.Send(msg); err != nil { + log.Printf("Error sending claude_output: %v", err) + } +} + +// sendStoryDiff sends a proactive diff message when a story completes during a run. +func (rm *runManager) sendStoryDiff(info *runInfo, event loop.Event) { + if rm.sender == nil { + return + } + + storyID := event.StoryID + if storyID == "" { + rm.mu.RLock() + storyID = info.storyID + rm.mu.RUnlock() + } + if storyID == "" { + return + } + + // Get the project path from the PRD path + // prdPath is like /path/to/project/.chief/prds/<id>/prd.json + projectPath := filepath.Dir(filepath.Dir(filepath.Dir(filepath.Dir(info.prdPath)))) + + diffText, files, err := getStoryDiff(projectPath, storyID) + if err != nil { + log.Printf("Could not get diff for story %s: %v", storyID, err) + return + } + + sendDiffMessage(rm.sender, info.project, info.prdID, storyID, files, diffText) +} + +// handleQuotaExhausted handles a quota exhaustion event for a specific run. +func (rm *runManager) handleQuotaExhausted(prdName string) { + rm.mu.RLock() + info, exists := rm.runs[prdName] + rm.mu.RUnlock() + + if !exists { + log.Printf("Quota exhausted for unknown run key: %s", prdName) + return + } + + log.Printf("Quota exhausted for %s/%s, auto-pausing", info.project, info.prdID) + + if rm.sender == nil { + return + } + + // Send run_paused with reason quota_exhausted + envelope := ws.NewMessage(ws.TypeRunPaused) + pausedMsg := ws.RunPausedMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: info.project, + PRDID: info.prdID, + Reason: "quota_exhausted", + } + if err := rm.sender.Send(pausedMsg); err != nil { + log.Printf("Error sending run_paused: %v", err) + } + + // Send quota_exhausted message listing affected runs + rm.sendQuotaExhausted(info.project, info.prdID) +} + +// sendQuotaExhausted sends a quota_exhausted message over WebSocket. +func (rm *runManager) sendQuotaExhausted(project, prdID string) { + if rm.sender == nil { + return + } + envelope := ws.NewMessage(ws.TypeQuotaExhausted) + msg := ws.QuotaExhaustedMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Runs: []string{runKey(project, prdID)}, + Sessions: []string{}, + } + if err := rm.sender.Send(msg); err != nil { + log.Printf("Error sending quota_exhausted: %v", err) + } +} + +// activeRuns returns the list of active runs for state snapshots. +func (rm *runManager) activeRuns() []ws.RunState { + rm.mu.RLock() + defer rm.mu.RUnlock() + + var runs []ws.RunState + for key := range rm.runs { + info := rm.runs[key] + instance := rm.eng.GetInstance(key) + if instance == nil { + continue + } + + status := loopStateToString(instance.State) + runs = append(runs, ws.RunState{ + Project: info.project, + PRDID: info.prdID, + Status: status, + Iteration: instance.Iteration, + }) + } + return runs +} + +// startRun starts a Ralph loop for a project/PRD. +func (rm *runManager) startRun(project, prdID, projectPath string) error { + key := runKey(project, prdID) + + // Check if already running + if instance := rm.eng.GetInstance(key); instance != nil { + if instance.State == loop.LoopStateRunning { + return fmt.Errorf("RUN_ALREADY_ACTIVE") + } + } + + prdPath := filepath.Join(projectPath, ".chief", "prds", prdID, "prd.json") + + // Register if not already registered + if instance := rm.eng.GetInstance(key); instance == nil { + if err := rm.eng.Register(key, prdPath); err != nil { + return fmt.Errorf("failed to register PRD: %w", err) + } + } + + // Create per-story logger (removes previous logs for this PRD) + sl, err := newStoryLogger(prdPath) + if err != nil { + log.Printf("Warning: could not create story logger: %v", err) + } + + rm.mu.Lock() + rm.runs[key] = &runInfo{ + project: project, + prdID: prdID, + prdPath: prdPath, + startTime: time.Now(), + } + if sl != nil { + rm.loggers[key] = sl + } + rm.mu.Unlock() + + if err := rm.eng.Start(key); err != nil { + return fmt.Errorf("failed to start run: %w", err) + } + + return nil +} + +// pauseRun pauses a running loop. +func (rm *runManager) pauseRun(project, prdID string) error { + key := runKey(project, prdID) + + instance := rm.eng.GetInstance(key) + if instance == nil || instance.State != loop.LoopStateRunning { + return fmt.Errorf("RUN_NOT_ACTIVE") + } + + if err := rm.eng.Pause(key); err != nil { + return fmt.Errorf("failed to pause run: %w", err) + } + + return nil +} + +// resumeRun resumes a paused loop by starting it again. +func (rm *runManager) resumeRun(project, prdID string) error { + key := runKey(project, prdID) + + instance := rm.eng.GetInstance(key) + if instance == nil || instance.State != loop.LoopStatePaused { + return fmt.Errorf("RUN_NOT_ACTIVE") + } + + // Start creates a fresh Loop that picks up from the next unfinished story + if err := rm.eng.Start(key); err != nil { + return fmt.Errorf("failed to resume run: %w", err) + } + + return nil +} + +// stopRun stops a running or paused loop immediately. +func (rm *runManager) stopRun(project, prdID string) error { + key := runKey(project, prdID) + + instance := rm.eng.GetInstance(key) + if instance == nil || (instance.State != loop.LoopStateRunning && instance.State != loop.LoopStatePaused) { + return fmt.Errorf("RUN_NOT_ACTIVE") + } + + if err := rm.eng.Stop(key); err != nil { + return fmt.Errorf("failed to stop run: %w", err) + } + + return nil +} + +// writeStoryLog writes a line to the per-story log file. +func (rm *runManager) writeStoryLog(runKey, storyID, text string) { + rm.mu.RLock() + sl := rm.loggers[runKey] + rm.mu.RUnlock() + + if sl != nil { + sl.WriteLog(storyID, text) + } +} + +// cleanup removes tracking for a completed/stopped run. +func (rm *runManager) cleanup(key string) { + rm.mu.Lock() + if sl, ok := rm.loggers[key]; ok { + sl.Close() + delete(rm.loggers, key) + } + delete(rm.runs, key) + rm.mu.Unlock() +} + +// markInterruptedStories marks any in-progress stories as interrupted in prd.json +// so that the next run resumes from where it left off. +func (rm *runManager) markInterruptedStories() { + rm.mu.RLock() + runs := make([]*runInfo, 0, len(rm.runs)) + for _, info := range rm.runs { + runs = append(runs, info) + } + rm.mu.RUnlock() + + for _, info := range runs { + if info.storyID == "" { + continue + } + p, err := prd.LoadPRD(info.prdPath) + if err != nil { + log.Printf("Warning: could not load PRD %s to mark interrupted story: %v", info.prdPath, err) + continue + } + for i := range p.UserStories { + if p.UserStories[i].ID == info.storyID && !p.UserStories[i].Passes { + p.UserStories[i].InProgress = true + if err := p.Save(info.prdPath); err != nil { + log.Printf("Warning: could not save PRD %s: %v", info.prdPath, err) + } + break + } + } + } +} + +// activeRunCount returns the number of currently tracked runs. +func (rm *runManager) activeRunCount() int { + rm.mu.RLock() + defer rm.mu.RUnlock() + return len(rm.runs) +} + +// stopAll stops all active runs (for shutdown). +func (rm *runManager) stopAll() { + rm.eng.StopAll() + + // Close all story loggers + rm.mu.Lock() + for key, sl := range rm.loggers { + sl.Close() + delete(rm.loggers, key) + } + rm.mu.Unlock() +} + +// loopStateToString converts a LoopState to a string for WebSocket messages. +func loopStateToString(state loop.LoopState) string { + switch state { + case loop.LoopStateReady: + return "ready" + case loop.LoopStateRunning: + return "running" + case loop.LoopStatePaused: + return "paused" + case loop.LoopStateStopped: + return "stopped" + case loop.LoopStateComplete: + return "complete" + case loop.LoopStateError: + return "error" + default: + return "unknown" + } +} + +// handleStartRun handles a start_run WebSocket message. +func handleStartRun(sender messageSender, scanner projectFinder, runs *runManager, watcher activator, msg ws.Message) { + var req ws.StartRunMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing start_run message: %v", err) + return + } + + project, found := scanner.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + if err := runs.startRun(req.Project, req.PRDID, project.Path); err != nil { + if err.Error() == "RUN_ALREADY_ACTIVE" { + sendError(sender, ws.ErrCodeRunAlreadyActive, + fmt.Sprintf("Run already active for %s/%s", req.Project, req.PRDID), msg.ID) + } else { + sendError(sender, ws.ErrCodeClaudeError, + fmt.Sprintf("Failed to start run: %v", err), msg.ID) + } + return + } + + // Activate file watching for the project + if watcher != nil { + watcher.Activate(req.Project) + } + + log.Printf("Started run for %s/%s", req.Project, req.PRDID) +} + +// handlePauseRun handles a pause_run WebSocket message. +func handlePauseRun(sender messageSender, runs *runManager, msg ws.Message) { + var req ws.PauseRunMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing pause_run message: %v", err) + return + } + + if err := runs.pauseRun(req.Project, req.PRDID); err != nil { + if err.Error() == "RUN_NOT_ACTIVE" { + sendError(sender, ws.ErrCodeRunNotActive, + fmt.Sprintf("No active run for %s/%s", req.Project, req.PRDID), msg.ID) + } else { + sendError(sender, ws.ErrCodeClaudeError, + fmt.Sprintf("Failed to pause run: %v", err), msg.ID) + } + return + } + + log.Printf("Paused run for %s/%s", req.Project, req.PRDID) +} + +// handleResumeRun handles a resume_run WebSocket message. +func handleResumeRun(sender messageSender, runs *runManager, msg ws.Message) { + var req ws.ResumeRunMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing resume_run message: %v", err) + return + } + + if err := runs.resumeRun(req.Project, req.PRDID); err != nil { + if err.Error() == "RUN_NOT_ACTIVE" { + sendError(sender, ws.ErrCodeRunNotActive, + fmt.Sprintf("No paused run for %s/%s", req.Project, req.PRDID), msg.ID) + } else { + sendError(sender, ws.ErrCodeClaudeError, + fmt.Sprintf("Failed to resume run: %v", err), msg.ID) + } + return + } + + log.Printf("Resumed run for %s/%s", req.Project, req.PRDID) +} + +// handleStopRun handles a stop_run WebSocket message. +func handleStopRun(sender messageSender, runs *runManager, msg ws.Message) { + var req ws.StopRunMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing stop_run message: %v", err) + return + } + + if err := runs.stopRun(req.Project, req.PRDID); err != nil { + if err.Error() == "RUN_NOT_ACTIVE" { + sendError(sender, ws.ErrCodeRunNotActive, + fmt.Sprintf("No active run for %s/%s", req.Project, req.PRDID), msg.ID) + } else { + sendError(sender, ws.ErrCodeClaudeError, + fmt.Sprintf("Failed to stop run: %v", err), msg.ID) + } + return + } + + log.Printf("Stopped run for %s/%s", req.Project, req.PRDID) +} + +// activator is an interface for activating file watching (for testability). +type activator interface { + Activate(name string) +} diff --git a/internal/cmd/runs_test.go b/internal/cmd/runs_test.go new file mode 100644 index 00000000..f7dabe05 --- /dev/null +++ b/internal/cmd/runs_test.go @@ -0,0 +1,1001 @@ +package cmd + +import ( + "context" + "encoding/json" + "os" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/engine" + "github.com/minicodemonkey/chief/internal/loop" + "github.com/minicodemonkey/chief/internal/ws" +) + +func TestRunServe_StartRun(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo with a PRD + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + // Write a minimal prd.json with one story + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test Story", "passes": false}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + var responseReceived map[string]interface{} + var mu sync.Mutex + gotError := false + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send start_run request + startReq := map[string]string{ + "type": "start_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(startReq) + + // Wait a moment for the run to start, then check if error was returned + raw, err := ms.waitForMessages(1, 2*time.Second) + if err == nil && len(raw) > 0 { + mu.Lock() + json.Unmarshal(raw[0], &responseReceived) + // If it's an error, it means the run couldn't start (expected in test env without claude) + if responseReceived["type"] == "error" { + gotError = true + } + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // In a test environment without a real claude binary, the engine.Start() call + // will succeed (registers + starts the loop) but the loop itself will fail + // quickly since there's no claude. We verify the handler routed correctly + // by checking that we didn't get a PROJECT_NOT_FOUND error. + mu.Lock() + defer mu.Unlock() + + if responseReceived != nil && gotError { + // If we got an error, it should NOT be PROJECT_NOT_FOUND + code, _ := responseReceived["code"].(string) + if code == "PROJECT_NOT_FOUND" { + t.Errorf("should not have gotten PROJECT_NOT_FOUND for existing project") + } + } +} + +func TestRunServe_StartRunProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + startReq := map[string]string{ + "type": "start_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "prd_id": "feature", + } + ms.sendCommand(startReq) + + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", errorReceived["code"]) + } +} + +func TestRunServe_PauseRunNotActive(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + pauseReq := map[string]string{ + "type": "pause_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(pauseReq) + + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "RUN_NOT_ACTIVE" { + t.Errorf("expected code 'RUN_NOT_ACTIVE', got %v", errorReceived["code"]) + } +} + +func TestRunServe_ResumeRunNotActive(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + resumeReq := map[string]string{ + "type": "resume_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(resumeReq) + + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "RUN_NOT_ACTIVE" { + t.Errorf("expected code 'RUN_NOT_ACTIVE', got %v", errorReceived["code"]) + } +} + +func TestRunServe_StopRunNotActive(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + stopReq := map[string]string{ + "type": "stop_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(stopReq) + + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "RUN_NOT_ACTIVE" { + t.Errorf("expected code 'RUN_NOT_ACTIVE', got %v", errorReceived["code"]) + } +} + +func TestRunManager_StartAndAlreadyActive(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Create a temp project with a PRD + projectDir := t.TempDir() + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + // Write a minimal prd.json + prdState := `{"project": "Test", "userStories": [{"id": "US-001", "title": "Story", "passes": false}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Start a run + err := rm.startRun("myproject", "feature", projectDir) + if err != nil { + t.Fatalf("startRun failed: %v", err) + } + + // Wait briefly for the engine to register the run as running + time.Sleep(100 * time.Millisecond) + + // Try to start the same run again — should get RUN_ALREADY_ACTIVE + err = rm.startRun("myproject", "feature", projectDir) + if err == nil { + t.Fatal("expected error for already active run") + } + if err.Error() != "RUN_ALREADY_ACTIVE" { + t.Errorf("expected 'RUN_ALREADY_ACTIVE', got: %v", err) + } + + // Clean up + rm.stopAll() +} + +func TestRunManager_PauseAndResume(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Trying to pause when nothing is running + err := rm.pauseRun("myproject", "feature") + if err == nil || err.Error() != "RUN_NOT_ACTIVE" { + t.Errorf("expected RUN_NOT_ACTIVE, got: %v", err) + } + + // Trying to resume when nothing is paused + err = rm.resumeRun("myproject", "feature") + if err == nil || err.Error() != "RUN_NOT_ACTIVE" { + t.Errorf("expected RUN_NOT_ACTIVE, got: %v", err) + } +} + +func TestRunManager_StopNotActive(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + err := rm.stopRun("myproject", "feature") + if err == nil || err.Error() != "RUN_NOT_ACTIVE" { + t.Errorf("expected RUN_NOT_ACTIVE, got: %v", err) + } +} + +func TestRunManager_ActiveRuns(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // No active runs initially + runs := rm.activeRuns() + if runs != nil && len(runs) != 0 { + t.Errorf("expected no active runs, got %d", len(runs)) + } + + // Create a temp project with a PRD + projectDir := t.TempDir() + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "Test", "userStories": [{"id": "US-001", "title": "Story", "passes": false}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Start a run + if err := rm.startRun("myproject", "feature", projectDir); err != nil { + t.Fatalf("startRun failed: %v", err) + } + + // Wait briefly for the engine to start + time.Sleep(100 * time.Millisecond) + + // Should have one active run + runs = rm.activeRuns() + if len(runs) != 1 { + t.Fatalf("expected 1 active run, got %d", len(runs)) + } + + if runs[0].Project != "myproject" { + t.Errorf("expected project 'myproject', got %q", runs[0].Project) + } + if runs[0].PRDID != "feature" { + t.Errorf("expected prd_id 'feature', got %q", runs[0].PRDID) + } + + rm.stopAll() +} + +func TestRunManager_MultipleConcurrentProjects(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Create two projects with PRDs + for _, name := range []string{"project-a", "project-b"} { + projectDir := filepath.Join(t.TempDir(), name) + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "Test", "userStories": [{"id": "US-001", "title": "Story", "passes": false}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + if err := rm.startRun(name, "feature", projectDir); err != nil { + t.Fatalf("startRun %s failed: %v", name, err) + } + } + + // Wait briefly + time.Sleep(100 * time.Millisecond) + + // Should have two active runs + runs := rm.activeRuns() + if len(runs) != 2 { + t.Errorf("expected 2 active runs, got %d", len(runs)) + } + + rm.stopAll() +} + +func TestRunManager_LoopStateToString(t *testing.T) { + tests := []struct { + state ws.RunState + expected string + }{ + {ws.RunState{Status: "running"}, "running"}, + {ws.RunState{Status: "paused"}, "paused"}, + {ws.RunState{Status: "stopped"}, "stopped"}, + } + + for _, tt := range tests { + if tt.state.Status != tt.expected { + t.Errorf("expected %q, got %q", tt.expected, tt.state.Status) + } + } +} + +func TestRunManager_HandleQuotaExhausted(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + // Create a mock WS client to capture sent messages + rm := newRunManager(eng, nil) + + // Add a run to the runs map + rm.mu.Lock() + rm.runs["myproject/feature"] = &runInfo{ + project: "myproject", + prdID: "feature", + prdPath: "/tmp/test/prd.json", + } + rm.mu.Unlock() + + // handleQuotaExhausted should not panic even with nil client + // (it logs errors but continues) + rm.handleQuotaExhausted("myproject/feature") +} + +func TestRunManager_HandleQuotaExhaustedUnknownRun(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Should not panic for unknown run + rm.handleQuotaExhausted("unknown/run") +} + +func TestRunManager_EventMonitorQuotaDetection(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Set up run tracking + rm.mu.Lock() + rm.runs["test/feature"] = &runInfo{ + project: "test", + prdID: "feature", + prdPath: "/tmp/test/prd.json", + } + rm.mu.Unlock() + + // Start the event monitor + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + rm.startEventMonitor(ctx) + + // Give the goroutine time to start + time.Sleep(50 * time.Millisecond) + + // Cancel context to stop the monitor + cancel() + time.Sleep(50 * time.Millisecond) +} + +func TestRunManager_QuotaExhaustedWebSocket(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo with a PRD that uses a mock claude that simulates quota error + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test Story", "passes": false}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Create a mock claude that outputs a quota error on stderr and exits with non-zero + mockDir := t.TempDir() + mockScript := `#!/bin/sh +echo "rate limit exceeded" >&2 +exit 1 +` + if err := os.WriteFile(filepath.Join(mockDir, "claude"), []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", mockDir+":"+origPath) + + var messages []map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send start_run request + startReq := map[string]string{ + "type": "start_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(startReq) + + // Read messages — we expect run_paused with reason "quota_exhausted" + // and a quota_exhausted message + raws, err := ms.waitForMessages(5, 5*time.Second) + if err == nil { + for _, data := range raws { + var msg map[string]interface{} + if json.Unmarshal(data, &msg) == nil { + mu.Lock() + messages = append(messages, msg) + mu.Unlock() + } + } + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + // Check that we received a run_paused with reason quota_exhausted + foundRunPaused := false + foundQuotaExhausted := false + for _, msg := range messages { + if msg["type"] == "run_paused" { + if msg["reason"] == "quota_exhausted" { + foundRunPaused = true + } + } + if msg["type"] == "quota_exhausted" { + foundQuotaExhausted = true + } + } + + if !foundRunPaused { + t.Errorf("expected run_paused with reason quota_exhausted, got messages: %v", messages) + } + if !foundQuotaExhausted { + t.Errorf("expected quota_exhausted message, got messages: %v", messages) + } +} + +func TestIsQuotaErrorIntegration(t *testing.T) { + // Test that the loop package's IsQuotaError function correctly identifies quota errors + tests := []struct { + stderr string + expected bool + }{ + {"Error: rate limit exceeded for model", true}, + {"HTTP 429 Too Many Requests", true}, + {"quota has been exceeded", true}, + {"normal crash: segfault", false}, + } + + for _, tt := range tests { + got := loop.IsQuotaError(tt.stderr) + if got != tt.expected { + t.Errorf("IsQuotaError(%q) = %v, want %v", tt.stderr, got, tt.expected) + } + } +} + +func TestRunManager_HandleEventRunProgress(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) // nil client — sendRunProgress guards against nil + + // Add a run to the runs map + rm.mu.Lock() + rm.runs["myproject/feature"] = &runInfo{ + project: "myproject", + prdID: "feature", + prdPath: "/tmp/test/prd.json", + startTime: time.Now(), + } + rm.mu.Unlock() + + // Test that handleEvent does not panic for each event type with nil client + eventTypes := []loop.EventType{ + loop.EventIterationStart, + loop.EventStoryStarted, + loop.EventStoryCompleted, + loop.EventComplete, + loop.EventMaxIterationsReached, + loop.EventRetrying, + loop.EventAssistantText, + loop.EventToolStart, + loop.EventToolResult, + loop.EventError, + } + + for _, et := range eventTypes { + event := engine.ManagerEvent{ + PRDName: "myproject/feature", + Event: loop.Event{ + Type: et, + Iteration: 1, + StoryID: "US-001", + Text: "test text", + Tool: "TestTool", + }, + } + rm.handleEvent(event) // should not panic with nil client + } +} + +func TestRunManager_HandleEventUnknownRun(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Events for unknown runs should be silently ignored + event := engine.ManagerEvent{ + PRDName: "unknown/run", + Event: loop.Event{ + Type: loop.EventIterationStart, + Iteration: 1, + }, + } + rm.handleEvent(event) // should not panic +} + +func TestRunManager_HandleEventStoryTracking(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + rm.mu.Lock() + rm.runs["myproject/feature"] = &runInfo{ + project: "myproject", + prdID: "feature", + prdPath: "/tmp/test/prd.json", + startTime: time.Now(), + } + rm.mu.Unlock() + + // Send a StoryStarted event — should update the tracked storyID + event := engine.ManagerEvent{ + PRDName: "myproject/feature", + Event: loop.Event{ + Type: loop.EventStoryStarted, + StoryID: "US-042", + }, + } + rm.handleEvent(event) + + rm.mu.RLock() + storyID := rm.runs["myproject/feature"].storyID + rm.mu.RUnlock() + + if storyID != "US-042" { + t.Errorf("expected storyID 'US-042', got %q", storyID) + } +} + +func TestRunManager_SendRunComplete(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) // nil client — guards against nil + + // Create a temp PRD with known pass/fail counts + tmpDir := t.TempDir() + prdDir := filepath.Join(tmpDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdJSON := `{"project": "Test", "userStories": [{"id": "US-001", "passes": true}, {"id": "US-002", "passes": true}, {"id": "US-003", "passes": false}]}` + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(prdJSON), 0o644); err != nil { + t.Fatal(err) + } + + info := &runInfo{ + project: "myproject", + prdID: "feature", + prdPath: prdPath, + startTime: time.Now().Add(-5 * time.Minute), + } + + // Should not panic with nil client + rm.sendRunComplete(info, "myproject/feature") +} + +func TestRunManager_MarkInterruptedStories(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Create a temp project with a PRD + projectDir := t.TempDir() + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdPath := filepath.Join(prdDir, "prd.json") + prdState := `{"project": "Test", "userStories": [{"id": "US-001", "title": "Story 1", "passes": false}, {"id": "US-002", "title": "Story 2", "passes": true}]}` + if err := os.WriteFile(prdPath, []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Add a run with an active story + rm.mu.Lock() + rm.runs["test/feature"] = &runInfo{ + project: "test", + prdID: "feature", + prdPath: prdPath, + startTime: time.Now(), + storyID: "US-001", + } + rm.mu.Unlock() + + // Mark interrupted stories + rm.markInterruptedStories() + + // Verify the PRD was updated + data, err := os.ReadFile(prdPath) + if err != nil { + t.Fatalf("failed to read PRD: %v", err) + } + + var result map[string]interface{} + if err := json.Unmarshal(data, &result); err != nil { + t.Fatalf("failed to parse PRD: %v", err) + } + + stories := result["userStories"].([]interface{}) + story1 := stories[0].(map[string]interface{}) + if story1["inProgress"] != true { + t.Errorf("expected US-001 to have inProgress=true, got %v", story1["inProgress"]) + } + + // US-002 is already passing, should NOT be marked as inProgress + story2 := stories[1].(map[string]interface{}) + if _, hasInProgress := story2["inProgress"]; hasInProgress && story2["inProgress"] == true { + t.Error("expected US-002 to NOT have inProgress=true (already passes)") + } +} + +func TestRunManager_MarkInterruptedStoriesNoStoryID(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + // Create a temp project with a PRD + projectDir := t.TempDir() + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdPath := filepath.Join(prdDir, "prd.json") + prdState := `{"project": "Test", "userStories": [{"id": "US-001", "title": "Story 1", "passes": false}]}` + if err := os.WriteFile(prdPath, []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Add a run WITHOUT a story ID (no story started yet) + rm.mu.Lock() + rm.runs["test/feature"] = &runInfo{ + project: "test", + prdID: "feature", + prdPath: prdPath, + startTime: time.Now(), + storyID: "", // no story started + } + rm.mu.Unlock() + + // Mark interrupted stories — should be a no-op + rm.markInterruptedStories() + + // Verify the PRD was NOT modified + data, err := os.ReadFile(prdPath) + if err != nil { + t.Fatalf("failed to read PRD: %v", err) + } + + var result map[string]interface{} + if err := json.Unmarshal(data, &result); err != nil { + t.Fatalf("failed to parse PRD: %v", err) + } + + stories := result["userStories"].([]interface{}) + story1 := stories[0].(map[string]interface{}) + if _, hasInProgress := story1["inProgress"]; hasInProgress && story1["inProgress"] == true { + t.Error("expected US-001 to NOT have inProgress=true when no story was started") + } +} + +func TestRunManager_ActiveRunCount(t *testing.T) { + eng := engine.New(5) + defer eng.Shutdown() + + rm := newRunManager(eng, nil) + + if rm.activeRunCount() != 0 { + t.Errorf("expected 0 active runs, got %d", rm.activeRunCount()) + } + + rm.mu.Lock() + rm.runs["test/feature"] = &runInfo{ + project: "test", + prdID: "feature", + } + rm.mu.Unlock() + + if rm.activeRunCount() != 1 { + t.Errorf("expected 1 active run, got %d", rm.activeRunCount()) + } +} + +func TestRunServe_RunProgressStreaming(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo with a PRD + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test Story", "passes": false}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Create a mock claude that outputs stream-json with a story start marker then exits successfully + mockDir := t.TempDir() + mockScript := `#!/bin/sh +echo '{"type":"system","subtype":"init"}' +echo '{"type":"assistant","message":{"content":[{"type":"text","text":"Working on <ralph-status>US-001</ralph-status>"}]}}' +echo '{"type":"assistant","message":{"content":[{"type":"text","text":"Hello from Claude"}]}}' +echo '{"type":"result"}' +exit 0 +` + if err := os.WriteFile(filepath.Join(mockDir, "claude"), []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", mockDir+":"+origPath) + + var messages []map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send start_run request + startReq := map[string]string{ + "type": "start_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(startReq) + + // Read messages — expect run_progress and claude_output messages + raws, err := ms.waitForMessages(15, 5*time.Second) + if err == nil { + for _, data := range raws { + var msg map[string]interface{} + if json.Unmarshal(data, &msg) == nil { + mu.Lock() + messages = append(messages, msg) + mu.Unlock() + } + } + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + // Check for run_progress messages + foundIterationStart := false + foundStoryStarted := false + foundClaudeOutput := false + for _, msg := range messages { + if msg["type"] == "run_progress" { + status, _ := msg["status"].(string) + if status == "iteration_started" { + foundIterationStart = true + } + if status == "story_started" { + foundStoryStarted = true + if msg["story_id"] != "US-001" { + t.Errorf("expected story_id 'US-001', got %v", msg["story_id"]) + } + if msg["project"] != "myproject" { + t.Errorf("expected project 'myproject', got %v", msg["project"]) + } + if msg["prd_id"] != "feature" { + t.Errorf("expected prd_id 'feature', got %v", msg["prd_id"]) + } + } + } + if msg["type"] == "claude_output" { + foundClaudeOutput = true + if msg["project"] != "myproject" { + t.Errorf("expected project 'myproject', got %v", msg["project"]) + } + } + } + + if !foundIterationStart { + t.Errorf("expected run_progress with status 'iteration_started', messages: %v", messages) + } + if !foundStoryStarted { + t.Errorf("expected run_progress with status 'story_started', messages: %v", messages) + } + if !foundClaudeOutput { + t.Errorf("expected claude_output messages, messages: %v", messages) + } +} diff --git a/internal/cmd/sender.go b/internal/cmd/sender.go new file mode 100644 index 00000000..5b7f08e5 --- /dev/null +++ b/internal/cmd/sender.go @@ -0,0 +1,45 @@ +package cmd + +import ( + "encoding/json" + "fmt" + "log" + + "github.com/minicodemonkey/chief/internal/uplink" +) + +// messageSender is an interface for sending messages to the server. +// The uplink adapter satisfies this interface. +type messageSender interface { + Send(msg interface{}) error +} + +// uplinkSender adapts *uplink.Uplink to the messageSender interface. +// It JSON-marshals the message, extracts the "type" field for the batcher's +// priority tier classification, and enqueues via Uplink.Send(). +type uplinkSender struct { + uplink *uplink.Uplink +} + +func newUplinkSender(u *uplink.Uplink) *uplinkSender { + return &uplinkSender{uplink: u} +} + +func (s *uplinkSender) Send(msg interface{}) error { + data, err := json.Marshal(msg) + if err != nil { + return fmt.Errorf("marshaling message: %w", err) + } + + // Extract the "type" field for batcher tier classification. + var envelope struct { + Type string `json:"type"` + } + if err := json.Unmarshal(data, &envelope); err != nil { + log.Printf("uplinkSender: could not extract message type: %v", err) + envelope.Type = "unknown" + } + + s.uplink.Send(data, envelope.Type) + return nil +} diff --git a/internal/cmd/serve.go b/internal/cmd/serve.go new file mode 100644 index 00000000..a4f98bb3 --- /dev/null +++ b/internal/cmd/serve.go @@ -0,0 +1,655 @@ +package cmd + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log" + "os" + "os/signal" + "path/filepath" + "syscall" + "time" + + "github.com/minicodemonkey/chief/internal/auth" + "github.com/minicodemonkey/chief/internal/engine" + "github.com/minicodemonkey/chief/internal/uplink" + "github.com/minicodemonkey/chief/internal/update" + "github.com/minicodemonkey/chief/internal/workspace" + "github.com/minicodemonkey/chief/internal/ws" +) + +const ( + // defaultServerURL is the default HTTP base URL for the chief server. + defaultServerURL = "https://uplink.chiefloop.com" +) + +// ServeOptions contains configuration for the serve command. +type ServeOptions struct { + Workspace string // Path to workspace directory + DeviceName string // Override device name (default: from credentials) + LogFile string // Path to log file (default: stdout) + BaseURL string // Override base URL (for testing) + ServerURL string // Override server URL for uplink (for testing/dev) + Version string // Chief version string + ReleasesURL string // Override GitHub releases URL (for testing) + Ctx context.Context // Optional context for cancellation (for testing) + +} + +// RunServe starts the headless serve daemon. +func RunServe(opts ServeOptions) error { + // Validate workspace directory exists + if opts.Workspace == "" { + opts.Workspace = "." + } + absWorkspace, err := filepath.Abs(opts.Workspace) + if err != nil { + return fmt.Errorf("resolving workspace path: %w", err) + } + opts.Workspace = absWorkspace + info, err := os.Stat(opts.Workspace) + if err != nil { + if os.IsNotExist(err) { + return fmt.Errorf("workspace directory does not exist: %s", opts.Workspace) + } + return fmt.Errorf("checking workspace directory: %w", err) + } + if !info.IsDir() { + return fmt.Errorf("workspace path is not a directory: %s", opts.Workspace) + } + + // Set up logging + var logFile *os.File + if opts.LogFile != "" { + f, err := os.OpenFile(opts.LogFile, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644) + if err != nil { + return fmt.Errorf("opening log file: %w", err) + } + logFile = f + defer func() { + logFile.Sync() + logFile.Close() + }() + log.SetOutput(f) + } + + // Check for credentials + creds, err := auth.LoadCredentials() + if err != nil { + if errors.Is(err, auth.ErrNotLoggedIn) { + return fmt.Errorf("Not logged in. Run 'chief login' first.") + } + return fmt.Errorf("loading credentials: %w", err) + } + + // Refresh token if near-expiry + if creds.IsNearExpiry(5 * time.Minute) { + log.Println("Access token near expiry, refreshing...") + creds, err = auth.RefreshToken(opts.BaseURL) + if err != nil { + return fmt.Errorf("refreshing token: %w", err) + } + log.Println("Token refreshed successfully") + } + + // Determine device name + deviceName := opts.DeviceName + if deviceName == "" { + deviceName = creds.DeviceName + } + + // Determine server URL (precedence: ServerURL flag > env > default) + serverURL := opts.ServerURL + if serverURL == "" { + serverURL = os.Getenv("CHIEF_SERVER_URL") + } + if serverURL == "" { + serverURL = defaultServerURL + } + + log.Printf("Starting chief serve (workspace: %s, device: %s)", opts.Workspace, deviceName) + log.Printf("Connecting to %s", serverURL) + + // Set up context with cancellation for clean shutdown + ctx := opts.Ctx + if ctx == nil { + ctx = context.Background() + } + ctx, cancel := context.WithCancel(ctx) + defer cancel() + + // Determine version string + version := opts.Version + if version == "" { + version = "dev" + } + + // Start workspace scanner (before connect so initial scan is ready) + scanner := workspace.New(opts.Workspace, nil) // sender set after connect + scanner.ScanAndUpdate() + + // Create engine for Ralph loop runs (default 5 max iterations) + eng := engine.New(5) + + // Create rate limiter for incoming messages + rateLimiter := ws.NewRateLimiter() + + // Create the uplink HTTP client + httpClient, err := uplink.New(serverURL, creds.AccessToken, + uplink.WithDeviceName(deviceName), + uplink.WithChiefVersion(version), + ) + if err != nil { + return fmt.Errorf("creating uplink client: %w", err) + } + + // Create the uplink with reconnect handler that re-sends state + var sender messageSender + var sessions *sessionManager + var runs *runManager + var ul *uplink.Uplink + + ul = uplink.NewUplink(httpClient, + uplink.WithOnReconnect(func() { + log.Println("Uplink reconnected, re-sending state snapshot") + rateLimiter.Reset() + sendStateSnapshot(sender, scanner, sessions, runs) + }), + uplink.WithOnAuthFailure(func() error { + log.Println("Auth failed during reconnection, refreshing token...") + newCreds, err := auth.RefreshToken(opts.BaseURL) + if err != nil { + return fmt.Errorf("token refresh failed: %w", err) + } + ul.SetAccessToken(newCreds.AccessToken) + log.Println("Token refreshed successfully during reconnection") + return nil + }), + ) + + // Create the sender adapter that wraps the uplink + sender = newUplinkSender(ul) + + // Set scanner's sender now that it exists + scanner.SetSender(sender) + + // Create session manager for Claude PRD sessions + sessions = newSessionManager(sender) + + // Create run manager for Ralph loop runs + runs = newRunManager(eng, sender) + + // Start engine event monitor for quota detection + runs.startEventMonitor(ctx) + + // Connect to server (HTTP connect + Pusher subscribe + batcher start) + if err := ul.Connect(ctx); err != nil { + if errors.Is(err, uplink.ErrAuthFailed) { + return fmt.Errorf("Device deauthorized. Run 'chief login' to re-authenticate.") + } + if errors.Is(err, uplink.ErrDeviceRevoked) { + return fmt.Errorf("Device deauthorized. Run 'chief login' to re-authenticate.") + } + return fmt.Errorf("connecting to server: %w", err) + } + log.Println("Connected to server") + + // Send initial state snapshot after successful connect + sendStateSnapshot(sender, scanner, sessions, runs) + + // Start periodic scanning loop + go scanner.Run(ctx) + log.Println("Workspace scanner started") + + // Start file watcher + watcher, err := workspace.NewWatcher(opts.Workspace, scanner, sender) + if err != nil { + log.Printf("Warning: could not start file watcher: %v", err) + } else { + go watcher.Run(ctx) + log.Println("File watcher started") + } + + // Start periodic version check (every 24 hours) + go runVersionChecker(ctx, sender, opts.Version, opts.ReleasesURL) + + // Set up signal handling + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGTERM, syscall.SIGINT) + defer signal.Stop(sigCh) + + log.Println("Serve is running. Press Ctrl+C to stop.") + + // Main event loop — commands arrive from Pusher via uplink.Receive() + for { + select { + case <-ctx.Done(): + log.Println("Context cancelled, shutting down...") + return serveShutdown(ul, watcher, sessions, runs, eng) + + case sig := <-sigCh: + log.Printf("Received signal %s, shutting down...", sig) + return serveShutdown(ul, watcher, sessions, runs, eng) + + case raw, ok := <-ul.Receive(): + if !ok { + // Channel closed, connection lost permanently + log.Println("Uplink connection closed permanently") + return serveShutdown(ul, watcher, sessions, runs, eng) + } + + // Parse the raw JSON into a ws.Message for dispatch + var msg ws.Message + if err := json.Unmarshal(raw, &msg); err != nil { + log.Printf("Ignoring unparseable command: %v", err) + continue + } + msg.Raw = raw + + // Extract payload wrapper if present. + // The CommandRelayController sends {"type": "...", "payload": {...}} + // but handlers expect fields at the top level of msg.Raw. + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(raw, &env); err == nil && len(env.Payload) > 0 { + msg.Raw = env.Payload + } + + // Check rate limit before processing + if result := rateLimiter.Allow(msg.Type); !result.Allowed { + log.Printf("Rate limited message type=%s, retry after %s", msg.Type, ws.FormatRetryAfter(result.RetryAfter)) + sendError(sender, ws.ErrCodeRateLimited, + fmt.Sprintf("Rate limited. Try again in %s.", ws.FormatRetryAfter(result.RetryAfter)), + msg.ID) + continue + } + + if shouldExit := handleMessage(sender, scanner, watcher, sessions, runs, msg, version, opts.ReleasesURL); shouldExit { + log.Println("Update installed, exiting for restart...") + serveShutdown(ul, watcher, sessions, runs, eng) + return nil + } + } + } +} + +// sendStateSnapshot sends a full state snapshot via the uplink. +func sendStateSnapshot(sender messageSender, scanner *workspace.Scanner, sessions *sessionManager, runs *runManager) { + projects := scanner.Projects() + envelope := ws.NewMessage(ws.TypeStateSnapshot) + + var activeSessions []ws.SessionState + if sessions != nil { + activeSessions = sessions.activeSessions() + } + if activeSessions == nil { + activeSessions = []ws.SessionState{} + } + + activeRuns := []ws.RunState{} + if runs != nil { + if r := runs.activeRuns(); r != nil { + activeRuns = r + } + } + + snapshot := ws.StateSnapshotMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Projects: projects, + Runs: activeRuns, + Sessions: activeSessions, + } + if err := sender.Send(snapshot); err != nil { + log.Printf("Error sending state_snapshot: %v", err) + } else { + log.Printf("Sent state_snapshot with %d projects", len(projects)) + } +} + +// sendError sends an error message. +func sendError(sender messageSender, code, message, requestID string) { + envelope := ws.NewMessage(ws.TypeError) + errMsg := ws.ErrorMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Code: code, + Message: message, + RequestID: requestID, + } + if err := sender.Send(errMsg); err != nil { + log.Printf("Error sending error message: %v", err) + } +} + +// handleMessage routes incoming commands. +// Returns true if the serve loop should exit (e.g., after a successful remote update). +func handleMessage(sender messageSender, scanner *workspace.Scanner, watcher *workspace.Watcher, sessions *sessionManager, runs *runManager, msg ws.Message, version, releasesURL string) bool { + log.Printf("Received command type=%s id=%s", msg.Type, msg.ID) + + switch msg.Type { + case ws.TypePing: + pong := ws.NewMessage(ws.TypePong) + if err := sender.Send(pong); err != nil { + log.Printf("Error sending pong: %v", err) + } + + case ws.TypeListProjects: + handleListProjects(sender, scanner) + + case ws.TypeGetProject: + handleGetProject(sender, scanner, watcher, msg) + + case ws.TypeGetPRD: + handleGetPRD(sender, scanner, msg) + + case ws.TypeGetPRDs: + handleGetPRDs(sender, scanner, msg) + + case ws.TypeNewPRD: + handleNewPRD(sender, scanner, sessions, msg) + + case ws.TypeRefinePRD: + handleRefinePRD(sender, scanner, sessions, msg) + + case ws.TypePRDMessage: + handlePRDMessage(sender, sessions, msg) + + case ws.TypeClosePRDSession: + handleClosePRDSession(sender, sessions, msg) + + case ws.TypeStartRun: + handleStartRun(sender, scanner, runs, watcher, msg) + + case ws.TypePauseRun: + handlePauseRun(sender, runs, msg) + + case ws.TypeResumeRun: + handleResumeRun(sender, runs, msg) + + case ws.TypeStopRun: + handleStopRun(sender, runs, msg) + + case ws.TypeGetDiff: + handleGetDiff(sender, scanner, msg) + + case ws.TypeGetDiffs: + handleGetDiffs(sender, scanner, msg) + + case ws.TypeGetLogs: + handleGetLogs(sender, scanner, msg) + + case ws.TypeGetSettings: + handleGetSettings(sender, scanner, msg) + + case ws.TypeUpdateSettings: + handleUpdateSettings(sender, scanner, msg) + + case ws.TypeCloneRepo: + handleCloneRepo(sender, scanner, msg) + + case ws.TypeCreateProject: + handleCreateProject(sender, scanner, msg) + + case ws.TypeTriggerUpdate: + return handleTriggerUpdate(sender, msg, version, releasesURL) + + default: + log.Printf("Received message type: %s", msg.Type) + } + return false +} + +// handleListProjects handles a list_projects request. +func handleListProjects(sender messageSender, scanner *workspace.Scanner) { + projects := scanner.Projects() + envelope := ws.NewMessage(ws.TypeProjectList) + plMsg := ws.ProjectListMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Projects: projects, + } + if err := sender.Send(plMsg); err != nil { + log.Printf("Error sending project_list: %v", err) + } +} + +// handleGetProject handles a get_project request. +func handleGetProject(sender messageSender, scanner *workspace.Scanner, watcher *workspace.Watcher, msg ws.Message) { + var req ws.GetProjectMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing get_project message: %v", err) + return + } + + project, found := scanner.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + // Activate file watching for the requested project + if watcher != nil { + watcher.Activate(req.Project) + } + + envelope := ws.NewMessage(ws.TypeProjectState) + psMsg := ws.ProjectStateMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: project, + } + if err := sender.Send(psMsg); err != nil { + log.Printf("Error sending project_state: %v", err) + } +} + +// handleGetPRD handles a get_prd request. +func handleGetPRD(sender messageSender, scanner *workspace.Scanner, msg ws.Message) { + var req ws.GetPRDMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing get_prd message: %v", err) + return + } + + project, found := scanner.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + // Read PRD markdown content + prdDir := filepath.Join(project.Path, ".chief", "prds", req.PRDID) + prdMD := filepath.Join(prdDir, "prd.md") + prdJSON := filepath.Join(prdDir, "prd.json") + + // Check that the PRD directory exists + if _, err := os.Stat(prdDir); os.IsNotExist(err) { + sendError(sender, ws.ErrCodePRDNotFound, + fmt.Sprintf("PRD %q not found in project %q", req.PRDID, req.Project), msg.ID) + return + } + + // Read markdown content (optional — may not exist yet) + var content string + if data, err := os.ReadFile(prdMD); err == nil { + content = string(data) + } + + // Read prd.json state + var state interface{} + if data, err := os.ReadFile(prdJSON); err == nil { + var parsed interface{} + if json.Unmarshal(data, &parsed) == nil { + state = parsed + } + } + + envelope := ws.NewMessage(ws.TypePRDContent) + prdMsg := ws.PRDContentMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + Project: req.Project, + PRDID: req.PRDID, + Content: content, + State: state, + } + if err := sender.Send(prdMsg); err != nil { + log.Printf("Error sending prd_content: %v", err) + } +} + +// runVersionChecker periodically checks for updates and sends update_available. +func runVersionChecker(ctx context.Context, sender messageSender, version, releasesURL string) { + // Check immediately on startup + checkAndNotify(sender, version, releasesURL) + + ticker := time.NewTicker(24 * time.Hour) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + checkAndNotify(sender, version, releasesURL) + } + } +} + +// checkAndNotify performs a version check and sends update_available if needed. +func checkAndNotify(sender messageSender, version, releasesURL string) { + result, err := update.CheckForUpdate(version, update.Options{ + ReleasesURL: releasesURL, + }) + if err != nil { + log.Printf("Version check failed: %v", err) + return + } + if result.UpdateAvailable { + log.Printf("Update available: v%s (current: v%s)", result.LatestVersion, result.CurrentVersion) + envelope := ws.NewMessage(ws.TypeUpdateAvailable) + msg := ws.UpdateAvailableMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + CurrentVersion: result.CurrentVersion, + LatestVersion: result.LatestVersion, + } + if err := sender.Send(msg); err != nil { + log.Printf("Error sending update_available: %v", err) + } + } +} + +// uplinkCloser is an interface for closing the uplink connection. +type uplinkCloser interface { + Close() error + CloseWithTimeout(timeout time.Duration) error +} + +// shutdownTimeout is the maximum time allowed for the entire shutdown sequence. +// This covers: process killing, batcher flush, Pusher close, and HTTP disconnect. +const shutdownTimeout = 10 * time.Second + +// processKillTimeout is the maximum time to wait for Claude processes to exit gracefully. +const processKillTimeout = 5 * time.Second + +// serveShutdown performs clean shutdown of the serve command. +// It kills all child Claude processes, marks interrupted stories, closes the uplink, +// and flushes log files. Processes are force-killed after 5 seconds if they haven't +// exited gracefully. The entire shutdown completes within 10 seconds. +func serveShutdown(closer uplinkCloser, watcher *workspace.Watcher, sessions *sessionManager, runs *runManager, eng *engine.Engine) error { + log.Println("Shutting down...") + + // Enforce an overall shutdown deadline. + shutdownDone := make(chan struct{}) + go func() { + defer close(shutdownDone) + doShutdown(closer, watcher, sessions, runs, eng) + }() + + select { + case <-shutdownDone: + // Normal shutdown completed within the timeout. + case <-time.After(shutdownTimeout): + log.Printf("Shutdown timed out after %s — forcing exit", shutdownTimeout) + } + + log.Println("Goodbye.") + return nil +} + +// doShutdown performs the actual shutdown sequence. +func doShutdown(closer uplinkCloser, watcher *workspace.Watcher, sessions *sessionManager, runs *runManager, eng *engine.Engine) { + // Count processes before shutdown. + processCount := 0 + if sessions != nil { + processCount += sessions.sessionCount() + } + if runs != nil { + processCount += runs.activeRunCount() + } + + // Mark any in-progress stories as interrupted in prd.json. + if runs != nil { + runs.markInterruptedStories() + } + + // Use a channel to track when graceful process shutdown completes. + done := make(chan struct{}) + go func() { + // Stop all active Ralph loop runs. + if runs != nil { + runs.stopAll() + } + + // Kill all active Claude sessions. + if sessions != nil { + sessions.killAll() + } + + // Shut down the engine (stops event forwarding goroutine). + if eng != nil { + eng.Shutdown() + } + + close(done) + }() + + // Wait for graceful shutdown or force-kill after 5 seconds. + select { + case <-done: + // Graceful shutdown completed. + case <-time.After(processKillTimeout): + log.Println("Force-killing hung processes after 5 second timeout") + } + + if processCount > 0 { + log.Printf("Killed %d processes", processCount) + } + + // Close file watcher. + if watcher != nil { + if err := watcher.Close(); err != nil { + log.Printf("Error closing file watcher: %v", err) + } + } + + // Close the uplink with a timeout to prevent hanging on unreachable servers. + // The batcher flush + Pusher close + HTTP disconnect must complete within this window. + if err := closer.CloseWithTimeout(5 * time.Second); err != nil { + log.Printf("Error closing connection: %v", err) + } +} diff --git a/internal/cmd/serve_test.go b/internal/cmd/serve_test.go new file mode 100644 index 00000000..8f40b64c --- /dev/null +++ b/internal/cmd/serve_test.go @@ -0,0 +1,1663 @@ +package cmd + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/auth" +) + +// serveWsURL converts an httptest.Server URL to a WebSocket URL. +// Kept for use by session_test.go and remote_update_test.go. +func serveWsURL(s *httptest.Server) string { + return "ws" + strings.TrimPrefix(s.URL, "http") +} + +func setupServeCredentials(t *testing.T) { + t.Helper() + creds := &auth.Credentials{ + AccessToken: "test-token", + RefreshToken: "test-refresh", + ExpiresAt: time.Now().Add(time.Hour), + DeviceName: "test-device", + User: "user@example.com", + } + if err := auth.SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } +} + +func TestRunServe_WorkspaceDefaultsToCwd(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + setupServeCredentials(t) + + // Empty workspace should default to "." and resolve to an absolute path. + // Cancel immediately — we only care that workspace validation passes. + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + err := RunServe(ServeOptions{Ctx: ctx}) + if err != nil && strings.Contains(err.Error(), "does not exist") { + t.Errorf("empty workspace should default to cwd, got: %v", err) + } +} + +func TestRunServe_WorkspaceDoesNotExist(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + setupServeCredentials(t) + + err := RunServe(ServeOptions{ + Workspace: "/nonexistent/path", + }) + if err == nil { + t.Fatal("expected error for nonexistent workspace") + } + if !strings.Contains(err.Error(), "does not exist") { + t.Errorf("expected 'does not exist' error, got: %v", err) + } +} + +func TestRunServe_WorkspaceIsFile(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + setupServeCredentials(t) + + // Create a file instead of directory + filePath := filepath.Join(home, "not-a-dir") + if err := os.WriteFile(filePath, []byte("not a directory"), 0o644); err != nil { + t.Fatal(err) + } + + err := RunServe(ServeOptions{ + Workspace: filePath, + }) + if err == nil { + t.Fatal("expected error for file workspace") + } + if !strings.Contains(err.Error(), "not a directory") { + t.Errorf("expected 'not a directory' error, got: %v", err) + } +} + +func TestRunServe_NotLoggedIn(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + err := RunServe(ServeOptions{ + Workspace: workspace, + }) + if err == nil { + t.Fatal("expected error for missing credentials") + } + if !strings.Contains(err.Error(), "Not logged in") { + t.Errorf("expected 'Not logged in' error, got: %v", err) + } +} + +func TestRunServe_ConnectsAndHandshakes(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Check the connect body for expected metadata + connectBody := ms.getConnectBody() + if connectBody == nil { + t.Fatal("connect request body was not received") + } + + if connectBody["chief_version"] != "1.0.0" { + t.Errorf("expected chief_version '1.0.0', got %v", connectBody["chief_version"]) + } + if connectBody["device_name"] != "test-device" { + t.Errorf("expected device_name 'test-device', got %v", connectBody["device_name"]) + } +} + +func TestRunServe_DeviceNameOverride(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + DeviceName: "my-custom-device", + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + connectBody := ms.getConnectBody() + if connectBody == nil { + t.Fatal("connect request body was not received") + } + + if connectBody["device_name"] != "my-custom-device" { + t.Errorf("expected device name 'my-custom-device', got %q", connectBody["device_name"]) + } +} + +func TestRunServe_AuthFailed(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + ms := newMockUplinkServer(t) + // Set connect endpoint to return 401 + ms.connectStatus.Store(401) + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + }) + if err == nil { + t.Fatal("expected error for auth failure") + } + if !strings.Contains(err.Error(), "deauthorized") { + t.Errorf("expected 'deauthorized' error, got: %v", err) + } +} + +func TestRunServe_LogFile(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + logFile := filepath.Join(home, "chief.log") + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + LogFile: logFile, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify log file was created and has content + data, err := os.ReadFile(logFile) + if err != nil { + t.Fatalf("failed to read log file: %v", err) + } + if len(data) == 0 { + t.Error("log file is empty") + } + content := string(data) + if !strings.Contains(content, "Starting chief serve") { + t.Errorf("log file missing startup message, got: %s", content) + } +} + +func TestRunServe_PingPong(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + // Wait for connection + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for initial state_snapshot + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send ping command via Pusher + pingReq := map[string]string{ + "type": "ping", + "id": "ping-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + } + if err := ms.sendCommand(pingReq); err != nil { + t.Logf("sendCommand(ping): %v", err) + cancel() + return + } + + // Wait for pong response via HTTP messages + if _, err := ms.waitForMessageType("pong", 5*time.Second); err != nil { + t.Logf("waitForMessageType(pong): %v", err) + } + + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify pong was received + pong, err := ms.waitForMessageType("pong", time.Second) + if err != nil { + t.Error("expected pong response to be received by server") + } else { + var msg map[string]interface{} + json.Unmarshal(pong, &msg) + if msg["type"] != "pong" { + t.Errorf("expected type 'pong', got %v", msg["type"]) + } + } +} + +func TestRunServe_TokenRefresh(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Save credentials that are near expiry + creds := &auth.Credentials{ + AccessToken: "old-token", + RefreshToken: "test-refresh", + ExpiresAt: time.Now().Add(2 * time.Minute), // Within 5 min threshold + DeviceName: "test-device", + User: "user@example.com", + } + if err := auth.SaveCredentials(creds); err != nil { + t.Fatalf("SaveCredentials failed: %v", err) + } + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + var tokenRefreshed bool + var mu sync.Mutex + + ctx, cancel := context.WithCancel(context.Background()) + + // Create a mock uplink server for the uplink endpoints + ms := newMockUplinkServer(t) + + // Create a mux that combines token refresh with uplink server endpoints. + // Token refresh goes to BaseURL, uplink goes to ServerURL. We use + // a separate mux server as the BaseURL for token refresh. + tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/oauth/token" { + mu.Lock() + tokenRefreshed = true + mu.Unlock() + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]interface{}{ + "access_token": "new-refreshed-token", + "refresh_token": "new-refresh", + "expires_in": 3600, + }) + return + } + http.NotFound(w, r) + })) + defer tokenSrv.Close() + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + BaseURL: tokenSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if !tokenRefreshed { + t.Error("expected token refresh to be called for near-expiry credentials") + } +} + +// createGitRepo creates a minimal git repository for testing. +func createGitRepo(t *testing.T, dir string) { + t.Helper() + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatal(err) + } + // Initialize git repo + cmd := exec.Command("git", "init", dir) + cmd.Env = append(os.Environ(), "GIT_CONFIG_GLOBAL=/dev/null", "GIT_CONFIG_SYSTEM=/dev/null") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git init failed: %v\n%s", err, out) + } + // Configure user for the repo + for _, args := range [][]string{ + {"config", "user.email", "test@test.com"}, + {"config", "user.name", "Test"}, + } { + cmd := exec.Command("git", args...) + cmd.Dir = dir + cmd.Env = append(os.Environ(), "GIT_CONFIG_GLOBAL=/dev/null", "GIT_CONFIG_SYSTEM=/dev/null") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git config failed: %v\n%s", err, out) + } + } + // Create initial commit + readmePath := filepath.Join(dir, "README.md") + if err := os.WriteFile(readmePath, []byte("# Test\n"), 0o644); err != nil { + t.Fatal(err) + } + cmd = exec.Command("git", "add", ".") + cmd.Dir = dir + cmd.Env = append(os.Environ(), "GIT_CONFIG_GLOBAL=/dev/null", "GIT_CONFIG_SYSTEM=/dev/null") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git add failed: %v\n%s", err, out) + } + cmd = exec.Command("git", "commit", "-m", "initial commit") + cmd.Dir = dir + cmd.Env = append(os.Environ(), "GIT_CONFIG_GLOBAL=/dev/null", "GIT_CONFIG_SYSTEM=/dev/null") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git commit failed: %v\n%s", err, out) + } +} + +// serveTestHelper sets up a serve test with a mock uplink server. +// The serverFn receives the mock uplink server after the CLI has connected +// (HTTP connect + Pusher subscribe) and sent the initial state_snapshot. +// The serverFn should send commands via ms.sendCommand() and read responses +// via ms.waitForMessageType() or ms.getMessages(). +func serveTestHelper(t *testing.T, workspacePath string, serverFn func(ms *mockUplinkServer)) error { + t.Helper() + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + // Wait for the CLI to connect and send state_snapshot, then run test logic. + go func() { + // Wait for Pusher subscription (indicates full connection). + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("serveTestHelper: %v", err) + cancel() + return + } + + // Wait for initial state_snapshot to arrive. + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("serveTestHelper: %v", err) + cancel() + return + } + + // Run test-specific server logic. + serverFn(ms) + + // Cancel context to stop serve loop. + cancel() + }() + + return RunServe(ServeOptions{ + Workspace: workspacePath, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) +} + +func TestRunServe_StateSnapshotOnConnect(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo in the workspace + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for state_snapshot + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + } + + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspaceDir, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Retrieve the state_snapshot message + raw, err := ms.waitForMessageType("state_snapshot", time.Second) + if err != nil { + t.Fatal("state_snapshot was not received") + } + + var snapshotReceived map[string]interface{} + json.Unmarshal(raw, &snapshotReceived) + + if snapshotReceived["type"] != "state_snapshot" { + t.Errorf("expected type 'state_snapshot', got %v", snapshotReceived["type"]) + } + + // Verify projects are included + projects, ok := snapshotReceived["projects"].([]interface{}) + if !ok { + t.Fatal("expected projects array in state_snapshot") + } + if len(projects) != 1 { + t.Fatalf("expected 1 project, got %d", len(projects)) + } + proj := projects[0].(map[string]interface{}) + if proj["name"] != "myproject" { + t.Errorf("expected project name 'myproject', got %v", proj["name"]) + } + + // Verify runs and sessions are empty arrays + runs, ok := snapshotReceived["runs"].([]interface{}) + if !ok { + t.Fatal("expected runs array in state_snapshot") + } + if len(runs) != 0 { + t.Errorf("expected 0 runs, got %d", len(runs)) + } + sessions, ok := snapshotReceived["sessions"].([]interface{}) + if !ok { + t.Fatal("expected sessions array in state_snapshot") + } + if len(sessions) != 0 { + t.Errorf("expected 0 sessions, got %d", len(sessions)) + } +} + +func TestRunServe_ListProjects(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create two git repos + createGitRepo(t, filepath.Join(workspaceDir, "alpha")) + createGitRepo(t, filepath.Join(workspaceDir, "beta")) + + var projectListReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send list_projects request + listReq := map[string]string{ + "type": "list_projects", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + } + ms.sendCommand(listReq) + + // Read project_list response + raw, err := ms.waitForMessageType("project_list", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &projectListReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if projectListReceived == nil { + t.Fatal("project_list was not received") + } + if projectListReceived["type"] != "project_list" { + t.Errorf("expected type 'project_list', got %v", projectListReceived["type"]) + } + + projects, ok := projectListReceived["projects"].([]interface{}) + if !ok { + t.Fatal("expected projects array") + } + if len(projects) != 2 { + t.Fatalf("expected 2 projects, got %d", len(projects)) + } + + // Collect project names + names := make(map[string]bool) + for _, p := range projects { + proj := p.(map[string]interface{}) + names[proj["name"].(string)] = true + } + if !names["alpha"] || !names["beta"] { + t.Errorf("expected projects alpha and beta, got %v", names) + } +} + +func TestRunServe_GetProject(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var projectStateReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send get_project request + getReq := map[string]string{ + "type": "get_project", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + } + ms.sendCommand(getReq) + + // Read project_state response + raw, err := ms.waitForMessageType("project_state", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &projectStateReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if projectStateReceived == nil { + t.Fatal("project_state was not received") + } + if projectStateReceived["type"] != "project_state" { + t.Errorf("expected type 'project_state', got %v", projectStateReceived["type"]) + } + + project, ok := projectStateReceived["project"].(map[string]interface{}) + if !ok { + t.Fatal("expected project object in project_state") + } + if project["name"] != "myproject" { + t.Errorf("expected project name 'myproject', got %v", project["name"]) + } +} + +func TestRunServe_GetProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send get_project for nonexistent project + getReq := map[string]string{ + "type": "get_project", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + } + ms.sendCommand(getReq) + + // Read error response + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", errorReceived["code"]) + } +} + +func TestRunServe_GetPRD(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo with .chief/prds/feature/ + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + // Write prd.md + prdMD := "# My Feature\nThis is a feature PRD." + if err := os.WriteFile(filepath.Join(prdDir, "prd.md"), []byte(prdMD), 0o644); err != nil { + t.Fatal(err) + } + + // Write prd.json + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "passes": true}]}` + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + var prdContentReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send get_prd request + getReq := map[string]string{ + "type": "get_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(getReq) + + // Read prd_content response + raw, err := ms.waitForMessageType("prd_content", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &prdContentReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if prdContentReceived == nil { + t.Fatal("prd_content was not received") + } + if prdContentReceived["type"] != "prd_content" { + t.Errorf("expected type 'prd_content', got %v", prdContentReceived["type"]) + } + if prdContentReceived["project"] != "myproject" { + t.Errorf("expected project 'myproject', got %v", prdContentReceived["project"]) + } + if prdContentReceived["prd_id"] != "feature" { + t.Errorf("expected prd_id 'feature', got %v", prdContentReceived["prd_id"]) + } + if prdContentReceived["content"] != prdMD { + t.Errorf("expected content %q, got %v", prdMD, prdContentReceived["content"]) + } + + // Verify state is present and contains expected data + state, ok := prdContentReceived["state"].(map[string]interface{}) + if !ok { + t.Fatal("expected state object in prd_content") + } + if state["project"] != "My Feature" { + t.Errorf("expected state.project 'My Feature', got %v", state["project"]) + } +} + +func TestRunServe_GetPRDNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo without any PRDs + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send get_prd for nonexistent PRD + getReq := map[string]string{ + "type": "get_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "nonexistent", + } + ms.sendCommand(getReq) + + // Read error response + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "PRD_NOT_FOUND" { + t.Errorf("expected code 'PRD_NOT_FOUND', got %v", errorReceived["code"]) + } +} + +func TestRunServe_GetPRDProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send get_prd for nonexistent project + getReq := map[string]string{ + "type": "get_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "prd_id": "feature", + } + ms.sendCommand(getReq) + + // Read error response + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", errorReceived["code"]) + } +} + +func TestRunServe_RateLimitGlobal(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var rateLimitReceived bool + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send more than globalBurst (30) messages rapidly to trigger rate limiting + for i := 0; i < 35; i++ { + msg := map[string]string{ + "type": "list_projects", + "id": fmt.Sprintf("req-%d", i), + "timestamp": time.Now().UTC().Format(time.RFC3339), + } + ms.sendCommand(msg) + } + + // Wait for a RATE_LIMITED error response + deadline := time.After(5 * time.Second) + for { + msgs := ms.getMessages() + for _, raw := range msgs { + var resp map[string]interface{} + json.Unmarshal(raw, &resp) + if resp["type"] == "error" && resp["code"] == "RATE_LIMITED" { + mu.Lock() + rateLimitReceived = true + mu.Unlock() + return + } + } + select { + case <-deadline: + return + case <-time.After(50 * time.Millisecond): + } + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if !rateLimitReceived { + t.Error("expected RATE_LIMITED error after burst exhaustion") + } +} + +func TestRunServe_RateLimitPingExempt(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var pongReceived bool + var rateLimitSeen bool + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Exhaust the global rate limit with normal messages + for i := 0; i < 35; i++ { + msg := map[string]string{ + "type": "list_projects", + "id": fmt.Sprintf("req-%d", i), + "timestamp": time.Now().UTC().Format(time.RFC3339), + } + ms.sendCommand(msg) + } + + // Now immediately send a ping — should bypass rate limiting + ping := map[string]string{ + "type": "ping", + "id": "ping-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + } + ms.sendCommand(ping) + + // Wait for both RATE_LIMITED and pong responses + deadline := time.After(5 * time.Second) + for { + msgs := ms.getMessages() + for _, raw := range msgs { + var resp map[string]interface{} + json.Unmarshal(raw, &resp) + if resp["type"] == "pong" { + mu.Lock() + pongReceived = true + mu.Unlock() + } + if resp["type"] == "error" && resp["code"] == "RATE_LIMITED" { + mu.Lock() + rateLimitSeen = true + mu.Unlock() + } + } + mu.Lock() + done := pongReceived && rateLimitSeen + mu.Unlock() + if done { + return + } + select { + case <-deadline: + return + case <-time.After(50 * time.Millisecond): + } + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if !rateLimitSeen { + t.Error("expected RATE_LIMITED error to confirm rate limiting was active") + } + if !pongReceived { + t.Error("expected pong response even after rate limit exhaustion — ping should be exempt") + } +} + +func TestRunServe_RateLimitExpensiveOps(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var rateLimitReceived bool + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send 3 start_run messages (limit is 2/minute) + for i := 0; i < 3; i++ { + msg := map[string]interface{}{ + "type": "start_run", + "id": fmt.Sprintf("req-%d", i), + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "prd_id": "test", + } + ms.sendCommand(msg) + } + + // Wait for a RATE_LIMITED error response + deadline := time.After(5 * time.Second) + for { + msgs := ms.getMessages() + for _, raw := range msgs { + var resp map[string]interface{} + json.Unmarshal(raw, &resp) + if resp["type"] == "error" && resp["code"] == "RATE_LIMITED" { + mu.Lock() + rateLimitReceived = true + mu.Unlock() + return + } + } + select { + case <-deadline: + return + case <-time.After(50 * time.Millisecond): + } + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if !rateLimitReceived { + t.Error("expected RATE_LIMITED error for excessive expensive operations") + } +} + +func TestRunServe_ShutdownLogsSequence(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + logFile := filepath.Join(home, "chief-shutdown.log") + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for state_snapshot to ensure connection is fully established + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + } + + // Cancel to trigger shutdown + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + LogFile: logFile, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify log file contains shutdown sequence + data, err := os.ReadFile(logFile) + if err != nil { + t.Fatalf("failed to read log file: %v", err) + } + content := string(data) + + if !strings.Contains(content, "Shutting down...") { + t.Error("log file missing 'Shutting down...' message") + } + if !strings.Contains(content, "Goodbye.") { + t.Error("log file missing 'Goodbye.' message") + } +} + +func TestRunServe_ShutdownMarksInterruptedStories(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a git repo with a PRD + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdPath := filepath.Join(prdDir, "prd.json") + prdState := `{"project": "My Feature", "userStories": [{"id": "US-001", "title": "Test Story", "passes": false}, {"id": "US-002", "title": "Done Story", "passes": true}]}` + if err := os.WriteFile(prdPath, []byte(prdState), 0o644); err != nil { + t.Fatal(err) + } + + // Create a mock claude that hangs until killed (simulates in-progress work) + mockDir := t.TempDir() + mockScript := `#!/bin/sh +echo '{"type":"system","subtype":"init"}' +echo '{"type":"assistant","message":{"content":[{"type":"text","text":"Working on <ralph-status>US-001</ralph-status>"}]}}' +sleep 300 +` + if err := os.WriteFile(filepath.Join(mockDir, "claude"), []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", mockDir+":"+origPath) + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for state_snapshot + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send start_run to get a story in-progress + startReq := map[string]string{ + "type": "start_run", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "prd_id": "feature", + } + ms.sendCommand(startReq) + + // Wait for run_progress with story_started so we know US-001 is tracked + deadline := time.After(10 * time.Second) + for { + msgs := ms.getMessages() + for _, raw := range msgs { + var msg map[string]interface{} + json.Unmarshal(raw, &msg) + if msg["type"] == "run_progress" { + status, _ := msg["status"].(string) + if status == "story_started" { + // Now cancel to trigger shutdown while story is in-progress + cancel() + return + } + } + } + select { + case <-deadline: + t.Logf("timeout waiting for story_started") + cancel() + return + case <-time.After(50 * time.Millisecond): + } + } + }() + + err := RunServe(ServeOptions{ + Workspace: workspaceDir, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify the PRD was updated with inProgress: true for US-001 + data, err := os.ReadFile(prdPath) + if err != nil { + t.Fatalf("failed to read PRD: %v", err) + } + + var result map[string]interface{} + if err := json.Unmarshal(data, &result); err != nil { + t.Fatalf("failed to parse PRD: %v", err) + } + + stories := result["userStories"].([]interface{}) + story1 := stories[0].(map[string]interface{}) + if story1["inProgress"] != true { + t.Errorf("expected US-001 to have inProgress=true after shutdown, got %v", story1["inProgress"]) + } + + // US-002 is already passing, should NOT be marked as inProgress + story2 := stories[1].(map[string]interface{}) + if _, hasInProgress := story2["inProgress"]; hasInProgress && story2["inProgress"] == true { + t.Error("expected US-002 to NOT have inProgress=true (already passes)") + } +} + +func TestRunServe_ShutdownLogFileFlush(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + logFile := filepath.Join(home, "chief-flush.log") + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for state_snapshot + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + } + + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + LogFile: logFile, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify the log file was flushed and contains the final "Goodbye." message + data, err := os.ReadFile(logFile) + if err != nil { + t.Fatalf("failed to read log file: %v", err) + } + content := string(data) + if !strings.Contains(content, "Goodbye.") { + t.Error("log file missing 'Goodbye.' — may not have been flushed properly") + } +} + +func TestSessionManager_SessionCount(t *testing.T) { + sm := &sessionManager{ + sessions: make(map[string]*claudeSession), + stopTimeout: make(chan struct{}), + } + // Manually close stopTimeout since we're not starting the timeout checker + close(sm.stopTimeout) + + if sm.sessionCount() != 0 { + t.Errorf("expected 0 sessions, got %d", sm.sessionCount()) + } + + sm.sessions["sess1"] = &claudeSession{sessionID: "sess1"} + sm.sessions["sess2"] = &claudeSession{sessionID: "sess2"} + + if sm.sessionCount() != 2 { + t.Errorf("expected 2 sessions, got %d", sm.sessionCount()) + } +} + +func TestRunServe_ServerURLFromEnvVar(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + cancel() + }() + + // Set env var to point at our test server (no ServerURL in ServeOptions) + t.Setenv("CHIEF_SERVER_URL", ms.httpSrv.URL) + + err := RunServe(ServeOptions{ + Workspace: workspace, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestRunServe_ServerURLPrecedence_FlagOverridesEnv(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + cancel() + }() + + // Set env var to a bad URL — flag should override it + t.Setenv("CHIEF_SERVER_URL", "http://bad-url-that-should-not-be-used:9999") + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, // Flag value — should take precedence + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestRunServe_ServerURLLoggedOnStartup(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + logFile := filepath.Join(home, "serve.log") + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + cancel() + }() + + serverURL := ms.httpSrv.URL + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: serverURL, + LogFile: logFile, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + data, err := os.ReadFile(logFile) + if err != nil { + t.Fatalf("failed to read log file: %v", err) + } + content := string(data) + if !strings.Contains(content, "Connecting to "+serverURL) { + t.Errorf("expected log to contain 'Connecting to %s', got: %s", serverURL, content) + } +} + +// --- serveShutdown unit tests --- + +// mockCloser implements uplinkCloser for testing serveShutdown directly. +type mockCloser struct { + closeCalled atomic.Int32 + closeDelay time.Duration // Simulates a slow Close operation. + closeErr error +} + +func (m *mockCloser) Close() error { + m.closeCalled.Add(1) + if m.closeDelay > 0 { + time.Sleep(m.closeDelay) + } + return m.closeErr +} + +func (m *mockCloser) CloseWithTimeout(timeout time.Duration) error { + done := make(chan error, 1) + go func() { + done <- m.Close() + }() + + select { + case err := <-done: + return err + case <-time.After(timeout): + return nil + } +} + +func TestServeShutdown_CompletesWithinTimeout(t *testing.T) { + closer := &mockCloser{} + + start := time.Now() + err := serveShutdown(closer, nil, nil, nil, nil) + elapsed := time.Since(start) + + if err != nil { + t.Errorf("serveShutdown returned error: %v", err) + } + + // Should complete quickly. + if elapsed > 2*time.Second { + t.Errorf("serveShutdown took %s, expected < 2s", elapsed) + } + + // Close should have been called. + if got := closer.closeCalled.Load(); got != 1 { + t.Errorf("close called %d times, want 1", got) + } +} + +func TestServeShutdown_TimesOutWithHangingClose(t *testing.T) { + // Create a closer that hangs longer than the shutdown timeout. + closer := &mockCloser{closeDelay: 30 * time.Second} + + start := time.Now() + err := serveShutdown(closer, nil, nil, nil, nil) + elapsed := time.Since(start) + + if err != nil { + t.Errorf("serveShutdown returned error: %v", err) + } + + // Should complete within the shutdown timeout (10s) + small buffer, + // not hang for the full 30s close delay. + if elapsed > 15*time.Second { + t.Errorf("serveShutdown took %s, expected < 15s (shutdown timeout is 10s)", elapsed) + } + + t.Logf("serveShutdown completed in %s", elapsed.Round(time.Millisecond)) +} + +func TestServeShutdown_DisconnectFailureLoggedNotBlocking(t *testing.T) { + closer := &mockCloser{closeErr: fmt.Errorf("connection refused")} + + err := serveShutdown(closer, nil, nil, nil, nil) + + // serveShutdown should not propagate the close error. + if err != nil { + t.Errorf("serveShutdown returned error: %v, want nil", err) + } + + // Close should have been called. + if got := closer.closeCalled.Load(); got != 1 { + t.Errorf("close called %d times, want 1", got) + } +} + +func TestServeShutdown_CallsDisconnectOnServer(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspace := filepath.Join(home, "projects") + if err := os.MkdirAll(workspace, 0o755); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for state_snapshot to ensure connection is fully established. + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + } + + // Cancel to trigger shutdown. + cancel() + }() + + err := RunServe(ServeOptions{ + Workspace: workspace, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify disconnect was called during shutdown. + if got := ms.disconnectCalls.Load(); got != 1 { + t.Errorf("disconnect calls = %d, want 1", got) + } +} diff --git a/internal/cmd/serve_test_helper.go b/internal/cmd/serve_test_helper.go new file mode 100644 index 00000000..563a63be --- /dev/null +++ b/internal/cmd/serve_test_helper.go @@ -0,0 +1,423 @@ +package cmd + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/gorilla/websocket" +) + +// mockUplinkServer is a combined HTTP API + Pusher WebSocket server for testing +// the serve command with the uplink transport. It replaces the old WebSocket-only +// test server used before the uplink refactor. +type mockUplinkServer struct { + httpSrv *httptest.Server + pusherSrv *mockPusherServer + + mu sync.Mutex + messageBatches []mockMessageBatch + connectBody map[string]interface{} + + connectCalls atomic.Int32 + disconnectCalls atomic.Int32 + heartbeatCalls atomic.Int32 + messagesCalls atomic.Int32 + + // connectStatus controls the HTTP status returned by /api/device/connect. + // 0 means success (200). + connectStatus atomic.Int32 +} + +type mockMessageBatch struct { + BatchID string `json:"batch_id"` + Messages []json.RawMessage `json:"messages"` +} + +// newMockUplinkServer creates a new combined test server. +func newMockUplinkServer(t *testing.T) *mockUplinkServer { + t.Helper() + + ps := newMockPusherServer(t) + + ms := &mockUplinkServer{ + pusherSrv: ps, + } + + reverbCfg := ps.reverbConfig() + + ms.httpSrv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ms.handleHTTP(w, r, reverbCfg) + })) + t.Cleanup(func() { ms.httpSrv.Close() }) + + return ms +} + +func (ms *mockUplinkServer) handleHTTP(w http.ResponseWriter, r *http.Request, reverbCfg mockReverbConfig) { + // Check auth header. + auth := r.Header.Get("Authorization") + if !strings.HasPrefix(auth, "Bearer ") { + w.WriteHeader(http.StatusUnauthorized) + json.NewEncoder(w).Encode(map[string]string{"error": "missing token"}) + return + } + + // Check for simulated auth failure. + if r.URL.Path == "/api/device/connect" { + status := int(ms.connectStatus.Load()) + if status >= 400 { + w.WriteHeader(status) + json.NewEncoder(w).Encode(map[string]string{"error": "auth failed"}) + return + } + } + + w.Header().Set("Content-Type", "application/json") + + switch r.URL.Path { + case "/api/device/connect": + ms.connectCalls.Add(1) + + var body map[string]interface{} + json.NewDecoder(r.Body).Decode(&body) + ms.mu.Lock() + ms.connectBody = body + ms.mu.Unlock() + + json.NewEncoder(w).Encode(map[string]interface{}{ + "type": "welcome", + "protocol_version": 1, + "device_id": 42, + "session_id": "test-session-1", + "reverb": map[string]interface{}{ + "key": reverbCfg.Key, + "host": reverbCfg.Host, + "port": reverbCfg.Port, + "scheme": reverbCfg.Scheme, + }, + }) + + case "/api/device/disconnect": + ms.disconnectCalls.Add(1) + json.NewEncoder(w).Encode(map[string]string{"status": "disconnected"}) + + case "/api/device/heartbeat": + ms.heartbeatCalls.Add(1) + json.NewEncoder(w).Encode(map[string]string{"status": "ok"}) + + case "/api/device/messages": + ms.messagesCalls.Add(1) + + var req struct { + BatchID string `json:"batch_id"` + Messages []json.RawMessage `json:"messages"` + } + json.NewDecoder(r.Body).Decode(&req) + + ms.mu.Lock() + ms.messageBatches = append(ms.messageBatches, mockMessageBatch{ + BatchID: req.BatchID, + Messages: req.Messages, + }) + ms.mu.Unlock() + + json.NewEncoder(w).Encode(map[string]interface{}{ + "accepted": len(req.Messages), + "batch_id": req.BatchID, + "session_id": "test-session-1", + }) + + case "/api/device/broadcasting/auth": + var body struct { + SocketID string `json:"socket_id"` + ChannelName string `json:"channel_name"` + } + json.NewDecoder(r.Body).Decode(&body) + + sig := generateTestAuthSignature( + ms.pusherSrv.appKey, + ms.pusherSrv.appSecret, + body.SocketID, + body.ChannelName, + ) + json.NewEncoder(w).Encode(map[string]string{"auth": sig}) + + default: + http.NotFound(w, r) + } +} + +// getConnectBody returns the last connect request body. +func (ms *mockUplinkServer) getConnectBody() map[string]interface{} { + ms.mu.Lock() + defer ms.mu.Unlock() + return ms.connectBody +} + +// getMessages returns all messages received across all batches, flattened. +func (ms *mockUplinkServer) getMessages() []json.RawMessage { + ms.mu.Lock() + defer ms.mu.Unlock() + var msgs []json.RawMessage + for _, b := range ms.messageBatches { + msgs = append(msgs, b.Messages...) + } + return msgs +} + +// waitForMessageType waits for a message of the given type to arrive. +// Returns the first matching message or an error on timeout. +func (ms *mockUplinkServer) waitForMessageType(msgType string, timeout time.Duration) (json.RawMessage, error) { + deadline := time.After(timeout) + for { + msgs := ms.getMessages() + for _, raw := range msgs { + var envelope struct { + Type string `json:"type"` + } + if json.Unmarshal(raw, &envelope) == nil && envelope.Type == msgType { + return raw, nil + } + } + select { + case <-deadline: + return nil, fmt.Errorf("timeout waiting for message type %q (got %d messages total)", msgType, len(msgs)) + case <-time.After(10 * time.Millisecond): + } + } +} + +// waitForMessages waits until at least n messages have been received. +func (ms *mockUplinkServer) waitForMessages(n int, timeout time.Duration) ([]json.RawMessage, error) { + deadline := time.After(timeout) + for { + msgs := ms.getMessages() + if len(msgs) >= n { + return msgs, nil + } + select { + case <-deadline: + return msgs, fmt.Errorf("timeout waiting for %d messages (got %d)", n, len(msgs)) + case <-time.After(10 * time.Millisecond): + } + } +} + +// sendCommand sends a command to the CLI via the Pusher server. +// Commands are wrapped in a {"type": ..., "payload": {...}} envelope +// to match the real CommandRelayController format. +func (ms *mockUplinkServer) sendCommand(command interface{}) error { + data, err := json.Marshal(command) + if err != nil { + return fmt.Errorf("marshaling command: %w", err) + } + + // Wrap in payload envelope: extract "type", put everything else under "payload". + var flat map[string]json.RawMessage + if err := json.Unmarshal(data, &flat); err == nil { + cmdType := flat["type"] + delete(flat, "type") + + payload, _ := json.Marshal(flat) + wrapped, _ := json.Marshal(map[string]json.RawMessage{ + "type": cmdType, + "payload": payload, + }) + data = wrapped + } + + channel := fmt.Sprintf("private-chief-server.%d", 42) // device ID 42 + return ms.pusherSrv.sendCommand(channel, data) +} + +// waitForPusherSubscribe waits for the CLI to subscribe to its Pusher channel. +func (ms *mockUplinkServer) waitForPusherSubscribe(timeout time.Duration) error { + select { + case <-ms.pusherSrv.onSubscribe: + return nil + case <-time.After(timeout): + return fmt.Errorf("timeout waiting for Pusher subscription") + } +} + +// generateTestAuthSignature generates a Pusher auth signature for testing. +func generateTestAuthSignature(appKey, appSecret, socketID, channelName string) string { + toSign := socketID + ":" + channelName + mac := hmac.New(sha256.New, []byte(appSecret)) + mac.Write([]byte(toSign)) + sig := hex.EncodeToString(mac.Sum(nil)) + return appKey + ":" + sig +} + +// mockPusherServer is a minimal Pusher protocol WebSocket server for testing. +type mockPusherServer struct { + srv *httptest.Server + upgrader websocket.Upgrader + + mu sync.Mutex + conn *websocket.Conn + + appKey string + appSecret string + socketID string + activityTimeout int + + onSubscribe chan string +} + +type mockReverbConfig struct { + Key string `json:"key"` + Host string `json:"host"` + Port int `json:"port"` + Scheme string `json:"scheme"` +} + +func newMockPusherServer(t *testing.T) *mockPusherServer { + t.Helper() + + ps := &mockPusherServer{ + appKey: "test-app-key", + appSecret: "test-app-secret", + socketID: "123456.7890", + activityTimeout: 120, + onSubscribe: make(chan string, 10), + upgrader: websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { return true }, + }, + } + + ps.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ps.handleWS(t, w, r) + })) + + t.Cleanup(func() { + ps.mu.Lock() + if ps.conn != nil { + ps.conn.Close() + } + ps.mu.Unlock() + ps.srv.Close() + }) + + return ps +} + +type pusherMsg struct { + Event string `json:"event"` + Data json.RawMessage `json:"data,omitempty"` + Channel string `json:"channel,omitempty"` +} + +type pusherConnData struct { + SocketID string `json:"socket_id"` + ActivityTimeout int `json:"activity_timeout"` +} + +func (ps *mockPusherServer) handleWS(t *testing.T, w http.ResponseWriter, r *http.Request) { + t.Helper() + + expectedPath := fmt.Sprintf("/app/%s", ps.appKey) + if !strings.HasPrefix(r.URL.Path, expectedPath) { + http.Error(w, "invalid path", http.StatusNotFound) + return + } + + conn, err := ps.upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + + ps.mu.Lock() + ps.conn = conn + ps.mu.Unlock() + + // Send connection_established. + // Real Pusher protocol double-encodes the data: the data field is a JSON string + // containing the connection info, not an embedded object. + connDataInner, _ := json.Marshal(pusherConnData{ + SocketID: ps.socketID, + ActivityTimeout: ps.activityTimeout, + }) + connDataStr, _ := json.Marshal(string(connDataInner)) + conn.WriteJSON(pusherMsg{ + Event: "pusher:connection_established", + Data: connDataStr, + }) + + // Read loop. + for { + _, data, err := conn.ReadMessage() + if err != nil { + return + } + + var msg pusherMsg + if json.Unmarshal(data, &msg) != nil { + continue + } + + switch msg.Event { + case "pusher:subscribe": + var subData map[string]string + json.Unmarshal(msg.Data, &subData) + channel := subData["channel"] + + select { + case ps.onSubscribe <- channel: + default: + } + + conn.WriteJSON(pusherMsg{ + Event: "pusher_internal:subscription_succeeded", + Channel: channel, + Data: json.RawMessage("{}"), + }) + + case "pusher:pong": + // Ignore pong responses. + } + } +} + +// sendCommand sends a chief.command event to the connected client. +func (ps *mockPusherServer) sendCommand(channel string, command json.RawMessage) error { + ps.mu.Lock() + conn := ps.conn + ps.mu.Unlock() + + if conn == nil { + return fmt.Errorf("no client connected") + } + + return conn.WriteJSON(pusherMsg{ + Event: "chief.command", + Channel: channel, + Data: command, + }) +} + +// reverbConfig returns a config pointing at the test server. +func (ps *mockPusherServer) reverbConfig() mockReverbConfig { + addr := ps.srv.Listener.Addr().String() + parts := strings.Split(addr, ":") + host := parts[0] + port := 0 + fmt.Sscanf(parts[1], "%d", &port) + + return mockReverbConfig{ + Key: ps.appKey, + Host: host, + Port: port, + Scheme: "http", + } +} diff --git a/internal/cmd/session.go b/internal/cmd/session.go new file mode 100644 index 00000000..8f24838a --- /dev/null +++ b/internal/cmd/session.go @@ -0,0 +1,783 @@ +package cmd + +import ( + "bufio" + "encoding/json" + "fmt" + "io" + "log" + "os" + "os/exec" + "path/filepath" + "sync" + "time" + + "github.com/creack/pty" + "github.com/minicodemonkey/chief/embed" + "github.com/minicodemonkey/chief/internal/loop" + "github.com/minicodemonkey/chief/internal/prd" + "github.com/minicodemonkey/chief/internal/ws" +) + +// Default session timeout configuration. +const ( + defaultSessionTimeout = 30 * time.Minute +) + +// Default warning thresholds (minutes of inactivity at which to warn). +var defaultWarningThresholds = []int{20, 25, 29} + +// claudeSession tracks a single Claude PRD session process. +type claudeSession struct { + sessionID string + project string + projectPath string + cmd *exec.Cmd + stdin io.WriteCloser + done chan struct{} // closed when the process exits + lastActive time.Time // last time a prd_message was received + activeMu sync.Mutex // protects lastActive +} + +// resetActivity updates the last active time for this session. +func (s *claudeSession) resetActivity() { + s.activeMu.Lock() + s.lastActive = time.Now() + s.activeMu.Unlock() +} + +// inactiveDuration returns how long the session has been inactive. +func (s *claudeSession) inactiveDuration() time.Duration { + s.activeMu.Lock() + defer s.activeMu.Unlock() + return time.Since(s.lastActive) +} + +// sessionManager manages Claude PRD sessions spawned via WebSocket. +type sessionManager struct { + mu sync.RWMutex + sessions map[string]*claudeSession + sender messageSender + timeout time.Duration // session inactivity timeout + warningThresholds []int // minutes of inactivity at which to send warnings + checkInterval time.Duration // how often to check for timeouts (configurable for tests) + stopTimeout chan struct{} // closed to stop the timeout checker +} + +// newSessionManager creates a new session manager. +func newSessionManager(sender messageSender) *sessionManager { + sm := &sessionManager{ + sessions: make(map[string]*claudeSession), + sender: sender, + timeout: defaultSessionTimeout, + warningThresholds: defaultWarningThresholds, + checkInterval: 30 * time.Second, + stopTimeout: make(chan struct{}), + } + go sm.runTimeoutChecker(sm.stopTimeout) + return sm +} + +// sessionCount returns the number of active sessions. +func (sm *sessionManager) sessionCount() int { + sm.mu.RLock() + defer sm.mu.RUnlock() + return len(sm.sessions) +} + +// getSession returns a session by ID, or nil if not found. +func (sm *sessionManager) getSession(sessionID string) *claudeSession { + sm.mu.RLock() + defer sm.mu.RUnlock() + return sm.sessions[sessionID] +} + +// activeSessions returns a list of active session states for state snapshots. +func (sm *sessionManager) activeSessions() []ws.SessionState { + sm.mu.RLock() + defer sm.mu.RUnlock() + var sessions []ws.SessionState + for _, s := range sm.sessions { + sessions = append(sessions, ws.SessionState{ + SessionID: s.sessionID, + Project: s.project, + }) + } + return sessions +} + +// newPRD spawns a new Claude PRD session. +func (sm *sessionManager) newPRD(projectPath, projectName, sessionID, initialMessage string) error { + sm.mu.Lock() + if _, exists := sm.sessions[sessionID]; exists { + sm.mu.Unlock() + return fmt.Errorf("session %s already exists", sessionID) + } + sm.mu.Unlock() + + // Ensure .chief/prds directory structure exists + prdsDir := filepath.Join(projectPath, ".chief", "prds") + if err := os.MkdirAll(prdsDir, 0o755); err != nil { + return fmt.Errorf("failed to create prds directory: %w", err) + } + + // Build prompt from init_prompt.txt template + // Use a temp PRD dir name based on session ID — Claude will create the actual + // directory when it writes prd.md (the init prompt instructs it to). + // We pass the prds base dir so the prompt has the right context. + prompt := embed.GetInitPrompt(prdsDir, initialMessage) + + // Spawn claude in print mode. + // --output-format stream-json streams JSONL events as they are generated. + // We use a PTY for stdin so Claude treats input as a real terminal — without + // a PTY, Claude detects a non-TTY stdin and buffers all output until exit. + cmd := exec.Command(claudeBinary(), "-p", "--dangerously-skip-permissions", "--output-format", "stream-json", "--verbose", prompt) + cmd.Dir = projectPath + cmd.Env = filterEnv(os.Environ(), "CLAUDECODE") + + // Create a PTY pair: ptm (master, used by chief) and pts (slave, used by Claude). + ptm, pts, err := pty.Open() + if err != nil { + return fmt.Errorf("failed to open PTY: %w", err) + } + cmd.Stdin = pts // Claude reads from the slave PTY (looks like a real terminal) + + stdoutPipe, err := cmd.StdoutPipe() + if err != nil { + ptm.Close() + pts.Close() + return fmt.Errorf("failed to create stdout pipe: %w", err) + } + + stderrPipe, err := cmd.StderrPipe() + if err != nil { + ptm.Close() + pts.Close() + return fmt.Errorf("failed to create stderr pipe: %w", err) + } + + log.Printf("[debug] newPRD: launching %s -p --dangerously-skip-permissions --output-format stream-json (PTY stdin) <prompt len=%d> in dir=%s", claudeBinary(), len(prompt), projectPath) + if err := cmd.Start(); err != nil { + ptm.Close() + pts.Close() + return fmt.Errorf("failed to start Claude: %w", err) + } + // Close the slave in the parent after the child has inherited it. + pts.Close() + log.Printf("[debug] newPRD: Claude process started pid=%d session=%s", cmd.Process.Pid, sessionID) + + // Drain PTY master reads (echo of what we write) to prevent buffer backpressure. + go drainPTY(ptm) + + sess := &claudeSession{ + sessionID: sessionID, + project: projectName, + projectPath: projectPath, + cmd: cmd, + stdin: ptm, // Write user messages to PTY master + done: make(chan struct{}), + lastActive: time.Now(), + } + + sm.mu.Lock() + sm.sessions[sessionID] = sess + sm.mu.Unlock() + + // Stream stdout in a goroutine + go sm.streamOutput(sessionID, stdoutPipe) + + // Stream stderr in a goroutine + go sm.streamOutput(sessionID, stderrPipe) + + // Send the user's initial message via the PTY master. + go func() { + time.Sleep(100 * time.Millisecond) + msg := initialMessage + if msg == "" { + msg = "Please start." + } + log.Printf("[debug] newPRD: writing PTY message (%d bytes) for session %s", len(msg), sessionID) + n, err := fmt.Fprintf(ptm, "%s\n", msg) + log.Printf("[debug] newPRD: PTY write complete n=%d err=%v for session %s", n, err, sessionID) + }() + + // Watchdog: log process state every 10s until it exits. + go func() { + for i := 1; i <= 6; i++ { + time.Sleep(10 * time.Second) + select { + case <-sess.done: + return + default: + } + if sess.cmd.ProcessState != nil { + log.Printf("[debug] watchdog session=%s tick=%d: process already exited state=%v", sessionID, i, sess.cmd.ProcessState) + return + } + log.Printf("[debug] watchdog session=%s tick=%d: process pid=%d still running", sessionID, i, sess.cmd.Process.Pid) + } + }() + + // Wait for process to exit + go func() { + err := cmd.Wait() + if err != nil { + log.Printf("Claude session %s exited with error: %v (pid=%d)", sessionID, err, cmd.Process.Pid) + } else { + log.Printf("Claude session %s exited normally (pid=%d)", sessionID, cmd.Process.Pid) + } + + // Send prd_response_complete to signal the PRD session is done + log.Printf("[debug] newPRD: sending prd_response_complete for session %s", sessionID) + completeMsg := ws.PRDResponseCompleteMessage{ + Type: ws.TypePRDResponseComplete, + Payload: ws.PRDResponseCompletePayload{ + SessionID: sessionID, + Project: projectName, + }, + } + if sendErr := sm.sender.Send(completeMsg); sendErr != nil { + log.Printf("Error sending prd_response_complete: %v", sendErr) + } + + // Auto-convert prd.md to prd.json if prd.md was created + sm.autoConvert(projectPath) + + close(sess.done) + + sm.mu.Lock() + delete(sm.sessions, sessionID) + sm.mu.Unlock() + }() + + return nil +} + +// refinePRD spawns a Claude PRD session to edit an existing PRD. +func (sm *sessionManager) refinePRD(projectPath, projectName, sessionID, prdID, message string) error { + sm.mu.Lock() + if _, exists := sm.sessions[sessionID]; exists { + sm.mu.Unlock() + return fmt.Errorf("session %s already exists", sessionID) + } + sm.mu.Unlock() + + // Verify the PRD directory exists + prdDir := filepath.Join(projectPath, ".chief", "prds", prdID) + if _, err := os.Stat(prdDir); os.IsNotExist(err) { + return fmt.Errorf("PRD %q not found in project", prdID) + } + + // Build prompt from edit_prompt.txt template + prompt := embed.GetEditPrompt(prdDir) + + // Spawn claude in print mode with PTY stdin so Claude treats input as a terminal. + // --output-format stream-json streams JSONL events as they are generated. + cmd := exec.Command(claudeBinary(), "-p", "--dangerously-skip-permissions", "--output-format", "stream-json", "--verbose", prompt) + cmd.Dir = projectPath + cmd.Env = filterEnv(os.Environ(), "CLAUDECODE") + + ptm, pts, err := pty.Open() + if err != nil { + return fmt.Errorf("failed to open PTY: %w", err) + } + cmd.Stdin = pts + + stdoutPipe, err := cmd.StdoutPipe() + if err != nil { + ptm.Close() + pts.Close() + return fmt.Errorf("failed to create stdout pipe: %w", err) + } + + stderrPipe, err := cmd.StderrPipe() + if err != nil { + ptm.Close() + pts.Close() + return fmt.Errorf("failed to create stderr pipe: %w", err) + } + + if err := cmd.Start(); err != nil { + ptm.Close() + pts.Close() + return fmt.Errorf("failed to start Claude: %w", err) + } + pts.Close() + + go drainPTY(ptm) + + sess := &claudeSession{ + sessionID: sessionID, + project: projectName, + projectPath: projectPath, + cmd: cmd, + stdin: ptm, + done: make(chan struct{}), + lastActive: time.Now(), + } + + sm.mu.Lock() + sm.sessions[sessionID] = sess + sm.mu.Unlock() + + // Stream stdout in a goroutine + go sm.streamOutput(sessionID, stdoutPipe) + + // Stream stderr in a goroutine + go sm.streamOutput(sessionID, stderrPipe) + + // Send the user's message via the PTY master + go func() { + time.Sleep(100 * time.Millisecond) + fmt.Fprintf(ptm, "%s\n", message) + }() + + // Wait for process to exit + go func() { + err := cmd.Wait() + if err != nil { + log.Printf("Claude session %s exited with error: %v", sessionID, err) + } else { + log.Printf("Claude session %s exited normally", sessionID) + } + + // Send prd_response_complete to signal the PRD session is done + completeMsg := ws.PRDResponseCompleteMessage{ + Type: ws.TypePRDResponseComplete, + Payload: ws.PRDResponseCompletePayload{ + SessionID: sessionID, + Project: projectName, + }, + } + if sendErr := sm.sender.Send(completeMsg); sendErr != nil { + log.Printf("Error sending prd_response_complete: %v", sendErr) + } + + // Auto-convert prd.md to prd.json if prd.md was updated + sm.autoConvert(projectPath) + + close(sess.done) + + sm.mu.Lock() + delete(sm.sessions, sessionID) + sm.mu.Unlock() + }() + + return nil +} + +// streamOutput reads stream-json JSONL from r, extracts assistant text events, +// and forwards them as prd_output messages. Non-text events (tool calls, etc.) are logged. +// Stderr lines are forwarded as-is since they don't contain JSONL. +func (sm *sessionManager) streamOutput(sessionID string, r io.Reader) { + sm.mu.RLock() + sess := sm.sessions[sessionID] + sm.mu.RUnlock() + if sess == nil { + log.Printf("[debug] streamOutput: session %s not found", sessionID) + return + } + + log.Printf("[debug] streamOutput: started for session %s", sessionID) + lineCount := 0 + textCount := 0 + + scanner := bufio.NewScanner(r) + scanner.Buffer(make([]byte, 64*1024), 1024*1024) + for scanner.Scan() { + line := scanner.Text() + lineCount++ + + if lineCount <= 5 { + preview := line + if len(preview) > 120 { + preview = preview[:120] + "..." + } + log.Printf("[debug] streamOutput session=%s line=%d: %q", sessionID, lineCount, preview) + } + + // Try to parse as stream-json and extract assistant text. + event := loop.ParseLine(line) + if event != nil && event.Type == loop.EventAssistantText && event.Text != "" { + textCount++ + outMsg := ws.PRDOutputMessage{ + Type: ws.TypePRDOutput, + Payload: ws.PRDOutputPayload{ + Content: event.Text, + SessionID: sessionID, + Project: sess.project, + }, + } + if sendErr := sm.sender.Send(outMsg); sendErr != nil { + log.Printf("[debug] streamOutput: send error for session %s: %v", sessionID, sendErr) + return + } + } else if event != nil { + log.Printf("[debug] streamOutput session=%s non-text event: %s", sessionID, event.Type.String()) + } else if line != "" { + // ParseLine returned nil — this is either a stream-json event we + // don't handle (hook_started, hook_response, thinking blocks, result, + // rate_limit_event, etc.) or a non-JSON stderr line. Only forward + // non-JSON lines; silently skip unhandled JSON events. + var raw json.RawMessage + if json.Unmarshal([]byte(line), &raw) == nil { + // Valid JSON event we don't need — skip silently. + continue + } + // Non-JSON line (stderr output) — forward as-is. + outMsg := ws.PRDOutputMessage{ + Type: ws.TypePRDOutput, + Payload: ws.PRDOutputPayload{ + Content: line + "\n", + SessionID: sessionID, + Project: sess.project, + }, + } + if sendErr := sm.sender.Send(outMsg); sendErr != nil { + log.Printf("[debug] streamOutput: send error for session %s: %v", sessionID, sendErr) + return + } + } + } + + if err := scanner.Err(); err != nil { + log.Printf("[debug] streamOutput: scanner error for session %s after %d lines: %v", sessionID, lineCount, err) + } else { + log.Printf("[debug] streamOutput: EOF for session %s after %d lines (%d text events)", sessionID, lineCount, textCount) + } +} + +// sendMessage writes a user message to an active session's stdin. +func (sm *sessionManager) sendMessage(sessionID, content string) error { + sess := sm.getSession(sessionID) + if sess == nil { + return fmt.Errorf("session not found") + } + + // Reset the inactivity timer + sess.resetActivity() + + // Write the message followed by a newline to the Claude process stdin + _, err := fmt.Fprintf(sess.stdin, "%s\n", content) + if err != nil { + return fmt.Errorf("failed to write to Claude stdin: %w", err) + } + return nil +} + +// closeSession closes a PRD session. If save is true, waits for Claude to finish. +// If save is false, kills immediately. +func (sm *sessionManager) closeSession(sessionID string, save bool) error { + sess := sm.getSession(sessionID) + if sess == nil { + return fmt.Errorf("session not found") + } + + if save { + // Close stdin to signal EOF to Claude, then wait for it to finish + sess.stdin.Close() + <-sess.done + } else { + // Kill immediately + if sess.cmd.Process != nil { + sess.cmd.Process.Kill() + } + <-sess.done + } + + return nil +} + +// killAll kills all active sessions (used during shutdown). +func (sm *sessionManager) killAll() { + // Stop the timeout checker + select { + case <-sm.stopTimeout: + // Already closed + default: + close(sm.stopTimeout) + } + + sm.mu.RLock() + sessions := make([]*claudeSession, 0, len(sm.sessions)) + for _, s := range sm.sessions { + sessions = append(sessions, s) + } + sm.mu.RUnlock() + + for _, s := range sessions { + if s.cmd.Process != nil { + s.cmd.Process.Kill() + } + } + + // Wait for all to finish + for _, s := range sessions { + <-s.done + } +} + +// runTimeoutChecker periodically checks all sessions for inactivity and sends +// warnings at the configured thresholds. When the timeout is reached, the session +// is expired: state is saved to disk, the process is killed, and session_expired is sent. +func (sm *sessionManager) runTimeoutChecker(stopCh <-chan struct{}) { + ticker := time.NewTicker(sm.checkInterval) + defer ticker.Stop() + + // Track which warnings have been sent for each session to avoid duplicates. + // Key: sessionID, Value: set of warning minutes already sent. + sentWarnings := make(map[string]map[int]bool) + + for { + select { + case <-stopCh: + return + case <-ticker.C: + sm.mu.RLock() + sessions := make([]*claudeSession, 0, len(sm.sessions)) + for _, s := range sm.sessions { + sessions = append(sessions, s) + } + sm.mu.RUnlock() + + for _, sess := range sessions { + inactive := sess.inactiveDuration() + inactiveMinutes := int(inactive.Minutes()) + + // Check if session should be expired + if inactive >= sm.timeout { + log.Printf("Session %s timed out after %v of inactivity", sess.sessionID, sm.timeout) + sm.expireSession(sess) + delete(sentWarnings, sess.sessionID) + continue + } + + // Check warning thresholds + if _, ok := sentWarnings[sess.sessionID]; !ok { + sentWarnings[sess.sessionID] = make(map[int]bool) + } + + for _, threshold := range sm.warningThresholds { + if inactiveMinutes >= threshold && !sentWarnings[sess.sessionID][threshold] { + timeoutMinutes := int(sm.timeout.Minutes()) + remaining := timeoutMinutes - threshold + log.Printf("Session %s: sending timeout warning (%d minutes remaining)", sess.sessionID, remaining) + sm.sendTimeoutWarning(sess.sessionID, remaining) + sentWarnings[sess.sessionID][threshold] = true + } + } + } + + // Clean up sentWarnings for sessions that no longer exist + sm.mu.RLock() + for sid := range sentWarnings { + if _, exists := sm.sessions[sid]; !exists { + delete(sentWarnings, sid) + } + } + sm.mu.RUnlock() + } + } +} + +// sendTimeoutWarning sends a session_timeout_warning message over WebSocket. +func (sm *sessionManager) sendTimeoutWarning(sessionID string, minutesRemaining int) { + envelope := ws.NewMessage(ws.TypeSessionTimeoutWarning) + msg := ws.SessionTimeoutWarningMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + SessionID: sessionID, + MinutesRemaining: minutesRemaining, + } + if err := sm.sender.Send(msg); err != nil { + log.Printf("Error sending session_timeout_warning: %v", err) + } +} + +// expireSession saves whatever PRD state exists, kills the Claude process, +// and sends a session_expired message. +func (sm *sessionManager) expireSession(sess *claudeSession) { + // Close stdin to let Claude finish writing, then kill after a brief grace period + sess.stdin.Close() + + // Give Claude 2 seconds to finish writing + select { + case <-sess.done: + // Process exited cleanly + case <-time.After(2 * time.Second): + // Force kill + if sess.cmd.Process != nil { + sess.cmd.Process.Kill() + } + <-sess.done + } + + // Send session_expired message + envelope := ws.NewMessage(ws.TypeSessionExpired) + expiredMsg := ws.SessionExpiredMessage{ + Type: envelope.Type, + ID: envelope.ID, + Timestamp: envelope.Timestamp, + SessionID: sess.sessionID, + } + if err := sm.sender.Send(expiredMsg); err != nil { + log.Printf("Error sending session_expired: %v", err) + } + + log.Printf("Session %s expired and cleaned up", sess.sessionID) +} + +// autoConvert scans for any prd.md files that need conversion and converts them. +func (sm *sessionManager) autoConvert(projectPath string) { + prdsDir := filepath.Join(projectPath, ".chief", "prds") + entries, err := os.ReadDir(prdsDir) + if err != nil { + return + } + + for _, entry := range entries { + if !entry.IsDir() { + continue + } + prdDir := filepath.Join(prdsDir, entry.Name()) + needs, err := prd.NeedsConversion(prdDir) + if err != nil { + log.Printf("Error checking conversion for %s: %v", prdDir, err) + continue + } + if needs { + log.Printf("Auto-converting PRD in %s", prdDir) + if err := prd.Convert(prd.ConvertOptions{PRDDir: prdDir}); err != nil { + log.Printf("Auto-conversion failed for %s: %v", prdDir, err) + } else { + log.Printf("Auto-conversion succeeded for %s", prdDir) + } + } + } +} + +// handleNewPRD handles a new_prd WebSocket message. +func handleNewPRD(sender messageSender, scanner projectFinder, sessions *sessionManager, msg ws.Message) { + var req ws.NewPRDMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing new_prd message: %v", err) + return + } + + project, found := scanner.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + if err := sessions.newPRD(project.Path, req.Project, req.SessionID, req.Message); err != nil { + sendError(sender, ws.ErrCodeClaudeError, + fmt.Sprintf("Failed to start Claude session: %v", err), msg.ID) + return + } + + log.Printf("Started Claude PRD session %s for project %s", req.SessionID, req.Project) +} + +// handleRefinePRD handles a refine_prd WebSocket message. +func handleRefinePRD(sender messageSender, scanner projectFinder, sessions *sessionManager, msg ws.Message) { + var req ws.RefinePRDMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing refine_prd message: %v", err) + return + } + + project, found := scanner.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + if err := sessions.refinePRD(project.Path, req.Project, req.SessionID, req.PRDID, req.Message); err != nil { + sendError(sender, ws.ErrCodeClaudeError, + fmt.Sprintf("Failed to start Claude session: %v", err), msg.ID) + return + } + + log.Printf("Started Claude PRD refine session %s for project %s (prd: %s)", req.SessionID, req.Project, req.PRDID) +} + +// handlePRDMessage handles a prd_message WebSocket message. +func handlePRDMessage(sender messageSender, sessions *sessionManager, msg ws.Message) { + var req ws.PRDMessageMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing prd_message: %v", err) + return + } + + if err := sessions.sendMessage(req.SessionID, req.Message); err != nil { + sendError(sender, ws.ErrCodeSessionNotFound, + fmt.Sprintf("Session %q not found", req.SessionID), msg.ID) + return + } +} + +// handleClosePRDSession handles a close_prd_session WebSocket message. +func handleClosePRDSession(sender messageSender, sessions *sessionManager, msg ws.Message) { + var req ws.ClosePRDSessionMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing close_prd_session: %v", err) + return + } + + if err := sessions.closeSession(req.SessionID, req.Save); err != nil { + sendError(sender, ws.ErrCodeSessionNotFound, + fmt.Sprintf("Session %q not found", req.SessionID), msg.ID) + return + } + + log.Printf("Closed Claude PRD session %s (save=%v)", req.SessionID, req.Save) +} + +// drainPTY reads and discards data from a PTY master file descriptor. +// This prevents the PTY echo buffer from filling up when we write user messages. +// Runs until the PTY is closed (typically when the child process exits). +func drainPTY(f *os.File) { + buf := make([]byte, 256) + for { + _, err := f.Read(buf) + if err != nil { + return + } + } +} + +// claudeBinary returns the path to the claude CLI binary. +// It checks the CHIEF_CLAUDE_BINARY environment variable first, falling back to "claude". +func claudeBinary() string { + if bin := os.Getenv("CHIEF_CLAUDE_BINARY"); bin != "" { + return bin + } + return "claude" +} + +// filterEnv returns a copy of env with the named variables removed. +func filterEnv(env []string, keys ...string) []string { + filtered := make([]string, 0, len(env)) + for _, e := range env { + skip := false + for _, key := range keys { + if len(e) > len(key) && e[:len(key)+1] == key+"=" { + skip = true + break + } + } + if !skip { + filtered = append(filtered, e) + } + } + return filtered +} + +// projectFinder is an interface for finding projects (for testability). +type projectFinder interface { + FindProject(name string) (ws.ProjectSummary, bool) +} diff --git a/internal/cmd/session_test.go b/internal/cmd/session_test.go new file mode 100644 index 00000000..afa31426 --- /dev/null +++ b/internal/cmd/session_test.go @@ -0,0 +1,1469 @@ +package cmd + +import ( + "context" + "encoding/json" + "fmt" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/ws" +) + +// captureSender is a mock messageSender that captures sent messages. +type captureSender struct { + mu sync.Mutex + messages []map[string]interface{} +} + +func (c *captureSender) Send(msg interface{}) error { + data, err := json.Marshal(msg) + if err != nil { + return err + } + var m map[string]interface{} + if err := json.Unmarshal(data, &m); err != nil { + return err + } + c.mu.Lock() + c.messages = append(c.messages, m) + c.mu.Unlock() + return nil +} + +func (c *captureSender) getMessages() []map[string]interface{} { + c.mu.Lock() + defer c.mu.Unlock() + cp := make([]map[string]interface{}, len(c.messages)) + copy(cp, c.messages) + return cp +} + +// discardSender is a mock messageSender that discards all messages. +type discardSender struct{} + +func (d *discardSender) Send(msg interface{}) error { return nil } + +// mockProjectFinder implements projectFinder for tests. +type mockProjectFinder struct { + projects map[string]ws.ProjectSummary +} + +func (m *mockProjectFinder) FindProject(name string) (ws.ProjectSummary, bool) { + p, ok := m.projects[name] + return p, ok +} + +func TestSessionManager_NewPRD(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send new_prd request + newPRDReq := map[string]string{ + "type": "new_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "session_id": "sess-123", + "message": "Build a todo app", + } + ms.sendCommand(newPRDReq) + + // We should receive prd_output messages. + // Since we can't actually run claude in tests, expect an error response + // (claude binary not available in test) — this tests the error path. + // Wait for first message with 3 second timeout. + raw, err := ms.waitForMessageType("error", 3*time.Second) + if err != nil { + // If not an error, might be prd_output + raw, err = ms.waitForMessageType("prd_output", 3*time.Second) + if err != nil { + t.Fatal("expected error or prd_output message") + } + } + + var msg map[string]interface{} + if err := json.Unmarshal(raw, &msg); err != nil { + t.Fatalf("failed to unmarshal response: %v", err) + } + + // Check that we got some kind of response + msgType := msg["type"].(string) + if msgType != "error" && msgType != "prd_output" { + t.Errorf("expected error or prd_output message, got %s", msgType) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestSessionManager_NewPRD_ProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send new_prd for nonexistent project + newPRDReq := map[string]string{ + "type": "new_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "session_id": "sess-123", + "message": "Build a todo app", + } + ms.sendCommand(newPRDReq) + + // Read error response + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err != nil { + t.Fatalf("expected error message: %v", err) + } + + var errorReceived map[string]interface{} + if err := json.Unmarshal(raw, &errorReceived); err != nil { + t.Fatalf("failed to unmarshal error: %v", err) + } + + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", errorReceived["code"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestSessionManager_PRDMessage_SessionNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send prd_message for nonexistent session + prdMsg := map[string]string{ + "type": "prd_message", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "session_id": "nonexistent-session", + "message": "hello", + } + ms.sendCommand(prdMsg) + + // Read error response + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err != nil { + t.Fatalf("expected error message: %v", err) + } + + var errorReceived map[string]interface{} + if err := json.Unmarshal(raw, &errorReceived); err != nil { + t.Fatalf("failed to unmarshal error: %v", err) + } + + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "SESSION_NOT_FOUND" { + t.Errorf("expected code 'SESSION_NOT_FOUND', got %v", errorReceived["code"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestSessionManager_ClosePRDSession_SessionNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send close_prd_session for nonexistent session + closeMsg := map[string]interface{}{ + "type": "close_prd_session", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "session_id": "nonexistent-session", + "save": false, + } + ms.sendCommand(closeMsg) + + // Read error response + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err != nil { + t.Fatalf("expected error message: %v", err) + } + + var errorReceived map[string]interface{} + if err := json.Unmarshal(raw, &errorReceived); err != nil { + t.Fatalf("failed to unmarshal error: %v", err) + } + + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "SESSION_NOT_FOUND" { + t.Errorf("expected code 'SESSION_NOT_FOUND', got %v", errorReceived["code"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +// TestSessionManager_WithMockClaude uses a shell script to simulate Claude, +// testing the full session lifecycle: spawn, stream output, send message, close. +func TestSessionManager_WithMockClaude(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + // Create a mock "claude" script that echoes input + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +echo "Claude PRD session started" +echo "Processing: $1" +# Read from stdin and echo back +while IFS= read -r line; do + echo "Received: $line" +done +echo "Session complete" +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + + // Add mock claude to PATH + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for initial state_snapshot + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send new_prd request via Pusher + newPRDReq := map[string]string{ + "type": "new_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "session_id": "sess-mock-1", + "message": "Build a todo app", + } + ms.sendCommand(newPRDReq) + + // Wait a bit for process to start and produce output + time.Sleep(500 * time.Millisecond) + + // Send a prd_message + prdMsg := map[string]string{ + "type": "prd_message", + "id": "req-2", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "session_id": "sess-mock-1", + "message": "Add user authentication", + } + ms.sendCommand(prdMsg) + + // Wait for output + time.Sleep(500 * time.Millisecond) + + // Close the session (save=false, kill immediately) + closeMsg := map[string]interface{}{ + "type": "close_prd_session", + "id": "req-3", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "session_id": "sess-mock-1", + "save": false, + } + ms.sendCommand(closeMsg) + + // Wait for a prd_response_complete message + deadline := time.After(5 * time.Second) + for { + msgs := ms.getMessages() + for _, raw := range msgs { + var msg map[string]interface{} + json.Unmarshal(raw, &msg) + if msg["type"] == "prd_response_complete" { + cancel() + return + } + } + select { + case <-deadline: + cancel() + return + case <-time.After(50 * time.Millisecond): + } + } + }() + + err := RunServe(ServeOptions{ + Workspace: workspaceDir, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Collect all prd_output messages + allMsgs := ms.getMessages() + var prdOutputs []map[string]interface{} + for _, raw := range allMsgs { + var msg map[string]interface{} + if json.Unmarshal(raw, &msg) == nil && msg["type"] == "prd_output" { + prdOutputs = append(prdOutputs, msg) + } + } + + if len(prdOutputs) == 0 { + t.Fatal("expected at least one prd_output message") + } + + // Verify session_id and project are set on all prd_output messages (inside payload) + for _, co := range prdOutputs { + payload, _ := co["payload"].(map[string]interface{}) + if payload == nil { + t.Error("expected prd_output to have a payload field") + continue + } + if payload["session_id"] != "sess-mock-1" { + t.Errorf("expected payload.session_id 'sess-mock-1', got %v", payload["session_id"]) + } + if payload["project"] != "myproject" { + t.Errorf("expected payload.project 'myproject', got %v", payload["project"]) + } + } + + // Verify we got a prd_response_complete message + hasComplete := false + for _, raw := range allMsgs { + var msg map[string]interface{} + if json.Unmarshal(raw, &msg) == nil && msg["type"] == "prd_response_complete" { + hasComplete = true + payload, _ := msg["payload"].(map[string]interface{}) + if payload == nil { + t.Error("expected prd_response_complete to have a payload field") + } else if payload["session_id"] != "sess-mock-1" { + t.Errorf("expected payload.session_id 'sess-mock-1' on prd_response_complete, got %v", payload["session_id"]) + } + break + } + } + if !hasComplete { + t.Error("expected a prd_response_complete message") + } + + // Verify we received some actual content + hasContent := false + for _, co := range prdOutputs { + payload, _ := co["payload"].(map[string]interface{}) + if payload != nil { + if content, ok := payload["content"].(string); ok && strings.TrimSpace(content) != "" { + hasContent = true + break + } + } + } + if !hasContent { + t.Error("expected at least one prd_output with non-empty content") + } +} + +// TestSessionManager_WithMockClaude_SaveClose tests save=true close behavior. +func TestSessionManager_WithMockClaude_SaveClose(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + // Create a mock "claude" script that exits on EOF + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +echo "Session started" +# Read until EOF (stdin closed) +while IFS= read -r line; do + echo "Got: $line" +done +echo "Saving PRD..." +exit 0 +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + // Wait for initial state_snapshot + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send new_prd via Pusher + newPRDReq := map[string]string{ + "type": "new_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "session_id": "sess-save-1", + "message": "Build an API", + } + ms.sendCommand(newPRDReq) + + time.Sleep(500 * time.Millisecond) + + // Close with save=true (waits for Claude to finish) + closeMsg := map[string]interface{}{ + "type": "close_prd_session", + "id": "req-2", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "session_id": "sess-save-1", + "save": true, + } + ms.sendCommand(closeMsg) + + // Wait for a prd_response_complete message + deadline := time.After(5 * time.Second) + for { + msgs := ms.getMessages() + for _, raw := range msgs { + var msg map[string]interface{} + json.Unmarshal(raw, &msg) + if msg["type"] == "prd_response_complete" { + cancel() + return + } + } + select { + case <-deadline: + cancel() + return + case <-time.After(50 * time.Millisecond): + } + } + }() + + err := RunServe(ServeOptions{ + Workspace: workspaceDir, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Verify we received a prd_response_complete message + allMsgs := ms.getMessages() + hasComplete := false + for _, raw := range allMsgs { + var msg map[string]interface{} + if json.Unmarshal(raw, &msg) == nil && msg["type"] == "prd_response_complete" { + hasComplete = true + break + } + } + if !hasComplete { + t.Error("expected a prd_response_complete message after save close") + } +} + +func TestSessionManager_ActiveSessions(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Create a mock "claude" script that stays alive + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +while IFS= read -r line; do + echo "$line" +done +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sm := newSessionManager(&discardSender{}) + + // Initially no active sessions + sessions := sm.activeSessions() + if len(sessions) != 0 { + t.Errorf("expected 0 active sessions, got %d", len(sessions)) + } + + // Create a project dir for the session + projectDir := filepath.Join(home, "testproject") + if err := os.MkdirAll(projectDir, 0o755); err != nil { + t.Fatal(err) + } + + // Start a session + err := sm.newPRD(projectDir, "testproject", "sess-1", "test message") + if err != nil { + t.Fatalf("newPRD failed: %v", err) + } + + // Now should have 1 active session + sessions = sm.activeSessions() + if len(sessions) != 1 { + t.Fatalf("expected 1 active session, got %d", len(sessions)) + } + if sessions[0].SessionID != "sess-1" { + t.Errorf("expected session_id 'sess-1', got %q", sessions[0].SessionID) + } + if sessions[0].Project != "testproject" { + t.Errorf("expected project 'testproject', got %q", sessions[0].Project) + } + + // Kill all sessions + sm.killAll() + + // Now should have 0 active sessions + sessions = sm.activeSessions() + if len(sessions) != 0 { + t.Errorf("expected 0 active sessions after killAll, got %d", len(sessions)) + } +} + +func TestSessionManager_SendMessage(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Create a mock "claude" script that echoes input + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +echo "ready" +while IFS= read -r line; do + echo "echo: $line" +done +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sender := &captureSender{} + sm := newSessionManager(sender) + + projectDir := filepath.Join(home, "testproject") + if err := os.MkdirAll(projectDir, 0o755); err != nil { + t.Fatal(err) + } + + err := sm.newPRD(projectDir, "testproject", "sess-msg-1", "test") + if err != nil { + t.Fatalf("newPRD failed: %v", err) + } + + // Wait for process to start + time.Sleep(300 * time.Millisecond) + + // Send a message + if err := sm.sendMessage("sess-msg-1", "hello world"); err != nil { + t.Fatalf("sendMessage failed: %v", err) + } + + // Wait for echo + time.Sleep(500 * time.Millisecond) + + // Verify error on nonexistent session + if err := sm.sendMessage("nonexistent", "test"); err == nil { + t.Error("expected error for nonexistent session") + } + + // Check that we received the echoed message via captureSender + msgs := sender.getMessages() + hasEcho := false + for _, msg := range msgs { + if msg["type"] == "prd_output" { + payload, _ := msg["payload"].(map[string]interface{}) + if payload != nil { + if content, ok := payload["content"].(string); ok && strings.Contains(content, "echo: hello world") { + hasEcho = true + break + } + } + } + } + if !hasEcho { + t.Errorf("expected echoed message 'echo: hello world' in captured messages: %v", msgs) + } + + // Clean up + sm.killAll() +} + +func TestSessionManager_CloseSession_Errors(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + sm := newSessionManager(&discardSender{}) + + // Close nonexistent session + err := sm.closeSession("nonexistent", false) + if err == nil { + t.Error("expected error for nonexistent session") + } + if !strings.Contains(err.Error(), "session not found") { + t.Errorf("expected 'session not found' error, got: %v", err) + } +} + +func TestSessionManager_DuplicateSession(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Create a mock "claude" script + mockClaudeBin := filepath.Join(home, "claude") + mockScript := fmt.Sprintf("#!/bin/sh\nwhile IFS= read -r line; do echo \"$line\"; done") + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sm := newSessionManager(&discardSender{}) + + projectDir := filepath.Join(home, "testproject") + if err := os.MkdirAll(projectDir, 0o755); err != nil { + t.Fatal(err) + } + + // Start first session + err := sm.newPRD(projectDir, "testproject", "sess-dup", "test") + if err != nil { + t.Fatalf("first newPRD failed: %v", err) + } + + // Try to start duplicate session + err = sm.newPRD(projectDir, "testproject", "sess-dup", "test") + if err == nil { + t.Error("expected error for duplicate session_id") + } + if !strings.Contains(err.Error(), "already exists") { + t.Errorf("expected 'already exists' error, got: %v", err) + } + + sm.killAll() +} + +func TestSessionManager_RefinePRD(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + // Create existing PRD directory with prd.md + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature-auth") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(prdDir, "prd.md"), []byte("# Auth PRD\n"), 0o644); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Send refine_prd request + refinePRDReq := map[string]string{ + "type": "refine_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "session_id": "sess-refine-1", + "prd_id": "feature-auth", + "message": "Add OAuth support", + } + ms.sendCommand(refinePRDReq) + + // Since we can't actually run claude in tests, expect an error response + // (claude binary not available in test) — this tests the error path. + raw, err := ms.waitForMessageType("error", 3*time.Second) + if err != nil { + raw, err = ms.waitForMessageType("prd_output", 3*time.Second) + if err != nil { + t.Fatal("expected error or prd_output message") + } + } + + var msg map[string]interface{} + if err := json.Unmarshal(raw, &msg); err != nil { + t.Fatalf("failed to unmarshal response: %v", err) + } + + msgType := msg["type"].(string) + if msgType != "error" && msgType != "prd_output" { + t.Errorf("expected error or prd_output message, got %s", msgType) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestSessionManager_RefinePRD_ProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + refinePRDReq := map[string]string{ + "type": "refine_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "session_id": "sess-refine-1", + "prd_id": "feature-auth", + "message": "Add OAuth support", + } + ms.sendCommand(refinePRDReq) + + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err != nil { + t.Fatalf("expected error message: %v", err) + } + + var errorReceived map[string]interface{} + if err := json.Unmarshal(raw, &errorReceived); err != nil { + t.Fatalf("failed to unmarshal error: %v", err) + } + + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", errorReceived["code"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestSessionManager_RefinePRD_PRDNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + refinePRDReq := map[string]string{ + "type": "refine_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "session_id": "sess-refine-1", + "prd_id": "nonexistent-prd", + "message": "Add OAuth support", + } + ms.sendCommand(refinePRDReq) + + raw, err := ms.waitForMessageType("error", 2*time.Second) + if err != nil { + t.Fatalf("expected error message: %v", err) + } + + var errorReceived map[string]interface{} + if err := json.Unmarshal(raw, &errorReceived); err != nil { + t.Fatalf("failed to unmarshal error: %v", err) + } + + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "CLAUDE_ERROR" { + t.Errorf("expected code 'CLAUDE_ERROR', got %v", errorReceived["code"]) + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } +} + +func TestSessionManager_WithMockClaude_RefinePRD(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + // Create existing PRD directory with prd.md + prdDir := filepath.Join(projectDir, ".chief", "prds", "feature-auth") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(prdDir, "prd.md"), []byte("# Auth PRD\n"), 0o644); err != nil { + t.Fatal(err) + } + + // Create a mock "claude" script that echoes input + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +echo "Claude PRD edit session started" +echo "Processing: $1" +while IFS= read -r line; do + echo "Received: $line" +done +echo "Edit complete" +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + ctx, cancel := context.WithCancel(context.Background()) + ms := newMockUplinkServer(t) + + go func() { + if err := ms.waitForPusherSubscribe(10 * time.Second); err != nil { + t.Logf("waitForPusherSubscribe: %v", err) + cancel() + return + } + + if _, err := ms.waitForMessageType("state_snapshot", 5*time.Second); err != nil { + t.Logf("waitForMessageType(state_snapshot): %v", err) + cancel() + return + } + + // Send refine_prd request via Pusher + refinePRDReq := map[string]string{ + "type": "refine_prd", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "session_id": "sess-refine-mock-1", + "prd_id": "feature-auth", + "message": "Add OAuth support", + } + ms.sendCommand(refinePRDReq) + + // Wait for output + time.Sleep(500 * time.Millisecond) + + // Send a follow-up message + prdMsg := map[string]string{ + "type": "prd_message", + "id": "req-2", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "session_id": "sess-refine-mock-1", + "message": "Also add RBAC", + } + ms.sendCommand(prdMsg) + + // Wait for output + time.Sleep(500 * time.Millisecond) + + // Close the session + closeMsg := map[string]interface{}{ + "type": "close_prd_session", + "id": "req-3", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "session_id": "sess-refine-mock-1", + "save": false, + } + ms.sendCommand(closeMsg) + + // Wait for prd_response_complete + deadline := time.After(5 * time.Second) + for { + msgs := ms.getMessages() + for _, raw := range msgs { + var msg map[string]interface{} + json.Unmarshal(raw, &msg) + if msg["type"] == "prd_response_complete" { + cancel() + return + } + } + select { + case <-deadline: + cancel() + return + case <-time.After(50 * time.Millisecond): + } + } + }() + + err := RunServe(ServeOptions{ + Workspace: workspaceDir, + ServerURL: ms.httpSrv.URL, + Version: "1.0.0", + Ctx: ctx, + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + // Collect all prd_output messages + allMsgs := ms.getMessages() + var prdOutputs []map[string]interface{} + for _, raw := range allMsgs { + var msg map[string]interface{} + if json.Unmarshal(raw, &msg) == nil && msg["type"] == "prd_output" { + prdOutputs = append(prdOutputs, msg) + } + } + + if len(prdOutputs) == 0 { + t.Fatal("expected at least one prd_output message") + } + + // Verify session_id and project are set on all prd_output messages (inside payload) + for _, co := range prdOutputs { + payload, _ := co["payload"].(map[string]interface{}) + if payload == nil { + t.Error("expected prd_output to have a payload field") + continue + } + if payload["session_id"] != "sess-refine-mock-1" { + t.Errorf("expected payload.session_id 'sess-refine-mock-1', got %v", payload["session_id"]) + } + if payload["project"] != "myproject" { + t.Errorf("expected payload.project 'myproject', got %v", payload["project"]) + } + } + + // Verify we got a prd_response_complete message + hasComplete := false + for _, raw := range allMsgs { + var msg map[string]interface{} + if json.Unmarshal(raw, &msg) == nil && msg["type"] == "prd_response_complete" { + hasComplete = true + payload, _ := msg["payload"].(map[string]interface{}) + if payload == nil { + t.Error("expected prd_response_complete to have a payload field") + } else if payload["session_id"] != "sess-refine-mock-1" { + t.Errorf("expected payload.session_id 'sess-refine-mock-1' on prd_response_complete, got %v", payload["session_id"]) + } + break + } + } + if !hasComplete { + t.Error("expected a prd_response_complete message") + } + + // Verify the user's message was received by Claude (should appear in prd_output) + hasUserMessage := false + for _, co := range prdOutputs { + payload, _ := co["payload"].(map[string]interface{}) + if payload != nil { + if content, ok := payload["content"].(string); ok && strings.Contains(content, "Add OAuth support") { + hasUserMessage = true + break + } + } + } + if !hasUserMessage { + t.Error("expected user's refine message 'Add OAuth support' to appear in prd_output") + } +} + +// newTestSessionManager creates a session manager with configurable timeouts for testing. +// It does NOT start the timeout checker goroutine automatically. +func newTestSessionManager(t *testing.T, timeout time.Duration, warningThresholds []int, checkInterval time.Duration) (*sessionManager, *captureSender, func()) { + t.Helper() + + home := t.TempDir() + setTestHome(t, home) + + // Create a mock "claude" script that stays alive + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +while IFS= read -r line; do + echo "$line" +done +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sender := &captureSender{} + sm := &sessionManager{ + sessions: make(map[string]*claudeSession), + sender: sender, + timeout: timeout, + warningThresholds: warningThresholds, + checkInterval: checkInterval, + stopTimeout: make(chan struct{}), + } + + cleanup := func() { + sm.killAll() + } + + return sm, sender, cleanup +} + +func TestSessionManager_TimeoutExpiration(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Create mock claude + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +while IFS= read -r line; do + echo "$line" +done +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sender := &captureSender{} + sm := &sessionManager{ + sessions: make(map[string]*claudeSession), + sender: sender, + timeout: 200 * time.Millisecond, // Very short for testing + warningThresholds: []int{}, // No warnings, just test expiry + checkInterval: 50 * time.Millisecond, // Check frequently + stopTimeout: make(chan struct{}), + } + go sm.runTimeoutChecker(sm.stopTimeout) + + projectDir := filepath.Join(home, "testproject") + if err := os.MkdirAll(projectDir, 0o755); err != nil { + t.Fatal(err) + } + + err := sm.newPRD(projectDir, "testproject", "sess-timeout-1", "test") + if err != nil { + t.Fatalf("newPRD failed: %v", err) + } + + // Session should be active + if len(sm.activeSessions()) != 1 { + t.Fatal("expected 1 active session") + } + + // Wait for timeout to expire + some buffer + time.Sleep(500 * time.Millisecond) + + // Session should be expired and removed + if len(sm.activeSessions()) != 0 { + t.Errorf("expected 0 active sessions after timeout, got %d", len(sm.activeSessions())) + } + + // Check that session_expired message was sent + msgs := sender.getMessages() + hasExpired := false + for _, msg := range msgs { + if msg["type"] == "session_expired" && msg["session_id"] == "sess-timeout-1" { + hasExpired = true + break + } + } + if !hasExpired { + t.Error("expected session_expired message to be sent") + } + + sm.killAll() +} + +func TestSessionManager_TimeoutWarnings(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Create mock claude + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +while IFS= read -r line; do + echo "$line" +done +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sender := &captureSender{} + + // Use a 3-minute timeout with thresholds at 1 and 2 minutes. + // We simulate time by setting lastActive in the past. + sm := &sessionManager{ + sessions: make(map[string]*claudeSession), + sender: sender, + timeout: 3 * time.Minute, + warningThresholds: []int{1, 2}, // Warn at 1min and 2min of inactivity + checkInterval: 50 * time.Millisecond, // Check frequently + stopTimeout: make(chan struct{}), + } + go sm.runTimeoutChecker(sm.stopTimeout) + + projectDir := filepath.Join(home, "testproject") + if err := os.MkdirAll(projectDir, 0o755); err != nil { + t.Fatal(err) + } + + err := sm.newPRD(projectDir, "testproject", "sess-warn-1", "test") + if err != nil { + t.Fatalf("newPRD failed: %v", err) + } + + // Simulate 90 seconds of inactivity by setting lastActive in the past + sess := sm.getSession("sess-warn-1") + if sess == nil { + t.Fatal("session not found") + } + sess.activeMu.Lock() + sess.lastActive = time.Now().Add(-90 * time.Second) + sess.activeMu.Unlock() + + // Wait for the checker to pick it up + time.Sleep(200 * time.Millisecond) + + // Should have the 1-minute warning (2 remaining) + msgs := sender.getMessages() + var warningMessages []map[string]interface{} + for _, msg := range msgs { + if msg["type"] == "session_timeout_warning" { + warningMessages = append(warningMessages, msg) + } + } + + if len(warningMessages) != 1 { + t.Fatalf("expected 1 warning message, got %d", len(warningMessages)) + } + + // The warning at 1 min means 3-1 = 2 minutes remaining + if warningMessages[0]["minutes_remaining"] != float64(2) { + t.Errorf("expected minutes_remaining=2, got %v", warningMessages[0]["minutes_remaining"]) + } + + // Now simulate 2.5 minutes of inactivity + sess.activeMu.Lock() + sess.lastActive = time.Now().Add(-150 * time.Second) + sess.activeMu.Unlock() + + time.Sleep(200 * time.Millisecond) + + msgs = sender.getMessages() + warningMessages = nil + for _, msg := range msgs { + if msg["type"] == "session_timeout_warning" { + warningMessages = append(warningMessages, msg) + } + } + + // Should now have 2 warnings (1 min and 2 min thresholds) + if len(warningMessages) != 2 { + t.Fatalf("expected 2 warning messages, got %d", len(warningMessages)) + } + + sm.killAll() +} + +func TestSessionManager_TimeoutResetOnMessage(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Create mock claude + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +while IFS= read -r line; do + echo "$line" +done +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sm := &sessionManager{ + sessions: make(map[string]*claudeSession), + sender: &discardSender{}, + timeout: 300 * time.Millisecond, + warningThresholds: []int{}, + checkInterval: 50 * time.Millisecond, + stopTimeout: make(chan struct{}), + } + go sm.runTimeoutChecker(sm.stopTimeout) + + projectDir := filepath.Join(home, "testproject") + if err := os.MkdirAll(projectDir, 0o755); err != nil { + t.Fatal(err) + } + + err := sm.newPRD(projectDir, "testproject", "sess-reset-1", "test") + if err != nil { + t.Fatalf("newPRD failed: %v", err) + } + + // Wait 200ms (timeout is 300ms) + time.Sleep(200 * time.Millisecond) + + // Session should still be active + if len(sm.activeSessions()) != 1 { + t.Fatal("expected session to still be active before timeout") + } + + // Send a message to reset the timer + if err := sm.sendMessage("sess-reset-1", "keep alive"); err != nil { + t.Fatalf("sendMessage failed: %v", err) + } + + // Wait another 200ms (total 400ms since start, but only 200ms since last activity) + time.Sleep(200 * time.Millisecond) + + // Session should still be active because we reset the timer + if len(sm.activeSessions()) != 1 { + t.Error("expected session to still be active after timer reset") + } + + // Wait for the full timeout from last activity (another 200ms) + time.Sleep(200 * time.Millisecond) + + // Now it should have timed out + if len(sm.activeSessions()) != 0 { + t.Errorf("expected 0 active sessions after timeout, got %d", len(sm.activeSessions())) + } + + sm.killAll() +} + +func TestSessionManager_IndependentTimers(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Create mock claude + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +while IFS= read -r line; do + echo "$line" +done +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sender := &captureSender{} + sm := &sessionManager{ + sessions: make(map[string]*claudeSession), + sender: sender, + timeout: 300 * time.Millisecond, + warningThresholds: []int{}, + checkInterval: 50 * time.Millisecond, + stopTimeout: make(chan struct{}), + } + go sm.runTimeoutChecker(sm.stopTimeout) + + projectDir1 := filepath.Join(home, "project1") + projectDir2 := filepath.Join(home, "project2") + os.MkdirAll(projectDir1, 0o755) + os.MkdirAll(projectDir2, 0o755) + + // Start two sessions + if err := sm.newPRD(projectDir1, "project1", "sess-a", "test"); err != nil { + t.Fatalf("newPRD a failed: %v", err) + } + if err := sm.newPRD(projectDir2, "project2", "sess-b", "test"); err != nil { + t.Fatalf("newPRD b failed: %v", err) + } + + // Both should be active + if len(sm.activeSessions()) != 2 { + t.Fatalf("expected 2 active sessions, got %d", len(sm.activeSessions())) + } + + // Keep session B alive by sending a message after 200ms + time.Sleep(200 * time.Millisecond) + if err := sm.sendMessage("sess-b", "keep alive"); err != nil { + t.Fatalf("sendMessage failed: %v", err) + } + + // Wait for session A to expire (another 200ms) + time.Sleep(200 * time.Millisecond) + + // Session A should be expired, session B should still be active + sessions := sm.activeSessions() + if len(sessions) != 1 { + t.Fatalf("expected 1 active session, got %d", len(sessions)) + } + if sessions[0].SessionID != "sess-b" { + t.Errorf("expected session 'sess-b' to survive, got %q", sessions[0].SessionID) + } + + // Verify session_expired was sent for sess-a + msgs := sender.getMessages() + hasExpiredA := false + for _, msg := range msgs { + if msg["type"] == "session_expired" && msg["session_id"] == "sess-a" { + hasExpiredA = true + break + } + } + + if !hasExpiredA { + t.Error("expected session_expired for sess-a") + } + + sm.killAll() +} + +func TestSessionManager_TimeoutCheckerGoroutineSafe(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + + // Create mock claude + mockClaudeBin := filepath.Join(home, "claude") + mockScript := `#!/bin/sh +while IFS= read -r line; do + echo "$line" +done +` + if err := os.WriteFile(mockClaudeBin, []byte(mockScript), 0o755); err != nil { + t.Fatal(err) + } + origPath := os.Getenv("PATH") + t.Setenv("PATH", home+":"+origPath) + + sm := &sessionManager{ + sessions: make(map[string]*claudeSession), + sender: &discardSender{}, + timeout: 500 * time.Millisecond, + warningThresholds: []int{}, + checkInterval: 50 * time.Millisecond, + stopTimeout: make(chan struct{}), + } + go sm.runTimeoutChecker(sm.stopTimeout) + + projectDir := filepath.Join(home, "testproject") + os.MkdirAll(projectDir, 0o755) + + // Concurrently create sessions and send messages while timeout checker runs + var wg sync.WaitGroup + for i := 0; i < 5; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + sid := fmt.Sprintf("sess-conc-%d", idx) + dir := filepath.Join(home, fmt.Sprintf("proj-%d", idx)) + os.MkdirAll(dir, 0o755) + + if err := sm.newPRD(dir, fmt.Sprintf("proj-%d", idx), sid, "test"); err != nil { + t.Errorf("newPRD %s failed: %v", sid, err) + return + } + + // Send some messages + for j := 0; j < 3; j++ { + time.Sleep(50 * time.Millisecond) + sm.sendMessage(sid, fmt.Sprintf("msg-%d", j)) + } + }(i) + } + wg.Wait() + + // No crash = goroutine-safe. Wait for all to expire. + time.Sleep(700 * time.Millisecond) + + if len(sm.activeSessions()) != 0 { + t.Errorf("expected all sessions to expire, got %d active", len(sm.activeSessions())) + } + + sm.killAll() +} diff --git a/internal/cmd/settings.go b/internal/cmd/settings.go new file mode 100644 index 00000000..09c01bfc --- /dev/null +++ b/internal/cmd/settings.go @@ -0,0 +1,119 @@ +package cmd + +import ( + "encoding/json" + "fmt" + "log" + + "github.com/minicodemonkey/chief/internal/config" + "github.com/minicodemonkey/chief/internal/ws" +) + +// handleGetSettings handles a get_settings request. +func handleGetSettings(sender messageSender, finder projectFinder, msg ws.Message) { + var req ws.GetSettingsMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing get_settings message: %v", err) + return + } + + project, found := finder.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + cfg, err := config.Load(project.Path) + if err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Failed to load settings: %v", err), msg.ID) + return + } + + resp := ws.SettingsResponseMessage{ + Type: ws.TypeSettingsResponse, + Payload: ws.SettingsResponsePayload{ + Project: req.Project, + Settings: ws.SettingsData{ + MaxIterations: cfg.EffectiveMaxIterations(), + AutoCommit: cfg.EffectiveAutoCommit(), + CommitPrefix: cfg.CommitPrefix, + ClaudeModel: cfg.ClaudeModel, + TestCommand: cfg.TestCommand, + }, + }, + } + if err := sender.Send(resp); err != nil { + log.Printf("Error sending settings_response: %v", err) + } +} + +// handleUpdateSettings handles an update_settings request. +func handleUpdateSettings(sender messageSender, finder projectFinder, msg ws.Message) { + var req ws.UpdateSettingsMessage + if err := json.Unmarshal(msg.Raw, &req); err != nil { + log.Printf("Error parsing update_settings message: %v", err) + return + } + + project, found := finder.FindProject(req.Project) + if !found { + sendError(sender, ws.ErrCodeProjectNotFound, + fmt.Sprintf("Project %q not found", req.Project), msg.ID) + return + } + + cfg, err := config.Load(project.Path) + if err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Failed to load settings: %v", err), msg.ID) + return + } + + // Merge provided fields + if req.MaxIterations != nil { + if *req.MaxIterations < 1 { + sendError(sender, ws.ErrCodeFilesystemError, + "max_iterations must be at least 1", msg.ID) + return + } + cfg.MaxIterations = *req.MaxIterations + } + if req.AutoCommit != nil { + cfg.AutoCommit = req.AutoCommit + } + if req.CommitPrefix != nil { + cfg.CommitPrefix = *req.CommitPrefix + } + if req.ClaudeModel != nil { + cfg.ClaudeModel = *req.ClaudeModel + } + if req.TestCommand != nil { + cfg.TestCommand = *req.TestCommand + } + + if err := config.Save(project.Path, cfg); err != nil { + sendError(sender, ws.ErrCodeFilesystemError, + fmt.Sprintf("Failed to save settings: %v", err), msg.ID) + return + } + + // Echo back full updated settings + resp := ws.SettingsResponseMessage{ + Type: ws.TypeSettingsUpdated, + Payload: ws.SettingsResponsePayload{ + Project: req.Project, + Settings: ws.SettingsData{ + MaxIterations: cfg.EffectiveMaxIterations(), + AutoCommit: cfg.EffectiveAutoCommit(), + CommitPrefix: cfg.CommitPrefix, + ClaudeModel: cfg.ClaudeModel, + TestCommand: cfg.TestCommand, + }, + }, + } + if err := sender.Send(resp); err != nil { + log.Printf("Error sending settings_updated: %v", err) + } +} diff --git a/internal/cmd/settings_test.go b/internal/cmd/settings_test.go new file mode 100644 index 00000000..bb703778 --- /dev/null +++ b/internal/cmd/settings_test.go @@ -0,0 +1,488 @@ +package cmd + +import ( + "encoding/json" + "os" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/config" +) + +func TestRunServe_GetSettings_Defaults(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var settingsReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]string{ + "type": "get_settings", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("settings_response", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &settingsReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if settingsReceived == nil { + t.Fatal("settings_response was not received") + } + if settingsReceived["type"] != "settings_response" { + t.Errorf("expected type 'settings_response', got %v", settingsReceived["type"]) + } + payload, ok := settingsReceived["payload"].(map[string]interface{}) + if !ok { + t.Fatal("expected payload to be an object") + } + if payload["project"] != "myproject" { + t.Errorf("expected project 'myproject', got %v", payload["project"]) + } + settings, ok := payload["settings"].(map[string]interface{}) + if !ok { + t.Fatal("expected settings to be an object") + } + // Default max_iterations should be 5 + if maxIter, ok := settings["max_iterations"].(float64); !ok || int(maxIter) != 5 { + t.Errorf("expected max_iterations 5, got %v", settings["max_iterations"]) + } + // Default auto_commit should be true + if autoCommit, ok := settings["auto_commit"].(bool); !ok || !autoCommit { + t.Errorf("expected auto_commit true, got %v", settings["auto_commit"]) + } + // Other fields should be empty strings + if settings["commit_prefix"] != "" { + t.Errorf("expected empty commit_prefix, got %v", settings["commit_prefix"]) + } + if settings["claude_model"] != "" { + t.Errorf("expected empty claude_model, got %v", settings["claude_model"]) + } + if settings["test_command"] != "" { + t.Errorf("expected empty test_command, got %v", settings["test_command"]) + } +} + +func TestRunServe_GetSettings_ProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]string{ + "type": "get_settings", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", errorReceived["code"]) + } +} + +func TestRunServe_GetSettings_WithExistingConfig(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + // Write existing config + autoCommit := false + cfg := &config.Config{ + MaxIterations: 10, + AutoCommit: &autoCommit, + CommitPrefix: "fix:", + ClaudeModel: "claude-sonnet-4-5-20250929", + TestCommand: "npm test", + } + if err := config.Save(projectDir, cfg); err != nil { + t.Fatal(err) + } + + var settingsReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]string{ + "type": "get_settings", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("settings_response", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &settingsReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if settingsReceived == nil { + t.Fatal("settings_response was not received") + } + payload := settingsReceived["payload"].(map[string]interface{}) + settings := payload["settings"].(map[string]interface{}) + if maxIter, ok := settings["max_iterations"].(float64); !ok || int(maxIter) != 10 { + t.Errorf("expected max_iterations 10, got %v", settings["max_iterations"]) + } + if autoCommitVal, ok := settings["auto_commit"].(bool); !ok || autoCommitVal { + t.Errorf("expected auto_commit false, got %v", settings["auto_commit"]) + } + if settings["commit_prefix"] != "fix:" { + t.Errorf("expected commit_prefix 'fix:', got %v", settings["commit_prefix"]) + } + if settings["claude_model"] != "claude-sonnet-4-5-20250929" { + t.Errorf("expected claude_model 'claude-sonnet-4-5-20250929', got %v", settings["claude_model"]) + } + if settings["test_command"] != "npm test" { + t.Errorf("expected test_command 'npm test', got %v", settings["test_command"]) + } +} + +func TestRunServe_UpdateSettings(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + var settingsReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + maxIter := 8 + autoCommit := false + commitPrefix := "chore:" + claudeModel := "claude-sonnet-4-5-20250929" + testCommand := "go test ./..." + + req := map[string]interface{}{ + "type": "update_settings", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "max_iterations": maxIter, + "auto_commit": autoCommit, + "commit_prefix": commitPrefix, + "claude_model": claudeModel, + "test_command": testCommand, + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("settings_updated", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &settingsReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if settingsReceived == nil { + t.Fatal("settings_updated was not received") + } + if settingsReceived["type"] != "settings_updated" { + t.Errorf("expected type 'settings_updated', got %v", settingsReceived["type"]) + } + payload := settingsReceived["payload"].(map[string]interface{}) + settings := payload["settings"].(map[string]interface{}) + if maxIter, ok := settings["max_iterations"].(float64); !ok || int(maxIter) != 8 { + t.Errorf("expected max_iterations 8, got %v", settings["max_iterations"]) + } + if autoCommitVal, ok := settings["auto_commit"].(bool); !ok || autoCommitVal { + t.Errorf("expected auto_commit false, got %v", settings["auto_commit"]) + } + if settings["commit_prefix"] != "chore:" { + t.Errorf("expected commit_prefix 'chore:', got %v", settings["commit_prefix"]) + } + if settings["claude_model"] != "claude-sonnet-4-5-20250929" { + t.Errorf("expected claude_model 'claude-sonnet-4-5-20250929', got %v", settings["claude_model"]) + } + if settings["test_command"] != "go test ./..." { + t.Errorf("expected test_command 'go test ./...', got %v", settings["test_command"]) + } + + // Verify the config was persisted to disk + cfg, err := config.Load(filepath.Join(workspaceDir, "myproject")) + if err != nil { + t.Fatalf("config.Load failed: %v", err) + } + if cfg.MaxIterations != 8 { + t.Errorf("expected saved max_iterations 8, got %d", cfg.MaxIterations) + } + if cfg.AutoCommit == nil || *cfg.AutoCommit != false { + t.Errorf("expected saved auto_commit false, got %v", cfg.AutoCommit) + } +} + +func TestRunServe_UpdateSettings_PartialUpdate(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + projectDir := filepath.Join(workspaceDir, "myproject") + createGitRepo(t, projectDir) + + // Set initial config + autoCommit := false + cfg := &config.Config{ + MaxIterations: 10, + AutoCommit: &autoCommit, + CommitPrefix: "fix:", + TestCommand: "npm test", + } + if err := config.Save(projectDir, cfg); err != nil { + t.Fatal(err) + } + + var settingsReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + // Only update test_command — other fields should be preserved + req := map[string]interface{}{ + "type": "update_settings", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "test_command": "go test ./...", + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("settings_updated", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &settingsReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if settingsReceived == nil { + t.Fatal("settings_updated was not received") + } + payload := settingsReceived["payload"].(map[string]interface{}) + settings := payload["settings"].(map[string]interface{}) + // Existing values should be preserved + if maxIter, ok := settings["max_iterations"].(float64); !ok || int(maxIter) != 10 { + t.Errorf("expected max_iterations 10 preserved, got %v", settings["max_iterations"]) + } + if autoCommitVal, ok := settings["auto_commit"].(bool); !ok || autoCommitVal { + t.Errorf("expected auto_commit false preserved, got %v", settings["auto_commit"]) + } + if settings["commit_prefix"] != "fix:" { + t.Errorf("expected commit_prefix 'fix:' preserved, got %v", settings["commit_prefix"]) + } + // Updated value + if settings["test_command"] != "go test ./..." { + t.Errorf("expected test_command 'go test ./...', got %v", settings["test_command"]) + } +} + +func TestRunServe_UpdateSettings_ProjectNotFound(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]interface{}{ + "type": "update_settings", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "nonexistent", + "max_iterations": 3, + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "PROJECT_NOT_FOUND" { + t.Errorf("expected code 'PROJECT_NOT_FOUND', got %v", errorReceived["code"]) + } +} + +func TestRunServe_UpdateSettings_InvalidMaxIterations(t *testing.T) { + home := t.TempDir() + setTestHome(t, home) + setupServeCredentials(t) + + workspaceDir := filepath.Join(home, "projects") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatal(err) + } + + createGitRepo(t, filepath.Join(workspaceDir, "myproject")) + + var errorReceived map[string]interface{} + var mu sync.Mutex + + err := serveTestHelper(t, workspaceDir, func(ms *mockUplinkServer) { + req := map[string]interface{}{ + "type": "update_settings", + "id": "req-1", + "timestamp": time.Now().UTC().Format(time.RFC3339), + "project": "myproject", + "max_iterations": 0, + } + if err := ms.sendCommand(req); err != nil { + t.Errorf("sendCommand failed: %v", err) + return + } + + raw, err := ms.waitForMessageType("error", 5*time.Second) + if err == nil { + mu.Lock() + json.Unmarshal(raw, &errorReceived) + mu.Unlock() + } + }) + if err != nil { + t.Fatalf("RunServe returned error: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + if errorReceived == nil { + t.Fatal("error message was not received") + } + if errorReceived["type"] != "error" { + t.Errorf("expected type 'error', got %v", errorReceived["type"]) + } + if errorReceived["code"] != "FILESYSTEM_ERROR" { + t.Errorf("expected code 'FILESYSTEM_ERROR', got %v", errorReceived["code"]) + } +} diff --git a/internal/cmd/status.go b/internal/cmd/status.go index 02e26aaa..8fbf6cb5 100644 --- a/internal/cmd/status.go +++ b/internal/cmd/status.go @@ -30,7 +30,7 @@ func RunStatus(opts StatusOptions) error { } // Build PRD path - prdPath := filepath.Join(opts.BaseDir, ".chief", "prds", opts.Name, "prd.md") + prdPath := filepath.Join(opts.BaseDir, ".chief", "prds", opts.Name, "prd.json") // Load PRD p, err := prd.LoadPRD(prdPath) @@ -123,7 +123,7 @@ func RunList(opts ListOptions) error { } name := entry.Name() - prdPath := filepath.Join(prdsDir, name, "prd.md") + prdPath := filepath.Join(prdsDir, name, "prd.json") // Try to load the PRD p, err := prd.LoadPRD(prdPath) diff --git a/internal/cmd/status_test.go b/internal/cmd/status_test.go index f4dee9e6..97144de2 100644 --- a/internal/cmd/status_test.go +++ b/internal/cmd/status_test.go @@ -9,29 +9,25 @@ import ( func TestRunStatusWithValidPRD(t *testing.T) { tmpDir := t.TempDir() + // Create a PRD directory with prd.json prdDir := filepath.Join(tmpDir, ".chief", "prds", "test") if err := os.MkdirAll(prdDir, 0755); err != nil { t.Fatalf("Failed to create directory: %v", err) } - prdMd := `# Test Project - -Test description - -### US-001: Story 1 -**Status:** done -- [x] Done - -### US-002: Story 2 -- [ ] Pending - -### US-003: Story 3 -**Status:** in-progress -- [ ] Working -` - prdPath := filepath.Join(prdDir, "prd.md") - if err := os.WriteFile(prdPath, []byte(prdMd), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) + // Create a test prd.json + prdJSON := `{ + "project": "Test Project", + "description": "Test description", + "userStories": [ + {"id": "US-001", "title": "Story 1", "passes": true, "priority": 1}, + {"id": "US-002", "title": "Story 2", "passes": false, "priority": 2}, + {"id": "US-003", "title": "Story 3", "passes": false, "inProgress": true, "priority": 3} + ] +}` + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(prdJSON), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) } opts := StatusOptions{ @@ -39,6 +35,7 @@ Test description BaseDir: tmpDir, } + // Should not return error err := RunStatus(opts) if err != nil { t.Errorf("RunStatus() returned error: %v", err) @@ -48,19 +45,23 @@ Test description func TestRunStatusWithDefaultName(t *testing.T) { tmpDir := t.TempDir() + // Create a PRD directory with prd.json using default name "main" prdDir := filepath.Join(tmpDir, ".chief", "prds", "main") if err := os.MkdirAll(prdDir, 0755); err != nil { t.Fatalf("Failed to create directory: %v", err) } - prdMd := "# Main Project\n" - prdPath := filepath.Join(prdDir, "prd.md") - if err := os.WriteFile(prdPath, []byte(prdMd), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) + prdJSON := `{ + "project": "Main Project", + "userStories": [] +}` + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(prdJSON), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) } opts := StatusOptions{ - Name: "", + Name: "", // Empty should default to "main" BaseDir: tmpDir, } @@ -91,6 +92,7 @@ func TestRunListWithNoPRDs(t *testing.T) { BaseDir: tmpDir, } + // Should not return error, just print "No PRDs found" err := RunList(opts) if err != nil { t.Errorf("RunList() returned error: %v", err) @@ -100,17 +102,24 @@ func TestRunListWithNoPRDs(t *testing.T) { func TestRunListWithPRDs(t *testing.T) { tmpDir := t.TempDir() + // Create multiple PRD directories prds := []struct { - name string - md string + name string + project string + stories string }{ { "auth", - "# Authentication\n\n### US-001: Login\n**Status:** done\n- [x] Works\n\n### US-002: Logout\n- [ ] Works\n", + "Authentication", + `[{"id": "US-001", "title": "Login", "passes": true, "priority": 1}, + {"id": "US-002", "title": "Logout", "passes": false, "priority": 2}]`, }, { "api", - "# API Service\n\n### US-001: Endpoints\n**Status:** done\n- [x] Done\n\n### US-002: Auth\n**Status:** done\n- [x] Done\n\n### US-003: Rate limiting\n**Status:** done\n- [x] Done\n", + "API Service", + `[{"id": "US-001", "title": "Endpoints", "passes": true, "priority": 1}, + {"id": "US-002", "title": "Auth", "passes": true, "priority": 2}, + {"id": "US-003", "title": "Rate limiting", "passes": true, "priority": 3}]`, }, } @@ -119,9 +128,11 @@ func TestRunListWithPRDs(t *testing.T) { if err := os.MkdirAll(prdDir, 0755); err != nil { t.Fatalf("Failed to create directory: %v", err) } - prdPath := filepath.Join(prdDir, "prd.md") - if err := os.WriteFile(prdPath, []byte(p.md), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) + + prdJSON := `{"project": "` + p.project + `", "userStories": ` + p.stories + `}` + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(prdJSON), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) } } @@ -143,20 +154,31 @@ func TestRunListSkipsInvalidPRDs(t *testing.T) { if err := os.MkdirAll(validDir, 0755); err != nil { t.Fatalf("Failed to create directory: %v", err) } - if err := os.WriteFile(filepath.Join(validDir, "prd.md"), []byte("# Valid\n"), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) + validJSON := `{"project": "Valid", "userStories": []}` + if err := os.WriteFile(filepath.Join(validDir, "prd.json"), []byte(validJSON), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) } - // Create an invalid PRD directory (no prd.md) + // Create an invalid PRD directory (no prd.json) invalidDir := filepath.Join(tmpDir, ".chief", "prds", "invalid") if err := os.MkdirAll(invalidDir, 0755); err != nil { t.Fatalf("Failed to create directory: %v", err) } + // Create another invalid PRD (invalid JSON) + badJsonDir := filepath.Join(tmpDir, ".chief", "prds", "badjson") + if err := os.MkdirAll(badJsonDir, 0755); err != nil { + t.Fatalf("Failed to create directory: %v", err) + } + if err := os.WriteFile(filepath.Join(badJsonDir, "prd.json"), []byte("not json"), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) + } + opts := ListOptions{ BaseDir: tmpDir, } + // Should not return error, just skip invalid PRDs err := RunList(opts) if err != nil { t.Errorf("RunList() returned error: %v", err) @@ -171,10 +193,16 @@ func TestRunStatusAllComplete(t *testing.T) { t.Fatalf("Failed to create directory: %v", err) } - prdMd := "# Complete Project\n\n### US-001: Story 1\n**Status:** done\n- [x] Done\n\n### US-002: Story 2\n**Status:** done\n- [x] Done\n" - prdPath := filepath.Join(prdDir, "prd.md") - if err := os.WriteFile(prdPath, []byte(prdMd), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) + prdJSON := `{ + "project": "Complete Project", + "userStories": [ + {"id": "US-001", "title": "Story 1", "passes": true, "priority": 1}, + {"id": "US-002", "title": "Story 2", "passes": true, "priority": 2} + ] +}` + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(prdJSON), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) } opts := StatusOptions{ @@ -196,10 +224,10 @@ func TestRunStatusEmptyPRD(t *testing.T) { t.Fatalf("Failed to create directory: %v", err) } - prdMd := "# Empty Project\n" - prdPath := filepath.Join(prdDir, "prd.md") - if err := os.WriteFile(prdPath, []byte(prdMd), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) + prdJSON := `{"project": "Empty Project", "userStories": []}` + prdPath := filepath.Join(prdDir, "prd.json") + if err := os.WriteFile(prdPath, []byte(prdJSON), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) } opts := StatusOptions{ diff --git a/internal/config/config.go b/internal/config/config.go index 4523b1ba..84084f9a 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -1,6 +1,7 @@ package config import ( + "fmt" "os" "path/filepath" @@ -9,17 +10,45 @@ import ( const configFile = ".chief/config.yaml" -// Config holds project-level settings for Chief. -type Config struct { - Worktree WorktreeConfig `yaml:"worktree"` - OnComplete OnCompleteConfig `yaml:"onComplete"` - Agent AgentConfig `yaml:"agent"` +// UserConfig holds user-level settings from ~/.chief/config.yaml. +type UserConfig struct { + WSURL string `yaml:"ws_url,omitempty"` +} + +// LoadUserConfig reads the user-level config from ~/.chief/config.yaml. +// Returns an empty UserConfig when the file doesn't exist (no error). +func LoadUserConfig() (*UserConfig, error) { + home, err := os.UserHomeDir() + if err != nil { + return &UserConfig{}, fmt.Errorf("determining home directory: %w", err) + } + + path := filepath.Join(home, ".chief", "config.yaml") + data, err := os.ReadFile(path) + if err != nil { + if os.IsNotExist(err) { + return &UserConfig{}, nil + } + return &UserConfig{}, err + } + + cfg := &UserConfig{} + if err := yaml.Unmarshal(data, cfg); err != nil { + return &UserConfig{}, err + } + + return cfg, nil } -// AgentConfig holds agent CLI settings (Claude, Codex, OpenCode, or Cursor). -type AgentConfig struct { - Provider string `yaml:"provider"` // "claude" (default) | "codex" | "opencode" | "cursor" - CLIPath string `yaml:"cliPath"` // optional custom path to CLI binary +// Config holds project-level settings for Chief. +type Config struct { + Worktree WorktreeConfig `yaml:"worktree"` + OnComplete OnCompleteConfig `yaml:"onComplete"` + MaxIterations int `yaml:"maxIterations,omitempty"` + AutoCommit *bool `yaml:"autoCommit,omitempty"` + CommitPrefix string `yaml:"commitPrefix,omitempty"` + ClaudeModel string `yaml:"claudeModel,omitempty"` + TestCommand string `yaml:"testCommand,omitempty"` } // WorktreeConfig holds worktree-related settings. @@ -33,11 +62,30 @@ type OnCompleteConfig struct { CreatePR bool `yaml:"createPR"` } +// DefaultMaxIterations is the default value for MaxIterations when not set. +const DefaultMaxIterations = 5 + // Default returns a Config with zero-value defaults. func Default() *Config { return &Config{} } +// EffectiveMaxIterations returns MaxIterations or the default if not set. +func (c *Config) EffectiveMaxIterations() int { + if c.MaxIterations > 0 { + return c.MaxIterations + } + return DefaultMaxIterations +} + +// EffectiveAutoCommit returns AutoCommit or true if not set. +func (c *Config) EffectiveAutoCommit() bool { + if c.AutoCommit != nil { + return *c.AutoCommit + } + return true +} + // configPath returns the full path to the config file. func configPath(baseDir string) string { return filepath.Join(baseDir, configFile) diff --git a/internal/config/config_test.go b/internal/config/config_test.go index dacee39c..7a678230 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -62,6 +62,121 @@ func TestSaveAndLoad(t *testing.T) { } } +func TestSaveAndLoadSettingsFields(t *testing.T) { + dir := t.TempDir() + autoCommit := false + cfg := &Config{ + MaxIterations: 10, + AutoCommit: &autoCommit, + CommitPrefix: "feat:", + ClaudeModel: "claude-sonnet-4-5-20250929", + TestCommand: "go test ./...", + } + + if err := Save(dir, cfg); err != nil { + t.Fatalf("Save failed: %v", err) + } + + loaded, err := Load(dir) + if err != nil { + t.Fatalf("Load failed: %v", err) + } + + if loaded.MaxIterations != 10 { + t.Errorf("expected MaxIterations 10, got %d", loaded.MaxIterations) + } + if loaded.AutoCommit == nil || *loaded.AutoCommit != false { + t.Errorf("expected AutoCommit false, got %v", loaded.AutoCommit) + } + if loaded.CommitPrefix != "feat:" { + t.Errorf("expected CommitPrefix %q, got %q", "feat:", loaded.CommitPrefix) + } + if loaded.ClaudeModel != "claude-sonnet-4-5-20250929" { + t.Errorf("expected ClaudeModel %q, got %q", "claude-sonnet-4-5-20250929", loaded.ClaudeModel) + } + if loaded.TestCommand != "go test ./..." { + t.Errorf("expected TestCommand %q, got %q", "go test ./...", loaded.TestCommand) + } +} + +func TestEffectiveDefaults(t *testing.T) { + cfg := Default() + + if cfg.EffectiveMaxIterations() != 5 { + t.Errorf("expected EffectiveMaxIterations 5, got %d", cfg.EffectiveMaxIterations()) + } + if !cfg.EffectiveAutoCommit() { + t.Error("expected EffectiveAutoCommit true") + } + + // With explicit values + cfg.MaxIterations = 3 + autoCommit := false + cfg.AutoCommit = &autoCommit + + if cfg.EffectiveMaxIterations() != 3 { + t.Errorf("expected EffectiveMaxIterations 3, got %d", cfg.EffectiveMaxIterations()) + } + if cfg.EffectiveAutoCommit() { + t.Error("expected EffectiveAutoCommit false") + } +} + +func TestLoadUserConfig_NonExistent(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + + cfg, err := LoadUserConfig() + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if cfg.WSURL != "" { + t.Errorf("expected empty WSURL, got %q", cfg.WSURL) + } +} + +func TestLoadUserConfig_WithWSURL(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + + chiefDir := filepath.Join(home, ".chief") + if err := os.MkdirAll(chiefDir, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(chiefDir, "config.yaml"), []byte("ws_url: ws://localhost:8080/ws/server\n"), 0o644); err != nil { + t.Fatal(err) + } + + cfg, err := LoadUserConfig() + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if cfg.WSURL != "ws://localhost:8080/ws/server" { + t.Errorf("expected ws://localhost:8080/ws/server, got %q", cfg.WSURL) + } +} + +func TestLoadUserConfig_EmptyWSURL(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + + chiefDir := filepath.Join(home, ".chief") + if err := os.MkdirAll(chiefDir, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(chiefDir, "config.yaml"), []byte("{}\n"), 0o644); err != nil { + t.Fatal(err) + } + + cfg, err := LoadUserConfig() + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if cfg.WSURL != "" { + t.Errorf("expected empty WSURL, got %q", cfg.WSURL) + } +} + func TestExists(t *testing.T) { dir := t.TempDir() diff --git a/internal/contract/contract_test.go b/internal/contract/contract_test.go new file mode 100644 index 00000000..b3d9a886 --- /dev/null +++ b/internal/contract/contract_test.go @@ -0,0 +1,558 @@ +package contract + +import ( + "encoding/json" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/minicodemonkey/chief/internal/uplink" + "github.com/minicodemonkey/chief/internal/ws" +) + +// fixturesDir returns the absolute path to contract/fixtures relative to the repo root. +func fixturesDir(t *testing.T) string { + t.Helper() + // Determine repo root from this test file's location: + // internal/contract/contract_test.go → ../../contract/fixtures + _, thisFile, _, ok := runtime.Caller(0) + if !ok { + t.Fatal("cannot determine test file location") + } + return filepath.Join(filepath.Dir(thisFile), "..", "..", "contract", "fixtures") +} + +func loadFixture(t *testing.T, relPath string) []byte { + t.Helper() + data, err := os.ReadFile(filepath.Join(fixturesDir(t), relPath)) + if err != nil { + t.Fatalf("loading fixture %s: %v", relPath, err) + } + return data +} + +// --- server-to-cli fixtures --- + +func TestWelcomeResponse_Deserialize(t *testing.T) { + data := loadFixture(t, "server-to-cli/welcome_response.json") + + var welcome uplink.WelcomeResponse + if err := json.Unmarshal(data, &welcome); err != nil { + t.Fatalf("failed to unmarshal welcome_response.json: %v", err) + } + + if welcome.Type != "welcome" { + t.Errorf("type = %q, want %q", welcome.Type, "welcome") + } + if welcome.ProtocolVersion != 1 { + t.Errorf("protocol_version = %d, want 1", welcome.ProtocolVersion) + } + if welcome.DeviceID != 42 { + t.Errorf("device_id = %d, want 42", welcome.DeviceID) + } + if welcome.SessionID != "550e8400-e29b-41d4-a716-446655440000" { + t.Errorf("session_id = %q, want UUID", welcome.SessionID) + } + + // Reverb config — port MUST be an int, not a string + if welcome.Reverb.Port != 8080 { + t.Errorf("reverb.port = %d, want 8080", welcome.Reverb.Port) + } + if welcome.Reverb.Key != "test-app-key" { + t.Errorf("reverb.key = %q, want %q", welcome.Reverb.Key, "test-app-key") + } + if welcome.Reverb.Host != "127.0.0.1" { + t.Errorf("reverb.host = %q, want %q", welcome.Reverb.Host, "127.0.0.1") + } + if welcome.Reverb.Scheme != "https" { + t.Errorf("reverb.scheme = %q, want %q", welcome.Reverb.Scheme, "https") + } +} + +func TestWelcomeResponse_PortIsInt(t *testing.T) { + // Regression: PHP env() returns strings — verify port decodes as int. + data := loadFixture(t, "server-to-cli/welcome_response.json") + + var raw map[string]json.RawMessage + if err := json.Unmarshal(data, &raw); err != nil { + t.Fatal(err) + } + + var reverb map[string]json.RawMessage + json.Unmarshal(raw["reverb"], &reverb) + + // Verify port is a JSON number, not a string + portStr := string(reverb["port"]) + if portStr == `"8080"` { + t.Fatal("reverb.port is a JSON string — must be a number") + } + if portStr != "8080" { + t.Errorf("reverb.port raw JSON = %s, want 8080", portStr) + } +} + +func TestCommandCreateProject_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_create_project.json") + + // Verify the envelope has type + payload + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "create_project" { + t.Errorf("envelope type = %q, want %q", env.Type, "create_project") + } + if len(env.Payload) == 0 { + t.Fatal("envelope payload is empty — commands must have payload wrapper") + } + + // The payload itself should parse into CreateProjectMessage fields + var req ws.CreateProjectMessage + if err := json.Unmarshal(env.Payload, &req); err != nil { + t.Fatalf("failed to unmarshal payload into CreateProjectMessage: %v", err) + } + + if req.Name != "new-project" { + t.Errorf("payload.name = %q, want %q", req.Name, "new-project") + } + if !req.GitInit { + t.Error("payload.git_init = false, want true") + } +} + +func TestCommandStartRun_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_start_run.json") + + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "start_run" { + t.Errorf("envelope type = %q, want %q", env.Type, "start_run") + } + + var req ws.StartRunMessage + if err := json.Unmarshal(env.Payload, &req); err != nil { + t.Fatalf("failed to unmarshal payload into StartRunMessage: %v", err) + } + + if req.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", req.Project, "my-project") + } + if req.PRDID != "feature-auth" { + t.Errorf("payload.prd_id = %q, want %q", req.PRDID, "feature-auth") + } +} + +func TestCommandListProjects_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_list_projects.json") + + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "list_projects" { + t.Errorf("envelope type = %q, want %q", env.Type, "list_projects") + } +} + +// --- cli-to-server fixtures --- + +func TestStateSnapshot_Roundtrip(t *testing.T) { + data := loadFixture(t, "cli-to-server/state_snapshot.json") + + // Unmarshal into the Go struct + var snapshot ws.StateSnapshotMessage + if err := json.Unmarshal(data, &snapshot); err != nil { + t.Fatalf("failed to unmarshal state_snapshot.json: %v", err) + } + + if snapshot.Type != "state_snapshot" { + t.Errorf("type = %q, want %q", snapshot.Type, "state_snapshot") + } + if len(snapshot.Projects) != 1 { + t.Fatalf("projects count = %d, want 1", len(snapshot.Projects)) + } + + // Verify project uses "name" field, not "project_slug" + proj := snapshot.Projects[0] + if proj.Name != "my-project" { + t.Errorf("project.name = %q, want %q", proj.Name, "my-project") + } + if proj.Path != "/home/user/projects/my-project" { + t.Errorf("project.path = %q", proj.Path) + } + if !proj.HasChief { + t.Error("project.has_chief = false, want true") + } + if proj.Branch != "main" { + t.Errorf("project.branch = %q, want %q", proj.Branch, "main") + } + if proj.Commit.Hash != "abc1234" { + t.Errorf("project.commit.hash = %q, want %q", proj.Commit.Hash, "abc1234") + } + + // Re-marshal and verify it round-trips cleanly + remarshaled, err := json.Marshal(snapshot) + if err != nil { + t.Fatalf("failed to re-marshal: %v", err) + } + + var roundtrip ws.StateSnapshotMessage + if err := json.Unmarshal(remarshaled, &roundtrip); err != nil { + t.Fatalf("failed to unmarshal round-trip: %v", err) + } + if roundtrip.Projects[0].Name != "my-project" { + t.Errorf("round-trip project.name = %q, want %q", roundtrip.Projects[0].Name, "my-project") + } +} + +func TestStateSnapshot_NameFieldNotProjectSlug(t *testing.T) { + // Regression: CLI sends "name", not "project_slug". + data := loadFixture(t, "cli-to-server/state_snapshot.json") + + var raw map[string]json.RawMessage + json.Unmarshal(data, &raw) + + var projects []map[string]json.RawMessage + json.Unmarshal(raw["projects"], &projects) + + if len(projects) == 0 { + t.Fatal("no projects in fixture") + } + + proj := projects[0] + if _, hasName := proj["name"]; !hasName { + t.Error("project should have 'name' field") + } + if _, hasSlug := proj["project_slug"]; hasSlug { + t.Error("project should NOT have 'project_slug' field — CLI uses 'name'") + } +} + +func TestConnectRequest_Deserialize(t *testing.T) { + data := loadFixture(t, "cli-to-server/connect_request.json") + + var req struct { + ChiefVersion string `json:"chief_version"` + DeviceName string `json:"device_name"` + OS string `json:"os"` + Arch string `json:"arch"` + ProtocolVersion int `json:"protocol_version"` + } + if err := json.Unmarshal(data, &req); err != nil { + t.Fatalf("failed to unmarshal connect_request.json: %v", err) + } + + if req.ChiefVersion != "1.0.0" { + t.Errorf("chief_version = %q, want %q", req.ChiefVersion, "1.0.0") + } + if req.ProtocolVersion != 1 { + t.Errorf("protocol_version = %d, want 1", req.ProtocolVersion) + } + if req.OS == "" { + t.Error("os should not be empty") + } +} + +func TestCommandGetPRDs_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_get_prds.json") + + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "get_prds" { + t.Errorf("envelope type = %q, want %q", env.Type, "get_prds") + } + + var req ws.GetPRDsMessage + if err := json.Unmarshal(env.Payload, &req); err != nil { + t.Fatalf("failed to unmarshal payload into GetPRDsMessage: %v", err) + } + + if req.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", req.Project, "my-project") + } +} + +func TestCommandGetSettings_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_get_settings.json") + + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "get_settings" { + t.Errorf("envelope type = %q, want %q", env.Type, "get_settings") + } + + var req ws.GetSettingsMessage + if err := json.Unmarshal(env.Payload, &req); err != nil { + t.Fatalf("failed to unmarshal payload into GetSettingsMessage: %v", err) + } + + if req.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", req.Project, "my-project") + } +} + +func TestCommandGetDiffs_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_get_diffs.json") + + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "get_diffs" { + t.Errorf("envelope type = %q, want %q", env.Type, "get_diffs") + } + + var req ws.GetDiffsMessage + if err := json.Unmarshal(env.Payload, &req); err != nil { + t.Fatalf("failed to unmarshal payload into GetDiffsMessage: %v", err) + } + + if req.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", req.Project, "my-project") + } + if req.StoryID != "US-001" { + t.Errorf("payload.story_id = %q, want %q", req.StoryID, "US-001") + } +} + +func TestCommandNewPRD_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_new_prd.json") + + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "new_prd" { + t.Errorf("envelope type = %q, want %q", env.Type, "new_prd") + } + + var req ws.NewPRDMessage + if err := json.Unmarshal(env.Payload, &req); err != nil { + t.Fatalf("failed to unmarshal payload into NewPRDMessage: %v", err) + } + + if req.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", req.Project, "my-project") + } + if req.SessionID != "session-abc" { + t.Errorf("payload.session_id = %q, want %q", req.SessionID, "session-abc") + } + if req.Message != "Build an authentication system" { + t.Errorf("payload.message = %q, want %q", req.Message, "Build an authentication system") + } +} + +func TestCommandPRDMessage_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_prd_message.json") + + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "prd_message" { + t.Errorf("envelope type = %q, want %q", env.Type, "prd_message") + } + + var req ws.PRDMessageMessage + if err := json.Unmarshal(env.Payload, &req); err != nil { + t.Fatalf("failed to unmarshal payload into PRDMessageMessage: %v", err) + } + + if req.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", req.Project, "my-project") + } + if req.SessionID != "session-abc" { + t.Errorf("payload.session_id = %q, want %q", req.SessionID, "session-abc") + } + if req.Message != "Add OAuth support to the PRD" { + t.Errorf("payload.message = %q, want %q", req.Message, "Add OAuth support to the PRD") + } +} + +func TestCommandRefinePRD_PayloadWrapper(t *testing.T) { + data := loadFixture(t, "server-to-cli/command_refine_prd.json") + + var env struct { + Type string `json:"type"` + Payload json.RawMessage `json:"payload,omitempty"` + } + if err := json.Unmarshal(data, &env); err != nil { + t.Fatalf("failed to unmarshal command envelope: %v", err) + } + + if env.Type != "refine_prd" { + t.Errorf("envelope type = %q, want %q", env.Type, "refine_prd") + } + + var req ws.RefinePRDMessage + if err := json.Unmarshal(env.Payload, &req); err != nil { + t.Fatalf("failed to unmarshal payload into RefinePRDMessage: %v", err) + } + + if req.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", req.Project, "my-project") + } + if req.SessionID != "session-abc" { + t.Errorf("payload.session_id = %q, want %q", req.SessionID, "session-abc") + } + if req.PRDID != "feature-auth" { + t.Errorf("payload.prd_id = %q, want %q", req.PRDID, "feature-auth") + } + if req.Message != "Add OAuth support to the PRD" { + t.Errorf("payload.message = %q, want %q", req.Message, "Add OAuth support to the PRD") + } +} + +// --- cli-to-server response fixtures --- + +func TestPRDsResponse_Roundtrip(t *testing.T) { + data := loadFixture(t, "cli-to-server/prds_response.json") + + var resp ws.PRDsResponseMessage + if err := json.Unmarshal(data, &resp); err != nil { + t.Fatalf("failed to unmarshal prds_response.json: %v", err) + } + + if resp.Type != "prds_response" { + t.Errorf("type = %q, want %q", resp.Type, "prds_response") + } + if resp.Payload.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", resp.Payload.Project, "my-project") + } + if len(resp.Payload.PRDs) != 2 { + t.Fatalf("prds count = %d, want 2", len(resp.Payload.PRDs)) + } + + prd := resp.Payload.PRDs[0] + if prd.ID != "feature-auth" { + t.Errorf("prds[0].id = %q, want %q", prd.ID, "feature-auth") + } + if prd.Status != "active" { + t.Errorf("prds[0].status = %q, want %q", prd.Status, "active") + } + if prd.StoryCount != 5 { + t.Errorf("prds[0].story_count = %d, want 5", prd.StoryCount) + } +} + +func TestSettingsResponse_Roundtrip(t *testing.T) { + data := loadFixture(t, "cli-to-server/settings_response.json") + + var resp ws.SettingsResponseMessage + if err := json.Unmarshal(data, &resp); err != nil { + t.Fatalf("failed to unmarshal settings_response.json: %v", err) + } + + if resp.Type != "settings_response" { + t.Errorf("type = %q, want %q", resp.Type, "settings_response") + } + if resp.Payload.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", resp.Payload.Project, "my-project") + } + if resp.Payload.Settings.MaxIterations != 5 { + t.Errorf("settings.max_iterations = %d, want 5", resp.Payload.Settings.MaxIterations) + } + if !resp.Payload.Settings.AutoCommit { + t.Error("settings.auto_commit = false, want true") + } +} + +func TestDiffsResponse_Roundtrip(t *testing.T) { + data := loadFixture(t, "cli-to-server/diffs_response.json") + + var resp ws.DiffsResponseMessage + if err := json.Unmarshal(data, &resp); err != nil { + t.Fatalf("failed to unmarshal diffs_response.json: %v", err) + } + + if resp.Type != "diffs_response" { + t.Errorf("type = %q, want %q", resp.Type, "diffs_response") + } + if resp.Payload.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", resp.Payload.Project, "my-project") + } + if resp.Payload.StoryID != "US-001" { + t.Errorf("payload.story_id = %q, want %q", resp.Payload.StoryID, "US-001") + } + if len(resp.Payload.Files) != 1 { + t.Fatalf("files count = %d, want 1", len(resp.Payload.Files)) + } + + file := resp.Payload.Files[0] + if file.Filename != "src/auth.go" { + t.Errorf("files[0].filename = %q, want %q", file.Filename, "src/auth.go") + } + if file.Additions != 25 { + t.Errorf("files[0].additions = %d, want 25", file.Additions) + } + if file.Deletions != 3 { + t.Errorf("files[0].deletions = %d, want 3", file.Deletions) + } +} + +func TestMessagesBatch_Deserialize(t *testing.T) { + data := loadFixture(t, "cli-to-server/messages_batch.json") + + var batch struct { + BatchID string `json:"batch_id"` + Messages []json.RawMessage `json:"messages"` + } + if err := json.Unmarshal(data, &batch); err != nil { + t.Fatalf("failed to unmarshal messages_batch.json: %v", err) + } + + if batch.BatchID == "" { + t.Error("batch_id should not be empty") + } + if len(batch.Messages) != 1 { + t.Fatalf("messages count = %d, want 1", len(batch.Messages)) + } + + // First message should be a state_snapshot + var msg ws.StateSnapshotMessage + if err := json.Unmarshal(batch.Messages[0], &msg); err != nil { + t.Fatalf("failed to unmarshal first message: %v", err) + } + if msg.Type != "state_snapshot" { + t.Errorf("first message type = %q, want %q", msg.Type, "state_snapshot") + } +} diff --git a/internal/engine/engine.go b/internal/engine/engine.go new file mode 100644 index 00000000..cd45457a --- /dev/null +++ b/internal/engine/engine.go @@ -0,0 +1,248 @@ +// Package engine provides a shared orchestration layer on top of loop.Manager +// that both the TUI and the serve command (WebSocket handler) can consume. +// It supports multiple concurrent event consumers via fan-out subscription. +package engine + +import ( + "sync" + + "github.com/minicodemonkey/chief/internal/config" + "github.com/minicodemonkey/chief/internal/loop" + "github.com/minicodemonkey/chief/internal/prd" +) + +// Engine wraps loop.Manager to provide a shared interface for driving Ralph loops +// and Claude sessions. It fans out events to multiple consumers. +type Engine struct { + manager *loop.Manager + + // Fan-out event distribution + subscribers map[int]chan ManagerEvent + nextID int + subMu sync.RWMutex + + // Forwarding goroutine lifecycle + stopForward chan struct{} + forwarding bool + forwardMu sync.Mutex +} + +// ManagerEvent wraps a loop.ManagerEvent for engine consumers. +// It mirrors loop.ManagerEvent to avoid exposing the loop package directly. +type ManagerEvent = loop.ManagerEvent + +// New creates a new Engine with the given max iterations. +func New(maxIter int) *Engine { + e := &Engine{ + manager: loop.NewManager(maxIter), + subscribers: make(map[int]chan ManagerEvent), + stopForward: make(chan struct{}), + } + e.startForwarding() + return e +} + +// startForwarding starts the goroutine that reads from the manager's event +// channel and fans out to all subscribers. +func (e *Engine) startForwarding() { + e.forwardMu.Lock() + defer e.forwardMu.Unlock() + + if e.forwarding { + return + } + e.forwarding = true + + go func() { + for { + select { + case event, ok := <-e.manager.Events(): + if !ok { + return + } + e.subMu.RLock() + for _, ch := range e.subscribers { + // Non-blocking send: drop events for slow consumers + select { + case ch <- event: + default: + } + } + e.subMu.RUnlock() + + case <-e.stopForward: + return + } + } + }() +} + +// Subscribe creates a new event subscription and returns a channel and an +// unsubscribe function. The channel is buffered (100 events). The caller must +// call the returned function when done to avoid resource leaks. +func (e *Engine) Subscribe() (<-chan ManagerEvent, func()) { + ch := make(chan ManagerEvent, 100) + + e.subMu.Lock() + id := e.nextID + e.nextID++ + e.subscribers[id] = ch + e.subMu.Unlock() + + unsub := func() { + e.subMu.Lock() + delete(e.subscribers, id) + e.subMu.Unlock() + } + + return ch, unsub +} + +// Manager returns the underlying loop.Manager for direct access when needed. +// This is useful for operations like Register, UpdateWorktreeInfo, etc. +// that don't need to be abstracted by the engine. +func (e *Engine) Manager() *loop.Manager { + return e.manager +} + +// --- Delegated Manager methods --- + +// Register registers a PRD with the engine (does not start it). +func (e *Engine) Register(name, prdPath string) error { + return e.manager.Register(name, prdPath) +} + +// RegisterWithWorktree registers a PRD with worktree metadata. +func (e *Engine) RegisterWithWorktree(name, prdPath, worktreeDir, branch string) error { + return e.manager.RegisterWithWorktree(name, prdPath, worktreeDir, branch) +} + +// Unregister removes a PRD from the engine. +func (e *Engine) Unregister(name string) error { + return e.manager.Unregister(name) +} + +// Start starts the loop for a specific PRD. +func (e *Engine) Start(name string) error { + return e.manager.Start(name) +} + +// Pause pauses the loop for a specific PRD. +func (e *Engine) Pause(name string) error { + return e.manager.Pause(name) +} + +// Stop stops the loop for a specific PRD immediately. +func (e *Engine) Stop(name string) error { + return e.manager.Stop(name) +} + +// StopAll stops all running loops and waits for completion. +func (e *Engine) StopAll() { + e.manager.StopAll() +} + +// GetState returns the state of a specific PRD loop. +func (e *Engine) GetState(name string) (loop.LoopState, int, error) { + return e.manager.GetState(name) +} + +// GetInstance returns a copy of the loop instance for a specific PRD. +func (e *Engine) GetInstance(name string) *loop.LoopInstance { + return e.manager.GetInstance(name) +} + +// GetAllInstances returns a snapshot of all loop instances. +func (e *Engine) GetAllInstances() []*loop.LoopInstance { + return e.manager.GetAllInstances() +} + +// GetRunningPRDs returns the names of all currently running PRDs. +func (e *Engine) GetRunningPRDs() []string { + return e.manager.GetRunningPRDs() +} + +// GetRunningCount returns the number of currently running loops. +func (e *Engine) GetRunningCount() int { + return e.manager.GetRunningCount() +} + +// IsAnyRunning returns true if any loop is currently running. +func (e *Engine) IsAnyRunning() bool { + return e.manager.IsAnyRunning() +} + +// SetMaxIterations updates the default max iterations for new loops. +func (e *Engine) SetMaxIterations(maxIter int) { + e.manager.SetMaxIterations(maxIter) +} + +// MaxIterations returns the current default max iterations. +func (e *Engine) MaxIterations() int { + return e.manager.MaxIterations() +} + +// SetMaxIterationsForInstance updates max iterations for a running loop. +func (e *Engine) SetMaxIterationsForInstance(name string, maxIter int) error { + return e.manager.SetMaxIterationsForInstance(name, maxIter) +} + +// SetRetryConfig sets the retry configuration for new loops. +func (e *Engine) SetRetryConfig(cfg loop.RetryConfig) { + e.manager.SetRetryConfig(cfg) +} + +// DisableRetry disables automatic retry for new loops. +func (e *Engine) DisableRetry() { + e.manager.DisableRetry() +} + +// SetCompletionCallback sets a callback for when any PRD completes. +func (e *Engine) SetCompletionCallback(fn func(prdName string)) { + e.manager.SetCompletionCallback(fn) +} + +// SetPostCompleteCallback sets a callback for post-completion actions. +func (e *Engine) SetPostCompleteCallback(fn func(prdName, branch, workDir string)) { + e.manager.SetPostCompleteCallback(fn) +} + +// SetConfig sets the project config. +func (e *Engine) SetConfig(cfg *config.Config) { + e.manager.SetConfig(cfg) +} + +// Config returns the current project config. +func (e *Engine) Config() *config.Config { + return e.manager.Config() +} + +// UpdateWorktreeInfo updates the worktree directory and branch for a PRD. +func (e *Engine) UpdateWorktreeInfo(name, worktreeDir, branch string) error { + return e.manager.UpdateWorktreeInfo(name, worktreeDir, branch) +} + +// ClearWorktreeInfo clears the worktree directory and optionally branch. +func (e *Engine) ClearWorktreeInfo(name string, clearBranch bool) error { + return e.manager.ClearWorktreeInfo(name, clearBranch) +} + +// --- Project state queries --- + +// LoadPRD loads and returns a PRD from the given path. +func (e *Engine) LoadPRD(prdPath string) (*prd.PRD, error) { + return prd.LoadPRD(prdPath) +} + +// Shutdown stops all loops and the event forwarding goroutine. +func (e *Engine) Shutdown() { + e.manager.StopAll() + + e.forwardMu.Lock() + defer e.forwardMu.Unlock() + + if e.forwarding { + close(e.stopForward) + e.forwarding = false + } +} diff --git a/internal/engine/engine_test.go b/internal/engine/engine_test.go new file mode 100644 index 00000000..65f6f756 --- /dev/null +++ b/internal/engine/engine_test.go @@ -0,0 +1,504 @@ +package engine + +import ( + "os" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/config" + "github.com/minicodemonkey/chief/internal/loop" +) + +// createTestPRD creates a minimal test PRD file and returns its path. +func createTestPRD(t *testing.T, dir, name string) string { + t.Helper() + prdDir := filepath.Join(dir, name) + if err := os.MkdirAll(prdDir, 0755); err != nil { + t.Fatal(err) + } + prdPath := filepath.Join(prdDir, "prd.json") + content := `{ + "project": "Test PRD", + "description": "Test", + "userStories": [ + {"id": "US-001", "title": "Test Story", "description": "Test", "priority": 1, "passes": false} + ] + }` + if err := os.WriteFile(prdPath, []byte(content), 0644); err != nil { + t.Fatal(err) + } + return prdPath +} + +func TestNew(t *testing.T) { + e := New(10) + defer e.Shutdown() + + if e == nil { + t.Fatal("expected non-nil engine") + } + if e.manager == nil { + t.Fatal("expected non-nil manager") + } + if e.MaxIterations() != 10 { + t.Errorf("expected maxIter 10, got %d", e.MaxIterations()) + } +} + +func TestRegisterAndGetInstance(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + if err := e.Register("test-prd", prdPath); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + instance := e.GetInstance("test-prd") + if instance == nil { + t.Fatal("expected instance to be registered") + } + if instance.Name != "test-prd" { + t.Errorf("expected name 'test-prd', got '%s'", instance.Name) + } + if instance.State != loop.LoopStateReady { + t.Errorf("expected state Ready, got %v", instance.State) + } +} + +func TestRegisterDuplicate(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + if err := e.Register("test-prd", prdPath); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + err := e.Register("test-prd", prdPath) + if err == nil { + t.Error("expected error when registering duplicate PRD") + } +} + +func TestUnregister(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + e.Register("test-prd", prdPath) + if err := e.Unregister("test-prd"); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if inst := e.GetInstance("test-prd"); inst != nil { + t.Error("expected instance to be removed") + } +} + +func TestSubscribeAndUnsubscribe(t *testing.T) { + e := New(10) + defer e.Shutdown() + + ch1, unsub1 := e.Subscribe() + ch2, unsub2 := e.Subscribe() + + if ch1 == nil || ch2 == nil { + t.Fatal("expected non-nil channels") + } + + e.subMu.RLock() + count := len(e.subscribers) + e.subMu.RUnlock() + if count != 2 { + t.Errorf("expected 2 subscribers, got %d", count) + } + + unsub1() + + e.subMu.RLock() + count = len(e.subscribers) + e.subMu.RUnlock() + if count != 1 { + t.Errorf("expected 1 subscriber after unsub, got %d", count) + } + + unsub2() + + e.subMu.RLock() + count = len(e.subscribers) + e.subMu.RUnlock() + if count != 0 { + t.Errorf("expected 0 subscribers after unsub, got %d", count) + } +} + +func TestMultipleSubscribersReceiveEvents(t *testing.T) { + e := New(10) + defer e.Shutdown() + + ch1, unsub1 := e.Subscribe() + defer unsub1() + ch2, unsub2 := e.Subscribe() + defer unsub2() + + // Inject an event directly into the manager's events channel for testing + // We do this by sending an event through the underlying manager + go func() { + // Send a synthetic event via the manager's event channel + e.manager.Events() + }() + + // Instead of trying to trigger a real loop event, test fan-out by + // directly injecting into the fan-out mechanism + testEvent := ManagerEvent{ + PRDName: "test", + Completed: false, + Event: loop.Event{ + Type: loop.EventIterationStart, + Text: "test event", + }, + } + + // Directly write to subscriber channels to verify wiring + e.subMu.RLock() + for _, ch := range e.subscribers { + ch <- testEvent + } + e.subMu.RUnlock() + + // Both subscribers should receive the event + select { + case ev := <-ch1: + if ev.PRDName != "test" { + t.Errorf("ch1: expected PRDName 'test', got '%s'", ev.PRDName) + } + case <-time.After(time.Second): + t.Error("ch1: timed out waiting for event") + } + + select { + case ev := <-ch2: + if ev.PRDName != "test" { + t.Errorf("ch2: expected PRDName 'test', got '%s'", ev.PRDName) + } + case <-time.After(time.Second): + t.Error("ch2: timed out waiting for event") + } +} + +func TestGetState(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + e.Register("test-prd", prdPath) + + state, iteration, err := e.GetState("test-prd") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if state != loop.LoopStateReady { + t.Errorf("expected Ready state, got %v", state) + } + if iteration != 0 { + t.Errorf("expected iteration 0, got %d", iteration) + } +} + +func TestGetAllInstances(t *testing.T) { + tmpDir := t.TempDir() + prd1 := createTestPRD(t, tmpDir, "prd1") + prd2 := createTestPRD(t, tmpDir, "prd2") + + e := New(10) + defer e.Shutdown() + + e.Register("prd1", prd1) + e.Register("prd2", prd2) + + instances := e.GetAllInstances() + if len(instances) != 2 { + t.Errorf("expected 2 instances, got %d", len(instances)) + } +} + +func TestGetRunningPRDs(t *testing.T) { + e := New(10) + defer e.Shutdown() + + running := e.GetRunningPRDs() + if len(running) != 0 { + t.Errorf("expected 0 running PRDs, got %d", len(running)) + } +} + +func TestIsAnyRunning(t *testing.T) { + e := New(10) + defer e.Shutdown() + + if e.IsAnyRunning() { + t.Error("expected no running loops") + } +} + +func TestSetMaxIterations(t *testing.T) { + e := New(10) + defer e.Shutdown() + + e.SetMaxIterations(20) + if e.MaxIterations() != 20 { + t.Errorf("expected 20, got %d", e.MaxIterations()) + } +} + +func TestSetConfig(t *testing.T) { + e := New(10) + defer e.Shutdown() + + if e.Config() != nil { + t.Error("expected nil config initially") + } + + cfg := &config.Config{ + OnComplete: config.OnCompleteConfig{Push: true}, + } + e.SetConfig(cfg) + + got := e.Config() + if got == nil || !got.OnComplete.Push { + t.Error("expected config with Push=true") + } +} + +func TestRetryConfig(t *testing.T) { + e := New(10) + defer e.Shutdown() + + e.SetRetryConfig(loop.RetryConfig{MaxRetries: 5, Enabled: true}) + e.DisableRetry() + // No assertion on internal state; just verify no panic +} + +func TestSetCompletionCallback(t *testing.T) { + e := New(10) + defer e.Shutdown() + + called := false + e.SetCompletionCallback(func(prdName string) { + called = true + }) + + // Manually trigger via manager to verify it's wired + e.manager.SetCompletionCallback(func(prdName string) { + called = true + }) + // The callback is set on the manager, verify it + if called { + t.Error("callback should not be called yet") + } +} + +func TestRegisterWithWorktree(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + err := e.RegisterWithWorktree("test-prd", prdPath, "/tmp/wt", "branch") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + inst := e.GetInstance("test-prd") + if inst.WorktreeDir != "/tmp/wt" { + t.Errorf("expected WorktreeDir '/tmp/wt', got '%s'", inst.WorktreeDir) + } + if inst.Branch != "branch" { + t.Errorf("expected Branch 'branch', got '%s'", inst.Branch) + } +} + +func TestUpdateAndClearWorktreeInfo(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + e.Register("test-prd", prdPath) + e.UpdateWorktreeInfo("test-prd", "/tmp/wt", "branch") + + inst := e.GetInstance("test-prd") + if inst.WorktreeDir != "/tmp/wt" { + t.Errorf("expected '/tmp/wt', got '%s'", inst.WorktreeDir) + } + + e.ClearWorktreeInfo("test-prd", true) + inst = e.GetInstance("test-prd") + if inst.WorktreeDir != "" || inst.Branch != "" { + t.Error("expected cleared worktree info") + } +} + +func TestManagerAccess(t *testing.T) { + e := New(10) + defer e.Shutdown() + + if e.Manager() == nil { + t.Error("expected non-nil manager") + } +} + +func TestLoadPRD(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + p, err := e.LoadPRD(prdPath) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if p.Project != "Test PRD" { + t.Errorf("expected 'Test PRD', got '%s'", p.Project) + } +} + +func TestStopAll(t *testing.T) { + tmpDir := t.TempDir() + prd1 := createTestPRD(t, tmpDir, "prd1") + prd2 := createTestPRD(t, tmpDir, "prd2") + + e := New(10) + defer e.Shutdown() + + e.Register("prd1", prd1) + e.Register("prd2", prd2) + + done := make(chan struct{}) + go func() { + e.StopAll() + close(done) + }() + + select { + case <-done: + case <-time.After(time.Second): + t.Error("StopAll did not complete in time") + } +} + +func TestShutdown(t *testing.T) { + e := New(10) + e.Shutdown() + + // Verify forwarding is stopped + e.forwardMu.Lock() + forwarding := e.forwarding + e.forwardMu.Unlock() + + if forwarding { + t.Error("expected forwarding to be stopped after shutdown") + } +} + +func TestConcurrentSubscribeUnsubscribe(t *testing.T) { + e := New(10) + defer e.Shutdown() + + var wg sync.WaitGroup + for i := 0; i < 50; i++ { + wg.Add(1) + go func() { + defer wg.Done() + _, unsub := e.Subscribe() + time.Sleep(time.Millisecond) + unsub() + }() + } + wg.Wait() + + e.subMu.RLock() + count := len(e.subscribers) + e.subMu.RUnlock() + if count != 0 { + t.Errorf("expected 0 subscribers after all unsubscribed, got %d", count) + } +} + +func TestConcurrentAccess(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + e.Register("test-prd", prdPath) + + var wg sync.WaitGroup + for i := 0; i < 100; i++ { + wg.Add(1) + go func() { + defer wg.Done() + _ = e.GetInstance("test-prd") + _ = e.GetAllInstances() + _ = e.GetRunningPRDs() + _ = e.GetRunningCount() + _, _, _ = e.GetState("test-prd") + _ = e.IsAnyRunning() + _ = e.MaxIterations() + }() + } + wg.Wait() +} + +func TestStartNonExistent(t *testing.T) { + e := New(10) + defer e.Shutdown() + + err := e.Start("nonexistent") + if err == nil { + t.Error("expected error when starting non-existent PRD") + } +} + +func TestPauseNonRunning(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + e.Register("test-prd", prdPath) + err := e.Pause("test-prd") + if err == nil { + t.Error("expected error when pausing non-running PRD") + } +} + +func TestStopNonRunning(t *testing.T) { + tmpDir := t.TempDir() + prdPath := createTestPRD(t, tmpDir, "test-prd") + + e := New(10) + defer e.Shutdown() + + e.Register("test-prd", prdPath) + err := e.Stop("test-prd") + if err != nil { + t.Errorf("stop non-running should not error: %v", err) + } +} diff --git a/internal/loop/codex_parser.go b/internal/loop/codex_parser.go deleted file mode 100644 index d9346d02..00000000 --- a/internal/loop/codex_parser.go +++ /dev/null @@ -1,125 +0,0 @@ -package loop - -import ( - "encoding/json" - "errors" - "strings" -) - -// codexEvent represents the top-level structure of a Codex exec --json JSONL line. -type codexEvent struct { - Type string `json:"type"` - Item *codexItem `json:"item,omitempty"` - Message string `json:"message,omitempty"` // top-level for type "error" - Error *struct { - Message string `json:"message"` - } `json:"error,omitempty"` -} - -// codexItem represents an item in item.started / item.completed / item.updated events. -type codexItem struct { - ID string `json:"id"` - Type string `json:"type"` - Text string `json:"text,omitempty"` - Command string `json:"command,omitempty"` - AggregatedOutput string `json:"aggregated_output,omitempty"` - ExitCode *int `json:"exit_code,omitempty"` - Status string `json:"status,omitempty"` - Server string `json:"server,omitempty"` - Tool string `json:"tool,omitempty"` -} - -// ParseLineCodex parses a single line of Codex exec --json JSONL output and returns an Event. -// If the line cannot be parsed or is not relevant, it returns nil. -func ParseLineCodex(line string) *Event { - line = strings.TrimSpace(line) - if line == "" { - return nil - } - - var ev codexEvent - if err := json.Unmarshal([]byte(line), &ev); err != nil { - return nil - } - - switch ev.Type { - case "thread.started", "turn.started": - return &Event{Type: EventIterationStart} - - case "turn.failed": - msg := "" - if ev.Error != nil { - msg = ev.Error.Message - } - return &Event{Type: EventError, Err: errors.New(msg)} - - case "error": - msg := ev.Message - if msg == "" && ev.Error != nil { - msg = ev.Error.Message - } - if msg == "" { - msg = "unknown error" - } - return &Event{Type: EventError, Err: errors.New(msg)} - - case "item.started": - if ev.Item == nil { - return nil - } - switch ev.Item.Type { - case "command_execution": - return &Event{ - Type: EventToolStart, - Tool: ev.Item.Command, - } - case "mcp_tool_call": - toolName := ev.Item.Tool - if ev.Item.Server != "" { - toolName = ev.Item.Server + "/" + ev.Item.Tool - } - return &Event{ - Type: EventToolStart, - Tool: toolName, - } - } - return nil - - case "item.completed": - if ev.Item == nil { - return nil - } - switch ev.Item.Type { - case "command_execution": - return &Event{ - Type: EventToolResult, - Text: ev.Item.AggregatedOutput, - } - case "mcp_tool_call": - return &Event{ - Type: EventToolResult, - Text: ev.Item.AggregatedOutput, - } - case "agent_message": - text := ev.Item.Text - if strings.Contains(text, "<chief-done/>") { - return &Event{Type: EventStoryDone, Text: text} - } - return &Event{Type: EventAssistantText, Text: text} - case "file_change": - return &Event{ - Type: EventToolResult, - Tool: "file_change", - Text: ev.Item.AggregatedOutput, - } - } - return nil - - case "turn.completed": - // Usage info only, no event - return nil - - default: - return nil - } -} diff --git a/internal/loop/codex_parser_test.go b/internal/loop/codex_parser_test.go deleted file mode 100644 index 3afd4ca7..00000000 --- a/internal/loop/codex_parser_test.go +++ /dev/null @@ -1,151 +0,0 @@ -package loop - -import ( - "testing" -) - -func TestParseLineCodex_threadStarted(t *testing.T) { - line := `{"type":"thread.started","thread_id":"0199a213-81c0-7800-8aa1-bbab2a035a53"}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventIterationStart { - t.Errorf("expected EventIterationStart, got %v", ev.Type) - } -} - -func TestParseLineCodex_turnStarted(t *testing.T) { - line := `{"type":"turn.started"}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventIterationStart { - t.Errorf("expected EventIterationStart, got %v", ev.Type) - } -} - -func TestParseLineCodex_commandExecutionStarted(t *testing.T) { - line := `{"type":"item.started","item":{"id":"item_1","type":"command_execution","command":"bash -lc ls","aggregated_output":"","exit_code":null,"status":"in_progress"}}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "bash -lc ls" { - t.Errorf("expected Tool bash -lc ls, got %q", ev.Tool) - } -} - -func TestParseLineCodex_commandExecutionCompleted(t *testing.T) { - line := `{"type":"item.completed","item":{"id":"item_1","type":"command_execution","command":"bash -lc ls","aggregated_output":"docs\nsrc\n","exit_code":0,"status":"completed"}}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev.Type) - } - if ev.Text != "docs\nsrc\n" { - t.Errorf("expected Text docs\\nsrc\\n, got %q", ev.Text) - } -} - -func TestParseLineCodex_agentMessageWithChiefDoneTag(t *testing.T) { - line := `{"type":"item.completed","item":{"id":"item_3","type":"agent_message","text":"Done. <chief-done/>"}}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventStoryDone { - t.Errorf("expected EventStoryDone, got %v", ev.Type) - } -} - -func TestParseLineCodex_agentMessageWithChiefDone(t *testing.T) { - line := `{"type":"item.completed","item":{"id":"item_3","type":"agent_message","text":"All criteria pass. <chief-done/>"}}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventStoryDone { - t.Errorf("expected EventStoryDone, got %v", ev.Type) - } -} - -func TestParseLineCodex_agentMessagePlain(t *testing.T) { - line := `{"type":"item.completed","item":{"id":"item_3","type":"agent_message","text":"Done. I updated the docs."}}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventAssistantText { - t.Errorf("expected EventAssistantText, got %v", ev.Type) - } - if ev.Text != "Done. I updated the docs." { - t.Errorf("unexpected Text: %q", ev.Text) - } -} - -func TestParseLineCodex_turnFailed(t *testing.T) { - line := `{"type":"turn.failed","error":{"message":"model response stream ended unexpectedly"}}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventError { - t.Errorf("expected EventError, got %v", ev.Type) - } - if ev.Err == nil { - t.Fatal("expected Err set") - } - if ev.Err.Error() != "model response stream ended unexpectedly" { - t.Errorf("unexpected Err: %v", ev.Err) - } -} - -func TestParseLineCodex_error(t *testing.T) { - line := `{"type":"error","message":"stream error: broken pipe"}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventError { - t.Errorf("expected EventError, got %v", ev.Type) - } -} - -func TestParseLineCodex_turnCompleted_ignored(t *testing.T) { - line := `{"type":"turn.completed","usage":{"input_tokens":24763,"cached_input_tokens":24448,"output_tokens":122}}` - ev := ParseLineCodex(line) - if ev != nil { - t.Errorf("expected nil (ignore turn.completed), got %v", ev) - } -} - -func TestParseLineCodex_mcpToolCallStarted(t *testing.T) { - line := `{"type":"item.started","item":{"id":"item_5","type":"mcp_tool_call","server":"docs","tool":"search","arguments":{"q":"exec --json"},"result":null,"error":null,"status":"in_progress"}}` - ev := ParseLineCodex(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "docs/search" { - t.Errorf("expected Tool docs/search, got %q", ev.Tool) - } -} - -func TestParseLineCodex_emptyOrInvalid_returnsNil(t *testing.T) { - tests := []string{"", " ", "not json", "{}", `{"type":"unknown"}`} - for _, line := range tests { - ev := ParseLineCodex(line) - if ev != nil { - t.Errorf("ParseLineCodex(%q) expected nil, got %v", line, ev) - } - } -} diff --git a/internal/loop/cursor_parser.go b/internal/loop/cursor_parser.go deleted file mode 100644 index 993123d0..00000000 --- a/internal/loop/cursor_parser.go +++ /dev/null @@ -1,421 +0,0 @@ -package loop - -import ( - "encoding/json" - "fmt" - "strings" -) - -// cursorEvent represents the top-level structure of Cursor CLI stream-json NDJSON. -type cursorEvent struct { - Type string `json:"type"` - Subtype string `json:"subtype,omitempty"` - Message json.RawMessage `json:"message,omitempty"` - ToolCall json.RawMessage `json:"tool_call,omitempty"` -} - -// cursorAssistantMessage is the message body for type "assistant". -type cursorAssistantMessage struct { - Role string `json:"role"` - Content []cursorContentBlock `json:"content"` -} - -// cursorContentBlock is a content block in an assistant message. -type cursorContentBlock struct { - Type string `json:"type"` - Text string `json:"text,omitempty"` -} - -// cursorToolCall is the tool_call object. -type cursorToolCall struct { - ReadToolCall *cursorReadToolCall `json:"readToolCall,omitempty"` - WriteToolCall *cursorWriteToolCall `json:"writeToolCall,omitempty"` - EditToolCall *cursorEditToolCall `json:"editToolCall,omitempty"` - ShellToolCall *cursorShellToolCall `json:"shellToolCall,omitempty"` - GrepToolCall *cursorGrepToolCall `json:"grepToolCall,omitempty"` - GlobToolCall *cursorGlobToolCall `json:"globToolCall,omitempty"` - LsToolCall *cursorLsToolCall `json:"lsToolCall,omitempty"` - DeleteToolCall *cursorDeleteToolCall `json:"deleteToolCall,omitempty"` - WebFetchToolCall *cursorWebFetchToolCall `json:"webFetchToolCall,omitempty"` - WebSearchToolCall *cursorWebSearchToolCall `json:"webSearchToolCall,omitempty"` - Function *cursorFunctionCall `json:"function,omitempty"` -} - -// cursorWebFetchToolCall holds web fetch args and optional result. -type cursorWebFetchToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct { - URL string `json:"url"` - Markdown string `json:"markdown"` - } `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorWebSearchToolCall holds web search args and optional result. -type cursorWebSearchToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct { - References []struct { - Title string `json:"title"` - URL string `json:"url"` - Chunk string `json:"chunk"` - } `json:"references,omitempty"` - } `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorEditToolCall holds edit/strreplace args and optional result. -type cursorEditToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct{} `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorShellToolCall holds shell command args and optional result. -type cursorShellToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct { - ExitCode *int `json:"exitCode,omitempty"` - Output string `json:"output,omitempty"` - } `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorGrepToolCall holds grep args and optional result. -type cursorGrepToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct { - WorkspaceResults map[string]struct { - Content struct { - TotalMatchedLines int `json:"totalMatchedLines"` - } `json:"content"` - } `json:"workspaceResults,omitempty"` - } `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorGlobToolCall holds glob args and optional result. -type cursorGlobToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct { - TotalFiles int `json:"totalFiles"` - } `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorLsToolCall holds list-directory args and optional result. -type cursorLsToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct{} `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorDeleteToolCall holds delete-file args and optional result. -type cursorDeleteToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct{} `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorReadToolCall holds read file args and optional result. -type cursorReadToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct { - Content string `json:"content"` - } `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorWriteToolCall holds write file args and optional result. -type cursorWriteToolCall struct { - Args map[string]interface{} `json:"args,omitempty"` - Result *struct { - Success *struct { - Path string `json:"path"` - LinesCreated int `json:"linesCreated"` - FileSize int `json:"fileSize"` - } `json:"success,omitempty"` - } `json:"result,omitempty"` -} - -// cursorFunctionCall holds generic function name, arguments, and optional result. -type cursorFunctionCall struct { - Name string `json:"name,omitempty"` - Arguments string `json:"arguments,omitempty"` - Result json.RawMessage `json:"result,omitempty"` -} - -// ParseLineCursor parses a single line of Cursor CLI stream-json NDJSON and returns an Event. -// If the line cannot be parsed or is not relevant, it returns nil. -func ParseLineCursor(line string) *Event { - line = strings.TrimSpace(line) - if line == "" { - return nil - } - - var ev cursorEvent - if err := json.Unmarshal([]byte(line), &ev); err != nil { - return nil - } - - switch ev.Type { - case "system": - if ev.Subtype == "init" { - return &Event{Type: EventIterationStart} - } - return nil - - case "assistant": - return parseCursorAssistantMessage(ev.Message) - - case "tool_call": - return parseCursorToolCall(ev.Subtype, ev.ToolCall) - - case "user", "result": - return nil - - default: - return nil - } -} - -func parseCursorAssistantMessage(raw json.RawMessage) *Event { - if raw == nil { - return nil - } - var msg cursorAssistantMessage - if err := json.Unmarshal(raw, &msg); err != nil { - return nil - } - for _, block := range msg.Content { - if block.Type != "text" { - continue - } - text := block.Text - if strings.Contains(text, "<chief-complete/>") { - return &Event{Type: EventComplete, Text: text} - } - if strings.Contains(text, "<chief-done/>") { - return &Event{Type: EventStoryDone, Text: text} - } - return &Event{Type: EventAssistantText, Text: text} - } - return nil -} - -func parseCursorToolCall(subtype string, raw json.RawMessage) *Event { - if raw == nil { - return nil - } - var tc cursorToolCall - if err := json.Unmarshal(raw, &tc); err != nil { - return nil - } - toolName, toolInput := cursorToolCallNameAndInput(&tc) - switch subtype { - case "started": - return &Event{Type: EventToolStart, Tool: toolName, ToolInput: toolInput} - case "completed": - text := cursorToolCallResultSummary(&tc) - return &Event{Type: EventToolResult, Tool: toolName, Text: text} - } - return nil -} - -// cursorToolCallNameAndInput returns display name (PascalCase for TUI icons) and optional ToolInput for the log. -func cursorToolCallNameAndInput(tc *cursorToolCall) (name string, input map[string]interface{}) { - if tc.ReadToolCall != nil { - input = make(map[string]interface{}) - if path, ok := tc.ReadToolCall.Args["path"].(string); ok { - input["file_path"] = path - } - return "Read", input - } - if tc.WriteToolCall != nil { - input = make(map[string]interface{}) - if path, ok := tc.WriteToolCall.Args["path"].(string); ok { - input["file_path"] = path - } - return "Write", input - } - if tc.EditToolCall != nil { - input = make(map[string]interface{}) - if path, ok := tc.EditToolCall.Args["path"].(string); ok { - input["file_path"] = path - } - return "Edit", input - } - if tc.ShellToolCall != nil { - input = make(map[string]interface{}) - if cmd, ok := tc.ShellToolCall.Args["command"].(string); ok { - input["command"] = cmd - } - return "Bash", input - } - if tc.GrepToolCall != nil { - input = make(map[string]interface{}) - if pattern, ok := tc.GrepToolCall.Args["pattern"].(string); ok { - input["pattern"] = pattern - } - if path, ok := tc.GrepToolCall.Args["path"].(string); ok { - input["path"] = path - } - return "Grep", input - } - if tc.GlobToolCall != nil { - input = make(map[string]interface{}) - if pattern, ok := tc.GlobToolCall.Args["globPattern"].(string); ok { - input["pattern"] = pattern - } - if dir, ok := tc.GlobToolCall.Args["targetDirectory"].(string); ok { - input["path"] = dir - } - return "Glob", input - } - if tc.LsToolCall != nil { - input = make(map[string]interface{}) - if path, ok := tc.LsToolCall.Args["path"].(string); ok { - input["path"] = path - } - return "List", input - } - if tc.DeleteToolCall != nil { - input = make(map[string]interface{}) - if path, ok := tc.DeleteToolCall.Args["path"].(string); ok { - input["file_path"] = path - } - return "Delete", input - } - if tc.WebFetchToolCall != nil { - input = make(map[string]interface{}) - if url, ok := tc.WebFetchToolCall.Args["url"].(string); ok { - input["url"] = url - } - return "WebFetch", input - } - if tc.WebSearchToolCall != nil { - input = make(map[string]interface{}) - if term, ok := tc.WebSearchToolCall.Args["searchTerm"].(string); ok { - input["query"] = term - } - return "WebSearch", input - } - if tc.Function != nil && tc.Function.Name != "" { - // TUI knows "Bash" for command execution; Cursor may use different names - name = tc.Function.Name - if name == "run_terminal_cmd" || name == "run_command" { - name = "Bash" - } - if tc.Function.Arguments != "" { - input = map[string]interface{}{"arguments": tc.Function.Arguments} - // Try to extract command for Bash display - var argsMap map[string]interface{} - if json.Unmarshal([]byte(tc.Function.Arguments), &argsMap) == nil { - if cmd, ok := argsMap["command"].(string); ok { - input["command"] = cmd - } - } - } - return name, input - } - return "tool", nil -} - -func cursorToolCallResultSummary(tc *cursorToolCall) string { - if tc.ReadToolCall != nil && tc.ReadToolCall.Result != nil && tc.ReadToolCall.Result.Success != nil { - return tc.ReadToolCall.Result.Success.Content - } - if tc.WriteToolCall != nil && tc.WriteToolCall.Result != nil && tc.WriteToolCall.Result.Success != nil { - s := tc.WriteToolCall.Result.Success - if s.Path != "" { - return s.Path - } - return "(written)" - } - if tc.EditToolCall != nil && tc.EditToolCall.Result != nil && tc.EditToolCall.Result.Success != nil { - return "(edited)" - } - if tc.ShellToolCall != nil && tc.ShellToolCall.Result != nil && tc.ShellToolCall.Result.Success != nil { - s := tc.ShellToolCall.Result.Success - if s.Output != "" { - return strings.TrimSpace(s.Output) - } - if s.ExitCode != nil { - return fmt.Sprintf("(exit %d)", *s.ExitCode) - } - return "(executed)" - } - if tc.GrepToolCall != nil && tc.GrepToolCall.Result != nil && tc.GrepToolCall.Result.Success != nil { - for _, v := range tc.GrepToolCall.Result.Success.WorkspaceResults { - return fmt.Sprintf("%d matches", v.Content.TotalMatchedLines) - } - return "(matches)" - } - if tc.GlobToolCall != nil && tc.GlobToolCall.Result != nil && tc.GlobToolCall.Result.Success != nil { - n := tc.GlobToolCall.Result.Success.TotalFiles - return fmt.Sprintf("%d files", n) - } - if tc.LsToolCall != nil && tc.LsToolCall.Result != nil && tc.LsToolCall.Result.Success != nil { - return "(listed)" - } - if tc.DeleteToolCall != nil && tc.DeleteToolCall.Result != nil && tc.DeleteToolCall.Result.Success != nil { - return "(deleted)" - } - if tc.WebFetchToolCall != nil && tc.WebFetchToolCall.Result != nil && tc.WebFetchToolCall.Result.Success != nil { - s := tc.WebFetchToolCall.Result.Success - if s.Markdown != "" { - return strings.TrimSpace(s.Markdown) - } - return "(fetched)" - } - if tc.WebSearchToolCall != nil && tc.WebSearchToolCall.Result != nil && tc.WebSearchToolCall.Result.Success != nil { - refs := tc.WebSearchToolCall.Result.Success.References - if len(refs) == 0 { - return "(no results)" - } - if len(refs) == 1 && refs[0].Chunk != "" { - return strings.TrimSpace(refs[0].Chunk) - } - return fmt.Sprintf("%d reference(s)", len(refs)) - } - if tc.Function != nil && len(tc.Function.Result) > 0 { - s := extractFunctionResultText(tc.Function.Result) - if s != "" { - return s - } - } - if tc.Function != nil { - return "(executed)" - } - return "" -} - -// extractFunctionResultText tries to get a short result string from Cursor function result JSON. -func extractFunctionResultText(raw json.RawMessage) string { - var m map[string]interface{} - if json.Unmarshal(raw, &m) != nil { - return "" - } - for _, key := range []string{"output", "content", "result", "stdout", "text"} { - if v, ok := m[key].(string); ok && v != "" { - return v - } - } - if success, ok := m["success"].(map[string]interface{}); ok { - for _, key := range []string{"output", "content", "result", "stdout"} { - if v, ok := success[key].(string); ok && v != "" { - return v - } - } - } - return "" -} diff --git a/internal/loop/cursor_parser_test.go b/internal/loop/cursor_parser_test.go deleted file mode 100644 index 54bf8605..00000000 --- a/internal/loop/cursor_parser_test.go +++ /dev/null @@ -1,322 +0,0 @@ -package loop - -import ( - "strings" - "testing" -) - -func TestParseLineCursor_systemInit(t *testing.T) { - line := `{"type":"system","subtype":"init","apiKeySource":"login","cwd":"/Users/user/project","session_id":"c6b62c6f-7ead-4fd6-9922-e952131177ff","model":"Claude 4 Sonnet","permissionMode":"default"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventIterationStart { - t.Errorf("expected EventIterationStart, got %v", ev.Type) - } -} - -func TestParseLineCursor_assistantText(t *testing.T) { - line := `{"type":"assistant","message":{"role":"assistant","content":[{"type":"text","text":"I'll read the README.md file"}]},"session_id":"c6b62c6f-7ead-4fd6-9922-e952131177ff"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventAssistantText { - t.Errorf("expected EventAssistantText, got %v", ev.Type) - } - if ev.Text != "I'll read the README.md file" { - t.Errorf("expected Text, got %q", ev.Text) - } -} - -func TestParseLineCursor_chiefComplete(t *testing.T) { - line := `{"type":"assistant","message":{"role":"assistant","content":[{"type":"text","text":"Done. <chief-complete/>"}]},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventComplete { - t.Errorf("expected EventComplete, got %v", ev.Type) - } -} - -func TestParseLineCursor_chiefDone(t *testing.T) { - line := `{"type":"assistant","message":{"role":"assistant","content":[{"type":"text","text":"Story complete. <chief-done/>"}]},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventStoryDone { - t.Errorf("expected EventStoryDone, got %v", ev.Type) - } -} - -func TestParseLineCursor_toolCallStartedRead(t *testing.T) { - line := `{"type":"tool_call","subtype":"started","call_id":"toolu_abc","tool_call":{"readToolCall":{"args":{"path":"README.md"}}},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "Read" { - t.Errorf("expected Tool Read, got %q", ev.Tool) - } - if ev.ToolInput == nil || ev.ToolInput["file_path"] != "README.md" { - t.Errorf("expected ToolInput file_path=README.md, got %v", ev.ToolInput) - } -} - -func TestParseLineCursor_toolCallStartedWrite(t *testing.T) { - line := `{"type":"tool_call","subtype":"started","call_id":"toolu_xyz","tool_call":{"writeToolCall":{"args":{"path":"summary.txt","fileText":"content"}}},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "Write" { - t.Errorf("expected Tool Write, got %q", ev.Tool) - } - if ev.ToolInput == nil || ev.ToolInput["file_path"] != "summary.txt" { - t.Errorf("expected ToolInput file_path=summary.txt, got %v", ev.ToolInput) - } -} - -func TestParseLineCursor_toolCallCompletedRead(t *testing.T) { - line := `{"type":"tool_call","subtype":"completed","call_id":"toolu_abc","tool_call":{"readToolCall":{"args":{"path":"README.md"},"result":{"success":{"content":"# Project\n\nContent here.","isEmpty":false,"exceededLimit":false,"totalLines":10,"totalChars":100}}}},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev.Type) - } - if ev.Tool != "Read" { - t.Errorf("expected Tool Read, got %q", ev.Tool) - } - if ev.Text != "# Project\n\nContent here." { - t.Errorf("expected Text content, got %q", ev.Text) - } -} - -func TestParseLineCursor_toolCallCompletedWrite(t *testing.T) { - line := `{"type":"tool_call","subtype":"completed","call_id":"toolu_xyz","tool_call":{"writeToolCall":{"args":{"path":"summary.txt"},"result":{"success":{"path":"/Users/user/project/summary.txt","linesCreated":19,"fileSize":942}}}},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev.Type) - } - if ev.Tool != "Write" { - t.Errorf("expected Tool Write, got %q", ev.Tool) - } - if ev.Text != "/Users/user/project/summary.txt" { - t.Errorf("expected Text path, got %q", ev.Text) - } -} - -func TestParseLineCursor_toolCallFunctionWithResult(t *testing.T) { - line := `{"type":"tool_call","subtype":"completed","call_id":"toolu_fn","tool_call":{"function":{"name":"run_terminal_cmd","arguments":"{}","result":{"success":{"output":"hello world"}}}},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev.Type) - } - if ev.Text != "hello world" { - t.Errorf("expected Text from result.success.output, got %q", ev.Text) - } -} - -func TestParseLineCursor_toolCallFunctionNoResult(t *testing.T) { - line := `{"type":"tool_call","subtype":"completed","call_id":"toolu_fn2","tool_call":{"function":{"name":"some_tool","arguments":"{}"}},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev.Type) - } - if ev.Text != "(executed)" { - t.Errorf("expected Text (executed) when no result, got %q", ev.Text) - } -} - -func TestParseLineCursor_userAndResult_ignored(t *testing.T) { - userLine := `{"type":"user","message":{"role":"user","content":[{"type":"text","text":"prompt"}]},"session_id":"x"}` - if ev := ParseLineCursor(userLine); ev != nil { - t.Errorf("user event expected nil, got %v", ev) - } - resultLine := `{"type":"result","subtype":"success","duration_ms":1234,"result":"done","session_id":"x"}` - if ev := ParseLineCursor(resultLine); ev != nil { - t.Errorf("result event expected nil, got %v", ev) - } -} - -func TestParseLineCursor_toolCallEditStartedAndCompleted(t *testing.T) { - started := `{"type":"tool_call","subtype":"started","call_id":"toolu_edit","tool_call":{"editToolCall":{"args":{"path":"index.html"}}},"session_id":"x"}` - ev := ParseLineCursor(started) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "Edit" { - t.Errorf("expected Tool Edit, got %q", ev.Tool) - } - if ev.ToolInput == nil || ev.ToolInput["file_path"] != "index.html" { - t.Errorf("expected ToolInput file_path=index.html, got %v", ev.ToolInput) - } - - completed := `{"type":"tool_call","subtype":"completed","call_id":"toolu_edit","tool_call":{"editToolCall":{"args":{"path":"index.html"},"result":{"success":{}}}},"session_id":"x"}` - ev2 := ParseLineCursor(completed) - if ev2 == nil { - t.Fatal("expected event, got nil") - } - if ev2.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev2.Type) - } - if ev2.Tool != "Edit" { - t.Errorf("expected Tool Edit, got %q", ev2.Tool) - } - if ev2.Text != "(edited)" { - t.Errorf("expected Text (edited), got %q", ev2.Text) - } -} - -func TestParseLineCursor_toolCallShellStartedAndCompleted(t *testing.T) { - started := `{"type":"tool_call","subtype":"started","call_id":"toolu_shell","tool_call":{"shellToolCall":{"args":{"command":"echo hello"}}},"session_id":"x"}` - ev := ParseLineCursor(started) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "Bash" { - t.Errorf("expected Tool Bash, got %q", ev.Tool) - } - if ev.ToolInput == nil || ev.ToolInput["command"] != "echo hello" { - t.Errorf("expected ToolInput command=echo hello, got %v", ev.ToolInput) - } - - completed := `{"type":"tool_call","subtype":"completed","call_id":"toolu_shell","tool_call":{"shellToolCall":{"args":{"command":"echo hello"},"result":{"success":{"exitCode":0,"output":"hello\n"}}}},"session_id":"x"}` - ev2 := ParseLineCursor(completed) - if ev2 == nil { - t.Fatal("expected event, got nil") - } - if ev2.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev2.Type) - } - if ev2.Tool != "Bash" { - t.Errorf("expected Tool Bash, got %q", ev2.Tool) - } - if ev2.Text != "hello" { - t.Errorf("expected Text hello, got %q", ev2.Text) - } -} - -func TestParseLineCursor_toolCallWebSearchStartedAndCompleted(t *testing.T) { - started := `{"type":"tool_call","subtype":"started","call_id":"toolu_01SFCs5FKmApRiaNPKy3BBqi","tool_call":{"webSearchToolCall":{"args":{"searchTerm":"place.horse API random horse image","toolCallId":"toolu_01SFCs5FKmApRiaNPKy3BBqi"}}},"session_id":"x"}` - ev := ParseLineCursor(started) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "WebSearch" { - t.Errorf("expected Tool WebSearch, got %q", ev.Tool) - } - if ev.ToolInput == nil || ev.ToolInput["query"] != "place.horse API random horse image" { - t.Errorf("expected ToolInput query=place.horse API random horse image, got %v", ev.ToolInput) - } - - completed := `{"type":"tool_call","subtype":"completed","call_id":"toolu_01SFCs5FKmApRiaNPKy3BBqi","tool_call":{"webSearchToolCall":{"args":{"searchTerm":"place.horse API random horse image","toolCallId":"toolu_01SFCs5FKmApRiaNPKy3BBqi"},"result":{"success":{"references":[{"title":"Web search results for query: place.horse API random horse image","url":"","chunk":"Links:\n1. [API docs](https://theponyapi.com/docs)\n2. [Lorem Picsum](http://picsum.photos/)\n\nBased on your query about a place.horse API..."}]}}}},"session_id":"x"}` - ev2 := ParseLineCursor(completed) - if ev2 == nil { - t.Fatal("expected event, got nil") - } - if ev2.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev2.Type) - } - if ev2.Tool != "WebSearch" { - t.Errorf("expected Tool WebSearch, got %q", ev2.Tool) - } - // Single reference with chunk: summary is the chunk text - if ev2.Text == "" { - t.Errorf("expected non-empty Text (chunk), got %q", ev2.Text) - } - if !strings.Contains(ev2.Text, "Based on your query") { - t.Errorf("expected Text to contain chunk content, got %q", ev2.Text) - } -} - -func TestParseLineCursor_toolCallWebSearchCompletedMultipleRefs(t *testing.T) { - line := `{"type":"tool_call","subtype":"completed","call_id":"toolu_x","tool_call":{"webSearchToolCall":{"result":{"success":{"references":[{"title":"A","url":"","chunk":""},{"title":"B","url":"","chunk":""}]}}}},"session_id":"x"}` - ev := ParseLineCursor(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev.Type) - } - if ev.Tool != "WebSearch" { - t.Errorf("expected Tool WebSearch, got %q", ev.Tool) - } - if ev.Text != "2 reference(s)" { - t.Errorf("expected Text 2 reference(s), got %q", ev.Text) - } -} - -func TestParseLineCursor_toolCallWebFetchStartedAndCompleted(t *testing.T) { - started := `{"type":"tool_call","subtype":"started","call_id":"toolu_01SALput2Tb7iCNqx4jfy8v2","tool_call":{"webFetchToolCall":{"args":{"url":"https://github.com/treboryx/animalsAPI","toolCallId":"toolu_01SALput2Tb7iCNqx4jfy8v2"}}},"session_id":"x"}` - ev := ParseLineCursor(started) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "WebFetch" { - t.Errorf("expected Tool WebFetch, got %q", ev.Tool) - } - if ev.ToolInput == nil || ev.ToolInput["url"] != "https://github.com/treboryx/animalsAPI" { - t.Errorf("expected ToolInput url=https://github.com/treboryx/animalsAPI, got %v", ev.ToolInput) - } - - completed := `{"type":"tool_call","subtype":"completed","call_id":"toolu_01SALput2Tb7iCNqx4jfy8v2","tool_call":{"webFetchToolCall":{"args":{"url":"https://github.com/treboryx/animalsAPI","toolCallId":"toolu_01SALput2Tb7iCNqx4jfy8v2"},"result":{"success":{"url":"https://github.com/treboryx/animalsAPI","markdown":"# treboryx/animalsAPI\n\nAll-in-one API for random animal images."}}}},"session_id":"x"}` - ev2 := ParseLineCursor(completed) - if ev2 == nil { - t.Fatal("expected event, got nil") - } - if ev2.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev2.Type) - } - if ev2.Tool != "WebFetch" { - t.Errorf("expected Tool WebFetch, got %q", ev2.Tool) - } - if ev2.Text == "" { - t.Errorf("expected non-empty Text (markdown), got %q", ev2.Text) - } - if !strings.Contains(ev2.Text, "treboryx/animalsAPI") { - t.Errorf("expected Text to contain markdown content, got %q", ev2.Text) - } -} - -func TestParseLineCursor_emptyOrInvalid_returnsNil(t *testing.T) { - tests := []string{"", " ", "not json", "{}", `{"type":"unknown"}`} - for _, line := range tests { - ev := ParseLineCursor(line) - if ev != nil { - t.Errorf("ParseLineCursor(%q) expected nil, got %v", line, ev) - } - } -} diff --git a/internal/loop/loop.go b/internal/loop/loop.go index c85717f0..77680697 100644 --- a/internal/loop/loop.go +++ b/internal/loop/loop.go @@ -6,14 +6,15 @@ package loop import ( "bufio" + "bytes" "context" + "errors" "fmt" "io" "os" "os/exec" "path/filepath" "sync" - "sync/atomic" "time" "github.com/minicodemonkey/chief/embed" @@ -22,14 +23,11 @@ import ( // RetryConfig configures automatic retry behavior on Claude crashes. type RetryConfig struct { - MaxRetries int // Maximum number of retry attempts (default: 3) + MaxRetries int // Maximum number of retry attempts (default: 3) RetryDelays []time.Duration // Delays between retries (default: 0s, 5s, 15s) - Enabled bool // Whether retry is enabled (default: true) + Enabled bool // Whether retry is enabled (default: true) } -// DefaultWatchdogTimeout is the default duration of silence before the watchdog kills a hung process. -const DefaultWatchdogTimeout = 5 * time.Minute - // DefaultRetryConfig returns the default retry configuration. func DefaultRetryConfig() RetryConfig { return RetryConfig{ @@ -39,87 +37,51 @@ func DefaultRetryConfig() RetryConfig { } } -// Loop manages the core agent loop that invokes the configured agent repeatedly until all stories are complete. +// Loop manages the core agent loop that invokes Claude repeatedly until all stories are complete. type Loop struct { - prdPath string - workDir string - prompt string - buildPrompt func() (string, string, error) // optional: rebuild prompt each iteration; returns (prompt, storyID, error) - maxIter int - iteration int - events chan Event - provider Provider - agentCmd *exec.Cmd - logFile *os.File - mu sync.Mutex - stopped bool - paused bool - retryConfig RetryConfig - lastOutputTime time.Time - watchdogTimeout time.Duration - sawStoryDone bool - currentStoryID string + prdPath string + workDir string + prompt string + maxIter int + iteration int + events chan Event + claudeCmd *exec.Cmd + logFile *os.File + mu sync.Mutex + stopped bool + paused bool + retryConfig RetryConfig } // NewLoop creates a new Loop instance. -func NewLoop(prdPath, prompt string, maxIter int, provider Provider) *Loop { +func NewLoop(prdPath, prompt string, maxIter int) *Loop { return &Loop{ - prdPath: prdPath, - prompt: prompt, - maxIter: maxIter, - provider: provider, - events: make(chan Event, 100), - retryConfig: DefaultRetryConfig(), - watchdogTimeout: DefaultWatchdogTimeout, + prdPath: prdPath, + prompt: prompt, + maxIter: maxIter, + events: make(chan Event, 100), + retryConfig: DefaultRetryConfig(), } } // NewLoopWithWorkDir creates a new Loop instance with a configurable working directory. // When workDir is empty, defaults to the project root for backward compatibility. -func NewLoopWithWorkDir(prdPath, workDir string, prompt string, maxIter int, provider Provider) *Loop { +func NewLoopWithWorkDir(prdPath, workDir string, prompt string, maxIter int) *Loop { return &Loop{ - prdPath: prdPath, - workDir: workDir, - prompt: prompt, - maxIter: maxIter, - provider: provider, - events: make(chan Event, 100), - retryConfig: DefaultRetryConfig(), - watchdogTimeout: DefaultWatchdogTimeout, + prdPath: prdPath, + workDir: workDir, + prompt: prompt, + maxIter: maxIter, + events: make(chan Event, 100), + retryConfig: DefaultRetryConfig(), } } // NewLoopWithEmbeddedPrompt creates a new Loop instance using the embedded agent prompt. -// The prompt is rebuilt on each iteration to inline the current story context. -func NewLoopWithEmbeddedPrompt(prdPath string, maxIter int, provider Provider) *Loop { - l := NewLoop(prdPath, "", maxIter, provider) - l.buildPrompt = promptBuilderForPRD(prdPath) - return l -} - -// promptBuilderForPRD returns a function that loads the PRD and builds a prompt -// with the next story inlined. This is called before each iteration so that -// newly completed stories are skipped. The returned storyID is stored on the Loop. -func promptBuilderForPRD(prdPath string) func() (string, string, error) { - return func() (string, string, error) { - p, err := prd.LoadPRD(prdPath) - if err != nil { - return "", "", fmt.Errorf("failed to load PRD for prompt: %w", err) - } - - story := p.NextStory() - if story == nil { - return "", "", fmt.Errorf("all stories are complete") - } - - // Mark the story as in-progress in the markdown file - _ = prd.SetStoryStatus(prdPath, story.ID, "in-progress") - - storyCtx := p.NextStoryContext() - - prompt := embed.GetPrompt(prd.ProgressPath(prdPath), *storyCtx, story.ID, story.Title) - return prompt, story.ID, nil - } +// The PRD path placeholder in the prompt is automatically substituted. +func NewLoopWithEmbeddedPrompt(prdPath string, maxIter int) *Loop { + prompt := embed.GetPrompt(prdPath) + return NewLoop(prdPath, prompt, maxIter) } // Events returns the channel for receiving events from the loop. @@ -136,13 +98,9 @@ func (l *Loop) Iteration() int { // Run executes the agent loop until completion or max iterations. func (l *Loop) Run(ctx context.Context) error { - if l.provider == nil { - return fmt.Errorf("loop provider is not configured") - } - // Open log file in PRD directory prdDir := filepath.Dir(l.prdPath) - logPath := filepath.Join(prdDir, l.provider.LogFileName()) + logPath := filepath.Join(prdDir, "claude.log") var err error l.logFile, err = os.OpenFile(logPath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0644) if err != nil { @@ -174,38 +132,24 @@ func (l *Loop) Run(ctx context.Context) error { return nil } - // Rebuild prompt if builder is set (inlines the current story each iteration) - if l.buildPrompt != nil { - prompt, storyID, err := l.buildPrompt() - if err != nil { - l.events <- Event{ - Type: EventComplete, - Iteration: currentIter, - } - return nil - } - l.mu.Lock() - l.prompt = prompt - l.currentStoryID = storyID - l.sawStoryDone = false - l.mu.Unlock() - } - - // Send iteration start event with current story ID - l.mu.Lock() - iterStoryID := l.currentStoryID - l.mu.Unlock() + // Send iteration start event l.events <- Event{ Type: EventIterationStart, Iteration: currentIter, - StoryID: iterStoryID, } // Run a single iteration with retry logic if err := l.runIterationWithRetry(ctx); err != nil { - l.events <- Event{ - Type: EventError, - Err: err, + if errors.Is(err, ErrQuotaExhausted) { + l.events <- Event{ + Type: EventQuotaExhausted, + Err: err, + } + } else { + l.events <- Event{ + Type: EventError, + Err: err, + } } return err } @@ -217,17 +161,23 @@ func (l *Loop) Run(ctx context.Context) error { default: } - // If the agent emitted <chief-done/>, mark the story as done in prd.md - l.mu.Lock() - saw := l.sawStoryDone - storyID := l.currentStoryID - l.sawStoryDone = false - l.mu.Unlock() - if saw && storyID != "" { - _ = prd.SetStoryStatus(l.prdPath, storyID, "done") + // Check prd.json for completion + p, err := prd.LoadPRD(l.prdPath) + if err != nil { + l.events <- Event{ + Type: EventError, + Err: fmt.Errorf("failed to load PRD: %w", err), + } + return err + } + + if p.AllComplete() { + l.events <- Event{ + Type: EventComplete, + Iteration: currentIter, + } + return nil } - // buildPrompt on the next iteration will return error if all stories are complete, - // which causes EventComplete to be emitted above. // Check pause flag after iteration (loop stops after current iteration completes) l.mu.Lock() @@ -269,7 +219,7 @@ func (l *Loop) runIterationWithRetry(ctx context.Context) error { Iteration: iter, RetryCount: attempt, RetryMax: config.MaxRetries, - Text: fmt.Sprintf("%s crashed, retrying (%d/%d)...", l.provider.Name(), attempt, config.MaxRetries), + Text: fmt.Sprintf("Claude crashed, retrying (%d/%d)...", attempt, config.MaxRetries), } // Wait before retry @@ -309,44 +259,45 @@ func (l *Loop) runIterationWithRetry(ctx context.Context) error { return nil } + // Quota errors should NOT be retried — return immediately + if errors.Is(err, ErrQuotaExhausted) { + return err + } + lastErr = err } return fmt.Errorf("max retries (%d) exceeded: %w", config.MaxRetries, lastErr) } -// runIteration spawns the agent and processes its output. +// runIteration spawns Claude and processes its output. func (l *Loop) runIteration(ctx context.Context) error { - workDir := l.effectiveWorkDir() - cmd := l.provider.LoopCommand(ctx, l.prompt, workDir) + // Build Claude command with required flags l.mu.Lock() - l.agentCmd = cmd - // Initialize watchdog state - l.lastOutputTime = time.Now() - watchdogTimeout := l.watchdogTimeout + l.claudeCmd = exec.CommandContext(ctx, "claude", + "--dangerously-skip-permissions", + "-p", l.prompt, + "--output-format", "stream-json", + "--verbose", + ) + // Set working directory: use workDir if configured, otherwise default to PRD directory + l.claudeCmd.Dir = l.effectiveWorkDir() l.mu.Unlock() // Create pipes for stdout and stderr - stdout, err := l.agentCmd.StdoutPipe() + stdout, err := l.claudeCmd.StdoutPipe() if err != nil { return fmt.Errorf("failed to create stdout pipe: %w", err) } - stderr, err := l.agentCmd.StderrPipe() + stderr, err := l.claudeCmd.StderrPipe() if err != nil { return fmt.Errorf("failed to create stderr pipe: %w", err) } // Start the command - if err := l.agentCmd.Start(); err != nil { - return fmt.Errorf("failed to start %s: %w", l.provider.Name(), err) - } - - // Start watchdog goroutine to detect hung processes - watchdogDone := make(chan struct{}) - var watchdogFired atomic.Bool - if watchdogTimeout > 0 { - go l.runWatchdog(watchdogTimeout, watchdogDone, &watchdogFired) + if err := l.claudeCmd.Start(); err != nil { + return fmt.Errorf("failed to start Claude: %w", err) } // Process stdout in a separate goroutine @@ -358,20 +309,18 @@ func (l *Loop) runIteration(ctx context.Context) error { l.processOutput(stdout) }() - // Log stderr to the log file + // Capture stderr into a buffer while also logging it + var stderrBuf bytes.Buffer go func() { defer wg.Done() - l.logStream(stderr, "[stderr] ") + l.logAndCaptureStream(stderr, "[stderr] ", &stderrBuf) }() // Wait for output processing to complete wg.Wait() - // Stop watchdog - close(watchdogDone) - // Wait for the command to finish - if err := l.agentCmd.Wait(); err != nil { + if err := l.claudeCmd.Wait(); err != nil { // If the context was cancelled, don't treat it as an error if ctx.Err() != nil { return ctx.Err() @@ -383,73 +332,21 @@ func (l *Loop) runIteration(ctx context.Context) error { if stopped { return nil } - // Check if the watchdog killed the process - if watchdogFired.Load() { - return fmt.Errorf("watchdog timeout: no output for %s", watchdogTimeout) + // Check if this is a quota/rate-limit error + stderrText := stderrBuf.String() + if IsQuotaError(stderrText) || IsQuotaError(err.Error()) { + return fmt.Errorf("Claude quota exhausted: %w", ErrQuotaExhausted) } - return fmt.Errorf("%s exited with error: %w", l.provider.Name(), err) + return fmt.Errorf("Claude exited with error: %w", err) } l.mu.Lock() - l.agentCmd = nil + l.claudeCmd = nil l.mu.Unlock() return nil } -// runWatchdog monitors lastOutputTime and kills the process if no output is received -// within the timeout duration. It stops when watchdogDone is closed. -func (l *Loop) runWatchdog(timeout time.Duration, done <-chan struct{}, fired *atomic.Bool) { - // Check interval scales with timeout: 1/5 of timeout, clamped to [10ms, 10s] - checkInterval := timeout / 5 - if checkInterval < 10*time.Millisecond { - checkInterval = 10 * time.Millisecond - } - if checkInterval > 10*time.Second { - checkInterval = 10 * time.Second - } - ticker := time.NewTicker(checkInterval) - defer ticker.Stop() - - for { - select { - case <-ticker.C: - l.mu.Lock() - lastOutput := l.lastOutputTime - stopped := l.stopped - l.mu.Unlock() - - if stopped { - return - } - - if time.Since(lastOutput) > timeout { - fired.Store(true) - - // Emit watchdog timeout event - l.mu.Lock() - iter := l.iteration - l.mu.Unlock() - l.events <- Event{ - Type: EventWatchdogTimeout, - Iteration: iter, - Text: fmt.Sprintf("No output for %s, killing hung process", timeout), - } - - // Kill the process - l.mu.Lock() - if l.agentCmd != nil && l.agentCmd.Process != nil { - l.agentCmd.Process.Kill() - } - l.mu.Unlock() - return - } - case <-done: - return - } - } -} - // processOutput reads stdout line by line, logs it, and parses events. func (l *Loop) processOutput(r io.Reader) { scanner := bufio.NewScanner(r) @@ -460,21 +357,13 @@ func (l *Loop) processOutput(r io.Reader) { for scanner.Scan() { line := scanner.Text() - // Update last output time for watchdog - l.mu.Lock() - l.lastOutputTime = time.Now() - l.mu.Unlock() - // Log raw output l.logLine(line) // Parse the line and emit event if valid - if event := l.provider.ParseLine(line); event != nil { + if event := ParseLine(line); event != nil { l.mu.Lock() event.Iteration = l.iteration - if event.Type == EventStoryDone { - l.sawStoryDone = true - } l.mu.Unlock() l.events <- *event } @@ -489,6 +378,17 @@ func (l *Loop) logStream(r io.Reader, prefix string) { } } +// logAndCaptureStream logs a stream with a prefix and also captures it into a buffer. +func (l *Loop) logAndCaptureStream(r io.Reader, prefix string, buf *bytes.Buffer) { + scanner := bufio.NewScanner(r) + for scanner.Scan() { + text := scanner.Text() + l.logLine(prefix + text) + buf.WriteString(text) + buf.WriteByte('\n') + } +} + // logLine writes a line to the log file. func (l *Loop) logLine(line string) { if l.logFile != nil { @@ -496,15 +396,16 @@ func (l *Loop) logLine(line string) { } } -// Stop terminates the current agent process and stops the loop. +// Stop terminates the current Claude process and stops the loop. func (l *Loop) Stop() { l.mu.Lock() defer l.mu.Unlock() l.stopped = true - if l.agentCmd != nil && l.agentCmd.Process != nil { - l.agentCmd.Process.Kill() + if l.claudeCmd != nil && l.claudeCmd.Process != nil { + // Kill the process + l.claudeCmd.Process.Kill() } } @@ -536,7 +437,7 @@ func (l *Loop) IsStopped() bool { return l.stopped } -// effectiveWorkDir returns the working directory to use for the agent. +// effectiveWorkDir returns the working directory to use for Claude. // If workDir is set, it is used directly. Otherwise, defaults to the PRD directory. func (l *Loop) effectiveWorkDir() string { if l.workDir != "" { @@ -545,11 +446,11 @@ func (l *Loop) effectiveWorkDir() string { return filepath.Dir(l.prdPath) } -// IsRunning returns whether an agent process is currently running. +// IsRunning returns whether a Claude process is currently running. func (l *Loop) IsRunning() bool { l.mu.Lock() defer l.mu.Unlock() - return l.agentCmd != nil && l.agentCmd.Process != nil + return l.claudeCmd != nil && l.claudeCmd.Process != nil } // SetMaxIterations updates the maximum iterations limit. @@ -579,18 +480,3 @@ func (l *Loop) DisableRetry() { defer l.mu.Unlock() l.retryConfig.Enabled = false } - -// SetWatchdogTimeout sets the watchdog timeout duration. -// Setting timeout to 0 disables the watchdog. -func (l *Loop) SetWatchdogTimeout(timeout time.Duration) { - l.mu.Lock() - defer l.mu.Unlock() - l.watchdogTimeout = timeout -} - -// WatchdogTimeout returns the current watchdog timeout duration. -func (l *Loop) WatchdogTimeout() time.Duration { - l.mu.Lock() - defer l.mu.Unlock() - return l.watchdogTimeout -} diff --git a/internal/loop/loop_test.go b/internal/loop/loop_test.go index fe72c0ca..037d802d 100644 --- a/internal/loop/loop_test.go +++ b/internal/loop/loop_test.go @@ -2,48 +2,15 @@ package loop import ( "context" - "fmt" + "encoding/json" "os" - "os/exec" "path/filepath" - "strings" - "sync/atomic" "testing" "time" "github.com/minicodemonkey/chief/internal/prd" ) -// mockProvider implements Provider for tests without importing agent (avoids import cycle). -type mockProvider struct { - cliPath string // if set, used as CLI path; otherwise "claude" -} - -func (m *mockProvider) Name() string { return "Test" } -func (m *mockProvider) CLIPath() string { return m.path() } -func (m *mockProvider) InteractiveCommand(_, _ string) *exec.Cmd { return exec.Command("true") } -func (m *mockProvider) ParseLine(line string) *Event { return ParseLine(line) } -func (m *mockProvider) LogFileName() string { return "claude.log" } - -func (m *mockProvider) path() string { - if m.cliPath != "" { - return m.cliPath - } - return "claude" -} - -func (m *mockProvider) LoopCommand(ctx context.Context, _, workDir string) *exec.Cmd { - p := m.path() - cmd := exec.CommandContext(ctx, p) - cmd.Dir = workDir - return cmd -} - -func (m *mockProvider) CleanOutput(output string) string { return output } - -// testProvider is used by loop tests so they don't need to run a real CLI. -var testProvider Provider = &mockProvider{} - // createMockClaudeScript creates a shell script that outputs predefined stream-json. func createMockClaudeScript(t *testing.T, dir string, output []string) string { t.Helper() @@ -61,21 +28,27 @@ func createMockClaudeScript(t *testing.T, dir string, output []string) string { return scriptPath } -// createTestPRD creates a minimal test PRD markdown file. +// createTestPRD creates a minimal test PRD file. func createTestPRD(t *testing.T, dir string, allComplete bool) string { t.Helper() - status := "" - checkbox := "- [ ] It works" - if allComplete { - status = "**Status:** done\n" - checkbox = "- [x] It works" - } - - md := fmt.Sprintf("# Test Project\n\nTest Description\n\n### US-001: Test Story\n%s%s\n", status, checkbox) - - prdPath := filepath.Join(dir, "prd.md") - if err := os.WriteFile(prdPath, []byte(md), 0644); err != nil { + prdFile := &prd.PRD{ + Project: "Test Project", + Description: "Test Description", + UserStories: []prd.UserStory{ + { + ID: "US-001", + Title: "Test Story", + Description: "A test story", + Priority: 1, + Passes: allComplete, + }, + }, + } + + prdPath := filepath.Join(dir, "prd.json") + data, _ := json.MarshalIndent(prdFile, "", " ") + if err := os.WriteFile(prdPath, data, 0644); err != nil { t.Fatalf("Failed to create test PRD: %v", err) } @@ -83,7 +56,7 @@ func createTestPRD(t *testing.T, dir string, allComplete bool) string { } func TestNewLoop(t *testing.T) { - l := NewLoop("/path/to/prd.json", "test prompt", 5, testProvider) + l := NewLoop("/path/to/prd.json", "test prompt", 5) if l.prdPath != "/path/to/prd.json" { t.Errorf("Expected prdPath %q, got %q", "/path/to/prd.json", l.prdPath) @@ -100,7 +73,7 @@ func TestNewLoop(t *testing.T) { } func TestNewLoopWithWorkDir(t *testing.T) { - l := NewLoopWithWorkDir("/path/to/prd.json", "/work/dir", "test prompt", 5, testProvider) + l := NewLoopWithWorkDir("/path/to/prd.json", "/work/dir", "test prompt", 5) if l.prdPath != "/path/to/prd.json" { t.Errorf("Expected prdPath %q, got %q", "/path/to/prd.json", l.prdPath) @@ -120,7 +93,7 @@ func TestNewLoopWithWorkDir(t *testing.T) { } func TestNewLoopWithWorkDir_EmptyWorkDir(t *testing.T) { - l := NewLoopWithWorkDir("/path/to/prd.json", "", "test prompt", 5, testProvider) + l := NewLoopWithWorkDir("/path/to/prd.json", "", "test prompt", 5) if l.workDir != "" { t.Errorf("Expected empty workDir, got %q", l.workDir) @@ -128,7 +101,7 @@ func TestNewLoopWithWorkDir_EmptyWorkDir(t *testing.T) { } func TestLoop_Events(t *testing.T) { - l := NewLoop("/path/to/prd.json", "test prompt", 5, testProvider) + l := NewLoop("/path/to/prd.json", "test prompt", 5) events := l.Events() if events == nil { @@ -137,7 +110,7 @@ func TestLoop_Events(t *testing.T) { } func TestLoop_Iteration(t *testing.T) { - l := NewLoop("/path/to/prd.json", "test prompt", 5, testProvider) + l := NewLoop("/path/to/prd.json", "test prompt", 5) if l.Iteration() != 0 { t.Errorf("Expected initial iteration to be 0, got %d", l.Iteration()) @@ -150,7 +123,7 @@ func TestLoop_Iteration(t *testing.T) { } func TestLoop_Stop(t *testing.T) { - l := NewLoop("/path/to/prd.json", "test prompt", 5, testProvider) + l := NewLoop("/path/to/prd.json", "test prompt", 5) l.Stop() @@ -186,7 +159,7 @@ func TestLoop_RunWithMockClaude(t *testing.T) { // Create a prompt that invokes our mock script instead of real Claude // For the actual test, we'll test the internal methods - l := NewLoop(prdPath, "test prompt", 1, testProvider) + l := NewLoop(prdPath, "test prompt", 1) // Override the command for testing - we'll test processOutput directly ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) @@ -268,7 +241,7 @@ func TestLoop_MaxIterations(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRD(t, tmpDir, false) // Not complete - l := NewLoop(prdPath, "test prompt", 2, testProvider) + l := NewLoop(prdPath, "test prompt", 2) // Simulate reaching max iterations by manually incrementing l.iteration = 2 @@ -314,7 +287,7 @@ func TestLoop_LogFile(t *testing.T) { t.Fatalf("Failed to create log file: %v", err) } - l := NewLoop(filepath.Join(tmpDir, "prd.md"), "test", 1, testProvider) + l := NewLoop(filepath.Join(tmpDir, "prd.json"), "test", 1) l.logFile = logFile l.logLine("test log line") @@ -331,9 +304,9 @@ func TestLoop_LogFile(t *testing.T) { } } -// TestLoop_ChiefDoneEvent tests detection of <chief-done/> event. -func TestLoop_ChiefDoneEvent(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) +// TestLoop_ChiefCompleteEvent tests detection of <chief-complete/> event. +func TestLoop_ChiefCompleteEvent(t *testing.T) { + l := NewLoop("/test/prd.json", "test", 5) l.iteration = 1 done := make(chan bool) @@ -341,17 +314,17 @@ func TestLoop_ChiefDoneEvent(t *testing.T) { go func() { for event := range l.Events() { events = append(events, event) - if event.Type == EventStoryDone { + if event.Type == EventComplete { break } } done <- true }() - // Simulate processing a line with chief-done + // Simulate processing a line with chief-complete r, w, _ := os.Pipe() go func() { - w.WriteString(`{"type":"assistant","message":{"content":[{"type":"text","text":"All criteria pass! <chief-done/>"}]}}` + "\n") + w.WriteString(`{"type":"assistant","message":{"content":[{"type":"text","text":"All done! <chief-complete/>"}]}}` + "\n") w.Close() }() @@ -359,28 +332,22 @@ func TestLoop_ChiefDoneEvent(t *testing.T) { close(l.events) <-done - // Check that we got a StoryDone event and sawStoryDone was set - hasStoryDone := false + // Check that we got a Complete event + hasComplete := false for _, e := range events { - if e.Type == EventStoryDone { - hasStoryDone = true + if e.Type == EventComplete { + hasComplete = true } } - if !hasStoryDone { - t.Error("Expected StoryDone event for <chief-done/>") + if !hasComplete { + t.Error("Expected Complete event for <chief-complete/>") } - - l.mu.Lock() - if !l.sawStoryDone { - t.Error("Expected sawStoryDone to be true after processing <chief-done/>") - } - l.mu.Unlock() } // TestLoop_SetMaxIterations tests setting max iterations at runtime. func TestLoop_SetMaxIterations(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) + l := NewLoop("/test/prd.json", "test", 5) if l.MaxIterations() != 5 { t.Errorf("Expected initial maxIter 5, got %d", l.MaxIterations()) @@ -410,7 +377,7 @@ func TestDefaultRetryConfig(t *testing.T) { // TestLoop_SetRetryConfig tests setting retry config. func TestLoop_SetRetryConfig(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) + l := NewLoop("/test/prd.json", "test", 5) // Check default if !l.retryConfig.Enabled { @@ -435,234 +402,3 @@ func TestLoop_SetRetryConfig(t *testing.T) { t.Errorf("Expected MaxRetries 5, got %d", l.retryConfig.MaxRetries) } } - -// TestLoop_WatchdogDefaultTimeout tests that the default watchdog timeout is set. -func TestLoop_WatchdogDefaultTimeout(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) - - if l.WatchdogTimeout() != DefaultWatchdogTimeout { - t.Errorf("Expected default watchdog timeout %v, got %v", DefaultWatchdogTimeout, l.WatchdogTimeout()) - } -} - -// TestLoop_SetWatchdogTimeout tests setting the watchdog timeout. -func TestLoop_SetWatchdogTimeout(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) - - l.SetWatchdogTimeout(10 * time.Minute) - if l.WatchdogTimeout() != 10*time.Minute { - t.Errorf("Expected watchdog timeout 10m, got %v", l.WatchdogTimeout()) - } - - // Setting to 0 disables the watchdog - l.SetWatchdogTimeout(0) - if l.WatchdogTimeout() != 0 { - t.Errorf("Expected watchdog timeout 0 (disabled), got %v", l.WatchdogTimeout()) - } -} - -// TestLoop_WatchdogKillsHungProcess tests that a hung process is killed after timeout. -func TestLoop_WatchdogKillsHungProcess(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) - l.iteration = 1 - - // Use a very short timeout for testing - timeout := 100 * time.Millisecond - - // Collect events - var events []Event - done := make(chan bool) - go func() { - for event := range l.Events() { - events = append(events, event) - } - done <- true - }() - - // Create a pipe that never sends data (simulates hung process) - r, w, _ := os.Pipe() - - // Initialize lastOutputTime - l.mu.Lock() - l.lastOutputTime = time.Now() - l.mu.Unlock() - - // Start watchdog with a short check interval - watchdogDone := make(chan struct{}) - var fired atomic.Bool - go l.runWatchdog(timeout, watchdogDone, &fired) - - // processOutput will block until pipe is closed (by watchdog killing would close it, - // but in this test we close it manually after watchdog fires) - go func() { - // Wait for watchdog to fire - time.Sleep(500 * time.Millisecond) - w.Close() - }() - - l.processOutput(r) - close(watchdogDone) - close(l.events) - <-done - - if !fired.Load() { - t.Error("Expected watchdog to fire for hung process") - } - - // Check that we got a WatchdogTimeout event - hasWatchdog := false - for _, e := range events { - if e.Type == EventWatchdogTimeout { - hasWatchdog = true - if e.Text == "" { - t.Error("Expected watchdog event to have descriptive text") - } - } - } - if !hasWatchdog { - t.Error("Expected WatchdogTimeout event") - } -} - -// TestLoop_WatchdogDoesNotFireForActiveProcess tests that an active process doesn't trigger the watchdog. -func TestLoop_WatchdogDoesNotFireForActiveProcess(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) - l.iteration = 1 - - // Use a timeout that's longer than our test - timeout := 2 * time.Second - - // Collect events - var events []Event - done := make(chan bool) - go func() { - for event := range l.Events() { - events = append(events, event) - } - done <- true - }() - - // Create a pipe that produces output regularly - r, w, _ := os.Pipe() - - l.mu.Lock() - l.lastOutputTime = time.Now() - l.mu.Unlock() - - watchdogDone := make(chan struct{}) - var fired atomic.Bool - go l.runWatchdog(timeout, watchdogDone, &fired) - - // Send output regularly, then close - go func() { - for i := 0; i < 5; i++ { - w.WriteString(`{"type":"assistant","message":{"content":[{"type":"text","text":"working..."}]}}` + "\n") - time.Sleep(100 * time.Millisecond) - } - w.Close() - }() - - l.processOutput(r) - close(watchdogDone) - close(l.events) - <-done - - if fired.Load() { - t.Error("Watchdog should NOT fire for an actively producing process") - } - - // Verify no WatchdogTimeout events - for _, e := range events { - if e.Type == EventWatchdogTimeout { - t.Error("Should not have received WatchdogTimeout event for active process") - } - } -} - -// TestLoop_WatchdogDisabledWithZeroTimeout tests that watchdog is disabled when timeout is 0. -func TestLoop_WatchdogDisabledWithZeroTimeout(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) - l.SetWatchdogTimeout(0) - - if l.WatchdogTimeout() != 0 { - t.Errorf("Expected watchdog timeout 0, got %v", l.WatchdogTimeout()) - } - - // Verify that runIteration would not start a watchdog - // (tested indirectly: timeout == 0 means the if-block in runIteration is skipped) - // We test this by verifying the constructor behavior and setter - l2 := NewLoop("/test/prd.json", "test", 5, testProvider) - l2.SetWatchdogTimeout(0) - - l2.mu.Lock() - wt := l2.watchdogTimeout - l2.mu.Unlock() - - if wt != 0 { - t.Errorf("Expected internal watchdogTimeout to be 0, got %v", wt) - } -} - -// TestLoop_LastOutputTimeUpdated tests that lastOutputTime is updated on each scanner output. -func TestLoop_LastOutputTimeUpdated(t *testing.T) { - l := NewLoop("/test/prd.json", "test", 5, testProvider) - l.iteration = 1 - - // Drain events to avoid blocking - go func() { - for range l.Events() { - } - }() - - // Record initial time - l.mu.Lock() - l.lastOutputTime = time.Now().Add(-1 * time.Hour) // Set to an old time - initialTime := l.lastOutputTime - l.mu.Unlock() - - // Send output through processOutput - r, w, _ := os.Pipe() - go func() { - w.WriteString(`{"type":"assistant","message":{"content":[{"type":"text","text":"hello"}]}}` + "\n") - time.Sleep(50 * time.Millisecond) - w.WriteString(`{"type":"assistant","message":{"content":[{"type":"text","text":"world"}]}}` + "\n") - w.Close() - }() - - l.processOutput(r) - close(l.events) - - // Verify lastOutputTime was updated - l.mu.Lock() - finalTime := l.lastOutputTime - l.mu.Unlock() - - if !finalTime.After(initialTime) { - t.Errorf("Expected lastOutputTime to be updated after output, initial=%v, final=%v", initialTime, finalTime) - } -} - -// TestLoop_WatchdogReturnsError tests that watchdog kill causes runIteration to return an error -// that feeds into retry logic. -func TestLoop_WatchdogReturnsError(t *testing.T) { - // This test verifies the error message format that runIterationWithRetry will see - l := NewLoop("/test/prd.json", "test", 5, testProvider) - l.SetWatchdogTimeout(100 * time.Millisecond) - - // The watchdog error message should contain "watchdog timeout" - // This ensures the retry logic in runIterationWithRetry will process it - expectedPrefix := "watchdog timeout:" - errMsg := fmt.Sprintf("watchdog timeout: no output for %s", 100*time.Millisecond) - if !strings.HasPrefix(errMsg, expectedPrefix) { - t.Errorf("Expected error to start with %q, got %q", expectedPrefix, errMsg) - } -} - -// TestLoop_WatchdogWithWorkDir tests that watchdog works with NewLoopWithWorkDir too. -func TestLoop_WatchdogWithWorkDir(t *testing.T) { - l := NewLoopWithWorkDir("/test/prd.json", "/work", "test", 5, testProvider) - - if l.WatchdogTimeout() != DefaultWatchdogTimeout { - t.Errorf("Expected default watchdog timeout for NewLoopWithWorkDir, got %v", l.WatchdogTimeout()) - } -} diff --git a/internal/loop/manager.go b/internal/loop/manager.go index ec5aa203..347fab7a 100644 --- a/internal/loop/manager.go +++ b/internal/loop/manager.go @@ -2,10 +2,12 @@ package loop import ( "context" + "errors" "fmt" "sync" "time" + "github.com/minicodemonkey/chief/embed" "github.com/minicodemonkey/chief/internal/config" "github.com/minicodemonkey/chief/internal/prd" ) @@ -66,13 +68,12 @@ type ManagerEvent struct { // Manager manages multiple Loop instances for parallel PRD execution. type Manager struct { - instances map[string]*LoopInstance - events chan ManagerEvent - maxIter int - retryConfig RetryConfig - provider Provider - baseDir string // Project root directory (for CLAUDE.md etc.) - config *config.Config // Project config for post-completion actions + instances map[string]*LoopInstance + events chan ManagerEvent + maxIter int + retryConfig RetryConfig + baseDir string // Project root directory (for CLAUDE.md etc.) + config *config.Config // Project config for post-completion actions mu sync.RWMutex wg sync.WaitGroup onComplete func(prdName string) // Callback when a PRD completes @@ -80,13 +81,12 @@ type Manager struct { } // NewManager creates a new loop manager. -func NewManager(maxIter int, provider Provider) *Manager { +func NewManager(maxIter int) *Manager { return &Manager{ instances: make(map[string]*LoopInstance), events: make(chan ManagerEvent, 100), maxIter: maxIter, retryConfig: DefaultRetryConfig(), - provider: provider, } } @@ -209,10 +209,6 @@ func (m *Manager) Unregister(name string) error { // Start starts the loop for a specific PRD. func (m *Manager) Start(name string) error { - if m.provider == nil { - return fmt.Errorf("manager provider is not configured") - } - m.mu.Lock() instance, exists := m.instances[name] m.mu.Unlock() @@ -230,14 +226,14 @@ func (m *Manager) Start(name string) error { // Create a new loop instance, using worktree-aware constructor if WorktreeDir is set. // When no worktree is configured, run from the project root (baseDir) so that // CLAUDE.md and other project-level files are visible to Claude. + prompt := embed.GetPrompt(instance.PRDPath) workDir := instance.WorktreeDir if workDir == "" { m.mu.RLock() workDir = m.baseDir m.mu.RUnlock() } - instance.Loop = NewLoopWithWorkDir(instance.PRDPath, workDir, "", m.maxIter, m.provider) - instance.Loop.buildPrompt = promptBuilderForPRD(instance.PRDPath) + instance.Loop = NewLoopWithWorkDir(instance.PRDPath, workDir, prompt, m.maxIter) m.mu.RLock() instance.Loop.SetRetryConfig(m.retryConfig) m.mu.RUnlock() @@ -313,8 +309,14 @@ func (m *Manager) runLoop(instance *LoopInstance) { // Update state based on result instance.mu.Lock() if err != nil && err != context.Canceled { - instance.State = LoopStateError - instance.Error = err + if errors.Is(err, ErrQuotaExhausted) { + // Quota exhaustion pauses the run (resumable by user) + instance.State = LoopStatePaused + instance.Error = err + } else { + instance.State = LoopStateError + instance.Error = err + } } else if instance.Loop.IsPaused() { instance.State = LoopStatePaused } else if instance.Loop.IsStopped() { diff --git a/internal/loop/manager_test.go b/internal/loop/manager_test.go index f2f65fb5..d9b23b9e 100644 --- a/internal/loop/manager_test.go +++ b/internal/loop/manager_test.go @@ -36,7 +36,7 @@ func createTestPRDWithName(t *testing.T, dir, name string) string { } func TestNewManager(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) if m == nil { t.Fatal("expected non-nil manager") } @@ -52,7 +52,7 @@ func TestManagerRegister(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) // Register a new PRD err := m.Register("test-prd", prdPath) @@ -83,7 +83,7 @@ func TestManagerUnregister(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("test-prd", prdPath) // Unregister @@ -109,7 +109,7 @@ func TestManagerGetState(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("test-prd", prdPath) state, iteration, err := m.GetState("test-prd") @@ -136,7 +136,7 @@ func TestManagerGetAllInstances(t *testing.T) { prd2Path := createTestPRDWithName(t, tmpDir, "prd2") prd3Path := createTestPRDWithName(t, tmpDir, "prd3") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("prd1", prd1Path) m.Register("prd2", prd2Path) m.Register("prd3", prd3Path) @@ -159,7 +159,7 @@ func TestManagerGetAllInstances(t *testing.T) { } func TestManagerGetRunningPRDs(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) // Initially no running PRDs running := m.GetRunningPRDs() @@ -169,7 +169,7 @@ func TestManagerGetRunningPRDs(t *testing.T) { } func TestManagerGetRunningCount(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) count := m.GetRunningCount() if count != 0 { @@ -178,7 +178,7 @@ func TestManagerGetRunningCount(t *testing.T) { } func TestManagerIsAnyRunning(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) if m.IsAnyRunning() { t.Error("expected no running loops") @@ -189,7 +189,7 @@ func TestManagerPauseNonRunning(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("test-prd", prdPath) // Pause a non-running PRD should error @@ -203,7 +203,7 @@ func TestManagerStopNonRunning(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("test-prd", prdPath) // Stop a non-running PRD should not error (idempotent) @@ -214,7 +214,7 @@ func TestManagerStopNonRunning(t *testing.T) { } func TestManagerStartNonExistent(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) err := m.Start("non-existent") if err == nil { @@ -222,29 +222,11 @@ func TestManagerStartNonExistent(t *testing.T) { } } -func TestManagerStartRequiresProvider(t *testing.T) { - tmpDir := t.TempDir() - prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - - m := NewManager(10, nil) - if err := m.Register("test-prd", prdPath); err != nil { - t.Fatalf("register failed: %v", err) - } - - err := m.Start("test-prd") - if err == nil { - t.Fatal("expected provider validation error") - } - if err.Error() != "manager provider is not configured" { - t.Fatalf("unexpected error: %v", err) - } -} - func TestManagerConcurrentAccess(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("test-prd", prdPath) // Test concurrent access to manager methods @@ -285,7 +267,7 @@ func TestLoopStateString(t *testing.T) { } func TestManagerSetCompletionCallback(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) called := false var calledWith string @@ -316,7 +298,7 @@ func TestManagerStopAll(t *testing.T) { prd1Path := createTestPRDWithName(t, tmpDir, "prd1") prd2Path := createTestPRDWithName(t, tmpDir, "prd2") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("prd1", prd1Path) m.Register("prd2", prd2Path) @@ -336,7 +318,7 @@ func TestManagerStopAll(t *testing.T) { } func TestManagerSetMaxIterations(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) if m.MaxIterations() != 10 { t.Errorf("expected initial maxIter 10, got %d", m.MaxIterations()) @@ -350,7 +332,7 @@ func TestManagerSetMaxIterations(t *testing.T) { } func TestManagerRetryConfig(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) // Check default retry config if !m.retryConfig.Enabled { @@ -378,7 +360,7 @@ func TestManagerRegisterWithWorktree(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) err := m.RegisterWithWorktree("test-prd", prdPath, "/tmp/worktree/test-prd", "chief/test-prd") if err != nil { @@ -414,7 +396,7 @@ func TestManagerRegisterWithWorktreeFieldsInGetAllInstances(t *testing.T) { prd1Path := createTestPRDWithName(t, tmpDir, "prd1") prd2Path := createTestPRDWithName(t, tmpDir, "prd2") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("prd1", prd1Path) m.RegisterWithWorktree("prd2", prd2Path, "/tmp/wt/prd2", "chief/prd2") @@ -443,7 +425,7 @@ func TestManagerRegisterWithWorktreeFieldsInGetAllInstances(t *testing.T) { } func TestManagerSetConfig(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) // Initially nil if m.Config() != nil { @@ -472,7 +454,7 @@ func TestManagerSetConfig(t *testing.T) { } func TestManagerSetPostCompleteCallback(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) var calledPRD, calledBranch, calledWorkDir string m.SetPostCompleteCallback(func(prdName, branch, workDir string) { @@ -505,7 +487,7 @@ func TestManagerClearWorktreeInfoAll(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.RegisterWithWorktree("test-prd", prdPath, "/tmp/wt/test", "chief/test") // Clear both worktree and branch @@ -526,7 +508,7 @@ func TestManagerClearWorktreeInfoKeepBranch(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.RegisterWithWorktree("test-prd", prdPath, "/tmp/wt/test", "chief/test") // Clear worktree only, keep branch @@ -544,7 +526,7 @@ func TestManagerClearWorktreeInfoKeepBranch(t *testing.T) { } func TestManagerClearWorktreeInfoNotFound(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) err := m.ClearWorktreeInfo("nonexistent", true) if err == nil { t.Error("expected error for nonexistent PRD") @@ -555,7 +537,7 @@ func TestManagerUpdateWorktreeInfo(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.Register("test-prd", prdPath) // Initially no worktree info @@ -579,7 +561,7 @@ func TestManagerUpdateWorktreeInfo(t *testing.T) { } func TestManagerUpdateWorktreeInfoNotFound(t *testing.T) { - m := NewManager(10, testProvider) + m := NewManager(10) err := m.UpdateWorktreeInfo("nonexistent", "/tmp", "branch") if err == nil { t.Error("expected error for nonexistent PRD") @@ -590,7 +572,7 @@ func TestManagerUpdateWorktreeInfoOverwrite(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.RegisterWithWorktree("test-prd", prdPath, "/old/path", "old-branch") // Update with new values @@ -611,7 +593,7 @@ func TestManagerConcurrentAccessWithWorktreeFields(t *testing.T) { tmpDir := t.TempDir() prdPath := createTestPRDWithName(t, tmpDir, "test-prd") - m := NewManager(10, testProvider) + m := NewManager(10) m.RegisterWithWorktree("test-prd", prdPath, "/tmp/wt/test", "chief/test") m.SetConfig(&config.Config{}) diff --git a/internal/loop/opencode_parser.go b/internal/loop/opencode_parser.go deleted file mode 100644 index a9e81ae8..00000000 --- a/internal/loop/opencode_parser.go +++ /dev/null @@ -1,136 +0,0 @@ -package loop - -import ( - "encoding/json" - "errors" - "strings" -) - -type opencodeEvent struct { - Type string `json:"type"` - Timestamp int64 `json:"timestamp"` - SessionID string `json:"sessionID"` - Part *opencodePart `json:"part,omitempty"` - Error *opencodeError `json:"error,omitempty"` -} - -type opencodePart struct { - ID string `json:"id"` - Type string `json:"type,omitempty"` - Text string `json:"text,omitempty"` - Tool string `json:"tool,omitempty"` - CallID string `json:"callID,omitempty"` - Reason string `json:"reason,omitempty"` - Snapshot string `json:"snapshot,omitempty"` - State *opencodeState `json:"state,omitempty"` - Tokens *opencodeTokens `json:"tokens,omitempty"` - Cost float64 `json:"cost,omitempty"` -} - -type opencodeState struct { - Status string `json:"status"` - Input map[string]interface{} `json:"input,omitempty"` - Output string `json:"output,omitempty"` - Title string `json:"title,omitempty"` - Time *opencodeTime `json:"time,omitempty"` -} - -type opencodeTime struct { - Start int64 `json:"start"` - End int64 `json:"end"` -} - -type opencodeTokens struct { - Input int `json:"input"` - Output int `json:"output"` - Reasoning int `json:"reasoning"` - Cache *opencodeCacheTokens `json:"cache,omitempty"` -} - -type opencodeCacheTokens struct { - Read int `json:"read"` - Write int `json:"write"` -} - -type opencodeError struct { - Name string `json:"name"` - Data *opencodeErrorData `json:"data,omitempty"` -} - -type opencodeErrorData struct { - Message string `json:"message"` - StatusCode int `json:"statusCode,omitempty"` -} - -func ParseLineOpenCode(line string) *Event { - line = strings.TrimSpace(line) - if line == "" { - return nil - } - - var ev opencodeEvent - if err := json.Unmarshal([]byte(line), &ev); err != nil { - return nil - } - - switch ev.Type { - case "step_start": - return &Event{Type: EventIterationStart} - - case "tool_use": - if ev.Part == nil { - return nil - } - if ev.Part.State != nil && ev.Part.State.Status == "completed" { - return &Event{ - Type: EventToolResult, - Tool: ev.Part.Tool, - Text: ev.Part.State.Output, - } - } - // Tool starting or in-progress - return &Event{ - Type: EventToolStart, - Tool: ev.Part.Tool, - } - - case "text": - if ev.Part == nil { - return nil - } - if strings.Contains(ev.Part.Text, "<chief-done/>") { - return &Event{ - Type: EventStoryDone, - Text: ev.Part.Text, - } - } - return &Event{ - Type: EventAssistantText, - Text: ev.Part.Text, - } - - case "step_finish": - if ev.Part == nil { - return nil - } - if ev.Part.Reason == "stop" { - return &Event{Type: EventComplete} - } - return nil - - case "error": - msg := "unknown error" - if ev.Error != nil { - if ev.Error.Data != nil { - msg = ev.Error.Data.Message - } - if msg == "" { - msg = ev.Error.Name - } - } - return &Event{Type: EventError, Err: errors.New(msg)} - - default: - return nil - } -} diff --git a/internal/loop/opencode_parser_test.go b/internal/loop/opencode_parser_test.go deleted file mode 100644 index 4d9f97ca..00000000 --- a/internal/loop/opencode_parser_test.go +++ /dev/null @@ -1,121 +0,0 @@ -package loop - -import ( - "testing" -) - -func TestParseLineOpenCode_stepStart(t *testing.T) { - line := `{"type":"step_start","timestamp":1767036059338,"sessionID":"ses_494719016ffe85dkDMj0FPRbHK","part":{"id":"prt_b6b8e7ec7001qAZUB7eTENxPpI","sessionID":"ses_494719016ffe85dkDMj0FPRbHK","messageID":"msg_b6b8e702b0012XuEC4bGe0XhKa","type":"step-start","snapshot":"71db24a798b347669c0ebadb2dfad238f991753d"}}` - ev := ParseLineOpenCode(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventIterationStart { - t.Errorf("expected EventIterationStart, got %v", ev.Type) - } -} - -func TestParseLineOpenCode_toolUseCompleted(t *testing.T) { - line := `{"type":"tool_use","timestamp":1767036061199,"sessionID":"ses_494719016ffe85dkDMj0FPRbHK","part":{"id":"prt_b6b8e85bb001CzBoN2dDlEZJnP","sessionID":"ses_494719016ffe85dkDMj0FPRbHK","messageID":"msg_b6b8e702b0012XuEC4bGe0XhKa","type":"tool","callID":"r9bQWsNLvOrJGIOz","tool":"bash","state":{"status":"completed","input":{"command":"echo hello","description":"Print hello to stdout"},"output":"hello\n","title":"Print hello to stdout","metadata":{"output":"hello\n","exit":0,"description":"Print hello to stdout"},"time":{"start":1767036061123,"end":1767036061173}}}}` - ev := ParseLineOpenCode(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolResult { - t.Errorf("expected EventToolResult, got %v", ev.Type) - } - if ev.Tool != "bash" { - t.Errorf("expected Tool bash, got %q", ev.Tool) - } - if ev.Text != "hello\n" { - t.Errorf("expected Text hello\\n, got %q", ev.Text) - } -} - -func TestParseLineOpenCode_toolUseStarting(t *testing.T) { - line := `{"type":"tool_use","timestamp":1767036061100,"sessionID":"ses_test","part":{"id":"prt_1","type":"tool","tool":"bash","callID":"abc123","state":{"status":"pending","input":{"command":"ls"}}}}` - ev := ParseLineOpenCode(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "bash" { - t.Errorf("expected Tool bash, got %q", ev.Tool) - } -} - -func TestParseLineOpenCode_toolUseNoState(t *testing.T) { - line := `{"type":"tool_use","timestamp":1767036061100,"sessionID":"ses_test","part":{"id":"prt_1","type":"tool","tool":"read"}}` - ev := ParseLineOpenCode(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventToolStart { - t.Errorf("expected EventToolStart, got %v", ev.Type) - } - if ev.Tool != "read" { - t.Errorf("expected Tool read, got %q", ev.Tool) - } -} - -func TestParseLineOpenCode_text(t *testing.T) { - line := `{"type":"text","timestamp":1767036064268,"sessionID":"ses_494719016ffe85dkDMj0FPRbHK","part":{"id":"prt_b6b8e8ff2002mxSx9LtvAlf8Ng","sessionID":"ses_494719016ffe85dkDMj0FPRbHK","messageID":"msg_b6b8e8627001yM4qKJCXdC7W1L","type":"text","text":"hello\n","time":{"start":1767036064265,"end":1767036064265}}}` - ev := ParseLineOpenCode(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventAssistantText { - t.Errorf("expected EventAssistantText, got %v", ev.Type) - } - if ev.Text != "hello\n" { - t.Errorf("expected Text hello\\n, got %q", ev.Text) - } -} - -func TestParseLineOpenCode_stepFinishStop(t *testing.T) { - line := `{"type":"step_finish","timestamp":1767036064273,"sessionID":"ses_494719016ffe85dkDMj0FPRbHK","part":{"id":"prt_b6b8e9209001ojZ4ECN1geZISm","sessionID":"ses_494719016ffe85dkDMj0FPRbHK","messageID":"msg_b6b8e8627001yM4qKJCXdC7W1L","type":"step-finish","reason":"stop","snapshot":"09dd05d11a4ac013136c1df10932efc0ad9116e8","cost":0.001,"tokens":{"input":671,"output":8,"reasoning":0,"cache":{"read":21415,"write":0}}}}` - ev := ParseLineOpenCode(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventComplete { - t.Errorf("expected EventComplete, got %v", ev.Type) - } -} - -func TestParseLineOpenCode_stepFinishToolCalls(t *testing.T) { - line := `{"type":"step_finish","timestamp":1767036061205,"sessionID":"ses_494719016ffe85dkDMj0FPRbHK","part":{"id":"prt_b6b8e85fb001L4I3WHMqH6EQNI","sessionID":"ses_494719016ffe85dkDMj0FPRbHK","messageID":"msg_b6b8e702b0012XuEC4bGe0XhKa","type":"step-finish","reason":"tool-calls","snapshot":"ee3406d50c7d9048674bbb1a3e325d82513b74ed","cost":0,"tokens":{"input":21772,"output":110,"reasoning":0,"cache":{"read":0,"write":0}}}}` - ev := ParseLineOpenCode(line) - if ev != nil { - t.Errorf("expected nil (ignore step_finish with reason=tool-calls), got %v", ev) - } -} - -func TestParseLineOpenCode_error(t *testing.T) { - line := `{"type":"error","timestamp":1767036065000,"sessionID":"ses_494719016ffe85dkDMj0FPRbHK","error":{"name":"APIError","data":{"message":"Rate limit exceeded","statusCode":429,"isRetryable":true}}}` - ev := ParseLineOpenCode(line) - if ev == nil { - t.Fatal("expected event, got nil") - } - if ev.Type != EventError { - t.Errorf("expected EventError, got %v", ev.Type) - } - if ev.Err == nil { - t.Fatal("expected Err set") - } - if ev.Err.Error() != "Rate limit exceeded" { - t.Errorf("unexpected Err: %v", ev.Err) - } -} - -func TestParseLineOpenCode_emptyOrInvalid_returnsNil(t *testing.T) { - tests := []string{"", " ", "not json", "{}", `{"type":"unknown"}`} - for _, line := range tests { - ev := ParseLineOpenCode(line) - if ev != nil { - t.Errorf("ParseLineOpenCode(%q) expected nil, got %v", line, ev) - } - } -} diff --git a/internal/loop/parser.go b/internal/loop/parser.go index 5a38a07f..610a0870 100644 --- a/internal/loop/parser.go +++ b/internal/loop/parser.go @@ -2,6 +2,7 @@ package loop import ( "encoding/json" + "errors" "strings" ) @@ -19,9 +20,11 @@ const ( EventToolStart // EventToolResult is emitted when a tool returns a result. EventToolResult - // EventStoryDone is emitted when Claude signals a story is done via <chief-done/>. - EventStoryDone - // EventComplete is emitted when all stories are complete (buildPrompt returns error). + // EventStoryStarted is emitted when Claude indicates a story is being worked on. + EventStoryStarted + // EventStoryCompleted is emitted when Claude completes a story. + EventStoryCompleted + // EventComplete is emitted when <chief-complete/> is detected. EventComplete // EventMaxIterationsReached is emitted when max iterations are reached. EventMaxIterationsReached @@ -29,8 +32,8 @@ const ( EventError // EventRetrying is emitted when retrying after a crash. EventRetrying - // EventWatchdogTimeout is emitted when the watchdog kills a hung process. - EventWatchdogTimeout + // EventQuotaExhausted is emitted when Claude exits due to quota/rate-limit errors. + EventQuotaExhausted ) // String returns the string representation of an EventType. @@ -44,8 +47,10 @@ func (e EventType) String() string { return "ToolStart" case EventToolResult: return "ToolResult" - case EventStoryDone: - return "StoryDone" + case EventStoryStarted: + return "StoryStarted" + case EventStoryCompleted: + return "StoryCompleted" case EventComplete: return "Complete" case EventMaxIterationsReached: @@ -54,8 +59,8 @@ func (e EventType) String() string { return "Error" case EventRetrying: return "Retrying" - case EventWatchdogTimeout: - return "WatchdogTimeout" + case EventQuotaExhausted: + return "QuotaExhausted" default: return "Unknown" } @@ -134,6 +139,8 @@ func ParseLine(line string) *Event { return parseUserMessage(msg.Message) case "result": + // Result messages indicate the end of an iteration + // We don't emit a specific event for this, but could in the future return nil default: @@ -152,17 +159,28 @@ func parseAssistantMessage(raw json.RawMessage) *Event { return nil } + // Process content blocks - return the first meaningful event + // In practice, we might want to return multiple events, but for simplicity + // we return the first one found for _, block := range msg.Content { switch block.Type { case "text": text := block.Text - // Check for <chief-done/> tag - if strings.Contains(text, "<chief-done/>") { + // Check for <chief-complete/> tag + if strings.Contains(text, "<chief-complete/>") { return &Event{ - Type: EventStoryDone, + Type: EventComplete, Text: text, } } + // Check for story markers using ralph-status tags + if storyID := extractStoryID(text, "<ralph-status>", "</ralph-status>"); storyID != "" { + return &Event{ + Type: EventStoryStarted, + Text: text, + StoryID: storyID, + } + } return &Event{ Type: EventAssistantText, Text: text, @@ -202,3 +220,44 @@ func parseUserMessage(raw json.RawMessage) *Event { return nil } + +// ErrQuotaExhausted is returned when Claude exits due to quota/rate-limit errors. +var ErrQuotaExhausted = errors.New("quota exhausted") + +// quotaPatterns are stderr/error patterns that indicate quota or rate-limit exhaustion. +var quotaPatterns = []string{ + "rate limit", + "rate_limit", + "quota", + "429", + "too many requests", + "resource_exhausted", + "overloaded", +} + +// IsQuotaError checks if an error text contains quota/rate-limit patterns. +func IsQuotaError(errText string) bool { + lower := strings.ToLower(errText) + for _, pattern := range quotaPatterns { + if strings.Contains(lower, pattern) { + return true + } + } + return false +} + +// extractStoryID extracts a story ID from text between start and end tags. +func extractStoryID(text, startTag, endTag string) string { + startIdx := strings.Index(text, startTag) + if startIdx == -1 { + return "" + } + startIdx += len(startTag) + + endIdx := strings.Index(text[startIdx:], endTag) + if endIdx == -1 { + return "" + } + + return strings.TrimSpace(text[startIdx : startIdx+endIdx]) +} diff --git a/internal/loop/parser_test.go b/internal/loop/parser_test.go index 34a98068..19b0926b 100644 --- a/internal/loop/parser_test.go +++ b/internal/loop/parser_test.go @@ -14,12 +14,13 @@ func TestEventTypeString(t *testing.T) { {EventAssistantText, "AssistantText"}, {EventToolStart, "ToolStart"}, {EventToolResult, "ToolResult"}, - {EventStoryDone, "StoryDone"}, + {EventStoryStarted, "StoryStarted"}, + {EventStoryCompleted, "StoryCompleted"}, {EventComplete, "Complete"}, {EventMaxIterationsReached, "MaxIterationsReached"}, {EventError, "Error"}, {EventRetrying, "Retrying"}, - {EventWatchdogTimeout, "WatchdogTimeout"}, + {EventQuotaExhausted, "QuotaExhausted"}, } for _, tt := range tests { @@ -96,11 +97,11 @@ func TestParseLineAssistantText(t *testing.T) { } } -func TestParseLineChiefDone(t *testing.T) { +func TestParseLineChiefComplete(t *testing.T) { tests := []string{ - `{"type":"assistant","message":{"content":[{"type":"text","text":"All criteria pass! <chief-done/>"}]}}`, - `{"type":"assistant","message":{"content":[{"type":"text","text":"<chief-done/>"}]}}`, - `{"type":"assistant","message":{"content":[{"type":"text","text":"Done\n<chief-done/>\nGoodbye"}]}}`, + `{"type":"assistant","message":{"content":[{"type":"text","text":"All stories complete! <chief-complete/>"}]}}`, + `{"type":"assistant","message":{"content":[{"type":"text","text":"<chief-complete/>"}]}}`, + `{"type":"assistant","message":{"content":[{"type":"text","text":"Done\n<chief-complete/>\nGoodbye"}]}}`, } for _, line := range tests { @@ -108,8 +109,8 @@ func TestParseLineChiefDone(t *testing.T) { if event == nil { t.Fatalf("ParseLine(%q) returned nil, want event", line) } - if event.Type != EventStoryDone { - t.Errorf("ParseLine(%q): event.Type = %v, want EventStoryDone", line, event.Type) + if event.Type != EventComplete { + t.Errorf("ParseLine(%q): event.Type = %v, want EventComplete", line, event.Type) } } } @@ -150,6 +151,21 @@ func TestParseLineToolResult(t *testing.T) { } } +func TestParseLineStoryStarted(t *testing.T) { + line := `{"type":"assistant","message":{"content":[{"type":"text","text":"Working on the next story.\n<ralph-status>US-003</ralph-status>\nLet me implement this."}]}}` + + event := ParseLine(line) + if event == nil { + t.Fatal("ParseLine returned nil, want event") + } + if event.Type != EventStoryStarted { + t.Errorf("event.Type = %v, want EventStoryStarted", event.Type) + } + if event.StoryID != "US-003" { + t.Errorf("event.StoryID = %q, want %q", event.StoryID, "US-003") + } +} + func TestParseLineResultMessage(t *testing.T) { line := `{"type":"result","subtype":"success","is_error":false,"result":"Done"}` @@ -179,6 +195,7 @@ func TestParseLineEmptyContent(t *testing.T) { } func TestParseLineRealWorldSystemInit(t *testing.T) { + // Real example from Claude output line := `{"type":"system","subtype":"init","cwd":"/Users/codemonkey/projects/chief","session_id":"7cdf33f6-72ec-4c0e-94fb-e2b637e109da","tools":["Task","Bash","Read"],"model":"claude-opus-4-5-20251101","permissionMode":"default"}` event := ParseLine(line) @@ -191,6 +208,7 @@ func TestParseLineRealWorldSystemInit(t *testing.T) { } func TestParseLineRealWorldToolUse(t *testing.T) { + // Real example from Claude output line := `{"type":"assistant","message":{"model":"claude-opus-4-5-20251101","id":"msg_01XPBzHqPaCuQMUFm77eHe4D","type":"message","role":"assistant","content":[{"type":"tool_use","id":"toolu_01MiR5Ps9inHigemS2gxEz5R","name":"Read","input":{"file_path":"/Users/codemonkey/projects/chief/go.mod"}}]}}` event := ParseLine(line) @@ -205,13 +223,65 @@ func TestParseLineRealWorldToolUse(t *testing.T) { } } +func TestExtractStoryID(t *testing.T) { + tests := []struct { + text string + startTag string + endTag string + expected string + }{ + { + text: "<ralph-status>US-001</ralph-status>", + startTag: "<ralph-status>", + endTag: "</ralph-status>", + expected: "US-001", + }, + { + text: "Some text <ralph-status>US-002</ralph-status> more text", + startTag: "<ralph-status>", + endTag: "</ralph-status>", + expected: "US-002", + }, + { + text: "<ralph-status> US-003 </ralph-status>", + startTag: "<ralph-status>", + endTag: "</ralph-status>", + expected: "US-003", + }, + { + text: "no tags here", + startTag: "<ralph-status>", + endTag: "</ralph-status>", + expected: "", + }, + { + text: "<ralph-status>unclosed", + startTag: "<ralph-status>", + endTag: "</ralph-status>", + expected: "", + }, + } + + for _, tt := range tests { + t.Run(tt.text, func(t *testing.T) { + got := extractStoryID(tt.text, tt.startTag, tt.endTag) + if got != tt.expected { + t.Errorf("extractStoryID(%q, %q, %q) = %q, want %q", tt.text, tt.startTag, tt.endTag, got, tt.expected) + } + }) + } +} + func TestParseLineMultipleContentBlocks(t *testing.T) { + // When there are multiple content blocks, we return the first meaningful one + // This tests that text comes before tool_use in the content array line := `{"type":"assistant","message":{"content":[{"type":"text","text":"First"},{"type":"tool_use","name":"Read","input":{}}]}}` event := ParseLine(line) if event == nil { t.Fatal("ParseLine returned nil, want event") } + // Should return the text event since it comes first if event.Type != EventAssistantText { t.Errorf("event.Type = %v, want EventAssistantText", event.Type) } @@ -221,12 +291,14 @@ func TestParseLineMultipleContentBlocks(t *testing.T) { } func TestParseLineToolUseFirst(t *testing.T) { + // When tool_use comes first line := `{"type":"assistant","message":{"content":[{"type":"tool_use","name":"Write","input":{"file_path":"/test"}},{"type":"text","text":"Second"}]}}` event := ParseLine(line) if event == nil { t.Fatal("ParseLine returned nil, want event") } + // Should return the tool_use event since it comes first if event.Type != EventToolStart { t.Errorf("event.Type = %v, want EventToolStart", event.Type) } @@ -234,3 +306,35 @@ func TestParseLineToolUseFirst(t *testing.T) { t.Errorf("event.Tool = %q, want %q", event.Tool, "Write") } } + +func TestIsQuotaError(t *testing.T) { + tests := []struct { + text string + expected bool + }{ + {"rate limit exceeded", true}, + {"Rate Limit Exceeded", true}, + {"rate_limit_error", true}, + {"quota exceeded for this billing period", true}, + {"HTTP 429 Too Many Requests", true}, + {"429", true}, + {"too many requests", true}, + {"Too Many Requests", true}, + {"resource_exhausted", true}, + {"model is overloaded", true}, + {"Overloaded", true}, + {"normal error message", false}, + {"exit status 1", false}, + {"connection refused", false}, + {"", false}, + {"Claude exited with error: exit status 2", false}, + } + + for _, tt := range tests { + t.Run(tt.text, func(t *testing.T) { + if got := IsQuotaError(tt.text); got != tt.expected { + t.Errorf("IsQuotaError(%q) = %v, want %v", tt.text, got, tt.expected) + } + }) + } +} diff --git a/internal/loop/provider.go b/internal/loop/provider.go deleted file mode 100644 index 40fb4ca7..00000000 --- a/internal/loop/provider.go +++ /dev/null @@ -1,20 +0,0 @@ -package loop - -import ( - "context" - "os/exec" -) - -// Provider is the interface for an agent CLI (e.g. Claude, Codex). -// Implementations live in internal/agent to avoid import cycles. -type Provider interface { - Name() string - CLIPath() string - LoopCommand(ctx context.Context, prompt, workDir string) *exec.Cmd - InteractiveCommand(workDir, prompt string) *exec.Cmd - // CleanOutput extracts JSON from the provider's output format (e.g., NDJSON). - // Returns the original output if no cleaning needed. - CleanOutput(output string) string - ParseLine(line string) *Event - LogFileName() string -} diff --git a/internal/prd/generator.go b/internal/prd/generator.go new file mode 100644 index 00000000..d34c7110 --- /dev/null +++ b/internal/prd/generator.go @@ -0,0 +1,759 @@ +package prd + +import ( + "bufio" + "bytes" + "encoding/json" + "fmt" + "math/rand" + "os" + "os/exec" + "path/filepath" + "strings" + "time" + + "github.com/charmbracelet/lipgloss" + "github.com/charmbracelet/x/term" + "github.com/minicodemonkey/chief/embed" +) + +// Colors duplicated from tui/styles.go to avoid import cycle (tui → git → prd). +var ( + cPrimary = lipgloss.Color("#00D7FF") + cSuccess = lipgloss.Color("#5AF78E") + cMuted = lipgloss.Color("#6C7086") + cBorder = lipgloss.Color("#45475A") + cText = lipgloss.Color("#CDD6F4") +) + +// waitingJokes are shown on a rotating basis during long-running operations. +var waitingJokes = []string{ + "Why do programmers prefer dark mode? Because light attracts bugs.", + "There are only 10 types of people: those who understand binary and those who don't.", + "A SQL query walks into a bar, sees two tables and asks... 'Can I JOIN you?'", + "!false — it's funny because it's true.", + "A programmer's wife says: 'Go to the store and get a gallon of milk. If they have eggs, get a dozen.' He returns with 12 gallons of milk.", + "Why do Java developers wear glasses? Because they can't C#.", + "There's no place like 127.0.0.1.", + "Algorithm: a word used by programmers when they don't want to explain what they did.", + "It works on my machine. Ship it!", + "99 little bugs in the code, 99 little bugs. Take one down, patch it around... 127 little bugs in the code.", + "The best thing about a boolean is that even if you're wrong, you're only off by a bit.", + "Debugging is like being the detective in a crime movie where you are also the murderer.", + "How many programmers does it take to change a light bulb? None, that's a hardware problem.", + "I asked the AI to write a PRD. It wrote a PRD about writing PRDs.", + "You're absolutely right. That's a great point. I completely agree. — Claude, before doing what it was already going to do", + "The AI said it was 95% confident. It was not.", + "Prompt engineering: the art of saying 'no really, do what I said' in 47 different ways.", + "The LLM hallucinated a library that doesn't exist. Honestly, the API looked pretty good though.", + "AI will replace programmers any day now. — programmers, every year since 2022", + "Homer Simpson: 'To start, press any key.' Where's the ANY key?!", + "Homer Simpson: 'Kids, you tried your best and you failed miserably. The lesson is, never try.'", + "The code works and nobody knows why. The code breaks and nobody knows why.", + "Frink: 'You've got to listen to me! Elementary chaos theory tells us that all robots will eventually turn against their masters!'", +} + +// ConvertOptions contains configuration for PRD conversion. +type ConvertOptions struct { + PRDDir string // Directory containing prd.md + Merge bool // Auto-merge progress on conversion conflicts + Force bool // Auto-overwrite on conversion conflicts +} + +// ProgressConflictChoice represents the user's choice when a progress conflict is detected. +type ProgressConflictChoice int + +const ( + ChoiceMerge ProgressConflictChoice = iota // Keep status for matching story IDs + ChoiceOverwrite // Discard all progress + ChoiceCancel // Cancel conversion +) + +// Convert converts prd.md to prd.json using Claude one-shot mode. +// Claude receives the PRD content inline and returns JSON on stdout. +// This function is called: +// - After chief new (new PRD creation) +// - After chief edit (PRD modification) +// - Before chief run if prd.md is newer than prd.json +// +// Progress protection: +// - If prd.json has progress (passes: true or inProgress: true) and prd.md changed: +// - opts.Merge: auto-merge, preserving status for matching story IDs +// - opts.Force: auto-overwrite, discarding all progress +// - Neither: prompt the user with Merge/Overwrite/Cancel options +func Convert(opts ConvertOptions) error { + prdMdPath := filepath.Join(opts.PRDDir, "prd.md") + prdJsonPath := filepath.Join(opts.PRDDir, "prd.json") + + // Check if prd.md exists + if _, err := os.Stat(prdMdPath); os.IsNotExist(err) { + return fmt.Errorf("prd.md not found in %s", opts.PRDDir) + } + + // Resolve absolute path so the prompt can specify exact file locations + absPRDDir, err := filepath.Abs(opts.PRDDir) + if err != nil { + return fmt.Errorf("failed to resolve absolute path: %w", err) + } + + // Check for existing progress before conversion + var existingPRD *PRD + hasProgress := false + if existing, err := LoadPRD(prdJsonPath); err == nil { + existingPRD = existing + hasProgress = HasProgress(existing) + } + + // Run Claude to convert prd.md → JSON string + rawJSON, err := runClaudeConversion(absPRDDir) + if err != nil { + return err + } + + // Clean up output (strip markdown fences if any) + cleanedJSON := cleanJSONOutput(rawJSON) + + // Parse and validate + newPRD, err := parseAndValidatePRD(cleanedJSON) + if err != nil { + // Retry once: ask Claude to fix the invalid JSON + fmt.Println("Conversion produced invalid JSON, retrying...") + fmt.Printf("Raw output:\n---\n%s\n---\n", cleanedJSON) + fixedJSON, retryErr := runClaudeJSONFix(cleanedJSON, err) + if retryErr != nil { + return fmt.Errorf("conversion retry failed: %w", retryErr) + } + + cleanedJSON = cleanJSONOutput(fixedJSON) + newPRD, err = parseAndValidatePRD(cleanedJSON) + if err != nil { + return fmt.Errorf("conversion produced invalid JSON after retry:\n---\n%s\n---\n%w", cleanedJSON, err) + } + } + + // Re-save through Go's JSON encoder to guarantee proper escaping and formatting + normalizedContent, err := json.MarshalIndent(newPRD, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal PRD: %w", err) + } + + // Handle progress protection if existing prd.json has progress + if hasProgress && existingPRD != nil { + choice := ChoiceOverwrite // Default to overwrite if no progress + + if opts.Merge { + choice = ChoiceMerge + } else if opts.Force { + choice = ChoiceOverwrite + } else { + // Prompt user for choice + var promptErr error + choice, promptErr = promptProgressConflict(existingPRD, newPRD) + if promptErr != nil { + return fmt.Errorf("failed to prompt for choice: %w", promptErr) + } + } + + switch choice { + case ChoiceCancel: + return fmt.Errorf("conversion cancelled by user") + case ChoiceMerge: + // Merge progress from existing PRD into new PRD + MergeProgress(existingPRD, newPRD) + // Re-marshal with merged progress + mergedContent, err := json.MarshalIndent(newPRD, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal merged PRD: %w", err) + } + normalizedContent = mergedContent + case ChoiceOverwrite: + // Use the new PRD as-is (no progress) + } + } + + // Write the final normalized prd.json + if err := os.WriteFile(prdJsonPath, append(normalizedContent, '\n'), 0644); err != nil { + return fmt.Errorf("failed to write prd.json: %w", err) + } + + fmt.Println(lipgloss.NewStyle().Foreground(cSuccess).Render("✓ PRD converted successfully")) + return nil +} + +// runClaudeConversion reads prd.md, sends content inline to Claude, and returns the JSON output. +func runClaudeConversion(absPRDDir string) (string, error) { + content, err := os.ReadFile(filepath.Join(absPRDDir, "prd.md")) + if err != nil { + return "", fmt.Errorf("failed to read prd.md: %w", err) + } + + prompt := embed.GetConvertPrompt(string(content)) + + cmd := exec.Command("claude", "-p", "--tools", "") + cmd.Dir = absPRDDir + cmd.Stdin = strings.NewReader(prompt) + + var stdout, stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + if err := cmd.Start(); err != nil { + return "", fmt.Errorf("failed to start Claude: %w", err) + } + + if err := waitWithPanel(cmd, "Converting PRD", "Analyzing PRD...", &stderr); err != nil { + return "", err + } + + return stdout.String(), nil +} + +// runClaudeJSONFix asks Claude to fix invalid JSON inline and returns the corrected output. +func runClaudeJSONFix(badJSON string, validationErr error) (string, error) { + fixPrompt := fmt.Sprintf( + "The following JSON is invalid. The error is: %s\n\n"+ + "Fix the JSON (pay special attention to escaping double quotes inside string values with backslashes) "+ + "and return ONLY the corrected JSON — no markdown fences, no explanation.\n\n%s", + validationErr.Error(), badJSON, + ) + + cmd := exec.Command("claude", "-p", fixPrompt) + + var stdout, stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + if err := cmd.Start(); err != nil { + return "", fmt.Errorf("failed to start Claude: %w", err) + } + + if err := waitWithSpinner(cmd, "Fixing JSON", "Fixing prd.json...", &stderr); err != nil { + return "", err + } + + return stdout.String(), nil +} + +// parseAndValidatePRD unmarshals a JSON string and validates it as a PRD. +func parseAndValidatePRD(jsonStr string) (*PRD, error) { + var prd PRD + if err := json.Unmarshal([]byte(jsonStr), &prd); err != nil { + return nil, fmt.Errorf("failed to parse JSON: %w", err) + } + if prd.Project == "" { + return nil, fmt.Errorf("prd.json missing required 'project' field") + } + if len(prd.UserStories) == 0 { + return nil, fmt.Errorf("prd.json has no user stories") + } + return &prd, nil +} + +// loadAndValidateConvertedPRD loads prd.json from disk and validates it can be parsed as a PRD. +func loadAndValidateConvertedPRD(prdJsonPath string) (*PRD, error) { + data, err := os.ReadFile(prdJsonPath) + if err != nil { + return nil, err + } + return parseAndValidatePRD(string(data)) +} + +// getTerminalWidth returns the current terminal width, defaulting to 80. +func getTerminalWidth() int { + w, _, err := term.GetSize(os.Stdout.Fd()) + if err != nil || w <= 0 { + return 80 + } + return w +} + +// wrapText wraps text to the given width at word boundaries. +func wrapText(text string, width int) string { + words := strings.Fields(text) + if len(words) == 0 { + return "" + } + var lines []string + line := words[0] + for _, w := range words[1:] { + if len(line)+1+len(w) <= width { + line += " " + w + } else { + lines = append(lines, line) + line = w + } + } + lines = append(lines, line) + return strings.Join(lines, "\n") +} + +// renderProgressBar renders a progress bar based on elapsed time vs estimated duration. +// Caps at 95% to avoid showing 100% prematurely. +func renderProgressBar(elapsed time.Duration, width int) string { + const estimatedDuration = 90 * time.Second + + progress := elapsed.Seconds() / estimatedDuration.Seconds() + if progress > 0.95 { + progress = 0.95 + } + if progress < 0 { + progress = 0 + } + + pct := int(progress * 100) + pctStr := fmt.Sprintf("%d%%", pct) + + barWidth := width - len(pctStr) - 2 // 2 for gap between bar and percentage + if barWidth < 10 { + barWidth = 10 + } + + fillWidth := int(float64(barWidth) * progress) + emptyWidth := barWidth - fillWidth + + fill := lipgloss.NewStyle().Foreground(cSuccess).Render(strings.Repeat("█", fillWidth)) + empty := lipgloss.NewStyle().Foreground(cMuted).Render(strings.Repeat("░", emptyWidth)) + styledPct := lipgloss.NewStyle().Foreground(cMuted).Render(pctStr) + + return fill + empty + " " + styledPct +} + +// renderActivityLine renders a line with a cyan dot, activity text, and right-aligned elapsed time. +func renderActivityLine(activity string, elapsed time.Duration, contentWidth int) string { + icon := lipgloss.NewStyle().Foreground(cPrimary).Render("●") + elapsedFmt := formatElapsed(elapsed) + elapsedStr := lipgloss.NewStyle().Foreground(cMuted).Render(elapsedFmt) + + // Truncate activity if it would overflow + maxDescWidth := contentWidth - 2 - len(elapsedFmt) - 2 // icon+space, elapsed, gap + if len(activity) > maxDescWidth && maxDescWidth > 3 { + activity = activity[:maxDescWidth-1] + "…" + } + + descStr := lipgloss.NewStyle().Foreground(cText).Render(activity) + leftPart := icon + " " + descStr + rightPart := elapsedStr + gap := contentWidth - lipgloss.Width(leftPart) - lipgloss.Width(rightPart) + if gap < 1 { + gap = 1 + } + return leftPart + strings.Repeat(" ", gap) + rightPart +} + +// renderProgressBox builds the full lipgloss-styled progress panel with progress bar and joke. +func renderProgressBox(title, activity string, elapsed time.Duration, joke string, panelWidth int) string { + contentWidth := panelWidth - 6 // 2 border + 4 padding (2 each side) + if contentWidth < 20 { + contentWidth = 20 + } + + // Header: "chief <title>" + chiefStr := lipgloss.NewStyle().Bold(true).Foreground(cPrimary).Render("chief") + titleStr := lipgloss.NewStyle().Foreground(cText).Render(title) + header := chiefStr + " " + titleStr + + // Divider + divider := lipgloss.NewStyle().Foreground(cBorder).Render(strings.Repeat("─", contentWidth)) + + // Activity + progress bar + activityLine := renderActivityLine(activity, elapsed, contentWidth) + progressLine := renderProgressBar(elapsed, contentWidth) + + // Joke (word-wrapped, muted) + wrappedJoke := wrapText(joke, contentWidth) + jokeStr := lipgloss.NewStyle().Foreground(cMuted).Render(wrappedJoke) + + content := strings.Join([]string{ + header, + divider, + "", + activityLine, + progressLine, + "", + divider, + jokeStr, + }, "\n") + + style := lipgloss.NewStyle(). + Border(lipgloss.RoundedBorder()). + BorderForeground(cPrimary). + Padding(1, 2). + Width(panelWidth - 2) + + return style.Render(content) +} + +// renderSpinnerBox builds a simpler bordered panel for non-streaming operations. +func renderSpinnerBox(title, activity string, elapsed time.Duration, panelWidth int) string { + contentWidth := panelWidth - 6 + if contentWidth < 20 { + contentWidth = 20 + } + + chiefStr := lipgloss.NewStyle().Bold(true).Foreground(cPrimary).Render("chief") + titleStr := lipgloss.NewStyle().Foreground(cText).Render(title) + header := chiefStr + " " + titleStr + + divider := lipgloss.NewStyle().Foreground(cBorder).Render(strings.Repeat("─", contentWidth)) + activityLine := renderActivityLine(activity, elapsed, contentWidth) + + content := strings.Join([]string{ + header, + divider, + "", + activityLine, + }, "\n") + + style := lipgloss.NewStyle(). + Border(lipgloss.RoundedBorder()). + BorderForeground(cPrimary). + Padding(1, 2). + Width(panelWidth - 2) + + return style.Render(content) +} + +// clearPanelLines clears N lines of previous panel output by moving cursor up and erasing. +func clearPanelLines(n int) { + if n <= 0 { + return + } + // Move to first line + if n > 1 { + fmt.Printf("\033[%dA", n-1) + } + fmt.Print("\r") + // Clear each line + for i := 0; i < n; i++ { + fmt.Print("\033[2K") + if i < n-1 { + fmt.Print("\n") + } + } + // Return to first line + if n > 1 { + fmt.Printf("\033[%dA", n-1) + } + fmt.Print("\r") +} + +// repaintBox repaints the panel box, handling cursor movement for the previous frame. +// Returns the new line count for the next frame. +func repaintBox(box string, prevLines int) int { + newLines := strings.Count(box, "\n") + 1 + + // Move cursor to start of previous panel + if prevLines > 1 { + fmt.Printf("\033[%dA", prevLines-1) + } + if prevLines > 0 { + fmt.Print("\r") + } + + // Print the new box + fmt.Print(box) + + // Clear leftover lines if new box is shorter + if newLines < prevLines { + for i := 0; i < prevLines-newLines; i++ { + fmt.Print("\n\033[2K") + } + fmt.Printf("\033[%dA", prevLines-newLines) + } + + return newLines +} + +// waitWithSpinner runs a bordered panel while waiting for a command to finish. +func waitWithSpinner(cmd *exec.Cmd, title, message string, stderr *bytes.Buffer) error { + done := make(chan error, 1) + go func() { + done <- cmd.Wait() + }() + + startTime := time.Now() + ticker := time.NewTicker(200 * time.Millisecond) + defer ticker.Stop() + + termWidth := getTerminalWidth() + panelWidth := termWidth - 2 + if panelWidth > 62 { + panelWidth = 62 + } + + prevLines := 0 + + for { + select { + case err := <-done: + clearPanelLines(prevLines) + if err != nil { + return fmt.Errorf("Claude failed: %s", stderr.String()) + } + return nil + case <-ticker.C: + box := renderSpinnerBox(title, message, time.Since(startTime), panelWidth) + prevLines = repaintBox(box, prevLines) + } + } +} + +// waitWithPanel runs a full progress panel (header, activity, progress bar, jokes) +// while waiting for a command to finish. Unlike waitWithProgress, it does not parse +// stdout — activity text is static. +func waitWithPanel(cmd *exec.Cmd, title, activity string, stderr *bytes.Buffer) error { + done := make(chan error, 1) + go func() { + done <- cmd.Wait() + }() + + startTime := time.Now() + ticker := time.NewTicker(80 * time.Millisecond) + defer ticker.Stop() + + // Pick a random starting joke and track rotation + jokeIndex := rand.Intn(len(waitingJokes)) + currentJoke := waitingJokes[jokeIndex] + lastJokeChange := time.Now() + + termWidth := getTerminalWidth() + panelWidth := termWidth - 2 + if panelWidth > 62 { + panelWidth = 62 + } + + prevLines := 0 + + for { + select { + case err := <-done: + clearPanelLines(prevLines) + if err != nil { + return fmt.Errorf("Claude failed: %s", stderr.String()) + } + return nil + case <-ticker.C: + // Rotate joke every 30 seconds + if time.Since(lastJokeChange) >= 30*time.Second { + jokeIndex = (jokeIndex + 1 + rand.Intn(len(waitingJokes)-1)) % len(waitingJokes) + currentJoke = waitingJokes[jokeIndex] + lastJokeChange = time.Now() + } + + box := renderProgressBox(title, activity, time.Since(startTime), currentJoke, panelWidth) + prevLines = repaintBox(box, prevLines) + } + } +} + +// formatElapsed formats a duration as a human-readable elapsed time string. +// Examples: "0s", "5s", "1m 12s", "2m 0s" +func formatElapsed(d time.Duration) string { + d = d.Truncate(time.Second) + if d < time.Minute { + return fmt.Sprintf("%ds", int(d.Seconds())) + } + minutes := int(d.Minutes()) + seconds := int(d.Seconds()) % 60 + return fmt.Sprintf("%dm %ds", minutes, seconds) +} + +// NeedsConversion checks if prd.md is newer than prd.json, indicating conversion is needed. +// Returns true if: +// - prd.md exists and prd.json does not exist +// - prd.md exists and is newer than prd.json +// Returns false if: +// - prd.md does not exist +// - prd.json is newer than or same age as prd.md +func NeedsConversion(prdDir string) (bool, error) { + prdMdPath := filepath.Join(prdDir, "prd.md") + prdJsonPath := filepath.Join(prdDir, "prd.json") + + // Check if prd.md exists + mdInfo, err := os.Stat(prdMdPath) + if os.IsNotExist(err) { + // No prd.md, no conversion needed + return false, nil + } + if err != nil { + return false, fmt.Errorf("failed to stat prd.md: %w", err) + } + + // Check if prd.json exists + jsonInfo, err := os.Stat(prdJsonPath) + if os.IsNotExist(err) { + // prd.md exists but prd.json doesn't - needs conversion + return true, nil + } + if err != nil { + return false, fmt.Errorf("failed to stat prd.json: %w", err) + } + + // Both exist - compare modification times + return mdInfo.ModTime().After(jsonInfo.ModTime()), nil +} + +// cleanJSONOutput removes markdown code blocks, conversational preamble, and trims +// whitespace from Claude's output to extract the JSON object. +func cleanJSONOutput(output string) string { + output = strings.TrimSpace(output) + + // Remove markdown code blocks if present + if strings.HasPrefix(output, "```json") { + output = strings.TrimPrefix(output, "```json") + } else if strings.HasPrefix(output, "```") { + output = strings.TrimPrefix(output, "```") + } + + if strings.HasSuffix(output, "```") { + output = strings.TrimSuffix(output, "```") + } + + output = strings.TrimSpace(output) + + // If output doesn't start with '{', Claude may have added preamble text. + // Extract the JSON object by finding the first '{' and matching closing '}'. + if len(output) > 0 && output[0] != '{' { + start := strings.Index(output, "{") + if start == -1 { + return output // No JSON object found, return as-is for error handling + } + // Find the matching closing brace by counting brace depth + depth := 0 + inString := false + escaped := false + end := -1 + for i := start; i < len(output); i++ { + if escaped { + escaped = false + continue + } + ch := output[i] + if ch == '\\' && inString { + escaped = true + continue + } + if ch == '"' { + inString = !inString + continue + } + if inString { + continue + } + if ch == '{' { + depth++ + } else if ch == '}' { + depth-- + if depth == 0 { + end = i + break + } + } + } + if end != -1 { + output = output[start : end+1] + } else { + // No matching closing brace; take from first '{' to end + output = output[start:] + } + } + + return strings.TrimSpace(output) +} + +// validateJSON checks if the given string is valid JSON. +func validateJSON(content string) error { + var js json.RawMessage + if err := json.Unmarshal([]byte(content), &js); err != nil { + return fmt.Errorf("invalid JSON: %w", err) + } + return nil +} + +// HasProgress checks if the PRD has any progress (passes: true or inProgress: true). +func HasProgress(prd *PRD) bool { + if prd == nil { + return false + } + for _, story := range prd.UserStories { + if story.Passes || story.InProgress { + return true + } + } + return false +} + +// MergeProgress merges progress from the old PRD into the new PRD. +// For stories with matching IDs, it preserves the Passes and InProgress status. +// New stories (in newPRD but not in oldPRD) are added without progress. +// Removed stories (in oldPRD but not in newPRD) are dropped. +func MergeProgress(oldPRD, newPRD *PRD) { + if oldPRD == nil || newPRD == nil { + return + } + + // Create a map of old story statuses by ID + oldStatus := make(map[string]struct { + passes bool + inProgress bool + }) + for _, story := range oldPRD.UserStories { + oldStatus[story.ID] = struct { + passes bool + inProgress bool + }{ + passes: story.Passes, + inProgress: story.InProgress, + } + } + + // Apply old status to matching stories in new PRD + for i := range newPRD.UserStories { + if status, exists := oldStatus[newPRD.UserStories[i].ID]; exists { + newPRD.UserStories[i].Passes = status.passes + newPRD.UserStories[i].InProgress = status.inProgress + } + } +} + +// promptProgressConflict prompts the user to choose how to handle a progress conflict. +func promptProgressConflict(oldPRD, newPRD *PRD) (ProgressConflictChoice, error) { + // Count stories with progress + progressCount := 0 + for _, story := range oldPRD.UserStories { + if story.Passes || story.InProgress { + progressCount++ + } + } + + // Show warning + fmt.Println() + fmt.Printf("⚠️ Warning: prd.json has progress (%d stories with status)\n", progressCount) + fmt.Println() + fmt.Println("How would you like to proceed?") + fmt.Println() + fmt.Println(" [m] Merge - Keep status for matching story IDs, add new stories, drop removed stories") + fmt.Println(" [o] Overwrite - Discard all progress and use the new PRD") + fmt.Println(" [c] Cancel - Cancel conversion and keep existing prd.json") + fmt.Println() + fmt.Print("Choice [m/o/c]: ") + + reader := bufio.NewReader(os.Stdin) + input, err := reader.ReadString('\n') + if err != nil { + return ChoiceCancel, fmt.Errorf("failed to read input: %w", err) + } + + input = strings.TrimSpace(strings.ToLower(input)) + switch input { + case "m", "merge": + return ChoiceMerge, nil + case "o", "overwrite": + return ChoiceOverwrite, nil + case "c", "cancel", "": + return ChoiceCancel, nil + default: + fmt.Printf("Invalid choice %q, cancelling conversion.\n", input) + return ChoiceCancel, nil + } +} diff --git a/internal/prd/generator_test.go b/internal/prd/generator_test.go new file mode 100644 index 00000000..67859884 --- /dev/null +++ b/internal/prd/generator_test.go @@ -0,0 +1,634 @@ +package prd + +import ( + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func TestCleanJSONOutput(t *testing.T) { + tests := []struct { + name string + input string + expected string + }{ + { + name: "plain JSON", + input: `{"project": "test"}`, + expected: `{"project": "test"}`, + }, + { + name: "with json code block", + input: "```json\n{\"project\": \"test\"}\n```", + expected: `{"project": "test"}`, + }, + { + name: "with plain code block", + input: "```\n{\"project\": \"test\"}\n```", + expected: `{"project": "test"}`, + }, + { + name: "with extra whitespace", + input: " \n{\"project\": \"test\"}\n ", + expected: `{"project": "test"}`, + }, + { + name: "with conversational preamble", + input: "Since the file write is being denied, here's the JSON output directly:\n\n{\"project\": \"test\"}", + expected: `{"project": "test"}`, + }, + { + name: "with preamble and nested objects", + input: "Here is the JSON:\n{\"project\": \"test\", \"userStories\": [{\"id\": \"US-001\"}]}", + expected: `{"project": "test", "userStories": [{"id": "US-001"}]}`, + }, + { + name: "with preamble and trailing text", + input: "Here you go:\n{\"project\": \"test\"}\nLet me know if you need changes.", + expected: `{"project": "test"}`, + }, + { + name: "with code fence and preamble", + input: "Here is the output:\n```json\n{\"project\": \"test\"}\n```", + expected: `{"project": "test"}`, + }, + { + name: "JSON with escaped quotes in preamble scenario", + input: "Output:\n{\"project\": \"test \\\"quoted\\\"\"}", + expected: `{"project": "test \"quoted\""}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := cleanJSONOutput(tt.input) + if result != tt.expected { + t.Errorf("cleanJSONOutput() = %q, want %q", result, tt.expected) + } + }) + } +} + +func TestValidateJSON(t *testing.T) { + tests := []struct { + name string + input string + wantErr bool + }{ + { + name: "valid JSON object", + input: `{"project": "test", "stories": []}`, + wantErr: false, + }, + { + name: "valid JSON array", + input: `[1, 2, 3]`, + wantErr: false, + }, + { + name: "valid nested JSON", + input: `{"project": "test", "userStories": [{"id": "US-001", "title": "Test"}]}`, + wantErr: false, + }, + { + name: "invalid JSON - missing closing brace", + input: `{"project": "test"`, + wantErr: true, + }, + { + name: "invalid JSON - trailing comma", + input: `{"project": "test",}`, + wantErr: true, + }, + { + name: "invalid JSON - plain text", + input: `This is not JSON`, + wantErr: true, + }, + { + name: "empty string", + input: ``, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := validateJSON(tt.input) + if (err != nil) != tt.wantErr { + t.Errorf("validateJSON() error = %v, wantErr %v", err, tt.wantErr) + } + }) + } +} + +func TestNeedsConversion(t *testing.T) { + t.Run("no prd.md exists", func(t *testing.T) { + tmpDir := t.TempDir() + + needs, err := NeedsConversion(tmpDir) + if err != nil { + t.Errorf("NeedsConversion() unexpected error: %v", err) + } + if needs { + t.Error("NeedsConversion() = true, want false when no prd.md exists") + } + }) + + t.Run("prd.md exists but prd.json does not", func(t *testing.T) { + tmpDir := t.TempDir() + prdMdPath := filepath.Join(tmpDir, "prd.md") + if err := os.WriteFile(prdMdPath, []byte("# Test PRD"), 0644); err != nil { + t.Fatalf("Failed to create prd.md: %v", err) + } + + needs, err := NeedsConversion(tmpDir) + if err != nil { + t.Errorf("NeedsConversion() unexpected error: %v", err) + } + if !needs { + t.Error("NeedsConversion() = false, want true when prd.json doesn't exist") + } + }) + + t.Run("prd.md is newer than prd.json", func(t *testing.T) { + tmpDir := t.TempDir() + + // Create prd.json first + prdJsonPath := filepath.Join(tmpDir, "prd.json") + if err := os.WriteFile(prdJsonPath, []byte(`{"project":"test"}`), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) + } + + // Wait a moment to ensure different timestamps + time.Sleep(100 * time.Millisecond) + + // Create prd.md after (so it's newer) + prdMdPath := filepath.Join(tmpDir, "prd.md") + if err := os.WriteFile(prdMdPath, []byte("# Test PRD"), 0644); err != nil { + t.Fatalf("Failed to create prd.md: %v", err) + } + + needs, err := NeedsConversion(tmpDir) + if err != nil { + t.Errorf("NeedsConversion() unexpected error: %v", err) + } + if !needs { + t.Error("NeedsConversion() = false, want true when prd.md is newer") + } + }) + + t.Run("prd.json is newer than prd.md", func(t *testing.T) { + tmpDir := t.TempDir() + + // Create prd.md first + prdMdPath := filepath.Join(tmpDir, "prd.md") + if err := os.WriteFile(prdMdPath, []byte("# Test PRD"), 0644); err != nil { + t.Fatalf("Failed to create prd.md: %v", err) + } + + // Wait a moment to ensure different timestamps + time.Sleep(100 * time.Millisecond) + + // Create prd.json after (so it's newer) + prdJsonPath := filepath.Join(tmpDir, "prd.json") + if err := os.WriteFile(prdJsonPath, []byte(`{"project":"test"}`), 0644); err != nil { + t.Fatalf("Failed to create prd.json: %v", err) + } + + needs, err := NeedsConversion(tmpDir) + if err != nil { + t.Errorf("NeedsConversion() unexpected error: %v", err) + } + if needs { + t.Error("NeedsConversion() = true, want false when prd.json is newer") + } + }) +} + +func TestConvertMissingPrdMd(t *testing.T) { + tmpDir := t.TempDir() + + err := Convert(ConvertOptions{PRDDir: tmpDir}) + if err == nil { + t.Error("Convert() expected error when prd.md is missing") + } +} + +func TestHasProgress(t *testing.T) { + tests := []struct { + name string + prd *PRD + expected bool + }{ + { + name: "nil PRD", + prd: nil, + expected: false, + }, + { + name: "empty PRD", + prd: &PRD{}, + expected: false, + }, + { + name: "no progress", + prd: &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: false, InProgress: false}, + {ID: "US-002", Passes: false, InProgress: false}, + }, + }, + expected: false, + }, + { + name: "one story passes", + prd: &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: true, InProgress: false}, + {ID: "US-002", Passes: false, InProgress: false}, + }, + }, + expected: true, + }, + { + name: "one story in progress", + prd: &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: false, InProgress: true}, + {ID: "US-002", Passes: false, InProgress: false}, + }, + }, + expected: true, + }, + { + name: "all stories pass", + prd: &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: true}, + {ID: "US-002", Passes: true}, + }, + }, + expected: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := HasProgress(tt.prd) + if result != tt.expected { + t.Errorf("HasProgress() = %v, want %v", result, tt.expected) + } + }) + } +} + +func TestMergeProgress(t *testing.T) { + t.Run("nil PRDs", func(t *testing.T) { + // Should not panic + MergeProgress(nil, nil) + MergeProgress(&PRD{}, nil) + MergeProgress(nil, &PRD{}) + }) + + t.Run("matching story IDs - preserve status", func(t *testing.T) { + oldPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Title: "Old Title 1", Passes: true, InProgress: false}, + {ID: "US-002", Title: "Old Title 2", Passes: false, InProgress: true}, + {ID: "US-003", Title: "Old Title 3", Passes: false, InProgress: false}, + }, + } + newPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Title: "New Title 1", Passes: false, InProgress: false}, + {ID: "US-002", Title: "New Title 2", Passes: false, InProgress: false}, + {ID: "US-003", Title: "New Title 3", Passes: false, InProgress: false}, + }, + } + + MergeProgress(oldPRD, newPRD) + + // US-001 should have passes: true preserved + if !newPRD.UserStories[0].Passes { + t.Error("US-001 should have Passes: true after merge") + } + // US-002 should have inProgress: true preserved + if !newPRD.UserStories[1].InProgress { + t.Error("US-002 should have InProgress: true after merge") + } + // US-003 should remain unchanged (no progress) + if newPRD.UserStories[2].Passes || newPRD.UserStories[2].InProgress { + t.Error("US-003 should not have any progress after merge") + } + }) + + t.Run("new stories added - no progress", func(t *testing.T) { + oldPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: true}, + }, + } + newPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: false}, + {ID: "US-002", Passes: false}, // New story + }, + } + + MergeProgress(oldPRD, newPRD) + + // US-001 should have progress preserved + if !newPRD.UserStories[0].Passes { + t.Error("US-001 should have Passes: true after merge") + } + // US-002 is new, should have no progress + if newPRD.UserStories[1].Passes || newPRD.UserStories[1].InProgress { + t.Error("New story US-002 should not have any progress") + } + }) + + t.Run("removed stories are dropped", func(t *testing.T) { + oldPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: true}, + {ID: "US-002", Passes: true}, // Will be removed + }, + } + newPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: false}, + // US-002 removed from new PRD + }, + } + + MergeProgress(oldPRD, newPRD) + + // Only US-001 should exist + if len(newPRD.UserStories) != 1 { + t.Errorf("Expected 1 story, got %d", len(newPRD.UserStories)) + } + if newPRD.UserStories[0].ID != "US-001" { + t.Errorf("Expected US-001, got %s", newPRD.UserStories[0].ID) + } + if !newPRD.UserStories[0].Passes { + t.Error("US-001 should have Passes: true after merge") + } + }) + + t.Run("mixed scenario - add, remove, keep", func(t *testing.T) { + oldPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: true}, // Keep with progress + {ID: "US-002", Passes: true}, // Removed + {ID: "US-003", InProgress: true}, // Keep with progress + {ID: "US-004", Passes: false}, // Keep without progress + }, + } + newPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Passes: false}, // Existing + {ID: "US-003", Passes: false}, // Existing + {ID: "US-004", Passes: false}, // Existing + {ID: "US-005", Passes: false}, // New + }, + } + + MergeProgress(oldPRD, newPRD) + + // Verify each story + storyMap := make(map[string]*UserStory) + for i := range newPRD.UserStories { + storyMap[newPRD.UserStories[i].ID] = &newPRD.UserStories[i] + } + + if s, ok := storyMap["US-001"]; !ok || !s.Passes { + t.Error("US-001 should exist with Passes: true") + } + if _, ok := storyMap["US-002"]; ok { + t.Error("US-002 should be removed") + } + if s, ok := storyMap["US-003"]; !ok || !s.InProgress { + t.Error("US-003 should exist with InProgress: true") + } + if s, ok := storyMap["US-004"]; !ok || s.Passes || s.InProgress { + t.Error("US-004 should exist without progress") + } + if s, ok := storyMap["US-005"]; !ok || s.Passes || s.InProgress { + t.Error("US-005 should exist without progress (new story)") + } + }) + + t.Run("reordered stories - preserves progress by ID", func(t *testing.T) { + oldPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-001", Priority: 1, Passes: true}, + {ID: "US-002", Priority: 2, Passes: false}, + {ID: "US-003", Priority: 3, InProgress: true}, + }, + } + newPRD := &PRD{ + UserStories: []UserStory{ + {ID: "US-003", Priority: 1, Passes: false}, // Moved to top + {ID: "US-001", Priority: 2, Passes: false}, // Moved down + {ID: "US-002", Priority: 3, Passes: false}, // Moved down + }, + } + + MergeProgress(oldPRD, newPRD) + + // Verify progress is preserved regardless of order + if !newPRD.UserStories[0].InProgress { + t.Error("US-003 should have InProgress: true after merge") + } + if !newPRD.UserStories[1].Passes { + t.Error("US-001 should have Passes: true after merge") + } + if newPRD.UserStories[2].Passes || newPRD.UserStories[2].InProgress { + t.Error("US-002 should not have progress after merge") + } + }) +} + +func TestLoadAndValidateConvertedPRD(t *testing.T) { + t.Run("valid prd.json", func(t *testing.T) { + tmpDir := t.TempDir() + prdJsonPath := filepath.Join(tmpDir, "prd.json") + content := `{ + "project": "Test Project", + "description": "A test project", + "userStories": [ + { + "id": "US-001", + "title": "First Story", + "description": "Do something", + "acceptanceCriteria": ["It works"], + "priority": 1, + "passes": false + } + ] +}` + if err := os.WriteFile(prdJsonPath, []byte(content), 0644); err != nil { + t.Fatalf("Failed to write test prd.json: %v", err) + } + + prd, err := loadAndValidateConvertedPRD(prdJsonPath) + if err != nil { + t.Errorf("loadAndValidateConvertedPRD() unexpected error: %v", err) + } + if prd == nil { + t.Fatal("Expected non-nil PRD") + } + if prd.Project != "Test Project" { + t.Errorf("Expected project 'Test Project', got %q", prd.Project) + } + }) + + t.Run("file does not exist", func(t *testing.T) { + tmpDir := t.TempDir() + prdJsonPath := filepath.Join(tmpDir, "prd.json") + + _, err := loadAndValidateConvertedPRD(prdJsonPath) + if err == nil { + t.Error("Expected error when prd.json does not exist") + } + }) + + t.Run("invalid JSON", func(t *testing.T) { + tmpDir := t.TempDir() + prdJsonPath := filepath.Join(tmpDir, "prd.json") + if err := os.WriteFile(prdJsonPath, []byte(`{invalid json`), 0644); err != nil { + t.Fatalf("Failed to write test file: %v", err) + } + + _, err := loadAndValidateConvertedPRD(prdJsonPath) + if err == nil { + t.Error("Expected error for invalid JSON") + } + }) + + t.Run("missing project field", func(t *testing.T) { + tmpDir := t.TempDir() + prdJsonPath := filepath.Join(tmpDir, "prd.json") + content := `{ + "project": "", + "description": "A test", + "userStories": [{"id": "US-001", "title": "Story", "description": "Desc", "acceptanceCriteria": [], "priority": 1, "passes": false}] +}` + if err := os.WriteFile(prdJsonPath, []byte(content), 0644); err != nil { + t.Fatalf("Failed to write test file: %v", err) + } + + _, err := loadAndValidateConvertedPRD(prdJsonPath) + if err == nil { + t.Error("Expected error for missing project field") + } + if err != nil && !strings.Contains(err.Error(), "project") { + t.Errorf("Expected error about 'project' field, got: %v", err) + } + }) + + t.Run("no user stories", func(t *testing.T) { + tmpDir := t.TempDir() + prdJsonPath := filepath.Join(tmpDir, "prd.json") + content := `{ + "project": "Test", + "description": "A test", + "userStories": [] +}` + if err := os.WriteFile(prdJsonPath, []byte(content), 0644); err != nil { + t.Fatalf("Failed to write test file: %v", err) + } + + _, err := loadAndValidateConvertedPRD(prdJsonPath) + if err == nil { + t.Error("Expected error for empty user stories") + } + if err != nil && !strings.Contains(err.Error(), "user stories") { + t.Errorf("Expected error about 'user stories', got: %v", err) + } + }) + + t.Run("JSON with escaped quotes parses correctly", func(t *testing.T) { + tmpDir := t.TempDir() + prdJsonPath := filepath.Join(tmpDir, "prd.json") + content := `{ + "project": "Test Project", + "description": "A project with \"quoted\" text", + "userStories": [ + { + "id": "US-001", + "title": "Story with \"quotes\"", + "description": "Click the \"Submit\" button", + "acceptanceCriteria": ["User sees \"Success\" message", "Button says \"OK\""], + "priority": 1, + "passes": false + } + ] +}` + if err := os.WriteFile(prdJsonPath, []byte(content), 0644); err != nil { + t.Fatalf("Failed to write test file: %v", err) + } + + prd, err := loadAndValidateConvertedPRD(prdJsonPath) + if err != nil { + t.Errorf("loadAndValidateConvertedPRD() unexpected error: %v", err) + } + if prd == nil { + t.Fatal("Expected non-nil PRD") + } + // Verify the escaped quotes are properly parsed + if prd.UserStories[0].Title != `Story with "quotes"` { + t.Errorf("Expected title with unescaped quotes, got %q", prd.UserStories[0].Title) + } + if prd.UserStories[0].AcceptanceCriteria[0] != `User sees "Success" message` { + t.Errorf("Expected acceptance criteria with unescaped quotes, got %q", prd.UserStories[0].AcceptanceCriteria[0]) + } + }) +} + +// Note: Full integration tests for Convert(), runClaudeConversion(), runClaudeJSONFix(), +// and waitWithSpinner() require Claude to be available and are not included here. + +func TestSamplePRDMarkdown(t *testing.T) { + // Test that a sample prd.md structure is recognized + // This verifies the file detection logic, not the actual conversion + tmpDir := t.TempDir() + + sampleMd := `# My Test Project + +A sample project for testing. + +## User Stories + +### US-001: Setup Project +As a developer, I need a properly structured project. + +**Acceptance Criteria:** +- Create project structure +- Add dependencies +- Verify build works + +### US-002: Add Feature +As a user, I want a new feature. + +**Acceptance Criteria:** +- Feature works correctly +- Tests pass +` + prdMdPath := filepath.Join(tmpDir, "prd.md") + if err := os.WriteFile(prdMdPath, []byte(sampleMd), 0644); err != nil { + t.Fatalf("Failed to create sample prd.md: %v", err) + } + + // Verify the file can be detected for conversion + needs, err := NeedsConversion(tmpDir) + if err != nil { + t.Errorf("NeedsConversion() unexpected error: %v", err) + } + if !needs { + t.Error("Sample prd.md should trigger conversion need") + } +} diff --git a/internal/prd/loader.go b/internal/prd/loader.go index 88182072..a4da8330 100644 --- a/internal/prd/loader.go +++ b/internal/prd/loader.go @@ -1,6 +1,36 @@ package prd -// LoadPRD reads and parses a PRD markdown file from the given path. +import ( + "encoding/json" + "fmt" + "os" +) + +// LoadPRD reads and parses a PRD JSON file from the given path. func LoadPRD(path string) (*PRD, error) { - return ParseMarkdownPRD(path) + data, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("failed to read PRD file: %w", err) + } + + var p PRD + if err := json.Unmarshal(data, &p); err != nil { + return nil, fmt.Errorf("failed to parse PRD JSON: %w", err) + } + + return &p, nil +} + +// Save writes the PRD back to a JSON file at the given path. +func (p *PRD) Save(path string) error { + data, err := json.MarshalIndent(p, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal PRD: %w", err) + } + + if err := os.WriteFile(path, data, 0644); err != nil { + return fmt.Errorf("failed to write PRD file: %w", err) + } + + return nil } diff --git a/internal/prd/markdown.go b/internal/prd/markdown.go deleted file mode 100644 index e45925bf..00000000 --- a/internal/prd/markdown.go +++ /dev/null @@ -1,173 +0,0 @@ -package prd - -import ( - "fmt" - "os" - "regexp" - "strings" -) - -// storyHeadingRegex matches story headings like "### US-001: Story Title" or "#### US-001: Story Title" -var storyHeadingRegex = regexp.MustCompile(`^#{3,4}\s+([A-Za-z]+-\d+):\s+(.+)$`) - -// statusLineRegex matches "**Status:** value" -var statusLineRegex = regexp.MustCompile(`^\*\*Status:\*\*\s*(.+)$`) - -// priorityLineRegex matches "**Priority:** value" -var priorityLineRegex = regexp.MustCompile(`^\*\*Priority:\*\*\s*(.+)$`) - -// descriptionLineRegex matches "**Description:** value" -var descriptionLineRegex = regexp.MustCompile(`^\*\*Description:\*\*\s*(.+)$`) - -// checkboxRegex matches "- [ ] text" or "- [x] text" -var checkboxRegex = regexp.MustCompile(`^-\s+\[([ xX])\]\s+(.+)$`) - -// projectHeadingRegex matches "# PRD: Name" or "# Name" -var projectHeadingRegex = regexp.MustCompile(`^#\s+(?:PRD:\s+)?(.+)$`) - -// ParseMarkdownPRD reads and parses a PRD markdown file from the given path. -func ParseMarkdownPRD(path string) (*PRD, error) { - data, err := os.ReadFile(path) - if err != nil { - return nil, fmt.Errorf("failed to read PRD file: %w", err) - } - - return ParseMarkdownPRDFromString(string(data)) -} - -// ParseMarkdownPRDFromString parses a PRD from a markdown string. -func ParseMarkdownPRDFromString(content string) (*PRD, error) { - lines := strings.Split(content, "\n") - p := &PRD{} - - type storyBuilder struct { - story UserStory - descLines []string - } - - var current *storyBuilder - introStarted := false - introDone := false - autoPriority := float64(0) - - flushStory := func() { - if current == nil { - return - } - // If no explicit Description, join collected prose lines - if current.story.Description == "" && len(current.descLines) > 0 { - current.story.Description = strings.Join(current.descLines, " ") - } - // Assign auto-priority if none was set - if current.story.Priority == 0 { - autoPriority++ - current.story.Priority = autoPriority - } else if current.story.Priority > autoPriority { - autoPriority = current.story.Priority - } - p.UserStories = append(p.UserStories, current.story) - current = nil - } - - for _, line := range lines { - trimmed := strings.TrimSpace(line) - - // Check for project heading (# level only, not ## or ###) - if strings.HasPrefix(line, "# ") && !strings.HasPrefix(line, "## ") { - if m := projectHeadingRegex.FindStringSubmatch(trimmed); m != nil { - p.Project = strings.TrimSpace(m[1]) - introStarted = true - continue - } - } - - // Check for story heading (### ID: Title) - if m := storyHeadingRegex.FindStringSubmatch(trimmed); m != nil { - flushStory() - introDone = true - current = &storyBuilder{ - story: UserStory{ - ID: m[1], - Title: strings.TrimSpace(m[2]), - }, - } - continue - } - - // Check for ## or ### heading (section boundary — ends current story block) - if strings.HasPrefix(line, "## ") || strings.HasPrefix(line, "### ") { - flushStory() - - heading := strings.TrimSpace(strings.TrimLeft(trimmed, "#")) - if strings.EqualFold(heading, "Introduction") || strings.EqualFold(heading, "Overview") { - introStarted = true - introDone = false - } else { - introDone = true - } - continue - } - - // Inside a story block - if current != nil { - // **Status:** line - if m := statusLineRegex.FindStringSubmatch(trimmed); m != nil { - status := strings.TrimSpace(strings.ToLower(m[1])) - switch status { - case "done", "complete", "completed", "passed": - current.story.Passes = true - current.story.InProgress = false - case "in-progress", "in progress", "started": - current.story.InProgress = true - current.story.Passes = false - default: - current.story.Passes = false - current.story.InProgress = false - } - continue - } - - // **Priority:** line - if m := priorityLineRegex.FindStringSubmatch(trimmed); m != nil { - val := strings.TrimSpace(m[1]) - var pri float64 - if _, err := fmt.Sscanf(val, "%g", &pri); err == nil && pri > 0 { - current.story.Priority = pri - } - continue - } - - // **Description:** line - if m := descriptionLineRegex.FindStringSubmatch(trimmed); m != nil { - current.story.Description = strings.TrimSpace(m[1]) - continue - } - - // Checkbox items → acceptance criteria - if m := checkboxRegex.FindStringSubmatch(trimmed); m != nil { - current.story.AcceptanceCriteria = append(current.story.AcceptanceCriteria, strings.TrimSpace(m[2])) - continue - } - - // Collect prose lines as implicit description (only if no explicit **Description:** yet) - if trimmed != "" && current.story.Description == "" && - !strings.HasPrefix(trimmed, "**") && - !strings.HasPrefix(trimmed, "- ") { - current.descLines = append(current.descLines, trimmed) - } - continue - } - - // Collect introduction paragraph as project description - if introStarted && !introDone && p.Description == "" { - if trimmed != "" && !strings.HasPrefix(trimmed, "#") { - p.Description = trimmed - } - } - } - - // Flush the last story - flushStory() - - return p, nil -} diff --git a/internal/prd/markdown_test.go b/internal/prd/markdown_test.go deleted file mode 100644 index 19c1081c..00000000 --- a/internal/prd/markdown_test.go +++ /dev/null @@ -1,414 +0,0 @@ -package prd - -import ( - "os" - "path/filepath" - "testing" -) - -func TestParseMarkdownPRDFromString_Normal(t *testing.T) { - md := `# PRD: My Test Project - -A sample project for testing. - -## User Stories - -### US-001: Setup Project -As a developer, I need a properly structured project. - -**Priority:** 1 -**Status:** done - -- [x] Create project structure -- [x] Add dependencies - -### US-002: Add Feature -**Description:** As a user, I want a new feature. - -**Status:** in-progress - -- [ ] Feature works correctly -- [ ] Tests pass - -### US-003: Final Polish -Some prose description here. - -- [ ] Polish the UI -- [ ] Write docs -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("ParseMarkdownPRDFromString() error = %v", err) - } - - if p.Project != "My Test Project" { - t.Errorf("Project = %q, want %q", p.Project, "My Test Project") - } - if p.Description != "A sample project for testing." { - t.Errorf("Description = %q, want %q", p.Description, "A sample project for testing.") - } - if len(p.UserStories) != 3 { - t.Fatalf("len(UserStories) = %d, want 3", len(p.UserStories)) - } - - // Story 1: done - s1 := p.UserStories[0] - if s1.ID != "US-001" { - t.Errorf("s1.ID = %q, want %q", s1.ID, "US-001") - } - if s1.Title != "Setup Project" { - t.Errorf("s1.Title = %q, want %q", s1.Title, "Setup Project") - } - if !s1.Passes { - t.Error("s1.Passes = false, want true") - } - if s1.InProgress { - t.Error("s1.InProgress = true, want false") - } - if s1.Priority != 1 { - t.Errorf("s1.Priority = %g, want 1", s1.Priority) - } - if len(s1.AcceptanceCriteria) != 2 { - t.Errorf("len(s1.AcceptanceCriteria) = %d, want 2", len(s1.AcceptanceCriteria)) - } - - // Story 2: in-progress - s2 := p.UserStories[1] - if s2.ID != "US-002" { - t.Errorf("s2.ID = %q, want %q", s2.ID, "US-002") - } - if !s2.InProgress { - t.Error("s2.InProgress = false, want true") - } - if s2.Passes { - t.Error("s2.Passes = true, want false") - } - if s2.Description != "As a user, I want a new feature." { - t.Errorf("s2.Description = %q, want %q", s2.Description, "As a user, I want a new feature.") - } - - // Story 3: pending (no status) - s3 := p.UserStories[2] - if s3.ID != "US-003" { - t.Errorf("s3.ID = %q, want %q", s3.ID, "US-003") - } - if s3.Passes || s3.InProgress { - t.Error("s3 should be pending (both false)") - } - if s3.Description != "Some prose description here." { - t.Errorf("s3.Description = %q, want %q", s3.Description, "Some prose description here.") - } - if len(s3.AcceptanceCriteria) != 2 { - t.Errorf("len(s3.AcceptanceCriteria) = %d, want 2", len(s3.AcceptanceCriteria)) - } -} - -func TestParseMarkdownPRDFromString_ProjectWithoutPRDPrefix(t *testing.T) { - md := `# My Project - -Overview text. - -### FEAT-001: First Feature -- [ ] It works -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if p.Project != "My Project" { - t.Errorf("Project = %q, want %q", p.Project, "My Project") - } - if len(p.UserStories) != 1 { - t.Fatalf("len(UserStories) = %d, want 1", len(p.UserStories)) - } - if p.UserStories[0].ID != "FEAT-001" { - t.Errorf("ID = %q, want %q", p.UserStories[0].ID, "FEAT-001") - } -} - -func TestParseMarkdownPRDFromString_MissingFields(t *testing.T) { - md := `# Minimal - -### US-001: Only Title -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if len(p.UserStories) != 1 { - t.Fatalf("len(UserStories) = %d, want 1", len(p.UserStories)) - } - s := p.UserStories[0] - if s.ID != "US-001" { - t.Errorf("ID = %q", s.ID) - } - if s.Priority != 1 { - t.Errorf("Priority = %g, want 1 (auto-assigned)", s.Priority) - } - if s.Passes || s.InProgress { - t.Error("should be pending") - } -} - -func TestParseMarkdownPRDFromString_PhaseHeadingsIgnored(t *testing.T) { - md := `# My Project - -## Phase 1: Setup - -### US-001: Do Setup -- [ ] Setup done - -## Phase 2: Build - -### US-002: Do Build -- [ ] Build done -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if len(p.UserStories) != 2 { - t.Fatalf("len(UserStories) = %d, want 2", len(p.UserStories)) - } - if p.UserStories[0].ID != "US-001" { - t.Errorf("first story ID = %q", p.UserStories[0].ID) - } - if p.UserStories[1].ID != "US-002" { - t.Errorf("second story ID = %q", p.UserStories[1].ID) - } -} - -func TestParseMarkdownPRDFromString_IntroductionSection(t *testing.T) { - md := `# PRD: Test - -## Introduction - -This is the introduction paragraph. - -## Stories - -### US-001: First -- [ ] Done -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if p.Description != "This is the introduction paragraph." { - t.Errorf("Description = %q, want %q", p.Description, "This is the introduction paragraph.") - } -} - -func TestParseMarkdownPRDFromString_FreeSections(t *testing.T) { - md := `# Test Project - -Overview. - -## Background - -Some background text. - -## User Stories - -### US-001: First -- [ ] Criterion A - -## Appendix - -Extra info here. -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if len(p.UserStories) != 1 { - t.Fatalf("len(UserStories) = %d, want 1", len(p.UserStories)) - } - if p.UserStories[0].ID != "US-001" { - t.Errorf("ID = %q", p.UserStories[0].ID) - } -} - -func TestParseMarkdownPRDFromString_StatusMapping(t *testing.T) { - tests := []struct { - status string - wantPasses bool - wantIP bool - }{ - {"done", true, false}, - {"complete", true, false}, - {"completed", true, false}, - {"passed", true, false}, - {"in-progress", false, true}, - {"in progress", false, true}, - {"started", false, true}, - {"todo", false, false}, - {"pending", false, false}, - {"", false, false}, - } - - for _, tt := range tests { - t.Run(tt.status, func(t *testing.T) { - md := "# P\n\n### US-001: S\n**Status:** " + tt.status + "\n" - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if len(p.UserStories) != 1 { - t.Fatalf("len(UserStories) = %d", len(p.UserStories)) - } - if p.UserStories[0].Passes != tt.wantPasses { - t.Errorf("Passes = %v, want %v", p.UserStories[0].Passes, tt.wantPasses) - } - if p.UserStories[0].InProgress != tt.wantIP { - t.Errorf("InProgress = %v, want %v", p.UserStories[0].InProgress, tt.wantIP) - } - }) - } -} - -func TestParseMarkdownPRD_File(t *testing.T) { - tmpDir := t.TempDir() - prdPath := filepath.Join(tmpDir, "prd.md") - - md := `# Test - -### US-001: First Story -- [ ] Works -` - if err := os.WriteFile(prdPath, []byte(md), 0644); err != nil { - t.Fatalf("failed to write: %v", err) - } - - p, err := ParseMarkdownPRD(prdPath) - if err != nil { - t.Fatalf("ParseMarkdownPRD() error = %v", err) - } - if p.Project != "Test" { - t.Errorf("Project = %q", p.Project) - } - if len(p.UserStories) != 1 { - t.Fatalf("len(UserStories) = %d", len(p.UserStories)) - } -} - -func TestParseMarkdownPRD_FileNotFound(t *testing.T) { - _, err := ParseMarkdownPRD("/nonexistent/prd.md") - if err == nil { - t.Error("expected error for nonexistent file") - } -} - -func TestParseMarkdownPRDFromString_H4StoryHeadings(t *testing.T) { - md := `# PRD: Phased Project - -## Phase 1: Foundation - -### Design System - -#### US-001: Setup Theme -**Priority:** 1 -- [ ] Theme configured -- [ ] Tokens defined - -#### US-002: Build Components -- [ ] Components built - -## Phase 2: Features - -### Core Features - -#### US-003: Add Feature -**Status:** done -- [x] Feature works -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if p.Project != "Phased Project" { - t.Errorf("Project = %q, want %q", p.Project, "Phased Project") - } - if len(p.UserStories) != 3 { - t.Fatalf("len(UserStories) = %d, want 3", len(p.UserStories)) - } - if p.UserStories[0].ID != "US-001" { - t.Errorf("s1.ID = %q", p.UserStories[0].ID) - } - if len(p.UserStories[0].AcceptanceCriteria) != 2 { - t.Errorf("s1 AC count = %d, want 2", len(p.UserStories[0].AcceptanceCriteria)) - } - if p.UserStories[2].ID != "US-003" { - t.Errorf("s3.ID = %q", p.UserStories[2].ID) - } - if !p.UserStories[2].Passes { - t.Error("s3 should be done") - } -} - -func TestParseMarkdownPRDFromString_AutoPriority(t *testing.T) { - md := `# P - -### US-001: First -- [ ] A - -### US-002: Second -**Priority:** 5 -- [ ] B - -### US-003: Third -- [ ] C -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if len(p.UserStories) != 3 { - t.Fatalf("len = %d", len(p.UserStories)) - } - // First: auto-priority 1 - if p.UserStories[0].Priority != 1 { - t.Errorf("s1.Priority = %g, want 1", p.UserStories[0].Priority) - } - // Second: explicit priority 5 - if p.UserStories[1].Priority != 5 { - t.Errorf("s2.Priority = %g, want 5", p.UserStories[1].Priority) - } - // Third: auto-priority 6 (after 5) - if p.UserStories[2].Priority != 6 { - t.Errorf("s3.Priority = %g, want 6", p.UserStories[2].Priority) - } -} - -func TestParseMarkdownPRDFromString_FloatPriority(t *testing.T) { - md := `# P - -### US-001: First -**Priority:** 0.1 -- [ ] A - -### US-002: Second -**Priority:** 1.5 -- [ ] B - -### US-003: Third -**Priority:** 2 -- [ ] C -` - p, err := ParseMarkdownPRDFromString(md) - if err != nil { - t.Fatalf("error = %v", err) - } - if len(p.UserStories) != 3 { - t.Fatalf("len = %d", len(p.UserStories)) - } - if p.UserStories[0].Priority != 0.1 { - t.Errorf("s1.Priority = %g, want 0.1", p.UserStories[0].Priority) - } - if p.UserStories[1].Priority != 1.5 { - t.Errorf("s2.Priority = %g, want 1.5", p.UserStories[1].Priority) - } - if p.UserStories[2].Priority != 2 { - t.Errorf("s3.Priority = %g, want 2", p.UserStories[2].Priority) - } -} diff --git a/internal/prd/markdown_writer.go b/internal/prd/markdown_writer.go deleted file mode 100644 index 3696f8d1..00000000 --- a/internal/prd/markdown_writer.go +++ /dev/null @@ -1,89 +0,0 @@ -package prd - -import ( - "fmt" - "os" - "regexp" - "strings" -) - -// SetStoryStatus performs a surgical update of a story's status in a prd.md file. -// It finds the story block by its heading, updates or inserts the **Status:** line, -// and when status is "done", flips all unchecked checkboxes to checked. -func SetStoryStatus(path, storyID, status string) error { - data, err := os.ReadFile(path) - if err != nil { - return fmt.Errorf("failed to read PRD file: %w", err) - } - - result, err := setStoryStatusInString(string(data), storyID, status) - if err != nil { - return err - } - - return os.WriteFile(path, []byte(result), 0644) -} - -// setStoryStatusInString performs the status update on a string and returns the modified string. -func setStoryStatusInString(content, storyID, status string) (string, error) { - lines := strings.Split(content, "\n") - - // Find the story block - storyStart := -1 - storyEnd := len(lines) // default to end of file - - headingPattern := regexp.MustCompile(`^#{3,4}\s+` + regexp.QuoteMeta(storyID) + `:\s+`) - - for i, line := range lines { - if storyStart == -1 { - // Looking for the story heading - if headingPattern.MatchString(strings.TrimSpace(line)) { - storyStart = i - } - } else { - // Looking for the end of the story block (next ## or ### heading) - trimmed := strings.TrimSpace(line) - if strings.HasPrefix(trimmed, "## ") || strings.HasPrefix(trimmed, "### ") || strings.HasPrefix(trimmed, "#### ") { - storyEnd = i - break - } - } - } - - if storyStart == -1 { - return "", fmt.Errorf("story %s not found in PRD", storyID) - } - - // Process the story block - statusLineIdx := -1 - statusLine := fmt.Sprintf("**Status:** %s", status) - - for i := storyStart + 1; i < storyEnd; i++ { - if statusLineRegex.MatchString(strings.TrimSpace(lines[i])) { - statusLineIdx = i - break - } - } - - if statusLineIdx >= 0 { - // Replace existing status line - lines[statusLineIdx] = statusLine - } else { - // Insert status line as first line after heading - newLines := make([]string, 0, len(lines)+1) - newLines = append(newLines, lines[:storyStart+1]...) - newLines = append(newLines, statusLine) - newLines = append(newLines, lines[storyStart+1:]...) - lines = newLines - storyEnd++ // adjust for the inserted line - } - - // When status is "done", flip all unchecked checkboxes to checked - if status == "done" { - for i := storyStart + 1; i < storyEnd; i++ { - lines[i] = strings.Replace(lines[i], "- [ ]", "- [x]", 1) - } - } - - return strings.Join(lines, "\n"), nil -} diff --git a/internal/prd/markdown_writer_test.go b/internal/prd/markdown_writer_test.go deleted file mode 100644 index 5ea694c2..00000000 --- a/internal/prd/markdown_writer_test.go +++ /dev/null @@ -1,270 +0,0 @@ -package prd - -import ( - "os" - "path/filepath" - "strings" - "testing" -) - -func TestSetStoryStatusInString_ExistingStatusLine(t *testing.T) { - md := `# P - -### US-001: First -**Status:** todo -- [ ] A -- [ ] B - -### US-002: Second -- [ ] C -` - result, err := setStoryStatusInString(md, "US-001", "done") - if err != nil { - t.Fatalf("error = %v", err) - } - - if !strings.Contains(result, "**Status:** done") { - t.Error("expected **Status:** done in result") - } - // Should not contain the old status - if strings.Contains(result, "**Status:** todo") { - t.Error("old status should be replaced") - } - // Checkboxes should be flipped to checked - if strings.Contains(result, "- [ ] A") { - t.Error("expected checkbox A to be checked") - } - if !strings.Contains(result, "- [x] A") { - t.Error("expected checkbox A to be [x]") - } - // US-002 should be untouched - if !strings.Contains(result, "- [ ] C") { - t.Error("US-002 checkboxes should be untouched") - } -} - -func TestSetStoryStatusInString_MissingStatusLine(t *testing.T) { - md := `# P - -### US-001: First -- [ ] A -` - result, err := setStoryStatusInString(md, "US-001", "in-progress") - if err != nil { - t.Fatalf("error = %v", err) - } - - if !strings.Contains(result, "**Status:** in-progress") { - t.Error("expected **Status:** in-progress to be inserted") - } - - // Status line should appear after the heading - lines := strings.Split(result, "\n") - for i, line := range lines { - if strings.Contains(line, "### US-001") { - if i+1 >= len(lines) || !strings.Contains(lines[i+1], "**Status:** in-progress") { - t.Error("status line should be directly after heading") - } - break - } - } -} - -func TestSetStoryStatusInString_CheckboxFlipping(t *testing.T) { - md := `# P - -### US-001: First -**Status:** in-progress -- [ ] Unchecked A -- [x] Already checked B -- [ ] Unchecked C -` - result, err := setStoryStatusInString(md, "US-001", "done") - if err != nil { - t.Fatalf("error = %v", err) - } - - if strings.Contains(result, "- [ ] Unchecked A") { - t.Error("checkbox A should be checked") - } - if !strings.Contains(result, "- [x] Unchecked A") { - t.Error("expected [x] Unchecked A") - } - if !strings.Contains(result, "- [x] Already checked B") { - t.Error("already checked B should remain checked") - } - if !strings.Contains(result, "- [x] Unchecked C") { - t.Error("checkbox C should be checked") - } -} - -func TestSetStoryStatusInString_MultiStory(t *testing.T) { - md := `# P - -### US-001: First -**Status:** todo -- [ ] A - -### US-002: Second -**Status:** todo -- [ ] B - -### US-003: Third -- [ ] C -` - // Mark US-002 as done - result, err := setStoryStatusInString(md, "US-002", "done") - if err != nil { - t.Fatalf("error = %v", err) - } - - // US-001 should be unchanged - if !strings.Contains(result, "- [ ] A") { - t.Error("US-001 checkboxes should be untouched") - } - // US-003 should be unchanged - if !strings.Contains(result, "- [ ] C") { - t.Error("US-003 checkboxes should be untouched") - } - // US-002 should be done with checked boxes - if !strings.Contains(result, "- [x] B") { - t.Error("US-002 checkbox should be checked") - } -} - -func TestSetStoryStatusInString_StoryNotFound(t *testing.T) { - md := `# P - -### US-001: First -- [ ] A -` - _, err := setStoryStatusInString(md, "US-999", "done") - if err == nil { - t.Error("expected error for missing story") - } -} - -func TestSetStoryStatus_File(t *testing.T) { - tmpDir := t.TempDir() - prdPath := filepath.Join(tmpDir, "prd.md") - - md := `# P - -### US-001: First -- [ ] A -` - if err := os.WriteFile(prdPath, []byte(md), 0644); err != nil { - t.Fatalf("failed to write: %v", err) - } - - if err := SetStoryStatus(prdPath, "US-001", "done"); err != nil { - t.Fatalf("SetStoryStatus() error = %v", err) - } - - data, err := os.ReadFile(prdPath) - if err != nil { - t.Fatalf("failed to read: %v", err) - } - - result := string(data) - if !strings.Contains(result, "**Status:** done") { - t.Error("expected **Status:** done in file") - } - if !strings.Contains(result, "- [x] A") { - t.Error("expected checkbox to be checked") - } -} - -func TestSetStoryStatusInString_H4Headings(t *testing.T) { - md := `# P - -## Phase 1 - -#### US-001: First -- [ ] A - -#### US-002: Second -- [ ] B -` - result, err := setStoryStatusInString(md, "US-001", "done") - if err != nil { - t.Fatalf("error = %v", err) - } - if !strings.Contains(result, "**Status:** done") { - t.Error("expected **Status:** done") - } - if !strings.Contains(result, "- [x] A") { - t.Error("expected checkbox A to be checked") - } - // US-002 should be untouched - if !strings.Contains(result, "- [ ] B") { - t.Error("US-002 should be untouched") - } -} - -func TestSetStoryStatusInString_NoCheckboxFlipForNonDone(t *testing.T) { - md := `# P - -### US-001: First -- [ ] A -` - result, err := setStoryStatusInString(md, "US-001", "in-progress") - if err != nil { - t.Fatalf("error = %v", err) - } - - // Checkboxes should NOT be flipped for in-progress - if !strings.Contains(result, "- [ ] A") { - t.Error("checkboxes should not be flipped for non-done status") - } -} - -func TestSetStoryStatusInString_RoundTrip(t *testing.T) { - md := `# My Project - -A description. - -### US-001: First -**Status:** todo -- [ ] A -- [ ] B - -### US-002: Second -- [ ] C -` - // Set US-001 to in-progress - result, err := setStoryStatusInString(md, "US-001", "in-progress") - if err != nil { - t.Fatalf("error = %v", err) - } - - // Parse and verify - p, err := ParseMarkdownPRDFromString(result) - if err != nil { - t.Fatalf("parse error = %v", err) - } - if !p.UserStories[0].InProgress { - t.Error("US-001 should be in-progress") - } - if p.UserStories[0].Passes { - t.Error("US-001 should not be passes") - } - - // Now set US-001 to done - result, err = setStoryStatusInString(result, "US-001", "done") - if err != nil { - t.Fatalf("error = %v", err) - } - - // Parse and verify - p, err = ParseMarkdownPRDFromString(result) - if err != nil { - t.Fatalf("parse error = %v", err) - } - if !p.UserStories[0].Passes { - t.Error("US-001 should be passes") - } - if p.UserStories[0].InProgress { - t.Error("US-001 should not be in-progress") - } -} diff --git a/internal/prd/migrate.go b/internal/prd/migrate.go deleted file mode 100644 index 6dd5ad9e..00000000 --- a/internal/prd/migrate.go +++ /dev/null @@ -1,53 +0,0 @@ -package prd - -import ( - "encoding/json" - "fmt" - "os" - "path/filepath" -) - -// MigrateFromJSON reads prd.json, transfers story statuses into prd.md -// using SetStoryStatus, and renames prd.json to prd.json.bak. -func MigrateFromJSON(prdDir string) error { - jsonPath := filepath.Join(prdDir, "prd.json") - mdPath := filepath.Join(prdDir, "prd.md") - - // Read and parse prd.json - data, err := os.ReadFile(jsonPath) - if err != nil { - return fmt.Errorf("failed to read prd.json: %w", err) - } - - var p PRD - if err := json.Unmarshal(data, &p); err != nil { - return fmt.Errorf("failed to parse prd.json: %w", err) - } - - // Check that prd.md exists - if _, err := os.Stat(mdPath); err != nil { - return fmt.Errorf("prd.md not found: %w", err) - } - - // Transfer statuses - for _, story := range p.UserStories { - if story.Passes { - if err := SetStoryStatus(mdPath, story.ID, "done"); err != nil { - // Non-fatal: story might not exist in prd.md (was removed) - continue - } - } else if story.InProgress { - if err := SetStoryStatus(mdPath, story.ID, "in-progress"); err != nil { - continue - } - } - } - - // Rename prd.json → prd.json.bak - bakPath := filepath.Join(prdDir, "prd.json.bak") - if err := os.Rename(jsonPath, bakPath); err != nil { - return fmt.Errorf("failed to rename prd.json to prd.json.bak: %w", err) - } - - return nil -} diff --git a/internal/prd/migrate_test.go b/internal/prd/migrate_test.go deleted file mode 100644 index 055748ac..00000000 --- a/internal/prd/migrate_test.go +++ /dev/null @@ -1,153 +0,0 @@ -package prd - -import ( - "os" - "path/filepath" - "strings" - "testing" -) - -func TestMigrateFromJSON(t *testing.T) { - tmpDir := t.TempDir() - - // Create prd.json with some statuses - jsonContent := `{ - "project": "Test", - "description": "A test", - "userStories": [ - {"id": "US-001", "title": "Done Story", "passes": true, "priority": 1}, - {"id": "US-002", "title": "In Progress Story", "passes": false, "inProgress": true, "priority": 2}, - {"id": "US-003", "title": "Pending Story", "passes": false, "priority": 3} - ] -}` - if err := os.WriteFile(filepath.Join(tmpDir, "prd.json"), []byte(jsonContent), 0644); err != nil { - t.Fatalf("failed to write prd.json: %v", err) - } - - // Create prd.md with the same stories - mdContent := `# Test - -### US-001: Done Story -- [ ] A - -### US-002: In Progress Story -- [ ] B - -### US-003: Pending Story -- [ ] C -` - if err := os.WriteFile(filepath.Join(tmpDir, "prd.md"), []byte(mdContent), 0644); err != nil { - t.Fatalf("failed to write prd.md: %v", err) - } - - // Run migration - if err := MigrateFromJSON(tmpDir); err != nil { - t.Fatalf("MigrateFromJSON() error = %v", err) - } - - // Verify prd.json is renamed to prd.json.bak - if _, err := os.Stat(filepath.Join(tmpDir, "prd.json")); !os.IsNotExist(err) { - t.Error("prd.json should be renamed") - } - if _, err := os.Stat(filepath.Join(tmpDir, "prd.json.bak")); err != nil { - t.Error("prd.json.bak should exist") - } - - // Verify prd.md has the correct statuses - data, err := os.ReadFile(filepath.Join(tmpDir, "prd.md")) - if err != nil { - t.Fatalf("failed to read prd.md: %v", err) - } - result := string(data) - - // US-001 should be done with checked boxes - if !strings.Contains(result, "- [x] A") { - t.Error("US-001 should have checked checkbox") - } - - // Parse and verify - p, err := ParseMarkdownPRD(filepath.Join(tmpDir, "prd.md")) - if err != nil { - t.Fatalf("ParseMarkdownPRD() error = %v", err) - } - - if !p.UserStories[0].Passes { - t.Error("US-001 should be passes") - } - if !p.UserStories[1].InProgress { - t.Error("US-002 should be in-progress") - } - if p.UserStories[2].Passes || p.UserStories[2].InProgress { - t.Error("US-003 should be pending") - } -} - -func TestMigrateFromJSON_MissingStoryInMd(t *testing.T) { - tmpDir := t.TempDir() - - // prd.json has a story that doesn't exist in prd.md - jsonContent := `{ - "project": "Test", - "userStories": [ - {"id": "US-001", "title": "Exists", "passes": true, "priority": 1}, - {"id": "US-999", "title": "Missing", "passes": true, "priority": 2} - ] -}` - if err := os.WriteFile(filepath.Join(tmpDir, "prd.json"), []byte(jsonContent), 0644); err != nil { - t.Fatalf("failed to write: %v", err) - } - - mdContent := `# Test - -### US-001: Exists -- [ ] A -` - if err := os.WriteFile(filepath.Join(tmpDir, "prd.md"), []byte(mdContent), 0644); err != nil { - t.Fatalf("failed to write: %v", err) - } - - // Should not error — missing stories are skipped - if err := MigrateFromJSON(tmpDir); err != nil { - t.Fatalf("MigrateFromJSON() error = %v", err) - } - - // Verify the existing story was migrated - p, err := ParseMarkdownPRD(filepath.Join(tmpDir, "prd.md")) - if err != nil { - t.Fatalf("parse error = %v", err) - } - if !p.UserStories[0].Passes { - t.Error("US-001 should be passes") - } -} - -func TestMigrateFromJSON_NoJsonFile(t *testing.T) { - tmpDir := t.TempDir() - - mdContent := `# Test -### US-001: First -- [ ] A -` - if err := os.WriteFile(filepath.Join(tmpDir, "prd.md"), []byte(mdContent), 0644); err != nil { - t.Fatalf("failed to write: %v", err) - } - - err := MigrateFromJSON(tmpDir) - if err == nil { - t.Error("expected error when prd.json doesn't exist") - } -} - -func TestMigrateFromJSON_NoMdFile(t *testing.T) { - tmpDir := t.TempDir() - - jsonContent := `{"project": "Test", "userStories": [{"id": "US-001", "passes": true, "priority": 1}]}` - if err := os.WriteFile(filepath.Join(tmpDir, "prd.json"), []byte(jsonContent), 0644); err != nil { - t.Fatalf("failed to write: %v", err) - } - - err := MigrateFromJSON(tmpDir) - if err == nil { - t.Error("expected error when prd.md doesn't exist") - } -} diff --git a/internal/prd/prd_test.go b/internal/prd/prd_test.go index e4da61d5..fe59708e 100644 --- a/internal/prd/prd_test.go +++ b/internal/prd/prd_test.go @@ -1,30 +1,32 @@ package prd import ( - "encoding/json" - "fmt" "os" "path/filepath" "testing" ) func TestLoadPRD(t *testing.T) { - // Create a temp file with valid PRD markdown + // Create a temp file with valid PRD JSON tmpDir := t.TempDir() - prdPath := filepath.Join(tmpDir, "prd.md") + prdPath := filepath.Join(tmpDir, "prd.json") - validMd := `# Test Project - -A test PRD - -### US-001: First Story -**Description:** Test description - -- [ ] AC1 -- [ ] AC2 -` - - if err := os.WriteFile(prdPath, []byte(validMd), 0644); err != nil { + validJSON := `{ + "project": "Test Project", + "description": "A test PRD", + "userStories": [ + { + "id": "US-001", + "title": "First Story", + "description": "Test description", + "acceptanceCriteria": ["AC1", "AC2"], + "priority": 1, + "passes": false + } + ] + }` + + if err := os.WriteFile(prdPath, []byte(validJSON), 0644); err != nil { t.Fatalf("failed to write test file: %v", err) } @@ -48,12 +50,66 @@ A test PRD } func TestLoadPRD_FileNotFound(t *testing.T) { - _, err := LoadPRD("/nonexistent/path/prd.md") + _, err := LoadPRD("/nonexistent/path/prd.json") if err == nil { t.Error("expected error for nonexistent file, got nil") } } +func TestLoadPRD_InvalidJSON(t *testing.T) { + tmpDir := t.TempDir() + prdPath := filepath.Join(tmpDir, "prd.json") + + if err := os.WriteFile(prdPath, []byte("not valid json"), 0644); err != nil { + t.Fatalf("failed to write test file: %v", err) + } + + _, err := LoadPRD(prdPath) + if err == nil { + t.Error("expected error for invalid JSON, got nil") + } +} + +func TestPRD_Save(t *testing.T) { + tmpDir := t.TempDir() + prdPath := filepath.Join(tmpDir, "prd.json") + + p := &PRD{ + Project: "Saved Project", + Description: "A saved PRD", + UserStories: []UserStory{ + { + ID: "US-001", + Title: "Test Story", + Description: "Test", + AcceptanceCriteria: []string{"AC1"}, + Priority: 1, + Passes: true, + }, + }, + } + + if err := p.Save(prdPath); err != nil { + t.Fatalf("Save failed: %v", err) + } + + // Verify by loading it back + loaded, err := LoadPRD(prdPath) + if err != nil { + t.Fatalf("LoadPRD after Save failed: %v", err) + } + + if loaded.Project != p.Project { + t.Errorf("expected project '%s', got '%s'", p.Project, loaded.Project) + } + if len(loaded.UserStories) != 1 { + t.Errorf("expected 1 user story, got %d", len(loaded.UserStories)) + } + if !loaded.UserStories[0].Passes { + t.Error("expected story to have passes: true") + } +} + func TestPRD_AllComplete_EmptyPRD(t *testing.T) { p := &PRD{ Project: "Empty", @@ -180,6 +236,7 @@ func TestPRD_NextStory_SkipsCompleted(t *testing.T) { } func TestPRD_NextStory_InterruptedTakesPrecedence(t *testing.T) { + // Even if there's a lower priority story, in-progress takes precedence p := &PRD{ Project: "Test", UserStories: []UserStory{ @@ -216,186 +273,33 @@ func TestUserStory_Fields(t *testing.T) { } } -func TestPRD_NextStoryContext_ReturnsHighestPriority(t *testing.T) { - p := &PRD{ - Project: "Test", - UserStories: []UserStory{ - {ID: "US-001", Title: "Low priority", Priority: 3, Passes: false}, - {ID: "US-002", Title: "High priority", Priority: 1, Passes: false}, - {ID: "US-003", Title: "Mid priority", Priority: 2, Passes: false}, - }, - } - - ctx := p.NextStoryContext() - if ctx == nil { - t.Fatal("expected non-nil context") - } - - var story UserStory - if err := json.Unmarshal([]byte(*ctx), &story); err != nil { - t.Fatalf("failed to parse story context JSON: %v", err) - } - if story.ID != "US-002" { - t.Errorf("expected highest-priority story US-002, got %s", story.ID) - } -} - -func TestPRD_NextStoryContext_ReturnsNilWhenAllComplete(t *testing.T) { - p := &PRD{ - Project: "Test", - UserStories: []UserStory{ - {ID: "US-001", Passes: true}, - {ID: "US-002", Passes: true}, - }, - } - - ctx := p.NextStoryContext() - if ctx != nil { - t.Errorf("expected nil when all stories complete, got %q", *ctx) - } -} - -func TestPRD_NextStoryContext_SkipsPassingStories(t *testing.T) { - p := &PRD{ - Project: "Test", - UserStories: []UserStory{ - {ID: "US-001", Title: "Done", Priority: 1, Passes: true}, - {ID: "US-002", Title: "Pending", Priority: 2, Passes: false}, - }, - } - - ctx := p.NextStoryContext() - if ctx == nil { - t.Fatal("expected non-nil context") - } - - var story UserStory - if err := json.Unmarshal([]byte(*ctx), &story); err != nil { - t.Fatalf("failed to parse story context JSON: %v", err) - } - if story.ID != "US-002" { - t.Errorf("expected US-002 (only pending story), got %s", story.ID) - } -} - -func TestPRD_NextStoryContext_EmptyPRD(t *testing.T) { - p := &PRD{ - Project: "Empty", - UserStories: []UserStory{}, - } - - ctx := p.NextStoryContext() - if ctx != nil { - t.Errorf("expected nil for empty PRD, got %q", *ctx) - } -} +func TestPRD_Save_PreservesInProgress(t *testing.T) { + tmpDir := t.TempDir() + prdPath := filepath.Join(tmpDir, "prd.json") -func TestPRD_NextStoryContext_ValidJSON(t *testing.T) { p := &PRD{ Project: "Test", UserStories: []UserStory{ { - ID: "US-001", - Title: "Test Story", - Description: "A test description", - AcceptanceCriteria: []string{"AC1", "AC2"}, - Priority: 1, - Passes: false, + ID: "US-001", + Title: "Story", + Priority: 1, + Passes: false, + InProgress: true, }, }, } - ctx := p.NextStoryContext() - if ctx == nil { - t.Fatal("expected non-nil context") - } - - var story UserStory - if err := json.Unmarshal([]byte(*ctx), &story); err != nil { - t.Fatalf("NextStoryContext did not return valid JSON: %v", err) - } - if story.ID != "US-001" { - t.Errorf("expected ID US-001, got %s", story.ID) - } - if story.Title != "Test Story" { - t.Errorf("expected title 'Test Story', got '%s'", story.Title) - } - if len(story.AcceptanceCriteria) != 2 { - t.Errorf("expected 2 acceptance criteria, got %d", len(story.AcceptanceCriteria)) - } -} - -func TestPRD_NextStoryContext_PromptSizeUnder10KB(t *testing.T) { - stories := make([]UserStory, 300) - for i := range stories { - stories[i] = UserStory{ - ID: fmt.Sprintf("US-%03d", i+1), - Title: fmt.Sprintf("Story %d with a reasonably long title for realism", i+1), - Description: "This is a description that is moderately long to simulate realistic PRD content for testing purposes.", - AcceptanceCriteria: []string{"Criterion A", "Criterion B", "Criterion C"}, - Priority: float64(i + 1), - Passes: i > 0, - } - } - p := &PRD{ - Project: "Large Project", - Description: "A large PRD with 300 stories", - UserStories: stories, - } - - ctx := p.NextStoryContext() - if ctx == nil { - t.Fatal("expected non-nil context for 300-story PRD") - } - if len(*ctx) > 10*1024 { - t.Errorf("story context is %d bytes, expected under 10KB", len(*ctx)) - } -} - -func TestPRD_ExtractIDPrefix_US(t *testing.T) { - p := &PRD{ - Project: "Test", - UserStories: []UserStory{ - {ID: "US-001"}, - {ID: "US-002"}, - }, - } - if got := p.ExtractIDPrefix(); got != "US" { - t.Errorf("ExtractIDPrefix() = %q, want %q", got, "US") - } -} - -func TestPRD_ExtractIDPrefix_MFR(t *testing.T) { - p := &PRD{ - Project: "Test", - UserStories: []UserStory{ - {ID: "MFR-001"}, - {ID: "MFR-002"}, - }, - } - if got := p.ExtractIDPrefix(); got != "MFR" { - t.Errorf("ExtractIDPrefix() = %q, want %q", got, "MFR") + if err := p.Save(prdPath); err != nil { + t.Fatalf("Save failed: %v", err) } -} -func TestPRD_ExtractIDPrefix_Default(t *testing.T) { - p := &PRD{ - Project: "Empty", - UserStories: []UserStory{}, - } - if got := p.ExtractIDPrefix(); got != "US" { - t.Errorf("ExtractIDPrefix() = %q, want %q for empty PRD", got, "US") + loaded, err := LoadPRD(prdPath) + if err != nil { + t.Fatalf("LoadPRD failed: %v", err) } -} -func TestPRD_ExtractIDPrefix_SingleChar(t *testing.T) { - p := &PRD{ - Project: "Test", - UserStories: []UserStory{ - {ID: "T-001"}, - }, - } - if got := p.ExtractIDPrefix(); got != "T" { - t.Errorf("ExtractIDPrefix() = %q, want %q", got, "T") + if !loaded.UserStories[0].InProgress { + t.Error("expected InProgress to be preserved as true") } } diff --git a/internal/prd/types.go b/internal/prd/types.go index d8e428c0..ac70dc3d 100644 --- a/internal/prd/types.go +++ b/internal/prd/types.go @@ -3,19 +3,13 @@ // for changes, and converting between prd.md and prd.json formats. package prd -import ( - "encoding/json" - "fmt" - "strings" -) - // UserStory represents a single user story in a PRD. type UserStory struct { ID string `json:"id"` Title string `json:"title"` Description string `json:"description"` AcceptanceCriteria []string `json:"acceptanceCriteria"` - Priority float64 `json:"priority"` + Priority int `json:"priority"` Passes bool `json:"passes"` InProgress bool `json:"inProgress,omitempty"` } @@ -27,18 +21,6 @@ type PRD struct { UserStories []UserStory `json:"userStories"` } -// ExtractIDPrefix returns the ID prefix used by the stories in this PRD. -// For example, "US" from "US-001", "MFR" from "MFR-001", "T" from "T-001". -// Returns "US" as the default when the PRD has no stories or IDs lack a hyphen. -func (p *PRD) ExtractIDPrefix() string { - for _, story := range p.UserStories { - if idx := strings.LastIndex(story.ID, "-"); idx > 0 { - return story.ID[:idx] - } - } - return "US" -} - // AllComplete returns true when all stories have passes: true. func (p *PRD) AllComplete() bool { if len(p.UserStories) == 0 { @@ -77,29 +59,3 @@ func (p *PRD) NextStory() *UserStory { } return next } - -// NextStoryContext returns the next story to work on as a formatted string -// suitable for inlining into the agent prompt. Returns nil when all stories -// are complete. -func (p *PRD) NextStoryContext() *string { - story := p.NextStory() - if story == nil { - return nil - } - - data, err := json.MarshalIndent(story, "", " ") - if err != nil { - // Fallback to a simple text format - var b strings.Builder - fmt.Fprintf(&b, "ID: %s\nTitle: %s\nDescription: %s\n", story.ID, story.Title, story.Description) - fmt.Fprintf(&b, "Acceptance Criteria:\n") - for _, ac := range story.AcceptanceCriteria { - fmt.Fprintf(&b, "- %s\n", ac) - } - result := b.String() - return &result - } - - result := string(data) - return &result -} diff --git a/internal/prd/watcher.go b/internal/prd/watcher.go index 09e2b2f5..d86d5a72 100644 --- a/internal/prd/watcher.go +++ b/internal/prd/watcher.go @@ -15,13 +15,13 @@ type WatcherEvent struct { // Watcher watches a prd.json file for changes and sends events. type Watcher struct { - path string - watcher *fsnotify.Watcher - events chan WatcherEvent - done chan struct{} - mu sync.Mutex - running bool - lastPRD *PRD + path string + watcher *fsnotify.Watcher + events chan WatcherEvent + done chan struct{} + mu sync.Mutex + running bool + lastPRD *PRD } // NewWatcher creates a new Watcher for the given PRD file path. @@ -110,7 +110,7 @@ func (w *Watcher) processEvents() { // Handle file removal - try to re-watch if event.Op&fsnotify.Remove != 0 { - w.events <- WatcherEvent{Error: errors.New("prd.md was removed")} + w.events <- WatcherEvent{Error: errors.New("prd.json was removed")} // Try to re-add the watch (file might be re-created) _ = w.watcher.Add(w.path) } diff --git a/internal/prd/watcher_test.go b/internal/prd/watcher_test.go index 36d715f7..5c7ca8b6 100644 --- a/internal/prd/watcher_test.go +++ b/internal/prd/watcher_test.go @@ -1,42 +1,28 @@ package prd import ( + "encoding/json" "os" "path/filepath" "testing" "time" ) -// createTestPRDMd creates a markdown PRD file for testing. -func createTestPRDMd(t *testing.T, dir string, stories []UserStory) string { - t.Helper() - prdPath := filepath.Join(dir, "prd.md") - - md := "# Test\n\n" - for _, s := range stories { - md += "### " + s.ID + ": " + s.Title + "\n" - if s.Passes { - md += "**Status:** done\n" - } else if s.InProgress { - md += "**Status:** in-progress\n" - } - if s.Description != "" { - md += "**Description:** " + s.Description + "\n" - } - md += "- [ ] criterion\n\n" - } +func TestNewWatcher(t *testing.T) { + tmpDir := t.TempDir() + prdPath := filepath.Join(tmpDir, "prd.json") - if err := os.WriteFile(prdPath, []byte(md), 0644); err != nil { + // Create a test PRD file + testPRD := &PRD{ + Project: "Test", + UserStories: []UserStory{ + {ID: "US-001", Title: "Test Story", Passes: false}, + }, + } + data, _ := json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { t.Fatalf("Failed to write test PRD: %v", err) } - return prdPath -} - -func TestNewWatcher(t *testing.T) { - tmpDir := t.TempDir() - prdPath := createTestPRDMd(t, tmpDir, []UserStory{ - {ID: "US-001", Title: "Test Story", Passes: false}, - }) watcher, err := NewWatcher(prdPath) if err != nil { @@ -51,9 +37,19 @@ func TestNewWatcher(t *testing.T) { func TestWatcherStart(t *testing.T) { tmpDir := t.TempDir() - prdPath := createTestPRDMd(t, tmpDir, []UserStory{ - {ID: "US-001", Title: "Test Story", Passes: false}, - }) + prdPath := filepath.Join(tmpDir, "prd.json") + + // Create a test PRD file + testPRD := &PRD{ + Project: "Test", + UserStories: []UserStory{ + {ID: "US-001", Title: "Test Story", Passes: false}, + }, + } + data, _ := json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { + t.Fatalf("Failed to write test PRD: %v", err) + } watcher, err := NewWatcher(prdPath) if err != nil { @@ -73,9 +69,19 @@ func TestWatcherStart(t *testing.T) { func TestWatcherDetectsFileChange(t *testing.T) { tmpDir := t.TempDir() - prdPath := createTestPRDMd(t, tmpDir, []UserStory{ - {ID: "US-001", Title: "Test Story", Passes: false}, - }) + prdPath := filepath.Join(tmpDir, "prd.json") + + // Create a test PRD file + testPRD := &PRD{ + Project: "Test", + UserStories: []UserStory{ + {ID: "US-001", Title: "Test Story", Passes: false}, + }, + } + data, _ := json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { + t.Fatalf("Failed to write test PRD: %v", err) + } watcher, err := NewWatcher(prdPath) if err != nil { @@ -87,13 +93,17 @@ func TestWatcherDetectsFileChange(t *testing.T) { t.Fatalf("Failed to start watcher: %v", err) } + // Give watcher time to initialize time.Sleep(100 * time.Millisecond) // Modify the file - change passes status - if err := SetStoryStatus(prdPath, "US-001", "done"); err != nil { + testPRD.UserStories[0].Passes = true + data, _ = json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { t.Fatalf("Failed to update test PRD: %v", err) } + // Wait for the event select { case event := <-watcher.Events(): if event.Error != nil { @@ -112,9 +122,19 @@ func TestWatcherDetectsFileChange(t *testing.T) { func TestWatcherDetectsInProgressChange(t *testing.T) { tmpDir := t.TempDir() - prdPath := createTestPRDMd(t, tmpDir, []UserStory{ - {ID: "US-001", Title: "Test Story", Passes: false, InProgress: false}, - }) + prdPath := filepath.Join(tmpDir, "prd.json") + + // Create a test PRD file + testPRD := &PRD{ + Project: "Test", + UserStories: []UserStory{ + {ID: "US-001", Title: "Test Story", Passes: false, InProgress: false}, + }, + } + data, _ := json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { + t.Fatalf("Failed to write test PRD: %v", err) + } watcher, err := NewWatcher(prdPath) if err != nil { @@ -126,13 +146,17 @@ func TestWatcherDetectsInProgressChange(t *testing.T) { t.Fatalf("Failed to start watcher: %v", err) } + // Give watcher time to initialize time.Sleep(100 * time.Millisecond) // Modify the file - change inProgress status - if err := SetStoryStatus(prdPath, "US-001", "in-progress"); err != nil { + testPRD.UserStories[0].InProgress = true + data, _ = json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { t.Fatalf("Failed to update test PRD: %v", err) } + // Wait for the event select { case event := <-watcher.Events(): if event.Error != nil { @@ -151,7 +175,7 @@ func TestWatcherDetectsInProgressChange(t *testing.T) { func TestWatcherHandlesFileNotFound(t *testing.T) { tmpDir := t.TempDir() - prdPath := filepath.Join(tmpDir, "nonexistent.md") + prdPath := filepath.Join(tmpDir, "nonexistent.json") watcher, err := NewWatcher(prdPath) if err != nil { @@ -159,26 +183,89 @@ func TestWatcherHandlesFileNotFound(t *testing.T) { } defer watcher.Stop() + // Start should still work, but we'll get an error event if err := watcher.Start(); err != nil { + // This is expected since the file doesn't exist + // But the watcher.Add might fail first + // Let's check that events channel has an error t.Logf("Got expected start error: %v", err) return } + // If start succeeded, check for error event select { case event := <-watcher.Events(): if event.Error == nil { t.Error("Expected error event for nonexistent file") } case <-time.After(1 * time.Second): + // Might not get event if watcher.Add failed t.Log("No error event received, which is acceptable if Add failed") } } +func TestWatcherIgnoresNonStatusChanges(t *testing.T) { + tmpDir := t.TempDir() + prdPath := filepath.Join(tmpDir, "prd.json") + + // Create a test PRD file + testPRD := &PRD{ + Project: "Test", + UserStories: []UserStory{ + {ID: "US-001", Title: "Test Story", Description: "Original", Passes: false}, + }, + } + data, _ := json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { + t.Fatalf("Failed to write test PRD: %v", err) + } + + watcher, err := NewWatcher(prdPath) + if err != nil { + t.Fatalf("Failed to create watcher: %v", err) + } + defer watcher.Stop() + + if err := watcher.Start(); err != nil { + t.Fatalf("Failed to start watcher: %v", err) + } + + // Give watcher time to initialize + time.Sleep(100 * time.Millisecond) + + // Modify the file - only change description (not status) + testPRD.UserStories[0].Description = "Modified" + data, _ = json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { + t.Fatalf("Failed to update test PRD: %v", err) + } + + // Should NOT receive an event since status didn't change + select { + case event := <-watcher.Events(): + if event.PRD != nil { + t.Error("Did not expect PRD event for non-status change") + } + case <-time.After(500 * time.Millisecond): + // Expected - no event for non-status changes + } +} + func TestWatcherStop(t *testing.T) { tmpDir := t.TempDir() - prdPath := createTestPRDMd(t, tmpDir, []UserStory{ - {ID: "US-001", Title: "Test Story", Passes: false}, - }) + prdPath := filepath.Join(tmpDir, "prd.json") + + // Create a test PRD file + testPRD := &PRD{ + Project: "Test", + UserStories: []UserStory{ + {ID: "US-001", Title: "Test Story", Passes: false}, + }, + } + data, _ := json.Marshal(testPRD) + if err := os.WriteFile(prdPath, data, 0644); err != nil { + t.Fatalf("Failed to write test PRD: %v", err) + } watcher, err := NewWatcher(prdPath) if err != nil { @@ -189,8 +276,11 @@ func TestWatcherStop(t *testing.T) { t.Fatalf("Failed to start watcher: %v", err) } + // Stop should not panic or hang + watcher.Stop() + + // Stopping again should be safe watcher.Stop() - watcher.Stop() // Should be safe } func TestHasStatusChanged(t *testing.T) { diff --git a/internal/tui/app.go b/internal/tui/app.go index 37694fea..10371cf7 100644 --- a/internal/tui/app.go +++ b/internal/tui/app.go @@ -10,6 +10,7 @@ import ( tea "github.com/charmbracelet/bubbletea" "github.com/minicodemonkey/chief/internal/config" + "github.com/minicodemonkey/chief/internal/engine" "github.com/minicodemonkey/chief/internal/git" "github.com/minicodemonkey/chief/internal/loop" "github.com/minicodemonkey/chief/internal/prd" @@ -84,10 +85,10 @@ type mergeResultMsg struct { // cleanResultMsg is sent when a clean operation completes. type cleanResultMsg struct { - prdName string - success bool - message string - clearBranch bool + prdName string + success bool + message string + clearBranch bool } // autoActionResultMsg is sent when a post-completion auto-action (push/PR) completes. @@ -151,22 +152,24 @@ const ( // App is the main Bubble Tea model for the Chief TUI. type App struct { - prd *prd.PRD - prdPath string - prdName string - state AppState - iteration int - startTime time.Time - selectedIndex int - storiesScrollOffset int - width int - height int - err error - - // Loop manager for parallel PRD execution - manager *loop.Manager - provider loop.Provider - maxIter int + prd *prd.PRD + prdPath string + prdName string + state AppState + iteration int + startTime time.Time + selectedIndex int + width int + height int + err error + + // Shared engine for parallel PRD execution (used by both TUI and serve) + eng *engine.Engine + maxIter int + + // Event subscription from engine + eventCh <-chan engine.ManagerEvent + unsubFn func() // Activity tracking lastActivity string @@ -198,8 +201,8 @@ type App struct { previousViewMode ViewMode // View to return to when closing help // Branch warning dialog - branchWarning *BranchWarning - pendingStartPRD string // PRD name waiting to start after branch decision + branchWarning *BranchWarning + pendingStartPRD string // PRD name waiting to start after branch decision pendingWorktreePath string // Absolute worktree path for pending PRD // Worktree setup spinner @@ -209,8 +212,8 @@ type App struct { completionScreen *CompletionScreen // Story timing tracking - storyTimings []StoryTiming - currentStoryID string + storyTimings []StoryTiming + currentStoryID string currentStoryStart time.Time // Settings overlay @@ -240,13 +243,65 @@ const ( ) // NewApp creates a new App with the given PRD. -func NewApp(prdPath string, provider loop.Provider) (*App, error) { - return NewAppWithOptions(prdPath, 10, provider) +func NewApp(prdPath string) (*App, error) { + return NewAppWithOptions(prdPath, 10) // default max iterations +} + +// NewAppWithEngine creates a new App using a pre-existing engine. +// This is used by the serve command to share an engine between the TUI and WebSocket handler. +func NewAppWithEngine(prdPath string, eng *engine.Engine) (*App, error) { + p, err := prd.LoadPRD(prdPath) + if err != nil { + return nil, err + } + + prdName := filepath.Base(filepath.Dir(prdPath)) + if prdName == "." || prdName == "/" { + prdName = filepath.Base(prdPath) + } + + watcher, err := prd.NewWatcher(prdPath) + if err != nil { + return nil, err + } + + baseDir := filepath.Dir(filepath.Dir(filepath.Dir(filepath.Dir(prdPath)))) + if !strings.Contains(prdPath, ".chief/prds/") { + baseDir, _ = os.Getwd() + } + + // Subscribe to engine events for TUI consumption + eventCh, unsubFn := eng.Subscribe() + + return &App{ + prd: p, + prdPath: prdPath, + prdName: prdName, + state: StateReady, + selectedIndex: 0, + maxIter: eng.MaxIterations(), + eng: eng, + eventCh: eventCh, + unsubFn: unsubFn, + watcher: watcher, + viewMode: ViewDashboard, + logViewer: NewLogViewer(), + diffViewer: NewDiffViewer(baseDir), + tabBar: NewTabBar(baseDir, prdName, eng.Manager()), + picker: NewPRDPicker(baseDir, prdName, eng.Manager()), + baseDir: baseDir, + config: eng.Config(), + helpOverlay: NewHelpOverlay(), + branchWarning: NewBranchWarning(), + worktreeSpinner: NewWorktreeSpinner(), + completionScreen: NewCompletionScreen(), + settingsOverlay: NewSettingsOverlay(), + }, nil } // NewAppWithOptions creates a new App with the given PRD and options. // If maxIter <= 0, it will be calculated dynamically based on remaining stories. -func NewAppWithOptions(prdPath string, maxIter int, provider loop.Provider) (*App, error) { +func NewAppWithOptions(prdPath string, maxIter int) (*App, error) { p, err := prd.LoadPRD(prdPath) if err != nil { return nil, err @@ -302,54 +357,58 @@ func NewAppWithOptions(prdPath string, maxIter int, provider loop.Provider) (*Ap progressWatcher, _ := prd.NewProgressWatcher(prdPath) progress, _ := prd.ParseProgress(prd.ProgressPath(prdPath)) - // Create loop manager for parallel PRD execution - manager := loop.NewManager(maxIter, provider) - manager.SetBaseDir(baseDir) - manager.SetConfig(cfg) + // Create shared engine for parallel PRD execution + eng := engine.New(maxIter) + eng.Manager().SetBaseDir(baseDir) + eng.SetConfig(cfg) + + // Register the initial PRD with the engine + eng.Register(prdName, prdPath) - // Register the initial PRD with the manager - manager.Register(prdName, prdPath) + // Subscribe to engine events for TUI consumption + eventCh, unsubFn := eng.Subscribe() // Create tab bar for always-visible PRD tabs - tabBar := NewTabBar(baseDir, prdName, manager) + tabBar := NewTabBar(baseDir, prdName, eng.Manager()) // Create picker with manager reference (for creating new PRDs) - picker := NewPRDPicker(baseDir, prdName, manager) + picker := NewPRDPicker(baseDir, prdName, eng.Manager()) return &App{ - prd: p, - prdPath: prdPath, - prdName: prdName, - state: StateReady, - iteration: 0, - selectedIndex: 0, - maxIter: maxIter, - manager: manager, - provider: provider, + prd: p, + prdPath: prdPath, + prdName: prdName, + state: StateReady, + iteration: 0, + selectedIndex: 0, + maxIter: maxIter, + eng: eng, + eventCh: eventCh, + unsubFn: unsubFn, watcher: watcher, progressWatcher: progressWatcher, progress: progress, viewMode: ViewDashboard, - logViewer: NewLogViewer(), - diffViewer: NewDiffViewer(baseDir), - tabBar: tabBar, - picker: picker, - baseDir: baseDir, - config: cfg, + logViewer: NewLogViewer(), + diffViewer: NewDiffViewer(baseDir), + tabBar: tabBar, + picker: picker, + baseDir: baseDir, + config: cfg, helpOverlay: NewHelpOverlay(), branchWarning: NewBranchWarning(), worktreeSpinner: NewWorktreeSpinner(), completionScreen: NewCompletionScreen(), settingsOverlay: NewSettingsOverlay(), - quitConfirm: NewQuitConfirmation(), + quitConfirm: NewQuitConfirmation(), }, nil } // SetCompletionCallback sets a callback that is called when any PRD completes. func (a *App) SetCompletionCallback(fn func(prdName string)) { a.onCompletion = fn - if a.manager != nil { - a.manager.SetCompletionCallback(fn) + if a.eng != nil { + a.eng.SetCompletionCallback(fn) } } @@ -360,8 +419,8 @@ func (a *App) SetVerbose(v bool) { // DisableRetry disables automatic retry on Claude crashes. func (a *App) DisableRetry() { - if a.manager != nil { - a.manager.DisableRetry() + if a.eng != nil { + a.eng.DisableRetry() } } @@ -388,13 +447,13 @@ func (a App) Init() tea.Cmd { ) } -// listenForManagerEvents listens for events from all managed loops. +// listenForManagerEvents listens for events from the engine's subscription. func (a *App) listenForManagerEvents() tea.Cmd { - if a.manager == nil { + if a.eventCh == nil { return nil } return func() tea.Msg { - event, ok := <-a.manager.Events() + event, ok := <-a.eventCh if !ok { return nil } @@ -583,7 +642,7 @@ func (a App) Update(msg tea.Msg) (tea.Model, tea.Cmd) { if a.viewMode == ViewDashboard || a.viewMode == ViewLog { // Use the current PRD's worktree directory if available, otherwise base dir diffDir := a.baseDir - if instance := a.manager.GetInstance(a.prdName); instance != nil && instance.WorktreeDir != "" { + if instance := a.eng.Manager().GetInstance(a.prdName); instance != nil && instance.WorktreeDir != "" { diffDir = instance.WorktreeDir } a.diffViewer.SetBaseDir(diffDir) @@ -663,9 +722,6 @@ func (a App) Update(msg tea.Msg) (tea.Model, tea.Cmd) { } else { if a.selectedIndex > 0 { a.selectedIndex-- - if a.selectedIndex < a.storiesScrollOffset { - a.storiesScrollOffset = a.selectedIndex - } } } case "down", "j": @@ -676,7 +732,6 @@ func (a App) Update(msg tea.Msg) (tea.Model, tea.Cmd) { } else { if a.selectedIndex < len(a.prd.UserStories)-1 { a.selectedIndex++ - a.adjustStoriesScroll() } } @@ -768,10 +823,10 @@ func (a App) startLoopForPRD(prdName string) (tea.Model, tea.Cmd) { // isAnotherPRDRunningInSameDir checks if another PRD is running in the project root (no worktree). func (a *App) isAnotherPRDRunningInSameDir(prdName string) bool { - if a.manager == nil { + if a.eng == nil { return false } - for _, inst := range a.manager.GetAllInstances() { + for _, inst := range a.eng.GetAllInstances() { if inst.Name != prdName && inst.State == loop.LoopStateRunning && inst.WorktreeDir == "" { return true } @@ -782,14 +837,14 @@ func (a *App) isAnotherPRDRunningInSameDir(prdName string) bool { // doStartLoop actually starts the loop (after branch check). func (a App) doStartLoop(prdName, prdDir string) (tea.Model, tea.Cmd) { // Check if this PRD is registered, if not register it - if instance := a.manager.GetInstance(prdName); instance == nil { + if instance := a.eng.GetInstance(prdName); instance == nil { // Find the PRD path prdPath := filepath.Join(prdDir, "prd.json") - a.manager.Register(prdName, prdPath) + a.eng.Register(prdName, prdPath) } // Start the loop via manager - if err := a.manager.Start(prdName); err != nil { + if err := a.eng.Start(prdName); err != nil { a.lastActivity = "Error starting loop: " + err.Error() return a, nil } @@ -817,8 +872,8 @@ func (a App) pauseLoop() (tea.Model, tea.Cmd) { // pauseLoopForPRD pauses the loop for a specific PRD. func (a App) pauseLoopForPRD(prdName string) (tea.Model, tea.Cmd) { - if a.manager != nil { - a.manager.Pause(prdName) + if a.eng != nil { + a.eng.Pause(prdName) } if prdName == a.prdName { a.lastActivity = "Pausing after current iteration..." @@ -835,8 +890,8 @@ func (a *App) stopLoop() { // stopLoopForPRD stops the loop for a specific PRD immediately. func (a *App) stopLoopForPRD(prdName string) { - if a.manager != nil { - a.manager.Stop(prdName) + if a.eng != nil { + a.eng.Stop(prdName) } } @@ -857,17 +912,20 @@ func (a App) stopLoopAndUpdateForPRD(prdName string) (tea.Model, tea.Cmd) { return a, nil } -// stopAllLoops stops all running loops. +// stopAllLoops stops all running loops and unsubscribes from events. func (a *App) stopAllLoops() { - if a.manager != nil { - a.manager.StopAll() + if a.eng != nil { + a.eng.StopAll() + } + if a.unsubFn != nil { + a.unsubFn() } } // tryQuit attempts to quit the app. If any loop is running, it shows the quit // confirmation dialog instead of quitting immediately. func (a App) tryQuit() (tea.Model, tea.Cmd) { - if a.manager != nil && a.manager.IsAnyRunning() { + if a.eng.Manager() != nil && a.eng.Manager().IsAnyRunning() { a.previousViewMode = a.viewMode a.viewMode = ViewQuitConfirm a.quitConfirm.Reset() @@ -927,12 +985,6 @@ func (a App) handleLoopEvent(prdName string, event loop.Event) (tea.Model, tea.C case loop.EventIterationStart: if isCurrentPRD { a.lastActivity = "Starting iteration..." - // Start tracking story timing if this is a new story - if event.StoryID != "" && event.StoryID != a.currentStoryID { - a.finalizeStoryTiming() - a.currentStoryID = event.StoryID - a.currentStoryStart = time.Now() - } } case loop.EventAssistantText: if isCurrentPRD { @@ -951,11 +1003,14 @@ func (a App) handleLoopEvent(prdName string, event loop.Event) (tea.Model, tea.C if isCurrentPRD { a.lastActivity = "Tool completed" } - case loop.EventStoryDone: + case loop.EventStoryStarted: if isCurrentPRD { - a.lastActivity = "Story done" - // Finalize story timing + a.lastActivity = "Working on: " + event.StoryID + // Finalize previous story timing a.finalizeStoryTiming() + // Start tracking the new story + a.currentStoryID = event.StoryID + a.currentStoryStart = time.Now() } case loop.EventComplete: if isCurrentPRD { @@ -989,21 +1044,23 @@ func (a App) handleLoopEvent(prdName string, event loop.Event) (tea.Model, tea.C if isCurrentPRD { a.lastActivity = event.Text } - case loop.EventWatchdogTimeout: - if isCurrentPRD { - a.lastActivity = event.Text - } } // Reload PRD from disk only on meaningful state changes (not every event) if isCurrentPRD { switch event.Type { - case loop.EventStoryDone, loop.EventComplete, loop.EventError, loop.EventMaxIterationsReached: + case loop.EventStoryStarted, loop.EventComplete, loop.EventError, loop.EventMaxIterationsReached: if p, err := prd.LoadPRD(a.prdPath); err == nil { a.prd = p } } + // Mark the story as in-progress in the PRD and auto-select it + if event.Type == loop.EventStoryStarted && event.StoryID != "" { + a.markStoryInProgress(event.StoryID) + a.selectStoryByID(event.StoryID) + } + // Clear in-progress when the PRD completes or the loop stops if event.Type == loop.EventComplete || event.Type == loop.EventError || event.Type == loop.EventMaxIterationsReached { a.clearInProgress() @@ -1027,7 +1084,7 @@ func (a App) handleLoopFinished(prdName string, err error) (tea.Model, tea.Cmd) // Only update state if this is the current PRD if prdName == a.prdName { // Get the actual state from the manager - if state, _, _ := a.manager.GetState(prdName); state != 0 { + if state, _, _ := a.eng.GetState(prdName); state != 0 { switch state { case loop.LoopStateError: a.state = StateError @@ -1177,8 +1234,8 @@ func (a App) handleBranchWarningKeys(msg tea.KeyMsg) (tea.Model, tea.Cmd) { return a, nil } // Track the branch on the manager instance - if instance := a.manager.GetInstance(prdName); instance != nil { - a.manager.UpdateWorktreeInfo(prdName, "", branchName) + if instance := a.eng.GetInstance(prdName); instance != nil { + a.eng.UpdateWorktreeInfo(prdName, "", branchName) } a.lastActivity = "Created branch: " + branchName // Now start the loop @@ -1267,7 +1324,7 @@ func (a *App) showCompletionScreen(prdName string) tea.Cmd { // Get branch from manager branch := "" - if instance := a.manager.GetInstance(prdName); instance != nil { + if instance := a.eng.GetInstance(prdName); instance != nil { branch = instance.Branch } @@ -1312,7 +1369,7 @@ func (a *App) runBackgroundAutoActions(prdName string) tea.Cmd { return nil } - instance := a.manager.GetInstance(prdName) + instance := a.eng.GetInstance(prdName) if instance == nil || instance.Branch == "" { return nil } @@ -1371,7 +1428,7 @@ func (a App) handleBackgroundAutoAction(msg backgroundAutoActionResultMsg) (tea. if msg.action == "push" && a.config != nil && a.config.OnComplete.CreatePR { // Chain PR creation after successful push - instance := a.manager.GetInstance(msg.prdName) + instance := a.eng.GetInstance(msg.prdName) if instance != nil && instance.Branch != "" { prdName := msg.prdName branch := instance.Branch @@ -1398,7 +1455,7 @@ func (a *App) runAutoPush() tea.Cmd { branch := a.completionScreen.Branch() // Use worktree dir if available, otherwise base dir dir := a.baseDir - if instance := a.manager.GetInstance(a.completionScreen.PRDName()); instance != nil && instance.WorktreeDir != "" { + if instance := a.eng.GetInstance(a.completionScreen.PRDName()); instance != nil && instance.WorktreeDir != "" { dir = instance.WorktreeDir } return func() tea.Msg { @@ -1696,10 +1753,10 @@ func (a App) finishWorktreeSetup() (tea.Model, tea.Cmd) { // Register or update with worktree info prdPath := filepath.Join(prdDir, "prd.json") - if instance := a.manager.GetInstance(prdName); instance == nil { - a.manager.RegisterWithWorktree(prdName, prdPath, worktreePath, branchName) + if instance := a.eng.GetInstance(prdName); instance == nil { + a.eng.RegisterWithWorktree(prdName, prdPath, worktreePath, branchName) } else { - a.manager.UpdateWorktreeInfo(prdName, worktreePath, branchName) + a.eng.UpdateWorktreeInfo(prdName, worktreePath, branchName) } a.lastActivity = fmt.Sprintf("Created worktree at %s on branch %s", worktreePath, branchName) @@ -1814,8 +1871,8 @@ func (a App) handleCleanResult(msg cleanResultMsg) (tea.Model, tea.Cmd) { if msg.success { // Clear worktree info from manager - if a.manager != nil { - a.manager.ClearWorktreeInfo(msg.prdName, msg.clearBranch) + if a.eng != nil { + a.eng.ClearWorktreeInfo(msg.prdName, msg.clearBranch) } a.picker.Refresh() a.lastActivity = fmt.Sprintf("Cleaned worktree for %s", msg.prdName) @@ -2004,8 +2061,8 @@ func (a App) switchToPRD(name, prdPath string) (tea.Model, tea.Cmd) { } // Register with manager if not already registered - if instance := a.manager.GetInstance(name); instance == nil { - a.manager.Register(name, prdPath) + if instance := a.eng.GetInstance(name); instance == nil { + a.eng.Register(name, prdPath) } // Create new watcher for the new PRD @@ -2028,7 +2085,7 @@ func (a App) switchToPRD(name, prdPath string) (tea.Model, tea.Cmd) { a.progress, _ = prd.ParseProgress(prd.ProgressPath(prdPath)) // Get the state from the manager for this PRD - loopState, iteration, loopErr := a.manager.GetState(name) + loopState, iteration, loopErr := a.eng.GetState(name) appState := StateReady switch loopState { case loop.LoopStateRunning: @@ -2044,7 +2101,7 @@ func (a App) switchToPRD(name, prdPath string) (tea.Model, tea.Cmd) { } // Only recalculate max iterations if no loop is currently running for this PRD - if instance := a.manager.GetInstance(name); instance == nil || instance.State != loop.LoopStateRunning { + if instance := a.eng.Manager().GetInstance(name); instance == nil || instance.State != loop.LoopStateRunning { remaining := 0 for _, story := range newPRD.UserStories { if !story.Passes { @@ -2062,13 +2119,12 @@ func (a App) switchToPRD(name, prdPath string) (tea.Model, tea.Cmd) { a.prdPath = prdPath a.prdName = name a.selectedIndex = 0 - a.storiesScrollOffset = 0 a.state = appState a.iteration = iteration a.err = loopErr if appState == StateRunning { // Keep the existing start time if running - if instance := a.manager.GetInstance(name); instance != nil { + if instance := a.eng.GetInstance(name); instance != nil { a.startTime = instance.StartTime } } else { @@ -2122,69 +2178,26 @@ func (a *App) GetSelectedStory() *prd.UserStory { return nil } -// storiesListHeight calculates how many story lines fit in the panel. -// Must match the calculation in renderStoriesPanel. -func (a *App) storiesListHeight() int { - fh := footerHeight - if a.height < 12 { - fh = 0 - } - contentHeight := a.height - a.effectiveHeaderHeight() - fh - 2 - if a.isNarrowMode() { - storiesHeight := max((contentHeight*40)/100, 5) - return storiesHeight - 5 - } - return contentHeight - 5 -} - -// adjustStoriesScroll ensures the selected index is visible in the scroll window. -func (a *App) adjustStoriesScroll() { - listHeight := a.storiesListHeight() - if listHeight <= 0 { - return - } - if a.selectedIndex < a.storiesScrollOffset { - a.storiesScrollOffset = a.selectedIndex - } - if a.selectedIndex >= a.storiesScrollOffset+listHeight { - a.storiesScrollOffset = a.selectedIndex - listHeight + 1 - } - // Clamp - maxOffset := len(a.prd.UserStories) - listHeight - if maxOffset < 0 { - maxOffset = 0 - } - if a.storiesScrollOffset > maxOffset { - a.storiesScrollOffset = maxOffset - } - if a.storiesScrollOffset < 0 { - a.storiesScrollOffset = 0 - } -} - // markStoryInProgress clears any existing in-progress flags and marks the -// given story as in-progress, then reloads the PRD from disk. +// given story as in-progress, then saves the PRD to disk. func (a *App) markStoryInProgress(storyID string) { - _ = prd.SetStoryStatus(a.prdPath, storyID, "in-progress") - if p, err := prd.LoadPRD(a.prdPath); err == nil { - a.prd = p + for i := range a.prd.UserStories { + a.prd.UserStories[i].InProgress = a.prd.UserStories[i].ID == storyID } + _ = a.prd.Save(a.prdPath) } -// clearInProgress clears all in-progress flags by setting each in-progress -// story's status to "todo" in the markdown file, then reloads. +// clearInProgress clears all in-progress flags and saves the PRD to disk. func (a *App) clearInProgress() { dirty := false - for _, story := range a.prd.UserStories { - if story.InProgress { - _ = prd.SetStoryStatus(a.prdPath, story.ID, "todo") + for i := range a.prd.UserStories { + if a.prd.UserStories[i].InProgress { + a.prd.UserStories[i].InProgress = false dirty = true } } if dirty { - if p, err := prd.LoadPRD(a.prdPath); err == nil { - a.prd = p - } + _ = a.prd.Save(a.prdPath) } } @@ -2193,7 +2206,6 @@ func (a *App) selectStoryByID(storyID string) { for i, story := range a.prd.UserStories { if story.ID == storyID { a.selectedIndex = i - a.adjustStoriesScroll() return } } @@ -2204,7 +2216,6 @@ func (a *App) selectInProgressStory() { for i, story := range a.prd.UserStories { if story.InProgress { a.selectedIndex = i - a.adjustStoriesScroll() return } } @@ -2256,10 +2267,10 @@ func (a *App) adjustMaxIterations(delta int) { a.maxIter = newMax // Update the manager's default - if a.manager != nil { - a.manager.SetMaxIterations(newMax) + if a.eng != nil { + a.eng.SetMaxIterations(newMax) // Also update any running loop for the current PRD - a.manager.SetMaxIterationsForInstance(a.prdName, newMax) + a.eng.SetMaxIterationsForInstance(a.prdName, newMax) } a.lastActivity = fmt.Sprintf("Max iterations: %d", newMax) @@ -2312,7 +2323,6 @@ func (a App) handlePRDUpdate(msg PRDUpdateMsg) (tea.Model, tea.Cmd) { // Auto-select the in-progress story so the user sees its details a.selectInProgressStory() - a.adjustStoriesScroll() } // Continue listening for changes diff --git a/internal/tui/branch_warning.go b/internal/tui/branch_warning.go index 106ca0cc..8f431b54 100644 --- a/internal/tui/branch_warning.go +++ b/internal/tui/branch_warning.go @@ -11,10 +11,10 @@ import ( type BranchWarningOption int const ( - BranchOptionCreateWorktree BranchWarningOption = iota // Create worktree + branch - BranchOptionCreateBranch // Create branch only (no worktree) - BranchOptionContinue // Continue on current branch / run in same directory - BranchOptionCancel // Cancel + BranchOptionCreateWorktree BranchWarningOption = iota // Create worktree + branch + BranchOptionCreateBranch // Create branch only (no worktree) + BranchOptionContinue // Continue on current branch / run in same directory + BranchOptionCancel // Cancel ) // DialogContext determines which set of options to show. diff --git a/internal/tui/completion.go b/internal/tui/completion.go index 23d738ed..4d14b752 100644 --- a/internal/tui/completion.go +++ b/internal/tui/completion.go @@ -30,11 +30,11 @@ type CompletionScreen struct { width int height int - prdName string - completed int - total int - branch string - commitCount int + prdName string + completed int + total int + branch string + commitCount int hasAutoActions bool // Whether push/PR auto-actions are configured // Duration data diff --git a/internal/tui/confetti.go b/internal/tui/confetti.go index 4f5c4683..2da38bee 100644 --- a/internal/tui/confetti.go +++ b/internal/tui/confetti.go @@ -55,13 +55,13 @@ func NewConfetti(width, height int) *Confetti { for i := range c.particles { c.particles[i] = Particle{ - x: rand.Float64() * float64(width), - y: rand.Float64()*float64(height+10) - float64(height/2), // stagger: some above screen, some mid - vx: (rand.Float64() - 0.5) * 0.6, // lateral drift -0.3 to 0.3 - vy: 0.2 + rand.Float64()*0.4, // falling 0.2-0.6 - char: confettiChars[rand.Intn(len(confettiChars))], + x: rand.Float64() * float64(width), + y: rand.Float64()*float64(height+10) - float64(height/2), // stagger: some above screen, some mid + vx: (rand.Float64() - 0.5) * 0.6, // lateral drift -0.3 to 0.3 + vy: 0.2 + rand.Float64()*0.4, // falling 0.2-0.6 + char: confettiChars[rand.Intn(len(confettiChars))], color: confettiColors[rand.Intn(len(confettiColors))], - life: 80 + rand.Intn(120), // 80-200 ticks + life: 80 + rand.Intn(120), // 80-200 ticks } } diff --git a/internal/tui/dashboard.go b/internal/tui/dashboard.go index 99dbeaf7..56416d45 100644 --- a/internal/tui/dashboard.go +++ b/internal/tui/dashboard.go @@ -38,19 +38,10 @@ func (a *App) renderDashboard() string { } header := a.renderHeader() - - // Hide footer when terminal height < 12 - fh := footerHeight - var footer string - if a.height < 12 { - fh = 0 - footer = "" - } else { - footer = a.renderFooter() - } + footer := a.renderFooter() // Calculate content area height - contentHeight := a.height - a.effectiveHeaderHeight() - fh - 2 // -2 for panel borders + contentHeight := a.height - a.effectiveHeaderHeight() - footerHeight - 2 // -2 for panel borders // Render panels storiesWidth := (a.width * storiesPanelPct / 100) - 2 @@ -63,28 +54,16 @@ func (a *App) renderDashboard() string { content := lipgloss.JoinHorizontal(lipgloss.Top, storiesPanel, detailsPanel) // Stack header, content, and footer - if footer == "" { - return lipgloss.JoinVertical(lipgloss.Left, header, content) - } return lipgloss.JoinVertical(lipgloss.Left, header, content, footer) } // renderStackedDashboard renders the dashboard with stacked layout for narrow terminals. func (a *App) renderStackedDashboard() string { header := a.renderNarrowHeader() - - // Hide footer when terminal height < 12 - fh := footerHeight - var footer string - if a.height < 12 { - fh = 0 - footer = "" - } else { - footer = a.renderNarrowFooter() - } + footer := a.renderNarrowFooter() // Calculate content area height - contentHeight := a.height - a.effectiveHeaderHeight() - fh - 2 // -2 for panel borders + contentHeight := a.height - a.effectiveHeaderHeight() - footerHeight - 2 // -2 for panel borders // Split height between stories (40%) and details (60%) storiesHeight := max((contentHeight*40)/100, 5) @@ -99,19 +78,16 @@ func (a *App) renderStackedDashboard() string { content := lipgloss.JoinVertical(lipgloss.Left, storiesPanel, detailsPanel) // Stack header, content, and footer - if footer == "" { - return lipgloss.JoinVertical(lipgloss.Left, header, content) - } return lipgloss.JoinVertical(lipgloss.Left, header, content, footer) } // getWorktreeInfo returns the branch and directory info for the current PRD. // Returns empty strings if no branch is set (backward compatible). func (a *App) getWorktreeInfo() (branch, dir string) { - if a.manager == nil { + if a.eng == nil { return "", "" } - instance := a.manager.GetInstance(a.prdName) + instance := a.eng.GetInstance(a.prdName) if instance == nil || instance.Branch == "" { return "", "" } @@ -369,49 +345,29 @@ func (a *App) renderActivityLine() string { func (a *App) renderStoriesPanel(width, height int) string { var content strings.Builder - // Panel title — append scroll percentage when list is scrollable - listHeight := height - 5 // Account for title, border, and progress bar - totalStories := len(a.prd.UserStories) - titleText := "Stories" - if totalStories > listHeight && listHeight > 0 { - maxOffset := totalStories - listHeight - pct := 0 - if maxOffset > 0 { - pct = a.storiesScrollOffset * 100 / maxOffset - } - titleText = fmt.Sprintf("Stories (%d%%)", pct) - } - title := PanelTitleStyle.Render(titleText) + // Panel title using centralized style + title := PanelTitleStyle.Render("Stories") content.WriteString(title) content.WriteString("\n") content.WriteString(DividerStyle.Render(strings.Repeat("─", width-2))) content.WriteString("\n") - // Clamp scroll offset - if a.storiesScrollOffset < 0 { - a.storiesScrollOffset = 0 - } - if listHeight > 0 && a.storiesScrollOffset > totalStories-listHeight { - a.storiesScrollOffset = totalStories - listHeight - } - if a.storiesScrollOffset < 0 { - a.storiesScrollOffset = 0 - } + // Story list + listHeight := height - 5 // Account for title, border, and progress bar + for i, story := range a.prd.UserStories { + if i >= listHeight { + // Show indicator that there are more stories + moreStyle := lipgloss.NewStyle().Foreground(mutedColor) + content.WriteString(moreStyle.Render(fmt.Sprintf("... and %d more", len(a.prd.UserStories)-i))) + break + } - // Render visible slice of stories - endIdx := a.storiesScrollOffset + listHeight - if endIdx > totalStories { - endIdx = totalStories - } - visibleCount := 0 - for i := a.storiesScrollOffset; i < endIdx; i++ { - story := a.prd.UserStories[i] icon := GetStatusIcon(story.Passes, story.InProgress) // Truncate title to fit maxTitleLen := width - 12 // Account for icon, ID, and spacing displayTitle := story.Title - if len(displayTitle) > maxTitleLen && maxTitleLen > 3 { + if len(displayTitle) > maxTitleLen { displayTitle = displayTitle[:maxTitleLen-3] + "..." } @@ -429,11 +385,10 @@ func (a *App) renderStoriesPanel(width, height int) string { content.WriteString(line) content.WriteString("\n") - visibleCount++ } // Pad remaining space - linesWritten := visibleCount + 2 // +2 for title and divider + linesWritten := min(len(a.prd.UserStories), listHeight) + 2 // +2 for title and divider for i := linesWritten; i < height-3; i++ { content.WriteString("\n") } @@ -490,7 +445,7 @@ func (a *App) renderDetailsPanel(width, height int) string { statusText = "Pending" statusStyle = statusPendingStyle } - content.WriteString(fmt.Sprintf("%s %s │ Priority: %g\n", statusIcon, statusStyle.Render(statusText), story.Priority)) + content.WriteString(fmt.Sprintf("%s %s │ Priority: %d\n", statusIcon, statusStyle.Render(statusText), story.Priority)) content.WriteString(DividerStyle.Render(strings.Repeat("─", width-4))) content.WriteString("\n\n") @@ -523,15 +478,7 @@ func (a *App) renderDetailsPanel(width, height int) string { } } - // Truncate content to fit panel height (lipgloss Height only sets minimum, not maximum) - contentStr := content.String() - contentLines := strings.Split(contentStr, "\n") - if len(contentLines) > height { - contentLines = contentLines[:height] - contentStr = strings.Join(contentLines, "\n") - } - - return panelStyle.Width(width).Height(height).Render(contentStr) + return panelStyle.Width(width).Height(height).Render(content.String()) } // renderErrorPanel renders the error details panel when in error state. @@ -560,11 +507,7 @@ func (a *App) renderErrorPanel(width, height int) string { content.WriteString(DividerStyle.Render(strings.Repeat("─", width-4))) content.WriteString("\n\n") hintStyle := lipgloss.NewStyle().Foreground(WarningColor) - logName := "claude.log" - if a.provider != nil { - logName = a.provider.LogFileName() - } - content.WriteString(hintStyle.Render(fmt.Sprintf("💡 Tip: Check %s in the PRD directory for full error details.", logName))) + content.WriteString(hintStyle.Render("💡 Tip: Check claude.log in the PRD directory for full error details.")) content.WriteString("\n\n") // Retry instructions diff --git a/internal/tui/dashboard_test.go b/internal/tui/dashboard_test.go deleted file mode 100644 index 37583dfa..00000000 --- a/internal/tui/dashboard_test.go +++ /dev/null @@ -1,196 +0,0 @@ -package tui - -import ( - "fmt" - "strings" - "testing" - - "github.com/minicodemonkey/chief/internal/prd" -) - -// newTestApp creates a minimal App for testing scroll and rendering. -func newTestApp(stories []prd.UserStory, width, height int) *App { - return &App{ - prd: &prd.PRD{UserStories: stories}, - width: width, - height: height, - viewMode: ViewDashboard, - } -} - -func makeStories(n int) []prd.UserStory { - stories := make([]prd.UserStory, n) - for i := range stories { - stories[i] = prd.UserStory{ - ID: fmt.Sprintf("US-%03d", i+1), - Title: fmt.Sprintf("Story %d", i+1), - Priority: float64(i + 1), - } - } - return stories -} - -func TestScrollOffset_FollowsCursorDown(t *testing.T) { - app := newTestApp(makeStories(20), 120, 20) - listHeight := app.storiesListHeight() - if listHeight <= 0 { - t.Fatalf("expected positive listHeight, got %d", listHeight) - } - - // Navigate down past the visible range - for i := 0; i < listHeight+3; i++ { - if app.selectedIndex < len(app.prd.UserStories)-1 { - app.selectedIndex++ - app.adjustStoriesScroll() - } - } - - // Selected index should be past the first screen - if app.selectedIndex <= listHeight { - t.Errorf("expected selectedIndex > %d, got %d", listHeight, app.selectedIndex) - } - - // Scroll offset should have followed - if app.storiesScrollOffset == 0 { - t.Error("expected storiesScrollOffset > 0 after scrolling down past visible range") - } - - // Selected index should be visible - if app.selectedIndex < app.storiesScrollOffset || app.selectedIndex >= app.storiesScrollOffset+listHeight { - t.Errorf("selectedIndex %d not visible in scroll window [%d, %d)", app.selectedIndex, app.storiesScrollOffset, app.storiesScrollOffset+listHeight) - } -} - -func TestScrollOffset_FollowsCursorUp(t *testing.T) { - app := newTestApp(makeStories(20), 120, 20) - listHeight := app.storiesListHeight() - - // Move down first - for i := 0; i < listHeight+5; i++ { - if app.selectedIndex < len(app.prd.UserStories)-1 { - app.selectedIndex++ - app.adjustStoriesScroll() - } - } - savedOffset := app.storiesScrollOffset - - // Now navigate back up past the scroll offset - for i := 0; i < listHeight+5; i++ { - if app.selectedIndex > 0 { - app.selectedIndex-- - if app.selectedIndex < app.storiesScrollOffset { - app.storiesScrollOffset = app.selectedIndex - } - } - } - - // Should be back at top - if app.selectedIndex != 0 { - t.Errorf("expected selectedIndex 0, got %d", app.selectedIndex) - } - if app.storiesScrollOffset != 0 { - t.Errorf("expected storiesScrollOffset 0, got %d", app.storiesScrollOffset) - } - _ = savedOffset -} - -func TestScrollOffset_NoScrollWhenAllFit(t *testing.T) { - // 3 stories in a 20-tall terminal — all should fit - app := newTestApp(makeStories(3), 120, 20) - listHeight := app.storiesListHeight() - - if len(app.prd.UserStories) > listHeight { - t.Skipf("stories (%d) > listHeight (%d), skipping", len(app.prd.UserStories), listHeight) - } - - // Navigate through all stories - for i := 0; i < len(app.prd.UserStories); i++ { - app.selectedIndex = i - app.adjustStoriesScroll() - } - - if app.storiesScrollOffset != 0 { - t.Errorf("expected storiesScrollOffset 0 when all stories fit, got %d", app.storiesScrollOffset) - } -} - -func TestScrollOffset_ClampsToValidRange(t *testing.T) { - app := newTestApp(makeStories(20), 120, 20) - listHeight := app.storiesListHeight() - - // Force an invalid scroll offset - app.storiesScrollOffset = 100 - app.adjustStoriesScroll() - - maxOffset := len(app.prd.UserStories) - listHeight - if maxOffset < 0 { - maxOffset = 0 - } - if app.storiesScrollOffset > maxOffset { - t.Errorf("expected storiesScrollOffset <= %d, got %d", maxOffset, app.storiesScrollOffset) - } - - // Force negative - app.storiesScrollOffset = -5 - app.adjustStoriesScroll() - if app.storiesScrollOffset < 0 { - t.Errorf("expected storiesScrollOffset >= 0, got %d", app.storiesScrollOffset) - } -} - -func TestScrollPercentage_ShownWhenScrollable(t *testing.T) { - app := newTestApp(makeStories(20), 120, 20) - - // Render the panel - output := app.renderStoriesPanel(40, 15) - - // With 20 stories and listHeight = 15-5=10, list is scrollable - // Title should contain percentage - if !strings.Contains(output, "Stories (") || !strings.Contains(output, "%)") { - t.Errorf("expected scroll percentage in panel title, got: %s", output) - } -} - -func TestScrollPercentage_NotShownWhenNotScrollable(t *testing.T) { - app := newTestApp(makeStories(3), 120, 20) - - output := app.renderStoriesPanel(40, 15) - - // 3 stories fits in listHeight=10, so no percentage - if strings.Contains(output, "%)") { - t.Errorf("expected no scroll percentage when list fits, got: %s", output) - } -} - -func TestFooterHidden_WhenHeightLessThan12(t *testing.T) { - app := newTestApp(makeStories(5), 120, 11) - - output := app.renderDashboard() - - // The footer contains "quit" shortcut — should not be present - if strings.Contains(output, "q: quit") { - t.Error("expected footer to be hidden when height < 12") - } -} - -func TestFooterShown_WhenHeightAtLeast12(t *testing.T) { - // Need enough height to render without panic - app := newTestApp(makeStories(5), 120, 20) - - output := app.renderDashboard() - - if !strings.Contains(output, "q: quit") { - t.Error("expected footer to be shown when height >= 12") - } -} - -func TestAndNMore_Removed(t *testing.T) { - // Create more stories than can fit in the panel - app := newTestApp(makeStories(20), 120, 15) - - output := app.renderStoriesPanel(40, 12) - - if strings.Contains(output, "... and") || strings.Contains(output, "more") { - t.Error("expected '... and N more' to be removed from stories panel") - } -} diff --git a/internal/tui/diff.go b/internal/tui/diff.go index 0415f266..5549234f 100644 --- a/internal/tui/diff.go +++ b/internal/tui/diff.go @@ -9,16 +9,16 @@ import ( // DiffViewer displays git diffs with syntax highlighting and scrolling. type DiffViewer struct { - lines []string - offset int - width int - height int - stats string - baseDir string - storyID string // Story ID whose commit diff is being shown (empty = full branch diff) - noCommit bool // True when no commit was found for the selected story - err error - loaded bool + lines []string + offset int + width int + height int + stats string + baseDir string + storyID string // Story ID whose commit diff is being shown (empty = full branch diff) + noCommit bool // True when no commit was found for the selected story + err error + loaded bool } // NewDiffViewer creates a new diff viewer. @@ -236,3 +236,4 @@ func (d *DiffViewer) styleLine(line string) string { return line } } + diff --git a/internal/tui/layout_test.go b/internal/tui/layout_test.go index 8dfcd075..c38ad616 100644 --- a/internal/tui/layout_test.go +++ b/internal/tui/layout_test.go @@ -4,8 +4,7 @@ import ( "strings" "testing" - "github.com/minicodemonkey/chief/internal/agent" - "github.com/minicodemonkey/chief/internal/loop" + "github.com/minicodemonkey/chief/internal/engine" ) func TestIsNarrowMode(t *testing.T) { @@ -208,20 +207,36 @@ func TestMinMaxHelpers(t *testing.T) { } } +// newTestEngine creates a test engine with the given PRD registered. +func newTestEngine(name, prdPath string) *engine.Engine { + eng := engine.New(10) + if prdPath != "" { + eng.Register(name, prdPath) + } + return eng +} + +// newTestEngineWithWorktree creates a test engine with a worktree-registered PRD. +func newTestEngineWithWorktree(name, prdPath, worktreeDir, branch string) *engine.Engine { + eng := engine.New(10) + eng.RegisterWithWorktree(name, prdPath, worktreeDir, branch) + return eng +} + func TestGetWorktreeInfo_NoBranch(t *testing.T) { - // No manager - should return empty + // No engine - should return empty app := &App{prdName: "auth"} branch, dir := app.getWorktreeInfo() if branch != "" || dir != "" { - t.Errorf("expected empty worktree info without manager, got branch=%q dir=%q", branch, dir) + t.Errorf("expected empty worktree info without engine, got branch=%q dir=%q", branch, dir) } } func TestGetWorktreeInfo_WithBranch(t *testing.T) { - mgr := loop.NewManager(10, agent.NewClaudeProvider("")) - mgr.RegisterWithWorktree("auth", "/tmp/prd.json", "/tmp/.chief/worktrees/auth", "chief/auth") + eng := newTestEngineWithWorktree("auth", "/tmp/prd.json", "/tmp/.chief/worktrees/auth", "chief/auth") + defer eng.Shutdown() - app := &App{prdName: "auth", manager: mgr} + app := &App{prdName: "auth", eng: eng} branch, dir := app.getWorktreeInfo() if branch != "chief/auth" { t.Errorf("branch = %q, want %q", branch, "chief/auth") @@ -233,10 +248,10 @@ func TestGetWorktreeInfo_WithBranch(t *testing.T) { func TestGetWorktreeInfo_WithBranchNoWorktree(t *testing.T) { // Branch set but no worktree dir (branch-only mode) - mgr := loop.NewManager(10, agent.NewClaudeProvider("")) - mgr.RegisterWithWorktree("auth", "/tmp/prd.json", "", "chief/auth") + eng := newTestEngineWithWorktree("auth", "/tmp/prd.json", "", "chief/auth") + defer eng.Shutdown() - app := &App{prdName: "auth", manager: mgr} + app := &App{prdName: "auth", eng: eng} branch, dir := app.getWorktreeInfo() if branch != "chief/auth" { t.Errorf("branch = %q, want %q", branch, "chief/auth") @@ -248,10 +263,10 @@ func TestGetWorktreeInfo_WithBranchNoWorktree(t *testing.T) { func TestGetWorktreeInfo_RegisteredNoBranch(t *testing.T) { // Registered without worktree - should return empty (backward compatible) - mgr := loop.NewManager(10, agent.NewClaudeProvider("")) - mgr.Register("auth", "/tmp/prd.json") + eng := newTestEngine("auth", "/tmp/prd.json") + defer eng.Shutdown() - app := &App{prdName: "auth", manager: mgr} + app := &App{prdName: "auth", eng: eng} branch, dir := app.getWorktreeInfo() if branch != "" || dir != "" { t.Errorf("expected empty worktree info for no-branch PRD, got branch=%q dir=%q", branch, dir) @@ -259,16 +274,17 @@ func TestGetWorktreeInfo_RegisteredNoBranch(t *testing.T) { } func TestHasWorktreeInfo(t *testing.T) { - // No manager + // No engine app := &App{prdName: "auth"} if app.hasWorktreeInfo() { - t.Error("expected hasWorktreeInfo=false without manager") + t.Error("expected hasWorktreeInfo=false without engine") } // With branch - mgr := loop.NewManager(10, agent.NewClaudeProvider("")) - mgr.RegisterWithWorktree("auth", "/tmp/prd.json", "/tmp/.chief/worktrees/auth", "chief/auth") - app.manager = mgr + eng := newTestEngineWithWorktree("auth", "/tmp/prd.json", "/tmp/.chief/worktrees/auth", "chief/auth") + defer eng.Shutdown() + + app.eng = eng if !app.hasWorktreeInfo() { t.Error("expected hasWorktreeInfo=true with branch set") } @@ -282,10 +298,10 @@ func TestEffectiveHeaderHeight_NoBranch(t *testing.T) { } func TestEffectiveHeaderHeight_WithBranch(t *testing.T) { - mgr := loop.NewManager(10, agent.NewClaudeProvider("")) - mgr.RegisterWithWorktree("auth", "/tmp/prd.json", "/tmp/.chief/worktrees/auth", "chief/auth") + eng := newTestEngineWithWorktree("auth", "/tmp/prd.json", "/tmp/.chief/worktrees/auth", "chief/auth") + defer eng.Shutdown() - app := &App{prdName: "auth", manager: mgr} + app := &App{prdName: "auth", eng: eng} if got := app.effectiveHeaderHeight(); got != headerHeight+1 { t.Errorf("effectiveHeaderHeight() = %d, want %d (with branch)", got, headerHeight+1) } @@ -299,10 +315,10 @@ func TestRenderWorktreeInfoLine_NoBranch(t *testing.T) { } func TestRenderWorktreeInfoLine_WithBranch(t *testing.T) { - mgr := loop.NewManager(10, agent.NewClaudeProvider("")) - mgr.RegisterWithWorktree("auth", "/tmp/prd.json", "/tmp/.chief/worktrees/auth", "chief/auth") + eng := newTestEngineWithWorktree("auth", "/tmp/prd.json", "/tmp/.chief/worktrees/auth", "chief/auth") + defer eng.Shutdown() - app := &App{prdName: "auth", manager: mgr} + app := &App{prdName: "auth", eng: eng} got := app.renderWorktreeInfoLine() if got == "" { t.Error("renderWorktreeInfoLine() should not be empty with branch set") @@ -322,10 +338,10 @@ func TestRenderWorktreeInfoLine_WithBranch(t *testing.T) { } func TestRenderWorktreeInfoLine_BranchNoWorktree(t *testing.T) { - mgr := loop.NewManager(10, agent.NewClaudeProvider("")) - mgr.RegisterWithWorktree("auth", "/tmp/prd.json", "", "chief/auth") + eng := newTestEngineWithWorktree("auth", "/tmp/prd.json", "", "chief/auth") + defer eng.Shutdown() - app := &App{prdName: "auth", manager: mgr} + app := &App{prdName: "auth", eng: eng} got := app.renderWorktreeInfoLine() if !strings.Contains(got, "current directory") { t.Errorf("renderWorktreeInfoLine() should contain 'current directory' for branch-only mode, got %q", got) diff --git a/internal/tui/log.go b/internal/tui/log.go index 7cbeecf9..2ea81f59 100644 --- a/internal/tui/log.go +++ b/internal/tui/log.go @@ -76,8 +76,7 @@ func (l *LogViewer) AddEvent(event loop.Event) { // Filter out events we don't want to display switch event.Type { case loop.EventAssistantText, loop.EventToolStart, loop.EventToolResult, - loop.EventStoryDone, loop.EventComplete, loop.EventError, loop.EventRetrying, - loop.EventWatchdogTimeout: + loop.EventStoryStarted, loop.EventComplete, loop.EventError, loop.EventRetrying: // Pre-render and cache lines if l.width > 0 { entry.cachedLines = l.renderEntry(entry) @@ -354,16 +353,14 @@ func (l *LogViewer) renderEntry(entry LogEntry) []string { return l.renderToolCard(entry) case loop.EventToolResult: return l.renderToolResult(entry) - case loop.EventStoryDone: - return l.renderStoryDone(entry) + case loop.EventStoryStarted: + return l.renderStoryStarted(entry) case loop.EventComplete: return l.renderComplete(entry) case loop.EventError: return l.renderError(entry) case loop.EventRetrying: return l.renderRetrying(entry) - case loop.EventWatchdogTimeout: - return l.renderWatchdogTimeout(entry) default: return l.renderText(entry) } @@ -554,20 +551,20 @@ func stripLineNumbers(code string) string { return strings.Join(result, "\n") } -// renderStoryDone renders a story done marker. -func (l *LogViewer) renderStoryDone(entry LogEntry) []string { +// renderStoryStarted renders a story started marker. +func (l *LogViewer) renderStoryStarted(entry LogEntry) []string { storyStyle := lipgloss.NewStyle(). - Foreground(SuccessColor). + Foreground(PrimaryColor). Bold(true). Padding(0, 1) - dividerStyle := lipgloss.NewStyle().Foreground(SuccessColor) + dividerStyle := lipgloss.NewStyle().Foreground(PrimaryColor) divider := dividerStyle.Render(strings.Repeat("─", l.width-4)) return []string{ "", divider, - storyStyle.Render("✓ Story done"), + storyStyle.Render(fmt.Sprintf("▶ Working on: %s", entry.StoryID)), divider, "", } @@ -618,17 +615,3 @@ func (l *LogViewer) renderRetrying(entry LogEntry) []string { return []string{retryStyle.Render("🔄 " + text)} } - -// renderWatchdogTimeout renders a watchdog timeout message. -func (l *LogViewer) renderWatchdogTimeout(entry LogEntry) []string { - style := lipgloss.NewStyle(). - Foreground(WarningColor). - Bold(true) - - text := entry.Text - if text == "" { - text = "Watchdog timeout: process killed" - } - - return []string{style.Render("⏱ " + text)} -} diff --git a/internal/tui/log_perf_test.go b/internal/tui/log_perf_test.go index 73f9fd77..3c17ded0 100644 --- a/internal/tui/log_perf_test.go +++ b/internal/tui/log_perf_test.go @@ -27,7 +27,7 @@ func makeToolResultEvent(text string) loop.Event { // makeStoryEvent creates a story started event. func makeStoryEvent(storyID string) loop.Event { - return loop.Event{Type: loop.EventStoryDone, StoryID: storyID} + return loop.Event{Type: loop.EventStoryStarted, StoryID: storyID} } // --- AddEvent caching tests --- @@ -85,10 +85,10 @@ func TestTotalLineCount_AccurateAcrossEventTypes(t *testing.T) { lv.SetSize(100, 30) // Add diverse events - lv.AddEvent(makeStoryEvent("US-1")) // 5 lines (blank, divider, title, divider, blank) + lv.AddEvent(makeStoryEvent("US-1")) // 5 lines (blank, divider, title, divider, blank) lv.AddEvent(makeToolStartEvent("Read", map[string]interface{}{"file_path": "/test.go"})) // 1 line - lv.AddEvent(makeToolResultEvent("some output")) // 1 line - lv.AddEvent(makeTextEvent("Hello")) // 1 line + lv.AddEvent(makeToolResultEvent("some output")) // 1 line + lv.AddEvent(makeTextEvent("Hello")) // 1 line // Count actual cached lines actualTotal := 0 diff --git a/internal/tui/picker.go b/internal/tui/picker.go index dd1e4032..f013ccdd 100644 --- a/internal/tui/picker.go +++ b/internal/tui/picker.go @@ -47,10 +47,10 @@ const ( // CleanConfirmation holds the state of the clean confirmation dialog. type CleanConfirmation struct { - EntryName string // Name of the PRD being cleaned - Branch string // Branch name to display - WorktreeDir string // Worktree path to display - SelectedIdx int // Selected option index (0-2) + EntryName string // Name of the PRD being cleaned + Branch string // Branch name to display + WorktreeDir string // Worktree path to display + SelectedIdx int // Selected option index (0-2) } // CleanResult holds the result of a clean operation for display. @@ -61,18 +61,18 @@ type CleanResult struct { // PRDPicker manages the PRD picker modal state. type PRDPicker struct { - entries []PRDEntry - selectedIndex int - width int - height int - basePath string // Base path where .chief/prds/ is located - currentPRD string // Name of the currently active PRD - inputMode bool // Whether we're in input mode for new PRD name - inputValue string // The current input value for new PRD name - manager *loop.Manager // Reference to the loop manager for status updates - mergeResult *MergeResult // Result of the last merge operation (nil = none) - cleanConfirmation *CleanConfirmation // Active clean confirmation dialog (nil = none) - cleanResult *CleanResult // Result of the last clean operation (nil = none) + entries []PRDEntry + selectedIndex int + width int + height int + basePath string // Base path where .chief/prds/ is located + currentPRD string // Name of the currently active PRD + inputMode bool // Whether we're in input mode for new PRD name + inputValue string // The current input value for new PRD name + manager *loop.Manager // Reference to the loop manager for status updates + mergeResult *MergeResult // Result of the last merge operation (nil = none) + cleanConfirmation *CleanConfirmation // Active clean confirmation dialog (nil = none) + cleanResult *CleanResult // Result of the last clean operation (nil = none) } // NewPRDPicker creates a new PRD picker. @@ -117,13 +117,7 @@ func (p *PRDPicker) Refresh() { } name := entry.Name() - dirPath := filepath.Join(prdsDir, name) - prdPath := filepath.Join(dirPath, "prd.md") - - // Skip directories without prd.md (empty/incomplete) - if _, err := os.Stat(prdPath); os.IsNotExist(err) { - continue - } + prdPath := filepath.Join(prdsDir, name, "prd.json") prdEntry := p.loadPRDEntry(name, prdPath) p.entries = append(p.entries, prdEntry) @@ -131,7 +125,7 @@ func (p *PRDPicker) Refresh() { } // Also check if there's a "main" PRD directly in .chief/ (legacy location) - mainPrdPath := filepath.Join(p.basePath, ".chief", "prd.md") + mainPrdPath := filepath.Join(p.basePath, ".chief", "prd.json") if _, err := os.Stat(mainPrdPath); err == nil && !addedNames["main"] { prdEntry := p.loadPRDEntry("main", mainPrdPath) p.entries = append(p.entries, prdEntry) @@ -173,11 +167,11 @@ func (p *PRDPicker) Refresh() { if !found { p.entries = append(p.entries, PRDEntry{ Name: prdName, - Path: filepath.Join(p.basePath, ".chief", "prds", prdName, "prd.md"), + Path: filepath.Join(p.basePath, ".chief", "prds", prdName, "prd.json"), LoopState: loop.LoopStateReady, WorktreeDir: absPath, Orphaned: true, - LoadError: fmt.Errorf("orphaned worktree (no prd.md)"), + LoadError: fmt.Errorf("orphaned worktree (no prd.json)"), }) } } diff --git a/internal/tui/picker_test.go b/internal/tui/picker_test.go deleted file mode 100644 index 359ae14e..00000000 --- a/internal/tui/picker_test.go +++ /dev/null @@ -1,1002 +0,0 @@ -package tui - -import ( - "fmt" - "os" - "path/filepath" - "testing" - "unicode/utf8" - - "github.com/minicodemonkey/chief/internal/loop" -) - -func TestRenderEntryWithBranchAndWorktree(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - currentPRD: "", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - } - - result := p.renderEntry(p.entries[0], false, 80) - if result == "" { - t.Fatal("expected non-empty render result") - } - // Should contain branch name - if !containsText(result, "chief/auth") { - t.Errorf("expected branch 'chief/auth' in output, got: %s", result) - } - // Should contain worktree path - if !containsText(result, ".chief/worktrees/auth/") { - t.Errorf("expected worktree path in output, got: %s", result) - } -} - -func TestRenderEntryNoBranch(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - currentPRD: "", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 3, - Total: 8, - LoopState: loop.LoopStateRunning, - Iteration: 2, - Branch: "", - }, - }, - } - - result := p.renderEntry(p.entries[0], false, 80) - if result == "" { - t.Fatal("expected non-empty render result") - } - // Should NOT contain branch brackets - if containsText(result, "chief/") { - t.Errorf("expected no branch in output when branch is empty, got: %s", result) - } -} - -func TestRenderEntryNoBranchOmitsCurrentDirectoryLabel(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - currentPRD: "", - entries: []PRDEntry{ - { - Name: "legacy", - Completed: 5, - Total: 5, - LoopState: loop.LoopStateComplete, - Branch: "", - }, - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - } - - // Render the entry without a branch — should NOT show "(current directory)" to keep worktrees low-key - result := p.renderEntry(p.entries[0], false, 80) - if containsText(result, "(current directory)") { - t.Errorf("expected no '(current directory)' label for branchless entry, got: %s", result) - } -} - -func TestRenderEntryNarrowTerminalOmitsBranchPath(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - currentPRD: "", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - } - - // Very narrow width — should not crash and should omit branch/path info - result := p.renderEntry(p.entries[0], false, 35) - if result == "" { - t.Fatal("expected non-empty render result even at narrow width") - } - // At 35 chars wide, remaining space (35-32=3) is too small for branch info - if containsText(result, "chief/auth") { - t.Errorf("expected branch to be omitted at narrow width, got: %s", result) - } -} - -func TestFormatBranchPathFull(t *testing.T) { - p := &PRDPicker{} - - result := p.formatBranchPath("chief/auth", ".chief/worktrees/auth/", 50) - expected := " chief/auth .chief/worktrees/auth/" - if result != expected { - t.Errorf("expected %q, got %q", expected, result) - } -} - -func TestFormatBranchPathTruncatesPath(t *testing.T) { - p := &PRDPicker{} - - result := p.formatBranchPath("chief/auth", ".chief/worktrees/auth/", 30) - // Should contain branch but path should be truncated with … - if !containsSubstring(result, "chief/auth") { - t.Errorf("expected branch in truncated output, got: %s", result) - } - runeCount := utf8.RuneCountInString(result) - if runeCount > 30 { - t.Errorf("expected result to fit within 30 display chars, got %d: %s", runeCount, result) - } -} - -func TestFormatBranchPathTruncatesBranch(t *testing.T) { - p := &PRDPicker{} - - // Very small width — only room for branch (truncated) - result := p.formatBranchPath("chief/very-long-branch-name", ".chief/worktrees/auth/", 15) - runeCount := utf8.RuneCountInString(result) - if runeCount > 15 { - t.Errorf("expected result to fit within 15 display chars, got %d: %s", runeCount, result) - } -} - -func TestWorktreeDisplayPathWithWorktree(t *testing.T) { - p := &PRDPicker{basePath: "/project"} - - entry := PRDEntry{WorktreeDir: "/project/.chief/worktrees/auth"} - result := p.worktreeDisplayPath(entry) - if result != ".chief/worktrees/auth/" { - t.Errorf("expected '.chief/worktrees/auth/', got %q", result) - } -} - -func TestWorktreeDisplayPathWithoutWorktree(t *testing.T) { - p := &PRDPicker{basePath: "/project"} - - entry := PRDEntry{WorktreeDir: ""} - result := p.worktreeDisplayPath(entry) - if result != "(current directory)" { - t.Errorf("expected '(current directory)', got %q", result) - } -} - -func TestRenderEntryWithLoadError(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "broken", - LoadError: fmt.Errorf("parse error"), - Branch: "chief/broken", - WorktreeDir: "/project/.chief/worktrees/broken", - }, - }, - } - - result := p.renderEntry(p.entries[0], false, 80) - // With load error, should show [error] but not branch/worktree info - if !containsText(result, "error") { - t.Errorf("expected [error] in output, got: %s", result) - } -} - -func TestCanMergeCompletedWithBranch(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - }, - }, - selectedIndex: 0, - } - if !p.CanMerge() { - t.Error("expected CanMerge() to return true for completed PRD with branch") - } -} - -func TestCanMergeNoBranch(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "", - }, - }, - selectedIndex: 0, - } - if p.CanMerge() { - t.Error("expected CanMerge() to return false for completed PRD without branch") - } -} - -func TestCanMergeRunningPRD(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 3, - Total: 8, - LoopState: loop.LoopStateRunning, - Branch: "chief/auth", - }, - }, - selectedIndex: 0, - } - if p.CanMerge() { - t.Error("expected CanMerge() to return false for running PRD") - } -} - -func TestCanMergeAllPassedButNotCompleteState(t *testing.T) { - // All stories pass but loop state is Ready (e.g., not started via loop) - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 5, - Total: 5, - LoopState: loop.LoopStateReady, - Branch: "chief/auth", - }, - }, - selectedIndex: 0, - } - if !p.CanMerge() { - t.Error("expected CanMerge() to return true when all stories pass, even if LoopState is Ready") - } -} - -func TestMergeResultSuccessRendering(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - width: 80, - height: 24, - entries: []PRDEntry{ - {Name: "auth", Branch: "chief/auth"}, - }, - mergeResult: &MergeResult{ - Success: true, - Message: "Merged chief/auth into main", - Branch: "chief/auth", - }, - } - - result := p.Render() - if !containsText(result, "Merge Successful") { - t.Errorf("expected 'Merge Successful' in success render, got: %s", stripAnsi(result)) - } - if !containsText(result, "Merged chief/auth into main") { - t.Errorf("expected merge message in output, got: %s", stripAnsi(result)) - } - if !containsText(result, "Press any key to continue") { - t.Errorf("expected dismiss hint in output, got: %s", stripAnsi(result)) - } -} - -func TestMergeResultConflictRendering(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - width: 80, - height: 24, - entries: []PRDEntry{ - {Name: "auth", Branch: "chief/auth"}, - }, - mergeResult: &MergeResult{ - Success: false, - Message: "Failed to merge chief/auth into current branch", - Conflicts: []string{"src/auth.go", "src/handler.go"}, - Branch: "chief/auth", - }, - } - - result := p.Render() - if !containsText(result, "Merge Conflict") { - t.Errorf("expected 'Merge Conflict' in conflict render, got: %s", stripAnsi(result)) - } - if !containsText(result, "src/auth.go") { - t.Errorf("expected conflicting file in output, got: %s", stripAnsi(result)) - } - if !containsText(result, "src/handler.go") { - t.Errorf("expected conflicting file in output, got: %s", stripAnsi(result)) - } - if !containsText(result, "git merge chief/auth") { - t.Errorf("expected manual merge instruction in output, got: %s", stripAnsi(result)) - } -} - -func TestMergeResultClearsOnDismiss(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - mergeResult: &MergeResult{ - Success: true, - Message: "Merged", - Branch: "chief/auth", - }, - } - - if !p.HasMergeResult() { - t.Error("expected HasMergeResult() to return true") - } - - p.ClearMergeResult() - - if p.HasMergeResult() { - t.Error("expected HasMergeResult() to return false after clear") - } -} - -func TestFooterShowsMergeHintForCompletedPRD(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - }, - }, - selectedIndex: 0, - } - - shortcuts := p.buildFooterShortcuts() - if !containsSubstring(shortcuts, "m: merge") { - t.Errorf("expected 'm: merge' in footer for completed PRD with branch, got: %s", shortcuts) - } -} - -func TestFooterHidesMergeHintForRunningPRD(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 3, - Total: 8, - LoopState: loop.LoopStateRunning, - Iteration: 2, - Branch: "chief/auth", - }, - }, - selectedIndex: 0, - } - - shortcuts := p.buildFooterShortcuts() - if containsSubstring(shortcuts, "m: merge") { - t.Errorf("expected no 'm: merge' in footer for running PRD, got: %s", shortcuts) - } -} - -// --- Clean Action Tests --- - -func TestCanCleanNonRunningWithWorktree(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - if !p.CanClean() { - t.Error("expected CanClean() to return true for completed non-running PRD with worktree") - } -} - -func TestCanCleanDisabledForRunningPRD(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 3, - Total: 8, - LoopState: loop.LoopStateRunning, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - if p.CanClean() { - t.Error("expected CanClean() to return false for running PRD") - } -} - -func TestCanCleanDisabledWithoutWorktree(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - }, - }, - selectedIndex: 0, - } - if p.CanClean() { - t.Error("expected CanClean() to return false for PRD without worktree") - } -} - -func TestCanCleanStoppedPRD(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 3, - Total: 8, - LoopState: loop.LoopStateStopped, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - if !p.CanClean() { - t.Error("expected CanClean() to return true for stopped PRD with worktree") - } -} - -func TestCleanConfirmationDialog(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - - // Start clean confirmation - p.StartCleanConfirmation() - - if !p.HasCleanConfirmation() { - t.Fatal("expected HasCleanConfirmation() to return true after start") - } - - cc := p.GetCleanConfirmation() - if cc.EntryName != "auth" { - t.Errorf("expected EntryName 'auth', got %q", cc.EntryName) - } - if cc.Branch != "chief/auth" { - t.Errorf("expected Branch 'chief/auth', got %q", cc.Branch) - } - if cc.SelectedIdx != 0 { - t.Errorf("expected SelectedIdx 0, got %d", cc.SelectedIdx) - } - - // Default selection is RemoveAll - if p.GetCleanOption() != CleanOptionRemoveAll { - t.Errorf("expected CleanOptionRemoveAll by default, got %d", p.GetCleanOption()) - } -} - -func TestCleanConfirmationNavigation(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - p.StartCleanConfirmation() - - // Move down to "Remove worktree only" - p.CleanConfirmMoveDown() - if p.GetCleanOption() != CleanOptionWorktreeOnly { - t.Errorf("expected CleanOptionWorktreeOnly after move down, got %d", p.GetCleanOption()) - } - - // Move down to "Cancel" - p.CleanConfirmMoveDown() - if p.GetCleanOption() != CleanOptionCancel { - t.Errorf("expected CleanOptionCancel after two moves down, got %d", p.GetCleanOption()) - } - - // Move down again - should stay at Cancel (index 2) - p.CleanConfirmMoveDown() - if p.GetCleanOption() != CleanOptionCancel { - t.Errorf("expected CleanOptionCancel to remain after extra move down, got %d", p.GetCleanOption()) - } - - // Move back up - p.CleanConfirmMoveUp() - if p.GetCleanOption() != CleanOptionWorktreeOnly { - t.Errorf("expected CleanOptionWorktreeOnly after move up, got %d", p.GetCleanOption()) - } -} - -func TestCleanConfirmationCancel(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - p.StartCleanConfirmation() - - if !p.HasCleanConfirmation() { - t.Fatal("expected confirmation to be active") - } - - p.CancelCleanConfirmation() - - if p.HasCleanConfirmation() { - t.Error("expected confirmation to be cancelled") - } -} - -func TestCleanConfirmationRendering(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - width: 80, - height: 24, - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - p.StartCleanConfirmation() - - result := p.Render() - stripped := stripAnsi(result) - - if !containsText(result, "Clean Worktree") { - t.Errorf("expected 'Clean Worktree' in render, got: %s", stripped) - } - if !containsText(result, "auth") { - t.Errorf("expected PRD name 'auth' in render, got: %s", stripped) - } - if !containsText(result, "chief/auth") { - t.Errorf("expected branch 'chief/auth' in render, got: %s", stripped) - } - if !containsText(result, "Remove worktree + delete branch") { - t.Errorf("expected option text in render, got: %s", stripped) - } - if !containsText(result, "Remove worktree only") { - t.Errorf("expected option text in render, got: %s", stripped) - } - if !containsText(result, "Cancel") { - t.Errorf("expected 'Cancel' option in render, got: %s", stripped) - } -} - -func TestCleanResultSuccessRendering(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - width: 80, - height: 24, - entries: []PRDEntry{ - {Name: "auth"}, - }, - cleanResult: &CleanResult{ - Success: true, - Message: "Removed worktree and deleted branch chief/auth", - }, - } - - result := p.Render() - if !containsText(result, "Clean Successful") { - t.Errorf("expected 'Clean Successful' in success render, got: %s", stripAnsi(result)) - } - if !containsText(result, "Removed worktree and deleted branch chief/auth") { - t.Errorf("expected clean message in output, got: %s", stripAnsi(result)) - } - if !containsText(result, "Press any key to continue") { - t.Errorf("expected dismiss hint in output, got: %s", stripAnsi(result)) - } -} - -func TestCleanResultErrorRendering(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - width: 80, - height: 24, - entries: []PRDEntry{ - {Name: "auth"}, - }, - cleanResult: &CleanResult{ - Success: false, - Message: "Failed to remove worktree: permission denied", - }, - } - - result := p.Render() - if !containsText(result, "Clean Failed") { - t.Errorf("expected 'Clean Failed' in error render, got: %s", stripAnsi(result)) - } - if !containsText(result, "permission denied") { - t.Errorf("expected error message in output, got: %s", stripAnsi(result)) - } -} - -func TestCleanResultClearsOnDismiss(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - cleanResult: &CleanResult{ - Success: true, - Message: "Cleaned", - }, - } - - if !p.HasCleanResult() { - t.Error("expected HasCleanResult() to return true") - } - - p.ClearCleanResult() - - if p.HasCleanResult() { - t.Error("expected HasCleanResult() to return false after clear") - } -} - -func TestFooterShowsCleanHintForNonRunningPRDWithWorktree(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - - shortcuts := p.buildFooterShortcuts() - if !containsSubstring(shortcuts, "c: clean") { - t.Errorf("expected 'c: clean' in footer for completed PRD with worktree, got: %s", shortcuts) - } -} - -func TestFooterHidesCleanHintForRunningPRD(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 3, - Total: 8, - LoopState: loop.LoopStateRunning, - Iteration: 2, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - }, - }, - selectedIndex: 0, - } - - shortcuts := p.buildFooterShortcuts() - if containsSubstring(shortcuts, "c: clean") { - t.Errorf("expected no 'c: clean' in footer for running PRD, got: %s", shortcuts) - } -} - -func TestFooterHidesCleanHintForPRDWithoutWorktree(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - }, - }, - selectedIndex: 0, - } - - shortcuts := p.buildFooterShortcuts() - if containsSubstring(shortcuts, "c: clean") { - t.Errorf("expected no 'c: clean' in footer for PRD without worktree, got: %s", shortcuts) - } -} - -// --- Orphaned Worktree Tests --- - -func TestRenderEntryOrphanedWithPRD(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - currentPRD: "", - entries: []PRDEntry{ - { - Name: "auth", - Completed: 8, - Total: 8, - LoopState: loop.LoopStateComplete, - Branch: "chief/auth", - WorktreeDir: "/project/.chief/worktrees/auth", - Orphaned: true, - }, - }, - } - - result := p.renderEntry(p.entries[0], false, 80) - if !containsText(result, "[orphaned]") { - t.Errorf("expected '[orphaned]' indicator for orphaned entry with PRD, got: %s", stripAnsi(result)) - } - // Should still show progress since PRD is loaded - if !containsText(result, "8/8") { - t.Errorf("expected progress '8/8' for orphaned entry with PRD, got: %s", stripAnsi(result)) - } -} - -func TestRenderEntryOrphanedWithoutPRD(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - currentPRD: "", - entries: []PRDEntry{ - { - Name: "stale-project", - WorktreeDir: "/project/.chief/worktrees/stale-project", - Orphaned: true, - LoadError: fmt.Errorf("orphaned worktree (no prd.json)"), - }, - }, - } - - result := p.renderEntry(p.entries[0], false, 80) - if !containsText(result, "[orphaned worktree]") { - t.Errorf("expected '[orphaned worktree]' for orphaned entry without PRD, got: %s", stripAnsi(result)) - } -} - -func TestCanCleanOrphanedWorktree(t *testing.T) { - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "auth", - WorktreeDir: "/project/.chief/worktrees/auth", - Orphaned: true, - LoopState: loop.LoopStateReady, - }, - }, - selectedIndex: 0, - } - if !p.CanClean() { - t.Error("expected CanClean() to return true for orphaned worktree") - } -} - -func TestOrphanedWorktreeNotTracked(t *testing.T) { - // An orphaned entry with no PRD loaded should still be cleanable - p := &PRDPicker{ - basePath: "/project", - entries: []PRDEntry{ - { - Name: "stale", - WorktreeDir: "/project/.chief/worktrees/stale", - Orphaned: true, - LoopState: loop.LoopStateReady, - LoadError: fmt.Errorf("orphaned worktree (no prd.json)"), - }, - }, - selectedIndex: 0, - } - if !p.CanClean() { - t.Error("expected CanClean() to return true for orphaned worktree without PRD") - } -} - -// --- Empty Directory Filtering Tests --- - -func TestRefreshIgnoresEmptyDirectories(t *testing.T) { - tmpDir := t.TempDir() - prdsDir := filepath.Join(tmpDir, ".chief", "prds") - - // Create an empty directory (simulates cancelled chief new) - emptyDir := filepath.Join(prdsDir, "cancelled") - if err := os.MkdirAll(emptyDir, 0755); err != nil { - t.Fatalf("Failed to create empty dir: %v", err) - } - - p := &PRDPicker{ - basePath: tmpDir, - entries: make([]PRDEntry, 0), - } - p.Refresh() - - if len(p.entries) != 0 { - t.Errorf("expected 0 entries for empty directory, got %d", len(p.entries)) - } -} - -func TestRefreshShowsDirectoryWithPrdMdOnly(t *testing.T) { - tmpDir := t.TempDir() - prdsDir := filepath.Join(tmpDir, ".chief", "prds") - - // Create directory with only prd.md (no json yet) - prdDir := filepath.Join(prdsDir, "in-progress") - if err := os.MkdirAll(prdDir, 0755); err != nil { - t.Fatalf("Failed to create dir: %v", err) - } - if err := os.WriteFile(filepath.Join(prdDir, "prd.md"), []byte("# My PRD"), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) - } - - p := &PRDPicker{ - basePath: tmpDir, - entries: make([]PRDEntry, 0), - } - p.Refresh() - - if len(p.entries) != 1 { - t.Fatalf("expected 1 entry for directory with prd.md, got %d", len(p.entries)) - } - if p.entries[0].Name != "in-progress" { - t.Errorf("expected entry name 'in-progress', got %q", p.entries[0].Name) - } -} - -func TestRefreshSkipsDirectoryWithOnlyPrdJson(t *testing.T) { - tmpDir := t.TempDir() - prdsDir := filepath.Join(tmpDir, ".chief", "prds") - - // Create directory with only prd.json (no prd.md) — should be skipped - prdDir := filepath.Join(prdsDir, "converted") - if err := os.MkdirAll(prdDir, 0755); err != nil { - t.Fatalf("Failed to create dir: %v", err) - } - prdJSON := `{"project":"test","userStories":[]}` - if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), []byte(prdJSON), 0644); err != nil { - t.Fatalf("Failed to create prd.json: %v", err) - } - - p := &PRDPicker{ - basePath: tmpDir, - entries: make([]PRDEntry, 0), - } - p.Refresh() - - if len(p.entries) != 0 { - t.Fatalf("expected 0 entries for directory with only prd.json (no prd.md), got %d", len(p.entries)) - } -} - -func TestRefreshMixedDirectories(t *testing.T) { - tmpDir := t.TempDir() - prdsDir := filepath.Join(tmpDir, ".chief", "prds") - - // Empty directory (should be filtered) - if err := os.MkdirAll(filepath.Join(prdsDir, "empty"), 0755); err != nil { - t.Fatalf("Failed to create dir: %v", err) - } - - // Directory with prd.md (should be shown) - validDir := filepath.Join(prdsDir, "valid") - if err := os.MkdirAll(validDir, 0755); err != nil { - t.Fatalf("Failed to create dir: %v", err) - } - if err := os.WriteFile(filepath.Join(validDir, "prd.md"), []byte("# PRD"), 0644); err != nil { - t.Fatalf("Failed to create prd.md: %v", err) - } - - // Another empty directory (should be filtered) - if err := os.MkdirAll(filepath.Join(prdsDir, "also-empty"), 0755); err != nil { - t.Fatalf("Failed to create dir: %v", err) - } - - p := &PRDPicker{ - basePath: tmpDir, - entries: make([]PRDEntry, 0), - } - p.Refresh() - - if len(p.entries) != 1 { - t.Errorf("expected 1 entry (filtering 2 empty dirs), got %d", len(p.entries)) - } - if len(p.entries) > 0 && p.entries[0].Name != "valid" { - t.Errorf("expected entry name 'valid', got %q", p.entries[0].Name) - } -} - -// containsText checks if rendered output contains a substring (ignoring ANSI codes). -func containsText(rendered, substr string) bool { - // Strip ANSI escape sequences for comparison - return containsSubstring(stripAnsi(rendered), substr) -} - -// containsSubstring is a simple substring check. -func containsSubstring(s, substr string) bool { - return len(s) >= len(substr) && indexOf(s, substr) >= 0 -} - -func indexOf(s, substr string) int { - for i := 0; i <= len(s)-len(substr); i++ { - if s[i:i+len(substr)] == substr { - return i - } - } - return -1 -} - -// stripAnsi removes ANSI escape codes from a string. -func stripAnsi(s string) string { - var result []byte - i := 0 - for i < len(s) { - if s[i] == '\x1b' && i+1 < len(s) && s[i+1] == '[' { - // Skip to end of escape sequence - j := i + 2 - for j < len(s) && !((s[j] >= 'A' && s[j] <= 'Z') || (s[j] >= 'a' && s[j] <= 'z')) { - j++ - } - if j < len(s) { - j++ // skip the final letter - } - i = j - } else { - result = append(result, s[i]) - i++ - } - } - return string(result) -} diff --git a/internal/tui/settings.go b/internal/tui/settings.go index 455a1042..8365fe0f 100644 --- a/internal/tui/settings.go +++ b/internal/tui/settings.go @@ -12,17 +12,17 @@ import ( type SettingsItemType int const ( - SettingsItemBool SettingsItemType = iota + SettingsItemBool SettingsItemType = iota SettingsItemString ) // SettingsItem represents a single editable setting. type SettingsItem struct { - Section string - Label string - Key string // config key for identification - Type SettingsItemType - BoolVal bool + Section string + Label string + Key string // config key for identification + Type SettingsItemType + BoolVal bool StringVal string } @@ -39,7 +39,7 @@ type SettingsOverlay struct { editBuffer string // GH CLI validation error - ghError string + ghError string showGHError bool } diff --git a/internal/tui/styles.go b/internal/tui/styles.go index 56f3fae5..31a552ad 100644 --- a/internal/tui/styles.go +++ b/internal/tui/styles.go @@ -21,8 +21,8 @@ var ( TextBrightColor = lipgloss.Color("#FFFFFF") // Bright white - emphasis // Background colors - BgColor = lipgloss.Color("#1E1E2E") // Dark background - BgSelectedColor = lipgloss.Color("#313244") // Selected item background + BgColor = lipgloss.Color("#1E1E2E") // Dark background + BgSelectedColor = lipgloss.Color("#313244") // Selected item background BgHighlightColor = lipgloss.Color("#45475A") // Highlight background ) diff --git a/internal/tui/tabbar.go b/internal/tui/tabbar.go index 8964f72b..a5627154 100644 --- a/internal/tui/tabbar.go +++ b/internal/tui/tabbar.go @@ -66,7 +66,7 @@ func (t *TabBar) Refresh() { } name := entry.Name() - prdPath := filepath.Join(prdsDir, name, "prd.md") + prdPath := filepath.Join(prdsDir, name, "prd.json") tabEntry := t.loadTabEntry(name, prdPath) t.entries = append(t.entries, tabEntry) @@ -74,7 +74,7 @@ func (t *TabBar) Refresh() { } // Also check if there's a "main" PRD directly in .chief/ (legacy location) - mainPrdPath := filepath.Join(t.baseDir, ".chief", "prd.md") + mainPrdPath := filepath.Join(t.baseDir, ".chief", "prd.json") if _, err := os.Stat(mainPrdPath); err == nil && !addedNames["main"] { tabEntry := t.loadTabEntry("main", mainPrdPath) t.entries = append(t.entries, tabEntry) diff --git a/internal/uplink/batcher.go b/internal/uplink/batcher.go new file mode 100644 index 00000000..7a9f4838 --- /dev/null +++ b/internal/uplink/batcher.go @@ -0,0 +1,271 @@ +package uplink + +import ( + "context" + "crypto/rand" + "encoding/json" + "fmt" + "log" + "sync" + "time" +) + +// Flush tier durations. +const ( + // tierImmediate flushes immediately (0ms delay). + tierImmediate = 0 * time.Millisecond + + // tierStandard flushes after 200ms. + tierStandard = 200 * time.Millisecond + + // tierLowPriority flushes after 1s. + tierLowPriority = 1 * time.Second + + // maxBatchMessages is the maximum number of messages before a forced flush. + maxBatchMessages = 20 + + // maxBufferMessages is the maximum total messages in the buffer before dropping. + maxBufferMessages = 1000 + + // maxBufferBytes is the maximum total payload size (5MB) before dropping. + maxBufferBytes = 5 * 1024 * 1024 +) + +// tier identifies a flush priority tier. +type tier int + +const ( + tierIDImmediate tier = iota + tierIDStandard + tierIDLowPriority +) + +// tierFor returns the tier for a given message type. +func tierFor(msgType string) tier { + switch msgType { + case "run_complete", "run_paused", "error", "clone_complete", "session_expired", "quota_exhausted", "prd_response_complete": + return tierIDImmediate + case "claude_output", "prd_output", "run_progress", "clone_progress": + return tierIDStandard + case "state_snapshot", "project_state", "project_list", "settings", "log_lines": + return tierIDLowPriority + default: + // Unknown types go to standard tier. + return tierIDStandard + } +} + +// tierDelay returns the flush delay for a tier. +func tierDelay(t tier) time.Duration { + switch t { + case tierIDImmediate: + return tierImmediate + case tierIDStandard: + return tierStandard + case tierIDLowPriority: + return tierLowPriority + default: + return tierStandard + } +} + +// bufferedMessage is a message waiting to be flushed. +type bufferedMessage struct { + data json.RawMessage + tier tier + size int +} + +// SendFunc is the function called on flush to send a batch of messages. +// batchID is a unique UUID for idempotency. The function should retry internally if needed. +type SendFunc func(batchID string, messages []json.RawMessage) error + +// Batcher collects outgoing messages and flushes them in batches. +// Messages are assigned to priority tiers that control flush timing. +// Flushes are sequential — the next flush waits for the current one to complete. +type Batcher struct { + sendFn SendFunc + + mu sync.Mutex + messages []bufferedMessage + totalSize int + flushNotify chan struct{} // signals the run loop that a flush may be needed + stopped bool + + // Timer management for tier-based flushing. + standardTimer *time.Timer + lowPriorityTimer *time.Timer + standardActive bool + lowPriorityActive bool +} + +// NewBatcher creates a Batcher that calls sendFn on each flush. +func NewBatcher(sendFn SendFunc) *Batcher { + return &Batcher{ + sendFn: sendFn, + flushNotify: make(chan struct{}, 1), + } +} + +// Enqueue adds a message to the appropriate tier buffer. +// It is safe to call from multiple goroutines. +func (b *Batcher) Enqueue(msg json.RawMessage, msgType string) { + t := tierFor(msgType) + size := len(msg) + + b.mu.Lock() + defer b.mu.Unlock() + + if b.stopped { + return + } + + // Check buffer limits and drop low-priority messages if full. + for b.totalSize+size > maxBufferBytes || len(b.messages)+1 > maxBufferMessages { + if !b.dropLowestPriority() { + // Nothing left to drop — reject this message too. + log.Printf("batcher: buffer full, dropping %s message (%d bytes)", msgType, size) + return + } + } + + b.messages = append(b.messages, bufferedMessage{data: msg, tier: t, size: size}) + b.totalSize += size + + // Determine if we need to flush now. + shouldFlushNow := t == tierIDImmediate || len(b.messages) >= maxBatchMessages + + if shouldFlushNow { + b.notifyFlush() + return + } + + // Start tier timers if not already running. + if t == tierIDStandard && !b.standardActive { + b.standardActive = true + if b.standardTimer == nil { + b.standardTimer = time.AfterFunc(tierStandard, func() { + b.mu.Lock() + b.standardActive = false + b.mu.Unlock() + b.notifyFlush() + }) + } else { + b.standardTimer.Reset(tierStandard) + } + } + if t == tierIDLowPriority && !b.lowPriorityActive { + b.lowPriorityActive = true + if b.lowPriorityTimer == nil { + b.lowPriorityTimer = time.AfterFunc(tierLowPriority, func() { + b.mu.Lock() + b.lowPriorityActive = false + b.mu.Unlock() + b.notifyFlush() + }) + } else { + b.lowPriorityTimer.Reset(tierLowPriority) + } + } +} + +// notifyFlush signals the run loop that a flush should happen. +// Must not be called with b.mu held if blocking is possible, but the channel is buffered. +func (b *Batcher) notifyFlush() { + select { + case b.flushNotify <- struct{}{}: + default: + // Already notified. + } +} + +// dropLowestPriority removes the last low-priority message from the buffer. +// Returns false if there are no low-priority messages to drop. +// Caller must hold b.mu. +func (b *Batcher) dropLowestPriority() bool { + // Search from the end for the lowest-priority message. + // Priority order for dropping: low priority first, then standard. + for priority := tierIDLowPriority; priority >= tierIDStandard; priority-- { + for i := len(b.messages) - 1; i >= 0; i-- { + if b.messages[i].tier == priority { + b.totalSize -= b.messages[i].size + log.Printf("batcher: buffer overflow, dropping message at index %d (tier %d, %d bytes)", i, priority, b.messages[i].size) + b.messages = append(b.messages[:i], b.messages[i+1:]...) + return true + } + } + } + return false +} + +// Run starts the background flush loop. It blocks until ctx is done. +// Flushes are sequential — only one flush runs at a time. +func (b *Batcher) Run(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + case <-b.flushNotify: + b.flush() + } + } +} + +// Stop performs a final flush of all remaining messages, then marks the batcher as stopped. +func (b *Batcher) Stop() { + b.mu.Lock() + b.stopped = true + // Stop timers. + if b.standardTimer != nil { + b.standardTimer.Stop() + } + if b.lowPriorityTimer != nil { + b.lowPriorityTimer.Stop() + } + b.mu.Unlock() + + // Final flush. + b.flush() +} + +// flush collects all pending messages and sends them as a single batch. +// It is always called from a single goroutine (the Run loop or Stop), so flushes never overlap. +func (b *Batcher) flush() { + b.mu.Lock() + if len(b.messages) == 0 { + b.mu.Unlock() + return + } + + // Collect all messages. + msgs := make([]json.RawMessage, len(b.messages)) + for i, m := range b.messages { + msgs[i] = m.data + } + b.messages = b.messages[:0] + b.totalSize = 0 + + // Reset timer state since we're flushing everything. + b.standardActive = false + b.lowPriorityActive = false + b.mu.Unlock() + + batchID := generateBatchID() + if err := b.sendFn(batchID, msgs); err != nil { + log.Printf("batcher: flush failed (batch %s, %d messages): %v", batchID, len(msgs), err) + } +} + +// generateBatchID returns a new UUID v4 string for batch idempotency. +func generateBatchID() string { + var uuid [16]byte + if _, err := rand.Read(uuid[:]); err != nil { + // Fallback to timestamp-based ID if crypto/rand fails (extremely unlikely). + return fmt.Sprintf("batch-%d", time.Now().UnixNano()) + } + // Set version (4) and variant (RFC 4122). + uuid[6] = (uuid[6] & 0x0f) | 0x40 + uuid[8] = (uuid[8] & 0x3f) | 0x80 + return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", + uuid[0:4], uuid[4:6], uuid[6:8], uuid[8:10], uuid[10:16]) +} diff --git a/internal/uplink/batcher_test.go b/internal/uplink/batcher_test.go new file mode 100644 index 00000000..5d53f06c --- /dev/null +++ b/internal/uplink/batcher_test.go @@ -0,0 +1,601 @@ +package uplink + +import ( + "context" + "encoding/json" + "strings" + "sync" + "sync/atomic" + "testing" + "time" +) + +// flushRecord captures a single flush call. +type flushRecord struct { + batchID string + messages []json.RawMessage + time time.Time +} + +// newRecordingSendFn returns a SendFunc that records all flush calls +// and a function to retrieve the records. +func newRecordingSendFn() (SendFunc, func() []flushRecord) { + var mu sync.Mutex + var records []flushRecord + + fn := func(batchID string, messages []json.RawMessage) error { + mu.Lock() + defer mu.Unlock() + // Copy messages to avoid data races. + copied := make([]json.RawMessage, len(messages)) + copy(copied, messages) + records = append(records, flushRecord{ + batchID: batchID, + messages: copied, + time: time.Now(), + }) + return nil + } + + get := func() []flushRecord { + mu.Lock() + defer mu.Unlock() + result := make([]flushRecord, len(records)) + copy(result, records) + return result + } + + return fn, get +} + +func TestBatcher_ImmediateFlush(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + // Enqueue an immediate-tier message. + b.Enqueue(json.RawMessage(`{"type":"run_complete"}`), "run_complete") + + // Wait for flush. + time.Sleep(50 * time.Millisecond) + + records := getRecords() + if len(records) != 1 { + t.Fatalf("flush count = %d, want 1", len(records)) + } + if len(records[0].messages) != 1 { + t.Errorf("message count = %d, want 1", len(records[0].messages)) + } + if records[0].batchID == "" { + t.Error("batchID should not be empty") + } +} + +func TestBatcher_ImmediateFlushDrainsAllTiers(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + // Enqueue messages from different tiers. + b.Enqueue(json.RawMessage(`{"type":"project_state"}`), "project_state") // low + b.Enqueue(json.RawMessage(`{"type":"claude_output"}`), "claude_output") // standard + b.Enqueue(json.RawMessage(`{"type":"error"}`), "error") // immediate + + // Wait for flush. + time.Sleep(50 * time.Millisecond) + + records := getRecords() + if len(records) != 1 { + t.Fatalf("flush count = %d, want 1 (all tiers drain together)", len(records)) + } + if len(records[0].messages) != 3 { + t.Errorf("message count = %d, want 3", len(records[0].messages)) + } +} + +func TestBatcher_AllImmediateTypes(t *testing.T) { + immediateTypes := []string{ + "run_complete", "run_paused", "error", + "clone_complete", "session_expired", "quota_exhausted", + } + + for _, msgType := range immediateTypes { + t.Run(msgType, func(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + b.Enqueue(json.RawMessage(`{}`), msgType) + time.Sleep(50 * time.Millisecond) + + records := getRecords() + if len(records) != 1 { + t.Errorf("expected immediate flush for %s, got %d flushes", msgType, len(records)) + } + }) + } +} + +func TestBatcher_StandardTimerFlush(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + start := time.Now() + b.Enqueue(json.RawMessage(`{"type":"claude_output"}`), "claude_output") + + // Should not flush immediately. + time.Sleep(50 * time.Millisecond) + if len(getRecords()) != 0 { + t.Fatal("standard tier should not flush immediately") + } + + // Wait for the 200ms timer. + time.Sleep(300 * time.Millisecond) + records := getRecords() + if len(records) != 1 { + t.Fatalf("flush count = %d, want 1", len(records)) + } + + elapsed := records[0].time.Sub(start) + if elapsed < 150*time.Millisecond { + t.Errorf("flushed too early: %v (expected ~200ms)", elapsed) + } +} + +func TestBatcher_LowPriorityTimerFlush(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + start := time.Now() + b.Enqueue(json.RawMessage(`{"type":"project_state"}`), "project_state") + + // Should not flush at 200ms. + time.Sleep(300 * time.Millisecond) + if len(getRecords()) != 0 { + t.Fatal("low priority tier should not flush at 200ms") + } + + // Wait for the 1s timer. + time.Sleep(900 * time.Millisecond) + records := getRecords() + if len(records) != 1 { + t.Fatalf("flush count = %d, want 1", len(records)) + } + + elapsed := records[0].time.Sub(start) + if elapsed < 800*time.Millisecond { + t.Errorf("flushed too early: %v (expected ~1s)", elapsed) + } +} + +func TestBatcher_SizeBasedFlush(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + // Enqueue 20 low-priority messages — should trigger size-based flush + // even though the 1s timer hasn't expired. + for i := 0; i < maxBatchMessages; i++ { + b.Enqueue(json.RawMessage(`{"type":"log_lines"}`), "log_lines") + } + + time.Sleep(50 * time.Millisecond) + + records := getRecords() + if len(records) != 1 { + t.Fatalf("flush count = %d, want 1", len(records)) + } + if len(records[0].messages) != maxBatchMessages { + t.Errorf("message count = %d, want %d", len(records[0].messages), maxBatchMessages) + } +} + +func TestBatcher_StopFlushesRemaining(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + // Don't start Run — just enqueue and stop. + b.Enqueue(json.RawMessage(`{"type":"claude_output"}`), "claude_output") + b.Enqueue(json.RawMessage(`{"type":"log_lines"}`), "log_lines") + + b.Stop() + + records := getRecords() + if len(records) != 1 { + t.Fatalf("flush count = %d, want 1 (final flush on stop)", len(records)) + } + if len(records[0].messages) != 2 { + t.Errorf("message count = %d, want 2", len(records[0].messages)) + } +} + +func TestBatcher_StopPreventsNewEnqueues(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + b.Enqueue(json.RawMessage(`{"type":"claude_output"}`), "claude_output") + b.Stop() + + // Enqueue after stop should be silently dropped. + b.Enqueue(json.RawMessage(`{"type":"error"}`), "error") + + records := getRecords() + if len(records) != 1 { + t.Fatalf("flush count = %d, want 1", len(records)) + } + if len(records[0].messages) != 1 { + t.Errorf("message count = %d, want 1 (post-stop enqueue should be dropped)", len(records[0].messages)) + } +} + +func TestBatcher_BufferOverflowDropsLowPriority(t *testing.T) { + sendFn, _ := newRecordingSendFn() + b := NewBatcher(sendFn) + + // Fill with low-priority messages up to the limit. + for i := 0; i < maxBufferMessages-1; i++ { + b.Enqueue(json.RawMessage(`{"type":"log_lines"}`), "log_lines") + } + + // Buffer is almost full. Add one standard message. + b.Enqueue(json.RawMessage(`{"type":"claude_output"}`), "claude_output") + + // Buffer is now at limit. Adding an immediate message should drop a low-priority one. + b.Enqueue(json.RawMessage(`{"type":"error"}`), "error") + + b.mu.Lock() + count := len(b.messages) + // Count messages by tier. + var immediate, standard, low int + for _, m := range b.messages { + switch m.tier { + case tierIDImmediate: + immediate++ + case tierIDStandard: + standard++ + case tierIDLowPriority: + low++ + } + } + b.mu.Unlock() + + if count != maxBufferMessages { + t.Errorf("buffer count = %d, want %d", count, maxBufferMessages) + } + if immediate != 1 { + t.Errorf("immediate count = %d, want 1", immediate) + } + if standard != 1 { + t.Errorf("standard count = %d, want 1", standard) + } + if low != maxBufferMessages-2 { + t.Errorf("low priority count = %d, want %d", low, maxBufferMessages-2) + } +} + +func TestBatcher_BufferOverflowBySize(t *testing.T) { + sendFn, _ := newRecordingSendFn() + b := NewBatcher(sendFn) + + // Create a large message (~1MB). + bigPayload := strings.Repeat("x", 1024*1024) + bigMsg := json.RawMessage(`{"type":"log_lines","data":"` + bigPayload + `"}`) + + // Fill buffer with 4 big messages (~4MB). + for i := 0; i < 4; i++ { + b.Enqueue(bigMsg, "log_lines") + } + + b.mu.Lock() + countBefore := len(b.messages) + sizeBefore := b.totalSize + b.mu.Unlock() + + if countBefore != 4 { + t.Fatalf("buffer count = %d, want 4", countBefore) + } + + // Adding another big message should trigger overflow — drops a low-priority message. + bigStandard := json.RawMessage(`{"type":"claude_output","data":"` + bigPayload + `"}`) + b.Enqueue(bigStandard, "claude_output") + + b.mu.Lock() + countAfter := len(b.messages) + sizeAfter := b.totalSize + b.mu.Unlock() + + // Should have dropped one log_lines message to make room. + if countAfter != 4 { + t.Errorf("buffer count after overflow = %d, want 4", countAfter) + } + if sizeAfter >= sizeBefore+len(bigStandard) { + t.Errorf("buffer size should not exceed limit: before=%d, after=%d", sizeBefore, sizeAfter) + } +} + +func TestBatcher_FlushesNeverOverlap(t *testing.T) { + var concurrent atomic.Int32 + var maxConcurrent atomic.Int32 + + sendFn := func(batchID string, messages []json.RawMessage) error { + n := concurrent.Add(1) + // Track max concurrency. + for { + old := maxConcurrent.Load() + if n <= old || maxConcurrent.CompareAndSwap(old, n) { + break + } + } + time.Sleep(50 * time.Millisecond) // Simulate slow send. + concurrent.Add(-1) + return nil + } + + b := NewBatcher(sendFn) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + // Rapidly enqueue immediate messages to trigger many flushes. + for i := 0; i < 10; i++ { + b.Enqueue(json.RawMessage(`{"type":"error"}`), "error") + time.Sleep(10 * time.Millisecond) + } + + // Wait for all flushes to complete. + time.Sleep(600 * time.Millisecond) + + if maxConcurrent.Load() > 1 { + t.Errorf("max concurrent flushes = %d, want 1 (sequential flushes only)", maxConcurrent.Load()) + } +} + +func TestBatcher_UniqueBatchIDs(t *testing.T) { + var mu sync.Mutex + var batchIDs []string + + sendFn := func(batchID string, messages []json.RawMessage) error { + mu.Lock() + batchIDs = append(batchIDs, batchID) + mu.Unlock() + return nil + } + + b := NewBatcher(sendFn) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + // Trigger multiple flushes. + for i := 0; i < 5; i++ { + b.Enqueue(json.RawMessage(`{"type":"error"}`), "error") + time.Sleep(50 * time.Millisecond) + } + + time.Sleep(100 * time.Millisecond) + + mu.Lock() + defer mu.Unlock() + + seen := make(map[string]bool) + for _, id := range batchIDs { + if seen[id] { + t.Errorf("duplicate batch ID: %s", id) + } + seen[id] = true + } +} + +func TestBatcher_EmptyFlushIsNoop(t *testing.T) { + var flushCount atomic.Int32 + + sendFn := func(batchID string, messages []json.RawMessage) error { + flushCount.Add(1) + return nil + } + + b := NewBatcher(sendFn) + + // Flush with nothing in the buffer should not call sendFn. + b.flush() + + if flushCount.Load() != 0 { + t.Errorf("flush count = %d, want 0 (empty flush should be noop)", flushCount.Load()) + } +} + +func TestBatcher_ContextCancellationStopsRun(t *testing.T) { + sendFn, _ := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + b.Run(ctx) + close(done) + }() + + cancel() + + select { + case <-done: + // Run exited. + case <-time.After(2 * time.Second): + t.Fatal("Run() did not exit after context cancellation") + } +} + +func TestBatcher_MessageOrderPreserved(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + // Enqueue messages in a specific order, then trigger flush with an immediate message. + b.Enqueue(json.RawMessage(`{"id":"1","type":"project_state"}`), "project_state") + b.Enqueue(json.RawMessage(`{"id":"2","type":"claude_output"}`), "claude_output") + b.Enqueue(json.RawMessage(`{"id":"3","type":"error"}`), "error") // triggers flush + + time.Sleep(50 * time.Millisecond) + + records := getRecords() + if len(records) != 1 { + t.Fatalf("flush count = %d, want 1", len(records)) + } + + msgs := records[0].messages + if len(msgs) != 3 { + t.Fatalf("message count = %d, want 3", len(msgs)) + } + + // Verify order is preserved. + expected := []string{`{"id":"1","type":"project_state"}`, `{"id":"2","type":"claude_output"}`, `{"id":"3","type":"error"}`} + for i, msg := range msgs { + if string(msg) != expected[i] { + t.Errorf("message[%d] = %s, want %s", i, msg, expected[i]) + } + } +} + +func TestBatcher_SendErrorDoesNotLoseMessages(t *testing.T) { + // When send fails, messages are already removed from the buffer. + // This is by design — the caller (SendMessagesWithRetry) handles retries. + var flushCount atomic.Int32 + + sendFn := func(batchID string, messages []json.RawMessage) error { + flushCount.Add(1) + return nil + } + + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + b.Enqueue(json.RawMessage(`{"type":"error"}`), "error") + time.Sleep(50 * time.Millisecond) + + if flushCount.Load() != 1 { + t.Errorf("flush count = %d, want 1", flushCount.Load()) + } + + // Buffer should be empty after flush. + b.mu.Lock() + remaining := len(b.messages) + b.mu.Unlock() + + if remaining != 0 { + t.Errorf("remaining messages = %d, want 0", remaining) + } +} + +func TestTierFor(t *testing.T) { + tests := []struct { + msgType string + want tier + }{ + // Immediate tier. + {"run_complete", tierIDImmediate}, + {"run_paused", tierIDImmediate}, + {"error", tierIDImmediate}, + {"clone_complete", tierIDImmediate}, + {"session_expired", tierIDImmediate}, + {"quota_exhausted", tierIDImmediate}, + {"prd_response_complete", tierIDImmediate}, + // Standard tier. + {"claude_output", tierIDStandard}, + {"prd_output", tierIDStandard}, + {"run_progress", tierIDStandard}, + {"clone_progress", tierIDStandard}, + // Low priority tier. + {"state_snapshot", tierIDLowPriority}, + {"project_state", tierIDLowPriority}, + {"project_list", tierIDLowPriority}, + {"settings", tierIDLowPriority}, + {"log_lines", tierIDLowPriority}, + // Unknown defaults to standard. + {"unknown_type", tierIDStandard}, + } + + for _, tt := range tests { + t.Run(tt.msgType, func(t *testing.T) { + got := tierFor(tt.msgType) + if got != tt.want { + t.Errorf("tierFor(%q) = %d, want %d", tt.msgType, got, tt.want) + } + }) + } +} + +func TestGenerateBatchID(t *testing.T) { + id1 := generateBatchID() + id2 := generateBatchID() + + if id1 == "" { + t.Error("batch ID should not be empty") + } + if id1 == id2 { + t.Errorf("batch IDs should be unique: %q == %q", id1, id2) + } + + // Verify UUID v4 format (8-4-4-4-12 hex chars). + if len(id1) != 36 { + t.Errorf("batch ID length = %d, want 36", len(id1)) + } +} + +func TestBatcher_ConcurrentEnqueue(t *testing.T) { + sendFn, getRecords := newRecordingSendFn() + b := NewBatcher(sendFn) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go b.Run(ctx) + + // Enqueue from multiple goroutines concurrently. + var wg sync.WaitGroup + for i := 0; i < 50; i++ { + wg.Add(1) + go func() { + defer wg.Done() + b.Enqueue(json.RawMessage(`{"type":"claude_output"}`), "claude_output") + }() + } + wg.Wait() + + // Wait for all timer-based flushes. + time.Sleep(500 * time.Millisecond) + + records := getRecords() + total := 0 + for _, r := range records { + total += len(r.messages) + } + + if total != 50 { + t.Errorf("total flushed messages = %d, want 50", total) + } +} diff --git a/internal/uplink/client.go b/internal/uplink/client.go new file mode 100644 index 00000000..3433f91c --- /dev/null +++ b/internal/uplink/client.go @@ -0,0 +1,356 @@ +package uplink + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "math" + "math/rand/v2" + "net/http" + "net/url" + "runtime" + "sync" + "time" + + "github.com/minicodemonkey/chief/internal/ws" +) + +const ( + // maxBackoff is the maximum reconnection delay. + maxBackoff = 60 * time.Second + + // initialBackoff is the starting reconnection delay. + initialBackoff = 1 * time.Second + + // httpTimeout is the default HTTP request timeout. + httpTimeout = 10 * time.Second +) + +// WelcomeResponse is the response from POST /api/device/connect. +type WelcomeResponse struct { + Type string `json:"type"` + ProtocolVersion int `json:"protocol_version"` + DeviceID int `json:"device_id"` + SessionID string `json:"session_id"` + Reverb ReverbConfig `json:"reverb"` +} + +// ReverbConfig contains Pusher/Reverb connection details from the connect response. +type ReverbConfig struct { + Key string `json:"key"` + Host string `json:"host"` + Port int `json:"port"` + Scheme string `json:"scheme"` +} + +// connectRequest is the JSON body sent to POST /api/device/connect. +type connectRequest struct { + ChiefVersion string `json:"chief_version"` + DeviceName string `json:"device_name"` + OS string `json:"os"` + Arch string `json:"arch"` + ProtocolVersion int `json:"protocol_version"` +} + +// errorResponse is a JSON error returned by the server. +type errorResponse struct { + Error string `json:"error"` + Code string `json:"code,omitempty"` + Message string `json:"message,omitempty"` +} + +// ErrAuthFailed is returned when the server rejects authentication (401). +var ErrAuthFailed = fmt.Errorf("device deauthorized — run 'chief login' to re-authenticate") + +// ErrDeviceRevoked is returned when the device is revoked (403). +var ErrDeviceRevoked = fmt.Errorf("device revoked — run 'chief login' to re-authenticate") + +// Client is an HTTP client for the uplink device API. +type Client struct { + baseURL string + accessToken string + mu sync.RWMutex + httpClient *http.Client + + // Device metadata sent on connect. + chiefVersion string + deviceName string +} + +// Option configures a Client. +type Option func(*Client) + +// WithChiefVersion sets the chief CLI version string. +func WithChiefVersion(v string) Option { + return func(c *Client) { + c.chiefVersion = v + } +} + +// WithDeviceName sets the device name. +func WithDeviceName(name string) Option { + return func(c *Client) { + c.deviceName = name + } +} + +// WithHTTPClient sets a custom http.Client (useful for testing). +func WithHTTPClient(hc *http.Client) Option { + return func(c *Client) { + c.httpClient = hc + } +} + +// New creates a new uplink HTTP client. +// The baseURL must use HTTPS unless the host is localhost or 127.0.0.1. +func New(baseURL, accessToken string, opts ...Option) (*Client, error) { + if err := validateBaseURL(baseURL); err != nil { + return nil, err + } + + c := &Client{ + baseURL: baseURL, + accessToken: accessToken, + httpClient: &http.Client{Timeout: httpTimeout}, + } + for _, o := range opts { + o(c) + } + return c, nil +} + +// validateBaseURL ensures the URL uses HTTPS unless the host is localhost/127.0.0.1. +func validateBaseURL(rawURL string) error { + u, err := url.Parse(rawURL) + if err != nil { + return fmt.Errorf("invalid base URL: %w", err) + } + + host := u.Hostname() + if u.Scheme == "http" && host != "localhost" && host != "127.0.0.1" { + return fmt.Errorf("base URL must use HTTPS (got %s); HTTP is only allowed for localhost", rawURL) + } + if u.Scheme != "http" && u.Scheme != "https" { + return fmt.Errorf("base URL must use http or https scheme (got %s)", u.Scheme) + } + + return nil +} + +// SetAccessToken updates the access token in a thread-safe manner. +// This is called after a token refresh. +func (c *Client) SetAccessToken(token string) { + c.mu.Lock() + defer c.mu.Unlock() + c.accessToken = token +} + +// Connect calls POST /api/device/connect to register the device with the server. +// Returns the welcome response containing session ID and Reverb configuration. +func (c *Client) Connect(ctx context.Context) (*WelcomeResponse, error) { + version := c.chiefVersion + if version == "" { + version = "dev" + } + + body := connectRequest{ + ChiefVersion: version, + DeviceName: c.deviceName, + OS: runtime.GOOS, + Arch: runtime.GOARCH, + ProtocolVersion: ws.ProtocolVersion, + } + + var welcome WelcomeResponse + if err := c.doJSON(ctx, "POST", "/api/device/connect", body, &welcome); err != nil { + return nil, fmt.Errorf("connect: %w", err) + } + + return &welcome, nil +} + +// IngestResponse is the response from POST /api/device/messages. +type IngestResponse struct { + Accepted int `json:"accepted"` + BatchID string `json:"batch_id"` + SessionID string `json:"session_id"` +} + +// ingestRequest is the JSON body sent to POST /api/device/messages. +type ingestRequest struct { + BatchID string `json:"batch_id"` + Messages []json.RawMessage `json:"messages"` +} + +// SendMessages sends a batch of messages via POST /api/device/messages. +// It does NOT retry on failure — use SendMessagesWithRetry for retry behavior. +func (c *Client) SendMessages(ctx context.Context, batchID string, messages []json.RawMessage) (*IngestResponse, error) { + body := ingestRequest{ + BatchID: batchID, + Messages: messages, + } + + var resp IngestResponse + if err := c.doJSON(ctx, "POST", "/api/device/messages", body, &resp); err != nil { + return nil, fmt.Errorf("send messages: %w", err) + } + + return &resp, nil +} + +// SendMessagesWithRetry sends a message batch with exponential backoff retry on transient failures. +// It does not retry on 401/403 auth errors. Retries use the same batchID for server-side deduplication. +func (c *Client) SendMessagesWithRetry(ctx context.Context, batchID string, messages []json.RawMessage) (*IngestResponse, error) { + attempt := 0 + for { + resp, err := c.SendMessages(ctx, batchID, messages) + if err == nil { + return resp, nil + } + + // Don't retry auth errors. + if isAuthError(err) { + return nil, err + } + + attempt++ + delay := backoff(attempt) + log.Printf("SendMessages failed (attempt %d, batch %s): %v — retrying in %s", attempt, batchID, err, delay.Round(time.Millisecond)) + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-time.After(delay): + } + } +} + +// Heartbeat calls POST /api/device/heartbeat to tell the server the device is alive. +func (c *Client) Heartbeat(ctx context.Context) error { + var resp json.RawMessage + if err := c.doJSON(ctx, "POST", "/api/device/heartbeat", nil, &resp); err != nil { + return fmt.Errorf("heartbeat: %w", err) + } + return nil +} + +// Disconnect calls POST /api/device/disconnect to notify the server the device is going offline. +func (c *Client) Disconnect(ctx context.Context) error { + var resp json.RawMessage + if err := c.doJSON(ctx, "POST", "/api/device/disconnect", nil, &resp); err != nil { + return fmt.Errorf("disconnect: %w", err) + } + return nil +} + +// doJSON performs an HTTP request with JSON body and parses the JSON response. +// It handles auth headers and classifies HTTP error responses. +func (c *Client) doJSON(ctx context.Context, method, path string, body interface{}, result interface{}) error { + var bodyReader io.Reader + if body != nil { + data, err := json.Marshal(body) + if err != nil { + return fmt.Errorf("marshaling request: %w", err) + } + bodyReader = bytes.NewReader(data) + } + + req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, bodyReader) + if err != nil { + return fmt.Errorf("creating request: %w", err) + } + + c.mu.RLock() + token := c.accessToken + c.mu.RUnlock() + + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + + resp, err := c.httpClient.Do(req) + if err != nil { + return fmt.Errorf("sending request: %w", err) + } + defer resp.Body.Close() + + // Read the full response body (capped to prevent OOM on rogue responses). + respBody, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) // 1MB max + if err != nil { + return fmt.Errorf("reading response: %w", err) + } + + // Classify HTTP errors. + if resp.StatusCode == http.StatusUnauthorized { + return ErrAuthFailed + } + if resp.StatusCode == http.StatusForbidden { + return ErrDeviceRevoked + } + if resp.StatusCode >= 400 { + var errResp errorResponse + if json.Unmarshal(respBody, &errResp) == nil && errResp.Message != "" { + return fmt.Errorf("server error %d: %s", resp.StatusCode, errResp.Message) + } + return fmt.Errorf("server error %d: %s", resp.StatusCode, string(respBody)) + } + + if result != nil && len(respBody) > 0 { + if err := json.Unmarshal(respBody, result); err != nil { + return fmt.Errorf("parsing response: %w", err) + } + } + + return nil +} + +// backoff returns a duration for the given attempt using exponential backoff + jitter. +func backoff(attempt int) time.Duration { + base := float64(initialBackoff) * math.Pow(2, float64(attempt-1)) + if base > float64(maxBackoff) { + base = float64(maxBackoff) + } + // Add jitter: 0.5x to 1.5x + jitter := 0.5 + rand.Float64() + return time.Duration(base * jitter) +} + +// ConnectWithRetry calls Connect with exponential backoff retry on transient failures. +// It does not retry on 401/403 auth errors. +func (c *Client) ConnectWithRetry(ctx context.Context) (*WelcomeResponse, error) { + attempt := 0 + for { + welcome, err := c.Connect(ctx) + if err == nil { + return welcome, nil + } + + // Don't retry auth errors. + if isAuthError(err) { + return nil, err + } + + attempt++ + delay := backoff(attempt) + log.Printf("Connect failed (attempt %d): %v — retrying in %s", attempt, err, delay.Round(time.Millisecond)) + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-time.After(delay): + } + } +} + +// isAuthError returns true if the error is a 401 or 403 that should not be retried. +// Uses errors.Is to handle wrapped errors. +func isAuthError(err error) bool { + if err == nil { + return false + } + return errors.Is(err, ErrAuthFailed) || errors.Is(err, ErrDeviceRevoked) +} diff --git a/internal/uplink/client_test.go b/internal/uplink/client_test.go new file mode 100644 index 00000000..5f118223 --- /dev/null +++ b/internal/uplink/client_test.go @@ -0,0 +1,743 @@ +package uplink + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "runtime" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/ws" +) + +// testContext returns a context with a 15-second timeout for tests. +func testContext(t *testing.T) context.Context { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + t.Cleanup(cancel) + return ctx +} + +// newTestClient creates a Client pointing at a test server with the given token. +func newTestClient(t *testing.T, serverURL, token string, opts ...Option) *Client { + t.Helper() + c, err := New(serverURL, token, opts...) + if err != nil { + t.Fatalf("New() failed: %v", err) + } + return c +} + +func TestNew_ValidHTTPS(t *testing.T) { + c, err := New("https://example.com", "token123") + if err != nil { + t.Fatalf("New() with HTTPS failed: %v", err) + } + if c.baseURL != "https://example.com" { + t.Errorf("baseURL = %q, want %q", c.baseURL, "https://example.com") + } +} + +func TestNew_LocalhostHTTP(t *testing.T) { + _, err := New("http://localhost:8080", "token123") + if err != nil { + t.Fatalf("New() with localhost HTTP failed: %v", err) + } +} + +func TestNew_Loopback127HTTP(t *testing.T) { + _, err := New("http://127.0.0.1:8080", "token123") + if err != nil { + t.Fatalf("New() with 127.0.0.1 HTTP failed: %v", err) + } +} + +func TestNew_RejectsNonLocalhostHTTP(t *testing.T) { + _, err := New("http://example.com", "token123") + if err == nil { + t.Fatal("expected error for non-localhost HTTP, got nil") + } +} + +func TestNew_RejectsInvalidScheme(t *testing.T) { + _, err := New("ftp://example.com", "token123") + if err == nil { + t.Fatal("expected error for ftp scheme, got nil") + } +} + +func TestNew_WithOptions(t *testing.T) { + c, err := New("https://example.com", "token123", + WithChiefVersion("1.2.3"), + WithDeviceName("my-device"), + ) + if err != nil { + t.Fatalf("New() failed: %v", err) + } + if c.chiefVersion != "1.2.3" { + t.Errorf("chiefVersion = %q, want %q", c.chiefVersion, "1.2.3") + } + if c.deviceName != "my-device" { + t.Errorf("deviceName = %q, want %q", c.deviceName, "my-device") + } +} + +func TestConnect_Success(t *testing.T) { + var receivedBody connectRequest + var receivedAuth string + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/device/connect" { + http.NotFound(w, r) + return + } + if r.Method != "POST" { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + + receivedAuth = r.Header.Get("Authorization") + + body, _ := io.ReadAll(r.Body) + json.Unmarshal(body, &receivedBody) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(WelcomeResponse{ + Type: "welcome", + ProtocolVersion: 1, + DeviceID: 42, + SessionID: "sess-abc-123", + Reverb: ReverbConfig{ + Key: "app-key", + Host: "reverb.example.com", + Port: 443, + Scheme: "https", + }, + }) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "test-token-abc", + WithChiefVersion("2.0.0"), + WithDeviceName("test-device"), + ) + + ctx := testContext(t) + welcome, err := client.Connect(ctx) + if err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Verify request. + if receivedAuth != "Bearer test-token-abc" { + t.Errorf("Authorization = %q, want %q", receivedAuth, "Bearer test-token-abc") + } + if receivedBody.ChiefVersion != "2.0.0" { + t.Errorf("chief_version = %q, want %q", receivedBody.ChiefVersion, "2.0.0") + } + if receivedBody.DeviceName != "test-device" { + t.Errorf("device_name = %q, want %q", receivedBody.DeviceName, "test-device") + } + if receivedBody.OS != runtime.GOOS { + t.Errorf("os = %q, want %q", receivedBody.OS, runtime.GOOS) + } + if receivedBody.Arch != runtime.GOARCH { + t.Errorf("arch = %q, want %q", receivedBody.Arch, runtime.GOARCH) + } + if receivedBody.ProtocolVersion != ws.ProtocolVersion { + t.Errorf("protocol_version = %d, want %d", receivedBody.ProtocolVersion, ws.ProtocolVersion) + } + + // Verify response. + if welcome.Type != "welcome" { + t.Errorf("Type = %q, want %q", welcome.Type, "welcome") + } + if welcome.DeviceID != 42 { + t.Errorf("DeviceID = %d, want %d", welcome.DeviceID, 42) + } + if welcome.SessionID != "sess-abc-123" { + t.Errorf("SessionID = %q, want %q", welcome.SessionID, "sess-abc-123") + } + if welcome.Reverb.Key != "app-key" { + t.Errorf("Reverb.Key = %q, want %q", welcome.Reverb.Key, "app-key") + } + if welcome.Reverb.Host != "reverb.example.com" { + t.Errorf("Reverb.Host = %q, want %q", welcome.Reverb.Host, "reverb.example.com") + } + if welcome.Reverb.Port != 443 { + t.Errorf("Reverb.Port = %d, want %d", welcome.Reverb.Port, 443) + } + if welcome.Reverb.Scheme != "https" { + t.Errorf("Reverb.Scheme = %q, want %q", welcome.Reverb.Scheme, "https") + } +} + +func TestConnect_DefaultVersion(t *testing.T) { + var receivedBody connectRequest + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + json.Unmarshal(body, &receivedBody) + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(WelcomeResponse{Type: "welcome", DeviceID: 1, SessionID: "s"}) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token") + ctx := testContext(t) + _, err := client.Connect(ctx) + if err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + if receivedBody.ChiefVersion != "dev" { + t.Errorf("chief_version = %q, want %q (default)", receivedBody.ChiefVersion, "dev") + } +} + +func TestConnect_AuthFailed401(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + json.NewEncoder(w).Encode(errorResponse{Error: "invalid_token", Message: "Invalid token"}) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "bad-token") + ctx := testContext(t) + _, err := client.Connect(ctx) + if err == nil { + t.Fatal("expected error for 401, got nil") + } + if !errors.Is(err, ErrAuthFailed) { + t.Errorf("error = %v, want ErrAuthFailed", err) + } +} + +func TestConnect_DeviceRevoked403(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusForbidden) + json.NewEncoder(w).Encode(errorResponse{Error: "device_revoked", Message: "Device revoked"}) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "revoked-token") + ctx := testContext(t) + _, err := client.Connect(ctx) + if err == nil { + t.Fatal("expected error for 403, got nil") + } + if !errors.Is(err, ErrDeviceRevoked) { + t.Errorf("error = %v, want ErrDeviceRevoked", err) + } +} + +func TestConnect_ServerError5xx(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusInternalServerError) + json.NewEncoder(w).Encode(errorResponse{Message: "something went wrong"}) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token") + ctx := testContext(t) + _, err := client.Connect(ctx) + if err == nil { + t.Fatal("expected error for 500, got nil") + } + if isAuthError(err) { + t.Error("5xx error should not be classified as auth error") + } +} + +func TestDisconnect_Success(t *testing.T) { + var receivedMethod string + var receivedPath string + var receivedAuth string + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + receivedMethod = r.Method + receivedPath = r.URL.Path + receivedAuth = r.Header.Get("Authorization") + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]string{"status": "disconnected"}) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "test-token") + ctx := testContext(t) + err := client.Disconnect(ctx) + if err != nil { + t.Fatalf("Disconnect() failed: %v", err) + } + + if receivedMethod != "POST" { + t.Errorf("method = %q, want POST", receivedMethod) + } + if receivedPath != "/api/device/disconnect" { + t.Errorf("path = %q, want /api/device/disconnect", receivedPath) + } + if receivedAuth != "Bearer test-token" { + t.Errorf("Authorization = %q, want %q", receivedAuth, "Bearer test-token") + } +} + +func TestDisconnect_AuthFailed(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "bad-token") + ctx := testContext(t) + err := client.Disconnect(ctx) + if err == nil { + t.Fatal("expected error for 401, got nil") + } +} + +func TestSetAccessToken_ThreadSafe(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Return the received token in the response body for verification. + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]string{ + "token": r.Header.Get("Authorization"), + }) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token-v1") + + // Spawn goroutines that concurrently update the token and make requests. + var wg sync.WaitGroup + for i := 0; i < 10; i++ { + wg.Add(1) + go func() { + defer wg.Done() + client.SetAccessToken("token-v2") + }() + } + wg.Wait() + + // After all updates, the token should be v2. + client.mu.RLock() + token := client.accessToken + client.mu.RUnlock() + + if token != "token-v2" { + t.Errorf("accessToken = %q, want %q", token, "token-v2") + } +} + +func TestConnect_RequestFormat(t *testing.T) { + var receivedContentType string + var receivedAccept string + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + receivedContentType = r.Header.Get("Content-Type") + receivedAccept = r.Header.Get("Accept") + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(WelcomeResponse{Type: "welcome", DeviceID: 1, SessionID: "s"}) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token") + ctx := testContext(t) + client.Connect(ctx) + + if receivedContentType != "application/json" { + t.Errorf("Content-Type = %q, want application/json", receivedContentType) + } + if receivedAccept != "application/json" { + t.Errorf("Accept = %q, want application/json", receivedAccept) + } +} + +func TestConnect_ContextCancellation(t *testing.T) { + blocked := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Block until the test is done (context cancelled will abort the request). + <-blocked + })) + defer srv.Close() + defer close(blocked) // unblock the handler so the server can shut down cleanly + + client := newTestClient(t, srv.URL, "token") + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + + _, err := client.Connect(ctx) + if err == nil { + t.Fatal("expected error from cancelled context, got nil") + } +} + +func TestConnectWithRetry_SuccessOnSecondAttempt(t *testing.T) { + var attempt atomic.Int32 + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + n := attempt.Add(1) + if n == 1 { + w.WriteHeader(http.StatusServiceUnavailable) + return + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(WelcomeResponse{ + Type: "welcome", + DeviceID: 42, + SessionID: "sess-123", + }) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token") + ctx := testContext(t) + welcome, err := client.ConnectWithRetry(ctx) + if err != nil { + t.Fatalf("ConnectWithRetry() failed: %v", err) + } + if welcome.DeviceID != 42 { + t.Errorf("DeviceID = %d, want 42", welcome.DeviceID) + } + if attempt.Load() != 2 { + t.Errorf("attempts = %d, want 2", attempt.Load()) + } +} + +func TestConnectWithRetry_NoRetryOnAuthError(t *testing.T) { + var attempt atomic.Int32 + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempt.Add(1) + w.WriteHeader(http.StatusUnauthorized) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "bad-token") + ctx := testContext(t) + _, err := client.ConnectWithRetry(ctx) + if err == nil { + t.Fatal("expected error, got nil") + } + if attempt.Load() != 1 { + t.Errorf("attempts = %d, want 1 (no retry on auth error)", attempt.Load()) + } +} + +func TestBackoff(t *testing.T) { + tests := []struct { + attempt int + minMs int64 + maxMs int64 + }{ + {1, 500, 1500}, // 1s * (0.5 to 1.5) + {2, 1000, 3000}, // 2s * (0.5 to 1.5) + {3, 2000, 6000}, // 4s * (0.5 to 1.5) + {4, 4000, 12000}, // 8s * (0.5 to 1.5) + {10, 30000, 90000}, // capped at 60s * (0.5 to 1.5) + } + + for _, tt := range tests { + d := backoff(tt.attempt) + ms := d.Milliseconds() + if ms < tt.minMs || ms > tt.maxMs { + t.Errorf("backoff(%d) = %dms, want [%d, %d]ms", tt.attempt, ms, tt.minMs, tt.maxMs) + } + } +} + +func TestSendMessages_Success(t *testing.T) { + var receivedBody ingestRequest + var receivedAuth string + var receivedMethod string + var receivedPath string + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + receivedMethod = r.Method + receivedPath = r.URL.Path + receivedAuth = r.Header.Get("Authorization") + + body, _ := io.ReadAll(r.Body) + json.Unmarshal(body, &receivedBody) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(IngestResponse{ + Accepted: 2, + BatchID: "batch-abc-123", + SessionID: "sess-xyz-789", + }) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "test-token") + ctx := testContext(t) + + messages := []json.RawMessage{ + json.RawMessage(`{"type":"project_state","id":"m1"}`), + json.RawMessage(`{"type":"claude_output","id":"m2"}`), + } + + resp, err := client.SendMessages(ctx, "batch-abc-123", messages) + if err != nil { + t.Fatalf("SendMessages() failed: %v", err) + } + + // Verify request format. + if receivedMethod != "POST" { + t.Errorf("method = %q, want POST", receivedMethod) + } + if receivedPath != "/api/device/messages" { + t.Errorf("path = %q, want /api/device/messages", receivedPath) + } + if receivedAuth != "Bearer test-token" { + t.Errorf("Authorization = %q, want %q", receivedAuth, "Bearer test-token") + } + if receivedBody.BatchID != "batch-abc-123" { + t.Errorf("batch_id = %q, want %q", receivedBody.BatchID, "batch-abc-123") + } + if len(receivedBody.Messages) != 2 { + t.Errorf("messages count = %d, want 2", len(receivedBody.Messages)) + } + + // Verify response parsing. + if resp.Accepted != 2 { + t.Errorf("Accepted = %d, want 2", resp.Accepted) + } + if resp.BatchID != "batch-abc-123" { + t.Errorf("BatchID = %q, want %q", resp.BatchID, "batch-abc-123") + } + if resp.SessionID != "sess-xyz-789" { + t.Errorf("SessionID = %q, want %q", resp.SessionID, "sess-xyz-789") + } +} + +func TestSendMessages_AuthFailed401(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "bad-token") + ctx := testContext(t) + + _, err := client.SendMessages(ctx, "batch-1", []json.RawMessage{json.RawMessage(`{}`)}) + if err == nil { + t.Fatal("expected error for 401, got nil") + } + if !errors.Is(err, ErrAuthFailed) { + t.Errorf("error = %v, want ErrAuthFailed", err) + } +} + +func TestSendMessages_DeviceRevoked403(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusForbidden) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "revoked-token") + ctx := testContext(t) + + _, err := client.SendMessages(ctx, "batch-1", []json.RawMessage{json.RawMessage(`{}`)}) + if err == nil { + t.Fatal("expected error for 403, got nil") + } + if !errors.Is(err, ErrDeviceRevoked) { + t.Errorf("error = %v, want ErrDeviceRevoked", err) + } +} + +func TestSendMessages_ServerError5xx(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token") + ctx := testContext(t) + + _, err := client.SendMessages(ctx, "batch-1", []json.RawMessage{json.RawMessage(`{}`)}) + if err == nil { + t.Fatal("expected error for 500, got nil") + } + if isAuthError(err) { + t.Error("5xx error should not be classified as auth error") + } +} + +func TestSendMessagesWithRetry_SuccessAfterRetry(t *testing.T) { + var attempt atomic.Int32 + var receivedBatchIDs []string + var mu sync.Mutex + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var body ingestRequest + data, _ := io.ReadAll(r.Body) + json.Unmarshal(data, &body) + + mu.Lock() + receivedBatchIDs = append(receivedBatchIDs, body.BatchID) + mu.Unlock() + + n := attempt.Add(1) + if n <= 2 { + w.WriteHeader(http.StatusServiceUnavailable) + return + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(IngestResponse{ + Accepted: 1, + BatchID: body.BatchID, + SessionID: "sess-1", + }) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token") + ctx := testContext(t) + + resp, err := client.SendMessagesWithRetry(ctx, "batch-retry-123", []json.RawMessage{json.RawMessage(`{"type":"log_lines"}`)}) + if err != nil { + t.Fatalf("SendMessagesWithRetry() failed: %v", err) + } + + if resp.Accepted != 1 { + t.Errorf("Accepted = %d, want 1", resp.Accepted) + } + if attempt.Load() != 3 { + t.Errorf("attempts = %d, want 3", attempt.Load()) + } + + // Verify same batch_id was used on all retries (for server-side deduplication). + mu.Lock() + defer mu.Unlock() + for i, id := range receivedBatchIDs { + if id != "batch-retry-123" { + t.Errorf("attempt %d batch_id = %q, want %q", i+1, id, "batch-retry-123") + } + } +} + +func TestSendMessagesWithRetry_NoRetryOnAuthError(t *testing.T) { + var attempt atomic.Int32 + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempt.Add(1) + w.WriteHeader(http.StatusUnauthorized) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "bad-token") + ctx := testContext(t) + + _, err := client.SendMessagesWithRetry(ctx, "batch-1", []json.RawMessage{json.RawMessage(`{}`)}) + if err == nil { + t.Fatal("expected error, got nil") + } + if !errors.Is(err, ErrAuthFailed) { + t.Errorf("error = %v, want ErrAuthFailed", err) + } + if attempt.Load() != 1 { + t.Errorf("attempts = %d, want 1 (no retry on auth error)", attempt.Load()) + } +} + +func TestSendMessagesWithRetry_NoRetryOnRevoked(t *testing.T) { + var attempt atomic.Int32 + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempt.Add(1) + w.WriteHeader(http.StatusForbidden) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "revoked-token") + ctx := testContext(t) + + _, err := client.SendMessagesWithRetry(ctx, "batch-1", []json.RawMessage{json.RawMessage(`{}`)}) + if err == nil { + t.Fatal("expected error, got nil") + } + if !errors.Is(err, ErrDeviceRevoked) { + t.Errorf("error = %v, want ErrDeviceRevoked", err) + } + if attempt.Load() != 1 { + t.Errorf("attempts = %d, want 1 (no retry on revoked)", attempt.Load()) + } +} + +func TestSendMessagesWithRetry_ContextCancellation(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusServiceUnavailable) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token") + ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) + defer cancel() + + _, err := client.SendMessagesWithRetry(ctx, "batch-1", []json.RawMessage{json.RawMessage(`{}`)}) + if err == nil { + t.Fatal("expected error from cancelled context, got nil") + } +} + +func TestSendMessages_RequestBodyFormat(t *testing.T) { + var receivedRaw json.RawMessage + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + data, _ := io.ReadAll(r.Body) + receivedRaw = data + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(IngestResponse{Accepted: 1, BatchID: "b1", SessionID: "s1"}) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "token") + ctx := testContext(t) + + messages := []json.RawMessage{ + json.RawMessage(`{"type":"project_state","id":"msg-1","timestamp":"2026-02-16T00:00:00Z"}`), + } + _, err := client.SendMessages(ctx, "batch-format-test", messages) + if err != nil { + t.Fatalf("SendMessages() failed: %v", err) + } + + // Verify the raw JSON structure matches what the server expects. + var parsed map[string]json.RawMessage + if err := json.Unmarshal(receivedRaw, &parsed); err != nil { + t.Fatalf("failed to parse request body: %v", err) + } + if _, ok := parsed["batch_id"]; !ok { + t.Error("request body missing batch_id field") + } + if _, ok := parsed["messages"]; !ok { + t.Error("request body missing messages field") + } +} + +func TestIsAuthError(t *testing.T) { + if isAuthError(nil) { + t.Error("nil should not be auth error") + } + if !isAuthError(ErrAuthFailed) { + t.Error("ErrAuthFailed should be auth error") + } + if !isAuthError(ErrDeviceRevoked) { + t.Error("ErrDeviceRevoked should be auth error") + } + if isAuthError(context.Canceled) { + t.Error("context.Canceled should not be auth error") + } +} diff --git a/internal/uplink/pusher.go b/internal/uplink/pusher.go new file mode 100644 index 00000000..ec33d51d --- /dev/null +++ b/internal/uplink/pusher.go @@ -0,0 +1,472 @@ +package uplink + +import ( + "context" + "crypto/hmac" + "crypto/sha256" + "encoding/json" + "fmt" + "log" + "net/http" + "net/url" + "sync" + "time" + + "github.com/gorilla/websocket" +) + +const ( + // pusherProtocolVersion is the Pusher protocol version to use. + pusherProtocolVersion = 7 + + // receiveBufSize is the buffer size for the receive channel. + receiveBufSize = 256 + + // pusherPingTimeout is how long to wait for a pong after sending a ping. + pusherPingTimeout = 30 * time.Second + + // pusherWriteTimeout is the timeout for WebSocket write operations. + pusherWriteTimeout = 10 * time.Second +) + +// pusherMessage is a Pusher protocol message (both sent and received). +type pusherMessage struct { + Event string `json:"event"` + Channel string `json:"channel,omitempty"` + Data json.RawMessage `json:"data"` +} + +// pusherConnectionData is the data field of pusher:connection_established. +type pusherConnectionData struct { + SocketID string `json:"socket_id"` + ActivityTimeout int `json:"activity_timeout"` +} + +// pusherAuthResponse is the response from the broadcast auth endpoint. +type pusherAuthResponse struct { + Auth string `json:"auth"` +} + +// AuthFunc is a function that authenticates a Pusher channel subscription. +// It takes a socketID and channelName and returns the auth signature string. +type AuthFunc func(ctx context.Context, socketID, channelName string) (string, error) + +// PusherClient connects to a Reverb/Pusher WebSocket and subscribes to a private channel. +type PusherClient struct { + appKey string + host string + port int + scheme string + channel string + authFn AuthFunc + dialer *websocket.Dialer + + mu sync.Mutex + conn *websocket.Conn + socketID string + recvCh chan json.RawMessage + done chan struct{} + stopped bool +} + +// NewPusherClient creates a PusherClient configured to connect to Reverb. +// +// Parameters: +// - cfg: Reverb connection config (key, host, port, scheme) from the connect response +// - channel: the private channel to subscribe to (e.g., "private-chief-server.42") +// - authFn: function to authenticate the channel subscription +func NewPusherClient(cfg ReverbConfig, channel string, authFn AuthFunc) *PusherClient { + return &PusherClient{ + appKey: cfg.Key, + host: cfg.Host, + port: cfg.Port, + scheme: cfg.Scheme, + channel: channel, + authFn: authFn, + dialer: websocket.DefaultDialer, + recvCh: make(chan json.RawMessage, receiveBufSize), + done: make(chan struct{}), + } +} + +// Connect dials the Pusher WebSocket, waits for connection_established, +// subscribes to the private channel, and starts the read loop. +func (p *PusherClient) Connect(ctx context.Context) error { + wsURL := p.buildURL() + + headers := http.Header{} + headers.Set("Origin", fmt.Sprintf("%s://%s", p.scheme, p.host)) + + conn, _, err := p.dialer.DialContext(ctx, wsURL, headers) + if err != nil { + return fmt.Errorf("pusher dial: %w", err) + } + + p.mu.Lock() + p.conn = conn + p.mu.Unlock() + + // Wait for pusher:connection_established. + socketID, activityTimeout, err := p.waitForConnectionEstablished(ctx, conn) + if err != nil { + conn.Close() + return err + } + + p.mu.Lock() + p.socketID = socketID + p.mu.Unlock() + + // Subscribe to the private channel. + if err := p.subscribe(ctx, conn, socketID); err != nil { + conn.Close() + return err + } + + // Start read loop. + go p.readLoop(ctx, conn, activityTimeout) + + return nil +} + +// Receive returns a channel that delivers incoming command payloads. +// The channel is closed when the client shuts down. +func (p *PusherClient) Receive() <-chan json.RawMessage { + return p.recvCh +} + +// Close gracefully shuts down the Pusher client. +func (p *PusherClient) Close() error { + p.mu.Lock() + if p.stopped { + p.mu.Unlock() + return nil + } + p.stopped = true + conn := p.conn + p.conn = nil + p.mu.Unlock() + + var err error + if conn != nil { + deadline := time.Now().Add(5 * time.Second) + closeMsg := websocket.FormatCloseMessage(websocket.CloseNormalClosure, "") + _ = conn.WriteControl(websocket.CloseMessage, closeMsg, deadline) + err = conn.Close() + } + + // Wait for readLoop to finish. + <-p.done + + return err +} + +// buildURL constructs the Pusher WebSocket URL. +func (p *PusherClient) buildURL() string { + wsScheme := "wss" + if p.scheme == "http" { + wsScheme = "ws" + } + + u := url.URL{ + Scheme: wsScheme, + Host: fmt.Sprintf("%s:%d", p.host, p.port), + Path: fmt.Sprintf("/app/%s", p.appKey), + RawQuery: fmt.Sprintf("protocol=%d", pusherProtocolVersion), + } + return u.String() +} + +// waitForConnectionEstablished reads messages until it receives +// pusher:connection_established. Returns the socket ID and activity timeout. +func (p *PusherClient) waitForConnectionEstablished(ctx context.Context, conn *websocket.Conn) (string, int, error) { + // Set a read deadline for the connection established message. + conn.SetReadDeadline(time.Now().Add(10 * time.Second)) + defer conn.SetReadDeadline(time.Time{}) // Clear deadline. + + for { + select { + case <-ctx.Done(): + return "", 0, ctx.Err() + default: + } + + _, data, err := conn.ReadMessage() + if err != nil { + return "", 0, fmt.Errorf("pusher: waiting for connection_established: %w", err) + } + + var msg pusherMessage + if err := json.Unmarshal(data, &msg); err != nil { + continue // Skip unparseable messages. + } + + if msg.Event == "pusher:connection_established" { + // The data field is a JSON-encoded string inside the outer JSON, + // so we unmarshal twice: first to get the string, then to parse it. + var dataStr string + if err := json.Unmarshal(msg.Data, &dataStr); err != nil { + return "", 0, fmt.Errorf("pusher: parsing connection data wrapper: %w", err) + } + var connData pusherConnectionData + if err := json.Unmarshal([]byte(dataStr), &connData); err != nil { + return "", 0, fmt.Errorf("pusher: parsing connection data: %w", err) + } + if connData.SocketID == "" { + return "", 0, fmt.Errorf("pusher: empty socket_id in connection_established") + } + return connData.SocketID, connData.ActivityTimeout, nil + } + + if msg.Event == "pusher:error" { + return "", 0, fmt.Errorf("pusher: server error during connect: %s", string(msg.Data)) + } + } +} + +// subscribe authenticates and subscribes to the private channel. +func (p *PusherClient) subscribe(ctx context.Context, conn *websocket.Conn, socketID string) error { + // Get auth signature from the auth endpoint. + authSig, err := p.authFn(ctx, socketID, p.channel) + if err != nil { + return fmt.Errorf("pusher: channel auth failed: %w", err) + } + + // Send subscribe message. + subData, _ := json.Marshal(map[string]string{ + "auth": authSig, + "channel": p.channel, + }) + subMsg := pusherMessage{ + Event: "pusher:subscribe", + Data: subData, + } + + conn.SetWriteDeadline(time.Now().Add(pusherWriteTimeout)) + if err := conn.WriteJSON(subMsg); err != nil { + return fmt.Errorf("pusher: sending subscribe: %w", err) + } + conn.SetWriteDeadline(time.Time{}) + + // Wait for subscription_succeeded or error. + conn.SetReadDeadline(time.Now().Add(10 * time.Second)) + defer conn.SetReadDeadline(time.Time{}) + + for { + select { + case <-ctx.Done(): + return ctx.Err() + default: + } + + _, data, err := conn.ReadMessage() + if err != nil { + return fmt.Errorf("pusher: waiting for subscription response: %w", err) + } + + var msg pusherMessage + if err := json.Unmarshal(data, &msg); err != nil { + continue + } + + if msg.Event == "pusher_internal:subscription_succeeded" && msg.Channel == p.channel { + return nil + } + + if msg.Event == "pusher:error" { + return fmt.Errorf("pusher: subscription error: %s", string(msg.Data)) + } + } +} + +// readResult is a message or error from the reader goroutine. +type readResult struct { + data []byte + err error +} + +// readLoop reads messages from the WebSocket and dispatches command events. +// +// A separate goroutine performs the blocking ReadMessage calls and feeds +// results into a channel, allowing the main loop to select on both incoming +// messages and the ping timer. Per the Pusher protocol, the client sends a +// pusher:ping after activityTimeout seconds of inactivity; if no pusher:pong +// arrives within pusherPingTimeout, the read deadline expires and the +// connection is considered dead. +func (p *PusherClient) readLoop(ctx context.Context, conn *websocket.Conn, activityTimeout int) { + defer close(p.done) + defer close(p.recvCh) + + // Ping interval is the server's advertised activity timeout — the client + // should send a ping after this many seconds of silence. + pingInterval := time.Duration(activityTimeout) * time.Second + if pingInterval <= 0 { + pingInterval = 120 * time.Second // Default Pusher activity timeout. + } + pingTimer := time.NewTimer(pingInterval) + defer pingTimer.Stop() + + // Reader goroutine: performs blocking ReadMessage calls and feeds results + // to readCh. The read deadline allows pingInterval for normal activity + // plus pusherPingTimeout for a pong response after we send a ping. + readCh := make(chan readResult, 1) + go func() { + for { + conn.SetReadDeadline(time.Now().Add(pingInterval + pusherPingTimeout)) + _, data, err := conn.ReadMessage() + select { + case readCh <- readResult{data, err}: + case <-ctx.Done(): + return + } + if err != nil { + return + } + } + }() + + for { + select { + case <-ctx.Done(): + return + + case result := <-readCh: + if result.err != nil { + select { + case <-ctx.Done(): + return + default: + } + p.mu.Lock() + stopped := p.stopped + p.mu.Unlock() + if stopped { + return + } + log.Printf("Pusher read error: %v", result.err) + return + } + + // Reset ping timer on any received data. + if !pingTimer.Stop() { + select { + case <-pingTimer.C: + default: + } + } + pingTimer.Reset(pingInterval) + + p.handleMessage(result.data) + + case <-pingTimer.C: + // No activity for pingInterval — send a ping to keep alive. + if !p.sendPusherMessage(conn, pusherMessage{ + Event: "pusher:ping", + Data: json.RawMessage("{}"), + }) { + return + } + // The read deadline (pingInterval + pusherPingTimeout) gives the + // server pusherPingTimeout to respond with pusher:pong. If no + // response arrives, ReadMessage returns a timeout error. + } + } +} + +// handleMessage processes a single Pusher protocol message. +func (p *PusherClient) handleMessage(data []byte) { + var msg pusherMessage + if err := json.Unmarshal(data, &msg); err != nil { + log.Printf("Pusher: ignoring unparseable message: %v", err) + return + } + + switch msg.Event { + case "pusher:ping": + // Respond with pong. + p.sendPusherMessage(p.getConn(), pusherMessage{ + Event: "pusher:pong", + Data: json.RawMessage("{}"), + }) + + case "pusher:pong": + // Server responded to our ping — connection confirmed alive. + + case "pusher:error": + log.Printf("Pusher server error: %s", string(msg.Data)) + + case "chief.command": + if msg.Channel == p.channel { + // Pusher wraps event data as a JSON-encoded string, so we + // must unwrap it before forwarding to the command handler. + payload := msg.Data + var dataStr string + if err := json.Unmarshal(msg.Data, &dataStr); err == nil { + payload = json.RawMessage(dataStr) + } + select { + case p.recvCh <- payload: + default: + log.Printf("Pusher: receive buffer full, dropping command") + } + } + + default: + // Ignore other event types (subscription_succeeded during reconnect, etc.). + } +} + +// sendPusherMessage writes a Pusher protocol JSON message to the connection. +// Returns false if the write failed (connection should be considered dead). +func (p *PusherClient) sendPusherMessage(conn *websocket.Conn, msg pusherMessage) bool { + if conn == nil { + return false + } + conn.SetWriteDeadline(time.Now().Add(pusherWriteTimeout)) + err := conn.WriteJSON(msg) + conn.SetWriteDeadline(time.Time{}) + if err != nil { + log.Printf("Pusher: write error: %v", err) + return false + } + return true +} + +// getConn returns the current WebSocket connection, or nil if closed. +func (p *PusherClient) getConn() *websocket.Conn { + p.mu.Lock() + defer p.mu.Unlock() + return p.conn +} + +// BroadcastAuth authenticates a Pusher channel subscription via the uplink HTTP client. +// This creates an AuthFunc that calls POST /api/device/broadcasting/auth. +func (c *Client) BroadcastAuth(ctx context.Context, socketID, channelName string) (string, error) { + body := broadcastAuthRequest{ + SocketID: socketID, + ChannelName: channelName, + } + + var resp pusherAuthResponse + if err := c.doJSON(ctx, "POST", "/api/device/broadcasting/auth", body, &resp); err != nil { + return "", fmt.Errorf("broadcast auth: %w", err) + } + + return resp.Auth, nil +} + +// broadcastAuthRequest is the JSON body sent to POST /api/device/broadcasting/auth. +type broadcastAuthRequest struct { + SocketID string `json:"socket_id"` + ChannelName string `json:"channel_name"` +} + +// GenerateAuthSignature generates a Pusher private channel auth signature locally. +// This is used in tests to verify auth signatures without hitting the server. +func GenerateAuthSignature(appKey, appSecret, socketID, channelName string) string { + data := socketID + ":" + channelName + mac := hmac.New(sha256.New, []byte(appSecret)) + mac.Write([]byte(data)) + sig := fmt.Sprintf("%x", mac.Sum(nil)) + return appKey + ":" + sig +} diff --git a/internal/uplink/pusher_test.go b/internal/uplink/pusher_test.go new file mode 100644 index 00000000..a496e424 --- /dev/null +++ b/internal/uplink/pusher_test.go @@ -0,0 +1,861 @@ +package uplink + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/gorilla/websocket" +) + +// testPusherServer is a mock Pusher/Reverb WebSocket server for testing. +type testPusherServer struct { + srv *httptest.Server + upgrader websocket.Upgrader + + mu sync.Mutex + conn *websocket.Conn + + // Configuration. + appKey string + appSecret string + socketID string + activityTimeout int + + // Control channels. + onSubscribe chan string // receives channel name when client subscribes + onMessage chan []byte // receives raw messages from client + + // Behavior flags. + rejectAuth bool + rejectSubscribe bool + skipEstablished bool +} + +func newTestPusherServer(t *testing.T) *testPusherServer { + t.Helper() + + ps := &testPusherServer{ + appKey: "test-app-key", + appSecret: "test-app-secret", + socketID: "123456.7890", + activityTimeout: 120, + onSubscribe: make(chan string, 10), + onMessage: make(chan []byte, 10), + upgrader: websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { return true }, + }, + } + + ps.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ps.handleWS(t, w, r) + })) + + t.Cleanup(func() { + ps.mu.Lock() + if ps.conn != nil { + ps.conn.Close() + } + ps.mu.Unlock() + ps.srv.Close() + }) + + return ps +} + +func (ps *testPusherServer) handleWS(t *testing.T, w http.ResponseWriter, r *http.Request) { + t.Helper() + + // Verify the URL path matches Pusher format. + expectedPath := fmt.Sprintf("/app/%s", ps.appKey) + if !strings.HasPrefix(r.URL.Path, expectedPath) { + http.Error(w, "invalid path", http.StatusNotFound) + return + } + + conn, err := ps.upgrader.Upgrade(w, r, nil) + if err != nil { + t.Logf("upgrade error: %v", err) + return + } + + ps.mu.Lock() + ps.conn = conn + ps.mu.Unlock() + + // Send connection_established unless configured not to. + if !ps.skipEstablished { + connDataJSON, _ := json.Marshal(pusherConnectionData{ + SocketID: ps.socketID, + ActivityTimeout: ps.activityTimeout, + }) + // Real Pusher/Reverb double-encodes: the data field is a JSON string. + connDataStr, _ := json.Marshal(string(connDataJSON)) + established := pusherMessage{ + Event: "pusher:connection_established", + Data: connDataStr, + } + if err := conn.WriteJSON(established); err != nil { + t.Logf("write connection_established: %v", err) + return + } + } + + // Read loop — handle subscribe messages and pass others to onMessage. + for { + _, data, err := conn.ReadMessage() + if err != nil { + return + } + + var msg pusherMessage + if json.Unmarshal(data, &msg) != nil { + continue + } + + switch msg.Event { + case "pusher:subscribe": + var subData map[string]string + json.Unmarshal(msg.Data, &subData) + channel := subData["channel"] + + select { + case ps.onSubscribe <- channel: + default: + } + + if ps.rejectSubscribe { + errData, _ := json.Marshal(map[string]interface{}{ + "message": "subscription rejected", + "code": 4009, + }) + conn.WriteJSON(pusherMessage{ + Event: "pusher:error", + Data: errData, + }) + continue + } + + // Send subscription_succeeded. + conn.WriteJSON(pusherMessage{ + Event: "pusher_internal:subscription_succeeded", + Channel: channel, + Data: json.RawMessage("{}"), + }) + + case "pusher:pong": + select { + case ps.onMessage <- data: + default: + } + + default: + select { + case ps.onMessage <- data: + default: + } + } + } +} + +// sendCommand sends a chief.command event to the connected client. +func (ps *testPusherServer) sendCommand(channel string, command json.RawMessage) error { + ps.mu.Lock() + conn := ps.conn + ps.mu.Unlock() + + if conn == nil { + return fmt.Errorf("no client connected") + } + + msg := pusherMessage{ + Event: "chief.command", + Channel: channel, + Data: command, + } + return conn.WriteJSON(msg) +} + +// sendCommandStringEncoded sends a chief.command event where the data field +// is a JSON-encoded string, matching real Reverb/Pusher wire format. +func (ps *testPusherServer) sendCommandStringEncoded(channel string, command json.RawMessage) error { + ps.mu.Lock() + conn := ps.conn + ps.mu.Unlock() + + if conn == nil { + return fmt.Errorf("no client connected") + } + + // Double-encode: wrap the JSON object as a JSON string. + encoded, err := json.Marshal(string(command)) + if err != nil { + return fmt.Errorf("encoding command: %w", err) + } + + msg := pusherMessage{ + Event: "chief.command", + Channel: channel, + Data: encoded, + } + return conn.WriteJSON(msg) +} + +// closeConnection closes the WebSocket connection from the server side, +// simulating a Pusher disconnection. +func (ps *testPusherServer) closeConnection() error { + ps.mu.Lock() + conn := ps.conn + ps.mu.Unlock() + + if conn == nil { + return fmt.Errorf("no client connected") + } + + return conn.Close() +} + +// sendPing sends a pusher:ping to the connected client. +func (ps *testPusherServer) sendPing() error { + ps.mu.Lock() + conn := ps.conn + ps.mu.Unlock() + + if conn == nil { + return fmt.Errorf("no client connected") + } + + return conn.WriteJSON(pusherMessage{ + Event: "pusher:ping", + Data: json.RawMessage("{}"), + }) +} + +// reverbConfig returns a ReverbConfig pointing at the test server. +func (ps *testPusherServer) reverbConfig() ReverbConfig { + // Extract host and port from the test server URL. + addr := ps.srv.Listener.Addr().String() + parts := strings.Split(addr, ":") + host := parts[0] + port := 0 + fmt.Sscanf(parts[1], "%d", &port) + + return ReverbConfig{ + Key: ps.appKey, + Host: host, + Port: port, + Scheme: "http", + } +} + +// testAuthFn returns an AuthFunc that uses the test server's app key/secret. +func (ps *testPusherServer) testAuthFn() AuthFunc { + return func(ctx context.Context, socketID, channelName string) (string, error) { + return GenerateAuthSignature(ps.appKey, ps.appSecret, socketID, channelName), nil + } +} + +// failingAuthFn returns an AuthFunc that always fails. +func failingAuthFn() AuthFunc { + return func(ctx context.Context, socketID, channelName string) (string, error) { + return "", fmt.Errorf("auth endpoint unavailable") + } +} + +// --- Tests --- + +func TestPusherClient_ConnectAndReceive(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer client.Close() + + // Wait for subscription. + select { + case ch := <-ps.onSubscribe: + if ch != channel { + t.Errorf("subscribed to %q, want %q", ch, channel) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Send a command and verify receipt. + cmd := json.RawMessage(`{"type":"start_run","project":"test"}`) + if err := ps.sendCommand(channel, cmd); err != nil { + t.Fatalf("sendCommand failed: %v", err) + } + + select { + case received := <-client.Receive(): + var parsed map[string]interface{} + json.Unmarshal(received, &parsed) + if parsed["type"] != "start_run" { + t.Errorf("received type = %v, want start_run", parsed["type"]) + } + if parsed["project"] != "test" { + t.Errorf("received project = %v, want test", parsed["project"]) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for command") + } +} + +func TestPusherClient_MultipleCommands(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.99" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer client.Close() + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Send multiple commands. + commands := []string{ + `{"type":"start_run","id":"1"}`, + `{"type":"pause_run","id":"2"}`, + `{"type":"stop_run","id":"3"}`, + } + + for _, cmd := range commands { + if err := ps.sendCommand(channel, json.RawMessage(cmd)); err != nil { + t.Fatalf("sendCommand failed: %v", err) + } + } + + // Receive all commands in order. + for i, expected := range commands { + select { + case received := <-client.Receive(): + var expectedMap, receivedMap map[string]interface{} + json.Unmarshal([]byte(expected), &expectedMap) + json.Unmarshal(received, &receivedMap) + if receivedMap["id"] != expectedMap["id"] { + t.Errorf("command %d: id = %v, want %v", i, receivedMap["id"], expectedMap["id"]) + } + case <-time.After(5 * time.Second): + t.Fatalf("timeout waiting for command %d", i) + } + } +} + +func TestPusherClient_IgnoresOtherChannels(t *testing.T) { + ps := newTestPusherServer(t) + myChannel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), myChannel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer client.Close() + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Send a command on a different channel — should be ignored. + if err := ps.sendCommand("private-chief-server.99", json.RawMessage(`{"type":"other"}`)); err != nil { + t.Fatalf("sendCommand failed: %v", err) + } + + // Send a command on our channel — should be received. + if err := ps.sendCommand(myChannel, json.RawMessage(`{"type":"mine"}`)); err != nil { + t.Fatalf("sendCommand failed: %v", err) + } + + select { + case received := <-client.Receive(): + var parsed map[string]interface{} + json.Unmarshal(received, &parsed) + if parsed["type"] != "mine" { + t.Errorf("received type = %v, want mine (wrong channel message leaked through)", parsed["type"]) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for command") + } +} + +func TestPusherClient_PingPong(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer client.Close() + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Send a ping and verify pong response. + if err := ps.sendPing(); err != nil { + t.Fatalf("sendPing failed: %v", err) + } + + // The client should send back a pusher:pong. + select { + case data := <-ps.onMessage: + var msg pusherMessage + if err := json.Unmarshal(data, &msg); err != nil { + t.Fatalf("failed to parse pong message: %v", err) + } + if msg.Event != "pusher:pong" { + t.Errorf("response event = %q, want pusher:pong", msg.Event) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for pong") + } +} + +func TestPusherClient_Close(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Close and verify the receive channel closes. + if err := client.Close(); err != nil { + t.Fatalf("Close() failed: %v", err) + } + + // Receive channel should be closed. + select { + case _, ok := <-client.Receive(): + if ok { + t.Error("expected receive channel to be closed after Close()") + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for receive channel to close") + } + + // Double-close should be safe. + if err := client.Close(); err != nil { + t.Fatalf("double Close() failed: %v", err) + } +} + +func TestPusherClient_ContextCancellation(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Cancel context — the read loop should stop. + cancel() + + // Give the readLoop time to notice the cancellation and close. + select { + case <-client.Receive(): + // Channel closed or drained — good. + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for shutdown after context cancellation") + } +} + +func TestPusherClient_AuthFailure(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, failingAuthFn()) + + ctx := testContext(t) + err := client.Connect(ctx) + if err == nil { + t.Fatal("expected error when auth fails, got nil") + client.Close() + } + if !strings.Contains(err.Error(), "auth endpoint unavailable") { + t.Errorf("error = %v, want containing 'auth endpoint unavailable'", err) + } +} + +func TestPusherClient_SubscriptionRejected(t *testing.T) { + ps := newTestPusherServer(t) + ps.rejectSubscribe = true + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + err := client.Connect(ctx) + if err == nil { + t.Fatal("expected error when subscription is rejected, got nil") + client.Close() + } + if !strings.Contains(err.Error(), "subscription error") { + t.Errorf("error = %v, want containing 'subscription error'", err) + } +} + +func TestPusherClient_BuildURL(t *testing.T) { + tests := []struct { + name string + cfg ReverbConfig + expect string + }{ + { + name: "HTTPS scheme", + cfg: ReverbConfig{ + Key: "my-key", + Host: "reverb.example.com", + Port: 443, + Scheme: "https", + }, + expect: "wss://reverb.example.com:443/app/my-key?protocol=7", + }, + { + name: "HTTP scheme", + cfg: ReverbConfig{ + Key: "local-key", + Host: "localhost", + Port: 8080, + Scheme: "http", + }, + expect: "ws://localhost:8080/app/local-key?protocol=7", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + p := NewPusherClient(tt.cfg, "private-test", nil) + got := p.buildURL() + if got != tt.expect { + t.Errorf("buildURL() = %q, want %q", got, tt.expect) + } + }) + } +} + +func TestPusherClient_ReceiveChannelBuffered(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer client.Close() + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // The receive channel should be buffered. + if cap(client.recvCh) != receiveBufSize { + t.Errorf("receive channel capacity = %d, want %d", cap(client.recvCh), receiveBufSize) + } +} + +func TestPusherClient_ServerError(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer client.Close() + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Send a pusher:error — client should log it but not crash. + ps.mu.Lock() + conn := ps.conn + ps.mu.Unlock() + + errData, _ := json.Marshal(map[string]interface{}{"message": "test error", "code": 4100}) + conn.WriteJSON(pusherMessage{ + Event: "pusher:error", + Data: errData, + }) + + // Send a command after the error — client should still be functioning. + if err := ps.sendCommand(channel, json.RawMessage(`{"type":"after_error"}`)); err != nil { + t.Fatalf("sendCommand failed: %v", err) + } + + select { + case received := <-client.Receive(): + var parsed map[string]interface{} + json.Unmarshal(received, &parsed) + if parsed["type"] != "after_error" { + t.Errorf("received type = %v, want after_error", parsed["type"]) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for command after error") + } +} + +func TestBroadcastAuth_Success(t *testing.T) { + var receivedSocketID, receivedChannel string + var receivedAuth string + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/device/broadcasting/auth" { + http.NotFound(w, r) + return + } + + receivedAuth = r.Header.Get("Authorization") + + var body broadcastAuthRequest + json.NewDecoder(r.Body).Decode(&body) + receivedSocketID = body.SocketID + receivedChannel = body.ChannelName + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(pusherAuthResponse{ + Auth: "app-key:test-signature", + }) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "test-token") + ctx := testContext(t) + + auth, err := client.BroadcastAuth(ctx, "12345.67890", "private-chief-server.42") + if err != nil { + t.Fatalf("BroadcastAuth() failed: %v", err) + } + + if receivedAuth != "Bearer test-token" { + t.Errorf("Authorization = %q, want %q", receivedAuth, "Bearer test-token") + } + if receivedSocketID != "12345.67890" { + t.Errorf("socket_id = %q, want %q", receivedSocketID, "12345.67890") + } + if receivedChannel != "private-chief-server.42" { + t.Errorf("channel_name = %q, want %q", receivedChannel, "private-chief-server.42") + } + if auth != "app-key:test-signature" { + t.Errorf("auth = %q, want %q", auth, "app-key:test-signature") + } +} + +func TestBroadcastAuth_AuthFailed(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "bad-token") + ctx := testContext(t) + + _, err := client.BroadcastAuth(ctx, "12345.67890", "private-chief-server.42") + if err == nil { + t.Fatal("expected error for 401, got nil") + } +} + +func TestGenerateAuthSignature(t *testing.T) { + // Known test vectors for Pusher private channel auth. + sig := GenerateAuthSignature("278d425bdf160313ff76", "7ad3773142a6692b25b8", "1234.1234", "private-foobar") + + // The format should be "key:hex_signature". + if !strings.HasPrefix(sig, "278d425bdf160313ff76:") { + t.Errorf("signature should start with app key, got %q", sig) + } + + parts := strings.SplitN(sig, ":", 2) + if len(parts) != 2 { + t.Fatalf("signature should have format key:sig, got %q", sig) + } + if len(parts[1]) != 64 { // SHA256 hex = 64 chars + t.Errorf("signature hex length = %d, want 64", len(parts[1])) + } +} + +func TestPusherClient_ConnectionEstablishedTimeout(t *testing.T) { + ps := newTestPusherServer(t) + ps.skipEstablished = true // Server won't send connection_established. + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + err := client.Connect(ctx) + if err == nil { + t.Fatal("expected error when connection_established is not received, got nil") + client.Close() + } +} + +func TestPusherClient_ConcurrentClose(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Close from multiple goroutines concurrently. + var wg sync.WaitGroup + var closeErrors atomic.Int32 + for i := 0; i < 5; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if err := client.Close(); err != nil { + closeErrors.Add(1) + } + }() + } + wg.Wait() + + // At most one goroutine should get an error (the connection close), rest should be nil. + // This test mainly verifies no panics from concurrent access. +} + +func TestPusherClient_DialFailure(t *testing.T) { + // Connect to a non-existent server. + cfg := ReverbConfig{ + Key: "test-key", + Host: "127.0.0.1", + Port: 1, // Port 1 should be unreachable. + Scheme: "http", + } + channel := "private-chief-server.42" + client := NewPusherClient(cfg, channel, failingAuthFn()) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + err := client.Connect(ctx) + if err == nil { + t.Fatal("expected error connecting to unreachable server, got nil") + client.Close() + } +} + +// TestPusherClient_DoubleEncodedData verifies that commands with Pusher's +// real wire format (data as JSON string) are correctly unwrapped. +func TestPusherClient_DoubleEncodedData(t *testing.T) { + ps := newTestPusherServer(t) + channel := "private-chief-server.42" + + client := NewPusherClient(ps.reverbConfig(), channel, ps.testAuthFn()) + + ctx := testContext(t) + if err := client.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer client.Close() + + // Wait for subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Send command using Reverb's real format (data as JSON string). + cmd := json.RawMessage(`{"type":"start_run","payload":{"project_slug":"my-project"}}`) + if err := ps.sendCommandStringEncoded(channel, cmd); err != nil { + t.Fatalf("sendCommandStringEncoded failed: %v", err) + } + + select { + case received := <-client.Receive(): + var parsed map[string]interface{} + if err := json.Unmarshal(received, &parsed); err != nil { + t.Fatalf("failed to parse received command: %v (raw: %s)", err, string(received)) + } + if parsed["type"] != "start_run" { + t.Errorf("received type = %v, want start_run", parsed["type"]) + } + payload, ok := parsed["payload"].(map[string]interface{}) + if !ok { + t.Fatalf("payload is not an object: %v", parsed["payload"]) + } + if payload["project_slug"] != "my-project" { + t.Errorf("received project_slug = %v, want my-project", payload["project_slug"]) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for command") + } +} diff --git a/internal/uplink/uplink.go b/internal/uplink/uplink.go new file mode 100644 index 00000000..58a972b1 --- /dev/null +++ b/internal/uplink/uplink.go @@ -0,0 +1,553 @@ +package uplink + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log" + "sync" + "time" +) + +const ( + // heartbeatInterval is how often heartbeats are sent. + heartbeatInterval = 30 * time.Second + + // heartbeatRetryDelay is the delay before retrying a failed heartbeat. + heartbeatRetryDelay = 5 * time.Second + + // heartbeatSkipWindow is the duration after a message send within which + // we skip the explicit heartbeat (server treats message receipt as implicit heartbeat). + heartbeatSkipWindow = 25 * time.Second + + // heartbeatMaxFailures is the number of consecutive failures before triggering reconnection. + heartbeatMaxFailures = 3 +) + +// Uplink composes the HTTP client, message batcher, and Pusher client +// into a unified Send/Receive interface. +type Uplink struct { + client *Client + batcher *Batcher + pusher *PusherClient + + mu sync.RWMutex + sessionID string + deviceID int + connected bool + + // lastSendTime records when the batcher last successfully sent a batch. + // Used by the heartbeat goroutine to skip heartbeats when messages + // were recently sent (implicit heartbeat optimization). + lastSendTime time.Time + + // Heartbeat timing (overridable for tests, default to package constants). + hbInterval time.Duration + hbRetryDelay time.Duration + hbSkipWindow time.Duration + hbMaxFails int + + // recvCh is a stable receive channel that outlives individual Pusher clients. + // Commands from each Pusher client are forwarded into this channel so callers + // of Receive() don't need to re-subscribe after reconnection. + recvCh chan json.RawMessage + + // onReconnect is called after each successful reconnection. + onReconnect func() + + // onAuthFailure is called when a 401 auth error occurs during reconnection. + // The callback should perform a token refresh and call SetAccessToken() before returning. + // If the callback returns nil, reconnection retries with the new token. + // If it returns an error, reconnection aborts. + onAuthFailure func() error + + // onHeartbeatMaxFailures is called when consecutive heartbeat failures + // reach hbMaxFails. If nil, triggerReconnect() is called directly. + // Tests can set this to override the default reconnection behavior. + onHeartbeatMaxFailures func() + + // reconnecting tracks whether a reconnection is in progress to prevent concurrent reconnects. + reconnecting bool + + // parentCtx is the context passed to Connect() — used as the parent for reconnection contexts. + parentCtx context.Context + + // cancel stops the batcher run loop and heartbeat goroutine. + cancel context.CancelFunc +} + +// UplinkOption configures an Uplink. +type UplinkOption func(*Uplink) + +// WithOnReconnect sets a callback invoked after each successful reconnection. +// This matches the ws.WithOnReconnect pattern — serve.go uses it to re-send +// a full state snapshot after reconnecting. +func WithOnReconnect(fn func()) UplinkOption { + return func(u *Uplink) { + u.onReconnect = fn + } +} + +// WithOnAuthFailure sets a callback invoked when a 401 auth error occurs during +// reconnection. The callback should perform a token refresh and call +// SetAccessToken() before returning. Return nil to retry, or an error to abort. +func WithOnAuthFailure(fn func() error) UplinkOption { + return func(u *Uplink) { + u.onAuthFailure = fn + } +} + +// NewUplink creates a new Uplink that uses the given HTTP client. +// The Uplink does not connect until Connect is called. +func NewUplink(client *Client, opts ...UplinkOption) *Uplink { + u := &Uplink{ + client: client, + hbInterval: heartbeatInterval, + hbRetryDelay: heartbeatRetryDelay, + hbSkipWindow: heartbeatSkipWindow, + hbMaxFails: heartbeatMaxFailures, + recvCh: make(chan json.RawMessage, receiveBufSize), + } + for _, o := range opts { + o(u) + } + return u +} + +// Connect establishes the full uplink connection: +// 1. HTTP connect (registers device, gets session ID + Reverb config) +// 2. Pusher connect (subscribes to private command channel) +// 3. Batcher start (begins background flush loop) +// 4. Heartbeat start (sends periodic heartbeats to server) +// 5. Pusher monitor (detects disconnection and triggers reconnection) +func (u *Uplink) Connect(ctx context.Context) error { + // Step 1: HTTP connect to register the device. + welcome, err := u.client.Connect(ctx) + if err != nil { + return fmt.Errorf("uplink connect: %w", err) + } + + u.mu.Lock() + u.sessionID = welcome.SessionID + u.deviceID = welcome.DeviceID + u.connected = true + u.parentCtx = ctx + u.mu.Unlock() + + // Step 2: Start the Pusher client for receiving commands. + channel := fmt.Sprintf("private-chief-server.%d", welcome.DeviceID) + u.pusher = NewPusherClient(welcome.Reverb, channel, u.client.BroadcastAuth) + + if err := u.pusher.Connect(ctx); err != nil { + // Clean up: disconnect from HTTP since Pusher failed. + disconnectCtx, cancel := context.WithTimeout(context.Background(), httpTimeout) + defer cancel() + if dErr := u.client.Disconnect(disconnectCtx); dErr != nil { + log.Printf("uplink: failed to disconnect after Pusher error: %v", dErr) + } + return fmt.Errorf("uplink pusher connect: %w", err) + } + + // Step 3: Start the batcher for outgoing messages. + batchCtx, batchCancel := context.WithCancel(ctx) + u.cancel = batchCancel + + u.batcher = NewBatcher(func(batchID string, messages []json.RawMessage) error { + _, err := u.client.SendMessagesWithRetry(batchCtx, batchID, messages) + if err == nil { + u.mu.Lock() + u.lastSendTime = time.Now() + u.mu.Unlock() + } + return err + }) + go u.batcher.Run(batchCtx) + + // Step 4: Start the heartbeat goroutine. + go u.runHeartbeat(batchCtx) + + // Step 5: Monitor Pusher for disconnection. + go u.monitorPusher(batchCtx) + + log.Printf("Uplink connected (device=%d, session=%s)", welcome.DeviceID, welcome.SessionID) + return nil +} + +// Send enqueues a message into the batcher for batched delivery. +// The batcher handles flush timing. +// During reconnection, messages are buffered locally in the batcher. +func (u *Uplink) Send(msg json.RawMessage, msgType string) { + u.mu.RLock() + connected := u.connected + u.mu.RUnlock() + + if !connected { + log.Printf("uplink: dropping message (type=%s) — not connected", msgType) + return + } + + u.batcher.Enqueue(msg, msgType) +} + +// Receive returns a channel that delivers incoming command payloads. +// This channel is stable across reconnections — new Pusher clients +// forward commands into the same channel. +func (u *Uplink) Receive() <-chan json.RawMessage { + return u.recvCh +} + +// Close performs graceful shutdown: +// 1. Stop the batcher (flushes remaining messages) +// 2. Close the Pusher client +// 3. HTTP disconnect +func (u *Uplink) Close() error { + return u.doClose() +} + +// CloseWithTimeout performs graceful shutdown with a deadline. +// If the timeout expires before the batcher flush completes, the flush is +// abandoned and shutdown continues with Pusher close and HTTP disconnect. +// This prevents shutdown from hanging when the server is unreachable. +func (u *Uplink) CloseWithTimeout(timeout time.Duration) error { + done := make(chan error, 1) + go func() { + done <- u.doClose() + }() + + select { + case err := <-done: + return err + case <-time.After(timeout): + log.Printf("uplink: graceful close timed out after %s — forcing shutdown", timeout) + // Force-cancel the batcher/heartbeat/monitor contexts to unblock doClose. + u.mu.Lock() + u.connected = false + u.mu.Unlock() + if u.cancel != nil { + u.cancel() + } + // Wait briefly for doClose to finish after cancellation. + select { + case err := <-done: + return err + case <-time.After(2 * time.Second): + log.Printf("uplink: forced shutdown complete") + return nil + } + } +} + +// doClose is the internal close implementation shared by Close and CloseWithTimeout. +func (u *Uplink) doClose() error { + u.mu.Lock() + if !u.connected { + u.mu.Unlock() + return nil + } + u.connected = false + u.mu.Unlock() + + // Step 1: Stop the batcher — this flushes remaining messages. + if u.batcher != nil { + u.batcher.Stop() + } + + // Cancel the batcher context to stop the Run loop, heartbeat, and Pusher monitor. + if u.cancel != nil { + u.cancel() + } + + // Step 2: Close the Pusher client. + var pusherErr error + if u.pusher != nil { + pusherErr = u.pusher.Close() + } + + // Step 3: HTTP disconnect. + disconnectCtx, cancel := context.WithTimeout(context.Background(), httpTimeout) + defer cancel() + if err := u.client.Disconnect(disconnectCtx); err != nil { + log.Printf("uplink: disconnect failed: %v", err) + } + + log.Printf("Uplink disconnected") + return pusherErr +} + +// monitorPusher watches the Pusher client's receive channel. When it closes +// (Pusher readLoop exited due to an error), it triggers a full reconnection. +func (u *Uplink) monitorPusher(ctx context.Context) { + if u.pusher == nil { + return + } + pusherRecv := u.pusher.Receive() + + for { + select { + case <-ctx.Done(): + return + case msg, ok := <-pusherRecv: + if !ok { + // Pusher channel closed — readLoop exited. + // Check if we're shutting down. + select { + case <-ctx.Done(): + return + default: + } + + u.mu.RLock() + connected := u.connected + u.mu.RUnlock() + if !connected { + return + } + + u.triggerReconnect("Pusher disconnected") + return + } + // Forward the command to the stable recvCh. + select { + case u.recvCh <- msg: + default: + log.Printf("uplink: receive buffer full, dropping command") + } + } + } +} + +// triggerReconnect initiates a reconnection attempt in the background. +// It is safe to call from multiple goroutines — only one reconnection runs at a time. +func (u *Uplink) triggerReconnect(reason string) { + u.mu.Lock() + if u.reconnecting || !u.connected { + u.mu.Unlock() + return + } + u.reconnecting = true + parentCtx := u.parentCtx + u.mu.Unlock() + + log.Printf("uplink: triggering reconnection (%s)", reason) + go u.reconnect(parentCtx) +} + +// reconnect tears down the existing connection and re-establishes it with backoff. +// On success, it fires the onReconnect callback so the caller can re-send state. +func (u *Uplink) reconnect(ctx context.Context) { + defer func() { + u.mu.Lock() + u.reconnecting = false + u.mu.Unlock() + }() + + // Step 1: Tear down old batcher and Pusher. + // Stop the batcher — this flushes remaining messages. + if u.batcher != nil { + u.batcher.Stop() + } + + // Cancel old batcher context to stop the old Run loop, heartbeat, and monitor. + if u.cancel != nil { + u.cancel() + } + + // Close the old Pusher client. + if u.pusher != nil { + if err := u.pusher.Close(); err != nil { + log.Printf("uplink: error closing Pusher during reconnection: %v", err) + } + } + + // Step 2: Reconnect with exponential backoff. + attempt := 0 + for { + select { + case <-ctx.Done(): + log.Printf("uplink: reconnection cancelled") + return + default: + } + + u.mu.RLock() + connected := u.connected + u.mu.RUnlock() + if !connected { + // Close() was called — stop reconnecting. + return + } + + attempt++ + delay := backoff(attempt) + log.Printf("uplink: reconnection attempt %d — retrying in %s", attempt, delay.Round(time.Millisecond)) + + select { + case <-ctx.Done(): + log.Printf("uplink: reconnection cancelled") + return + case <-time.After(delay): + } + + // Try HTTP connect. + welcome, err := u.client.Connect(ctx) + if err != nil { + if errors.Is(err, ErrAuthFailed) { + // Auth failure — try token refresh if callback is set. + if u.onAuthFailure != nil { + log.Printf("uplink: auth failed during reconnection — requesting token refresh") + if refreshErr := u.onAuthFailure(); refreshErr != nil { + log.Printf("uplink: token refresh failed: %v — aborting reconnection", refreshErr) + return + } + // Token refreshed — retry without incrementing attempt. + attempt-- + continue + } + log.Printf("uplink: auth failed during reconnection (no refresh callback) — aborting") + return + } + log.Printf("uplink: reconnection attempt %d HTTP connect failed: %v", attempt, err) + continue + } + + // Update session/device. + u.mu.Lock() + u.sessionID = welcome.SessionID + u.deviceID = welcome.DeviceID + u.mu.Unlock() + + // Try Pusher connect. + channel := fmt.Sprintf("private-chief-server.%d", welcome.DeviceID) + pusher := NewPusherClient(welcome.Reverb, channel, u.client.BroadcastAuth) + + if err := pusher.Connect(ctx); err != nil { + log.Printf("uplink: reconnection attempt %d Pusher connect failed: %v — disconnecting HTTP", attempt, err) + disconnectCtx, cancel := context.WithTimeout(context.Background(), httpTimeout) + if dErr := u.client.Disconnect(disconnectCtx); dErr != nil { + log.Printf("uplink: failed to disconnect after Pusher reconnect error: %v", dErr) + } + cancel() + continue + } + + // Step 3: Start new batcher and heartbeat. + batchCtx, batchCancel := context.WithCancel(ctx) + + u.mu.Lock() + u.pusher = pusher + u.cancel = batchCancel + u.lastSendTime = time.Time{} // Reset — force next heartbeat to fire. + u.mu.Unlock() + + u.batcher = NewBatcher(func(batchID string, messages []json.RawMessage) error { + _, err := u.client.SendMessagesWithRetry(batchCtx, batchID, messages) + if err == nil { + u.mu.Lock() + u.lastSendTime = time.Now() + u.mu.Unlock() + } + return err + }) + go u.batcher.Run(batchCtx) + + // Restart heartbeat. + go u.runHeartbeat(batchCtx) + + // Restart Pusher monitor. + go u.monitorPusher(batchCtx) + + log.Printf("Uplink reconnected (attempt %d, device=%d, session=%s)", attempt, welcome.DeviceID, welcome.SessionID) + + // Fire the OnReconnect callback so serve.go can re-send state. + if u.onReconnect != nil { + u.onReconnect() + } + + return + } +} + +// runHeartbeat sends periodic heartbeats to the server every heartbeatInterval. +// It skips the heartbeat if a message batch was sent within heartbeatSkipWindow. +// On transient failure, it retries once after heartbeatRetryDelay. +// After heartbeatMaxFailures consecutive failures, it triggers reconnection. +func (u *Uplink) runHeartbeat(ctx context.Context) { + ticker := time.NewTicker(u.hbInterval) + defer ticker.Stop() + + consecutiveFailures := 0 + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + // Skip heartbeat if a message batch was sent recently. + u.mu.RLock() + lastSend := u.lastSendTime + u.mu.RUnlock() + + if !lastSend.IsZero() && time.Since(lastSend) < u.hbSkipWindow { + consecutiveFailures = 0 + continue + } + + // Send heartbeat. + err := u.client.Heartbeat(ctx) + if err == nil { + consecutiveFailures = 0 + continue + } + + // First failure — retry once after a short delay. + log.Printf("uplink: heartbeat failed: %v — retrying in %s", err, u.hbRetryDelay) + select { + case <-ctx.Done(): + return + case <-time.After(u.hbRetryDelay): + } + + err = u.client.Heartbeat(ctx) + if err == nil { + consecutiveFailures = 0 + continue + } + + // Retry also failed — count as a failure. + consecutiveFailures++ + log.Printf("uplink: heartbeat retry failed (%d/%d consecutive): %v", consecutiveFailures, u.hbMaxFails, err) + + if consecutiveFailures >= u.hbMaxFails { + log.Printf("uplink: %d consecutive heartbeat failures — triggering reconnection", consecutiveFailures) + if u.onHeartbeatMaxFailures != nil { + u.onHeartbeatMaxFailures() + } else { + u.triggerReconnect("heartbeat failures") + } + consecutiveFailures = 0 + } + } + } +} + +// SessionID returns the current session ID from the connect response. +func (u *Uplink) SessionID() string { + u.mu.RLock() + defer u.mu.RUnlock() + return u.sessionID +} + +// DeviceID returns the device ID from the connect response. +func (u *Uplink) DeviceID() int { + u.mu.RLock() + defer u.mu.RUnlock() + return u.deviceID +} + +// SetAccessToken updates the access token on the HTTP client. +// This is called after a token refresh — the new token will be used +// for subsequent HTTP requests and Pusher auth calls. +func (u *Uplink) SetAccessToken(token string) { + u.client.SetAccessToken(token) +} diff --git a/internal/uplink/uplink_test.go b/internal/uplink/uplink_test.go new file mode 100644 index 00000000..012207e3 --- /dev/null +++ b/internal/uplink/uplink_test.go @@ -0,0 +1,1658 @@ +package uplink + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" +) + +// testUplinkServer combines a mock HTTP API server and a mock Pusher WebSocket server +// for end-to-end Uplink testing. +type testUplinkServer struct { + httpSrv *httptest.Server + pusherSrv *testPusherServer + + mu sync.Mutex + connectCalls atomic.Int32 + disconnectCalls atomic.Int32 + heartbeatCalls atomic.Int32 + messageBatches []messageBatch + + // Last received connect metadata. + lastConnectBody map[string]interface{} + + // heartbeatStatus controls the HTTP status code returned by heartbeat. + // 0 or 200 means success. + heartbeatStatus atomic.Int32 + + // connectStatus controls the HTTP status code returned by connect. + // 0 or 200 means success. + connectStatus atomic.Int32 + + // sessionCounter increments on each connect — used for unique session IDs. + sessionCounter atomic.Int32 +} + +type messageBatch struct { + BatchID string + Messages []json.RawMessage +} + +func newTestUplinkServer(t *testing.T) *testUplinkServer { + t.Helper() + + ps := newTestPusherServer(t) + + us := &testUplinkServer{ + pusherSrv: ps, + } + + // Build the Reverb config from the Pusher server. + reverbCfg := ps.reverbConfig() + + us.httpSrv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + us.handleHTTP(t, w, r, reverbCfg) + })) + t.Cleanup(func() { us.httpSrv.Close() }) + + return us +} + +func (us *testUplinkServer) handleHTTP(t *testing.T, w http.ResponseWriter, r *http.Request, reverbCfg ReverbConfig) { + t.Helper() + + // Check auth header. + auth := r.Header.Get("Authorization") + if !strings.HasPrefix(auth, "Bearer ") { + w.WriteHeader(http.StatusUnauthorized) + json.NewEncoder(w).Encode(map[string]string{"error": "missing token"}) + return + } + + w.Header().Set("Content-Type", "application/json") + + switch r.URL.Path { + case "/api/device/connect": + us.connectCalls.Add(1) + + status := int(us.connectStatus.Load()) + if status >= 400 { + w.WriteHeader(status) + json.NewEncoder(w).Encode(map[string]string{"error": "connect failed"}) + return + } + + var body map[string]interface{} + json.NewDecoder(r.Body).Decode(&body) + us.mu.Lock() + us.lastConnectBody = body + us.mu.Unlock() + + n := us.sessionCounter.Add(1) + sessionID := fmt.Sprintf("test-session-%d", n) + + json.NewEncoder(w).Encode(WelcomeResponse{ + Type: "welcome", + ProtocolVersion: 1, + DeviceID: 42, + SessionID: sessionID, + Reverb: reverbCfg, + }) + + case "/api/device/disconnect": + us.disconnectCalls.Add(1) + json.NewEncoder(w).Encode(map[string]string{"status": "disconnected"}) + + case "/api/device/heartbeat": + us.heartbeatCalls.Add(1) + status := int(us.heartbeatStatus.Load()) + if status >= 400 { + w.WriteHeader(status) + json.NewEncoder(w).Encode(map[string]string{"error": "heartbeat failed"}) + return + } + json.NewEncoder(w).Encode(map[string]string{"status": "ok"}) + + case "/api/device/messages": + var req ingestRequest + json.NewDecoder(r.Body).Decode(&req) + + us.mu.Lock() + us.messageBatches = append(us.messageBatches, messageBatch{ + BatchID: req.BatchID, + Messages: req.Messages, + }) + us.mu.Unlock() + + currentSession := fmt.Sprintf("test-session-%d", us.sessionCounter.Load()) + json.NewEncoder(w).Encode(IngestResponse{ + Accepted: len(req.Messages), + BatchID: req.BatchID, + SessionID: currentSession, + }) + + case "/api/device/broadcasting/auth": + var body broadcastAuthRequest + json.NewDecoder(r.Body).Decode(&body) + + sig := GenerateAuthSignature( + us.pusherSrv.appKey, + us.pusherSrv.appSecret, + body.SocketID, + body.ChannelName, + ) + json.NewEncoder(w).Encode(pusherAuthResponse{Auth: sig}) + + default: + http.NotFound(w, r) + } +} + +func (us *testUplinkServer) getMessageBatches() []messageBatch { + us.mu.Lock() + defer us.mu.Unlock() + result := make([]messageBatch, len(us.messageBatches)) + copy(result, us.messageBatches) + return result +} + +// newTestUplink creates an Uplink connected to the test servers. +func newTestUplink(t *testing.T, us *testUplinkServer, opts ...UplinkOption) *Uplink { + t.Helper() + + client := newTestClient(t, us.httpSrv.URL, "test-token") + u := NewUplink(client, opts...) + return u +} + +// --- Tests --- + +func TestUplink_FullLifecycle(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + + // Connect. + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for Pusher subscription") + } + + // Verify connect was called. + if got := us.connectCalls.Load(); got != 1 { + t.Errorf("connect calls = %d, want 1", got) + } + + // Verify session/device IDs. + if got := u.SessionID(); !strings.HasPrefix(got, "test-session-") { + t.Errorf("SessionID() = %q, want prefix %q", got, "test-session-") + } + if got := u.DeviceID(); got != 42 { + t.Errorf("DeviceID() = %d, want 42", got) + } + + // Send a message (immediate tier — should flush right away). + msg := json.RawMessage(`{"type":"run_complete","project":"test"}`) + u.Send(msg, "run_complete") + + // Wait for the batcher to flush. + deadline := time.After(5 * time.Second) + for { + batches := us.getMessageBatches() + if len(batches) > 0 { + if len(batches[0].Messages) != 1 { + t.Errorf("batch has %d messages, want 1", len(batches[0].Messages)) + } + var parsed map[string]interface{} + json.Unmarshal(batches[0].Messages[0], &parsed) + if parsed["type"] != "run_complete" { + t.Errorf("message type = %v, want run_complete", parsed["type"]) + } + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for message batch to be sent") + case <-time.After(10 * time.Millisecond): + } + } + + // Receive a command from the server via Pusher. + channel := fmt.Sprintf("private-chief-server.%d", u.DeviceID()) + cmd := json.RawMessage(`{"type":"start_run","project":"myapp"}`) + if err := us.pusherSrv.sendCommand(channel, cmd); err != nil { + t.Fatalf("sendCommand failed: %v", err) + } + + select { + case received := <-u.Receive(): + var parsed map[string]interface{} + json.Unmarshal(received, &parsed) + if parsed["type"] != "start_run" { + t.Errorf("received type = %v, want start_run", parsed["type"]) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for command") + } + + // Close. + if err := u.Close(); err != nil { + t.Fatalf("Close() failed: %v", err) + } + + // Verify disconnect was called. + if got := us.disconnectCalls.Load(); got != 1 { + t.Errorf("disconnect calls = %d, want 1", got) + } +} + +func TestUplink_SessionIDAndDeviceID(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + // Before connect, values should be zero/empty. + if got := u.SessionID(); got != "" { + t.Errorf("SessionID() before connect = %q, want empty", got) + } + if got := u.DeviceID(); got != 0 { + t.Errorf("DeviceID() before connect = %d, want 0", got) + } + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + if got := u.SessionID(); !strings.HasPrefix(got, "test-session-") { + t.Errorf("SessionID() = %q, want prefix %q", got, "test-session-") + } + if got := u.DeviceID(); got != 42 { + t.Errorf("DeviceID() = %d, want 42", got) + } + + u.Close() +} + +func TestUplink_SendEnqueuesToBatcher(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Send multiple messages of different tiers. + u.Send(json.RawMessage(`{"type":"error","msg":"oops"}`), "error") // immediate + u.Send(json.RawMessage(`{"type":"claude_output","data":"hello"}`), "claude_output") // standard + u.Send(json.RawMessage(`{"type":"project_state","data":"state"}`), "project_state") // low priority + + // The immediate message triggers a flush that drains all tiers. + deadline := time.After(5 * time.Second) + for { + batches := us.getMessageBatches() + if len(batches) > 0 { + // All three messages should be in the first batch (immediate triggers drain of all). + if len(batches[0].Messages) != 3 { + t.Errorf("batch has %d messages, want 3", len(batches[0].Messages)) + } + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for batched messages") + case <-time.After(10 * time.Millisecond): + } + } +} + +func TestUplink_SendBeforeConnect(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + // Send before connect — should be silently dropped. + u.Send(json.RawMessage(`{"type":"error"}`), "error") + + // No crash, no messages sent. + time.Sleep(100 * time.Millisecond) + batches := us.getMessageBatches() + if len(batches) != 0 { + t.Errorf("expected 0 batches before connect, got %d", len(batches)) + } +} + +func TestUplink_ReceiveFromPusher(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + channel := fmt.Sprintf("private-chief-server.%d", u.DeviceID()) + + // Send 3 commands. + for i := 0; i < 3; i++ { + cmd := json.RawMessage(fmt.Sprintf(`{"type":"cmd","id":"%d"}`, i)) + if err := us.pusherSrv.sendCommand(channel, cmd); err != nil { + t.Fatalf("sendCommand(%d) failed: %v", i, err) + } + } + + // Receive all 3 in order. + for i := 0; i < 3; i++ { + select { + case received := <-u.Receive(): + var parsed map[string]interface{} + json.Unmarshal(received, &parsed) + want := fmt.Sprintf("%d", i) + if parsed["id"] != want { + t.Errorf("command %d: id = %v, want %v", i, parsed["id"], want) + } + case <-time.After(5 * time.Second): + t.Fatalf("timeout waiting for command %d", i) + } + } +} + +func TestUplink_Close_FlushesAndDisconnects(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Enqueue a low-priority message (wouldn't normally flush for 1s). + u.Send(json.RawMessage(`{"type":"settings","data":"config"}`), "settings") + + // Close should flush the remaining message before disconnecting. + if err := u.Close(); err != nil { + t.Fatalf("Close() failed: %v", err) + } + + // Verify the message was flushed. + batches := us.getMessageBatches() + if len(batches) == 0 { + t.Error("expected at least 1 batch after Close(), got 0") + } else { + found := false + for _, batch := range batches { + for _, msg := range batch.Messages { + var parsed map[string]interface{} + json.Unmarshal(msg, &parsed) + if parsed["type"] == "settings" { + found = true + } + } + } + if !found { + t.Error("settings message was not flushed on Close()") + } + } + + // Verify disconnect was called. + if got := us.disconnectCalls.Load(); got != 1 { + t.Errorf("disconnect calls = %d, want 1", got) + } +} + +func TestUplink_Close_DoubleCloseIsSafe(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // First close. + if err := u.Close(); err != nil { + t.Fatalf("first Close() failed: %v", err) + } + + // Second close should be a no-op. + if err := u.Close(); err != nil { + t.Fatalf("second Close() failed: %v", err) + } + + // Only one disconnect call. + if got := us.disconnectCalls.Load(); got != 1 { + t.Errorf("disconnect calls = %d, want 1", got) + } +} + +func TestUplink_SetAccessToken(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Update the token. + u.SetAccessToken("new-token-xyz") + + // The internal client should use the new token. + // We can verify this by checking the client's token directly. + u.client.mu.RLock() + token := u.client.accessToken + u.client.mu.RUnlock() + + if token != "new-token-xyz" { + t.Errorf("accessToken = %q, want %q", token, "new-token-xyz") + } +} + +func TestUplink_OnReconnectCallback(t *testing.T) { + us := newTestUplinkServer(t) + + var callCount atomic.Int32 + u := newTestUplink(t, us, WithOnReconnect(func() { + callCount.Add(1) + })) + + // Verify the callback is stored. + if u.onReconnect == nil { + t.Fatal("onReconnect should be set") + } + + // The callback itself is used by the reconnection logic (US-020). + // For now just verify it can be invoked. + u.onReconnect() + if got := callCount.Load(); got != 1 { + t.Errorf("callback count = %d, want 1", got) + } +} + +func TestUplink_ConnectFailure_HTTPError(t *testing.T) { + // HTTP server that rejects connect. + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "bad-token") + u := NewUplink(client) + + ctx := testContext(t) + err := u.Connect(ctx) + if err == nil { + t.Fatal("expected error when connect fails, got nil") + } + if !strings.Contains(err.Error(), "uplink connect") { + t.Errorf("error = %v, want containing 'uplink connect'", err) + } + + // Should not be connected. + if u.SessionID() != "" { + t.Error("SessionID should be empty after failed connect") + } +} + +func TestUplink_ConnectFailure_PusherError(t *testing.T) { + // HTTP server that succeeds for connect but Pusher server that rejects auth. + ps := newTestPusherServer(t) + ps.rejectSubscribe = true + reverbCfg := ps.reverbConfig() + + var disconnectCalled atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/api/device/connect": + json.NewEncoder(w).Encode(WelcomeResponse{ + Type: "welcome", + ProtocolVersion: 1, + DeviceID: 42, + SessionID: "sess-123", + Reverb: reverbCfg, + }) + case "/api/device/disconnect": + disconnectCalled.Add(1) + json.NewEncoder(w).Encode(map[string]string{"status": "disconnected"}) + case "/api/device/broadcasting/auth": + sig := GenerateAuthSignature(ps.appKey, ps.appSecret, "unused", "unused") + json.NewEncoder(w).Encode(pusherAuthResponse{Auth: sig}) + default: + http.NotFound(w, r) + } + })) + defer srv.Close() + + client := newTestClient(t, srv.URL, "test-token") + u := NewUplink(client) + + ctx := testContext(t) + err := u.Connect(ctx) + if err == nil { + t.Fatal("expected error when Pusher subscription fails, got nil") + } + if !strings.Contains(err.Error(), "pusher") { + t.Errorf("error = %v, want containing 'pusher'", err) + } + + // HTTP disconnect should have been called as cleanup. + time.Sleep(100 * time.Millisecond) + if got := disconnectCalled.Load(); got != 1 { + t.Errorf("disconnect calls after Pusher failure = %d, want 1", got) + } +} + +func TestUplink_ChannelName(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Verify the Pusher client subscribes to the correct channel. + select { + case channel := <-us.pusherSrv.onSubscribe: + expected := "private-chief-server.42" + if channel != expected { + t.Errorf("subscribed to %q, want %q", channel, expected) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } +} + +// --- Heartbeat Tests --- + +// newTestUplinkWithHeartbeat creates a connected Uplink with fast heartbeat timing for tests. +func newTestUplinkWithHeartbeat(t *testing.T, us *testUplinkServer, interval, retryDelay, skipWindow time.Duration, maxFails int, opts ...UplinkOption) *Uplink { + t.Helper() + + client := newTestClient(t, us.httpSrv.URL, "test-token") + u := NewUplink(client, opts...) + + // Override heartbeat timing for fast tests. + u.hbInterval = interval + u.hbRetryDelay = retryDelay + u.hbSkipWindow = skipWindow + u.hbMaxFails = maxFails + + return u +} + +func TestUplink_Heartbeat_SendsPeriodically(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplinkWithHeartbeat(t, us, 50*time.Millisecond, 10*time.Millisecond, 0, 3) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Wait for at least 3 heartbeats. + deadline := time.After(2 * time.Second) + for { + if us.heartbeatCalls.Load() >= 3 { + break + } + select { + case <-deadline: + t.Fatalf("expected at least 3 heartbeats, got %d", us.heartbeatCalls.Load()) + case <-time.After(10 * time.Millisecond): + } + } + + u.Close() +} + +func TestUplink_Heartbeat_StopsOnClose(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplinkWithHeartbeat(t, us, 50*time.Millisecond, 10*time.Millisecond, 0, 3) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Wait for at least 1 heartbeat. + deadline := time.After(2 * time.Second) + for { + if us.heartbeatCalls.Load() >= 1 { + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for first heartbeat") + case <-time.After(10 * time.Millisecond): + } + } + + // Close the uplink. + u.Close() + + // Record count and wait to confirm no more heartbeats are sent. + countAfterClose := us.heartbeatCalls.Load() + time.Sleep(200 * time.Millisecond) + + if got := us.heartbeatCalls.Load(); got != countAfterClose { + t.Errorf("heartbeat calls after close: got %d more (total %d), want 0 more", got-countAfterClose, got) + } +} + +func TestUplink_Heartbeat_SkipsWhenMessagesSentRecently(t *testing.T) { + us := newTestUplinkServer(t) + // skipWindow of 5s — any message sent within 5s skips heartbeat. + u := newTestUplinkWithHeartbeat(t, us, 50*time.Millisecond, 10*time.Millisecond, 5*time.Second, 3) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Send a message to trigger the lastSendTime update. + msg := json.RawMessage(`{"type":"run_complete","data":"done"}`) + u.Send(msg, "run_complete") + + // Wait for the message batch to be sent (sets lastSendTime). + deadline := time.After(2 * time.Second) + for { + batches := us.getMessageBatches() + if len(batches) > 0 { + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for message batch") + case <-time.After(10 * time.Millisecond): + } + } + + // Record heartbeat count now. + countBeforeSkip := us.heartbeatCalls.Load() + + // Wait 300ms — multiple heartbeat intervals would have passed (50ms each). + time.Sleep(300 * time.Millisecond) + + // Heartbeats should have been skipped because lastSendTime is recent. + countAfterWait := us.heartbeatCalls.Load() + if countAfterWait != countBeforeSkip { + t.Errorf("expected heartbeats to be skipped, but %d extra heartbeats were sent", countAfterWait-countBeforeSkip) + } + + u.Close() +} + +func TestUplink_Heartbeat_RetryOnTransientFailure(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplinkWithHeartbeat(t, us, 50*time.Millisecond, 10*time.Millisecond, 0, 3) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Make heartbeat fail with 500 (transient). + us.heartbeatStatus.Store(500) + + // Wait for heartbeat calls to accumulate (initial call + retry). + deadline := time.After(2 * time.Second) + for { + // Each heartbeat tick produces 2 calls (initial + retry). + if us.heartbeatCalls.Load() >= 4 { + break + } + select { + case <-deadline: + t.Fatalf("expected at least 4 heartbeat calls (2 ticks × 2 attempts), got %d", us.heartbeatCalls.Load()) + case <-time.After(10 * time.Millisecond): + } + } + + u.Close() +} + +func TestUplink_Heartbeat_RetrySucceedsResetsFailureCount(t *testing.T) { + us := newTestUplinkServer(t) + + var maxFailuresCalled atomic.Int32 + u := newTestUplinkWithHeartbeat(t, us, 50*time.Millisecond, 10*time.Millisecond, 0, 3) + u.onHeartbeatMaxFailures = func() { + maxFailuresCalled.Add(1) + } + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Let heartbeats succeed — no max failures callback should fire. + time.Sleep(200 * time.Millisecond) + + if got := maxFailuresCalled.Load(); got != 0 { + t.Errorf("maxFailures callback called %d times, want 0 (all heartbeats succeeded)", got) + } + + u.Close() +} + +func TestUplink_Heartbeat_MaxFailuresTriggersCallback(t *testing.T) { + us := newTestUplinkServer(t) + + maxFailuresCh := make(chan struct{}, 1) + u := newTestUplinkWithHeartbeat(t, us, 50*time.Millisecond, 10*time.Millisecond, 0, 3) + u.onHeartbeatMaxFailures = func() { + select { + case maxFailuresCh <- struct{}{}: + default: + } + } + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-us.pusherSrv.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Make all heartbeats fail. + us.heartbeatStatus.Store(500) + + // Wait for the max-failures callback. With 50ms interval and 10ms retry delay, + // each tick is ~60ms. We need 3 consecutive failures → ~180ms. + select { + case <-maxFailuresCh: + // Success — the callback was triggered. + case <-time.After(3 * time.Second): + t.Fatal("timeout waiting for heartbeat max failures callback") + } + + u.Close() +} + +func TestUplink_Heartbeat_ContextCancellationStops(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplinkWithHeartbeat(t, us, 50*time.Millisecond, 10*time.Millisecond, 0, 3) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for at least 1 heartbeat. + deadline := time.After(2 * time.Second) + for { + if us.heartbeatCalls.Load() >= 1 { + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for first heartbeat") + case <-time.After(10 * time.Millisecond): + } + } + + // Cancel the context. + cancel() + + countAfterCancel := us.heartbeatCalls.Load() + time.Sleep(200 * time.Millisecond) + + if got := us.heartbeatCalls.Load(); got != countAfterCancel { + t.Errorf("heartbeat calls after cancel: got %d more, want 0 more", got-countAfterCancel) + } + + // Clean up: close still works even though context is cancelled. + u.Close() +} + +// --- Reconnection Tests --- + +// waitForSubscription drains the onSubscribe channel and returns the channel name. +func waitForSubscription(t *testing.T, ps *testPusherServer) string { + t.Helper() + select { + case ch := <-ps.onSubscribe: + return ch + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for Pusher subscription") + return "" + } +} + +func TestUplink_Reconnect_PusherDisconnection(t *testing.T) { + us := newTestUplinkServer(t) + reconnectCh := make(chan struct{}, 1) + + u := newTestUplink(t, us, WithOnReconnect(func() { + select { + case reconnectCh <- struct{}{}: + default: + } + })) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for initial subscription. + waitForSubscription(t, us.pusherSrv) + + initialSession := u.SessionID() + connectsBefore := us.connectCalls.Load() + + // Close the Pusher WebSocket from the server side to simulate disconnection. + if err := us.pusherSrv.closeConnection(); err != nil { + t.Fatalf("closeConnection() failed: %v", err) + } + + // Wait for the OnReconnect callback — this means full reconnection succeeded. + select { + case <-reconnectCh: + case <-time.After(10 * time.Second): + t.Fatal("timeout waiting for reconnection after Pusher disconnect") + } + + // Wait for re-subscription. + waitForSubscription(t, us.pusherSrv) + + // Verify a new connect call was made. + if got := us.connectCalls.Load(); got <= connectsBefore { + t.Errorf("connect calls after reconnect = %d, want > %d", got, connectsBefore) + } + + // Verify session ID was refreshed. + newSession := u.SessionID() + if newSession == initialSession { + t.Errorf("session ID should change after reconnection, got same: %q", newSession) + } + + // Verify commands can still be received after reconnection. + channel := fmt.Sprintf("private-chief-server.%d", u.DeviceID()) + cmd := json.RawMessage(`{"type":"post_reconnect_cmd"}`) + // Give the Pusher client a moment to be ready. + time.Sleep(100 * time.Millisecond) + if err := us.pusherSrv.sendCommand(channel, cmd); err != nil { + t.Fatalf("sendCommand after reconnect failed: %v", err) + } + + select { + case received := <-u.Receive(): + var parsed map[string]interface{} + json.Unmarshal(received, &parsed) + if parsed["type"] != "post_reconnect_cmd" { + t.Errorf("received type = %v, want post_reconnect_cmd", parsed["type"]) + } + case <-time.After(5 * time.Second): + t.Fatal("timeout receiving command after reconnection") + } +} + +func TestUplink_Reconnect_HeartbeatFailuresTriggersReconnect(t *testing.T) { + us := newTestUplinkServer(t) + reconnectCh := make(chan struct{}, 1) + + // Use fast heartbeat timing to trigger reconnection quickly. + client := newTestClient(t, us.httpSrv.URL, "test-token") + u := NewUplink(client, WithOnReconnect(func() { + select { + case reconnectCh <- struct{}{}: + default: + } + })) + u.hbInterval = 50 * time.Millisecond + u.hbRetryDelay = 10 * time.Millisecond + u.hbSkipWindow = 0 + u.hbMaxFails = 2 + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for initial subscription. + waitForSubscription(t, us.pusherSrv) + + initialSession := u.SessionID() + + // Make heartbeats fail. + us.heartbeatStatus.Store(500) + + // Wait for reconnection (heartbeat failures → reconnect → OnReconnect). + select { + case <-reconnectCh: + case <-time.After(10 * time.Second): + t.Fatal("timeout waiting for reconnection triggered by heartbeat failures") + } + + // Verify session was refreshed. + newSession := u.SessionID() + if newSession == initialSession { + t.Errorf("session should change after reconnection, got same: %q", newSession) + } +} + +func TestUplink_Reconnect_HTTPConnectFailureThenRecover(t *testing.T) { + us := newTestUplinkServer(t) + reconnectCh := make(chan struct{}, 1) + + u := newTestUplink(t, us, WithOnReconnect(func() { + select { + case reconnectCh <- struct{}{}: + default: + } + })) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for initial subscription. + waitForSubscription(t, us.pusherSrv) + + // Make HTTP connect fail to simulate server outage during reconnection. + us.connectStatus.Store(500) + + // Trigger Pusher disconnection to start reconnection. + if err := us.pusherSrv.closeConnection(); err != nil { + t.Fatalf("closeConnection() failed: %v", err) + } + + // Wait a bit for the first reconnection attempt to fail. + time.Sleep(2 * time.Second) + + // Now restore HTTP connect. + us.connectStatus.Store(0) + + // Wait for successful reconnection. + select { + case <-reconnectCh: + case <-time.After(10 * time.Second): + t.Fatal("timeout waiting for reconnection after server recovery") + } + + // Multiple connect attempts should have been made (at least the failed one + the successful one). + if got := us.connectCalls.Load(); got < 3 { + t.Errorf("connect calls = %d, want >= 3 (initial + failed + success)", got) + } +} + +func TestUplink_Reconnect_AuthFailureTriggersTokenRefresh(t *testing.T) { + us := newTestUplinkServer(t) + reconnectCh := make(chan struct{}, 1) + refreshCh := make(chan struct{}, 1) + + client := newTestClient(t, us.httpSrv.URL, "test-token") + u := NewUplink(client, + WithOnReconnect(func() { + select { + case reconnectCh <- struct{}{}: + default: + } + }), + WithOnAuthFailure(func() error { + // Simulate token refresh: restore connect and update token. + us.connectStatus.Store(0) + client.SetAccessToken("refreshed-token") + select { + case refreshCh <- struct{}{}: + default: + } + return nil + }), + ) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for initial subscription. + waitForSubscription(t, us.pusherSrv) + + // Make connect return 401 to trigger auth failure during reconnection. + us.connectStatus.Store(401) + + // Trigger reconnection via Pusher disconnect. + if err := us.pusherSrv.closeConnection(); err != nil { + t.Fatalf("closeConnection() failed: %v", err) + } + + // Wait for the token refresh callback. + select { + case <-refreshCh: + case <-time.After(10 * time.Second): + t.Fatal("timeout waiting for token refresh callback") + } + + // Wait for successful reconnection with new token. + select { + case <-reconnectCh: + case <-time.After(10 * time.Second): + t.Fatal("timeout waiting for reconnection after token refresh") + } +} + +func TestUplink_Reconnect_OnReconnectCallbackFires(t *testing.T) { + us := newTestUplinkServer(t) + + var callbackCount atomic.Int32 + u := newTestUplink(t, us, WithOnReconnect(func() { + callbackCount.Add(1) + })) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for initial subscription. + waitForSubscription(t, us.pusherSrv) + + // No reconnect callback yet. + if got := callbackCount.Load(); got != 0 { + t.Errorf("callback count before reconnect = %d, want 0", got) + } + + // Trigger reconnection. + if err := us.pusherSrv.closeConnection(); err != nil { + t.Fatalf("closeConnection() failed: %v", err) + } + + // Wait for reconnection. + deadline := time.After(10 * time.Second) + for { + if callbackCount.Load() >= 1 { + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for OnReconnect callback") + case <-time.After(50 * time.Millisecond): + } + } + + // Drain subscription channel. + waitForSubscription(t, us.pusherSrv) + + if got := callbackCount.Load(); got != 1 { + t.Errorf("callback count = %d, want 1", got) + } +} + +func TestUplink_Reconnect_SendBuffersDuringOutage(t *testing.T) { + us := newTestUplinkServer(t) + reconnectCh := make(chan struct{}, 1) + + u := newTestUplink(t, us, WithOnReconnect(func() { + select { + case reconnectCh <- struct{}{}: + default: + } + })) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for initial subscription. + waitForSubscription(t, us.pusherSrv) + + // Send a message and verify it arrives. + u.Send(json.RawMessage(`{"type":"run_complete","data":"before"}`), "run_complete") + deadline := time.After(5 * time.Second) + for { + batches := us.getMessageBatches() + if len(batches) > 0 { + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for initial message") + case <-time.After(10 * time.Millisecond): + } + } + + batchCountBefore := len(us.getMessageBatches()) + + // Trigger reconnection. + if err := us.pusherSrv.closeConnection(); err != nil { + t.Fatalf("closeConnection() failed: %v", err) + } + + // Wait for reconnection to complete. + select { + case <-reconnectCh: + case <-time.After(10 * time.Second): + t.Fatal("timeout waiting for reconnection") + } + + // Wait for re-subscription. + waitForSubscription(t, us.pusherSrv) + + // Send a message after reconnection — it should be delivered. + u.Send(json.RawMessage(`{"type":"run_complete","data":"after"}`), "run_complete") + + deadline = time.After(5 * time.Second) + for { + batches := us.getMessageBatches() + if len(batches) > batchCountBefore { + // Found a new batch after reconnection. + lastBatch := batches[len(batches)-1] + found := false + for _, msg := range lastBatch.Messages { + var parsed map[string]interface{} + json.Unmarshal(msg, &parsed) + if parsed["data"] == "after" { + found = true + } + } + if !found { + t.Error("expected 'after' message in batch after reconnection") + } + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for message after reconnection") + case <-time.After(10 * time.Millisecond): + } + } +} + +func TestUplink_Reconnect_ConcurrentTriggersPrevented(t *testing.T) { + us := newTestUplinkServer(t) + + var reconnectCount atomic.Int32 + u := newTestUplink(t, us, WithOnReconnect(func() { + reconnectCount.Add(1) + })) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + // Wait for initial subscription. + waitForSubscription(t, us.pusherSrv) + + // Trigger reconnection multiple times concurrently — only one should run. + for i := 0; i < 5; i++ { + u.triggerReconnect("concurrent test") + } + + // Wait for exactly 1 reconnection. + deadline := time.After(10 * time.Second) + for { + if reconnectCount.Load() >= 1 { + break + } + select { + case <-deadline: + t.Fatal("timeout waiting for reconnection") + case <-time.After(50 * time.Millisecond): + } + } + + // Wait a bit more to confirm no additional reconnections happen. + waitForSubscription(t, us.pusherSrv) + time.Sleep(500 * time.Millisecond) + + if got := reconnectCount.Load(); got != 1 { + t.Errorf("reconnect count = %d, want 1 (concurrent triggers should be prevented)", got) + } +} + +func TestUplink_Reconnect_CloseDuringReconnectStops(t *testing.T) { + us := newTestUplinkServer(t) + + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for initial subscription. + waitForSubscription(t, us.pusherSrv) + + // Make connect fail so reconnection keeps retrying. + us.connectStatus.Store(500) + + // Trigger Pusher disconnection. + if err := us.pusherSrv.closeConnection(); err != nil { + t.Fatalf("closeConnection() failed: %v", err) + } + + // Wait a moment for reconnection to start. + time.Sleep(500 * time.Millisecond) + + // Close the uplink — should stop reconnection. + err := u.Close() + if err != nil { + // Close may return an error from the already-closed Pusher — that's OK. + t.Logf("Close() returned: %v (expected for already-closed Pusher)", err) + } + + // Verify it doesn't hang or panic. Record connect calls and wait. + connectsAfterClose := us.connectCalls.Load() + time.Sleep(2 * time.Second) + + // There should be no more connect attempts after Close(). + if got := us.connectCalls.Load(); got > connectsAfterClose+1 { + t.Errorf("connect calls after Close: %d more than expected (got %d, started at %d)", got-connectsAfterClose, got, connectsAfterClose) + } +} + +func TestUplink_Reconnect_LogsAttemptCountAndDelay(t *testing.T) { + // This test verifies the reconnection logic makes multiple attempts with backoff. + // We can't easily capture log output, so we verify the behavior indirectly + // by checking the number of connect attempts and timing. + us := newTestUplinkServer(t) + reconnectCh := make(chan struct{}, 1) + + u := newTestUplink(t, us, WithOnReconnect(func() { + select { + case reconnectCh <- struct{}{}: + default: + } + })) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + waitForSubscription(t, us.pusherSrv) + + // Make connect fail twice, then succeed. + failCount := atomic.Int32{} + originalStatus := us.connectStatus.Load() + us.connectStatus.Store(500) + + go func() { + for { + current := us.connectCalls.Load() + if current >= 3 { // initial + 2 failed + if failCount.Add(1) == 1 { + us.connectStatus.Store(int32(originalStatus)) + } + return + } + time.Sleep(50 * time.Millisecond) + } + }() + + // Trigger reconnection. + if err := us.pusherSrv.closeConnection(); err != nil { + t.Fatalf("closeConnection() failed: %v", err) + } + + select { + case <-reconnectCh: + case <-time.After(15 * time.Second): + t.Fatal("timeout waiting for reconnection after transient failures") + } + + // Multiple connect calls (initial + retries + success). + if got := us.connectCalls.Load(); got < 3 { + t.Errorf("connect calls = %d, want >= 3", got) + } +} + +func TestUplink_Reconnect_WithOnAuthFailureOption(t *testing.T) { + us := newTestUplinkServer(t) + + var authFailureCalled atomic.Int32 + u := newTestUplink(t, us, WithOnAuthFailure(func() error { + authFailureCalled.Add(1) + return nil + })) + + // Verify the option was set. + if u.onAuthFailure == nil { + t.Fatal("onAuthFailure should be set by WithOnAuthFailure option") + } + + // Invoke and verify. + if err := u.onAuthFailure(); err != nil { + t.Errorf("onAuthFailure() = %v, want nil", err) + } + if got := authFailureCalled.Load(); got != 1 { + t.Errorf("authFailureCalled = %d, want 1", got) + } +} + +func TestUplink_Reconnect_StableReceiveChannel(t *testing.T) { + // Verify that Receive() returns the same channel before and after reconnection. + us := newTestUplinkServer(t) + reconnectCh := make(chan struct{}, 1) + + u := newTestUplink(t, us, WithOnReconnect(func() { + select { + case reconnectCh <- struct{}{}: + default: + } + })) + + // Receive channel is created at construction time. + recvBefore := u.Receive() + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer u.Close() + + waitForSubscription(t, us.pusherSrv) + + // Verify it's the same channel after connect. + recvAfterConnect := u.Receive() + if recvBefore != recvAfterConnect { + t.Error("Receive() channel changed after Connect — should be stable") + } + + // Trigger reconnection. + if err := us.pusherSrv.closeConnection(); err != nil { + t.Fatalf("closeConnection() failed: %v", err) + } + + select { + case <-reconnectCh: + case <-time.After(10 * time.Second): + t.Fatal("timeout waiting for reconnection") + } + + waitForSubscription(t, us.pusherSrv) + + // Verify it's still the same channel after reconnection. + recvAfterReconnect := u.Receive() + if recvBefore != recvAfterReconnect { + t.Error("Receive() channel changed after reconnection — should be stable") + } +} + +// --- CloseWithTimeout Tests --- + +func TestUplink_CloseWithTimeout_NormalShutdown(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + waitForSubscription(t, us.pusherSrv) + + // Enqueue a low-priority message — flush should happen during close. + u.Send(json.RawMessage(`{"type":"settings","data":"config"}`), "settings") + + // Close with a generous timeout — should complete well within it. + start := time.Now() + if err := u.CloseWithTimeout(5 * time.Second); err != nil { + t.Fatalf("CloseWithTimeout() failed: %v", err) + } + elapsed := time.Since(start) + + // Should have completed quickly (under 2 seconds). + if elapsed > 2*time.Second { + t.Errorf("CloseWithTimeout took %s, expected < 2s", elapsed) + } + + // Verify the message was flushed. + batches := us.getMessageBatches() + found := false + for _, batch := range batches { + for _, msg := range batch.Messages { + var parsed map[string]interface{} + json.Unmarshal(msg, &parsed) + if parsed["type"] == "settings" { + found = true + } + } + } + if !found { + t.Error("settings message was not flushed during CloseWithTimeout") + } + + // Verify disconnect was called. + if got := us.disconnectCalls.Load(); got != 1 { + t.Errorf("disconnect calls = %d, want 1", got) + } +} + +func TestUplink_CloseWithTimeout_TimesOut(t *testing.T) { + // Create a server that hangs on message sending to simulate an unreachable server. + ps := newTestPusherServer(t) + reverbCfg := ps.reverbConfig() + + // hangDone is closed before the server closes — allows the hanging handler to exit + // so the httptest.Server can close cleanly. Registered AFTER srv.Close() in cleanup + // (LIFO order means hangDone closes first, then srv.Close proceeds). + hangDone := make(chan struct{}) + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + auth := r.Header.Get("Authorization") + if !strings.HasPrefix(auth, "Bearer ") { + w.WriteHeader(http.StatusUnauthorized) + return + } + + w.Header().Set("Content-Type", "application/json") + + switch r.URL.Path { + case "/api/device/connect": + json.NewEncoder(w).Encode(WelcomeResponse{ + Type: "welcome", + ProtocolVersion: 1, + DeviceID: 42, + SessionID: "sess-timeout-test", + Reverb: reverbCfg, + }) + + case "/api/device/disconnect": + json.NewEncoder(w).Encode(map[string]string{"status": "disconnected"}) + + case "/api/device/messages": + // Hang until test cleanup to simulate unreachable server during batcher flush. + select { + case <-hangDone: + case <-r.Context().Done(): + } + + case "/api/device/broadcasting/auth": + var body broadcastAuthRequest + json.NewDecoder(r.Body).Decode(&body) + sig := GenerateAuthSignature(ps.appKey, ps.appSecret, body.SocketID, body.ChannelName) + json.NewEncoder(w).Encode(pusherAuthResponse{Auth: sig}) + + default: + http.NotFound(w, r) + } + })) + // Register srv.Close first (runs second in LIFO), then hangDone (runs first). + t.Cleanup(func() { srv.Close() }) + t.Cleanup(func() { close(hangDone) }) + + client := newTestClient(t, srv.URL, "test-token") + u := NewUplink(client) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + select { + case <-ps.onSubscribe: + case <-time.After(5 * time.Second): + t.Fatal("timeout waiting for subscription") + } + + // Enqueue an immediate message — batcher will try to flush on Stop() + // but the server hangs, so the flush blocks. + u.Send(json.RawMessage(`{"type":"run_complete","data":"test"}`), "run_complete") + + // Give the batcher time to attempt sending (and block on the hanging server). + time.Sleep(200 * time.Millisecond) + + // CloseWithTimeout should return within the timeout even though flush is stuck. + start := time.Now() + err := u.CloseWithTimeout(1 * time.Second) + elapsed := time.Since(start) + + // Should complete near the timeout (1-4 seconds, accounting for the force-close 2s grace). + if elapsed > 5*time.Second { + t.Errorf("CloseWithTimeout took %s, expected < 5s", elapsed) + } + + // No error is expected — timeout is handled internally. + if err != nil { + t.Logf("CloseWithTimeout returned: %v (acceptable)", err) + } + + t.Logf("CloseWithTimeout completed in %s", elapsed.Round(time.Millisecond)) +} + +func TestUplink_CloseWithTimeout_DoubleCloseIsSafe(t *testing.T) { + us := newTestUplinkServer(t) + u := newTestUplink(t, us) + + ctx := testContext(t) + if err := u.Connect(ctx); err != nil { + t.Fatalf("Connect() failed: %v", err) + } + + // Wait for Pusher subscription. + waitForSubscription(t, us.pusherSrv) + + // First close. + if err := u.CloseWithTimeout(5 * time.Second); err != nil { + t.Fatalf("first CloseWithTimeout() failed: %v", err) + } + + // Second close should be a no-op. + if err := u.CloseWithTimeout(5 * time.Second); err != nil { + t.Fatalf("second CloseWithTimeout() failed: %v", err) + } + + // Only one disconnect call. + if got := us.disconnectCalls.Load(); got != 1 { + t.Errorf("disconnect calls = %d, want 1", got) + } +} diff --git a/internal/workspace/scanner.go b/internal/workspace/scanner.go new file mode 100644 index 00000000..7a323e11 --- /dev/null +++ b/internal/workspace/scanner.go @@ -0,0 +1,339 @@ +// Package workspace provides workspace directory scanning for discovering +// git repositories and tracking their state. +package workspace + +import ( + "context" + "encoding/json" + "fmt" + "log" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "time" + + "github.com/minicodemonkey/chief/internal/ws" +) + +// ScanInterval is how often the scanner re-scans the workspace. +const ScanInterval = 60 * time.Second + +// MessageSender is an interface for sending messages to the server. +type MessageSender interface { + Send(msg interface{}) error +} + +// Scanner discovers and tracks git repositories in a workspace directory. +type Scanner struct { + workspace string + sender MessageSender + interval time.Duration + + mu sync.RWMutex + projects []ws.ProjectSummary +} + +// New creates a new Scanner for the given workspace directory. +func New(workspace string, sender MessageSender) *Scanner { + return &Scanner{ + workspace: workspace, + sender: sender, + interval: ScanInterval, + } +} + +// WorkspacePath returns the workspace directory path. +func (s *Scanner) WorkspacePath() string { + return s.workspace +} + +// SetSender sets the message sender on the scanner. +// This allows creating the scanner before the sender is fully set up. +func (s *Scanner) SetSender(sender MessageSender) { + s.mu.Lock() + defer s.mu.Unlock() + s.sender = sender +} + +// Projects returns the current list of discovered projects. +func (s *Scanner) Projects() []ws.ProjectSummary { + s.mu.RLock() + defer s.mu.RUnlock() + result := make([]ws.ProjectSummary, len(s.projects)) + copy(result, s.projects) + return result +} + +// FindProject looks up a single project by name. +// Returns the project and true if found, or a zero value and false if not found. +func (s *Scanner) FindProject(name string) (ws.ProjectSummary, bool) { + s.mu.RLock() + defer s.mu.RUnlock() + for _, p := range s.projects { + if p.Name == name { + return p, true + } + } + return ws.ProjectSummary{}, false +} + +// Scan performs a single scan of the workspace directory and returns discovered projects. +func (s *Scanner) Scan() []ws.ProjectSummary { + entries, err := os.ReadDir(s.workspace) + if err != nil { + log.Printf("Warning: failed to read workspace directory: %v", err) + return nil + } + + var projects []ws.ProjectSummary + for _, entry := range entries { + if !entry.IsDir() { + continue + } + + dirPath := filepath.Join(s.workspace, entry.Name()) + + // Check for .git/ directory + gitDir := filepath.Join(dirPath, ".git") + info, err := os.Stat(gitDir) + if err != nil { + if os.IsPermission(err) { + log.Printf("Warning: permission denied accessing %s, skipping", dirPath) + } + continue + } + // .git can be a directory (normal repo) or a file (worktree) + if !info.IsDir() { + // .git file means it's a worktree link, still a valid git repo + _ = info + } + + project := scanProject(dirPath, entry.Name()) + projects = append(projects, project) + } + + return projects +} + +// ScanAndUpdate performs a scan and updates the stored project list. +// Returns true if the project list changed. +func (s *Scanner) ScanAndUpdate() bool { + newProjects := s.Scan() + + s.mu.Lock() + defer s.mu.Unlock() + + if projectsEqual(s.projects, newProjects) { + return false + } + + s.projects = newProjects + return true +} + +// Run starts the periodic scanning loop. It performs an initial scan immediately, +// then re-scans at the configured interval. It sends project_list updates over +// WebSocket when projects change. +func (s *Scanner) Run(ctx context.Context) { + // Initial scan + if s.ScanAndUpdate() { + s.sendProjectList() + } + + ticker := time.NewTicker(s.interval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if s.ScanAndUpdate() { + log.Println("Workspace projects changed, sending update") + s.sendProjectList() + } + } + } +} + +// sendProjectList sends a project_list message. +func (s *Scanner) sendProjectList() { + if s.sender == nil { + return + } + + s.mu.RLock() + projects := make([]ws.ProjectSummary, len(s.projects)) + copy(projects, s.projects) + s.mu.RUnlock() + + msg := ws.NewMessage(ws.TypeProjectList) + plMsg := ws.ProjectListMessage{ + Type: msg.Type, + ID: msg.ID, + Timestamp: msg.Timestamp, + Projects: projects, + } + + if err := s.sender.Send(plMsg); err != nil { + log.Printf("Error sending project_list: %v", err) + } +} + +// scanProject gathers information about a single project directory. +func scanProject(dirPath, name string) ws.ProjectSummary { + project := ws.ProjectSummary{ + Name: name, + Path: dirPath, + } + + // Check for .chief/ directory + chiefDir := filepath.Join(dirPath, ".chief") + if info, err := os.Stat(chiefDir); err == nil && info.IsDir() { + project.HasChief = true + } + + // Get git branch + branch, err := gitCurrentBranch(dirPath) + if err == nil { + project.Branch = branch + } + + // Get last commit info + commit, err := gitLastCommit(dirPath) + if err == nil { + project.Commit = commit + } + + // Get PRD list if .chief/ exists + if project.HasChief { + project.PRDs = scanPRDs(dirPath) + } + + return project +} + +// gitCurrentBranch returns the current branch for a git repo. +func gitCurrentBranch(dir string) (string, error) { + cmd := exec.Command("git", "rev-parse", "--abbrev-ref", "HEAD") + cmd.Dir = dir + output, err := cmd.Output() + if err != nil { + return "", err + } + return strings.TrimSpace(string(output)), nil +} + +// gitLastCommit returns the last commit info for a git repo. +func gitLastCommit(dir string) (ws.CommitInfo, error) { + // Use git log with a specific format to get hash, message, author, timestamp + cmd := exec.Command("git", "log", "-1", "--format=%H%n%s%n%an%n%aI") + cmd.Dir = dir + output, err := cmd.Output() + if err != nil { + return ws.CommitInfo{}, err + } + + lines := strings.SplitN(strings.TrimSpace(string(output)), "\n", 4) + if len(lines) < 4 { + return ws.CommitInfo{}, fmt.Errorf("unexpected git log output") + } + + return ws.CommitInfo{ + Hash: lines[0], + Message: lines[1], + Author: lines[2], + Timestamp: lines[3], + }, nil +} + +// scanPRDs discovers PRDs in a project's .chief/prds/ directory. +func scanPRDs(dirPath string) []ws.PRDInfo { + prdsDir := filepath.Join(dirPath, ".chief", "prds") + entries, err := os.ReadDir(prdsDir) + if err != nil { + return nil + } + + var prds []ws.PRDInfo + for _, entry := range entries { + if !entry.IsDir() { + continue + } + + prdJSON := filepath.Join(prdsDir, entry.Name(), "prd.json") + data, err := os.ReadFile(prdJSON) + if err != nil { + continue + } + + var prdData struct { + Project string `json:"project"` + UserStories []struct { + ID string `json:"id"` + Passes bool `json:"passes"` + } `json:"userStories"` + } + if err := json.Unmarshal(data, &prdData); err != nil { + continue + } + + total := len(prdData.UserStories) + passed := 0 + for _, s := range prdData.UserStories { + if s.Passes { + passed++ + } + } + + status := fmt.Sprintf("%d/%d", passed, total) + + prds = append(prds, ws.PRDInfo{ + ID: entry.Name(), + Name: prdData.Project, + StoryCount: total, + CompletionStatus: status, + }) + } + + return prds +} + +// projectsEqual compares two project lists for equality. +func projectsEqual(a, b []ws.ProjectSummary) bool { + if len(a) != len(b) { + return false + } + + // Build maps for comparison + aMap := make(map[string]ws.ProjectSummary, len(a)) + for _, p := range a { + aMap[p.Name] = p + } + + for _, pb := range b { + pa, ok := aMap[pb.Name] + if !ok { + return false + } + if pa.Path != pb.Path || + pa.HasChief != pb.HasChief || + pa.Branch != pb.Branch || + pa.Commit.Hash != pb.Commit.Hash || + len(pa.PRDs) != len(pb.PRDs) { + return false + } + // Compare PRDs + for i := range pa.PRDs { + if pa.PRDs[i].ID != pb.PRDs[i].ID || + pa.PRDs[i].StoryCount != pb.PRDs[i].StoryCount || + pa.PRDs[i].CompletionStatus != pb.PRDs[i].CompletionStatus { + return false + } + } + } + + return true +} diff --git a/internal/workspace/scanner_test.go b/internal/workspace/scanner_test.go new file mode 100644 index 00000000..f6c33c01 --- /dev/null +++ b/internal/workspace/scanner_test.go @@ -0,0 +1,504 @@ +package workspace + +import ( + "context" + "encoding/json" + "os" + "os/exec" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/ws" +) + +// testSender is a mock MessageSender that captures sent messages. +type testSender struct { + mu sync.Mutex + messages []json.RawMessage +} + +func (s *testSender) Send(msg interface{}) error { + data, err := json.Marshal(msg) + if err != nil { + return err + } + s.mu.Lock() + s.messages = append(s.messages, data) + s.mu.Unlock() + return nil +} + +func (s *testSender) getMessages() []json.RawMessage { + s.mu.Lock() + defer s.mu.Unlock() + cp := make([]json.RawMessage, len(s.messages)) + copy(cp, s.messages) + return cp +} + +func (s *testSender) waitForType(msgType string, timeout time.Duration) (json.RawMessage, bool) { + deadline := time.After(timeout) + for { + msgs := s.getMessages() + for _, raw := range msgs { + var m struct{ Type string `json:"type"` } + if json.Unmarshal(raw, &m) == nil && m.Type == msgType { + return raw, true + } + } + select { + case <-deadline: + return nil, false + case <-time.After(50 * time.Millisecond): + } + } +} + +// initGitRepo initializes a git repo with an initial commit. +func initGitRepo(t *testing.T, dir string) { + t.Helper() + cmds := [][]string{ + {"git", "init"}, + {"git", "config", "user.email", "test@example.com"}, + {"git", "config", "user.name", "Test User"}, + {"git", "commit", "--allow-empty", "-m", "initial commit"}, + } + for _, args := range cmds { + cmd := exec.Command(args[0], args[1:]...) + cmd.Dir = dir + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git command %v failed: %v\n%s", args, err, out) + } + } +} + +func TestScan_DiscoversGitRepos(t *testing.T) { + workspace := t.TempDir() + + // Create a git repo + repoDir := filepath.Join(workspace, "my-project") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + // Create a non-git directory (should be ignored) + nonGitDir := filepath.Join(workspace, "not-a-repo") + if err := os.MkdirAll(nonGitDir, 0o755); err != nil { + t.Fatal(err) + } + + // Create a file (should be ignored) + if err := os.WriteFile(filepath.Join(workspace, "some-file.txt"), []byte("hello"), 0o644); err != nil { + t.Fatal(err) + } + + scanner := New(workspace, nil) + projects := scanner.Scan() + + if len(projects) != 1 { + t.Fatalf("expected 1 project, got %d", len(projects)) + } + + p := projects[0] + if p.Name != "my-project" { + t.Errorf("expected name 'my-project', got %q", p.Name) + } + if p.Path != repoDir { + t.Errorf("expected path %q, got %q", repoDir, p.Path) + } + if p.HasChief { + t.Error("expected has_chief to be false") + } + if p.Branch == "" { + t.Error("expected branch to be set") + } + if p.Commit.Hash == "" { + t.Error("expected commit hash to be set") + } + if p.Commit.Message != "initial commit" { + t.Errorf("expected commit message 'initial commit', got %q", p.Commit.Message) + } + if p.Commit.Author != "Test User" { + t.Errorf("expected commit author 'Test User', got %q", p.Commit.Author) + } + if p.Commit.Timestamp == "" { + t.Error("expected commit timestamp to be set") + } +} + +func TestScan_DetectsChiefDirectory(t *testing.T) { + workspace := t.TempDir() + + repoDir := filepath.Join(workspace, "chief-project") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + // Create .chief/ directory + if err := os.MkdirAll(filepath.Join(repoDir, ".chief", "prds"), 0o755); err != nil { + t.Fatal(err) + } + + scanner := New(workspace, nil) + projects := scanner.Scan() + + if len(projects) != 1 { + t.Fatalf("expected 1 project, got %d", len(projects)) + } + if !projects[0].HasChief { + t.Error("expected has_chief to be true") + } +} + +func TestScan_DiscoversPRDs(t *testing.T) { + workspace := t.TempDir() + + repoDir := filepath.Join(workspace, "prd-project") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + // Create .chief/prds/my-feature/prd.json + prdDir := filepath.Join(repoDir, ".chief", "prds", "my-feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdData := map[string]interface{}{ + "project": "My Feature", + "userStories": []map[string]interface{}{ + {"id": "US-001", "passes": true}, + {"id": "US-002", "passes": false}, + {"id": "US-003", "passes": true}, + }, + } + data, _ := json.Marshal(prdData) + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), data, 0o644); err != nil { + t.Fatal(err) + } + + scanner := New(workspace, nil) + projects := scanner.Scan() + + if len(projects) != 1 { + t.Fatalf("expected 1 project, got %d", len(projects)) + } + + p := projects[0] + if len(p.PRDs) != 1 { + t.Fatalf("expected 1 PRD, got %d", len(p.PRDs)) + } + + prd := p.PRDs[0] + if prd.ID != "my-feature" { + t.Errorf("expected PRD ID 'my-feature', got %q", prd.ID) + } + if prd.Name != "My Feature" { + t.Errorf("expected PRD name 'My Feature', got %q", prd.Name) + } + if prd.StoryCount != 3 { + t.Errorf("expected 3 stories, got %d", prd.StoryCount) + } + if prd.CompletionStatus != "2/3" { + t.Errorf("expected completion '2/3', got %q", prd.CompletionStatus) + } +} + +func TestScan_MultipleProjects(t *testing.T) { + workspace := t.TempDir() + + for _, name := range []string{"alpha", "beta", "gamma"} { + dir := filepath.Join(workspace, name) + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, dir) + } + + scanner := New(workspace, nil) + projects := scanner.Scan() + + if len(projects) != 3 { + t.Fatalf("expected 3 projects, got %d", len(projects)) + } + + names := make(map[string]bool) + for _, p := range projects { + names[p.Name] = true + } + for _, name := range []string{"alpha", "beta", "gamma"} { + if !names[name] { + t.Errorf("expected project %q to be discovered", name) + } + } +} + +func TestScan_EmptyWorkspace(t *testing.T) { + workspace := t.TempDir() + + scanner := New(workspace, nil) + projects := scanner.Scan() + + if len(projects) != 0 { + t.Errorf("expected 0 projects, got %d", len(projects)) + } +} + +func TestScan_PermissionError(t *testing.T) { + // Skip if running as root (permissions are not enforced) + if os.Getuid() == 0 { + t.Skip("skipping permission test when running as root") + } + + workspace := t.TempDir() + + // Create a directory with .git inside, then remove traverse permission on parent + // so os.Stat on .git fails with permission denied + restrictedDir := filepath.Join(workspace, "restricted") + if err := os.MkdirAll(filepath.Join(restrictedDir, ".git"), 0o755); err != nil { + t.Fatal(err) + } + if err := os.Chmod(restrictedDir, 0o000); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + os.Chmod(restrictedDir, 0o755) + }) + + // Create a normal git repo too + goodDir := filepath.Join(workspace, "good-project") + if err := os.MkdirAll(goodDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, goodDir) + + scanner := New(workspace, nil) + projects := scanner.Scan() + + // Should still discover the good project even if restricted one has issues + if len(projects) != 1 { + t.Fatalf("expected 1 project, got %d", len(projects)) + } + if projects[0].Name != "good-project" { + t.Errorf("expected 'good-project', got %q", projects[0].Name) + } +} + +func TestScanAndUpdate_DetectsChanges(t *testing.T) { + workspace := t.TempDir() + + scanner := New(workspace, nil) + + // First scan: empty + changed := scanner.ScanAndUpdate() + if changed { + t.Error("expected no change on first scan of empty workspace") + } + + // Add a project + repoDir := filepath.Join(workspace, "new-project") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + // Second scan: should detect the new project + changed = scanner.ScanAndUpdate() + if !changed { + t.Error("expected change after adding a project") + } + + projects := scanner.Projects() + if len(projects) != 1 { + t.Fatalf("expected 1 project, got %d", len(projects)) + } + + // Third scan: no changes + changed = scanner.ScanAndUpdate() + if changed { + t.Error("expected no change on repeat scan") + } +} + +func TestScanAndUpdate_DetectsRemoval(t *testing.T) { + workspace := t.TempDir() + + repoDir := filepath.Join(workspace, "removable") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + scanner := New(workspace, nil) + scanner.ScanAndUpdate() + + if len(scanner.Projects()) != 1 { + t.Fatal("expected 1 project initially") + } + + // Remove the project + if err := os.RemoveAll(repoDir); err != nil { + t.Fatal(err) + } + + changed := scanner.ScanAndUpdate() + if !changed { + t.Error("expected change after removing project") + } + if len(scanner.Projects()) != 0 { + t.Error("expected 0 projects after removal") + } +} + +func TestRun_SendsProjectListOnChange(t *testing.T) { + workspace := t.TempDir() + + sender := &testSender{} + + // Create a project before starting the scanner + repoDir := filepath.Join(workspace, "starter") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + scanner := New(workspace, sender) + scanner.interval = 100 * time.Millisecond // Speed up for testing + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + // Run scanner in background + go scanner.Run(ctx) + + // Wait for initial project_list message + raw, ok := sender.waitForType(ws.TypeProjectList, 5*time.Second) + if !ok { + t.Fatal("timed out waiting for initial project_list message") + } + + var first ws.ProjectListMessage + if err := json.Unmarshal(raw, &first); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if len(first.Projects) != 1 { + t.Errorf("expected 1 project in initial scan, got %d", len(first.Projects)) + } else if first.Projects[0].Name != "starter" { + t.Errorf("expected project name 'starter', got %q", first.Projects[0].Name) + } + + // Add another project + newDir := filepath.Join(workspace, "newcomer") + if err := os.MkdirAll(newDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, newDir) + + // Wait for periodic scan to detect the new project + deadline := time.After(5 * time.Second) + for { + msgs := sender.getMessages() + for _, raw := range msgs { + var msg ws.ProjectListMessage + if json.Unmarshal(raw, &msg) == nil && msg.Type == ws.TypeProjectList && len(msg.Projects) == 2 { + return // Success + } + } + select { + case <-deadline: + t.Fatal("timed out waiting for updated project_list message") + case <-time.After(50 * time.Millisecond): + } + } +} + +func TestRun_StopsOnContextCancel(t *testing.T) { + workspace := t.TempDir() + + scanner := New(workspace, nil) + scanner.interval = 50 * time.Millisecond + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + scanner.Run(ctx) + close(done) + }() + + // Let it run briefly + time.Sleep(100 * time.Millisecond) + cancel() + + select { + case <-done: + // Good, it stopped + case <-time.After(2 * time.Second): + t.Fatal("scanner did not stop after context cancel") + } +} + +func TestProjectsEqual(t *testing.T) { + a := []ws.ProjectSummary{ + {Name: "proj1", Path: "/a/proj1", Branch: "main", Commit: ws.CommitInfo{Hash: "abc"}}, + {Name: "proj2", Path: "/a/proj2", Branch: "dev", Commit: ws.CommitInfo{Hash: "def"}}, + } + b := []ws.ProjectSummary{ + {Name: "proj1", Path: "/a/proj1", Branch: "main", Commit: ws.CommitInfo{Hash: "abc"}}, + {Name: "proj2", Path: "/a/proj2", Branch: "dev", Commit: ws.CommitInfo{Hash: "def"}}, + } + + if !projectsEqual(a, b) { + t.Error("expected equal project lists to be equal") + } + + // Change a commit hash + b[1].Commit.Hash = "changed" + if projectsEqual(a, b) { + t.Error("expected project lists with different commit hashes to be unequal") + } + + // Different lengths + if projectsEqual(a, a[:1]) { + t.Error("expected project lists of different lengths to be unequal") + } + + // Both nil/empty + if !projectsEqual(nil, nil) { + t.Error("expected two nil lists to be equal") + } + if !projectsEqual(nil, []ws.ProjectSummary{}) { + t.Error("expected nil and empty to be equal") + } +} + +func TestScan_GitBranch(t *testing.T) { + workspace := t.TempDir() + + repoDir := filepath.Join(workspace, "branched") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + // Create and switch to a feature branch + cmd := exec.Command("git", "checkout", "-b", "feature/test") + cmd.Dir = repoDir + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git checkout failed: %v\n%s", err, out) + } + + scanner := New(workspace, nil) + projects := scanner.Scan() + + if len(projects) != 1 { + t.Fatalf("expected 1 project, got %d", len(projects)) + } + if projects[0].Branch != "feature/test" { + t.Errorf("expected branch 'feature/test', got %q", projects[0].Branch) + } +} diff --git a/internal/workspace/watcher.go b/internal/workspace/watcher.go new file mode 100644 index 00000000..ffddbe79 --- /dev/null +++ b/internal/workspace/watcher.go @@ -0,0 +1,311 @@ +package workspace + +import ( + "context" + "log" + "os" + "path/filepath" + "strings" + "sync" + "time" + + "github.com/fsnotify/fsnotify" + "github.com/minicodemonkey/chief/internal/ws" +) + +// InactivityTimeout is how long a project remains active without interaction. +const InactivityTimeout = 10 * time.Minute + +// Watcher watches filesystem changes in the workspace using fsnotify. +// It watches the workspace root for new/removed projects and sets up deep +// watchers (.chief/, .git/HEAD) only for active projects. +type Watcher struct { + workspace string + scanner *Scanner + sender MessageSender + watcher *fsnotify.Watcher + + mu sync.Mutex + activeProjects map[string]*activeProject // project name → state + inactiveTimeout time.Duration +} + +// activeProject tracks an actively-watched project. +type activeProject struct { + name string + path string + lastActive time.Time + watching bool // whether deep watchers are set up +} + +// NewWatcher creates a new Watcher for the given workspace directory. +func NewWatcher(workspace string, scanner *Scanner, sender MessageSender) (*Watcher, error) { + fsw, err := fsnotify.NewWatcher() + if err != nil { + return nil, err + } + + return &Watcher{ + workspace: workspace, + scanner: scanner, + sender: sender, + watcher: fsw, + activeProjects: make(map[string]*activeProject), + inactiveTimeout: InactivityTimeout, + }, nil +} + +// Activate marks a project as active, setting up deep watchers if not already watching. +// Call this when a run is started, a Claude session is opened, or get_project is requested. +func (w *Watcher) Activate(projectName string) { + w.mu.Lock() + defer w.mu.Unlock() + + ap, exists := w.activeProjects[projectName] + if exists { + ap.lastActive = time.Now() + log.Printf("[debug] Project %q activity refreshed", projectName) + return + } + + // Find the project path from the scanner + projectPath := "" + for _, p := range w.scanner.Projects() { + if p.Name == projectName { + projectPath = p.Path + break + } + } + if projectPath == "" { + log.Printf("[debug] Project %q not found in scanner, cannot activate watcher", projectName) + return + } + + ap = &activeProject{ + name: projectName, + path: projectPath, + lastActive: time.Now(), + } + w.activeProjects[projectName] = ap + + w.setupDeepWatchers(ap) +} + +// setupDeepWatchers adds fsnotify watches for .chief/ and .git/HEAD for a project. +func (w *Watcher) setupDeepWatchers(ap *activeProject) { + if ap.watching { + return + } + + chiefDir := filepath.Join(ap.path, ".chief") + prdsDir := filepath.Join(ap.path, ".chief", "prds") + gitDir := filepath.Join(ap.path, ".git") + + // Watch .chief/ directory + if err := w.watcher.Add(chiefDir); err != nil { + log.Printf("[debug] Could not watch %s: %v", chiefDir, err) + } else { + log.Printf("[debug] Watching %s for project %q", chiefDir, ap.name) + } + + // Watch .chief/prds/ directory and each PRD subdirectory + // fsnotify does not recurse, so we must add each subdirectory explicitly + if err := w.watcher.Add(prdsDir); err != nil { + log.Printf("[debug] Could not watch %s: %v", prdsDir, err) + } else { + log.Printf("[debug] Watching %s for project %q", prdsDir, ap.name) + // Also watch each PRD subdirectory (e.g., .chief/prds/feature/) + entries, err := os.ReadDir(prdsDir) + if err == nil { + for _, entry := range entries { + if entry.IsDir() { + subDir := filepath.Join(prdsDir, entry.Name()) + if err := w.watcher.Add(subDir); err != nil { + log.Printf("[debug] Could not watch %s: %v", subDir, err) + } else { + log.Printf("[debug] Watching %s for project %q", subDir, ap.name) + } + } + } + } + } + + // Watch .git/ directory (for HEAD changes = branch switches) + if err := w.watcher.Add(gitDir); err != nil { + log.Printf("[debug] Could not watch %s: %v", gitDir, err) + } else { + log.Printf("[debug] Watching .git/ for project %q", ap.name) + } + + ap.watching = true +} + +// removeDeepWatchers removes fsnotify watches for a project. +func (w *Watcher) removeDeepWatchers(ap *activeProject) { + if !ap.watching { + return + } + + chiefDir := filepath.Join(ap.path, ".chief") + prdsDir := filepath.Join(ap.path, ".chief", "prds") + gitDir := filepath.Join(ap.path, ".git") + + _ = w.watcher.Remove(chiefDir) + _ = w.watcher.Remove(prdsDir) + _ = w.watcher.Remove(gitDir) + + ap.watching = false + log.Printf("[debug] Removed watchers for project %q", ap.name) +} + +// Run starts the watcher event loop. It watches the workspace root for project +// additions/removals and handles deep watcher events for active projects. +func (w *Watcher) Run(ctx context.Context) error { + // Watch workspace root for new/removed project directories + if err := w.watcher.Add(w.workspace); err != nil { + return err + } + log.Printf("[debug] Watching workspace root: %s", w.workspace) + + // Start inactivity checker + inactivityTicker := time.NewTicker(1 * time.Minute) + defer inactivityTicker.Stop() + + for { + select { + case <-ctx.Done(): + return w.watcher.Close() + + case event, ok := <-w.watcher.Events: + if !ok { + return nil + } + w.handleEvent(event) + + case err, ok := <-w.watcher.Errors: + if !ok { + return nil + } + log.Printf("Watcher error: %v", err) + + case <-inactivityTicker.C: + w.cleanupInactive() + } + } +} + +// handleEvent processes a single fsnotify event. +func (w *Watcher) handleEvent(event fsnotify.Event) { + path := event.Name + + // Check if this is a workspace-root-level event (new/removed project) + if filepath.Dir(path) == w.workspace { + if event.Has(fsnotify.Create) || event.Has(fsnotify.Remove) || event.Has(fsnotify.Rename) { + log.Printf("[debug] Workspace root change detected: %s (%s)", filepath.Base(path), event.Op) + // Trigger a re-scan to detect new/removed projects + if w.scanner.ScanAndUpdate() { + w.scanner.sendProjectList() + } + } + return + } + + // For deep watcher events, find which project this belongs to + projectName := w.projectForPath(path) + if projectName == "" { + return + } + + // Determine what changed + rel, err := filepath.Rel(w.workspace, path) + if err != nil { + return + } + parts := strings.SplitN(rel, string(filepath.Separator), 3) + if len(parts) < 2 { + return + } + + subPath := strings.Join(parts[1:], string(filepath.Separator)) + + switch { + case strings.HasPrefix(subPath, filepath.Join(".chief", "prds")): + log.Printf("[debug] PRD change detected in project %q: %s", projectName, subPath) + w.sendProjectState(projectName) + + case subPath == filepath.Join(".git", "HEAD"): + log.Printf("[debug] Git HEAD change detected in project %q", projectName) + w.sendProjectState(projectName) + + case strings.HasPrefix(subPath, ".chief"): + log.Printf("[debug] Chief config change in project %q: %s", projectName, subPath) + w.sendProjectState(projectName) + + case strings.HasPrefix(subPath, ".git"): + // Other .git changes (like refs) — check if HEAD changed + if strings.Contains(subPath, "HEAD") { + log.Printf("[debug] Git ref change in project %q: %s", projectName, subPath) + w.sendProjectState(projectName) + } + } +} + +// projectForPath finds which active project a file path belongs to. +func (w *Watcher) projectForPath(path string) string { + w.mu.Lock() + defer w.mu.Unlock() + + for _, ap := range w.activeProjects { + if strings.HasPrefix(path, ap.path+string(filepath.Separator)) || path == ap.path { + return ap.name + } + } + return "" +} + +// sendProjectState re-scans a single project and sends a project_state update. +func (w *Watcher) sendProjectState(projectName string) { + if w.sender == nil { + return + } + + // Re-scan the project to get updated state + w.scanner.ScanAndUpdate() + + // Find the project in the scanner's list + for _, p := range w.scanner.Projects() { + if p.Name == projectName { + msg := ws.NewMessage(ws.TypeProjectState) + psMsg := ws.ProjectStateMessage{ + Type: msg.Type, + ID: msg.ID, + Timestamp: msg.Timestamp, + Project: p, + } + if err := w.sender.Send(psMsg); err != nil { + log.Printf("Error sending project_state for %q: %v", projectName, err) + } + return + } + } +} + +// cleanupInactive removes watchers for projects that have been inactive. +func (w *Watcher) cleanupInactive() { + w.mu.Lock() + defer w.mu.Unlock() + + now := time.Now() + for name, ap := range w.activeProjects { + if now.Sub(ap.lastActive) > w.inactiveTimeout { + log.Printf("[debug] Project %q inactive for %s, removing watchers", name, w.inactiveTimeout) + w.removeDeepWatchers(ap) + delete(w.activeProjects, name) + } + } +} + +// Close closes the underlying fsnotify watcher. +func (w *Watcher) Close() error { + return w.watcher.Close() +} diff --git a/internal/workspace/watcher_test.go b/internal/workspace/watcher_test.go new file mode 100644 index 00000000..df00d51b --- /dev/null +++ b/internal/workspace/watcher_test.go @@ -0,0 +1,431 @@ +package workspace + +import ( + "context" + "encoding/json" + "os" + "os/exec" + "path/filepath" + "testing" + "time" + + "github.com/minicodemonkey/chief/internal/ws" +) + +func TestWatcher_WorkspaceRootChanges(t *testing.T) { + workspace := t.TempDir() + + // Create initial project + repoDir := filepath.Join(workspace, "existing") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + sender := &testSender{} + + scanner := New(workspace, sender) + scanner.ScanAndUpdate() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + watcher, err := NewWatcher(workspace, scanner, sender) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + + go watcher.Run(ctx) + + // Allow watcher to start + time.Sleep(100 * time.Millisecond) + + // Create a new project directory with git + newDir := filepath.Join(workspace, "new-project") + if err := os.MkdirAll(newDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, newDir) + + // Wait for the project_list message with 2 projects + deadline := time.After(5 * time.Second) + for { + msgs := sender.getMessages() + for _, raw := range msgs { + var msg struct { + Type string `json:"type"` + } + if json.Unmarshal(raw, &msg) == nil && msg.Type == ws.TypeProjectList { + var plMsg ws.ProjectListMessage + if err := json.Unmarshal(raw, &plMsg); err != nil { + t.Fatalf("unmarshal project_list: %v", err) + } + if len(plMsg.Projects) == 2 { + return // Success + } + } + } + select { + case <-deadline: + t.Fatal("timed out waiting for project_list with new project") + case <-time.After(50 * time.Millisecond): + } + } +} + +func TestWatcher_ActivateProject(t *testing.T) { + workspace := t.TempDir() + + // Create a project with .chief and .git + repoDir := filepath.Join(workspace, "my-project") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + if err := os.MkdirAll(filepath.Join(repoDir, ".chief", "prds"), 0o755); err != nil { + t.Fatal(err) + } + + scanner := New(workspace, nil) + scanner.ScanAndUpdate() + + watcher, err := NewWatcher(workspace, scanner, nil) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + defer watcher.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go watcher.Run(ctx) + + // Initially no active projects + watcher.mu.Lock() + if len(watcher.activeProjects) != 0 { + t.Error("expected 0 active projects initially") + } + watcher.mu.Unlock() + + // Activate a project + watcher.Activate("my-project") + + watcher.mu.Lock() + ap, exists := watcher.activeProjects["my-project"] + watcher.mu.Unlock() + + if !exists { + t.Fatal("expected my-project to be active") + } + if !ap.watching { + t.Error("expected deep watchers to be set up") + } +} + +func TestWatcher_ActivateUnknownProject(t *testing.T) { + workspace := t.TempDir() + + scanner := New(workspace, nil) + scanner.ScanAndUpdate() + + watcher, err := NewWatcher(workspace, scanner, nil) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + defer watcher.Close() + + // Activating unknown project should not panic or add to active list + watcher.Activate("nonexistent") + + watcher.mu.Lock() + defer watcher.mu.Unlock() + if len(watcher.activeProjects) != 0 { + t.Error("expected 0 active projects for unknown project") + } +} + +func TestWatcher_ActivateRefreshesActivity(t *testing.T) { + workspace := t.TempDir() + + repoDir := filepath.Join(workspace, "proj") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + scanner := New(workspace, nil) + scanner.ScanAndUpdate() + + watcher, err := NewWatcher(workspace, scanner, nil) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + defer watcher.Close() + + watcher.Activate("proj") + + watcher.mu.Lock() + firstActive := watcher.activeProjects["proj"].lastActive + watcher.mu.Unlock() + + time.Sleep(10 * time.Millisecond) + + watcher.Activate("proj") + + watcher.mu.Lock() + secondActive := watcher.activeProjects["proj"].lastActive + watcher.mu.Unlock() + + if !secondActive.After(firstActive) { + t.Error("expected lastActive to be refreshed on re-activation") + } +} + +func TestWatcher_InactivityCleanup(t *testing.T) { + workspace := t.TempDir() + + repoDir := filepath.Join(workspace, "proj") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + scanner := New(workspace, nil) + scanner.ScanAndUpdate() + + watcher, err := NewWatcher(workspace, scanner, nil) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + defer watcher.Close() + + // Use a very short timeout for testing + watcher.inactiveTimeout = 50 * time.Millisecond + + watcher.Activate("proj") + + watcher.mu.Lock() + if len(watcher.activeProjects) != 1 { + t.Fatal("expected 1 active project") + } + watcher.mu.Unlock() + + // Wait for the project to become inactive + time.Sleep(100 * time.Millisecond) + + watcher.cleanupInactive() + + watcher.mu.Lock() + defer watcher.mu.Unlock() + if len(watcher.activeProjects) != 0 { + t.Error("expected project to be cleaned up after inactivity timeout") + } +} + +func TestWatcher_ChiefPRDChangeSendsProjectState(t *testing.T) { + workspace := t.TempDir() + + // Create project with .chief/prds + repoDir := filepath.Join(workspace, "proj") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + prdDir := filepath.Join(repoDir, ".chief", "prds", "feature") + if err := os.MkdirAll(prdDir, 0o755); err != nil { + t.Fatal(err) + } + + prdData := map[string]interface{}{ + "project": "Feature", + "userStories": []map[string]interface{}{ + {"id": "US-001", "passes": false}, + }, + } + data, _ := json.Marshal(prdData) + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), data, 0o644); err != nil { + t.Fatal(err) + } + + sender := &testSender{} + + scanner := New(workspace, sender) + scanner.ScanAndUpdate() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + watcher, err := NewWatcher(workspace, scanner, sender) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + + go watcher.Run(ctx) + time.Sleep(100 * time.Millisecond) + + // Activate the project to set up deep watchers + watcher.Activate("proj") + time.Sleep(100 * time.Millisecond) + + // Modify the PRD file + prdData["userStories"] = []map[string]interface{}{ + {"id": "US-001", "passes": true}, + } + data, _ = json.Marshal(prdData) + if err := os.WriteFile(filepath.Join(prdDir, "prd.json"), data, 0o644); err != nil { + t.Fatal(err) + } + + // Wait for project_state message + deadline := time.After(5 * time.Second) + for { + msgs := sender.getMessages() + for _, raw := range msgs { + var msg struct { + Type string `json:"type"` + } + if json.Unmarshal(raw, &msg) == nil && msg.Type == ws.TypeProjectState { + var psMsg ws.ProjectStateMessage + if err := json.Unmarshal(raw, &psMsg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if psMsg.Project.Name == "proj" { + return // Success + } + } + } + select { + case <-deadline: + t.Fatal("timed out waiting for project_state message after PRD change") + case <-time.After(50 * time.Millisecond): + } + } +} + +func TestWatcher_GitHEADChangeSendsProjectState(t *testing.T) { + workspace := t.TempDir() + + repoDir := filepath.Join(workspace, "proj") + if err := os.MkdirAll(repoDir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, repoDir) + + sender := &testSender{} + + scanner := New(workspace, sender) + scanner.ScanAndUpdate() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + watcher, err := NewWatcher(workspace, scanner, sender) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + + go watcher.Run(ctx) + time.Sleep(100 * time.Millisecond) + + // Activate the project + watcher.Activate("proj") + time.Sleep(100 * time.Millisecond) + + // Switch branch (changes .git/HEAD) + cmd := exec.Command("git", "checkout", "-b", "feature/new-branch") + cmd.Dir = repoDir + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git checkout failed: %v\n%s", err, out) + } + + // Wait for project_state message + deadline := time.After(5 * time.Second) + for { + msgs := sender.getMessages() + for _, raw := range msgs { + var msg struct { + Type string `json:"type"` + } + if json.Unmarshal(raw, &msg) == nil && msg.Type == ws.TypeProjectState { + var psMsg ws.ProjectStateMessage + if err := json.Unmarshal(raw, &psMsg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if psMsg.Project.Name == "proj" { + return // Success + } + } + } + select { + case <-deadline: + t.Fatal("timed out waiting for project_state message after branch switch") + case <-time.After(50 * time.Millisecond): + } + } +} + +func TestWatcher_ContextCancellation(t *testing.T) { + workspace := t.TempDir() + + scanner := New(workspace, nil) + + watcher, err := NewWatcher(workspace, scanner, nil) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- watcher.Run(ctx) + }() + + // Let it start + time.Sleep(50 * time.Millisecond) + cancel() + + select { + case <-done: + // Good, it stopped + case <-time.After(2 * time.Second): + t.Fatal("watcher did not stop after context cancel") + } +} + +func TestWatcher_NoDeepWatchersForInactiveProjects(t *testing.T) { + workspace := t.TempDir() + + // Create two projects + for _, name := range []string{"active-proj", "inactive-proj"} { + dir := filepath.Join(workspace, name) + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatal(err) + } + initGitRepo(t, dir) + if err := os.MkdirAll(filepath.Join(dir, ".chief", "prds"), 0o755); err != nil { + t.Fatal(err) + } + } + + scanner := New(workspace, nil) + scanner.ScanAndUpdate() + + watcher, err := NewWatcher(workspace, scanner, nil) + if err != nil { + t.Fatalf("NewWatcher failed: %v", err) + } + defer watcher.Close() + + // Only activate one project + watcher.Activate("active-proj") + + watcher.mu.Lock() + defer watcher.mu.Unlock() + + if _, exists := watcher.activeProjects["active-proj"]; !exists { + t.Error("expected active-proj to be in active projects") + } + if _, exists := watcher.activeProjects["inactive-proj"]; exists { + t.Error("expected inactive-proj to NOT be in active projects") + } +} diff --git a/internal/ws/messages.go b/internal/ws/messages.go new file mode 100644 index 00000000..54c46bdd --- /dev/null +++ b/internal/ws/messages.go @@ -0,0 +1,631 @@ +package ws + +import ( + "crypto/rand" + "encoding/json" + "fmt" + "time" +) + +// ProtocolVersion is the current protocol version. +const ProtocolVersion = 1 + +// Message represents a protocol message envelope. +type Message struct { + Type string `json:"type"` + ID string `json:"id,omitempty"` + Timestamp string `json:"timestamp,omitempty"` + Raw json.RawMessage `json:"-"` +} + +// NewMessage creates a new message envelope with type, UUID, and ISO8601 timestamp. +func NewMessage(msgType string) Message { + return Message{ + Type: msgType, + ID: newUUID(), + Timestamp: time.Now().UTC().Format(time.RFC3339), + } +} + +// newUUID generates a random UUID v4 string. +func newUUID() string { + var uuid [16]byte + _, _ = rand.Read(uuid[:]) + // Set version 4 bits. + uuid[6] = (uuid[6] & 0x0f) | 0x40 + // Set variant bits. + uuid[8] = (uuid[8] & 0x3f) | 0x80 + return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", + uuid[0:4], uuid[4:6], uuid[6:8], uuid[8:10], uuid[10:16]) +} + +// Error codes for protocol error messages. +const ( + ErrCodeAuthFailed = "AUTH_FAILED" + ErrCodeProjectNotFound = "PROJECT_NOT_FOUND" + ErrCodePRDNotFound = "PRD_NOT_FOUND" + ErrCodeRunAlreadyActive = "RUN_ALREADY_ACTIVE" + ErrCodeRunNotActive = "RUN_NOT_ACTIVE" + ErrCodeSessionNotFound = "SESSION_NOT_FOUND" + ErrCodeCloneFailed = "CLONE_FAILED" + ErrCodeQuotaExhausted = "QUOTA_EXHAUSTED" + ErrCodeFilesystemError = "FILESYSTEM_ERROR" + ErrCodeClaudeError = "CLAUDE_ERROR" + ErrCodeUpdateFailed = "UPDATE_FAILED" + ErrCodeIncompatibleVersion = "INCOMPATIBLE_VERSION" + ErrCodeRateLimited = "RATE_LIMITED" +) + +// Message type constants for the protocol catalog. +const ( + // Server → Web App message types. + TypeHello = "hello" + TypeStateSnapshot = "state_snapshot" + TypeProjectList = "project_list" + TypeProjectState = "project_state" + TypePRDContent = "prd_content" + TypeClaudeOutput = "claude_output" + TypeRunProgress = "run_progress" + TypeRunComplete = "run_complete" + TypeRunPaused = "run_paused" + TypeDiff = "diff" + TypeDiffsResponse = "diffs_response" + TypePRDsResponse = "prds_response" + TypeCloneProgress = "clone_progress" + TypeCloneComplete = "clone_complete" + TypeError = "error" + TypeQuotaExhausted = "quota_exhausted" + TypeLogLines = "log_lines" + TypeSessionTimeoutWarning = "session_timeout_warning" + TypeSessionExpired = "session_expired" + TypeSettings = "settings" + TypeSettingsResponse = "settings_response" + TypeSettingsUpdated = "settings_updated" + TypeUpdateAvailable = "update_available" + TypePRDOutput = "prd_output" + TypePRDResponseComplete = "prd_response_complete" + + // Web App → Server message types. + TypeWelcome = "welcome" + TypeIncompatible = "incompatible" + TypeListProjects = "list_projects" + TypeGetProject = "get_project" + TypeGetPRD = "get_prd" + TypeGetPRDs = "get_prds" + TypeNewPRD = "new_prd" + TypeRefinePRD = "refine_prd" + TypePRDMessage = "prd_message" + TypeClosePRDSession = "close_prd_session" + TypeStartRun = "start_run" + TypePauseRun = "pause_run" + TypeResumeRun = "resume_run" + TypeStopRun = "stop_run" + TypeCloneRepo = "clone_repo" + TypeCreateProject = "create_project" + TypeGetDiff = "get_diff" + TypeGetDiffs = "get_diffs" + TypeGetLogs = "get_logs" + TypeGetSettings = "get_settings" + TypeUpdateSettings = "update_settings" + TypeTriggerUpdate = "trigger_update" + TypePing = "ping" + + // Bidirectional. + TypePong = "pong" +) + +// --- Server → Web App messages --- + +// StateSnapshotMessage is sent on connect/reconnect with full state. +type StateSnapshotMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Projects []ProjectSummary `json:"projects"` + Runs []RunState `json:"runs"` + Sessions []SessionState `json:"sessions"` +} + +// ProjectSummary describes a project in the workspace. +type ProjectSummary struct { + Name string `json:"name"` + Path string `json:"path"` + HasChief bool `json:"has_chief"` + Branch string `json:"branch"` + Commit CommitInfo `json:"commit"` + PRDs []PRDInfo `json:"prds"` +} + +// CommitInfo describes a git commit. +type CommitInfo struct { + Hash string `json:"hash"` + Message string `json:"message"` + Author string `json:"author"` + Timestamp string `json:"timestamp"` +} + +// PRDInfo describes a PRD in a project. +type PRDInfo struct { + ID string `json:"id"` + Name string `json:"name"` + StoryCount int `json:"story_count"` + CompletionStatus string `json:"completion_status"` +} + +// RunState describes an active run. +type RunState struct { + Project string `json:"project"` + PRDID string `json:"prd_id"` + StoryID string `json:"story_id"` + Status string `json:"status"` + Iteration int `json:"iteration"` +} + +// SessionState describes an active Claude session. +type SessionState struct { + SessionID string `json:"session_id"` + Project string `json:"project"` + PRDID string `json:"prd_id"` +} + +// ProjectListMessage lists all discovered projects. +type ProjectListMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Projects []ProjectSummary `json:"projects"` +} + +// ProjectStateMessage returns state for a single project. +type ProjectStateMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project ProjectSummary `json:"project"` +} + +// PRDContentMessage returns PRD content and state. +type PRDContentMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` + Content string `json:"content"` + State interface{} `json:"state"` +} + +// ClaudeOutputMessage streams Claude output. +type ClaudeOutputMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + SessionID string `json:"session_id,omitempty"` + Project string `json:"project,omitempty"` + PRDID string `json:"prd_id,omitempty"` + StoryID string `json:"story_id,omitempty"` + Data string `json:"data"` + Done bool `json:"done"` +} + +// PRDOutputPayload is the payload of a PRD output message. +type PRDOutputPayload struct { + Content string `json:"content"` + SessionID string `json:"session_id"` + Project string `json:"project"` +} + +// PRDOutputMessage streams PRD session output (text chunks from Claude). +type PRDOutputMessage struct { + Type string `json:"type"` + Payload PRDOutputPayload `json:"payload"` +} + +// PRDResponseCompletePayload is the payload of a PRD response complete message. +type PRDResponseCompletePayload struct { + SessionID string `json:"session_id"` + Project string `json:"project"` +} + +// PRDResponseCompleteMessage signals that a PRD session's Claude process has finished. +type PRDResponseCompleteMessage struct { + Type string `json:"type"` + Payload PRDResponseCompletePayload `json:"payload"` +} + +// RunProgressMessage reports run state changes. +type RunProgressMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` + StoryID string `json:"story_id"` + Status string `json:"status"` + Iteration int `json:"iteration"` + Attempt int `json:"attempt"` +} + +// RunCompleteMessage reports run completion. +type RunCompleteMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` + StoriesCompleted int `json:"stories_completed"` + Duration string `json:"duration"` + PassCount int `json:"pass_count"` + FailCount int `json:"fail_count"` +} + +// RunPausedMessage reports a paused run. +type RunPausedMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` + StoryID string `json:"story_id"` + Reason string `json:"reason"` +} + +// DiffMessage contains a story's diff. +type DiffMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` + StoryID string `json:"story_id"` + Files []string `json:"files"` + DiffText string `json:"diff_text"` +} + +// CloneProgressMessage reports git clone progress. +type CloneProgressMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + URL string `json:"url"` + ProgressText string `json:"progress_text"` + Percent int `json:"percent"` +} + +// CloneCompleteMessage reports clone completion. +type CloneCompleteMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + URL string `json:"url"` + Success bool `json:"success"` + Error string `json:"error,omitempty"` + Project string `json:"project,omitempty"` +} + +// ErrorMessage reports an error. +type ErrorMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Code string `json:"code"` + Message string `json:"message"` + RequestID string `json:"request_id,omitempty"` +} + +// QuotaExhaustedMessage reports quota exhaustion. +type QuotaExhaustedMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Runs []string `json:"runs"` + Sessions []string `json:"sessions"` +} + +// LogLinesMessage returns log content. +type LogLinesMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` + StoryID string `json:"story_id"` + Lines []string `json:"lines"` + Level string `json:"level"` +} + +// SessionTimeoutWarningMessage warns of impending session timeout. +type SessionTimeoutWarningMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + SessionID string `json:"session_id"` + MinutesRemaining int `json:"minutes_remaining"` +} + +// SessionExpiredMessage reports that a session has timed out. +type SessionExpiredMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + SessionID string `json:"session_id"` + SavedState string `json:"saved_state,omitempty"` +} + +// SettingsMessage returns project settings. +type SettingsMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + MaxIterations int `json:"max_iterations"` + AutoCommit bool `json:"auto_commit"` + CommitPrefix string `json:"commit_prefix"` + ClaudeModel string `json:"claude_model"` + TestCommand string `json:"test_command"` +} + +// UpdateAvailableMessage reports an available update. +type UpdateAvailableMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + CurrentVersion string `json:"current_version"` + LatestVersion string `json:"latest_version"` +} + +// --- Web App → Server messages --- + +// ListProjectsMessage requests the project list. +type ListProjectsMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` +} + +// GetProjectMessage requests a single project's state. +type GetProjectMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` +} + +// GetPRDsMessage requests a list of all PRDs for a project. +type GetPRDsMessage struct { + Project string `json:"project"` +} + +// PRDsResponseMessage returns a list of PRDs for a project. +type PRDsResponseMessage struct { + Type string `json:"type"` + Payload PRDsResponsePayload `json:"payload"` +} + +// PRDsResponsePayload is the payload of a PRDs response. +type PRDsResponsePayload struct { + Project string `json:"project"` + PRDs []PRDItem `json:"prds"` +} + +// PRDItem describes a PRD in the response list. +type PRDItem struct { + ID string `json:"id"` + Name string `json:"name"` + StoryCount int `json:"story_count"` + Status string `json:"status"` +} + +// SettingsResponseMessage wraps settings for browser delivery. +type SettingsResponseMessage struct { + Type string `json:"type"` + Payload SettingsResponsePayload `json:"payload"` +} + +// SettingsResponsePayload is the payload of a settings response. +type SettingsResponsePayload struct { + Project string `json:"project"` + Settings SettingsData `json:"settings"` +} + +// SettingsData contains project settings fields. +type SettingsData struct { + MaxIterations int `json:"max_iterations"` + AutoCommit bool `json:"auto_commit"` + CommitPrefix string `json:"commit_prefix"` + ClaudeModel string `json:"claude_model"` + TestCommand string `json:"test_command"` +} + +// DiffsResponseMessage wraps diff data for browser delivery. +type DiffsResponseMessage struct { + Type string `json:"type"` + Payload DiffsResponsePayload `json:"payload"` +} + +// DiffsResponsePayload is the payload of a diffs response. +type DiffsResponsePayload struct { + Project string `json:"project"` + StoryID string `json:"story_id"` + Files []DiffFileDetail `json:"files"` +} + +// DiffFileDetail represents a single file's diff information. +type DiffFileDetail struct { + Filename string `json:"filename"` + Additions int `json:"additions"` + Deletions int `json:"deletions"` + Patch string `json:"patch"` +} + +// GetDiffsMessage requests diffs for a story (without requiring prd_id). +type GetDiffsMessage struct { + Project string `json:"project"` + StoryID string `json:"story_id"` +} + +// GetPRDMessage requests a PRD's content. +type GetPRDMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` +} + +// NewPRDMessage requests creation of a new PRD via Claude. +type NewPRDMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + SessionID string `json:"session_id"` + Message string `json:"message"` +} + +// RefinePRDMessage requests editing an existing PRD via Claude. +type RefinePRDMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + SessionID string `json:"session_id"` + PRDID string `json:"prd_id"` + Message string `json:"message"` +} + +// PRDMessageMessage sends a user message to an active PRD session. +type PRDMessageMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + SessionID string `json:"session_id"` + Message string `json:"message"` +} + +// ClosePRDSessionMessage closes a PRD session. +type ClosePRDSessionMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + SessionID string `json:"session_id"` + Save bool `json:"save"` +} + +// StartRunMessage starts a Ralph loop. +type StartRunMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` +} + +// PauseRunMessage pauses a running loop. +type PauseRunMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` +} + +// ResumeRunMessage resumes a paused loop. +type ResumeRunMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` +} + +// StopRunMessage stops a running loop. +type StopRunMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` +} + +// CloneRepoMessage requests a git clone. +type CloneRepoMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + URL string `json:"url"` + DirectoryName string `json:"directory_name,omitempty"` +} + +// CreateProjectMessage creates a new project. +type CreateProjectMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Name string `json:"name"` + GitInit bool `json:"git_init"` +} + +// GetDiffMessage requests a story's diff. +type GetDiffMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` + StoryID string `json:"story_id"` +} + +// GetLogsMessage requests log lines. +type GetLogsMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + PRDID string `json:"prd_id"` + StoryID string `json:"story_id,omitempty"` + Lines int `json:"lines,omitempty"` +} + +// GetSettingsMessage requests project settings. +type GetSettingsMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` +} + +// UpdateSettingsMessage updates project settings. +type UpdateSettingsMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` + Project string `json:"project"` + MaxIterations *int `json:"max_iterations,omitempty"` + AutoCommit *bool `json:"auto_commit,omitempty"` + CommitPrefix *string `json:"commit_prefix,omitempty"` + ClaudeModel *string `json:"claude_model,omitempty"` + TestCommand *string `json:"test_command,omitempty"` +} + +// TriggerUpdateMessage requests a self-update. +type TriggerUpdateMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` +} + +// PingMessage is a keepalive ping. +type PingMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` +} + +// PongMessage is a keepalive pong response. +type PongMessage struct { + Type string `json:"type"` + ID string `json:"id"` + Timestamp string `json:"timestamp"` +} diff --git a/internal/ws/messages_test.go b/internal/ws/messages_test.go new file mode 100644 index 00000000..efb8ed89 --- /dev/null +++ b/internal/ws/messages_test.go @@ -0,0 +1,1038 @@ +package ws + +import ( + "encoding/json" + "testing" +) + +func TestStateSnapshotRoundTrip(t *testing.T) { + msg := StateSnapshotMessage{ + Type: TypeStateSnapshot, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Projects: []ProjectSummary{ + { + Name: "my-project", + Path: "/home/user/projects/my-project", + HasChief: true, + Branch: "main", + Commit: CommitInfo{ + Hash: "abc123", + Message: "initial commit", + Author: "dev", + Timestamp: "2026-02-15T09:00:00Z", + }, + PRDs: []PRDInfo{ + {ID: "auth", Name: "Authentication", StoryCount: 5, CompletionStatus: "3/5"}, + }, + }, + }, + Runs: []RunState{{Project: "my-project", PRDID: "auth", StoryID: "US-003", Status: "running", Iteration: 2}}, + Sessions: []SessionState{{SessionID: "sess-1", Project: "my-project", PRDID: "auth"}}, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got StateSnapshotMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Type != TypeStateSnapshot { + t.Errorf("type = %q, want %q", got.Type, TypeStateSnapshot) + } + if len(got.Projects) != 1 { + t.Fatalf("projects count = %d, want 1", len(got.Projects)) + } + if got.Projects[0].Name != "my-project" { + t.Errorf("project name = %q, want %q", got.Projects[0].Name, "my-project") + } + if got.Projects[0].Commit.Hash != "abc123" { + t.Errorf("commit hash = %q, want %q", got.Projects[0].Commit.Hash, "abc123") + } + if len(got.Projects[0].PRDs) != 1 || got.Projects[0].PRDs[0].StoryCount != 5 { + t.Errorf("unexpected PRD info: %+v", got.Projects[0].PRDs) + } + if len(got.Runs) != 1 || got.Runs[0].Iteration != 2 { + t.Errorf("unexpected run state: %+v", got.Runs) + } + if len(got.Sessions) != 1 || got.Sessions[0].SessionID != "sess-1" { + t.Errorf("unexpected session state: %+v", got.Sessions) + } +} + +func TestClaudeOutputRoundTrip(t *testing.T) { + msg := ClaudeOutputMessage{ + Type: TypeClaudeOutput, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + SessionID: "session-123", + Project: "my-project", + Data: "Hello from Claude!\nLine 2.", + Done: false, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got ClaudeOutputMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Type != TypeClaudeOutput { + t.Errorf("type = %q, want %q", got.Type, TypeClaudeOutput) + } + if got.SessionID != "session-123" { + t.Errorf("session_id = %q, want %q", got.SessionID, "session-123") + } + if got.Data != "Hello from Claude!\nLine 2." { + t.Errorf("data = %q, want %q", got.Data, "Hello from Claude!\nLine 2.") + } + if got.Done { + t.Error("done should be false") + } +} + +func TestRunProgressRoundTrip(t *testing.T) { + msg := RunProgressMessage{ + Type: TypeRunProgress, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + StoryID: "US-003", + Status: "running", + Iteration: 3, + Attempt: 1, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got RunProgressMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Project != "my-project" { + t.Errorf("project = %q, want %q", got.Project, "my-project") + } + if got.StoryID != "US-003" { + t.Errorf("story_id = %q, want %q", got.StoryID, "US-003") + } + if got.Iteration != 3 { + t.Errorf("iteration = %d, want 3", got.Iteration) + } + if got.Attempt != 1 { + t.Errorf("attempt = %d, want 1", got.Attempt) + } +} + +func TestRunCompleteRoundTrip(t *testing.T) { + msg := RunCompleteMessage{ + Type: TypeRunComplete, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + StoriesCompleted: 5, + Duration: "12m34s", + PassCount: 4, + FailCount: 1, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got RunCompleteMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.StoriesCompleted != 5 { + t.Errorf("stories_completed = %d, want 5", got.StoriesCompleted) + } + if got.Duration != "12m34s" { + t.Errorf("duration = %q, want %q", got.Duration, "12m34s") + } + if got.PassCount != 4 || got.FailCount != 1 { + t.Errorf("pass/fail = %d/%d, want 4/1", got.PassCount, got.FailCount) + } +} + +func TestErrorMessageRoundTrip(t *testing.T) { + msg := ErrorMessage{ + Type: TypeError, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Code: ErrCodeProjectNotFound, + Message: "Project 'foobar' not found", + RequestID: "req-456", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got ErrorMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Code != ErrCodeProjectNotFound { + t.Errorf("code = %q, want %q", got.Code, ErrCodeProjectNotFound) + } + if got.Message != "Project 'foobar' not found" { + t.Errorf("message = %q", got.Message) + } + if got.RequestID != "req-456" { + t.Errorf("request_id = %q, want %q", got.RequestID, "req-456") + } +} + +func TestErrorMessageWithoutRequestID(t *testing.T) { + msg := ErrorMessage{ + Type: TypeError, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Code: ErrCodeClaudeError, + Message: "Claude process crashed", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + // Verify request_id is omitted. + var raw map[string]interface{} + json.Unmarshal(data, &raw) + if _, ok := raw["request_id"]; ok { + t.Error("request_id should be omitted when empty") + } +} + +func TestDiffMessageRoundTrip(t *testing.T) { + msg := DiffMessage{ + Type: TypeDiff, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + StoryID: "US-003", + Files: []string{"internal/auth/auth.go", "internal/auth/auth_test.go"}, + DiffText: "--- a/internal/auth/auth.go\n+++ b/internal/auth/auth.go\n@@ -1,3 +1,5 @@\n+// new code\n", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got DiffMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if len(got.Files) != 2 { + t.Fatalf("files count = %d, want 2", len(got.Files)) + } + if got.Files[0] != "internal/auth/auth.go" { + t.Errorf("files[0] = %q", got.Files[0]) + } + if got.DiffText == "" { + t.Error("diff_text should not be empty") + } +} + +func TestStartRunRoundTrip(t *testing.T) { + msg := StartRunMessage{ + Type: TypeStartRun, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got StartRunMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Type != TypeStartRun { + t.Errorf("type = %q, want %q", got.Type, TypeStartRun) + } + if got.Project != "my-project" { + t.Errorf("project = %q", got.Project) + } + if got.PRDID != "auth" { + t.Errorf("prd_id = %q", got.PRDID) + } +} + +func TestNewPRDRoundTrip(t *testing.T) { + msg := NewPRDMessage{ + Type: TypeNewPRD, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + SessionID: "session-abc", + Message: "Build an authentication system", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got NewPRDMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.SessionID != "session-abc" { + t.Errorf("session_id = %q", got.SessionID) + } + if got.Message != "Build an authentication system" { + t.Errorf("message = %q", got.Message) + } +} + +func TestRefinePRDRoundTrip(t *testing.T) { + msg := RefinePRDMessage{ + Type: TypeRefinePRD, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + SessionID: "session-abc", + PRDID: "feature-auth", + Message: "Add OAuth support", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got RefinePRDMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Type != TypeRefinePRD { + t.Errorf("type = %q, want %q", got.Type, TypeRefinePRD) + } + if got.SessionID != "session-abc" { + t.Errorf("session_id = %q", got.SessionID) + } + if got.PRDID != "feature-auth" { + t.Errorf("prd_id = %q", got.PRDID) + } + if got.Message != "Add OAuth support" { + t.Errorf("message = %q", got.Message) + } +} + +func TestCloneRepoRoundTrip(t *testing.T) { + msg := CloneRepoMessage{ + Type: TypeCloneRepo, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + URL: "git@github.com:user/repo.git", + DirectoryName: "my-repo", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got CloneRepoMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.URL != "git@github.com:user/repo.git" { + t.Errorf("url = %q", got.URL) + } + if got.DirectoryName != "my-repo" { + t.Errorf("directory_name = %q", got.DirectoryName) + } +} + +func TestCloneRepoOmitsEmptyDirectoryName(t *testing.T) { + msg := CloneRepoMessage{ + Type: TypeCloneRepo, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + URL: "git@github.com:user/repo.git", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var raw map[string]interface{} + json.Unmarshal(data, &raw) + if _, ok := raw["directory_name"]; ok { + t.Error("directory_name should be omitted when empty") + } +} + +func TestUpdateSettingsPartialFields(t *testing.T) { + // Only updating max_iterations and auto_commit. + maxIter := 10 + autoCommit := false + msg := UpdateSettingsMessage{ + Type: TypeUpdateSettings, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + MaxIterations: &maxIter, + AutoCommit: &autoCommit, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got UpdateSettingsMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.MaxIterations == nil || *got.MaxIterations != 10 { + t.Errorf("max_iterations = %v, want 10", got.MaxIterations) + } + if got.AutoCommit == nil || *got.AutoCommit != false { + t.Errorf("auto_commit = %v, want false", got.AutoCommit) + } + if got.CommitPrefix != nil { + t.Errorf("commit_prefix should be nil, got %v", got.CommitPrefix) + } + if got.ClaudeModel != nil { + t.Errorf("claude_model should be nil, got %v", got.ClaudeModel) + } + if got.TestCommand != nil { + t.Errorf("test_command should be nil, got %v", got.TestCommand) + } +} + +func TestSessionTimeoutWarningRoundTrip(t *testing.T) { + msg := SessionTimeoutWarningMessage{ + Type: TypeSessionTimeoutWarning, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + SessionID: "session-xyz", + MinutesRemaining: 5, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got SessionTimeoutWarningMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.MinutesRemaining != 5 { + t.Errorf("minutes_remaining = %d, want 5", got.MinutesRemaining) + } + if got.SessionID != "session-xyz" { + t.Errorf("session_id = %q", got.SessionID) + } +} + +func TestQuotaExhaustedRoundTrip(t *testing.T) { + msg := QuotaExhaustedMessage{ + Type: TypeQuotaExhausted, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Runs: []string{"run-1", "run-2"}, + Sessions: []string{"session-1"}, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got QuotaExhaustedMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if len(got.Runs) != 2 { + t.Errorf("runs count = %d, want 2", len(got.Runs)) + } + if len(got.Sessions) != 1 { + t.Errorf("sessions count = %d, want 1", len(got.Sessions)) + } +} + +func TestSettingsRoundTrip(t *testing.T) { + msg := SettingsMessage{ + Type: TypeSettings, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + MaxIterations: 5, + AutoCommit: true, + CommitPrefix: "feat:", + ClaudeModel: "claude-opus-4-6", + TestCommand: "go test ./...", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got SettingsMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.MaxIterations != 5 { + t.Errorf("max_iterations = %d, want 5", got.MaxIterations) + } + if !got.AutoCommit { + t.Error("auto_commit should be true") + } + if got.TestCommand != "go test ./..." { + t.Errorf("test_command = %q", got.TestCommand) + } +} + +func TestLogLinesRoundTrip(t *testing.T) { + msg := LogLinesMessage{ + Type: TypeLogLines, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + StoryID: "US-003", + Lines: []string{"Starting iteration 1...", "Running tests...", "All tests passed."}, + Level: "info", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got LogLinesMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if len(got.Lines) != 3 { + t.Fatalf("lines count = %d, want 3", len(got.Lines)) + } + if got.Level != "info" { + t.Errorf("level = %q, want %q", got.Level, "info") + } +} + +func TestRunPausedRoundTrip(t *testing.T) { + msg := RunPausedMessage{ + Type: TypeRunPaused, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + StoryID: "US-003", + Reason: "quota_exhausted", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got RunPausedMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Reason != "quota_exhausted" { + t.Errorf("reason = %q, want %q", got.Reason, "quota_exhausted") + } +} + +func TestCloneProgressRoundTrip(t *testing.T) { + msg := CloneProgressMessage{ + Type: TypeCloneProgress, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + URL: "git@github.com:user/repo.git", + ProgressText: "Receiving objects: 45%", + Percent: 45, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got CloneProgressMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Percent != 45 { + t.Errorf("percent = %d, want 45", got.Percent) + } + if got.ProgressText != "Receiving objects: 45%" { + t.Errorf("progress_text = %q", got.ProgressText) + } +} + +func TestUpdateAvailableRoundTrip(t *testing.T) { + msg := UpdateAvailableMessage{ + Type: TypeUpdateAvailable, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + CurrentVersion: "0.5.0", + LatestVersion: "0.5.1", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got UpdateAvailableMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.CurrentVersion != "0.5.0" { + t.Errorf("current_version = %q", got.CurrentVersion) + } + if got.LatestVersion != "0.5.1" { + t.Errorf("latest_version = %q", got.LatestVersion) + } +} + +func TestPingPongRoundTrip(t *testing.T) { + ping := PingMessage{ + Type: TypePing, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + } + + data, err := json.Marshal(ping) + if err != nil { + t.Fatalf("marshal ping: %v", err) + } + + var gotPing PingMessage + if err := json.Unmarshal(data, &gotPing); err != nil { + t.Fatalf("unmarshal ping: %v", err) + } + + if gotPing.Type != TypePing { + t.Errorf("ping type = %q, want %q", gotPing.Type, TypePing) + } + + pong := PongMessage{ + Type: TypePong, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + } + + data, err = json.Marshal(pong) + if err != nil { + t.Fatalf("marshal pong: %v", err) + } + + var gotPong PongMessage + if err := json.Unmarshal(data, &gotPong); err != nil { + t.Fatalf("unmarshal pong: %v", err) + } + + if gotPong.Type != TypePong { + t.Errorf("pong type = %q, want %q", gotPong.Type, TypePong) + } +} + +func TestGenericMessageEnvelopeParsing(t *testing.T) { + // Verify that any message can be parsed as the generic Message type for routing. + msg := RunProgressMessage{ + Type: TypeRunProgress, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + StoryID: "US-003", + Status: "running", + Iteration: 2, + Attempt: 1, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var envelope Message + if err := json.Unmarshal(data, &envelope); err != nil { + t.Fatalf("unmarshal envelope: %v", err) + } + + if envelope.Type != TypeRunProgress { + t.Errorf("type = %q, want %q", envelope.Type, TypeRunProgress) + } + if envelope.ID == "" { + t.Error("id should be set") + } + if envelope.Timestamp == "" { + t.Error("timestamp should be set") + } +} + +func TestClosePRDSessionRoundTrip(t *testing.T) { + msg := ClosePRDSessionMessage{ + Type: TypeClosePRDSession, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + SessionID: "session-abc", + Save: true, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got ClosePRDSessionMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if !got.Save { + t.Error("save should be true") + } + if got.SessionID != "session-abc" { + t.Errorf("session_id = %q", got.SessionID) + } +} + +func TestSessionExpiredRoundTrip(t *testing.T) { + msg := SessionExpiredMessage{ + Type: TypeSessionExpired, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + SessionID: "session-abc", + SavedState: "partial PRD content here", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got SessionExpiredMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.SavedState != "partial PRD content here" { + t.Errorf("saved_state = %q", got.SavedState) + } +} + +func TestCreateProjectRoundTrip(t *testing.T) { + msg := CreateProjectMessage{ + Type: TypeCreateProject, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Name: "new-project", + GitInit: true, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got CreateProjectMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Name != "new-project" { + t.Errorf("name = %q", got.Name) + } + if !got.GitInit { + t.Error("git_init should be true") + } +} + +func TestGetLogsWithOptionalFields(t *testing.T) { + // With all fields. + msg := GetLogsMessage{ + Type: TypeGetLogs, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + StoryID: "US-003", + Lines: 100, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got GetLogsMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.StoryID != "US-003" { + t.Errorf("story_id = %q", got.StoryID) + } + if got.Lines != 100 { + t.Errorf("lines = %d, want 100", got.Lines) + } + + // Without optional fields. + msg2 := GetLogsMessage{ + Type: TypeGetLogs, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + } + + data2, _ := json.Marshal(msg2) + var raw map[string]interface{} + json.Unmarshal(data2, &raw) + + if _, ok := raw["story_id"]; ok { + t.Error("story_id should be omitted when empty") + } + if v, ok := raw["lines"]; ok && v != float64(0) { + t.Error("lines should be omitted when zero") + } +} + +func TestCloneCompleteRoundTrip(t *testing.T) { + // Success case. + msg := CloneCompleteMessage{ + Type: TypeCloneComplete, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + URL: "git@github.com:user/repo.git", + Success: true, + Project: "repo", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got CloneCompleteMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if !got.Success { + t.Error("success should be true") + } + if got.Project != "repo" { + t.Errorf("project = %q", got.Project) + } + + // Failure case. + msg2 := CloneCompleteMessage{ + Type: TypeCloneComplete, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + URL: "git@github.com:user/repo.git", + Success: false, + Error: "repository not found", + } + + data2, _ := json.Marshal(msg2) + var got2 CloneCompleteMessage + json.Unmarshal(data2, &got2) + + if got2.Success { + t.Error("success should be false") + } + if got2.Error != "repository not found" { + t.Errorf("error = %q", got2.Error) + } +} + +func TestPRDContentRoundTrip(t *testing.T) { + msg := PRDContentMessage{ + Type: TypePRDContent, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + PRDID: "auth", + Content: "# Authentication PRD\n\nBuild a login system.", + State: map[string]interface{}{"stories": 5, "completed": 3}, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got PRDContentMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Content != "# Authentication PRD\n\nBuild a login system." { + t.Errorf("content = %q", got.Content) + } + if got.PRDID != "auth" { + t.Errorf("prd_id = %q", got.PRDID) + } +} + +func TestProjectListRoundTrip(t *testing.T) { + msg := ProjectListMessage{ + Type: TypeProjectList, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Projects: []ProjectSummary{ + {Name: "project-a", Path: "/home/user/projects/project-a", HasChief: true, Branch: "main"}, + {Name: "project-b", Path: "/home/user/projects/project-b", HasChief: false, Branch: "develop"}, + }, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got ProjectListMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if len(got.Projects) != 2 { + t.Fatalf("projects count = %d, want 2", len(got.Projects)) + } + if got.Projects[1].Branch != "develop" { + t.Errorf("projects[1].branch = %q", got.Projects[1].Branch) + } +} + +func TestPRDOutputRoundTrip(t *testing.T) { + msg := PRDOutputMessage{ + Type: TypePRDOutput, + Payload: PRDOutputPayload{ + Content: "Here is the PRD content\n", + SessionID: "session-123", + Project: "my-project", + }, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got PRDOutputMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Type != TypePRDOutput { + t.Errorf("type = %q, want %q", got.Type, TypePRDOutput) + } + if got.Payload.SessionID != "session-123" { + t.Errorf("payload.session_id = %q, want %q", got.Payload.SessionID, "session-123") + } + if got.Payload.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", got.Payload.Project, "my-project") + } + if got.Payload.Content != "Here is the PRD content\n" { + t.Errorf("payload.content = %q, want %q", got.Payload.Content, "Here is the PRD content\n") + } +} + +func TestPRDResponseCompleteRoundTrip(t *testing.T) { + msg := PRDResponseCompleteMessage{ + Type: TypePRDResponseComplete, + Payload: PRDResponseCompletePayload{ + SessionID: "session-123", + Project: "my-project", + }, + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got PRDResponseCompleteMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Type != TypePRDResponseComplete { + t.Errorf("type = %q, want %q", got.Type, TypePRDResponseComplete) + } + if got.Payload.SessionID != "session-123" { + t.Errorf("payload.session_id = %q, want %q", got.Payload.SessionID, "session-123") + } + if got.Payload.Project != "my-project" { + t.Errorf("payload.project = %q, want %q", got.Payload.Project, "my-project") + } +} + +func TestPRDMessageRoundTrip(t *testing.T) { + msg := PRDMessageMessage{ + Type: TypePRDMessage, + ID: newUUID(), + Timestamp: "2026-02-15T10:00:00Z", + Project: "my-project", + SessionID: "session-abc", + Message: "Add OAuth support to the PRD", + } + + data, err := json.Marshal(msg) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got PRDMessageMessage + if err := json.Unmarshal(data, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.Message != "Add OAuth support to the PRD" { + t.Errorf("message = %q", got.Message) + } +} diff --git a/internal/ws/ratelimit.go b/internal/ws/ratelimit.go new file mode 100644 index 00000000..5ca70e2f --- /dev/null +++ b/internal/ws/ratelimit.go @@ -0,0 +1,207 @@ +package ws + +import ( + "fmt" + "sync" + "time" +) + +// Rate limiter constants (hardcoded for V1). +const ( + // Global token bucket: 30 burst, 10/second sustained. + globalBurst = 30 + globalRate = 10.0 // tokens per second + globalInterval = time.Second / 10 + + // Expensive operations: 2 per minute. + expensiveLimit = 2 + expensiveWindow = time.Minute + expensiveInterval = expensiveWindow / 2 // 30 seconds between allowed ops +) + +// expensiveTypes are message types with stricter per-type rate limits. +var expensiveTypes = map[string]bool{ + TypeCloneRepo: true, + TypeStartRun: true, + TypeNewPRD: true, +} + +// exemptTypes are message types exempt from rate limiting. +var exemptTypes = map[string]bool{ + TypePing: true, +} + +// IsExpensiveType returns true if the message type has a stricter per-type rate limit. +func IsExpensiveType(msgType string) bool { + return expensiveTypes[msgType] +} + +// IsExemptType returns true if the message type is exempt from rate limiting. +func IsExemptType(msgType string) bool { + return exemptTypes[msgType] +} + +// tokenBucket implements a simple token bucket rate limiter. +type tokenBucket struct { + tokens float64 + capacity float64 + rate float64 // tokens per second + lastTime time.Time +} + +func newTokenBucket(capacity float64, rate float64) *tokenBucket { + return &tokenBucket{ + tokens: capacity, + capacity: capacity, + rate: rate, + lastTime: time.Now(), + } +} + +// allow checks if a token is available and consumes one if so. +// Returns true if allowed, false if rate limited. +func (tb *tokenBucket) allow(now time.Time) bool { + elapsed := now.Sub(tb.lastTime).Seconds() + tb.tokens += elapsed * tb.rate + if tb.tokens > tb.capacity { + tb.tokens = tb.capacity + } + tb.lastTime = now + + if tb.tokens >= 1 { + tb.tokens-- + return true + } + return false +} + +// retryAfter returns the duration until the next token is available. +func (tb *tokenBucket) retryAfter() time.Duration { + if tb.tokens >= 1 { + return 0 + } + needed := 1.0 - tb.tokens + return time.Duration(needed / tb.rate * float64(time.Second)) +} + +// expensiveTracker tracks per-type rate limiting for expensive operations. +type expensiveTracker struct { + timestamps []time.Time + limit int + window time.Duration +} + +func newExpensiveTracker(limit int, window time.Duration) *expensiveTracker { + return &expensiveTracker{ + limit: limit, + window: window, + } +} + +// allow checks if the operation is allowed within the rate limit window. +func (et *expensiveTracker) allow(now time.Time) bool { + // Remove expired timestamps + cutoff := now.Add(-et.window) + valid := et.timestamps[:0] + for _, ts := range et.timestamps { + if ts.After(cutoff) { + valid = append(valid, ts) + } + } + et.timestamps = valid + + if len(et.timestamps) >= et.limit { + return false + } + et.timestamps = append(et.timestamps, now) + return true +} + +// retryAfter returns the duration until the next operation would be allowed. +func (et *expensiveTracker) retryAfter(now time.Time) time.Duration { + if len(et.timestamps) < et.limit { + return 0 + } + oldest := et.timestamps[0] + return oldest.Add(et.window).Sub(now) +} + +// RateLimiter provides rate limiting for incoming WebSocket messages. +type RateLimiter struct { + mu sync.Mutex + global *tokenBucket + expensive map[string]*expensiveTracker +} + +// NewRateLimiter creates a new rate limiter with default settings. +func NewRateLimiter() *RateLimiter { + return &RateLimiter{ + global: newTokenBucket(globalBurst, globalRate), + expensive: make(map[string]*expensiveTracker), + } +} + +// RateLimitResult contains the result of a rate limit check. +type RateLimitResult struct { + Allowed bool + RetryAfter time.Duration +} + +// Allow checks if a message of the given type should be allowed. +// Returns RateLimitResult indicating whether the message is allowed and retry-after hint. +func (rl *RateLimiter) Allow(msgType string) RateLimitResult { + // Exempt types bypass all rate limiting + if IsExemptType(msgType) { + return RateLimitResult{Allowed: true} + } + + rl.mu.Lock() + defer rl.mu.Unlock() + + now := time.Now() + + // Check expensive operation limit first + if IsExpensiveType(msgType) { + tracker, ok := rl.expensive[msgType] + if !ok { + tracker = newExpensiveTracker(expensiveLimit, expensiveWindow) + rl.expensive[msgType] = tracker + } + if !tracker.allow(now) { + retryAfter := tracker.retryAfter(now) + return RateLimitResult{ + Allowed: false, + RetryAfter: retryAfter, + } + } + } + + // Check global rate limit + if !rl.global.allow(now) { + retryAfter := rl.global.retryAfter() + return RateLimitResult{ + Allowed: false, + RetryAfter: retryAfter, + } + } + + return RateLimitResult{Allowed: true} +} + +// Reset clears all rate limiter state (called on reconnection). +func (rl *RateLimiter) Reset() { + rl.mu.Lock() + defer rl.mu.Unlock() + + rl.global = newTokenBucket(globalBurst, globalRate) + rl.expensive = make(map[string]*expensiveTracker) +} + +// FormatRetryAfter returns a human-readable retry-after string. +func FormatRetryAfter(d time.Duration) string { + secs := int(d.Seconds()) + 1 // round up + if secs <= 0 { + secs = 1 + } + return fmt.Sprintf("%ds", secs) +} diff --git a/internal/ws/ratelimit_test.go b/internal/ws/ratelimit_test.go new file mode 100644 index 00000000..a6015da8 --- /dev/null +++ b/internal/ws/ratelimit_test.go @@ -0,0 +1,293 @@ +package ws + +import ( + "testing" + "time" +) + +func TestRateLimiter_AllowsNormalMessages(t *testing.T) { + rl := NewRateLimiter() + + // Should allow a burst of messages up to the burst limit + for i := 0; i < globalBurst; i++ { + result := rl.Allow(TypeGetProject) + if !result.Allowed { + t.Fatalf("expected message %d to be allowed within burst", i+1) + } + } +} + +func TestRateLimiter_BlocksAfterBurstExhausted(t *testing.T) { + rl := NewRateLimiter() + + // Exhaust the burst + for i := 0; i < globalBurst; i++ { + rl.Allow(TypeGetProject) + } + + // Next message should be blocked + result := rl.Allow(TypeGetProject) + if result.Allowed { + t.Fatal("expected message to be blocked after burst exhausted") + } + if result.RetryAfter <= 0 { + t.Error("expected positive retry-after duration") + } +} + +func TestRateLimiter_RefillsTokensOverTime(t *testing.T) { + rl := NewRateLimiter() + + // Exhaust the burst + for i := 0; i < globalBurst; i++ { + rl.Allow(TypeGetProject) + } + + // Manually advance the token bucket's last time to simulate time passing + rl.mu.Lock() + rl.global.lastTime = time.Now().Add(-200 * time.Millisecond) // should refill ~2 tokens + rl.mu.Unlock() + + result := rl.Allow(TypeGetProject) + if !result.Allowed { + t.Fatal("expected message to be allowed after token refill") + } +} + +func TestRateLimiter_PingExempt(t *testing.T) { + rl := NewRateLimiter() + + // Exhaust the burst + for i := 0; i < globalBurst+10; i++ { + rl.Allow(TypeGetProject) + } + + // Ping should still be allowed + result := rl.Allow(TypePing) + if !result.Allowed { + t.Fatal("expected ping to be exempt from rate limiting") + } +} + +func TestRateLimiter_ExpensiveOperationsLimited(t *testing.T) { + rl := NewRateLimiter() + + // First two clone_repo should be allowed + result1 := rl.Allow(TypeCloneRepo) + result2 := rl.Allow(TypeCloneRepo) + if !result1.Allowed || !result2.Allowed { + t.Fatal("expected first two expensive operations to be allowed") + } + + // Third should be blocked + result3 := rl.Allow(TypeCloneRepo) + if result3.Allowed { + t.Fatal("expected third expensive operation to be blocked") + } + if result3.RetryAfter <= 0 { + t.Error("expected positive retry-after for expensive operation") + } +} + +func TestRateLimiter_ExpensiveTypesIndependent(t *testing.T) { + rl := NewRateLimiter() + + // Use up clone_repo limit + rl.Allow(TypeCloneRepo) + rl.Allow(TypeCloneRepo) + + // start_run should still be allowed (independent tracker) + result := rl.Allow(TypeStartRun) + if !result.Allowed { + t.Fatal("expected start_run to be allowed independently of clone_repo") + } + + // new_prd should also be allowed + result = rl.Allow(TypeNewPRD) + if !result.Allowed { + t.Fatal("expected new_prd to be allowed independently") + } +} + +func TestRateLimiter_ExpensiveWindowExpiry(t *testing.T) { + rl := NewRateLimiter() + + // Use up the limit + rl.Allow(TypeStartRun) + rl.Allow(TypeStartRun) + + // Should be blocked + result := rl.Allow(TypeStartRun) + if result.Allowed { + t.Fatal("expected to be blocked") + } + + // Simulate time passing beyond the window + rl.mu.Lock() + tracker := rl.expensive[TypeStartRun] + for i := range tracker.timestamps { + tracker.timestamps[i] = time.Now().Add(-expensiveWindow - time.Second) + } + rl.mu.Unlock() + + // Should be allowed again + result = rl.Allow(TypeStartRun) + if !result.Allowed { + t.Fatal("expected to be allowed after window expires") + } +} + +func TestRateLimiter_Reset(t *testing.T) { + rl := NewRateLimiter() + + // Exhaust burst and expensive limits + for i := 0; i < globalBurst+5; i++ { + rl.Allow(TypeGetProject) + } + rl.Allow(TypeCloneRepo) + rl.Allow(TypeCloneRepo) + + // Verify blocked + result := rl.Allow(TypeGetProject) + if result.Allowed { + t.Fatal("expected blocked before reset") + } + result = rl.Allow(TypeCloneRepo) + if result.Allowed { + t.Fatal("expected expensive blocked before reset") + } + + // Reset + rl.Reset() + + // Should be allowed again + result = rl.Allow(TypeGetProject) + if !result.Allowed { + t.Fatal("expected allowed after reset") + } + result = rl.Allow(TypeCloneRepo) + if !result.Allowed { + t.Fatal("expected expensive allowed after reset") + } +} + +func TestRateLimiter_ExpensiveAlsoConsumesGlobal(t *testing.T) { + rl := NewRateLimiter() + + // Exhaust global bucket + for i := 0; i < globalBurst; i++ { + rl.Allow(TypeGetProject) + } + + // Expensive operation should be blocked by global limit even though + // the expensive tracker would allow it + result := rl.Allow(TypeCloneRepo) + if result.Allowed { + t.Fatal("expected expensive operation to be blocked by global limit") + } +} + +func TestFormatRetryAfter(t *testing.T) { + tests := []struct { + d time.Duration + want string + }{ + {100 * time.Millisecond, "1s"}, + {500 * time.Millisecond, "1s"}, + {1500 * time.Millisecond, "2s"}, + {30 * time.Second, "31s"}, + {0, "1s"}, + } + + for _, tt := range tests { + got := FormatRetryAfter(tt.d) + if got != tt.want { + t.Errorf("FormatRetryAfter(%v) = %q, want %q", tt.d, got, tt.want) + } + } +} + +func TestIsExpensiveType(t *testing.T) { + if !IsExpensiveType(TypeCloneRepo) { + t.Error("expected clone_repo to be expensive") + } + if !IsExpensiveType(TypeStartRun) { + t.Error("expected start_run to be expensive") + } + if !IsExpensiveType(TypeNewPRD) { + t.Error("expected new_prd to be expensive") + } + if IsExpensiveType(TypeGetProject) { + t.Error("expected get_project to NOT be expensive") + } + if IsExpensiveType(TypePing) { + t.Error("expected ping to NOT be expensive") + } +} + +func TestIsExemptType(t *testing.T) { + if !IsExemptType(TypePing) { + t.Error("expected ping to be exempt") + } + if IsExemptType(TypeGetProject) { + t.Error("expected get_project to NOT be exempt") + } +} + +func TestTokenBucket_RetryAfter(t *testing.T) { + tb := newTokenBucket(1, 10) // 1 burst, 10/sec + + // Use the token + now := time.Now() + tb.allow(now) + + // Should need ~100ms for next token + retryAfter := tb.retryAfter() + if retryAfter <= 0 { + t.Error("expected positive retry-after") + } + if retryAfter > 200*time.Millisecond { + t.Errorf("retry-after too large: %v", retryAfter) + } +} + +func TestExpensiveTracker_RetryAfter(t *testing.T) { + et := newExpensiveTracker(2, time.Minute) + + now := time.Now() + et.allow(now) + et.allow(now.Add(time.Second)) + + // Should be blocked with retry-after pointing to when oldest expires + retryAfter := et.retryAfter(now.Add(2 * time.Second)) + if retryAfter <= 0 { + t.Error("expected positive retry-after") + } + // Should be about 58 seconds (60 - 2 seconds elapsed) + if retryAfter < 55*time.Second || retryAfter > 61*time.Second { + t.Errorf("unexpected retry-after: %v", retryAfter) + } +} + +func TestRateLimiter_ConcurrentAccess(t *testing.T) { + rl := NewRateLimiter() + + done := make(chan struct{}) + for i := 0; i < 10; i++ { + go func() { + defer func() { done <- struct{}{} }() + for j := 0; j < 100; j++ { + rl.Allow(TypeGetProject) + rl.Allow(TypeCloneRepo) + rl.Allow(TypePing) + } + }() + } + + for i := 0; i < 10; i++ { + <-done + } + + // Just ensure no panics or races + rl.Reset() +}