diff --git a/.env.example b/.env.example index 92166a36..aa3aa4d8 100644 --- a/.env.example +++ b/.env.example @@ -34,14 +34,14 @@ CODEX_TOKEN= # --- OpenCode Go Configuration --- # Track OpenCode Go subscription quotas. Full walkthrough: docs/OPENCODE_SETUP.md -# Recommended: a console service-account key with usage read access (oc_sk_...). -# onWatch reads the plan's own 5-hour/weekly/monthly meters. No browser cookie needed. +# onWatch reads the plan's own 5-hour/weekly/monthly meters from the console API. +# Option 1: a console service-account key with usage read access (oc_sk_...). OPENCODE_GO_API_KEY= -# Legacy fallback (dashboard scrape) — used only when OPENCODE_GO_API_KEY is empty. -# Both values are required for scrape mode. -# Workspace ID from https://opencode.ai/workspace/wrk_.../go +# Option 2 (used when OPENCODE_GO_API_KEY is empty): your browser session. +# Both values are required. +# Workspace ID (wrk_...), sent as the x-org-id header OPENCODE_GO_WORKSPACE_ID= -# Browser cookie value named "auth" from opencode.ai (value only, no auth= prefix) +# Value of the __Host-console_session cookie from opencode.ai (the old "auth" cookie no longer works) OPENCODE_GO_AUTH_COOKIE= # --- GitHub Copilot Configuration (Beta) --- diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c8a57efc..ba82d9f4 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -1,5 +1,9 @@ name: CI +# Every job runs in parallel on its own runner, so wall time is the slowest +# job, not the sum. "CI OK" at the bottom depends on all of them and is the +# single check to require on main. + on: push: branches: [main] @@ -7,10 +11,16 @@ on: branches: [main] workflow_dispatch: +# A new push to the same PR cancels the older, now-stale run. +concurrency: + group: ci-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: ${{ github.event_name == 'pull_request' }} + jobs: - test: + lint: runs-on: ubuntu-latest - name: Test + name: Lint + timeout-minutes: 15 steps: - uses: actions/checkout@v4 @@ -20,15 +30,69 @@ jobs: with: go-version-file: go.mod - - name: Lint + - name: gofmt + run: | + unformatted=$(gofmt -l .) + if [ -n "$unformatted" ]; then + echo "These files need gofmt:" + echo "$unformatted" + exit 1 + fi + + # Vet every shipped OS and build-tag combination from one runner, so a + # file that only compiles on one platform cannot slip through. + - name: go vet (all platforms) run: | - go fmt ./... - go vet ./... + for os in linux darwin windows; do + echo "== GOOS=$os" + GOOS=$os go vet ./... + done + GOOS=linux go vet -tags menubar ./... + GOOS=windows go vet -tags menubar ./... + + test: + name: Test (${{ matrix.name }}) + runs-on: ${{ matrix.os }} + timeout-minutes: 30 + strategy: + fail-fast: false + matrix: + include: + - name: Linux + os: ubuntu-latest + race: "-race" + tags: menubar + - name: macOS + os: macos-15 + race: "-race" + tags: menubar + # The race detector needs cgo, which the Windows runner lacks. + - name: Windows + os: windows-latest + race: "" + tags: menubar + + steps: + - uses: actions/checkout@v4 - - name: Test with coverage - run: go test -race -coverprofile=coverage.out -covermode=atomic -count=1 ./... + - name: Setup Go + uses: actions/setup-go@v5 + with: + go-version-file: go.mod + + - name: Test + shell: bash + run: go test ${{ matrix.race }} -timeout 15m -coverprofile=coverage.out -covermode=atomic -count=1 ./... + + # The tray companion is compiled only with -tags menubar. + - name: Test tray packages + shell: bash + env: + CGO_LDFLAGS: ${{ runner.os == 'macOS' && '-framework UniformTypeIdentifiers' || '' }} + run: go test ${{ matrix.race }} -timeout 15m -tags ${{ matrix.tags }} -count=1 ./internal/menubar ./internal/web ./cmd/onwatch - name: Upload coverage to Codecov + if: runner.os == 'Linux' uses: codecov/codecov-action@v4 with: files: ./coverage.out @@ -37,12 +101,45 @@ jobs: env: CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} - - name: Build - run: go build -o onwatch ./cmd/onwatch - - tray-linux: + # Alpine is what the shell Docker image runs on: musl libc and busybox + # instead of glibc and coreutils. + test-alpine: runs-on: ubuntu-latest - name: Tray Linux + name: Test (Alpine) + timeout-minutes: 30 + container: golang:1.25-alpine + + steps: + - uses: actions/checkout@v4 + + # Test the musl/busybox userland as a normal host: onWatch switches to + # its Docker defaults (/data, foreground only) when /.dockerenv exists, + # and root bypasses the permission checks some tests rely on. + - name: Prepare non-root host + run: | + rm -f /.dockerenv + adduser -D tester + chown -R tester "$GITHUB_WORKSPACE" + + - name: Test + run: su tester -s /bin/sh -c "export PATH=/usr/local/go/bin:\$PATH; cd '$GITHUB_WORKSPACE' && go test -timeout 15m -count=1 ./..." + + # Build every release artifact exactly as release.yml does, so a PR cannot + # break the release pipeline. + build: + name: Build ${{ matrix.goos }}/${{ matrix.goarch }} + runs-on: ${{ matrix.os }} + timeout-minutes: 20 + strategy: + fail-fast: false + matrix: + include: + - { os: ubuntu-latest, goos: linux, goarch: amd64, cgo: "0", tags: menubar } + - { os: ubuntu-latest, goos: linux, goarch: arm64, cgo: "0", tags: menubar } + - { os: ubuntu-latest, goos: windows, goarch: amd64, cgo: "0", tags: menubar } + - { os: ubuntu-latest, goos: windows, goarch: arm64, cgo: "0", tags: menubar } + - { os: macos-15, goos: darwin, goarch: amd64, cgo: "1", tags: "menubar,desktop,production" } + - { os: macos-15, goos: darwin, goarch: arm64, cgo: "1", tags: "menubar,desktop,production" } steps: - uses: actions/checkout@v4 @@ -52,24 +149,24 @@ jobs: with: go-version-file: go.mod - - name: Test tagged tray packages - # -race needs cgo for the test binary; the tray itself stays pure Go, - # which the CGO_ENABLED=0 cross-compile step below proves. - run: | - go vet -tags menubar ./... - go test -race -tags menubar -count=1 ./internal/menubar ./internal/web ./cmd/onwatch - - - name: Cross-compile tray binaries + - name: Build env: - CGO_ENABLED: "0" - run: | - GOOS=linux GOARCH=amd64 go build -tags menubar -o /tmp/onwatch-linux-amd64 ./cmd/onwatch - GOOS=linux GOARCH=arm64 go build -tags menubar -o /tmp/onwatch-linux-arm64 ./cmd/onwatch - GOOS=windows GOARCH=amd64 go build -tags menubar -o /tmp/onwatch-windows-amd64.exe ./cmd/onwatch + CGO_ENABLED: ${{ matrix.cgo }} + GOOS: ${{ matrix.goos }} + GOARCH: ${{ matrix.goarch }} + run: go build -tags ${{ matrix.tags }} -ldflags="-s -w" -o onwatch-${{ matrix.goos }}-${{ matrix.goarch }} ./cmd/onwatch - tray-windows: - runs-on: windows-latest - name: Tray Windows + e2e: + name: E2E (${{ matrix.name }}) + runs-on: ${{ matrix.os }} + timeout-minutes: 30 + strategy: + fail-fast: false + matrix: + include: + - { name: Linux, os: ubuntu-latest, tags: menubar } + - { name: macOS, os: macos-15, tags: "menubar,desktop,production" } + - { name: Windows, os: windows-latest, tags: menubar } steps: - uses: actions/checkout@v4 @@ -79,64 +176,138 @@ jobs: with: go-version-file: go.mod - - name: Test tagged tray packages + - name: Setup Python + uses: actions/setup-python@v5 + with: + python-version: '3.11' + + - name: Install E2E dependencies + run: | + python -m pip install --upgrade pip + python -m pip install -r tests/e2e/requirements.txt + python -m playwright install --with-deps chromium + + - name: Run E2E suite + shell: bash env: - CGO_ENABLED: "0" + CGO_LDFLAGS: ${{ runner.os == 'macOS' && '-framework UniformTypeIdentifiers' || '' }} + ONWATCH_E2E_GO_BUILD_TAGS: ${{ matrix.tags }} run: | - go test -tags menubar -count=1 ./internal/menubar - go build -tags menubar -o onwatch-tray.exe ./cmd/onwatch + cd tests/e2e + pytest -v --tracing retain-on-failure --output test-results + + - name: Show onWatch logs + if: failure() + shell: bash + run: | + python - <<'EOF' + import glob, os, tempfile + tmp = tempfile.gettempdir() + # Daemon stdout, plus the log file it writes next to its database. + paths = glob.glob(os.path.join(tmp, "onwatch-e2e-*.log")) + glob.glob(os.path.join(tmp, ".onwatch-test.log")) + for path in sorted(paths): + print(f"===== {path}") + with open(path, errors="replace") as f: + print("".join(f.readlines()[-150:])) + EOF + + - name: Upload Playwright traces + if: failure() + uses: actions/upload-artifact@v4 + with: + name: e2e-traces-${{ matrix.name }} + path: tests/e2e/test-results + if-no-files-found: ignore - installer-windows: - runs-on: windows-latest - name: Installer Windows + installer: + name: Installer (${{ matrix.name }}) + runs-on: ${{ matrix.os }} + timeout-minutes: 15 + strategy: + fail-fast: false + matrix: + include: + - { name: Linux, os: ubuntu-latest } + - { name: macOS, os: macos-15 } + - { name: Windows, os: windows-latest } steps: - uses: actions/checkout@v4 + - name: Test install.sh + if: runner.os != 'Windows' + run: bash tests/test_install.sh + # 5.1 is what ships with Windows and what `irm ... | iex` runs, and it is # the host where redirected native stderr under ErrorActionPreference=Stop # becomes a terminating NativeCommandError. 7 keeps the script honest for # anyone who upgraded. - name: Test install.ps1 under Windows PowerShell 5.1 + if: runner.os == 'Windows' shell: powershell run: .\tests\test_install_ps1.ps1 - name: Test install.ps1 under PowerShell 7 + if: runner.os == 'Windows' shell: pwsh run: .\tests\test_install_ps1.ps1 - menubar-macos: - runs-on: macos-15 - name: Menubar macOS + docker: + runs-on: ubuntu-latest + name: Docker + timeout-minutes: 20 steps: - uses: actions/checkout@v4 - - name: Setup Go - uses: actions/setup-go@v5 - with: - go-version-file: go.mod + - name: Build images + run: | + docker build --target runtime-shell -t onwatch:ci-shell . + docker build --target runtime -t onwatch:ci . - - name: Compile tagged menubar packages + # Start each image and wait for the dashboard. Z.ai points at a closed + # local port, so the container makes no outbound provider calls. + - name: Smoke test images run: | - go test -tags menubar ./internal/menubar ./internal/web - CGO_LDFLAGS="-framework UniformTypeIdentifiers" go build -tags menubar,desktop,production -o /tmp/onwatch-menubar ./cmd/onwatch + for image in onwatch:ci-shell onwatch:ci; do + name=smoke-${image//[:.]/-} + docker run -d --name "$name" -p 19300:9211 \ + -e ONWATCH_ADMIN_PASS=ci-smoke -e ZAI_API_KEY=ci -e ZAI_BASE_URL=http://127.0.0.1:1 \ + "$image" + ok="" + for i in $(seq 1 30); do + if curl -fsS -o /dev/null http://localhost:19300/login; then ok=1; break; fi + sleep 1 + done + docker logs "$name" | tail -40 + docker rm -f "$name" + if [ -z "$ok" ]; then echo "$image did not serve /login"; exit 1; fi + done - - name: Setup Python - uses: actions/setup-python@v5 - with: - python-version: '3.11' + # Catches flake.nix drift such as a stale vendorHash after go.sum changes. + nix: + runs-on: ubuntu-latest + name: Nix + timeout-minutes: 30 - - name: Install E2E dependencies - run: | - python -m pip install --upgrade pip - python -m pip install -r tests/e2e/requirements.txt - python -m playwright install chromium + steps: + - uses: actions/checkout@v4 - - name: Run menubar browser tests - env: - CGO_LDFLAGS: -framework UniformTypeIdentifiers - ONWATCH_E2E_GO_BUILD_TAGS: menubar,desktop,production + - uses: DeterminateSystems/nix-installer-action@v16 + + - name: Build + run: nix build .#onwatch --print-build-logs + + ci-ok: + name: CI OK + if: always() + needs: [lint, test, test-alpine, build, e2e, installer, docker, nix] + runs-on: ubuntu-latest + steps: + - name: All jobs passed run: | - cd tests/e2e - pytest tests/test_menubar.py -v + if [ "${{ contains(needs.*.result, 'failure') || contains(needs.*.result, 'cancelled') || contains(needs.*.result, 'skipped') }}" = "true" ]; then + echo "A required job did not pass:" + echo '${{ toJSON(needs) }}' + exit 1 + fi diff --git a/.gitignore b/.gitignore index 71904829..ba46be06 100644 --- a/.gitignore +++ b/.gitignore @@ -95,3 +95,4 @@ onwatch-test # Agent worktrees, local scratch state .claude/worktrees/ +tests/e2e/test-results/ diff --git a/README.md b/README.md index 167d7ad7..d30a72d8 100644 --- a/README.md +++ b/README.md @@ -186,7 +186,7 @@ It is a thin client: it needs the onWatch daemon running (steps above) and finds - **Grok** -- xAI Grok Build / SuperGrok credits tracking via local `~/.grok/auth.json` (or `$GROK_HOME`), optional `grok agent stdio` RPC, and grok.com gRPC-web bearer probe (no browser cookie import). Primary "Credits" utilization against plan limit with reset countdown. Informational local session token stats also captured. - **Moonshot** -- Balance-based tracking for the Moonshot (Kimi) open-platform API. Available, Voucher, and Cash balance cards with drop-rate trends. Set `MOONSHOT_API_KEY`. See [Moonshot Setup](docs/MOONSHOT_SETUP.md). - **DeepSeek** -- Balance-based tracking for the DeepSeek platform API. Total, Granted, and Topped-Up balance cards with drop-rate trends. Set `DEEPSEEK_API_KEY`. See [DeepSeek Setup](docs/DEEPSEEK_SETUP.md). -- **OpenCode Go** -- Subscription quota cards (5-Hour, Weekly, and Monthly when present) read from the plan's own meters via the OpenCode console API (set `OPENCODE_GO_API_KEY`, a service-account key) or, as a legacy fallback, scraped from the authenticated dashboard (`OPENCODE_GO_WORKSPACE_ID` + `OPENCODE_GO_AUTH_COOKIE`), with cycle history and deep insights. Separate from `OPENCODE_ENABLED`, which only feeds ChatGPT credentials into the Codex provider. See [OpenCode Setup](docs/OPENCODE_SETUP.md). +- **OpenCode Go** -- Subscription quota cards (5-Hour, Weekly, and Monthly when present) read from the plan's own meters (used and limit in USD) via the OpenCode console API, using a service-account key (`OPENCODE_GO_API_KEY`) or your browser session (`OPENCODE_GO_WORKSPACE_ID` + `OPENCODE_GO_AUTH_COOKIE`), with cycle history and deep insights. Separate from `OPENCODE_ENABLED`, which only feeds ChatGPT credentials into the Codex provider. See [OpenCode Setup](docs/OPENCODE_SETUP.md). - **Mistral** (beta) - Separate included API and Vibe Code allowances plus pay-as-you-go charges, with browser-cookie import or manual authentication. See [Mistral Setup](docs/MISTRAL_SETUP.md). - **Ollama Cloud** (beta) -- Included monthly usage in USD from the ollama.com API with plan-derived caps, per-model request counts, extra-usage spend, cycle history and insights. Set `OLLAMA_API_KEY`. See [Ollama Setup](docs/OLLAMA_SETUP.md). - **Muse** -- Meta Muse coding-plan quota tracking (5-hour prompts + weekly usage) from the same subscription snapshot `muse /usage` shows, via one minimal probe per poll. Opt-in: set `MUSE_ENABLED=true` and onWatch uses the key `muse login` stored (macOS Keychain / login file), or set `META_API_KEY` directly. Tracking stays off until you opt in, because each poll spends a prompt from your own 5h window. See [Muse Setup](docs/MUSE_SETUP.md). @@ -363,10 +363,13 @@ Additional environment variables: | `KIMI_CODE_ENABLED` | Enable Kimi Code provider (default: auto when credentials present)| | `KIMI_CODE_CREDENTIALS` | Path to kimi-code.json (default ~/.kimi-code/credentials/kimi-code.json)| | `MOONSHOT_API_KEY` | Moonshot (Kimi) open-platform API key (enables balance tracking)| +| `MOONSHOT_BASE_URL` | Override the Moonshot API base URL (default `https://api.moonshot.ai`)| | `DEEPSEEK_API_KEY` | DeepSeek platform API key (enables balance tracking) | +| `DEEPSEEK_BASE_URL` | Override the DeepSeek API base URL (default `https://api.deepseek.com`)| | `OPENCODE_GO_API_KEY` | OpenCode console service-account key with usage read access (enables quota tracking; preferred)| -| `OPENCODE_GO_WORKSPACE_ID` | Legacy scrape mode: OpenCode Go workspace ID (`wrk_...`) from the dashboard URL| -| `OPENCODE_GO_AUTH_COOKIE` | Legacy scrape mode: OpenCode Go `auth` cookie value| +| `OPENCODE_GO_WORKSPACE_ID` | Session mode: OpenCode Go workspace ID (`wrk_...`)| +| `OPENCODE_GO_AUTH_COOKIE` | Session mode: `__Host-console_session` cookie value from opencode.ai| +| `OPENCODE_GO_BASE_URL` | Override the OpenCode console base URL (default `https://opencode.ai`)| | `MISTRAL_ENABLED` | Enable Mistral subscription and pay-as-you-go tracking (default: false) | | `MISTRAL_AUTH_COOKIE` | Manual Mistral Cookie header; keep private | | `MISTRAL_BROWSER` | auto, chrome, firefox, safari (macOS), or edge | diff --git a/cmd/onwatch/codex_profiles.go b/cmd/onwatch/codex_profiles.go index 97019011..8ff1110a 100644 --- a/cmd/onwatch/codex_profiles.go +++ b/cmd/onwatch/codex_profiles.go @@ -817,10 +817,13 @@ func listCodexProfiles() ([]CodexProfile, error) { } entries, err := os.ReadDir(profilesDir) - if os.IsNotExist(err) { - return nil, nil - } if err != nil { + // Only a missing directory means "no profiles yet". Windows reports a + // file in the directory's place as "path not found" too, which must + // not be mistaken for an empty profile list. + if _, statErr := os.Stat(profilesDir); os.IsNotExist(statErr) { + return nil, nil + } return nil, fmt.Errorf("failed to read profiles directory: %w", err) } diff --git a/cmd/onwatch/codex_profiles_test.go b/cmd/onwatch/codex_profiles_test.go index f9581a19..8b84103a 100644 --- a/cmd/onwatch/codex_profiles_test.go +++ b/cmd/onwatch/codex_profiles_test.go @@ -10,6 +10,7 @@ import ( "time" "github.com/onllm-dev/onwatch/v2/internal/api" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func writeRefreshAuthJSON(t *testing.T, home, access, refresh, idToken, account string) { @@ -113,7 +114,7 @@ func loadProfileForTest(t *testing.T, home, name string) *CodexProfile { func TestRefreshCodexProfile_SameAccount(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeRefreshAuthJSON(t, home, "new_access", "new_refresh", "new_id", "acc_same") @@ -145,14 +146,12 @@ func TestRefreshCodexProfile_SameAccount(t *testing.T) { if err != nil { t.Fatalf("stat profile: %v", err) } - if info.Mode().Perm() != 0o600 { - t.Fatalf("profile permissions = %o, want 600", info.Mode().Perm()) - } + assertPerm(t, info, 0o600) } func TestRefreshCodexProfile_DifferentAccount_UserConfirms(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeRefreshAuthJSON(t, home, "new_access", "new_refresh", "new_id", "acc_new") @@ -176,7 +175,7 @@ func TestRefreshCodexProfile_DifferentAccount_UserConfirms(t *testing.T) { func TestRefreshCodexProfile_DifferentAccount_UserDeclines(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeRefreshAuthJSON(t, home, "new_access", "new_refresh", "new_id", "acc_new") @@ -205,7 +204,7 @@ func TestRefreshCodexProfile_DifferentAccount_UserDeclines(t *testing.T) { func TestRefreshCodexProfile_NewProfile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeRefreshAuthJSON(t, home, "new_access", "new_refresh", "new_id", "acc_new") @@ -228,7 +227,7 @@ func TestRefreshCodexProfile_NewProfile(t *testing.T) { func TestRefreshCodexProfile_NoAuthJSON(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") err := codexProfileRefresh("work", "") @@ -242,7 +241,7 @@ func TestRefreshCodexProfile_NoAuthJSON(t *testing.T) { func TestCodexProfilesDir(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) if got := codexProfilesDir(); got != filepath.Join(home, ".onwatch", "data", "codex-profiles") { t.Fatalf("codexProfilesDir() = %q", got) @@ -271,7 +270,7 @@ func TestPrintCodexHelp(t *testing.T) { func TestRunCodexCommand(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") origArgs := os.Args @@ -352,7 +351,7 @@ func TestRunCodexCommand(t *testing.T) { func TestCodexProfileSaveListStatusDeleteFlow(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeRefreshAuthJSON(t, home, "save_access", "save_refresh", "save_id", "acct_one") @@ -408,7 +407,7 @@ func TestCodexProfileSaveListStatusDeleteFlow(t *testing.T) { func TestCodexProfileSave_BlocksDuplicateAccount(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") // Both profile and new auth have the same account AND same user_id -> duplicate. @@ -429,7 +428,7 @@ func TestCodexProfileSave_BlocksDuplicateAccount(t *testing.T) { func TestCodexProfileSave_InvalidNameAndMissingCredentials(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") if err := codexProfileSave("bad name", ""); err == nil || !strings.Contains(err.Error(), "invalid profile name") { @@ -443,7 +442,7 @@ func TestCodexProfileSave_InvalidNameAndMissingCredentials(t *testing.T) { func TestListCodexProfiles_SkipsInvalidFilesAndDerivesName(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") profilesDir := filepath.Join(home, ".onwatch", "data", "codex-profiles") @@ -472,7 +471,7 @@ func TestListCodexProfiles_SkipsInvalidFilesAndDerivesName(t *testing.T) { func TestCodexProfileStatus_NoCredentials(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") profilesDir := filepath.Join(home, ".onwatch", "data", "codex-profiles") @@ -501,7 +500,7 @@ func TestCodexAuthRefreshPath_UsesCODEXHOMEAndDeleteMissingProfile(t *testing.T) } home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") if err := codexProfileDelete("missing"); err == nil || !strings.Contains(err.Error(), `profile "missing" not found`) { t.Fatalf("codexProfileDelete(missing) = %v", err) @@ -511,7 +510,7 @@ func TestCodexAuthRefreshPath_UsesCODEXHOMEAndDeleteMissingProfile(t *testing.T) func TestLoadCodexAuthForRefresh_FlatShapeAndErrors(t *testing.T) { t.Run("supports flat auth.json shape", func(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") codexDir := filepath.Join(home, ".codex") @@ -597,7 +596,7 @@ func TestRunCodexCommand_AdditionalHelpPaths(t *testing.T) { func TestListCodexProfiles_ReadDirError(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") dataDir := filepath.Join(home, ".onwatch", "data") @@ -617,7 +616,7 @@ func TestListCodexProfiles_ReadDirError(t *testing.T) { func TestCodexProfileSave_WarnsOnSameProfileAccountChange(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeRefreshAuthJSON(t, home, "new_access", "new_refresh", "new_id", "acct_new") @@ -642,7 +641,7 @@ func TestCodexProfileRefresh_InvalidName(t *testing.T) { func TestCodexProfileSave_AllowsSameAccountDifferentUser(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeProfileFileWithUser(t, home, "personal", "old_access", "old_refresh", "acct_team", "user-one") @@ -663,7 +662,7 @@ func TestCodexProfileSave_AllowsSameAccountDifferentUser(t *testing.T) { func TestCodexProfileSave_StoresUserIDFromIDToken(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeRefreshAuthJSONWithUser(t, home, "save_access", "save_refresh", "acct_one", "user-one") @@ -680,7 +679,7 @@ func TestCodexProfileSave_StoresUserIDFromIDToken(t *testing.T) { func TestCodexProfileRefresh_UpdatesUserIDFromIDToken(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") writeProfileFileWithUser(t, home, "work", "old_access", "old_refresh", "acct_team", "user-one") @@ -717,7 +716,7 @@ func TestIsDuplicateCodexProfile_Direct(t *testing.T) { func TestCodexProfileSave_AllowsSameAccountNoUserIDRegression(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") // Existing profile: same account, no user_id in JWT (legacy tokens) @@ -748,7 +747,7 @@ func TestCodexProfileSave_AllowsSameAccountNoUserIDRegression(t *testing.T) { // account has no user_id. This is the Team upgrade scenario. func TestCodexProfileSave_AllowsNewUserAlongsideLegacyProfile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") // Legacy profile: same account, no user_id in JWT @@ -962,7 +961,7 @@ func TestLoadCodexAuthFromFile(t *testing.T) { func TestCodexProfileSaveWithAuthFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") t.Setenv("CODEX_TOKEN", "") @@ -1002,7 +1001,7 @@ func TestCodexProfileSaveWithAuthFile(t *testing.T) { func TestCodexProfileRefreshWithAuthFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") t.Setenv("CODEX_TOKEN", "") diff --git a/cmd/onwatch/first_run_test.go b/cmd/onwatch/first_run_test.go index 32f6832a..6225ffc3 100644 --- a/cmd/onwatch/first_run_test.go +++ b/cmd/onwatch/first_run_test.go @@ -2,8 +2,6 @@ package main import ( "bufio" - "os/exec" - "runtime" "strings" "testing" "time" @@ -126,9 +124,6 @@ func indexOf(s, sub string) int { } func TestStopProcessAndProcessAlive(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("uses sleep(1)") - } if processAlive(0) { t.Fatal("processAlive(0) must be false") } @@ -136,10 +131,9 @@ func TestStopProcessAndProcessAlive(t *testing.T) { t.Fatal("stopProcess(0) must be false") } - cmd := exec.Command("sleep", "30") - if err := cmd.Start(); err != nil { - t.Skipf("cannot start sleep: %v", err) - } + // Re-exec the test binary as the long-running child rather than sleep(1), + // which does not exist on Windows. + cmd := startSleepSubprocess(t) pid := cmd.Process.Pid t.Cleanup(func() { _ = cmd.Process.Kill(); _, _ = cmd.Process.Wait() }) diff --git a/cmd/onwatch/main.go b/cmd/onwatch/main.go index e9ba5f69..94db7543 100644 --- a/cmd/onwatch/main.go +++ b/cmd/onwatch/main.go @@ -255,13 +255,14 @@ func findOnwatchOnPort(port int) []int { return pids } -// isOnwatchProcess checks if a PID belongs to an onwatch (or legacy syntrack) binary. +// isOnwatchProcess checks if a PID belongs to an onwatch (or legacy syntrack) +// binary. processCommandName is per-platform: ps on Unix, the process image +// path on Windows, which has no ps. func isOnwatchProcess(pid int) bool { - out, err := exec.Command("ps", "-p", strconv.Itoa(pid), "-o", "comm=").Output() - if err != nil { + if pid <= 0 { return false } - cmd := strings.ToLower(strings.TrimSpace(string(out))) + cmd := strings.ToLower(processCommandName(pid)) return strings.Contains(cmd, "onwatch") || strings.Contains(cmd, "syntrack") } @@ -390,8 +391,8 @@ func migrateDBLocation(newPath string, logger *slog.Logger) { oldPaths := []string{ "./onwatch.db", } - oldHome := os.Getenv("HOME") - if oldHome != "" { + // os.UserHomeDir, not $HOME: Windows keeps the profile in USERPROFILE. + if oldHome, err := os.UserHomeDir(); err == nil && oldHome != "" { oldPaths = append(oldPaths, filepath.Join(oldHome, ".onwatch", "onwatch.db"), ) @@ -523,8 +524,12 @@ func testDaemonIsolationEnv(exe string) []string { return nil } _ = os.MkdirAll(filepath.Join(dir, ".onwatch", "data"), 0o755) + // HOME is the home directory on Unix, USERPROFILE on Windows, and + // LOCALAPPDATA holds the Windows PID directory. env := []string{ "HOME=" + dir, + "USERPROFILE=" + dir, + "LOCALAPPDATA=" + dir, "ONWATCH_DB_PATH=" + filepath.Join(dir, "onwatch.db"), } if port, err := freeLocalPort(); err == nil { @@ -1107,14 +1112,22 @@ func run() error { var moonshotClient *api.MoonshotClient if cfg.HasProvider("moonshot") { - moonshotClient = api.NewMoonshotClient(cfg.MoonshotAPIKey, logger) - logger.Info("Moonshot API client configured") + var moonshotOpts []api.MoonshotOption + if cfg.MoonshotBaseURL != "" { + moonshotOpts = append(moonshotOpts, api.WithMoonshotBaseURL(strings.TrimRight(cfg.MoonshotBaseURL, "/"))) + } + moonshotClient = api.NewMoonshotClient(cfg.MoonshotAPIKey, logger, moonshotOpts...) + logger.Info("Moonshot API client configured", "base_url_override", cfg.MoonshotBaseURL != "") } var deepseekClient *api.DeepSeekClient if cfg.HasProvider("deepseek") { - deepseekClient = api.NewDeepSeekClient(cfg.DeepSeekAPIKey, logger) - logger.Info("DeepSeek API client configured") + var deepseekOpts []api.DeepSeekOption + if cfg.DeepSeekBaseURL != "" { + deepseekOpts = append(deepseekOpts, api.WithDeepSeekBaseURL(strings.TrimRight(cfg.DeepSeekBaseURL, "/"))) + } + deepseekClient = api.NewDeepSeekClient(cfg.DeepSeekAPIKey, logger, deepseekOpts...) + logger.Info("DeepSeek API client configured", "base_url_override", cfg.DeepSeekBaseURL != "") } // Gemini provider - env vars or auto-detect from ~/.gemini/oauth_creds.json @@ -1498,7 +1511,11 @@ func run() error { } var opencodeAg *agent.OpenCodeAgent if cfg.HasProvider("opencode") { - opencodeClient := api.NewOpenCodeClient(logger) + var opencodeOpts []api.OpenCodeClientOption + if cfg.OpenCodeGoBaseURL != "" { + opencodeOpts = append(opencodeOpts, api.WithOpenCodeBaseURL(cfg.OpenCodeGoBaseURL)) + } + opencodeClient := api.NewOpenCodeClient(logger, opencodeOpts...) opencodeSm := agent.NewSessionManager(db, "opencode", idleTimeout, logger) opencodeAg = agent.NewOpenCodeAgent(opencodeClient, db, opencodeTr, cfg, cfg.PollInterval, logger, opencodeSm) } diff --git a/cmd/onwatch/main_test.go b/cmd/onwatch/main_test.go index 08089ec4..b92ce252 100644 --- a/cmd/onwatch/main_test.go +++ b/cmd/onwatch/main_test.go @@ -10,6 +10,7 @@ import ( "testing" "github.com/onllm-dev/onwatch/v2/internal/config" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" "github.com/onllm-dev/onwatch/v2/internal/web" ) @@ -17,7 +18,7 @@ func TestConfigLoad_WithOnlyCodexAuthFile_AllowsEmptyProviderConfig(t *testing.T homeDir := t.TempDir() codexHome := t.TempDir() t.Chdir(t.TempDir()) - t.Setenv("HOME", homeDir) + testhome.SetTestHome(t, homeDir) t.Setenv("CODEX_HOME", codexHome) t.Setenv("SYNTHETIC_API_KEY", "") t.Setenv("ZAI_API_KEY", "") @@ -117,7 +118,7 @@ func TestDeriveEncryptionKey_UsesEncryptionSalt(t *testing.T) { func TestStatusLogCandidates(t *testing.T) { t.Run("prefers db directory then home then cwd", func(t *testing.T) { homeDir := t.TempDir() - t.Setenv("HOME", homeDir) + testhome.SetTestHome(t, homeDir) dbPath := filepath.Join(t.TempDir(), "data", "onwatch.db") got := statusLogCandidates(dbPath, "main.log", "menubar.log") @@ -142,7 +143,7 @@ func TestStatusLogCandidates(t *testing.T) { t.Run("adds pid dir when db path missing", func(t *testing.T) { homeDir := t.TempDir() - t.Setenv("HOME", homeDir) + testhome.SetTestHome(t, homeDir) oldPIDDir := pidDir pidDir = t.TempDir() @@ -166,7 +167,7 @@ func TestStatusLogCandidates(t *testing.T) { t.Run("deduplicates repeated names", func(t *testing.T) { homeDir := t.TempDir() - t.Setenv("HOME", homeDir) + testhome.SetTestHome(t, homeDir) dbPath := filepath.Join(t.TempDir(), "data", "onwatch.db") got := statusLogCandidates(dbPath, "main.log", "main.log") @@ -345,9 +346,7 @@ func TestWriteEnvFile(t *testing.T) { if err != nil { t.Fatalf("stat env file: %v", err) } - if stat.Mode().Perm() != 0o600 { - t.Fatalf("expected mode 0600, got %o", stat.Mode().Perm()) - } + assertPerm(t, stat, 0o600) } func TestMaskValue(t *testing.T) { diff --git a/cmd/onwatch/main_testmain_test.go b/cmd/onwatch/main_testmain_test.go index dd87b0c0..447f2558 100644 --- a/cmd/onwatch/main_testmain_test.go +++ b/cmd/onwatch/main_testmain_test.go @@ -1,27 +1,33 @@ package main import ( - "log/slog" - + "errors" "fmt" - "github.com/onllm-dev/onwatch/v2/internal/api" + "log/slog" "net" "os" "path/filepath" + "runtime" "testing" + "github.com/onllm-dev/onwatch/v2/internal/api" "github.com/onllm-dev/onwatch/v2/internal/service" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" + "github.com/onllm-dev/onwatch/v2/internal/update" ) -// TestMain runs before all tests in the main package. It unsets -// OPENCODE_HOME/XDG_DATA_HOME so the interactive setup flow's codex credential -// auto-detection never resolves to the host's real -// ~/.local/share/opencode/auth.json. Setup tests set a temp HOME and drive the -// prompts with fixed input; reading a real OpenCode ChatGPT login would shift -// those input sequences and make the tests environment-dependent (flaky). +// TestMain runs before all tests in the main package. It enables api test +// mode and sandboxes the home directory, clearing provider location overrides +// such as OPENCODE_HOME/XDG_DATA_HOME, so the interactive setup flow's +// credential auto-detection never resolves to the host's real files (e.g. +// ~/.local/share/opencode/auth.json). Setup tests set a temp HOME and drive +// the prompts with fixed input; reading a real login would shift those input +// sequences and make the tests environment-dependent (flaky). func TestMain(m *testing.M) { - os.Unsetenv("OPENCODE_HOME") - os.Unsetenv("XDG_DATA_HOME") + // SetTestMode must run before sandboxTestHome redirects HOME/USERPROFILE: + // its first enable records the real home, which the credential-file guard + // then refuses, and it keeps every keychain/keyring operation off. + api.SetTestMode(true) // GitHub Actions runs jobs under systemd, so INVOCATION_ID is set on the // runner and update.IsSystemd() reports true there but not on a developer @@ -53,7 +59,84 @@ func TestMain(m *testing.M) { // scripted prompt input in the wizard tests desyncs from there on. detectMuseCredentialsFunc = func(*slog.Logger) *api.MuseCredentials { return nil } - os.Exit(m.Run()) + // `onwatch update` tests must never reach GitHub: with a real updater a + // test that sets an old version downloads the latest release and replaces + // the running test binary, after which every os.Args[0] helper spawn + // launches a real onWatch. Tests that need other answers stub their own. + newCLIUpdater = func(v string, _ *slog.Logger) cliUpdater { return offlineCLIUpdater{version: v} } + + // Setup tests must never star the repo through a developer's logged-in + // gh CLI; the star tests opt back in with t.Setenv. + os.Setenv("ONWATCH_STAR", "no") + + cleanupHome := sandboxTestHome() + + code := m.Run() + cleanupHome() + os.Exit(code) +} + +// offlineCLIUpdater is the network-free default updater for tests. A dev +// build is always current, matching the real updater; any other version +// reports a failed check, which the update tests accept as the offline result. +type offlineCLIUpdater struct{ version string } + +func (o offlineCLIUpdater) Check() (update.UpdateInfo, error) { + if o.version == "dev" || o.version == "" { + return update.UpdateInfo{CurrentVersion: o.version, LatestVersion: o.version}, nil + } + return update.UpdateInfo{}, errors.New("network access is disabled in tests") +} + +func (o offlineCLIUpdater) Apply() error { + return errors.New("network access is disabled in tests") +} + +// testScratchHomeEnv marks a process tree whose home is already sandboxed. +const testScratchHomeEnv = "_ONWATCH_TEST_SCRATCH_HOME" + +// sandboxTestHome is the home-directory safety net: it points the whole test +// process at a scratch home (testhome.SandboxHome: HOME, USERPROFILE and +// LOCALAPPDATA, with provider location overrides cleared) so a test that +// forgets to set its own never reads or writes the developer's (or CI +// runner's) real ~/.onwatch, ~/.codex and so on. +// +// Helper subprocesses re-run TestMain. They inherit the marker and keep the +// home and environment their parent test gave them (itself a sandbox), so a +// test that seeds a fixture home for a child still has it seen. Returns the +// cleanup for the directory this process created (a no-op when it inherited +// one). +func sandboxTestHome() func() { + if os.Getenv(testScratchHomeEnv) != "" { + return func() {} + } + dir, cleanup, err := testhome.SandboxHome() + if err != nil { + fmt.Fprintf(os.Stderr, "cannot create scratch home: %v\n", err) + os.Exit(1) + } + os.Setenv(testScratchHomeEnv, dir) + // pidDir was resolved at package init from the real home. Re-resolve it + // so runStop/runStatus never read (or stop) a real daemon or menubar + // companion. initialPIDFilePath keeps the real path for the off-limits + // guard. + pidDir = defaultPIDDir() + pidFile = filepath.Join(pidDir, "onwatch.pid") + return cleanup +} + +// assertPerm checks a file's Unix permission bits. Windows has no such bits: +// os.Chmod only toggles the read-only attribute and Stat reports a writable +// file as 0666, so access there is governed by the profile directory's ACL +// and the check is Unix-only. +func assertPerm(t testing.TB, info os.FileInfo, want os.FileMode) { + t.Helper() + if runtime.GOOS == "windows" { + return + } + if got := info.Mode().Perm(); got != want { + t.Fatalf("%s permissions = %o, want %o", info.Name(), got, want) + } } // This runs during package initialisation rather than from TestMain. A @@ -95,6 +178,8 @@ func isolateSpawnedDaemonChild() { } _ = os.MkdirAll(filepath.Join(dir, ".onwatch", "data"), 0o755) os.Setenv("HOME", dir) + os.Setenv("USERPROFILE", dir) + os.Setenv("LOCALAPPDATA", dir) os.Setenv("ONWATCH_DB_PATH", filepath.Join(dir, "onwatch.db")) if port, err := freePort(); err == nil { os.Setenv("ONWATCH_PORT", fmt.Sprintf("%d", port)) @@ -111,3 +196,11 @@ func freePort() (int, error) { defer ln.Close() return ln.Addr().(*net.TCPAddr).Port, nil } + +// TestMain must enable api test mode so no cmd test can reach the keychain, +// keyring or the real Claude credentials file. +func TestTestMainEnablesAPITestMode(t *testing.T) { + if !api.IsTestMode() { + t.Fatal("TestMain must call api.SetTestMode(true)") + } +} diff --git a/cmd/onwatch/platform_unix.go b/cmd/onwatch/platform_unix.go index 74a2a3f0..53990910 100644 --- a/cmd/onwatch/platform_unix.go +++ b/cmd/onwatch/platform_unix.go @@ -7,6 +7,8 @@ import ( "os" "os/exec" "path/filepath" + "runtime" + "strconv" "strings" "syscall" ) @@ -43,13 +45,62 @@ func processAlive(pid int) bool { return proc.Signal(syscall.Signal(0)) == nil } +// processCommandName returns the executable base name of pid ("" when +// unknown). macOS ps prints the full path for comm, and only the base name may +// count: otherwise any binary under a directory named onwatch would pass +// isOnwatchProcess. This matches the Windows variant. +func processCommandName(pid int) string { + if name, ok := procExeName(pid); ok { + return name + } + out, err := exec.Command("ps", "-p", strconv.Itoa(pid), "-o", "comm=").Output() + if err != nil { + return "" + } + name := strings.TrimSpace(string(out)) + if name == "" { + return "" + } + return filepath.Base(name) +} + func processZombie(pid int) bool { if pid <= 0 { return false } + if runtime.GOOS == "linux" { + // /proc//stat: "pid (comm) state ..."; comm may contain spaces + // or parentheses, so the state follows the last ')'. + if data, err := os.ReadFile(fmt.Sprintf("/proc/%d/stat", pid)); err == nil { + if i := strings.LastIndexByte(string(data), ')'); i >= 0 && i+2 < len(data) { + return data[i+2] == 'Z' + } + } + return false + } out, err := exec.Command("ps", "-p", fmt.Sprintf("%d", pid), "-o", "stat=").Output() if err != nil { return false } return strings.Contains(strings.TrimSpace(string(out)), "Z") } + +// procExeName reads the process name from /proc on Linux, where ps may be +// missing (Nix build sandbox, distroless image) or busybox's, which has no +// -p. ok is false when /proc has no entry for pid. +func procExeName(pid int) (name string, ok bool) { + if runtime.GOOS != "linux" { + return "", false + } + dir := "/proc/" + strconv.Itoa(pid) + // comm is the name the process was started as (what ps -o comm= shows), + // so a binary launched through a symlink named onwatch still matches. + // The kernel truncates it to 15 characters, which "onwatch" fits. + if comm, err := os.ReadFile(dir + "/comm"); err == nil { + return strings.TrimSpace(string(comm)), true + } + if exe, err := os.Readlink(dir + "/exe"); err == nil { + return filepath.Base(strings.TrimSuffix(exe, " (deleted)")), true + } + return "", false +} diff --git a/cmd/onwatch/platform_unix_test.go b/cmd/onwatch/platform_unix_test.go new file mode 100644 index 00000000..d8a20be0 --- /dev/null +++ b/cmd/onwatch/platform_unix_test.go @@ -0,0 +1,58 @@ +//go:build !windows + +package main + +import ( + "os" + "os/exec" + "path/filepath" + "testing" +) + +func TestDaemonSysProcAttr_UnixSetsid(t *testing.T) { + attr := daemonSysProcAttr() + if attr == nil { + t.Fatal("expected non-nil SysProcAttr") + } + if !attr.Setsid { + t.Fatal("expected Setsid=true") + } +} + +// macOS ps prints the full executable path for comm, so a binary that merely +// lives under a directory called onwatch must not pass for onWatch. Only the +// base name counts, as on Windows. +func TestProcessCommandName_BaseNameOnly(t *testing.T) { + // A copy of this test binary, renamed and placed under an onwatch + // directory, runs the idle sleep helper. + self, err := os.Executable() + if err != nil { + t.Fatalf("locate test binary: %v", err) + } + data, err := os.ReadFile(self) + if err != nil { + t.Fatalf("read test binary: %v", err) + } + dir := filepath.Join(t.TempDir(), "onwatch", "bin") + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatal(err) + } + bin := filepath.Join(dir, "idler") + if err := os.WriteFile(bin, data, 0o755); err != nil { + t.Fatal(err) + } + + cmd := exec.Command(bin, "-test.run=^TestSleepHelperProcess_NeverRun$") + cmd.Env = append(os.Environ(), "GO_SLEEP_HELPER=1") + if err := cmd.Start(); err != nil { + t.Fatalf("start helper: %v", err) + } + t.Cleanup(func() { _ = cmd.Process.Kill(); _ = cmd.Wait() }) + + if got := processCommandName(cmd.Process.Pid); got != "idler" { + t.Errorf("processCommandName() = %q, want %q", got, "idler") + } + if isOnwatchProcess(cmd.Process.Pid) { + t.Error("a binary under an onwatch directory must not be treated as onWatch") + } +} diff --git a/cmd/onwatch/platform_windows.go b/cmd/onwatch/platform_windows.go index 5cdb4571..bfc9e7c4 100644 --- a/cmd/onwatch/platform_windows.go +++ b/cmd/onwatch/platform_windows.go @@ -6,6 +6,7 @@ import ( "os" "path/filepath" "syscall" + "unsafe" ) const createNoWindow = 0x08000000 @@ -14,6 +15,12 @@ const createNoWindow = 0x08000000 // process is still running. const waitTimeout = uint32(0x00000102) +// processQueryLimitedInformation is PROCESS_QUERY_LIMITED_INFORMATION, the +// least access right that allows reading a process's image path. +const processQueryLimitedInformation = 0x1000 + +var procQueryFullProcessImageNameW = syscall.NewLazyDLL("kernel32.dll").NewProc("QueryFullProcessImageNameW") + func daemonSysProcAttr() *syscall.SysProcAttr { return &syscall.SysProcAttr{ HideWindow: true, @@ -64,3 +71,30 @@ func processAlive(pid int) bool { } return state == waitTimeout } + +// processCommandName returns the image file name of pid, e.g. onwatch.exe +// ("" when unknown). Windows has no ps, so ask the process object itself. Only +// the base name counts: a full path such as C:\Users\x\.onwatch\bin\... would +// let any binary under an "onwatch" directory pass for onWatch. +func processCommandName(pid int) string { + if pid <= 0 { + return "" + } + handle, err := syscall.OpenProcess(processQueryLimitedInformation, false, uint32(pid)) + if err != nil { + return "" + } + defer syscall.CloseHandle(handle) + buf := make([]uint16, 1024) + size := uint32(len(buf)) + r, _, _ := procQueryFullProcessImageNameW.Call( + uintptr(handle), + 0, + uintptr(unsafe.Pointer(&buf[0])), + uintptr(unsafe.Pointer(&size)), + ) + if r == 0 { + return "" + } + return filepath.Base(syscall.UTF16ToString(buf[:size])) +} diff --git a/cmd/onwatch/root_coverage3_test.go b/cmd/onwatch/root_coverage3_test.go index d587eff2..1cbe643f 100644 --- a/cmd/onwatch/root_coverage3_test.go +++ b/cmd/onwatch/root_coverage3_test.go @@ -15,6 +15,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/config" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) // --------------------------------------------------------------------------- @@ -134,7 +135,7 @@ func TestRun_DaemonChildStartupError(t *testing.T) { t.Setenv("ANTIGRAVITY_BASE_URL", "") t.Setenv("ANTIGRAVITY_CSRF_TOKEN", "") home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // The data directory is created on demand now, so block it with a regular // file to keep exercising the "logging setup fails" path. if err := os.MkdirAll(filepath.Join(home, ".onwatch"), 0o755); err != nil { @@ -220,7 +221,7 @@ func TestFreshSetup_ZaiOnly(t *testing.T) { func TestFreshSetup_AllProviders(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) // Disable keychain tools so no auto-detect occurs for anthropic/codex t.Setenv("PATH", "") @@ -272,7 +273,7 @@ func TestFreshSetup_AllProviders(t *testing.T) { func TestFreshSetup_MultipleProviders_Choice6(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) // Disable keychain tools so no auto-detect occurs t.Setenv("PATH", "") @@ -316,7 +317,7 @@ func TestFreshSetup_MultipleProviders_Choice6(t *testing.T) { func TestFreshSetup_AnthropicOnly(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Disable keychain tools and PATH so no auto-detect occurs t.Setenv("PATH", "") @@ -342,7 +343,7 @@ func TestFreshSetup_AnthropicOnly(t *testing.T) { func TestFreshSetup_CodexOnly(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) // No codex auth file -> fallback to manual entry @@ -371,7 +372,7 @@ func TestFreshSetup_CodexOnly(t *testing.T) { func TestAddMissingProviders_AllSkipped(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) // Disable keychain tools so no auto-detect occurs t.Setenv("PATH", "") @@ -418,7 +419,7 @@ func TestAddMissingProviders_AllSkipped(t *testing.T) { func TestAddMissingProviders_ZaiSkippedAnthropicAdded(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) // Disable keychain tools so no auto-detect occurs t.Setenv("PATH", "") @@ -458,12 +459,9 @@ func TestAddMissingProviders_ZaiSkippedAnthropicAdded(t *testing.T) { // --------------------------------------------------------------------------- func TestCollectAnthropicToken_AutoDetect_Accept(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("unix credential file path test") - } home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Disable keychain lookup tools so file fallback is used t.Setenv("PATH", "") @@ -487,12 +485,9 @@ func TestCollectAnthropicToken_AutoDetect_Accept(t *testing.T) { } func TestCollectAnthropicToken_AutoDetect_Decline(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("unix credential file path test") - } home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Disable keychain lookup tools so file fallback is used t.Setenv("PATH", "") @@ -787,12 +782,9 @@ func TestRunStatus_SelfPIDRunning(t *testing.T) { // --------------------------------------------------------------------------- func TestAddMissingProviders_AnthropicAutoDetected(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("unix credential file path test") - } home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) // Disable keychain lookup tools so file fallback is used t.Setenv("PATH", "") @@ -842,7 +834,7 @@ func TestAddMissingProviders_AnthropicAutoDetected(t *testing.T) { func TestAddMissingProviders_FileOpenError(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) t.Setenv("PATH", "") @@ -866,7 +858,7 @@ func TestAddMissingProviders_FileOpenError(t *testing.T) { func TestAddMissingProviders_ZaiAdded(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) t.Setenv("PATH", "") @@ -907,7 +899,7 @@ func TestAddMissingProviders_ZaiAdded(t *testing.T) { func TestAddMissingProviders_AntigravityAdded(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) t.Setenv("PATH", "") @@ -942,13 +934,77 @@ func TestAddMissingProviders_AntigravityAdded(t *testing.T) { } } +// --------------------------------------------------------------------------- +// addMissingProviders() - Gemini auto-detected from the home directory +// --------------------------------------------------------------------------- + +// The Gemini CLI credential lookup must resolve the home directory the same +// way everywhere: on Windows os.UserHomeDir reads USERPROFILE and $HOME is +// normally unset, so reading $HOME directly never detected Gemini there. +func TestAddMissingProviders_GeminiAutoDetectedFromHome(t *testing.T) { + for _, tc := range []struct { + name string + withCreds bool + }{ + {"credentials present", true}, + {"credentials absent", false}, + } { + t.Run(tc.name, func(t *testing.T) { + home := t.TempDir() + testhome.SetTestHome(t, home) + if tc.withCreds { + credPath := filepath.Join(home, ".gemini", "oauth_creds.json") + if err := os.MkdirAll(filepath.Dir(credPath), 0o700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(credPath, []byte(`{}`), 0o600); err != nil { + t.Fatal(err) + } + } + + envFile := filepath.Join(t.TempDir(), ".env") + if err := os.WriteFile(envFile, []byte("SYNTHETIC_API_KEY=syn_existing\n"), 0o600); err != nil { + t.Fatalf("write env: %v", err) + } + + // Every other provider is already configured, so the Gemini prompt + // is the first one asked. An empty answer takes the default: yes + // when credentials were detected, no otherwise. + existing := &existingEnv{ + syntheticKey: "syn_existing", + zaiKey: "zai", + anthropicToken: "anthropic", + codexToken: "codex", + openCodeEnabled: true, + antigravityEnabled: true, + grokEnabled: true, + ollamaKey: "ollama", + } + reader := bufio.NewReader(strings.NewReader("\n")) + captureStdout(t, func() { + if err := addMissingProviders(reader, envFile, existing); err != nil { + t.Fatalf("addMissingProviders error: %v", err) + } + }) + + data, err := os.ReadFile(envFile) + if err != nil { + t.Fatalf("read env: %v", err) + } + if got := strings.Contains(string(data), "GEMINI_ENABLED=true"); got != tc.withCreds { + t.Fatalf("GEMINI_ENABLED written = %v, want %v:\n%s", got, tc.withCreds, data) + } + }) + } +} + // --------------------------------------------------------------------------- // addMissingProviders() - codex manual path (no auto-detect) // --------------------------------------------------------------------------- func TestAddMissingProviders_CodexManualPath(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) t.Setenv("PATH", "") @@ -987,12 +1043,9 @@ func TestAddMissingProviders_CodexManualPath(t *testing.T) { // --------------------------------------------------------------------------- func TestAddMissingProviders_AnthropicAutoDetectDeclined(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("unix credential file path test") - } home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) t.Setenv("PATH", "") @@ -1042,7 +1095,7 @@ func TestAddMissingProviders_AnthropicAutoDetectDeclined(t *testing.T) { func TestAddMissingProviders_CodexAutoDetectDeclined(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) codexHome := t.TempDir() t.Setenv("CODEX_HOME", codexHome) @@ -1088,7 +1141,7 @@ func TestAddMissingProviders_CodexAutoDetectDeclined(t *testing.T) { func TestAddMissingProviders_SyntheticAdded(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) t.Setenv("PATH", "") @@ -1128,7 +1181,7 @@ func TestAddMissingProviders_SyntheticAdded(t *testing.T) { func TestAddMissingProviders_CodexAutoDetected(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) codexHome := t.TempDir() t.Setenv("CODEX_HOME", codexHome) @@ -1181,9 +1234,6 @@ func TestStopPreviousInstance_NonTestModeNoPIDFile(t *testing.T) { } func TestStopPreviousInstance_WithPIDFilePortAndListener(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only available on macOS/Linux") - } oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") @@ -1232,7 +1282,7 @@ func TestStopPreviousInstance_WithSelfPIDFile(t *testing.T) { func TestMigrateDBLocation_NewAlreadyExists(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Create both old and new DB oldDB := filepath.Join(home, ".onwatch", "onwatch.db") @@ -1261,7 +1311,7 @@ func TestMigrateDBLocation_NewAlreadyExists(t *testing.T) { func TestMigrateDBLocation_OldPathEqualsNew(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // newPath == one of the oldPaths -> should skip (continue branch) newDB := filepath.Join(home, ".onwatch", "onwatch.db") @@ -1286,7 +1336,7 @@ func TestMigrateDBLocation_OldPathEqualsNew(t *testing.T) { func TestFreshSetup_NoProviderSelected_ReturnsError(t *testing.T) { // Provide choice 7 (Multiple), answer "n" to everything. home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) // Disable keychain tools so no auto-detect occurs t.Setenv("PATH", "") @@ -1394,6 +1444,10 @@ func TestDaemonChildRun_HelperProcess(t *testing.T) { port := ln.Addr().(*net.TCPAddr).Port // Keep ln open so server can't bind - set the env var BEFORE calling main() _ = os.Setenv("ONWATCH_PORT", strconv.Itoa(port)) + // Bind the server to the exact address held above. The default + // 0.0.0.0 would still succeed on Windows, which lets a wildcard bind + // coexist with a 127.0.0.1 listener on the same port. + _ = os.Setenv("ONWATCH_HOST", "127.0.0.1") os.Args = []string{"onwatch", "--debug", "--test"} main() // server fails to bind → serverErr → run() logs error and exits ln.Close() @@ -1405,6 +1459,11 @@ func TestDaemonChildRun_HelperProcess(t *testing.T) { // runDaemonSubprocess starts a test subprocess in daemon/debug mode, // waits briefly, then sends SIGINT and waits for exit. +// +// Callers sandbox the child with HOME, USERPROFILE and LOCALAPPDATA all set to +// a temp dir: os.UserHomeDir reads USERPROFILE on Windows and the Windows PID +// directory lives under LOCALAPPDATA, so HOME alone would leave a Windows +// child writing PID files into the real profile. func runDaemonSubprocess(t *testing.T, env []string, waitMs int) { t.Helper() @@ -1416,7 +1475,12 @@ func runDaemonSubprocess(t *testing.T, env []string, waitMs int) { } time.Sleep(time.Duration(waitMs) * time.Millisecond) - _ = cmd.Process.Signal(os.Interrupt) + // Windows cannot deliver os.Interrupt to another process; stop it the way + // onWatch itself does there (terminateProcess) instead of waiting out the + // timeout below. + if err := cmd.Process.Signal(os.Interrupt); err != nil { + _ = cmd.Process.Kill() + } done := make(chan error, 1) go func() { done <- cmd.Wait() }() @@ -1457,6 +1521,8 @@ func TestDaemonChildRun_DebugModeAntigravity(t *testing.T) { "COPILOT_TOKEN=", "CODEX_TOKEN=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, "ONWATCH_ADMIN_PASS=testpass", @@ -1490,6 +1556,8 @@ func TestDaemonChildRun_DebugModeSyntheticProvider(t *testing.T) { "CODEX_TOKEN=", "ANTIGRAVITY_ENABLED=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, "ONWATCH_ADMIN_PASS=testpass", @@ -1524,6 +1592,8 @@ func TestDaemonChildRun_DebugModeAllProviders(t *testing.T) { "CODEX_TOKEN=fake-codex-token", "ANTIGRAVITY_ENABLED=true", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, "ONWATCH_ADMIN_PASS=testpass", @@ -1536,9 +1606,6 @@ func TestDaemonChildRun_DebugModeAllProviders(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStop_NonTestMode_WithPIDFilePort_LocalListener(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only available on macOS/Linux") - } // Skip if real onwatch is running - runStop(false) scans default ports as fallback for _, p := range []int{9211, 8932} { conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", p), 200*time.Millisecond) @@ -1578,9 +1645,6 @@ func TestRunStop_NonTestMode_WithPIDFilePort_LocalListener(t *testing.T) { } func TestRunStop_NonTestMode_DefaultPortScan(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only available on macOS/Linux") - } // Skip if a real onwatch is running on default ports for _, p := range []int{9211, 8932} { @@ -1611,9 +1675,6 @@ func TestRunStop_NonTestMode_DefaultPortScan(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStatus_NonTestMode_WithLocalListener(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only available on macOS/Linux") - } // Skip if a real onwatch is running on default ports for _, p := range []int{9211, 8932} { conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", p), 500*time.Millisecond) @@ -1702,7 +1763,7 @@ func TestRunStatus_LegacyPIDFormat(t *testing.T) { func TestRunSetup_ExistingEnvNoProviders_FreshSetup(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("PATH", "") t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) @@ -1748,7 +1809,7 @@ func TestRunSetup_ExistingEnvNoProviders_FreshSetup(t *testing.T) { func TestRunSetup_ExistingEnvSomeProviders_AddsMore(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("PATH", "") t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) @@ -1786,7 +1847,7 @@ func TestRunSetup_ExistingEnvSomeProviders_AddsMore(t *testing.T) { func TestCollectMultipleProviders_AllNo(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("PATH", "") t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) @@ -1815,7 +1876,7 @@ func TestCollectMultipleProviders_AllNo(t *testing.T) { func TestCollectMultipleProviders_AnthropicAndCodexAdded(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("PATH", "") t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) @@ -1891,6 +1952,8 @@ func TestDaemonChildRun_DebugModeLogLevels(t *testing.T) { "COPILOT_TOKEN=", "CODEX_TOKEN=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, "ONWATCH_ADMIN_PASS=testpass", @@ -1941,9 +2004,6 @@ func TestDaemonChildRun_AutoTokenDetection(t *testing.T) { if testing.Short() { t.Skip("skipping daemon subprocess test in short mode") } - if runtime.GOOS == "windows" { - t.Skip("unix credential file path test") - } home := t.TempDir() dbPath := filepath.Join(home, "onwatch.db") @@ -1987,6 +2047,8 @@ func TestDaemonChildRun_AutoTokenDetection(t *testing.T) { "COPILOT_TOKEN=", "ANTIGRAVITY_ENABLED=true", // Need at least one provider "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, "CODEX_HOME="+codexHome, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, @@ -2029,6 +2091,8 @@ func TestDaemonize_ViaSubprocess(t *testing.T) { "COPILOT_TOKEN=", "CODEX_TOKEN=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, "ONWATCH_ADMIN_PASS=testpass", @@ -2051,14 +2115,25 @@ func TestDaemonize_ViaSubprocess(t *testing.T) { t.Log("subprocess timed out - daemon may have started but took too long") } - // Kill any spawned daemon children - // PID file goes to $HOME/.onwatch/onwatch.pid (not dbDir) + // Kill any spawned daemon children. The PID file goes to the child's + // defaultPIDDir (not dbDir): $HOME/.onwatch on Unix, %LOCALAPPDATA%\onwatch + // on Windows. pidPath := filepath.Join(home, ".onwatch", "onwatch.pid") + if runtime.GOOS == "windows" { + pidPath = filepath.Join(home, "onwatch", "onwatch.pid") + } if data, err := os.ReadFile(pidPath); err == nil { if pid, err := strconv.Atoi(strings.Split(strings.TrimSpace(string(data)), ":")[0]); err == nil && pid > 0 { if proc, err := os.FindProcess(pid); err == nil { proc.Kill() } + // The daemon holds its log file (inside home) open as stdout, and + // Windows cannot delete an open file, so wait for it to exit before + // t.TempDir cleanup. It is not our child, so poll instead of Wait. + deadline := time.Now().Add(5 * time.Second) + for processAlive(pid) && time.Now().Before(deadline) { + time.Sleep(50 * time.Millisecond) + } } } } @@ -2090,6 +2165,8 @@ func TestDaemonChildRun_MigrateDBPath(t *testing.T) { "COPILOT_TOKEN=", "CODEX_TOKEN=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_ADMIN_PASS=testpass", // No ONWATCH_DB_PATH -> uses default -> DBPathExplicit=false -> migrateDBLocation @@ -2124,6 +2201,8 @@ func TestDaemonChildRun_DefaultPassword(t *testing.T) { "COPILOT_TOKEN=", "CODEX_TOKEN=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, // No ONWATCH_ADMIN_PASS -> default password @@ -2159,6 +2238,8 @@ func TestDaemonChildRun_WithAntigravityManualURL(t *testing.T) { "COPILOT_TOKEN=", "CODEX_TOKEN=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, "ONWATCH_ADMIN_PASS=testpass", @@ -2219,10 +2300,6 @@ func TestSleepHelperProcess_NeverRun(t *testing.T) { } func TestRunStop_WithLivePIDAndPort(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") t.Cleanup(func() { pidFile = oldPIDFile }) @@ -2250,10 +2327,6 @@ func TestRunStop_WithLivePIDAndPort(t *testing.T) { } func TestRunStop_WithLivePIDLegacyFormat(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") t.Cleanup(func() { pidFile = oldPIDFile }) @@ -2284,10 +2357,6 @@ func TestRunStop_WithLivePIDLegacyFormat(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStatus_WithLivePIDAndPort(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") t.Cleanup(func() { pidFile = oldPIDFile }) @@ -2315,10 +2384,6 @@ func TestRunStatus_WithLivePIDAndPort(t *testing.T) { } func TestRunStatus_WithLivePIDLegacyFormat(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") t.Cleanup(func() { pidFile = oldPIDFile }) @@ -2349,10 +2414,6 @@ func TestRunStatus_WithLivePIDLegacyFormat(t *testing.T) { // --------------------------------------------------------------------------- func TestStopPreviousInstance_WithLivePID(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") t.Cleanup(func() { pidFile = oldPIDFile }) @@ -2380,10 +2441,6 @@ func TestStopPreviousInstance_WithLivePID(t *testing.T) { } func TestStopPreviousInstance_WithLivePIDLegacyFormat(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") t.Cleanup(func() { pidFile = oldPIDFile }) @@ -2410,7 +2467,7 @@ func TestStopPreviousInstance_WithLivePIDLegacyFormat(t *testing.T) { func TestMigrateDBLocation_MkdirFails(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Create a file where the old DB is expected - to simulate "file exists" but in wrong place oldDB := filepath.Join(home, ".onwatch", "onwatch.db") @@ -2435,7 +2492,7 @@ func TestMigrateDBLocation_MkdirFails(t *testing.T) { func TestMigrateDBLocation_RenameFails(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Create old DB oldDB := filepath.Join(home, ".onwatch", "onwatch.db") @@ -2542,6 +2599,8 @@ func TestDaemonChildRun_DebugModeAllProvidersWithCopilot(t *testing.T) { "CODEX_TOKEN=codex-test-token", "ANTIGRAVITY_ENABLED=true", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+dbPath, "ONWATCH_ADMIN_PASS=testpass", @@ -2554,10 +2613,6 @@ func TestDaemonChildRun_DebugModeAllProvidersWithCopilot(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStatus_ShowsDashboardURL(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") t.Cleanup(func() { pidFile = oldPIDFile }) @@ -2591,10 +2646,6 @@ func TestRunStatus_ShowsDashboardURL(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStop_ShowsPortInOutput(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch.pid") t.Cleanup(func() { pidFile = oldPIDFile }) @@ -2628,9 +2679,6 @@ func TestDaemonChildRun_ServerError(t *testing.T) { if testing.Short() { t.Skip("skipping server error subprocess test in short mode") } - if runtime.GOOS == "windows" { - t.Skip("server bind error test is unix-specific") - } home := t.TempDir() dbPath := filepath.Join(home, "onwatch.db") @@ -2649,6 +2697,8 @@ func TestDaemonChildRun_ServerError(t *testing.T) { "COPILOT_TOKEN=", "CODEX_TOKEN=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, "ONWATCH_DB_PATH="+dbPath, "ONWATCH_ADMIN_PASS=testpass", // ONWATCH_PORT is set by the subprocess itself (server_error_test mode) @@ -2702,6 +2752,8 @@ func TestDaemonChildRun_FixExplicitDBPath(t *testing.T) { "COPILOT_TOKEN=", "CODEX_TOKEN=", "HOME="+home, + "USERPROFILE="+home, + "LOCALAPPDATA="+home, fmt.Sprintf("ONWATCH_PORT=%d", port), "ONWATCH_DB_PATH="+explicitDB, // explicit path → DBPathExplicit=true → fixExplicitDBPath "ONWATCH_ADMIN_PASS=testpass", @@ -2713,10 +2765,6 @@ func TestDaemonChildRun_FixExplicitDBPath(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStatus_WithLogAndDBFiles(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("SIGTERM not supported the same way on Windows") - } - oldPIDFile := pidFile pidDir := t.TempDir() pidFile = filepath.Join(pidDir, "onwatch.pid") @@ -2734,15 +2782,17 @@ func TestRunStatus_WithLogAndDBFiles(t *testing.T) { } // Create the log file in the current dir so the stat check succeeds - // (testMode=true → logPath = ".onwatch-test.log") + // (testMode=true → logPath = ".onwatch-test.log"). Work in a temp dir so + // the test never writes into the package directory. + t.Chdir(t.TempDir()) logFile := ".onwatch-test.log" if err := os.WriteFile(logFile, []byte("log data\n"), 0o600); err != nil { t.Fatalf("write log file: %v", err) } - t.Cleanup(func() { os.Remove(logFile) }) // Create DB file in home/.onwatch/data/onwatch.db - home, _ := os.UserHomeDir() + home := t.TempDir() + testhome.SetTestHome(t, home) dbDir := filepath.Join(home, ".onwatch", "data") if mkErr := os.MkdirAll(dbDir, 0o755); mkErr == nil { dbFile := filepath.Join(dbDir, "onwatch.db") @@ -3305,7 +3355,7 @@ func TestFixExplicitDBPath_ExplicitNotExist(t *testing.T) { // test the case where explicit path doesn't exist by pointing cfg.DBPath // to a nonexistent file within the temp home. tmpHome := filepath.Dir(filepath.Dir(canonDir)) // the temp dir itself - t.Setenv("HOME", tmpHome) + testhome.SetTestHome(t, tmpHome) cfg := &config.Config{ DBPath: filepath.Join(t.TempDir(), "nonexistent.db"), @@ -3380,9 +3430,6 @@ func TestStopPreviousInstance_EmptyPIDFile(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStop_NonTestMode_PIDFilePortBranch(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only on macOS/Linux") - } // Skip if real onwatch on default ports for _, p := range []int{9211, 8932} { conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", p), 500*time.Millisecond) @@ -3426,9 +3473,6 @@ func TestRunStop_NonTestMode_PIDFilePortBranch(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStatus_NonTestMode_PIDFilePortBranch(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only on macOS/Linux") - } // Skip if real onwatch on default ports for _, p := range []int{9211, 8932} { conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", p), 500*time.Millisecond) @@ -3476,9 +3520,6 @@ func TestRunStatus_NonTestMode_PIDFilePortBranch(t *testing.T) { } func TestRunStatus_NonTestMode_NoPIDFileFallback(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only on macOS/Linux") - } // Skip if real onwatch on default ports for _, p := range []int{9211, 8932} { conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", p), 500*time.Millisecond) @@ -3508,9 +3549,6 @@ func TestRunStatus_NonTestMode_NoPIDFileFallback(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStatus_NonTestMode_RunningProcessNoPort(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only on macOS/Linux") - } oldPIDFile := pidFile tmpDir := t.TempDir() @@ -3544,9 +3582,6 @@ func TestRunStatus_NonTestMode_RunningProcessNoPort(t *testing.T) { // --------------------------------------------------------------------------- func TestStopPreviousInstance_NonTestMode_PortFallback(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only on macOS/Linux") - } oldPIDFile := pidFile oldPIDDir := pidDir @@ -3574,9 +3609,6 @@ func TestStopPreviousInstance_NonTestMode_PortFallback(t *testing.T) { } func TestStopPreviousInstance_NonTestMode_PIDFileWithPort(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only on macOS/Linux") - } oldPIDFile := pidFile oldPIDDir := pidDir @@ -3615,7 +3647,7 @@ func TestRun_SetupCommand(t *testing.T) { // "all providers configured" early return instead of entering the // interactive wizard (which loops forever on EOF stdin in CI). home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) installDir := filepath.Join(home, ".onwatch") if err := os.MkdirAll(filepath.Join(installDir, "data"), 0755); err != nil { t.Fatalf("mkdir: %v", err) @@ -3669,7 +3701,7 @@ func TestRun_InProcessDaemonChild_ServerBindFails(t *testing.T) { t.Setenv("ONWATCH_ADMIN_PASS", "testpass123") t.Setenv("ONWATCH_PORT", strconv.Itoa(port)) t.Setenv("ONWATCH_LOG_LEVEL", "error") - t.Setenv("HOME", tmpDir) + testhome.SetTestHome(t, tmpDir) oldPIDFile := pidFile oldPIDDir := pidDir @@ -3731,7 +3763,7 @@ func TestRun_InProcessDaemonChild_AllProviders(t *testing.T) { t.Setenv("ONWATCH_ADMIN_PASS", "testpass456") t.Setenv("ONWATCH_PORT", strconv.Itoa(port)) t.Setenv("ONWATCH_LOG_LEVEL", "error") - t.Setenv("HOME", tmpDir) + testhome.SetTestHome(t, tmpDir) oldPIDFile := pidFile oldPIDDir := pidDir @@ -3779,7 +3811,7 @@ func TestCollectSyntheticKey_EmptyThenValid(t *testing.T) { func TestAddMissingProviders_AntigravityAlreadyEnabled(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "no-codex")) t.Setenv("PATH", "") @@ -3832,7 +3864,7 @@ func TestFixExplicitDBPath_AlreadyCanonical(t *testing.T) { func TestFixExplicitDBPath_CanonicalHasMoreData(t *testing.T) { tmpHome := t.TempDir() - t.Setenv("HOME", tmpHome) + testhome.SetTestHome(t, tmpHome) // Create canonical path with large data canonDir := filepath.Join(tmpHome, ".onwatch", "data") @@ -3868,7 +3900,7 @@ func TestFixExplicitDBPath_CanonicalHasMoreData(t *testing.T) { func TestFixExplicitDBPath_CanonicalDoesNotExist(t *testing.T) { tmpHome := t.TempDir() - t.Setenv("HOME", tmpHome) + testhome.SetTestHome(t, tmpHome) // No canonical path created @@ -3893,7 +3925,7 @@ func TestFixExplicitDBPath_CanonicalDoesNotExist(t *testing.T) { func TestFixExplicitDBPath_ExplicitMissingCanonicalExists(t *testing.T) { tmpHome := t.TempDir() - t.Setenv("HOME", tmpHome) + testhome.SetTestHome(t, tmpHome) // Create canonical with data canonDir := filepath.Join(tmpHome, ".onwatch", "data") @@ -3984,9 +4016,6 @@ func TestInitEncryptionSalt_InvalidSaltInDB(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStop_NonTestMode_StalePIDWithPortFallback(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only on macOS/Linux") - } for _, p := range []int{9211, 8932} { conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", p), 500*time.Millisecond) if err == nil { @@ -4033,9 +4062,6 @@ func TestRunStop_NonTestMode_StalePIDWithPortFallback(t *testing.T) { // --------------------------------------------------------------------------- func TestRunStatus_NonTestMode_StalePIDNoPort(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("lsof only on macOS/Linux") - } for _, p := range []int{9211, 8932} { conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", p), 500*time.Millisecond) if err == nil { diff --git a/cmd/onwatch/root_coverage_test.go b/cmd/onwatch/root_coverage_test.go index d85a5383..7088605a 100644 --- a/cmd/onwatch/root_coverage_test.go +++ b/cmd/onwatch/root_coverage_test.go @@ -14,12 +14,12 @@ import ( "runtime" "strconv" "strings" - "syscall" "testing" "time" "github.com/onllm-dev/onwatch/v2/internal/config" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" "github.com/onllm-dev/onwatch/v2/internal/update" "github.com/onllm-dev/onwatch/v2/internal/web" ) @@ -31,18 +31,37 @@ func captureStdout(t *testing.T, fn func()) string { if err != nil { t.Fatalf("create stdout pipe: %v", err) } - defer r.Close() + + // Drain the pipe while fn runs: a pipe buffer is small (a few KB on + // Windows), so a chatty fn would otherwise block forever on a full pipe. + var out []byte + var readErr error + done := make(chan struct{}) + go func() { + out, readErr = io.ReadAll(r) + close(done) + }() + os.Stdout = w - defer func() { os.Stdout = oldStdout }() + // Runs on every exit, including a t.Fatal inside fn. Close the writer + // before the reader: on Windows closing a pipe handle waits for the + // blocked read on it, which only ends once the writer is closed. + defer func() { + os.Stdout = oldStdout + _ = w.Close() + <-done + _ = r.Close() + }() fn() + os.Stdout = oldStdout if err := w.Close(); err != nil { t.Fatalf("close writer: %v", err) } - out, err := io.ReadAll(r) - if err != nil { - t.Fatalf("read stdout: %v", err) + <-done + if readErr != nil { + t.Fatalf("read stdout: %v", readErr) } return string(out) } @@ -110,7 +129,7 @@ func TestPIDFileLifecycle(t *testing.T) { func TestMigrateDBLocation_MovesDBAndSidecars(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) oldDB := filepath.Join(home, ".onwatch", "onwatch.db") newDB := filepath.Join(home, ".onwatch", "data", "onwatch.db") @@ -145,7 +164,7 @@ func TestMigrateDBLocation_MovesDBAndSidecars(t *testing.T) { func TestFixExplicitDBPath_RedirectsToCanonicalWhenBetter(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) canonical := filepath.Join(home, ".onwatch", "data", "onwatch.db") if err := os.MkdirAll(filepath.Dir(canonical), 0o755); err != nil { t.Fatalf("mkdir canonical dir: %v", err) @@ -416,7 +435,7 @@ func TestPrintSummaryAndNextSteps(t *testing.T) { func TestRunSetupEarlyPathsAndSafeRunCommands(t *testing.T) { t.Run("runSetup returns early when all providers already configured", func(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) installDir := filepath.Join(home, ".onwatch") envFile := filepath.Join(installDir, ".env") if err := os.MkdirAll(filepath.Join(installDir, "data"), 0o755); err != nil { @@ -441,7 +460,7 @@ func TestRunSetupEarlyPathsAndSafeRunCommands(t *testing.T) { t.Run("runSetup fresh safe path", func(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) input := strings.Join([]string{ "6", // antigravity only "1", // antigravity source: both @@ -491,22 +510,22 @@ func TestRunSetupEarlyPathsAndSafeRunCommands(t *testing.T) { } func TestRunStopAndStatus_WithPIDFileProcess(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("sleep process helper not used on windows") - } - oldPIDFile := pidFile pidFile = filepath.Join(t.TempDir(), "onwatch-test.pid") t.Cleanup(func() { pidFile = oldPIDFile }) - cmd := exec.Command("sleep", "30") - if err := cmd.Start(); err != nil { - t.Fatalf("start sleep process: %v", err) - } + // Re-exec the test binary as the long-running child rather than sleep(1), + // which does not exist on Windows. + cmd := startSleepSubprocess(t) + waitDone := make(chan error, 1) + go func() { + waitDone <- cmd.Wait() + }() t.Cleanup(func() { - if cmd.Process != nil { - _ = cmd.Process.Kill() - _, _ = cmd.Process.Wait() + _ = cmd.Process.Kill() + select { + case <-waitDone: + case <-time.After(5 * time.Second): } }) @@ -535,13 +554,9 @@ func TestRunStopAndStatus_WithPIDFileProcess(t *testing.T) { t.Fatalf("unexpected runStop output: %s", stopOut) } - waitDone := make(chan error, 1) - go func() { - waitDone <- cmd.Wait() - }() select { case <-waitDone: - case <-time.After(2 * time.Second): + case <-time.After(5 * time.Second): t.Fatal("process did not stop after runStop") } } @@ -740,8 +755,14 @@ func TestDaemonize_SuccessAndLogOpenError(t *testing.T) { return } if pid := parsePIDContent(string(data)); pid > 0 && pid != os.Getpid() { - if proc, err := os.FindProcess(pid); err == nil { - _ = proc.Signal(syscall.SIGTERM) + // stopProcess, not SIGTERM: Windows cannot deliver signals. + // The child holds the log in tmp open as stdout and Windows + // cannot delete an open file, so wait for it to exit before + // tmp is removed. It is not our child, so poll instead of Wait. + stopProcess(pid) + deadline := time.Now().Add(5 * time.Second) + for processAlive(pid) && time.Now().Before(deadline) { + time.Sleep(50 * time.Millisecond) } } }) diff --git a/cmd/onwatch/root_more_coverage_test.go b/cmd/onwatch/root_more_coverage_test.go index 2d2eb3bb..4fedcd4f 100644 --- a/cmd/onwatch/root_more_coverage_test.go +++ b/cmd/onwatch/root_more_coverage_test.go @@ -4,10 +4,11 @@ import ( "bufio" "os" "path/filepath" - "runtime" "strconv" "strings" "testing" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func setTestArgs(t *testing.T, args []string) { @@ -155,7 +156,7 @@ func TestSetupHelpers_AddMissingProvidersAndTokenCollectors(t *testing.T) { t.Run("collectAnthropicToken and collectCodexToken stay deterministic", func(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "missing-codex")) anthReader := bufio.NewReader(strings.NewReader("\nmanual-anth-token\n")) @@ -172,19 +173,6 @@ func TestSetupHelpers_AddMissingProvidersAndTokenCollectors(t *testing.T) { }) } -func TestDaemonSysProcAttr_UnixSetsid(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("unix-only test") - } - attr := daemonSysProcAttr() - if attr == nil { - t.Fatal("expected non-nil SysProcAttr") - } - if !attr.Setsid { - t.Fatal("expected Setsid=true") - } -} - func TestRun_HelpCommand(t *testing.T) { setTestArgs(t, []string{"onwatch", "--help"}) out := captureStdout(t, func() { @@ -199,7 +187,7 @@ func TestRun_HelpCommand(t *testing.T) { func TestMain_ErrorPath(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("ONWATCH_PORT", "1") // Clear all API keys for _, key := range []string{ diff --git a/cmd/onwatch/service_cmd.go b/cmd/onwatch/service_cmd.go index 1157cb92..ba5eaf5f 100644 --- a/cmd/onwatch/service_cmd.go +++ b/cmd/onwatch/service_cmd.go @@ -8,7 +8,6 @@ import ( "path/filepath" "strconv" "strings" - "syscall" "time" "github.com/mattn/go-isatty" @@ -296,11 +295,9 @@ func runningDaemonPID() (int, bool) { if pid <= 0 || pid == os.Getpid() { return 0, false } - proc, err := os.FindProcess(pid) - if err != nil { - return 0, false - } - if err := proc.Signal(syscall.Signal(0)); err != nil { + // processAlive, not proc.Signal(0): Windows cannot deliver signals, so a + // signal probe reports every daemon there as not running. + if !processAlive(pid) { return 0, false } // The PID file can name a PID the OS has since recycled onto an unrelated @@ -315,11 +312,7 @@ func runningDaemonPID() (int, bool) { func waitForExit(pid int, timeout time.Duration) { deadline := time.Now().Add(timeout) for time.Now().Before(deadline) { - proc, err := os.FindProcess(pid) - if err != nil { - return - } - if err := proc.Signal(syscall.Signal(0)); err != nil { + if !processAlive(pid) { return } time.Sleep(100 * time.Millisecond) @@ -373,8 +366,9 @@ func restartAfterUpdate() { if pid, running := runningDaemonPID(); running { fmt.Println("Restarting daemon...") - if proc, err := os.FindProcess(pid); err == nil { - _ = proc.Signal(syscall.SIGTERM) + // stopProcess: SIGTERM on Unix, TerminateProcess on Windows, where + // proc.Signal(SIGTERM) always fails and the old daemon kept running. + if stopProcess(pid) { waitForExit(pid, 5*time.Second) } } else { diff --git a/cmd/onwatch/service_cmd_test.go b/cmd/onwatch/service_cmd_test.go index 9456b005..3019e97f 100644 --- a/cmd/onwatch/service_cmd_test.go +++ b/cmd/onwatch/service_cmd_test.go @@ -10,8 +10,10 @@ import ( "path/filepath" "strings" "testing" + "time" "github.com/onllm-dev/onwatch/v2/internal/service" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" "github.com/onllm-dev/onwatch/v2/internal/update" ) @@ -75,7 +77,7 @@ func newAutostartHarness(t *testing.T) *autostartHarness { func isolateHome(t *testing.T) string { t.Helper() home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) prevPID := pidFile pidFile = filepath.Join(home, "onwatch.pid") t.Cleanup(func() { pidFile = prevPID }) @@ -120,6 +122,57 @@ func TestRunningDaemonPID(t *testing.T) { if _, ok := runningDaemonPID(); ok { t.Error("own PID should report not running") } + + // A live onwatch process is reported running. The helper child is this + // test binary, onwatch.test(.exe), so it passes the onwatch-name check. + // This must hold on Windows too, where a signal-0 probe always fails. + cmd := startSleepSubprocess(t) + t.Cleanup(func() { _ = cmd.Process.Kill(); _ = cmd.Wait() }) + if err := os.WriteFile(pidFile, []byte(fmt.Sprintf("%d:9211", cmd.Process.Pid)), 0o644); err != nil { + t.Fatal(err) + } + if pid, ok := runningDaemonPID(); !ok || pid != cmd.Process.Pid { + t.Errorf("runningDaemonPID() = %d, %v; want %d, true", pid, ok, cmd.Process.Pid) + } +} + +func TestIsOnwatchProcess(t *testing.T) { + if !isOnwatchProcess(os.Getpid()) { + t.Error("the onwatch.test binary itself must be recognised as onwatch") + } + if isOnwatchProcess(0) || isOnwatchProcess(-1) { + t.Error("non-positive PIDs are never onwatch") + } +} + +// A running daemon must be stopped before the updated one starts. On Windows +// the old SIGTERM was a no-op and the daemon was never even seen as running. +func TestRestartAfterUpdateStopsRunningDaemon(t *testing.T) { + isolateHome(t) + h := newAutostartHarness(t) + h.spawnPID = 4242 + + cmd := startSleepSubprocess(t) + exited := make(chan struct{}) + go func() { _ = cmd.Wait(); close(exited) }() + t.Cleanup(func() { _ = cmd.Process.Kill(); <-exited }) + if err := os.WriteFile(pidFile, []byte(fmt.Sprintf("%d:9211", cmd.Process.Pid)), 0o644); err != nil { + t.Fatal(err) + } + + out := captureStdout(t, restartAfterUpdate) + + if !strings.Contains(out, "Restarting daemon") { + t.Errorf("expected the running daemon to be detected, got: %s", out) + } + select { + case <-exited: + case <-time.After(5 * time.Second): + t.Fatal("the old daemon was not stopped") + } + if len(h.spawns) != 1 { + t.Errorf("expected exactly one daemon spawn, got %d", len(h.spawns)) + } } // The reported bug: after a reboot nothing is running, and `onwatch update` @@ -633,11 +686,13 @@ func TestTestDaemonIsolationEnv(t *testing.T) { t.Fatal("a test binary must get isolation overrides") } - var home, db, port string + var home, profile, db, port string for _, kv := range env { switch { case strings.HasPrefix(kv, "HOME="): home = strings.TrimPrefix(kv, "HOME=") + case strings.HasPrefix(kv, "USERPROFILE="): + profile = strings.TrimPrefix(kv, "USERPROFILE=") case strings.HasPrefix(kv, "ONWATCH_DB_PATH="): db = strings.TrimPrefix(kv, "ONWATCH_DB_PATH=") case strings.HasPrefix(kv, "ONWATCH_PORT="): @@ -654,6 +709,11 @@ func TestTestDaemonIsolationEnv(t *testing.T) { if home == "" || home == realHome { t.Errorf("HOME override = %q, must be a scratch directory", home) } + // os.UserHomeDir reads USERPROFILE on Windows, so HOME alone would leave a + // Windows child in the real profile. + if profile != home { + t.Errorf("USERPROFILE override = %q, must match HOME %q", profile, home) + } if db == "" || strings.HasPrefix(db, realHome) { t.Errorf("ONWATCH_DB_PATH = %q, must not point into the real install", db) } diff --git a/cmd/onwatch/setup.go b/cmd/onwatch/setup.go index dc3dd3c8..a4832014 100644 --- a/cmd/onwatch/setup.go +++ b/cmd/onwatch/setup.go @@ -935,8 +935,14 @@ func addMissingProviders(reader *bufio.Reader, envFile string, existing *existin } if !existing.geminiEnabled { - // Try to detect Gemini CLI credentials - if _, err := os.Stat(filepath.Join(os.Getenv("HOME"), ".gemini", "oauth_creds.json")); err == nil { + // Try to detect Gemini CLI credentials. Use os.UserHomeDir, not $HOME: + // Windows keeps the profile in USERPROFILE. + geminiDetected := false + if home, err := os.UserHomeDir(); err == nil && home != "" { + _, statErr := os.Stat(filepath.Join(home, ".gemini", "oauth_creds.json")) + geminiDetected = statErr == nil + } + if geminiDetected { fmt.Printf(" %s ok %s Gemini CLI credentials detected on this system\n", colorGreen, colorReset) if promptYesNo(reader, "Enable Gemini tracking?", true) { fmt.Fprintf(f, "\n# Gemini CLI - auto-detected from ~/.gemini/oauth_creds.json\nGEMINI_ENABLED=true\n") diff --git a/cmd/onwatch/setup_commandcode_test.go b/cmd/onwatch/setup_commandcode_test.go index 29b3fd84..2c0ca82b 100644 --- a/cmd/onwatch/setup_commandcode_test.go +++ b/cmd/onwatch/setup_commandcode_test.go @@ -10,6 +10,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/api" "github.com/onllm-dev/onwatch/v2/internal/config" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func stubCommandCodeDetect(t *testing.T, creds *api.CommandCodeCredentials) { @@ -98,7 +99,7 @@ func TestCollectCommandCodeVerifyFailureSavesKey(t *testing.T) { func TestCommandCodeEitherSourceSatisfies(t *testing.T) { isolate := func(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("COMMAND_CODE_API_KEY", "") t.Setenv("COMMANDCODE_API_KEY", "") t.Setenv("COMMANDCODE_ENABLED", "") diff --git a/cmd/onwatch/setup_muse_test.go b/cmd/onwatch/setup_muse_test.go index b3fefc8b..10b2530e 100644 --- a/cmd/onwatch/setup_muse_test.go +++ b/cmd/onwatch/setup_muse_test.go @@ -8,6 +8,7 @@ import ( "testing" "github.com/onllm-dev/onwatch/v2/internal/api" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func stubMuseDetect(t *testing.T, creds *api.MuseCredentials) { @@ -107,7 +108,7 @@ func readSetupTestFile(t *testing.T, path string) string { func TestFreshSetup_MuseOnly(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("PATH", "") stubMuseDetect(t, &api.MuseCredentials{APIKey: "login-key", Model: "m", Source: "keychain"}) diff --git a/docs/OPENCODE_SETUP.md b/docs/OPENCODE_SETUP.md index 04a705a6..9a7faa65 100644 --- a/docs/OPENCODE_SETUP.md +++ b/docs/OPENCODE_SETUP.md @@ -2,111 +2,57 @@ Track OpenCode Go subscription quotas in onWatch. -OpenCode Go has no documented quota API. onWatch supports two modes: +OpenCode Go has no documented quota API. onWatch reads the subscription's own meters from the console endpoint `GET /console/api/go/status`, and it can authenticate in either of two ways: -- **Usage API (recommended).** A console service-account key reads the Go subscription's own meters from `/console/api/go/status`. No browser cookie, and it works with the current OpenCode console. -- **Dashboard scrape (legacy).** Your workspace ID and browser `auth` cookie are used to scrape `/workspace/{id}/go`. Used only when no usage API key is set. The dashboard moved to a new console, so scraping may no longer find the usage data. +- **Service-account key (recommended).** A console key with usage read access. No browser cookie, and it does not expire when you log out. +- **Browser session.** Your workspace ID plus the `__Host-console_session` cookie from opencode.ai. Used only when no key is set. ---- +This is separate from `OPENCODE_ENABLED`, which only feeds ChatGPT credentials from OpenCode into the **Codex** provider. -## Usage API mode (recommended) +> **Upgrading from the old scrape mode?** The `/workspace/{id}/go` page and the `auth` cookie no longer work ([issue #134](https://github.com/onllm-dev/onWatch/issues/134)). Keep your workspace ID, but replace `OPENCODE_GO_AUTH_COOKIE` with the value of the `__Host-console_session` cookie, or switch to a service-account key. -1. In the OpenCode console, open **Keys → Add Service Account** and create a key with **usage read** access. Keys look like `oc_sk_...`. -2. Set it (or paste it into **Settings → OpenCode Go → Usage API Key**): +--- - ```bash - OPENCODE_GO_API_KEY=oc_sk_... - ``` +## What you see -**Where the numbers come from.** `GET /console/api/go/status` returns the subscription's three meters, each with a used and a limit amount. They are the figures OpenCode enforces and shows on the Go page. onWatch displays used ÷ limit as the **5-Hour**, **Weekly** and **Monthly** cards, the same cards the scrape mode produces. Nothing is estimated, and there is no price table to keep in sync. +`go/status` returns three meters, each with a used and a limit amount. They are the figures OpenCode enforces and shows on the Go page. onWatch shows them as the **5-Hour**, **Weekly** and **Monthly** cards, with the percentage and the dollar amount (for example `$3.84 / $30.00`). Nothing is estimated, and there is no price table to keep in sync. - **5-Hour** is a session window. While a session is open it carries its reset time. With no open session it shows 0% and no reset time. - **Weekly** resets at the time the API reports. - **Monthly** renews with the subscription period. Its reset is the period end the API reports, so no reset day needs to be configured. -The meters belong to the Go subscription the key resolves to. To track a different subscription, use a key from that subscription's workspace. +Amounts are stored in USD. Snapshots taken before this change were stored as percentages (limit 100), so the logging history shows a short window of mixed units after upgrading. Charts and cycle history use the percentage and are unaffected. If you set an **absolute** notification threshold for an OpenCode quota, it is now compared in dollars; percentage thresholds are unchanged. onWatch reuses a status response for 60 seconds, so a poll interval shorter than that does not add requests. -This endpoint is undocumented. The console usage export (`/console/api/v1/usage/export`) is not used: its Go rows carry no cost, and since 25 September 2026 it answers 403 to service-account keys while OpenCode migrates its usage records. See [anomalyco/opencode#50912](https://github.com/anomalyco/opencode/issues/50912), which also describes `go/status`. - ---- - -## Dashboard scrape mode (legacy) - -### Prerequisites - -- An active [OpenCode Go](https://opencode.ai) subscription -- Access to the OpenCode Go dashboard in a browser -- onWatch installed ([Quick Start](../README.md#quick-start)) - --- -### How It Works +## Option 1: Service-account key (recommended) -onWatch polls: +1. In the OpenCode console, open **Keys -> Add Service Account** and create a key with **usage read** access. Keys look like `oc_sk_...`. +2. Set it (or paste it into **Settings -> Providers -> OpenCode Go -> Usage API Key**): -```text -https://opencode.ai/workspace/{workspaceId}/go -``` - -using your session cookie, then extracts utilization and reset countdowns for: - -- **5-Hour** (rolling / session window) -- **Weekly** -- **Monthly** (when present on the dashboard) - -Parsing tries SolidJS SSR hydration data first, then falls back to the newer `data-slot="usage-item"` HTML layout. Snapshots are stored locally in SQLite like every other provider. + ```bash + OPENCODE_GO_API_KEY=oc_sk_... + ``` -This is separate from `OPENCODE_ENABLED`, which only feeds ChatGPT credentials from OpenCode into the **Codex** provider. +The meters belong to the Go subscription the key resolves to. To track a different subscription, use a key from that subscription's workspace. --- -### 1. Find Your Workspace ID - -Open https://opencode.ai and sign in, then use either method below. +## Option 2: Browser session -#### From the browser URL +### 1. Copy the session cookie -Open your OpenCode Go usage page. The URL looks like: +1. Sign in at https://opencode.ai and open the console +2. Open browser Developer Tools -> **Application** / **Storage** -> **Cookies** -> `https://opencode.ai` +3. Find the `__Host-console_session` cookie and copy its **value** -```text -https://opencode.ai/workspace/wrk_xxxxxxxx/go -``` - -Copy the `wrk_...` segment. That is your `OPENCODE_GO_WORKSPACE_ID`. +You can paste just the value, `__Host-console_session=`, or a full `Cookie:` header copied from DevTools. Treat it like a password: logging out, rotating sessions, or clearing cookies invalidates it. -#### From the authenticated Go page +### 2. Find your workspace ID -Copy the `auth` cookie value as described in the next section, then run: - -```bash -curl -sS --compressed \ - -H 'Cookie: auth=' \ - https://opencode.ai/go | - grep -oE 'wrk_[A-Za-z0-9]+' | - sort -u -``` - -The authenticated `/go` page contains the workspace ID. The site root does not. -`https://opencode.ai/zen` can also expose it, but `/go` is preferred for an -OpenCode Go subscription. - -If the command returns multiple workspace IDs, use the one whose -`https://opencode.ai/workspace//go` page shows your Go usage. - ---- - -### 2. Copy the Auth Cookie - -1. While signed in on opencode.ai, open your browser Developer Tools -2. Go to **Application** / **Storage** → **Cookies** → `https://opencode.ai` -3. Find the `auth` cookie -4. Copy its **value** (not the `auth=` name prefix) - -Treat this cookie like a password. Logging out of OpenCode, rotating sessions, or clearing cookies will invalidate it. - ---- +With Developer Tools open on the **Network** tab, open your Go usage page in the console. Select the `status` request to `/console/api/go/status` and copy the value of its `x-org-id` request header. It looks like `wrk_...`. ### 3. Configure onWatch @@ -114,76 +60,46 @@ Add both values to `~/.onwatch/.env` (or your project `.env`): ```bash OPENCODE_GO_WORKSPACE_ID=wrk_xxxxxxxx -OPENCODE_GO_AUTH_COOKIE=your_auth_cookie_value +OPENCODE_GO_AUTH_COOKIE=your_console_session_value ``` -In scrape mode both are required. Without them (and without `OPENCODE_GO_API_KEY`) the OpenCode Go provider stays disabled. - -You can also set them in the dashboard: - -1. Open **Settings → Providers → OpenCode Go** -2. Paste **Workspace ID** and **Auth Cookie** -3. Save +Both are required. Without them (and without `OPENCODE_GO_API_KEY`) the OpenCode Go provider stays disabled. -Dashboard values override `.env` for the running process. A daemon restart may still be needed depending on how the agent was started. +You can also set them in **Settings -> Providers -> OpenCode Go** (**Workspace ID** and **Session Cookie**). --- -### 4. Reload / Restart +## Reload / Restart -Reload providers from Settings if available, or restart onWatch: +Settings changes take effect after a daemon restart: ```bash onwatch stop onwatch ``` -Or verify in the foreground: - -```bash -onwatch --debug -``` - -You should see the OpenCode agent start once it is configured. - ---- - -### 5. Verify - -- Open http://localhost:9211 -- Switch to the **OpenCode** tab -- Confirm 5-Hour / Weekly cards populate (Monthly appears when OpenCode returns it) -- Charts, cycle overview, and insights begin filling after a few polls - ---- - -## Dashboard - -The OpenCode Go tab shows: +Or verify in the foreground with `onwatch --debug`. You should see `OpenCode poll complete` with `quota_count=3`. -- Quota cards with utilization, remaining countdown, and status -- Historical chart across tracked windows -- Billing-cycle / usage-sample tables -- Burn-rate insights for the active windows +Then open http://localhost:9211, switch to the **OpenCode** tab, and confirm the 5-Hour, Weekly and Monthly cards populate. Charts, cycle overview and insights fill in after a few polls. --- ## Security Notes -- Never commit `.env` or paste the cookie into issue reports / logs +- Never commit `.env` or paste the key or cookie into issue reports / logs - onWatch redacts `api_key` and `auth_cookie` from `/api/settings` responses -- Scraped HTML and Go status responses are not written to logs +- `go/status` responses are never written to logs or echoed in errors (they contain account IDs) +- The session cookie is only sent to `opencode.ai`, and redirects are never followed - All processing stays local on your machine --- ## Limitations & Notes -- Scrape mode depends on undocumented dashboard HTML. OpenCode UI changes can break parsing until onWatch is updated. -- Auth failures and parse failures are surfaced as errors. onWatch does **not** invent fake currency quotas when scraping fails. -- Usage API mode depends on the undocumented `go/status` response. If OpenCode changes its shape, onWatch reports a parse failure rather than showing wrong numbers. -- Cookie lifetime is controlled by OpenCode. Expect to refresh the cookie after logout or session rotation. -- In scrape mode the workspace ID is required; onWatch does not auto-discover workspaces. +- `go/status` is undocumented. If OpenCode changes its shape, onWatch reports a parse failure rather than showing wrong numbers. +- Session cookie lifetime is controlled by OpenCode. Expect to refresh it after logout or session rotation; a service-account key avoids this. +- The workspace ID is required in session mode; onWatch does not auto-discover workspaces. +- The console usage export (`/console/api/v1/usage/export`) is not used: its Go rows carry no cost, and since 25 September 2026 it answers 403 to service-account keys while OpenCode migrates its usage records. See [anomalyco/opencode#50912](https://github.com/anomalyco/opencode/issues/50912), which also describes `go/status`. --- @@ -193,21 +109,20 @@ The OpenCode Go tab shows: - Confirm `OPENCODE_GO_API_KEY` is set, or both `OPENCODE_GO_WORKSPACE_ID` and `OPENCODE_GO_AUTH_COOKIE` - Restart onWatch and check `--debug` logs for missing-config messages -- In Settings → Providers, confirm OpenCode Go shows as configured / polling -### Unauthorized / forbidden / empty data +### Unauthorized / forbidden -- Usage API mode: check the service-account key still exists and has usage read access. A 403 on a key that used to work usually means OpenCode changed what service-account keys may read. A 429 means the API is rate limiting; onWatch skips that poll and tries again on the next one. -- Scrape mode: re-copy a fresh `auth` cookie while signed in, confirm the workspace ID matches the `/go` URL, and check the Go dashboard still loads in your browser. -- Restart onWatch. +- **Key:** check the service-account key still exists and has usage read access. A 403 on a key that used to work usually means OpenCode changed what service-account keys may read. +- **Session:** the log says `paste the __Host-console_session cookie`. Re-copy a fresh `__Host-console_session` value while signed in; the old `auth` cookie is rejected. A `400` / invalid response usually means the workspace ID is wrong. +- A 429 means the API is rate limiting; onWatch skips that poll and tries again on the next one. ### Parse failed / response format changed -In usage API mode, OpenCode likely changed the `go/status` response. In scrape mode, it likely changed the dashboard markup. File an issue with: +OpenCode likely changed the `go/status` response. File an issue with: - Approximate time of failure -- Whether the browser dashboard still shows 5h / weekly / monthly -- **Do not** attach keys, cookies, full HTML dumps or raw `go/status` responses (they contain account IDs) +- Whether the browser console still shows 5h / weekly / monthly usage +- **Do not** attach keys, cookies or raw `go/status` responses (they contain account IDs) ### Docker / headless @@ -215,9 +130,9 @@ Pass the env vars into the container. There is no local credential auto-detectio ```bash OPENCODE_GO_API_KEY=oc_sk_... -# or, scrape mode: +# or, browser session: OPENCODE_GO_WORKSPACE_ID=wrk_xxxxxxxx -OPENCODE_GO_AUTH_COOKIE=your_auth_cookie_value +OPENCODE_GO_AUTH_COOKIE=your_console_session_value ``` --- diff --git a/install.sh b/install.sh index e180d425..c213a0aa 100755 --- a/install.sh +++ b/install.sh @@ -64,9 +64,12 @@ _jctl_cmd() { # ─── Input Helpers ────────────────────────────────────────────────── # Generate a random 12-char alphanumeric password +# Reads a bounded 512 bytes (~124 alphanumerics expected) rather than piping +# the endless /dev/urandom through tr: when SIGPIPE is ignored (CI runners, +# Node-spawned terminals) BSD tr never exits after head closes the pipe. generate_password() { local bytes - bytes=$(LC_ALL=C tr -dc 'A-Za-z0-9' /dev/null | head -c 12) || true + bytes=$(head -c 512 /dev/urandom 2>/dev/null | LC_ALL=C tr -dc 'A-Za-z0-9' 2>/dev/null | head -c 12) || true printf '%s' "$bytes" } diff --git a/internal/agent/anthropic_authrecovery_test.go b/internal/agent/anthropic_authrecovery_test.go index bf982ab8..1a449c3e 100644 --- a/internal/agent/anthropic_authrecovery_test.go +++ b/internal/agent/anthropic_authrecovery_test.go @@ -14,6 +14,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/api" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" "github.com/onllm-dev/onwatch/v2/internal/tracker" ) @@ -31,7 +32,7 @@ func newAuthRecoveryFixture(t *testing.T, oauthHandler http.HandlerFunc) *authRe t.Helper() // Isolate credential writes from the developer's real Claude Code session. - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) var apiCalls, oauthCalls atomic.Int32 apiServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -301,7 +302,7 @@ func TestAnthropicAgent_ProactiveRefresh_AppliesTokenWhenSaveFails(t *testing.T) // Corrupt credentials file on disk makes WriteAnthropicCredentials fail. home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) if err := os.MkdirAll(filepath.Join(home, ".claude"), 0o700); err != nil { t.Fatalf("mkdir: %v", err) } diff --git a/internal/agent/anthropic_statusline.go b/internal/agent/anthropic_statusline.go index cb65d876..03688d4a 100644 --- a/internal/agent/anthropic_statusline.go +++ b/internal/agent/anthropic_statusline.go @@ -5,7 +5,9 @@ import ( "fmt" "log/slog" "os" + "os/exec" "path/filepath" + "runtime" "sort" "strings" "sync" @@ -34,7 +36,55 @@ const statuslineFileName = "anthropic-statusline.json" // 1. Reads all of stdin into $I // 2. Saves $I to ~/.onwatch/data/anthropic-statusline.json (atomic via temp+mv) // 3. Pipes $I to stdout (so the next command in the pipe gets it) -const bridgeSnippet = `bash -c 'I=$(cat);D=$HOME/.onwatch/data;mkdir -p "$D" 2>/dev/null;T="$D/.sl-$$";printf "%s" "$I">"$T"&&mv -f "$T" "$D/anthropic-statusline.json" 2>/dev/null||rm -f "$T" 2>/dev/null;printf "%s" "$I"'` +// +// This is the exact text written on macOS and Linux. It must stay byte-for-byte +// stable: existing installs are recognised (and removed) by matching it. +const bridgeSnippet = bridgeSnippetHead + `$HOME/.onwatch/data` + bridgeSnippetTail + +// bridgeSnippetHead and bridgeSnippetTail surround the data directory in the +// bridge snippet. Splitting the snippet here lets Windows embed an absolute +// data directory while every variant stays recognisable for removal. +const ( + bridgeSnippetHead = `bash -c 'I=$(cat);D=` + bridgeSnippetTail = `;mkdir -p "$D" 2>/dev/null;T="$D/.sl-$$";printf "%s" "$I">"$T"&&mv -f "$T" "$D/anthropic-statusline.json" 2>/dev/null||rm -f "$T" 2>/dev/null;printf "%s" "$I"'` +) + +// bridgeStandaloneSuffix discards the snippet's stdout when the user has no +// statusline command of their own. +const bridgeStandaloneSuffix = " > /dev/null" + +// bridgeSnippetFor returns the bridge snippet for the given platform and +// onWatch data directory. +// +// On Windows, Claude Code runs statusline commands through Git Bash, where +// $HOME is not guaranteed to match the directory onWatch reads from: +// os.UserHomeDir uses %USERPROFILE%, while Git Bash derives HOME from an +// existing HOME variable or %HOMEDRIVE%%HOMEPATH% first. The snippet therefore +// embeds the absolute data directory, with forward slashes because Git Bash +// treats backslashes as escapes. +func bridgeSnippetFor(goos, dataDir string) string { + if goos != "windows" || dataDir == "" { + return bridgeSnippet + } + dir := strings.ReplaceAll(dataDir, `\`, "/") + // Double-quote for the inner bash, then escape for the outer single quotes. + quoted := `"` + bashDoubleQuoteEscaper.Replace(dir) + `"` + quoted = strings.ReplaceAll(quoted, "'", `'\''`) + return bridgeSnippetHead + quoted + bridgeSnippetTail +} + +// bashDoubleQuoteEscaper escapes the characters that stay special inside a +// bash double-quoted string. +var bashDoubleQuoteEscaper = strings.NewReplacer(`\`, `\\`, `"`, `\"`, "$", `\$`, "`", "\\`") + +// bridgeGOOS is the platform whose bridge snippet this process writes. It is a +// variable only so tests can exercise the Windows and Unix handling on any OS. +var bridgeGOOS = runtime.GOOS + +// currentBridgeSnippet returns the bridge snippet for this platform. +func currentBridgeSnippet() string { + return bridgeSnippetFor(bridgeGOOS, onwatchDataDir()) +} // bridgeMarker is a substring used to detect if the bridge snippet is already // present in the user's statusline command. @@ -333,27 +383,161 @@ func hasBridgeSnippet(command string) bool { // addBridgeSnippet prepends the save snippet to the user's command via a pipe. // If the user has no command, returns just the save snippet (no pipe). func addBridgeSnippet(userCommand string) string { + snippet := currentBridgeSnippet() if userCommand == "" { // No user command - standalone: save data, no display output - return bridgeSnippet + " > /dev/null" + return snippet + bridgeStandaloneSuffix } // Prepend: save stdin to file, then pipe original stdin to user's command - return bridgeSnippet + " | " + userCommand + return snippet + " | " + userCommand } // removeBridgeSnippet strips our snippet from the command, returning the // user's original command. Returns empty string if nothing remains. func removeBridgeSnippet(command string) string { - // Remove "snippet | user-cmd" → "user-cmd" - if idx := strings.Index(command, bridgeSnippet+" | "); idx == 0 { - return strings.TrimSpace(command[len(bridgeSnippet+" | "):]) + userCmd, _ := stripBridgeSnippet(command) + return userCmd +} + +// stripBridgeSnippet removes a bridge snippet written by any onWatch version or +// platform variant. ok is false when the command does not start with a bridge +// snippet, in which case command is returned unchanged. +func stripBridgeSnippet(command string) (userCmd string, ok bool) { + if !strings.HasPrefix(command, bridgeSnippetHead) { + return command, false + } + end := strings.Index(command, bridgeSnippetTail) + if end < 0 { + return command, false + } + rest := command[end+len(bridgeSnippetTail):] + switch { + case strings.HasPrefix(rest, " | "): + // "snippet | user-cmd" -> "user-cmd" + return strings.TrimSpace(rest[len(" | "):]), true + case rest == bridgeStandaloneSuffix: + // "snippet > /dev/null" -> "" (standalone mode) + return "", true + } + return command, false +} + +// bridgedCommand returns the statusline command with the current bridge +// snippet in front, and whether it differs from currentCmd. A bridge that is +// outdated for this platform (written for another data directory, or the +// $HOME form older Windows builds wrote) is replaced in place. A command that +// mentions the bridge file but was not written by onWatch is left alone. +// +// A recognised bridge from the other platform family is also left alone: a +// settings.json synced between Windows and macOS/Linux would otherwise be +// rewritten by each machine in turn, forever. Only Windows writes the quoted +// absolute-path form, so macOS/Linux never rewrite it; Windows replaces the +// $HOME form, which older Windows builds wrote, after which neither side +// changes the synced command again. +func bridgedCommand(currentCmd string) (string, bool) { + if !hasBridgeSnippet(currentCmd) { + return addBridgeSnippet(currentCmd), true } - // Remove "snippet > /dev/null" → "" (standalone mode) - if command == bridgeSnippet+" > /dev/null" { - return "" + userCmd, ok := stripBridgeSnippet(currentCmd) + if !ok { + return currentCmd, false + } + if bridgeGOOS != "windows" && isWindowsBridgeSnippet(currentCmd) { + return currentCmd, false + } + newCmd := addBridgeSnippet(userCmd) + return newCmd, newCmd != currentCmd +} + +// isWindowsBridgeSnippet reports whether command starts with the Windows form +// of the bridge snippet, which embeds a quoted absolute data directory where +// the macOS/Linux form has $HOME/.onwatch/data. +func isWindowsBridgeSnippet(command string) bool { + if !strings.HasPrefix(command, bridgeSnippetHead) { + return false + } + return strings.HasPrefix(command[len(bridgeSnippetHead):], `"`) +} + +// removeUnrunnableBridge strips an existing bridge snippet from Claude Code's +// settings when no shell that can run it is available (Windows without Git +// Bash, where Claude Code runs the statusline in PowerShell and the snippet +// takes the user's own statusline down with it). The user's command is +// restored; a standalone bridge leaves no statusline. Settings without a +// bridge are not touched. +func removeUnrunnableBridge(logger *slog.Logger) { + if logger == nil { + logger = slog.Default() + } + settings, err := readClaudeSettings() + if err != nil { + return } - // Not our command - return command + userCmd, ok := stripBridgeSnippet(getCurrentStatusLineCommand(settings)) + if !ok { + return + } + if userCmd == "" { + delete(settings, "statusLine") + } else { + setStatusLineCommand(settings, userCmd) + } + if err := writeClaudeSettings(settings); err != nil { + logger.Warn("Failed to remove statusline bridge that PowerShell cannot run", "error", err) + return + } + logger.Info("Removed statusline bridge from Claude Code settings: Git Bash not found, so Claude Code runs the statusline in PowerShell, which cannot run it") +} + +// bridgeShellAvailable reports whether Claude Code will run the statusline +// command in a shell that understands the bash snippet. On Windows, Claude +// Code uses Git Bash when it is installed and PowerShell otherwise; under +// PowerShell the snippet fails and takes the user's own statusline down with +// it, so the bridge is only configured when Git Bash is present. +var bridgeShellAvailable = func() bool { + if runtime.GOOS != "windows" { + return true + } + return findGitBash(os.Getenv, exec.LookPath, isRegularFile) != "" +} + +// findGitBash locates Git for Windows' bash.exe the way a Windows user would +// have it installed: an explicit CLAUDE_CODE_GIT_BASH_PATH, next to git.exe on +// PATH, or a standard install location. Returns "" if none is found. +func findGitBash(getenv func(string) string, lookPath func(string) (string, error), exists func(string) bool) string { + if p := strings.TrimSpace(getenv("CLAUDE_CODE_GIT_BASH_PATH")); p != "" && exists(p) { + return p + } + // git.exe lives in \cmd, \bin or \mingw64\bin. + if gitPath, err := lookPath("git"); err == nil && gitPath != "" { + dir := filepath.Dir(gitPath) + for _, up := range []string{"..", filepath.Join("..", "..")} { + if c := filepath.Join(dir, up, "bin", "bash.exe"); exists(c) { + return c + } + } + } + var candidates []string + for _, env := range []string{"ProgramFiles", "ProgramW6432", "ProgramFiles(x86)"} { + if root := getenv(env); root != "" { + candidates = append(candidates, filepath.Join(root, "Git", "bin", "bash.exe")) + } + } + if root := getenv("LOCALAPPDATA"); root != "" { + candidates = append(candidates, filepath.Join(root, "Programs", "Git", "bin", "bash.exe")) + } + for _, c := range candidates { + if exists(c) { + return c + } + } + return "" +} + +// isRegularFile reports whether path exists and is not a directory. +func isRegularFile(path string) bool { + info, err := os.Stat(path) + return err == nil && !info.IsDir() } // readClaudeSettings reads and parses ~/.claude/settings.json. @@ -466,6 +650,12 @@ func SetupStatuslineBridge(logger *slog.Logger) error { return nil } + if !bridgeShellAvailable() { + removeUnrunnableBridge(logger) + logger.Info("Claude Code statusline bridge disabled: Git Bash not found; statusline data unavailable on this machine (API polling still runs unless ANTHROPIC_SOURCE=statusline)") + return nil + } + // Ensure data directory exists for the statusline file dataDir := onwatchDataDir() if dataDir != "" { @@ -480,20 +670,22 @@ func SetupStatuslineBridge(logger *slog.Logger) error { currentCmd := getCurrentStatusLineCommand(settings) - if hasBridgeSnippet(currentCmd) { + // Prepend our snippet to whatever the user has (or standalone if empty). + // An outdated snippet is replaced in place. + newCmd, changed := bridgedCommand(currentCmd) + if !changed { logger.Debug("Statusline bridge already configured") return nil } - - // Prepend our snippet to whatever the user has (or standalone if empty) - newCmd := addBridgeSnippet(currentCmd) setStatusLineCommand(settings, newCmd) if err := writeClaudeSettings(settings); err != nil { logger.Warn("Failed to configure statusline bridge", "error", err) return nil } - if currentCmd == "" { + if hasBridgeSnippet(currentCmd) { + logger.Info("Updated statusline bridge") + } else if currentCmd == "" { logger.Info("Configured statusline bridge (standalone)") } else { logger.Info("Configured statusline bridge (prepended to existing command)") @@ -515,6 +707,10 @@ func EnsureStatuslineBridge(logger *slog.Logger) { if !isClaudeCodeInstalled() || isBridgeDisabled() { return } + if !bridgeShellAvailable() { + removeUnrunnableBridge(logger) + return + } settings, err := readClaudeSettings() if err != nil { @@ -522,15 +718,17 @@ func EnsureStatuslineBridge(logger *slog.Logger) { } currentCmd := getCurrentStatusLineCommand(settings) - if hasBridgeSnippet(currentCmd) { + newCmd, changed := bridgedCommand(currentCmd) + if !changed { return // Still healthy } - // Bridge was removed (user changed their statusline) - re-prepend - newCmd := addBridgeSnippet(currentCmd) + // Bridge was removed (user changed their statusline) or is outdated - re-prepend setStatusLineCommand(settings, newCmd) if err := writeClaudeSettings(settings); err == nil { - if currentCmd == "" { + if hasBridgeSnippet(currentCmd) { + logger.Info("Statusline bridge updated") + } else if currentCmd == "" { logger.Info("Statusline bridge re-established (standalone)") } else { logger.Info("Statusline bridge re-prepended to user command") diff --git a/internal/agent/anthropic_statusline_platform_test.go b/internal/agent/anthropic_statusline_platform_test.go new file mode 100644 index 00000000..fa14cf39 --- /dev/null +++ b/internal/agent/anthropic_statusline_platform_test.go @@ -0,0 +1,474 @@ +package agent + +import ( + "context" + "encoding/json" + "errors" + "log/slog" + "os" + "os/exec" + "path/filepath" + "runtime" + "strings" + "testing" + "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" +) + +// stubBridgeShell overrides the Git Bash check so bridge tests behave the same +// on every OS, regardless of what the machine running them has installed. +func stubBridgeShell(t *testing.T, available bool) { + t.Helper() + prev := bridgeShellAvailable + bridgeShellAvailable = func() bool { return available } + t.Cleanup(func() { bridgeShellAvailable = prev }) +} + +// stubBridgeGOOS makes the bridge code behave as it does on goos, so the +// Windows and Unix settings.json handling can be tested on any OS. +func stubBridgeGOOS(t *testing.T, goos string) { + t.Helper() + prev := bridgeGOOS + bridgeGOOS = goos + t.Cleanup(func() { bridgeGOOS = prev }) +} + +// The macOS/Linux snippet is how existing installs are recognised, so its text +// must never drift. +func TestBridgeSnippet_UnixTextUnchanged(t *testing.T) { + const want = `bash -c 'I=$(cat);D=$HOME/.onwatch/data;mkdir -p "$D" 2>/dev/null;T="$D/.sl-$$";printf "%s" "$I">"$T"&&mv -f "$T" "$D/anthropic-statusline.json" 2>/dev/null||rm -f "$T" 2>/dev/null;printf "%s" "$I"'` + if bridgeSnippet != want { + t.Fatalf("bridgeSnippet changed:\n got %s\nwant %s", bridgeSnippet, want) + } + for _, goos := range []string{"darwin", "linux", "freebsd"} { + if got := bridgeSnippetFor(goos, "/home/u/.onwatch/data"); got != bridgeSnippet { + t.Errorf("bridgeSnippetFor(%s) = %s, want the $HOME snippet", goos, got) + } + } +} + +func TestBridgeSnippetFor_WindowsEmbedsForwardSlashDataDir(t *testing.T) { + got := bridgeSnippetFor("windows", `C:\Users\O'Brien $x\.onwatch\data`) + want := `D="C:/Users/O'\''Brien \$x/.onwatch/data";` + if !strings.Contains(got, want) { + t.Fatalf("windows snippet = %s\nwant it to contain %s", got, want) + } + if strings.Contains(got, `\.onwatch`) || strings.Contains(got, "$HOME") { + t.Fatalf("windows snippet must not use backslashes or $HOME: %s", got) + } + if !hasBridgeSnippet(got) { + t.Fatal("windows snippet must carry the bridge marker") + } +} + +func TestStripBridgeSnippet_AllVariants(t *testing.T) { + variants := map[string]string{ + "unix": bridgeSnippet, + "windows": bridgeSnippetFor("windows", `C:\Users\me\.onwatch\data`), + } + for name, snippet := range variants { + if got, ok := stripBridgeSnippet(snippet + " | ~/.claude/statusline.sh"); !ok || got != "~/.claude/statusline.sh" { + t.Errorf("%s piped: got (%q, %v)", name, got, ok) + } + if got, ok := stripBridgeSnippet(snippet + bridgeStandaloneSuffix); !ok || got != "" { + t.Errorf("%s standalone: got (%q, %v)", name, got, ok) + } + if got, ok := stripBridgeSnippet(snippet); ok || got != snippet { + t.Errorf("%s bare: got (%q, %v), want unchanged", name, got, ok) + } + } + foreign := "cat ~/.onwatch/data/anthropic-statusline.json" + if got, ok := stripBridgeSnippet(foreign); ok || got != foreign { + t.Errorf("foreign command: got (%q, %v), want unchanged", got, ok) + } +} + +func TestBridgedCommand(t *testing.T) { + for _, goos := range []string{"darwin", "linux", "windows"} { + t.Run(goos, func(t *testing.T) { + stubBridgeGOOS(t, goos) + current := addBridgeSnippet("~/sl.sh") + if got, changed := bridgedCommand(current); changed || got != current { + t.Errorf("current bridge: got (%q, %v), want unchanged", got, changed) + } + + foreign := "cat ~/.onwatch/data/anthropic-statusline.json" + if got, changed := bridgedCommand(foreign); changed || got != foreign { + t.Errorf("foreign command: got (%q, %v), want unchanged", got, changed) + } + + if got, changed := bridgedCommand(""); !changed || got != addBridgeSnippet("") { + t.Errorf("empty: got (%q, %v)", got, changed) + } + if got, changed := bridgedCommand("~/sl.sh"); !changed || got != current { + t.Errorf("no bridge: got (%q, %v), want %q", got, changed, current) + } + }) + } +} + +// On Windows, a bridge written by an older Windows build (the $HOME form) or +// for another data directory is outdated for this platform and replaced in +// place. +func TestBridgedCommand_WindowsReplacesStaleWindowsVariants(t *testing.T) { + stubBridgeGOOS(t, "windows") + current := addBridgeSnippet("~/sl.sh") + stale := map[string]string{ + "$HOME form": bridgeSnippet + " | ~/sl.sh", + "other data dir": bridgeSnippetFor("windows", filepath.Join(t.TempDir(), "elsewhere")) + " | ~/sl.sh", + } + for name, cmd := range stale { + if got, changed := bridgedCommand(cmd); !changed || got != current { + t.Errorf("%s: got (%q, %v), want %q", name, got, changed, current) + } + } +} + +// A settings.json synced between a Windows machine and a macOS/Linux one +// carries the Windows bridge. macOS/Linux must leave that recognised bridge +// alone rather than rewrite it, or each machine would rewrite the other's +// bridge forever. +func TestBridgedCommand_UnixLeavesWindowsVariant(t *testing.T) { + for _, goos := range []string{"darwin", "linux"} { + t.Run(goos, func(t *testing.T) { + stubBridgeGOOS(t, goos) + for _, cmd := range []string{ + bridgeSnippetFor("windows", `C:\Users\me\.onwatch\data`) + " | ~/sl.sh", + bridgeSnippetFor("windows", `C:\Users\me\.onwatch\data`) + bridgeStandaloneSuffix, + } { + if got, changed := bridgedCommand(cmd); changed || got != cmd { + t.Errorf("windows bridge on %s: got (%q, %v), want unchanged", goos, got, changed) + } + } + }) + } +} + +// Once Windows has written its bridge, neither side rewrites it again. +func TestBridgedCommand_SyncedSettingsSettle(t *testing.T) { + home := t.TempDir() + testhome.SetTestHome(t, home) + cmd := bridgeSnippet + " | ~/sl.sh" // written on macOS + + stubBridgeGOOS(t, "windows") + cmd, _ = bridgedCommand(cmd) + for i := 0; i < 3; i++ { + for _, goos := range []string{"darwin", "windows"} { + bridgeGOOS = goos + if got, changed := bridgedCommand(cmd); changed { + t.Fatalf("round %d on %s rewrote the bridge:\n from %s\n to %s", i, goos, cmd, got) + } + } + } +} + +// A bridge written for another data directory (for example the $HOME form an +// older Windows build wrote) is replaced, not stacked. +func TestSetupStatuslineBridge_ReplacesStaleSnippet(t *testing.T) { + home := t.TempDir() + testhome.SetTestHome(t, home) + stubBridgeShell(t, true) + stubBridgeGOOS(t, "windows") + claudeDir := filepath.Join(home, ".claude") + if err := os.MkdirAll(claudeDir, 0o700); err != nil { + t.Fatalf("mkdir: %v", err) + } + stale := bridgeSnippetFor("windows", filepath.Join(t.TempDir(), "elsewhere")) + " | ~/sl.sh" + writeStatusLineSettings(t, claudeDir, stale) + + if err := SetupStatuslineBridge(slog.Default()); err != nil { + t.Fatalf("SetupStatuslineBridge: %v", err) + } + + cmd := readStatusLineCommand(t, claudeDir) + if cmd != addBridgeSnippet("~/sl.sh") { + t.Fatalf("command = %s\nwant %s", cmd, addBridgeSnippet("~/sl.sh")) + } + if n := strings.Count(cmd, bridgeMarker); n != 1 { + t.Fatalf("bridge marker appears %d times, want 1", n) + } +} + +// Setup and the health check on macOS/Linux leave a synced Windows bridge in +// settings.json untouched. +func TestStatuslineBridge_UnixLeavesSyncedWindowsBridge(t *testing.T) { + home := t.TempDir() + testhome.SetTestHome(t, home) + stubBridgeShell(t, true) + stubBridgeGOOS(t, "darwin") + claudeDir := filepath.Join(home, ".claude") + if err := os.MkdirAll(claudeDir, 0o700); err != nil { + t.Fatalf("mkdir: %v", err) + } + windowsCmd := bridgeSnippetFor("windows", `C:\Users\me\.onwatch\data`) + " | ~/sl.sh" + writeStatusLineSettings(t, claudeDir, windowsCmd) + + if err := SetupStatuslineBridge(slog.Default()); err != nil { + t.Fatalf("SetupStatuslineBridge: %v", err) + } + resetBridgeCheck() + EnsureStatuslineBridge(slog.Default()) + + if cmd := readStatusLineCommand(t, claudeDir); cmd != windowsCmd { + t.Fatalf("command = %s\nwant untouched %s", cmd, windowsCmd) + } +} + +// Without Git Bash, Claude Code on Windows runs the statusline in PowerShell, +// where the bash snippet would break the user's statusline. Setup and the +// health check must leave a bridge-free settings.json alone. +func TestStatuslineBridge_NoBashShellLeavesSettingsAlone(t *testing.T) { + home := t.TempDir() + testhome.SetTestHome(t, home) + stubBridgeShell(t, false) + stubBridgeGOOS(t, "windows") + claudeDir := filepath.Join(home, ".claude") + if err := os.MkdirAll(claudeDir, 0o700); err != nil { + t.Fatalf("mkdir: %v", err) + } + userCmd := "powershell -NoProfile -File C:/Users/me/.claude/statusline.ps1" + writeStatusLineSettings(t, claudeDir, userCmd) + + if err := SetupStatuslineBridge(slog.Default()); err != nil { + t.Fatalf("SetupStatuslineBridge: %v", err) + } + resetBridgeCheck() + EnsureStatuslineBridge(slog.Default()) + + if cmd := readStatusLineCommand(t, claudeDir); cmd != userCmd { + t.Fatalf("command = %s, want untouched %s", cmd, userCmd) + } +} + +// A bridge that an older onWatch wrote before it checked for Git Bash keeps +// breaking the statusline under PowerShell. Without Git Bash, Setup and the +// health check remove it and restore the user's own command. +func TestStatuslineBridge_NoBashShellRemovesExistingBridge(t *testing.T) { + userCmd := "powershell -NoProfile -File C:/Users/me/.claude/statusline.ps1" + variants := map[string]string{ + "older $HOME form piped": bridgeSnippet + " | " + userCmd, + "windows form piped": bridgeSnippetFor("windows", `C:\Users\me\.onwatch\data`) + " | " + userCmd, + "older $HOME form standalone": bridgeSnippet + bridgeStandaloneSuffix, + } + for name, bridged := range variants { + want := userCmd + if strings.HasSuffix(bridged, bridgeStandaloneSuffix) { + want = "" + } + for _, entry := range []string{"setup", "ensure"} { + t.Run(name+"/"+entry, func(t *testing.T) { + home := t.TempDir() + testhome.SetTestHome(t, home) + stubBridgeShell(t, false) + stubBridgeGOOS(t, "windows") + claudeDir := filepath.Join(home, ".claude") + if err := os.MkdirAll(claudeDir, 0o700); err != nil { + t.Fatalf("mkdir: %v", err) + } + writeStatusLineSettings(t, claudeDir, bridged) + + if entry == "setup" { + if err := SetupStatuslineBridge(slog.Default()); err != nil { + t.Fatalf("SetupStatuslineBridge: %v", err) + } + } else { + resetBridgeCheck() + EnsureStatuslineBridge(slog.Default()) + } + + if cmd := readStatusLineCommand(t, claudeDir); cmd != want { + t.Fatalf("command = %q, want %q", cmd, want) + } + }) + } + } +} + +// resetBridgeCheck lets the next EnsureStatuslineBridge call run its check. +func resetBridgeCheck() { + bridgeSetup.mu.Lock() + bridgeSetup.lastCheck = time.Time{} + bridgeSetup.mu.Unlock() +} + +func TestFindGitBash(t *testing.T) { + root := filepath.Join("X", "Git") + bash := filepath.Join(root, "bin", "bash.exe") + + tests := []struct { + name string + env map[string]string + git string + files []string + want string + }{ + { + name: "explicit env var", + env: map[string]string{"CLAUDE_CODE_GIT_BASH_PATH": filepath.Join("Y", "bash.exe")}, + files: []string{filepath.Join("Y", "bash.exe")}, + want: filepath.Join("Y", "bash.exe"), + }, + { + name: "env var pointing nowhere falls through", + env: map[string]string{"CLAUDE_CODE_GIT_BASH_PATH": filepath.Join("Y", "bash.exe"), "ProgramFiles": "X"}, + files: []string{bash}, + want: bash, + }, + { + name: "git in cmd dir", + git: filepath.Join(root, "cmd", "git.exe"), + files: []string{bash}, + want: bash, + }, + { + name: "git in mingw64 bin", + git: filepath.Join(root, "mingw64", "bin", "git.exe"), + files: []string{bash}, + want: bash, + }, + { + name: "program files", + env: map[string]string{"ProgramFiles": "X"}, + files: []string{bash}, + want: bash, + }, + { + name: "per-user install", + env: map[string]string{"LOCALAPPDATA": "L"}, + files: []string{filepath.Join("L", "Programs", "Git", "bin", "bash.exe")}, + want: filepath.Join("L", "Programs", "Git", "bin", "bash.exe"), + }, + { + name: "not installed", + env: map[string]string{"ProgramFiles": "X"}, + git: filepath.Join("Z", "shims", "git.exe"), + want: "", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + getenv := func(k string) string { return tt.env[k] } + lookPath := func(string) (string, error) { + if tt.git == "" { + return "", errors.New("not found") + } + return tt.git, nil + } + exists := func(p string) bool { + for _, f := range tt.files { + if filepath.Clean(p) == filepath.Clean(f) { + return true + } + } + return false + } + if got := findGitBash(getenv, lookPath, exists); filepath.Clean(got) != filepath.Clean(tt.want) { + t.Fatalf("findGitBash() = %q, want %q", got, tt.want) + } + }) + } +} + +// TestBridgeSnippet_RunsUnderBash executes the generated statusline command in +// a real bash, the way Claude Code does, and checks that stdin is passed +// through untouched and saved exactly where onWatch reads it. The data +// directory contains a space, a single quote and a dollar sign to exercise the +// quoting of the Windows variant. +func TestBridgeSnippet_RunsUnderBash(t *testing.T) { + bash := testBashPath(t) + payload := `{"rate_limits":{"five_hour":{"used_percentage":12.5,"resets_at":1790000000}}}` + + run := func(t *testing.T, command string, env []string) string { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second) + defer cancel() + cmd := exec.CommandContext(ctx, bash, "-c", command) + cmd.Stdin = strings.NewReader(payload) + cmd.Env = append(os.Environ(), env...) + out, err := cmd.Output() + if err != nil { + t.Fatalf("run statusline command: %v", err) + } + return string(out) + } + + t.Run("windows variant", func(t *testing.T) { + dataDir := filepath.Join(t.TempDir(), "it's $odd", "data") + command := bridgeSnippetFor("windows", dataDir) + " | cat" + if out := run(t, command, nil); out != payload { + t.Fatalf("stdout = %q, want payload passed through", out) + } + assertStatuslineFile(t, filepath.Join(dataDir, statuslineFileName), payload) + + standalone := bridgeSnippetFor("windows", dataDir) + bridgeStandaloneSuffix + if out := run(t, standalone, nil); out != "" { + t.Fatalf("standalone stdout = %q, want empty", out) + } + }) + + t.Run("unix variant", func(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("the $HOME snippet is only written on macOS and Linux") + } + home := t.TempDir() + if out := run(t, bridgeSnippet+" | cat", []string{"HOME=" + home}); out != payload { + t.Fatalf("stdout = %q, want payload passed through", out) + } + assertStatuslineFile(t, filepath.Join(home, ".onwatch", "data", statuslineFileName), payload) + }) +} + +// testBashPath returns the bash Claude Code would use: Git Bash on Windows, +// bash from PATH elsewhere. +func testBashPath(t *testing.T) string { + t.Helper() + if runtime.GOOS == "windows" { + if p := findGitBash(os.Getenv, exec.LookPath, isRegularFile); p != "" { + return p + } + t.Skip("Git Bash not installed; the bridge is not configured without it") + } + p, err := exec.LookPath("bash") + if err != nil { + t.Skip("bash not installed") + } + return p +} + +func assertStatuslineFile(t *testing.T, path, want string) { + t.Helper() + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("statusline file not written where onWatch reads it: %v", err) + } + if string(data) != want { + t.Fatalf("statusline file = %q, want %q", data, want) + } +} + +func writeStatusLineSettings(t *testing.T, claudeDir, command string) { + t.Helper() + data, err := json.MarshalIndent(map[string]interface{}{ + "statusLine": map[string]interface{}{"type": "command", "command": command}, + }, "", " ") + if err != nil { + t.Fatalf("marshal settings: %v", err) + } + if err := os.WriteFile(filepath.Join(claudeDir, "settings.json"), data, 0o600); err != nil { + t.Fatalf("write settings: %v", err) + } +} + +func readStatusLineCommand(t *testing.T, claudeDir string) string { + t.Helper() + data, err := os.ReadFile(filepath.Join(claudeDir, "settings.json")) + if err != nil { + t.Fatalf("read settings: %v", err) + } + var settings map[string]interface{} + if err := json.Unmarshal(data, &settings); err != nil { + t.Fatalf("parse settings: %v", err) + } + return getCurrentStatusLineCommand(settings) +} diff --git a/internal/agent/anthropic_statusline_test.go b/internal/agent/anthropic_statusline_test.go index acdcf289..e51e001e 100644 --- a/internal/agent/anthropic_statusline_test.go +++ b/internal/agent/anthropic_statusline_test.go @@ -8,6 +8,8 @@ import ( "strings" "testing" "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) // --- readStatuslineData tests --- @@ -354,7 +356,7 @@ func TestRemoveBridgeSnippet_NotOurCommand(t *testing.T) { func TestSetupStatuslineBridge_CCNotInstalled(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) logger := slog.Default() if err := SetupStatuslineBridge(logger); err != nil { t.Fatalf("expected no error when CC not installed: %v", err) @@ -367,7 +369,8 @@ func TestSetupStatuslineBridge_CCNotInstalled(t *testing.T) { func TestSetupStatuslineBridge_NoExistingStatusline(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) + stubBridgeShell(t, true) if err := os.MkdirAll(filepath.Join(home, ".claude"), 0o700); err != nil { t.Fatalf("mkdir: %v", err) } @@ -394,7 +397,8 @@ func TestSetupStatuslineBridge_NoExistingStatusline(t *testing.T) { func TestSetupStatuslineBridge_PrependsToExistingCommand(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) + stubBridgeShell(t, true) claudeDir := filepath.Join(home, ".claude") if err := os.MkdirAll(claudeDir, 0o700); err != nil { t.Fatalf("mkdir: %v", err) @@ -437,7 +441,8 @@ func TestSetupStatuslineBridge_PrependsToExistingCommand(t *testing.T) { func TestSetupStatuslineBridge_Idempotent(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) + stubBridgeShell(t, true) if err := os.MkdirAll(filepath.Join(home, ".claude"), 0o700); err != nil { t.Fatalf("mkdir: %v", err) } @@ -458,7 +463,7 @@ func TestSetupStatuslineBridge_Idempotent(t *testing.T) { func TestSetupStatuslineBridge_Disabled(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) if err := os.MkdirAll(filepath.Join(home, ".claude"), 0o700); err != nil { t.Fatalf("mkdir: %v", err) } @@ -478,7 +483,7 @@ func TestSetupStatuslineBridge_Disabled(t *testing.T) { func TestSetupStatuslineBridge_MalformedSettings(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) claudeDir := filepath.Join(home, ".claude") os.MkdirAll(claudeDir, 0o700) os.WriteFile(filepath.Join(claudeDir, "settings.json"), []byte("{bad json"), 0o600) @@ -493,7 +498,8 @@ func TestSetupStatuslineBridge_MalformedSettings(t *testing.T) { func TestDisableStatuslineBridge_RestoresOriginalCommand(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) + stubBridgeShell(t, true) claudeDir := filepath.Join(home, ".claude") os.MkdirAll(claudeDir, 0o700) @@ -522,7 +528,8 @@ func TestDisableStatuslineBridge_RestoresOriginalCommand(t *testing.T) { func TestDisableStatuslineBridge_RemovesStandalone(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) + stubBridgeShell(t, true) claudeDir := filepath.Join(home, ".claude") os.MkdirAll(claudeDir, 0o700) @@ -542,7 +549,8 @@ func TestDisableStatuslineBridge_RemovesStandalone(t *testing.T) { func TestEnsureStatuslineBridge_DetectsUserChange(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) + stubBridgeShell(t, true) claudeDir := filepath.Join(home, ".claude") os.MkdirAll(claudeDir, 0o700) diff --git a/internal/agent/codex_agent_manager_test.go b/internal/agent/codex_agent_manager_test.go index 038b4e7e..0d34fc3d 100644 --- a/internal/agent/codex_agent_manager_test.go +++ b/internal/agent/codex_agent_manager_test.go @@ -15,6 +15,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/notify" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" "github.com/onllm-dev/onwatch/v2/internal/tracker" ) @@ -29,7 +30,7 @@ func newCodexManagerFixture(t *testing.T) *codexManagerFixture { t.Helper() home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") // Pin OpenCode detection under the temp HOME so DetectCodexCredentials never // reads the host's real ~/.local/share/opencode/auth.json (issue #78 path). @@ -120,7 +121,7 @@ func makeCodexIDToken(t *testing.T, exp time.Time, accountID, userID string) str func TestNewCodexAgentManager_Defaults(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) manager := NewCodexAgentManager(nil, nil, 15*time.Second, nil) if manager.logger == nil { @@ -269,7 +270,7 @@ func TestCodexAgentManager_StartAgentForProfile_WiresNotifierChecksAndRefresh(t func TestCodexAgentManager_StartDefaultAgent(t *testing.T) { fx := newCodexManagerFixture(t) - authDir := filepath.Join(os.Getenv("HOME"), ".codex") + authDir := filepath.Join(testHomeDir(t), ".codex") if err := os.MkdirAll(authDir, 0o700); err != nil { t.Fatalf("mkdir .codex: %v", err) } @@ -365,7 +366,7 @@ func TestCodexAgentManager_Run_LoadsProfilesAndStopsOnCancel(t *testing.T) { func TestCodexAgentManager_Run_UsesDefaultCredentialsWhenNoProfiles(t *testing.T) { fx := newCodexManagerFixture(t) - authDir := filepath.Join(os.Getenv("HOME"), ".codex") + authDir := filepath.Join(testHomeDir(t), ".codex") if err := os.MkdirAll(authDir, 0o700); err != nil { t.Fatalf("mkdir .codex: %v", err) } @@ -506,7 +507,7 @@ func TestCodexAgentManager_StartAgentForProfile_UsesAuthJSONWhenProfileTokenStal profile.Tokens.IDToken = staleToken profilePath := fx.writeProfile(t, profile) - authDir := filepath.Join(os.Getenv("HOME"), ".codex") + authDir := filepath.Join(testHomeDir(t), ".codex") if err := os.MkdirAll(authDir, 0o700); err != nil { t.Fatalf("mkdir .codex: %v", err) } @@ -560,7 +561,7 @@ func TestCodexAgentManager_StartAgentForProfile_TokenSaveScopedToProfile(t *test profilePath := fx.writeProfile(t, profile) // Write something to global auth.json so we can verify it's NOT modified - authDir := filepath.Join(os.Getenv("HOME"), ".codex") + authDir := filepath.Join(testHomeDir(t), ".codex") if err := os.MkdirAll(authDir, 0o700); err != nil { t.Fatalf("mkdir .codex: %v", err) } @@ -653,7 +654,7 @@ func TestCodexAgentManager_LoadAndStartProfiles_TeamUsersGetDistinctAccounts(t * func TestCodexAgentManager_StartDefaultAgent_UsesCompositeExternalID(t *testing.T) { fx := newCodexManagerFixture(t) - authDir := filepath.Join(os.Getenv("HOME"), ".codex") + authDir := filepath.Join(testHomeDir(t), ".codex") if err := os.MkdirAll(authDir, 0o700); err != nil { t.Fatalf("mkdir .codex: %v", err) } @@ -699,7 +700,7 @@ func TestCodexAgentManager_TeamProfileRejectsSystemCredsFromDifferentUser(t *tes profile.Tokens.IDToken = tokenA fx.writeProfile(t, profile) - authDir := filepath.Join(os.Getenv("HOME"), ".codex") + authDir := filepath.Join(testHomeDir(t), ".codex") if err := os.MkdirAll(authDir, 0o700); err != nil { t.Fatalf("mkdir .codex: %v", err) } diff --git a/internal/agent/coverage_final_test.go b/internal/agent/coverage_final_test.go index 8258903b..e7daa4cb 100644 --- a/internal/agent/coverage_final_test.go +++ b/internal/agent/coverage_final_test.go @@ -14,6 +14,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/api" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" "github.com/onllm-dev/onwatch/v2/internal/tracker" ) @@ -371,7 +372,7 @@ func TestAnthropicAgent_PollAuthPauseAndResume(t *testing.T) { } func TestAnthropicAgent_PollRateLimitBypassWithOAuthRefresh(t *testing.T) { - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) oauthServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") @@ -430,7 +431,7 @@ func TestAnthropicAgent_PollRateLimitBypassWithOAuthRefresh(t *testing.T) { } func TestAnthropicAgent_PollProactiveOAuthRefresh(t *testing.T) { - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) oauthServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") diff --git a/internal/agent/deepseek_agent.go b/internal/agent/deepseek_agent.go index 497ef436..a82c453d 100644 --- a/internal/agent/deepseek_agent.go +++ b/internal/agent/deepseek_agent.go @@ -88,7 +88,7 @@ func (a *DeepSeekAgent) poll(ctx context.Context) { a.logger.Error("Failed to fetch DeepSeek balance", "error", err) return } - + if !resp.IsAvailable { a.logger.Info("DeepSeek service is currently not available") return diff --git a/internal/agent/gemini_agent_test.go b/internal/agent/gemini_agent_test.go index 0a03408f..63104ed4 100644 --- a/internal/agent/gemini_agent_test.go +++ b/internal/agent/gemini_agent_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "sync/atomic" "testing" "time" @@ -80,9 +81,10 @@ func TestGeminiAgent_Poll(t *testing.T) { func TestGeminiAgent_AuthFailurePause(t *testing.T) { t.Parallel() - callCount := 0 + // The handler can still be serving a request after Run returns. + var callCount atomic.Int32 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - callCount++ + callCount.Add(1) if r.URL.Path == "/v1internal:loadCodeAssist" { w.WriteHeader(http.StatusUnauthorized) return @@ -105,7 +107,7 @@ func TestGeminiAgent_AuthFailurePause(t *testing.T) { _ = agent.Run(ctx) // Should have attempted multiple polls - if callCount == 0 { + if callCount.Load() == 0 { t.Error("expected at least 1 API call") } } @@ -124,7 +126,7 @@ func TestGeminiAgent_TokenPersistenceOnRefresh(t *testing.T) { t.Parallel() refreshedAccessToken := "refreshed-access-token-xyz" originalRefreshToken := "original-refresh-token-abc" - quotaCallCount := 0 + var quotaCallCount atomic.Int32 // Mock server: first quota call returns 401, OAuth refresh succeeds, retry succeeds srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -133,8 +135,7 @@ func TestGeminiAgent_TokenPersistenceOnRefresh(t *testing.T) { case "/v1internal:loadCodeAssist": json.NewEncoder(w).Encode(api.GeminiTierResponse{Tier: "free"}) case "/v1internal:retrieveUserQuota": - quotaCallCount++ - if quotaCallCount == 1 { + if quotaCallCount.Add(1) == 1 { w.WriteHeader(http.StatusUnauthorized) return } diff --git a/internal/agent/opencode_agent.go b/internal/agent/opencode_agent.go index 71cee09d..e44e7d0c 100644 --- a/internal/agent/opencode_agent.go +++ b/internal/agent/opencode_agent.go @@ -136,9 +136,8 @@ func (a *OpenCodeAgent) poll(ctx context.Context) { ) } -// fetch prefers the Go status API (service-account key), which reads the -// plan's own meters, and falls back to scraping the Go dashboard with the -// workspace ID + auth cookie. +// fetch reads the Go status API with the service-account key when one is set, +// otherwise with the workspace ID + console session cookie. func (a *OpenCodeAgent) fetch(ctx context.Context) (*api.OpenCodeSnapshot, error) { if key := a.cfg.OpenCodeGoAPIKey; key != "" { return a.client.FetchUsageSnapshot(ctx, key) diff --git a/internal/agent/test_main_test.go b/internal/agent/test_main_test.go index 87c2709e..0dbb7381 100644 --- a/internal/agent/test_main_test.go +++ b/internal/agent/test_main_test.go @@ -1,22 +1,53 @@ package agent import ( + "fmt" "os" "testing" "github.com/onllm-dev/onwatch/v2/internal/api" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) // TestMain runs before all tests in the agent package. It enables test mode // on the api package to prevent any keychain/keyring operations during tests. // This ensures tests never read or write real Claude Code OAuth tokens. // -// It also unsets OPENCODE_HOME/XDG_DATA_HOME so codex credential detection -// never resolves to the host's real ~/.local/share/opencode/auth.json; tests -// that set a temp HOME stay fully isolated regardless of host env. +// It then points the home directory at a throwaway sandbox for the whole run +// and clears provider location overrides (CODEX_HOME, OPENCODE_HOME, +// XDG_DATA_HOME, ...), so codex/opencode credential detection never resolves +// to the host's real files. os.UserHomeDir reads HOME on Unix but USERPROFILE +// on Windows, so a test that only overrides HOME would otherwise read and +// write the real Windows profile - including ~/.claude/.credentials.json and +// ~/.claude/settings.json. Tests that need their own home call +// testhome.SetTestHome, which sets both. func TestMain(m *testing.M) { + os.Exit(runTests(m)) +} + +func runTests(m *testing.M) int { + // SetTestMode must run before the home is redirected: its first enable + // records the real home that the credential-file guard refuses. api.SetTestMode(true) - os.Unsetenv("OPENCODE_HOME") - os.Unsetenv("XDG_DATA_HOME") - os.Exit(m.Run()) + + _, cleanup, err := testhome.SandboxHome() + if err != nil { + fmt.Fprintf(os.Stderr, "agent tests: %v\n", err) + return 1 + } + defer cleanup() + + return m.Run() +} + +// testHomeDir returns the home directory the code under test will resolve, +// using the same lookup as production (os.UserHomeDir) so paths built from it +// are correct on every platform. +func testHomeDir(t *testing.T) string { + t.Helper() + home, err := os.UserHomeDir() + if err != nil { + t.Fatalf("UserHomeDir: %v", err) + } + return home } diff --git a/internal/api/anthropic_oauth.go b/internal/api/anthropic_oauth.go index 527d72ac..1d18da70 100644 --- a/internal/api/anthropic_oauth.go +++ b/internal/api/anthropic_oauth.go @@ -47,7 +47,7 @@ type oauthRateLimitedError struct { RetryAfter time.Duration } -func (e *oauthRateLimitedError) Error() string { return ErrOAuthRateLimited.Error() } +func (e *oauthRateLimitedError) Error() string { return ErrOAuthRateLimited.Error() } func (e *oauthRateLimitedError) Is(target error) bool { return errors.Is(target, ErrOAuthRateLimited) } // RetryAfter returns the Retry-After duration, or 0 if not available. diff --git a/internal/api/anthropic_token.go b/internal/api/anthropic_token.go index 375f9c5f..30e41527 100644 --- a/internal/api/anthropic_token.go +++ b/internal/api/anthropic_token.go @@ -2,10 +2,107 @@ package api import ( "encoding/json" + "errors" "log/slog" + "os" + "os/user" + "path/filepath" + "runtime" + "strings" + "sync" "time" ) +// testMode disables all keychain/keyring operations. Set to true in tests +// to prevent tests from reading or writing real Claude Code credentials. +// This is a critical safety guard - without it, tests can overwrite the user's +// real OAuth tokens in the macOS Keychain, logging them out of Claude Code. +// It also makes the credentials-file helpers refuse the real account home +// (see anthropicHomeBlocked), which matters most on Windows where the file is +// the only store. +var testMode bool + +// SetTestMode enables or disables test mode. When enabled, all keychain and +// keyring operations are skipped, and only file-based credential storage is used. +// Files are redirected by pointing the home directory (HOME, and USERPROFILE on +// Windows) at a temp dir in tests. +// +// The first enable records the account's real home directories, before a test +// harness redirects HOME/USERPROFILE, so later file access can be checked +// against them. +func SetTestMode(enabled bool) { + if enabled { + captureRealHomes() + } + testMode = enabled +} + +// IsTestMode reports whether SetTestMode(true) is in effect. Test harnesses in +// other packages use it to assert their TestMain enabled the guard. +func IsTestMode() bool { + return testMode +} + +// errRealCredentialsInTestMode is returned when test mode would touch the real +// account's ~/.claude/.credentials.json. +var errRealCredentialsInTestMode = errors.New("anthropic: test mode refuses the real Claude credentials file") + +var ( + realHomesMu sync.Mutex + realHomesCaptured bool + realHomes []string +) + +// captureRealHomes records the home directory as seen at first test-mode +// enable, both from the environment (os.UserHomeDir) and from the OS account +// database (user.Current, which ignores HOME and USERPROFILE). +func captureRealHomes() { + realHomesMu.Lock() + defer realHomesMu.Unlock() + if realHomesCaptured { + return + } + realHomesCaptured = true + if home, err := os.UserHomeDir(); err == nil && home != "" { + realHomes = append(realHomes, home) + } + if u, err := user.Current(); err == nil && u.HomeDir != "" { + realHomes = append(realHomes, u.HomeDir) + } +} + +// anthropicHomeBlocked reports whether home is the account's real home while +// test mode is on. Credential-file helpers use it so a test that failed to +// redirect the home directory can never read or rotate the developer's real +// Claude Code tokens (a rotated refresh token logs Claude Code out). On +// Windows this is the only guard: there is no keychain, the file is the store. +func anthropicHomeBlocked(home string) bool { + if !testMode { + return false + } + realHomesMu.Lock() + homes := realHomes + realHomesMu.Unlock() + for _, real := range homes { + if sameDirPath(home, real) { + return true + } + } + return false +} + +// sameDirPath reports whether a and b name the same directory, tolerating +// case differences on Windows and aliases such as symlinks. +func sameDirPath(a, b string) bool { + a, b = filepath.Clean(a), filepath.Clean(b) + if a == b || (runtime.GOOS == "windows" && strings.EqualFold(a, b)) { + return true + } + sa, errA := os.Stat(a) + sb, errB := os.Stat(b) + return errA == nil && errB == nil && os.SameFile(sa, sb) +} + // claudeCredentials represents the Claude Code credentials JSON structure. type claudeCredentials struct { ClaudeAiOauth struct { diff --git a/internal/api/anthropic_token_testmode_test.go b/internal/api/anthropic_token_testmode_test.go new file mode 100644 index 00000000..d8b2c2d4 --- /dev/null +++ b/internal/api/anthropic_token_testmode_test.go @@ -0,0 +1,130 @@ +package api + +import ( + "errors" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" +) + +// fakeRealHome makes dir stand in for the account's real home directory for +// the duration of the test, so the test-mode guard can be exercised without +// ever pointing a test at the developer's actual profile. The recorded real +// homes stay guarded: dir is added to them, never swapped in for them. +func fakeRealHome(t *testing.T, dir string) { + t.Helper() + realHomesMu.Lock() + saved := realHomes + realHomes = append(append([]string(nil), saved...), dir) + realHomesMu.Unlock() + t.Cleanup(func() { + realHomesMu.Lock() + realHomes = saved + realHomesMu.Unlock() + }) +} + +func TestSetTestModeCapturesRealHome(t *testing.T) { + realHomesMu.Lock() + captured, homes := realHomesCaptured, len(realHomes) + realHomesMu.Unlock() + if !captured || homes == 0 { + t.Fatalf("TestMain enabled test mode, so the real home must be recorded (captured=%v, homes=%d)", captured, homes) + } +} + +func TestAnthropicCredentialsFileRefusesRealHomeInTestMode(t *testing.T) { + realHome := t.TempDir() + fakeRealHome(t, realHome) + testhome.SetTestHome(t, realHome) + + claudeDir := filepath.Join(realHome, ".claude") + if err := os.MkdirAll(claudeDir, 0o700); err != nil { + t.Fatal(err) + } + credPath := filepath.Join(claudeDir, ".credentials.json") + original := `{"claudeAiOauth":{"accessToken":"real-access","refreshToken":"real-refresh","expiresAt":4102444800000}}` + if err := os.WriteFile(credPath, []byte(original), 0o600); err != nil { + t.Fatal(err) + } + + if got := DetectAnthropicToken(nil); got != "" { + t.Errorf("DetectAnthropicToken read the real credentials file in test mode: %q", got) + } + if got := DetectAnthropicCredentials(nil); got != nil { + t.Errorf("DetectAnthropicCredentials read the real credentials file in test mode: %+v", got) + } + if err := WriteAnthropicCredentials("new-access", "new-refresh", 3600); err == nil { + t.Error("WriteAnthropicCredentials must fail rather than rotate the real credentials file") + } else if runtime.GOOS == "windows" && !errors.Is(err, errRealCredentialsInTestMode) { + t.Errorf("err = %v, want errRealCredentialsInTestMode", err) + } + + data, err := os.ReadFile(credPath) + if err != nil { + t.Fatal(err) + } + if string(data) != original { + t.Fatalf("real credentials file was modified in test mode: %s", data) + } + if _, err := os.Stat(credPath + ".bak"); !os.IsNotExist(err) { + t.Errorf("no backup may be written next to the real credentials file, stat err = %v", err) + } +} + +func TestAnthropicCredentialsFileAllowsSandboxHomeInTestMode(t *testing.T) { + fakeRealHome(t, t.TempDir()) + home := t.TempDir() + testhome.SetTestHome(t, home) + + claudeDir := filepath.Join(home, ".claude") + if err := os.MkdirAll(claudeDir, 0o700); err != nil { + t.Fatal(err) + } + credPath := filepath.Join(claudeDir, ".credentials.json") + if err := os.WriteFile(credPath, []byte(`{"claudeAiOauth":{"accessToken":"sandbox-access","refreshToken":"r","expiresAt":4102444800000}}`), 0o600); err != nil { + t.Fatal(err) + } + if got := DetectAnthropicToken(nil); got != "sandbox-access" { + t.Fatalf("DetectAnthropicToken() = %q, want sandbox-access", got) + } +} + +func TestAnthropicHomeBlocked(t *testing.T) { + // Test mode stays on here: turning it off, even briefly, would let any + // concurrently running test reach the real keychain. + realHome := t.TempDir() + fakeRealHome(t, realHome) + + if !anthropicHomeBlocked(realHome) { + t.Fatal("real home must be blocked in test mode") + } + if !anthropicHomeBlocked(realHome + string(filepath.Separator)) { + t.Fatal("real home with a trailing separator must be blocked") + } + if anthropicHomeBlocked(t.TempDir()) { + t.Fatal("a sandbox home must not be blocked") + } +} + +// fakeRealHome must add to the recorded real homes, not replace them: +// otherwise the developer's actual home is unguarded while such a test runs. +func TestFakeRealHomeKeepsRecordedRealHomesBlocked(t *testing.T) { + realHomesMu.Lock() + recorded := append([]string(nil), realHomes...) + realHomesMu.Unlock() + if len(recorded) == 0 { + t.Fatal("TestMain enabled test mode, so a real home must be recorded") + } + + fakeRealHome(t, t.TempDir()) + + for _, home := range recorded { + if !anthropicHomeBlocked(home) { + t.Errorf("recorded real home %q is no longer blocked while a fake real home is set", home) + } + } +} diff --git a/internal/api/anthropic_token_unix.go b/internal/api/anthropic_token_unix.go index d6d2a045..4eb1ca35 100644 --- a/internal/api/anthropic_token_unix.go +++ b/internal/api/anthropic_token_unix.go @@ -16,20 +16,8 @@ import ( "time" ) -// testMode disables all keychain/keyring operations. Set to true in tests -// to prevent tests from reading or writing real Claude Code credentials. -// This is a critical safety guard - without it, tests can overwrite the user's -// real OAuth tokens in the macOS Keychain, logging them out of Claude Code. -var testMode bool - -// SetTestMode enables or disables test mode. When enabled, all keychain and -// keyring operations are skipped, and only file-based credential storage is used. -// Files are redirected by setting HOME to a temp dir in tests. -func SetTestMode(enabled bool) { - testMode = enabled -} - // getCredentialsFilePath returns the path to the Claude credentials file. +// It returns "" in test mode when the path would be the real account's file. func getCredentialsFilePath() string { home, err := os.UserHomeDir() if err != nil { @@ -37,7 +25,7 @@ func getCredentialsFilePath() string { home = u.HomeDir } } - if home == "" { + if home == "" || anthropicHomeBlocked(home) { return "" } return filepath.Join(home, ".claude", ".credentials.json") @@ -96,6 +84,10 @@ func detectAnthropicTokenPlatform(logger *slog.Logger) string { logger.Debug("Cannot determine home directory for credential file lookup") return "" } + if anthropicHomeBlocked(home) { + logger.Debug("Test mode: skipping the real credentials file") + return "" + } credPath := filepath.Join(home, ".claude", ".credentials.json") data, err := os.ReadFile(credPath) if err != nil { diff --git a/internal/api/anthropic_token_unix_coverage_test.go b/internal/api/anthropic_token_unix_coverage_test.go index 7dd26ae5..e3755d02 100644 --- a/internal/api/anthropic_token_unix_coverage_test.go +++ b/internal/api/anthropic_token_unix_coverage_test.go @@ -7,6 +7,7 @@ import ( "log/slog" "os" "path/filepath" + "strings" "testing" "time" ) @@ -269,3 +270,50 @@ func TestWriteAnthropicCredentials_ReturnsErrorForInvalidJSON(t *testing.T) { t.Fatal("expected error for invalid JSON, got nil") } } + +// Moved from extra_coverage_test.go: getCredentialsFilePath is Unix-only. + +func TestGetCredentialsFilePath_WithHome(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + + path := getCredentialsFilePath() + expected := filepath.Join(home, ".claude", ".credentials.json") + if path != expected { + t.Errorf("getCredentialsFilePath() = %q, want %q", path, expected) + } +} + +func TestGetCredentialsFilePath_ReturnsNonEmpty(t *testing.T) { + // Regardless of platform, should return a non-empty path if home exists + path := getCredentialsFilePath() + // Could be empty in some edge cases, but should not panic + _ = path +} + +func TestGetCredentialsFilePath_EmptyHOME(t *testing.T) { + // When HOME is not set, getCredentialsFilePath may return "" or + // use user.Current() as a fallback. Either way it must not panic. + t.Setenv("HOME", "") + + path := getCredentialsFilePath() + // The function returns "" or a valid path via user.Current() fallback. + // We just verify no panic and correct format if non-empty. + if path != "" { + // path should end with .claude/.credentials.json + if !strings.HasSuffix(path, ".credentials.json") { + t.Errorf("getCredentialsFilePath() = %q, should end with .credentials.json", path) + } + } +} + +func TestGetCredentialsFilePath_ValidHOME(t *testing.T) { + dir := t.TempDir() + t.Setenv("HOME", dir) + + path := getCredentialsFilePath() + expected := filepath.Join(dir, ".claude", ".credentials.json") + if path != expected { + t.Errorf("getCredentialsFilePath() = %q, want %q", path, expected) + } +} diff --git a/internal/api/anthropic_token_windows.go b/internal/api/anthropic_token_windows.go index e9feb49d..71df71bb 100644 --- a/internal/api/anthropic_token_windows.go +++ b/internal/api/anthropic_token_windows.go @@ -11,14 +11,18 @@ import ( "time" ) -// testMode disables keychain/keyring operations during tests. -// On Windows this is a no-op (no keychain), but the variable must exist -// for cross-platform compilation. -var testMode bool - -// SetTestMode enables or disables test mode. -func SetTestMode(enabled bool) { - testMode = enabled +// anthropicCredentialsFilePath returns %USERPROFILE%\.claude\.credentials.json, +// the only place Claude Code keeps its OAuth tokens on Windows. In test mode +// it refuses the real account profile (see anthropicHomeBlocked). +func anthropicCredentialsFilePath() (string, error) { + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + if anthropicHomeBlocked(home) { + return "", errRealCredentialsInTestMode + } + return filepath.Join(home, ".claude", ".credentials.json"), nil } // detectAnthropicCredentialsPlatform tries to detect full OAuth credentials on Windows. @@ -27,11 +31,10 @@ func detectAnthropicCredentialsPlatform(logger *slog.Logger) *AnthropicCredentia logger = slog.Default() } - home, err := os.UserHomeDir() + credPath, err := anthropicCredentialsFilePath() if err != nil { return nil } - credPath := filepath.Join(home, ".claude", ".credentials.json") data, err := os.ReadFile(credPath) if err != nil { return nil @@ -57,11 +60,10 @@ func detectAnthropicCredentialsPlatform(logger *slog.Logger) *AnthropicCredentia // // Related: https://github.com/onllm-dev/onWatch/issues/16 func WriteAnthropicCredentials(accessToken, refreshToken string, expiresIn int) error { - home, err := os.UserHomeDir() + credPath, err := anthropicCredentialsFilePath() if err != nil { return err } - credPath := filepath.Join(home, ".claude", ".credentials.json") data, err := os.ReadFile(credPath) if err != nil { return err @@ -109,11 +111,10 @@ func detectAnthropicTokenPlatform(logger *slog.Logger) string { logger = slog.Default() } - home, err := os.UserHomeDir() + credPath, err := anthropicCredentialsFilePath() if err != nil { return "" } - credPath := filepath.Join(home, ".claude", ".credentials.json") data, err := os.ReadFile(credPath) if err != nil { return "" diff --git a/internal/api/antigravity_cli.go b/internal/api/antigravity_cli.go index 7cbb9959..1d436679 100644 --- a/internal/api/antigravity_cli.go +++ b/internal/api/antigravity_cli.go @@ -56,11 +56,11 @@ type AntigravityCLIRunner struct { rootCtx context.Context rootCancel context.CancelFunc - mu sync.Mutex - sess *agySession - lastUsed time.Time - failures int - watchdog sync.Once + mu sync.Mutex + sess *agySession + lastUsed time.Time + failures int + watchdog sync.Once } // NewAntigravityCLIRunner creates a runner. It does not launch agy until the diff --git a/internal/api/antigravity_client.go b/internal/api/antigravity_client.go index 407614a1..92e420e5 100644 --- a/internal/api/antigravity_client.go +++ b/internal/api/antigravity_client.go @@ -47,6 +47,15 @@ type AntigravityClient struct { httpClient *http.Client connection *AntigravityConnection logger *slog.Logger + // runCommand runs a discovery tool (ps, lsof, ss, netstat, PowerShell, + // WMIC) and returns its stdout. Tests swap it to feed canned output on + // any host OS. + runCommand func(ctx context.Context, name string, args ...string) ([]byte, error) +} + +// runExternalCommand is the default AntigravityClient.runCommand. +func runExternalCommand(ctx context.Context, name string, args ...string) ([]byte, error) { + return exec.CommandContext(ctx, name, args...).Output() } // AntigravityOption configures an AntigravityClient. @@ -82,7 +91,8 @@ func NewAntigravityClient(logger *slog.Logger, opts ...AntigravityOption) *Antig }, }, }, - logger: logger, + logger: logger, + runCommand: runExternalCommand, } for _, opt := range opts { @@ -235,8 +245,7 @@ func (c *AntigravityClient) detectProcess(ctx context.Context) (*AntigravityProc // detectProcessUnix finds the process on Unix-like systems. func (c *AntigravityClient) detectProcessUnix(ctx context.Context) (*AntigravityProcessInfo, error) { - cmd := exec.CommandContext(ctx, "ps", "aux") - output, err := cmd.Output() + output, err := c.runCommand(ctx, "ps", "aux") if err != nil { return nil, fmt.Errorf("antigravity: ps command failed: %w", err) } @@ -308,11 +317,9 @@ func (c *AntigravityClient) detectProcessWindows(ctx context.Context) (*Antigrav } // Fallback 2: WMIC (deprecated on newer Windows 11 but works on older builds) - cmd := exec.CommandContext(ctx, "wmic", "process", "where", + output, err := c.runCommand(ctx, "wmic", "process", "where", "name like '%antigravity%' or commandline like '%antigravity%'", "get", "processid,commandline", "/format:csv") - - output, err := cmd.Output() if err == nil { if info := c.parseWMICOutput(string(output)); info != nil { return info, nil @@ -326,15 +333,17 @@ func (c *AntigravityClient) detectProcessWindows(ctx context.Context) (*Antigrav // This is the most reliable method on modern Windows as it searches command lines // for both "antigravity" and "language_server" process names. func (c *AntigravityClient) detectProcessWindowsCIM(ctx context.Context) (*AntigravityProcessInfo, error) { - // Single PowerShell command that finds all candidate processes by command line content + // Single PowerShell command that finds all candidate processes by command line content. + // The filter text itself contains "antigravity" and "language_server", so the + // querying powershell.exe would match its own filter; $PID excludes it here and + // isAntigravityProbeProcess drops any other probe (e.g. a concurrent instance). psCmd := `Get-CimInstance Win32_Process | Where-Object {` + - ` $_.CommandLine -and (` + + ` $_.ProcessId -ne $PID -and $_.CommandLine -and (` + `$_.CommandLine -like '*antigravity*' -or ` + `$_.Name -like '*language_server*'` + `)} | Select-Object ProcessId, Name, CommandLine | ConvertTo-Json` - cmd := exec.CommandContext(ctx, "powershell", "-NoProfile", "-Command", psCmd) - output, err := cmd.Output() + output, err := c.runCommand(ctx, "powershell", "-NoProfile", "-Command", psCmd) if err != nil { return nil, fmt.Errorf("antigravity: CIM query failed: %w", err) } @@ -365,7 +374,7 @@ func (c *AntigravityClient) detectProcessWindowsCIM(ctx context.Context) (*Antig for _, proc := range processes { cmdLine := proc.CommandLine - if cmdLine == "" { + if cmdLine == "" || isAntigravityProbeProcess(cmdLine) { continue } @@ -413,7 +422,7 @@ func (c *AntigravityClient) parseWMICOutput(output string) *AntigravityProcessIn } commandLine := strings.Join(parts[1:len(parts)-1], ",") - if !strings.Contains(strings.ToLower(commandLine), "antigravity") { + if !strings.Contains(strings.ToLower(commandLine), "antigravity") || isAntigravityProbeProcess(commandLine) { continue } @@ -442,10 +451,8 @@ func (c *AntigravityClient) parseWMICOutput(output string) *AntigravityProcessIn // detectProcessWindowsPowerShell uses PowerShell as fallback. func (c *AntigravityClient) detectProcessWindowsPowerShell(ctx context.Context) (*AntigravityProcessInfo, error) { // Search both "antigravity" and "language_server" process names - cmd := exec.CommandContext(ctx, "powershell", "-NoProfile", "-Command", + output, err := c.runCommand(ctx, "powershell", "-NoProfile", "-Command", "Get-Process | Where-Object { $_.ProcessName -like '*antigravity*' -or $_.ProcessName -like '*language_server*' } | Select-Object Id, ProcessName | ConvertTo-Json") - - output, err := cmd.Output() if err != nil { return nil, ErrAntigravityProcessNotFound } @@ -473,16 +480,16 @@ func (c *AntigravityClient) detectProcessWindowsPowerShell(ctx context.Context) bestScore := -1 for _, proc := range processes { - cmdLineCmd := exec.CommandContext(ctx, "powershell", "-Command", + // -NoProfile: output from a user profile script would otherwise be + // prepended to the command line read back here. + cmdOutput, err := c.runCommand(ctx, "powershell", "-NoProfile", "-Command", fmt.Sprintf("(Get-CimInstance Win32_Process -Filter 'ProcessId = %d').CommandLine", proc.Id)) - - cmdOutput, err := cmdLineCmd.Output() if err != nil { continue } commandLine := strings.TrimSpace(string(cmdOutput)) - if !strings.Contains(strings.ToLower(commandLine), "antigravity") { + if !strings.Contains(strings.ToLower(commandLine), "antigravity") || isAntigravityProbeProcess(commandLine) { continue } @@ -523,8 +530,7 @@ func (c *AntigravityClient) discoverPorts(ctx context.Context, pid int) ([]int, // discoverPortsMacOS uses lsof to find listening ports. func (c *AntigravityClient) discoverPortsMacOS(ctx context.Context, pid int) ([]int, error) { - cmd := exec.CommandContext(ctx, "lsof", "-nP", "-iTCP", "-sTCP:LISTEN", "-a", "-p", strconv.Itoa(pid)) - output, err := cmd.Output() + output, err := c.runCommand(ctx, "lsof", "-nP", "-iTCP", "-sTCP:LISTEN", "-a", "-p", strconv.Itoa(pid)) if err != nil { return nil, err } @@ -535,8 +541,7 @@ func (c *AntigravityClient) discoverPortsMacOS(ctx context.Context, pid int) ([] // discoverPortsLinux uses ss or netstat to find listening ports. func (c *AntigravityClient) discoverPortsLinux(ctx context.Context, pid int) ([]int, error) { // Try ss first - cmd := exec.CommandContext(ctx, "ss", "-tlnp") - output, err := cmd.Output() + output, err := c.runCommand(ctx, "ss", "-tlnp") if err == nil { ports := parsePortsFromSS(string(output), pid) if len(ports) > 0 { @@ -545,8 +550,7 @@ func (c *AntigravityClient) discoverPortsLinux(ctx context.Context, pid int) ([] } // Fallback to netstat - cmd = exec.CommandContext(ctx, "netstat", "-tlnp") - output, err = cmd.Output() + output, err = c.runCommand(ctx, "netstat", "-tlnp") if err != nil { return nil, err } @@ -556,8 +560,7 @@ func (c *AntigravityClient) discoverPortsLinux(ctx context.Context, pid int) ([] // discoverPortsWindows uses netstat to find listening ports. func (c *AntigravityClient) discoverPortsWindows(ctx context.Context, pid int) ([]int, error) { - cmd := exec.CommandContext(ctx, "netstat", "-ano") - output, err := cmd.Output() + output, err := c.runCommand(ctx, "netstat", "-ano") if err != nil { return nil, err } @@ -672,6 +675,19 @@ func scoreWindowsCandidate(info *AntigravityProcessInfo) int { return score } +// isAntigravityProbeProcess reports whether a candidate command line is one of +// onWatch's own discovery queries rather than Antigravity. The CIM and WMIC +// filters carry the literal "antigravity" (and "language_server") in their own +// command lines, so the querying powershell.exe or wmic.exe - or a concurrent +// probe from another onWatch instance - otherwise matches itself. Left in, it +// outscores nothing-found, which masks "not running" as a port failure and +// skips the fallbacks; with an equally scored real server it can win outright. +func isAntigravityProbeProcess(commandLine string) bool { + lower := strings.ToLower(commandLine) + return strings.Contains(lower, "win32_process") || + (strings.Contains(lower, "wmic") && strings.Contains(lower, "process where")) +} + func parsePortsFromLsof(output string) []int { var ports []int portPattern := regexp.MustCompile(`:(\d+)\s+\(LISTEN\)`) @@ -725,17 +741,21 @@ func parsePortsFromNetstat(output string, pid int) []int { return ports } +// parsePortsFromWindowsNetstat reads `netstat -ano` TCP rows: +// Proto, Local Address, Foreign Address, State, PID. Rows whose state is +// LISTENING are preferred. The State column is localized (e.g. "ABHÖREN" on +// German Windows) and printed in the OEM code page, so when no row for the PID +// says LISTENING, a listener is recognized by its unconnected foreign address +// (port 0, as in 0.0.0.0:0 or [::]:0) instead. Bound-but-not-listening sockets +// also show port 0 there, so rows in a known English non-listening state +// (BOUND, CLOSED, ...) are skipped by that fallback. func parsePortsFromWindowsNetstat(output string, pid int) []int { - var ports []int + var listening, fallback []int portPattern := regexp.MustCompile(`:(\d+)$`) for _, line := range strings.Split(output, "\n") { - if !strings.Contains(line, "LISTENING") { - continue - } - parts := strings.Fields(line) - if len(parts) < 5 { + if len(parts) < 5 || !strings.EqualFold(parts[0], "TCP") { continue } @@ -744,13 +764,43 @@ func parsePortsFromWindowsNetstat(output string, pid int) []int { continue } - localAddr := parts[1] - if match := portPattern.FindStringSubmatch(localAddr); len(match) > 1 { - if port, err := strconv.Atoi(match[1]); err == nil { - ports = append(ports, port) - } + match := portPattern.FindStringSubmatch(parts[1]) + if len(match) < 2 { + continue + } + port, err := strconv.Atoi(match[1]) + if err != nil { + continue + } + + state := strings.ToUpper(strings.Join(parts[3:len(parts)-1], " ")) + switch { + case state == "LISTENING": + listening = append(listening, port) + case strings.HasSuffix(parts[2], ":0") && !windowsNetstatNonListeningStates[state]: + fallback = append(fallback, port) } } - return ports + if len(listening) > 0 { + return listening + } + return fallback +} + +// windowsNetstatNonListeningStates are the English netstat TCP states other +// than LISTENING. The foreign-port-0 fallback skips them. +var windowsNetstatNonListeningStates = map[string]bool{ + "BOUND": true, + "CLOSED": true, + "CLOSE_WAIT": true, + "CLOSING": true, + "DELETE_TCB": true, + "ESTABLISHED": true, + "FIN_WAIT_1": true, + "FIN_WAIT_2": true, + "LAST_ACK": true, + "SYN_RECEIVED": true, + "SYN_SENT": true, + "TIME_WAIT": true, } diff --git a/internal/api/antigravity_client_coverage_test.go b/internal/api/antigravity_client_coverage_test.go index 5a160a9a..b6d3767f 100644 --- a/internal/api/antigravity_client_coverage_test.go +++ b/internal/api/antigravity_client_coverage_test.go @@ -7,6 +7,7 @@ import ( "net/http" "net/http/httptest" "os" + "slices" "strings" "testing" "time" @@ -193,7 +194,12 @@ func TestParsePortsFromLsof_Empty(t *testing.T) { } func TestParsePortsFromWindowsNetstat_NoListening(t *testing.T) { - output := ` TCP 0.0.0.0:42100 0.0.0.0:0 ESTABLISHED 1234 + // Non-listening TCP rows always have a connected foreign address; only a + // listener shows port 0 there (0.0.0.0:0 / [::]:0). + output := ` TCP 127.0.0.1:42100 10.0.0.1:443 ESTABLISHED 1234 + TCP 127.0.0.1:42101 127.0.0.1:50000 TIME_WAIT 1234 + TCP 127.0.0.1:42102 10.0.0.1:443 SYN_SENT 1234 + UDP 0.0.0.0:42103 *:* 1234 ` ports := parsePortsFromWindowsNetstat(output, 1234) if len(ports) != 0 { @@ -201,6 +207,33 @@ func TestParsePortsFromWindowsNetstat_NoListening(t *testing.T) { } } +// A socket that is bound but not listening (BOUND, or CLOSED after close) +// also shows a foreign address of port 0. When the state column says +// LISTENING somewhere, only those rows count. +func TestParsePortsFromWindowsNetstat_PrefersListeningOverBound(t *testing.T) { + output := " TCP 0.0.0.0:50001 0.0.0.0:0 BOUND 1234\r\n" + + " TCP 127.0.0.1:42100 0.0.0.0:0 LISTENING 1234\r\n" + + " TCP 127.0.0.1:50002 0.0.0.0:0 CLOSED 1234\r\n" + + " TCP [::]:42101 [::]:0 LISTENING 1234\r\n" + ports := parsePortsFromWindowsNetstat(output, 1234) + if !slices.Equal(ports, []int{42100, 42101}) { + t.Fatalf("ports = %v, want [42100 42101]", ports) + } +} + +// Without any LISTENING row (localized Windows), the foreign-port-0 fallback +// still skips rows whose state is a known English non-listening state. +func TestParsePortsFromWindowsNetstat_FallbackSkipsKnownNonListeningStates(t *testing.T) { + output := " TCP 0.0.0.0:50001 0.0.0.0:0 BOUND 1234\r\n" + + " TCP 127.0.0.1:50002 0.0.0.0:0 CLOSED 1234\r\n" + + " TCP 127.0.0.1:50003 0.0.0.0:0 SYN_SENT 1234\r\n" + + " TCP 127.0.0.1:7007 0.0.0.0:0 ECOUTE 1234\r\n" + ports := parsePortsFromWindowsNetstat(output, 1234) + if !slices.Equal(ports, []int{7007}) { + t.Fatalf("ports = %v, want [7007]", ports) + } +} + func TestProbePort_Success200(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // Verify the probe request diff --git a/internal/api/antigravity_command_coverage_test.go b/internal/api/antigravity_command_coverage_test.go index ad8e0aec..a9748dc3 100644 --- a/internal/api/antigravity_command_coverage_test.go +++ b/internal/api/antigravity_command_coverage_test.go @@ -2,45 +2,98 @@ package api import ( "context" + "errors" "io" "log/slog" - "os" - "path/filepath" "runtime" "slices" + "strings" "testing" ) -func writeExecutable(t *testing.T, dir, name, content string) { - t.Helper() - path := filepath.Join(dir, name) - if err := os.WriteFile(path, []byte(content), 0o755); err != nil { - t.Fatalf("write executable %s: %v", name, err) +// fakeAntigravityCommands replaces the client's command runner with handler, +// so discovery parsing is exercised with canned tool output on every host OS +// (shell-script fakes on PATH cannot run on Windows, and the real netstat, +// PowerShell and WMIC would answer instead). +func fakeAntigravityCommands(client *AntigravityClient, handler func(name string, args []string) (string, error)) { + client.runCommand = func(_ context.Context, name string, args ...string) ([]byte, error) { + out, err := handler(name, args) + return []byte(out), err } } -func withPathDir(t *testing.T, dir string) { - t.Helper() - oldPath := os.Getenv("PATH") - t.Cleanup(func() { _ = os.Setenv("PATH", oldPath) }) - if err := os.Setenv("PATH", dir+string(os.PathListSeparator)+oldPath); err != nil { - t.Fatalf("set PATH: %v", err) - } -} +var errFakeCommand = errors.New("fake command failed") + +// windowsCIMSelfRow is what Get-CimInstance returned on a real Windows runner +// with no Antigravity running: the querying powershell.exe, which matches its +// own '*antigravity*' filter. +const windowsCIMSelfRow = `{"ProcessId":2880,"Name":"powershell.exe","CommandLine":"powershell -NoProfile -Command \"Get-CimInstance Win32_Process | Where-Object { $_.CommandLine -and ($_.CommandLine -like '*antigravity*' -or $_.Name -like '*language_server*')} | Select-Object ProcessId, Name, CommandLine | ConvertTo-Json\""}` func discardLoggerCommands() *slog.Logger { return slog.New(slog.NewTextHandler(io.Discard, nil)) } func TestAntigravityCommandHelpers(t *testing.T) { - client := NewAntigravityClient(discardLoggerCommands()) ctx := context.Background() + t.Run("default runner executes real commands", func(t *testing.T) { + client := NewAntigravityClient(discardLoggerCommands()) + if client.runCommand == nil { + t.Fatal("NewAntigravityClient must install a command runner") + } + if _, err := client.runCommand(ctx, "onwatch-no-such-binary-for-test"); err == nil { + t.Fatal("running a missing binary must fail") + } + }) + + t.Run("detect process unix parses ps", func(t *testing.T) { + client := NewAntigravityClient(discardLoggerCommands()) + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + if name != "ps" || !slices.Equal(args, []string{"aux"}) { + return "", errFakeCommand + } + return "USER PID %CPU %MEM VSZ RSS TTY STAT START TIME COMMAND\n" + + "me 555 0.0 0.1 1 1 ?? S 10:00 0:00 /Applications/Antigravity.app/language_server_macos --csrf_token unix --extension_server_port 6336\n", nil + }) + + info, err := client.detectProcessUnix(ctx) + if err != nil { + t.Fatalf("detectProcessUnix: %v", err) + } + if info.PID != 555 || info.CSRFToken != "unix" || info.ExtensionServerPort != 6336 { + t.Fatalf("unexpected ps info: %+v", info) + } + }) + + t.Run("discover ports macos uses lsof", func(t *testing.T) { + client := NewAntigravityClient(discardLoggerCommands()) + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + if name != "lsof" || !slices.Contains(args, "555") { + return "", errFakeCommand + } + return "COMMAND PID USER FD TYPE DEVICE SIZE/OFF NODE NAME\n" + + "language_ 555 me 9u IPv4 0x1 0t0 TCP 127.0.0.1:6337 (LISTEN)\n", nil + }) + + ports, err := client.discoverPortsMacOS(ctx, 555) + if err != nil { + t.Fatalf("discoverPortsMacOS: %v", err) + } + if !slices.Equal(ports, []int{6337}) { + t.Fatalf("discoverPortsMacOS() = %v, want [6337]", ports) + } + }) + t.Run("discover ports linux uses ss", func(t *testing.T) { - dir := t.TempDir() - writeExecutable(t, dir, "ss", "#!/bin/sh\ncat <<'EOF'\nLISTEN 0 4096 127.0.0.1:4242 0.0.0.0:* users:((\"language_server\",pid=777,fd=9))\nEOF\n") - writeExecutable(t, dir, "netstat", "#!/bin/sh\nexit 1\n") - withPathDir(t, dir) + client := NewAntigravityClient(discardLoggerCommands()) + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + switch name { + case "ss": + return "LISTEN 0 4096 127.0.0.1:4242 0.0.0.0:* users:((\"language_server\",pid=777,fd=9))\n", nil + default: + return "", errFakeCommand + } + }) ports, err := client.discoverPortsLinux(ctx, 777) if err != nil { @@ -52,10 +105,17 @@ func TestAntigravityCommandHelpers(t *testing.T) { }) t.Run("discover ports linux falls back to netstat", func(t *testing.T) { - dir := t.TempDir() - writeExecutable(t, dir, "ss", "#!/bin/sh\ncat <<'EOF'\nLISTEN 0 4096 127.0.0.1:9999 0.0.0.0:* users:((\"other\",pid=1,fd=1))\nEOF\n") - writeExecutable(t, dir, "netstat", "#!/bin/sh\ncat <<'EOF'\ntcp 0 0 127.0.0.1:5151 0.0.0.0:* LISTEN 777/language_server\nEOF\n") - withPathDir(t, dir) + client := NewAntigravityClient(discardLoggerCommands()) + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + switch name { + case "ss": + return "LISTEN 0 4096 127.0.0.1:9999 0.0.0.0:* users:((\"other\",pid=1,fd=1))\n", nil + case "netstat": + return "tcp 0 0 127.0.0.1:5151 0.0.0.0:* LISTEN 777/language_server\n", nil + default: + return "", errFakeCommand + } + }) ports, err := client.discoverPortsLinux(ctx, 777) if err != nil { @@ -67,23 +127,51 @@ func TestAntigravityCommandHelpers(t *testing.T) { }) t.Run("discover ports windows parses netstat", func(t *testing.T) { - dir := t.TempDir() - writeExecutable(t, dir, "netstat", "#!/bin/sh\ncat <<'EOF'\n TCP 127.0.0.1:7007 0.0.0.0:0 LISTENING 888\nEOF\n") - withPathDir(t, dir) + client := NewAntigravityClient(discardLoggerCommands()) + // Real netstat -ano output: CRLF line endings, header rows, IPv4 and + // IPv6 listeners, an established connection and UDP rows. + out := "\r\nActive Connections\r\n\r\n" + + " Proto Local Address Foreign Address State PID\r\n" + + " TCP 0.0.0.0:135 0.0.0.0:0 LISTENING 1000\r\n" + + " TCP 127.0.0.1:7007 0.0.0.0:0 LISTENING 888\r\n" + + " TCP 127.0.0.1:7008 127.0.0.1:50000 ESTABLISHED 888\r\n" + + " TCP [::1]:7009 [::]:0 LISTENING 888\r\n" + + " UDP 0.0.0.0:5353 *:* 888\r\n" + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + if name != "netstat" || !slices.Equal(args, []string{"-ano"}) { + return "", errFakeCommand + } + return out, nil + }) ports, err := client.discoverPortsWindows(ctx, 888) if err != nil { t.Fatalf("discoverPortsWindows: %v", err) } + if !slices.Equal(ports, []int{7007, 7009}) { + t.Fatalf("discoverPortsWindows() = %v, want [7007 7009]", ports) + } + }) + + t.Run("discover ports windows ignores localized state", func(t *testing.T) { + // German Windows prints ABHÖREN/HERGESTELLT instead of LISTENING/ESTABLISHED. + out := " Proto Lokale Adresse Remoteadresse Status PID\r\n" + + " TCP 127.0.0.1:7007 0.0.0.0:0 ABH\x99REN 888\r\n" + + " TCP 127.0.0.1:7008 127.0.0.1:50000 HERGESTELLT 888\r\n" + ports := parsePortsFromWindowsNetstat(out, 888) if !slices.Equal(ports, []int{7007}) { - t.Fatalf("discoverPortsWindows() = %v, want [7007]", ports) + t.Fatalf("parsePortsFromWindowsNetstat(localized) = %v, want [7007]", ports) } }) t.Run("detect process windows cim handles single object", func(t *testing.T) { - dir := t.TempDir() - writeExecutable(t, dir, "powershell", "#!/bin/sh\ncat <<'EOF'\n{\"ProcessId\":1234,\"Name\":\"language_server_windows_x64\",\"CommandLine\":\"C:/antigravity/language_server_windows_x64.exe --csrf_token tok --extension_server_port 7447\"}\nEOF\n") - withPathDir(t, dir) + client := NewAntigravityClient(discardLoggerCommands()) + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + if name != "powershell" { + return "", errFakeCommand + } + return "{\"ProcessId\":1234,\"Name\":\"language_server_windows_x64\",\"CommandLine\":\"C:/antigravity/language_server_windows_x64.exe --csrf_token tok --extension_server_port 7447\"}\r\n", nil + }) info, err := client.detectProcessWindowsCIM(ctx) if err != nil { @@ -94,10 +182,59 @@ func TestAntigravityCommandHelpers(t *testing.T) { } }) + t.Run("detect process windows cim query excludes itself", func(t *testing.T) { + client := NewAntigravityClient(discardLoggerCommands()) + var query string + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + query = strings.Join(args, " ") + return windowsCIMSelfRow, nil + }) + + if _, err := client.detectProcessWindowsCIM(ctx); !errors.Is(err, ErrAntigravityProcessNotFound) { + t.Fatalf("own query process must not be detected as Antigravity, err = %v", err) + } + if !strings.Contains(query, "$_.ProcessId -ne $PID") { + t.Fatalf("CIM query must exclude the querying PowerShell process, got %q", query) + } + }) + + t.Run("detect process windows cim prefers server over own query", func(t *testing.T) { + client := NewAntigravityClient(discardLoggerCommands()) + // A server whose command line carries no flags scores the same as the + // self row (antigravity + language_server), and the self row comes first. + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + return "[\r\n" + windowsCIMSelfRow + ",\r\n" + + "{\"ProcessId\":4242,\"Name\":\"language_server_windows_x64.exe\",\"CommandLine\":\"C:\\\\Antigravity\\\\language_server_windows_x64.exe\"}\r\n]\r\n", nil + }) + + info, err := client.detectProcessWindowsCIM(ctx) + if err != nil { + t.Fatalf("detectProcessWindowsCIM: %v", err) + } + if info.PID != 4242 { + t.Fatalf("detected PID %d, want the language server 4242", info.PID) + } + }) + t.Run("detect process windows powershell uses process lookup", func(t *testing.T) { - dir := t.TempDir() - writeExecutable(t, dir, "powershell", "#!/bin/sh\ncase \"$*\" in\n *\"Get-Process\"*)\n printf '[{\"Id\":4321}]'\n ;;\n *\"ProcessId = 4321\"*)\n printf 'C:/Users/test/antigravity/language_server.exe --csrf_token ps --extension_server_port 8558'\n ;;\n *)\n exit 1\n ;;\nesac\n") - withPathDir(t, dir) + client := NewAntigravityClient(discardLoggerCommands()) + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + if name != "powershell" { + return "", errFakeCommand + } + joined := strings.Join(args, " ") + switch { + case strings.Contains(joined, "Get-Process"): + return "[{\"Id\":4321}]\r\n", nil + case strings.Contains(joined, "ProcessId = 4321"): + if !slices.Contains(args, "-NoProfile") { + return "profile noise\r\n", nil + } + return "C:/Users/test/antigravity/language_server.exe --csrf_token ps --extension_server_port 8558\r\n", nil + default: + return "", errFakeCommand + } + }) info, err := client.detectProcessWindowsPowerShell(ctx) if err != nil { @@ -109,10 +246,17 @@ func TestAntigravityCommandHelpers(t *testing.T) { }) t.Run("detect process windows falls back to wmic", func(t *testing.T) { - dir := t.TempDir() - writeExecutable(t, dir, "powershell", "#!/bin/sh\nexit 1\n") - writeExecutable(t, dir, "wmic", "#!/bin/sh\ncat <<'EOF'\nNode,CommandLine,ProcessId\nHOST,C:/antigravity/language_server.exe --csrf_token wmic --extension_server_port 9669,2468\nEOF\n") - withPathDir(t, dir) + client := NewAntigravityClient(discardLoggerCommands()) + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + if name != "wmic" { + return "", errFakeCommand + } + // WMIC CSV uses \r\r\n line endings and lists its own process, + // whose command line matches the '%antigravity%' filter. + return "\r\r\nNode,CommandLine,ProcessId\r\r\n" + + "HOST,wmic process where \"name like '%antigravity%' or commandline like '%antigravity%'\" get processid,commandline /format:csv,1357\r\r\n" + + "HOST,C:/antigravity/language_server.exe --csrf_token wmic --extension_server_port 9669,2468\r\r\n", nil + }) info, err := client.detectProcessWindows(ctx) if err != nil { @@ -122,6 +266,27 @@ func TestAntigravityCommandHelpers(t *testing.T) { t.Fatalf("unexpected WMIC info: %+v", info) } }) + + t.Run("detect process windows reports not found when only probes match", func(t *testing.T) { + client := NewAntigravityClient(discardLoggerCommands()) + fakeAntigravityCommands(client, func(name string, args []string) (string, error) { + switch { + case name == "powershell" && strings.Contains(strings.Join(args, " "), "Win32_Process |"): + return windowsCIMSelfRow, nil + case name == "powershell": + return "", nil + case name == "wmic": + return "Node,CommandLine,ProcessId\r\r\n" + + "HOST,wmic process where \"name like '%antigravity%' or commandline like '%antigravity%'\" get processid,commandline /format:csv,1357\r\r\n", nil + default: + return "", errFakeCommand + } + }) + + if _, err := client.detectProcessWindows(ctx); !errors.Is(err, ErrAntigravityProcessNotFound) { + t.Fatalf("err = %v, want ErrAntigravityProcessNotFound", err) + } + }) } func TestMiniMaxDisplayName_DefaultAndKnown(t *testing.T) { diff --git a/internal/api/codex_credentials_test.go b/internal/api/codex_credentials_test.go index e430df53..2db667be 100644 --- a/internal/api/codex_credentials_test.go +++ b/internal/api/codex_credentials_test.go @@ -10,6 +10,8 @@ import ( "strings" "testing" "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func discardLoggerCredentials() *slog.Logger { @@ -29,7 +31,7 @@ func isolateOpenCodeEnv(t *testing.T) { func TestDetectCodexCredentials_ParsesOAuthTokens(t *testing.T) { t.Setenv("CODEX_HOME", t.TempDir()) - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) authPath := filepath.Join(os.Getenv("CODEX_HOME"), "auth.json") if err := os.WriteFile(authPath, []byte(`{ @@ -63,7 +65,7 @@ func TestDetectCodexCredentials_ParsesOAuthTokens(t *testing.T) { func TestDetectCodexCredentials_ParsesAPIKey(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") codexDir := filepath.Join(home, ".codex") @@ -90,7 +92,7 @@ func TestDetectCodexCredentials_ParsesAPIKey(t *testing.T) { // This ensures the fallback path works when no auth file is available. func TestDetectCodexCredentials_EnvVarFallback(t *testing.T) { t.Setenv("CODEX_HOME", t.TempDir()) - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) isolateOpenCodeEnv(t) t.Setenv("CODEX_TOKEN", "env_access_token") @@ -114,7 +116,7 @@ func TestDetectCodexCredentials_EnvVarFallback(t *testing.T) { func TestDetectCodexToken_PrefersAccessToken(t *testing.T) { t.Setenv("CODEX_HOME", t.TempDir()) - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) authPath := filepath.Join(os.Getenv("CODEX_HOME"), "auth.json") if err := os.WriteFile(authPath, []byte(`{ @@ -132,7 +134,7 @@ func TestDetectCodexToken_PrefersAccessToken(t *testing.T) { func TestDetectCodexToken_RejectsAPIKeyOnly(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") codexDir := filepath.Join(home, ".codex") @@ -325,7 +327,7 @@ func TestWriteCodexCredentials_NewFile(t *testing.T) { func TestDetectCodexCredentials_ParsesUserIDFromIDToken(t *testing.T) { t.Setenv("CODEX_HOME", t.TempDir()) - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) header := "eyJhbGciOiJub25lIn0" payloadJSON := `{"https://api.openai.com/auth":{"chatgpt_user_id":"user-123"}}` @@ -370,7 +372,7 @@ func openCodeAuthJSON(access, refresh, accountID string, expiresMs int64) string func setOpenCodeOnly(t *testing.T) string { t.Helper() home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") t.Setenv("CODEX_TOKEN", "") t.Setenv("OPENCODE_HOME", "") @@ -419,7 +421,7 @@ func TestDetectCodexCredentials_OpenCodeFormat(t *testing.T) { func TestDetectCodexCredentials_CodexPriorityOverOpenCode(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") t.Setenv("CODEX_TOKEN", "") t.Setenv("OPENCODE_HOME", "") @@ -455,7 +457,7 @@ func TestDetectCodexCredentials_CodexPriorityOverOpenCode(t *testing.T) { func TestDetectCodexCredentials_OpenCodeHomeOverride(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", "") t.Setenv("CODEX_TOKEN", "") t.Setenv("XDG_DATA_HOME", "") @@ -613,7 +615,7 @@ func TestCodexCredentials_CompositeExternalID(t *testing.T) { want: "", // user_id missing -> ambiguous identity, caller must dedupe at account level }, { - name: "neither present", + name: "neither present", creds: CodexCredentials{}, want: "", }, diff --git a/internal/api/codex_oauth_test.go b/internal/api/codex_oauth_test.go index 40aebd85..46748036 100644 --- a/internal/api/codex_oauth_test.go +++ b/internal/api/codex_oauth_test.go @@ -151,8 +151,8 @@ func TestParseIDTokenExpiry_ValidJWT(t *testing.T) { // Header: {"alg":"none"} // Payload: {"exp":1893456000} (2030-01-01 00:00:00 UTC) // Signature: (empty) - header := "eyJhbGciOiJub25lIn0" // {"alg":"none"} base64url encoded - payload := "eyJleHAiOjE4OTM0NTYwMDB9" // {"exp":1893456000} base64url encoded + header := "eyJhbGciOiJub25lIn0" // {"alg":"none"} base64url encoded + payload := "eyJleHAiOjE4OTM0NTYwMDB9" // {"exp":1893456000} base64url encoded idToken := header + "." + payload + "." expiry := ParseIDTokenExpiry(idToken) diff --git a/internal/api/commandcode_credentials_test.go b/internal/api/commandcode_credentials_test.go index 3d1f6959..97a74c9c 100644 --- a/internal/api/commandcode_credentials_test.go +++ b/internal/api/commandcode_credentials_test.go @@ -5,6 +5,8 @@ import ( "path/filepath" "runtime" "testing" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) // isolateCommandCodeCredentials points every auth-file source at an empty temp @@ -12,7 +14,7 @@ import ( func isolateCommandCodeCredentials(t *testing.T) string { t.Helper() home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("COMMAND_CODE_API_KEY", "") t.Setenv("COMMANDCODE_API_KEY", "") t.Setenv("COMMANDCODE_AUTH_PATH", "") @@ -123,16 +125,23 @@ func TestDetectCommandCodeCredentialsNone(t *testing.T) { } func TestReadCommandCodeAuthFileRejectsPermissiveMode(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("Windows reports plain files as 0666; the unix permission check does not apply") - } home := isolateCommandCodeCredentials(t) path := filepath.Join(home, "auth.json") writeCommandCodeAuthFile(t, path, `{"apiKey":"user_secret"}`) if err := os.Chmod(path, 0o644); err != nil { t.Fatal(err) } - if got := readCommandCodeAuthFile(path); got != "" { + got := readCommandCodeAuthFile(path) + if runtime.GOOS == "windows" { + // Windows reports plain files as 0666 and guards them with ACLs that + // fs.FileMode cannot express, so the product deliberately skips the + // unix mode check there (commandCodeAuthFilePermsOK is a no-op). + if got != "user_secret" { + t.Fatalf("Windows must accept the auth file regardless of mode bits, got %q", got) + } + return + } + if got != "" { t.Fatalf("group/world-readable auth file must be ignored, got %q", got) } } diff --git a/internal/api/copilot_types.go b/internal/api/copilot_types.go index 7971d3bb..0fd4fc3a 100644 --- a/internal/api/copilot_types.go +++ b/internal/api/copilot_types.go @@ -31,10 +31,10 @@ type CopilotUserResponse struct { QuotaSnapshots map[string]*CopilotQuotaSnapshot `json:"quota_snapshots"` // New format fields (free_limited_copilot plans) - LimitedUserQuotas map[string]int `json:"limited_user_quotas"` - MonthlyQuotas map[string]int `json:"monthly_quotas"` - LimitedUserSubscribedDay int `json:"limited_user_subscribed_day"` - LimitedUserResetDate string `json:"limited_user_reset_date"` + LimitedUserQuotas map[string]int `json:"limited_user_quotas"` + MonthlyQuotas map[string]int `json:"monthly_quotas"` + LimitedUserSubscribedDay int `json:"limited_user_subscribed_day"` + LimitedUserResetDate string `json:"limited_user_reset_date"` } // normalize synthesizes QuotaSnapshots from the new limited_user_quotas/monthly_quotas diff --git a/internal/api/cursor_token_test.go b/internal/api/cursor_token_test.go index 553bc019..1839658c 100644 --- a/internal/api/cursor_token_test.go +++ b/internal/api/cursor_token_test.go @@ -2,6 +2,7 @@ package api import ( "log/slog" + "path/filepath" "testing" ) @@ -107,7 +108,9 @@ func TestCursorStateDBPathForOS(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - got := cursorStateDBPathForOS(home, tt.goos) + // filepath.Join uses the host separator, so compare in slash form: + // the layout under home is what is pinned here, on every host. + got := filepath.ToSlash(cursorStateDBPathForOS(home, tt.goos)) if got != tt.want { t.Fatalf("cursorStateDBPathForOS() = %q, want %q", got, tt.want) } diff --git a/internal/api/deepseek_client_test.go b/internal/api/deepseek_client_test.go index df6cbc7b..754f57d4 100644 --- a/internal/api/deepseek_client_test.go +++ b/internal/api/deepseek_client_test.go @@ -111,9 +111,9 @@ func TestDeepSeekTypes(t *testing.T) { if err != nil { t.Fatal(err) } - + snap := resp.ToSnapshot(time.Now()) - + if snap.Currency != "CNY" { t.Errorf("expected priority currency CNY, got %s", snap.Currency) } diff --git a/internal/api/deepseek_types.go b/internal/api/deepseek_types.go index 2c67d69c..60af75b1 100644 --- a/internal/api/deepseek_types.go +++ b/internal/api/deepseek_types.go @@ -8,9 +8,9 @@ import ( // DeepSeekBalanceInfo represents the balance data from DeepSeek API. type DeepSeekBalanceInfo struct { - Currency string `json:"currency"` - TotalBalance string `json:"total_balance"` - GrantedBalance string `json:"granted_balance"` + Currency string `json:"currency"` + TotalBalance string `json:"total_balance"` + GrantedBalance string `json:"granted_balance"` ToppedUpBalance string `json:"topped_up_balance"` } diff --git a/internal/api/extra_coverage_test.go b/internal/api/extra_coverage_test.go index e5eaae3c..f72e88c6 100644 --- a/internal/api/extra_coverage_test.go +++ b/internal/api/extra_coverage_test.go @@ -3,15 +3,19 @@ package api import ( "context" "encoding/json" + "errors" "fmt" "io" "net/http" "net/http/httptest" "os" "path/filepath" + "runtime" "strings" "testing" "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) // --------------------------------------------------------------------------- @@ -108,7 +112,7 @@ func TestDetectCodexCredentials_APIKeyOnly_ReturnsCredentials(t *testing.T) { // When only APIKey is set (no access_token), the credentials should be returned. home := t.TempDir() t.Setenv("CODEX_HOME", "") - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) isolateOpenCodeEnv(t) codexDir := filepath.Join(home, ".codex") @@ -138,7 +142,7 @@ func TestDetectCodexCredentials_BothEmpty_ReturnsNil(t *testing.T) { // When both access_token and OPENAI_API_KEY are empty, nil should be returned. home := t.TempDir() t.Setenv("CODEX_HOME", "") - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_TOKEN", "") isolateOpenCodeEnv(t) @@ -162,7 +166,7 @@ func TestDetectCodexCredentials_BothEmpty_ReturnsNil(t *testing.T) { func TestDetectCodexCredentials_InvalidJSON_ReturnsNil(t *testing.T) { home := t.TempDir() t.Setenv("CODEX_HOME", home) - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) t.Setenv("CODEX_TOKEN", "") isolateOpenCodeEnv(t) @@ -180,7 +184,7 @@ func TestDetectCodexCredentials_NoFile_ReturnsNil(t *testing.T) { // Set CODEX_HOME to a temp dir that has no auth.json home := t.TempDir() t.Setenv("CODEX_HOME", home) - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) t.Setenv("CODEX_TOKEN", "") isolateOpenCodeEnv(t) @@ -602,7 +606,7 @@ func TestCopilotToSnapshot_MultipleMixedQuotas(t *testing.T) { func TestDetectAnthropicToken_ReturnsStringOrEmpty(t *testing.T) { // When no credentials file exists, should return empty string without panic. home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Don't create any .claude directory - should return empty gracefully token := DetectAnthropicToken(nil) @@ -612,7 +616,7 @@ func TestDetectAnthropicToken_ReturnsStringOrEmpty(t *testing.T) { func TestDetectAnthropicCredentials_NoFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) creds := DetectAnthropicCredentials(nil) if creds != nil { @@ -627,7 +631,7 @@ func TestDetectAnthropicCredentials_NoFile(t *testing.T) { func TestWriteAnthropicCredentials_Success(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Create the .claude directory and a credentials file claudeDir := filepath.Join(home, ".claude") @@ -672,15 +676,34 @@ func TestWriteAnthropicCredentials_Success(t *testing.T) { func TestWriteAnthropicCredentials_NoFile(t *testing.T) { // No credentials file exists - on macOS/Linux this is OK because // Keychain/keyring is the primary store. File write is skipped. + // On Windows the file is the only store, so see assertMissingCredentialsFileResult. home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // Don't create .claude directory err := WriteAnthropicCredentials("token", "refresh", 3600) - // File not existing is OK (Keychain/keyring is primary). - if err != nil { + assertMissingCredentialsFileResult(t, home, err) +} + +// assertMissingCredentialsFileResult checks WriteAnthropicCredentials when +// ~/.claude/.credentials.json does not exist. On macOS/Linux the keychain or +// keyring is the primary store, so a missing file is skipped silently. On +// Windows the file is the only store Claude Code reads: a missing file means +// the rotated refresh token could not be persisted, which must surface as an +// error (the Anthropic agent logs it and guards against re-reading the stale +// token). In both cases no file may be created from scratch. +func assertMissingCredentialsFileResult(t *testing.T, home string, err error) { + t.Helper() + if runtime.GOOS == "windows" { + if !errors.Is(err, os.ErrNotExist) { + t.Errorf("err = %v, want a not-exist error on Windows", err) + } + } else if err != nil { t.Errorf("unexpected error: %v", err) } + if _, statErr := os.Stat(filepath.Join(home, ".claude", ".credentials.json")); !os.IsNotExist(statErr) { + t.Errorf("credentials file must not be created, stat err = %v", statErr) + } } // --------------------------------------------------------------------------- @@ -802,7 +825,7 @@ func TestAnthropicClient_SetAndGetToken(t *testing.T) { func TestDetectAnthropicToken_FromCredentialsFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) claudeDir := filepath.Join(home, ".claude") if err := os.MkdirAll(claudeDir, 0o755); err != nil { @@ -823,7 +846,7 @@ func TestDetectAnthropicToken_FromCredentialsFile(t *testing.T) { func TestDetectAnthropicToken_InvalidCredentialsFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) claudeDir := filepath.Join(home, ".claude") if err := os.MkdirAll(claudeDir, 0o755); err != nil { @@ -844,7 +867,7 @@ func TestDetectAnthropicToken_InvalidCredentialsFile(t *testing.T) { func TestDetectAnthropicToken_EmptyTokenInFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) claudeDir := filepath.Join(home, ".claude") if err := os.MkdirAll(claudeDir, 0o755); err != nil { @@ -870,7 +893,7 @@ func TestDetectAnthropicToken_EmptyTokenInFile(t *testing.T) { func TestDetectAnthropicCredentials_FromFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) claudeDir := filepath.Join(home, ".claude") if err := os.MkdirAll(claudeDir, 0o755); err != nil { @@ -897,7 +920,7 @@ func TestDetectAnthropicCredentials_FromFile(t *testing.T) { func TestDetectAnthropicCredentials_EmptyTokenInFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) claudeDir := filepath.Join(home, ".claude") if err := os.MkdirAll(claudeDir, 0o755); err != nil { @@ -918,7 +941,7 @@ func TestDetectAnthropicCredentials_EmptyTokenInFile(t *testing.T) { func TestDetectAnthropicCredentials_InvalidJSON(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) claudeDir := filepath.Join(home, ".claude") if err := os.MkdirAll(claudeDir, 0o755); err != nil { @@ -1138,7 +1161,7 @@ func TestCopilotToSnapshot_EmptyQuotaSnapshots(t *testing.T) { func TestDetectCodexToken_APIKeyOnly_ReturnsEmpty(t *testing.T) { home := t.TempDir() t.Setenv("CODEX_HOME", "") - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) codexDir := filepath.Join(home, ".codex") if err := os.MkdirAll(codexDir, 0o755); err != nil { @@ -1160,7 +1183,7 @@ func TestDetectCodexToken_APIKeyOnly_ReturnsEmpty(t *testing.T) { func TestDetectCodexToken_WithAccessToken(t *testing.T) { home := t.TempDir() t.Setenv("CODEX_HOME", "") - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) codexDir := filepath.Join(home, ".codex") if err := os.MkdirAll(codexDir, 0o755); err != nil { @@ -1184,7 +1207,7 @@ func TestDetectCodexToken_WithAccessToken(t *testing.T) { func TestDetectCodexCredentials_EmptyCodexHome_NoHomeDir(t *testing.T) { t.Setenv("CODEX_HOME", "") - t.Setenv("HOME", "") + testhome.SetTestHome(t, "") // codexAuthPath should return "" when HOME is unset // On macOS, os.UserHomeDir may still succeed, so we just verify no panic creds := DetectCodexCredentials(nil) @@ -1391,7 +1414,7 @@ func TestAntigravityClient_ResetClearsConnection(t *testing.T) { func TestDetectCodexToken_NilCreds_ReturnsEmpty(t *testing.T) { home := t.TempDir() t.Setenv("CODEX_HOME", "") - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) // No .codex/auth.json exists, so DetectCodexCredentials returns nil token := DetectCodexToken(nil) if token != "" { @@ -1681,7 +1704,7 @@ func TestAnthropicFetchQuotas_CreateRequestError(t *testing.T) { func TestWriteAnthropicCredentials_NoOAuthSection(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) claudeDir := filepath.Join(home, ".claude") if err := os.MkdirAll(claudeDir, 0o755); err != nil { @@ -1772,24 +1795,6 @@ func TestCodexQuotaSortOrder_Default(t *testing.T) { // getCredentialsFilePath - covers the home dir lookup // --------------------------------------------------------------------------- -func TestGetCredentialsFilePath_WithHome(t *testing.T) { - home := t.TempDir() - t.Setenv("HOME", home) - - path := getCredentialsFilePath() - expected := filepath.Join(home, ".claude", ".credentials.json") - if path != expected { - t.Errorf("getCredentialsFilePath() = %q, want %q", path, expected) - } -} - -func TestGetCredentialsFilePath_ReturnsNonEmpty(t *testing.T) { - // Regardless of platform, should return a non-empty path if home exists - path := getCredentialsFilePath() - // Could be empty in some edge cases, but should not panic - _ = path -} - // --------------------------------------------------------------------------- // codexAuthPath - cover HOME-based path // --------------------------------------------------------------------------- @@ -1808,7 +1813,7 @@ func TestCodexAuthPath_WithCODEX_HOME(t *testing.T) { func TestCodexAuthPath_WithoutCODEX_HOME(t *testing.T) { home := t.TempDir() t.Setenv("CODEX_HOME", "") - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) path := codexAuthPath() expected := filepath.Join(home, ".codex", "auth.json") @@ -2336,7 +2341,7 @@ func TestCodexAuthPath_EmptyHOME_ReturnsEmpty(t *testing.T) { // When CODEX_HOME is unset and HOME is empty, codexAuthPath returns "" // because os.UserHomeDir() returns an error when HOME is not set. t.Setenv("CODEX_HOME", "") - t.Setenv("HOME", "") + testhome.SetTestHome(t, "") path := codexAuthPath() if path != "" { @@ -2354,22 +2359,6 @@ func TestCodexAuthPath_EmptyHOME_ReturnsEmpty(t *testing.T) { // getCredentialsFilePath - HOME empty or error path // --------------------------------------------------------------------------- -func TestGetCredentialsFilePath_EmptyHOME(t *testing.T) { - // When HOME is not set, getCredentialsFilePath may return "" or - // use user.Current() as a fallback. Either way it must not panic. - t.Setenv("HOME", "") - - path := getCredentialsFilePath() - // The function returns "" or a valid path via user.Current() fallback. - // We just verify no panic and correct format if non-empty. - if path != "" { - // path should end with .claude/.credentials.json - if !strings.HasSuffix(path, ".credentials.json") { - t.Errorf("getCredentialsFilePath() = %q, should end with .credentials.json", path) - } - } -} - // --------------------------------------------------------------------------- // detectAnthropicTokenPlatform - empty home path // --------------------------------------------------------------------------- @@ -2377,7 +2366,7 @@ func TestGetCredentialsFilePath_EmptyHOME(t *testing.T) { func TestDetectAnthropicTokenPlatform_EmptyHOME_ReturnsEmpty(t *testing.T) { // When HOME is unset and platform keychain lookups fail, the function // logs "Cannot determine home directory" and returns "". - t.Setenv("HOME", "") + testhome.SetTestHome(t, "") // This will attempt keychain (which will likely fail), then try to read // the credentials file. With HOME="", os.UserHomeDir() returns an error, @@ -2456,14 +2445,13 @@ func TestWriteAnthropicCredentials_FileNotFound(t *testing.T) { // has no .claude/.credentials.json. On macOS/Linux, this is OK because // Keychain/keyring is the primary store - file write is skipped silently. dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) err := WriteAnthropicCredentials("access_token", "refresh_token", 3600) // File not existing is OK (Keychain/keyring is primary on macOS/Linux). - // writeCredentialsToFile returns nil when file doesn't exist. - if err != nil { - t.Errorf("unexpected error: %v", err) - } + // writeCredentialsToFile returns nil when file doesn't exist. On Windows + // the file is the only store, so the write must report not-exist. + assertMissingCredentialsFileResult(t, dir, err) } // --------------------------------------------------------------------------- @@ -2642,7 +2630,7 @@ func TestAntigravityClient_FetchQuotas_500Response(t *testing.T) { func TestDetectAnthropicTokenPlatform_MalformedCredentialsFile(t *testing.T) { dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) claudeDir := filepath.Join(dir, ".claude") if err := os.MkdirAll(claudeDir, 0755); err != nil { @@ -2664,7 +2652,7 @@ func TestDetectAnthropicTokenPlatform_MalformedCredentialsFile(t *testing.T) { func TestDetectAnthropicTokenPlatform_EmptyAccessToken(t *testing.T) { dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) claudeDir := filepath.Join(dir, ".claude") if err := os.MkdirAll(claudeDir, 0755); err != nil { @@ -2691,7 +2679,7 @@ func TestDetectAnthropicTokenPlatform_EmptyAccessToken(t *testing.T) { func TestDetectAnthropicCredentialsPlatform_MalformedJSON(t *testing.T) { dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) claudeDir := filepath.Join(dir, ".claude") if err := os.MkdirAll(claudeDir, 0755); err != nil { @@ -2716,7 +2704,7 @@ func TestDetectAnthropicCredentialsPlatform_MalformedJSON(t *testing.T) { func TestDetectAnthropicCredentialsPlatform_NoOAuthSection(t *testing.T) { dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) claudeDir := filepath.Join(dir, ".claude") if err := os.MkdirAll(claudeDir, 0755); err != nil { @@ -2741,7 +2729,7 @@ func TestDetectAnthropicCredentialsPlatform_NoOAuthSection(t *testing.T) { func TestWriteAnthropicCredentials_CreatesBackup(t *testing.T) { dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) claudeDir := filepath.Join(dir, ".claude") if err := os.MkdirAll(claudeDir, 0755); err != nil { @@ -2787,7 +2775,7 @@ func TestDetectCodexCredentials_EmptyAuthFile(t *testing.T) { isolateOpenCodeEnv(t) dir := t.TempDir() t.Setenv("CODEX_HOME", dir) - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) // Write an auth file with all empty fields authData := `{"OPENAI_API_KEY":"","tokens":{"access_token":"","refresh_token":"","id_token":"","account_id":""}}` @@ -3019,7 +3007,7 @@ func TestAntigravityToSnapshot_ModelWithNilQuotaInfo(t *testing.T) { func TestDetectAnthropicCredentialsPlatform_ValidCredentials(t *testing.T) { dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) claudeDir := filepath.Join(dir, ".claude") if err := os.MkdirAll(claudeDir, 0755); err != nil { @@ -3053,7 +3041,7 @@ func TestDetectAnthropicCredentialsPlatform_ValidCredentials(t *testing.T) { func TestDetectAnthropicTokenPlatform_ValidFile_ReturnsToken(t *testing.T) { dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) claudeDir := filepath.Join(dir, ".claude") if err := os.MkdirAll(claudeDir, 0755); err != nil { @@ -3080,7 +3068,7 @@ func TestDetectAnthropicTokenPlatform_ValidFile_ReturnsToken(t *testing.T) { func TestWriteAnthropicCredentials_InvalidJSONFile(t *testing.T) { dir := t.TempDir() - t.Setenv("HOME", dir) + testhome.SetTestHome(t, dir) claudeDir := filepath.Join(dir, ".claude") if err := os.MkdirAll(claudeDir, 0755); err != nil { @@ -3102,17 +3090,6 @@ func TestWriteAnthropicCredentials_InvalidJSONFile(t *testing.T) { // getCredentialsFilePath - HOME set to temp dir (covers normal path fully) // --------------------------------------------------------------------------- -func TestGetCredentialsFilePath_ValidHOME(t *testing.T) { - dir := t.TempDir() - t.Setenv("HOME", dir) - - path := getCredentialsFilePath() - expected := filepath.Join(dir, ".claude", ".credentials.json") - if path != expected { - t.Errorf("getCredentialsFilePath() = %q, want %q", path, expected) - } -} - // --------------------------------------------------------------------------- // AntigravityClient - FetchQuotas with context cancel resets connection // --------------------------------------------------------------------------- diff --git a/internal/api/gemini_types.go b/internal/api/gemini_types.go index ef135065..ee48106c 100644 --- a/internal/api/gemini_types.go +++ b/internal/api/gemini_types.go @@ -21,9 +21,9 @@ type GeminiQuotaResponse struct { // GeminiTierResponse is the response from loadCodeAssist. type GeminiTierResponse struct { - Tier string `json:"tier"` - CloudAICompanionProject string `json:"cloudaicompanionProject"` - PlanName string `json:"planName,omitempty"` + Tier string `json:"tier"` + CloudAICompanionProject string `json:"cloudaicompanionProject"` + PlanName string `json:"planName,omitempty"` } // GeminiQuota is a normalized per-model quota for storage. @@ -168,11 +168,11 @@ func AggregateGeminiByFamily(quotas []GeminiQuota) []GeminiFamilyQuota { // geminiDisplayNames maps model IDs to human-readable labels. var geminiDisplayNames = map[string]string{ - "gemini-2.5-pro": "Gemini 2.5 Pro", - "gemini-2.5-flash": "Gemini 2.5 Flash", - "gemini-2.5-flash-lite": "Gemini 2.5 Flash Lite", - "gemini-3-pro-preview": "Gemini 3 Pro", - "gemini-3-flash-preview": "Gemini 3 Flash", + "gemini-2.5-pro": "Gemini 2.5 Pro", + "gemini-2.5-flash": "Gemini 2.5 Flash", + "gemini-2.5-flash-lite": "Gemini 2.5 Flash Lite", + "gemini-3-pro-preview": "Gemini 3 Pro", + "gemini-3-flash-preview": "Gemini 3 Flash", "gemini-3.1-flash-lite-preview": "Gemini 3.1 Flash Lite", } diff --git a/internal/api/grok_client.go b/internal/api/grok_client.go index 9730fa57..3878620d 100644 --- a/internal/api/grok_client.go +++ b/internal/api/grok_client.go @@ -14,6 +14,7 @@ import ( "os" "os/exec" "path/filepath" + "runtime" "strconv" "strings" "sync" @@ -188,11 +189,7 @@ func (c *GrokClient) tryRPC(ctx context.Context, creds *GrokCredentials) (*GrokB resolved, err := exec.LookPath(bin) if err != nil { // Try common install locations quickly - for _, cand := range []string{ - filepath.Join(os.Getenv("HOME"), ".local", "bin", "grok"), - "/usr/local/bin/grok", - "/opt/homebrew/bin/grok", - } { + for _, cand := range grokBinaryFallbackPaths() { if _, statErr := os.Stat(cand); statErr == nil { resolved = cand break @@ -580,8 +577,8 @@ type ProtobufScanGo struct { order int } varints []struct { - path []uint64 - val uint64 + path []uint64 + val uint64 } } @@ -611,8 +608,8 @@ func scanProtobufGo(data []byte, depth int, path []uint64, order int) (ProtobufS case 0: if v, ok := readVarintGo(b, &idx); ok { scan.varints = append(scan.varints, struct { - path []uint64 - val uint64 + path []uint64 + val uint64 }{fpath, v}) } else { idx = start + 1 @@ -683,6 +680,22 @@ func (c *GrokClient) scanLocalSessions() *GrokLocalSessionSummary { return scanGrokSessionsDir(root, time.Now().AddDate(0, 0, -30)) } +// grokBinaryFallbackPaths lists install locations to try when grok is not on +// PATH. The per-user location comes from os.UserHomeDir (USERPROFILE on +// Windows, where HOME is normally unset) and is skipped when no home is known, +// so an empty home never turns into a cwd-relative ".local/bin/grok" lookup. +func grokBinaryFallbackPaths() []string { + name := "grok" + if runtime.GOOS == "windows" { + name = "grok.exe" + } + var paths []string + if home, err := os.UserHomeDir(); err == nil && home != "" { + paths = append(paths, filepath.Join(home, ".local", "bin", name)) + } + return append(paths, "/usr/local/bin/grok", "/opt/homebrew/bin/grok") +} + func GrokHomeDir() string { if h := strings.TrimSpace(os.Getenv("GROK_HOME")); h != "" { return h diff --git a/internal/api/grok_client_test.go b/internal/api/grok_client_test.go index b113c647..b1be1977 100644 --- a/internal/api/grok_client_test.go +++ b/internal/api/grok_client_test.go @@ -8,8 +8,11 @@ import ( "net/http/httptest" "os" "path/filepath" + "strings" "testing" "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func TestNewGrokClient_Basic(t *testing.T) { @@ -64,7 +67,7 @@ func buildTestPayloadForNoUsageYet(t *testing.T) []byte { // Field 1 (len) containing sub with field 6 (varint 1) and field 5 len containing field 1 (varint future ts) // Rough wire that triggers hasUsagePeriod + future reset at preferred path. // Use the frame builder from the code paths. - future := uint64(time.Now().Add(24*time.Hour).Unix()) + future := uint64(time.Now().Add(24 * time.Hour).Unix()) // Build a tiny message: 1:{ 6: varint(1), 5: {1: varint(future)} } inner5 := appendVarint(nil, (1<<3)|0, future) inner5field := appendLenField(nil, 5, inner5) @@ -157,3 +160,26 @@ func TestRPC_NoBinary(t *testing.T) { t.Error("expected binary not found err") } } + +func TestGrokBinaryFallbackPaths_UsesUserHome(t *testing.T) { + home := t.TempDir() + testhome.SetTestHome(t, home) + + paths := grokBinaryFallbackPaths() + if len(paths) == 0 { + t.Fatal("expected fallback paths") + } + if dir := filepath.Dir(paths[0]); dir != filepath.Join(home, ".local", "bin") { + t.Fatalf("per-user fallback dir = %q, want %q", dir, filepath.Join(home, ".local", "bin")) + } +} + +func TestGrokBinaryFallbackPaths_NoHomeSkipsRelativePath(t *testing.T) { + testhome.SetTestHome(t, "") + + for _, p := range grokBinaryFallbackPaths() { + if strings.Contains(filepath.ToSlash(p), ".local/bin") { + t.Fatalf("empty home must not produce a cwd-relative candidate, got %q", p) + } + } +} diff --git a/internal/api/grok_types.go b/internal/api/grok_types.go index 1da6c7b9..49bace1d 100644 --- a/internal/api/grok_types.go +++ b/internal/api/grok_types.go @@ -27,12 +27,12 @@ type GrokBillingUsage struct { // GrokBillingResponse is the shape returned by `x.ai/billing` RPC (and synthesized from web probe). // All monetary are cents via GrokCent. type GrokBillingResponse struct { - BillingCycle *GrokBillingCycle `json:"billingCycle"` - MonthlyLimit *GrokCent `json:"monthlyLimit"` - OnDemandCap *GrokCent `json:"onDemandCap"` - OnDemandEnabled *bool `json:"on_demand_enabled"` - DisabledByConfig *bool `json:"disabledByConfig"` - Usage *GrokBillingUsage `json:"usage"` + BillingCycle *GrokBillingCycle `json:"billingCycle"` + MonthlyLimit *GrokCent `json:"monthlyLimit"` + OnDemandCap *GrokCent `json:"onDemandCap"` + OnDemandEnabled *bool `json:"on_demand_enabled"` + DisabledByConfig *bool `json:"disabledByConfig"` + Usage *GrokBillingUsage `json:"usage"` } // GrokWebBillingSnapshot is the normalized result from the gRPC-web fallback. @@ -60,14 +60,14 @@ type GrokQuota struct { // GrokSnapshot is the storage + UI representation. type GrokSnapshot struct { - ID int64 - CapturedAt time.Time - AccountID int64 // default 1 for single-account - Email string - TeamID string - LoginMethod string - Quotas []GrokQuota - RawJSON string + ID int64 + CapturedAt time.Time + AccountID int64 // default 1 for single-account + Email string + TeamID string + LoginMethod string + Quotas []GrokQuota + RawJSON string // LocalSessions is informational fallback data (may be nil). LocalSessions *GrokLocalSessionSummary } diff --git a/internal/api/kimi_client_test.go b/internal/api/kimi_client_test.go index 9360a665..d283c0e7 100644 --- a/internal/api/kimi_client_test.go +++ b/internal/api/kimi_client_test.go @@ -13,6 +13,8 @@ import ( "sync/atomic" "testing" "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func TestKimiClientFetchSnapshot(t *testing.T) { @@ -99,7 +101,7 @@ func TestKimiClientFetchSnapshot_ForceRefreshOn401UnexpiredAccess(t *testing.T) defer srv.Close() home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("KIMI_CODE_HOME", "") t.Setenv("KIMI_CODE_CREDENTIALS", "") t.Setenv("KIMI_CREDENTIALS", "") @@ -166,7 +168,7 @@ func kimiUsagesOK(w http.ResponseWriter, used string) { func setupKimiCodeCreds(t *testing.T, access, refresh string, expiresAt float64) string { t.Helper() home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("KIMI_CODE_HOME", "") t.Setenv("KIMI_CODE_CREDENTIALS", "") t.Setenv("KIMI_CREDENTIALS", "") diff --git a/internal/api/kimi_credentials_test.go b/internal/api/kimi_credentials_test.go index d1dc12d8..c6a69426 100644 --- a/internal/api/kimi_credentials_test.go +++ b/internal/api/kimi_credentials_test.go @@ -6,6 +6,8 @@ import ( "path/filepath" "testing" "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func writeKimiCred(t *testing.T, dir string, access, refresh string, expiresAt float64) string { @@ -32,7 +34,7 @@ func writeKimiCred(t *testing.T, dir string, access, refresh string, expiresAt f func TestDetectKimiCredentials_KimiCodeOnly(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("KIMI_CODE_CREDENTIALS", "") t.Setenv("KIMI_CREDENTIALS", "") t.Setenv("KIMI_CODE_HOME", "") @@ -63,7 +65,7 @@ func TestDetectKimiCredentials_KimiCodeOnly(t *testing.T) { func TestDetectKimiCredentials_IgnoresKimiCLIAlone(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("KIMI_CODE_CREDENTIALS", "") t.Setenv("KIMI_CREDENTIALS", "") t.Setenv("KIMI_CODE_HOME", filepath.Join(home, "no-code")) @@ -83,7 +85,7 @@ func TestDetectKimiCredentials_IgnoresKimiCLIAlone(t *testing.T) { func TestDetectKimiCredentials_ExplicitEnvFile(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) path := writeKimiCred(t, filepath.Join(home, "custom"), "env-access", "env-refresh", float64(time.Now().Unix()+3600)) t.Setenv("KIMI_CODE_CREDENTIALS", path) t.Setenv("KIMI_CODE_HOME", "") @@ -110,7 +112,7 @@ func TestKimiCredentials_ExpiredSkew(t *testing.T) { func TestLoadKimiCredentialsCached_ReloadsWhenFileChanges(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("KIMI_CODE_HOME", "") t.Setenv("KIMI_CODE_CREDENTIALS", "") t.Setenv("KIMI_CREDENTIALS", "") diff --git a/internal/api/kimi_types.go b/internal/api/kimi_types.go index 8dc6eb1b..d7b2303a 100644 --- a/internal/api/kimi_types.go +++ b/internal/api/kimi_types.go @@ -11,21 +11,21 @@ import ( // KimiUsagesResponse is the JSON shape from GET /coding/v1/usages // (same endpoint used by the official kimi-code CLI). type KimiUsagesResponse struct { - User *KimiUser `json:"user"` - Usage *KimiUsageDetail `json:"usage"` - Limits []KimiWindowLimit `json:"limits"` - Parallel *KimiParallel `json:"parallel"` - Total *KimiUsageDetail `json:"totalQuota"` - Auth *KimiAuthentication `json:"authentication"` - SubType string `json:"subType"` + User *KimiUser `json:"user"` + Usage *KimiUsageDetail `json:"usage"` + Limits []KimiWindowLimit `json:"limits"` + Parallel *KimiParallel `json:"parallel"` + Total *KimiUsageDetail `json:"totalQuota"` + Auth *KimiAuthentication `json:"authentication"` + SubType string `json:"subType"` } // KimiUser holds identity/membership metadata from the usages endpoint. type KimiUser struct { - UserID string `json:"userId"` - Region string `json:"region"` - BusinessID string `json:"businessId"` - Membership *KimiMembership `json:"membership"` + UserID string `json:"userId"` + Region string `json:"region"` + BusinessID string `json:"businessId"` + Membership *KimiMembership `json:"membership"` } // KimiMembership holds plan level (e.g. LEVEL_INTERMEDIATE). @@ -43,7 +43,7 @@ type KimiUsageDetail struct { // KimiWindowLimit is a time-windowed rate limit (e.g. 300 minutes → 5h). type KimiWindowLimit struct { - Window *KimiWindow `json:"window"` + Window *KimiWindow `json:"window"` Detail *KimiUsageDetail `json:"detail"` } diff --git a/internal/api/kimi_types_test.go b/internal/api/kimi_types_test.go index 75d0d3aa..a03516ab 100644 --- a/internal/api/kimi_types_test.go +++ b/internal/api/kimi_types_test.go @@ -73,13 +73,13 @@ func TestKimiDisplayName(t *testing.T) { func TestKimiMembershipDisplayName(t *testing.T) { cases := map[string]string{ - "LEVEL_FREE": "Free", - "LEVEL_BASIC": "Adagio", - "LEVEL_STANDARD": "Moderato", - "LEVEL_INTERMEDIATE": "Allegretto", - "LEVEL_ADVANCED": "Allegro", - "LEVEL_PREMIUM": "Vivace", - "": "", + "LEVEL_FREE": "Free", + "LEVEL_BASIC": "Adagio", + "LEVEL_STANDARD": "Moderato", + "LEVEL_INTERMEDIATE": "Allegretto", + "LEVEL_ADVANCED": "Allegro", + "LEVEL_PREMIUM": "Vivace", + "": "", "LEVEL_UNKNOWN_FUTURE": "LEVEL_UNKNOWN_FUTURE", } for in, want := range cases { diff --git a/internal/api/minimax_types.go b/internal/api/minimax_types.go index 4f65935f..80f76b7f 100644 --- a/internal/api/minimax_types.go +++ b/internal/api/minimax_types.go @@ -54,9 +54,11 @@ func clampPercent(v int) int { // minimaxIntervalActive reports whether a percentage-based quota window belongs // to an active/subscribed plan whose percentage should be tracked. MiniMax uses: -// status 1 = active with quota remaining -// status 2 = active but exhausted (0% remaining = 100% used) -// status 3 = model not part of the subscription +// +// status 1 = active with quota remaining +// status 2 = active but exhausted (0% remaining = 100% used) +// status 3 = model not part of the subscription +// // Both 1 and 2 are live windows that must be recorded - status 2 is exactly when // the user has hit their limit and most needs the reading; status 3 is dropped. func minimaxIntervalActive(status *int) bool { diff --git a/internal/api/mistral_cookie_test.go b/internal/api/mistral_cookie_test.go index fbf4eddd..a306511d 100644 --- a/internal/api/mistral_cookie_test.go +++ b/internal/api/mistral_cookie_test.go @@ -296,3 +296,15 @@ func TestMistralScanErrorOmitsPath(t *testing.T) { t.Fatalf("error leaks a filesystem path: %v", err) } } + +// Browser stores are opened through a file: URI. On Windows the drive letter +// must follow a slash, or SQLite reads "C:" as the URI authority and fails. +func TestReadOnlySQLiteURI(t *testing.T) { + path, want := "/tmp/a b/Cookies", "file:///tmp/a%20b/Cookies?mode=ro" + if runtime.GOOS == "windows" { + path, want = `C:\Users\a b\Cookies`, "file:///C:/Users/a%20b/Cookies?mode=ro" + } + if got := readOnlySQLiteURI(path); got != want { + t.Fatalf("readOnlySQLiteURI(%q)=%q, want %q", path, got, want) + } +} diff --git a/internal/api/mistral_scope.go b/internal/api/mistral_scope.go index b2b5fcff..fadac877 100644 --- a/internal/api/mistral_scope.go +++ b/internal/api/mistral_scope.go @@ -8,6 +8,7 @@ import ( "io" "net/url" "os" + "path/filepath" "strconv" "strings" @@ -71,8 +72,7 @@ func readMistralScopes(ctx context.Context, path string, browser sweetcookie.Bro if browser == sweetcookie.BrowserSafari { return readMistralSafariScopes(ctx, path) } - u := url.URL{Scheme: "file", Path: path, RawQuery: "mode=ro"} - db, e := sql.Open("sqlite", u.String()) + db, e := sql.Open("sqlite", readOnlySQLiteURI(path)) if e != nil { return nil, e } @@ -104,6 +104,17 @@ func readMistralScopes(ctx context.Context, path string, browser sweetcookie.Bro } return scopes, rows.Err() } + +// readOnlySQLiteURI opens a browser cookie store read-only. SQLite URIs take +// forward slashes, and a Windows drive path needs a leading slash +// (file:///C:/...), or the drive letter is read as the URI authority. +func readOnlySQLiteURI(path string) string { + p := filepath.ToSlash(path) + if !strings.HasPrefix(p, "/") { + p = "/" + p + } + return (&url.URL{Scheme: "file", Path: p, RawQuery: "mode=ro"}).String() +} func readMistralSafariScopes(ctx context.Context, path string) (map[string]string, error) { f, e := os.Open(path) if e != nil { diff --git a/internal/api/muse_credentials_cache_test.go b/internal/api/muse_credentials_cache_test.go index 77e1ae56..3e57eb7a 100644 --- a/internal/api/muse_credentials_cache_test.go +++ b/internal/api/muse_credentials_cache_test.go @@ -6,6 +6,8 @@ import ( "context" "path/filepath" "testing" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) // isolateMuseCredentials points every Muse credential source at an empty temp @@ -13,7 +15,7 @@ import ( func isolateMuseCredentials(t *testing.T) { t.Helper() home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("META_API_KEY", "") t.Setenv("MUSE_AUTH_PATH", filepath.Join(home, "missing-auth.json")) // MuseSettingsPath checks XDG_CONFIG_HOME before HOME, so without this diff --git a/internal/api/muse_credentials_test.go b/internal/api/muse_credentials_test.go index 0f3b03c8..7aebe220 100644 --- a/internal/api/muse_credentials_test.go +++ b/internal/api/muse_credentials_test.go @@ -4,6 +4,7 @@ import ( "log/slog" "os" "path/filepath" + "runtime" "testing" ) @@ -49,11 +50,26 @@ func TestDetectMuseCredentialsFromAuthFile(t *testing.T) { func TestMuseAuthFileRejectsPermissiveMode(t *testing.T) { dir := t.TempDir() authPath := filepath.Join(dir, "auth.json") - if err := os.WriteFile(authPath, []byte(`{"providers":{"meta":{"api_key":"x"}}}`), 0o644); err != nil { + if err := os.WriteFile(authPath, []byte(`{"providers":{"meta":{"api_key":"x"}}}`), 0o600); err != nil { + t.Fatal(err) + } + // Chmod explicitly so the mode does not depend on the process umask. + if err := os.Chmod(authPath, 0o644); err != nil { t.Fatal(err) } t.Setenv("MUSE_AUTH_PATH", authPath) - if got := readMuseAuthFileKey(); got != "" { + got := readMuseAuthFileKey() + if runtime.GOOS == "windows" { + // Windows has no unix mode bits: Go reports every writable file as + // 0666, and access is governed by ACLs that fs.FileMode cannot express. + // A mode check there would reject every `muse login` file, so the + // product deliberately trusts the file (museAuthFilePermsOK is a no-op). + if got != "x" { + t.Fatalf("Windows must accept the login file regardless of mode bits, got %q", got) + } + return + } + if got != "" { t.Fatalf("permissive auth file must be ignored, got %q", got) } } diff --git a/internal/api/opencode_client.go b/internal/api/opencode_client.go index 265456f6..1dfd4180 100644 --- a/internal/api/opencode_client.go +++ b/internal/api/opencode_client.go @@ -4,25 +4,14 @@ import ( "context" "errors" "fmt" - "io" "log/slog" "net/http" - "net/url" - "regexp" - "strconv" "strings" "sync" "time" ) -const ( - openCodeDefaultBaseURL = "https://opencode.ai" - openCodeDashboardURLPrefix = openCodeDefaultBaseURL + "/workspace/" - openCodeDashboardURLSuffix = "/go" - openCodeUserAgent = "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) Gecko/20100101 Firefox/148.0" - openCodeScrapeTimeout = 10 * time.Second - openCodeMaxBodyBytes = 2 << 20 // 2 MiB -) +const openCodeDefaultBaseURL = "https://opencode.ai" var ( ErrOpenCodeUnauthorized = errors.New("opencode: unauthorized") @@ -31,19 +20,17 @@ var ( ErrOpenCodeNetworkError = errors.New("opencode: network error") ErrOpenCodeInvalidResponse = errors.New("opencode: invalid response") ErrOpenCodeParseFailed = errors.New("opencode: parse failed") - ErrOpenCodeMissingConfig = errors.New("opencode: missing usage api key, or workspace id and auth cookie") + ErrOpenCodeMissingConfig = errors.New("opencode: missing usage api key, or workspace id and session cookie") ) type OpenCodeClient struct { - httpClient *http.Client - logger *slog.Logger - dashboardURLPrefix string - goStatusURL string - usageHTTPClient *http.Client // longer timeout; derived from httpClient after options + httpClient *http.Client + logger *slog.Logger + goStatusURL string usageMu sync.Mutex usageStatus *openCodeGoStatus // last good go/status, reused for openCodeUsageMinInterval - usageKey string + usageKey string // credential the cached status belongs to usageAt time.Time } @@ -64,7 +51,6 @@ func WithOpenCodeTimeout(timeout time.Duration) OpenCodeClientOption { func WithOpenCodeBaseURL(baseURL string) OpenCodeClientOption { return func(c *OpenCodeClient) { baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/") - c.dashboardURLPrefix = baseURL + "/workspace/" c.goStatusURL = baseURL + openCodeGoStatusPath } } @@ -75,318 +61,63 @@ func NewOpenCodeClient(logger *slog.Logger, opts ...OpenCodeClientOption) *OpenC } c := &OpenCodeClient{ httpClient: &http.Client{ - Timeout: openCodeScrapeTimeout, + Timeout: openCodeUsageTimeout, + // A redirect is the console sending an unauthenticated request to + // login; never follow it with the session cookie attached. CheckRedirect: func(_ *http.Request, _ []*http.Request) error { return http.ErrUseLastResponse }, Transport: &http.Transport{ MaxIdleConns: 1, MaxIdleConnsPerHost: 1, - ResponseHeaderTimeout: openCodeScrapeTimeout, + ResponseHeaderTimeout: openCodeUsageTimeout, IdleConnTimeout: 30 * time.Second, TLSHandshakeTimeout: 10 * time.Second, ForceAttemptHTTP2: true, }, }, - logger: logger, - dashboardURLPrefix: openCodeDashboardURLPrefix, - goStatusURL: openCodeDefaultBaseURL + openCodeGoStatusPath, + logger: logger, + goStatusURL: openCodeDefaultBaseURL + openCodeGoStatusPath, } for _, o := range opts { o(c) } - usage := *c.httpClient - if usage.Timeout < openCodeUsageTimeout { - usage.Timeout = openCodeUsageTimeout - } - if t, ok := usage.Transport.(*http.Transport); ok { - t = t.Clone() - t.ResponseHeaderTimeout = openCodeUsageTimeout - usage.Transport = t - } - c.usageHTTPClient = &usage return c } -type scrapedWindowUsage struct { - usagePercent float64 - resetInSec float64 -} - -func (c *OpenCodeClient) FetchSnapshot(ctx context.Context, workspaceID, authCookie string) (*OpenCodeSnapshot, error) { +// FetchSnapshot reads the Go meters with the browser's console session. The +// console authenticates with the __Host-console_session cookie and needs the +// workspace ID as x-org-id; the old auth cookie and the /workspace//go +// page no longer work (issue #134). +func (c *OpenCodeClient) FetchSnapshot(ctx context.Context, workspaceID, sessionCookie string) (*OpenCodeSnapshot, error) { workspaceID = strings.TrimSpace(workspaceID) - authCookie = strings.TrimSpace(authCookie) - if workspaceID == "" || authCookie == "" { + sessionCookie = strings.TrimSpace(sessionCookie) + if workspaceID == "" || sessionCookie == "" { return nil, ErrOpenCodeMissingConfig } - - capturedAt := time.Now().UTC() - html, err := c.fetchDashboardHTML(ctx, workspaceID, authCookie) - if err != nil { - return nil, err - } - - quotas, err := parseOpenCodeQuotas(html, capturedAt) - if err != nil { - return nil, err - } - - return &OpenCodeSnapshot{ - CapturedAt: capturedAt, - AccountType: OpenCodeAccountTypePro, - PlanName: "OpenCode Go", - Quotas: quotas, - }, nil -} - -func (c *OpenCodeClient) fetchDashboardHTML(ctx context.Context, workspaceID, authCookie string) (string, error) { - dashboardURL := c.dashboardURLPrefix + url.PathEscape(workspaceID) + openCodeDashboardURLSuffix - - req, err := http.NewRequestWithContext(ctx, http.MethodGet, dashboardURL, nil) - if err != nil { - return "", fmt.Errorf("%w: build request: %v", ErrOpenCodeNetworkError, err) - } - req.Header.Set("User-Agent", openCodeUserAgent) - req.Header.Set("Accept", "text/html") - req.Header.Set("Cookie", openCodeAuthCookieHeader(authCookie)) - - resp, err := c.httpClient.Do(req) - if err != nil { - if ctx.Err() != nil { - return "", ctx.Err() - } - return "", fmt.Errorf("%w: %v", ErrOpenCodeNetworkError, err) - } - defer resp.Body.Close() - - body, err := io.ReadAll(io.LimitReader(resp.Body, openCodeMaxBodyBytes)) - if err != nil { - return "", fmt.Errorf("%w: read body: %v", ErrOpenCodeNetworkError, err) - } - if resp.StatusCode >= http.StatusMultipleChoices && resp.StatusCode < http.StatusBadRequest { - return "", ErrOpenCodeUnauthorized - } - - switch resp.StatusCode { - case http.StatusOK: - return string(body), nil - case http.StatusUnauthorized: - return "", ErrOpenCodeUnauthorized - case http.StatusForbidden: - return "", ErrOpenCodeForbidden - default: - if resp.StatusCode >= 500 { - return "", fmt.Errorf("%w: http %d", ErrOpenCodeServerError, resp.StatusCode) - } - return "", fmt.Errorf("%w: http %d: %s", ErrOpenCodeInvalidResponse, resp.StatusCode, sanitizeOpenCodeMessage(string(body))) - } -} - -func openCodeAuthCookieHeader(authCookie string) string { - if strings.HasPrefix(authCookie, "auth=") { - return authCookie - } - return "auth=" + authCookie -} - -func sanitizeOpenCodeMessage(text string) string { - s := strings.TrimSpace(text) - if s == "" { - return "unknown" - } - s = strings.Join(strings.Fields(s), " ") - if len(s) > 120 { - s = s[:120] - } - return s -} - -func parseOpenCodeQuotas(html string, capturedAt time.Time) ([]OpenCodeQuota, error) { - rolling := parseSSRWindowUsage(html, "rollingUsage") - weekly := parseSSRWindowUsage(html, "weeklyUsage") - monthly := parseSSRWindowUsage(html, "monthlyUsage") - - if rolling == nil && weekly == nil && monthly == nil { - dataSlot := parseDataSlotFormat(html) - rolling = dataSlot["rolling"] - weekly = dataSlot["weekly"] - monthly = dataSlot["monthly"] - } - - var quotas []OpenCodeQuota - if rolling != nil { - quotas = append(quotas, windowToQuota("five_hour", *rolling, capturedAt)) - } - if weekly != nil { - quotas = append(quotas, windowToQuota("weekly", *weekly, capturedAt)) + cookieHeader := openCodeConsoleCookieHeader(sessionCookie) + snap, err := c.goStatusSnapshot(ctx, "cookie\x00"+workspaceID+"\x00"+cookieHeader, func(req *http.Request) { + req.Header.Set("Cookie", cookieHeader) + req.Header.Set("x-org-id", workspaceID) + }) + if errors.Is(err, ErrOpenCodeUnauthorized) { + return nil, fmt.Errorf("%w: paste the __Host-console_session cookie from opencode.ai/console", err) } - if monthly != nil { - quotas = append(quotas, windowToQuota("monthly", *monthly, capturedAt)) - } - - if len(quotas) == 0 { - return nil, fmt.Errorf("%w: could not parse rollingUsage, weeklyUsage, or monthlyUsage", ErrOpenCodeParseFailed) - } - return quotas, nil -} - -func windowToQuota(name string, window scrapedWindowUsage, capturedAt time.Time) OpenCodeQuota { - pct := window.usagePercent - if pct < 0 { - pct = 0 - } - resetSec := window.resetInSec - if resetSec < 0 { - resetSec = 0 - } - resetsAt := capturedAt.Add(time.Duration(resetSec) * time.Second) - return OpenCodeQuota{ - Name: name, - Used: pct, - Limit: 100, - Utilization: pct, - Format: OpenCodeQuotaFormatPercent, - ResetsAt: &resetsAt, - } -} - -var openCodeScrapedNumberPattern = `(-?\d+(?:\.\d+)?)` - -func parseSSRWindowUsage(html, prefix string) *scrapedWindowUsage { - pattern1 := prefix + `:\$R\[\d+\]=\{[^}]*usagePercent:` + openCodeScrapedNumberPattern + `[^}]*resetInSec:` + openCodeScrapedNumberPattern - pattern2 := prefix + `:\$R\[\d+\]=\{[^}]*resetInSec:` + openCodeScrapedNumberPattern + `[^}]*usagePercent:` + openCodeScrapedNumberPattern - - re1 := regexp.MustCompile(pattern1) - re2 := regexp.MustCompile(pattern2) - - if m := re1.FindStringSubmatch(html); len(m) >= 3 { - if pct, reset, ok := parseScrapedNumbers(m[1], m[2]); ok { - return &scrapedWindowUsage{usagePercent: pct, resetInSec: reset} - } - } - if m := re2.FindStringSubmatch(html); len(m) >= 3 { - if reset, pct, ok := parseScrapedNumbers(m[1], m[2]); ok { - return &scrapedWindowUsage{usagePercent: pct, resetInSec: reset} - } - } - return nil -} - -func parseScrapedNumbers(a, b string) (float64, float64, bool) { - first, err1 := strconv.ParseFloat(a, 64) - second, err2 := strconv.ParseFloat(b, 64) - if err1 != nil || err2 != nil { - return 0, 0, false - } - return first, second, true + return snap, err } -func parseDataSlotFormat(html string) map[string]*scrapedWindowUsage { - result := make(map[string]*scrapedWindowUsage) - parts := strings.Split(html, `data-slot="usage-item"`) - for i := 1; i < len(parts); i++ { - content := parts[i] - - labelRe := regexp.MustCompile(`data-slot="usage-label">([^<]+)<`) - labelMatch := labelRe.FindStringSubmatch(content) - if len(labelMatch) < 2 { - continue - } - label := strings.ToLower(strings.TrimSpace(labelMatch[1])) - - usageRe := regexp.MustCompile(`data-slot="usage-value">[^0-9]*(\d+(?:\.\d+)?)`) - usageMatch := usageRe.FindStringSubmatch(content) - if len(usageMatch) < 2 { - continue - } - usagePercent, err := strconv.ParseFloat(usageMatch[1], 64) - if err != nil { - continue - } - - resetRe := regexp.MustCompile(`data-slot="(reset-time|reset-now)">([\s\S]*?)`) - resetMatch := resetRe.FindStringSubmatch(content) - if len(resetMatch) < 3 { - continue - } - - var resetInSec float64 - if resetMatch[1] == "reset-now" { - resetInSec = 0 - } else { - resetContent := resetMatch[2] - resetContent = regexp.MustCompile(``).ReplaceAllString(resetContent, "") - resetContent = strings.TrimSpace(resetContent) - resetContent = regexp.MustCompile(`(?i)Resets?\s*in\s*`).ReplaceAllString(resetContent, "") - parsed, ok := parseHumanReadableTime(resetContent) - if !ok { - continue - } - resetInSec = parsed - } - - var windowKey string - switch { - case strings.Contains(label, "rolling"): - windowKey = "rolling" - case strings.Contains(label, "weekly"): - windowKey = "weekly" - case strings.Contains(label, "monthly"): - windowKey = "monthly" - default: - continue - } - - result[windowKey] = &scrapedWindowUsage{ - usagePercent: usagePercent, - resetInSec: resetInSec, - } +// openCodeConsoleCookieHeader sends a bare value as __Host-console_session and +// a pasted cookie header verbatim (minus a copied "Cookie:" prefix). A bare +// value may itself end in "=" padding, so only a named session cookie or a +// multi-cookie header is taken as-is. +func openCodeConsoleCookieHeader(value string) string { + if len(value) > len("cookie:") && strings.EqualFold(value[:len("cookie:")], "cookie:") { + value = strings.TrimSpace(value[len("cookie:"):]) } - return result -} - -func parseHumanReadableTime(timeStr string) (float64, bool) { - normalized := strings.ToLower(strings.TrimSpace(timeStr)) - normalized = strings.Join(strings.Fields(normalized), " ") - switch normalized { - case "reset-now", "reset now", "now", "resets now": - return 0, true + if strings.Contains(value, "__Host-console_session=") || strings.Contains(value, ";") { + return value } - - var totalSeconds float64 - hasDuration := false - - dayRe := regexp.MustCompile(`(\d+(?:\.\d+)?)\s*days?`) - hourRe := regexp.MustCompile(`(\d+(?:\.\d+)?)\s*hours?`) - minuteRe := regexp.MustCompile(`(\d+(?:\.\d+)?)\s*minutes?`) - secondRe := regexp.MustCompile(`(\d+(?:\.\d+)?)\s*seconds?`) - - if m := dayRe.FindStringSubmatch(normalized); len(m) >= 2 { - if v, err := strconv.ParseFloat(m[1], 64); err == nil { - totalSeconds += v * 86400 - hasDuration = true - } - } - if m := hourRe.FindStringSubmatch(normalized); len(m) >= 2 { - if v, err := strconv.ParseFloat(m[1], 64); err == nil { - totalSeconds += v * 3600 - hasDuration = true - } - } - if m := minuteRe.FindStringSubmatch(normalized); len(m) >= 2 { - if v, err := strconv.ParseFloat(m[1], 64); err == nil { - totalSeconds += v * 60 - hasDuration = true - } - } - if m := secondRe.FindStringSubmatch(normalized); len(m) >= 2 { - if v, err := strconv.ParseFloat(m[1], 64); err == nil { - totalSeconds += v - hasDuration = true - } - } - - return totalSeconds, hasDuration + return "__Host-console_session=" + value } func IsOpenCodeAuthError(err error) bool { diff --git a/internal/api/opencode_client_test.go b/internal/api/opencode_client_test.go index 8cdaa5cf..437ae83a 100644 --- a/internal/api/opencode_client_test.go +++ b/internal/api/opencode_client_test.go @@ -6,245 +6,152 @@ import ( "net/http" "net/http/httptest" "strings" + "sync/atomic" "testing" - "time" ) -const ssrFixtureHTML = ` -rollingUsage:$R[123]={usagePercent:4.5,resetInSec:6000} -weeklyUsage:$R[124]={resetInSec:1209600,usagePercent:12.3} -monthlyUsage:$R[125]={usagePercent:25.0,resetInSec:2592000} -` - -const dataSlotFixtureHTML = ` -
- Rolling Usage - 15% - Resets in 1 hour 30 minutes -
-
- Weekly Usage - 22.5% - Reset now -
-
- Monthly Usage - 40% - Resets in 6 days 2 hours -
-` - -func TestOpenCodeClient_FetchSnapshot_SSR(t *testing.T) { +func TestOpenCodeClient_FetchSnapshot_CookieModeReadsGoStatus(t *testing.T) { + var gotPath, gotCookie, gotOrg, gotAuth, gotAccept string srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/workspace/ws-123/go" { - t.Errorf("unexpected path: %s", r.URL.Path) - } - if cookie := r.Header.Get("Cookie"); cookie != "auth=secret-cookie" { - t.Errorf("unexpected cookie: %q", cookie) - } - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(ssrFixtureHTML)) + gotPath, gotCookie, gotOrg = r.URL.Path, r.Header.Get("Cookie"), r.Header.Get("x-org-id") + gotAuth, gotAccept = r.Header.Get("Authorization"), r.Header.Get("Accept") + _, _ = w.Write([]byte(goStatusBody)) })) defer srv.Close() - client := newTestOpenCodeClient(t, srv) - snap, err := client.FetchSnapshot(context.Background(), "ws-123", "secret-cookie") + snap, err := newTestOpenCodeClient(t, srv).FetchSnapshot(context.Background(), " wrk_123 ", " sess-value ") if err != nil { t.Fatalf("FetchSnapshot: %v", err) } - if len(snap.Quotas) != 3 { - t.Fatalf("quotas = %d, want 3", len(snap.Quotas)) + if gotPath != "/console/api/go/status" { + t.Fatalf("path = %q, want /console/api/go/status", gotPath) } - - byName := mapQuotasByName(snap.Quotas) - if byName["five_hour"].Utilization != 4.5 { - t.Errorf("five_hour util = %v, want 4.5", byName["five_hour"].Utilization) + if gotCookie != "__Host-console_session=sess-value" || gotOrg != "wrk_123" { + t.Fatalf("cookie=%q x-org-id=%q", gotCookie, gotOrg) } - if byName["weekly"].Utilization != 12.3 { - t.Errorf("weekly util = %v, want 12.3", byName["weekly"].Utilization) + if gotAuth != "" || gotAccept != "application/json" { + t.Fatalf("authorization=%q accept=%q, want no bearer and JSON", gotAuth, gotAccept) } - if byName["monthly"].Utilization != 25.0 { - t.Errorf("monthly util = %v, want 25.0", byName["monthly"].Utilization) + if snap.PlanName != "OpenCode Go" || len(snap.Quotas) != 3 { + t.Fatalf("snapshot = %+v", snap) } - for _, q := range snap.Quotas { - if q.Format != OpenCodeQuotaFormatPercent { - t.Errorf("quota %s format = %q, want percent", q.Name, q.Format) - } - if q.ResetsAt == nil { - t.Errorf("quota %s missing resetsAt", q.Name) - } + weekly := quotaByName(t, snap.Quotas, "weekly") + if weekly.Format != OpenCodeQuotaFormatCurrency || weekly.Limit != 30 || weekly.Utilization != 12.8 { + t.Fatalf("weekly = %+v, want $30 currency quota at 12.8%%", weekly) } + wantReset(t, quotaByName(t, snap.Quotas, "monthly"), "2026-10-24T15:00:25Z") } -func TestOpenCodeClient_FetchSnapshot_DataSlotFallback(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(dataSlotFixtureHTML)) - })) - defer srv.Close() - - client := newTestOpenCodeClient(t, srv) - snap, err := client.FetchSnapshot(context.Background(), "ws-abc", "tok") - if err != nil { - t.Fatalf("FetchSnapshot: %v", err) - } - if len(snap.Quotas) != 3 { - t.Fatalf("quotas = %d, want 3", len(snap.Quotas)) - } - byName := mapQuotasByName(snap.Quotas) - if byName["five_hour"].Utilization != 15 { - t.Errorf("five_hour util = %v, want 15", byName["five_hour"].Utilization) - } - if byName["weekly"].Utilization != 22.5 { - t.Errorf("weekly util = %v, want 22.5", byName["weekly"].Utilization) - } - if byName["monthly"].Utilization != 40 { - t.Errorf("monthly util = %v, want 40", byName["monthly"].Utilization) - } - - rollingReset := byName["five_hour"].ResetsAt.Sub(snap.CapturedAt) - if rollingReset < 89*time.Minute || rollingReset > 91*time.Minute { - t.Errorf("five_hour reset offset = %v, want ~90m", rollingReset) +func TestOpenCodeClient_FetchSnapshot_CookieHeader(t *testing.T) { + for _, tt := range []struct { + name, value, want string + }{ + {"bare value", "abc", "__Host-console_session=abc"}, + {"bare value with base64 padding", "token==", "__Host-console_session=token=="}, + {"named cookie", "__Host-console_session=abc", "__Host-console_session=abc"}, + {"full cookie header", "theme=dark; __Host-console_session=abc", "theme=dark; __Host-console_session=abc"}, + {"copied header line", "Cookie: __Host-console_session=abc", "__Host-console_session=abc"}, + } { + t.Run(tt.name, func(t *testing.T) { + if got := openCodeConsoleCookieHeader(tt.value); got != tt.want { + t.Fatalf("cookie header = %q, want %q", got, tt.want) + } + }) } } -func TestOpenCodeClient_FetchSnapshot_401(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusUnauthorized) - _, _ = w.Write([]byte("login required secret")) - })) - defer srv.Close() - - client := newTestOpenCodeClient(t, srv) - _, err := client.FetchSnapshot(context.Background(), "ws", "cookie") - if !errors.Is(err, ErrOpenCodeUnauthorized) { - t.Fatalf("err = %v, want ErrOpenCodeUnauthorized", err) - } - if strings.Contains(err.Error(), "secret") { - t.Fatalf("error leaked response body: %v", err) +func TestOpenCodeClient_FetchSnapshot_RejectedCookieExplainsWhichCookie(t *testing.T) { + for _, status := range []int{http.StatusUnauthorized, http.StatusFound} { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if status == http.StatusFound { + http.Redirect(w, r, "/auth/authorize", status) + return + } + w.WriteHeader(status) + _, _ = w.Write([]byte(`{"_tag":"Unauthorized","secret":"BODY-MARKER"}`)) + })) + _, err := newTestOpenCodeClient(t, srv).FetchSnapshot(context.Background(), "ws", "auth=Fe26.2**old") + srv.Close() + if !errors.Is(err, ErrOpenCodeUnauthorized) { + t.Fatalf("status %d: err = %v, want ErrOpenCodeUnauthorized", status, err) + } + if !strings.Contains(err.Error(), "__Host-console_session") { + t.Fatalf("status %d: error %q does not name the cookie to paste", status, err) + } + if strings.Contains(err.Error(), "BODY-MARKER") || strings.Contains(err.Error(), "Fe26") { + t.Fatalf("status %d: error leaks the body or cookie: %v", status, err) + } } } -func TestOpenCodeClient_FetchSnapshot_RedirectIsUnauthorizedAndNotFollowed(t *testing.T) { - var redirectTargetHits int - target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - redirectTargetHits++ - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(ssrFixtureHTML)) +func TestOpenCodeClient_FetchSnapshot_RedirectIsNotFollowed(t *testing.T) { + var targetHits atomic.Int32 + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + targetHits.Add(1) + _, _ = w.Write([]byte(goStatusBody)) })) defer target.Close() - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, target.URL+"/login", http.StatusFound) })) defer srv.Close() - client := newTestOpenCodeClient(t, srv) - _, err := client.FetchSnapshot(context.Background(), "ws", "secret-cookie") - if !errors.Is(err, ErrOpenCodeUnauthorized) { + if _, err := newTestOpenCodeClient(t, srv).FetchSnapshot(context.Background(), "ws", "cookie"); !errors.Is(err, ErrOpenCodeUnauthorized) { t.Fatalf("err = %v, want ErrOpenCodeUnauthorized", err) } - if redirectTargetHits != 0 { - t.Fatalf("redirect target received %d request(s), want 0", redirectTargetHits) + if targetHits.Load() != 0 { + t.Fatalf("redirect target received %d request(s), want 0", targetHits.Load()) } } func TestOpenCodeClient_FetchSnapshot_Malformed(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("no usage data")) - })) - defer srv.Close() - - client := newTestOpenCodeClient(t, srv) - _, err := client.FetchSnapshot(context.Background(), "ws", "cookie") - if !errors.Is(err, ErrOpenCodeParseFailed) { + srv := goStatusServer(t, `no usage data`) + if _, err := newTestOpenCodeClient(t, srv).FetchSnapshot(context.Background(), "ws", "cookie"); !errors.Is(err, ErrOpenCodeParseFailed) { t.Fatalf("err = %v, want ErrOpenCodeParseFailed", err) } } -func TestOpenCodeClient_FetchSnapshot_MissingConfig(t *testing.T) { - client := NewOpenCodeClient(nil) - _, err := client.FetchSnapshot(context.Background(), "", "cookie") - if !errors.Is(err, ErrOpenCodeMissingConfig) { +func TestOpenCodeClient_FetchSnapshot_MissingConfigMakesNoRequest(t *testing.T) { + var calls atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { calls.Add(1) })) + defer srv.Close() + client := newTestOpenCodeClient(t, srv) + if _, err := client.FetchSnapshot(context.Background(), "", "cookie"); !errors.Is(err, ErrOpenCodeMissingConfig) { t.Fatalf("empty workspace err = %v", err) } - _, err = client.FetchSnapshot(context.Background(), "ws", "") - if !errors.Is(err, ErrOpenCodeMissingConfig) { + if _, err := client.FetchSnapshot(context.Background(), "ws", " "); !errors.Is(err, ErrOpenCodeMissingConfig) { t.Fatalf("empty cookie err = %v", err) } -} - -func TestOpenCodeClient_FetchSnapshot_CookieHeader(t *testing.T) { - var gotCookies []string - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotCookies = append(gotCookies, r.Header.Get("Cookie")) - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(ssrFixtureHTML)) - })) - defer srv.Close() - - client := newTestOpenCodeClient(t, srv) - for _, tt := range []struct { - name string - value string - want string - }{ - {name: "prefixed", value: "auth=already-set", want: "auth=already-set"}, - {name: "raw padded", value: "token==", want: "auth=token=="}, - } { - t.Run(tt.name, func(t *testing.T) { - _, err := client.FetchSnapshot(context.Background(), "ws", tt.value) - if err != nil { - t.Fatalf("FetchSnapshot: %v", err) - } - if got := gotCookies[len(gotCookies)-1]; got != tt.want { - t.Errorf("cookie = %q, want %q", got, tt.want) - } - }) + if calls.Load() != 0 { + t.Fatalf("requests = %d, want 0", calls.Load()) } } -func TestOpenCodeClient_FetchSnapshot_WorkspaceURLEncoded(t *testing.T) { - var gotRequestURI string - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - gotRequestURI = r.RequestURI - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(ssrFixtureHTML)) +// The reuse window is per credential: a cookie never serves a cached API-key +// status (or another workspace's), and vice versa. +func TestOpenCodeClient_StatusReuseIsPerCredential(t *testing.T) { + var calls atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + calls.Add(1) + _, _ = w.Write([]byte(goStatusBody)) })) defer srv.Close() - - client := newTestOpenCodeClient(t, srv) - _, err := client.FetchSnapshot(context.Background(), "ws/special id", "cookie") - if err != nil { - t.Fatalf("FetchSnapshot: %v", err) - } - if !strings.Contains(gotRequestURI, "ws%2Fspecial%20id") { - t.Errorf("request URI = %q, want encoded workspace id", gotRequestURI) - } -} - -func TestParseHumanReadableTime(t *testing.T) { - tests := []struct { - in string - want float64 - ok bool - }{ - {"1 hour 56 minutes", 6960, true}, - {"6 days 2 hours", 525600, true}, - {"reset now", 0, true}, - {"not a duration", 0, false}, - } - for _, tc := range tests { - got, ok := parseHumanReadableTime(tc.in) - if ok != tc.ok { - t.Errorf("parseHumanReadableTime(%q) ok = %v, want %v", tc.in, ok, tc.ok) - continue - } - if ok && got != tc.want { - t.Errorf("parseHumanReadableTime(%q) = %v, want %v", tc.in, got, tc.want) + c := newTestOpenCodeClient(t, srv) + ctx := context.Background() + steps := []func() error{ + func() error { _, err := c.FetchUsageSnapshot(ctx, "k"); return err }, + func() error { _, err := c.FetchSnapshot(ctx, "ws", "k"); return err }, + func() error { _, err := c.FetchSnapshot(ctx, "ws", "k"); return err }, + func() error { _, err := c.FetchSnapshot(ctx, "ws2", "k"); return err }, + } + for i, step := range steps { + if err := step(); err != nil { + t.Fatalf("step %d: %v", i, err) } } + if calls.Load() != 3 { + t.Fatalf("status requests = %d, want 3 (key, cookie ws, cookie ws2)", calls.Load()) + } } func TestIsOpenCodeAuthError(t *testing.T) { @@ -263,11 +170,3 @@ func newTestOpenCodeClient(t *testing.T, srv *httptest.Server) *OpenCodeClient { t.Helper() return NewOpenCodeClient(nil, WithOpenCodeBaseURL(srv.URL)) } - -func mapQuotasByName(quotas []OpenCodeQuota) map[string]OpenCodeQuota { - out := make(map[string]OpenCodeQuota, len(quotas)) - for _, q := range quotas { - out[q.Name] = q - } - return out -} diff --git a/internal/api/opencode_usage.go b/internal/api/opencode_usage.go index fc90d3cb..e2b50bbb 100644 --- a/internal/api/opencode_usage.go +++ b/internal/api/opencode_usage.go @@ -15,22 +15,24 @@ import ( // OpenCode Go usage via the console's Go status API. // -// The Go dashboard (/workspace//go) has no public quota API and moved to a -// new console, which breaks HTML scraping. GET /console/api/go/status accepts a -// console service-account key and returns the subscription's own meters: the -// 5-hour, weekly and monthly used/limit amounts OpenCode enforces and the Go -// page shows. The amounts are micro-cents of the plan's base allowance, so the -// bar is simply used/limit. See anomalyco/opencode#50912. +// The Go dashboard moved to a single-page console and /workspace//go now +// redirects to login, so there is no HTML left to scrape. GET +// /console/api/go/status returns the subscription's own meters: the 5-hour, +// weekly and monthly used/limit amounts OpenCode enforces and the Go page +// shows. It accepts either a console service-account key (Bearer) or the +// browser's __Host-console_session cookie plus an x-org-id header. Amounts +// are micro-cents; onWatch stores them as USD. See anomalyco/opencode#50912. const ( openCodeGoStatusPath = "/console/api/go/status" openCodeUsageUserAgent = "onwatch-opencode-usage/1" openCodeUsageMaxBodyBytes = 1 << 20 // the status document is well under 1 KiB // openCodeUsageTimeout: go/status is a live billing lookup behind - // Cloudflare rather than a cached page, so it gets more headroom than the - // 10s scrape; 20s still fails a stuck request well before the next poll. + // Cloudflare rather than a cached page, so it gets some headroom; 20s + // still fails a stuck request well before the next poll. openCodeUsageTimeout = 20 * time.Second openCodeUsageMinInterval = 60 * time.Second + openCodeMicroCentsPerUSD = 1e8 // 100 cents x 1e6 micro-cents ) // ErrOpenCodeMissingAPIKey is returned when usage-API mode has no key. @@ -86,7 +88,15 @@ func (c *OpenCodeClient) FetchUsageSnapshot(ctx context.Context, apiKey string) if apiKey == "" { return nil, ErrOpenCodeMissingAPIKey } - status, err := c.goStatus(ctx, apiKey) + return c.goStatusSnapshot(ctx, "key\x00"+apiKey, func(req *http.Request) { + req.Header.Set("Authorization", "Bearer "+apiKey) + }) +} + +// goStatusSnapshot fetches go/status with the given credential and maps it. +// credKey identifies the credential for the reuse window. +func (c *OpenCodeClient) goStatusSnapshot(ctx context.Context, credKey string, authorize func(*http.Request)) (*OpenCodeSnapshot, error) { + status, err := c.goStatus(ctx, credKey, authorize) if err != nil { return nil, err } @@ -102,23 +112,23 @@ func (c *OpenCodeClient) FetchUsageSnapshot(ctx context.Context, apiKey string) }, nil } -// goStatus returns the status, reusing a recent one so a short poll interval -// does not query the console every few seconds. -func (c *OpenCodeClient) goStatus(ctx context.Context, apiKey string) (*openCodeGoStatus, error) { +// goStatus returns the status, reusing a recent one for the same credential so +// a short poll interval does not query the console every few seconds. +func (c *OpenCodeClient) goStatus(ctx context.Context, credKey string, authorize func(*http.Request)) (*openCodeGoStatus, error) { c.usageMu.Lock() defer c.usageMu.Unlock() - if c.usageStatus != nil && c.usageKey == apiKey && time.Since(c.usageAt) < openCodeUsageMinInterval { + if c.usageStatus != nil && c.usageKey == credKey && time.Since(c.usageAt) < openCodeUsageMinInterval { return c.usageStatus, nil } - status, err := c.fetchGoStatus(ctx, apiKey) + status, err := c.fetchGoStatus(ctx, authorize) if err != nil { return nil, err } - c.usageStatus, c.usageKey, c.usageAt = status, apiKey, time.Now() + c.usageStatus, c.usageKey, c.usageAt = status, credKey, time.Now() return status, nil } -func (c *OpenCodeClient) fetchGoStatus(ctx context.Context, apiKey string) (*openCodeGoStatus, error) { +func (c *OpenCodeClient) fetchGoStatus(ctx context.Context, authorize func(*http.Request)) (*openCodeGoStatus, error) { req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.goStatusURL, nil) if err != nil { return nil, fmt.Errorf("%w: build request: %v", ErrOpenCodeNetworkError, err) @@ -126,9 +136,9 @@ func (c *OpenCodeClient) fetchGoStatus(ctx context.Context, apiKey string) (*ope // Cloudflare in front of opencode.ai rejects default client User-Agents. req.Header.Set("User-Agent", openCodeUsageUserAgent) req.Header.Set("Accept", "application/json") - req.Header.Set("Authorization", "Bearer "+apiKey) + authorize(req) - resp, err := c.usageHTTPClient.Do(req) + resp, err := c.httpClient.Do(req) if err != nil { if ctx.Err() != nil { return nil, ctx.Err() @@ -139,7 +149,8 @@ func (c *OpenCodeClient) fetchGoStatus(ctx context.Context, apiKey string) (*ope // Bodies are never logged or echoed: the document carries account IDs. switch { case resp.StatusCode == http.StatusOK: - case resp.StatusCode == http.StatusUnauthorized: + case resp.StatusCode == http.StatusUnauthorized, + resp.StatusCode >= 300 && resp.StatusCode < 400: // redirect to login return nil, ErrOpenCodeUnauthorized case resp.StatusCode == http.StatusForbidden: return nil, ErrOpenCodeForbidden @@ -167,7 +178,7 @@ func (c *OpenCodeClient) fetchGoStatus(ctx context.Context, apiKey string) (*ope return &status, nil } -// quotas maps the three meters to the quota names scrape mode produces. Every +// quotas maps the three meters to onWatch's quota names. Every // meter must be present with a positive limit: a partial document is a format // change, not zero usage. func (s *openCodeGoStatus) quotas() ([]OpenCodeQuota, error) { @@ -201,13 +212,12 @@ func (s *openCodeGoStatus) quotas() ([]OpenCodeQuota, error) { t := resetsAt.UTC() reset = &t } - pct := math.Round(used/limit*1000) / 10 quotas = append(quotas, OpenCodeQuota{ Name: w.name, - Used: pct, - Limit: 100, - Utilization: pct, - Format: OpenCodeQuotaFormatPercent, + Used: used / openCodeMicroCentsPerUSD, + Limit: limit / openCodeMicroCentsPerUSD, + Utilization: math.Round(used/limit*1000) / 10, + Format: OpenCodeQuotaFormatCurrency, ResetsAt: reset, }) } diff --git a/internal/api/opencode_usage_test.go b/internal/api/opencode_usage_test.go index 48539d2d..caf2b442 100644 --- a/internal/api/opencode_usage_test.go +++ b/internal/api/opencode_usage_test.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "math" "net/http" "net/http/httptest" "strings" @@ -88,13 +89,24 @@ func TestFetchUsageSnapshot_SendsHeadersAndMapsMeters(t *testing.T) { names := []string{} for _, q := range snap.Quotas { names = append(names, q.Name) - if q.Limit != 100 || q.Used != q.Utilization || q.Format != OpenCodeQuotaFormatPercent { - t.Fatalf("quota %+v, want percent of 100 with Used == Utilization", q) + if q.Format != OpenCodeQuotaFormatCurrency { + t.Fatalf("quota %+v, want currency format", q) } } if strings.Join(names, ",") != "five_hour,weekly,monthly" { t.Fatalf("quota names = %v", names) } + // Micro-cents to USD: 1 USD = 100 cents = 1e8 micro-cents. + for name, want := range map[string][2]float64{ + "five_hour": {0, 12}, + "weekly": {3.84204992, 30}, + "monthly": {3.84204992, 60}, + } { + q := quotaByName(t, snap.Quotas, name) + if math.Abs(q.Used-want[0]) > 1e-9 || q.Limit != want[1] { + t.Fatalf("%s used/limit = %v/%v, want %v/%v USD", name, q.Used, q.Limit, want[0], want[1]) + } + } // 384204992 / 3000000000 = 12.807% and / 6000000000 = 6.403% (console: 12.81%, 6.40%). for name, want := range map[string]float64{"five_hour": 0, "weekly": 12.8, "monthly": 6.4} { if got := quotaByName(t, snap.Quotas, name).Utilization; got != want { @@ -184,7 +196,7 @@ func TestFetchUsageSnapshot_MapsHTTPErrorsWithoutEchoingBody(t *testing.T) { {http.StatusBadGateway, ErrOpenCodeServerError}, {http.StatusBadRequest, ErrOpenCodeInvalidResponse}, {http.StatusNotFound, ErrOpenCodeInvalidResponse}, - {http.StatusFound, ErrOpenCodeInvalidResponse}, + {http.StatusFound, ErrOpenCodeUnauthorized}, // login redirect } for _, tc := range cases { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -271,18 +283,15 @@ func TestFetchUsageSnapshot_FailuresAreNotReused(t *testing.T) { } } -func TestNewOpenCodeClient_StatusGetsItsOwnTimeout(t *testing.T) { +func TestNewOpenCodeClient_Timeout(t *testing.T) { c := NewOpenCodeClient(nil) - if c.httpClient.Timeout != openCodeScrapeTimeout || c.usageHTTPClient.Timeout != openCodeUsageTimeout { - t.Fatalf("timeouts: scrape=%v status=%v", c.httpClient.Timeout, c.usageHTTPClient.Timeout) - } - if tr, ok := c.usageHTTPClient.Transport.(*http.Transport); !ok || tr.ResponseHeaderTimeout != openCodeUsageTimeout { - t.Fatalf("status transport = %+v", c.usageHTTPClient.Transport) + if c.httpClient.Timeout != openCodeUsageTimeout { + t.Fatalf("timeout = %v, want %v", c.httpClient.Timeout, openCodeUsageTimeout) } - if tr := c.httpClient.Transport.(*http.Transport); tr.ResponseHeaderTimeout != openCodeScrapeTimeout { - t.Fatalf("scrape transport header timeout changed to %v", tr.ResponseHeaderTimeout) + if tr, ok := c.httpClient.Transport.(*http.Transport); !ok || tr.ResponseHeaderTimeout != openCodeUsageTimeout { + t.Fatalf("transport = %+v", c.httpClient.Transport) } - if long := NewOpenCodeClient(nil, WithOpenCodeTimeout(45*time.Second)); long.usageHTTPClient.Timeout != 45*time.Second { - t.Fatalf("a longer configured timeout was shortened to %v", long.usageHTTPClient.Timeout) + if long := NewOpenCodeClient(nil, WithOpenCodeTimeout(45*time.Second)); long.httpClient.Timeout != 45*time.Second { + t.Fatalf("configured timeout = %v, want 45s", long.httpClient.Timeout) } } diff --git a/internal/api/openrouter_client.go b/internal/api/openrouter_client.go index 593ecab3..11fccba0 100644 --- a/internal/api/openrouter_client.go +++ b/internal/api/openrouter_client.go @@ -12,10 +12,10 @@ import ( // Custom errors for OpenRouter API failures. var ( - ErrOpenRouterUnauthorized = errors.New("openrouter: unauthorized - invalid API key") - ErrOpenRouterRateLimited = errors.New("openrouter: rate limited") - ErrOpenRouterServerError = errors.New("openrouter: server error") - ErrOpenRouterNetworkError = errors.New("openrouter: network error") + ErrOpenRouterUnauthorized = errors.New("openrouter: unauthorized - invalid API key") + ErrOpenRouterRateLimited = errors.New("openrouter: rate limited") + ErrOpenRouterServerError = errors.New("openrouter: server error") + ErrOpenRouterNetworkError = errors.New("openrouter: network error") ErrOpenRouterInvalidResponse = errors.New("openrouter: invalid response") ) diff --git a/internal/api/test_main_test.go b/internal/api/test_main_test.go index 214f3027..2110c69a 100644 --- a/internal/api/test_main_test.go +++ b/internal/api/test_main_test.go @@ -1,8 +1,11 @@ package api import ( + "fmt" "os" "testing" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) // TestMain runs before all tests in the api package. It enables test mode @@ -11,13 +14,28 @@ import ( // WriteAnthropicCredentials or DetectAnthropicToken can overwrite the user's // real Claude Code OAuth tokens, causing Claude Code to be logged out. // -// It also clears OPENCODE_HOME/XDG_DATA_HOME so package-level detection does -// not inherit a developer's override. Individual tests that must stay hermetic -// should call isolateOpenCodeEnv (pins both to empty temp dirs) because -// clearing alone can still fall through to the real UserHomeDir path. +// It also points the home directory at an empty sandbox for the whole run and +// clears provider location overrides (CODEX_HOME, OPENCODE_HOME, XDG_*, ...), +// so a test that forgets to isolate itself can never read or write the real +// ~/.claude, ~/.codex, ~/.kimi-code and so on. Individual tests that must stay +// hermetic should call isolateOpenCodeEnv (pins OPENCODE_HOME/XDG_DATA_HOME to +// empty temp dirs) because clearing alone can still fall through to the +// UserHomeDir path. func TestMain(m *testing.M) { + os.Exit(runTests(m)) +} + +func runTests(m *testing.M) int { + // SetTestMode must run before HOME/USERPROFILE are redirected: its first + // enable records the real home that the credential-file guard refuses. SetTestMode(true) - os.Unsetenv("OPENCODE_HOME") - os.Unsetenv("XDG_DATA_HOME") - os.Exit(m.Run()) + + _, cleanup, err := testhome.SandboxHome() + if err != nil { + fmt.Fprintf(os.Stderr, "api tests: %v\n", err) + return 1 + } + defer cleanup() + + return m.Run() } diff --git a/internal/config/auth_mode_test.go b/internal/config/auth_mode_test.go index b22d7a37..bf9c5fec 100644 --- a/internal/config/auth_mode_test.go +++ b/internal/config/auth_mode_test.go @@ -17,7 +17,7 @@ func mustParseIP(t *testing.T, s string) net.IP { func TestAuthMode_DefaultLocal(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := loadWithArgs(nil) if err != nil { @@ -36,7 +36,7 @@ func TestAuthMode_TrustedProxy(t *testing.T) { os.Setenv("ONWATCH_AUTH_MODE", "trusted_proxy") os.Setenv("ONWATCH_TRUSTED_PROXY_CIDRS", "172.30.0.0/16, 127.0.0.1") os.Setenv("ONWATCH_TRUSTED_USER_HEADER", "X-authentik-username") - defer os.Clearenv() + defer clearTestEnv() cfg, err := loadWithArgs(nil) if err != nil { @@ -56,7 +56,7 @@ func TestAuthMode_TrustedProxy(t *testing.T) { func TestAuthMode_TrustedProxyWithoutCIDRsFailsClosed(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ONWATCH_AUTH_MODE", "trusted_proxy") - defer os.Clearenv() + defer clearTestEnv() if _, err := loadWithArgs(nil); err == nil { t.Fatal("expected error when trusted_proxy mode has no CIDRs, got nil") @@ -66,7 +66,7 @@ func TestAuthMode_TrustedProxyWithoutCIDRsFailsClosed(t *testing.T) { func TestAuthMode_InvalidValueRejected(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ONWATCH_AUTH_MODE", "oidc") - defer os.Clearenv() + defer clearTestEnv() if _, err := loadWithArgs(nil); err == nil { t.Fatal("expected error for invalid ONWATCH_AUTH_MODE, got nil") @@ -77,7 +77,7 @@ func TestAuthMode_InvalidCIDRRejected(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ONWATCH_AUTH_MODE", "trusted_proxy") os.Setenv("ONWATCH_TRUSTED_PROXY_CIDRS", "not-a-cidr") - defer os.Clearenv() + defer clearTestEnv() if _, err := loadWithArgs(nil); err == nil { t.Fatal("expected error for invalid CIDR, got nil") diff --git a/internal/config/config.go b/internal/config/config.go index 28588f52..2bcd5538 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -54,8 +54,9 @@ type Config struct { OpenCodeEnabled bool // OPENCODE_ENABLED=true: track ChatGPT via OpenCode auth.json (feeds Codex) // OpenCode Go provider configuration OpenCodeGoWorkspaceID string // OPENCODE_GO_WORKSPACE_ID - OpenCodeGoAuthCookie string // OPENCODE_GO_AUTH_COOKIE - OpenCodeGoAPIKey string // OPENCODE_GO_API_KEY: console service-account key (reads the plan's meters via go/status; preferred over the cookie scrape) + OpenCodeGoAuthCookie string // OPENCODE_GO_AUTH_COOKIE: __Host-console_session cookie value (or a full cookie header) + OpenCodeGoAPIKey string // OPENCODE_GO_API_KEY: console service-account key (preferred over the session cookie) + OpenCodeGoBaseURL string // OPENCODE_GO_BASE_URL override for proxy setups (default https://opencode.ai) // Ollama Cloud provider configuration OllamaAPIKey string // OLLAMA_API_KEY from ollama.com/settings/keys OllamaMonthlyLimit float64 // OLLAMA_MONTHLY_LIMIT: included usage cap in USD (overrides the plan default; 0 = derive from plan) @@ -86,10 +87,12 @@ type Config struct { OpenRouterAPIKey string // OPENROUTER_API_KEY // Moonshot provider configuration - MoonshotAPIKey string // MOONSHOT_API_KEY + MoonshotAPIKey string // MOONSHOT_API_KEY + MoonshotBaseURL string // MOONSHOT_BASE_URL override for proxy setups (default https://api.moonshot.ai) // DeepSeek provider configuration - DeepSeekAPIKey string // DEEPSEEK_API_KEY + DeepSeekAPIKey string // DEEPSEEK_API_KEY + DeepSeekBaseURL string // DEEPSEEK_BASE_URL override for proxy setups (default https://api.deepseek.com) // Gemini provider configuration (auto-detected from ~/.gemini/oauth_creds.json or env vars) GeminiEnabled bool // true if auto-detected or GEMINI_ENABLED=true @@ -256,6 +259,7 @@ var onwatchEnvKeys = []string{ "OPENCODE_GO_WORKSPACE_ID", "OPENCODE_GO_AUTH_COOKIE", "OPENCODE_GO_API_KEY", + "OPENCODE_GO_BASE_URL", "OPENCODE_HOME", "OLLAMA_API_KEY", "OLLAMA_MONTHLY_LIMIT", @@ -273,7 +277,9 @@ var onwatchEnvKeys = []string{ "MINIMAX_API_KEY", "OPENROUTER_API_KEY", "MOONSHOT_API_KEY", + "MOONSHOT_BASE_URL", "DEEPSEEK_API_KEY", + "DEEPSEEK_BASE_URL", "CURSOR_TOKEN", "GROK_TOKEN", "GROK_ENABLED", @@ -395,6 +401,7 @@ func loadFromEnvAndFlags(flags *flagValues) (*Config, error) { cfg.OpenCodeGoWorkspaceID = strings.TrimSpace(os.Getenv("OPENCODE_GO_WORKSPACE_ID")) cfg.OpenCodeGoAuthCookie = strings.TrimSpace(os.Getenv("OPENCODE_GO_AUTH_COOKIE")) cfg.OpenCodeGoAPIKey = strings.TrimSpace(os.Getenv("OPENCODE_GO_API_KEY")) + cfg.OpenCodeGoBaseURL = strings.TrimSpace(os.Getenv("OPENCODE_GO_BASE_URL")) cfg.OllamaAPIKey = strings.TrimSpace(os.Getenv("OLLAMA_API_KEY")) if v := strings.TrimSpace(os.Getenv("OLLAMA_MONTHLY_LIMIT")); v != "" { if f, err := strconv.ParseFloat(strings.TrimPrefix(v, "$"), 64); err == nil && f > 0 { @@ -446,9 +453,11 @@ func loadFromEnvAndFlags(flags *flagValues) (*Config, error) { // Moonshot provider cfg.MoonshotAPIKey = strings.TrimSpace(os.Getenv("MOONSHOT_API_KEY")) + cfg.MoonshotBaseURL = strings.TrimSpace(os.Getenv("MOONSHOT_BASE_URL")) // DeepSeek provider cfg.DeepSeekAPIKey = strings.TrimSpace(os.Getenv("DEEPSEEK_API_KEY")) + cfg.DeepSeekBaseURL = strings.TrimSpace(os.Getenv("DEEPSEEK_BASE_URL")) // Gemini provider (auto-detected, env vars, or opt-out via GEMINI_ENABLED=false) cfg.GeminiRefreshToken = strings.TrimSpace(os.Getenv("GEMINI_REFRESH_TOKEN")) @@ -1210,7 +1219,7 @@ func (c *Config) IsDefaultPassword() bool { } // OpenCodeGoConfigured reports whether OpenCode Go tracking has credentials: -// a usage-API key (preferred) or the legacy workspace ID + auth cookie pair. +// a usage-API key (preferred) or the workspace ID + console session cookie pair. func (c *Config) OpenCodeGoConfigured() bool { return c.OpenCodeGoAPIKey != "" || (c.OpenCodeGoWorkspaceID != "" && c.OpenCodeGoAuthCookie != "") } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index d4f7023c..b3d162dd 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -8,6 +8,37 @@ import ( "time" ) +// preservedTestEnv holds the process variables clearTestEnv keeps: the ones +// that locate the OS temp directory (os.TempDir, t.TempDir) plus the Windows +// system root. Without TMP/TEMP, Windows falls back to C:\Windows as the temp +// directory. HOME and USERPROFILE are deliberately NOT kept, so loadEnvFile +// can never pick up the developer's real ~/.onwatch/.env during a test. +var preservedTestEnv = func() map[string]string { + keep := map[string]string{} + for _, key := range []string{"TMPDIR", "TMP", "TEMP", "SystemRoot", "windir"} { + if v, ok := os.LookupEnv(key); ok { + keep[key] = v + } + } + return keep +}() + +// clearTestEnv empties the process environment like os.Clearenv, but keeps +// the temp-dir and system variables in preservedTestEnv. +func clearTestEnv() { + os.Clearenv() + for key, v := range preservedTestEnv { + os.Setenv(key, v) + } +} + +// setTestHome points os.UserHomeDir at dir on every OS: it reads HOME on +// Unix and USERPROFILE on Windows. +func setTestHome(dir string) { + os.Setenv("HOME", dir) + os.Setenv("USERPROFILE", dir) +} + func TestConfig_LoadsFromEnv(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key_123") os.Setenv("ONWATCH_POLL_INTERVAL", "120") @@ -16,7 +47,7 @@ func TestConfig_LoadsFromEnv(t *testing.T) { os.Setenv("ONWATCH_ADMIN_PASS", "mypass") os.Setenv("ONWATCH_DB_PATH", "/tmp/test.db") os.Setenv("ONWATCH_LOG_LEVEL", "debug") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -47,9 +78,9 @@ func TestConfig_LoadsFromEnv(t *testing.T) { } func TestConfig_LoadsMetricsTokenFromEnv(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("ONWATCH_METRICS_TOKEN", "metrics-secret") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -63,7 +94,7 @@ func TestConfig_LoadsMetricsTokenFromEnv(t *testing.T) { func TestConfig_LoadsZaiFromEnv(t *testing.T) { os.Setenv("ZAI_API_KEY", "zai_test_key_456") os.Setenv("ZAI_BASE_URL", "https://custom.z.ai/api") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -80,7 +111,7 @@ func TestConfig_LoadsZaiFromEnv(t *testing.T) { func TestConfig_ZaiDefaults(t *testing.T) { os.Setenv("ZAI_API_KEY", "zai_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -93,10 +124,10 @@ func TestConfig_ZaiDefaults(t *testing.T) { } func TestConfig_ZaiRegion_LoadsFromEnv(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("ZAI_API_KEY", "zai_test_key") os.Setenv("ZAI_REGION", "cn") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -108,9 +139,9 @@ func TestConfig_ZaiRegion_LoadsFromEnv(t *testing.T) { } func TestConfig_ZaiRegion_DefaultsToGlobal(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("ZAI_API_KEY", "zai_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -122,10 +153,10 @@ func TestConfig_ZaiRegion_DefaultsToGlobal(t *testing.T) { } func TestConfig_ZaiRegion_NormalizesToLowercase(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("ZAI_API_KEY", "zai_test_key") os.Setenv("ZAI_REGION", "CN") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -137,10 +168,10 @@ func TestConfig_ZaiRegion_NormalizesToLowercase(t *testing.T) { } func TestConfig_ZaiRegion_SelectsCNBaseURL(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("ZAI_API_KEY", "zai_test_key") os.Setenv("ZAI_REGION", "cn") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -153,7 +184,7 @@ func TestConfig_ZaiRegion_SelectsCNBaseURL(t *testing.T) { func TestConfig_DefaultValues(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key_123") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -193,9 +224,9 @@ func TestConfig_DefaultValues(t *testing.T) { } func TestConfig_APIIntegrationsRetention_LoadsFromEnv(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("ONWATCH_API_INTEGRATIONS_RETENTION", "168h") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -207,9 +238,9 @@ func TestConfig_APIIntegrationsRetention_LoadsFromEnv(t *testing.T) { } func TestConfig_APIIntegrationsRetention_Disabled(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("ONWATCH_API_INTEGRATIONS_RETENTION", "0") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -222,7 +253,7 @@ func TestConfig_APIIntegrationsRetention_Disabled(t *testing.T) { func TestConfig_OnlySyntheticProvider(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -245,7 +276,7 @@ func TestConfig_OnlySyntheticProvider(t *testing.T) { func TestConfig_OnlyZaiProvider(t *testing.T) { os.Setenv("ZAI_API_KEY", "zai_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -269,7 +300,7 @@ func TestConfig_OnlyZaiProvider(t *testing.T) { func TestConfig_BothProviders(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ZAI_API_KEY", "zai_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -291,9 +322,9 @@ func TestConfig_BothProviders(t *testing.T) { } func TestConfig_MiniMaxProvider(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("MINIMAX_API_KEY", "sk-cp-test-key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -309,10 +340,10 @@ func TestConfig_MiniMaxProvider(t *testing.T) { } func TestConfig_MiniMaxRegion_LoadsFromEnv(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("MINIMAX_API_KEY", "sk-cp-test-key") os.Setenv("MINIMAX_REGION", "cn") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -324,9 +355,9 @@ func TestConfig_MiniMaxRegion_LoadsFromEnv(t *testing.T) { } func TestConfig_MiniMaxRegion_DefaultsToGlobal(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("MINIMAX_API_KEY", "sk-cp-test-key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -338,10 +369,10 @@ func TestConfig_MiniMaxRegion_DefaultsToGlobal(t *testing.T) { } func TestConfig_MiniMaxRegion_NormalizesToLowercase(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("MINIMAX_API_KEY", "sk-cp-test-key") os.Setenv("MINIMAX_REGION", "CN") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -353,7 +384,7 @@ func TestConfig_MiniMaxRegion_NormalizesToLowercase(t *testing.T) { } func TestConfig_AllowsNoProvidersConfigured(t *testing.T) { - os.Clearenv() + clearTestEnv() cfg, err := Load() if err != nil { @@ -380,9 +411,9 @@ func TestConfig_ValidatesSyntheticAPIKey_Format(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("SYNTHETIC_API_KEY", tt.apiKey) - defer os.Clearenv() + defer clearTestEnv() _, err := Load() if tt.wantErr && err == nil { @@ -398,7 +429,7 @@ func TestConfig_ValidatesSyntheticAPIKey_Format(t *testing.T) { func TestConfig_ValidatesInterval_Minimum(t *testing.T) { os.Setenv("ZAI_API_KEY", "zai_test_key") os.Setenv("ONWATCH_POLL_INTERVAL", "5") - defer os.Clearenv() + defer clearTestEnv() _, err := Load() if err == nil { @@ -409,7 +440,7 @@ func TestConfig_ValidatesInterval_Minimum(t *testing.T) { func TestConfig_ValidatesInterval_Maximum(t *testing.T) { os.Setenv("ZAI_API_KEY", "zai_test_key") os.Setenv("ONWATCH_POLL_INTERVAL", "7200") - defer os.Clearenv() + defer clearTestEnv() _, err := Load() if err == nil { @@ -434,10 +465,10 @@ func TestConfig_ValidatesPort_Range(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("ZAI_API_KEY", "zai_test_key") os.Setenv("ONWATCH_PORT", tt.port) - defer os.Clearenv() + defer clearTestEnv() _, err := Load() if tt.wantOK && err != nil { @@ -474,7 +505,7 @@ func TestConfig_RedactsZaiAPIKey(t *testing.T) { func TestConfig_DebugMode_Default(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -491,7 +522,7 @@ func TestConfig_LoadWithArgs_FlagOverridesEnv(t *testing.T) { os.Setenv("ONWATCH_POLL_INTERVAL", "120") os.Setenv("ONWATCH_PORT", "8080") os.Setenv("ONWATCH_DB_PATH", "/tmp/env.db") - defer os.Clearenv() + defer clearTestEnv() cfg, err := loadWithArgs([]string{"--interval", "30", "--port", "9000", "--db", "/tmp/flag.db"}) if err != nil { @@ -511,7 +542,7 @@ func TestConfig_LoadWithArgs_FlagOverridesEnv(t *testing.T) { func TestConfig_LoadWithArgs_EqualsSyntax(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := loadWithArgs([]string{"--interval=45", "--port=7777"}) if err != nil { @@ -528,7 +559,7 @@ func TestConfig_LoadWithArgs_EqualsSyntax(t *testing.T) { func TestConfig_DebugMode_Flag(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := loadWithArgs([]string{"--debug"}) if err != nil { @@ -732,7 +763,7 @@ func TestConfig_LogWriter_RotatesFileWhenAtLimit(t *testing.T) { func TestConfig_LoadsAnthropicFromEnv(t *testing.T) { os.Setenv("ANTHROPIC_TOKEN", "sk-ant-test-token-123") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -746,7 +777,7 @@ func TestConfig_LoadsAnthropicFromEnv(t *testing.T) { func TestConfig_OnlyAnthropicProvider(t *testing.T) { os.Setenv("ANTHROPIC_TOKEN", "sk-ant-test-token") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -774,7 +805,7 @@ func TestConfig_OnlyAnthropicProvider(t *testing.T) { func TestConfig_AnthropicWithOtherProviders(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ANTHROPIC_TOKEN", "sk-ant-test-token") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -1090,6 +1121,11 @@ func TestConfig_LogWriter_TestMode(t *testing.T) { if err != nil { t.Fatalf("LogWriter() failed: %v", err) } + // Close the log file before TempDir cleanup: Windows cannot delete a + // file that still has an open handle. + if file, ok := writer.(*os.File); ok && file != os.Stdout { + t.Cleanup(func() { _ = file.Close() }) + } if writer == os.Stdout { t.Error("TestMode background should not return os.Stdout") } @@ -1102,7 +1138,7 @@ func TestConfig_LogWriter_TestMode(t *testing.T) { func TestConfig_LoadWithArgs_TestFlag(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := loadWithArgs([]string{"--test"}) if err != nil { @@ -1115,7 +1151,7 @@ func TestConfig_LoadWithArgs_TestFlag(t *testing.T) { func TestConfig_LoadWithArgs_DbEqualsSyntax(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := loadWithArgs([]string{"--db=/tmp/equals.db"}) if err != nil { @@ -1128,7 +1164,7 @@ func TestConfig_LoadWithArgs_DbEqualsSyntax(t *testing.T) { func TestConfig_LoadAntigravityFromEnv(t *testing.T) { os.Setenv("ANTIGRAVITY_ENABLED", "true") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -1144,7 +1180,7 @@ func TestConfig_LoadAntigravityFromEnv(t *testing.T) { func TestConfig_LoadCopilotFromEnv(t *testing.T) { os.Setenv("COPILOT_TOKEN", "ghp_test_copilot_token") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -1158,7 +1194,7 @@ func TestConfig_LoadCopilotFromEnv(t *testing.T) { func TestConfig_SecureCookiesFromEnv(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ONWATCH_SECURE_COOKIES", "true") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -1172,7 +1208,7 @@ func TestConfig_SecureCookiesFromEnv(t *testing.T) { func TestConfig_SessionIdleTimeoutFromEnv(t *testing.T) { os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ONWATCH_SESSION_IDLE_TIMEOUT", "300") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -1277,13 +1313,11 @@ func TestIsOnwatchEnvFile_NonexistentFile(t *testing.T) { } func TestLoadEnvFile_PrefersStandardLocation(t *testing.T) { - // Save original HOME and restore after test - origHome := os.Getenv("HOME") - defer os.Setenv("HOME", origHome) - - // Create temp directory structure + // Point the home directory (HOME on Unix, USERPROFILE on Windows) at a + // temp dir; t.Setenv restores the originals after the test. tmpDir := t.TempDir() - os.Setenv("HOME", tmpDir) + t.Setenv("HOME", tmpDir) + t.Setenv("USERPROFILE", tmpDir) // Create ~/.onwatch/.env onwatchDir := filepath.Join(tmpDir, ".onwatch") @@ -1297,8 +1331,8 @@ func TestLoadEnvFile_PrefersStandardLocation(t *testing.T) { } // Clear env and load - os.Clearenv() - os.Setenv("HOME", tmpDir) + clearTestEnv() + setTestHome(tmpDir) loadEnvFile() // Verify the standard location was loaded @@ -1311,17 +1345,14 @@ func TestLoadEnvFile_PrefersStandardLocation(t *testing.T) { } func TestLoadEnvFile_FallsBackToLocalOnwatchEnv(t *testing.T) { - // Save original HOME and cwd - origHome := os.Getenv("HOME") + // Save original cwd; t.Setenv restores HOME/USERPROFILE after the test. origDir, _ := os.Getwd() - defer func() { - os.Setenv("HOME", origHome) - os.Chdir(origDir) - }() + defer os.Chdir(origDir) // Create temp directory with NO ~/.onwatch/.env tmpDir := t.TempDir() - os.Setenv("HOME", tmpDir) + t.Setenv("HOME", tmpDir) + t.Setenv("USERPROFILE", tmpDir) // Create local .env with onwatch-specific keys localDir := filepath.Join(tmpDir, "project") @@ -1340,8 +1371,8 @@ func TestLoadEnvFile_FallsBackToLocalOnwatchEnv(t *testing.T) { } // Clear env and load - os.Clearenv() - os.Setenv("HOME", tmpDir) + clearTestEnv() + setTestHome(tmpDir) loadEnvFile() // Verify the local .env was loaded (because standard location doesn't exist) @@ -1351,17 +1382,14 @@ func TestLoadEnvFile_FallsBackToLocalOnwatchEnv(t *testing.T) { } func TestLoadEnvFile_IgnoresNonOnwatchLocalEnv(t *testing.T) { - // Save original HOME and cwd - origHome := os.Getenv("HOME") + // Save original cwd; t.Setenv restores HOME/USERPROFILE after the test. origDir, _ := os.Getwd() - defer func() { - os.Setenv("HOME", origHome) - os.Chdir(origDir) - }() + defer os.Chdir(origDir) // Create temp directory with NO ~/.onwatch/.env tmpDir := t.TempDir() - os.Setenv("HOME", tmpDir) + t.Setenv("HOME", tmpDir) + t.Setenv("USERPROFILE", tmpDir) // Create local .env WITHOUT onwatch-specific keys (generic env file) localDir := filepath.Join(tmpDir, "project") @@ -1381,8 +1409,8 @@ func TestLoadEnvFile_IgnoresNonOnwatchLocalEnv(t *testing.T) { } // Clear env and load - os.Clearenv() - os.Setenv("HOME", tmpDir) + clearTestEnv() + setTestHome(tmpDir) loadEnvFile() // Verify the local .env was NOT loaded (because it's not onwatch-specific) @@ -1437,9 +1465,9 @@ func TestConfig_CodexShowAvailable(t *testing.T) { } func TestConfig_LogFormat_DefaultsToText(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -1451,10 +1479,10 @@ func TestConfig_LogFormat_DefaultsToText(t *testing.T) { } func TestConfig_LogFormat_LoadsFromEnv(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ONWATCH_LOG_FORMAT", "json") - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { @@ -1466,10 +1494,10 @@ func TestConfig_LogFormat_LoadsFromEnv(t *testing.T) { } func TestConfig_LogFormat_FlagOverridesEnv(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") os.Setenv("ONWATCH_LOG_FORMAT", "text") - defer os.Clearenv() + defer clearTestEnv() cfg, err := loadWithArgs([]string{"--log-format", "json"}) if err != nil { @@ -1513,12 +1541,12 @@ func TestConfig_LogFormat_AliasesAndCaseInsensitive(t *testing.T) { for _, tt := range tests { t.Run("input_"+tt.input, func(t *testing.T) { - os.Clearenv() + clearTestEnv() os.Setenv("SYNTHETIC_API_KEY", "syn_test_key") if tt.input != "" { os.Setenv("ONWATCH_LOG_FORMAT", tt.input) } - defer os.Clearenv() + defer clearTestEnv() cfg, err := Load() if err != nil { diff --git a/internal/config/opencode_go_test.go b/internal/config/opencode_go_test.go index a9df6a46..a995fabc 100644 --- a/internal/config/opencode_go_test.go +++ b/internal/config/opencode_go_test.go @@ -33,3 +33,18 @@ func TestOpenCodeGoConfiguredLegacyPairStillWorks(t *testing.T) { t.Fatal("legacy workspace+cookie detection changed") } } + +// Base-URL overrides point the balance and OpenCode clients at a proxy or a +// mock server (the e2e suite uses them). +func TestLoadProviderBaseURLOverrides(t *testing.T) { + t.Setenv("OPENCODE_GO_BASE_URL", " http://127.0.0.1:19212 ") + t.Setenv("DEEPSEEK_BASE_URL", "http://127.0.0.1:19213") + t.Setenv("MOONSHOT_BASE_URL", "http://127.0.0.1:19214") + cfg, err := Load() + if err != nil { + t.Fatalf("Load: %v", err) + } + if cfg.OpenCodeGoBaseURL != "http://127.0.0.1:19212" || cfg.DeepSeekBaseURL != "http://127.0.0.1:19213" || cfg.MoonshotBaseURL != "http://127.0.0.1:19214" { + t.Fatalf("base URLs = %q %q %q", cfg.OpenCodeGoBaseURL, cfg.DeepSeekBaseURL, cfg.MoonshotBaseURL) + } +} diff --git a/internal/menubar/assets.go b/internal/menubar/assets.go index 31543556..61df80fa 100644 --- a/internal/menubar/assets.go +++ b/internal/menubar/assets.go @@ -1,8 +1,8 @@ package menubar import ( - "encoding/json" "embed" + "encoding/json" "fmt" "io/fs" "strings" diff --git a/internal/menubar/browser_access_test.go b/internal/menubar/browser_access_test.go index 1adf8e20..a957fd47 100644 --- a/internal/menubar/browser_access_test.go +++ b/internal/menubar/browser_access_test.go @@ -45,7 +45,7 @@ func TestBlockedBrowserRootReportsUnreadable(t *testing.T) { } func TestBrowserDataRootsCoverChromiumAndFirefox(t *testing.T) { - roots := browserDataRoots("/Users/example") + roots := browserDataRoots(t.TempDir()) if len(roots) == 0 { t.Skip("no browser data roots on this platform") } diff --git a/internal/menubar/session.go b/internal/menubar/session.go index ffae8f00..7d5df971 100644 --- a/internal/menubar/session.go +++ b/internal/menubar/session.go @@ -2,7 +2,7 @@ package menubar import ( "os" - "path/filepath" + "path" "runtime" "strings" ) @@ -54,8 +54,10 @@ func linuxSessionAvailable(getenv func(string) string, exists func(string) bool) if strings.TrimSpace(getenv("DBUS_SESSION_BUS_ADDRESS")) != "" { return true } + // XDG_RUNTIME_DIR is always a slash-separated Linux path, so join with + // path (not filepath) to keep this check host-independent. if dir := strings.TrimSpace(getenv("XDG_RUNTIME_DIR")); dir != "" { - return exists(filepath.Join(dir, "bus")) + return exists(path.Join(dir, "bus")) } return false } diff --git a/internal/procscan/procscan.go b/internal/procscan/procscan.go index 676dc0af..89d6445a 100644 --- a/internal/procscan/procscan.go +++ b/internal/procscan/procscan.go @@ -49,7 +49,7 @@ func RunningContext(ctx context.Context, windowsImage string, match func(cmdline } // tasklist always exits 0; findstr verifies a real match. query := `tasklist /FI "IMAGENAME eq ` + windowsImage + `" /NH 2>nul | findstr /I "` + windowsImage + `"` - return exec.CommandContext(ctx, "cmd", "/C", query).Run() == nil + return runCommandContext(ctx, "cmd", "/C", query) == nil } if match == nil { return false @@ -81,12 +81,20 @@ func validWindowsImage(name string) bool { return true } -// execCommandContext runs the process listing. A variable so tests can assert -// the deadline the scan actually receives. +// execCommandContext runs the unix process listing and returns its output. A +// variable so tests can assert the deadline the scan actually receives. var execCommandContext = func(ctx context.Context, name string, args ...string) ([]byte, error) { return exec.CommandContext(ctx, name, args...).Output() } +// runCommandContext runs the Windows tasklist|findstr check, which only needs +// the exit status. Run (not Output) leaves stdout on the null device, so no +// pipe is held open by the tasklist/findstr grandchildren after the deadline +// kills cmd.exe. A variable so tests can assert the deadline on Windows too. +var runCommandContext = func(ctx context.Context, name string, args ...string) error { + return exec.CommandContext(ctx, name, args...).Run() +} + // Scan reports whether any line of a process listing satisfies match. func Scan(psOutput []byte, match func(cmdline string) bool) bool { if match == nil { diff --git a/internal/procscan/procscan_test.go b/internal/procscan/procscan_test.go index db7ab5fe..ba18c982 100644 --- a/internal/procscan/procscan_test.go +++ b/internal/procscan/procscan_test.go @@ -82,13 +82,20 @@ func TestRunningContextBoundsAnUnboundedCallerContext(t *testing.T) { t.Fatal("test precondition: caller context must have no deadline") } + // Stub both seams: unix lists processes via ps (execCommandContext), and + // Windows runs tasklist|findstr (runCommandContext). Whichever the host + // uses must receive a context bounded by ScanTimeout. var seen context.Context - restore := execCommandContext + restoreExec, restoreRun := execCommandContext, runCommandContext execCommandContext = func(c context.Context, name string, args ...string) ([]byte, error) { seen = c return nil, context.Canceled } - t.Cleanup(func() { execCommandContext = restore }) + runCommandContext = func(c context.Context, name string, args ...string) error { + seen = c + return context.Canceled + } + t.Cleanup(func() { execCommandContext, runCommandContext = restoreExec, restoreRun }) RunningContext(ctx, "x.exe", func(string) bool { return false }) if seen == nil { diff --git a/internal/service/launchd_test.go b/internal/service/launchd_test.go index 8749b096..315441c4 100644 --- a/internal/service/launchd_test.go +++ b/internal/service/launchd_test.go @@ -56,7 +56,9 @@ func TestPlistPathUsesLaunchAgents(t *testing.T) { if err != nil { t.Fatalf("PlistPath: %v", err) } - want := "/Users/tester/Library/LaunchAgents/dev.onllm.onwatch.plist" + // PlistPath joins with the host separator; the path is only used on + // macOS, but the helper must stay correct when the suite runs elsewhere. + want := filepath.Join("/Users/tester", "Library", "LaunchAgents", "dev.onllm.onwatch.plist") if got != want { t.Errorf("PlistPath = %q, want %q", got, want) } diff --git a/internal/store/connection_pragmas_test.go b/internal/store/connection_pragmas_test.go index 2b581427..a0fe507f 100644 --- a/internal/store/connection_pragmas_test.go +++ b/internal/store/connection_pragmas_test.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "path/filepath" + "strings" "testing" ) @@ -42,3 +43,19 @@ func TestConnectionPragmasApplyToEveryConnection(t *testing.T) { } } } + +func TestSQLiteDSNKeepsCallerQueryAndAddsMissingPragmas(t *testing.T) { + if got := sqliteDSN("/data/onwatch.db"); got != "/data/onwatch.db?"+sqliteConnectionPragmas { + t.Fatalf("plain path DSN = %q", got) + } + got := sqliteDSN("file:/data/onwatch.db?_txlock=immediate&_pragma=busy_timeout(9000)") + if !strings.HasPrefix(got, "file:/data/onwatch.db?_txlock=immediate&_pragma=busy_timeout(9000)&") { + t.Fatalf("caller query not kept: %q", got) + } + if strings.Contains(got, "busy_timeout(5000)") { + t.Fatalf("caller busy_timeout overridden: %q", got) + } + if !strings.Contains(got, "_pragma=foreign_keys(1)") { + t.Fatalf("missing pragmas not added: %q", got) + } +} diff --git a/internal/store/grok_store.go b/internal/store/grok_store.go index 1149bdef..84cb5d5d 100644 --- a/internal/store/grok_store.go +++ b/internal/store/grok_store.go @@ -185,13 +185,13 @@ func (s *Store) QueryGrokRange(accountID int64, start, end time.Time, limit ...i byID := make(map[int64]*api.GrokSnapshot) for rows.Next() { var ( - id int64 - capturedAt string - email, teamID, loginMethod sql.NullString - rawJSON string - accID int64 - qName, qResets, qStatus sql.NullString - qUtil sql.NullFloat64 + id int64 + capturedAt string + email, teamID, loginMethod sql.NullString + rawJSON string + accID int64 + qName, qResets, qStatus sql.NullString + qUtil sql.NullFloat64 ) if err := rows.Scan(&id, &capturedAt, &email, &teamID, &loginMethod, &rawJSON, &accID, &qName, &qUtil, &qResets, &qStatus); err != nil { @@ -234,14 +234,14 @@ func (s *Store) QueryGrokRange(accountID int64, start, end time.Time, limit ...i // GrokResetCycle mirrors the reset cycle for the grok provider. type GrokResetCycle struct { - ID int64 - AccountID int64 - QuotaName string - CycleStart time.Time - CycleEnd *time.Time - ResetsAt *time.Time - PeakUtilization float64 - TotalDelta float64 + ID int64 + AccountID int64 + QuotaName string + CycleStart time.Time + CycleEnd *time.Time + ResetsAt *time.Time + PeakUtilization float64 + TotalDelta float64 } // InsertGrokResetCycle creates a new cycle row. diff --git a/internal/store/minimax_store_error_test.go b/internal/store/minimax_store_error_test.go index ec383f8b..b2716d69 100644 --- a/internal/store/minimax_store_error_test.go +++ b/internal/store/minimax_store_error_test.go @@ -53,7 +53,7 @@ func TestClosedDB_MiniMaxStoreFunctions(t *testing.T) { }) t.Run("QueryActiveMiniMaxCycle", func(t *testing.T) { - _, err := s.QueryActiveMiniMaxCycle("MiniMax-M2", 2) + _, err := s.QueryActiveMiniMaxCycle("MiniMax-M2", 2) if err == nil { t.Fatal("expected error from QueryActiveMiniMaxCycle on closed DB") } diff --git a/internal/store/opencode_store_test.go b/internal/store/opencode_store_test.go index a706d48b..834379cd 100644 --- a/internal/store/opencode_store_test.go +++ b/internal/store/opencode_store_test.go @@ -54,7 +54,7 @@ func TestOpenCodeStore_QueryRangeLoadsQuotas(t *testing.T) { snap := &api.OpenCodeSnapshot{ CapturedAt: base.Add(time.Duration(i) * time.Minute), Quotas: []api.OpenCodeQuota{ - {Name: "five_hour", Utilization: float64(i+1)*10, Format: api.OpenCodeQuotaFormatPercent}, + {Name: "five_hour", Utilization: float64(i+1) * 10, Format: api.OpenCodeQuotaFormatPercent}, }, } if _, err := s.InsertOpenCodeSnapshot(snap); err != nil { diff --git a/internal/store/store.go b/internal/store/store.go index 641c8174..f4c865cd 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -159,6 +159,27 @@ func preflightDatabasePath(dbPath string) error { // database file itself, so it stays with the one-off pragmas in New. const sqliteConnectionPragmas = "_pragma=busy_timeout(5000)&_pragma=foreign_keys(1)&_pragma=synchronous(NORMAL)&_pragma=cache_size(-500)" +// sqliteDSN appends the connection pragmas to dbPath. A path that already +// carries a query string keeps it, and only pragmas it does not set itself +// are added, so a caller's parameters cannot silently drop busy_timeout. +func sqliteDSN(dbPath string) string { + sep, existing := "?", "" + if i := strings.IndexByte(dbPath, '?'); i >= 0 { + sep, existing = "&", strings.ToLower(dbPath[i+1:]) + } + var add []string + for _, p := range strings.Split(sqliteConnectionPragmas, "&") { + name := p[:strings.IndexByte(p, '(')] // "_pragma=busy_timeout" + if !strings.Contains(existing, name) { + add = append(add, p) + } + } + if len(add) == 0 { + return dbPath + } + return dbPath + sep + strings.Join(add, "&") +} + // New creates a new Store with the given database path func New(dbPath string) (*Store, error) { if err := preflightDatabasePath(dbPath); err != nil { @@ -169,13 +190,8 @@ func New(dbPath string) (*Store, error) { // every connection the pool opens. Applied with db.Exec below they reached // only one connection, leaving the pool's second one with foreign_keys off // (ON DELETE CASCADE did nothing there), busy_timeout 0 (contended writes - // failed with SQLITE_BUSY at once) and SQLite's default 2MB page cache. A - // path that already carries a query string is left untouched. - dsn := dbPath - if !strings.Contains(dbPath, "?") { - dsn += "?" + sqliteConnectionPragmas - } - db, err := sql.Open("sqlite", dsn) + // failed with SQLITE_BUSY at once) and SQLite's default 2MB page cache. + db, err := sql.Open("sqlite", sqliteDSN(dbPath)) if err != nil { return nil, fmt.Errorf("failed to open database: %w", err) } diff --git a/internal/testutil/cmd/mockserver/main.go b/internal/testutil/cmd/mockserver/main.go index 1bf6ab65..33606c54 100644 --- a/internal/testutil/cmd/mockserver/main.go +++ b/internal/testutil/cmd/mockserver/main.go @@ -11,6 +11,9 @@ // --syn-key Expected Synthetic API key (default: syn_test_e2e_key) // --zai-key Expected Z.ai API key (default: zai_test_e2e_key) // --anth-token Expected Anthropic OAuth token (default: anth_test_e2e_token) +// +// It also mocks the OpenCode Go console status API and the DeepSeek and +// Moonshot balance APIs (see providers.go). package main import ( @@ -48,31 +51,51 @@ func main() { log.Fatalf("failed to listen on %s: %v", addr, err) } + log.Printf("mock server listening on http://localhost:%d", ln.Addr().(*net.TCPAddr).Port) + log.Printf(" Synthetic key: %s", *synKey) + log.Printf(" Z.ai key: %s", *zaiKey) + log.Printf(" Anthropic tok: %s", *anthToken) + + // Serve until interrupted. On Windows os.Interrupt arrives as a console + // Ctrl+C/Ctrl+Break; the e2e harness instead stops the process with + // TerminateProcess, which needs no handling here. + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + if err := serve(ctx, ln, srv.mux); err != nil { + log.Fatalf("server error: %v", err) + } +} + +// serve runs handler on ln until ctx is cancelled, then shuts the server down +// gracefully. It returns nil on a clean shutdown. +func serve(ctx context.Context, ln net.Listener, handler http.Handler) error { httpSrv := &http.Server{ - Handler: srv.mux, + Handler: handler, ReadTimeout: 5 * time.Second, WriteTimeout: 5 * time.Second, } + errCh := make(chan error, 1) go func() { - log.Printf("mock server listening on http://localhost:%d", *port) - log.Printf(" Synthetic key: %s", *synKey) - log.Printf(" Z.ai key: %s", *zaiKey) - log.Printf(" Anthropic tok: %s", *anthToken) - if err := httpSrv.Serve(ln); err != nil && err != http.ErrServerClosed { - log.Fatalf("server error: %v", err) - } + errCh <- httpSrv.Serve(ln) }() - // Wait for interrupt - sigCh := make(chan os.Signal, 1) - signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) - <-sigCh + select { + case err := <-errCh: + return fmt.Errorf("serve: %w", err) + case <-ctx.Done(): + } log.Println("shutting down...") - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - httpSrv.Shutdown(ctx) + if err := httpSrv.Shutdown(shutdownCtx); err != nil { + return fmt.Errorf("shutdown: %w", err) + } + if err := <-errCh; err != nil && err != http.ErrServerClosed { + return fmt.Errorf("serve: %w", err) + } + return nil } // standaloneServer wraps the mock server logic without using httptest.Server, @@ -99,6 +122,8 @@ type standaloneServer struct { anthropicError atomic.Int32 anthropicIdx atomic.Int64 anthropicCount atomic.Int64 + + providers providerMocks } func newStandaloneServer(synKey, zaiKey, anthToken string) *standaloneServer { @@ -119,6 +144,7 @@ func newStandaloneServer(synKey, zaiKey, anthToken string) *standaloneServer { srv.mux.HandleFunc("/admin/error", srv.handleAdminError) srv.mux.HandleFunc("/admin/requests", srv.handleAdminRequests) srv.mux.HandleFunc("/admin/reset", srv.handleAdminReset) + srv.providers.register(srv.mux) return srv } @@ -303,6 +329,9 @@ func (s *standaloneServer) handleAdminRequests(w http.ResponseWriter, _ *http.Re "zai": s.zaiCount.Load(), "anthropic": s.anthropicCount.Load(), } + for name, n := range s.providers.counts() { + counts[name] = n + } w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(counts) @@ -317,6 +346,7 @@ func (s *standaloneServer) handleAdminReset(w http.ResponseWriter, r *http.Reque s.syntheticError.Store(0) s.zaiError.Store(0) s.anthropicError.Store(0) + s.providers.reset() s.syntheticCount.Store(0) s.zaiCount.Store(0) s.anthropicCount.Store(0) diff --git a/internal/testutil/cmd/mockserver/main_test.go b/internal/testutil/cmd/mockserver/main_test.go index d5b382dd..fafdb3af 100644 --- a/internal/testutil/cmd/mockserver/main_test.go +++ b/internal/testutil/cmd/mockserver/main_test.go @@ -2,12 +2,17 @@ package main import ( "bytes" + "context" "encoding/json" "io" + "net" "net/http" "net/http/httptest" "os" "os/exec" + "runtime" + "strconv" + "syscall" "testing" "time" ) @@ -192,26 +197,112 @@ func TestStandaloneServer_AdminEndpointsValidateMethodBodyAndProvider(t *testing } } +// TestServe_ShutsDownCleanlyOnCancel covers the serve loop in-process on +// every OS: it answers real HTTP requests and returns nil once its context is +// cancelled (which is what SIGINT/SIGTERM do via signal.NotifyContext). +func TestServe_ShutsDownCleanlyOnCancel(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + srv := newStandaloneServer("syn-key", "zai-key", "anth-token") + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan error, 1) + go func() { done <- serve(ctx, ln, srv.mux) }() + + resp, err := http.Get("http://" + ln.Addr().String() + "/admin/requests") + if err != nil { + t.Fatalf("GET /admin/requests: %v", err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("expected /admin/requests 200, got %d", resp.StatusCode) + } + + cancel() + select { + case err := <-done: + if err != nil { + t.Fatalf("serve returned %v, want nil on clean shutdown", err) + } + case <-time.After(10 * time.Second): + t.Fatal("serve did not return after its context was cancelled") + } +} + +// TestMain_StartsAndShutsDownOnSignal runs main() in a child process and stops +// it the way tests/e2e/conftest.py does: SIGTERM on unix (graceful shutdown, +// clean exit) and TerminateProcess on Windows, where a child cannot be sent +// os.Interrupt. On both it asserts the binary really serves before it stops. func TestMain_StartsAndShutsDownOnSignal(t *testing.T) { if os.Getenv("ONWATCH_MOCKSERVER_MAIN_HELPER") == "1" { - os.Args = []string{"mockserver", "-port=0", "-syn-key=helper-syn", "-zai-key=helper-zai", "-anth-token=helper-anth"} + os.Args = []string{"mockserver", "-port=" + os.Getenv("ONWATCH_MOCKSERVER_MAIN_PORT"), "-syn-key=helper-syn", "-zai-key=helper-zai", "-anth-token=helper-anth"} main() return } - cmd := exec.Command(os.Args[0], "-test.run=TestMain_StartsAndShutsDownOnSignal") - cmd.Env = append(os.Environ(), "ONWATCH_MOCKSERVER_MAIN_HELPER=1") + port := freePort(t) + cmd := exec.Command(os.Args[0], "-test.run=^TestMain_StartsAndShutsDownOnSignal$") + cmd.Env = append(os.Environ(), "ONWATCH_MOCKSERVER_MAIN_HELPER=1", "ONWATCH_MOCKSERVER_MAIN_PORT="+strconv.Itoa(port)) cmd.Stdout = io.Discard cmd.Stderr = io.Discard if err := cmd.Start(); err != nil { t.Fatalf("start helper process: %v", err) } - - time.Sleep(300 * time.Millisecond) - if err := cmd.Process.Signal(os.Interrupt); err != nil { + waitErr := make(chan error, 1) + go func() { waitErr <- cmd.Wait() }() + t.Cleanup(func() { _ = cmd.Process.Kill() }) + + url := "http://127.0.0.1:" + strconv.Itoa(port) + "/admin/requests" + deadline := time.Now().Add(10 * time.Second) + for { + resp, err := http.Get(url) + if err == nil { + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("expected /admin/requests 200, got %d", resp.StatusCode) + } + break + } + select { + case err := <-waitErr: + t.Fatalf("helper process exited before serving: %v", err) + default: + } + if time.Now().After(deadline) { + t.Fatalf("helper process never served %s: %v", url, err) + } + time.Sleep(50 * time.Millisecond) + } + + if runtime.GOOS == "windows" { + if err := cmd.Process.Kill(); err != nil { + t.Fatalf("terminate helper process: %v", err) + } + } else if err := cmd.Process.Signal(syscall.SIGTERM); err != nil { t.Fatalf("signal helper process: %v", err) } - if err := cmd.Wait(); err != nil { - t.Fatalf("wait helper process: %v", err) + + select { + case err := <-waitErr: + // TerminateProcess forces a non-zero exit; only unix exits cleanly. + if runtime.GOOS != "windows" && err != nil { + t.Fatalf("helper process did not exit cleanly on SIGTERM: %v", err) + } + case <-time.After(10 * time.Second): + t.Fatal("helper process did not exit after being stopped") + } +} + +// freePort returns a TCP port that was free a moment ago. +func freePort(t *testing.T) int { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) } + defer ln.Close() + return ln.Addr().(*net.TCPAddr).Port } diff --git a/internal/testutil/cmd/mockserver/providers.go b/internal/testutil/cmd/mockserver/providers.go new file mode 100644 index 00000000..b912ea57 --- /dev/null +++ b/internal/testutil/cmd/mockserver/providers.go @@ -0,0 +1,95 @@ +package main + +import ( + "net/http" + "strings" + "sync/atomic" +) + +// Mocks for the OpenCode Go console status API and the DeepSeek and Moonshot +// balance APIs. onWatch reaches them through OPENCODE_GO_BASE_URL, +// DEEPSEEK_BASE_URL and MOONSHOT_BASE_URL. + +const ( + e2eOpenCodeSession = "oc_session_e2e" + e2eOpenCodeWorkspace = "wrk_e2e" + e2eBalanceKey = "sk_balance_e2e" +) + +// Weekly: 384204992 / 3000000000 = 12.8%; monthly: 6.4%. Amounts are +// micro-cents, sent as strings like the live API. +const openCodeGoStatusBody = `{"access":{"endsAt":"2099-01-01T00:00:00.000Z","meters":{ +"fiveHour":{"resetsAt":null,"limitMicroCents":"1200000000","usedMicroCents":"0"}, +"week":{"resetsAt":"2099-01-01T00:00:00.000Z","limitMicroCents":"3000000000","usedMicroCents":"384204992"}, +"month":{"limitMicroCents":"6000000000","usedMicroCents":"384204992"}}}}` + +const deepSeekBalanceBody = `{"is_available":true,"balance_infos":[{"currency":"USD","total_balance":"4.02","granted_balance":"0.32","topped_up_balance":"3.70"}]}` + +const moonshotBalanceBody = `{"code":0,"data":{"available_balance":19.47,"voucher_balance":5,"cash_balance":14.47}}` + +type providerMocks struct { + openCodeCount atomic.Int64 + deepSeekCount atomic.Int64 + moonshotCount atomic.Int64 +} + +func (p *providerMocks) register(mux *http.ServeMux) { + mux.HandleFunc("/console/api/go/status", p.handleOpenCodeGoStatus) + mux.HandleFunc("/user/balance", p.handleDeepSeekBalance) + mux.HandleFunc("/v1/users/me/balance", p.handleMoonshotBalance) +} + +func (p *providerMocks) counts() map[string]int64 { + return map[string]int64{ + "opencode": p.openCodeCount.Load(), + "deepseek": p.deepSeekCount.Load(), + "moonshot": p.moonshotCount.Load(), + } +} + +func (p *providerMocks) reset() { + p.openCodeCount.Store(0) + p.deepSeekCount.Store(0) + p.moonshotCount.Store(0) +} + +// handleOpenCodeGoStatus mirrors the console: a session cookie needs the +// workspace as x-org-id, and the retired "auth" cookie is rejected. +func (p *providerMocks) handleOpenCodeGoStatus(w http.ResponseWriter, r *http.Request) { + p.openCodeCount.Add(1) + w.Header().Set("Content-Type", "application/json") + bearer := r.Header.Get("Authorization") == "Bearer "+e2eBalanceKey + session := strings.Contains(r.Header.Get("Cookie"), "__Host-console_session="+e2eOpenCodeSession) + switch { + case bearer: + case !session: + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"_tag":"Unauthorized"}`)) + return + case r.Header.Get("x-org-id") != e2eOpenCodeWorkspace: + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"_tag":"OrgRequired","message":"x-org-id is required"}`)) + return + } + _, _ = w.Write([]byte(openCodeGoStatusBody)) +} + +func (p *providerMocks) handleDeepSeekBalance(w http.ResponseWriter, r *http.Request) { + p.deepSeekCount.Add(1) + writeBalance(w, r, deepSeekBalanceBody) +} + +func (p *providerMocks) handleMoonshotBalance(w http.ResponseWriter, r *http.Request) { + p.moonshotCount.Add(1) + writeBalance(w, r, moonshotBalanceBody) +} + +func writeBalance(w http.ResponseWriter, r *http.Request, body string) { + w.Header().Set("Content-Type", "application/json") + if r.Header.Get("Authorization") != "Bearer "+e2eBalanceKey { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"error":"invalid api key"}`)) + return + } + _, _ = w.Write([]byte(body)) +} diff --git a/internal/testutil/testhome/testhome.go b/internal/testutil/testhome/testhome.go new file mode 100644 index 00000000..b6af14ea --- /dev/null +++ b/internal/testutil/testhome/testhome.go @@ -0,0 +1,90 @@ +// Package testhome points a test process at a throwaway home directory so +// tests never read or write the developer's (or CI runner's) real ~/.claude, +// ~/.codex, ~/.onwatch and so on. +// +// It is a leaf package (standard library only) so that the in-package tests +// of api, agent, web and cmd/onwatch can import it without an import cycle; +// internal/testutil itself imports those packages and cannot be used there. +package testhome + +import ( + "fmt" + "os" + "path/filepath" + "testing" +) + +// OverrideEnv lists the environment variables that redirect provider +// credential or data lookups away from the home directory. A developer's shell +// value for any of them would bypass the sandbox home, so SandboxHome clears +// them all. +var OverrideEnv = []string{ + "CODEX_HOME", + "OPENCODE_HOME", + "XDG_DATA_HOME", + "XDG_CONFIG_HOME", + "KIMI_CODE_HOME", + "KIMI_CODE_CREDENTIALS", + "KIMI_CREDENTIALS", + "MUSE_AUTH_PATH", + "COMMANDCODE_AUTH_PATH", + "GROK_HOME", +} + +// SandboxHome is meant for TestMain. It creates an empty temp home, points +// HOME, USERPROFILE and LOCALAPPDATA at it, and unsets OverrideEnv plus any +// extraUnset variables. os.UserHomeDir reads HOME on Unix but USERPROFILE on +// Windows, and onWatch's Windows state lives under LOCALAPPDATA, so all three +// move. LOCALAPPDATA is set on every OS so behavior is uniform. +// +// Callers that must record the real home first (api.SetTestMode(true)) have to +// do so before calling SandboxHome. cleanup removes the sandbox. +func SandboxHome(extraUnset ...string) (home string, cleanup func(), err error) { + for _, env := range OverrideEnv { + os.Unsetenv(env) + } + for _, env := range extraUnset { + os.Unsetenv(env) + } + + home, err = os.MkdirTemp("", "onwatch-test-home-") + if err != nil { + return "", func() {}, fmt.Errorf("create sandbox home: %w", err) + } + if err := setHomeEnv(os.Setenv, home); err != nil { + os.RemoveAll(home) + return "", func() {}, err + } + return home, func() { os.RemoveAll(home) }, nil +} + +// SetTestHome points the user's home directory at dir for the duration of the +// test via t.Setenv: HOME (Unix, Git Bash), USERPROFILE (os.UserHomeDir on +// Windows) and LOCALAPPDATA (/AppData/Local). An empty dir clears HOME +// and USERPROFILE, which makes os.UserHomeDir fail on every platform. +func SetTestHome(t testing.TB, dir string) { + t.Helper() + _ = setHomeEnv(func(k, v string) error { t.Setenv(k, v); return nil }, dir) +} + +// LocalAppData returns the LOCALAPPDATA value SandboxHome and SetTestHome use +// for home ("" for an empty home). +func LocalAppData(home string) string { + if home == "" { + return "" + } + return filepath.Join(home, "AppData", "Local") +} + +func setHomeEnv(setenv func(k, v string) error, home string) error { + for k, v := range map[string]string{ + "HOME": home, + "USERPROFILE": home, + "LOCALAPPDATA": LocalAppData(home), + } { + if err := setenv(k, v); err != nil { + return fmt.Errorf("set %s: %w", k, err) + } + } + return nil +} diff --git a/internal/testutil/testhome/testhome_test.go b/internal/testutil/testhome/testhome_test.go new file mode 100644 index 00000000..75d9d309 --- /dev/null +++ b/internal/testutil/testhome/testhome_test.go @@ -0,0 +1,65 @@ +package testhome + +import ( + "os" + "path/filepath" + "testing" +) + +func TestSandboxHome(t *testing.T) { + for _, env := range append([]string{"HOME", "USERPROFILE", "LOCALAPPDATA", "ONWATCH_EXTRA_TEST_VAR"}, OverrideEnv...) { + t.Setenv(env, "/developer/value") + } + + home, cleanup, err := SandboxHome("ONWATCH_EXTRA_TEST_VAR") + if err != nil { + t.Fatalf("SandboxHome: %v", err) + } + if info, err := os.Stat(home); err != nil || !info.IsDir() { + t.Fatalf("sandbox home %q not created: %v", home, err) + } + for _, env := range []string{"HOME", "USERPROFILE"} { + if got := os.Getenv(env); got != home { + t.Errorf("%s = %q, want %q", env, got, home) + } + } + if got, want := os.Getenv("LOCALAPPDATA"), filepath.Join(home, "AppData", "Local"); got != want { + t.Errorf("LOCALAPPDATA = %q, want %q", got, want) + } + if got, err := os.UserHomeDir(); err != nil || got != home { + t.Errorf("os.UserHomeDir() = %q, %v; want %q", got, err, home) + } + for _, env := range append([]string{"ONWATCH_EXTRA_TEST_VAR"}, OverrideEnv...) { + if v, ok := os.LookupEnv(env); ok { + t.Errorf("%s still set to %q", env, v) + } + } + + cleanup() + if _, err := os.Stat(home); !os.IsNotExist(err) { + t.Errorf("cleanup left %q behind (stat err %v)", home, err) + } +} + +func TestSetTestHome(t *testing.T) { + dir := t.TempDir() + SetTestHome(t, dir) + for _, env := range []string{"HOME", "USERPROFILE"} { + if got := os.Getenv(env); got != dir { + t.Errorf("%s = %q, want %q", env, got, dir) + } + } + if got, want := os.Getenv("LOCALAPPDATA"), filepath.Join(dir, "AppData", "Local"); got != want { + t.Errorf("LOCALAPPDATA = %q, want %q", got, want) + } + if got, err := os.UserHomeDir(); err != nil || got != dir { + t.Errorf("os.UserHomeDir() = %q, %v; want %q", got, err, dir) + } +} + +func TestSetTestHome_EmptyMakesUserHomeDirFail(t *testing.T) { + SetTestHome(t, "") + if got, err := os.UserHomeDir(); err == nil { + t.Errorf("os.UserHomeDir() = %q, want an error for an empty home", got) + } +} diff --git a/internal/tracker/grok_tracker.go b/internal/tracker/grok_tracker.go index 4cf3db8e..f989d68d 100644 --- a/internal/tracker/grok_tracker.go +++ b/internal/tracker/grok_tracker.go @@ -13,7 +13,7 @@ import ( type GrokTracker struct { store *store.Store logger *slog.Logger - lastValues map[int64]map[string]float64 // account -> quota -> last util + lastValues map[int64]map[string]float64 // account -> quota -> last util lastResets map[int64]map[string]time.Time hasLast map[int64]bool diff --git a/internal/update/update.go b/internal/update/update.go index d4c27bcd..77a95757 100644 --- a/internal/update/update.go +++ b/internal/update/update.go @@ -41,6 +41,9 @@ var ( execCommand = exec.Command sleepFn = time.Sleep exitFn = os.Exit + // executablePath resolves the binary Apply replaces. Tests point it at a + // temp copy so they never overwrite the running test binary. + executablePath = os.Executable ) // UpdateInfo holds the result of a version check. @@ -354,7 +357,7 @@ func (u *Updater) Apply() error { } // Get current binary path - exePath, err := os.Executable() + exePath, err := executablePath() if err != nil { return fmt.Errorf("update.Apply: os.Executable: %w", err) } @@ -719,7 +722,7 @@ func (u *Updater) Restart() error { if exePath == "" { var err error - exePath, err = os.Executable() + exePath, err = executablePath() if err != nil { return fmt.Errorf("update.Restart: %w", err) } diff --git a/internal/update/update_more_coverage_test.go b/internal/update/update_more_coverage_test.go index 46bb25f7..c262c516 100644 --- a/internal/update/update_more_coverage_test.go +++ b/internal/update/update_more_coverage_test.go @@ -16,6 +16,8 @@ import ( "sync/atomic" "testing" "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) type updateExitPanic struct { @@ -274,11 +276,7 @@ func TestMigrateSystemdUnit_ReadAndWriteFailuresAndNoop(t *testing.T) { t.Run("missing unit file is noop", func(t *testing.T) { tmpHome := t.TempDir() - oldHome := os.Getenv("HOME") - t.Cleanup(func() { _ = os.Setenv("HOME", oldHome) }) - if err := os.Setenv("HOME", tmpHome); err != nil { - t.Fatalf("set HOME: %v", err) - } + testhome.SetTestHome(t, tmpHome) readCgroupFile = func() ([]byte, error) { return []byte("0::/user.slice/user-501.slice/user@501.service/app.slice/missing.service"), nil } @@ -287,11 +285,7 @@ func TestMigrateSystemdUnit_ReadAndWriteFailuresAndNoop(t *testing.T) { t.Run("read failure is noop", func(t *testing.T) { tmpHome := t.TempDir() - oldHome := os.Getenv("HOME") - t.Cleanup(func() { _ = os.Setenv("HOME", oldHome) }) - if err := os.Setenv("HOME", tmpHome); err != nil { - t.Fatalf("set HOME: %v", err) - } + testhome.SetTestHome(t, tmpHome) userDir := filepath.Join(tmpHome, ".config", "systemd", "user") if err := os.MkdirAll(userDir, 0o755); err != nil { t.Fatalf("mkdir user dir: %v", err) @@ -308,11 +302,7 @@ func TestMigrateSystemdUnit_ReadAndWriteFailuresAndNoop(t *testing.T) { t.Run("already up to date is noop", func(t *testing.T) { tmpHome := t.TempDir() - oldHome := os.Getenv("HOME") - t.Cleanup(func() { _ = os.Setenv("HOME", oldHome) }) - if err := os.Setenv("HOME", tmpHome); err != nil { - t.Fatalf("set HOME: %v", err) - } + testhome.SetTestHome(t, tmpHome) serviceName := "noop.service" userDir := filepath.Join(tmpHome, ".config", "systemd", "user") if err := os.MkdirAll(userDir, 0o755); err != nil { diff --git a/internal/update/update_test.go b/internal/update/update_test.go index fd1b09ef..861271c3 100644 --- a/internal/update/update_test.go +++ b/internal/update/update_test.go @@ -3,16 +3,20 @@ package update import ( "encoding/json" "fmt" + "io" "log/slog" "net/http" "net/http/httptest" "os" + "os/exec" "path/filepath" "runtime" "strings" "sync/atomic" "testing" "time" + + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) func TestCompareVersions(t *testing.T) { @@ -593,33 +597,92 @@ func TestApply_EmptyVersion(t *testing.T) { } } -func TestApply_DownloadAndReplace(t *testing.T) { - // Create a mock server that serves the release API and a binary download - var currentExe string - var err error - currentExe, err = os.Executable() +// runningBinaryHelperEnv turns a copy of this test binary into an idle +// process, so Apply can be exercised against a binary that is genuinely +// executing - the real self-update situation. Windows refuses to delete a +// running .exe, which is what drives replaceBinary's backup-rename path. +const runningBinaryHelperEnv = "ONWATCH_UPDATE_TEST_RUNNING_BINARY" + +func TestHelperRunningBinary(t *testing.T) { + if os.Getenv(runningBinaryHelperEnv) != "1" { + return + } + time.Sleep(2 * time.Minute) +} + +// startRunningBinaryCopy copies this test binary into a temp dir, starts the +// copy as an idle process and returns its path. The process is killed and +// reaped before the temp dir is removed, so Windows can delete the image. +func startRunningBinaryCopy(t *testing.T) string { + t.Helper() + self, err := os.Executable() if err != nil { t.Fatalf("os.Executable: %v", err) } + name := "onwatch" + if runtime.GOOS == "windows" { + name += ".exe" + } + exePath := filepath.Join(t.TempDir(), name) - // Read real binary magic bytes from the current executable for validation - magic := make([]byte, 8) - f, err := os.Open(currentExe) + src, err := os.Open(self) if err != nil { - t.Fatalf("open current exe: %v", err) + t.Fatalf("open test binary: %v", err) + } + defer src.Close() + dst, err := os.OpenFile(exePath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o755) + if err != nil { + t.Fatalf("create binary copy: %v", err) + } + if _, err := io.Copy(dst, src); err != nil { + dst.Close() + t.Fatalf("copy test binary: %v", err) } - f.Read(magic) - f.Close() + if err := dst.Close(); err != nil { + t.Fatalf("close binary copy: %v", err) + } + + cmd := exec.Command(exePath, "-test.run=^TestHelperRunningBinary$") + cmd.Env = append(os.Environ(), runningBinaryHelperEnv+"=1") + if err := cmd.Start(); err != nil { + t.Fatalf("start binary copy: %v", err) + } + t.Cleanup(func() { + _ = cmd.Process.Kill() + _ = cmd.Wait() + }) + return exePath +} - // Create download server serving a valid binary (using real magic bytes) +// Apply must replace the binary it is running from. The target is a running +// temp copy, never this test binary: replacing the test binary would corrupt +// the package's own executable mid-run. +func TestApply_DownloadAndReplace(t *testing.T) { + exePath := startRunningBinaryCopy(t) + oldExecutablePath := executablePath + t.Cleanup(func() { executablePath = oldExecutablePath }) + executablePath = func() (string, error) { return exePath, nil } + + // Serve a payload that passes validateBinary: this platform's magic bytes. + current, err := os.ReadFile(exePath) + if err != nil { + t.Fatalf("read binary copy: %v", err) + } + payload := append(append([]byte(nil), current[:8]...), []byte("rest-of-binary-content-padded-to-be-non-empty")...) + + wantAsset := fmt.Sprintf("/v99.0.0/onwatch-%s-%s", runtime.GOOS, runtime.GOARCH) + if runtime.GOOS == "windows" { + wantAsset += ".exe" + } dlSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Serve a file with valid magic bytes - w.Write(magic) - w.Write([]byte("rest-of-binary-content-padded-to-be-non-empty")) + if r.URL.Path != wantAsset { + http.NotFound(w, r) + return + } + w.Write(payload) })) defer dlSrv.Close() - // Create API server that returns a newer version apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { json.NewEncoder(w).Encode(githubRelease{TagName: "v99.0.0"}) })) @@ -629,15 +692,32 @@ func TestApply_DownloadAndReplace(t *testing.T) { u.apiURL = apiSrv.URL u.downloadURL = dlSrv.URL - // Apply will try to replace the current executable, which we can't really do in test. - // But we can verify it gets past the download and validation steps. - err = u.Apply() - // We expect an error because either: - // 1. The download URL format won't match the mock server, or - // 2. We can't actually replace the running test binary - // The key is we exercised more of the Apply() code path. - if err == nil { - t.Log("Apply succeeded unexpectedly (may be OK on some platforms)") + if err := u.Apply(); err != nil { + t.Fatalf("Apply() = %v", err) + } + + got, err := os.ReadFile(exePath) + if err != nil { + t.Fatalf("read replaced binary: %v", err) + } + if string(got) != string(payload) { + t.Fatalf("binary not replaced: got %d bytes, want the %d-byte download", len(got), len(payload)) + } + + leftovers, err := filepath.Glob(filepath.Join(filepath.Dir(exePath), "onwatch.tmp.*")) + if err != nil { + t.Fatalf("glob temp downloads: %v", err) + } + if len(leftovers) != 0 { + t.Fatalf("temp download left behind: %v", leftovers) + } + + wantApplied, err := filepath.EvalSymlinks(exePath) + if err != nil { + t.Fatalf("EvalSymlinks: %v", err) + } + if u.lastAppliedPath != wantApplied { + t.Fatalf("lastAppliedPath = %q, want %q", u.lastAppliedPath, wantApplied) } } @@ -1003,12 +1083,7 @@ func TestCheck_RateLimitErrorMessage(t *testing.T) { func TestFindUnitFile_UserLevelPath(t *testing.T) { serviceName := "onwatch-user-level-test.service" tmpHome := t.TempDir() - - origHome := os.Getenv("HOME") - defer os.Setenv("HOME", origHome) - if err := os.Setenv("HOME", tmpHome); err != nil { - t.Fatalf("Setenv HOME: %v", err) - } + testhome.SetTestHome(t, tmpHome) userDir := filepath.Join(tmpHome, ".config", "systemd", "user") if err := os.MkdirAll(userDir, 0755); err != nil { @@ -1186,7 +1261,7 @@ func TestMigrateSystemdUnit_UpdatesUserUnitAndReloads(t *testing.T) { } tmpHome := t.TempDir() - t.Setenv("HOME", tmpHome) + testhome.SetTestHome(t, tmpHome) t.Setenv("INVOCATION_ID", "invocation-test-id") unitDir := filepath.Join(tmpHome, ".config", "systemd", "user") @@ -1199,22 +1274,17 @@ func TestMigrateSystemdUnit_UpdatesUserUnitAndReloads(t *testing.T) { t.Fatalf("WriteFile unitPath: %v", err) } - binDir := t.TempDir() - markerFile := filepath.Join(binDir, "systemctl.called") - scriptPath := filepath.Join(binDir, "systemctl") - script := "#!/bin/sh\n" + - "echo \"$@\" >> \"" + markerFile + "\"\n" + - "exit 0\n" - if err := os.WriteFile(scriptPath, []byte(script), 0755); err != nil { - t.Fatalf("WriteFile systemctl stub: %v", err) + // Record the reload through the execCommand hook rather than a PATH stub: + // a shell-script stub is not executable on Windows. The recorded command + // runs this test binary with no tests selected, which exits 0 everywhere. + oldExecCommand := execCommand + t.Cleanup(func() { execCommand = oldExecCommand }) + var calls []string + execCommand = func(name string, args ...string) *exec.Cmd { + calls = append(calls, strings.Join(append([]string{name}, args...), " ")) + return exec.Command(os.Args[0], "-test.run=^$") } - pathSep := ":" - if runtime.GOOS == "windows" { - pathSep = ";" - } - t.Setenv("PATH", binDir+pathSep+os.Getenv("PATH")) - MigrateSystemdUnit(slog.Default()) updatedBytes, err := os.ReadFile(unitPath) @@ -1229,12 +1299,8 @@ func TestMigrateSystemdUnit_UpdatesUserUnitAndReloads(t *testing.T) { t.Fatalf("expected RestartSec=5 in unit file, got:\n%s", updated) } - calls, err := os.ReadFile(markerFile) - if err != nil { - t.Fatalf("expected systemctl to be called, read marker: %v", err) - } - if !strings.Contains(string(calls), "--user daemon-reload") { - t.Fatalf("expected user-level daemon-reload call, got: %s", string(calls)) + if len(calls) != 1 || calls[0] != "systemctl --user daemon-reload" { + t.Fatalf("expected one user-level daemon-reload call, got: %q", calls) } } diff --git a/internal/web/balance_providers_test.go b/internal/web/balance_providers_test.go new file mode 100644 index 00000000..1a6faa5f --- /dev/null +++ b/internal/web/balance_providers_test.go @@ -0,0 +1,205 @@ +package web + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/onllm-dev/onwatch/v2/internal/api" + "github.com/onllm-dev/onwatch/v2/internal/config" + "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/tracker" +) + +// Issue #137: the DeepSeek and Moonshot tabs rendered no statistics. + +func newBalanceTestHandler(t *testing.T) (*Handler, *store.Store) { + t.Helper() + s, err := store.New(":memory:") + if err != nil { + t.Fatalf("store.New: %v", err) + } + t.Cleanup(func() { s.Close() }) + cfg := &config.Config{ + DeepSeekAPIKey: "sk-test", + MoonshotAPIKey: "sk-test", + PollInterval: 60 * time.Second, + AdminUser: "admin", + AdminPass: "test", + } + h := NewHandler(s, nil, nil, nil, cfg) + h.SetDeepSeekTracker(tracker.NewDeepSeekTracker(s, nil)) + h.SetMoonshotTracker(tracker.NewMoonshotTracker(s, nil)) + return h, s +} + +func insertDeepSeekUSD(t *testing.T, h *Handler, s *store.Store, at time.Time, total float64) { + t.Helper() + snap := &api.DeepSeekSnapshot{CapturedAt: at, IsAvailable: true, Currency: "USD", TotalBalance: total, ToppedUpBalance: total} + if _, err := s.InsertDeepSeekSnapshot(snap); err != nil { + t.Fatalf("insert DeepSeek snapshot: %v", err) + } + if err := h.deepseekTracker.Process(snap); err != nil { + t.Fatalf("track DeepSeek snapshot: %v", err) + } +} + +func getJSON(t *testing.T, h *Handler, target string, out interface{}) { + t.Helper() + routes := map[string]http.HandlerFunc{ + "/api/insights": h.Insights, + "/api/summary": h.Summary, + "/api/cycles": h.Cycles, + "/api/cycle-overview": h.CycleOverview, + "/api/logging-history": h.LoggingHistory, + } + req := httptest.NewRequest(http.MethodGet, target, nil) + route, ok := routes[req.URL.Path] + if !ok { + t.Fatalf("no route for %s", req.URL.Path) + } + rec := httptest.NewRecorder() + route(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("GET %s = %d: %s", target, rec.Code, rec.Body.String()) + } + if err := json.Unmarshal(rec.Body.Bytes(), out); err != nil { + t.Fatalf("GET %s: decode: %v", target, err) + } +} + +// A USD account must not be looked up as CNY when no currency is requested. +func TestDeepSeekDefaultsToTheTrackedCurrency(t *testing.T) { + h, s := newBalanceTestHandler(t) + now := time.Now().UTC() + insertDeepSeekUSD(t, h, s, now.Add(-2*time.Hour), 4.10) + insertDeepSeekUSD(t, h, s, now.Add(-time.Hour), 4.02) + + var insights insightsResponse + getJSON(t, h, "/api/insights?provider=deepseek", &insights) + if len(insights.Stats) == 0 || !strings.HasPrefix(insights.Stats[0].Value, "$") { + t.Fatalf("insights stats = %+v, want USD balance stats", insights.Stats) + } + + var summary map[string]map[string]interface{} + getJSON(t, h, "/api/summary?provider=deepseek", &summary) + if summary["balance"]["currency"] != "USD" || summary["balance"]["currentBalance"] != 4.02 { + t.Fatalf("summary balance = %v, want the USD balance 4.02", summary["balance"]) + } + + var overview map[string]interface{} + getJSON(t, h, "/api/cycle-overview?provider=deepseek&groupBy=balance", &overview) + cyclesOut, _ := overview["cycles"].([]interface{}) + if len(cyclesOut) == 0 || overview["currency"] != "USD" { + t.Fatalf("cycle overview = %v, want the USD cycle", overview) + } + if first, _ := cyclesOut[0].(map[string]interface{}); first["cycleId"] == nil { + t.Fatalf("cycle row %v has no cycleId for the dashboard table", first) + } + // Balance cycles carry only a spend delta: no per-quota columns, and the + // combined view must agree with the provider view. + if names, _ := overview["quotaNames"].([]interface{}); len(names) != 0 { + t.Fatalf("cycle overview quotaNames = %v, want none", names) + } + if both := h.deepseekCycleOverview(h.deepseekCurrency("")); len(both["quotaNames"].([]string)) != 0 { + t.Fatalf("combined overview quotaNames = %v, want none", both["quotaNames"]) + } + + var logs map[string]interface{} + getJSON(t, h, "/api/logging-history?provider=deepseek&range=1", &logs) + if logs["currency"] != "USD" { + t.Fatalf("logging history currency = %v, want USD", logs["currency"]) + } + + var cycles []interface{} + getJSON(t, h, "/api/cycles?provider=deepseek", &cycles) + if len(cycles) == 0 { + t.Fatal("cycles empty, want the USD cycle") + } +} + +// Logging history must use the shared crossQuotas row shape the dashboard +// table reads. +func TestBalanceLoggingHistoryUsesCrossQuotas(t *testing.T) { + h, s := newBalanceTestHandler(t) + now := time.Now().UTC().Add(-time.Hour) + insertDeepSeekUSD(t, h, s, now, 4.02) + if _, err := s.InsertMoonshotSnapshot(&api.MoonshotSnapshot{CapturedAt: now, AvailableBalance: 19.47, VoucherBalance: 5, CashBalance: 14.47}); err != nil { + t.Fatalf("insert Moonshot snapshot: %v", err) + } + + for provider, want := range map[string]map[string]float64{ + "deepseek": {"total_balance": 4.02, "granted_balance": 0, "topped_up_balance": 4.02}, + "moonshot": {"available_balance": 19.47, "voucher_balance": 5, "cash_balance": 14.47}, + } { + var resp struct { + QuotaNames []string `json:"quotaNames"` + Logs []struct { + CrossQuotas []struct { + Name string `json:"name"` + Value float64 `json:"value"` + } `json:"crossQuotas"` + } `json:"logs"` + } + getJSON(t, h, "/api/logging-history?provider="+provider+"&range=1", &resp) + if len(resp.QuotaNames) != len(want) { + t.Fatalf("%s quotaNames = %v, want %d balance fields", provider, resp.QuotaNames, len(want)) + } + if len(resp.Logs) != 1 { + t.Fatalf("%s logs = %d, want 1", provider, len(resp.Logs)) + } + got := map[string]float64{} + for _, cq := range resp.Logs[0].CrossQuotas { + got[cq.Name] = cq.Value + } + for name, v := range want { + if gv, ok := got[name]; !ok || gv != v { + t.Errorf("%s %s = %v (present %v), want %v", provider, name, gv, ok, v) + } + } + } +} + +func TestBalanceProvidersAreWiredIntoTheDashboard(t *testing.T) { + js := readStaticFile(t, "static/app.js") + for _, want := range []string{ + "document.getElementById('quota-grid-deepseek')", + "document.getElementById('quota-grid-moonshot')", + "renderBalanceCards(", + } { + if !strings.Contains(js, want) { + t.Errorf("app.js is missing %q", want) + } + } +} + +// Logging rows carry one currency label, so rows in another currency are +// left out rather than shown with the wrong symbol. +func TestDeepSeekLoggingHistorySkipsOtherCurrencies(t *testing.T) { + h, s := newBalanceTestHandler(t) + now := time.Now().UTC() + cny := &api.DeepSeekSnapshot{CapturedAt: now.Add(-2 * time.Hour), IsAvailable: true, Currency: "CNY", TotalBalance: 30} + if _, err := s.InsertDeepSeekSnapshot(cny); err != nil { + t.Fatalf("insert: %v", err) + } + insertDeepSeekUSD(t, h, s, now.Add(-time.Hour), 4.02) + + var resp struct { + Currency string `json:"currency"` + Logs []interface{} `json:"logs"` + } + getJSON(t, h, "/api/logging-history?provider=deepseek&range=1", &resp) + if resp.Currency != "USD" || len(resp.Logs) != 1 { + t.Fatalf("currency=%q logs=%d, want only the USD row", resp.Currency, len(resp.Logs)) + } +} + +func TestBalanceCardsTreatPlaceholderAsNoData(t *testing.T) { + js := readStaticFile(t, "static/app.js") + if !strings.Contains(js, "if (!balance || !balance.status) {") { + t.Fatal("renderBalanceCards must not render the pre-poll zero placeholder as a healthy balance") + } +} diff --git a/internal/web/commandcode_handlers_test.go b/internal/web/commandcode_handlers_test.go index 71ffb2ce..de09a1d3 100644 --- a/internal/web/commandcode_handlers_test.go +++ b/internal/web/commandcode_handlers_test.go @@ -11,6 +11,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/api" "github.com/onllm-dev/onwatch/v2/internal/config" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" "github.com/onllm-dev/onwatch/v2/internal/tracker" ) @@ -556,7 +557,7 @@ func TestIsProviderConfiguredCommandCode(t *testing.T) { h := NewHandler(nil, nil, nil, nil, tc.cfg) // Point detection at an empty home so the developer's real auth // files cannot make the "nothing" case pass by accident. - t.Setenv("HOME", t.TempDir()) + testhome.SetTestHome(t, t.TempDir()) t.Setenv("COMMAND_CODE_API_KEY", "") t.Setenv("COMMANDCODE_API_KEY", "") t.Setenv("COMMANDCODE_AUTH_PATH", "/nonexistent/auth.json") diff --git a/internal/web/deepseek_handlers.go b/internal/web/deepseek_handlers.go index 7bebefd7..c53f1b21 100644 --- a/internal/web/deepseek_handlers.go +++ b/internal/web/deepseek_handlers.go @@ -3,6 +3,7 @@ package web import ( "fmt" "net/http" + "strings" "time" "github.com/onllm-dev/onwatch/v2/internal/store" @@ -39,12 +40,12 @@ func (h *Handler) buildDeepSeekCurrent() map[string]interface{} { if latest != nil { response["capturedAt"] = latest.CapturedAt.Format(time.RFC3339) - + status := "healthy" if latest.TotalBalance == 0 { status = "critical" } - + balance := map[string]interface{}{ "name": "Balance", "description": "DeepSeek API balance", @@ -132,10 +133,7 @@ func (h *Handler) cyclesDeepSeek(w http.ResponseWriter, r *http.Request) { } quotaType := "balance" - currency := r.URL.Query().Get("currency") - if currency == "" { - currency = "CNY" // Default - } + currency := h.deepseekCurrency(r.URL.Query().Get("currency")) response := make([]map[string]interface{}, 0) active, err := h.store.QueryActiveDeepSeekCycle(quotaType, currency) @@ -166,6 +164,7 @@ func (h *Handler) cyclesDeepSeek(w http.ResponseWriter, r *http.Request) { func deepseekCycleToMap(cycle *store.DeepSeekResetCycle) map[string]interface{} { result := map[string]interface{}{ "id": cycle.ID, + "cycleId": cycle.ID, // the dashboard cycle tables read cycleId "quotaType": cycle.QuotaType, "currency": cycle.Currency, "cycleStart": cycle.CycleStart.Format(time.RFC3339), @@ -181,12 +180,24 @@ func deepseekCycleToMap(cycle *store.DeepSeekResetCycle) map[string]interface{} return result } +// deepseekCurrency resolves the currency to report: the requested one, else +// the currency of the latest snapshot. DeepSeek accounts hold either CNY or +// USD, and a view queried in the wrong currency comes back empty (#137). +func (h *Handler) deepseekCurrency(requested string) string { + if requested = strings.ToUpper(strings.TrimSpace(requested)); requested != "" { + return requested + } + if h.store != nil { + if latest, err := h.store.QueryLatestDeepSeek(); err == nil && latest != nil && latest.Currency != "" { + return latest.Currency + } + } + return "CNY" +} + // summaryDeepSeek returns DeepSeek usage summary func (h *Handler) summaryDeepSeek(w http.ResponseWriter, r *http.Request) { - currency := r.URL.Query().Get("currency") - if currency == "" { - currency = "CNY" - } + currency := h.deepseekCurrency(r.URL.Query().Get("currency")) respondJSON(w, http.StatusOK, h.buildDeepSeekSummaryMap(currency)) } @@ -244,10 +255,7 @@ func (h *Handler) buildDeepSeekSummaryMap(currency string) map[string]interface{ // insightsDeepSeek returns DeepSeek insights func (h *Handler) insightsDeepSeek(w http.ResponseWriter, r *http.Request, rangeDur time.Duration) { hidden := h.getHiddenInsightKeys() - currency := r.URL.Query().Get("currency") - if currency == "" { - currency = "CNY" - } + currency := h.deepseekCurrency(r.URL.Query().Get("currency")) respondJSON(w, http.StatusOK, h.buildDeepSeekInsights(currency, hidden)) } @@ -268,7 +276,7 @@ func (h *Handler) buildDeepSeekInsights(currency string, hidden map[string]bool) }) return resp } - + if latest.Currency != currency { // Only reporting for currently tracked currency return resp @@ -306,7 +314,7 @@ func (h *Handler) buildDeepSeekInsights(currency string, hidden map[string]bool) } } } - + if !latest.IsAvailable { resp.Insights = append(resp.Insights, insightItem{ Type: "warning", Severity: "high", @@ -320,44 +328,42 @@ func (h *Handler) buildDeepSeekInsights(currency string, hidden map[string]bool) // cycleOverviewDeepSeek returns DeepSeek cycle overview. func (h *Handler) cycleOverviewDeepSeek(w http.ResponseWriter, r *http.Request) { - if h.store == nil { - respondJSON(w, http.StatusOK, map[string]interface{}{"cycles": []interface{}{}}) - return - } + respondJSON(w, http.StatusOK, h.deepseekCycleOverview(h.deepseekCurrency(r.URL.Query().Get("currency")))) +} +// deepseekCycleOverview lists the active and recent balance cycles. Balance +// cycles carry only a spend delta, so there are no per-quota columns. +func (h *Handler) deepseekCycleOverview(currency string) map[string]interface{} { quotaType := "balance" - currency := r.URL.Query().Get("currency") - if currency == "" { - currency = "CNY" - } - var cycles []map[string]interface{} - - if active, err := h.store.QueryActiveDeepSeekCycle(quotaType, currency); err == nil && active != nil { - cycles = append(cycles, deepseekCycleToMap(active)) - } - if history, err := h.store.QueryDeepSeekCycleHistory(quotaType, currency, 50); err == nil { - for _, c := range history { - cycles = append(cycles, deepseekCycleToMap(c)) + cycles := []map[string]interface{}{} + if h.store != nil { + if active, err := h.store.QueryActiveDeepSeekCycle(quotaType, currency); err == nil && active != nil { + cycles = append(cycles, deepseekCycleToMap(active)) + } + if history, err := h.store.QueryDeepSeekCycleHistory(quotaType, currency, 50); err == nil { + for _, c := range history { + cycles = append(cycles, deepseekCycleToMap(c)) + } } } - - respondJSON(w, http.StatusOK, map[string]interface{}{ + return map[string]interface{}{ "groupBy": quotaType, "provider": "deepseek", - "quotaNames": []string{"balance"}, + "currency": currency, + "quotaNames": []string{}, "cycles": cycles, - }) + } } // loggingHistoryDeepSeek returns DeepSeek polling history. func (h *Handler) loggingHistoryDeepSeek(w http.ResponseWriter, r *http.Request) { + quotaNames := []string{"total_balance", "granted_balance", "topped_up_balance"} if h.store == nil { - respondJSON(w, http.StatusOK, map[string]interface{}{"provider": "deepseek", "quotaNames": []string{}, "logs": []interface{}{}}) + respondJSON(w, http.StatusOK, map[string]interface{}{"provider": "deepseek", "quotaNames": quotaNames, "logs": []interface{}{}}) return } start, end, limit := h.loggingHistoryRangeAndLimit(r) - snapshots, err := h.store.QueryDeepSeekRange(start, end, limit) if err != nil { h.logger.Error("failed to query DeepSeek logging history", "error", err) @@ -365,55 +371,35 @@ func (h *Handler) loggingHistoryDeepSeek(w http.ResponseWriter, r *http.Request) return } - quotaNames := []string{"balance"} - type quotaVal struct { - Name string - Value float64 - HasValue bool - } - - capturedAt := make([]string, 0, len(snapshots)) + // Rows are labelled with one currency, so skip snapshots taken in another + // (an account that switched between CNY and USD). + currency := h.deepseekCurrency("") + capturedAt := make([]time.Time, 0, len(snapshots)) ids := make([]int64, 0, len(snapshots)) - series := make([]map[string]quotaVal, 0, len(snapshots)) - + series := make([]map[string]loggingHistoryCrossQuota, 0, len(snapshots)) for _, snap := range snapshots { - capturedAt = append(capturedAt, snap.CapturedAt.Format(time.RFC3339)) - ids = append(ids, snap.ID) - - row := map[string]quotaVal{ - "balance": { - Name: "balance", - Value: snap.TotalBalance, - HasValue: true, - }, - } - series = append(series, row) - } - - logs := make([]map[string]interface{}, 0, len(snapshots)) - for i := range snapshots { - entry := map[string]interface{}{ - "capturedAt": capturedAt[i], - "id": ids[i], - "quotas": map[string]interface{}{}, - } - quotas := map[string]interface{}{} - for _, qn := range quotaNames { - if qv, ok := series[i][qn]; ok { - quotas[qn] = map[string]interface{}{ - "name": qv.Name, - "value": qv.Value, - "hasValue": qv.HasValue, - } - } + if snap.Currency != "" && snap.Currency != currency { + continue } - entry["quotas"] = quotas - logs = append(logs, entry) + capturedAt = append(capturedAt, snap.CapturedAt) + ids = append(ids, snap.ID) + series = append(series, balanceCrossQuotas(quotaNames, snap.TotalBalance, snap.GrantedBalance, snap.ToppedUpBalance)) } respondJSON(w, http.StatusOK, map[string]interface{}{ "provider": "deepseek", + "currency": currency, "quotaNames": quotaNames, - "logs": logs, + "logs": loggingHistoryRowsFromSnapshots(capturedAt, ids, quotaNames, series), }) } + +// balanceCrossQuotas maps balance amounts to logging-history cells. Balances +// have no limit, so only the value is set. +func balanceCrossQuotas(names []string, values ...float64) map[string]loggingHistoryCrossQuota { + row := make(map[string]loggingHistoryCrossQuota, len(names)) + for i, name := range names { + row[name] = loggingHistoryCrossQuota{Name: name, Value: values[i], HasValue: true} + } + return row +} diff --git a/internal/web/handlers.go b/internal/web/handlers.go index 42036e69..533cd540 100644 --- a/internal/web/handlers.go +++ b/internal/web/handlers.go @@ -4213,46 +4213,11 @@ func (h *Handler) cyclesBoth(w http.ResponseWriter, r *http.Request) { } if h.config.HasProvider("moonshot") { - quotaType := "balance" - var msCycles []map[string]interface{} - if active, err := h.store.QueryActiveMoonshotCycle(quotaType); err == nil && active != nil { - msCycles = append(msCycles, moonshotCycleToMap(active)) - } - if history, err := h.store.QueryMoonshotCycleHistory(quotaType, 50); err == nil { - for _, c := range history { - msCycles = append(msCycles, moonshotCycleToMap(c)) - } - } - response["moonshot"] = map[string]interface{}{ - "groupBy": quotaType, - "provider": "moonshot", - "quotaNames": []string{"balance"}, - "cycles": msCycles, - } + response["moonshot"] = h.moonshotCycleOverview() } if h.config.HasProvider("deepseek") { - quotaType := "balance" - var dsCycles []map[string]interface{} - - // Use CNY as default if not specified elsewhere. DeepSeek could use USD, - // but tracking one primary currency for UI is sufficient for summary. - currency := "CNY" - - if active, err := h.store.QueryActiveDeepSeekCycle(quotaType, currency); err == nil && active != nil { - dsCycles = append(dsCycles, deepseekCycleToMap(active)) - } - if history, err := h.store.QueryDeepSeekCycleHistory(quotaType, currency, 50); err == nil { - for _, c := range history { - dsCycles = append(dsCycles, deepseekCycleToMap(c)) - } - } - response["deepseek"] = map[string]interface{}{ - "groupBy": quotaType, - "provider": "deepseek", - "quotaNames": []string{"balance"}, - "cycles": dsCycles, - } + response["deepseek"] = h.deepseekCycleOverview(h.deepseekCurrency("")) } if h.config.HasProvider("gemini") { @@ -4490,8 +4455,7 @@ func (h *Handler) summaryBoth(w http.ResponseWriter, r *http.Request) { response["moonshot"] = h.buildMoonshotSummaryMap() } if h.config.HasProvider("deepseek") { - // DeepSeek could use either currency. Use CNY by default for summary view unless we know better - response["deepseek"] = h.buildDeepSeekSummaryMap("CNY") + response["deepseek"] = h.buildDeepSeekSummaryMap(h.deepseekCurrency("")) } if h.config.HasProvider("anthropic") { response["anthropic"] = h.buildAnthropicSummaryMap() @@ -5366,8 +5330,7 @@ func (h *Handler) insightsBoth(w http.ResponseWriter, r *http.Request, rangeDur response["moonshot"] = h.buildMoonshotInsights(hidden) } if h.config.HasProvider("deepseek") && providerTelemetryEnabled(visibility, "deepseek") { - // Use CNY for deepseek overall insights if not explicitly asked - response["deepseek"] = h.buildDeepSeekInsights("CNY", hidden) + response["deepseek"] = h.buildDeepSeekInsights(h.deepseekCurrency(""), hidden) } if h.config.HasProvider("gemini") && providerTelemetryEnabled(visibility, "gemini") { response["gemini"] = insightsResponse{Stats: []insightStat{}, Insights: []insightItem{}} @@ -8295,42 +8258,11 @@ func (h *Handler) cycleOverviewBoth(w http.ResponseWriter, r *http.Request) { } if h.config.HasProvider("moonshot") { - quotaType := "balance" - var msCycles []map[string]interface{} - if active, err := h.store.QueryActiveMoonshotCycle(quotaType); err == nil && active != nil { - msCycles = append(msCycles, moonshotCycleToMap(active)) - } - if history, err := h.store.QueryMoonshotCycleHistory(quotaType, 50); err == nil { - for _, c := range history { - msCycles = append(msCycles, moonshotCycleToMap(c)) - } - } - response["moonshot"] = map[string]interface{}{ - "groupBy": quotaType, - "provider": "moonshot", - "quotaNames": []string{"balance"}, - "cycles": msCycles, - } + response["moonshot"] = h.moonshotCycleOverview() } if h.config.HasProvider("deepseek") { - quotaType := "balance" - currency := "CNY" // Could be made dynamic - var dsCycles []map[string]interface{} - if active, err := h.store.QueryActiveDeepSeekCycle(quotaType, currency); err == nil && active != nil { - dsCycles = append(dsCycles, deepseekCycleToMap(active)) - } - if history, err := h.store.QueryDeepSeekCycleHistory(quotaType, currency, 50); err == nil { - for _, c := range history { - dsCycles = append(dsCycles, deepseekCycleToMap(c)) - } - } - response["deepseek"] = map[string]interface{}{ - "groupBy": quotaType, - "provider": "deepseek", - "quotaNames": []string{"balance"}, - "cycles": dsCycles, - } + response["deepseek"] = h.deepseekCycleOverview(h.deepseekCurrency("")) } if h.config.HasProvider("gemini") { diff --git a/internal/web/handlers_coverage_test.go b/internal/web/handlers_coverage_test.go index bad16caf..92d060cb 100644 --- a/internal/web/handlers_coverage_test.go +++ b/internal/web/handlers_coverage_test.go @@ -14,6 +14,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/api" "github.com/onllm-dev/onwatch/v2/internal/config" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" "github.com/onllm-dev/onwatch/v2/internal/tracker" ) @@ -257,7 +258,7 @@ func TestHandlerTryAutoDetectAdditionalCoverage(t *testing.T) { t.Run("anthropic and codex miss", func(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", filepath.Join(home, "codex-home")) h := NewHandler(nil, nil, nil, nil, &config.Config{}) @@ -271,7 +272,7 @@ func TestHandlerTryAutoDetectAdditionalCoverage(t *testing.T) { t.Run("anthropic success from credentials file", func(t *testing.T) { home := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) credsDir := filepath.Join(home, ".claude") if err := os.MkdirAll(credsDir, 0o755); err != nil { t.Fatalf("MkdirAll: %v", err) @@ -343,7 +344,7 @@ func TestHandlerReloadProvidersCoverage(t *testing.T) { home := t.TempDir() codexHome := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", codexHome) t.Setenv("SYNTHETIC_API_KEY", "syn_reload_key") t.Setenv("ZAI_API_KEY", "") diff --git a/internal/web/handlers_extra_test.go b/internal/web/handlers_extra_test.go index b89bd72d..8ab9cf62 100644 --- a/internal/web/handlers_extra_test.go +++ b/internal/web/handlers_extra_test.go @@ -12123,8 +12123,8 @@ func TestSanitizeProviderSettings_NonEnumFieldsUntouched(t *testing.T) { }, "anthropic": map[string]interface{}{ "api_poll_cycle_interval": float64(20), - "staleness_minutes": float64(10), - "source": "api", + "staleness_minutes": float64(10), + "source": "api", }, } @@ -12200,8 +12200,8 @@ func TestCountWorkTime(t *testing.T) { {"empty range", mon, mon, "5-day", 0.0}, {"partial first day", mon, monPartial, "5-day", 0.25}, {"2 full + partial day", mon, wedNoon, "5-day", 2.5}, - {"5-day to Sat noon", mon, satNoon, "5-day", 5.0}, // Sat doesn't count - {"6-day to Sat noon", mon, satNoon, "6-day", 5.5}, // Sat counts, partial + {"5-day to Sat noon", mon, satNoon, "5-day", 5.0}, // Sat doesn't count + {"6-day to Sat noon", mon, satNoon, "6-day", 5.5}, // Sat counts, partial {"calendar to Sat noon", mon, satNoon, "calendar", 5.5}, } diff --git a/internal/web/kimi_handlers.go b/internal/web/kimi_handlers.go index b320943f..25898bd4 100644 --- a/internal/web/kimi_handlers.go +++ b/internal/web/kimi_handlers.go @@ -353,4 +353,3 @@ func (h *Handler) buildKimiInsights(hidden map[string]bool) insightsResponse { } return resp } - diff --git a/internal/web/minimax_handlers_test.go b/internal/web/minimax_handlers_test.go index b004cd4e..61ac7186 100644 --- a/internal/web/minimax_handlers_test.go +++ b/internal/web/minimax_handlers_test.go @@ -87,12 +87,12 @@ func TestBuildMiniMaxCurrent_SharedQuota(t *testing.T) { var resp struct { SharedQuota bool `json:"sharedQuota"` Quotas []struct { - Name string `json:"name"` - DisplayName string `json:"displayName"` - Used int `json:"used"` - Remaining int `json:"remaining"` - Total int `json:"total"` - UsagePercent float64 `json:"usagePercent"` + Name string `json:"name"` + DisplayName string `json:"displayName"` + Used int `json:"used"` + Remaining int `json:"remaining"` + Total int `json:"total"` + UsagePercent float64 `json:"usagePercent"` } `json:"quotas"` } if err := json.Unmarshal(body, &resp); err != nil { diff --git a/internal/web/moonshot_handlers.go b/internal/web/moonshot_handlers.go index 4ec5dff9..2f398efb 100644 --- a/internal/web/moonshot_handlers.go +++ b/internal/web/moonshot_handlers.go @@ -37,12 +37,12 @@ func (h *Handler) buildMoonshotCurrent() map[string]interface{} { if latest != nil { response["capturedAt"] = latest.CapturedAt.Format(time.RFC3339) - + status := "healthy" if latest.AvailableBalance == 0 { status = "critical" } - + balance := map[string]interface{}{ "name": "Balance", "description": "Moonshot Kimi API balance", @@ -156,6 +156,7 @@ func (h *Handler) cyclesMoonshot(w http.ResponseWriter, r *http.Request) { func moonshotCycleToMap(cycle *store.MoonshotResetCycle) map[string]interface{} { result := map[string]interface{}{ "id": cycle.ID, + "cycleId": cycle.ID, // the dashboard cycle tables read cycleId "quotaType": cycle.QuotaType, "cycleStart": cycle.CycleStart.Format(time.RFC3339), "cycleEnd": nil, @@ -279,40 +280,41 @@ func (h *Handler) buildMoonshotInsights(hidden map[string]bool) insightsResponse // cycleOverviewMoonshot returns Moonshot cycle overview. func (h *Handler) cycleOverviewMoonshot(w http.ResponseWriter, r *http.Request) { - if h.store == nil { - respondJSON(w, http.StatusOK, map[string]interface{}{"cycles": []interface{}{}}) - return - } + respondJSON(w, http.StatusOK, h.moonshotCycleOverview()) +} +// moonshotCycleOverview lists the active and recent balance cycles. Balance +// cycles carry only a spend delta, so there are no per-quota columns. +func (h *Handler) moonshotCycleOverview() map[string]interface{} { quotaType := "balance" - var cycles []map[string]interface{} - - if active, err := h.store.QueryActiveMoonshotCycle(quotaType); err == nil && active != nil { - cycles = append(cycles, moonshotCycleToMap(active)) - } - if history, err := h.store.QueryMoonshotCycleHistory(quotaType, 50); err == nil { - for _, c := range history { - cycles = append(cycles, moonshotCycleToMap(c)) + cycles := []map[string]interface{}{} + if h.store != nil { + if active, err := h.store.QueryActiveMoonshotCycle(quotaType); err == nil && active != nil { + cycles = append(cycles, moonshotCycleToMap(active)) + } + if history, err := h.store.QueryMoonshotCycleHistory(quotaType, 50); err == nil { + for _, c := range history { + cycles = append(cycles, moonshotCycleToMap(c)) + } } } - - respondJSON(w, http.StatusOK, map[string]interface{}{ + return map[string]interface{}{ "groupBy": quotaType, "provider": "moonshot", - "quotaNames": []string{"balance"}, + "quotaNames": []string{}, "cycles": cycles, - }) + } } // loggingHistoryMoonshot returns Moonshot polling history. func (h *Handler) loggingHistoryMoonshot(w http.ResponseWriter, r *http.Request) { + quotaNames := []string{"available_balance", "voucher_balance", "cash_balance"} if h.store == nil { - respondJSON(w, http.StatusOK, map[string]interface{}{"provider": "moonshot", "quotaNames": []string{}, "logs": []interface{}{}}) + respondJSON(w, http.StatusOK, map[string]interface{}{"provider": "moonshot", "quotaNames": quotaNames, "logs": []interface{}{}}) return } start, end, limit := h.loggingHistoryRangeAndLimit(r) - snapshots, err := h.store.QueryMoonshotRange(start, end, limit) if err != nil { h.logger.Error("failed to query Moonshot logging history", "error", err) @@ -320,55 +322,18 @@ func (h *Handler) loggingHistoryMoonshot(w http.ResponseWriter, r *http.Request) return } - quotaNames := []string{"balance"} - type quotaVal struct { - Name string - Value float64 - HasValue bool - } - - capturedAt := make([]string, 0, len(snapshots)) + capturedAt := make([]time.Time, 0, len(snapshots)) ids := make([]int64, 0, len(snapshots)) - series := make([]map[string]quotaVal, 0, len(snapshots)) - + series := make([]map[string]loggingHistoryCrossQuota, 0, len(snapshots)) for _, snap := range snapshots { - capturedAt = append(capturedAt, snap.CapturedAt.Format(time.RFC3339)) + capturedAt = append(capturedAt, snap.CapturedAt) ids = append(ids, snap.ID) - - row := map[string]quotaVal{ - "balance": { - Name: "balance", - Value: snap.AvailableBalance, - HasValue: true, - }, - } - series = append(series, row) - } - - logs := make([]map[string]interface{}, 0, len(snapshots)) - for i := range snapshots { - entry := map[string]interface{}{ - "capturedAt": capturedAt[i], - "id": ids[i], - "quotas": map[string]interface{}{}, - } - quotas := map[string]interface{}{} - for _, qn := range quotaNames { - if qv, ok := series[i][qn]; ok { - quotas[qn] = map[string]interface{}{ - "name": qv.Name, - "value": qv.Value, - "hasValue": qv.HasValue, - } - } - } - entry["quotas"] = quotas - logs = append(logs, entry) + series = append(series, balanceCrossQuotas(quotaNames, snap.AvailableBalance, snap.VoucherBalance, snap.CashBalance)) } respondJSON(w, http.StatusOK, map[string]interface{}{ "provider": "moonshot", "quotaNames": quotaNames, - "logs": logs, + "logs": loggingHistoryRowsFromSnapshots(capturedAt, ids, quotaNames, series), }) } diff --git a/internal/web/provider_management_test.go b/internal/web/provider_management_test.go index 5005a589..5f4af05d 100644 --- a/internal/web/provider_management_test.go +++ b/internal/web/provider_management_test.go @@ -14,6 +14,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/api" "github.com/onllm-dev/onwatch/v2/internal/config" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) type mockProviderAgentController struct { @@ -137,7 +138,7 @@ func TestHandler_ProviderVisibilityHelpers(t *testing.T) { func TestHandler_IsProviderConfiguredAndTryAutoDetect(t *testing.T) { home := t.TempDir() codexHome := t.TempDir() - t.Setenv("HOME", home) + testhome.SetTestHome(t, home) t.Setenv("CODEX_HOME", codexHome) if err := os.WriteFile(filepath.Join(codexHome, "auth.json"), []byte(`{"tokens":{"access_token":"codex-auto"}}`), 0o600); err != nil { diff --git a/internal/web/server_test.go b/internal/web/server_test.go index f03174cb..e23c9487 100644 --- a/internal/web/server_test.go +++ b/internal/web/server_test.go @@ -2,6 +2,7 @@ package web import ( "context" + "fmt" "io" "log/slog" "net" @@ -15,6 +16,7 @@ import ( "github.com/onllm-dev/onwatch/v2/internal/api" "github.com/onllm-dev/onwatch/v2/internal/config" "github.com/onllm-dev/onwatch/v2/internal/store" + "github.com/onllm-dev/onwatch/v2/internal/testutil/testhome" ) // freePort returns an available TCP port for testing @@ -274,9 +276,28 @@ func TestServer_EmbeddedAssets(t *testing.T) { server.Shutdown(ctx) } +// TestMain isolates the package from the developer's real credentials and +// state. api.SetTestMode(true) stops provider auto-detection from reading the +// macOS Keychain or Linux keyring; it runs before the home is redirected so +// the real home is recorded for the credential-file guard. The home directory +// (HOME, USERPROFILE and LOCALAPPDATA) is then pointed at a throwaway sandbox +// for the whole run and provider location overrides are cleared. Tests that +// need their own home call testhome.SetTestHome. func TestMain(m *testing.M) { - // Ensure templates directory exists for tests - os.Exit(m.Run()) + os.Exit(runTests(m)) +} + +func runTests(m *testing.M) int { + api.SetTestMode(true) + + _, cleanup, err := testhome.SandboxHome() + if err != nil { + fmt.Fprintf(os.Stderr, "web tests: %v\n", err) + return 1 + } + defer cleanup() + + return m.Run() } func TestServer_RequiresCSRFHeader_OnPost(t *testing.T) { diff --git a/internal/web/settings_init_static_test.go b/internal/web/settings_init_static_test.go new file mode 100644 index 00000000..15e349ba --- /dev/null +++ b/internal/web/settings_init_static_test.go @@ -0,0 +1,31 @@ +package web + +import ( + "strings" + "testing" +) + +// The password form must be wired before initSettingsPage awaits the menubar +// and settings loads; otherwise a click during a slow load does nothing. The +// data-ready marker tells the e2e suite every control is wired. +func TestSettingsInitWiresPasswordBeforeAsyncLoads(t *testing.T) { + // A Windows checkout may convert app.js to CRLF line endings. + js := strings.ReplaceAll(readStaticFile(t, "static/app.js"), "\r\n", "\n") + start := strings.Index(js, "async function initSettingsPage() {") + if start < 0 { + t.Fatal("initSettingsPage not found") + } + end := strings.Index(js[start:], "\n}\n") + if end < 0 { + t.Fatal("end of initSettingsPage not found") + } + body := js[start : start+end] + password := strings.Index(body, "setupSettingsPassword();") + firstAwait := strings.Index(body, "await ") + if password < 0 || firstAwait < 0 || password > firstAwait { + t.Fatalf("setupSettingsPassword must run before the first await in initSettingsPage:\n%s", body) + } + if !strings.Contains(body, "setAttribute('data-ready', 'true')") { + t.Fatal("initSettingsPage must mark .settings-page data-ready when wiring completes") + } +} diff --git a/internal/web/static/app.js b/internal/web/static/app.js index cd5c7130..0b7bbb49 100644 --- a/internal/web/static/app.js +++ b/internal/web/static/app.js @@ -81,6 +81,10 @@ function getCurrentProvider() { if (museGrid) return 'muse'; const commandCodeGrid = document.getElementById('quota-grid-commandcode'); if (commandCodeGrid) return 'commandcode'; + const moonshotGrid = document.getElementById('quota-grid-moonshot'); + if (moonshotGrid) return 'moonshot'; + const deepseekGrid = document.getElementById('quota-grid-deepseek'); + if (deepseekGrid) return 'deepseek'; const grid = document.getElementById('quota-grid'); return (grid && grid.dataset.provider) || 'synthetic'; } @@ -1220,6 +1224,12 @@ const renewalCategories = { openrouter: [ { label: 'Credits', groupBy: 'credits' } ], + moonshot: [ + { label: 'Balance', groupBy: 'balance' } + ], + deepseek: [ + { label: 'Balance', groupBy: 'balance' } + ], grok: [ { label: 'Credits', groupBy: 'credits' } ], @@ -1302,6 +1312,16 @@ const providerQuotaDisplayOverrides = { five_hour: '5-Hour Credits', weekly: 'Weekly Credits', monthly: 'Monthly Credits' + }, + moonshot: { + available_balance: 'Available', + voucher_balance: 'Voucher', + cash_balance: 'Cash' + }, + deepseek: { + total_balance: 'Total Balance', + granted_balance: 'Granted', + topped_up_balance: 'Topped Up' } }; @@ -2514,6 +2534,104 @@ function updateOpenRouterCard(credits) { if (resetEl) resetEl.textContent = hasLimit ? 'Remaining: ' + remainStr : ''; } +// ── Balance providers (Moonshot, DeepSeek) ── +// These report a remaining balance with no limit, so their cards show the +// amount instead of a utilization bar. +const balanceProviderFields = { + moonshot: [ + { key: 'available', historyKey: 'available_balance', label: 'Available', primary: true }, + { key: 'voucher', historyKey: 'voucher_balance', label: 'Voucher' }, + { key: 'cash', historyKey: 'cash_balance', label: 'Cash' } + ], + deepseek: [ + { key: 'total', historyKey: 'total_balance', label: 'Total Balance', primary: true }, + { key: 'granted', historyKey: 'granted_balance', label: 'Granted' }, + { key: 'toppedUp', historyKey: 'topped_up_balance', label: 'Topped Up' } + ] +}; +const balanceChartColors = [ + { border: '#0D9488', bg: 'rgba(13, 148, 136, 0.06)' }, + { border: '#F59E0B', bg: 'rgba(245, 158, 11, 0.06)' }, + { border: '#3B82F6', bg: 'rgba(59, 130, 246, 0.06)' } +]; + +function isBalanceProvider(provider) { + return Object.prototype.hasOwnProperty.call(balanceProviderFields, provider); +} + +// Moonshot does not report a currency, so its amounts are shown bare. +function formatBalanceAmount(value, currency) { + const amount = Number(value || 0).toLocaleString('en-US', { minimumFractionDigits: 2, maximumFractionDigits: 2 }); + if (currency === 'USD') return '$' + amount; + if (currency === 'CNY') return '\u00A5' + amount; + return currency ? `${amount} ${currency}` : amount; +} + +function renderBalanceCards(provider, balance, containerId) { + const container = document.getElementById(containerId); + if (!container) return; + // Before the first poll the API returns a zero placeholder with no status; + // show that as "no data" rather than a healthy zero balance. + if (!balance || !balance.status) { + container.innerHTML = '

No balance data yet

'; + return; + } + State.balanceCurrency = balance.currency || ''; + const fields = balanceProviderFields[provider] || []; + if (container.querySelectorAll('.balance-card').length !== fields.length) { + const icon = ''; + container.innerHTML = fields.map((field, i) => `
+
+

+ ${icon} + ${field.label} +

+
+
+ + +
+ ${field.primary ? `
+ +
` : ''} +
`).join(''); + } + updateBalanceCards(provider, balance, container); +} + +function updateBalanceCards(provider, balance, container) { + const currency = balance.currency || ''; + const rate = Number(balance.rate || 0); + const cards = container.querySelectorAll('.balance-card'); + (balanceProviderFields[provider] || []).forEach((field, i) => { + const card = cards[i]; + if (!card) return; + card.querySelector('[data-balance-amount]').textContent = formatBalanceAmount(balance[field.key], currency); + card.querySelector('[data-balance-detail]').textContent = field.primary && rate > 0 + ? `Spending ${formatBalanceAmount(rate, currency)}/h` + : ''; + const statusEl = card.querySelector('[data-balance-status]'); + if (statusEl) { + const status = balance.status || 'healthy'; + const statusCfg = statusConfig[status] || statusConfig.healthy; + statusEl.setAttribute('data-status', status); + statusEl.innerHTML = `${statusCfg.label}`; + } + }); +} + +function buildBalanceDatasets(provider, rows, range) { + const fields = balanceProviderFields[provider] || []; + const datasets = buildFixedDatasetsForRows(rows, range, fields.map((field, i) => ({ + label: field.label, + key: field.historyKey, + color: balanceChartColors[i % balanceChartColors.length].border, + bg: balanceChartColors[i % balanceChartColors.length].bg + }))); + datasets.forEach((ds, i) => { ds.hidden = State.hiddenQuotas.has(fields[i].historyKey); }); + return datasets; +} + function getQuotaStatus(percent) { if (percent >= 90) return 'critical'; if (percent >= 75) return 'warning'; @@ -4746,6 +4864,9 @@ async function fetchCurrent() { } } + } else if (isBalanceProvider(provider)) { + renderBalanceCards(provider, data.balance, `quota-grid-${provider}`); + } else if (provider === 'zai') { updateCard('tokensLimit', data.tokensLimit); updateCard('timeLimit', data.timeLimit); @@ -5754,6 +5875,8 @@ function initChart() { defaultDatasets = []; // OpenCode datasets are dynamic } else if (provider === 'ollama') { defaultDatasets = []; // Ollama datasets are dynamic + } else if (isBalanceProvider(provider)) { + defaultDatasets = []; // Balance datasets are built when history data arrives } else if (provider === 'zai') { defaultDatasets = [ { label: zaiQuotaLabel('tokensLimit', 'Tokens Limit'), data: [], borderColor: getComputedStyle(document.documentElement).getPropertyValue('--chart-subscription').trim() || '#0D9488', backgroundColor: 'rgba(13, 148, 136, 0.06)', fill: true, tension: 0.4, borderWidth: 2, pointRadius: 0, pointHoverRadius: 4, hidden: State.hiddenQuotas.has('tokensLimit') }, @@ -5787,11 +5910,14 @@ function initChart() { ? [] : provider === 'ollama' ? [] + : isBalanceProvider(provider) + ? [] : provider === 'api-integrations' ? [] : ['subscription', 'search', 'toolCalls']; const isAPIIntegrations = provider === 'api-integrations'; + const isBalance = isBalanceProvider(provider); State.chart = new Chart(ctx, { type: 'line', data: { @@ -5813,7 +5939,7 @@ function initChart() { meta.hidden = meta.hidden === null ? !ci.data.datasets[index].hidden : null; ci.update('none'); // Recalculate Y-axis based on visible datasets - State.chartYMax = computeYMax(ci.data.datasets, ci); + State.chartYMax = computeYMax(ci.data.datasets, ci, isBalance ? { cap: false } : {}); ci.options.scales.y.max = State.chartYMax; ci.update(); } @@ -5842,6 +5968,9 @@ function initChart() { } return `${ctx.dataset.label}: ${formatNumber(Number(ctx.parsed.y || 0))}`; } + if (isBalance) { + return `${ctx.dataset.label}: ${formatBalanceAmount(ctx.parsed.y, State.balanceCurrency)}`; + } return `${ctx.dataset.label}: ${ctx.parsed.y.toFixed(1)}%`; } } @@ -5864,7 +5993,7 @@ function initChart() { : ((State.apiIntegrationsSelectedMetric || 'tokenPerCall') === 'tokenPerCall' ? formatNumber(Number(v || 0).toFixed(1)) : formatNumber(Number(v || 0)))) - : v + '%' + : (isBalance ? formatBalanceAmount(v, State.balanceCurrency) : v + '%') }, title: { display: isAPIIntegrations, @@ -6373,6 +6502,16 @@ async function fetchHistory(range) { return; } + if (isBalanceProvider(provider)) { + const lastRow = historyRows[historyRows.length - 1]; + if (lastRow) State.balanceCurrency = lastRow.currency || ''; + State.chart.data.datasets = buildBalanceDatasets(provider, historyRows, range); + updateTimeScale(State.chart, range); + State.chartYMax = computeYMax(State.chart.data.datasets, State.chart, { cap: false }); + State.chart.options.scales.y.max = State.chartYMax; + State.chart.update(); + return; + } if (provider === 'codex') { // Codex history: array of { capturedAt, five_hour, seven_day, ... } @@ -7376,25 +7515,8 @@ function buildProviderCardDatasets(provider, rows, range) { const orFallback = [{ border: '#8B5CF6', bg: 'rgba(139, 92, 246, 0.06)' }]; return buildDynamicDatasetsForRows(rows, range, orDisplayNames, orColors, orFallback, 'openrouter'); } - if (provider === 'moonshot') { - const msDisplayNames = { available_balance: 'Available', voucher_balance: 'Voucher', cash_balance: 'Cash' }; - const msColors = { - available_balance: { border: '#0D9488', bg: 'rgba(13, 148, 136, 0.06)' }, - voucher_balance: { border: '#F59E0B', bg: 'rgba(245, 158, 11, 0.06)' }, - cash_balance: { border: '#3B82F6', bg: 'rgba(59, 130, 246, 0.06)' } - }; - const msFallback = [{ border: '#8B5CF6', bg: 'rgba(139, 92, 246, 0.06)' }]; - return buildDynamicDatasetsForRows(rows, range, msDisplayNames, msColors, msFallback, 'moonshot'); - } - if (provider === 'deepseek') { - const dsDisplayNames = { total_balance: 'Total Balance', granted_balance: 'Granted', topped_up_balance: 'Topped Up' }; - const dsColors = { - total_balance: { border: '#0D9488', bg: 'rgba(13, 148, 136, 0.06)' }, - granted_balance: { border: '#F59E0B', bg: 'rgba(245, 158, 11, 0.06)' }, - topped_up_balance: { border: '#3B82F6', bg: 'rgba(59, 130, 246, 0.06)' } - }; - const dsFallback = [{ border: '#8B5CF6', bg: 'rgba(139, 92, 246, 0.06)' }]; - return buildDynamicDatasetsForRows(rows, range, dsDisplayNames, dsColors, dsFallback, 'deepseek'); + if (isBalanceProvider(provider)) { + return buildBalanceDatasets(provider, rows, range); } if (provider === 'grok') { const grokDisplay = { credits: 'Credits' }; @@ -7883,7 +8005,7 @@ async function fetchCycles() { const requestSeq = (State.cyclesRequestSeq || 0) + 1; State.cyclesRequestSeq = requestSeq; const provider = requestProvider; - const loggingHistoryProviders = new Set(['synthetic', 'zai', 'anthropic', 'copilot', 'codex', 'antigravity', 'minimax', 'gemini', 'cursor', 'grok', 'kimi', 'mistral', 'opencode', 'ollama', 'muse', 'commandcode']); + const loggingHistoryProviders = new Set(['synthetic', 'zai', 'anthropic', 'copilot', 'codex', 'antigravity', 'minimax', 'gemini', 'cursor', 'grok', 'kimi', 'mistral', 'opencode', 'ollama', 'muse', 'commandcode', 'moonshot', 'deepseek']); // All-accounts overview: fetch each account's logging history and merge, // tagging every row with its account name for the combined table. @@ -7953,6 +8075,7 @@ async function fetchCycles() { crossQuotas: log.crossQuotas || [], })); State.cyclesQuotaNames = data.quotaNames || []; + if (isBalanceProvider(requestProvider)) State.balanceCurrency = data.currency || ''; State.cyclesPage = 1; State.isLoggingHistory = true; renderCyclesTable(); @@ -8077,8 +8200,8 @@ function renderCyclesTable() { const provider = getCurrentProvider(); const quotaNames = State.cyclesQuotaNames; - const usePercent = provider === 'anthropic' || provider === 'copilot' || provider === 'codex' || provider === 'antigravity' || provider === 'minimax' || provider === 'gemini' || provider === 'openrouter' || provider === 'cursor' || provider === 'grok' || provider === 'kimi' || provider === 'moonshot' || provider === 'deepseek' || provider === 'mistral' || provider === 'opencode' || provider === 'ollama'; - const deltaUsesPercent = usePercent && provider !== 'minimax' && provider !== 'moonshot' && provider !== 'deepseek'; + const usePercent = provider === 'anthropic' || provider === 'copilot' || provider === 'codex' || provider === 'antigravity' || provider === 'minimax' || provider === 'gemini' || provider === 'openrouter' || provider === 'cursor' || provider === 'grok' || provider === 'kimi' || provider === 'mistral' || provider === 'opencode' || provider === 'ollama'; + const deltaUsesPercent = usePercent && provider !== 'minimax'; const isLoggingHistory = State.isLoggingHistory === true; const showAccount = isAccountsOverviewMode(provider); const accountTh = showAccount ? 'Account ' : ''; @@ -8279,6 +8402,9 @@ function renderCyclesTable() { cellVal = limit > 0 ? `${formatNumber(used)} / ${formatNumber(limit)} (${percentText})${deltaText}` : `${formatNumber(used)} (${percentText})${deltaText}`; + } else if (isBalanceProvider(provider)) { + const cq = getCrossQuotaValue(row, qn); + cellVal = cq ? escapeHTML(formatBalanceAmount(cq.value, State.balanceCurrency)) : '--'; } else if (usePercent) { cellVal = fmtPctWithDelta(pct, delta); } else { @@ -9425,6 +9551,7 @@ async function fetchCycleOverview() { State.allOverviewData = data.cycles || []; State.overviewQuotaNames = data.quotaNames || []; + if (isBalanceProvider(requestProvider)) State.balanceCurrency = data.currency || ''; renderOverviewTable(); } catch (e) { // cycle overview fetch error - non-critical @@ -9559,7 +9686,7 @@ function renderOverviewTable() { if (showDurationDelta) { html += ` ${duration} - ${fmtOverviewWithRate(row.totalDelta, durationHrs, suffix)}`; + ${isBalanceProvider(overviewProv) ? escapeHTML(formatBalanceAmount(row.totalDelta, State.balanceCurrency)) : fmtOverviewWithRate(row.totalDelta, durationHrs, suffix)}`; } quotaNames.forEach(qn => { @@ -9909,6 +10036,9 @@ function isSettingsPage() { async function initSettingsPage() { setupSettingsTabs(); + // The password form does not depend on loaded settings. Wire it before the + // awaits below, or a click during a slow load silently does nothing. + setupSettingsPassword(); await setupMenubarSettings(); populateTimezoneSelect(); await loadSettings(); @@ -9918,9 +10048,10 @@ async function initSettingsPage() { setupSMTPTest(); setupWebhookTest(); setupPushNotifications(); - setupSettingsPassword(); setupThresholdSliders(); setupOverrides(); + // Signals that every settings control is wired (used by the e2e suite). + document.querySelector('.settings-page')?.setAttribute('data-ready', 'true'); } function activateSettingsTab(tabName) { @@ -11309,11 +11440,11 @@ const providerSettingsConfig = { }, opencode: { title: 'OpenCode Go', - desc: 'Configure OpenCode Go quota tracking. Recommended: a console service-account key, which reads your plan\'s own 5-hour, weekly and monthly meters. The workspace ID + auth cookie scrape is the legacy fallback. Changes take effect after daemon restart.', + desc: 'Configure OpenCode Go quota tracking. onWatch reads your plan\'s 5-hour, weekly and monthly meters from the OpenCode console, using a service-account key or your browser session. Changes take effect after daemon restart.', fields: [ - { id: 'api_key', label: 'Usage API Key', type: 'password', placeholder: 'Not configured', hint: 'OpenCode console service-account key (oc_sk_...) of the account with the Go subscription. Reads the plan meters from the Go status API; preferred over the cookie. Overrides OPENCODE_GO_API_KEY from .env.', sensitive: true }, - { id: 'workspace_id', label: 'Workspace ID (legacy)', type: 'text', placeholder: 'wrk_...', hint: 'Legacy scrape mode only: your OpenCode Go workspace ID. Overrides OPENCODE_GO_WORKSPACE_ID from .env.' }, - { id: 'auth_cookie', label: 'Auth Cookie (legacy)', type: 'password', placeholder: 'Not configured', hint: 'Legacy scrape mode only: the auth cookie used to scrape the dashboard. Overrides OPENCODE_GO_AUTH_COOKIE from .env.', sensitive: true }, + { id: 'api_key', label: 'Usage API Key', type: 'password', placeholder: 'Not configured', hint: 'OpenCode console service-account key (oc_sk_...) of the account with the Go subscription. Used instead of the session cookie when set. Overrides OPENCODE_GO_API_KEY from .env.', sensitive: true }, + { id: 'workspace_id', label: 'Workspace ID', type: 'text', placeholder: 'wrk_...', hint: 'Session mode: your OpenCode Go workspace ID (wrk_...). Overrides OPENCODE_GO_WORKSPACE_ID from .env.' }, + { id: 'auth_cookie', label: 'Session Cookie', type: 'password', placeholder: 'Not configured', hint: 'Session mode: the __Host-console_session cookie value from opencode.ai (the old auth cookie no longer works). Overrides OPENCODE_GO_AUTH_COOKIE from .env.', sensitive: true }, ], }, ollama: { diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 50f7762b..727c9de8 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -4,8 +4,9 @@ then tear them down after all tests complete. """ import os -import signal import subprocess +import sys +import tempfile import time from pathlib import Path from typing import Generator @@ -24,15 +25,33 @@ USERNAME = "admin" PASSWORD = "testpass123" -# Paths +# Paths (tempdir + .exe so the suite also runs on Windows) PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent -MOCK_BINARY = "/tmp/mockserver-test" -ONWATCH_BINARY = "/tmp/onwatch-test" +TMP_DIR = Path(tempfile.gettempdir()) +EXE_SUFFIX = ".exe" if os.name == "nt" else "" +MOCK_BINARY = str(TMP_DIR / f"mockserver-test{EXE_SUFFIX}") +ONWATCH_BINARY = str(TMP_DIR / f"onwatch-test{EXE_SUFFIX}") # E2E isolation: override HOME so the canonical DB path (~/.onwatch/data/onwatch.db) # does not exist. This prevents main.go's fixExplicitDBPath() from redirecting to # the production database. -E2E_HOME = "/tmp/onwatch-e2e-home" -DB_PATH = "/tmp/onwatch-e2e.db" +E2E_HOME = str(TMP_DIR / "onwatch-e2e-home") +DB_PATH = str(TMP_DIR / "onwatch-e2e.db") + + +def pytest_configure(config) -> None: + """Refuse to run on a developer Mac. + + onWatch auto-detects Cursor credentials from the macOS Keychain (not from + HOME) and can refresh and rewrite them, so an e2e daemon on a signed-in + Mac could rotate real tokens. CI runners are clean; set + ONWATCH_E2E_ALLOW_HOST=1 only on a machine with no real credentials. + """ + if sys.platform == "darwin" and not os.environ.get("CI") and os.environ.get("ONWATCH_E2E_ALLOW_HOST") != "1": + pytest.exit( + "e2e suite is CI-only on macOS: the daemon can read and refresh real " + "Keychain credentials. Set ONWATCH_E2E_ALLOW_HOST=1 on a clean machine.", + returncode=2, + ) def _wait_for_http(url: str, timeout: float = 30.0, interval: float = 0.5) -> bool: @@ -54,13 +73,74 @@ def _kill_process(proc: subprocess.Popen) -> None: """Kill a subprocess and wait for it to exit.""" if proc.poll() is None: try: - proc.send_signal(signal.SIGTERM) + proc.terminate() # SIGTERM on Unix, TerminateProcess on Windows proc.wait(timeout=5) except (subprocess.TimeoutExpired, OSError): proc.kill() proc.wait(timeout=5) +def remove_instance_files(db_path: str, home: str) -> None: + """Remove an onwatch instance's database files and HOME directory.""" + import shutil + for path in [db_path, f"{db_path}-journal", f"{db_path}-wal", f"{db_path}-shm"]: + try: + os.unlink(path) + except OSError: + # Already gone, or still locked on Windows while the daemon exits; + # start_onwatch clears leftovers before the next run. + continue + if os.path.exists(home): + shutil.rmtree(home, ignore_errors=True) + + +def start_onwatch(port: int, db_path: str, home: str, provider_env: dict) -> subprocess.Popen: + """Start the built onwatch binary in an isolated HOME and wait for /login. + + Anthropic is pinned to a fake token in statusline mode: onWatch never + reads the real Claude Code keychain entry and never calls the usage or + OAuth refresh endpoints, so running the suite on a developer machine + cannot rotate (and log out) real credentials. + """ + remove_instance_files(db_path, home) + os.makedirs(home, exist_ok=True) + env = os.environ.copy() + env.update({ + "HOME": home, + "USERPROFILE": home, # Windows home directory + # Windows keeps the test PID file under LOCALAPPDATA; a shared one + # lets a second daemon stop the first on startup. + "LOCALAPPDATA": os.path.join(home, "AppData", "Local"), + "ONWATCH_ADMIN_PASS": PASSWORD, + "ONWATCH_TEST_MODE": "1", + "ANTHROPIC_TOKEN": "anth_test_e2e_token", + "ANTHROPIC_SOURCE": "statusline", + }) + env.update(provider_env) + # A file, not a pipe: an unread pipe fills up and blocks the daemon. CI + # prints these logs when a job fails. The child keeps its own handle, so + # ours can be closed as soon as the process starts. + with open(TMP_DIR / f"onwatch-e2e-{port}.log", "w") as log: + proc = subprocess.Popen( + [ + ONWATCH_BINARY, + "--debug", + f"--port={port}", + "--interval=10", + "--test", + f"--db={db_path}", + ], + env=env, + stdout=log, + stderr=subprocess.STDOUT, + ) + ready = _wait_for_http(f"http://localhost:{port}/login", timeout=30) + if not ready: + _kill_process(proc) + assert ready, f"onWatch on port {port} did not start in time" + return proc + + @pytest.fixture(scope="session") def mock_server() -> Generator[subprocess.Popen, None, None]: """Build and start the mock server binary.""" @@ -104,17 +184,6 @@ def mock_server() -> Generator[subprocess.Popen, None, None]: @pytest.fixture(scope="session") def onwatch_server(mock_server: subprocess.Popen) -> Generator[subprocess.Popen, None, None]: """Build and start the onwatch binary.""" - # Clean up any stale DB and home directory - import shutil - for path in [DB_PATH, f"{DB_PATH}-journal", f"{DB_PATH}-wal", f"{DB_PATH}-shm"]: - try: - os.unlink(path) - except OSError: - pass - if os.path.exists(E2E_HOME): - shutil.rmtree(E2E_HOME) - os.makedirs(E2E_HOME, exist_ok=True) - # Build onwatch build_cmd = ["go", "build"] build_tags = os.environ.get("ONWATCH_E2E_GO_BUILD_TAGS", "").strip() @@ -131,35 +200,12 @@ def onwatch_server(mock_server: subprocess.Popen) -> Generator[subprocess.Popen, ) assert result.returncode == 0, f"onWatch build failed: {result.stderr}" - env = os.environ.copy() - env.update({ - "HOME": E2E_HOME, - "ONWATCH_ADMIN_PASS": PASSWORD, - "ONWATCH_TEST_MODE": "1", + proc = start_onwatch(ONWATCH_PORT, DB_PATH, E2E_HOME, { "SYNTHETIC_API_KEY": "syn_test_e2e_key", "ZAI_API_KEY": "zai_test_e2e_key", "ZAI_BASE_URL": f"http://localhost:{MOCK_PORT}", - "ANTHROPIC_TOKEN": "anth_test_e2e_token", }) - proc = subprocess.Popen( - [ - ONWATCH_BINARY, - "--debug", - f"--port={ONWATCH_PORT}", - "--interval=10", - "--test", - f"--db={DB_PATH}", - ], - env=env, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - ) - - # Wait for onwatch to be ready (login page returns 200) - ready = _wait_for_http(f"{BASE_URL}/login", timeout=30) - assert ready, "onWatch server did not start in time" - yield proc _kill_process(proc) @@ -168,14 +214,7 @@ def onwatch_server(mock_server: subprocess.Popen) -> Generator[subprocess.Popen, os.unlink(ONWATCH_BINARY) except OSError: pass - for path in [DB_PATH, f"{DB_PATH}-journal", f"{DB_PATH}-wal", f"{DB_PATH}-shm"]: - try: - os.unlink(path) - except OSError: - pass - import shutil - if os.path.exists(E2E_HOME): - shutil.rmtree(E2E_HOME, ignore_errors=True) + remove_instance_files(DB_PATH, E2E_HOME) @pytest.fixture(autouse=True, scope="session") @@ -208,5 +247,6 @@ def dashboard_page(authenticated_page): def settings_page(authenticated_page): """Navigate to the settings page and return the page.""" authenticated_page.goto(f"{BASE_URL}/settings") - authenticated_page.wait_for_selector(".settings-page", timeout=10000) + # data-ready: every settings control is wired, not just rendered. + authenticated_page.wait_for_selector(".settings-page[data-ready]", timeout=15000) return authenticated_page diff --git a/tests/e2e/page_objects/dashboard_page.py b/tests/e2e/page_objects/dashboard_page.py index cd1193f6..29db262b 100644 --- a/tests/e2e/page_objects/dashboard_page.py +++ b/tests/e2e/page_objects/dashboard_page.py @@ -148,7 +148,7 @@ def has_settings_link(self) -> bool: def navigate_to_settings_password(self) -> None: """Navigate to the settings page and open the General tab for password.""" self.page.click("#settings-btn") - self.page.wait_for_selector(".settings-page", timeout=5000) + self.page.wait_for_selector(".settings-page[data-ready]", timeout=15000) self.page.click('.settings-tab[data-tab="general"]') self.page.wait_for_selector("#panel-general:not([hidden])", timeout=5000) diff --git a/tests/e2e/page_objects/settings_page.py b/tests/e2e/page_objects/settings_page.py index 1f4e600d..e84a40bd 100644 --- a/tests/e2e/page_objects/settings_page.py +++ b/tests/e2e/page_objects/settings_page.py @@ -16,7 +16,7 @@ def __init__(self, page: Page) -> None: def goto(self) -> None: """Navigate to the settings page.""" self.page.goto(f"{BASE_URL}/settings") - self.page.wait_for_selector(".settings-page", timeout=10000) + self.page.wait_for_selector(".settings-page[data-ready]", timeout=15000) def select_tab(self, tab_name: str) -> None: """Click a settings tab by its data-tab attribute.""" diff --git a/tests/e2e/tests/test_dashboard.py b/tests/e2e/tests/test_dashboard.py index b975591f..90af67a1 100644 --- a/tests/e2e/tests/test_dashboard.py +++ b/tests/e2e/tests/test_dashboard.py @@ -80,4 +80,4 @@ def test_last_updated_displays(self, dashboard_page: Page) -> None: dashboard_page.wait_for_timeout(3000) text = dash.get_last_updated() assert text != "" - assert "Last updated" in text + assert text.startswith("Updated ") diff --git a/tests/e2e/tests/test_data_tables.py b/tests/e2e/tests/test_data_tables.py index 08e031ea..26286a9c 100644 --- a/tests/e2e/tests/test_data_tables.py +++ b/tests/e2e/tests/test_data_tables.py @@ -29,11 +29,12 @@ def test_cycles_table_has_sort_headers(self, dashboard_page: Page) -> None: headers = dashboard_page.query_selector_all( "#cycles-table thead th[data-sort-key]" ) - assert len(headers) >= 5 + # Logging history: row number, time, then one column per quota. + assert len(headers) >= 3 sort_keys = [h.get_attribute("data-sort-key") for h in headers] + assert "id" in sort_keys assert "start" in sort_keys - assert "peak" in sort_keys - assert "total" in sort_keys + assert any(k.startswith("cq_") for k in sort_keys) def test_cycles_pagination_controls(self, dashboard_page: Page) -> None: """The cycles section should have pagination controls.""" diff --git a/tests/e2e/tests/test_provider_balances.py b/tests/e2e/tests/test_provider_balances.py new file mode 100644 index 00000000..262004da --- /dev/null +++ b/tests/e2e/tests/test_provider_balances.py @@ -0,0 +1,142 @@ +"""E2E tests for OpenCode Go session-cookie mode and the DeepSeek and Moonshot tabs. + +A dedicated onwatch instance polls the mock server with real agents: +OpenCode Go through the console status API with a session cookie (#134), +and the DeepSeek and Moonshot balance APIs (#137). +""" +import json +import urllib.request +from typing import Generator + +import pytest +from playwright.sync_api import Page, expect + +from conftest import ( + MOCK_URL, + PASSWORD, + TMP_DIR, + USERNAME, + _kill_process, + remove_instance_files, + start_onwatch, +) + +PORT = 19215 +BASE = f"http://localhost:{PORT}" +DB = str(TMP_DIR / "onwatch-e2e-providers.db") +HOME = str(TMP_DIR / "onwatch-e2e-providers-home") + + +def _mock_counts() -> dict: + with urllib.request.urlopen(f"{MOCK_URL}/admin/requests", timeout=5) as resp: + return json.load(resp) + + +@pytest.fixture(scope="module") +def provider_server(servers) -> Generator[str, None, None]: + """onwatch with only OpenCode Go (session cookie), DeepSeek and Moonshot.""" + proc = start_onwatch(PORT, DB, HOME, { + "OPENCODE_GO_BASE_URL": MOCK_URL, + "OPENCODE_GO_WORKSPACE_ID": "wrk_e2e", + # A bare value: onWatch must send it as __Host-console_session. + "OPENCODE_GO_AUTH_COOKIE": "oc_session_e2e", + "DEEPSEEK_API_KEY": "sk_balance_e2e", + "DEEPSEEK_BASE_URL": MOCK_URL, + "MOONSHOT_API_KEY": "sk_balance_e2e", + "MOONSHOT_BASE_URL": MOCK_URL, + }) + yield BASE + _kill_process(proc) + remove_instance_files(DB, HOME) + + +@pytest.fixture +def logged_in(page: Page, provider_server: str) -> Page: + # Kept for failure messages: a render that never happens is otherwise + # impossible to diagnose from CI output. + page.console_errors = [] + page.on("pageerror", lambda e: page.console_errors.append(f"pageerror: {e}")) + page.on("console", lambda m: m.type == "error" and page.console_errors.append(f"console: {m.text}")) + page.goto(f"{provider_server}/login") + page.fill("#username", USERNAME) + page.fill("#password", PASSWORD) + page.click("button.login-button") + page.wait_for_url(f"{provider_server}/", timeout=10000) + return page + + +def _current(page: Page, provider: str) -> str: + """The provider's /api/current body, for failure messages.""" + return page.evaluate( + "async (p) => JSON.stringify(await (await fetch(`/api/current?provider=${p}`)).json())", + provider, + ) + + +def _open_tab(page: Page, provider: str, card_selector: str, count: int) -> None: + # The first poll runs at startup and then every 10s; reload until the + # stored snapshot renders the cards. + last_error = None + for _ in range(12): + page.goto(f"{BASE}/?provider={provider}") + try: + # Functions, not bare expressions: the dashboard's CSP forbids + # unsafe-eval, which Playwright needs for an expression string. + page.wait_for_function( + "([sel, n]) => document.querySelectorAll(sel).length === n", + arg=[card_selector, count], + timeout=4000, + ) + return + except Exception as e: # timeout or a navigation mid-wait; retry + last_error = e + page.wait_for_timeout(1000) + grid = page.evaluate( + "(id) => { const g = document.getElementById(id); return g ? g.outerHTML.slice(0, 400) : 'missing'; }", + f"quota-grid-{provider}", + ) + raise AssertionError( + f"{provider}: {count} cards never rendered ({str(last_error)[:300]}); url={page.url}; grid={grid}; " + f"errors={getattr(page, 'console_errors', [])}; mock requests={_mock_counts()}; " + f"/api/current={_current(page, provider)}" + ) + + +def _chart_labels(page: Page) -> list: + page.wait_for_function("() => State.chart && State.chart.data.datasets.length > 0", timeout=10000) + return page.evaluate("() => State.chart.data.datasets.map(d => d.label)") + + +def _expect_logging_row(page: Page, text: str) -> None: + # The table loads when scrolled into view; until then it shows the + # template's placeholder row, so wait for the value itself. + page.locator(".cycles-section").scroll_into_view_if_needed() + expect(page.locator("#cycles-tbody tr").first).to_contain_text(text, timeout=15000) + + +class TestOpenCodeSessionCookie: + def test_cards_show_dollar_meters(self, logged_in: Page) -> None: + _open_tab(logged_in, "opencode", "#quota-grid-opencode .opencode-card", 3) + expect(logged_in.locator("#fraction-opencode-weekly")).to_have_text("$3.84 / $30.00") + expect(logged_in.locator("#percent-opencode-weekly")).to_have_text("12.8%") + expect(logged_in.locator("#fraction-opencode-monthly")).to_have_text("$3.84 / $60.00") + # The mock only answers with __Host-console_session + x-org-id. + assert _mock_counts()["opencode"] >= 1 + + +class TestBalanceTabs: + def test_deepseek_tab_renders_balances(self, logged_in: Page) -> None: + _open_tab(logged_in, "deepseek", "#quota-grid-deepseek .balance-card", 3) + amounts = logged_in.locator("#quota-grid-deepseek .usage-percent").all_inner_texts() + assert amounts == ["$4.02", "$0.32", "$3.70"] + labels = _chart_labels(logged_in) + assert labels == ["Total Balance", "Granted", "Topped Up"], "DeepSeek chart fell back to another provider" + _expect_logging_row(logged_in, "$4.02") + + def test_moonshot_tab_renders_balances(self, logged_in: Page) -> None: + _open_tab(logged_in, "moonshot", "#quota-grid-moonshot .balance-card", 3) + amounts = logged_in.locator("#quota-grid-moonshot .usage-percent").all_inner_texts() + assert amounts == ["19.47", "5.00", "14.47"] + labels = _chart_labels(logged_in) + assert labels == ["Available", "Voucher", "Cash"], "Moonshot chart fell back to another provider" + _expect_logging_row(logged_in, "19.47") diff --git a/tools/perf-monitor/main.go b/tools/perf-monitor/main.go index ef041787..8164295a 100644 --- a/tools/perf-monitor/main.go +++ b/tools/perf-monitor/main.go @@ -204,16 +204,9 @@ func runMonitoring(pid, port int, totalDuration time.Duration) *Report { func findonWatchProcess(port int) int { // Try PID file - pidFile := filepath.Join(os.Getenv("HOME"), "Library", "Application Support", "onwatch", "onwatch.pid") - if runtime.GOOS != "darwin" { - pidFile = filepath.Join(os.Getenv("HOME"), ".local", "share", "onwatch", "onwatch.pid") - } - - if data, err := os.ReadFile(pidFile); err == nil { - if pid, err := strconv.Atoi(strings.TrimSpace(string(data))); err == nil && pid > 0 { - if isProcessRunning(pid) { - return pid - } + if data, err := os.ReadFile(pidFilePath()); err == nil { + if pid := parsePIDFile(data); pid > 0 && isProcessRunning(pid) { + return pid } } @@ -234,21 +227,24 @@ func findonWatchProcess(port int) int { return 0 } -func isProcessRunning(pid int) bool { - proc, _ := os.FindProcess(pid) - if proc == nil { - return false +// parsePIDFile extracts the PID from onWatch's PID file, which holds +// "pid:port" (older builds wrote a bare "pid"). Returns 0 if unparseable. +func parsePIDFile(data []byte) int { + s := strings.TrimSpace(string(data)) + if idx := strings.IndexByte(s, ':'); idx >= 0 { + s = s[:idx] } - // Signal 0 check - return proc.Signal(os.Signal(nil)) == nil + pid, err := strconv.Atoi(s) + if err != nil || pid <= 0 { + return 0 + } + return pid } +// isOnwatchProcess reports whether pid runs an onWatch binary, judged by the +// executable's base name (see processCommandName). func isOnwatchProcess(pid int) bool { - out, err := exec.Command("ps", "-p", strconv.Itoa(pid), "-o", "comm=").Output() - if err != nil { - return false - } - return strings.Contains(strings.ToLower(string(out)), "onwatch") + return strings.Contains(strings.ToLower(processCommandName(pid)), "onwatch") } // parseArgs reads the optional [port] [duration] positionals and the --restart @@ -284,16 +280,16 @@ func parseArgs(args []string) (port int, duration time.Duration, shouldRestart b func stoponWatch(port int) { // Try PID file first - pidFile := filepath.Join(os.Getenv("HOME"), "Library", "Application Support", "onwatch", "onwatch.pid") - if runtime.GOOS != "darwin" { - pidFile = filepath.Join(os.Getenv("HOME"), ".local", "share", "onwatch", "onwatch.pid") - } - + pidFile := pidFilePath() if data, err := os.ReadFile(pidFile); err == nil { - if pid, err := strconv.Atoi(strings.TrimSpace(string(data))); err == nil && pid > 0 { + // A PID file left by a crashed onWatch can name a PID the OS has + // since reused, so only a live onWatch process is stopped; anything + // else means the file is stale. + if pid := parsePIDFile(data); pid > 0 && isProcessRunning(pid) && isOnwatchProcess(pid) { if proc, err := os.FindProcess(pid); err == nil { - proc.Signal(os.Interrupt) - fmt.Printf(" Stopped process (PID: %d) via PID file\n", pid) + if err := stopProcess(proc); err == nil { + fmt.Printf(" Stopped process (PID: %d) via PID file\n", pid) + } time.Sleep(500 * time.Millisecond) } } @@ -308,7 +304,7 @@ func stoponWatch(port int) { if pid, err := strconv.Atoi(strings.TrimSpace(line)); err == nil && pid > 0 { if isOnwatchProcess(pid) { if proc, err := os.FindProcess(pid); err == nil { - proc.Signal(os.Interrupt) + _ = stopProcess(proc) fmt.Printf(" Stopped process (PID: %d) on port %d\n", pid, port) } } @@ -319,36 +315,31 @@ func stoponWatch(port int) { } func startonWatch(port int) int { - // Find onwatch binary in various locations - possiblePaths := []string{ - "./onwatch", - "../onwatch", - "../../onwatch", - "/Users/prakersh/project./onwatch/onwatch", - } - + // Look for the onwatch binary near the working directory, then on PATH. binaryPath := "" - for _, path := range possiblePaths { - if _, err := os.Stat(path); err == nil { - binaryPath = path + for _, dir := range []string{".", "..", filepath.Join("..", "..")} { + candidate := filepath.Join(dir, onwatchBinaryName) + if _, err := os.Stat(candidate); err == nil { + binaryPath = candidate break } } - if binaryPath == "" { - // Try PATH + workDir := "" + if binaryPath != "" { + if abs, err := filepath.Abs(binaryPath); err == nil { + binaryPath = abs + } + // Run from the binary's directory so it can find .env and database. + workDir = filepath.Dir(binaryPath) + } else { + // Try PATH (LookPath adds .exe on Windows) binaryPath = "onwatch" } - // Change to the binary's directory so it can find .env and database - binaryDir := filepath.Dir(binaryPath) - if binaryDir != "." && binaryDir != "" { - os.Chdir(binaryDir) - binaryPath = "./onwatch" - } - // Start onwatch in debug mode cmd := exec.Command(binaryPath, "--debug", "--port", strconv.Itoa(port)) + cmd.Dir = workDir cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr cmd.Env = os.Environ() @@ -360,15 +351,22 @@ func startonWatch(port int) int { pid := cmd.Process.Pid + // Reap the child so an early exit is seen as an exit: an unreaped child + // is a zombie on Unix and would still look alive to a signal-0 probe. + exited := make(chan struct{}) + go func() { + _ = cmd.Wait() + close(exited) + }() + // Wait for it to be ready fmt.Println(" Waiting for onWatch to be ready...") for i := 0; i < 30; i++ { - time.Sleep(200 * time.Millisecond) - - // Check if process is still running - if !isProcessRunning(pid) { + select { + case <-exited: fmt.Println(" ❌ onWatch process died during startup") return 0 + case <-time.After(200 * time.Millisecond): } // Check if port is listening diff --git a/tools/perf-monitor/main_test.go b/tools/perf-monitor/main_test.go index 06c7f211..d6c5ee36 100644 --- a/tools/perf-monitor/main_test.go +++ b/tools/perf-monitor/main_test.go @@ -8,6 +8,7 @@ import ( "net/http/httptest" "os" "os/exec" + "os/signal" "path/filepath" "runtime" "strconv" @@ -25,19 +26,144 @@ func captureStdout(t *testing.T, fn func()) string { t.Fatalf("create stdout pipe: %v", err) } defer r.Close() - os.Stdout = w - defer func() { os.Stdout = oldStdout }() + // Closes the writer if fn fails the test before the explicit Close below, + // so the reader goroutine still sees EOF. + defer w.Close() + + // Drain the pipe concurrently. Pipe buffers are small (a few KB on + // Windows), so reading only after fn returns deadlocks once fn writes + // more than the buffer holds. + type readResult struct { + out []byte + err error + } + done := make(chan readResult, 1) + go func() { + out, err := io.ReadAll(r) + done <- readResult{out: out, err: err} + }() - fn() + os.Stdout = w + func() { + defer func() { os.Stdout = oldStdout }() + fn() + }() if err := w.Close(); err != nil { t.Fatalf("close writer: %v", err) } - out, err := io.ReadAll(r) + res := <-done + if res.err != nil { + t.Fatalf("read stdout: %v", res.err) + } + return string(res.out) +} + +// isolatePIDFile points the tool's PID file lookup at a fresh temp home and +// returns the path it will read, with its directory created. HOME (Unix), +// USERPROFILE and LOCALAPPDATA (Windows) are all redirected so the real +// onWatch PID file is never read, signalled or removed. +func isolatePIDFile(t *testing.T) string { + t.Helper() + home := t.TempDir() + t.Setenv("HOME", home) + t.Setenv("USERPROFILE", home) + t.Setenv("LOCALAPPDATA", filepath.Join(home, "AppData", "Local")) + path := pidFilePath() + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + t.Fatalf("mkdir pid dir: %v", err) + } + return path +} + +// exitAtStartEnv makes a copy of this test binary exit 0 before flag parsing, +// so it can stand in for an onwatch binary that dies during startup. +const exitAtStartEnv = "PERF_MONITOR_EXIT_AT_START" + +func init() { + if os.Getenv(exitAtStartEnv) == "1" { + os.Exit(0) + } +} + +// idleHelperEnv turns a re-executed copy of this test binary into an idle +// process that exits cleanly on os.Interrupt, like the onWatch daemon. +const idleHelperEnv = "PERF_MONITOR_IDLE_HELPER" + +func TestHelperIdleProcess(t *testing.T) { + if os.Getenv(idleHelperEnv) != "1" { + return + } + // Notify also re-enables SIGINT if the test run started with it ignored. + interrupted := make(chan os.Signal, 1) + signal.Notify(interrupted, os.Interrupt) + select { + case <-interrupted: + os.Exit(0) + case <-time.After(2 * time.Minute): + os.Exit(3) + } +} + +// startIdleHelper starts an idle helper process running this test binary +// (perf-monitor.test, so not an onWatch process). The returned channel closes +// once the process has exited and been reaped. +func startIdleHelper(t *testing.T) (*exec.Cmd, <-chan struct{}) { + t.Helper() + return startIdleHelperBinary(t, os.Args[0]) +} + +// startOnwatchNamedIdleHelper starts an idle helper from a copy of this test +// binary named like the onWatch executable, so it passes isOnwatchProcess. +func startOnwatchNamedIdleHelper(t *testing.T) (*exec.Cmd, <-chan struct{}) { + t.Helper() + selfPath, err := os.Executable() if err != nil { - t.Fatalf("read stdout: %v", err) + t.Fatalf("locate test binary: %v", err) + } + self, err := os.ReadFile(selfPath) + if err != nil { + t.Fatalf("read test binary: %v", err) + } + bin := filepath.Join(t.TempDir(), onwatchBinaryName) + if err := os.WriteFile(bin, self, 0o755); err != nil { + t.Fatalf("write onwatch-named helper: %v", err) + } + return startIdleHelperBinary(t, bin) +} + +func startIdleHelperBinary(t *testing.T, bin string) (*exec.Cmd, <-chan struct{}) { + t.Helper() + cmd := exec.Command(bin, "-test.run=^TestHelperIdleProcess$") + cmd.Env = append(os.Environ(), idleHelperEnv+"=1") + if err := cmd.Start(); err != nil { + t.Fatalf("start helper process: %v", err) + } + exited := make(chan struct{}) + go func() { + _ = cmd.Wait() + close(exited) + }() + t.Cleanup(func() { + _ = cmd.Process.Kill() + <-exited + }) + return cmd, exited +} + +func TestParsePIDFile(t *testing.T) { + cases := map[string]int{ + "1234:9211\n": 1234, // current onWatch format: pid:port + "1234\n": 1234, // older bare-pid format + "not-a-pid": 0, + "": 0, + "-5:9211": 0, + } + for in, want := range cases { + if got := parsePIDFile([]byte(in)); got != want { + t.Errorf("parsePIDFile(%q) = %d, want %d", in, got, want) + } } - return string(out) } func TestCalculateStats_EmptySamples(t *testing.T) { @@ -187,21 +313,7 @@ func TestSaveReport_WritesJSONFile(t *testing.T) { } func TestFindOnWatchProcess_InvalidPidFileFallsBackToPortScanAndReturnsZero(t *testing.T) { - home := t.TempDir() - if err := os.Setenv("HOME", home); err != nil { - t.Fatalf("set HOME: %v", err) - } - - var pidDir string - if runtime.GOOS == "darwin" { - pidDir = filepath.Join(home, "Library", "Application Support", "onwatch") - } else { - pidDir = filepath.Join(home, ".local", "share", "onwatch") - } - if err := os.MkdirAll(pidDir, 0o755); err != nil { - t.Fatalf("mkdir pid dir: %v", err) - } - pidFile := filepath.Join(pidDir, "onwatch.pid") + pidFile := isolatePIDFile(t) if err := os.WriteFile(pidFile, []byte("not-a-pid"), 0o644); err != nil { t.Fatalf("write pid file: %v", err) } @@ -242,19 +354,15 @@ func TestIsOnwatchProcess_UnknownPidReturnsFalse(t *testing.T) { } } -func TestIsOnwatchProcess_ProcessNameCoverageForCurrentProcess(t *testing.T) { - if runtime.GOOS != "darwin" && runtime.GOOS != "linux" { - t.Skip("ps-based process checks are only used on darwin/linux") - } - - currentIsOnwatch := isOnwatchProcess(os.Getpid()) - out, err := exec.Command("ps", "-p", strconv.Itoa(os.Getpid()), "-o", "comm=").Output() - if err != nil { - t.Fatalf("read current process name: %v", err) +// Only the executable's base name counts: the test binary (perf-monitor.test) +// is not onWatch, a binary named onwatch is, wherever it lives. +func TestIsOnwatchProcess_MatchesExecutableBaseName(t *testing.T) { + if isOnwatchProcess(os.Getpid()) { + t.Fatalf("the perf-monitor test binary must not be identified as onwatch (name %q)", processCommandName(os.Getpid())) } - want := strings.Contains(strings.ToLower(string(out)), "onwatch") - if currentIsOnwatch != want { - t.Fatalf("expected %v for current process name %q, got %v", want, string(out), currentIsOnwatch) + helper, _ := startOnwatchNamedIdleHelper(t) + if !isOnwatchProcess(helper.Process.Pid) { + t.Fatalf("a process running %s must be identified as onwatch (name %q)", onwatchBinaryName, processCommandName(helper.Process.Pid)) } } @@ -283,8 +391,10 @@ func TestGenerateLoad_CollectsMetricsDeterministically(t *testing.T) { if m.Count < 1 { t.Fatalf("expected at least one request for %s", m.Endpoint) } - if m.MinTime <= 0 || m.MaxTime <= 0 || m.AvgTime <= 0 { - t.Fatalf("expected positive durations for %s, got min=%v avg=%v max=%v", m.Endpoint, m.MinTime, m.AvgTime, m.MaxTime) + // A local request can measure 0s on a coarse clock (Windows), so + // check the aggregates are consistent rather than strictly positive. + if m.MinTime < 0 || m.MinTime > m.AvgTime || m.AvgTime > m.MaxTime { + t.Fatalf("inconsistent durations for %s: min=%v avg=%v max=%v", m.Endpoint, m.MinTime, m.AvgTime, m.MaxTime) } } }) @@ -295,71 +405,45 @@ func TestGenerateLoad_CollectsMetricsDeterministically(t *testing.T) { } func TestIsProcessRunning_CurrentAndNonexistentPID(t *testing.T) { - gotCurrent := isProcessRunning(os.Getpid()) - proc, err := os.FindProcess(os.Getpid()) - if err != nil { - t.Fatalf("find current process: %v", err) - } - wantCurrent := proc.Signal(os.Signal(nil)) == nil - if gotCurrent != wantCurrent { - t.Fatalf("expected current process running=%v, got %v", wantCurrent, gotCurrent) + if !isProcessRunning(os.Getpid()) { + t.Fatal("expected current process to be running") } - if isProcessRunning(999999) { t.Fatal("expected nonexistent pid to not be running") } + if isProcessRunning(0) || isProcessRunning(-1) { + t.Fatal("expected non-positive pids to not be running") + } } -func TestFindOnWatchProcess_ValidPIDInFileUsesIsProcessRunningBranch(t *testing.T) { - home := t.TempDir() - oldHome := os.Getenv("HOME") - t.Cleanup(func() { - _ = os.Setenv("HOME", oldHome) - }) - if err := os.Setenv("HOME", home); err != nil { - t.Fatalf("set HOME: %v", err) +func TestIsProcessRunning_ExitedProcess(t *testing.T) { + cmd, exited := startIdleHelper(t) + if !isProcessRunning(cmd.Process.Pid) { + t.Fatalf("expected helper %d to be running", cmd.Process.Pid) } - - var pidDir string - if runtime.GOOS == "darwin" { - pidDir = filepath.Join(home, "Library", "Application Support", "onwatch") - } else { - pidDir = filepath.Join(home, ".local", "share", "onwatch") + if err := cmd.Process.Kill(); err != nil { + t.Fatalf("kill helper: %v", err) } - if err := os.MkdirAll(pidDir, 0o755); err != nil { - t.Fatalf("mkdir pid dir: %v", err) + <-exited + if isProcessRunning(cmd.Process.Pid) { + t.Fatalf("expected exited helper %d to not be running", cmd.Process.Pid) } - pidFile := filepath.Join(pidDir, "onwatch.pid") - if err := os.WriteFile(pidFile, []byte(strconv.Itoa(os.Getpid())), 0o644); err != nil { +} + +func TestFindOnWatchProcess_ValidPIDInFileUsesIsProcessRunningBranch(t *testing.T) { + pidFile := isolatePIDFile(t) + // onWatch writes "pid:port". + if err := os.WriteFile(pidFile, []byte(strconv.Itoa(os.Getpid())+":65529"), 0o644); err != nil { t.Fatalf("write pid file: %v", err) } - got := findonWatchProcess(65529) - if got != 0 && got != os.Getpid() { - t.Fatalf("expected pid file branch to return 0 or current pid, got %d", got) + if got := findonWatchProcess(65529); got != os.Getpid() { + t.Fatalf("expected pid file branch to return current pid %d, got %d", os.Getpid(), got) } } func TestStopOnWatch_RemovesInvalidPIDFileSafely(t *testing.T) { - home := t.TempDir() - oldHome := os.Getenv("HOME") - t.Cleanup(func() { - _ = os.Setenv("HOME", oldHome) - }) - if err := os.Setenv("HOME", home); err != nil { - t.Fatalf("set HOME: %v", err) - } - - var pidDir string - if runtime.GOOS == "darwin" { - pidDir = filepath.Join(home, "Library", "Application Support", "onwatch") - } else { - pidDir = filepath.Join(home, ".local", "share", "onwatch") - } - if err := os.MkdirAll(pidDir, 0o755); err != nil { - t.Fatalf("mkdir pid dir: %v", err) - } - pidFile := filepath.Join(pidDir, "onwatch.pid") + pidFile := isolatePIDFile(t) if err := os.WriteFile(pidFile, []byte("invalid-pid"), 0o644); err != nil { t.Fatalf("write pid file: %v", err) } @@ -489,8 +573,19 @@ func TestStartOnWatch_ProcessDiesDuringStartupReturnsZero(t *testing.T) { } defer func() { _ = os.Chdir(oldWD) }() - if err := os.WriteFile(filepath.Join(tempDir, "onwatch"), []byte("#!/bin/sh\nexit 0\n"), 0o755); err != nil { - t.Fatalf("write failing onwatch script: %v", err) + // A copy of this test binary stands in for onwatch and exits at once on + // every platform. A shell script would not be executable on Windows. + t.Setenv(exitAtStartEnv, "1") + selfPath, err := os.Executable() + if err != nil { + t.Fatalf("locate test binary: %v", err) + } + self, err := os.ReadFile(selfPath) + if err != nil { + t.Fatalf("read test binary: %v", err) + } + if err := os.WriteFile(filepath.Join(tempDir, onwatchBinaryName), self, 0o755); err != nil { + t.Fatalf("write failing onwatch binary: %v", err) } pid := startonWatch(65524) @@ -500,52 +595,56 @@ func TestStartOnWatch_ProcessDiesDuringStartupReturnsZero(t *testing.T) { } func TestStopOnWatch_ValidPIDFileSignalsProcess(t *testing.T) { - home := t.TempDir() - oldHome := os.Getenv("HOME") - t.Cleanup(func() { - _ = os.Setenv("HOME", oldHome) - }) - if err := os.Setenv("HOME", home); err != nil { - t.Fatalf("set HOME: %v", err) + pidFile := isolatePIDFile(t) + helper, exited := startOnwatchNamedIdleHelper(t) + if err := os.WriteFile(pidFile, []byte(strconv.Itoa(helper.Process.Pid)+":65523"), 0o644); err != nil { + t.Fatalf("write pid file: %v", err) } - var pidDir string - if runtime.GOOS == "darwin" { - pidDir = filepath.Join(home, "Library", "Application Support", "onwatch") - } else { - pidDir = filepath.Join(home, ".local", "share", "onwatch") + stoponWatch(65523) + + select { + case <-exited: + case <-time.After(10 * time.Second): + t.Fatalf("expected helper process %d to be stopped", helper.Process.Pid) } - if err := os.MkdirAll(pidDir, 0o755); err != nil { - t.Fatalf("mkdir pid dir: %v", err) + if isProcessRunning(helper.Process.Pid) { + t.Fatalf("expected helper process %d to be gone", helper.Process.Pid) } - - helpCmd := exec.Command("sh", "-c", "trap 'exit 0' INT TERM; while true; do sleep 1; done") - if err := helpCmd.Start(); err != nil { - t.Fatalf("start helper process: %v", err) + if _, err := os.Stat(pidFile); !os.IsNotExist(err) { + t.Fatalf("expected pid file removed, stat err=%v", err) } - t.Cleanup(func() { - if helpCmd.Process != nil { - _ = helpCmd.Process.Kill() - _, _ = helpCmd.Process.Wait() - } - }) +} - pidFile := filepath.Join(pidDir, "onwatch.pid") - if err := os.WriteFile(pidFile, []byte(strconv.Itoa(helpCmd.Process.Pid)), 0o644); err != nil { +// A PID file left behind by a crashed onWatch can name a PID that the OS has +// since reused for an unrelated process. stoponWatch must treat it as stale: +// never signal or kill it, and still remove the file. +func TestStopOnWatch_StalePIDFileDoesNotKillOtherProcess(t *testing.T) { + pidFile := isolatePIDFile(t) + helper, exited := startIdleHelper(t) + if err := os.WriteFile(pidFile, []byte(strconv.Itoa(helper.Process.Pid)+":65522"), 0o644); err != nil { t.Fatalf("write pid file: %v", err) } - stoponWatch(65523) + stoponWatch(65522) - if isProcessRunning(helpCmd.Process.Pid) { - t.Fatalf("expected helper process %d to be stopped", helpCmd.Process.Pid) + select { + case <-exited: + t.Fatalf("stoponWatch stopped non-onwatch process %d named in a stale PID file", helper.Process.Pid) + case <-time.After(1 * time.Second): + } + if !isProcessRunning(helper.Process.Pid) { + t.Fatalf("expected non-onwatch process %d to keep running", helper.Process.Pid) + } + if _, err := os.Stat(pidFile); !os.IsNotExist(err) { + t.Fatalf("expected stale pid file removed, stat err=%v", err) } } // runMainHelper re-executes this test binary as a child running main(). // // The child is deliberately isolated: it runs in an empty directory with an -// empty PATH, so startonWatch's binary search (./onwatch, ../onwatch, +// empty PATH and a temp home, so startonWatch's binary search (./onwatch, ../onwatch, // ../../onwatch, then PATH) genuinely finds nothing. Without that isolation the // child locates the repo-root binary built by `app.sh --build` (or an installed // onwatch on PATH), starts a real daemon instead of failing, and then the @@ -560,7 +659,11 @@ func runMainHelper(t *testing.T, envVar string) ([]byte, error) { cmd := exec.CommandContext(ctx, os.Args[0], "-test.run="+t.Name()) cmd.Dir = t.TempDir() - cmd.Env = append(os.Environ(), envVar+"=1", "PATH="+t.TempDir()) + // A temp home keeps the child away from the real onWatch PID file, which + // --restart would otherwise use to stop the developer's running daemon. + home := t.TempDir() + cmd.Env = append(os.Environ(), envVar+"=1", "PATH="+t.TempDir(), + "HOME="+home, "USERPROFILE="+home, "LOCALAPPDATA="+filepath.Join(home, "AppData", "Local")) output, err := cmd.CombinedOutput() if ctx.Err() != nil { diff --git a/tools/perf-monitor/process_unix.go b/tools/perf-monitor/process_unix.go new file mode 100644 index 00000000..65146989 --- /dev/null +++ b/tools/perf-monitor/process_unix.go @@ -0,0 +1,81 @@ +//go:build !windows + +package main + +import ( + "os" + "os/exec" + "path/filepath" + "runtime" + "strconv" + "strings" + "syscall" +) + +// onwatchBinaryName is the file name of the onWatch executable. +const onwatchBinaryName = "onwatch" + +// pidFilePath mirrors onWatch's own PID file location on Unix. +func pidFilePath() string { + return filepath.Join(os.Getenv("HOME"), ".onwatch", "onwatch.pid") +} + +// isProcessRunning reports whether pid names a live process. Signal 0 checks +// existence without delivering anything. A nil os.Signal is not a signal-0 +// probe - os.Process.Signal rejects it - so it must be syscall.Signal(0). +func isProcessRunning(pid int) bool { + if pid <= 0 { + return false + } + proc, err := os.FindProcess(pid) + if err != nil { + return false + } + return proc.Signal(syscall.Signal(0)) == nil +} + +// stopProcess asks the process to shut down gracefully. +func stopProcess(proc *os.Process) error { + return proc.Signal(os.Interrupt) +} + +// processCommandName returns the executable base name of pid ("" when +// unknown). macOS ps prints the full path for comm; only the base name may +// count, or any binary under a directory named onwatch would match. +func processCommandName(pid int) string { + if pid <= 0 { + return "" + } + if name, ok := procExeName(pid); ok { + return name + } + out, err := exec.Command("ps", "-p", strconv.Itoa(pid), "-o", "comm=").Output() + if err != nil { + return "" + } + name := strings.TrimSpace(string(out)) + if name == "" { + return "" + } + return filepath.Base(name) +} + +// procExeName reads the process name from /proc on Linux, where ps may be +// missing (Nix build sandbox, distroless image) or busybox's, which has no +// -p. ok is false when /proc has no entry for pid. +func procExeName(pid int) (name string, ok bool) { + if runtime.GOOS != "linux" { + return "", false + } + dir := "/proc/" + strconv.Itoa(pid) + // comm is the name the process was started as (what ps -o comm= shows), + // so a binary launched through a symlink named onwatch still matches. + // The kernel truncates it to 15 characters, which "onwatch" fits. + if comm, err := os.ReadFile(dir + "/comm"); err == nil { + return strings.TrimSpace(string(comm)), true + } + if exe, err := os.Readlink(dir + "/exe"); err == nil { + return filepath.Base(strings.TrimSuffix(exe, " (deleted)")), true + } + return "", false +} diff --git a/tools/perf-monitor/process_windows.go b/tools/perf-monitor/process_windows.go new file mode 100644 index 00000000..c23a1e88 --- /dev/null +++ b/tools/perf-monitor/process_windows.go @@ -0,0 +1,84 @@ +//go:build windows + +package main + +import ( + "os" + "path/filepath" + "syscall" + "unsafe" +) + +// onwatchBinaryName is the file name of the onWatch executable. +const onwatchBinaryName = "onwatch.exe" + +// waitTimeout is WAIT_TIMEOUT: the process object is not signalled, so the +// process is still running. +const waitTimeout = uint32(0x00000102) + +// processQueryLimitedInformation is PROCESS_QUERY_LIMITED_INFORMATION, the +// least access right that allows reading a process's image path. +const processQueryLimitedInformation = 0x1000 + +var procQueryFullProcessImageNameW = syscall.NewLazyDLL("kernel32.dll").NewProc("QueryFullProcessImageNameW") + +// pidFilePath mirrors onWatch's own PID file location on Windows. +func pidFilePath() string { + dir := os.Getenv("LOCALAPPDATA") + if dir != "" { + return filepath.Join(dir, "onwatch", "onwatch.pid") + } + return filepath.Join(os.Getenv("USERPROFILE"), ".onwatch", "onwatch.pid") +} + +// isProcessRunning reports whether pid names a running process. Windows has +// no signal-0 probe, and an exited process keeps an openable handle while +// anything holds one, so ask whether the process object has been signalled +// (which happens exactly when the process ends). +func isProcessRunning(pid int) bool { + if pid <= 0 { + return false + } + handle, err := syscall.OpenProcess(syscall.SYNCHRONIZE, false, uint32(pid)) + if err != nil { + // A process we may not synchronize on still exists. + return err == syscall.ERROR_ACCESS_DENIED + } + defer syscall.CloseHandle(handle) + state, err := syscall.WaitForSingleObject(handle, 0) + if err != nil { + return false + } + return state == waitTimeout +} + +// stopProcess terminates the process. Windows cannot deliver os.Interrupt to +// another process (os.Process.Signal only supports Kill there). +func stopProcess(proc *os.Process) error { + return proc.Kill() +} + +// processCommandName returns the image file base name of pid, e.g. +// onwatch.exe ("" when unknown). Windows has no ps, so ask the process object. +func processCommandName(pid int) string { + if pid <= 0 { + return "" + } + handle, err := syscall.OpenProcess(processQueryLimitedInformation, false, uint32(pid)) + if err != nil { + return "" + } + defer syscall.CloseHandle(handle) + buf := make([]uint16, 1024) + size := uint32(len(buf)) + r, _, _ := procQueryFullProcessImageNameW.Call( + uintptr(handle), + 0, + uintptr(unsafe.Pointer(&buf[0])), + uintptr(unsafe.Pointer(&size)), + ) + if r == 0 { + return "" + } + return filepath.Base(syscall.UTF16ToString(buf[:size])) +}