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
-
-
-
-
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//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 US-001"}]}}'
+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//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 US-001"}]}}'
+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 US-001"}]}}'
+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) 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, "") {
- 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. "}}`
- 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. "}}`
- 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, "") {
- return &Event{Type: EventComplete, Text: text}
- }
- if strings.Contains(text, "") {
- 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. "}]},"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. "}]},"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 , 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 event.
-func TestLoop_ChiefDoneEvent(t *testing.T) {
- l := NewLoop("/test/prd.json", "test", 5, testProvider)
+// TestLoop_ChiefCompleteEvent tests detection of 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! "}]}}` + "\n")
+ w.WriteString(`{"type":"assistant","message":{"content":[{"type":"text","text":"All done! "}]}}` + "\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 ")
+ if !hasComplete {
+ t.Error("Expected Complete event for ")
}
-
- l.mu.Lock()
- if !l.sawStoryDone {
- t.Error("Expected sawStoryDone to be true after processing ")
- }
- 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, "") {
- 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 .
- 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 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 tag
- if strings.Contains(text, "") {
+ // Check for tag
+ if strings.Contains(text, "") {
return &Event{
- Type: EventStoryDone,
+ Type: EventComplete,
Text: text,
}
}
+ // Check for story markers using ralph-status tags
+ if storyID := extractStoryID(text, "", ""); 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! "}]}}`,
- `{"type":"assistant","message":{"content":[{"type":"text","text":""}]}}`,
- `{"type":"assistant","message":{"content":[{"type":"text","text":"Done\n\nGoodbye"}]}}`,
+ `{"type":"assistant","message":{"content":[{"type":"text","text":"All stories complete! "}]}}`,
+ `{"type":"assistant","message":{"content":[{"type":"text","text":""}]}}`,
+ `{"type":"assistant","message":{"content":[{"type":"text","text":"Done\n\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.\nUS-003\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: "US-001",
+ startTag: "",
+ endTag: "",
+ expected: "US-001",
+ },
+ {
+ text: "Some text US-002 more text",
+ startTag: "",
+ endTag: "",
+ expected: "US-002",
+ },
+ {
+ text: " US-003 ",
+ startTag: "",
+ endTag: "",
+ expected: "US-003",
+ },
+ {
+ text: "no tags here",
+ startTag: "",
+ endTag: "",
+ expected: "",
+ },
+ {
+ text: "unclosed",
+ startTag: "",
+ endTag: "",
+ 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 "
+ 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()
+}