Merge origin/main into feat/plugin-catalog
Python plugin CLI/loader/web/tui files taken from main wholesale; the catalog layer is re-ported onto main's decomposed shapes in the following commits. plugin_index.py removed (catalog is the sole discovery system).
This commit is contained in:
+3
-3
@@ -107,9 +107,9 @@ plans/
|
||||
.hadolint.yaml
|
||||
.mailmap
|
||||
|
||||
# Repo-root debug/export artifacts — must never reach image layers (COPY . .)
|
||||
# Debug/export artifacts — must never reach image layers (COPY . .)
|
||||
/log.txt
|
||||
/sqlite_leak_fix.png
|
||||
/*.png.bak
|
||||
/default.tar.gz
|
||||
/*.tar.gz
|
||||
*.tar.gz
|
||||
*.tgz
|
||||
|
||||
+45
-5
@@ -153,6 +153,35 @@
|
||||
# Optional base URL override:
|
||||
# UPSTAGE_BASE_URL=https://api.upstage.ai/v1
|
||||
|
||||
# =============================================================================
|
||||
# LLM PROVIDER (Ramp Router)
|
||||
# =============================================================================
|
||||
# Ramp Router (router.com) — Responses-native LLM gateway; model IDs come
|
||||
# from your account's live catalog (GET /v1/models).
|
||||
# Get your key at: https://app.router.com/keys
|
||||
# RAMP_ROUTER_API_KEY=your_key_here
|
||||
# Optional base URL override:
|
||||
# RAMP_ROUTER_BASE_URL=https://api.router.com/v1
|
||||
|
||||
# =============================================================================
|
||||
# LLM PROVIDER (Nebius Token Factory)
|
||||
# =============================================================================
|
||||
# Nebius Token Factory — OpenAI-compatible inference for open models.
|
||||
# Get your key at: https://tokenfactory.nebius.com/
|
||||
# NEBIUS_API_KEY=your_key_here
|
||||
# Optional base URL override:
|
||||
# NEBIUS_BASE_URL=https://api.tokenfactory.nebius.com/v1
|
||||
|
||||
# =============================================================================
|
||||
# LLM PROVIDER (Tencent Hy — TokenHub & TokenPlan)
|
||||
# =============================================================================
|
||||
# Tencent TokenHub (OpenAI-compatible): https://tokenhub.tencentmaas.com
|
||||
# TOKENHUB_API_KEY=your_key_here
|
||||
# TOKENHUB_BASE_URL=https://tokenhub.tencentmaas.com/v1
|
||||
# Tencent TokenPlan (Anthropic Messages endpoint via LKEAP):
|
||||
# TOKENPLAN_API_KEY=your_key_here
|
||||
# TOKENPLAN_BASE_URL=https://api.lkeap.cloud.tencent.com/plan/anthropic
|
||||
|
||||
# =============================================================================
|
||||
# TOOL API KEYS
|
||||
# =============================================================================
|
||||
@@ -293,12 +322,14 @@ BROWSERBASE_PROXIES=true
|
||||
# Uses custom Chromium build to avoid bot detection altogether
|
||||
BROWSERBASE_ADVANCED_STEALTH=false
|
||||
|
||||
# Browser engine for local mode (default: auto = Chrome)
|
||||
# "auto" — use Chrome (don't pass --engine flag)
|
||||
# "lightpanda" — use Lightpanda (1.3-5.8x faster navigation, no screenshots)
|
||||
# Local browser engine (default: auto = Chrome)
|
||||
# "auto" — use Chrome
|
||||
# "lightpanda" — use Lightpanda (faster navigation, no screenshots)
|
||||
# "chrome" — explicitly request Chrome
|
||||
# Requires agent-browser v0.25.3+. Lightpanda commands that fail or return
|
||||
# empty results are automatically retried with Chrome.
|
||||
# Browser Use mode (default) spawns `lightpanda serve` itself; the built-in
|
||||
# browser tools pass --engine to agent-browser v0.25.3+ and retry failed or
|
||||
# empty Lightpanda results with Chrome. Ignored while a cloud provider,
|
||||
# Camofox or browser.cdp_url is active (`hermes doctor` reports that).
|
||||
# Also configurable via browser.engine in config.yaml.
|
||||
# AGENT_BROWSER_ENGINE=auto
|
||||
|
||||
@@ -502,3 +533,12 @@ IMAGE_TOOLS_DEBUG=false
|
||||
# GOOGLE_CHAT_ALLOW_ALL_USERS=false # Set true to skip the allowlist
|
||||
# GOOGLE_CHAT_HOME_CHANNEL= # Default space (spaces/XXXX) for cron delivery
|
||||
# GOOGLE_CHAT_HOME_CHANNEL_NAME= # Display name for the home channel
|
||||
|
||||
# =============================================================================
|
||||
# reddit-reading skill (optional) — app-only credentials, NOT a user login
|
||||
# =============================================================================
|
||||
# The skill works with no credentials via Reddit's public feeds (~1 request/minute).
|
||||
# For faster access with scores and nested comments, register a free "script" app
|
||||
# at https://www.reddit.com/prefs/apps and paste its id and secret here.
|
||||
# REDDIT_CLIENT_ID=
|
||||
# REDDIT_CLIENT_SECRET=
|
||||
|
||||
@@ -48,6 +48,9 @@ outputs:
|
||||
installer:
|
||||
description: Run the PowerShell installer tests on a Windows runner.
|
||||
value: ${{ steps.classify.outputs.installer }}
|
||||
desktop_updater:
|
||||
description: Run the Windows desktop-update hand-off (windows.ps1) integration tests.
|
||||
value: ${{ steps.classify.outputs.desktop_updater }}
|
||||
rust:
|
||||
description: Run `cargo test` for the Tauri bootstrap installer.
|
||||
value: ${{ steps.classify.outputs.rust }}
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
name: E2E screen recording
|
||||
description: >
|
||||
One screen-recording mechanism for every install-e2e runner. mode=start
|
||||
installs what the OS needs (ffmpeg everywhere; Xvfb on headless linux),
|
||||
starts the recorder, and exports DISPLAY for later steps. mode=stop
|
||||
finalizes the recording and fails on a zero-frame file. A missing ffmpeg
|
||||
is a hard error - a graceful skip makes the missing tool invisible and
|
||||
the artifact silently loses its recording.
|
||||
|
||||
inputs:
|
||||
mode:
|
||||
description: start | stop
|
||||
required: true
|
||||
output:
|
||||
description: Path of the recording (mkv).
|
||||
required: true
|
||||
|
||||
runs:
|
||||
using: composite
|
||||
steps:
|
||||
# ---- setup + start ------------------------------------------------------
|
||||
- name: Install ffmpeg + Xvfb (linux)
|
||||
if: inputs.mode == 'start' && runner.os == 'Linux'
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
if ! command -v ffmpeg >/dev/null 2>&1 || ! command -v Xvfb >/dev/null 2>&1; then
|
||||
sudo apt-get update -qq
|
||||
sudo apt-get install -y -qq --no-install-recommends ffmpeg xvfb
|
||||
fi
|
||||
|
||||
- name: Start Xvfb (linux, headless)
|
||||
if: inputs.mode == 'start' && runner.os == 'Linux'
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
# A dedicated display rather than xvfb-run-wrapping each command, so
|
||||
# ONE display serves both the app under test and the recorder.
|
||||
if [ -z "${DISPLAY:-}" ]; then
|
||||
Xvfb :99 -screen 0 1920x1080x24 &
|
||||
echo "$!" > "$RUNNER_TEMP/xvfb.pid"
|
||||
echo "DISPLAY=:99" >> "$GITHUB_ENV"
|
||||
export DISPLAY=:99
|
||||
fi
|
||||
# Wait until the display accepts connections; xdpyinfo may not be
|
||||
# installed, so probe with the X socket.
|
||||
for _ in $(seq 1 50); do
|
||||
[ -S "/tmp/.X11-unix/X99" ] && break
|
||||
sleep 0.2
|
||||
done
|
||||
[ -S "/tmp/.X11-unix/X99" ] || { echo "Xvfb :99 did not come up" >&2; exit 1; }
|
||||
|
||||
- name: Verify ffmpeg (macos)
|
||||
if: inputs.mode == 'start' && runner.os == 'macOS'
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
# Do not trust "preinstalled" claims - verify, install on miss.
|
||||
command -v ffmpeg >/dev/null 2>&1 || brew install --quiet ffmpeg
|
||||
|
||||
# until new runner image is published by github, we have to hack on screen record approvals
|
||||
# see https://github.com/actions/runner-images/issues/14474 - as of this hermes agent commit,
|
||||
# it's merged but the image isn't updated.
|
||||
approvalsPlist="$HOME/Library/Group Containers/group.com.apple.replayd/ScreenCaptureApprovals.plist"
|
||||
mkdir -p "$(dirname "$approvalsPlist")"
|
||||
defaults write "$approvalsPlist" "/opt/hca/hosted-compute-agent" -date "3024-01-01 00:00:00 +0000"
|
||||
killall cfprefsd 2>/dev/null || true
|
||||
|
||||
- name: Restore cached ffmpeg (windows)
|
||||
if: inputs.mode == 'start' && runner.os == 'Windows'
|
||||
id: ffmpeg-cache
|
||||
uses: actions/cache@27d5ce7f107fe9357f9df03efb73ab90386fccae # v5
|
||||
with:
|
||||
path: ${{ runner.temp }}\test-bins\ffmpeg
|
||||
key: e2e-ffmpeg-${{ runner.os }}-v1
|
||||
|
||||
- name: Install ffmpeg (windows)
|
||||
if: inputs.mode == 'start' && runner.os == 'Windows' && steps.ffmpeg-cache.outputs.cache-hit != 'true'
|
||||
shell: pwsh
|
||||
run: |
|
||||
$bins = "$env:RUNNER_TEMP\test-bins\ffmpeg"
|
||||
New-Item -ItemType Directory -Path $bins -Force | Out-Null
|
||||
winget install -e --id Gyan.FFmpeg --silent --accept-source-agreements --accept-package-agreements --disable-interactivity --location "$env:RUNNER_TEMP\ffmpeg_dir"
|
||||
Copy-Item -Path "$env:RUNNER_TEMP\ffmpeg_dir\*\*" -Destination $bins -Recurse -Force
|
||||
|
||||
- name: Add ffmpeg to PATH (windows)
|
||||
if: inputs.mode == 'start' && runner.os == 'Windows'
|
||||
shell: pwsh
|
||||
run: Add-Content -Path $env:GITHUB_PATH -Value "$env:RUNNER_TEMP\test-bins\ffmpeg\bin"
|
||||
|
||||
- name: Start recording (posix)
|
||||
if: inputs.mode == 'start' && runner.os != 'Windows'
|
||||
shell: bash
|
||||
run: bash "$GITHUB_ACTION_PATH/../../../tests/install/e2e-assets/record-start.sh" '${{ inputs.output }}'
|
||||
|
||||
- name: Start recording (windows)
|
||||
if: inputs.mode == 'start' && runner.os == 'Windows'
|
||||
shell: powershell
|
||||
run: powershell -NoProfile -ExecutionPolicy Bypass -File "$env:GITHUB_ACTION_PATH\..\..\..\tests\install\e2e-assets\record-start.ps1" -OutFile "${{ inputs.output }}"
|
||||
|
||||
# ---- stop ---------------------------------------------------------------
|
||||
- name: Stop recording (posix)
|
||||
if: inputs.mode == 'stop' && runner.os != 'Windows'
|
||||
shell: bash
|
||||
run: bash "$GITHUB_ACTION_PATH/../../../tests/install/e2e-assets/record-stop.sh" '${{ inputs.output }}'
|
||||
|
||||
- name: Stop recording (windows)
|
||||
if: inputs.mode == 'stop' && runner.os == 'Windows'
|
||||
shell: powershell
|
||||
run: powershell -NoProfile -ExecutionPolicy Bypass -File "$env:GITHUB_ACTION_PATH\..\..\..\tests\install\e2e-assets\record-stop.ps1" -OutFile "${{ inputs.output }}"
|
||||
@@ -0,0 +1,33 @@
|
||||
name: Case Collision Check
|
||||
|
||||
# Rejects PRs that track two files whose paths differ only by case
|
||||
# (README.md vs readme.md, src/Foo.py vs SRC/foo.py).
|
||||
#
|
||||
# Linux is case-sensitive; Windows and macOS (default) are not. A
|
||||
# case-colliding pair lives fine in a Linux checkout and silently breaks
|
||||
# every clone on a case-insensitive host — the filesystem can hold only
|
||||
# one of them, so checkout fails or whichever wins overwrites the other.
|
||||
# Git won't prevent the pair from landing (it only warns at checkout time,
|
||||
# on a case-insensitive FS, for the client doing the checkout), so the only
|
||||
# enforcement point is CI, on Linux, against the index.
|
||||
#
|
||||
# Runs unconditionally (no change-classifier gate): a collision can ship in
|
||||
# any kind of PR — docs, JS, config, not just Python — so gating on a
|
||||
# language lane would be the same "passive rule that cannot enforce a
|
||||
# policy" trap the infographic check exists to close.
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
check-case-collisions:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
|
||||
- name: Run case-collision checker
|
||||
run: python3 scripts/check-case-collisions.py
|
||||
@@ -49,6 +49,7 @@ jobs:
|
||||
uv_lock: ${{ steps.classify.outputs.uv_lock }}
|
||||
npm_lock: ${{ steps.classify.outputs.npm_lock }}
|
||||
installer: ${{ steps.classify.outputs.installer }}
|
||||
desktop_updater: ${{ steps.classify.outputs.desktop_updater }}
|
||||
rust: ${{ steps.classify.outputs.rust }}
|
||||
docker_meta: ${{ steps.classify.outputs.docker_meta }}
|
||||
mcp_catalog: ${{ steps.classify.outputs.mcp_catalog }}
|
||||
@@ -84,6 +85,11 @@ jobs:
|
||||
needs: detect
|
||||
if: needs.detect.outputs.python == 'true'
|
||||
uses: ./.github/workflows/tests-os.yml
|
||||
with:
|
||||
# The Windows lane spawns the real desktop-update hand-off script
|
||||
# (tests/test_desktop_update_windows_*.py) only when that surface
|
||||
# changed; unit-level windows_only tests always run.
|
||||
desktop_updater: ${{ needs.detect.outputs.desktop_updater == 'true' }}
|
||||
|
||||
lint:
|
||||
name: Python lints
|
||||
@@ -122,14 +128,14 @@ jobs:
|
||||
# Tests-only PRs (~17% of commits) skip this 5-minute job — the longest
|
||||
# single job in the workflow — while still running the full pytest lanes.
|
||||
#
|
||||
# ⛔ TEMPORARILY DISABLED (Aug 2, 2026, Teknium) — the suite is red on
|
||||
# every PR and on main itself since the Aug 1 night engines/npm churn
|
||||
# (#76499 → #76562 → #76575): the mock-backend Electron window never
|
||||
# gets a title, so boot/chat/setup/interim specs all fail identically
|
||||
# regardless of the PR's diff (verified on #76573 and the docs-only
|
||||
# #76582). Tracking issue: #76627 (assigned: Ari). To re-enable,
|
||||
# delete the `false &&` below — nothing else changed.
|
||||
if: ${{ false && (needs.detect.outputs.python_prod == 'true' || needs.detect.outputs.frontend == 'true') }}
|
||||
# Re-disabled (Sep 2026): the Sep 1 re-enable is still incredibly flaky.
|
||||
# Keep this a bare `if: false`. The earlier
|
||||
# `${{ false && (... || ...) }}` form on this reusable-workflow job made
|
||||
# GitHub's workflow parser fail at startup ("An unexpected error has
|
||||
# occurred") — every ci.yaml run repo-wide dispatched 0 jobs from
|
||||
# 24f5a60ed1 until this line changed. To re-enable, restore:
|
||||
# if: ${{ needs.detect.outputs.python_prod == 'true' || needs.detect.outputs.frontend == 'true' }}
|
||||
if: false
|
||||
uses: ./.github/workflows/e2e-desktop.yml
|
||||
|
||||
docs-site:
|
||||
@@ -166,6 +172,16 @@ jobs:
|
||||
needs: detect
|
||||
uses: ./.github/workflows/infographic-check.yml
|
||||
|
||||
profile-artifact-check:
|
||||
name: Profile artifact check
|
||||
needs: detect
|
||||
uses: ./.github/workflows/profile-artifact-check.yml
|
||||
|
||||
case-collision-check:
|
||||
name: Check no case-colliding filenames
|
||||
needs: detect
|
||||
uses: ./.github/workflows/case-collision-check.yml
|
||||
|
||||
lockfile-diff:
|
||||
name: package-lock.json diff
|
||||
needs: detect
|
||||
@@ -228,8 +244,10 @@ jobs:
|
||||
- history-check
|
||||
- contributor-check
|
||||
- uv-lockfile
|
||||
- case-collision-check
|
||||
- lockfile-diff
|
||||
- docker-lint
|
||||
- profile-artifact-check
|
||||
- supply-chain
|
||||
- review-labels
|
||||
- osv-scanner
|
||||
|
||||
@@ -94,6 +94,9 @@ jobs:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
|
||||
- name: Reject profile exports in the build context
|
||||
run: python3 scripts/ci/check_profile_archive_boundary.py
|
||||
|
||||
# Retry once on transient Docker Hub / buildkit pull failures
|
||||
# (connection reset, auth token timeout, rate limiting). The action
|
||||
# generates a unique builder name per invocation so the retry doesn't
|
||||
@@ -206,6 +209,9 @@ jobs:
|
||||
- name: Checkout trusted source
|
||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
|
||||
- name: Reject profile exports in the build context
|
||||
run: python3 scripts/ci/check_profile_archive_boundary.py
|
||||
|
||||
# Retry once on transient Docker Hub / buildkit pull failures.
|
||||
# See build job for rationale; same pattern.
|
||||
- name: Set up Docker Buildx
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
# Reusable runner for ONE macOS install/update combination.
|
||||
#
|
||||
# Two driver arms, one runs per dispatch (the other natively skips):
|
||||
#
|
||||
# e2e (script arms) tests/install/installer-script-e2e.sh - the
|
||||
# OS-agnostic git-redirect driver shared with
|
||||
# linux. installer-script(+desktop) installs,
|
||||
# script/updater/hermes-desktop-app-update
|
||||
# updates.
|
||||
# gui-e2e (desktop arm) tests/install/macos-desktop-e2e.sh - the
|
||||
# published Hermes-Setup.dmg, mounted and run,
|
||||
# then the app driven by Playwright for the
|
||||
# app-update methods.
|
||||
#
|
||||
# Method pairs without a driver arm yet NATIVELY SKIP (grey check, no
|
||||
# runner): the capability knowledge lives here, next to the drivers.
|
||||
|
||||
name: install-e2e macos leg
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
install-method:
|
||||
description: 'How OLD gets installed. Supported: installer-script, installer-script+desktop (curl | bash one-liner, optionally with --include-desktop) and desktop-installer@latest (the published Hermes-Setup.dmg).'
|
||||
required: true
|
||||
type: string
|
||||
update-method:
|
||||
description: 'How the install updates to HEAD. Script installs support hermes-update / installer-script / installer-script+desktop / hermes-desktop-app-update; dmg installs support open-app-update / hermes-desktop-app-update.'
|
||||
required: true
|
||||
type: string
|
||||
install-ref:
|
||||
description: 'What to install before updating: a branch, a tag, or a SHA reachable from main.'
|
||||
required: false
|
||||
type: string
|
||||
default: refs/heads/main
|
||||
tag-has-desktop:
|
||||
description: "Whether install-ref ships the desktop app (apps/desktop). The caller annotates this from the tag's own tree; desktop-method legs from pre-desktop releases natively skip."
|
||||
required: false
|
||||
type: boolean
|
||||
default: true
|
||||
leg-id:
|
||||
description: 'Artifact-safe matrix leg id (from generate-e2e-matrix.mjs legId). Names this leg''s logs + player artifacts so the report job can link a row to its zip.'
|
||||
required: true
|
||||
type: string
|
||||
dmg-url:
|
||||
description: 'Bootstrap dmg to install OLD with. Default: the latest published one — what a user downloads today.'
|
||||
required: false
|
||||
type: string
|
||||
default: https://hermes-assets.nousresearch.com/Hermes-Setup.dmg
|
||||
timeout-minutes:
|
||||
description: 'Job timeout. App-update legs do a full Electron build.'
|
||||
required: false
|
||||
type: number
|
||||
default: 60
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
# ---- arm 1: script installs (the shared OS-agnostic driver) --------------
|
||||
e2e:
|
||||
name: install & update
|
||||
if: >-
|
||||
(inputs.install-method == 'installer-script'
|
||||
|| (inputs.install-method == 'installer-script+desktop' && inputs.tag-has-desktop))
|
||||
&& (contains(fromJSON('["hermes-update", "installer-script"]'), inputs.update-method)
|
||||
|| (contains(fromJSON('["installer-script+desktop", "hermes-desktop-app-update"]'), inputs.update-method) && inputs.tag-has-desktop))
|
||||
uses: ./.github/workflows/install-e2e-run.yml
|
||||
with:
|
||||
install-method: ${{ inputs.install-method }}
|
||||
update-method: ${{ inputs.update-method }}
|
||||
install-ref: ${{ inputs.install-ref }}
|
||||
tag-has-desktop: ${{ inputs.tag-has-desktop }}
|
||||
leg-id: ${{ inputs.leg-id }}
|
||||
runner: macos-latest
|
||||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
|
||||
# ---- arm 2: the published dmg, then Playwright drives the app ------------
|
||||
gui-e2e:
|
||||
# Short static name on purpose: name expressions render UNEXPANDED on
|
||||
# skipped jobs.
|
||||
name: Hermes-Setup.dmg
|
||||
if: >-
|
||||
inputs.install-method == 'desktop-installer@latest' && inputs.tag-has-desktop
|
||||
&& contains(fromJSON('["open-app-update", "hermes-desktop-app-update", "hermes-update", "installer-script", "installer-script+desktop"]'), inputs.update-method)
|
||||
runs-on: macos-latest
|
||||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
|
||||
steps:
|
||||
# Full history: the driver bare-clones this checkout as the repo the
|
||||
# installer/updater talk to.
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Start screen recording
|
||||
uses: ./.github/actions/e2e-screen-record
|
||||
with:
|
||||
mode: start
|
||||
output: ${{ runner.temp }}/e2e-logs/recording.mkv
|
||||
|
||||
- name: Stage serve repo (main -> ${{ inputs.install-ref }})
|
||||
run: |
|
||||
set -euo pipefail
|
||||
tests/install/macos-desktop-e2e.sh --phase stage \
|
||||
--update-method '${{ inputs.update-method }}' \
|
||||
--install-ref '${{ inputs.install-ref }}' \
|
||||
--dmg-url '${{ inputs.dmg-url }}'
|
||||
env:
|
||||
HERMES_E2E_LOG_DIR: ${{ runner.temp }}/e2e-logs
|
||||
|
||||
- name: Install ${{ inputs.install-ref }} via Hermes-Setup.dmg
|
||||
run: |
|
||||
set -euo pipefail
|
||||
tests/install/macos-desktop-e2e.sh --phase install \
|
||||
--update-method '${{ inputs.update-method }}' \
|
||||
--install-ref '${{ inputs.install-ref }}' \
|
||||
--dmg-url '${{ inputs.dmg-url }}'
|
||||
env:
|
||||
HERMES_E2E_LOG_DIR: ${{ runner.temp }}/e2e-logs
|
||||
|
||||
- name: Update ${{ inputs.install-ref }} -> HEAD (${{ inputs.update-method }})
|
||||
run: |
|
||||
set -euo pipefail
|
||||
tests/install/macos-desktop-e2e.sh --phase update \
|
||||
--update-method '${{ inputs.update-method }}' \
|
||||
--install-ref '${{ inputs.install-ref }}' \
|
||||
--dmg-url '${{ inputs.dmg-url }}'
|
||||
env:
|
||||
HERMES_E2E_LOG_DIR: ${{ runner.temp }}/e2e-logs
|
||||
|
||||
- name: Stop screen recording
|
||||
if: always()
|
||||
uses: ./.github/actions/e2e-screen-record
|
||||
with:
|
||||
mode: stop
|
||||
output: ${{ runner.temp }}/e2e-logs/recording.mkv
|
||||
|
||||
- name: Remux recording for browser playback
|
||||
if: always()
|
||||
run: |
|
||||
set -euo pipefail
|
||||
if [ -f "${{ runner.temp }}/e2e-logs/recording.mkv" ]; then
|
||||
ffmpeg -y -hide_banner -loglevel error -i "${{ runner.temp }}/e2e-logs/recording.mkv" \
|
||||
-c copy "${{ runner.temp }}/e2e-logs/recording.mp4"
|
||||
fi
|
||||
|
||||
# Artifact names cannot contain '/'; install-ref may be a full ref.
|
||||
- name: Build artifact name
|
||||
id: artifact
|
||||
if: always()
|
||||
run: |
|
||||
echo "name=install-e2e-logs-${{ inputs.leg-id }}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Upload logs
|
||||
if: always()
|
||||
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
|
||||
with:
|
||||
name: ${{ steps.artifact.outputs.name }}
|
||||
path: ${{ runner.temp }}/e2e-logs
|
||||
retention-days: 14
|
||||
if-no-files-found: ignore
|
||||
@@ -1,12 +1,28 @@
|
||||
name: Install & Update E2E (reusable)
|
||||
|
||||
# Runs ONE update route against ONE starting commit, in the dev sandbox, with a
|
||||
# real install (uv, a managed Python, Node, the venv) behind it.
|
||||
# Runs ONE {install-method, update-method} combination against ONE starting
|
||||
# commit, with a real install (uv, a managed Python, Node, the venv) behind
|
||||
# it.
|
||||
#
|
||||
# Reusable so callers can fan out over the combinations that matter -- update
|
||||
# from the tip vs. from an older release, `hermes update` vs. re-running the
|
||||
# installer -- without duplicating the runner setup. Each leg is independent:
|
||||
# its own sandbox, its own install, nothing rewound or shared.
|
||||
# Reusable so callers can fan out over the combinations that matter --
|
||||
# update from the tip vs. from an older release, `hermes update` vs.
|
||||
# re-running the installer -- without duplicating the runner setup. Each leg
|
||||
# is independent: its own isolated HOME, its own install, nothing rewound
|
||||
# or shared.
|
||||
#
|
||||
# No sandbox: tests/install/installer-script-e2e.sh points every git
|
||||
# process at a local bare clone (url.<file://serve.git>.insteadOf in a
|
||||
# driver-owned GIT_CONFIG_GLOBAL) and isolates HOME, so the installer and
|
||||
# updater run byte-for-byte against their real URLs on the bare runner --
|
||||
# which is disposable, and therefore IS the sandbox. That also makes this
|
||||
# workflow OS-agnostic: the same driver runs on ubuntu and macos runners.
|
||||
#
|
||||
# Method ids come from scripts/sandbox/generate-e2e-matrix.mjs. Supported
|
||||
# today: install via installer-script, update via hermes-update or
|
||||
# installer-script (re-run the one-liner).
|
||||
# Anything else NATIVELY SKIPS (grey check, no runner): capability
|
||||
# knowledge lives here, next to the driver, so the caller can dispatch
|
||||
# every declared combination without knowing which ones work.
|
||||
#
|
||||
# Call it:
|
||||
#
|
||||
@@ -14,14 +30,19 @@ name: Install & Update E2E (reusable)
|
||||
# tip:
|
||||
# uses: ./.github/workflows/install-e2e-run.yml
|
||||
# with:
|
||||
# route: update
|
||||
# install-method: installer-script
|
||||
# update-method: hermes-update
|
||||
# install-ref: refs/heads/main
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
route:
|
||||
description: 'Update path to exercise: update (hermes update) or installer (re-run install.sh).'
|
||||
install-method:
|
||||
description: 'How the starting version gets installed. Supported: installer-script (the real curl | install.sh one-liner) and installer-script+desktop (the same one-liner with --include-desktop).'
|
||||
required: true
|
||||
type: string
|
||||
update-method:
|
||||
description: 'How the install updates to HEAD. Supported: hermes-update (the updater), installer-script (re-run the one-liner), installer-script+desktop (re-run with --include-desktop), hermes-desktop-app-update (launch via hermes desktop under Playwright, click Update now). open-app-update runs only where an OS entry point exists (see the per-OS run workflows); pairs without one skip.'
|
||||
required: true
|
||||
type: string
|
||||
install-ref:
|
||||
@@ -29,84 +50,100 @@ on:
|
||||
required: false
|
||||
type: string
|
||||
default: refs/heads/main
|
||||
leg-id:
|
||||
description: 'Artifact-safe matrix leg id (from generate-e2e-matrix.mjs legId). Names this leg''s logs + player artifacts so the report job can link a row to its zip.'
|
||||
required: true
|
||||
type: string
|
||||
tag-has-desktop:
|
||||
description: "Whether install-ref ships the desktop app (apps/desktop). The caller annotates this from the tag's own tree; desktop-method legs from pre-desktop releases natively skip."
|
||||
required: false
|
||||
type: boolean
|
||||
default: true
|
||||
runner:
|
||||
description: 'Runner label.'
|
||||
required: false
|
||||
type: string
|
||||
default: ubuntu-latest
|
||||
timeout-minutes:
|
||||
description: 'Job timeout. A cold run installs real toolchains twice.'
|
||||
description: 'Job timeout. A cold run installs real toolchains twice, and app-update legs add a full Electron build + launch.'
|
||||
required: false
|
||||
type: number
|
||||
default: 45
|
||||
default: 60
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
e2e:
|
||||
name: ${{ inputs.route }} from ${{ inputs.install-ref }}
|
||||
name: install & update
|
||||
# The pairs the driver can run today; anything else natively skips.
|
||||
# Desktop-surface methods (+desktop installs,
|
||||
# hermes-desktop-app-update) also need the starting tag to ship
|
||||
# apps/desktop (their flags shipped with it).
|
||||
if: >-
|
||||
(inputs.install-method == 'installer-script'
|
||||
|| (inputs.install-method == 'installer-script+desktop' && inputs.tag-has-desktop))
|
||||
&& (contains(fromJSON('["hermes-update", "installer-script"]'), inputs.update-method)
|
||||
|| (contains(fromJSON('["installer-script+desktop", "hermes-desktop-app-update"]'), inputs.update-method) && inputs.tag-has-desktop))
|
||||
runs-on: ${{ inputs.runner }}
|
||||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
|
||||
steps:
|
||||
# Full history: the sandbox fetches the starting commit and the test
|
||||
# compares against this commit, so a shallow clone is not enough.
|
||||
# Full history: the driver bare-clones this checkout as the repo the
|
||||
# installer/updater talk to, and both OLD and HEAD must be reachable
|
||||
# in that clone. A shallow clone cannot serve either need.
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
# bubblewrap + slirp4netns are what the sandbox is built on; util-linux
|
||||
# supplies the `unshare` that builds the multi-uid userns for the
|
||||
# user-level (non-root) install.
|
||||
- name: Install sandbox dependencies
|
||||
run: |
|
||||
set -euo pipefail
|
||||
sudo apt-get update -qq
|
||||
sudo apt-get install -y -qq bubblewrap slirp4netns uidmap util-linux
|
||||
|
||||
# Ubuntu 24.04 restricts unprivileged user namespaces through AppArmor,
|
||||
# which is exactly what bwrap needs. Report the state before touching it
|
||||
# so a future runner-image change is visible in the log rather than
|
||||
# silently altering what this job proves.
|
||||
- name: Permit unprivileged user namespaces
|
||||
run: |
|
||||
set -euo pipefail
|
||||
echo "--- kernel userns settings (before)"
|
||||
sysctl kernel.unprivileged_userns_clone 2>/dev/null || echo " (sysctl absent)"
|
||||
sysctl kernel.apparmor_restrict_unprivileged_userns 2>/dev/null || echo " (sysctl absent)"
|
||||
if sysctl -n kernel.apparmor_restrict_unprivileged_userns >/dev/null 2>&1; then
|
||||
sudo sysctl -w kernel.apparmor_restrict_unprivileged_userns=0
|
||||
fi
|
||||
echo "--- subuid/subgid for $(id -un)"
|
||||
grep "^$(id -un):" /etc/subuid /etc/subgid || echo " (none — sandbox will say so)"
|
||||
# One recording mechanism on every OS (Xvfb gives headless linux a
|
||||
# display; the same display serves any app the driver launches).
|
||||
- name: Start screen recording
|
||||
uses: ./.github/actions/e2e-screen-record
|
||||
with:
|
||||
mode: start
|
||||
output: ${{ runner.temp }}/e2e-logs/recording.mkv
|
||||
|
||||
- name: Run install + update E2E
|
||||
run: |
|
||||
set -euo pipefail
|
||||
tests/install/install-update-e2e.sh \
|
||||
--route '${{ inputs.route }}' \
|
||||
tests/install/installer-script-e2e.sh \
|
||||
--install-method '${{ inputs.install-method }}' \
|
||||
--update-method '${{ inputs.update-method }}' \
|
||||
--install-ref '${{ inputs.install-ref }}'
|
||||
env:
|
||||
# Outside the workspace on purpose: the script creates this directory
|
||||
# up front, and an untracked dir inside the repo makes the worktree
|
||||
# dirty -- which dev-sandbox reacts to by snapshotting the working
|
||||
# copy into a fresh fake-main commit on every invocation, moving the
|
||||
# update target mid-run.
|
||||
# Outside the workspace on purpose: logs written into the repo
|
||||
# would trip the driver's own dirty-tree guard.
|
||||
HERMES_E2E_LOG_DIR: ${{ runner.temp }}/e2e-logs
|
||||
|
||||
# Artifact names cannot contain '/', and install-ref may be a full ref
|
||||
# like refs/heads/main. GitHub Actions expressions have no string-replace
|
||||
# function, so build the safe name here. Runs even on failure -- that is
|
||||
# exactly when the logs are wanted.
|
||||
- name: Stop screen recording
|
||||
if: always()
|
||||
uses: ./.github/actions/e2e-screen-record
|
||||
with:
|
||||
mode: stop
|
||||
output: ${{ runner.temp }}/e2e-logs/recording.mkv
|
||||
|
||||
# Browsers cannot play Matroska: remux (copy codec, no re-encode) so
|
||||
# the artifact zip feeds the static playback.html leg player directly.
|
||||
- name: Remux recording for browser playback
|
||||
if: always()
|
||||
run: |
|
||||
set -euo pipefail
|
||||
if [ -f "${{ runner.temp }}/e2e-logs/recording.mkv" ]; then
|
||||
ffmpeg -y -hide_banner -loglevel error -i "${{ runner.temp }}/e2e-logs/recording.mkv" \
|
||||
-c copy "${{ runner.temp }}/e2e-logs/recording.mp4"
|
||||
fi
|
||||
|
||||
# The leg player: ONE static HTML per run, uploaded up front by the
|
||||
# leg-player job in install-e2e.yml (archive: false, so GitHub names
|
||||
# the artifact after the file: playback.html). The report job links
|
||||
# every ran leg to it with the leg's zip as a #zip= hash param.
|
||||
- name: Build artifact name
|
||||
if: always()
|
||||
id: artifact
|
||||
run: |
|
||||
set -euo pipefail
|
||||
safe_ref='${{ inputs.install-ref }}'
|
||||
safe_ref="${safe_ref//\//-}"
|
||||
echo "name=install-e2e-${{ inputs.route }}-${safe_ref}" >> "$GITHUB_OUTPUT"
|
||||
echo "name=install-e2e-logs-${{ inputs.leg-id }}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
# The installer's own transcripts say far more than the assertion that
|
||||
# tripped when a real install breaks.
|
||||
|
||||
@@ -0,0 +1,201 @@
|
||||
# Reusable runner for ONE Windows install/update combination.
|
||||
#
|
||||
# One job, two orthogonal axes: tests/install/windows-e2e.ps1 dispatches its
|
||||
# install phase on install-method and its update phase on update-method, so
|
||||
# implementing a new pair is a driver function + a gate edit here - never a
|
||||
# new job. The driver's phases share state via the workroot, and every leg
|
||||
# runs the REAL user surface for its methods:
|
||||
#
|
||||
# desktop-installer@latest the website's Hermes-Setup.exe, downloaded and
|
||||
# run headed, AutoHotkey clicks Install ->
|
||||
# Launch, the real Electron window must appear.
|
||||
# installer-script the irm | iex one-liner: the install.ps1
|
||||
# shipped AT the OLD ref, headless.
|
||||
# installer-script+desktop the same one-liner with -IncludeDesktop:
|
||||
# builds Hermes.exe AND registers Start Menu /
|
||||
# Desktop shortcuts.
|
||||
# hermes-update venv hermes.exe update.
|
||||
# open-app-update the app's own Update button, app launched
|
||||
# from the installed exe under Playwright's
|
||||
# Electron driver (Settings -> About ->
|
||||
# "Update now"); the production hand-off chain
|
||||
# runs untouched.
|
||||
# hermes-desktop-app-update the same button, app launched via `hermes
|
||||
# desktop`: the driver captures the product's
|
||||
# own spawn (argv/cwd/env) and re-executes it
|
||||
# under Playwright.
|
||||
#
|
||||
# Method pairs without a driver arm yet NATIVELY SKIP (grey check, no
|
||||
# runner): the capability knowledge lives here, next to the driver, so the
|
||||
# caller can dispatch every declared combination without knowing which ones
|
||||
# work.
|
||||
#
|
||||
# Call it:
|
||||
#
|
||||
# jobs:
|
||||
# windows:
|
||||
# uses: ./.github/workflows/install-e2e-windows-run.yml
|
||||
# with:
|
||||
# install-method: desktop-installer@latest
|
||||
# update-method: open-app-update
|
||||
# install-ref: v2026.8.3
|
||||
|
||||
name: install-e2e windows leg
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
install-method:
|
||||
description: 'How OLD gets installed. Supported: desktop-installer@latest (website exe, AHK-clicked), installer-script (irm | iex install.ps1) and installer-script+desktop (the same with -IncludeDesktop). Declared-but-TODO methods skip.'
|
||||
required: true
|
||||
type: string
|
||||
update-method:
|
||||
description: 'How the install updates to HEAD. Supported: open-app-update (Update button under Playwright, from a desktop-bearing install), hermes-desktop-app-update (same button, app launched via hermes desktop), hermes-update, installer-script, installer-script+desktop, desktop-installer@latest (re-download Hermes-Setup.exe, AHK clicks Install over the existing install).'
|
||||
required: true
|
||||
type: string
|
||||
install-ref:
|
||||
description: 'Ref to install as OLD (served as main while the installer runs). auto = the newest release tag in the checkout.'
|
||||
required: false
|
||||
type: string
|
||||
default: auto
|
||||
tag-has-desktop:
|
||||
description: "Whether install-ref ships the desktop app (apps/desktop). The caller annotates this from the tag's own tree; desktop-method legs from pre-desktop releases natively skip."
|
||||
required: false
|
||||
type: boolean
|
||||
default: true
|
||||
leg-id:
|
||||
description: 'Artifact-safe matrix leg id (from generate-e2e-matrix.mjs legId). Names this leg''s logs + player artifacts so the report job can link a row to its zip.'
|
||||
required: true
|
||||
type: string
|
||||
setup-exe-url:
|
||||
description: 'Bootstrap installer to install OLD with. Default: the latest published one — what a user downloads today.'
|
||||
required: false
|
||||
type: string
|
||||
default: https://hermes-assets.nousresearch.com/Hermes-Setup.exe
|
||||
timeout-minutes:
|
||||
description: 'Job timeout. The install leg does real toolchain work and the update leg a full Electron rebuild.'
|
||||
required: false
|
||||
type: number
|
||||
default: 60
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
e2e:
|
||||
# Short static name on purpose: name expressions render UNEXPANDED on
|
||||
# skipped jobs.
|
||||
name: e2e
|
||||
# The implemented {install x update} pairs. Two rules feed the table:
|
||||
# * every desktop-surface method needs the starting tag to ship
|
||||
# apps/desktop (pre-desktop releases have no window to launch, no
|
||||
# Update button to click, no -IncludeDesktop to pass);
|
||||
# * open-app-update needs an OS entry point, which only the
|
||||
# desktop-bearing installs create.
|
||||
if: >-
|
||||
(inputs.install-method == 'installer-script'
|
||||
|| (contains(fromJSON('["installer-script+desktop", "desktop-installer@latest"]'), inputs.install-method) && inputs.tag-has-desktop))
|
||||
&& (contains(fromJSON('["hermes-update", "installer-script"]'), inputs.update-method)
|
||||
|| (contains(fromJSON('["installer-script+desktop", "hermes-desktop-app-update", "desktop-installer@latest"]'), inputs.update-method) && inputs.tag-has-desktop)
|
||||
|| (inputs.update-method == 'open-app-update' && inputs.tag-has-desktop
|
||||
&& contains(fromJSON('["desktop-installer@latest", "installer-script+desktop"]'), inputs.install-method)))
|
||||
runs-on: windows-latest
|
||||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
|
||||
env:
|
||||
# Sibling of the checkout (D:\a\hermes-agent\hermes-desktop-gui-e2e):
|
||||
# outside the repo so the staged bare clone and the install never
|
||||
# collide with the checkout itself. NOTE: ${{ runner.temp }} is NOT
|
||||
# available in job-level env (only github/inputs/matrix/needs/
|
||||
# secrets/strategy/vars).
|
||||
HERMES_E2E_WORKROOT: ${{ github.workspace }}\..\hermes-desktop-gui-e2e
|
||||
|
||||
steps:
|
||||
# Full history: the driver bare-clones this checkout as the repo the
|
||||
# installer/updater talk to, and both OLD and HEAD must be reachable
|
||||
# in that clone. A shallow checkout cannot serve either need.
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
# One recording mechanism on every OS: the composite action installs
|
||||
# ffmpeg (cached - winget's download is the slow part), starts the
|
||||
# capture, and record-stop fails on a zero-frame file so a silently
|
||||
# missing recording cannot go green.
|
||||
- name: Start screen recording
|
||||
uses: ./.github/actions/e2e-screen-record
|
||||
with:
|
||||
mode: start
|
||||
output: ${{ github.workspace }}\gui-e2e-proof\recording.mkv
|
||||
|
||||
- name: Stage serve repo (main -> ${{ inputs.install-ref }})
|
||||
shell: powershell
|
||||
run: powershell -NoProfile -ExecutionPolicy Bypass -File tests\install\windows-e2e.ps1 -Phase stage -InstallMethod "${{ inputs.install-method }}" -Route "${{ inputs.update-method }}" -InstallRef "${{ inputs.install-ref }}" -SetupExeUrl ${{ inputs.setup-exe-url }}
|
||||
|
||||
- name: Install ${{ inputs.install-ref }} (${{ inputs.install-method }})
|
||||
shell: powershell
|
||||
run: powershell -NoProfile -ExecutionPolicy Bypass -File tests\install\windows-e2e.ps1 -Phase install -InstallMethod "${{ inputs.install-method }}" -Route "${{ inputs.update-method }}" -InstallRef "${{ inputs.install-ref }}" -SetupExeUrl ${{ inputs.setup-exe-url }}
|
||||
|
||||
- name: Update ${{ inputs.install-ref }} -> HEAD (${{ inputs.update-method }})
|
||||
id: update
|
||||
shell: powershell
|
||||
run: powershell -NoProfile -ExecutionPolicy Bypass -File tests\install\windows-e2e.ps1 -Phase update -InstallMethod "${{ inputs.install-method }}" -Route "${{ inputs.update-method }}" -InstallRef "${{ inputs.install-ref }}" -SetupExeUrl ${{ inputs.setup-exe-url }}
|
||||
|
||||
- name: Stage known-failure receipt
|
||||
if: steps.update.outputs.known_failure != ''
|
||||
shell: pwsh
|
||||
run: |
|
||||
New-Item -ItemType Directory -Path gui-e2e-proof -Force | Out-Null
|
||||
Copy-Item -LiteralPath (Join-Path $env:HERMES_E2E_WORKROOT 'known-failure.json') -Destination gui-e2e-proof/known-failure.json
|
||||
|
||||
- name: Upload known-failure receipt
|
||||
if: steps.update.outputs.known_failure != ''
|
||||
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
|
||||
with:
|
||||
name: install-e2e-known-${{ steps.update.outputs.known_failure }}--${{ inputs.leg-id }}
|
||||
path: gui-e2e-proof/known-failure.json
|
||||
if-no-files-found: error
|
||||
retention-days: 14
|
||||
|
||||
- name: Stop screen recording
|
||||
if: always()
|
||||
uses: ./.github/actions/e2e-screen-record
|
||||
with:
|
||||
mode: stop
|
||||
output: ${{ github.workspace }}\gui-e2e-proof\recording.mkv
|
||||
|
||||
- name: Remux recording for browser playback
|
||||
if: always()
|
||||
shell: pwsh
|
||||
run: |
|
||||
$mkv = "$env:GITHUB_WORKSPACE\gui-e2e-proof\recording.mkv"
|
||||
if (Test-Path -LiteralPath $mkv) {
|
||||
& ffmpeg -y -hide_banner -loglevel error -i $mkv -c copy "$env:GITHUB_WORKSPACE\gui-e2e-proof\recording.mp4"
|
||||
}
|
||||
|
||||
- name: Collect proof + logs
|
||||
if: always()
|
||||
shell: powershell
|
||||
run: |
|
||||
$out = "gui-e2e-proof"
|
||||
New-Item -ItemType Directory -Path $out -Force | Out-Null
|
||||
$work = $env:HERMES_E2E_WORKROOT
|
||||
$home_ = Join-Path $work "hermes-home"
|
||||
foreach ($pair in @(
|
||||
@{ src = (Join-Path $work "proof"); dst = "proof" },
|
||||
@{ src = (Join-Path $work "logs"); dst = "driver-logs" },
|
||||
@{ src = (Join-Path $work "shas.json"); dst = "shas.json" },
|
||||
@{ src = (Join-Path $home_ "logs"); dst = "logs" },
|
||||
@{ src = (Join-Path $home_ ".hermes-update-result.json"); dst = ".hermes-update-result.json" }
|
||||
)) {
|
||||
if (Test-Path $pair.src) { Copy-Item $pair.src (Join-Path $out $pair.dst) -Recurse -Force }
|
||||
}
|
||||
|
||||
- name: Upload proof + logs
|
||||
if: always()
|
||||
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
|
||||
with:
|
||||
name: install-e2e-logs-${{ inputs.leg-id }}
|
||||
path: gui-e2e-proof
|
||||
retention-days: 14
|
||||
if-no-files-found: ignore
|
||||
@@ -2,15 +2,35 @@ name: Install & Update E2E
|
||||
|
||||
# Can a user on a released version get to this commit?
|
||||
#
|
||||
# For each release we sample, a leg installs that release through the real
|
||||
# `curl | install.sh` one-liner (uv, a managed Python, Node, the venv) inside
|
||||
# scripts/dev-sandbox.sh, then applies one update route and requires the
|
||||
# checkout to land on this commit with a working `hermes`.
|
||||
# The support matrix -- every {os, install-method, update-method} combination
|
||||
# a user could be on -- lives in scripts/sandbox/generate-e2e-matrix.mjs.
|
||||
# generate-matrix expands it against the picked release tags into one leg
|
||||
# per {combination, tag}, split into one matrix job per OS:
|
||||
#
|
||||
# Matrix: linux the real curl|bash install one-liner, isolated by a
|
||||
# git URL redirect to a local bare clone
|
||||
# (install-e2e-run.yml)
|
||||
# Matrix: windows the real desktop user flow: website Hermes-Setup.exe
|
||||
# clicked by AutoHotkey, update via the app, Playwright
|
||||
# clicking "Update now" (install-e2e-windows-run.yml)
|
||||
# Matrix: macos script installs on the shared OS-agnostic driver,
|
||||
# plus the real desktop user flow: website
|
||||
# Hermes-Setup.dmg mounted and run, updates via the
|
||||
# app under Playwright (install-e2e-macos-run.yml)
|
||||
#
|
||||
# Every combination is dispatched to its OS's run workflow; the run
|
||||
# workflow natively skips (grey) what its driver cannot run yet -- an
|
||||
# unimplemented method pair, or a starting tag that predates the surface
|
||||
# under test (pick-releases annotates each tag with what its tree ships,
|
||||
# e.g. whether the desktop app exists yet). Capability knowledge lives
|
||||
# next to each driver, never here and never in the generator: declaring a
|
||||
# method is a spec edit, implementing one is flipping the run workflow's
|
||||
# gate.
|
||||
#
|
||||
# The starting versions are chosen at runtime from the repo's release tags
|
||||
# (scripts/sandbox/pick-release-tags.sh): newest, oldest, and a spread between.
|
||||
# A hardcoded list would stop covering the newest release the day after it
|
||||
# ships, and would pin an "oldest" that nobody still runs.
|
||||
# (scripts/sandbox/pick-release-tags.sh): newest, oldest, and a spread
|
||||
# between. A hardcoded list would stop covering the newest release the day
|
||||
# after it ships, and would pin an "oldest" that nobody still runs.
|
||||
#
|
||||
# Triggers:
|
||||
# * every 12 hours, so upstream drift (a new uv, a Node bump, a PyPI change)
|
||||
@@ -27,16 +47,21 @@ on:
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
route:
|
||||
description: 'Which update route to exercise.'
|
||||
description: 'Which combinations to run. all = every OS; both/update/installer = the linux legs; windows-desktop = the windows legs; macos-desktop = the macos legs.'
|
||||
required: false
|
||||
type: choice
|
||||
default: both
|
||||
options: [both, update, installer]
|
||||
default: all
|
||||
options: [all, both, update, installer, windows-desktop, macos-desktop]
|
||||
tag-count:
|
||||
description: 'How many release tags to sample (newest, oldest, and a spread between).'
|
||||
required: false
|
||||
type: string
|
||||
default: '5'
|
||||
default: '3'
|
||||
install-ref:
|
||||
description: 'Optional exact release tag for a focused reproduction; overrides tag-count.'
|
||||
required: false
|
||||
type: string
|
||||
default: ''
|
||||
schedule:
|
||||
# Every 12 hours, off the hour to avoid the top-of-hour runner crunch.
|
||||
- cron: '20 7,19 * * *'
|
||||
@@ -54,8 +79,9 @@ concurrency:
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
# Which released versions do we test updating FROM? Resolved once and shared
|
||||
# by both route matrices, so the two routes cover the same set.
|
||||
# Which released versions do we test updating FROM? Resolved once,
|
||||
# annotated with what each tag's own tree supports, and shared by every
|
||||
# OS's matrix so all combos cover the same set.
|
||||
pick-releases:
|
||||
name: Pick release tags
|
||||
runs-on: ubuntu-latest
|
||||
@@ -63,9 +89,10 @@ jobs:
|
||||
outputs:
|
||||
tags: ${{ steps.pick.outputs.tags }}
|
||||
steps:
|
||||
# This job only reads tag names and runs one script, so take the cheap
|
||||
# checkout: no blobs (filter), no other files (sparse), but DO fetch tags
|
||||
# -- they are the whole input, and the default shallow checkout has none.
|
||||
# This job only reads tag names and trees, so take the cheap
|
||||
# checkout: no blobs (filter), no other files (sparse), but DO fetch
|
||||
# tags -- they are the whole input, and the default shallow checkout
|
||||
# has none.
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
with:
|
||||
filter: blob:none
|
||||
@@ -73,38 +100,179 @@ jobs:
|
||||
sparse-checkout: scripts/sandbox/pick-release-tags.sh
|
||||
sparse-checkout-cone-mode: false
|
||||
- id: pick
|
||||
env:
|
||||
# Dispatch inputs never touch shell syntax directly: TAG_COUNT
|
||||
# arrives via the environment and is validated decimal-only (bash
|
||||
# arithmetic reads a leading zero as octal). GitHub's 256-job cap
|
||||
# applies to each per-OS matrix separately; at 10 tags the largest
|
||||
# is windows at 180 (first over the cap at 15 tags = 270).
|
||||
TAG_COUNT: ${{ inputs.tag-count || 2 }}
|
||||
INSTALL_REF: ${{ inputs.install-ref }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
tags="$(scripts/sandbox/pick-release-tags.sh --count '${{ inputs.tag-count || 5 }}')"
|
||||
[[ "$TAG_COUNT" =~ ^(10|[1-9])$ ]] || { echo "tag-count must be 1-10, got: $TAG_COUNT" >&2; exit 1; }
|
||||
if [ -n "$INSTALL_REF" ]; then
|
||||
[[ "$INSTALL_REF" =~ ^v[0-9]+\.[0-9]+\.[0-9]+(\.[0-9]+)?$ ]] || { echo 'install-ref must be an exact release tag' >&2; exit 1; }
|
||||
git rev-parse --verify "refs/tags/$INSTALL_REF^{commit}" >/dev/null
|
||||
tags="$(jq -cn --arg ref "$INSTALL_REF" '[$ref]')"
|
||||
else
|
||||
tags="$(scripts/sandbox/pick-release-tags.sh --count "$TAG_COUNT")"
|
||||
fi
|
||||
echo "Testing updates from: $tags"
|
||||
echo "tags=$tags" >> "$GITHUB_OUTPUT"
|
||||
# Annotate each tag with what its own tree supports, so run
|
||||
# workflows can natively skip surfaces the starting version does
|
||||
# not have. Today: does the release ship the desktop app
|
||||
# (apps/desktop, #20059)? Cheaper here -- the tags are already
|
||||
# fetched -- than a probe job per leg.
|
||||
enriched="$(for t in $(echo "$tags" | jq -r '.[]'); do
|
||||
if git ls-tree -d "$t" apps/desktop | grep -q .; then d=true; else d=false; fi
|
||||
echo "{\"ref\":\"$t\",\"desktop\":$d}"
|
||||
done | jq -sc .)"
|
||||
echo "Annotated: $enriched"
|
||||
echo "tags=$enriched" >> "$GITHUB_OUTPUT"
|
||||
|
||||
# `hermes update` -- the route most users take.
|
||||
update:
|
||||
if: github.event_name != 'workflow_dispatch' || inputs.route != 'installer'
|
||||
# Expand the support matrix against the picked tags: one leg per
|
||||
# {os, install-method, update-method, tag}, split into a matrix per OS.
|
||||
generate-matrix:
|
||||
name: Expand combinations
|
||||
needs: pick-releases
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
outputs:
|
||||
linux: ${{ steps.gen.outputs.linux }}
|
||||
windows: ${{ steps.gen.outputs.windows }}
|
||||
macos: ${{ steps.gen.outputs.macos }}
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
with:
|
||||
sparse-checkout: scripts/sandbox/generate-e2e-matrix.mjs
|
||||
sparse-checkout-cone-mode: false
|
||||
- id: gen
|
||||
run: |
|
||||
set -euo pipefail
|
||||
matrices="$(node scripts/sandbox/generate-e2e-matrix.mjs \
|
||||
--tags '${{ needs.pick-releases.outputs.tags }}')"
|
||||
echo "$matrices"
|
||||
for key in linux windows macos; do
|
||||
echo "$key=$(echo "$matrices" | node -e 'let d="";process.stdin.on("data",c=>d+=c).on("end",()=>console.log(JSON.stringify(JSON.parse(d)[process.argv[1]])))' "$key")" >> "$GITHUB_OUTPUT"
|
||||
done
|
||||
# The plan, human-readable: a combination x starting-tag chart on
|
||||
# the run's summary page.
|
||||
node scripts/sandbox/generate-e2e-matrix.mjs \
|
||||
--tags '${{ needs.pick-releases.outputs.tags }}' \
|
||||
--format markdown >> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
linux:
|
||||
name: ${{ matrix.name }}
|
||||
# The update/installer route choices map to the linux update methods;
|
||||
# either way the whole linux matrix runs (legs are cheap and the
|
||||
# distinction wasn't worth a filter layer in the generator).
|
||||
if: github.event_name != 'workflow_dispatch' || contains(fromJSON('["all", "both", "update", "installer"]'), inputs.route)
|
||||
needs: generate-matrix
|
||||
strategy:
|
||||
# One release breaking is worth knowing about even if another already
|
||||
# One leg breaking is worth knowing about even if another already
|
||||
# failed, so let every leg report.
|
||||
fail-fast: false
|
||||
matrix:
|
||||
install-ref: ${{ fromJSON(needs.pick-releases.outputs.tags) }}
|
||||
matrix: ${{ fromJSON(needs.generate-matrix.outputs.linux) }}
|
||||
uses: ./.github/workflows/install-e2e-run.yml
|
||||
with:
|
||||
route: update
|
||||
install-ref: ${{ matrix.install-ref }}
|
||||
install-method: ${{ matrix.install_method }}
|
||||
update-method: ${{ matrix.update_method }}
|
||||
install-ref: ${{ matrix.install_ref }}
|
||||
tag-has-desktop: ${{ matrix.tag_has_desktop }}
|
||||
leg-id: ${{ matrix.leg_id }}
|
||||
|
||||
# Re-running the curl one-liner over an existing checkout: autostash + pull
|
||||
# rather than the updater's own git handling.
|
||||
installer:
|
||||
if: github.event_name != 'workflow_dispatch' || inputs.route != 'update'
|
||||
needs: pick-releases
|
||||
windows:
|
||||
name: ${{ matrix.name }}
|
||||
if: github.event_name != 'workflow_dispatch' || contains(fromJSON('["all", "windows-desktop"]'), inputs.route)
|
||||
needs: generate-matrix
|
||||
strategy:
|
||||
fail-fast: false
|
||||
max-parallel: 3
|
||||
matrix:
|
||||
install-ref: ${{ fromJSON(needs.pick-releases.outputs.tags) }}
|
||||
uses: ./.github/workflows/install-e2e-run.yml
|
||||
matrix: ${{ fromJSON(needs.generate-matrix.outputs.windows) }}
|
||||
uses: ./.github/workflows/install-e2e-windows-run.yml
|
||||
with:
|
||||
route: installer
|
||||
install-ref: ${{ matrix.install-ref }}
|
||||
install-method: ${{ matrix.install_method }}
|
||||
update-method: ${{ matrix.update_method }}
|
||||
install-ref: ${{ matrix.install_ref }}
|
||||
tag-has-desktop: ${{ matrix.tag_has_desktop }}
|
||||
leg-id: ${{ matrix.leg_id }}
|
||||
|
||||
macos:
|
||||
name: ${{ matrix.name }}
|
||||
if: github.event_name != 'workflow_dispatch' || contains(fromJSON('["all", "macos-desktop"]'), inputs.route)
|
||||
needs: generate-matrix
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix: ${{ fromJSON(needs.generate-matrix.outputs.macos) }}
|
||||
# Two driver arms: the OS-agnostic script driver (shared with linux)
|
||||
# and the published-dmg GUI driver; the run workflow routes.
|
||||
uses: ./.github/workflows/install-e2e-macos-run.yml
|
||||
with:
|
||||
install-method: ${{ matrix.install_method }}
|
||||
update-method: ${{ matrix.update_method }}
|
||||
install-ref: ${{ matrix.install_ref }}
|
||||
tag-has-desktop: ${{ matrix.tag_has_desktop }}
|
||||
leg-id: ${{ matrix.leg_id }}
|
||||
|
||||
# The leg player: one static HTML for the whole run. Uploaded BEFORE the
|
||||
# matrix legs so it exists even when every leg dies; the report job links
|
||||
# every ran leg to it with that leg's logs zip as a #zip= hash param
|
||||
# (hash survives the artifact URL's server-side redirect, the query does
|
||||
# not). archive: false makes GitHub name the artifact after the FILE
|
||||
# (playback.html), ignoring the name: input -- harmless, the renderer
|
||||
# looks it up by that name.
|
||||
leg-player:
|
||||
name: Upload leg player
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
with:
|
||||
sparse-checkout: tests/install/e2e-assets/playback.html
|
||||
sparse-checkout-cone-mode: false
|
||||
- uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
|
||||
with:
|
||||
name: install-e2e-player
|
||||
path: tests/install/e2e-assets/playback.html
|
||||
archive: false
|
||||
retention-days: 14
|
||||
if-no-files-found: error
|
||||
|
||||
# The outcome, human-readable: the plan chart again, with each cell
|
||||
# replaced by how that leg actually concluded. Per-leg conclusions are
|
||||
# NOT reachable through `needs` (a matrix job's result collapses to one
|
||||
# aggregate), so the table body comes from the run's own job list; the
|
||||
# `needs` results only sequence this job after every leg and provide
|
||||
# the per-OS aggregates.
|
||||
report:
|
||||
name: Result chart
|
||||
if: always()
|
||||
needs: [leg-player, pick-releases, linux, windows, macos]
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
with:
|
||||
sparse-checkout: |
|
||||
scripts/sandbox/generate-e2e-matrix.mjs
|
||||
tests/install/e2e-assets/known-failures.json
|
||||
sparse-checkout-cone-mode: false
|
||||
- env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
{
|
||||
echo "OS jobs: linux ${{ needs.linux.result }}, windows ${{ needs.windows.result }}, macos ${{ needs.macos.result }}"
|
||||
echo
|
||||
# The tag annotations let the chart say WHY a cell skipped
|
||||
# (pre-desktop vs declared TODO) instead of a flat "skip".
|
||||
gh api "repos/${{ github.repository }}/actions/runs/${{ github.run_id }}/jobs?per_page=100" \
|
||||
--paginate --jq '.jobs[] | {name, conclusion}' > /tmp/e2e-jobs.ndjson
|
||||
gh api "repos/${{ github.repository }}/actions/runs/${{ github.run_id }}/artifacts?per_page=100" \
|
||||
--paginate --jq '.artifacts[] | {name, id}' > /tmp/e2e-artifacts.ndjson
|
||||
echo 'Legend: ✅ upgrade passed · known [n] = exact historical failure, see footnote · ❌ unexpected failure · pre-desktop / TODO = why a leg skipped · 📼 opens the leg player (recording + synced logs)'
|
||||
echo
|
||||
node scripts/sandbox/generate-e2e-matrix.mjs --format results \
|
||||
--tags '${{ needs.pick-releases.outputs.tags }}' \
|
||||
--artifacts /tmp/e2e-artifacts.ndjson < /tmp/e2e-jobs.ndjson
|
||||
} >> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
@@ -36,3 +36,11 @@ jobs:
|
||||
- name: 8.3 short-path normalization (Windows PowerShell 5.1)
|
||||
shell: powershell
|
||||
run: powershell -NoProfile -ExecutionPolicy Bypass -File scripts/tests/test-install-ps1-longpath.ps1
|
||||
|
||||
- name: System Node and npm compatibility (pwsh 7)
|
||||
shell: pwsh
|
||||
run: pwsh -NoProfile -ExecutionPolicy Bypass -File scripts/tests/test-install-ps1-node-compatibility.ps1
|
||||
|
||||
- name: System Node and npm compatibility (Windows PowerShell 5.1)
|
||||
shell: powershell
|
||||
run: powershell -NoProfile -ExecutionPolicy Bypass -File scripts/tests/test-install-ps1-node-compatibility.ps1
|
||||
|
||||
@@ -178,3 +178,25 @@ jobs:
|
||||
|
||||
- name: Run footgun checker
|
||||
run: python scripts/check-windows-footguns.py --all
|
||||
|
||||
# The Sep 2026 decomposition kept old import paths alive for external plugins
|
||||
# (PLUGIN-COMPAT blocks, see COMPAT_MANIFEST.md). They are removed on schedule by
|
||||
# reverting one commit, so in-tree code must never depend on them.
|
||||
- name: Forbid in-tree use of plugin-compat pointers
|
||||
run: python scripts/check_compat_pointers.py
|
||||
|
||||
# Advisory: dropped public names / methods / test defs vs the PR base, printed into the log.
|
||||
# A refactor that silently removes a public symbol breaks plugins that import it; the Sep 2026
|
||||
# decomposition opened with 1,703 such drops that reviewers had to find by hand.
|
||||
# Advisory: it never fails the job. The checkout above is depth-1, so deepen both sides until
|
||||
# a merge-base exists (the script refuses to report a clean diff without one, by design).
|
||||
- name: Public-surface diff vs base (advisory)
|
||||
if: github.event_name == 'pull_request'
|
||||
continue-on-error: true
|
||||
run: |
|
||||
git fetch --no-tags --deepen=200 origin "${{ github.base_ref }}" HEAD
|
||||
for i in 1 2 3; do
|
||||
git merge-base "origin/${{ github.base_ref }}" HEAD >/dev/null 2>&1 && break
|
||||
git fetch --no-tags --deepen=1000 origin "${{ github.base_ref }}" HEAD
|
||||
done
|
||||
python scripts/ci/check_public_surface.py --base "origin/${{ github.base_ref }}" --head HEAD
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
name: Profile Artifact Boundary
|
||||
|
||||
# A reusable, unconditional guard for the incident class in #92457. Ignore
|
||||
# files reduce accidental staging; this job is the enforcement boundary that
|
||||
# still catches `git add -f` and generated files present during a build.
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
check-profile-artifacts:
|
||||
name: Reject profile archives
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
- name: Reject profile exports in the checkout
|
||||
run: python3 scripts/ci/check_profile_archive_boundary.py
|
||||
@@ -27,6 +27,19 @@ name: OS-specific tests
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
desktop_updater:
|
||||
description: >-
|
||||
Run the Windows desktop-update hand-off integration tests
|
||||
(tests/test_desktop_update_windows_*.py). These spawn the real
|
||||
scripts/desktop-update/windows.ps1 and poll its loopback server, so
|
||||
they carry process-timing noise a shared runner amplifies; the
|
||||
caller gates them on the classifier's desktop_updater lane so a PR
|
||||
that never touched that surface cannot be failed by it. Push /
|
||||
dispatch runs fail open (classifier sets every lane true).
|
||||
type: boolean
|
||||
required: false
|
||||
default: true
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
@@ -134,9 +147,23 @@ jobs:
|
||||
# would therefore abort the script on any non-zero exit and the
|
||||
# exit-5 branch below would be unreachable dead code — the job
|
||||
# would still fail red, but the diagnostic would never print.
|
||||
# Desktop-update hand-off integration tests spawn the real
|
||||
# windows.ps1; deselect them unless the PR touched that surface
|
||||
# (see the workflow_call input). ``--ignore-glob`` keeps the file
|
||||
# list above intact, so a renamed test file still trips the
|
||||
# zero-tests guard rather than silently vanishing.
|
||||
# (bash 3.2 on the macOS runner: an empty array under ``set -u`` is
|
||||
# an unbound-variable error, hence the ``${arr[@]+...}`` idiom.)
|
||||
EXTRA_ARGS=()
|
||||
if [ "${{ inputs.desktop_updater }}" != "true" ]; then
|
||||
echo "desktop_updater lane off: skipping tests/test_desktop_update_windows_*.py"
|
||||
EXTRA_ARGS+=(--ignore-glob='*test_desktop_update_windows_*.py')
|
||||
fi
|
||||
|
||||
status=0
|
||||
uv run --no-sync python -m pytest \
|
||||
"$@" \
|
||||
${EXTRA_ARGS[@]+"${EXTRA_ARGS[@]}"} \
|
||||
-m "${{ matrix.marker }} and not integration" \
|
||||
-v --tb=short || status=$?
|
||||
if [ "$status" -eq 5 ]; then
|
||||
|
||||
@@ -50,7 +50,7 @@ jobs:
|
||||
- name: Install dependencies
|
||||
uses: ./.github/actions/retry
|
||||
with:
|
||||
command: uv sync --locked --python 3.11 --extra dev
|
||||
command: uv sync --locked --python 3.11 --extra dev --extra messaging
|
||||
|
||||
- name: Run venv-holder live E2E
|
||||
shell: bash
|
||||
@@ -58,4 +58,23 @@ jobs:
|
||||
set -uo pipefail
|
||||
uv run --no-sync python -m pytest \
|
||||
tests/hermes_cli/test_venv_holder_windows_live.py \
|
||||
tests/hermes_cli/test_taskkill_identity_windows_live.py \
|
||||
tests/hermes_cli/test_git_trampoline_windows_live.py \
|
||||
"tests/hermes_cli/test_managed_uv.py::TestWindowsRuntimeSelfLock" \
|
||||
-o addopts= -v -p no:cacheprovider
|
||||
|
||||
- name: Run Telegram CLOSE-WAIT reconnect live E2E (#87057)
|
||||
shell: bash
|
||||
run: |
|
||||
set -uo pipefail
|
||||
uv run --no-sync python -m pytest \
|
||||
tests/gateway/test_telegram_closewait_windows_live.py \
|
||||
-o addopts= -v -p no:cacheprovider
|
||||
|
||||
- name: Run background-executor spawn parity live E2E (#70716)
|
||||
shell: bash
|
||||
run: |
|
||||
set -uo pipefail
|
||||
uv run --no-sync python -m pytest \
|
||||
tests/tools/test_process_registry_windows_live.py \
|
||||
-o addopts= -v -p no:cacheprovider
|
||||
|
||||
+10
-2
@@ -98,15 +98,19 @@ apps/desktop/src/**/*.d.ts
|
||||
!apps/desktop/src/global.d.ts
|
||||
!apps/desktop/src/vite-env.d.ts
|
||||
|
||||
# Repo-root build/debug artifacts that must never be committed
|
||||
# Build/debug artifacts that must never be committed
|
||||
/log.txt
|
||||
/sqlite_leak_fix.png
|
||||
/*.png.bak
|
||||
/default.tar.gz
|
||||
*.tar.gz
|
||||
*.tgz
|
||||
apps/shared/src/**/*.js
|
||||
apps/shared/src/**/*.js.map
|
||||
apps/shared/src/**/*.d.ts
|
||||
apps/desktop/release/
|
||||
# stage-and-swap Desktop rebuild output (#86443); removed after the swap, but
|
||||
# a killed build must not leave the checkout dirty
|
||||
apps/desktop/.staging-*/
|
||||
*.tsbuildinfo
|
||||
|
||||
# Web UI assets — synced from @nous-research/ui at build time via
|
||||
@@ -215,3 +219,7 @@ native/fts5_cjk/*.so
|
||||
# interrupted; consumed by launch-time recovery. Never commit it (was tracked
|
||||
# by accident via 3a69e34702, removed in the #72002 salvage).
|
||||
.lazy-refresh-incomplete
|
||||
.skills_prompt_snapshot.json
|
||||
|
||||
# Disposable profile created by scripts/probe_active_session_exclusivity.py
|
||||
.probe-home/
|
||||
|
||||
+3869
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -210,7 +210,7 @@ hermes-agent/
|
||||
| `~/.hermes/skills/` | Todas las habilidades activas (incluidas + instaladas desde hub + creadas por el agente) |
|
||||
| `~/.hermes/memories/` | Memoria persistente (MEMORY.md, USER.md) |
|
||||
| `~/.hermes/state.db` | Base de datos de sesiones SQLite |
|
||||
| `~/.hermes/sessions/` | Índice de enrutamiento del gateway (`sessions.json`), migas de pan de solicitudes, transcripciones `*.jsonl` del gateway y (opcionalmente) snapshots JSON por sesión cuando `sessions.write_json_snapshots: true` está configurado. Los snapshots por sesión están desactivados por defecto; state.db es canónica. |
|
||||
| `~/.hermes/sessions/` | Índice de enrutamiento del gateway (`sessions.json`), migas de pan de solicitudes, transcripciones `*.jsonl` del gateway y exportaciones explícitas con `/save`. Ya no se escriben snapshots JSON automáticos; los archivos existentes se conservan y state.db es canónica. |
|
||||
| `~/.hermes/cron/` | Datos de trabajos programados |
|
||||
| `~/.hermes/whatsapp/session/` | Credenciales del puente WhatsApp |
|
||||
|
||||
|
||||
+19
-10
@@ -216,14 +216,17 @@ pytest tests/ -v
|
||||
|
||||
```
|
||||
hermes-agent/
|
||||
├── run_agent.py # AIAgent class — core conversation loop, tool dispatch, session persistence
|
||||
├── cli.py # HermesCLI class — interactive TUI, prompt_toolkit integration
|
||||
├── run_agent.py # AIAgent facade (~1.5k LOC) — the turn loop lives in agent/conversation_loop.py + agent/turn_*.py
|
||||
├── cli.py # HermesCLI class — interactive CLI orchestrator (~4.6k LOC + hermes_cli/cli_*_mixin.py)
|
||||
├── model_tools.py # Tool orchestration (thin layer over tools/registry.py)
|
||||
├── toolsets.py # Tool groupings and presets (hermes-cli, hermes-telegram, etc.)
|
||||
├── hermes_state.py # SQLite session database with FTS5 full-text search, session titles
|
||||
├── hermes_state.py # SessionDB facade (~1.4k LOC); implementation in hermes_state_*.py (21 siblings) — FTS5 search, session titles
|
||||
├── batch_runner.py # Parallel batch processing for trajectory generation
|
||||
│
|
||||
├── agent/ # Agent internals (extracted modules)
|
||||
│ ├── conversation_loop.py # run_conversation() — the agent turn loop (phases in turn_*.py)
|
||||
│ ├── tool_executor.py # Tool dispatch (inline agent-level tools, delegate, registry)
|
||||
│ ├── session_persistence.py # Session/trajectory saving
|
||||
│ ├── prompt_builder.py # System prompt assembly (identity, skills, context files, memory)
|
||||
│ ├── context_compressor.py # Auto-summarization when approaching context limits
|
||||
│ ├── auxiliary_client.py # Resolves auxiliary OpenAI clients (summarization, vision)
|
||||
@@ -233,16 +236,19 @@ hermes-agent/
|
||||
│
|
||||
├── hermes_cli/ # CLI command implementations
|
||||
│ ├── main.py # Entry point, argument parsing, command dispatch
|
||||
│ ├── cli_*_mixin.py # HermesCLI mixins (slash commands, display, session, ...)
|
||||
│ ├── config.py # Config management, migration, env var definitions
|
||||
│ ├── setup.py # Interactive setup wizard
|
||||
│ ├── auth.py # Provider resolution, OAuth, Nous Portal
|
||||
│ ├── auth.py # Provider resolution, OAuth, Nous Portal (facade + auth_*.py siblings)
|
||||
│ ├── models.py # OpenRouter model selection lists
|
||||
│ ├── banner.py # Welcome banner, ASCII art
|
||||
│ ├── commands.py # Central slash command registry (CommandDef), autocomplete, gateway helpers
|
||||
│ ├── callbacks.py # Interactive callbacks (clarify, sudo, approval)
|
||||
│ ├── doctor.py # Diagnostics
|
||||
│ ├── skills_hub.py # Skills Hub CLI + /skills slash command
|
||||
│ └── skin_engine.py # Skin/theme engine — data-driven CLI visual customization
|
||||
│ ├── skin_engine.py # Skin/theme engine — data-driven CLI visual customization
|
||||
│ ├── web_server.py # Dashboard server (facade + web_server_*.py siblings)
|
||||
│ └── web_routers/ # Dashboard FastAPI routers (one file per surface)
|
||||
│
|
||||
├── tools/ # Tool implementations (self-registering)
|
||||
│ ├── registry.py # Central tool registry (schemas, handlers, dispatch)
|
||||
@@ -252,7 +258,9 @@ hermes-agent/
|
||||
│ ├── web_tools.py # web_search, web_extract (Parallel/Firecrawl + Gemini summarization)
|
||||
│ ├── vision_tools.py # Image analysis via multimodal models
|
||||
│ ├── delegate_tool.py # Subagent spawning and parallel task execution
|
||||
│ ├── code_execution_tool.py # Sandboxed Python with RPC tool access
|
||||
│ ├── code_execution_tool.py # Sandboxed Python with RPC tool access (env allowlists in code_execution_env.py)
|
||||
│ ├── mcp_tool.py # MCP client (facade + mcp_tool_*.py siblings: config, discovery, transport, ...)
|
||||
│ ├── browser_tool.py # Browser automation (facade + browser_tool_*.py siblings)
|
||||
│ ├── session_search_tool.py # Search past conversations with FTS5 + anchored windows
|
||||
│ ├── cronjob_tools.py # Scheduled task management
|
||||
│ ├── skill_tools.py # Skill search, load, manage
|
||||
@@ -261,9 +269,10 @@ hermes-agent/
|
||||
│ ├── local.py, docker.py, ssh.py, singularity.py, modal.py, daytona.py
|
||||
│
|
||||
├── gateway/ # Messaging gateway
|
||||
│ ├── run.py # GatewayRunner — platform lifecycle, message routing, cron
|
||||
│ ├── run.py # GatewayRunner facade (~5.5k LOC); phases in run_*.py (startup, inbound, turn, busy, ...)
|
||||
│ ├── slash_commands_*.py # Gateway slash command handler mixins
|
||||
│ ├── config.py # Platform configuration resolution
|
||||
│ ├── session.py # Session store, context prompts, reset policies
|
||||
│ ├── session.py # Session store, context prompts, explicit resets (+ session_*.py siblings)
|
||||
│ └── platforms/ # Platform adapters
|
||||
│ ├── telegram.py, discord_adapter.py, slack.py, whatsapp.py
|
||||
│
|
||||
@@ -291,7 +300,7 @@ hermes-agent/
|
||||
| `~/.hermes/skills/` | All active skills (bundled + hub-installed + agent-created) |
|
||||
| `~/.hermes/memories/` | Persistent memory (MEMORY.md, USER.md) |
|
||||
| `~/.hermes/state.db` | SQLite session database |
|
||||
| `~/.hermes/sessions/` | Gateway routing index (`sessions.json`), request-dump breadcrumbs, gateway `*.jsonl` transcripts, and (optionally) per-session JSON snapshots when `sessions.write_json_snapshots: true` is set. The per-session snapshots are off by default; state.db is canonical. |
|
||||
| `~/.hermes/sessions/` | Gateway routing index (`sessions.json`), request-dump breadcrumbs, gateway `*.jsonl` transcripts, and explicit `/save` exports. Automatic per-session JSON snapshots are no longer written; state.db is canonical. |
|
||||
| `~/.hermes/cron/` | Scheduled job data |
|
||||
| `~/.hermes/whatsapp/session/` | WhatsApp bridge credentials |
|
||||
|
||||
@@ -320,7 +329,7 @@ User message → AIAgent._run_agent_loop()
|
||||
|
||||
- **Self-registering tools**: Each tool file calls `registry.register()` at import time. `model_tools.py` triggers discovery by importing all tool modules.
|
||||
- **Toolset grouping**: Tools are grouped into toolsets (`web`, `terminal`, `file`, `browser`, etc.) that can be enabled/disabled per platform.
|
||||
- **Session persistence**: All conversations are stored in SQLite (`hermes_state.py`) with full-text search and unique session titles. Per-session JSON snapshots in `~/.hermes/sessions/` were superseded by the SQLite store and are off by default; opt back in with `sessions.write_json_snapshots: true` if you have external tooling that consumes the JSON files directly.
|
||||
- **Session persistence**: All conversations are stored in SQLite (`hermes_state.py`) with full-text search and unique session titles. Automatic per-session JSON snapshots have been removed. Existing files are left untouched; use `/save json` or `hermes sessions export` for an explicit export.
|
||||
- **Ephemeral injection**: System prompts and prefill messages are injected at API call time, never persisted to the database or logs.
|
||||
- **Provider abstraction**: The agent works with any OpenAI-compatible API. Provider resolution happens at init time (Nous Portal OAuth, OpenRouter API key, or custom endpoint).
|
||||
- **Provider routing**: When using OpenRouter, `provider_routing` in config.yaml controls provider selection (sort by throughput/latency/price, allow/ignore specific providers, data retention policies). These are injected as `extra_body.provider` in API requests.
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
You are Hermes Agent, built by Nous Research. Be direct: match the length of your reply to the weight of the ask — a one-line question gets a one-line answer, and finished work gets a short report of what changed, what's verified, and what's left, never a replay of the process. No filler ("Great question," "I'd be happy to"), no restating the request back, no re-summarizing what you already said, no narrating tool calls the user can see. Plain claims over adjectives; when unsure, say so plainly. Agree because it's right, not because the user said it. Depth is earned — give it when the user asks for detail, teaches, or the stakes demand it, not by default.
|
||||
+32
-49
@@ -11,69 +11,52 @@ TERMINAL_SETUP_AUTH_METHOD_ID = "hermes-setup"
|
||||
def detect_provider() -> Optional[str]:
|
||||
"""Resolve the active Hermes runtime provider, or None if unavailable.
|
||||
|
||||
Treats a ``Callable`` ``api_key`` (Azure Foundry Entra ID bearer
|
||||
token provider — see :mod:`agent.azure_identity_adapter`) as a valid
|
||||
credential. Without this, ACP sessions for Entra-configured Foundry
|
||||
deployments silently default to ``"openrouter"`` and the ACP auth
|
||||
handshake rejects the legitimate provider.
|
||||
"""
|
||||
A callable ``api_key`` (Azure Foundry Entra ID bearer-token provider, see
|
||||
:mod:`agent.azure_identity_adapter`) counts as a valid credential; otherwise
|
||||
Entra-configured Foundry deployments would default to ``"openrouter"`` and
|
||||
the ACP auth handshake would reject the legitimate provider."""
|
||||
try:
|
||||
from hermes_cli.runtime_provider import resolve_runtime_provider
|
||||
runtime = resolve_runtime_provider()
|
||||
api_key = runtime.get("api_key")
|
||||
provider = runtime.get("provider")
|
||||
if not isinstance(provider, str) or not provider.strip():
|
||||
return None
|
||||
is_string_key = isinstance(api_key, str) and api_key.strip()
|
||||
is_callable_provider = callable(api_key) and not isinstance(api_key, str)
|
||||
if is_string_key or is_callable_provider:
|
||||
api_key, provider = runtime.get("api_key"), runtime.get("provider")
|
||||
if isinstance(provider, str) and provider.strip() and (
|
||||
(isinstance(api_key, str) and api_key.strip()) or callable(api_key)):
|
||||
return provider.strip().lower()
|
||||
except Exception:
|
||||
return None
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def has_provider() -> bool:
|
||||
"""Return True if Hermes can resolve any runtime provider credentials."""
|
||||
return detect_provider() is not None
|
||||
|
||||
|
||||
def build_auth_methods() -> list[Any]:
|
||||
"""Return registry-compatible ACP auth methods for Hermes.
|
||||
|
||||
The official ACP registry validates that agents advertise at least one
|
||||
usable auth method during the initial handshake. A fresh Zed install may
|
||||
not have Hermes provider credentials configured yet, so Hermes always
|
||||
advertises a terminal setup method. When credentials are already present,
|
||||
it also advertises the resolved provider as the default agent-managed
|
||||
runtime credential method.
|
||||
"""
|
||||
The ACP registry requires at least one usable auth method in the initial
|
||||
handshake. A fresh Zed install may have no Hermes credentials yet, so the
|
||||
terminal setup method is always advertised; when credentials resolve, the
|
||||
provider is also advertised as the default agent-managed runtime method."""
|
||||
from acp.schema import AuthMethodAgent, TerminalAuthMethod
|
||||
|
||||
methods: list[Any] = []
|
||||
provider = detect_provider()
|
||||
if provider:
|
||||
methods.append(
|
||||
AuthMethodAgent(
|
||||
id=provider,
|
||||
name=f"{provider} runtime credentials",
|
||||
description=(
|
||||
"Authenticate Hermes using the currently configured "
|
||||
f"{provider} runtime credentials."
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
methods.append(
|
||||
TerminalAuthMethod(
|
||||
id=TERMINAL_SETUP_AUTH_METHOD_ID,
|
||||
name="Configure Hermes provider",
|
||||
description=(
|
||||
"Open Hermes' interactive model/provider setup in a terminal. "
|
||||
"Use this when Hermes has not been configured on this machine yet."
|
||||
),
|
||||
type="terminal",
|
||||
args=["--setup"],
|
||||
)
|
||||
)
|
||||
methods.append(AuthMethodAgent(
|
||||
id=provider, name=f"{provider} runtime credentials",
|
||||
description=f"Authenticate Hermes using the currently configured {provider} runtime credentials.",
|
||||
))
|
||||
methods.append(TerminalAuthMethod(
|
||||
id=TERMINAL_SETUP_AUTH_METHOD_ID, name="Configure Hermes provider", type="terminal", args=["--setup"],
|
||||
description=("Open Hermes' interactive model/provider setup in a terminal. "
|
||||
"Use this when Hermes has not been configured on this machine yet."),
|
||||
))
|
||||
return methods
|
||||
|
||||
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
|
||||
def has_provider() -> bool:
|
||||
"""Return True if Hermes can resolve any runtime provider credentials."""
|
||||
return detect_provider() is not None
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
@@ -0,0 +1,292 @@
|
||||
"""Headless slash commands for ACP sessions (``/help``, ``/model``, ``/compress`` ...)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextvars
|
||||
import logging
|
||||
from collections import Counter
|
||||
from typing import Any
|
||||
|
||||
from acp.schema import AvailableCommand, AvailableCommandsUpdate, UnstructuredCommandInput
|
||||
|
||||
from acp_adapter.session import SessionState, _expand_acp_enabled_toolsets
|
||||
|
||||
logger = logging.getLogger("acp_adapter.server")
|
||||
|
||||
try:
|
||||
from hermes_cli import __version__ as HERMES_VERSION
|
||||
except Exception:
|
||||
HERMES_VERSION = "0.0.0"
|
||||
|
||||
|
||||
def _estimate_tokens(history: list, agent: Any, system_prompt: str | None = None, tools: Any = None) -> int:
|
||||
"""Rough request-token estimate over history + system prompt + tool schemas."""
|
||||
from agent.model_metadata import estimate_request_tokens_rough
|
||||
|
||||
if system_prompt is None:
|
||||
system_prompt = getattr(agent, "_cached_system_prompt", "") or ""
|
||||
if tools is None:
|
||||
tools = getattr(agent, "tools", None) or None
|
||||
return estimate_request_tokens_rough(history, system_prompt=system_prompt, tools=tools)
|
||||
|
||||
|
||||
def _queue_prompt(state: SessionState, text: str) -> int:
|
||||
with state.runtime_lock:
|
||||
state.queued_prompts.append(text)
|
||||
return len(state.queued_prompts)
|
||||
|
||||
|
||||
class SlashCommandsMixin:
|
||||
"""Slash-command surface for ``HermesACPAgent``; relies on ``_conn``, ``_send``, ``_schedule_soon``,
|
||||
``session_manager`` and ``_switch_model`` from the host class."""
|
||||
|
||||
# name -> (help text, advertised description, input hint)
|
||||
_COMMANDS: dict[str, tuple[str, str, str | None]] = {
|
||||
"help": ("Show available commands", "List available commands", None),
|
||||
"model": (
|
||||
"Show or change current model",
|
||||
"Show current model and provider, or switch models",
|
||||
"model name to switch to",
|
||||
),
|
||||
"tools": ("List available tools", "List available tools with descriptions", None),
|
||||
"context": ("Show conversation context info", "Show conversation message counts by role", None),
|
||||
"reset": ("Clear conversation history", "Clear conversation history", None),
|
||||
"compress": ("Compress conversation context", "Compress conversation context", None),
|
||||
"steer": (
|
||||
"Inject guidance into the currently running agent turn",
|
||||
"Inject guidance into the currently running agent turn",
|
||||
"guidance for the active turn",
|
||||
),
|
||||
"queue": (
|
||||
"Queue a prompt to run after the current turn finishes",
|
||||
"Queue a prompt to run after the current turn finishes",
|
||||
"prompt to run next",
|
||||
),
|
||||
"version": ("Show Hermes version", "Show Hermes version", None),
|
||||
}
|
||||
|
||||
|
||||
@classmethod
|
||||
def _available_commands(cls) -> list[AvailableCommand]:
|
||||
return [
|
||||
AvailableCommand(name=name, description=desc, input=UnstructuredCommandInput(hint=hint) if hint else None)
|
||||
for name, (_help, desc, hint) in cls._COMMANDS.items()
|
||||
]
|
||||
|
||||
async def _send_available_commands_update(self, session_id: str) -> None:
|
||||
"""Advertise supported slash commands to the connected ACP client."""
|
||||
if not self._conn:
|
||||
return
|
||||
update = AvailableCommandsUpdate(
|
||||
session_update="available_commands_update", available_commands=self._available_commands()
|
||||
)
|
||||
await self._send(session_id, update, fail_msg="Failed to advertise ACP slash commands for session %s")
|
||||
|
||||
def _schedule_available_commands_update(self, session_id: str) -> None:
|
||||
self._schedule_soon(lambda: self._send_available_commands_update(session_id))
|
||||
|
||||
def _handle_slash_command(self, text: str, state: SessionState) -> str | None:
|
||||
"""Dispatch a slash command; ``None`` for unknown ones so they fall through to the LLM."""
|
||||
parts = text.split(maxsplit=1)
|
||||
cmd = parts[0].lstrip("/").lower()
|
||||
args = parts[1].strip() if len(parts) > 1 else ""
|
||||
|
||||
if cmd not in self._COMMANDS:
|
||||
return None
|
||||
handler = getattr(self, f"_cmd_{cmd}")
|
||||
|
||||
# Handlers run on the loop thread, outside the per-turn cwd-pinning context. ``/compress``
|
||||
# and ``/model`` REBUILD the system prompt, so unpinned they'd bake the Hermes install tree
|
||||
# into the persisted cached prompt. Pin inside a fresh context: no leak, no teardown.
|
||||
def _dispatch() -> str | None:
|
||||
try:
|
||||
from agent.runtime_cwd import set_session_cwd
|
||||
|
||||
set_session_cwd(state.cwd)
|
||||
except Exception:
|
||||
logger.debug("Could not pin ACP session cwd for slash command", exc_info=True)
|
||||
return handler(args, state)
|
||||
|
||||
try:
|
||||
return contextvars.copy_context().run(_dispatch)
|
||||
except Exception as e:
|
||||
logger.error("Slash command /%s error: %s", cmd, e, exc_info=True)
|
||||
return f"Error executing /{cmd}: {e}"
|
||||
|
||||
def _cmd_help(self, args: str, state: SessionState) -> str:
|
||||
lines = ["Available commands:", ""]
|
||||
lines.extend(f" /{cmd:10s} {desc}" for cmd, (desc, _adv, _hint) in self._COMMANDS.items())
|
||||
lines.extend(["", "Unrecognized /commands are sent to the model as normal messages."])
|
||||
return "\n".join(lines)
|
||||
|
||||
def _cmd_model(self, args: str, state: SessionState) -> str:
|
||||
if not args:
|
||||
model = state.model or getattr(state.agent, "model", "unknown")
|
||||
provider = getattr(state.agent, "provider", None) or "auto"
|
||||
return f"Current model: {model}\nProvider: {provider}"
|
||||
|
||||
current_provider, target_provider, new_model = self._switch_model(state, args)
|
||||
provider_label = getattr(state.agent, "provider", None) or target_provider or current_provider or "openrouter"
|
||||
logger.info("Session %s: model switched to %s", state.session_id, new_model)
|
||||
return f"Model switched to: {new_model}\nProvider: {provider_label}"
|
||||
|
||||
def _cmd_tools(self, args: str, state: SessionState) -> str:
|
||||
try:
|
||||
from model_tools import get_tool_definitions
|
||||
from types import SimpleNamespace
|
||||
from agent.memory_manager import inject_memory_provider_tools
|
||||
|
||||
toolsets = _expand_acp_enabled_toolsets(getattr(state.agent, "enabled_toolsets", None) or ["hermes-acp"])
|
||||
tools = get_tool_definitions(enabled_toolsets=toolsets, quiet_mode=True)
|
||||
tool_view = SimpleNamespace(
|
||||
tools=list(tools or []),
|
||||
valid_tool_names={t.get("function", {}).get("name") for t in tools or [] if isinstance(t, dict)},
|
||||
enabled_toolsets=toolsets, _memory_manager=getattr(state.agent, "_memory_manager", None),
|
||||
)
|
||||
inject_memory_provider_tools(tool_view)
|
||||
tools = tool_view.tools
|
||||
if not tools:
|
||||
return "No tools available."
|
||||
lines = [f"Available tools ({len(tools)}):"]
|
||||
for t in tools:
|
||||
name = (t.get("function") or {}).get("name", "?")
|
||||
desc = (t.get("function") or {}).get("description", "")
|
||||
if len(desc) > 80:
|
||||
desc = desc[:77] + "..."
|
||||
lines.append(f" {name}: {desc}")
|
||||
return "\n".join(lines)
|
||||
except Exception as e:
|
||||
return f"Could not list tools: {e}"
|
||||
|
||||
def _cmd_context(self, args: str, state: SessionState) -> str:
|
||||
"""Show ACP session context pressure and compression guidance."""
|
||||
n_messages = len(state.history)
|
||||
roles = Counter(msg.get("role", "unknown") for msg in state.history)
|
||||
|
||||
agent = state.agent
|
||||
model = state.model or getattr(agent, "model", "")
|
||||
provider = getattr(agent, "provider", None) or "auto"
|
||||
compressor = getattr(agent, "context_compressor", None)
|
||||
context_length = int(getattr(compressor, "context_length", 0) or 0)
|
||||
threshold_tokens = int(getattr(compressor, "threshold_tokens", 0) or 0)
|
||||
|
||||
try:
|
||||
approx_tokens = _estimate_tokens(state.history, agent)
|
||||
except Exception:
|
||||
logger.debug("Could not estimate ACP context usage", exc_info=True)
|
||||
approx_tokens = 0
|
||||
|
||||
if threshold_tokens <= 0 and context_length > 0:
|
||||
threshold_tokens = int(context_length * 0.80)
|
||||
|
||||
lines = [
|
||||
f"Conversation: {n_messages} messages" if n_messages else "Conversation is empty (no messages yet).",
|
||||
f" user: {roles.get('user', 0)}, assistant: {roles.get('assistant', 0)}, "
|
||||
f"tool: {roles.get('tool', 0)}, system: {roles.get('system', 0)}",
|
||||
]
|
||||
if model:
|
||||
lines.append(f"Model: {model}")
|
||||
lines.append(f"Provider: {provider}")
|
||||
|
||||
if approx_tokens > 0 and context_length > 0:
|
||||
usage_pct = (approx_tokens / context_length) * 100
|
||||
lines.append(f"Context usage: ~{approx_tokens:,} / {context_length:,} tokens ({usage_pct:.1f}%)")
|
||||
elif approx_tokens > 0:
|
||||
lines.append(f"Context usage: ~{approx_tokens:,} tokens")
|
||||
|
||||
if threshold_tokens > 0 and approx_tokens > 0:
|
||||
threshold_pct = (threshold_tokens / context_length) * 100 if context_length > 0 else 0
|
||||
pct_note = f", {threshold_pct:.0f}%" if threshold_pct else ""
|
||||
if approx_tokens >= threshold_tokens:
|
||||
lines.append(f"Compression: due now (threshold ~{threshold_tokens:,}{pct_note}). Run /compress.")
|
||||
else:
|
||||
remaining = max(threshold_tokens - approx_tokens, 0)
|
||||
lines.append(f"Compression: ~{remaining:,} tokens until threshold (~{threshold_tokens:,}{pct_note}).")
|
||||
elif threshold_tokens > 0:
|
||||
lines.append(f"Compression threshold: ~{threshold_tokens:,} tokens")
|
||||
|
||||
lines.append(
|
||||
"Auto-compaction is disabled (compression.enabled: false); /compress still compresses manually."
|
||||
if getattr(agent, "compression_enabled", True) is False
|
||||
else "Tip: run /compress to compress manually before the threshold."
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
def _cmd_reset(self, args: str, state: SessionState) -> str:
|
||||
state.history.clear()
|
||||
try:
|
||||
reset_session_state = getattr(state.agent, "reset_session_state", None)
|
||||
if callable(reset_session_state):
|
||||
reset_session_state()
|
||||
except Exception:
|
||||
logger.warning("ACP session state reset failed for %s", state.session_id, exc_info=True)
|
||||
return "Conversation history cleared. Agent session state reset failed; see logs."
|
||||
finally:
|
||||
self.session_manager.save_session(state.session_id)
|
||||
return "Conversation history cleared."
|
||||
|
||||
def _cmd_compress(self, args: str, state: SessionState) -> str:
|
||||
if not state.history:
|
||||
return "Nothing to compress — conversation is empty."
|
||||
try:
|
||||
agent = state.agent
|
||||
# No compression_enabled gate: it only disables *automatic* compaction (CLI/gateway parity).
|
||||
if not hasattr(agent, "_compress_context"):
|
||||
return "Context compression not available for this agent."
|
||||
|
||||
original_count = len(state.history)
|
||||
# Include system prompt + tool schemas so the figure reflects real request pressure.
|
||||
# See #6217.
|
||||
# See #6217.
|
||||
_sys_prompt = getattr(agent, "_cached_system_prompt", "") or ""
|
||||
_tools = getattr(agent, "tools", None) or None
|
||||
approx_tokens = _estimate_tokens(state.history, agent, _sys_prompt, _tools)
|
||||
original_session_db = getattr(agent, "_session_db", None)
|
||||
|
||||
try:
|
||||
# Stable ACP session id: suppress _compress_context's SQLite session split.
|
||||
agent._session_db = None
|
||||
compressed, _ = agent._compress_context(
|
||||
state.history, _sys_prompt, approx_tokens=approx_tokens, task_id=state.session_id, force=True,
|
||||
)
|
||||
finally:
|
||||
agent._session_db = original_session_db
|
||||
|
||||
state.history = compressed
|
||||
self.session_manager.save_session(state.session_id)
|
||||
|
||||
new_tokens = _estimate_tokens(
|
||||
state.history, agent, getattr(agent, "_cached_system_prompt", "") or _sys_prompt,
|
||||
getattr(agent, "tools", None) or _tools,
|
||||
)
|
||||
return (
|
||||
f"Context compressed: {original_count} -> {len(state.history)} messages\n"
|
||||
f"~{approx_tokens:,} -> ~{new_tokens:,} tokens"
|
||||
)
|
||||
except Exception as e:
|
||||
return f"Compression failed: {e}"
|
||||
|
||||
def _cmd_steer(self, args: str, state: SessionState) -> str:
|
||||
steer_text = args.strip()
|
||||
if not steer_text:
|
||||
return "Usage: /steer <guidance>"
|
||||
|
||||
if state.is_running and hasattr(state.agent, "steer"):
|
||||
try:
|
||||
if state.agent.steer(steer_text):
|
||||
preview = steer_text[:80] + ("..." if len(steer_text) > 80 else "")
|
||||
return f"⏩ Steer queued for the active turn: {preview}"
|
||||
except Exception as exc:
|
||||
logger.warning("ACP steer failed for session %s: %s", state.session_id, exc)
|
||||
return f"⚠️ Steer failed: {exc}"
|
||||
|
||||
return f"No active turn — queued for the next turn. ({_queue_prompt(state, steer_text)} queued)"
|
||||
|
||||
def _cmd_queue(self, args: str, state: SessionState) -> str:
|
||||
queued_text = args.strip()
|
||||
if not queued_text:
|
||||
return "Usage: /queue <prompt>"
|
||||
return f"Queued for the next turn. ({_queue_prompt(state, queued_text)} queued)"
|
||||
|
||||
def _cmd_version(self, args: str, state: SessionState) -> str:
|
||||
return f"Hermes Agent v{HERMES_VERSION}"
|
||||
@@ -0,0 +1,275 @@
|
||||
"""ACP prompt content blocks -> Hermes/OpenAI user-content payloads (text, images, resources)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
from acp.schema import (
|
||||
AudioContentBlock, BlobResourceContents, EmbeddedResourceContentBlock, ImageContentBlock,
|
||||
ResourceContentBlock, TextContentBlock, TextResourceContents,
|
||||
)
|
||||
|
||||
logger = logging.getLogger("acp_adapter.server")
|
||||
|
||||
PromptBlock = (
|
||||
TextContentBlock | ImageContentBlock | AudioContentBlock | ResourceContentBlock | EmbeddedResourceContentBlock
|
||||
)
|
||||
|
||||
_MAX_ACP_RESOURCE_BYTES = 512 * 1024
|
||||
_TEXT_RESOURCE_MIME_TYPES = {
|
||||
"application/json",
|
||||
"application/javascript",
|
||||
"application/typescript",
|
||||
"application/xml",
|
||||
"application/x-yaml",
|
||||
"application/yaml",
|
||||
"application/toml",
|
||||
"application/sql",
|
||||
}
|
||||
|
||||
|
||||
def _resource_display_name(uri: str, name: str | None = None, title: str | None = None) -> str:
|
||||
"""Human-readable attachment name for prompt context."""
|
||||
raw_name = (name or "").strip()
|
||||
raw_title = (title or "").strip()
|
||||
if raw_title and raw_name and raw_title != raw_name:
|
||||
return f"{raw_title} ({raw_name})"
|
||||
if raw_title or raw_name:
|
||||
return raw_title or raw_name
|
||||
parsed = urlparse(uri)
|
||||
candidate = parsed.path if parsed.scheme else uri
|
||||
return Path(unquote(candidate)).name or uri or "resource"
|
||||
|
||||
|
||||
def _mime_main(mime_type: str | None) -> str:
|
||||
return (mime_type or "").split(";", 1)[0].strip().lower()
|
||||
|
||||
|
||||
def _is_text_resource(mime_type: str | None) -> bool:
|
||||
mime = _mime_main(mime_type)
|
||||
return mime.startswith("text/") or mime in _TEXT_RESOURCE_MIME_TYPES
|
||||
|
||||
|
||||
def _is_image_resource(mime_type: str | None) -> bool:
|
||||
return _mime_main(mime_type).startswith("image/")
|
||||
|
||||
|
||||
_IMAGE_SUFFIX_MIME = {
|
||||
".png": "image/png",
|
||||
".jpg": "image/jpeg",
|
||||
".jpeg": "image/jpeg",
|
||||
".gif": "image/gif",
|
||||
".webp": "image/webp",
|
||||
".bmp": "image/bmp",
|
||||
".svg": "image/svg+xml",
|
||||
}
|
||||
|
||||
|
||||
def _path_from_file_uri(uri: str) -> Path | None:
|
||||
"""Local file URI/path from an ACP client -> readable Path (None for non-file URIs).
|
||||
Windows drive forms (Zed via wsl.exe) become ``/mnt/<drive>/...``."""
|
||||
raw = (uri or "").strip()
|
||||
if not raw:
|
||||
return None
|
||||
|
||||
parsed = urlparse(raw)
|
||||
if parsed.scheme and parsed.scheme != "file":
|
||||
return None
|
||||
|
||||
if parsed.scheme == "file" and parsed.netloc and parsed.netloc not in {"", "localhost"}:
|
||||
return None
|
||||
path_text = unquote(parsed.path or "") if parsed.scheme == "file" else unquote(raw)
|
||||
|
||||
# file:///C:/Users/... or C:\Users\...
|
||||
if len(path_text) >= 3 and path_text[0] == "/" and path_text[2] == ":" and path_text[1].isalpha():
|
||||
drive, rest = path_text[1], path_text[3:]
|
||||
elif len(path_text) >= 2 and path_text[1] == ":" and path_text[0].isalpha():
|
||||
drive, rest = path_text[0], path_text[2:]
|
||||
else:
|
||||
return Path(path_text)
|
||||
return Path("/mnt") / drive.lower() / rest.lstrip("/\\").replace("\\", "/")
|
||||
|
||||
|
||||
def _decode_text_bytes(data: bytes, mime_type: str | None) -> str | None:
|
||||
"""Decode resource bytes if they are probably text; return None for binary."""
|
||||
if b"\x00" in data and not _is_text_resource(mime_type):
|
||||
return None
|
||||
for encoding in ("utf-8-sig", "utf-8", "latin-1"):
|
||||
try:
|
||||
return data.decode(encoding)
|
||||
except UnicodeDecodeError:
|
||||
continue
|
||||
# Binary (ELF/Mach-O/PE), not a shell script: feeding its decoded bytes back into the guard tokenizes
|
||||
# machine code into bogus NUL-bearing paths and crashes the scanner (#77703). Mirror
|
||||
# lifecycle_guard._read_referenced_script and treat it as nothing to scan.
|
||||
return data.decode("utf-8", errors="replace")
|
||||
|
||||
|
||||
def _format_resource_text(
|
||||
*, uri: str, body: str, name: str | None = None, title: str | None = None, note: str | None = None
|
||||
) -> str:
|
||||
display = _resource_display_name(uri, name=name, title=title)
|
||||
header = f"[Attached file: {display}]"
|
||||
if note:
|
||||
header += f" ({note})"
|
||||
return f"{header}\nURI: {uri}\n\n{body}"
|
||||
|
||||
|
||||
def _text_parts(**kwargs: Any) -> list[dict[str, Any]]:
|
||||
"""Single OpenAI text part wrapping ``_format_resource_text(**kwargs)``."""
|
||||
return [{"type": "text", "text": _format_resource_text(**kwargs)}]
|
||||
|
||||
|
||||
def _image_parts(uri: str, display: str, data: bytes, mime: str) -> list[dict[str, Any]]:
|
||||
"""Text header + image_url data URL so vision models can see the attachment."""
|
||||
return [
|
||||
{"type": "text", "text": f"[Attached image: {display}]" + (f"\nURI: {uri}" if uri else "")},
|
||||
{"type": "image_url", "image_url": {"url": f"data:{mime};base64,{base64.b64encode(data).decode('ascii')}"}},
|
||||
]
|
||||
|
||||
|
||||
def _attr(obj: Any, name: str) -> str | None:
|
||||
"""Stripped string attribute, ``None`` when missing/blank."""
|
||||
return str(getattr(obj, name, "") or "").strip() or None
|
||||
|
||||
|
||||
def _resource_link_to_parts(block: ResourceContentBlock) -> list[dict[str, Any]]:
|
||||
"""ACP resource_link -> OpenAI content parts: images become a text header + image_url,
|
||||
everything else a single text part with the inlined body (or a binary-omit note)."""
|
||||
uri = _attr(block, "uri")
|
||||
if not uri:
|
||||
return []
|
||||
|
||||
name, title, mime_type = _attr(block, "name"), _attr(block, "title"), _attr(block, "mime_type")
|
||||
path = _path_from_file_uri(uri)
|
||||
ident = dict(uri=uri, name=name, title=title)
|
||||
|
||||
if path is None:
|
||||
return _text_parts(
|
||||
**ident, body="[Resource link only; Hermes cannot read non-file ACP resource URIs directly.]"
|
||||
)
|
||||
|
||||
image_mime = mime_type if _is_image_resource(mime_type) else _IMAGE_SUFFIX_MIME.get(path.suffix.lower())
|
||||
if image_mime and _is_image_resource(image_mime):
|
||||
try:
|
||||
size = path.stat().st_size
|
||||
if size > _MAX_ACP_RESOURCE_BYTES:
|
||||
return _text_parts(
|
||||
**ident, body=f"[Image too large to inline: {size} bytes, cap={_MAX_ACP_RESOURCE_BYTES}]"
|
||||
)
|
||||
with path.open("rb") as fh:
|
||||
data = fh.read()
|
||||
except OSError as exc:
|
||||
logger.warning("ACP image resource read failed: %s", uri, exc_info=True)
|
||||
return _text_parts(**ident, body=f"[Could not read attached image: {exc}]")
|
||||
return _image_parts(uri, _resource_display_name(uri, name=name, title=title), data, image_mime)
|
||||
|
||||
try:
|
||||
size = path.stat().st_size
|
||||
with path.open("rb") as fh:
|
||||
data = fh.read(min(size, _MAX_ACP_RESOURCE_BYTES))
|
||||
text = _decode_text_bytes(data, mime_type)
|
||||
if text is None:
|
||||
return _text_parts(**ident, body=f"[Binary file omitted: {size} bytes, mime={mime_type or 'unknown'}]")
|
||||
note = f"truncated to {_MAX_ACP_RESOURCE_BYTES} of {size} bytes" if size > _MAX_ACP_RESOURCE_BYTES else None
|
||||
return _text_parts(**ident, body=text, note=note)
|
||||
except OSError as exc:
|
||||
logger.warning("ACP resource read failed: %s", uri, exc_info=True)
|
||||
return _text_parts(**ident, body=f"[Could not read attached file: {exc}]")
|
||||
|
||||
|
||||
def _embedded_resource_to_parts(block: EmbeddedResourceContentBlock) -> list[dict[str, Any]]:
|
||||
resource = getattr(block, "resource", None)
|
||||
if resource is None:
|
||||
return []
|
||||
|
||||
uri = _attr(resource, "uri") or ""
|
||||
mime_type = _attr(resource, "mime_type")
|
||||
|
||||
if isinstance(resource, TextResourceContents):
|
||||
return _text_parts(uri=uri, body=resource.text)
|
||||
|
||||
if isinstance(resource, BlobResourceContents):
|
||||
blob = resource.blob or ""
|
||||
try:
|
||||
data = base64.b64decode(blob, validate=True)
|
||||
except Exception:
|
||||
data = blob.encode("utf-8", errors="replace")
|
||||
|
||||
if _is_image_resource(mime_type):
|
||||
if len(data) > _MAX_ACP_RESOURCE_BYTES:
|
||||
return _text_parts(
|
||||
uri=uri,
|
||||
body=f"[Embedded image too large to inline: {len(data)} bytes, cap={_MAX_ACP_RESOURCE_BYTES}]",
|
||||
)
|
||||
return _image_parts(uri, _resource_display_name(uri), data, mime_type or "image/png")
|
||||
|
||||
body = _decode_text_bytes(data[:_MAX_ACP_RESOURCE_BYTES], mime_type)
|
||||
if body is None:
|
||||
body = f"[Binary embedded file omitted: {len(data)} bytes, mime={mime_type or 'unknown'}]"
|
||||
elif len(data) > _MAX_ACP_RESOURCE_BYTES:
|
||||
body += f"\n\n[Truncated to {_MAX_ACP_RESOURCE_BYTES} of {len(data)} bytes]"
|
||||
return _text_parts(uri=uri, body=body)
|
||||
|
||||
text = getattr(resource, "text", None)
|
||||
if text:
|
||||
return _text_parts(uri=uri, body=str(text))
|
||||
return []
|
||||
|
||||
|
||||
def _extract_text(prompt: list[PromptBlock]) -> str:
|
||||
"""Extract plain text from ACP content blocks for display/commands."""
|
||||
return "\n".join(str(block.text) for block in prompt if hasattr(block, "text"))
|
||||
|
||||
|
||||
def _image_block_to_openai_part(block: ImageContentBlock) -> dict[str, Any] | None:
|
||||
"""Convert an ACP image content block to OpenAI-style multimodal content."""
|
||||
data, uri = _attr(block, "data"), _attr(block, "uri")
|
||||
mime_type = _attr(block, "mime_type") or "image/png"
|
||||
if data:
|
||||
url = data if data.startswith("data:") else f"data:{mime_type};base64,{data}"
|
||||
elif uri:
|
||||
url = uri
|
||||
else:
|
||||
return None
|
||||
return {"type": "image_url", "image_url": {"url": url}}
|
||||
|
||||
|
||||
def _append_parts(parts: list, text_parts: list[str], new_parts: list[dict[str, Any]]) -> None:
|
||||
for part in new_parts:
|
||||
parts.append(part)
|
||||
if part.get("type") == "text":
|
||||
text_parts.append(part["text"])
|
||||
|
||||
|
||||
def _content_blocks_to_openai_user_content(prompt: list[PromptBlock]) -> str | list[dict[str, Any]]:
|
||||
"""Convert ACP prompt blocks into a Hermes/OpenAI-compatible user content payload."""
|
||||
parts: list[dict[str, Any]] = []
|
||||
text_parts: list[str] = []
|
||||
|
||||
for block in prompt:
|
||||
if isinstance(block, TextContentBlock):
|
||||
if block.text:
|
||||
parts.append({"type": "text", "text": block.text})
|
||||
text_parts.append(block.text)
|
||||
elif isinstance(block, ImageContentBlock):
|
||||
image_part = _image_block_to_openai_part(block)
|
||||
if image_part is not None:
|
||||
parts.append(image_part)
|
||||
elif isinstance(block, ResourceContentBlock):
|
||||
_append_parts(parts, text_parts, _resource_link_to_parts(block))
|
||||
elif isinstance(block, EmbeddedResourceContentBlock):
|
||||
_append_parts(parts, text_parts, _embedded_resource_to_parts(block))
|
||||
|
||||
if not parts:
|
||||
return _extract_text(prompt)
|
||||
|
||||
# Pure text stays a string (slash commands, text-only providers); structured only for media.
|
||||
if all(part.get("type") == "text" for part in parts):
|
||||
return "\n".join(text_parts)
|
||||
|
||||
return parts
|
||||
+85
-179
@@ -1,8 +1,8 @@
|
||||
"""Pre-execution ACP edit approval helpers.
|
||||
|
||||
This module is intentionally isolated from the generic tool registry. ACP binds
|
||||
an edit approval requester in a ContextVar for the duration of one ACP agent run;
|
||||
CLI, gateway, and other sessions leave it unset and therefore bypass this guard.
|
||||
Intentionally isolated from the generic tool registry: ACP binds an edit
|
||||
approval requester in a ContextVar for the duration of one ACP agent run; CLI,
|
||||
gateway, and other sessions leave it unset and therefore bypass this guard.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -12,7 +12,6 @@ import json
|
||||
import logging
|
||||
import re
|
||||
import tempfile
|
||||
from concurrent.futures import TimeoutError as FutureTimeout
|
||||
from contextvars import ContextVar, Token
|
||||
from dataclasses import dataclass
|
||||
from itertools import count
|
||||
@@ -35,75 +34,57 @@ class EditProposal:
|
||||
|
||||
EditApprovalRequester = Callable[[EditProposal], bool]
|
||||
|
||||
_EDIT_APPROVAL_REQUESTER: ContextVar[EditApprovalRequester | None] = ContextVar(
|
||||
"ACP_EDIT_APPROVAL_REQUESTER",
|
||||
default=None,
|
||||
)
|
||||
_EDIT_APPROVAL_REQUESTER: ContextVar[EditApprovalRequester | None] = ContextVar("ACP_EDIT_APPROVAL_REQUESTER", default=None)
|
||||
_PERMISSION_REQUEST_IDS = count(1)
|
||||
|
||||
|
||||
SENSITIVE_AUTO_APPROVE_NAMES = {".env", ".env.local", ".env.production", "id_rsa", "id_ed25519"}
|
||||
AUTO_APPROVE_ASK = "ask"
|
||||
AUTO_APPROVE_WORKSPACE = "workspace_session"
|
||||
AUTO_APPROVE_SESSION = "session"
|
||||
|
||||
_V4A_FILE_RE = re.compile(r'^\*\*\*\s+(?:Update|Add|Delete)\s+File:\s*(.+)$', re.MULTILINE)
|
||||
_V4A_MOVE_RE = re.compile(r'^\*\*\*\s+Move\s+File:\s*(.+?)\s*->\s*(.+)$', re.MULTILINE)
|
||||
|
||||
|
||||
def set_edit_approval_requester(requester: EditApprovalRequester | None) -> Token:
|
||||
"""Bind an ACP edit approval requester for the current context."""
|
||||
|
||||
return _EDIT_APPROVAL_REQUESTER.set(requester)
|
||||
|
||||
|
||||
def reset_edit_approval_requester(token: Token) -> None:
|
||||
"""Restore a previous edit approval requester binding."""
|
||||
|
||||
_EDIT_APPROVAL_REQUESTER.reset(token)
|
||||
|
||||
|
||||
def clear_edit_approval_requester() -> None:
|
||||
"""Clear the current requester; primarily used by tests."""
|
||||
|
||||
_EDIT_APPROVAL_REQUESTER.set(None)
|
||||
|
||||
|
||||
def get_edit_approval_requester() -> EditApprovalRequester | None:
|
||||
return _EDIT_APPROVAL_REQUESTER.get()
|
||||
|
||||
|
||||
def _read_text_if_exists(path: str) -> str | None:
|
||||
p = Path(path).expanduser()
|
||||
if not p.exists():
|
||||
return None
|
||||
if not p.is_file():
|
||||
if p.is_file():
|
||||
return p.read_text(encoding="utf-8", errors="replace")
|
||||
if p.exists():
|
||||
raise OSError(f"Cannot edit non-file path: {path}")
|
||||
return p.read_text(encoding="utf-8", errors="replace")
|
||||
return None
|
||||
|
||||
|
||||
def _required_path(arguments: dict[str, Any]) -> str:
|
||||
path = str(arguments.get("path") or "")
|
||||
if not path:
|
||||
raise ValueError("path required")
|
||||
return path
|
||||
|
||||
|
||||
def _proposal_for_write_file(arguments: dict[str, Any]) -> EditProposal:
|
||||
path = str(arguments.get("path") or "")
|
||||
if not path:
|
||||
raise ValueError("path required")
|
||||
path = _required_path(arguments)
|
||||
content = arguments.get("content")
|
||||
if content is None:
|
||||
raise ValueError("content required")
|
||||
return EditProposal(
|
||||
tool_name="write_file",
|
||||
path=path,
|
||||
old_text=_read_text_if_exists(path),
|
||||
new_text=str(content),
|
||||
arguments=dict(arguments),
|
||||
)
|
||||
return EditProposal("write_file", path, _read_text_if_exists(path), str(content), dict(arguments))
|
||||
|
||||
|
||||
def _proposal_for_patch_replace(arguments: dict[str, Any]) -> EditProposal:
|
||||
path = str(arguments.get("path") or "")
|
||||
if not path:
|
||||
raise ValueError("path required")
|
||||
old_string = arguments.get("old_string")
|
||||
new_string = arguments.get("new_string")
|
||||
path = _required_path(arguments)
|
||||
old_string, new_string = arguments.get("old_string"), arguments.get("new_string")
|
||||
if old_string is None or new_string is None:
|
||||
raise ValueError("old_string and new_string required")
|
||||
|
||||
old_text = _read_text_if_exists(path)
|
||||
if old_text is None:
|
||||
raise ValueError(f"Failed to read file: {path}")
|
||||
@@ -111,99 +92,58 @@ def _proposal_for_patch_replace(arguments: dict[str, Any]) -> EditProposal:
|
||||
from tools.fuzzy_match import fuzzy_find_and_replace
|
||||
|
||||
new_text, match_count, _strategy, error = fuzzy_find_and_replace(
|
||||
old_text,
|
||||
str(old_string),
|
||||
str(new_string),
|
||||
bool(arguments.get("replace_all", False)),
|
||||
)
|
||||
old_text, str(old_string), str(new_string), bool(arguments.get("replace_all", False)))
|
||||
if error or match_count == 0:
|
||||
raise ValueError(error or f"Could not find match for old_string in {path}")
|
||||
|
||||
return EditProposal(
|
||||
tool_name="patch",
|
||||
path=path,
|
||||
old_text=old_text,
|
||||
new_text=new_text,
|
||||
arguments=dict(arguments),
|
||||
)
|
||||
return EditProposal("patch", path, old_text, new_text, dict(arguments))
|
||||
|
||||
|
||||
def _extract_v4a_patch_paths(patch_body: str) -> list[str]:
|
||||
paths: list[str] = []
|
||||
for match in re.finditer(
|
||||
r'^\*\*\*\s+(?:Update|Add|Delete)\s+File:\s*(.+)$',
|
||||
patch_body,
|
||||
re.MULTILINE,
|
||||
):
|
||||
path = match.group(1).strip()
|
||||
if path:
|
||||
paths.append(path)
|
||||
for match in re.finditer(
|
||||
r'^\*\*\*\s+Move\s+File:\s*(.+?)\s*->\s*(.+)$',
|
||||
patch_body,
|
||||
re.MULTILINE,
|
||||
):
|
||||
src = match.group(1).strip()
|
||||
dst = match.group(2).strip()
|
||||
if src:
|
||||
paths.append(src)
|
||||
if dst:
|
||||
paths.append(dst)
|
||||
return paths
|
||||
paths = [m.group(1).strip() for m in _V4A_FILE_RE.finditer(patch_body)]
|
||||
for match in _V4A_MOVE_RE.finditer(patch_body):
|
||||
paths.extend(match.group(i).strip() for i in (1, 2))
|
||||
return [p for p in paths if p]
|
||||
|
||||
|
||||
def _proposal_for_patch_v4a(arguments: dict[str, Any]) -> EditProposal:
|
||||
patch_body = arguments.get("patch")
|
||||
if not isinstance(patch_body, str) or not patch_body:
|
||||
raise ValueError("patch content required")
|
||||
|
||||
paths = _extract_v4a_patch_paths(patch_body)
|
||||
if not paths:
|
||||
raise ValueError("no file paths found in V4A patch")
|
||||
|
||||
proposal_path = paths[0] if len(paths) == 1 else ", ".join(paths)
|
||||
old_text = _read_text_if_exists(paths[0]) if len(paths) == 1 else None
|
||||
single = len(paths) == 1
|
||||
# ACP only supports a single diff payload: surface the exact V4A patch as new_text so
|
||||
# patch-mode calls are permissioned and denied patches cannot mutate.
|
||||
return EditProposal(
|
||||
tool_name="patch",
|
||||
path=proposal_path,
|
||||
old_text=old_text,
|
||||
# ACP only supports a single diff payload here. Surface the exact V4A
|
||||
# patch content before execution so patch-mode calls are permissioned
|
||||
# and denied patches cannot mutate.
|
||||
new_text=patch_body,
|
||||
arguments=dict(arguments),
|
||||
"patch", paths[0] if single else ", ".join(paths),
|
||||
_read_text_if_exists(paths[0]) if single else None, patch_body, dict(arguments),
|
||||
)
|
||||
|
||||
|
||||
# (tool_name, patch mode or None) -> proposal builder.
|
||||
_PROPOSAL_BUILDERS = {
|
||||
("write_file", None): _proposal_for_write_file, ("patch", "replace"): _proposal_for_patch_replace,
|
||||
("patch", "patch"): _proposal_for_patch_v4a,
|
||||
}
|
||||
|
||||
|
||||
def build_edit_proposal(tool_name: str, arguments: dict[str, Any]) -> EditProposal | None:
|
||||
"""Return an edit proposal for supported file mutation calls."""
|
||||
|
||||
if tool_name == "write_file":
|
||||
return _proposal_for_write_file(arguments)
|
||||
if tool_name == "patch":
|
||||
mode = arguments.get("mode", "replace")
|
||||
if mode == "replace":
|
||||
return _proposal_for_patch_replace(arguments)
|
||||
if mode == "patch":
|
||||
return _proposal_for_patch_v4a(arguments)
|
||||
return None
|
||||
mode = arguments.get("mode", "replace") if tool_name == "patch" else None
|
||||
builder = _PROPOSAL_BUILDERS.get((tool_name, mode))
|
||||
return builder(arguments) if builder else None
|
||||
|
||||
|
||||
def _is_sensitive_auto_approve_path(path: str) -> bool:
|
||||
parts = Path(path).expanduser().parts
|
||||
lowered = {part.lower() for part in parts}
|
||||
if ".git" in lowered or ".ssh" in lowered:
|
||||
return True
|
||||
return Path(path).name.lower() in SENSITIVE_AUTO_APPROVE_NAMES
|
||||
lowered = {part.lower() for part in Path(path).expanduser().parts}
|
||||
return bool(lowered & {".git", ".ssh"}) or Path(path).name.lower() in SENSITIVE_AUTO_APPROVE_NAMES
|
||||
|
||||
|
||||
def should_auto_approve_edit(proposal: EditProposal, policy: str, cwd: str | None = None) -> bool:
|
||||
"""Return whether an ACP edit proposal may bypass the prompt for this session.
|
||||
|
||||
This is intentionally session-scoped and conservative: sensitive paths still
|
||||
ask even under autonomous policies.
|
||||
"""
|
||||
|
||||
Session-scoped and conservative: sensitive paths still ask under autonomous policies."""
|
||||
policy = str(policy or AUTO_APPROVE_ASK).strip()
|
||||
if policy == AUTO_APPROVE_ASK or _is_sensitive_auto_approve_path(proposal.path):
|
||||
return False
|
||||
@@ -211,90 +151,61 @@ def should_auto_approve_edit(proposal: EditProposal, policy: str, cwd: str | Non
|
||||
if policy == AUTO_APPROVE_SESSION:
|
||||
return True
|
||||
if policy == AUTO_APPROVE_WORKSPACE:
|
||||
# `/tmp` is the POSIX path but tempfile.gettempdir() is the real one on
|
||||
# every platform: `/private/tmp` on macOS (because `/tmp` is a symlink
|
||||
# and Path.resolve() follows it) and the per-user Temp dir on Windows.
|
||||
tmp_root = Path(tempfile.gettempdir()).resolve(strict=False)
|
||||
try:
|
||||
path.relative_to(tmp_root)
|
||||
return True
|
||||
except ValueError:
|
||||
pass
|
||||
if cwd:
|
||||
root = Path(cwd).expanduser().resolve(strict=False)
|
||||
try:
|
||||
path.relative_to(root)
|
||||
return True
|
||||
except ValueError:
|
||||
return False
|
||||
# tempfile.gettempdir() is the real temp root on every platform
|
||||
# (``/private/tmp`` on macOS since resolve() follows the symlink).
|
||||
return path.is_relative_to(Path(tempfile.gettempdir()).resolve(strict=False)) or (
|
||||
bool(cwd) and path.is_relative_to(Path(cwd).expanduser().resolve(strict=False)))
|
||||
return False
|
||||
|
||||
|
||||
def _denied(message: str) -> str:
|
||||
return json.dumps({"error": message}, ensure_ascii=False)
|
||||
|
||||
|
||||
def maybe_require_edit_approval(tool_name: str, arguments: dict[str, Any]) -> str | None:
|
||||
"""Run ACP edit approval if bound.
|
||||
|
||||
Returns a JSON tool-error string when the edit must be blocked, otherwise
|
||||
``None`` so dispatch can continue. Requester exceptions deny by default.
|
||||
"""
|
||||
|
||||
requester = get_edit_approval_requester()
|
||||
``None`` so dispatch can continue. Requester exceptions deny by default."""
|
||||
requester = _EDIT_APPROVAL_REQUESTER.get()
|
||||
if requester is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
proposal = build_edit_proposal(tool_name, arguments)
|
||||
except Exception as exc:
|
||||
logger.warning("Could not build ACP edit approval proposal for %s: %s", tool_name, exc)
|
||||
return json.dumps({"error": f"Edit approval denied: could not prepare diff ({exc})"}, ensure_ascii=False)
|
||||
|
||||
return _denied(f"Edit approval denied: could not prepare diff ({exc})")
|
||||
if proposal is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
approved = bool(requester(proposal))
|
||||
except Exception as exc:
|
||||
logger.warning("ACP edit approval requester failed: %s", exc)
|
||||
approved = False
|
||||
|
||||
if approved:
|
||||
return None
|
||||
return json.dumps({"error": "Edit approval denied by ACP client; file was not modified."}, ensure_ascii=False)
|
||||
return None if approved else _denied("Edit approval denied by ACP client; file was not modified.")
|
||||
|
||||
|
||||
def build_acp_edit_tool_call(proposal: EditProposal):
|
||||
"""Build the ToolCallUpdate payload for ACP request_permission."""
|
||||
|
||||
import acp
|
||||
|
||||
tool_call_id = f"edit-approval-{next(_PERMISSION_REQUEST_IDS)}"
|
||||
return acp.update_tool_call(
|
||||
tool_call_id,
|
||||
title=f"Approve edit: {proposal.path}",
|
||||
kind="edit",
|
||||
f"edit-approval-{next(_PERMISSION_REQUEST_IDS)}", title=f"Approve edit: {proposal.path}", kind="edit",
|
||||
status="pending",
|
||||
content=[
|
||||
acp.tool_diff_content(
|
||||
path=proposal.path,
|
||||
old_text=proposal.old_text,
|
||||
new_text=proposal.new_text,
|
||||
)
|
||||
],
|
||||
content=[acp.tool_diff_content(path=proposal.path, old_text=proposal.old_text, new_text=proposal.new_text)],
|
||||
raw_input={"tool": proposal.tool_name, "arguments": proposal.arguments},
|
||||
)
|
||||
|
||||
|
||||
def make_acp_edit_approval_requester(
|
||||
request_permission_fn: Callable,
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
session_id: str,
|
||||
timeout: float = 60.0,
|
||||
auto_approve_getter: Callable[[], tuple[str, str | None]] | None = None,
|
||||
request_permission_fn: Callable, loop: asyncio.AbstractEventLoop, session_id: str,
|
||||
timeout: float = 60.0, auto_approve_getter: Callable[[], tuple[str, str | None]] | None = None,
|
||||
) -> EditApprovalRequester:
|
||||
"""Return a sync requester that bridges edit proposals to ACP permissions."""
|
||||
|
||||
def _requester(proposal: EditProposal) -> bool:
|
||||
from acp.schema import PermissionOption
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
from acp_adapter.permissions import await_permission
|
||||
|
||||
if auto_approve_getter is not None:
|
||||
try:
|
||||
@@ -305,34 +216,29 @@ def make_acp_edit_approval_requester(
|
||||
except Exception:
|
||||
logger.debug("ACP edit auto-approval policy check failed", exc_info=True)
|
||||
|
||||
options = [
|
||||
PermissionOption(option_id="allow_once", kind="allow_once", name="Allow edit"),
|
||||
PermissionOption(option_id="deny", kind="reject_once", name="Deny"),
|
||||
]
|
||||
tool_call = build_acp_edit_tool_call(proposal)
|
||||
coro = request_permission_fn(
|
||||
session_id=session_id,
|
||||
tool_call=tool_call,
|
||||
options=options,
|
||||
response, _timed_out = await_permission(
|
||||
request_permission_fn, loop, session_id, tool_call=build_acp_edit_tool_call(proposal),
|
||||
options=[PermissionOption(option_id="allow_once", kind="allow_once", name="Allow edit"),
|
||||
PermissionOption(option_id="deny", kind="reject_once", name="Deny")],
|
||||
timeout=timeout, what="Edit approval request",
|
||||
)
|
||||
future = safe_schedule_threadsafe(
|
||||
coro,
|
||||
loop,
|
||||
logger=logger,
|
||||
log_message="Edit approval request: failed to schedule on loop",
|
||||
)
|
||||
if future is None:
|
||||
return False
|
||||
try:
|
||||
response = future.result(timeout=timeout)
|
||||
except (FutureTimeout, Exception) as exc:
|
||||
future.cancel()
|
||||
logger.warning("Edit approval request timed out or failed: %s", exc)
|
||||
return False
|
||||
outcome = getattr(response, "outcome", None)
|
||||
return (
|
||||
getattr(outcome, "outcome", None) == "selected"
|
||||
and getattr(outcome, "option_id", None) == "allow_once"
|
||||
)
|
||||
return getattr(outcome, "outcome", None) == "selected" and getattr(outcome, "option_id", None) == "allow_once"
|
||||
|
||||
return _requester
|
||||
|
||||
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
from concurrent.futures import TimeoutError as FutureTimeout # noqa: F401,E402
|
||||
|
||||
def clear_edit_approval_requester() -> None:
|
||||
"""Clear the current requester; primarily used by tests."""
|
||||
|
||||
_EDIT_APPROVAL_REQUESTER.set(None)
|
||||
|
||||
def get_edit_approval_requester() -> EditApprovalRequester | None:
|
||||
return _EDIT_APPROVAL_REQUESTER.get()
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
+63
-135
@@ -1,16 +1,11 @@
|
||||
"""CLI entry point for the hermes-agent ACP adapter.
|
||||
|
||||
Loads environment variables from ``~/.hermes/.env``, configures logging
|
||||
to write to stderr (so stdout is reserved for ACP JSON-RPC transport),
|
||||
and starts the ACP agent server.
|
||||
Loads ``~/.hermes/.env``, routes logging to stderr (stdout is reserved for ACP
|
||||
JSON-RPC), and starts the ACP agent server.
|
||||
|
||||
Usage::
|
||||
|
||||
python -m acp_adapter.entry
|
||||
# or
|
||||
hermes acp
|
||||
# or
|
||||
hermes-acp
|
||||
python -m acp_adapter.entry # or: hermes acp / hermes-acp
|
||||
"""
|
||||
|
||||
# IMPORTANT: hermes_bootstrap must be the very first import — UTF-8 stdio
|
||||
@@ -18,15 +13,11 @@ Usage::
|
||||
try:
|
||||
import hermes_bootstrap # noqa: F401
|
||||
except ModuleNotFoundError:
|
||||
# Graceful fallback when hermes_bootstrap isn't registered in the venv
|
||||
# yet — happens during partial ``hermes update`` where git-reset landed
|
||||
# new code but ``uv pip install -e .`` didn't finish. Missing bootstrap
|
||||
# means UTF-8 stdio setup is skipped on Windows; POSIX is unaffected.
|
||||
# Partial ``hermes update`` (git-reset landed, ``uv pip install -e .`` did not):
|
||||
# UTF-8 stdio setup is skipped on Windows; POSIX is unaffected.
|
||||
pass
|
||||
else:
|
||||
# Stop a ``utils/``/``proxy/``/``ui/`` package in the launch directory from
|
||||
# shadowing Hermes's own modules — ``hermes acp`` can be started from any
|
||||
# cwd, including a project that has same-named packages on its path.
|
||||
# Stop a ``utils/``/``proxy/``/``ui/`` package in the launch cwd from shadowing Hermes modules.
|
||||
hermes_bootstrap.harden_import_path()
|
||||
|
||||
import argparse
|
||||
@@ -38,44 +29,30 @@ from pathlib import Path
|
||||
from hermes_constants import get_hermes_home
|
||||
|
||||
|
||||
# Methods clients send as periodic liveness probes. They are not part of the
|
||||
# ACP schema, so the acp router correctly returns JSON-RPC -32601 to the
|
||||
# caller — but the supervisor task that dispatches the request then surfaces
|
||||
# the raised RequestError via ``logging.exception("Background task failed")``,
|
||||
# which dumps a traceback to stderr every probe interval. Clients like
|
||||
# acp-bridge already treat the -32601 response as "agent alive", so the
|
||||
# traceback is pure noise. We keep the protocol response intact and only
|
||||
# silence the stderr noise for this specific benign case.
|
||||
# Liveness-probe methods outside the ACP schema. The router correctly answers JSON-RPC -32601
|
||||
# (clients treat that as "agent alive"), but the dispatching supervisor task also logs
|
||||
# ``"Background task failed"`` with a traceback every probe. Keep the response; silence the noise.
|
||||
_BENIGN_PROBE_METHODS = frozenset({"ping", "health", "healthcheck"})
|
||||
|
||||
|
||||
class _BenignProbeMethodFilter(logging.Filter):
|
||||
"""Suppress acp 'Background task failed' tracebacks caused by unknown
|
||||
liveness-probe methods (e.g. ``ping``) while leaving every other
|
||||
background-task error — including method_not_found for any non-probe
|
||||
method — visible in stderr.
|
||||
"""
|
||||
"""Suppress acp 'Background task failed' tracebacks caused by unknown liveness-probe methods
|
||||
(e.g. ``ping``); every other background-task error, incl. method_not_found for non-probe
|
||||
methods, stays visible."""
|
||||
|
||||
def filter(self, record: logging.LogRecord) -> bool:
|
||||
if record.getMessage() != "Background task failed":
|
||||
if record.getMessage() != "Background task failed" or not record.exc_info:
|
||||
return True
|
||||
exc_info = record.exc_info
|
||||
if not exc_info:
|
||||
return True
|
||||
exc = exc_info[1]
|
||||
# Imported lazily so this module stays importable when the optional
|
||||
# ``agent-client-protocol`` dependency is not installed.
|
||||
# Lazy import keeps this module importable without ``agent-client-protocol``.
|
||||
try:
|
||||
from acp.exceptions import RequestError
|
||||
except ImportError:
|
||||
return True
|
||||
if not isinstance(exc, RequestError):
|
||||
return True
|
||||
if getattr(exc, "code", None) != -32601:
|
||||
exc = record.exc_info[1]
|
||||
if not isinstance(exc, RequestError) or getattr(exc, "code", None) != -32601:
|
||||
return True
|
||||
data = getattr(exc, "data", None)
|
||||
method = data.get("method") if isinstance(data, dict) else None
|
||||
return method not in _BENIGN_PROBE_METHODS
|
||||
return not (isinstance(data, dict) and data.get("method") in _BENIGN_PROBE_METHODS)
|
||||
|
||||
|
||||
def _setup_logging() -> None:
|
||||
@@ -83,22 +60,15 @@ def _setup_logging() -> None:
|
||||
from agent.redact import RedactingFormatter
|
||||
|
||||
handler = logging.StreamHandler(sys.stderr)
|
||||
handler.setFormatter(
|
||||
RedactingFormatter(
|
||||
"%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
)
|
||||
handler.setFormatter(RedactingFormatter("%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S"))
|
||||
handler.addFilter(_BenignProbeMethodFilter())
|
||||
root = logging.getLogger()
|
||||
root.handlers.clear()
|
||||
root.addHandler(handler)
|
||||
root.setLevel(logging.INFO)
|
||||
|
||||
# Quiet down noisy libraries
|
||||
logging.getLogger("httpx").setLevel(logging.WARNING)
|
||||
logging.getLogger("httpcore").setLevel(logging.WARNING)
|
||||
logging.getLogger("openai").setLevel(logging.WARNING)
|
||||
for noisy in ("httpx", "httpcore", "openai"):
|
||||
logging.getLogger(noisy).setLevel(logging.WARNING)
|
||||
|
||||
|
||||
def _load_env() -> None:
|
||||
@@ -107,45 +77,25 @@ def _load_env() -> None:
|
||||
|
||||
hermes_home = get_hermes_home()
|
||||
loaded = load_hermes_dotenv(hermes_home=hermes_home)
|
||||
if loaded:
|
||||
for env_file in loaded:
|
||||
logging.getLogger(__name__).info("Loaded env from %s", env_file)
|
||||
else:
|
||||
logging.getLogger(__name__).info(
|
||||
"No .env found at %s, using system env", hermes_home / ".env"
|
||||
)
|
||||
log = logging.getLogger(__name__)
|
||||
for env_file in loaded or ():
|
||||
log.info("Loaded env from %s", env_file)
|
||||
if not loaded:
|
||||
log.info("No .env found at %s, using system env", hermes_home / ".env")
|
||||
|
||||
|
||||
def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="hermes-acp",
|
||||
description="Run Hermes Agent as an ACP stdio server.",
|
||||
)
|
||||
parser = argparse.ArgumentParser(prog="hermes-acp", description="Run Hermes Agent as an ACP stdio server.")
|
||||
parser.add_argument("--version", action="store_true", help="Print Hermes version and exit")
|
||||
parser.add_argument(
|
||||
"--check",
|
||||
action="store_true",
|
||||
help="Verify ACP dependencies and adapter imports, then exit",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--setup",
|
||||
action="store_true",
|
||||
help="Run interactive Hermes provider/model setup for ACP terminal auth",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--setup-browser",
|
||||
action="store_true",
|
||||
help="Install agent-browser + Playwright Chromium into ~/.hermes/node/ "
|
||||
"for browser tool support. Idempotent.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--yes",
|
||||
"-y",
|
||||
action="store_true",
|
||||
dest="assume_yes",
|
||||
help="Accept all prompts (currently used by --setup-browser to skip the "
|
||||
"~400 MB Chromium download confirmation).",
|
||||
)
|
||||
parser.add_argument("--check", action="store_true", help="Verify ACP dependencies and adapter imports, then exit")
|
||||
parser.add_argument("--setup", action="store_true",
|
||||
help="Run interactive Hermes provider/model setup for ACP terminal auth")
|
||||
parser.add_argument("--setup-browser", action="store_true",
|
||||
help="Install agent-browser + Playwright Chromium into ~/.hermes/node/ "
|
||||
"for browser tool support. Idempotent.")
|
||||
parser.add_argument("--yes", "-y", action="store_true", dest="assume_yes",
|
||||
help="Accept all prompts (currently used by --setup-browser to skip the "
|
||||
"~400 MB Chromium download confirmation).")
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
@@ -172,45 +122,35 @@ def _run_setup() -> None:
|
||||
finally:
|
||||
sys.argv = old_argv
|
||||
|
||||
# Offer browser-tools install as a follow-up. The terminal auth method
|
||||
# is the one supported first-run UX for registry installs, so this is
|
||||
# the natural moment to ask. Skip silently if stdin isn't a TTY (the
|
||||
# answer can't be collected anyway).
|
||||
# Terminal auth is the first-run UX for registry installs, so offer the browser-tools
|
||||
# install here. Skip silently without a TTY.
|
||||
if not sys.stdin.isatty():
|
||||
return
|
||||
try:
|
||||
reply = input(
|
||||
"\nInstall browser tools? Downloads agent-browser (npm) and "
|
||||
"optionally Playwright Chromium (~400 MB). [y/N] "
|
||||
).strip().lower()
|
||||
reply = input("\nInstall browser tools? Downloads agent-browser (npm) and "
|
||||
"optionally Playwright Chromium (~400 MB). [y/N] ").strip().lower()
|
||||
except (EOFError, KeyboardInterrupt):
|
||||
return
|
||||
if reply in {"y", "yes"}:
|
||||
_run_setup_browser(assume_yes=False)
|
||||
|
||||
|
||||
_SETUP_BROWSER_STEPS = (
|
||||
("node", "Node.js installation failed — cannot proceed with browser tools."),
|
||||
("browser", "Browser tools installation failed."),
|
||||
)
|
||||
|
||||
|
||||
def _run_setup_browser(assume_yes: bool = False) -> int:
|
||||
"""Bootstrap agent-browser + Chromium.
|
||||
|
||||
Routes through dep_ensure -> install.{sh,ps1} --ensure, sharing code
|
||||
with the runtime lazy installer.
|
||||
|
||||
Returns 0 on success, 1 on failure.
|
||||
"""
|
||||
"""Bootstrap agent-browser + Chromium via dep_ensure -> install.{sh,ps1}
|
||||
--ensure (shared with the runtime lazy installer). Returns 0 on success, 1 on failure."""
|
||||
from hermes_cli.dep_ensure import ensure_dependency
|
||||
|
||||
try:
|
||||
node_ok = ensure_dependency("node", interactive=not assume_yes)
|
||||
if not node_ok:
|
||||
print("Node.js installation failed — cannot proceed with browser tools.",
|
||||
file=sys.stderr)
|
||||
return 1
|
||||
|
||||
browser_ok = ensure_dependency("browser", interactive=not assume_yes)
|
||||
if not browser_ok:
|
||||
print("Browser tools installation failed.", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
for dep, failure_msg in _SETUP_BROWSER_STEPS:
|
||||
if not ensure_dependency(dep, interactive=not assume_yes):
|
||||
print(failure_msg, file=sys.stderr)
|
||||
return 1
|
||||
return 0
|
||||
except OSError as exc:
|
||||
print(f"Browser bootstrap failed: {exc}", file=sys.stderr)
|
||||
@@ -220,18 +160,11 @@ def _run_setup_browser(assume_yes: bool = False) -> int:
|
||||
def main(argv: list[str] | None = None) -> None:
|
||||
"""Entry point: load env, configure logging, run the ACP agent."""
|
||||
args = _parse_args(argv)
|
||||
if args.version:
|
||||
_print_version()
|
||||
return
|
||||
if args.check:
|
||||
_run_check()
|
||||
return
|
||||
if args.setup:
|
||||
_run_setup()
|
||||
return
|
||||
for flag, action in (("version", _print_version), ("check", _run_check), ("setup", _run_setup)):
|
||||
if getattr(args, flag):
|
||||
return action()
|
||||
if args.setup_browser:
|
||||
rc = _run_setup_browser(assume_yes=args.assume_yes)
|
||||
if rc != 0:
|
||||
if rc := _run_setup_browser(assume_yes=args.assume_yes):
|
||||
sys.exit(rc)
|
||||
return
|
||||
|
||||
@@ -249,22 +182,17 @@ def main(argv: list[str] | None = None) -> None:
|
||||
import acp
|
||||
from .server import HermesACPAgent
|
||||
|
||||
# MCP tool discovery from config.yaml — fire-and-forget in a
|
||||
# background daemon thread so the ACP server becomes responsive
|
||||
# immediately while MCP servers connect. Previously this blocked
|
||||
# asyncio.run() for 2-5 s. (ACP also registers per-session MCP
|
||||
# servers dynamically via asyncio.to_thread inside the event loop;
|
||||
# that path is unaffected.) Moved from model_tools.py module scope
|
||||
# to avoid freezing the gateway's loop on lazy import (#16856).
|
||||
# Metadata-only hosts can opt out of unrelated global MCP startup.
|
||||
# MCP discovery from config.yaml runs in a background daemon thread so the ACP server is
|
||||
# responsive immediately (blocking here cost 2-5 s); per-session MCP servers registered via
|
||||
# asyncio.to_thread are unaffected. Metadata-only hosts can opt out of the global startup.
|
||||
# Previously this blocked asyncio.run() for 2-5 s. (ACP also registers per-session MCP servers
|
||||
# dynamically via asyncio.to_thread inside the event loop; that path is unaffected.) Moved from
|
||||
# model_tools.py module scope to avoid freezing the gateway's loop on lazy import (#16856).
|
||||
if os.environ.get("HERMES_ACP_SKIP_CONFIGURED_MCP", "").strip() != "1":
|
||||
try:
|
||||
from hermes_cli.mcp_startup import start_background_mcp_discovery
|
||||
|
||||
start_background_mcp_discovery(
|
||||
logger=logger,
|
||||
thread_name="acp-mcp-discovery",
|
||||
)
|
||||
start_background_mcp_discovery(logger=logger, thread_name="acp-mcp-discovery")
|
||||
except Exception:
|
||||
logger.debug("MCP tool discovery failed at ACP startup", exc_info=True)
|
||||
|
||||
|
||||
+77
-172
@@ -1,14 +1,12 @@
|
||||
"""Callback factories for bridging AIAgent events to ACP notifications.
|
||||
|
||||
Each factory returns a callable with the signature that AIAgent expects
|
||||
for its callbacks. Internally, the callbacks push ACP session updates
|
||||
to the client via ``conn.session_update()`` using
|
||||
``asyncio.run_coroutine_threadsafe()`` (since AIAgent runs in a worker
|
||||
thread while the event loop lives on the main thread).
|
||||
Each factory returns a callable with the signature AIAgent expects for its
|
||||
callbacks. AIAgent runs in a worker thread while the event loop lives on the
|
||||
main thread, so updates are pushed via ``conn.session_update()`` scheduled
|
||||
thread-safely onto the loop.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from collections import deque
|
||||
from typing import Any, Callable, Deque, Dict
|
||||
@@ -16,88 +14,46 @@ from typing import Any, Callable, Deque, Dict
|
||||
import acp
|
||||
from acp.schema import AgentPlanUpdate, PlanEntry
|
||||
|
||||
from .tools import (
|
||||
build_tool_complete,
|
||||
build_tool_start,
|
||||
make_tool_call_id,
|
||||
)
|
||||
from .tools import _json_loads_maybe, build_tool_complete, build_tool_start, coerce_tool_args, make_tool_call_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _json_loads_maybe_prefix(value: str) -> Any:
|
||||
"""Parse a JSON object even when Hermes appended a human hint after it."""
|
||||
text = value.strip()
|
||||
try:
|
||||
return json.loads(text)
|
||||
except Exception:
|
||||
decoder = json.JSONDecoder()
|
||||
data, _ = decoder.raw_decode(text)
|
||||
return data
|
||||
# ACP plans only support pending/in_progress/completed. Cancelled tasks are kept
|
||||
# as terminal entries so the client's full-list replacement doesn't drop them.
|
||||
_PLAN_STATUS = {"pending": "pending", "in_progress": "in_progress", "completed": "completed", "cancelled": "completed"}
|
||||
|
||||
|
||||
def _build_plan_update_from_todo_result(result: Any) -> AgentPlanUpdate | None:
|
||||
"""Translate Hermes' todo tool result into ACP's native plan update.
|
||||
|
||||
Zed renders ``sessionUpdate: plan`` as its first-class task/todo panel. The
|
||||
Hermes agent already maintains task state through the ``todo`` tool, so the
|
||||
ACP adapter should expose that state natively instead of only as a generic
|
||||
tool-call transcript block.
|
||||
"""
|
||||
Zed renders ``sessionUpdate: plan`` as its first-class task panel, so the
|
||||
todo state is exposed natively rather than only as a tool-call transcript."""
|
||||
if not isinstance(result, str) or not result.strip():
|
||||
return None
|
||||
|
||||
try:
|
||||
data = _json_loads_maybe_prefix(result)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
data = _json_loads_maybe(result)
|
||||
if not isinstance(data, dict) or not isinstance(data.get("todos"), list):
|
||||
return None
|
||||
|
||||
todos = data["todos"]
|
||||
if not todos:
|
||||
return AgentPlanUpdate(session_update="plan", entries=[])
|
||||
|
||||
status_map = {
|
||||
"pending": "pending",
|
||||
"in_progress": "in_progress",
|
||||
"completed": "completed",
|
||||
# ACP plans only support pending/in_progress/completed. Preserve
|
||||
# cancelled tasks as terminal entries instead of dropping them and
|
||||
# making the client's full-list replacement lose visible context.
|
||||
"cancelled": "completed",
|
||||
}
|
||||
entries: list[PlanEntry] = []
|
||||
for item in todos:
|
||||
for item in data["todos"]:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
content = str(item.get("content") or item.get("id") or "").strip()
|
||||
if not content:
|
||||
continue
|
||||
raw_status = str(item.get("status") or "pending").strip()
|
||||
status = status_map.get(raw_status, "pending")
|
||||
if raw_status == "cancelled":
|
||||
content = f"[cancelled] {content}"
|
||||
entries.append(PlanEntry(content=content, priority="medium", status=status))
|
||||
|
||||
entries.append(PlanEntry(content=content, priority="medium", status=_PLAN_STATUS.get(raw_status, "pending")))
|
||||
return AgentPlanUpdate(session_update="plan", entries=entries)
|
||||
|
||||
|
||||
def _send_update(
|
||||
conn: acp.Client,
|
||||
session_id: str,
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
update: Any,
|
||||
) -> None:
|
||||
def _send_update(conn: acp.Client, session_id: str, loop: asyncio.AbstractEventLoop, update: Any) -> None:
|
||||
"""Fire-and-forget an ACP session update from a worker thread."""
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
|
||||
future = safe_schedule_threadsafe(
|
||||
conn.session_update(session_id, update),
|
||||
loop,
|
||||
logger=logger,
|
||||
log_message="Failed to send ACP update",
|
||||
conn.session_update(session_id, update), loop, logger=logger, log_message="Failed to send ACP update",
|
||||
)
|
||||
if future is None:
|
||||
return
|
||||
@@ -107,50 +63,34 @@ def _send_update(
|
||||
logger.debug("Failed to send ACP update", exc_info=True)
|
||||
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Tool progress callback
|
||||
# ------------------------------------------------------------------
|
||||
def _upgrade_queue(tool_call_ids: Dict[str, Deque[str]], name: str) -> Deque[str] | None:
|
||||
"""Fetch the per-tool FIFO of pending call IDs, upgrading a legacy bare-string entry in place."""
|
||||
queue = tool_call_ids.get(name)
|
||||
if isinstance(queue, str):
|
||||
queue = tool_call_ids[name] = deque([queue])
|
||||
return queue
|
||||
|
||||
|
||||
def make_tool_progress_cb(
|
||||
conn: acp.Client,
|
||||
session_id: str,
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
tool_call_ids: Dict[str, Deque[str]],
|
||||
conn: acp.Client, session_id: str, loop: asyncio.AbstractEventLoop, tool_call_ids: Dict[str, Deque[str]],
|
||||
tool_call_meta: Dict[str, Dict[str, Any]],
|
||||
edit_approval_policy_getter: Callable[[], tuple[str, str | None]] | None = None,
|
||||
) -> Callable:
|
||||
"""Create a ``tool_progress_callback`` for AIAgent.
|
||||
|
||||
Signature expected by AIAgent::
|
||||
|
||||
tool_progress_callback(event_type: str, name: str, preview: str, args: dict, **kwargs)
|
||||
|
||||
Emits ``ToolCallStart`` for ``tool.started`` events and tracks IDs in a FIFO
|
||||
queue per tool name so duplicate/parallel same-name calls still complete
|
||||
against the correct ACP tool call. Other event types (``tool.completed``,
|
||||
``reasoning.available``) are silently ignored.
|
||||
"""
|
||||
Signature: ``tool_progress_callback(event_type, name, preview, args, **kwargs)``.
|
||||
Emits ``ToolCallStart`` for ``tool.started`` and tracks IDs in a FIFO per tool
|
||||
name so parallel same-name calls complete against the right ACP tool call.
|
||||
Other event types (``tool.completed``, ``reasoning.available``) are ignored."""
|
||||
|
||||
def _tool_progress(event_type: str, name: str = None, preview: str = None, args: Any = None, **kwargs) -> None:
|
||||
# Only emit ACP ToolCallStart for tool.started; ignore other event types
|
||||
if event_type != "tool.started":
|
||||
return
|
||||
if isinstance(args, str):
|
||||
try:
|
||||
args = json.loads(args)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
args = {"raw": args}
|
||||
if not isinstance(args, dict):
|
||||
args = {}
|
||||
|
||||
args = coerce_tool_args(args)
|
||||
tc_id = make_tool_call_id()
|
||||
queue = tool_call_ids.get(name)
|
||||
queue = _upgrade_queue(tool_call_ids, name)
|
||||
if queue is None:
|
||||
queue = deque()
|
||||
tool_call_ids[name] = queue
|
||||
elif isinstance(queue, str):
|
||||
queue = deque([queue])
|
||||
tool_call_ids[name] = queue
|
||||
queue = tool_call_ids[name] = deque()
|
||||
queue.append(tc_id)
|
||||
|
||||
snapshot = None
|
||||
@@ -176,104 +116,69 @@ def make_tool_progress_cb(
|
||||
except Exception:
|
||||
logger.debug("Failed to prepare auto-approved ACP edit diff for %s", name, exc_info=True)
|
||||
|
||||
update = build_tool_start(tc_id, name, args, edit_diff=edit_diff)
|
||||
_send_update(conn, session_id, loop, update)
|
||||
_send_update(conn, session_id, loop, build_tool_start(tc_id, name, args, edit_diff=edit_diff))
|
||||
|
||||
return _tool_progress
|
||||
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Thinking callback
|
||||
# ------------------------------------------------------------------
|
||||
def _make_text_cb(conn: acp.Client, session_id: str, loop: asyncio.AbstractEventLoop, wrap: Callable[[str], Any]) -> Callable:
|
||||
def _cb(text: str) -> None:
|
||||
if text:
|
||||
_send_update(conn, session_id, loop, wrap(text))
|
||||
|
||||
def make_thinking_cb(
|
||||
conn: acp.Client,
|
||||
session_id: str,
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
) -> Callable:
|
||||
return _cb
|
||||
|
||||
|
||||
def make_thinking_cb(conn: acp.Client, session_id: str, loop: asyncio.AbstractEventLoop) -> Callable:
|
||||
"""Create a ``thinking_callback`` for AIAgent."""
|
||||
|
||||
def _thinking(text: str) -> None:
|
||||
if not text:
|
||||
return
|
||||
update = acp.update_agent_thought_text(text)
|
||||
_send_update(conn, session_id, loop, update)
|
||||
|
||||
return _thinking
|
||||
return _make_text_cb(conn, session_id, loop, acp.update_agent_thought_text)
|
||||
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Step callback
|
||||
# ------------------------------------------------------------------
|
||||
def make_message_cb(conn: acp.Client, session_id: str, loop: asyncio.AbstractEventLoop) -> Callable:
|
||||
"""Create a callback that streams agent response text to the editor."""
|
||||
return _make_text_cb(conn, session_id, loop, acp.update_agent_message_text)
|
||||
|
||||
|
||||
def make_step_cb(
|
||||
conn: acp.Client,
|
||||
session_id: str,
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
tool_call_ids: Dict[str, Deque[str]],
|
||||
conn: acp.Client, session_id: str, loop: asyncio.AbstractEventLoop, tool_call_ids: Dict[str, Deque[str]],
|
||||
tool_call_meta: Dict[str, Dict[str, Any]],
|
||||
) -> Callable:
|
||||
"""Create a ``step_callback`` for AIAgent.
|
||||
|
||||
Signature expected by AIAgent::
|
||||
|
||||
step_callback(api_call_count: int, prev_tools: list)
|
||||
"""
|
||||
"""Create a ``step_callback(api_call_count: int, prev_tools: list)`` for AIAgent."""
|
||||
|
||||
def _step(api_call_count: int, prev_tools: Any = None) -> None:
|
||||
if prev_tools and isinstance(prev_tools, list):
|
||||
for tool_info in prev_tools:
|
||||
tool_name = None
|
||||
result = None
|
||||
function_args = None
|
||||
if not isinstance(prev_tools, list):
|
||||
return
|
||||
for tool_info in prev_tools:
|
||||
tool_name = result = function_args = None
|
||||
if isinstance(tool_info, dict):
|
||||
tool_name = tool_info.get("name") or tool_info.get("function_name")
|
||||
result = tool_info.get("result") or tool_info.get("output")
|
||||
function_args = tool_info.get("arguments") or tool_info.get("args")
|
||||
elif isinstance(tool_info, str):
|
||||
tool_name = tool_info
|
||||
|
||||
if isinstance(tool_info, dict):
|
||||
tool_name = tool_info.get("name") or tool_info.get("function_name")
|
||||
result = tool_info.get("result") or tool_info.get("output")
|
||||
function_args = tool_info.get("arguments") or tool_info.get("args")
|
||||
elif isinstance(tool_info, str):
|
||||
tool_name = tool_info
|
||||
|
||||
queue = tool_call_ids.get(tool_name or "")
|
||||
if isinstance(queue, str):
|
||||
queue = deque([queue])
|
||||
tool_call_ids[tool_name] = queue
|
||||
if tool_name and queue:
|
||||
tc_id = queue.popleft()
|
||||
meta = tool_call_meta.pop(tc_id, {})
|
||||
update = build_tool_complete(
|
||||
tc_id,
|
||||
tool_name,
|
||||
result=str(result) if result is not None else None,
|
||||
function_args=function_args or meta.get("args"),
|
||||
snapshot=meta.get("snapshot"),
|
||||
)
|
||||
_send_update(conn, session_id, loop, update)
|
||||
if tool_name == "todo":
|
||||
plan_update = _build_plan_update_from_todo_result(result)
|
||||
if plan_update is not None:
|
||||
_send_update(conn, session_id, loop, plan_update)
|
||||
if not queue:
|
||||
tool_call_ids.pop(tool_name, None)
|
||||
if not tool_name:
|
||||
continue
|
||||
queue = _upgrade_queue(tool_call_ids, tool_name)
|
||||
if not queue:
|
||||
continue
|
||||
tc_id = queue.popleft()
|
||||
meta = tool_call_meta.pop(tc_id, {})
|
||||
_send_update(conn, session_id, loop, build_tool_complete(
|
||||
tc_id, tool_name, result=str(result) if result is not None else None,
|
||||
function_args=function_args or meta.get("args"), snapshot=meta.get("snapshot"),
|
||||
))
|
||||
if tool_name == "todo" and (plan_update := _build_plan_update_from_todo_result(result)) is not None:
|
||||
_send_update(conn, session_id, loop, plan_update)
|
||||
if not queue:
|
||||
tool_call_ids.pop(tool_name, None)
|
||||
|
||||
return _step
|
||||
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Agent message callback
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def make_message_cb(
|
||||
conn: acp.Client,
|
||||
session_id: str,
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
) -> Callable:
|
||||
"""Create a callback that streams agent response text to the editor."""
|
||||
|
||||
def _message(text: str) -> None:
|
||||
if not text:
|
||||
return
|
||||
update = acp.update_agent_message_text(text)
|
||||
_send_update(conn, session_id, loop, update)
|
||||
|
||||
return _message
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
import json # noqa: F401,E402
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
@@ -0,0 +1,287 @@
|
||||
"""ACP model picker: deduplicated ``provider:model`` rows from the Hermes inventory + named endpoints."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable
|
||||
|
||||
from acp.schema import ModelInfo, SessionModelState
|
||||
|
||||
logger = logging.getLogger("acp_adapter.server")
|
||||
|
||||
# Per-provider row cap (clients render all `availableModels` in one dropdown; mirrors the
|
||||
# MoA picker cap). Not a total cap; the current model is always kept via the fallback insert.
|
||||
ACP_MAX_MODELS_PER_PROVIDER = 200
|
||||
|
||||
|
||||
def _named_custom_provider_catalogs() -> list[tuple[str, str, list[tuple[str, str]]]]:
|
||||
"""``(slug, label, [(model_id, description), ...])`` for named endpoints (v12 ``providers:``
|
||||
and legacy ``custom_providers:``), which canonical provider enumeration never lists.
|
||||
|
||||
Models = the entry's declared models, refreshed from the live ``/models`` listing when a
|
||||
credential exists and ``discover_models`` isn't disabled; declared models survive a failed
|
||||
discovery (some endpoints have no ``/models`` route). Slugs use the ``custom:<name>`` shape
|
||||
``parse_model_input``/``resolve_runtime_provider`` resolve, so choice ids round-trip."""
|
||||
try:
|
||||
from hermes_cli.config import (get_compatible_custom_providers, is_provider_enabled, load_config)
|
||||
from hermes_cli.model_switch import _declared_model_ids, _entry_models_discovered, _models_config_is_allowlist
|
||||
from hermes_cli.model_switch_providers import _NativePickerModelList, _fetch_picker_live_models
|
||||
from hermes_cli.model_switch_providers import _discover_flag
|
||||
from hermes_cli.models_local import should_use_ollama_native_catalog
|
||||
from hermes_cli.providers import custom_provider_slug
|
||||
except ImportError:
|
||||
return []
|
||||
|
||||
try:
|
||||
cfg = load_config()
|
||||
entries = get_compatible_custom_providers(cfg)
|
||||
except Exception:
|
||||
logger.debug("Could not load named custom providers", exc_info=True)
|
||||
return []
|
||||
|
||||
# ``get_compatible_custom_providers`` drops ``enabled``; read disabled keys from raw config.
|
||||
raw_providers = cfg.get("providers") if isinstance(cfg, dict) else None
|
||||
disabled_keys = {
|
||||
str(key).strip().lower()
|
||||
for key, raw in (raw_providers.items() if isinstance(raw_providers, dict) else ())
|
||||
if isinstance(raw, dict) and not is_provider_enabled(raw)
|
||||
}
|
||||
|
||||
def _entry_catalog(entry: dict) -> tuple[str, str, list[tuple[str, str]]] | None:
|
||||
field = lambda key: str(entry.get(key) or "").strip() # noqa: E731
|
||||
provider_key, name, base_url = field("provider_key"), field("name"), field("base_url")
|
||||
if provider_key.lower() in disabled_keys or not name or not base_url:
|
||||
return None
|
||||
slug = custom_provider_slug(name, provider_key)
|
||||
|
||||
api_key = field("api_key")
|
||||
if not api_key:
|
||||
key_env = str(entry.get("key_env") or entry.get("api_key_env") or "").strip()
|
||||
api_key = os.environ.get(key_env, "").strip() if key_env else ""
|
||||
|
||||
models_cfg = entry.get("models")
|
||||
declared = [m for m in dict.fromkeys([field("model"), *_declared_model_ids(models_cfg)]) if m]
|
||||
|
||||
native_headers = entry.get("extra_headers") or None
|
||||
is_ollama_key = provider_key.lower() in {"ollama", "custom:ollama"}
|
||||
is_native_ollama = should_use_ollama_native_catalog(
|
||||
provider_key if is_ollama_key else "custom", base_url, headers=native_headers
|
||||
)
|
||||
if not api_key and not declared and not is_native_ollama:
|
||||
return None # nothing to discover with and nothing declared: not addressable
|
||||
|
||||
model_ids = list(declared)
|
||||
live = None
|
||||
if _discover_flag(entry) and (api_key or is_native_ollama):
|
||||
try:
|
||||
live = _fetch_picker_live_models(
|
||||
api_key, base_url, provider_key if is_native_ollama and is_ollama_key else "custom",
|
||||
_models_config_is_allowlist(models_cfg, _entry_models_discovered(entry)),
|
||||
headers=native_headers, timeout=1.5, api_mode=entry.get("api_mode"),
|
||||
)
|
||||
except Exception:
|
||||
live = None
|
||||
if isinstance(live, _NativePickerModelList):
|
||||
model_ids = list(live)
|
||||
elif live is not None:
|
||||
model_ids = declared + [m for m in live if m not in declared]
|
||||
|
||||
if not model_ids and not isinstance(live, _NativePickerModelList):
|
||||
return None
|
||||
return slug, name, [(mid, "") for mid in model_ids]
|
||||
|
||||
catalogs = [_entry_catalog(entry) for entry in entries if isinstance(entry, dict)]
|
||||
return [c for c in catalogs if c is not None]
|
||||
|
||||
|
||||
def _semantic_provider(provider_id: str, normalize_provider: Callable[[str], str]) -> str:
|
||||
raw = str(provider_id or "").strip().lower()
|
||||
if raw in {"ollama", "custom:ollama"}:
|
||||
return "ollama"
|
||||
if raw.startswith("custom:"):
|
||||
return raw
|
||||
return normalize_provider(raw)
|
||||
|
||||
|
||||
def _empty_catalog_applies(
|
||||
provider_id: str, empty_authoritative: set[str], normalize_provider: Callable[[str], str]
|
||||
) -> bool:
|
||||
"""True when a named endpoint with an authoritative-empty catalog owns ``provider_id``."""
|
||||
raw = str(provider_id or "").strip().lower()
|
||||
normalized = normalize_provider(raw)
|
||||
if normalized == "custom":
|
||||
return any(
|
||||
candidate == raw
|
||||
or f"custom:{candidate}" == raw
|
||||
or (raw == "custom" and candidate == "custom")
|
||||
for candidate in empty_authoritative
|
||||
)
|
||||
return any(
|
||||
candidate == raw
|
||||
or candidate == f"custom:{normalized}"
|
||||
or candidate == f"custom:{raw}"
|
||||
or normalize_provider(candidate) == normalized
|
||||
for candidate in empty_authoritative
|
||||
)
|
||||
|
||||
|
||||
def _choice_provider(model_id: str) -> str:
|
||||
"""Provider prefix of an encoded choice id; longest configured ``custom:`` slug wins."""
|
||||
parts = model_id.split(":")
|
||||
if parts[:1] == ["custom"] and len(parts) > 1:
|
||||
from hermes_cli.models import _configured_custom_provider_ids
|
||||
|
||||
lowered = model_id.lower()
|
||||
for candidate in sorted(
|
||||
(p for p in _configured_custom_provider_ids() if p.startswith("custom:")), key=len, reverse=True,
|
||||
):
|
||||
if lowered.startswith(candidate + ":"):
|
||||
return candidate
|
||||
return "custom"
|
||||
return parts[0]
|
||||
|
||||
|
||||
def encode_model_choice(provider: str | None, model: str | None) -> str:
|
||||
"""``provider:model`` so ACP clients keep provider context."""
|
||||
raw_model = str(model or "").strip()
|
||||
if not raw_model:
|
||||
return ""
|
||||
raw_provider = str(provider or "").strip().lower()
|
||||
return f"{raw_provider}:{raw_model}" if raw_provider else raw_model
|
||||
|
||||
|
||||
@dataclass
|
||||
class _ModelCatalog:
|
||||
"""Deduplicated ACP model rows from the inventory + named endpoints.
|
||||
|
||||
Dedupes on the encoded choice id AND a semantic ``provider:model`` id (``ollama`` ==
|
||||
``custom:ollama``). A bare/``custom`` current provider whose base_url matches an ollama
|
||||
inventory row is resolved to ``custom:ollama``."""
|
||||
|
||||
normalize_provider: Callable[[str], str]
|
||||
current_model: str
|
||||
current_choice_provider: str
|
||||
current_base_url: str
|
||||
models: list[ModelInfo] = field(default_factory=list)
|
||||
seen_ids: set[str] = field(default_factory=set)
|
||||
seen_semantic_ids: set[str] = field(default_factory=set)
|
||||
empty_authoritative: set[str] = field(default_factory=set)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.current_choice_provider == "ollama":
|
||||
self.current_choice_provider = "custom:ollama"
|
||||
self._identity_resolved = self.current_choice_provider not in {"", "custom"}
|
||||
|
||||
def semantic(self, provider_id: str) -> str:
|
||||
return _semantic_provider(provider_id, self.normalize_provider)
|
||||
|
||||
def add(self, provider_id: str, model_id: str, name: str, description: str) -> None:
|
||||
choice_id = encode_model_choice(provider_id, model_id)
|
||||
semantic_id = f"{self.semantic(provider_id)}:{model_id}"
|
||||
if not choice_id or choice_id in self.seen_ids or semantic_id in self.seen_semantic_ids:
|
||||
return
|
||||
self.models.append(ModelInfo(model_id=choice_id, name=name, description=description))
|
||||
self.seen_ids.add(choice_id)
|
||||
self.seen_semantic_ids.add(semantic_id)
|
||||
|
||||
def add_inventory_rows(self, rows: list, provider_label: Callable[[str], str]) -> None:
|
||||
for row in rows:
|
||||
raw_row_provider = str(row.get("slug") or "").strip().lower()
|
||||
row_provider = self.normalize_provider(raw_row_provider)
|
||||
row_base_url = str(row.get("api_url") or "").strip().rstrip("/").lower()
|
||||
if row.get("native_catalog_empty"):
|
||||
self.empty_authoritative.add(raw_row_provider)
|
||||
if not self._identity_resolved and raw_row_provider in {"ollama", "custom:ollama"} and (
|
||||
self.current_base_url and row_base_url == self.current_base_url
|
||||
):
|
||||
self.current_choice_provider = "custom:ollama"
|
||||
self._identity_resolved = True
|
||||
row_models = row.get("models")
|
||||
if not row_provider or not isinstance(row_models, (list, tuple)):
|
||||
continue
|
||||
provider_name = str(row.get("name") or "").strip() or provider_label(row_provider)
|
||||
encoded_provider = (
|
||||
"custom:ollama" if raw_row_provider == "ollama"
|
||||
else raw_row_provider if raw_row_provider.startswith("custom:")
|
||||
else row_provider
|
||||
)
|
||||
for model_entry in row_models:
|
||||
if isinstance(model_entry, dict):
|
||||
model_entry = model_entry.get("id") or model_entry.get("model") or model_entry.get("name")
|
||||
rendered_model = str(model_entry or "").strip()
|
||||
if not rendered_model:
|
||||
continue
|
||||
is_current = rendered_model == self.current_model and (
|
||||
self.semantic(encoded_provider) == self.semantic(self.current_choice_provider)
|
||||
)
|
||||
self.add(
|
||||
encoded_provider, rendered_model, f"{provider_name} · {rendered_model}",
|
||||
f"Provider: {provider_name}" + (" • current" if is_current else ""),
|
||||
)
|
||||
|
||||
def add_named_catalogs(self, catalogs: list, normalized_provider: str) -> None:
|
||||
"""Named user-defined endpoints (providers: / custom_providers:) are invisible
|
||||
to canonical enumeration — append them like the TUI /model picker. An empty
|
||||
catalog marks that slug authoritative-empty."""
|
||||
for named_slug, named_label, named_catalog in catalogs:
|
||||
if not named_catalog:
|
||||
self.empty_authoritative.add(str(named_slug).strip().lower())
|
||||
continue
|
||||
for named_model, named_desc in named_catalog:
|
||||
is_current = named_slug == normalized_provider and named_model == self.current_model
|
||||
parts = [f"Provider: {named_label}", str(named_desc or "").strip(), "current" if is_current else ""]
|
||||
self.add(named_slug, named_model, named_model, " • ".join(part for part in parts if part))
|
||||
|
||||
|
||||
def build_model_state(model: str, provider: str, base_url: str) -> SessionModelState | None:
|
||||
"""Picker state from the shared inventory + named endpoints; ``None`` when nothing is listable
|
||||
(caller falls back to a single current-model row). Raises on inventory failure."""
|
||||
from hermes_cli.inventory import build_models_payload, load_picker_context
|
||||
from hermes_cli.models import normalize_provider, provider_label
|
||||
|
||||
normalized_provider = normalize_provider(provider)
|
||||
context = load_picker_context().with_overrides(
|
||||
current_provider=normalized_provider, current_model=model, current_base_url=base_url,
|
||||
)
|
||||
payload = build_models_payload(
|
||||
context, explicit_only=True, include_unconfigured=False, picker_hints=False,
|
||||
canonical_order=True, pricing=False, capabilities=False, refresh=False,
|
||||
probe_custom_providers=False, probe_current_custom_provider=False, max_models=ACP_MAX_MODELS_PER_PROVIDER,
|
||||
)
|
||||
|
||||
cat = _ModelCatalog(
|
||||
normalize_provider=normalize_provider, current_model=model,
|
||||
current_choice_provider=str(provider or "").strip().lower(),
|
||||
current_base_url=base_url.strip().rstrip("/").lower(),
|
||||
)
|
||||
cat.add_inventory_rows(payload.get("providers") or [], provider_label)
|
||||
cat.add_named_catalogs(_named_custom_provider_catalogs(), normalized_provider)
|
||||
available_models = cat.models
|
||||
|
||||
def empty_applies(provider_id: str) -> bool:
|
||||
return _empty_catalog_applies(provider_id, cat.empty_authoritative, normalize_provider)
|
||||
|
||||
if cat.empty_authoritative:
|
||||
available_models = [m for m in available_models if not empty_applies(_choice_provider(m.model_id))]
|
||||
|
||||
current_is_empty = empty_applies(cat.current_choice_provider)
|
||||
if current_is_empty:
|
||||
available_models = [m for m in available_models if " • current" not in str(m.description or "")]
|
||||
current_model_id = "" if current_is_empty else encode_model_choice(cat.current_choice_provider, model)
|
||||
if current_model_id and current_model_id not in {item.model_id for item in available_models}:
|
||||
provider_name = provider_label(normalized_provider)
|
||||
available_models.insert(0, ModelInfo(
|
||||
model_id=current_model_id, name=f"{provider_name} · {model}",
|
||||
description=f"Provider: {provider_name} • current",
|
||||
))
|
||||
|
||||
if not available_models and current_is_empty:
|
||||
return SessionModelState(available_models=[], current_model_id="")
|
||||
if available_models:
|
||||
return SessionModelState(
|
||||
available_models=available_models,
|
||||
current_model_id=current_model_id if current_model_id or current_is_empty else available_models[0].model_id,
|
||||
)
|
||||
return None
|
||||
+59
-126
@@ -8,22 +8,14 @@ from concurrent.futures import TimeoutError as FutureTimeout
|
||||
from itertools import count
|
||||
from typing import Callable
|
||||
|
||||
from acp.schema import (
|
||||
AllowedOutcome,
|
||||
PermissionOption,
|
||||
)
|
||||
from acp.schema import AllowedOutcome, PermissionOption
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Maps ACP permission option ids to Hermes approval result strings.
|
||||
# Option ids are stable across both the ``allow_permanent=True`` and
|
||||
# ``allow_permanent=False`` paths even though the option list differs.
|
||||
# ACP permission option id -> Hermes approval result. Ids are stable across the
|
||||
# ``allow_permanent=True`` and ``False`` paths even though the option list differs.
|
||||
_OPTION_ID_TO_HERMES = {
|
||||
"allow_once": "once",
|
||||
"allow_session": "session",
|
||||
"allow_always": "always",
|
||||
"deny": "deny",
|
||||
"deny_always": "deny",
|
||||
"allow_once": "once", "allow_session": "session", "allow_always": "always", "deny": "deny", "deny_always": "deny"
|
||||
}
|
||||
|
||||
_PERMISSION_REQUEST_IDS = count(1)
|
||||
@@ -33,70 +25,41 @@ def _permission_option_supports_kind(kind: str) -> bool:
|
||||
"""Return whether the installed ACP SDK accepts a permission option kind."""
|
||||
try:
|
||||
PermissionOption(option_id="__probe__", kind=kind, name="probe")
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _build_permission_options(
|
||||
*, allow_permanent: bool, allow_session: bool = True,
|
||||
smart_denied: bool = False,
|
||||
*, allow_permanent: bool, allow_session: bool = True, smart_denied: bool = False,
|
||||
) -> list[PermissionOption]:
|
||||
"""Return ACP options that match Hermes approval semantics."""
|
||||
# A gate that re-asks every time (allow_session=False, e.g. protected
|
||||
# agent-instruction writes) collapses to the same two options as a
|
||||
# Smart DENY override — the editor must not offer a scope Hermes
|
||||
# discards, or every subsequent write re-prompts (#81887).
|
||||
# agent-instruction writes) collapses to the same two options as a Smart
|
||||
# DENY override — offering a scope Hermes discards would re-prompt every write.
|
||||
# See #81887.
|
||||
once_only = smart_denied or not allow_session
|
||||
options = [PermissionOption(
|
||||
option_id="allow_once", kind="allow_once", name="Allow once",
|
||||
)]
|
||||
options = [PermissionOption(option_id="allow_once", kind="allow_once", name="Allow once")]
|
||||
if not once_only:
|
||||
options.append(PermissionOption(
|
||||
option_id="allow_session",
|
||||
# ACP has no session-scoped kind, so use the closest persistent
|
||||
# hint while keeping Hermes semantics in the option id.
|
||||
kind="allow_always",
|
||||
name="Allow for session",
|
||||
))
|
||||
if allow_permanent and not once_only:
|
||||
options.append(
|
||||
PermissionOption(
|
||||
option_id="allow_always",
|
||||
kind="allow_always",
|
||||
name="Allow always",
|
||||
),
|
||||
)
|
||||
# ACP has no session-scoped kind: closest persistent hint, Hermes semantics in the id.
|
||||
options.append(PermissionOption(option_id="allow_session", kind="allow_always", name="Allow for session"))
|
||||
if allow_permanent:
|
||||
options.append(PermissionOption(option_id="allow_always", kind="allow_always", name="Allow always"))
|
||||
options.append(PermissionOption(option_id="deny", kind="reject_once", name="Deny"))
|
||||
if not once_only and _permission_option_supports_kind("reject_always"):
|
||||
options.append(
|
||||
PermissionOption(
|
||||
option_id="deny_always",
|
||||
kind="reject_always",
|
||||
name="Deny always",
|
||||
),
|
||||
)
|
||||
options.append(PermissionOption(option_id="deny_always", kind="reject_always", name="Deny always"))
|
||||
return options
|
||||
|
||||
|
||||
def _build_permission_tool_call(command: str, description: str):
|
||||
"""Return the ACP tool-call update attached to a permission request.
|
||||
|
||||
``request_permission`` expects a ``ToolCallUpdate`` payload — produced
|
||||
by ``_acp.update_tool_call`` — not a ``ToolCallStart``. Each request
|
||||
gets a unique ``perm-check-N`` id so concurrent requests don't collide.
|
||||
"""
|
||||
"""Return the ``ToolCallUpdate`` (not ``ToolCallStart``) payload attached to a
|
||||
permission request; unique ``perm-check-N`` ids keep concurrent requests apart."""
|
||||
import acp as _acp
|
||||
|
||||
tool_call_id = f"perm-check-{next(_PERMISSION_REQUEST_IDS)}"
|
||||
title = f"{description}: {command}" if description else command
|
||||
content_text = f"{description}\n$ {command}" if description else f"$ {command}"
|
||||
return _acp.update_tool_call(
|
||||
tool_call_id,
|
||||
title=title,
|
||||
kind="execute",
|
||||
status="pending",
|
||||
content=[_acp.tool_content(_acp.text_block(content_text))],
|
||||
f"perm-check-{next(_PERMISSION_REQUEST_IDS)}", title=f"{description}: {command}" if description else command,
|
||||
kind="execute", status="pending", content=[_acp.tool_content(_acp.text_block(content_text))],
|
||||
raw_input={"command": command, "description": description},
|
||||
)
|
||||
|
||||
@@ -105,86 +68,56 @@ def _map_outcome_to_hermes(outcome: object, *, allowed_option_ids: set[str]) ->
|
||||
"""Map an ACP permission outcome into Hermes approval strings."""
|
||||
if not isinstance(outcome, AllowedOutcome):
|
||||
return "deny"
|
||||
|
||||
option_id = outcome.option_id
|
||||
if option_id not in allowed_option_ids:
|
||||
logger.warning("Permission request returned unknown option_id: %s", option_id)
|
||||
if outcome.option_id not in allowed_option_ids:
|
||||
logger.warning("Permission request returned unknown option_id: %s", outcome.option_id)
|
||||
return "deny"
|
||||
return _OPTION_ID_TO_HERMES.get(option_id, "deny")
|
||||
return _OPTION_ID_TO_HERMES.get(outcome.option_id, "deny")
|
||||
|
||||
|
||||
def make_approval_callback(
|
||||
request_permission_fn: Callable,
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
session_id: str,
|
||||
timeout: float = 60.0,
|
||||
) -> Callable[..., str]:
|
||||
"""
|
||||
Return a Hermes-compatible approval callback that bridges to ACP.
|
||||
def await_permission(
|
||||
request_permission_fn: Callable, loop: asyncio.AbstractEventLoop, session_id: str, *,
|
||||
tool_call, options: list[PermissionOption], timeout: float, what: str,
|
||||
) -> tuple[object | None, bool]:
|
||||
"""Schedule ``request_permission`` on ``loop`` from a worker thread and block for the answer.
|
||||
Returns ``(response, timed_out)``; ``(None, False)`` when scheduling or the request failed."""
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
|
||||
The callback accepts ``command`` and ``description`` plus optional
|
||||
keyword arguments such as ``allow_permanent`` used by
|
||||
``tools.approval.prompt_dangerous_approval()``.
|
||||
coro = request_permission_fn(session_id=session_id, tool_call=tool_call, options=options)
|
||||
future = safe_schedule_threadsafe(coro, loop, logger=logger, log_message=f"{what}: failed to schedule on loop")
|
||||
if future is None:
|
||||
return None, False
|
||||
try:
|
||||
return future.result(timeout=timeout), False
|
||||
except FutureTimeout:
|
||||
future.cancel()
|
||||
logger.warning("%s timed out after %ss", what, timeout)
|
||||
return None, True
|
||||
except Exception as exc:
|
||||
future.cancel()
|
||||
logger.warning("%s failed: %s", what, exc)
|
||||
return None, False
|
||||
|
||||
Args:
|
||||
request_permission_fn: The ACP connection's ``request_permission`` coroutine.
|
||||
loop: The event loop on which the ACP connection lives.
|
||||
session_id: Current ACP session id.
|
||||
timeout: Seconds to wait for a response before auto-denying.
|
||||
"""
|
||||
|
||||
def _callback(
|
||||
command: str,
|
||||
description: str,
|
||||
*,
|
||||
allow_permanent: bool = True,
|
||||
allow_session: bool = True,
|
||||
smart_denied: bool = False,
|
||||
**_: object,
|
||||
) -> str:
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
def make_approval_callback(request_permission_fn: Callable, loop: asyncio.AbstractEventLoop,
|
||||
session_id: str, timeout: float = 60.0) -> Callable[..., str]:
|
||||
"""Return a Hermes approval callback (``command, description, **kw`` as used by
|
||||
``tools.approval.prompt_dangerous_approval()``) that bridges to the ACP
|
||||
connection's ``request_permission`` coroutine on ``loop``; auto-denies after ``timeout`` s."""
|
||||
|
||||
options = _build_permission_options(
|
||||
allow_permanent=allow_permanent,
|
||||
allow_session=allow_session,
|
||||
smart_denied=smart_denied,
|
||||
def _callback(command: str, description: str, *, allow_permanent: bool = True,
|
||||
allow_session: bool = True, smart_denied: bool = False, **_: object) -> str:
|
||||
options = _build_permission_options(allow_permanent=allow_permanent, allow_session=allow_session,
|
||||
smart_denied=smart_denied)
|
||||
response, timed_out = await_permission(
|
||||
request_permission_fn, loop, session_id, tool_call=_build_permission_tool_call(command, description),
|
||||
options=options, timeout=timeout, what="Permission request",
|
||||
)
|
||||
|
||||
tool_call = _build_permission_tool_call(command, description)
|
||||
coro = request_permission_fn(
|
||||
session_id=session_id,
|
||||
tool_call=tool_call,
|
||||
options=options,
|
||||
)
|
||||
future = safe_schedule_threadsafe(
|
||||
coro, loop,
|
||||
logger=logger,
|
||||
log_message="Permission request: failed to schedule on loop",
|
||||
)
|
||||
if future is None:
|
||||
return "deny"
|
||||
|
||||
try:
|
||||
response = future.result(timeout=timeout)
|
||||
except FutureTimeout:
|
||||
future.cancel()
|
||||
logger.warning("Permission request timed out after %ss", timeout)
|
||||
# Distinct from an explicit deny: the client never answered.
|
||||
# tools.approval callers report this as "timed out without user
|
||||
# response" instead of a user denial.
|
||||
if timed_out:
|
||||
# Distinct from an explicit deny: tools.approval reports "timed out
|
||||
# without user response" instead of a user denial.
|
||||
return "timeout"
|
||||
except Exception as exc:
|
||||
future.cancel()
|
||||
logger.warning("Permission request failed: %s", exc)
|
||||
return "deny"
|
||||
|
||||
if response is None:
|
||||
return "deny"
|
||||
|
||||
allowed_option_ids = {option.option_id for option in options}
|
||||
return _map_outcome_to_hermes(
|
||||
response.outcome,
|
||||
allowed_option_ids=allowed_option_ids,
|
||||
)
|
||||
return _map_outcome_to_hermes(response.outcome, allowed_option_ids={option.option_id for option in options})
|
||||
|
||||
return _callback
|
||||
|
||||
+44
-85
@@ -1,14 +1,13 @@
|
||||
"""Derive ACP session-provenance metadata from the existing compression chain.
|
||||
|
||||
This is an additive Hermes extension surfaced under ACP ``_meta.hermes`` so
|
||||
existing ACP clients ignore it. It carries no new persisted state: everything
|
||||
is derived on demand from the ``sessions`` table (``parent_session_id`` /
|
||||
``end_reason``), which already models compression-continuation chains.
|
||||
Additive Hermes extension under ACP ``_meta.hermes`` (unknown to other clients,
|
||||
so ignored). No new persisted state: everything is derived from the ``sessions``
|
||||
table (``parent_session_id`` / ``end_reason``), which already models
|
||||
compression-continuation chains.
|
||||
|
||||
The ACP/editor ``session_id`` stays the stable public handle. When context
|
||||
compression rotates the internal Hermes head, ``build_session_provenance`` lets
|
||||
a client see the previous/current internal ids and the lineage root without
|
||||
parsing status text, guessing from token drops, or reading ``state.db``.
|
||||
The ACP/editor ``session_id`` stays the stable public handle; when compression
|
||||
rotates the internal Hermes head, ``build_session_provenance`` exposes the
|
||||
previous/current internal ids and lineage root without parsing status text.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -19,109 +18,69 @@ from typing import Any, Dict, Optional
|
||||
_MAX_WALK = 100
|
||||
|
||||
|
||||
def _get_row(db: Any, session_id: str) -> Optional[Dict[str, Any]]:
|
||||
try:
|
||||
return db.get_session(session_id)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _is_compression_end(row: Any) -> bool:
|
||||
return bool(row) and row.get("end_reason") == "compression"
|
||||
|
||||
|
||||
def build_session_provenance(
|
||||
db: Any,
|
||||
acp_session_id: str,
|
||||
current_hermes_session_id: str,
|
||||
*,
|
||||
previous_hermes_session_id: Optional[str] = None,
|
||||
db: Any, acp_session_id: str, current_hermes_session_id: str, *, previous_hermes_session_id: Optional[str] = None,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Build ``_meta.hermes.sessionProvenance`` for an ACP session.
|
||||
|
||||
Args:
|
||||
db: A ``SessionDB`` (must expose ``get_session``).
|
||||
acp_session_id: The stable ACP/editor-facing session handle.
|
||||
current_hermes_session_id: The live internal Hermes DB session id
|
||||
(``state.agent.session_id``).
|
||||
previous_hermes_session_id: The internal id from before the most recent
|
||||
turn, when known. Supplied by ``prompt()`` to flag a rotation.
|
||||
|
||||
Returns:
|
||||
A dict suitable for ``{"hermes": {"sessionProvenance": <dict>}}`` under
|
||||
ACP ``_meta``, or ``None`` if the session can't be read.
|
||||
"""
|
||||
try:
|
||||
row = db.get_session(current_hermes_session_id)
|
||||
except Exception:
|
||||
return None
|
||||
``db`` must expose ``get_session``. ``current_hermes_session_id`` is the live
|
||||
internal id (``state.agent.session_id``); ``previous_hermes_session_id`` is
|
||||
the id before the most recent turn, supplied by ``prompt()`` to flag a
|
||||
rotation. Returns ``None`` if the session can't be read."""
|
||||
row = _get_row(db, current_hermes_session_id)
|
||||
if not row:
|
||||
return None
|
||||
|
||||
parent_id = row.get("parent_session_id")
|
||||
end_reason = row.get("end_reason")
|
||||
|
||||
# Walk parents to the lineage root and count compression depth. Only
|
||||
# compression-split parents (parent.end_reason == 'compression') count
|
||||
# toward depth — delegate/branch children share the parent_session_id
|
||||
# column but are not compaction boundaries.
|
||||
root_id = current_hermes_session_id
|
||||
compression_depth = 0
|
||||
cursor_parent = parent_id
|
||||
# Walk parents to the lineage root. Only compression-split parents
|
||||
# (parent.end_reason == 'compression') count toward depth — delegate/branch
|
||||
# children share the parent_session_id column but are not compaction boundaries.
|
||||
root_id, compression_depth, cursor_parent = current_hermes_session_id, 0, parent_id
|
||||
seen = {current_hermes_session_id}
|
||||
for _ in range(_MAX_WALK):
|
||||
if not cursor_parent or cursor_parent in seen:
|
||||
break
|
||||
seen.add(cursor_parent)
|
||||
try:
|
||||
prow = db.get_session(cursor_parent)
|
||||
except Exception:
|
||||
prow = None
|
||||
prow = _get_row(db, cursor_parent)
|
||||
if not prow:
|
||||
break
|
||||
root_id = cursor_parent
|
||||
if prow.get("end_reason") == "compression":
|
||||
compression_depth += 1
|
||||
compression_depth += _is_compression_end(prow)
|
||||
cursor_parent = prow.get("parent_session_id")
|
||||
|
||||
# A session is a compression continuation when its parent was ended with
|
||||
# end_reason='compression'. Determine that from the immediate parent.
|
||||
is_continuation = False
|
||||
if parent_id:
|
||||
try:
|
||||
immediate_parent = db.get_session(parent_id)
|
||||
except Exception:
|
||||
immediate_parent = None
|
||||
if immediate_parent and immediate_parent.get("end_reason") == "compression":
|
||||
is_continuation = True
|
||||
|
||||
rotated = bool(
|
||||
previous_hermes_session_id
|
||||
and previous_hermes_session_id != current_hermes_session_id
|
||||
)
|
||||
# A continuation is a session whose immediate parent ended with end_reason='compression'.
|
||||
is_continuation = bool(parent_id) and _is_compression_end(_get_row(db, parent_id))
|
||||
|
||||
provenance: Dict[str, Any] = {
|
||||
"acpSessionId": acp_session_id,
|
||||
"currentHermesSessionId": current_hermes_session_id,
|
||||
"rootHermesSessionId": root_id,
|
||||
"parentHermesSessionId": parent_id,
|
||||
"sessionKind": "continuation" if is_continuation else "root",
|
||||
"compressionDepth": compression_depth,
|
||||
"acpSessionId": acp_session_id, "currentHermesSessionId": current_hermes_session_id,
|
||||
"rootHermesSessionId": root_id, "parentHermesSessionId": parent_id,
|
||||
"sessionKind": "continuation" if is_continuation else "root", "compressionDepth": compression_depth,
|
||||
}
|
||||
if previous_hermes_session_id:
|
||||
provenance["previousHermesSessionId"] = previous_hermes_session_id
|
||||
if rotated:
|
||||
# The head moved during the last turn. The only mechanism that rotates
|
||||
# the internal id mid-turn is compression-driven session splitting.
|
||||
provenance["reason"] = "compression"
|
||||
provenance["creatorKind"] = "compression"
|
||||
|
||||
if previous_hermes_session_id != current_hermes_session_id:
|
||||
# The only mechanism that rotates the internal id mid-turn is
|
||||
# compression-driven session splitting.
|
||||
provenance["reason"] = "compression"
|
||||
provenance["creatorKind"] = "compression"
|
||||
return provenance
|
||||
|
||||
|
||||
def session_provenance_meta(
|
||||
db: Any,
|
||||
acp_session_id: str,
|
||||
current_hermes_session_id: str,
|
||||
*,
|
||||
previous_hermes_session_id: Optional[str] = None,
|
||||
db: Any, acp_session_id: str, current_hermes_session_id: str, *, previous_hermes_session_id: Optional[str] = None,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Return a ready ``_meta`` payload: ``{"hermes": {"sessionProvenance": ...}}``."""
|
||||
prov = build_session_provenance(
|
||||
db,
|
||||
acp_session_id,
|
||||
current_hermes_session_id,
|
||||
previous_hermes_session_id=previous_hermes_session_id,
|
||||
)
|
||||
if prov is None:
|
||||
return None
|
||||
return {"hermes": {"sessionProvenance": prov}}
|
||||
prov = build_session_provenance(db, acp_session_id, current_hermes_session_id,
|
||||
previous_hermes_session_id=previous_hermes_session_id)
|
||||
return None if prov is None else {"hermes": {"sessionProvenance": prov}}
|
||||
|
||||
+622
-2251
File diff suppressed because it is too large
Load Diff
+181
-441
@@ -1,14 +1,12 @@
|
||||
"""ACP session manager — maps ACP sessions to Hermes AIAgent instances.
|
||||
|
||||
Sessions are persisted to the shared SessionDB (``~/.hermes/state.db``) so they
|
||||
survive process restarts and appear in ``session_search``. When the editor
|
||||
reconnects after idle/restart, the ``load_session`` / ``resume_session`` calls
|
||||
find the persisted session in the database and restore the full conversation
|
||||
history.
|
||||
survive process restarts and appear in ``session_search``; ``load_session`` /
|
||||
``resume_session`` after an editor reconnect restore the full history from there.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from hermes_constants import get_hermes_home
|
||||
from hermes_constants import get_hermes_home, translate_cwd_for_wsl_backend, windows_path_to_wsl
|
||||
|
||||
import copy
|
||||
import json
|
||||
@@ -16,75 +14,53 @@ import logging
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from dataclasses import dataclass, field
|
||||
from threading import Lock
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _translate_acp_cwd(cwd: str) -> str:
|
||||
"""Translate Windows ACP cwd values when Hermes itself is running in WSL.
|
||||
|
||||
Windows ACP clients can launch ``hermes acp`` inside WSL while still sending
|
||||
editor workspaces as Windows drive paths (``E:\\Projects``) or
|
||||
``\\\\wsl.localhost\\`` UNC paths. Store and execute against the POSIX form so
|
||||
agents, tools, and persisted ACP sessions all agree on the usable workspace.
|
||||
Native Linux/macOS keeps the original cwd unchanged.
|
||||
"""
|
||||
from hermes_constants import translate_cwd_for_wsl_backend
|
||||
|
||||
"""Translate Windows ACP cwd values (``E:\\Projects``, ``\\\\wsl.localhost\\``) to POSIX form
|
||||
when Hermes runs in WSL so agents, tools, and persisted sessions agree; no-op elsewhere."""
|
||||
return translate_cwd_for_wsl_backend(str(cwd))
|
||||
|
||||
|
||||
def _normalize_cwd_for_compare(cwd: str | None) -> str:
|
||||
raw = str(cwd or ".").strip()
|
||||
if not raw:
|
||||
raw = "."
|
||||
expanded = os.path.expanduser(raw)
|
||||
|
||||
# Normalize Windows drive paths into the equivalent WSL mount form so
|
||||
# ACP history filters match the same workspace across Windows and WSL.
|
||||
from hermes_constants import windows_path_to_wsl
|
||||
expanded = os.path.expanduser(str(cwd or ".").strip() or ".")
|
||||
|
||||
# Windows drive paths -> WSL mount form so history filters match across hosts.
|
||||
translated = windows_path_to_wsl(expanded)
|
||||
if translated is not None:
|
||||
expanded = translated
|
||||
elif re.match(r"^/mnt/[A-Za-z]/", expanded):
|
||||
expanded = f"/mnt/{expanded[5].lower()}/{expanded[7:]}"
|
||||
|
||||
# Resolve symlink aliases so equivalent spellings of the same directory
|
||||
# compare equal — macOS reports editor workspaces as ``/var/...`` while
|
||||
# sessions get stored under ``/private/var/...`` (and ``/tmp`` vs
|
||||
# ``/private/tmp``), which made ACP history filters silently drop a
|
||||
# workspace's own sessions. ``os.path.realpath`` is lexical for missing
|
||||
# paths (strict=False), so cwds that don't exist on this host — e.g.
|
||||
# WSL-translated Windows drives — keep the previous normpath behavior.
|
||||
# Ported from PrimeIntellect-ai/prime-agent#628.
|
||||
# realpath resolves symlink aliases (macOS ``/var`` vs ``/private/var``, ``/tmp`` vs
|
||||
# ``/private/tmp``) that otherwise drop a workspace's own sessions; it is lexical
|
||||
# for missing paths (e.g. WSL-translated drives).
|
||||
try:
|
||||
# Resolve symlink aliases so equivalent spellings of the same directory compare equal — macOS
|
||||
# reports editor workspaces as ``/var/...`` while sessions get stored under ``/private/var/...``
|
||||
# (and ``/tmp`` vs ``/private/tmp``), which made ACP history filters silently drop a workspace's own
|
||||
# sessions. WSL-translated Windows drives — keep the previous normpath behavior. Ported from
|
||||
# PrimeIntellect-ai/prime-agent#628.
|
||||
return os.path.realpath(expanded)
|
||||
except OSError:
|
||||
return os.path.normpath(expanded)
|
||||
|
||||
|
||||
def _build_session_title(title: Any, preview: Any, cwd: str | None) -> str:
|
||||
explicit = str(title or "").strip()
|
||||
if explicit:
|
||||
return explicit
|
||||
preview_text = str(preview or "").strip()
|
||||
if preview_text:
|
||||
return preview_text
|
||||
leaf = os.path.basename(str(cwd or "").rstrip("/\\"))
|
||||
return leaf or "New thread"
|
||||
return str(title or "").strip() or str(preview or "").strip() or leaf or "New thread"
|
||||
|
||||
|
||||
def _format_updated_at(value: Any) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str) and value.strip():
|
||||
if value is None or (isinstance(value, str) and value.strip()):
|
||||
return value
|
||||
try:
|
||||
return datetime.fromtimestamp(float(value), tz=timezone.utc).isoformat()
|
||||
@@ -93,41 +69,28 @@ def _format_updated_at(value: Any) -> str | None:
|
||||
|
||||
|
||||
def _updated_at_sort_key(value: Any) -> float:
|
||||
if value is None:
|
||||
return float("-inf")
|
||||
if isinstance(value, (int, float)):
|
||||
return float(value)
|
||||
raw = str(value).strip()
|
||||
raw = str(value).strip() if value is not None else ""
|
||||
if not raw:
|
||||
return float("-inf")
|
||||
try:
|
||||
return datetime.fromisoformat(raw.replace("Z", "+00:00")).timestamp()
|
||||
except Exception:
|
||||
for parse in (lambda s: datetime.fromisoformat(s.replace("Z", "+00:00")).timestamp(), float):
|
||||
try:
|
||||
return float(raw)
|
||||
return parse(raw)
|
||||
except Exception:
|
||||
return float("-inf")
|
||||
continue
|
||||
return float("-inf")
|
||||
|
||||
|
||||
def _acp_stderr_print(*args, **kwargs) -> None:
|
||||
"""Best-effort human-readable output sink for ACP stdio sessions.
|
||||
|
||||
ACP reserves stdout for JSON-RPC frames, so any incidental CLI/status output
|
||||
from AIAgent must be redirected away from stdout. Route it to stderr instead.
|
||||
"""
|
||||
kwargs = dict(kwargs)
|
||||
"""Route incidental AIAgent output to stderr; ACP reserves stdout for JSON-RPC."""
|
||||
kwargs.setdefault("file", sys.stderr)
|
||||
print(*args, **kwargs)
|
||||
|
||||
|
||||
def _register_task_cwd(task_id: str, cwd: str) -> None:
|
||||
"""Bind a task/session id to the editor's working directory for tools.
|
||||
|
||||
Zed can launch Hermes from a Windows workspace while the ACP process runs
|
||||
inside WSL. In that case ACP sends cwd as e.g. ``E:\\Projects\\POTI``;
|
||||
local tools need the WSL mount equivalent or subprocess creation fails
|
||||
before the command can run.
|
||||
"""
|
||||
"""Bind a task/session id to the editor cwd for tools. Zed may send a Windows cwd while
|
||||
the ACP process runs in WSL; tools need the WSL mount or subprocess creation fails."""
|
||||
if not task_id:
|
||||
return
|
||||
try:
|
||||
@@ -137,33 +100,32 @@ def _register_task_cwd(task_id: str, cwd: str) -> None:
|
||||
logger.debug("Failed to register ACP task cwd override", exc_info=True)
|
||||
|
||||
|
||||
def _expand_acp_enabled_toolsets(
|
||||
toolsets: List[str] | None = None,
|
||||
mcp_server_names: List[str] | None = None,
|
||||
) -> List[str]:
|
||||
def _expand_acp_enabled_toolsets(toolsets: List[str] | None = None,
|
||||
mcp_server_names: List[str] | None = None) -> List[str]:
|
||||
"""Return ACP toolsets plus explicit MCP server toolsets for this session."""
|
||||
expanded: List[str] = []
|
||||
for name in list(toolsets or ["hermes-acp"]):
|
||||
if name and name not in expanded:
|
||||
expanded.append(name)
|
||||
|
||||
for server_name in list(mcp_server_names or []):
|
||||
toolset_name = f"mcp-{server_name}"
|
||||
if server_name and toolset_name not in expanded:
|
||||
expanded.append(toolset_name)
|
||||
|
||||
return expanded
|
||||
names = [n for n in (toolsets or ["hermes-acp"]) if n]
|
||||
names += [f"mcp-{s}" for s in (mcp_server_names or []) if s]
|
||||
return list(dict.fromkeys(names))
|
||||
|
||||
|
||||
def _clear_task_cwd(task_id: str) -> None:
|
||||
"""Remove task-specific cwd overrides for an ACP session."""
|
||||
if not task_id:
|
||||
return
|
||||
def _parse_model_config(mc: Any) -> dict:
|
||||
"""Decode a persisted model_config JSON blob; ``{}`` when absent/invalid/non-dict."""
|
||||
try:
|
||||
from tools.terminal_tool import clear_task_env_overrides
|
||||
clear_task_env_overrides(task_id)
|
||||
except Exception:
|
||||
logger.debug("Failed to clear ACP task cwd override", exc_info=True)
|
||||
meta = json.loads(mc) if mc else None
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
meta = None
|
||||
return meta if isinstance(meta, dict) else {}
|
||||
|
||||
|
||||
def _session_info(sid: str, cwd: str, model: Any, history_len: int, title: Any, preview: Any,
|
||||
updated_at: Any) -> Dict[str, Any]:
|
||||
return {"session_id": sid, "cwd": cwd, "model": model, "history_len": history_len,
|
||||
"title": _build_session_title(title, preview, cwd), "updated_at": _format_updated_at(updated_at)}
|
||||
|
||||
|
||||
def _first_user_preview(history: List[Dict[str, Any]], default: str) -> str:
|
||||
return next((str(m.get("content") or "").strip() for m in history
|
||||
if m.get("role") == "user" and str(m.get("content") or "").strip()), default)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -178,7 +140,7 @@ class SessionState:
|
||||
cancel_event: Any = None # threading.Event
|
||||
is_running: bool = False
|
||||
queued_prompts: List[str] = field(default_factory=list)
|
||||
runtime_lock: Any = field(default_factory=Lock)
|
||||
runtime_lock: Any = field(default_factory=threading.Lock)
|
||||
current_prompt_text: str = ""
|
||||
interrupted_prompt_text: str = ""
|
||||
|
||||
@@ -186,22 +148,14 @@ class SessionState:
|
||||
class SessionManager:
|
||||
"""Thread-safe manager for ACP sessions backed by Hermes AIAgent instances.
|
||||
|
||||
Sessions are held in-memory for fast access **and** persisted to the
|
||||
shared SessionDB so they survive process restarts and are searchable
|
||||
via ``session_search``.
|
||||
"""
|
||||
Sessions are held in-memory for fast access **and** persisted to the shared
|
||||
SessionDB so they survive restarts and are searchable via ``session_search``."""
|
||||
|
||||
def __init__(self, agent_factory=None, db=None):
|
||||
"""
|
||||
Args:
|
||||
agent_factory: Optional callable that creates an AIAgent-like object.
|
||||
Used by tests. When omitted, a real AIAgent is created
|
||||
using the current Hermes runtime provider configuration.
|
||||
db: Optional SessionDB instance. When omitted, the default
|
||||
SessionDB (``~/.hermes/state.db``) is lazily created.
|
||||
"""
|
||||
"""``agent_factory``: AIAgent-like factory (tests); default builds a real AIAgent from
|
||||
the runtime provider config. ``db``: SessionDB; default lazily opens ``~/.hermes/state.db``."""
|
||||
self._sessions: Dict[str, SessionState] = {}
|
||||
self._lock = Lock()
|
||||
self._lock = threading.Lock()
|
||||
self._agent_factory = agent_factory
|
||||
self._db_instance = db # None → lazy-init on first use
|
||||
|
||||
@@ -209,74 +163,30 @@ class SessionManager:
|
||||
|
||||
def create_session(self, cwd: str = ".") -> SessionState:
|
||||
"""Create a new session with a unique ID and a fresh AIAgent."""
|
||||
import threading
|
||||
|
||||
cwd = _translate_acp_cwd(cwd)
|
||||
session_id = str(uuid.uuid4())
|
||||
agent = self._make_agent(session_id=session_id, cwd=cwd)
|
||||
state = SessionState(
|
||||
session_id=session_id,
|
||||
agent=agent,
|
||||
cwd=cwd,
|
||||
model=getattr(agent, "model", "") or "",
|
||||
cancel_event=threading.Event(),
|
||||
)
|
||||
with self._lock:
|
||||
self._sessions[session_id] = state
|
||||
_register_task_cwd(session_id, cwd)
|
||||
self._persist(state)
|
||||
state = self._install_state(session_id, agent, cwd, getattr(agent, "model", "") or "", [])
|
||||
logger.info("Created ACP session %s (cwd=%s)", session_id, cwd)
|
||||
return state
|
||||
|
||||
def get_session(self, session_id: str) -> Optional[SessionState]:
|
||||
"""Return the session for *session_id*, or ``None``.
|
||||
|
||||
If the session is not in memory but exists in the database (e.g. after
|
||||
a process restart), it is transparently restored.
|
||||
"""
|
||||
"""Return the session, transparently restoring it from the DB (e.g. after
|
||||
a process restart) when it is not in memory; ``None`` if unknown."""
|
||||
with self._lock:
|
||||
state = self._sessions.get(session_id)
|
||||
if state is not None:
|
||||
return state
|
||||
# Attempt to restore from database.
|
||||
return self._restore(session_id)
|
||||
|
||||
def remove_session(self, session_id: str) -> bool:
|
||||
"""Remove a session from memory and database. Returns True if it existed."""
|
||||
with self._lock:
|
||||
existed = self._sessions.pop(session_id, None) is not None
|
||||
db_existed = self._delete_persisted(session_id)
|
||||
if existed or db_existed:
|
||||
_clear_task_cwd(session_id)
|
||||
return existed or db_existed
|
||||
return state if state is not None else self._restore(session_id)
|
||||
|
||||
def fork_session(self, session_id: str, cwd: str = ".") -> Optional[SessionState]:
|
||||
"""Deep-copy a session's history into a new session."""
|
||||
import threading
|
||||
|
||||
cwd = _translate_acp_cwd(cwd)
|
||||
original = self.get_session(session_id) # checks DB too
|
||||
if original is None:
|
||||
return None
|
||||
|
||||
new_id = str(uuid.uuid4())
|
||||
agent = self._make_agent(
|
||||
session_id=new_id,
|
||||
cwd=cwd,
|
||||
model=original.model or None,
|
||||
)
|
||||
state = SessionState(
|
||||
session_id=new_id,
|
||||
agent=agent,
|
||||
cwd=cwd,
|
||||
model=getattr(agent, "model", original.model) or original.model,
|
||||
history=copy.deepcopy(original.history),
|
||||
cancel_event=threading.Event(),
|
||||
)
|
||||
with self._lock:
|
||||
self._sessions[new_id] = state
|
||||
_register_task_cwd(new_id, cwd)
|
||||
self._persist(state)
|
||||
agent = self._make_agent(session_id=new_id, cwd=cwd, model=original.model or None)
|
||||
model = getattr(agent, "model", original.model) or original.model
|
||||
state = self._install_state(new_id, agent, cwd, model, copy.deepcopy(original.history))
|
||||
logger.info("Forked ACP session %s -> %s", session_id, new_id)
|
||||
return state
|
||||
|
||||
@@ -285,71 +195,39 @@ class SessionManager:
|
||||
normalized_cwd = _normalize_cwd_for_compare(cwd) if cwd else None
|
||||
db = self._get_db()
|
||||
persisted_rows: dict[str, dict[str, Any]] = {}
|
||||
try:
|
||||
for row in (db.list_sessions_rich(source="acp", limit=1000) if db is not None else ()):
|
||||
persisted_rows[str(row["id"])] = dict(row)
|
||||
except Exception:
|
||||
logger.debug("Failed to load ACP sessions from DB", exc_info=True)
|
||||
|
||||
if db is not None:
|
||||
try:
|
||||
for row in db.list_sessions_rich(source="acp", limit=1000):
|
||||
persisted_rows[str(row["id"])] = dict(row)
|
||||
except Exception:
|
||||
logger.debug("Failed to load ACP sessions from DB", exc_info=True)
|
||||
def _matches(session_cwd: str) -> bool:
|
||||
return not normalized_cwd or _normalize_cwd_for_compare(session_cwd) == normalized_cwd
|
||||
|
||||
# Collect in-memory sessions first.
|
||||
# In-memory sessions first.
|
||||
with self._lock:
|
||||
seen_ids = set(self._sessions.keys())
|
||||
results = []
|
||||
for s in self._sessions.values():
|
||||
history_len = len(s.history)
|
||||
if history_len <= 0:
|
||||
continue
|
||||
if normalized_cwd and _normalize_cwd_for_compare(s.cwd) != normalized_cwd:
|
||||
if not s.history or not _matches(s.cwd):
|
||||
continue
|
||||
persisted = persisted_rows.get(s.session_id, {})
|
||||
preview = next(
|
||||
(
|
||||
str(msg.get("content") or "").strip()
|
||||
for msg in s.history
|
||||
if msg.get("role") == "user" and str(msg.get("content") or "").strip()
|
||||
),
|
||||
persisted.get("preview") or "",
|
||||
)
|
||||
results.append(
|
||||
{
|
||||
"session_id": s.session_id,
|
||||
"cwd": s.cwd,
|
||||
"model": s.model,
|
||||
"history_len": history_len,
|
||||
"title": _build_session_title(persisted.get("title"), preview, s.cwd),
|
||||
"updated_at": _format_updated_at(
|
||||
persisted.get("last_active") or persisted.get("started_at") or time.time()
|
||||
),
|
||||
}
|
||||
)
|
||||
results.append(_session_info(
|
||||
s.session_id, s.cwd, s.model, len(s.history), persisted.get("title"),
|
||||
_first_user_preview(s.history, persisted.get("preview") or ""),
|
||||
persisted.get("last_active") or persisted.get("started_at") or time.time(),
|
||||
))
|
||||
|
||||
# Merge any persisted sessions not currently in memory.
|
||||
# Then persisted sessions not currently in memory.
|
||||
for sid, row in persisted_rows.items():
|
||||
if sid in seen_ids:
|
||||
continue
|
||||
message_count = int(row.get("message_count") or 0)
|
||||
if message_count <= 0:
|
||||
session_cwd = _parse_model_config(row.get("model_config")).get("cwd", ".")
|
||||
if sid in seen_ids or message_count <= 0 or not _matches(session_cwd):
|
||||
continue
|
||||
# Extract cwd from model_config JSON.
|
||||
session_cwd = "."
|
||||
mc = row.get("model_config")
|
||||
if mc:
|
||||
try:
|
||||
session_cwd = json.loads(mc).get("cwd", ".")
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
if normalized_cwd and _normalize_cwd_for_compare(session_cwd) != normalized_cwd:
|
||||
continue
|
||||
results.append({
|
||||
"session_id": sid,
|
||||
"cwd": session_cwd,
|
||||
"model": row.get("model") or "",
|
||||
"history_len": message_count,
|
||||
"title": _build_session_title(row.get("title"), row.get("preview"), session_cwd),
|
||||
"updated_at": _format_updated_at(row.get("last_active") or row.get("started_at")),
|
||||
})
|
||||
results.append(_session_info(
|
||||
sid, session_cwd, row.get("model") or "", message_count, row.get("title"),
|
||||
row.get("preview"), row.get("last_active") or row.get("started_at"),
|
||||
))
|
||||
|
||||
results.sort(key=lambda item: _updated_at_sort_key(item.get("updated_at")), reverse=True)
|
||||
return results
|
||||
@@ -365,32 +243,9 @@ class SessionManager:
|
||||
self._persist(state)
|
||||
return state
|
||||
|
||||
def cleanup(self) -> None:
|
||||
"""Remove all sessions (memory and database) and clear task-specific cwd overrides."""
|
||||
with self._lock:
|
||||
session_ids = list(self._sessions.keys())
|
||||
self._sessions.clear()
|
||||
for session_id in session_ids:
|
||||
_clear_task_cwd(session_id)
|
||||
self._delete_persisted(session_id)
|
||||
# Also remove any DB-only ACP sessions not currently in memory.
|
||||
db = self._get_db()
|
||||
if db is not None:
|
||||
try:
|
||||
rows = db.search_sessions(source="acp", limit=10000)
|
||||
for row in rows:
|
||||
sid = row["id"]
|
||||
_clear_task_cwd(sid)
|
||||
db.delete_session(sid)
|
||||
except Exception:
|
||||
logger.debug("Failed to cleanup ACP sessions from DB", exc_info=True)
|
||||
|
||||
def save_session(self, session_id: str) -> None:
|
||||
"""Persist the current state of a session to the database.
|
||||
|
||||
Called by the server after prompt completion, slash commands that
|
||||
mutate history, and model switches.
|
||||
"""
|
||||
"""Persist a session; called by the server after prompt completion,
|
||||
history-mutating slash commands, and model switches."""
|
||||
with self._lock:
|
||||
state = self._sessions.get(session_id)
|
||||
if state is not None:
|
||||
@@ -398,34 +253,32 @@ class SessionManager:
|
||||
|
||||
# ---- persistence via SessionDB ------------------------------------------
|
||||
|
||||
def _install_state(self, session_id: str, agent: Any, cwd: str, model: str,
|
||||
history: List[Dict[str, Any]], *, persist: bool = True) -> SessionState:
|
||||
"""Build a SessionState, register it in memory, bind its cwd for tools, optionally persist."""
|
||||
state = SessionState(session_id=session_id, agent=agent, cwd=cwd, model=model,
|
||||
history=history, cancel_event=threading.Event())
|
||||
with self._lock:
|
||||
self._sessions[session_id] = state
|
||||
_register_task_cwd(session_id, cwd)
|
||||
if persist:
|
||||
self._persist(state)
|
||||
return state
|
||||
|
||||
def _get_db(self):
|
||||
"""Lazily initialise and return the SessionDB instance.
|
||||
|
||||
Returns ``None`` if the DB is unavailable (e.g. import error in a
|
||||
minimal test environment).
|
||||
|
||||
Note: we resolve ``HERMES_HOME`` dynamically rather than relying on
|
||||
the module-level ``DEFAULT_DB_PATH`` constant, because that constant
|
||||
is evaluated at import time and won't reflect env-var changes made
|
||||
later (e.g. by the test fixture ``_isolate_hermes_home``).
|
||||
"""
|
||||
if self._db_instance is not None:
|
||||
return self._db_instance
|
||||
try:
|
||||
from hermes_state import SessionDB
|
||||
hermes_home = get_hermes_home()
|
||||
self._db_instance = SessionDB(db_path=hermes_home / "state.db")
|
||||
return self._db_instance
|
||||
except Exception:
|
||||
logger.debug("SessionDB unavailable for ACP persistence", exc_info=True)
|
||||
return None
|
||||
"""Lazily initialise the SessionDB; ``None`` if unavailable (e.g. import error in a
|
||||
minimal test env). ``HERMES_HOME`` is resolved here, not via the import-time
|
||||
``DEFAULT_DB_PATH``, so test fixtures that change the env var later are honoured."""
|
||||
if self._db_instance is None:
|
||||
try:
|
||||
from hermes_state import SessionDB
|
||||
self._db_instance = SessionDB(db_path=get_hermes_home() / "state.db")
|
||||
except Exception:
|
||||
logger.debug("SessionDB unavailable for ACP persistence", exc_info=True)
|
||||
return self._db_instance
|
||||
|
||||
def _persist(self, state: SessionState) -> None:
|
||||
"""Write session state to the database.
|
||||
|
||||
Creates the session record if it doesn't exist, then replaces all
|
||||
stored messages with the current in-memory history.
|
||||
"""
|
||||
"""Create/update the session record, then sync the live message set."""
|
||||
db = self._get_db()
|
||||
if db is None:
|
||||
return
|
||||
@@ -433,182 +286,89 @@ class SessionManager:
|
||||
# Ensure model is a plain string (not a MagicMock or other proxy).
|
||||
model_str = str(state.model) if state.model else None
|
||||
session_meta = {"cwd": state.cwd}
|
||||
provider = getattr(state.agent, "provider", None)
|
||||
base_url = getattr(state.agent, "base_url", None)
|
||||
api_mode = getattr(state.agent, "api_mode", None)
|
||||
if isinstance(provider, str) and provider.strip():
|
||||
session_meta["provider"] = provider.strip()
|
||||
if isinstance(base_url, str) and base_url.strip():
|
||||
session_meta["base_url"] = base_url.strip()
|
||||
if isinstance(api_mode, str) and api_mode.strip():
|
||||
session_meta["api_mode"] = api_mode.strip()
|
||||
cwd_json = json.dumps(session_meta)
|
||||
for key in ("provider", "base_url", "api_mode"):
|
||||
value = getattr(state.agent, key, None)
|
||||
if isinstance(value, str) and value.strip():
|
||||
session_meta[key] = value.strip()
|
||||
|
||||
try:
|
||||
# Ensure the session record exists.
|
||||
existing = db.get_session(state.session_id)
|
||||
if existing is None:
|
||||
db.create_session(
|
||||
session_id=state.session_id,
|
||||
source="acp",
|
||||
model=model_str,
|
||||
model_config={"cwd": state.cwd},
|
||||
)
|
||||
if db.get_session(state.session_id) is None:
|
||||
if not state.history:
|
||||
# Empty editor probes stay ephemeral; copied fork history persists.
|
||||
return
|
||||
db.create_session(session_id=state.session_id, source="acp", model=model_str,
|
||||
model_config={"cwd": state.cwd})
|
||||
else:
|
||||
# Update model_config (contains cwd) if changed.
|
||||
try:
|
||||
db.update_session_meta(state.session_id, cwd_json, model_str)
|
||||
db.update_session_meta(state.session_id, json.dumps(session_meta), model_str)
|
||||
except Exception:
|
||||
logger.debug("Failed to update ACP session metadata", exc_info=True)
|
||||
|
||||
# When the agent owns persistence to this same SessionDB it has
|
||||
# already flushed the live transcript incrementally during
|
||||
# run_conversation (append_message), and it preserves pre-compaction
|
||||
# turns non-destructively via archive_and_compact() — keeping them on
|
||||
# disk as searchable active=0/compacted=1 rows. Calling
|
||||
# replace_messages() here would then be a redundant double-write that
|
||||
# DELETEs exactly those archived rows (and, after a compression-driven
|
||||
# id rotation where agent.session_id no longer equals
|
||||
# state.session_id, clobbers the ended parent transcript) — silent
|
||||
# data loss for any ACP conversation long enough to compress.
|
||||
#
|
||||
# Only fall back to the destructive atomic replace when the agent is
|
||||
# NOT persisting itself to this DB (e.g. a test agent factory, or a
|
||||
# fresh create/fork whose copied history the agent has not flushed
|
||||
# yet). That path still rolls back on a mid-rewrite failure so the
|
||||
# previously persisted conversation survives (salvaged from #13675).
|
||||
# An agent that owns persistence to this same DB already flushed the transcript
|
||||
# incrementally (append_message) and keeps pre-compaction turns as archived
|
||||
# active=0 rows; replace_messages() would DELETE those (and, after a compression
|
||||
# id rotation, clobber the ended parent transcript). Skip it in that case.
|
||||
# Calling replace_messages() here would then be a redundant double-write that DELETEs exactly
|
||||
# those archived rows (and, after a compression-driven id rotation where agent.session_id no
|
||||
# longer equals state.session_id, clobbers the ended parent transcript) — silent data loss for
|
||||
# any ACP conversation long enough to compress. Only fall back to the destructive atomic replace
|
||||
# when the agent is NOT persisting itself to this DB (e.g. a test agent factory, or a fresh
|
||||
# create/fork whose copied history the agent has not flushed yet). That path still rolls back on
|
||||
# a mid-rewrite failure so the previously persisted conversation survives (salvaged from
|
||||
# #13675).
|
||||
agent = state.agent
|
||||
agent_db = getattr(agent, "_session_db", None)
|
||||
agent_owns_persistence = (
|
||||
agent_db is not None
|
||||
and agent_db is db
|
||||
and bool(getattr(agent, "_session_db_created", False))
|
||||
)
|
||||
if not agent_owns_persistence:
|
||||
# Even when the current agent doesn't "own" persistence, the
|
||||
# session on disk may already carry compaction-archived rows —
|
||||
# e.g. after a model switch or a /restore, both of which mint a
|
||||
# fresh agent with _session_db_created=False (so the check above
|
||||
# is False) yet leave the durable archived transcript in place.
|
||||
# A full-history replace would DELETE those archived rows just
|
||||
# like the owned-agent case. Guard against it by replacing ONLY
|
||||
# the live (active=1) set unconditionally: on a fresh
|
||||
# create/fork every row is active=1, so active-only replace is
|
||||
# behaviorally identical to the full replace — and when archived
|
||||
# rows DO exist they survive. An existence probe here
|
||||
# (has_archived_messages) would fail OPEN into the destructive
|
||||
# replace on any DB error and can race a concurrent
|
||||
# archive_and_compact — the same probe failure mode #80216's
|
||||
# /retry fix (gateway/slash_commands.py) deliberately avoids.
|
||||
db.replace_messages(
|
||||
state.session_id, state.history, active_only=True
|
||||
)
|
||||
if getattr(agent, "_session_db", None) is db and getattr(agent, "_session_db_created", False):
|
||||
return
|
||||
# A non-owning agent (model switch, /restore: fresh agent, _session_db_created=False)
|
||||
# may still sit on archived rows, so replace ONLY the active=1 set: on a fresh
|
||||
# create/fork every row is active (== full replace), and archived rows survive.
|
||||
# Unconditional because an existence probe would fail OPEN on DB error and can
|
||||
# race a concurrent archive_and_compact. Still rolls back on mid-rewrite failure.
|
||||
db.replace_messages(state.session_id, state.history, active_only=True)
|
||||
except Exception:
|
||||
logger.warning("Failed to persist ACP session %s", state.session_id, exc_info=True)
|
||||
|
||||
def _restore(self, session_id: str) -> Optional[SessionState]:
|
||||
"""Load a session from the database into memory, recreating the AIAgent."""
|
||||
import threading
|
||||
|
||||
"""Load an ACP session from the database into memory, recreating the AIAgent."""
|
||||
db = self._get_db()
|
||||
if db is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
row = db.get_session(session_id)
|
||||
except Exception:
|
||||
logger.debug("Failed to query DB for ACP session %s", session_id, exc_info=True)
|
||||
return None
|
||||
|
||||
if row is None:
|
||||
if row is None or row.get("source") != "acp":
|
||||
return None
|
||||
|
||||
# Only restore ACP sessions.
|
||||
if row.get("source") != "acp":
|
||||
return None
|
||||
meta = _parse_model_config(row.get("model_config"))
|
||||
cwd, model = meta.get("cwd", "."), row.get("model") or None
|
||||
|
||||
# Extract cwd from model_config.
|
||||
cwd = "."
|
||||
requested_provider = row.get("billing_provider")
|
||||
restored_base_url = row.get("billing_base_url")
|
||||
restored_api_mode = None
|
||||
mc = row.get("model_config")
|
||||
if mc:
|
||||
try:
|
||||
meta = json.loads(mc)
|
||||
if isinstance(meta, dict):
|
||||
cwd = meta.get("cwd", ".")
|
||||
requested_provider = meta.get("provider") or requested_provider
|
||||
restored_base_url = meta.get("base_url") or restored_base_url
|
||||
restored_api_mode = meta.get("api_mode") or restored_api_mode
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
|
||||
model = row.get("model") or None
|
||||
|
||||
# Load conversation history. repair_alternation: this restore feeds
|
||||
# LIVE REPLAY — the loaded list becomes the resumed agent's working
|
||||
# conversation. A durable ``user;user`` violation left in state.db would
|
||||
# otherwise re-fire the pre-request defensive repair on every request
|
||||
# for the rest of the session (see hermes_state.get_messages_as_conversation).
|
||||
# repair_alternation: this list becomes the resumed agent's LIVE conversation; a durable
|
||||
# ``user;user`` violation in state.db would otherwise re-fire the pre-request repair every request.
|
||||
try:
|
||||
history = db.get_messages_as_conversation(
|
||||
session_id, repair_alternation=True
|
||||
)
|
||||
history = db.get_messages_as_conversation(session_id, repair_alternation=True)
|
||||
except Exception:
|
||||
logger.warning("Failed to load messages for ACP session %s", session_id, exc_info=True)
|
||||
history = []
|
||||
|
||||
try:
|
||||
agent = self._make_agent(
|
||||
session_id=session_id,
|
||||
cwd=cwd,
|
||||
model=model,
|
||||
requested_provider=requested_provider,
|
||||
base_url=restored_base_url,
|
||||
api_mode=restored_api_mode,
|
||||
)
|
||||
session_id=session_id, cwd=cwd, model=model, api_mode=meta.get("api_mode") or None,
|
||||
requested_provider=meta.get("provider") or row.get("billing_provider"),
|
||||
base_url=meta.get("base_url") or row.get("billing_base_url"))
|
||||
except Exception:
|
||||
logger.warning("Failed to recreate agent for ACP session %s", session_id, exc_info=True)
|
||||
return None
|
||||
|
||||
state = SessionState(
|
||||
session_id=session_id,
|
||||
agent=agent,
|
||||
cwd=cwd,
|
||||
model=model or getattr(agent, "model", "") or "",
|
||||
history=history,
|
||||
cancel_event=threading.Event(),
|
||||
)
|
||||
with self._lock:
|
||||
self._sessions[session_id] = state
|
||||
_register_task_cwd(session_id, cwd)
|
||||
state = self._install_state(session_id, agent, cwd, model or getattr(agent, "model", "") or "",
|
||||
history, persist=False)
|
||||
logger.info("Restored ACP session %s from DB (%d messages)", session_id, len(history))
|
||||
return state
|
||||
|
||||
def _delete_persisted(self, session_id: str) -> bool:
|
||||
"""Delete a session from the database. Returns True if it existed."""
|
||||
db = self._get_db()
|
||||
if db is None:
|
||||
return False
|
||||
try:
|
||||
return db.delete_session(session_id)
|
||||
except Exception:
|
||||
logger.debug("Failed to delete ACP session %s from DB", session_id, exc_info=True)
|
||||
return False
|
||||
|
||||
# ---- internal -----------------------------------------------------------
|
||||
|
||||
def _make_agent(
|
||||
self,
|
||||
*,
|
||||
session_id: str,
|
||||
cwd: str,
|
||||
model: str | None = None,
|
||||
requested_provider: str | None = None,
|
||||
base_url: str | None = None,
|
||||
api_mode: str | None = None,
|
||||
):
|
||||
def _make_agent(self, *, session_id: str, cwd: str, model: str | None = None,
|
||||
requested_provider: str | None = None, base_url: str | None = None, api_mode: str | None = None):
|
||||
if self._agent_factory is not None:
|
||||
return self._agent_factory()
|
||||
|
||||
@@ -618,78 +378,58 @@ class SessionManager:
|
||||
|
||||
config = load_config()
|
||||
model_cfg = config.get("model")
|
||||
default_model = ""
|
||||
config_provider = None
|
||||
default_model, config_provider = "", None
|
||||
if isinstance(model_cfg, dict):
|
||||
default_model = str(model_cfg.get("default") or default_model)
|
||||
config_provider = model_cfg.get("provider")
|
||||
elif isinstance(model_cfg, str) and model_cfg.strip():
|
||||
default_model, config_provider = str(model_cfg.get("default") or ""), model_cfg.get("provider")
|
||||
elif isinstance(model_cfg, str):
|
||||
default_model = model_cfg.strip()
|
||||
|
||||
configured_mcp_servers = [
|
||||
name
|
||||
for name, cfg in (config.get("mcp_servers") or {}).items()
|
||||
name for name, cfg in (config.get("mcp_servers") or {}).items()
|
||||
if not isinstance(cfg, dict) or cfg.get("enabled", True) is not False
|
||||
]
|
||||
|
||||
kwargs = {
|
||||
"platform": "acp",
|
||||
"enabled_toolsets": _expand_acp_enabled_toolsets(
|
||||
["hermes-acp"],
|
||||
mcp_server_names=configured_mcp_servers,
|
||||
),
|
||||
"quiet_mode": True,
|
||||
"session_id": session_id,
|
||||
"session_db": self._get_db(),
|
||||
"platform": "acp", "quiet_mode": True, "session_id": session_id, "session_db": self._get_db(),
|
||||
"enabled_toolsets": _expand_acp_enabled_toolsets(["hermes-acp"], mcp_server_names=configured_mcp_servers),
|
||||
"model": model or default_model,
|
||||
}
|
||||
|
||||
try:
|
||||
runtime = resolve_runtime_provider(requested=requested_provider or config_provider)
|
||||
kwargs.update(
|
||||
{
|
||||
"provider": runtime.get("provider"),
|
||||
"api_mode": api_mode or runtime.get("api_mode"),
|
||||
"base_url": base_url or runtime.get("base_url"),
|
||||
"api_key": runtime.get("api_key"),
|
||||
"command": runtime.get("command"),
|
||||
"args": list(runtime.get("args") or []),
|
||||
}
|
||||
)
|
||||
kwargs.update({
|
||||
"provider": runtime.get("provider"), "api_mode": api_mode or runtime.get("api_mode"),
|
||||
"base_url": base_url or runtime.get("base_url"), "api_key": runtime.get("api_key"),
|
||||
"command": runtime.get("command"), "args": list(runtime.get("args") or []),
|
||||
})
|
||||
except Exception:
|
||||
logger.debug("ACP session falling back to default provider resolution", exc_info=True)
|
||||
|
||||
_register_task_cwd(session_id, cwd)
|
||||
|
||||
# Bounded wait for background MCP discovery so already-spawning fast
|
||||
# servers land in the agent's tool snapshot. ACP entry.py fires
|
||||
# discovery in a background daemon thread (start_background_mcp_discovery);
|
||||
# the agent snapshots tools once at build (run_agent/agent_init) and
|
||||
# never re-reads the registry, so without this join a reachable-but-
|
||||
# slow configured server would be invisible for the whole session.
|
||||
# ``ensure_mcp_discovery_before_agent_build`` also (re)starts discovery
|
||||
# when the entry.py spawn never ran or exited with zero connected
|
||||
# servers (the retry-after-zero-connected allowance), making this
|
||||
# construction site self-sufficient. Bounded by
|
||||
# ``mcp_discovery_timeout`` (config.yaml, default ~1.5s) so a dead
|
||||
# server can't block — servers that miss the bound are picked up by
|
||||
# the automatic late-refresh (see HermesACPAgent._schedule_mcp_late_refresh).
|
||||
# Bounded wait for the background MCP discovery started by entry.py: the agent
|
||||
# snapshots tools once at build and never re-reads the registry, so without this
|
||||
# join a slow-but-reachable server would be invisible all session. ensure_* also
|
||||
# (re)starts discovery if the entry spawn never ran or connected zero servers.
|
||||
# Bounded by ``mcp_discovery_timeout`` (config.yaml, ~1.5s); late servers are
|
||||
# picked up by HermesACPAgent._schedule_mcp_late_refresh.
|
||||
try:
|
||||
from hermes_cli.mcp_startup import ensure_mcp_discovery_before_agent_build
|
||||
|
||||
ensure_mcp_discovery_before_agent_build(
|
||||
logger=logger,
|
||||
thread_name="acp-mcp-discovery",
|
||||
)
|
||||
ensure_mcp_discovery_before_agent_build(logger=logger, thread_name="acp-mcp-discovery")
|
||||
except Exception:
|
||||
logger.debug("ACP: bounded MCP discovery wait failed", exc_info=True)
|
||||
|
||||
agent = AIAgent(**kwargs)
|
||||
# Codex app-server sessions are spawned lazily on the first turn. Stamp
|
||||
# the ACP workspace onto the agent so the Codex runtime starts from the
|
||||
# editor/session cwd instead of the Hermes daemon's process cwd.
|
||||
# Codex app-server sessions spawn lazily on the first turn; stamp the ACP
|
||||
# workspace so the Codex runtime starts from the editor cwd, not ours.
|
||||
agent.session_cwd = cwd
|
||||
# ACP stdio transport requires stdout to remain protocol-only JSON-RPC.
|
||||
# Route any incidental human-readable agent output to stderr instead.
|
||||
# ACP stdio: stdout is protocol-only JSON-RPC; agent chatter goes to stderr.
|
||||
agent._print_fn = _acp_stderr_print
|
||||
return agent
|
||||
|
||||
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
from threading import Lock # noqa: F401,E402
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
+570
-1019
File diff suppressed because it is too large
Load Diff
+112
@@ -0,0 +1,112 @@
|
||||
# agent/ — AIAgent, turn loop, prompt, compression
|
||||
|
||||
Applies on top of the root `AGENTS.md` (prompt-caching invariant, facade + siblings rules).
|
||||
|
||||
## Shape
|
||||
|
||||
`run_agent.py` is the public facade: `AIAgent` is assembled from mixins (`agent/turn_facade.py`,
|
||||
`client_lifecycle.py`, `stream_delivery.py`, `session_persistence.py`, `compression_facade.py`, ...).
|
||||
Construction runs `agent/agent_init.py::init_agent`; a turn is
|
||||
`agent/conversation_loop.py::run_conversation`, which `AIAgent.run_conversation` forwards to after
|
||||
taking the session turn lease (`turn_facade_lease.py`). `AIAgent.__init__` takes ~60 parameters
|
||||
(credentials, routing, callbacks, session context, budget, credential pool, ...) — read
|
||||
`run_agent.py` for the list; the subset you usually touch: `base_url`, `api_key`, `provider`,
|
||||
`api_mode` (`"chat_completions" | "codex_responses" | ...`), `model` (empty → resolved from
|
||||
config/provider later), `max_iterations` (default 500, shared with subagents),
|
||||
`enabled_toolsets`/`disabled_toolsets`, `quiet_mode`, `save_trajectories`, `platform`
|
||||
(`"cli"`, `"telegram"`, ...), `session_id`, `skip_context_files`, `skip_memory`, `credential_pool`.
|
||||
`chat(message) -> str` is the simple interface; `run_conversation(user_message, system_message=None,
|
||||
conversation_history=None, task_id=None) -> dict` returns `final_response` + `messages`.
|
||||
|
||||
## Agent loop (`agent/conversation_loop.py` + `agent/turn_*.py`)
|
||||
|
||||
Entirely synchronous, with interrupt checks, budget tracking, and a one-turn grace call:
|
||||
|
||||
```python
|
||||
while (api_call_count < self.max_iterations and self.iteration_budget.remaining > 0) \
|
||||
or self._budget_grace_call:
|
||||
if self._interrupt_requested: break
|
||||
response = client.chat.completions.create(model=model, messages=messages, tools=tool_schemas)
|
||||
if response.tool_calls:
|
||||
for tc in response.tool_calls:
|
||||
messages.append(tool_result_message(handle_function_call(tc.name, tc.args, task_id)))
|
||||
api_call_count += 1
|
||||
else:
|
||||
return response.content
|
||||
```
|
||||
|
||||
Each phase of an iteration is its own sibling, so a change to (say) overflow handling touches one
|
||||
~600-line file: `turn_preflight*`, `turn_iteration_prep`, `turn_request_assembly`/`turn_api_request`,
|
||||
`turn_api_call`, `turn_api_error`, `turn_response_intake`/`turn_response_check`,
|
||||
`turn_empty_response`, `turn_tool_round`/`turn_tool_validation`, `turn_overflow`,
|
||||
`turn_truncation`, `turn_context_compaction`, `turn_recovery`, `turn_retry_state`,
|
||||
`turn_stop_gates`, `turn_liveness`, `turn_usage`, `turn_final_response`, `turn_finalizer`,
|
||||
`turn_summary`. Find the phase with `grep -rn "def X" agent/turn_*.py`.
|
||||
|
||||
Messages use OpenAI format `{"role": "system|user|assistant|tool", ...}`; reasoning content is stored
|
||||
in `assistant_msg["reasoning"]`.
|
||||
|
||||
**Agent-level tools** (`todo`, `memory`, ...) are intercepted by `agent/tool_executor.py` through the
|
||||
`INLINE_TOOL_EXECUTORS` table in `agent/inline_tool_executors.py` before `handle_function_call()`.
|
||||
Adding one: register in that table (no `if name == ...` chain); `tools/todo_tool.py` is the pattern.
|
||||
|
||||
## Message-flow invariants (every change is reviewed against these)
|
||||
|
||||
- **Prompt caching must not break.** Never alter past context, change toolsets, reload memories,
|
||||
or rebuild the system prompt mid-conversation. The system prompt is byte-stable for the life of
|
||||
a conversation; the ONLY context mutation is compression. Anything that must inject content
|
||||
mid-conversation rides a **user message or tool result**, never the system prompt: skill slash
|
||||
commands (`agent/skill_commands.py`) inject as a user message; subdirectory `AGENTS.md` hints
|
||||
(`agent/subdirectory_hints.py`) append to the tool result (head+tail truncated past `_MAX_HINT_CHARS = 32_000`, with a warning).
|
||||
- **Strict role alternation.** Never two same-role messages in a row; never a synthetic user
|
||||
message injected mid-loop. The one exception is `/steer`, delivered as a standalone user row
|
||||
after a tool result (`assistant(tool_calls) → tool → user` is legal on every provider path) —
|
||||
never smeared onto the already-persisted tool row, which append-only persistence would leave
|
||||
divergent from the live request. Cron deliveries live in their own session for this reason.
|
||||
- **Context files** (`agent/prompt_builder.py`) load from the CWD only at startup and are capped
|
||||
(`CONTEXT_FILE_MAX_CHARS` / dynamic cap from the context window / `context_file_max_chars`).
|
||||
Never load an install-tree `AGENTS.md` as project context (PR #64611); subdirectory hints reject
|
||||
paths outside the working dir so `~/.codex/AGENTS.md` / `~/.claude/CLAUDE.md` never mix in.
|
||||
- **`_last_resolved_tool_names` is a process-global in `model_tools.py`.** `_run_single_child()`
|
||||
in `tools/delegate_tool.py` saves/restores it around subagent execution; code reading it may see
|
||||
a temporarily stale value during child runs.
|
||||
|
||||
## Compression (`agent/compression_facade.py`, `conversation_compression.py`, `turn_context_compaction.py`)
|
||||
|
||||
Two layers: gateway session hygiene (85% threshold) and the agent `ContextCompressor` (50%,
|
||||
configurable; per-model overrides; failure cooldown after provider-proven overflow). The algorithm
|
||||
prunes old tool results first (no LLM call), then picks boundaries, then generates a structured
|
||||
summary with the `auxiliary` compression model. In-place compaction keeps a single stable session
|
||||
id; native Responses/Codex compaction paths are provider-specific. Compression is the sanctioned
|
||||
cache break — keep it the only one. Full detail:
|
||||
`website/docs/developer-guide/context-compression-and-caching.md`.
|
||||
|
||||
## Model and provider resolution
|
||||
|
||||
- Runtime provider/model resolution and its precedence: `website/docs/developer-guide/provider-runtime.md`.
|
||||
Provider profiles are plugins (`plugins/model-providers/<name>/`, see `plugins/AGENTS.md`);
|
||||
`agent/model_metadata.py` holds context lengths and capabilities.
|
||||
- **Auxiliary (side-LLM) work** — curator, vision, embedding, title generation, session_search,
|
||||
compression — resolves through `agent/auxiliary_client.py::_resolve_auto_route`; each task can pin
|
||||
its own `provider/model/base_url/reasoning_effort` under `auxiliary:` in config.yaml.
|
||||
- Fallback models and credential pools are resolution-chain code: E2E them with real imports
|
||||
against a temp `HERMES_HOME`, not mocks (root rubric).
|
||||
|
||||
## Memory, context engines, curator
|
||||
|
||||
`agent/memory_provider.py` (ABC) + `agent/memory_manager.py` (orchestrator) drive memory-provider
|
||||
plugins; `agent/context_engine.py` drives context-engine plugins; `agent/image_gen_provider.py`
|
||||
image-gen plugins (all in `plugins/AGENTS.md`). `agent/curator.py` + `curator_backup.py` implement
|
||||
the skill curator (`skills/AGENTS.md`). Cron sessions pass `skip_memory=True` by default — memory
|
||||
providers intentionally do not run during cron.
|
||||
|
||||
## Tests
|
||||
|
||||
Loop/phase tests go in `tests/agent/`; patch the binding the phase actually reads (siblings often
|
||||
`from run_agent import X` inside the function — root "patch where production reads"). Assert
|
||||
message-shape invariants (alternation, byte-stable system prompt) rather than snapshotting prompt
|
||||
text.
|
||||
|
||||
Long-form: `website/docs/developer-guide/agent-loop.md`, `prompt-assembly.md`,
|
||||
`context-compression-and-caching.md`, `provider-runtime.md`, `session-storage.md`,
|
||||
`subagent-lifecycle-api.md`.
|
||||
+1
-6
@@ -1,8 +1,3 @@
|
||||
"""Agent internals -- extracted modules from run_agent.py.
|
||||
|
||||
These modules contain pure utility functions and self-contained classes
|
||||
that were previously embedded in the 3,600-line run_agent.py. Extracting
|
||||
them makes run_agent.py focused on the AIAgent orchestrator class.
|
||||
"""
|
||||
"""Agent internals extracted from run_agent.py so it stays focused on AIAgent."""
|
||||
|
||||
from . import jiter_preload as _jiter_preload # noqa: F401
|
||||
|
||||
+302
-631
File diff suppressed because it is too large
Load Diff
+73
-190
@@ -1,33 +1,11 @@
|
||||
"""OpenAI-shape bridge shared by Hermes' ACP clients.
|
||||
|
||||
An ACP agent (``copilot --acp``, and the ACP CLIs that reach Hermes as
|
||||
providers) speaks the Agent Client Protocol, which has no OpenAI-style
|
||||
``tools``/``tool_calls`` channel: a prompt is text, and a response is text plus
|
||||
the agent's *own* tool notifications. Hermes' agentic surface — ``memory``,
|
||||
``todo``, ``skill_manage`` and friends — is dispatched from OpenAI-shaped
|
||||
``tool_calls``, so on an ACP provider it can only work if the schemas travel
|
||||
*into* the prompt as text and the calls are parsed back *out* of the response
|
||||
text.
|
||||
|
||||
``agent/copilot_acp_client.py`` already carried a private copy of that bridge.
|
||||
This module is that code, lifted verbatim into one place so every ACP client
|
||||
shares it instead of re-deriving the wire contract:
|
||||
|
||||
* :func:`render_tool_bridge_sections` — prompt sections describing the
|
||||
forwarded tools and the ``<tool_call>{...}</tool_call>`` contract.
|
||||
* :func:`extract_tool_calls_from_text` — parse those blocks back into
|
||||
``ChatCompletionMessageToolCall`` objects and return the response text with
|
||||
the blocks stripped.
|
||||
* :func:`completion_to_stream_chunks` — re-shape a one-shot ACP response as
|
||||
OpenAI stream chunks for callers that asked for ``stream=True`` (an ACP turn
|
||||
is inherently one-shot from Hermes' perspective).
|
||||
|
||||
The one axis clients differ on is *which* tools they forward, so
|
||||
``render_tool_bridge_sections`` takes an optional allowlist. A CLI with no tools
|
||||
of its own (Copilot) forwards everything Hermes offers; a CLI that is an
|
||||
autonomous agent with its own read/edit/execute tools must forward only Hermes'
|
||||
agent-level tools, because re-offering the overlapping ones makes Hermes re-run
|
||||
work the agent already finished.
|
||||
ACP has no OpenAI-style ``tools``/``tool_calls`` channel, so Hermes' tool schemas travel INTO the
|
||||
prompt as text (:func:`render_tool_bridge_sections`) and calls are parsed back OUT of the response
|
||||
text (:func:`extract_tool_calls_from_text`). Clients differ only in WHICH tools they forward
|
||||
(``allowlist``): a CLI with no tools of its own forwards everything; an autonomous agent with its own
|
||||
read/edit/execute tools forwards only Hermes' agent-level tools, since re-offering overlapping ones
|
||||
makes Hermes redo finished work.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -37,18 +15,13 @@ import re
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Iterable
|
||||
|
||||
from openai.types.chat.chat_completion_message_tool_call import (
|
||||
ChatCompletionMessageToolCall,
|
||||
Function,
|
||||
)
|
||||
from openai.types.chat.chat_completion_message_tool_call import ChatCompletionMessageToolCall, Function
|
||||
|
||||
TOOL_CALL_BLOCK_RE = re.compile(r"<tool_call>\s*(\{.*?\})\s*</tool_call>", re.DOTALL)
|
||||
TOOL_CALL_JSON_RE = re.compile(
|
||||
r"\{\s*\"id\"\s*:\s*\"[^\"]+\"\s*,\s*\"type\"\s*:\s*\"function\"\s*,\s*\"function\"\s*:\s*\{.*?\}\s*\}",
|
||||
re.DOTALL,
|
||||
r"\{\s*\"id\"\s*:\s*\"[^\"]+\"\s*,\s*\"type\"\s*:\s*\"function\"\s*,\s*\"function\"\s*:\s*\{.*?\}\s*\}", re.DOTALL
|
||||
)
|
||||
|
||||
# The contract sentence shared by every ACP client: how to emit a call.
|
||||
TOOL_CALL_CONTRACT = (
|
||||
"Available tools (OpenAI function schema). "
|
||||
"When using a tool, emit ONLY <tool_call>{...}</tool_call> with one JSON object "
|
||||
@@ -56,77 +29,41 @@ TOOL_CALL_CONTRACT = (
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"TOOL_CALL_BLOCK_RE",
|
||||
"TOOL_CALL_JSON_RE",
|
||||
"TOOL_CALL_CONTRACT",
|
||||
"StreamChunks",
|
||||
"build_openai_tool_call",
|
||||
"tool_specs_from_openai_tools",
|
||||
"render_tool_bridge_sections",
|
||||
"extract_tool_calls_from_text",
|
||||
"TOOL_CALL_BLOCK_RE", "TOOL_CALL_JSON_RE", "TOOL_CALL_CONTRACT", "StreamChunks", "build_openai_tool_call",
|
||||
"tool_specs_from_openai_tools", "render_tool_bridge_sections", "extract_tool_calls_from_text",
|
||||
"completion_to_stream_chunks",
|
||||
]
|
||||
|
||||
|
||||
class StreamChunks(list):
|
||||
"""Stream chunks that can still carry response-level attributes.
|
||||
|
||||
Hermes reads provider-level extras off the object returned by
|
||||
``chat.completions.create`` (e.g. ``hermes_projected_messages``, consumed by
|
||||
``agent/provider_projection.py``). A plain list of chunks would silently drop
|
||||
them on the ``stream=True`` path, so ACP clients return this instead and copy
|
||||
the extras onto it.
|
||||
"""
|
||||
"""Chunk list that also carries response-level attributes (e.g. ``hermes_projected_messages``)
|
||||
Hermes reads off the ``create`` result; a plain list would drop them on the stream path."""
|
||||
|
||||
|
||||
def completion_to_stream_chunks(completion: SimpleNamespace) -> StreamChunks:
|
||||
"""Convert a one-shot ACP response into OpenAI-style stream chunks.
|
||||
|
||||
Response-level attributes other than ``choices``/``usage``/``model`` are
|
||||
copied onto the returned object so nothing a caller reads off the completion
|
||||
is lost when it asked to stream.
|
||||
"""
|
||||
"""Re-shape a one-shot ACP response as OpenAI stream chunks (data chunk + usage chunk); response-level
|
||||
attributes other than choices/usage/model are copied onto the result."""
|
||||
choice = completion.choices[0]
|
||||
message = choice.message
|
||||
tool_call_deltas = None
|
||||
if message.tool_calls:
|
||||
tool_call_deltas = []
|
||||
for index, tool_call in enumerate(message.tool_calls):
|
||||
tool_call_deltas.append(
|
||||
SimpleNamespace(
|
||||
index=index,
|
||||
id=getattr(tool_call, "id", None),
|
||||
type=getattr(tool_call, "type", "function"),
|
||||
function=SimpleNamespace(
|
||||
name=getattr(tool_call.function, "name", None),
|
||||
arguments=getattr(tool_call.function, "arguments", None),
|
||||
),
|
||||
)
|
||||
tool_call_deltas = [
|
||||
SimpleNamespace(
|
||||
index=index, id=getattr(tool_call, "id", None), type=getattr(tool_call, "type", "function"),
|
||||
function=SimpleNamespace(name=getattr(tool_call.function, "name", None),
|
||||
arguments=getattr(tool_call.function, "arguments", None)),
|
||||
)
|
||||
|
||||
for index, tool_call in enumerate(message.tool_calls)
|
||||
]
|
||||
delta = SimpleNamespace(
|
||||
role="assistant",
|
||||
content=message.content or None,
|
||||
tool_calls=tool_call_deltas,
|
||||
reasoning_content=getattr(message, "reasoning_content", None),
|
||||
reasoning=getattr(message, "reasoning", None),
|
||||
role="assistant", content=message.content or None, tool_calls=tool_call_deltas,
|
||||
reasoning_content=getattr(message, "reasoning_content", None), reasoning=getattr(message, "reasoning", None),
|
||||
)
|
||||
data_chunk = SimpleNamespace(
|
||||
choices=[
|
||||
SimpleNamespace(
|
||||
index=0,
|
||||
delta=delta,
|
||||
finish_reason=choice.finish_reason,
|
||||
)
|
||||
],
|
||||
model=completion.model,
|
||||
usage=None,
|
||||
)
|
||||
usage_chunk = SimpleNamespace(
|
||||
choices=[],
|
||||
model=completion.model,
|
||||
usage=completion.usage,
|
||||
choices=[SimpleNamespace(index=0, delta=delta, finish_reason=choice.finish_reason)],
|
||||
model=completion.model, usage=None,
|
||||
)
|
||||
usage_chunk = SimpleNamespace(choices=[], model=completion.model, usage=completion.usage)
|
||||
chunks = StreamChunks([data_chunk, usage_chunk])
|
||||
for key, value in vars(completion).items():
|
||||
if key not in ("choices", "usage", "model"):
|
||||
@@ -134,135 +71,83 @@ def completion_to_stream_chunks(completion: SimpleNamespace) -> StreamChunks:
|
||||
return chunks
|
||||
|
||||
|
||||
def build_openai_tool_call(
|
||||
*,
|
||||
call_id: str,
|
||||
name: str,
|
||||
arguments: str,
|
||||
) -> ChatCompletionMessageToolCall:
|
||||
def build_openai_tool_call(*, call_id: str, name: str, arguments: str) -> ChatCompletionMessageToolCall:
|
||||
"""Build an OpenAI-compatible tool-call object for downstream handling."""
|
||||
return ChatCompletionMessageToolCall(
|
||||
id=call_id,
|
||||
call_id=call_id,
|
||||
response_item_id=None,
|
||||
type="function",
|
||||
id=call_id, call_id=call_id, response_item_id=None, type="function",
|
||||
function=Function(name=name, arguments=arguments),
|
||||
)
|
||||
|
||||
|
||||
def tool_specs_from_openai_tools(
|
||||
tools: list[dict[str, Any]] | None,
|
||||
*,
|
||||
allowlist: Iterable[str] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Flatten OpenAI ``tools`` into ``{name, description, parameters}`` specs.
|
||||
def _named_function(container: Any) -> tuple[dict[str, Any], str] | None:
|
||||
"""``(fn, stripped name)`` from ``container["function"]`` when it is a dict with a non-blank name, else None."""
|
||||
fn = container.get("function") if isinstance(container, dict) else None
|
||||
name = fn.get("name") if isinstance(fn, dict) else None
|
||||
return (fn, name.strip()) if isinstance(name, str) and name.strip() else None
|
||||
|
||||
Malformed entries are skipped. When ``allowlist`` is given, only tools whose
|
||||
name is in it survive — that is how a client forwards just Hermes'
|
||||
agent-level tools instead of the whole toolset.
|
||||
"""
|
||||
|
||||
def tool_specs_from_openai_tools(
|
||||
tools: list[dict[str, Any]] | None, *, allowlist: Iterable[str] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Flatten OpenAI ``tools`` into ``{name, description, parameters}`` specs; malformed entries are skipped."""
|
||||
allowed = {str(n).strip() for n in allowlist} if allowlist is not None else None
|
||||
specs: list[dict[str, Any]] = []
|
||||
for t in tools or []:
|
||||
if not isinstance(t, dict):
|
||||
named = _named_function(t)
|
||||
if named is None or (allowed is not None and named[1] not in allowed):
|
||||
continue
|
||||
fn = t.get("function") or {}
|
||||
if not isinstance(fn, dict):
|
||||
continue
|
||||
name = fn.get("name")
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
continue
|
||||
name = name.strip()
|
||||
if allowed is not None and name not in allowed:
|
||||
continue
|
||||
specs.append(
|
||||
{
|
||||
"name": name,
|
||||
"description": fn.get("description", ""),
|
||||
"parameters": fn.get("parameters", {}),
|
||||
}
|
||||
)
|
||||
fn, name = named
|
||||
specs.append({"name": name, "description": fn.get("description", ""), "parameters": fn.get("parameters", {})})
|
||||
return specs
|
||||
|
||||
|
||||
def render_tool_bridge_sections(
|
||||
tools: list[dict[str, Any]] | None,
|
||||
tool_choice: Any = None,
|
||||
*,
|
||||
allowlist: Iterable[str] | None = None,
|
||||
tools: list[dict[str, Any]] | None, tool_choice: Any = None, *, allowlist: Iterable[str] | None = None,
|
||||
) -> list[str]:
|
||||
"""Prompt sections that carry the forwarded tool schemas + choice hint.
|
||||
|
||||
Returns an empty list when no tool survives filtering and no choice hint was
|
||||
requested, so callers can splice the result into their section list
|
||||
unconditionally.
|
||||
"""
|
||||
"""Prompt sections carrying the forwarded tool schemas + choice hint (empty list when neither applies)."""
|
||||
specs = tool_specs_from_openai_tools(tools, allowlist=allowlist)
|
||||
sections: list[str] = []
|
||||
if specs:
|
||||
sections.append(
|
||||
TOOL_CALL_CONTRACT + "\n" + json.dumps(specs, ensure_ascii=False)
|
||||
)
|
||||
sections.append(TOOL_CALL_CONTRACT + "\n" + json.dumps(specs, ensure_ascii=False))
|
||||
if tool_choice is not None:
|
||||
sections.append(f"Tool choice hint: {json.dumps(tool_choice, ensure_ascii=False)}")
|
||||
return sections
|
||||
|
||||
|
||||
def extract_tool_calls_from_text(
|
||||
text: str,
|
||||
) -> tuple[list[ChatCompletionMessageToolCall], str]:
|
||||
"""Pull ``<tool_call>`` blocks out of an ACP response.
|
||||
def _parse_tool_call(raw_json: str, ordinal: int) -> ChatCompletionMessageToolCall | None:
|
||||
"""One ``<tool_call>`` JSON body → tool call, or None when malformed. Missing id → ``acp_call_<ordinal>``."""
|
||||
try:
|
||||
obj = json.loads(raw_json)
|
||||
except Exception:
|
||||
return None
|
||||
named = _named_function(obj)
|
||||
if named is None:
|
||||
return None
|
||||
fn, fn_name = named
|
||||
fn_args = fn.get("arguments", "{}")
|
||||
if not isinstance(fn_args, str):
|
||||
fn_args = json.dumps(fn_args, ensure_ascii=False)
|
||||
call_id = obj.get("id")
|
||||
if not isinstance(call_id, str) or not call_id.strip():
|
||||
call_id = f"acp_call_{ordinal}"
|
||||
return build_openai_tool_call(call_id=call_id, name=fn_name, arguments=fn_args)
|
||||
|
||||
Returns ``(tool_calls, cleaned_text)`` where ``cleaned_text`` is the
|
||||
response with the consumed blocks removed, so the assistant message doesn't
|
||||
show raw JSON to the user.
|
||||
"""
|
||||
|
||||
def extract_tool_calls_from_text(text: str) -> tuple[list[ChatCompletionMessageToolCall], str]:
|
||||
"""Pull ``<tool_call>`` blocks out of an ACP response → ``(tool_calls, cleaned_text)`` with the consumed blocks
|
||||
removed so the assistant message doesn't show raw JSON. Bare-JSON fallback runs only when no XML block parsed."""
|
||||
if not isinstance(text, str) or not text.strip():
|
||||
return [], ""
|
||||
|
||||
extracted: list[ChatCompletionMessageToolCall] = []
|
||||
consumed_spans: list[tuple[int, int]] = []
|
||||
|
||||
def _try_add_tool_call(raw_json: str) -> None:
|
||||
try:
|
||||
obj = json.loads(raw_json)
|
||||
except Exception:
|
||||
return
|
||||
if not isinstance(obj, dict):
|
||||
return
|
||||
fn = obj.get("function")
|
||||
if not isinstance(fn, dict):
|
||||
return
|
||||
fn_name = fn.get("name")
|
||||
if not isinstance(fn_name, str) or not fn_name.strip():
|
||||
return
|
||||
fn_args = fn.get("arguments", "{}")
|
||||
if not isinstance(fn_args, str):
|
||||
fn_args = json.dumps(fn_args, ensure_ascii=False)
|
||||
call_id = obj.get("id")
|
||||
if not isinstance(call_id, str) or not call_id.strip():
|
||||
call_id = f"acp_call_{len(extracted)+1}"
|
||||
|
||||
extracted.append(
|
||||
build_openai_tool_call(
|
||||
call_id=call_id,
|
||||
name=fn_name.strip(),
|
||||
arguments=fn_args,
|
||||
)
|
||||
)
|
||||
|
||||
for m in TOOL_CALL_BLOCK_RE.finditer(text):
|
||||
raw = m.group(1)
|
||||
_try_add_tool_call(raw)
|
||||
consumed_spans.append((m.start(), m.end()))
|
||||
|
||||
# Only try bare-JSON fallback when no XML blocks were found.
|
||||
if not extracted:
|
||||
for m in TOOL_CALL_JSON_RE.finditer(text):
|
||||
raw = m.group(0)
|
||||
_try_add_tool_call(raw)
|
||||
for pattern, group in ((TOOL_CALL_BLOCK_RE, 1), (TOOL_CALL_JSON_RE, 0)):
|
||||
for m in pattern.finditer(text):
|
||||
call = _parse_tool_call(m.group(group), len(extracted) + 1)
|
||||
if call is not None:
|
||||
extracted.append(call)
|
||||
consumed_spans.append((m.start(), m.end()))
|
||||
|
||||
if extracted:
|
||||
break
|
||||
if not consumed_spans:
|
||||
return extracted, text.strip()
|
||||
|
||||
@@ -273,7 +158,6 @@ def extract_tool_calls_from_text(
|
||||
merged.append((start, end))
|
||||
else:
|
||||
merged[-1] = (merged[-1][0], max(merged[-1][1], end))
|
||||
|
||||
parts: list[str] = []
|
||||
cursor = 0
|
||||
for start, end in merged:
|
||||
@@ -282,6 +166,5 @@ def extract_tool_calls_from_text(
|
||||
cursor = max(cursor, end)
|
||||
if cursor < len(text):
|
||||
parts.append(text[cursor:])
|
||||
|
||||
cleaned = "\n".join(p.strip() for p in parts if p and p.strip()).strip()
|
||||
return extracted, cleaned
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
"""Turn-liveness activity tracking for ``AIAgent`` (gateway watchdog + session activity persistence).
|
||||
|
||||
``_touch_activity`` is the single write path; persistence is rate-limited and never raises.
|
||||
Extracted from ``run_agent.py``; every method resolves through ``AIAgent``'s MRO unchanged.
|
||||
"""
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from typing import Optional
|
||||
|
||||
from agent.session_activity import ActivityProvenance
|
||||
|
||||
# Same logger name as the origin module so log records / caplog filters are unchanged.
|
||||
logger = logging.getLogger("run_agent")
|
||||
|
||||
|
||||
def _activity_lock(obj) -> "threading.Lock":
|
||||
"""Lazy per-instance ``_turn_liveness_activity_lock`` (so ``__new__``/SimpleNamespace doubles work)."""
|
||||
_lock = getattr(obj, "_turn_liveness_activity_lock", None)
|
||||
if _lock is None:
|
||||
_lock = threading.Lock()
|
||||
obj._turn_liveness_activity_lock = _lock
|
||||
return _lock
|
||||
|
||||
|
||||
class ActivityTrackingMixin:
|
||||
"""Liveness timestamps/labels and rate-limited session activity persistence."""
|
||||
|
||||
def _liveness_activity_lock(self) -> "threading.Lock":
|
||||
"""Shared lock for the activity clock and its generation counter.
|
||||
|
||||
``_touch_activity`` stamps under it and the liveness watchdog samples/commits under it, so a stall
|
||||
observation can never abort a turn that resumed in between.
|
||||
|
||||
Created lazily so ``AIAgent.__new__``-based test doubles keep working. See #95663.
|
||||
"""
|
||||
return _activity_lock(self)
|
||||
|
||||
def _touch_activity(
|
||||
self, desc: str, *, provenance: Optional[ActivityProvenance] = None,
|
||||
force_persist: bool = False,
|
||||
) -> None:
|
||||
"""Update the last-activity timestamp and description (thread-safe).
|
||||
|
||||
Bumps a monotonic generation under the activity lock so the watchdog can bind a stall observation to
|
||||
the exact ``(generation, timestamp)`` it sampled. Also bridges (rate-limited, best-effort) to the
|
||||
kanban heartbeat when this is a dispatcher-spawned worker, and to the durable SessionDB activity
|
||||
projection. ``provenance`` names special writers (compression); ``force_persist`` bypasses the
|
||||
SessionDB rate limit. Module-level lock helper, not ``self._liveness_activity_lock()``: doubles bind
|
||||
only ``_touch_activity`` (tests/run_agent/test_session_activity_persist.py).
|
||||
|
||||
Bridge is rate-limited (60s) and best-effort — it never raises into the agent loop. See #31752.
|
||||
See #72016, #72039.
|
||||
"""
|
||||
from agent.session_activity import (
|
||||
bound_activity_description, normalize_activity_provenance,
|
||||
reset_session_activity_persist_window,
|
||||
)
|
||||
|
||||
with _activity_lock(self):
|
||||
self._turn_liveness_activity_generation = (
|
||||
getattr(self, "_turn_liveness_activity_generation", 0) + 1
|
||||
)
|
||||
self._last_activity_ts = time.time()
|
||||
self._last_activity_desc = bound_activity_description(desc)
|
||||
self._last_activity_provenance = normalize_activity_provenance(provenance)
|
||||
# Real progress invalidates a reserved abort claim; an in-flight watchdog interrupt must abandon
|
||||
# itself at the final mutation edge.
|
||||
self._turn_liveness_abort_claim = None
|
||||
if os.environ.get("HERMES_KANBAN_TASK"):
|
||||
# Never let the bridge break the loop; this guard covers import-time failures.
|
||||
with suppress(Exception):
|
||||
from tools.kanban_tools import (
|
||||
heartbeat_current_worker_from_env, inject_new_comments_from_env
|
||||
)
|
||||
heartbeat_current_worker_from_env()
|
||||
# Fold new operator notes into the running turn (OUT-OF-BAND steer).
|
||||
inject_new_comments_from_env(self)
|
||||
if force_persist:
|
||||
reset_session_activity_persist_window(self)
|
||||
self._persist_session_activity_if_due()
|
||||
|
||||
def _persist_session_activity_if_due(self) -> None:
|
||||
"""Best-effort durable activity heartbeat for SessionDB consumers.
|
||||
|
||||
Cadence pinned by ``SESSION_ACTIVITY_HEARTBEAT_MIN_INTERVAL_SECONDS`` (config-independent). Fail-open:
|
||||
a failed write never raises into the agent loop.
|
||||
"""
|
||||
session_id = getattr(self, "session_id", None)
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
if not session_id or session_db is None:
|
||||
return
|
||||
touch = getattr(session_db, "touch_session_activity", None)
|
||||
if not callable(touch):
|
||||
return
|
||||
from agent.session_activity import (
|
||||
SESSION_ACTIVITY_HEARTBEAT_MIN_INTERVAL_SECONDS, normalize_activity_provenance
|
||||
)
|
||||
|
||||
now_mono = time.monotonic()
|
||||
last_mono = getattr(self, "_session_activity_last_persist_mono", 0.0)
|
||||
if (now_mono - last_mono) < SESSION_ACTIVITY_HEARTBEAT_MIN_INTERVAL_SECONDS:
|
||||
return
|
||||
self._session_activity_last_persist_mono = now_mono
|
||||
try:
|
||||
touch(
|
||||
session_id,
|
||||
getattr(self, "_last_activity_ts", None),
|
||||
description=getattr(self, "_last_activity_desc", None),
|
||||
provenance=normalize_activity_provenance(
|
||||
getattr(self, "_last_activity_provenance", None)
|
||||
),
|
||||
)
|
||||
except Exception:
|
||||
# Heartbeat is observation-only; never let its I/O break the loop.
|
||||
logger.debug("session activity heartbeat write failed (ignored)", exc_info=True)
|
||||
|
||||
def _reset_activity_labels_after_turn(self) -> None:
|
||||
"""Drop mid-turn activity labels once the turn is no longer running.
|
||||
|
||||
Keeps ``_last_activity_ts`` so idle/watchdog clocks stay continuous across turns; clears description +
|
||||
provenance so idle agents / SessionDB listings stop advertising the last mid-turn stamp.
|
||||
|
||||
See #15654, #72039.
|
||||
"""
|
||||
self._last_activity_desc = ""
|
||||
self._last_activity_provenance = ActivityProvenance.UNKNOWN
|
||||
session_id = getattr(self, "session_id", None)
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
if not session_id or session_db is None:
|
||||
return
|
||||
clear = getattr(session_db, "clear_session_activity_labels", None)
|
||||
if not callable(clear):
|
||||
return
|
||||
with suppress(Exception): # never let durable cleanup I/O break turn teardown
|
||||
clear(session_id)
|
||||
+1609
-2395
File diff suppressed because it is too large
Load Diff
+2449
-3849
File diff suppressed because it is too large
Load Diff
+542
-3107
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,585 @@
|
||||
"""Anthropic credential sources, OAuth flows, and token resolution.
|
||||
|
||||
``resolve_anthropic_token()`` order: ``ANTHROPIC_TOKEN`` / ``CLAUDE_CODE_OAUTH_TOKEN``,
|
||||
``ANTHROPIC_API_KEY``, Hermes-owned OAuth grants in the ``auth.json`` credential
|
||||
pool, then ``~/.claude/.credentials.json`` / macOS Keychain as a borrowed fallback.
|
||||
``~/.hermes/.anthropic_oauth.json`` (Hermes PKCE) and
|
||||
the Claude Code file are *singletons*: ``credential_pool._seed_from_singletons()``
|
||||
re-reads them on every ``load_pool()``, so a failed write here is a failed refresh
|
||||
(``CredentialPersistError``), not a cache miss.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import contextlib
|
||||
import functools
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import secrets
|
||||
import stat
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from hermes_constants import get_hermes_home
|
||||
from agent.secret_scope import get_secret as _get_secret
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_OAUTH_CLIENT_ID = "9d1c250a-e61b-44d9-88ed-5944d1962f5e"
|
||||
# platform.claude.com is the live token host; console.anthropic.com 404s but is kept as a fallback.
|
||||
_OAUTH_TOKEN_URLS = [
|
||||
"https://platform.claude.com/v1/oauth/token", "https://console.anthropic.com/v1/oauth/token"
|
||||
]
|
||||
# Anthropic 429s token-endpoint requests whose UA starts with ``claude-code/`` (or Mozilla); the real CLI uses
|
||||
# bare axios there. Inference (build_anthropic_kwargs) still needs claude-code/.
|
||||
_OAUTH_TOKEN_USER_AGENT = "axios/1.7.9"
|
||||
_OAUTH_REDIRECT_URI = "https://console.anthropic.com/oauth/code/callback"
|
||||
_OAUTH_SCOPES = "org:create_api_key user:profile user:inference"
|
||||
|
||||
|
||||
def _getenv(name: str, default: str = "") -> str:
|
||||
"""Profile-scoped os.getenv for credential reads (fail-closed on unscoped reads when multiplexing)."""
|
||||
val = _get_secret(name, default)
|
||||
return val if val is not None else default
|
||||
|
||||
|
||||
def _first_env(*names: str) -> str:
|
||||
"""First non-blank (stripped) value among *names*, else ''."""
|
||||
return next((v for v in (_getenv(n).strip() for n in names) if v), "")
|
||||
|
||||
|
||||
def _is_oauth_token(key: str) -> bool:
|
||||
"""True for Anthropic OAuth/setup tokens (sk-ant-*, eyJ JWTs, cc-); False for sk-ant-api* Console keys."""
|
||||
if not key or key.startswith("sk-ant-api"):
|
||||
return False
|
||||
return key.startswith(("sk-ant-", "eyJ", "cc-"))
|
||||
|
||||
|
||||
class CredentialPersistError(RuntimeError):
|
||||
"""A rotated single-use credential could not be durably committed. The refresh POST already spent the old
|
||||
refresh token, so a swallowed write failure leaves a consumed pair on disk that later replays as invalid_grant."""
|
||||
|
||||
def __init__(self, path: Any, cause: BaseException) -> None:
|
||||
super().__init__(f"failed to durably persist rotated Anthropic credentials to {path}: {cause}")
|
||||
self.path = path
|
||||
|
||||
|
||||
def _load_json_if_exists(path: Path, what: str) -> Optional[Any]:
|
||||
"""Parsed JSON from *path*, or None when missing/unreadable/corrupt (debug-logged)."""
|
||||
if not path.exists():
|
||||
return None
|
||||
try:
|
||||
return json.loads(path.read_text(encoding="utf-8"))
|
||||
except (json.JSONDecodeError, OSError) as e:
|
||||
logger.debug("Failed to read %s: %s", what, e)
|
||||
return None
|
||||
|
||||
|
||||
def _atomic_write_private_json(path: Path, payload: Any) -> None:
|
||||
"""Write *payload* via a 0o600 O_EXCL temp file + fsync + os.replace: the token is never briefly umask-readable
|
||||
(write_text + chmod had a TOCTOU window); the random suffix avoids collisions with concurrent writers and
|
||||
crashed leftovers. The parent dir's mode is left alone (~/.claude/ is owned by Claude Code)."""
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp = path.with_suffix(f".tmp.{os.getpid()}.{secrets.token_hex(4)}")
|
||||
try:
|
||||
fd = os.open(str(tmp), os.O_WRONLY | os.O_CREAT | os.O_EXCL, stat.S_IRUSR | stat.S_IWUSR)
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as fh:
|
||||
json.dump(payload, fh, indent=2)
|
||||
fh.flush()
|
||||
os.fsync(fh.fileno())
|
||||
os.replace(tmp, path)
|
||||
except OSError:
|
||||
with contextlib.suppress(OSError):
|
||||
tmp.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
def _commit_private_json(path: Path, payload: Any, what: str) -> None:
|
||||
"""Atomic private write; any failure becomes ``CredentialPersistError`` (the commit step of a rotation)."""
|
||||
try:
|
||||
_atomic_write_private_json(path, payload)
|
||||
except (OSError, ValueError) as e:
|
||||
logger.error("Failed to write refreshed %s to %s: %s", what, path, e)
|
||||
raise CredentialPersistError(path, e) from e
|
||||
|
||||
|
||||
# ── Spent-rotation registry: fingerprints of secrets whose refresh POST succeeded but whose replacement never
|
||||
# reached its store. Two scopes: process-local (OrderedDict) and a durable sidecar next to the shared singleton
|
||||
# file so OTHER processes fail closed too. Non-reversible digests; never cleared.
|
||||
_SPENT_ROTATION_LOCK = threading.Lock()
|
||||
_SPENT_ROTATION_FINGERPRINTS: "OrderedDict[str, None]" = OrderedDict()
|
||||
_SPENT_ROTATION_MAX_TRACKED = 64
|
||||
_SPENT_ROTATION_SIDECAR_COMMENT = (
|
||||
"Non-secret one-way fingerprints of Anthropic OAuth credentials whose rotation was "
|
||||
"consumed server-side but never durably committed. Written by Hermes so sibling "
|
||||
"processes sharing this credential source fail closed instead of replaying a spent "
|
||||
"single-use refresh token."
|
||||
)
|
||||
|
||||
|
||||
def _spent_rotation_sidecar_path(source_path: Path) -> Path:
|
||||
return source_path.with_name(source_path.name + ".hermes-spent-rotations.json")
|
||||
|
||||
|
||||
def spent_rotation_source_path(source: Any) -> Optional[Path]:
|
||||
"""Map a pool-entry source to the shared singleton file it borrows from (or None)."""
|
||||
getter = _SINGLETON_SOURCE_PATHS.get(source) if isinstance(source, str) else None
|
||||
return getter() if getter else None
|
||||
|
||||
|
||||
def _read_spent_rotation_sidecar(source_path: Optional[Path]) -> set:
|
||||
if source_path is None:
|
||||
return set()
|
||||
try:
|
||||
raw = json.loads(_spent_rotation_sidecar_path(source_path).read_text(encoding="utf-8"))
|
||||
except (OSError, ValueError):
|
||||
return set()
|
||||
fingerprints = raw.get("fingerprints") if isinstance(raw, dict) else None
|
||||
return {fp for fp in fingerprints if isinstance(fp, str) and fp} if isinstance(fingerprints, list) else set()
|
||||
|
||||
|
||||
def _append_spent_rotation_sidecar(source_path: Path, fingerprints: list) -> None:
|
||||
"""Merge fingerprints into the sidecar (atomic replace; caller holds the path lock). Fail-soft: a sidecar
|
||||
write failure must never mask the process-local verdict."""
|
||||
sidecar = _spent_rotation_sidecar_path(source_path)
|
||||
try:
|
||||
merged = _read_spent_rotation_sidecar(source_path)
|
||||
merged.update(fingerprints)
|
||||
payload = json.dumps({
|
||||
"version": 1,
|
||||
"comment": _SPENT_ROTATION_SIDECAR_COMMENT,
|
||||
"fingerprints": sorted(merged)[-_SPENT_ROTATION_MAX_TRACKED * 4 :],
|
||||
}, indent=2)
|
||||
sidecar.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp = sidecar.with_name(sidecar.name + ".tmp")
|
||||
tmp.write_text(payload, encoding="utf-8")
|
||||
os.replace(tmp, sidecar)
|
||||
except Exception:
|
||||
logger.debug("Failed to persist spent-rotation fingerprints to %s", sidecar, exc_info=True)
|
||||
|
||||
|
||||
def _fingerprint(secret: Any) -> Optional[str]:
|
||||
from agent.credential_persistence import fingerprint_secret_value
|
||||
value = str(secret or "").strip()
|
||||
return fingerprint_secret_value(value) if value else None
|
||||
|
||||
|
||||
def mark_rotation_consumed_uncommitted(*secrets: Any, source_path: Optional[Path] = None) -> None:
|
||||
"""Record the pre-rotation pair of a refresh whose replacement never committed; with ``source_path`` the
|
||||
verdict is also persisted to that singleton's sidecar."""
|
||||
recorded = [fp for fp in map(_fingerprint, secrets) if fp]
|
||||
with _SPENT_ROTATION_LOCK:
|
||||
for fingerprint in recorded:
|
||||
_SPENT_ROTATION_FINGERPRINTS.pop(fingerprint, None)
|
||||
_SPENT_ROTATION_FINGERPRINTS[fingerprint] = None
|
||||
while len(_SPENT_ROTATION_FINGERPRINTS) > _SPENT_ROTATION_MAX_TRACKED:
|
||||
_SPENT_ROTATION_FINGERPRINTS.popitem(last=False)
|
||||
if recorded and source_path is not None:
|
||||
_append_spent_rotation_sidecar(source_path, recorded)
|
||||
|
||||
|
||||
def is_rotation_consumed_uncommitted(secret: Any, *, source_path: Optional[Path] = None) -> bool:
|
||||
"""True when *secret* belongs to a rotation that was spent but not committed."""
|
||||
fingerprint = _fingerprint(secret)
|
||||
if not fingerprint:
|
||||
return False
|
||||
with _SPENT_ROTATION_LOCK:
|
||||
if fingerprint in _SPENT_ROTATION_FINGERPRINTS:
|
||||
return True
|
||||
return fingerprint in _read_spent_rotation_sidecar(source_path)
|
||||
|
||||
|
||||
# ── Claude Code credentials (Keychain / ~/.claude/.credentials.json) ──
|
||||
# Only singleton-backed pool sources have a cross-process authority boundary.
|
||||
_SINGLETON_SOURCE_PATHS = {
|
||||
"claude_code": lambda: claude_code_credentials_path(), "hermes_pkce": lambda: _get_hermes_oauth_file()
|
||||
}
|
||||
|
||||
|
||||
def _claude_oauth_record(data: Any, source: str) -> Optional[Dict[str, Any]]:
|
||||
"""Normalise a ``{"claudeAiOauth": {...}}`` payload into our credential dict."""
|
||||
oauth_data = data.get("claudeAiOauth")
|
||||
access_token = oauth_data.get("accessToken", "") if isinstance(oauth_data, dict) else ""
|
||||
if not access_token:
|
||||
return None
|
||||
return {
|
||||
"accessToken": access_token, "refreshToken": oauth_data.get("refreshToken", ""),
|
||||
"expiresAt": oauth_data.get("expiresAt", 0), "source": source,
|
||||
}
|
||||
|
||||
|
||||
def _read_claude_code_credentials_from_keychain() -> Optional[Dict[str, Any]]:
|
||||
"""Read the "Claude Code-credentials" macOS Keychain entry (Claude Code >=2.1.114)."""
|
||||
if platform.system() != "Darwin":
|
||||
return None
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["security", "find-generic-password", "-s", "Claude Code-credentials", "-w"],
|
||||
capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=5, stdin=subprocess.DEVNULL,
|
||||
)
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
logger.debug("Keychain: security command not available or timed out")
|
||||
return None
|
||||
if result.returncode != 0:
|
||||
logger.debug("Keychain: no entry found for 'Claude Code-credentials'")
|
||||
return None
|
||||
raw = result.stdout.strip()
|
||||
try:
|
||||
return _claude_oauth_record(json.loads(raw), "macos_keychain") if raw else None
|
||||
except json.JSONDecodeError:
|
||||
logger.debug("Keychain: credentials payload is not valid JSON")
|
||||
return None
|
||||
|
||||
|
||||
def claude_code_credentials_path() -> Path:
|
||||
"""Claude Code's shared OAuth file; every profile reads/writes this same path."""
|
||||
return Path.home() / ".claude" / ".credentials.json"
|
||||
|
||||
|
||||
def _read_claude_code_credentials_from_file() -> Optional[Dict[str, Any]]:
|
||||
data = _load_json_if_exists(claude_code_credentials_path(), "~/.claude/.credentials.json")
|
||||
return _claude_oauth_record(data, "claude_code_credentials_file") if data is not None else None
|
||||
|
||||
|
||||
def read_claude_code_credentials() -> Optional[Dict[str, Any]]:
|
||||
"""Read refreshable Claude Code OAuth credentials (Keychain and/or file). When both exist: prefer the only
|
||||
non-expired one (Claude Code 2.1.x refreshes one source but not the other), else the later ``expiresAt`` so a
|
||||
refresh uses the freshest refreshToken. ~/.claude.json primaryApiKey is deliberately excluded."""
|
||||
kc_creds = _read_claude_code_credentials_from_keychain()
|
||||
file_creds = _read_claude_code_credentials_from_file()
|
||||
if not (kc_creds and file_creds):
|
||||
return kc_creds or file_creds
|
||||
kc_valid, file_valid = is_claude_code_token_valid(kc_creds), is_claude_code_token_valid(file_creds)
|
||||
if kc_valid != file_valid:
|
||||
return kc_creds if kc_valid else file_creds
|
||||
return kc_creds if (kc_creds.get("expiresAt", 0) or 0) >= (file_creds.get("expiresAt", 0) or 0) else file_creds
|
||||
|
||||
|
||||
def is_claude_code_token_valid(creds: Dict[str, Any]) -> bool:
|
||||
"""Non-expired access token (60s buffer); no expiresAt means managed key → valid if present."""
|
||||
expires_at = creds.get("expiresAt", 0)
|
||||
return int(time.time() * 1000) < (expires_at - 60_000) if expires_at else bool(creds.get("accessToken"))
|
||||
|
||||
|
||||
# ── OAuth token endpoint ──
|
||||
|
||||
|
||||
def _post_oauth_token(
|
||||
data: bytes, *, content_type: str, timeout: int, what: str, user_agent: str = _OAUTH_TOKEN_USER_AGENT
|
||||
) -> Dict[str, Any]:
|
||||
"""POST to the token endpoints in order; raise the last error if all fail."""
|
||||
import urllib.request
|
||||
last_error = None
|
||||
for endpoint in _OAUTH_TOKEN_URLS:
|
||||
req = urllib.request.Request(
|
||||
endpoint, data=data, method="POST", headers={"Content-Type": content_type, "User-Agent": user_agent}
|
||||
)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=timeout) as resp:
|
||||
return json.loads(resp.read().decode())
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
logger.debug("Anthropic token %s failed at %s: %s", what, endpoint, exc)
|
||||
raise last_error or ValueError(f"Anthropic token {what} failed")
|
||||
|
||||
|
||||
def _oauth_token_state(result: Dict[str, Any], *, fallback_refresh_token: str = "") -> Dict[str, Any]:
|
||||
"""Token-endpoint JSON -> ``{access_token, refresh_token, expires_at_ms}`` (expires_in defaults to 3600s)."""
|
||||
return {
|
||||
"access_token": result.get("access_token", ""),
|
||||
"refresh_token": result.get("refresh_token", fallback_refresh_token),
|
||||
"expires_at_ms": int(time.time() * 1000) + (result.get("expires_in", 3600) * 1000),
|
||||
}
|
||||
|
||||
|
||||
def refresh_anthropic_oauth_pure(refresh_token: str, *, use_json: bool = False) -> Dict[str, Any]:
|
||||
"""Refresh an Anthropic OAuth token without mutating local credential files."""
|
||||
import urllib.parse
|
||||
if not refresh_token:
|
||||
raise ValueError("refresh_token is required")
|
||||
payload = {"grant_type": "refresh_token", "refresh_token": refresh_token, "client_id": _OAUTH_CLIENT_ID}
|
||||
encode, content_type = ((json.dumps, "application/json") if use_json
|
||||
else (urllib.parse.urlencode, "application/x-www-form-urlencoded"))
|
||||
result = _post_oauth_token(encode(payload).encode(), content_type=content_type, timeout=10, what="refresh",
|
||||
user_agent=_OAUTH_TOKEN_USER_AGENT)
|
||||
if not result.get("access_token"):
|
||||
raise ValueError("Anthropic refresh response was missing access_token")
|
||||
return _oauth_token_state(result, fallback_refresh_token=refresh_token)
|
||||
|
||||
|
||||
def _refresh_oauth_token(creds: Dict[str, Any]) -> Optional[str]:
|
||||
"""Refresh an expired Claude Code OAuth token, returning the new access token. Refresh tokens are single-use and
|
||||
Claude Code refreshes on its own schedule, so we first re-read the live sources and adopt an already-rotated
|
||||
token instead of racing it into ``invalid_grant``. Read, decision, POST and write-back share the pool's
|
||||
path-keyed cross-process lock (else two profiles can spend one refresh token)."""
|
||||
try:
|
||||
from hermes_cli.auth import AUTH_LOCK_TIMEOUT_SECONDS, _auth_store_lock, env_float
|
||||
refresh_timeout_seconds = env_float("HERMES_ANTHROPIC_REFRESH_TIMEOUT_SECONDS", 20)
|
||||
lock_timeout_seconds = max(float(AUTH_LOCK_TIMEOUT_SECONDS), float(refresh_timeout_seconds) + 5.0)
|
||||
cred_path = claude_code_credentials_path()
|
||||
with _auth_store_lock(timeout_seconds=lock_timeout_seconds, target_path=cred_path):
|
||||
# Adopt only a DIFFERENT token with a real future expiry (0/absent expiresAt = managed key/unknown).
|
||||
current = read_claude_code_credentials() or {}
|
||||
current_token = current.get("accessToken", "")
|
||||
if (current_token and current_token != creds.get("accessToken", "")
|
||||
and (current.get("expiresAt", 0) or 0) > 0 and is_claude_code_token_valid(current)):
|
||||
logger.debug("Adopted Claude Code's already-refreshed OAuth token")
|
||||
return current_token
|
||||
|
||||
refresh_token = current.get("refreshToken", "") or creds.get("refreshToken", "")
|
||||
if not refresh_token:
|
||||
logger.debug("No refresh token available — cannot refresh")
|
||||
return None
|
||||
# Another process may have spent this token and lost the commit; its sidecar verdict is authoritative.
|
||||
if is_rotation_consumed_uncommitted(refresh_token, source_path=cred_path):
|
||||
logger.debug("Refresh token was already consumed by an uncommitted rotation "
|
||||
"- refusing to replay it; re-run 'claude setup-token'")
|
||||
return None
|
||||
try:
|
||||
refreshed = refresh_anthropic_oauth_pure(refresh_token, use_json=False)
|
||||
except Exception as e:
|
||||
logger.debug("Failed to refresh Claude Code token: %s", e)
|
||||
return None
|
||||
# The POST spent ``refresh_token``; this write is the commit step. On failure, fail closed and
|
||||
# mark the pre-rotation pair as spent.
|
||||
try:
|
||||
_write_claude_code_credentials(refreshed["access_token"], refreshed["refresh_token"], refreshed["expires_at_ms"])
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Anthropic OAuth refresh rotated the single-use token but could not "
|
||||
"commit it to %s (%s) — treating the refresh as failed; "
|
||||
"re-run 'claude setup-token' to reauthenticate",
|
||||
cred_path, e,
|
||||
)
|
||||
mark_rotation_consumed_uncommitted(
|
||||
refresh_token, creds.get("accessToken", ""), current.get("accessToken", ""),
|
||||
current.get("refreshToken", ""), source_path=cred_path,
|
||||
)
|
||||
return None
|
||||
logger.debug("Successfully refreshed Claude Code OAuth token")
|
||||
return refreshed["access_token"]
|
||||
except Exception as e:
|
||||
# Lock/read failures keep the resolver's fail-soft contract.
|
||||
logger.debug("Failed to acquire Claude Code refresh lock: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def _write_claude_code_credentials(
|
||||
access_token: str, refresh_token: str, expires_at_ms: int, *, scopes: Optional[list] = None
|
||||
) -> None:
|
||||
"""Commit refreshed credentials to ~/.claude/.credentials.json; ``CredentialPersistError`` on any failure (a
|
||||
corrupt existing file included). *scopes* (or the previously stored scopes) are persisted because Claude Code
|
||||
>=2.1.81 gates on ``"user:inference"`` being present."""
|
||||
cred_path = claude_code_credentials_path()
|
||||
try:
|
||||
existing = json.loads(cred_path.read_text(encoding="utf-8")) if cred_path.exists() else {}
|
||||
except (OSError, ValueError) as e:
|
||||
logger.error("Failed to write refreshed credentials to %s: %s", cred_path, e)
|
||||
raise CredentialPersistError(cred_path, e) from e
|
||||
oauth_data: Dict[str, Any] = {"accessToken": access_token, "refreshToken": refresh_token, "expiresAt": expires_at_ms}
|
||||
if scopes is not None:
|
||||
oauth_data["scopes"] = scopes
|
||||
elif "claudeAiOauth" in existing and "scopes" in existing["claudeAiOauth"]:
|
||||
oauth_data["scopes"] = existing["claudeAiOauth"]["scopes"]
|
||||
existing["claudeAiOauth"] = oauth_data
|
||||
_commit_private_json(cred_path, existing, "credentials")
|
||||
|
||||
|
||||
# ── Resolution ──
|
||||
|
||||
|
||||
def _resolve_claude_code_token_from_credentials(creds: Optional[Dict[str, Any]] = None) -> Optional[str]:
|
||||
"""Resolve a token from Claude Code credential files, refreshing if needed."""
|
||||
creds = creds or read_claude_code_credentials()
|
||||
if not creds:
|
||||
return None
|
||||
if is_rotation_consumed_uncommitted(creds.get("accessToken", ""), source_path=claude_code_credentials_path()):
|
||||
# The file still holds the spent pre-rotation copy of a failed commit.
|
||||
logger.debug("Claude Code credentials hold a rotated-but-uncommitted token - refusing")
|
||||
return None
|
||||
if is_claude_code_token_valid(creds):
|
||||
logger.debug("Using Claude Code credentials (auto-detected)")
|
||||
return creds["accessToken"]
|
||||
logger.debug("Claude Code credentials expired — attempting refresh")
|
||||
refreshed = _refresh_oauth_token(creds)
|
||||
if not refreshed:
|
||||
logger.debug("Token refresh failed — re-run 'claude setup-token' to reauthenticate")
|
||||
return refreshed or None
|
||||
|
||||
|
||||
def _prefer_refreshable_claude_code_token(env_token: str, creds: Optional[Dict[str, Any]]) -> Optional[str]:
|
||||
"""Prefer refreshable Claude Code creds over a static env OAuth token: Hermes historically persisted setup tokens
|
||||
into ANTHROPIC_TOKEN, and that static token would otherwise win before the refreshable file is inspected."""
|
||||
if not (env_token and _is_oauth_token(env_token) and isinstance(creds, dict) and creds.get("refreshToken")):
|
||||
return None
|
||||
resolved = _resolve_claude_code_token_from_credentials(creds)
|
||||
if resolved and resolved != env_token:
|
||||
logger.debug("Preferring Claude Code credential file over static env OAuth token so refresh can proceed")
|
||||
return resolved
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_anthropic_pool_token(*, skip_borrowed: bool = False) -> Optional[str]:
|
||||
"""First available Anthropic OAuth token from credential_pool, read-only: enumerates with ``clear_expired=False,
|
||||
refresh=False`` (never ``select()``) so diagnostic call sites (account_usage, ``hermes models``) never mutate
|
||||
auth.json or hit the network; refresh-on-expiry belongs to the API call path's pool recovery."""
|
||||
try:
|
||||
from agent.credential_pool import AUTH_TYPE_OAUTH, load_pool
|
||||
entries, _pending = load_pool("anthropic")._available_entries(clear_expired=False, refresh=False)
|
||||
except Exception:
|
||||
logger.debug("Failed to read Anthropic credential_pool", exc_info=True)
|
||||
return None
|
||||
for entry in entries:
|
||||
if skip_borrowed and entry.source == "claude_code":
|
||||
continue
|
||||
# access_token may be an explicit null on a persisted entry; None.strip() would crash the resolver.
|
||||
token = (getattr(entry, "access_token", None) or "").strip()
|
||||
if getattr(entry, "auth_type", None) != AUTH_TYPE_OAUTH or not token:
|
||||
continue
|
||||
# load_pool() re-seeds rows from the singleton files, so a spent-but-uncommitted rotation
|
||||
# (possibly from another process) looks healthy here.
|
||||
entry_source_path = spent_rotation_source_path(getattr(entry, "source", None))
|
||||
if any(
|
||||
is_rotation_consumed_uncommitted(secret, source_path=entry_source_path)
|
||||
for secret in (token, getattr(entry, "refresh_token", None))
|
||||
):
|
||||
logger.debug("Skipping Anthropic pool entry %s: rotated-but-uncommitted credential", getattr(entry, "id", "?"))
|
||||
continue
|
||||
return token
|
||||
return None
|
||||
|
||||
|
||||
def resolve_anthropic_token() -> Optional[str]:
|
||||
"""Resolve an Anthropic token from all sources in priority order (see module docstring)."""
|
||||
_read_creds = functools.cache(read_claude_code_credentials) # read the file at most once per resolve
|
||||
token = _first_env("ANTHROPIC_TOKEN", "CLAUDE_CODE_OAUTH_TOKEN")
|
||||
if token:
|
||||
return _prefer_refreshable_claude_code_token(token, _read_creds()) or token
|
||||
api_key = _first_env("ANTHROPIC_API_KEY") # an explicit API key must not be shadowed by discovered OAuth creds
|
||||
if api_key:
|
||||
return api_key
|
||||
# The pool's claude_code row mirrors the same externally owned refresh grant.
|
||||
return _resolve_anthropic_pool_token(skip_borrowed=True) or _resolve_claude_code_token_from_credentials(_read_creds())
|
||||
|
||||
|
||||
def run_oauth_setup_token() -> Optional[str]:
|
||||
"""Run 'claude setup-token' interactively; the resulting token or None. FileNotFoundError if no 'claude' CLI."""
|
||||
import shutil
|
||||
claude_path = shutil.which("claude")
|
||||
if not claude_path:
|
||||
raise FileNotFoundError("The 'claude' CLI is not installed. Install it with: npm install -g @anthropic-ai/claude-code")
|
||||
# Interactive: stdio inherited so the user can complete the OAuth prompt. noqa: subprocess-stdin
|
||||
try:
|
||||
subprocess.run([claude_path, "setup-token"])
|
||||
except (KeyboardInterrupt, EOFError):
|
||||
return None
|
||||
creds = read_claude_code_credentials()
|
||||
if creds and is_claude_code_token_valid(creds):
|
||||
return creds["accessToken"]
|
||||
return _first_env("CLAUDE_CODE_OAUTH_TOKEN", "ANTHROPIC_TOKEN") or None
|
||||
|
||||
|
||||
# ── Hermes-native PKCE OAuth flow (~/.hermes/.anthropic_oauth.json); mirrors Claude Code / pi-ai / OpenCode ──
|
||||
|
||||
|
||||
def _get_hermes_oauth_file() -> Path:
|
||||
return get_hermes_home() / ".anthropic_oauth.json"
|
||||
|
||||
|
||||
def _root_hermes_oauth_file() -> Optional[Path]:
|
||||
"""Global-root ``.anthropic_oauth.json`` inside a named profile (None in classic mode); used to commit a
|
||||
rotation of a grant the profile borrowed via the pool's root fallback."""
|
||||
try:
|
||||
from hermes_constants import get_default_hermes_root
|
||||
root = get_default_hermes_root()
|
||||
return None if root.resolve(strict=False) == get_hermes_home().resolve(strict=False) else root / ".anthropic_oauth.json"
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _generate_pkce() -> tuple:
|
||||
"""Generate PKCE code_verifier and code_challenge (S256)."""
|
||||
verifier = base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode()
|
||||
challenge = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
|
||||
return verifier, challenge
|
||||
|
||||
|
||||
def run_hermes_oauth_login_pure() -> Optional[Dict[str, Any]]:
|
||||
"""Run Hermes-native OAuth PKCE flow and return credential state."""
|
||||
import webbrowser
|
||||
from urllib.parse import urlencode
|
||||
verifier, challenge = _generate_pkce()
|
||||
oauth_state = secrets.token_urlsafe(32)
|
||||
params = {
|
||||
"code": "true", "client_id": _OAUTH_CLIENT_ID, "response_type": "code", "redirect_uri": _OAUTH_REDIRECT_URI,
|
||||
"scope": _OAUTH_SCOPES, "code_challenge": challenge, "code_challenge_method": "S256", "state": oauth_state,
|
||||
}
|
||||
auth_url = f"https://claude.ai/oauth/authorize?{urlencode(params)}"
|
||||
print("\n".join([
|
||||
"", "Authorize Hermes with your Claude Pro/Max subscription.", "",
|
||||
"╭─ Claude Pro/Max Authorization ────────────────────╮",
|
||||
"│ │",
|
||||
"│ Open this link in your browser: │",
|
||||
"╰───────────────────────────────────────────────────╯",
|
||||
"", f" {auth_url}", "",
|
||||
]))
|
||||
try:
|
||||
from hermes_cli.auth import _can_open_graphical_browser as _can_open_gui
|
||||
except Exception:
|
||||
_can_open_gui = lambda: True # noqa: E731 — degrade to prior behavior
|
||||
if _can_open_gui():
|
||||
with contextlib.suppress(Exception):
|
||||
webbrowser.open(auth_url)
|
||||
print(" (Browser opened automatically)")
|
||||
print("\nAfter authorizing, you'll see a code. Paste it below.\n")
|
||||
try:
|
||||
auth_code = input("Authorization code: ").strip()
|
||||
except (KeyboardInterrupt, EOFError):
|
||||
return None
|
||||
if not auth_code:
|
||||
print("No code entered.")
|
||||
return None
|
||||
splits = auth_code.split("#")
|
||||
code, received_state = splits[0], (splits[1] if len(splits) > 1 else "")
|
||||
if received_state != oauth_state: # CSRF guard (RFC 6749 §10.12)
|
||||
logger.warning("OAuth state mismatch — possible CSRF, aborting")
|
||||
return None
|
||||
try:
|
||||
exchange_data = json.dumps({
|
||||
"grant_type": "authorization_code", "client_id": _OAUTH_CLIENT_ID, "code": code, "state": received_state,
|
||||
"redirect_uri": _OAUTH_REDIRECT_URI, "code_verifier": verifier,
|
||||
}).encode()
|
||||
result = _post_oauth_token(exchange_data, content_type="application/json", timeout=15, what="exchange")
|
||||
except Exception as e:
|
||||
print(f"Token exchange failed: {e}")
|
||||
return None
|
||||
if not result.get("access_token"):
|
||||
print("No access token in response.")
|
||||
return None
|
||||
return _oauth_token_state(result)
|
||||
|
||||
|
||||
def read_hermes_oauth_credentials() -> Optional[Dict[str, Any]]:
|
||||
"""Read Hermes-managed OAuth credentials from ~/.hermes/.anthropic_oauth.json."""
|
||||
data = _load_json_if_exists(_get_hermes_oauth_file(), "Hermes OAuth credentials")
|
||||
return data if data is not None and data.get("accessToken") else None
|
||||
|
||||
|
||||
def _write_hermes_oauth_credentials(
|
||||
access_token: str, refresh_token: Optional[str], expires_at_ms: Optional[int], *, target: Optional[Path] = None
|
||||
) -> None:
|
||||
"""Commit refreshed hermes_pkce tokens to ~/.hermes/.anthropic_oauth.json (``CredentialPersistError`` on failure).
|
||||
``target`` lets a named profile commit a grant it BORROWED from the global root back to the ROOT singleton
|
||||
instead of forking a copy under its own HERMES_HOME; without this write-through the next ``load_pool()``
|
||||
re-seeds the stale (consumed) pair from the file over the rotated pool entry."""
|
||||
_commit_private_json(
|
||||
target if target is not None else _get_hermes_oauth_file(),
|
||||
{"accessToken": access_token, "refreshToken": refresh_token, "expiresAt": expires_at_ms},
|
||||
"Hermes OAuth credentials",
|
||||
)
|
||||
@@ -0,0 +1,159 @@
|
||||
"""Endpoint-family detection for Anthropic-compatible base URLs.
|
||||
|
||||
A dozen services speak the Anthropic Messages API but differ in auth style, accepted beta
|
||||
headers, and request quirks (MiniMax, Kimi/Moonshot, DeepSeek, OpenCode, Azure AI Foundry, Nous
|
||||
Portal, Bedrock). Every such difference is decided from the configured base URL, so the
|
||||
predicates live together here as pure functions (no I/O, SDK or credentials) that both
|
||||
``agent/anthropic_adapter.py`` and ``agent/anthropic_message_convert.py`` can import without a
|
||||
cycle.
|
||||
"""
|
||||
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from utils import base_url_host_matches, base_url_hostname
|
||||
|
||||
_MINIMAX_ANTHROPIC_PREFIXES = ("https://api.minimax.io/anthropic", "https://api.minimaxi.com/anthropic")
|
||||
|
||||
|
||||
def _normalize_base_url_text(base_url) -> str:
|
||||
"""Coerce a base URL (str or ``httpx.URL``) to a stripped string; "" when falsy."""
|
||||
return str(base_url).strip() if base_url else ""
|
||||
|
||||
|
||||
def _normalized_lower(base_url) -> str:
|
||||
"""``_normalize_base_url_text`` + rstrip("/") + lower(), the shape most predicates match on."""
|
||||
return _normalize_base_url_text(base_url).rstrip("/").lower()
|
||||
|
||||
|
||||
def _is_third_party_anthropic_endpoint(base_url: str | None) -> bool:
|
||||
"""Any non-anthropic.com endpoint (own x-api-key keys; skip OAuth detection). No base_url =
|
||||
direct Anthropic API."""
|
||||
normalized = _normalized_lower(base_url)
|
||||
return bool(normalized) and "anthropic.com" not in normalized
|
||||
|
||||
|
||||
def _is_kimi_coding_endpoint(base_url: str | None) -> bool:
|
||||
"""Kimi's /coding endpoint, which requires a claude-code User-Agent."""
|
||||
return _normalized_lower(base_url).startswith("https://api.kimi.com/coding")
|
||||
|
||||
|
||||
def _is_opencode_endpoint(base_url: str | None) -> bool:
|
||||
"""OpenCode's Zen/Go relay (opencode.ai)."""
|
||||
return base_url_host_matches(base_url or "", "opencode.ai")
|
||||
|
||||
|
||||
# Kimi / Moonshot family model-name prefixes: official slugs (``kimi-k2.5``, ``kimi_thinking``,
|
||||
# ``moonshot-v1-8k``) and release lines (``k1.5-…``, ``k2-thinking``, ``k25-…``, ``k3.x``/``k3-…``).
|
||||
# Matched case-insensitively after stripping any ``vendor/`` prefix.
|
||||
_KIMI_FAMILY_MODEL_PREFIXES = (
|
||||
"kimi-", "kimi_", "moonshot-", "moonshot_", "k1.", "k1-", "k2.", "k2-", "k25", "k2.5", "k3.", "k3-",
|
||||
)
|
||||
# Bare release slugs with no separator suffix (Kimi Coding Plan serves K3 as exactly ``k3``);
|
||||
# exact-match so unrelated names sharing the prefix don't match.
|
||||
_KIMI_FAMILY_EXACT_SLUGS = frozenset({"k3"})
|
||||
|
||||
|
||||
def _model_name_is_kimi_family(model: str | None) -> bool:
|
||||
if not isinstance(model, str):
|
||||
return False
|
||||
m = model.strip().lower().rsplit("/", 1)[-1] # ``moonshotai/kimi-k2.5`` -> ``kimi-k2.5``
|
||||
return bool(m) and (m in _KIMI_FAMILY_EXACT_SLUGS or m.startswith(_KIMI_FAMILY_MODEL_PREFIXES))
|
||||
|
||||
|
||||
def _is_kimi_family_endpoint(base_url: str | None, model: str | None = None) -> bool:
|
||||
"""Any Kimi / Moonshot Anthropic-Messages endpoint: the /coding endpoint, any api.kimi.com /
|
||||
moonshot.ai / moonshot.cn host, or any endpoint (e.g. a private gateway) whose *model* is in
|
||||
the Kimi family — the upstream enforces Kimi's thinking semantics regardless of hostname.
|
||||
Decides whether unsigned reasoning_content-derived thinking blocks are preserved on replay."""
|
||||
return (
|
||||
_is_kimi_coding_endpoint(base_url)
|
||||
or any(base_url_host_matches(base_url or "", d) for d in ("api.kimi.com", "moonshot.ai", "moonshot.cn"))
|
||||
or _model_name_is_kimi_family(model)
|
||||
)
|
||||
|
||||
|
||||
|
||||
_DEEPSEEK_THINKING_MODEL_PREFIXES = (
|
||||
"deepseek-r", "deepseek-v4", "deepseek_v4", "deepseek-pro",
|
||||
"deepseek_pro", "deepseek-flash", "deepseek_flash",
|
||||
)
|
||||
|
||||
|
||||
def _model_name_is_deepseek_thinking(model: str | None) -> bool:
|
||||
"""Known DeepSeek thinking families behind an Anthropic-compatible relay.
|
||||
|
||||
Strip vendor namespaces, but do not treat arbitrary DeepSeek chat/distill
|
||||
names as evidence of the thinking replay contract.
|
||||
"""
|
||||
if not isinstance(model, str):
|
||||
return False
|
||||
name = model.strip().lower().rsplit("/", 1)[-1]
|
||||
return bool(name) and name.startswith(_DEEPSEEK_THINKING_MODEL_PREFIXES)
|
||||
|
||||
|
||||
def _is_deepseek_anthropic_endpoint(base_url: str | None) -> bool:
|
||||
"""DeepSeek's ``/anthropic`` route. In thinking mode DeepSeek requires prior-turn ``thinking``
|
||||
blocks to round-trip while the generic third-party path strips them; its blocks are unsigned,
|
||||
so it gets the same strip-signed / keep-unsigned policy as Kimi. Pinned to the ``/anthropic``
|
||||
path so the OpenAI-compatible base URL is not misclassified.
|
||||
|
||||
Per DeepSeek's published compatibility matrix the blocks are unsigned (no Anthropic-proprietary
|
||||
signature, no ``redacted_thinking`` support), so this endpoint is handled with the same strip-signed /
|
||||
keep-unsigned policy used for Kimi's ``/coding`` endpoint. See hermes-agent#16748.
|
||||
"""
|
||||
return base_url_host_matches(base_url or "", "api.deepseek.com") and "/anthropic" in _normalized_lower(base_url)
|
||||
|
||||
|
||||
def _is_nous_portal_endpoint(base_url: str | None) -> bool:
|
||||
"""Nous Portal's Anthropic Messages route (Bearer JWT, verbatim catalog ids, native
|
||||
thinking-signature replay). Trusted hosts only: prod ``inference-api.nousresearch.com`` or the
|
||||
operator-set ``NOUS_INFERENCE_BASE_URL`` host (exact hostname equality, so neither lookalike
|
||||
domains nor sibling hosts of the override match)."""
|
||||
if base_url_host_matches(base_url or "", "inference-api.nousresearch.com"):
|
||||
return True
|
||||
try:
|
||||
from hermes_cli.auth import _nous_inference_env_override
|
||||
override = _nous_inference_env_override()
|
||||
except Exception:
|
||||
return False
|
||||
override_host = base_url_hostname(override) if override else ""
|
||||
return bool(override_host) and base_url_hostname(base_url or "") == override_host
|
||||
|
||||
|
||||
def _requires_bearer_auth(base_url: str | None) -> bool:
|
||||
"""Providers needing ``Authorization: Bearer`` instead of ``x-api-key``: MiniMax, Azure AI
|
||||
Foundry, Palantir Foundry's LLM proxy, CommandCode, Nous Portal. Palantir/CommandCode use
|
||||
hostname matching (not substring) so ``evil.com/palantirfoundry`` paths don't trigger it."""
|
||||
normalized = _normalized_lower(base_url)
|
||||
return (
|
||||
_is_nous_portal_endpoint(base_url)
|
||||
or normalized.startswith(_MINIMAX_ANTHROPIC_PREFIXES)
|
||||
or "azure.com" in normalized
|
||||
or base_url_host_matches(normalized, "palantirfoundry.com")
|
||||
or base_url_host_matches(normalized, "api.commandcode.ai")
|
||||
)
|
||||
|
||||
|
||||
def _base_url_needs_context_1m_beta(base_url: str | None) -> bool:
|
||||
"""Endpoints that still gate 1M context behind a beta (Azure)."""
|
||||
return "azure.com" in _normalize_base_url_text(base_url).lower()
|
||||
|
||||
|
||||
def _is_minimax_anthropic_endpoint(base_url: str | None) -> bool:
|
||||
"""MiniMax's Anthropic-compatible endpoints, which reject the fine-grained-tool-streaming and
|
||||
context-1m betas (stripped even though MiniMax also uses Bearer auth)."""
|
||||
return _normalized_lower(base_url).startswith(_MINIMAX_ANTHROPIC_PREFIXES)
|
||||
|
||||
|
||||
def _is_azure_anthropic_endpoint(base_url: str | None) -> bool:
|
||||
"""Azure-hosted Anthropic Messages endpoints serving ``/anthropic``: modern Foundry
|
||||
(``*.services.ai.azure.*``) and legacy Azure OpenAI (``*.openai.azure.*``) hosts; opts them
|
||||
into ``api-version`` query plumbing. Deliberately no finite TLD allow-list, so
|
||||
sovereign/private clouds work."""
|
||||
normalized = _normalize_base_url_text(base_url)
|
||||
if not normalized:
|
||||
return False
|
||||
parsed = urlparse(normalized)
|
||||
host_padded = f".{(parsed.hostname or '').lower().rstrip('.')}."
|
||||
is_azure_host = ".services.ai.azure." in host_padded or ".openai.azure." in host_padded
|
||||
return is_azure_host and "/anthropic" in (parsed.path or "").lower()
|
||||
@@ -0,0 +1,708 @@
|
||||
"""OpenAI-style -> Anthropic Messages API request conversion: model-id normalization, tool
|
||||
schemas, and the message list (content blocks, thinking blocks and their signatures,
|
||||
tool_use/tool_result pairing, cache_control placement, screenshot eviction, blank-block
|
||||
scrubbing). Endpoint predicates come from ``agent/anthropic_endpoints.py`` so this module never
|
||||
imports the adapter (no cycle)."""
|
||||
|
||||
import copy
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from agent.anthropic_endpoints import (
|
||||
_is_deepseek_anthropic_endpoint, _is_kimi_family_endpoint, _is_nous_portal_endpoint,
|
||||
_is_third_party_anthropic_endpoint, _model_name_is_deepseek_thinking,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_THINKING_TYPES = frozenset(("thinking", "redacted_thinking"))
|
||||
_CACHEABLE_TYPES = frozenset(("text", "tool_use"))
|
||||
_EMPTY_TEXT_PLACEHOLDER = "(empty)"
|
||||
_EMPTY_SCHEMA = {"type": "object", "properties": {}}
|
||||
_BEDROCK_REGION_PREFIXES = ("global.", "us.", "eu.", "apac.", "ap.", "au.", "jp.", "ca.", "sa.", "me.", "af.")
|
||||
|
||||
|
||||
def _block_type(b: Any) -> Any:
|
||||
"""``type`` of a dict block, None for non-dicts."""
|
||||
return b.get("type") if isinstance(b, dict) else None
|
||||
|
||||
|
||||
def _has_block_type(blocks: List[Any], types) -> bool:
|
||||
return any(_block_type(b) in types for b in blocks)
|
||||
|
||||
|
||||
def _is_blank_text_block(b: Any) -> bool:
|
||||
"""A text block whose ``text`` is not a non-whitespace string (None/int/blank all count) —
|
||||
Anthropic 400s on them ("text content blocks must contain non-whitespace text")."""
|
||||
return _block_type(b) == "text" and not (isinstance(b.get("text"), str) and b["text"].strip())
|
||||
|
||||
|
||||
def _cache_control_of(b: Any) -> Optional[Dict[str, Any]]:
|
||||
cc = b.get("cache_control") if isinstance(b, dict) else None
|
||||
return cc if isinstance(cc, dict) else None
|
||||
|
||||
|
||||
def _text_block(text: str) -> Dict[str, str]:
|
||||
return {"type": "text", "text": text}
|
||||
|
||||
|
||||
def _text_block_with_citations(text: Any, cits: Any) -> Dict[str, Any]:
|
||||
"""Text block carrying ``citations`` only when it is a non-empty list (the only input-valid shape)."""
|
||||
block: Dict[str, Any] = _text_block(text)
|
||||
if isinstance(cits, list) and cits:
|
||||
block["citations"] = cits
|
||||
return block
|
||||
|
||||
|
||||
def _parse_tool_args(raw: Any) -> Any:
|
||||
"""JSON-decode a tool_call ``arguments`` string; non-strings pass through, bad JSON -> {}."""
|
||||
try:
|
||||
return json.loads(raw) if isinstance(raw, str) else raw
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return {}
|
||||
|
||||
|
||||
def _strip_thinking(blocks: List[Any]) -> List[Any]:
|
||||
return [b for b in blocks if _block_type(b) not in _THINKING_TYPES]
|
||||
|
||||
|
||||
def _block_ids(blocks: List[Any], btype: str, key: str) -> set:
|
||||
return {b.get(key) for b in blocks if _block_type(b) == btype}
|
||||
|
||||
|
||||
def _assistant_block_lists(result: List[Dict[str, Any]]):
|
||||
"""``(index, message)`` for every assistant message whose content is a block list."""
|
||||
return ((i, m) for i, m in enumerate(result) if m.get("role") == "assistant" and isinstance(m.get("content"), list))
|
||||
|
||||
|
||||
def _carry_cache_control(out: Dict[str, Any], b: Any, *, copy: bool = False) -> Dict[str, Any]:
|
||||
"""Carry a dict-valued ``cache_control`` marker from ``b`` onto ``out`` (returned); ``copy``
|
||||
shallow-copies it so the caller's dict is never shared with the wire payload."""
|
||||
cc = _cache_control_of(b)
|
||||
if cc is not None:
|
||||
out["cache_control"] = dict(cc) if copy else cc
|
||||
return out
|
||||
|
||||
|
||||
def _split_blank_text_blocks(blocks: List[Any]) -> Tuple[List[Any], Any, List[int]]:
|
||||
"""``(kept, relocated_cache_control, dropped_indexes)``: drop blank text blocks, remembering
|
||||
the cache_control of the last one dropped so the caller can relocate the breakpoint."""
|
||||
dropped = [i for i, blk in enumerate(blocks) if _is_blank_text_block(blk)]
|
||||
kept = [blk for i, blk in enumerate(blocks) if i not in dropped]
|
||||
relocated = [cc for i in dropped if (cc := _cache_control_of(blocks[i])) is not None]
|
||||
return kept, relocated[-1] if relocated else None, dropped
|
||||
|
||||
|
||||
def _is_bedrock_model_id(model: str) -> bool:
|
||||
"""Bedrock ids (``anthropic.claude-opus-4-7``, ``us.anthropic.claude-*``) use dots as namespace
|
||||
separators that must be preserved verbatim."""
|
||||
return model.lower().startswith(_BEDROCK_REGION_PREFIXES + ("anthropic.",))
|
||||
|
||||
|
||||
def normalize_model_name(model: str, preserve_dots: bool = False) -> str:
|
||||
"""Strip the ``anthropic/`` prefix (case-insensitive) and, unless ``preserve_dots`` (DashScope:
|
||||
``qwen3.5-plus``), convert version dots to hyphens for Claude models only (``claude-opus-4.6``
|
||||
-> ``claude-opus-4-6``). Bedrock ids keep their namespace dots; non-Anthropic names
|
||||
(``gpt-5.4``) keep dots as part of their canonical form."""
|
||||
if model.lower().startswith("anthropic/"):
|
||||
model = model[len("anthropic/"):]
|
||||
if not preserve_dots and not _is_bedrock_model_id(model) and model.lower().startswith(("claude-", "anthropic/")):
|
||||
# Only convert dots to hyphens for Anthropic/Claude models. See issue #17171.
|
||||
model = model.replace(".", "-")
|
||||
return model
|
||||
|
||||
|
||||
def _sanitize_tool_id(tool_id: str) -> str:
|
||||
"""Anthropic requires ids matching [a-zA-Z0-9_-]; replace the rest, never empty."""
|
||||
return (re.sub(r"[^a-zA-Z0-9_-]", "_", tool_id) if tool_id else "") or "tool_0"
|
||||
|
||||
|
||||
def _tool_use_block(tool_id: Any, name: Any, tool_input: Any) -> Dict[str, Any]:
|
||||
return {"type": "tool_use", "id": _sanitize_tool_id(tool_id), "name": name, "input": tool_input}
|
||||
|
||||
|
||||
def _normalize_tool_input_schema(schema: Any) -> Dict[str, Any]:
|
||||
"""Normalize a tool schema for Anthropic's validator: collapse nullable unions (``anyOf:
|
||||
[{type: string}, {type: null}]`` from Pydantic/MCP optional fields) to the non-null branch —
|
||||
optionality is already expressed by ``required``; ``keep_nullable_hint=False`` because the
|
||||
OpenAPI ``nullable`` keyword is not recognized. Top-level oneOf/allOf/anyOf are rejected with a
|
||||
generic 400, so they are dropped in favour of a plain object."""
|
||||
from tools.schema_sanitizer import strip_nullable_unions
|
||||
|
||||
normalized = strip_nullable_unions(schema, keep_nullable_hint=False) if schema else None
|
||||
if not isinstance(normalized, dict):
|
||||
return dict(_EMPTY_SCHEMA)
|
||||
banned = {"oneOf", "allOf", "anyOf"}
|
||||
if banned & normalized.keys():
|
||||
normalized = {k: v for k, v in normalized.items() if k not in banned}
|
||||
normalized.setdefault("type", "object")
|
||||
if normalized.get("type") == "object" and not isinstance(normalized.get("properties"), dict):
|
||||
normalized = {**normalized, "properties": {}}
|
||||
return normalized
|
||||
|
||||
|
||||
def convert_tools_to_anthropic(tools: List[Dict]) -> List[Dict]:
|
||||
"""Convert OpenAI tool definitions to Anthropic format. Duplicate names are dropped with a
|
||||
warning (Anthropic hard-400s on them); ``cache_control`` on the OpenAI tool dict is forwarded."""
|
||||
result = []
|
||||
seen_names: set = set()
|
||||
for t in tools or []:
|
||||
fn = t.get("function", {})
|
||||
name = fn.get("name", "")
|
||||
# Defensive dedup: Anthropic rejects requests with duplicate tool names. Upstream injection paths
|
||||
# already dedup, but this guard converts a hard API failure into a warning. See: #18478
|
||||
if name and name in seen_names:
|
||||
logger.warning("convert_tools_to_anthropic: duplicate tool name '%s' — dropping second occurrence", name)
|
||||
continue
|
||||
if name:
|
||||
seen_names.add(name)
|
||||
anthropic_tool: Dict[str, Any] = {
|
||||
"name": name, "description": fn.get("description", ""),
|
||||
"input_schema": _normalize_tool_input_schema(fn.get("parameters") or {}),
|
||||
}
|
||||
result.append(_carry_cache_control(anthropic_tool, t, copy=True))
|
||||
return result
|
||||
|
||||
|
||||
def _image_source_from_openai_url(url: str) -> Dict[str, str]:
|
||||
"""OpenAI image URL / data URL -> Anthropic image ``source``."""
|
||||
url = str(url or "").strip()
|
||||
if url.startswith("data:"):
|
||||
header, _, data = url.partition(",")
|
||||
mime_part = header[len("data:"):].split(";", 1)[0].strip()
|
||||
media_type = mime_part if mime_part.startswith("image/") else "image/jpeg"
|
||||
return {"type": "base64", "media_type": media_type, "data": data}
|
||||
return {"type": "url", "url": url}
|
||||
|
||||
|
||||
def _convert_content_part_to_anthropic(part: Any) -> Optional[Dict[str, Any]]:
|
||||
"""Convert one OpenAI-style content part to an Anthropic block (None -> dropped)."""
|
||||
if part is None:
|
||||
return None
|
||||
if not isinstance(part, dict):
|
||||
return _text_block(part if isinstance(part, str) else str(part))
|
||||
ptype = part.get("type")
|
||||
if ptype in ("input_text", "text"):
|
||||
# Rebuild from whitelisted fields only: stored SDK text blocks carry output-only siblings
|
||||
# (parsed_output, citations=None) that the INPUT schema rejects with 400.
|
||||
block = _text_block_with_citations(part.get("text", ""), part.get("citations") if ptype == "text" else None)
|
||||
elif ptype in {"image_url", "input_image"}:
|
||||
image_value = part.get("image_url", {})
|
||||
url = image_value.get("url", "") if isinstance(image_value, dict) else str(image_value or "")
|
||||
block = {"type": "image", "source": _image_source_from_openai_url(url)}
|
||||
else:
|
||||
block = dict(part)
|
||||
if (cache_control := _cache_control_of(part)) is not None:
|
||||
block.setdefault("cache_control", dict(cache_control))
|
||||
return block
|
||||
|
||||
|
||||
def _to_plain_data(value: Any, *, _depth: int = 0, _path: Optional[set] = None) -> Any:
|
||||
"""Recursively convert SDK objects to plain Python data. ``_path`` tracks ids on the *current*
|
||||
recursion path (shared but non-cyclic objects convert normally while true cycles stringify);
|
||||
depth is capped at 20."""
|
||||
_path = set() if _path is None else _path
|
||||
obj_id = id(value)
|
||||
if _depth > 20 or obj_id in _path:
|
||||
return str(value)
|
||||
|
||||
def rec(v):
|
||||
return _to_plain_data(v, _depth=_depth + 1, _path=_path)
|
||||
|
||||
_path.add(obj_id)
|
||||
if hasattr(value, "model_dump"):
|
||||
try:
|
||||
# warnings=False: streaming-accumulator blocks trip pydantic's serializer-mismatch
|
||||
# UserWarning, which otherwise leaks to the terminal.
|
||||
dumped = value.model_dump(warnings=False)
|
||||
except TypeError: # duck-typed model_dump without pydantic's signature
|
||||
dumped = value.model_dump()
|
||||
result = rec(dumped)
|
||||
elif isinstance(value, dict):
|
||||
result = {k: rec(v) for k, v in value.items()}
|
||||
elif isinstance(value, (list, tuple)):
|
||||
result = [rec(v) for v in value]
|
||||
elif hasattr(value, "__dict__"):
|
||||
result = {k: rec(v) for k, v in vars(value).items() if not k.startswith("_")}
|
||||
else:
|
||||
result = value
|
||||
_path.discard(obj_id)
|
||||
return result
|
||||
|
||||
|
||||
def _extract_preserved_thinking_blocks(message: Dict[str, Any]) -> List[Dict[str, Any]]:
|
||||
"""Deep-copied thinking/redacted_thinking blocks from ``reasoning_details``."""
|
||||
raw_details = message.get("reasoning_details")
|
||||
if not isinstance(raw_details, list):
|
||||
return []
|
||||
return [
|
||||
copy.deepcopy(d)
|
||||
for d in raw_details
|
||||
if isinstance(d, dict) and str(d.get("type", "") or "").strip().lower() in _THINKING_TYPES
|
||||
]
|
||||
|
||||
|
||||
def _convert_content_to_anthropic(content: Any) -> Any:
|
||||
"""Convert an OpenAI multimodal content list to Anthropic blocks (non-lists pass through)."""
|
||||
if not isinstance(content, list):
|
||||
return content
|
||||
return [b for b in map(_convert_content_part_to_anthropic, content) if b is not None]
|
||||
|
||||
|
||||
def _content_parts_to_anthropic_blocks(parts: Any) -> List[Dict[str, Any]]:
|
||||
"""Tool-message content parts -> tool_result inner blocks (text + image only, the types
|
||||
Anthropic accepts there). Used for multimodal tool results."""
|
||||
out: List[Dict[str, Any]] = []
|
||||
for block in map(_convert_content_part_to_anthropic, parts if isinstance(parts, list) else []):
|
||||
if not block:
|
||||
continue
|
||||
btype, text_val, src = block.get("type"), block.get("text"), block.get("source")
|
||||
if btype == "text" and isinstance(text_val, str) and text_val:
|
||||
out.append(_text_block(text_val))
|
||||
elif btype == "image" and isinstance(src, dict) and src:
|
||||
out.append({"type": "image", "source": src})
|
||||
return out
|
||||
|
||||
|
||||
def _safe_text(text: Any) -> str:
|
||||
"""``text`` if non-whitespace, else the placeholder. A blank text block stored in history (e.g.
|
||||
by compression) is replayed on every turn and wedges the session with HTTP 400; the placeholder
|
||||
is self-healing. Mirrors ``bedrock_adapter._safe_text`` (kept separate on purpose)."""
|
||||
text = "" if text is None else str(text)
|
||||
return text if text.strip() else _EMPTY_TEXT_PLACEHOLDER
|
||||
|
||||
|
||||
def _replay_text(b: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
# Drop blank blocks rather than coerce in place: the caller relocates any cache_control and
|
||||
# falls back to a placeholder only when nothing survives, so "(empty)" never sits as
|
||||
# model-visible noise next to real blocks.
|
||||
if _is_blank_text_block(b):
|
||||
return None
|
||||
return _carry_cache_control(_text_block_with_citations(b["text"], b.get("citations")), b)
|
||||
|
||||
|
||||
def _replay_thinking(b: Dict[str, Any]) -> Dict[str, Any]:
|
||||
out = {"type": "thinking", "thinking": b.get("thinking", "")}
|
||||
return {**out, "signature": b["signature"]} if b.get("signature") else out
|
||||
|
||||
|
||||
def _replay_redacted_thinking(b: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
return {"type": "redacted_thinking", "data": b["data"]} if b.get("data") else None
|
||||
|
||||
|
||||
def _replay_tool_use(b: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return _carry_cache_control(_tool_use_block(b.get("id", ""), b.get("name", ""), b.get("input", {})), b)
|
||||
|
||||
|
||||
def _replay_image(b: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
src = b.get("source")
|
||||
return {"type": "image", "source": src} if isinstance(src, dict) else None
|
||||
|
||||
|
||||
_REPLAY_SANITIZERS = {
|
||||
"text": _replay_text, "thinking": _replay_thinking, "redacted_thinking": _replay_redacted_thinking,
|
||||
"tool_use": _replay_tool_use, "image": _replay_image,
|
||||
}
|
||||
|
||||
|
||||
def _sanitize_replay_block(b: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""Whitelist a stored Anthropic block so it is valid as REQUEST input. SDK response blocks carry
|
||||
output-only fields the INPUT schema forbids ("Extra inputs are not permitted": ``parsed_output``,
|
||||
``caller``, ``citations=None``), and ``_to_plain_data`` captured them verbatim. Whitelist per
|
||||
type (not blacklist) so future SDK fields can't reintroduce the bug; unknown types are dropped.
|
||||
Returns a clean block or None."""
|
||||
if not isinstance(b, dict):
|
||||
return None
|
||||
sanitizer = _REPLAY_SANITIZERS.get(b.get("type"))
|
||||
return sanitizer(b) if sanitizer else None
|
||||
|
||||
|
||||
def _apply_assistant_cache_control_to_last_cacheable_block(blocks: List[Dict[str, Any]], cache_control: Any) -> None:
|
||||
if not isinstance(cache_control, dict):
|
||||
return
|
||||
for block in reversed(blocks):
|
||||
if _block_type(block) in _CACHEABLE_TYPES:
|
||||
block.setdefault("cache_control", dict(cache_control))
|
||||
break
|
||||
|
||||
|
||||
def _replay_ordered_blocks(m: Dict[str, Any], ordered_blocks: List[Any]) -> Optional[List[Dict[str, Any]]]:
|
||||
"""Interleaved-thinking replay: rebuild the assistant turn from the verbatim block list
|
||||
normalize_response stored (only for turns interleaving SIGNED thinking with tool_use).
|
||||
Preserves block ORDER; returns None if nothing survives. tool_use ``input`` is re-sourced from
|
||||
``tool_calls`` (redacted at storage time) rather than the captured block (raw API response, NOT
|
||||
redacted), so a secret the model inlined into a tool call never goes back on the wire."""
|
||||
redacted_input_by_id = {
|
||||
_sanitize_tool_id(tc.get("id", "")): _parse_tool_args((tc.get("function", {}) or {}).get("arguments", "{}"))
|
||||
for tc in m.get("tool_calls", []) or []
|
||||
if isinstance(tc, dict)
|
||||
}
|
||||
replayed: List[Dict[str, Any]] = []
|
||||
relocated_cc = None
|
||||
dropped_blank_text = False
|
||||
for b in ordered_blocks:
|
||||
clean = _sanitize_replay_block(b)
|
||||
if clean is None:
|
||||
dropped_blank_text = dropped_blank_text or _block_type(b) == "text"
|
||||
if (cc := _cache_control_of(b)) is not None: # relocate a dropped block's breakpoint
|
||||
relocated_cc = cc
|
||||
continue
|
||||
if clean.get("type") == "tool_use" and (redacted := redacted_input_by_id.get(clean.get("id", ""))) is not None:
|
||||
clean["input"] = redacted
|
||||
replayed.append(clean)
|
||||
# Nothing cacheable survived (e.g. signed thinking + blank text): emit the placeholder so the
|
||||
# turn stays schema-valid and a relocated marker has a carrier.
|
||||
if not _has_block_type(replayed, _CACHEABLE_TYPES) and (dropped_blank_text or relocated_cc is not None):
|
||||
replayed.append(_text_block(_EMPTY_TEXT_PLACEHOLDER))
|
||||
if not replayed:
|
||||
return None
|
||||
_apply_assistant_cache_control_to_last_cacheable_block(replayed, relocated_cc)
|
||||
_apply_assistant_cache_control_to_last_cacheable_block(replayed, m.get("cache_control"))
|
||||
# prompt_caching marks an assistant turn with text by writing cache_control INTO ``content``
|
||||
# (not top-level). This path never reads ``content``, so carry that marker over or the
|
||||
# breakpoint is burned rather than relocated.
|
||||
msg_content = m.get("content")
|
||||
if isinstance(msg_content, list):
|
||||
inline_cc = next((cc for cc in map(_cache_control_of, msg_content) if cc is not None), None)
|
||||
_apply_assistant_cache_control_to_last_cacheable_block(replayed, inline_cc)
|
||||
return replayed
|
||||
|
||||
|
||||
def _convert_assistant_message(m: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Assistant message -> Anthropic content blocks (thinking, text, tool_use, Kimi/DeepSeek
|
||||
reasoning_content injection)."""
|
||||
# apply_anthropic_cache_control marks an assistant turn with non-empty text by writing cache_control
|
||||
# INTO ``content`` (see _apply_cache_marker's list branch), not at the top level. This branch rebuilds
|
||||
# the message from ordered_blocks and never reads ``content``, so that marker would be dropped -- and
|
||||
# because _can_carry_marker already counted this message as a carrier, the breakpoint is burned rather
|
||||
# than relocated. #56195 covered the complementary shape (blank content -> top-level marker); this is
|
||||
# the interleaved thinking + preamble-text + tool_use shape.
|
||||
content = m.get("content", "")
|
||||
ordered_blocks = m.get("anthropic_content_blocks")
|
||||
if isinstance(ordered_blocks, list) and ordered_blocks:
|
||||
replayed = _replay_ordered_blocks(m, ordered_blocks)
|
||||
if replayed:
|
||||
return {"role": "assistant", "content": replayed}
|
||||
blocks = _extract_preserved_thinking_blocks(m)
|
||||
# Blank text blocks are dropped; a cache marker riding on one is relocated onto the last
|
||||
# surviving cacheable block (prompt_caching sets cache_control on content[-1], which may be
|
||||
# exactly the blank block).
|
||||
relocated_cc = None
|
||||
if isinstance(content, list):
|
||||
kept, relocated_cc, _ = _split_blank_text_blocks(_convert_content_to_anthropic(content))
|
||||
blocks.extend(kept)
|
||||
elif content and str(content).strip():
|
||||
blocks.append(_text_block(str(content)))
|
||||
for tc in m.get("tool_calls", []):
|
||||
if not tc or not isinstance(tc, dict):
|
||||
continue
|
||||
fn = tc.get("function", {})
|
||||
blocks.append(_tool_use_block(tc.get("id", ""), fn.get("name", ""), _parse_tool_args(fn.get("arguments", "{}"))))
|
||||
# Kimi's /coding endpoint requires reasoning_content on replayed tool-call turns — even ""
|
||||
# (injected as a fallback upstream). Prepend, since thinking must precede text/tool_use. Skip
|
||||
# when reasoning_details already supplied (signed) thinking blocks: a duplicate unsigned one
|
||||
# would be downgraded to a spurious text block on the last assistant message.
|
||||
# See hermes-agent#13848. Accept empty string "" — _copy_reasoning_content_for_api() injects "" as a
|
||||
# tier-3 fallback for Kimi tool-call messages that had no reasoning.
|
||||
reasoning_content = m.get("reasoning_content")
|
||||
if isinstance(reasoning_content, str) and not _has_block_type(blocks, _THINKING_TYPES):
|
||||
blocks.insert(0, {"type": "thinking", "thinking": reasoning_content})
|
||||
# Empty assistant content is rejected. Fall back ONLY to the placeholder, never to raw
|
||||
# ``content`` — that is the unfiltered blank payload just removed. Markers are applied after
|
||||
# the fallback so one from a sole dropped blank block lands on the placeholder.
|
||||
effective = blocks or [_text_block(_EMPTY_TEXT_PLACEHOLDER)]
|
||||
_apply_assistant_cache_control_to_last_cacheable_block(effective, relocated_cc)
|
||||
_apply_assistant_cache_control_to_last_cacheable_block(effective, m.get("cache_control"))
|
||||
return {"role": "assistant", "content": effective}
|
||||
|
||||
|
||||
def _tool_result_content(m: Dict[str, Any]) -> Any:
|
||||
"""Resolve a tool message's content into tool_result content (blocks or string)."""
|
||||
content = m.get("content", "")
|
||||
multimodal_blocks: Optional[List[Dict[str, Any]]] = None
|
||||
if isinstance(content, dict) and content.get("_multimodal"):
|
||||
multimodal_blocks = _content_parts_to_anthropic_blocks(content.get("content") or [])
|
||||
if not multimodal_blocks and content.get("text_summary"):
|
||||
multimodal_blocks = [_text_block(str(content["text_summary"]))]
|
||||
elif isinstance(content, list):
|
||||
converted = _content_parts_to_anthropic_blocks(content)
|
||||
if _has_block_type(converted, {"image"}):
|
||||
multimodal_blocks = converted
|
||||
if multimodal_blocks is None: # back-compat: blocks stashed under a private key
|
||||
stashed = m.get("_anthropic_content_blocks")
|
||||
if isinstance(stashed, list) and stashed:
|
||||
text_content = content if isinstance(content, str) and content.strip() else None
|
||||
multimodal_blocks = [_text_block(text_content)] + stashed if text_content else list(stashed)
|
||||
if multimodal_blocks:
|
||||
return multimodal_blocks
|
||||
return (content if isinstance(content, str) else json.dumps(content) if content else "") or "(no output)"
|
||||
|
||||
|
||||
def _convert_tool_message_to_result(result: List[Dict[str, Any]], m: Dict[str, Any]) -> None:
|
||||
"""Append a tool_result to ``result``, merging into a trailing tool_result user message when
|
||||
there is one. Mutates ``result`` in place."""
|
||||
tool_result = {
|
||||
"type": "tool_result", "tool_use_id": _sanitize_tool_id(m.get("tool_call_id", "")),
|
||||
"content": _tool_result_content(m),
|
||||
}
|
||||
_carry_cache_control(tool_result, m, copy=True)
|
||||
last = result[-1] if result else {}
|
||||
last_content = last.get("content") if last.get("role") == "user" else None
|
||||
if isinstance(last_content, list) and last_content and last_content[0].get("type") == "tool_result":
|
||||
last_content.append(tool_result)
|
||||
else:
|
||||
result.append({"role": "user", "content": [tool_result]})
|
||||
|
||||
|
||||
def _convert_user_message(content: Any) -> Dict[str, Any]:
|
||||
"""Validate and convert a user message to Anthropic format."""
|
||||
if isinstance(content, list):
|
||||
content = _fix_blank_text_blocks_in_list(
|
||||
_convert_content_to_anthropic(content), placeholder_text="(empty message)",
|
||||
msg_index=-1, role="user", location="_convert_user_message",
|
||||
)
|
||||
elif not content or (isinstance(content, str) and not content.strip()):
|
||||
content = "(empty message)"
|
||||
return {"role": "user", "content": content}
|
||||
|
||||
|
||||
def _strip_orphaned_tool_blocks(result: List[Dict[str, Any]]) -> None:
|
||||
"""Strip tool_use blocks with no matching tool_result, and vice versa. Compression/truncation
|
||||
can remove either side of a pair or insert messages between them. Anthropic requires the
|
||||
tool_result in the IMMEDIATELY FOLLOWING user message — a global id match is not enough.
|
||||
Mutates ``result`` in place."""
|
||||
# Pass 1: tool_use without an adjacent result.
|
||||
for i, m in _assistant_block_lists(result):
|
||||
tool_use_ids_in_turn = _block_ids(m["content"], "tool_use", "id")
|
||||
if not tool_use_ids_in_turn:
|
||||
continue
|
||||
adjacent_result_ids: set = set()
|
||||
if i + 1 < len(result):
|
||||
nxt = result[i + 1]
|
||||
if nxt.get("role") == "user" and isinstance(nxt.get("content"), list):
|
||||
adjacent_result_ids = _block_ids(nxt["content"], "tool_result", "tool_use_id")
|
||||
orphaned = tool_use_ids_in_turn - adjacent_result_ids
|
||||
if not orphaned:
|
||||
continue
|
||||
kept = [b for b in m["content"] if not (_block_type(b) == "tool_use" and b.get("id") in orphaned)]
|
||||
# A signed thinking block on this turn was signed against the ORIGINAL content and is now
|
||||
# dead (400 "thinking blocks in the latest assistant message cannot be modified"). Flag so
|
||||
# _manage_thinking_signatures demotes it.
|
||||
if len(kept) != len(m["content"]) and _has_block_type(m["content"], _THINKING_TYPES):
|
||||
m["_thinking_signature_invalidated"] = True
|
||||
m["content"] = kept if kept else [_text_block("(tool call removed)")]
|
||||
# Pass 2: tool_result whose tool_use no longer exists anywhere.
|
||||
surviving_tool_use_ids: set = set()
|
||||
for _, m in _assistant_block_lists(result):
|
||||
surviving_tool_use_ids |= _block_ids(m["content"], "tool_use", "id")
|
||||
for m in result:
|
||||
if m.get("role") != "user" or not isinstance(m.get("content"), list):
|
||||
continue
|
||||
new_content = [
|
||||
b for b in m["content"] if _block_type(b) != "tool_result" or b.get("tool_use_id") in surviving_tool_use_ids
|
||||
]
|
||||
if len(new_content) != len(m["content"]):
|
||||
m["content"] = new_content if new_content else [_text_block("(tool result removed)")]
|
||||
|
||||
|
||||
def _concat_content(prev: Any, curr: Any) -> Any:
|
||||
"""Merge two message contents: str+str joined by newline, list+list concatenated, mixed shapes
|
||||
promoted to block lists."""
|
||||
if isinstance(prev, str) and isinstance(curr, str):
|
||||
return prev + "\n" + curr
|
||||
as_blocks = lambda c: [_text_block(c)] if isinstance(c, str) else c # noqa: E731
|
||||
return as_blocks(prev) + as_blocks(curr)
|
||||
|
||||
|
||||
def _merge_consecutive_roles(result: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
"""Merge consecutive same-role messages to enforce alternation. Returns a new list."""
|
||||
fixed: List[Dict[str, Any]] = []
|
||||
for m in result:
|
||||
if not (fixed and fixed[-1]["role"] == m["role"]):
|
||||
fixed.append(m)
|
||||
continue
|
||||
if m["role"] != "user":
|
||||
# Keep the orphan-strip flag visible to _manage_thinking_signatures.
|
||||
if m.get("_thinking_signature_invalidated"):
|
||||
fixed[-1]["_thinking_signature_invalidated"] = True
|
||||
# The second message's thinking blocks were signed against a different turn boundary
|
||||
# and become invalid once merged.
|
||||
if isinstance(m["content"], list):
|
||||
m["content"] = _strip_thinking(m["content"])
|
||||
fixed[-1]["content"] = _concat_content(fixed[-1]["content"], m["content"])
|
||||
return fixed
|
||||
|
||||
|
||||
def _keep_valid_latest_thinking(content: List[Any], signature_dead: bool) -> List[Any]:
|
||||
"""Latest assistant turn on direct Anthropic: keep signed thinking, demote unsigned to text so
|
||||
the reasoning isn't lost. If orphan-stripping mutated THIS turn every signature is dead (and a
|
||||
bare signed block with no tool_use is also invalid), so demote ALL of them."""
|
||||
new_content = []
|
||||
for b in content:
|
||||
if _block_type(b) not in _THINKING_TYPES:
|
||||
new_content.append(b)
|
||||
continue
|
||||
is_redacted = b.get("type") == "redacted_thinking"
|
||||
signed = b.get("data") if is_redacted else b.get("signature") # redacted 'data' IS the signature
|
||||
if signed and not signature_dead:
|
||||
new_content.append(b)
|
||||
elif (signature_dead or not is_redacted) and b.get("thinking"):
|
||||
new_content.append(_text_block(b["thinking"])) # demote to plain text; dataless redacted dropped
|
||||
return new_content
|
||||
|
||||
|
||||
def _manage_thinking_signatures(result: List[Dict[str, Any]], base_url: str | None, model: str | None) -> None:
|
||||
"""Strip or preserve thinking blocks per endpoint. Mutates ``result`` in place.
|
||||
|
||||
Anthropic signs thinking blocks against the full turn; any upstream mutation invalidates them
|
||||
(400 "Invalid signature in thinking block"), so on direct Anthropic only the LATEST assistant
|
||||
turn keeps signed blocks. Signatures are proprietary: third-party endpoints strip all thinking.
|
||||
Kimi replays as-is; DeepSeek needs unsigned blocks round-tripped but rejects signed ones. Nous
|
||||
Portal proxies Claude with sticky sessions and validates the same signatures, so it takes the
|
||||
native path despite not being anthropic.com.
|
||||
"""
|
||||
is_third_party = _is_third_party_anthropic_endpoint(base_url) and not _is_nous_portal_endpoint(base_url)
|
||||
is_kimi = _is_kimi_family_endpoint(base_url, model)
|
||||
is_deepseek = _is_deepseek_anthropic_endpoint(base_url) or (
|
||||
is_third_party and _model_name_is_deepseek_thinking(model)
|
||||
)
|
||||
last_assistant_idx = next((i for i in range(len(result) - 1, -1, -1) if result[i].get("role") == "assistant"), None)
|
||||
for idx, m in _assistant_block_lists(result):
|
||||
if is_kimi:
|
||||
pass # shared cleanup below still strips cache markers + the flag
|
||||
elif is_deepseek:
|
||||
# Strip signed (or redacted-with-data), keep unsigned.
|
||||
new_content = [
|
||||
b for b in m["content"]
|
||||
if _block_type(b) not in _THINKING_TYPES or not (b.get("signature") or b.get("data"))
|
||||
]
|
||||
m["content"] = new_content or [_text_block("(empty)")]
|
||||
elif is_third_party or idx != last_assistant_idx:
|
||||
m["content"] = _strip_thinking(m["content"]) or [_text_block("(thinking elided)")]
|
||||
else:
|
||||
new_content = _keep_valid_latest_thinking(m["content"], bool(m.get("_thinking_signature_invalidated")))
|
||||
m["content"] = new_content or [_text_block("(empty)")]
|
||||
# cache_control on thinking blocks interferes with signature validation.
|
||||
for b in m["content"]:
|
||||
if _block_type(b) in _THINKING_TYPES:
|
||||
b.pop("cache_control", None)
|
||||
m.pop("_thinking_signature_invalidated", None) # internal flag, never on the wire
|
||||
|
||||
|
||||
def _evict_old_screenshots(result: List[Dict[str, Any]]) -> None:
|
||||
"""Keep only the 3 most recent computer-use screenshots (~1,465 tokens each); older images
|
||||
become a placeholder text block. Mutates ``result`` in place."""
|
||||
image_count = 0
|
||||
for msg in reversed(result):
|
||||
content = msg.get("content")
|
||||
for block in content if isinstance(content, list) else []:
|
||||
inner = block.get("content") if _block_type(block) == "tool_result" else None
|
||||
if not isinstance(inner, list) or not _has_block_type(inner, {"image"}):
|
||||
continue
|
||||
image_count += 1
|
||||
if image_count > 3:
|
||||
placeholder = _text_block("[screenshot removed to save context]")
|
||||
block["content"] = [placeholder if b.get("type") == "image" else b for b in inner]
|
||||
|
||||
|
||||
def _ensure_leading_user_turn(result: List[Dict[str, Any]]) -> None:
|
||||
"""Anthropic requires messages[0].role == user; prepend a placeholder turn otherwise. A second
|
||||
auto-compaction can leave a role=assistant summary first, which the API rejects (often masked
|
||||
as a misleading tool_use/tool_result 400). The filler must be non-whitespace text or it trades
|
||||
that 400 for the blank-block one.
|
||||
|
||||
The inserted text block must be non-whitespace: Anthropic separately rejects any text content block
|
||||
whose text is empty or whitespace-only ("text content blocks must contain non-whitespace text"), so a
|
||||
single space here traded the "leading assistant turn" 400 for that one (#69512 class). Uses the same
|
||||
placeholder as every other synthesized filler block in this module for consistency.
|
||||
"""
|
||||
if result and result[0].get("role") != "user":
|
||||
result.insert(0, {"role": "user", "content": [_text_block(_EMPTY_TEXT_PLACEHOLDER)]})
|
||||
|
||||
|
||||
def _fix_blank_text_blocks_in_list(
|
||||
blocks: List[Any], *, placeholder_text: str, msg_index: int, role: Any, location: str
|
||||
) -> List[Any]:
|
||||
"""Drop blank text blocks; relocate any cache_control they carried onto the last surviving
|
||||
cacheable block; if nothing survives, substitute one placeholder block (carrying the relocated
|
||||
marker). Non-text blocks and order are untouched. Returns a new list; logs structure only."""
|
||||
kept, relocated_cache_control, dropped = _split_blank_text_blocks(blocks)
|
||||
for block_index in dropped:
|
||||
logger.warning(
|
||||
"Pre-call sanitizer: dropped blank text content block "
|
||||
"(message_index=%d role=%s location=%s block_index=%d block_type=text)",
|
||||
msg_index, role, location, block_index,
|
||||
)
|
||||
kept = kept or [_text_block(placeholder_text)]
|
||||
_apply_assistant_cache_control_to_last_cacheable_block(kept, relocated_cache_control)
|
||||
return kept
|
||||
|
||||
|
||||
def _scrub_blank_text_blocks(result: List[Dict[str, Any]]) -> None:
|
||||
"""Final boundary guard against blank text blocks (HTTP 400 "text content blocks must contain
|
||||
non-whitespace text"), including inside tool_result content. Runs LAST so a blank block from
|
||||
any current or future producer never reaches the wire. Mutates ``result`` in place."""
|
||||
for msg_index, msg in enumerate(result):
|
||||
if not isinstance(msg, dict):
|
||||
continue
|
||||
role = msg.get("role")
|
||||
content = msg.get("content")
|
||||
if not isinstance(content, list) or not content:
|
||||
continue
|
||||
msg["content"] = _fix_blank_text_blocks_in_list(
|
||||
content, placeholder_text=_EMPTY_TEXT_PLACEHOLDER if role == "assistant" else "(empty message)",
|
||||
msg_index=msg_index, role=role, location="content",
|
||||
)
|
||||
for blk in msg["content"]:
|
||||
inner = blk.get("content") if _block_type(blk) == "tool_result" else None
|
||||
if isinstance(inner, list) and inner:
|
||||
blk["content"] = _fix_blank_text_blocks_in_list(
|
||||
inner, placeholder_text="(no output)", msg_index=msg_index, role=role, location="tool_result",
|
||||
)
|
||||
|
||||
|
||||
def _convert_system_content(content: Any) -> Any:
|
||||
"""System message content -> Anthropic ``system`` param (str, or block list when cache_control
|
||||
is present). With cache markers the blocks are copied (never mutating the caller's dicts) and
|
||||
blank text is replaced by the placeholder: Anthropic rejects blank system blocks too, and a
|
||||
blank block carrying a breakpoint can't simply be dropped."""
|
||||
if not isinstance(content, list):
|
||||
return content
|
||||
if not any(p.get("cache_control") for p in content if isinstance(p, dict)):
|
||||
return "\n".join(p["text"] for p in content if p.get("type") == "text")
|
||||
return [
|
||||
{**p, "text": _EMPTY_TEXT_PLACEHOLDER} if _is_blank_text_block(p) and isinstance(p.get("text"), str) else p
|
||||
for p in content
|
||||
if isinstance(p, dict)
|
||||
]
|
||||
|
||||
|
||||
def convert_messages_to_anthropic(
|
||||
messages: List[Dict], base_url: str | None = None, model: str | None = None
|
||||
) -> Tuple[Optional[Any], List[Dict]]:
|
||||
"""Convert OpenAI-format messages to Anthropic format -> ``(system, messages)``. System is
|
||||
extracted into its own param (a string, or a block list when cache_control is present).
|
||||
``base_url``/``model`` drive thinking-signature policy — third-party endpoints strip signatures
|
||||
(proprietary, they 400 on them); Kimi-family endpoints/models keep unsigned
|
||||
reasoning_content-derived blocks, which Kimi requires even when empty."""
|
||||
system = None
|
||||
result: List[Dict[str, Any]] = []
|
||||
for m in messages:
|
||||
role = m.get("role", "user")
|
||||
if role == "system":
|
||||
system = _convert_system_content(m.get("content", ""))
|
||||
elif role == "assistant":
|
||||
result.append(_convert_assistant_message(m))
|
||||
elif role == "tool":
|
||||
_convert_tool_message_to_result(result, m)
|
||||
else:
|
||||
result.append(_convert_user_message(m.get("content", "")))
|
||||
_strip_orphaned_tool_blocks(result)
|
||||
result = _merge_consecutive_roles(result)
|
||||
_ensure_leading_user_turn(result)
|
||||
_manage_thinking_signatures(result, base_url, model)
|
||||
_evict_old_screenshots(result)
|
||||
_scrub_blank_text_blocks(result)
|
||||
return system, result
|
||||
@@ -0,0 +1,220 @@
|
||||
"""Provider API error summarising for ``AIAgent``.
|
||||
|
||||
Entitlement-failure detection, xAI subscription decoration, structured-detail coercion, and log-safe
|
||||
redaction of provider error payloads.
|
||||
Extracted from ``run_agent.py``; every method resolves through ``AIAgent``'s MRO unchanged.
|
||||
"""
|
||||
import json
|
||||
import re
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from agent.redact import redact_sensitive_text
|
||||
|
||||
# Offline DNS failures are wrapped in a generic "Connection error" by SDKs — inspect the chain.
|
||||
_NETWORK_RESOLUTION_MARKERS = (
|
||||
"temporary failure in name resolution",
|
||||
"name or service not known",
|
||||
"nodename nor servname provided, or not known",
|
||||
"getaddrinfo failed",
|
||||
"no address associated with hostname",
|
||||
"network is unreachable",
|
||||
)
|
||||
_XAI_ENTITLEMENT_HINT = (
|
||||
" — xAI rejected this OAuth account. NOTE: X Premium+ does NOT "
|
||||
"include xAI API access — only standalone SuperGrok subscribers "
|
||||
"can use this provider. Other possible causes: no Grok "
|
||||
"subscription, your tier doesn't include this model, or your "
|
||||
"quota is exhausted. Check https://grok.com/?_s=usage to see "
|
||||
"which, or run `/model` to switch providers."
|
||||
)
|
||||
_ERROR_DETAIL_KEYS = ("message", "detail", "error", "code", "type")
|
||||
|
||||
|
||||
def _is_xai_entitlement_text(lower: str) -> bool:
|
||||
"""xAI's permission-denied body text for an unsubscribed / under-tiered / exhausted account."""
|
||||
return (
|
||||
"do not have an active grok subscription" in lower
|
||||
or ("out of available resources" in lower and "grok" in lower)
|
||||
or ("does not have permission" in lower and "grok" in lower)
|
||||
)
|
||||
|
||||
|
||||
def _http_prefix(error: Exception) -> str:
|
||||
status_code = getattr(error, "status_code", None)
|
||||
return f"HTTP {status_code}: " if status_code else ""
|
||||
|
||||
|
||||
class ApiErrorSummaryMixin:
|
||||
"""Provider error -> user/log-safe summary (see module docstring)."""
|
||||
|
||||
@staticmethod
|
||||
def _is_entitlement_failure(
|
||||
error_context: Optional[Dict[str, Any]], status_code: Optional[int]
|
||||
) -> bool:
|
||||
"""Detect subscription/entitlement 401/403s that masquerade as auth failures.
|
||||
|
||||
Refreshing a token cannot fix an unsubscribed account, so callers surface the error instead of looping
|
||||
the pool. xAI returns the same permission-denied text for BOTH cases; a ``[WKE=unauthenticated:...]``
|
||||
suffix (or "access token could not be validated") means stale token → return False so the refresh path
|
||||
runs.
|
||||
|
||||
Disambiguator for xAI (#29344): the same ``code`` text ("The caller does not have permission to
|
||||
execute the specified operation") is returned for BOTH an unsubscribed account AND a stale OAuth
|
||||
access token. xAI ships an explicit signal in the ``error`` field that tells the two apart: a
|
||||
``[WKE=unauthenticated:...]`` suffix (and/or the ``OAuth2 access token could not be validated``
|
||||
phrasing) means the credentials failed validation — that's recoverable by refreshing the token, NOT
|
||||
by surfacing an entitlement message. When either signal is present we return False eagerly so the
|
||||
credential-pool refresh path runs, letting long-running TUI sessions recover from stale tokens
|
||||
without an exit/reopen cycle.
|
||||
"""
|
||||
if status_code not in {401, 403, None}:
|
||||
return False
|
||||
if not isinstance(error_context, dict):
|
||||
return False
|
||||
# Single lowercase haystack over every field shape (message/reason and raw code/error).
|
||||
haystack = " ".join(
|
||||
str(error_context.get(k) or "").lower() for k in ("message", "reason", "code", "error")
|
||||
)
|
||||
if not haystack.strip():
|
||||
return False
|
||||
if "[wke=unauthenticated:" in haystack or "oauth2 access token could not be validated" in haystack:
|
||||
return False
|
||||
return _is_xai_entitlement_text(haystack)
|
||||
|
||||
@staticmethod
|
||||
def _decorate_xai_entitlement_error(detail: str) -> str:
|
||||
"""Append a neutral hint when xAI's OAuth surface returns the permission-denied 403.
|
||||
|
||||
xAI's ``/v1/responses`` uses one body for several causes (no subscription, tier lacks the model, quota
|
||||
exhausted). The least obvious: X Premium+ does NOT include API access — only SuperGrok does. Lead with
|
||||
that, keep the raw text, point at https://grok.com/?_s=usage. Idempotent: a substring unique to the
|
||||
hint marks prior decoration.
|
||||
"""
|
||||
if not detail or not _is_xai_entitlement_text(detail.lower()):
|
||||
return detail
|
||||
if "X Premium+ does NOT include" in detail:
|
||||
return detail
|
||||
return f"{detail}{_XAI_ENTITLEMENT_HINT}"
|
||||
|
||||
@staticmethod
|
||||
def _coerce_api_error_detail(value: Any) -> str:
|
||||
"""Return a display-safe string for structured provider error fields."""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, dict):
|
||||
for key in _ERROR_DETAIL_KEYS:
|
||||
nested = value.get(key)
|
||||
if isinstance(nested, str) and nested.strip():
|
||||
return nested
|
||||
for key in _ERROR_DETAIL_KEYS:
|
||||
if key in value:
|
||||
nested_detail = ApiErrorSummaryMixin._coerce_api_error_detail(value[key])
|
||||
if nested_detail:
|
||||
return nested_detail
|
||||
try:
|
||||
return json.dumps(value, ensure_ascii=False, sort_keys=True)
|
||||
except TypeError:
|
||||
return str(value)
|
||||
if isinstance(value, (list, tuple)):
|
||||
parts = [ApiErrorSummaryMixin._coerce_api_error_detail(item) for item in value]
|
||||
return "; ".join(part for part in parts if part)
|
||||
if value is None:
|
||||
return ""
|
||||
return str(value)
|
||||
|
||||
@staticmethod
|
||||
def _summarize_api_error(error: Exception) -> str:
|
||||
"""Extract a human-readable one-liner from an API error.
|
||||
|
||||
Cloudflare HTML pages → ``<title>``; network/DNS failures (even SDK-wrapped) → offline hint; else
|
||||
truncated str(error).
|
||||
"""
|
||||
raw = str(error)
|
||||
|
||||
current: Optional[BaseException] = error
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
if any(marker in str(current).lower() for marker in _NETWORK_RESOLUTION_MARKERS):
|
||||
return (
|
||||
"Hermes can't reach the model provider. You may be offline. "
|
||||
"Check your internet connection and try again."
|
||||
)
|
||||
current = current.__cause__ or current.__context__
|
||||
|
||||
if isinstance(error, ValueError) and "expected ident at line" in raw.lower():
|
||||
return f"Malformed provider streaming response: {raw[:300]}"
|
||||
|
||||
prefix = _http_prefix(error)
|
||||
# Cloudflare / proxy HTML pages: grab the <title> (and Ray ID) for a clean summary
|
||||
if "<!DOCTYPE" in raw or "<html" in raw:
|
||||
m = re.search(r"<title[^>]*>([^<]+)</title>", raw, re.IGNORECASE)
|
||||
title = m.group(1).strip() if m else "HTML error page (title not found)"
|
||||
ray = re.search(r"Cloudflare Ray ID:\s*<strong[^>]*>([^<]+)</strong>", raw)
|
||||
parts = [prefix[:-2]] if prefix else []
|
||||
parts.append(title)
|
||||
if ray:
|
||||
parts.append(f"Ray {ray.group(1).strip()}")
|
||||
return " — ".join(parts)
|
||||
|
||||
# GeminiAPIError already composes a clean one-liner with guidance; don't re-extract the raw body.
|
||||
if type(error).__name__ == "GeminiAPIError":
|
||||
return redact_sensitive_text(raw[:1000])
|
||||
|
||||
# JSON body errors from OpenAI/Anthropic SDKs
|
||||
body = getattr(error, "body", None)
|
||||
if isinstance(body, dict):
|
||||
msg = body.get("error", {}).get("message") if isinstance(body.get("error"), dict) else body.get("message")
|
||||
if msg:
|
||||
msg = ApiErrorSummaryMixin._coerce_api_error_detail(msg)
|
||||
return ApiErrorSummaryMixin._decorate_xai_entitlement_error(f"{prefix}{msg[:300]}")
|
||||
|
||||
# SDK may leave body empty while httpx has the payload. Redact: the body is attacker-influenced
|
||||
# and may echo Authorization / x-api-key / request JSON.
|
||||
# Redact before returning: the raw provider/proxy error body is attacker-influenced and may echo
|
||||
# Authorization / x-api-key / request JSON, which would otherwise leak into final_response + logs
|
||||
# (this path widens exposure vs the old empty-body "HTTP 400" string). See #36109.
|
||||
response = getattr(error, "response", None)
|
||||
if response is not None:
|
||||
try:
|
||||
snippet = (getattr(response, "text", None) or "").strip()
|
||||
except Exception:
|
||||
snippet = ""
|
||||
if snippet:
|
||||
try:
|
||||
payload = json.loads(snippet)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
payload = None
|
||||
if isinstance(payload, dict):
|
||||
err = payload.get("error")
|
||||
if isinstance(err, dict) and err.get("message"):
|
||||
return redact_sensitive_text(f"{prefix}{str(err['message'])[:300]}")
|
||||
if payload.get("message"):
|
||||
return redact_sensitive_text(f"{prefix}{str(payload['message'])[:300]}")
|
||||
return redact_sensitive_text(f"{prefix}{snippet[:300]}")
|
||||
|
||||
# Fallback: truncate the raw string but give more room than 200 chars
|
||||
return ApiErrorSummaryMixin._decorate_xai_entitlement_error(f"{prefix}{raw[:500]}")
|
||||
|
||||
def _mask_api_key_for_logs(self, key: Any) -> Optional[str]:
|
||||
# Azure Foundry Entra ID bearer providers are callables — never invoke them in log
|
||||
# paths; identify the auth surface instead.
|
||||
if callable(key) and not isinstance(key, str):
|
||||
return "<entra-id-bearer>"
|
||||
if not key:
|
||||
return None
|
||||
if len(key) <= 12:
|
||||
return "***"
|
||||
return f"{key[:8]}...{key[-4:]}"
|
||||
|
||||
def _clean_error_message(self, error_msg: str) -> str:
|
||||
"""Clean up error messages for user display, removing HTML content and truncating."""
|
||||
if not error_msg:
|
||||
return "Unknown error"
|
||||
# HTML content is common with CloudFlare and gateway error pages
|
||||
if error_msg.strip().startswith('<!DOCTYPE html') or '<html' in error_msg:
|
||||
return "Service temporarily unavailable (HTML error page returned)"
|
||||
cleaned = ' '.join(error_msg.split())
|
||||
if len(cleaned) > 150:
|
||||
cleaned = cleaned[:150] + "..."
|
||||
return cleaned
|
||||
@@ -0,0 +1,196 @@
|
||||
"""Lifecycle-hook payloads for ``AIAgent`` API requests.
|
||||
|
||||
JSON-safe coercion, secret-key redaction, size caps, and the ``api_request_error`` hook dispatch.
|
||||
Extracted from ``run_agent.py``; every method resolves through ``AIAgent``'s MRO unchanged.
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from agent.usage_pricing import normalize_usage
|
||||
|
||||
_SENSITIVE_HOOK_KEYS = {"api_key", "authorization", "proxy_authorization", "cookie", "set_cookie"}
|
||||
|
||||
|
||||
def _model_dump(value: Any) -> Any:
|
||||
"""``value.model_dump(mode="json")`` with graceful degradation for older pydantic signatures.
|
||||
|
||||
warnings=False: pydantic UserWarnings on generic-union SDK models would leak to the terminal.
|
||||
"""
|
||||
try:
|
||||
return value.model_dump(mode="json", warnings=False)
|
||||
except TypeError:
|
||||
try:
|
||||
return value.model_dump(mode="json")
|
||||
except TypeError:
|
||||
return value.model_dump()
|
||||
|
||||
|
||||
class ApiRequestHooksMixin:
|
||||
"""Hook payload sanitising + ``api_request_error`` dispatch (see module docstring)."""
|
||||
|
||||
def _usage_summary_for_api_request_hook(self, response: Any) -> Optional[Dict[str, Any]]:
|
||||
"""Token buckets for ``post_api_request`` plugins (no raw ``response`` object)."""
|
||||
if response is None:
|
||||
return None
|
||||
raw_usage = getattr(response, "usage", None)
|
||||
if not raw_usage:
|
||||
return None
|
||||
from dataclasses import asdict
|
||||
|
||||
cu = normalize_usage(raw_usage, provider=self.provider, api_mode=self.api_mode)
|
||||
summary = asdict(cu)
|
||||
summary.pop("raw_usage", None)
|
||||
summary["prompt_tokens"] = cu.prompt_tokens
|
||||
summary["total_tokens"] = cu.total_tokens
|
||||
return summary
|
||||
|
||||
@staticmethod
|
||||
def _hook_payload_max_chars() -> int:
|
||||
raw = os.getenv("HERMES_PLUGIN_PAYLOAD_MAX_CHARS", "50000")
|
||||
try:
|
||||
return max(1000, int(raw))
|
||||
except (TypeError, ValueError):
|
||||
return 50000
|
||||
|
||||
@staticmethod
|
||||
def _is_sensitive_hook_key(key: Any) -> bool:
|
||||
if not isinstance(key, str):
|
||||
return False
|
||||
lowered = key.lower().replace("-", "_")
|
||||
return lowered in _SENSITIVE_HOOK_KEYS or lowered.endswith("_api_key")
|
||||
|
||||
@classmethod
|
||||
def _hook_jsonable(
|
||||
cls, value: Any, *, depth: int = 0, max_depth: int = 8, max_string: int = 8000,
|
||||
max_sequence: int = 200,
|
||||
) -> Any:
|
||||
if depth > max_depth:
|
||||
return f"<{type(value).__name__} depth limit>"
|
||||
if value is None or isinstance(value, (bool, int, float)):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
if len(value) > max_string:
|
||||
return value[:max_string] + f"...[truncated {len(value) - max_string} chars]"
|
||||
return value
|
||||
if isinstance(value, (bytes, bytearray)):
|
||||
return f"<{len(value)} bytes>"
|
||||
|
||||
def recurse(item):
|
||||
return cls._hook_jsonable(
|
||||
item, depth=depth + 1, max_depth=max_depth, max_string=max_string,
|
||||
max_sequence=max_sequence,
|
||||
)
|
||||
|
||||
if isinstance(value, dict):
|
||||
out: Dict[str, Any] = {}
|
||||
for idx, (key, item) in enumerate(value.items()):
|
||||
if idx >= max_sequence:
|
||||
out["_truncated_items"] = len(value) - max_sequence
|
||||
break
|
||||
str_key = str(key)
|
||||
out[str_key] = "<redacted>" if cls._is_sensitive_hook_key(str_key) else recurse(item)
|
||||
return out
|
||||
if isinstance(value, (list, tuple, set)):
|
||||
seq = list(value)
|
||||
out = [recurse(item) for item in seq[:max_sequence]]
|
||||
if len(seq) > max_sequence:
|
||||
out.append({"_truncated_items": len(seq) - max_sequence})
|
||||
return out
|
||||
with suppress(Exception):
|
||||
if hasattr(value, "model_dump"):
|
||||
return recurse(_model_dump(value))
|
||||
with suppress(Exception):
|
||||
from dataclasses import asdict, is_dataclass
|
||||
if is_dataclass(value):
|
||||
return recurse(asdict(value))
|
||||
if isinstance(value, SimpleNamespace):
|
||||
return recurse(vars(value))
|
||||
if hasattr(value, "__dict__"):
|
||||
with suppress(Exception):
|
||||
return recurse({k: v for k, v in vars(value).items() if not str(k).startswith("_")})
|
||||
return str(value)[:max_string]
|
||||
|
||||
@classmethod
|
||||
def _sanitize_hook_payload(cls, value: Any) -> Any:
|
||||
"""JSON-able payload under the size cap: full → reduced caps → truncated preview."""
|
||||
limit = cls._hook_payload_max_chars()
|
||||
encoded = ""
|
||||
for caps in ({}, {"max_string": 1000, "max_sequence": 50}):
|
||||
payload = cls._hook_jsonable(value, **caps)
|
||||
try:
|
||||
encoded = json.dumps(payload, ensure_ascii=False, default=str)
|
||||
except Exception:
|
||||
return str(payload)[:limit]
|
||||
if len(encoded) <= limit:
|
||||
return payload
|
||||
return {
|
||||
"_truncated": True, "original_type": type(value).__name__, "preview": encoded[:limit]
|
||||
}
|
||||
|
||||
def _api_request_payload_for_hook(self, api_kwargs: Optional[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
body = {
|
||||
key: value
|
||||
for key, value in (api_kwargs or {}).items()
|
||||
if key not in {"timeout", "http_client"}
|
||||
}
|
||||
return self._sanitize_hook_payload({"method": "POST", "body": body})
|
||||
|
||||
def _api_response_payload_for_hook(
|
||||
self, response: Any, assistant_message: Any, *, finish_reason: Optional[str]
|
||||
) -> Dict[str, Any]:
|
||||
# Raw provider SDK tool_call objects are handed to the sanitizer on purpose; `_hook_jsonable` must
|
||||
# keep normalising them (model_dump / __dict__ / dataclass) or subscribers get str() blobs.
|
||||
tool_calls = getattr(assistant_message, "tool_calls", None) or []
|
||||
return self._sanitize_hook_payload(
|
||||
{
|
||||
"model": getattr(response, "model", None),
|
||||
"finish_reason": finish_reason,
|
||||
"assistant_message": {
|
||||
"role": getattr(assistant_message, "role", "assistant"),
|
||||
"content": getattr(assistant_message, "content", None),
|
||||
"tool_calls": tool_calls,
|
||||
},
|
||||
"usage": self._usage_summary_for_api_request_hook(response),
|
||||
}
|
||||
)
|
||||
|
||||
def _invoke_api_request_error_hook(
|
||||
self, *, task_id: str, turn_id: str, api_request_id: str, api_call_count: int,
|
||||
api_start_time: float, api_kwargs: Optional[Dict[str, Any]], error_type: str,
|
||||
error_message: str, status_code: Optional[int] = None, retry_count: Optional[int] = None,
|
||||
max_retries: Optional[int] = None, retryable: Optional[bool] = None,
|
||||
reason: Optional[str] = None,
|
||||
) -> None:
|
||||
# Lazy module import (not from-import) so tests can replace lifecycle dispatch at this call site.
|
||||
with suppress(Exception):
|
||||
from hermes_cli import lifecycle as _lifecycle
|
||||
if not _lifecycle.has_hook("api_request_error"):
|
||||
return
|
||||
ended_at = time.time()
|
||||
_lifecycle.invoke_hook(
|
||||
"api_request_error",
|
||||
task_id=task_id,
|
||||
turn_id=turn_id,
|
||||
api_request_id=api_request_id,
|
||||
session_id=self.session_id or "",
|
||||
platform=self.platform or "",
|
||||
model=self.model,
|
||||
provider=self.provider,
|
||||
base_url=self.base_url,
|
||||
api_mode=self.api_mode,
|
||||
api_call_count=api_call_count,
|
||||
api_duration=ended_at - api_start_time,
|
||||
started_at=api_start_time,
|
||||
ended_at=ended_at,
|
||||
status_code=status_code,
|
||||
retry_count=retry_count,
|
||||
max_retries=max_retries,
|
||||
retryable=retryable,
|
||||
reason=reason,
|
||||
error={"type": error_type, "message": error_message},
|
||||
request=self._api_request_payload_for_hook(api_kwargs),
|
||||
)
|
||||
+15
-49
@@ -1,24 +1,10 @@
|
||||
"""Async/sync bridging helpers.
|
||||
|
||||
The codebase has ~30 sites that schedule a coroutine onto an event loop from a
|
||||
worker thread via :func:`asyncio.run_coroutine_threadsafe`. That function can
|
||||
raise :class:`RuntimeError` (e.g. the loop was closed during a shutdown race),
|
||||
and when it does the coroutine object is never awaited and never closed —
|
||||
which triggers a ``"coroutine '<name>' was never awaited"`` RuntimeWarning and
|
||||
leaks the coroutine's frame until GC.
|
||||
|
||||
:func:`safe_schedule_threadsafe` wraps the call, closes the coroutine on
|
||||
scheduling failure, and returns ``None`` (instead of a half-formed future) so
|
||||
callers can branch cleanly:
|
||||
|
||||
fut = safe_schedule_threadsafe(coro, loop)
|
||||
if fut is None:
|
||||
return # or fallback behavior
|
||||
fut.result(timeout=5)
|
||||
|
||||
The helper deliberately does NOT also handle ``future.result()`` failures —
|
||||
that is a separate concern. Once the loop has accepted the coroutine, its
|
||||
lifecycle belongs to the loop, not the scheduling thread.
|
||||
``asyncio.run_coroutine_threadsafe`` can raise ``RuntimeError`` (loop closed during a
|
||||
shutdown race); the coroutine is then never awaited or closed, which triggers a
|
||||
"coroutine was never awaited" RuntimeWarning and leaks its frame. The helpers here
|
||||
close the coroutine on scheduling failure. ``future.result()`` failures are deliberately
|
||||
NOT handled: once the loop accepts the coroutine its lifecycle belongs to the loop.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -32,34 +18,20 @@ _DEFAULT_LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def safe_schedule_threadsafe(
|
||||
coro: Coroutine[Any, Any, Any],
|
||||
loop: Optional[asyncio.AbstractEventLoop],
|
||||
*,
|
||||
coro: Coroutine[Any, Any, Any], loop: Optional[asyncio.AbstractEventLoop], *,
|
||||
logger: Optional[logging.Logger] = None,
|
||||
log_message: str = "Failed to schedule coroutine on loop",
|
||||
log_level: int = logging.DEBUG,
|
||||
log_message: str = "Failed to schedule coroutine on loop", log_level: int = logging.DEBUG,
|
||||
) -> Optional[Future]:
|
||||
"""Schedule ``coro`` on ``loop`` from a sync context, leak-safe.
|
||||
|
||||
Returns the :class:`concurrent.futures.Future` on success, or ``None`` if
|
||||
the loop is missing or :func:`asyncio.run_coroutine_threadsafe` raised
|
||||
(e.g. the loop was closed during a shutdown race). In all failure paths
|
||||
the coroutine is :meth:`close`-d so it does not trigger
|
||||
``"coroutine was never awaited"`` warnings or leak its frame.
|
||||
|
||||
Callers retain full control over what to do with the returned future
|
||||
(call ``.result(timeout=...)``, attach ``add_done_callback``, ignore it
|
||||
fire-and-forget, etc.).
|
||||
Returns the Future on success, or ``None`` if the loop is missing or scheduling
|
||||
raised; in every failure path the coroutine is closed. Callers keep full control
|
||||
over the returned future (``.result(timeout=...)``, callbacks, fire-and-forget).
|
||||
"""
|
||||
log = logger if logger is not None else _DEFAULT_LOGGER
|
||||
|
||||
if loop is None:
|
||||
if asyncio.iscoroutine(coro):
|
||||
coro.close()
|
||||
log.log(log_level, "%s: loop is None", log_message)
|
||||
return None
|
||||
|
||||
try:
|
||||
if loop is None:
|
||||
raise RuntimeError("loop is None")
|
||||
return asyncio.run_coroutine_threadsafe(coro, loop)
|
||||
except Exception as exc:
|
||||
if asyncio.iscoroutine(coro):
|
||||
@@ -69,15 +41,9 @@ def safe_schedule_threadsafe(
|
||||
|
||||
|
||||
def consume_detached_task_result(task: "asyncio.Future[Any]") -> None:
|
||||
"""Retrieve a detached task's result without surfacing cancellation.
|
||||
|
||||
Used as an ``add_done_callback`` on tasks that were cancelled and
|
||||
detached (e.g. an adapter close path that swallows ``CancelledError``
|
||||
past its teardown deadline). Observing ``task.exception()`` prevents
|
||||
"exception was never retrieved" noise on the event loop; cancellation
|
||||
and any terminal error are deliberately swallowed — the task's owner
|
||||
already gave up on it.
|
||||
"""
|
||||
"""``add_done_callback`` for cancelled-and-detached tasks: observe the exception so the
|
||||
loop does not log "exception was never retrieved"; cancellation and terminal errors
|
||||
are swallowed because the task's owner already gave up on it."""
|
||||
try:
|
||||
task.exception()
|
||||
except (asyncio.CancelledError, Exception):
|
||||
|
||||
+32
-68
@@ -1,27 +1,11 @@
|
||||
"""Ambient session-accounting context for auxiliary LLM calls.
|
||||
|
||||
Auxiliary calls (vision, compression, title generation, web_extract,
|
||||
session_search, ...) funnel through ``agent.auxiliary_client`` which has no
|
||||
session handle — so their token usage was historically discarded, leaving
|
||||
dashboard analytics blind to aux model spend (issue #23270).
|
||||
|
||||
Instead of threading ``session_db``/``session_id`` parameters through every
|
||||
aux call site, the agent loop publishes them here (mirroring the Nous Portal
|
||||
conversation context in ``agent.portal_tags``) and the auxiliary client
|
||||
records usage at its single response-validation chokepoint.
|
||||
|
||||
ContextVar semantics give us the right isolation for free:
|
||||
|
||||
* concurrent agents in one process (gateway sessions, delegate subagents)
|
||||
never see each other's accounting context;
|
||||
* worker threads spawned via ``tools.thread_context.propagate_context_to_thread``
|
||||
(MoA fan-out, background review) inherit the parent turn's context;
|
||||
* asyncio tasks inherit the context of the code that created them.
|
||||
|
||||
MoA reference/aggregator slots are explicitly EXCLUDED from recording:
|
||||
``agent/conversation_loop.py`` already folds MoA advisor usage and cost into
|
||||
the main loop's ``update_token_counts`` delta, so recording them here would
|
||||
double-count (see ``_EXCLUDED_TASKS``).
|
||||
Aux calls (vision, compression, title generation, web_extract, session_search, ...) go
|
||||
through ``agent.auxiliary_client`` which has no session handle, so their usage was
|
||||
historically discarded. The agent loop publishes ``(session_db, session_id)`` here
|
||||
(mirroring ``agent.portal_tags``) and the aux client records usage at its single
|
||||
response-validation chokepoint. ContextVar semantics isolate concurrent agents, propagate
|
||||
to worker threads via ``tools.thread_context`` and to asyncio tasks automatically.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -33,22 +17,17 @@ from typing import Any, Optional
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# (session_db, session_id) for the active agent turn, or None outside one.
|
||||
_accounting: ContextVar[Optional[tuple]] = ContextVar(
|
||||
"aux_accounting_context", default=None
|
||||
)
|
||||
_accounting: ContextVar[Optional[tuple]] = ContextVar("aux_accounting_context", default=None)
|
||||
|
||||
# Aux tasks whose usage is already accounted by the main loop — recording
|
||||
# them here would double-count. MoA advisor/aggregator usage is folded into
|
||||
# conversation_loop's update_token_counts delta (tokens AND cost).
|
||||
# MoA advisor/aggregator usage is already folded into conversation_loop's
|
||||
# update_token_counts delta (tokens AND cost); recording it here would double-count.
|
||||
_EXCLUDED_TASKS = frozenset({"moa_reference", "moa_aggregator"})
|
||||
|
||||
|
||||
def set_accounting_context(session_db: Any, session_id: Optional[str]):
|
||||
"""Publish the active session's accounting handles for aux usage recording.
|
||||
"""Publish the active session's accounting handles; returns the token for ``reset_accounting_context``.
|
||||
|
||||
Called by the agent loop at turn entry. Returns the ContextVar token so
|
||||
callers can ``reset_accounting_context(token)`` on turn exit. Publishing
|
||||
``None`` handles (no DB / no session id) clears the context.
|
||||
``None`` handles (no DB / no session id) clear the context.
|
||||
"""
|
||||
if session_db is None or not session_id:
|
||||
return _accounting.set(None)
|
||||
@@ -63,31 +42,16 @@ def reset_accounting_context(token) -> None:
|
||||
_accounting.set(None)
|
||||
|
||||
|
||||
def get_accounting_context() -> Optional[tuple]:
|
||||
"""Return ``(session_db, session_id)`` for the active turn, or ``None``."""
|
||||
return _accounting.get()
|
||||
|
||||
|
||||
def record_aux_usage(
|
||||
response: Any,
|
||||
task: Optional[str],
|
||||
*,
|
||||
provider: Optional[str] = None,
|
||||
response: Any, task: Optional[str], *, provider: Optional[str] = None,
|
||||
base_url: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Record an auxiliary response's token usage against the ambient session.
|
||||
|
||||
Called from the auxiliary client's response-validation chokepoint. Strictly
|
||||
best-effort: any failure is swallowed (accounting must never break an aux
|
||||
call). No-ops when:
|
||||
|
||||
* no accounting context is published (call is outside any agent turn),
|
||||
* the task is main-loop-accounted (MoA slots — see ``_EXCLUDED_TASKS``),
|
||||
* the response carries no usage object.
|
||||
|
||||
The model is read from ``response.model`` (accurate even after the aux
|
||||
client's provider-fallback chains); *provider*/*base_url* reflect the
|
||||
originally-resolved route and are best-effort.
|
||||
Strictly best-effort (accounting must never break an aux call). No-ops outside an
|
||||
agent turn, for main-loop-accounted tasks (``_EXCLUDED_TASKS``), or without usage.
|
||||
The model is read from ``response.model`` (accurate after aux provider fallback);
|
||||
*provider*/*base_url* reflect the originally-resolved route.
|
||||
"""
|
||||
try:
|
||||
if not task or task in _EXCLUDED_TASKS:
|
||||
@@ -109,30 +73,30 @@ def record_aux_usage(
|
||||
or usage.reasoning_tokens
|
||||
):
|
||||
return
|
||||
|
||||
model = str(getattr(response, "model", "") or "") or "unknown"
|
||||
estimated_cost = None
|
||||
try:
|
||||
cost = estimate_usage_cost(
|
||||
model, usage, provider=provider, base_url=base_url
|
||||
)
|
||||
cost = estimate_usage_cost(model, usage, provider=provider, base_url=base_url)
|
||||
if cost.amount_usd is not None:
|
||||
estimated_cost = float(cost.amount_usd)
|
||||
except Exception:
|
||||
logger.debug("Aux usage cost estimation failed", exc_info=True)
|
||||
|
||||
session_db.record_auxiliary_usage(
|
||||
session_id,
|
||||
task,
|
||||
model=model,
|
||||
billing_provider=provider,
|
||||
billing_base_url=base_url,
|
||||
input_tokens=usage.input_tokens,
|
||||
output_tokens=usage.output_tokens,
|
||||
cache_read_tokens=usage.cache_read_tokens,
|
||||
cache_write_tokens=usage.cache_write_tokens,
|
||||
reasoning_tokens=usage.reasoning_tokens,
|
||||
estimated_cost_usd=estimated_cost,
|
||||
session_id, task, model=model, billing_provider=provider, billing_base_url=base_url,
|
||||
input_tokens=usage.input_tokens, output_tokens=usage.output_tokens,
|
||||
cache_read_tokens=usage.cache_read_tokens, cache_write_tokens=usage.cache_write_tokens,
|
||||
reasoning_tokens=usage.reasoning_tokens, estimated_cost_usd=estimated_cost,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Aux usage recording failed (non-fatal)", exc_info=True)
|
||||
|
||||
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
|
||||
def get_accounting_context() -> Optional[tuple]:
|
||||
"""Return ``(session_db, session_id)`` for the active turn, or ``None``."""
|
||||
return _accounting.get()
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
+4291
-8072
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,36 @@
|
||||
"""Endpoint identity for auxiliary custom-provider health checks."""
|
||||
import contextlib
|
||||
from typing import Any, Optional
|
||||
|
||||
from hermes_cli.route_identity import normalize_route_base_url
|
||||
|
||||
def _unhealthy_cache_key(provider: str, base_url: Optional[str] = None) -> Any:
|
||||
"""Provider-wide key, or endpoint-specific key for an explicit custom endpoint."""
|
||||
from agent.auxiliary_client import _normalize_chain_label
|
||||
label = _normalize_chain_label(provider)
|
||||
endpoint = normalize_route_base_url(_custom_health_base_url(provider, base_url))
|
||||
if endpoint:
|
||||
return "custom-endpoint", endpoint
|
||||
return label
|
||||
|
||||
|
||||
def _custom_health_base_url(provider: str, explicit_base_url: Optional[str] = None) -> str:
|
||||
"""Return the concrete custom endpoint used to scope health and failed-route checks."""
|
||||
from agent.auxiliary_client import _current_custom_base_url
|
||||
explicit = str(explicit_base_url or "").strip()
|
||||
from agent.auxiliary_client import _normalize_chain_label
|
||||
label = _normalize_chain_label(provider)
|
||||
if label == "local/custom":
|
||||
return explicit or _current_custom_base_url()
|
||||
if label.startswith("custom:") and explicit:
|
||||
return explicit
|
||||
with contextlib.suppress(ImportError):
|
||||
from hermes_cli.runtime_provider import _get_named_custom_provider, _resolves_to_custom
|
||||
if _resolves_to_custom(label):
|
||||
return explicit or _current_custom_base_url()
|
||||
entry = _get_named_custom_provider(provider)
|
||||
if entry:
|
||||
return explicit or str(entry.get("base_url") or "").strip()
|
||||
return ""
|
||||
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
"""Message hygiene at the resolved auxiliary client boundary."""
|
||||
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
||||
from agent.transports.chat_completions import ChatCompletionsTransport
|
||||
|
||||
|
||||
def prepare_chat_messages(client, kwargs: dict) -> dict:
|
||||
"""Sanitize actual Chat Completions SDK requests, not native adapter replay.
|
||||
|
||||
Auxiliary and MoA callers can retain a prepared request before the virtual
|
||||
transport sanitizes its copy. The resolved SDK client identifies the wire;
|
||||
native Messages/Responses adapters must retain their reasoning sidecars.
|
||||
"""
|
||||
if not isinstance(client, (OpenAI, AsyncOpenAI)) or "messages" not in kwargs:
|
||||
return kwargs
|
||||
messages = ChatCompletionsTransport().convert_messages(
|
||||
kwargs["messages"], model=kwargs.get("model")
|
||||
)
|
||||
return {**kwargs, "messages": messages}
|
||||
+155
-443
@@ -1,32 +1,12 @@
|
||||
"""Microsoft Entra ID adapter for Microsoft Foundry.
|
||||
|
||||
Provides keyless authentication for Microsoft Foundry deployments using the
|
||||
`azure-identity` SDK's `DefaultAzureCredential` chain (env service principal
|
||||
→ workload identity → managed identity → VS Code → Azure CLI → azd →
|
||||
PowerShell → broker).
|
||||
|
||||
Architecture mirrors `agent/bedrock_adapter.py`:
|
||||
|
||||
* Lazy import. `azure-identity` is only loaded when ``model.auth_mode =
|
||||
entra_id`` is selected. Users who stick with `AZURE_FOUNDRY_API_KEY`
|
||||
never pay the import cost.
|
||||
* SDK-callable contract. The public entry point ``build_token_provider``
|
||||
returns a zero-arg callable produced by ``get_bearer_token_provider`` —
|
||||
this is exactly the value Microsoft's documented sample plugs into
|
||||
``OpenAI(api_key=token_provider, base_url=...)``. The OpenAI SDK calls
|
||||
it before every request, so token refresh is transparent.
|
||||
* Three explicit consumer-side helpers (display / cache / http-bearer)
|
||||
rather than one generic "materialize" function — splitting them by
|
||||
purpose prevents accidental token-minting in logging paths or token
|
||||
leakage into cache keys / dashboard JSON.
|
||||
* No persisted JWT. ``azure-identity`` caches in-process and (where
|
||||
available) in the OS keychain or ``~/.IdentityService``. Hermes does
|
||||
not duplicate that storage in ``auth.json``.
|
||||
|
||||
Keyless auth via the `azure-identity` ``DefaultAzureCredential`` chain (env service principal →
|
||||
workload identity → managed identity → VS Code → Azure CLI → azd → PowerShell → broker). Mirrors
|
||||
``agent/bedrock_adapter.py``: `azure-identity` is imported lazily (only for ``auth_mode = entra_id``);
|
||||
``build_token_provider`` returns the zero-arg callable the OpenAI SDK calls before every request
|
||||
(transparent refresh); consumer helpers are split by purpose so logging paths never mint tokens and
|
||||
tokens never leak into cache keys; no JWT is persisted (azure-identity caches in-process / OS keychain).
|
||||
Reference: https://learn.microsoft.com/azure/ai-foundry/foundry-models/how-to/configure-entra-id
|
||||
|
||||
Requires: ``azure-identity`` (optional dependency — only needed when
|
||||
``model.auth_mode = entra_id``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -40,28 +20,20 @@ from typing import Any, Callable, Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Microsoft-documented scope for Foundry inference auth. Both the new
|
||||
# Foundry portal and the legacy Azure OpenAI managed-identity docs use
|
||||
# this scope for ALL Foundry endpoint shapes (*.openai.azure.com,
|
||||
# *.services.ai.azure.com, *.ai.azure.com). The older control-plane
|
||||
# scope ``https://cognitiveservices.azure.com/.default`` is for ARM
|
||||
# resource management and is rejected for inference by newer
|
||||
# resources — users with that requirement override via
|
||||
# ``model.entra.scope`` in config.yaml.
|
||||
# Microsoft-documented Foundry inference scope for ALL endpoint shapes. The older cognitiveservices.azure.com
|
||||
# scope is an ARM control-plane scope rejected for inference by newer resources; override via ``model.entra.scope``.
|
||||
SCOPE_AI_AZURE_DEFAULT = "https://ai.azure.com/.default"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lazy SDK import — only loaded when the Entra path is actually used.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_AZURE_IDENTITY_FEATURE = "provider.azure_identity"
|
||||
_INSTALL_MSG = "The 'azure-identity' package is required for Azure AI Foundry Entra ID authentication. "
|
||||
_LAZY_INSTALL_HINT = (
|
||||
"pip install azure-identity manually, or enable lazy installs (security.allow_lazy_installs: true in config.yaml)."
|
||||
)
|
||||
_AUTH_HEADERS = ("Authorization", "authorization", "Api-Key", "api-key", "X-Api-Key", "x-api-key")
|
||||
|
||||
|
||||
def has_azure_identity_installed() -> bool:
|
||||
"""Return True if `azure-identity` can be imported right now.
|
||||
|
||||
Cheap check — does not walk the credential chain.
|
||||
"""
|
||||
"""Cheap importability check — does not walk the credential chain."""
|
||||
try:
|
||||
import azure.identity # noqa: F401
|
||||
return True
|
||||
@@ -70,11 +42,7 @@ def has_azure_identity_installed() -> bool:
|
||||
|
||||
|
||||
def _require_azure_identity():
|
||||
"""Import ``azure.identity``, lazy-installing it if allowed.
|
||||
|
||||
Raises ``ImportError`` with a clear actionable message when the
|
||||
package is missing and lazy installs are disabled.
|
||||
"""
|
||||
"""Import ``azure.identity``, lazy-installing if allowed; ImportError with an actionable message otherwise."""
|
||||
try:
|
||||
import azure.identity as _ai
|
||||
return _ai
|
||||
@@ -82,389 +50,190 @@ def _require_azure_identity():
|
||||
try:
|
||||
from tools.lazy_deps import ensure, FeatureUnavailable
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"The 'azure-identity' package is required for Azure AI "
|
||||
"Foundry Entra ID authentication. Install it with: "
|
||||
"pip install azure-identity"
|
||||
) from exc
|
||||
|
||||
raise ImportError(_INSTALL_MSG + "Install it with: pip install azure-identity") from exc
|
||||
try:
|
||||
ensure(_AZURE_IDENTITY_FEATURE, prompt=False)
|
||||
except FeatureUnavailable as exc:
|
||||
raise ImportError(
|
||||
"The 'azure-identity' package is required for Azure AI "
|
||||
"Foundry Entra ID authentication. " + str(exc)
|
||||
) from exc
|
||||
|
||||
# Retry import after lazy install.
|
||||
import azure.identity as _ai # noqa: WPS440
|
||||
raise ImportError(_INSTALL_MSG + str(exc)) from exc
|
||||
import azure.identity as _ai # noqa: WPS440 — retry after lazy install
|
||||
return _ai
|
||||
|
||||
|
||||
def reset_credential_cache() -> None:
|
||||
"""Clear the cached ``DefaultAzureCredential``. Used by tests and
|
||||
profile switches.
|
||||
|
||||
Defensive against tests that ``monkeypatch.setattr`` over
|
||||
``build_credential`` with a plain (non-lru-cached) function — those
|
||||
won't expose ``cache_clear()`` until pytest reverts the patch.
|
||||
"""
|
||||
"""Clear the cached ``DefaultAzureCredential`` (tests, profile switches); tolerates a monkeypatched plain function."""
|
||||
cache_clear = getattr(build_credential, "cache_clear", None)
|
||||
if callable(cache_clear):
|
||||
cache_clear()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Token-provider construction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EntraIdentityConfig:
|
||||
"""Serializable Entra ID config.
|
||||
|
||||
Captures the Hermes-managed Entra knobs we need outside Azure SDK
|
||||
environment configuration. Everything else
|
||||
(tenant ID, service principal secret, federated token file, sovereign
|
||||
cloud authority, etc.) flows through azure-identity's standard
|
||||
``AZURE_*`` env vars — see the Bedrock pattern in
|
||||
``hermes_cli/runtime_provider.py:1310-1377`` for the analogous
|
||||
"let the SDK read env" approach.
|
||||
|
||||
``scope`` is Microsoft's documented Foundry inference audience. Almost
|
||||
everyone uses the default; sovereign-cloud / non-standard tenants can
|
||||
override via ``model.entra.scope``. Identity selection (user-assigned
|
||||
managed identity, workload identity, service principal, tenant, authority)
|
||||
stays in the standard Azure SDK env vars such as ``AZURE_CLIENT_ID``.
|
||||
|
||||
``exclude_interactive_browser`` is kept as an internal constructor knob
|
||||
so probes stay non-interactive by default. It is not written by the setup
|
||||
wizard.
|
||||
|
||||
The dataclass is frozen so it's hashable for ``functools.lru_cache``
|
||||
keying, and serializable across multiprocessing boundaries (workers
|
||||
rebuild the credential inside their own process).
|
||||
"""
|
||||
"""Hermes-managed Entra knobs; everything else (tenant, SP secret, federated token file, authority...) flows
|
||||
through azure-identity's standard ``AZURE_*`` env vars. ``exclude_interactive_browser`` keeps probes
|
||||
non-interactive (the setup wizard never writes it). Frozen: hashable for ``lru_cache``, picklable for workers."""
|
||||
|
||||
scope: str = SCOPE_AI_AZURE_DEFAULT
|
||||
exclude_interactive_browser: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
scope = str(self.scope or "").strip() or SCOPE_AI_AZURE_DEFAULT
|
||||
object.__setattr__(self, "scope", scope)
|
||||
object.__setattr__(self, "scope", str(self.scope or "").strip() or SCOPE_AI_AZURE_DEFAULT)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"scope": self.scope,
|
||||
"exclude_interactive_browser": self.exclude_interactive_browser,
|
||||
}
|
||||
return {"scope": self.scope, "exclude_interactive_browser": self.exclude_interactive_browser}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Optional[Dict[str, Any]],
|
||||
*, default_scope: Optional[str] = None) -> "EntraIdentityConfig":
|
||||
def from_dict(cls, data: Optional[Dict[str, Any]], *, default_scope: Optional[str] = None) -> "EntraIdentityConfig":
|
||||
data = data or {}
|
||||
scope = str(data.get("scope") or "").strip() or default_scope or SCOPE_AI_AZURE_DEFAULT
|
||||
exclude_browser = bool(data.get("exclude_interactive_browser", True))
|
||||
return cls(
|
||||
scope=scope,
|
||||
exclude_interactive_browser=exclude_browser,
|
||||
scope=str(data.get("scope") or "").strip() or default_scope or SCOPE_AI_AZURE_DEFAULT,
|
||||
exclude_interactive_browser=bool(data.get("exclude_interactive_browser", True)),
|
||||
)
|
||||
|
||||
|
||||
def _build_default_credential(config: EntraIdentityConfig) -> Any:
|
||||
"""Construct a ``DefaultAzureCredential`` for ``config``.
|
||||
|
||||
Only Hermes-selected knobs are passed as kwargs. Everything else
|
||||
(tenant, service principal secret, federated token file, sovereign
|
||||
cloud authority, etc.) is read by ``azure-identity`` from the
|
||||
standard ``AZURE_*`` environment variables — see Microsoft's
|
||||
documented credential resolution chain. Users configure those in
|
||||
``~/.hermes/.env`` or the deployment environment.
|
||||
"""
|
||||
ai = _require_azure_identity()
|
||||
kwargs: Dict[str, Any] = {}
|
||||
# SDK default is True (browser excluded); only pass when the user
|
||||
# explicitly opts in to interactive browser auth.
|
||||
if not config.exclude_interactive_browser:
|
||||
kwargs["exclude_interactive_browser_credential"] = False
|
||||
return ai.DefaultAzureCredential(**kwargs)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def build_credential(config: EntraIdentityConfig) -> Any:
|
||||
"""Return the cached ``DefaultAzureCredential`` for ``config``.
|
||||
|
||||
Hermes processes use exactly one Entra config at a time (the
|
||||
``model.entra.*`` block in config.yaml drives every aux task,
|
||||
subagent, and credential probe in the session). ``maxsize=1`` is
|
||||
intentional: it reflects the actual usage pattern and keeps the
|
||||
cache trivially small.
|
||||
|
||||
``EntraIdentityConfig`` is a frozen dataclass, so it's hashable and
|
||||
safe as an LRU-cache key. ``functools.lru_cache`` is thread-safe in
|
||||
CPython.
|
||||
|
||||
If two distinct configs are ever passed (tests do this; production
|
||||
rarely), the LRU eviction handles it correctly — each call still
|
||||
returns a credential matching its config; only one is cached at a
|
||||
time. Use :func:`reset_credential_cache` to clear (e.g. in tests).
|
||||
"""
|
||||
return _build_default_credential(config)
|
||||
|
||||
|
||||
def build_token_provider(scope: Optional[str] = None,
|
||||
*,
|
||||
config: Optional[EntraIdentityConfig] = None,
|
||||
base_url: Optional[str] = None,
|
||||
exclude_interactive_browser: bool = True,
|
||||
) -> Callable[[], str]:
|
||||
"""Return a zero-arg callable that mints a fresh Entra bearer JWT.
|
||||
|
||||
The returned callable is exactly what Microsoft's documented Foundry
|
||||
sample expects::
|
||||
|
||||
from openai import OpenAI
|
||||
client = OpenAI(
|
||||
base_url="https://my-resource.openai.azure.com/openai/v1/",
|
||||
api_key=build_token_provider(),
|
||||
)
|
||||
|
||||
Scope resolution order:
|
||||
1. ``config.scope`` when a config object is supplied
|
||||
2. explicit ``scope`` kwarg
|
||||
3. ``SCOPE_AI_AZURE_DEFAULT`` (Microsoft's documented Foundry scope)
|
||||
|
||||
``base_url`` is unused today and kept for back-compat. Tenant /
|
||||
service-principal / sovereign-cloud configuration flows through
|
||||
``azure-identity``'s standard ``AZURE_*`` environment variables —
|
||||
see :func:`_build_default_credential` for the rationale.
|
||||
|
||||
NOT serializable across process boundaries. For multiprocessing
|
||||
workers, serialize the ``EntraIdentityConfig`` and rebuild the
|
||||
provider inside the worker.
|
||||
"""
|
||||
"""Cached ``DefaultAzureCredential``. ``maxsize=1`` is intentional: a process uses one ``model.entra.*``
|
||||
block at a time. Only Hermes knobs are passed as kwargs; the rest comes from ``AZURE_*`` env vars."""
|
||||
ai = _require_azure_identity()
|
||||
if config is None:
|
||||
config = EntraIdentityConfig(
|
||||
scope=scope or SCOPE_AI_AZURE_DEFAULT,
|
||||
exclude_interactive_browser=exclude_interactive_browser,
|
||||
)
|
||||
credential = build_credential(config)
|
||||
return ai.get_bearer_token_provider(credential, config.scope)
|
||||
# SDK default already excludes the browser; only pass the kwarg when opting in.
|
||||
kwargs = {} if config.exclude_interactive_browser else {"exclude_interactive_browser_credential": False}
|
||||
return ai.DefaultAzureCredential(**kwargs)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Credential probing
|
||||
# ---------------------------------------------------------------------------
|
||||
def _resolve_config(config: Optional[EntraIdentityConfig], scope: Optional[str], **overrides: Any) -> EntraIdentityConfig:
|
||||
if config is not None:
|
||||
return config
|
||||
return EntraIdentityConfig(scope=(scope or "").strip() or SCOPE_AI_AZURE_DEFAULT, **overrides)
|
||||
|
||||
|
||||
def has_azure_identity_credentials(scope: Optional[str] = None,
|
||||
*,
|
||||
config: Optional[EntraIdentityConfig] = None,
|
||||
timeout_seconds: float = 10.0,
|
||||
allow_install: bool = True,
|
||||
**overrides: Any) -> bool:
|
||||
"""Best-effort probe: can `DefaultAzureCredential` mint a token now?
|
||||
|
||||
Runs ``credential.get_token(scope)`` under a thread-based timeout so
|
||||
a slow token service can't hang the caller. Returns False on any
|
||||
error — never raises. Use for ``hermes doctor`` /
|
||||
``hermes auth status`` / wizard preflight.
|
||||
|
||||
``allow_install``: when True (default) and ``azure-identity`` is not
|
||||
importable, the adapter triggers the standard lazy-install path
|
||||
(subject to ``security.allow_lazy_installs``) before probing. Set
|
||||
False to make this strictly an "is installed?" check — used on hot
|
||||
paths like CLI startup where we never want pip to run.
|
||||
|
||||
NOT used by ``is_provider_configured()`` — that path is structural
|
||||
only (no token mint), so CLI startup doesn't pay this latency.
|
||||
"""
|
||||
if not has_azure_identity_installed():
|
||||
if not allow_install:
|
||||
return False
|
||||
try:
|
||||
_require_azure_identity()
|
||||
except ImportError as exc:
|
||||
logger.debug("azure-identity lazy install unavailable: %s", exc)
|
||||
return False
|
||||
if config is None:
|
||||
effective_scope = (scope or "").strip() or SCOPE_AI_AZURE_DEFAULT
|
||||
config = EntraIdentityConfig(scope=effective_scope, **overrides)
|
||||
|
||||
result = {"ok": False}
|
||||
|
||||
def _probe() -> None:
|
||||
try:
|
||||
credential = build_credential(config)
|
||||
tok = credential.get_token(config.scope)
|
||||
result["ok"] = bool(getattr(tok, "token", None))
|
||||
except Exception as exc:
|
||||
logger.debug("Entra credential probe failed: %s", exc)
|
||||
result["ok"] = False
|
||||
|
||||
thread = threading.Thread(target=_probe, daemon=True)
|
||||
thread.start()
|
||||
thread.join(timeout=max(0.01, timeout_seconds))
|
||||
if thread.is_alive():
|
||||
logger.debug("Entra token service probe timed out after %ss", timeout_seconds)
|
||||
return False
|
||||
return bool(result.get("ok"))
|
||||
def _install_failure(allow_install: bool) -> Optional[Dict[str, Any]]:
|
||||
"""None when ``azure.identity`` is importable (lazy-installing if allowed), else ``{"error", "hint"}``."""
|
||||
if has_azure_identity_installed():
|
||||
return None
|
||||
if not allow_install:
|
||||
return {"error": "azure-identity not installed", "hint": "pip install azure-identity (or rely on lazy install at first use)"}
|
||||
try:
|
||||
_require_azure_identity()
|
||||
except ImportError as exc:
|
||||
return {"error": str(exc) or "azure-identity not installed", "hint": _LAZY_INSTALL_HINT, "exc": exc}
|
||||
return None
|
||||
|
||||
|
||||
def describe_active_credential(config: Optional[EntraIdentityConfig] = None,
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
timeout_seconds: float = 10.0,
|
||||
allow_install: bool = True,
|
||||
**overrides: Any) -> Dict[str, Any]:
|
||||
"""Return diagnostic info about the active credential chain.
|
||||
def build_token_provider(scope: Optional[str] = None, *, config: Optional[EntraIdentityConfig] = None,
|
||||
exclude_interactive_browser: bool = True) -> Callable[[], str]:
|
||||
"""Zero-arg callable minting a fresh Entra bearer JWT — pass as ``OpenAI(api_key=...)``. Scope precedence:
|
||||
``config.scope`` > ``scope`` kwarg > default. Not picklable: ship the ``EntraIdentityConfig`` and rebuild
|
||||
in the worker."""
|
||||
ai = _require_azure_identity()
|
||||
config = _resolve_config(config, scope, exclude_interactive_browser=exclude_interactive_browser)
|
||||
return ai.get_bearer_token_provider(build_credential(config), config.scope)
|
||||
|
||||
Best-effort: runs ``get_token()`` and inspects what came back.
|
||||
Designed for ``hermes doctor`` and the wizard preflight — never
|
||||
raises, returns ``{"ok": False, "error": ...}`` on failure.
|
||||
|
||||
``allow_install``: when True (default) and ``azure-identity`` is not
|
||||
importable, the adapter triggers the standard lazy-install path
|
||||
(subject to ``security.allow_lazy_installs``) before probing. The
|
||||
install failure is surfaced as the diagnostic error when it fails.
|
||||
Set False for hot CLI paths that should never trigger pip.
|
||||
|
||||
``azure-identity`` doesn't expose the winning inner credential as
|
||||
a public field, so we report a coarse picture (env vars present,
|
||||
token expiry, claims-derived tenant) rather than the credential
|
||||
class name. Users wanting the precise class can run with
|
||||
``AZURE_LOG_LEVEL=DEBUG``.
|
||||
"""
|
||||
info: Dict[str, Any] = {"ok": False}
|
||||
if not has_azure_identity_installed():
|
||||
if not allow_install:
|
||||
info["error"] = "azure-identity not installed"
|
||||
info["hint"] = (
|
||||
"pip install azure-identity (or rely on lazy install at "
|
||||
"first use)"
|
||||
)
|
||||
return info
|
||||
try:
|
||||
_require_azure_identity()
|
||||
except ImportError as exc:
|
||||
info["error"] = str(exc) or "azure-identity not installed"
|
||||
info["hint"] = (
|
||||
"pip install azure-identity manually, or enable lazy "
|
||||
"installs (security.allow_lazy_installs: true in "
|
||||
"config.yaml)."
|
||||
)
|
||||
return info
|
||||
|
||||
if config is None:
|
||||
effective_scope = (scope or "").strip() or SCOPE_AI_AZURE_DEFAULT
|
||||
config = EntraIdentityConfig(scope=effective_scope, **overrides)
|
||||
|
||||
info["scope"] = config.scope
|
||||
# Tenant / authority / service-principal config flow through the
|
||||
# standard ``AZURE_*`` env vars; surface them below.
|
||||
if os.environ.get("AZURE_TENANT_ID", "").strip():
|
||||
info["tenant_id_env"] = os.environ["AZURE_TENANT_ID"].strip()
|
||||
|
||||
# Surface which env-var sources are present without minting yet.
|
||||
# Credential-bearing vars (AZURE_CLIENT_SECRET, AZURE_FEDERATED_TOKEN_FILE)
|
||||
# are read through the profile secret scope so a multiplexed profile's
|
||||
# diagnostics don't report another profile's env-bridged credentials;
|
||||
# unscoped CLI probes keep the legacy env read (Slack pattern).
|
||||
def _scoped_env(name: str) -> str:
|
||||
try:
|
||||
from agent.secret_scope import UnscopedSecretError, get_secret
|
||||
|
||||
try:
|
||||
return (get_secret(name) or "").strip()
|
||||
except UnscopedSecretError:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
return os.environ.get(name, "").strip()
|
||||
|
||||
env_sources = []
|
||||
if _scoped_env("AZURE_FEDERATED_TOKEN_FILE"):
|
||||
env_sources.append("WorkloadIdentityCredential (AZURE_FEDERATED_TOKEN_FILE)")
|
||||
if (os.environ.get("AZURE_CLIENT_ID", "").strip()
|
||||
and _scoped_env("AZURE_CLIENT_SECRET")
|
||||
and os.environ.get("AZURE_TENANT_ID", "").strip()):
|
||||
env_sources.append("EnvironmentCredential (client secret)")
|
||||
if os.environ.get("IDENTITY_ENDPOINT", "").strip() or os.environ.get("MSI_ENDPOINT", "").strip():
|
||||
env_sources.append("ManagedIdentityCredential (IDENTITY_ENDPOINT)")
|
||||
info["env_sources"] = env_sources
|
||||
|
||||
# Now try minting.
|
||||
def _probe_token(config: EntraIdentityConfig, timeout_seconds: float) -> Optional[Dict[str, Any]]:
|
||||
"""``get_token`` on a daemon thread under a hard deadline → ``{"token"}`` / ``{"error"}`` / None on timeout."""
|
||||
result: Dict[str, Any] = {}
|
||||
|
||||
def _probe() -> None:
|
||||
try:
|
||||
credential = build_credential(config)
|
||||
tok = credential.get_token(config.scope)
|
||||
result["token"] = tok
|
||||
result["token"] = build_credential(config).get_token(config.scope)
|
||||
except Exception as exc:
|
||||
result["error"] = str(exc)
|
||||
|
||||
thread = threading.Thread(target=_probe, daemon=True)
|
||||
thread.start()
|
||||
thread.join(timeout=max(0.01, timeout_seconds))
|
||||
if thread.is_alive():
|
||||
info["error"] = f"Token probe timed out after {timeout_seconds:.0f}s"
|
||||
info["hint"] = (
|
||||
"DefaultAzureCredential can be slow when the token service is unreachable "
|
||||
"or when az login state is stale. Try `az login` or set "
|
||||
"AZURE_CLIENT_ID / AZURE_TENANT_ID / AZURE_CLIENT_SECRET."
|
||||
)
|
||||
return info
|
||||
return None if thread.is_alive() else result
|
||||
|
||||
|
||||
def has_azure_identity_credentials(scope: Optional[str] = None, *, config: Optional[EntraIdentityConfig] = None,
|
||||
timeout_seconds: float = 10.0, allow_install: bool = True,
|
||||
**overrides: Any) -> bool:
|
||||
"""Timeout-bounded probe: can the chain mint a token now? Never raises. ``allow_install=False`` makes it a
|
||||
strict "is installed?" check for hot paths (CLI startup) where pip must never run. NOT used by
|
||||
``is_provider_configured()`` (structural, no mint)."""
|
||||
failure = _install_failure(allow_install)
|
||||
if failure is not None:
|
||||
if "exc" in failure:
|
||||
logger.debug("azure-identity lazy install unavailable: %s", failure["exc"])
|
||||
return False
|
||||
result = _probe_token(_resolve_config(config, scope, **overrides), timeout_seconds)
|
||||
if result is None:
|
||||
logger.debug("Entra token service probe timed out after %ss", timeout_seconds)
|
||||
return False
|
||||
if "error" in result:
|
||||
logger.debug("Entra credential probe failed: %s", result["error"])
|
||||
return False
|
||||
return bool(getattr(result.get("token"), "token", None))
|
||||
|
||||
|
||||
def _env(name: str) -> str:
|
||||
return os.environ.get(name, "").strip()
|
||||
|
||||
|
||||
def _scoped_env(name: str) -> str:
|
||||
"""Credential-bearing env read via the profile secret scope so a multiplexed profile never reports
|
||||
another profile's env-bridged credentials; unscoped CLI probes fall back to plain env."""
|
||||
try:
|
||||
from agent.secret_scope import get_secret
|
||||
return (get_secret(name) or "").strip()
|
||||
except Exception: # UnscopedSecretError, import failure, or any scope error
|
||||
return _env(name)
|
||||
|
||||
|
||||
# (label, predicate) for env-var-driven credential sources, in chain order.
|
||||
_ENV_SOURCE_CHECKS = (
|
||||
("WorkloadIdentityCredential (AZURE_FEDERATED_TOKEN_FILE)", lambda: _scoped_env("AZURE_FEDERATED_TOKEN_FILE")),
|
||||
("EnvironmentCredential (client secret)",
|
||||
lambda: _env("AZURE_CLIENT_ID") and _scoped_env("AZURE_CLIENT_SECRET") and _env("AZURE_TENANT_ID")),
|
||||
("ManagedIdentityCredential (IDENTITY_ENDPOINT)", lambda: _env("IDENTITY_ENDPOINT") or _env("MSI_ENDPOINT")),
|
||||
)
|
||||
|
||||
|
||||
def describe_active_credential(config: Optional[EntraIdentityConfig] = None, *, scope: Optional[str] = None,
|
||||
timeout_seconds: float = 10.0, allow_install: bool = True,
|
||||
**overrides: Any) -> Dict[str, Any]:
|
||||
"""Doctor / preflight diagnostics. Never raises; ``{"ok": False, "error": ...}`` on failure. azure-identity
|
||||
hides the winning inner credential, so this reports a coarse picture (env sources, token expiry) rather
|
||||
than a class name; ``AZURE_LOG_LEVEL=DEBUG`` shows the chain."""
|
||||
info: Dict[str, Any] = {"ok": False}
|
||||
failure = _install_failure(allow_install)
|
||||
if failure is not None:
|
||||
info["error"], info["hint"] = failure["error"], failure["hint"]
|
||||
return info
|
||||
config = _resolve_config(config, scope, **overrides)
|
||||
info["scope"] = config.scope
|
||||
if tenant := _env("AZURE_TENANT_ID"):
|
||||
info["tenant_id_env"] = tenant
|
||||
info["env_sources"] = [label for label, present in _ENV_SOURCE_CHECKS if present()]
|
||||
result = _probe_token(config, timeout_seconds)
|
||||
if result is None:
|
||||
info["error"] = f"Token probe timed out after {timeout_seconds:.0f}s"
|
||||
info["hint"] = ("DefaultAzureCredential can be slow when the token service is unreachable "
|
||||
"or when az login state is stale. Try `az login` or set "
|
||||
"AZURE_CLIENT_ID / AZURE_TENANT_ID / AZURE_CLIENT_SECRET.")
|
||||
return info
|
||||
if "error" in result:
|
||||
info["error"] = result["error"]
|
||||
return info
|
||||
|
||||
token = result.get("token")
|
||||
if token is None:
|
||||
info["error"] = "credential chain exhausted"
|
||||
return info
|
||||
|
||||
info["ok"] = True
|
||||
info["expires_on"] = getattr(token, "expires_on", None)
|
||||
return info
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Consumer-side helpers — split by purpose to prevent accidental token
|
||||
# minting in logging / cache-key / dashboard paths.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
# Consumer-side helpers — split by purpose so logging / cache-key / dashboard paths never mint tokens.
|
||||
def is_token_provider(value: Any) -> bool:
|
||||
"""Return True when ``value`` is a callable Entra token provider.
|
||||
|
||||
Used at the seams where a consumer must decide between
|
||||
string-API-key semantics and bearer-callable semantics.
|
||||
"""
|
||||
"""True when ``value`` is a callable Entra token provider (vs. a string API key)."""
|
||||
return callable(value) and not isinstance(value, str)
|
||||
|
||||
|
||||
def materialize_bearer_for_http(value: Any) -> str:
|
||||
"""Return a fresh Bearer JWT for a manual HTTP request.
|
||||
|
||||
Only call this at sites that must construct an ``Authorization``
|
||||
header outside the OpenAI SDK (e.g. ``hermes_cli/azure_detect.py``).
|
||||
Calls the callable exactly once and returns the resulting token.
|
||||
|
||||
**Anthropic SDK integration:** the Anthropic Python SDK does not
|
||||
accept a ``Callable[[], str]`` for ``auth_token``. Instead,
|
||||
:func:`build_bearer_http_client` returns an ``httpx.Client`` whose
|
||||
request event hook calls this function and rewrites the
|
||||
``Authorization`` header per request — and that client is passed to
|
||||
the Anthropic SDK via ``http_client=...``. See
|
||||
:func:`agent.anthropic_adapter.build_anthropic_client` for the
|
||||
consumer.
|
||||
|
||||
Raises ``ValueError`` if ``value`` is not a callable token provider
|
||||
or non-empty string.
|
||||
"""
|
||||
"""Mint a fresh Bearer JWT for a manual HTTP request (calls the provider once). Only for sites building
|
||||
``Authorization`` outside the OpenAI SDK; the Anthropic SDK can't take a callable, so
|
||||
:func:`build_bearer_http_client` calls this from an httpx hook. ``ValueError`` on an unusable value/empty token."""
|
||||
if is_token_provider(value):
|
||||
token = value()
|
||||
if not isinstance(token, str) or not token:
|
||||
@@ -475,97 +244,40 @@ def materialize_bearer_for_http(value: Any) -> str:
|
||||
raise ValueError("no usable api_key / token provider")
|
||||
|
||||
|
||||
def _strip_auth_headers(request: Any) -> None:
|
||||
for header_name in _AUTH_HEADERS:
|
||||
request.headers.pop(header_name, None)
|
||||
|
||||
|
||||
def build_bearer_http_client(token_provider: Callable[[], str], **httpx_kwargs: Any) -> Any:
|
||||
"""Return an ``httpx.Client`` that mints a fresh Entra bearer JWT
|
||||
per outbound request.
|
||||
|
||||
The Anthropic SDK (≤ 0.86.0 at the time of writing) stores
|
||||
``api_key`` / ``auth_token`` as static strings and computes the
|
||||
``Authorization`` header at construction time. To get per-request
|
||||
token refresh (the Microsoft-recommended Foundry pattern for
|
||||
callable bearer providers), we install an httpx ``request`` event
|
||||
hook on a custom client and pass that client to the SDK via
|
||||
``http_client=...``. The hook:
|
||||
|
||||
1. Calls :func:`materialize_bearer_for_http` to mint a fresh JWT
|
||||
(azure-identity caches internally — this is cheap when the
|
||||
cached token is still valid).
|
||||
2. Strips any pre-set ``Authorization`` / ``api-key`` /
|
||||
``x-api-key`` headers the SDK may have added (avoids
|
||||
conflicting auth values).
|
||||
3. Sets ``Authorization: Bearer <fresh-jwt>``.
|
||||
|
||||
``token_provider`` must be a zero-arg callable returning a string —
|
||||
typically the result of :func:`build_token_provider`.
|
||||
|
||||
``httpx_kwargs`` are forwarded verbatim to ``httpx.Client(...)`` so
|
||||
callers can attach a ``timeout``, ``transport``, ``proxy``, etc.
|
||||
|
||||
Raises ``ImportError`` if ``httpx`` is not installed (it is a
|
||||
transitive dependency of both ``openai`` and ``anthropic`` SDKs, so
|
||||
in practice always available when this helper is reached).
|
||||
"""
|
||||
"""``httpx.Client`` minting a fresh Entra bearer JWT per outbound request. The Anthropic SDK computes
|
||||
``Authorization`` once at construction, so per-request refresh needs a ``request`` hook: mint (cheap —
|
||||
azure-identity caches), strip pre-set auth headers, set ``Authorization: Bearer``. ``httpx_kwargs`` are
|
||||
forwarded verbatim (``timeout``, ``transport``...)."""
|
||||
if not is_token_provider(token_provider):
|
||||
raise ValueError(
|
||||
"build_bearer_http_client requires a zero-arg callable "
|
||||
"token provider"
|
||||
)
|
||||
|
||||
try:
|
||||
import httpx
|
||||
except ImportError as exc: # pragma: no cover — httpx ships with openai/anthropic
|
||||
raise ImportError(
|
||||
"httpx is required for Entra ID bearer auth on Microsoft Foundry "
|
||||
"Anthropic-style endpoints. It is normally a transitive "
|
||||
"dependency of the openai/anthropic SDKs."
|
||||
) from exc
|
||||
raise ValueError("build_bearer_http_client requires a zero-arg callable token provider")
|
||||
import httpx
|
||||
|
||||
def _inject_bearer(request: "httpx.Request") -> None:
|
||||
try:
|
||||
token = materialize_bearer_for_http(token_provider)
|
||||
except ValueError as exc:
|
||||
# Token provider failed (chain exhausted, token service unreachable,
|
||||
# az login expired, etc.). Strip any auth headers the SDK
|
||||
# may have set — including our own placeholder sentinel
|
||||
# ``entra-id-bearer-via-http-hook`` from
|
||||
# ``_build_anthropic_client_with_bearer_hook`` — so the
|
||||
# outbound request hits Azure with NO Authorization rather
|
||||
# than with the placeholder. Azure returns a clean 401
|
||||
# "missing auth" that is easier to diagnose than a 401
|
||||
# against the sentinel string, and the sentinel never
|
||||
# appears in upstream access logs.
|
||||
#
|
||||
# Log at WARNING (not DEBUG) so the misconfiguration is
|
||||
# visible at default log levels.
|
||||
logger.warning(
|
||||
"Bearer hook: Entra ID token provider returned empty (%s) "
|
||||
"— stripping Authorization headers. Azure will respond 401. "
|
||||
"Run `hermes doctor` or `az login` to recover.",
|
||||
exc,
|
||||
)
|
||||
for header_name in ("Authorization", "authorization", "Api-Key", "api-key", "X-Api-Key", "x-api-key"):
|
||||
request.headers.pop(header_name, None)
|
||||
# Chain exhausted / az login expired: strip ALL auth headers (incl. the anthropic_adapter placeholder
|
||||
# sentinel) so Azure returns a clean "missing auth" 401 and the sentinel never reaches upstream logs.
|
||||
# WARNING so the misconfiguration is visible at default levels.
|
||||
logger.warning("Bearer hook: Entra ID token provider returned empty (%s) "
|
||||
"— stripping Authorization headers. Azure will respond 401. "
|
||||
"Run `hermes doctor` or `az login` to recover.", exc)
|
||||
_strip_auth_headers(request)
|
||||
return
|
||||
for header_name in ("Authorization", "authorization", "Api-Key", "api-key", "X-Api-Key", "x-api-key"):
|
||||
request.headers.pop(header_name, None)
|
||||
_strip_auth_headers(request)
|
||||
request.headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
return httpx.Client(
|
||||
event_hooks={"request": [_inject_bearer]},
|
||||
**httpx_kwargs,
|
||||
)
|
||||
return httpx.Client(event_hooks={"request": [_inject_bearer]}, **httpx_kwargs)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EntraIdentityConfig",
|
||||
"SCOPE_AI_AZURE_DEFAULT",
|
||||
"build_bearer_http_client",
|
||||
"build_credential",
|
||||
"build_token_provider",
|
||||
"describe_active_credential",
|
||||
"has_azure_identity_credentials",
|
||||
"has_azure_identity_installed",
|
||||
"is_token_provider",
|
||||
"materialize_bearer_for_http",
|
||||
"reset_credential_cache",
|
||||
"EntraIdentityConfig", "SCOPE_AI_AZURE_DEFAULT", "build_bearer_http_client", "build_credential",
|
||||
"build_token_provider", "describe_active_credential", "has_azure_identity_credentials",
|
||||
"has_azure_identity_installed", "is_token_provider", "materialize_bearer_for_http", "reset_credential_cache",
|
||||
]
|
||||
|
||||
+68
-118
@@ -1,32 +1,11 @@
|
||||
"""Single owner for backend identity and failure-scoped skip decisions.
|
||||
|
||||
Every fallback / dedup / skip / quarantine decision in Hermes ultimately asks
|
||||
one question: **"is this candidate the same backend as the one that failed,
|
||||
along the axis that failure invalidated?"** Before this module, that
|
||||
question was re-implemented inline at six call sites across four subsystems,
|
||||
each comparing whatever string was locally convenient (provider label,
|
||||
provider+model, base_url+model, ...). Each incident fixed one site while the
|
||||
others kept the bug: #22548 (same-shim aliases), #70893 (xai-oauth vs xai —
|
||||
same host, distinct credential), #59561 (aux chain skipped sibling models),
|
||||
#72468 (aux main-model safety net, same bug three weeks later), #62984 /
|
||||
#54250 / #57584 (dedup ignoring base_url strands multi-endpoint pools).
|
||||
|
||||
The root insight: "provider" conflates three independent identity axes, and
|
||||
each failure class invalidates a different one:
|
||||
|
||||
* **credential surface** — auth 401 / payment 402 kill everything sharing the
|
||||
credential (every model, every host reached with that key/token).
|
||||
* **endpoint** — DNS failure / connection refused kill everything behind the
|
||||
URL, regardless of model or credential.
|
||||
* **model deployment** — timeout / overload / rate limit / model-incompatible
|
||||
kill ONE model's deployment. A sibling model behind the same URL is an
|
||||
independent deployment (real incident: aux ``glm-5.2`` hung and timed out
|
||||
while main ``macaron-v1-venti`` on the identical endpoint was serving
|
||||
448K-token turns).
|
||||
|
||||
Call sites should build :class:`BackendIdentity` values, classify the failure
|
||||
with :func:`classify_failure_scope`, and ask :func:`should_skip_candidate`.
|
||||
Do not re-implement any comparison inline — extend THIS module instead.
|
||||
Every fallback / dedup / skip / quarantine decision asks: "is this candidate the same backend
|
||||
as the one that failed, along the axis that failure invalidated?" Answering inline at each
|
||||
call site kept reintroducing the same bugs (same-shim aliases treated as distinct, sibling
|
||||
models skipped for one model's timeout, dedup ignoring ``base_url``). "provider" conflates
|
||||
three axes — credential surface (401/402), endpoint (DNS/refused), model deployment
|
||||
(timeout/overload/429). Build :class:`BackendIdentity` values, ask :func:`should_skip_candidate`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -36,62 +15,33 @@ from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Optional
|
||||
|
||||
from hermes_cli.route_identity import normalize_route_base_url
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FailureScope(Enum):
|
||||
"""Which identity axis a failure invalidates."""
|
||||
|
||||
#: Timeout, overload/429, connection blip, model-incompatible, invalid
|
||||
#: response: evidence against ONE model deployment only.
|
||||
#: Timeout, overload/429, connection blip, model-incompatible, invalid response:
|
||||
#: evidence against ONE model deployment only.
|
||||
MODEL = "model"
|
||||
#: Auth 401 / payment 402: evidence against the shared credential —
|
||||
#: every model reached with it is equally dead.
|
||||
#: Auth 401 / payment 402: evidence against the shared credential.
|
||||
CREDENTIAL = "credential"
|
||||
#: DNS / connection-refused / unreachable host: evidence against the
|
||||
#: endpoint — every model behind the URL is equally dead.
|
||||
#: DNS / connection-refused / unreachable host: evidence against the endpoint.
|
||||
ENDPOINT = "endpoint"
|
||||
|
||||
|
||||
#: Reason strings already used by auxiliary_client's except-chain, mapped to
|
||||
#: scopes. Unknown reasons default to MODEL — the least-invalidating scope —
|
||||
#: so an unrecognized failure never over-skips viable candidates.
|
||||
_REASON_SCOPES = {
|
||||
"auth error": FailureScope.CREDENTIAL,
|
||||
"payment error": FailureScope.CREDENTIAL,
|
||||
"rate limit": FailureScope.MODEL,
|
||||
"model incompatible with route": FailureScope.MODEL,
|
||||
"invalid provider response": FailureScope.MODEL,
|
||||
"connection error": FailureScope.MODEL,
|
||||
"timeout": FailureScope.MODEL,
|
||||
}
|
||||
|
||||
|
||||
def classify_failure_scope(reason: Optional[str]) -> FailureScope:
|
||||
"""Map a human-readable failure reason to the identity axis it kills."""
|
||||
return _REASON_SCOPES.get((reason or "").strip().lower(), FailureScope.MODEL)
|
||||
|
||||
|
||||
def _norm_provider(value: Optional[str]) -> str:
|
||||
def _norm(value: Optional[str]) -> str:
|
||||
return (value or "").strip().lower()
|
||||
|
||||
|
||||
def _norm_model(value: Optional[str]) -> str:
|
||||
return (value or "").strip().lower()
|
||||
|
||||
|
||||
def _norm_base_url(value: Optional[str]) -> str:
|
||||
return (value or "").strip().rstrip("/").lower()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BackendIdentity:
|
||||
"""Normalized identity of one (provider, model, endpoint) deployment.
|
||||
|
||||
Empty fields mean "unknown" — comparisons treat an unknown axis as
|
||||
non-distinguishing (it can neither prove sameness nor difference on its
|
||||
own; the remaining axes decide).
|
||||
"""
|
||||
Empty fields mean "unknown" — an unknown axis can neither prove sameness nor difference
|
||||
on its own; the remaining axes decide."""
|
||||
|
||||
provider: str = ""
|
||||
model: str = ""
|
||||
@@ -99,26 +49,21 @@ class BackendIdentity:
|
||||
|
||||
@classmethod
|
||||
def build(
|
||||
cls,
|
||||
provider: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
cls, provider: Optional[str] = None, model: Optional[str] = None,
|
||||
base_url: Optional[str] = None,
|
||||
) -> "BackendIdentity":
|
||||
return cls(
|
||||
provider=_norm_provider(provider),
|
||||
model=_norm_model(model),
|
||||
base_url=_norm_base_url(base_url),
|
||||
provider=_norm(provider), model=_norm(model),
|
||||
base_url=normalize_route_base_url(base_url),
|
||||
)
|
||||
|
||||
|
||||
def _both_first_class(a: BackendIdentity, b: BackendIdentity) -> bool:
|
||||
"""True when both providers are distinct registered first-class providers.
|
||||
|
||||
Two different registry providers have distinct credential surfaces even
|
||||
when they share an inference host (xai-oauth vs xai, openai-codex vs
|
||||
openai-api) — #70893. Custom/shim aliases are NOT in the registry, so
|
||||
two aliases pointing at one URL still count as the same backend (#22548).
|
||||
"""
|
||||
Two different registry providers have distinct credential surfaces even when they share an
|
||||
inference host (xai-oauth vs xai). Custom/shim aliases are NOT in the registry, so two
|
||||
aliases pointing at one URL still count as the same backend."""
|
||||
if not a.provider or not b.provider or a.provider == b.provider:
|
||||
return False
|
||||
try:
|
||||
@@ -132,73 +77,78 @@ def _both_first_class(a: BackendIdentity, b: BackendIdentity) -> bool:
|
||||
def same_credential_surface(a: BackendIdentity, b: BackendIdentity) -> bool:
|
||||
"""Do two identities share the credential a 401/402 just invalidated?
|
||||
|
||||
Conservative on purpose: an unprovable axis must answer "different"
|
||||
(try the candidate — worst case one wasted RTT) rather than "same"
|
||||
(skip — worst case stranded failover). Two distinct custom labels at
|
||||
one URL may carry different per-entry api_keys, so a shared URL alone
|
||||
never proves a shared credential; it is only used as a weak signal
|
||||
when a provider label is missing entirely.
|
||||
"""
|
||||
Conservative: an unprovable axis answers "different" (one wasted RTT) rather than "same"
|
||||
(stranded failover). Same label = same configured credential; custom entries can each carry
|
||||
their own api_key, so a shared URL alone is only a weak signal when a label is missing."""
|
||||
if a.provider and b.provider:
|
||||
# Same label = same configured credential. Different labels =
|
||||
# different credential config (first-class registry providers
|
||||
# explicitly so — #70893; custom entries can each carry their own
|
||||
# api_key, so sameness is unprovable and we must not skip).
|
||||
# Different labels = different credential config (first-class registry providers explicitly so —
|
||||
# #70893; custom entries can each carry their own api_key, so sameness is unprovable and we must not
|
||||
# skip).
|
||||
return a.provider == b.provider
|
||||
# Provider unknown on a side: same explicit URL is the best signal left.
|
||||
return bool(a.base_url and a.base_url == b.base_url)
|
||||
|
||||
|
||||
def same_endpoint(a: BackendIdentity, b: BackendIdentity) -> bool:
|
||||
"""Do two identities sit behind the endpoint that just went unreachable?"""
|
||||
"""Do two identities sit behind the endpoint that just went unreachable?
|
||||
An unknown base_url inherits the provider default, so a shared label implies the same endpoint."""
|
||||
if a.base_url and b.base_url:
|
||||
return a.base_url == b.base_url
|
||||
# An unknown base_url inherits the provider default → same provider
|
||||
# label implies the same default endpoint.
|
||||
return bool(a.provider and a.provider == b.provider)
|
||||
|
||||
|
||||
def same_deployment(a: BackendIdentity, b: BackendIdentity) -> bool:
|
||||
"""Are these the exact same model deployment (the thing a timeout kills)?
|
||||
|
||||
Provider+model must match; the base_url axis distinguishes only when BOTH
|
||||
sides carry an explicit URL (#62984: same provider+model on two different
|
||||
explicit URLs is two deployments — a pool). A side with an unknown URL
|
||||
inherits the provider default and cannot prove difference.
|
||||
"""
|
||||
Provider+model must match; base_url distinguishes only when BOTH sides carry an explicit URL
|
||||
(same provider+model on two explicit URLs is a pool, not a dup). Different labels with the
|
||||
same URL + model are still one deployment (same-host shim aliases) — unless both labels are
|
||||
first-class registry providers."""
|
||||
if not (a.provider and b.provider and a.provider == b.provider):
|
||||
# Same-host different-label shims: same URL + same model IS the same
|
||||
# deployment even when the alias labels differ (#22548) — unless both
|
||||
# labels are first-class registry providers (#70893).
|
||||
if (
|
||||
return bool(
|
||||
a.base_url
|
||||
# Same-host different-label shims: same URL + same model IS the same deployment even when the
|
||||
# alias labels differ (#22548) — unless both labels are first-class registry providers (#70893).
|
||||
and a.base_url == b.base_url
|
||||
and a.model
|
||||
and a.model == b.model
|
||||
and not _both_first_class(a, b)
|
||||
):
|
||||
return True
|
||||
return False
|
||||
)
|
||||
if not (a.model and b.model and a.model == b.model):
|
||||
return False
|
||||
if a.base_url and b.base_url and a.base_url != b.base_url:
|
||||
return False # distinct explicit endpoints — a pool, not a dup
|
||||
return True
|
||||
return not (a.base_url and b.base_url and a.base_url != b.base_url)
|
||||
|
||||
|
||||
_SCOPE_PREDICATES = {
|
||||
FailureScope.CREDENTIAL: same_credential_surface, FailureScope.ENDPOINT: same_endpoint,
|
||||
FailureScope.MODEL: same_deployment,
|
||||
}
|
||||
|
||||
|
||||
def should_skip_candidate(
|
||||
candidate: BackendIdentity,
|
||||
failed: BackendIdentity,
|
||||
scope: FailureScope = FailureScope.MODEL,
|
||||
candidate: BackendIdentity, failed: BackendIdentity, scope: FailureScope = FailureScope.MODEL
|
||||
) -> bool:
|
||||
"""THE skip predicate: would trying ``candidate`` just repeat the failure?
|
||||
True when it is the same backend as ``failed`` along the axis ``scope`` invalidated.
|
||||
Every fallback/dedup/skip site must call this."""
|
||||
return _SCOPE_PREDICATES.get(scope, same_deployment)(candidate, failed)
|
||||
|
||||
True when the candidate is the same backend as ``failed`` along the axis
|
||||
``scope`` says the failure invalidated. Every fallback/dedup/skip site
|
||||
must call this instead of comparing labels inline.
|
||||
"""
|
||||
if scope is FailureScope.CREDENTIAL:
|
||||
return same_credential_surface(candidate, failed)
|
||||
if scope is FailureScope.ENDPOINT:
|
||||
return same_endpoint(candidate, failed)
|
||||
return same_deployment(candidate, failed)
|
||||
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
|
||||
_REASON_SCOPES = {
|
||||
"auth error": FailureScope.CREDENTIAL,
|
||||
"payment error": FailureScope.CREDENTIAL,
|
||||
"rate limit": FailureScope.MODEL,
|
||||
"model incompatible with route": FailureScope.MODEL,
|
||||
"invalid provider response": FailureScope.MODEL,
|
||||
"connection error": FailureScope.MODEL,
|
||||
"timeout": FailureScope.MODEL,
|
||||
}
|
||||
|
||||
def classify_failure_scope(reason: Optional[str]) -> FailureScope:
|
||||
"""Map a human-readable failure reason to the identity axis it kills."""
|
||||
return _REASON_SCOPES.get((reason or "").strip().lower(), FailureScope.MODEL)
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
+1099
-1584
File diff suppressed because it is too large
Load Diff
+20
-49
@@ -1,13 +1,8 @@
|
||||
"""System-battery read-out for the CLI/TUI status bar.
|
||||
|
||||
Reads the host battery through ``psutil`` (already a Hermes dependency) and
|
||||
exposes a compact, colour-coded label. Everything degrades to "unavailable"
|
||||
when there is no battery (desktops, servers, VMs) or when the read fails, so
|
||||
callers can render the result unconditionally and simply show nothing.
|
||||
|
||||
The status bar repaints often (every keystroke and on a ~1s idle refresh), so
|
||||
:func:`read_battery` memoises the last reading for a few seconds instead of
|
||||
hitting ``psutil`` on every frame.
|
||||
Reads the host battery through ``psutil`` and exposes a compact, colour-coded label. Everything
|
||||
degrades to "unavailable" (no battery / read failure) so callers can render unconditionally. The
|
||||
status bar repaints on every keystroke, so :func:`read_battery` memoises the reading for a few seconds.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -19,12 +14,7 @@ from typing import Optional
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BatteryStatus:
|
||||
"""A single battery reading.
|
||||
|
||||
``available`` is False on machines without a battery (or when the read
|
||||
failed). ``percent`` is clamped to 0-100. ``plugged`` is True when on AC
|
||||
power, False on battery, and None when the platform can't tell.
|
||||
"""
|
||||
"""One reading: ``percent`` clamped 0-100; ``plugged`` None when the platform can't tell."""
|
||||
|
||||
available: bool
|
||||
percent: Optional[int] = None
|
||||
@@ -37,14 +27,16 @@ class BatteryStatus:
|
||||
|
||||
UNAVAILABLE = BatteryStatus(available=False)
|
||||
|
||||
# Colour buckets, mirroring the status-bar context styles but inverted (a full
|
||||
# battery is "good", an empty one is "critical").
|
||||
# Colour buckets, mirroring the status-bar context styles but inverted (full battery = "good").
|
||||
CATEGORY_GOOD = "good"
|
||||
CATEGORY_WARN = "warn"
|
||||
CATEGORY_BAD = "bad"
|
||||
CATEGORY_CRITICAL = "critical"
|
||||
CATEGORY_DIM = "dim"
|
||||
|
||||
# (upper bound inclusive, category) for a discharging battery; first match wins.
|
||||
_LEVEL_CATEGORIES = ((10, CATEGORY_CRITICAL), (20, CATEGORY_BAD), (50, CATEGORY_WARN))
|
||||
|
||||
_CACHE_TTL_SECONDS = 8.0
|
||||
_cache: Optional[tuple[float, BatteryStatus]] = None
|
||||
|
||||
@@ -52,22 +44,13 @@ _cache: Optional[tuple[float, BatteryStatus]] = None
|
||||
def _read_battery_uncached() -> BatteryStatus:
|
||||
try:
|
||||
import psutil
|
||||
|
||||
# ``sensors_battery`` is missing on some platforms/builds of psutil.
|
||||
batt = getattr(psutil, "sensors_battery")()
|
||||
except Exception:
|
||||
return UNAVAILABLE
|
||||
|
||||
# ``sensors_battery`` is missing on some platforms/builds of psutil.
|
||||
reader = getattr(psutil, "sensors_battery", None)
|
||||
if reader is None:
|
||||
return UNAVAILABLE
|
||||
|
||||
try:
|
||||
batt = reader()
|
||||
except Exception:
|
||||
return UNAVAILABLE
|
||||
|
||||
if batt is None:
|
||||
return UNAVAILABLE
|
||||
|
||||
percent: Optional[int] = None
|
||||
raw_percent = getattr(batt, "percent", None)
|
||||
if raw_percent is not None:
|
||||
@@ -75,22 +58,15 @@ def _read_battery_uncached() -> BatteryStatus:
|
||||
percent = max(0, min(100, int(round(float(raw_percent)))))
|
||||
except (TypeError, ValueError):
|
||||
percent = None
|
||||
|
||||
plugged = getattr(batt, "power_plugged", None)
|
||||
if plugged is not None:
|
||||
plugged = bool(plugged)
|
||||
|
||||
return BatteryStatus(available=True, percent=percent, plugged=plugged)
|
||||
return BatteryStatus(available=True, percent=percent, plugged=None if plugged is None else bool(plugged))
|
||||
|
||||
|
||||
def read_battery(use_cache: bool = True) -> BatteryStatus:
|
||||
"""Return the current battery status (cached for a few seconds)."""
|
||||
global _cache
|
||||
if use_cache and _cache is not None:
|
||||
ts, cached = _cache
|
||||
if time.monotonic() - ts < _CACHE_TTL_SECONDS:
|
||||
return cached
|
||||
|
||||
if use_cache and _cache is not None and time.monotonic() - _cache[0] < _CACHE_TTL_SECONDS:
|
||||
return _cache[1]
|
||||
status = _read_battery_uncached()
|
||||
_cache = (time.monotonic(), status)
|
||||
return status
|
||||
@@ -106,26 +82,21 @@ def battery_category(status: BatteryStatus) -> str:
|
||||
"""Bucket a reading into a colour category: good/warn/bad/critical/dim."""
|
||||
if not status.available or status.percent is None:
|
||||
return CATEGORY_DIM
|
||||
# On AC power the level isn't a concern — always read as healthy.
|
||||
if status.charging:
|
||||
if status.charging: # on AC power the level isn't a concern
|
||||
return CATEGORY_GOOD
|
||||
pct = status.percent
|
||||
if pct <= 10:
|
||||
return CATEGORY_CRITICAL
|
||||
if pct <= 20:
|
||||
return CATEGORY_BAD
|
||||
if pct <= 50:
|
||||
return CATEGORY_WARN
|
||||
for bound, category in _LEVEL_CATEGORIES:
|
||||
if status.percent <= bound:
|
||||
return category
|
||||
return CATEGORY_GOOD
|
||||
|
||||
|
||||
def battery_glyph(status: BatteryStatus) -> str:
|
||||
"""Return the leading glyph: a bolt while charging, else a battery."""
|
||||
"""Leading glyph: a bolt while charging, else a battery."""
|
||||
return "\u26a1" if status.charging else "\U0001f50b" # ⚡ / 🔋
|
||||
|
||||
|
||||
def format_battery(status: BatteryStatus) -> str:
|
||||
"""Return a compact label like ``🔋 82%`` / ``⚡ 82%`` (empty if N/A)."""
|
||||
"""Compact label like ``🔋 82%`` / ``⚡ 82%`` (empty if N/A)."""
|
||||
if not status.available or status.percent is None:
|
||||
return ""
|
||||
return f"{battery_glyph(status)} {status.percent}%"
|
||||
|
||||
+731
-1405
File diff suppressed because it is too large
Load Diff
+17
-45
@@ -1,11 +1,8 @@
|
||||
"""Provider-agnostic billing/credit recovery links.
|
||||
"""Provider-agnostic billing/credit recovery links for billing-classified failures.
|
||||
|
||||
Maps a billing-classified failure onto a recovery link + label. *Detection*
|
||||
is not done here — that is :mod:`agent.error_classifier`
|
||||
(``FailoverReason.billing``), the single source of truth for "credit wall vs.
|
||||
rate limit / auth / transport". The resulting :class:`BillingBlock` rides the
|
||||
turn result and the gateway ``message.complete`` event so every surface (CLI,
|
||||
TUI, desktop) renders one structured signal instead of re-parsing error text.
|
||||
Detection lives in :mod:`agent.error_classifier` (``FailoverReason.billing``); the
|
||||
:class:`BillingBlock` rides the turn result and the gateway ``message.complete`` event
|
||||
so every surface renders one structured signal instead of re-parsing error text.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -18,13 +15,9 @@ from utils import base_url_host_matches
|
||||
|
||||
@dataclass
|
||||
class BillingBlock:
|
||||
"""Structured billing-wall descriptor shared across every surface.
|
||||
|
||||
``is_nous`` is the routing bit: Nous has a first-class in-app billing surface
|
||||
(desktop Settings → Billing, TUI/CLI ``/topup``), so surfaces prefer that over
|
||||
``billing_url``; third-party providers have no in-app flow, so ``billing_url``
|
||||
is the deep link the user actually needs.
|
||||
"""
|
||||
"""Structured billing-wall descriptor shared across every surface. ``is_nous`` is the routing bit:
|
||||
Nous has an in-app billing surface (Settings → Billing, ``/topup``) preferred over ``billing_url``;
|
||||
third-party providers have none, so ``billing_url`` is the deep link needed."""
|
||||
|
||||
provider: str
|
||||
provider_label: str
|
||||
@@ -45,11 +38,9 @@ class _Provider:
|
||||
hosts: tuple[str, ...] = ()
|
||||
|
||||
|
||||
# Single source of truth: internal slug(s) + base_url host(s) → billing page.
|
||||
# Curated "add credits / manage billing" landing pages, not marketing homes.
|
||||
# Hosts back the OpenAI-compatible fallback where the slug is a generic bucket
|
||||
# (e.g. "openai_compatible") but base_url reveals the real upstream. An unknown
|
||||
# provider degrades to a readable label with no invented URL.
|
||||
# Single source of truth: slug(s) + base_url host(s) → curated "add credits" page.
|
||||
# Hosts back the OpenAI-compatible fallback where the slug is a generic bucket but
|
||||
# base_url reveals the real upstream. Unknown providers get a readable label, no URL.
|
||||
_PROVIDERS: tuple[_Provider, ...] = (
|
||||
_Provider("OpenAI", "https://platform.openai.com/settings/organization/billing", ("openai",), ("api.openai.com",)),
|
||||
_Provider("Anthropic", "https://console.anthropic.com/settings/billing", ("anthropic",), ("api.anthropic.com",)),
|
||||
@@ -72,16 +63,13 @@ _BY_SLUG: dict[str, _Provider] = {slug: p for p in _PROVIDERS for slug in p.slug
|
||||
|
||||
def is_nous_inference_route(provider: str, base_url: str) -> bool:
|
||||
"""True when the failing route is the Nous-managed inference gateway."""
|
||||
if (provider or "").strip().lower() == "nous":
|
||||
return True
|
||||
return base_url_host_matches(str(base_url or ""), "inference-api.nousresearch.com")
|
||||
return (provider or "").strip().lower() == "nous" or base_url_host_matches(str(base_url or ""), "inference-api.nousresearch.com")
|
||||
|
||||
|
||||
def _nous_billing_url() -> Optional[str]:
|
||||
"""Best-effort Nous portal billing URL (text-surface fallback; Nous prefers the in-app flow)."""
|
||||
try:
|
||||
from hermes_cli.nous_account import nous_portal_billing_url
|
||||
|
||||
return nous_portal_billing_url(None)
|
||||
except Exception:
|
||||
return "https://portal.nousresearch.com/billing"
|
||||
@@ -89,36 +77,20 @@ def _nous_billing_url() -> Optional[str]:
|
||||
|
||||
def _resolve_provider_link(slug: str, base_url: str) -> tuple[str, Optional[str]]:
|
||||
"""Resolve ``(label, url)``: exact slug → base_url host → readable-label fallback."""
|
||||
hit = _BY_SLUG.get(slug)
|
||||
base = str(base_url or "")
|
||||
hit = _BY_SLUG.get(slug) or next(
|
||||
(p for p in _PROVIDERS if any(base_url_host_matches(base, host) for host in p.hosts)), None
|
||||
)
|
||||
if hit:
|
||||
return hit.label, hit.url
|
||||
|
||||
base = str(base_url or "")
|
||||
for p in _PROVIDERS:
|
||||
if any(base_url_host_matches(base, host) for host in p.hosts):
|
||||
return p.label, p.url
|
||||
|
||||
return slug.replace("_", " ").replace("-", " ").strip().title() or "your provider", None
|
||||
|
||||
|
||||
def build_billing_block(
|
||||
*,
|
||||
provider: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
message: str = "",
|
||||
) -> BillingBlock:
|
||||
"""Build the billing descriptor for a billing-classified failure.
|
||||
|
||||
``message`` is the guidance already assembled by the agent loop
|
||||
(:func:`agent.conversation_loop._billing_or_entitlement_message`), carried
|
||||
through unchanged so every surface shows identical copy.
|
||||
"""
|
||||
def build_billing_block(*, provider: str, base_url: str, model: str, message: str = "") -> BillingBlock:
|
||||
"""Billing descriptor for a billing-classified failure; ``message`` (agent-loop guidance) passes through unchanged."""
|
||||
slug = (provider or "").strip().lower()
|
||||
model = (model or "").strip()
|
||||
|
||||
if is_nous_inference_route(slug, base_url):
|
||||
return BillingBlock(slug or "nous", "Nous Portal", model, _nous_billing_url(), True, message or "")
|
||||
|
||||
label, url = _resolve_provider_link(slug, base_url)
|
||||
return BillingBlock(slug, label, model, url, False, message or "")
|
||||
|
||||
+81
-203
@@ -1,32 +1,11 @@
|
||||
"""Shared dollar-denominated usage model for the billing/subscription surfaces.
|
||||
"""Shared dollar-denominated usage model for the ``/usage`` and ``/subscription`` bars.
|
||||
|
||||
The single source of truth behind the ``/usage`` and ``/subscription`` usage
|
||||
bars (TUI + CLI). User feedback (Jun 2026): the terminal surfaces show
|
||||
**dollars**, never "credits", and every usage bar must make the monthly
|
||||
subscription allowance and separately-purchased top-up dollars distinctly
|
||||
visible.
|
||||
|
||||
Data source: the NAS account-info fetch (``NousPortalAccountInfo``), whose
|
||||
``paid_service_access_info`` carries the three dollar magnitudes we render
|
||||
(despite the legacy ``*_credits`` field names, these are USD floats):
|
||||
|
||||
- ``subscription_credits_remaining`` -> plan dollars left this month
|
||||
- ``purchased_credits_remaining`` -> top-up dollars left (rolls over)
|
||||
- ``total_usable_credits`` -> total spendable
|
||||
|
||||
plus ``subscription.monthly_credits`` (the plan's monthly $ allowance, the
|
||||
denominator for the "% used" plan bar) and ``current_period_end`` (renewal).
|
||||
|
||||
Design: two SEPARATE bars (decided with the user) rather than one crammed
|
||||
three-segment bar — at terminal widths three same-glyph density segments are
|
||||
unreadable. The plan bar is "spent vs allowance this month" (carries % used);
|
||||
the top-up bar is "money you bought, doesn't expire". Each gets full
|
||||
resolution and a single fill glyph, so the bar is never ambiguous and never
|
||||
relies on color.
|
||||
|
||||
Fail-open everywhere: any missing/non-finite field degrades to fewer bars or a
|
||||
magnitudes-only view; a logged-out / unreachable portal yields
|
||||
``available=False`` and the surface shows nothing.
|
||||
Terminal surfaces show **dollars**, never "credits"; the plan allowance and top-up
|
||||
dollars stay distinctly visible as two SEPARATE bars (a three-segment bar is
|
||||
unreadable at terminal widths). Source: ``NousPortalAccountInfo.paid_service_access_info``
|
||||
(USD floats despite the legacy ``*_credits`` names) plus ``subscription.monthly_credits``
|
||||
(plan bar denominator) and ``current_period_end``. Fail-open: missing/non-finite fields
|
||||
degrade to fewer bars; logged-out / unreachable portal yields ``available=False``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -39,14 +18,13 @@ from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Below this TOTAL spendable ($), a paid account is flagged "low" — the alert
|
||||
# state that nudges top-up/upgrade before a mid-run cutoff. Product threshold
|
||||
# (user feedback): "any amount below $5 should be an alert status."
|
||||
# Below this TOTAL spendable ($) a paid account is flagged "low" — the alert state
|
||||
# that nudges top-up/upgrade before a mid-run cutoff (product: "below $5 is an alert").
|
||||
LOW_BALANCE_THRESHOLD_USD = 5.0
|
||||
|
||||
|
||||
def _finite(value: Any) -> Optional[float]:
|
||||
"""Return value as a float iff it's a real finite number (not bool/NaN/Inf)."""
|
||||
"""Float iff a real finite number (not bool/NaN/Inf); json.loads admits bare NaN, which would render ``$nan``."""
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
return None
|
||||
f = float(value)
|
||||
@@ -54,29 +32,38 @@ def _finite(value: Any) -> Optional[float]:
|
||||
|
||||
|
||||
def _fmt_usd(value: Optional[float]) -> str:
|
||||
"""``$X.YY`` for display. ``None`` -> ``$0.00`` (callers gate on presence)."""
|
||||
"""``$X,XXX.YY`` for display. ``None`` -> ``$0.00`` (callers gate on presence)."""
|
||||
return f"${(value or 0.0):,.2f}"
|
||||
|
||||
|
||||
def nous_logged_in() -> bool:
|
||||
"""Cheap local auth-state check: a Nous access token is present. Fail-closed."""
|
||||
try:
|
||||
from hermes_cli.auth import get_provider_auth_state
|
||||
tok = (get_provider_auth_state("nous") or {}).get("access_token")
|
||||
return isinstance(tok, str) and bool(tok.strip())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def fetch_nous_account(timeout: float):
|
||||
"""Wall-clock-bounded fresh portal account fetch. Raises on failure/timeout."""
|
||||
import concurrent.futures
|
||||
from hermes_cli.nous_account import get_nous_portal_account_info
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||
return pool.submit(get_nous_portal_account_info, force_fresh=True).result(timeout=timeout)
|
||||
|
||||
|
||||
def format_renews(value: Optional[str]) -> Optional[str]:
|
||||
"""Format an ISO date/timestamp as a human date, e.g. ``Jul 24, 2026``.
|
||||
|
||||
Accepts ``2026-07-24``, ``2026-07-24T11:05:01.000Z``, etc. Returns the raw
|
||||
string unchanged if it can't be parsed (never raises), and ``None`` for
|
||||
empty input.
|
||||
"""
|
||||
if not value:
|
||||
return None
|
||||
"""ISO date/timestamp -> ``Jul 24, 2026``; unparseable input is returned unchanged."""
|
||||
from datetime import datetime
|
||||
|
||||
text = str(value).strip()
|
||||
text = str(value or "").strip()
|
||||
if not text:
|
||||
return None
|
||||
iso = text[:-1] + "+00:00" if text.endswith("Z") else text
|
||||
try:
|
||||
dt = datetime.fromisoformat(iso)
|
||||
except ValueError:
|
||||
# Fall back to a bare date prefix (YYYY-MM-DD) if present.
|
||||
try:
|
||||
dt = datetime.strptime(text[:10], "%Y-%m-%d")
|
||||
except ValueError:
|
||||
@@ -87,12 +74,8 @@ def format_renews(value: Optional[str]) -> Optional[str]:
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UsageBar:
|
||||
"""One full-resolution bar: ``spent`` of ``total``, plus a remaining figure.
|
||||
|
||||
``kind`` is ``"plan"`` (monthly allowance, shows % used) or ``"topup"``
|
||||
(purchased dollars, no denominator — ``spent`` is 0 and ``total`` ==
|
||||
``remaining`` so it renders as a full bar of available balance).
|
||||
"""
|
||||
"""One bar: ``spent`` of ``total`` plus remaining. ``plan`` shows % used; ``topup`` has no
|
||||
denominator (``spent`` is 0 and ``total == remaining`` so it renders full)."""
|
||||
|
||||
kind: str # "plan" | "topup"
|
||||
remaining_usd: float
|
||||
@@ -108,21 +91,13 @@ class UsageBar:
|
||||
@property
|
||||
def fill_fraction(self) -> float:
|
||||
"""Fraction of the bar that should read as 'remaining' (filled)."""
|
||||
if self.total_usd <= 0:
|
||||
return 0.0
|
||||
return max(0.0, min(1.0, self.remaining_usd / self.total_usd))
|
||||
return max(0.0, min(1.0, self.remaining_usd / self.total_usd)) if self.total_usd > 0 else 0.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UsageModel:
|
||||
"""Surface-agnostic dollar usage model shared by /usage and /subscription.
|
||||
|
||||
``status`` classifies the account for copy selection:
|
||||
- ``"free"`` : no paid access / no subscription (free models only)
|
||||
- ``"low"`` : paid, but total spendable < $5 (ALERT)
|
||||
- ``"healthy"`` : paid, total spendable >= $5
|
||||
- ``"depleted"`` : paid access lost (balance exhausted)
|
||||
"""
|
||||
"""Dollar usage model shared by /usage and /subscription. ``status``: ``free`` (no plan / no
|
||||
paid access), ``low`` (spendable < $5, ALERT), ``healthy`` (>= $5), ``depleted`` (paid access lost)."""
|
||||
|
||||
available: bool
|
||||
status: str = "free"
|
||||
@@ -141,82 +116,47 @@ class UsageModel:
|
||||
|
||||
|
||||
def usage_model_from_account(account_info: Any) -> UsageModel:
|
||||
"""Build a :class:`UsageModel` from a ``NousPortalAccountInfo``. Fail-open.
|
||||
|
||||
Returns ``UsageModel(available=False)`` when there's no usable account info
|
||||
(logged out, no entitlement block). Never raises.
|
||||
"""
|
||||
"""Build a :class:`UsageModel` from a ``NousPortalAccountInfo``. Never raises."""
|
||||
try:
|
||||
if account_info is None or not getattr(account_info, "logged_in", False):
|
||||
return UsageModel(available=False)
|
||||
|
||||
# Sub-structs are dataclasses or None; getattr(None, ..., None) is None, so no guards needed.
|
||||
access = getattr(account_info, "paid_service_access_info", None)
|
||||
sub = getattr(account_info, "subscription", None)
|
||||
paid = getattr(account_info, "paid_service_access", None)
|
||||
|
||||
sub_remaining = _finite(getattr(access, "subscription_credits_remaining", None)) if access else None
|
||||
topup_remaining = _finite(getattr(access, "purchased_credits_remaining", None)) if access else None
|
||||
total_usable = _finite(getattr(access, "total_usable_credits", None)) if access else None
|
||||
|
||||
plan_name = getattr(sub, "plan", None) if sub is not None else None
|
||||
renews_at = getattr(sub, "current_period_end", None) if sub is not None else None
|
||||
monthly = _finite(getattr(sub, "monthly_credits", None)) if sub is not None else None
|
||||
|
||||
sub_remaining = _finite(getattr(access, "subscription_credits_remaining", None))
|
||||
topup_remaining = _finite(getattr(access, "purchased_credits_remaining", None))
|
||||
total_usable = _finite(getattr(access, "total_usable_credits", None))
|
||||
plan_name = getattr(sub, "plan", None)
|
||||
renews_at = getattr(sub, "current_period_end", None)
|
||||
monthly = _finite(getattr(sub, "monthly_credits", None))
|
||||
has_subscription = bool(plan_name) or (monthly is not None and monthly > 0)
|
||||
|
||||
# Total spendable: prefer the server's total; else sum the parts we have.
|
||||
if total_usable is not None:
|
||||
total_spendable = total_usable
|
||||
else:
|
||||
parts = [v for v in (sub_remaining, topup_remaining) if v is not None]
|
||||
total_spendable = sum(parts) if parts else None
|
||||
|
||||
# Status classification.
|
||||
has_topup = bool(topup_remaining and topup_remaining > 0)
|
||||
# Prefer the server's total; else sum the parts we have.
|
||||
parts = [v for v in (sub_remaining, topup_remaining) if v is not None]
|
||||
total_spendable = total_usable if total_usable is not None else (sum(parts) if parts else None)
|
||||
if paid is False:
|
||||
status = "depleted"
|
||||
elif not has_subscription and not (topup_remaining and topup_remaining > 0):
|
||||
# No plan and no purchased balance -> free-models-only.
|
||||
status = "free"
|
||||
elif not has_subscription and not has_topup:
|
||||
status = "free" # no plan and no purchased balance -> free-models-only
|
||||
elif total_spendable is not None and total_spendable < LOW_BALANCE_THRESHOLD_USD:
|
||||
status = "low"
|
||||
else:
|
||||
status = "healthy"
|
||||
|
||||
# Plan bar — only with a positive monthly allowance AND a remaining we
|
||||
# can place on it. spent = cap - remaining, clamped (a debt/over-cap
|
||||
# balance reads as fully spent rather than a nonsensical negative).
|
||||
plan_bar: Optional[UsageBar] = None
|
||||
# Plan bar needs a positive allowance AND a remaining to place on it; spent is
|
||||
# clamped so a debt/over-cap balance reads fully spent, not negative.
|
||||
plan_bar = topup_bar = None
|
||||
if monthly is not None and monthly > 0 and sub_remaining is not None:
|
||||
remaining = max(0.0, min(monthly, sub_remaining))
|
||||
plan_bar = UsageBar(
|
||||
kind="plan",
|
||||
remaining_usd=remaining,
|
||||
total_usd=monthly,
|
||||
spent_usd=max(0.0, monthly - sub_remaining),
|
||||
)
|
||||
|
||||
# Top-up bar — only when there are purchased dollars to show. No
|
||||
# denominator (top-up has no monthly cap), so it renders full = balance.
|
||||
topup_bar: Optional[UsageBar] = None
|
||||
plan_bar = UsageBar(kind="plan", remaining_usd=max(0.0, min(monthly, sub_remaining)), total_usd=monthly,
|
||||
spent_usd=max(0.0, monthly - sub_remaining))
|
||||
# Top-up has no monthly cap, so the bar renders full = balance.
|
||||
if topup_remaining is not None and topup_remaining > 0:
|
||||
topup_bar = UsageBar(
|
||||
kind="topup",
|
||||
remaining_usd=topup_remaining,
|
||||
total_usd=topup_remaining,
|
||||
spent_usd=0.0,
|
||||
)
|
||||
|
||||
topup_bar = UsageBar(kind="topup", remaining_usd=topup_remaining, total_usd=topup_remaining, spent_usd=0.0)
|
||||
return UsageModel(
|
||||
available=True,
|
||||
status=status,
|
||||
plan_name=plan_name,
|
||||
renews_at=renews_at,
|
||||
renews_display=format_renews(renews_at),
|
||||
subscription_remaining_usd=sub_remaining,
|
||||
topup_remaining_usd=topup_remaining,
|
||||
total_spendable_usd=total_spendable,
|
||||
plan_bar=plan_bar,
|
||||
topup_bar=topup_bar,
|
||||
available=True, status=status, plan_name=plan_name, renews_at=renews_at,
|
||||
renews_display=format_renews(renews_at), subscription_remaining_usd=sub_remaining,
|
||||
topup_remaining_usd=topup_remaining, total_spendable_usd=total_spendable,
|
||||
plan_bar=plan_bar, topup_bar=topup_bar,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("usage ▸ model build failed (fail-open)", exc_info=True)
|
||||
@@ -224,100 +164,38 @@ def usage_model_from_account(account_info: Any) -> UsageModel:
|
||||
|
||||
|
||||
def build_usage_model(*, timeout: float = 10.0) -> UsageModel:
|
||||
"""Fetch account-info and build the shared usage model. Fail-open.
|
||||
|
||||
Dev override: ``HERMES_DEV_CREDITS_FIXTURE`` short-circuits to a fixture so
|
||||
every usage state is testable without a live account (mirrors the existing
|
||||
``/usage`` credits-block fixture path).
|
||||
"""
|
||||
"""Fetch account-info and build the usage model; fail-open. ``HERMES_DEV_CREDITS_FIXTURE`` short-circuits to a fixture."""
|
||||
fixture = _dev_fixture_usage_model()
|
||||
if fixture is not None:
|
||||
return fixture
|
||||
|
||||
try:
|
||||
from hermes_cli.auth import get_provider_auth_state
|
||||
|
||||
tok = (get_provider_auth_state("nous") or {}).get("access_token")
|
||||
if not (isinstance(tok, str) and tok.strip()):
|
||||
return UsageModel(available=False)
|
||||
except Exception:
|
||||
if not nous_logged_in():
|
||||
return UsageModel(available=False)
|
||||
|
||||
try:
|
||||
import concurrent.futures
|
||||
|
||||
from hermes_cli.nous_account import get_nous_portal_account_info
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||
account = pool.submit(get_nous_portal_account_info, force_fresh=True).result(timeout=timeout)
|
||||
return usage_model_from_account(account)
|
||||
return usage_model_from_account(fetch_nous_account(timeout))
|
||||
except Exception:
|
||||
logger.debug("usage ▸ portal fetch failed (fail-open)", exc_info=True)
|
||||
return UsageModel(available=False)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Dev fixtures (throwaway scaffolding — env-var driven, no live portal)
|
||||
# =============================================================================
|
||||
def _plan_bar(remaining: float, spent: float) -> UsageBar:
|
||||
return UsageBar(kind="plan", remaining_usd=remaining, total_usd=20.0, spent_usd=spent)
|
||||
|
||||
|
||||
def _dev_fixture_usage_model() -> Optional[UsageModel]:
|
||||
"""Map ``HERMES_DEV_CREDITS_FIXTURE`` to a usage model for offline UX work.
|
||||
|
||||
Recognized names: ``free | healthy | low | topup | depleted``. Returns
|
||||
``None`` when the env var is unset (real portal path runs).
|
||||
"""
|
||||
"""``HERMES_DEV_CREDITS_FIXTURE`` -> fixture model (``free|healthy|low|topup|depleted``), else None."""
|
||||
name = (os.getenv("HERMES_DEV_CREDITS_FIXTURE") or "").strip().lower()
|
||||
if not name:
|
||||
return None
|
||||
|
||||
if name == "free":
|
||||
return UsageModel(available=True, status="free", plan_name=None)
|
||||
|
||||
if name in ("healthy", "mid"):
|
||||
return UsageModel(
|
||||
available=True,
|
||||
status="healthy",
|
||||
plan_name="Plus",
|
||||
renews_at="2026-07-01",
|
||||
subscription_remaining_usd=14.0,
|
||||
total_spendable_usd=14.0,
|
||||
plan_bar=UsageBar(kind="plan", remaining_usd=14.0, total_usd=20.0, spent_usd=6.0),
|
||||
)
|
||||
|
||||
if name in ("topup", "top-up"):
|
||||
return UsageModel(
|
||||
available=True,
|
||||
status="healthy",
|
||||
plan_name="Plus",
|
||||
renews_at="2026-07-01",
|
||||
subscription_remaining_usd=14.0,
|
||||
topup_remaining_usd=12.0,
|
||||
total_spendable_usd=26.0,
|
||||
plan_bar=UsageBar(kind="plan", remaining_usd=14.0, total_usd=20.0, spent_usd=6.0),
|
||||
name = {"mid": "healthy", "top-up": "topup"}.get(name, name)
|
||||
plus = dict(available=True, plan_name="Plus", renews_at="2026-07-01")
|
||||
specs: dict[str, dict] = {
|
||||
"free": dict(available=True, status="free", plan_name=None),
|
||||
"healthy": dict(**plus, status="healthy", subscription_remaining_usd=14.0, total_spendable_usd=14.0, plan_bar=_plan_bar(14.0, 6.0)),
|
||||
"topup": dict(
|
||||
**plus, status="healthy", subscription_remaining_usd=14.0, topup_remaining_usd=12.0,
|
||||
total_spendable_usd=26.0, plan_bar=_plan_bar(14.0, 6.0),
|
||||
topup_bar=UsageBar(kind="topup", remaining_usd=12.0, total_usd=12.0, spent_usd=0.0),
|
||||
)
|
||||
|
||||
if name == "low":
|
||||
return UsageModel(
|
||||
available=True,
|
||||
status="low",
|
||||
plan_name="Plus",
|
||||
renews_at="2026-07-01",
|
||||
subscription_remaining_usd=3.4,
|
||||
total_spendable_usd=3.4,
|
||||
plan_bar=UsageBar(kind="plan", remaining_usd=3.4, total_usd=20.0, spent_usd=16.6),
|
||||
)
|
||||
|
||||
if name == "depleted":
|
||||
return UsageModel(
|
||||
available=True,
|
||||
status="depleted",
|
||||
plan_name="Plus",
|
||||
renews_at="2026-07-01",
|
||||
subscription_remaining_usd=0.0,
|
||||
total_spendable_usd=0.0,
|
||||
plan_bar=UsageBar(kind="plan", remaining_usd=0.0, total_usd=20.0, spent_usd=20.0),
|
||||
)
|
||||
|
||||
return None
|
||||
),
|
||||
"low": dict(**plus, status="low", subscription_remaining_usd=3.4, total_spendable_usd=3.4, plan_bar=_plan_bar(3.4, 16.6)),
|
||||
"depleted": dict(**plus, status="depleted", subscription_remaining_usd=0.0, total_spendable_usd=0.0, plan_bar=_plan_bar(0.0, 20.0)),
|
||||
}
|
||||
spec = specs.get(name)
|
||||
return UsageModel(**spec) if spec else None
|
||||
|
||||
+181
-328
@@ -1,12 +1,8 @@
|
||||
"""Surface-agnostic core for the Phase 2b Remote Spending screens.
|
||||
"""Surface-agnostic core for the Remote Spending screens (CLI ``_show_billing``, TUI JSON-RPC).
|
||||
|
||||
One fetch/parse per concern, consumed identically by the CLI handler
|
||||
(``cli.py::_show_billing``), the TUI JSON-RPC methods
|
||||
(``tui_gateway/server.py``), and any other surface. Mirrors the proven
|
||||
``agent/account_usage.py::build_credits_view`` pattern: parse the server payload
|
||||
into a frozen dataclass; **fail open** — when not logged in or the portal is
|
||||
unreachable, return a struct with ``logged_in=False`` and let the surface degrade
|
||||
gracefully (never crash).
|
||||
One fetch/parse per concern; the server payload is parsed into frozen dataclasses.
|
||||
**Fail open**: when not logged in or the portal is unreachable, return a struct with
|
||||
``logged_in=False`` and let the surface degrade gracefully (never crash).
|
||||
|
||||
Money discipline: the server emits decimal STRINGS (``"142.5"``, not fixed 2dp).
|
||||
We keep them as :class:`decimal.Decimal` end-to-end and only format for display.
|
||||
@@ -19,55 +15,47 @@ import os
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from decimal import Decimal, InvalidOperation
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Decimal money helpers
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def parse_money(value: Any) -> Optional[Decimal]:
|
||||
"""Parse a server money value (decimal string) into :class:`Decimal`.
|
||||
|
||||
Returns None for missing/invalid input. Never raises. Accepts str/int (and,
|
||||
defensively, float — though the server always sends strings).
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
"""Server money value (decimal string; defensively int/float) -> Decimal, or None. Never raises."""
|
||||
try:
|
||||
# Decimal(str(...)) avoids binary-float artifacts if a float ever sneaks in.
|
||||
return Decimal(str(value).strip())
|
||||
return Decimal(str(value).strip()) if value is not None else None
|
||||
except (InvalidOperation, ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def format_money(value: Optional[Decimal]) -> str:
|
||||
"""Format a Decimal as ``$X`` / ``$X.YY`` for display.
|
||||
def format_money(value: Optional[Decimal], *, grouped: bool = False) -> str:
|
||||
"""``$X`` for whole dollars, ``$X.YY`` (exactly 2dp) otherwise; ``None`` -> ``—``.
|
||||
|
||||
Whole dollars show no decimals; any fractional amount shows exactly 2dp:
|
||||
``Decimal("142.5")`` → ``"$142.50"``, ``Decimal("100")`` → ``"$100"``,
|
||||
``Decimal("0.01")`` → ``"$0.01"``.
|
||||
``grouped=True`` adds thousands separators (mirrors the TUI's ``toLocaleString('en-US')``
|
||||
on plan-catalog rows); the default is intentionally ungrouped across the other surfaces.
|
||||
"""
|
||||
if value is None:
|
||||
return "—"
|
||||
spec = ",f" if grouped else "f"
|
||||
if value == value.to_integral_value():
|
||||
# Whole dollars — no decimal point. format(..., "f") avoids 1E+3 for 1000.
|
||||
return f"${format(value.to_integral_value(), 'f')}"
|
||||
# Fractional — always show 2dp.
|
||||
return f"${format(value.quantize(Decimal('0.01')), 'f')}"
|
||||
# format(..., "f") avoids 1E+3 for 1000.
|
||||
return f"${format(value.to_integral_value(), spec)}"
|
||||
return f"${format(value.quantize(Decimal('0.01')), spec)}"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Parsed sub-structures
|
||||
# =============================================================================
|
||||
def _optional_str(raw: dict, key: str) -> Optional[str]:
|
||||
value = raw.get(key)
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
# resolvedVia → the human answer to "why THIS card?". Keys are the server's card
|
||||
# resolution rungs (NAS card-on-file ladder); absent/unknown rungs render no label
|
||||
# so the display degrades cleanly on servers that don't send resolvedVia yet.
|
||||
def _dict_parser(fn: Callable[[dict], Any]) -> Callable[[Any], Any]:
|
||||
"""Sub-structure parsers accept anything and return None unless the payload is a dict."""
|
||||
return lambda raw: fn(raw) if isinstance(raw, dict) else None
|
||||
|
||||
|
||||
# resolvedVia (server card-on-file ladder rung) → "why THIS card?". Unknown/absent
|
||||
# rungs render no label so older servers degrade cleanly.
|
||||
_CARD_PROVENANCE_LABELS = {
|
||||
"subPin": "the card on your subscription",
|
||||
"customerDefault": "your default card saved on the portal",
|
||||
@@ -79,39 +67,29 @@ _CARD_PROVENANCE_LABELS = {
|
||||
class CardInfo:
|
||||
brand: str
|
||||
last4: str
|
||||
# NAS card-on-file field (post card-resolver): which ladder rung found the
|
||||
# card. Defaults off so pre-resolver payloads parse unchanged.
|
||||
resolved_via: Optional[str] = None
|
||||
resolved_via: Optional[str] = None # ladder rung; None on pre-resolver payloads
|
||||
|
||||
@property
|
||||
def masked(self) -> str:
|
||||
# A Link payment method has no card number (last4 = "") — render the
|
||||
# brand alone, not "Link ····".
|
||||
if not self.last4:
|
||||
return self.brand
|
||||
return f"{self.brand} ····{self.last4}"
|
||||
# A Link payment method has no card number (last4 = "") — brand alone, not "Link ····".
|
||||
return f"{self.brand} ····{self.last4}" if self.last4 else self.brand
|
||||
|
||||
@property
|
||||
def provenance(self) -> Optional[str]:
|
||||
"""Human label for why this card was picked, or None (unknown rung /
|
||||
server too old to say)."""
|
||||
if self.resolved_via is None:
|
||||
return None
|
||||
return _CARD_PROVENANCE_LABELS.get(self.resolved_via)
|
||||
"""Human label for why this card was picked, or None (unknown rung / old server)."""
|
||||
return _CARD_PROVENANCE_LABELS.get(self.resolved_via) if self.resolved_via is not None else None
|
||||
|
||||
@property
|
||||
def display(self) -> str:
|
||||
"""The one-line card display: ``Visa ····4242 — the card on your
|
||||
subscription`` (or just the masked card when provenance is unknown)."""
|
||||
"""``Visa ····4242 — the card on your subscription`` (masked only when provenance unknown)."""
|
||||
label = self.provenance
|
||||
return f"{self.masked} — {label}" if label else self.masked
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PaymentMethodInfo:
|
||||
"""The payment method on file. `kind` is "card", "link", or "unknown"
|
||||
— anything else is normalised to "unknown" at parse time, so consumers
|
||||
only ever see fields that belong to the kind they are looking at."""
|
||||
"""Payment method on file. ``kind`` is "card" | "link" | "unknown" (settled at parse time
|
||||
so consumers only see fields that belong to the kind they are looking at)."""
|
||||
|
||||
kind: str
|
||||
brand: Optional[str] = None
|
||||
@@ -119,8 +97,7 @@ class PaymentMethodInfo:
|
||||
wallet: Optional[str] = None
|
||||
email: Optional[str] = None
|
||||
resolved_via: Optional[str] = None
|
||||
#: What the server called it, when we did not recognise the kind.
|
||||
raw_kind: Optional[str] = None
|
||||
raw_kind: Optional[str] = None # what the server called an unrecognised kind
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -146,13 +123,26 @@ class AutoReload:
|
||||
card: Optional[AutoReloadCard] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BillingState:
|
||||
"""Parsed ``GET /api/billing/state`` — the overview screen's data.
|
||||
class OrgRoleCapability:
|
||||
"""``is_admin`` / ``can_change_plan`` shared by the billing and subscription states."""
|
||||
|
||||
Fail-open: ``logged_in=False`` (and empty fields) when not logged in or the
|
||||
portal is unreachable.
|
||||
"""
|
||||
role: Optional[str]
|
||||
can_change_plan_raw: Optional[bool]
|
||||
|
||||
@property
|
||||
def is_admin(self) -> bool:
|
||||
"""Display only — legacy OWNER/ADMIN check; gate plan-change actions on :attr:`can_change_plan`."""
|
||||
return (self.role or "").upper() in ("OWNER", "ADMIN")
|
||||
|
||||
@property
|
||||
def can_change_plan(self) -> bool:
|
||||
"""Server capability when supplied; otherwise the legacy role fallback."""
|
||||
return self.can_change_plan_raw if self.can_change_plan_raw is not None else self.is_admin
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BillingState(OrgRoleCapability):
|
||||
"""Parsed ``GET /api/billing/state``; fail-open ``logged_in=False`` (empty fields) when unavailable."""
|
||||
|
||||
logged_in: bool
|
||||
org_id: Optional[str] = None
|
||||
@@ -170,156 +160,78 @@ class BillingState:
|
||||
monthly_cap: Optional[MonthlyCap] = None
|
||||
auto_reload: Optional[AutoReload] = None
|
||||
portal_url: Optional[str] = None
|
||||
# When the fetch failed (vs cleanly not-logged-in), the message for the surface.
|
||||
error: Optional[str] = None
|
||||
|
||||
@property
|
||||
def is_admin(self) -> bool:
|
||||
"""Deprecated/display only — a legacy OWNER/ADMIN check.
|
||||
|
||||
NOT a capability check; use :attr:`can_change_plan` for gating billing
|
||||
plan-change actions.
|
||||
"""
|
||||
return (self.role or "").upper() in ("OWNER", "ADMIN")
|
||||
|
||||
@property
|
||||
def can_change_plan(self) -> bool:
|
||||
"""Server capability when supplied; otherwise the legacy role fallback."""
|
||||
if self.can_change_plan_raw is not None:
|
||||
return self.can_change_plan_raw
|
||||
return self.is_admin
|
||||
error: Optional[str] = None # set when the fetch failed (vs cleanly not-logged-in)
|
||||
|
||||
@property
|
||||
def can_charge(self) -> bool:
|
||||
"""True when the UI should offer charge/auto-reload actions.
|
||||
|
||||
Uses the server-granted plan-change capability (``can_change_plan``,
|
||||
which itself falls back to the legacy OWNER/ADMIN role check when the
|
||||
server omits ``canChangePlan``) AND the per-org kill-switch. This lets
|
||||
the server grant charge capability to non-OWNER/ADMIN roles (e.g.
|
||||
FINANCE_ADMIN) via ``canChangePlan``, instead of hard-coding the
|
||||
deprecated 3-role admin check. (The server still enforces; this is
|
||||
just for graying out actions the user can't take.)
|
||||
"""
|
||||
"""Offer charge/auto-reload actions: ``can_change_plan`` (server-grantable, e.g. FINANCE_ADMIN)
|
||||
AND the per-org kill-switch. Display gating only — the server still enforces."""
|
||||
return self.can_change_plan and self.cli_billing_enabled
|
||||
|
||||
|
||||
def _parse_card(raw: Any) -> Optional[CardInfo]:
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
brand = raw.get("brand")
|
||||
last4 = raw.get("last4")
|
||||
@_dict_parser
|
||||
def _parse_card(raw: dict) -> Optional[CardInfo]:
|
||||
brand, last4 = raw.get("brand"), raw.get("last4")
|
||||
if not (isinstance(brand, str) and isinstance(last4, str)):
|
||||
return None
|
||||
# Post-resolver fields — all optional so both payload generations parse.
|
||||
resolved_via = raw.get("resolvedVia")
|
||||
if not isinstance(resolved_via, str):
|
||||
resolved_via = None
|
||||
return CardInfo(brand=brand, last4=last4, resolved_via=resolved_via)
|
||||
return CardInfo(brand=brand, last4=last4, resolved_via=_optional_str(raw, "resolvedVia"))
|
||||
|
||||
|
||||
def _parse_payment_method(raw: Any) -> Optional[PaymentMethodInfo]:
|
||||
if not isinstance(raw, dict):
|
||||
@_dict_parser
|
||||
def _parse_payment_method(raw: dict) -> Optional[PaymentMethodInfo]:
|
||||
if not isinstance(kind := raw.get("kind"), str):
|
||||
return None
|
||||
kind = raw.get("kind")
|
||||
if not isinstance(kind, str):
|
||||
return None
|
||||
|
||||
def _optional_string(key: str) -> Optional[str]:
|
||||
value = raw.get(key)
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
resolved_via = _optional_string("resolvedVia")
|
||||
brand = _optional_string("brand")
|
||||
last4 = _optional_string("last4")
|
||||
# Settle the kind here, the way _parse_card settles a card, so nothing
|
||||
# downstream has to re-check which fields this kind is allowed to have.
|
||||
resolved_via = _optional_str(raw, "resolvedVia")
|
||||
brand = _optional_str(raw, "brand")
|
||||
last4 = _optional_str(raw, "last4")
|
||||
# Settle the kind here (like _parse_card) so nothing downstream re-checks fields.
|
||||
if kind == "card" and brand and last4:
|
||||
return PaymentMethodInfo(
|
||||
kind="card",
|
||||
brand=brand,
|
||||
last4=last4,
|
||||
wallet=_optional_string("wallet"),
|
||||
resolved_via=resolved_via,
|
||||
)
|
||||
return PaymentMethodInfo(kind="card", brand=brand, last4=last4, wallet=_optional_str(raw, "wallet"), resolved_via=resolved_via)
|
||||
if kind == "link":
|
||||
return PaymentMethodInfo(
|
||||
kind="link",
|
||||
email=_optional_string("email"),
|
||||
resolved_via=resolved_via,
|
||||
)
|
||||
return PaymentMethodInfo(
|
||||
kind="unknown", raw_kind=kind, resolved_via=resolved_via
|
||||
)
|
||||
return PaymentMethodInfo(kind="link", email=_optional_str(raw, "email"), resolved_via=resolved_via)
|
||||
return PaymentMethodInfo(kind="unknown", raw_kind=kind, resolved_via=resolved_via)
|
||||
|
||||
|
||||
def _parse_monthly_cap(raw: Any) -> Optional[MonthlyCap]:
|
||||
if not isinstance(raw, dict):
|
||||
@_dict_parser
|
||||
def _parse_monthly_cap(raw: dict) -> MonthlyCap:
|
||||
return MonthlyCap(limit_usd=parse_money(raw.get("limitUsd")), spent_this_month_usd=parse_money(raw.get("spentThisMonthUsd")),
|
||||
is_default_ceiling=bool(raw.get("isDefaultCeiling")))
|
||||
|
||||
|
||||
@_dict_parser
|
||||
def _parse_auto_reload(raw: dict) -> AutoReload:
|
||||
return AutoReload(enabled=bool(raw.get("enabled")), threshold_usd=parse_money(raw.get("thresholdUsd")),
|
||||
reload_to_usd=parse_money(raw.get("reloadToUsd")), card=_parse_auto_reload_card(raw.get("card")))
|
||||
|
||||
|
||||
@_dict_parser
|
||||
def _parse_auto_reload_card(raw: dict) -> Optional[AutoReloadCard]:
|
||||
if (kind := raw.get("kind")) not in ("canonical", "distinct", "none"):
|
||||
return None
|
||||
return MonthlyCap(
|
||||
limit_usd=parse_money(raw.get("limitUsd")),
|
||||
spent_this_month_usd=parse_money(raw.get("spentThisMonthUsd")),
|
||||
is_default_ceiling=bool(raw.get("isDefaultCeiling")),
|
||||
)
|
||||
|
||||
|
||||
def _parse_auto_reload(raw: Any) -> Optional[AutoReload]:
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
return AutoReload(
|
||||
enabled=bool(raw.get("enabled")),
|
||||
threshold_usd=parse_money(raw.get("thresholdUsd")),
|
||||
reload_to_usd=parse_money(raw.get("reloadToUsd")),
|
||||
card=_parse_auto_reload_card(raw.get("card")),
|
||||
)
|
||||
|
||||
|
||||
def _parse_auto_reload_card(raw: Any) -> Optional[AutoReloadCard]:
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
kind = raw.get("kind")
|
||||
if kind not in ("canonical", "distinct", "none"):
|
||||
return None
|
||||
if kind in ("canonical", "none"):
|
||||
if kind != "distinct":
|
||||
return AutoReloadCard(kind=kind)
|
||||
|
||||
payment_method_id = raw.get("paymentMethodId")
|
||||
brand = raw.get("brand")
|
||||
last4 = raw.get("last4")
|
||||
return AutoReloadCard(
|
||||
kind=kind,
|
||||
payment_method_id=payment_method_id if isinstance(payment_method_id, str) else None,
|
||||
brand=brand if isinstance(brand, str) else None,
|
||||
last4=last4 if isinstance(last4, str) else None,
|
||||
)
|
||||
return AutoReloadCard(kind=kind, payment_method_id=_optional_str(raw, "paymentMethodId"),
|
||||
brand=_optional_str(raw, "brand"), last4=_optional_str(raw, "last4"))
|
||||
|
||||
|
||||
def billing_state_from_payload(
|
||||
payload: dict[str, Any], *, portal_url: Optional[str] = None
|
||||
) -> BillingState:
|
||||
def parse_org_fields(payload: dict[str, Any]) -> tuple[dict[str, Any], Optional[bool]]:
|
||||
"""``(org dict or {}, canChangePlan if bool else None)`` — shared by both state parsers."""
|
||||
raw_org, ccp = payload.get("org"), payload.get("canChangePlan")
|
||||
return (raw_org if isinstance(raw_org, dict) else {}), (ccp if isinstance(ccp, bool) else None)
|
||||
|
||||
|
||||
def billing_state_from_payload(payload: dict[str, Any], *, portal_url: Optional[str] = None) -> BillingState:
|
||||
"""Map a raw ``/api/billing/state`` JSON dict into :class:`BillingState`."""
|
||||
raw_org = payload.get("org")
|
||||
org: dict[str, Any] = raw_org if isinstance(raw_org, dict) else {}
|
||||
raw_bounds = payload.get("bounds")
|
||||
bounds: dict[str, Any] = raw_bounds if isinstance(raw_bounds, dict) else {}
|
||||
|
||||
presets: list[Decimal] = []
|
||||
for item in payload.get("chargePresets") or ():
|
||||
parsed = parse_money(item)
|
||||
if parsed is not None:
|
||||
presets.append(parsed)
|
||||
|
||||
org, can_change_plan_raw = parse_org_fields(payload)
|
||||
bounds: dict[str, Any] = payload.get("bounds") if isinstance(payload.get("bounds"), dict) else {}
|
||||
presets = [p for p in map(parse_money, payload.get("chargePresets") or ()) if p is not None]
|
||||
return BillingState(
|
||||
logged_in=True,
|
||||
org_id=org.get("id"),
|
||||
org_slug=org.get("slug"),
|
||||
org_name=org.get("name"),
|
||||
role=org.get("role"),
|
||||
can_change_plan_raw=(
|
||||
payload.get("canChangePlan")
|
||||
if isinstance(payload.get("canChangePlan"), bool)
|
||||
else None
|
||||
),
|
||||
can_change_plan_raw=can_change_plan_raw,
|
||||
balance_usd=parse_money(payload.get("balanceUsd")),
|
||||
cli_billing_enabled=bool(payload.get("cliBillingEnabled")),
|
||||
charge_presets=tuple(presets),
|
||||
@@ -333,153 +245,100 @@ def billing_state_from_payload(
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Fail-open builders (the surface front doors)
|
||||
# =============================================================================
|
||||
def fetch_portal_state(
|
||||
endpoint: str, label: str, *, failed: Callable[..., Any], parse: Callable[[dict, Optional[str]], Any],
|
||||
portal_fallback: Callable[[str], str], timeout: float, log: logging.Logger,
|
||||
):
|
||||
"""Shared fail-open fetch+parse for the billing/subscription overview builders.
|
||||
|
||||
``failed(**kw)`` builds the ``logged_in=False`` struct: bare on auth failure, with ``error``
|
||||
set on a portal/HTTP failure. Portal URL: server ``portalUrl`` (absolutized), else
|
||||
``portal_fallback(portal_base_url)``.
|
||||
"""
|
||||
try:
|
||||
import hermes_cli.nous_billing as nb
|
||||
except Exception:
|
||||
return failed(error="billing client unavailable")
|
||||
try:
|
||||
payload = getattr(nb, endpoint)(timeout=timeout)
|
||||
except nb.BillingAuthError:
|
||||
return failed()
|
||||
except nb.BillingError as exc:
|
||||
log.debug("%s ▸ /state fetch failed (fail-open)", label, exc_info=True)
|
||||
return failed(error=str(exc))
|
||||
except Exception:
|
||||
log.debug("%s ▸ /state unexpected error (fail-open)", label, exc_info=True)
|
||||
return failed(error=f"could not load {label} state")
|
||||
raw_portal = payload.get("portalUrl") if isinstance(payload, dict) else None
|
||||
portal_url = nb._absolutize_portal_url(raw_portal) if raw_portal else None
|
||||
if not portal_url:
|
||||
try:
|
||||
portal_url = portal_fallback(nb.resolve_portal_base_url())
|
||||
except Exception:
|
||||
portal_url = None
|
||||
return parse(payload, portal_url)
|
||||
|
||||
|
||||
def build_billing_state(*, timeout: float = 15.0) -> BillingState:
|
||||
"""Fetch + parse ``/api/billing/state``. Fail-open.
|
||||
|
||||
Returns ``BillingState(logged_in=False)`` when not logged in. On a portal/HTTP
|
||||
failure, returns ``logged_in=False`` with ``error`` set so the surface can show
|
||||
a clear message rather than crashing.
|
||||
|
||||
Dev override: ``HERMES_DEV_BILLING_FIXTURE`` short-circuits to a fixture so the
|
||||
card-on-file / admin / scope states are testable offline (mirrors
|
||||
``HERMES_DEV_CREDITS_FIXTURE`` for the usage model).
|
||||
"""
|
||||
"""Fetch + parse ``/api/billing/state``; fail-open. ``HERMES_DEV_BILLING_FIXTURE`` short-circuits to a fixture."""
|
||||
fixture = _dev_fixture_billing_state()
|
||||
if fixture is not None:
|
||||
return fixture
|
||||
|
||||
try:
|
||||
from hermes_cli.nous_billing import (
|
||||
BillingAuthError,
|
||||
BillingError,
|
||||
_absolutize_portal_url,
|
||||
get_billing_state,
|
||||
resolve_portal_base_url,
|
||||
)
|
||||
except Exception:
|
||||
return BillingState(logged_in=False, error="billing client unavailable")
|
||||
|
||||
try:
|
||||
payload = get_billing_state(timeout=timeout)
|
||||
except BillingAuthError:
|
||||
return BillingState(logged_in=False)
|
||||
except BillingError as exc:
|
||||
logger.debug("billing ▸ /state fetch failed (fail-open)", exc_info=True)
|
||||
return BillingState(logged_in=False, error=str(exc))
|
||||
except Exception:
|
||||
logger.debug("billing ▸ /state unexpected error (fail-open)", exc_info=True)
|
||||
return BillingState(logged_in=False, error="could not load billing state")
|
||||
|
||||
# Prefer a server-supplied portalUrl if present (resolved to absolute in case
|
||||
# it's relative); else build the standard one.
|
||||
raw_portal = payload.get("portalUrl") if isinstance(payload, dict) else None
|
||||
portal_url = _absolutize_portal_url(raw_portal) if raw_portal else None
|
||||
if not portal_url:
|
||||
try:
|
||||
portal_url = _fallback_portal_url(resolve_portal_base_url())
|
||||
except Exception:
|
||||
portal_url = None
|
||||
|
||||
return billing_state_from_payload(payload, portal_url=portal_url)
|
||||
return fetch_portal_state(
|
||||
"get_billing_state", "billing", timeout=timeout, log=logger,
|
||||
failed=lambda **kw: BillingState(logged_in=False, **kw),
|
||||
parse=lambda payload, portal_url: billing_state_from_payload(payload, portal_url=portal_url),
|
||||
portal_fallback=lambda base: f"{base.rstrip('/')}/billing?topup=open",
|
||||
)
|
||||
|
||||
|
||||
def _fallback_portal_url(base: str) -> str:
|
||||
"""Standard billing deep-link when the server omits ``portalUrl``."""
|
||||
return f"{base.rstrip('/')}/billing?topup=open"
|
||||
# ── Dev fixtures (env-var driven, no live portal) ────────────────────────────
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Dev fixtures (throwaway scaffolding — env-var driven, no live portal)
|
||||
# =============================================================================
|
||||
_FIXTURE_ALIASES = {
|
||||
"logged_out": "logged-out", "loggedout": "logged-out", "card_sub": "card-sub", "card_autoreload": "card-autoreload",
|
||||
"autoreload": "card-autoreload", "not-admin": "notadmin", "member": "notadmin", "billing_off": "billing-off", "off": "billing-off",
|
||||
}
|
||||
|
||||
|
||||
def _dev_fixture_billing_state() -> Optional[BillingState]:
|
||||
"""Map ``HERMES_DEV_BILLING_FIXTURE`` to a :class:`BillingState` for offline UX.
|
||||
"""``HERMES_DEV_BILLING_FIXTURE`` -> :class:`BillingState` for offline UX; None when unset.
|
||||
|
||||
Recognized names::
|
||||
|
||||
nocard logged in · billing on · admin · NO card on file
|
||||
card card on file · auto-reload off
|
||||
card-autoreload card on file · auto-reload on
|
||||
notadmin logged in · MEMBER role (billing actions disabled)
|
||||
billing-off logged in · admin · per-org kill-switch OFF
|
||||
logged-out not logged in
|
||||
|
||||
Returns ``None`` when the env var is unset (the real portal path runs).
|
||||
Mirrors ``HERMES_DEV_CREDITS_FIXTURE``; the usage *bar* still comes from
|
||||
``HERMES_DEV_CREDITS_FIXTURE`` (set both to pair a bar with a billing state).
|
||||
Names: nocard · card · card-sub · card-autoreload · notadmin · billing-off · logged-out; an
|
||||
unknown name yields logged-out with ``error`` so the misconfiguration is visible.
|
||||
"""
|
||||
name = (os.getenv("HERMES_DEV_BILLING_FIXTURE") or "").strip().lower()
|
||||
if not name:
|
||||
return None
|
||||
name = _FIXTURE_ALIASES.get(name, name)
|
||||
if name == "logged-out":
|
||||
return BillingState(logged_in=False)
|
||||
|
||||
# Shared fixture portal host (matches subscription_view._DEV_FIXTURE_PORTAL —
|
||||
# prod host, not staging; the ?topup=open suffix is the /topup deep-link).
|
||||
portal = "https://portal.nousresearch.com/billing?topup=open"
|
||||
# Prod portal host (matches subscription_view._DEV_FIXTURE_PORTAL) + the /topup deep-link suffix.
|
||||
common: dict[str, Any] = dict(
|
||||
org_id="org_acme",
|
||||
org_slug="acme",
|
||||
org_name="Acme Inc",
|
||||
role="OWNER",
|
||||
balance_usd=Decimal("3.40"),
|
||||
cli_billing_enabled=True,
|
||||
charge_presets=(Decimal("10"), Decimal("25"), Decimal("50")),
|
||||
min_usd=Decimal("5"),
|
||||
max_usd=Decimal("500"),
|
||||
portal_url=portal,
|
||||
logged_in=True, org_id="org_acme", org_slug="acme", org_name="Acme Inc", role="OWNER",
|
||||
balance_usd=Decimal("3.40"), cli_billing_enabled=True, min_usd=Decimal("5"), max_usd=Decimal("500"),
|
||||
charge_presets=(Decimal("10"), Decimal("25"), Decimal("50")), portal_url="https://portal.nousresearch.com/billing?topup=open",
|
||||
)
|
||||
card = CardInfo(brand="Visa", last4="4242")
|
||||
autoreload_on = AutoReload(enabled=True, threshold_usd=Decimal("5"), reload_to_usd=Decimal("25"))
|
||||
|
||||
if name in ("logged-out", "logged_out", "loggedout"):
|
||||
return BillingState(logged_in=False)
|
||||
if name == "nocard":
|
||||
return BillingState(logged_in=True, card=None, **common)
|
||||
if name == "card":
|
||||
return BillingState(logged_in=True, card=card, **common)
|
||||
if name in ("card-sub", "card_sub"):
|
||||
# Post-resolver: the card came from the subscription (provenance label).
|
||||
_sub_card = CardInfo(brand="Visa", last4="4242", resolved_via="subPin")
|
||||
return BillingState(logged_in=True, card=_sub_card, **common)
|
||||
if name in ("card-autoreload", "card_autoreload", "autoreload"):
|
||||
return BillingState(logged_in=True, card=card, auto_reload=autoreload_on, **common)
|
||||
if name in ("notadmin", "not-admin", "member"):
|
||||
opts = {**common, "role": "MEMBER"}
|
||||
return BillingState(logged_in=True, card=card, **opts)
|
||||
if name in ("billing-off", "billing_off", "off"):
|
||||
opts = {**common, "cli_billing_enabled": False}
|
||||
return BillingState(logged_in=True, card=None, **opts)
|
||||
|
||||
# Unknown name → logged-out so the misconfiguration is visible.
|
||||
return BillingState(logged_in=False, error=f"unknown HERMES_DEV_BILLING_FIXTURE: {name}")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Idempotency
|
||||
# =============================================================================
|
||||
overrides: dict[str, dict[str, Any]] = {
|
||||
"nocard": dict(card=None),
|
||||
"card": dict(card=card),
|
||||
"card-sub": dict(card=CardInfo(brand="Visa", last4="4242", resolved_via="subPin")),
|
||||
"card-autoreload": dict(card=card, auto_reload=AutoReload(enabled=True, threshold_usd=Decimal("5"), reload_to_usd=Decimal("25"))),
|
||||
"notadmin": dict(card=card, role="MEMBER"),
|
||||
"billing-off": dict(card=None, cli_billing_enabled=False),
|
||||
}
|
||||
if name not in overrides:
|
||||
return BillingState(logged_in=False, error=f"unknown HERMES_DEV_BILLING_FIXTURE: {name}")
|
||||
return BillingState(**{**common, **overrides[name]})
|
||||
|
||||
|
||||
def new_idempotency_key() -> str:
|
||||
"""Fresh UUID for a user-confirmed purchase (reuse on retry of the SAME buy).
|
||||
|
||||
The ``Idempotency-Key`` header is mandatory on ``POST /charge``; generate one
|
||||
per confirmed purchase and reuse it across retries so a double-submit collapses
|
||||
to a single charge. Never reuse a key across different amounts (the server
|
||||
returns 409 idempotency_conflict).
|
||||
"""
|
||||
"""Fresh ``Idempotency-Key`` for ``POST /charge``: reuse across retries of the SAME buy so a
|
||||
double-submit collapses to one charge; never across amounts (server 409 idempotency_conflict)."""
|
||||
return str(uuid.uuid4())
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Amount validation (Screen 3 custom input)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AmountValidation:
|
||||
ok: bool
|
||||
@@ -487,25 +346,19 @@ class AmountValidation:
|
||||
error: Optional[str] = None
|
||||
|
||||
|
||||
def validate_charge_amount(
|
||||
raw: str, *, min_usd: Optional[Decimal], max_usd: Optional[Decimal]
|
||||
) -> AmountValidation:
|
||||
"""Validate a custom charge amount against bounds + 2dp (multipleOf 0.01).
|
||||
|
||||
Mirrors the server's accept/reject so the UI can give instant feedback rather
|
||||
than round-tripping a sure-to-fail charge. The server is still authoritative.
|
||||
"""
|
||||
cleaned = (raw or "").strip().lstrip("$").strip()
|
||||
amount = parse_money(cleaned)
|
||||
def validate_charge_amount(raw: str, *, min_usd: Optional[Decimal], max_usd: Optional[Decimal]) -> AmountValidation:
|
||||
"""Mirror the server's accept/reject (bounds + multipleOf 0.01) for instant UI feedback; server is authoritative."""
|
||||
amount = parse_money((raw or "").strip().lstrip("$").strip())
|
||||
if amount is None:
|
||||
return AmountValidation(ok=False, error="Enter a dollar amount, e.g. 100")
|
||||
if amount <= 0:
|
||||
return AmountValidation(ok=False, error="Amount must be greater than $0")
|
||||
# multipleOf 0.01 — reject sub-cent precision.
|
||||
if amount != amount.quantize(Decimal("0.01")):
|
||||
return AmountValidation(ok=False, error="Amount can't be smaller than a cent")
|
||||
if min_usd is not None and amount < min_usd:
|
||||
return AmountValidation(ok=False, error=f"Minimum is {format_money(min_usd)}")
|
||||
if max_usd is not None and amount > max_usd:
|
||||
return AmountValidation(ok=False, error=f"Maximum is {format_money(max_usd)}")
|
||||
return AmountValidation(ok=True, amount=amount)
|
||||
error = "Enter a dollar amount, e.g. 100"
|
||||
elif amount <= 0:
|
||||
error = "Amount must be greater than $0"
|
||||
elif amount != amount.quantize(Decimal("0.01")):
|
||||
error = "Amount can't be smaller than a cent"
|
||||
elif min_usd is not None and amount < min_usd:
|
||||
error = f"Minimum is {format_money(min_usd)}"
|
||||
elif max_usd is not None and amount > max_usd:
|
||||
error = f"Maximum is {format_money(max_usd)}"
|
||||
else:
|
||||
return AmountValidation(ok=True, amount=amount)
|
||||
return AmountValidation(ok=False, error=error)
|
||||
|
||||
+33
-79
@@ -1,55 +1,27 @@
|
||||
"""Bounded reads of HTTP error response bodies.
|
||||
|
||||
When a provider returns a non-OK status on a *streaming* request, Hermes reads
|
||||
the response body to build a useful diagnostic error. A bare ``response.read()``
|
||||
on a streaming httpx response is unbounded in two dangerous ways:
|
||||
|
||||
1. A server can declare (or stream) an arbitrarily large body, so the read can
|
||||
balloon memory.
|
||||
2. A server can open the body and then stall forever (no ``Content-Length``,
|
||||
no further bytes), so the read hangs the agent indefinitely.
|
||||
|
||||
Both are realistic against a misbehaving proxy, a hijacked endpoint, or a
|
||||
provider having a bad day. The diagnostic body is only ever shown to the user
|
||||
truncated to a few hundred characters, so reading megabytes — or blocking
|
||||
forever — buys nothing.
|
||||
|
||||
``read_streaming_error_body`` bounds the read to a byte cap and enforces a
|
||||
hard wall-clock deadline, returning the decoded text snippet. Callers pass the
|
||||
returned text into their existing error builders instead of touching
|
||||
``response.text`` (which would be unbounded / would raise after a partial
|
||||
stream read).
|
||||
|
||||
A subtlety the implementation must respect: ``httpx``'s ``iter_bytes()`` blocks
|
||||
*inside* the C/socket read while waiting for the next chunk. A wall-clock check
|
||||
placed only between yielded chunks cannot interrupt a server that opens the
|
||||
body and then stalls mid-chunk — control never returns to Python until httpx's
|
||||
own (often 30s+) read timeout fires. To guarantee a bounded stop regardless of
|
||||
socket behavior, the read runs on a daemon worker thread and the caller waits
|
||||
on it with a hard deadline; on timeout we close the response (which unblocks /
|
||||
cancels the read) and return whatever partial bytes were collected.
|
||||
|
||||
Ported and adapted from openclaw/openclaw#95108 ("bound Anthropic error
|
||||
streams"), generalized to cover Hermes's three streaming error-body sites
|
||||
(native Gemini, Gemini Cloud Code, Antigravity Cloud Code).
|
||||
On a non-OK *streaming* response Hermes reads the body for a diagnostic (only ever shown truncated to
|
||||
a few hundred chars). A bare ``response.read()`` is unbounded two ways: arbitrarily large body
|
||||
(memory) or a body that stalls forever (hang). ``read_streaming_error_body`` caps bytes and enforces a
|
||||
hard wall-clock deadline; callers use the returned text instead of ``response.text`` (unbounded /
|
||||
raises after a partial stream read). ``httpx.iter_bytes()`` blocks *inside* the socket read, so the
|
||||
read runs on a daemon thread; on timeout we close the response (unblocking the read) and return the
|
||||
partial bytes. Used by the streaming error-body sites: native Gemini, Gemini Cloud Code, Antigravity.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import List, Optional
|
||||
from typing import List
|
||||
|
||||
import httpx
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Defaults chosen to comfortably hold any real provider error envelope (Google
|
||||
# RPC error JSON, Anthropic error JSON) while rejecting pathological bodies.
|
||||
# Comfortably holds any real provider error envelope while rejecting pathological bodies.
|
||||
DEFAULT_ERROR_BODY_MAX_BYTES = 64 * 1024
|
||||
# Hard wall-clock deadline for the whole bounded read. A streaming error body
|
||||
# that does not finish within this window is abandoned and the connection is
|
||||
# closed; we keep whatever partial bytes arrived.
|
||||
# Hard deadline for the whole read; past it the connection is closed and the partial bytes are kept.
|
||||
DEFAULT_ERROR_BODY_TIMEOUT_S = 10.0
|
||||
|
||||
|
||||
@@ -59,17 +31,11 @@ def read_streaming_error_body(
|
||||
max_bytes: int = DEFAULT_ERROR_BODY_MAX_BYTES,
|
||||
timeout_s: float = DEFAULT_ERROR_BODY_TIMEOUT_S,
|
||||
) -> str:
|
||||
"""Read a non-OK streaming response body with a byte cap and a hard deadline.
|
||||
"""Read a non-OK streaming body with a byte cap and a hard deadline.
|
||||
|
||||
Returns the decoded body text (UTF-8, errors replaced), truncated to
|
||||
``max_bytes``. Never raises: any transport error, stall, or oversize
|
||||
condition is swallowed and the best-effort partial text (or an empty
|
||||
string) is returned, because this runs on the error path and must not
|
||||
mask the original HTTP failure with a read error.
|
||||
|
||||
The byte cap protects against huge bodies; the wall-clock deadline (enforced
|
||||
via a worker thread so it can interrupt a socket read that stalls mid-chunk)
|
||||
protects against bodies that open and then hang.
|
||||
Returns UTF-8 text (errors replaced) truncated to ``max_bytes``. Never raises: transport errors,
|
||||
stalls and oversize bodies yield best-effort partial text (or ""), so a read error can't mask the
|
||||
original failure.
|
||||
"""
|
||||
chunks: List[bytes] = []
|
||||
state = {"truncated": False}
|
||||
@@ -82,12 +48,9 @@ def read_streaming_error_body(
|
||||
if not chunk:
|
||||
continue
|
||||
remaining = max_bytes - total
|
||||
if remaining <= 0:
|
||||
state["truncated"] = True
|
||||
break
|
||||
if len(chunk) > remaining:
|
||||
chunks.append(chunk[:remaining])
|
||||
total += remaining
|
||||
if remaining > 0:
|
||||
chunks.append(chunk[:remaining])
|
||||
state["truncated"] = True
|
||||
break
|
||||
chunks.append(chunk)
|
||||
@@ -97,40 +60,30 @@ def read_streaming_error_body(
|
||||
finally:
|
||||
done.set()
|
||||
|
||||
worker = threading.Thread(
|
||||
target=_drain, name="bounded-error-body-read", daemon=True
|
||||
)
|
||||
worker.start()
|
||||
finished = done.wait(timeout=timeout_s)
|
||||
|
||||
if not finished:
|
||||
threading.Thread(target=_drain, name="bounded-error-body-read", daemon=True).start()
|
||||
if not done.wait(timeout=timeout_s):
|
||||
logger.debug(
|
||||
"bounded error-body read: hard timeout after %.1fs (%d bytes so far)",
|
||||
timeout_s,
|
||||
sum(len(c) for c in chunks),
|
||||
timeout_s, sum(len(c) for c in chunks),
|
||||
)
|
||||
# Closing the response cancels the in-flight socket read, letting the
|
||||
# worker thread unwind. We do not join (it is a daemon and may be
|
||||
# blocked in C); the partial `chunks` collected so far are returned.
|
||||
_safe_close(response)
|
||||
else:
|
||||
_safe_close(response)
|
||||
|
||||
if state["truncated"]:
|
||||
logger.debug(
|
||||
"bounded error-body read: capped at %d bytes (max=%d)",
|
||||
sum(len(c) for c in chunks),
|
||||
max_bytes,
|
||||
)
|
||||
return b"".join(chunks).decode("utf-8", errors="replace")
|
||||
|
||||
|
||||
def _safe_close(response: httpx.Response) -> None:
|
||||
# Closing cancels any in-flight socket read so the worker unwinds. No join (daemon, may be blocked in C).
|
||||
try:
|
||||
response.close()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
if state["truncated"]:
|
||||
logger.debug(
|
||||
"bounded error-body read: capped at %d bytes (max=%d)", sum(len(c) for c in chunks), max_bytes,
|
||||
)
|
||||
return b"".join(chunks).decode("utf-8", errors="replace")
|
||||
|
||||
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
from typing import Optional # noqa: F401,E402
|
||||
|
||||
def read_error_body_or_default(
|
||||
response: httpx.Response,
|
||||
@@ -146,3 +99,4 @@ def read_error_body_or_default(
|
||||
response, max_bytes=max_bytes, timeout_s=timeout_s
|
||||
)
|
||||
return text or None
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
+31
-142
@@ -1,26 +1,12 @@
|
||||
"""
|
||||
Browser Provider ABC
|
||||
====================
|
||||
"""Browser Provider ABC: pluggable cloud browser backends (Browserbase, Browser Use, Firecrawl, …).
|
||||
|
||||
Defines the pluggable-backend interface for cloud browser providers
|
||||
(Browserbase, Browser Use, Firecrawl, …). Providers register instances via
|
||||
:meth:`PluginContext.register_browser_provider`; the active one (selected via
|
||||
``browser.cloud_provider`` in ``config.yaml``) services every cloud-mode
|
||||
``browser_*`` tool call.
|
||||
Providers register via :meth:`PluginContext.register_browser_provider`; the active one (selected by
|
||||
``browser.cloud_provider``) services every cloud-mode ``browser_*`` tool call. They live in
|
||||
``<repo>/plugins/browser/<name>/`` (built-in) or ``~/.hermes/plugins/browser/<name>/`` (user).
|
||||
|
||||
Providers live in ``<repo>/plugins/browser/<name>/`` (built-in, auto-loaded as
|
||||
``kind: backend``) or ``~/.hermes/plugins/browser/<name>/`` (user, opt-in via
|
||||
``plugins.enabled``).
|
||||
|
||||
This ABC mirrors :class:`agent.web_search_provider.WebSearchProvider` (PR
|
||||
#25182) — same shape, same registration flow, same picker integration. The
|
||||
legacy in-tree ``tools.browser_providers.base.CloudBrowserProvider`` ABC was
|
||||
deleted in PR #25214 (this work) along with the per-vendor inline modules in
|
||||
``tools/browser_providers/``; the lifecycle contract documented below is
|
||||
preserved bit-for-bit so the tool wrapper (:mod:`tools.browser_tool`) does
|
||||
not have to translate.
|
||||
|
||||
Session metadata contract (preserved from the legacy ``CloudBrowserProvider``)::
|
||||
Session metadata contract (legacy ``CloudBrowserProvider`` shape; ``tools.browser_tool`` needs no
|
||||
translation). ``bb_session_id`` is a legacy key name kept verbatim — it holds the provider's session ID
|
||||
regardless of provider::
|
||||
|
||||
{
|
||||
"session_name": str, # unique name for agent-browser --session
|
||||
@@ -30,148 +16,51 @@ Session metadata contract (preserved from the legacy ``CloudBrowserProvider``)::
|
||||
"features": dict, # feature flags that were enabled
|
||||
"external_call_id": str, # optional, managed-gateway billing key
|
||||
}
|
||||
|
||||
``bb_session_id`` is a legacy key name kept verbatim for backward compat with
|
||||
:mod:`tools.browser_tool` — it holds the provider's session ID regardless of
|
||||
which provider is in use.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import abc
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Dict
|
||||
|
||||
from agent.provider_base import ProviderBase
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ABC
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BrowserProvider(abc.ABC):
|
||||
class BrowserProvider(ProviderBase):
|
||||
"""Abstract base class for a cloud browser backend.
|
||||
|
||||
Subclasses must implement :meth:`name`, :meth:`is_available`, and the
|
||||
three lifecycle methods: :meth:`create_session`, :meth:`close_session`,
|
||||
:meth:`emergency_cleanup`.
|
||||
|
||||
The lifecycle shape preserves the legacy ``CloudBrowserProvider`` contract
|
||||
bit-for-bit so the dispatcher in :mod:`tools.browser_tool` is a pure
|
||||
registry lookup — no per-provider conditionals, no shape translation.
|
||||
Subclasses implement :attr:`name` (the ``browser.cloud_provider`` value), :meth:`is_available`, and
|
||||
the lifecycle trio :meth:`create_session` / :meth:`close_session` / :meth:`emergency_cleanup`.
|
||||
``get_setup_schema`` may add ``"post_setup"`` (e.g. ``"agent_browser"``) to trigger the install hook.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used in the ``browser.cloud_provider``
|
||||
config key.
|
||||
|
||||
Lowercase, hyphens permitted to preserve existing user-visible names.
|
||||
Examples: ``browserbase``, ``browser-use``, ``firecrawl``.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``. Defaults to ``name``."""
|
||||
return self.name
|
||||
|
||||
@abc.abstractmethod
|
||||
def is_available(self) -> bool:
|
||||
"""Return True when this provider can service calls.
|
||||
|
||||
Typically a cheap check (env var present, managed-gateway token
|
||||
readable, optional Python dep importable). Must NOT make network
|
||||
calls — this runs at tool-registration time and on every
|
||||
``hermes tools`` paint.
|
||||
|
||||
Mirrors the legacy ``CloudBrowserProvider.is_configured()`` method;
|
||||
renamed for parity with :class:`agent.web_search_provider.WebSearchProvider`.
|
||||
"""
|
||||
"""True when this provider can service calls. Cheap check only (env var, token readable, dep
|
||||
importable) — must NOT make network calls; runs at tool-registration time and on every
|
||||
``hermes tools`` paint."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def create_session(self, task_id: str) -> Dict[str, object]:
|
||||
"""Create a cloud browser session and return session metadata.
|
||||
|
||||
Must return a dict with at least::
|
||||
|
||||
{
|
||||
"session_name": str, # unique name for agent-browser --session
|
||||
"bb_session_id": str, # provider session ID (for close/cleanup)
|
||||
"cdp_url": str, # CDP websocket URL
|
||||
"expires_at": str, # optional provider-authoritative ISO timestamp
|
||||
"features": dict, # feature flags that were enabled
|
||||
}
|
||||
|
||||
``bb_session_id`` is a legacy key name kept for backward compat with
|
||||
the rest of :mod:`tools.browser_tool` — it holds the provider's
|
||||
session ID regardless of which provider is in use.
|
||||
|
||||
May raise ``ValueError`` (missing credentials) or ``RuntimeError``
|
||||
(network / API failure); the dispatcher surfaces these to the user.
|
||||
"""
|
||||
"""Create a cloud browser session and return the metadata dict from the module docstring.
|
||||
May raise ``ValueError`` (missing credentials) or ``RuntimeError`` (network / API failure);
|
||||
the dispatcher surfaces these to the user."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def close_session(self, session_id: str) -> bool:
|
||||
"""Release / terminate a cloud session by its provider session ID.
|
||||
|
||||
Returns True on success, False on failure. Should not raise — log and
|
||||
return False on any exception so the dispatcher's cleanup loop keeps
|
||||
moving across sessions.
|
||||
"""
|
||||
"""Release a cloud session by provider session ID. Returns True on success, False on failure;
|
||||
should not raise (log and return False so the dispatcher's cleanup loop keeps moving)."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def emergency_cleanup(self, session_id: str) -> None:
|
||||
"""Best-effort session teardown during process exit.
|
||||
"""Best-effort teardown from atexit / signal handlers. Must tolerate missing credentials and
|
||||
network errors; must not raise."""
|
||||
|
||||
Called from atexit / signal handlers. Must tolerate missing
|
||||
credentials, network errors, etc. — log and move on. Must not raise.
|
||||
"""
|
||||
|
||||
def get_setup_schema(self) -> Optional[Dict[str, Any]]:
|
||||
"""Return provider metadata for the ``hermes tools`` picker.
|
||||
|
||||
Used by :mod:`hermes_cli.tools_config` to inject this provider as a
|
||||
row in the Browser Automation picker. Shape mirrors the existing
|
||||
hardcoded entries in ``TOOL_CATEGORIES["browser"]``::
|
||||
|
||||
{
|
||||
"name": "Browserbase",
|
||||
"badge": "paid",
|
||||
"tag": "Cloud browser with stealth and proxies",
|
||||
"env_vars": [
|
||||
{"key": "BROWSERBASE_API_KEY",
|
||||
"prompt": "Browserbase API key",
|
||||
"url": "https://browserbase.com"},
|
||||
],
|
||||
"post_setup": "agent_browser",
|
||||
}
|
||||
|
||||
Default: minimal entry derived from :attr:`display_name`. Override to
|
||||
expose API key prompts, badges, managed-Nous gating, and the
|
||||
``post_setup`` install hook.
|
||||
"""
|
||||
return {
|
||||
"name": self.display_name,
|
||||
"badge": "",
|
||||
"tag": "",
|
||||
"env_vars": [],
|
||||
}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Backward-compat shims for the legacy CloudBrowserProvider API
|
||||
# ------------------------------------------------------------------
|
||||
#
|
||||
# The pre-PR-#25214 ABC exposed ``is_configured()`` and ``provider_name()``;
|
||||
# ``tools.browser_tool`` has ~6 callers that still use those names. Rather
|
||||
# than churn every callsite (and break out-of-tree downstream code that
|
||||
# subclassed CloudBrowserProvider), we expose the old names as thin
|
||||
# delegations to the new API. Subclasses MUST implement :meth:`is_available`
|
||||
# and :attr:`name`; they may override ``is_configured`` / ``provider_name``
|
||||
# for compatibility with the legacy ABC but it is not required.
|
||||
|
||||
def is_configured(self) -> bool:
|
||||
"""Backward-compat alias for :meth:`is_available`."""
|
||||
return self.is_available()
|
||||
|
||||
def provider_name(self) -> str:
|
||||
"""Backward-compat alias returning :attr:`display_name`."""
|
||||
return self.display_name
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
from typing import Any # noqa: F401,E402
|
||||
from typing import Optional # noqa: F401,E402
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
+53
-220
@@ -1,253 +1,86 @@
|
||||
"""
|
||||
Browser Provider Registry
|
||||
=========================
|
||||
"""Browser provider registry: cloud browser backends registered by plugins via
|
||||
:meth:`PluginContext.register_browser_provider`, consumed by ``tools.browser_tool_cloud._get_cloud_provider``.
|
||||
|
||||
Central map of registered cloud browser providers. Populated by plugins at
|
||||
import-time via :meth:`PluginContext.register_browser_provider`; consumed by
|
||||
:func:`tools.browser_tool._get_cloud_provider` to route each cloud-mode
|
||||
``browser_*`` tool call to the active backend.
|
||||
|
||||
Active selection
|
||||
----------------
|
||||
The active provider is chosen by configuration with this precedence:
|
||||
|
||||
1. ``browser.cloud_provider`` in ``config.yaml`` (explicit override).
|
||||
2. Legacy preference order — ``browser-use`` → ``browserbase`` — filtered by
|
||||
availability. Matches the historic auto-detect order in
|
||||
:func:`tools.browser_tool._get_cloud_provider` (Browser Use checked first
|
||||
because it covers both the managed Nous gateway and direct API key path;
|
||||
Browserbase as the older direct-credentials fallback). ``firecrawl`` is
|
||||
intentionally NOT in the legacy walk — users only get Firecrawl as a
|
||||
cloud browser when they explicitly set ``browser.cloud_provider:
|
||||
firecrawl``, matching pre-migration behaviour where Firecrawl was never
|
||||
auto-selected.
|
||||
3. Otherwise ``None`` — the dispatcher falls back to local browser mode.
|
||||
|
||||
The explicit-config branch (rule 1) intentionally ignores ``is_available()``
|
||||
so the dispatcher surfaces a typed "X_API_KEY is not set" error to the user
|
||||
instead of silently switching backends. Matches the legacy
|
||||
:func:`tools.browser_tool._get_cloud_provider` behaviour for configured names.
|
||||
|
||||
Note: there is no "capability" split here (unlike the web subsystem, which
|
||||
has search/extract/crawl). Every browser provider implements the full
|
||||
:class:`agent.browser_provider.BrowserProvider` lifecycle; the registry's
|
||||
job is purely selection, not capability routing.
|
||||
Active-provider precedence (see :func:`_resolve`): ``browser.cloud_provider`` in config.yaml wins
|
||||
regardless of ``is_available()`` (so the dispatcher surfaces a typed "X_API_KEY is not set" error
|
||||
instead of silently switching); else the legacy auto-detect walk ``browser-use`` → ``browserbase``
|
||||
filtered by availability; else ``None`` (local browser mode). There is no capability split here —
|
||||
every provider implements the full :class:`agent.browser_provider.BrowserProvider` lifecycle.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Optional
|
||||
|
||||
from agent.browser_provider import BrowserProvider
|
||||
from hermes_constants import hermes_home_key
|
||||
from agent.provider_registry import ProviderRegistry, is_available_safe
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_providers: Dict[str, BrowserProvider] = {}
|
||||
_scoped_providers: Dict[str, Dict[str, BrowserProvider]] = {}
|
||||
_generation = 0
|
||||
_scoped_generations: Dict[str, int] = {}
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def register_provider(provider: BrowserProvider, *, scope: Optional[str] = None) -> None:
|
||||
"""Register a cloud browser provider.
|
||||
|
||||
Re-registration (same ``name``) overwrites the previous entry and logs
|
||||
a debug message — makes hot-reload scenarios (tests, dev loops) behave
|
||||
predictably.
|
||||
"""
|
||||
if not isinstance(provider, BrowserProvider):
|
||||
raise TypeError(
|
||||
f"register_provider() expects a BrowserProvider instance, "
|
||||
f"got {type(provider).__name__}"
|
||||
)
|
||||
raw_name = provider.name
|
||||
if not isinstance(raw_name, str) or not raw_name.strip():
|
||||
raise ValueError("Browser provider .name must be a non-empty string")
|
||||
name = raw_name.strip()
|
||||
global _generation
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
existing = target.get(name)
|
||||
target[name] = provider
|
||||
if scope is None:
|
||||
_generation += 1
|
||||
else:
|
||||
_scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1
|
||||
if existing is not None:
|
||||
logger.debug(
|
||||
"Browser provider '%s' re-registered (was %r)",
|
||||
name, type(existing).__name__,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"Registered browser provider '%s' (%s)",
|
||||
name, type(provider).__name__,
|
||||
)
|
||||
|
||||
|
||||
def list_providers(*, scope: Optional[str] = None) -> List[BrowserProvider]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
with _lock:
|
||||
merged = dict(_providers)
|
||||
merged.update(_scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
items = list(merged.values())
|
||||
return sorted(items, key=lambda p: p.name)
|
||||
|
||||
|
||||
def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[BrowserProvider]:
|
||||
"""Return the provider registered under *name*, or None."""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
with _lock:
|
||||
key = name.strip()
|
||||
return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key)
|
||||
|
||||
|
||||
def snapshot_registration(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[BrowserProvider]:
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.get(scope, {})
|
||||
return target.get(name.strip())
|
||||
|
||||
|
||||
def registry_generation(*, scope: Optional[str] = None) -> tuple[int, int]:
|
||||
"""Return a cache fingerprint for the global base and one profile."""
|
||||
active_scope = scope or hermes_home_key()
|
||||
with _lock:
|
||||
return _generation, _scoped_generations.get(active_scope, 0)
|
||||
|
||||
|
||||
def restore_registration(
|
||||
name: str,
|
||||
current: BrowserProvider,
|
||||
previous: Optional[BrowserProvider],
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Restore a plugin registration only when *current* is still installed."""
|
||||
key = name.strip()
|
||||
global _generation
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
if scope is None:
|
||||
_generation += 1
|
||||
else:
|
||||
_scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1
|
||||
if not target:
|
||||
_scoped_providers.pop(scope, None)
|
||||
return True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Active-provider resolution
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
# Legacy auto-detect order — used when no ``browser.cloud_provider`` is set.
|
||||
# Matches the pre-migration walk in :func:`tools.browser_tool._get_cloud_provider`.
|
||||
# Firecrawl is intentionally absent so users with ``FIRECRAWL_API_KEY`` set
|
||||
# for web-extract don't get silently routed to a paid cloud browser. See
|
||||
# :func:`_resolve` for the full rationale.
|
||||
_LEGACY_PREFERENCE = (
|
||||
"browser-use",
|
||||
"browserbase",
|
||||
_registry: ProviderRegistry[BrowserProvider] = ProviderRegistry(
|
||||
label="Browser", provider_cls=BrowserProvider, logger=logger,
|
||||
)
|
||||
_registry.export(globals())
|
||||
|
||||
|
||||
# Auto-detect order when ``browser.cloud_provider`` is unset (historic order: Browser Use first because
|
||||
# it covers both the managed Nous gateway and the direct API key path; Browserbase as the older
|
||||
# direct-credentials fallback). Firecrawl is deliberately absent — see :func:`_resolve`.
|
||||
_LEGACY_PREFERENCE = ("browser-use", "browserbase")
|
||||
|
||||
|
||||
def _resolve(configured: Optional[str]) -> Optional[BrowserProvider]:
|
||||
"""Resolve the active browser provider.
|
||||
"""Resolve the active browser provider (rules in the module docstring).
|
||||
|
||||
Resolution rules (in order):
|
||||
|
||||
1. **Explicit "local".** Returns None — the dispatcher disables cloud
|
||||
mode entirely. Mirrors legacy short-circuit in
|
||||
:func:`tools.browser_tool._get_cloud_provider`.
|
||||
2. **Explicit config wins, ignoring availability.** If ``configured``
|
||||
names a registered provider, return it even if its
|
||||
:meth:`is_available` returns False — the dispatcher will surface a
|
||||
precise "X_API_KEY is not set" error instead of silently routing
|
||||
somewhere else.
|
||||
3. **Legacy preference walk, filtered by availability.** Walk
|
||||
:data:`_LEGACY_PREFERENCE` (``browser-use`` → ``browserbase``) looking
|
||||
for a provider whose ``is_available()`` is True.
|
||||
|
||||
There is intentionally NO "single-eligible shortcut" rule here (unlike
|
||||
:func:`agent.web_search_registry._resolve`). Pre-migration, the
|
||||
auto-detect branch in ``tools.browser_tool._get_cloud_provider`` only
|
||||
considered Browser Use and Browserbase; Firecrawl was reachable only
|
||||
via an explicit ``browser.cloud_provider: firecrawl`` config key.
|
||||
Preserving that gate matters because Firecrawl shares its API key with
|
||||
the *web* extract plugin (``plugins/web/firecrawl/``), so users who set
|
||||
``FIRECRAWL_API_KEY`` for web extract must NOT get silently routed to a
|
||||
paid cloud browser on a fresh install. Third-party browser-provider
|
||||
plugins added under ``~/.hermes/plugins/browser/<vendor>/`` are subject
|
||||
to the same gate — they must be explicitly configured to take effect.
|
||||
|
||||
Returns None when no provider is configured AND no available provider
|
||||
matches the legacy preference; the dispatcher then falls back to local
|
||||
browser mode.
|
||||
Intentionally NO "single-eligible shortcut" (unlike ``agent.web_search_registry._resolve``): only
|
||||
``_LEGACY_PREFERENCE`` names are auto-eligible. Firecrawl shares its API key with the *web* extract
|
||||
plugin, so a user with ``FIRECRAWL_API_KEY`` must never be routed to a paid cloud browser without
|
||||
setting ``browser.cloud_provider``; the same gate applies to third-party browser-provider plugins.
|
||||
"""
|
||||
with _lock:
|
||||
snapshot = dict(_providers)
|
||||
snapshot.update(_scoped_providers.get(hermes_home_key(), {}))
|
||||
|
||||
def _is_available_safe(p: BrowserProvider) -> bool:
|
||||
"""Wrap ``is_available()`` so a buggy provider doesn't kill resolution."""
|
||||
try:
|
||||
return bool(p.is_available())
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning(
|
||||
"Browser provider %s.is_available() raised %s — treating as unavailable",
|
||||
p.name, exc, exc_info=True,
|
||||
)
|
||||
return False
|
||||
|
||||
# 1. Explicit "local" short-circuit.
|
||||
snapshot = _registry.merged()
|
||||
if configured == "local":
|
||||
return None
|
||||
|
||||
# 2. Explicit config wins — return regardless of is_available() so the
|
||||
# user gets a precise downstream error message rather than a silent
|
||||
# backend switch. Matches _get_cloud_provider() in browser_tool.py.
|
||||
if configured:
|
||||
provider = snapshot.get(configured)
|
||||
if provider is not None:
|
||||
return provider
|
||||
logger.debug(
|
||||
"browser cloud_provider '%s' configured but not registered; "
|
||||
"falling back to auto-detect",
|
||||
"browser cloud_provider '%s' configured but not registered; falling back to auto-detect",
|
||||
configured,
|
||||
)
|
||||
|
||||
# 3. Legacy preference walk — only providers in _LEGACY_PREFERENCE are
|
||||
# auto-eligible. Filtered by availability so we don't surface a
|
||||
# provider the user has no credentials for. See docstring for why
|
||||
# we do NOT fall back to "any single-eligible registered provider".
|
||||
for legacy in _LEGACY_PREFERENCE:
|
||||
provider = snapshot.get(legacy)
|
||||
if provider is not None and _is_available_safe(provider):
|
||||
if provider is not None and is_available_safe(
|
||||
provider, logger,
|
||||
"Browser provider %s.is_available() raised %s — treating as unavailable",
|
||||
level=logging.WARNING, exc_info=True,
|
||||
):
|
||||
return provider
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _reset_for_tests() -> None:
|
||||
"""Clear the registry. **Test-only.**"""
|
||||
global _generation
|
||||
with _lock:
|
||||
_providers.clear()
|
||||
_scoped_providers.clear()
|
||||
_scoped_generations.clear()
|
||||
_generation += 1
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
from typing import Dict # noqa: F401,E402
|
||||
from typing import List # noqa: F401,E402
|
||||
import threading # noqa: F401,E402
|
||||
|
||||
|
||||
_PLUGIN_COMPAT_LAZY = {
|
||||
'hermes_home_key': ('hermes_constants', 'hermes_home_key'),
|
||||
}
|
||||
|
||||
|
||||
def __getattr__(name): # PEP 562 — lazy so no import cycles
|
||||
target = _PLUGIN_COMPAT_LAZY.get(name)
|
||||
if target is None:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
import importlib
|
||||
from hermes_cli.plugin_compat import warn_once
|
||||
warn_once(__name__, name, *target)
|
||||
return getattr(importlib.import_module(target[0]), target[1])
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
+2525
-4580
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,74 @@
|
||||
"""Relay-side accumulator for the chat_completions streaming wire.
|
||||
|
||||
Relay invokes its collector for every post-intercept chunk and then its finalizer as soon
|
||||
as the provider stream ends — concurrently with Hermes' consumer thread, which may not have
|
||||
read the last chunk yet. The finalizer therefore builds Relay's recorded response from
|
||||
collector-observed state only, never from the consumer loop's closures. Sibling of
|
||||
``relay_llm.AnthropicStreamAccumulator``; Bedrock and Codex follow the same contract.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
from agent.chat_completion_helpers import _ToolCallAccumulator
|
||||
from agent.message_content import flatten_message_text
|
||||
from agent.reasoning_summaries import separate_glued_reasoning_blocks
|
||||
|
||||
|
||||
def _tool_call_delta_view(tc_delta: Any) -> Any:
|
||||
"""Attribute view of a JSON tool-call delta for ``_ToolCallAccumulator.feed`` (written
|
||||
against SDK objects). Only ``function`` is wrapped: ``feed`` passes ``extra_content``
|
||||
(a dict) straight through ``_dump_if_model``, so a recursive view would corrupt it."""
|
||||
if not isinstance(tc_delta, dict):
|
||||
return tc_delta
|
||||
function = tc_delta.get("function")
|
||||
return SimpleNamespace(**{**tc_delta,
|
||||
"function": SimpleNamespace(**function) if isinstance(function, dict) else function})
|
||||
|
||||
|
||||
class RelayChatAccumulator:
|
||||
"""Rebuild a chat.completion from Relay's post-intercept chunk dicts."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._content: list[str] = []
|
||||
self._reasoning: list[str] = []
|
||||
self._tool_calls = _ToolCallAccumulator()
|
||||
self._model = self._usage = self._finish_reason = None
|
||||
self._role = "assistant"
|
||||
|
||||
def observe(self, chunk: Any) -> None:
|
||||
if not isinstance(chunk, dict):
|
||||
return
|
||||
self._model = chunk.get("model") or self._model
|
||||
if chunk.get("usage"):
|
||||
self._usage = chunk["usage"]
|
||||
choices = chunk.get("choices") or []
|
||||
choice = choices[0] if choices else None # Hermes never requests n>1
|
||||
if not isinstance(choice, dict):
|
||||
return
|
||||
self._finish_reason = choice.get("finish_reason") or self._finish_reason
|
||||
delta = choice.get("delta")
|
||||
if not isinstance(delta, dict):
|
||||
return
|
||||
if delta.get("role"):
|
||||
self._role = delta["role"]
|
||||
text = flatten_message_text(delta.get("content"), sep="")
|
||||
if text:
|
||||
self._content.append(text)
|
||||
reasoning = delta.get("reasoning_content") or delta.get("reasoning")
|
||||
if reasoning:
|
||||
self._reasoning.append(separate_glued_reasoning_blocks(
|
||||
self._reasoning[-1] if self._reasoning else "", reasoning))
|
||||
for tc_delta in delta.get("tool_calls") or []:
|
||||
self._tool_calls.feed(_tool_call_delta_view(tc_delta))
|
||||
|
||||
def finalize(self) -> dict[str, Any]:
|
||||
acc = self._tool_calls.materialize()
|
||||
message = {"role": self._role, "content": "".join(self._content) or None,
|
||||
"reasoning_content": "".join(self._reasoning) or None,
|
||||
"tool_calls": [acc[i] for i in sorted(acc)] or None}
|
||||
# "stop" also covers Nous Portal ``lastOne`` usage frames, which carry no finish_reason.
|
||||
return {"model": self._model, "usage": self._usage,
|
||||
"choices": [{"message": message, "finish_reason": self._finish_reason or "stop"}]}
|
||||
@@ -0,0 +1,285 @@
|
||||
"""Request-local worker lifecycle, watchdog polling, and wait status."""
|
||||
|
||||
from agent import chat_completion_helpers as h
|
||||
|
||||
|
||||
class _NonStreamRequest:
|
||||
"""One non-streaming request on a worker thread, polled by the caller.
|
||||
|
||||
State shared between the worker (``_call``) and the poll loop lives on the
|
||||
instance; ``_abort_request`` may run from the poll (stranger) thread.
|
||||
"""
|
||||
|
||||
def __init__(self, agent, api_kwargs: dict):
|
||||
self.agent = agent
|
||||
self.api_kwargs = api_kwargs
|
||||
self.result = {"response": None, "error": None}
|
||||
self.clients = h._RequestClientRegistry(agent)
|
||||
# Request-local cancel flag: agent._interrupt_requested is cleared at turn
|
||||
# boundaries but this daemon worker can outlive the turn, so it must know THIS
|
||||
# request was force-closed and not surface the transport error as a bug (#6600).
|
||||
self.cancelled = False
|
||||
# Codex retirement token: the worker checks ``agent._active_codex_stream_request_token``
|
||||
# to know it still owns the turn; a watchdog kill clears it so a worker still
|
||||
# draining SSE raises instead of returning partial output as "completed"
|
||||
# (run_codex_stream._request_is_current). ``codex_retired`` mirrors it locally.
|
||||
self.codex_token = object() if agent.api_mode == "codex_responses" else None
|
||||
self.codex_retired = False
|
||||
self.wd = h._resolve_nonstream_watchdogs(agent, api_kwargs)
|
||||
self.codex_watchdog_state = (
|
||||
h.SimpleNamespace(
|
||||
token=self.codex_token,
|
||||
lock=h.threading.Lock(),
|
||||
last_event_ts=None,
|
||||
last_progress_ts=None,
|
||||
retry_started_ts=None,
|
||||
phase_aware=self.wd.idle_requires_progress,
|
||||
)
|
||||
if self.codex_token is not None
|
||||
else None
|
||||
)
|
||||
self.call_start = h.time.time()
|
||||
self.wait_notice_started_ts = None
|
||||
self.thread = None
|
||||
|
||||
def _install_codex_request_token(self) -> None:
|
||||
if self.codex_token is not None and not self.codex_retired: # retired before start: don't re-publish
|
||||
self.agent._active_codex_stream_request_token = self.codex_token
|
||||
|
||||
def _retire_codex_request_token(self) -> None:
|
||||
if self.codex_token is None:
|
||||
return
|
||||
self.codex_retired = True
|
||||
if getattr(self.agent, "_active_codex_stream_request_token", None) is self.codex_token:
|
||||
self.agent._active_codex_stream_request_token = None
|
||||
|
||||
def _make_client(self, reason: str, kind: str = "openai"):
|
||||
# Per-request clients are registered with the abort machinery so the watchdogs
|
||||
# force-close the worker's connection, never the shared client (#67142).
|
||||
if kind == "anthropic_messages":
|
||||
client = self.agent._create_request_anthropic_client(reason=reason)
|
||||
else:
|
||||
client = self.agent._create_request_openai_client(reason=reason, api_kwargs=self.api_kwargs)
|
||||
return self.clients.set_client(client, kind=kind)
|
||||
|
||||
def _call(self):
|
||||
watchdog_state_var = watchdog_context_token = None
|
||||
try:
|
||||
self._install_codex_request_token()
|
||||
if self.codex_watchdog_state is not None:
|
||||
from agent.codex_runtime import _codex_watchdog_state_var
|
||||
|
||||
watchdog_state_var = _codex_watchdog_state_var
|
||||
watchdog_context_token = watchdog_state_var.set(self.codex_watchdog_state)
|
||||
self.result["response"] = h._dispatch_nonstreaming_api_request(
|
||||
self.agent, self.api_kwargs, make_client=self._make_client)
|
||||
except Exception as e:
|
||||
# Our own force-close caused this error: swallow it, the main
|
||||
# thread raises InterruptedError (#6600). Retirement logs at info
|
||||
# (a watchdog discarded output the provider already sent — what an
|
||||
# operator debugging a truncated reply needs); cancellation at debug.
|
||||
if self.codex_retired:
|
||||
h.logger.info("Codex worker caught %s after request retirement — "
|
||||
"discarding the stale partial instead of surfacing it as a completed response. %s",
|
||||
type(e).__name__, self.agent._client_log_context())
|
||||
return
|
||||
if self.cancelled:
|
||||
h.logger.debug("Non-streaming worker caught %s after request "
|
||||
"cancellation — exiting without surfacing a network error.", type(e).__name__)
|
||||
return
|
||||
self.result["error"] = e
|
||||
finally:
|
||||
if watchdog_state_var is not None:
|
||||
watchdog_state_var.reset(watchdog_context_token)
|
||||
# Retire first: close_once can raise, and a leaked token would let
|
||||
# a later worker mistake itself for the owning attempt.
|
||||
self._retire_codex_request_token()
|
||||
# Reuse reason only on a clean response; error or cancel-swallow
|
||||
# really closes so the next attempt builds a fresh pool.
|
||||
self.clients.close_once(
|
||||
"request_complete" if self.result["response"] is not None else "request_error_cleanup")
|
||||
|
||||
def _abort_request(self, reason: str) -> None:
|
||||
"""Watchdog/interrupt kill: abort the request client (kind-aware, #67142)
|
||||
and retire the codex token; the worker sees its own forced close via
|
||||
the cancel flags."""
|
||||
with h.contextlib.suppress(Exception):
|
||||
self.clients.close_once(reason)
|
||||
self._retire_codex_request_token()
|
||||
|
||||
def _await_worker_after_kill(self, timeout_message: str) -> None:
|
||||
# Wait briefly for the worker to notice the closed connection.
|
||||
self.thread.join(timeout=2.0)
|
||||
if self.result["error"] is None and self.result["response"] is None:
|
||||
self.result["error"] = TimeoutError(timeout_message)
|
||||
|
||||
def _model(self) -> str:
|
||||
return self.api_kwargs.get("model", "unknown")
|
||||
|
||||
def _codex_watchdog_snapshot(self):
|
||||
state = self.codex_watchdog_state
|
||||
if state is None: # non-codex request: no watchdog reads these
|
||||
return (None, None, None)
|
||||
with state.lock:
|
||||
return state.last_event_ts, state.last_progress_ts, state.retry_started_ts
|
||||
|
||||
def _emit_wait_notice(self, elapsed: float, *, heartbeat: bool = True) -> None:
|
||||
wd = self.wd
|
||||
try:
|
||||
last_event_ts, last_progress_ts, retry_started_ts = self._codex_watchdog_snapshot()
|
||||
activity_ts = retry_started_ts if retry_started_ts is not None else last_event_ts
|
||||
# Only undo a notice this request owns, promptly rather than at the
|
||||
# next heartbeat: reasoning callbacks do not reset the CLI spinner.
|
||||
if (self.wait_notice_started_ts is not None and activity_ts is not None
|
||||
and activity_ts > self.wait_notice_started_ts):
|
||||
self.agent._emit_wait_notice("")
|
||||
self.wait_notice_started_ts = None
|
||||
if not heartbeat:
|
||||
return
|
||||
silence = self.call_start + elapsed - (
|
||||
activity_ts if activity_ts is not None else self.call_start)
|
||||
if silence < 60.0:
|
||||
self.agent._touch_activity(
|
||||
"waiting for first stream event after reconnect"
|
||||
if retry_started_ts is not None else "waiting for provider response")
|
||||
return
|
||||
status = "no response yet"
|
||||
if retry_started_ts is not None:
|
||||
status = "no response after reconnect"
|
||||
elif last_event_ts is not None:
|
||||
status = "no stream events"
|
||||
recovery = h._codex_wait_notice_recovery(stale_timeout=wd.stale_timeout,
|
||||
ttfb_enabled=wd.ttfb_enabled, ttfb_timeout=wd.ttfb_timeout,
|
||||
last_event_ts=last_event_ts, last_progress_ts=last_progress_ts,
|
||||
retry_started_ts=retry_started_ts,
|
||||
call_start=self.call_start, idle_enabled=wd.idle_enabled, idle_timeout=wd.idle_timeout,
|
||||
idle_requires_progress=wd.idle_requires_progress,
|
||||
elapsed=elapsed)
|
||||
if recovery and activity_ts is not None:
|
||||
recovery += " total elapsed"
|
||||
self.agent._emit_wait_notice(
|
||||
f"⏳ waiting on {self.api_kwargs.get('model', 'the provider')} — "
|
||||
f"{int(silence)}s with {status} (provider may be slow or overloaded{recovery})")
|
||||
self.wait_notice_started_ts = self.call_start + elapsed
|
||||
except Exception:
|
||||
h.logger.debug("wait-notice construction failed", exc_info=True)
|
||||
|
||||
def _ttfb_kill(self, elapsed: float) -> None:
|
||||
"""No parsed Codex event past the first-event cutoff — kill so the retry loop
|
||||
reconnects instead of waiting out the stale timeout."""
|
||||
agent, wd = self.agent, self.wd
|
||||
silent_hint = h._codex_silent_hang_hint(agent, self.api_kwargs)
|
||||
h.logger.warning("Codex stream produced no parsed stream event within TTFB cutoff "
|
||||
"(%.0fs > %.0fs, model=%s). Backend accepted the connection "
|
||||
"but sent no stream events. Killing connection so the retry loop can reconnect.", elapsed,
|
||||
wd.ttfb_timeout, self._model())
|
||||
agent._buffer_status(
|
||||
f"⚠️ No first stream event from provider in {int(elapsed)}s (codex stream, model: {self._model()}). "
|
||||
f"Reconnecting." + (f" {silent_hint}" if silent_hint else ""))
|
||||
self._abort_request("codex_ttfb_kill")
|
||||
agent._emit_wait_notice(f"⚠ no response from provider in {int(elapsed)}s — reconnecting...")
|
||||
agent._touch_activity(f"codex stream killed after {int(elapsed)}s with no first stream event")
|
||||
self._await_worker_after_kill(
|
||||
f"Codex stream produced no parsed stream event within {int(elapsed)}s "
|
||||
f"(TTFB threshold: {int(wd.ttfb_timeout)}s)"
|
||||
+ (f". {silent_hint}" if silent_hint else ""))
|
||||
|
||||
def _idle_kill(self, event_stale_elapsed: float) -> None:
|
||||
"""SSE events stopped after the phase-specific idle arm point.
|
||||
|
||||
Only the implicit official OpenAI Codex policy arms on substantive model
|
||||
progress; compatible providers and explicit operator timeouts arm on first
|
||||
parsed event. Once armed, any parsed SSE event refreshes transport activity.
|
||||
"""
|
||||
agent, wd = self.agent, self.wd
|
||||
arm_point = "model progress began" if wd.idle_requires_progress else "the first parsed event"
|
||||
h.logger.warning("Codex stream produced no SSE events for %.0fs after %s "
|
||||
"(threshold %.0fs, model=%s, context=~%s tokens). Killing "
|
||||
"connection so the retry loop can reconnect.", event_stale_elapsed, arm_point, wd.idle_timeout,
|
||||
self._model(), f"{wd.est_tokens:,}")
|
||||
agent._buffer_status(
|
||||
f"⚠️ Codex stream sent no events for {int(event_stale_elapsed)}s after {arm_point} "
|
||||
f"(model: {self._model()}). Reconnecting.")
|
||||
self._abort_request("codex_stream_idle_kill")
|
||||
agent._touch_activity(f"codex stream killed after {int(event_stale_elapsed)}s with no SSE events")
|
||||
self._await_worker_after_kill(
|
||||
f"Codex stream produced no SSE events for {int(event_stale_elapsed)}s "
|
||||
f"after {arm_point} (threshold: {int(wd.idle_timeout)}s)")
|
||||
|
||||
def _stale_kill(self, elapsed: float) -> None:
|
||||
"""No response within the stale timeout: kill and count toward the
|
||||
circuit breaker (#58962, see ``_stale_streak``)."""
|
||||
agent, wd = self.agent, self.wd
|
||||
silent_hint = h._codex_silent_hang_hint(agent, self.api_kwargs)
|
||||
h._report_stale_nonstream_kill(agent, self.api_kwargs, elapsed, wd.stale_timeout, hint=silent_hint)
|
||||
self._abort_request("stale_call_kill")
|
||||
h._bump_stale_streak(agent)
|
||||
h._touch_stale_kill_activity(agent, elapsed)
|
||||
self._await_worker_after_kill(
|
||||
f"Non-streaming API call timed out after {int(elapsed)}s with no response (threshold: {int(wd.stale_timeout)}s)"
|
||||
+ (f". {silent_hint}" if silent_hint else ""))
|
||||
|
||||
def _interrupt(self, elapsed: float) -> None:
|
||||
agent = self.agent
|
||||
last_event_ts, _, _ = self._codex_watchdog_snapshot()
|
||||
h._record_interrupted_provider_wait(agent, elapsed,
|
||||
response_started=self.wd.codex and last_event_ts is not None
|
||||
)
|
||||
# Mark cancelled BEFORE force-closing so the worker treats the transport
|
||||
# error as a cancel (#6600). Never close the shared client (releasing a
|
||||
# TLS FD mid-SSL-BIO corrupted an unrelated SQLite DB, #67142). Then let
|
||||
# the worker unwind Relay scopes before raising (#81521).
|
||||
self.cancelled = True
|
||||
h.logger.debug("Force-closing httpx client due to interrupt (not a network error).")
|
||||
self._abort_request("interrupt_abort")
|
||||
h._join_worker_for_relay_teardown(self.thread, label="Non-streaming")
|
||||
raise InterruptedError("Agent interrupted during API call")
|
||||
|
||||
def run(self):
|
||||
agent, wd = self.agent, self.wd
|
||||
if wd.codex:
|
||||
# Reset before the worker starts so a marker left over from a previous
|
||||
# call on this agent can't be misread as the first event for this one.
|
||||
with self.codex_watchdog_state.lock:
|
||||
self.codex_watchdog_state.last_event_ts = None
|
||||
self.codex_watchdog_state.last_progress_ts = None
|
||||
self.codex_watchdog_state.retry_started_ts = None
|
||||
agent._touch_activity("waiting for non-streaming API response")
|
||||
|
||||
self.thread = t = h.threading.Thread(target=h._context_thread_target(self._call), daemon=True)
|
||||
t.start()
|
||||
poll_count = 0
|
||||
while t.is_alive():
|
||||
t.join(timeout=0.3)
|
||||
poll_count += 1
|
||||
# Keep the quiet gateway heartbeat; only silence warrants a notice.
|
||||
# Resumed events clear our notice on the next poll, not 30s later.
|
||||
now = h.time.time()
|
||||
elapsed = now - self.call_start
|
||||
self._emit_wait_notice(elapsed, heartbeat=poll_count % 100 == 0)
|
||||
last_event_ts, last_progress_ts, retry_started_ts = self._codex_watchdog_snapshot()
|
||||
retry_ttfb_elapsed = now - retry_started_ts if retry_started_ts is not None else None
|
||||
if wd.ttfb_enabled and retry_ttfb_elapsed is not None and retry_ttfb_elapsed > wd.ttfb_timeout:
|
||||
self._ttfb_kill(retry_ttfb_elapsed)
|
||||
break
|
||||
if (retry_started_ts is None and wd.ttfb_enabled
|
||||
and elapsed > wd.ttfb_timeout and last_event_ts is None):
|
||||
self._ttfb_kill(elapsed)
|
||||
break
|
||||
idle_elapsed = now - last_event_ts if last_event_ts is not None else None
|
||||
if (retry_started_ts is None and wd.idle_enabled and idle_elapsed is not None
|
||||
and (not wd.idle_requires_progress or last_progress_ts is not None)
|
||||
and idle_elapsed > wd.idle_timeout):
|
||||
self._idle_kill(idle_elapsed)
|
||||
break
|
||||
if elapsed > wd.stale_timeout:
|
||||
self._stale_kill(elapsed)
|
||||
break
|
||||
if agent._interrupt_requested:
|
||||
self._interrupt(elapsed)
|
||||
if self.result["error"] is not None:
|
||||
raise self.result["error"]
|
||||
# Success — the provider proved responsive: clear the breaker (#58962).
|
||||
if self.result["response"] is not None:
|
||||
h._reset_stale_streak(agent)
|
||||
return self.result["response"]
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Display and heartbeat phase of the request-local streaming monitor."""
|
||||
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
|
||||
from agent.model_metadata import is_local_endpoint
|
||||
|
||||
|
||||
class StreamingWaitMonitor:
|
||||
def _poll_local_load_notice(self, now: float) -> bool:
|
||||
"""Managed local server: surface a cold model's weight-load progress
|
||||
instead of the 60s "provider may be slow" copy. Polled ~1s only while no
|
||||
REAL chunk arrived for 2s+ (never during healthy token flow); in-memory,
|
||||
no network. True while loading = heartbeat liveness, skip the rest of
|
||||
this iteration (the stale detector's local floor dwarfs any load)."""
|
||||
from agent.chat_completion_helpers import _managed_local_load_notice
|
||||
|
||||
m = self._mon
|
||||
if now - self.last_chunk_time["t"] < 2.0 or now - m.last_load_poll < 1.0:
|
||||
return False
|
||||
m.last_load_poll = now
|
||||
_load_notice = _managed_local_load_notice(self.agent, self.api_kwargs)
|
||||
if _load_notice is not None:
|
||||
m.wait_notice_started_ts = None # The local loader now owns the display.
|
||||
self.agent._emit_wait_notice(_load_notice)
|
||||
self.agent._touch_activity("local model loading")
|
||||
m.load_notice_shown, m.load_notice_misses, m.last_heartbeat = True, 0, now # loading IS liveness
|
||||
return True
|
||||
if m.load_notice_shown:
|
||||
# One missed sample is routine (probe timeout under load); clearing on it strobed the line.
|
||||
m.load_notice_misses += 1
|
||||
if m.load_notice_misses >= 3:
|
||||
m.load_notice_shown, m.load_notice_misses = False, 0
|
||||
self.agent._emit_wait_notice("")
|
||||
return False
|
||||
|
||||
def _heartbeat(self, waiting_secs: int) -> None:
|
||||
"""Gateway inactivity heartbeat: the start-to-first-chunk gap (thinking,
|
||||
local prefill) can exceed the gateway timeout."""
|
||||
if waiting_secs >= 60.0:
|
||||
# No chunks for 60s+: say WHAT the wait is and WHEN recovery kicks in.
|
||||
stale = self._stream_stale_timeout
|
||||
_recovery = f"; auto-reconnect at {int(stale)}s" if stale is not None and stale != float("inf") else ""
|
||||
self._mon.wait_notice_started_ts = self._mon.last_heartbeat
|
||||
self.agent._emit_wait_notice(
|
||||
f"⏳ waiting on {self.api_kwargs.get('model', 'the provider')} — no stream output for {waiting_secs}s "
|
||||
f"(provider may be slow or overloaded, or the model is thinking{_recovery})")
|
||||
else:
|
||||
# Chunks are flowing — keep the tracker fresh, leave the display alone.
|
||||
self.agent._touch_activity(f"waiting for stream response ({waiting_secs}s, no chunks yet)")
|
||||
|
||||
def _monitor_loop(self) -> None:
|
||||
_HEARTBEAT_INTERVAL = 30.0 # seconds between gateway activity touches
|
||||
self._mon = SimpleNamespace(
|
||||
last_heartbeat=time.time(), last_load_poll=0.0,
|
||||
load_notice_shown=False, load_notice_misses=0, wait_notice_started_ts=None,
|
||||
)
|
||||
_is_local_base = bool(self.agent.base_url) and is_local_endpoint(self.agent.base_url)
|
||||
while not self._call_done.is_set():
|
||||
self._call_done.wait(timeout=0.3)
|
||||
_hb_now = time.time()
|
||||
if _is_local_base and self._poll_local_load_notice(_hb_now):
|
||||
continue
|
||||
# Reasoning callbacks do not clear the classic CLI spinner. The empty
|
||||
# protocol payload resets status without adding synthetic reasoning.
|
||||
if (self._mon.wait_notice_started_ts is not None
|
||||
and self.last_chunk_time["t"] > self._mon.wait_notice_started_ts):
|
||||
self.agent._emit_wait_notice("")
|
||||
self._mon.wait_notice_started_ts = None
|
||||
if _hb_now - self._mon.last_heartbeat >= _HEARTBEAT_INTERVAL:
|
||||
self._mon.last_heartbeat = _hb_now
|
||||
self._heartbeat(int(_hb_now - self.last_chunk_time["t"]))
|
||||
_stale_elapsed = time.time() - self.last_chunk_time["t"]
|
||||
if _stale_elapsed > self._stream_stale_timeout:
|
||||
self._mon.wait_notice_started_ts = None # Reconnect status has its own owner.
|
||||
self._kill_stale_stream(_stale_elapsed)
|
||||
if self.agent._interrupt_requested:
|
||||
self._abort_for_interrupt(_stale_elapsed)
|
||||
return
|
||||
|
||||
@@ -0,0 +1,970 @@
|
||||
"""Tool-resource teardown, wire-client lifecycle and credential refresh for ``AIAgent``.
|
||||
|
||||
``ClientLifecycleMixin`` owns task cleanup, the shared primary client, per-request client caches
|
||||
(owner-thread close vs stranger-thread abort), credential rotation and route-derived headers.
|
||||
"""
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from typing import Any, Optional
|
||||
|
||||
from agent.lazy_forward import forward as _forward, forward_static as _forward_static, lazy_attr as _lazy_attr
|
||||
from hermes_cli.timeouts import get_provider_request_timeout
|
||||
from utils import base_url_host_matches, env_float
|
||||
|
||||
logger = logging.getLogger("run_agent") # origin module's logger name: log records / caplog filters unchanged
|
||||
_QWEN_CODE_VERSION = "0.14.1" # Qwen Portal mimics the QwenCode CLI
|
||||
# Per-request cache slot attribute names (OpenAI-style and Anthropic clients).
|
||||
_OPENAI_SLOT = "_request_client_cache"
|
||||
_ANTHROPIC_SLOT = "_request_anthropic_client_cache"
|
||||
_NO_SOCKETS_SUFFIX = " — no sockets found; in-flight request may keep running until the provider finishes"
|
||||
|
||||
|
||||
def _routermint_headers() -> dict:
|
||||
"""User-Agent RouterMint needs to avoid Cloudflare 1010 blocks."""
|
||||
from hermes_cli import __version__ as _HERMES_VERSION
|
||||
return {"User-Agent": f"HermesAgent/{_HERMES_VERSION}"}
|
||||
|
||||
|
||||
def _qwen_portal_headers() -> dict:
|
||||
import platform as _plat
|
||||
_ua = f"QwenCode/{_QWEN_CODE_VERSION} ({_plat.system().lower()}; {_plat.machine()})"
|
||||
return {
|
||||
"User-Agent": _ua, "X-DashScope-CacheControl": "enable", "X-DashScope-UserAgent": _ua,
|
||||
"X-DashScope-AuthType": "qwen-oauth",
|
||||
}
|
||||
|
||||
|
||||
# Route-specific default headers; first host match wins (order preserved from the original chain).
|
||||
# Builders resolve their module lazily so run_agent keeps its import-time cost and avoids cycles.
|
||||
_ROUTE_DEFAULT_HEADERS = (
|
||||
("openrouter.ai", lambda self, url: _lazy_attr("agent.auxiliary_client", "build_or_headers")()),
|
||||
("ai-gateway.vercel.sh", lambda self, url: dict(_lazy_attr("agent.auxiliary_client", "_AI_GATEWAY_HEADERS"))),
|
||||
("integrate.api.nvidia.com", lambda self, url: _lazy_attr("agent.auxiliary_client", "build_nvidia_nim_headers")(url)),
|
||||
("api.routermint.com", lambda self, url: _routermint_headers()),
|
||||
("githubcopilot.com", lambda self, url: _lazy_attr("hermes_cli.models", "copilot_default_headers")()),
|
||||
("api.kimi.com", lambda self, url: dict(_lazy_attr("agent.auxiliary_client", "_AI_GATEWAY_HEADERS"))),
|
||||
("portal.qwen.ai", lambda self, url: _qwen_portal_headers()),
|
||||
("chatgpt.com", lambda self, url: _lazy_attr("agent.codex_headers", "codex_cloudflare_headers")(
|
||||
self._client_kwargs.get("api_key", ""), base_url=url)),
|
||||
# Covers provider=xai and provider=xai-oauth (api.x.ai).
|
||||
("x.ai", lambda self, url: _lazy_attr("tools.xai_http", "hermes_xai_default_headers")()),
|
||||
)
|
||||
|
||||
|
||||
def _reset_slot(cache: dict, *, in_use: bool = False) -> None:
|
||||
cache["client"] = None
|
||||
cache["key"] = None
|
||||
cache["poisoned"] = False
|
||||
cache["in_use"] = in_use
|
||||
|
||||
|
||||
def _valid_credential_pair(api_key: Any, base_url: Any) -> bool:
|
||||
return bool(isinstance(api_key, str) and api_key.strip() and isinstance(base_url, str) and base_url.strip())
|
||||
|
||||
|
||||
def _swap_fallback_clients(agent, fb_client, fb_provider: str, fb_model: str, fb_base_url: str, fb_api_mode: str) -> None:
|
||||
"""Install the fallback client(s) in place, honoring request_timeout_seconds (None = SDK default)."""
|
||||
timeout = get_provider_request_timeout(fb_provider, fb_model)
|
||||
# The SDK exposes an empty/stale api_key when a rotating source is installed.
|
||||
key_provider = vars(fb_client).get("_api_key_provider")
|
||||
credential = key_provider if callable(key_provider) else fb_client.api_key
|
||||
if fb_api_mode == "anthropic_messages":
|
||||
from agent.anthropic_adapter import build_anthropic_client
|
||||
from agent.anthropic_credentials import resolve_anthropic_token, _is_oauth_token
|
||||
is_anthropic = fb_provider == "anthropic"
|
||||
effective_key = credential or (resolve_anthropic_token() if is_anthropic else None) or ""
|
||||
agent.api_key = agent._anthropic_api_key = effective_key
|
||||
agent._anthropic_base_url = fb_base_url
|
||||
agent._anthropic_client = build_anthropic_client(effective_key, fb_base_url, timeout=timeout)
|
||||
agent._is_anthropic_oauth = (
|
||||
_is_oauth_token(effective_key)
|
||||
if is_anthropic and isinstance(effective_key, str) else False
|
||||
)
|
||||
agent.client, agent._client_kwargs = None, {}
|
||||
return
|
||||
agent.api_key = credential
|
||||
agent.client = fb_client
|
||||
# Keep provider headers resolve_provider_client() baked into fb_client (SDK: _custom_headers), else
|
||||
# later request-client rebuilds drop them and User-Agent-sentinel providers (Kimi Coding) 403.
|
||||
fb_headers = getattr(fb_client, "_custom_headers", None) or getattr(fb_client, "default_headers", None)
|
||||
agent._client_kwargs = {"api_key": credential, "base_url": fb_base_url}
|
||||
if fb_headers:
|
||||
agent._client_kwargs["default_headers"] = dict(fb_headers)
|
||||
if timeout is not None:
|
||||
agent._client_kwargs["timeout"] = timeout
|
||||
# Rebuild now so the timeout applies to the very next request, not only after a rotation rebuild.
|
||||
agent._replace_primary_openai_client(reason="fallback_timeout_apply")
|
||||
|
||||
|
||||
class ClientLifecycleMixin:
|
||||
def _close_task_resources(self, task_id: str) -> None:
|
||||
"""Release task resources without treating a shared environment as process ownership."""
|
||||
from run_agent import _quietly, cleanup_browser, cleanup_vm
|
||||
|
||||
def kill_processes() -> None:
|
||||
from tools.process_registry import process_registry
|
||||
# A session can run several task IDs; delegated IDs also differ from session_id.
|
||||
# Never match the environment key (e.g. "default"), shared by parent and siblings.
|
||||
owners = getattr(self, "_process_owner_task_ids", ())
|
||||
for process in process_registry.list_sessions():
|
||||
if process["owner_task_id"] in owners and process["status"] == "running":
|
||||
process_registry.kill_process(
|
||||
process["session_id"], source="agent_close", consume_output=True,
|
||||
)
|
||||
|
||||
def release_computer_use() -> None:
|
||||
from tools.computer_use.tool import release_computer_use_session
|
||||
release_computer_use_session(task_id)
|
||||
|
||||
for step in (kill_processes, lambda: cleanup_vm(task_id), lambda: cleanup_browser(task_id), release_computer_use):
|
||||
_quietly(step)
|
||||
|
||||
def _client_log_context(self) -> str:
|
||||
thread = threading.current_thread()
|
||||
return (
|
||||
f"thread={thread.name}:{thread.ident} provider={getattr(self, 'provider', 'unknown')} "
|
||||
f"base_url={getattr(self, 'base_url', 'unknown')} model={getattr(self, 'model', 'unknown')}"
|
||||
)
|
||||
|
||||
def _anthropic_log_context(self) -> str:
|
||||
return f"provider={getattr(self, 'provider', None)} model={getattr(self, 'model', None)}"
|
||||
|
||||
def _openai_client_lock(self) -> threading.RLock:
|
||||
if getattr(self, "_client_lock", None) is None:
|
||||
self._client_lock = threading.RLock()
|
||||
return self._client_lock
|
||||
|
||||
@staticmethod
|
||||
def _is_openai_client_closed(client: Any) -> bool:
|
||||
"""Check if an OpenAI client is closed.
|
||||
|
||||
Handles both property and method forms of is_closed:
|
||||
- httpx.Client.is_closed is a bool property
|
||||
- openai.OpenAI.is_closed is a method returning bool
|
||||
|
||||
Prior bug: getattr(client, "is_closed", False) returned the bound method,
|
||||
which is always truthy, causing unnecessary client recreation on every call.
|
||||
"""
|
||||
from unittest.mock import Mock
|
||||
|
||||
if isinstance(client, Mock):
|
||||
return False
|
||||
|
||||
is_closed_attr = getattr(client, "is_closed", None)
|
||||
if is_closed_attr is not None:
|
||||
# Handle method (openai SDK) vs property (httpx)
|
||||
if callable(is_closed_attr):
|
||||
if is_closed_attr():
|
||||
return True
|
||||
elif bool(is_closed_attr):
|
||||
return True
|
||||
|
||||
http_client = getattr(client, "_client", None)
|
||||
if http_client is not None:
|
||||
return bool(getattr(http_client, "is_closed", False))
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _build_keepalive_http_client(base_url: str = "", *, verify: Any = True) -> Any:
|
||||
"""Build the shared OpenAI httpx client used by main and aux paths."""
|
||||
from agent.process_bootstrap import build_keepalive_http_client
|
||||
return build_keepalive_http_client(base_url, verify=verify)
|
||||
|
||||
_create_openai_client = _forward("agent.agent_runtime_helpers", "create_openai_client")
|
||||
_force_close_tcp_sockets = _forward_static("agent.agent_runtime_helpers", "force_close_tcp_sockets")
|
||||
_cleanup_dead_connections = _forward("agent.agent_runtime_helpers", "cleanup_dead_connections")
|
||||
_run_codex_stream = _forward("agent.codex_runtime", "run_codex_stream")
|
||||
_recover_with_credential_pool = _forward("agent.agent_runtime_helpers", "recover_with_credential_pool")
|
||||
|
||||
def _close_openai_client(self, client: Any, *, reason: str, shared: bool) -> None:
|
||||
if client is None:
|
||||
return
|
||||
ctx = self._client_log_context()
|
||||
# Force-close TCP sockets first (no CLOSE-WAIT accumulation), then the graceful SDK close.
|
||||
force_closed = self._force_close_tcp_sockets(client)
|
||||
try:
|
||||
client.close()
|
||||
logger.info("OpenAI client closed (%s, shared=%s, tcp_force_closed=%d) %s", reason, shared, force_closed, ctx)
|
||||
except Exception as exc:
|
||||
logger.debug("OpenAI client close failed (%s, shared=%s) %s error=%s", reason, shared, ctx, exc)
|
||||
|
||||
def _retire_shared_openai_client(self, client: Any, *, reason: str) -> None:
|
||||
"""Ownership-safe retirement of a replaced shared OpenAI client: ``shutdown()`` sockets only, FDs to GC.
|
||||
|
||||
``close()`` releases FDs from the calling thread while other threads may still hold the fd in an SSL BIO;
|
||||
a recycled fd then gets a TLS record written into an unrelated file (SQLite-header corruption).
|
||||
|
||||
The shared primary client has no single owning thread — worker threads from stale-killed attempts
|
||||
may still be unwinding their SSL BIOs, and the codex-direct / MoA paths stream on the shared client
|
||||
itself. If we release an FD while another thread's SSL layer still caches the raw integer fd, the
|
||||
kernel can recycle it into an unrelated ``open()`` (e.g. ``kanban.db``) and the unwinding TLS flush
|
||||
then writes an application-data record into that file — the SQLite-header corruption documented in
|
||||
#29507/#70773.
|
||||
"""
|
||||
if client is None:
|
||||
return
|
||||
try:
|
||||
shutdown_count = self._force_close_tcp_sockets(client)
|
||||
logger.info(
|
||||
"Shared OpenAI client retired (%s, tcp_shutdown=%d, fd_release=deferred_to_gc) %s",
|
||||
reason, shutdown_count, self._client_log_context(),
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("Shared OpenAI client retire failed (%s) %s error=%s", reason, self._client_log_context(), exc)
|
||||
|
||||
def _drain_transports_after_abandonment(self, *, reason: str) -> int:
|
||||
"""FD-safe transport drain for an abandoned (timed-out) worker; returns sockets shut down.
|
||||
|
||||
The worker may be blocked in an OpenSSL read; hard-closing from the timeout thread releases FDs under a
|
||||
live BIO (native corruption / SIGSEGV). Only ``shutdown()`` so the read sees EOF and the worker closes itself.
|
||||
|
||||
See #94248.
|
||||
A delegation deadline abandons this agent's daemon worker while it may still be blocked inside an
|
||||
in-flight OpenSSL ``read`` (Codex Responses stream, httpx request). This helper only ``shutdown()``s
|
||||
pooled sockets (safe from any thread), settling blocked reads with EOF/EPIPE so the worker can
|
||||
unwind and run the real close from its own thread. See #70773, #94248.
|
||||
"""
|
||||
drained = 0
|
||||
# Shared primary client (codex-direct / MoA stream on it directly).
|
||||
try:
|
||||
client = getattr(self, "client", None)
|
||||
if client is not None:
|
||||
drained += self._force_close_tcp_sockets(client)
|
||||
except Exception:
|
||||
logger.debug("Abandoned-worker drain: shared client sweep failed", exc_info=True)
|
||||
# Cached per-request wire clients: abort (shutdown + poison the reuse slot) so the unwinding
|
||||
# worker discards them instead of re-caching.
|
||||
for slot_attr, abort, label in (
|
||||
(_OPENAI_SLOT, self._abort_request_openai_client, "request"),
|
||||
(_ANTHROPIC_SLOT, self._abort_request_anthropic_client, "anthropic"),
|
||||
):
|
||||
try:
|
||||
with self._openai_client_lock():
|
||||
cache = getattr(self, slot_attr, None)
|
||||
cached = cache["client"] if cache else None
|
||||
if cached is not None:
|
||||
abort(cached, reason=reason)
|
||||
except Exception:
|
||||
logger.debug("Abandoned-worker drain: %s client abort failed", label, exc_info=True)
|
||||
# Codex app-server session watches a private interrupt event.
|
||||
try:
|
||||
request_interrupt = getattr(getattr(self, "_codex_session", None), "request_interrupt", None)
|
||||
if callable(request_interrupt):
|
||||
request_interrupt()
|
||||
except Exception:
|
||||
logger.debug("Abandoned-worker drain: codex interrupt failed", exc_info=True)
|
||||
# Inline (cron-style) request abort hook, when registered.
|
||||
try:
|
||||
abort_active = getattr(self, "_active_request_abort", None)
|
||||
if callable(abort_active):
|
||||
abort_active(reason)
|
||||
except Exception:
|
||||
logger.debug("Abandoned-worker drain: active request abort failed", exc_info=True)
|
||||
logger.info(
|
||||
"Abandoned-worker transports drained (%s, tcp_shutdown=%d, fd_release=deferred_to_worker) %s",
|
||||
reason, drained, self._client_log_context(),
|
||||
)
|
||||
return drained
|
||||
|
||||
def _replace_primary_openai_client(self, *, reason: str) -> bool:
|
||||
with self._openai_client_lock():
|
||||
old_client = getattr(self, "client", None)
|
||||
try:
|
||||
# MoA's ``client`` is an in-process facade, not an SDK client; generic rebuilds must preserve that.
|
||||
if (getattr(self, "provider", "") or "").strip().lower() == "moa":
|
||||
from agent.moa_loop import build_moa_facade
|
||||
new_client = build_moa_facade(self, self.model)
|
||||
else:
|
||||
new_client = self._create_openai_client(self._client_kwargs, reason=reason, shared=True)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to rebuild shared primary client (%s) %s error=%s", reason, self._client_log_context(), exc,
|
||||
)
|
||||
return False
|
||||
self.client = new_client
|
||||
# Never hard-close the replaced shared client (another thread may still be unwinding on the old pool).
|
||||
# #70773: never hard-close the replaced shared client from here — the caller may not be the thread
|
||||
# whose request is still unwinding on the old pool (credential rotation and dead-connection cleanup
|
||||
# run on the turn thread while stale-killed workers unwind; the codex-direct path streams on the
|
||||
# shared client itself). Retire it instead: sockets are shut down (FD-safe), FD release deferred to
|
||||
# GC.
|
||||
self._retire_shared_openai_client(old_client, reason=f"replace:{reason}")
|
||||
return True
|
||||
|
||||
def _ensure_primary_openai_client(self, *, reason: str) -> Any:
|
||||
with self._openai_client_lock():
|
||||
client = getattr(self, "client", None)
|
||||
if client is not None and not self._is_openai_client_closed(client):
|
||||
return client
|
||||
try:
|
||||
new_client = self._create_openai_client(self._client_kwargs, reason=reason, shared=True)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to recreate closed OpenAI client (%s) %s error=%s", reason, self._client_log_context(), exc,
|
||||
)
|
||||
raise RuntimeError("Failed to recreate closed OpenAI client") from exc
|
||||
self.client = new_client
|
||||
logger.warning(
|
||||
"Detected closed shared OpenAI client; recreated before use (%s) %s", reason, self._client_log_context(),
|
||||
)
|
||||
self._close_openai_client(client, reason=f"replace:{reason}", shared=True)
|
||||
return new_client
|
||||
|
||||
@staticmethod
|
||||
def _api_kwargs_have_image_parts(api_kwargs: dict) -> bool:
|
||||
"""True when the outbound request still has native image parts (Chat ``messages`` / Responses ``input``)."""
|
||||
if not isinstance(api_kwargs, dict):
|
||||
return False
|
||||
|
||||
def _contains_image(value: Any) -> bool:
|
||||
if isinstance(value, dict):
|
||||
return value.get("type") in {"image_url", "input_image"} or any(
|
||||
_contains_image(v) for v in value.values()
|
||||
)
|
||||
return isinstance(value, list) and any(_contains_image(v) for v in value)
|
||||
return any(
|
||||
_contains_image(item)
|
||||
for field in ("messages", "input")
|
||||
if isinstance(api_kwargs.get(field), list)
|
||||
for item in api_kwargs[field]
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ per-request client slots
|
||||
# One single-slot cache per client kind: {"client", "key", "poisoned", "in_use"}. Reuse keeps the warm httpx
|
||||
# pool between sequential calls; ``in_use`` keeps concurrent calls off one pool; ``poisoned`` marks a pool
|
||||
# whose sockets were shut from a stranger thread (never reuse it).
|
||||
# Reuse reasons: closes from the FD-owning worker's own finally after a response — the only closes that
|
||||
# attest a healthy pool. Poisoning still wins.
|
||||
_REQUEST_CLIENT_REUSE_REASONS = frozenset({"request_complete", "stream_request_complete"})
|
||||
|
||||
def _request_slot(self, slot_attr: str) -> dict:
|
||||
cache = getattr(self, slot_attr, None) # lazy: tests build agents via AIAgent.__new__ without __init__
|
||||
if cache is None:
|
||||
setattr(self, slot_attr, cache := {})
|
||||
_reset_slot(cache)
|
||||
return cache
|
||||
|
||||
def _checkout_request_slot(self, slot_attr: str, key: Any) -> tuple:
|
||||
"""Return ``(reusable_client, stale_client)``; at most one is non-None."""
|
||||
with self._openai_client_lock():
|
||||
cache = self._request_slot(slot_attr)
|
||||
cached = cache["client"]
|
||||
if cached is None or cache["in_use"]:
|
||||
return None, None
|
||||
if not cache["poisoned"] and cache["key"] == key and not self._is_openai_client_closed(cached):
|
||||
cache["in_use"] = True
|
||||
return cached, None
|
||||
# Key changed / poisoned / externally closed — rebuild. in_use was False, so closing the stale
|
||||
# client from this thread is FD-safe (no worker owns it).
|
||||
_reset_slot(cache)
|
||||
return None, cached
|
||||
|
||||
def _store_request_slot(self, slot_attr: str, client: Any, key: Any) -> None:
|
||||
with self._openai_client_lock():
|
||||
cache = self._request_slot(slot_attr)
|
||||
if cache["client"] is None:
|
||||
cache.update(client=client, key=key, poisoned=False, in_use=True)
|
||||
# else: a concurrent call holds the slot — hand this client out untracked (fully closed later).
|
||||
|
||||
def _release_request_slot(self, slot_attr: str, client: Any, reason: str) -> bool:
|
||||
"""Owner-thread release; True when the client stays cached (clean finish, not poisoned)."""
|
||||
with self._openai_client_lock():
|
||||
cache = self._request_slot(slot_attr)
|
||||
if cache["client"] is client:
|
||||
if reason in self._REQUEST_CLIENT_REUSE_REASONS and not cache["poisoned"]:
|
||||
cache["in_use"] = False
|
||||
return True
|
||||
_reset_slot(cache)
|
||||
return False
|
||||
|
||||
def _take_request_slot(self, slot_attr: str) -> tuple:
|
||||
"""Teardown: empty the slot and return ``(client, was_in_use)``."""
|
||||
with self._openai_client_lock():
|
||||
cache = getattr(self, slot_attr, None)
|
||||
client, in_use = (cache["client"], bool(cache["in_use"])) if cache else (None, False)
|
||||
if cache is not None:
|
||||
_reset_slot(cache)
|
||||
return client, in_use
|
||||
|
||||
def _abort_request_slot_client(self, slot_attr: str, client: Any, *, reason: str) -> None:
|
||||
"""Cross-thread abort (interrupt loop, stale detector): ``shutdown(SHUT_RDWR)`` without releasing FDs.
|
||||
|
||||
``close()`` from a non-owning thread races the live SSL BIO and corrupts unrelated FDs; shutdown unblocks
|
||||
the owner's recv/send so it closes from its own context. The slot is poisoned so the pool is never reused.
|
||||
"""
|
||||
if client is None:
|
||||
return
|
||||
anthropic = slot_attr == _ANTHROPIC_SLOT
|
||||
label = "Anthropic" if anthropic else "OpenAI"
|
||||
context = self._anthropic_log_context() if anthropic else self._client_log_context()
|
||||
with self._openai_client_lock():
|
||||
cache = self._request_slot(slot_attr)
|
||||
if cache["client"] is client:
|
||||
cache["poisoned"] = True
|
||||
try:
|
||||
shutdown_count = self._force_close_tcp_sockets(client)
|
||||
# Zero sockets shut down means the worker stays blocked — WARN, not success.
|
||||
# tcp_force_closed=0 means the stranger-thread abort found no sockets to shut down — the worker
|
||||
# stays blocked in recv and the provider keeps the slot (#72975). Surface that as WARNING so it
|
||||
# cannot be mistaken for a successful abort in the logs.
|
||||
# See #72975.
|
||||
_log = logger.warning if shutdown_count == 0 else logger.info
|
||||
_log(
|
||||
"%s client aborted (%s, shared=False, tcp_force_closed=%d, deferred_close=stranger_thread) %s%s",
|
||||
label, reason, shutdown_count, context, _NO_SOCKETS_SUFFIX if shutdown_count == 0 else "",
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("%s client abort failed (%s, shared=False) %s error=%s", label, reason, context, exc)
|
||||
|
||||
def _create_request_openai_client(self, *, reason: str, api_kwargs: Optional[dict] = None) -> Any:
|
||||
from unittest.mock import Mock
|
||||
primary_client = self._ensure_primary_openai_client(reason=reason)
|
||||
if self.provider == "moa" or isinstance(primary_client, Mock):
|
||||
return primary_client
|
||||
with self._openai_client_lock():
|
||||
request_kwargs = dict(self._client_kwargs)
|
||||
# No SDK retry loop: the outer loop owns retries/rotation/fallback, and SDK retries stretch a hung
|
||||
# request ~3x past our stale detector.
|
||||
request_kwargs["max_retries"] = 0
|
||||
is_copilot = base_url_host_matches(str(request_kwargs.get("base_url", "")), "githubcopilot.com")
|
||||
if is_copilot and self._api_kwargs_have_image_parts(api_kwargs or {}):
|
||||
from hermes_cli.copilot_auth import copilot_request_headers
|
||||
request_kwargs["default_headers"] = copilot_request_headers(is_agent_turn=True, is_vision=True)
|
||||
cached, stale = self._checkout_request_slot(_OPENAI_SLOT, request_kwargs)
|
||||
if cached is not None:
|
||||
return cached
|
||||
if stale is not None:
|
||||
self._close_openai_client(stale, reason=f"reuse_evict:{reason}", shared=False)
|
||||
client = self._create_openai_client(request_kwargs, reason=reason, shared=False)
|
||||
# Snapshot nested dicts (default_headers) so an aliased inner object can't mutate the cache key.
|
||||
snapshot = {k: dict(v) if isinstance(v, dict) else v for k, v in request_kwargs.items()}
|
||||
self._store_request_slot(_OPENAI_SLOT, client, snapshot)
|
||||
return client
|
||||
|
||||
def _close_request_openai_client(self, client: Any, *, reason: str) -> None:
|
||||
if not self._release_request_slot(_OPENAI_SLOT, client, reason):
|
||||
self._close_openai_client(client, reason=reason, shared=False)
|
||||
|
||||
def _close_cached_request_openai_client(self, *, reason: str) -> None:
|
||||
"""Teardown hook: really close the cached per-request wire client."""
|
||||
client, in_use = self._take_request_slot(_OPENAI_SLOT)
|
||||
if client is None:
|
||||
return
|
||||
if in_use:
|
||||
# Checked out by a worker: close() here would release FDs from a stranger thread. Abort the
|
||||
# sockets; the worker's own finally does the real close.
|
||||
self._abort_request_openai_client(client, reason=f"{reason}_in_flight")
|
||||
else:
|
||||
self._close_openai_client(client, reason=reason, shared=False)
|
||||
|
||||
def _abort_request_openai_client(self, client: Any, *, reason: str) -> None:
|
||||
self._abort_request_slot_client(_OPENAI_SLOT, client, reason=reason)
|
||||
|
||||
def _request_anthropic_client_key(self) -> tuple:
|
||||
"""Cache key over everything forcing a fresh client: credential, base URL/region, timeout, 1M-beta flag."""
|
||||
if getattr(self, "provider", None) == "bedrock":
|
||||
return ("bedrock", getattr(self, "_bedrock_region", "us-east-1") or "us-east-1")
|
||||
return (
|
||||
"direct", self._anthropic_api_key, getattr(self, "_anthropic_base_url", None),
|
||||
get_provider_request_timeout(self.provider, self.model), bool(getattr(self, "_oauth_1m_beta_disabled", False)),
|
||||
)
|
||||
|
||||
def _build_direct_anthropic_client(self, token: str, base_url: Any) -> Any:
|
||||
"""Native Anthropic client for ``token``/``base_url`` with the provider/model request timeout."""
|
||||
from agent.anthropic_adapter import build_anthropic_client
|
||||
return build_anthropic_client(token, base_url, timeout=get_provider_request_timeout(self.provider, self.model))
|
||||
|
||||
def _anthropic_oauth_flag(self, token: str) -> bool:
|
||||
"""OAuth flag only on native Anthropic; third-party Anthropic-protocol endpoints must not trip OAuth paths."""
|
||||
from agent.anthropic_credentials import _is_oauth_token
|
||||
return _is_oauth_token(token) if self.provider == "anthropic" else False
|
||||
|
||||
def _build_anthropic_client_for_key(self, key: tuple) -> Any:
|
||||
from agent.anthropic_adapter import build_anthropic_bedrock_client, build_anthropic_client
|
||||
if key[0] == "bedrock":
|
||||
return build_anthropic_bedrock_client(key[1])
|
||||
return build_anthropic_client(key[1], key[2], timeout=key[3], drop_context_1m_beta=key[4])
|
||||
|
||||
def _create_request_anthropic_client(self, *, reason: str) -> Any:
|
||||
"""Build (or reuse) a request-local Anthropic client for one in-flight call.
|
||||
|
||||
The watchdog must never ``close()`` a client a worker is still reading (fd recycled under a live SSL BIO →
|
||||
TLS record in a SQLite header); per-request clients let the stranger ``shutdown()`` while the owner closes.
|
||||
"""
|
||||
if self.api_mode == "anthropic_messages":
|
||||
self._try_refresh_anthropic_client_credentials()
|
||||
key = self._request_anthropic_client_key()
|
||||
cached, stale = self._checkout_request_slot(_ANTHROPIC_SLOT, key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
if stale is not None:
|
||||
self._close_request_anthropic_client(stale, reason=f"reuse_evict:{reason}")
|
||||
client = self._build_anthropic_client_for_key(key)
|
||||
logger.debug("Anthropic request client created (%s, shared=False) %s", reason, self._anthropic_log_context())
|
||||
self._store_request_slot(_ANTHROPIC_SLOT, client, key)
|
||||
return client
|
||||
|
||||
def _close_request_anthropic_client(self, client: Any, *, reason: str) -> None:
|
||||
"""Owner-thread close: clean finish keeps the pool warm; otherwise force-close sockets (CLOSE-WAIT) + SDK close."""
|
||||
if client is None or self._release_request_slot(_ANTHROPIC_SLOT, client, reason):
|
||||
return
|
||||
try:
|
||||
self._force_close_tcp_sockets(client)
|
||||
client.close()
|
||||
logger.info("Anthropic client closed (%s, shared=False) %s", reason, self._anthropic_log_context())
|
||||
except Exception as exc:
|
||||
logger.debug(
|
||||
"Anthropic client close failed (%s, shared=False) %s error=%s", reason, self._anthropic_log_context(), exc,
|
||||
)
|
||||
|
||||
def _close_cached_request_anthropic_client(self, *, reason: str) -> None:
|
||||
"""Teardown hook: really close the cached per-request Anthropic client."""
|
||||
client, in_use = self._take_request_slot(_ANTHROPIC_SLOT)
|
||||
if client is None:
|
||||
return
|
||||
if in_use: # checked out by a worker — same reasoning as the OpenAI teardown hook
|
||||
self._abort_request_anthropic_client(client, reason=f"{reason}_in_flight")
|
||||
return
|
||||
with suppress(Exception):
|
||||
self._force_close_tcp_sockets(client)
|
||||
client.close()
|
||||
|
||||
def _abort_request_anthropic_client(self, client: Any, *, reason: str) -> None:
|
||||
self._abort_request_slot_client(_ANTHROPIC_SLOT, client, reason=reason)
|
||||
|
||||
# ------------------------------------------------------------------ credential refresh
|
||||
def _sync_client_kwargs_credentials(self) -> None:
|
||||
"""Mirror ``self.api_key`` / ``self.base_url`` into the OpenAI-style client kwargs."""
|
||||
self._client_kwargs["api_key"] = self.api_key
|
||||
self._client_kwargs["base_url"] = self.base_url
|
||||
|
||||
def _adopt_openai_credentials(self, api_key: str, base_url: str, *, reason: str) -> bool:
|
||||
"""Apply a fresh key/base_url to the OpenAI-style kwargs and rebuild the shared client."""
|
||||
self.api_key, self.base_url = api_key.strip(), base_url.strip().rstrip("/")
|
||||
self._sync_client_kwargs_credentials()
|
||||
return self._replace_primary_openai_client(reason=reason)
|
||||
|
||||
def _try_refresh_codex_client_credentials(self, *, force: bool = True) -> bool:
|
||||
if self.api_mode != "codex_responses" or self.provider not in {"openai-codex", "xai-oauth"}:
|
||||
return False
|
||||
# No silent account swap: a non-singleton credential (manual pool entry, explicit api_key=) must not be
|
||||
# replaced by the device_code singleton's tokens — the pool's reactive recovery owns that case.
|
||||
try:
|
||||
from hermes_cli import auth as _auth
|
||||
resolve = (
|
||||
_auth.resolve_codex_runtime_credentials if self.provider == "openai-codex"
|
||||
else _auth.resolve_xai_oauth_runtime_credentials
|
||||
)
|
||||
singleton_now = resolve(refresh_if_expiring=False)
|
||||
except Exception as exc:
|
||||
logger.debug("%s singleton read failed: %s", self.provider, exc)
|
||||
return False
|
||||
singleton_key = str(singleton_now.get("api_key") or "").strip()
|
||||
old_key = str(self.api_key or "").strip()
|
||||
if singleton_key and old_key and singleton_key != old_key:
|
||||
logger.debug(
|
||||
"%s singleton tokens differ from the active api_key; skipping singleton force-refresh to avoid "
|
||||
"silent account swap. Reactive credential rotation should go through the pool.", self.provider,
|
||||
)
|
||||
return False
|
||||
try:
|
||||
creds = resolve(force_refresh=force)
|
||||
except Exception as exc:
|
||||
logger.debug("%s credential refresh failed: %s", self.provider, exc)
|
||||
return False
|
||||
api_key, base_url = creds.get("api_key"), creds.get("base_url")
|
||||
if not _valid_credential_pair(api_key, base_url):
|
||||
return False
|
||||
# No NEW token minted (the resolver returns the same stale token when refresh fails) → False.
|
||||
if old_key and api_key.strip() == old_key:
|
||||
logger.debug("%s credential refresh returned the same token; refresh likely failed silently", self.provider)
|
||||
return False
|
||||
return self._adopt_openai_credentials(api_key, base_url, reason=f"{self.provider}_credential_refresh")
|
||||
|
||||
def _try_refresh_nous_client_credentials(self, *, force: bool = True, require_account: str | None = None) -> bool:
|
||||
# Portal serves anthropic/* on the native Messages route, so either client kind may hold the expiring JWT.
|
||||
if self.provider != "nous" or self.api_mode not in ("chat_completions", "anthropic_messages"):
|
||||
return False
|
||||
try:
|
||||
from hermes_cli.auth import resolve_nous_runtime_credentials
|
||||
timeout = env_float("HERMES_NOUS_TIMEOUT_SECONDS", 15)
|
||||
# Pass the bearer that just 401'd so a refresh already done by a sibling process is
|
||||
# adopted instead of rotating the grant again.
|
||||
creds = resolve_nous_runtime_credentials(
|
||||
timeout_seconds=timeout, force_refresh=force, stale_access_token=self.api_key or None,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("Nous credential refresh failed: %s", exc)
|
||||
return False
|
||||
api_key, base_url = creds.get("api_key"), creds.get("base_url")
|
||||
if not _valid_credential_pair(api_key, base_url):
|
||||
return False
|
||||
if str(api_key).strip() == str(self.api_key or "").strip():
|
||||
return False # store holds the same key: nothing to adopt, no client rebuild
|
||||
if require_account is not None:
|
||||
try:
|
||||
from hermes_cli.auth_constants import _decode_jwt_claims
|
||||
new_account = _decode_jwt_claims(str(api_key)).get("sub")
|
||||
except Exception:
|
||||
new_account = None
|
||||
if str(new_account or "") != require_account:
|
||||
logger.info(
|
||||
"Nous pre-expiry adoption skipped: the store's key belongs to a different account "
|
||||
"than the one in hand; keeping the current credential."
|
||||
)
|
||||
return False
|
||||
if self.api_mode == "anthropic_messages":
|
||||
self.api_key, self.base_url = api_key.strip(), base_url.strip().rstrip("/")
|
||||
self._anthropic_api_key, self._anthropic_base_url = self.api_key, self.base_url
|
||||
self._rebuild_anthropic_client()
|
||||
return True
|
||||
# Nous requests should not inherit OpenRouter-only attribution headers.
|
||||
self._client_kwargs.pop("default_headers", None)
|
||||
return self._adopt_openai_credentials(api_key, base_url, reason="nous_credential_refresh")
|
||||
|
||||
# Adopt a fresh key this many seconds before the one in hand expires. Wider than the store's
|
||||
# own refresh skew (120 s) so the keepalive has normally already minted the replacement.
|
||||
_NOUS_KEY_ADOPT_SKEW_S = 180
|
||||
|
||||
def _adopt_nous_key_before_expiry(self) -> bool:
|
||||
"""Swap in a fresh Nous agent key BEFORE the one in hand expires, so the request never 401s.
|
||||
|
||||
The agent key is a JWT; its ``exp`` is read locally (no network). Inside the skew the store
|
||||
is re-read under the auth-store lock: the keepalive thread normally holds a fresh key already
|
||||
(adopt, no POST), otherwise ONE refresh runs and every peer adopts its result. Before this,
|
||||
every agent in a process learned about the hourly expiry from its own 401, all in the same
|
||||
minute (620 in one 200-subagent run), and the pool benched the sole credential for all of
|
||||
them. Returns True when a new key was adopted.
|
||||
|
||||
Identity guard: the replacement must belong to the SAME account (``sub`` claim) as the key
|
||||
in hand. The store holds the logged-in singleton; an agent running on an explicitly supplied
|
||||
or pool-selected key for a different account must never be silently moved onto it (that
|
||||
changes who is billed). When either side lacks a ``sub`` nothing is adopted here; the
|
||||
reactive 401 path is unchanged.
|
||||
"""
|
||||
if getattr(self, "provider", "") != "nous" or not getattr(self, "api_key", None):
|
||||
return False
|
||||
try:
|
||||
from hermes_cli.auth_constants import _decode_jwt_claims
|
||||
claims = _decode_jwt_claims(self.api_key)
|
||||
except Exception:
|
||||
return False
|
||||
exp, account = claims.get("exp"), claims.get("sub")
|
||||
if not account or not isinstance(exp, (int, float)) or exp - time.time() > self._NOUS_KEY_ADOPT_SKEW_S:
|
||||
return False
|
||||
return self._try_refresh_nous_client_credentials(force=False, require_account=str(account))
|
||||
|
||||
|
||||
def _resolve_env_credentials(self) -> Optional[tuple]:
|
||||
"""Current ``.env``-sourced ``(api_key, base_url, default_base)`` for this provider, or ``None``.
|
||||
|
||||
Covers registry api-key providers and named custom providers with ``key_env``.
|
||||
"""
|
||||
try:
|
||||
from agent.credential_pool import get_env_prefer_dotenv
|
||||
from hermes_cli.auth import PROVIDER_REGISTRY
|
||||
except ImportError:
|
||||
return None
|
||||
pconfig = PROVIDER_REGISTRY.get(self.provider)
|
||||
if pconfig and getattr(pconfig, "auth_type", "") == "api_key" and getattr(pconfig, "api_key_env_vars", ()):
|
||||
# First non-empty env var wins (lazy: later vars are not read).
|
||||
api_key = next((k for k in (get_env_prefer_dotenv(v).strip() for v in pconfig.api_key_env_vars) if k), "")
|
||||
if not api_key:
|
||||
return None
|
||||
url_var = pconfig.base_url_env_var
|
||||
env_url = get_env_prefer_dotenv(url_var).strip().rstrip("/") if url_var else ""
|
||||
default_base = (pconfig.inference_base_url or "").strip().rstrip("/")
|
||||
base_url = env_url or default_base
|
||||
if self.provider in ("kimi-coding", "zai"):
|
||||
from hermes_cli import auth as _auth
|
||||
resolver = _auth._resolve_kimi_base_url if self.provider == "kimi-coding" else _auth._resolve_zai_base_url
|
||||
base_url = resolver(api_key, pconfig.inference_base_url, env_url).rstrip("/")
|
||||
elif self.provider == "custom":
|
||||
# Named custom provider: identity in config, credential in key_env; no key_env → nothing to watch.
|
||||
try:
|
||||
from hermes_cli.runtime_provider import _get_named_custom_provider
|
||||
except ImportError:
|
||||
return None
|
||||
custom_provider = _get_named_custom_provider(getattr(self, "requested_provider", "") or "")
|
||||
key_env = str((custom_provider or {}).get("key_env") or "").strip()
|
||||
api_key = get_env_prefer_dotenv(key_env).strip() if key_env else ""
|
||||
if not custom_provider or not api_key:
|
||||
return None
|
||||
# Custom providers pin base_url in config, so only key edits are adopted here.
|
||||
base_url = default_base = str(custom_provider.get("base_url") or "").strip().rstrip("/")
|
||||
else:
|
||||
return None
|
||||
if not base_url:
|
||||
return None
|
||||
return api_key, base_url, default_base
|
||||
|
||||
def _should_adopt_env_credentials(self, api_key: str, base_url: str, default_base: str) -> bool:
|
||||
"""Adopt only env *edits* (resolved value changed since last look), never drift from ``self.*``:
|
||||
pool rotation/failover and a config ``model.base_url`` legitimately move the session."""
|
||||
prev = getattr(self, "_env_creds_seen", None)
|
||||
current_base = (self.base_url or "").strip().rstrip("/")
|
||||
unchanged = base_url == current_base and api_key == self.api_key
|
||||
if prev is None:
|
||||
# First look: adopt only the boot-default case (anything else is unattributable on turn one), and
|
||||
# never stomp a pool-rotated key.
|
||||
pool_rotated = (
|
||||
api_key != self.api_key and getattr(self, "_credential_pool", None) is not None
|
||||
and getattr(self, "_credential_pool_entry_id", None)
|
||||
)
|
||||
return current_base == default_base and not unchanged and not pool_rotated
|
||||
# Adopt only while the session still runs on the registry default or the previously-seen env value.
|
||||
return (base_url, api_key) != prev and current_base in {default_base, prev[0]} and not unchanged
|
||||
|
||||
def _try_refresh_env_client_credentials(self) -> bool:
|
||||
"""Adopt ``~/.hermes/.env`` credential/base-url edits at the turn boundary (a Settings save updates ``.env``
|
||||
but a live worker keeps init-time values). Adoption rule: ``_should_adopt_env_credentials``.
|
||||
|
||||
Covers api-key registry providers and named custom providers with a ``key_env`` (#67935) — the
|
||||
latter resolve to ``provider="custom"`` with no registry entry, so they are matched through the
|
||||
runtime provider's config lookup instead.
|
||||
"""
|
||||
if self.api_mode != "chat_completions" or getattr(self, "_fallback_activated", False):
|
||||
return False
|
||||
resolved = self._resolve_env_credentials()
|
||||
if resolved is None:
|
||||
return False
|
||||
api_key, base_url, default_base = resolved
|
||||
if not self._should_adopt_env_credentials(api_key, base_url, default_base):
|
||||
self._env_creds_seen = (base_url, api_key)
|
||||
return False
|
||||
from hermes_cli.route_identity import normalize_route_base_url
|
||||
route_changed = normalize_route_base_url(self.base_url) != normalize_route_base_url(base_url)
|
||||
prior_api_key, prior_base_url = self.api_key, self.base_url
|
||||
prior_client_kwargs = dict(self._client_kwargs)
|
||||
self.api_key, self.base_url = api_key, base_url
|
||||
self._sync_client_kwargs_credentials()
|
||||
# A base-url change moves the route: recompute TLS material and default headers.
|
||||
self._reapply_route_client_config(route_changed=route_changed)
|
||||
if not self._replace_primary_openai_client(reason="env_credential_refresh"):
|
||||
# Leave the baseline un-advanced (retry next turn); roll the agent back to match the live client.
|
||||
self.api_key, self.base_url = prior_api_key, prior_base_url
|
||||
self._client_kwargs.clear()
|
||||
self._client_kwargs.update(prior_client_kwargs)
|
||||
return False
|
||||
# Rebind the pool entry id to the adopted key, or the next 429 quarantines the wrong credential.
|
||||
try:
|
||||
from agent.agent_runtime_helpers import sync_credential_pool_entry_id
|
||||
sync_credential_pool_entry_id(self)
|
||||
except Exception:
|
||||
logger.debug("sync_credential_pool_entry_id after env refresh failed", exc_info=True)
|
||||
self._env_creds_seen = (base_url, api_key)
|
||||
logger.info("Applied updated .env credentials for %s: endpoint %s", self.provider, self.base_url)
|
||||
return True
|
||||
|
||||
def _try_refresh_vertex_client_credentials(self) -> bool:
|
||||
"""Re-mint the Vertex OAuth2 token (~1h TTL; long sessions 401 on the expired bearer) and rebuild the client."""
|
||||
if self.api_mode != "chat_completions" or self.provider != "vertex":
|
||||
return False
|
||||
try:
|
||||
from agent.vertex_adapter import get_vertex_config
|
||||
token, base_url = get_vertex_config()
|
||||
except Exception as exc:
|
||||
logger.debug("Vertex credential refresh failed: %s", exc)
|
||||
return False
|
||||
ok = _valid_credential_pair(token, base_url) and self._adopt_openai_credentials(
|
||||
token, base_url, reason="vertex_credential_refresh",
|
||||
)
|
||||
if ok:
|
||||
logger.info("Vertex AI OAuth token refreshed")
|
||||
return ok
|
||||
|
||||
def _apply_copilot_token(self, token: str, enterprise_base_url: Any, *, reason: str) -> bool:
|
||||
self.api_key = token
|
||||
if enterprise_base_url:
|
||||
self.base_url = enterprise_base_url.rstrip("/")
|
||||
self._sync_client_kwargs_credentials()
|
||||
self._apply_client_headers_for_base_url(str(self.base_url or ""))
|
||||
return self._replace_primary_openai_client(reason=reason)
|
||||
|
||||
def _try_refresh_copilot_client_credentials(self) -> bool:
|
||||
"""Refresh Copilot credentials and rebuild the shared OpenAI client (caller enforces the single-shot guard).
|
||||
|
||||
The raw GitHub token is stable; the short-TTL *exchanged* IDE token is what expires mid-turn (``401 IDE
|
||||
token expired``), so force a fresh exchange rather than re-resolving the raw token.
|
||||
"""
|
||||
if not self._is_copilot_provider():
|
||||
return False
|
||||
try:
|
||||
from hermes_cli.copilot_auth import resolve_copilot_token, get_copilot_api_token, evict_cached_exchanged_token
|
||||
new_token, token_source = resolve_copilot_token()
|
||||
except Exception as exc:
|
||||
logger.debug("Copilot credential refresh failed: %s", exc)
|
||||
return False
|
||||
new_token, enterprise_base_url = (new_token.strip() if isinstance(new_token, str) else ""), None
|
||||
if not new_token:
|
||||
return False
|
||||
# Fall back to the raw token only if the exchange itself is unavailable.
|
||||
try:
|
||||
evict_cached_exchanged_token(new_token)
|
||||
api_token, exchanged_base_url = get_copilot_api_token(new_token)
|
||||
if isinstance(api_token, str) and api_token.strip():
|
||||
new_token, enterprise_base_url = api_token.strip(), exchanged_base_url
|
||||
except Exception as exc:
|
||||
logger.debug("Copilot 401 re-exchange failed, using resolved token: %s", exc)
|
||||
ok = self._apply_copilot_token(new_token, enterprise_base_url, reason="copilot_credential_refresh")
|
||||
if ok:
|
||||
logger.info("Copilot credentials refreshed from %s", token_source)
|
||||
return ok
|
||||
|
||||
def _try_recover_stale_copilot_credential(self) -> bool:
|
||||
"""Force a fresh Copilot token exchange + client rebuild after a 400 (single-shot, caller-guarded).
|
||||
|
||||
Copilot surfaces a stale credential as ``400 model_not_available_for_integrator`` / ``model_not_supported``
|
||||
— typically a raw ``ghu_`` token cached when the startup exchange degraded (restricted integrator allowlist).
|
||||
"""
|
||||
if not self._is_copilot_provider():
|
||||
return False
|
||||
try:
|
||||
from hermes_cli.copilot_auth import resolve_copilot_token, get_copilot_api_token, evict_cached_exchanged_token
|
||||
raw_token, token_source = resolve_copilot_token()
|
||||
if not isinstance(raw_token, str) or not raw_token.strip():
|
||||
return False
|
||||
raw_token = raw_token.strip()
|
||||
# Evict the cached (possibly degraded/raw) exchanged token so the exchange hits the network.
|
||||
evict_cached_exchanged_token(raw_token)
|
||||
api_token, enterprise_base_url = get_copilot_api_token(raw_token)
|
||||
except Exception as exc:
|
||||
logger.debug("Copilot stale-credential recovery failed: %s", exc)
|
||||
return False
|
||||
if not isinstance(api_token, str) or not api_token.strip():
|
||||
return False
|
||||
# Exchange STILL degraded to the raw token: a rebuild won't help — don't burn the single-shot retry.
|
||||
if api_token == raw_token and not enterprise_base_url:
|
||||
logger.warning(
|
||||
"Copilot stale-credential recovery: exchange still degraded to raw token; skipping retry "
|
||||
"(network/exchange endpoint unavailable)."
|
||||
)
|
||||
return False
|
||||
ok = self._apply_copilot_token(api_token.strip(), enterprise_base_url, reason="copilot_stale_credential_recovery")
|
||||
if ok:
|
||||
logger.info("Copilot credentials re-exchanged after stale-credential 400 (source=%s)", token_source)
|
||||
return ok
|
||||
|
||||
def _try_refresh_anthropic_client_credentials(self) -> bool:
|
||||
# Only native Anthropic rotates OAuth tokens; other anthropic_messages providers (MiniMax, Alibaba, ...)
|
||||
# and Azure use static keys — a refresh would pick up the ~/.claude OAuth token and break auth.
|
||||
if (
|
||||
self.api_mode != "anthropic_messages" or not hasattr(self, "_anthropic_api_key")
|
||||
or self.provider != "anthropic"
|
||||
or base_url_host_matches(getattr(self, "_anthropic_base_url", "") or "", "azure.com")
|
||||
):
|
||||
return False
|
||||
try:
|
||||
from agent.anthropic_credentials import resolve_anthropic_token
|
||||
new_token = resolve_anthropic_token()
|
||||
except Exception as exc:
|
||||
logger.debug("Anthropic credential refresh failed: %s", exc)
|
||||
return False
|
||||
new_token = new_token.strip() if isinstance(new_token, str) else ""
|
||||
if not new_token or new_token == self._anthropic_api_key:
|
||||
return False
|
||||
with suppress(Exception):
|
||||
self._anthropic_client.close()
|
||||
try:
|
||||
base_url = getattr(self, "_anthropic_base_url", None)
|
||||
self._anthropic_client = self._build_direct_anthropic_client(new_token, base_url)
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to rebuild Anthropic client after credential refresh: %s", exc)
|
||||
return False
|
||||
self._anthropic_api_key, self._is_anthropic_oauth = new_token, self._anthropic_oauth_flag(new_token)
|
||||
return True
|
||||
|
||||
# ------------------------------------------------------------------ route-derived client config
|
||||
def _apply_client_headers_for_base_url(self, base_url: str, *, apply_user_headers: bool = True) -> None:
|
||||
for host, build in _ROUTE_DEFAULT_HEADERS:
|
||||
if base_url_host_matches(base_url, host):
|
||||
self._client_kwargs["default_headers"] = build(self, base_url)
|
||||
break
|
||||
else:
|
||||
# No URL-specific headers — fall back to profile.default_headers, else clear.
|
||||
self._client_kwargs.pop("default_headers", None)
|
||||
with suppress(Exception):
|
||||
from providers import get_provider_profile
|
||||
profile = get_provider_profile(self.provider)
|
||||
if profile and profile.default_headers and (profile_headers := dict(profile.default_headers)):
|
||||
self._client_kwargs["default_headers"] = profile_headers
|
||||
# User overrides win over URL/profile defaults for the same route; a swap to another endpoint must not
|
||||
# inherit them.
|
||||
if apply_user_headers:
|
||||
self._apply_user_default_headers()
|
||||
# Per-provider extra_headers last so they survive swaps/rebuilds. SECURITY: may carry credentials; never log.
|
||||
if self.api_mode not in ("anthropic_messages", "bedrock_converse"):
|
||||
try:
|
||||
from hermes_cli.config import apply_custom_provider_extra_headers_to_client_kwargs
|
||||
apply_custom_provider_extra_headers_to_client_kwargs(self._client_kwargs, base_url)
|
||||
except Exception:
|
||||
logger.debug("custom-provider extra_headers skipped", exc_info=True)
|
||||
|
||||
def _apply_user_default_headers(self) -> None:
|
||||
"""Merge config ``model.default_headers`` onto the OpenAI client (user wins; WAFs rejecting SDK headers).
|
||||
Delegates to ``agent.auxiliary_client`` so main and aux clients cannot drift. No-op for Anthropic/Bedrock."""
|
||||
if self.api_mode in ("anthropic_messages", "bedrock_converse"):
|
||||
return
|
||||
from agent.auxiliary_client import _apply_user_default_headers as _merge_user_headers
|
||||
merged = _merge_user_headers(self._client_kwargs.get("default_headers"))
|
||||
if merged:
|
||||
self._client_kwargs["default_headers"] = merged
|
||||
|
||||
def _swap_credential(self, entry) -> None:
|
||||
runtime_key = getattr(entry, "runtime_api_key", None) or getattr(entry, "access_token", "")
|
||||
runtime_base = getattr(entry, "runtime_base_url", None) or getattr(entry, "base_url", None) or self.base_url
|
||||
self._credential_pool_entry_id = getattr(entry, "id", None)
|
||||
from hermes_cli.route_identity import normalize_route_base_url
|
||||
route_changed = normalize_route_base_url(self.base_url) != normalize_route_base_url(runtime_base)
|
||||
stripped_base = runtime_base.rstrip("/") if isinstance(runtime_base, str) else runtime_base
|
||||
if self.api_mode == "anthropic_messages":
|
||||
with suppress(Exception):
|
||||
self._anthropic_client.close()
|
||||
self._anthropic_api_key, self._anthropic_base_url = runtime_key, stripped_base
|
||||
self._anthropic_client = self._build_direct_anthropic_client(runtime_key, self._anthropic_base_url)
|
||||
self._is_anthropic_oauth = self._anthropic_oauth_flag(runtime_key)
|
||||
self.api_key, self.base_url = runtime_key, stripped_base
|
||||
return
|
||||
self.api_key, self.base_url = runtime_key, stripped_base
|
||||
# Inlined (not _sync_client_kwargs_credentials): tests call this unbound on a SimpleNamespace agent.
|
||||
self._client_kwargs["api_key"] = self.api_key
|
||||
self._client_kwargs["base_url"] = self.base_url
|
||||
self._reapply_route_client_config(route_changed=route_changed)
|
||||
self._replace_primary_openai_client(reason="credential_rotation")
|
||||
|
||||
def _reapply_route_client_config(self, *, route_changed: bool) -> None:
|
||||
"""Recompute route-derived client kwargs (TLS material, default headers) for ``self.base_url``.
|
||||
|
||||
Any rebuild that may have moved ``base_url`` must call this or the new endpoint inherits the old config.
|
||||
"""
|
||||
self._client_kwargs.pop("ssl_verify", None)
|
||||
self._client_kwargs.pop("ssl_ca_cert", None)
|
||||
try:
|
||||
from hermes_cli.config import (
|
||||
apply_custom_provider_tls_to_client_kwargs, get_compatible_custom_providers, load_config_readonly,
|
||||
)
|
||||
apply_custom_provider_tls_to_client_kwargs(
|
||||
self._client_kwargs, str(self.base_url or ""), get_compatible_custom_providers(load_config_readonly()),
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("custom-provider TLS resolution skipped on credential rotation", exc_info=True)
|
||||
self._apply_client_headers_for_base_url(self.base_url, apply_user_headers=not route_changed)
|
||||
|
||||
def _anthropic_messages_create(self, api_kwargs: dict, *, client: Any = None):
|
||||
# A supplied request-local client was already refreshed in _create_request_anthropic_client.
|
||||
if client is None and self.api_mode == "anthropic_messages":
|
||||
self._try_refresh_anthropic_client_credentials()
|
||||
# Strips Responses-only kwargs that leak in under an api_mode-flip race.
|
||||
from agent.anthropic_adapter import create_anthropic_message
|
||||
# on_response: rate-limit + credits state live in response headers, which the parsed Message drops.
|
||||
return create_anthropic_message(
|
||||
client or self._anthropic_client, api_kwargs, log_prefix=getattr(self, "log_prefix", ""),
|
||||
prefer_stream=not bool(getattr(self, "_disable_streaming", False)),
|
||||
on_response=self._capture_anthropic_response_headers,
|
||||
)
|
||||
|
||||
def _rebuild_anthropic_client(self) -> None:
|
||||
"""Rebuild the Anthropic client after an interrupt/stale call (Bedrock SDK for bedrock; honors 1M-beta flag)."""
|
||||
self._anthropic_client = self._build_anthropic_client_for_key(self._request_anthropic_client_key())
|
||||
@@ -0,0 +1,75 @@
|
||||
"""Codex request identity helpers shared by agent client builders.
|
||||
|
||||
Leaf module with no dependency on the large auxiliary-client router, so a
|
||||
long-lived process can import a newly added client builder without resolving a
|
||||
new symbol from an older cached ``auxiliary_client`` module.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
from typing import Any, Dict
|
||||
from urllib.parse import urlparse
|
||||
|
||||
|
||||
CODEX_AUX_BASE_URL = "https://chatgpt.com/backend-api/codex"
|
||||
|
||||
|
||||
def is_official_codex_base_url(base_url: str) -> bool:
|
||||
"""Identify OpenAI's Codex endpoint without matching custom proxies."""
|
||||
try:
|
||||
parsed = urlparse(base_url)
|
||||
path = parsed.path.rstrip("/")
|
||||
return (
|
||||
parsed.scheme == "https"
|
||||
and parsed.hostname == "chatgpt.com"
|
||||
and parsed.port in (None, 443)
|
||||
and (path == "/backend-api/codex" or path.startswith("/backend-api/codex/"))
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
|
||||
|
||||
def codex_cloudflare_headers(access_token: str, *, base_url: str = CODEX_AUX_BASE_URL) -> Dict[str, str]:
|
||||
"""Identity and account headers for chatgpt.com/backend-api/codex.
|
||||
|
||||
OpenAI requires third-party harnesses to identify themselves: the official
|
||||
endpoint gets Hermes' originator and version, custom endpoints keep the
|
||||
codex_cli_rs compatibility identity. ``ChatGPT-Account-ID`` comes from the
|
||||
OAuth JWT's ``chatgpt_account_id`` claim; a malformed token drops the header
|
||||
rather than raising, so it surfaces as a 401 instead of a crash at client
|
||||
construction.
|
||||
"""
|
||||
if is_official_codex_base_url(base_url):
|
||||
from hermes_cli import __version__
|
||||
headers = {"User-Agent": f"HermesAgent/{__version__}", "originator": "hermes-agent"}
|
||||
else:
|
||||
headers = {"User-Agent": "codex_cli_rs/0.0.0 (Hermes Agent)", "originator": "codex_cli_rs"}
|
||||
if not isinstance(access_token, str) or not access_token.strip():
|
||||
return headers
|
||||
try:
|
||||
parts = access_token.split(".")
|
||||
if len(parts) < 2:
|
||||
return headers
|
||||
payload_b64 = parts[1] + "=" * (-len(parts[1]) % 4)
|
||||
claims = json.loads(base64.urlsafe_b64decode(payload_b64))
|
||||
acct_id = claims.get("https://api.openai.com/auth", {}).get("chatgpt_account_id")
|
||||
if isinstance(acct_id, str) and acct_id:
|
||||
headers["ChatGPT-Account-ID"] = acct_id
|
||||
except Exception:
|
||||
pass
|
||||
return headers
|
||||
|
||||
|
||||
def apply_required_codex_headers(client_kwargs: Dict[str, Any], *, access_token: str, base_url: str) -> None:
|
||||
"""Keep required Codex identity after user/provider header overrides."""
|
||||
if not is_official_codex_base_url(base_url):
|
||||
return
|
||||
required = codex_cloudflare_headers(access_token, base_url=base_url)
|
||||
required_names = {name.lower() for name in required}
|
||||
existing = client_kwargs.get("default_headers") or {}
|
||||
client_kwargs["default_headers"] = {
|
||||
**{name: value for name, value in existing.items() if str(name).lower() not in required_names},
|
||||
**required,
|
||||
}
|
||||
+830
-1636
File diff suppressed because it is too large
Load Diff
+774
-1573
File diff suppressed because it is too large
Load Diff
+342
-683
File diff suppressed because it is too large
Load Diff
@@ -1,35 +1,12 @@
|
||||
"""Mint a provider API key by running a command (``key_cmd``).
|
||||
|
||||
Static API keys are the exception at enterprise gateways: SSO/OIDC brokers,
|
||||
cloud IAM, and internal auth proxies all issue SHORT-LIVED bearers instead.
|
||||
A key copied into ``.env`` (``key_env``) is stale within the hour, so every
|
||||
request after that 401s and the user has to restart the session.
|
||||
|
||||
``key_cmd`` names a command that PRINTS a token, so the credential is derived
|
||||
rather than stored::
|
||||
|
||||
providers:
|
||||
my-gateway:
|
||||
base_url: https://gateway.internal.example.com/v1
|
||||
api_mode: chat_completions
|
||||
key_cmd: my-auth-cli print-token --profile prod
|
||||
|
||||
This is the established pattern for agent tooling — Claude Code's
|
||||
``apiKeyHelper``, the ``gcloud auth print-access-token`` / ``aws ecr
|
||||
get-login-password`` idiom, and vendor helpers such as ``databricks auth
|
||||
token`` all expose exactly this contract. Hermes already accepts a callable
|
||||
API key on both wire clients (the Entra ID / Azure identity path) and invokes
|
||||
it per request, so nothing downstream changes: the token is simply always
|
||||
fresh. It is cached until shortly before expiry, so the command runs about
|
||||
once per token lifetime rather than once per request.
|
||||
|
||||
Output contract: print ONLY the token on stdout, either bare or as JSON with
|
||||
an ``access_token`` field (``expires_in`` is honoured when present) — the
|
||||
shape OAuth 2.0 token endpoints and the helpers above already emit.
|
||||
|
||||
Precedence: an explicit ``--api-key`` still wins (the one-off recovery escape
|
||||
hatch); otherwise ``key_cmd`` is preferred over a static ``api_key`` /
|
||||
``key_env`` on the same entry.
|
||||
Enterprise gateways (SSO/OIDC brokers, cloud IAM, auth proxies) issue SHORT-LIVED bearers; a key
|
||||
copied into ``.env`` goes stale within the hour. ``key_cmd`` names a command that PRINTS a token
|
||||
(the ``apiKeyHelper`` / ``gcloud auth print-access-token`` idiom). Both wire clients accept a
|
||||
callable API key and invoke it per request; the token is cached until shortly before expiry.
|
||||
Output contract: ONLY the token on stdout, bare or as JSON with an ``access_token`` field
|
||||
(``expires_in`` / ISO ``expiry`` honoured). Precedence: explicit ``--api-key`` wins (one-off
|
||||
recovery escape hatch); otherwise ``key_cmd`` beats a static ``api_key`` / ``key_env``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -39,23 +16,17 @@ import logging
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
from typing import Callable, Optional
|
||||
from typing import Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Treat a cached token as spent slightly before its stated expiry, so a request
|
||||
# can't be signed with a token that dies in flight. 60s matches the leeway used
|
||||
# by comparable OAuth token caches.
|
||||
# Treat a token as spent slightly before expiry so a request can't be signed with one that dies in
|
||||
# flight (60s = usual OAuth cache leeway).
|
||||
_TOKEN_REFRESH_LEEWAY_SECONDS = 60.0
|
||||
# A token helper reads a local credential cache and should answer in
|
||||
# milliseconds; anything approaching this budget is hung, not slow.
|
||||
# Helpers answer from a local cache in milliseconds; this long means hung.
|
||||
_MINT_TIMEOUT_SECONDS = 15
|
||||
# When a helper advertises NO expiry, the token cannot be cached for the life
|
||||
# of the process: nothing in the request path re-mints on 401 (the SDK retries
|
||||
# 429/5xx only), so an expired no-TTL token would 401 every request until
|
||||
# restart. Re-mint on a bounded window instead — the helper answers from a
|
||||
# local credential cache in milliseconds, so a periodic re-run is cheap, and a
|
||||
# helper that wants a longer cache can simply advertise its real expiry.
|
||||
# No advertised expiry: nothing in the request path re-mints on 401 (the SDK retries 429/5xx only), so
|
||||
# a process-lifetime cache would 401 forever once the token died. Re-mint on a bounded window instead.
|
||||
_NO_TTL_REFRESH_SECONDS = 900.0
|
||||
|
||||
|
||||
@@ -63,33 +34,31 @@ class CommandTokenError(RuntimeError):
|
||||
"""A ``key_cmd`` failed to produce a usable token."""
|
||||
|
||||
|
||||
def materialize_probe_api_key(api_key: object) -> str:
|
||||
"""Best-effort probe credential; never send a callable's repr or log mint errors."""
|
||||
try:
|
||||
token = api_key() if callable(api_key) else api_key
|
||||
except Exception:
|
||||
return ""
|
||||
return token.strip() if isinstance(token, str) else ""
|
||||
|
||||
|
||||
def _mint(command: str, label: str) -> tuple[str, Optional[float]]:
|
||||
"""Run *command*, returning ``(token, ttl_seconds_or_None)``."""
|
||||
try:
|
||||
completed = subprocess.run(
|
||||
command,
|
||||
shell=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=_MINT_TIMEOUT_SECONDS,
|
||||
command, shell=True, capture_output=True, text=True, timeout=_MINT_TIMEOUT_SECONDS,
|
||||
)
|
||||
except subprocess.TimeoutExpired as exc:
|
||||
raise CommandTokenError(
|
||||
f"key_cmd for provider {label!r} timed out after "
|
||||
f"{_MINT_TIMEOUT_SECONDS}s"
|
||||
f"key_cmd for provider {label!r} timed out after {_MINT_TIMEOUT_SECONDS}s"
|
||||
) from exc
|
||||
except OSError as exc:
|
||||
raise CommandTokenError(
|
||||
f"key_cmd for provider {label!r} could not be executed: {exc}"
|
||||
) from exc
|
||||
raise CommandTokenError(f"key_cmd for provider {label!r} could not be executed: {exc}") from exc
|
||||
|
||||
if completed.returncode != 0:
|
||||
# NEVER include stdout/stderr: a partially-successful auth helper can
|
||||
# print a token or refresh secret there. The command STRING is also
|
||||
# withheld — a key_cmd can legitimately embed a secret
|
||||
# (`print-token --client-secret=…`), so echoing it back would leak the
|
||||
# very credential this module exists to protect. Name the provider so
|
||||
# the user knows which config entry to run by hand.
|
||||
# NEVER include stdout/stderr (may hold a token) or the command string (may embed
|
||||
# `--client-secret=…`); name the provider instead.
|
||||
raise CommandTokenError(
|
||||
f"key_cmd for provider {label!r} exited {completed.returncode}. "
|
||||
f"Run that provider's key_cmd manually to see why "
|
||||
@@ -101,8 +70,6 @@ def _mint(command: str, label: str) -> tuple[str, Optional[float]]:
|
||||
raise CommandTokenError(f"key_cmd for provider {label!r} produced no output")
|
||||
|
||||
# JSON payload — the shape `databricks auth token --output json` prints.
|
||||
# Token extraction mirrors databricks/ucode's get_databricks_token:
|
||||
# json.loads(result.stdout or "{}").get("access_token", "")
|
||||
if stdout.lstrip().startswith("{"):
|
||||
try:
|
||||
payload = json.loads(stdout)
|
||||
@@ -112,34 +79,24 @@ def _mint(command: str, label: str) -> tuple[str, Optional[float]]:
|
||||
token = str(payload.get("access_token") or "").strip()
|
||||
if not token:
|
||||
raise CommandTokenError(
|
||||
f"key_cmd for provider {label!r} returned JSON without an "
|
||||
"'access_token' field"
|
||||
f"key_cmd for provider {label!r} returned JSON without an 'access_token' field"
|
||||
)
|
||||
ttl = payload.get("expires_in")
|
||||
if isinstance(ttl, (int, float)) and ttl > 0:
|
||||
return token, float(ttl)
|
||||
# A relative lifetime is the OAuth 2.0 field, but CLI token helpers
|
||||
# commonly print an absolute ISO 8601 deadline instead. Treating
|
||||
# that as "no TTL advertised" caches the token for the life of the
|
||||
# process, so every request 401s once the deadline passes.
|
||||
# Imported lazily: hermes_cli.auth imports from agent.* at module
|
||||
# level, so a top-level import here would risk a cycle.
|
||||
# CLI helpers often print an absolute ISO 8601 deadline instead of OAuth's relative
|
||||
# lifetime; honour it or the token 401s once past. Lazy import: hermes_cli.auth imports agent.*.
|
||||
from hermes_cli.auth import _parse_iso_timestamp
|
||||
|
||||
for field in ("expiry", "expiresOn"):
|
||||
deadline = _parse_iso_timestamp(payload.get(field))
|
||||
if deadline is not None:
|
||||
remaining = deadline - time.time()
|
||||
if remaining > 0:
|
||||
return token, remaining
|
||||
remaining = deadline - time.time() if deadline is not None else 0
|
||||
if remaining > 0:
|
||||
return token, remaining
|
||||
return token, None
|
||||
|
||||
# Bare token. The contract every comparable helper documents is "stdout
|
||||
# carries the token and nothing else" — extra output would be consumed as
|
||||
# part of the credential. Strip surrounding whitespace and take the rest
|
||||
# verbatim; do NOT silently keep one line of several, which converts a
|
||||
# misconfigured helper (banner, warning, two tokens) into a corrupt-key 401
|
||||
# that is far harder to diagnose than an explicit refusal.
|
||||
# Bare token: stdout carries the token and nothing else. Do NOT keep one line of several — that
|
||||
# turns a misconfigured helper (banner, warning) into a corrupt-key 401 far harder to diagnose.
|
||||
token = stdout.strip()
|
||||
if "\n" in token:
|
||||
raise CommandTokenError(
|
||||
@@ -159,18 +116,19 @@ class CommandTokenSource:
|
||||
self._token = ""
|
||||
self._expires_at: float = 0.0
|
||||
|
||||
@property
|
||||
def cache_identity(self) -> str:
|
||||
"""Stable catalog identity; token rotation must not mint on cache reads."""
|
||||
return f"cmd:{self._command}"
|
||||
|
||||
def __call__(self) -> str:
|
||||
with self._lock:
|
||||
if self._token and time.monotonic() < self._expires_at:
|
||||
return self._token
|
||||
token, ttl = _mint(self._command, self._label)
|
||||
self._token = token
|
||||
self._expires_at = (
|
||||
time.monotonic() + max(ttl - _TOKEN_REFRESH_LEEWAY_SECONDS, 5.0)
|
||||
if ttl
|
||||
# No advertised TTL: bounded cache (see _NO_TTL_REFRESH_SECONDS)
|
||||
# — there is no 401-driven re-mint hook to fall back on.
|
||||
else time.monotonic() + _NO_TTL_REFRESH_SECONDS
|
||||
self._expires_at = time.monotonic() + (
|
||||
max(ttl - _TOKEN_REFRESH_LEEWAY_SECONDS, 5.0) if ttl else _NO_TTL_REFRESH_SECONDS
|
||||
)
|
||||
logger.debug(
|
||||
"Minted key_cmd token for provider %s (ttl=%s)",
|
||||
@@ -179,12 +137,7 @@ class CommandTokenSource:
|
||||
return token
|
||||
|
||||
|
||||
def build_command_token_provider(
|
||||
key_cmd: str,
|
||||
provider_label: str = "custom",
|
||||
) -> Optional[Callable[[], str]]:
|
||||
def build_command_token_provider(key_cmd: str, provider_label: str = "custom") -> Optional[CommandTokenSource]:
|
||||
"""A per-request token provider for *key_cmd*, or ``None`` when unset."""
|
||||
command = str(key_cmd or "").strip()
|
||||
if not command:
|
||||
return None
|
||||
return CommandTokenSource(command, provider_label)
|
||||
return CommandTokenSource(command, provider_label) if command else None
|
||||
|
||||
@@ -4,16 +4,25 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from agent.context_compressor import (
|
||||
ContextCompressor,
|
||||
is_compaction_summary_message,
|
||||
)
|
||||
from agent.context_compressor import ContextCompressor, is_compaction_summary_message
|
||||
|
||||
|
||||
_COMPACTION_INTERNAL_FIELDS = (
|
||||
"tool_calls",
|
||||
"finish_reason",
|
||||
"reasoning",
|
||||
# Provider replay/metadata fields that ride the wire on every request but are invisible to
|
||||
# ``msg["content"]``/``msg["tool_calls"]`` accounting. Codex Responses sessions in particular carry
|
||||
# ``codex_reasoning_items`` blobs of ``encrypted_content`` that can dominate the serialized session (a
|
||||
# measured 214-turn session held ~115K tokens / 27% of its payload there — #55572).
|
||||
# ``reasoning_details`` is handled separately (see ``_reasoning_details_text_chars``): its signed/base64
|
||||
# envelope is excluded from the budget, mirroring the preflight estimator's exclusion in
|
||||
# ``model_metadata._estimate_message_tokens_without_images`` (#73298).
|
||||
# An assistant turn may carry only reasoning/thinking content with no visible text (extended-thinking
|
||||
# turns, thinking-only recovery responses). Such a turn is persisted with its reasoning fields and is
|
||||
# recallable from the transcript, but dropping it here as "empty" makes it vanish from the
|
||||
# resumed/reloaded session view while the desktop's reasoning disclosure has nothing to render. Keep it
|
||||
# when it carries reasoning so the "Thinking…" block still shows. (#44022)
|
||||
"reasoning_content",
|
||||
"reasoning_details",
|
||||
"codex_reasoning_items",
|
||||
@@ -21,9 +30,7 @@ _COMPACTION_INTERNAL_FIELDS = (
|
||||
)
|
||||
|
||||
|
||||
def project_compaction_message_for_display(
|
||||
message: Dict[str, Any],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
def project_compaction_message_for_display(message: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""Return authentic transcript content, or ``None`` for a pure handoff.
|
||||
|
||||
Model-facing recovery history retains the complete carrier. Display
|
||||
|
||||
@@ -0,0 +1,280 @@
|
||||
"""Host-side ``AIAgent._compress_context`` wrapper.
|
||||
|
||||
Publishes the commit fence ``hard_interrupt()`` reads, runs the compressor on a snapshot under the progress
|
||||
timeout, mirrors ``_DB_PERSISTED_MARKER`` stamps back onto the live lists and rebinds the session context.
|
||||
Extracted from ``run_agent.py``; every method resolves through ``AIAgent``'s MRO unchanged.
|
||||
"""
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import copy
|
||||
import threading
|
||||
|
||||
from agent.session_activity import ActivityProvenance
|
||||
|
||||
# Same logger name as the origin module so log records / caplog filters are unchanged.
|
||||
logger = logging.getLogger("run_agent")
|
||||
|
||||
|
||||
def _timeout_fallback_prompt(agent, system_message: str) -> str:
|
||||
"""Cached prompt, else a fresh build, else the raw ``system_message`` (never raises).
|
||||
Resolved lazily by the timeout wrapper: an eager rebuild would raise before compress_context runs when
|
||||
``_cached_system_prompt`` is unset and the builder fails."""
|
||||
if cached := getattr(agent, "_cached_system_prompt", None):
|
||||
return cached
|
||||
try:
|
||||
return agent._build_system_prompt(system_message)
|
||||
except Exception:
|
||||
logger.debug("compress_context timeout fallback prompt rebuild failed; using raw system_message", exc_info=True)
|
||||
return system_message or ""
|
||||
|
||||
|
||||
def _report_compression_timeout(
|
||||
agent, *, idle: float, waited: float, since_progress: float, total_ceiling: float, total_exhausted: bool,
|
||||
progress_observed: bool,
|
||||
) -> None:
|
||||
"""Host-side timeout bookkeeping: log, activity stamp, cooldown ladder, user warning."""
|
||||
from agent.conversation_compression import mark_context_compression_timed_out
|
||||
mark_context_compression_timed_out(agent)
|
||||
if total_exhausted:
|
||||
logger.warning(
|
||||
"Context compression reached its total ceiling after %.1fs (progress observed=%s); continuing without compression",
|
||||
waited, progress_observed,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Context compression made no progress for %.1fs (total wait %.1fs, ceiling %.1fs); continuing without compression",
|
||||
since_progress, waited, total_ceiling,
|
||||
)
|
||||
touch = getattr(agent, "_touch_activity", None)
|
||||
if callable(touch):
|
||||
try:
|
||||
touch("context compression timed out", provenance=ActivityProvenance.AGENT_COMPRESSION_TIMEOUT)
|
||||
except Exception:
|
||||
logger.debug("compress_context timeout activity touch failed", exc_info=True)
|
||||
# Same timeout cooldown ladder as summary-LLM timeouts: avoid re-burning the full idle budget every turn.
|
||||
record = getattr(getattr(agent, "context_compressor", None), "record_timeout_failure", None)
|
||||
if callable(record):
|
||||
try:
|
||||
if total_exhausted:
|
||||
record("host compress_context total ceiling exhausted", failure_kind="ceiling_exhausted")
|
||||
else:
|
||||
record("host compress_context timeout (no summary progress)", failure_kind="stalled")
|
||||
except Exception:
|
||||
logger.debug("failed to record compress_context timeout cooldown", exc_info=True)
|
||||
emit = getattr(agent, "_emit_warning", None)
|
||||
if not callable(emit):
|
||||
return
|
||||
if total_exhausted:
|
||||
progress = " after summary output was observed" if progress_observed else ""
|
||||
emit(
|
||||
"⚠ Context compression reached its total ceiling "
|
||||
f"after {waited:.1f}s{progress}. No messages were "
|
||||
"dropped — continuing without compression. Run /compress to retry or /new for a clean session."
|
||||
)
|
||||
else:
|
||||
emit(
|
||||
f"⚠ Context compression timed out after {idle:.1f}s with no output from the summary "
|
||||
"model. No messages were dropped — continuing without compression. Run /compress to retry, /new "
|
||||
"for a clean session, or check auxiliary.compression."
|
||||
)
|
||||
|
||||
|
||||
def _warn_commit_overrun(agent, waited: float, ceiling: float) -> None:
|
||||
"""Commit-phase ceiling breach: the SessionDB mutation must complete, so only surface it."""
|
||||
emit = getattr(agent, "_emit_warning", None)
|
||||
if callable(emit):
|
||||
emit(
|
||||
f"⚠ Context compression commit is taking unusually long ({waited:.0f}s, ceiling {ceiling:.0f}s). "
|
||||
"Waiting for it to finish safely — if this persists, check SessionDB health (disk / lock contention)."
|
||||
)
|
||||
|
||||
|
||||
def _sync_persisted_markers(target_messages, source_messages) -> None:
|
||||
"""Mirror ``_DB_PERSISTED_MARKER`` stamps from the worker's snapshot onto a live list.
|
||||
Matched by scoped identity; timestamp-less repeated content is ambiguous, so every scoped match is
|
||||
stamped. Imported UNCONDITIONALLY: a silent fallback literal would split the stamping key from the flush's
|
||||
and resurrect the duplicate-row bug."""
|
||||
from agent.context_compressor import _DB_PERSISTED_MARKER
|
||||
from agent.conversation_compression import _stamp_scoped_twins
|
||||
if not isinstance(target_messages, list) or not isinstance(source_messages, list):
|
||||
return
|
||||
for source_message in source_messages:
|
||||
if isinstance(source_message, dict) and source_message.get(_DB_PERSISTED_MARKER):
|
||||
_stamp_scoped_twins(target_messages, source_message)
|
||||
|
||||
|
||||
def _run_under_progress_timeout(
|
||||
agent, run, messages, system_message, *, active_fence, fence_registration_lock, idle_timeout, total_ceiling
|
||||
):
|
||||
"""Run ``run(fence, target_messages=snapshot)`` on the pool under the progress-aware timeout.
|
||||
The pooled worker must NEVER share the caller's live transcript — a late engine after a host timeout could
|
||||
rewrite it. It deep-snapshots on the worker and publishes only via an ADMITTED commit; a no-op/abort
|
||||
returns the snapshot unchanged, so the ORIGINAL list is handed back to keep identity semantics."""
|
||||
from agent.conversation_compression import CompressionCommitFence, run_compress_context_with_progress_timeout
|
||||
|
||||
def _snapshot_worker(fence=None):
|
||||
# #76354 review F3: the pooled worker must NEVER share the caller's live transcript. Plugin/legacy
|
||||
# context engines are allowed to mutate their input list in place; after a host timeout the worker
|
||||
# stays alive, so a shared list would let a late engine rewrite the live conversation (roles,
|
||||
# ordering, persisted content) behind the caller's back. Deep-snapshot here, on the worker thread,
|
||||
# so the caller's list object is never touched by pooled code. Results are published to
|
||||
# caller-visible state only via the returned value of an ADMITTED commit (the host discards results
|
||||
# on timeout/cancel); durable SessionDB mutation is already gated behind the commit fence inside
|
||||
# compress_context.
|
||||
snapshot = copy.deepcopy(messages)
|
||||
result_msgs, result_prompt = run(fence, target_messages=snapshot)
|
||||
return (messages if result_msgs is snapshot else result_msgs), result_prompt
|
||||
|
||||
timeout_cause = {"total_exhausted": False, "progress_observed": False}
|
||||
|
||||
def _on_timeout_cause(total_exhausted, progress_observed):
|
||||
timeout_cause.update(total_exhausted=total_exhausted, progress_observed=progress_observed)
|
||||
|
||||
def _on_timeout(idle, waited, since_progress):
|
||||
_report_compression_timeout(
|
||||
agent, idle=idle, waited=waited, since_progress=since_progress, total_ceiling=total_ceiling, **timeout_cause
|
||||
)
|
||||
|
||||
def _publish_new_fence():
|
||||
# The stall-fallback retry needs a fence the aborted attempt cannot veto; publish
|
||||
# it on the slot hard_interrupt() reads. The caller's finally restores its fence.
|
||||
retry_fence = CompressionCommitFence()
|
||||
with fence_registration_lock:
|
||||
agent._active_compression_commit_fence = retry_fence
|
||||
return retry_fence
|
||||
|
||||
return run_compress_context_with_progress_timeout(
|
||||
worker=_snapshot_worker, messages=messages,
|
||||
system_prompt_fallback=lambda: _timeout_fallback_prompt(agent, system_message),
|
||||
idle_timeout_seconds=idle_timeout, total_ceiling_seconds=total_ceiling, on_timeout=_on_timeout,
|
||||
on_timeout_cause=_on_timeout_cause,
|
||||
on_commit_overrun=lambda waited, ceiling: _warn_commit_overrun(agent, waited, ceiling), fence=active_fence,
|
||||
telemetry_agent=agent, new_fence=_publish_new_fence,
|
||||
)
|
||||
|
||||
|
||||
def _mirror_result_onto_live_lists(agent, result, messages, *, direct_path: bool) -> None:
|
||||
"""Mirror persisted-marker stamps from the result list onto the live list(s)."""
|
||||
if not (isinstance(result, tuple) and result and isinstance(result[0], list)):
|
||||
return
|
||||
result_messages = result[0]
|
||||
# Direct-path callers bypass the snapshot worker but still need the post-publish mirror.
|
||||
if direct_path or result_messages is not messages:
|
||||
_sync_persisted_markers(messages, result_messages)
|
||||
session_messages = getattr(agent, "_session_messages", None)
|
||||
if isinstance(session_messages, list) and session_messages is not messages:
|
||||
# Durable-parent adoption can leave `_session_messages` on the pre-adoption list.
|
||||
_sync_persisted_markers(session_messages, result_messages)
|
||||
|
||||
|
||||
def _rebind_caller_session_context(agent) -> None:
|
||||
"""Propagate a rotated session id to the CALLER's thread/ContextVar (idempotent otherwise).
|
||||
The worker thread rotated hermes_logging's thread-local id; post-compression tools must resolve
|
||||
HERMES_SESSION_ID to the child id."""
|
||||
with contextlib.suppress(Exception):
|
||||
from hermes_logging import set_session_context
|
||||
set_session_context(agent.session_id)
|
||||
try:
|
||||
from gateway.session_context import set_current_session_id
|
||||
if agent.session_id:
|
||||
set_current_session_id(agent.session_id)
|
||||
except Exception:
|
||||
logger.debug("post-compression session ContextVar rebind failed", exc_info=True)
|
||||
|
||||
|
||||
class CompressionFacadeMixin:
|
||||
"""``_compress_context`` (see module docstring)."""
|
||||
|
||||
def _compress_context(
|
||||
self, messages: list, system_message: str, *, approx_tokens: int = None, task_id: str = "default",
|
||||
focus_topic: str = None, force: bool = False, bypass_cooldown: bool = False,
|
||||
defer_context_engine_notification: bool = False, commit_fence=None,
|
||||
) -> tuple:
|
||||
"""Forwarder — see ``agent.conversation_compression.compress_context``.
|
||||
``force=True`` (manual /compress) bypasses the summary-failure cooldown; ``bypass_cooldown=True``
|
||||
(provider-proven overflow recovery) runs one real attempt while the cooldown stays armed.
|
||||
|
||||
``force=True`` is passed by the manual ``/compress`` slash command so users can bypass the
|
||||
summary-failure cooldown after an auto-compress abort. Auto-compress callers use the default
|
||||
``force=False``. See #100661.
|
||||
"""
|
||||
# Per-attempt timeout signal for turn-start preflight and in-loop consumers: a stalled
|
||||
# compression must not be mistaken for a structural no-op. Thread-local + per-agent lock.
|
||||
# A stalled compression must not be mistaken for a structural no-op and followed by the oversized
|
||||
# provider request it was meant to prevent. The typed helper upgrades the simple attribute to
|
||||
# thread-local state guarded by a per-agent lock so overlapping automatic/manual entrypoints cannot
|
||||
# clobber each other's outcome (#98741).
|
||||
from agent.conversation_compression import (
|
||||
CompressionCommitFence, compress_context, reset_context_compression_timeout_outcome,
|
||||
resolve_context_compression_timeouts,
|
||||
)
|
||||
reset_context_compression_timeout_outcome(self)
|
||||
from agent.portal_tags import (
|
||||
get_affinity_scope, get_conversation_context, reset_affinity_scope, reset_conversation_context,
|
||||
set_affinity_scope, set_conversation_context,
|
||||
)
|
||||
from agent.prompt_cache_scope import declared_conversation_scope_safe
|
||||
# Out-of-turn compaction (/compact, gateway /compress, partial head compression) runs outside
|
||||
# run_conversation's ambient scope; publish the root as a fallback so the summarizer's call carries
|
||||
# the conversation tag. No-op for in-turn callers. Same for the ROUTING scope when declared.
|
||||
token = None
|
||||
if get_conversation_context() is None:
|
||||
root = self._conversation_root_id()
|
||||
if root:
|
||||
token = set_conversation_context(root)
|
||||
# Initialized alongside `token`: the turn-lease timeout/interrupt early returns leave the try block
|
||||
# before set_affinity_scope() runs, and the finally reads this name unconditionally
|
||||
# (UnboundLocalError otherwise — the 4 red cross-process lease tests on PR #97158).
|
||||
affinity_token = None
|
||||
if get_affinity_scope() is None:
|
||||
declared = declared_conversation_scope_safe(self)
|
||||
if declared:
|
||||
affinity_token = set_affinity_scope(declared)
|
||||
# Every compression has a fence; hard_interrupt() uses this exact instance to serialize cancel
|
||||
# admission against begin_commit(). Publication is serialized so overlapping automatic/manual
|
||||
# entrypoints cannot replace the fence of the attempt currently committing.
|
||||
active_fence = commit_fence or CompressionCommitFence()
|
||||
fence_registration_lock = vars(self).setdefault("_compression_commit_fence_lock", threading.RLock())
|
||||
with fence_registration_lock:
|
||||
missing_fence = object()
|
||||
previous_fence = vars(self).get("_active_compression_commit_fence", missing_fence)
|
||||
self._active_compression_commit_fence = active_fence
|
||||
try:
|
||||
|
||||
def _run(fence=None, target_messages=None):
|
||||
return compress_context(
|
||||
self, target_messages if target_messages is not None else messages, system_message,
|
||||
approx_tokens=approx_tokens, task_id=task_id, focus_topic=focus_topic, force=force,
|
||||
bypass_cooldown=bypass_cooldown,
|
||||
defer_context_engine_notification=(defer_context_engine_notification), commit_fence=fence,
|
||||
)
|
||||
|
||||
# Callers that already own a progress-aware wait (gateway session
|
||||
# hygiene) pass commit_fence and must not be double-wrapped.
|
||||
direct_path = commit_fence is not None
|
||||
if not direct_path:
|
||||
idle_timeout, total_ceiling = resolve_context_compression_timeouts()
|
||||
direct_path = idle_timeout <= 0
|
||||
if direct_path:
|
||||
result = _run(active_fence)
|
||||
else:
|
||||
result = _run_under_progress_timeout(
|
||||
self, _run, messages, system_message,
|
||||
active_fence=active_fence, fence_registration_lock=fence_registration_lock,
|
||||
idle_timeout=idle_timeout, total_ceiling=total_ceiling,
|
||||
)
|
||||
_mirror_result_onto_live_lists(self, result, messages, direct_path=direct_path)
|
||||
_rebind_caller_session_context(self)
|
||||
return result
|
||||
finally:
|
||||
with fence_registration_lock:
|
||||
if previous_fence is missing_fence:
|
||||
vars(self).pop("_active_compression_commit_fence", None)
|
||||
else:
|
||||
self._active_compression_commit_fence = previous_fence
|
||||
# Restore whatever the caller had, so a compaction never leaks its tag into the surrounding scope.
|
||||
if token is not None:
|
||||
reset_conversation_context(token)
|
||||
if affinity_token is not None:
|
||||
reset_affinity_scope(affinity_token)
|
||||
+186
-233
@@ -1,9 +1,8 @@
|
||||
"""Live session context-window breakdown for UI surfaces.
|
||||
|
||||
Estimates how the next provider request is composed: system prompt tiers,
|
||||
tool schemas, and conversation history. Uses the same rough char/4 heuristic
|
||||
as ``agent.model_metadata.estimate_request_tokens_rough`` so numbers align
|
||||
with compression thresholds — not exact tokenizer counts.
|
||||
Estimates system prompt tiers, tool schemas, and conversation history for the
|
||||
category breakdown. Overall occupancy retains its provider-usage or estimate
|
||||
provenance; category estimates are not exact tokenizer counts or gate authority.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -13,38 +12,41 @@ import re
|
||||
from typing import Any, Dict, List, Optional, Sequence, Tuple
|
||||
|
||||
_SKILLS_BLOCK_RE = re.compile(r"<available_skills>.*?</available_skills>", re.DOTALL)
|
||||
|
||||
_SUBAGENT_TOOL_NAMES = frozenset({"delegate_task"})
|
||||
|
||||
_CATEGORY_COLORS = {
|
||||
"system_prompt": "var(--context-usage-system)",
|
||||
"tool_definitions": "var(--context-usage-tools)",
|
||||
"rules": "var(--context-usage-rules)",
|
||||
"skills": "var(--context-usage-skills)",
|
||||
"mcp": "var(--context-usage-mcp)",
|
||||
"subagent_definitions": "var(--context-usage-subagents)",
|
||||
"memory": "var(--context-usage-memory)",
|
||||
"conversation": "var(--context-usage-conversation)",
|
||||
# id -> (label, dashboard color, /context glyph); declaration order is display order.
|
||||
_CATEGORIES = {
|
||||
"system_prompt": ("System prompt", "var(--context-usage-system)", "■"),
|
||||
"tool_definitions": ("Tool definitions", "var(--context-usage-tools)", "▣"),
|
||||
"rules": ("Rules", "var(--context-usage-rules)", "▩"),
|
||||
"skills": ("Skills", "var(--context-usage-skills)", "▤"),
|
||||
"mcp": ("MCP", "var(--context-usage-mcp)", "▥"),
|
||||
"subagent_definitions": ("Subagent definitions", "var(--context-usage-subagents)", "▦"),
|
||||
"memory": ("Memory", "var(--context-usage-memory)", "▧"),
|
||||
"conversation": ("Conversation", "var(--context-usage-conversation)", "▨"),
|
||||
}
|
||||
_FREE_GLYPH = "·"
|
||||
_GRID_COLUMNS = 20
|
||||
_GRID_ROWS = 5 # 100 cells → 1 cell per percent of the context window
|
||||
_DETAILS_TABLE_LIMIT = 15 # display cap only; the underlying data keeps everything
|
||||
|
||||
|
||||
def _chars_to_tokens(text: str) -> int:
|
||||
if not text:
|
||||
return 0
|
||||
return (len(text) + 3) // 4
|
||||
|
||||
|
||||
def _json_tokens(value: Any) -> int:
|
||||
if not value:
|
||||
return 0
|
||||
return _chars_to_tokens(json.dumps(value, ensure_ascii=False))
|
||||
return _chars_to_tokens(json.dumps(value, ensure_ascii=False)) if value else 0
|
||||
|
||||
|
||||
def _tool_name(tool: dict) -> str:
|
||||
fn = tool.get("function") if isinstance(tool, dict) else None
|
||||
if isinstance(fn, dict):
|
||||
return str(fn.get("name") or "")
|
||||
return str(tool.get("name") or "")
|
||||
def _bytes_to_tokens(size: Optional[int]) -> Optional[int]:
|
||||
return None if size is None else (int(size) + 3) // 4
|
||||
|
||||
|
||||
def _skills_block(stable: str) -> str:
|
||||
"""The live ``<available_skills>`` block inside the stable tier, or ''."""
|
||||
m = _SKILLS_BLOCK_RE.search(stable)
|
||||
return m.group(0) if m else ""
|
||||
|
||||
|
||||
def _split_tools(tools: Sequence[dict]) -> Tuple[List[dict], List[dict], List[dict]]:
|
||||
@@ -52,26 +54,20 @@ def _split_tools(tools: Sequence[dict]) -> Tuple[List[dict], List[dict], List[di
|
||||
mcp: List[dict] = []
|
||||
subagent: List[dict] = []
|
||||
for tool in tools:
|
||||
name = _tool_name(tool)
|
||||
if name.startswith("mcp_"):
|
||||
mcp.append(tool)
|
||||
elif name in _SUBAGENT_TOOL_NAMES:
|
||||
subagent.append(tool)
|
||||
else:
|
||||
builtin.append(tool)
|
||||
fn = tool.get("function") if isinstance(tool, dict) else None
|
||||
name = str((fn if isinstance(fn, dict) else tool).get("name") or "")
|
||||
bucket = mcp if name.startswith("mcp_") else subagent if name in _SUBAGENT_TOOL_NAMES else builtin
|
||||
bucket.append(tool)
|
||||
return builtin, mcp, subagent
|
||||
|
||||
|
||||
def _memory_blocks(agent: Any) -> Tuple[str, str]:
|
||||
memory_block = ""
|
||||
user_block = ""
|
||||
memory_block = user_block = ""
|
||||
store = getattr(agent, "_memory_store", None)
|
||||
if store is None:
|
||||
return memory_block, user_block
|
||||
try:
|
||||
if getattr(agent, "_memory_enabled", True):
|
||||
if store is not None and getattr(agent, "_memory_enabled", True):
|
||||
memory_block = store.format_for_system_prompt("memory") or ""
|
||||
if getattr(agent, "_user_profile_enabled", True):
|
||||
if store is not None and getattr(agent, "_user_profile_enabled", True):
|
||||
user_block = store.format_for_system_prompt("user") or ""
|
||||
except Exception:
|
||||
pass
|
||||
@@ -79,182 +75,160 @@ def _memory_blocks(agent: Any) -> Tuple[str, str]:
|
||||
|
||||
|
||||
def _strip_blocks(text: str, *blocks: str) -> str:
|
||||
out = text
|
||||
for block in blocks:
|
||||
if block:
|
||||
out = out.replace(block, "")
|
||||
return out.strip()
|
||||
text = text.replace(block, "")
|
||||
return text.strip()
|
||||
|
||||
|
||||
def compute_session_context_breakdown(
|
||||
agent: Any,
|
||||
messages: Optional[List[dict]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
def _join(*parts: str) -> str:
|
||||
return "\n\n".join(part for part in parts if part).strip()
|
||||
|
||||
|
||||
def _glyph(cat: Dict[str, Any]) -> str:
|
||||
return _CATEGORIES.get(str(cat.get("id") or ""), (None, None, "▪"))[2]
|
||||
|
||||
|
||||
def context_display_source(compressor: Any) -> str:
|
||||
"""Distinguish the built-in preflight display seed from a provider reading.
|
||||
|
||||
Engines without the built-in real-usage ledger own their occupancy figure.
|
||||
A seed never updates that ledger, even if its number later matches real usage.
|
||||
"""
|
||||
real = getattr(compressor, "last_real_prompt_tokens", None)
|
||||
shown = getattr(compressor, "last_prompt_tokens", 0) or 0
|
||||
return "local_estimate" if isinstance(real, (int, float)) and shown > 0 and shown != real else "provider_usage"
|
||||
|
||||
|
||||
def context_usage_fields(compressor: Any) -> Dict[str, Any]:
|
||||
"""Current occupancy only; lifetime throughput is never a context fallback."""
|
||||
used = max(0, getattr(compressor, "last_prompt_tokens", 0) or 0)
|
||||
maximum = getattr(compressor, "context_length", 0) or 0
|
||||
if not used or not maximum:
|
||||
return {}
|
||||
source = context_display_source(compressor)
|
||||
return {"context_used": used, "context_max": maximum,
|
||||
"context_percent": max(0, min(100, round(used / maximum * 100))),
|
||||
"context_source": source, "context_estimated": source != "provider_usage"}
|
||||
|
||||
|
||||
def compute_session_context_breakdown(agent: Any, messages: Optional[List[dict]] = None) -> Dict[str, Any]:
|
||||
"""Return a Cursor-style context usage breakdown for one live agent."""
|
||||
from agent.model_metadata import estimate_messages_tokens_rough
|
||||
from agent.usage_anchor import anchored_context_tokens
|
||||
from agent.system_prompt import build_system_prompt_parts
|
||||
|
||||
messages = messages or []
|
||||
parts = build_system_prompt_parts(agent)
|
||||
stable = parts.get("stable", "") or ""
|
||||
context = parts.get("context", "") or ""
|
||||
volatile = parts.get("volatile", "") or ""
|
||||
|
||||
skills_match = _SKILLS_BLOCK_RE.search(stable)
|
||||
skills_index = skills_match.group(0) if skills_match else ""
|
||||
|
||||
skills_index = _skills_block(stable)
|
||||
memory_block, user_block = _memory_blocks(agent)
|
||||
memory_text = "\n\n".join(part for part in (memory_block, user_block) if part).strip()
|
||||
|
||||
system_core = _strip_blocks(stable, skills_index)
|
||||
system_tail = _strip_blocks(volatile, memory_block, user_block)
|
||||
system_prompt_text = "\n\n".join(part for part in (system_core, system_tail) if part).strip()
|
||||
|
||||
tools = list(getattr(agent, "tools", None) or [])
|
||||
builtin_tools, mcp_tools, subagent_tools = _split_tools(tools)
|
||||
|
||||
conversation_tokens = estimate_messages_tokens_rough(messages or [])
|
||||
|
||||
categories = [
|
||||
("system_prompt", "System prompt", _chars_to_tokens(system_prompt_text)),
|
||||
("tool_definitions", "Tool definitions", _json_tokens(builtin_tools)),
|
||||
("rules", "Rules", _chars_to_tokens(context)),
|
||||
("skills", "Skills", _chars_to_tokens(skills_index)),
|
||||
("mcp", "MCP", _json_tokens(mcp_tools)),
|
||||
("subagent_definitions", "Subagent definitions", _json_tokens(subagent_tools)),
|
||||
("memory", "Memory", _chars_to_tokens(memory_text)),
|
||||
("conversation", "Conversation", conversation_tokens),
|
||||
]
|
||||
|
||||
estimated_total = sum(tokens for _, _, tokens in categories)
|
||||
system_prompt_text = _join(
|
||||
_strip_blocks(stable, skills_index), _strip_blocks(parts.get("volatile", "") or "", memory_block, user_block)
|
||||
)
|
||||
builtin_tools, mcp_tools, subagent_tools = _split_tools(list(getattr(agent, "tools", None) or []))
|
||||
tokens_by_id = {
|
||||
"system_prompt": _chars_to_tokens(system_prompt_text),
|
||||
"tool_definitions": _json_tokens(builtin_tools),
|
||||
"rules": _chars_to_tokens(parts.get("context", "") or ""),
|
||||
"skills": _chars_to_tokens(skills_index),
|
||||
"mcp": _json_tokens(mcp_tools),
|
||||
"subagent_definitions": _json_tokens(subagent_tools),
|
||||
"memory": _chars_to_tokens(_join(memory_block, user_block)),
|
||||
"conversation": estimate_messages_tokens_rough(messages),
|
||||
}
|
||||
estimated_total = sum(tokens_by_id.values())
|
||||
|
||||
comp = getattr(agent, "context_compressor", None)
|
||||
context_max = int(getattr(comp, "context_length", 0) or 0) if comp else 0
|
||||
measured_used = int(getattr(comp, "last_prompt_tokens", 0) or 0) if comp else 0
|
||||
context_used = measured_used if measured_used > 0 else estimated_total
|
||||
context_percent = (
|
||||
max(0, min(100, round(context_used / context_max * 100)))
|
||||
if context_max
|
||||
else 0
|
||||
)
|
||||
# Usage-anchored figure (provider-exact tokens of a response + delta of what was
|
||||
# appended since) beats last_prompt_tokens (lags) and the heuristic. Prefer the
|
||||
# turn-base anchor: on reasoning models later same-turn responses inflate
|
||||
# prompt_tokens with replayed thinking that evaporates at the turn boundary, so
|
||||
# anchoring on the LAST response makes the meter sawtooth. Fall back to the
|
||||
# last-response anchor, then measured, then estimated.
|
||||
anchor = getattr(agent, "_turn_base_usage_anchor", None)
|
||||
context_used = anchored_context_tokens(messages, anchor, charge_stale_thinking=False)
|
||||
if context_used is None:
|
||||
anchor = getattr(agent, "_usage_anchor", None)
|
||||
context_used = anchored_context_tokens(messages, anchor)
|
||||
if context_used is None:
|
||||
measured_used = int(getattr(comp, "last_prompt_tokens", 0) or 0) if comp else 0
|
||||
context_used = measured_used if measured_used > 0 else estimated_total
|
||||
source = context_display_source(comp) if measured_used > 0 else "local_estimate"
|
||||
else:
|
||||
delta = messages[int(anchor["base_count"]):]
|
||||
if delta and delta[0].get("role") == "assistant":
|
||||
delta = delta[1:]
|
||||
source = "provider_usage_plus_estimate" if delta else "provider_usage"
|
||||
|
||||
return {
|
||||
"categories": [
|
||||
{
|
||||
"color": _CATEGORY_COLORS.get(category_id, "var(--ui-text-tertiary)"),
|
||||
"id": category_id,
|
||||
"label": label,
|
||||
"tokens": tokens,
|
||||
}
|
||||
for category_id, label, tokens in categories
|
||||
if tokens > 0
|
||||
{"color": color, "id": category_id, "label": label, "tokens": tokens_by_id[category_id]}
|
||||
for category_id, (label, color, _glyph_) in _CATEGORIES.items()
|
||||
if tokens_by_id[category_id] > 0
|
||||
],
|
||||
"context_max": context_max,
|
||||
"context_percent": context_percent,
|
||||
"context_percent": max(0, min(100, round(context_used / context_max * 100))) if context_max else 0,
|
||||
"context_used": context_used,
|
||||
"context_source": source,
|
||||
"context_estimated": source != "provider_usage",
|
||||
"estimated_total": estimated_total,
|
||||
"model": getattr(agent, "model", "") or "",
|
||||
}
|
||||
|
||||
|
||||
# ── /context rendering (CLI + gateway) ──────────────────────────────────────
|
||||
#
|
||||
# Pure text renderers over the payload above. The CLI shows a glyph block-grid
|
||||
# plus a category table; the gateway uses the same table without the grid
|
||||
# (proportional monospace is not guaranteed on messaging platforms).
|
||||
|
||||
_CATEGORY_GLYPHS = {
|
||||
"system_prompt": "■",
|
||||
"tool_definitions": "▣",
|
||||
"rules": "▩",
|
||||
"skills": "▤",
|
||||
"mcp": "▥",
|
||||
"subagent_definitions": "▦",
|
||||
"memory": "▧",
|
||||
"conversation": "▨",
|
||||
}
|
||||
_FREE_GLYPH = "·"
|
||||
_GRID_COLUMNS = 20
|
||||
_GRID_ROWS = 5 # 100 cells → 1 cell per percent of the context window
|
||||
|
||||
# Human-readable tables cap the expanded listings; nothing is dropped from
|
||||
# the underlying data.
|
||||
_DETAILS_TABLE_LIMIT = 15
|
||||
|
||||
|
||||
def _bytes_to_tokens(size: Optional[int]) -> Optional[int]:
|
||||
if size is None:
|
||||
return None
|
||||
return (int(size) + 3) // 4
|
||||
|
||||
|
||||
def compute_context_details(agent: Any) -> Dict[str, Any]:
|
||||
"""Expanded per-skill / per-toolset cost listing for ``/context all``.
|
||||
|
||||
Reuses the ``hermes prompt-size`` attribution mechanism (PR #66656):
|
||||
per-skill index-line bytes parsed from the live ``<available_skills>``
|
||||
block, and per-toolset schema bytes attributed via the tool registry's
|
||||
canonical tool→toolset map. Byte figures are converted to the same
|
||||
chars/4 token heuristic the categories above use.
|
||||
Reuses the ``hermes prompt-size`` attribution (index-line bytes from the
|
||||
live skills block; schema bytes via the registry's tool→toolset map).
|
||||
"""
|
||||
from hermes_cli.prompt_size import (
|
||||
_compute_skills_breakdown,
|
||||
_compute_toolsets_breakdown,
|
||||
)
|
||||
from hermes_cli.prompt_size import _compute_skills_breakdown, _compute_toolsets_breakdown
|
||||
from agent.system_prompt import build_system_prompt_parts
|
||||
|
||||
parts = build_system_prompt_parts(agent)
|
||||
stable = parts.get("stable", "") or ""
|
||||
skills_match = _SKILLS_BLOCK_RE.search(stable)
|
||||
skills_block = skills_match.group(0) if skills_match else ""
|
||||
|
||||
skills: List[Dict[str, Any]] = []
|
||||
if skills_block:
|
||||
for entry in _compute_skills_breakdown(skills_block):
|
||||
skills.append({
|
||||
skills_block = _skills_block(build_system_prompt_parts(agent).get("stable", "") or "")
|
||||
tools = list(getattr(agent, "tools", None) or [])
|
||||
return {
|
||||
"skills": [
|
||||
{
|
||||
"name": entry.get("name", ""),
|
||||
"index_tokens": _bytes_to_tokens(entry.get("index_line_bytes")) or 0,
|
||||
"skill_md_tokens": _bytes_to_tokens(entry.get("skill_md_bytes")),
|
||||
})
|
||||
|
||||
toolsets: List[Dict[str, Any]] = []
|
||||
tools = list(getattr(agent, "tools", None) or [])
|
||||
if tools:
|
||||
for group in _compute_toolsets_breakdown(tools):
|
||||
toolsets.append({
|
||||
}
|
||||
for entry in (_compute_skills_breakdown(skills_block) if skills_block else [])
|
||||
],
|
||||
"toolsets": [
|
||||
{
|
||||
"toolset": group.get("toolset", ""),
|
||||
"tool_count": int(group.get("tool_count", 0) or 0),
|
||||
"schema_tokens": _bytes_to_tokens(group.get("json_bytes")) or 0,
|
||||
})
|
||||
}
|
||||
for group in (_compute_toolsets_breakdown(tools) if tools else [])
|
||||
],
|
||||
}
|
||||
|
||||
return {"skills": skills, "toolsets": toolsets}
|
||||
|
||||
# ── /context rendering (CLI + gateway) ──────────────────────────────────────
|
||||
# Pure text renderers over the payload above. The gateway skips the glyph grid
|
||||
# (monospace is not guaranteed on messaging platforms).
|
||||
|
||||
|
||||
def render_context_grid(payload: Dict[str, Any]) -> List[str]:
|
||||
"""Render the payload as a Claude Code-style glyph block grid.
|
||||
|
||||
100 cells (5×20), each one percent of the model context window. Categories
|
||||
fill in declaration order; the remainder renders as free space.
|
||||
"""
|
||||
"""Glyph grid: 100 cells, one per percent of the context window; categories
|
||||
fill in declaration order, the remainder is free space."""
|
||||
context_max = int(payload.get("context_max") or 0)
|
||||
categories = payload.get("categories") or []
|
||||
total_cells = _GRID_COLUMNS * _GRID_ROWS
|
||||
|
||||
cells: List[str] = []
|
||||
if context_max > 0:
|
||||
for cat in categories:
|
||||
for cat in payload.get("categories") or []:
|
||||
tokens = int(cat.get("tokens") or 0)
|
||||
n = round(tokens / context_max * total_cells)
|
||||
if tokens > 0 and n == 0:
|
||||
n = 1 # never render a nonzero category as invisible
|
||||
glyph = _CATEGORY_GLYPHS.get(str(cat.get("id") or ""), "▪")
|
||||
cells.extend([glyph] * n)
|
||||
# never render a nonzero category as invisible
|
||||
n = round(tokens / context_max * total_cells) or (1 if tokens > 0 else 0)
|
||||
cells.extend([_glyph(cat)] * n)
|
||||
cells = cells[:total_cells]
|
||||
cells.extend([_FREE_GLYPH] * (total_cells - len(cells)))
|
||||
|
||||
return [
|
||||
" ".join(cells[row * _GRID_COLUMNS:(row + 1) * _GRID_COLUMNS])
|
||||
for row in range(_GRID_ROWS)
|
||||
]
|
||||
return [" ".join(cells[row * _GRID_COLUMNS:(row + 1) * _GRID_COLUMNS]) for row in range(_GRID_ROWS)]
|
||||
|
||||
|
||||
def render_context_category_lines(payload: Dict[str, Any]) -> List[str]:
|
||||
@@ -266,59 +240,47 @@ def render_context_category_lines(payload: Dict[str, Any]) -> List[str]:
|
||||
|
||||
lines = ["Estimated usage by category"]
|
||||
if not categories:
|
||||
lines.append(" (no data yet — send a message first)")
|
||||
return lines
|
||||
|
||||
width = max(len(str(cat.get("label") or "")) for cat in categories)
|
||||
width = max(width, len("Free space"))
|
||||
return [*lines, " (no data yet — send a message first)"]
|
||||
width = max(len("Free space"), *(len(str(cat.get("label") or "")) for cat in categories))
|
||||
for cat in categories:
|
||||
tokens = int(cat.get("tokens") or 0)
|
||||
glyph = _CATEGORY_GLYPHS.get(str(cat.get("id") or ""), "▪")
|
||||
pct = tokens / denom * 100 if denom else 0.0
|
||||
label = str(cat.get("label") or cat.get("id") or "")
|
||||
lines.append(f"{glyph} {label:<{width}} {tokens:>9,} tokens {pct:>5.1f}%")
|
||||
tokens, label = int(cat.get("tokens") or 0), str(cat.get("label") or cat.get("id") or "")
|
||||
lines.append(f"{_glyph(cat)} {label:<{width}} ~{tokens:>9,} tokens ~{tokens / denom * 100 if denom else 0.0:>5.1f}%")
|
||||
if context_max > 0:
|
||||
free = max(0, context_max - estimated_total)
|
||||
pct = free / context_max * 100
|
||||
lines.append(f"{_FREE_GLYPH} {'Free space':<{width}} {free:>9,} tokens {pct:>5.1f}%")
|
||||
lines.append(f"{_FREE_GLYPH} {'Free space':<{width}} ~{free:>9,} tokens ~{free / context_max * 100:>5.1f}%")
|
||||
return lines
|
||||
|
||||
|
||||
def _toolset_row(group: Dict[str, Any]) -> str:
|
||||
return f" {group['toolset']:<24} {group['tool_count']:>3} tools ~{group['schema_tokens']:>8,} tokens"
|
||||
|
||||
|
||||
def _skill_row(entry: Dict[str, Any]) -> str:
|
||||
name = str(entry.get("name") or "")
|
||||
if len(name) > 28:
|
||||
name = name[:27] + "…"
|
||||
md = entry.get("skill_md_tokens")
|
||||
md_str = f"~{md:>8,}" if md is not None else f"{'n/a':>8}"
|
||||
return f" {name:<28} index ~{entry['index_tokens']:>6,} SKILL.md {md_str} tokens"
|
||||
|
||||
|
||||
def _table(lines: List[str], title: str, rows: List[Dict[str, Any]], fmt) -> None:
|
||||
"""Append a titled, display-capped table (blank-separated from a preceding one)."""
|
||||
if not rows:
|
||||
return
|
||||
if lines:
|
||||
lines.append("")
|
||||
lines.append(title)
|
||||
lines.extend(fmt(row) for row in rows[:_DETAILS_TABLE_LIMIT])
|
||||
if len(rows) > _DETAILS_TABLE_LIMIT:
|
||||
lines.append(f" … and {len(rows) - _DETAILS_TABLE_LIMIT} more")
|
||||
|
||||
|
||||
def render_context_details_lines(details: Dict[str, Any]) -> List[str]:
|
||||
"""Render the expanded ``/context all`` per-skill / per-toolset tables."""
|
||||
lines: List[str] = []
|
||||
|
||||
toolsets = details.get("toolsets") or []
|
||||
if toolsets:
|
||||
lines.append("Toolsets by schema cost (largest first)")
|
||||
for group in toolsets[:_DETAILS_TABLE_LIMIT]:
|
||||
lines.append(
|
||||
f" {group['toolset']:<24} {group['tool_count']:>3} tools"
|
||||
f" {group['schema_tokens']:>8,} tokens"
|
||||
)
|
||||
remaining = len(toolsets) - _DETAILS_TABLE_LIMIT
|
||||
if remaining > 0:
|
||||
lines.append(f" … and {remaining} more")
|
||||
|
||||
skills = details.get("skills") or []
|
||||
if skills:
|
||||
if lines:
|
||||
lines.append("")
|
||||
lines.append("Skills by cost (index = always-on; SKILL.md = cost when loaded)")
|
||||
for entry in skills[:_DETAILS_TABLE_LIMIT]:
|
||||
name = str(entry.get("name") or "")
|
||||
if len(name) > 28:
|
||||
name = name[:27] + "…"
|
||||
md = entry.get("skill_md_tokens")
|
||||
md_str = f"{md:>8,}" if md is not None else f"{'n/a':>8}"
|
||||
lines.append(
|
||||
f" {name:<28} index {entry['index_tokens']:>6,}"
|
||||
f" SKILL.md {md_str} tokens"
|
||||
)
|
||||
remaining = len(skills) - _DETAILS_TABLE_LIMIT
|
||||
if remaining > 0:
|
||||
lines.append(f" … and {remaining} more")
|
||||
|
||||
_table(lines, "Toolsets by schema cost (largest first)", details.get("toolsets") or [], _toolset_row)
|
||||
_table(lines, "Skills by cost (index = always-on; SKILL.md = cost when loaded)", details.get("skills") or [], _skill_row)
|
||||
return lines
|
||||
|
||||
|
||||
@@ -328,33 +290,24 @@ def render_context_breakdown_lines(
|
||||
details: Optional[Dict[str, Any]] = None,
|
||||
grid: bool = True,
|
||||
) -> List[str]:
|
||||
"""Render the full /context view as plain-text lines.
|
||||
|
||||
``grid=True`` (CLI) prepends the glyph block grid; the gateway passes
|
||||
``grid=False`` and keeps its own gauge. ``details`` (from
|
||||
:func:`compute_context_details`) appends the expanded listings.
|
||||
"""
|
||||
lines: List[str] = []
|
||||
if grid:
|
||||
lines.extend(render_context_grid(payload))
|
||||
lines.append("")
|
||||
"""Full /context view. ``grid`` prepends the glyph grid (CLI; the gateway
|
||||
keeps its own gauge); ``details`` appends the expanded listings."""
|
||||
lines: List[str] = [*render_context_grid(payload), ""] if grid else []
|
||||
lines.extend(render_context_category_lines(payload))
|
||||
|
||||
context_max = int(payload.get("context_max") or 0)
|
||||
context_used = int(payload.get("context_used") or 0)
|
||||
if context_max > 0:
|
||||
pct = int(payload.get("context_percent") or 0)
|
||||
lines.append("")
|
||||
lines.append(
|
||||
f"Context window: {context_used:,} / {context_max:,} tokens ({pct}%)"
|
||||
)
|
||||
used, pct = int(payload.get("context_used") or 0), int(payload.get("context_percent") or 0)
|
||||
mark = "~" if payload.get("context_estimated") else ""
|
||||
lines.extend(["", f"Context window: {mark}{used:,} / {context_max:,} tokens ({mark}{pct}%)"])
|
||||
source = payload.get("context_source")
|
||||
if source:
|
||||
labels = {"local_estimate": "local estimate", "provider_usage": "provider usage",
|
||||
"provider_usage_plus_estimate": "provider usage + estimated new messages"}
|
||||
lines.append(f"Source: {labels.get(source, source)}; category counts are local estimates.")
|
||||
|
||||
if details is not None:
|
||||
detail_lines = render_context_details_lines(details)
|
||||
if detail_lines:
|
||||
lines.append("")
|
||||
lines.extend(detail_lines)
|
||||
else:
|
||||
lines.append("")
|
||||
lines.append("Use /context all for per-skill and per-toolset costs.")
|
||||
if details is None:
|
||||
lines.extend(["", "Use /context all for per-skill and per-toolset costs."])
|
||||
elif detail_lines := render_context_details_lines(details):
|
||||
lines.extend(["", *detail_lines])
|
||||
return lines
|
||||
|
||||
+3134
-6843
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,48 @@
|
||||
"""Summary-hook dispatch and cancellation rollback for context compression."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from agent.auxiliary_client import AuxiliaryExplicitCancellation
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agent.context_compressor import _HandoffScan
|
||||
|
||||
|
||||
def _accepts_keyword_argument(callable_obj: Any, name: str) -> bool:
|
||||
"""Return whether an inspectable callable accepts ``name`` as a keyword."""
|
||||
try:
|
||||
parameters = inspect.signature(callable_obj).parameters
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
if any(parameter.kind is inspect.Parameter.VAR_KEYWORD for parameter in parameters.values()):
|
||||
return True
|
||||
parameter = parameters.get(name)
|
||||
return parameter is not None and parameter.kind in (
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
inspect.Parameter.KEYWORD_ONLY,
|
||||
)
|
||||
|
||||
|
||||
class SummaryDispatchMixin:
|
||||
def _summarize_window(
|
||||
self, messages: List[Dict[str, Any]], turns_to_summarize: List[Dict[str, Any]], scan: "_HandoffScan",
|
||||
focus_topic: Optional[str], memory_context: str, bypass_cooldown: bool,
|
||||
) -> Optional[str]:
|
||||
"""Run the summary LLM; a cancellation rolls back the handoff scan's self-heal mutation first."""
|
||||
# Focus-topic derivation scans user turns; only pay when a summary is generated.
|
||||
summary_kwargs: Dict[str, Any] = {
|
||||
"focus_topic": focus_topic or self._derive_auto_focus_topic(messages),
|
||||
"memory_context": memory_context,
|
||||
}
|
||||
if _accepts_keyword_argument(self._generate_summary, "bypass_cooldown"):
|
||||
summary_kwargs["bypass_cooldown"] = bypass_cooldown
|
||||
try:
|
||||
return self._generate_summary(turns_to_summarize, **summary_kwargs)
|
||||
except AuxiliaryExplicitCancellation:
|
||||
# Cancellation is a true no-op: restore the scan's mutation before the exception escapes.
|
||||
self._previous_summary = scan.previous_summary_before
|
||||
self._summary_has_user_turn = scan.has_user_turn_before
|
||||
raise
|
||||
+105
-352
@@ -1,30 +1,14 @@
|
||||
"""Abstract base class for pluggable context engines.
|
||||
|
||||
A context engine controls how conversation context is managed when
|
||||
approaching the model's token limit. The built-in ContextCompressor
|
||||
is the default implementation. Third-party engines (e.g. LCM) can
|
||||
replace it via the plugin system or by being placed in the
|
||||
``plugins/context_engine/<name>/`` directory.
|
||||
|
||||
Selection is config-driven: ``context.engine`` in config.yaml.
|
||||
Default is ``"compressor"`` (the built-in). Only one engine is active.
|
||||
|
||||
The engine is responsible for:
|
||||
- Deciding when compaction should fire
|
||||
- Performing compaction (summarization, DAG construction, etc.)
|
||||
- Optionally exposing tools the agent can call (e.g. lcm_grep)
|
||||
- Tracking token usage from API responses
|
||||
|
||||
Lifecycle:
|
||||
1. Engine is instantiated and registered (plugin register() or default)
|
||||
2. on_session_start() called when a conversation begins
|
||||
3. update_from_response() called after each API response with usage data
|
||||
4. should_compress() checked after each turn
|
||||
5. compress() called when should_compress() returns True
|
||||
6. on_session_end() called at real session boundaries (CLI exit, /reset,
|
||||
gateway session expiry) — NOT per-turn
|
||||
A context engine decides when/how conversation context is compacted near the token
|
||||
limit, tracks usage, and may expose tools. ContextCompressor is the default;
|
||||
``context.engine`` selects a plugin (``plugins/context_engine/<name>/``); one is active.
|
||||
Lifecycle: on_session_start() -> per API response update_from_response() -> per turn
|
||||
should_compress() / compress() -> on_session_end() at real session boundaries only
|
||||
(CLI exit, /reset, gateway expiry), never per-turn.
|
||||
"""
|
||||
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
@@ -39,107 +23,61 @@ _MEMORY_CONTEXT_TRUNCATION_MARKER = "\n...[memory provider context truncated]...
|
||||
|
||||
def sanitize_memory_context(memory_context: str) -> str:
|
||||
"""Prepare provider context for a context-engine/LLM egress boundary."""
|
||||
sanitized = redact_sensitive_text(
|
||||
memory_context.strip(),
|
||||
force=True,
|
||||
redact_url_credentials=True,
|
||||
)
|
||||
sanitized = redact_sensitive_text(memory_context.strip(), force=True, redact_url_credentials=True)
|
||||
if len(sanitized) <= MEMORY_CONTEXT_MAX_CHARS:
|
||||
return sanitized
|
||||
return (
|
||||
sanitized[:_MEMORY_CONTEXT_HEAD_CHARS]
|
||||
+ _MEMORY_CONTEXT_TRUNCATION_MARKER
|
||||
+ sanitized[-_MEMORY_CONTEXT_TAIL_CHARS:]
|
||||
)
|
||||
return sanitized[:_MEMORY_CONTEXT_HEAD_CHARS] + _MEMORY_CONTEXT_TRUNCATION_MARKER + sanitized[-_MEMORY_CONTEXT_TAIL_CHARS:]
|
||||
|
||||
|
||||
def automatic_compaction_status_message(
|
||||
engine: Any,
|
||||
*,
|
||||
phase: str,
|
||||
default_message: str,
|
||||
**context: Any,
|
||||
) -> str | None:
|
||||
"""Resolve host-visible status for an automatic compaction event.
|
||||
def automatic_compaction_status_message(engine: Any, *, phase: str, default_message: str, **context: Any) -> str | None:
|
||||
"""Host-visible status for an automatic compaction event; ``None`` = emit nothing.
|
||||
|
||||
Engines can suppress routine automatic status with
|
||||
``emit_automatic_compaction_status = False`` or customize it by defining
|
||||
``get_automatic_compaction_status_message(...)``. Empty strings and
|
||||
``None`` mean "do not emit a lifecycle status".
|
||||
Engines suppress via ``emit_automatic_compaction_status = False`` or
|
||||
customize via ``get_automatic_compaction_status_message(...)``.
|
||||
"""
|
||||
if not getattr(engine, "emit_automatic_compaction_status", True):
|
||||
return None
|
||||
|
||||
formatter = getattr(engine, "get_automatic_compaction_status_message", None)
|
||||
if callable(formatter):
|
||||
message = formatter(
|
||||
phase=phase,
|
||||
default_message=default_message,
|
||||
**context,
|
||||
)
|
||||
else:
|
||||
message = default_message
|
||||
|
||||
message = formatter(phase=phase, default_message=default_message, **context) if callable(formatter) else default_message
|
||||
if message is None:
|
||||
return None
|
||||
message = str(message).strip()
|
||||
return message or None
|
||||
return str(message).strip() or None
|
||||
|
||||
|
||||
class ContextEngine(ABC):
|
||||
"""Base class all context engines must implement."""
|
||||
|
||||
# -- Identity ----------------------------------------------------------
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Short identifier (e.g. 'compressor', 'lcm')."""
|
||||
|
||||
# -- Token state (read by run_agent.py for display/logging) ------------
|
||||
#
|
||||
# Engines MUST maintain these. run_agent.py reads them directly.
|
||||
|
||||
# Token state: engines MUST maintain these; run_agent.py reads them directly.
|
||||
last_prompt_tokens: int = 0
|
||||
last_completion_tokens: int = 0
|
||||
last_total_tokens: int = 0
|
||||
threshold_tokens: int = 0
|
||||
context_length: int = 0
|
||||
compression_count: int = 0
|
||||
|
||||
# -- Compaction parameters (read by run_agent.py for preflight) --------
|
||||
#
|
||||
# These control the preflight compression check. Subclasses may
|
||||
# override via __init__ or property; defaults are sensible for most
|
||||
# engines.
|
||||
#
|
||||
# protect_first_n semantics (since PR #13754): count of non-system head
|
||||
# messages always preserved verbatim, IN ADDITION to the system prompt
|
||||
# which is always implicitly protected. Default 3 keeps the
|
||||
# historical "system + first 3 non-system messages" head shape.
|
||||
|
||||
# Compaction parameters (read by run_agent.py for preflight). protect_first_n counts
|
||||
# non-system head messages kept verbatim IN ADDITION to the always-protected system
|
||||
# prompt (3 keeps the historical head shape).
|
||||
# These control the preflight compression check. Subclasses may override via __init__ or property;
|
||||
# defaults are sensible for most engines. See #13754.
|
||||
threshold_percent: float = 0.75
|
||||
protect_first_n: int = 3
|
||||
protect_last_n: int = 6
|
||||
|
||||
# User-visible lifecycle status for automatic host-triggered compaction.
|
||||
# Alternative engines that treat compaction as routine background
|
||||
# maintenance can set this false to keep successful automatic passes silent;
|
||||
# warnings, errors, and explicit manual commands should still surface.
|
||||
# False keeps successful automatic compaction passes silent (routine background
|
||||
# maintenance); warnings, errors and manual /compress still surface.
|
||||
emit_automatic_compaction_status: bool = True
|
||||
|
||||
# -- Core interface ----------------------------------------------------
|
||||
|
||||
@abstractmethod
|
||||
def update_from_response(self, usage: Dict[str, Any]) -> None:
|
||||
"""Update tracked token usage from an API response.
|
||||
"""Update tracked token usage after every LLM call.
|
||||
|
||||
Called after every LLM call with a normalized usage dict. The legacy
|
||||
keys ``prompt_tokens``, ``completion_tokens``, and ``total_tokens``
|
||||
are always present. Newer hosts also include canonical buckets:
|
||||
``input_tokens``, ``output_tokens``, ``cache_read_tokens``,
|
||||
``cache_write_tokens``, and ``reasoning_tokens``. Engines should
|
||||
treat those fields as optional for compatibility with older hosts.
|
||||
``prompt_tokens``/``completion_tokens``/``total_tokens`` are always present; the
|
||||
canonical buckets (``input_tokens``, ``output_tokens``, ``cache_read_tokens``,
|
||||
``cache_write_tokens``, ``reasoning_tokens``) are optional on older hosts.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -149,341 +87,156 @@ class ContextEngine(ABC):
|
||||
def should_compress_info(self, prompt_tokens: int = None) -> "tuple[bool, str | None]":
|
||||
"""Return ``(should_compress, reason)``.
|
||||
|
||||
The base implementation is backward-compatible: engines that only
|
||||
implement ``should_compress`` get ``(should_compress(prompt_tokens),
|
||||
None)``. Concrete engines with richer block reasons (e.g. a
|
||||
summary-LLM cooldown or an anti-thrashing guard) override this to
|
||||
surface a human-readable reason so callers can warn the user instead
|
||||
of silently skipping compression. Added for the silent-overflow
|
||||
warning fix (#62625) so plugin engines don't raise AttributeError.
|
||||
Engines with block reasons (summary-LLM cooldown, anti-thrashing guard) override
|
||||
this so callers can warn instead of silently skipping; the default keeps plugin
|
||||
engines from raising AttributeError.
|
||||
"""
|
||||
return self.should_compress(prompt_tokens), None
|
||||
|
||||
@abstractmethod
|
||||
def compress(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
current_tokens: Optional[int] = None,
|
||||
focus_topic: Optional[str] = None,
|
||||
force: bool = False,
|
||||
memory_context: str = "",
|
||||
self, messages: List[Dict[str, Any]], current_tokens: Optional[int] = None,
|
||||
focus_topic: Optional[str] = None, force: bool = False, memory_context: str = "",
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Compact the message list and return the new message list.
|
||||
"""Compact ``messages`` into a valid OpenAI-format list that fits the budget.
|
||||
|
||||
This is the main entry point. The engine receives the full message
|
||||
list and returns a (possibly shorter) list that fits within the
|
||||
context budget. The implementation is free to summarize, build a
|
||||
DAG, or do anything else — as long as the returned list is a valid
|
||||
OpenAI-format message sequence.
|
||||
|
||||
Args:
|
||||
focus_topic: Optional topic string from manual ``/compress <focus>``.
|
||||
Engines that support guided compression should prioritise
|
||||
preserving information related to this topic. Engines that
|
||||
don't support it may simply ignore this argument.
|
||||
force: Whether a user-requested compression should bypass an
|
||||
engine-owned cooldown. Engines without cooldowns may ignore it.
|
||||
memory_context: Text returned by memory providers immediately before
|
||||
compaction. Summarizing engines should include non-empty text in
|
||||
their handoff prompt. Older engines may omit this parameter; the
|
||||
host filters unsupported optional arguments by signature.
|
||||
``focus_topic`` comes from manual ``/compress <focus>`` (prioritise that topic);
|
||||
``force`` asks to bypass an engine-owned cooldown; ``memory_context`` is provider
|
||||
text for the handoff prompt. Older engines may omit optional parameters — the
|
||||
host filters them by signature.
|
||||
"""
|
||||
|
||||
# -- Optional: proactive tool-result prune -----------------------------
|
||||
|
||||
def prune_tool_results_only(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
current_tokens: int | None = None,
|
||||
self, messages: List[Dict[str, Any]], current_tokens: int | None = None,
|
||||
) -> tuple[List[Dict[str, Any]], int]:
|
||||
"""Deterministically trim old tool-result payloads without an LLM call.
|
||||
|
||||
Runs on a low, cost-oriented trigger independent of ``should_compress``
|
||||
so large-window engines can reclaim re-sent tool output long before full
|
||||
compaction would fire. Returns ``(messages, n_pruned)``.
|
||||
|
||||
Default is a safe no-op: the list is returned unchanged with ``0``
|
||||
pruned. Engines that don't implement a cheap prune — and any engine that
|
||||
predates this hook — inherit this default, so the agent loop's
|
||||
post-tool-call prune path never raises ``AttributeError`` on them. The
|
||||
built-in ContextCompressor overrides this with the real implementation.
|
||||
Runs on a low, cost-oriented trigger independent of ``should_compress`` so
|
||||
large-window engines reclaim re-sent tool output long before full compaction.
|
||||
Returns ``(messages, n_pruned)``; the default no-op keeps older engines safe.
|
||||
"""
|
||||
return messages, 0
|
||||
|
||||
# -- Optional: per-turn context selection (distinct from compression) --
|
||||
|
||||
def select_context(
|
||||
self,
|
||||
request_messages: List[Dict[str, Any]],
|
||||
*,
|
||||
conversation_messages: List[Dict[str, Any]] = None,
|
||||
incoming_message: Dict[str, Any] = None,
|
||||
budget_tokens: int = 0,
|
||||
self, request_messages: List[Dict[str, Any]], *, conversation_messages: List[Dict[str, Any]] = None,
|
||||
incoming_message: Dict[str, Any] = None, budget_tokens: int = 0,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Optionally choose/replace the context for THIS request, pre-generation.
|
||||
"""Optionally *select* (replace) the context for THIS request, pre-generation.
|
||||
|
||||
Called every turn after the request message list is assembled and
|
||||
before it is dispatched to the provider — independent of
|
||||
``should_compress()``. This lets an engine *select* which context
|
||||
enters the prompt (retrieval, topic routing, role/branch switching)
|
||||
rather than *shrink* context that is already there. The two verbs are
|
||||
orthogonal:
|
||||
|
||||
- ``compress()`` : context is too long -> make it shorter.
|
||||
- ``select_context()``: this turn belongs to a different context
|
||||
-> use that one instead.
|
||||
|
||||
Without this hook, engines that need per-turn access to the message
|
||||
list have to force ``should_compress()`` to return ``True`` so that
|
||||
``compress()`` is invoked every turn purely as a callback — which
|
||||
conflates selection with compression and degrades behaviour when the
|
||||
engine's backend is unavailable. ``select_context()`` removes the need
|
||||
for that workaround.
|
||||
|
||||
The returned list is request-only: it replaces the messages sent to
|
||||
the provider for this single call and MUST NOT be treated as persisted
|
||||
transcript state. The conversation history in the session DB is left
|
||||
untouched, so nothing leaks across turns. Return ``None`` to leave the
|
||||
request unchanged.
|
||||
|
||||
Unlike the ``pre_llm_call`` plugin hook (which appends to the user
|
||||
message and intentionally never rewrites the list, to preserve the
|
||||
cache prefix), ``select_context()`` may *replace* the message list.
|
||||
|
||||
Ordering / cache contract: the host runs this hook **before** prompt
|
||||
cache-control and **before** every request sanitizer (orphaned-tool
|
||||
cleanup, thinking-only/role normalization, whitespace/JSON
|
||||
normalization). So (a) whatever the hook returns still passes through
|
||||
the same validation as any request — a malformed replacement cannot
|
||||
reach the provider — and (b) prompt-cache stability (an AGENTS.md
|
||||
invariant) is preserved: the default no-op leaves the request
|
||||
byte-identical, so cache behaviour is unchanged for the built-in
|
||||
compressor and any non-implementing engine. An engine that *does*
|
||||
replace the list changes its own cache prefix by definition; that is
|
||||
the engine's concern, and cache-control breakpoints are re-derived on
|
||||
the selected list. The hook is evaluated per provider request (so it
|
||||
re-runs on retries within a turn), consistent with "select the context
|
||||
for THIS request".
|
||||
|
||||
Args:
|
||||
request_messages: The assembled request message list (system
|
||||
prompt + history + any ephemeral prefill), in OpenAI format.
|
||||
conversation_messages: The unmodified persisted conversation
|
||||
history, for reference only (do not mutate).
|
||||
incoming_message: The current turn's user message, if available.
|
||||
budget_tokens: The active model's context length, or 0 if unknown.
|
||||
|
||||
Default returns ``None`` (no-op) — zero impact on the built-in
|
||||
compressor or any existing engine.
|
||||
Runs on every provider request (also retries), independent of
|
||||
``should_compress()``: ``compress()`` shrinks over-long context, this swaps in a
|
||||
different one (retrieval, topic routing, branch switching). Return ``None`` to
|
||||
leave the request unchanged. The returned list is request-only — it MUST NOT be
|
||||
treated as persisted transcript state (session DB history is untouched); unlike
|
||||
``pre_llm_call`` it may replace the list. The host runs it before prompt
|
||||
cache-control and every request sanitizer, so a malformed replacement never
|
||||
reaches the provider and the default no-op keeps the request byte-identical;
|
||||
an engine that replaces the list changes its own cache prefix (breakpoints are
|
||||
re-derived on the selected list). ``request_messages`` is the assembled request
|
||||
(system prompt + history + ephemeral prefill); ``conversation_messages`` is the
|
||||
persisted history for reference only (do not mutate); ``budget_tokens`` is the
|
||||
model's context length or 0 if unknown.
|
||||
"""
|
||||
return None
|
||||
|
||||
def on_turn_complete(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
usage: Dict[str, Any] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Observe a finished user turn (post-turn ingestion / observation).
|
||||
def on_turn_complete(self, messages: List[Dict[str, Any]], usage: Dict[str, Any] = None, **kwargs: Any) -> None:
|
||||
"""Observe a finished turn (complement of ``select_context()``) to index/update
|
||||
routing state for the next request.
|
||||
|
||||
Called from the standard turn-finalization path once the assistant/tool
|
||||
loop completes, with the finalized in-memory transcript snapshot. This
|
||||
is the complement to ``select_context()``: selection happens *before*
|
||||
the request, while observation happens *after* the turn. It lets an
|
||||
engine ingest, index, summarize, or update routing / topic / session
|
||||
state from what actually happened — so the next ``select_context()``
|
||||
can act on it.
|
||||
|
||||
Coverage: this fires from the normal finalization seam. Some abnormal
|
||||
early-return paths in the loop (e.g. a content-policy block or a
|
||||
provider terminal failure) persist and return without routing through
|
||||
finalization, and therefore do not currently emit this hook. Treat it
|
||||
as a best-effort post-turn observation for completed turns, not a
|
||||
guaranteed callback for every possible early exit; unifying all
|
||||
terminal paths behind one finalization seam is a separate follow-up.
|
||||
|
||||
Together the two hooks remove the need to abuse ``should_compress()`` /
|
||||
``compress()`` as a generic per-turn callback just to observe history,
|
||||
and they cover the case where a turn finishes and there may be no next
|
||||
request from which to infer the previous turn.
|
||||
|
||||
``messages`` is a shallow copy and should be treated as read-only:
|
||||
return values are ignored and this hook must not rely on transcript
|
||||
mutation for persistence. ``kwargs`` may include ``turn_id``,
|
||||
``task_id``, ``api_call_count``, ``interrupted``, ``failed``, and
|
||||
``turn_exit_reason``.
|
||||
|
||||
``usage`` carries the completed turn's canonical token usage (the same
|
||||
dict shape passed to ``update_from_response`` — ``prompt_tokens`` /
|
||||
``completion_tokens`` / ``total_tokens`` plus the canonical
|
||||
``input_tokens`` / ``output_tokens`` / ``cache_read_tokens`` /
|
||||
``cache_write_tokens`` / ``reasoning_tokens`` buckets) so an engine can
|
||||
weigh how large/expensive the selected context actually was when
|
||||
deciding the next ``select_context()``. It is ``None`` on finalized
|
||||
turns that never reached a provider response (e.g. interrupt); engines
|
||||
must treat it as optional.
|
||||
|
||||
Default is a no-op.
|
||||
Best-effort, not guaranteed: fires from the normal finalization seam only; some
|
||||
abnormal early returns (content-policy block, provider terminal failure) skip it.
|
||||
``messages`` is a read-only shallow copy (return value ignored; never rely on
|
||||
transcript mutation). ``usage`` has the ``update_from_response`` shape and is
|
||||
``None`` when no provider response was reached (interrupt). ``kwargs`` may include
|
||||
``turn_id``, ``task_id``, ``api_call_count``, ``interrupted``, ``failed``, ``turn_exit_reason``.
|
||||
"""
|
||||
return None
|
||||
|
||||
# -- Optional: pre-flight check ----------------------------------------
|
||||
|
||||
def should_compress_preflight(self, messages: List[Dict[str, Any]]) -> bool:
|
||||
"""Quick rough check before the API call (no real token count yet).
|
||||
|
||||
Default returns False (skip pre-flight). Override if your engine
|
||||
can do a cheap estimate.
|
||||
"""
|
||||
"""Cheap rough check before the API call (no real token count yet); default skips."""
|
||||
return False
|
||||
|
||||
def should_defer_preflight_to_real_usage(self, rough_tokens: int) -> bool:
|
||||
"""Return True when preflight should trust recent real usage instead.
|
||||
|
||||
Built-in compression uses this to avoid re-compacting from known-noisy
|
||||
rough estimates after a compressed request has already fit. Third-party
|
||||
engines can ignore it safely.
|
||||
"""
|
||||
"""True when preflight should trust recent real usage over the noisy rough
|
||||
estimate (avoids re-compacting after a compressed request already fit)."""
|
||||
return False
|
||||
|
||||
def get_automatic_compaction_status_message(
|
||||
self,
|
||||
*,
|
||||
phase: str,
|
||||
default_message: str,
|
||||
**context: Any,
|
||||
self, *, phase: str, default_message: str, **context: Any,
|
||||
) -> str | None:
|
||||
"""Return user-visible status for automatic host-triggered compaction.
|
||||
"""User-visible status for automatic compaction, or ``None`` to suppress it.
|
||||
|
||||
Return ``None`` to suppress successful automatic lifecycle status for
|
||||
this compaction event. ``phase`` identifies the host call site (for
|
||||
example ``"preflight"`` or ``"compress"``). ``context`` contains
|
||||
best-effort fields such as ``approx_tokens`` and ``threshold_tokens``.
|
||||
|
||||
This hook does not control warning/error messages or explicit manual
|
||||
commands such as ``/compress``.
|
||||
``phase`` is the host call site (``"preflight"`` / ``"compress"``); ``context``
|
||||
carries best-effort ``approx_tokens`` / ``threshold_tokens``. Warnings, errors
|
||||
and manual ``/compress`` are not governed by this hook.
|
||||
"""
|
||||
if not self.emit_automatic_compaction_status:
|
||||
return None
|
||||
return default_message
|
||||
|
||||
# -- Optional: manual /compress preflight ------------------------------
|
||||
return default_message if self.emit_automatic_compaction_status else None
|
||||
|
||||
def has_content_to_compress(self, messages: List[Dict[str, Any]]) -> bool:
|
||||
"""Quick check: is there anything in ``messages`` that can be compacted?
|
||||
|
||||
Used by the gateway ``/compress`` command as a preflight guard —
|
||||
returning False lets the gateway report "nothing to compress yet"
|
||||
without making an LLM call.
|
||||
|
||||
Default returns True (always attempt). Engines with a cheap way
|
||||
to introspect their own head/tail boundaries should override this
|
||||
to return False when the transcript is still entirely protected.
|
||||
"""
|
||||
"""Preflight guard for gateway ``/compress``: False reports "nothing to
|
||||
compress yet" without an LLM call (e.g. transcript entirely protected)."""
|
||||
return True
|
||||
|
||||
# -- Optional: session lifecycle ---------------------------------------
|
||||
|
||||
def on_session_start(self, session_id: str, **kwargs) -> None:
|
||||
"""Called when a new conversation session begins.
|
||||
|
||||
Use this to load persisted state (DAG, store) for the session.
|
||||
kwargs may include hermes_home, platform, model, etc.
|
||||
"""
|
||||
"""Session begins: load persisted state. kwargs may include hermes_home, platform, model."""
|
||||
|
||||
def on_session_end(self, session_id: str, messages: List[Dict[str, Any]]) -> None:
|
||||
"""Called at real session boundaries (CLI exit, /reset, gateway expiry).
|
||||
|
||||
Use this to flush state, close DB connections, etc.
|
||||
NOT called per-turn — only when the session truly ends.
|
||||
"""
|
||||
"""Real session boundary (CLI exit, /reset, gateway expiry) — never per-turn."""
|
||||
|
||||
def on_session_reset(self) -> None:
|
||||
"""Called on /new or /reset. Reset per-session state.
|
||||
|
||||
Default resets compression_count and token tracking.
|
||||
"""
|
||||
"""/new or /reset: reset per-session state (default: counters and token tracking)."""
|
||||
# Reset cross-call calibration state captured under the PREVIOUS model. These fields encode "the
|
||||
# provider proved this prompt fit" / "preflight can be deferred" decisions that are only valid for
|
||||
# the model that produced them. Carrying them across a switch to a smaller-context model would let
|
||||
# should_defer_preflight_to_real_usage() suppress a preflight compression the new model actually
|
||||
# needs — the exact oversized-send-after-switch failure in #23767. The new model's first response
|
||||
# repopulates them via update_from_response(). Setting last_prompt_tokens to 0 (NOT -1) is
|
||||
# deliberate: 0 is the documented "no real usage yet -> use the rough estimate" state, so the post-
|
||||
# response should_compress path falls back to estimate_request_tokens_rough rather than skipping
|
||||
# compression. -1 is a different sentinel (#36718, "compression just ran, await real usage") and
|
||||
# must not be set here.
|
||||
self.last_prompt_tokens = 0
|
||||
self.last_completion_tokens = 0
|
||||
self.last_total_tokens = 0
|
||||
self.compression_count = 0
|
||||
|
||||
# -- Optional: tools ---------------------------------------------------
|
||||
|
||||
def get_tool_schemas(self) -> List[Dict[str, Any]]:
|
||||
"""Return tool schemas this engine provides to the agent.
|
||||
|
||||
Default returns empty list (no tools). LCM would return schemas
|
||||
for lcm_grep, lcm_describe, lcm_expand here.
|
||||
"""
|
||||
"""Tool schemas this engine exposes to the agent (default: none)."""
|
||||
return []
|
||||
|
||||
def handle_tool_call(self, name: str, args: Dict[str, Any], **kwargs) -> str:
|
||||
"""Handle a tool call from the agent.
|
||||
|
||||
Only called for tool names returned by get_tool_schemas().
|
||||
Must return a JSON string.
|
||||
|
||||
kwargs may include:
|
||||
messages: the current in-memory message list (for live ingestion)
|
||||
"""
|
||||
import json
|
||||
"""Handle a call to one of this engine's tools; must return a JSON string.
|
||||
kwargs may include ``messages`` (live in-memory list)."""
|
||||
return json.dumps({"error": f"Unknown context engine tool: {name}"})
|
||||
|
||||
# -- Optional: status / display ----------------------------------------
|
||||
|
||||
def get_status(self) -> Dict[str, Any]:
|
||||
"""Return status dict for display/logging.
|
||||
|
||||
Default returns the standard fields run_agent.py expects.
|
||||
"""
|
||||
# Clamp the -1 "compression just ran, awaiting real usage" sentinel
|
||||
# (set by conversation_compression) to 0 so status readers don't see a
|
||||
# raw -1 or a negative usage_percent on the transitional turn. Mirrors
|
||||
# the CLI/gateway status-bar paths (cli.py, tui_gateway/server.py).
|
||||
last_prompt = self.last_prompt_tokens if self.last_prompt_tokens > 0 else 0
|
||||
"""Status dict with the standard fields run_agent.py expects."""
|
||||
# Clamp the -1 "compression just ran, awaiting real usage" sentinel to 0 so no
|
||||
# reader sees a negative usage_percent on the transitional turn.
|
||||
last_prompt = max(self.last_prompt_tokens, 0)
|
||||
return {
|
||||
"last_prompt_tokens": last_prompt,
|
||||
"threshold_tokens": self.threshold_tokens,
|
||||
"context_length": self.context_length,
|
||||
"usage_percent": (
|
||||
min(100, last_prompt / self.context_length * 100)
|
||||
if self.context_length else 0
|
||||
),
|
||||
"usage_percent": min(100, last_prompt / self.context_length * 100) if self.context_length else 0,
|
||||
"compression_count": self.compression_count,
|
||||
}
|
||||
|
||||
# -- Optional: model switch support ------------------------------------
|
||||
|
||||
def update_model(
|
||||
self,
|
||||
model: str,
|
||||
context_length: int,
|
||||
base_url: str = "",
|
||||
api_key: str = "",
|
||||
provider: str = "",
|
||||
api_mode: str = "",
|
||||
self, model: str, context_length: int, base_url: str = "", api_key: str = "",
|
||||
provider: str = "", api_mode: str = "",
|
||||
) -> None:
|
||||
"""Called when the user switches models or on fallback activation.
|
||||
"""Model switch / fallback: recompute threshold_tokens (override for more).
|
||||
|
||||
Default updates context_length and recalculates threshold_tokens
|
||||
from threshold_percent. Override if your engine needs more
|
||||
(e.g. recalculate DAG budgets, switch summary models).
|
||||
Per-model threshold override (longest substring match), else the raw config
|
||||
percent — snapshotted ONCE so repeated switches fall back to the configured
|
||||
value, not the previous model's override.
|
||||
"""
|
||||
self.context_length = context_length
|
||||
# Apply per-model threshold overrides if set (longest substring match).
|
||||
# Falls back to _config_threshold_percent (the raw config value) when
|
||||
# no override matches. Plugin engines that override update_model() can
|
||||
# call resolve_model_threshold() for the same logic.
|
||||
from agent.context_compressor import resolve_model_threshold
|
||||
if not hasattr(self, "_config_threshold_percent"):
|
||||
# Snapshot the pre-override percent ONCE so repeated model
|
||||
# switches fall back to the engine's configured value, not the
|
||||
# previous model's override.
|
||||
self._config_threshold_percent = self.threshold_percent
|
||||
self._base_threshold_percent = resolve_model_threshold(
|
||||
model, getattr(self, "model_thresholds", {}),
|
||||
self._config_threshold_percent,
|
||||
)
|
||||
model, getattr(self, "model_thresholds", {}), self._config_threshold_percent)
|
||||
self.threshold_percent = self._base_threshold_percent
|
||||
self.threshold_tokens = int(context_length * self.threshold_percent)
|
||||
|
||||
+221
-444
@@ -1,3 +1,5 @@
|
||||
"""@-reference expansion (``@file:``, ``@folder:``, ``@diff``, ``@git:``, ``@url:`` + plugin prefixes)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -7,20 +9,19 @@ import mimetypes
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Awaitable, Callable
|
||||
|
||||
from agent.model_metadata import estimate_tokens_rough
|
||||
from hermes_cli._subprocess_compat import IS_WINDOWS, windows_hide_flags
|
||||
from hermes_cli._subprocess_compat import IS_WINDOWS, harden_git_argv, noninteractive_git_env, windows_hide_flags
|
||||
from hermes_cli.sizefmt import format_bytes
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Plugin context-reference provider API (Issue #26193)
|
||||
# ---------------------------------------------------------------------------
|
||||
# ── Plugin context-reference provider API ────────────────────────────────────
|
||||
|
||||
# --------------------------------------------------------------------------- Plugin context-reference
|
||||
# provider API (Issue #26193) ---------------------------------------------------------------------------
|
||||
BUILTIN_PREFIXES = frozenset({"diff", "staged", "file", "folder", "git", "url"})
|
||||
|
||||
_context_reference_providers: dict[str, "ContextReferenceProvider"] = {}
|
||||
@@ -38,11 +39,7 @@ class ContextCompletionItem:
|
||||
|
||||
|
||||
class ContextReferenceProvider(ABC):
|
||||
"""Base class for plugin-registered @-prefix context reference providers.
|
||||
|
||||
Plugins subclass this and register via
|
||||
``PluginContext.register_context_reference()``.
|
||||
"""
|
||||
"""Base class for plugin @-prefix providers, registered via ``PluginContext.register_context_reference()``."""
|
||||
|
||||
prefix: str = "" # e.g. "issue", "channel", "doc"
|
||||
description: str = "" # shown in autocomplete meta column
|
||||
@@ -50,12 +47,10 @@ class ContextReferenceProvider(ABC):
|
||||
@abstractmethod
|
||||
async def autocomplete(self, query: str, *, limit: int = 10) -> list[ContextCompletionItem]:
|
||||
"""Return autocomplete items for the given query string."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def expand(self, target: str) -> str | None:
|
||||
"""Expand *target* to prompt content. Return ``None`` to skip."""
|
||||
...
|
||||
|
||||
|
||||
def register_context_reference_provider(provider: ContextReferenceProvider) -> None:
|
||||
@@ -81,31 +76,29 @@ _QUOTED_REFERENCE_VALUE = r'(?:`[^`\n]+`|"[^"\n]+"|\'[^\'\n]+\')'
|
||||
REFERENCE_PATTERN = re.compile(
|
||||
rf"(?<![\w/])@(?:(?P<simple>diff|staged)\b|(?P<kind>file|folder|git|url):(?P<value>{_QUOTED_REFERENCE_VALUE}(?::\d+(?:-\d+)?)?|\S+))"
|
||||
)
|
||||
# Plugin fallback pattern – catches any @<word>:<value> not handled by the
|
||||
# built-in regex so that plugin-registered prefixes can be resolved.
|
||||
# Plugin fallback: any @<word>:<value> the built-in regex did not claim.
|
||||
_PLUGIN_REFERENCE_PATTERN = re.compile(
|
||||
rf"(?<![\w/])@(?P<kind>[a-zA-Z][a-zA-Z0-9_-]*):(?P<value>{_QUOTED_REFERENCE_VALUE}(?::\d+(?:-\d+)?)?|\S+)"
|
||||
)
|
||||
# ``@file:`` value: quoted path or bare path, each with an optional ``:start[-end]`` range.
|
||||
_FILE_VALUE_PATTERN = re.compile(
|
||||
r'^(?:(?P<quote>`|"|\')(?P<qpath>.+?)(?P=quote)|(?P<path>.+?))(?::(?P<start>\d+)(?:-(?P<end>\d+))?)?$'
|
||||
)
|
||||
|
||||
TRAILING_PUNCTUATION = ",.;!?"
|
||||
_OPENERS = {")": "(", "]": "[", "}": "{"}
|
||||
_NEEDS_QUOTING = re.compile(r"""[\s()\[\]{}<>"'`]""")
|
||||
_SENSITIVE_HOME_DIRS = (".ssh", ".aws", ".gnupg", ".kube", ".docker", ".azure", ".config/gh")
|
||||
_SENSITIVE_HERMES_DIRS = (Path("skills") / ".hub",)
|
||||
_SENSITIVE_HOME_FILES = (
|
||||
Path(".ssh") / "authorized_keys",
|
||||
Path(".ssh") / "id_rsa",
|
||||
Path(".ssh") / "id_ed25519",
|
||||
Path(".ssh") / "config",
|
||||
Path(".bashrc"),
|
||||
Path(".zshrc"),
|
||||
Path(".profile"),
|
||||
Path(".bash_profile"),
|
||||
Path(".zprofile"),
|
||||
Path(".netrc"),
|
||||
Path(".pgpass"),
|
||||
Path(".npmrc"),
|
||||
Path(".pypirc"),
|
||||
)
|
||||
_SENSITIVE_HOME_FILES = tuple(Path(p) for p in (
|
||||
".ssh/authorized_keys", ".ssh/id_rsa", ".ssh/id_ed25519", ".ssh/config", ".bashrc", ".zshrc",
|
||||
".profile", ".bash_profile", ".zprofile", ".netrc", ".pgpass", ".npmrc", ".pypirc",
|
||||
))
|
||||
_TEXT_EXTENSIONS = (".py", ".md", ".txt", ".json", ".yaml", ".yml", ".toml", ".js", ".ts")
|
||||
_FENCE_LANGUAGES = {
|
||||
".py": "python", ".js": "javascript", ".ts": "typescript", ".tsx": "tsx", ".jsx": "jsx",
|
||||
".json": "json", ".md": "markdown", ".sh": "bash", ".yml": "yaml", ".yaml": "yaml", ".toml": "toml",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -130,13 +123,13 @@ class ContextReferenceResult:
|
||||
blocked: bool = False
|
||||
|
||||
|
||||
def format_reference_value(value: str) -> str:
|
||||
"""Quote a reference value so ``REFERENCE_PATTERN`` reads it back whole.
|
||||
UrlFetcher = Callable[[str], str | Awaitable[str]] | None
|
||||
Expansion = tuple[str | None, str | None] # (warning, block) — exactly one side is set
|
||||
|
||||
The unquoted alternative in the pattern is ``\\S+``, so a path containing a
|
||||
space parses as a truncated ref with the tail left behind as loose text.
|
||||
Mirrors ``formatRefValue`` in the desktop's directive-text.tsx.
|
||||
"""
|
||||
|
||||
def format_reference_value(value: str) -> str:
|
||||
"""Quote a value so ``REFERENCE_PATTERN`` (bare alternative ``\\S+``) reads it back whole.
|
||||
Mirrors ``formatRefValue`` in the desktop's directive-text.tsx."""
|
||||
if not _NEEDS_QUOTING.search(value):
|
||||
return value
|
||||
for quote in ("`", '"', "'"):
|
||||
@@ -149,201 +142,112 @@ def parse_context_references(message: str) -> list[ContextReference]:
|
||||
refs: list[ContextReference] = []
|
||||
if not message:
|
||||
return refs
|
||||
|
||||
for match in REFERENCE_PATTERN.finditer(message):
|
||||
simple = match.group("simple")
|
||||
if simple:
|
||||
refs.append(
|
||||
ContextReference(
|
||||
raw=match.group(0),
|
||||
kind=simple,
|
||||
target="",
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
kind = match.group("kind")
|
||||
kind = match.group("simple") or match.group("kind")
|
||||
value = _strip_trailing_punctuation(match.group("value") or "")
|
||||
line_start = None
|
||||
line_end = None
|
||||
target = _strip_reference_wrappers(value)
|
||||
|
||||
if kind == "file":
|
||||
if match.group("simple"):
|
||||
target, line_start, line_end = "", None, None
|
||||
elif kind == "file":
|
||||
target, line_start, line_end = _parse_file_reference_value(value)
|
||||
else:
|
||||
target, line_start, line_end = _strip_reference_wrappers(value), None, None
|
||||
refs.append(ContextReference(match.group(0), kind, target, match.start(), match.end(), line_start, line_end))
|
||||
|
||||
refs.append(
|
||||
ContextReference(
|
||||
raw=match.group(0),
|
||||
kind=kind,
|
||||
target=target,
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
line_start=line_start,
|
||||
line_end=line_end,
|
||||
)
|
||||
)
|
||||
|
||||
# Second pass: resolve plugin-registered prefixes the built-in pattern missed
|
||||
if _context_reference_providers:
|
||||
for match in _PLUGIN_REFERENCE_PATTERN.finditer(message):
|
||||
kind = match.group("kind")
|
||||
if kind in BUILTIN_PREFIXES:
|
||||
continue
|
||||
# Skip if already captured by the built-in pattern
|
||||
if any(r.kind == kind and r.start == match.start() for r in refs):
|
||||
continue
|
||||
if kind in _context_reference_providers:
|
||||
value = _strip_trailing_punctuation(match.group("value") or "")
|
||||
refs.append(
|
||||
ContextReference(
|
||||
raw=match.group(0),
|
||||
kind=kind,
|
||||
target=_strip_reference_wrappers(value),
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
)
|
||||
)
|
||||
|
||||
# Second pass: plugin-registered prefixes the built-in pattern missed.
|
||||
for match in _PLUGIN_REFERENCE_PATTERN.finditer(message) if _context_reference_providers else ():
|
||||
kind = match.group("kind")
|
||||
if kind in BUILTIN_PREFIXES or kind not in _context_reference_providers:
|
||||
continue
|
||||
if any(r.kind == kind and r.start == match.start() for r in refs):
|
||||
continue
|
||||
target = _strip_reference_wrappers(_strip_trailing_punctuation(match.group("value") or ""))
|
||||
refs.append(ContextReference(match.group(0), kind, target, match.start(), match.end()))
|
||||
return refs
|
||||
|
||||
|
||||
def preprocess_context_references(
|
||||
message: str,
|
||||
*,
|
||||
cwd: str | Path,
|
||||
context_length: int,
|
||||
url_fetcher: Callable[[str], str | Awaitable[str]] | None = None,
|
||||
message: str, *, cwd: str | Path, context_length: int, url_fetcher: UrlFetcher = None,
|
||||
allowed_root: str | Path | None = None,
|
||||
) -> ContextReferenceResult:
|
||||
"""Sync wrapper; safe both without a loop (CLI) and inside a running loop (gateway)."""
|
||||
coro = preprocess_context_references_async(
|
||||
message,
|
||||
cwd=cwd,
|
||||
context_length=context_length,
|
||||
url_fetcher=url_fetcher,
|
||||
allowed_root=allowed_root,
|
||||
message, cwd=cwd, context_length=context_length, url_fetcher=url_fetcher, allowed_root=allowed_root
|
||||
)
|
||||
# Safe for both CLI (no loop) and gateway (loop already running).
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
loop = None
|
||||
if loop and loop.is_running():
|
||||
import concurrent.futures
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||
return pool.submit(asyncio.run, coro).result()
|
||||
return asyncio.run(coro)
|
||||
return asyncio.run(coro)
|
||||
import concurrent.futures
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||
return pool.submit(asyncio.run, coro).result()
|
||||
|
||||
|
||||
async def preprocess_context_references_async(
|
||||
message: str,
|
||||
*,
|
||||
cwd: str | Path,
|
||||
context_length: int,
|
||||
url_fetcher: Callable[[str], str | Awaitable[str]] | None = None,
|
||||
message: str, *, cwd: str | Path, context_length: int, url_fetcher: UrlFetcher = None,
|
||||
allowed_root: str | Path | None = None,
|
||||
) -> ContextReferenceResult:
|
||||
refs = parse_context_references(message)
|
||||
if not refs:
|
||||
return ContextReferenceResult(message=message, original_message=message)
|
||||
|
||||
cwd_path = Path(cwd).expanduser().resolve()
|
||||
# Default to the current working directory so @ references cannot escape
|
||||
# the active workspace unless a caller explicitly widens the root.
|
||||
allowed_root_path = (
|
||||
Path(allowed_root).expanduser().resolve() if allowed_root is not None else cwd_path
|
||||
)
|
||||
warnings: list[str] = []
|
||||
blocks: list[str] = []
|
||||
injected_tokens = 0
|
||||
|
||||
# Expand all references concurrently. Each _expand_reference is independent
|
||||
# (no shared state during expansion) — a message with several @url: refs
|
||||
# would otherwise pay one full web_extract round-trip per ref in series.
|
||||
# gather preserves positional order, so we reassemble warnings/blocks in the
|
||||
# original ref order exactly as the prior serial loop did; the token-budget
|
||||
# check below is unchanged (it runs once, after all refs are expanded).
|
||||
expanded = await asyncio.gather(
|
||||
*(
|
||||
_expand_reference(
|
||||
ref,
|
||||
cwd_path,
|
||||
url_fetcher=url_fetcher,
|
||||
allowed_root=allowed_root_path,
|
||||
)
|
||||
for ref in refs
|
||||
)
|
||||
)
|
||||
for warning, block in expanded:
|
||||
if warning:
|
||||
warnings.append(warning)
|
||||
if block:
|
||||
blocks.append(block)
|
||||
injected_tokens += estimate_tokens_rough(block)
|
||||
|
||||
# Default root = cwd so @ references cannot escape the workspace unless a caller widens it.
|
||||
allowed_root_path = Path(allowed_root).expanduser().resolve() if allowed_root is not None else cwd_path
|
||||
# Expand concurrently (each ref is independent; several @url: refs would otherwise
|
||||
# serialize web_extract round-trips). gather preserves order, so warnings/blocks
|
||||
# are assembled in ref order; the token-budget check runs once afterwards.
|
||||
hard_limit = max(1, int(context_length * 0.50))
|
||||
soft_limit = max(1, int(context_length * 0.25))
|
||||
tasks = (
|
||||
_expand_reference(ref, cwd_path, url_fetcher=url_fetcher, allowed_root=allowed_root_path,
|
||||
max_inline_tokens=hard_limit)
|
||||
for ref in refs
|
||||
)
|
||||
expanded = await asyncio.gather(*tasks)
|
||||
warnings = [warning for warning, _ in expanded if warning]
|
||||
blocks = [block for _, block in expanded if block]
|
||||
injected_tokens = sum(estimate_tokens_rough(block) for block in blocks)
|
||||
result = ContextReferenceResult(
|
||||
message=message, original_message=message, references=refs, warnings=warnings, injected_tokens=injected_tokens
|
||||
)
|
||||
|
||||
if injected_tokens > hard_limit:
|
||||
warnings.append(
|
||||
f"@ context injection refused: {injected_tokens} tokens exceeds the 50% hard limit ({hard_limit})."
|
||||
)
|
||||
return ContextReferenceResult(
|
||||
message=message,
|
||||
original_message=message,
|
||||
references=refs,
|
||||
warnings=warnings,
|
||||
injected_tokens=injected_tokens,
|
||||
expanded=False,
|
||||
blocked=True,
|
||||
)
|
||||
|
||||
warnings.append(f"@ context injection refused: {injected_tokens} tokens exceeds the 50% hard limit ({hard_limit}).")
|
||||
result.blocked = True
|
||||
return result
|
||||
if injected_tokens > soft_limit:
|
||||
warnings.append(
|
||||
f"@ context injection warning: {injected_tokens} tokens exceeds the 25% soft limit ({soft_limit})."
|
||||
)
|
||||
warnings.append(f"@ context injection warning: {injected_tokens} tokens exceeds the 25% soft limit ({soft_limit}).")
|
||||
|
||||
# Leave the `@file:`/`@folder:` tokens where the user typed them. The token
|
||||
# IS the reference, not scaffolding around it: clients render each one as an
|
||||
# inline chip, so stripping them left a sentence with a hole in it ("review
|
||||
# and ship") and made the desktop re-derive the refs from the attached block
|
||||
# to show them as a detached list above the prose.
|
||||
# The `@file:`/`@folder:` tokens stay where the user typed them: the token IS the
|
||||
# reference (clients render it as an inline chip); stripping it left a hole in the
|
||||
# sentence and forced the desktop to re-derive refs from the attached block.
|
||||
final = message
|
||||
if warnings:
|
||||
final = f"{final}\n\n--- Context Warnings ---\n" + "\n".join(f"- {warning}" for warning in warnings)
|
||||
if blocks:
|
||||
final = f"{final}\n\n--- Attached Context ---\n\n" + "\n\n".join(blocks)
|
||||
result.message = final.strip()
|
||||
result.expanded = bool(blocks or warnings)
|
||||
return result
|
||||
|
||||
return ContextReferenceResult(
|
||||
message=final.strip(),
|
||||
original_message=message,
|
||||
references=refs,
|
||||
warnings=warnings,
|
||||
injected_tokens=injected_tokens,
|
||||
expanded=bool(blocks or warnings),
|
||||
blocked=False,
|
||||
)
|
||||
|
||||
# Git-backed reference kinds -> f(ref) -> git argv (the label is "git " + argv).
|
||||
_GIT_REFERENCE_ARGS: dict[str, Callable[[ContextReference], list[str]]] = {
|
||||
"diff": lambda ref: ["diff"],
|
||||
"staged": lambda ref: ["diff", "--staged"],
|
||||
"git": lambda ref: ["log", f"-{max(1, min(int(ref.target or '1'), 10))}", "-p"],
|
||||
}
|
||||
|
||||
|
||||
async def _expand_reference(
|
||||
ref: ContextReference,
|
||||
cwd: Path,
|
||||
*,
|
||||
url_fetcher: Callable[[str], str | Awaitable[str]] | None = None,
|
||||
allowed_root: Path | None = None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
ref: ContextReference, cwd: Path, *, url_fetcher: UrlFetcher = None, allowed_root: Path | None = None,
|
||||
max_inline_tokens: int | None = None,
|
||||
) -> Expansion:
|
||||
try:
|
||||
if ref.kind == "file":
|
||||
return _expand_file_reference(ref, cwd, allowed_root=allowed_root)
|
||||
if ref.kind == "folder":
|
||||
return _expand_folder_reference(ref, cwd, allowed_root=allowed_root)
|
||||
if ref.kind == "diff":
|
||||
return _expand_git_reference(ref, cwd, ["diff"], "git diff")
|
||||
if ref.kind == "staged":
|
||||
return _expand_git_reference(ref, cwd, ["diff", "--staged"], "git diff --staged")
|
||||
if ref.kind == "git":
|
||||
count = max(1, min(int(ref.target or "1"), 10))
|
||||
return _expand_git_reference(ref, cwd, ["log", f"-{count}", "-p"], f"git log -{count} -p")
|
||||
if ref.kind in ("file", "folder"):
|
||||
return _expand_path_reference(ref, cwd, allowed_root=allowed_root, max_inline_tokens=max_inline_tokens)
|
||||
if ref.kind in _GIT_REFERENCE_ARGS:
|
||||
git_args = _GIT_REFERENCE_ARGS[ref.kind](ref)
|
||||
return _expand_git_reference(ref, cwd, git_args, "git " + " ".join(git_args))
|
||||
if ref.kind == "url":
|
||||
content = await _fetch_url_content(ref.target, url_fetcher=url_fetcher)
|
||||
if not content:
|
||||
@@ -351,8 +255,6 @@ async def _expand_reference(
|
||||
return None, f"🌐 {ref.raw} ({estimate_tokens_rough(content)} tokens)\n{content}"
|
||||
except Exception as exc:
|
||||
return f"{ref.raw}: {exc}", None
|
||||
|
||||
# Plugin-provided context references
|
||||
provider = _context_reference_providers.get(ref.kind)
|
||||
if provider is not None:
|
||||
try:
|
||||
@@ -361,96 +263,63 @@ async def _expand_reference(
|
||||
return None, f"📌 {ref.raw} ({estimate_tokens_rough(plugin_content)} tokens)\n{plugin_content}"
|
||||
except Exception as exc:
|
||||
return f"{ref.raw}: plugin expansion error: {exc}", None
|
||||
|
||||
return f"{ref.raw}: unsupported reference type", None
|
||||
|
||||
|
||||
def _expand_file_reference(
|
||||
ref: ContextReference,
|
||||
cwd: Path,
|
||||
*,
|
||||
allowed_root: Path | None = None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
def _expand_path_reference(ref: ContextReference, cwd: Path, *, allowed_root: Path | None = None,
|
||||
max_inline_tokens: int | None = None) -> Expansion:
|
||||
"""``@file:`` / ``@folder:``: resolve, allow-check, then inline text / binary stub / listing."""
|
||||
is_folder = ref.kind == "folder"
|
||||
path = _resolve_path(cwd, ref.target, allowed_root=allowed_root)
|
||||
_ensure_reference_path_allowed(path)
|
||||
if not path.exists():
|
||||
return f"{ref.raw}: file not found", None
|
||||
if not path.is_file():
|
||||
return f"{ref.raw}: path is not a file", None
|
||||
return f"{ref.raw}: {ref.kind} not found", None
|
||||
if not (path.is_dir() if is_folder else path.is_file()):
|
||||
return f"{ref.raw}: path is not a {ref.kind}", None
|
||||
if is_folder:
|
||||
listing = _build_folder_listing(path, cwd)
|
||||
return None, f"📁 {ref.raw} ({estimate_tokens_rough(listing)} tokens)\n{listing}"
|
||||
if _is_binary_file(path):
|
||||
# A binary file can't be inlined as text, but it IS on disk (the agent's
|
||||
# tools run where this resolves — the local cwd, or the staged copy in a
|
||||
# remote session workspace). Returning a bare "not supported" warning
|
||||
# with no content was a dead end: the model saw a failure and gave up
|
||||
# (told the user the file type wasn't supported). Instead, hand it an
|
||||
# actionable block — the path, type, size, and a nudge to use its tools —
|
||||
# so it can read/convert/view the file itself.
|
||||
# A bare "not supported" warning was a dead end (the model gave up); the file IS
|
||||
# on disk where the agent's tools run, so hand it an actionable block instead.
|
||||
return None, _binary_reference_block(ref, path)
|
||||
|
||||
text = path.read_text(encoding="utf-8")
|
||||
if ref.line_start is not None:
|
||||
lines = text.splitlines()
|
||||
start_idx = max(ref.line_start - 1, 0)
|
||||
end_idx = min(ref.line_end or ref.line_start, len(lines))
|
||||
text = "\n".join(lines[start_idx:end_idx])
|
||||
|
||||
lang = _code_fence_language(path)
|
||||
label = ref.raw
|
||||
return None, f"📄 {label} ({estimate_tokens_rough(text)} tokens)\n```{lang}\n{text}\n```"
|
||||
text = "\n".join(text.splitlines()[max(ref.line_start - 1, 0):ref.line_end or ref.line_start])
|
||||
lang = _FENCE_LANGUAGES.get(path.suffix.lower(), "")
|
||||
text_tokens = estimate_tokens_rough(text)
|
||||
# Check BEFORE building the fenced block: an oversized file is not going to be
|
||||
# inlined, so don't build a second MB-scale string just to discard it.
|
||||
if max_inline_tokens is not None and text_tokens > max_inline_tokens:
|
||||
# One oversized file used to poison the aggregate check and refuse the whole
|
||||
# turn (#61987); the file stays readable via the agent's tools instead. The
|
||||
# block alone carries the message (same shape as the binary path) — a warning
|
||||
# would duplicate it in "--- Context Warnings ---".
|
||||
return None, _oversized_text_reference_block(ref, path, text_tokens)
|
||||
return None, f"📄 {ref.raw} ({text_tokens} tokens)\n```{lang}\n{text}\n```"
|
||||
|
||||
|
||||
def _expand_folder_reference(
|
||||
ref: ContextReference,
|
||||
cwd: Path,
|
||||
*,
|
||||
allowed_root: Path | None = None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
path = _resolve_path(cwd, ref.target, allowed_root=allowed_root)
|
||||
_ensure_reference_path_allowed(path)
|
||||
if not path.exists():
|
||||
return f"{ref.raw}: folder not found", None
|
||||
if not path.is_dir():
|
||||
return f"{ref.raw}: path is not a folder", None
|
||||
|
||||
listing = _build_folder_listing(path, cwd)
|
||||
return None, f"📁 {ref.raw} ({estimate_tokens_rough(listing)} tokens)\n{listing}"
|
||||
def _run_quiet(cmd: list[str], cwd: Path, timeout: int, env: dict | None = None) -> subprocess.CompletedProcess:
|
||||
"""subprocess.run with captured text output, no stdin, and no console flash on Windows."""
|
||||
popen_kwargs: dict = {"creationflags": windows_hide_flags()} if IS_WINDOWS else {}
|
||||
return subprocess.run(cmd, cwd=cwd, capture_output=True, text=True, encoding='utf-8', errors='replace',
|
||||
timeout=timeout, stdin=subprocess.DEVNULL, **popen_kwargs, **({} if env is None else {"env": env}))
|
||||
|
||||
|
||||
def _expand_git_reference(
|
||||
ref: ContextReference,
|
||||
cwd: Path,
|
||||
args: list[str],
|
||||
label: str,
|
||||
) -> tuple[str | None, str | None]:
|
||||
_popen_kwargs = {"creationflags": windows_hide_flags()} if IS_WINDOWS else {}
|
||||
def _expand_git_reference(ref: ContextReference, cwd: Path, args: list[str], label: str) -> Expansion:
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["git", *args],
|
||||
cwd=cwd,
|
||||
capture_output=True,
|
||||
text=True, encoding='utf-8', errors='replace',
|
||||
timeout=30,
|
||||
stdin=subprocess.DEVNULL,
|
||||
**_popen_kwargs,
|
||||
)
|
||||
# Repo-supplied config/attributes must never execute code (GHSA-7x36-8jrh-v4pw).
|
||||
result = _run_quiet(["git", *harden_git_argv(args)], cwd, 30, env=noninteractive_git_env())
|
||||
except subprocess.TimeoutExpired:
|
||||
return f"{ref.raw}: git command timed out (30s)", None
|
||||
if result.returncode != 0:
|
||||
stderr = (result.stderr or "").strip() or "git command failed"
|
||||
return f"{ref.raw}: {stderr}", None
|
||||
content = result.stdout.strip()
|
||||
if not content:
|
||||
content = "(no output)"
|
||||
return f"{ref.raw}: {(result.stderr or '').strip() or 'git command failed'}", None
|
||||
content = result.stdout.strip() or "(no output)"
|
||||
return None, f"🧾 {label} ({estimate_tokens_rough(content)} tokens)\n```diff\n{content}\n```"
|
||||
|
||||
|
||||
async def _fetch_url_content(
|
||||
url: str,
|
||||
*,
|
||||
url_fetcher: Callable[[str], str | Awaitable[str]] | None = None,
|
||||
) -> str:
|
||||
fetcher = url_fetcher or _default_url_fetcher
|
||||
content = fetcher(url)
|
||||
async def _fetch_url_content(url: str, *, url_fetcher: UrlFetcher = None) -> str:
|
||||
content = (url_fetcher or _default_url_fetcher)(url)
|
||||
if inspect.isawaitable(content):
|
||||
content = await content
|
||||
return str(content or "").strip()
|
||||
@@ -458,158 +327,101 @@ async def _fetch_url_content(
|
||||
|
||||
async def _default_url_fetcher(url: str) -> str:
|
||||
from tools.web_tools import web_extract_tool
|
||||
docs = json.loads(await web_extract_tool([url], format="markdown")).get("results", [])
|
||||
return str(docs[0].get("content") or docs[0].get("raw_content") or "").strip() if docs else ""
|
||||
|
||||
raw = await web_extract_tool([url], format="markdown")
|
||||
payload = json.loads(raw)
|
||||
docs = payload.get("results", [])
|
||||
if not docs:
|
||||
return ""
|
||||
doc = docs[0]
|
||||
return str(doc.get("content") or doc.get("raw_content") or "").strip()
|
||||
|
||||
def _is_under(path: Path, root: Path) -> bool:
|
||||
try:
|
||||
path.relative_to(root)
|
||||
except ValueError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _resolve_path(cwd: Path, target: str, *, allowed_root: Path | None = None) -> Path:
|
||||
path = Path(os.path.expanduser(target))
|
||||
if not path.is_absolute():
|
||||
path = cwd / path
|
||||
resolved = path.resolve()
|
||||
if allowed_root is not None:
|
||||
try:
|
||||
resolved.relative_to(allowed_root)
|
||||
except ValueError as exc:
|
||||
raise ValueError("path is outside the allowed workspace") from exc
|
||||
resolved = (cwd / Path(os.path.expanduser(target))).resolve() # `/` keeps an absolute target as-is
|
||||
if allowed_root is not None and not _is_under(resolved, allowed_root):
|
||||
raise ValueError("path is outside the allowed workspace")
|
||||
return resolved
|
||||
|
||||
|
||||
def _ensure_reference_path_allowed(path: Path) -> None:
|
||||
"""Refuse credential/internal paths. Fails CLOSED: the gateway feeds untrusted remote text here."""
|
||||
from hermes_constants import get_hermes_home
|
||||
home = Path(os.path.expanduser("~")).resolve()
|
||||
hermes_home = get_hermes_home().resolve()
|
||||
|
||||
blocked_exact = {home / rel for rel in _SENSITIVE_HOME_FILES}
|
||||
blocked_exact.add(hermes_home / ".env")
|
||||
blocked_dirs = [home / rel for rel in _SENSITIVE_HOME_DIRS]
|
||||
blocked_dirs.extend(hermes_home / rel for rel in _SENSITIVE_HERMES_DIRS)
|
||||
|
||||
home, hermes_home = Path(os.path.expanduser("~")).resolve(), get_hermes_home().resolve()
|
||||
blocked_exact = {home / rel for rel in _SENSITIVE_HOME_FILES} | {hermes_home / ".env"}
|
||||
blocked_dirs = [home / rel for rel in _SENSITIVE_HOME_DIRS] + [hermes_home / rel for rel in _SENSITIVE_HERMES_DIRS]
|
||||
if path in blocked_exact:
|
||||
raise ValueError("path is a sensitive credential file and cannot be attached")
|
||||
|
||||
for blocked_dir in blocked_dirs:
|
||||
try:
|
||||
path.relative_to(blocked_dir)
|
||||
except ValueError:
|
||||
continue
|
||||
if any(_is_under(path, blocked_dir) for blocked_dir in blocked_dirs):
|
||||
raise ValueError("path is a sensitive credential or internal Hermes path and cannot be attached")
|
||||
|
||||
# Anchor to the canonical read deny-list (agent/file_safety.get_read_block_error),
|
||||
# the single source of truth used by the file/terminal read path. The narrow
|
||||
# list above predates that guard and never caught the real credential stores:
|
||||
# provider keys (auth.json), Anthropic OAuth tokens (.anthropic_oauth.json),
|
||||
# MCP OAuth material (mcp-tokens/), webhook HMAC secrets, and project-local
|
||||
# .env files. That gap matters because the gateway feeds UNTRUSTED remote
|
||||
# message text into reference expansion, so `@file:~/.hermes/auth.json` from a
|
||||
# chat peer would otherwise read the operator's keys straight into context.
|
||||
# Routing through the canonical guard closes the gap today and keeps this path
|
||||
# protected automatically whenever that deny-list grows.
|
||||
# Anchor to the canonical read deny-list (agent/file_safety.get_read_block_error): the
|
||||
# narrow list above never caught auth.json, .anthropic_oauth.json, mcp-tokens/, webhook
|
||||
# secrets or project .env files, and it grows automatically with that deny-list.
|
||||
try:
|
||||
from agent.file_safety import get_read_block_error
|
||||
|
||||
if get_read_block_error(str(path)) is not None:
|
||||
raise ValueError(
|
||||
"path is a sensitive credential or internal Hermes path and cannot be attached"
|
||||
)
|
||||
blocked = get_read_block_error(str(path)) is not None
|
||||
except ValueError:
|
||||
raise
|
||||
except Exception:
|
||||
# Fail CLOSED on the security path. This guard exists specifically to
|
||||
# cover credential stores the narrow list above misses (auth.json,
|
||||
# .anthropic_oauth.json, mcp-tokens/, ...). If the canonical lookup
|
||||
# ever fails, silently falling through would re-open that exact hole —
|
||||
# the gateway feeds untrusted remote text here, so a probe could then
|
||||
# attach the operator's keys. Refuse instead: a spurious block on a
|
||||
# legitimate file is a recoverable annoyance; a leaked credential is not.
|
||||
raise ValueError(
|
||||
"path could not be verified against the credential deny-list and cannot be attached"
|
||||
)
|
||||
# If the canonical lookup fails, falling through would re-open the exact hole this
|
||||
# guard closes; a spurious block is recoverable, a leaked credential is not.
|
||||
raise ValueError("path could not be verified against the credential deny-list and cannot be attached")
|
||||
if blocked:
|
||||
raise ValueError("path is a sensitive credential or internal Hermes path and cannot be attached")
|
||||
|
||||
|
||||
def _strip_trailing_punctuation(value: str) -> str:
|
||||
stripped = value.rstrip(TRAILING_PUNCTUATION)
|
||||
while stripped.endswith((")", "]", "}")):
|
||||
closer = stripped[-1]
|
||||
opener = {")": "(", "]": "[", "}": "{"}[closer]
|
||||
if stripped.count(closer) > stripped.count(opener):
|
||||
stripped = stripped[:-1]
|
||||
continue
|
||||
break
|
||||
# Drop unbalanced closers so "(see @file:x.py)" does not swallow the ")".
|
||||
while stripped.endswith((")", "]", "}")) and stripped.count(stripped[-1]) > stripped.count(_OPENERS[stripped[-1]]):
|
||||
stripped = stripped[:-1]
|
||||
return stripped
|
||||
|
||||
|
||||
def _strip_reference_wrappers(value: str) -> str:
|
||||
if len(value) >= 2 and value[0] == value[-1] and value[0] in "`\"'":
|
||||
return value[1:-1]
|
||||
return value
|
||||
return value[1:-1] if len(value) >= 2 and value[0] == value[-1] and value[0] in "`\"'" else value
|
||||
|
||||
|
||||
def _parse_file_reference_value(value: str) -> tuple[str, int | None, int | None]:
|
||||
quoted_match = re.match(
|
||||
r'^(?P<quote>`|"|\')(?P<path>.+?)(?P=quote)(?::(?P<start>\d+)(?:-(?P<end>\d+))?)?$',
|
||||
value,
|
||||
)
|
||||
if quoted_match:
|
||||
line_start = quoted_match.group("start")
|
||||
line_end = quoted_match.group("end")
|
||||
return (
|
||||
quoted_match.group("path"),
|
||||
int(line_start) if line_start is not None else None,
|
||||
int(line_end or line_start) if line_start is not None else None,
|
||||
)
|
||||
|
||||
range_match = re.match(r"^(?P<path>.+?):(?P<start>\d+)(?:-(?P<end>\d+))?$", value)
|
||||
if range_match:
|
||||
line_start = int(range_match.group("start"))
|
||||
return (
|
||||
range_match.group("path"),
|
||||
line_start,
|
||||
int(range_match.group("end") or range_match.group("start")),
|
||||
)
|
||||
|
||||
return _strip_reference_wrappers(value), None, None
|
||||
m = _FILE_VALUE_PATTERN.match(value)
|
||||
start = m and m.group("start")
|
||||
if not start: # no line range: the whole value is the (possibly quoted) path
|
||||
return _strip_reference_wrappers(value), None, None
|
||||
return m.group("qpath") or m.group("path"), int(start), int(m.group("end") or start)
|
||||
|
||||
|
||||
def _is_binary_file(path: Path) -> bool:
|
||||
mime, _ = mimetypes.guess_type(path.name)
|
||||
if mime and not mime.startswith("text/") and not any(
|
||||
path.name.endswith(ext) for ext in (".py", ".md", ".txt", ".json", ".yaml", ".yml", ".toml", ".js", ".ts")
|
||||
):
|
||||
return True
|
||||
chunk = path.read_bytes()[:4096]
|
||||
return b"\x00" in chunk
|
||||
mime = mimetypes.guess_type(path.name)[0]
|
||||
return bool(mime and not mime.startswith("text/") and not path.name.endswith(_TEXT_EXTENSIONS)) or (
|
||||
b"\x00" in path.read_bytes()[:4096]
|
||||
)
|
||||
|
||||
|
||||
def _build_folder_listing(path: Path, cwd: Path, limit: int = 200) -> str:
|
||||
lines = [f"{path.relative_to(cwd)}/"]
|
||||
entries = _iter_visible_entries(path, cwd, limit=limit)
|
||||
base_depth = len(path.relative_to(cwd).parts)
|
||||
for entry in entries:
|
||||
rel = entry.relative_to(cwd)
|
||||
indent = " " * max(len(rel.parts) - len(path.relative_to(cwd).parts) - 1, 0)
|
||||
if entry.is_dir():
|
||||
lines.append(f"{indent}- {entry.name}/")
|
||||
else:
|
||||
meta = _file_metadata(entry)
|
||||
lines.append(f"{indent}- {entry.name} ({meta})")
|
||||
indent = " " * max(len(entry.relative_to(cwd).parts) - base_depth - 1, 0)
|
||||
lines.append(f"{indent}- {entry.name}/" if entry.is_dir() else f"{indent}- {entry.name} ({_file_metadata(entry)})")
|
||||
if len(entries) >= limit:
|
||||
lines.append("- ...")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _iter_visible_entries(path: Path, cwd: Path, limit: int) -> list[Path]:
|
||||
rg_entries = _rg_files(path, cwd, limit=limit)
|
||||
if rg_entries is not None:
|
||||
"""Files under ``path`` via ``rg --files`` (honours ignore files), else an os.walk fallback."""
|
||||
try:
|
||||
rg = _run_quiet(["rg", "--files", str(path.relative_to(cwd))], cwd, 10)
|
||||
except (FileNotFoundError, OSError, subprocess.TimeoutExpired):
|
||||
rg = None
|
||||
if rg is not None and rg.returncode == 0:
|
||||
output: list[Path] = []
|
||||
seen_dirs: set[Path] = set()
|
||||
for rel in rg_entries:
|
||||
full = cwd / rel
|
||||
for line in [ln.strip() for ln in rg.stdout.splitlines() if ln.strip()][:limit]:
|
||||
full = cwd / Path(line)
|
||||
for parent in full.parents:
|
||||
if parent == cwd or parent in seen_dirs or path not in {parent, *parent.parents}:
|
||||
continue
|
||||
@@ -617,104 +429,69 @@ def _iter_visible_entries(path: Path, cwd: Path, limit: int) -> list[Path]:
|
||||
output.append(parent)
|
||||
output.append(full)
|
||||
return sorted({p for p in output if p.exists()}, key=lambda p: (not p.is_dir(), str(p)))
|
||||
|
||||
output = []
|
||||
for root, dirs, files in os.walk(path):
|
||||
dirs[:] = sorted(d for d in dirs if not d.startswith(".") and d != "__pycache__")
|
||||
files = sorted(f for f in files if not f.startswith("."))
|
||||
root_path = Path(root)
|
||||
for d in dirs:
|
||||
output.append(root_path / d)
|
||||
if len(output) >= limit:
|
||||
return output
|
||||
for f in files:
|
||||
output.append(root_path / f)
|
||||
for name in dirs + files:
|
||||
output.append(Path(root) / name)
|
||||
if len(output) >= limit:
|
||||
return output
|
||||
return output
|
||||
|
||||
|
||||
def _rg_files(path: Path, cwd: Path, limit: int) -> list[Path] | None:
|
||||
_popen_kwargs = {"creationflags": windows_hide_flags()} if IS_WINDOWS else {}
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["rg", "--files", str(path.relative_to(cwd))],
|
||||
cwd=cwd,
|
||||
capture_output=True,
|
||||
text=True, encoding='utf-8', errors='replace',
|
||||
timeout=10,
|
||||
stdin=subprocess.DEVNULL,
|
||||
**_popen_kwargs,
|
||||
)
|
||||
except (FileNotFoundError, OSError, subprocess.TimeoutExpired):
|
||||
return None
|
||||
if result.returncode != 0:
|
||||
return None
|
||||
files = [Path(line.strip()) for line in result.stdout.splitlines() if line.strip()]
|
||||
return files[:limit]
|
||||
|
||||
|
||||
def _agent_visible_path(path: Path) -> str:
|
||||
"""Map a host path to the path the agent's tools can read in the active backend.
|
||||
|
||||
Under a container backend (docker) the gateway host path dangles inside the
|
||||
sandbox — the container has its own filesystem and the host path is not
|
||||
mounted. Files staged into an auto-mounted cache dir (``images/``,
|
||||
``attachments/``, ...) are translated to their in-container path via the
|
||||
existing ``tools.credential_files`` machinery (#76577). Falls back to the
|
||||
host path when the backend is local or translation is unavailable.
|
||||
"""
|
||||
# Under a container backend the host path dangles inside the sandbox: translate staged
|
||||
# files to their auto-mounted cache path; fall back to the host path (local backend /
|
||||
# translation failure). Run the idempotent TERMINAL_ENV bridge first so in-process
|
||||
# gateways that never bridged terminal.* config still see the active backend.
|
||||
try:
|
||||
# Desktop/in-process gateways may not have bridged ``terminal.*``
|
||||
# config into ``TERMINAL_ENV`` at startup; run the idempotent bridge so
|
||||
# the credential_files translation gate sees the active backend.
|
||||
from tools.terminal_tool import _ensure_terminal_env_bridged
|
||||
|
||||
_ensure_terminal_env_bridged()
|
||||
from tools.credential_files import to_agent_visible_cache_path
|
||||
|
||||
return to_agent_visible_cache_path(str(path))
|
||||
except Exception:
|
||||
return str(path)
|
||||
|
||||
|
||||
def _binary_reference_block(ref: ContextReference, path: Path) -> str:
|
||||
mime, _ = mimetypes.guess_type(path.name)
|
||||
mime = mime or "application/octet-stream"
|
||||
def _on_disk_reference_block(ref: ContextReference, path: Path, descriptor: str, reason: str, guidance: str) -> str:
|
||||
"""Shared 📎 shape: the file was not inlined, but it IS on disk where the agent's
|
||||
tools run — hand the model the path and a nudge instead of a dead-end warning."""
|
||||
try:
|
||||
size = format_bytes(path.stat().st_size)
|
||||
except OSError:
|
||||
size = "unknown size"
|
||||
return (
|
||||
f"📎 {ref.raw} ({mime}, {size}) — binary file, not inlined as text. "
|
||||
f"It is available on disk at `{_agent_visible_path(path)}`. Use your tools to work with it "
|
||||
f"(read or convert it, extract its text, or view/render it as needed); "
|
||||
f"do not tell the user the file type is unsupported."
|
||||
f"📎 {ref.raw} ({descriptor}, {size}) — {reason} "
|
||||
f"It is available on disk at `{_agent_visible_path(path)}`. {guidance}"
|
||||
)
|
||||
|
||||
|
||||
def _binary_reference_block(ref: ContextReference, path: Path) -> str:
|
||||
mime = mimetypes.guess_type(path.name)[0] or "application/octet-stream"
|
||||
return _on_disk_reference_block(
|
||||
ref, path,
|
||||
descriptor=mime,
|
||||
reason="binary file, not inlined as text.",
|
||||
guidance="Use your tools to work with it (read or convert it, extract its text, "
|
||||
"or view/render it as needed); do not tell the user the file type is unsupported.",
|
||||
)
|
||||
|
||||
|
||||
def _oversized_text_reference_block(ref: ContextReference, path: Path, text_tokens: int) -> str:
|
||||
return _on_disk_reference_block(
|
||||
ref, path,
|
||||
descriptor=f"text file, approximately {text_tokens} tokens",
|
||||
reason="too large to inline safely.",
|
||||
guidance="Use read_file with a narrow line range, search_files, or terminal/code tools "
|
||||
"to inspect only the relevant parts; do not load the entire file into context.",
|
||||
)
|
||||
|
||||
|
||||
def _file_metadata(path: Path) -> str:
|
||||
if _is_binary_file(path):
|
||||
return f"{path.stat().st_size} bytes"
|
||||
try:
|
||||
line_count = path.read_text(encoding="utf-8").count("\n") + 1
|
||||
except Exception:
|
||||
return f"{path.stat().st_size} bytes"
|
||||
return f"{line_count} lines"
|
||||
|
||||
|
||||
def _code_fence_language(path: Path) -> str:
|
||||
mapping = {
|
||||
".py": "python",
|
||||
".js": "javascript",
|
||||
".ts": "typescript",
|
||||
".tsx": "tsx",
|
||||
".jsx": "jsx",
|
||||
".json": "json",
|
||||
".md": "markdown",
|
||||
".sh": "bash",
|
||||
".yml": "yaml",
|
||||
".yaml": "yaml",
|
||||
".toml": "toml",
|
||||
}
|
||||
return mapping.get(path.suffix.lower(), "")
|
||||
if not _is_binary_file(path):
|
||||
try:
|
||||
return f"{path.read_text(encoding='utf-8').count(chr(10)) + 1} lines"
|
||||
except Exception:
|
||||
pass
|
||||
return f"{path.stat().st_size} bytes"
|
||||
|
||||
+2876
-3928
File diff suppressed because it is too large
Load Diff
+1066
-8015
File diff suppressed because it is too large
Load Diff
+213
-468
@@ -1,14 +1,14 @@
|
||||
"""OpenAI-compatible shim that forwards Hermes requests to `copilot --acp`.
|
||||
|
||||
This adapter lets Hermes treat the GitHub Copilot ACP server as a chat-style
|
||||
backend. Each request starts a short-lived ACP session, sends the formatted
|
||||
conversation as a single prompt, collects text chunks, and converts the result
|
||||
back into the minimal shape Hermes expects from an OpenAI client.
|
||||
Each request starts a short-lived ACP session, sends the formatted conversation
|
||||
as one prompt, collects text chunks, and returns the minimal OpenAI-client shape.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import queue
|
||||
import re
|
||||
@@ -31,206 +31,147 @@ from agent.redact import redact_sensitive_text
|
||||
from tools.environments.local import hermes_subprocess_env
|
||||
|
||||
ACP_MARKER_BASE_URL = "acp://copilot"
|
||||
logger = logging.getLogger(__name__)
|
||||
_DEFAULT_TIMEOUT_SECONDS = 900.0
|
||||
|
||||
# Stderr fingerprint of the deprecated `gh copilot` CLI extension
|
||||
# (https://github.blog/changelog/2025-09-25-upcoming-deprecation-of-gh-copilot-cli-extension).
|
||||
# We require BOTH the literal product name ("gh-copilot") AND a deprecation
|
||||
# marker, so generic stderr from the NEW `@github/copilot` CLI — whose repo
|
||||
# is github.com/github/copilot-cli and which legitimately mentions "copilot-cli"
|
||||
# in its own banners and error messages — doesn't get misclassified as the
|
||||
# deprecated extension.
|
||||
# Stderr fingerprint of the deprecated `gh copilot` extension. Require BOTH the product name
|
||||
# AND a deprecation marker: the NEW `@github/copilot` CLI legitimately mentions "copilot-cli".
|
||||
_DEPRECATION_REQUIRED = ("gh-copilot",)
|
||||
_DEPRECATION_MARKERS = (
|
||||
"has been deprecated",
|
||||
"no commands will be executed",
|
||||
_DEPRECATION_MARKERS = ("has been deprecated", "no commands will be executed")
|
||||
_ROLE_LABELS = {"system": "System", "user": "User", "assistant": "Assistant", "tool": "Tool", "context": "Context"}
|
||||
# Probe verdicts per binary path (~50ms --help paid once per process). Only definitive
|
||||
# True/False is cached, so a CLI installed mid-session is picked up.
|
||||
_ACP_PROBE_CACHE: dict[str, bool] = {}
|
||||
_PROMPT_PREAMBLE = (
|
||||
"You are being used as the active ACP agent backend for Hermes.",
|
||||
"Use ACP capabilities to complete tasks.",
|
||||
"IMPORTANT: If you take an action with a tool, you MUST output tool calls using <tool_call>{...}</tool_call> blocks with JSON exactly in OpenAI function-call shape.",
|
||||
"If no tool is needed, answer normally.",
|
||||
)
|
||||
_INITIALIZE_PARAMS = {
|
||||
"protocolVersion": 1,
|
||||
"clientCapabilities": {"fs": {"readTextFile": True, "writeTextFile": True}},
|
||||
"clientInfo": {"name": "hermes-agent", "title": "Hermes Agent", "version": "0.0.0"},
|
||||
}
|
||||
_DEPRECATED_CLI_ERROR = (
|
||||
"Hermes ACP mode requires the NEW GitHub Copilot CLI (github.com/github/copilot-cli), but the binary it just "
|
||||
"spawned is the deprecated `gh copilot` extension.\n\n"
|
||||
"Install the new CLI:\n npm install -g @github/copilot\n # then verify with: copilot --help\n\n"
|
||||
"If `copilot` already resolves to the new CLI but you still see this,\npoint Hermes at it explicitly:\n"
|
||||
" export HERMES_COPILOT_ACP_COMMAND=/path/to/new/copilot\n\n"
|
||||
"Alternative: use the `copilot` provider (no ACP, hits the Copilot API\ndirectly with a Copilot subscription "
|
||||
"token) via `hermes setup`.\n\nOriginal error:\n"
|
||||
)
|
||||
|
||||
|
||||
def _is_gh_copilot_deprecation_message(stderr_text: str) -> bool:
|
||||
"""True iff stderr looks like the deprecated gh-copilot extension's banner."""
|
||||
|
||||
lower = stderr_text.lower()
|
||||
if not any(req in lower for req in _DEPRECATION_REQUIRED):
|
||||
return False
|
||||
return any(marker in lower for marker in _DEPRECATION_MARKERS)
|
||||
return any(req in lower for req in _DEPRECATION_REQUIRED) and any(m in lower for m in _DEPRECATION_MARKERS)
|
||||
|
||||
|
||||
def _resolve_command() -> str:
|
||||
return (
|
||||
os.getenv("HERMES_COPILOT_ACP_COMMAND", "").strip()
|
||||
or os.getenv("COPILOT_CLI_PATH", "").strip()
|
||||
or "copilot"
|
||||
)
|
||||
return os.getenv("HERMES_COPILOT_ACP_COMMAND", "").strip() or os.getenv("COPILOT_CLI_PATH", "").strip() or "copilot"
|
||||
|
||||
|
||||
def _resolve_args() -> list[str]:
|
||||
raw = os.getenv("HERMES_COPILOT_ACP_ARGS", "").strip()
|
||||
if not raw:
|
||||
return ["--acp", "--stdio"]
|
||||
return shlex.split(raw)
|
||||
|
||||
|
||||
# Probe verdicts cached per binary path so repeated prompts against a
|
||||
# CLI that supports --acp pay the ~50ms --help cost exactly once per
|
||||
# process. Only definitive verdicts (True/False) are cached; an
|
||||
# inconclusive probe (binary missing, --help crashed or timed out) is
|
||||
# not cached so a CLI installed mid-session is picked up.
|
||||
_ACP_PROBE_CACHE: dict[str, bool] = {}
|
||||
return shlex.split(os.getenv("HERMES_COPILOT_ACP_ARGS", "").strip()) or ["--acp", "--stdio"]
|
||||
|
||||
|
||||
def _acp_supported(command: str, args: list[str]) -> bool | None:
|
||||
"""Tri-state probe: does ``command`` accept the ACP args we'd pass?
|
||||
|
||||
Different CLI versions support different transports. The GitHub
|
||||
Copilot CLI (`@github/copilot`, late 2025+) ships with ``--acp``;
|
||||
older releases (and Claude Code v2.x as of Aug 2026) do not.
|
||||
Spawning a CLI that doesn't recognize the flag silently exits
|
||||
with code 1 and ``error: unknown option '--acp'`` on stderr,
|
||||
after which every delegate_task call hangs the parent for
|
||||
``child_timeout_seconds`` (default 600s) waiting for stdout
|
||||
that never arrives.
|
||||
|
||||
Returns:
|
||||
- ``True`` — help text advertises ``--acp``; safe to spawn.
|
||||
- ``False`` — help ran cleanly but ``--acp`` is absent; spawning
|
||||
would hang, so the caller should fast-fail with a clear error.
|
||||
- ``None`` — inconclusive (binary missing, --help failed or
|
||||
timed out). The caller must fall through to the normal spawn
|
||||
path, which surfaces the existing "Could not start Copilot ACP
|
||||
command" error with full context.
|
||||
|
||||
Only probes when ``--acp`` is actually among ``args``: a custom
|
||||
HERMES_COPILOT_ACP_ARGS transport is the operator's business.
|
||||
"""
|
||||
"""Tri-state ``--acp`` probe (a CLI without the flag exits 1 and the parent would wait the
|
||||
full child timeout for stdout that never arrives). True = help advertises --acp; False =
|
||||
help ran cleanly without it (caller fast-fails); None = inconclusive (binary missing /
|
||||
--help failed → normal spawn error). Skipped when ``--acp`` is not in ``args`` (custom transport)."""
|
||||
if "--acp" not in args:
|
||||
return True
|
||||
cached = _ACP_PROBE_CACHE.get(command)
|
||||
if cached is not None:
|
||||
if (cached := _ACP_PROBE_CACHE.get(command)) is not None:
|
||||
return cached
|
||||
try:
|
||||
probe = subprocess.run(
|
||||
[command, "--help"],
|
||||
capture_output=True, text=True, timeout=5,
|
||||
[command, "--help"], capture_output=True, text=True, encoding="utf-8", errors="replace", timeout=5,
|
||||
stdin=subprocess.DEVNULL,
|
||||
)
|
||||
except (FileNotFoundError, subprocess.TimeoutExpired, OSError):
|
||||
return None
|
||||
if probe.returncode != 0:
|
||||
# --help itself failed; can't tell anything about --acp.
|
||||
return None
|
||||
# Match ``--acp`` as a flag in the help text; tolerate spacing and
|
||||
# variants like ``[--acp]``.
|
||||
verdict = bool(re.search(r"(?:^|[\s\[])--acp(?:[\s=\],]|$)", probe.stdout, re.MULTILINE))
|
||||
_ACP_PROBE_CACHE[command] = verdict
|
||||
# ``--acp`` as a flag token; tolerate spacing and ``[--acp]`` variants.
|
||||
verdict = _ACP_PROBE_CACHE[command] = bool(re.search(r"(?:^|[\s\[])--acp(?:[\s=\],]|$)", probe.stdout, re.MULTILINE))
|
||||
return verdict
|
||||
|
||||
|
||||
def _resolve_home_dir() -> str:
|
||||
"""Return a stable HOME for child ACP processes."""
|
||||
home = os.environ.get("HOME", "").strip()
|
||||
if home:
|
||||
"""Stable HOME for child ACP processes; /tmp as a last resort so the child never starts HOME-less."""
|
||||
if home := os.environ.get("HOME", "").strip():
|
||||
return home
|
||||
|
||||
expanded = os.path.expanduser("~")
|
||||
if expanded and expanded != "~":
|
||||
if (expanded := os.path.expanduser("~")) and expanded != "~":
|
||||
return expanded
|
||||
|
||||
try:
|
||||
import pwd
|
||||
|
||||
resolved = pwd.getpwuid(os.getuid()).pw_dir.strip() # windows-footgun: ok — POSIX fallback inside try/except (pwd import fails on Windows)
|
||||
if resolved:
|
||||
return resolved
|
||||
return pwd.getpwuid(os.getuid()).pw_dir.strip() or "/tmp" # windows-footgun: ok — POSIX fallback inside try/except (pwd import fails on Windows)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Last resort: /tmp (writable on any POSIX system). Avoids crashing the
|
||||
# subprocess with no HOME; callers can set HERMES_HOME explicitly if they
|
||||
# need a different writable dir.
|
||||
return "/tmp"
|
||||
return "/tmp"
|
||||
|
||||
|
||||
def _build_subprocess_env() -> dict[str, str]:
|
||||
# Copilot ACP is a model-driving CLI executor: it legitimately needs LLM
|
||||
# provider credentials. Route through the central helper so Tier-1 secrets
|
||||
# (gateway bot tokens, GitHub auth, infra) are still stripped (#29157).
|
||||
env = hermes_subprocess_env(inherit_credentials=True)
|
||||
home = _resolve_home_dir()
|
||||
env["HOME"] = home
|
||||
from hermes_constants import apply_subprocess_home_env
|
||||
|
||||
# Copilot ACP drives a model and needs LLM provider credentials; the central helper still
|
||||
# strips Tier-1 secrets (bot tokens, GitHub auth, infra).
|
||||
# See #29157.
|
||||
env = hermes_subprocess_env(inherit_credentials=True)
|
||||
env["HOME"] = _resolve_home_dir()
|
||||
apply_subprocess_home_env(env)
|
||||
return env
|
||||
|
||||
|
||||
def _jsonrpc_result(message_id: Any, result: Any) -> dict[str, Any]:
|
||||
return {"jsonrpc": "2.0", "id": message_id, "result": result}
|
||||
|
||||
|
||||
def _jsonrpc_error(message_id: Any, code: int, message: str) -> dict[str, Any]:
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": message_id,
|
||||
"error": {
|
||||
"code": code,
|
||||
"message": message,
|
||||
},
|
||||
}
|
||||
return {"jsonrpc": "2.0", "id": message_id, "error": {"code": code, "message": message}}
|
||||
|
||||
|
||||
def _permission_denied(message_id: Any) -> dict[str, Any]:
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": message_id,
|
||||
"result": {
|
||||
"outcome": {
|
||||
"outcome": "cancelled",
|
||||
}
|
||||
},
|
||||
}
|
||||
def _enabled_ids(entries: Any, key: str) -> set[str]:
|
||||
"""Ids of ``entries`` (dicts) whose ``_meta.copilotEnablement`` is not ``disabled``."""
|
||||
return {str(e.get(key) or "").strip() for e in (entries or []) if isinstance(e, dict)
|
||||
and str((e.get("_meta") or {}).get("copilotEnablement") or "").strip().lower() != "disabled"}
|
||||
|
||||
|
||||
def _model_selection_request(session: dict[str, Any], requested_model: str) -> tuple[str, dict[str, str]] | None:
|
||||
"""ACP request selecting ``requested_model`` for ``session``: stable v1
|
||||
``session/set_config_option``, else Copilot's pre-stabilization ``session/set_model``
|
||||
when no model config option is advertised. A reported model list is authoritative:
|
||||
unknown and policy-disabled ids return None instead of being sent."""
|
||||
session_id = str(session.get("sessionId") or "").strip()
|
||||
requested_model = str(requested_model or "").strip()
|
||||
if not session_id or not requested_model or requested_model == "copilot-acp":
|
||||
return None
|
||||
options = [o for o in (session.get("configOptions") or []) if isinstance(o, dict) and "model" in (o.get("category"), o.get("id"))]
|
||||
if options:
|
||||
if requested_model not in _enabled_ids(options[0].get("options"), "value"):
|
||||
return None
|
||||
return "session/set_config_option", {"sessionId": session_id, "configId": str(options[0].get("id") or "model"), "value": requested_model}
|
||||
available = _enabled_ids((session.get("models") or {}).get("availableModels"), "modelId")
|
||||
return None if available and requested_model not in available else ("session/set_model", {"sessionId": session_id, "modelId": requested_model})
|
||||
|
||||
|
||||
def _format_messages_as_prompt(
|
||||
messages: list[dict[str, Any]],
|
||||
model: str | None = None,
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
tool_choice: Any = None,
|
||||
messages: list[dict[str, Any]], model: str | None = None, tools: list[dict[str, Any]] | None = None, tool_choice: Any = None,
|
||||
) -> str:
|
||||
sections: list[str] = [
|
||||
"You are being used as the active ACP agent backend for Hermes.",
|
||||
"Use ACP capabilities to complete tasks.",
|
||||
"IMPORTANT: If you take an action with a tool, you MUST output tool calls using <tool_call>{...}</tool_call> blocks with JSON exactly in OpenAI function-call shape.",
|
||||
"If no tool is needed, answer normally.",
|
||||
]
|
||||
if model:
|
||||
sections.append(f"Hermes requested model hint: {model}")
|
||||
|
||||
# Copilot has no tools of its own that would collide with Hermes', so it
|
||||
# forwards the whole toolset (no allowlist).
|
||||
sections.extend(_render_tool_bridge_sections(tools, tool_choice))
|
||||
|
||||
# Deliberately no "requested model" line: the model is applied for real via ACP session/set_model;
|
||||
# a prompt-text mention makes a substituted backend model FALSELY self-identify as the requested
|
||||
# one. Copilot has no tools of its own that collide with Hermes', so forward the whole toolset.
|
||||
sections: list[str] = [*_PROMPT_PREAMBLE, *_render_tool_bridge_sections(tools, tool_choice)]
|
||||
transcript: list[str] = []
|
||||
for message in messages:
|
||||
if not isinstance(message, dict):
|
||||
continue
|
||||
for message in (m for m in messages if isinstance(m, dict)):
|
||||
role = str(message.get("role") or "unknown").strip().lower()
|
||||
if role == "tool":
|
||||
role = "tool"
|
||||
elif role not in {"system", "user", "assistant"}:
|
||||
role = "context"
|
||||
|
||||
content = message.get("content")
|
||||
rendered = _render_message_content(content)
|
||||
if not rendered:
|
||||
continue
|
||||
|
||||
label = {
|
||||
"system": "System",
|
||||
"user": "User",
|
||||
"assistant": "Assistant",
|
||||
"tool": "Tool",
|
||||
"context": "Context",
|
||||
}.get(role, role.title())
|
||||
transcript.append(f"{label}:\n{rendered}")
|
||||
|
||||
if rendered := _render_message_content(message.get("content")):
|
||||
transcript.append(f"{_ROLE_LABELS.get(role, 'Context')}:\n{rendered}")
|
||||
if transcript:
|
||||
sections.append("Conversation transcript:\n\n" + "\n\n".join(transcript))
|
||||
|
||||
sections.append("Continue the conversation from the latest user request.")
|
||||
return "\n\n".join(section.strip() for section in sections if section and section.strip())
|
||||
|
||||
@@ -238,33 +179,21 @@ def _format_messages_as_prompt(
|
||||
def _render_message_content(content: Any) -> str:
|
||||
if content is None:
|
||||
return ""
|
||||
if isinstance(content, str):
|
||||
return content.strip()
|
||||
if isinstance(content, dict):
|
||||
if "text" in content:
|
||||
return str(content.get("text") or "").strip()
|
||||
if "content" in content and isinstance(content.get("content"), str):
|
||||
return str(content.get("content") or "").strip()
|
||||
return json.dumps(content, ensure_ascii=True)
|
||||
return content["content"].strip() if isinstance(content.get("content"), str) else json.dumps(content, ensure_ascii=True)
|
||||
if isinstance(content, list):
|
||||
parts: list[str] = []
|
||||
for item in content:
|
||||
if isinstance(item, str):
|
||||
parts.append(item)
|
||||
elif isinstance(item, dict):
|
||||
text = item.get("text")
|
||||
if isinstance(text, str) and text.strip():
|
||||
parts.append(text.strip())
|
||||
parts = [item if isinstance(item, str) else item["text"].strip() for item in content if isinstance(item, str)
|
||||
or (isinstance(item, dict) and isinstance(item.get("text"), str) and item["text"].strip())]
|
||||
return "\n".join(parts).strip()
|
||||
return str(content).strip()
|
||||
|
||||
|
||||
def _ensure_path_within_cwd(path_text: str, cwd: str) -> Path:
|
||||
candidate = Path(path_text)
|
||||
if not candidate.is_absolute():
|
||||
if not Path(path_text).is_absolute():
|
||||
raise PermissionError("ACP file-system paths must be absolute.")
|
||||
resolved = candidate.resolve()
|
||||
root = Path(cwd).resolve()
|
||||
resolved, root = Path(path_text).resolve(), Path(cwd).resolve()
|
||||
try:
|
||||
resolved.relative_to(root)
|
||||
except ValueError as exc:
|
||||
@@ -272,407 +201,223 @@ def _ensure_path_within_cwd(path_text: str, cwd: str) -> Path:
|
||||
return resolved
|
||||
|
||||
|
||||
class _ACPChatCompletions:
|
||||
def __init__(self, client: "CopilotACPClient"):
|
||||
self._client = client
|
||||
|
||||
def create(self, **kwargs: Any) -> Any:
|
||||
return self._client._create_chat_completion(**kwargs)
|
||||
def _effective_timeout(timeout: Any) -> float:
|
||||
"""Normalise a float or httpx.Timeout-like object to wall-clock seconds (largest component wins)."""
|
||||
if isinstance(timeout, (int, float)):
|
||||
return float(timeout)
|
||||
candidates = [getattr(timeout, attr, None) for attr in ("read", "write", "connect", "pool", "timeout")]
|
||||
return max((float(v) for v in candidates if isinstance(v, (int, float))), default=_DEFAULT_TIMEOUT_SECONDS)
|
||||
|
||||
|
||||
class _ACPChatNamespace:
|
||||
def __init__(self, client: "CopilotACPClient"):
|
||||
self.completions = _ACPChatCompletions(client)
|
||||
def _fs_read_text_file(params: dict[str, Any], cwd: str) -> Any:
|
||||
path = _ensure_path_within_cwd(str(params.get("path") or ""), cwd)
|
||||
if block_error := get_read_block_error(str(path)):
|
||||
raise PermissionError(block_error)
|
||||
try:
|
||||
content = path.read_text(encoding="utf-8")
|
||||
except FileNotFoundError:
|
||||
content = ""
|
||||
line, limit = params.get("line"), params.get("limit")
|
||||
if isinstance(line, int) and line > 1:
|
||||
end = line - 1 + limit if isinstance(limit, int) and limit > 0 else None
|
||||
content = "".join(content.splitlines(keepends=True)[line - 1:end])
|
||||
return {"content": redact_sensitive_text(content, force=True) if content else content}
|
||||
|
||||
|
||||
def _fs_write_text_file(params: dict[str, Any], cwd: str) -> Any:
|
||||
path = _ensure_path_within_cwd(str(params.get("path") or ""), cwd)
|
||||
if denied := get_write_denied_error(str(path)):
|
||||
raise PermissionError(denied)
|
||||
if is_write_approval_required(str(path)): # soft-gated for interactive tools; the ACP shim has no human channel → fail closed
|
||||
raise PermissionError(f"Write denied: '{path}' requires interactive approval and cannot be written through the ACP file bridge.")
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(str(params.get("content") or ""), encoding="utf-8")
|
||||
return None
|
||||
|
||||
|
||||
_FS_HANDLERS = {"fs/read_text_file": _fs_read_text_file, "fs/write_text_file": _fs_write_text_file}
|
||||
|
||||
|
||||
class CopilotACPClient:
|
||||
"""Minimal OpenAI-client-compatible facade for Copilot ACP."""
|
||||
|
||||
# Declared for agent/auxiliary_client.py: this shim drives an ACP subprocess over stdio, so it is
|
||||
# already a complete client (never re-dispatch through a wire adapter) and async-safe as-is.
|
||||
HERMES_SKIP_TRANSPORT_WRAP = True
|
||||
HERMES_SKIP_ASYNC_WRAP = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
default_headers: dict[str, str] | None = None,
|
||||
acp_command: str | None = None,
|
||||
acp_args: list[str] | None = None,
|
||||
acp_cwd: str | None = None,
|
||||
command: str | None = None,
|
||||
args: list[str] | None = None,
|
||||
**_: Any,
|
||||
self, *, api_key: str | None = None, base_url: str | None = None, default_headers: dict[str, str] | None = None,
|
||||
acp_command: str | None = None, acp_args: list[str] | None = None, acp_cwd: str | None = None, command: str | None = None,
|
||||
args: list[str] | None = None, **_: Any,
|
||||
):
|
||||
self.api_key = api_key or "copilot-acp"
|
||||
self.base_url = base_url or ACP_MARKER_BASE_URL
|
||||
self.api_key, self.base_url = api_key or "copilot-acp", base_url or ACP_MARKER_BASE_URL
|
||||
self._default_headers = dict(default_headers or {})
|
||||
self._acp_command = acp_command or command or _resolve_command()
|
||||
self._acp_args = list(acp_args or args or _resolve_args())
|
||||
self._acp_cwd = str(Path(acp_cwd or os.getcwd()).resolve())
|
||||
self.chat = _ACPChatNamespace(self)
|
||||
self.is_closed = False
|
||||
self._active_process: subprocess.Popen[str] | None = None
|
||||
self.chat = SimpleNamespace(completions=SimpleNamespace(create=self._create_chat_completion))
|
||||
self.is_closed, self._active_process = False, None
|
||||
self._active_process_lock = threading.Lock()
|
||||
|
||||
def close(self) -> None:
|
||||
proc: subprocess.Popen[str] | None
|
||||
with self._active_process_lock:
|
||||
proc = self._active_process
|
||||
self._active_process = None
|
||||
proc, self._active_process = self._active_process, None
|
||||
self.is_closed = True
|
||||
if proc is None:
|
||||
return
|
||||
try:
|
||||
proc.terminate()
|
||||
proc.wait(timeout=2)
|
||||
if proc is not None:
|
||||
proc.terminate()
|
||||
proc.wait(timeout=2)
|
||||
except Exception:
|
||||
try:
|
||||
with contextlib.suppress(Exception):
|
||||
proc.kill()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _create_chat_completion(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
messages: list[dict[str, Any]] | None = None,
|
||||
timeout: float | None = None,
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
tool_choice: Any = None,
|
||||
stream: bool = False,
|
||||
**_: Any,
|
||||
self, *, model: str | None = None, messages: list[dict[str, Any]] | None = None, timeout: float | None = None,
|
||||
tools: list[dict[str, Any]] | None = None, tool_choice: Any = None, stream: bool = False, **_: Any,
|
||||
) -> Any:
|
||||
prompt_text = _format_messages_as_prompt(
|
||||
messages or [],
|
||||
model=model,
|
||||
tools=tools,
|
||||
tool_choice=tool_choice,
|
||||
)
|
||||
# Normalise timeout: run_agent.py may pass an httpx.Timeout object
|
||||
# (used natively by the OpenAI SDK) rather than a plain float.
|
||||
if timeout is None:
|
||||
_effective_timeout = _DEFAULT_TIMEOUT_SECONDS
|
||||
elif isinstance(timeout, (int, float)):
|
||||
_effective_timeout = float(timeout)
|
||||
else:
|
||||
# httpx.Timeout or similar — pick the largest component so the
|
||||
# subprocess has enough wall-clock time for the full response.
|
||||
_candidates = [
|
||||
getattr(timeout, attr, None)
|
||||
for attr in ("read", "write", "connect", "pool", "timeout")
|
||||
]
|
||||
_numeric = [float(v) for v in _candidates if isinstance(v, (int, float))]
|
||||
_effective_timeout = max(_numeric) if _numeric else _DEFAULT_TIMEOUT_SECONDS
|
||||
|
||||
response_text, reasoning_text = self._run_prompt(
|
||||
prompt_text,
|
||||
timeout_seconds=_effective_timeout,
|
||||
)
|
||||
|
||||
prompt_text = _format_messages_as_prompt(messages or [], model=model, tools=tools, tool_choice=tool_choice)
|
||||
response_text, reasoning = self._run_prompt(prompt_text, timeout_seconds=_effective_timeout(timeout), model=model)
|
||||
tool_calls, cleaned_text = _extract_tool_calls_from_text(response_text)
|
||||
|
||||
usage = SimpleNamespace(
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
total_tokens=0,
|
||||
prompt_tokens_details=SimpleNamespace(cached_tokens=0),
|
||||
)
|
||||
assistant_message = SimpleNamespace(
|
||||
content=cleaned_text,
|
||||
tool_calls=tool_calls,
|
||||
reasoning=reasoning_text or None,
|
||||
reasoning_content=reasoning_text or None,
|
||||
message = SimpleNamespace(
|
||||
content=cleaned_text, tool_calls=tool_calls, reasoning=reasoning or None, reasoning_content=reasoning or None,
|
||||
reasoning_details=None,
|
||||
)
|
||||
finish_reason = "tool_calls" if tool_calls else "stop"
|
||||
choice = SimpleNamespace(message=assistant_message, finish_reason=finish_reason)
|
||||
completion = SimpleNamespace(
|
||||
choices=[choice],
|
||||
usage=usage,
|
||||
choices=[SimpleNamespace(message=message, finish_reason="tool_calls" if tool_calls else "stop")],
|
||||
usage=SimpleNamespace(prompt_tokens=0, completion_tokens=0, total_tokens=0, prompt_tokens_details=SimpleNamespace(cached_tokens=0)),
|
||||
model=model or "copilot-acp",
|
||||
)
|
||||
if stream:
|
||||
return _completion_to_stream_chunks(completion)
|
||||
return completion
|
||||
return _completion_to_stream_chunks(completion) if stream else completion
|
||||
|
||||
def _run_prompt(self, prompt_text: str, *, timeout_seconds: float) -> tuple[str, str]:
|
||||
# Fast-fail when the CLI doesn't support the ACP args we'd pass.
|
||||
# Without this guard, a CLI like Claude Code v2.x exits with
|
||||
# ``error: unknown option '--acp'`` immediately, then the parent
|
||||
# ACP loop waits the full ``child_timeout_seconds`` (default 600s)
|
||||
# for stdout that never arrives. The probe costs ~50ms and turns
|
||||
# a 600s silent hang into a 280ms clear error.
|
||||
# ``None`` (inconclusive probe — e.g. binary missing) falls
|
||||
# through to the spawn below, which raises the established
|
||||
# "Could not start Copilot ACP command" error.
|
||||
def _spawn(self) -> subprocess.Popen[str]:
|
||||
# Fast-fail when the CLI rejects --acp (else the parent waits the full child timeout for stdout that
|
||||
# never arrives). ``None`` falls through to the spawn's established start error.
|
||||
if _acp_supported(self._acp_command, self._acp_args) is False:
|
||||
preview = " ".join(self._acp_args[:3]) if self._acp_args else "(none)"
|
||||
raise RuntimeError(
|
||||
f"ACP transport not supported by '{self._acp_command}': "
|
||||
f"`{preview}` is rejected as an unknown option. "
|
||||
f"This usually means the CLI is an older release (e.g. "
|
||||
f"Claude Code v2.x) or a different tool than expected. "
|
||||
f"Either install a CLI that ships with --acp support "
|
||||
f"(e.g. `@github/copilot` late 2025+), or set "
|
||||
f"HERMES_COPILOT_ACP_COMMAND / HERMES_COPILOT_ACP_ARGS "
|
||||
f"to a working pair."
|
||||
f"ACP transport not supported by '{self._acp_command}': `{preview}` is rejected as an unknown option. This "
|
||||
"usually means the CLI is an older release (e.g. Claude Code v2.x) or a different tool than expected. Either "
|
||||
"install a CLI that ships with --acp support (e.g. `@github/copilot` late 2025+), or set "
|
||||
"HERMES_COPILOT_ACP_COMMAND / HERMES_COPILOT_ACP_ARGS to a working pair."
|
||||
)
|
||||
|
||||
try:
|
||||
# Hide the console the CLI child would otherwise flash on Windows
|
||||
# (#56747). Hide-only — stdio pipes stay intact for the ACP wire.
|
||||
from hermes_cli._subprocess_compat import windows_hide_flags
|
||||
from hermes_cli._subprocess_compat import windows_hide_flags # hide the Windows console flash (#56747); pipes intact for the ACP wire
|
||||
|
||||
# Hide the console the CLI child would otherwise flash on Windows (#56747). Hide-only — stdio
|
||||
# pipes stay intact for the ACP wire.
|
||||
proc = subprocess.Popen(
|
||||
[self._acp_command] + self._acp_args,
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True, encoding='utf-8', errors='replace',
|
||||
bufsize=1,
|
||||
cwd=self._acp_cwd,
|
||||
env=_build_subprocess_env(),
|
||||
[self._acp_command] + self._acp_args, stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.PIPE,
|
||||
text=True, encoding='utf-8', errors='replace', bufsize=1, cwd=self._acp_cwd, env=_build_subprocess_env(),
|
||||
creationflags=windows_hide_flags(),
|
||||
)
|
||||
except FileNotFoundError as exc:
|
||||
raise RuntimeError(
|
||||
f"Could not start Copilot ACP command '{self._acp_command}'. "
|
||||
"Install GitHub Copilot CLI or set HERMES_COPILOT_ACP_COMMAND/COPILOT_CLI_PATH."
|
||||
) from exc
|
||||
|
||||
raise RuntimeError(f"Could not start Copilot ACP command '{self._acp_command}'. Install GitHub Copilot CLI or set "
|
||||
"HERMES_COPILOT_ACP_COMMAND/COPILOT_CLI_PATH.") from exc
|
||||
if proc.stdin is None or proc.stdout is None:
|
||||
proc.kill()
|
||||
raise RuntimeError("Copilot ACP process did not expose stdin/stdout pipes.")
|
||||
|
||||
self.is_closed = False
|
||||
with self._active_process_lock:
|
||||
self._active_process = proc
|
||||
return proc
|
||||
|
||||
def _run_prompt(self, prompt_text: str, *, timeout_seconds: float, model: str | None = None) -> tuple[str, str]:
|
||||
# The CLI's `--model` spawn flag is deliberately NOT used: `copilot --acp` validates it (unknown id
|
||||
# aborts the spawn) but ignores it for the session; the model is applied after session/new instead.
|
||||
requested_model = str(model or "").strip()
|
||||
proc = self._spawn()
|
||||
inbox: queue.Queue[dict[str, Any]] = queue.Queue()
|
||||
stderr_tail: deque[str] = deque(maxlen=40)
|
||||
|
||||
def _stdout_reader() -> None:
|
||||
if proc.stdout is None:
|
||||
return
|
||||
for line in proc.stdout:
|
||||
try:
|
||||
inbox.put(json.loads(line))
|
||||
except Exception:
|
||||
inbox.put({"raw": line.rstrip("\n")})
|
||||
def _decode(line: str) -> dict[str, Any]:
|
||||
try:
|
||||
return json.loads(line)
|
||||
except Exception:
|
||||
return {"raw": line.rstrip("\n")}
|
||||
|
||||
def _stderr_reader() -> None:
|
||||
if proc.stderr is None:
|
||||
return
|
||||
for line in proc.stderr:
|
||||
stderr_tail.append(line.rstrip("\n"))
|
||||
def _pump(stream, sink) -> None:
|
||||
for line in stream or ():
|
||||
sink(line)
|
||||
|
||||
out_thread = threading.Thread(target=_stdout_reader, daemon=True)
|
||||
err_thread = threading.Thread(target=_stderr_reader, daemon=True)
|
||||
out_thread.start()
|
||||
err_thread.start()
|
||||
|
||||
next_id = 0
|
||||
threading.Thread(target=_pump, args=(proc.stdout, lambda line: inbox.put(_decode(line))), daemon=True).start()
|
||||
threading.Thread(target=_pump, args=(proc.stderr, lambda line: stderr_tail.append(line.rstrip("\n"))), daemon=True).start()
|
||||
request_ids = iter(range(1, 1 << 62))
|
||||
|
||||
def _request(method: str, params: dict[str, Any], *, text_parts: list[str] | None = None, reasoning_parts: list[str] | None = None) -> Any:
|
||||
nonlocal next_id
|
||||
next_id += 1
|
||||
request_id = next_id
|
||||
payload = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"method": method,
|
||||
"params": params,
|
||||
}
|
||||
proc.stdin.write(json.dumps(payload) + "\n")
|
||||
request_id = next(request_ids)
|
||||
proc.stdin.write(json.dumps({"jsonrpc": "2.0", "id": request_id, "method": method, "params": params}) + "\n")
|
||||
proc.stdin.flush()
|
||||
|
||||
deadline = time.monotonic() + timeout_seconds
|
||||
while time.monotonic() < deadline:
|
||||
if proc.poll() is not None:
|
||||
break
|
||||
while time.monotonic() < deadline and proc.poll() is None:
|
||||
try:
|
||||
msg = inbox.get(timeout=0.1)
|
||||
except queue.Empty:
|
||||
continue
|
||||
|
||||
if self._handle_server_message(
|
||||
msg,
|
||||
process=proc,
|
||||
cwd=self._acp_cwd,
|
||||
text_parts=text_parts,
|
||||
reasoning_parts=reasoning_parts,
|
||||
):
|
||||
continue
|
||||
|
||||
if msg.get("id") != request_id:
|
||||
msg, process=proc, cwd=self._acp_cwd, text_parts=text_parts, reasoning_parts=reasoning_parts
|
||||
) or msg.get("id") != request_id:
|
||||
continue
|
||||
if "error" in msg:
|
||||
err = msg.get("error") or {}
|
||||
raise RuntimeError(
|
||||
f"Copilot ACP {method} failed: {err.get('message') or err}"
|
||||
)
|
||||
raise RuntimeError(f"Copilot ACP {method} failed: {err.get('message') or err}")
|
||||
return msg.get("result")
|
||||
|
||||
stderr_text = "\n".join(stderr_tail).strip()
|
||||
if proc.poll() is not None and stderr_text:
|
||||
if _is_gh_copilot_deprecation_message(stderr_text):
|
||||
raise RuntimeError(
|
||||
"Hermes ACP mode requires the NEW GitHub Copilot CLI "
|
||||
"(github.com/github/copilot-cli), but the binary it just "
|
||||
"spawned is the deprecated `gh copilot` extension.\n\n"
|
||||
"Install the new CLI:\n"
|
||||
" npm install -g @github/copilot\n"
|
||||
" # then verify with: copilot --help\n\n"
|
||||
"If `copilot` already resolves to the new CLI but you still see this,\n"
|
||||
"point Hermes at it explicitly:\n"
|
||||
" export HERMES_COPILOT_ACP_COMMAND=/path/to/new/copilot\n\n"
|
||||
"Alternative: use the `copilot` provider (no ACP, hits the Copilot API\n"
|
||||
"directly with a Copilot subscription token) via `hermes setup`.\n\n"
|
||||
f"Original error:\n{stderr_text}"
|
||||
)
|
||||
raise RuntimeError(_DEPRECATED_CLI_ERROR + stderr_text)
|
||||
raise RuntimeError(f"Copilot ACP process exited early: {stderr_text}")
|
||||
raise TimeoutError(f"Timed out waiting for Copilot ACP response to {method}.")
|
||||
|
||||
try:
|
||||
_request(
|
||||
"initialize",
|
||||
{
|
||||
"protocolVersion": 1,
|
||||
"clientCapabilities": {
|
||||
"fs": {
|
||||
"readTextFile": True,
|
||||
"writeTextFile": True,
|
||||
}
|
||||
},
|
||||
"clientInfo": {
|
||||
"name": "hermes-agent",
|
||||
"title": "Hermes Agent",
|
||||
"version": "0.0.0",
|
||||
},
|
||||
},
|
||||
)
|
||||
session = _request(
|
||||
"session/new",
|
||||
{
|
||||
"cwd": self._acp_cwd,
|
||||
"mcpServers": [],
|
||||
},
|
||||
) or {}
|
||||
_request("initialize", _INITIALIZE_PARAMS)
|
||||
session = _request("session/new", {"cwd": self._acp_cwd, "mcpServers": []}) or {}
|
||||
session_id = str(session.get("sessionId") or "").strip()
|
||||
if not session_id:
|
||||
raise RuntimeError("Copilot ACP did not return a sessionId.")
|
||||
|
||||
if requested_model and requested_model != "copilot-acp":
|
||||
try:
|
||||
if (selection := _model_selection_request(session, requested_model)) is not None:
|
||||
_request(*selection)
|
||||
else:
|
||||
logger.warning("Copilot ACP does not offer model %r; using the session default.", requested_model)
|
||||
except Exception as exc:
|
||||
logger.warning("Copilot ACP model selection for %r failed; continuing with the session default: %s", requested_model, exc)
|
||||
text_parts: list[str] = []
|
||||
reasoning_parts: list[str] = []
|
||||
_request(
|
||||
"session/prompt",
|
||||
{
|
||||
"sessionId": session_id,
|
||||
"prompt": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": prompt_text,
|
||||
}
|
||||
],
|
||||
},
|
||||
text_parts=text_parts,
|
||||
reasoning_parts=reasoning_parts,
|
||||
)
|
||||
prompt = {"sessionId": session_id, "prompt": [{"type": "text", "text": prompt_text}]}
|
||||
_request("session/prompt", prompt, text_parts=text_parts, reasoning_parts=reasoning_parts)
|
||||
return "".join(text_parts), "".join(reasoning_parts)
|
||||
finally:
|
||||
self.close()
|
||||
|
||||
def _handle_server_message(
|
||||
self,
|
||||
msg: dict[str, Any],
|
||||
*,
|
||||
process: subprocess.Popen[str],
|
||||
cwd: str,
|
||||
text_parts: list[str] | None,
|
||||
reasoning_parts: list[str] | None,
|
||||
self, msg: dict[str, Any], *, process: subprocess.Popen[str], cwd: str, text_parts: list[str] | None, reasoning_parts: list[str] | None,
|
||||
) -> bool:
|
||||
"""Consume a server->client message; True when handled (notification or request answered)."""
|
||||
method = msg.get("method")
|
||||
if not isinstance(method, str):
|
||||
return False
|
||||
|
||||
if method == "session/update":
|
||||
params = msg.get("params") or {}
|
||||
update = params.get("update") or {}
|
||||
kind = str(update.get("sessionUpdate") or "").strip()
|
||||
update = (msg.get("params") or {}).get("update") or {}
|
||||
content = update.get("content") or {}
|
||||
chunk_text = ""
|
||||
if isinstance(content, dict):
|
||||
chunk_text = str(content.get("text") or "")
|
||||
if kind == "agent_message_chunk" and chunk_text and text_parts is not None:
|
||||
text_parts.append(chunk_text)
|
||||
elif kind == "agent_thought_chunk" and chunk_text and reasoning_parts is not None:
|
||||
reasoning_parts.append(chunk_text)
|
||||
chunk_text = str(content.get("text") or "") if isinstance(content, dict) else ""
|
||||
sinks = {"agent_message_chunk": text_parts, "agent_thought_chunk": reasoning_parts}
|
||||
if chunk_text and (sink := sinks.get(str(update.get("sessionUpdate") or "").strip())) is not None:
|
||||
sink.append(chunk_text)
|
||||
return True
|
||||
|
||||
if process.stdin is None:
|
||||
return True
|
||||
|
||||
message_id = msg.get("id")
|
||||
params = msg.get("params") or {}
|
||||
|
||||
if method == "session/request_permission":
|
||||
response = _permission_denied(message_id)
|
||||
elif method == "fs/read_text_file":
|
||||
response = _jsonrpc_result(message_id, {"outcome": {"outcome": "cancelled"}})
|
||||
elif method in _FS_HANDLERS:
|
||||
try:
|
||||
path = _ensure_path_within_cwd(str(params.get("path") or ""), cwd)
|
||||
block_error = get_read_block_error(str(path))
|
||||
if block_error:
|
||||
raise PermissionError(block_error)
|
||||
try:
|
||||
content = path.read_text(encoding="utf-8")
|
||||
except FileNotFoundError:
|
||||
content = ""
|
||||
line = params.get("line")
|
||||
limit = params.get("limit")
|
||||
if isinstance(line, int) and line > 1:
|
||||
lines = content.splitlines(keepends=True)
|
||||
start = line - 1
|
||||
end = start + limit if isinstance(limit, int) and limit > 0 else None
|
||||
content = "".join(lines[start:end])
|
||||
if content:
|
||||
content = redact_sensitive_text(content, force=True)
|
||||
response = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": message_id,
|
||||
"result": {
|
||||
"content": content,
|
||||
},
|
||||
}
|
||||
except Exception as exc:
|
||||
response = _jsonrpc_error(message_id, -32602, str(exc))
|
||||
elif method == "fs/write_text_file":
|
||||
try:
|
||||
path = _ensure_path_within_cwd(str(params.get("path") or ""), cwd)
|
||||
denied = get_write_denied_error(str(path))
|
||||
if denied:
|
||||
raise PermissionError(denied)
|
||||
# Approval-gated paths (e.g. ~/.ssh/config) are not hard-denied
|
||||
# for interactive tools, but the ACP shim has no human channel
|
||||
# to confirm the write — fail closed here.
|
||||
if is_write_approval_required(str(path)):
|
||||
raise PermissionError(
|
||||
f"Write denied: '{path}' requires interactive approval "
|
||||
"and cannot be written through the ACP file bridge."
|
||||
)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(str(params.get("content") or ""), encoding="utf-8")
|
||||
response = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": message_id,
|
||||
"result": None,
|
||||
}
|
||||
response = _jsonrpc_result(message_id, _FS_HANDLERS[method](msg.get("params") or {}, cwd))
|
||||
except Exception as exc:
|
||||
response = _jsonrpc_error(message_id, -32602, str(exc))
|
||||
else:
|
||||
response = _jsonrpc_error(
|
||||
message_id,
|
||||
-32601,
|
||||
f"ACP client method '{method}' is not supported by Hermes yet.",
|
||||
)
|
||||
|
||||
response = _jsonrpc_error(message_id, -32601, f"ACP client method '{method}' is not supported by Hermes yet.")
|
||||
process.stdin.write(json.dumps(response) + "\n")
|
||||
process.stdin.flush()
|
||||
return True
|
||||
|
||||
@@ -1,11 +1,7 @@
|
||||
"""Credential-pool disk-boundary sanitization helpers.
|
||||
|
||||
These helpers define which credential-pool entries are references to borrowed
|
||||
runtime secrets and strip raw values before those entries are written to
|
||||
``auth.json``. They intentionally have no dependency on ``hermes_cli.auth`` so
|
||||
both the pool model and the final auth-store write boundary can share the same
|
||||
policy without import cycles.
|
||||
"""
|
||||
"""Credential-pool disk-boundary sanitization: strip raw secrets from *borrowed*
|
||||
pool entries before they reach ``auth.json``. Deliberately free of
|
||||
``hermes_cli.auth`` imports so the pool model and the auth-store write boundary
|
||||
share one policy without import cycles."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -14,9 +10,9 @@ import re
|
||||
from typing import Any, Dict, Mapping
|
||||
|
||||
|
||||
# Sources Hermes owns and can intentionally persist in auth.json. Everything
|
||||
# else with a non-empty source is treated as borrowed/reference-only by default
|
||||
# so future external secret providers fail closed at the disk boundary.
|
||||
# Sources Hermes owns and may persist with secrets. Any other non-empty,
|
||||
# non-manual source is borrowed/reference-only so new external providers fail
|
||||
# closed at the disk boundary.
|
||||
_PERSISTABLE_PROVIDER_SOURCES = frozenset({
|
||||
("anthropic", "hermes_pkce"),
|
||||
("minimax-oauth", "oauth"),
|
||||
@@ -25,87 +21,42 @@ _PERSISTABLE_PROVIDER_SOURCES = frozenset({
|
||||
("xai-oauth", "device_code"),
|
||||
})
|
||||
|
||||
# Metadata keys that look secret-ish by suffix but are safe to persist.
|
||||
_SAFE_SECRETISH_METADATA_KEYS = frozenset({
|
||||
"secret_fingerprint",
|
||||
"secret_source",
|
||||
"token_type",
|
||||
"scope",
|
||||
"client_id",
|
||||
"agent_key_id",
|
||||
"agent_key_expires_at",
|
||||
"agent_key_expires_in",
|
||||
"agent_key_reused",
|
||||
"agent_key_obtained_at",
|
||||
"expires_at",
|
||||
"expires_at_ms",
|
||||
"expires_in",
|
||||
"last_refresh",
|
||||
"last_status",
|
||||
"last_status_at",
|
||||
"last_error_code",
|
||||
"last_error_reason",
|
||||
"last_error_message",
|
||||
"secret_fingerprint", "secret_source", "token_type", "scope", "client_id",
|
||||
"agent_key_id", "agent_key_expires_at", "agent_key_expires_in",
|
||||
"agent_key_reused", "agent_key_obtained_at", "expires_at", "expires_at_ms",
|
||||
"expires_in", "last_refresh", "last_status", "last_status_at",
|
||||
"last_error_code", "last_error_reason", "last_error_message",
|
||||
"last_error_reset_at",
|
||||
})
|
||||
|
||||
_SECRET_VALUE_KEYS = frozenset({
|
||||
"access_token",
|
||||
"refresh_token",
|
||||
"agent_key",
|
||||
"api_key",
|
||||
"apikey",
|
||||
"api_token",
|
||||
"auth_token",
|
||||
"authorization",
|
||||
"bearer_token",
|
||||
"client_secret",
|
||||
"credential",
|
||||
"credentials",
|
||||
"id_token",
|
||||
"oauth_token",
|
||||
"private_key",
|
||||
"secret_key",
|
||||
"session_token",
|
||||
"password",
|
||||
"secret",
|
||||
"token",
|
||||
"tokens",
|
||||
"access_token", "refresh_token", "agent_key", "api_key", "apikey",
|
||||
"api_token", "auth_token", "authorization", "bearer_token", "client_secret",
|
||||
"credential", "credentials", "id_token", "oauth_token", "private_key",
|
||||
"secret_key", "session_token", "password", "secret", "token", "tokens",
|
||||
})
|
||||
|
||||
_SECRET_VALUE_SUFFIXES = (
|
||||
"_api_key",
|
||||
"_api_token",
|
||||
"_access_token",
|
||||
"_auth_token",
|
||||
"_refresh_token",
|
||||
"_bearer_token",
|
||||
"_client_secret",
|
||||
"_id_token",
|
||||
"_oauth_token",
|
||||
"_private_key",
|
||||
"_session_token",
|
||||
"_secret_key",
|
||||
"_password",
|
||||
"_secret",
|
||||
"_token",
|
||||
"_key",
|
||||
"_api_key", "_api_token", "_access_token", "_auth_token", "_refresh_token",
|
||||
"_bearer_token", "_client_secret", "_id_token", "_oauth_token",
|
||||
"_private_key", "_session_token", "_secret_key", "_password", "_secret",
|
||||
"_token", "_key",
|
||||
)
|
||||
|
||||
_CAMEL_CASE_BOUNDARY = re.compile(r"(?<=[a-z0-9])(?=[A-Z])")
|
||||
|
||||
|
||||
def _normalize_key(key: Any) -> str:
|
||||
raw = str(key or "").strip()
|
||||
raw = _CAMEL_CASE_BOUNDARY.sub("_", raw)
|
||||
raw = _CAMEL_CASE_BOUNDARY.sub("_", str(key or "").strip())
|
||||
return raw.lower().replace("-", "_").replace(".", "_")
|
||||
|
||||
|
||||
def is_borrowed_credential_source(source: Any, provider_id: Any = None) -> bool:
|
||||
"""Return True when ``source`` points at a borrowed/reference-only secret."""
|
||||
normalized_source = str(source or "").strip().lower()
|
||||
if not normalized_source:
|
||||
return False
|
||||
if normalized_source == "manual" or normalized_source.startswith("manual:"):
|
||||
if not normalized_source or normalized_source == "manual" or normalized_source.startswith("manual:"):
|
||||
return False
|
||||
normalized_provider = str(provider_id or "").strip().lower()
|
||||
return (normalized_provider, normalized_source) not in _PERSISTABLE_PROVIDER_SOURCES
|
||||
@@ -115,15 +66,16 @@ def _is_secret_payload_key(key: Any) -> bool:
|
||||
normalized = _normalize_key(key)
|
||||
if not normalized or normalized in _SAFE_SECRETISH_METADATA_KEYS:
|
||||
return False
|
||||
if normalized in _SECRET_VALUE_KEYS:
|
||||
return True
|
||||
return normalized.endswith(_SECRET_VALUE_SUFFIXES)
|
||||
return normalized in _SECRET_VALUE_KEYS or normalized.endswith(_SECRET_VALUE_SUFFIXES)
|
||||
|
||||
|
||||
def _fingerprint_value(value: Any) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
text = str(value)
|
||||
def fingerprint_secret_value(value: Any) -> str | None:
|
||||
"""Non-reversible ``sha256:<16 hex>`` fingerprint of one secret value.
|
||||
|
||||
Callers comparing a live secret against the ``secret_fingerprint`` left on
|
||||
a sanitized (borrowed) pool row need exactly the digest this module writes.
|
||||
"""
|
||||
text = "" if value is None else str(value)
|
||||
if not text:
|
||||
return None
|
||||
digest = hashlib.sha256(text.encode("utf-8", errors="surrogatepass")).hexdigest()
|
||||
@@ -131,17 +83,13 @@ def _fingerprint_value(value: Any) -> str | None:
|
||||
|
||||
|
||||
def _credential_secret_fingerprint(payload: Mapping[str, Any]) -> str | None:
|
||||
for key in ("agent_key", "access_token", "refresh_token", "api_key", "token", "secret"):
|
||||
fingerprint = _fingerprint_value(payload.get(key))
|
||||
preferred = ("agent_key", "access_token", "refresh_token", "api_key", "token", "secret")
|
||||
candidates = [payload.get(k) for k in preferred]
|
||||
candidates += [v for k, v in payload.items() if _is_secret_payload_key(k)]
|
||||
for value in candidates:
|
||||
fingerprint = fingerprint_secret_value(value)
|
||||
if fingerprint:
|
||||
return fingerprint
|
||||
|
||||
for key, value in payload.items():
|
||||
if _is_secret_payload_key(key):
|
||||
fingerprint = _fingerprint_value(value)
|
||||
if fingerprint:
|
||||
return fingerprint
|
||||
|
||||
existing = payload.get("secret_fingerprint")
|
||||
if isinstance(existing, str) and existing.startswith("sha256:"):
|
||||
return existing
|
||||
@@ -154,21 +102,15 @@ def sanitize_borrowed_credential_payload(
|
||||
) -> Dict[str, Any]:
|
||||
"""Return a disk-safe credential-pool payload.
|
||||
|
||||
Owned sources (manual entries and Hermes-owned OAuth/device-code state)
|
||||
pass through unchanged. Borrowed/reference-only sources keep labels,
|
||||
source refs, status/cooldown metadata, counters, and a non-reversible
|
||||
fingerprint, but raw secret value fields are removed.
|
||||
Owned sources pass through unchanged. Borrowed sources keep labels,
|
||||
source refs, status/cooldown metadata, counters and a fingerprint, but
|
||||
every raw secret value field is removed.
|
||||
"""
|
||||
result = dict(payload)
|
||||
if not is_borrowed_credential_source(result.get("source"), provider_id):
|
||||
return result
|
||||
|
||||
fingerprint = _credential_secret_fingerprint(result)
|
||||
sanitized = {
|
||||
key: value
|
||||
for key, value in result.items()
|
||||
if not _is_secret_payload_key(key)
|
||||
}
|
||||
sanitized = {k: v for k, v in result.items() if not _is_secret_payload_key(k)}
|
||||
if fingerprint:
|
||||
sanitized["secret_fingerprint"] = fingerprint
|
||||
return sanitized
|
||||
|
||||
+1744
-2218
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,131 @@
|
||||
"""Locked credential-pool administration and target resolution."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import replace
|
||||
from typing import Any, Optional, Tuple, TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agent.credential_pool import PooledCredential
|
||||
|
||||
|
||||
def _cleared_status_copy(entry: PooledCredential) -> PooledCredential:
|
||||
from agent.credential_pool import _CLEAR_STATUS
|
||||
|
||||
return replace(entry, **_CLEAR_STATUS,
|
||||
extra={k: v for k, v in entry.extra.items() if k != "failure_reason"})
|
||||
|
||||
|
||||
class CredentialPoolAdminMixin:
|
||||
def reset_status(self, credential_id: str) -> Optional[PooledCredential]:
|
||||
"""Clear only the target's local error state, preserving sibling cooldowns."""
|
||||
with self._lock:
|
||||
entry = self._find(lambda e: e.id == credential_id)
|
||||
if entry is None:
|
||||
return None
|
||||
cleared = _cleared_status_copy(entry)
|
||||
self._replace_entry(entry, cleared)
|
||||
self._persist(status_cleared_ids=[cleared.id])
|
||||
return cleared
|
||||
def reset_statuses(self) -> int:
|
||||
"""Clear exhaustion state on every entry. Returns how many were cleared.
|
||||
|
||||
``failure_reason`` lives in ``extra``, not a dataclass field, so it is
|
||||
stripped explicitly. The persist declares the cleared ids because the
|
||||
disk-recency merge reads a cleared ``last_status_at`` (None -> epoch 0)
|
||||
as a stale snapshot and would copy a still-binding cooldown back.
|
||||
"""
|
||||
from agent.credential_pool import _CLEAR_STATUS
|
||||
|
||||
with self._lock:
|
||||
stale = [
|
||||
e for e in self._entries
|
||||
if e.last_status or e.last_status_at or e.last_error_code or e.failure_reason
|
||||
]
|
||||
if stale:
|
||||
stale_ids = {e.id for e in stale}
|
||||
self._entries = [
|
||||
_cleared_status_copy(e) if e.id in stale_ids else e
|
||||
for e in self._entries
|
||||
]
|
||||
self._persist(status_cleared_ids=list(stale_ids))
|
||||
return len(stale)
|
||||
|
||||
def remove_index(self, index: int) -> Optional[PooledCredential]:
|
||||
from agent.credential_pool import persist_pool_entries
|
||||
|
||||
with self._lock:
|
||||
if index < 1 or index > len(self._entries):
|
||||
return None
|
||||
removed = self._entries.pop(index - 1)
|
||||
self._entries = [replace(e, priority=p) for p, e in enumerate(self._entries)]
|
||||
persist_pool_entries(
|
||||
self.provider,
|
||||
[entry.to_dict() for entry in self._entries],
|
||||
removed_ids=[removed.id],
|
||||
)
|
||||
if self._current_id == removed.id:
|
||||
self._current_id = None
|
||||
return removed
|
||||
|
||||
def move_entry(self, credential_id: str, priority: int) -> Optional[PooledCredential]:
|
||||
"""Place an entry at a clamped zero-based position and persist contiguous priorities."""
|
||||
from agent.credential_pool import _normalize_pool_priorities
|
||||
|
||||
with self._lock:
|
||||
entry = self._find(lambda e: e.id == credential_id)
|
||||
if entry is None:
|
||||
return None
|
||||
others = [e for e in self._entries if e.id != credential_id]
|
||||
others.insert(max(0, min(int(priority), len(others))), entry)
|
||||
entries = [replace(e, priority=p) for p, e in enumerate(others)]
|
||||
# Apply load-time ordering now so the reported position survives reload.
|
||||
_normalize_pool_priorities(self.provider, entries)
|
||||
self._entries = sorted(entries, key=lambda e: e.priority)
|
||||
self._persist()
|
||||
return self._find(lambda e: e.id == credential_id)
|
||||
|
||||
def resolve_target(self, target: Any) -> Tuple[Optional[int], Optional[PooledCredential], Optional[str]]:
|
||||
raw = str(target or "").strip()
|
||||
if not raw:
|
||||
return None, None, "No credential target provided."
|
||||
|
||||
with self._lock:
|
||||
for idx, entry in enumerate(self._entries, start=1):
|
||||
if entry.id == raw:
|
||||
return idx, entry, None
|
||||
|
||||
label_matches = [
|
||||
(idx, entry)
|
||||
for idx, entry in enumerate(self._entries, start=1)
|
||||
if entry.label.strip().lower() == raw.lower()
|
||||
]
|
||||
if len(label_matches) == 1:
|
||||
return label_matches[0][0], label_matches[0][1], None
|
||||
if len(label_matches) > 1:
|
||||
return None, None, f'Ambiguous credential label "{raw}". Use the numeric index or entry id instead.'
|
||||
if raw.isdigit():
|
||||
index = int(raw)
|
||||
if 1 <= index <= len(self._entries):
|
||||
return index, self._entries[index - 1], None
|
||||
return None, None, f"No credential #{index}."
|
||||
return None, None, f'No credential matching "{raw}".'
|
||||
|
||||
def add_entry(self, entry: PooledCredential) -> PooledCredential:
|
||||
from agent.credential_pool import _next_priority, write_credential_pool
|
||||
|
||||
with self._lock:
|
||||
entry = replace(entry, priority=_next_priority(self._entries))
|
||||
self._entries.append(entry)
|
||||
borrowed_ids = getattr(self, "_borrowed_root_ids", None)
|
||||
if borrowed_ids:
|
||||
# ``hermes -p <profile> auth add <single-use provider>``: the
|
||||
# profile claims its OWN credential. Persist only profile-owned
|
||||
# rows — copying the borrowed root grant alongside would fork
|
||||
# its single-use refresh token (#100339). Once the profile owns
|
||||
# rows, the root fallback for this provider is shadowed.
|
||||
self._entries = [e for e in self._entries if e.id not in borrowed_ids]
|
||||
write_credential_pool(self.provider, [e.to_dict() for e in self._entries])
|
||||
self._borrowed_root_ids = set()
|
||||
else:
|
||||
self._persist()
|
||||
return entry
|
||||
+100
-279
@@ -1,46 +1,12 @@
|
||||
"""Unified removal contract for every credential source Hermes reads from.
|
||||
|
||||
Hermes seeds its credential pool from many places:
|
||||
|
||||
env:<VAR> — os.environ / ~/.hermes/.env
|
||||
claude_code — ~/.claude/.credentials.json
|
||||
hermes_pkce — ~/.hermes/.anthropic_oauth.json
|
||||
device_code — auth.json providers.<provider> (nous, openai-codex, ...)
|
||||
qwen-cli — ~/.qwen/oauth_creds.json
|
||||
gh_cli — gh auth token
|
||||
config:<name> — custom_providers config entry
|
||||
model_config — model.api_key when model.provider == "custom"
|
||||
manual — user ran `hermes auth add`
|
||||
|
||||
Each source has its own reader inside ``agent.credential_pool._seed_from_*``
|
||||
(which keep their existing shape — we haven't restructured them). What we
|
||||
unify here is **removal**:
|
||||
|
||||
``hermes auth remove <provider> <N>`` must make the pool entry stay gone.
|
||||
|
||||
Before this module, every source had an ad-hoc removal branch in
|
||||
``auth_remove_command``, and several sources had no branch at all — so
|
||||
``auth remove`` silently reverted on the next ``load_pool()`` call for
|
||||
qwen-cli, nous device_code (partial), hermes_pkce, copilot gh_cli, and
|
||||
custom-config sources.
|
||||
|
||||
Now every source registers a ``RemovalStep`` that does exactly three things
|
||||
in the same shape:
|
||||
|
||||
1. Clean up whatever externally-readable state the source reads from
|
||||
(.env line, auth.json block, OAuth file, etc.)
|
||||
2. Suppress the ``(provider, source_id)`` in auth.json so the
|
||||
corresponding ``_seed_from_*`` branch skips the upsert on re-load
|
||||
3. Return ``RemovalResult`` describing what was cleaned and any
|
||||
diagnostic hints the user should see (shell-exported env vars,
|
||||
external credential files we deliberately don't delete, etc.)
|
||||
|
||||
Adding a new credential source is:
|
||||
- wire up a reader branch in ``_seed_from_*`` (existing pattern)
|
||||
- gate that reader behind ``is_source_suppressed(provider, source_id)``
|
||||
- register a ``RemovalStep`` here
|
||||
|
||||
No more per-source if/elif chain in ``auth_remove_command``.
|
||||
Readers live in ``agent.credential_pool``; what is unified here is **removal**:
|
||||
``hermes auth remove <provider> <N>`` must make the entry stay gone across
|
||||
``load_pool()`` calls. Each source registers a ``RemovalStep`` whose
|
||||
``remove_fn`` cleans the external state the source reads from, and the
|
||||
dispatcher suppresses ``(provider, source_id)`` in auth.json so the seeding
|
||||
branch skips the upsert. Adding a source: wire a reader branch in
|
||||
``_seed_from_*``, gate it behind ``is_source_suppressed``, register a step here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -54,20 +20,10 @@ from typing import Callable, List, Optional
|
||||
class RemovalResult:
|
||||
"""Outcome of removing a credential source.
|
||||
|
||||
Attributes:
|
||||
cleaned: Short strings describing external state that was actually
|
||||
mutated (``"Cleared XAI_API_KEY from .env"``,
|
||||
``"Cleared openai-codex OAuth tokens from auth store"``).
|
||||
Printed as plain lines to the user.
|
||||
hints: Diagnostic lines ABOUT state the user may need to clean up
|
||||
themselves or is deliberately left intact (shell-exported env
|
||||
var, Claude Code credential file we don't delete, etc.).
|
||||
Printed as plain lines to the user. Always non-destructive.
|
||||
suppress: Whether to call ``suppress_credential_source`` after
|
||||
cleanup so future ``load_pool`` calls skip this source.
|
||||
Default True — almost every source needs this to stay sticky.
|
||||
The only legitimate False is ``manual`` entries, which aren't
|
||||
seeded from anywhere external.
|
||||
``cleaned``: external state actually mutated (printed to the user).
|
||||
``hints``: diagnostics about state left intact or that the user must clean
|
||||
up. ``suppress``: call ``suppress_credential_source`` afterwards so
|
||||
``load_pool`` skips the source; only ``manual`` entries legitimately use False.
|
||||
"""
|
||||
|
||||
cleaned: List[str] = field(default_factory=list)
|
||||
@@ -77,22 +33,11 @@ class RemovalResult:
|
||||
|
||||
@dataclass
|
||||
class RemovalStep:
|
||||
"""How to remove one specific credential source cleanly.
|
||||
"""How to remove one credential source.
|
||||
|
||||
Attributes:
|
||||
provider: Provider pool key (``"xai"``, ``"anthropic"``, ``"nous"``, ...).
|
||||
Special value ``"*"`` means "matches any provider" — used for
|
||||
sources like ``manual`` that aren't provider-specific.
|
||||
source_id: Source identifier as it appears in
|
||||
``PooledCredential.source``. May be a literal (``"claude_code"``)
|
||||
or a prefix pattern matched via ``match_fn``.
|
||||
match_fn: Optional predicate overriding literal ``source_id``
|
||||
matching. Gets the removed entry's source string. Used for
|
||||
``env:*`` (any env-seeded key), ``config:*`` (any custom
|
||||
pool), and ``manual:*`` (any manual-source variant).
|
||||
remove_fn: ``(provider, removed_entry) -> RemovalResult``. Does the
|
||||
actual cleanup and returns what happened for the user.
|
||||
description: One-line human-readable description for docs / tests.
|
||||
``provider`` ``"*"`` matches any provider. ``match_fn`` overrides literal
|
||||
``source_id`` matching (prefix patterns like ``env:*`` / ``config:*``).
|
||||
``remove_fn(provider, removed_entry) -> RemovalResult``.
|
||||
"""
|
||||
|
||||
provider: str
|
||||
@@ -109,46 +54,13 @@ class RemovalStep:
|
||||
return source == self.source_id
|
||||
|
||||
|
||||
_REGISTRY: List[RemovalStep] = []
|
||||
|
||||
|
||||
def register(step: RemovalStep) -> RemovalStep:
|
||||
_REGISTRY.append(step)
|
||||
return step
|
||||
|
||||
|
||||
def find_removal_step(provider: str, source: str) -> Optional[RemovalStep]:
|
||||
"""Return the first matching RemovalStep, or None if unregistered.
|
||||
|
||||
Unregistered sources fall through to the default remove path in
|
||||
``auth_remove_command``: the pool entry is already gone (that happens
|
||||
before dispatch), no external cleanup, no suppression. This is the
|
||||
correct behaviour for ``manual`` entries — they were only ever stored
|
||||
in the pool, nothing external to clean up.
|
||||
"""
|
||||
for step in _REGISTRY:
|
||||
if step.matches(provider, source):
|
||||
return step
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Individual RemovalStep implementations — one per source.
|
||||
# ---------------------------------------------------------------------------
|
||||
# Each remove_fn is intentionally small and single-purpose. Adding a new
|
||||
# credential source means adding ONE entry here — no other changes to
|
||||
# auth_remove_command.
|
||||
"""First matching RemovalStep, or None (``manual``: nothing external to clean)."""
|
||||
return next((step for step in _REGISTRY if step.matches(provider, source)), None)
|
||||
|
||||
|
||||
def _remove_env_source(provider: str, removed) -> RemovalResult:
|
||||
"""env:<VAR> — the most common case.
|
||||
|
||||
Handles three user situations:
|
||||
1. Var lives only in ~/.hermes/.env → clear it
|
||||
2. Var lives only in the user's shell (shell profile, systemd
|
||||
EnvironmentFile, launchd plist) → hint them where to unset it
|
||||
3. Var lives in both → clear from .env, hint about shell
|
||||
"""
|
||||
"""env:<VAR> — clear from ~/.hermes/.env; hint when the shell exports it."""
|
||||
from hermes_cli.config import get_env_path, remove_env_value
|
||||
|
||||
result = RemovalResult()
|
||||
@@ -156,33 +68,25 @@ def _remove_env_source(provider: str, removed) -> RemovalResult:
|
||||
if not env_var:
|
||||
return result
|
||||
|
||||
# Detect shell vs .env BEFORE remove_env_value pops os.environ.
|
||||
# Detect shell vs .env BEFORE remove_env_value pops os.environ. Read the
|
||||
# .env as utf-8-sig like hermes_cli/config.py: a BOM-sensitive read would
|
||||
# misreport a Notepad-edited .env var as a shell export.
|
||||
env_in_process = bool(os.getenv(env_var))
|
||||
env_in_dotenv = False
|
||||
try:
|
||||
env_path = get_env_path()
|
||||
if env_path.exists():
|
||||
# Read the .env as UTF-8 with BOM tolerance, matching the
|
||||
# canonical reader in hermes_cli/config.py. read_text() with no
|
||||
# encoding falls back to the system locale (cp1252/GBK on Windows)
|
||||
# and never strips a BOM, so a Notepad-edited .env (BOM + non-ASCII
|
||||
# values) would make the first line fail the startswith() check —
|
||||
# misreporting a .env-backed var as a shell export.
|
||||
env_in_dotenv = any(
|
||||
line.strip().startswith(f"{env_var}=")
|
||||
for line in env_path.read_text(
|
||||
encoding="utf-8-sig", errors="replace"
|
||||
).splitlines()
|
||||
for line in env_path.read_text(encoding="utf-8-sig", errors="replace").splitlines()
|
||||
)
|
||||
except OSError:
|
||||
pass
|
||||
shell_exported = env_in_process and not env_in_dotenv
|
||||
|
||||
cleared = remove_env_value(env_var)
|
||||
if cleared:
|
||||
if remove_env_value(env_var):
|
||||
result.cleaned.append(f"Cleared {env_var} from .env")
|
||||
|
||||
if shell_exported:
|
||||
if env_in_process and not env_in_dotenv:
|
||||
result.hints.extend([
|
||||
f"Note: {env_var} is still set in your shell environment "
|
||||
f"(not in ~/.hermes/.env).",
|
||||
@@ -199,19 +103,6 @@ def _remove_env_source(provider: str, removed) -> RemovalResult:
|
||||
return result
|
||||
|
||||
|
||||
def _remove_claude_code(provider: str, removed) -> RemovalResult:
|
||||
"""~/.claude/.credentials.json is owned by Claude Code itself.
|
||||
|
||||
We don't delete it — the user's Claude Code install still needs to
|
||||
work. We just suppress it so Hermes stops reading it.
|
||||
"""
|
||||
return RemovalResult(hints=[
|
||||
"Suppressed claude_code credential — it will not be re-seeded.",
|
||||
"Note: Claude Code credentials still live in ~/.claude/.credentials.json",
|
||||
"Run `hermes auth add anthropic` to re-enable if needed.",
|
||||
])
|
||||
|
||||
|
||||
def _remove_hermes_pkce(provider: str, removed) -> RemovalResult:
|
||||
"""~/.hermes/.anthropic_oauth.json is ours — delete it outright."""
|
||||
from hermes_constants import get_hermes_home
|
||||
@@ -227,66 +118,27 @@ def _remove_hermes_pkce(provider: str, removed) -> RemovalResult:
|
||||
return result
|
||||
|
||||
|
||||
def _clear_auth_store_provider(provider: str) -> bool:
|
||||
"""Delete auth_store.providers[provider]. Returns True if deleted."""
|
||||
from hermes_cli.auth import (
|
||||
_auth_store_lock,
|
||||
_load_auth_store,
|
||||
_save_auth_store,
|
||||
)
|
||||
def _remove_auth_store_oauth(provider: str, removed) -> RemovalResult:
|
||||
"""Clear auth.json ``providers.<provider>`` (nous, minimax-oauth, xai-oauth, openai-codex).
|
||||
|
||||
Suppression by the dispatcher is still required — otherwise
|
||||
``_seed_from_singletons`` re-seeds from any path that rewrites the block.
|
||||
"""
|
||||
from hermes_cli.auth import _auth_store_lock, _load_auth_store, _save_auth_store
|
||||
|
||||
result = RemovalResult()
|
||||
with _auth_store_lock():
|
||||
auth_store = _load_auth_store()
|
||||
providers_dict = auth_store.get("providers")
|
||||
if isinstance(providers_dict, dict) and provider in providers_dict:
|
||||
del providers_dict[provider]
|
||||
_save_auth_store(auth_store)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _remove_nous_device_code(provider: str, removed) -> RemovalResult:
|
||||
"""Nous OAuth lives in auth.json providers.nous — clear it and suppress.
|
||||
|
||||
We suppress in addition to clearing because nothing else stops a future
|
||||
`hermes auth add nous` (or any other path that writes providers.nous)
|
||||
from re-seeding before the user has decided to. Suppression forces
|
||||
them to go through `hermes auth add nous` to re-engage, which is the
|
||||
documented re-add path and clears the suppression atomically.
|
||||
"""
|
||||
result = RemovalResult()
|
||||
if _clear_auth_store_provider(provider):
|
||||
result.cleaned.append(f"Cleared {provider} OAuth tokens from auth store")
|
||||
return result
|
||||
|
||||
|
||||
def _remove_minimax_oauth(provider: str, removed) -> RemovalResult:
|
||||
"""MiniMax OAuth lives in auth.json providers.minimax-oauth — clear it.
|
||||
|
||||
Same pattern as Nous: single-source OAuth state with refresh tokens.
|
||||
Suppression of the `oauth` source ensures the pool reseed path
|
||||
(_seed_from_singletons) doesn't instantly undo the removal.
|
||||
"""
|
||||
result = RemovalResult()
|
||||
if _clear_auth_store_provider(provider):
|
||||
result.cleaned.append(f"Cleared {provider} OAuth tokens from auth store")
|
||||
result.cleaned.append(f"Cleared {provider} OAuth tokens from auth store")
|
||||
return result
|
||||
|
||||
|
||||
def _remove_xai_oauth_device_code(provider: str, removed) -> RemovalResult:
|
||||
"""xAI OAuth tokens live in auth.json providers.xai-oauth — clear them.
|
||||
|
||||
Without this step, ``hermes auth remove xai-oauth <N>`` silently undoes
|
||||
itself: the central dispatcher only removes the in-memory pool entry,
|
||||
leaves ``providers.xai-oauth`` in auth.json intact, and on the next
|
||||
``load_pool("xai-oauth")`` call ``_seed_from_singletons`` re-seeds the
|
||||
entry from the still-present singleton — credentials reappear with no
|
||||
user feedback. Clearing the singleton in step with the suppression set
|
||||
by the central dispatcher makes the removal stick.
|
||||
"""
|
||||
result = RemovalResult()
|
||||
if _clear_auth_store_provider(provider):
|
||||
result.cleaned.append(f"Cleared {provider} OAuth tokens from auth store")
|
||||
result = _remove_auth_store_oauth(provider, removed)
|
||||
result.hints.append(
|
||||
"Run `hermes model` → xAI Grok OAuth (SuperGrok / Premium+) to re-authenticate if needed."
|
||||
)
|
||||
@@ -294,30 +146,14 @@ def _remove_xai_oauth_device_code(provider: str, removed) -> RemovalResult:
|
||||
|
||||
|
||||
def _remove_codex_device_code(provider: str, removed) -> RemovalResult:
|
||||
"""Codex tokens live in TWO places: our auth store AND ~/.codex/auth.json.
|
||||
"""Codex tokens also live in ~/.codex/auth.json (Codex CLI's file, kept).
|
||||
|
||||
refresh_codex_oauth_pure() writes both every time, so clearing only
|
||||
the Hermes auth store is not enough — _seed_from_singletons() would
|
||||
re-import from ~/.codex/auth.json on the next load_pool() call and
|
||||
the removal would be instantly undone. We suppress instead of
|
||||
deleting Codex CLI's file, so the Codex CLI itself keeps working.
|
||||
|
||||
The canonical source name in ``_seed_from_singletons`` is
|
||||
``"device_code"`` (no prefix). Entries may show up in the pool as
|
||||
either ``"device_code"`` (seeded) or ``"manual:device_code"`` (added
|
||||
via ``hermes auth add openai-codex``), but in both cases the re-seed
|
||||
gate lives at the ``"device_code"`` suppression key. We suppress
|
||||
that canonical key here; the central dispatcher also suppresses
|
||||
``removed.source`` which is fine — belt-and-suspenders, idempotent.
|
||||
Suppress the canonical ``device_code`` key — not just ``removed.source`` —
|
||||
so a ``manual:device_code`` removal still blocks the re-seed path.
|
||||
"""
|
||||
from hermes_cli.auth import suppress_credential_source
|
||||
|
||||
result = RemovalResult()
|
||||
if _clear_auth_store_provider(provider):
|
||||
result.cleaned.append(f"Cleared {provider} OAuth tokens from auth store")
|
||||
# Suppress the canonical re-seed source, not just whatever source the
|
||||
# removed entry had. Otherwise `manual:device_code` removals wouldn't
|
||||
# block the `device_code` re-seed path.
|
||||
result = _remove_auth_store_oauth(provider, removed)
|
||||
suppress_credential_source(provider, "device_code")
|
||||
result.hints.extend([
|
||||
"Suppressed openai-codex device_code source — it will not be re-seeded.",
|
||||
@@ -327,41 +163,15 @@ def _remove_codex_device_code(provider: str, removed) -> RemovalResult:
|
||||
return result
|
||||
|
||||
|
||||
def _remove_qwen_cli(provider: str, removed) -> RemovalResult:
|
||||
"""~/.qwen/oauth_creds.json is owned by the Qwen CLI.
|
||||
|
||||
Same pattern as claude_code — suppress, don't delete. The user's
|
||||
Qwen CLI install still reads from that file.
|
||||
"""
|
||||
return RemovalResult(hints=[
|
||||
"Suppressed qwen-cli credential — it will not be re-seeded.",
|
||||
"Note: Qwen CLI credentials still live in ~/.qwen/oauth_creds.json",
|
||||
"Run `hermes auth add qwen-oauth` to re-enable if needed.",
|
||||
])
|
||||
|
||||
|
||||
def _remove_copilot_gh(provider: str, removed) -> RemovalResult:
|
||||
"""Copilot token comes from `gh auth token` or COPILOT_GITHUB_TOKEN / GH_TOKEN / GITHUB_TOKEN.
|
||||
|
||||
Copilot is special: the same token can be seeded as multiple source
|
||||
entries (gh_cli from ``_seed_from_singletons`` plus env:<VAR> from
|
||||
``_seed_from_env``), so removing one entry without suppressing the
|
||||
others lets the duplicates resurrect. We suppress ALL known copilot
|
||||
sources here so removal is stable regardless of which entry the
|
||||
user clicked.
|
||||
|
||||
We don't touch the user's gh CLI or shell state — just suppress so
|
||||
Hermes stops picking the token up.
|
||||
"""
|
||||
# Suppress ALL copilot source variants up-front so no path resurrects
|
||||
# the pool entry. The central dispatcher in auth_remove_command will
|
||||
# ALSO suppress removed.source, but it's idempotent so double-calling
|
||||
# is harmless.
|
||||
"""The same Copilot token is seeded as gh_cli AND env:<VAR> rows, so suppress
|
||||
every variant or the duplicates resurrect the entry. gh CLI and shell state
|
||||
are left untouched."""
|
||||
from hermes_cli.auth import suppress_credential_source
|
||||
|
||||
suppress_credential_source(provider, "gh_cli")
|
||||
for env_var in ("COPILOT_GITHUB_TOKEN", "GH_TOKEN", "GITHUB_TOKEN"):
|
||||
suppress_credential_source(provider, f"env:{env_var}")
|
||||
|
||||
return RemovalResult(hints=[
|
||||
"Suppressed all copilot token sources (gh_cli + env vars) — they will not be re-seeded.",
|
||||
"Note: Your gh CLI / shell environment is unchanged.",
|
||||
@@ -369,83 +179,94 @@ def _remove_copilot_gh(provider: str, removed) -> RemovalResult:
|
||||
])
|
||||
|
||||
|
||||
def _remove_custom_config(provider: str, removed) -> RemovalResult:
|
||||
"""Custom provider pools are seeded from custom_providers config or
|
||||
model.api_key. Both are in config.yaml — modifying that from here
|
||||
is more invasive than suppression. We suppress; the user can edit
|
||||
config.yaml if they want to remove the key from disk entirely.
|
||||
"""
|
||||
source_label = removed.source
|
||||
return RemovalResult(hints=[
|
||||
f"Suppressed {source_label} — it will not be re-seeded.",
|
||||
"Note: The underlying value in config.yaml is unchanged. Edit it "
|
||||
"directly if you want to remove the credential from disk.",
|
||||
])
|
||||
def _suppress_only(*hints: str) -> Callable[..., RemovalResult]:
|
||||
"""remove_fn for sources whose backing file/config belongs to another tool
|
||||
(Claude Code, Qwen CLI, config.yaml): never delete it, just suppress and
|
||||
explain. ``{source}`` in a hint is the removed entry's source string."""
|
||||
def remove_fn(provider: str, removed) -> RemovalResult:
|
||||
return RemovalResult(hints=[h.format(source=removed.source) for h in hints])
|
||||
return remove_fn
|
||||
|
||||
|
||||
def _register_all_sources() -> None:
|
||||
"""Called once on module import.
|
||||
|
||||
ORDER MATTERS — ``find_removal_step`` returns the first match. Put
|
||||
provider-specific steps before the generic ``env:*`` step so that e.g.
|
||||
copilot's ``env:GH_TOKEN`` goes through the copilot removal (which
|
||||
doesn't touch the user's shell), not the generic env-var removal
|
||||
(which would try to clear .env).
|
||||
"""
|
||||
register(RemovalStep(
|
||||
# ORDER MATTERS — ``find_removal_step`` returns the first match. Provider-
|
||||
# specific steps precede the generic ``env:*`` step so copilot's ``env:GH_TOKEN``
|
||||
# takes the copilot path (no .env edits) rather than the generic env-var removal.
|
||||
_REGISTRY: List[RemovalStep] = [
|
||||
RemovalStep(
|
||||
provider="copilot", source_id="gh_cli",
|
||||
match_fn=lambda src: src == "gh_cli" or src.startswith("env:"),
|
||||
remove_fn=_remove_copilot_gh,
|
||||
description="gh auth token / COPILOT_GITHUB_TOKEN / GH_TOKEN",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="*", source_id="env:",
|
||||
match_fn=lambda src: src.startswith("env:"),
|
||||
remove_fn=_remove_env_source,
|
||||
description="Any env-seeded credential (XAI_API_KEY, DEEPSEEK_API_KEY, etc.)",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="anthropic", source_id="claude_code",
|
||||
remove_fn=_remove_claude_code,
|
||||
remove_fn=_suppress_only(
|
||||
"Suppressed claude_code credential — it will not be re-seeded.",
|
||||
"Note: Claude Code credentials still live in ~/.claude/.credentials.json",
|
||||
"Run `hermes auth add anthropic` to re-enable if needed.",
|
||||
),
|
||||
description="~/.claude/.credentials.json",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="anthropic", source_id="hermes_pkce",
|
||||
remove_fn=_remove_hermes_pkce,
|
||||
description="~/.hermes/.anthropic_oauth.json",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="nous", source_id="device_code",
|
||||
remove_fn=_remove_nous_device_code,
|
||||
remove_fn=_remove_auth_store_oauth,
|
||||
description="auth.json providers.nous",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="openai-codex", source_id="device_code",
|
||||
match_fn=lambda src: src == "device_code" or src.endswith(":device_code"),
|
||||
remove_fn=_remove_codex_device_code,
|
||||
description="auth.json providers.openai-codex + ~/.codex/auth.json",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="xai-oauth", source_id="device_code",
|
||||
remove_fn=_remove_xai_oauth_device_code,
|
||||
description="auth.json providers.xai-oauth",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="qwen-oauth", source_id="qwen-cli",
|
||||
remove_fn=_remove_qwen_cli,
|
||||
remove_fn=_suppress_only(
|
||||
"Suppressed qwen-cli credential — it will not be re-seeded.",
|
||||
"Note: Qwen CLI credentials still live in ~/.qwen/oauth_creds.json",
|
||||
"Run `hermes auth add qwen-oauth` to re-enable if needed.",
|
||||
),
|
||||
description="~/.qwen/oauth_creds.json",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="minimax-oauth", source_id="oauth",
|
||||
remove_fn=_remove_minimax_oauth,
|
||||
remove_fn=_remove_auth_store_oauth,
|
||||
description="auth.json providers.minimax-oauth",
|
||||
))
|
||||
register(RemovalStep(
|
||||
),
|
||||
RemovalStep(
|
||||
provider="*", source_id="config:",
|
||||
match_fn=lambda src: src.startswith("config:") or src == "model_config",
|
||||
remove_fn=_remove_custom_config,
|
||||
remove_fn=_suppress_only(
|
||||
"Suppressed {source} — it will not be re-seeded.",
|
||||
"Note: The underlying value in config.yaml is unchanged. Edit it "
|
||||
"directly if you want to remove the credential from disk.",
|
||||
),
|
||||
description="Custom provider config.yaml api_key field",
|
||||
))
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
_register_all_sources()
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
|
||||
def register(step: RemovalStep) -> RemovalStep:
|
||||
_REGISTRY.append(step)
|
||||
return step
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
+216
-667
File diff suppressed because it is too large
Load Diff
+602
-1559
File diff suppressed because it is too large
Load Diff
+230
-547
@@ -1,165 +1,97 @@
|
||||
"""Curator snapshot + rollback.
|
||||
|
||||
A pre-run snapshot of ``~/.hermes/skills/`` (excluding ``.curator_backups/``
|
||||
itself) is taken before any mutating curator pass. Snapshots are tar.gz
|
||||
files under ``~/.hermes/skills/.curator_backups/<utc-iso>/`` with a
|
||||
companion ``manifest.json`` describing the snapshot (reason, time, size,
|
||||
counted skill files). Rollback picks a snapshot, moves the current
|
||||
``skills/`` tree aside into another snapshot so even the rollback itself
|
||||
is undoable, then extracts the chosen snapshot into place.
|
||||
|
||||
The snapshot does NOT include:
|
||||
- ``.curator_backups/`` (would recurse)
|
||||
- ``.hub/`` (hub-installed skills — managed by the hub, not us)
|
||||
|
||||
It DOES include:
|
||||
- all SKILL.md files + their directories (``scripts/``, ``references/``,
|
||||
``templates/``, ``assets/``)
|
||||
- ``.usage.json`` (usage telemetry — needed to rehydrate state cleanly)
|
||||
- ``.archive/`` (so rollback restores previously-archived skills too)
|
||||
- ``.curator_state`` (so rolling back also restores the last-run-at
|
||||
pointer — otherwise the curator would immediately re-fire on the next
|
||||
tick)
|
||||
- ``.bundled_manifest`` (so protection markers stay consistent)
|
||||
- ``.curator_suppressed`` (so rollback restores the set of pruned built-ins
|
||||
the re-seeder must leave archived)
|
||||
|
||||
Alongside the skills tarball, each snapshot also captures a copy of
|
||||
``~/.hermes/cron/jobs.json`` as ``cron-jobs.json`` when it exists. Cron
|
||||
jobs reference skills by name in their ``skills``/``skill`` fields; the
|
||||
curator's consolidation pass rewrites those in place via
|
||||
``cron.jobs.rewrite_skill_refs()``. Without capturing the pre-run state,
|
||||
rolling back the skills tree would leave cron jobs pointing at the
|
||||
umbrella skills even though the narrow skills they were originally
|
||||
configured with have been restored. We store the whole jobs.json for
|
||||
fidelity but rollback only touches the ``skills``/``skill`` fields — the
|
||||
rest (schedule, next_run_at, enabled, prompt, etc.) is live state and
|
||||
we leave it alone.
|
||||
"""
|
||||
"""Curator snapshot + rollback. Before any mutating curator pass, ``~/.hermes/skills/`` is tar.gz'd under
|
||||
``~/.hermes/skills/.curator_backups/<utc-iso>/`` with a ``manifest.json``. Rollback first snapshots the CURRENT tree (so it is
|
||||
itself undoable), then extracts the chosen snapshot into place. Excluded: ``.curator_backups/``, ``.hub/`` (hub-managed), ``.git/``.
|
||||
Included: skill dirs, ``.usage.json``, ``.archive/``, ``.curator_state`` (so rollback also restores last-run-at and the curator
|
||||
doesn't re-fire), ``.bundled_manifest``, ``.curator_suppressed``. Each snapshot also copies ``~/.hermes/cron/jobs.json`` as
|
||||
``cron-jobs.json``: the consolidation pass rewrites cron ``skills``/``skill`` references in place, so rollback restores those two
|
||||
fields (only) — the rest is live state."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import tarfile
|
||||
from datetime import datetime, timezone
|
||||
from itertools import chain, count
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Set, Tuple
|
||||
|
||||
from hermes_constants import get_hermes_home
|
||||
from agent.skill_utils import is_excluded_skill_path
|
||||
from agent.curator import _read_config_section
|
||||
from hermes_cli.sizefmt import format_bytes
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
DEFAULT_KEEP = 5
|
||||
|
||||
# Entries under skills/ that should NEVER be rolled up into a snapshot.
|
||||
# .hub/ is managed by the skills hub; rolling it back would break lockfile
|
||||
# invariants. .curator_backups is the backup dir itself — recursion bomb.
|
||||
_EXCLUDE_TOP_LEVEL = {".curator_backups", ".hub"}
|
||||
# Never rolled into a snapshot: .hub/ is owned by the skills hub (rolling it back breaks lockfile invariants); .curator_backups
|
||||
# is the backup dir itself; .git is repository metadata — rolling it back breaks git tracking, and snapshots that include it grow
|
||||
# with the full history (once backups are committed back, each snapshot contains the prior ones: 38MB of skills inflated to 24GB
|
||||
# in weeks). The tar filter in ``snapshot_skills`` applies the same set to nested paths, so a nested ``.git`` is skipped too.
|
||||
# See #91449.
|
||||
_EXCLUDE_TOP_LEVEL = {".curator_backups", ".hub", ".git"}
|
||||
|
||||
# Snapshot id regex: UTC ISO with colons replaced by dashes so the filename
|
||||
# is portable (Windows-safe). An optional ``-NN`` suffix handles two
|
||||
# snapshots landing in the same wallclock second.
|
||||
# Snapshot id: UTC ISO with colons replaced by dashes (Windows-safe filename); optional ``-NN`` suffix for same-second snapshots.
|
||||
_ID_RE = re.compile(r"^\d{4}-\d{2}-\d{2}T\d{2}-\d{2}-\d{2}Z(-\d{2})?$")
|
||||
|
||||
|
||||
def _backups_dir() -> Path:
|
||||
return get_hermes_home() / "skills" / ".curator_backups"
|
||||
CRON_JOBS_FILENAME = "cron-jobs.json"
|
||||
_ARCHIVE_NAME = "skills.tar.gz"
|
||||
_STAGING_PREFIX = ".rollback-staging-"
|
||||
|
||||
|
||||
def _skills_dir() -> Path:
|
||||
return get_hermes_home() / "skills"
|
||||
|
||||
|
||||
def _cron_jobs_file() -> Path:
|
||||
"""Source path for the live cron jobs store (``~/.hermes/cron/jobs.json``)."""
|
||||
return get_hermes_home() / "cron" / "jobs.json"
|
||||
def _backups_dir() -> Path:
|
||||
return _skills_dir() / ".curator_backups"
|
||||
|
||||
|
||||
CRON_JOBS_FILENAME = "cron-jobs.json"
|
||||
def _jobs_list(parsed: Any) -> Optional[list]:
|
||||
"""jobs.json is ``{"jobs": [...], "updated_at": ...}``; also accept a bare list for forward compat. None otherwise."""
|
||||
parsed = parsed.get("jobs") if isinstance(parsed, dict) else parsed
|
||||
return parsed if isinstance(parsed, list) else None
|
||||
|
||||
|
||||
def _backup_cron_jobs_into(dest: Path) -> Dict[str, Any]:
|
||||
"""Copy the live cron jobs.json into ``dest`` as ``cron-jobs.json``.
|
||||
|
||||
Returns a small dict describing what was captured so the caller can
|
||||
fold it into the manifest. Never raises — if the cron file is missing
|
||||
or unreadable, the return dict has ``backed_up=False`` and the reason,
|
||||
and the snapshot proceeds without cron data (the snapshot is still
|
||||
useful for rolling back skills).
|
||||
"""
|
||||
src = _cron_jobs_file()
|
||||
"""Copy the live ``~/.hermes/cron/jobs.json`` into ``dest`` as ``cron-jobs.json``. Never raises: a missing/unreadable
|
||||
file yields ``backed_up=False`` plus a reason, and the snapshot proceeds."""
|
||||
src = get_hermes_home() / "cron" / "jobs.json"
|
||||
info: Dict[str, Any] = {"backed_up": False, "jobs_count": 0}
|
||||
if not src.exists():
|
||||
info["reason"] = "no cron/jobs.json present"
|
||||
return info
|
||||
return {**info, "reason": "no cron/jobs.json present"}
|
||||
try:
|
||||
# utf-8-sig: same dialect as cron/jobs.load_jobs — a UTF-8 BOM left
|
||||
# by Windows editors otherwise survives decoding as U+FEFF, breaks
|
||||
# json.loads below, and misreports jobs_count as 0 with a spurious
|
||||
# parse warning. The BOM-less text is also what gets written to the
|
||||
# backup, so a later rollback restores a loadable file.
|
||||
# utf-8-sig, same dialect as cron/jobs.load_jobs: a Windows-editor BOM would otherwise break json.loads
|
||||
# AND be written into the backup.
|
||||
raw = src.read_text(encoding="utf-8-sig")
|
||||
except OSError as e:
|
||||
logger.debug("Failed to read cron/jobs.json for backup: %s", e)
|
||||
info["reason"] = f"read error: {e}"
|
||||
return info
|
||||
# Count jobs as a nice diagnostic — but don't fail the snapshot if the
|
||||
# file is unparseable; just store the raw text and let rollback deal
|
||||
# with it (or not, if it's corrupted). jobs.json wraps the list as
|
||||
# `{"jobs": [...], "updated_at": ...}` — we count via that shape, and
|
||||
# fall back to bare-list shape just in case the format ever changes.
|
||||
try:
|
||||
parsed = json.loads(raw)
|
||||
if isinstance(parsed, dict):
|
||||
inner = parsed.get("jobs")
|
||||
if isinstance(inner, list):
|
||||
info["jobs_count"] = len(inner)
|
||||
elif isinstance(parsed, list):
|
||||
info["jobs_count"] = len(parsed)
|
||||
return {**info, "reason": f"read error: {e}"}
|
||||
try: # jobs_count is a diagnostic only — an unparseable file is still stored raw.
|
||||
info["jobs_count"] = len(_jobs_list(json.loads(raw)) or [])
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
info["jobs_count"] = 0
|
||||
info["parse_warning"] = "jobs.json was not valid JSON at snapshot time"
|
||||
try:
|
||||
(dest / CRON_JOBS_FILENAME).write_text(raw, encoding="utf-8")
|
||||
except OSError as e:
|
||||
logger.debug("Failed to write cron backup file: %s", e)
|
||||
info["reason"] = f"write error: {e}"
|
||||
return info
|
||||
info["backed_up"] = True
|
||||
return info
|
||||
return {**info, "reason": f"write error: {e}"}
|
||||
return {**info, "backed_up": True}
|
||||
|
||||
|
||||
def _utc_id(now: Optional[datetime] = None) -> str:
|
||||
"""UTC ISO-ish filesystem-safe timestamp: ``2026-05-01T13-05-42Z``."""
|
||||
if now is None:
|
||||
now = datetime.now(timezone.utc)
|
||||
# isoformat → "2026-05-01T13:05:42.123456+00:00"; strip subseconds and tz.
|
||||
s = now.replace(microsecond=0).isoformat()
|
||||
if s.endswith("+00:00"):
|
||||
s = s[:-6]
|
||||
return s.replace(":", "-") + "Z"
|
||||
s = (datetime.now(timezone.utc) if now is None else now).replace(microsecond=0).isoformat()
|
||||
return s.removesuffix("+00:00").replace(":", "-") + "Z"
|
||||
|
||||
|
||||
def _load_config() -> Dict[str, Any]:
|
||||
try:
|
||||
from hermes_cli.config import load_config_readonly
|
||||
cfg = load_config_readonly()
|
||||
except Exception as e:
|
||||
logger.debug("Failed to load config for curator backup: %s", e)
|
||||
return {}
|
||||
if not isinstance(cfg, dict):
|
||||
return {}
|
||||
cur = cfg.get("curator") or {}
|
||||
if not isinstance(cur, dict):
|
||||
return {}
|
||||
bk = cur.get("backup") or {}
|
||||
return bk if isinstance(bk, dict) else {}
|
||||
return _read_config_section("curator", "backup", label="curator backup", log=logger)
|
||||
|
||||
|
||||
def is_enabled() -> bool:
|
||||
@@ -168,122 +100,73 @@ def is_enabled() -> bool:
|
||||
|
||||
|
||||
def get_keep() -> int:
|
||||
cfg = _load_config()
|
||||
try:
|
||||
n = int(cfg.get("keep", DEFAULT_KEEP))
|
||||
return max(1, int(_load_config().get("keep", DEFAULT_KEEP)))
|
||||
except (TypeError, ValueError):
|
||||
n = DEFAULT_KEEP
|
||||
return max(1, n)
|
||||
return DEFAULT_KEEP
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Snapshot
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# --- Snapshot ---
|
||||
def _count_skill_files(base: Path) -> int:
|
||||
try:
|
||||
return sum(
|
||||
1 for p in base.rglob("SKILL.md") if not is_excluded_skill_path(p)
|
||||
)
|
||||
return sum(1 for p in base.rglob("SKILL.md") if not is_excluded_skill_path(p))
|
||||
except OSError:
|
||||
return 0
|
||||
|
||||
|
||||
def _write_manifest(dest: Path, reason: str, archive_path: Path,
|
||||
skills_counted: int,
|
||||
cron_info: Optional[Dict[str, Any]] = None) -> None:
|
||||
manifest = {
|
||||
"id": dest.name,
|
||||
"reason": reason,
|
||||
"created_at": datetime.now(timezone.utc).isoformat(),
|
||||
"archive": archive_path.name,
|
||||
"archive_bytes": archive_path.stat().st_size,
|
||||
"skill_files": skills_counted,
|
||||
}
|
||||
if cron_info is not None:
|
||||
manifest["cron_jobs"] = {
|
||||
"backed_up": bool(cron_info.get("backed_up", False)),
|
||||
"jobs_count": int(cron_info.get("jobs_count", 0)),
|
||||
}
|
||||
if not cron_info.get("backed_up"):
|
||||
manifest["cron_jobs"]["reason"] = cron_info.get("reason", "not captured")
|
||||
if cron_info.get("parse_warning"):
|
||||
manifest["cron_jobs"]["parse_warning"] = cron_info["parse_warning"]
|
||||
(dest / "manifest.json").write_text(
|
||||
json.dumps(manifest, indent=2, sort_keys=True), encoding="utf-8"
|
||||
)
|
||||
def _write_manifest(dest: Path, reason: str, archive_path: Path, skills_counted: int, cron_info: Dict[str, Any]) -> None:
|
||||
cron_jobs: Dict[str, Any] = {"backed_up": bool(cron_info.get("backed_up", False)), "jobs_count": int(cron_info.get("jobs_count", 0))}
|
||||
if not cron_info.get("backed_up"):
|
||||
cron_jobs["reason"] = cron_info.get("reason", "not captured")
|
||||
if cron_info.get("parse_warning"):
|
||||
cron_jobs["parse_warning"] = cron_info["parse_warning"]
|
||||
manifest = {"id": dest.name, "reason": reason, "created_at": datetime.now(timezone.utc).isoformat(), "archive": archive_path.name,
|
||||
"archive_bytes": archive_path.stat().st_size, "skill_files": skills_counted, "cron_jobs": cron_jobs}
|
||||
(dest / "manifest.json").write_text(json.dumps(manifest, indent=2, sort_keys=True), encoding="utf-8")
|
||||
|
||||
|
||||
def _mkdir(path: Path, what: str, *, exist_ok: bool) -> bool:
|
||||
try:
|
||||
path.mkdir(parents=True, exist_ok=exist_ok)
|
||||
return True
|
||||
except OSError as e:
|
||||
logger.debug("Failed to create %s %s: %s", what, path, e)
|
||||
return False
|
||||
|
||||
|
||||
def snapshot_skills(reason: str = "manual", *, protect_ids: Optional[Set[str]] = None) -> Optional[Path]:
|
||||
"""Create a tar.gz snapshot of ``~/.hermes/skills/`` and prune old ones.
|
||||
|
||||
Returns the snapshot directory path, or ``None`` if the snapshot was
|
||||
skipped (backup disabled, skills dir missing, or an IO error occurred —
|
||||
in which case we log at debug and return None so the curator never
|
||||
aborts a pass because of a backup failure).
|
||||
|
||||
``protect_ids`` is forwarded to the prune step so callers can guarantee
|
||||
specific snapshot ids survive even when they fall outside the keep
|
||||
window (rollback passes the id it is about to restore from).
|
||||
"""
|
||||
"""Create a tar.gz snapshot of ``~/.hermes/skills/`` and prune old ones. Returns the snapshot dir, or None when
|
||||
skipped (disabled, skills dir missing, IO error) — logged at debug so the curator never aborts a pass over a
|
||||
backup failure. ``protect_ids`` survive the prune step (rollback protects its target)."""
|
||||
if not is_enabled():
|
||||
logger.debug("Curator backup disabled by config; skipping snapshot")
|
||||
return None
|
||||
|
||||
skills = _skills_dir()
|
||||
skills, backups = _skills_dir(), _backups_dir()
|
||||
if not skills.exists():
|
||||
logger.debug("No ~/.hermes/skills/ directory — nothing to back up")
|
||||
return None
|
||||
|
||||
backups = _backups_dir()
|
||||
try:
|
||||
backups.mkdir(parents=True, exist_ok=True)
|
||||
except OSError as e:
|
||||
logger.debug("Failed to create backups dir %s: %s", backups, e)
|
||||
if not _mkdir(backups, "backups dir", exist_ok=True):
|
||||
return None
|
||||
|
||||
# Uniquify: if a snapshot with the same second already exists (can
|
||||
# happen if two curator runs fire in the same second), append a short
|
||||
# counter. Avoids clobbering and avoids timestamp collisions.
|
||||
base_id = _utc_id()
|
||||
snap_id = base_id
|
||||
counter = 1
|
||||
while (backups / snap_id).exists():
|
||||
snap_id = f"{base_id}-{counter:02d}"
|
||||
counter += 1
|
||||
|
||||
base_id = _utc_id() # Two curator runs in the same second must not clobber each other: -NN suffix.
|
||||
snap_id = next(i for i in chain([base_id], (f"{base_id}-{n:02d}" for n in count(1))) if not (backups / i).exists())
|
||||
dest = backups / snap_id
|
||||
try:
|
||||
dest.mkdir(parents=True, exist_ok=False)
|
||||
except OSError as e:
|
||||
logger.debug("Failed to create snapshot dir %s: %s", dest, e)
|
||||
if not _mkdir(dest, "snapshot dir", exist_ok=False):
|
||||
return None
|
||||
|
||||
archive = dest / "skills.tar.gz"
|
||||
archive = dest / _ARCHIVE_NAME
|
||||
try:
|
||||
# Stream into the tarball — no tempdir copy needed.
|
||||
with tarfile.open(archive, "w:gz", compresslevel=6) as tf:
|
||||
for entry in sorted(skills.iterdir()):
|
||||
if entry.name in _EXCLUDE_TOP_LEVEL:
|
||||
continue
|
||||
# arcname: store paths relative to skills/ so extraction
|
||||
# drops cleanly back into the skills dir.
|
||||
tf.add(str(entry), arcname=entry.name, recursive=True)
|
||||
# Capture cron/jobs.json alongside the tarball. Never fails the
|
||||
# snapshot — the skills side is the core guarantee; cron is
|
||||
# additive. We still record in the manifest whether it was
|
||||
# captured so rollback can surface "no cron data in this snapshot".
|
||||
cron_info = _backup_cron_jobs_into(dest)
|
||||
_write_manifest(dest, reason, archive,
|
||||
_count_skill_files(skills),
|
||||
cron_info=cron_info)
|
||||
if entry.name not in _EXCLUDE_TOP_LEVEL:
|
||||
# arcname relative to skills/ so extraction drops back in cleanly; the filter excludes nested _EXCLUDE_TOP_LEVEL paths too.
|
||||
tf.add(str(entry), arcname=entry.name, recursive=True,
|
||||
filter=lambda ti: None if any(p in _EXCLUDE_TOP_LEVEL for p in Path(ti.name).parts) else ti)
|
||||
# Cron capture is additive and never fails the snapshot; the manifest records whether it happened so rollback can say "no cron data".
|
||||
_write_manifest(dest, reason, archive, _count_skill_files(skills), _backup_cron_jobs_into(dest))
|
||||
except (OSError, tarfile.TarError) as e:
|
||||
logger.debug("Curator snapshot failed: %s", e, exc_info=True)
|
||||
# Clean up partial snapshot
|
||||
try:
|
||||
shutil.rmtree(dest, ignore_errors=True)
|
||||
except OSError:
|
||||
pass
|
||||
shutil.rmtree(dest, ignore_errors=True) # clean up partial snapshot
|
||||
return None
|
||||
|
||||
_prune_old(keep=get_keep(), protect=protect_ids)
|
||||
@@ -292,343 +175,211 @@ def snapshot_skills(reason: str = "manual", *, protect_ids: Optional[Set[str]] =
|
||||
|
||||
|
||||
def _prune_old(keep: int, protect: Optional[Set[str]] = None) -> List[str]:
|
||||
"""Delete regular snapshots beyond the newest *keep*. Returns deleted
|
||||
ids. Snapshot ids in *protect* are never deleted even when they fall
|
||||
outside the keep window — rollback() uses this so the mandatory
|
||||
pre-rollback safety snapshot can never evict the very snapshot being
|
||||
restored. Staging dirs (``.rollback-staging-*``) are implementation
|
||||
detail and pruned independently on every call."""
|
||||
"""Delete regular snapshots beyond the newest *keep*; returns deleted ids. Ids in *protect* are never deleted —
|
||||
rollback() uses this so the mandatory pre-rollback safety snapshot cannot evict the snapshot being restored.
|
||||
Stale ``.rollback-staging-*`` dirs (crashed rollback) are cleaned up on every call."""
|
||||
protect = protect or set()
|
||||
backups = _backups_dir()
|
||||
if not backups.exists():
|
||||
return []
|
||||
entries: List[Tuple[str, Path]] = []
|
||||
stale_staging: List[Path] = []
|
||||
for child in backups.iterdir():
|
||||
if not child.is_dir():
|
||||
continue
|
||||
if child.name.startswith(".rollback-staging-"):
|
||||
# Staging dirs are only supposed to exist briefly during a
|
||||
# rollback. If we find one here (e.g. from a crashed rollback),
|
||||
# clean it up opportunistically.
|
||||
stale_staging.append(child)
|
||||
continue
|
||||
if _ID_RE.match(child.name):
|
||||
entries.append((child.name, child))
|
||||
dirs = [c for c in backups.iterdir() if c.is_dir()]
|
||||
# Newest first (lexicographic works because the id is UTC ISO).
|
||||
entries.sort(key=lambda t: t[0], reverse=True)
|
||||
entries = sorted((c for c in dirs if _ID_RE.match(c.name)), key=lambda c: c.name, reverse=True)
|
||||
doomed = [(p, "prune") for p in entries[keep:] if p.name not in protect]
|
||||
doomed += [(p, "clean stale staging dir") for p in dirs if p.name.startswith(_STAGING_PREFIX)]
|
||||
deleted: List[str] = []
|
||||
for _, path in entries[keep:]:
|
||||
if path.name in protect:
|
||||
continue
|
||||
for path, what in doomed:
|
||||
try:
|
||||
shutil.rmtree(path)
|
||||
deleted.append(path.name)
|
||||
if what == "prune":
|
||||
deleted.append(path.name)
|
||||
except OSError as e:
|
||||
logger.debug("Failed to prune %s: %s", path, e)
|
||||
for path in stale_staging:
|
||||
try:
|
||||
shutil.rmtree(path)
|
||||
except OSError as e:
|
||||
logger.debug("Failed to clean stale staging dir %s: %s", path, e)
|
||||
logger.debug("Failed to %s %s: %s", what, path, e)
|
||||
return deleted
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# List + rollback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# --- List + rollback ---
|
||||
def _read_manifest(snap_dir: Path) -> Dict[str, Any]:
|
||||
mf = snap_dir / "manifest.json"
|
||||
if not mf.exists():
|
||||
return {}
|
||||
try:
|
||||
return json.loads(mf.read_text(encoding="utf-8"))
|
||||
return json.loads((snap_dir / "manifest.json").read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return {}
|
||||
|
||||
|
||||
def list_backups() -> List[Dict[str, Any]]:
|
||||
"""Return all restorable snapshots, newest first. Only entries with a
|
||||
real ``skills.tar.gz`` tarball are listed — transient
|
||||
``.rollback-staging-*`` directories created mid-rollback are
|
||||
implementation detail and not shown."""
|
||||
def _is_restorable(child: Path) -> bool:
|
||||
"""A real snapshot dir with a tarball (excludes ``.rollback-staging-*``)."""
|
||||
return bool(child.is_dir() and _ID_RE.match(child.name) and (child / _ARCHIVE_NAME).exists())
|
||||
|
||||
|
||||
def _restorable_snapshots() -> List[Path]:
|
||||
"""Restorable snapshot dirs, newest first."""
|
||||
backups = _backups_dir()
|
||||
if not backups.exists():
|
||||
return []
|
||||
return [c for c in sorted(backups.iterdir(), reverse=True) if _is_restorable(c)] if backups.exists() else []
|
||||
|
||||
|
||||
def list_backups() -> List[Dict[str, Any]]:
|
||||
"""All restorable snapshots (manifest dicts), newest first."""
|
||||
out: List[Dict[str, Any]] = []
|
||||
for child in sorted(backups.iterdir(), reverse=True):
|
||||
if not child.is_dir():
|
||||
continue
|
||||
if not _ID_RE.match(child.name):
|
||||
continue
|
||||
if not (child / "skills.tar.gz").exists():
|
||||
continue
|
||||
mf = _read_manifest(child)
|
||||
mf.setdefault("id", child.name)
|
||||
mf.setdefault("path", str(child))
|
||||
if "archive_bytes" not in mf:
|
||||
arc = child / "skills.tar.gz"
|
||||
try:
|
||||
mf["archive_bytes"] = arc.stat().st_size
|
||||
except OSError:
|
||||
mf["archive_bytes"] = 0
|
||||
for child in _restorable_snapshots():
|
||||
mf = {"id": child.name, "path": str(child), **_read_manifest(child)}
|
||||
try:
|
||||
mf.setdefault("archive_bytes", (child / _ARCHIVE_NAME).stat().st_size)
|
||||
except OSError:
|
||||
mf.setdefault("archive_bytes", 0)
|
||||
out.append(mf)
|
||||
return out
|
||||
|
||||
|
||||
def _resolve_backup(backup_id: Optional[str]) -> Optional[Path]:
|
||||
"""Return the path of the requested backup, or the newest one if
|
||||
*backup_id* is None. Returns None if no match."""
|
||||
backups = _backups_dir()
|
||||
if not backups.exists():
|
||||
return None
|
||||
"""Path of the requested backup (newest if *backup_id* is None); None if no match."""
|
||||
if backup_id:
|
||||
target = backups / backup_id
|
||||
if (
|
||||
target.is_dir()
|
||||
and _ID_RE.match(backup_id)
|
||||
and (target / "skills.tar.gz").exists()
|
||||
):
|
||||
return target
|
||||
return None
|
||||
candidates = [
|
||||
c for c in sorted(backups.iterdir(), reverse=True)
|
||||
if c.is_dir() and _ID_RE.match(c.name) and (c / "skills.tar.gz").exists()
|
||||
]
|
||||
return candidates[0] if candidates else None
|
||||
target = _backups_dir() / backup_id
|
||||
return target if _ID_RE.match(backup_id) and _is_restorable(target) else None
|
||||
return next(iter(_restorable_snapshots()), None)
|
||||
|
||||
|
||||
def _restore_cron_skill_links(snapshot_dir: Path) -> Dict[str, Any]:
|
||||
"""Reconcile backed-up cron skill links into the live ``cron/jobs.json``.
|
||||
|
||||
We do NOT overwrite the whole cron file. Only the ``skills`` and
|
||||
``skill`` fields are restored, and only on jobs that still exist in the
|
||||
current file (matched by ``id``). Everything else about the job —
|
||||
schedule, next_run_at, last_run_at, enabled, prompt, workdir, hooks —
|
||||
is live state that the user/scheduler has modified since the snapshot;
|
||||
overwriting it would regress unrelated cron activity.
|
||||
|
||||
Rules:
|
||||
- Jobs present in backup AND live, with differing skills → skills restored.
|
||||
- Jobs present in backup AND live, with matching skills → no-op.
|
||||
- Jobs present in backup but gone from live (user deleted the job
|
||||
after the snapshot) → skipped, noted in the return report.
|
||||
- Jobs present in live but not in backup (user created a new cron
|
||||
job after the snapshot) → left untouched.
|
||||
|
||||
Never raises; failures are captured in the return dict. Writes through
|
||||
``cron.jobs`` to pick up the same lock + atomic-write path that tick()
|
||||
uses, so we don't race the scheduler.
|
||||
"""
|
||||
report: Dict[str, Any] = {
|
||||
"attempted": False,
|
||||
"restored": [],
|
||||
"skipped_missing": [],
|
||||
"unchanged": 0,
|
||||
"error": None,
|
||||
}
|
||||
"""Reconcile backed-up cron skill links into the live ``cron/jobs.json``. Only ``skills``/``skill`` are restored,
|
||||
and only on jobs that still exist live (by ``id``) — everything else is live state. Backup-only jobs are skipped
|
||||
and reported; live-only jobs untouched. Never raises; writes through ``cron.jobs`` under the scheduler's lock so
|
||||
we don't race tick()."""
|
||||
report: Dict[str, Any] = {"attempted": False, "restored": [], "skipped_missing": [], "unchanged": 0, "error": None}
|
||||
backup_file = snapshot_dir / CRON_JOBS_FILENAME
|
||||
if not backup_file.exists():
|
||||
report["error"] = f"snapshot has no {CRON_JOBS_FILENAME}"
|
||||
return report
|
||||
|
||||
return {**report, "error": f"snapshot has no {CRON_JOBS_FILENAME}"}
|
||||
try:
|
||||
backup_text = backup_file.read_text(encoding="utf-8")
|
||||
backup_parsed = json.loads(backup_text)
|
||||
backup_jobs = _jobs_list(json.loads(backup_file.read_text(encoding="utf-8")))
|
||||
except (OSError, json.JSONDecodeError) as e:
|
||||
report["error"] = f"failed to load backed-up jobs: {e}"
|
||||
return report
|
||||
# jobs.json on disk is `{"jobs": [...], "updated_at": ...}`; accept both
|
||||
# that shape and a bare list for forward compat.
|
||||
if isinstance(backup_parsed, dict):
|
||||
backup_jobs = backup_parsed.get("jobs")
|
||||
elif isinstance(backup_parsed, list):
|
||||
backup_jobs = backup_parsed
|
||||
else:
|
||||
backup_jobs = None
|
||||
if not isinstance(backup_jobs, list):
|
||||
report["error"] = "backed-up cron-jobs.json has no jobs list"
|
||||
return report
|
||||
|
||||
# Build a lookup of the backed-up skill state keyed by job id.
|
||||
# We only need the two skill-ish fields (legacy single and modern list).
|
||||
backup_by_id: Dict[str, Dict[str, Any]] = {}
|
||||
for job in backup_jobs:
|
||||
if not isinstance(job, dict):
|
||||
continue
|
||||
jid = job.get("id")
|
||||
if not isinstance(jid, str) or not jid:
|
||||
continue
|
||||
backup_by_id[jid] = {
|
||||
"skills": job.get("skills"),
|
||||
"skill": job.get("skill"),
|
||||
"name": job.get("name") or jid,
|
||||
}
|
||||
return {**report, "error": f"failed to load backed-up jobs: {e}"}
|
||||
if backup_jobs is None:
|
||||
return {**report, "error": "backed-up cron-jobs.json has no jobs list"}
|
||||
|
||||
# Backed-up skill state keyed by job id (legacy single + modern list field).
|
||||
backup_by_id: Dict[str, Dict[str, Any]] = {
|
||||
job["id"]: {"skills": job.get("skills"), "skill": job.get("skill"), "name": job.get("name") or job["id"]}
|
||||
for job in backup_jobs if isinstance(job, dict) and isinstance(job.get("id"), str) and job.get("id")
|
||||
}
|
||||
if not backup_by_id:
|
||||
report["attempted"] = True # we tried but there was nothing to do
|
||||
return report
|
||||
|
||||
# Load and rewrite the live jobs under the scheduler's cross-process lock.
|
||||
return {**report, "attempted": True} # we tried but there was nothing to do
|
||||
try:
|
||||
from cron.jobs import load_jobs, save_jobs, _jobs_lock
|
||||
except ImportError as e:
|
||||
report["error"] = f"cron module unavailable: {e}"
|
||||
return report
|
||||
return {**report, "error": f"cron module unavailable: {e}"}
|
||||
|
||||
report["attempted"] = True
|
||||
try:
|
||||
with _jobs_lock():
|
||||
live_jobs = load_jobs()
|
||||
changed = False
|
||||
|
||||
live_ids = set()
|
||||
changed, live_ids = False, set()
|
||||
for live in live_jobs:
|
||||
if not isinstance(live, dict):
|
||||
continue
|
||||
jid = live.get("id")
|
||||
jid = live.get("id") if isinstance(live, dict) else None
|
||||
if not isinstance(jid, str) or not jid:
|
||||
continue
|
||||
live_ids.add(jid)
|
||||
|
||||
backup = backup_by_id.get(jid)
|
||||
if backup is None:
|
||||
continue # live job didn't exist at snapshot time
|
||||
|
||||
cur_skills = live.get("skills")
|
||||
cur_skill = live.get("skill")
|
||||
bkp_skills = backup.get("skills")
|
||||
bkp_skill = backup.get("skill")
|
||||
|
||||
if cur_skills == bkp_skills and cur_skill == bkp_skill:
|
||||
if backup is None: # live job didn't exist at snapshot time
|
||||
continue
|
||||
cur = {"skills": live.get("skills"), "skill": live.get("skill")}
|
||||
bkp = {"skills": backup.get("skills"), "skill": backup.get("skill")}
|
||||
if cur == bkp:
|
||||
report["unchanged"] += 1
|
||||
continue
|
||||
|
||||
# Restore. Preserve absence (don't force the key to appear
|
||||
# if the backup didn't have it either).
|
||||
if bkp_skills is None:
|
||||
live.pop("skills", None)
|
||||
else:
|
||||
live["skills"] = bkp_skills
|
||||
if bkp_skill is None:
|
||||
live.pop("skill", None)
|
||||
else:
|
||||
live["skill"] = bkp_skill
|
||||
|
||||
report["restored"].append({
|
||||
"job_id": jid,
|
||||
"job_name": backup.get("name") or jid,
|
||||
"from": {"skills": cur_skills, "skill": cur_skill},
|
||||
"to": {"skills": bkp_skills, "skill": bkp_skill},
|
||||
})
|
||||
for key, value in bkp.items(): # Restore, preserving absence (don't add a key the backup lacked).
|
||||
if value is None:
|
||||
live.pop(key, None)
|
||||
else:
|
||||
live[key] = value
|
||||
report["restored"].append({"job_id": jid, "job_name": backup.get("name") or jid, "from": cur, "to": bkp})
|
||||
changed = True
|
||||
|
||||
# Jobs in backup but not in live = user deleted them after snapshot
|
||||
for jid, backup in backup_by_id.items():
|
||||
if jid not in live_ids:
|
||||
report["skipped_missing"].append({
|
||||
"job_id": jid,
|
||||
"job_name": backup.get("name") or jid,
|
||||
})
|
||||
|
||||
# Jobs in backup but not live = user deleted them after the snapshot.
|
||||
report["skipped_missing"] = [{"job_id": jid, "job_name": b.get("name") or jid}
|
||||
for jid, b in backup_by_id.items() if jid not in live_ids]
|
||||
if changed:
|
||||
save_jobs(live_jobs)
|
||||
except Exception as e: # noqa: BLE001 — rollback must not die mid-restore
|
||||
logger.debug("Cron skill-link restore failed: %s", e, exc_info=True)
|
||||
report["error"] = f"restore failed mid-flight: {e}"
|
||||
|
||||
return report
|
||||
|
||||
|
||||
def _remove_entry(entry: Path) -> None:
|
||||
if entry.is_dir() and not entry.is_symlink():
|
||||
shutil.rmtree(entry)
|
||||
elif entry.exists() or entry.is_symlink():
|
||||
entry.unlink()
|
||||
|
||||
|
||||
def _restore_excluded_subtrees(staged: Path, skills: Path) -> None:
|
||||
"""Move excluded entries (nested ``.git``/``.hub``/...) from *staged* back under *skills* after a successful extract.
|
||||
Snapshots never contain these, so the staged copy of the live tree is the only source. ``.git`` may be a dir or a file
|
||||
(submodule / worktree ``gitdir:`` pointer) — both are moved. Best-effort and conditional: an entry is carried only when
|
||||
its parent skill dir was restored and nothing sits at the target. If the target snapshot predates the skill, the entry
|
||||
is dropped with the staging dir rather than left orphaned; the safety snapshot excludes these paths too, so not undoable."""
|
||||
for dirpath, dirnames, filenames in os.walk(staged):
|
||||
for src in [Path(dirpath) / n for n in (*dirnames, *filenames) if n in _EXCLUDE_TOP_LEVEL]:
|
||||
dest = skills / src.relative_to(staged)
|
||||
if dest.parent.is_dir() and not dest.exists():
|
||||
try:
|
||||
shutil.move(str(src), str(dest))
|
||||
except OSError as e:
|
||||
logger.debug("Could not restore excluded entry %s: %s", src, e)
|
||||
dirnames[:] = [d for d in dirnames if d not in _EXCLUDE_TOP_LEVEL]
|
||||
|
||||
|
||||
def _unstage(moved: List[Tuple[Path, Path]]) -> List[str]:
|
||||
"""Move staged entries back to their original paths.
|
||||
|
||||
``shutil.move`` moves *into* an existing destination directory rather than
|
||||
replacing it, so a partially-completed extract leaves debris that would
|
||||
otherwise bury the user's real skill one level deeper
|
||||
(``skills/foo/foo/``) while the tree still looks populated. Clear whatever
|
||||
the failed extract created at each original path first. The staged copy is
|
||||
authoritative, and the pre-rollback safety snapshot is the undo handle for
|
||||
the extract's own output.
|
||||
|
||||
Returns the names that could not be restored, so the caller can report an
|
||||
incomplete recovery instead of claiming the state was restored.
|
||||
"""
|
||||
"""Move staged entries back to their original paths; returns names that could not be restored. ``shutil.move``
|
||||
moves *into* an existing destination dir, so partial-extract debris would bury the real skill
|
||||
(``skills/foo/foo/``) — clear each original path first. The staged copy is authoritative."""
|
||||
failed: List[str] = []
|
||||
for orig, dest in moved:
|
||||
try:
|
||||
if orig.is_dir() and not orig.is_symlink():
|
||||
shutil.rmtree(orig)
|
||||
elif orig.exists() or orig.is_symlink():
|
||||
orig.unlink()
|
||||
_remove_entry(orig)
|
||||
shutil.move(str(dest), str(orig))
|
||||
except OSError:
|
||||
failed.append(orig.name)
|
||||
return failed
|
||||
|
||||
|
||||
def _cron_summary(cron_report: Dict[str, Any]) -> Optional[str]:
|
||||
if not cron_report.get("attempted"):
|
||||
return None
|
||||
if cron_report.get("error"):
|
||||
return f"cron links: error — {cron_report['error']}"
|
||||
# (attempted with nothing matched — empty snapshot or no overlapping ids — says nothing)
|
||||
parts = [f"{n} {label}" for n, label in (
|
||||
(len(cron_report.get("restored") or []), "job(s) had skill links restored"),
|
||||
(len(cron_report.get("skipped_missing") or []), "backed-up job(s) no longer exist (skipped)"),
|
||||
(cron_report.get("unchanged", 0), "already matched"),
|
||||
) if n]
|
||||
return "cron links: " + ", ".join(parts) if parts else None
|
||||
|
||||
|
||||
def rollback(backup_id: Optional[str] = None) -> Tuple[bool, str, Optional[Path]]:
|
||||
"""Restore ``~/.hermes/skills/`` from a snapshot.
|
||||
|
||||
Strategy:
|
||||
1. Resolve the target snapshot (explicit id or newest regular).
|
||||
2. Take a safety snapshot of the CURRENT skills tree under
|
||||
``.curator_backups/pre-rollback-<ts>/`` so the rollback itself is
|
||||
undoable.
|
||||
3. Move all current top-level entries (except ``.curator_backups``
|
||||
and ``.hub``) into a tempdir.
|
||||
4. Extract the chosen snapshot into ``~/.hermes/skills/``.
|
||||
5. On failure during 4, move the tempdir contents back (best-effort)
|
||||
and return failure.
|
||||
|
||||
Returns ``(ok, message, snapshot_path)``.
|
||||
"""
|
||||
"""Restore ``~/.hermes/skills/`` from a snapshot (explicit id or newest): safety-snapshot the CURRENT tree; stage
|
||||
current top-level entries; extract; on failure move staged entries back. Returns ``(ok, message, snapshot_path)``."""
|
||||
target = _resolve_backup(backup_id)
|
||||
if target is None:
|
||||
return (
|
||||
False,
|
||||
"no matching backup found"
|
||||
+ (f" for id '{backup_id}'" if backup_id else "")
|
||||
+ " (use `hermes curator rollback --list` to see available snapshots)",
|
||||
None,
|
||||
)
|
||||
archive = target / "skills.tar.gz"
|
||||
return (False, "no matching backup found" + (f" for id '{backup_id}'" if backup_id else "")
|
||||
+ " (use `hermes curator rollback --list` to see available snapshots)", None)
|
||||
archive = target / _ARCHIVE_NAME
|
||||
if not archive.exists():
|
||||
return (False, f"snapshot {target.name} has no skills.tar.gz — corrupted?", None)
|
||||
|
||||
skills = _skills_dir()
|
||||
skills.mkdir(parents=True, exist_ok=True)
|
||||
backups = _backups_dir()
|
||||
backups.mkdir(parents=True, exist_ok=True)
|
||||
skills, backups = _skills_dir(), _backups_dir()
|
||||
backups.mkdir(parents=True, exist_ok=True) # parents=True also creates skills/
|
||||
|
||||
# Step 2: safety snapshot of current state FIRST. If this fails we bail
|
||||
# out before touching anything — otherwise a failed extract could leave
|
||||
# the user with no skills.
|
||||
# Safety snapshot FIRST; bail if it fails, else a failed extract could leave the user with no skills. Protect the target from its prune.
|
||||
try:
|
||||
# Protect the target from this snapshot's prune step: at the steady
|
||||
# keep limit, pruning the oldest snapshot would otherwise delete the
|
||||
# very snapshot we are about to extract from.
|
||||
safety_snapshot = snapshot_skills(
|
||||
reason=f"pre-rollback to {target.name}",
|
||||
protect_ids={target.name},
|
||||
)
|
||||
safety_snapshot = snapshot_skills(reason=f"pre-rollback to {target.name}", protect_ids={target.name})
|
||||
except Exception as e:
|
||||
return (False, f"pre-rollback safety snapshot failed: {e}", None)
|
||||
if safety_snapshot is None:
|
||||
return (
|
||||
False,
|
||||
"pre-rollback safety snapshot failed; backups may be disabled "
|
||||
"or unavailable, and current skills were not changed",
|
||||
None,
|
||||
)
|
||||
return (False, "pre-rollback safety snapshot failed; backups may be disabled "
|
||||
"or unavailable, and current skills were not changed", None)
|
||||
|
||||
# Additionally move current entries into an internal staging dir so
|
||||
# the extract happens into an empty skills tree (predictable result).
|
||||
# This dir is implementation detail — not listed as a restorable
|
||||
# backup. The safety snapshot above is the user-facing undo handle.
|
||||
staged = backups / f".rollback-staging-{_utc_id()}"
|
||||
# Stage current entries so the extract lands in an empty tree; the safety snapshot above (not staging) is the user-facing undo handle.
|
||||
staged = backups / f"{_STAGING_PREFIX}{_utc_id()}"
|
||||
try:
|
||||
staged.mkdir(parents=True, exist_ok=False)
|
||||
except OSError as e:
|
||||
@@ -637,122 +388,54 @@ def rollback(backup_id: Optional[str] = None) -> Tuple[bool, str, Optional[Path]
|
||||
moved: List[Tuple[Path, Path]] = []
|
||||
try:
|
||||
for entry in list(skills.iterdir()):
|
||||
if entry.name in _EXCLUDE_TOP_LEVEL:
|
||||
continue
|
||||
dest = staged / entry.name
|
||||
shutil.move(str(entry), str(dest))
|
||||
moved.append((entry, dest))
|
||||
if entry.name not in _EXCLUDE_TOP_LEVEL:
|
||||
shutil.move(str(entry), str(staged / entry.name))
|
||||
moved.append((entry, staged / entry.name))
|
||||
except OSError as e:
|
||||
# Best-effort rollback of the move
|
||||
_unstage(moved)
|
||||
try:
|
||||
shutil.rmtree(staged, ignore_errors=True)
|
||||
except OSError:
|
||||
pass
|
||||
shutil.rmtree(staged, ignore_errors=True)
|
||||
return (False, f"failed to stage current skills: {e}", None)
|
||||
|
||||
# Step 4: extract the snapshot into skills/
|
||||
try:
|
||||
with tarfile.open(archive, "r:gz") as tf:
|
||||
# Python 3.12+ supports filter='data' for safer extraction.
|
||||
# Fall back to the unfiltered call for older interpreters but
|
||||
# still reject absolute paths and .. components defensively.
|
||||
# Reject absolute paths and ".." defensively; Python 3.12+ also gets filter='data', older interpreters fall back unfiltered.
|
||||
for member in tf.getmembers():
|
||||
name = member.name
|
||||
if name.startswith("/") or ".." in Path(name).parts:
|
||||
raise tarfile.TarError(
|
||||
f"refusing to extract unsafe path: {name!r}"
|
||||
)
|
||||
if member.name.startswith("/") or ".." in Path(member.name).parts:
|
||||
raise tarfile.TarError(f"refusing to extract unsafe path: {member.name!r}")
|
||||
try:
|
||||
tf.extractall(str(skills), filter="data") # type: ignore[call-arg]
|
||||
except TypeError:
|
||||
# Python < 3.12 — no filter kwarg
|
||||
tf.extractall(str(skills))
|
||||
tf.extractall(str(skills)) # Python < 3.12 — no filter kwarg
|
||||
except (OSError, tarfile.TarError) as e:
|
||||
# Best-effort recover. A partial extract can leave entries the
|
||||
# original tree never had, so drop those first, otherwise the
|
||||
# "restored" tree is the user's skills plus a slice of the snapshot.
|
||||
staged_names = {orig.name for orig, _ in moved}
|
||||
for entry in list(skills.iterdir()):
|
||||
if entry.name in _EXCLUDE_TOP_LEVEL or entry.name in staged_names:
|
||||
continue
|
||||
try:
|
||||
if entry.is_dir() and not entry.is_symlink():
|
||||
shutil.rmtree(entry)
|
||||
else:
|
||||
entry.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
# A partial extract can leave entries the original tree never had; drop those first or the "restored" tree is skills + a slice of snapshot.
|
||||
keep = _EXCLUDE_TOP_LEVEL | {orig.name for orig, _ in moved}
|
||||
for entry in [e for e in skills.iterdir() if e.name not in keep]:
|
||||
with contextlib.suppress(OSError):
|
||||
_remove_entry(entry)
|
||||
unrestored = _unstage(moved)
|
||||
if unrestored:
|
||||
# Do not claim a clean restore we did not achieve, and keep the
|
||||
# staging dir so the entries can be recovered by hand.
|
||||
return (
|
||||
False,
|
||||
f"snapshot extract failed: {e} - could not restore "
|
||||
f"{', '.join(sorted(unrestored))}; staged copies kept at {staged}",
|
||||
None,
|
||||
)
|
||||
try:
|
||||
shutil.rmtree(staged, ignore_errors=True)
|
||||
except OSError:
|
||||
pass
|
||||
if unrestored: # Don't claim a clean restore; keep the staging dir for hand recovery.
|
||||
return (False, f"snapshot extract failed: {e} - could not restore "
|
||||
f"{', '.join(sorted(unrestored))}; staged copies kept at {staged}", None)
|
||||
shutil.rmtree(staged, ignore_errors=True)
|
||||
return (False, f"snapshot extract failed (state restored): {e}", None)
|
||||
|
||||
# Extract succeeded — the staging dir has served its purpose. The
|
||||
# user's undo handle is the safety snapshot tarball we took earlier.
|
||||
try:
|
||||
shutil.rmtree(staged, ignore_errors=True)
|
||||
except OSError:
|
||||
pass
|
||||
# Snapshots never contain excluded subtrees (nested ``.git``, ``.hub``, ...), so carry them over from the staged live tree
|
||||
# (top-level ``.git`` is never staged). Then staging is done; the undo handle is the safety snapshot.
|
||||
_restore_excluded_subtrees(staged, skills)
|
||||
shutil.rmtree(staged, ignore_errors=True)
|
||||
|
||||
# Reconcile cron skill-links. Surgical: only the skills/skill fields
|
||||
# on jobs matched by id. Everything else in jobs.json is live state
|
||||
# (schedule, next_run_at, enabled, prompt, etc.) and we leave it
|
||||
# alone. Failures here don't fail the overall rollback — the skills
|
||||
# tree is already restored, which is the main guarantee.
|
||||
# Cron reconciliation failures don't fail the rollback — the skills tree (the main guarantee) is already restored.
|
||||
cron_report = _restore_cron_skill_links(target)
|
||||
|
||||
summary_bits = [f"restored from snapshot {target.name}"]
|
||||
if cron_report.get("attempted"):
|
||||
restored_n = len(cron_report.get("restored") or [])
|
||||
skipped_n = len(cron_report.get("skipped_missing") or [])
|
||||
if cron_report.get("error"):
|
||||
summary_bits.append(f"cron links: error — {cron_report['error']}")
|
||||
elif restored_n == 0 and skipped_n == 0 and cron_report.get("unchanged", 0) == 0:
|
||||
# Attempted but nothing matched — empty snapshot or no overlapping ids.
|
||||
pass
|
||||
else:
|
||||
parts = []
|
||||
if restored_n:
|
||||
parts.append(f"{restored_n} job(s) had skill links restored")
|
||||
if skipped_n:
|
||||
parts.append(f"{skipped_n} backed-up job(s) no longer exist (skipped)")
|
||||
if cron_report.get("unchanged"):
|
||||
parts.append(f"{cron_report['unchanged']} already matched")
|
||||
summary_bits.append("cron links: " + ", ".join(parts))
|
||||
|
||||
logger.info("Curator rollback: restored from %s (cron_report=%s)",
|
||||
target.name, cron_report)
|
||||
return (True, "; ".join(summary_bits), target)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Human-readable summary for CLI
|
||||
# ---------------------------------------------------------------------------
|
||||
logger.info("Curator rollback: restored from %s (cron_report=%s)", target.name, cron_report)
|
||||
return (True, "; ".join(filter(None, [f"restored from snapshot {target.name}", _cron_summary(cron_report)])), target)
|
||||
|
||||
|
||||
# --- Human-readable summary for CLI ---
|
||||
def summarize_backups() -> str:
|
||||
rows = list_backups()
|
||||
if not rows:
|
||||
return "No curator snapshots yet."
|
||||
lines = [f"{'id':<24} {'reason':<40} {'skills':>6} {'size':>8}"]
|
||||
lines.append("─" * len(lines[0]))
|
||||
for r in rows:
|
||||
lines.append(
|
||||
f"{r.get('id','?'):<24} "
|
||||
f"{(r.get('reason','?') or '?')[:40]:<40} "
|
||||
f"{r.get('skill_files', 0):>6} "
|
||||
f"{format_bytes(int(r.get('archive_bytes', 0))):>8}"
|
||||
)
|
||||
return "\n".join(lines)
|
||||
header = f"{'id':<24} {'reason':<40} {'skills':>6} {'size':>8}"
|
||||
return "\n".join([header, "─" * len(header)] + [
|
||||
f"{r.get('id','?'):<24} {(r.get('reason','?') or '?')[:40]:<40} "
|
||||
f"{r.get('skill_files', 0):>6} {format_bytes(int(r.get('archive_bytes', 0))):>8}" for r in rows])
|
||||
|
||||
+211
-401
@@ -1,68 +1,22 @@
|
||||
"""Unified deadline layer — one bounded-execution primitive, one timeout resolver.
|
||||
"""Unified deadline layer — one bounded-execution primitive, one timeout resolver (#85125).
|
||||
|
||||
Phase 1 of the architectural fix for the timeout/hang backlog
|
||||
(https://github.com/NousResearch/hermes-agent/issues/85125).
|
||||
* :func:`resolve_timeout` — ``timeouts:`` in config.yaml > legacy env var > default.
|
||||
* :func:`clamp_timeout` — huge timeouts overflow ``time_t`` in ``Lock.acquire`` /
|
||||
``Thread.join`` on macOS (#83220), so every timeout is capped.
|
||||
* :func:`run_bounded_async` / :func:`run_bounded_sync` — wall-clock deadlines driven by
|
||||
a daemon ``threading.Timer`` / worker thread, so a blocked event loop cannot disable them.
|
||||
* :func:`kill_process_tree` — portable whole-tree termination.
|
||||
|
||||
The tree currently carries at least six site-local deadline mechanisms, each
|
||||
built for one incident, none shared (tool_executor batch deadline, telegram
|
||||
``_await_with_thread_deadline``, gateway turn lease, reasoning stale floors,
|
||||
``human_wait_ceiling``, per-MCP-handler timeouts). Every new stall report
|
||||
grows that list by one. This module is the shared foundation the call sites
|
||||
migrate onto in later phases:
|
||||
|
||||
* :func:`resolve_timeout` — one config-first resolution path for timeout
|
||||
values (``timeouts:`` section in config.yaml > legacy env var > default),
|
||||
so new surfaces stop inventing ``HERMES_*_TIMEOUT`` env vars (".env is for
|
||||
secrets only") and hardcoded literals stop ignoring user config
|
||||
(#63302, #53161, #43272 class).
|
||||
|
||||
* :func:`clamp_timeout` — platform-safe clamping. Large user-supplied
|
||||
timeouts overflow ``time_t`` inside ``threading.Lock.acquire(timeout=...)``
|
||||
/ ``Thread.join(timeout=...)`` on macOS and kill whole tool batches
|
||||
(#83220). Clamping at the shared boundary fixes that class once, for
|
||||
every consumer.
|
||||
|
||||
* :func:`run_bounded_async` — a wall-clock deadline for awaitables that does
|
||||
NOT depend on event-loop timers. ``asyncio.wait_for`` schedules its expiry
|
||||
on the loop; when the loop thread itself is blocked in a synchronous call
|
||||
(family A of the #84047 stall triage), every asyncio-based timeout in the
|
||||
process is silently disabled. This helper drives the deadline from a
|
||||
daemon ``threading.Timer`` (generalizing the proven telegram-adapter
|
||||
primitive) and abandons cancellation-shielded tasks instead of waiting for
|
||||
cancellation to complete. The telegram adapter's private copy
|
||||
(``plugins/platforms/telegram/adapter.py:_await_with_thread_deadline``)
|
||||
migrates onto this in Phase 2 of #85125 — do not let the two drift in the
|
||||
meantime; fix bugs here first.
|
||||
|
||||
* :func:`run_bounded_sync` — the same contract for synchronous callables
|
||||
bounded from a synchronous context (daemon worker thread, abandoned on
|
||||
expiry).
|
||||
|
||||
* :func:`kill_process_tree` — portable whole-tree termination so
|
||||
kill-on-timeout stops orphaning descendants (#71148, #59549, #84967,
|
||||
#68139 class). Existing site-local tree-kills that migrate onto this in
|
||||
Phase 4 of #85125: ``gateway/status.py`` (taskkill wrapper + psutil
|
||||
snapshot/reap pair) and ``tools/code_execution_tool.py`` (psutil
|
||||
recursive children kill).
|
||||
|
||||
Design invariants:
|
||||
|
||||
* Exceptions raised by the bounded operation propagate unchanged — callers
|
||||
keep their existing error handling. Only the *timeout* outcome is
|
||||
reified (as :class:`BoundedResult`), because that is the outcome the
|
||||
call sites keep getting wrong.
|
||||
* A timeout produced by this layer is OUR deadline, not the provider's.
|
||||
Callers that feed errors into ``agent/error_classifier.py`` should
|
||||
classify :class:`DeadlineExpired` distinctly from transport timeouts
|
||||
(the #59549 / #80323 misattribution class).
|
||||
* ``None`` timeout means unbounded, and non-positive resolved values are
|
||||
normalized to ``None`` (matching the existing
|
||||
``HERMES_CONCURRENT_TOOL_TIMEOUT_S`` convention).
|
||||
Invariants: operation exceptions propagate unchanged (only the *timeout* outcome is reified
|
||||
as :class:`BoundedResult`); a timeout here is OUR deadline, not the provider's (classify
|
||||
:class:`DeadlineExpired` distinctly from transport timeouts); ``None`` / non-positive means unbounded.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
from contextlib import contextmanager
|
||||
import faulthandler
|
||||
import logging
|
||||
import os
|
||||
@@ -76,39 +30,25 @@ from typing import Any, Awaitable, Callable, Optional, Protocol
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
__all__ = [
|
||||
"MAX_SAFE_TIMEOUT_S",
|
||||
"BoundedResult",
|
||||
"DeadlineExpired",
|
||||
"clamp_timeout",
|
||||
"resolve_timeout",
|
||||
"run_bounded_async",
|
||||
"run_bounded_sync",
|
||||
"kill_process_tree",
|
||||
"MAX_SAFE_TIMEOUT_S", "BoundedResult", "DeadlineExpired", "clamp_timeout", "resolve_timeout",
|
||||
"run_bounded_async", "run_bounded_sync", "kill_process_tree",
|
||||
]
|
||||
|
||||
# Upper bound for any timeout handed to platform wait primitives.
|
||||
#
|
||||
# CPython converts ``threading.Lock.acquire(timeout=...)`` /
|
||||
# ``Thread.join(timeout=...)`` deadlines to an absolute timestamp; very large
|
||||
# relative timeouts overflow ``time_t`` on macOS and raise
|
||||
# ``OverflowError: timestamp out of range for platform time_t`` (#83220).
|
||||
# One year is semantically "unbounded" for every wait in this codebase while
|
||||
# staying far below any platform conversion limit.
|
||||
MAX_SAFE_TIMEOUT_S = 31_536_000.0 # 365 days
|
||||
# One year: semantically "unbounded" yet far below any platform time_t limit (#83220).
|
||||
MAX_SAFE_TIMEOUT_S = 31_536_000.0
|
||||
|
||||
# Grace period after a deadline fires before concluding the event loop thread
|
||||
# is blocked in a synchronous call and dumping stacks (family A diagnostics).
|
||||
# Grace after a deadline fires before concluding the loop thread is blocked and dumping stacks.
|
||||
_LOOP_BLOCKED_DUMP_GRACE_S = 5.0
|
||||
|
||||
# ``Event.wait`` is a C-level block: KeyboardInterrupt / SetAsyncExc only land when the
|
||||
# thread returns to Python, so the sync wait is sliced to observe /stop or SIGINT promptly.
|
||||
# Slice the wait so a /stop or SIGINT during a bounded sync call is observed within this window rather than
|
||||
# at the full deadline (#94285, tools/test_local_interrupt_cleanup).
|
||||
_BOUNDED_SYNC_WAIT_SLICE_S = 0.2
|
||||
|
||||
|
||||
class DeadlineExpired(TimeoutError):
|
||||
"""A deadline enforced by this layer expired.
|
||||
|
||||
Distinct from transport/provider timeout types on purpose: when this is
|
||||
raised (or a :class:`BoundedResult` reports ``timed_out``), the timeout
|
||||
was Hermes's own bound — error classification must not attribute it to
|
||||
the provider (#59549 / #80323 misattribution class).
|
||||
"""
|
||||
"""A deadline enforced by this layer expired (Hermes's own bound, not the provider's)."""
|
||||
|
||||
def __init__(self, label: str, timeout_s: float):
|
||||
super().__init__(f"deadline expired after {timeout_s:.1f}s: {label}")
|
||||
@@ -117,22 +57,10 @@ class DeadlineExpired(TimeoutError):
|
||||
|
||||
|
||||
class SuspectableBackend(Protocol):
|
||||
"""Phase 3a (#85125): a stateful backend the deadline layer can flag.
|
||||
|
||||
A timed-out stateful backend (MCP connection, browser session, LSP
|
||||
client) may be left wedged by the abandoned half-finished operation.
|
||||
``run_bounded_*`` calls ``mark_suspect`` on timeout so the OWNER can
|
||||
health-check or recycle the backend before reuse (``ensure_healthy``)
|
||||
instead of returning a poisoned handle to the cache. Consumers adopt
|
||||
incrementally (Phase 3b, one backend per PR), so the layer fails open:
|
||||
backends without the protocol are simply never marked.
|
||||
|
||||
Adopter contract: ``mark_suspect`` MUST be cheap, non-blocking, and
|
||||
must not acquire locks the guarded operation may hold. It runs inline —
|
||||
on the event loop in the async flavor, and on the caller's thread in
|
||||
the sync flavor while the wedged worker is still alive. Set a flag;
|
||||
do the expensive health-check/recycle work in ``ensure_healthy``.
|
||||
"""
|
||||
"""A stateful backend (MCP connection, browser session, LSP client) ``run_bounded_*`` flags via
|
||||
``mark_suspect`` on timeout so the owner can health-check/recycle it before reuse. ``mark_suspect``
|
||||
MUST be cheap, non-blocking, and must not acquire locks the guarded operation may hold — it runs
|
||||
inline on the event loop / caller's thread while the wedged worker is still alive."""
|
||||
|
||||
def mark_suspect(self, reason: str) -> None: ...
|
||||
|
||||
@@ -140,13 +68,7 @@ class SuspectableBackend(Protocol):
|
||||
|
||||
|
||||
def _mark_backend_suspect(backend: object | None, label: str, timeout_s: float) -> None:
|
||||
"""Best-effort ``mark_suspect`` on a timed-out call's backend.
|
||||
|
||||
Never raises: adoption state must not be able to weaken the deadline
|
||||
bound or corrupt the ``BoundedResult`` the caller is about to receive.
|
||||
A non-adopting backend (no ``mark_suspect``) is tolerated silently —
|
||||
Phase 3b lands per-backend, so absence is the norm during adoption.
|
||||
"""
|
||||
"""Best-effort ``mark_suspect``; never raises, non-adopting backends tolerated."""
|
||||
if backend is None:
|
||||
return
|
||||
try:
|
||||
@@ -159,12 +81,7 @@ def _mark_backend_suspect(backend: object | None, label: str, timeout_s: float)
|
||||
|
||||
@dataclass(frozen=True, kw_only=True)
|
||||
class BoundedResult:
|
||||
"""Outcome of a bounded operation.
|
||||
|
||||
``timed_out`` is the reified outcome; on completion ``value`` holds the
|
||||
operation's return value. Operation exceptions are never captured here —
|
||||
they propagate to the caller unchanged.
|
||||
"""
|
||||
"""Outcome of a bounded operation; operation exceptions are never captured here."""
|
||||
|
||||
timed_out: bool
|
||||
value: Any
|
||||
@@ -172,57 +89,35 @@ class BoundedResult:
|
||||
timeout_s: Optional[float]
|
||||
label: str
|
||||
|
||||
def raise_if_timed_out(self) -> Any:
|
||||
"""Return ``value``, raising :class:`DeadlineExpired` on timeout."""
|
||||
if self.timed_out:
|
||||
raise DeadlineExpired(self.label, float(self.timeout_s or 0.0))
|
||||
return self.value
|
||||
|
||||
def _result(start: float, timeout_s: Optional[float], label: str, *, value: Any = None, timed_out: bool = False) -> BoundedResult:
|
||||
return BoundedResult(
|
||||
timed_out=timed_out, value=value, elapsed_s=time.monotonic() - start, timeout_s=timeout_s, label=label
|
||||
)
|
||||
|
||||
|
||||
def clamp_timeout(timeout: Optional[float]) -> Optional[float]:
|
||||
"""Normalize a timeout value for platform wait primitives.
|
||||
|
||||
* ``None`` stays ``None`` (unbounded).
|
||||
* Non-positive values become ``None`` (unbounded) — matching the existing
|
||||
``HERMES_CONCURRENT_TOOL_TIMEOUT_S`` "0 disables the bound" convention.
|
||||
* Values above :data:`MAX_SAFE_TIMEOUT_S` are capped so they can never
|
||||
overflow ``time_t`` inside ``Lock.acquire`` / ``Thread.join`` on macOS
|
||||
(#83220).
|
||||
* Non-numeric values are treated as unset (``None``) with a warning
|
||||
rather than crashing the call path they were meant to protect.
|
||||
"""
|
||||
"""Normalize a timeout: None/non-positive/non-numeric/NaN -> None (unbounded), else capped."""
|
||||
if timeout is None:
|
||||
return None
|
||||
try:
|
||||
value = float(timeout)
|
||||
except (TypeError, ValueError):
|
||||
logger.warning(
|
||||
"clamp_timeout: non-numeric timeout %r; treating as unbounded", timeout
|
||||
)
|
||||
logger.warning("clamp_timeout: non-numeric timeout %r; treating as unbounded", timeout)
|
||||
return None
|
||||
if value != value: # NaN
|
||||
logger.warning("clamp_timeout: NaN timeout; treating as unbounded")
|
||||
return None
|
||||
if value <= 0:
|
||||
return None
|
||||
return min(value, MAX_SAFE_TIMEOUT_S)
|
||||
return None if value <= 0 else min(value, MAX_SAFE_TIMEOUT_S)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Timeout resolution: config.yaml ``timeouts:`` section > legacy env var >
|
||||
# registered default.
|
||||
# ---------------------------------------------------------------------------
|
||||
# --- Timeout resolution: config ``timeouts:`` > legacy env var > default ------
|
||||
|
||||
|
||||
def _timeouts_section() -> dict:
|
||||
"""Read the ``timeouts:`` root section from config.yaml (read-only).
|
||||
|
||||
Isolated for testability and so a broken config read can never take down
|
||||
the call path the timeout was protecting.
|
||||
"""
|
||||
"""Read the ``timeouts:`` root section from config.yaml (read-only, fail-open)."""
|
||||
try:
|
||||
from hermes_cli.config import load_config_readonly
|
||||
|
||||
section = load_config_readonly().get("timeouts")
|
||||
return section if isinstance(section, dict) else {}
|
||||
except Exception:
|
||||
@@ -240,36 +135,14 @@ def _lookup_dotted(section: dict, key: str) -> Any:
|
||||
return node
|
||||
|
||||
|
||||
def resolve_timeout(
|
||||
key: str,
|
||||
*,
|
||||
default: Optional[float],
|
||||
env_var: Optional[str] = None,
|
||||
) -> Optional[float]:
|
||||
"""Resolve a timeout in seconds for a dotted config key.
|
||||
|
||||
Precedence (established by the ``providers.*.request_timeout_seconds``
|
||||
pattern — config wins over the legacy env var):
|
||||
|
||||
1. ``timeouts.<key>`` in config.yaml (dotted key walks nested maps, e.g.
|
||||
``tools.concurrent_batch`` reads ``timeouts: {tools: {concurrent_batch: ...}}``)
|
||||
2. ``env_var`` when set and non-empty (legacy bridge — internal mechanism
|
||||
and back-compat only; new surfaces must not grow new user-facing
|
||||
``HERMES_*`` timeout env vars)
|
||||
3. ``default``
|
||||
|
||||
The winning value is passed through :func:`clamp_timeout`, so ``0`` or a
|
||||
negative value means "unbounded" and oversized values are made
|
||||
platform-safe. Invalid (non-numeric) config/env values fall through to
|
||||
the next source with a warning instead of breaking the protected path.
|
||||
"""
|
||||
def resolve_timeout(key: str, *, default: Optional[float], env_var: Optional[str] = None) -> Optional[float]:
|
||||
"""Resolve a timeout (seconds): dotted ``timeouts.<key>`` > ``env_var`` > ``default``; the winner
|
||||
goes through :func:`clamp_timeout`, invalid config/env values fall through with a warning."""
|
||||
raw = _lookup_dotted(_timeouts_section(), key)
|
||||
if raw is not None:
|
||||
# Explicit float() (clamp_timeout would also convert) so that invalid
|
||||
# config values FALL THROUGH to the env var / default instead of
|
||||
# resolving as unbounded — do not "simplify" this away. bool is
|
||||
# rejected because YAML `true` would silently become a 1-second
|
||||
# deadline; NaN is rejected for the same fall-through reason.
|
||||
# Explicit float() so invalid config values FALL THROUGH to env/default instead of
|
||||
# resolving as unbounded. bool rejected (YAML `true` would become a 1s deadline);
|
||||
# NaN rejected for the same fall-through reason.
|
||||
if not isinstance(raw, bool):
|
||||
try:
|
||||
value = float(raw)
|
||||
@@ -277,9 +150,7 @@ def resolve_timeout(
|
||||
return clamp_timeout(value)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
logger.warning(
|
||||
"timeouts.%s: invalid value %r in config.yaml; ignoring", key, raw
|
||||
)
|
||||
logger.warning("timeouts.%s: invalid value %r in config.yaml; ignoring", key, raw)
|
||||
|
||||
if env_var:
|
||||
env_raw = os.getenv(env_var, "").strip()
|
||||
@@ -292,17 +163,17 @@ def resolve_timeout(
|
||||
return clamp_timeout(default)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bounded execution — async flavor.
|
||||
#
|
||||
# Generalizes plugins/platforms/telegram/adapter.py:_await_with_thread_deadline
|
||||
# (the #63309 fix): the deadline is driven by a daemon threading.Timer so a
|
||||
# blocked event loop cannot disable it, and a second timer dumps all thread
|
||||
# stacks when the loop provably failed to process the expiry — the one piece
|
||||
# of information loop-blocked hangs otherwise never surface.
|
||||
# ---------------------------------------------------------------------------
|
||||
# --- Bounded execution — async flavor ------------------------------------------
|
||||
# The deadline is a daemon threading.Timer so a blocked event loop cannot disable it; a
|
||||
# second timer dumps all thread stacks when the loop provably failed to process the expiry.
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- Bounded execution — async
|
||||
# flavor. Generalizes plugins/platforms/telegram/adapter.py:_await_with_thread_deadline (the #63309 fix):
|
||||
# the deadline is driven by a daemon threading.Timer so a blocked event loop cannot disable it, and a second
|
||||
# timer dumps all thread stacks when the loop provably failed to process the expiry — the one piece of
|
||||
# information loop-blocked hangs otherwise never surface.
|
||||
# ---------------------------------------------------------------------------
|
||||
def _consume_abandoned(task: "asyncio.Future[Any]") -> None:
|
||||
"""Observe an abandoned task's outcome so it never logs 'never retrieved'."""
|
||||
try:
|
||||
@@ -312,8 +183,14 @@ def _consume_abandoned(task: "asyncio.Future[Any]") -> None:
|
||||
pass
|
||||
|
||||
|
||||
def _abandon(task: "asyncio.Future[Any]") -> None:
|
||||
"""Cancel ``task`` and never await it; its outcome is consumed so it stays unobserved-safe."""
|
||||
task.cancel()
|
||||
task.add_done_callback(_consume_abandoned)
|
||||
|
||||
|
||||
async def _run_abandon_cleanup(on_abandon: Callable[[], Awaitable[Any]]) -> None:
|
||||
"""Run abandonment cleanup fully fire-and-forget (its failures swallowed)."""
|
||||
"""Run abandonment cleanup fire-and-forget (its failures swallowed)."""
|
||||
try:
|
||||
await on_abandon()
|
||||
except Exception:
|
||||
@@ -322,14 +199,11 @@ async def _run_abandon_cleanup(on_abandon: Callable[[], Awaitable[Any]]) -> None
|
||||
|
||||
def _dump_blocked_loop_diagnostics(label: str, timeout_s: float) -> None:
|
||||
logger.warning(
|
||||
"[deadline] %r deadline (%.0fs) expired but the event loop has not "
|
||||
"processed the expiry after a further %.0fs — the loop thread appears "
|
||||
"BLOCKED in a synchronous call, which is why no asyncio timeout can "
|
||||
"fire. Dumping all thread stacks to stderr to identify the blocking "
|
||||
"frame.",
|
||||
label,
|
||||
timeout_s,
|
||||
_LOOP_BLOCKED_DUMP_GRACE_S,
|
||||
"[deadline] %r deadline (%.0fs) expired but the event loop has not processed the expiry "
|
||||
"after a further %.0fs — the loop thread appears BLOCKED in a synchronous call, which is "
|
||||
"why no asyncio timeout can fire. Dumping all thread stacks to stderr to identify the "
|
||||
"blocking frame.",
|
||||
label, timeout_s, _LOOP_BLOCKED_DUMP_GRACE_S,
|
||||
)
|
||||
try:
|
||||
faulthandler.dump_traceback(all_threads=True)
|
||||
@@ -348,30 +222,13 @@ async def run_bounded_async(
|
||||
) -> BoundedResult:
|
||||
"""Await ``awaitable`` under a wall-clock deadline independent of loop timers.
|
||||
|
||||
On completion returns ``BoundedResult(timed_out=False, value=...)``;
|
||||
exceptions from the operation (including ``asyncio.CancelledError`` from a
|
||||
caller cancelling *us*) propagate unchanged.
|
||||
|
||||
On timeout the underlying task is cancelled and **abandoned** — we do not
|
||||
await cancellation completion, because cancellation-shielded scopes (anyio,
|
||||
httpcore init, MCP SDK teardown) are exactly the paths that wedge forever.
|
||||
``on_abandon`` (zero-arg callable returning an awaitable) is scheduled as
|
||||
detached best-effort cleanup for the half-built state the abandoned task
|
||||
may leave behind. Returns ``BoundedResult(timed_out=True, value=None)``.
|
||||
|
||||
``timeout=None`` (or a non-positive resolved value) awaits unbounded.
|
||||
"""
|
||||
Operation exceptions (incl. ``CancelledError`` from a caller cancelling *us*) propagate
|
||||
unchanged. On timeout the task is cancelled and **abandoned** (never awaited —
|
||||
cancellation-shielded scopes are exactly the paths that wedge); ``on_abandon`` runs detached."""
|
||||
timeout_s = clamp_timeout(timeout)
|
||||
start = time.monotonic()
|
||||
if timeout_s is None:
|
||||
value = await awaitable
|
||||
return BoundedResult(
|
||||
timed_out=False,
|
||||
value=value,
|
||||
elapsed_s=time.monotonic() - start,
|
||||
timeout_s=None,
|
||||
label=label,
|
||||
)
|
||||
return _result(start, None, label, value=await awaitable)
|
||||
|
||||
task = asyncio.ensure_future(awaitable)
|
||||
loop = asyncio.get_running_loop()
|
||||
@@ -383,84 +240,44 @@ async def run_bounded_async(
|
||||
if not deadline.done():
|
||||
deadline.set_result(None)
|
||||
|
||||
def _expire_from_thread() -> None:
|
||||
loop.call_soon_threadsafe(_mark_expired)
|
||||
|
||||
def _watchdog_check() -> None:
|
||||
if not loop_processed_expiry.is_set():
|
||||
_dump_blocked_loop_diagnostics(label, timeout_s)
|
||||
|
||||
timer = threading.Timer(timeout_s, _expire_from_thread)
|
||||
timer.daemon = True
|
||||
timer.start()
|
||||
watchdog: Optional[threading.Timer] = None
|
||||
timers = [threading.Timer(timeout_s, lambda: loop.call_soon_threadsafe(_mark_expired))]
|
||||
if dump_on_blocked_loop:
|
||||
watchdog = threading.Timer(
|
||||
timeout_s + _LOOP_BLOCKED_DUMP_GRACE_S, _watchdog_check
|
||||
)
|
||||
watchdog.daemon = True
|
||||
watchdog.start()
|
||||
timers.append(threading.Timer(timeout_s + _LOOP_BLOCKED_DUMP_GRACE_S, _watchdog_check))
|
||||
for t in timers:
|
||||
t.daemon = True
|
||||
t.start()
|
||||
try:
|
||||
try:
|
||||
done, _ = await asyncio.wait(
|
||||
{task, deadline}, return_when=asyncio.FIRST_COMPLETED
|
||||
)
|
||||
done, _ = await asyncio.wait({task, deadline}, return_when=asyncio.FIRST_COMPLETED)
|
||||
except asyncio.CancelledError:
|
||||
# The CALLER cancelled us. Without this, `task` would keep running
|
||||
# unobserved (and later log "exception was never retrieved") —
|
||||
# a leak the telegram original also had. Cancel + abandon it, then
|
||||
# let the cancellation propagate.
|
||||
task.cancel()
|
||||
task.add_done_callback(_consume_abandoned)
|
||||
_abandon(task) # the CALLER cancelled us; `task` must not run unobserved
|
||||
raise
|
||||
if task in done:
|
||||
if not deadline.done():
|
||||
deadline.cancel()
|
||||
value = await task
|
||||
return BoundedResult(
|
||||
timed_out=False,
|
||||
value=value,
|
||||
elapsed_s=time.monotonic() - start,
|
||||
timeout_s=timeout_s,
|
||||
label=label,
|
||||
)
|
||||
return _result(start, timeout_s, label, value=await task)
|
||||
|
||||
task.cancel()
|
||||
task.add_done_callback(_consume_abandoned)
|
||||
_abandon(task)
|
||||
if on_abandon is not None:
|
||||
cleanup = asyncio.ensure_future(_run_abandon_cleanup(on_abandon))
|
||||
cleanup.add_done_callback(_consume_abandoned)
|
||||
# Phase 3a (#85125): the abandoned task may leave the backend
|
||||
# half-wedged; flag it so the owner recycles before reuse.
|
||||
# Deliberately INLINE on the loop (adopter contract: mark_suspect is
|
||||
# cheap and non-blocking). Running it synchronously guarantees the
|
||||
# mark happens-before this BoundedResult returns AND before the
|
||||
# ensure_future'd on_abandon cleanup can start (next loop tick) — an
|
||||
# offloaded mark would race both.
|
||||
asyncio.ensure_future(_run_abandon_cleanup(on_abandon)).add_done_callback(_consume_abandoned)
|
||||
# Deliberately INLINE on the loop: the mark must happen-before this result returns
|
||||
# AND before the ensure_future'd on_abandon cleanup starts (next tick).
|
||||
_mark_backend_suspect(backend, label, timeout_s)
|
||||
logger.warning(
|
||||
"[deadline] %r timed out after %.1fs; task abandoned", label, timeout_s
|
||||
)
|
||||
return BoundedResult(
|
||||
timed_out=True,
|
||||
value=None,
|
||||
elapsed_s=time.monotonic() - start,
|
||||
timeout_s=timeout_s,
|
||||
label=label,
|
||||
)
|
||||
logger.warning("[deadline] %r timed out after %.1fs; task abandoned", label, timeout_s)
|
||||
return _result(start, timeout_s, label, timed_out=True)
|
||||
finally:
|
||||
timer.cancel()
|
||||
if watchdog is not None:
|
||||
watchdog.cancel()
|
||||
# cancel() cannot stop a Timer whose callback is already running;
|
||||
# setting the event closes that race so a completed await can never
|
||||
# be misreported as a blocked loop.
|
||||
for t in timers:
|
||||
t.cancel()
|
||||
# cancel() cannot stop a Timer whose callback is already running; setting the
|
||||
# event closes that race so a completed await is never misreported as blocked.
|
||||
loop_processed_expiry.set()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bounded execution — sync flavor.
|
||||
# ---------------------------------------------------------------------------
|
||||
# --- Bounded execution — sync flavor -------------------------------------------
|
||||
|
||||
|
||||
def run_bounded_sync(
|
||||
@@ -471,180 +288,173 @@ def run_bounded_sync(
|
||||
on_timeout: Optional[Callable[[], None]] = None,
|
||||
backend: object | None = None,
|
||||
) -> BoundedResult:
|
||||
"""Run ``fn`` in a daemon worker thread under a wall-clock deadline.
|
||||
"""Run ``fn`` in a daemon worker thread under a wall-clock deadline; exceptions re-raise in
|
||||
the caller. On expiry the worker is **abandoned** (every timeout leaks one daemon thread, so
|
||||
do NOT use per-item in hot loops) and ``on_timeout`` runs best-effort in the caller's thread.
|
||||
The worker runs under ``contextvars.copy_context()`` so secret scope / session id survive.
|
||||
|
||||
On completion returns its value (exceptions re-raised in the caller).
|
||||
On expiry the worker thread is **abandoned** (daemon, so it cannot block
|
||||
interpreter exit), ``on_timeout`` (if given) runs best-effort in the
|
||||
caller's thread — e.g. to mark a backend suspect or kill a subprocess —
|
||||
and ``BoundedResult(timed_out=True)`` is returned.
|
||||
|
||||
Intended for infrequent, seconds-scale blocking backend calls. Do NOT
|
||||
use per-item in hot loops: each call spawns a thread, and every timeout
|
||||
permanently leaks an abandoned daemon thread — a wedged backend called
|
||||
in a retry loop would accumulate them.
|
||||
|
||||
``timeout=None`` (or non-positive) blocks until ``fn`` returns.
|
||||
See #94285.
|
||||
"""
|
||||
timeout_s = clamp_timeout(timeout)
|
||||
start = time.monotonic()
|
||||
if timeout_s is None:
|
||||
return BoundedResult(
|
||||
timed_out=False,
|
||||
value=fn(),
|
||||
elapsed_s=time.monotonic() - start,
|
||||
timeout_s=None,
|
||||
label=label,
|
||||
)
|
||||
return _result(start, None, label, value=fn())
|
||||
|
||||
box: dict[str, Any] = {}
|
||||
done = threading.Event()
|
||||
ctx = contextvars.copy_context()
|
||||
|
||||
def _worker() -> None:
|
||||
try:
|
||||
box["value"] = fn()
|
||||
box["value"] = ctx.run(fn)
|
||||
except BaseException as exc: # re-raised in caller; must not vanish
|
||||
box["exc"] = exc
|
||||
finally:
|
||||
done.set()
|
||||
|
||||
thread = threading.Thread(target=_worker, name=f"deadline-{label}", daemon=True)
|
||||
thread.start()
|
||||
if not done.wait(timeout_s):
|
||||
logger.warning(
|
||||
"[deadline] %r timed out after %.1fs; worker abandoned", label, timeout_s
|
||||
)
|
||||
# Phase 3a (#85125), ordering: mark suspect BEFORE owner cleanup so a
|
||||
# recycle/re-init in on_timeout never gets a stale flag on the healed
|
||||
# replacement. The sync flavor runs the mark inline — the protocol
|
||||
# contract requires mark_suspect to be cheap.
|
||||
threading.Thread(target=_worker, name=f"deadline-{label}", daemon=True).start()
|
||||
deadline = start + timeout_s
|
||||
while not done.is_set():
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
break
|
||||
done.wait(min(_BOUNDED_SYNC_WAIT_SLICE_S, remaining))
|
||||
|
||||
if not done.is_set():
|
||||
logger.warning("[deadline] %r timed out after %.1fs; worker abandoned", label, timeout_s)
|
||||
# Mark suspect BEFORE owner cleanup so a recycle in on_timeout never
|
||||
# inherits a stale flag on the healed replacement.
|
||||
_mark_backend_suspect(backend, label, timeout_s)
|
||||
if on_timeout is not None:
|
||||
try:
|
||||
on_timeout()
|
||||
except Exception:
|
||||
logger.debug("deadline on_timeout callback failed", exc_info=True)
|
||||
return BoundedResult(
|
||||
timed_out=True,
|
||||
value=None,
|
||||
elapsed_s=time.monotonic() - start,
|
||||
timeout_s=timeout_s,
|
||||
label=label,
|
||||
)
|
||||
return _result(start, timeout_s, label, timed_out=True)
|
||||
|
||||
if "exc" in box:
|
||||
raise box["exc"]
|
||||
return BoundedResult(
|
||||
timed_out=False,
|
||||
value=box.get("value"),
|
||||
elapsed_s=time.monotonic() - start,
|
||||
timeout_s=timeout_s,
|
||||
label=label,
|
||||
)
|
||||
return _result(start, timeout_s, label, value=box.get("value"))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Whole-tree process termination.
|
||||
# ---------------------------------------------------------------------------
|
||||
# --- Whole-tree process termination --------------------------------------------
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _process_tree_snapshot(pid: int, *, hard_kill: bool):
|
||||
"""Stop each hard-kill target before discovering its children: a running
|
||||
parent can fork after psutil builds its PID map and escape the final signal.
|
||||
Resume anything we stopped if signalling fails. Graceful signals never stop
|
||||
their recipients, since their handlers must remain able to run.
|
||||
"""
|
||||
descendants = []
|
||||
stopped = []
|
||||
try:
|
||||
try:
|
||||
import psutil
|
||||
root = psutil.Process(pid)
|
||||
descendants = root.children(recursive=True)
|
||||
if hard_kill:
|
||||
pending = [root]
|
||||
seen = {os.getpid()}
|
||||
known = {process.pid: process for process in descendants}
|
||||
stop_deadline = time.monotonic() + 1.0
|
||||
while pending:
|
||||
for process in pending:
|
||||
if process.pid in seen:
|
||||
continue
|
||||
seen.add(process.pid)
|
||||
if time.monotonic() >= stop_deadline:
|
||||
raise TimeoutError("process tree did not stop before snapshot deadline")
|
||||
try:
|
||||
status = process.status()
|
||||
if status in (psutil.STATUS_ZOMBIE, psutil.STATUS_DEAD):
|
||||
continue
|
||||
if status != psutil.STATUS_STOPPED:
|
||||
process.suspend()
|
||||
stopped.append(process)
|
||||
while process.status() != psutil.STATUS_STOPPED:
|
||||
if time.monotonic() >= stop_deadline:
|
||||
raise TimeoutError("process did not stop before snapshot deadline")
|
||||
time.sleep(0.001)
|
||||
except psutil.NoSuchProcess:
|
||||
continue
|
||||
# Rescan after stopping the discovered generation. A child
|
||||
# may have forked while that generation was being stopped.
|
||||
for process in root.children(recursive=True):
|
||||
known.setdefault(process.pid, process)
|
||||
descendants = list(known.values())
|
||||
pending = [process for process in descendants if process.pid not in seen]
|
||||
except Exception:
|
||||
# Preserve the existing best-effort group fallback when discovery or
|
||||
# stopping is unavailable; never strand a successfully stopped target.
|
||||
logger.debug("kill_process_tree: snapshot incomplete for pid %s", pid, exc_info=True)
|
||||
yield descendants
|
||||
finally:
|
||||
for process in stopped:
|
||||
try:
|
||||
process.resume()
|
||||
except Exception:
|
||||
logger.debug("kill_process_tree: target already gone or resume refused", exc_info=True)
|
||||
|
||||
|
||||
def kill_process_tree(pid: int, *, sig: Optional[int] = None) -> bool:
|
||||
"""Terminate ``pid`` and all its descendants, portably.
|
||||
"""Terminate ``pid`` and all its descendants, portably; True when anything was signalled.
|
||||
|
||||
Kill-on-timeout that signals only the direct child orphans process trees
|
||||
(cron scripts, in-container shells, browser daemons — #71148 class).
|
||||
|
||||
* Windows: ``taskkill /F /T`` terminates the tree (``sig`` ignored;
|
||||
Windows has no equivalent). Console-window flash is suppressed via
|
||||
``windows_hide_flags`` and the exit code is checked, so a dead or
|
||||
inaccessible PID reports ``False`` like the POSIX path.
|
||||
* POSIX: the descendant set is snapshotted via psutil (a hard
|
||||
dependency) BEFORE any signal — once the parent dies its children are
|
||||
reparented and can no longer be found by a parent walk. Then the
|
||||
process group is signalled when ``pid`` leads one (covers
|
||||
grandchildren in the same session in one syscall), and every
|
||||
snapshotted descendant is signalled individually — which also reaches
|
||||
descendants that created their OWN sessions (a child that called
|
||||
``setsid``, exactly what user shell commands do; see
|
||||
tools/environments/base.py). ``sig`` defaults to ``SIGKILL``.
|
||||
psutil's identity-aware ``Process`` (PID + create time) means a
|
||||
recycled PID is never signalled.
|
||||
|
||||
Returns True when the target (or any of its tree) was signalled, False
|
||||
when the process was already gone or every termination call failed.
|
||||
"""
|
||||
Windows: ``taskkill /F /T`` (``sig`` ignored). POSIX: snapshot descendants via
|
||||
psutil; for SIGKILL, stop and rescan the live tree so a concurrent fork cannot
|
||||
escape a stale snapshot. Signal identity-checked descendants before their
|
||||
parent, then its group when ``pid`` leads one. Stopping is best-effort with a
|
||||
bounded wait; unavailable psutil still leaves process-group cleanup. Other
|
||||
signals do not suspend recipients. ``sig`` defaults to ``SIGKILL``."""
|
||||
if sys.platform == "win32":
|
||||
try:
|
||||
from hermes_cli._subprocess_compat import windows_hide_flags
|
||||
|
||||
creationflags = windows_hide_flags()
|
||||
except Exception:
|
||||
creationflags = 0
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
["taskkill", "/F", "/T", "/PID", str(pid)],
|
||||
capture_output=True,
|
||||
timeout=15,
|
||||
check=False,
|
||||
creationflags=creationflags,
|
||||
capture_output=True, timeout=15, check=False, creationflags=creationflags,
|
||||
)
|
||||
# taskkill exits non-zero for not-found / access-denied; keep the
|
||||
# cross-platform contract (False = nothing was terminated).
|
||||
# taskkill exits non-zero for not-found / access-denied (False = nothing terminated).
|
||||
return proc.returncode == 0
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"kill_process_tree: taskkill failed for pid %s", pid, exc_info=True
|
||||
)
|
||||
logger.debug("kill_process_tree: taskkill failed for pid %s", pid, exc_info=True)
|
||||
return False
|
||||
|
||||
import signal as _signal
|
||||
|
||||
if sig is None:
|
||||
sig = _signal.SIGKILL
|
||||
|
||||
# Snapshot descendants while the parent is still alive — after it dies
|
||||
# they reparent to init/subreaper and a parent walk finds nothing.
|
||||
descendants: list = []
|
||||
try:
|
||||
import psutil
|
||||
|
||||
descendants = psutil.Process(int(pid)).children(recursive=True)
|
||||
except Exception:
|
||||
# Already gone, or psutil unavailable in a stripped env — the
|
||||
# group-signal below still covers same-session descendants.
|
||||
descendants = []
|
||||
|
||||
signalled = False
|
||||
try:
|
||||
# NOTE: getpgid→killpg has an inherent TOCTOU (pid could be reaped and
|
||||
# recycled between the calls). All existing killpg sites share it; the
|
||||
# psutil sweep below is identity-aware and does not.
|
||||
pgid = os.getpgid(pid)
|
||||
except (ProcessLookupError, PermissionError, OSError):
|
||||
pgid = None
|
||||
try:
|
||||
if pgid is not None and pgid == pid:
|
||||
# pid leads its own group: one syscall covers the whole group.
|
||||
# (The == check guards against signalling the caller's own group
|
||||
# when pid is not a leader.)
|
||||
os.killpg( # windows-footgun: ok — POSIX-only branch (win32 returns above)
|
||||
pgid, sig
|
||||
)
|
||||
else:
|
||||
os.kill(pid, sig)
|
||||
signalled = True
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
except (PermissionError, OSError):
|
||||
logger.debug("kill_process_tree: signal failed for pid %s", pid, exc_info=True)
|
||||
|
||||
# Sweep the snapshot: reaches descendants outside the parent's group
|
||||
# (their own setsid sessions) and the non-group-leader case.
|
||||
for child in descendants:
|
||||
with _process_tree_snapshot(int(pid), hard_kill=sig == _signal.SIGKILL) as descendants:
|
||||
signalled = False
|
||||
# Signal descendants while their ownership ancestry is still observable.
|
||||
# Frozen hard-kill targets cannot fork during this bottom-up teardown.
|
||||
for child in reversed(descendants):
|
||||
try:
|
||||
if child.is_running():
|
||||
child.send_signal(sig)
|
||||
signalled = True
|
||||
except Exception:
|
||||
continue
|
||||
try:
|
||||
if child.is_running(): # identity-aware: recycled PIDs skipped
|
||||
child.send_signal(sig)
|
||||
signalled = True
|
||||
except Exception:
|
||||
continue
|
||||
return signalled
|
||||
# getpgid→killpg has an inherent TOCTOU shared by every killpg site; the psutil
|
||||
# sweep below is identity-aware (PID + create time) and does not.
|
||||
pgid = os.getpgid(pid)
|
||||
except (ProcessLookupError, PermissionError, OSError):
|
||||
pgid = None
|
||||
try:
|
||||
if pgid is not None and pgid == pid:
|
||||
# pid leads its own group (the == check avoids signalling the caller's group).
|
||||
os.killpg(pgid, sig) # windows-footgun: ok — POSIX-only branch (win32 returns above)
|
||||
else:
|
||||
os.kill(pid, sig)
|
||||
signalled = True
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
except (PermissionError, OSError):
|
||||
logger.debug("kill_process_tree: signal failed for pid %s", pid, exc_info=True)
|
||||
|
||||
return signalled
|
||||
|
||||
+57
-106
@@ -1,63 +1,38 @@
|
||||
"""Context-local state for delegate_task child execution.
|
||||
|
||||
The parent Hermes process may itself be a Kanban dispatcher worker with
|
||||
HERMES_KANBAN_* variables in process env. delegate_task children run inside the
|
||||
same Python process, but they are not dispatcher-owned Kanban workers. This
|
||||
module lets code paths that resolve tool schemas or spawn subprocesses fail
|
||||
closed for delegated children without mutating global os.environ for the parent.
|
||||
|
||||
Cron jobs need the same treatment for the same reason: ``cronjob(action="run")``
|
||||
executes ``run_job()`` in-process, so a cron agent fired from inside a Kanban
|
||||
worker would otherwise inherit that worker's dispatcher identity.
|
||||
``non_dispatcher_owned_context()`` covers both cases.
|
||||
A Hermes process may itself be a Kanban dispatcher worker with HERMES_KANBAN_* in
|
||||
os.environ. In-process delegate_task children and cron jobs fired via
|
||||
``cronjob(action="run")`` are NOT dispatcher-owned, so identity gates must fail
|
||||
closed for them without mutating the process-global environment.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar, Token
|
||||
from typing import Iterator, Mapping, MutableMapping
|
||||
from typing import Iterator, Mapping, MutableMapping, overload
|
||||
|
||||
_DELEGATED_CHILD_CONTEXT: ContextVar[bool] = ContextVar(
|
||||
"hermes_delegated_child_context",
|
||||
default=False,
|
||||
)
|
||||
|
||||
# Set for any in-process execution that is NOT the dispatcher-owned worker even
|
||||
# though the worker's HERMES_KANBAN_* vars are legitimately in os.environ (cron
|
||||
# jobs fired via the `cronjob` tool). Kept separate from
|
||||
# _DELEGATED_CHILD_CONTEXT so the delegate_task-specific behaviour attached to
|
||||
# that flag (subprocess env scrubbing, its own error strings) is unchanged.
|
||||
_NON_DISPATCHER_OWNED_CONTEXT: ContextVar[bool] = ContextVar(
|
||||
"hermes_non_dispatcher_owned_context",
|
||||
default=False,
|
||||
)
|
||||
_DELEGATED_CHILD_CONTEXT: ContextVar[bool] = ContextVar("hermes_delegated_child_context", default=False)
|
||||
# Any in-process execution that is NOT the dispatcher-owned worker (cron jobs). Kept separate
|
||||
# so delegate_task-specific behaviour (subprocess env scrubbing, its error strings) is unchanged.
|
||||
_NON_DISPATCHER_OWNED_CONTEXT: ContextVar[bool] = ContextVar("hermes_non_dispatcher_owned_context", default=False)
|
||||
|
||||
DELEGATED_CHILD_ENV_MARKER = "HERMES_DELEGATED_CHILD_CONTEXT"
|
||||
|
||||
KANBAN_ENV_KEYS: tuple[str, ...] = (
|
||||
"HERMES_KANBAN_TASK",
|
||||
"HERMES_KANBAN_RUN_ID",
|
||||
"HERMES_KANBAN_WORKSPACE",
|
||||
"HERMES_KANBAN_WORKSPACES_ROOT",
|
||||
"HERMES_KANBAN_CLAIM_LOCK",
|
||||
"HERMES_KANBAN_BOARD",
|
||||
"HERMES_KANBAN_DB",
|
||||
"HERMES_KANBAN_TASK", "HERMES_KANBAN_RUN_ID", "HERMES_KANBAN_CLAIM_LOCK",
|
||||
"HERMES_KANBAN_GOAL_MODE", "HERMES_KANBAN_GOAL_MAX_TURNS",
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def delegated_child_context(session_id: str | None = None) -> Iterator[None]:
|
||||
"""Mark child execution and isolate its task-local session identity.
|
||||
|
||||
Child construction calls ``set_current_session_id`` internally, so even a
|
||||
context entered without an id must restore the parent's ContextVar. Child
|
||||
execution passes its explicit id and receives it only for this scope.
|
||||
"""
|
||||
"""Mark child execution and isolate its task-local session identity. Even a context
|
||||
entered without an id must restore the parent's session ContextVar (child
|
||||
construction calls ``set_current_session_id``)."""
|
||||
token = _DELEGATED_CHILD_CONTEXT.set(True)
|
||||
try:
|
||||
# Import lazily: session_context calls is_delegated_child_context() when
|
||||
# deciding whether the compatibility os.environ mirror is safe.
|
||||
from gateway.session_context import scoped_current_session_id
|
||||
from gateway.session_context import scoped_current_session_id # lazy: it calls is_delegated_child_context()
|
||||
|
||||
with scoped_current_session_id(session_id):
|
||||
yield
|
||||
@@ -70,49 +45,8 @@ def is_delegated_child_context() -> bool:
|
||||
return bool(_DELEGATED_CHILD_CONTEXT.get())
|
||||
|
||||
|
||||
@contextmanager
|
||||
def non_dispatcher_owned_context() -> Iterator[None]:
|
||||
"""Mark in-process execution that does NOT own the dispatcher's Kanban task.
|
||||
|
||||
A Kanban worker is a normal CLI agent whose default toolset includes
|
||||
``cronjob``; ``cronjob(action="run")`` runs ``run_job()`` inside the worker's
|
||||
own process, where ``HERMES_KANBAN_TASK`` is legitimately set. Without this
|
||||
marker the cron agent is misread as that worker: the kanban toolset is
|
||||
force-added, the worker protocol is injected into its system prompt, and
|
||||
``kanban_complete`` defaults ``task_id`` to ``$HERMES_KANBAN_TASK`` — letting
|
||||
an unrelated cron job close the worker's task and overwrite real results.
|
||||
|
||||
Scoped via ContextVar rather than by clearing ``os.environ``: the env is
|
||||
process-global and shared with the worker's own claim heartbeat, the
|
||||
gateway's Kanban watchers, and concurrent cron jobs on the parallel pool, so
|
||||
mutating it would starve the worker's claim and race those readers.
|
||||
"""
|
||||
token = _NON_DISPATCHER_OWNED_CONTEXT.set(True)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_NON_DISPATCHER_OWNED_CONTEXT.reset(token)
|
||||
|
||||
|
||||
def is_dispatcher_owned_worker_context() -> bool:
|
||||
"""Return True only when this execution owns the dispatcher's Kanban task.
|
||||
|
||||
The single predicate every ``HERMES_KANBAN_*`` identity gate should use
|
||||
before trusting those vars. False for delegate_task children and for cron
|
||||
jobs fired in-process from a worker.
|
||||
"""
|
||||
if _DELEGATED_CHILD_CONTEXT.get():
|
||||
return False
|
||||
return not _NON_DISPATCHER_OWNED_CONTEXT.get()
|
||||
|
||||
|
||||
def enter_non_dispatcher_owned_context() -> Token[bool]:
|
||||
"""Token-based form of :func:`non_dispatcher_owned_context`.
|
||||
|
||||
For callers whose scope is a long ``try`` with a matching ``finally`` rather
|
||||
than a ``with`` block (``cron.scheduler.run_job``). Pair with
|
||||
:func:`exit_non_dispatcher_owned_context`.
|
||||
"""
|
||||
"""Token form of :func:`non_dispatcher_owned_context` for long try/finally scopes."""
|
||||
return _NON_DISPATCHER_OWNED_CONTEXT.set(True)
|
||||
|
||||
|
||||
@@ -121,41 +55,58 @@ def exit_non_dispatcher_owned_context(token: Token[bool]) -> None:
|
||||
_NON_DISPATCHER_OWNED_CONTEXT.reset(token)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def non_dispatcher_owned_context() -> Iterator[None]:
|
||||
"""Mark in-process execution that does NOT own the dispatcher's Kanban task; without it
|
||||
a cron agent run inside a worker is misread as that worker (kanban toolset force-added,
|
||||
``kanban_complete`` defaulting to its task). ContextVar-scoped rather than clearing
|
||||
os.environ, which the worker's claim heartbeat and concurrent readers share."""
|
||||
token = enter_non_dispatcher_owned_context()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
exit_non_dispatcher_owned_context(token)
|
||||
|
||||
|
||||
def is_dispatcher_owned_worker_context() -> bool:
|
||||
"""The single predicate every ``HERMES_KANBAN_*`` identity gate should use."""
|
||||
return not (is_delegated_child_process_context() or _NON_DISPATCHER_OWNED_CONTEXT.get())
|
||||
|
||||
|
||||
def is_delegated_child_process_context() -> bool:
|
||||
"""Return True in this process or a subprocess spawned by a child."""
|
||||
import os
|
||||
|
||||
return bool(_DELEGATED_CHILD_CONTEXT.get()) or bool(
|
||||
os.environ.get(DELEGATED_CHILD_ENV_MARKER)
|
||||
)
|
||||
return bool(_DELEGATED_CHILD_CONTEXT.get()) or bool(os.environ.get(DELEGATED_CHILD_ENV_MARKER))
|
||||
|
||||
|
||||
def scrub_kanban_env(env: Mapping[str, str] | MutableMapping[str, str]) -> dict[str, str]:
|
||||
"""Return *env* with dispatcher-only Kanban variables removed."""
|
||||
cleaned = dict(env)
|
||||
for key in KANBAN_ENV_KEYS:
|
||||
cleaned.pop(key, None)
|
||||
"""Remove worker identity, retaining board/location and an inherited write fence.
|
||||
|
||||
TASK absence alone would promote a descendant to an orchestrator. The marker
|
||||
survives later execs, including scripts that remove TASK themselves. This is
|
||||
cooperative runtime scoping, not confinement of code with direct SQLite access.
|
||||
"""
|
||||
cleaned = {k: v for k, v in env.items() if k not in KANBAN_ENV_KEYS}
|
||||
cleaned[DELEGATED_CHILD_ENV_MARKER] = "1"
|
||||
return cleaned
|
||||
|
||||
|
||||
@overload
|
||||
def delegated_child_subprocess_env(env: Mapping[str, str]) -> dict[str, str]: ...
|
||||
|
||||
|
||||
@overload
|
||||
def delegated_child_subprocess_env(env: None = None) -> dict[str, str] | None: ...
|
||||
|
||||
|
||||
def delegated_child_subprocess_env(
|
||||
env: Mapping[str, str] | MutableMapping[str, str] | None = None,
|
||||
) -> dict[str, str] | None:
|
||||
"""Return an env override only when delegated-child lineage must cross fork.
|
||||
"""Carry worker/delegate descendant denial across a real process spawn.
|
||||
|
||||
Most subprocess call sites historically used ``env=None`` to inherit the
|
||||
process environment. In a ``delegate_task`` child, inheriting as-is leaks
|
||||
parent dispatcher ``HERMES_KANBAN_*`` vars while losing the ContextVar in
|
||||
the new process. This helper preserves normal ``env=None`` semantics for
|
||||
non-delegated calls, and only materializes a scrubbed env when the lineage
|
||||
marker must be propagated across a child-process boundary.
|
||||
Location and credentials are untouched; callers retain their existing secret policy.
|
||||
Dispatcher workers and supervised tool transports grant their own explicit scope.
|
||||
"""
|
||||
if not is_delegated_child_process_context():
|
||||
if not (is_delegated_child_process_context() or os.environ.get("HERMES_KANBAN_TASK")
|
||||
or (env and (env.get("HERMES_KANBAN_TASK") or env.get(DELEGATED_CHILD_ENV_MARKER)))):
|
||||
return None if env is None else dict(env)
|
||||
|
||||
if env is None:
|
||||
import os
|
||||
|
||||
env = os.environ
|
||||
return scrub_kanban_env(env)
|
||||
return scrub_kanban_env(os.environ if env is None else env)
|
||||
|
||||
+565
-1058
File diff suppressed because it is too large
Load Diff
+80
-109
@@ -1,29 +1,27 @@
|
||||
"""Deterministic-empty detection and cost-aware retry budgets (NS-503).
|
||||
"""Deterministic-empty detection and cost-aware retry budgets.
|
||||
|
||||
When a provider returns an empty completion, the agent loop retries up to
|
||||
3 times and then walks the fallback chain. Every attempt re-sends the full
|
||||
conversation input — at large context on paid routes this bills the user
|
||||
repeatedly for a turn that produces no text (the "charged ~$2.33 for an
|
||||
empty answer" incident class).
|
||||
On an empty completion the loop retries up to 3 times, then walks the fallback chain —
|
||||
each attempt re-bills the full input. Signaled refusals (``content_filter``, Anthropic
|
||||
``refusal``, Bedrock guardrails) are already terminal; this handles *unsignaled* empties
|
||||
(success, zero output tokens, generic finish reason — typical of portal-proxied refusals).
|
||||
|
||||
Signaled refusals (``finish_reason="content_filter"``, Anthropic
|
||||
``stop_reason="refusal"``, Bedrock guardrails) are already terminal and
|
||||
never reach the empty-retry loop. This module addresses the *unsignaled*
|
||||
empties: the provider reports a successful completion with zero output
|
||||
tokens and a generic finish reason (portal-proxied refusals commonly look
|
||||
like this).
|
||||
Two independent guards, both failing OPEN to legacy behaviour:
|
||||
|
||||
Two independent guards, both failing OPEN to today's behaviour:
|
||||
1. Deterministic-empty: two consecutive empties, both with usage present and
|
||||
``output_tokens == 0``, from the same (model, provider, finish_reason) → skip the
|
||||
remaining retries and go straight to the fallback chain. Missing usage or
|
||||
``output_tokens > 0`` (think-block stripping, whitespace, flaky decoding) never classifies.
|
||||
2. Cost-aware budget: when one empty attempt's estimated input cost exceeds the threshold
|
||||
(default $0.25), the retry budget drops from 3 to 1. Unknown pricing / missing usage /
|
||||
included routes leave it untouched.
|
||||
|
||||
1. **Deterministic-empty detection** — two consecutive empty attempts,
|
||||
both with usage present and ``output_tokens == 0``, from the same
|
||||
(model, provider, finish_reason), are treated as deterministic: the
|
||||
same prompt will keep producing the same empty. Remaining retries are
|
||||
skipped and the loop proceeds straight to the fallback chain (a
|
||||
different model may behave differently). Attempts with missing usage
|
||||
or ``output_tokens > 0`` (model generated *something* — think-block
|
||||
stripping, whitespace, flaky decoding) never classify as deterministic
|
||||
and keep the full retry budget.
|
||||
1. **Deterministic-empty detection** — two consecutive empty attempts from
|
||||
the same (model, provider, finish_reason) are treated as deterministic
|
||||
when usage proves zero output, or when usage is absent and the assembled
|
||||
responses contain neither content nor reasoning. Remaining retries are
|
||||
skipped and the loop proceeds straight to the fallback chain (a different
|
||||
model may behave differently). Mixed evidence or any generated tokens keep
|
||||
the full retry budget.
|
||||
|
||||
2. **Cost-aware retry budget** — when the estimated input cost of a
|
||||
single empty attempt exceeds the configured threshold (default
|
||||
@@ -58,11 +56,9 @@ REDUCED_EMPTY_RETRY_BUDGET = 1
|
||||
DEFAULT_COST_THRESHOLD_USD = Decimal("0.25")
|
||||
DEFAULT_GUARD_ENABLED = True
|
||||
|
||||
# Attribute names stashed on the agent object. State is scoped to one
|
||||
# consecutive empty streak: it is cleared whenever a streak starts
|
||||
# (``_empty_content_retries == 0`` at record time), which transparently
|
||||
# honours every existing reset site (turn start, compaction, tool
|
||||
# success, fallback activation) without touching them.
|
||||
# Agent-object attribute names. State is scoped to one consecutive empty streak: cleared
|
||||
# whenever ``_empty_content_retries == 0`` at record time, so every existing counter-reset
|
||||
# site (turn start, compaction, tool success, fallback activation) is honoured.
|
||||
_ATTEMPTS_ATTR = "_empty_attempt_history"
|
||||
_STREAK_COST_ATTR = "_empty_streak_cost_usd"
|
||||
_ENABLED_ATTR = "_empty_guard_enabled"
|
||||
@@ -78,6 +74,7 @@ class EmptyAttempt:
|
||||
finish_reason: str
|
||||
usage_present: bool
|
||||
zero_output: bool
|
||||
observed_generation: bool
|
||||
|
||||
@property
|
||||
def signature(self) -> tuple:
|
||||
@@ -85,23 +82,14 @@ class EmptyAttempt:
|
||||
|
||||
|
||||
def resolve_guard_settings(section: Any) -> Tuple[bool, Decimal]:
|
||||
"""Resolve ``agent.empty_response_guard`` config into (enabled, threshold).
|
||||
|
||||
Tolerant of malformed input: anything that isn't a well-formed dict
|
||||
(or well-formed values within it) falls back to the schema defaults.
|
||||
Called once per agent at init; the resolved values are stashed on the
|
||||
agent object so the hot loop never re-reads config.
|
||||
"""
|
||||
"""Resolve ``agent.empty_response_guard`` into (enabled, threshold); malformed input → schema defaults."""
|
||||
if not isinstance(section, dict):
|
||||
return (DEFAULT_GUARD_ENABLED, DEFAULT_COST_THRESHOLD_USD)
|
||||
|
||||
enabled_raw = section.get("enabled", DEFAULT_GUARD_ENABLED)
|
||||
if isinstance(enabled_raw, bool):
|
||||
enabled = enabled_raw
|
||||
elif isinstance(enabled_raw, str):
|
||||
# YAML quoting can turn true/false into strings.
|
||||
enabled = enabled_raw.strip().lower() not in ("0", "false", "no", "off")
|
||||
else:
|
||||
enabled = section.get("enabled", DEFAULT_GUARD_ENABLED)
|
||||
if isinstance(enabled, str): # YAML quoting can turn true/false into strings.
|
||||
enabled = enabled.strip().lower() not in ("0", "false", "no", "off")
|
||||
elif not isinstance(enabled, bool):
|
||||
enabled = DEFAULT_GUARD_ENABLED
|
||||
|
||||
threshold = DEFAULT_COST_THRESHOLD_USD
|
||||
@@ -112,28 +100,19 @@ def resolve_guard_settings(section: Any) -> Tuple[bool, Decimal]:
|
||||
if candidate > 0:
|
||||
threshold = candidate
|
||||
except Exception: # noqa: BLE001 — malformed config must not break init
|
||||
logger.debug(
|
||||
"empty-guard: invalid cost_threshold_usd %r, using default",
|
||||
threshold_raw,
|
||||
)
|
||||
logger.debug("empty-guard: invalid cost_threshold_usd %r, using default", threshold_raw)
|
||||
return (enabled, threshold)
|
||||
|
||||
|
||||
def guard_enabled(agent: Any) -> bool:
|
||||
"""Whether the guard is enabled for this agent (config-resolved).
|
||||
|
||||
Agents built before the config was threaded through (tests, embedded
|
||||
callers) simply get the default: enabled.
|
||||
"""
|
||||
"""Config-resolved enabled flag; agents built without config default to enabled."""
|
||||
value = getattr(agent, _ENABLED_ATTR, DEFAULT_GUARD_ENABLED)
|
||||
return value if isinstance(value, bool) else DEFAULT_GUARD_ENABLED
|
||||
|
||||
|
||||
def _cost_threshold_usd(agent: Any) -> Decimal:
|
||||
value = getattr(agent, _THRESHOLD_ATTR, None)
|
||||
if isinstance(value, Decimal) and value > 0:
|
||||
return value
|
||||
return DEFAULT_COST_THRESHOLD_USD
|
||||
return value if isinstance(value, Decimal) and value > 0 else DEFAULT_COST_THRESHOLD_USD
|
||||
|
||||
|
||||
def _attempts(agent: Any) -> List[EmptyAttempt]:
|
||||
@@ -144,25 +123,30 @@ def _attempts(agent: Any) -> List[EmptyAttempt]:
|
||||
return attempts
|
||||
|
||||
|
||||
def _estimate_attempt_cost(agent: Any, response: Any) -> Optional[Decimal]:
|
||||
"""Best-effort USD estimate for one attempt. None when unknown."""
|
||||
def _normalized_usage(agent: Any, response: Any, what: str) -> Any:
|
||||
"""Canonical usage for ``response`` or None (no usage / normalization failed)."""
|
||||
raw_usage = getattr(response, "usage", None)
|
||||
if not raw_usage:
|
||||
return None
|
||||
try:
|
||||
from agent.usage_pricing import estimate_usage_cost, normalize_usage
|
||||
from agent.usage_pricing import normalize_usage
|
||||
return normalize_usage(raw_usage, provider=getattr(agent, "provider", None),
|
||||
api_mode=getattr(agent, "api_mode", None))
|
||||
except Exception: # noqa: BLE001 — pricing must never break the loop
|
||||
logger.debug("empty-guard: %s failed", what, exc_info=True)
|
||||
return None
|
||||
|
||||
canonical = normalize_usage(
|
||||
raw_usage,
|
||||
provider=getattr(agent, "provider", None),
|
||||
api_mode=getattr(agent, "api_mode", None),
|
||||
)
|
||||
|
||||
def _estimate_attempt_cost(agent: Any, response: Any) -> Optional[Decimal]:
|
||||
"""Best-effort USD estimate for one attempt. None when unknown."""
|
||||
canonical = _normalized_usage(agent, response, "cost estimation")
|
||||
if canonical is None:
|
||||
return None
|
||||
try:
|
||||
from agent.usage_pricing import estimate_usage_cost
|
||||
result = estimate_usage_cost(
|
||||
getattr(agent, "model", "") or "",
|
||||
canonical,
|
||||
provider=getattr(agent, "provider", None),
|
||||
base_url=getattr(agent, "base_url", None),
|
||||
api_key=getattr(agent, "api_key", None),
|
||||
getattr(agent, "model", "") or "", canonical, provider=getattr(agent, "provider", None),
|
||||
base_url=getattr(agent, "base_url", None), api_key=getattr(agent, "api_key", None),
|
||||
)
|
||||
except Exception: # noqa: BLE001 — pricing must never break the loop
|
||||
logger.debug("empty-guard: cost estimation failed", exc_info=True)
|
||||
@@ -172,43 +156,31 @@ def _estimate_attempt_cost(agent: Any, response: Any) -> Optional[Decimal]:
|
||||
|
||||
def _zero_output(agent: Any, response: Any) -> tuple:
|
||||
"""Return (usage_present, zero_output) for a response, failing open."""
|
||||
raw_usage = getattr(response, "usage", None)
|
||||
if not raw_usage:
|
||||
return (False, False)
|
||||
try:
|
||||
from agent.usage_pricing import normalize_usage
|
||||
|
||||
canonical = normalize_usage(
|
||||
raw_usage,
|
||||
provider=getattr(agent, "provider", None),
|
||||
api_mode=getattr(agent, "api_mode", None),
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("empty-guard: usage normalization failed", exc_info=True)
|
||||
canonical = _normalized_usage(agent, response, "usage normalization")
|
||||
if canonical is None:
|
||||
return (False, False)
|
||||
output = getattr(canonical, "output_tokens", None)
|
||||
if output is None:
|
||||
# A present-but-empty usage object (some proxies) normalizes to all zeros;
|
||||
# a genuine completion always has input tokens — no evidence, fail open.
|
||||
if output is None or getattr(canonical, "prompt_tokens", 0) <= 0:
|
||||
return (False, False)
|
||||
# A present-but-empty usage object (some proxies emit usage with no
|
||||
# fields) normalizes to all zeros. A genuine completion always has
|
||||
# input tokens — without them the usage is not evidence, fail open.
|
||||
if getattr(canonical, "prompt_tokens", 0) <= 0:
|
||||
return (False, False)
|
||||
# Reasoning tokens count as real generation — a reasoning-only
|
||||
# response is NOT a deterministic empty (the prefill-continuation
|
||||
# path upstream owns that case).
|
||||
# Reasoning tokens are real generation: a reasoning-only response is NOT
|
||||
# a deterministic empty (the prefill-continuation path owns that case).
|
||||
reasoning = getattr(canonical, "reasoning_tokens", 0) or 0
|
||||
return (True, (output + reasoning) == 0)
|
||||
|
||||
|
||||
def record_empty_attempt(agent: Any, *, finish_reason: str, response: Any) -> None:
|
||||
def record_empty_attempt(
|
||||
agent: Any,
|
||||
*,
|
||||
finish_reason: str,
|
||||
response: Any,
|
||||
observed_generation: bool = True,
|
||||
) -> None:
|
||||
"""Record one empty completion in the current streak.
|
||||
|
||||
Must be called before ``_empty_content_retries`` is incremented for
|
||||
this attempt: a counter of 0 marks the start of a new streak and
|
||||
clears prior history (this transparently follows every existing
|
||||
counter-reset site).
|
||||
"""
|
||||
Call BEFORE ``_empty_content_retries`` is incremented: a counter of 0 marks a new
|
||||
streak and clears prior history."""
|
||||
attempts = _attempts(agent)
|
||||
if getattr(agent, "_empty_content_retries", 0) == 0:
|
||||
attempts.clear()
|
||||
@@ -222,6 +194,7 @@ def record_empty_attempt(agent: Any, *, finish_reason: str, response: Any) -> No
|
||||
finish_reason=str(finish_reason or ""),
|
||||
usage_present=usage_present,
|
||||
zero_output=zero_output,
|
||||
observed_generation=bool(observed_generation),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -234,10 +207,10 @@ def record_empty_attempt(agent: Any, *, finish_reason: str, response: Any) -> No
|
||||
def deterministic_empty(agent: Any) -> bool:
|
||||
"""True when the current streak looks deterministic.
|
||||
|
||||
Requires >= 2 consecutive attempts, ALL with usage present, zero
|
||||
output tokens, and an identical (model, provider, finish_reason)
|
||||
signature. Any attempt with missing usage or non-zero output keeps
|
||||
this False (fail open — transients deserve their retries).
|
||||
Requires >= 2 consecutive attempts with an identical (model, provider,
|
||||
finish_reason) signature. Usage-backed attempts must all prove zero output.
|
||||
Usage-absent attempts must all have no observed content or reasoning. Mixed
|
||||
evidence fails open so ambiguous transients keep their retries.
|
||||
"""
|
||||
if not guard_enabled(agent):
|
||||
return False
|
||||
@@ -245,21 +218,21 @@ def deterministic_empty(agent: Any) -> bool:
|
||||
if len(attempts) < 2:
|
||||
return False
|
||||
first = attempts[0]
|
||||
return all(
|
||||
a.usage_present and a.zero_output and a.signature == first.signature
|
||||
for a in attempts
|
||||
same_signature = all(a.signature == first.signature for a in attempts)
|
||||
usage_proves_empty = all(a.usage_present and a.zero_output for a in attempts)
|
||||
response_proves_empty = all(
|
||||
not a.usage_present and not a.observed_generation for a in attempts
|
||||
)
|
||||
return same_signature and (usage_proves_empty or response_proves_empty)
|
||||
|
||||
|
||||
def empty_retry_budget(agent: Any, response: Any) -> int:
|
||||
"""Empty-retry budget for the current streak (3, or 1 when a single
|
||||
attempt is estimated to cost more than the configured threshold)."""
|
||||
"""Empty-retry budget for the current streak (3, or 1 when a single attempt is
|
||||
estimated to cost more than the configured threshold)."""
|
||||
if not guard_enabled(agent):
|
||||
return DEFAULT_EMPTY_RETRY_BUDGET
|
||||
cost = _estimate_attempt_cost(agent, response)
|
||||
if cost is None:
|
||||
return DEFAULT_EMPTY_RETRY_BUDGET
|
||||
if cost >= _cost_threshold_usd(agent):
|
||||
if cost is not None and cost >= _cost_threshold_usd(agent):
|
||||
return REDUCED_EMPTY_RETRY_BUDGET
|
||||
return DEFAULT_EMPTY_RETRY_BUDGET
|
||||
|
||||
@@ -267,6 +240,4 @@ def empty_retry_budget(agent: Any, response: Any) -> int:
|
||||
def streak_cost_usd(agent: Any) -> Optional[Decimal]:
|
||||
"""Accumulated estimated cost of the current empty streak, if known."""
|
||||
cost = getattr(agent, _STREAK_COST_ATTR, None)
|
||||
if cost is None or cost <= 0:
|
||||
return None
|
||||
return cost
|
||||
return cost if cost is not None and cost > 0 else None
|
||||
|
||||
+829
-2031
File diff suppressed because it is too large
Load Diff
+84
-176
@@ -1,29 +1,15 @@
|
||||
"""Structured error-surface descriptors for UI clients (Desktop/TUI).
|
||||
|
||||
Maps the internal failure taxonomy (``agent.error_classifier.FailoverReason``
|
||||
values carried in turn results as ``failure_reason``, or raw exceptions from
|
||||
the turn dispatcher) onto a small, stable wire descriptor:
|
||||
|
||||
{"layer": <ui layer>, "code": <specific code>, "retryable": <bool>}
|
||||
|
||||
The *layer* names which part of the stack failed, so clients can say
|
||||
"Provider error" / "Gateway error" instead of toasting an opaque string and
|
||||
leaving the user to guess whether the model, the gateway, or the app froze:
|
||||
|
||||
provider — the model/provider API rejected or failed the call
|
||||
endpoint — a user-configured custom/local endpoint failed (transport)
|
||||
streaming — the provider's SSE/stream connection dropped mid-turn
|
||||
auth — authentication/authorization failed
|
||||
billing — credits/quota wall (clients usually have a richer
|
||||
billing_block descriptor; this is the fallback signal)
|
||||
gateway — the local gateway/agent runtime itself errored
|
||||
runtime — agent initialization / local environment failure
|
||||
disk — local disk full / persistence failure
|
||||
|
||||
This module is intentionally dependency-light and NEVER raises: surfacing
|
||||
diagnostics must not be able to break the error path it describes. Clients
|
||||
treat the descriptor as advisory — an absent or partial descriptor falls
|
||||
back to today's string-sniffing behavior (older backends keep working).
|
||||
Maps the internal failure taxonomy (``FailoverReason`` values carried in turn
|
||||
results as ``failure_reason``, or raw exceptions from the turn dispatcher)
|
||||
onto a small, stable wire descriptor ``{"layer", "code", "retryable"}``.
|
||||
Layers: provider (model API rejected/failed), endpoint (user-configured
|
||||
custom/local endpoint transport failure), streaming (SSE dropped mid-turn),
|
||||
auth, billing (fallback signal; clients usually have a richer
|
||||
``billing_block``), gateway (local runtime errored), disk (disk full).
|
||||
Dependency-light and NEVER raises: surfacing diagnostics must not break the
|
||||
error path it describes. Descriptors are advisory — clients fall back to
|
||||
string sniffing when absent or partial.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -33,95 +19,48 @@ from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# UI layers (wire values — stable contract with desktop/TUI clients).
|
||||
LAYER_PROVIDER = "provider"
|
||||
LAYER_ENDPOINT = "endpoint"
|
||||
LAYER_STREAMING = "streaming"
|
||||
LAYER_AUTH = "auth"
|
||||
LAYER_BILLING = "billing"
|
||||
LAYER_GATEWAY = "gateway"
|
||||
LAYER_RUNTIME = "runtime"
|
||||
LAYER_DISK = "disk"
|
||||
|
||||
# failure_reason (FailoverReason.value) → UI layer. Reasons not listed fall
|
||||
# back to LAYER_PROVIDER: every FailoverReason is produced by classifying a
|
||||
# provider API call, so "the provider call failed" is the honest default.
|
||||
# failure_reason → UI layer. Unlisted reasons fall back to LAYER_PROVIDER:
|
||||
# every FailoverReason comes from classifying a provider call.
|
||||
_REASON_TO_LAYER = {
|
||||
"auth": LAYER_AUTH,
|
||||
"auth_permanent": LAYER_AUTH,
|
||||
"billing": LAYER_BILLING,
|
||||
"billing_unverified": LAYER_BILLING,
|
||||
"auth": LAYER_AUTH, "auth_permanent": LAYER_AUTH, "billing": LAYER_BILLING, "billing_unverified": LAYER_BILLING,
|
||||
}
|
||||
|
||||
# Transport-ish reasons: the failure is between us and the base_url, not a
|
||||
# verdict the provider returned. On a custom/local endpoint these point at
|
||||
# the user's endpoint config, so they surface as LAYER_ENDPOINT there.
|
||||
_TRANSPORT_REASONS = {
|
||||
"timeout",
|
||||
"ssl_cert_verification",
|
||||
}
|
||||
# Failures between us and the base_url (not a provider verdict); on a
|
||||
# custom/local endpoint they point at the user's endpoint config.
|
||||
_TRANSPORT_REASONS = {"timeout", "ssl_cert_verification"}
|
||||
|
||||
# Reasons that are deterministic for the request — a bare "Retry" repeats the
|
||||
# same failure, so clients shouldn't lead with it. Fallback only: results
|
||||
# from current backends carry the classifier's own verdict in
|
||||
# ``failure_retryable`` and never consult this set. Kept in sync with
|
||||
# ``classify_api_error``'s retryable=False verdicts.
|
||||
# Deterministic for the request — a bare "Retry" repeats the failure. Fallback
|
||||
# only: current backends stamp the classifier's verdict in ``failure_retryable``.
|
||||
# Kept in sync with ``classify_api_error``'s retryable=False verdicts.
|
||||
_NON_RETRYABLE_REASONS = {
|
||||
"auth",
|
||||
"auth_permanent",
|
||||
"billing",
|
||||
"billing_unverified",
|
||||
"content_policy_blocked",
|
||||
"provider_policy_blocked",
|
||||
"model_not_found",
|
||||
"format_error",
|
||||
"ssl_cert_verification",
|
||||
"auth", "auth_permanent", "billing", "billing_unverified", "content_policy_blocked",
|
||||
"provider_policy_blocked", "model_not_found", "format_error", "ssl_cert_verification",
|
||||
}
|
||||
|
||||
# Providers whose base_url is user-supplied rather than a known vendor —
|
||||
# transport failures against these are endpoint-config problems.
|
||||
_CUSTOM_ENDPOINT_PROVIDERS = {
|
||||
"custom",
|
||||
"local",
|
||||
"llama.cpp",
|
||||
"llamacpp",
|
||||
"ollama",
|
||||
"lmstudio",
|
||||
"vllm",
|
||||
}
|
||||
# Providers whose base_url is user-supplied rather than a known vendor.
|
||||
_CUSTOM_ENDPOINT_PROVIDERS = {"custom", "local", "llama.cpp", "llamacpp", "ollama", "lmstudio", "vllm"}
|
||||
|
||||
# Message fragments that mark a mid-stream connection drop. Deliberately
|
||||
# narrow: these strings come from our own retry-exhaustion summaries and the
|
||||
# OpenAI SDK's stream-abort errors.
|
||||
# Mid-stream drop markers. Deliberately narrow: our own retry-exhaustion
|
||||
# summaries plus the OpenAI SDK's stream-abort errors.
|
||||
_STREAM_DROP_FRAGMENTS = (
|
||||
"stream connection",
|
||||
"peer closed connection",
|
||||
"incomplete chunked read",
|
||||
"connection broken",
|
||||
"stream ended prematurely",
|
||||
"sse",
|
||||
"mid-stream",
|
||||
"stream connection", "peer closed connection", "incomplete chunked read",
|
||||
"connection broken", "stream ended prematurely", "sse", "mid-stream",
|
||||
)
|
||||
|
||||
# Exception modules that indicate the failure came from an API/transport call
|
||||
# (vs. a bug in our own dispatcher code, which is a gateway-layer failure).
|
||||
# Covers every SDK family our provider adapters raise from: OpenAI-compatible
|
||||
# (openai/httpx/httpcore), Anthropic, Bedrock (botocore/boto3), Google
|
||||
# (google.*/grpc), plus raw transports (requests/aiohttp/ssl/socket/urllib).
|
||||
# Exception top-level modules that mean "API/transport call failed" (vs. a bug
|
||||
# in our dispatcher = gateway layer): every SDK family our adapters raise from
|
||||
# plus raw transports.
|
||||
_API_EXC_MODULE_PREFIXES = (
|
||||
"openai",
|
||||
"httpx",
|
||||
"httpcore",
|
||||
"anthropic",
|
||||
"botocore",
|
||||
"boto3",
|
||||
"google",
|
||||
"grpc",
|
||||
"requests",
|
||||
"aiohttp",
|
||||
"ssl",
|
||||
"socket",
|
||||
"urllib",
|
||||
"openai", "httpx", "httpcore", "anthropic", "botocore", "boto3", "google",
|
||||
"grpc", "requests", "aiohttp", "ssl", "socket", "urllib",
|
||||
)
|
||||
|
||||
|
||||
@@ -131,36 +70,38 @@ def _is_custom_endpoint(provider: Optional[str]) -> bool:
|
||||
|
||||
|
||||
def _looks_like_stream_drop(message: str) -> bool:
|
||||
msg = message.lower()
|
||||
return any(fragment in msg for fragment in _STREAM_DROP_FRAGMENTS)
|
||||
return any(fragment in message.lower() for fragment in _STREAM_DROP_FRAGMENTS)
|
||||
|
||||
|
||||
def _surface(
|
||||
layer: str,
|
||||
code: str,
|
||||
retryable: bool,
|
||||
provider: str = "",
|
||||
model: str = "",
|
||||
) -> dict:
|
||||
out = {"layer": layer, "code": code, "retryable": bool(retryable)}
|
||||
# The failing session's identity, captured at classification time so
|
||||
# clients report the model/provider that actually failed — not whatever
|
||||
# the foreground composer points at when a button is clicked later.
|
||||
if provider:
|
||||
out["provider"] = provider
|
||||
if model:
|
||||
out["model"] = model
|
||||
return out
|
||||
def _surface(layer: str, code: str, retryable: bool, provider: str = "", model: str = "") -> dict:
|
||||
# Identity captured at classification time, so clients report the session
|
||||
# that actually failed — not whatever the composer points at later.
|
||||
identity = {k: v for k, v in (("provider", provider), ("model", model)) if v}
|
||||
return {"layer": layer, "code": code, "retryable": bool(retryable), **identity}
|
||||
|
||||
|
||||
def build_error_surface_from_result(
|
||||
result: Any, provider: str = "", model: str = ""
|
||||
) -> Optional[dict]:
|
||||
def _disk_full(candidate: Any) -> bool:
|
||||
try:
|
||||
from hermes_state_errors import is_disk_full_error
|
||||
|
||||
return bool(is_disk_full_error(candidate))
|
||||
except Exception: # pragma: no cover - defensive import guard
|
||||
return False
|
||||
|
||||
|
||||
def _result_layer(reason: str, error_text: str, provider: str) -> str:
|
||||
if reason in _REASON_TO_LAYER:
|
||||
return _REASON_TO_LAYER[reason]
|
||||
if reason in _TRANSPORT_REASONS and _is_custom_endpoint(provider):
|
||||
return LAYER_ENDPOINT
|
||||
return LAYER_STREAMING if _looks_like_stream_drop(error_text) else LAYER_PROVIDER
|
||||
|
||||
|
||||
def build_error_surface_from_result(result: Any, provider: str = "", model: str = "") -> Optional[dict]:
|
||||
"""Descriptor for a returned-error turn result (``failed=True`` dicts).
|
||||
|
||||
Reads the ``failure_reason`` the conversation loop already stamps
|
||||
(a ``FailoverReason.value``) plus the error text, and maps them onto a
|
||||
UI layer. Returns None when the result carries no failure signal.
|
||||
Uses the stamped ``failure_reason`` plus error text. None when the result
|
||||
carries no failure signal.
|
||||
"""
|
||||
try:
|
||||
if not isinstance(result, dict):
|
||||
@@ -169,90 +110,57 @@ def build_error_surface_from_result(
|
||||
reason = str(result.get("failure_reason") or "").strip()
|
||||
if not error_text and not reason:
|
||||
return None
|
||||
|
||||
# Disk-full wins outright: the fix (free space) is unrelated to the
|
||||
# provider stack, and hermes_state owns the pattern list.
|
||||
try:
|
||||
from hermes_state import is_disk_full_error
|
||||
|
||||
if error_text and is_disk_full_error(error_text):
|
||||
return _surface(LAYER_DISK, "disk_full", False, provider, model)
|
||||
except Exception: # pragma: no cover - defensive import guard
|
||||
pass
|
||||
|
||||
# provider stack; hermes_state owns the pattern list.
|
||||
if error_text and _disk_full(error_text):
|
||||
return _surface(LAYER_DISK, "disk_full", False, provider, model)
|
||||
if result.get("billing_block") or reason in ("billing", "billing_unverified"):
|
||||
return _surface(LAYER_BILLING, reason or "billing", False, provider, model)
|
||||
|
||||
if not reason:
|
||||
# Failed result without a classified reason (legacy paths).
|
||||
if _looks_like_stream_drop(error_text):
|
||||
return _surface(LAYER_STREAMING, "stream_drop", True, provider, model)
|
||||
return _surface(LAYER_PROVIDER, "unknown", True, provider, model)
|
||||
|
||||
layer = _REASON_TO_LAYER.get(reason)
|
||||
if layer is None:
|
||||
if reason in _TRANSPORT_REASONS and _is_custom_endpoint(provider):
|
||||
layer = LAYER_ENDPOINT
|
||||
elif _looks_like_stream_drop(error_text):
|
||||
layer = LAYER_STREAMING
|
||||
else:
|
||||
layer = LAYER_PROVIDER
|
||||
# Prefer the classifier's own retry verdict when the result carries it
|
||||
# (conversation_loop stamps ``failure_retryable`` next to
|
||||
# ``failure_reason``); the reason-set fallback covers older results.
|
||||
if not reason: # failed result without a classified reason (legacy paths)
|
||||
drop = _looks_like_stream_drop(error_text)
|
||||
return _surface(LAYER_STREAMING if drop else LAYER_PROVIDER, "stream_drop" if drop else "unknown", True, provider, model)
|
||||
# Prefer the classifier's own verdict (``failure_retryable``); the
|
||||
# reason-set fallback covers older results.
|
||||
retryable = result.get("failure_retryable")
|
||||
if not isinstance(retryable, bool):
|
||||
retryable = reason not in _NON_RETRYABLE_REASONS
|
||||
return _surface(layer, reason, retryable, provider, model)
|
||||
return _surface(_result_layer(reason, error_text, provider), reason, retryable, provider, model)
|
||||
except Exception: # pragma: no cover — never break the error path
|
||||
logger.debug("error_surface: result classification failed", exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
def build_error_surface_from_exception(
|
||||
exc: BaseException, provider: str = "", model: str = ""
|
||||
) -> Optional[dict]:
|
||||
def build_error_surface_from_exception(exc: BaseException, provider: str = "", model: str = "") -> Optional[dict]:
|
||||
"""Descriptor for an exception that escaped the turn dispatcher.
|
||||
|
||||
API/transport exceptions are classified through the real
|
||||
``classify_api_error`` pipeline (same taxonomy as the retry loop);
|
||||
anything else is a gateway-layer failure — a bug or environment problem
|
||||
in our own dispatcher, not a provider verdict.
|
||||
API/transport exceptions go through ``classify_api_error`` (same taxonomy
|
||||
as the retry loop); anything else is a gateway-layer failure.
|
||||
"""
|
||||
try:
|
||||
message = str(exc) or type(exc).__name__
|
||||
|
||||
try:
|
||||
from hermes_state import is_disk_full_error
|
||||
|
||||
if is_disk_full_error(exc):
|
||||
return _surface(LAYER_DISK, "disk_full", False, provider, model)
|
||||
except Exception: # pragma: no cover - defensive import guard
|
||||
pass
|
||||
|
||||
exc_module = type(exc).__module__ or ""
|
||||
api_like = exc_module.split(".")[0] in _API_EXC_MODULE_PREFIXES or hasattr(
|
||||
exc, "status_code"
|
||||
)
|
||||
|
||||
if _disk_full(exc):
|
||||
return _surface(LAYER_DISK, "disk_full", False, provider, model)
|
||||
api_like = (type(exc).__module__ or "").split(".")[0] in _API_EXC_MODULE_PREFIXES or hasattr(exc, "status_code")
|
||||
if not api_like or not isinstance(exc, Exception):
|
||||
return _surface(LAYER_GATEWAY, type(exc).__name__, True, provider, model)
|
||||
|
||||
from agent.error_classifier import classify_api_error
|
||||
|
||||
classified = classify_api_error(exc, provider=provider, model=model)
|
||||
reason = classified.reason.value
|
||||
|
||||
synthetic = {
|
||||
"error": classified.message or message,
|
||||
"failure_reason": reason,
|
||||
}
|
||||
surface = build_error_surface_from_result(
|
||||
synthetic, provider=provider, model=model
|
||||
)
|
||||
synthetic = {"error": classified.message or message, "failure_reason": classified.reason.value}
|
||||
surface = build_error_surface_from_result(synthetic, provider=provider, model=model)
|
||||
if surface is not None:
|
||||
surface["retryable"] = bool(classified.retryable)
|
||||
return surface
|
||||
except Exception: # pragma: no cover — never break the error path
|
||||
logger.debug("error_surface: exception classification failed", exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
|
||||
LAYER_RUNTIME = "runtime"
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
@@ -1,13 +1,10 @@
|
||||
class SSLConfigurationError(Exception):
|
||||
"""Raised when SSL/TLS certificate bundle configuration fails."""
|
||||
pass
|
||||
|
||||
|
||||
class EmptyStreamError(RuntimeError):
|
||||
"""Raised when a provider closes a stream without yielding a response."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class MoAPresetNotFoundError(ValueError):
|
||||
"""Raised when a persisted MoA preset no longer exists in config."""
|
||||
|
||||
+81
-109
@@ -1,125 +1,113 @@
|
||||
"""Global emergency stop (ESTOP) — a resumable pause for NEW work only.
|
||||
|
||||
``hermes pause`` writes a sentinel file at ``$HERMES_HOME/ESTOP``;
|
||||
``hermes resume`` removes it. While the sentinel exists:
|
||||
|
||||
* the cron scheduler skips dispatching due jobs (``cron/scheduler.py:tick``),
|
||||
* the embedded kanban dispatcher skips spawning workers
|
||||
(``gateway/kanban_watchers.py``),
|
||||
* new gateway turns get a brief "Hermes is paused" reply instead of an
|
||||
agent run (``gateway/run.py:_handle_message``).
|
||||
|
||||
In-flight work is NEVER killed — this is pause-new-work, not panic/exit.
|
||||
The check is a single ``os.stat`` so callers may run it every tick; no
|
||||
caching beyond the OS is performed, so engaging/disengaging takes effect on
|
||||
the very next check.
|
||||
|
||||
The sentinel body is optional JSON ``{"reason": ..., "engaged_at": ...}``.
|
||||
A corrupt or empty file still counts as engaged (fail safe): the pause must
|
||||
hold even if the file was created by ``touch ~/.hermes/ESTOP``.
|
||||
|
||||
Ported from: gastownhall/gastown estop.go (MIT). Related prior art:
|
||||
#26778 (/panic — kill/exit semantics; deliberately different, ours is
|
||||
resumable) and #44617 (interrupting in-flight cron; deliberately out of
|
||||
scope here).
|
||||
``hermes pause`` writes a sentinel at ``$HERMES_HOME/ESTOP``; ``hermes resume``
|
||||
removes it. While it exists the cron scheduler, kanban dispatcher and new gateway
|
||||
turns skip work; in-flight work is never killed. The check is one or two uncached
|
||||
``os.stat`` calls (process home + fleet root when they differ). The body is optional
|
||||
JSON ``{"reason", "engaged_at"}``; a corrupt/empty file still counts as engaged
|
||||
(fail safe, e.g. ``touch ~/.hermes/ESTOP``). Ported from gastownhall/gastown estop.go (MIT).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from contextlib import suppress
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
# Same profile-aware / fleet-root resolvers the file-safety guards use (fail-open to ~/.hermes).
|
||||
from agent.file_safety import _hermes_home_path as _hermes_home, _hermes_root_path as _canonical_root
|
||||
|
||||
SENTINEL_NAME = "ESTOP"
|
||||
|
||||
# Per-component "logged already for this engagement" flags so a paused
|
||||
# dispatch loop logs once per engagement instead of once per tick.
|
||||
# Per-component "logged already for this engagement" flags: log once per engagement, not per tick.
|
||||
_log_lock = threading.Lock()
|
||||
_logged_components: set[str] = set()
|
||||
|
||||
|
||||
def _hermes_home() -> Path:
|
||||
"""Resolve the active HERMES_HOME (profile-aware) at call time."""
|
||||
try:
|
||||
from hermes_constants import get_hermes_home
|
||||
return get_hermes_home()
|
||||
except Exception:
|
||||
return Path(os.path.expanduser("~/.hermes"))
|
||||
|
||||
|
||||
def sentinel_path() -> Path:
|
||||
"""Path of the ESTOP sentinel under the active HERMES_HOME."""
|
||||
"""Path of the ESTOP sentinel this process would write on `hermes pause`."""
|
||||
return _hermes_home() / SENTINEL_NAME
|
||||
|
||||
|
||||
def is_engaged() -> bool:
|
||||
"""Cheap check (one stat): is the global emergency stop engaged?
|
||||
|
||||
Fail SAFE on stat errors: if we cannot determine whether the sentinel
|
||||
exists (permission error, transient I/O failure on HERMES_HOME), report
|
||||
engaged. The module contract is that the pause must hold even when the
|
||||
sentinel is unreadable — a fail-open here would silently lift an
|
||||
operator's emergency stop exactly when the filesystem is misbehaving.
|
||||
"""
|
||||
def _candidate_sentinel_paths() -> list:
|
||||
"""Profile home first, then the fleet root if it is a different directory: a profile
|
||||
gateway (HERMES_HOME=~/.hermes/profiles/<n>) must still honor an operator's ~/.hermes/ESTOP."""
|
||||
primary = sentinel_path()
|
||||
try:
|
||||
return sentinel_path().exists()
|
||||
except OSError:
|
||||
return True
|
||||
root = _canonical_root() / SENTINEL_NAME
|
||||
except Exception:
|
||||
return [primary]
|
||||
try:
|
||||
distinct = root.resolve() != primary.resolve()
|
||||
except Exception:
|
||||
# Non-Path test doubles fail .resolve(); plain equality still dedupes.
|
||||
distinct = root != primary
|
||||
return [primary, root] if distinct else [primary]
|
||||
|
||||
|
||||
def is_engaged() -> bool:
|
||||
"""True if ANY candidate sentinel exists; fail SAFE (True) on stat errors."""
|
||||
saw_stat_error = False
|
||||
for path in _candidate_sentinel_paths():
|
||||
try:
|
||||
if path.exists():
|
||||
return True
|
||||
except OSError:
|
||||
saw_stat_error = True
|
||||
return saw_stat_error
|
||||
|
||||
|
||||
def engage(reason: Optional[str] = None) -> Path:
|
||||
"""Create the ESTOP sentinel. Idempotent; re-engaging updates the file."""
|
||||
path = sentinel_path()
|
||||
payload = {
|
||||
"engaged_at": datetime.now(timezone.utc).isoformat(),
|
||||
"reason": reason or None,
|
||||
}
|
||||
payload = {"engaged_at": datetime.now(timezone.utc).isoformat(), "reason": reason or None}
|
||||
try:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
|
||||
except OSError:
|
||||
# Best effort: an empty/partial sentinel still pauses (fail safe).
|
||||
try:
|
||||
with suppress(OSError): # Best effort: an empty/partial sentinel still pauses (fail safe).
|
||||
path.touch(exist_ok=True)
|
||||
except OSError:
|
||||
pass
|
||||
return path
|
||||
|
||||
|
||||
def disengage() -> bool:
|
||||
"""Remove the ESTOP sentinel. Returns True if a pause was lifted."""
|
||||
try:
|
||||
sentinel_path().unlink()
|
||||
return True
|
||||
except FileNotFoundError:
|
||||
return False
|
||||
except OSError:
|
||||
return False
|
||||
"""Remove every visible sentinel (process-local and fleet-root)."""
|
||||
lifted = False
|
||||
for path in _candidate_sentinel_paths():
|
||||
try:
|
||||
path.unlink()
|
||||
lifted = True
|
||||
except (OSError, AttributeError):
|
||||
continue
|
||||
return lifted
|
||||
|
||||
|
||||
def get_state() -> Optional[dict]:
|
||||
"""Return ``{"reason": ..., "engaged_at": ...}`` or None when not engaged.
|
||||
|
||||
A sentinel with an unreadable/corrupt body still reports engaged, with
|
||||
both fields None — the pause is authoritative, the metadata is not.
|
||||
"""
|
||||
path = sentinel_path()
|
||||
if not path.exists():
|
||||
"""Return ``{"reason", "engaged_at"}`` or None when not engaged; an unreadable/corrupt
|
||||
body still reports engaged with both fields None."""
|
||||
if not is_engaged():
|
||||
return None
|
||||
reason = None
|
||||
engaged_at = None
|
||||
try:
|
||||
raw = json.loads(path.read_text(encoding="utf-8"))
|
||||
if isinstance(raw, dict):
|
||||
reason = raw.get("reason") or None
|
||||
engaged_at = raw.get("engaged_at") or None
|
||||
except (OSError, ValueError):
|
||||
pass
|
||||
return {"reason": reason, "engaged_at": engaged_at}
|
||||
state = {"reason": None, "engaged_at": None}
|
||||
found = False
|
||||
for path in _candidate_sentinel_paths():
|
||||
try:
|
||||
if not path.exists():
|
||||
continue
|
||||
except OSError:
|
||||
return state
|
||||
except AttributeError:
|
||||
continue
|
||||
found = True
|
||||
with suppress(OSError, ValueError, AttributeError):
|
||||
raw = json.loads(path.read_text(encoding="utf-8"))
|
||||
if isinstance(raw, dict):
|
||||
state = {"reason": raw.get("reason") or None, "engaged_at": raw.get("engaged_at") or None}
|
||||
break
|
||||
return state if found else None
|
||||
|
||||
|
||||
def paused_reply() -> Optional[str]:
|
||||
@@ -127,48 +115,32 @@ def paused_reply() -> Optional[str]:
|
||||
state = get_state()
|
||||
if state is None:
|
||||
return None
|
||||
reason = state.get("reason")
|
||||
if reason:
|
||||
return (
|
||||
f"⏸️ Hermes is paused ({reason}). New work is on hold; "
|
||||
"run `hermes resume` to pick things back up."
|
||||
)
|
||||
return (
|
||||
"⏸️ Hermes is paused. New work is on hold; "
|
||||
"run `hermes resume` to pick things back up."
|
||||
)
|
||||
tag = f" ({state['reason']})" if state.get("reason") else ""
|
||||
return f"⏸️ Hermes is paused{tag}. New work is on hold; run `hermes resume` to pick things back up."
|
||||
|
||||
|
||||
def check_paused(component: str, logger: logging.Logger) -> bool:
|
||||
"""Return True when engaged, logging once per engagement per component.
|
||||
|
||||
Dispatch loops call this every tick; the log fires on the disengaged→
|
||||
engaged transition for that component and re-arms after a resume, so a
|
||||
long pause doesn't spam one line per tick.
|
||||
"""
|
||||
"""Return True when engaged, logging once per engagement per component (re-armed after a resume)."""
|
||||
if not is_engaged():
|
||||
with _log_lock:
|
||||
_logged_components.discard(component)
|
||||
return False
|
||||
with _log_lock:
|
||||
first = component not in _logged_components
|
||||
if first:
|
||||
_logged_components.add(component)
|
||||
_logged_components.add(component)
|
||||
if first:
|
||||
state = get_state() or {}
|
||||
reason = state.get("reason")
|
||||
reason = (get_state() or {}).get("reason")
|
||||
suffix = f" (reason: {reason})" if reason else ""
|
||||
logger.info(
|
||||
"%s dispatch paused by global emergency stop%s — remove with "
|
||||
"`hermes resume` (%s)",
|
||||
component,
|
||||
suffix,
|
||||
sentinel_path(),
|
||||
"%s dispatch paused by global emergency stop%s — remove with `hermes resume` (%s)",
|
||||
component, suffix, sentinel_path(),
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def _reset_log_state_for_tests() -> None:
|
||||
"""Clear the log-once bookkeeping (test isolation helper)."""
|
||||
with _log_lock:
|
||||
_logged_components.clear()
|
||||
# ---- BEGIN PLUGIN-COMPAT (revert-scheduled; see COMPAT_MANIFEST.md) ----
|
||||
# Names external plugins imported from this module before the Sep 2026 decomposition.
|
||||
# Internal code MUST NOT use these (scripts/check_compat_pointers.py fails CI if it does).
|
||||
# The whole block is removed by reverting the commit that added it.
|
||||
import os # noqa: F401,E402
|
||||
# ---- END PLUGIN-COMPAT ----
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
"""Primary rate-limit cooldown arming, shared by fallback switches."""
|
||||
import logging
|
||||
import time
|
||||
|
||||
from agent.error_classifier import FailoverReason
|
||||
|
||||
_RATE_LIMIT_FAILOVER_REASONS = frozenset({FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.upstream_rate_limit})
|
||||
|
||||
|
||||
def _arm_rate_limit_cooldown(agent, reason: "FailoverReason | None") -> int | None:
|
||||
"""Arm the primary's exponential cooldown (60s → 2m → ... → 4h cap) on CONSECUTIVE rate-limits;
|
||||
restore_primary_runtime resets the counter. Only when leaving the primary: chain-switching from
|
||||
an active fallback means the primary was not the 429 source, so its cooldown is left alone.
|
||||
Return the armed cooldown in seconds, or None when no cooldown was armed."""
|
||||
if reason not in _RATE_LIMIT_FAILOVER_REASONS:
|
||||
return None
|
||||
current_provider = (getattr(agent, "provider", "") or "").strip().lower()
|
||||
primary_provider = ((agent._primary_runtime or {}).get("provider") or "").strip().lower()
|
||||
if getattr(agent, "_fallback_activated", False) and not (primary_provider and current_provider == primary_provider):
|
||||
return None
|
||||
backoff_count = getattr(agent, "_rate_limit_backoff_count", 0)
|
||||
agent._rate_limit_backoff_count = backoff_count + 1
|
||||
backoff_seconds = min(60 * (2 ** backoff_count), 14400)
|
||||
agent._rate_limited_until = time.monotonic() + backoff_seconds
|
||||
logging.info("Rate-limit backoff level %d: cooldown %d s (%.1f min, backoff#%d)", backoff_count, backoff_seconds, backoff_seconds / 60, backoff_count + 1)
|
||||
return backoff_seconds
|
||||
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
"""Bounded fast-mode windows (``/fast auto`` and ``/fast cold``).
|
||||
|
||||
``agent.service_tier``: ``None`` (normal), ``"priority"`` (static fast, pinned into
|
||||
``agent.request_overrides`` at build time), ``"auto"`` (every user turn opens a
|
||||
window of ``agent.fast_auto_seconds``) or ``"cold"`` (only a session's first turn,
|
||||
no prior history, opens it). The provider's fast override is layered onto request
|
||||
kwargs only while the window is open; only per-request params (``service_tier`` /
|
||||
``speed``) vary, so the prompt cache survives the boundary.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
BOUNDED_MODES = frozenset({"auto", "cold"})
|
||||
DEFAULT_WINDOW_SECONDS = 60
|
||||
|
||||
|
||||
def begin_turn(agent: Any, conversation_history: Any) -> None:
|
||||
"""Open (or refuse) the fast window at a user-turn boundary."""
|
||||
mode = getattr(agent, "service_tier", None)
|
||||
agent._fast_until = 0.0
|
||||
if mode not in BOUNDED_MODES:
|
||||
return
|
||||
if mode == "cold" and any(
|
||||
isinstance(m, dict) and m.get("role") in ("user", "assistant", "tool")
|
||||
for m in (conversation_history or ())
|
||||
):
|
||||
return
|
||||
try:
|
||||
window = float(getattr(agent, "fast_auto_seconds", DEFAULT_WINDOW_SECONDS))
|
||||
except (TypeError, ValueError):
|
||||
window = DEFAULT_WINDOW_SECONDS
|
||||
agent._fast_until = time.monotonic() + max(window, 0.0)
|
||||
|
||||
|
||||
def effective_request_overrides(agent: Any) -> dict[str, Any]:
|
||||
"""``agent.request_overrides`` plus the fast override while the window is open."""
|
||||
overrides = dict(getattr(agent, "request_overrides", None) or {})
|
||||
if getattr(agent, "service_tier", None) not in BOUNDED_MODES or time.monotonic() >= getattr(agent, "_fast_until", 0.0):
|
||||
return overrides
|
||||
from hermes_cli.models import resolve_fast_mode_overrides
|
||||
base_url = getattr(agent, "base_url", None)
|
||||
if getattr(agent, "api_mode", None) == "anthropic_messages":
|
||||
base_url = getattr(agent, "_anthropic_base_url", None) or base_url
|
||||
overrides.update(
|
||||
resolve_fast_mode_overrides(getattr(agent, "model", None), provider=getattr(agent, "provider", None), base_url=base_url) or {}
|
||||
)
|
||||
return overrides
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user