diff --git a/.config/nextest.toml b/.config/nextest.toml new file mode 100644 index 000000000..ea686cd7c --- /dev/null +++ b/.config/nextest.toml @@ -0,0 +1,6 @@ +# 不固定 test-threads:nextest 默认按 num-cpus 并发,固定值会在更大规格的 +# runner 或本地开发机上主动压低并发、反而变慢,且无法表达 min(4, num-cpus)。 +# 这里只保留卡死保护,避免单个挂起用例拖满整个 job。 +[profile.default] +# 60 秒后标记慢测试,连续两轮仍未结束则终止;超时结果保持失败,不隐藏回归。 +slow-timeout = { period = "60s", terminate-after = 2, grace-period = "10s" } diff --git a/.github/workflows/nightly.yml b/.github/workflows/nightly.yml index 109a38e6c..802663651 100644 --- a/.github/workflows/nightly.yml +++ b/.github/workflows/nightly.yml @@ -63,6 +63,8 @@ jobs: name: Rust CI needs: source uses: ./.github/workflows/rust-ci.yml + with: + full_scope: true rust_extended: name: Rust extended checks diff --git a/.github/workflows/rust-ci.yml b/.github/workflows/rust-ci.yml index a01af779f..a4fa22eb0 100644 --- a/.github/workflows/rust-ci.yml +++ b/.github/workflows/rust-ci.yml @@ -2,6 +2,12 @@ name: Rust CI on: workflow_call: + inputs: + full_scope: + description: "Run all Rust and shell scopes, used by Nightly" + required: false + type: boolean + default: false push: branches: - master @@ -9,8 +15,11 @@ on: paths: - "Cargo.toml" - "Cargo.lock" + - "rust-toolchain.toml" + - ".cargo/**" - "crates/**" - "apps/**" + - "*.sql" - "install.sh" - "deploy.sh" - "update.sh" @@ -28,17 +37,17 @@ on: - "tests/update_*_test.sh" - "tests/release_supply_chain_test.sh" - "tests/tunnel_installer_config_security_test.sh" - - ".github/workflows/build-tunnel.yml" - - ".github/workflows/deploy-pages.yml" - - ".github/workflows/release.yml" - - ".github/workflows/rust-ci.yml" - - ".github/workflows/nightly.yml" + - ".github/workflows/*.yml" + - ".github/workflows/*.yaml" pull_request: paths: - "Cargo.toml" - "Cargo.lock" + - "rust-toolchain.toml" + - ".cargo/**" - "crates/**" - "apps/**" + - "*.sql" - "install.sh" - "deploy.sh" - "update.sh" @@ -56,11 +65,8 @@ on: - "tests/update_*_test.sh" - "tests/release_supply_chain_test.sh" - "tests/tunnel_installer_config_security_test.sh" - - ".github/workflows/build-tunnel.yml" - - ".github/workflows/deploy-pages.yml" - - ".github/workflows/release.yml" - - ".github/workflows/rust-ci.yml" - - ".github/workflows/nightly.yml" + - ".github/workflows/*.yml" + - ".github/workflows/*.yaml" concurrency: group: rust-ci-${{ github.event_name }}-${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} @@ -76,8 +82,77 @@ env: CARGO_TERM_COLOR: always jobs: + changes: + name: Detect Rust CI scope + runs-on: ubuntu-latest + outputs: + rust: ${{ steps.scope.outputs.rust }} + shell: ${{ steps.scope.outputs.shell }} + steps: + - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 + with: + fetch-depth: 0 + + - name: Classify changed paths + id: scope + shell: bash + env: + RUST_CI_FULL_SCOPE: ${{ inputs.full_scope || false }} + run: | + # 任何命令失败都必须让本 job 失败,否则 git fetch/diff 出错后仍会写出 + # rust=false/shell=false,下游会误判为“无需测试”而假绿放行。 + set -euo pipefail + + # Nightly 通过 workflow_call 显式传入 full_scope;普通 push/PR 只按源码和构建 + # 指纹触发 Rust jobs,安装脚本、Compose、README 等由 shell scope 覆盖。 + if [ "$RUST_CI_FULL_SCOPE" = "true" ]; then + echo "rust=true" >> "$GITHUB_OUTPUT" + echo "shell=true" >> "$GITHUB_OUTPUT" + exit 0 + fi + + if [ "$GITHUB_EVENT_NAME" = "pull_request" ] \ + && [ -n "${GITHUB_BASE_REF:-}" ] \ + && [ -n "${GITHUB_SHA:-}" ]; then + git fetch --no-tags origin "$GITHUB_BASE_REF" --depth=1 + changed_paths=$(git diff --name-only "origin/$GITHUB_BASE_REF...$GITHUB_SHA") + elif [ "$GITHUB_EVENT_NAME" = "push" ] \ + && [ -n "${GITHUB_EVENT_BEFORE:-}" ] \ + && [ "$GITHUB_EVENT_BEFORE" != "0000000000000000000000000000000000000000" ] \ + && [ -n "${GITHUB_SHA:-}" ]; then + changed_paths=$(git diff --name-only "$GITHUB_EVENT_BEFORE" "$GITHUB_SHA") + else + changed_paths=$(git ls-files) + fi + + # 防御性兜底:diff 结果为空(异常事件或比较失败)时按全量运行, + # 宁可多跑也不能漏测。 + if [ -z "$changed_paths" ]; then + echo "rust=true" >> "$GITHUB_OUTPUT" + echo "shell=true" >> "$GITHUB_OUTPUT" + exit 0 + fi + + rust=false + shell=false + while IFS= read -r path; do + case "$path" in + Cargo.toml|Cargo.lock|rust-toolchain.toml|.cargo/*|*.rs|*/Cargo.toml|*/build.rs|*.sql|.github/workflows/*.yml|.github/workflows/*.yaml) + rust=true + ;; + *.sh|*.py|README.md|*/README.md|.env.example|Dockerfile*|docker-compose*.yml|docker-compose*.yaml) + shell=true + ;; + esac + done <<< "$changed_paths" + + echo "rust=$rust" >> "$GITHUB_OUTPUT" + echo "shell=$shell" >> "$GITHUB_OUTPUT" + shell_security: name: Shell security fixtures + needs: changes + if: ${{ needs.changes.outputs.shell == 'true' }} runs-on: ubuntu-latest steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 @@ -99,6 +174,8 @@ jobs: fmt: name: Format + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 @@ -114,6 +191,8 @@ jobs: clippy_gateway: name: Clippy (Gateway) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 @@ -127,7 +206,9 @@ jobs: - name: Rust cache uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: - shared-key: rust-ci-${{ runner.os }} + # Gateway lint 与 Gateway 测试都可能触发 mold/大型链接依赖,单独隔离缓存 + # 指纹,避免不同 job 的构建产物互相驱逐或复用错误的链接参数。 + shared-key: rust-ci-gateway-clippy-${{ runner.os }} workspaces: . -> target - name: Setup sccache @@ -148,6 +229,8 @@ jobs: clippy_data: name: Clippy (Data) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 @@ -182,6 +265,8 @@ jobs: clippy_rest: name: Clippy (Workspace Rest) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 @@ -218,6 +303,7 @@ jobs: name: Clippy runs-on: ubuntu-latest needs: + - changes - clippy_gateway - clippy_data - clippy_rest @@ -225,6 +311,10 @@ jobs: steps: - name: Verify clippy jobs run: | + if [ "${{ needs.changes.outputs.rust }}" != "true" ]; then + echo "Rust scope unchanged; clippy jobs skipped" + exit 0 + fi if [ "${{ needs.clippy_gateway.result }}" != "success" ] || \ [ "${{ needs.clippy_data.result }}" != "success" ] || \ [ "${{ needs.clippy_rest.result }}" != "success" ]; then @@ -234,12 +324,24 @@ jobs: test_gateway: name: Test (Gateway) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest + # 构建指纹提到 job 级:mold RUSTFLAGS / 栈 / sccache 对 lib、bins、integration 三步保持一致, + # 避免 step 级 env 漂移导致同 job 内 rustc 指纹不一致。 + env: + RUSTC_WRAPPER: sccache + SCCACHE_GHA_ENABLED: "true" + RUST_MIN_STACK: "16777216" + RUSTFLAGS: "-C link-arg=-fuse-ld=mold" steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable + with: + # 与 rust-toolchain.toml、fmt/clippy 钉在同一版本,避免浮动 stable 换指纹导致全量重编 + toolchain: 1.95.0 - name: Show Rust toolchain run: rustup show active-toolchain @@ -247,7 +349,8 @@ jobs: - name: Rust cache uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 with: - shared-key: rust-ci-${{ runner.os }} + # mold RUSTFLAGS 只在本 job 生效:独立 cache key,避免与无 mold 的 job 互相污染指纹 + shared-key: rust-ci-gateway-test-${{ runner.os }} workspaces: . -> target - name: Setup sccache @@ -263,36 +366,34 @@ jobs: run: pg_config --bindir >> "$GITHUB_PATH" - name: Test lib - env: - RUSTC_WRAPPER: sccache - SCCACHE_GHA_ENABLED: "true" - RUST_MIN_STACK: "16777216" - RUSTFLAGS: "-C link-arg=-fuse-ld=mold" run: cargo nextest run -p aether-gateway --lib - name: Test bins - env: - RUSTC_WRAPPER: sccache - SCCACHE_GHA_ENABLED: "true" - RUST_MIN_STACK: "16777216" - RUSTFLAGS: "-C link-arg=-fuse-ld=mold" run: cargo nextest run -p aether-gateway --bins + # 只运行独立 integration targets;显式列出目标,避免 --tests 再次执行 lib/bin 测试。 + - name: Test integration targets + run: >- + cargo nextest run -p aether-gateway + --test admin_unsigned_identity_headers + --test architecture_guard + - name: Show sccache stats if: always() - env: - RUSTC_WRAPPER: sccache - SCCACHE_GHA_ENABLED: "true" run: sccache --show-stats test_data: name: Test (Data) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable + with: + toolchain: 1.95.0 - name: Show Rust toolchain run: rustup show active-toolchain @@ -328,6 +429,8 @@ jobs: check_data_features: name: Check (Data Feature - ${{ matrix.feature }}) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest strategy: fail-fast: false @@ -340,6 +443,8 @@ jobs: - name: Install Rust toolchain uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable + with: + toolchain: 1.95.0 - name: Rust cache uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 @@ -365,12 +470,16 @@ jobs: test_rest: name: Test (Workspace Rest) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable + with: + toolchain: 1.95.0 - name: Show Rust toolchain run: rustup show active-toolchain @@ -402,6 +511,8 @@ jobs: test_data_adapters: name: Test (Data Adapter - ${{ matrix.package }}) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest strategy: fail-fast: false @@ -413,6 +524,8 @@ jobs: - name: Install Rust toolchain uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable + with: + toolchain: 1.95.0 - name: Rust cache uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 @@ -441,12 +554,16 @@ jobs: check_integration_scenarios: name: Test (Integration Scenarios) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest steps: - uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5 - name: Install Rust toolchain uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable + with: + toolchain: 1.95.0 - name: Rust cache uses: Swatinem/rust-cache@49a0bdc70d2e1b713ca9e2869b211fcce03d3c1c # v2 @@ -477,6 +594,7 @@ jobs: name: Test runs-on: ubuntu-latest needs: + - changes - test_gateway - test_data - check_data_features @@ -487,6 +605,10 @@ jobs: steps: - name: Verify test jobs run: | + if [ "${{ needs.changes.outputs.rust }}" != "true" ]; then + echo "Rust scope unchanged; test jobs skipped" + exit 0 + fi if [ "${{ needs.test_gateway.result }}" != "success" ] || \ [ "${{ needs.test_data.result }}" != "success" ] || \ [ "${{ needs.check_data_features.result }}" != "success" ] || \ @@ -499,6 +621,8 @@ jobs: data_db_smoke_postgres: name: Data DB Smoke (Postgres) + needs: changes + if: ${{ needs.changes.outputs.rust == 'true' }} runs-on: ubuntu-latest services: postgres: @@ -519,6 +643,8 @@ jobs: - name: Install Rust toolchain uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable + with: + toolchain: 1.95.0 - name: Show Rust toolchain run: rustup show active-toolchain @@ -556,6 +682,13 @@ jobs: AETHER_TEST_DATABASE_URL: postgres://aether:aether@127.0.0.1:5432/aether_test run: cargo test -p aether-data-postgres live_payment_callback --lib -- --ignored --nocapture + - name: Run Postgres batch wallet deduction regression + env: + RUSTC_WRAPPER: sccache + SCCACHE_GHA_ENABLED: "true" + AETHER_TEST_DATABASE_URL: postgres://aether:aether@127.0.0.1:5432/aether_test + run: cargo test -p aether-data-postgres live_bulk_wallet_adjustment_persists_actual_delta_and_skips_zero_ledger --lib -- --ignored --nocapture + - name: Run Postgres API key lifecycle tests env: RUSTC_WRAPPER: sccache @@ -586,11 +719,20 @@ jobs: name: Data DB Smoke runs-on: ubuntu-latest needs: + - changes - data_db_smoke_postgres if: ${{ always() }} steps: - name: Verify database smoke jobs run: | + if [ "${{ needs.changes.result }}" != "success" ]; then + echo "Scope detection failed" + exit 1 + fi + if [ "${{ needs.changes.outputs.rust }}" != "true" ]; then + echo "Rust scope unchanged; database smoke jobs skipped" + exit 0 + fi if [ "${{ needs.data_db_smoke_postgres.result }}" != "success" ]; then echo "Data DB smoke failed" exit 1 @@ -600,6 +742,7 @@ jobs: name: check runs-on: ubuntu-latest needs: + - changes - fmt - clippy - test @@ -609,11 +752,30 @@ jobs: steps: - name: Verify required jobs run: | - if [ "${{ needs.fmt.result }}" != "success" ] || \ - [ "${{ needs.clippy.result }}" != "success" ] || \ - [ "${{ needs.test.result }}" != "success" ] || \ - [ "${{ needs.data_db_smoke.result }}" != "success" ] || \ - [ "${{ needs.shell_security.result }}" != "success" ]; then + # changes 失败或未产出 scope 时不允许直接放行,避免假绿。 + if [ "${{ needs.changes.result }}" != "success" ]; then + echo "Scope detection failed" + exit 1 + fi + rust="${{ needs.changes.outputs.rust }}" + shell="${{ needs.changes.outputs.shell }}" + + if [ "$rust" != "true" ] && [ "$shell" != "true" ]; then + echo "No Rust or shell scope changed" + exit 0 + fi + + if [ "$rust" = "true" ] && { + [ "${{ needs.fmt.result }}" != "success" ] || + [ "${{ needs.clippy.result }}" != "success" ] || + [ "${{ needs.test.result }}" != "success" ] || + [ "${{ needs.data_db_smoke.result }}" != "success" ]; + }; then + echo "Rust CI failed" + exit 1 + fi + + if [ "$shell" = "true" ] && [ "${{ needs.shell_security.result }}" != "success" ]; then echo "Rust CI failed" exit 1 fi diff --git a/Cargo.lock b/Cargo.lock index ab6a7e347..5d6742d56 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -431,11 +431,14 @@ dependencies = [ "aether-runtime", "aether-runtime-state", "aether-testkit", + "aether-tunnel", + "arc-swap", "async-stream", "axum", "futures-util", "http", "reqwest 0.12.28", + "rustls", "serde", "serde_json", "sha2", @@ -526,6 +529,7 @@ dependencies = [ "aether-data-contracts", "aether-pool-core", "aether-provider-transport", + "chrono", "serde_json", "url", "uuid", @@ -673,7 +677,6 @@ name = "aether-tunnel" version = "0.3.17" dependencies = [ "aether-contracts", - "aether-gateway", "aether-gateway-tunnel", "aether-http", "aether-runtime", @@ -750,6 +753,7 @@ dependencies = [ "async-trait", "serde", "serde_json", + "sha2", "url", "uuid", ] diff --git a/Cargo.toml b/Cargo.toml index 96272433b..9bc1273aa 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -92,6 +92,7 @@ aether-usage-core = { path = "crates/aether-usage/core" } aether-usage-runtime = { path = "crates/aether-usage/runtime" } aether-video-tasks-core = { path = "crates/aether-video-tasks-core" } aether-gateway = { path = "apps/aether-gateway" } +aether-tunnel = { path = "apps/aether-tunnel" } aether-http = { path = "crates/aether-http" } aether-runtime = { path = "crates/aether-runtime/base" } aether-testkit = { path = "crates/aether-testing/testkit" } diff --git a/apps/aether-gateway/.aether-windsurf-binary-test-03d740bc-64f1-49cb-b13d-4e5234e5de31/language-server b/apps/aether-gateway/.aether-windsurf-binary-test-03d740bc-64f1-49cb-b13d-4e5234e5de31/language-server new file mode 100755 index 000000000..f8b18e900 --- /dev/null +++ b/apps/aether-gateway/.aether-windsurf-binary-test-03d740bc-64f1-49cb-b13d-4e5234e5de31/language-server @@ -0,0 +1 @@ +test binary \ No newline at end of file diff --git a/apps/aether-gateway/src/ai_serving/api.rs b/apps/aether-gateway/src/ai_serving/api.rs index 9ff10edec..3225fa7af 100644 --- a/apps/aether-gateway/src/ai_serving/api.rs +++ b/apps/aether-gateway/src/ai_serving/api.rs @@ -69,9 +69,14 @@ pub(crate) use aether_ai_formats::api::{ OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND, }; pub(crate) use aether_ai_formats::protocol::stream::CanonicalUsage as StreamingCanonicalUsage; -/// Codex client identity headers re-exported for out-of-crate probe binaries, -/// which must reach `aether_ai_formats` through this seam. -pub use aether_ai_formats::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT}; +/// Codex client identity accessors re-exported for out-of-crate probe binaries, +/// which must reach the runtime profile through this seam. +pub use aether_ai_formats::{codex_client_originator, codex_client_user_agent}; +/// Codex 动态客户端画像 API 只允许经此根缝进入 gateway,避免其它模块直接依赖 formats crate。 +pub(crate) use aether_ai_formats::{ + codex_client_profile, codex_client_version, set_codex_cli_version, set_codex_client_profile, + CodexClientProfile, +}; pub(crate) use aether_ai_formats::{CODEX_RESPONSES_LITE_HEADER, UPSTREAM_IS_STREAM_KEY}; pub(crate) fn parse_direct_request_body( diff --git a/apps/aether-gateway/src/ai_serving/finalize/internal/stream_rewrite.rs b/apps/aether-gateway/src/ai_serving/finalize/internal/stream_rewrite.rs index c1b8c0425..4355437ea 100644 --- a/apps/aether-gateway/src/ai_serving/finalize/internal/stream_rewrite.rs +++ b/apps/aether-gateway/src/ai_serving/finalize/internal/stream_rewrite.rs @@ -18,6 +18,12 @@ pub(crate) fn maybe_build_local_stream_rewriter<'a>( } impl LocalStreamRewriter<'_> { + pub(crate) fn into_owned(self) -> LocalStreamRewriter<'static> { + LocalStreamRewriter { + inner: self.inner.into_owned(), + } + } + pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result, GatewayError> { self.inner.push_chunk(chunk).map_err(map_surface_error) } diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_resolution.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_resolution.rs index abac4c3e6..1c6b92f20 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_resolution.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_resolution.rs @@ -93,6 +93,16 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> { &self, candidate: Self::Candidate, ) -> Self::Skipped { + warn!( + event_name = "local_candidate_skipped", + log_type = "event", + provider_id = %candidate.provider_id, + endpoint_id = %candidate.endpoint_id, + key_id = %candidate.key_id, + api_format = %candidate.endpoint_api_format, + skip_reason = "transport_snapshot_missing", + "local execution candidate skipped during planning" + ); SkippedLocalExecutionCandidate { candidate, skip_reason: "transport_snapshot_missing", @@ -145,6 +155,16 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> { transport: Self::Transport, skip_reason: &'static str, ) -> Self::Skipped { + warn!( + event_name = "local_candidate_skipped", + log_type = "event", + provider_id = %candidate.provider_id, + endpoint_id = %candidate.endpoint_id, + key_id = %candidate.key_id, + api_format = %candidate.endpoint_api_format, + skip_reason, + "local execution candidate skipped during planning" + ); SkippedLocalExecutionCandidate { candidate, skip_reason, diff --git a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs index 41008896e..4bd50da98 100644 --- a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs @@ -314,6 +314,20 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts( return Ok(None); } + // Same-format requests skip `apply_transport_request_body_semantics`, so the opt-in + // Claude Code body mimicry has to be applied here as well. + if crate::ai_serving::transport::claude_code::apply_claude_code_body_mimicry_for_transport( + &mut base_provider_request_body, + &transport, + prepared.provider_api_format.as_str(), + ) { + compatibility_edits.push(SameFormatProviderCompatibilityEdit { + field: "body".to_string(), + action: SameFormatProviderCompatibilityEditAction::ProviderCompatibilityRewrite, + detail: "applied Claude Code body mimicry for provider compatibility".to_string(), + }); + } + let antigravity_auth = if prepared.is_antigravity { let mut antigravity_support = classify_local_antigravity_request_support( &transport, @@ -583,6 +597,11 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts( source_model, codex_model_capabilities.as_ref(), ); + crate::ai_serving::transport::xai::insert_cli_identity_headers_if_needed( + transport.as_ref(), + prepared.provider_api_format.as_str(), + &mut provider_request_headers, + ); request_identity_response_encoding_when_redacted( &mut provider_request_headers, redaction.redacted, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/image/request.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/image/request.rs index 3ca11f0d2..be3ff2fc0 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/image/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/image/request.rs @@ -17,8 +17,8 @@ use crate::ai_serving::transport::{ ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH, }; use crate::ai_serving::{ - apply_codex_openai_special_headers, build_chatgpt_web_image_request_body, - build_codex_openai_image_api_provider_request_body, + apply_codex_openai_special_headers, apply_xai_upstream_payload_edits, + build_chatgpt_web_image_request_body, build_codex_openai_image_api_provider_request_body, build_gemini_image_request_body_from_openai_image_request, build_openai_image_api_provider_request_body, build_openai_image_provider_request_body, default_model_for_openai_image_operation, normalize_openai_image_request, @@ -211,7 +211,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts( upstream_is_stream, ) }; - let Some(provider_request_body) = provider_request_body else { + let Some(mut provider_request_body) = provider_request_body else { mark_skipped_local_openai_image_candidate_with_failure_diagnostic( state, input, @@ -229,6 +229,11 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts( .await; return None; }; + apply_xai_upstream_payload_edits( + &mut provider_request_body, + transport.provider.provider_type.as_str(), + provider_api_format, + ); let Some(mut provider_request_headers) = (if is_grok { build_grok_browser_headers(GrokHeaderInput { transport, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs index c1207443e..8120ca3d7 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs @@ -8,6 +8,7 @@ use crate::ai_serving::planner::{ build_ai_execution_decision_response, resolve_transport_request_encoding_policy, AiExecutionDecisionResponseParts, }; +use crate::ai_serving::transport::xai::video::is_native_video_request; use crate::ai_serving::transport::{ resolve_transport_execution_timeouts, resolve_transport_profile, }; @@ -33,7 +34,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat let Some(resolved) = resolve_local_video_create_candidate_payload_parts( state, parts, body_json, trace_id, input, &attempt, spec, ) - .await + .await? else { return Ok(None); }; @@ -52,9 +53,32 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat .await; let transport_profile = resolve_transport_profile(&transport); let mut extra_fields = serde_json::Map::new(); + if is_native_video_request(&transport.provider.provider_type, parts.uri.path()) { + extra_fields.insert( + "video_client_protocol".to_string(), + serde_json::json!("xai"), + ); + } + if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) { extra_fields.insert("proxy".to_string(), proxy_value); } + if transport.provider.provider_type.eq_ignore_ascii_case("xai") { + extra_fields.insert("video_provider_xai".into(), serde_json::json!(true)); + if let Some(duration) = resolved.provider_request_body.get("duration") { + extra_fields.insert("video_duration".into(), duration.clone()); + } + if parts.uri.path() == "/openai/v1/videos" { + extra_fields.insert( + "video_size".into(), + body_json + .get("size") + .filter(|v| v.as_str().is_some_and(|s| !s.trim().is_empty())) + .cloned() + .unwrap_or_else(|| serde_json::json!("720x1280")), + ); + } + } let effective_headers = input.effective_headers(&parts.headers); let report_context = build_local_execution_report_context(LocalExecutionReportContextParts { auth_context: &input.auth_context, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/video/request.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/video/request.rs index 795384f94..5910ba566 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/video/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/video/request.rs @@ -3,15 +3,23 @@ use std::sync::Arc; use serde_json::Value; -use crate::ai_serving::planner::candidate_preparation::resolve_candidate_mapped_model; +use crate::ai_serving::planner::candidate_preparation::{ + prepare_header_authenticated_candidate, resolve_candidate_mapped_model, OauthPreparationContext, +}; use crate::ai_serving::planner::spec_metadata::local_video_create_spec_metadata; +use crate::ai_serving::transport::xai::video::{ + convert_openai_video_request, is_explicit_native_video_path, is_native_video_request, +}; use crate::ai_serving::transport::{ build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url, resolve_video_create_auth, video_create_transport_unsupported_reason, ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, }; -use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot}; -use crate::AppState; +use crate::ai_serving::{ + apply_xai_upstream_payload_edits, CandidateFailureDiagnostic, GatewayProviderTransportSnapshot, + PlannerAppState, +}; +use crate::{AppState, GatewayError}; use super::support::{ mark_skipped_local_video_candidate, mark_skipped_local_video_candidate_with_failure_diagnostic, @@ -37,11 +45,16 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( input: &LocalVideoCreateDecisionInput, attempt: &LocalVideoCreateCandidateAttempt, spec: LocalVideoCreateSpec, -) -> Option { +) -> Result, GatewayError> { let spec_metadata = local_video_create_spec_metadata(spec); let candidate = &attempt.eligible.candidate; let transport = &attempt.eligible.transport; let effective_headers = input.effective_headers(&parts.headers); + if is_explicit_native_video_path(parts.uri.path()) + && !transport.provider.provider_type.eq_ignore_ascii_case("xai") + { + return Ok(None); + } let provider_family = provider_video_create_family(spec.family); let transport_unsupported_reason = video_create_transport_unsupported_reason( @@ -60,23 +73,39 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( skip_reason, ) .await; - return None; + return Ok(None); } - let auth = resolve_video_create_auth(transport, provider_family); - let Some((auth_header, auth_value)) = auth else { - mark_skipped_local_video_candidate( - state, - input, + let prepared_candidate = match prepare_header_authenticated_candidate( + PlannerAppState::new(state), + transport, + candidate, + resolve_video_create_auth(transport, provider_family), + OauthPreparationContext { trace_id, - candidate, - attempt.candidate_index, - &attempt.candidate_id, - "transport_auth_unavailable", - ) - .await; - return None; + api_format: spec_metadata.api_format, + operation: "video_create_candidate_request", + }, + ) + .await + { + Ok(prepared) => prepared, + Err(skip_reason) => { + mark_skipped_local_video_candidate( + state, + input, + trace_id, + candidate, + attempt.candidate_index, + &attempt.candidate_id, + skip_reason, + ) + .await; + return Ok(None); + } }; + let auth_header = prepared_candidate.auth_header; + let auth_value = prepared_candidate.auth_value; let mapped_model = match resolve_candidate_mapped_model(candidate) { Ok(mapped_model) => mapped_model, @@ -91,7 +120,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( skip_reason, ) .await; - return None; + return Ok(None); } }; @@ -117,10 +146,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( ), ) .await; - return None; + return Ok(None); }; - let Some(provider_request_body) = build_video_create_request_body( + let Some(mut provider_request_body) = build_video_create_request_body( body_json, provider_family, &mapped_model, @@ -142,11 +171,28 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( ), ) .await; - return None; + return Ok(None); }; + if transport.provider.provider_type.eq_ignore_ascii_case("xai") + && !is_native_video_request(&transport.provider.provider_type, parts.uri.path()) + { + provider_request_body = + convert_openai_video_request(&provider_request_body).map_err(|message| { + GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + message: message.to_string(), + } + })?; + } + apply_xai_upstream_payload_edits( + &mut provider_request_body, + transport.provider.provider_type.as_str(), + spec_metadata.api_format, + ); let Some(provider_request_headers) = build_video_create_headers(ProviderVideoCreateHeadersInput { + transport, headers: effective_headers, auth_header: &auth_header, auth_value: &auth_value, @@ -170,10 +216,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( ), ) .await; - return None; + return Ok(None); }; - Some(LocalVideoCreateCandidatePayloadParts { + Ok(Some(LocalVideoCreateCandidatePayloadParts { transport: Arc::clone(transport), auth_header, auth_value, @@ -181,7 +227,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( provider_request_headers, provider_request_body, upstream_url, - }) + })) } fn provider_video_create_family(family: LocalVideoCreateFamily) -> ProviderVideoCreateFamily { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/codex/tests.rs b/apps/aether-gateway/src/ai_serving/planner/standard/codex/tests.rs index 52b3ef76f..30b51f6c6 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/codex/tests.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/codex/tests.rs @@ -505,9 +505,12 @@ fn projects_uuid_prompt_cache_identity_into_missing_session_headers() { assert_eq!(headers.get("x-client-request-id"), None); assert_eq!( headers.get("user-agent").map(String::as_str), - Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT) + Some(aether_ai_formats::codex_client_user_agent().as_str()) + ); + assert_eq!( + headers.get("originator"), + Some(&aether_ai_formats::codex_client_originator()) ); - assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string())); assert!(!headers.contains_key("version")); assert_eq!(headers.get("x-openai-fedramp"), Some(&"true".to_string())); assert_eq!( @@ -615,9 +618,12 @@ fn injects_only_codex_client_headers_for_images_requests() { ); assert_eq!( headers.get("user-agent").map(String::as_str), - Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT) + Some(aether_ai_formats::codex_client_user_agent().as_str()) + ); + assert_eq!( + headers.get("originator"), + Some(&aether_ai_formats::codex_client_originator()) ); - assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string())); assert!(!headers.contains_key("version")); assert_eq!(headers.get("x-openai-fedramp"), Some(&"true".to_string())); for name in ["x-client-request-id", "session-id", "thread-id"] { @@ -699,9 +705,12 @@ fn preserves_client_context_headers_and_enforces_codex_provider_identity() { ); assert_eq!( headers.get("user-agent").map(String::as_str), - Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT) + Some(aether_ai_formats::codex_client_user_agent().as_str()) + ); + assert_eq!( + headers.get("originator"), + Some(&aether_ai_formats::codex_client_originator()) ); - assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string())); assert_eq!( headers .keys() @@ -763,9 +772,12 @@ fn compact_projects_uuid_prompt_cache_identity_into_session_headers() { assert_eq!(headers.get("x-client-request-id"), None); assert_eq!( headers.get("user-agent").map(String::as_str), - Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT) + Some(aether_ai_formats::codex_client_user_agent().as_str()) + ); + assert_eq!( + headers.get("originator"), + Some(&aether_ai_formats::codex_client_originator()) ); - assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string())); assert!(!headers.contains_key("version")); assert_eq!(headers.get("x-openai-fedramp"), Some(&"true".to_string())); assert_eq!( diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs b/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs index 6bb5ac92a..bc0eb87d6 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs @@ -13,7 +13,9 @@ pub(crate) fn openai_responses_reasoning_replay_policy( base_url: &str, _provider_model: &str, ) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy { - if is_deepseek_provider(provider_type, base_url) { + if provider_type.trim().eq_ignore_ascii_case("xai") { + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + } else if is_deepseek_provider(provider_type, base_url) { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque } else { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds @@ -238,6 +240,27 @@ mod tests { openai_responses_reasoning_replay_policy, }; + #[test] + fn xai_reasoning_policy_comes_from_provider_type() { + use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy; + assert_eq!( + openai_responses_reasoning_replay_policy( + "xai", + "https://custom.example/v1", + "grok-4.6" + ), + OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + ); + assert_eq!( + openai_responses_reasoning_replay_policy( + "openai", + "https://custom.example/v1", + "grok-4.6" + ), + OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds + ); + } + #[test] fn detects_deepseek_provider_only_by_official_host() { assert!(!is_deepseek_provider( diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/normalize/chat.rs b/apps/aether-gateway/src/ai_serving/planner/standard/normalize/chat.rs index c113e88f6..5c2fbfe31 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/normalize/chat.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/normalize/chat.rs @@ -4,7 +4,7 @@ use crate::ai_serving::transport::apply_standard_provider_request_body_rules_wit use crate::ai_serving::{ apply_codex_openai_responses_chat_body_edits, apply_openai_responses_compact_special_body_edits, - build_cross_format_openai_chat_request_body_with_model_directives as surface_build_cross_format_openai_chat_request_body, + build_cross_format_openai_chat_request_body_with_provider_context as surface_build_cross_format_openai_chat_request_body, build_local_openai_chat_request_body_with_model_directives as surface_build_local_openai_chat_request_body, GatewayProviderTransportSnapshot, }; @@ -73,9 +73,11 @@ pub(crate) fn build_cross_format_openai_chat_request_body( let provider_request_body = surface_build_cross_format_openai_chat_request_body( body_json, mapped_model, + provider_type, provider_api_format, upstream_is_stream, enable_model_directives, + user_api_key_id, )?; let mut provider_request_body = apply_standard_provider_request_body_rules_with_request_headers( @@ -110,6 +112,42 @@ pub(crate) fn build_cross_format_openai_chat_request_body( Some(provider_request_body) } +#[cfg(test)] +mod antigravity_schema_tests { + use super::*; + use serde_json::json; + + #[test] + fn antigravity_chat_route_preserves_tool_schema_and_alternate_responses_shape() { + let schema = json!({"type": "object", "properties": {"mode": {"const": "fast"}}}); + let body = json!({"model": "client", "messages": [{"role": "user", "content": "hi"}], + "tools": [{"type": "function", "function": {"name": "probe", "parameters": schema}}]}); + let responses_body = json!({"model": "client", "input": "hi", + "tools": [{"type": "function", "name": "probe", "parameters": schema}]}); + for input in [body, responses_body] { + for provider in ["antigravity", "gemini"] { + let output = build_cross_format_openai_chat_request_body( + &input, + "claude-test", + provider, + "gemini:generate_content", + true, + false, + None, + None, + &http::HeaderMap::new(), + false, + ) + .unwrap(); + let parameters = &output["tools"][0]["functionDeclarations"][0]["parameters"]; + assert_eq!(parameters == &schema, provider == "antigravity"); + assert!(output.get("stream").is_none()); + assert_eq!(output["contents"][0]["parts"][0]["text"], "hi"); + } + } + } +} + pub(crate) fn build_cross_format_openai_chat_upstream_url( parts: &http::request::Parts, transport: &GatewayProviderTransportSnapshot, diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/normalize/responses.rs b/apps/aether-gateway/src/ai_serving/planner/standard/normalize/responses.rs index 47c91f1b9..5a138ac83 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/normalize/responses.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/normalize/responses.rs @@ -3,7 +3,7 @@ use serde_json::Value; use crate::ai_serving::transport::apply_standard_provider_request_body_rules_with_request_headers; use crate::ai_serving::{ apply_openai_responses_compact_special_body_edits, - build_cross_format_openai_responses_request_body_with_model_directives_and_history_scope as surface_build_cross_format_openai_responses_request_body, + build_cross_format_openai_responses_request_body_with_provider_context as surface_build_cross_format_openai_responses_request_body, build_local_openai_responses_request_body_with_model_directives as surface_build_local_openai_responses_request_body, GatewayProviderTransportSnapshot, }; @@ -218,6 +218,7 @@ pub(crate) fn build_cross_format_openai_responses_request_body_with_codex_model_ body_json, mapped_model, client_api_format, + provider_type, provider_api_format, upstream_is_stream, enable_model_directives, @@ -274,6 +275,41 @@ pub(crate) fn build_local_openai_responses_upstream_url( ) } +#[cfg(test)] +mod antigravity_schema_tests { + use super::*; + use serde_json::json; + + #[test] + fn antigravity_responses_route_preserves_tool_schema_without_changing_public_gemini() { + let schema = json!({"type": "object", "properties": {"mode": {"const": "fast"}}}); + let input = json!({"model": "client", "input": "hi", + "tools": [{"type": "function", "name": "probe", "parameters": schema}]}); + for provider in ["antigravity", "gemini"] { + let output = + build_cross_format_openai_responses_request_body_with_codex_model_capabilities( + &input, + "claude-test", + "openai:responses", + "gemini:generate_content", + true, + false, + provider, + None, + &http::HeaderMap::new(), + Some("antigravity-schema-test"), + None, + false, + ) + .unwrap(); + let parameters = &output["tools"][0]["functionDeclarations"][0]["parameters"]; + assert_eq!(parameters == &schema, provider == "antigravity"); + assert!(output.get("stream").is_none()); + assert_eq!(output["contents"][0]["parts"][0]["text"], "hi"); + } + } +} + pub(crate) fn build_cross_format_openai_responses_upstream_url( parts: &http::request::Parts, transport: &GatewayProviderTransportSnapshot, diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/request.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/request.rs index 4bfa07dea..87d17fd6d 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/request.rs @@ -140,7 +140,7 @@ fn finalize_openai_chat_provider_request_body( mapped_model, source_model, ); - crate::ai_serving::finalize_openai_provider_request_with_codex_model_capabilities_and_reasoning_replay_policy( + let finalization_failure = crate::ai_serving::finalize_openai_provider_request_with_codex_model_capabilities_and_reasoning_replay_policy( provider_request_body, crate::ai_serving::OpenAiProviderRequestFinalization { source_api_format: "openai:chat", @@ -170,7 +170,17 @@ fn finalize_openai_chat_provider_request_body( provider_api_format, "openai_chat_request_finalization", ) - }) + }); + if finalization_failure.is_none() { + // This builder does not go through `apply_transport_request_body_semantics`, so the + // Claude Code body mimicry must be applied here for Chat -> claude_code requests. + crate::ai_serving::transport::claude_code::apply_claude_code_body_mimicry_for_transport( + provider_request_body, + transport, + provider_api_format, + ); + } + finalization_failure } #[allow(clippy::too_many_arguments)] @@ -2763,7 +2773,7 @@ mod tests { payload.provider_request_body["userAgent"], "vscode/1.X.X (Antigravity/4.3.0)" ); - assert_eq!(payload.provider_request_body["requestType"], "agent"); + assert!(payload.provider_request_body.get("requestType").is_none()); assert!(payload.provider_request_body.get("contents").is_none()); assert!(payload.provider_request_body["request"] .get("contents") diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/stream.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/stream.rs index 0f6ae98c6..c846b8df8 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/stream.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/stream.rs @@ -1,4 +1,4 @@ -use aether_routing_core::RoutingExecutionPolicy; +use aether_routing_core::{RoutingExecutionPolicy, RoutingSchedulingMode}; use async_trait::async_trait; use std::collections::VecDeque; use tracing::warn; @@ -207,7 +207,12 @@ impl LocalOpenAiChatStreamAttemptSource<'_> { async fn next_raw_attempt_with_target_select( &mut self, ) -> Result, GatewayError> { - let select_window = openai_chat_stream_target_select_window(); + let select_window = openai_chat_stream_target_select_window_for_mode( + self.input + .routing_policy + .as_ref() + .map(|policy| policy.scheduling_mode), + ); if select_window <= 1 { return self.next_raw_attempt_linear().await; } @@ -365,6 +370,15 @@ fn openai_chat_stream_target_select_window() -> usize { .clamp(1, MAX_OPENAI_CHAT_STREAM_TARGET_SELECT_WINDOW) } +fn openai_chat_stream_target_select_window_for_mode( + scheduling_mode: Option, +) -> usize { + if scheduling_mode == Some(RoutingSchedulingMode::FixedOrder) { + return 1; + } + openai_chat_stream_target_select_window() +} + #[derive(Clone, Copy)] struct TargetSelectCandidateIdentity<'a> { provider_id: &'a str, @@ -574,4 +588,14 @@ mod tests { assert_eq!(select_target_index(19, &choices), 1); } + + #[test] + fn fixed_order_disables_stream_target_selection() { + assert_eq!( + openai_chat_stream_target_select_window_for_mode(Some( + RoutingSchedulingMode::FixedOrder, + )), + 1 + ); + } } diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/request.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/request.rs index a9170ce59..0a52493bc 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/request.rs @@ -635,6 +635,13 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_ { log_responses_to_chat_tool_conversion(trace_id, body_json, &base_provider_request_body); } + // This builder does not go through `apply_transport_request_body_semantics`, so the + // Claude Code body mimicry must be applied here for Responses -> claude_code requests. + crate::ai_serving::transport::claude_code::apply_claude_code_body_mimicry_for_transport( + &mut base_provider_request_body, + &transport, + provider_api_format, + ); let provider_request_body = base_provider_request_body; if let Some(kiro_auth) = kiro_auth.as_ref() { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/mod.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/mod.rs index 7ea4f4bdf..c6fcc2779 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/mod.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/mod.rs @@ -488,6 +488,7 @@ impl ResponsesWebSocketBodyNormalization { digest.update([match self.reasoning_replay_policy { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds => 0, crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque => 1, + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted => 2, }]); update_normalization_optional_json_digest(&mut digest, self.model_directive_patch.as_ref()); digest.finalize().into() diff --git a/apps/aether-gateway/src/ai_serving/pure/mod.rs b/apps/aether-gateway/src/ai_serving/pure/mod.rs index 91115566b..2b3eb5d7a 100644 --- a/apps/aether-gateway/src/ai_serving/pure/mod.rs +++ b/apps/aether-gateway/src/ai_serving/pure/mod.rs @@ -12,13 +12,16 @@ pub(crate) use aether_ai_formats::api::{ apply_codex_openai_responses_websocket_continuation_body_edits_with_source_model_and_capabilities, apply_codex_openai_special_headers, apply_model_directive_mapping_patch, apply_model_directive_overrides_from_model, apply_model_directive_overrides_from_request, - apply_openai_responses_compact_special_body_edits, build_chatgpt_web_image_request_body, + apply_openai_responses_compact_special_body_edits, apply_xai_upstream_payload_edits, + apply_xai_upstream_payload_edits_with_client, build_chatgpt_web_image_request_body, build_codex_model_catalog_metadata, build_codex_openai_image_api_provider_request_body, build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body, build_cross_format_openai_chat_request_body_with_model_directives, + build_cross_format_openai_chat_request_body_with_provider_context, build_cross_format_openai_responses_request_body, build_cross_format_openai_responses_request_body_with_model_directives, build_cross_format_openai_responses_request_body_with_model_directives_and_history_scope, + build_cross_format_openai_responses_request_body_with_provider_context, build_gemini_image_request_body_from_openai_image_request, build_gemini_image_response_from_openai_image_response, build_gemini_image_response_from_openai_responses_image_response, build_generated_tool_call_id, @@ -181,7 +184,7 @@ pub(crate) use aether_ai_formats::{ openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id, strip_incompatible_openai_responses_reasoning_items, strip_incompatible_openai_responses_reasoning_items_with_policy, ApiOperation, ClientSurface, - CODEX_CLIENT_VERSION, OPENAI_RESPONSES_OPERATION_COMPACT, + OPENAI_RESPONSES_OPERATION_COMPACT, }; pub(crate) fn plan_kind_matches_api_operation( diff --git a/apps/aether-gateway/src/ai_serving/transport.rs b/apps/aether-gateway/src/ai_serving/transport.rs index 97fe5434a..8c81cb17d 100644 --- a/apps/aether-gateway/src/ai_serving/transport.rs +++ b/apps/aether-gateway/src/ai_serving/transport.rs @@ -58,6 +58,10 @@ pub(crate) mod windsurf { pub(crate) use aether_provider_transport::windsurf::*; } +pub(crate) mod xai { + pub(crate) use aether_provider_transport::xai::*; +} + pub(crate) use aether_provider_transport::{ append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence, apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides, diff --git a/apps/aether-gateway/src/api/ai/registry.rs b/apps/aether-gateway/src/api/ai/registry.rs index 4fe6cddbf..685e0d956 100644 --- a/apps/aether-gateway/src/api/ai/registry.rs +++ b/apps/aether-gateway/src/api/ai/registry.rs @@ -53,6 +53,8 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[ "/v1beta/operations/{*operation_path}", "/v1/videos", "/v1/videos/{*video_path}", + "/openai/v1/videos", + "/openai/v1/videos/{*video_path}", "/upload/v1beta/files", "/v1beta/files", "/v1beta/files/{*file_path}", diff --git a/apps/aether-gateway/src/async_task/runtime.rs b/apps/aether-gateway/src/async_task/runtime.rs index d1e4a64b5..f5d301ee2 100644 --- a/apps/aether-gateway/src/async_task/runtime.rs +++ b/apps/aether-gateway/src/async_task/runtime.rs @@ -536,6 +536,9 @@ mod tests { fn sample_sparse_stored_task() -> StoredVideoTask { let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-1".to_string(), upstream_task_id: "ext-1".to_string(), created_at_unix_ms: 1, diff --git a/apps/aether-gateway/src/backup/executor.rs b/apps/aether-gateway/src/backup/executor.rs index a7bc20baf..5b518b916 100644 --- a/apps/aether-gateway/src/backup/executor.rs +++ b/apps/aether-gateway/src/backup/executor.rs @@ -799,6 +799,7 @@ mod tests { use aes_gcm::aead::{Aead, AeadCore, KeyInit, OsRng, Payload}; use aes_gcm::Aes256Gcm; use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use base64::Engine as _; use bytes::Bytes; use chrono::{DateTime, Utc}; use serde_json::json; @@ -1243,8 +1244,18 @@ mod tests { assert_eq!(restored.key_id, None); assert_eq!(restored.export_version.as_deref(), Some("2.3")); + // 17 个互不相同的合法 base64-32 字节直接密钥:本段只验证“legacy 候选 >16 → TooManyLegacyKeys”, + // 不测口令强度、不解密。直接密钥走 decode_direct_fernet_key(生产已支持路径),跳过 PBKDF2, + // 避免本用例为计数语义再付 17×10 万次迭代;上半段 DEVELOPMENT_ENCRYPTION_KEY 真实 v1 兼容 + // 与 wrong-legacy-secret 派生路径保持不变。 let too_many: Vec<_> = (0..17) - .map(|index| BackupDecryptionKey::historical(format!("legacy-{index}")).unwrap()) + .map(|index| { + let mut material = [0u8; 32]; + material[0] = index as u8 + 1; + material[31] = index as u8 + 1; + let secret = base64::engine::general_purpose::STANDARD.encode(material); + BackupDecryptionKey::historical(secret).unwrap() + }) .collect(); assert!(matches!( restore_backup_json( diff --git a/apps/aether-gateway/src/bin/aether-codex-ws-probe.rs b/apps/aether-gateway/src/bin/aether-codex-ws-probe.rs index b9def32c9..ec9437da9 100644 --- a/apps/aether-gateway/src/bin/aether-codex-ws-probe.rs +++ b/apps/aether-gateway/src/bin/aether-codex-ws-probe.rs @@ -7,7 +7,7 @@ #[path = "support/responses_ws_probe.rs"] mod responses_ws_probe; -use aether_gateway::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT}; +use aether_gateway::{codex_client_originator, codex_client_user_agent}; use clap::Parser; use http::header::{AUTHORIZATION, USER_AGENT}; use http::{HeaderMap, HeaderName, HeaderValue}; @@ -78,14 +78,12 @@ fn handshake_headers(access_token: &str, account_id: &str) -> Result, +} + +#[derive(Debug, Deserialize, Serialize)] +struct CachedProfile { + version: String, + verified_at_unix_secs: u64, +} + +#[derive(Debug, thiserror::Error)] +enum ProfileRefreshError { + #[error("Codex CLI release client initialization failed: {0}")] + Client(#[from] reqwest::Error), + #[error("Codex CLI release request returned HTTP {0}")] + HttpStatus(u16), + #[error("Codex CLI release response exceeded {MAX_RELEASE_BYTES} bytes")] + ResponseTooLarge, + #[error("Codex CLI release metadata is invalid")] + InvalidMetadata, + #[error("Codex CLI release version is older than the active profile")] + Rollback, + #[error("Codex CLI profile cache operation failed: {0}")] + Cache(String), +} + +fn version_sequence(version: &str) -> Result { + let parsed = Version::parse(version).map_err(|_| ProfileRefreshError::InvalidMetadata)?; + if !parsed.pre.is_empty() + || !parsed.build.is_empty() + || parsed.major > 999 + || parsed.minor > 999 + || parsed.patch > 999 + { + return Err(ProfileRefreshError::InvalidMetadata); + } + Ok(1 + parsed.major * 1_000_000 + parsed.minor * 1_000 + parsed.patch) +} + +/// 校验官方 npm stable 标签及六个平台依赖来自同一版本发布。 +fn parse_cli_release(bytes: &[u8]) -> Result { + if bytes.len() > MAX_RELEASE_BYTES { + return Err(ProfileRefreshError::ResponseTooLarge); + } + let release = serde_json::from_slice::(bytes) + .map_err(|_| ProfileRefreshError::InvalidMetadata)?; + let sequence = version_sequence(&release.version)?; + if sequence == 0 + || release.name != "@openai/codex" + || CLI_TARGETS.iter().any(|target| { + release + .optional_dependencies + .get(&format!("@openai/codex-{target}")) + != Some(&format!("npm:@openai/codex@{}-{target}", release.version)) + }) + { + return Err(ProfileRefreshError::InvalidMetadata); + } + Ok(release.version) +} + +fn refresh_enabled_from(value: Option<&str>) -> bool { + !value.is_some_and(|value| { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "0" | "false" | "off" + ) + }) +} + +fn refresh_enabled() -> bool { + refresh_enabled_from( + std::env::var("AETHER_CODEX_CLIENT_PROFILE_REFRESH") + .ok() + .as_deref(), + ) +} + +fn fixed_version_from(value: Option<&str>) -> Option { + let value = value?.trim(); + if value.is_empty() || version_sequence(value).is_err() { + None + } else { + Some(value.to_owned()) + } +} + +fn fixed_version_override() -> Option { + let value = std::env::var("AETHER_CODEX_CLIENT_VERSION").ok()?; + let version = fixed_version_from(Some(&value)); + if version.is_none() { + warn!( + event_name = "codex_client_profile_fixed_version_invalid", + "AETHER_CODEX_CLIENT_VERSION is invalid; using cached or built-in profile" + ); + } + version +} + +fn build_release_client() -> Result { + Client::builder() + .https_only(true) + .no_proxy() + .redirect(Policy::none()) + .connect_timeout(RELEASE_CONNECT_TIMEOUT) + .timeout(RELEASE_REQUEST_TIMEOUT) + .build() + .map_err(ProfileRefreshError::Client) +} + +async fn fetch_latest_cli_version(client: &Client) -> Result { + let response = client + .get(CLI_RELEASE_ENDPOINT) + .send() + .await + .map_err(ProfileRefreshError::Client)?; + if !response.status().is_success() { + return Err(ProfileRefreshError::HttpStatus(response.status().as_u16())); + } + if response + .content_length() + .is_some_and(|length| length > MAX_RELEASE_BYTES as u64) + { + return Err(ProfileRefreshError::ResponseTooLarge); + } + + let mut bytes = Vec::new(); + let mut stream = response.bytes_stream(); + while let Some(chunk) = stream.next().await { + let chunk = chunk.map_err(ProfileRefreshError::Client)?; + if bytes.len().saturating_add(chunk.len()) > MAX_RELEASE_BYTES { + return Err(ProfileRefreshError::ResponseTooLarge); + } + bytes.extend_from_slice(&chunk); + } + parse_cli_release(&bytes) +} + +async fn restore_cached_profile(runtime: &RuntimeState) -> Result<(), ProfileRefreshError> { + let Some(raw) = runtime + .kv_get(PROFILE_CACHE_KEY) + .await + .map_err(|err| ProfileRefreshError::Cache(err.to_string()))? + else { + return Ok(()); + }; + let cached = serde_json::from_str::(&raw) + .map_err(|_| ProfileRefreshError::InvalidMetadata)?; + if let Some(version) = cached_version_to_restore(&cached, &codex_client_version())? { + set_codex_cli_version(&version).map_err(|_| ProfileRefreshError::InvalidMetadata)?; + info!( + event_name = "codex_client_profile_restored", + version = %version, + verified_at_unix_secs = cached.verified_at_unix_secs, + "restored cached Codex CLI profile" + ); + } + Ok(()) +} + +fn cached_version_to_restore( + cached: &CachedProfile, + active_version: &str, +) -> Result, ProfileRefreshError> { + let cached_sequence = version_sequence(&cached.version)?; + let active_sequence = version_sequence(active_version)?; + Ok((cached_sequence >= active_sequence).then(|| cached.version.clone())) +} + +async fn refresh_once_with_fetch( + runtime: &RuntimeState, + fixed_version: Option<&str>, + refresh_is_enabled: bool, + fetch_latest: F, +) -> Result +where + F: FnOnce() -> Fut, + Fut: Future>, +{ + if let Some(version) = fixed_version { + set_codex_cli_version(version).map_err(|_| ProfileRefreshError::InvalidMetadata)?; + return Ok(version.to_owned()); + } + + if let Err(error) = restore_cached_profile(runtime).await { + // 缓存损坏或暂时不可用不应阻断官方版本检查;当前进程继续使用旧画像。 + warn!( + event_name = "codex_client_profile_cache_restore_failed", + error = %error, + "could not restore cached Codex CLI profile" + ); + } + if !refresh_is_enabled { + return Ok(codex_client_version()); + } + + let version = fetch_latest().await?; + let current = codex_client_version(); + if version_sequence(&version)? < version_sequence(¤t)? { + return Err(ProfileRefreshError::Rollback); + } + + let cached = CachedProfile { + version: version.clone(), + verified_at_unix_secs: chrono::Utc::now().timestamp().max(0) as u64, + }; + let serialized = + serde_json::to_string(&cached).map_err(|_| ProfileRefreshError::InvalidMetadata)?; + set_codex_cli_version(&version).map_err(|_| ProfileRefreshError::InvalidMetadata)?; + if let Err(error) = runtime + .kv_set(PROFILE_CACHE_KEY, serialized, Some(PROFILE_CACHE_TTL)) + .await + { + // 本地画像已经完成原子替换;缓存写失败只影响下次进程启动的恢复。 + warn!( + event_name = "codex_client_profile_cache_write_failed", + error = %error, + "published Codex CLI profile locally but could not persist the cache" + ); + } + Ok(version) +} + +async fn refresh_once(runtime: &RuntimeState) -> Result { + let fixed_version = fixed_version_override(); + refresh_once_with_fetch( + runtime, + fixed_version.as_deref(), + refresh_enabled(), + || async { + let client = build_release_client()?; + fetch_latest_cli_version(&client).await + }, + ) + .await +} + +pub(crate) async fn prewarm(runtime: &RuntimeState) -> Result { + refresh_once(runtime).await.map_err(|err| err.to_string()) +} + +pub(crate) fn spawn_worker(app: AppState) -> tokio::task::JoinHandle<()> { + crate::task_runtime::spawn_singleton_worker( + app, + crate::task_runtime::TASK_KEY_CODEX_CLIENT_PROFILE, + |app| async move { + let mut interval = tokio::time::interval(PROFILE_REFRESH_INTERVAL); + interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + // 启动阶段由 prewarm 完成一次检查;后台任务只负责后续每日刷新,避免重复建连。 + interval.tick().await; + loop { + interval.tick().await; + match refresh_once(app.runtime_state()).await { + Ok(version) => info!( + event_name = "codex_client_profile_refreshed", + version = %version, + "refreshed Codex CLI profile" + ), + Err(error) => warn!( + event_name = "codex_client_profile_refresh_failed", + error = %error, + "keeping the previous Codex CLI profile after refresh failure" + ), + } + } + }, + ) +} + +#[cfg(test)] +mod tests { + use std::sync::{ + atomic::{AtomicBool, Ordering}, + Mutex, OnceLock, + }; + use std::time::Duration; + + use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState}; + + use super::{ + cached_version_to_restore, fixed_version_from, parse_cli_release, refresh_enabled_from, + refresh_once_with_fetch, CachedProfile, ProfileRefreshError, PROFILE_CACHE_KEY, + }; + use crate::ai_serving::api::{ + codex_client_profile, codex_client_version, set_codex_cli_version, + set_codex_client_profile, CodexClientProfile, + }; + + static PROFILE_TEST_LOCK: OnceLock> = OnceLock::new(); + + struct ProfileRestore(CodexClientProfile); + + impl Drop for ProfileRestore { + fn drop(&mut self) { + set_codex_client_profile(self.0.clone()); + } + } + + fn profile_restore_guard() -> (std::sync::MutexGuard<'static, ()>, ProfileRestore) { + let lock = PROFILE_TEST_LOCK.get_or_init(|| Mutex::new(())); + let guard = lock.lock().expect("profile test lock"); + let restore = ProfileRestore(codex_client_profile()); + (guard, restore) + } + + #[test] + fn accepts_only_one_verified_cli_release_for_all_targets() { + let body = serde_json::json!({ + "name": "@openai/codex", + "version": "0.200.1", + "optionalDependencies": { + "@openai/codex-darwin-arm64": "npm:@openai/codex@0.200.1-darwin-arm64", + "@openai/codex-darwin-x64": "npm:@openai/codex@0.200.1-darwin-x64", + "@openai/codex-linux-arm64": "npm:@openai/codex@0.200.1-linux-arm64", + "@openai/codex-linux-x64": "npm:@openai/codex@0.200.1-linux-x64", + "@openai/codex-win32-arm64": "npm:@openai/codex@0.200.1-win32-arm64", + "@openai/codex-win32-x64": "npm:@openai/codex@0.200.1-win32-x64" + } + }); + assert_eq!( + parse_cli_release(&serde_json::to_vec(&body).unwrap()).unwrap(), + "0.200.1" + ); + } + + #[test] + fn rejects_incomplete_platform_release() { + let body = serde_json::json!({ + "name": "@openai/codex", + "version": "0.200.1", + "optionalDependencies": {} + }); + assert!(parse_cli_release(&serde_json::to_vec(&body).unwrap()).is_err()); + } + + #[test] + fn refresh_and_fixed_version_environment_policies_are_strict() { + assert!(!refresh_enabled_from(Some("off"))); + assert!(!refresh_enabled_from(Some(" FALSE "))); + assert!(refresh_enabled_from(None)); + assert_eq!( + fixed_version_from(Some(" 0.200.1 ")).as_deref(), + Some("0.200.1") + ); + assert!(fixed_version_from(Some("0.200.1-beta.1")).is_none()); + assert!(fixed_version_from(Some("1.2")).is_none()); + } + + #[test] + fn cached_profile_never_rewinds_active_profile() { + let cached = CachedProfile { + version: "0.200.1".to_string(), + verified_at_unix_secs: 1, + }; + assert_eq!( + cached_version_to_restore(&cached, "0.200.0").unwrap(), + Some("0.200.1".to_string()) + ); + assert_eq!(cached_version_to_restore(&cached, "0.201.0").unwrap(), None); + } + + #[tokio::test] + async fn cache_hit_is_restored_without_network_when_refresh_is_disabled() { + let (_lock, _restore) = profile_restore_guard(); + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + runtime + .kv_set( + PROFILE_CACHE_KEY, + serde_json::to_string(&CachedProfile { + version: "0.200.1".to_string(), + verified_at_unix_secs: 1, + }) + .unwrap(), + Some(Duration::from_secs(60)), + ) + .await + .unwrap(); + + let result = refresh_once_with_fetch(&runtime, None, false, || async { + Err(ProfileRefreshError::HttpStatus(599)) + }) + .await + .unwrap(); + + assert_eq!(result, "0.200.1"); + assert_eq!(codex_client_version(), "0.200.1"); + } + + #[tokio::test] + async fn refresh_failure_keeps_previous_profile() { + let (_lock, _restore) = profile_restore_guard(); + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let before = codex_client_profile(); + let result = refresh_once_with_fetch(&runtime, None, true, || async { + Err(ProfileRefreshError::HttpStatus(503)) + }) + .await; + + assert!(matches!(result, Err(ProfileRefreshError::HttpStatus(503)))); + assert_eq!(codex_client_profile(), before); + } + + #[tokio::test] + async fn fixed_version_override_skips_network_and_publishes_profile() { + let (_lock, _restore) = profile_restore_guard(); + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let fetch_called = AtomicBool::new(false); + let result = refresh_once_with_fetch(&runtime, Some("0.220.0"), true, || async { + fetch_called.store(true, Ordering::SeqCst); + Ok("0.221.0".to_string()) + }) + .await + .unwrap(); + + assert_eq!(result, "0.220.0"); + assert!(!fetch_called.load(Ordering::SeqCst)); + assert_eq!(codex_client_version(), "0.220.0"); + } + + #[tokio::test] + async fn rollback_is_rejected_without_replacing_profile() { + let (_lock, _restore) = profile_restore_guard(); + set_codex_cli_version("0.220.0").unwrap(); + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let result = + refresh_once_with_fetch(&runtime, None, true, || async { Ok("0.219.9".to_string()) }) + .await; + + assert!(matches!(result, Err(ProfileRefreshError::Rollback))); + assert_eq!(codex_client_version(), "0.220.0"); + } +} diff --git a/apps/aether-gateway/src/constants.rs b/apps/aether-gateway/src/constants.rs index e1edeb5d9..9a2bcbe14 100644 --- a/apps/aether-gateway/src/constants.rs +++ b/apps/aether-gateway/src/constants.rs @@ -140,6 +140,8 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[ "/v1beta/models/{model}/operations/{id}", "/v1beta/operations", "/v1beta/operations/{id}", + "/openai/v1/videos", + "/openai/v1/videos/{path...}", "/v1/videos", "/v1/videos/{path...}", "/upload/v1beta/files", diff --git a/apps/aether-gateway/src/control/route/admin/observability_families.rs b/apps/aether-gateway/src/control/route/admin/observability_families.rs index 3c287781a..0276c479d 100644 --- a/apps/aether-gateway/src/control/route/admin/observability_families.rs +++ b/apps/aether-gateway/src/control/route/admin/observability_families.rs @@ -605,6 +605,20 @@ pub(super) fn classify_admin_observability_family_route( "admin:stats", false, )) + } else if method == http::Method::GET + && matches!( + normalized_path, + "/api/admin/stats/leaderboard/user-groups" + | "/api/admin/stats/leaderboard/user-groups/" + ) + { + Some(classified( + "admin_proxy", + "stats_manage", + "leaderboard_user_groups", + "admin:stats", + false, + )) } else if method == http::Method::GET && matches!( normalized_path, diff --git a/apps/aether-gateway/src/control/route/ai.rs b/apps/aether-gateway/src/control/route/ai.rs index f6ac5c83b..856f62543 100644 --- a/apps/aether-gateway/src/control/route/ai.rs +++ b/apps/aether-gateway/src/control/route/ai.rs @@ -137,7 +137,11 @@ pub(super) fn classify_ai_public_route( .with_client_surface(detect_claude_client_surface(headers)) .with_api_operation(ApiOperation::ClaudeMessagesCreate), ) - } else if normalized_path.starts_with("/v1/videos") { + } else if normalized_path == "/v1/videos" + || normalized_path.starts_with("/v1/videos/") + || normalized_path == "/openai/v1/videos" + || normalized_path.starts_with("/openai/v1/videos/") + { Some(classified( "ai_public", "openai", diff --git a/apps/aether-gateway/src/control/tests/admin_core.rs b/apps/aether-gateway/src/control/tests/admin_core.rs index 1b8808419..f4979763e 100644 --- a/apps/aether-gateway/src/control/tests/admin_core.rs +++ b/apps/aether-gateway/src/control/tests/admin_core.rs @@ -1,6 +1,7 @@ use http::Uri; -use crate::control::management_token_required_permission; +use crate::control::{management_token_required_permission, GatewayPublicRequestContext}; +use crate::handlers::shared::local_proxy_route_requires_buffered_body; use super::{classify_control_route, headers}; @@ -206,6 +207,10 @@ fn classifies_admin_system_maintenance_write_routes_as_admin_proxy_route() { "/api/admin/system/important-notification/test", "important_notification_test", ), + ( + "/api/admin/system/cleanup/usage/manual", + "cleanup_usage_manual", + ), ("/api/admin/system/cleanup", "cleanup"), ("/api/admin/system/purge/config", "purge_config"), ("/api/admin/system/purge/users", "purge_users"), @@ -235,6 +240,28 @@ fn classifies_admin_system_maintenance_write_routes_as_admin_proxy_route() { Some("admin:system") ); assert!(!decision.is_execution_runtime_candidate()); + + if matches!( + expected_kind, + "config_import" + | "users_import" + | "data_import" + | "smtp_test" + | "important_notification_test" + | "cleanup_usage_manual" + ) { + let context = GatewayPublicRequestContext::from_request_parts( + "trace-system-maintenance-write", + &http::Method::POST, + &uri, + &headers, + Some(decision), + ); + assert!( + local_proxy_route_requires_buffered_body(&context), + "POST {path} should buffer request body" + ); + } } } @@ -303,6 +330,20 @@ fn classifies_admin_system_update_routes_as_admin_proxy_routes() { Some("admin:system") ); assert!(!decision.is_execution_runtime_candidate()); + + if matches!(expected_kind, "prepare_update" | "apply_update") { + let context = GatewayPublicRequestContext::from_request_parts( + "trace-system-update-write", + &method, + &uri, + &headers, + Some(decision), + ); + assert!( + local_proxy_route_requires_buffered_body(&context), + "{method} {path} should buffer request body" + ); + } } } diff --git a/apps/aether-gateway/src/control/tests/admin_stats.rs b/apps/aether-gateway/src/control/tests/admin_stats.rs index 7f12d4b62..edde61a7b 100644 --- a/apps/aether-gateway/src/control/tests/admin_stats.rs +++ b/apps/aether-gateway/src/control/tests/admin_stats.rs @@ -197,6 +197,28 @@ fn classifies_admin_stats_leaderboard_models_as_admin_proxy_route() { assert!(!decision.is_execution_runtime_candidate()); } +#[test] +fn classifies_admin_stats_leaderboard_user_groups_as_admin_proxy_route() { + let headers = headers(&[]); + let uri: Uri = "/api/admin/stats/leaderboard/user-groups" + .parse() + .expect("uri should parse"); + let decision = + classify_control_route(&http::Method::GET, &uri, &headers).expect("route should classify"); + + assert_eq!(decision.route_class.as_deref(), Some("admin_proxy")); + assert_eq!(decision.route_family.as_deref(), Some("stats_manage")); + assert_eq!( + decision.route_kind.as_deref(), + Some("leaderboard_user_groups") + ); + assert_eq!( + decision.auth_endpoint_signature.as_deref(), + Some("admin:stats") + ); + assert!(!decision.is_execution_runtime_candidate()); +} + #[test] fn classifies_admin_stats_leaderboard_users_as_admin_proxy_route() { let headers = headers(&[]); diff --git a/apps/aether-gateway/src/data/state/mod.rs b/apps/aether-gateway/src/data/state/mod.rs index 0a455684b..6837220c7 100644 --- a/apps/aether-gateway/src/data/state/mod.rs +++ b/apps/aether-gateway/src/data/state/mod.rs @@ -62,8 +62,9 @@ pub(crate) use aether_data::repository::users::{ StoredUserPreferenceRecord, StoredUserSessionRecord, }; use aether_data::repository::wallet::{ - AdjustWalletBalanceInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, - AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery, + AdjustWalletBalanceInBatchInput, AdjustWalletBalanceInput, AdminPaymentOrderListQuery, + AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, + AdminUserWalletBalanceBatchUserOutcome, AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput, CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, @@ -72,13 +73,14 @@ use aether_data::repository::wallet::{ CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, + PrepareAdminUserWalletBalanceBatchInput, PrepareAdminUserWalletBalanceBatchOutcome, ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, - StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage, - StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage, - StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, + StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminUserWalletBalanceBatch, + StoredAdminWalletLedgerPage, StoredAdminWalletListPage, StoredAdminWalletRefund, + StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, diff --git a/apps/aether-gateway/src/data/state/runtime.rs b/apps/aether-gateway/src/data/state/runtime.rs index c11ac96d2..a58a1960c 100644 --- a/apps/aether-gateway/src/data/state/runtime.rs +++ b/apps/aether-gateway/src/data/state/runtime.rs @@ -1,9 +1,10 @@ use super::{ read_decision_trace, read_provider_transport_snapshot, read_request_candidate_trace, - AdjustWalletBalanceInput, AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, - AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord, - AdminBillingRuleWriteInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, - AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery, + AdjustWalletBalanceInBatchInput, AdjustWalletBalanceInput, AdminBillingCollectorRecord, + AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult, + AdminBillingRuleRecord, AdminBillingRuleWriteInput, AdminPaymentOrderListQuery, + AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, + AdminUserWalletBalanceBatchUserOutcome, AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletRefundRequestListQuery, AnnouncementListQuery, AuditLogListQuery, BackgroundTaskListQuery, BackgroundTaskSummary, BillingModelContextCacheKey, BillingModelContextCacheState, BillingModelContextInflightState, BillingPlanRecord, @@ -17,24 +18,25 @@ use super::{ DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, FailWalletRechargeCheckoutInput, GatewayDataState, GatewayProviderTransportSnapshot, LocalVideoTaskReadResponse, PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, - PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, ProcessAdminWalletRefundInput, - ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, - ReconcileUsagePolicyCostInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, - ReleaseUsagePolicyRequestAdmissionInput, RequestAuditBundle, RequestCandidateTrace, - ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput, - ReserveUsagePolicyRequestOutcome, StoredAdminAuditLogPage, StoredAdminPaymentCallbackPage, - StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCodeBatch, - StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage, - StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage, - StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, - StoredAdminWalletTransactionPage, StoredAnnouncement, StoredAnnouncementPage, - StoredBackgroundTaskEvent, StoredBackgroundTaskRun, StoredBackgroundTaskRunPage, - StoredBillingModelContext, StoredProviderQuotaSnapshot, StoredProviderUsageSummary, - StoredRequestUsageAudit, StoredSuspiciousActivity, StoredUsagePolicyCostReservation, - StoredUsagePolicyRequestAdmission, StoredUsageSettlement, StoredUserAuditLogPage, - StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, StoredVideoTask, - StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, - UpdateAdminWalletRefundGatewayInput, UpdateAnnouncementRecord, + PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, + PrepareAdminUserWalletBalanceBatchInput, PrepareAdminUserWalletBalanceBatchOutcome, + ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + ReclaimWalletRechargeCheckoutInput, ReconcileUsagePolicyCostInput, RedeemWalletCodeInput, + RedeemWalletCodeOutcome, ReleaseUsagePolicyRequestAdmissionInput, RequestAuditBundle, + RequestCandidateTrace, ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, + ReserveUsagePolicyRequestInput, ReserveUsagePolicyRequestOutcome, StoredAdminAuditLogPage, + StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, + StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, + StoredAdminUserWalletBalanceBatch, StoredAdminWalletLedgerPage, StoredAdminWalletListPage, + StoredAdminWalletRefund, StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage, + StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredAnnouncement, + StoredAnnouncementPage, StoredBackgroundTaskEvent, StoredBackgroundTaskRun, + StoredBackgroundTaskRunPage, StoredBillingModelContext, StoredProviderQuotaSnapshot, + StoredProviderUsageSummary, StoredRequestUsageAudit, StoredSuspiciousActivity, + StoredUsagePolicyCostReservation, StoredUsagePolicyRequestAdmission, StoredUsageSettlement, + StoredUserAuditLogPage, StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, + StoredVideoTask, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, + StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, UpdateAnnouncementRecord, UpdateWalletRechargeCheckoutInput, UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun, UpsertUsageRecord, UpsertVideoTask, UsageSettlementInput, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, VideoTaskLookupKey, VideoTaskModelCount, VideoTaskQueryFilter, @@ -1101,13 +1103,81 @@ impl GatewayDataState { pub(crate) async fn adjust_wallet_balance( &self, input: AdjustWalletBalanceInput, - ) -> Result, DataLayerError> { + ) -> Result)>, DataLayerError> + { match &self.wallet_writer { Some(repository) => repository.adjust_wallet_balance(input).await, None => Ok(None), } } + pub(crate) async fn prepare_admin_user_wallet_balance_batch( + &self, + input: PrepareAdminUserWalletBalanceBatchInput, + ) -> Result, DataLayerError> { + match &self.wallet_writer { + Some(repository) => repository + .prepare_admin_user_wallet_balance_batch(input) + .await + .map(Some), + None => Ok(None), + } + } + + pub(crate) async fn get_admin_user_wallet_balance_batch( + &self, + admin_user_id: &str, + idempotency_key: &str, + request_fingerprint: &str, + ) -> Result, DataLayerError> { + match &self.wallet_writer { + Some(repository) => { + repository + .get_admin_user_wallet_balance_batch( + admin_user_id, + idempotency_key, + request_fingerprint, + ) + .await + } + None => Ok(None), + } + } + + pub(crate) async fn adjust_admin_user_wallet_balance_batch_user( + &self, + input: AdjustWalletBalanceInBatchInput, + ) -> Result, DataLayerError> { + match &self.wallet_writer { + Some(repository) => repository + .adjust_admin_user_wallet_balance_batch_user(input) + .await + .map(Some), + None => Ok(None), + } + } + + pub(crate) async fn record_admin_user_wallet_balance_batch_failure( + &self, + admin_user_id: &str, + idempotency_key: &str, + user_id: &str, + reason: &str, + ) -> Result, DataLayerError> { + match &self.wallet_writer { + Some(repository) => repository + .record_admin_user_wallet_balance_batch_failure( + admin_user_id, + idempotency_key, + user_id, + reason, + ) + .await + .map(Some), + None => Ok(None), + } + } + pub(crate) async fn create_manual_wallet_recharge( &self, input: CreateManualWalletRechargeInput, diff --git a/apps/aether-gateway/src/data/state/testing/video_tasks.rs b/apps/aether-gateway/src/data/state/testing/video_tasks.rs index 9db104d98..24952c220 100644 --- a/apps/aether-gateway/src/data/state/testing/video_tasks.rs +++ b/apps/aether-gateway/src/data/state/testing/video_tasks.rs @@ -123,6 +123,15 @@ impl GatewayDataState { } #[cfg(test)] + pub(crate) fn attach_video_task_repository_for_tests(mut self, repository: Arc) -> Self + where + T: VideoTaskRepository + 'static, + { + self.video_task_reader = Some(repository.clone()); + self.video_task_writer = Some(repository); + self + } + pub(crate) fn with_video_task_repository_for_tests(repository: Arc) -> Self where T: VideoTaskRepository + 'static, diff --git a/apps/aether-gateway/src/dispatch/pool_scheduler.rs b/apps/aether-gateway/src/dispatch/pool_scheduler.rs index 6338b93ff..dcb60a8ee 100644 --- a/apps/aether-gateway/src/dispatch/pool_scheduler.rs +++ b/apps/aether-gateway/src/dispatch/pool_scheduler.rs @@ -131,8 +131,13 @@ async fn schedule_pool_page_candidates( entry.1.insert(candidate.candidate.key_id.clone()); } - let key_context_by_id = - read_pool_catalog_key_contexts_by_id(state, &candidates, provider_model_name).await; + let key_context_by_id = read_pool_catalog_key_contexts_by_id( + state, + &candidates, + provider_model_name, + effective_pool_config, + ) + .await; let mut runtime_by_provider = BTreeMap::new(); let mut pool_config_by_provider = BTreeMap::new(); @@ -634,7 +639,9 @@ impl<'a> PoolKeyCursor<'a> { if !self.score_phase_exhausted { if let Some(score_candidates) = self.next_score_candidates().await { - return Some(score_candidates); + if !score_candidates.is_empty() { + return Some(score_candidates); + } } } @@ -1027,6 +1034,26 @@ impl<'a> PoolKeyCursor<'a> { return None; } + if pool_config.reserve_minimum_quota + && admin_provider_pool_pure::admin_pool_key_minimum_quota_reached( + &key, + self.group.candidate.provider_type.as_str(), + Some(self.group.candidate.selected_provider_model_name.as_str()), + ) + { + self.seen_key_ids.insert(key.id.clone()); + self.record_skip_reason(POOL_ACCOUNT_EXHAUSTED_SKIP_REASON); + self.skipped_candidates + .push(SkippedLocalExecutionCandidate { + candidate: pool_candidate_from_catalog_key(&self.group, key), + skip_reason: POOL_ACCOUNT_EXHAUSTED_SKIP_REASON, + transport: None, + ranking: self.group.ranking.clone(), + extra_data: None, + }); + return None; + } + let candidate = pool_candidate_from_catalog_key(&self.group, key); self.build_eligible_candidate(candidate).await } @@ -1427,15 +1454,23 @@ async fn read_pool_catalog_key_contexts_by_id( state: PlannerAppState<'_>, candidates: &[EligibleLocalExecutionCandidate], provider_model_name: Option<&str>, + effective_pool_config: Option<&AdminProviderPoolConfig>, ) -> BTreeMap { let mut key_ids = Vec::new(); let mut provider_type_by_key_id = BTreeMap::::new(); + let mut reserve_minimum_quota_key_ids = BTreeSet::new(); for candidate in candidates { - if pool_config_for_candidate(candidate).is_none() { + let Some(pool_config) = effective_pool_config + .cloned() + .or_else(|| pool_config_for_candidate(candidate)) + else { continue; - } + }; let key_id = candidate.candidate.key_id.clone(); + if pool_config.reserve_minimum_quota { + reserve_minimum_quota_key_ids.insert(key_id.clone()); + } if let Entry::Vacant(entry) = provider_type_by_key_id.entry(key_id.clone()) { entry.insert(candidate.transport.provider.provider_type.clone()); key_ids.push(key_id); @@ -1487,16 +1522,20 @@ async fn read_pool_catalog_key_contexts_by_id( .get(&key.id) .map(String::as_str) .unwrap_or_default(); - ( - key.id.clone(), - build_pool_catalog_key_context( - state, - &provider_pool_service, + let mut context = build_pool_catalog_key_context( + state, + &provider_pool_service, + &key, + provider_type, + provider_model_name, + ); + context.quota_exhausted |= reserve_minimum_quota_key_ids.contains(&key.id) + && admin_provider_pool_pure::admin_pool_key_minimum_quota_reached( &key, provider_type, provider_model_name, - ), - ) + ); + (key.id.clone(), context) }) .collect::>(); // A key can disappear between the candidate-row and catalog reads. Keep @@ -3962,6 +4001,110 @@ mod tests { })); } + #[tokio::test] + async fn pool_key_cursor_reserve_minimum_quota_filters_pages_and_sticky_hits() { + for reserve_enabled in [false, true] { + for sticky in [false, true] { + for used_percent in [99.0, 98.0, 83.0] { + let provider_config = Some(json!({ + "pool_advanced": { + "reserve_minimum_quota": reserve_enabled, + "skip_exhausted_accounts": false + } + })); + let provider = + sample_codex_pool_provider("provider-pool", 0, provider_config.clone()); + let endpoint = sample_codex_pool_endpoint("provider-pool", "endpoint-1"); + let mut reserved = sample_codex_pool_key("provider-pool", "key-low"); + reserved.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "updated_at": 100, + "allowed": false, + "exhausted": true, + "code": "exhausted", + "windows": [{ + "code": "weekly", + "scope": "account", + "used_ratio": 1.0, + "reset_at": 4_102_444_800u64 + }] + } + })); + reserved.upstream_metadata = Some(json!({ + "codex": { + "updated_at": 200, + "primary_used_percent": used_percent, + "primary_reset_at": 4_102_444_800u64 + } + })); + let ready = sample_codex_pool_key("provider-pool", "key-ready"); + let rows = vec![ + sample_codex_pool_row("provider-pool", "endpoint-1", "key-low", 0), + sample_codex_pool_row("provider-pool", "endpoint-1", "key-ready", 0), + ]; + let data_state = GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], vec![endpoint], vec![reserved, ready], + )), + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)), + ) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state); + let group = + sample_codex_pool_group("provider-pool", "endpoint-1", 0, provider_config); + let pool_config = + pool_config_for_candidate(&group).expect("pool config should parse"); + let sticky_token = sticky.then_some("reserve-session"); + if sticky { + record_admin_provider_pool_success( + app.runtime_state.as_ref(), + "provider-pool", + "key-low", + &pool_config, + sticky_token, + 0, + None, + ) + .await; + } + let mut cursor = PoolKeyCursor::new( + PlannerAppState::new(&app), + group, + sticky_token, + None, + None, + ); + cursor.window_size = 1; + cursor.page_size = 1; + let mut returned = Vec::new(); + while let Some(candidate) = cursor.next_key().await { + returned.push(candidate.candidate.key_id); + } + let reserve_reached = reserve_enabled && used_percent >= 99.0; + assert_eq!( + returned.contains(&"key-low".to_string()), + !reserve_reached, + "reserve={reserve_enabled}, sticky={sticky}, used={used_percent}" + ); + assert!(returned.contains(&"key-ready".to_string())); + if reserve_reached { + assert_eq!( + cursor + .skip_reason_counts + .get(POOL_ACCOUNT_EXHAUSTED_SKIP_REASON), + Some(&1) + ); + } else if sticky { + assert_eq!(returned.first().map(String::as_str), Some("key-low")); + } + } + } + } + } + #[tokio::test] async fn pool_key_cursor_does_not_spend_effective_scan_budget_on_exhausted_accounts() { let provider_config = Some(json!({ @@ -4168,6 +4311,115 @@ mod tests { ); } + #[tokio::test] + async fn inactive_pool_key_with_stale_score_does_not_exhaust_pool() { + let provider_config = Some(json!({ + "pool_advanced": { + "score_top_n": 128, + "scheduling_presets": [ + {"preset": "single_account", "enabled": true}, + {"preset": "priority_first", "enabled": true} + ] + } + })); + let (provider, endpoint, mut keys, mut rows) = + large_pool_fixture(2, provider_config.clone()); + keys[1].is_active = false; + rows.retain(|row| row.key_id != "key-00001"); + let scores = vec![ + sample_provider_key_pool_score("provider-pool", "key-00000", 5.0), + sample_provider_key_pool_score("provider-pool", "key-00001", 20.0), + ]; + let data_state = + GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + keys, + )), + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)), + ) + .with_pool_score_repository_for_tests(Arc::new( + InMemoryPoolMemberScoreRepository::seed(scores), + )) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state); + let group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + provider_config, + ); + let mut cursor = PoolKeyCursor::new(PlannerAppState::new(&app), group, None, None, None); + + let candidate = cursor + .next_key() + .await + .expect("active key must stay schedulable beside a stale inactive score"); + + assert_eq!(candidate.candidate.key_id, "key-00000"); + assert_eq!( + cursor.skip_reason_counts.get("pool_score_member_missing"), + Some(&1) + ); + } + + #[tokio::test] + async fn stale_inactive_score_only_does_not_exhaust_pool() { + let provider_config = Some(json!({ + "pool_advanced": { + "score_top_n": 128, + "scheduling_presets": [ + {"preset": "single_account", "enabled": true}, + {"preset": "priority_first", "enabled": true} + ] + } + })); + let (provider, endpoint, mut keys, mut rows) = + large_pool_fixture(2, provider_config.clone()); + keys[1].is_active = false; + rows.retain(|row| row.key_id != "key-00001"); + let scores = vec![sample_provider_key_pool_score( + "provider-pool", + "key-00001", + 20.0, + )]; + let data_state = + GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + keys, + )), + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)), + ) + .with_pool_score_repository_for_tests(Arc::new( + InMemoryPoolMemberScoreRepository::seed(scores), + )) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state); + let group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + provider_config, + ); + let mut cursor = PoolKeyCursor::new(PlannerAppState::new(&app), group, None, None, None); + + let candidate = cursor + .next_key() + .await + .expect("catalog rows must remain schedulable when the only score is stale"); + + assert_eq!(candidate.candidate.key_id, "key-00000"); + } + #[tokio::test] async fn score_candidates_continue_across_pool_windows() { let provider_config = Some(json!({ @@ -4913,15 +5165,6 @@ mod tests { )) } - fn provider_catalog_credential_state() -> AppState { - AppState::new() - .expect("credential state should build") - .with_data_state_for_tests( - GatewayDataState::disabled() - .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY), - ) - } - fn large_pool_fixture( key_count: usize, provider_config: Option, @@ -4972,18 +5215,12 @@ mod tests { ) .expect("endpoint transport should build"); - let credential_state = provider_catalog_credential_state(); + // 这些用例只验证池扫描、跳过计数和游标预算,不会发起请求或读取凭据。 + // 留空凭据可跳过无关的 Fernet 加解密,同时避免复用绑定密文破坏 key_id AAD。 let mut keys = Vec::with_capacity(key_count); let mut rows = Vec::with_capacity(key_count); for index in 0..key_count { let key_id = format!("key-{index:05}"); - let encrypted_api_key = credential_state - .seal_provider_catalog_key_api_key( - "provider-pool", - &key_id, - &format!("secret-{index}"), - ) - .expect("api key should encrypt"); let mut key = StoredProviderCatalogKey::new( key_id.clone(), "provider-pool".to_string(), @@ -4995,7 +5232,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(json!(["openai:chat"])), - encrypted_api_key, + None, None, None, None, @@ -5130,10 +5367,8 @@ mod tests { .expect("endpoint transport should build") } + /// 这些测试只检查池调度状态,不涉及凭据解密,因此不构造无关的密文。 fn sample_codex_pool_key(provider_id: &str, key_id: &str) -> StoredProviderCatalogKey { - let encrypted_api_key = provider_catalog_credential_state() - .seal_provider_catalog_key_api_key(provider_id, key_id, &format!("secret-{key_id}")) - .expect("api key should encrypt"); let mut key = StoredProviderCatalogKey::new( key_id.to_string(), provider_id.to_string(), @@ -5145,7 +5380,7 @@ mod tests { .expect("key should build") .with_transport_fields( Some(json!(["openai:responses"])), - encrypted_api_key, + None, None, None, Some(json!({"openai:responses": 1})), diff --git a/apps/aether-gateway/src/execution_runtime/grok.rs b/apps/aether-gateway/src/execution_runtime/grok.rs index ab95832f8..2f4500068 100644 --- a/apps/aether-gateway/src/execution_runtime/grok.rs +++ b/apps/aether-gateway/src/execution_runtime/grok.rs @@ -3199,13 +3199,15 @@ fn openai_responses_body( let response_id = format!("resp_{}", Uuid::new_v4()); let mut output = Vec::new(); if !collected.thinking.trim().is_empty() { + let thinking = collected.thinking.trim(); output.push(json!({ "id": openai_responses_synthetic_reasoning_item_id(&response_id, 0), "type": "reasoning", "status": "completed", - "summary": [{ - "type": "summary_text", - "text": collected.thinking.trim(), + "summary": [], + "content": [{ + "type": "reasoning_text", + "text": thinking, }], })); } @@ -4723,6 +4725,15 @@ mod tests { serde_json::json!(usage.reasoning_tokens) ); assert_eq!(body["output"][0]["type"], serde_json::json!("reasoning")); + assert_eq!( + body["output"][0]["content"][0]["type"], + serde_json::json!("reasoning_text") + ); + assert_eq!( + body["output"][0]["content"][0]["text"], + serde_json::json!("short reasoning") + ); + assert_eq!(body["output"][0]["summary"], serde_json::json!([])); assert_eq!(body["output"][1]["type"], serde_json::json!("message")); assert!(body["output"][1]["id"] .as_str() @@ -4906,7 +4917,12 @@ mod tests { assert!(body.contains("event: response.created")); assert!(body.contains("event: response.in_progress")); - assert!(body.contains("event: response.reasoning_summary_part.added")); + // Thinking must stay off the summary channel or clients that render + // both (Codex) print the raw chain-of-thought twice. + assert!(!body.contains("event: response.reasoning_summary_part.added")); + assert!(!body.contains("event: response.reasoning_summary_text.delta")); + assert!(!body.contains("event: response.reasoning_summary_text.done")); + assert!(body.contains("\"type\":\"reasoning_text\"")); assert!(body.contains("event: response.content_part.added")); assert!(body.contains("event: response.output_text.done")); assert!(body.contains("event: response.completed")); diff --git a/apps/aether-gateway/src/execution_runtime/stream/execution.rs b/apps/aether-gateway/src/execution_runtime/stream/execution.rs index 2ae4e283e..f6eb44e44 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/execution.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/execution.rs @@ -6513,6 +6513,21 @@ async fn execute_stream_from_frame_stream_with_retry_scope( let normalized_stream_report_context = normalize_provider_private_report_context(report_context.as_ref()); + // Observers follow the live protocol stream across prefetch and transfer. + // Diagnostic capture limits must never determine parser state. + let stream_usage_report_context = normalized_stream_report_context.clone().or_else(|| { + Some(json!({ + "provider_api_format": plan.provider_api_format.as_str(), + "client_api_format": plan.client_api_format.as_str(), + })) + }); + let mut stream_usage_observer = stream_usage_report_context + .as_ref() + .map(|_| StreamingStandardTerminalObserver::default()); + let mut stream_usage_observer_buffered = + StreamUsageObservationBuffer::new(max_stream_body_buffer_bytes); + let mut provider_error_inspection = ProviderStreamErrorInspection::default(); + let mut prefetched_provider_error = None; let upstream_headers = headers.clone(); let mut private_stream_normalizer = maybe_build_provider_private_stream_normalizer(report_context.as_ref()); @@ -6613,7 +6628,8 @@ async fn execute_stream_from_frame_stream_with_retry_scope( stream_commit_gate.commit(); } let mut prefetched_chunks: Vec = Vec::new(); - let mut provider_prefetched_body = Vec::new(); + let mut provider_prefetched_body = StreamBodyCapture::default(); + let mut provider_prefetched_bytes = 0_u64; let mut provider_prefetched_body_truncated = false; let mut prefetched_body = Vec::new(); let mut prefetched_inspection_body = Vec::new(); @@ -6866,10 +6882,12 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } } - append_stream_capture_bytes( + provider_prefetched_bytes = + provider_prefetched_bytes.saturating_add(chunk.len() as u64); + append_budgeted_stream_capture_bytes( &mut provider_prefetched_body, &chunk, - MAX_STREAM_PREFETCH_BYTES, + max_stream_body_buffer_bytes, &mut provider_prefetched_body_truncated, ); append_stream_capture_bytes( @@ -7114,6 +7132,22 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } else { chunk }; + if let Some(error) = provider_error_inspection + .observe(stream_usage_report_context.as_ref(), &normalized_chunk) + { + prefetched_provider_error.get_or_insert(error); + } + if let (Some(observer), Some(context)) = ( + stream_usage_observer.as_mut(), + stream_usage_report_context.as_ref(), + ) { + observe_stream_usage_bytes( + observer, + context, + &mut stream_usage_observer_buffered, + &normalized_chunk, + ); + } let rewritten_chunk = if let Some(rewriter) = local_stream_rewriter.as_mut() { match rewriter.push_chunk(&normalized_chunk) { Ok(rewritten_chunk) => rewritten_chunk, @@ -7244,17 +7278,21 @@ async fn execute_stream_from_frame_stream_with_retry_scope( if stream_commit_gate.is_uncommitted() { stream_commit_gate.commit(); } - let prefetched_response_history_persisted = if let Some(record) = local_stream_rewriter + if let Some(record) = local_stream_rewriter .as_mut() .and_then(|rewriter| rewriter.take_response_history_record()) { crate::ai_serving::persist_response_history_record(state, record).await; - true - } else { - false - }; - drop(private_stream_normalizer); - drop(local_stream_rewriter); + } + // Keep partial records and conversion state; replaying the bounded + // inspection/capture prefix loses any bytes consumed beyond that prefix. + let mut private_stream_normalizer = private_stream_normalizer.map(|parser| parser.into_owned()); + let mut local_stream_rewriter = local_stream_rewriter.map(|parser| parser.into_owned()); + if sync_json_stream_bridge_active { + private_stream_normalizer = None; + local_stream_rewriter = None; + stream_usage_observer = None; + } let initial_usage_telemetry = prefetched_usage_telemetry.clone().or_else(|| { prefetched_telemetry @@ -7297,7 +7335,6 @@ async fn execute_stream_from_frame_stream_with_retry_scope( let headers_for_report = headers.clone(); let report_kind_owned = report_kind; let report_context_owned = report_context; - let normalized_stream_report_context_owned = normalized_stream_report_context; let lifecycle_seed_for_report = lifecycle_seed; let provider_prefetched_body_for_report = provider_prefetched_body; let prefetched_body_for_report = prefetched_body; @@ -7339,40 +7376,10 @@ async fn execute_stream_from_frame_stream_with_retry_scope( let _stream_total_guard = StageElapsedGuard::from_started_at("stream_total", stream_started_at_for_report); let _provider_pool_in_flight_guard = provider_pool_in_flight_guard_for_report; - let mut provider_buffered_body = StreamBodyCapture::default(); + let mut provider_buffered_body = provider_prefetched_body_for_report; let mut buffered_body = StreamBodyCapture::default(); - let mut provider_body_truncated = false; + let mut provider_body_truncated = provider_prefetched_body_truncated; let mut client_body_truncated = false; - let mut private_stream_normalizer = if sync_json_stream_bridge_active_for_report { - None - } else { - maybe_build_provider_private_stream_normalizer(report_context_owned.as_ref()) - }; - let mut local_stream_rewriter = if sync_json_stream_bridge_active_for_report { - None - } else { - maybe_build_stream_response_rewriter(normalized_stream_report_context_owned.as_ref()) - }; - let stream_usage_report_context = - normalized_stream_report_context_owned.clone().or_else(|| { - Some(serde_json::json!({ - "provider_api_format": plan_for_report.provider_api_format.as_str(), - "client_api_format": plan_for_report.client_api_format.as_str(), - })) - }); - let mut stream_usage_observer = stream_usage_report_context - .as_ref() - .filter(|_| !sync_json_stream_bridge_active_for_report) - .map(|_| StreamingStandardTerminalObserver::default()); - let mut stream_usage_observer_buffered = - StreamUsageObservationBuffer::new(max_stream_body_buffer_bytes); - let mut provider_error_inspection = ProviderStreamErrorInspection::default(); - append_budgeted_stream_capture_bytes( - &mut provider_buffered_body, - &provider_prefetched_body_for_report, - max_stream_body_buffer_bytes, - &mut provider_body_truncated, - ); append_budgeted_stream_capture_bytes( &mut buffered_body, &prefetched_body_for_report, @@ -7416,9 +7423,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } else { initial_elapsed_ms })); - let provider_stream_bytes = Arc::new(AtomicU64::new( - u64::try_from(provider_prefetched_body_for_report.len()).unwrap_or(u64::MAX), - )); + let provider_stream_bytes = Arc::new(AtomicU64::new(provider_prefetched_bytes)); let client_stream_bytes = Arc::new(AtomicU64::new( u64::try_from(prefetched_body_for_report.len()).unwrap_or(u64::MAX), )); @@ -7514,96 +7519,20 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } }) }; - if !provider_prefetched_body_for_report.is_empty() { - let normalized_prefetched_chunk = if let Some(normalizer) = - private_stream_normalizer.as_mut() - { - match normalizer.push_chunk(&provider_prefetched_body_for_report) { - Ok(normalized_chunk) => Some(normalized_chunk), - Err(err) => { - warn!( - event_name = "stream_execution_prefetch_normalize_restore_failed", - log_type = "ops", - trace_id = %trace_id_owned, - request_id = %request_id_for_report_log, - candidate_id = ?candidate_id_for_report.as_deref(), - error_category = "stream_normalization_restore_failed", - "gateway failed to restore private stream normalization state after prefetch" - ); - terminal_failure = Some(build_stream_failure_report( - "execution_runtime_stream_rewrite_error", - format!( - "failed to restore private stream normalization state after prefetch: {err:?}" - ), - 502, - )); - None - } - } - } else { - None - }; - let replay_chunk = normalized_prefetched_chunk - .as_deref() - .unwrap_or(provider_prefetched_body_for_report.as_slice()); - if let Some(error_body_json) = provider_error_inspection - .observe(stream_usage_report_context.as_ref(), replay_chunk) - { - provider_error_forwarded_to_client = !prefetched_body_for_report.is_empty(); - let error_status_code = resolve_provider_stream_error_status_code( - plan_for_report.provider_api_format.as_str(), - status_code, - &error_body_json, - ); - terminal_failure = Some(build_stream_failure_from_provider_error_body( - error_status_code, - &error_body_json, - )); - } - if let (Some(observer), Some(report_context)) = ( - stream_usage_observer.as_mut(), - stream_usage_report_context.as_ref(), - ) { - observe_stream_usage_bytes( - observer, - report_context, - &mut stream_usage_observer_buffered, - replay_chunk, - ); - } - if terminal_failure.is_none() { - if let Some(rewriter) = local_stream_rewriter.as_mut() { - if let Err(err) = rewriter.push_chunk(replay_chunk) { - warn!( - event_name = "stream_execution_prefetch_rewrite_restore_failed", - log_type = "ops", - trace_id = %trace_id_owned, - request_id = %request_id_for_report_log, - candidate_id = ?candidate_id_for_report.as_deref(), - error_category = "stream_rewrite_restore_failed", - "gateway failed to restore local stream rewrite state after prefetch" - ); - terminal_failure = Some(build_stream_failure_report( - "execution_runtime_stream_rewrite_error", - format!( - "failed to restore local stream rewrite state after prefetch: {err:?}" - ), - 502, - )); - } - } - } - if prefetched_response_history_persisted { - if let Some(rewriter) = local_stream_rewriter.as_mut() { - let _ = rewriter.take_response_history_record(); - } - } + if let Some(error_body_json) = prefetched_provider_error { + provider_error_forwarded_to_client = !prefetched_body_for_report.is_empty(); + let error_status_code = resolve_provider_stream_error_status_code( + plan_for_report.provider_api_format.as_str(), + status_code, + &error_body_json, + ); + terminal_failure = Some(build_stream_failure_from_provider_error_body( + error_status_code, + &error_body_json, + )); } - - // These buffers restore parser/rewriter state above. Audit capture owns - // its budgeted copies; retaining semantic prefetch duplicates for the - // rest of the stream would bypass the capture memory limit. - drop(provider_prefetched_body_for_report); + // Parser state is already current and capture owns its budgeted bytes. + // This output prefix is needed only to initialize client-side trackers. drop(prefetched_body_for_report); if terminal_failure.is_none() && !reached_eof { @@ -9363,6 +9292,188 @@ mod tests { .unwrap() } + #[tokio::test] + async fn prefetch_handoff_preserves_large_responses_setup_event() { + let event = format!( + "event: response.created\ndata: {}\n\n", + json!({"type":"response.created", "response": { + "id":"resp-large-setup", "status":"in_progress", "output":[], + "tools":[{"name":"write", "description":"x".repeat(64 * 1024)}] + }}) + ); + let done = "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-large-setup\",\"status\":\"completed\",\"output\":[],\"usage\":{\"input_tokens\":7,\"output_tokens\":2}}}\n\n"; + // Include the two observed transport boundaries, exact/near budget + // boundaries, and multiple prefetch chunks crossing the budget. + for cuts in [ + vec![16_383], + vec![16_384], + vec![17_735], + vec![17_741], + vec![8_192, 17_735], + ] { + let mut chunks = Vec::new(); + let mut start = 0; + for end in cuts { + chunks.push(&event[start..end]); + start = end; + } + chunks.push(&event[start..]); + chunks.push(done); + let response = execute_generic_sse_precommit(chunks, json!({}), None, false) + .await + .expect("large setup should commit at the bounded prefetch limit"); + let body = to_bytes(response.into_body(), usize::MAX).await.unwrap(); + let body = String::from_utf8(body.to_vec()).unwrap(); + assert!( + body.starts_with(&event), + "setup bytes lost or duplicated at split {start}" + ); + let events: Vec = body + .lines() + .filter_map(|line| line.strip_prefix("data: ")) + .filter(|payload| *payload != "[DONE]") + .map(|payload| { + serde_json::from_str(payload).expect("every SSE payload must be valid JSON") + }) + .collect(); + assert_eq!(events.len(), 2, "events must be forwarded exactly once"); + assert_eq!(events[1]["type"], "response.completed"); + } + } + + #[tokio::test] + async fn prefetch_handoff_keeps_audit_usage_and_private_conversion() { + for private in [false, true] { + let request_id = format!("handoff-audit-{}", uuid::Uuid::new_v4()); + let mut plan = if private { + antigravity_gemini_stream_plan(&request_id) + } else { + native_anthropic_stream_plan(&request_id) + }; + if !private { + plan.provider_api_format = "openai:responses".into(); + plan.client_api_format = "openai:responses".into(); + } + let context = json!({ + "request_id": request_id, "candidate_id": plan.candidate_id, + "candidate_index":0, "retry_index":0, + "provider_api_format": plan.provider_api_format, + "client_api_format": plan.client_api_format, + "needs_conversion": private, "has_envelope": private, + "envelope_name": if private { "antigravity:v1internal" } else { "" }, + }); + let repository = Arc::new(InMemoryUsageReadRepository::default()); + let catalog = provider_catalog_for_plan(&plan, None); + let state = AppState::new() + .unwrap() + .with_data_state_for_tests( + crate::data::GatewayDataState::with_usage_repository_for_tests(Arc::clone( + &repository, + )) + .with_provider_catalog_reader(Arc::new(catalog)) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_system_config_values_for_tests([( + "request_record_level".into(), + json!("full"), + )]), + ) + .with_usage_runtime_for_tests(UsageRuntimeConfig { + enabled: true, + ..Default::default() + }); + let text = "hello".repeat(12_000); + let payload = if private { + json!({"response":{"candidates":[{"content":{"role":"model","parts":[{"text":text}]}, + "finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":1234,"candidatesTokenCount":567}, + "modelVersion":"gemini-3.7-flash-tiered"}}) + } else { + json!({"type":"response.completed","response":{"id":"resp-handoff-usage","status":"completed", + "output":[{"type":"message","id":"msg-handoff","role":"assistant","status":"completed", + "content":[{"type":"output_text","text":text,"annotations":[]}]}], + "usage":{"input_tokens":1234,"output_tokens":567,"total_tokens":1801}}}) + }; + let input = format!("data: {payload}\n\n"); + // One complete large chunk exercises an already-emitted prefetch + // result; the private path exercises incomplete normalization too. + let chunks = if private { + vec![input[..17_735].to_string(), input[17_735..].to_string()] + } else { + vec![input.clone()] + }; + let frames = stream! { + yield Ok::(ndjson_frame(StreamFrame { + frame_type:StreamFrameType::Headers, + payload:StreamFramePayload::Headers { status_code:200, + headers:BTreeMap::from([("content-type".into(),"text/event-stream".into())]), + response_observation:None }, + })); + for chunk in chunks { + yield Ok(ndjson_frame(StreamFrame { frame_type:StreamFrameType::Data, + payload:StreamFramePayload::Data { text:Some(chunk),chunk_b64:None } })); + } + yield Ok(ndjson_frame(StreamFrame::eof())); + }.boxed(); + let response = execute_stream_from_frame_stream( + &state, + plan, + "trace-handoff-audit", + &test_decision(), + OPENAI_RESPONSES_STREAM_PLAN_KIND, + Some("openai_responses_stream_success".into()), + Some(context), + crate::clock::current_unix_ms(), + Instant::now(), + RequestStageTrace::from_env(), + false, + frames, + None, + ) + .await + .unwrap() + .unwrap(); + let body = to_bytes(response.into_body(), usize::MAX).await.unwrap(); + let body = String::from_utf8(body.to_vec()).unwrap(); + let events: Vec = body + .lines() + .filter_map(|l| l.strip_prefix("data: ")) + .filter(|p| *p != "[DONE]") + .map(|p| serde_json::from_str(p).unwrap()) + .collect(); + assert_eq!( + events + .iter() + .filter(|e| e["type"] == "response.completed") + .count(), + 1 + ); + assert!(body.contains(&text)); + let usage = tokio::time::timeout(Duration::from_secs(3), async { + loop { + if let Some(u) = repository + .find_by_request_id(&request_id) + .await + .unwrap() + .filter(|u| u.status == "completed" || u.status == "failed") + { + break u; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("usage should finalize"); + assert_eq!(usage.status, "completed", "{:?}", usage.error_message); + assert_eq!(usage.input_tokens, 1234); + assert_eq!(usage.output_tokens, 567); + let captured = usage.response_body.as_ref().expect("provider capture"); + assert!( + captured["metadata"].get("dropped_chunks").is_none(), + "{captured}" + ); + assert_eq!(captured["chunks"].as_array().unwrap(), &vec![payload]); + } + } + #[tokio::test] async fn generic_stream_success_regex_matches_fragmented_plain_body() { for chunks in [ @@ -9901,7 +10012,7 @@ mod tests { let mut buffer = super::StreamUsageObservationBuffer::new(32 * 1024); let mut rewriter = super::maybe_build_stream_response_rewriter(Some(&context)).unwrap(); let mut delivered = Vec::new(); - for chunk in chunks { + for (index, chunk) in chunks.into_iter().enumerate() { provider.append(chunk, 32 * 1024, &mut provider_truncated); super::observe_stream_usage_bytes( observer.as_mut().unwrap(), @@ -9912,6 +10023,10 @@ mod tests { let output = rewriter.push_chunk(chunk).unwrap(); client.append(&output, 32 * 1024, &mut client_truncated); delivered.extend(output); + if index == 0 { + // Task handoff must also work when audit admits no bytes. + rewriter = rewriter.into_owned(); + } } let tail = rewriter.finish().unwrap(); client.append(&tail, 32 * 1024, &mut client_truncated); @@ -11931,7 +12046,11 @@ mod tests { .expect("response body should read"); let body = String::from_utf8(body.to_vec()).expect("response body should be utf8"); assert!( - body.contains("event: response.reasoning_summary_text.delta\n"), + body.contains("event: response.reasoning_text.delta\n"), + "{body}" + ); + assert!( + !body.contains("event: response.reasoning_summary_text.delta\n"), "{body}" ); assert!( diff --git a/apps/aether-gateway/src/execution_runtime/submission.rs b/apps/aether-gateway/src/execution_runtime/submission.rs index 3ee844002..ff1d72fcf 100644 --- a/apps/aether-gateway/src/execution_runtime/submission.rs +++ b/apps/aether-gateway/src/execution_runtime/submission.rs @@ -196,6 +196,78 @@ fn maybe_build_invalid_provider_success_finalize_response( )?)) } +fn local_sync_needs_conversion(payload: &GatewaySyncReportRequest) -> bool { + payload + .report_context + .as_ref() + .and_then(|value| value.get("needs_conversion")) + .and_then(|value| value.as_bool()) + .unwrap_or(false) +} + +/// A successful upstream response that needed conversion but could not be +/// converted must not reach the client in the provider's own format. +fn maybe_build_unconverted_cross_format_success_response( + trace_id: &str, + decision: &GatewayControlDecision, + payload: &GatewaySyncReportRequest, +) -> Result>, GatewayError> { + if payload.status_code >= 400 + || !local_sync_needs_conversion(payload) + || !is_core_error_finalize_kind(payload.report_kind.as_str()) + { + return Ok(None); + } + + let client_api_format = resolve_local_sync_client_api_format(payload); + let provider_api_format = resolve_local_sync_provider_api_format(payload); + warn!( + event_name = "local_core_finalize_cross_format_success_unconverted", + log_type = "event", + trace_id = %trace_id, + report_kind = %payload.report_kind, + status_code = payload.status_code, + client_api_format = %client_api_format, + provider_api_format = %provider_api_format, + "gateway could not convert a successful provider response to the client format" + ); + let message = format!( + "Provider returned HTTP {} but its {provider_api_format} response could not be converted to {client_api_format}.", + payload.status_code + ); + let body_json = build_core_error_body_for_client_format( + &client_api_format, + &message, + Some("response_conversion_failed"), + LocalCoreSyncErrorKind::ServerError, + ) + .unwrap_or_else(|| { + serde_json::json!({ + "error": { + "message": message, + "type": "server_error", + "code": "response_conversion_failed" + } + }) + }); + + let mut response_headers = payload.headers.clone(); + response_headers.remove("content-encoding"); + response_headers.remove("content-length"); + response_headers.insert("content-type".to_string(), "application/json".to_string()); + let body_bytes = + serde_json::to_vec(&body_json).map_err(|err| GatewayError::Internal(err.to_string()))?; + response_headers.insert("content-length".to_string(), body_bytes.len().to_string()); + + Ok(Some(build_client_response_from_parts( + StatusCode::BAD_GATEWAY.as_u16(), + &response_headers, + Body::from(body_bytes), + trace_id, + Some(decision), + )?)) +} + fn local_core_sync_finalize_has_invalid_provider_success( payload: &GatewaySyncReportRequest, ) -> Result { @@ -274,6 +346,12 @@ pub(crate) fn resolve_local_core_error_response_body_json( return Ok(Some(body_json)); } + // A 2xx cross-format body that is not JSON (e.g. an aggregated SSE capture) + // carries no upstream error; wrapping it as one would ship raw provider + // bytes to the client under the success status. + if payload.status_code < 400 && local_sync_needs_conversion(payload) { + return Ok(None); + } let Some(body_text) = decode_local_sync_body_text(payload)? else { return Ok(None); }; @@ -626,6 +704,10 @@ pub(crate) async fn submit_local_core_error_or_sync_finalize( maybe_build_local_core_error_response(trace_id, decision, &payload)? { response + } else if let Some(response) = + maybe_build_unconverted_cross_format_success_response(trace_id, decision, &payload)? + { + response } else { warn!( event_name = "local_core_finalize_fallback_raw_response_body", @@ -937,6 +1019,128 @@ mod tests { ); } + #[tokio::test] + async fn local_core_sync_finalize_converts_forced_responses_stream_for_gemini_client() { + use base64::Engine as _; + + // Forced-stream xAI shape: the terminal response echoes request + // metadata and encrypted reasoning next to the real answer. + let raw_sse = concat!( + "event: response.created\n", + "data: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"id\":\"resp_xai_123\",\"object\":\"response\",\"status\":\"in_progress\",\"model\":\"grok-4.7-build\",\"output\":[],\"parallel_tool_calls\":true,\"tools\":[]}}\n\n", + "event: response.output_item.done\n", + "data: {\"type\":\"response.output_item.done\",\"sequence_number\":1,\"output_index\":0,\"item\":{\"id\":\"rs_xai_123\",\"type\":\"reasoning\",\"status\":\"completed\",\"summary\":[],\"encrypted_content\":\"opaque-xai-reasoning\"}}\n\n", + "event: response.output_text.delta\n", + "data: {\"type\":\"response.output_text.delta\",\"sequence_number\":2,\"item_id\":\"msg_xai_123\",\"output_index\":1,\"content_index\":0,\"delta\":\"Hi there, friend\"}\n\n", + "event: response.output_item.done\n", + "data: {\"type\":\"response.output_item.done\",\"sequence_number\":3,\"output_index\":1,\"item\":{\"id\":\"msg_xai_123\",\"type\":\"message\",\"status\":\"completed\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"Hi there, friend\",\"annotations\":[]}]}}\n\n", + "event: response.completed\n", + "data: {\"type\":\"response.completed\",\"sequence_number\":4,\"response\":{\"id\":\"resp_xai_123\",\"object\":\"response\",\"status\":\"completed\",\"model\":\"grok-4.7-build\",\"output\":[],\"parallel_tool_calls\":true,\"tool_choice\":\"auto\",\"tools\":[],\"text\":{\"format\":{\"type\":\"text\"}},\"temperature\":0.7,\"store\":false,\"usage\":{\"input_tokens\":1249,\"output_tokens\":12,\"total_tokens\":1261}}}\n\n", + ); + let mut payload = core_finalize_payload( + "gemini_chat_sync_finalize", + "gemini:generate_content", + "openai:responses", + 200, + json!(null), + ); + payload.body_json = None; + payload.body_base64 = Some(base64::engine::general_purpose::STANDARD.encode(raw_sse)); + payload.report_context = Some(json!({ + "client_api_format": "gemini:generate_content", + "provider_api_format": "openai:responses", + "provider_stream_event_api_format": "openai:responses", + "model": "grok-4.7", + "mapped_model": "grok-4.7", + "needs_conversion": true, + })); + + let state = AppState::new().expect("state should build"); + let response = submit_local_core_error_or_sync_finalize( + &state, + "trace-forced-responses-gemini", + &test_decision(), + payload, + ) + .await + .expect("finalize should build a response"); + + assert_eq!(response.status(), http::StatusCode::OK); + let body_bytes = to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"); + let body = + serde_json::from_slice::(&body_bytes).expect("body should decode"); + assert!(body.get("error").is_none(), "unexpected error body: {body}"); + let parts = body["candidates"][0]["content"]["parts"] + .as_array() + .expect("gemini parts"); + assert!(parts.iter().any(|part| part["text"] == "Hi there, friend")); + let text = String::from_utf8_lossy(&body_bytes); + assert!(!text.contains("opaque-xai-reasoning") && !text.contains("response.created")); + } + + #[tokio::test] + async fn local_core_sync_finalize_never_wraps_unconvertible_success_sse_as_client_error() { + use base64::Engine as _; + + // A complete stream whose output the Gemini client cannot represent. + let raw_sse = concat!( + "event: response.output_item.done\n", + "data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"future_item_123\",\"type\":\"future_output\",\"payload\":\"must-not-drop\"}}\n\n", + "event: response.completed\n", + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_raw_123\",\"object\":\"response\",\"status\":\"completed\",\"model\":\"grok-4.7\",\"output\":[]}}\n\n", + ); + let mut payload = core_finalize_payload( + "gemini_chat_sync_finalize", + "gemini:generate_content", + "openai:responses", + 200, + json!(null), + ); + payload.body_json = None; + payload.body_base64 = Some(base64::engine::general_purpose::STANDARD.encode(raw_sse)); + payload.report_context = Some(json!({ + "client_api_format": "gemini:generate_content", + "provider_api_format": "openai:responses", + "provider_stream_event_api_format": "openai:responses", + "needs_conversion": true, + })); + + assert!(maybe_build_local_core_error_response( + "trace-raw-success-sse", + &test_decision(), + &payload, + ) + .expect("response build should not error") + .is_none()); + + let state = AppState::new().expect("state should build"); + let response = submit_local_core_error_or_sync_finalize( + &state, + "trace-raw-success-sse", + &test_decision(), + payload, + ) + .await + .expect("finalize should build a response"); + + assert_eq!(response.status(), http::StatusCode::BAD_GATEWAY); + let body = serde_json::from_slice::( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"), + ) + .expect("body should decode"); + let message = body["error"]["message"] + .as_str() + .expect("error message should exist"); + assert!( + message.contains("could not be converted") && !message.contains("must-not-drop"), + "unexpected message: {message}" + ); + } + #[tokio::test] async fn submit_local_core_finalize_keeps_http_200_for_success_image_body() { let payload = core_finalize_payload( diff --git a/apps/aether-gateway/src/execution_runtime/transport.rs b/apps/aether-gateway/src/execution_runtime/transport.rs index 5cfeb3cf4..c5c52ae26 100644 --- a/apps/aether-gateway/src/execution_runtime/transport.rs +++ b/apps/aether-gateway/src/execution_runtime/transport.rs @@ -9737,7 +9737,10 @@ mod tests { headers: BTreeMap::from([("content-type".into(), "application/json".into())]), content_type: Some("application/json".into()), content_encoding: Some(encoding.into()), - body: RequestBody::from_json(json!({"model": "gpt-4.1"})), + body: RequestBody::from_json(json!({ + "model": "gpt-4.1", + "service_tier": "ultrafast" + })), stream: false, client_api_format: "openai:chat".into(), provider_api_format: "openai:chat".into(), @@ -9758,7 +9761,7 @@ mod tests { result.body.and_then(|body| body.json_body), Some(json!({ "content_encoding": encoding, - "body": {"model": "gpt-4.1"}, + "body": {"model": "gpt-4.1", "service_tier": "ultrafast"}, })) ); } diff --git a/apps/aether-gateway/src/executor/orchestration.rs b/apps/aether-gateway/src/executor/orchestration.rs index e79f96038..0473bfa3f 100644 --- a/apps/aether-gateway/src/executor/orchestration.rs +++ b/apps/aether-gateway/src/executor/orchestration.rs @@ -1464,6 +1464,22 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision( .await } +fn supports_local_video_get( + parts: &http::request::Parts, + decision: &GatewayControlDecision, +) -> bool { + parts.method == http::Method::GET + && decision.route_kind.as_deref() == Some("video") + && (crate::video_tasks::resolve_video_task_read_lookup_key( + decision.route_family.as_deref(), + parts.uri.path(), + ) + .is_some() + || (decision.route_family.as_deref() == Some("openai") + && crate::video_tasks::extract_openai_task_id_from_content_path(parts.uri.path()) + .is_some())) +} + pub(crate) fn maybe_execute_sync_request<'a>( state: &'a AppState, parts: &'a http::request::Parts, @@ -1477,7 +1493,7 @@ pub(crate) fn maybe_execute_sync_request<'a>( }; #[cfg(not(test))] { - if parts.method != http::Method::POST { + if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) { return Ok(LocalExecutionRequestOutcome::NoPath); } return maybe_execute_sync_local_path(state, parts, body_bytes, trace_id, decision) @@ -1490,6 +1506,7 @@ pub(crate) fn maybe_execute_sync_request<'a>( .unwrap_or_default() .is_empty() && parts.method != http::Method::POST + && !supports_local_video_get(parts, decision) { return Ok(LocalExecutionRequestOutcome::NoPath); } @@ -1511,7 +1528,7 @@ pub(crate) fn maybe_execute_stream_request<'a>( }; #[cfg(not(test))] { - if parts.method != http::Method::POST { + if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) { return Ok(LocalExecutionRequestOutcome::NoPath); } return maybe_execute_stream_local_path(state, parts, body_bytes, trace_id, decision) @@ -1524,6 +1541,7 @@ pub(crate) fn maybe_execute_stream_request<'a>( .unwrap_or_default() .is_empty() && parts.method != http::Method::POST + && !supports_local_video_get(parts, decision) { return Ok(LocalExecutionRequestOutcome::NoPath); } diff --git a/apps/aether-gateway/src/frontdoor_loop_guard.rs b/apps/aether-gateway/src/frontdoor_loop_guard.rs index 52c50feff..b3824be3d 100644 --- a/apps/aether-gateway/src/frontdoor_loop_guard.rs +++ b/apps/aether-gateway/src/frontdoor_loop_guard.rs @@ -32,6 +32,10 @@ fn request_has_execution_runtime_via_guard(headers: &HeaderMap) -> bool { } pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool { + let path = path + .strip_prefix("/openai") + .filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/")) + .unwrap_or(path); matches!( path, "/v1/messages" diff --git a/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/adjust.rs b/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/adjust.rs index 8efabaa0b..f3d2809af 100644 --- a/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/adjust.rs +++ b/apps/aether-gateway/src/handlers/admin/billing/wallets/mutations/adjust.rs @@ -63,13 +63,14 @@ pub(in super::super) async fn build_admin_wallet_adjust_response( } let operator_id = admin_wallet_operator_id(request_context); let has_wallet_writer = state.has_wallet_data_writer(); - let Some((wallet, transaction)) = state + let Some((wallet, Some(transaction))) = state .admin_adjust_wallet_balance( &wallet_id, amount_usd, &balance_type, operator_id.as_deref(), description.as_deref(), + false, ) .await? else { diff --git a/apps/aether-gateway/src/handlers/admin/observability/mod.rs b/apps/aether-gateway/src/handlers/admin/observability/mod.rs index c2bd1461c..118c6b748 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/mod.rs @@ -12,3 +12,138 @@ pub(crate) use self::stats::{ }; pub(crate) use self::stats::{AdminStatsTimeRange, AdminStatsUsageFilter}; pub(crate) use self::usage::maybe_build_local_admin_usage_response; + +pub(crate) async fn resolve_usage_user_group_scope( + state: &crate::handlers::admin::request::AdminAppState<'_>, + query: Option<&str>, + include_inactive: bool, + exclude_admin: bool, +) -> Result>, String>, crate::GatewayError> { + let group_id = crate::handlers::admin::shared::query_param_value(query, "user_group_id"); + let Some(group_id) = group_id else { + return Ok(Ok(None)); + }; + if crate::handlers::admin::shared::query_param_value(query, "user_id").is_some() { + return Ok(Err( + "user_id and user_group_id cannot be used together".to_string() + )); + } + if !state.has_user_data_reader() { + return Ok(Err("user group data is unavailable".to_string())); + } + + if group_id == UNGROUPED_USAGE_ID { + let ids = ungrouped_usage_users(state) + .await? + .into_iter() + .filter(|user| include_inactive || user.is_active) + .filter(|user| !exclude_admin || !user.role.eq_ignore_ascii_case("admin")) + .map(|user| user.id) + .collect(); + return Ok(Ok(Some(ids))); + } + match state + .resolve_usage_user_group_member_ids(&group_id, include_inactive, exclude_admin) + .await? + { + Some(user_ids) => Ok(Ok(Some(user_ids))), + None => Ok(Err("user_group_id does not exist".to_string())), + } +} + +/// Reserved statistics-only scope; never a permission group. +pub(crate) const UNGROUPED_USAGE_ID: &str = "__ungrouped__"; + +pub(crate) async fn ungrouped_usage_users( + state: &crate::handlers::admin::request::AdminAppState<'_>, +) -> Result, crate::GatewayError> { + use aether_data::repository::users::UserExportListQuery; + let mut users = Vec::new(); + let mut skip = 0; + loop { + let page = state + .list_export_users_page(&UserExportListQuery { + skip, + limit: 500, + ..Default::default() + }) + .await?; + let count = page.len(); + if count == 0 { + break; + } + let ids = page.into_iter().map(|user| user.id).collect::>(); + let grouped = state + .list_user_group_memberships_by_user_ids(&ids) + .await? + .into_iter() + .map(|membership| membership.user_id) + .collect::>(); + let ids = ids + .into_iter() + .filter(|id| !grouped.contains(id)) + .collect::>(); + users.extend( + state + .list_users_by_ids(&ids) + .await? + .into_iter() + .filter(|user| !user.is_deleted), + ); + skip += count; + if count < 500 { + break; + } + } + Ok(users) +} + +/// Current group provider policy, resolved to the provider-name dimension used by usage rollups. +/// None is unrestricted; Some(empty) deliberately matches no usage. +pub(crate) async fn usage_group_provider_names( + state: &crate::handlers::admin::request::AdminAppState<'_>, + group: &aether_data::repository::users::StoredUserGroup, +) -> Result>, crate::GatewayError> { + if matches!( + group.allowed_providers_mode.as_str(), + "unrestricted" | "inherit" + ) { + return Ok(None); + } + if group.allowed_providers_mode != "specific" { + return Ok(Some(Vec::new())); + } + let allowed = group.allowed_providers.as_deref().unwrap_or_default(); + let providers = state.list_provider_catalog_providers(false).await?; + let mut names = providers + .into_iter() + .filter(|provider| { + allowed.iter().any(|value| { + let value = value.trim(); + value.eq_ignore_ascii_case(&provider.id) + || value.eq_ignore_ascii_case(&provider.name) + || value.eq_ignore_ascii_case(&provider.provider_type) + }) + }) + .map(|provider| provider.name) + .collect::>(); + names.sort(); + names.dedup(); + Ok(Some(names)) +} + +pub(crate) async fn resolve_usage_group_provider_names( + state: &crate::handlers::admin::request::AdminAppState<'_>, + query: Option<&str>, +) -> Result>, crate::GatewayError> { + let Some(id) = crate::handlers::admin::shared::query_param_value(query, "user_group_id") else { + return Ok(None); + }; + if id == UNGROUPED_USAGE_ID { + return Ok(None); + } + let Some(group) = state.find_user_group_by_id(&id).await? else { + return Ok(Some(Vec::new())); + }; + usage_group_provider_names(state, &group).await +} diff --git a/apps/aether-gateway/src/handlers/admin/observability/monitoring/activity.rs b/apps/aether-gateway/src/handlers/admin/observability/monitoring/activity.rs index 14ed524dd..4e590277e 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/monitoring/activity.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/monitoring/activity.rs @@ -169,9 +169,11 @@ pub(super) async fn build_admin_monitoring_system_status_response( let today_usage = state .summarize_usage_audits(&UsageAuditSummaryQuery { + provider_names: None, created_from_unix_secs: today_start.timestamp().max(0) as u64, created_until_unix_secs: now_unix_secs.saturating_add(1), user_id: None, + user_ids: None, provider_name: None, model: None, }) diff --git a/apps/aether-gateway/src/handlers/admin/observability/stats/analytics_routes.rs b/apps/aether-gateway/src/handlers/admin/observability/stats/analytics_routes.rs index 45b59a005..a039afdcd 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/stats/analytics_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/stats/analytics_routes.rs @@ -1,3 +1,4 @@ +use super::super::resolve_usage_user_group_scope; use super::range::{build_comparison_range, parse_bounded_u32}; use super::resolve_admin_usage_time_range; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; @@ -109,6 +110,7 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response( }; let current_summary = state .summarize_usage_audits(&UsageAuditSummaryQuery { + provider_names: None, created_from_unix_secs: current_from_unix_secs, created_until_unix_secs: current_until_unix_secs, ..Default::default() @@ -116,6 +118,7 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response( .await?; let comparison_summary = state .summarize_usage_audits(&UsageAuditSummaryQuery { + provider_names: None, created_from_unix_secs: comparison_from_unix_secs, created_until_unix_secs: comparison_until_unix_secs, ..Default::default() @@ -294,6 +297,17 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response( } let filters = AdminStatsUsageFilter::from_query(request_context.query_string()); + let user_ids = match resolve_usage_user_group_scope( + state, + request_context.query_string(), + false, + false, + ) + .await? + { + Ok(value) => value, + Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))), + }; let query_granularity = match granularity { AdminStatsGranularity::Hour => UsageTimeSeriesGranularity::Hour, AdminStatsGranularity::Day @@ -306,11 +320,17 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response( }; let buckets = state .summarize_usage_time_series(&UsageTimeSeriesQuery { + provider_names: super::super::resolve_usage_group_provider_names( + state, + request_context.query_string(), + ) + .await?, created_from_unix_secs, created_until_unix_secs, granularity: query_granularity, tz_offset_minutes: time_range.tz_offset_minutes, user_id: filters.user_id, + user_ids, provider_name: filters.provider_name, model: filters.model, }) diff --git a/apps/aether-gateway/src/handlers/admin/observability/stats/cost_routes.rs b/apps/aether-gateway/src/handlers/admin/observability/stats/cost_routes.rs index f4c85089b..7867a1b2f 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/stats/cost_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/stats/cost_routes.rs @@ -72,11 +72,13 @@ pub(super) async fn maybe_build_local_admin_stats_cost_response( }; let buckets = state .summarize_usage_time_series(&UsageTimeSeriesQuery { + provider_names: None, created_from_unix_secs, created_until_unix_secs, granularity: UsageTimeSeriesGranularity::Day, tz_offset_minutes: time_range.tz_offset_minutes, user_id: None, + user_ids: None, provider_name: None, model: None, }) diff --git a/apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard.rs b/apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard.rs index 6ad30f30e..6b7b28089 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard.rs @@ -3,12 +3,13 @@ use crate::GatewayError; use aether_data_contracts::repository::usage::StoredRequestUsageAudit; pub(super) use aether_admin::observability::stats::{ - build_admin_stats_leaderboard_response, build_api_key_leaderboard_items, - build_api_key_leaderboard_items_from_summaries, build_model_leaderboard_items, - build_model_leaderboard_items_from_summaries, build_user_leaderboard_items, - build_user_leaderboard_items_from_summaries, compare_leaderboard_items, compute_dense_rank, - AdminStatsLeaderboardItem, AdminStatsLeaderboardMetric, AdminStatsLeaderboardNameMode, - AdminStatsSortOrder, AdminStatsUserMetadata, + build_admin_stats_leaderboard_response, build_admin_stats_user_group_leaderboard_response, + build_api_key_leaderboard_items, build_api_key_leaderboard_items_from_summaries, + build_model_leaderboard_items, build_model_leaderboard_items_from_summaries, + build_user_leaderboard_items, build_user_leaderboard_items_from_summaries, + compare_leaderboard_items, compute_dense_rank, AdminStatsLeaderboardItem, + AdminStatsLeaderboardMetric, AdminStatsLeaderboardNameMode, AdminStatsSortOrder, + AdminStatsUserMetadata, }; pub(super) async fn load_user_leaderboard_metadata( diff --git a/apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard_routes.rs b/apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard_routes.rs index 41e5a3b9a..49cd8a1a8 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/stats/leaderboard_routes.rs @@ -1,7 +1,9 @@ +use super::super::resolve_usage_user_group_scope; use super::leaderboard::{ - build_admin_stats_leaderboard_response, build_api_key_leaderboard_items_from_summaries, - build_model_leaderboard_items_from_summaries, build_user_leaderboard_items_from_summaries, - compare_leaderboard_items, load_user_leaderboard_metadata, AdminStatsLeaderboardNameMode, + build_admin_stats_leaderboard_response, build_admin_stats_user_group_leaderboard_response, + build_api_key_leaderboard_items_from_summaries, build_model_leaderboard_items_from_summaries, + build_user_leaderboard_items_from_summaries, compare_leaderboard_items, + load_user_leaderboard_metadata, AdminStatsLeaderboardItem, AdminStatsLeaderboardNameMode, }; use super::range::{parse_bounded_u32, parse_nonnegative_usize}; use super::resolve_admin_usage_time_range; @@ -14,6 +16,7 @@ use aether_admin::observability::stats::{ }; use aether_data_contracts::repository::usage::{UsageLeaderboardGroupBy, UsageLeaderboardQuery}; use axum::{body::Body, http, response::Response}; +use std::collections::{BTreeMap, BTreeSet}; pub(super) async fn maybe_build_local_admin_stats_leaderboard_response( state: &AdminAppState<'_>, @@ -75,10 +78,12 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response( }; let summaries = state .summarize_usage_leaderboard(&UsageLeaderboardQuery { + provider_names: None, created_from_unix_secs, created_until_unix_secs, group_by: UsageLeaderboardGroupBy::Model, user_id: filters.user_id, + user_ids: None, provider_name: filters.provider_name, model: filters.model, }) @@ -152,10 +157,12 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response( }; let summaries = state .summarize_usage_leaderboard(&UsageLeaderboardQuery { + provider_names: None, created_from_unix_secs, created_until_unix_secs, group_by: UsageLeaderboardGroupBy::ApiKey, user_id: filters.user_id, + user_ids: None, provider_name: filters.provider_name, model: filters.model, }) @@ -206,6 +213,181 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response( ))); } + if request_context + .decision() + .and_then(|decision| decision.route_kind.as_deref()) + == Some("leaderboard_user_groups") + && request_context.method() == http::Method::GET + && matches!( + request_context.path(), + "/api/admin/stats/leaderboard/user-groups" + | "/api/admin/stats/leaderboard/user-groups/" + ) + { + let time_range = match resolve_admin_usage_time_range(query) { + Ok(value) => value, + Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))), + }; + let metric = match AdminStatsLeaderboardMetric::parse(query) { + Ok(value) => value, + Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))), + }; + let order = match AdminStatsSortOrder::parse(query) { + Ok(value) => value, + Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))), + }; + let limit = match query_param_value(query, "limit") + .map(|value| parse_bounded_u32("limit", &value, 1, 100)) + .transpose() + { + Ok(Some(value)) => value as usize, + Ok(None) => 10, + Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))), + }; + let offset = match query_param_value(query, "offset") + .map(|value| parse_nonnegative_usize("offset", &value)) + .transpose() + { + Ok(Some(value)) => value, + Ok(None) => 0, + Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))), + }; + let empty_counts = BTreeMap::new(); + if !state.has_usage_data_reader() || !state.has_user_data_reader() { + return Ok(Some(build_admin_stats_user_group_leaderboard_response( + metric, + Some(&time_range), + &[], + &empty_counts, + &empty_counts, + offset, + limit, + ))); + } + let include_inactive = query_param_bool(query, "include_inactive", false); + let exclude_admin = query_param_bool(query, "exclude_admin", false); + let filters = AdminStatsUsageFilter::from_query(query); + if filters.user_id.is_some() { + return Ok(Some(admin_stats_bad_request_response( + "user_id is not supported for the user group leaderboard".to_string(), + ))); + } + let Some((created_from_unix_secs, created_until_unix_secs)) = time_range.to_unix_bounds() + else { + return Ok(Some(build_admin_stats_user_group_leaderboard_response( + metric, + Some(&time_range), + &[], + &empty_counts, + &empty_counts, + offset, + limit, + ))); + }; + + let mut leaderboard = Vec::new(); + let mut member_counts = BTreeMap::new(); + let mut active_member_counts = BTreeMap::new(); + for group in state.list_user_groups().await? { + let members = state.list_user_group_members(&group.id).await?; + let member_count = members.iter().filter(|member| !member.is_deleted).count(); + let active_member_count = members + .iter() + .filter(|member| !member.is_deleted && member.is_active) + .count(); + let user_ids = members + .iter() + .filter(|member| !member.is_deleted) + .filter(|member| include_inactive || member.is_active) + .filter(|member| !exclude_admin || !member.role.eq_ignore_ascii_case("admin")) + .map(|member| member.user_id.clone()) + .collect::>(); + let summaries = state + .summarize_usage_leaderboard(&UsageLeaderboardQuery { + created_from_unix_secs, + created_until_unix_secs, + group_by: UsageLeaderboardGroupBy::User, + user_id: None, + user_ids: Some(user_ids), + provider_names: super::super::usage_group_provider_names(state, &group).await?, + provider_name: filters.provider_name.clone(), + model: filters.model.clone(), + }) + .await?; + let user_ids = summaries + .iter() + .map(|row| row.group_key.clone()) + .collect::>(); + let metadata = load_user_leaderboard_metadata(state, &user_ids).await?; + let users = build_user_leaderboard_items_from_summaries( + &summaries, + &metadata, + state.has_auth_user_data_reader(), + state.has_user_data_reader(), + include_inactive, + exclude_admin, + ); + let mut item = AdminStatsLeaderboardItem { + id: group.id.clone(), + name: group.name, + requests: 0, + tokens: 0, + cost: 0.0, + }; + for user in users { + item.requests = item.requests.saturating_add(user.requests); + item.tokens = item.tokens.saturating_add(user.tokens); + item.cost += user.cost; + } + member_counts.insert(group.id.clone(), member_count); + active_member_counts.insert(group.id, active_member_count); + leaderboard.push(item); + } + let ungrouped = super::super::ungrouped_usage_users(state).await?; + let id = super::super::UNGROUPED_USAGE_ID.to_string(); + member_counts.insert(id.clone(), ungrouped.len()); + active_member_counts.insert( + id.clone(), + ungrouped.iter().filter(|user| user.is_active).count(), + ); + let user_ids = ungrouped + .into_iter() + .filter(|user| include_inactive || user.is_active) + .filter(|user| !exclude_admin || !user.role.eq_ignore_ascii_case("admin")) + .map(|user| user.id) + .collect(); + let rows = state + .summarize_usage_leaderboard(&UsageLeaderboardQuery { + created_from_unix_secs, + created_until_unix_secs, + group_by: UsageLeaderboardGroupBy::User, + user_id: None, + user_ids: Some(user_ids), + provider_names: None, + provider_name: filters.provider_name.clone(), + model: filters.model.clone(), + }) + .await?; + leaderboard.push(AdminStatsLeaderboardItem { + id, + name: "Ungrouped".to_string(), + requests: rows.iter().map(|row| row.request_count).sum(), + tokens: rows.iter().map(|row| row.total_tokens).sum(), + cost: rows.iter().map(|row| row.total_cost_usd).sum(), + }); + leaderboard.sort_by(|left, right| compare_leaderboard_items(metric, order, left, right)); + + return Ok(Some(build_admin_stats_user_group_leaderboard_response( + metric, + Some(&time_range), + &leaderboard, + &member_counts, + &active_member_counts, + offset, + limit, + ))); + } + if request_context .decision() .and_then(|decision| decision.route_kind.as_deref()) @@ -253,6 +435,13 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response( let include_inactive = query_param_bool(query, "include_inactive", false); let exclude_admin = query_param_bool(query, "exclude_admin", false); let filters = AdminStatsUsageFilter::from_query(query); + let scoped_user_ids = + match resolve_usage_user_group_scope(state, query, include_inactive, exclude_admin) + .await? + { + Ok(value) => value, + Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))), + }; let Some((created_from_unix_secs, created_until_unix_secs)) = time_range.to_unix_bounds() else { return Ok(Some(admin_stats_leaderboard_empty_response( @@ -262,10 +451,13 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response( }; let summaries = state .summarize_usage_leaderboard(&UsageLeaderboardQuery { + provider_names: super::super::resolve_usage_group_provider_names(state, query) + .await?, created_from_unix_secs, created_until_unix_secs, group_by: UsageLeaderboardGroupBy::User, user_id: filters.user_id, + user_ids: scoped_user_ids, provider_name: filters.provider_name, model: filters.model, }) diff --git a/apps/aether-gateway/src/handlers/admin/observability/usage/summary_routes.rs b/apps/aether-gateway/src/handlers/admin/observability/usage/summary_routes.rs index a197e8cce..8a4c1f48a 100644 --- a/apps/aether-gateway/src/handlers/admin/observability/usage/summary_routes.rs +++ b/apps/aether-gateway/src/handlers/admin/observability/usage/summary_routes.rs @@ -1,3 +1,4 @@ +use super::super::resolve_usage_user_group_scope; use super::super::stats::resolve_admin_usage_time_range; use super::analytics::admin_usage_api_key_names; use super::analytics::admin_usage_provider_key_names; @@ -121,21 +122,30 @@ fn apply_admin_usage_status_filter(query: &mut UsageAuditListQuery, status: Opti "pending" | "streaming" | "completed" | "cancelled" => { query.statuses = Some(vec![status]); } - "has_fallback" | "has_retry" => {} + "has_fallback" | "has_retry" | "has_skipped_candidate" => {} _ => {} } } -#[derive(Clone, Copy, Debug, Default)] +#[derive(Clone, Debug, Default)] struct AdminUsageAttemptFlags { has_fallback: bool, has_retry: bool, + /// 是否存在"被调度跳过"的候选(调度阶段判定本次不可用,从未向上游发起请求)。 + /// + /// 这是与 has_fallback 正交的信号:has_fallback 表示"更靠前的候选真的失败并被换掉", + /// 而本字段表示"更靠前的候选压根没被发出去"。两者在日志列表里观感都是"换了提供商", + /// 但用户拿不到 has_fallback 小图标时容易误判为调度错误,故单独暴露。 + has_skipped_candidate: bool, + /// 跳过原因(去重、保持出现顺序),用于前端 tooltip 直接说明"为什么没用它"。 + skipped_candidate_reasons: Vec, } fn admin_usage_attempt_status_filter(status: Option<&str>) -> Option<&'static str> { match status?.trim().to_ascii_lowercase().as_str() { "has_fallback" => Some("has_fallback"), "has_retry" => Some("has_retry"), + "has_skipped_candidate" => Some("has_skipped_candidate"), _ => None, } } @@ -190,25 +200,52 @@ fn admin_usage_attempt_flags_from_candidates( }) }); let has_retry = candidates.iter().any(admin_usage_candidate_was_retried); + let skipped_candidate_reasons = admin_usage_skipped_candidate_reasons(candidates); AdminUsageAttemptFlags { has_fallback, has_retry, + has_skipped_candidate: !skipped_candidate_reasons.is_empty(), + skipped_candidate_reasons, } } +/// 收集被跳过候选的原因,去重并保持候选顺序(决定性的在前,便于阅读)。 +fn admin_usage_skipped_candidate_reasons(candidates: &[StoredRequestCandidate]) -> Vec { + let mut reasons = Vec::new(); + for candidate in candidates + .iter() + .filter(|candidate| candidate.status == RequestCandidateStatus::Skipped) + { + let Some(reason) = candidate + .skip_reason + .as_deref() + .map(str::trim) + .filter(|reason| !reason.is_empty()) + else { + continue; + }; + if !reasons.iter().any(|existing| existing == reason) { + reasons.push(reason.to_string()); + } + } + reasons +} + fn admin_usage_attempt_flags_for_item( item: &StoredRequestUsageAudit, flags_by_usage_id: &BTreeMap, request_candidate_reader_available: bool, ) -> AdminUsageAttemptFlags { - flags_by_usage_id.get(&item.id).copied().unwrap_or_else(|| { + flags_by_usage_id.get(&item.id).cloned().unwrap_or_else(|| { if request_candidate_reader_available { AdminUsageAttemptFlags::default() } else { AdminUsageAttemptFlags { has_fallback: admin_usage_has_fallback(item), has_retry: false, + has_skipped_candidate: false, + skipped_candidate_reasons: Vec::new(), } } }) @@ -477,6 +514,8 @@ fn admin_usage_matches_attempt_status( match status { "has_fallback" => flags.has_fallback, "has_retry" => flags.has_retry, + // 与 has_fallback 区分:这里是"更靠前的候选被调度跳过、根本没发出去" + "has_skipped_candidate" => flags.has_skipped_candidate, _ => true, } } @@ -548,6 +587,9 @@ fn build_admin_usage_records_response_with_attempt_flags( ); record["has_fallback"] = json!(flags.has_fallback); record["has_retry"] = json!(flags.has_retry); + // 被跳过的候选:前端据此提示"这次没用某个提供商,是因为它在调度阶段就被排除了"。 + record["has_skipped_candidate"] = json!(flags.has_skipped_candidate); + record["skipped_candidate_reasons"] = json!(flags.skipped_candidate_reasons); record }) .collect(); @@ -799,11 +841,18 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response( &Default::default(), ))); }; + let user_ids = match resolve_usage_user_group_scope(state, query, false, false).await? { + Ok(value) => value, + Err(detail) => return Ok(Some(admin_usage_bad_request_response(detail))), + }; let summary = state .summarize_usage_audits(&UsageAuditSummaryQuery { + provider_names: super::super::resolve_usage_group_provider_names(state, query) + .await?, created_from_unix_secs, created_until_unix_secs, user_id: query_param_value(query, "user_id"), + user_ids, provider_name: query_param_value(query, "provider"), model: query_param_value(query, "model"), }) @@ -1202,9 +1251,11 @@ mod tests { use aether_data_contracts::repository::candidates::{ RequestCandidateStatus, StoredRequestCandidate, }; + use aether_data_contracts::repository::usage::StoredRequestUsageAudit; use serde_json::json; use super::{ + admin_usage_attempt_flags_from_candidates, admin_usage_skipped_candidate_reasons, admin_usage_terminal_candidate_state_override, build_admin_usage_keyword_search_query, build_admin_usage_records_query, latest_admin_usage_image_progress, AdminUsageSearchContext, @@ -1246,6 +1297,144 @@ mod tests { .expect("candidate should build") } + /// 构造一条"被调度跳过"的候选(从未向上游发起请求)。 + fn skipped_candidate(candidate_index: i32, reason: &str) -> StoredRequestCandidate { + let mut candidate = sample_candidate( + candidate_index, + RequestCandidateStatus::Skipped, + None, + None, + None, + ); + candidate.skip_reason = Some(reason.to_string()); + // 跳过候选没有开始时间,is_attempted 因此为 false + candidate.started_at_unix_ms = None; + candidate + } + + #[test] + fn skipped_candidate_reasons_are_deduplicated_in_candidate_order() { + let reasons = admin_usage_skipped_candidate_reasons(&[ + skipped_candidate(0, "key_rpm_exhausted"), + skipped_candidate(1, "provider_inactive"), + skipped_candidate(2, "key_rpm_exhausted"), + ]); + + assert_eq!( + reasons, + vec![ + "key_rpm_exhausted".to_string(), + "provider_inactive".to_string() + ] + ); + } + + #[test] + fn skipped_candidate_reasons_ignore_attempted_candidates() { + // 真正发起过请求的失败候选不属于"被跳过",避免与 has_fallback 语义混淆 + let failed = sample_candidate( + 0, + RequestCandidateStatus::Failed, + Some(503), + Some(1_000), + Some("upstream exploded"), + ); + assert!(admin_usage_skipped_candidate_reasons(&[failed]).is_empty()); + } + + #[test] + fn attempt_flags_report_skipped_candidates_without_fallback() { + let candidates = vec![ + skipped_candidate(0, "key_rpm_exhausted"), + sample_candidate( + 1, + RequestCandidateStatus::Success, + Some(200), + Some(900), + None, + ), + ]; + + let flags = admin_usage_attempt_flags_from_candidates(&sample_usage_audit(), &candidates); + + // 这正是用户遇到的场景:换了提供商,但没有任何候选失败过 + assert!(flags.has_skipped_candidate); + assert!(!flags.has_fallback); + assert_eq!( + flags.skipped_candidate_reasons, + vec!["key_rpm_exhausted".to_string()] + ); + } + + #[test] + fn attempt_flags_keep_fallback_and_skipped_candidate_independent() { + let candidates = vec![ + skipped_candidate(0, "provider_inactive"), + sample_candidate( + 1, + RequestCandidateStatus::Failed, + Some(503), + Some(500), + None, + ), + sample_candidate( + 2, + RequestCandidateStatus::Success, + Some(200), + Some(700), + None, + ), + ]; + + let flags = admin_usage_attempt_flags_from_candidates(&sample_usage_audit(), &candidates); + + assert!(flags.has_skipped_candidate); + assert!(flags.has_fallback); + } + + /// 最小可用的用量审计行,仅用于驱动 flags 计算(其中候选 id 为空即可)。 + fn sample_usage_audit() -> StoredRequestUsageAudit { + StoredRequestUsageAudit::new( + "usage-1".to_string(), + "req-1".to_string(), + Some("user-1".to_string()), + Some("api-key-1".to_string()), + Some("alice".to_string()), + Some("default".to_string()), + "OpenAI".to_string(), + "gpt-4.1".to_string(), + None, + None, + None, + None, + None, + Some("openai:chat".to_string()), + Some("openai".to_string()), + Some("chat".to_string()), + Some("openai:chat".to_string()), + Some("openai".to_string()), + Some("chat".to_string()), + false, + false, + 10, + 20, + 30, + 0.0, + 0.0, + Some(200), + None, + None, + None, + None, + "completed".to_string(), + "settled".to_string(), + 1_000, + 1_001, + None, + ) + .expect("usage should build") + } + #[test] fn admin_usage_active_override_uses_current_terminal_candidate_latency() { let candidate = sample_candidate( diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs index 42d9c0c15..f2a355f3d 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs @@ -74,7 +74,8 @@ fn validate_batch_access_token_import( ) -> Result<(), String> { if !provider_type_supports_access_token_import(provider_type) { return Err( - "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider".to_string(), + "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider" + .to_string(), ); } if provider_type.eq_ignore_ascii_case("claude_code") { diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs index 60cff9220..0ad07cd61 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs @@ -214,7 +214,12 @@ fn extract_admin_provider_oauth_batch_import_entry( } else { let sso_from_cookie = grok_cookie_session_token(provider_type, raw_token); let token_input = sso_from_cookie.as_deref().unwrap_or(raw_token); - let (refresh_token, access_token) = import_tokens_from_raw_token(token_input); + let (refresh_token, access_token) = + if provider_type.trim().eq_ignore_ascii_case("xai") { + (None, Some(token_input.to_string())) + } else { + import_tokens_from_raw_token(token_input) + }; let (refresh_token, access_token) = normalize_provider_import_tokens( provider_type, refresh_token.as_deref(), @@ -262,6 +267,7 @@ fn extract_admin_provider_oauth_batch_import_entry( let object = normalized_claude_object.as_ref().unwrap_or(object); let is_grok = provider_type.trim().eq_ignore_ascii_case("grok"); let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf"); + let is_xai = provider_type.trim().eq_ignore_ascii_case("xai"); let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex") && aether_provider_transport::is_codex_agent_identity_auth_config_value(item); if is_codex_agent_identity { @@ -336,14 +342,6 @@ fn extract_admin_provider_oauth_batch_import_entry( } else { None }; - let (refresh_token, access_token) = normalize_provider_import_tokens( - provider_type, - refresh_token.as_deref(), - access_token - .as_deref() - .or(session_token.as_deref()) - .or(header_bearer_token.as_deref()), - ); let windsurf_api_key = is_windsurf .then(|| { coerce_admin_provider_oauth_import_str( @@ -351,6 +349,22 @@ fn extract_admin_provider_oauth_batch_import_entry( ) }) .flatten(); + let xai_api_key = is_xai + .then(|| { + coerce_admin_provider_oauth_import_str( + object.get("api_key").or_else(|| object.get("apiKey")), + ) + }) + .flatten(); + let (refresh_token, access_token) = normalize_provider_import_tokens( + provider_type, + refresh_token.as_deref(), + access_token + .as_deref() + .or(session_token.as_deref()) + .or(header_bearer_token.as_deref()) + .or(xai_api_key.as_deref()), + ); let windsurf_token = is_windsurf .then(|| { coerce_admin_provider_oauth_import_str( @@ -1577,4 +1591,23 @@ mod tests { assert!(entries[1].access_token.is_none()); assert!(entries[1].raw_credentials.is_none()); } + + #[test] + fn parses_xai_api_key_json_and_raw_lines_as_access_token() { + let entries = parse_admin_provider_oauth_batch_import_entries( + "xai", + r#"{"api_key":"xai-api-key","email":"a@x.ai"} +{"refresh_token":"xai-refresh"} +xai-raw-api-key"#, + ); + + assert_eq!(entries.len(), 3); + assert!(entries[0].refresh_token.is_none()); + assert_eq!(entries[0].access_token.as_deref(), Some("xai-api-key")); + assert_eq!(entries[0].email.as_deref(), Some("a@x.ai")); + assert_eq!(entries[1].refresh_token.as_deref(), Some("xai-refresh")); + assert!(entries[1].access_token.is_none()); + assert!(entries[2].refresh_token.is_none()); + assert_eq!(entries[2].access_token.as_deref(), Some("xai-raw-api-key")); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/authorize.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/authorize.rs index b4ce3c5e5..fc92acb64 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/authorize.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/authorize.rs @@ -186,10 +186,10 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( )); }; let provider_type = provider.provider_type.trim().to_ascii_lowercase(); - if provider_type != "kiro" && provider_type != "windsurf" { + if provider_type != "kiro" && provider_type != "windsurf" && provider_type != "xai" { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - "设备授权仅支持 Kiro / Windsurf provider", + "设备授权仅支持 Kiro / Windsurf / xAI provider", )); } let Some(principal) = request_context @@ -219,6 +219,19 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( ) .await; + if provider_type == "xai" { + return super::xai::handle_admin_provider_oauth_xai_device_authorize( + state, + &provider_id, + &provider, + principal, + runtime_endpoint.as_ref(), + request_proxy, + payload.proxy_node_id.as_deref(), + ) + .await; + } + if provider_type == "windsurf" { let session_id = generate_provider_oauth_nonce(); let login_option = payload diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/mod.rs index 295a93336..172f99cde 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/mod.rs @@ -2,6 +2,7 @@ mod authorize; mod lease; mod poll; mod session; +mod xai; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::GatewayError; diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/poll.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/poll.rs index 8074a0024..649623f51 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/poll.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/poll.rs @@ -479,6 +479,18 @@ pub(super) async fn handle_admin_provider_oauth_device_poll( ) .await; + if provider_type == "xai" { + return super::xai::handle_admin_provider_oauth_xai_device_poll( + state, + &provider, + &endpoints, + request_proxy, + session_id, + session, + ) + .await; + } + if provider_type == "windsurf" { return handle_admin_provider_oauth_windsurf_browser_device_poll( state, diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/xai.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/xai.rs new file mode 100644 index 000000000..bf42d5db6 --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/xai.rs @@ -0,0 +1,368 @@ +use super::session::attach_admin_provider_oauth_device_poll_terminal_response; +use crate::control::GatewayAdminPrincipalContext; +use crate::handlers::admin::provider::oauth::dispatch::helpers::admin_provider_oauth_key_name_from_auth_config; +use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response; +use crate::handlers::admin::provider::oauth::provisioning::{ + provider_oauth_active_api_formats, provider_oauth_key_proxy_value, +}; +use crate::handlers::admin::provider::oauth::runtime::spawn_provider_oauth_account_state_refresh_after_update; +use crate::handlers::admin::provider::oauth::state::{ + current_unix_secs, generate_provider_oauth_nonce, +}; +use crate::handlers::admin::request::AdminAppState; +use crate::GatewayError; +use aether_contracts::ProxySnapshot; +use aether_data::repository::provider_oauth::{ + StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS, +}; +use aether_data_contracts::repository::provider_catalog::{ + StoredProviderCatalogEndpoint, StoredProviderCatalogProvider, +}; +use aether_oauth::core::OAuthError; +use aether_oauth::provider::providers::{ + XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_URL, + XAI_TOKEN_URL, +}; +use aether_oauth::provider::ProviderOAuthTransportContext; +use axum::{ + body::Body, + http, + response::{IntoResponse, Response}, + Json, +}; +use serde_json::{json, Value}; + +pub(super) async fn handle_admin_provider_oauth_xai_device_authorize( + state: &AdminAppState<'_>, + provider_id: &str, + provider: &StoredProviderCatalogProvider, + principal: &GatewayAdminPrincipalContext, + runtime_endpoint: Option<&StoredProviderCatalogEndpoint>, + request_proxy: Option, + proxy_node_id: Option<&str>, +) -> Result, GatewayError> { + let device_url = state.provider_oauth_token_url("xai_device", XAI_DEVICE_CODE_URL); + let token_url = state.provider_oauth_token_url("xai", XAI_TOKEN_URL); + let adapter = + XaiProviderOAuthAdapter::default().with_endpoint_overrides(&device_url, &token_url); + let ctx = ProviderOAuthTransportContext { + provider_id: provider_id.to_string(), + provider_type: provider.provider_type.clone(), + endpoint_id: runtime_endpoint.map(|endpoint| endpoint.id.clone()), + key_id: None, + auth_type: Some("oauth".to_string()), + decrypted_api_key: None, + decrypted_auth_config: None, + provider_config: provider.config.clone(), + endpoint_config: runtime_endpoint.and_then(|endpoint| endpoint.config.clone()), + key_config: None, + network: aether_oauth::network::OAuthNetworkContext::provider_operation( + request_proxy.clone(), + ), + }; + let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state); + let authorization = match adapter.start_device_flow(&executor, &ctx).await { + Ok(authorization) => authorization, + Err(error) => { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + sanitize_xai_oauth_error(&error), + )); + } + }; + + let now_unix_secs = current_unix_secs(); + let session_id = generate_provider_oauth_nonce(); + let session = StoredAdminProviderOAuthDeviceSession { + session_id: session_id.clone(), + provider_id: provider_id.to_string(), + initiated_by_user_id: principal.user_id.clone(), + initiated_by_session_id: principal.session_id.clone(), + initiated_by_management_token_id: principal.management_token_id.clone(), + region: String::new(), + client_id: XAI_CLIENT_ID.to_string(), + client_secret: String::new(), + device_code: authorization.device_code.clone(), + auth_type: Some("device".to_string()), + social_provider: None, + code_verifier: None, + redirect_uri: Some(token_url), + machine_id: None, + interval: authorization.interval, + expires_at_unix_secs: now_unix_secs.saturating_add(authorization.expires_in), + status: "pending".to_string(), + proxy_node_id: proxy_node_id + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned), + created_at_unix_ms: now_unix_secs, + key_id: None, + email: None, + replaced: false, + error_msg: None, + }; + if let Err(response) = state + .save_provider_oauth_device_session( + &session_id, + &session, + authorization + .expires_in + .saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS), + ) + .await + { + return Ok(response); + } + + Ok(Json(json!({ + "session_id": session_id, + "user_code": authorization.user_code, + "verification_uri": authorization.verification_uri, + "verification_uri_complete": authorization.verification_uri_complete, + "expires_in": authorization.expires_in, + "interval": authorization.interval, + "auth_type": "device", + })) + .into_response()) +} + +pub(super) async fn handle_admin_provider_oauth_xai_device_poll( + state: &AdminAppState<'_>, + provider: &StoredProviderCatalogProvider, + endpoints: &[StoredProviderCatalogEndpoint], + request_proxy: Option, + session_id: &str, + mut session: StoredAdminProviderOAuthDeviceSession, +) -> Result, GatewayError> { + let token_url = session + .redirect_uri + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .unwrap_or_else(|| state.provider_oauth_token_url("xai", XAI_TOKEN_URL)); + let adapter = + XaiProviderOAuthAdapter::default().with_endpoint_overrides(XAI_DEVICE_CODE_URL, token_url); + let ctx = ProviderOAuthTransportContext { + provider_id: provider.id.clone(), + provider_type: provider.provider_type.clone(), + endpoint_id: None, + key_id: None, + auth_type: Some("oauth".to_string()), + decrypted_api_key: None, + decrypted_auth_config: None, + provider_config: provider.config.clone(), + endpoint_config: None, + key_config: None, + network: aether_oauth::network::OAuthNetworkContext::provider_operation( + request_proxy.clone(), + ), + }; + let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state); + let outcome = match adapter + .poll_device_token(&executor, &ctx, &session.device_code) + .await + { + Ok(outcome) => outcome, + Err(error) => { + return Ok(xai_device_poll_terminal_from_error( + state, + session_id, + &mut session, + &error, + ) + .await); + } + }; + + match outcome { + XaiDevicePollOutcome::Pending => { + Ok(Json(json!({"status": "pending", "replaced": false})).into_response()) + } + XaiDevicePollOutcome::SlowDown => { + Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response()) + } + XaiDevicePollOutcome::Authorized(result) => { + persist_xai_device_authorization( + state, + provider, + endpoints, + request_proxy, + session_id, + session, + *result, + ) + .await + } + } +} + +async fn persist_xai_device_authorization( + state: &AdminAppState<'_>, + provider: &StoredProviderCatalogProvider, + endpoints: &[StoredProviderCatalogEndpoint], + request_proxy: Option, + session_id: &str, + mut session: StoredAdminProviderOAuthDeviceSession, + result: aether_oauth::provider::ProviderOAuthTokenSet, +) -> Result, GatewayError> { + let access_token = result.token_set.access_token.trim().to_string(); + if access_token.is_empty() { + return Ok(Json(json!({ + "status": "error", + "error": "xAI token 响应缺少 access_token", + "replaced": false, + })) + .into_response()); + } + let mut auth_config = result.auth_config.as_object().cloned().unwrap_or_default(); + auth_config.insert("provider_type".to_string(), json!("xai")); + auth_config.insert("auth_method".to_string(), json!("oauth")); + auth_config.insert("using_api".to_string(), json!(false)); + + let duplicate = match state + .find_duplicate_provider_oauth_key(&provider.id, &auth_config, None) + .await + { + Ok(duplicate) => duplicate, + Err(detail) => { + return Ok(Json(json!({ + "status": "error", + "error": detail, + "replaced": false, + })) + .into_response()); + } + }; + + let api_formats = provider_oauth_active_api_formats(endpoints); + let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref()); + let expires_at = result.token_set.expires_at_unix_secs; + let email = auth_config + .get("email") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let mut replaced = false; + let persisted_key = if let Some(existing_key) = duplicate { + replaced = true; + match state + .update_existing_provider_oauth_catalog_key( + &existing_key, + &provider.provider_type, + &access_token, + &auth_config, + &api_formats, + key_proxy.clone(), + expires_at, + ) + .await? + { + Some(key) => key, + None => { + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "provider oauth write unavailable", + )); + } + } + } else { + let key_name = admin_provider_oauth_key_name_from_auth_config( + &provider.provider_type, + &auth_config, + None, + ); + match state + .create_provider_oauth_catalog_key( + &provider.id, + &provider.provider_type, + &key_name, + &access_token, + &auth_config, + &api_formats, + key_proxy, + expires_at, + ) + .await? + { + Some(key) => key, + None => { + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "provider oauth write unavailable", + )); + } + } + }; + + spawn_provider_oauth_account_state_refresh_after_update( + state.cloned_app(), + provider.clone(), + persisted_key.id.clone(), + request_proxy.clone(), + ); + + session.status = "authorized".to_string(); + session.key_id = Some(persisted_key.id.clone()); + session.email = email.clone(); + session.replaced = replaced; + session.error_msg = None; + let _ = state + .save_provider_oauth_device_session(session_id, &session, 60) + .await; + + Ok(attach_admin_provider_oauth_device_poll_terminal_response( + session_id, + "authorized", + Json(json!({ + "status": "authorized", + "key_id": persisted_key.id, + "email": email, + "replaced": replaced, + })) + .into_response(), + )) +} + +async fn xai_device_poll_terminal_from_error( + state: &AdminAppState<'_>, + session_id: &str, + session: &mut StoredAdminProviderOAuthDeviceSession, + error: &OAuthError, +) -> Response { + let (status, message) = match error { + OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("expired") => { + ("expired", "设备码已过期".to_string()) + } + OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("denied") => { + ("error", "用户拒绝授权".to_string()) + } + _ => ("error", sanitize_xai_oauth_error(error)), + }; + session.status = status.to_string(); + session.error_msg = Some(message.clone()); + let _ = state + .save_provider_oauth_device_session(session_id, session, 30) + .await; + attach_admin_provider_oauth_device_poll_terminal_response( + session_id, + status, + Json(json!({ + "status": status, + "error": message, + "replaced": false, + })) + .into_response(), + ) +} + +fn sanitize_xai_oauth_error(error: &OAuthError) -> String { + match error { + OAuthError::InvalidRequest(_) => "xAI 设备授权失败: 请求参数无效".to_string(), + OAuthError::HttpStatus { status_code, .. } => { + format!("xAI 设备授权失败: HTTP {status_code}") + } + _ => "xAI 设备授权失败".to_string(), + } +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs index 5ef0ea8b9..c7707e786 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs @@ -715,7 +715,7 @@ async fn resolve_admin_provider_oauth_single_import_tokens( if !provider_type_supports_access_token_import(provider_type) { return Err(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider", + "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider", )); } @@ -867,7 +867,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( flatten_claude_code_credentials_payload(&mut raw_payload); } let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken"); - let access_token_input = import_payload_string_any( + let mut access_token_input = import_payload_string_any( &raw_payload, &[ "access_token", @@ -879,6 +879,9 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( ], ) .or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload)); + if provider_type == "xai" && access_token_input.is_none() { + access_token_input = import_payload_string(&raw_payload, "api_key", "apiKey"); + } let imported_expires_at = import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]); let (refresh_token_input, access_token_input) = normalize_provider_import_tokens( @@ -901,7 +904,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - "Refresh Token、Access Token 或 sso_token 不能为空", + if provider_type == "xai" { + "Refresh Token、Access Token 或 api_key 不能为空" + } else { + "Refresh Token、Access Token 或 sso_token 不能为空" + }, )); } if !is_fixed_provider_type_for_provider_oauth(&provider_type) { diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/start.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/start.rs index d909b0320..184ec1c1b 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/start.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/start.rs @@ -70,6 +70,12 @@ pub(super) async fn handle_admin_provider_oauth_start_key( "Windsurf 请使用浏览器登录或导入凭据。", )); } + if provider_type == "xai" { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "xAI 请使用设备授权或导入凭据。", + )); + } let Some(template) = admin_provider_oauth_template(&provider_type) else { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, @@ -167,6 +173,12 @@ pub(super) async fn handle_admin_provider_oauth_start_provider( "Windsurf 请使用浏览器登录或导入凭据。", )); } + if provider_type == "xai" { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "xAI 请使用设备授权或导入凭据。", + )); + } let Some(template) = admin_provider_oauth_template(&provider_type) else { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs index f446d4104..b6566f052 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs @@ -121,6 +121,9 @@ pub(super) fn normalize_provider_import_tokens( if provider_type == "grok" { return (None, access_token.or(refresh_token)); } + if provider_type == "xai" { + return (refresh_token, access_token); + } if provider_type == "claude_code" { if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) { return (None, refresh_token); @@ -237,7 +240,7 @@ pub(super) fn provider_oauth_import_authorization_bearer_token_from_object( pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool { matches!( provider_type.trim().to_ascii_lowercase().as_str(), - "claude_code" | "codex" | "chatgpt_web" | "grok" + "claude_code" | "codex" | "chatgpt_web" | "grok" | "xai" ) } @@ -331,6 +334,15 @@ pub(super) fn build_provider_access_token_import_auth_config( auth_config.insert("sso_token".to_string(), json!(access_token)); auth_config.insert("auth_method".to_string(), json!("sso_token")); } + if provider_type.trim().eq_ignore_ascii_case("xai") { + if refresh_token.is_some() { + auth_config.insert("auth_method".to_string(), json!("oauth")); + auth_config.insert("using_api".to_string(), json!(false)); + } else { + auth_config.insert("auth_method".to_string(), json!("api_key")); + auth_config.insert("using_api".to_string(), json!(true)); + } + } auth_config.insert( "access_token_import_temporary".to_string(), @@ -532,6 +544,41 @@ mod tests { ); } + #[test] + fn normalize_xai_import_keeps_refresh_token_separate_from_api_key() { + let (refresh_token, access_token) = + normalize_provider_import_tokens("xai", Some("xai-refresh-token"), None); + assert_eq!(refresh_token.as_deref(), Some("xai-refresh-token")); + assert!(access_token.is_none()); + + let (refresh_token, access_token) = + normalize_provider_import_tokens("xai", None, Some("xai-api-key")); + assert!(refresh_token.is_none()); + assert_eq!(access_token.as_deref(), Some("xai-api-key")); + } + + #[test] + fn builds_xai_auth_config_from_api_key_and_oauth_tokens() { + let (api_key_config, _) = + build_provider_access_token_import_auth_config("xai", "xai-api-key", None, None, None); + assert_eq!(api_key_config.get("auth_method"), Some(&json!("api_key"))); + assert_eq!(api_key_config.get("using_api"), Some(&json!(true))); + + let (oauth_config, _) = build_provider_access_token_import_auth_config( + "xai", + "xai-access-token", + Some("xai-refresh-token"), + None, + None, + ); + assert_eq!(oauth_config.get("auth_method"), Some(&json!("oauth"))); + assert_eq!(oauth_config.get("using_api"), Some(&json!(false))); + assert_eq!( + oauth_config.get("refresh_token"), + Some(&json!("xai-refresh-token")) + ); + } + #[test] fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() { let mut payload = json!({ diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/claude_code.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/claude_code.rs new file mode 100644 index 000000000..dd8d3a3c6 --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/claude_code.rs @@ -0,0 +1,221 @@ +use super::shared::{ + build_provider_quota_execution_plan, build_quota_snapshot_payload, + default_provider_quota_execution_timeouts, execute_provider_quota_plan, + extract_execution_error_message, oauth_refresh_auto_removed_result, + persist_provider_quota_refresh_state, quota_key_auto_removed, + quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome, +}; +use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot}; +use crate::GatewayError; +use aether_admin::provider::quota::parse_claude_code_oauth_usage_response; +use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json; +use aether_contracts::ProxySnapshot; +use aether_data_contracts::repository::provider_catalog::{ + StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, +}; +use aether_provider_pool::build_claude_code_pool_quota_request; +use serde_json::json; +use std::time::{SystemTime, UNIX_EPOCH}; + +async fn execute_claude_code_quota_plan( + state: &AdminAppState<'_>, + transport: &AdminGatewayProviderTransportSnapshot, + authorization: (String, String), + proxy_override: Option<&ProxySnapshot>, +) -> Result { + let proxy = match proxy_override { + Some(proxy) => Some(proxy.clone()), + None => { + state + .resolve_transport_proxy_snapshot_with_tunnel_affinity(transport) + .await + } + }; + let timeouts = state + .resolve_transport_execution_timeouts(transport) + .or(Some(default_provider_quota_execution_timeouts( + proxy.as_ref(), + ))); + let spec = build_claude_code_pool_quota_request(&transport.key.id, authorization); + let plan = build_provider_quota_execution_plan( + transport, + spec, + proxy, + state.resolve_transport_profile(transport), + timeouts, + ); + + execute_provider_quota_plan(state, transport, plan, "claude_code").await +} + +pub(crate) async fn refresh_claude_code_provider_quota_locally( + state: &AdminAppState<'_>, + provider: &StoredProviderCatalogProvider, + endpoint: &StoredProviderCatalogEndpoint, + keys: Vec, + proxy_override: Option, +) -> Result, GatewayError> { + let mut results = Vec::new(); + let mut success_count = 0usize; + let mut failed_count = 0usize; + let mut auto_removed_count = 0usize; + + for key in keys { + let transport = match state + .read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id) + .await? + { + Some(transport) => transport, + None => { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "Provider transport snapshot unavailable", + })); + continue; + } + }; + + let authorization = match state.resolve_local_oauth_header_auth(&transport).await? { + Some(auth) => auth, + _ => { + if quota_key_auto_removed(state, &key.id).await? { + auto_removed_count += 1; + results.push(oauth_refresh_auto_removed_result(&key)); + continue; + } + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "缺少 OAuth 认证信息,请先授权/刷新 Token", + })); + continue; + } + }; + + let result = match execute_claude_code_quota_plan( + state, + &transport, + authorization, + proxy_override.as_ref(), + ) + .await? + { + ProviderQuotaExecutionOutcome::Response(result) => result, + ProviderQuotaExecutionOutcome::Failure(_) => { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "oauth/usage 请求执行失败", + "status_code": 502, + })); + continue; + } + }; + + let now_unix_secs = SystemTime::now() + .duration_since(UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()) + .unwrap_or(0); + let mut metadata_update = None::; + let (oauth_invalid_at_unix_secs, oauth_invalid_reason) = + quota_refresh_success_invalid_state(&key); + let mut status = "error".to_string(); + let mut message = None::; + + if result.status_code == 200 { + if let Some(body_json) = result + .body + .as_ref() + .and_then(|body| body.json_body.as_ref()) + { + metadata_update = parse_claude_code_oauth_usage_response(body_json, now_unix_secs) + .map(|metadata| json!({ "claude_code": metadata })); + if metadata_update.is_some() { + status = "success".to_string(); + } else { + status = "no_metadata".to_string(); + message = Some("响应中未包含额度窗口".to_string()); + } + } else { + status = "no_metadata".to_string(); + message = Some("响应中未包含配额信息".to_string()); + } + } else { + message = Some(match result.status_code { + 401 => "oauth/usage 返回 401,Token 可能已失效,请刷新 Token".to_string(), + 403 => "oauth/usage 返回 403,该账号缺少 user:profile 权限(如 Setup Token),无法查询额度" + .to_string(), + 429 => "oauth/usage 被限流,请稍后重试".to_string(), + code => format!("oauth/usage 返回状态码 {code}"), + }); + } + + if !persist_provider_quota_refresh_state( + state, + &key.id, + metadata_update.as_ref(), + oauth_invalid_at_unix_secs, + oauth_invalid_reason, + None, + ) + .await? + { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "Key 状态写入失败", + })); + continue; + } + + if status == "success" { + success_count += 1; + } else { + failed_count += 1; + } + + let mut payload = serde_json::Map::new(); + payload.insert("key_id".to_string(), json!(key.id)); + payload.insert("key_name".to_string(), json!(key.name)); + payload.insert("status".to_string(), json!(status)); + if let Some(message) = message { + payload.insert("message".to_string(), json!(message)); + } + if let Some(metadata) = metadata_update + .as_ref() + .and_then(|value| value.get("claude_code")) + { + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("claude_code", Some(metadata)), + ); + } + if let Some(quota_snapshot) = build_quota_snapshot_payload( + "claude_code", + key.status_snapshot.as_ref(), + metadata_update.as_ref(), + ) { + payload.insert("quota_snapshot".to_string(), quota_snapshot); + } + results.push(serde_json::Value::Object(payload)); + } + + Ok(Some(json!({ + "success": success_count, + "failed": failed_count, + "total": results.len(), + "results": results, + "message": format!("已处理 {} 个 Key", results.len()), + "auto_removed": auto_removed_count, + }))) +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/dispatch.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/dispatch.rs index 2963ba5e7..fee4e20ef 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/dispatch.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/dispatch.rs @@ -3,11 +3,13 @@ use std::pin::Pin; use super::antigravity::refresh_antigravity_provider_quota_locally; use super::chatgpt_web::refresh_chatgpt_web_provider_quota_locally; +use super::claude_code::refresh_claude_code_provider_quota_locally; use super::codex::refresh_codex_provider_quota_locally; use super::gemini_cli::refresh_gemini_cli_provider_quota_locally; use super::grok::refresh_grok_provider_quota_locally; use super::kiro::refresh_kiro_provider_quota_locally; use super::windsurf::refresh_windsurf_provider_quota_locally; +use super::xai::refresh_xai_provider_quota_locally; use crate::handlers::admin::request::AdminAppState; use crate::GatewayError; use aether_contracts::ProxySnapshot; @@ -35,6 +37,10 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] = "chatgpt_web", refresh_chatgpt_web_provider_quota_locally_boxed, ), + ( + "claude_code", + refresh_claude_code_provider_quota_locally_boxed, + ), ("codex", refresh_codex_provider_quota_locally_boxed), ( "gemini_cli", @@ -43,6 +49,7 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] = ("grok", refresh_grok_provider_quota_locally_boxed), ("kiro", refresh_kiro_provider_quota_locally_boxed), ("windsurf", refresh_windsurf_provider_quota_locally_boxed), + ("xai", refresh_xai_provider_quota_locally_boxed), ]; pub(crate) async fn refresh_provider_pool_quota_locally( @@ -111,6 +118,22 @@ fn refresh_codex_provider_quota_locally_boxed<'a>( )) } +fn refresh_claude_code_provider_quota_locally_boxed<'a>( + state: &'a AdminAppState<'a>, + provider: &'a StoredProviderCatalogProvider, + endpoint: &'a StoredProviderCatalogEndpoint, + keys: Vec, + proxy_override: Option, +) -> ProviderQuotaRefreshFuture<'a> { + Box::pin(refresh_claude_code_provider_quota_locally( + state, + provider, + endpoint, + keys, + proxy_override, + )) +} + fn refresh_gemini_cli_provider_quota_locally_boxed<'a>( state: &'a AdminAppState<'a>, provider: &'a StoredProviderCatalogProvider, @@ -174,3 +197,19 @@ fn refresh_windsurf_provider_quota_locally_boxed<'a>( proxy_override, )) } + +fn refresh_xai_provider_quota_locally_boxed<'a>( + state: &'a AdminAppState<'a>, + provider: &'a StoredProviderCatalogProvider, + endpoint: &'a StoredProviderCatalogEndpoint, + keys: Vec, + proxy_override: Option, +) -> ProviderQuotaRefreshFuture<'a> { + Box::pin(refresh_xai_provider_quota_locally( + state, + provider, + endpoint, + keys, + proxy_override, + )) +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/mod.rs index 8629fff76..1ab388d30 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/mod.rs @@ -1,5 +1,6 @@ pub(crate) mod antigravity; pub(crate) mod chatgpt_web; +pub(crate) mod claude_code; pub(crate) mod codex; pub(crate) mod dispatch; pub(crate) mod gemini_cli; @@ -7,3 +8,4 @@ pub(crate) mod grok; pub(crate) mod kiro; pub(crate) mod shared; pub(crate) mod windsurf; +pub(crate) mod xai; diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/shared.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/shared.rs index 7f1129faf..5be9c5cb0 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/shared.rs @@ -1713,8 +1713,10 @@ fn provider_quota_url_has_allowed_origin(provider_name: &str, value: &str) -> bo | "daily-cloudcode-pa.sandbox.googleapis.com" ), "gemini_cli" => host == "cloudcode-pa.googleapis.com", + "claude_code" => host == "api.anthropic.com", "chatgpt_web" | "codex" => host == "chatgpt.com", "grok" => host == "grok.com", + "xai" => host == "cli-chat-proxy.grok.com", "windsurf" => host == "server.codeium.com", "kiro" => kiro_quota_host_is_allowed(host), _ => false, @@ -1814,6 +1816,14 @@ mod tests { ), ("codex", "https://chatgpt.com/backend-api/wham/usage"), ("grok", "https://grok.com/rest/rate-limits"), + ( + "xai", + "https://cli-chat-proxy.grok.com/v1/billing?format=credits", + ), + ( + "xai", + "https://cli-chat-proxy.grok.com/v1/user", + ), ( "windsurf", "https://server.codeium.com/exa.seat_management_pb.SeatManagementService/GetUserStatus", @@ -1847,6 +1857,11 @@ mod tests { "https://chatgpt.com.attacker.test/backend-api/wham/usage", ), ("grok", "https://grok.com.attacker.test/rest/rate-limits"), + ( + "xai", + "https://cli-chat-proxy.grok.com.attacker.test/v1/billing", + ), + ("xai", "https://api.x.ai/v1/billing?format=credits"), ("windsurf", "https://server.codeium.com.attacker.test/quota"), ( "gemini_cli", diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/xai.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/xai.rs new file mode 100644 index 000000000..32cae49fd --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/xai.rs @@ -0,0 +1,313 @@ +use super::shared::{ + build_provider_quota_execution_plan, build_quota_snapshot_payload, + default_provider_quota_execution_timeouts, execute_provider_quota_plan, + extract_execution_error_message, oauth_refresh_auto_removed_result, + persist_provider_quota_refresh_state, quota_key_auto_removed, + quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome, +}; +use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot}; +use crate::GatewayError; +use aether_admin::provider::quota::parse_xai_billing_response; +use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json; +use aether_contracts::ProxySnapshot; +use aether_data_contracts::repository::provider_catalog::{ + StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, +}; +use aether_provider_pool::{build_xai_pool_billing_request, build_xai_pool_user_request}; +use aether_provider_transport::xai::{ + extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, xai_auth_uses_api, +}; +use serde_json::{json, Value}; +use std::time::{SystemTime, UNIX_EPOCH}; + +async fn execute_xai_quota_plan( + state: &AdminAppState<'_>, + transport: &AdminGatewayProviderTransportSnapshot, + spec: aether_provider_pool::ProviderPoolQuotaRequestSpec, + proxy_override: Option<&ProxySnapshot>, +) -> Result { + let proxy = match proxy_override { + Some(proxy) => Some(proxy.clone()), + None => { + state + .resolve_transport_proxy_snapshot_with_tunnel_affinity(transport) + .await + } + }; + let timeouts = state + .resolve_transport_execution_timeouts(transport) + .or(Some(default_provider_quota_execution_timeouts( + proxy.as_ref(), + ))); + let plan = build_provider_quota_execution_plan( + transport, + spec, + proxy, + state.resolve_transport_profile(transport), + timeouts, + ); + + execute_provider_quota_plan(state, transport, plan, "xai").await +} + +fn xai_authorization_from_header(authorization: &(String, String)) -> (String, String) { + authorization.clone() +} + +fn enrich_xai_subscription_title(mut metadata: Value, auth_config: Option<&str>) -> Value { + if metadata + .get("subscription_title") + .and_then(Value::as_str) + .map(str::trim) + .is_some_and(|value| !value.is_empty()) + { + return metadata; + } + let Some(config) = auth_config + .map(str::trim) + .filter(|value| !value.is_empty()) + .and_then(|value| serde_json::from_str::(value).ok()) + else { + return metadata; + }; + let title = ["subscription_tier", "subscriptionTier", "tier", "plan"] + .iter() + .find_map(|field| { + config + .get(*field) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }); + if let Some(title) = title { + if let Some(object) = metadata.as_object_mut() { + object.insert("subscription_title".to_string(), json!(title)); + } + } + metadata +} + +pub(crate) async fn refresh_xai_provider_quota_locally( + state: &AdminAppState<'_>, + provider: &StoredProviderCatalogProvider, + endpoint: &StoredProviderCatalogEndpoint, + keys: Vec, + proxy_override: Option, +) -> Result, GatewayError> { + let mut results = Vec::new(); + let mut success_count = 0usize; + let mut failed_count = 0usize; + let mut auto_removed_count = 0usize; + + for key in keys { + let transport = match state + .read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id) + .await? + { + Some(transport) => transport, + None => { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "Provider transport snapshot unavailable", + })); + continue; + } + }; + + if xai_auth_uses_api( + transport.key.auth_type.as_str(), + transport.key.decrypted_auth_config.as_deref(), + ) { + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "skipped", + "message": "xAI API Key 账号没有 Grok Build 订阅额度接口,请使用设备授权账号查询额度。", + })); + continue; + } + + let authorization = match state.resolve_local_oauth_header_auth(&transport).await? { + Some(auth) => auth, + _ => { + if quota_key_auto_removed(state, &key.id).await? { + auto_removed_count += 1; + results.push(oauth_refresh_auto_removed_result(&key)); + continue; + } + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "缺少 OAuth 认证信息,请先授权/刷新 Token", + })); + continue; + } + }; + + let fallback_user_id = + extract_xai_user_id_from_auth_config(transport.key.decrypted_auth_config.as_deref()); + let user_id = match execute_xai_quota_plan( + state, + &transport, + build_xai_pool_user_request( + &transport.key.id, + xai_authorization_from_header(&authorization), + ), + proxy_override.as_ref(), + ) + .await? + { + ProviderQuotaExecutionOutcome::Response(result) if result.status_code == 200 => result + .body + .as_ref() + .and_then(|body| body.json_body.as_ref()) + .and_then(extract_xai_user_id_from_value) + .or(fallback_user_id), + _ => fallback_user_id, + }; + + let result = match execute_xai_quota_plan( + state, + &transport, + build_xai_pool_billing_request( + &transport.key.id, + xai_authorization_from_header(&authorization), + user_id.as_deref(), + ), + proxy_override.as_ref(), + ) + .await? + { + ProviderQuotaExecutionOutcome::Response(result) => result, + ProviderQuotaExecutionOutcome::Failure(_) => { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "xAI billing 请求执行失败", + "status_code": 502, + })); + continue; + } + }; + + let now_unix_secs = SystemTime::now() + .duration_since(UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()) + .unwrap_or(0); + let mut metadata_update = None::; + let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) = + quota_refresh_success_invalid_state(&key); + let mut status = "error".to_string(); + let mut message = None::; + + if result.status_code == 200 { + if let Some(body_json) = result + .body + .as_ref() + .and_then(|body| body.json_body.as_ref()) + { + metadata_update = + parse_xai_billing_response(body_json, now_unix_secs).map(|metadata| { + json!({ + "xai": enrich_xai_subscription_title( + metadata, + transport.key.decrypted_auth_config.as_deref(), + ) + }) + }); + if metadata_update.is_some() { + status = "success".to_string(); + } else { + status = "no_metadata".to_string(); + message = Some("响应中未包含可用的 Grok Build 额度信息".to_string()); + } + } else { + status = "no_metadata".to_string(); + message = Some("响应中未包含配额信息".to_string()); + } + } else { + message = Some( + extract_execution_error_message(&result) + .unwrap_or_else(|| format!("xAI billing 返回状态码 {}", result.status_code)), + ); + if result.status_code == 401 || result.status_code == 403 { + let reason = message + .clone() + .unwrap_or_else(|| "账户访问被禁止".to_string()); + oauth_invalid_at_unix_secs = Some(now_unix_secs); + oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}")); + status = if result.status_code == 401 { + "unauthorized".to_string() + } else { + "forbidden".to_string() + }; + } + } + + if !persist_provider_quota_refresh_state( + state, + &key.id, + metadata_update.as_ref(), + oauth_invalid_at_unix_secs, + oauth_invalid_reason, + None, + ) + .await? + { + failed_count += 1; + results.push(json!({ + "key_id": key.id, + "key_name": key.name, + "status": "error", + "message": "Key 状态写入失败", + })); + continue; + } + + if status == "success" { + success_count += 1; + } else { + failed_count += 1; + } + + let mut payload = serde_json::Map::new(); + payload.insert("key_id".to_string(), json!(key.id)); + payload.insert("key_name".to_string(), json!(key.name)); + payload.insert("status".to_string(), json!(status)); + if let Some(message) = message { + payload.insert("message".to_string(), json!(message)); + } + if let Some(metadata) = metadata_update.as_ref().and_then(|value| value.get("xai")) { + payload.insert( + "metadata".to_string(), + admin_provider_metadata_bucket_safe_json("xai", Some(metadata)), + ); + } + if let Some(quota_snapshot) = build_quota_snapshot_payload( + "xai", + key.status_snapshot.as_ref(), + metadata_update.as_ref(), + ) { + payload.insert("quota_snapshot".to_string(), quota_snapshot); + } + results.push(serde_json::Value::Object(payload)); + } + + Ok(Some(json!({ + "success": success_count, + "failed": failed_count, + "total": results.len(), + "results": results, + "message": format!("已处理 {} 个 Key", results.len()), + "auto_removed": auto_removed_count, + }))) +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool/config.rs b/apps/aether-gateway/src/handlers/admin/provider/pool/config.rs index 216f00405..779589534 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool/config.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool/config.rs @@ -411,6 +411,7 @@ pub(crate) fn admin_provider_pool_config_from_config_value( unschedulable_rules: Vec::new(), lru_enabled: false, skip_exhausted_accounts: false, + reserve_minimum_quota: false, sticky_session_ttl_seconds: 3600, latency_window_seconds: 3600, latency_sample_limit: 50, @@ -446,6 +447,10 @@ pub(crate) fn admin_provider_pool_config_from_config_value( .get("skip_exhausted_accounts") .and_then(Value::as_bool) .unwrap_or(false), + reserve_minimum_quota: pool_advanced + .get("reserve_minimum_quota") + .and_then(Value::as_bool) + .unwrap_or(false), sticky_session_ttl_seconds: pool_advanced .get("sticky_session_ttl_seconds") .and_then(json_u64) @@ -574,6 +579,22 @@ mod tests { let config = admin_provider_pool_config(&provider).expect("pool config should exist"); assert!(!config.skip_exhausted_accounts); + assert!(!config.reserve_minimum_quota); + } + + #[test] + fn parses_reserve_minimum_quota_independently_of_skip_exhausted_accounts() { + for enabled in [false, true] { + let provider = sample_provider(json!({ + "pool_advanced": { + "reserve_minimum_quota": enabled, + "skip_exhausted_accounts": false + } + })); + let config = admin_provider_pool_config(&provider).expect("pool config should exist"); + assert_eq!(config.reserve_minimum_quota, enabled); + assert!(!config.skip_exhausted_accounts); + } } #[test] diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/writes.rs b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/writes.rs index d25e675f8..2211d9e40 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/writes.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/writes.rs @@ -651,6 +651,7 @@ mod tests { unschedulable_rules: Vec::new(), lru_enabled: true, skip_exhausted_accounts: false, + reserve_minimum_quota: false, sticky_session_ttl_seconds: 120, latency_window_seconds: 600, latency_sample_limit: 10, diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/payloads.rs b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/payloads.rs index 5dec0f9d5..c31793973 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/payloads.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/payloads.rs @@ -932,6 +932,13 @@ fn admin_pool_build_account_quota( return Some(account_quota); } } + "xai" => { + if let Some(account_quota) = + admin_pool_build_kiro_account_quota_from_snapshot(quota_snapshot) + { + return Some(account_quota); + } + } "chatgpt_web" => { if let Some(account_quota) = admin_pool_build_chatgpt_web_account_quota_from_snapshot(quota_snapshot) @@ -1110,10 +1117,16 @@ pub(super) fn build_admin_pool_key_payload( let health_score = admin_pool_health_score(key); let circuit_breaker_open = false; let auth_semantics = provider_key_auth_semantics(key, provider_type); - let account_quota_exhausted = pool_config - .as_ref() - .is_some_and(|config| config.skip_exhausted_accounts) - && admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type); + let account_quota_exhausted = pool_config.as_ref().is_some_and(|config| { + (config.skip_exhausted_accounts + && admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type)) + || (config.reserve_minimum_quota + && admin_provider_pool_pure::admin_pool_key_minimum_quota_reached( + key, + provider_type, + None, + )) + }); let auth_config = state.parse_catalog_auth_config_json(key); let oauth_expires_at = admin_pool_derive_oauth_expires_at(provider_type, key, auth_config.as_ref()); @@ -1591,4 +1604,29 @@ mod tests { Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string()) ); } + + #[test] + fn xai_account_quota_is_rendered_as_remaining_percent() { + let quota_snapshot = json!({ + "provider_type": "xai", + "code": "ok", + "exhausted": false, + "plan_type": "SuperGrok", + "windows": [ + { + "code": "usage", + "label": "周额度", + "scope": "account", + "used_ratio": 0.46, + "remaining_ratio": 0.54 + } + ] + }); + let quota_snapshot = quota_snapshot.as_object().unwrap(); + + assert_eq!( + admin_pool_build_account_quota("xai", Some(quota_snapshot)), + Some("剩余 54.0%".to_string()) + ); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/read_routes/keys.rs b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/read_routes/keys.rs index 3d6b2ea80..96ced1b6c 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/read_routes/keys.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/read_routes/keys.rs @@ -488,9 +488,16 @@ pub(super) fn admin_pool_key_visible_status_filter( ) { return status; } - if pool_config.is_some_and(|config| config.skip_exhausted_accounts) - && admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type) - { + if pool_config.is_some_and(|config| { + (config.skip_exhausted_accounts + && admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type)) + || (config.reserve_minimum_quota + && admin_provider_pool_pure::admin_pool_key_minimum_quota_reached( + key, + provider_type, + None, + )) + }) { return "quota_exhausted"; } if !key.is_active { diff --git a/apps/aether-gateway/src/handlers/admin/provider/query/models/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/query/models/mod.rs index cb0cc44e5..e84bb48e8 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/query/models/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/query/models/mod.rs @@ -594,11 +594,12 @@ async fn provider_query_fetch_models_for_key( }); } + let dynamic_client_version = crate::ai_serving::api::codex_client_version(); let client_version = is_codex.then(|| { codex_catalog .as_ref() .map(|catalog| catalog.client_version.as_str()) - .unwrap_or(crate::ai_serving::CODEX_CLIENT_VERSION) + .unwrap_or(dynamic_client_version.as_str()) }); let outcome = match fetch_models_from_transports_for_management(state.app(), &transports, client_version) diff --git a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs index 8319b1409..0752cb059 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs @@ -1340,6 +1340,7 @@ async fn provider_query_apply_pool_scheduler_to_test_candidates( provider.id.clone(), provider_query_ai_pool_runtime_state(&runtime), ); + let reserve_minimum_quota = pool_config.reserve_minimum_quota; let pool_config = provider_query_ai_pool_scheduling_config(pool_config, provider.provider_type.as_str()); let inputs = keys @@ -1351,6 +1352,14 @@ async fn provider_query_apply_pool_scheduler_to_test_candidates( effective_model: effective_model.to_string(), scheduler_skip_reason: None, }; + let mut key_context = + provider_query_pool_catalog_key_context(state, &key, &provider.provider_type); + key_context.quota_exhausted |= reserve_minimum_quota + && admin_provider_pool_pure::admin_pool_key_minimum_quota_reached( + &key, + &provider.provider_type, + Some(effective_model), + ); AiPoolCandidateInput { facts: AiPoolCandidateFacts { provider_id: provider.id.clone(), @@ -1362,11 +1371,7 @@ async fn provider_query_apply_pool_scheduler_to_test_candidates( key_internal_priority: key.internal_priority, }, pool_config: Some(pool_config.clone()), - key_context: provider_query_pool_catalog_key_context( - state, - &key, - &provider.provider_type, - ), + key_context, candidate, } }) @@ -3511,6 +3516,11 @@ async fn provider_query_execute_standard_test_candidate( codex_model_capabilities.as_ref(), ); } + crate::provider_transport::insert_cli_identity_headers_if_needed( + &transport, + provider_api_format, + &mut request_headers, + ); if !uses_vertex_query_auth { if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) diff --git a/apps/aether-gateway/src/handlers/admin/provider/shared/support.rs b/apps/aether-gateway/src/handlers/admin/provider/shared/support.rs index a2daf220c..b1574f609 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/shared/support.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/shared/support.rs @@ -65,6 +65,7 @@ pub(crate) struct AdminProviderPoolConfig { pub(crate) unschedulable_rules: Vec, pub(crate) lru_enabled: bool, pub(crate) skip_exhausted_accounts: bool, + pub(crate) reserve_minimum_quota: bool, pub(crate) sticky_session_ttl_seconds: u64, pub(crate) latency_window_seconds: u64, pub(crate) latency_sample_limit: u64, diff --git a/apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs b/apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs index 65a6d364e..2649209f8 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs @@ -4,9 +4,9 @@ pub(crate) fn normalize_provider_type_input(value: &str) -> Result Ok(normalized), + | "antigravity" | "vertex_ai" | "grok" | "windsurf" | "xai" => Ok(normalized), _ => Err( - "provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf" + "provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf / xai" .to_string(), ), } @@ -405,6 +405,14 @@ mod tests { ); } + #[test] + fn normalize_provider_type_supports_xai() { + assert_eq!( + normalize_provider_type_input(" xAI ").expect("type should normalize"), + "xai" + ); + } + #[test] fn normalize_api_format_list_dedupes_canonical_formats() { assert_eq!( diff --git a/apps/aether-gateway/src/handlers/admin/request/billing.rs b/apps/aether-gateway/src/handlers/admin/request/billing.rs index 116369798..ad75dee7f 100644 --- a/apps/aether-gateway/src/handlers/admin/request/billing.rs +++ b/apps/aether-gateway/src/handlers/admin/request/billing.rs @@ -388,10 +388,11 @@ impl<'a> AdminAppState<'a> { balance_type: &str, operator_id: Option<&str>, description: Option<&str>, + clamp_deduction_to_available_balance: bool, ) -> Result< Option<( aether_data::repository::wallet::StoredWalletSnapshot, - crate::AdminWalletTransactionRecord, + Option, )>, GatewayError, > { @@ -402,6 +403,7 @@ impl<'a> AdminAppState<'a> { balance_type, operator_id, description, + clamp_deduction_to_available_balance, ) .await } diff --git a/apps/aether-gateway/src/handlers/admin/request/state.rs b/apps/aether-gateway/src/handlers/admin/request/state.rs index 7de1273e3..9564e03af 100644 --- a/apps/aether-gateway/src/handlers/admin/request/state.rs +++ b/apps/aether-gateway/src/handlers/admin/request/state.rs @@ -17,6 +17,64 @@ impl<'a> AdminAppState<'a> { pub(crate) fn cloned_app(&self) -> AppState { self.app.clone() } + + pub(crate) async fn get_admin_user_wallet_balance_batch( + &self, + admin_user_id: &str, + idempotency_key: &str, + request_fingerprint: &str, + ) -> Result< + Option, + GatewayError, + > { + self.app + .get_admin_user_wallet_balance_batch( + admin_user_id, + idempotency_key, + request_fingerprint, + ) + .await + } + + pub(crate) async fn prepare_admin_user_wallet_balance_batch( + &self, + input: aether_data::repository::wallet::PrepareAdminUserWalletBalanceBatchInput, + ) -> Result< + aether_data::repository::wallet::PrepareAdminUserWalletBalanceBatchOutcome, + GatewayError, + > { + self.app + .prepare_admin_user_wallet_balance_batch(input) + .await + } + + pub(crate) async fn record_admin_user_wallet_balance_batch_failure( + &self, + admin_user_id: &str, + idempotency_key: &str, + user_id: &str, + reason: &str, + ) -> Result + { + self.app + .record_admin_user_wallet_balance_batch_failure( + admin_user_id, + idempotency_key, + user_id, + reason, + ) + .await + } + + pub(crate) async fn adjust_admin_user_wallet_balance_batch_user( + &self, + input: aether_data::repository::wallet::AdjustWalletBalanceInBatchInput, + ) -> Result + { + self.app + .adjust_admin_user_wallet_balance_batch_user(input) + .await + } } impl<'a> AsRef for AdminAppState<'a> { diff --git a/apps/aether-gateway/src/handlers/admin/request/users.rs b/apps/aether-gateway/src/handlers/admin/request/users.rs index 39d8554a3..d0b163be6 100644 --- a/apps/aether-gateway/src/handlers/admin/request/users.rs +++ b/apps/aether-gateway/src/handlers/admin/request/users.rs @@ -126,6 +126,30 @@ impl<'a> AdminAppState<'a> { self.app.list_user_group_members(group_id).await } + pub(crate) async fn resolve_usage_user_group_member_ids( + &self, + group_id: &str, + include_inactive: bool, + exclude_admin: bool, + ) -> Result>, GatewayError> { + if self.find_user_group_by_id(group_id).await?.is_none() { + return Ok(None); + } + + let mut user_ids = self + .list_user_group_members(group_id) + .await? + .into_iter() + .filter(|member| !member.is_deleted) + .filter(|member| include_inactive || member.is_active) + .filter(|member| !exclude_admin || !member.role.eq_ignore_ascii_case("admin")) + .map(|member| member.user_id) + .collect::>(); + user_ids.sort(); + user_ids.dedup(); + Ok(Some(user_ids)) + } + pub(crate) async fn replace_user_group_members( &self, group_id: &str, diff --git a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/list.rs b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/list.rs index dc328761c..45932de1d 100644 --- a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/list.rs +++ b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/list.rs @@ -64,6 +64,7 @@ pub(crate) async fn build_admin_list_user_api_keys_response( "rate_limit": record.rate_limit, "concurrent_limit": record.concurrent_limit, "feature_settings": record.feature_settings, + "ip_rules": record.ip_rules, "expires_at": format_optional_unix_secs_iso8601(record.expires_at_unix_secs), "last_used_at": format_optional_unix_secs_iso8601(record.last_used_at_unix_secs), "created_at": format_optional_unix_secs_iso8601(record.created_at_unix_secs), diff --git a/apps/aether-gateway/src/handlers/admin/users/batch.rs b/apps/aether-gateway/src/handlers/admin/users/batch.rs index 75d945001..d56b3f08d 100644 --- a/apps/aether-gateway/src/handlers/admin/users/batch.rs +++ b/apps/aether-gateway/src/handlers/admin/users/batch.rs @@ -1,11 +1,16 @@ use super::{ build_admin_users_bad_request_response, build_admin_users_permission_denied_response, - build_admin_users_read_only_response, disabled_user_policy_detail, disabled_user_policy_field, + build_admin_users_read_only_response, build_admin_users_wallet_permission_denied_response, + disabled_user_policy_detail, disabled_user_policy_field, + management_token_may_adjust_admin_wallet_balance, management_token_may_administer_user_accounts, normalize_admin_user_role, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::attach_admin_audit_response; use crate::GatewayError; +use aether_data::repository::wallet::{ + AdminUserWalletBalanceBatchUserOutcome, PrepareAdminUserWalletBalanceBatchOutcome, +}; use axum::{ body::{Body, Bytes}, http, @@ -13,9 +18,10 @@ use axum::{ Json, }; use serde_json::{json, Value}; +use sha2::{Digest as _, Sha256}; use std::collections::{BTreeMap, BTreeSet}; -#[derive(Debug, Clone, Default, serde::Deserialize)] +#[derive(Debug, Clone, Default, serde::Deserialize, serde::Serialize)] struct AdminUserSelectionFilters { #[serde(default)] search: Option, @@ -27,7 +33,7 @@ struct AdminUserSelectionFilters { group_id: Option, } -#[derive(Debug, Clone, Default)] +#[derive(Debug, Clone, Default, serde::Serialize)] struct AdminUserSelectionRequest { user_ids: Vec, group_ids: Vec, @@ -40,6 +46,7 @@ struct AdminUserBatchActionRequest { selection: AdminUserSelectionRequest, action: String, payload: Option, + idempotency_key: Option, } #[derive(Debug, serde::Deserialize)] @@ -48,6 +55,8 @@ struct RawAdminUserBatchActionRequest { action: String, #[serde(default)] payload: Option, + #[serde(default)] + idempotency_key: Option, } #[derive(Debug, Clone, Default)] @@ -68,7 +77,7 @@ struct AdminUserSelectionItem { matched_by: Vec, } -#[derive(Debug, Clone, serde::Serialize)] +#[derive(Debug, Clone, serde::Deserialize, serde::Serialize)] struct AdminUserSelectionWarning { #[serde(rename = "type")] warning_type: String, @@ -88,6 +97,7 @@ struct AdminUserBatchMutation { role: Option, is_active: Option, unlimited: Option, + wallet_balance_adjustment: Option, modified_fields: Vec<&'static str>, } @@ -97,6 +107,28 @@ impl AdminUserBatchMutation { } } +#[derive(Debug, Clone, Copy)] +struct AdminUserWalletBalanceAdjustment { + operation: AdminUserWalletBalanceOperation, + amount: f64, +} + +#[derive(Debug, Clone, Copy)] +enum AdminUserWalletBalanceOperation { + Add, + Deduct, +} + +enum AdminBatchWalletBalanceAdjustmentError { + WalletLookup, + BalanceAdjustment, +} + +enum AdminBatchWalletLimitModeError { + WalletLookup, + Mutation, +} + pub(in super::super) async fn build_admin_resolve_user_selection_response( state: &AdminAppState<'_>, _request_context: &AdminRequestContext<'_>, @@ -128,10 +160,29 @@ pub(in super::super) async fn build_admin_user_batch_action_response( Ok(value) => value, Err(detail) => return Ok(build_admin_user_batch_bad_request_response(detail)), }; - let mutation = match parse_batch_mutation(&request.action, request.payload) { + let mutation = match parse_batch_mutation(&request.action, request.payload.clone()) { Ok(value) => value, Err(detail) => return Ok(build_admin_user_batch_bad_request_response(detail)), }; + if mutation.wallet_balance_adjustment.is_some() { + if !management_token_may_adjust_admin_wallet_balance(request_context) { + return Ok(build_admin_users_wallet_permission_denied_response( + request_context, + )); + } + if !state.has_auth_wallet_write_capability() { + return Ok(build_admin_users_read_only_response( + "当前为只读模式,无法批量调整用户钱包余额", + )); + } + return build_admin_user_wallet_balance_batch_response( + state, + request_context, + request, + mutation, + ) + .await; + } let resolved = match resolve_admin_user_selection(state, request.selection).await { Ok(value) => value, Err(detail) => return Ok(build_admin_user_batch_bad_request_response(detail)), @@ -177,9 +228,29 @@ pub(in super::super) async fn build_admin_user_batch_action_response( .iter() .map(|user_id| json!({ "user_id": user_id, "reason": "用户不存在或已删除" })) .collect::>(); + let mut completed_user_ids = Vec::new(); + let mut uncertain_user_ids = Vec::new(); + let mut unprocessed_user_ids = Vec::new(); + let mut interrupted = false; - for item in &resolved.items { - if state.find_user_auth_by_id(&item.user_id).await?.is_none() { + for (item_index, item) in resolved.items.iter().enumerate() { + let user = match state.find_user_auth_by_id(&item.user_id).await { + Ok(user) => user, + Err(_) => { + record_batch_action_interruption( + &resolved.items, + item_index, + false, + "读取用户状态失败,批次已中止,该用户未执行", + &mut failures, + &mut uncertain_user_ids, + &mut unprocessed_user_ids, + ); + interrupted = true; + break; + } + }; + if user.is_none() { failures.push(json!({ "user_id": item.user_id, "reason": "用户不存在或已删除", @@ -202,17 +273,92 @@ pub(in super::super) async fn build_admin_user_batch_action_response( } if let Some(unlimited) = mutation.unlimited { - if !apply_batch_user_wallet_limit_mode(state, &item.user_id, unlimited).await? { - failures.push(json!({ - "user_id": item.user_id, - "reason": "用户钱包不可用", - })); - continue; + match apply_batch_user_wallet_limit_mode(state, &item.user_id, unlimited).await { + Ok(true) => {} + Ok(false) => { + failures.push(json!({ + "user_id": item.user_id, + "reason": "用户钱包不可用", + })); + continue; + } + Err(AdminBatchWalletLimitModeError::WalletLookup) => { + record_batch_action_interruption( + &resolved.items, + item_index, + false, + "读取用户钱包失败,批次已中止,该用户未执行", + &mut failures, + &mut uncertain_user_ids, + &mut unprocessed_user_ids, + ); + interrupted = true; + break; + } + Err(AdminBatchWalletLimitModeError::Mutation) => { + record_batch_action_interruption( + &resolved.items, + item_index, + true, + "用户钱包更新结果未确认,批次已中止,请核对钱包后再重试", + &mut failures, + &mut uncertain_user_ids, + &mut unprocessed_user_ids, + ); + interrupted = true; + break; + } } } - if mutation.has_auth_user_fields() - && state + if let Some(adjustment) = mutation.wallet_balance_adjustment { + match apply_batch_user_wallet_balance_adjustment( + state, + &item.user_id, + adjustment, + current_admin_user_id, + ) + .await + { + Ok(true) => {} + Ok(false) => { + failures.push(json!({ + "user_id": item.user_id, + "reason": "用户钱包不可用", + })); + continue; + } + Err(AdminBatchWalletBalanceAdjustmentError::WalletLookup) => { + record_batch_action_interruption( + &resolved.items, + item_index, + false, + "读取用户钱包失败,批次已中止,该用户未执行", + &mut failures, + &mut uncertain_user_ids, + &mut unprocessed_user_ids, + ); + interrupted = true; + break; + } + Err(AdminBatchWalletBalanceAdjustmentError::BalanceAdjustment) => { + record_batch_action_interruption( + &resolved.items, + item_index, + true, + "余额调整结果未确认,批次已中止,请核对钱包后再重试", + &mut failures, + &mut uncertain_user_ids, + &mut unprocessed_user_ids, + ); + interrupted = true; + break; + } + } + } + + if mutation.has_auth_user_fields() { + let updated_user = match state .update_local_auth_user_admin_fields( &item.user_id, mutation.role.clone(), @@ -226,22 +372,39 @@ pub(in super::super) async fn build_admin_user_batch_action_response( None, mutation.is_active, ) - .await? - .is_none() - { - failures.push(json!({ - "user_id": item.user_id, - "reason": "用户不存在或已删除", - })); - continue; + .await + { + Ok(user) => user, + Err(_) => { + record_batch_action_interruption( + &resolved.items, + item_index, + true, + "用户更新结果未确认,批次已中止,请核对后再重试", + &mut failures, + &mut uncertain_user_ids, + &mut unprocessed_user_ids, + ); + interrupted = true; + break; + } + }; + if updated_user.is_none() { + failures.push(json!({ + "user_id": item.user_id, + "reason": "用户不存在或已删除", + })); + continue; + } } success += 1; + completed_user_ids.push(item.user_id.clone()); } let failed = failures.len(); let total = success + failed; - let response = Json(json!({ + let mut response_payload = json!({ "total": total, "success": success, "failed": failed, @@ -249,8 +412,14 @@ pub(in super::super) async fn build_admin_user_batch_action_response( "warnings": resolved.warnings, "action": request.action.trim().to_ascii_lowercase(), "modified_fields": mutation.modified_fields, - })) - .into_response(); + "interrupted": interrupted, + }); + if interrupted { + response_payload["completed_user_ids"] = json!(completed_user_ids); + response_payload["uncertain_user_ids"] = json!(uncertain_user_ids); + response_payload["unprocessed_user_ids"] = json!(unprocessed_user_ids); + } + let response = Json(response_payload).into_response(); Ok(attach_admin_audit_response( response, @@ -261,6 +430,408 @@ pub(in super::super) async fn build_admin_user_batch_action_response( )) } +fn record_batch_action_interruption( + items: &[AdminUserSelectionItem], + item_index: usize, + current_result_uncertain: bool, + reason: &str, + failures: &mut Vec, + uncertain_user_ids: &mut Vec, + unprocessed_user_ids: &mut Vec, +) { + let current_item = &items[item_index]; + failures.push(json!({ + "user_id": current_item.user_id, + "reason": reason, + })); + if current_result_uncertain { + uncertain_user_ids.push(current_item.user_id.clone()); + } else { + unprocessed_user_ids.push(current_item.user_id.clone()); + } + + for item in items.iter().skip(item_index + 1) { + failures.push(json!({ + "user_id": item.user_id, + "reason": "因前序错误未执行", + })); + unprocessed_user_ids.push(item.user_id.clone()); + } +} + +async fn build_admin_user_wallet_balance_batch_response( + state: &AdminAppState<'_>, + request_context: &AdminRequestContext<'_>, + request: AdminUserBatchActionRequest, + mutation: AdminUserBatchMutation, +) -> Result, GatewayError> { + let Some(idempotency_key) = request + .idempotency_key + .as_deref() + .map(str::trim) + .filter(|value| { + !value.is_empty() + && value.len() <= 128 + && value.bytes().all(|byte| (0x21..=0x7e).contains(&byte)) + }) + else { + return Ok(build_admin_user_batch_bad_request_response( + "余额批量操作必须提供有效的 idempotency_key".to_string(), + )); + }; + let Some(admin_user_id) = request_context + .decision() + .and_then(|decision| decision.admin_principal.as_ref()) + .map(|principal| principal.user_id.clone()) + else { + return Ok(build_admin_users_permission_denied_response( + request_context, + )); + }; + let action = request.action.trim().to_ascii_lowercase(); + let fingerprint_payload = json!({ + "selection": &request.selection, + "action": &action, + "payload": &request.payload, + }); + let encoded = serde_json::to_vec(&fingerprint_payload) + .map_err(|error| GatewayError::Internal(error.to_string()))?; + let request_fingerprint = format!("{:x}", Sha256::digest(encoded)); + + let existing = state + .get_admin_user_wallet_balance_batch(&admin_user_id, idempotency_key, &request_fingerprint) + .await?; + let batch = match existing { + Some(PrepareAdminUserWalletBalanceBatchOutcome::Conflict) => { + return Ok(build_admin_user_batch_idempotency_conflict_response()); + } + Some(PrepareAdminUserWalletBalanceBatchOutcome::Ready(batch)) => batch, + None => { + let resolved = + match resolve_admin_user_selection(state, request.selection.clone()).await { + Ok(value) => value, + Err(detail) => return Ok(build_admin_user_batch_bad_request_response(detail)), + }; + let warnings = serde_json::to_value(&resolved.warnings) + .ok() + .and_then(|value| value.as_array().cloned()) + .unwrap_or_default(); + let prepared = state + .prepare_admin_user_wallet_balance_batch( + aether_data::repository::wallet::PrepareAdminUserWalletBalanceBatchInput { + admin_user_id: admin_user_id.clone(), + idempotency_key: idempotency_key.to_string(), + request_fingerprint: request_fingerprint.clone(), + target_user_ids: resolved + .items + .iter() + .map(|item| item.user_id.clone()) + .collect(), + missing_user_ids: resolved.missing_user_ids, + warnings, + }, + ) + .await?; + match prepared { + PrepareAdminUserWalletBalanceBatchOutcome::Conflict => { + return Ok(build_admin_user_batch_idempotency_conflict_response()); + } + PrepareAdminUserWalletBalanceBatchOutcome::Ready(batch) => batch, + } + } + }; + let warnings: Vec = + serde_json::from_value(Value::Array(batch.warnings.clone())) + .map_err(|error| GatewayError::Internal(error.to_string()))?; + let resolved = ResolvedAdminUserSelection { + items: batch + .target_user_ids + .iter() + .map(|user_id| AdminUserSelectionItem { + user_id: user_id.clone(), + username: String::new(), + email: None, + role: "user".to_string(), + is_active: true, + matched_by: Vec::new(), + }) + .collect(), + missing_user_ids: batch.missing_user_ids.clone(), + warnings, + }; + let adjustment = mutation + .wallet_balance_adjustment + .expect("wallet balance action should have an adjustment"); + let signed_amount = match adjustment.operation { + AdminUserWalletBalanceOperation::Add => adjustment.amount, + AdminUserWalletBalanceOperation::Deduct => -adjustment.amount, + }; + + let mut outcomes = batch.user_outcomes; + let mut completed_user_ids = outcomes + .iter() + .filter_map(|(user_id, outcome)| { + matches!(outcome, AdminUserWalletBalanceBatchUserOutcome::Succeeded) + .then_some(user_id.clone()) + }) + .collect::>(); + let mut success = completed_user_ids.len(); + let mut failures = resolved + .missing_user_ids + .iter() + .map(|user_id| json!({ "user_id": user_id, "reason": "用户不存在或已删除" })) + .collect::>(); + for (user_id, outcome) in &outcomes { + if let AdminUserWalletBalanceBatchUserOutcome::Failed(reason) = outcome { + failures.push(json!({ "user_id": user_id, "reason": reason })); + } + } + let mut uncertain_user_ids = Vec::new(); + let mut unprocessed_user_ids = Vec::new(); + let mut interrupted = false; + + for (item_index, item) in resolved.items.iter().enumerate() { + if outcomes.contains_key(&item.user_id) { + continue; + } + match state.find_user_auth_by_id(&item.user_id).await { + Err(_) => { + failures.push(json!({ + "user_id": item.user_id, + "reason": "读取用户状态失败,批次已中止,该用户未执行", + })); + unprocessed_user_ids.push(item.user_id.clone()); + interrupted = true; + } + Ok(None) => { + let outcome = match state + .record_admin_user_wallet_balance_batch_failure( + &admin_user_id, + idempotency_key, + &item.user_id, + "用户不存在或已删除", + ) + .await + { + Ok(outcome) => outcome, + Err(_) => { + failures.push(json!({ + "user_id": item.user_id, + "reason": "记录用户状态失败,批次已中止,该用户未执行", + })); + unprocessed_user_ids.push(item.user_id.clone()); + interrupted = true; + append_wallet_batch_unprocessed_suffix( + &resolved.items, + item_index + 1, + &outcomes, + &mut failures, + &mut unprocessed_user_ids, + ); + break; + } + }; + outcomes.insert(item.user_id.clone(), outcome.clone()); + match outcome { + AdminUserWalletBalanceBatchUserOutcome::Succeeded => { + completed_user_ids.push(item.user_id.clone()); + success += 1; + } + AdminUserWalletBalanceBatchUserOutcome::Failed(reason) => { + failures.push(json!({ "user_id": item.user_id, "reason": reason })); + } + } + } + Ok(Some(_)) => { + let wallet = match state + .find_wallet(aether_data::repository::wallet::WalletLookupKey::UserId( + &item.user_id, + )) + .await + { + Ok(Some(wallet)) => wallet, + Ok(None) => { + let outcome = match state + .record_admin_user_wallet_balance_batch_failure( + &admin_user_id, + idempotency_key, + &item.user_id, + "用户钱包不可用", + ) + .await + { + Ok(outcome) => outcome, + Err(_) => { + failures.push(json!({ + "user_id": item.user_id, + "reason": "记录用户钱包状态失败,批次已中止,该用户未执行", + })); + unprocessed_user_ids.push(item.user_id.clone()); + interrupted = true; + append_wallet_batch_unprocessed_suffix( + &resolved.items, + item_index + 1, + &outcomes, + &mut failures, + &mut unprocessed_user_ids, + ); + break; + } + }; + outcomes.insert(item.user_id.clone(), outcome.clone()); + match outcome { + AdminUserWalletBalanceBatchUserOutcome::Succeeded => { + completed_user_ids.push(item.user_id.clone()); + success += 1; + } + AdminUserWalletBalanceBatchUserOutcome::Failed(reason) => { + failures.push(json!({ "user_id": item.user_id, "reason": reason })); + } + } + continue; + } + Err(_) => { + failures.push(json!({ + "user_id": item.user_id, + "reason": "读取用户钱包失败,批次已中止,该用户未执行", + })); + unprocessed_user_ids.push(item.user_id.clone()); + interrupted = true; + append_wallet_batch_unprocessed_suffix( + &resolved.items, + item_index + 1, + &outcomes, + &mut failures, + &mut unprocessed_user_ids, + ); + break; + } + }; + let result = state + .adjust_admin_user_wallet_balance_batch_user( + aether_data::repository::wallet::AdjustWalletBalanceInBatchInput { + admin_user_id: admin_user_id.clone(), + idempotency_key: idempotency_key.to_string(), + user_id: item.user_id.clone(), + adjustment: aether_data::repository::wallet::AdjustWalletBalanceInput { + wallet_id: wallet.id, + amount_usd: signed_amount, + balance_type: "recharge".to_string(), + operator_id: Some(admin_user_id.clone()), + description: Some("管理员批量调整用户余额".to_string()), + clamp_deduction_to_available_balance: true, + batch_context: None, + }, + }, + ) + .await; + match result { + Ok(AdminUserWalletBalanceBatchUserOutcome::Succeeded) => { + outcomes.insert( + item.user_id.clone(), + AdminUserWalletBalanceBatchUserOutcome::Succeeded, + ); + completed_user_ids.push(item.user_id.clone()); + success += 1; + } + Ok(AdminUserWalletBalanceBatchUserOutcome::Failed(reason)) => { + outcomes.insert( + item.user_id.clone(), + AdminUserWalletBalanceBatchUserOutcome::Failed(reason.clone()), + ); + failures.push(json!({ "user_id": item.user_id, "reason": reason })); + } + Err(_) => { + failures.push(json!({ + "user_id": item.user_id, + "reason": "余额调整结果未确认,批次已中止,请使用同一批次重试以核对结果", + })); + uncertain_user_ids.push(item.user_id.clone()); + interrupted = true; + append_wallet_batch_unprocessed_suffix( + &resolved.items, + item_index + 1, + &outcomes, + &mut failures, + &mut unprocessed_user_ids, + ); + break; + } + } + } + } + if interrupted { + for pending in resolved.items.iter().skip(item_index + 1) { + if outcomes.contains_key(&pending.user_id) { + continue; + } + failures.push(json!({ + "user_id": pending.user_id, + "reason": "因前序错误未执行", + })); + unprocessed_user_ids.push(pending.user_id.clone()); + } + break; + } + } + + let failed = failures.len(); + let total = success + failed; + let mut response_payload = json!({ + "total": total, + "success": success, + "failed": failed, + "failures": failures, + "warnings": resolved.warnings, + "action": action, + "modified_fields": mutation.modified_fields, + "interrupted": interrupted, + }); + if interrupted { + response_payload["completed_user_ids"] = json!(completed_user_ids); + response_payload["uncertain_user_ids"] = json!(uncertain_user_ids); + response_payload["unprocessed_user_ids"] = json!(unprocessed_user_ids); + } + let response = Json(response_payload).into_response(); + Ok(attach_admin_audit_response( + response, + "admin_users_batch_action_executed", + "batch_update_users", + "user_batch", + "users", + )) +} + +fn append_wallet_batch_unprocessed_suffix( + items: &[AdminUserSelectionItem], + start_index: usize, + outcomes: &BTreeMap, + failures: &mut Vec, + unprocessed_user_ids: &mut Vec, +) { + for pending in items.iter().skip(start_index) { + if outcomes.contains_key(&pending.user_id) { + continue; + } + failures.push(json!({ + "user_id": pending.user_id, + "reason": "因前序错误未执行", + })); + unprocessed_user_ids.push(pending.user_id.clone()); + } +} + +fn build_admin_user_batch_idempotency_conflict_response() -> Response { + ( + http::StatusCode::CONFLICT, + Json(json!({ + "detail": "idempotency_key was already used with a different request", + "error_code": "idempotency_key_conflict", + })), + ) + .into_response() +} + fn parse_resolve_selection_request( request_body: Option<&Bytes>, ) -> Result { @@ -286,6 +857,7 @@ fn parse_batch_action_request( selection: parse_selection_request_value(raw.selection)?, action: raw.action, payload: raw.payload, + idempotency_key: raw.idempotency_key, }) } _ => Err("Invalid JSON request body".to_string()), @@ -604,10 +1176,37 @@ fn parse_batch_mutation( }), "update_access_control" => parse_access_control_mutation(payload), "update_role" => parse_role_mutation(payload), + "adjust_wallet_balance" => parse_wallet_balance_adjustment_mutation(payload), _ => Err("不支持的批量操作".to_string()), } } +fn parse_wallet_balance_adjustment_mutation( + payload: Option, +) -> Result { + let Some(Value::Object(payload)) = payload else { + return Err("payload 必须是对象".to_string()); + }; + let operation = match payload.get("operation").and_then(Value::as_str) { + Some("add") => AdminUserWalletBalanceOperation::Add, + Some("deduct") => AdminUserWalletBalanceOperation::Deduct, + _ => return Err("operation 必须为 add 或 deduct".to_string()), + }; + let amount = payload + .get("amount") + .and_then(Value::as_f64) + .ok_or_else(|| "amount 必须为大于 0 的有限数字".to_string())?; + if !amount.is_finite() || amount <= 0.0 { + return Err("amount 必须为大于 0 的有限数字".to_string()); + } + + Ok(AdminUserBatchMutation { + wallet_balance_adjustment: Some(AdminUserWalletBalanceAdjustment { operation, amount }), + modified_fields: vec!["wallet_balance"], + ..AdminUserBatchMutation::default() + }) +} + fn parse_role_mutation(payload: Option) -> Result { let Some(Value::Object(payload)) = payload else { return Err("payload 必须是对象".to_string()); @@ -700,30 +1299,68 @@ async fn apply_batch_user_wallet_limit_mode( state: &AdminAppState<'_>, user_id: &str, unlimited: bool, -) -> Result { +) -> Result { let desired_limit_mode = if unlimited { "unlimited" } else { "finite" }; - match state + let wallet = state .find_wallet(aether_data::repository::wallet::WalletLookupKey::UserId( user_id, )) - .await? - { + .await + .map_err(|_| AdminBatchWalletLimitModeError::WalletLookup)?; + match wallet { Some(wallet) => { if wallet.limit_mode.eq_ignore_ascii_case(desired_limit_mode) { return Ok(true); } Ok(state .update_auth_user_wallet_limit_mode(user_id, desired_limit_mode) - .await? + .await + .map_err(|_| AdminBatchWalletLimitModeError::Mutation)? .is_some()) } None => Ok(state .initialize_auth_user_wallet(user_id, 0.0, unlimited) - .await? + .await + .map_err(|_| AdminBatchWalletLimitModeError::Mutation)? .is_some()), } } +async fn apply_batch_user_wallet_balance_adjustment( + state: &AdminAppState<'_>, + user_id: &str, + adjustment: AdminUserWalletBalanceAdjustment, + operator_id: Option<&str>, +) -> Result { + // Resolve only the wallet ID; the repository clamps the deduction under its row lock. + let Some(wallet) = state + .find_wallet(aether_data::repository::wallet::WalletLookupKey::UserId( + user_id, + )) + .await + .map_err(|_| AdminBatchWalletBalanceAdjustmentError::WalletLookup)? + else { + return Ok(false); + }; + let amount = match adjustment.operation { + AdminUserWalletBalanceOperation::Add => adjustment.amount, + AdminUserWalletBalanceOperation::Deduct => -adjustment.amount, + }; + + state + .admin_adjust_wallet_balance( + &wallet.id, + amount, + "recharge", + operator_id, + Some("管理员批量调整用户余额"), + true, + ) + .await + .map_err(|_| AdminBatchWalletBalanceAdjustmentError::BalanceAdjustment) + .map(|result| result.is_some()) +} + fn build_admin_user_batch_bad_request_response(detail: String) -> Response { if detail.as_str() == "缺少 user_id" { return build_admin_users_bad_request_response("缺少 user_id"); diff --git a/apps/aether-gateway/src/handlers/admin/users/mod.rs b/apps/aether-gateway/src/handlers/admin/users/mod.rs index e5c854748..5a5e9d16a 100644 --- a/apps/aether-gateway/src/handlers/admin/users/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/users/mod.rs @@ -49,12 +49,14 @@ use self::shared::AdminUpdateUserPatch; use self::shared::{ admin_default_user_initial_gift, build_admin_users_bad_request_response, build_admin_users_data_unavailable_response, build_admin_users_permission_denied_response, - build_admin_users_read_only_response, disabled_user_policy_detail, disabled_user_policy_field, - format_optional_datetime_iso8601, legacy_admin_list_policy_mode, - legacy_admin_rate_limit_policy_mode, management_token_may_administer_user_accounts, - normalize_admin_optional_user_email, normalize_admin_user_group_ids, normalize_admin_user_role, - normalize_admin_username, validate_admin_user_password, AdminCreateUserApiKeyRequest, - AdminCreateUserRequest, AdminToggleUserApiKeyLockRequest, AdminUpdateUserApiKeyRequest, + build_admin_users_read_only_response, build_admin_users_wallet_permission_denied_response, + disabled_user_policy_detail, disabled_user_policy_field, format_optional_datetime_iso8601, + legacy_admin_list_policy_mode, legacy_admin_rate_limit_policy_mode, + management_token_may_adjust_admin_wallet_balance, + management_token_may_administer_user_accounts, normalize_admin_optional_user_email, + normalize_admin_user_group_ids, normalize_admin_user_role, normalize_admin_username, + validate_admin_user_password, AdminCreateUserApiKeyRequest, AdminCreateUserRequest, + AdminToggleUserApiKeyLockRequest, AdminUpdateUserApiKeyRequest, }; pub(crate) use self::shared::{ normalize_admin_list_policy_mode, normalize_admin_rate_limit_policy_mode, diff --git a/apps/aether-gateway/src/handlers/admin/users/shared.rs b/apps/aether-gateway/src/handlers/admin/users/shared.rs index b304932c9..b74bd0c9a 100644 --- a/apps/aether-gateway/src/handlers/admin/users/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/users/shared.rs @@ -165,6 +165,18 @@ pub(super) fn management_token_may_administer_user_accounts( }) } +pub(super) fn management_token_may_adjust_admin_wallet_balance( + request_context: &crate::handlers::admin::request::AdminRequestContext<'_>, +) -> bool { + request_context.decision().is_some_and(|decision| { + crate::control::management_token_principal_has_permission(decision, "admin:wallets:write") + || crate::control::management_token_principal_has_permission( + decision, + "admin:wallets:admin", + ) + }) +} + pub(super) fn build_admin_users_permission_denied_response( request_context: &crate::handlers::admin::request::AdminRequestContext<'_>, ) -> Response { @@ -192,6 +204,34 @@ pub(super) fn build_admin_users_permission_denied_response( ) } +pub(super) fn build_admin_users_wallet_permission_denied_response( + request_context: &crate::handlers::admin::request::AdminRequestContext<'_>, +) -> Response { + let actor_id = request_context + .decision() + .and_then(|decision| decision.admin_principal.as_ref()) + .and_then(|principal| principal.management_token_id.as_deref()) + .unwrap_or("unknown"); + crate::handlers::admin::shared::attach_admin_audit_response( + ( + http::StatusCode::FORBIDDEN, + Json(json!({ + "detail": "management token permission denied", + "required_permissions": ["admin:wallets:write", "admin:wallets:admin"], + "permission_mode": "any_of", + "route_family": request_context.route_family(), + "route_kind": request_context.route_kind(), + "request_path": request_context.path(), + })), + ) + .into_response(), + "admin_user_wallet_balance_permission_denied", + "permission_denied", + "admin_user_wallet_balance", + actor_id, + ) +} + pub(super) fn normalize_admin_optional_user_email( value: Option<&str>, ) -> Result, String> { @@ -397,9 +437,52 @@ pub(super) fn format_optional_datetime_iso8601( #[cfg(test)] mod tests { - use super::{normalize_admin_user_api_formats, AdminUpdateUserApiKeyRequest}; + use super::{ + build_admin_users_wallet_permission_denied_response, normalize_admin_user_api_formats, + AdminUpdateUserApiKeyRequest, + }; + use crate::control::{GatewayControlDecision, GatewayPublicRequestContext}; + use crate::handlers::admin::request::AdminRequestContext; + use axum::http::{HeaderMap, Method, Uri}; use serde_json::json; + #[test] + fn wallet_permission_denial_uses_wallet_audit_category() { + let uri: Uri = "/api/admin/users/batch-action" + .parse() + .expect("uri should parse"); + let method = Method::POST; + let headers = HeaderMap::new(); + let decision = GatewayControlDecision::synthetic( + uri.path(), + Some("admin_proxy".to_string()), + Some("users_manage".to_string()), + Some("batch_user_action".to_string()), + Some("admin:users".to_string()), + ); + let context = GatewayPublicRequestContext::from_request_parts( + "trace-wallet-permission-denied", + &method, + &uri, + &headers, + Some(decision), + ); + let request_context = AdminRequestContext::new(&context); + + let response = build_admin_users_wallet_permission_denied_response(&request_context); + let event = response + .extensions() + .get::() + .expect("wallet denial should attach an audit event"); + + assert_eq!( + event.event_name, + "admin_user_wallet_balance_permission_denied" + ); + assert_eq!(event.action, "permission_denied"); + assert_eq!(event.target_type, "admin_user_wallet_balance"); + } + #[test] fn admin_user_api_formats_accept_current_canonical_signatures() { assert_eq!( diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/continuation.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/continuation.rs index d8de44134..7aaa58ba6 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/continuation.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/continuation.rs @@ -90,6 +90,8 @@ pub(super) struct ResponsesWebSocketContinuationRecord { /// request JSON can never set it. #[serde(default)] deepseek_opaque_reasoning_replay: bool, + #[serde(default)] + xai_encrypted_reasoning_replay: bool, /// A prior turn stored PII sentinels whose restore mapping exists only on /// the original downstream socket. Such a chain cannot safely resume on a /// new socket without leaking sentinels, so lookup succeeds but bootstrap @@ -122,6 +124,10 @@ impl ResponsesWebSocketContinuationRecord { normalization.reasoning_replay_policy(), crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque ), + xai_encrypted_reasoning_replay: matches!( + normalization.reasoning_replay_policy(), + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + ), has_connection_local_redaction, responses_lite_static_config, }; @@ -156,7 +162,9 @@ impl ResponsesWebSocketContinuationRecord { pub(super) fn reasoning_replay_policy( &self, ) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy { - if self.deepseek_opaque_reasoning_replay { + if self.xai_encrypted_reasoning_replay { + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + } else if self.deepseek_opaque_reasoning_replay { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque } else { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds @@ -476,6 +484,7 @@ mod tests { binding_fingerprint: [7; 32], normalization_fingerprint: [9; 32], deepseek_opaque_reasoning_replay: false, + xai_encrypted_reasoning_replay: false, has_connection_local_redaction: false, responses_lite_static_config: Some(ResponsesLiteStaticConfig::from_response_create( &json!({ @@ -714,6 +723,29 @@ mod tests { assert_eq!(decoded, record()); } + #[test] + fn serialized_record_preserves_xai_replay_policy_and_reads_legacy_records() { + let mut expected = record(); + expected.xai_encrypted_reasoning_replay = true; + let mut serialized = serde_json::to_value(&expected).unwrap(); + let decoded: ResponsesWebSocketContinuationRecord = + serde_json::from_value(serialized.clone()).unwrap(); + assert_eq!( + decoded.reasoning_replay_policy(), + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + ); + serialized + .as_object_mut() + .unwrap() + .remove("xai_encrypted_reasoning_replay"); + let legacy: ResponsesWebSocketContinuationRecord = + serde_json::from_value(serialized).unwrap(); + assert_eq!( + legacy.reasoning_replay_policy(), + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds + ); + } + #[test] fn serialized_record_preserves_only_the_server_derived_reasoning_replay_policy_bit() { let mut expected = record(); diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_usage.rs b/apps/aether-gateway/src/handlers/public/support/user_me_usage.rs index da54e73a8..fe0709c0f 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_usage.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_usage.rs @@ -567,6 +567,7 @@ fn build_users_me_usage_record_payload( "id": item.id, "model": item.model, "target_model": serde_json::Value::Null, + "response_model": item.provider_response_model(), "api_format": item.api_format, "endpoint_api_format": item.endpoint_api_format, "has_format_conversion": item.has_format_conversion, @@ -681,6 +682,7 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_ "client_ip": users_me_usage_metadata_string(item, "client_ip"), "user_agent": users_me_usage_metadata_string(item, "user_agent"), "target_model": item.target_model, + "response_model": item.provider_response_model(), "has_fallback": item.has_fallback(), }); payload["end_to_end_time_ms"] = json!(users_me_usage_metadata_u64(item, "end_to_end_time_ms")); @@ -1865,6 +1867,25 @@ mod tests { assert_eq!(payload["cache_creation_ephemeral_1h_input_tokens"], 6); } + #[test] + fn user_usage_payloads_expose_response_model_separately_from_mapping() { + let item = StoredRequestUsageAudit { + target_model: Some("provider-mapped-model".to_string()), + request_metadata: Some(json!({ + "provider_response_model": "gpt-5.1" + })), + ..sample_usage("completed") + }; + + let record = build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false); + let active = build_users_me_usage_active_payload(&item); + + for payload in [&record, &active] { + assert_eq!(payload["target_model"], "provider-mapped-model"); + assert_eq!(payload["response_model"], "gpt-5.1"); + } + } + #[test] fn user_usage_payloads_project_end_to_end_timings_from_metadata() { let item = StoredRequestUsageAudit { diff --git a/apps/aether-gateway/src/handlers/shared/catalog.rs b/apps/aether-gateway/src/handlers/shared/catalog.rs index 736640b3e..9d1936ece 100644 --- a/apps/aether-gateway/src/handlers/shared/catalog.rs +++ b/apps/aether-gateway/src/handlers/shared/catalog.rs @@ -21,6 +21,7 @@ use aether_crypto::{ use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey; use aether_provider_pool::{ grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier, + provider_pool_codex_metadata_has_account_quota, }; use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at; use serde_json::{json, Map, Value}; @@ -1111,9 +1112,8 @@ fn build_codex_quota_status_snapshot( source: &str, ) -> Option { let metadata = provider_quota_metadata_bucket(upstream_metadata, "codex")?; - let observed_at_unix_secs = metadata - .get("updated_at") - .and_then(admin_provider_quota_pure::coerce_json_u64); + let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("observed_at")) + .or_else(|| provider_quota_timestamp_unix_secs(metadata.get("updated_at"))); let plan_type = metadata .get("plan_type") .and_then(Value::as_str) @@ -1388,6 +1388,168 @@ fn build_kiro_quota_status_snapshot( })) } +fn build_xai_quota_status_snapshot( + upstream_metadata: Option<&Value>, + source: &str, +) -> Option { + let metadata = provider_quota_metadata_bucket(upstream_metadata, "xai")?; + let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at")); + let usage_limit = metadata + .get("usage_limit") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let current_usage = metadata + .get("current_usage") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let remaining = metadata + .get("remaining") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let usage_ratio = metadata + .get("usage_percentage") + .and_then(admin_provider_quota_pure::coerce_json_f64) + .map(|value| (value / 100.0).clamp(0.0, 1.0)) + .or_else(|| { + current_usage + .zip(usage_limit) + .and_then(|(current_usage, usage_limit)| { + (usage_limit > 0.0).then_some((current_usage / usage_limit).clamp(0.0, 1.0)) + }) + }); + let remaining_ratio = usage_ratio.map(|value| (1.0 - value).max(0.0)); + let next_reset_at = provider_quota_timestamp_unix_secs(metadata.get("next_reset_at")); + let reset_seconds = quota_window_reset_seconds(observed_at_unix_secs, next_reset_at); + let plan_type = metadata + .get("subscription_title") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let period_type = metadata + .get("period_type") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let usage_label = match period_type.as_deref() { + Some("monthly") => "月额度", + Some("weekly") => "周额度", + _ => "额度", + }; + + let mut windows = Vec::new(); + if usage_ratio.is_some() + || remaining.is_some() + || usage_limit.is_some() + || current_usage.is_some() + || next_reset_at.is_some() + { + windows.push(json!({ + "code": "usage", + "label": usage_label, + "scope": "account", + "unit": if usage_limit.is_some() { "usd" } else { "percent" }, + "used_ratio": usage_ratio, + "remaining_ratio": remaining_ratio, + "used_value": current_usage, + "remaining_value": remaining, + "limit_value": usage_limit, + "reset_at": next_reset_at, + "reset_seconds": reset_seconds, + })); + } + + let prepaid_balance = metadata + .get("prepaid_balance") + .and_then(admin_provider_quota_pure::coerce_json_f64); + if prepaid_balance.is_some_and(|value| value > 0.0) { + windows.push(json!({ + "code": "prepaid", + "label": "预付额度", + "scope": "account", + "unit": "usd", + "used_ratio": serde_json::Value::Null, + "remaining_ratio": serde_json::Value::Null, + "remaining_value": prepaid_balance, + "reset_at": serde_json::Value::Null, + "reset_seconds": serde_json::Value::Null, + })); + } + + let on_demand_cap = metadata + .get("on_demand_cap") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let on_demand_used = metadata + .get("on_demand_used") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let on_demand_enabled = metadata + .get("on_demand_enabled") + .and_then(admin_provider_quota_pure::coerce_json_bool) + != Some(false); + if on_demand_enabled && on_demand_cap.is_some_and(|value| value > 0.0) { + let on_demand_remaining = on_demand_cap + .zip(on_demand_used) + .map(|(cap, used)| (cap - used).max(0.0)); + let on_demand_ratio = on_demand_cap + .zip(on_demand_used) + .and_then(|(cap, used)| (cap > 0.0).then_some((used / cap).clamp(0.0, 1.0))); + windows.push(json!({ + "code": "on_demand", + "label": "按需额度", + "scope": "account", + "unit": "usd", + "used_ratio": on_demand_ratio, + "remaining_ratio": on_demand_ratio.map(|value| (1.0 - value).max(0.0)), + "used_value": on_demand_used, + "remaining_value": on_demand_remaining, + "limit_value": on_demand_cap, + "reset_at": serde_json::Value::Null, + "reset_seconds": serde_json::Value::Null, + })); + } + + if windows.is_empty() && plan_type.is_none() && observed_at_unix_secs.is_none() { + return None; + } + + let prepaid_available = prepaid_balance.is_some_and(|value| value > 0.0); + let on_demand_available = on_demand_enabled + && on_demand_cap.is_some_and(|value| value > 0.0) + && on_demand_used + .zip(on_demand_cap) + .is_some_and(|(used, cap)| used < cap); + let usage_exhausted = remaining.is_some_and(|value| value <= 0.0) + || usage_ratio.is_some_and(|value| value >= 1.0 - 1e-6); + let exhausted = usage_exhausted && !prepaid_available && !on_demand_available; + let reason = if exhausted { + Some("额度已耗尽".to_string()) + } else { + None + }; + let label = if exhausted { + Some("额度耗尽") + } else { + None + }; + let code = if exhausted { "exhausted" } else { "ok" }; + + Some(json!({ + "version": 2, + "provider_type": "xai", + "code": code, + "label": label, + "reason": reason, + "freshness": "fresh", + "source": source, + "observed_at": observed_at_unix_secs, + "exhausted": exhausted, + "usage_ratio": usage_ratio, + "updated_at": observed_at_unix_secs, + "reset_at": next_reset_at, + "reset_seconds": reset_seconds, + "plan_type": plan_type, + "windows": windows, + })) +} + fn build_chatgpt_web_quota_status_snapshot( upstream_metadata: Option<&Value>, source: &str, @@ -2132,6 +2294,111 @@ fn build_gemini_cli_quota_status_snapshot( })) } +fn build_claude_code_quota_status_snapshot( + upstream_metadata: Option<&Value>, + source: &str, +) -> Option { + let metadata = provider_quota_metadata_bucket(upstream_metadata, "claude_code")?; + let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at")); + // (metadata prefix, window code, window minutes, account-wide?). Display labels are + // resolved by the frontend from `code` so they follow the UI locale. + let definitions: [(&str, &str, u64, bool); 4] = [ + ("five_hour", "5h", 300, true), + ("seven_day", "weekly", 10_080, true), + ("seven_day_sonnet", "weekly_sonnet", 10_080, false), + ("seven_day_fable", "weekly_fable", 10_080, false), + ]; + let mut windows = Vec::new(); + for (prefix, code, window_minutes, account_wide) in definitions { + let used_percent = metadata + .get(&format!("{prefix}_used_percent")) + .and_then(Value::as_f64); + let reset_at = + provider_quota_timestamp_unix_secs(metadata.get(&format!("{prefix}_reset_at"))); + // A window whose reset time already passed no longer describes current usage. + let expired = reset_at + .zip(observed_at_unix_secs) + .is_some_and(|(reset_at, observed_at)| reset_at <= observed_at); + let Some(used_percent) = used_percent else { + continue; + }; + let used_ratio = if expired { + 0.0 + } else { + (used_percent / 100.0).clamp(0.0, 1.0) + }; + let reset_seconds = reset_at + .zip(observed_at_unix_secs) + .map(|(reset_at, observed_at)| reset_at.saturating_sub(observed_at)); + let mut window = json!({ + "code": code, + "scope": if account_wide { "account" } else { "model" }, + "unit": "percent", + "used_ratio": used_ratio, + "remaining_ratio": 1.0 - used_ratio, + "reset_at": reset_at, + "reset_seconds": reset_seconds, + "window_minutes": window_minutes, + "is_exhausted": used_ratio >= 1.0 - 1e-6, + }); + if !account_wide { + window["quota_group"] = json!(code); + } + windows.push(window); + } + if windows.is_empty() { + return None; + } + + let account_windows = windows + .iter() + .filter(|window| window.get("scope").and_then(Value::as_str) == Some("account")) + .cloned() + .collect::>(); + let blocking_windows = account_windows + .iter() + .filter(|window| window.get("is_exhausted").and_then(Value::as_bool) == Some(true)) + .cloned() + .collect::>(); + let exhausted = !blocking_windows.is_empty(); + // The account is usable again only once every exhausted window resets. + let reset_at = if exhausted { + blocking_windows + .iter() + .filter_map(|window| provider_quota_timestamp_unix_secs(window.get("reset_at"))) + .max() + } else { + None + }; + let reset_seconds = if exhausted { + blocking_windows + .iter() + .filter_map(|window| window.get("reset_seconds").and_then(Value::as_u64)) + .max() + } else { + None + }; + + Some(json!({ + "version": 2, + "provider_type": "claude_code", + "code": if exhausted { "exhausted" } else { "ok" }, + "freshness": "fresh", + "source": source, + "observed_at": observed_at_unix_secs, + "exhausted": exhausted, + "usage_ratio": quota_windows_usage_ratio(&account_windows), + "updated_at": observed_at_unix_secs, + "reset_at": reset_at, + "reset_seconds": reset_seconds, + "reset_credits": build_codex_reset_credits_status_snapshot( + metadata, + observed_at_unix_secs, + ), + "windows": windows, + })) +} + fn build_codex_reset_credits_status_snapshot( metadata: &Map, observed_at_unix_secs: Option, @@ -2255,11 +2522,13 @@ pub(crate) fn sync_provider_key_quota_status_snapshot( let mut quota = match normalized_provider_type.as_str() { "codex" => build_codex_quota_status_snapshot(upstream_metadata, source), "kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source), + "xai" => build_xai_quota_status_snapshot(upstream_metadata, source), "chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source), "windsurf" => build_windsurf_quota_status_snapshot(upstream_metadata, source), "antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source), "grok" => build_grok_quota_status_snapshot(upstream_metadata, source), "gemini_cli" => build_gemini_cli_quota_status_snapshot(upstream_metadata, source), + "claude_code" => build_claude_code_quota_status_snapshot(upstream_metadata, source), _ => None, }?; if normalized_provider_type == "codex" { @@ -2339,17 +2608,19 @@ fn codex_upstream_metadata_is_at_least_as_fresh( let Some(metadata) = provider_quota_metadata_bucket(upstream_metadata, "codex") else { return false; }; - let Some(metadata_updated_at) = metadata - .get("updated_at") - .and_then(admin_provider_quota_pure::coerce_json_u64) + // Identity, reset-credit, and model-only updates do not replace the + // account's quota observation, even when their timestamp is newer. + if !provider_pool_codex_metadata_has_account_quota(metadata) { + return false; + } + let Some(metadata_updated_at) = provider_quota_timestamp_unix_secs(metadata.get("observed_at")) + .or_else(|| provider_quota_timestamp_unix_secs(metadata.get("updated_at"))) else { return false; }; let snapshot_updated_at = quota_snapshot.and_then(|quota| { - quota - .get("updated_at") - .or_else(|| quota.get("observed_at")) - .and_then(admin_provider_quota_pure::coerce_json_u64) + provider_quota_timestamp_unix_secs(quota.get("observed_at")) + .or_else(|| provider_quota_timestamp_unix_secs(quota.get("updated_at"))) }); snapshot_updated_at.is_none_or(|updated_at| metadata_updated_at >= updated_at) @@ -2448,6 +2719,19 @@ pub(crate) fn provider_key_status_snapshot_payload( let mut snapshot = provider_key_status_snapshot_object(Some(&payload)) .or_else(|| default_provider_key_status_snapshot().as_object().cloned()) .unwrap_or_default(); + // Legacy snapshots can retain an exhausted summary after a window reset or + // newer quota observation. Use the same decision as scheduling so the + // account list and its status filter do not keep displaying that stale block. + if provider_type.trim().eq_ignore_ascii_case("codex") + && !aether_provider_pool::provider_pool_key_account_quota_exhausted(key, provider_type) + { + if let Some(quota) = snapshot.get_mut("quota").and_then(Value::as_object_mut) { + quota.insert("exhausted".to_string(), json!(false)); + if quota.get("code").and_then(Value::as_str) == Some("exhausted") { + quota.insert("code".to_string(), json!("ok")); + } + } + } snapshot.insert( "oauth".to_string(), build_provider_key_oauth_status_snapshot(key), @@ -3559,6 +3843,59 @@ mod tests { assert_eq!(window.get("used_value"), Some(&json!(0.0))); } + #[test] + fn provider_key_status_snapshot_payload_backfills_claude_code_usage_windows() { + let mut key = sample_catalog_key(); + key.upstream_metadata = Some(json!({ + "claude_code": { + "updated_at": 1_800_000_000u64, + "five_hour_used_percent": 100.0, + "five_hour_reset_at": 1_800_003_600u64, + "seven_day_used_percent": 40.0, + "seven_day_reset_at": 1_800_400_000u64, + "seven_day_sonnet_used_percent": 10.0, + "seven_day_sonnet_reset_at": 1_800_400_000u64, + "reset_credits": { + "available_count": 2, + "updated_at": 1_800_000_000u64, + "detail_source": "claude_oauth_usage", + "credits": [{ + "display_key": "Key-1", + "status": "available", + "expires_at": 1_800_144_000u64 + }] + } + } + })); + + let payload = provider_key_status_snapshot_payload(&key, "claude_code"); + let quota = payload + .get("quota") + .and_then(Value::as_object) + .expect("quota snapshot should be object"); + assert_eq!(quota.get("provider_type"), Some(&json!("claude_code"))); + // An exhausted 5h window blocks the whole account until it resets. + assert_eq!(quota.get("exhausted"), Some(&json!(true))); + assert_eq!(quota.get("reset_at"), Some(&json!(1_800_003_600u64))); + let windows = quota + .get("windows") + .and_then(Value::as_array) + .expect("windows should exist"); + assert_eq!(windows.len(), 3); + assert_eq!(windows[0]["code"], json!("5h")); + assert_eq!(windows[0]["scope"], json!("account")); + assert_eq!(windows[0]["window_minutes"], json!(300)); + assert_eq!(windows[1]["code"], json!("weekly")); + assert_eq!(windows[1]["used_ratio"], json!(0.4)); + assert_eq!(windows[2]["code"], json!("weekly_sonnet")); + assert_eq!(windows[2]["scope"], json!("model")); + assert_eq!(quota["reset_credits"]["available_count"], json!(2)); + assert_eq!( + quota["reset_credits"]["credits"][0]["remaining_seconds"], + json!(144_000u64) + ); + } + #[test] fn provider_key_status_snapshot_payload_backfills_grok_model_quota() { let mut key = sample_catalog_key(); @@ -3622,6 +3959,43 @@ mod tests { assert_eq!(auto.get("used_value"), Some(&json!(90.0))); } + #[test] + fn provider_key_status_snapshot_payload_backfills_xai_weekly_credits() { + let mut key = sample_catalog_key(); + key.upstream_metadata = Some(json!({ + "xai": { + "updated_at": 1_778_067_246u64, + "usage_percentage": 46.0, + "period_type": "weekly", + "next_reset_at": 1_778_157_172u64, + "subscription_title": "SuperGrok", + "prepaid_balance": 0.0, + "on_demand_cap": 0.0, + "on_demand_used": 0.0 + } + })); + + let payload = provider_key_status_snapshot_payload(&key, "xai"); + let quota = payload + .get("quota") + .and_then(Value::as_object) + .expect("quota snapshot should be object"); + let windows = quota + .get("windows") + .and_then(Value::as_array) + .expect("xai quota windows should exist"); + + assert_eq!(quota.get("provider_type"), Some(&json!("xai"))); + assert_eq!(quota.get("code"), Some(&json!("ok"))); + assert_eq!(quota.get("exhausted"), Some(&json!(false))); + assert_eq!(quota.get("plan_type"), Some(&json!("SuperGrok"))); + assert_eq!(quota.get("usage_ratio"), Some(&json!(0.46))); + assert_eq!(quota.get("reset_at"), Some(&json!(1_778_157_172u64))); + assert_eq!(windows.len(), 1); + assert_eq!(windows[0].get("code"), Some(&json!("usage"))); + assert_eq!(windows[0].get("label"), Some(&json!("周额度"))); + } + #[test] fn provider_key_status_snapshot_payload_backfills_gemini_cli_account_credits() { let mut key = sample_catalog_key(); @@ -4013,6 +4387,116 @@ mod tests { ); } + #[test] + fn provider_key_status_snapshot_payload_clears_stale_codex_exhaustion_summary() { + let mut key = sample_catalog_key(); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "code": "exhausted", + "exhausted": true, + "windows": [{ + "code": "weekly", + "scope": "account", + "used_ratio": 0.83, + "remaining_ratio": 0.17, + "reset_at": 4_102_444_800u64 + }] + } + })); + let payload = provider_key_status_snapshot_payload(&key, "codex"); + assert_eq!(payload["quota"]["code"], "ok"); + assert_eq!(payload["quota"]["exhausted"], false); + assert!(payload["quota"]["label"].is_null()); + + // An explicit current upstream refusal is not a stale percentage summary. + key.status_snapshot.as_mut().unwrap()["quota"]["allowed"] = json!(false); + let payload = provider_key_status_snapshot_payload(&key, "codex"); + assert_eq!(payload["quota"]["code"], "exhausted"); + assert_eq!(payload["quota"]["exhausted"], true); + + // Missing capacity evidence must not clear an exhausted summary either. + key.status_snapshot = Some(json!({"quota": { + "provider_type": "codex", "code": "exhausted", "exhausted": true + }})); + let payload = provider_key_status_snapshot_payload(&key, "codex"); + assert_eq!(payload["quota"]["code"], "exhausted"); + assert_eq!(payload["quota"]["exhausted"], true); + } + + #[test] + fn provider_key_status_snapshot_payload_refreshes_codex_timestamp_formats() { + for updated_at in [ + json!(1_900_000_000u64), + json!(1_900_000_000_000u64), + json!("2030-03-17T17:46:40Z"), + ] { + let mut key = sample_catalog_key(); + key.upstream_metadata = Some(json!({ + "codex": { + "updated_at": updated_at, + "primary_used_percent": 83.0, + "primary_reset_at": 4_102_444_800u64 + } + })); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "updated_at": 1_899_999_000u64, + "code": "exhausted", + "exhausted": true, + "allowed": false, + "windows": [{ + "code": "weekly", + "scope": "account", + "used_ratio": 1.0, + "remaining_ratio": 0.0, + "reset_at": 4_102_444_800u64 + }] + } + })); + let payload = provider_key_status_snapshot_payload(&key, "codex"); + assert_eq!(payload["quota"]["code"], "ok"); + assert_eq!(payload["quota"]["updated_at"], 1_900_000_000u64); + assert_eq!(payload["quota"]["windows"][0]["used_ratio"], 0.83); + assert!(payload["quota"]["allowed"].is_null()); + } + } + + #[test] + fn provider_key_status_snapshot_payload_preserves_codex_account_quota_on_unrelated_updates() { + for patch in [ + json!({"plan_type": "pro"}), + json!({"spark_primary_used_percent": 83.0}), + json!({"credits_unlimited": false}), + json!({"windows": [{"code": "weekly", "reset_at": 4_102_444_800u64}]}), + ] { + let mut metadata = patch; + metadata["updated_at"] = json!("2030-03-17T17:46:40Z"); + let mut key = sample_catalog_key(); + key.upstream_metadata = Some(json!({"codex": metadata})); + key.status_snapshot = Some(json!({"quota": { + "provider_type": "codex", "updated_at": 200, + "code": "exhausted", "exhausted": true, "allowed": false, + "windows": [{"code": "weekly", "scope": "account", "used_ratio": 1.0, + "remaining_ratio": 0.0, "reset_at": 4_102_444_800u64}] + }})); + let payload = provider_key_status_snapshot_payload(&key, "codex"); + assert_eq!(payload["quota"]["code"], "exhausted", "{metadata}"); + assert_eq!(payload["quota"]["exhausted"], true, "{metadata}"); + assert_eq!(payload["quota"]["allowed"], false, "{metadata}"); + assert_eq!(payload["quota"]["updated_at"], 200, "{metadata}"); + assert_eq!( + payload["quota"]["windows"][0]["code"], "weekly", + "{metadata}" + ); + assert_eq!( + payload["quota"]["windows"][0]["used_ratio"], 1.0, + "{metadata}" + ); + } + } + #[test] fn provider_key_status_snapshot_payload_restores_complete_codex_cache() { let mut key = sample_catalog_key(); diff --git a/apps/aether-gateway/src/handlers/shared/request_utils.rs b/apps/aether-gateway/src/handlers/shared/request_utils.rs index 603116ca3..44879fcad 100644 --- a/apps/aether-gateway/src/handlers/shared/request_utils.rs +++ b/apps/aether-gateway/src/handlers/shared/request_utils.rs @@ -276,6 +276,10 @@ pub(crate) fn admin_proxy_local_requires_buffered_body( | (Some("system_manage"), http::Method::POST, Some("config_import")) | (Some("system_manage"), http::Method::POST, Some("users_import")) | (Some("system_manage"), http::Method::POST, Some("data_import")) + | (Some("system_manage"), http::Method::POST, Some("cleanup_usage_manual")) + | (Some("system_manage"), http::Method::POST, Some("smtp_test")) + | (Some("system_manage"), http::Method::POST, Some("prepare_update")) + | (Some("system_manage"), http::Method::POST, Some("apply_update")) | (Some("system_manage"), http::Method::PUT, Some("settings_set")) | (Some("system_manage"), http::Method::PUT, Some("config_set")) | (Some("system_manage"), http::Method::PUT, Some("email_template_set")) @@ -611,4 +615,27 @@ mod tests { "/v1/chat/completions?key=passthrough" ); } + + #[test] + fn manual_cleanup_route_requires_buffered_body() { + use crate::control::GatewayPublicRequestContext; + let mut decision = GatewayControlDecision::synthetic( + "/api/admin/system/cleanup/usage/manual", + Some("admin_proxy".to_string()), + Some("system_manage".to_string()), + Some("cleanup_usage_manual".to_string()), + Some("system_manage:cleanup_usage_manual".to_string()), + ); + decision.route_class = Some("admin_proxy".to_string()); + let uri: http::Uri = "/api/admin/system/cleanup/usage/manual".parse().unwrap(); + let headers = http::HeaderMap::new(); + let context = GatewayPublicRequestContext::from_request_parts( + "trace-manual-cleanup", + &http::Method::POST, + &uri, + &headers, + Some(decision), + ); + assert!(super::admin_proxy_local_requires_buffered_body(&context)); + } } diff --git a/apps/aether-gateway/src/image_capabilities.rs b/apps/aether-gateway/src/image_capabilities.rs index 21dc5f183..0c5c10c1d 100644 --- a/apps/aether-gateway/src/image_capabilities.rs +++ b/apps/aether-gateway/src/image_capabilities.rs @@ -13,7 +13,7 @@ pub(crate) fn openai_image_provider_max_generation_count(provider_type: &str) -> GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT } else if matches!( provider_type.trim().to_ascii_lowercase().as_str(), - "openai" | "codex" + "openai" | "codex" | "xai" ) { OPENAI_IMAGE_MAX_GENERATION_COUNT } else { @@ -58,6 +58,7 @@ mod tests { assert_eq!(openai_image_provider_max_generation_count("grok"), 4); assert_eq!(openai_image_provider_max_generation_count("openai"), 10); assert_eq!(openai_image_provider_max_generation_count("codex"), 10); + assert_eq!(openai_image_provider_max_generation_count("xai"), 10); assert_eq!(openai_image_provider_max_generation_count("custom"), 1); assert_eq!( openai_image_provider_max_generation_count_for_model("openai", Some("dall-e-3")), diff --git a/apps/aether-gateway/src/lib.rs b/apps/aether-gateway/src/lib.rs index 308493329..7d5676c75 100644 --- a/apps/aether-gateway/src/lib.rs +++ b/apps/aether-gateway/src/lib.rs @@ -37,6 +37,7 @@ mod bark_push; mod cache; mod client_session_affinity; mod clock; +mod codex_profile; mod constants; mod control; mod data; @@ -92,12 +93,12 @@ mod usage; mod video_tasks; mod wallet_runtime; +pub use self::ai_serving::api::{codex_client_originator, codex_client_user_agent}; pub(crate) use self::ai_serving::api::{ AiControlPlanRequest, EXECUTION_RUNTIME_STREAM_DECISION_ACTION, EXECUTION_RUNTIME_SYNC_DECISION_ACTION, GEMINI_FILES_DOWNLOAD_PLAN_KIND, OPENAI_VIDEO_CONTENT_PLAN_KIND, }; -pub use self::ai_serving::api::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT}; pub(crate) use self::ai_serving::{ AiExecutionDecision, AiExecutionPlanPayload, AiStreamAttempt, AiSyncAttempt, }; diff --git a/apps/aether-gateway/src/main.rs b/apps/aether-gateway/src/main.rs index d765f7a83..c130fb53b 100644 --- a/apps/aether-gateway/src/main.rs +++ b/apps/aether-gateway/src/main.rs @@ -2513,6 +2513,20 @@ async fn run() -> Result<(), Box> { ); } } + match state.prewarm_codex_client_profile().await { + Ok(version) => { + info!( + codex_client_version = %version, + "prewarmed Codex client profile" + ); + } + Err(err) => { + warn!( + error = %err, + "failed to refresh Codex client profile; built-in or cached profile remains active" + ); + } + } match prewarm_direct_h2c_sender_cache_from_env_for_startup().await { Ok(Some(report)) => { if report.failed_targets > 0 { @@ -4963,7 +4977,11 @@ mod tests { builder .http1() .timer(TokioTimer::new()) - .header_read_timeout(std::time::Duration::from_millis(10)) + // Keep hyper's own header timeout far from the 5ms first-request + // deadline: when a slow runner lets both expire before the next + // poll, `select!` may pick the connection branch and surface + // hyper's header-timeout error instead of the clean deadline close. + .header_read_timeout(std::time::Duration::from_secs(30)) .max_buf_size(super::MIN_GATEWAY_HTTP_HEADER_MAX_BYTES) .max_headers(super::MIN_GATEWAY_HTTP_MAX_HEADERS); builder diff --git a/apps/aether-gateway/src/model_fetch/catalog.rs b/apps/aether-gateway/src/model_fetch/catalog.rs index 6b0d92e13..65e274e6f 100644 --- a/apps/aether-gateway/src/model_fetch/catalog.rs +++ b/apps/aether-gateway/src/model_fetch/catalog.rs @@ -19,6 +19,8 @@ use sha2::{Digest, Sha256}; use tokio::sync::{Mutex, Semaphore}; use tracing::{debug, info, warn}; +use crate::ai_serving::api::codex_client_version; + const CODEX_CATALOG_SCHEMA_VERSION: u32 = 2; const CODEX_CATALOG_CREDENTIAL_SCOPE_DOMAIN: &str = "aether-codex-catalog-credential-v2"; const CODEX_CLIENT_VERSION_MAX_LEN: usize = 64; @@ -99,7 +101,7 @@ pub(crate) fn normalize_codex_client_version(raw: Option<&str>) -> NormalizedCod used_fallback: false, }, None => NormalizedCodexClientVersion { - value: crate::ai_serving::CODEX_CLIENT_VERSION.to_string(), + value: codex_client_version(), used_fallback: true, }, } @@ -1521,7 +1523,7 @@ where .await?; let scope = target.credential_scope()?; let state = runtime.codex_catalog_runtime_state(); - let mut version = Version::parse(crate::ai_serving::CODEX_CLIENT_VERSION).ok()?; + let mut version = Version::parse(&codex_client_version()).ok()?; if let Some(recent) = read_recent_codex_catalog_client_version(state, provider_id, key_id, scope).await { @@ -2461,10 +2463,7 @@ mod tests { let initial = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID) .await .expect("management context"); - assert_eq!( - initial.client_version, - crate::ai_serving::CODEX_CLIENT_VERSION - ); + assert_eq!(initial.client_version, codex_client_version()); assert!(initial.models.is_none()); seed_catalog(&runtime, &version("0.200.0")).await; @@ -2498,10 +2497,7 @@ mod tests { let rebound = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID) .await .unwrap(); - assert_eq!( - rebound.client_version, - crate::ai_serving::CODEX_CLIENT_VERSION - ); + assert_eq!(rebound.client_version, codex_client_version()); assert!(rebound.models.is_none()); } @@ -2557,7 +2553,7 @@ mod tests { format!("1.2.3-{}", "x".repeat(CODEX_CLIENT_VERSION_MAX_LEN)), ] { let normalized = normalize_codex_client_version(Some(&raw)); - assert_eq!(normalized.as_str(), crate::ai_serving::CODEX_CLIENT_VERSION); + assert_eq!(normalized.as_str(), codex_client_version()); assert!(normalized.used_fallback()); assert!(!catalog_lkg_key(&target(), normalized.as_str()).contains(&raw)); } @@ -3753,7 +3749,7 @@ mod tests { .await .expect("seed legacy cache"); - let load = load_one(&runtime, &version(crate::ai_serving::CODEX_CLIENT_VERSION)).await; + let load = load_one(&runtime, &version(&codex_client_version())).await; assert!(load.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID).is_none()); assert_eq!(runtime.execution_count(), 1); } diff --git a/apps/aether-gateway/src/orchestration/report_effects.rs b/apps/aether-gateway/src/orchestration/report_effects.rs index 7e28a9e78..b86414637 100644 --- a/apps/aether-gateway/src/orchestration/report_effects.rs +++ b/apps/aether-gateway/src/orchestration/report_effects.rs @@ -551,6 +551,24 @@ async fn sync_grok_quota_from_report_context( async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncReportRequest) { apply_local_gemini_file_mapping_report_effect(state, payload).await; + if claude_code_quota_headers_reportable(payload.status_code) { + if let Err(err) = sync_claude_code_quota_from_response_headers( + state, + payload.report_context.as_ref(), + &payload.headers, + ) + .await + { + warn!( + event_name = "claude_code_realtime_quota_sync_failed", + log_type = "ops", + report_kind = %payload.report_kind, + report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())), + error = ?err, + "gateway failed to persist claude_code realtime quota from sync response headers" + ); + } + } if (200..300).contains(&payload.status_code) { if let Err(err) = sync_codex_quota_from_response_headers( state, @@ -641,6 +659,24 @@ async fn apply_local_stream_report_effect(state: &AppState, payload: &GatewayStr ); } } + if claude_code_quota_headers_reportable(payload.status_code) { + if let Err(err) = sync_claude_code_quota_from_response_headers( + state, + payload.report_context.as_ref(), + &payload.headers, + ) + .await + { + warn!( + event_name = "claude_code_realtime_quota_sync_failed", + log_type = "ops", + report_kind = %payload.report_kind, + report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())), + error = ?err, + "gateway failed to persist claude_code realtime quota from stream response headers" + ); + } + } if let Err(err) = sync_grok_quota_from_report_context( state, payload.report_context.as_ref(), @@ -896,6 +932,134 @@ async fn sync_codex_quota_from_response_headers( .await } +fn claude_code_quota_headers_reportable(status_code: u16) -> bool { + // Real limit 429s carry the unified headers (fingerprint-rejection 429s do not, and then + // the parser finds nothing to record). + (200..300).contains(&status_code) || status_code == 429 +} + +/// Passive sampling of Anthropic's `anthropic-ratelimit-unified-*` response headers into the +/// `claude_code` quota metadata, so the 5H / weekly windows stay fresh between active refreshes. +async fn sync_claude_code_quota_from_response_headers( + state: &AppState, + report_context: Option<&Value>, + headers: &BTreeMap, +) -> Result { + let observed_at_unix_secs = report_context_u64( + report_context, + "provider_response_headers_observed_at_unix_ms", + ) + .map(|value| value / 1_000) + .filter(|value| *value > 0) + .unwrap_or_else(current_unix_secs); + let parsed = report_context_provider_response_headers(report_context) + .and_then(|headers| { + admin_provider_quota_pure::parse_claude_code_usage_headers( + &headers, + observed_at_unix_secs, + ) + }) + .or_else(|| { + admin_provider_quota_pure::parse_claude_code_usage_headers( + headers, + observed_at_unix_secs, + ) + }); + let Some(parsed) = parsed else { + return Ok(false); + }; + let Some(key_id) = report_context_key_id(report_context) else { + return Ok(false); + }; + + for attempt in 0..RUNTIME_METADATA_CAS_MAX_ATTEMPTS { + let Some(key) = state + .read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id)) + .await? + .into_iter() + .next() + else { + return Ok(false); + }; + let Some(provider) = state + .read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id)) + .await? + .into_iter() + .next() + else { + return Ok(false); + }; + if !provider + .provider_type + .trim() + .eq_ignore_ascii_case("claude_code") + { + return Ok(false); + } + + let expected_namespace_value = + upstream_metadata_namespace_value(key.upstream_metadata.as_ref(), "claude_code"); + let mut bucket = expected_namespace_value + .as_ref() + .and_then(Value::as_object) + .cloned() + .unwrap_or_default(); + // Never let an older observation overwrite a newer active refresh. + if bucket + .get("updated_at") + .and_then(admin_provider_quota_pure::coerce_json_u64) + .is_some_and(|stored| stored > observed_at_unix_secs) + { + return Ok(false); + } + let Some(patch) = parsed.as_object() else { + return Ok(false); + }; + // Headers can describe only some windows; absent windows keep their stored value. + for (field, value) in patch { + bucket.insert(field.clone(), value.clone()); + } + let next_bucket = Value::Object(bucket); + if expected_namespace_value.as_ref() == Some(&next_bucket) { + return Ok(false); + } + + let updated_upstream_metadata = merge_metadata_object( + key.upstream_metadata.as_ref(), + "claude_code", + next_bucket.clone(), + ); + let updated_status_snapshot = sync_provider_key_quota_status_snapshot( + key.status_snapshot.as_ref(), + provider.provider_type.as_str(), + updated_upstream_metadata.as_ref(), + "response_headers", + ); + let updated = state + .update_provider_catalog_key_runtime_metadata( + &ProviderCatalogKeyRuntimeMetadataUpdate { + key_id: key_id.clone(), + namespace: "claude_code".to_string(), + expected_upstream_metadata_value: expected_namespace_value, + upstream_metadata_value: next_bucket, + status_snapshot_patch: quota_status_snapshot_patch( + updated_status_snapshot.as_ref(), + ), + updated_at_unix_secs: Some(observed_at_unix_secs), + }, + ) + .await?; + if updated { + return Ok(true); + } + if attempt + 1 < RUNTIME_METADATA_CAS_MAX_ATTEMPTS { + let backoff_us = 50_u64.saturating_mul((attempt + 1) as u64).min(1_000); + tokio::time::sleep(Duration::from_micros(backoff_us)).await; + } + } + Ok(false) +} + async fn sync_codex_websocket_quota_from_stream_payload( state: &AppState, payload: &GatewayStreamReportRequest, diff --git a/apps/aether-gateway/src/provider_key_auth.rs b/apps/aether-gateway/src/provider_key_auth.rs index 62379e62f..1d8021e7c 100644 --- a/apps/aether-gateway/src/provider_key_auth.rs +++ b/apps/aether-gateway/src/provider_key_auth.rs @@ -172,6 +172,7 @@ fn provider_uses_bearer_oauth_runtime(provider_type: &str) -> bool { | "antigravity" | "kiro" | "windsurf" + | "xai" ) } @@ -406,6 +407,22 @@ mod tests { ); } + #[test] + fn recognizes_xai_oauth_as_bearer_runtime() { + let semantics = provider_key_auth_semantics(&sample_key("oauth"), "xai"); + + assert!(semantics.oauth_managed()); + assert!(semantics.can_refresh_oauth()); + assert_eq!( + semantics.credential_kind(), + ProviderKeyCredentialKind::OAuthSession + ); + assert_eq!( + semantics.runtime_auth_kind(), + ProviderKeyRuntimeAuthKind::Bearer + ); + } + #[test] fn refresh_capability_requires_stored_refresh_token() { let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex"); diff --git a/apps/aether-gateway/src/router.rs b/apps/aether-gateway/src/router.rs index e7d48f818..157c46653 100644 --- a/apps/aether-gateway/src/router.rs +++ b/apps/aether-gateway/src/router.rs @@ -186,6 +186,8 @@ fn frontend_path_bypasses_static(path: &str) -> bool { "/health" | "/test-connection" | crate::constants::READYZ_PATH ) || path.starts_with("/api/") || path.starts_with("/v1/") + || path == "/openai/v1/videos" + || path.starts_with("/openai/v1/videos/") || path.starts_with("/v1beta/") || path.starts_with("/upload/") || path.starts_with("/_gateway/") diff --git a/apps/aether-gateway/src/scheduler/candidate/resolution.rs b/apps/aether-gateway/src/scheduler/candidate/resolution.rs index fe68b6a1b..70d3d4305 100644 --- a/apps/aether-gateway/src/scheduler/candidate/resolution.rs +++ b/apps/aether-gateway/src/scheduler/candidate/resolution.rs @@ -23,6 +23,16 @@ pub(super) fn resolve_scheduler_candidate_selectability( if let Some(skip_reason) = current_candidate_runtime_skip_reason(&candidate, runtime_snapshot, now_unix_secs) { + tracing::debug!( + event_name = "scheduler_candidate_skipped", + log_type = "event", + provider_id = %candidate.provider_id, + endpoint_id = %candidate.endpoint_id, + key_id = %candidate.key_id, + api_format = %candidate.endpoint_api_format, + skip_reason, + "scheduler candidate skipped during runtime selectability resolution" + ); if emitted_skipped_keys.insert(key) { skipped.push(SchedulerSkippedCandidate { candidate, diff --git a/apps/aether-gateway/src/scheduler/candidate/runtime.rs b/apps/aether-gateway/src/scheduler/candidate/runtime.rs index 51a2003f3..c6837019f 100644 --- a/apps/aether-gateway/src/scheduler/candidate/runtime.rs +++ b/apps/aether-gateway/src/scheduler/candidate/runtime.rs @@ -40,10 +40,6 @@ pub(super) async fn read_candidate_runtime_selection_snapshot( ) -> Result { let provider_concurrent_limits = read_provider_concurrent_limits(state, candidates).await?; let provider_pool_state = read_provider_pool_state_map(state, candidates).await?; - let provider_skip_exhausted_accounts = provider_pool_state - .iter() - .map(|(provider_id, state)| (provider_id.clone(), state.skip_exhausted_accounts)) - .collect::>(); let pool_provider_ids = provider_pool_state .iter() .filter_map(|(provider_id, state)| state.pool_enabled.then_some(provider_id.clone())) @@ -62,7 +58,7 @@ pub(super) async fn read_candidate_runtime_selection_snapshot( let key_account_quota_exhausted = read_key_account_quota_exhaustion_map( candidates, &provider_key_rpm_states, - &provider_skip_exhausted_accounts, + &provider_pool_state, ); let key_oauth_invalid = read_key_oauth_invalid_map(candidates, &provider_key_rpm_states, now_unix_secs); @@ -360,6 +356,7 @@ async fn read_provider_quota_block_map( struct ProviderPoolState { pool_enabled: bool, skip_exhausted_accounts: bool, + reserve_minimum_quota: bool, } async fn read_provider_pool_state_map( @@ -391,11 +388,17 @@ async fn read_provider_pool_state_map( .and_then(|value| value.get("skip_exhausted_accounts")) .and_then(serde_json::Value::as_bool) .unwrap_or(false); + let reserve_minimum_quota = pool_advanced + .and_then(serde_json::Value::as_object) + .and_then(|value| value.get("reserve_minimum_quota")) + .and_then(serde_json::Value::as_bool) + .unwrap_or(false); ( provider.id, ProviderPoolState { pool_enabled: pool_advanced.is_some(), skip_exhausted_accounts, + reserve_minimum_quota, }, ) }) @@ -405,7 +408,7 @@ async fn read_provider_pool_state_map( fn read_key_account_quota_exhaustion_map( candidates: &[SchedulerMinimalCandidateSelectionCandidate], provider_key_rpm_states: &BTreeMap, - provider_skip_exhausted_accounts: &BTreeMap, + provider_pool_state: &BTreeMap, ) -> BTreeMap { candidates .iter() @@ -431,11 +434,19 @@ fn read_key_account_quota_exhaustion_map( candidate.provider_type.as_str(), candidate.selected_provider_model_name.as_str(), ); - let skip_configured = provider_skip_exhausted_accounts + let pool_state = provider_pool_state .get(candidate.provider_id.as_str()) .copied() - .unwrap_or(false); - hard_blocked || (skip_configured && account_exhausted) + .unwrap_or_default(); + let reserve_reached = pool_state.reserve_minimum_quota + && admin_provider_pool_pure::admin_pool_key_minimum_quota_reached( + key, + candidate.provider_type.as_str(), + Some(candidate.selected_provider_model_name.as_str()), + ); + hard_blocked + || reserve_reached + || (pool_state.skip_exhausted_accounts && account_exhausted) }); (candidate.key_id.clone(), exhausted) }) @@ -566,3 +577,64 @@ fn read_provider_key_rpm_reset_at_map( }) .collect::>() } + +#[cfg(test)] +mod reserve_minimum_quota_tests { + use super::*; + use serde_json::json; + + #[test] + fn reserve_minimum_quota_is_independent_of_skip_exhausted_accounts() { + let candidate = SchedulerMinimalCandidateSelectionCandidate { + provider_id: "provider-codex".to_string(), + provider_name: "codex".to_string(), + provider_type: "codex".to_string(), + provider_priority: 0, + endpoint_id: "endpoint-codex".to_string(), + endpoint_api_format: "openai:responses".to_string(), + key_id: "key-codex".to_string(), + key_name: "codex".to_string(), + key_auth_type: "oauth".to_string(), + key_internal_priority: 0, + key_global_priority_for_format: None, + key_capabilities: None, + model_id: "model-codex".to_string(), + global_model_id: "global-model-codex".to_string(), + global_model_name: "gpt-5".to_string(), + selected_provider_model_name: "gpt-5".to_string(), + supports_streaming: true, + mapping_matched_model: None, + }; + let mut key = StoredProviderCatalogKey::new( + candidate.key_id.clone(), + candidate.provider_id.clone(), + "codex".to_string(), + "oauth".to_string(), + None, + true, + ) + .expect("key should build"); + for reserve_enabled in [false, true] { + for used_percent in [99.0, 98.0] { + key.upstream_metadata = + Some(json!({"codex": {"primary_used_percent": used_percent}})); + let exhausted = read_key_account_quota_exhaustion_map( + std::slice::from_ref(&candidate), + &BTreeMap::from([(key.id.clone(), key.clone())]), + &BTreeMap::from([( + candidate.provider_id.clone(), + ProviderPoolState { + pool_enabled: true, + reserve_minimum_quota: reserve_enabled, + skip_exhausted_accounts: false, + }, + )]), + ); + assert_eq!( + exhausted.get(&key.id), + Some(&(reserve_enabled && used_percent >= 99.0)) + ); + } + } + } +} diff --git a/apps/aether-gateway/src/state/app.rs b/apps/aether-gateway/src/state/app.rs index 6813bbf81..56757cb06 100644 --- a/apps/aether-gateway/src/state/app.rs +++ b/apps/aether-gateway/src/state/app.rs @@ -495,6 +495,25 @@ pub struct AppState { Arc>>, >, #[cfg(test)] + pub(crate) auth_wallet_adjustment_error_for_tests: Option, + #[cfg(test)] + pub(crate) auth_wallet_lookup_error_for_tests: Option, + #[cfg(test)] + pub(crate) auth_wallet_batch_store_for_tests: Option< + Arc< + StdMutex< + HashMap< + (String, String), + aether_data::repository::wallet::StoredAdminUserWalletBalanceBatch, + >, + >, + >, + >, + #[cfg(test)] + pub(crate) auth_wallet_batch_operation_lock_for_tests: Arc>, + #[cfg(test)] + pub(crate) auth_wallet_batch_failure_record_error_for_tests: Option, + #[cfg(test)] pub(crate) admin_wallet_payment_order_store: Option>>>, #[cfg(test)] diff --git a/apps/aether-gateway/src/state/catalog.rs b/apps/aether-gateway/src/state/catalog.rs index 6349545b0..baa8fb6ec 100644 --- a/apps/aether-gateway/src/state/catalog.rs +++ b/apps/aether-gateway/src/state/catalog.rs @@ -804,12 +804,43 @@ impl AppState { if updated.is_some() { self.invalidate_provider_routing_caches(); } + if let Some(key) = updated.as_ref().filter(|key| !key.is_active) { + self.delete_inactive_provider_catalog_key_pool_scores( + key.provider_id.as_str(), + key.id.as_str(), + ) + .await; + } match updated { Some(key) => self.open_provider_catalog_key(key).await.map(Some), None => Ok(None), } } + async fn delete_inactive_provider_catalog_key_pool_scores( + &self, + provider_id: &str, + key_id: &str, + ) { + if let Err(err) = self + .data + .delete_pool_member_scores_for_member( + &pool_scores::PoolMemberIdentity::provider_api_key( + provider_id.to_string(), + key_id.to_string(), + ), + ) + .await + { + warn!( + provider_id, + key_id, + error = ?err, + "gateway provider catalog key deactivate: failed to delete pool member scores" + ); + } + } + pub(crate) async fn compare_and_update_provider_catalog_key_admin_state( &self, update: &provider_catalog::ProviderCatalogKeyAdminCasUpdate, @@ -843,6 +874,15 @@ impl AppState { if updated.as_ref().is_some_and(|keys| !keys.is_empty()) { self.invalidate_provider_routing_caches(); } + if let Some(keys) = updated.as_ref() { + for key in keys.iter().filter(|key| !key.is_active) { + self.delete_inactive_provider_catalog_key_pool_scores( + key.provider_id.as_str(), + key.id.as_str(), + ) + .await; + } + } match updated { Some(keys) => self.open_provider_catalog_keys(keys).await.map(Some), None => Ok(None), diff --git a/apps/aether-gateway/src/state/core.rs b/apps/aether-gateway/src/state/core.rs index 3438d515c..495de7a10 100644 --- a/apps/aether-gateway/src/state/core.rs +++ b/apps/aether-gateway/src/state/core.rs @@ -54,6 +54,7 @@ use super::super::router::RequestAdmissionError; use super::super::{control::GatewayControlDecision, error::GatewayError}; use super::super::{provider_transport, usage}; +use crate::codex_profile::spawn_worker as spawn_codex_client_profile_worker; use crate::maintenance::spawn_account_self_check_worker; use crate::maintenance::spawn_audit_cleanup_worker; use crate::maintenance::spawn_db_maintenance_worker; @@ -149,6 +150,10 @@ fn system_config_key_affects_provider_transport_snapshot(key: &str) -> bool { } impl AppState { + pub async fn prewarm_codex_client_profile(&self) -> Result { + crate::codex_profile::prewarm(self.runtime_state()).await + } + pub async fn prewarm_chat_pii_redaction_runtime_config(&self) -> Result { crate::privacy::read_chat_pii_redaction_runtime_config(self) .await @@ -471,6 +476,16 @@ impl AppState { #[cfg(test)] auth_wallet_store: Some(Arc::new(StdMutex::new(HashMap::new()))), #[cfg(test)] + auth_wallet_adjustment_error_for_tests: None, + #[cfg(test)] + auth_wallet_lookup_error_for_tests: None, + #[cfg(test)] + auth_wallet_batch_store_for_tests: Some(Arc::new(StdMutex::new(HashMap::new()))), + #[cfg(test)] + auth_wallet_batch_operation_lock_for_tests: Arc::new(TokioMutex::new(())), + #[cfg(test)] + auth_wallet_batch_failure_record_error_for_tests: None, + #[cfg(test)] admin_wallet_payment_order_store: Some(Arc::new(StdMutex::new(HashMap::new()))), #[cfg(test)] admin_payment_callback_store: Some(Arc::new(StdMutex::new(HashMap::new()))), @@ -2337,6 +2352,10 @@ impl AppState { crate::task_runtime::TASK_KEY_MODEL_FETCH_WORKER, spawn_model_fetch_worker(background_state.clone()), ); + supervise_worker( + crate::task_runtime::TASK_KEY_CODEX_CLIENT_PROFILE, + Some(spawn_codex_client_profile_worker(background_state.clone())), + ); supervise_worker( crate::task_runtime::TASK_KEY_VIDEO_TASK_POLLER, spawn_video_task_poller(background_state.clone()), diff --git a/apps/aether-gateway/src/state/integrations.rs b/apps/aether-gateway/src/state/integrations.rs index db94e5786..caf97843a 100644 --- a/apps/aether-gateway/src/state/integrations.rs +++ b/apps/aether-gateway/src/state/integrations.rs @@ -290,6 +290,14 @@ impl provider_transport::VideoTaskTransportSnapshotLookup for AppState { .await .map_err(GatewayError::into_message) } + + async fn resolve_video_task_proxy( + &self, + transport: &GatewayProviderTransportSnapshot, + ) -> Option { + self.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport) + .await + } } #[async_trait] diff --git a/apps/aether-gateway/src/state/runtime/wallet/balance_mutations.rs b/apps/aether-gateway/src/state/runtime/wallet/balance_mutations.rs index 8501b155a..c6b6b1b39 100644 --- a/apps/aether-gateway/src/state/runtime/wallet/balance_mutations.rs +++ b/apps/aether-gateway/src/state/runtime/wallet/balance_mutations.rs @@ -1,8 +1,198 @@ use crate::{AdminWalletPaymentOrderRecord, AdminWalletTransactionRecord, AppState, GatewayError}; +use aether_data::repository::wallet::{ + AdjustWalletBalanceInBatchInput, AdminUserWalletBalanceBatchUserOutcome, + PrepareAdminUserWalletBalanceBatchInput, PrepareAdminUserWalletBalanceBatchOutcome, + StoredAdminUserWalletBalanceBatch, +}; +use std::collections::BTreeMap; use super::admin_wallet_build_order_no; impl AppState { + pub(crate) async fn prepare_admin_user_wallet_balance_batch( + &self, + input: PrepareAdminUserWalletBalanceBatchInput, + ) -> Result { + #[cfg(test)] + if let Some(store) = self.auth_wallet_batch_store_for_tests.as_ref() { + let mut batches = store.lock().expect("auth wallet batch store should lock"); + let key = (input.admin_user_id.clone(), input.idempotency_key.clone()); + if let Some(existing) = batches.get(&key) { + if existing.request_fingerprint != input.request_fingerprint { + return Ok(PrepareAdminUserWalletBalanceBatchOutcome::Conflict); + } + return Ok(PrepareAdminUserWalletBalanceBatchOutcome::Ready( + existing.clone(), + )); + } + let batch = StoredAdminUserWalletBalanceBatch { + admin_user_id: input.admin_user_id, + idempotency_key: input.idempotency_key, + request_fingerprint: input.request_fingerprint, + target_user_ids: input.target_user_ids, + missing_user_ids: input.missing_user_ids, + warnings: input.warnings, + user_outcomes: BTreeMap::new(), + }; + batches.insert(key, batch.clone()); + return Ok(PrepareAdminUserWalletBalanceBatchOutcome::Ready(batch)); + } + + self.data + .prepare_admin_user_wallet_balance_batch(input) + .await + .map_err(|error| GatewayError::Internal(error.to_string()))? + .ok_or_else(|| { + GatewayError::Internal("admin wallet batch storage is unavailable".to_string()) + }) + } + + pub(crate) async fn get_admin_user_wallet_balance_batch( + &self, + admin_user_id: &str, + idempotency_key: &str, + request_fingerprint: &str, + ) -> Result, GatewayError> { + #[cfg(test)] + if let Some(store) = self.auth_wallet_batch_store_for_tests.as_ref() { + let batches = store.lock().expect("auth wallet batch store should lock"); + return Ok(batches + .get(&(admin_user_id.to_string(), idempotency_key.to_string())) + .map(|existing| { + if existing.request_fingerprint != request_fingerprint { + PrepareAdminUserWalletBalanceBatchOutcome::Conflict + } else { + PrepareAdminUserWalletBalanceBatchOutcome::Ready(existing.clone()) + } + })); + } + + self.data + .get_admin_user_wallet_balance_batch( + admin_user_id, + idempotency_key, + request_fingerprint, + ) + .await + .map_err(|error| GatewayError::Internal(error.to_string())) + } + + pub(crate) async fn adjust_admin_user_wallet_balance_batch_user( + &self, + input: AdjustWalletBalanceInBatchInput, + ) -> Result { + #[cfg(test)] + if let Some(store) = self.auth_wallet_batch_store_for_tests.as_ref() { + let _operation = self.auth_wallet_batch_operation_lock_for_tests.lock().await; + let key = (input.admin_user_id.clone(), input.idempotency_key.clone()); + if let Some(existing) = store + .lock() + .expect("auth wallet batch store should lock") + .get(&key) + .and_then(|batch| batch.user_outcomes.get(&input.user_id)) + .cloned() + { + return Ok(existing); + } + let target_exists = store + .lock() + .expect("auth wallet batch store should lock") + .get(&key) + .is_some_and(|batch| batch.target_user_ids.contains(&input.user_id)); + if !target_exists { + return Err(GatewayError::Internal( + "user is outside the prepared admin wallet batch".to_string(), + )); + } + let adjustment = input.adjustment; + let result = self + .admin_adjust_wallet_balance( + &adjustment.wallet_id, + adjustment.amount_usd, + &adjustment.balance_type, + adjustment.operator_id.as_deref(), + adjustment.description.as_deref(), + adjustment.clamp_deduction_to_available_balance, + ) + .await?; + let outcome = if result.is_some() { + AdminUserWalletBalanceBatchUserOutcome::Succeeded + } else { + AdminUserWalletBalanceBatchUserOutcome::Failed("用户钱包不可用".to_string()) + }; + if let Some(batch) = store + .lock() + .expect("auth wallet batch store should lock") + .get_mut(&key) + { + batch.user_outcomes.insert(input.user_id, outcome.clone()); + } + return Ok(outcome); + } + + self.data + .adjust_admin_user_wallet_balance_batch_user(input) + .await + .map_err(|error| GatewayError::Internal(error.to_string()))? + .ok_or_else(|| { + GatewayError::Internal("admin wallet batch storage is unavailable".to_string()) + }) + } + + pub(crate) async fn record_admin_user_wallet_balance_batch_failure( + &self, + admin_user_id: &str, + idempotency_key: &str, + user_id: &str, + reason: &str, + ) -> Result { + #[cfg(test)] + if self + .auth_wallet_batch_failure_record_error_for_tests + .as_deref() + == Some(user_id) + { + return Err(GatewayError::Internal( + "injected wallet batch failure-record error".to_string(), + )); + } + + #[cfg(test)] + if let Some(store) = self.auth_wallet_batch_store_for_tests.as_ref() { + let _operation = self.auth_wallet_batch_operation_lock_for_tests.lock().await; + let key = (admin_user_id.to_string(), idempotency_key.to_string()); + let mut batches = store.lock().expect("auth wallet batch store should lock"); + let batch = batches.get_mut(&key).ok_or_else(|| { + GatewayError::Internal("admin wallet batch was not prepared".to_string()) + })?; + if !batch.target_user_ids.iter().any(|target| target == user_id) { + return Err(GatewayError::Internal( + "user is outside the prepared admin wallet batch".to_string(), + )); + } + return Ok(batch + .user_outcomes + .entry(user_id.to_string()) + .or_insert_with(|| { + AdminUserWalletBalanceBatchUserOutcome::Failed(reason.to_string()) + }) + .clone()); + } + + self.data + .record_admin_user_wallet_balance_batch_failure( + admin_user_id, + idempotency_key, + user_id, + reason, + ) + .await + .map_err(|error| GatewayError::Internal(error.to_string()))? + .ok_or_else(|| { + GatewayError::Internal("admin wallet batch storage is unavailable".to_string()) + }) + } + pub(crate) async fn admin_adjust_wallet_balance( &self, wallet_id: &str, @@ -10,13 +200,21 @@ impl AppState { balance_type: &str, operator_id: Option<&str>, description: Option<&str>, + clamp_deduction_to_available_balance: bool, ) -> Result< Option<( aether_data::repository::wallet::StoredWalletSnapshot, - AdminWalletTransactionRecord, + Option, )>, GatewayError, > { + #[cfg(test)] + if self.auth_wallet_adjustment_error_for_tests.as_deref() == Some(wallet_id) { + return Err(GatewayError::Internal( + "injected test wallet adjustment failure".to_string(), + )); + } + #[cfg(test)] if let Some(store) = self.auth_wallet_store.as_ref() { let mut guard = store.lock().expect("auth wallet store should lock"); @@ -27,6 +225,18 @@ impl AppState { let before_recharge = wallet.balance; let before_gift = wallet.gift_balance; let before_total = before_recharge + before_gift; + let amount_usd = if clamp_deduction_to_available_balance && amount_usd < 0.0 { + if before_total < 0.0 { + -before_total + } else { + -(-amount_usd).min(before_total) + } + } else { + amount_usd + }; + if amount_usd == 0.0 { + return Ok(Some((wallet.clone(), None))); + } let mut after_recharge = before_recharge; let mut after_gift = before_gift; @@ -90,7 +300,7 @@ impl AppState { let updated_wallet = wallet.clone(); drop(guard); self.invalidate_auth_context_cache(); - return Ok(Some((updated_wallet, transaction))); + return Ok(Some((updated_wallet, Some(transaction)))); } Ok(self @@ -100,10 +310,15 @@ impl AppState { balance_type: balance_type.to_string(), operator_id: operator_id.map(ToOwned::to_owned), description: description.map(ToOwned::to_owned), + clamp_deduction_to_available_balance, + batch_context: None, }) .await? .map(|(wallet, transaction)| { - (wallet, stored_wallet_transaction_to_gateway(transaction)) + ( + wallet, + transaction.map(stored_wallet_transaction_to_gateway), + ) })) } diff --git a/apps/aether-gateway/src/state/runtime/wallet/mutations.rs b/apps/aether-gateway/src/state/runtime/wallet/mutations.rs index 995ba626f..40b4960a8 100644 --- a/apps/aether-gateway/src/state/runtime/wallet/mutations.rs +++ b/apps/aether-gateway/src/state/runtime/wallet/mutations.rs @@ -112,7 +112,7 @@ impl AppState { ) -> Result< Option<( aether_data::repository::wallet::StoredWalletSnapshot, - aether_data::repository::wallet::StoredAdminWalletTransaction, + Option, )>, GatewayError, > { diff --git a/apps/aether-gateway/src/state/runtime/wallet/reads.rs b/apps/aether-gateway/src/state/runtime/wallet/reads.rs index 90896a419..a0bf2d187 100644 --- a/apps/aether-gateway/src/state/runtime/wallet/reads.rs +++ b/apps/aether-gateway/src/state/runtime/wallet/reads.rs @@ -5,6 +5,19 @@ impl AppState { &self, lookup: aether_data::repository::wallet::WalletLookupKey<'_>, ) -> Result, GatewayError> { + #[cfg(test)] + if let Some(failed_user_id) = self.auth_wallet_lookup_error_for_tests.as_deref() { + let lookup_user_id = match &lookup { + aether_data::repository::wallet::WalletLookupKey::UserId(user_id) => Some(*user_id), + _ => None, + }; + if lookup_user_id == Some(failed_user_id) { + return Err(GatewayError::Internal( + "injected test wallet lookup failure".to_string(), + )); + } + } + #[cfg(test)] if let Some(store) = self.auth_wallet_store.as_ref() { let wallet = { diff --git a/apps/aether-gateway/src/state/testing.rs b/apps/aether-gateway/src/state/testing.rs index bdbf662e9..ebc27b469 100644 --- a/apps/aether-gateway/src/state/testing.rs +++ b/apps/aether-gateway/src/state/testing.rs @@ -472,6 +472,27 @@ impl AppState { self } + pub(crate) fn fail_auth_wallet_adjustment_for_tests( + mut self, + wallet_id: impl Into, + ) -> Self { + self.auth_wallet_adjustment_error_for_tests = Some(wallet_id.into()); + self + } + + pub(crate) fn fail_auth_wallet_lookup_for_tests(mut self, user_id: impl Into) -> Self { + self.auth_wallet_lookup_error_for_tests = Some(user_id.into()); + self + } + + pub(crate) fn fail_auth_wallet_batch_failure_record_for_tests( + mut self, + user_id: impl Into, + ) -> Self { + self.auth_wallet_batch_failure_record_error_for_tests = Some(user_id.into()); + self + } + pub(crate) fn with_admin_wallet_payment_orders_for_tests(mut self, orders: I) -> Self where I: IntoIterator, diff --git a/apps/aether-gateway/src/task_runtime/mod.rs b/apps/aether-gateway/src/task_runtime/mod.rs index e8d4c95e9..226579c97 100644 --- a/apps/aether-gateway/src/task_runtime/mod.rs +++ b/apps/aether-gateway/src/task_runtime/mod.rs @@ -24,6 +24,7 @@ pub(crate) const TASK_KEY_USAGE_QUEUE_WORKER: &str = "usage.queue.worker"; pub(crate) const TASK_KEY_USAGE_COUNTER_FLUSH: &str = "usage.counter.flush.worker"; pub(crate) const TASK_KEY_VIDEO_TASK_POLLER: &str = "video.task.poller"; pub(crate) const TASK_KEY_MODEL_FETCH_WORKER: &str = "model.fetch.worker"; +pub(crate) const TASK_KEY_CODEX_CLIENT_PROFILE: &str = "maintenance.codex.client.profile"; pub(crate) const TASK_KEY_PROVIDER_QUOTA_RESET: &str = "provider.quota.reset.worker"; pub(crate) const TASK_KEY_ACCOUNT_SELF_CHECK: &str = "account.self_check.worker"; pub(crate) const TASK_KEY_POOL_SCORE_REBUILD: &str = "pool.score.rebuild.worker"; @@ -202,6 +203,14 @@ const TASK_DEFINITIONS: &[TaskDefinition] = &[ true, RETRY_ONCE, ), + TaskDefinition::new( + TASK_KEY_CODEX_CLIENT_PROFILE, + TaskKind::Scheduled, + "daily", + true, + true, + RETRY_ONCE, + ), TaskDefinition::new( TASK_KEY_PROVIDER_QUOTA_RESET, TaskKind::Scheduled, diff --git a/apps/aether-gateway/src/tests/ai_execute/finalize_local.rs b/apps/aether-gateway/src/tests/ai_execute/finalize_local.rs index 3bac90c5d..dd8c0350a 100644 --- a/apps/aether-gateway/src/tests/ai_execute/finalize_local.rs +++ b/apps/aether-gateway/src/tests/ai_execute/finalize_local.rs @@ -1972,7 +1972,7 @@ async fn gateway_executes_openai_chat_antigravity_cross_format_sync_via_local_fi seen_execution_runtime_request.user_agent, aether_provider_transport::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT ); - assert_eq!(seen_execution_runtime_request.request_type, "agent"); + assert_eq!(seen_execution_runtime_request.request_type, ""); assert_eq!(seen_execution_runtime_request.contents_len, 1); assert!(!seen_execution_runtime_request.request_has_model); diff --git a/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/cross_format.rs b/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/cross_format.rs index 0cedf3111..795829d64 100644 --- a/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/cross_format.rs +++ b/apps/aether-gateway/src/tests/ai_execute/finalize_local_cli/cross_format.rs @@ -893,8 +893,8 @@ async fn gateway_executes_openai_responses_cross_format_function_call_upstream_s }, { "type": "function_call", - "id": "call_auto_1", - "call_id": "call_auto_1", + "id": "call_auto_0", + "call_id": "call_auto_0", "name": "get_weather", "arguments": "{\"location\":\"Tokyo\"}" } @@ -1570,7 +1570,7 @@ async fn gateway_executes_openai_responses_antigravity_cross_format_upstream_str seen_remote_execution_runtime_request.user_agent, aether_provider_transport::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT ); - assert_eq!(seen_remote_execution_runtime_request.request_type, "agent"); + assert_eq!(seen_remote_execution_runtime_request.request_type, ""); assert_eq!(seen_remote_execution_runtime_request.contents_len, 1); assert!(!seen_remote_execution_runtime_request.request_has_model); diff --git a/apps/aether-gateway/src/tests/ai_execute/finalize_local_provider/gemini.rs b/apps/aether-gateway/src/tests/ai_execute/finalize_local_provider/gemini.rs index 7c3575c2f..6a84c0e24 100644 --- a/apps/aether-gateway/src/tests/ai_execute/finalize_local_provider/gemini.rs +++ b/apps/aether-gateway/src/tests/ai_execute/finalize_local_provider/gemini.rs @@ -2155,7 +2155,7 @@ async fn gateway_executes_antigravity_gemini_cli_sync_upstream_stream_via_local_ seen_remote_execution_runtime_request.user_agent, aether_provider_transport::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT ); - assert_eq!(seen_remote_execution_runtime_request.request_type, "agent"); + assert_eq!(seen_remote_execution_runtime_request.request_type, ""); assert_eq!(seen_remote_execution_runtime_request.contents_len, 0); assert!((seen_remote_execution_runtime_request.exact_temperature - 0.2).abs() < f64::EPSILON); assert!(!seen_remote_execution_runtime_request.request_has_model); diff --git a/apps/aether-gateway/src/tests/ai_execute/stream/image.rs b/apps/aether-gateway/src/tests/ai_execute/stream/image.rs index 1b76e7cff..477bed8d4 100644 --- a/apps/aether-gateway/src/tests/ai_execute/stream/image.rs +++ b/apps/aether-gateway/src/tests/ai_execute/stream/image.rs @@ -421,7 +421,7 @@ async fn gateway_executes_codex_image_stream_via_local_decision_gate_after_oauth ); assert_eq!( seen_execution_runtime_request.headers["user-agent"], - aether_ai_formats::CODEX_CLIENT_USER_AGENT + aether_ai_formats::codex_client_user_agent() ); assert_eq!( seen_execution_runtime_request.headers["originator"], diff --git a/apps/aether-gateway/src/tests/ai_execute/stream_cli/direct.rs b/apps/aether-gateway/src/tests/ai_execute/stream_cli/direct.rs index 8ad2bb958..f4224aced 100644 --- a/apps/aether-gateway/src/tests/ai_execute/stream_cli/direct.rs +++ b/apps/aether-gateway/src/tests/ai_execute/stream_cli/direct.rs @@ -38,6 +38,7 @@ async fn gateway_executes_codex_cli_stream_via_local_decision_gate_after_oauth_r trace_id: String, url: String, model: String, + service_tier: String, content_encoding: String, stream: bool, accept: String, @@ -383,6 +384,13 @@ async fn gateway_executes_codex_cli_stream_via_local_decision_gate_after_oauth_r .and_then(|value| value.as_str()) .unwrap_or_default() .to_string(), + service_tier: payload + .get("body") + .and_then(|value| value.get("json_body")) + .and_then(|value| value.get("service_tier")) + .and_then(|value| value.as_str()) + .unwrap_or_default() + .to_string(), content_encoding: payload .get("content_encoding") .and_then(|value| value.as_str()) @@ -605,7 +613,7 @@ async fn gateway_executes_codex_cli_stream_via_local_decision_gate_after_oauth_r ) .header(TRACE_ID_HEADER, "trace-codex-cli-stream-local-123") .body( - r#"{"model":"gpt-5.6-sol","instructions":"Use the configured tools.","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"compact"}]},{"type":"compaction_trigger"}],"tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],"context_management":[{"type":"compaction","compact_threshold":128000}],"parallel_tool_calls":true,"prompt_cache_key":"thread-codex-stream-local-123","client_metadata":{"session_id":"session-codex-stream-local-123","thread_id":"thread-codex-stream-local-123"},"stream":true}"#, + r#"{"model":"gpt-5.6-sol","service_tier":"ultrafast","instructions":"Use the configured tools.","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"compact"}]},{"type":"compaction_trigger"}],"tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],"context_management":[{"type":"compaction","compact_threshold":128000}],"parallel_tool_calls":true,"prompt_cache_key":"thread-codex-stream-local-123","client_metadata":{"session_id":"session-codex-stream-local-123","thread_id":"thread-codex-stream-local-123"},"stream":true}"#, ) .send() .await @@ -716,6 +724,7 @@ async fn gateway_executes_codex_cli_stream_via_local_decision_gate_after_oauth_r "https://chatgpt.com/backend-api/codex/responses" ); assert_eq!(seen_execution_runtime_request.model, "gpt-5.6-sol"); + assert_eq!(seen_execution_runtime_request.service_tier, "ultrafast"); assert_eq!(seen_execution_runtime_request.content_encoding, "zstd"); assert!(seen_execution_runtime_request.stream); assert_eq!(seen_execution_runtime_request.accept, "text/event-stream"); diff --git a/apps/aether-gateway/src/tests/ai_execute/stream_provider.rs b/apps/aether-gateway/src/tests/ai_execute/stream_provider.rs index 3f5963bc1..27c72990b 100644 --- a/apps/aether-gateway/src/tests/ai_execute/stream_provider.rs +++ b/apps/aether-gateway/src/tests/ai_execute/stream_provider.rs @@ -1585,7 +1585,7 @@ async fn gateway_executes_claude_code_cli_stream_via_local_decision_gate_with_lo ); assert_eq!( seen_execution_runtime_request.user_agent, - "claude-cli/2.1.161 (external, cli)" + "claude-cli/2.1.284 (external, cli)" ); assert_eq!( seen_execution_runtime_request.endpoint_tag, diff --git a/apps/aether-gateway/src/tests/ai_execute/stream_provider_gemini/local_cli.rs b/apps/aether-gateway/src/tests/ai_execute/stream_provider_gemini/local_cli.rs index 5c5d26cd0..603ba0d64 100644 --- a/apps/aether-gateway/src/tests/ai_execute/stream_provider_gemini/local_cli.rs +++ b/apps/aether-gateway/src/tests/ai_execute/stream_provider_gemini/local_cli.rs @@ -2060,7 +2060,7 @@ async fn gateway_executes_antigravity_gemini_cli_stream_via_local_decision_gate_ seen_execution_runtime_request.user_agent, aether_provider_transport::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT ); - assert_eq!(seen_execution_runtime_request.request_type, "agent"); + assert_eq!(seen_execution_runtime_request.request_type, ""); assert_eq!(seen_execution_runtime_request.contents_len, 0); assert!((seen_execution_runtime_request.exact_temperature - 0.2).abs() < f64::EPSILON); assert!(!seen_execution_runtime_request.request_has_model); diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/chat/local_decision.rs b/apps/aether-gateway/src/tests/ai_execute/sync/chat/local_decision.rs index 9d97ebd83..71465d9f3 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/chat/local_decision.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/chat/local_decision.rs @@ -2028,14 +2028,30 @@ async fn gateway_returns_openai_chat_error_for_local_cross_format_claude_cli_syn "claude-code-upstream" ); assert_eq!(seen_execution_runtime_request.body["max_tokens"], 64); + // claude_code providers get the Claude Code body shape: billing header + identity + + // generic prompt, with the client's own system moved into the message history. + let system = seen_execution_runtime_request.body["system"] + .as_array() + .expect("claude_code system should be rewritten into blocks"); + assert_eq!(system.len(), 3); + assert!(system[0]["text"] + .as_str() + .is_some_and(|text| text.starts_with("x-anthropic-billing-header: cc_version="))); assert_eq!( - seen_execution_runtime_request.body["system"], - "You are terse." + system[1]["text"], + "You are Claude Code, Anthropic's official CLI for Claude." ); assert_eq!( seen_execution_runtime_request.body["messages"], - json!([{"role":"user","content":"Say hello"}]) + json!([ + {"role":"user","content":[{"type":"text","text":"[System Instructions]\nYou are terse."}]}, + {"role":"assistant","content":[{"type":"text","text":"Understood. I will follow these instructions."}]}, + {"role":"user","content":"Say hello"} + ]) ); + assert!(seen_execution_runtime_request.body["metadata"]["user_id"] + .as_str() + .is_some_and(|user_id| user_id.contains("\"session_id\""))); assert_eq!( seen_execution_runtime_request.client_api_format, "openai:chat" diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/claude/claude_code.rs b/apps/aether-gateway/src/tests/ai_execute/sync/claude/claude_code.rs index 4665f0b51..cf62688b4 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/claude/claude_code.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/claude/claude_code.rs @@ -567,11 +567,11 @@ async fn gateway_executes_claude_code_cli_sync_via_local_decision_gate_with_loca assert_eq!(seen_execution_runtime_request.x_stainless_helper_method, ""); assert_eq!( seen_execution_runtime_request.x_stainless_package_version, - "0.94.0" + "0.112.1" ); assert_eq!( seen_execution_runtime_request.user_agent, - "claude-cli/2.1.161 (external, cli)" + "claude-cli/2.1.284 (external, cli)" ); assert_eq!( seen_execution_runtime_request.endpoint_tag, diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/cli.rs b/apps/aether-gateway/src/tests/ai_execute/sync/cli.rs index 3fce6a6c7..a8e5c8acd 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/cli.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/cli.rs @@ -2473,6 +2473,416 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_sy upstream_handle.abort(); } +#[test] +fn gateway_returns_openai_responses_to_claude_code_applies_body_mimicry() { + run_cli_sync_test( + "gateway_returns_openai_responses_to_claude_code_applies_body_mimicry", + gateway_returns_openai_responses_to_claude_code_applies_body_mimicry_impl, + ); +} + +async fn gateway_returns_openai_responses_to_claude_code_applies_body_mimicry_impl() { + #[derive(Debug, Clone)] + struct SeenExecutionRuntimeSyncRequest { + trace_id: String, + url: String, + authorization: String, + endpoint_tag: String, + body: serde_json::Value, + } + + fn hash_api_key(value: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) + } + + fn sample_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot { + StoredAuthApiKeySnapshot::new( + user_id.to_string(), + "alice".to_string(), + Some("alice@example.com".to_string()), + "user".to_string(), + "local".to_string(), + true, + false, + Some(serde_json::json!(["openai", "claude", "claude_code"])), + Some(serde_json::json!(["openai:responses"])), + Some(serde_json::json!(["gpt-5"])), + api_key_id.to_string(), + Some("default".to_string()), + true, + false, + false, + Some(60), + Some(5), + Some(4_102_444_800_i64), + Some(serde_json::json!(["openai", "claude", "claude_code"])), + Some(serde_json::json!(["openai:responses"])), + Some(serde_json::json!(["gpt-5"])), + ) + .expect("auth snapshot should build") + } + + fn sample_candidate_row() -> StoredMinimalCandidateSelectionRow { + StoredMinimalCandidateSelectionRow { + provider_id: "provider-openai-cli-claude-code-local-1".to_string(), + provider_name: "claude_code".to_string(), + provider_type: "claude_code".to_string(), + provider_priority: 10, + provider_is_active: true, + endpoint_id: "endpoint-openai-cli-claude-code-local-1".to_string(), + endpoint_api_format: "claude:messages".to_string(), + endpoint_api_family: Some("claude".to_string()), + endpoint_kind: Some("cli".to_string()), + endpoint_is_active: true, + key_id: "key-openai-cli-claude-code-local-1".to_string(), + key_name: "prod".to_string(), + key_auth_type: "oauth".to_string(), + key_is_active: true, + key_api_formats: Some(vec!["claude:messages".to_string()]), + key_allowed_models: None, + key_capabilities: None, + key_internal_priority: 5, + key_global_priority_by_format: Some(serde_json::json!({"claude:messages": 1})), + model_id: "model-openai-cli-claude-code-local-1".to_string(), + global_model_id: "global-model-openai-cli-claude-code-local-1".to_string(), + global_model_name: "gpt-5".to_string(), + global_model_mappings: None, + global_model_supports_streaming: Some(true), + model_provider_model_name: "claude-code-upstream".to_string(), + model_provider_model_mappings: Some(vec![StoredProviderModelMapping { + name: "claude-code-upstream".to_string(), + priority: 1, + api_formats: Some(vec!["claude:messages".to_string()]), + endpoint_ids: None, + operations: None, + }]), + model_supports_streaming: Some(true), + model_is_active: true, + model_is_available: true, + } + } + + fn sample_provider_catalog_provider() -> StoredProviderCatalogProvider { + StoredProviderCatalogProvider::new( + "provider-openai-cli-claude-code-local-1".to_string(), + "claude_code".to_string(), + Some("https://example.com".to_string()), + "claude_code".to_string(), + ) + .expect("provider should build") + .with_transport_fields( + true, + false, + true, + None, + Some(2), + None, + Some(20.0), + None, + Some(serde_json::json!({ + "claude_code_advanced": {"cli_only_enabled": false}, + "failover_rules": { + "stop_on_status_codes": [429] + } + })), + ) + } + + fn sample_provider_catalog_endpoint() -> StoredProviderCatalogEndpoint { + StoredProviderCatalogEndpoint::new( + "endpoint-openai-cli-claude-code-local-1".to_string(), + "provider-openai-cli-claude-code-local-1".to_string(), + "claude:messages".to_string(), + Some("claude".to_string()), + Some("cli".to_string()), + true, + ) + .expect("endpoint should build") + .with_transport_fields( + "https://api.anthropic.com/v1".to_string(), + Some(serde_json::json!([ + {"action":"set","key":"x-endpoint-tag","value":"openai-cli-claude-code-cross-format"} + ])), + None, + Some(2), + None, + None, + None, + None, + ) + .expect("endpoint transport should build") + } + + fn sample_provider_catalog_key() -> StoredProviderCatalogKey { + StoredProviderCatalogKey::new( + "key-openai-cli-claude-code-local-1".to_string(), + "provider-openai-cli-claude-code-local-1".to_string(), + "prod".to_string(), + "oauth".to_string(), + None, + true, + ) + .expect("key should build") + .with_transport_fields( + Some(serde_json::json!(["claude:messages"])), + encrypt_python_fernet_plaintext( + DEVELOPMENT_ENCRYPTION_KEY, + "sk-upstream-openai-cli-claude-code", + ) + .expect("api key should encrypt"), + None, + None, + Some(serde_json::json!({"claude:messages": 1})), + None, + None, + None, + None, + ) + .expect("key transport should build") + } + + let seen_execution_runtime = Arc::new(Mutex::new(None::)); + let seen_execution_runtime_clone = Arc::clone(&seen_execution_runtime); + let seen_report = Arc::new(Mutex::new(false)); + let seen_report_clone = Arc::clone(&seen_report); + let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default()); + + let upstream = Router::new() + .route( + "/api/internal/gateway/resolve", + any(|_request: Request| async move { + Json(json!({ + "action": "proxy_public", + "route_class": "ai_public", + "route_family": "openai", + "route_kind": "cli", + "auth_endpoint_signature": "openai:responses", + "execution_runtime_candidate": true, + "auth_context": { + "user_id": "user-openai-cli-claude-code-local-error-123", + "api_key_id": "key-openai-cli-claude-code-local-error-123", + "access_allowed": true + }, + "public_path": "/v1/responses" + })) + }), + ) + .route( + "/api/internal/gateway/report-sync", + any(move |request: Request| { + let seen_report_inner = Arc::clone(&seen_report_clone); + async move { + let (_parts, body) = request.into_parts(); + let _raw_body = to_bytes(body, usize::MAX).await.expect("body should read"); + *seen_report_inner.lock().expect("mutex should lock") = true; + Json(json!({"ok": true})) + } + }), + ); + + let execution_runtime = Router::new().route( + "/v1/execute/sync", + any(move |request: Request| { + let seen_execution_runtime_inner = Arc::clone(&seen_execution_runtime_clone); + async move { + let (parts, body) = request.into_parts(); + let raw_body = to_bytes(body, usize::MAX).await.expect("body should read"); + let payload: serde_json::Value = serde_json::from_slice(&raw_body) + .expect("execution runtime payload should parse"); + *seen_execution_runtime_inner + .lock() + .expect("mutex should lock") = Some(SeenExecutionRuntimeSyncRequest { + trace_id: parts + .headers + .get(TRACE_ID_HEADER) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(), + url: payload + .get("url") + .and_then(|value| value.as_str()) + .unwrap_or_default() + .to_string(), + authorization: payload + .get("headers") + .and_then(|value| value.get("authorization")) + .and_then(|value| value.as_str()) + .unwrap_or_default() + .to_string(), + endpoint_tag: payload + .get("headers") + .and_then(|value| value.get("x-endpoint-tag")) + .and_then(|value| value.as_str()) + .unwrap_or_default() + .to_string(), + body: payload + .get("body") + .and_then(|value| value.get("json_body")) + .cloned() + .unwrap_or(serde_json::Value::Null), + }); + Json(json!({ + "request_id": "trace-openai-cli-claude-code-local-error-123", + "status_code": 429, + "headers": { + "content-type": "application/json" + }, + "body": { + "json_body": { + "type": "error", + "error": { + "type": "rate_limit_error", + "message": "slow down" + } + } + }, + "telemetry": { + "elapsed_ms": 28 + } + })) + } + }), + ); + + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![( + Some(hash_api_key("sk-client-openai-cli-claude-code-error")), + sample_auth_snapshot( + "key-openai-cli-claude-code-local-error-123", + "user-openai-cli-claude-code-local-error-123", + ), + )])); + let candidate_selection_repository = + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![ + sample_candidate_row(), + ])); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![sample_provider_catalog_provider()], + vec![sample_provider_catalog_endpoint()], + vec![sample_provider_catalog_key()], + )); + + let (_upstream_url, upstream_handle) = start_server(upstream).await; + let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; + let gateway_state = + build_state_with_execution_runtime_override(execution_runtime_url.clone()) + .with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests( + auth_repository, + candidate_selection_repository, + provider_catalog_repository, + Arc::clone(&request_candidate_repository), + DEVELOPMENT_ENCRYPTION_KEY, + ) + .attach_proxy_node_repository_for_tests( + crate::tests::ai_execute::ai_execute_proxy_node_repository([ + "proxy-node-openai-cli-local", + ]), + ), + ); + let gateway = build_router_with_state(gateway_state); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/v1/responses")) + .header(http::header::CONTENT_TYPE, "application/json") + .header( + http::header::AUTHORIZATION, + "Bearer sk-client-openai-cli-claude-code-error", + ) + .header(TRACE_ID_HEADER, "trace-openai-cli-claude-code-local-error-123") + .body( + "{\"model\":\"gpt-5\",\"instructions\":\"You are terse.\",\"input\":\"hello\",\"max_output_tokens\":64,\"store\":false}", + ) + .send() + .await + .expect("request should succeed"); + + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + assert_eq!( + response + .headers() + .get(EXECUTION_PATH_HEADER) + .and_then(|value| value.to_str().ok()), + Some(EXECUTION_PATH_EXECUTION_RUNTIME_SYNC) + ); + let response_json: serde_json::Value = response.json().await.expect("body should parse"); + assert_eq!( + response_json, + json!({ + "error": { + "message": "slow down", + "type": "rate_limit_error" + } + }) + ); + + let seen_execution_runtime_request = seen_execution_runtime + .lock() + .expect("mutex should lock") + .clone() + .expect("execution runtime sync should be captured"); + assert_eq!( + seen_execution_runtime_request.trace_id, + "trace-openai-cli-claude-code-local-error-123" + ); + assert_eq!( + seen_execution_runtime_request.url, + "https://api.anthropic.com/v1/messages" + ); + assert_eq!( + seen_execution_runtime_request.authorization, + "Bearer sk-upstream-openai-cli-claude-code" + ); + assert_eq!( + seen_execution_runtime_request.endpoint_tag, + "openai-cli-claude-code-cross-format" + ); + let body = &seen_execution_runtime_request.body; + assert_eq!(body["model"], "claude-code-upstream"); + // claude_code providers get the Claude Code body shape regardless of client format. + let system = body["system"] + .as_array() + .expect("claude_code system should be rewritten into blocks"); + assert_eq!(system.len(), 3); + assert!(system[0]["text"] + .as_str() + .is_some_and(|text| text.starts_with("x-anthropic-billing-header: cc_version="))); + assert_eq!( + system[1]["text"], + "You are Claude Code, Anthropic's official CLI for Claude." + ); + let messages = body["messages"].as_array().expect("messages should exist"); + assert_eq!(messages.len(), 3); + assert_eq!(messages[0]["role"], "user"); + assert!(messages[0].to_string().contains("[System Instructions]")); + assert!(messages[0].to_string().contains("You are terse.")); + assert_eq!(messages[1]["role"], "assistant"); + assert_eq!(messages[2]["role"], "user"); + assert!(messages[2].to_string().contains("hello")); + assert!(body["metadata"]["user_id"] + .as_str() + .is_some_and(|user_id| user_id.contains("\"session_id\""))); + + let stored_candidates = request_candidate_repository + .list_by_request_id("trace-openai-cli-claude-code-local-error-123") + .await + .expect("request candidate trace should read"); + assert_eq!(stored_candidates.len(), 1); + assert_eq!(stored_candidates[0].status, RequestCandidateStatus::Failed); + + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + assert!( + !*seen_report.lock().expect("mutex should lock"), + "report-sync should stay local when request candidate persistence is available" + ); + + gateway_handle.abort(); + execution_runtime_handle.abort(); + upstream_handle.abort(); +} + #[test] fn gateway_returns_openai_responses_error_for_local_cross_format_claude_chat_sync_failure() { run_cli_sync_test( diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/gemini/cli.rs b/apps/aether-gateway/src/tests/ai_execute/sync/gemini/cli.rs index 8abd3a094..bf491eef1 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/gemini/cli.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/gemini/cli.rs @@ -2374,7 +2374,7 @@ async fn gateway_executes_antigravity_gemini_cli_sync_via_local_decision_gate_af seen_execution_runtime_request.user_agent, aether_provider_transport::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT ); - assert_eq!(seen_execution_runtime_request.request_type, "agent"); + assert_eq!(seen_execution_runtime_request.request_type, ""); assert_eq!(seen_execution_runtime_request.contents_len, 0); assert!((seen_execution_runtime_request.exact_temperature - 0.2).abs() < f64::EPSILON); assert!(!seen_execution_runtime_request.request_has_model); diff --git a/apps/aether-gateway/src/tests/ai_execute/sync/image.rs b/apps/aether-gateway/src/tests/ai_execute/sync/image.rs index 921e33eea..58e201f67 100644 --- a/apps/aether-gateway/src/tests/ai_execute/sync/image.rs +++ b/apps/aether-gateway/src/tests/ai_execute/sync/image.rs @@ -1119,7 +1119,7 @@ async fn gateway_executes_codex_image_sync_via_local_decision_gate_after_oauth_r ); assert_eq!( seen_execution_runtime_request.headers["user-agent"], - aether_ai_formats::CODEX_CLIENT_USER_AGENT + aether_ai_formats::codex_client_user_agent() ); assert_eq!( seen_execution_runtime_request.headers["originator"], diff --git a/apps/aether-gateway/src/tests/control/admin.rs b/apps/aether-gateway/src/tests/control/admin.rs index 442d422e7..ab54ff946 100644 --- a/apps/aether-gateway/src/tests/control/admin.rs +++ b/apps/aether-gateway/src/tests/control/admin.rs @@ -21,5 +21,6 @@ mod system; mod system_import; mod usage; mod users; +mod users_batch; mod video_tasks; mod wallets; diff --git a/apps/aether-gateway/src/tests/control/admin/endpoints/keys.rs b/apps/aether-gateway/src/tests/control/admin/endpoints/keys.rs index 1043eb702..7b9afb6dd 100644 --- a/apps/aether-gateway/src/tests/control/admin/endpoints/keys.rs +++ b/apps/aether-gateway/src/tests/control/admin/endpoints/keys.rs @@ -44,21 +44,11 @@ where F: FnOnce() -> Fut + Send + 'static, Fut: std::future::Future + 'static, { - let handle = std::thread::Builder::new() - .name(test_name.to_string()) - .stack_size(PROVIDER_KEYS_TEST_STACK_BYTES) - .spawn(move || { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("test runtime should build"); - runtime.block_on(make_future()); - }) - .expect("provider keys test thread should spawn"); - - if let Err(payload) = handle.join() { - std::panic::resume_unwind(payload); - } + crate::tests::run_async_test_on_large_stack( + test_name, + PROVIDER_KEYS_TEST_STACK_BYTES, + make_future, + ); } struct SummaryNullingProviderCatalogReadRepository { diff --git a/apps/aether-gateway/src/tests/control/admin/endpoints/quota.rs b/apps/aether-gateway/src/tests/control/admin/endpoints/quota.rs index 37beb9b9a..7b1c4372c 100644 --- a/apps/aether-gateway/src/tests/control/admin/endpoints/quota.rs +++ b/apps/aether-gateway/src/tests/control/admin/endpoints/quota.rs @@ -59,21 +59,11 @@ where F: FnOnce() -> Fut + Send + 'static, Fut: std::future::Future + 'static, { - let handle = std::thread::Builder::new() - .name(test_name.to_string()) - .stack_size(PROVIDER_QUOTA_TEST_STACK_BYTES) - .spawn(move || { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("test runtime should build"); - runtime.block_on(make_future()); - }) - .expect("provider quota test thread should spawn"); - - if let Err(payload) = handle.join() { - std::panic::resume_unwind(payload); - } + crate::tests::run_async_test_on_large_stack( + test_name, + PROVIDER_QUOTA_TEST_STACK_BYTES, + make_future, + ); } #[tokio::test] @@ -2051,24 +2041,16 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_gemini_cli_with_trus #[tokio::test] async fn gateway_refresh_quota_reconciles_unsupported_fixed_provider_endpoints_before_clear_message( ) { - let cases = [ - ( - "provider-claude-code-reconcile", - "claude_code", - 1usize, - "claude:messages", - "https://api.anthropic.com/v1", - "Claude Code 暂不支持自动刷新额度", - ), - ( - "provider-vertex-ai-reconcile", - "vertex_ai", - 2usize, - "gemini:generate_content", - "https://aiplatform.googleapis.com", - "Vertex AI 暂不支持自动刷新额度", - ), - ]; + // Claude Code supports quota refresh now, so Vertex AI is the remaining fixed provider + // whose refresh is unsupported. + let cases = [( + "provider-vertex-ai-reconcile", + "vertex_ai", + 2usize, + "gemini:generate_content", + "https://aiplatform.googleapis.com", + "Vertex AI 暂不支持自动刷新额度", + )]; let providers = cases .iter() diff --git a/apps/aether-gateway/src/tests/control/admin/oauth.rs b/apps/aether-gateway/src/tests/control/admin/oauth.rs index d75debf5b..5d307c626 100644 --- a/apps/aether-gateway/src/tests/control/admin/oauth.rs +++ b/apps/aether-gateway/src/tests/control/admin/oauth.rs @@ -56,21 +56,11 @@ where F: FnOnce() -> Fut + Send + 'static, Fut: std::future::Future + 'static, { - let handle = std::thread::Builder::new() - .name(test_name.to_string()) - .stack_size(ADMIN_OAUTH_TEST_STACK_BYTES) - .spawn(move || { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("test runtime should build"); - runtime.block_on(make_future()); - }) - .expect("admin oauth test thread should spawn"); - - if let Err(payload) = handle.join() { - std::panic::resume_unwind(payload); - } + crate::tests::run_async_test_on_large_stack( + test_name, + ADMIN_OAUTH_TEST_STACK_BYTES, + make_future, + ); } fn decrypt_persisted_provider_api_key(key: &StoredProviderCatalogKey) -> String { @@ -420,6 +410,10 @@ async fn gateway_authorizes_claude_cookie_without_persisting_cookie_impl() { "email_address": "claude@example.com" } }), + // Newly authorized accounts get their 5H/weekly quota fetched right away. + quota if quota.starts_with("claude-code-quota:") => { + json!({"five_hour": {"utilization": 10.0}}) + } unexpected => panic!("unexpected execution plan: {unexpected}"), }; Json(json!({ @@ -821,7 +815,15 @@ async fn gateway_batch_authorizes_claude_cookies_as_redacted_task_impl() { } let plans = execution_plans.lock().expect("mutex should lock"); - assert_eq!(plans.len(), 6); + // Post-authorization quota refreshes are fire-and-forget, so their count is not + // deterministic here; only the OAuth flow plans are asserted. + assert_eq!( + plans + .iter() + .filter(|plan| !plan.request_id.starts_with("claude-code-quota:")) + .count(), + 6 + ); assert_eq!( plans .iter() @@ -983,6 +985,297 @@ async fn gateway_rejects_generic_oauth_start_for_windsurf_provider_impl() { ); } +#[test] +fn gateway_rejects_generic_oauth_start_for_xai_provider() { + run_admin_oauth_test( + "gateway_rejects_generic_oauth_start_for_xai_provider", + gateway_rejects_generic_oauth_start_for_xai_provider_impl, + ); +} + +async fn gateway_rejects_generic_oauth_start_for_xai_provider_impl() { + let mut provider = sample_provider("provider-xai", "xai", 10); + provider.provider_type = "xai".to_string(); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![], + vec![], + )); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests( + provider_catalog_repository, + )); + + let response = local_admin_provider_oauth_response( + &state, + http::Method::POST, + "/api/admin/provider-oauth/providers/provider-xai/start", + None, + ) + .await; + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let body = to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"); + let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); + assert!( + payload["detail"].as_str().is_some_and(|detail| { + detail.contains("设备授权") || detail.contains("导入凭据") + }), + "payload={payload}" + ); +} + +#[test] +fn gateway_handles_admin_provider_oauth_device_authorize_for_xai() { + run_admin_oauth_test( + "gateway_handles_admin_provider_oauth_device_authorize_for_xai", + gateway_handles_admin_provider_oauth_device_authorize_for_xai_impl, + ); +} + +async fn gateway_handles_admin_provider_oauth_device_authorize_for_xai_impl() { + let authorize_hits = Arc::new(Mutex::new(0usize)); + let authorize_hits_clone = Arc::clone(&authorize_hits); + let oidc_server = Router::new().fallback(any(move |_request: Request| { + let authorize_hits_inner = Arc::clone(&authorize_hits_clone); + async move { + *authorize_hits_inner.lock().expect("mutex should lock") += 1; + Json(json!({ + "device_code": "xai-device-code", + "user_code": "XAI-CODE", + "verification_uri": "https://auth.x.ai/activate", + "verification_uri_complete": "https://auth.x.ai/activate?user_code=XAI-CODE", + "expires_in": 600, + "interval": 5, + })) + } + })); + + let mut provider = sample_provider("provider-xai", "xai", 10); + provider.provider_type = "xai".to_string(); + let endpoint = sample_endpoint( + "endpoint-xai-responses", + "provider-xai", + "openai:responses", + "https://cli-chat-proxy.grok.com/v1", + ); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + vec![], + )); + let (oidc_url, oidc_handle) = start_server(oidc_server).await; + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests( + provider_catalog_repository, + )) + .with_provider_oauth_token_url_for_tests( + "xai_device", + format!("{oidc_url}/oauth2/device/code"), + ) + .with_provider_oauth_token_url_for_tests("xai", format!("{oidc_url}/oauth2/token")); + + let response = local_admin_provider_oauth_response( + &state, + http::Method::POST, + "/api/admin/provider-oauth/providers/provider-xai/device-authorize", + Some(json!({ "proxy_node_id": "proxy-node-xai" })), + ) + .await; + let status = response.status(); + let body = to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"); + let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); + assert_eq!(status, StatusCode::OK, "payload={payload}"); + let session_id = payload["session_id"] + .as_str() + .expect("session_id should exist") + .to_string(); + assert_eq!(payload["user_code"], "XAI-CODE"); + assert_eq!(payload["verification_uri"], "https://auth.x.ai/activate"); + assert_eq!( + payload["verification_uri_complete"], + "https://auth.x.ai/activate?user_code=XAI-CODE" + ); + assert_eq!(payload["auth_type"], "device"); + assert!(payload.get("callback_required").is_none() || payload["callback_required"] == false); + assert_eq!(*authorize_hits.lock().expect("mutex should lock"), 1); + + let stored = state + .load_provider_oauth_device_session_for_tests(&format!("device_auth_session:{session_id}")) + .expect("device session should be stored"); + let stored: serde_json::Value = + serde_json::from_str(&stored).expect("device session json should parse"); + assert_eq!(stored["provider_id"], "provider-xai"); + assert_eq!(stored["device_code"], "xai-device-code"); + assert_eq!(stored["auth_type"], "device"); + assert_eq!(stored["redirect_uri"], format!("{oidc_url}/oauth2/token")); + assert_eq!(stored["proxy_node_id"], "proxy-node-xai"); + assert_eq!(stored["status"], "pending"); + + oidc_handle.abort(); +} + +#[test] +fn gateway_handles_admin_provider_oauth_device_poll_for_xai() { + run_admin_oauth_test( + "gateway_handles_admin_provider_oauth_device_poll_for_xai", + gateway_handles_admin_provider_oauth_device_poll_for_xai_impl, + ); +} + +async fn gateway_handles_admin_provider_oauth_device_poll_for_xai_impl() { + let token_hits = Arc::new(Mutex::new(0usize)); + let token_hits_clone = Arc::clone(&token_hits); + let access_token = sample_kiro_device_access_token("user@x.ai"); + let id_token = access_token.clone(); + let token_server = Router::new().fallback(any(move |_request: Request| { + let token_hits_inner = Arc::clone(&token_hits_clone); + let access_token = access_token.clone(); + let id_token = id_token.clone(); + async move { + let hit = { + let mut hits = token_hits_inner.lock().expect("mutex should lock"); + *hits += 1; + *hits + }; + if hit == 1 { + return ( + StatusCode::BAD_REQUEST, + Json(json!({ "error": "authorization_pending" })), + ) + .into_response(); + } + Json(json!({ + "access_token": access_token, + "refresh_token": "xai-refresh-token", + "token_type": "Bearer", + "expires_in": 3600, + "id_token": id_token, + })) + .into_response() + } + })); + + let mut provider = sample_provider("provider-xai", "xai", 10); + provider.provider_type = "xai".to_string(); + let endpoint = sample_endpoint( + "endpoint-xai-responses", + "provider-xai", + "openai:responses", + "https://cli-chat-proxy.grok.com/v1", + ); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + vec![], + )); + let (token_url, token_handle) = start_server(token_server).await; + let resolved_token_url = format!("{token_url}/oauth2/token"); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests( + provider_catalog_repository.clone(), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_provider_oauth_device_session_entry_for_tests( + "session-xai", + json!({ + "provider_id": "provider-xai", + "region": "", + "client_id": "b1a00492-073a-47ea-816f-4c329264a828", + "client_secret": "", + "device_code": "xai-device-code", + "auth_type": "device", + "social_provider": null, + "code_verifier": null, + "redirect_uri": resolved_token_url, + "machine_id": null, + "interval": 5, + "expires_at_unix_secs": 4_102_444_800u64, + "status": "pending", + "proxy_node_id": null, + "created_at_unix_ms": 1_711_000_000u64, + "key_id": null, + "email": null, + "replaced": false, + "error_msg": null, + }), + ) + .with_provider_oauth_token_url_for_tests("xai", resolved_token_url.clone()); + + let pending = local_admin_provider_oauth_response( + &state, + http::Method::POST, + "/api/admin/provider-oauth/providers/provider-xai/device-poll", + Some(json!({ "session_id": "session-xai" })), + ) + .await; + let pending_body = to_bytes(pending.into_body(), usize::MAX) + .await + .expect("pending body should read"); + let pending_payload: serde_json::Value = + serde_json::from_slice(&pending_body).expect("pending json should parse"); + assert_eq!( + pending_payload["status"], "pending", + "payload={pending_payload}" + ); + + let authorized = local_admin_provider_oauth_response( + &state, + http::Method::POST, + "/api/admin/provider-oauth/providers/provider-xai/device-poll", + Some(json!({ "session_id": "session-xai" })), + ) + .await; + let status = authorized.status(); + let body = to_bytes(authorized.into_body(), usize::MAX) + .await + .expect("authorized body should read"); + let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); + assert_eq!(status, StatusCode::OK, "payload={payload}"); + assert_eq!(payload["status"], "authorized"); + assert_eq!(payload["email"], "user@x.ai"); + assert_eq!(payload["replaced"], false); + assert_eq!(*token_hits.lock().expect("mutex should lock"), 2); + + let stored = state + .load_provider_oauth_device_session_for_tests("device_auth_session:session-xai") + .expect("device session should persist"); + let stored: serde_json::Value = + serde_json::from_str(&stored).expect("device session json should parse"); + assert_eq!(stored["status"], "authorized"); + let key_id = stored["key_id"] + .as_str() + .expect("key_id should be stored") + .to_string(); + assert_eq!(payload["key_id"], key_id); + + let persisted = provider_catalog_repository + .list_keys_by_ids(std::slice::from_ref(&key_id)) + .await + .expect("keys should load") + .into_iter() + .next() + .expect("persisted key should exist"); + assert_eq!(persisted.auth_type, "oauth"); + let decrypted_auth_config = decrypt_persisted_provider_auth_config(&persisted); + let auth_config: serde_json::Value = + serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); + assert_eq!(auth_config["provider_type"], "xai"); + assert_eq!(auth_config["auth_method"], "oauth"); + assert_eq!(auth_config["using_api"], false); + assert_eq!(auth_config["email"], "user@x.ai"); + + token_handle.abort(); +} + #[test] fn gateway_handles_admin_provider_oauth_device_poll_for_windsurf_one_time_token() { run_admin_oauth_test( diff --git a/apps/aether-gateway/src/tests/control/admin/pool.rs b/apps/aether-gateway/src/tests/control/admin/pool.rs index aec3221f9..21edb44b2 100644 --- a/apps/aether-gateway/src/tests/control/admin/pool.rs +++ b/apps/aether-gateway/src/tests/control/admin/pool.rs @@ -2364,6 +2364,206 @@ async fn gateway_marks_exhausted_codex_pool_key_as_blocked_when_flag_enabled() { assert_eq!(keys[0]["account_quota"], json!("5H剩余 0.0%")); } +#[tokio::test] +async fn gateway_reserve_minimum_quota_marks_and_filters_codex_pool_keys() { + for (reserve_enabled, used_percent, exhausted) in [ + (true, 99.0, true), + (false, 99.0, false), + (true, 98.0, false), + ] { + let mut provider = sample_provider("provider-codex", "codex", 10); + provider.provider_type = "codex".to_string(); + provider.config = Some(json!({ + "pool_advanced": { + "reserve_minimum_quota": reserve_enabled, + "skip_exhausted_accounts": false, + "auto_remove_quota_exhausted_keys": true + } + })); + let mut key = sample_key( + "key-codex-reserved", + "provider-codex", + "openai:responses", + "oauth-placeholder", + ); + key.auth_type = "oauth".to_string(); + key.upstream_metadata = Some(json!({ + "codex": {"primary_used_percent": used_percent} + })); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + Vec::new(), + vec![key], + )), + )); + for status in ["all", "quota_exhausted", "available"] { + let response = local_admin_pool_response( + &state, + http::Method::GET, + &format!("/api/admin/pool/provider-codex/keys?page=1&page_size=50&status={status}"), + None, + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + let payload: serde_json::Value = serde_json::from_slice( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"), + ) + .expect("json body should parse"); + let keys = payload["keys"].as_array().expect("keys should be array"); + let visible = status == "all" || (status == "quota_exhausted") == exhausted; + assert_eq!( + keys.len(), + usize::from(visible), + "reserve={reserve_enabled}, used={used_percent}, status={status}" + ); + if visible { + assert_eq!(keys[0]["key_id"], "key-codex-reserved"); + assert_eq!( + keys[0]["scheduling_reason"] == "account_quota_exhausted", + exhausted + ); + if exhausted { + assert_eq!(keys[0]["scheduling_status"], "blocked"); + assert_eq!(keys[0]["scheduling_label"], "额度耗尽"); + } + } + } + } +} + +#[tokio::test] +async fn gateway_reserve_minimum_quota_recovers_after_config_update_with_cached_catalog() { + let mut provider = sample_provider("provider-codex", "codex", 10); + provider.provider_type = "codex".to_string(); + provider.config = Some(json!({ + "pool_advanced": {"reserve_minimum_quota": true, "skip_exhausted_accounts": true} + })); + let mut key = sample_key( + "key-codex", + "provider-codex", + "openai:responses", + "oauth-placeholder", + ); + key.auth_type = "oauth".to_string(); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", "updated_at": 100, + "code": "exhausted", "exhausted": true, "allowed": false, + "windows": [{ + "code": "weekly", "scope": "account", "used_ratio": 1.0, + "remaining_ratio": 0.0, "reset_at": 4_102_444_800u64 + }] + } + })); + key.upstream_metadata = Some(json!({ + "codex": {"updated_at": 200, "primary_used_percent": 99.0, + "primary_reset_at": 4_102_444_800u64} + })); + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider.clone()], + Vec::new(), + vec![key], + )); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone(&repository)) + .with_cached_provider_catalog_reader_for_tests(repository), + ); + for reserve_enabled in [true, false, true] { + provider.config.as_mut().unwrap()["pool_advanced"]["reserve_minimum_quota"] = + json!(reserve_enabled); + state + .update_provider_catalog_provider(&provider) + .await + .expect("provider should update"); + for status in ["all", "quota_exhausted", "available"] { + let response = local_admin_pool_response( + &state, + http::Method::GET, + &format!("/api/admin/pool/provider-codex/keys?status={status}"), + None, + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + let payload: serde_json::Value = serde_json::from_slice( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"), + ) + .expect("json should parse"); + let keys = payload["keys"].as_array().expect("keys should be array"); + let visible = status == "all" || (status == "quota_exhausted") == reserve_enabled; + assert_eq!(keys.len(), usize::from(visible)); + if visible { + assert_eq!( + keys[0]["scheduling_reason"] == "account_quota_exhausted", + reserve_enabled + ); + assert_eq!(keys[0]["status_snapshot"]["quota"]["code"], "ok"); + } + } + } +} + +#[tokio::test] +async fn gateway_codex_pool_filter_clears_stale_exhausted_summary_with_remaining_quota() { + let mut provider = sample_provider("provider-codex", "codex", 10); + provider.provider_type = "codex".to_string(); + provider.config = Some(json!({"pool_advanced": { + "reserve_minimum_quota": false, "skip_exhausted_accounts": true + }})); + let mut key = sample_key( + "key-codex", + "provider-codex", + "openai:responses", + "oauth-placeholder", + ); + key.auth_type = "oauth".to_string(); + key.status_snapshot = Some(json!({"quota": { + "provider_type": "codex", "code": "exhausted", "exhausted": true, + "windows": [{"code": "weekly", "scope": "account", "used_ratio": 0.83, + "remaining_ratio": 0.17, "reset_at": 4_102_444_800u64}] + }})); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + Vec::new(), + vec![key], + )), + )); + for status in ["all", "available", "quota_exhausted"] { + let response = local_admin_pool_response( + &state, + http::Method::GET, + &format!("/api/admin/pool/provider-codex/keys?status={status}"), + None, + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + let payload: serde_json::Value = serde_json::from_slice( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"), + ) + .expect("json should parse"); + let keys = payload["keys"].as_array().expect("keys should be array"); + assert_eq!(keys.len(), usize::from(status != "quota_exhausted")); + if let Some(key) = keys.first() { + assert_eq!(key["scheduling_status"], "available"); + assert_eq!(key["status_snapshot"]["quota"]["code"], "ok"); + assert!(key["account_quota"].as_str().unwrap().contains("17.0%")); + } + } +} + #[tokio::test] async fn gateway_lists_inherited_fixed_provider_api_formats_for_pool_keys() { let mut provider = sample_provider("provider-codex", "codex", 10).with_transport_fields( @@ -2666,8 +2866,10 @@ async fn gateway_treats_stale_codex_exhausted_snapshot_as_available_when_windows assert_eq!(keys[0]["scheduling_reason"], json!("available")); assert_eq!( keys[0]["account_quota"], - json!("周剩余 100.0% (7天0小时后重置) | 5H剩余 100.0% (5小时0分钟后重置)") + json!("周剩余 100.0% | 5H剩余 100.0%") ); + assert_eq!(keys[0]["status_snapshot"]["quota"]["code"], "ok"); + assert_eq!(keys[0]["status_snapshot"]["quota"]["exhausted"], false); } #[tokio::test] diff --git a/apps/aether-gateway/src/tests/control/admin/provider_ops.rs b/apps/aether-gateway/src/tests/control/admin/provider_ops.rs index ccc2fee1f..a9b2f5647 100644 --- a/apps/aether-gateway/src/tests/control/admin/provider_ops.rs +++ b/apps/aether-gateway/src/tests/control/admin/provider_ops.rs @@ -68,21 +68,11 @@ where F: FnOnce() -> Fut + Send + 'static, Fut: std::future::Future + 'static, { - let handle = std::thread::Builder::new() - .name(test_name.to_string()) - .stack_size(PROVIDER_OPS_TEST_STACK_BYTES) - .spawn(move || { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("test runtime should build"); - runtime.block_on(make_future()); - }) - .expect("provider ops test thread should spawn"); - - if let Err(payload) = handle.join() { - std::panic::resume_unwind(payload); - } + crate::tests::run_async_test_on_large_stack( + test_name, + PROVIDER_OPS_TEST_STACK_BYTES, + make_future, + ); } async fn start_managed_redis_or_skip() -> Option { diff --git a/apps/aether-gateway/src/tests/control/admin/provider_query.rs b/apps/aether-gateway/src/tests/control/admin/provider_query.rs index 82af4621c..336063114 100644 --- a/apps/aether-gateway/src/tests/control/admin/provider_query.rs +++ b/apps/aether-gateway/src/tests/control/admin/provider_query.rs @@ -36,21 +36,11 @@ where F: FnOnce() -> Fut + Send + 'static, Fut: std::future::Future + 'static, { - let handle = std::thread::Builder::new() - .name(test_name.to_string()) - .stack_size(PROVIDER_QUERY_TEST_STACK_BYTES) - .spawn(move || { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("test runtime should build"); - runtime.block_on(make_future()); - }) - .expect("provider query test thread should spawn"); - - if let Err(payload) = handle.join() { - std::panic::resume_unwind(payload); - } + crate::tests::run_async_test_on_large_stack( + test_name, + PROVIDER_QUERY_TEST_STACK_BYTES, + make_future, + ); } fn crc32(data: &[u8]) -> u32 { @@ -560,12 +550,12 @@ async fn gateway_recovers_codex_slug_only_models_from_a_stale_legacy_cache_impl( plan.url, format!( "https://chatgpt.com/backend-api/codex/models?client_version={}", - aether_ai_formats::CODEX_CLIENT_VERSION + aether_ai_formats::codex_client_version() ) ); assert_eq!( plan.headers.get("user-agent").map(String::as_str), - Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT) + Some(aether_ai_formats::codex_client_user_agent().as_str()) ); assert_eq!(plan.provider_api_format, "openai:responses"); Json(json!({ @@ -764,7 +754,7 @@ async fn gateway_handles_admin_provider_query_models_falls_back_to_codex_preset_ plan.url, format!( "https://chatgpt.com/backend-api/codex/models?client_version={}", - aether_ai_formats::CODEX_CLIENT_VERSION + aether_ai_formats::codex_client_version() ) ); Json(json!({ diff --git a/apps/aether-gateway/src/tests/control/admin/security.rs b/apps/aether-gateway/src/tests/control/admin/security.rs index 2df981435..aa96d2d2d 100644 --- a/apps/aether-gateway/src/tests/control/admin/security.rs +++ b/apps/aether-gateway/src/tests/control/admin/security.rs @@ -239,11 +239,25 @@ async fn send_admin_security_request( method: reqwest::Method, path: &str, body: Option, +) -> (StatusCode, serde_json::Value, usize) { + let path = path.to_string(); + crate::tests::run_async_test_on_large_stack_with_result( + "admin-security-router-request", + 16 * 1024 * 1024, + move || send_admin_security_request_on_large_stack(gateway, method, path, body), + ) +} + +async fn send_admin_security_request_on_large_stack( + gateway: Router, + method: reqwest::Method, + path: String, + body: Option, ) -> (StatusCode, serde_json::Value, usize) { let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits_clone = Arc::clone(&upstream_hits); let upstream = Router::new().route( - path, + &path, any(move |_request: Request| { let upstream_hits_inner = Arc::clone(&upstream_hits_clone); async move { @@ -253,26 +267,52 @@ async fn send_admin_security_request( }), ); - let (upstream_url, upstream_handle) = start_server(upstream).await; - let (gateway_url, gateway_handle) = start_server(gateway).await; + let (_upstream_url, upstream_handle) = start_server(upstream).await; - let client = reqwest::Client::new(); - let mut request = client - .request(method, format!("{gateway_url}{path}")) + // 这些用例只验证本地安全路由和“不得转发”断言,不需要为 Gateway + // 再启动一个 TCP listener;send_request 会补齐 ConnectInfo,仍经过完整 Router。 + let mut request_builder = Request::builder() + .method(method.as_str()) + .uri(&path) .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") .header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123") .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123"); if let Some(body) = body { - request = request.json(&body); + request_builder = request_builder.header(http::header::CONTENT_TYPE, "application/json"); + let request = request_builder + .body(Body::from(body.to_string())) + .expect("request should build"); + let response = send_request(gateway, request).await; + let status = response.status(); + let payload = response + .into_body() + .collect() + .await + .expect("response body should collect") + .to_bytes(); + let payload: serde_json::Value = + serde_json::from_slice(&payload).expect("json body should parse"); + let upstream_count = *upstream_hits.lock().expect("mutex should lock"); + upstream_handle.abort(); + return (status, payload, upstream_count); } - let response = request.send().await.expect("request should succeed"); + let request = request_builder + .body(Body::empty()) + .expect("request should build"); + let response = send_request(gateway, request).await; let status = response.status(); - let payload: serde_json::Value = response.json().await.expect("json body should parse"); + let payload = response + .into_body() + .collect() + .await + .expect("response body should collect") + .to_bytes(); + let payload: serde_json::Value = + serde_json::from_slice(&payload).expect("json body should parse"); let upstream_count = *upstream_hits.lock().expect("mutex should lock"); - gateway_handle.abort(); upstream_handle.abort(); (status, payload, upstream_count) @@ -347,6 +387,42 @@ async fn gateway_handles_admin_security_blacklist_add_locally_with_trusted_admin assert_eq!(upstream_count, 0); } +/// 真实 TCP 冒烟测试:其余安全用例已改为进程内 Router 调用以提速,这里保留一条 +/// 覆盖网络层装配(真实监听端口、HTTP 请求头传递、JSON 收发)的端到端路径。 +/// +/// `/api/admin/security/*` 在路由分类中是本地管理端点 +/// (`execution_runtime_candidate: false`),架构上不经过任何可注入 base_url 的上游, +/// 因此这里不构造无意义的“上游计数器”,只验证真实链路下本地处理结果正确。 +#[tokio::test] +async fn gateway_serves_admin_security_blacklist_over_real_tcp() { + let gateway = build_router_with_state(AppState::new().expect("gateway should build")); + let (gateway_url, gateway_handle) = start_server(gateway).await; + + let response = reqwest::Client::new() + .post(format!("{gateway_url}/api/admin/security/ip/blacklist")) + .header(GATEWAY_HEADER, "rust-phase3b") + .header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123") + .header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") + .header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123") + .json(&json!({ "ip_address": "1.2.3.4", "reason": "manual", "ttl": 60 })) + .send() + .await + .expect("request should reach the gateway over TCP"); + + let status = response.status(); + let payload: serde_json::Value = response + .json() + .await + .expect("gateway response should be json"); + assert_eq!(status, StatusCode::OK); + assert_eq!(payload["success"], true); + assert_eq!(payload["message"], "IP 1.2.3.4 已加入黑名单"); + assert_eq!(payload["reason"], "manual"); + assert_eq!(payload["ttl"], 60); + + gateway_handle.abort(); +} + #[tokio::test] async fn gateway_rejects_invalid_admin_security_blacklist_ip() { let gateway = build_router_with_state(AppState::new().expect("gateway should build")); diff --git a/apps/aether-gateway/src/tests/control/admin/stats.rs b/apps/aether-gateway/src/tests/control/admin/stats.rs index 96da193bc..d7bfe0da8 100644 --- a/apps/aether-gateway/src/tests/control/admin/stats.rs +++ b/apps/aether-gateway/src/tests/control/admin/stats.rs @@ -9,7 +9,8 @@ use aether_data::repository::auth::{ use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::usage::InMemoryUsageReadRepository; use aether_data::repository::users::{ - InMemoryUserReadRepository, StoredUserAuthRecord, StoredUserSummary, + InMemoryUserReadRepository, StoredUserAuthRecord, StoredUserSummary, UpsertUserGroupRecord, + UserReadRepository, }; use aether_data_contracts::repository::usage::StoredRequestUsageAudit; use async_trait::async_trait; @@ -1939,6 +1940,356 @@ async fn gateway_handles_admin_stats_leaderboard_users_without_legacy_username_f upstream_handle.abort(); } +#[tokio::test] +async fn gateway_aggregates_admin_stats_by_current_user_group_membership() { + let (_upstream_url, upstream_hits, upstream_handle) = + start_stats_upstream("/api/admin/stats/leaderboard/user-groups").await; + let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![ + sample_usage_row( + "usage-group-a", + "req-group-a", + Some("user-1"), + Some("key-1"), + Some("key-1"), + "OpenAI", + "gpt-5", + 60, + 20, + 0.4, + 0.4, + DAY_1_UNIX_SECS, + ), + sample_usage_row( + "usage-group-b", + "req-group-b", + Some("user-2"), + Some("key-2"), + Some("key-2"), + "OpenAI", + "gpt-5", + 40, + 10, + 0.35, + 0.35, + DAY_1_UNIX_SECS + 10, + ), + sample_usage_row( + "usage-group-outside", + "req-group-outside", + Some("user-3"), + Some("key-3"), + Some("key-3"), + "OpenAI", + "gpt-5", + 100, + 50, + 1.5, + 1.5, + DAY_1_UNIX_SECS + 20, + ), + ])); + let user_repository = InMemoryUserReadRepository::seed_auth_users([ + sample_auth_user("user-1", "alice", "user", true), + sample_auth_user("user-2", "bob", "user", true), + sample_auth_user("user-3", "carol", "user", true), + ]); + let group = user_repository + .create_user_group(UpsertUserGroupRecord { + name: "Engineering".to_string(), + description: None, + priority: 0, + allowed_providers: None, + allowed_providers_mode: "inherit".to_string(), + allowed_api_formats: None, + allowed_api_formats_mode: "inherit".to_string(), + allowed_models: None, + allowed_models_mode: "inherit".to_string(), + rate_limit: None, + rate_limit_mode: "inherit".to_string(), + }) + .await + .expect("group creation should succeed") + .expect("group should be created"); + user_repository + .replace_user_group_members(&group.id, &["user-1".to_string(), "user-2".to_string()]) + .await + .expect("group members should be replaced"); + let data_state = GatewayDataState::with_usage_reader_for_tests(usage_repository) + .with_user_reader(Arc::new(user_repository)); + let gateway = build_router_with_state( + AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data_state), + ); + let (gateway_url, gateway_handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + + let response = admin_request(client.get(format!( + "{gateway_url}/api/admin/stats/leaderboard/user-groups?start_date=2024-03-21&end_date=2024-03-21&metric=cost&tz_offset_minutes=0" + ))) + .send() + .await + .expect("group leaderboard request should succeed"); + assert_eq!(response.status(), StatusCode::OK); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["attribution"], "current_membership"); + assert_eq!(payload["total"], 2); + let mut payload = payload; + payload["items"] + .as_array_mut() + .unwrap() + .retain(|item| item["id"] == group.id); + assert_eq!(payload["items"][0]["id"], group.id); + assert_eq!(payload["items"][0]["name"], "Engineering"); + assert_eq!(payload["items"][0]["requests"], 2); + assert_eq!(payload["items"][0]["cost"], 0.75); + assert_eq!(payload["items"][0]["member_count"], 2); + assert_eq!(payload["items"][0]["active_member_count"], 2); + + let response = admin_request(client.get(format!( + "{gateway_url}/api/admin/stats/leaderboard/users?start_date=2024-03-21&end_date=2024-03-21&metric=cost&tz_offset_minutes=0&user_group_id={}", + group.id + ))) + .send() + .await + .expect("group member leaderboard request should succeed"); + assert_eq!(response.status(), StatusCode::OK); + let payload: serde_json::Value = response.json().await.expect("json body should parse"); + assert_eq!(payload["total"], 2); + assert!(payload["items"] + .as_array() + .is_some_and(|items| { items.iter().all(|item| item["id"] != "user-3") })); + + let response = admin_request(client.get(format!( + "{gateway_url}/api/admin/stats/leaderboard/users?start_date=2024-03-21&end_date=2024-03-21&user_id=user-1&user_group_id={}", + group.id + ))) + .send() + .await + .expect("conflicting scope request should complete"); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); + + gateway_handle.abort(); + upstream_handle.abort(); +} + +#[tokio::test] +async fn gateway_group_usage_intersects_members_and_providers_with_shared_overlap() { + let rows = vec![ + sample_usage_row( + "g", + "g", + Some("user-1"), + None, + None, + "Gemini", + "model", + 10, + 2, + 0.4, + 0.4, + DAY_1_UNIX_SECS, + ), + sample_usage_row( + "s", + "s", + Some("user-1"), + None, + None, + "Shared", + "model", + 10, + 2, + 0.35, + 0.35, + DAY_1_UNIX_SECS, + ), + sample_usage_row( + "o", + "o", + Some("user-1"), + None, + None, + "Other", + "model", + 10, + 2, + 1.5, + 1.5, + DAY_1_UNIX_SECS, + ), + sample_usage_row( + "x", + "x", + Some("user-2"), + None, + None, + "Gemini", + "model", + 10, + 2, + 10.0, + 10.0, + DAY_1_UNIX_SECS, + ), + ]; + let users = InMemoryUserReadRepository::seed_auth_users([ + sample_auth_user("user-1", "alice", "user", true), + sample_auth_user("user-2", "bob", "user", true), + ]); + let export_users = users.list_export_users().await.unwrap(); + let users = users.with_export_users(export_users); + let mut group_ids = Vec::new(); + for (name, allowed, mode) in [ + ( + "Gemini group", + vec!["provider-gemini", "Shared", "provider-shared"], + "specific", + ), + ("Other group", vec!["Other", "shared-type"], "specific"), + ("Denied", vec!["Gemini"], "deny_all"), + ("Empty", vec![], "specific"), + ("Inherited", vec![], "inherit"), + ("Unrestricted", vec![], "unrestricted"), + ] { + let group = users + .create_user_group(UpsertUserGroupRecord { + name: name.to_string(), + description: None, + priority: 0, + allowed_providers: Some(allowed.into_iter().map(str::to_string).collect()), + allowed_providers_mode: mode.to_string(), + allowed_api_formats: None, + allowed_api_formats_mode: "inherit".to_string(), + allowed_models: None, + allowed_models_mode: "inherit".to_string(), + rate_limit: None, + rate_limit_mode: "inherit".to_string(), + }) + .await + .unwrap() + .unwrap(); + users + .replace_user_group_members(&group.id, &["user-1".to_string()]) + .await + .unwrap(); + group_ids.push(group.id); + } + let mut shared = sample_provider("provider-shared", "Shared", 0); + shared.provider_type = "shared-type".to_string(); + let providers = InMemoryProviderCatalogReadRepository::seed( + vec![ + sample_provider("provider-gemini", "Gemini", 0), + shared, + sample_provider("provider-other", "Other", 0), + ], + vec![], + vec![], + ); + let data = GatewayDataState::with_usage_reader_for_tests(Arc::new( + InMemoryUsageReadRepository::seed(rows), + )) + .with_user_reader(Arc::new(users)) + .with_provider_catalog_reader(Arc::new(providers)); + let gateway = build_router_with_state(AppState::new().unwrap().with_data_state_for_tests(data)); + let (url, handle) = start_server(gateway).await; + let client = reqwest::Client::new(); + let range = "start_date=2024-03-21&end_date=2024-03-21&tz_offset_minutes=0"; + let paths = [ + ( + format!("usage/stats?{range}&user_group_id=__ungrouped__"), + Some(10.0), + ), + ( + format!("stats/time-series?{range}&granularity=day&user_group_id=__ungrouped__"), + Some(10.0), + ), + ( + format!("stats/leaderboard/users?{range}&metric=cost&user_group_id=__ungrouped__"), + Some(10.0), + ), + ( + format!("stats/leaderboard/user-groups?{range}&metric=cost"), + None, + ), + ( + format!("usage/stats?{range}&user_group_id={}", group_ids[0]), + Some(0.75), + ), + ( + format!( + "stats/time-series?{range}&granularity=day&user_group_id={}", + group_ids[1] + ), + Some(1.85), + ), + ( + format!( + "stats/leaderboard/users?{range}&metric=cost&user_group_id={}", + group_ids[0] + ), + Some(0.75), + ), + (format!("usage/stats?{range}&user_id=user-1"), Some(2.25)), + ( + format!( + "usage/stats?{range}&user_group_id={}&provider=Other", + group_ids[0] + ), + Some(0.0), + ), + ]; + for (path, expected) in paths { + let response = admin_request(client.get(format!("{url}/api/admin/{path}"))) + .send() + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK, "{path}"); + let body: serde_json::Value = response.json().await.unwrap(); + if let Some(expected) = expected { + let value = if path.starts_with("stats/time-series") { + &body[0]["total_cost"] + } else if path.starts_with("stats/leaderboard") { + &body["items"][0]["cost"] + } else { + &body["total_cost"] + }; + assert!( + (value.as_f64().unwrap() - expected).abs() < 1e-9, + "{path}: {body}" + ); + } else { + let items = body["items"].as_array().unwrap(); + assert_eq!(items.len(), 7); + let ungrouped = items + .iter() + .find(|item| item["id"] == "__ungrouped__") + .unwrap(); + assert_eq!(ungrouped["cost"], 10.0); + assert_eq!(ungrouped["member_count"], 1); + assert_eq!(ungrouped["active_member_count"], 1); + + for (id, cost, requests) in [ + (&group_ids[0], 0.75, 2), + (&group_ids[1], 1.85, 2), + (&group_ids[2], 0.0, 0), + (&group_ids[3], 0.0, 0), + (&group_ids[4], 2.25, 3), + (&group_ids[5], 2.25, 3), + ] { + let row = items + .iter() + .find(|item| item["id"].as_str() == Some(id.as_str())) + .unwrap(); + assert!((row["cost"].as_f64().unwrap() - cost).abs() < 1e-9); + assert_eq!(row["requests"], requests); + } + } + } + handle.abort(); +} + #[tokio::test] async fn gateway_handles_admin_stats_leaderboard_users_locally_without_usage_reader() { let (upstream_url, upstream_hits, upstream_handle) = diff --git a/apps/aether-gateway/src/tests/control/admin/system_import.rs b/apps/aether-gateway/src/tests/control/admin/system_import.rs index 022a7ece6..8ad2871b1 100644 --- a/apps/aether-gateway/src/tests/control/admin/system_import.rs +++ b/apps/aether-gateway/src/tests/control/admin/system_import.rs @@ -341,21 +341,11 @@ where F: FnOnce() -> Fut + Send + 'static, Fut: std::future::Future + 'static, { - let handle = std::thread::Builder::new() - .name(test_name.to_string()) - .stack_size(ADMIN_SYSTEM_IMPORT_TEST_STACK_BYTES) - .spawn(move || { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("test runtime should build"); - runtime.block_on(make_future()); - }) - .expect("admin system import test thread should spawn"); - - if let Err(payload) = handle.join() { - std::panic::resume_unwind(payload); - } + crate::tests::run_async_test_on_large_stack( + test_name, + ADMIN_SYSTEM_IMPORT_TEST_STACK_BYTES, + make_future, + ); } #[test] diff --git a/apps/aether-gateway/src/tests/control/admin/users.rs b/apps/aether-gateway/src/tests/control/admin/users.rs index 8134f2106..0d9fcde6e 100644 --- a/apps/aether-gateway/src/tests/control/admin/users.rs +++ b/apps/aether-gateway/src/tests/control/admin/users.rs @@ -2732,7 +2732,9 @@ async fn gateway_lists_admin_user_api_keys_locally_with_trusted_admin_principal( Some(1_711_000_100), Some(1_711_000_101), ) - .expect("export activity timestamps should build")]), + .expect("export activity timestamps should build") + .with_ip_rules(Some(json!(["203.0.113.0/24"]))) + .expect("export IP rules should build")]), ); let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![ sample_admin_user("user-1"), @@ -2769,6 +2771,10 @@ async fn gateway_lists_admin_user_api_keys_locally_with_trusted_admin_principal( assert_eq!(payload["api_keys"][0]["key_display"], "sk...-1"); assert_eq!(payload["api_keys"][0]["is_active"], true); assert_eq!(payload["api_keys"][0]["is_locked"], false); + assert_eq!( + payload["api_keys"][0]["ip_rules"], + json!(["203.0.113.0/24"]) + ); assert_eq!(payload["api_keys"][0]["total_requests"], 9); assert_eq!(payload["api_keys"][0]["total_cost_usd"], 1.5); assert_eq!(payload["api_keys"][0]["rate_limit"], 60); diff --git a/apps/aether-gateway/src/tests/control/admin/users_batch.rs b/apps/aether-gateway/src/tests/control/admin/users_batch.rs new file mode 100644 index 000000000..35efa942e --- /dev/null +++ b/apps/aether-gateway/src/tests/control/admin/users_batch.rs @@ -0,0 +1,650 @@ +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::Arc; + +use aether_data::repository::management_tokens::InMemoryManagementTokenRepository; +use aether_data::repository::users::{InMemoryUserReadRepository, StoredUserAuthRecord}; +use aether_data::repository::wallet::StoredWalletSnapshot; +use axum::http::StatusCode; +use chrono::Utc; +use reqwest::{Client, RequestBuilder, Response}; +use serde_json::{json, Value}; + +use super::super::{ + build_router_with_state, hash_management_token, sample_management_token, start_server, AppState, +}; +use crate::data::GatewayDataState; + +fn admin_headers(request: RequestBuilder) -> RequestBuilder { + request + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .header(crate::constants::TRUSTED_ADMIN_USER_ID_HEADER, "admin-user") + .header(crate::constants::TRUSTED_ADMIN_USER_ROLE_HEADER, "admin") + .header( + crate::constants::TRUSTED_ADMIN_SESSION_ID_HEADER, + "session-admin", + ) +} + +fn sample_user(user_id: &str) -> StoredUserAuthRecord { + sample_user_with_role(user_id, "user") +} + +fn sample_user_with_role(user_id: &str, role: &str) -> StoredUserAuthRecord { + StoredUserAuthRecord::new( + user_id.to_string(), + Some(format!("{user_id}@example.com")), + true, + user_id.to_string(), + Some("hash".to_string()), + role.to_string(), + "local".to_string(), + Some(json!(["openai"])), + Some(json!(["openai:chat"])), + Some(json!(["gpt-4.1"])), + true, + false, + Some(Utc::now()), + Some(Utc::now()), + ) + .expect("test user should build") +} + +fn sample_wallet(user_id: &str, balance: f64, gift_balance: f64) -> StoredWalletSnapshot { + StoredWalletSnapshot::new( + format!("wallet-{user_id}"), + Some(user_id.to_string()), + None, + balance, + gift_balance, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + balance.max(0.0), + 0.0, + 0.0, + 0.0, + 1_710_000_000, + ) + .expect("test wallet should build") +} + +async fn post_batch_action(client: &Client, gateway_url: &str, payload: Value) -> Response { + static NEXT_TEST_IDEMPOTENCY_KEY: AtomicU64 = AtomicU64::new(1); + let mut payload = payload; + if payload.get("action").and_then(Value::as_str) == Some("adjust_wallet_balance") + && payload.get("idempotency_key").is_none() + { + let sequence = NEXT_TEST_IDEMPOTENCY_KEY.fetch_add(1, Ordering::Relaxed); + payload["idempotency_key"] = json!(format!("test-wallet-batch-{sequence}")); + } + admin_headers(client.post(format!("{gateway_url}/api/admin/users/batch-action"))) + .json(&payload) + .send() + .await + .expect("batch request should complete") +} + +async fn wallet_detail(client: &Client, gateway_url: &str, user_id: &str) -> Value { + admin_headers(client.get(format!("{gateway_url}/api/admin/wallets/wallet-{user_id}"))) + .send() + .await + .expect("wallet lookup should complete") + .json() + .await + .expect("wallet response should parse") +} + +#[tokio::test] +async fn gateway_batches_wallet_addition_deduction_and_clamped_deduction_per_user() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([sample_user("user-1"), sample_user("user-2")]) + .with_auth_wallets_for_tests([ + sample_wallet("user-1", 10.0, 3.0), + sample_wallet("user-2", 2.0, 1.0), + ]); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + let selection = json!({ "user_ids": ["user-1", "user-2"] }); + + let add_response = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": selection.clone(), + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 5.0 } + }), + ) + .await; + assert_eq!(add_response.status(), StatusCode::OK); + let add_result: Value = add_response.json().await.expect("response should parse"); + assert_eq!(add_result["success"], 2); + assert_eq!(add_result["failed"], 0); + assert_eq!(add_result["modified_fields"], json!(["wallet_balance"])); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 18.0 + ); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-2").await["balance"], + 8.0 + ); + + let deduct_response = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": selection.clone(), + "action": "adjust_wallet_balance", + "payload": { "operation": "deduct", "amount": 4.0 } + }), + ) + .await; + assert_eq!(deduct_response.status(), StatusCode::OK); + let deduct_result: Value = deduct_response.json().await.expect("response should parse"); + assert_eq!(deduct_result["success"], 2); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 14.0 + ); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-2").await["balance"], + 4.0 + ); + + let over_deduct_response = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": selection, + "action": "adjust_wallet_balance", + "payload": { "operation": "deduct", "amount": 100.0 } + }), + ) + .await; + assert_eq!(over_deduct_response.status(), StatusCode::OK); + let over_deduct_result: Value = over_deduct_response + .json() + .await + .expect("response should parse"); + assert_eq!(over_deduct_result["success"], 2); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 0.0 + ); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-2").await["balance"], + 0.0 + ); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_replays_wallet_batch_idempotently_and_rejects_key_reuse() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([sample_user("user-1")]) + .with_auth_wallets_for_tests([sample_wallet("user-1", 10.0, 0.0)]); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + let request = json!({ + "selection": { "user_ids": ["user-1"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 5.0 }, + "idempotency_key": "same-wallet-batch" + }); + + let first = post_batch_action(&client, &gateway_url, request.clone()).await; + assert_eq!(first.status(), StatusCode::OK); + let first_result: Value = first.json().await.expect("response should parse"); + assert_eq!(first_result["success"], 1); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 15.0 + ); + + let replay = post_batch_action(&client, &gateway_url, request.clone()).await; + assert_eq!(replay.status(), StatusCode::OK); + let replay_result: Value = replay.json().await.expect("response should parse"); + assert_eq!(replay_result["success"], 1); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 15.0 + ); + + let changed_request = json!({ + "selection": { "user_ids": ["user-1"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 50.0 }, + "idempotency_key": "same-wallet-batch" + }); + let conflict = post_batch_action(&client, &gateway_url, changed_request).await; + assert_eq!(conflict.status(), StatusCode::CONFLICT); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 15.0 + ); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_returns_partial_results_when_failure_recording_fails() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([sample_user("user-1"), sample_user("user-2")]) + .with_auth_wallets_for_tests([sample_wallet("user-1", 10.0, 0.0)]) + .fail_auth_wallet_batch_failure_record_for_tests("user-2"); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + let request = json!({ + "selection": { "user_ids": ["user-1", "user-2"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 5.0 }, + "idempotency_key": "failure-record-wallet-batch" + }); + + let response = post_batch_action(&client, &gateway_url, request).await; + assert_eq!(response.status(), StatusCode::OK); + let result: Value = response.json().await.expect("response should parse"); + assert_eq!(result["success"], 1); + assert_eq!(result["failed"], 1); + assert_eq!(result["interrupted"], true); + assert_eq!(result["completed_user_ids"], json!(["user-1"])); + assert_eq!(result["uncertain_user_ids"], json!([])); + assert_eq!(result["unprocessed_user_ids"], json!(["user-2"])); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 15.0 + ); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_records_zero_delta_batch_for_later_replay() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([sample_user("user-1")]) + .with_auth_wallets_for_tests([sample_wallet("user-1", 0.0, 0.0)]); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + let deduction = json!({ + "selection": { "user_ids": ["user-1"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "deduct", "amount": 5.0 }, + "idempotency_key": "zero-delta-wallet-batch" + }); + let first = post_batch_action(&client, &gateway_url, deduction.clone()).await; + assert_eq!(first.status(), StatusCode::OK); + + let top_up = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": { "user_ids": ["user-1"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 10.0 }, + "idempotency_key": "wallet-top-up-batch" + }), + ) + .await; + assert_eq!(top_up.status(), StatusCode::OK); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 10.0 + ); + + let replay = post_batch_action(&client, &gateway_url, deduction).await; + assert_eq!(replay.status(), StatusCode::OK); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 10.0 + ); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_requires_wallet_write_permission_for_batch_balance_adjustments() { + let users_write_token = "ae-batch-users-write-only"; + let wallet_write_token = "ae-batch-users-wallet-write"; + let wallet_admin_token = "ae-batch-users-wallet-admin"; + let token_owner = sample_user_with_role("token-owner", "admin"); + let target_user = sample_user("user-1"); + let mut users_only = sample_management_token( + "token-users-write-only", + &token_owner.id, + &token_owner.username, + true, + ); + users_only.token.allowed_ips = None; + users_only.token.permissions = Some(json!(["admin:users:write"])); + let mut users_and_wallets = sample_management_token( + "token-users-and-wallet-write", + &token_owner.id, + &token_owner.username, + true, + ); + users_and_wallets.token.allowed_ips = None; + users_and_wallets.token.permissions = Some(json!(["admin:users:write", "admin:wallets:write"])); + let mut wallets_admin = sample_management_token( + "token-users-wallet-admin", + &token_owner.id, + &token_owner.username, + true, + ); + wallets_admin.token.allowed_ips = None; + wallets_admin.token.permissions = Some(json!(["admin:users:write", "admin:wallets:admin"])); + let token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes( + vec![users_only, users_and_wallets, wallets_admin], + vec![ + ( + hash_management_token(users_write_token), + "token-users-write-only".to_string(), + ), + ( + hash_management_token(wallet_write_token), + "token-users-and-wallet-write".to_string(), + ), + ( + hash_management_token(wallet_admin_token), + "token-users-wallet-admin".to_string(), + ), + ], + )); + let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![ + token_owner.clone(), + target_user.clone(), + ])); + let data = GatewayDataState::with_management_token_repository_for_tests(token_repository) + .with_user_reader(user_repository); + let state = AppState::new() + .expect("gateway should build") + .with_data_state_for_tests(data) + .with_auth_users_for_tests([token_owner, target_user]) + .with_auth_wallets_for_tests([sample_wallet("user-1", 10.0, 0.0)]); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + let payload = json!({ + "selection": { "user_ids": ["user-1"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 5.0 }, + "idempotency_key": "wallet-write-batch" + }); + + let denied = client + .post(format!("{gateway_url}/api/admin/users/batch-action")) + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(users_write_token) + .json(&payload) + .send() + .await + .expect("users-only management token request should complete"); + assert_eq!(denied.status(), StatusCode::FORBIDDEN); + let denied_payload: Value = denied.json().await.expect("response should parse"); + assert_eq!( + denied_payload["required_permissions"], + json!(["admin:wallets:write", "admin:wallets:admin"]) + ); + assert_eq!(denied_payload["permission_mode"], "any_of"); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 10.0 + ); + + let allowed = client + .post(format!("{gateway_url}/api/admin/users/batch-action")) + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(wallet_write_token) + .json(&payload) + .send() + .await + .expect("wallet-write management token request should complete"); + assert_eq!(allowed.status(), StatusCode::OK); + let result: Value = allowed.json().await.expect("response should parse"); + assert_eq!(result["success"], 1); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 15.0 + ); + + let allowed_with_wallet_admin = client + .post(format!("{gateway_url}/api/admin/users/batch-action")) + .header(crate::constants::GATEWAY_HEADER, "rust-phase3b") + .bearer_auth(wallet_admin_token) + .json(&json!({ + "selection": payload["selection"], + "action": payload["action"], + "payload": payload["payload"], + "idempotency_key": "wallet-admin-batch" + })) + .send() + .await + .expect("wallet-admin management token request should complete"); + assert_eq!(allowed_with_wallet_admin.status(), StatusCode::OK); + let result: Value = allowed_with_wallet_admin + .json() + .await + .expect("response should parse"); + assert_eq!(result["success"], 1); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 20.0 + ); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_reports_completed_uncertain_and_unprocessed_users_after_adjustment_error() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([ + sample_user("user-1"), + sample_user("user-2"), + sample_user("user-3"), + ]) + .with_auth_wallets_for_tests([ + sample_wallet("user-1", 10.0, 0.0), + sample_wallet("user-2", 20.0, 0.0), + sample_wallet("user-3", 30.0, 0.0), + ]) + .fail_auth_wallet_adjustment_for_tests("wallet-user-2"); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + + let response = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": { "user_ids": ["user-1", "user-2", "user-3"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 5.0 } + }), + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + let result: Value = response.json().await.expect("response should parse"); + assert_eq!(result["interrupted"], true); + assert_eq!(result["success"], 1); + assert_eq!(result["failed"], 2); + assert_eq!(result["completed_user_ids"], json!(["user-1"])); + assert_eq!(result["uncertain_user_ids"], json!(["user-2"])); + assert_eq!(result["unprocessed_user_ids"], json!(["user-3"])); + assert_eq!(result["failures"][0]["user_id"], "user-2"); + assert_eq!(result["failures"][1]["user_id"], "user-3"); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 15.0 + ); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-2").await["balance"], + 20.0 + ); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-3").await["balance"], + 30.0 + ); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_reports_wallet_lookup_failure_as_unprocessed() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([ + sample_user("user-1"), + sample_user("user-2"), + sample_user("user-3"), + ]) + .with_auth_wallets_for_tests([ + sample_wallet("user-1", 10.0, 0.0), + sample_wallet("user-2", 20.0, 0.0), + sample_wallet("user-3", 30.0, 0.0), + ]) + .fail_auth_wallet_lookup_for_tests("user-2"); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + + let response = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": { "user_ids": ["user-1", "user-2", "user-3"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 5.0 } + }), + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + let result: Value = response.json().await.expect("response should parse"); + assert_eq!(result["interrupted"], true); + assert_eq!(result["completed_user_ids"], json!(["user-1"])); + assert_eq!(result["uncertain_user_ids"], json!([])); + assert_eq!(result["unprocessed_user_ids"], json!(["user-2", "user-3"])); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-1").await["balance"], + 15.0 + ); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-2").await["balance"], + 20.0 + ); + assert_eq!( + wallet_detail(&client, &gateway_url, "user-3").await["balance"], + 30.0 + ); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_reports_wallet_limit_lookup_failure_as_unprocessed() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([ + sample_user("user-1"), + sample_user("user-2"), + sample_user("user-3"), + ]) + .with_auth_wallets_for_tests([ + sample_wallet("user-1", 10.0, 0.0), + sample_wallet("user-2", 20.0, 0.0), + sample_wallet("user-3", 30.0, 0.0), + ]) + .fail_auth_wallet_lookup_for_tests("user-2"); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + + let response = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": { "user_ids": ["user-1", "user-2", "user-3"] }, + "action": "update_access_control", + "payload": { "unlimited": true } + }), + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + let result: Value = response.json().await.expect("response should parse"); + assert_eq!(result["interrupted"], true); + assert_eq!(result["completed_user_ids"], json!(["user-1"])); + assert_eq!(result["uncertain_user_ids"], json!([])); + assert_eq!(result["unprocessed_user_ids"], json!(["user-2", "user-3"])); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_reports_missing_wallet_and_floors_negative_balance_on_deduction() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([sample_user("user-negative"), sample_user("user-no-wallet")]) + .with_auth_wallets_for_tests([sample_wallet("user-negative", -2.0, 1.0)]); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + + let response = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": { "user_ids": ["user-negative", "user-no-wallet"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "deduct", "amount": 10.0 } + }), + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + let result: Value = response.json().await.expect("response should parse"); + assert_eq!(result["success"], 1); + assert_eq!(result["failed"], 1); + assert_eq!(result["failures"][0]["user_id"], "user-no-wallet"); + assert_eq!(result["failures"][0]["reason"], "用户钱包不可用"); + + let wallet = wallet_detail(&client, &gateway_url, "user-negative").await; + assert_eq!(wallet["balance"], 0.0); + assert_eq!(wallet["total_adjusted"], 1.0); + + gateway_handle.abort(); +} + +#[tokio::test] +async fn gateway_rejects_zero_and_non_finite_batch_wallet_adjustments() { + let state = AppState::new() + .expect("gateway should build") + .with_auth_users_for_tests([sample_user("user-1")]) + .with_auth_wallets_for_tests([sample_wallet("user-1", 10.0, 0.0)]); + let (gateway_url, gateway_handle) = start_server(build_router_with_state(state)).await; + let client = Client::new(); + + let zero_response = post_batch_action( + &client, + &gateway_url, + json!({ + "selection": { "user_ids": ["user-1"] }, + "action": "adjust_wallet_balance", + "payload": { "operation": "add", "amount": 0.0 } + }), + ) + .await; + assert_eq!(zero_response.status(), StatusCode::BAD_REQUEST); + + let non_finite_response = admin_headers( + client.post(format!("{gateway_url}/api/admin/users/batch-action")), + ) + .header(reqwest::header::CONTENT_TYPE, "application/json") + .body( + r#"{"selection":{"user_ids":["user-1"]},"action":"adjust_wallet_balance","payload":{"operation":"add","amount":1e999}}"#, + ) + .send() + .await + .expect("non-finite amount request should complete"); + assert_eq!(non_finite_response.status(), StatusCode::BAD_REQUEST); + + gateway_handle.abort(); +} diff --git a/apps/aether-gateway/src/tests/files/mod.rs b/apps/aether-gateway/src/tests/files/mod.rs index a9c9b1e7c..0acd7fb10 100644 --- a/apps/aether-gateway/src/tests/files/mod.rs +++ b/apps/aether-gateway/src/tests/files/mod.rs @@ -37,21 +37,7 @@ where F: FnOnce() -> Fut + Send + 'static, Fut: std::future::Future + 'static, { - let handle = std::thread::Builder::new() - .name(test_name.to_string()) - .stack_size(FILES_TEST_STACK_BYTES) - .spawn(move || { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("test runtime should build"); - runtime.block_on(make_future()); - }) - .expect("files test thread should spawn"); - - if let Err(payload) = handle.join() { - std::panic::resume_unwind(payload); - } + crate::tests::run_async_test_on_large_stack(test_name, FILES_TEST_STACK_BYTES, make_future); } fn hash_api_key(value: &str) -> String { diff --git a/apps/aether-gateway/src/tests/frontdoor.rs b/apps/aether-gateway/src/tests/frontdoor.rs index 15ae3a645..32a20d3e4 100644 --- a/apps/aether-gateway/src/tests/frontdoor.rs +++ b/apps/aether-gateway/src/tests/frontdoor.rs @@ -39,21 +39,7 @@ fn run_frontdoor_async_test(name: &'static str, future: F) where F: std::future::Future + Send + 'static, { - let handle = std::thread::Builder::new() - .name(name.to_string()) - .stack_size(16 * 1024 * 1024) - .spawn(move || { - tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("frontdoor test runtime should build") - .block_on(future); - }) - .expect("large-stack frontdoor test thread should spawn"); - - if let Err(payload) = handle.join() { - std::panic::resume_unwind(payload); - } + crate::tests::run_async_test_on_large_stack(name, 16 * 1024 * 1024, || future); } fn hash_api_key(value: &str) -> String { diff --git a/apps/aether-gateway/src/tests/mod.rs b/apps/aether-gateway/src/tests/mod.rs index 15d90bcf9..4e93f587a 100644 --- a/apps/aether-gateway/src/tests/mod.rs +++ b/apps/aether-gateway/src/tests/mod.rs @@ -10,7 +10,6 @@ pub(super) use http::StatusCode; pub(super) use serde_json::json; mod ai_execute; -mod architecture; mod async_task; mod audit; mod concurrency; @@ -46,6 +45,50 @@ pub(super) async fn start_server(app: Router) -> (String, tokio::task::JoinHandl (format!("http://{addr}"), handle) } +/// 在独立的大栈线程中运行需要深调用栈的异步测试。 +/// +/// 这些测试仍保留 16 MiB 栈空间;这里只统一线程和 runtime 的启动逻辑, +/// 避免每个测试分区各自复制一份 helper,降低维护时误改测试执行语义的风险。 +pub(crate) fn run_async_test_on_large_stack( + test_name: &'static str, + stack_size: usize, + make_future: F, +) where + F: FnOnce() -> Fut + Send + 'static, + Fut: std::future::Future + 'static, +{ + run_async_test_on_large_stack_with_result(test_name, stack_size, make_future); +} + +/// 与上面的 helper 相同,但允许深栈测试返回结果,供公共请求 helper 使用。 +pub(crate) fn run_async_test_on_large_stack_with_result( + test_name: &'static str, + stack_size: usize, + make_future: F, +) -> R +where + F: FnOnce() -> Fut + Send + 'static, + Fut: std::future::Future + 'static, + R: Send + 'static, +{ + let handle = std::thread::Builder::new() + .name(test_name.to_string()) + .stack_size(stack_size) + .spawn(move || { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .expect("test runtime should build"); + runtime.block_on(make_future()) + }) + .expect("large-stack test thread should spawn"); + + match handle.join() { + Ok(result) => result, + Err(payload) => std::panic::resume_unwind(payload), + } +} + pub(super) const OPERATIONAL_ADMIN_DEVICE_ID: &str = "device-operational-admin"; pub(super) async fn start_authenticated_operational_server( diff --git a/apps/aether-gateway/src/tests/video/mod.rs b/apps/aether-gateway/src/tests/video/mod.rs index a93ac3b2d..61b94cd08 100644 --- a/apps/aether-gateway/src/tests/video/mod.rs +++ b/apps/aether-gateway/src/tests/video/mod.rs @@ -36,6 +36,7 @@ mod openai_sync_task; mod registry_poller; mod routing; mod stream; +mod xai; /// Seed online manual proxy nodes for video execution fixtures. /// @@ -44,6 +45,17 @@ mod stream; /// the same deployment-state record; the loopback URL is never contacted when /// the execution-runtime override is active. pub(super) fn video_proxy_node_repository(node_ids: I) -> Arc +where + I: IntoIterator, + S: AsRef, +{ + video_proxy_node_repository_at_url(node_ids, "http://127.0.0.1:1") +} + +pub(super) fn video_proxy_node_repository_at_url( + node_ids: I, + proxy_url: &str, +) -> Arc where I: IntoIterator, S: AsRef, @@ -68,7 +80,7 @@ where 1, ) .expect("video test proxy node should build") - .with_manual_proxy_fields(Some("http://127.0.0.1:1".to_string()), None, None) + .with_manual_proxy_fields(Some(proxy_url.to_string()), None, None) .with_tunnel_generation(format!("video-test-generation-{node_id}")) }); Arc::new(InMemoryProxyNodeRepository::seed(nodes)) @@ -86,6 +98,28 @@ pub(super) fn video_provider_catalog_repository( endpoint_base_url: &str, key_id: &str, upstream_api_key: &str, +) -> Arc { + video_provider_catalog_repository_with_proxy( + provider_id, + provider_type, + endpoint_id, + api_format, + endpoint_base_url, + key_id, + upstream_api_key, + None, + ) +} + +pub(super) fn video_provider_catalog_repository_with_proxy( + provider_id: &str, + provider_type: &str, + endpoint_id: &str, + api_format: &str, + endpoint_base_url: &str, + key_id: &str, + upstream_api_key: &str, + proxy: Option, ) -> Arc { fn seal_bound_credential( provider_id: &str, @@ -117,7 +151,7 @@ pub(super) fn video_provider_catalog_repository( false, None, Some(2), - None, + proxy, Some(20.0), None, None, diff --git a/apps/aether-gateway/src/tests/video/registry_poller.rs b/apps/aether-gateway/src/tests/video/registry_poller.rs index b47b5938a..35bf4a0c3 100644 --- a/apps/aether-gateway/src/tests/video/registry_poller.rs +++ b/apps/aether-gateway/src/tests/video/registry_poller.rs @@ -13,7 +13,8 @@ use serde_json::json; use super::{ build_state_with_execution_runtime_override, start_server, video_provider_catalog_repository, - AppState, VideoTaskTruthSourceMode, + video_provider_catalog_repository_with_proxy, video_proxy_node_repository_at_url, AppState, + VideoTaskTruthSourceMode, }; fn sample_due_openai_task(upstream_base_url: &str) -> UpsertVideoTask { @@ -279,13 +280,13 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep ); let (upstream_url, upstream_handle) = start_server(upstream).await; - let upstream_api_root = format!("{upstream_url}/v1"); + let upstream_api_root = "http://video-provider.invalid/v1".to_string(); let repository = Arc::new(InMemoryVideoTaskRepository::default()); repository .upsert(sample_due_openai_task(&upstream_api_root)) .await .expect("task upsert should succeed"); - let provider_catalog_repository = video_provider_catalog_repository( + let provider_catalog_repository = video_provider_catalog_repository_with_proxy( "provider-openai-video-local-1", "openai", "endpoint-openai-video-local-1", @@ -293,6 +294,7 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep &upstream_api_root, "key-openai-video-local-1", "sk-upstream-openai-video", + Some(json!({"enabled":true,"node_id":"poller-video-proxy"})), ); let gateway_state = AppState::new() @@ -302,7 +304,7 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep Arc::clone(&repository), provider_catalog_repository, DEVELOPMENT_ENCRYPTION_KEY, - ), + ).attach_proxy_node_repository_for_tests(video_proxy_node_repository_at_url(["poller-video-proxy"], &upstream_url)), ) .with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative) .with_video_task_poller_config(std::time::Duration::from_millis(25), 8); diff --git a/apps/aether-gateway/src/tests/video/routing.rs b/apps/aether-gateway/src/tests/video/routing.rs index 07f551281..ba787c274 100644 --- a/apps/aether-gateway/src/tests/video/routing.rs +++ b/apps/aether-gateway/src/tests/video/routing.rs @@ -14,8 +14,7 @@ use crate::constants::{ use super::{build_router, start_server}; #[tokio::test] -async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when_execution_runtime_missing( -) { +async fn gateway_hides_video_task_from_unauthenticated_caller_with_opt_in_headers() { let execute_hits = Arc::new(Mutex::new(0usize)); let execute_hits_clone = Arc::clone(&execute_hits); let public_hits = Arc::new(Mutex::new(0usize)); @@ -66,13 +65,9 @@ async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(response.status(), StatusCode::NOT_FOUND); let payload: serde_json::Value = response.json().await.expect("body should parse"); - assert_eq!(payload["error"]["type"], "http_error"); - assert_eq!( - payload["error"]["message"], - "当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload, crate::video_tasks::not_found_body()); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); @@ -81,8 +76,7 @@ async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when } #[tokio::test] -async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_execution_runtime_missing( -) { +async fn gateway_hides_video_task_without_calling_public_or_control_upstream() { let execute_hits = Arc::new(Mutex::new(0usize)); let execute_hits_clone = Arc::clone(&execute_hits); let public_hits = Arc::new(Mutex::new(0usize)); @@ -142,13 +136,9 @@ async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_ex .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(response.status(), StatusCode::NOT_FOUND); let payload: serde_json::Value = response.json().await.expect("body should parse"); - assert_eq!(payload["error"]["type"], "http_error"); - assert_eq!( - payload["error"]["message"], - "当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload, crate::video_tasks::not_found_body()); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!( @@ -165,7 +155,7 @@ async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_ex } #[tokio::test] -async fn gateway_skips_video_get_control_sync_without_opt_in_header() { +async fn gateway_hides_video_task_from_unauthenticated_caller_without_opt_in_headers() { let execute_hits = Arc::new(Mutex::new(0usize)); let execute_hits_clone = Arc::clone(&execute_hits); let public_hits = Arc::new(Mutex::new(0usize)); @@ -211,13 +201,9 @@ async fn gateway_skips_video_get_control_sync_without_opt_in_header() { .await .expect("request should succeed"); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(response.status(), StatusCode::NOT_FOUND); let payload: serde_json::Value = response.json().await.expect("body should parse"); - assert_eq!(payload["error"]["type"], "http_error"); - assert_eq!( - payload["error"]["message"], - "当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload, crate::video_tasks::not_found_body()); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); diff --git a/apps/aether-gateway/src/tests/video/xai.rs b/apps/aether-gateway/src/tests/video/xai.rs new file mode 100644 index 000000000..059b4b844 --- /dev/null +++ b/apps/aether-gateway/src/tests/video/xai.rs @@ -0,0 +1,427 @@ +use super::*; +use aether_data::repository::auth::{ + InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, +}; +use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository; +use aether_data::repository::candidates::InMemoryRequestCandidateRepository; +use aether_data_contracts::repository::candidate_selection::{ + StoredMinimalCandidateSelectionRow, StoredProviderModelMapping, +}; +use sha2::{Digest, Sha256}; +use std::sync::atomic::{AtomicUsize, Ordering}; + +fn sample_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot { + StoredAuthApiKeySnapshot::new( + user_id.to_string(), + "video-user".to_string(), + Some("video@example.com".to_string()), + "user".to_string(), + "local".to_string(), + true, + false, + Some(json!(["openai"])), + Some(json!(["openai:video"])), + Some(json!(["video-model"])), + api_key_id.to_string(), + Some("default".to_string()), + true, + false, + false, + Some(60), + Some(5), + Some(4_102_444_800), + Some(json!(["openai"])), + Some(json!(["openai:video"])), + Some(json!(["video-model"])), + ) + .expect("auth snapshot should build") +} + +fn sample_candidate_row() -> StoredMinimalCandidateSelectionRow { + StoredMinimalCandidateSelectionRow { + provider_id: "provider-openai-video-local-1".to_string(), + provider_name: "openai".to_string(), + provider_type: "xai".to_string(), + provider_priority: 10, + provider_is_active: true, + endpoint_id: "endpoint-openai-video-local-1".to_string(), + endpoint_api_format: "openai:video".to_string(), + endpoint_api_family: Some("openai".to_string()), + endpoint_kind: Some("video".to_string()), + endpoint_is_active: true, + key_id: "key-openai-video-local-1".to_string(), + key_name: "prod".to_string(), + key_auth_type: "api_key".to_string(), + key_is_active: true, + key_api_formats: Some(vec!["openai:video".to_string()]), + key_allowed_models: None, + key_capabilities: None, + key_internal_priority: 5, + key_global_priority_by_format: Some(json!({"openai:video": 1})), + model_id: "model-openai-video-local-1".to_string(), + global_model_id: "global-model-openai-video-local-1".to_string(), + global_model_name: "video-model".to_string(), + global_model_mappings: None, + global_model_supports_streaming: Some(false), + model_provider_model_name: "grok-imagine-video".to_string(), + model_provider_model_mappings: Some(vec![StoredProviderModelMapping { + name: "grok-imagine-video".to_string(), + priority: 1, + api_formats: Some(vec!["openai:video".to_string()]), + endpoint_ids: None, + operations: None, + }]), + model_supports_streaming: Some(false), + model_is_active: true, + model_is_available: true, + } +} + +#[tokio::test] +async fn xai_video_native_and_compatibility_http_lifecycle() { + Box::pin(assert_xai_video_http_lifecycle(Arc::new( + InMemoryVideoTaskRepository::default(), + ))) + .await; +} + +#[tokio::test] +async fn xai_video_native_and_compatibility_http_lifecycle_postgres() { + let configured_database_url = std::env::var("AETHER_TEST_DATABASE_URL").ok(); + let managed_database = if configured_database_url.is_none() { + Some( + aether_testkit::ManagedPostgresServer::start() + .await + .expect("temporary PostgreSQL should start"), + ) + } else { + None + }; + let database_url = configured_database_url.unwrap_or_else(|| { + managed_database + .as_ref() + .expect("managed test database should exist") + .database_url() + .to_string() + }); + let pool = sqlx::postgres::PgPoolOptions::new() + .max_connections(1) + .connect(&database_url) + .await + .expect("test database should connect"); + aether_data::driver::postgres::run_migrations(&pool) + .await + .expect("test database should migrate"); + // Preserve the production column constraints and unique indexes while isolating test rows. + sqlx::query("CREATE TEMP TABLE video_tasks (LIKE public.video_tasks INCLUDING ALL)") + .execute(&pool) + .await + .expect("isolated video task table should be created"); + let repository = + Arc::new(aether_data::repository::video_tasks::SqlxVideoTaskRepository::new(pool.clone())); + Box::pin(assert_xai_video_http_lifecycle(repository)).await; + pool.close().await; +} + +async fn assert_xai_video_http_lifecycle(repository: Arc) +where + T: aether_data_contracts::repository::video_tasks::VideoTaskRepository + 'static, +{ + let static_dir = std::env::temp_dir().join(format!( + "aether-xai-video-static-{}", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos() + )); + std::fs::create_dir_all(&static_dir).unwrap(); + std::fs::write( + static_dir.join("index.html"), + "Aether test frontend", + ) + .unwrap(); + let seen = Arc::new(Mutex::new(Vec::::new())); + let calls = Arc::new(AtomicUsize::new(0)); + // Exercise the real HTTP executor, including production method gates, instead of + // the test execution-runtime override that used to hide rejected GET requests. + let video_url = Arc::new(Mutex::new(String::new())); + let runtime = Router::new() + .route("/v1/videos/{operation}", any({ + let seen = seen.clone(); + let calls = calls.clone(); + let video_url = video_url.clone(); + move |request: Request| { + let seen = seen.clone(); + let calls = calls.clone(); + let video_url = video_url.clone(); + async move { + let (parts, body) = request.into_parts(); + assert_eq!(parts.headers["authorization"], "Bearer upstream-video-key"); + let bytes = to_bytes(body, usize::MAX).await.unwrap(); + let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap_or(json!(null)); + seen.lock().unwrap().push(json!({ + "method": parts.method.as_str(), + "url": parts.uri.path(), + "body": {"json_body": body} + })); + let response = if parts.method == http::Method::POST { + json!({"request_id":"upstream-video-id", "provider_extension":{"accepted":true}}) + } else { + assert_eq!(parts.uri.path(), "/v1/videos/upstream-video-id"); + if calls.fetch_add(1, Ordering::SeqCst) == 0 { + json!({"status":"pending"}) + } else { + json!({"status":"done", "model":"grok-imagine-video", "video":{"url":video_url.lock().unwrap().clone(), "duration":6, "respect_moderation":true}, "provider_extension":"preserved"}) + } + }; + Json(response) + } + } + })) + .route("/test.mp4", any(|request: Request| async move { + assert!(request.headers().get("authorization").is_none()); + assert!(request.headers().get("x-xai-token-auth").is_none()); + ([("content-type", "video/mp4")], "test-video-bytes") + })); + let (runtime_url, runtime_handle) = start_server(runtime).await; + let expected_video_url = format!("{runtime_url}/test.mp4"); + *video_url.lock().unwrap() = expected_video_url.clone(); + let state_factory = || { + let auth = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![ + ( + Some(format!("{:x}", Sha256::digest(b"owner-key"))), + sample_auth_snapshot("owner-api-key", "owner"), + ), + ( + Some(format!("{:x}", Sha256::digest(b"foreign-key"))), + sample_auth_snapshot("foreign-api-key", "foreign"), + ), + ])); + let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![ + sample_candidate_row(), + ])); + let catalog = video_provider_catalog_repository_with_proxy( + "provider-openai-video-local-1", + "xai", + "endpoint-openai-video-local-1", + "openai:video", + "http://video-provider.invalid/v1", + "key-openai-video-local-1", + "upstream-video-key", + Some(json!({"enabled":true,"node_id":"video-proxy"})), + ); + AppState::new().expect("gateway should build").with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative).with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests( + auth, candidates, catalog, Arc::new(InMemoryRequestCandidateRepository::default()), DEVELOPMENT_ENCRYPTION_KEY + ).attach_video_task_repository_for_tests(repository.clone()) + .attach_proxy_node_repository_for_tests(video_proxy_node_repository_at_url(["video-proxy"], &runtime_url)) + ) + }; + let router_factory = + || crate::attach_static_frontend(build_router_with_state(state_factory()), &static_dir); + let (gateway_url, gateway_handle) = start_server(router_factory()).await; + let client = reqwest::Client::new(); + assert_eq!( + client + .get(&gateway_url) + .send() + .await + .unwrap() + .text() + .await + .unwrap(), + "Aether test frontend" + ); + for (path, native) in [ + ("/v1/videos/generations", true), + ("/v1/videos", true), + ("/v1/videos/edits", true), + ("/v1/videos/extensions", true), + ("/openai/v1/videos", false), + ] { + calls.store(0, Ordering::SeqCst); + let body = if native { + json!({"model":"video-model","prompt":"A cat","duration":6,"aspect_ratio":"1:1","video":{"url":"https://example.com/input.mp4"},"future_option":true}) + } else { + json!({"model":"video-model","prompt":"A cat","seconds":"6","size":"1280x720"}) + }; + let response = client + .post(format!("{gateway_url}{path}")) + .bearer_auth("owner-key") + .json(&body) + .send() + .await + .unwrap(); + let status = response.status(); + let result: serde_json::Value = response.json().await.unwrap(); + assert_eq!(status, StatusCode::OK, "{path}: {result}"); + let id = result[if native { "request_id" } else { "id" }] + .as_str() + .unwrap(); + assert_ne!(id, "upstream-video-id"); + if native { + assert!(result.get("id").is_none()); + assert_eq!(result["provider_extension"]["accepted"], true); + } else { + assert_eq!(result["status"], "queued"); + } + let request = seen.lock().unwrap().last().unwrap().clone(); + let suffix = if path.ends_with("/edits") { + "edits" + } else if path.ends_with("/extensions") { + "extensions" + } else { + "generations" + }; + assert_eq!(request["url"], format!("/v1/videos/{suffix}")); + assert_eq!(request["body"]["json_body"]["model"], "grok-imagine-video"); + assert_eq!(request["body"]["json_body"]["duration"], 6); + if native { + assert_eq!(request["body"]["json_body"]["future_option"], true); + } else { + assert_eq!(request["body"]["json_body"]["aspect_ratio"], "16:9"); + assert_eq!(request["body"]["json_body"]["resolution"], "720p"); + assert!(request["body"]["json_body"].get("seconds").is_none()); + assert!(request["body"]["json_body"].get("size").is_none()); + } + let query = format!( + "{gateway_url}{}/{id}", + if native { + "/v1/videos" + } else { + "/openai/v1/videos" + } + ); + let before = seen.lock().unwrap().len(); + let denied = client + .get(&query) + .bearer_auth("foreign-key") + .send() + .await + .unwrap(); + assert_eq!(denied.status(), StatusCode::NOT_FOUND); + assert_eq!(seen.lock().unwrap().len(), before); + let denied_content = client + .get(format!("{gateway_url}/openai/v1/videos/{id}/content")) + .bearer_auth("foreign-key") + .send() + .await + .unwrap(); + assert_eq!(denied_content.status(), StatusCode::NOT_FOUND); + assert_eq!(seen.lock().unwrap().len(), before); + let pending: serde_json::Value = client + .get(&query) + .bearer_auth("owner-key") + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!( + pending["status"], + if native { "pending" } else { "queued" }, + "{path}: {pending}" + ); + let done: serde_json::Value = client + .get(&query) + .bearer_auth("owner-key") + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(done["status"], if native { "done" } else { "completed" }); + if native { + assert_eq!(done["video"]["respect_moderation"], true); + assert_eq!(done["provider_extension"], "preserved"); + } else { + assert_eq!(done["video_url"], expected_video_url); + } + let stored = repository + .find(VideoTaskLookupKey::Id(id)) + .await + .unwrap() + .unwrap(); + assert_eq!( + stored.client_api_format.as_deref(), + Some(if native { "xai:video" } else { "openai:video" }) + ); + assert_eq!( + stored.external_task_id.as_deref(), + Some("upstream-video-id") + ); + assert!(stored.request_metadata.is_none()); + assert!(stored.original_request_body.is_none()); + // A new gateway instance must reconstruct the pinned provider/credential and protocol. + let (restart_url, restart_handle) = start_server(router_factory()).await; + let restored: serde_json::Value = client + .get(format!( + "{restart_url}{}/{id}", + if native { + "/v1/videos" + } else { + "/openai/v1/videos" + } + )) + .bearer_auth("owner-key") + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(restored["status"], done["status"]); + if native { + assert_eq!(restored["video"]["respect_moderation"], true); + } + let compat: serde_json::Value = client + .get(format!("{restart_url}/openai/v1/videos/{id}")) + .bearer_auth("owner-key") + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(compat["status"], "completed"); + assert_eq!(compat["video_url"], expected_video_url); + let native_view: serde_json::Value = client + .get(format!("{restart_url}/v1/videos/{id}")) + .bearer_auth("owner-key") + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(native_view["status"], "done"); + assert_eq!(native_view["video"]["respect_moderation"], true); + for prefix in ["/v1/videos", "/openai/v1/videos"] { + let content = client + .get(format!("{restart_url}{prefix}/{id}/content")) + .bearer_auth("owner-key") + .send() + .await + .unwrap(); + assert_eq!(content.status(), StatusCode::OK); + assert_eq!(content.headers()["content-type"], "video/mp4"); + assert_eq!(content.bytes().await.unwrap(), "test-video-bytes"); + } + restart_handle.abort(); + } + let before = seen.lock().unwrap().len(); + let bad = client + .post(format!("{gateway_url}/openai/v1/videos")) + .bearer_auth("owner-key") + .json(&json!({"model":"video-model","prompt":"cat","seconds":"wrong"})) + .send() + .await + .unwrap(); + assert_eq!(bad.status(), StatusCode::BAD_REQUEST); + assert_eq!(seen.lock().unwrap().len(), before); + gateway_handle.abort(); + runtime_handle.abort(); + std::fs::remove_dir_all(&static_dir).unwrap(); +} diff --git a/apps/aether-gateway/src/video_tasks/tests/plans.rs b/apps/aether-gateway/src/video_tasks/tests/plans.rs index e40ec578a..154910309 100644 --- a/apps/aether-gateway/src/video_tasks/tests/plans.rs +++ b/apps/aether-gateway/src/video_tasks/tests/plans.rs @@ -11,6 +11,9 @@ use super::{ fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, @@ -92,6 +95,9 @@ fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() { fn rust_authoritative_service_builds_openai_remix_follow_up_plan() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, @@ -177,6 +183,9 @@ fn rust_authoritative_service_builds_openai_remix_follow_up_plan() { fn rust_authoritative_service_builds_openai_delete_follow_up_plan() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, @@ -332,6 +341,9 @@ fn rust_authoritative_service_builds_gemini_cancel_follow_up_plan() { fn rust_authoritative_service_builds_openai_read_refresh_plan() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, @@ -407,6 +419,9 @@ fn rust_authoritative_service_builds_gemini_read_refresh_plan() { fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-active-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, @@ -428,6 +443,9 @@ fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only() transport: sample_transport("https://api.openai.example", "openai:video"), })); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-completed-123".to_string(), upstream_task_id: "ext-video-task-999".to_string(), created_at_unix_ms: 1712345678, @@ -471,6 +489,9 @@ fn file_video_task_store_persists_snapshots_across_service_rebuilds() { ) .expect("file-backed service should build"); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-file-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, diff --git a/apps/aether-gateway/src/video_tasks/tests/projection.rs b/apps/aether-gateway/src/video_tasks/tests/projection.rs index fa7978bfd..fed35f5d1 100644 --- a/apps/aether-gateway/src/video_tasks/tests/projection.rs +++ b/apps/aether-gateway/src/video_tasks/tests/projection.rs @@ -10,6 +10,9 @@ use super::{ fn rust_authoritative_service_projects_openai_status_into_local_read_response() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, @@ -93,6 +96,9 @@ fn rust_authoritative_service_projects_openai_status_into_local_read_response() fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_video_url() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, @@ -159,6 +165,9 @@ fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_vide fn rust_authoritative_service_returns_processing_content_response_for_pending_openai_task() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, diff --git a/apps/aether-gateway/src/video_tasks/tests/sync.rs b/apps/aether-gateway/src/video_tasks/tests/sync.rs index 9634d322a..7d85ef8d8 100644 --- a/apps/aether-gateway/src/video_tasks/tests/sync.rs +++ b/apps/aether-gateway/src/video_tasks/tests/sync.rs @@ -218,6 +218,9 @@ fn rust_authoritative_video_truth_source_can_background_success_report() { fn rust_authoritative_service_reads_openai_task_from_local_registry() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, @@ -266,6 +269,9 @@ fn rust_authoritative_service_reads_openai_task_from_local_registry() { fn rust_authoritative_service_applies_cancel_and_delete_mutations() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, diff --git a/apps/aether-gateway/src/tests/architecture/admin_billing.rs b/apps/aether-gateway/tests/architecture/admin_billing.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/admin_billing.rs rename to apps/aether-gateway/tests/architecture/admin_billing.rs diff --git a/apps/aether-gateway/src/tests/architecture/admin_model.rs b/apps/aether-gateway/tests/architecture/admin_model.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/admin_model.rs rename to apps/aether-gateway/tests/architecture/admin_model.rs diff --git a/apps/aether-gateway/src/tests/architecture/admin_observability.rs b/apps/aether-gateway/tests/architecture/admin_observability.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/admin_observability.rs rename to apps/aether-gateway/tests/architecture/admin_observability.rs diff --git a/apps/aether-gateway/src/tests/architecture/admin_provider.rs b/apps/aether-gateway/tests/architecture/admin_provider.rs similarity index 99% rename from apps/aether-gateway/src/tests/architecture/admin_provider.rs rename to apps/aether-gateway/tests/architecture/admin_provider.rs index 9c27d423d..5e95f5c57 100644 --- a/apps/aether-gateway/src/tests/architecture/admin_provider.rs +++ b/apps/aether-gateway/tests/architecture/admin_provider.rs @@ -1789,6 +1789,7 @@ fn admin_provider_oauth_quota_mod_stays_thin() { "pub(crate) mod dispatch;", "pub(crate) mod kiro;", "pub(crate) mod shared;", + "pub(crate) mod xai;", ] { assert!( quota_mod.contains(pattern), @@ -1861,6 +1862,7 @@ fn admin_provider_oauth_quota_mod_stays_thin() { "refresh_antigravity_provider_quota_locally", "refresh_gemini_cli_provider_quota_locally", "refresh_chatgpt_web_provider_quota_locally", + "refresh_xai_provider_quota_locally", ] { assert!( quota_dispatch.contains(pattern), diff --git a/apps/aether-gateway/src/tests/architecture/admin_shared.rs b/apps/aether-gateway/tests/architecture/admin_shared.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/admin_shared.rs rename to apps/aether-gateway/tests/architecture/admin_shared.rs diff --git a/apps/aether-gateway/src/tests/architecture/admin_system.rs b/apps/aether-gateway/tests/architecture/admin_system.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/admin_system.rs rename to apps/aether-gateway/tests/architecture/admin_system.rs diff --git a/apps/aether-gateway/src/tests/architecture/admin_users.rs b/apps/aether-gateway/tests/architecture/admin_users.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/admin_users.rs rename to apps/aether-gateway/tests/architecture/admin_users.rs diff --git a/apps/aether-gateway/src/tests/architecture/ai_serving.rs b/apps/aether-gateway/tests/architecture/ai_serving.rs similarity index 99% rename from apps/aether-gateway/src/tests/architecture/ai_serving.rs rename to apps/aether-gateway/tests/architecture/ai_serving.rs index b5ef1144f..555a460e7 100644 --- a/apps/aether-gateway/src/tests/architecture/ai_serving.rs +++ b/apps/aether-gateway/tests/architecture/ai_serving.rs @@ -1194,7 +1194,7 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() { let candidate_resolution = read_workspace_file("apps/aether-gateway/src/ai_serving/planner/candidate_resolution.rs"); - let ranking_call = candidate_resolution + candidate_resolution .find("rank_eligible_local_execution_candidates(") .expect("candidate_resolution.rs should call core-backed local candidate ranking"); assert!( @@ -1452,7 +1452,8 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() { "GeminiCliProviderPoolAdapter", "KiroProviderPoolAdapter", "ChatGptWebProviderPoolAdapter", - "CLAUDE_CODE_PROVIDER_POOL_ADAPTER", + "XaiProviderPoolAdapter", + "ClaudeCodeProviderPoolAdapter", "VERTEX_AI_PROVIDER_POOL_ADAPTER", "provider_types_for_capability", "supports_quota_refresh", @@ -1478,6 +1479,7 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() { "pub mod gemini_cli;", "pub mod kiro;", "pub mod chatgpt_web;", + "pub mod xai;", ] { assert!( provider_pool_providers.contains(pattern), @@ -1513,6 +1515,14 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() { "crates/aether-provider/pool/src/providers/kiro.rs", vec!["KiroProviderPoolAdapter", "quota_exhausted_from_bucket"], ), + ( + "crates/aether-provider/pool/src/providers/xai.rs", + vec![ + "XaiProviderPoolAdapter", + "build_xai_pool_billing_request", + "quota_exhausted_from_bucket", + ], + ), ( "crates/aether-provider/pool/src/providers/chatgpt_web.rs", vec![ @@ -1527,10 +1537,16 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() { "crates/aether-provider/pool/src/providers/unsupported.rs", vec![ "UnsupportedQuotaProviderPoolAdapter", - "CLAUDE_CODE_PROVIDER_POOL_ADAPTER", "VERTEX_AI_PROVIDER_POOL_ADAPTER", ], ), + ( + "crates/aether-provider/pool/src/providers/claude_code.rs", + vec![ + "ClaudeCodeProviderPoolAdapter", + "build_claude_code_pool_quota_request", + ], + ), ] { let source = read_workspace_file(path); for pattern in patterns { @@ -5035,7 +5051,9 @@ fn retired_api_format_occurrences_are_whitelisted() { .expect("file should be under workspace root") .to_string_lossy() .replace('\\', "/"); - if relative == "apps/aether-gateway/src/tests/architecture/ai_serving.rs" { + if relative == "apps/aether-gateway/tests/architecture/ai_serving.rs" + || relative == "apps/aether-gateway/src/tests/architecture/ai_serving.rs" + { continue; } diff --git a/apps/aether-gateway/src/tests/architecture/mod.rs b/apps/aether-gateway/tests/architecture/mod.rs similarity index 89% rename from apps/aether-gateway/src/tests/architecture/mod.rs rename to apps/aether-gateway/tests/architecture/mod.rs index 374e1074b..57ddab353 100644 --- a/apps/aether-gateway/src/tests/architecture/mod.rs +++ b/apps/aether-gateway/tests/architecture/mod.rs @@ -1,7 +1,9 @@ use std::fs; use std::path::{Path, PathBuf}; -pub(super) fn collect_rust_files(root: &Path, files: &mut Vec) { +// 架构守卫在独立 integration test 中是顶层模块;helper 统一 pub(crate), +// 子模块经 `use super::*` / `use super::{...}` 访问(与原 lib 内布局一致)。 +pub(crate) fn collect_rust_files(root: &Path, files: &mut Vec) { for entry in fs::read_dir(root).expect("directory should be readable") { let entry = entry.expect("directory entry should be readable"); let path = entry.path(); @@ -15,7 +17,7 @@ pub(super) fn collect_rust_files(root: &Path, files: &mut Vec) { } } -pub(super) fn assert_no_sqlx_queries(root_relative_path: &str) { +pub(crate) fn assert_no_sqlx_queries(root_relative_path: &str) { let root = Path::new(env!("CARGO_MANIFEST_DIR")).join(root_relative_path); let mut files = Vec::new(); collect_rust_files(&root, &mut files); @@ -80,7 +82,7 @@ fn sql_pool_scan_distinguishes_pool_types_from_repository_names() { )); } -pub(super) fn assert_no_sensitive_log_patterns(root_relative_path: &str, patterns: &[&str]) { +pub(crate) fn assert_no_sensitive_log_patterns(root_relative_path: &str, patterns: &[&str]) { let root = Path::new(env!("CARGO_MANIFEST_DIR")).join(root_relative_path); let mut files = Vec::new(); collect_rust_files(&root, &mut files); @@ -109,7 +111,7 @@ pub(super) fn assert_no_sensitive_log_patterns(root_relative_path: &str, pattern ); } -pub(super) fn assert_no_module_dependency_patterns(root_relative_path: &str, patterns: &[&str]) { +pub(crate) fn assert_no_module_dependency_patterns(root_relative_path: &str, patterns: &[&str]) { let root = Path::new(env!("CARGO_MANIFEST_DIR")).join(root_relative_path); let mut files = Vec::new(); collect_rust_files(&root, &mut files); @@ -138,14 +140,14 @@ pub(super) fn assert_no_module_dependency_patterns(root_relative_path: &str, pat ); } -pub(super) fn workspace_file_exists(root_relative_path: &str) -> bool { +pub(crate) fn workspace_file_exists(root_relative_path: &str) -> bool { Path::new(env!("CARGO_MANIFEST_DIR")) .join("../..") .join(root_relative_path) .exists() } -pub(super) fn workspace_files_with_extension( +pub(crate) fn workspace_files_with_extension( root_relative_path: &str, extension: &str, ) -> Vec { @@ -162,7 +164,7 @@ pub(super) fn workspace_files_with_extension( files } -pub(super) fn collect_workspace_rust_files(root_relative_path: &str) -> Vec { +pub(crate) fn collect_workspace_rust_files(root_relative_path: &str) -> Vec { let root = Path::new(env!("CARGO_MANIFEST_DIR")) .join("../..") .join(root_relative_path); @@ -172,7 +174,7 @@ pub(super) fn collect_workspace_rust_files(root_relative_path: &str) -> Vec String { +pub(crate) fn read_workspace_file(path: &str) -> String { let workspace_root = Path::new(env!("CARGO_MANIFEST_DIR")) .join("../..") .canonicalize() @@ -180,7 +182,7 @@ pub(super) fn read_workspace_file(path: &str) -> String { fs::read_to_string(workspace_root.join(path)).expect("source file should be readable") } -pub(super) fn read_workspace_module_tree(path: &str) -> String { +pub(crate) fn read_workspace_module_tree(path: &str) -> String { let workspace_root = Path::new(env!("CARGO_MANIFEST_DIR")) .join("../..") .canonicalize() diff --git a/apps/aether-gateway/src/tests/architecture/runtime_and_security.rs b/apps/aether-gateway/tests/architecture/runtime_and_security.rs similarity index 99% rename from apps/aether-gateway/src/tests/architecture/runtime_and_security.rs rename to apps/aether-gateway/tests/architecture/runtime_and_security.rs index 8772263c8..4cff5b71c 100644 --- a/apps/aether-gateway/src/tests/architecture/runtime_and_security.rs +++ b/apps/aether-gateway/tests/architecture/runtime_and_security.rs @@ -1,4 +1,4 @@ -use std::path::{Path, PathBuf}; +use std::path::Path; use super::*; diff --git a/apps/aether-gateway/src/tests/architecture/sql_and_data.rs b/apps/aether-gateway/tests/architecture/sql_and_data.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/sql_and_data.rs rename to apps/aether-gateway/tests/architecture/sql_and_data.rs diff --git a/apps/aether-gateway/src/tests/architecture/usage.rs b/apps/aether-gateway/tests/architecture/usage.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/usage.rs rename to apps/aether-gateway/tests/architecture/usage.rs diff --git a/apps/aether-gateway/src/tests/architecture/workspace_tiers.rs b/apps/aether-gateway/tests/architecture/workspace_tiers.rs similarity index 100% rename from apps/aether-gateway/src/tests/architecture/workspace_tiers.rs rename to apps/aether-gateway/tests/architecture/workspace_tiers.rs diff --git a/apps/aether-gateway/tests/architecture_guard.rs b/apps/aether-gateway/tests/architecture_guard.rs new file mode 100644 index 000000000..ebe87041c --- /dev/null +++ b/apps/aether-gateway/tests/architecture_guard.rs @@ -0,0 +1,5 @@ +//! 架构守卫独立测试目标。 +//! +//! 从 lib 的 `cfg(test)` 巨型编译单元迁出:只做源码/manifest 字符串断言, +//! 不启动 AppState、不依赖 gateway 私有类型,用于压低 lib test 编译面与 rustc 峰值。 +mod architecture; diff --git a/apps/aether-tunnel/Cargo.toml b/apps/aether-tunnel/Cargo.toml index 330b7127a..480f4f267 100644 --- a/apps/aether-tunnel/Cargo.toml +++ b/apps/aether-tunnel/Cargo.toml @@ -46,5 +46,4 @@ webpki-roots = "0.26" uuid.workspace = true [dev-dependencies] -aether-gateway = { workspace = true, features = ["testkit"] } tokio = { version = "1", features = ["test-util"] } diff --git a/apps/aether-tunnel/src/config.rs b/apps/aether-tunnel/src/config.rs index 5c45c495b..0cf824188 100644 --- a/apps/aether-tunnel/src/config.rs +++ b/apps/aether-tunnel/src/config.rs @@ -885,7 +885,7 @@ impl Config { Ok(Duration::from_millis(self.tunnel_connect_timeout_ms)) } - pub fn tunnel_ip_family(&self) -> crate::egress_proxy::IpFamily { + pub(crate) fn tunnel_ip_family(&self) -> crate::egress_proxy::IpFamily { if self.tunnel_ipv4_only { crate::egress_proxy::IpFamily::Ipv4Only } else if self.tunnel_ipv6_only { diff --git a/apps/aether-tunnel/src/hardware.rs b/apps/aether-tunnel/src/hardware.rs index d0e9a25fc..01192c96c 100644 --- a/apps/aether-tunnel/src/hardware.rs +++ b/apps/aether-tunnel/src/hardware.rs @@ -116,6 +116,13 @@ impl RuntimeResourceMonitor { } } +impl Default for RuntimeResourceMonitor { + fn default() -> Self { + // 默认构造与显式 new 保持一致,便于库目标和二进制目标共用监控器。 + Self::new() + } +} + /// Collect hardware information and estimate max concurrency. /// /// Should be called once at startup -- hardware does not change at runtime. diff --git a/apps/aether-tunnel/src/lib.rs b/apps/aether-tunnel/src/lib.rs new file mode 100644 index 000000000..98ef478a2 --- /dev/null +++ b/apps/aether-tunnel/src/lib.rs @@ -0,0 +1,17 @@ +#![allow(clippy::large_enum_variant)] + +// Tunnel 的运行模块作为库暴露给独立集成测试使用;生产二进制仍由 +// src/main.rs 负责命令行解析,避免端到端测试把 Gateway dev-dependency +// 带进 Workspace Rest 的默认测试目标。 +pub mod app; +pub mod config; +pub mod egress_proxy; +pub mod hardware; +mod net; +pub mod registration; +pub mod runtime; +pub mod setup; +pub mod state; +pub mod target_filter; +pub mod tunnel; +pub mod upstream_client; diff --git a/apps/aether-tunnel/src/main.rs b/apps/aether-tunnel/src/main.rs index d2b01e23d..9ea67bf51 100644 --- a/apps/aether-tunnel/src/main.rs +++ b/apps/aether-tunnel/src/main.rs @@ -1,20 +1,8 @@ #![allow(clippy::large_enum_variant)] -mod app; -mod config; -mod egress_proxy; -mod hardware; -mod net; -mod registration; -mod runtime; -mod setup; -mod state; -mod target_filter; -mod tunnel; -mod upstream_client; - use std::path::PathBuf; +use aether_tunnel::{app, config, setup}; use clap::{parser::ValueSource, CommandFactory, FromArgMatches, Parser}; use config::{Config, ServerEntry, TunnelSecurity}; diff --git a/apps/aether-tunnel/src/setup/mod.rs b/apps/aether-tunnel/src/setup/mod.rs index 3c9628c7c..0437863fb 100644 --- a/apps/aether-tunnel/src/setup/mod.rs +++ b/apps/aether-tunnel/src/setup/mod.rs @@ -1,5 +1,5 @@ -pub(crate) mod service; +pub mod service; mod tui; -pub(crate) mod upgrade; +pub mod upgrade; pub use self::tui::{run, SetupOutcome}; diff --git a/apps/aether-tunnel/src/state.rs b/apps/aether-tunnel/src/state.rs index 26b54a42d..df5959e28 100644 --- a/apps/aether-tunnel/src/state.rs +++ b/apps/aether-tunnel/src/state.rs @@ -226,6 +226,13 @@ impl TunnelRequestMetrics { } } +impl Default for TunnelRequestMetrics { + fn default() -> Self { + // 指标初始值全部为零,Default 与现有 new 语义完全一致。 + Self::new() + } +} + const RECENT_TUNNEL_ERROR_CAPACITY: usize = 64; const TUNNEL_ERROR_CATEGORY_MAX_CHARS: usize = 48; const TUNNEL_ERROR_MESSAGE_MAX_CHARS: usize = 320; @@ -534,6 +541,13 @@ impl TunnelMetrics { } } +impl Default for TunnelMetrics { + fn default() -> Self { + // 保留 recent_errors 的容量初始化,避免 Default 改变错误环形缓存行为。 + Self::new() + } +} + fn now_unix_secs() -> u64 { now_unix_ms() / 1_000 } diff --git a/apps/aether-tunnel/src/tunnel/mod.rs b/apps/aether-tunnel/src/tunnel/mod.rs index 810899bad..b0c46986f 100644 --- a/apps/aether-tunnel/src/tunnel/mod.rs +++ b/apps/aether-tunnel/src/tunnel/mod.rs @@ -1,10 +1,10 @@ pub mod client; -pub mod dispatcher; -pub mod heartbeat; +mod dispatcher; +mod heartbeat; pub mod protocol; -pub mod stream_handler; +mod stream_handler; mod task; -pub mod writer; +mod writer; use std::sync::Arc; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; @@ -230,34 +230,10 @@ fn mix_u64(mut x: u64) -> u64 { #[cfg(test)] mod tests { - use std::sync::atomic::AtomicU64; - use std::sync::{Arc, Once}; - use std::time::{Duration, SystemTime, UNIX_EPOCH}; - - use aether_contracts::tunnel::{ - sign_tunnel_relay_request, tunnel_relay_payload_digest, TUNNEL_RELAY_AUTH_NONCE_HEADER, - TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, TUNNEL_RELAY_AUTH_SENDER_HEADER, - TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, - TUNNEL_RELAY_OWNER_INSTANCE_HEADER, - }; - use aether_gateway::{build_router_with_state, AppState as GatewayAppState}; - use arc_swap::ArcSwap; - use axum::Router; - use reqwest::StatusCode; - use tokio::sync::watch; - - use crate::config::Config; - use crate::registration::client::AetherClient; - use crate::runtime::DynamicConfig; - use crate::state::{ - AppState as TunnelAppState, ServerContext, TunnelMetrics, TunnelRequestMetrics, - }; - use crate::target_filter::DnsCache; - use crate::tunnel::protocol; - use crate::upstream_client; + use std::time::Duration; use super::{ - compute_reconnect_cap_ms, compute_reconnect_delay, compute_startup_stagger, run, + compute_reconnect_cap_ms, compute_reconnect_delay, compute_startup_stagger, MAX_STARTUP_STAGGER_MS, RECONNECT_PROBE_MAX_DELAY_MS, STARTUP_STAGGER_STEP_MS, }; @@ -299,480 +275,4 @@ mod tests { let d = compute_reconnect_delay(500, 45_000, 100, 12345); assert!(d <= Duration::from_millis(RECONNECT_PROBE_MAX_DELAY_MS)); } - - #[tokio::test] - async fn tunnel_reconnects_after_gateway_restart() { - ensure_rustls_provider(); - - let gateway_port = reserve_local_port().expect("gateway port should reserve"); - let gateway_base_url = format!("http://127.0.0.1:{gateway_port}"); - let (gateway_state, mut gateway_handle) = start_gateway_on_port(gateway_port) - .await - .expect("gateway should start"); - - let mut tunnel_config = sample_config(&gateway_base_url); - tunnel_config.tunnel_security = crate::config::TunnelSecurity::NonTlsRequired; - tunnel_config.tunnel_encryption_key = - Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".to_string()); - let state = sample_state(tunnel_config); - let server = sample_server(&state, "node-recovery"); - let (shutdown_tx, shutdown_rx) = watch::channel(false); - let tunnel_task = tokio::spawn({ - let state = Arc::clone(&state); - let server = Arc::clone(&server); - let (_drain_tx, drain_rx) = watch::channel(false); - async move { - run(&state, &server, 0, shutdown_rx, drain_rx).await; - } - }); - - wait_until_relay_status( - &gateway_base_url, - "node-recovery", - StatusCode::GATEWAY_TIMEOUT, - ) - .await; - - gateway_handle.abort(); - let _ = (&mut gateway_handle).await; - assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1); - - let (_restarted_gateway_state, restarted_gateway_handle) = - start_gateway_on_port_retry(gateway_port) - .await - .expect("gateway should restart on fixed port"); - gateway_handle = restarted_gateway_handle; - - wait_until_relay_status( - &gateway_base_url, - "node-recovery", - StatusCode::GATEWAY_TIMEOUT, - ) - .await; - - assert!(server.tunnel_metrics.snapshot().connect_successes >= 2); - let _ = shutdown_tx.send(true); - tokio::time::timeout(Duration::from_secs(5), tunnel_task) - .await - .expect("tunnel task should stop") - .expect("tunnel task should join"); - gateway_handle.abort(); - } - - async fn wait_until_relay_status(gateway_base_url: &str, node_id: &str, expected: StatusCode) { - let deadline = tokio::time::Instant::now() + Duration::from_secs(10); - let mut last_observed = None::; - loop { - if let Some((status, body)) = probe_relay_status(gateway_base_url, node_id).await { - last_observed = Some(format!("{status} body={body}")); - if status == expected { - return; - } - } - assert!( - tokio::time::Instant::now() < deadline, - "relay status did not become {expected} within timeout; last={:?}", - last_observed - ); - tokio::time::sleep(Duration::from_millis(25)).await; - } - } - - async fn probe_relay_status( - gateway_base_url: &str, - node_id: &str, - ) -> Option<(StatusCode, String)> { - let response = relay_response(gateway_base_url, node_id, relay_probe_envelope()).await?; - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - Some((status, body)) - } - - async fn relay_response( - gateway_base_url: &str, - node_id: &str, - payload: Vec, - ) -> Option { - let timestamp = SystemTime::now() - .duration_since(UNIX_EPOCH) - .expect("test clock should be after epoch") - .as_secs(); - let nonce = uuid::Uuid::new_v4().simple().to_string(); - let digest = tunnel_relay_payload_digest(&payload, &[]); - let signature = sign_tunnel_relay_request( - b"tunnel-reconnect-test-secret-at-least-32-bytes", - "tunnel-reconnect-test-client", - "tunnel-reconnect-test-gateway", - node_id, - "", - false, - timestamp, - &nonce, - &digest, - ); - reqwest::Client::new() - .post(format!( - "{gateway_base_url}/api/internal/tunnel/relay/{node_id}" - )) - .header("content-type", "application/octet-stream") - .header( - TUNNEL_RELAY_AUTH_SENDER_HEADER, - "tunnel-reconnect-test-client", - ) - .header( - TUNNEL_RELAY_OWNER_INSTANCE_HEADER, - "tunnel-reconnect-test-gateway", - ) - .header(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, timestamp) - .header(TUNNEL_RELAY_AUTH_NONCE_HEADER, nonce) - .header( - TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, - digest.encode_header_value(), - ) - .header(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, signature) - .body(payload) - .send() - .await - .ok() - } - - fn relay_probe_envelope() -> Vec { - let meta = protocol::RequestMeta { - provider_id: None, - endpoint_id: None, - key_id: None, - method: "GET".to_string(), - url: "http://127.0.0.1:80/blocked".to_string(), - headers: std::collections::HashMap::new(), - stream: false, - request_timeout_ms: None, - stream_first_byte_timeout_ms: None, - timeout: 5, - follow_redirects: None, - http1_only: false, - transport_profile: None, - }; - let meta_json = - serde_json::to_vec(&meta).expect("tunnel relay probe metadata should serialize"); - let mut envelope = Vec::with_capacity(4 + meta_json.len()); - envelope.extend_from_slice(&(meta_json.len() as u32).to_be_bytes()); - envelope.extend_from_slice(&meta_json); - envelope - } - - async fn start_gateway_on_port( - port: u16, - ) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> { - // The embedded gateway now fails closed when relay authentication is - // not configured. Keep this integration fixture explicitly authenticated. - static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(()); - let state = { - let _guard = ENV_LOCK.lock().unwrap(); - let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET"); - let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID"); - std::env::set_var( - "AETHER_TUNNEL_RELAY_AUTH_SECRET", - "tunnel-reconnect-test-secret-at-least-32-bytes", - ); - std::env::set_var( - "AETHER_GATEWAY_INSTANCE_ID", - "tunnel-reconnect-test-gateway", - ); - let mut state = GatewayAppState::new().expect("gateway test state should build"); - aether_gateway::configure_test_tunnel_security( - &mut state, - "node-recovery", - "test-generation-1", - "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=", - ); - restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret); - restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance); - state - }; - let router = build_router_with_state(state.clone()); - let handle = spawn_router_on_port(port, router).await?; - Ok((state, handle)) - } - - #[tokio::test] - async fn negotiated_small_window_streams_large_responses_and_cancels_idle_upstream() { - use axum::body::{Body, Bytes}; - use axum::routing::get; - use futures_util::StreamExt; - - ensure_rustls_provider(); - let upstream_port = reserve_local_port().unwrap(); - let upstream = Router::new() - .route( - "/large", - get(|| async { Body::from(vec![b'x'; 2 * 1024 * 1024]) }), - ) - .route( - "/idle", - get(|| async { - let first = futures_util::stream::once(async { - Ok::<_, std::io::Error>(Bytes::from_static(b"data: started\n\n")) - }); - ( - [("content-type", "text/event-stream")], - Body::from_stream(first.chain(futures_util::stream::pending())), - ) - }), - ); - let upstream_task = super::task::SessionTask::new( - spawn_router_on_port(upstream_port, upstream).await.unwrap(), - ); - let gateway_port = reserve_local_port().unwrap(); - let gateway_url = format!("http://127.0.0.1:{gateway_port}"); - let (_, gateway_task) = start_gateway_on_port(gateway_port).await.unwrap(); - let gateway_task = super::task::SessionTask::new(gateway_task); - let mut config = sample_config(&gateway_url); - config.tunnel_security = crate::config::TunnelSecurity::NonTlsRequired; - config.tunnel_encryption_key = Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".into()); - config.tunnel_stream_initial_window_bytes = 512 * 1024; - config.tunnel_drain_deadline_ms = 100; - config.allow_private_targets = true; - config.allowed_ports.push(upstream_port); - let state = sample_state(config); - let server = sample_server(&state, "node-recovery"); - let (shutdown_tx, shutdown_rx) = watch::channel(false); - let (_drain_tx, drain_rx) = watch::channel(false); - let tunnel_task = super::task::SessionTask::new(tokio::spawn({ - let state = Arc::clone(&state); - let server = Arc::clone(&server); - async move { - run(&state, &server, 0, shutdown_rx, drain_rx).await; - } - })); - wait_until_relay_status(&gateway_url, "node-recovery", StatusCode::GATEWAY_TIMEOUT).await; - - let envelope = |path: &str| { - let mut meta: protocol::RequestMeta = - serde_json::from_slice(&relay_probe_envelope()[4..]).unwrap(); - meta.url = format!("http://127.0.0.1:{upstream_port}/{path}"); - meta.stream = true; - meta.timeout = 10; - meta.stream_first_byte_timeout_ms = Some(10_000); - let encoded = serde_json::to_vec(&meta).unwrap(); - let mut result = (encoded.len() as u32).to_be_bytes().to_vec(); - result.extend(encoded); - result - }; - let response = relay_response(&gateway_url, "node-recovery", envelope("large")) - .await - .unwrap(); - assert_eq!(response.status(), StatusCode::OK); - let body = tokio::time::timeout(Duration::from_secs(10), response.bytes()) - .await - .unwrap() - .unwrap(); - assert_eq!(body.len(), 2 * 1024 * 1024); - assert!(body.iter().all(|byte| *byte == b'x')); - - let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle")) - .await - .unwrap(); - assert_eq!( - response.chunk().await.unwrap().unwrap(), - "data: started\n\n" - ); - drop(response); - tokio::time::timeout(Duration::from_secs(3), async { - while server - .active_connections - .load(std::sync::atomic::Ordering::Acquire) - != 0 - { - tokio::task::yield_now().await; - } - }) - .await - .expect("cancelled SSE must release the upstream handler"); - - let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle")) - .await - .unwrap(); - assert!(response.chunk().await.unwrap().is_some()); - shutdown_tx.send(true).unwrap(); - tokio::time::timeout(Duration::from_secs(3), tunnel_task) - .await - .unwrap() - .unwrap(); - assert_eq!( - server - .active_connections - .load(std::sync::atomic::Ordering::Acquire), - 0 - ); - drop(response); - drop(gateway_task); - drop(upstream_task); - } - - fn restore_test_env(key: &str, value: Option) { - if let Some(value) = value { - std::env::set_var(key, value); - } else { - std::env::remove_var(key); - } - } - - async fn start_gateway_on_port_retry( - port: u16, - ) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> { - let mut attempts = 0usize; - loop { - match start_gateway_on_port(port).await { - Ok(server) => return Ok(server), - Err(err) => { - attempts += 1; - if attempts >= 20 { - return Err(err); - } - tokio::time::sleep(Duration::from_millis(50)).await; - } - } - } - } - - async fn spawn_router_on_port( - port: u16, - app: Router, - ) -> Result, std::io::Error> { - let listener = tokio::net::TcpListener::bind(("127.0.0.1", port)).await?; - Ok(tokio::spawn(async move { - axum::serve( - listener, - app.into_make_service_with_connect_info::(), - ) - .await - .expect("gateway test server should run"); - })) - } - - fn reserve_local_port() -> Result { - let listener = std::net::TcpListener::bind("127.0.0.1:0")?; - let port = listener.local_addr()?.port(); - drop(listener); - Ok(port) - } - - fn sample_state(config: Config) -> Arc { - let config = Arc::new(config); - let dns_cache = Arc::new(DnsCache::new(Duration::from_secs(60), 128)); - let upstream_client_pool = - upstream_client::UpstreamClientPool::new(Arc::clone(&config), Arc::clone(&dns_cache)); - Arc::new(TunnelAppState { - config, - dns_cache, - upstream_client_pool, - tunnel_tls_config: Arc::new(crate::tunnel::client::build_tls_config()), - resource_monitor: Arc::new(crate::hardware::RuntimeResourceMonitor::new()), - stream_gate: None, - distributed_stream_gate: None, - }) - } - - fn sample_server(state: &Arc, node_id: &str) -> Arc { - let config = Arc::clone(&state.config); - Arc::new(ServerContext { - server_label: "gateway-owned-tunnel".to_string(), - aether_url: config.aether_url.clone(), - management_token: config.management_token.clone(), - tunnel_security: config.tunnel_security, - tunnel_encryption_key: config.tunnel_encryption_key.clone(), - node_name: config.node_name.clone(), - node_id: Arc::new(std::sync::RwLock::new(node_id.to_string())), - tunnel_generation: "test-generation-1".to_string(), - aether_client: Arc::new(AetherClient::new( - &config, - &config.aether_url, - &config.management_token, - )), - dynamic: Arc::new(ArcSwap::from_pointee(DynamicConfig::from_config(&config))), - active_connections: Arc::new(AtomicU64::new(0)), - metrics: Arc::new(TunnelRequestMetrics::new()), - tunnel_metrics: Arc::new(TunnelMetrics::new()), - }) - } - - fn sample_config(aether_url: &str) -> Config { - Config { - aether_url: aether_url.to_string(), - management_token: "token".to_string(), - public_ip: None, - node_name: "tunnel-test".to_string(), - tunnel_security: crate::config::TunnelSecurity::Off, - tunnel_encryption_key: None, - node_region: None, - heartbeat_interval: 1, - allowed_ports: vec![80, 443], - allow_private_targets: false, - aether_request_timeout_secs: 10, - aether_connect_timeout_secs: 2, - aether_pool_max_idle_per_host: 8, - aether_pool_idle_timeout_secs: 90, - aether_tcp_keepalive_secs: 60, - aether_tcp_nodelay: true, - aether_http2: true, - aether_outbound_proxy_url: None, - aether_retry_max_attempts: 1, - aether_retry_base_delay_ms: 50, - aether_retry_max_delay_ms: 100, - diagnostics_bind: None, - max_concurrent_connections: None, - max_in_flight_streams: None, - distributed_stream_limit: None, - distributed_stream_redis_url: None, - distributed_stream_redis_key_prefix: None, - distributed_stream_lease_ttl_ms: 30_000, - distributed_stream_renew_interval_ms: 10_000, - distributed_stream_command_timeout_ms: 1_000, - dns_cache_ttl_secs: 60, - dns_cache_capacity: 128, - upstream_connect_timeout_secs: 30, - upstream_pool_max_idle_per_host: 4, - upstream_pool_idle_timeout_secs: 60, - upstream_client_pool_capacity: crate::config::DEFAULT_UPSTREAM_CLIENT_POOL_CAPACITY, - upstream_tcp_keepalive_secs: 60, - upstream_tcp_nodelay: true, - upstream_proxy_url: None, - upstream_proxy_remote_dns: false, - legacy_redirect_replay_budget_bytes_ignored: None, - emit_proxy_timing_header: true, - log_level: "info".to_string(), - log_destination: crate::config::TunnelLogDestinationArg::Stdout, - log_dir: None, - log_rotation: crate::config::TunnelLogRotationArg::Daily, - log_retention_days: 7, - log_max_files: 30, - tunnel_reconnect_base_ms: 50, - tunnel_reconnect_max_ms: 250, - tunnel_ping_interval_ms: 1_000, - tunnel_max_streams: Some(8), - tunnel_profile: crate::config::TunnelProfileArg::Lite, - tunnel_stream_initial_window_bytes: - crate::config::DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES, - tunnel_drain_deadline_ms: crate::config::DEFAULT_TUNNEL_DRAIN_DEADLINE_MS, - tunnel_connect_timeout_ms: 2_000, - tunnel_ipv4_only: false, - tunnel_ipv6_only: false, - tunnel_tcp_keepalive_secs: 30, - tunnel_tcp_nodelay: true, - tunnel_stale_timeout_ms: 5_000, - tunnel_connections: Some(1), - tunnel_connections_max: Some(1), - tunnel_scale_check_interval_ms: 1_000, - tunnel_scale_up_threshold_percent: 70, - tunnel_scale_down_threshold_percent: 35, - tunnel_scale_down_grace_secs: 15, - } - } - - fn ensure_rustls_provider() { - static INIT: Once = Once::new(); - INIT.call_once(|| { - let _ = rustls::crypto::ring::default_provider().install_default(); - }); - } } diff --git a/crates/aether-admin/src/observability/stats.rs b/crates/aether-admin/src/observability/stats.rs index a041b6524..cc2ff7630 100644 --- a/crates/aether-admin/src/observability/stats.rs +++ b/crates/aether-admin/src/observability/stats.rs @@ -841,6 +841,53 @@ pub fn build_admin_stats_leaderboard_response( .into_response() } +pub fn build_admin_stats_user_group_leaderboard_response( + metric: AdminStatsLeaderboardMetric, + time_range: Option<&AdminStatsTimeRange>, + leaderboard: &[AdminStatsLeaderboardItem], + member_counts: &std::collections::BTreeMap, + active_member_counts: &std::collections::BTreeMap, + offset: usize, + limit: usize, +) -> Response { + let total = leaderboard.len(); + let items: Vec<_> = leaderboard + .iter() + .enumerate() + .skip(offset) + .take(limit) + .map(|(index, item)| { + let rank = compute_dense_rank(metric, leaderboard, index); + let value = match metric { + AdminStatsLeaderboardMetric::Requests => json!(item.requests), + AdminStatsLeaderboardMetric::Tokens => json!(item.tokens), + AdminStatsLeaderboardMetric::Cost => json!(round_to(item.cost, 6)), + }; + json!({ + "rank": rank, + "id": item.id, + "name": item.name, + "value": value, + "requests": item.requests, + "tokens": item.tokens, + "cost": round_to(item.cost, 6), + "member_count": member_counts.get(&item.id).copied().unwrap_or(0), + "active_member_count": active_member_counts.get(&item.id).copied().unwrap_or(0), + }) + }) + .collect(); + + Json(json!({ + "items": items, + "total": total, + "metric": metric.as_str(), + "start_date": time_range.map(|value| value.start_date.to_string()), + "end_date": time_range.map(|value| value.end_date.to_string()), + "attribution": "current_membership", + })) + .into_response() +} + pub fn build_admin_stats_comparison_response( current_usage: &[StoredRequestUsageAudit], comparison_usage: &[StoredRequestUsageAudit], diff --git a/crates/aether-admin/src/observability/usage.rs b/crates/aether-admin/src/observability/usage.rs index 8a604590f..8b8465508 100644 --- a/crates/aether-admin/src/observability/usage.rs +++ b/crates/aether-admin/src/observability/usage.rs @@ -1349,6 +1349,7 @@ fn admin_usage_active_request_json( if let Some(target_model) = item.target_model.as_ref() { value["target_model"] = json!(target_model); } + value["response_model"] = json!(item.provider_response_model()); if let Some(reasoning_effort) = item.provider_reasoning_effort() { value["reasoning_effort"] = json!(reasoning_effort); } @@ -1445,6 +1446,11 @@ pub fn admin_usage_record_json( let object = payload .as_object_mut() .expect("admin usage record payload should be an object"); + // 大型 json! 宏接近 Rust 的递归展开上限,响应模型在宏展开后补入即可避免编译失败。 + object.insert( + "response_model".to_string(), + json!(item.provider_response_model()), + ); object.insert( "end_to_end_time_ms".to_string(), json!(admin_usage_metadata_u64(item, "end_to_end_time_ms")), @@ -2842,6 +2848,32 @@ mod tests { assert_eq!(record["client_is_stream"], false); } + #[test] + fn admin_usage_payloads_expose_response_model_separately_from_mapping() { + let item = StoredRequestUsageAudit { + target_model: Some("provider-mapped-model".to_string()), + request_metadata: Some(json!({ + "provider_response_model": "gpt-5.1" + })), + ..sample_usage("completed", Some(200), None) + }; + + let record = admin_usage_record_json( + &item, + &BTreeMap::new(), + &BTreeMap::new(), + false, + false, + None, + ); + let active = admin_usage_active_request_json(&item, None, None, None); + + for payload in [&record, &active] { + assert_eq!(payload["target_model"], "provider-mapped-model"); + assert_eq!(payload["response_model"], "gpt-5.1"); + } + } + #[test] fn admin_usage_payloads_project_end_to_end_timings_from_metadata() { let item = StoredRequestUsageAudit { diff --git a/crates/aether-admin/src/provider/pool.rs b/crates/aether-admin/src/provider/pool.rs index 18f18d2c7..91307a8d0 100644 --- a/crates/aether-admin/src/provider/pool.rs +++ b/crates/aether-admin/src/provider/pool.rs @@ -169,6 +169,18 @@ pub fn admin_pool_key_account_quota_exhausted( aether_provider_pool::provider_pool_key_account_quota_exhausted(key, provider_type) } +pub fn admin_pool_key_minimum_quota_reached( + key: &StoredProviderCatalogKey, + provider_type: &str, + provider_model_name: Option<&str>, +) -> bool { + aether_provider_pool::provider_pool_key_minimum_quota_reached( + key, + provider_type, + provider_model_name, + ) +} + pub fn admin_pool_key_quota_hard_blocked( key: &StoredProviderCatalogKey, provider_type: &str, diff --git a/crates/aether-admin/src/provider/quota.rs b/crates/aether-admin/src/provider/quota.rs index 2a4647159..7e4b5174b 100644 --- a/crates/aether-admin/src/provider/quota.rs +++ b/crates/aether-admin/src/provider/quota.rs @@ -404,6 +404,200 @@ pub fn parse_antigravity_quota_summary_response( (!parsed_groups.is_empty()).then_some(serde_json::Value::Array(parsed_groups)) } +/// Windows reported by `GET /api/oauth/usage`, as `(response key, metadata prefix)`. +pub const CLAUDE_CODE_USAGE_WINDOWS: [(&str, &str); 4] = [ + ("five_hour", "five_hour"), + ("seven_day", "seven_day"), + ("seven_day_sonnet", "seven_day_sonnet"), + ("seven_day_overage_included", "seven_day_fable"), +]; + +/// Parses the Anthropic OAuth usage response into the `claude_code` metadata bucket. +/// +/// Each window carries `utilization` (percent, 0-100) and `resets_at` (RFC 3339). +/// Windows that are absent or `null` (e.g. plans without a Sonnet/Fable window) are +/// skipped; `None` is returned when no window is present at all. +pub fn parse_claude_code_oauth_usage_response( + value: &serde_json::Value, + updated_at_unix_secs: u64, +) -> Option { + let root = value.as_object()?; + let mut bucket = serde_json::Map::new(); + for (response_key, prefix) in CLAUDE_CODE_USAGE_WINDOWS { + let Some(window) = root + .get(response_key) + .and_then(serde_json::Value::as_object) + else { + continue; + }; + let utilization = window + .get("utilization") + .and_then(|value| match value { + serde_json::Value::Number(number) => number.as_f64(), + serde_json::Value::String(text) => text.trim().parse::().ok(), + _ => None, + }) + .filter(|value| value.is_finite()); + let reset_at = window + .get("resets_at") + .and_then(serde_json::Value::as_str) + .and_then(|text| chrono::DateTime::parse_from_rfc3339(text.trim()).ok()) + .and_then(|time| u64::try_from(time.timestamp()).ok()); + if utilization.is_none() && reset_at.is_none() { + continue; + } + if let Some(utilization) = utilization { + bucket.insert( + format!("{prefix}_used_percent"), + serde_json::json!(utilization.clamp(0.0, 100.0)), + ); + } + if let Some(reset_at) = reset_at { + bucket.insert(format!("{prefix}_reset_at"), serde_json::json!(reset_at)); + } + } + if bucket.is_empty() { + return None; + } + // Explicit null so a refresh overwrites (rather than keeps) credits that were used up. + bucket.insert( + "reset_credits".to_string(), + parse_claude_code_reset_credits(root, updated_at_unix_secs) + .unwrap_or(serde_json::Value::Null), + ); + bucket.insert( + "updated_at".to_string(), + serde_json::json!(updated_at_unix_secs), + ); + Some(serde_json::Value::Object(bucket)) +} + +/// Projects the `cedar_ember` block (returned with `?cedar_ember=1`) into the same +/// `reset_credits` shape codex uses. Upstream grant/organization ids are never copied; each +/// usable grant becomes one entry, and `available_count` sums their remaining resets. +fn parse_claude_code_reset_credits( + root: &serde_json::Map, + now_unix_secs: u64, +) -> Option { + let grants = root + .get("cedar_ember") + .and_then(serde_json::Value::as_object)? + .get("grants") + .and_then(serde_json::Value::as_array)?; + let parse_time = |grant: &serde_json::Map, field: &str| { + grant + .get(field) + .and_then(serde_json::Value::as_str) + .and_then(|text| chrono::DateTime::parse_from_rfc3339(text.trim()).ok()) + .and_then(|time| u64::try_from(time.timestamp()).ok()) + }; + let mut available_count = 0u64; + let mut credits = Vec::new(); + for grant in grants.iter().filter_map(serde_json::Value::as_object) { + let resets_left = grant + .get("resets_left") + .and_then(serde_json::Value::as_u64) + .unwrap_or(0); + let has_clears = grant + .get("clears") + .and_then(serde_json::Value::as_array) + .is_some_and(|clears| !clears.is_empty()); + let paused = grant + .get("paused") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false); + let starts_at = parse_time(grant, "starts_at"); + let expires_at = parse_time(grant, "ends_at"); + if resets_left == 0 + || !has_clears + || paused + || starts_at.is_some_and(|value| value > now_unix_secs) + || expires_at.is_some_and(|value| value <= now_unix_secs) + { + continue; + } + available_count += resets_left; + let Some(expires_at) = expires_at else { + continue; + }; + credits.push(serde_json::json!({ + "display_key": format!("Key-{}", credits.len() + 1), + "status": "available", + "expires_at": expires_at, + "remaining_seconds": expires_at - now_unix_secs, + })); + } + if available_count == 0 { + return None; + } + Some(serde_json::json!({ + "available_count": available_count, + "updated_at": now_unix_secs, + "detail_source": "claude_oauth_usage", + "detail_status": "ok", + "credits": credits, + })) +} + +/// Parses the `anthropic-ratelimit-unified-*` response headers into a partial `claude_code` +/// metadata bucket (same field names as [`parse_claude_code_oauth_usage_response`]). +/// +/// Utilization headers are 0-1 fractions and are stored as percent; reset headers are Unix +/// seconds (millisecond values are normalized). Returns `None` when no window is reported. +pub fn parse_claude_code_usage_headers( + headers: &BTreeMap, + updated_at_unix_secs: u64, +) -> Option { + let normalized = headers + .iter() + .map(|(key, value)| (key.trim().to_ascii_lowercase(), value.trim().to_string())) + .collect::>(); + let mut bucket = serde_json::Map::new(); + for (header_window, prefix) in [ + ("5h", "five_hour"), + ("7d", "seven_day"), + ("7d_oi", "seven_day_fable"), + ] { + let header = |suffix: &str| { + normalized + .get(&format!( + "anthropic-ratelimit-unified-{header_window}-{suffix}" + )) + .map(String::as_str) + }; + if let Some(utilization) = header("utilization") + .and_then(|value| value.parse::().ok()) + .filter(|value| value.is_finite()) + { + bucket.insert( + format!("{prefix}_used_percent"), + serde_json::json!((utilization * 100.0).clamp(0.0, 100.0)), + ); + } + if let Some(reset_at) = header("reset") + .and_then(|value| value.parse::().ok()) + .map(|value| { + if value > 100_000_000_000 { + value / 1_000 + } else { + value + } + }) + .filter(|value| *value > 0) + { + bucket.insert(format!("{prefix}_reset_at"), serde_json::json!(reset_at)); + } + } + if bucket.is_empty() { + return None; + } + bucket.insert( + "updated_at".to_string(), + serde_json::json!(updated_at_unix_secs), + ); + Some(serde_json::Value::Object(bucket)) +} + pub fn parse_gemini_cli_retrieve_user_quota_response( value: &serde_json::Value, updated_at_unix_secs: u64, @@ -3546,6 +3740,258 @@ pub fn parse_kiro_usage_response( Some(serde_json::Value::Object(result)) } +pub fn parse_xai_billing_response( + value: &serde_json::Value, + updated_at_unix_secs: u64, +) -> Option { + let root = value.as_object()?; + let config = root + .get("config") + .and_then(serde_json::Value::as_object) + .unwrap_or(root); + + let usage_percentage = coerce_json_f64_from_map(config, "creditUsagePercent") + .or_else(|| extract_xai_product_usage_percent(config)); + let period = config.get("currentPeriod"); + let period_type = period + .and_then(|value| value.get("type").or_else(|| value.get("periodType"))) + .and_then(normalize_xai_period_type); + let next_reset_at = period + .and_then(|value| value.get("end")) + .and_then(parse_xai_timestamp) + .or_else(|| config.get("billingPeriodEnd").and_then(parse_xai_timestamp)); + let monthly_limit = + coerce_xai_cents_dollars(config.get("monthlyLimit")).filter(|value| *value > 0.0); + let current_usage = if monthly_limit.is_some() { + coerce_xai_cents_dollars(config.get("used")) + } else { + None + }; + let remaining = monthly_limit + .zip(current_usage) + .map(|(limit, used)| (limit - used).max(0.0)); + let usage_percentage = usage_percentage.or_else(|| { + monthly_limit + .zip(current_usage) + .map(|(limit, used)| ((used / limit) * 100.0).clamp(0.0, 100.0)) + }); + let usage_percentage = match usage_percentage { + Some(value) => Some(value.clamp(0.0, 100.0)), + None if period_type.is_some() || next_reset_at.is_some() => Some(0.0), + None => None, + }; + let prepaid_balance = coerce_xai_cents_dollars(config.get("prepaidBalance")); + let on_demand_cap = coerce_xai_cents_dollars(config.get("onDemandCap")); + let on_demand_used = coerce_xai_cents_dollars(config.get("onDemandUsed")); + let on_demand_enabled = coerce_json_bool_from_map(root, "onDemandEnabled") + .or_else(|| coerce_json_bool_from_map(config, "onDemandEnabled")); + let subscription_title = first_json_string_by_paths( + value, + &[ + &["subscriptionTier"], + &["subscription_tier"], + &["config", "subscriptionTier"], + &["config", "subscription_title"], + ], + ); + + if usage_percentage.is_none() + && monthly_limit.is_none() + && current_usage.is_none() + && prepaid_balance.is_none() + && on_demand_cap.is_none() + && next_reset_at.is_none() + && subscription_title.is_none() + { + return None; + } + + let mut result = serde_json::Map::new(); + result.insert("updated_at".to_string(), json!(updated_at_unix_secs)); + if let Some(value) = usage_percentage { + result.insert("usage_percentage".to_string(), json!(value)); + } + if let Some(value) = monthly_limit { + result.insert("usage_limit".to_string(), json!(value)); + } + if let Some(value) = current_usage { + result.insert("current_usage".to_string(), json!(value)); + } + if let Some(value) = remaining { + result.insert("remaining".to_string(), json!(value)); + } + if let Some(value) = next_reset_at { + result.insert("next_reset_at".to_string(), json!(value)); + } + if let Some(value) = period_type { + result.insert("period_type".to_string(), json!(value)); + } + if let Some(value) = prepaid_balance { + result.insert("prepaid_balance".to_string(), json!(value)); + } + if let Some(value) = on_demand_cap { + result.insert("on_demand_cap".to_string(), json!(value)); + } + if let Some(value) = on_demand_used { + result.insert("on_demand_used".to_string(), json!(value)); + } + if let Some(value) = on_demand_enabled { + result.insert("on_demand_enabled".to_string(), json!(value)); + } + if let Some(value) = subscription_title { + result.insert("subscription_title".to_string(), json!(value)); + } + Some(serde_json::Value::Object(result)) +} + +fn coerce_json_f64_from_map( + object: &serde_json::Map, + key: &str, +) -> Option { + object.get(key).and_then(coerce_json_f64) +} + +fn coerce_json_bool_from_map( + object: &serde_json::Map, + key: &str, +) -> Option { + object.get(key).and_then(coerce_json_bool) +} + +fn extract_xai_product_usage_percent( + config: &serde_json::Map, +) -> Option { + let items = config.get("productUsage")?.as_array()?; + let grok_build = items.iter().find(|item| { + item.get("product") + .and_then(serde_json::Value::as_str) + .is_some_and(|product| product.eq_ignore_ascii_case("GrokBuild")) + }); + grok_build + .or(items.first()) + .and_then(|item| item.get("usagePercent").and_then(coerce_json_f64)) +} + +fn coerce_xai_cents_dollars(value: Option<&serde_json::Value>) -> Option { + let value = value?; + let cents = match value { + serde_json::Value::Object(object) => object.get("val").and_then(coerce_json_f64)?, + other => coerce_json_f64(other)?, + }; + Some(cents / 100.0) +} + +fn normalize_xai_period_type(value: &serde_json::Value) -> Option { + let raw = value + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty())?; + let lowered = raw.to_ascii_lowercase(); + if lowered.contains("week") { + Some("weekly".to_string()) + } else if lowered.contains("month") { + Some("monthly".to_string()) + } else { + Some(raw.to_string()) + } +} + +fn parse_xai_timestamp(value: &serde_json::Value) -> Option { + if let Some(value) = coerce_json_u64(value) { + return Some(if value > 1_000_000_000_000 { + value / 1000 + } else { + value + }); + } + let raw = value.as_str()?.trim(); + if raw.is_empty() { + return None; + } + chrono::DateTime::parse_from_rfc3339(raw) + .ok() + .and_then(|timestamp| u64::try_from(timestamp.timestamp()).ok()) +} + +#[cfg(test)] +mod xai_quota_tests { + use super::parse_xai_billing_response; + use serde_json::json; + + #[test] + fn parse_xai_credits_percent_and_weekly_period() { + let metadata = parse_xai_billing_response( + &json!({ + "config": { + "currentPeriod": { + "type": "USAGE_PERIOD_TYPE_WEEKLY", + "start": "2026-08-08T01:53:09.930537+00:00", + "end": "2026-08-15T01:53:09.930537+00:00" + }, + "creditUsagePercent": 46.0, + "productUsage": [ + {"product": "GrokBuild", "usagePercent": 41.0}, + {"product": "GrokChat"} + ], + "onDemandCap": {"val": 0}, + "onDemandUsed": {"val": 0}, + "prepaidBalance": {"val": 0} + }, + "subscriptionTier": "SuperGrok" + }), + 1_775_000_000, + ) + .expect("credits payload should parse"); + + assert_eq!(metadata["usage_percentage"], json!(46.0)); + assert_eq!(metadata["period_type"], json!("weekly")); + assert_eq!(metadata["next_reset_at"], json!(1_786_758_789u64)); + assert_eq!(metadata["prepaid_balance"], json!(0.0)); + assert_eq!(metadata["on_demand_cap"], json!(0.0)); + assert_eq!(metadata["subscription_title"], json!("SuperGrok")); + } + + #[test] + fn parse_xai_omitted_percent_as_fresh_weekly_zero() { + let metadata = parse_xai_billing_response( + &json!({ + "config": { + "currentPeriod": { + "type": "USAGE_PERIOD_TYPE_WEEKLY", + "end": "2026-08-15T01:53:09.930537+00:00" + }, + "isUnifiedBillingUser": true + } + }), + 1_775_000_000, + ) + .expect("fresh weekly period should parse"); + + assert_eq!(metadata["usage_percentage"], json!(0.0)); + assert_eq!(metadata["period_type"], json!("weekly")); + } + + #[test] + fn parse_xai_legacy_monthly_cents() { + let metadata = parse_xai_billing_response( + &json!({ + "config": { + "monthlyLimit": {"val": 2500}, + "used": {"val": 1000}, + "billingPeriodEnd": "2026-09-01T00:00:00Z" + } + }), + 1_775_000_000, + ) + .expect("legacy monthly payload should parse"); + + assert_eq!(metadata["usage_limit"], json!(25.0)); + assert_eq!(metadata["current_usage"], json!(10.0)); + assert_eq!(metadata["remaining"], json!(15.0)); + assert_eq!(metadata["usage_percentage"], json!(40.0)); + } +} + pub fn parse_windsurf_user_status_response( value: &serde_json::Value, updated_at_unix_secs: u64, @@ -7452,3 +7898,163 @@ mod tests { assert!(!serialized.contains("user:password")); } } + +#[cfg(test)] +mod claude_code_quota_tests { + use super::{parse_claude_code_oauth_usage_response, parse_claude_code_usage_headers}; + use serde_json::json; + use std::collections::BTreeMap; + + #[test] + fn parses_cedar_ember_grants_into_reset_credits() { + let parsed = parse_claude_code_oauth_usage_response( + &json!({ + "five_hour": {"utilization": 65.0, "resets_at": "2027-01-15T08:00:00Z"}, + "cedar_ember": { + "eligible": true, + "grants": [ + {"id": "launch", "resets_left": 2, "clears": ["five_hour"], + "ends_at": "2027-01-20T00:00:00Z"}, + {"id": "later", "resets_left": 1, "clears": ["five_hour"], + "starts_at": "2027-02-01T00:00:00Z"}, + {"id": "paused", "resets_left": 1, "clears": ["five_hour"], "paused": true}, + {"id": "spent", "resets_left": 0, "clears": ["five_hour"]}, + {"id": "expired", "resets_left": 1, "clears": ["five_hour"], + "ends_at": "2027-01-01T00:00:00Z"} + ] + } + }), + 1_800_000_000, + ) + .expect("usage should parse"); + + let credits = &parsed["reset_credits"]; + assert_eq!(credits["available_count"], json!(2)); + assert_eq!(credits["credits"].as_array().map(Vec::len), Some(1)); + assert_eq!(credits["credits"][0]["display_key"], json!("Key-1")); + assert_eq!(credits["credits"][0]["expires_at"], json!(1_800_403_200u64)); + assert!(credits["credits"][0].get("id").is_none()); + assert!(!credits.to_string().contains("launch")); + } + + #[test] + fn real_cedar_ember_response_survives_metadata_redaction() { + // Shape captured from a live Claude Pro account (unrelated fields trimmed). + let parsed = parse_claude_code_oauth_usage_response( + &json!({ + "five_hour": {"utilization": 32.0, "resets_at": "2026-09-29T19:19:59.933008+00:00"}, + "cedar_ember": { + "eligible": true, + "at_limit": false, + "grants": [{ + "id": "opus55-launch-promax-20260921", + "label": "Claude Opus 5.5 launch", + "resets_total": 1, + "resets_left": 1, + "starts_at": "2026-09-22T16:00:00+00:00", + "ends_at": "2026-10-22T16:00:00+00:00", + "clears": ["five_hour", "seven_day"], + "paused": false, + "usable_now": true, + "use_requires_limit": false + }], + "next_grant_id": "opus55-launch-promax-20260921" + } + }), + 1_790_699_000, + ) + .expect("usage should parse"); + assert_eq!(parsed["reset_credits"]["available_count"], json!(1)); + + let safe = crate::provider::redaction::admin_provider_upstream_metadata_safe_json(Some( + &json!({ "claude_code": parsed }), + )); + assert_eq!( + safe["claude_code"]["reset_credits"]["available_count"], + json!(1), + "redaction dropped reset_credits: {safe}" + ); + } + + #[test] + fn cedar_ember_null_or_empty_omits_reset_credits() { + let parsed = parse_claude_code_oauth_usage_response( + &json!({"five_hour": {"utilization": 1.0}, "cedar_ember": null}), + 1, + ) + .expect("usage should parse"); + assert!(parsed["reset_credits"].is_null()); + } + + #[test] + fn parses_unified_ratelimit_headers_into_percent_and_reset() { + let headers = BTreeMap::from([ + ( + "Anthropic-Ratelimit-Unified-5h-Utilization".to_string(), + "0.42".to_string(), + ), + ( + "anthropic-ratelimit-unified-5h-reset".to_string(), + "1800003600".to_string(), + ), + ( + "anthropic-ratelimit-unified-7d-utilization".to_string(), + "1.0".to_string(), + ), + ( + "anthropic-ratelimit-unified-7d-reset".to_string(), + "1800400000000".to_string(), + ), + ( + "anthropic-ratelimit-unified-7d_oi-utilization".to_string(), + "0.1".to_string(), + ), + ]); + let parsed = + parse_claude_code_usage_headers(&headers, 1_800_000_000).expect("headers should parse"); + + assert_eq!(parsed["five_hour_used_percent"], json!(42.0)); + assert_eq!(parsed["five_hour_reset_at"], json!(1_800_003_600u64)); + assert_eq!(parsed["seven_day_used_percent"], json!(100.0)); + // Millisecond timestamps are normalized to seconds. + assert_eq!(parsed["seven_day_reset_at"], json!(1_800_400_000u64)); + assert_eq!(parsed["seven_day_fable_used_percent"], json!(10.0)); + assert!(parsed.get("seven_day_fable_reset_at").is_none()); + assert_eq!(parsed["updated_at"], json!(1_800_000_000u64)); + } + + #[test] + fn unified_ratelimit_headers_absent_yield_none() { + let headers = + BTreeMap::from([("content-type".to_string(), "application/json".to_string())]); + assert!(parse_claude_code_usage_headers(&headers, 1).is_none()); + } + + #[test] + fn parses_windows_and_skips_null_ones() { + let parsed = parse_claude_code_oauth_usage_response( + &json!({ + "five_hour": {"utilization": 37.5, "resets_at": "2027-01-15T08:00:00.000000+00:00"}, + "seven_day": {"utilization": 12, "resets_at": "2027-01-20T00:00:00Z"}, + "seven_day_sonnet": null, + "seven_day_overage_included": {"utilization": 3.0, "resets_at": null} + }), + 1_800_000_000, + ) + .expect("usage windows should parse"); + + assert_eq!(parsed["updated_at"], json!(1_800_000_000u64)); + assert_eq!(parsed["five_hour_used_percent"], json!(37.5)); + assert_eq!(parsed["five_hour_reset_at"], json!(1_800_000_000u64)); + assert_eq!(parsed["seven_day_used_percent"], json!(12.0)); + assert!(parsed.get("seven_day_sonnet_used_percent").is_none()); + assert_eq!(parsed["seven_day_fable_used_percent"], json!(3.0)); + assert!(parsed.get("seven_day_fable_reset_at").is_none()); + } + + #[test] + fn returns_none_without_any_window() { + assert!(parse_claude_code_oauth_usage_response(&json!({}), 1).is_none()); + assert!(parse_claude_code_oauth_usage_response(&json!({"five_hour": null}), 1).is_none()); + } +} diff --git a/crates/aether-admin/src/provider/state.rs b/crates/aether-admin/src/provider/state.rs index 1d4a40bf4..b7cbc2ade 100644 --- a/crates/aether-admin/src/provider/state.rs +++ b/crates/aether-admin/src/provider/state.rs @@ -285,6 +285,23 @@ pub fn enrich_admin_provider_oauth_auth_config( ], ); + if provider_type.trim().eq_ignore_ascii_case("xai") { + auth_config.insert("auth_method".to_string(), json!("oauth")); + auth_config.insert("using_api".to_string(), json!(false)); + if let Some(id_token) = ["id_token", "idToken"] + .iter() + .find_map(|field| json_non_empty_string(token_payload.get(field))) + { + auth_config + .entry("id_token".to_string()) + .or_insert_with(|| json!(id_token.clone())); + if let Some(claims) = decode_jwt_claims(&id_token) { + merge_missing_auth_config_fields(auth_config, &claims, &["email", "sub"]); + } + } + return; + } + if provider_type.trim().eq_ignore_ascii_case("claude_code") { if let Some(organization_uuid) = token_payload_object .get("organization") @@ -554,6 +571,28 @@ mod tests { assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true))); } + #[test] + fn xai_enrichment_marks_oauth_and_extracts_id_token_identity() { + let id_token = sample_unsigned_jwt(json!({ + "email": "grok@x.ai", + "sub": "user-xai-1", + })); + let token_payload = json!({ + "access_token": "access-token", + "refresh_token": "refresh-token", + "id_token": id_token, + }); + let mut auth_config = serde_json::Map::new(); + + enrich_admin_provider_oauth_auth_config("xai", &mut auth_config, &token_payload); + + assert_eq!(auth_config.get("auth_method"), Some(&json!("oauth"))); + assert_eq!(auth_config.get("using_api"), Some(&json!(false))); + assert_eq!(auth_config.get("email"), Some(&json!("grok@x.ai"))); + assert_eq!(auth_config.get("sub"), Some(&json!("user-xai-1"))); + assert_eq!(auth_config.get("id_token"), Some(&json!(id_token))); + } + #[test] fn decode_jwt_claims_rejects_oversized_payload_before_decode() { let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES diff --git a/crates/aether-ai/formats/src/api.rs b/crates/aether-ai/formats/src/api.rs index de48589f7..fbbd51933 100644 --- a/crates/aether-ai/formats/src/api.rs +++ b/crates/aether-ai/formats/src/api.rs @@ -208,6 +208,10 @@ pub use crate::formats::{ resolve_stream_spec as resolve_openai_responses_stream_spec, resolve_sync_spec as resolve_openai_responses_sync_spec, LocalOpenAiResponsesSpec, }, + xai::{ + apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client, + xai_supports_native_image_generation, + }, }, }, shared::{ @@ -220,9 +224,11 @@ pub use crate::formats::{ standard_normalize::{ build_cross_format_openai_chat_request_body, build_cross_format_openai_chat_request_body_with_model_directives, + build_cross_format_openai_chat_request_body_with_provider_context, build_cross_format_openai_responses_request_body, build_cross_format_openai_responses_request_body_with_model_directives, build_cross_format_openai_responses_request_body_with_model_directives_and_history_scope, + build_cross_format_openai_responses_request_body_with_provider_context, build_local_openai_chat_request_body, build_local_openai_chat_request_body_with_model_directives, build_local_openai_responses_request_body, diff --git a/crates/aether-ai/formats/src/codex_profile.rs b/crates/aether-ai/formats/src/codex_profile.rs new file mode 100644 index 000000000..1338a9f8c --- /dev/null +++ b/crates/aether-ai/formats/src/codex_profile.rs @@ -0,0 +1,108 @@ +use std::sync::{OnceLock, RwLock}; + +/// 当前支持的 Codex 客户端类型。 +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum CodexClientKind { + Cli, + Desktop, +} + +/// Codex 上游请求使用的客户端画像。 +/// +/// 画像由网关后台任务更新,格式转换层只读取不可变快照,避免在请求路径执行网络操作。 +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct CodexClientProfile { + pub client_kind: CodexClientKind, + pub codex_version: String, + pub originator: String, + pub user_agent: String, +} + +impl CodexClientProfile { + /// 从稳定版本号创建 CLI 画像;版本校验由发布检查器负责,构造器只拒绝明显非法值。 + pub fn cli(version: &str) -> Result { + let version = version.trim(); + if version.is_empty() + || version.len() > 64 + || !version.bytes().all(|byte| (32..=126).contains(&byte)) + { + return Err("invalid Codex CLI version"); + } + let originator = "codex_cli_rs".to_owned(); + Ok(Self { + client_kind: CodexClientKind::Cli, + codex_version: version.to_owned(), + user_agent: format!("{}/{}", originator, version), + originator, + }) + } +} + +impl Default for CodexClientProfile { + fn default() -> Self { + // 远程发布检查不可用时仍保持现有线上行为,避免启动或请求被版本服务拖住。 + Self::cli("0.153.4").expect("built-in Codex CLI profile must be valid") + } +} + +static ACTIVE_PROFILE: OnceLock> = OnceLock::new(); + +fn active_profile() -> &'static RwLock { + ACTIVE_PROFILE.get_or_init(|| RwLock::new(CodexClientProfile::default())) +} + +/// 返回当前画像的独立快照,调用方不会持有全局锁。 +pub fn codex_client_profile() -> CodexClientProfile { + active_profile() + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() +} + +/// 原子替换当前画像,并返回替换前的画像。 +pub fn set_codex_client_profile(profile: CodexClientProfile) -> CodexClientProfile { + let mut current = active_profile() + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + std::mem::replace(&mut *current, profile) +} + +/// 发布一份新的 CLI 画像。 +pub fn set_codex_cli_version(version: &str) -> Result { + let profile = CodexClientProfile::cli(version)?; + Ok(set_codex_client_profile(profile)) +} + +/// 返回当前画像的 Codex Core 版本。 +pub fn codex_client_version() -> String { + codex_client_profile().codex_version +} + +/// 返回当前画像的 User-Agent。 +pub fn codex_client_user_agent() -> String { + codex_client_profile().user_agent +} + +/// 返回当前画像的 originator。 +pub fn codex_client_originator() -> String { + codex_client_profile().originator +} + +#[cfg(test)] +mod tests { + use super::{CodexClientKind, CodexClientProfile}; + + #[test] + fn cli_profile_derives_wire_identity_from_version() { + let profile = CodexClientProfile::cli("0.200.1").expect("valid version"); + assert_eq!(profile.client_kind, CodexClientKind::Cli); + assert_eq!(profile.originator, "codex_cli_rs"); + assert_eq!(profile.user_agent, "codex_cli_rs/0.200.1"); + } + + #[test] + fn cli_profile_rejects_empty_or_control_values() { + assert!(CodexClientProfile::cli("").is_err()); + assert!(CodexClientProfile::cli("0.1.0\nspoof").is_err()); + } +} diff --git a/crates/aether-ai/formats/src/formats/claude/messages/stream.rs b/crates/aether-ai/formats/src/formats/claude/messages/stream.rs index af816734c..a20788a38 100644 --- a/crates/aether-ai/formats/src/formats/claude/messages/stream.rs +++ b/crates/aether-ai/formats/src/formats/claude/messages/stream.rs @@ -2,6 +2,7 @@ use std::collections::BTreeMap; use serde_json::{json, Map, Value}; +use crate::formats::shared::citations::canonical_citations_to_claude_citations; use crate::formats::shared::response::{ build_generated_tool_call_id, canonicalize_tool_arguments, remove_empty_pages_from_tool_arguments, @@ -773,6 +774,33 @@ impl ClaudeClientEmitter { name, content, } => self.emit_tool_result_block(index, tool_use_id, name, content), + CanonicalStreamEvent::Citations(citations) => { + let citations = canonical_citations_to_claude_citations(&citations); + if citations.is_empty() { + return Ok(Vec::new()); + } + // Citations belong to the answer text. If a tool call or a + // thinking block closed it, open a fresh text block rather than + // hang the evidence off an unrelated one. + let mut out = self.ensure_text_block()?; + let Some(ClaudeOpenBlock::Text { block_index }) = self.open_block else { + return Ok(out); + }; + for citation in citations { + out.extend(encode_json_sse( + Some("content_block_delta"), + &json!({ + "type": "content_block_delta", + "index": block_index, + "delta": { + "type": "citations_delta", + "citation": citation, + } + }), + )?); + } + Ok(out) + } CanonicalStreamEvent::UnknownEvent(_) => Ok(Vec::new()), CanonicalStreamEvent::Finish { finish_reason, diff --git a/crates/aether-ai/formats/src/formats/context.rs b/crates/aether-ai/formats/src/formats/context.rs index 02b189156..be7dee222 100644 --- a/crates/aether-ai/formats/src/formats/context.rs +++ b/crates/aether-ai/formats/src/formats/context.rs @@ -10,6 +10,8 @@ pub struct FormatContext { pub upstream_is_stream: bool, pub report_context: Option, pub history_scope: Option, + /// Defer tool schema lowering to the private provider transport boundary. + pub preserve_gemini_tool_schemas: bool, } impl FormatContext { @@ -45,6 +47,7 @@ impl FormatContext { upstream_is_stream: false, report_context: self.report_context.clone(), history_scope: self.history_scope.clone(), + preserve_gemini_tool_schemas: false, } } diff --git a/crates/aether-ai/formats/src/formats/conversion/response.rs b/crates/aether-ai/formats/src/formats/conversion/response.rs index 646b99a9f..41b7492d4 100644 --- a/crates/aether-ai/formats/src/formats/conversion/response.rs +++ b/crates/aether-ai/formats/src/formats/conversion/response.rs @@ -9,7 +9,8 @@ use serde_json::{json, Value}; use crate::formats::{ context::FormatContext, openai::responses::{ - openai_responses_message_item_id, openai_responses_synthetic_reasoning_item_id, + openai_responses_message_item_id, openai_responses_reasoning_text_parts, + openai_responses_synthetic_reasoning_item_id, response::ensure_modern_openai_responses_response_fields, }, registry, @@ -205,14 +206,13 @@ pub fn build_openai_responses_response_with_content( if trimmed.is_empty() { continue; } + let content = openai_responses_reasoning_text_parts(std::iter::once(trimmed)); output.push(json!({ "type": "reasoning", "id": openai_responses_synthetic_reasoning_item_id(response_id, index), "status": "completed", - "summary": [{ - "type": "summary_text", - "text": trimmed, - }] + "summary": [], + "content": content, })); } if !content.is_empty() { @@ -289,6 +289,31 @@ mod tests { assert!(converted["completed_at"].as_i64().is_some()); } + #[test] + fn manual_responses_response_builder_puts_reasoning_in_content() { + let response = super::build_openai_responses_response_with_reasoning( + "resp_manual_reason", + "gpt-5", + "answer", + vec!["raw thinking".to_string()], + Vec::new(), + super::OpenAiResponsesResponseUsage { + prompt_tokens: 1, + output_tokens: 2, + total_tokens: 3, + }, + ); + + assert_eq!(response["output"][0]["type"], "reasoning"); + assert_eq!( + response["output"][0]["content"][0]["type"], + "reasoning_text" + ); + assert_eq!(response["output"][0]["content"][0]["text"], "raw thinking"); + assert_eq!(response["output"][0]["summary"], json!([])); + assert_eq!(response["output"][1]["content"][0]["text"], "answer"); + } + #[test] fn manual_responses_response_builder_emits_modern_fields() { let response = super::build_openai_responses_response( @@ -306,6 +331,40 @@ mod tests { assert!(response["completed_at"].as_i64().is_some()); } + #[test] + fn chat_reasoning_content_maps_to_responses_content_and_summary() { + let body = json!({ + "id": "chatcmpl-reason", + "object": "chat.completion", + "model": "deepseek-reasoner", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "reasoning_content": "compare the decimals", + "content": "9.80 is larger" + }, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3} + }); + + let converted = convert_openai_chat_response_to_openai_responses(&body, &json!({}), false) + .expect("responses response"); + let item = &converted["output"][0]; + + assert_eq!(item["type"], "reasoning"); + assert_eq!(item["content"][0]["type"], "reasoning_text"); + assert_eq!(item["content"][0]["text"], "compare the decimals"); + assert_eq!(item["summary"], json!([])); + assert!(!item.get("content").unwrap().is_null()); + assert_eq!(converted["output"][1]["type"], "message"); + assert_eq!( + converted["output"][1]["content"][0]["text"], + "9.80 is larger" + ); + } + #[test] fn pairwise_response_helper_uses_report_context_model_fallback() { let body = json!({ diff --git a/crates/aether-ai/formats/src/formats/gemini/generate_content/request.rs b/crates/aether-ai/formats/src/formats/gemini/generate_content/request.rs index cfda9c244..c6d51c4e6 100644 --- a/crates/aether-ai/formats/src/formats/gemini/generate_content/request.rs +++ b/crates/aether-ai/formats/src/formats/gemini/generate_content/request.rs @@ -31,10 +31,11 @@ pub fn from(body: &Value, ctx: &FormatContext) -> Option { } pub fn to(request: &CanonicalRequest, ctx: &FormatContext) -> Option { - to_raw( + to_raw_with_schema_policy( request, ctx.mapped_model_or(request.model.as_str()), ctx.upstream_is_stream, + ctx.preserve_gemini_tool_schemas, ) } @@ -187,7 +188,21 @@ pub fn to_raw( mapped_model: &str, upstream_is_stream: bool, ) -> Option { - let mut output = canonical_to_gemini_request_body(canonical, mapped_model, upstream_is_stream)?; + to_raw_with_schema_policy(canonical, mapped_model, upstream_is_stream, false) +} + +fn to_raw_with_schema_policy( + canonical: &CanonicalRequest, + mapped_model: &str, + upstream_is_stream: bool, + preserve_tool_schemas: bool, +) -> Option { + let mut output = canonical_to_gemini_request_body( + canonical, + mapped_model, + upstream_is_stream, + preserve_tool_schemas, + )?; apply_gemini_request_extensions(&mut output, &canonical.extensions)?; if !canonical_has_raw_gemini_tools(canonical) { enable_server_side_tool_invocations_for_mixed_tools(&mut output, mapped_model)?; @@ -244,7 +259,8 @@ pub fn ensure_server_side_tool_invocations_for_mixed_tools(output: &mut Value) - } pub(crate) fn canonical_has_mixed_gemini_tools(canonical: &CanonicalRequest) -> bool { - canonical_tools_to_gemini(canonical) + // Only tool kinds matter here; do not lower/expand schemas just to count them. + canonical_tools_to_gemini(canonical, true) .and_then(|tools| tools.as_array().cloned()) .is_some_and(|tools| gemini_tools_are_mixed(&tools)) } @@ -275,6 +291,7 @@ fn canonical_to_gemini_request_body( canonical: &CanonicalRequest, mapped_model: &str, _upstream_is_stream: bool, + preserve_tool_schemas: bool, ) -> Option { let mut output = Map::new(); if !mapped_model.trim().is_empty() { @@ -297,7 +314,7 @@ fn canonical_to_gemini_request_body( { output.insert("generationConfig".to_string(), generation_config); } - if let Some(tools) = canonical_tools_to_gemini(canonical) { + if let Some(tools) = canonical_tools_to_gemini(canonical, preserve_tool_schemas) { output.insert("tools".to_string(), tools); } if let Some(tool_config) = canonical_tool_choice_to_gemini(canonical.tool_choice.as_ref()) { @@ -701,7 +718,10 @@ fn apply_response_format_to_gemini_generation_config( } } -fn canonical_tools_to_gemini(canonical: &CanonicalRequest) -> Option { +fn canonical_tools_to_gemini( + canonical: &CanonicalRequest, + preserve_tool_schemas: bool, +) -> Option { let mut declarations = Vec::new(); let mut tools = Vec::new(); let mut google_search = canonical @@ -717,7 +737,7 @@ fn canonical_tools_to_gemini(canonical: &CanonicalRequest) -> Option { let mut url_context = false; for tool in &canonical.tools { - match normalize_gemini_builtin_tool_name(&tool.name) { + match canonical_tool_builtin_gemini_name(tool) { Some("googleSearch") => { google_search = true; continue; @@ -747,7 +767,10 @@ fn canonical_tools_to_gemini(canonical: &CanonicalRequest) -> Option { google_search = true; continue; } - declarations.push(canonical_tool_to_gemini_declaration(tool)); + declarations.push(canonical_tool_to_gemini_declaration( + tool, + preserve_tool_schemas, + )); } let mut emitted_google_search = false; let mut emitted_code_execution = false; @@ -877,7 +900,10 @@ fn gemini_unhandled_builtin_tool_portion(tool_object: &Map) -> Op (!builtin.is_empty()).then_some(Value::Object(builtin)) } -fn canonical_tool_to_gemini_declaration(tool: &CanonicalToolDefinition) -> Value { +fn canonical_tool_to_gemini_declaration( + tool: &CanonicalToolDefinition, + preserve_tool_schema: bool, +) -> Value { let mut declaration = Map::new(); declaration.insert("name".to_string(), Value::String(tool.name.clone())); if let Some(description) = &tool.description { @@ -898,7 +924,7 @@ fn canonical_tool_to_gemini_declaration(tool: &CanonicalToolDefinition) -> Value .clone() .or_else(|| tool.parameters.clone()) .map(|mut schema| { - if raw_parameters.is_none() { + if raw_parameters.is_none() && !preserve_tool_schema { clean_gemini_schema(&mut schema); } schema @@ -980,6 +1006,25 @@ fn compact_gemini_contents(contents: Vec) -> Vec { compact } +/// Promote a canonical tool to a Gemini builtin only when it is a bare marker. +/// +/// Clients declare ordinary function tools whose names collide with the builtin +/// spellings — Claude Code ships a client-side `WebSearch` tool with a full +/// `input_schema`. Matching on the name alone dropped those declarations and +/// replaced them with server-side grounding, so the model could never call the +/// tool the client actually implements. A declared schema means the caller +/// expects to execute the call itself, so such tools stay function declarations. +fn canonical_tool_builtin_gemini_name(tool: &CanonicalToolDefinition) -> Option<&'static str> { + if tool + .parameters + .as_ref() + .is_some_and(|parameters| !parameters.is_null()) + { + return None; + } + normalize_gemini_builtin_tool_name(&tool.name) +} + fn normalize_gemini_builtin_tool_name(name: &str) -> Option<&'static str> { match name .trim() @@ -1198,43 +1243,46 @@ mod tests { #[test] fn canonical_tool_declaration_sanitizes_json_schema_for_gemini() { - let declaration = canonical_tool_to_gemini_declaration(&CanonicalToolDefinition { - name: "inspect".to_string(), - description: None, - parameters: Some(json!({ - "$defs": { - "Target": { - "type": "object", - "properties": { - "secret": { - "type": "string", - "encrypted": true - } + let declaration = canonical_tool_to_gemini_declaration( + &CanonicalToolDefinition { + name: "inspect".to_string(), + description: None, + parameters: Some(json!({ + "$defs": { + "Target": { + "type": "object", + "properties": { + "secret": { + "type": "string", + "encrypted": true + } + }, + "required": ["secret"], + "additionalProperties": false + } + }, + "type": "object", + "properties": { + "target": { + "oneOf": [ + {"$ref": "#/$defs/Target"}, + {"type": "null"} + ] }, - "required": ["secret"], - "additionalProperties": false + "mode": { + "type": ["string", "null"], + "enum": [1, "fast"] + }, + "value": { + "type": ["string", "integer"] + } } - }, - "type": "object", - "properties": { - "target": { - "oneOf": [ - {"$ref": "#/$defs/Target"}, - {"type": "null"} - ] - }, - "mode": { - "type": ["string", "null"], - "enum": [1, "fast"] - }, - "value": { - "type": ["string", "integer"] - } - } - })), - strict: None, - extensions: BTreeMap::new(), - }); + })), + strict: None, + extensions: BTreeMap::new(), + }, + false, + ); assert_eq!( declaration["parameters"], @@ -1323,4 +1371,62 @@ mod tests { assert!(to_raw(&canonical, "gemini-2.5-pro", false).is_none()); assert!(to_raw(&canonical, "gemini-3-flash-preview", false).is_some()); } + + #[test] + fn client_declared_web_search_tool_stays_a_function_declaration() { + let canonical = CanonicalRequest { + model: "gemini-3-flash-preview".to_string(), + tools: vec![CanonicalToolDefinition { + name: "WebSearch".to_string(), + description: Some("Search the web".to_string()), + parameters: Some(json!({ + "type": "object", + "properties": {"query": {"type": "string"}}, + "required": ["query"], + })), + strict: None, + extensions: BTreeMap::new(), + }], + ..CanonicalRequest::default() + }; + + for preserve_tool_schemas in [false, true] { + let tools = canonical_tools_to_gemini(&canonical, preserve_tool_schemas) + .expect("tools should be emitted"); + let tools = tools.as_array().expect("tools should be an array"); + + assert!( + tools.iter().all(|tool| tool.get("googleSearch").is_none()), + "a client tool named WebSearch must not become server-side grounding: {tools:?}" + ); + assert_eq!( + tools[0]["functionDeclarations"][0]["name"], "WebSearch", + "the client declaration must survive: {tools:?}" + ); + } + } + + #[test] + fn schemaless_builtin_tool_name_still_maps_to_google_search() { + let canonical = CanonicalRequest { + model: "gemini-3-flash-preview".to_string(), + tools: vec![CanonicalToolDefinition { + name: "google_search".to_string(), + description: None, + parameters: None, + strict: None, + extensions: BTreeMap::new(), + }], + ..CanonicalRequest::default() + }; + + for preserve_tool_schemas in [false, true] { + let tools = canonical_tools_to_gemini(&canonical, preserve_tool_schemas) + .expect("tools should be emitted"); + let tools = tools.as_array().expect("tools should be an array"); + + assert_eq!(tools.len(), 1, "{tools:?}"); + assert_eq!(tools[0]["googleSearch"], json!({})); + } + } } diff --git a/crates/aether-ai/formats/src/formats/gemini/generate_content/response.rs b/crates/aether-ai/formats/src/formats/gemini/generate_content/response.rs index 5b4208272..d622feb49 100644 --- a/crates/aether-ai/formats/src/formats/gemini/generate_content/response.rs +++ b/crates/aether-ai/formats/src/formats/gemini/generate_content/response.rs @@ -2,15 +2,176 @@ use serde_json::{json, Map, Value}; use crate::{ formats::context::FormatContext, + formats::shared::citations::{ + canonical_citation, canonical_citations_to_claude_citations, + canonical_citations_to_openai_annotations, + }, protocol::canonical::{ canonical_extension_object_mut, canonical_usage_total_input_tokens, canonical_usage_total_tokens_for_inclusive_input, gemini_extensions, gemini_part_to_canonical_block, gemini_stop_reason_to_canonical, gemini_usage_to_canonical, CanonicalContentBlock, CanonicalResponse, CanonicalResponseOutput, CanonicalRole, - CanonicalStopReason, CanonicalUsage, + CanonicalStopReason, CanonicalUsage, CLAUDE_EXTENSION_NAMESPACE, + OPENAI_RESPONSES_EXTENSION_NAMESPACE, }, }; +/// Project Gemini grounding metadata onto the answer text as structured +/// citations. +/// +/// Native `googleSearch` grounding runs inside Google, so there is no +/// client-visible tool call and the evidence only exists in +/// `candidates[].groundingMetadata`. Cross-format targets used to drop that +/// wholesale, leaving callers with prose that names its sources but nothing a +/// client can render or verify. Every grounded span is therefore emitted twice, +/// each time in the target family's own standard shape: OpenAI `url_citation` +/// annotations and Claude `web_search_result_location` citations. Both ride +/// extension namespaces the respective emitters already merge onto the text +/// block, so no target has to learn anything Gemini-specific. +fn attach_gemini_grounding_citations( + candidate: &Map, + content: &mut [CanonicalContentBlock], +) { + let Some(grounding) = gemini_candidate_grounding(candidate) else { + return; + }; + let Some(block) = content.iter_mut().find(|block| { + matches!(block, CanonicalContentBlock::Text { text, .. } if !text.trim().is_empty()) + }) else { + return; + }; + let CanonicalContentBlock::Text { text, extensions } = block else { + return; + }; + + let citations = gemini_grounding_citations(grounding, text); + if citations.is_empty() { + return; + } + let annotations = canonical_citations_to_openai_annotations(&citations); + let claude_citations = canonical_citations_to_claude_citations(&citations); + canonical_extension_object_mut(extensions, OPENAI_RESPONSES_EXTENSION_NAMESPACE) + .entry("annotations".to_string()) + .or_insert_with(|| Value::Array(annotations)); + canonical_extension_object_mut(extensions, CLAUDE_EXTENSION_NAMESPACE) + .entry("citations".to_string()) + .or_insert_with(|| Value::Array(claude_citations)); +} + +pub(crate) fn gemini_candidate_grounding(candidate: &Map) -> Option<&Value> { + candidate + .get("groundingMetadata") + .or_else(|| candidate.get("grounding_metadata")) +} + +/// Normalise `groundingMetadata` into neutral citations against `text`. +/// +/// Gemini reports segment bounds as UTF-8 byte offsets while every target +/// counts characters, so the bounds are converted rather than copied. +pub(crate) fn gemini_grounding_citations(grounding: &Value, text: &str) -> Vec { + let chunks = grounding + .get("groundingChunks") + .and_then(Value::as_array) + .map(Vec::as_slice) + .unwrap_or_default(); + if chunks.is_empty() { + return Vec::new(); + } + + let supports = grounding + .get("groundingSupports") + .and_then(Value::as_array) + .map(Vec::as_slice) + .unwrap_or_default(); + + let mut citations = Vec::new(); + for support in supports { + let segment = support.get("segment"); + let start = segment + .and_then(|segment| segment.get("startIndex")) + .and_then(Value::as_u64) + .unwrap_or(0); + let end = segment + .and_then(|segment| segment.get("endIndex")) + .and_then(Value::as_u64); + let start_byte = gemini_clamped_byte_offset(text, start); + let end_byte = end + .map(|end| gemini_clamped_byte_offset(text, end)) + .filter(|end| *end >= start_byte); + let cited_text = segment + .and_then(|segment| segment.get("text")) + .and_then(Value::as_str) + .or_else(|| end_byte.map(|end| &text[start_byte..end])) + .map(str::trim) + .filter(|cited_text| !cited_text.is_empty()); + let indices = support + .get("groundingChunkIndices") + .and_then(Value::as_array) + .map(Vec::as_slice) + .unwrap_or_default(); + for index in indices { + let Some(chunk) = index + .as_u64() + .and_then(|index| usize::try_from(index).ok()) + .and_then(|index| chunks.get(index)) + else { + continue; + }; + let Some((uri, title)) = gemini_grounding_chunk_source(chunk) else { + continue; + }; + citations.push(canonical_citation( + uri, + title, + Some(text[..start_byte].chars().count()), + end_byte.map(|end| text[..end].chars().count()), + cited_text, + )); + } + } + + // `groundingSupports` is optional; without it the chunks are still the + // evidence, just unanchored. + if citations.is_empty() { + for chunk in chunks { + let Some((uri, title)) = gemini_grounding_chunk_source(chunk) else { + continue; + }; + citations.push(canonical_citation(uri, title, None, None, None)); + } + } + citations +} + +fn gemini_grounding_chunk_source(chunk: &Value) -> Option<(&str, Option<&str>)> { + let source = chunk.get("web").or_else(|| chunk.get("retrievedContext"))?; + let uri = source + .get("uri") + .or_else(|| source.get("url")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|uri| !uri.is_empty())?; + let title = source + .get("title") + .and_then(Value::as_str) + .map(str::trim) + .filter(|title| !title.is_empty()); + Some((uri, title)) +} + +/// Gemini offsets are byte counts into the UTF-8 answer. A truncated or stale +/// offset must not panic the conversion, so snap it into range and back onto a +/// character boundary. +fn gemini_clamped_byte_offset(text: &str, byte_offset: u64) -> usize { + let mut offset = usize::try_from(byte_offset) + .unwrap_or(text.len()) + .min(text.len()); + while offset > 0 && !text.is_char_boundary(offset) { + offset -= 1; + } + offset +} + pub fn from(body: &Value, _ctx: &FormatContext) -> Option { from_raw(body) } @@ -37,11 +198,12 @@ pub fn from_raw(body_json: &Value) -> Option { .and_then(Value::as_array) .map(Vec::as_slice) .unwrap_or(&[]); - let content = parts + let mut content = parts .iter() .enumerate() .filter_map(|(index, part)| gemini_part_to_canonical_block(part, index)) .collect::>(); + attach_gemini_grounding_citations(candidate_object, &mut content); let mut stop_reason = candidate_object .get("finishReason") .or_else(|| candidate_object.get("finish_reason")) @@ -433,6 +595,49 @@ mod tests { use super::*; use crate::CanonicalContentBlock; + /// Gemini omits `groundingSupports` when it cannot anchor the answer to a + /// span. The sources are still real, so they must survive unanchored + /// rather than be dropped for lacking offsets. + #[test] + fn grounding_without_supports_still_yields_unanchored_citations() { + let body = json!({ + "responseId": "resp-unanchored", + "candidates": [{ + "content": {"role": "model", "parts": [{"text": "Rust 1.95 is current."}]}, + "finishReason": "STOP", + "groundingMetadata": { + "groundingChunks": [ + {"web": {"uri": "https://blog.rust-lang.org/", "title": "Rust Blog"}}, + {"web": {"title": "no uri here"}} + ] + } + }] + }); + + let canonical = from_raw(&body).expect("canonical"); + let CanonicalContentBlock::Text { extensions, .. } = &canonical.outputs[0].content[0] + else { + panic!("expected a text block"); + }; + + assert_eq!( + extensions["claude"]["citations"], + json!([{ + "type": "web_search_result_location", + "url": "https://blog.rust-lang.org/", + "title": "Rust Blog", + }]) + ); + assert_eq!( + extensions["openai_responses"]["annotations"], + json!([{ + "type": "url_citation", + "url": "https://blog.rust-lang.org/", + "title": "Rust Blog", + }]) + ); + } + #[test] fn gemini_response_without_visible_parts_is_not_success() { let body = json!({ diff --git a/crates/aether-ai/formats/src/formats/gemini/generate_content/stream.rs b/crates/aether-ai/formats/src/formats/gemini/generate_content/stream.rs index de5b46e6d..b751794d4 100644 --- a/crates/aether-ai/formats/src/formats/gemini/generate_content/stream.rs +++ b/crates/aether-ai/formats/src/formats/gemini/generate_content/stream.rs @@ -2,6 +2,9 @@ use std::collections::BTreeMap; use serde_json::{json, Map, Value}; +use crate::formats::gemini::generate_content::response::{ + gemini_candidate_grounding, gemini_grounding_citations, +}; use crate::formats::shared::response::{build_generated_tool_call_id, canonicalize_tool_arguments}; use crate::formats::shared::sse::encode_json_sse; use crate::formats::shared::stream_core::common::*; @@ -36,6 +39,10 @@ pub struct GeminiProviderState { content_parts: BTreeMap, tool_calls: BTreeMap, tool_results: BTreeMap, + /// Last `groundingMetadata` seen. Gemini resends it cumulatively, so the + /// newest copy is the complete one; citations are emitted once at finish, + /// when the answer text they index into is whole. + grounding: Option, } impl GeminiProviderState { @@ -68,6 +75,31 @@ impl GeminiProviderState { self.started = true; } + /// Turn the grounding metadata collected over the stream into citations. + /// + /// The offsets Gemini reports index into the finished answer, so this can + /// only run once the text is complete — hence a single frame just ahead of + /// `Finish` rather than a delta per chunk. + fn push_citations_frame(&mut self, id: &str, model: &str, out: &mut Vec) { + let Some(grounding) = self.grounding.take() else { + return; + }; + let text = self + .text_parts + .values() + .map(String::as_str) + .collect::(); + let citations = gemini_grounding_citations(&grounding, &text); + if citations.is_empty() { + return; + } + out.push(CanonicalStreamFrame { + id: id.to_string(), + model: model.to_string(), + event: CanonicalStreamEvent::Citations(citations), + }); + } + fn unknown_frame(&self, report_context: &Value, payload: Value) -> CanonicalStreamFrame { let (id, model) = self.identity(report_context); CanonicalStreamFrame { @@ -120,6 +152,11 @@ impl GeminiProviderState { response_model.as_str(), event_object.get("usageMetadata"), ); + if !self.terminal_observation_only { + if let Some(grounding) = gemini_candidate_grounding(candidate_object) { + self.grounding = Some(grounding.clone()); + } + } let Some(content) = candidate_object.get("content").and_then(Value::as_object) else { if let Some(payload) = terminal_error { out.push(self.unknown_frame(report_context, payload)); @@ -286,6 +323,12 @@ impl GeminiProviderState { self.observed_tool_calls = true; continue; } + // Gemini streams are incremental and every functionCall part is a + // complete call, so parallel calls arriving in separate chunks all + // sit at parts[0]. Key calls by arrival order, not part position. + // Ids cannot disambiguate: they are optional, and the Antigravity + // envelope synthesizes per-chunk ids that repeat across chunks. + let index = self.tool_calls.len(); let tool_state = self.tool_calls.entry(index).or_default(); tool_state.call_id = function_call .get("id") @@ -361,6 +404,7 @@ impl GeminiProviderState { if has_tool_calls && finish_reason.as_deref().is_none_or(|value| value == "stop") { finish_reason = Some("tool_calls".to_string()); } + self.push_citations_frame(&id, &model, &mut out); out.push(CanonicalStreamFrame { id, model, @@ -385,14 +429,17 @@ impl GeminiProviderState { } self.finished = true; let (id, model) = self.identity(report_context); - Ok(vec![CanonicalStreamFrame { + let mut out = Vec::new(); + self.push_citations_frame(&id, &model, &mut out); + out.push(CanonicalStreamFrame { id, model, event: CanonicalStreamEvent::Finish { finish_reason: None, usage: None, }, - }]) + }); + Ok(out) } } @@ -674,6 +721,9 @@ impl GeminiClientEmitter { None, None, ), + // Only Gemini produces citations today, and a Gemini-to-Gemini + // stream keeps its own `groundingMetadata` on the passthrough path. + CanonicalStreamEvent::Citations(_) => Ok(Vec::new()), CanonicalStreamEvent::UnknownEvent(_) => Ok(Vec::new()), CanonicalStreamEvent::Finish { finish_reason, @@ -1482,6 +1532,80 @@ mod tests { assert!(signature_index < call_index); } + #[test] + fn gemini_provider_state_keeps_parallel_function_calls_from_separate_chunks() { + let mut state = GeminiProviderState::default(); + let report_context = json!({}); + let chunk = |call: Value| { + data_line(json!({ + "responseId": "resp_parallel_123", + "modelVersion": "gemini-3.8-flash", + "candidates": [{ + "index": 0, + "content": {"role": "model", "parts": [{"functionCall": call}]} + }] + })) + }; + let mut frames = Vec::new(); + for call in [ + json!({"id": "call_a", "name": "get_weather", "args": {"city": "Paris"}}), + json!({"id": "call_b", "name": "get_weather", "args": {"city": "Tokyo"}}), + json!({"name": "get_time", "args": {"city": "Paris"}}), + json!({"name": "get_time", "args": {"city": "Tokyo"}}), + ] { + frames.extend( + state + .push_line(&report_context, chunk(call)) + .expect("function call chunk should parse"), + ); + } + + let starts = frames + .iter() + .filter_map(|frame| match &frame.event { + CanonicalStreamEvent::ToolCallStart { + index, + call_id, + name, + } => Some((*index, call_id.clone(), name.clone())), + _ => None, + }) + .collect::>(); + assert_eq!(starts.len(), 4); + assert_eq!( + starts + .iter() + .map(|(index, _, _)| *index) + .collect::>(), + vec![0, 1, 2, 3] + ); + assert_eq!(starts[0].1, "call_a"); + assert_eq!(starts[1].1, "call_b"); + assert_eq!( + starts + .iter() + .map(|(_, _, name)| name.as_str()) + .collect::>(), + vec!["get_weather", "get_weather", "get_time", "get_time"] + ); + assert_ne!(starts[2].1, starts[3].1); + + let mut arguments = BTreeMap::::new(); + for frame in &frames { + if let CanonicalStreamEvent::ToolCallArgumentsDelta { + index, + arguments: delta, + } = &frame.event + { + arguments.entry(*index).or_default().push_str(delta); + } + } + assert_eq!(arguments[&0], "{\"city\":\"Paris\"}"); + assert_eq!(arguments[&1], "{\"city\":\"Tokyo\"}"); + assert_eq!(arguments[&2], "{\"city\":\"Paris\"}"); + assert_eq!(arguments[&3], "{\"city\":\"Tokyo\"}"); + } + #[test] fn gemini_client_emitter_marks_reasoning_parts_as_thoughts() { let mut emitter = GeminiClientEmitter::default(); diff --git a/crates/aether-ai/formats/src/formats/openai/chat/response.rs b/crates/aether-ai/formats/src/formats/openai/chat/response.rs index 69d503169..64a6589a2 100644 --- a/crates/aether-ai/formats/src/formats/openai/chat/response.rs +++ b/crates/aether-ai/formats/src/formats/openai/chat/response.rs @@ -1,6 +1,6 @@ use std::collections::BTreeMap; -use serde_json::{json, Value}; +use serde_json::{json, Map, Value}; use crate::{ formats::context::FormatContext, @@ -13,6 +13,53 @@ use crate::{ }, }; +/// Reasoning text carried by one Chat Completions `message` or streaming +/// `delta`, paired with the provider's reasoning block index where one exists. +/// +/// The field name is not standardized. DeepSeek-style upstreams send +/// `reasoning_content`; OpenRouter sends `reasoning` alongside a structured +/// `reasoning_details` array. OpenRouter repeats the same text in both of its +/// fields, so exactly one source is read per object and `reasoning_details` +/// wins because only it carries the block index. +pub(crate) fn openai_chat_reasoning_texts( + object: &Map, +) -> Vec<(Option, String)> { + if let Some(details) = object.get("reasoning_details").and_then(Value::as_array) { + let texts = details + .iter() + .filter_map(Value::as_object) + .filter_map(|detail| { + // `reasoning.encrypted` carries opaque provider state rather + // than readable text, so it has nothing to hand downstream. + if detail.get("type").and_then(Value::as_str) == Some("reasoning.encrypted") { + return None; + } + let text = detail + .get("text") + .or_else(|| detail.get("summary")) + .and_then(Value::as_str) + .filter(|text| !text.is_empty())?; + let index = detail + .get("index") + .and_then(Value::as_u64) + .map(|index| index as usize); + Some((index, text.to_string())) + }) + .collect::>(); + if !texts.is_empty() { + return texts; + } + } + // A provider may null out one spelling while filling the other, so skip + // past any key that is present but carries no string. + ["reasoning_content", "reasoning"] + .iter() + .find_map(|key| object.get(*key).and_then(Value::as_str)) + .filter(|text| !text.is_empty()) + .map(|text| vec![(None, text.to_string())]) + .unwrap_or_default() +} + pub fn from(body: &Value, _ctx: &FormatContext) -> Option { from_raw(body) } @@ -40,21 +87,18 @@ pub fn from_raw(body_json: &Value) -> Option { .iter() .any(|block| matches!(block, CanonicalContentBlock::Thinking { .. })) { - if let Some(reasoning_content) = message - .get("reasoning_content") - .and_then(Value::as_str) - .filter(|value| !value.trim().is_empty()) - { - content.insert( - 0, - CanonicalContentBlock::Thinking { - text: reasoning_content.to_string(), - signature: None, - encrypted_content: None, - extensions: BTreeMap::new(), - }, - ); - } + let thinking = openai_chat_reasoning_texts(message) + .into_iter() + .map(|(_, text)| text) + .filter(|text| !text.trim().is_empty()) + .map(|text| CanonicalContentBlock::Thinking { + text, + signature: None, + encrypted_content: None, + extensions: BTreeMap::new(), + }) + .collect::>(); + content.splice(0..0, thinking); } let stop_reason = openai_finish_reason_to_canonical(choice.get("finish_reason").and_then(Value::as_str)); @@ -179,3 +223,161 @@ pub fn to_raw(canonical: &CanonicalResponse) -> Value { } response } + +#[cfg(test)] +mod tests { + use super::*; + use crate::protocol::canonical::CanonicalContentBlock; + + fn thinking_texts(response: &CanonicalResponse) -> Vec { + response + .content + .iter() + .filter_map(|block| match block { + CanonicalContentBlock::Thinking { text, .. } => Some(text.clone()), + _ => None, + }) + .collect() + } + + #[test] + fn openrouter_reasoning_details_become_thinking_blocks() { + let response = from_raw(&json!({ + "id": "gen-openrouter-123", + "model": "stealth/ox-alpha", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "done", + "reasoning": "step onestep two", + "reasoning_details": [ + {"type": "reasoning.text", "text": "step one", "index": 0}, + {"type": "reasoning.text", "text": "step two", "index": 1} + ] + }, + "finish_reason": "stop" + }] + })) + .expect("openrouter response should convert"); + + // `reasoning` repeats the same text the details already carry, so the + // details win and the provider's own segmentation survives. + assert_eq!(thinking_texts(&response), vec!["step one", "step two"]); + } + + #[test] + fn openrouter_reasoning_string_becomes_a_thinking_block() { + let response = from_raw(&json!({ + "id": "gen-openrouter-123", + "model": "stealth/ox-alpha", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "done", + "reasoning": "thought about it" + }, + "finish_reason": "stop" + }] + })) + .expect("openrouter response should convert"); + + assert_eq!(thinking_texts(&response), vec!["thought about it"]); + } + + #[test] + fn deepseek_reasoning_content_still_becomes_a_thinking_block() { + let response = from_raw(&json!({ + "id": "chatcmpl-deepseek", + "model": "deepseek-reasoner", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "42", + "reasoning_content": "let me work it out" + }, + "finish_reason": "stop" + }] + })) + .expect("deepseek response should convert"); + + assert_eq!(thinking_texts(&response), vec!["let me work it out"]); + } + + #[test] + fn deepseek_reasoning_content_wins_over_a_bare_reasoning_field() { + let response = from_raw(&json!({ + "id": "chatcmpl-deepseek", + "model": "deepseek-reasoner", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "42", + "reasoning_content": "the real one", + "reasoning": "the other spelling" + }, + "finish_reason": "stop" + }] + })) + .expect("deepseek response should convert"); + + assert_eq!(thinking_texts(&response), vec!["the real one"]); + } + + #[test] + fn blank_reasoning_content_produces_no_thinking_block() { + let response = from_raw(&json!({ + "id": "chatcmpl-deepseek", + "model": "deepseek-reasoner", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "42", "reasoning_content": " "}, + "finish_reason": "stop" + }] + })) + .expect("deepseek response should convert"); + + assert!(thinking_texts(&response).is_empty()); + } + + #[test] + fn plain_openai_response_without_reasoning_is_unchanged() { + let response = from_raw(&json!({ + "id": "chatcmpl-openai", + "model": "gpt-4o", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "hello"}, + "finish_reason": "stop" + }] + })) + .expect("openai response should convert"); + + assert!(thinking_texts(&response).is_empty()); + } + + #[test] + fn encrypted_reasoning_details_carry_no_thinking_text() { + let response = from_raw(&json!({ + "id": "gen-openrouter-123", + "model": "stealth/ox-alpha", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "done", + "reasoning_details": [ + {"type": "reasoning.encrypted", "data": "b3BhcXVl", "index": 0} + ] + }, + "finish_reason": "stop" + }] + })) + .expect("openrouter response should convert"); + + assert!(thinking_texts(&response).is_empty()); + } +} diff --git a/crates/aether-ai/formats/src/formats/openai/chat/stream.rs b/crates/aether-ai/formats/src/formats/openai/chat/stream.rs index a5a29524e..87d65f44f 100644 --- a/crates/aether-ai/formats/src/formats/openai/chat/stream.rs +++ b/crates/aether-ai/formats/src/formats/openai/chat/stream.rs @@ -1,17 +1,20 @@ -use std::collections::{BTreeMap, BTreeSet}; +use std::collections::{BTreeMap, BTreeSet, VecDeque}; use serde_json::{json, Map, Value}; use sha2::{Digest, Sha256}; +use crate::formats::openai::chat::response::openai_chat_reasoning_texts; use crate::formats::openai::namespace::NamespaceToolAliases; use crate::formats::openai::responses::{ - encode_gemini_tool_signature_carrier_with_direction, openai_responses_message_item_id, + decode_gemini_tool_signature_carrier, encode_gemini_tool_signature_carrier_with_direction, + openai_responses_message_item_id, openai_responses_reasoning_text_parts, openai_responses_synthetic_reasoning_item_id, response::{ ensure_modern_openai_responses_response_fields, openai_responses_current_timestamp, }, GeminiToolSignatureCarrierDirection, }; +use crate::formats::shared::citations::canonical_citations_to_openai_annotations; use crate::formats::shared::response::build_generated_tool_call_id; use crate::formats::shared::sse::{encode_done_sse, encode_json_sse}; use crate::formats::shared::stream_core::common::*; @@ -41,6 +44,7 @@ pub struct OpenAIChatProviderState { started: bool, finished: bool, pending_finish_reason: Option, + last_reasoning_index: Option, tool_calls: BTreeMap, } @@ -75,6 +79,8 @@ pub struct OpenAIResponsesProviderState { tool_index_by_key: BTreeMap, image_item_keys: BTreeSet, opaque_completed_item_keys: BTreeSet, + seen_tool_signature_carriers: BTreeSet, + pending_tool_signatures: VecDeque, last_tool_index: Option, } @@ -271,24 +277,39 @@ impl OpenAIChatProviderState { } else if delta.contains_key("content") { recognized_delta = true; } - if let Some(reasoning_content) = delta.get("reasoning_content").and_then(Value::as_str) + if delta.contains_key("reasoning_content") + || delta.contains_key("reasoning_details") + || delta.contains_key("reasoning") { recognized_delta = true; - if !reasoning_content.is_empty() { + for (reasoning_index, text) in openai_chat_reasoning_texts(delta) { self.ensure_started(report_context, &mut out); if !self.terminal_only { let (id, model) = self.identity(report_context); + // A change of reasoning block index closes the part that + // is open, so downstream summaries keep the provider's + // own segmentation instead of collapsing into one + // paragraph. + if let Some(reasoning_index) = reasoning_index { + if self + .last_reasoning_index + .is_some_and(|last| last != reasoning_index) + { + out.push(CanonicalStreamFrame { + id: id.clone(), + model: model.clone(), + event: CanonicalStreamEvent::ReasoningSummaryDone, + }); + } + self.last_reasoning_index = Some(reasoning_index); + } out.push(CanonicalStreamFrame { id, model, - event: CanonicalStreamEvent::ReasoningDelta( - reasoning_content.to_string(), - ), + event: CanonicalStreamEvent::ReasoningDelta(text), }); } } - } else if delta.contains_key("reasoning_content") { - recognized_delta = true; } if let Some(tool_calls) = delta.get("tool_calls").and_then(Value::as_array) { @@ -852,6 +873,7 @@ impl OpenAIResponsesProviderState { return; } }; + self.emit_pending_tool_signature(report_context, out, index); let state = self.tool_calls.entry(index).or_default(); state.call_id = item .get("call_id") @@ -1186,6 +1208,68 @@ impl OpenAIResponsesProviderState { } } + fn capture_tool_signature_carrier( + &mut self, + report_context: &Value, + out: &mut Vec, + item: &Map, + ) { + if self.terminal_only { + return; + } + let Some(carrier) = item + .get("encrypted_content") + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) + else { + return; + }; + let Some((signature, direction)) = decode_gemini_tool_signature_carrier(carrier) else { + return; + }; + if !self + .seen_tool_signature_carriers + .insert(Self::output_item_key(item)) + { + return; + } + match direction { + GeminiToolSignatureCarrierDirection::Next => { + self.pending_tool_signatures.push_back(signature); + } + GeminiToolSignatureCarrierDirection::Previous => { + let Some(index) = self.last_tool_index else { + return; + }; + self.ensure_started(report_context, out); + let (id, model) = self.identity(report_context); + out.push(CanonicalStreamFrame { + id, + model, + event: CanonicalStreamEvent::ToolCallSignature { index, signature }, + }); + } + } + } + + fn emit_pending_tool_signature( + &mut self, + report_context: &Value, + out: &mut Vec, + index: usize, + ) { + let Some(signature) = self.pending_tool_signatures.pop_front() else { + return; + }; + self.ensure_started(report_context, out); + let (id, model) = self.identity(report_context); + out.push(CanonicalStreamFrame { + id, + model, + event: CanonicalStreamEvent::ToolCallSignature { index, signature }, + }); + } + fn emit_reasoning_item( &mut self, report_context: &Value, @@ -1195,40 +1279,13 @@ impl OpenAIResponsesProviderState { if item.get("type").and_then(Value::as_str) != Some("reasoning") { return; } + let completed_reasoning = reasoning_item_text(item); if self.terminal_only { - if item - .get("summary") - .and_then(Value::as_array) - .is_some_and(|summary| { - summary.iter().any(|part| { - part.get("type").and_then(Value::as_str) == Some("summary_text") - && part - .get("text") - .and_then(Value::as_str) - .is_some_and(|text| !text.is_empty()) - }) - }) - { + if !completed_reasoning.is_empty() { self.ensure_started(report_context, out); } return; } - let mut completed_reasoning = String::new(); - for raw_summary in item - .get("summary") - .and_then(Value::as_array) - .into_iter() - .flatten() - { - let Some(summary) = raw_summary.as_object() else { - continue; - }; - if summary.get("type").and_then(Value::as_str) == Some("summary_text") { - if let Some(text) = summary.get("text").and_then(Value::as_str) { - completed_reasoning.push_str(text); - } - } - } if !completed_reasoning.is_empty() { self.emit_missing_reasoning(report_context, out, &completed_reasoning); } @@ -1345,7 +1402,14 @@ impl OpenAIResponsesProviderState { output_index: Option, final_item: bool, ) -> bool { - match item.get("type").and_then(Value::as_str).unwrap_or_default() { + let item_type = item.get("type").and_then(Value::as_str).unwrap_or_default(); + if item_type == "reasoning" { + // Aether carries Gemini function-call signatures through Responses as + // encrypted reasoning items. Recover the carrier before the following + // function item is emitted so Gemini clients can replay it verbatim. + self.capture_tool_signature_carrier(report_context, out, item); + } + match item_type { "function_call" => { self.emit_tool_call_item(report_context, out, item, output_index); true @@ -1645,7 +1709,8 @@ impl OpenAIResponsesProviderState { .unwrap_or_default(); if !piece.is_empty() { let summary_index = value - .get("summary_index") + .get("content_index") + .or_else(|| value.get("summary_index")) .and_then(Value::as_u64) .map(|value| value as usize) .unwrap_or(0); @@ -1680,7 +1745,8 @@ impl OpenAIResponsesProviderState { .unwrap_or_default(); if !text.is_empty() { let summary_index = value - .get("summary_index") + .get("content_index") + .or_else(|| value.get("summary_index")) .and_then(Value::as_u64) .map(|value| value as usize) .unwrap_or(0); @@ -2077,6 +2143,29 @@ impl OpenAIResponsesProviderState { } } +/// Reads a Responses reasoning item's raw chain-of-thought. +/// +/// Raw thinking lives on `content` (`reasoning_text` parts); `summary` is the +/// summarised view and is only consulted when `content` carries nothing, so +/// items produced by other Aether versions still yield their thinking. +fn reasoning_item_text(item: &Map) -> String { + let mut text = reasoning_item_parts_text(item.get("content"), "reasoning_text"); + if text.is_empty() { + text = reasoning_item_parts_text(item.get("summary"), "summary_text"); + } + text +} + +fn reasoning_item_parts_text(raw: Option<&Value>, expected_type: &str) -> String { + raw.and_then(Value::as_array) + .into_iter() + .flatten() + .filter_map(Value::as_object) + .filter(|part| part.get("type").and_then(Value::as_str) == Some(expected_type)) + .filter_map(|part| part.get("text").and_then(Value::as_str)) + .collect::() +} + #[derive(Default)] pub struct OpenAIChatClientEmitter { response_id: Option, @@ -2144,6 +2233,9 @@ pub struct OpenAIResponsesClientEmitter { text_item_started: bool, text_part_started: bool, message_output_index: Option, + /// Citations projected onto the answer text, kept on the finished message + /// item so non-incremental clients see them too. + annotations: Vec, text: String, reasoning: String, reasoning_part: String, @@ -2269,6 +2361,28 @@ impl OpenAIChatClientEmitter { }))?); Ok(out) } + CanonicalStreamEvent::Citations(citations) => { + let annotations = canonical_citations_to_openai_annotations(&citations); + if annotations.is_empty() { + return Ok(Vec::new()); + } + let mut out = self.ensure_started()?; + out.extend(self.encode_chunk(json!({ + "id": self.response_id + .as_deref() + .unwrap_or("chatcmpl-local-stream"), + "object": "chat.completion.chunk", + "model": self.model.as_deref().unwrap_or("unknown"), + "choices": [{ + "index": 0, + "delta": { + "annotations": annotations, + }, + "finish_reason": Value::Null + }] + }))?); + Ok(out) + } CanonicalStreamEvent::ReasoningSignature(_) => Ok(Vec::new()), CanonicalStreamEvent::ContentPart(part) => { let placeholder = openai_stream_placeholder_for_content_part(&part); @@ -2621,6 +2735,71 @@ impl OpenAIResponsesClientEmitter { self.reasoning_summary_parts.len() } + fn reasoning_texts(&self) -> Vec { + if self.reasoning_summary_parts.is_empty() { + if self.reasoning.trim().is_empty() { + Vec::new() + } else { + vec![self.reasoning.clone()] + } + } else { + self.reasoning_summary_parts.clone() + } + } + + fn reasoning_item_value(&self) -> Value { + let content = openai_responses_reasoning_text_parts(self.reasoning_texts()); + json!({ + "type": "reasoning", + "id": self.reasoning_item_id(), + "status": "completed", + // Raw thinking goes on `content` only. Mirroring it onto `summary` + // makes Codex (which renders both channels) print it twice. + "summary": [], + "content": content, + }) + } + + fn encode_reasoning_text_delta( + &mut self, + text: &str, + ) -> Result, AiSurfaceFinalizeError> { + let item_id = self.reasoning_item_id(); + let output_index = self.reasoning_output_index.unwrap_or(0); + let part_index = self.current_reasoning_summary_index(); + self.encode_response_event( + "response.reasoning_text.delta", + json!({ + "type": "response.reasoning_text.delta", + "response_id": self.response_id(), + "item_id": item_id, + "output_index": output_index, + "content_index": part_index, + "delta": text, + }), + ) + } + + fn encode_reasoning_text_done_events( + &mut self, + item_id: &str, + output_index: usize, + part_index: usize, + part_text: &str, + ) -> Result, AiSurfaceFinalizeError> { + self.encode_response_event( + "response.reasoning_text.done", + json!({ + "type": "response.reasoning_text.done", + "response_id": self.response_id(), + "item_id": item_id, + "output_index": output_index, + "content_index": part_index, + "text": part_text, + }), + ) + } + fn ensure_message_output_index(&mut self) -> usize { if let Some(output_index) = self.message_output_index { return output_index; @@ -2676,27 +2855,13 @@ impl OpenAIResponsesClientEmitter { "type": "reasoning", "id": item_id.clone(), "summary": [], + "content": [], } }), )?); self.reasoning_item_started = true; } if !self.reasoning_part_started { - let summary_index = self.current_reasoning_summary_index(); - out.extend(self.encode_response_event( - "response.reasoning_summary_part.added", - json!({ - "type": "response.reasoning_summary_part.added", - "response_id": self.response_id(), - "item_id": item_id, - "output_index": output_index, - "summary_index": summary_index, - "part": { - "type": "summary_text", - "text": "", - } - }), - )?); self.reasoning_part_started = true; } Ok(out) @@ -2775,7 +2940,7 @@ impl OpenAIResponsesClientEmitter { "part": { "type": "output_text", "text": self.text.as_str(), - "annotations": [], + "annotations": self.annotations.as_slice(), } }), )?); @@ -2794,7 +2959,7 @@ impl OpenAIResponsesClientEmitter { "content": [{ "type": "output_text", "text": self.text.as_str(), - "annotations": [], + "annotations": self.annotations.as_slice(), }], } }), @@ -2812,66 +2977,23 @@ impl OpenAIResponsesClientEmitter { if self.reasoning_part_started { let summary_index = self.current_reasoning_summary_index(); let part_text = self.reasoning_part.clone(); - out.extend(self.encode_response_event( - "response.reasoning_summary_text.done", - json!({ - "type": "response.reasoning_summary_text.done", - "response_id": self.response_id(), - "item_id": item_id.clone(), - "output_index": output_index, - "summary_index": summary_index, - "text": part_text.as_str(), - }), - )?); - out.extend(self.encode_response_event( - "response.reasoning_summary_part.done", - json!({ - "type": "response.reasoning_summary_part.done", - "response_id": self.response_id(), - "item_id": item_id.clone(), - "output_index": output_index, - "summary_index": summary_index, - "part": { - "type": "summary_text", - "text": part_text.as_str(), - } - }), + out.extend(self.encode_reasoning_text_done_events( + &item_id, + output_index, + summary_index, + part_text.as_str(), )?); self.reasoning_summary_parts.push(part_text); self.reasoning_part.clear(); self.reasoning_part_started = false; } - let summary = if self.reasoning_summary_parts.is_empty() { - if self.reasoning.trim().is_empty() { - Vec::new() - } else { - vec![json!({ - "type": "summary_text", - "text": self.reasoning.as_str(), - })] - } - } else { - self.reasoning_summary_parts - .iter() - .map(|text| { - json!({ - "type": "summary_text", - "text": text, - }) - }) - .collect::>() - }; out.extend(self.encode_response_event( "response.output_item.done", json!({ "type": "response.output_item.done", "response_id": self.response_id(), "output_index": output_index, - "item": { - "type": "reasoning", - "id": item_id, - "summary": summary, - } + "item": self.reasoning_item_value(), }), )?); Ok(out) @@ -3035,35 +3157,10 @@ impl OpenAIResponsesClientEmitter { incomplete_reason: Option<&str>, ) -> Value { let mut ordered_output = Vec::new(); - let summary = if self.reasoning_summary_parts.is_empty() { - if self.reasoning.trim().is_empty() { - Vec::new() - } else { - vec![json!({ - "type": "summary_text", - "text": self.reasoning.as_str(), - })] - } - } else { - self.reasoning_summary_parts - .iter() - .map(|text| { - json!({ - "type": "summary_text", - "text": text, - }) - }) - .collect::>() - }; - if !summary.is_empty() { + if !self.reasoning_texts().is_empty() { ordered_output.push(( self.reasoning_output_index.unwrap_or(0), - json!({ - "type": "reasoning", - "id": self.reasoning_item_id(), - "status": "completed", - "summary": summary, - }), + self.reasoning_item_value(), )); } if self.text_item_started || !self.text.is_empty() { @@ -3077,7 +3174,7 @@ impl OpenAIResponsesClientEmitter { "content": [{ "type": "output_text", "text": self.text.as_str(), - "annotations": [], + "annotations": self.annotations.as_slice(), }], }), )); @@ -3297,17 +3394,36 @@ impl OpenAIResponsesClientEmitter { let mut out = self.ensure_reasoning_item_started()?; self.reasoning.push_str(&text); self.reasoning_part.push_str(&text); - out.extend(self.encode_response_event( - "response.reasoning_summary_text.delta", - json!({ - "type": "response.reasoning_summary_text.delta", - "response_id": self.response_id(), - "item_id": self.reasoning_item_id(), - "output_index": self.reasoning_output_index.unwrap_or(0), - "summary_index": self.current_reasoning_summary_index(), - "delta": text, - }), - )?); + out.extend(self.encode_reasoning_text_delta(&text)?); + Ok(out) + } + CanonicalStreamEvent::Citations(citations) => { + let annotations = canonical_citations_to_openai_annotations(&citations); + if annotations.is_empty() { + return Ok(Vec::new()); + } + // The text item has to exist before an annotation can point at + // it, and the annotations are also kept on the item itself so + // clients that only read `response.completed` still see them. + let mut out = self.ensure_text_item_started()?; + let item_id = self.message_item_id(); + let output_index = self.message_output_index.unwrap_or(0); + for annotation in annotations { + let annotation_index = self.annotations.len(); + self.annotations.push(annotation.clone()); + out.extend(self.encode_response_event( + "response.output_text.annotation.added", + json!({ + "type": "response.output_text.annotation.added", + "response_id": self.response_id(), + "output_index": output_index, + "item_id": item_id, + "content_index": 0, + "annotation_index": annotation_index, + "annotation": annotation, + }), + )?); + } Ok(out) } CanonicalStreamEvent::ReasoningSummaryDone => { @@ -3320,32 +3436,12 @@ impl OpenAIResponsesClientEmitter { let item_id = self.reasoning_item_id(); let summary_index = self.current_reasoning_summary_index(); let part_text = self.reasoning_part.clone(); - let mut out = Vec::new(); - out.extend(self.encode_response_event( - "response.reasoning_summary_text.done", - json!({ - "type": "response.reasoning_summary_text.done", - "response_id": self.response_id(), - "item_id": item_id.clone(), - "output_index": output_index, - "summary_index": summary_index, - "text": part_text.as_str(), - }), - )?); - out.extend(self.encode_response_event( - "response.reasoning_summary_part.done", - json!({ - "type": "response.reasoning_summary_part.done", - "response_id": self.response_id(), - "item_id": item_id, - "output_index": output_index, - "summary_index": summary_index, - "part": { - "type": "summary_text", - "text": part_text.as_str(), - } - }), - )?); + let out = self.encode_reasoning_text_done_events( + &item_id, + output_index, + summary_index, + part_text.as_str(), + )?; self.reasoning_summary_parts.push(part_text); self.reasoning_part.clear(); self.reasoning_part_started = false; @@ -3908,6 +4004,7 @@ fn openai_responses_incomplete_finish_reason(payload: &Value) -> String { mod tests { use super::*; use crate::formats::claude::messages::stream::ClaudeClientEmitter; + use crate::formats::gemini::generate_content::stream::GeminiClientEmitter; use crate::formats::openai::responses::encode_gemini_tool_signature_carrier; fn data_line(value: Value) -> Vec { @@ -4215,7 +4312,7 @@ mod tests { data = Some(value); } } - if event_name != Some("response.reasoning_summary_text.done") { + if event_name != Some("response.reasoning_text.done") { continue; } let Some(data) = data else { @@ -4224,13 +4321,17 @@ mod tests { let Ok(value) = serde_json::from_str::(data) else { continue; }; - let Some(summary_index) = value.get("summary_index").and_then(Value::as_u64) else { + let Some(part_index) = value + .get("content_index") + .or_else(|| value.get("summary_index")) + .and_then(Value::as_u64) + else { continue; }; let Some(text) = value.get("text").and_then(Value::as_str) else { continue; }; - parts.push((summary_index, text.to_string())); + parts.push((part_index, text.to_string())); } parts } @@ -4287,6 +4388,176 @@ mod tests { ))); } + #[test] + fn openai_chat_provider_state_reads_openrouter_reasoning_once() { + let mut state = OpenAIChatProviderState::default(); + let report_context = json!({}); + let frames = state + .push_line( + &report_context, + data_line(json!({ + "id": "gen-openrouter-123", + "model": "stealth/ox-alpha", + "choices": [{ + "index": 0, + "delta": { + "content": "", + "role": "assistant", + "reasoning": " me translate the", + "reasoning_details": [{ + "type": "reasoning.text", + "text": " me translate the", + "format": "unknown", + "index": 0 + }] + }, + "finish_reason": Value::Null + }] + })), + ) + .expect("openrouter reasoning delta should parse"); + + let reasoning = frames + .iter() + .filter_map(|frame| match frame.event { + CanonicalStreamEvent::ReasoningDelta(ref text) => Some(text.as_str()), + _ => None, + }) + .collect::>(); + // `reasoning` and `reasoning_details` repeat the same text, so only one + // of the two may reach the client. + assert_eq!(reasoning, vec![" me translate the"]); + assert!(!frames + .iter() + .any(|frame| matches!(frame.event, CanonicalStreamEvent::UnknownEvent(_)))); + } + + #[test] + fn openai_chat_provider_state_still_reads_deepseek_reasoning_content() { + let mut state = OpenAIChatProviderState::default(); + let report_context = json!({}); + let frames = state + .push_line( + &report_context, + data_line(json!({ + "id": "chatcmpl-deepseek", + "model": "deepseek-reasoner", + "choices": [{ + "index": 0, + "delta": {"role": "assistant", "reasoning_content": "let me work it out"}, + "finish_reason": Value::Null + }] + })), + ) + .expect("deepseek reasoning delta should parse"); + + let reasoning = frames + .iter() + .filter_map(|frame| match frame.event { + CanonicalStreamEvent::ReasoningDelta(ref text) => Some(text.as_str()), + _ => None, + }) + .collect::>(); + assert_eq!(reasoning, vec!["let me work it out"]); + // A single reasoning block carries no index, so nothing may close a part. + assert!(!frames + .iter() + .any(|frame| matches!(frame.event, CanonicalStreamEvent::ReasoningSummaryDone))); + assert!(!frames + .iter() + .any(|frame| matches!(frame.event, CanonicalStreamEvent::UnknownEvent(_)))); + } + + #[test] + fn openai_chat_provider_state_splits_reasoning_details_on_block_index() { + let mut state = OpenAIChatProviderState::default(); + let report_context = json!({}); + let reasoning_chunk = |index: u64, text: &str| { + data_line(json!({ + "id": "gen-openrouter-123", + "model": "stealth/ox-alpha", + "choices": [{ + "index": 0, + "delta": { + "content": "", + "role": "assistant", + "reasoning_details": [{ + "type": "reasoning.text", + "text": text, + "index": index + }] + }, + "finish_reason": Value::Null + }] + })) + }; + let mut frames = state + .push_line(&report_context, reasoning_chunk(0, "first")) + .expect("first reasoning block should parse"); + frames.extend( + state + .push_line(&report_context, reasoning_chunk(1, "second")) + .expect("second reasoning block should parse"), + ); + + let reasoning = frames + .iter() + .filter_map(|frame| match frame.event { + CanonicalStreamEvent::ReasoningDelta(ref text) => Some(text.as_str()), + CanonicalStreamEvent::ReasoningSummaryDone => Some(""), + _ => None, + }) + .collect::>(); + assert_eq!(reasoning, vec!["first", "", "second"]); + } + + #[test] + fn openai_chat_reasoning_only_stream_reaches_responses_clients() { + // Regression: OpenRouter streams a long reasoning phase as chunks whose + // `delta.content` is an empty string and whose text sits under + // `reasoning`. Dropping those chunks left Responses clients with + // nothing after `response.in_progress` until they timed the stream out. + let mut state = OpenAIChatProviderState::default(); + let mut emitter = OpenAIResponsesClientEmitter::default(); + let report_context = json!({}); + let mut bytes = Vec::new(); + for piece in ["Let", " me think."] { + let frames = state + .push_line( + &report_context, + data_line(json!({ + "id": "gen-openrouter-123", + "model": "stealth/ox-alpha", + "choices": [{ + "index": 0, + "delta": { + "content": "", + "role": "assistant", + "reasoning": piece, + "reasoning_details": [{ + "type": "reasoning.text", + "text": piece, + "format": "unknown", + "index": 0 + }] + }, + "finish_reason": Value::Null + }] + })), + ) + .expect("reasoning chunk should parse"); + for frame in frames { + bytes.extend(emitter.emit(frame).expect("frame should encode")); + } + } + + let sse = String::from_utf8(bytes).expect("sse should be utf8"); + assert!(sse.contains("event: response.reasoning_text.delta\n")); + assert!(!sse.contains("event: response.reasoning_summary_text.delta\n")); + assert!(sse.contains("\"delta\":\"Let\"")); + assert!(sse.contains("\"delta\":\" me think.\"")); + } + #[test] fn openai_chat_provider_state_waits_for_real_tool_call_identity() { let mut state = OpenAIChatProviderState::default(); @@ -5197,6 +5468,100 @@ mod tests { assert_eq!(text, "First message.Second message."); } + #[test] + fn openai_responses_provider_state_replays_gemini_signature_carrier_to_client() { + let mut state = OpenAIResponsesProviderState::default(); + let report_context = json!({}); + let signature = "skip_thought_signature_validator"; + let carrier = encode_gemini_tool_signature_carrier(signature) + .expect("signature carrier should encode"); + let reasoning_item = json!({ + "type": "reasoning", + "id": "rs_signature_0", + "status": "completed", + "encrypted_content": carrier, + "summary": [] + }); + let events = [ + json!({ + "type": "response.output_item.added", + "response_id": "resp_signed_fallback", + "output_index": 0, + "item": reasoning_item + }), + json!({ + "type": "response.output_item.done", + "response_id": "resp_signed_fallback", + "output_index": 0, + "item": reasoning_item + }), + json!({ + "type": "response.output_item.added", + "response_id": "resp_signed_fallback", + "output_index": 1, + "item": { + "type": "function_call", + "id": "fc_signed_1", + "call_id": "call_signed_1", + "name": "fabric_exec", + "arguments": "{\"code\":\"return 1\"}", + "status": "completed" + } + }), + ]; + let mut frames = Vec::new(); + for event in events { + frames.extend( + state + .push_line(&report_context, data_line(event)) + .expect("Responses event should parse"), + ); + } + + let signature_index = frames + .iter() + .position(|frame| { + matches!( + frame.event, + CanonicalStreamEvent::ToolCallSignature { + index: 1, + ref signature + } if signature == "skip_thought_signature_validator" + ) + }) + .expect("signature event should be restored"); + let call_index = frames + .iter() + .position(|frame| { + matches!( + frame.event, + CanonicalStreamEvent::ToolCallStart { index: 1, .. } + ) + }) + .expect("tool call should be emitted"); + assert!(signature_index < call_index); + assert_eq!( + frames + .iter() + .filter(|frame| matches!( + frame.event, + CanonicalStreamEvent::ToolCallSignature { .. } + )) + .count(), + 1, + "added and done snapshots must not duplicate the signature" + ); + + let mut emitter = GeminiClientEmitter::default(); + let mut bytes = Vec::new(); + for frame in frames { + bytes.extend(emitter.emit(frame).expect("Gemini frame should encode")); + } + let sse = String::from_utf8(bytes).expect("Gemini SSE should be UTF-8"); + assert!(sse.contains("\"name\":\"fabric_exec\"")); + assert!(sse.contains("\"thoughtSignature\":\"skip_thought_signature_validator\"")); + } + #[test] fn openai_responses_provider_state_delays_arguments_until_tool_name_is_known() { let mut state = OpenAIResponsesProviderState::default(); @@ -6351,14 +6716,19 @@ mod tests { ); let sse = String::from_utf8(bytes).expect("sse should be utf8"); - assert!(sse.contains("event: response.reasoning_summary_part.added\n")); - assert!(sse.contains("event: response.reasoning_summary_text.delta\n")); - assert!(sse.contains("event: response.reasoning_summary_text.done\n")); - assert!(sse.contains("event: response.reasoning_summary_part.done\n")); + assert!(sse.contains("event: response.reasoning_text.delta\n")); + assert!(sse.contains("event: response.reasoning_text.done\n")); + assert!(sse.contains("\"type\":\"reasoning_text\"")); + // Raw chain-of-thought must not be duplicated onto the summary channel: + // Codex renders both, so emitting both makes the thinking panel repeat. + assert!(!sse.contains("event: response.reasoning_summary_text.delta\n")); + assert!(!sse.contains("event: response.reasoning_summary_text.done\n")); + assert!(!sse.contains("event: response.reasoning_summary_part.added\n")); + assert!(!sse.contains("event: response.reasoning_summary_part.done\n")); let reasoning_item_id = openai_responses_synthetic_reasoning_item_id("resp_456", 0); assert!(sse.contains(&format!("\"item_id\":\"{reasoning_item_id}\""))); assert!(sse.contains("\"type\":\"reasoning\"")); - assert_eq!(response_sequence_numbers(&sse), (1..=9).collect::>()); + assert_eq!(response_sequence_numbers(&sse), (1..=7).collect::>()); } #[test] @@ -6427,6 +6797,75 @@ mod tests { ); } + /// Regression: raw thinking must reach the client exactly once. + /// + /// Codex renders both the `content` (`reasoning_text`) and `summary` + /// (`summary_text`) channels, so emitting the same chain-of-thought on both + /// made its thinking panel print every line twice. + #[test] + fn openai_responses_client_emitter_sends_raw_thinking_once() { + let mut emitter = OpenAIResponsesClientEmitter::default(); + let mut bytes = emitter + .emit(CanonicalStreamFrame { + id: "resp_once".to_string(), + model: "gpt-5.4".to_string(), + event: CanonicalStreamEvent::Start, + }) + .expect("start should encode"); + for text in ["Let", " me", " think."] { + bytes.extend( + emitter + .emit(CanonicalStreamFrame { + id: "resp_once".to_string(), + model: "gpt-5.4".to_string(), + event: CanonicalStreamEvent::ReasoningDelta(text.to_string()), + }) + .expect("reasoning delta should encode"), + ); + } + bytes.extend( + emitter + .emit(CanonicalStreamFrame { + id: "resp_once".to_string(), + model: "gpt-5.4".to_string(), + event: CanonicalStreamEvent::ReasoningSummaryDone, + }) + .expect("reasoning boundary should encode"), + ); + bytes.extend( + emitter + .emit(CanonicalStreamFrame { + id: "resp_once".to_string(), + model: "gpt-5.4".to_string(), + event: CanonicalStreamEvent::Finish { + finish_reason: Some("stop".to_string()), + usage: None, + }, + }) + .expect("finish should encode"), + ); + + let sse = String::from_utf8(bytes).expect("sse should be utf8"); + // Each thinking chunk is streamed on exactly one channel. The same delta + // used to be mirrored onto `reasoning_summary_text.delta`, so clients that + // render both channels (Codex) printed every chunk twice. + assert_eq!( + sse.matches("event: response.reasoning_text.delta\n") + .count(), + 3, + "one delta event per thinking chunk: {sse}" + ); + assert!(!sse.contains("event: response.reasoning_summary_text.delta\n")); + assert!(!sse.contains("event: response.reasoning_summary_text.done\n")); + assert!(!sse.contains("\"type\":\"summary_text\"")); + // The completed item carries the thinking on `content`, not `summary`. + assert!( + sse.contains("\"content\":[{\"type\":\"reasoning_text\",\"text\":\"Let me think.\"}]"), + "{sse}" + ); + assert!(sse.contains("\"summary\":[]"), "{sse}"); + } + #[test] fn openai_responses_client_emitter_emits_failed_event_with_sequence_number() { let mut emitter = OpenAIResponsesClientEmitter::default(); diff --git a/crates/aether-ai/formats/src/formats/openai/request_contract.rs b/crates/aether-ai/formats/src/formats/openai/request_contract.rs index b145d7c69..8805cd279 100644 --- a/crates/aether-ai/formats/src/formats/openai/request_contract.rs +++ b/crates/aether-ai/formats/src/formats/openai/request_contract.rs @@ -147,6 +147,9 @@ fn finalize_openai_provider_request_with_codex_model_capabilities_and_reasoning_ finalization.provider_api_format, reasoning_replay_policy, ); + if crate::is_openai_responses_family_format(finalization.provider_api_format) { + super::responses::normalize_openai_responses_call_ids(body); + } if finalization .provider_api_format .trim() @@ -226,7 +229,7 @@ fn validate_final_openai_provider_request_contract( #[cfg(test)] mod tests { - use serde_json::json; + use serde_json::{json, Value}; use super::{ finalize_openai_provider_request, @@ -235,6 +238,79 @@ mod tests { }; use crate::CodexResponsesModelCapabilities; + #[test] + fn finalization_bounds_responses_call_ids_and_preserves_pairing() { + let long_id = format!("call_{}", "a".repeat(78)); + let original = json!({ + "model": "gpt-5.4", + "input": [ + {"type": "function_call", "call_id": long_id, "name": "lookup", "arguments": "{}"}, + {"type": "function_call_output", "call_id": long_id, "output": "result"} + ] + }); + + for (source_api_format, provider_type, provider_api_format, websocket_continuation) in [ + ("openai:responses", "codex", "openai:responses", false), + ( + "openai:responses", + "codex", + "openai:responses:compact", + false, + ), + ("openai:responses", "openai", "openai:responses", false), + ( + "openai:responses", + "openai", + "openai:responses:compact", + false, + ), + ("openai:responses", "codex", "openai:responses", true), + ("openai:chat", "codex", "openai:responses", false), + ("openai:chat", "openai", "openai:responses", false), + ("claude:messages", "codex", "openai:responses", false), + ("claude:messages", "openai", "openai:responses", false), + ( + "gemini:generate_content", + "codex", + "openai:responses", + false, + ), + ( + "gemini:generate_content", + "openai", + "openai:responses", + false, + ), + ] { + let mut body = original.clone(); + let finalization = OpenAiProviderRequestFinalization { + source_api_format, + provider_api_format, + provider_type, + provider_model: "gpt-5.4", + source_model: "gpt-5.4", + body_rules: None, + upstream_is_stream: false, + require_body_stream_field: true, + }; + if websocket_continuation { + super::finalize_openai_provider_request_with_codex_model_capabilities_and_reasoning_replay_policy_for_websocket_continuation( + &mut body, + finalization, + None, + crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds, + ) + } else { + finalize_openai_provider_request(&mut body, finalization) + } + .expect("request should finalize"); + + let call_id = body["input"][0]["call_id"].as_str().expect("call ID"); + assert!(call_id.len() <= 64, "call ID has {} bytes", call_id.len()); + assert_eq!(body["input"][1]["call_id"], call_id); + } + } + #[test] fn validates_reasoning_and_prompt_cache_against_the_final_provider_model() { let body = json!({ @@ -511,6 +587,61 @@ mod tests { } } + #[test] + fn codex_finalization_preserves_explicit_service_tiers() { + let capabilities = CodexResponsesModelCapabilities { + use_responses_lite: false, + supports_reasoning_summary_parameter: false, + default_reasoning_effort: None, + default_reasoning_summary: None, + supported_reasoning_efforts: Vec::new(), + supports_parallel_tool_calls: true, + support_verbosity: false, + default_verbosity: None, + supported_service_tiers: vec!["priority".to_string()], + }; + for source_api_format in ["openai:responses", "openai:chat"] { + for provider_api_format in ["openai:responses", "openai:responses:compact"] { + for model_capabilities in [None, Some(&capabilities)] { + for service_tier in [ + Some("ultrafast"), + Some("priority"), + Some("default"), + Some("auto"), + Some("flex"), + Some("future-tier"), + None, + ] { + let mut body = json!({"model": "gpt-5.6-sol", "input": []}); + if let Some(service_tier) = service_tier { + body["service_tier"] = json!(service_tier); + } + finalize_openai_provider_request_with_codex_model_capabilities( + &mut body, + OpenAiProviderRequestFinalization { + source_api_format, + provider_api_format, + provider_type: "codex", + provider_model: "gpt-5.6-sol", + source_model: "gpt-5.6-sol", + body_rules: None, + upstream_is_stream: true, + require_body_stream_field: true, + }, + model_capabilities, + ) + .expect("explicit service tiers should be validated by the upstream"); + assert_eq!( + body.get("service_tier").and_then(Value::as_str), + service_tier, + "{source_api_format} -> {provider_api_format}", + ); + } + } + } + } + } + #[test] fn dynamic_codex_card_preserves_default_effort_and_keeps_mode_model_specific() { let finalization = OpenAiProviderRequestFinalization { diff --git a/crates/aether-ai/formats/src/formats/openai/responses/codex.rs b/crates/aether-ai/formats/src/formats/openai/responses/codex.rs index 919b7f659..548a65894 100644 --- a/crates/aether-ai/formats/src/formats/openai/responses/codex.rs +++ b/crates/aether-ai/formats/src/formats/openai/responses/codex.rs @@ -1,6 +1,7 @@ use std::collections::BTreeMap; use std::sync::OnceLock; +use crate::codex_profile::codex_client_profile; use aether_ai_formats::provider_compat::proxy::rules::body_rules_handle_path; use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; @@ -36,9 +37,6 @@ const CODEX_OPENAI_RESPONSES_COMPACT_BODY_FIELDS: &[&str] = &[ "prompt_cache_key", "text", ]; -pub const CODEX_CLIENT_VERSION: &str = "0.153.4"; -pub const CODEX_CLIENT_USER_AGENT: &str = "codex_cli_rs/0.153.4"; -pub const CODEX_CLIENT_ORIGINATOR: &str = "codex_cli_rs"; pub const CODEX_OPENAI_IMAGE_INTERNAL_MODEL: &str = "gpt-5.4-mini"; pub const CODEX_OPENAI_IMAGE_DEFAULT_MODEL: &str = "gpt-image-2"; pub const CODEX_OPENAI_IMAGE_DEFAULT_VARIATION_MODEL: &str = "dall-e-2"; @@ -91,12 +89,6 @@ impl CodexResponsesModelCapabilities { .iter() .any(|candidate| candidate == effort.trim()) } - - fn supports_service_tier(&self, service_tier: &str) -> bool { - self.supported_service_tiers - .iter() - .any(|candidate| candidate == service_tier) - } } fn codex_namespaced_model_suffix(model: &str) -> Option<&str> { @@ -1214,18 +1206,6 @@ fn apply_codex_model_request_capabilities( } } } - - if !body_rules_handle_path(body_rules, "service_tier") { - let service_tier = body_object - .get("service_tier") - .and_then(Value::as_str) - .map(str::to_string); - if !service_tier.as_deref().is_some_and(|service_tier| { - service_tier != "default" && capabilities.supports_service_tier(service_tier) - }) { - body_object.remove("service_tier"); - } - } } fn ensure_codex_reasoning_defaults( @@ -2116,6 +2096,7 @@ pub fn apply_codex_openai_special_headers( }; let auth_identity = parse_codex_auth_identity(decrypted_auth_config_raw); + let client_profile = codex_client_profile(); remove_btree_header(provider_request_headers, "chatgpt-account-id"); remove_btree_header(provider_request_headers, "x-openai-fedramp"); @@ -2130,12 +2111,12 @@ pub fn apply_codex_openai_special_headers( set_codex_client_header( provider_request_headers, "user-agent", - CODEX_CLIENT_USER_AGENT, + &client_profile.user_agent, ); set_codex_client_header( provider_request_headers, "originator", - CODEX_CLIENT_ORIGINATOR, + &client_profile.originator, ); if endpoint_kind == CodexOpenAiEndpointKind::Search { remove_btree_header(provider_request_headers, CODEX_RESPONSES_LITE_HEADER); @@ -2193,17 +2174,18 @@ mod tests { build_codex_model_catalog_metadata, bundled_codex_model_cards, effective_codex_model_cards, parse_codex_auth_identity, project_codex_catalog_model_card, resolve_codex_responses_model_capabilities, - validate_codex_openai_responses_compact_request_contract, CODEX_CLIENT_ORIGINATOR, - CODEX_CLIENT_USER_AGENT, CODEX_CLIENT_VERSION, CODEX_OPENAI_IMAGE_INTERNAL_MODEL, - CODEX_OPENAI_RESPONSES_UNSUPPORTED_BODY_FIELDS, CODEX_RESPONSES_LITE_HEADER, + validate_codex_openai_responses_compact_request_contract, + CODEX_OPENAI_IMAGE_INTERNAL_MODEL, CODEX_OPENAI_RESPONSES_UNSUPPORTED_BODY_FIELDS, + CODEX_RESPONSES_LITE_HEADER, }; use serde_json::{json, Value}; #[test] fn codex_client_user_agent_matches_originator_and_version() { + let profile = crate::codex_client_profile(); assert_eq!( - CODEX_CLIENT_USER_AGENT, - format!("{CODEX_CLIENT_ORIGINATOR}/{CODEX_CLIENT_VERSION}") + profile.user_agent, + format!("{}/{}", profile.originator, profile.codex_version) ); } @@ -2654,7 +2636,7 @@ mod tests { assert_eq!(body["reasoning"]["effort"], "high"); assert_eq!(body["include"], json!(["reasoning.encrypted_content"])); assert_eq!(body["parallel_tool_calls"], false); - assert!(body.get("service_tier").is_none()); + assert_eq!(body["service_tier"], "priority"); assert!(body["text"].get("verbosity").is_none()); assert_eq!(body["text"]["format"]["type"], "json_schema"); @@ -2981,11 +2963,11 @@ mod tests { ); assert_eq!( headers.get("user-agent").map(String::as_str), - Some(CODEX_CLIENT_USER_AGENT) + Some(crate::codex_client_user_agent().as_str()) ); assert_eq!( headers.get("originator").map(String::as_str), - Some(CODEX_CLIENT_ORIGINATOR) + Some(crate::codex_client_originator().as_str()) ); assert!(!headers.contains_key(CODEX_RESPONSES_LITE_HEADER)); assert!(!headers.contains_key("openai-beta")); @@ -3019,11 +3001,11 @@ mod tests { ); assert_eq!( headers.get("user-agent").map(String::as_str), - Some(CODEX_CLIENT_USER_AGENT) + Some(crate::codex_client_user_agent().as_str()) ); assert_eq!( headers.get("originator").map(String::as_str), - Some(CODEX_CLIENT_ORIGINATOR) + Some(crate::codex_client_originator().as_str()) ); assert!(!headers.contains_key(CODEX_RESPONSES_LITE_HEADER)); } diff --git a/crates/aether-ai/formats/src/formats/openai/responses/mod.rs b/crates/aether-ai/formats/src/formats/openai/responses/mod.rs index 36053f67e..e963df4b6 100644 --- a/crates/aether-ai/formats/src/formats/openai/responses/mod.rs +++ b/crates/aether-ai/formats/src/formats/openai/responses/mod.rs @@ -1,5 +1,9 @@ -use base64::{engine::general_purpose::STANDARD_NO_PAD, Engine as _}; -use serde_json::Value; +use base64::{ + engine::general_purpose::{STANDARD_NO_PAD, URL_SAFE_NO_PAD}, + Engine as _, +}; +use serde_json::{json, Map, Value}; +use sha2::{Digest, Sha256}; pub mod codex; pub(crate) mod history; @@ -7,6 +11,7 @@ pub mod request; pub mod response; pub mod spec; pub mod stream; +pub mod xai; const TOOL_ERROR_PREFIX: &str = "[tool error]"; const AETHER_REASONING_ITEM_ID_PREFIX: &str = "rs_aether_"; @@ -85,6 +90,8 @@ pub enum OpenAiResponsesReasoningReplayPolicy { #[default] OpenAiItemIds, DeepSeekOpaque, + /// xAI replays encrypted state without requiring OpenAI's item-ID prefix. + XaiEncrypted, } /// Builds a stable, wire-compatible ID for a reasoning item synthesized by Aether. @@ -119,6 +126,50 @@ pub fn openai_responses_message_item_id(response_id: &str, output_index: usize) ) } +/// Builds the Responses reasoning `content` array from raw thinking text. +/// +/// Raw chain-of-thought belongs in `content` as `reasoning_text` parts. It is +/// deliberately *not* mirrored into `summary`: OpenAI keeps the two channels +/// distinct, and clients such as Codex render both, so duplicating the same +/// text onto `summary` made the thinking panel print everything twice. +pub(crate) fn openai_responses_reasoning_text_parts( + texts: impl IntoIterator>, +) -> Value { + Value::Array( + texts + .into_iter() + .map(|text| text.as_ref().to_string()) + .filter(|text| !text.trim().is_empty()) + .map(|text| json!({ "type": "reasoning_text", "text": text })) + .collect(), + ) +} + +/// Writes raw thinking onto a Responses reasoning item without clobbering an +/// existing provider-owned summary or content. +pub(crate) fn apply_openai_responses_reasoning_text(item: &mut Map, text: &str) { + if text.trim().is_empty() { + return; + } + if reasoning_item_field_is_empty(item.get("content")) { + let content = openai_responses_reasoning_text_parts(std::iter::once(text)); + item.insert("content".to_string(), content); + } + // `summary` stays a valid (empty) array so the item keeps its documented + // shape; a provider-supplied summary is preserved as-is. + item.entry("summary".to_string()) + .or_insert_with(|| Value::Array(Vec::new())); +} + +fn reasoning_item_field_is_empty(value: Option<&Value>) -> bool { + match value { + None | Some(Value::Null) => true, + Some(Value::Array(parts)) => parts.is_empty(), + Some(Value::String(text)) => text.trim().is_empty(), + _ => false, + } +} + /// Repairs legacy/non-OpenAI message IDs in a Responses request in place. /// /// Aether versions before the `msg_` contract emitted IDs such as @@ -164,6 +215,28 @@ pub fn normalize_openai_responses_message_item_ids(body: &mut Value) -> usize { repaired } +pub(crate) fn normalize_openai_responses_call_ids(body: &mut Value) { + let Some(input) = body.get_mut("input") else { + return; + }; + let items = match input { + Value::Array(items) => items.as_mut_slice(), + Value::Object(_) => std::slice::from_mut(input), + _ => return, + }; + for item in items { + let Some(Value::String(call_id)) = item.get_mut("call_id") else { + continue; + }; + if call_id.chars().take(65).count() > 64 { + *call_id = format!( + "call_{}", + URL_SAFE_NO_PAD.encode(Sha256::digest(call_id.as_bytes())) + ); + } + } +} + /// Removes reasoning history items that cannot be replayed against an OpenAI Responses backend. /// /// Reasoning IDs are opaque provider references and must never be repaired by changing their @@ -234,6 +307,14 @@ fn openai_responses_reasoning_item_is_replayable( { return true; } + if policy == OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + && object + .get("encrypted_content") + .and_then(Value::as_str) + .is_some_and(|value| !value.trim().is_empty()) + { + return true; + } let Some(id) = object .get("id") .and_then(Value::as_str) @@ -325,8 +406,9 @@ mod tests { use super::{ decode_gemini_tool_signature_carrier, encode_gemini_tool_signature_carrier_with_direction, - normalize_openai_responses_message_item_ids, openai_responses_message_item_id, - openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id, + normalize_openai_responses_call_ids, normalize_openai_responses_message_item_ids, + openai_responses_message_item_id, openai_responses_request_operation, + openai_responses_synthetic_reasoning_item_id, strip_incompatible_openai_responses_reasoning_items, strip_incompatible_openai_responses_reasoning_items_with_policy, GeminiToolSignatureCarrierDirection, OpenAiResponsesReasoningReplayPolicy, @@ -334,6 +416,36 @@ mod tests { OPENAI_RESPONSES_OPERATION_COMPACT, }; + #[test] + fn xai_encrypted_replay_accepts_native_ids_but_excludes_foreign_carriers() { + let body = serde_json::json!({"input": [ + {"type": "reasoning", "id": "native-xai-id", "encrypted_content": "opaque-xai-state"}, + {"type": "reasoning", "encrypted_content": "opaque-idless-state"}, + {"type": "reasoning", "id": "rs_foreign", "encrypted_content": "cpa-gemini-responses-carrier-v1:foreign"}, + {"type": "reasoning", "id": "foreign-id", "summary": []} + ]}); + let mut xai = body.clone(); + assert_eq!( + super::strip_incompatible_openai_responses_reasoning_items_with_policy( + &mut xai, + "openai:responses", + super::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted, + ), + 2 + ); + assert_eq!(xai["input"].as_array().unwrap().len(), 2); + assert_eq!(xai["input"][0], body["input"][0]); + assert_eq!(xai["input"][1], body["input"][1]); + let mut openai = body; + assert_eq!( + super::strip_incompatible_openai_responses_reasoning_items( + &mut openai, + "openai:responses" + ), + 4 + ); + } + #[test] fn gemini_tool_signature_carrier_roundtrips_direction_and_exact_value() { let signature = " opaque-signature-with-padding== "; @@ -419,6 +531,33 @@ mod tests { assert_ne!(first, other); } + #[test] + fn reasoning_text_parts_put_raw_thinking_in_content_only() { + let content = super::openai_responses_reasoning_text_parts(["raw chain"]); + assert_eq!( + content, + json!([{ "type": "reasoning_text", "text": "raw chain" }]) + ); + + let mut item = serde_json::Map::new(); + super::apply_openai_responses_reasoning_text(&mut item, "raw chain"); + assert_eq!(item["content"], content); + // Never mirrored onto `summary`: clients rendering both would repeat it. + assert_eq!(item["summary"], json!([])); + + item.insert( + "summary".to_string(), + json!([{ "type": "summary_text", "text": "kept" }]), + ); + item.insert("content".to_string(), json!([])); + super::apply_openai_responses_reasoning_text(&mut item, "replacement"); + assert_eq!( + item["content"], + json!([{ "type": "reasoning_text", "text": "replacement" }]) + ); + assert_eq!(item["summary"][0]["text"], "kept"); + } + #[test] fn synthetic_message_item_ids_are_stable_and_start_with_msg() { let first = openai_responses_message_item_id("1c938e58-32a8-4d28-9c34-538d78076895", 0); @@ -430,6 +569,77 @@ mod tests { assert_ne!(first, other); } + #[test] + fn normalizes_long_call_ids_stably_without_changing_item_ids_or_payloads() { + let long_id = format!("call_{}", "a".repeat(78)); + let other_id = format!("{long_id}b"); + let arguments = json!({"call_id": long_id}).to_string(); + let mut body = json!({"input": [ + {"type": "function_call", "id": "fc_provider", "call_id": long_id, "name": "lookup", "arguments": arguments}, + {"type": "function_call_output", "call_id": long_id, "output": {"call_id": long_id}}, + {"type": "custom_tool_call", "call_id": other_id, "name": "patch", "input": long_id}, + {"type": "custom_tool_call_output", "call_id": other_id, "output": "done"} + ]}); + + normalize_openai_responses_call_ids(&mut body); + + let first_id = body["input"][0]["call_id"].as_str().expect("first call ID"); + let second_id = body["input"][2]["call_id"] + .as_str() + .expect("second call ID"); + for call_id in [first_id, second_id] { + assert!(call_id.len() <= 64); + assert!(call_id.chars().all( + |character| character.is_ascii_alphanumeric() || matches!(character, '_' | '-') + )); + } + assert_ne!(first_id, second_id); + assert_eq!(body["input"][1]["call_id"], first_id); + assert_eq!(body["input"][3]["call_id"], second_id); + assert_eq!(body["input"][0]["id"], "fc_provider"); + assert_eq!(body["input"][0]["arguments"], arguments); + assert_eq!(body["input"][1]["output"]["call_id"], long_id); + assert_eq!(body["input"][2]["input"], long_id); + + let mut continuation = json!({"input": { + "type": "function_call_output", "call_id": long_id, "output": "later" + }}); + normalize_openai_responses_call_ids(&mut continuation); + assert_eq!(continuation["input"]["call_id"], first_id); + + let once = body.clone(); + normalize_openai_responses_call_ids(&mut body); + assert_eq!(body, once); + } + + #[test] + fn call_id_normalization_preserves_valid_boundaries_and_non_item_data() { + let mut body = json!({"input": [ + {"type": "function_call", "call_id": "call_short"}, + {"type": "function_call", "call_id": "a".repeat(64)}, + {"type": "function_call", "call_id": "\u{00e9}".repeat(64)}, + {"type": "message", "content": [{"call_id": "a".repeat(83)}]}, + {"type": "function_call_output", "call_id": null}, + {"type": "function_call_output", "call_id": 42}, + null + ]}); + let unchanged = body.clone(); + normalize_openai_responses_call_ids(&mut body); + assert_eq!(body, unchanged); + + for input in [json!("text"), json!(null)] { + let mut body = json!({"input": input}); + let unchanged = body.clone(); + normalize_openai_responses_call_ids(&mut body); + assert_eq!(body, unchanged); + } + for call_id in ["a".repeat(65), "\u{00e9}".repeat(65)] { + let mut body = json!({"input": [{"type": "function_call", "call_id": call_id}]}); + normalize_openai_responses_call_ids(&mut body); + assert!(body["input"][0]["call_id"].as_str().expect("call ID").len() <= 64); + } + } + #[test] fn normalizes_legacy_message_ids_but_preserves_valid_ids() { let mut body = json!({ diff --git a/crates/aether-ai/formats/src/formats/openai/responses/request.rs b/crates/aether-ai/formats/src/formats/openai/responses/request.rs index 519a2e505..cd5799d32 100644 --- a/crates/aether-ai/formats/src/formats/openai/responses/request.rs +++ b/crates/aether-ai/formats/src/formats/openai/responses/request.rs @@ -2,7 +2,7 @@ use std::collections::{BTreeMap, VecDeque}; use serde_json::{json, Map, Value}; -use super::encode_tool_result_error; +use super::{apply_openai_responses_reasoning_text, encode_tool_result_error}; use crate::{ formats::context::FormatContext, @@ -702,14 +702,7 @@ fn canonical_thinking_to_responses_reasoning_item( .unwrap_or_default(); item.remove("item_type"); item.insert("type".to_string(), Value::String("reasoning".to_string())); - if !text.trim().is_empty() { - item.entry("summary".to_string()).or_insert_with(|| { - json!([{ - "type": "summary_text", - "text": text, - }]) - }); - } + apply_openai_responses_reasoning_text(&mut item, text); if let Some(value) = encrypted_content.filter(|value| !value.is_empty()) { item.insert( "encrypted_content".to_string(), diff --git a/crates/aether-ai/formats/src/formats/openai/responses/response.rs b/crates/aether-ai/formats/src/formats/openai/responses/response.rs index 977d8b8f6..86d44dc4b 100644 --- a/crates/aether-ai/formats/src/formats/openai/responses/response.rs +++ b/crates/aether-ai/formats/src/formats/openai/responses/response.rs @@ -6,8 +6,9 @@ use std::{ use serde_json::{json, Map, Value}; use super::{ - encode_gemini_tool_signature_carrier, encode_tool_result_error, - history::record_converted_response_history, openai_responses_synthetic_reasoning_item_id, + apply_openai_responses_reasoning_text, encode_gemini_tool_signature_carrier, + encode_tool_result_error, history::record_converted_response_history, + openai_responses_synthetic_reasoning_item_id, }; use crate::{ @@ -113,6 +114,14 @@ fn openai_responses_incomplete_stop_reason(body: &Map) -> Canonic } } +fn canonical_incomplete_reason(canonical: &CanonicalResponse) -> Option<&'static str> { + match canonical.stop_reason.as_ref()? { + CanonicalStopReason::MaxTokens => Some("max_output_tokens"), + CanonicalStopReason::ContentFiltered => Some("content_filter"), + _ => None, + } +} + pub fn to_raw(canonical: &CanonicalResponse, report_context: &Value, compact: bool) -> Value { let namespace_tool_aliases = NamespaceToolAliases::from_report_context(report_context); let mut response = Map::new(); @@ -142,6 +151,17 @@ pub fn to_raw(canonical: &CanonicalResponse, report_context: &Value, compact: bo .cloned() { response.insert("status".to_string(), raw_status); + } else if let Some(reason) = canonical_incomplete_reason(canonical) { + // Cross-format sources carry no Responses status of their own; a + // truncated or filtered answer must not be reported as completed. + response.insert( + "status".to_string(), + Value::String("incomplete".to_string()), + ); + response.insert( + "incomplete_details".to_string(), + json!({ "reason": reason }), + ); } } @@ -217,15 +237,7 @@ pub fn to_raw(canonical: &CanonicalResponse, report_context: &Value, compact: bo Value::String(encrypted_content.clone()), ); } - if !text.trim().is_empty() { - item.insert( - "summary".to_string(), - Value::Array(vec![json!({ - "type": "summary_text", - "text": text, - })]), - ); - } + apply_openai_responses_reasoning_text(&mut item, text); output.push(Value::Object(item)); } CanonicalContentBlock::ToolUse { @@ -792,6 +804,65 @@ mod tests { ); } + #[test] + fn responses_response_builder_puts_raw_thinking_in_content_only() { + let response = CanonicalResponse { + id: "resp_think".to_string(), + model: "deepseek-reasoner".to_string(), + content: vec![ + CanonicalContentBlock::Thinking { + text: "first add one to one".to_string(), + signature: None, + encrypted_content: None, + extensions: BTreeMap::new(), + }, + CanonicalContentBlock::Text { + text: "2".to_string(), + extensions: BTreeMap::new(), + }, + ], + outputs: Vec::new(), + stop_reason: Some(CanonicalStopReason::EndTurn), + usage: None, + extensions: BTreeMap::new(), + }; + + let body = to_raw(&response, &json!({}), false); + let item = &body["output"][0]; + + assert_eq!(item["type"], "reasoning"); + assert_eq!(item["content"][0]["type"], "reasoning_text"); + assert_eq!(item["content"][0]["text"], "first add one to one"); + assert_eq!(item["summary"], json!([])); + assert!(!item["content"].is_null()); + assert_eq!(body["output"][1]["type"], "message"); + assert_eq!(body["output"][1]["content"][0]["text"], "2"); + } + + #[test] + fn responses_response_parser_prefers_content_over_summary_for_raw_reasoning() { + let body = json!({ + "id": "resp_test", + "model": "gpt-5", + "status": "completed", + "output": [{ + "type": "reasoning", + "id": "rs_1", + "status": "completed", + "summary": [{"type": "summary_text", "text": "short summary"}], + "content": [{"type": "reasoning_text", "text": "full chain of thought"}] + }] + }); + + let canonical = from_raw(&body).expect("response should parse"); + + assert!(matches!( + canonical.content.first(), + Some(CanonicalContentBlock::Thinking { text, .. }) + if text == "full chain of thought" + )); + } + #[test] fn responses_response_parser_preserves_encrypted_reasoning_without_summary() { let body = json!({ @@ -879,4 +950,73 @@ mod tests { }) if id == "call_ws_1" && name == "web_search" && input["query"] == "today tech") ); } + + #[test] + fn responses_response_builder_reports_cross_format_truncation_as_incomplete() { + let response = |stop_reason| CanonicalResponse { + id: "gemini-resp".to_string(), + model: "gemini-3.8-flash".to_string(), + content: vec![CanonicalContentBlock::Text { + text: "partial".to_string(), + extensions: BTreeMap::new(), + }], + outputs: Vec::new(), + stop_reason: Some(stop_reason), + usage: None, + extensions: BTreeMap::new(), + }; + + let truncated = to_raw(&response(CanonicalStopReason::MaxTokens), &json!({}), false); + assert_eq!(truncated["status"], "incomplete"); + assert_eq!( + truncated["incomplete_details"], + json!({"reason": "max_output_tokens"}) + ); + + let filtered = to_raw( + &response(CanonicalStopReason::ContentFiltered), + &json!({}), + false, + ); + assert_eq!(filtered["status"], "incomplete"); + assert_eq!( + filtered["incomplete_details"], + json!({"reason": "content_filter"}) + ); + + let finished = to_raw(&response(CanonicalStopReason::EndTurn), &json!({}), false); + assert_eq!(finished["status"], "completed"); + assert!(finished.get("incomplete_details").is_none()); + } + + #[test] + fn gemini_max_tokens_response_converts_to_incomplete_responses_body() { + let gemini = json!({ + "responseId": "gemini-trunc-123", + "modelVersion": "gemini-3.8-flash", + "candidates": [{ + "content": {"role": "model", "parts": [{"text": "The printing press"}]}, + "finishReason": "MAX_TOKENS" + }], + "usageMetadata": { + "promptTokenCount": 19, + "candidatesTokenCount": 256, + "totalTokenCount": 275 + } + }); + + let body = crate::formats::registry::convert_response( + "gemini:generate_content", + "openai:responses", + &gemini, + &FormatContext::default(), + ) + .expect("gemini response should convert"); + + assert_eq!(body["status"], "incomplete"); + assert_eq!( + body["incomplete_details"], + json!({"reason": "max_output_tokens"}) + ); + } } diff --git a/crates/aether-ai/formats/src/formats/openai/responses/xai.rs b/crates/aether-ai/formats/src/formats/openai/responses/xai.rs new file mode 100644 index 000000000..212f5212d --- /dev/null +++ b/crates/aether-ai/formats/src/formats/openai/responses/xai.rs @@ -0,0 +1,914 @@ +use serde_json::{json, Map, Value}; + +const XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS: &[&str] = &[ + "previous_response_id", + "prompt_cache_retention", + "safety_identifier", + "stream_options", + "stop", + "metadata", +]; +const XAI_WEB_SEARCH_TOOL_TYPE: &str = "web_search"; +const XAI_IMAGE_GENERATION_TOOL_TYPE: &str = "image_generation"; +const XAI_TOOL_SEARCH_TOOL_TYPE: &str = "tool_search"; +const XAI_GROK_IMAGE_GENERATION_MIN: XaiGrokVersion = XaiGrokVersion { major: 4, minor: 6 }; + +#[derive(Clone, Copy)] +struct XaiGrokVersion { + major: i32, + minor: i32, +} + +pub fn apply_xai_upstream_payload_edits( + body: &mut Value, + provider_type: &str, + provider_api_format: &str, +) { + apply_xai_upstream_payload_edits_with_client( + body, + provider_type, + provider_api_format, + None, + None, + ); +} + +pub fn apply_xai_upstream_payload_edits_with_client( + body: &mut Value, + provider_type: &str, + provider_api_format: &str, + client_api_format: Option<&str>, + client_body: Option<&Value>, +) { + if !provider_type.trim().eq_ignore_ascii_case("xai") { + return; + } + normalize_xai_image_refs(body); + if crate::is_openai_responses_family_format(provider_api_format) { + restore_xai_web_search_from_client(body, client_api_format, client_body); + sanitize_xai_responses_body(body); + } +} + +fn sanitize_xai_responses_body(body: &mut Value) { + let Some(object) = body.as_object_mut() else { + return; + }; + for field in XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS { + object.remove(*field); + } + let keep_image_generation = object + .get("model") + .and_then(Value::as_str) + .is_some_and(xai_supports_native_image_generation); + normalize_xai_tool_arrays(object, keep_image_generation); + rewrite_xai_web_search_tool_choice(object); + prune_xai_orphaned_tool_choice(object); + rewrite_xai_image_generation_tool_choice(object); + drop_tool_choice_without_tools(object); + strip_unsupported_reasoning_effort(object); + sanitize_xai_input_encrypted_content(object); +} + +fn restore_xai_web_search_from_client( + body: &mut Value, + client_api_format: Option<&str>, + client_body: Option<&Value>, +) { + let Some(client_api_format) = client_api_format else { + return; + }; + let Some(client_body) = client_body else { + return; + }; + if !client_requests_web_search(client_api_format, client_body) { + return; + } + ensure_xai_web_search_tool(body); + // Claude names a hosted tool in tool_choice just like a client function. + // Resolve that name against the original declaration, never by name alone. + if crate::normalize_api_format_alias(client_api_format) == "claude:messages" { + let choice = &client_body["tool_choice"]; + if choice["type"] == "tool" + && choice["name"].as_str().is_some_and(|name| { + request_tools(client_body) + .iter() + .any(|tool| is_web_search_tool(tool) && tool_name(tool) == Some(name)) + }) + { + body["tool_choice"] = json!({"type": XAI_WEB_SEARCH_TOOL_TYPE}); + } + } +} + +fn client_requests_web_search(client_api_format: &str, client_body: &Value) -> bool { + let format = crate::normalize_api_format_alias(client_api_format); + match format.as_str() { + "openai:chat" => { + object_has_non_null_field(client_body, "web_search_options") + || request_tools(client_body).iter().any(is_web_search_tool) + } + "claude:messages" => request_tools(client_body).iter().any(is_web_search_tool), + "gemini:generate_content" => gemini_request_has_google_search(client_body), + _ => false, + } +} + +fn gemini_request_has_google_search(body: &Value) -> bool { + request_tools(body).iter().any(|tool| { + tool.get("googleSearch").is_some() + || tool.get("google_search").is_some() + || tool + .get("googleSearchRetrieval") + .is_some_and(|value| !value.is_null()) + }) +} + +fn object_has_non_null_field(body: &Value, field: &str) -> bool { + body.get(field).is_some_and(|value| !value.is_null()) +} + +fn ensure_xai_web_search_tool(body: &mut Value) { + let Some(object) = body.as_object_mut() else { + return; + }; + if tools_array(object).iter().any(is_web_search_tool) { + return; + } + let tools = object + .entry("tools".to_string()) + .or_insert_with(|| Value::Array(Vec::new())); + if let Some(tools) = tools.as_array_mut() { + tools.push(json!({ "type": XAI_WEB_SEARCH_TOOL_TYPE })); + } +} + +fn normalize_xai_tool_arrays(object: &mut Map, keep_image_generation: bool) { + if let Some(tools) = object.get_mut("tools").and_then(Value::as_array_mut) { + *tools = normalize_xai_tool_list(tools, keep_image_generation); + if tools.is_empty() { + object.remove("tools"); + } + } + let Some(input) = object.get_mut("input").and_then(Value::as_array_mut) else { + return; + }; + for item in input { + let Some(item_object) = item.as_object_mut() else { + continue; + }; + if item_object.get("type").and_then(Value::as_str) != Some("additional_tools") { + continue; + } + if let Some(tools) = item_object.get_mut("tools").and_then(Value::as_array_mut) { + *tools = normalize_xai_tool_list(tools, keep_image_generation); + } + } +} + +fn normalize_xai_tool_list(tools: &[Value], keep_image_generation: bool) -> Vec { + tools + .iter() + .filter_map(|tool| normalize_xai_tool(tool, keep_image_generation)) + .collect() +} + +fn normalize_xai_tool(tool: &Value, keep_image_generation: bool) -> Option { + let Some(object) = tool.as_object() else { + return Some(tool.clone()); + }; + let tool_type = tool_type(tool).unwrap_or("function"); + if tool_type == XAI_TOOL_SEARCH_TOOL_TYPE { + return None; + } + if tool_type == XAI_IMAGE_GENERATION_TOOL_TYPE && !keep_image_generation { + return None; + } + if tool_type == "custom" && tool_name(tool).is_some_and(|name| name == "apply_patch") { + return None; + } + + let mut next = object.clone(); + if tool_type.starts_with("web_search") { + next.insert( + "type".to_string(), + Value::String(XAI_WEB_SEARCH_TOOL_TYPE.to_string()), + ); + next.remove("name"); + next.remove("external_web_access"); + return Some(Value::Object(next)); + } + if tool_type == "custom" { + next.insert("type".to_string(), Value::String("function".to_string())); + if let Some(custom) = next.remove("custom") { + if let Some(custom_object) = custom.as_object() { + for (key, value) in custom_object { + next.entry(key.clone()).or_insert_with(|| value.clone()); + } + } + } + if !next.contains_key("parameters") { + next.insert( + "parameters".to_string(), + json!({"type": "object", "properties": {}}), + ); + } + return Some(Value::Object(next)); + } + if tool_type == "function" && !next.contains_key("parameters") { + next.insert( + "parameters".to_string(), + json!({"type": "object", "properties": {}}), + ); + } + Some(Value::Object(next)) +} + +fn rewrite_xai_web_search_tool_choice(object: &mut Map) { + let Some(choice) = object.get("tool_choice").cloned() else { + return; + }; + let Some(choice_type) = choice.as_object().and_then(|value| { + value + .get("type") + .and_then(Value::as_str) + .map(str::trim) + .map(str::to_ascii_lowercase) + }) else { + return; + }; + if is_web_search_choice_type(&choice_type) { + object.insert( + "tool_choice".to_string(), + json!({ + "type": "allowed_tools", + "mode": "required", + "tools": [{ "type": XAI_WEB_SEARCH_TOOL_TYPE }] + }), + ); + } +} + +fn rewrite_xai_image_generation_tool_choice(object: &mut Map) { + let has_image_generation = tools_array(object) + .iter() + .any(|tool| tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE)); + if !has_image_generation { + return; + } + let Some(choice) = object.get("tool_choice").cloned() else { + return; + }; + // xAI's allowed_tools schema cannot contain image_generation. Preserve an + // image-only restriction before filtering image entries out of mixed lists. + let image_only = is_allowed_tools_image_generation_only(&choice); + if choice["type"] == XAI_IMAGE_GENERATION_TOOL_TYPE || image_only { + let mode = if image_only && choice["mode"] == "auto" { + "auto" + } else { + "required" + }; + keep_only_image_generation_tools(object); + object.insert("tool_choice".to_string(), Value::String(mode.to_string())); + } else if choice["type"] == "allowed_tools" { + filter_image_generation_from_allowed_tools(object); + } +} + +fn is_allowed_tools_image_generation_only(choice: &Value) -> bool { + let Some(object) = choice.as_object() else { + return false; + }; + if object.get("type").and_then(Value::as_str) != Some("allowed_tools") { + return false; + } + let Some(tools) = object.get("tools").and_then(Value::as_array) else { + return false; + }; + !tools.is_empty() + && tools.iter().all(|tool| { + tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE) + }) +} + +fn keep_only_image_generation_tools(object: &mut Map) { + let Some(tools) = object.get_mut("tools").and_then(Value::as_array_mut) else { + return; + }; + tools.retain(|tool| { + tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE) + }); +} + +fn filter_image_generation_from_allowed_tools(object: &mut Map) { + let Some(choice) = object.get_mut("tool_choice").and_then(Value::as_object_mut) else { + return; + }; + let Some(tools) = choice.get_mut("tools").and_then(Value::as_array_mut) else { + return; + }; + tools + .retain(|tool| tool_type(tool).is_none_or(|value| value != XAI_IMAGE_GENERATION_TOOL_TYPE)); +} + +fn is_web_search_choice_type(value: &str) -> bool { + value == XAI_WEB_SEARCH_TOOL_TYPE || value.starts_with("web_search") +} + +fn prune_xai_orphaned_tool_choice(object: &mut Map) { + let available = collect_available_tool_choice_keys(object); + let Some(choice) = object.get("tool_choice").cloned() else { + return; + }; + if choice.as_str().is_some() { + return; + } + let Some(choice_object) = choice.as_object() else { + object.remove("tool_choice"); + return; + }; + let choice_type = choice_object + .get("type") + .and_then(Value::as_str) + .unwrap_or_default() + .trim() + .to_ascii_lowercase(); + if choice_type == "allowed_tools" { + let Some(allowed) = choice_object.get("tools").and_then(Value::as_array) else { + object.remove("tool_choice"); + return; + }; + let kept = allowed + .iter() + .filter(|tool| tool_matches_available(tool, &available)) + .cloned() + .collect::>(); + if kept.is_empty() { + object.remove("tool_choice"); + return; + } + if let Some(choice) = object.get_mut("tool_choice").and_then(Value::as_object_mut) { + choice.insert("tools".to_string(), Value::Array(kept)); + } + return; + } + if choice_type.is_empty() { + return; + } + if !tool_matches_available(&choice, &available) { + object.remove("tool_choice"); + } +} + +fn collect_available_tool_choice_keys(object: &Map) -> Vec { + let mut keys = Vec::new(); + collect_tool_choice_keys(tools_array(object), &mut keys); + if let Some(input) = object.get("input").and_then(Value::as_array) { + for item in input { + if item.get("type").and_then(Value::as_str) == Some("additional_tools") { + collect_tool_choice_keys( + item.get("tools") + .and_then(Value::as_array) + .map(Vec::as_slice) + .unwrap_or(&[]), + &mut keys, + ); + } + } + } + keys +} + +fn collect_tool_choice_keys(tools: &[Value], keys: &mut Vec) { + for tool in tools { + let Some(tool_type) = tool_type(tool) else { + continue; + }; + if matches!(tool_type, "function" | "custom") { + if let Some(name) = tool_name(tool) { + keys.push(ToolChoiceKey::Named { + name: name.to_ascii_lowercase(), + }); + } + continue; + } + keys.push(ToolChoiceKey::Hosted(tool_type.to_ascii_lowercase())); + } +} + +fn tool_matches_available(choice: &Value, available: &[ToolChoiceKey]) -> bool { + let Some(object) = choice.as_object() else { + return false; + }; + let choice_type = object + .get("type") + .and_then(Value::as_str) + .unwrap_or_default() + .trim() + .to_ascii_lowercase(); + if matches!(choice_type.as_str(), "function" | "custom" | "tool") { + let Some(name) = tool_choice_name(object) else { + return false; + }; + return available.iter().any(|key| { + matches!( + key, + ToolChoiceKey::Named { name: available_name, .. } + if available_name == &name.to_ascii_lowercase() + ) + }); + } + if is_web_search_choice_type(&choice_type) { + return available.iter().any( + |key| matches!(key, ToolChoiceKey::Hosted(value) if value == XAI_WEB_SEARCH_TOOL_TYPE), + ); + } + available + .iter() + .any(|key| matches!(key, ToolChoiceKey::Hosted(value) if value == &choice_type)) +} + +#[derive(Clone, Debug)] +enum ToolChoiceKey { + Named { name: String }, + Hosted(String), +} + +fn drop_tool_choice_without_tools(object: &mut Map) { + if xai_request_has_tools(object) { + return; + } + object.remove("tools"); + object.remove("tool_choice"); + object.remove("parallel_tool_calls"); +} + +fn xai_request_has_tools(object: &Map) -> bool { + if !tools_array(object).is_empty() { + return true; + } + object + .get("input") + .and_then(Value::as_array) + .into_iter() + .flatten() + .any(|item| { + item.get("type") + .and_then(Value::as_str) + .is_some_and(|value| value == "additional_tools") + && item + .get("tools") + .and_then(Value::as_array) + .is_some_and(|tools| !tools.is_empty()) + }) +} + +fn strip_unsupported_reasoning_effort(object: &mut Map) { + let model = object + .get("model") + .and_then(Value::as_str) + .unwrap_or_default(); + if xai_model_supports_reasoning_effort(model) { + return; + } + let Some(reasoning) = object.get_mut("reasoning") else { + return; + }; + let Some(reasoning_object) = reasoning.as_object_mut() else { + return; + }; + reasoning_object.remove("effort"); + if reasoning_object.is_empty() { + object.remove("reasoning"); + } +} + +pub fn xai_model_supports_reasoning_effort(model: &str) -> bool { + let lowered = model.trim().to_ascii_lowercase(); + let name = lowered.rsplit('/').next().unwrap_or(lowered.as_str()); + if name.is_empty() || name.contains("non-reasoning") || name.contains("imagine") { + return false; + } + name.starts_with("grok-3-mini") + || name.starts_with("grok-4") + || name.starts_with("grok-build") + || name.starts_with("grok-composer") +} + +pub fn xai_supports_native_image_generation(model: &str) -> bool { + let lowered = model.trim().to_ascii_lowercase(); + let name = lowered.rsplit('/').next().unwrap_or(lowered.as_str()); + let Some(rest) = name.strip_prefix("grok-") else { + return false; + }; + if rest == "4.20" || rest.starts_with("4.20-") { + return false; + } + parse_grok_version_prefix(rest).is_some_and(grok_version_at_least_image_generation) +} + +fn parse_grok_version_prefix(rest: &str) -> Option { + let major_len = rest + .find(|ch: char| !ch.is_ascii_digit()) + .unwrap_or(rest.len()); + if major_len == 0 { + return None; + } + let major = rest[..major_len].parse().ok()?; + if major_len == rest.len() || !rest[major_len..].starts_with('.') { + return Some(XaiGrokVersion { major, minor: -1 }); + } + let after_dot = &rest[major_len + 1..]; + let minor_len = after_dot + .find(|ch: char| !ch.is_ascii_digit()) + .unwrap_or(after_dot.len()); + if minor_len == 0 { + return Some(XaiGrokVersion { major, minor: -1 }); + } + let minor = after_dot[..minor_len].parse().ok()?; + Some(XaiGrokVersion { major, minor }) +} + +fn grok_version_at_least_image_generation(version: XaiGrokVersion) -> bool { + let minor = if version.minor < 0 { 0 } else { version.minor }; + (version.major, minor) + >= ( + XAI_GROK_IMAGE_GENERATION_MIN.major, + XAI_GROK_IMAGE_GENERATION_MIN.minor, + ) +} + +fn sanitize_xai_input_encrypted_content(object: &mut Map) { + let Some(input) = object.get_mut("input").and_then(Value::as_array_mut) else { + return; + }; + let mut kept = Vec::new(); + for item in input.iter() { + let Some(item_object) = item.as_object() else { + kept.push(item.clone()); + continue; + }; + let item_type = item_object + .get("type") + .and_then(Value::as_str) + .unwrap_or_default(); + if item_type != "reasoning" && item_type != "compaction" { + kept.push(item.clone()); + continue; + } + let Some(encrypted) = item_object.get("encrypted_content") else { + kept.push(item.clone()); + continue; + }; + let valid = encrypted + .as_str() + .is_some_and(|value| !value.trim().is_empty()); + if valid { + kept.push(item.clone()); + continue; + } + if item_type == "compaction" { + continue; + } + let mut next = item_object.clone(); + next.remove("encrypted_content"); + kept.push(Value::Object(next)); + } + *input = kept; +} + +fn normalize_xai_image_refs(value: &mut Value) { + match value { + Value::Object(object) => { + for key in ["image", "images", "reference_images"] { + match object.get_mut(key) { + Some(Value::Array(items)) if key != "image" => { + for item in items { + normalize_xai_image_ref(item); + } + } + Some(item) if key == "image" => normalize_xai_image_ref(item), + _ => {} + } + } + for child in object.values_mut() { + normalize_xai_image_refs(child); + } + } + Value::Array(items) => { + for item in items { + normalize_xai_image_refs(item); + } + } + _ => {} + } +} + +fn normalize_xai_image_ref(value: &mut Value) { + let Some(object) = value.as_object_mut() else { + return; + }; + let original_url = object + .get("url") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let image_url = object.get("image_url").cloned(); + let resolved_url = original_url.clone().or_else(|| match image_url.as_ref() { + Some(Value::String(url)) => { + let trimmed = url.trim(); + (!trimmed.is_empty()).then(|| trimmed.to_string()) + } + Some(Value::Object(inner)) => inner + .get("url") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned), + _ => None, + }); + let Some(url) = resolved_url else { + return; + }; + if original_url.as_deref() == Some(url.as_str()) && image_url.is_none() { + return; + } + object.insert("url".to_string(), Value::String(url)); + object.remove("image_url"); +} + +fn request_tools(body: &Value) -> &[Value] { + body.get("tools") + .and_then(Value::as_array) + .map(Vec::as_slice) + .unwrap_or(&[]) +} + +fn tools_array(object: &Map) -> &[Value] { + object + .get("tools") + .and_then(Value::as_array) + .map(Vec::as_slice) + .unwrap_or(&[]) +} + +fn tool_type(tool: &Value) -> Option<&str> { + tool.get("type").and_then(Value::as_str).map(str::trim) +} + +fn tool_name(tool: &Value) -> Option<&str> { + tool.get("name") + .and_then(Value::as_str) + .or_else(|| { + tool.get("function") + .and_then(Value::as_object) + .and_then(|value| value.get("name")) + .and_then(Value::as_str) + }) + .or_else(|| { + tool.get("custom") + .and_then(Value::as_object) + .and_then(|value| value.get("name")) + .and_then(Value::as_str) + }) + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +fn tool_choice_name(choice: &Map) -> Option<&str> { + choice + .get("name") + .and_then(Value::as_str) + .or_else(|| { + choice + .get("function") + .and_then(Value::as_object) + .and_then(|value| value.get("name")) + .and_then(Value::as_str) + }) + .or_else(|| { + choice + .get("custom") + .and_then(Value::as_object) + .and_then(|value| value.get("name")) + .and_then(Value::as_str) + }) + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +fn is_web_search_tool(tool: &Value) -> bool { + tool_type(tool).is_some_and(is_web_search_choice_type) +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::{ + apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client, + xai_model_supports_reasoning_effort, xai_supports_native_image_generation, + XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS, + }; + + #[test] + fn xai_responses_edits_strip_continuation_fields_and_empty_tool_choice() { + let mut body = json!({ + "model": "grok-4.6", + "input": "hello", + "previous_response_id": "resp_123", + "prompt_cache_retention": "24h", + "safety_identifier": "user-1", + "stream_options": {"include_obfuscation": true}, + "stop": ["END"], + "metadata": { + "user_id": "{\"device_id\":\"dev-1\",\"account_uuid\":\"acct-1\",\"session_id\":\"sess-1\"}" + }, + "include": ["reasoning.encrypted_content", "file_search_call.results"], + "tool_choice": "auto", + "parallel_tool_calls": true, + "tools": [] + }); + + apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses"); + + for field in XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS { + assert!(body.get(*field).is_none(), "{field} should be stripped"); + } + assert!(body.get("tool_choice").is_none()); + assert!(body.get("parallel_tool_calls").is_none()); + assert!(body.get("tools").is_none()); + assert_eq!( + body["include"], + json!(["reasoning.encrypted_content", "file_search_call.results"]) + ); + assert_eq!(body["model"], "grok-4.6"); + assert_eq!(body["input"], "hello"); + } + + #[test] + fn xai_responses_edits_keep_reasoning_effort_for_thinking_models() { + let mut body = json!({ + "model": "grok-4.6", + "reasoning": {"effort": "high", "summary": "auto"} + }); + apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses"); + assert_eq!(body["reasoning"]["effort"], "high"); + assert_eq!(body["reasoning"]["summary"], "auto"); + } + + #[test] + fn xai_responses_edits_strip_reasoning_effort_for_non_thinking_models() { + let mut body = json!({ + "model": "grok-4.20-0309-non-reasoning", + "reasoning": {"effort": "high"} + }); + apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses"); + assert!(body.get("reasoning").is_none()); + assert!(!xai_model_supports_reasoning_effort( + "grok-4.20-0309-non-reasoning" + )); + assert!(xai_model_supports_reasoning_effort("xai/grok-4.5")); + assert!(!xai_model_supports_reasoning_effort("grok-imagine-image")); + } + + #[test] + fn xai_hosted_tool_choice_rewrites_web_search_and_image_generation() { + let mut web_search = json!({ + "model": "grok-4.6", + "tools": [{"type": "web_search_preview", "name": "web_search"}], + "tool_choice": {"type": "web_search"} + }); + apply_xai_upstream_payload_edits(&mut web_search, "xai", "openai:responses"); + assert_eq!(web_search["tools"][0]["type"], "web_search"); + assert!(web_search["tools"][0].get("name").is_none()); + assert_eq!(web_search["tool_choice"]["type"], "allowed_tools"); + assert_eq!(web_search["tool_choice"]["mode"], "required"); + assert_eq!(web_search["tool_choice"]["tools"][0]["type"], "web_search"); + + let mut image = json!({ + "model": "grok-4.6", + "tools": [ + {"type": "web_search"}, + {"type": "image_generation", "action": "generate"} + ], + "tool_choice": {"type": "image_generation"} + }); + apply_xai_upstream_payload_edits(&mut image, "xai", "openai:responses"); + assert_eq!(image["tool_choice"], "required"); + assert_eq!(image["tools"].as_array().map(Vec::len), Some(1)); + assert_eq!(image["tools"][0]["type"], "image_generation"); + } + + #[test] + fn xai_strips_image_generation_on_older_conversation_models() { + let mut body = json!({ + "model": "grok-4.5", + "tools": [ + {"type": "function", "name": "lookup", "parameters": {"type": "object"}}, + {"type": "image_generation"} + ], + "tool_choice": {"type": "image_generation"} + }); + apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses"); + assert_eq!(body["tools"].as_array().map(Vec::len), Some(1)); + assert_eq!(body["tools"][0]["name"], "lookup"); + assert!(body.get("tool_choice").is_none()); + assert!(xai_supports_native_image_generation("grok-4.6")); + assert!(!xai_supports_native_image_generation("grok-4.20-0309")); + assert!(!xai_supports_native_image_generation("grok-4.5")); + } + + #[test] + fn xai_restores_web_search_from_chat_and_claude_clients() { + let mut chat_body = json!({ + "model": "grok-4.6", + "input": "search this" + }); + apply_xai_upstream_payload_edits_with_client( + &mut chat_body, + "xai", + "openai:responses", + Some("openai:chat"), + Some(&json!({ + "messages": [{"role": "user", "content": "news"}], + "web_search_options": {"search_context_size": "high"} + })), + ); + assert_eq!(chat_body["tools"][0]["type"], "web_search"); + + let mut claude_body = json!({ + "model": "grok-4.6", + "input": "search this", + "tools": [{ + "type": "function", + "name": "lookup", + "parameters": {"type": "object", "properties": {}} + }], + "tool_choice": {"type": "function", "name": "web_search"} + }); + apply_xai_upstream_payload_edits_with_client( + &mut claude_body, + "xai", + "openai:responses", + Some("claude:messages"), + Some(&json!({ + "tools": [ + {"type": "web_search_20250305", "name": "web_search"}, + {"name": "lookup", "input_schema": {"type": "object"}} + ], + "tool_choice": {"type": "tool", "name": "web_search"} + })), + ); + assert!(claude_body["tools"] + .as_array() + .into_iter() + .flatten() + .any(|tool| tool["type"] == "web_search")); + assert_eq!(claude_body["tool_choice"]["type"], "allowed_tools"); + } + + #[test] + fn xai_image_refs_rewrite_openai_aliases_without_touching_chat_parts() { + let mut body = json!({ + "model": "grok-imagine-image", + "prompt": "edit this", + "image": {"image_url": "https://cdn.example/a.png"}, + "reference_images": [ + {"image_url": {"url": "https://cdn.example/b.png"}} + ], + "input": [{ + "type": "message", + "content": [{ + "type": "image_url", + "image_url": {"url": "https://cdn.example/chat.png"} + }] + }] + }); + + apply_xai_upstream_payload_edits(&mut body, "xai", "openai:image"); + + assert_eq!(body["image"]["url"], "https://cdn.example/a.png"); + assert!(body["image"].get("image_url").is_none()); + assert_eq!( + body["reference_images"][0]["url"], + "https://cdn.example/b.png" + ); + assert_eq!( + body["input"][0]["content"][0]["image_url"]["url"], + "https://cdn.example/chat.png" + ); + } + + #[test] + fn other_providers_are_left_untouched() { + let mut body = json!({ + "previous_response_id": "resp_123", + "image": {"image_url": "https://cdn.example/a.png"} + }); + apply_xai_upstream_payload_edits(&mut body, "codex", "openai:responses"); + assert_eq!(body["previous_response_id"], "resp_123"); + assert_eq!(body["image"]["image_url"], "https://cdn.example/a.png"); + } +} diff --git a/crates/aether-ai/formats/src/formats/registry.rs b/crates/aether-ai/formats/src/formats/registry.rs index 289a08614..a1d3a4313 100644 --- a/crates/aether-ai/formats/src/formats/registry.rs +++ b/crates/aether-ai/formats/src/formats/registry.rs @@ -3594,6 +3594,108 @@ mod tests { .any(|field| field.field == "messages")); } + /// Gemini runs `googleSearch` server-side, so the only trace of the search + /// is `groundingMetadata`. Clients on the other formats have to receive it + /// as their own native citations or the answer arrives unverifiable. + #[test] + fn gemini_grounding_reaches_every_cross_format_client_as_citations() { + let gemini = grounded_gemini_response(); + + for target in ["openai:chat", "openai:responses"] { + let converted = + convert_response_pure("gemini:generate_content", target, &gemini).expect(target); + let body = serde_json::to_string(&converted.value).expect("serialize"); + let annotations = find_first_array(&converted.value, "annotations") + .unwrap_or_else(|| panic!("{target} dropped the grounding metadata: {body}")); + assert_eq!( + annotations, + &json!([{ + "type": "url_citation", + "url": "https://time.gov/", + "title": "time.gov", + "start_index": 0, + "end_index": 9, + }]), + "{target} annotations" + ); + } + + let converted = + convert_response_pure("gemini:generate_content", "claude:messages", &gemini) + .expect("claude:messages"); + let body = serde_json::to_string(&converted.value).expect("serialize"); + let citations = find_first_array(&converted.value, "citations") + .unwrap_or_else(|| panic!("claude:messages dropped the grounding metadata: {body}")); + assert_eq!( + citations, + &json!([{ + "type": "web_search_result_location", + "url": "https://time.gov/", + "title": "time.gov", + "cited_text": "今天是 2026", + }]) + ); + } + + /// The grounded span is reported in UTF-8 bytes but every target counts + /// characters, so a multi-byte answer must not shift the citation. + #[test] + fn gemini_grounding_offsets_are_converted_from_bytes_to_characters() { + let converted = convert_response_pure( + "gemini:generate_content", + "openai:chat", + &grounded_gemini_response(), + ) + .expect("convert"); + let annotation = + &find_first_array(&converted.value, "annotations").expect("annotations")[0]; + + // "今天是 2026 " is 15 bytes but 9 characters. + assert_eq!(annotation["end_index"], json!(9)); + } + + fn grounded_gemini_response() -> serde_json::Value { + json!({ + "responseId": "resp_grounded", + "modelVersion": "gemini-3.8-flash", + "candidates": [{ + "index": 0, + "finishReason": "STOP", + "groundingMetadata": { + "webSearchQueries": ["current UTC date"], + "groundingChunks": [{ + "web": {"uri": "https://time.gov/", "title": "time.gov"} + }], + "groundingSupports": [{ + "segment": {"startIndex": 0, "endIndex": 15}, + "groundingChunkIndices": [0] + }] + }, + "content": {"parts": [{"text": "今天是 2026 年"}]} + }] + }) + } + + fn find_first_array<'a>( + value: &'a serde_json::Value, + key: &str, + ) -> Option<&'a serde_json::Value> { + match value { + serde_json::Value::Object(object) => { + if let Some(found) = object.get(key).filter(|found| found.is_array()) { + return Some(found); + } + object + .values() + .find_map(|value| find_first_array(value, key)) + } + serde_json::Value::Array(items) => { + items.iter().find_map(|item| find_first_array(item, key)) + } + _ => None, + } + } + #[test] fn runtime_responses_to_gemini_rejects_mixed_tools_for_gemini_two() { let body = json!({ diff --git a/crates/aether-ai/formats/src/formats/shared/citations.rs b/crates/aether-ai/formats/src/formats/shared/citations.rs new file mode 100644 index 000000000..561c268ac --- /dev/null +++ b/crates/aether-ai/formats/src/formats/shared/citations.rs @@ -0,0 +1,113 @@ +//! Provider-neutral source citations. +//! +//! Some providers ground an answer server-side (Gemini's native `googleSearch` +//! is the motivating case): the search leaves no client-visible tool call, and +//! the evidence arrives only as provider-specific metadata alongside the text. +//! Dropping it leaves callers with prose that names its sources but nothing +//! they can render, link, or verify. +//! +//! Adapters therefore normalise that metadata into the neutral citation shape +//! below, and each target renders it into its own family's standard shape. +//! Neither side has to learn the other's vocabulary. + +use serde_json::{Map, Value}; + +/// Build one neutral citation. +/// +/// `start_index` / `end_index` are character offsets into the answer text — +/// providers that report byte offsets convert before calling. Every field but +/// `url` is optional, because providers routinely ground an answer without +/// anchoring it to a span. +pub(crate) fn canonical_citation( + url: &str, + title: Option<&str>, + start_index: Option, + end_index: Option, + cited_text: Option<&str>, +) -> Value { + let mut citation = Map::new(); + citation.insert("url".to_string(), Value::String(url.to_string())); + if let Some(title) = title { + citation.insert("title".to_string(), Value::String(title.to_string())); + } + if let Some(start_index) = start_index { + citation.insert("start_index".to_string(), Value::from(start_index as u64)); + } + if let Some(end_index) = end_index { + citation.insert("end_index".to_string(), Value::from(end_index as u64)); + } + if let Some(cited_text) = cited_text { + citation.insert( + "cited_text".to_string(), + Value::String(cited_text.to_string()), + ); + } + Value::Object(citation) +} + +fn citation_string<'a>(citation: &'a Value, key: &str) -> Option<&'a str> { + citation + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +/// Render a neutral citation as an OpenAI `url_citation` annotation, the shape +/// both `chat.completions` and `responses` attach to assistant text. +pub(crate) fn canonical_citation_to_openai_annotation(citation: &Value) -> Option { + let url = citation_string(citation, "url")?; + let mut annotation = Map::new(); + annotation.insert( + "type".to_string(), + Value::String("url_citation".to_string()), + ); + annotation.insert("url".to_string(), Value::String(url.to_string())); + if let Some(title) = citation_string(citation, "title") { + annotation.insert("title".to_string(), Value::String(title.to_string())); + } + for key in ["start_index", "end_index"] { + if let Some(index) = citation.get(key).and_then(Value::as_u64) { + annotation.insert(key.to_string(), Value::from(index)); + } + } + Some(Value::Object(annotation)) +} + +/// Render a neutral citation as a Claude `web_search_result_location`, the +/// shape Claude puts in a text block's `citations`. +pub(crate) fn canonical_citation_to_claude_citation(citation: &Value) -> Option { + let url = citation_string(citation, "url")?; + let mut out = Map::new(); + out.insert( + "type".to_string(), + Value::String("web_search_result_location".to_string()), + ); + out.insert("url".to_string(), Value::String(url.to_string())); + if let Some(title) = citation_string(citation, "title") { + out.insert("title".to_string(), Value::String(title.to_string())); + } + if let Some(cited_text) = citation_string(citation, "cited_text") { + out.insert( + "cited_text".to_string(), + Value::String(cited_text.to_string()), + ); + } + Some(Value::Object(out)) +} + +/// Render every citation that carries a usable URL. +pub(crate) fn canonical_citations_to_openai_annotations(citations: &[Value]) -> Vec { + citations + .iter() + .filter_map(canonical_citation_to_openai_annotation) + .collect() +} + +/// Render every citation that carries a usable URL. +pub(crate) fn canonical_citations_to_claude_citations(citations: &[Value]) -> Vec { + citations + .iter() + .filter_map(canonical_citation_to_claude_citation) + .collect() +} diff --git a/crates/aether-ai/formats/src/formats/shared/mod.rs b/crates/aether-ai/formats/src/formats/shared/mod.rs index 00fc97d12..0a21b9bdb 100644 --- a/crates/aether-ai/formats/src/formats/shared/mod.rs +++ b/crates/aether-ai/formats/src/formats/shared/mod.rs @@ -6,6 +6,7 @@ use std::fmt; /// a base64 field cannot trigger an unchecked allocation before parsing. pub(crate) const MAX_SYNC_REPORT_BODY_BYTES: usize = 64 * 1024 * 1024; +pub mod citations; pub mod error_body; pub mod family; pub mod image_bridge; diff --git a/crates/aether-ai/formats/src/formats/shared/routing.rs b/crates/aether-ai/formats/src/formats/shared/routing.rs index 3dec3d5b7..c85aec418 100644 --- a/crates/aether-ai/formats/src/formats/shared/routing.rs +++ b/crates/aether-ai/formats/src/formats/shared/routing.rs @@ -49,6 +49,10 @@ pub fn resolve_execution_runtime_stream_plan_kind_with_client_surface( method: &Method, path: &str, ) -> Option<&'static str> { + let path = path + .strip_prefix("/openai") + .filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/")) + .unwrap_or(path); if route_class != Some("ai_public") { return None; } @@ -181,6 +185,10 @@ pub fn resolve_execution_runtime_sync_plan_kind_with_client_surface( method: &Method, path: &str, ) -> Option<&'static str> { + let path = path + .strip_prefix("/openai") + .filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/")) + .unwrap_or(path); if route_class != Some("ai_public") { return None; } @@ -206,7 +214,10 @@ pub fn resolve_execution_runtime_sync_plan_kind_with_client_surface( if route_family == Some("openai") && route_kind == Some("video") && *method == Method::POST - && path == "/v1/videos" + && matches!( + path, + "/v1/videos" | "/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions" + ) { return Some(OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND); } diff --git a/crates/aether-ai/formats/src/formats/shared/standard_matrix.rs b/crates/aether-ai/formats/src/formats/shared/standard_matrix.rs index e46b2d808..9bb764655 100644 --- a/crates/aether-ai/formats/src/formats/shared/standard_matrix.rs +++ b/crates/aether-ai/formats/src/formats/shared/standard_matrix.rs @@ -17,11 +17,23 @@ use crate::formats::openai::responses::codex::{ apply_codex_openai_responses_chat_body_edits, apply_codex_openai_responses_special_body_edits, apply_openai_responses_compact_special_body_edits, }; +use crate::formats::openai::responses::xai::apply_xai_upstream_payload_edits_with_client; use crate::formats::shared::standard_normalize::{ build_local_openai_chat_request_body_with_model_directives, is_claude_messages_shaped_body_on_openai_chat_endpoint, }; +/// Tool schema preservation is a format-conversion policy, shared by the +/// standard matrix and provider-aware Chat/Responses entry points. +pub(super) fn preserves_gemini_tool_schemas( + provider_type: &str, + provider_api_format: &str, +) -> bool { + provider_type.trim().eq_ignore_ascii_case("antigravity") + && aether_ai_formats::normalize_api_format_alias(provider_api_format) + == "gemini:generate_content" +} + #[allow(clippy::too_many_arguments)] pub fn build_standard_request_body( body_json: &Value, @@ -121,10 +133,17 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and enable_model_directives: bool, reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy, ) -> Option { + let reasoning_replay_policy = if provider_type.trim().eq_ignore_ascii_case("xai") { + crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + } else { + reasoning_replay_policy + }; let mut format_context = FormatContext::default() .with_mapped_model(mapped_model) .with_request_path(request_path) .with_upstream_stream(upstream_is_stream); + format_context.preserve_gemini_tool_schemas = + preserves_gemini_tool_schemas(provider_type, provider_api_format); if let Some(history_scope) = user_api_key_id { format_context = format_context.with_history_scope(history_scope); } @@ -133,13 +152,27 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and client_api_format, provider_api_format, ); - // DeepSeek's Responses continuation state is opaque. Parsing a same-wire-format - // request through the canonical model would discard its id-less `reasoning_text` - // items and future provider-owned fields even though no conversion is required. - // Keep that provider-specific route wire-preserving, while retaining canonical - // normalization for ordinary OpenAI Responses and for Responses/Compact - // cross-format conversions. - let mut provider_request_body = if is_wire_preserving_deepseek_responses_hop( + // Keep the specialized OpenAI builders' compatibility/history preprocessing + // when routing them through the provider-aware schema-preserving path. + let antigravity_chat_body = if format_context.preserve_gemini_tool_schemas + && matches!( + aether_ai_formats::normalize_api_format_alias(source_api_format.as_ref()).as_str(), + "openai:chat" | "openai:responses" | "openai:responses:compact" + ) { + Some( + crate::formats::shared::standard_normalize::chat_compatible_body_for_standard_source( + body_json, + source_api_format.as_ref(), + user_api_key_id, + )?, + ) + } else { + None + }; + // DeepSeek and xAI replay opaque provider state. Preserve their native + // Responses input items: canonical conversion can lose reasoning IDs and + // encrypted-only items even when source and destination formats are equal. + let mut provider_request_body = if is_wire_preserving_responses_hop( source_api_format.as_ref(), provider_api_format, reasoning_replay_policy, @@ -149,9 +182,13 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and Value::Object(object) } else { convert_request( - source_api_format.as_ref(), + if antigravity_chat_body.is_some() { + "openai:chat" + } else { + source_api_format.as_ref() + }, provider_api_format, - body_json, + antigravity_chat_body.as_deref().unwrap_or(body_json), &format_context, ) .ok()? @@ -200,6 +237,13 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and &mut provider_request_body, provider_api_format, ); + apply_xai_upstream_payload_edits_with_client( + &mut provider_request_body, + provider_type, + provider_api_format, + Some(client_api_format), + Some(body_json), + ); crate::formats::openai::responses::strip_incompatible_openai_responses_reasoning_items_with_policy( &mut provider_request_body, provider_api_format, @@ -224,14 +268,16 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and Some(provider_request_body) } -fn is_wire_preserving_deepseek_responses_hop( +fn is_wire_preserving_responses_hop( source_api_format: &str, provider_api_format: &str, reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy, ) -> bool { - if reasoning_replay_policy - != crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque - { + if !matches!( + reasoning_replay_policy, + crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque + | crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + ) { return false; } let source_api_format = aether_ai_formats::normalize_api_format_alias(source_api_format); @@ -1939,6 +1985,66 @@ mod tests { assert_eq!(converted["tools"][0]["googleSearch"], json!({})); } + #[test] + fn claude_client_web_search_tool_survives_conversion_to_gemini() { + let request = json!({ + "model": "gemini-3-flash-preview", + "max_tokens": 1024, + "messages": [{"role": "user", "content": "find the release notes"}], + "tools": [ + { + "name": "WebSearch", + "description": "Search the web and use the results to inform responses", + "input_schema": { + "type": "object", + "properties": {"query": {"type": "string"}}, + "required": ["query"] + } + }, + { + "name": "Read", + "description": "Read a file", + "input_schema": { + "type": "object", + "properties": {"file_path": {"type": "string"}}, + "required": ["file_path"] + } + } + ] + }); + + let converted = build_standard_request_body( + &request, + "claude:messages", + "gemini-3-flash-preview", + "google", + "gemini:generate_content", + "/v1/messages", + false, + None, + None, + ) + .expect("claude messages should convert to gemini"); + + let tools = converted["tools"] + .as_array() + .expect("tools should be an array"); + assert!( + tools.iter().all(|tool| tool.get("googleSearch").is_none() + && tool.get("googleSearchRetrieval").is_none()), + "a client-declared WebSearch tool must not become server-side grounding: {tools:?}" + ); + + let declared: Vec<&str> = tools + .iter() + .filter_map(|tool| tool.get("functionDeclarations")) + .filter_map(Value::as_array) + .flatten() + .filter_map(|declaration| declaration.get("name").and_then(Value::as_str)) + .collect(); + assert_eq!(declared, vec!["WebSearch", "Read"], "{tools:?}"); + } + #[test] fn builds_claude_request_from_openai_chat_with_thinking_and_data_url_image() { let request = json!({ @@ -2077,4 +2183,316 @@ mod tests { ); assert_eq!(gemini["toolConfig"]["functionCallingConfig"]["mode"], "ANY"); } + + #[test] + fn xai_keeps_client_search_functions_distinct_from_hosted_search() { + for name in ["web_search", "web_search_internal"] { + for hosted in [false, true] { + let mut tools = vec![json!({ + "name": name, + "description": "Search internal documents", + "input_schema": {"type": "object", "properties": {"query": {"type": "string"}}} + })]; + if hosted { + tools.push(json!({"type": "web_search_20260209", "name": "internet_search"})); + } + let request = json!({ + "model": "source", "max_tokens": 64, + "messages": [{"role": "user", "content": "Search internal documents"}], + "tools": tools, + "tool_choice": {"type": "tool", "name": name} + }); + let converted = build_standard_request_body( + &request, + "claude:messages", + "grok-4.6", + "xai", + "openai:responses", + "/v1/messages", + true, + None, + None, + ) + .unwrap(); + assert_eq!( + converted["tool_choice"], + json!({"type": "function", "name": name}) + ); + assert_eq!( + converted["tools"] + .as_array() + .unwrap() + .iter() + .any(|tool| tool["type"] == "web_search"), + hosted + ); + } + } + + let request = json!({ + "model": "source", "max_tokens": 64, + "messages": [{"role": "user", "content": "Search the internet"}], + "tools": [{"type": "web_search_20260209", "name": "internet_search"}], + "tool_choice": {"type": "tool", "name": "internet_search"} + }); + let converted = build_standard_request_body( + &request, + "claude:messages", + "grok-4.6", + "xai", + "openai:responses", + "/v1/messages", + true, + None, + None, + ) + .unwrap(); + assert_eq!( + converted["tool_choice"], + json!({ + "type": "allowed_tools", "mode": "required", "tools": [{"type": "web_search"}] + }) + ); + } + + #[test] + fn xai_preserves_function_choices_in_chat_and_responses_requests() { + for name in ["web_search", "web_search_internal"] { + for (client, request) in [ + ( + "openai:chat", + json!({ + "messages": [{"role": "user", "content": "search"}], + "tools": [{"type": "function", "function": {"name": name, "parameters": {"type": "object"}}}], + "tool_choice": {"type": "function", "function": {"name": name}} + }), + ), + ( + "openai:responses", + json!({ + "input": "search", + "tools": [{"type": "function", "name": name, "parameters": {"type": "object"}}], + "tool_choice": {"type": "function", "name": name} + }), + ), + ] { + let converted = build_standard_request_body( + &request, + client, + "grok-4.6", + "xai", + "openai:responses", + "/v1/responses", + true, + None, + None, + ) + .unwrap(); + assert_eq!( + converted["tool_choice"], + json!({"type": "function", "name": name}) + ); + assert_eq!(converted["tools"].as_array().unwrap().len(), 1); + } + } + } + + #[test] + fn xai_image_allowed_tools_preserves_mode_and_restricts_available_tools() { + for mode in ["auto", "required"] { + for mixed in [false, true] { + let mut allowed = vec![json!({"type": "image_generation"})]; + if mixed { + allowed.push(json!({"type": "function", "name": "lookup"})); + } + let request = json!({ + "input": "Draw a cat", + "tools": [ + {"type": "web_search"}, {"type": "image_generation"}, + {"type": "function", "name": "lookup", "parameters": {"type": "object"}} + ], + "tool_choice": {"type": "allowed_tools", "mode": mode, "tools": allowed} + }); + let converted = build_standard_request_body( + &request, + "openai:responses", + "grok-4.6", + "xai", + "openai:responses", + "/v1/responses", + true, + None, + None, + ) + .unwrap(); + if mixed { + assert_eq!( + converted["tool_choice"], + json!({ + "type": "allowed_tools", "mode": mode, + "tools": [{"type": "function", "name": "lookup"}] + }) + ); + assert_eq!(converted["tools"].as_array().unwrap().len(), 3); + } else { + assert_eq!(converted["tool_choice"], mode); + assert_eq!(converted["tools"], json!([{"type": "image_generation"}])); + } + } + } + } + + #[test] + fn xai_responses_preserves_requested_encrypted_reasoning_and_replayed_input() { + let reasoning = json!({"type": "reasoning", "id": "550e8400-e29b-41d4-a716-446655440000", "summary": [], "encrypted_content": "opaque-xai-state"}); + let request = json!({ + "input": [reasoning.clone(), {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "Previous answer"}]}, {"role": "user", "content": "Continue"}], + "include": ["reasoning.encrypted_content"], "store": false + }); + let converted = build_standard_request_body( + &request, + "openai:responses", + "grok-4.6", + "xai", + "openai:responses", + "/v1/responses", + true, + None, + None, + ) + .unwrap(); + assert_eq!(converted["include"], request["include"]); + assert_eq!(converted["input"][0], reasoning); + assert_eq!(converted["store"], false); + } + + #[test] + fn xai_standard_conversion_strips_unsupported_responses_fields() { + let request = json!({ + "model": "source-model", + "messages": [{"role": "user", "content": "Hello xAI"}], + "max_tokens": 128, + "stop": ["END"], + "stream_options": {"include_usage": true}, + "metadata": {"user_id": "claude-session"}, + "web_search_options": {"search_context_size": "high"} + }); + let converted = build_standard_request_body( + &request, + "openai:chat", + "grok-4.6", + "xai", + "openai:responses", + "/v1/chat/completions", + true, + None, + None, + ) + .expect("chat should convert onto xAI Responses"); + + assert_eq!(converted["model"], "grok-4.6"); + assert!(converted.get("stop").is_none()); + assert!(converted.get("stream_options").is_none()); + assert!(converted.get("previous_response_id").is_none()); + assert!(converted.get("metadata").is_none()); + assert!(converted.get("input").is_some() || converted.get("messages").is_none()); + assert_eq!(converted["max_output_tokens"], 128); + assert_eq!(converted["tools"][0]["type"], "web_search"); + } + + #[test] + fn xai_standard_conversion_covers_claude_and_gemini_clients() { + let claude = json!({ + "model": "claude-sonnet", + "max_tokens": 64, + "messages": [{"role": "user", "content": "Hello xAI"}], + "metadata": { + "user_id": "{\"device_id\":\"dev-1\",\"account_uuid\":\"acct-1\",\"session_id\":\"sess-1\"}" + }, + "tools": [ + {"type": "web_search_20250305", "name": "web_search"}, + { + "name": "lookup", + "description": "Look something up", + "input_schema": {"type": "object", "properties": {}} + } + ], + "tool_choice": {"type": "tool", "name": "web_search"} + }); + let converted = build_standard_request_body( + &claude, + "claude:messages", + "grok-4.6", + "xai", + "openai:responses", + "/v1/messages", + true, + None, + None, + ) + .expect("claude should convert onto xAI Responses"); + assert_eq!(converted["model"], "grok-4.6"); + assert!(converted.get("metadata").is_none()); + assert!(converted.get("context_management").is_none()); + assert!(converted + .get("include") + .and_then(Value::as_array) + .into_iter() + .flatten() + .any(|item| item == "reasoning.encrypted_content")); + assert!(converted["tools"] + .as_array() + .into_iter() + .flatten() + .any(|tool| tool["type"] == "web_search")); + assert_eq!(converted["tool_choice"]["type"], "allowed_tools"); + assert!(converted.get("input").is_some()); + + let gemini = json!({ + "model": "gemini-2.5-pro", + "contents": [{ + "role": "user", + "parts": [{"text": "Hello xAI"}] + }], + "tools": [{"googleSearch": {}}] + }); + let converted = build_standard_request_body( + &gemini, + "gemini:generate_content", + "grok-4.6", + "xai", + "openai:responses", + "/v1beta/models/gemini-2.5-pro:generateContent", + false, + None, + None, + ) + .expect("gemini should convert onto xAI Responses"); + assert_eq!(converted["model"], "grok-4.6"); + assert_eq!(converted["tools"][0]["type"], "web_search"); + assert!(converted.get("input").is_some()); + + let same_format = json!({ + "model": "grok-4.6", + "input": "hello", + "previous_response_id": "resp_123", + "stop": ["END"], + "metadata": {"user_id": "claude-session"} + }); + let converted = build_standard_request_body( + &same_format, + "openai:responses", + "grok-4.6", + "xai", + "openai:responses", + "/v1/responses", + true, + None, + None, + ) + .expect("same-format xAI Responses should sanitize in place"); + assert!(converted.get("previous_response_id").is_none()); + assert!(converted.get("stop").is_none()); + assert!(converted.get("metadata").is_none()); + } } diff --git a/crates/aether-ai/formats/src/formats/shared/standard_normalize.rs b/crates/aether-ai/formats/src/formats/shared/standard_normalize.rs index de6d1f070..45c344161 100644 --- a/crates/aether-ai/formats/src/formats/shared/standard_normalize.rs +++ b/crates/aether-ai/formats/src/formats/shared/standard_normalize.rs @@ -65,7 +65,7 @@ fn chat_compatible_body_for_openai_chat_endpoint(body_json: &Value) -> Option( +pub(crate) fn chat_compatible_body_for_standard_source<'a>( body_json: &'a Value, client_api_format: &str, history_scope: Option<&str>, @@ -167,6 +167,40 @@ pub fn build_cross_format_openai_chat_request_body( ) } +/// Provider-aware entry point for gateway Chat planners. Keep private schema +/// conversion policy in the format crate while retaining legacy behavior elsewhere. +pub fn build_cross_format_openai_chat_request_body_with_provider_context( + body_json: &Value, + mapped_model: &str, + provider_type: &str, + provider_api_format: &str, + upstream_is_stream: bool, + enable_model_directives: bool, + history_scope: Option<&str>, +) -> Option { + if super::standard_matrix::preserves_gemini_tool_schemas(provider_type, provider_api_format) { + return super::standard_matrix::build_standard_request_body_with_model_directives( + body_json, + "openai:chat", + mapped_model, + provider_type, + provider_api_format, + "", + upstream_is_stream, + None, + history_scope, + enable_model_directives, + ); + } + build_cross_format_openai_chat_request_body_with_model_directives( + body_json, + mapped_model, + provider_api_format, + upstream_is_stream, + enable_model_directives, + ) +} + pub fn build_cross_format_openai_chat_request_body_with_model_directives( body_json: &Value, mapped_model: &str, @@ -342,6 +376,44 @@ pub fn build_cross_format_openai_responses_request_body_with_model_directives( ) } +/// Provider-aware Responses entry point; preserve history scoping and defer +/// private tool schema lowering without exposing provider policy to the gateway. +#[allow(clippy::too_many_arguments)] +pub fn build_cross_format_openai_responses_request_body_with_provider_context( + body_json: &Value, + mapped_model: &str, + client_api_format: &str, + provider_type: &str, + provider_api_format: &str, + upstream_is_stream: bool, + enable_model_directives: bool, + history_scope: Option<&str>, +) -> Option { + if super::standard_matrix::preserves_gemini_tool_schemas(provider_type, provider_api_format) { + return super::standard_matrix::build_standard_request_body_with_model_directives( + body_json, + client_api_format, + mapped_model, + provider_type, + provider_api_format, + "", + upstream_is_stream, + None, + history_scope, + enable_model_directives, + ); + } + build_cross_format_openai_responses_request_body_with_model_directives_and_history_scope( + body_json, + mapped_model, + client_api_format, + provider_api_format, + upstream_is_stream, + enable_model_directives, + history_scope, + ) +} + pub fn build_cross_format_openai_responses_request_body_with_model_directives_and_history_scope( body_json: &Value, mapped_model: &str, @@ -433,6 +505,145 @@ mod tests { }; use serde_json::{json, Value}; + #[test] + fn provider_context_builders_preserve_private_schemas_and_legacy_routes() { + use crate::api::{ + build_cross_format_openai_chat_request_body_with_provider_context as chat, + build_cross_format_openai_responses_request_body_with_provider_context as responses, + }; + let schema = json!({"type":"object", "properties":{"mode":{"const":"fast"}}}); + let chat_input = json!({"model":"client", "messages":[{"role":"user","content":"hi"}], + "tools":[{"type":"function","function":{"name":"probe","parameters":schema}}]}); + let responses_input = json!({"model":"client", "input":"hi", + "tools":[{"type":"function","name":"probe","parameters":schema}]}); + for provider in ["antigravity", " AnTiGrAvItY ", "gemini", "openai"] { + for target in [ + "gemini:generate_content", + "claude:messages", + "openai:responses", + ] { + for stream in [false, true] { + for directives in [false, true] { + for input in [&chat_input, &responses_input] { + let actual = chat( + input, + "claude-test", + provider, + target, + stream, + directives, + Some("seam-test"), + ); + let expected = + if super::super::standard_matrix::preserves_gemini_tool_schemas( + provider, target, + ) { + super::super::standard_matrix::build_standard_request_body_with_model_directives( + input, "openai:chat", "claude-test", provider, target, "", stream, None, Some("seam-test"), directives) + } else { + super::build_cross_format_openai_chat_request_body_with_model_directives( + input, "claude-test", target, stream, directives) + }; + assert!(actual.is_some(), "chat {provider} {target}"); + assert_eq!(actual, expected); + if target == "gemini:generate_content" { + assert_eq!( + actual.unwrap()["tools"][0]["functionDeclarations"][0] + ["parameters"] + == schema, + provider.trim().eq_ignore_ascii_case("antigravity") + ); + } + } + let actual = responses( + &responses_input, + "claude-test", + "openai:responses", + provider, + target, + stream, + directives, + Some("seam-test"), + ); + let expected = + if super::super::standard_matrix::preserves_gemini_tool_schemas( + provider, target, + ) { + super::super::standard_matrix::build_standard_request_body_with_model_directives( + &responses_input, "openai:responses", "claude-test", provider, target, "", stream, None, Some("seam-test"), directives) + } else { + super::build_cross_format_openai_responses_request_body_with_model_directives_and_history_scope( + &responses_input, "claude-test", "openai:responses", target, stream, directives, Some("seam-test")) + }; + // Same-format Responses uses the local builder, not this cross-format API. + assert_eq!( + actual.is_some(), + target != "openai:responses", + "responses {provider} {target}" + ); + assert_eq!(actual, expected); + } + } + } + } + } + + #[test] + fn provider_context_builders_keep_scoped_responses_history() { + use crate::api::{ + build_cross_format_openai_chat_request_body_with_provider_context as chat, + build_cross_format_openai_responses_request_body_with_provider_context as responses, + record_converted_response_history, + }; + let response_id = "resp_provider_context_seam_history"; + let scope = "provider-context-seam-history"; + record_converted_response_history(&json!({ + "needs_conversion":true, "client_api_format":"openai:responses", + "provider_api_format":"openai:chat", "api_key_id":scope, + "original_request_body":{"model":"client", "input":"first"} + }), &json!({"id":response_id, "status":"completed", "output":[{ + "type":"message", "role":"assistant", "content":[{"type":"output_text", "text":"remembered"}] + }]})).expect("seed scoped history"); + let input = json!({"model":"client", "previous_response_id":response_id, "input":"second"}); + for use_chat in [false, true] { + let build = |history_scope| { + if use_chat { + chat( + &input, + "claude-test", + "antigravity", + "gemini:generate_content", + true, + false, + history_scope, + ) + } else { + responses( + &input, + "claude-test", + "openai:responses", + "antigravity", + "gemini:generate_content", + true, + false, + history_scope, + ) + } + }; + if use_chat { + // The legacy Chat alternate-shape path does not hydrate scoped + // Responses history. Preserve that behavior during this refactor. + assert!(build(Some(scope)).is_none()); + continue; + } + let output = build(Some(scope)).expect("expand scoped history"); + assert_eq!(output["contents"][0]["parts"][0]["text"], "first"); + assert_eq!(output["contents"][1]["parts"][0]["text"], "remembered"); + assert_eq!(output["contents"][2]["parts"][0]["text"], "second"); + assert!(build(Some("different-seam-key")).is_none()); + } + } + fn object_keys(value: &Value) -> Vec<&str> { value .as_object() diff --git a/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs b/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs index 9351e6a25..4b05215a4 100644 --- a/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs +++ b/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs @@ -875,6 +875,104 @@ mod tests { format!("event: {event}\n").into_bytes() } + /// Gemini runs `googleSearch` inside Google, so a grounded streaming answer + /// carries its evidence as `groundingMetadata` on the final chunk and never + /// as a tool call. Each client family has to receive it in its own citation + /// shape, or the answer streams out unverifiable. + #[test] + fn streams_gemini_grounding_to_every_client_as_native_citations() { + let text = "今天是 2026 年"; + let first = json!({ + "responseId": "resp_grounded", + "modelVersion": "gemini-3.8-flash", + "candidates": [{ + "index": 0, + "content": {"role": "model", "parts": [{"text": text}]} + }] + }); + let last = json!({ + "responseId": "resp_grounded", + "modelVersion": "gemini-3.8-flash", + "candidates": [{ + "index": 0, + "finishReason": "STOP", + "content": {"role": "model", "parts": [{"text": text}]}, + "groundingMetadata": { + "webSearchQueries": ["current UTC date"], + "groundingChunks": [{ + "web": {"uri": "https://time.gov/", "title": "time.gov"} + }], + "groundingSupports": [{ + "segment": {"startIndex": 0, "endIndex": 15}, + "groundingChunkIndices": [0] + }] + } + }] + }); + + for (client_api_format, marker) in [ + ("openai:chat", "\"annotations\":[{\"type\":\"url_citation\""), + ( + "openai:responses", + "event: response.output_text.annotation.added\n", + ), + ("claude:messages", "\"type\":\"citations_delta\""), + ] { + let context = report_context("gemini:generate_content", client_api_format); + let mut matrix = StreamingStandardFormatMatrix::default(); + let mut output = matrix + .transform_line(&context, data_line(first.clone())) + .expect("text chunk"); + output.extend( + matrix + .transform_line(&context, data_line(last.clone())) + .expect("grounded chunk"), + ); + output.extend(matrix.finish(&context).expect("finish")); + let sse = String::from_utf8(output).expect("valid SSE"); + + assert!( + sse.contains(marker), + "{client_api_format} missing citations: {sse}" + ); + assert!( + sse.contains("https://time.gov/"), + "{client_api_format} missing source url: {sse}" + ); + } + } + + /// The citation frame is emitted once the answer is whole, so a provider + /// that closes the stream without a `finishReason` must still deliver it. + #[test] + fn streams_gemini_grounding_even_when_the_provider_never_sends_a_finish_reason() { + let context = report_context("gemini:generate_content", "openai:chat"); + let mut matrix = StreamingStandardFormatMatrix::default(); + let mut output = matrix + .transform_line( + &context, + data_line(json!({ + "responseId": "resp_grounded", + "modelVersion": "gemini-3.8-flash", + "candidates": [{ + "index": 0, + "content": {"role": "model", "parts": [{"text": "grounded"}]}, + "groundingMetadata": { + "groundingChunks": [{"web": {"uri": "https://time.gov/"}}] + } + }] + })), + ) + .expect("grounded chunk"); + output.extend(matrix.finish(&context).expect("finish")); + let sse = String::from_utf8(output).expect("valid SSE"); + + assert!( + sse.contains("url_citation") && sse.contains("https://time.gov/"), + "{sse}" + ); + } + #[test] fn terminal_observer_marks_malformed_gemini_function_call_as_failure() { let context = report_context("gemini:generate_content", "openai:responses"); @@ -951,7 +1049,11 @@ mod tests { let sse = String::from_utf8(output).expect("reasoning SSE should be utf8"); assert!( - sse.contains("event: response.reasoning_summary_text.delta\n"), + sse.contains("event: response.reasoning_text.delta\n"), + "{sse}" + ); + assert!( + !sse.contains("event: response.reasoning_summary_text.delta\n"), "{sse}" ); assert!(sse.contains("\"delta\":\"checking\""), "{sse}"); diff --git a/crates/aether-ai/formats/src/formats/shared/stream_rewrite.rs b/crates/aether-ai/formats/src/formats/shared/stream_rewrite.rs index d766c5aa2..10a9f3808 100644 --- a/crates/aether-ai/formats/src/formats/shared/stream_rewrite.rs +++ b/crates/aether-ai/formats/src/formats/shared/stream_rewrite.rs @@ -1,3 +1,4 @@ +use std::borrow::Cow; use std::collections::BTreeMap; use serde_json::{json, Map, Value}; @@ -226,7 +227,7 @@ enum AiSurfaceStreamRewriteState { } pub struct AiSurfaceStreamRewriter<'a> { - report_context: &'a Value, + report_context: Cow<'a, Value>, buffered: Vec, state: AiSurfaceStreamRewriteState, } @@ -271,30 +272,39 @@ pub fn maybe_build_ai_surface_stream_rewriter<'a>( }; Some(AiSurfaceStreamRewriter { - report_context, + report_context: Cow::Borrowed(report_context), buffered: Vec::new(), state, }) } impl AiSurfaceStreamRewriter<'_> { + /// Move parser state across task boundaries without replaying captured bytes. + pub fn into_owned(self) -> AiSurfaceStreamRewriter<'static> { + AiSurfaceStreamRewriter { + report_context: Cow::Owned(self.report_context.into_owned()), + buffered: self.buffered, + state: self.state, + } + } + pub fn push_chunk(&mut self, chunk: &[u8]) -> Result, AiSurfaceFinalizeError> { match &mut self.state { AiSurfaceStreamRewriteState::OpenAiImage(state) => { - state.push_chunk(self.report_context, chunk) + state.push_chunk(self.report_context.as_ref(), chunk) } AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(state) => { - state.push_chunk(self.report_context, chunk) + state.push_chunk(self.report_context.as_ref(), chunk) } AiSurfaceStreamRewriteState::ClaudeReadToolSanitize(state) => { - state.push_chunk(self.report_context, chunk) + state.push_chunk(self.report_context.as_ref(), chunk) } AiSurfaceStreamRewriteState::KiroToClaudeCli(state) => { - state.push_chunk(self.report_context, chunk) + state.push_chunk(self.report_context.as_ref(), chunk) } AiSurfaceStreamRewriteState::KiroToClaudeCliThenStandard { kiro, standard } => { - let claude_bytes = kiro.push_chunk(self.report_context, chunk)?; - transform_standard_bytes(standard, self.report_context, claude_bytes) + let claude_bytes = kiro.push_chunk(self.report_context.as_ref(), chunk)?; + transform_standard_bytes(standard, self.report_context.as_ref(), claude_bytes) } AiSurfaceStreamRewriteState::EnvelopeUnwrap | AiSurfaceStreamRewriteState::ModelDirectiveDisplay @@ -313,23 +323,25 @@ impl AiSurfaceStreamRewriter<'_> { pub fn finish(&mut self) -> Result, AiSurfaceFinalizeError> { match &mut self.state { - AiSurfaceStreamRewriteState::OpenAiImage(state) => state.finish(self.report_context), + AiSurfaceStreamRewriteState::OpenAiImage(state) => { + state.finish(self.report_context.as_ref()) + } AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(state) => { - state.finish(self.report_context) + state.finish(self.report_context.as_ref()) } AiSurfaceStreamRewriteState::ClaudeReadToolSanitize(state) => { - state.finish(self.report_context) + state.finish(self.report_context.as_ref()) } AiSurfaceStreamRewriteState::KiroToClaudeCli(state) => { - state.finish(self.report_context) + state.finish(self.report_context.as_ref()) } AiSurfaceStreamRewriteState::KiroToClaudeCliThenStandard { kiro, standard } => { let mut output = transform_standard_bytes( standard, - self.report_context, - kiro.finish(self.report_context)?, + self.report_context.as_ref(), + kiro.finish(self.report_context.as_ref())?, )?; - output.extend(standard.finish(self.report_context)?); + output.extend(standard.finish(self.report_context.as_ref())?); Ok(output) } AiSurfaceStreamRewriteState::EnvelopeUnwrap @@ -338,14 +350,14 @@ impl AiSurfaceStreamRewriter<'_> { | AiSurfaceStreamRewriteState::Standard(_) => { if self.buffered.is_empty() { if let AiSurfaceStreamRewriteState::Standard(state) = &mut self.state { - return state.finish(self.report_context); + return state.finish(self.report_context.as_ref()); } return Ok(Vec::new()); } let line = std::mem::take(&mut self.buffered); let mut output = self.transform_line(line)?; if let AiSurfaceStreamRewriteState::Standard(state) = &mut self.state { - output.extend(state.finish(self.report_context)?); + output.extend(state.finish(self.report_context.as_ref())?); } Ok(output) } @@ -365,18 +377,19 @@ impl AiSurfaceStreamRewriter<'_> { fn transform_line(&mut self, line: Vec) -> Result, AiSurfaceFinalizeError> { match &mut self.state { AiSurfaceStreamRewriteState::EnvelopeUnwrap => { - let output = transform_provider_private_stream_line(self.report_context, line) - .map_err(AiSurfaceFinalizeError::from)?; - rewrite_model_directive_stream_line(self.report_context, output) + let output = + transform_provider_private_stream_line(self.report_context.as_ref(), line) + .map_err(AiSurfaceFinalizeError::from)?; + rewrite_model_directive_stream_line(self.report_context.as_ref(), output) } AiSurfaceStreamRewriteState::ModelDirectiveDisplay => { - rewrite_model_directive_stream_line(self.report_context, line) + rewrite_model_directive_stream_line(self.report_context.as_ref(), line) } AiSurfaceStreamRewriteState::OpenAiResponsesCompat => { - rewrite_openai_responses_compat_stream_line(self.report_context, line) + rewrite_openai_responses_compat_stream_line(self.report_context.as_ref(), line) } AiSurfaceStreamRewriteState::Standard(state) => { - transform_standard_line(state, self.report_context, line) + transform_standard_line(state, self.report_context.as_ref(), line) } AiSurfaceStreamRewriteState::OpenAiImage(_) | AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(_) @@ -892,7 +905,7 @@ fn is_standard_cli_client_api_format(api_format: &str) -> bool { #[cfg(test)] mod tests { - use serde_json::json; + use serde_json::{json, Value}; use super::{ maybe_build_ai_surface_stream_rewriter, resolve_finalize_stream_rewrite_mode, @@ -1067,6 +1080,50 @@ data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_123\",\"object\ assert!(!output.contains("\"model\":\"gpt-5.5\"")); } + #[test] + fn owned_handoff_preserves_partial_utf8_and_conversion_state() { + for client in ["openai:responses", "openai:chat"] { + let text = "界".repeat(12_000); + let delta = format!( + "data: {}\n\n", + json!({ + "type":"response.output_text.delta", "response_id":"resp_handoff", + "item_id":"msg_handoff", "output_index":0, "content_index":0, "delta":text, + }) + ); + let split = delta.find('界').unwrap() + 17_002; + assert!(!delta.is_char_boundary(split)); + let (mut owned, mut output) = { + let context = json!({"provider_api_format":"openai:responses", + "client_api_format":client, "needs_conversion":client == "openai:chat"}); + let mut parser = maybe_build_ai_surface_stream_rewriter(Some(&context)).unwrap(); + let output = parser.push_chunk(&delta.as_bytes()[..split]).unwrap(); + (parser.into_owned(), output) + }; + output.extend(owned.push_chunk(&delta.as_bytes()[split..]).unwrap()); + output.extend(owned.finish().unwrap()); + let output = String::from_utf8(output).unwrap(); + let events: Vec = output + .lines() + .filter_map(|l| l.strip_prefix("data: ")) + .filter(|p| *p != "[DONE]") + .map(|p| serde_json::from_str(p).unwrap()) + .collect(); + let recovered: String = events + .iter() + .filter_map(|e| { + if client == "openai:responses" { + e["delta"].as_str() + } else { + e.pointer("/choices/0/delta/content") + .and_then(Value::as_str) + } + }) + .collect(); + assert_eq!(recovered, text); + } + } + #[test] fn standard_rewriter_converts_openai_responses_reasoning_delta_to_chat() { let report_context = json!({ diff --git a/crates/aether-ai/formats/src/formats/shared/sync_products.rs b/crates/aether-ai/formats/src/formats/shared/sync_products.rs index 4a20b159f..93fe07764 100644 --- a/crates/aether-ai/formats/src/formats/shared/sync_products.rs +++ b/crates/aether-ai/formats/src/formats/shared/sync_products.rs @@ -9,7 +9,8 @@ use aether_ai_formats::formats::conversion::response::{ }; use aether_ai_formats::formats::openai::responses::response::ensure_modern_openai_responses_response_fields; use aether_ai_formats::formats::openai::responses::{ - openai_responses_message_item_id, openai_responses_synthetic_reasoning_item_id, + openai_responses_message_item_id, openai_responses_reasoning_text_parts, + openai_responses_synthetic_reasoning_item_id, }; use aether_ai_formats::formats::registry::{convert_response, FormatContext, FormatError}; use aether_ai_formats::{ @@ -26,6 +27,7 @@ use serde_json::{json, Map, Value}; use super::{decode_sync_report_body_base64, AiSurfaceFinalizeError}; use crate::formats::claude::messages::stream::ClaudeProviderState; use crate::formats::gemini::generate_content::stream::GeminiProviderState; +use crate::formats::openai::chat::response::openai_chat_reasoning_texts; use crate::formats::openai::chat::stream::{OpenAIChatProviderState, OpenAIResponsesProviderState}; use crate::formats::shared::model_directives::model_directive_display_model_from_report_context; use crate::formats::shared::response::sanitize_claude_read_tool_inputs; @@ -113,11 +115,26 @@ pub fn maybe_build_standard_cross_format_sync_product_from_normalized_payload( .as_deref() .unwrap_or(provider_api_format); + let aggregated_from_stream = aggregated_stream_body.is_some(); let Some(provider_body_json) = aggregated_stream_body.or_else(|| body_json.cloned()) else { return Ok(None); }; + let projection_fallback_body = aggregated_from_stream.then(|| provider_body_json.clone()); - Ok(maybe_build_standard_cross_format_sync_product( + let product = maybe_build_standard_cross_format_sync_product( + report_kind, + provider_body_api_format, + client_api_format, + report_context, + provider_body_json, + ); + if product.is_some() { + return Ok(product); + } + let Some(provider_body_json) = projection_fallback_body else { + return Ok(None); + }; + Ok(project_validated_openai_responses_stream_sync_product( report_kind, provider_body_api_format, client_api_format, @@ -126,6 +143,46 @@ pub fn maybe_build_standard_cross_format_sync_product_from_normalized_payload( )) } +/// Forced-stream Responses upstreams (Codex, xAI) echo request metadata such as +/// `parallel_tool_calls`, `tools` and encrypted reasoning back in the aggregated +/// body, which the strict cross-format response check refuses. The aggregated +/// body stays the provider body, so the client projection may drop those +/// provider-only fields — mirroring the OpenAI Chat client path. +fn project_validated_openai_responses_stream_sync_product( + report_kind: &str, + provider_api_format: &str, + client_api_format: &str, + report_context: &Value, + provider_body_json: Value, +) -> Option { + let provider_api_format = normalize_openai_responses_family_api_format(provider_api_format); + if !matches!( + provider_api_format.as_str(), + "openai:responses" | "openai:responses:compact" + ) { + return None; + } + let client_api_format = client_api_format.trim().to_ascii_lowercase(); + if is_standard_chat_finalize_kind(report_kind) { + sync_chat_response_conversion_kind(&provider_api_format, &client_api_format)?; + } else if is_standard_cli_finalize_kind(report_kind) { + sync_cli_response_conversion_kind(&provider_api_format, &client_api_format)?; + } else { + return None; + } + let client_body_json = project_validated_openai_responses_stream_to_client( + &provider_body_json, + &client_api_format, + report_context, + )?; + let client_body_json = + client_body_with_report_context_model(client_body_json, report_context, &client_api_format); + Some(StandardCrossFormatSyncProduct { + client_body_json, + provider_body_json, + }) +} + pub fn maybe_build_standard_same_format_sync_body_from_normalized_payload( report_kind: &str, status_code: u16, @@ -1472,6 +1529,14 @@ fn convert_openai_chat_canonical_response_to_openai_chat( fn project_validated_openai_responses_stream_to_openai_chat( body_json: &Value, report_context: &Value, +) -> Option { + project_validated_openai_responses_stream_to_client(body_json, "openai:chat", report_context) +} + +fn project_validated_openai_responses_stream_to_client( + body_json: &Value, + client_api_format: &str, + report_context: &Value, ) -> Option { // The caller retains body_json as provider_body_json. This projection is therefore allowed // to omit provider-only response metadata, but never unknown canonical output blocks. @@ -1488,7 +1553,12 @@ fn project_validated_openai_responses_stream_to_openai_chat( } apply_report_context_model_fallback(&mut canonical.model, report_context); - Some(canonical_to_openai_chat_response(&canonical)) + match client_api_format { + "openai:chat" => Some(canonical_to_openai_chat_response(&canonical)), + "claude:messages" => Some(canonical_to_claude_response(&canonical)), + "gemini:generate_content" => canonical_to_gemini_response(&canonical, report_context), + _ => None, + } } fn openai_chat_response_can_use_single_response_canonical(body_json: &Value) -> bool { @@ -1827,6 +1897,7 @@ fn apply_report_context_model_fallback(model: &mut String, report_context: &Valu struct OpenAIChatChoiceState { role: Option, content: String, + reasoning: String, finish_reason: Option, tool_calls: BTreeMap, } @@ -2195,6 +2266,9 @@ pub fn aggregate_openai_chat_stream_sync_response(body: &[u8]) -> Option if let Some(content) = delta.get("content").and_then(Value::as_str) { state.content.push_str(content); } + for (_, piece) in openai_chat_reasoning_texts(delta) { + state.reasoning.push_str(&piece); + } if let Some(tool_calls) = delta.get("tool_calls").and_then(Value::as_array) { for tool_call in tool_calls { let Some(tool_call_object) = tool_call.as_object() else { @@ -2254,6 +2328,14 @@ pub fn aggregate_openai_chat_stream_sync_response(body: &[u8]) -> Option "role".to_string(), Value::String(state.role.unwrap_or_else(|| "assistant".to_string())), ); + // Reassemble under the spelling this crate emits for Chat clients; the + // provider's own spelling was already normalized away by the parser. + if !state.reasoning.is_empty() { + message.insert( + "reasoning_content".to_string(), + Value::String(state.reasoning), + ); + } if state.tool_calls.is_empty() { message.insert("content".to_string(), Value::String(state.content)); } else { @@ -2457,7 +2539,7 @@ fn aggregate_openai_responses_stream_sync_response_from_validated_terminal( reasoning_states .entry(output_index) .or_default() - .summary_text + .reasoning_text .push_str(delta); } "response.reasoning_text.done" | "response.reasoning_summary_text.done" => { @@ -2785,7 +2867,7 @@ struct OpenAIResponsesSyncMessageState { #[derive(Default)] struct OpenAIResponsesSyncReasoningState { item: Map, - summary_text: String, + reasoning_text: String, } #[derive(Default)] @@ -3099,8 +3181,8 @@ fn merge_openai_responses_reasoning_text( if text.is_empty() { return; } - if state.summary_text.is_empty() || text.len() >= state.summary_text.len() { - state.summary_text = text.to_string(); + if state.reasoning_text.is_empty() || text.len() >= state.reasoning_text.len() { + state.reasoning_text = text.to_string(); } } @@ -3117,19 +3199,27 @@ fn merge_openai_responses_tool_arguments( } fn extract_openai_responses_reasoning_text(item: &Map) -> Option { - item.get("summary") - .and_then(Value::as_array) + extract_openai_responses_reasoning_parts(item.get("content"), "reasoning_text") + .or_else(|| extract_openai_responses_reasoning_parts(item.get("summary"), "summary_text")) +} + +fn extract_openai_responses_reasoning_parts( + raw: Option<&Value>, + expected_type: &str, +) -> Option { + raw.and_then(Value::as_array) .into_iter() .flatten() .find_map(|part| { let part = part.as_object()?; - (part.get("type").and_then(Value::as_str) == Some("summary_text")).then(|| { + (part.get("type").and_then(Value::as_str) == Some(expected_type)).then(|| { part.get("text") .and_then(Value::as_str) .unwrap_or_default() .to_string() }) }) + .filter(|text| !text.is_empty()) } fn merge_openai_responses_message_item( @@ -3275,18 +3365,28 @@ fn materialize_openai_responses_reasoning_item( }); item.entry("status".to_string()) .or_insert_with(|| Value::String("completed".to_string())); - if !state.summary_text.is_empty() { - item.insert( - "summary".to_string(), - Value::Array(vec![json!({ - "type": "summary_text", - "text": state.summary_text, - })]), - ); + if !state.reasoning_text.is_empty() + && reasoning_item_field_missing_or_empty(item.get("content")) + { + let content = openai_responses_reasoning_text_parts([&state.reasoning_text]); + item.insert("content".to_string(), content); } + // Raw chain-of-thought lives on `content` only; never mirror it onto + // `summary`, or clients that render both channels show it twice. + item.entry("summary".to_string()) + .or_insert_with(|| Value::Array(Vec::new())); Value::Object(item) } +fn reasoning_item_field_missing_or_empty(value: Option<&Value>) -> bool { + match value { + None | Some(Value::Null) => true, + Some(Value::Array(parts)) => parts.is_empty(), + Some(Value::String(text)) => text.trim().is_empty(), + _ => false, + } +} + fn materialize_openai_responses_tool_item( output_index: usize, state: OpenAIResponsesSyncToolState, @@ -3631,6 +3731,11 @@ fn try_aggregate_gemini_stream_sync_response( CanonicalStreamEvent::TextDelta(text) => { append_gemini_text_part(&mut parts, text, false); } + // This rebuilds a raw Gemini body, and every non-`content` + // candidate key — `groundingMetadata` included — is already + // copied across above. Projecting it into citations is the + // job of whoever converts that body onward. + CanonicalStreamEvent::Citations(_) => {} CanonicalStreamEvent::ReasoningDelta(text) => { append_gemini_text_part(&mut parts, text, true); } @@ -4139,6 +4244,92 @@ mod tests { ); } + #[test] + fn aggregates_openai_chat_stream_reasoning_into_sync_body() { + // The aggregator used to keep only `content` and `tool_calls`, so a + // stream downgraded to a sync response lost the reasoning entirely — + // for OpenRouter's `reasoning`/`reasoning_details` and for the + // DeepSeek-style `reasoning_content` alike. + let body = concat!( + "data: {\"id\":\"gen-openrouter-123\",\"object\":\"chat.completion.chunk\",\"model\":\"stealth/ox-alpha\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"\",\"role\":\"assistant\",\"reasoning\":\"Let me\",\"reasoning_details\":[{\"type\":\"reasoning.text\",\"text\":\"Let me\",\"index\":0}]},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"gen-openrouter-123\",\"object\":\"chat.completion.chunk\",\"model\":\"stealth/ox-alpha\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"\",\"role\":\"assistant\",\"reasoning\":\" think.\",\"reasoning_details\":[{\"type\":\"reasoning.text\",\"text\":\" think.\",\"index\":0}]},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"gen-openrouter-123\",\"object\":\"chat.completion.chunk\",\"model\":\"stealth/ox-alpha\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Done.\",\"role\":\"assistant\"},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"gen-openrouter-123\",\"object\":\"chat.completion.chunk\",\"model\":\"stealth/ox-alpha\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"\",\"role\":\"assistant\",\"reasoning\":null},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":2,\"total_tokens\":3}}\n\n", + ); + + let result = aggregate_openai_chat_stream_sync_response(body.as_bytes()) + .expect("openrouter chat stream should aggregate into a sync body"); + + let message = &result["choices"][0]["message"]; + assert_eq!(message["content"], "Done."); + // `reasoning` and `reasoning_details` repeat one another, so the + // reassembled text must not double up. + assert_eq!(message["reasoning_content"], "Let me think."); + assert_eq!(result["choices"][0]["finish_reason"], "stop"); + } + + #[test] + fn aggregates_deepseek_reasoning_content_into_sync_body() { + let body = concat!( + "data: {\"id\":\"chatcmpl-deepseek\",\"object\":\"chat.completion.chunk\",\"model\":\"deepseek-reasoner\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"reasoning_content\":\"Let me\"},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"chatcmpl-deepseek\",\"object\":\"chat.completion.chunk\",\"model\":\"deepseek-reasoner\",\"choices\":[{\"index\":0,\"delta\":{\"reasoning_content\":\" think.\"},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"chatcmpl-deepseek\",\"object\":\"chat.completion.chunk\",\"model\":\"deepseek-reasoner\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"42\"},\"finish_reason\":\"stop\"}]}\n\n", + ); + + let result = aggregate_openai_chat_stream_sync_response(body.as_bytes()) + .expect("deepseek chat stream should aggregate into a sync body"); + + let message = &result["choices"][0]["message"]; + assert_eq!(message["content"], "42"); + assert_eq!(message["reasoning_content"], "Let me think."); + assert_eq!(result["choices"][0]["finish_reason"], "stop"); + } + + #[test] + fn aggregated_chat_stream_without_reasoning_adds_no_reasoning_key() { + let body = "data: {\"id\":\"chatcmpl-openai\",\"object\":\"chat.completion.chunk\",\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"hi\"},\"finish_reason\":\"stop\"}]}\n\n"; + + let result = aggregate_openai_chat_stream_sync_response(body.as_bytes()) + .expect("plain chat stream should aggregate into a sync body"); + + let message = &result["choices"][0]["message"]; + assert_eq!(message["content"], "hi"); + assert!( + message.get("reasoning_content").is_none(), + "a stream with no reasoning must not gain a reasoning key: {message}" + ); + } + + #[test] + fn aggregated_openai_chat_reasoning_reaches_every_client_format() { + let body = concat!( + "data: {\"id\":\"gen-openrouter-123\",\"object\":\"chat.completion.chunk\",\"model\":\"stealth/ox-alpha\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"\",\"role\":\"assistant\",\"reasoning\":\"Thinking.\"},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"gen-openrouter-123\",\"object\":\"chat.completion.chunk\",\"model\":\"stealth/ox-alpha\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Done.\",\"role\":\"assistant\"},\"finish_reason\":\"stop\"}]}\n\n", + ); + let aggregated = aggregate_openai_chat_stream_sync_response(body.as_bytes()) + .expect("openrouter chat stream should aggregate into a sync body"); + let report_context = json!({}); + + for (client_api_format, marker) in [ + ("openai:responses", "\"type\":\"reasoning\""), + ("claude:messages", "\"type\":\"thinking\""), + ("gemini:generate_content", "\"thought\":true"), + ] { + let converted = convert_standard_chat_response( + &aggregated, + "openai:chat", + client_api_format, + &report_context, + ) + .unwrap_or_else(|| panic!("{client_api_format} should convert")); + let encoded = serde_json::to_string(&converted).expect("converted body should encode"); + assert!( + encoded.contains(marker), + "{client_api_format} dropped the reasoning block: {encoded}" + ); + } + } + #[test] fn aggregates_openai_chat_stream_tool_usage_and_finish_into_sync_body() { let body = concat!( @@ -5606,7 +5797,9 @@ mod tests { .expect("modern response.done stream should aggregate"); assert_eq!(result["output"][0]["type"], "reasoning"); - assert_eq!(result["output"][0]["summary"][0]["text"], "Need care"); + assert_eq!(result["output"][0]["summary"], json!([])); + assert_eq!(result["output"][0]["content"][0]["type"], "reasoning_text"); + assert_eq!(result["output"][0]["content"][0]["text"], "Need care"); assert!(result["output"].as_array().is_some()); assert_eq!(result["output_text"], ""); assert!(result["completed_at"].as_i64().is_some()); @@ -5632,7 +5825,7 @@ mod tests { .as_object() .expect("reasoning item should be an object") .clone(), - summary_text: "must not replace provider-owned state".to_string(), + reasoning_text: "must not replace provider-owned state".to_string(), }; let materialized = materialize_openai_responses_reasoning_item("resp_opaque_123", state); @@ -5673,13 +5866,13 @@ mod tests { } #[test] - fn synthesizes_wire_compatible_id_for_local_reasoning_summary() { + fn synthesizes_wire_compatible_id_for_local_reasoning_text() { let state = OpenAIResponsesSyncReasoningState { item: json!({"type": "reasoning"}) .as_object() .expect("reasoning item should be an object") .clone(), - summary_text: "Need care".to_string(), + reasoning_text: "Need care".to_string(), }; let materialized = materialize_openai_responses_reasoning_item("resp_summary_123", state); @@ -5688,7 +5881,9 @@ mod tests { materialized["id"], openai_responses_synthetic_reasoning_item_id("resp_summary_123", 0) ); - assert_eq!(materialized["summary"][0]["text"], "Need care"); + assert_eq!(materialized["summary"], json!([])); + assert_eq!(materialized["content"][0]["type"], "reasoning_text"); + assert_eq!(materialized["content"][0]["text"], "Need care"); } #[test] @@ -7016,6 +7211,84 @@ mod tests { ); } + #[test] + fn standard_sync_finalize_projects_forced_responses_stream_to_gemini_and_claude_clients() { + // Shape of a forced-stream xAI / Codex upstream: the terminal response + // echoes request metadata and carries encrypted reasoning, which the + // strict cross-format check refuses. + let stream_body = concat!( + "data: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"id\":\"resp_forced_123\",\"object\":\"response\",\"status\":\"in_progress\",\"model\":\"grok-4.7-build\",\"output\":[],\"parallel_tool_calls\":true,\"tool_choice\":\"auto\",\"tools\":[],\"temperature\":0.7}}\n\n", + "data: {\"type\":\"response.output_item.done\",\"sequence_number\":1,\"output_index\":0,\"item\":{\"id\":\"rs_forced_123\",\"type\":\"reasoning\",\"status\":\"completed\",\"summary\":[{\"type\":\"summary_text\",\"text\":\"greet briefly\"}],\"encrypted_content\":\"opaque-xai-reasoning\"}}\n\n", + "data: {\"type\":\"response.output_text.delta\",\"sequence_number\":2,\"item_id\":\"msg_forced_123\",\"output_index\":1,\"content_index\":0,\"delta\":\"Hello there friend\"}\n\n", + "data: {\"type\":\"response.output_item.done\",\"sequence_number\":3,\"output_index\":1,\"item\":{\"id\":\"msg_forced_123\",\"type\":\"message\",\"status\":\"completed\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"Hello there friend\",\"annotations\":[]}]}}\n\n", + "data: {\"type\":\"response.output_item.done\",\"sequence_number\":4,\"output_index\":2,\"item\":{\"id\":\"fc_forced_123\",\"type\":\"function_call\",\"status\":\"completed\",\"call_id\":\"call_forced_123\",\"name\":\"search\",\"arguments\":\"{\\\"q\\\":\\\"aether\\\"}\"}}\n\n", + "data: {\"type\":\"response.completed\",\"sequence_number\":5,\"response\":{\"id\":\"resp_forced_123\",\"object\":\"response\",\"status\":\"completed\",\"model\":\"grok-4.7-build\",\"output\":[],\"parallel_tool_calls\":true,\"tool_choice\":\"auto\",\"tools\":[{\"type\":\"function\",\"name\":\"search\",\"parameters\":{\"type\":\"object\"}}],\"text\":{\"format\":{\"type\":\"text\"}},\"reasoning\":{\"effort\":null,\"summary\":null},\"temperature\":0.7,\"top_p\":0.95,\"store\":false,\"usage\":{\"input_tokens\":1249,\"input_tokens_details\":{\"cached_tokens\":1152},\"output_tokens\":40,\"output_tokens_details\":{\"reasoning_tokens\":31},\"total_tokens\":1289}}}\n\n", + ); + let encoded = base64::engine::general_purpose::STANDARD.encode(stream_body); + + for (report_kind, client_api_format) in [ + ("gemini_chat_sync_finalize", "gemini:generate_content"), + ("gemini_cli_sync_finalize", "gemini:generate_content"), + ("claude_chat_sync_finalize", "claude:messages"), + ("claude_cli_sync_finalize", "claude:messages"), + ] { + let report_context = json!({ + "provider_api_format": "openai:responses", + "provider_stream_event_api_format": "openai:responses", + "client_api_format": client_api_format, + "model": "grok-4.7", + "mapped_model": "grok-4.7", + "needs_conversion": true, + }); + let product = maybe_build_standard_sync_finalize_product_from_normalized_payload( + report_kind, + 200, + Some(&report_context), + None, + Some(&encoded), + ) + .expect("forced Responses stream should aggregate") + .unwrap_or_else(|| panic!("{report_kind} should receive a projection")); + let StandardSyncFinalizeNormalizedProduct::CrossFormat(product) = product else { + panic!("{report_kind}: Responses stream should stay a cross-format product") + }; + assert_eq!(product.provider_body_json["parallel_tool_calls"], true); + let client = product.client_body_json.to_string(); + assert!( + !client.contains("opaque-xai-reasoning") && !client.contains("response.created"), + "{report_kind}: provider-only data leaked into the client body: {client}" + ); + if client_api_format == "gemini:generate_content" { + let parts = product.client_body_json["candidates"][0]["content"]["parts"] + .as_array() + .expect("gemini parts"); + assert!(parts + .iter() + .any(|part| part["text"] == "Hello there friend" + && part.get("thought").is_none())); + assert!(parts + .iter() + .any(|part| part["functionCall"]["name"] == "search" + && part["functionCall"]["args"]["q"] == "aether")); + assert_eq!( + product.client_body_json["usageMetadata"]["promptTokenCount"], + 1249 + ); + } else { + let content = product.client_body_json["content"] + .as_array() + .expect("claude content"); + assert!(content + .iter() + .any(|block| block["type"] == "text" && block["text"] == "Hello there friend")); + assert!(content.iter().any(|block| block["type"] == "tool_use" + && block["name"] == "search" + && block["input"]["q"] == "aether")); + assert_eq!(product.client_body_json["stop_reason"], "tool_use"); + } + } + } + #[test] fn standard_sync_finalize_projects_authoritative_incomplete_responses_stream() { let report_context = json!({ diff --git a/crates/aether-ai/formats/src/lib.rs b/crates/aether-ai/formats/src/lib.rs index be40c6d6f..a7af48b23 100644 --- a/crates/aether-ai/formats/src/lib.rs +++ b/crates/aether-ai/formats/src/lib.rs @@ -1,11 +1,16 @@ extern crate self as aether_ai_formats; pub mod api; +pub mod codex_profile; pub mod contracts; pub mod formats; pub mod protocol; pub mod provider_compat; +pub use codex_profile::{ + codex_client_originator, codex_client_profile, codex_client_user_agent, codex_client_version, + set_codex_cli_version, set_codex_client_profile, CodexClientKind, CodexClientProfile, +}; pub use contracts::{ApiOperation, ClientSurface}; pub use formats::context::{ @@ -50,12 +55,15 @@ pub use formats::openai::responses::codex::{ codex_responses_lite_tool_is_client_executed, effective_codex_model_cards, parse_codex_auth_identity, project_codex_catalog_model_card, resolve_codex_responses_model_capabilities, CodexAuthIdentity, CodexResponsesModelCapabilities, - CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT, CODEX_CLIENT_VERSION, CODEX_MODEL_CATALOG_METADATA_FIELD, CODEX_RESPONSES_LITE_HEADER, }; pub use formats::openai::responses::request::{ validate_openai_responses_request_contract, OpenAiResponsesRequestContractViolation, }; +pub use formats::openai::responses::xai::{ + apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client, + xai_model_supports_reasoning_effort, xai_supports_native_image_generation, +}; pub use formats::openai::responses::{ normalize_openai_responses_message_item_ids, openai_responses_message_item_id, openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id, diff --git a/crates/aether-ai/formats/src/protocol/canonical.rs b/crates/aether-ai/formats/src/protocol/canonical.rs index 55b0d7661..52db51749 100644 --- a/crates/aether-ai/formats/src/protocol/canonical.rs +++ b/crates/aether-ai/formats/src/protocol/canonical.rs @@ -16,6 +16,7 @@ pub use crate::protocol::stream::{CanonicalStreamEvent, CanonicalStreamFrame}; pub(crate) const OPENAI_RESPONSES_EXTENSION_NAMESPACE: &str = "openai_responses"; pub(crate) const OPENAI_RESPONSES_LEGACY_EXTENSION_NAMESPACE: &str = "openai_cli"; +pub(crate) const CLAUDE_EXTENSION_NAMESPACE: &str = "claude"; const AETHER_EXTENSION_NAMESPACE: &str = "aether"; const CLAUDE_MESSAGES_REQUEST_SOURCE_MARKER: &str = "claude_messages_request"; const CLAUDE_SYSTEM_SOURCE_MARKER: &str = "claude_system"; @@ -2860,9 +2861,9 @@ fn openai_responses_reasoning_block_from_item( } fn openai_responses_reasoning_text(item_object: &Map) -> String { - let mut parts = openai_responses_reasoning_text_parts(item_object.get("summary")); + let mut parts = openai_responses_reasoning_text_parts(item_object.get("content")); if parts.is_empty() { - parts = openai_responses_reasoning_text_parts(item_object.get("content")); + parts = openai_responses_reasoning_text_parts(item_object.get("summary")); } parts.join("\n") } @@ -2961,38 +2962,47 @@ pub(crate) fn openai_responses_output_to_canonical( .and_then(Value::as_str) .filter(|value| !value.is_empty()) .map(ToOwned::to_owned); - if let Some(summary_items) = item_object.get("summary").and_then(Value::as_array) { - for summary in summary_items { - let Some(summary_object) = summary.as_object() else { - continue; - }; - let text = summary_object - .get("text") - .and_then(Value::as_str) - .unwrap_or_default(); - if text.trim().is_empty() { - continue; - } - let mut extensions = openai_responses_extensions( - item_object, - &["type", "id", "status", "summary", "encrypted_content"], - ); - canonical_extension_object_mut(&mut extensions, "openai") - .insert("omit_reasoning_parts".to_string(), Value::Bool(true)); - let extensions = openai_thinking_extensions(extensions); - blocks.push(CanonicalContentBlock::Thinking { - text: text.to_string(), - signature: None, - encrypted_content: encrypted_content.clone(), - extensions, - }); - emitted = true; + let mut texts = openai_responses_reasoning_text_parts(item_object.get("content")); + if texts.is_empty() { + texts = openai_responses_reasoning_text_parts(item_object.get("summary")); + } + for text in texts { + if text.trim().is_empty() { + continue; } + let mut extensions = openai_responses_extensions( + item_object, + &[ + "type", + "id", + "status", + "summary", + "content", + "encrypted_content", + ], + ); + canonical_extension_object_mut(&mut extensions, "openai") + .insert("omit_reasoning_parts".to_string(), Value::Bool(true)); + let extensions = openai_thinking_extensions(extensions); + blocks.push(CanonicalContentBlock::Thinking { + text, + signature: None, + encrypted_content: encrypted_content.clone(), + extensions, + }); + emitted = true; } if !emitted && encrypted_content.is_some() { let mut extensions = openai_responses_extensions( item_object, - &["type", "id", "status", "summary", "encrypted_content"], + &[ + "type", + "id", + "status", + "summary", + "content", + "encrypted_content", + ], ); canonical_extension_object_mut(&mut extensions, "openai") .insert("omit_reasoning_parts".to_string(), Value::Bool(true)); @@ -4925,13 +4935,22 @@ pub(crate) fn gemini_response_format_to_canonical( if response_mime_type != "application/json" { return None; } - let json_schema = gemini_value_by_case(generation_config, "responseSchema", "response_schema") - .map(|schema| { - json!({ - "name": "response_schema", - "schema": schema, - }) - }); + let json_schema = gemini_value_by_case( + generation_config, + "responseJsonSchema", + "response_json_schema", + ) + .cloned() + .or_else(|| { + gemini_value_by_case(generation_config, "responseSchema", "response_schema") + .map(gemini_openapi_schema_to_json_schema) + }) + .map(|schema| { + json!({ + "name": "response_schema", + "schema": schema, + }) + }); Some(CanonicalResponseFormat { format_type: if json_schema.is_some() { "json_schema".to_string() @@ -4943,6 +4962,68 @@ pub(crate) fn gemini_response_format_to_canonical( }) } +/// `parametersJsonSchema` is already standard JSON Schema; the legacy +/// `parameters` field is Gemini's OpenAPI subset with upper-case type names. +fn gemini_declaration_parameters_to_json_schema(declaration: &Map) -> Option { + gemini_value_by_case( + declaration, + "parametersJsonSchema", + "parameters_json_schema", + ) + .cloned() + .or_else(|| { + declaration + .get("parameters") + .map(gemini_openapi_schema_to_json_schema) + }) +} + +/// Gemini's OpenAPI-style `Schema` spells types in upper case (`OBJECT`, +/// `STRING`, ...); other protocols expect JSON Schema's lower-case names. +pub(crate) fn gemini_openapi_schema_to_json_schema(schema: &Value) -> Value { + fn normalize(value: &mut Value) { + match value { + Value::Object(object) => { + for (key, child) in object.iter_mut() { + if key == "type" { + match child { + Value::String(type_name) => lowercase_schema_type(type_name), + Value::Array(type_names) => { + for type_name in type_names.iter_mut() { + if let Value::String(type_name) = type_name { + lowercase_schema_type(type_name); + } + } + } + other => normalize(other), + } + } else if key != "enum" + && key != "const" + && key != "default" + && key != "example" + { + normalize(child); + } + } + } + Value::Array(items) => items.iter_mut().for_each(normalize), + _ => {} + } + } + fn lowercase_schema_type(type_name: &mut String) { + if matches!( + type_name.as_str(), + "OBJECT" | "STRING" | "INTEGER" | "NUMBER" | "BOOLEAN" | "ARRAY" | "NULL" + ) { + *type_name = type_name.to_ascii_lowercase(); + } + } + + let mut schema = schema.clone(); + normalize(&mut schema); + schema +} + pub(crate) type GeminiCanonicalTools = ( Vec, Vec, @@ -5116,12 +5197,18 @@ pub(crate) fn gemini_tools_to_canonical(value: Option<&Value>) -> Option, content: String, }, + /// Provider-neutral source citations for the answer text streamed so far. + /// + /// Emitted once, just before `Finish`, by providers that ground an answer + /// server-side and report the evidence as metadata instead of a tool call. + /// Each entry carries `url` plus optional `title`, `cited_text` and + /// `start_index`/`end_index` character offsets; every target renders them + /// into its own family's citation shape. + Citations(Vec), UnknownEvent(Value), Finish { finish_reason: Option, diff --git a/crates/aether-ai/formats/src/provider_compat/private_envelope.rs b/crates/aether-ai/formats/src/provider_compat/private_envelope.rs index d63f0e858..4adbd1e2f 100644 --- a/crates/aether-ai/formats/src/provider_compat/private_envelope.rs +++ b/crates/aether-ai/formats/src/provider_compat/private_envelope.rs @@ -1,3 +1,4 @@ +use std::borrow::Cow; use std::collections::BTreeMap; use serde_json::Value; @@ -354,7 +355,7 @@ enum ProviderPrivateStreamNormalizeMode { } pub struct ProviderPrivateStreamNormalizer<'a> { - report_context: &'a Value, + report_context: Cow<'a, Value>, buffered: Vec, current_event_type: Option, mode: ProviderPrivateStreamNormalizeMode, @@ -401,7 +402,7 @@ pub fn maybe_build_provider_private_stream_normalizer<'a>( return None; }; Some(ProviderPrivateStreamNormalizer { - report_context, + report_context: Cow::Borrowed(report_context), buffered: Vec::new(), current_event_type: None, mode, @@ -422,10 +423,20 @@ pub fn extract_provider_private_stream_error_body( } impl ProviderPrivateStreamNormalizer<'_> { + /// Move parser state across task boundaries without replaying captured bytes. + pub fn into_owned(self) -> ProviderPrivateStreamNormalizer<'static> { + ProviderPrivateStreamNormalizer { + report_context: Cow::Owned(self.report_context.into_owned()), + buffered: self.buffered, + current_event_type: self.current_event_type, + mode: self.mode, + } + } + pub fn push_chunk(&mut self, chunk: &[u8]) -> Result, AiSurfaceFinalizeError> { match &mut self.mode { ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => { - state.push_chunk(self.report_context, chunk) + state.push_chunk(self.report_context.as_ref(), chunk) } ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => { let next_len = self @@ -441,7 +452,7 @@ impl ProviderPrivateStreamNormalizer<'_> { ))); } self.buffered.extend_from_slice(chunk); - if report_context_is_windsurf_envelope(self.report_context) + if report_context_is_windsurf_envelope(self.report_context.as_ref()) && buffer_looks_like_connect_frame(&self.buffered) { return drain_windsurf_connect_json_frames(&mut self.buffered); @@ -451,7 +462,7 @@ impl ProviderPrivateStreamNormalizer<'_> { let line = self.buffered.drain(..=line_end).collect::>(); output.extend( transform_provider_private_stream_line_with_event_state( - self.report_context, + self.report_context.as_ref(), line, &mut self.current_event_type, ) @@ -466,20 +477,20 @@ impl ProviderPrivateStreamNormalizer<'_> { pub fn finish(&mut self) -> Result, AiSurfaceFinalizeError> { match &mut self.mode { ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => { - state.finish(self.report_context) + state.finish(self.report_context.as_ref()) } ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => { if self.buffered.is_empty() { return Ok(Vec::new()); } - if report_context_is_windsurf_envelope(self.report_context) + if report_context_is_windsurf_envelope(self.report_context.as_ref()) && buffer_looks_like_connect_frame(&self.buffered) { return drain_windsurf_connect_json_frames(&mut self.buffered); } let line = std::mem::take(&mut self.buffered); transform_provider_private_stream_line_with_event_state( - self.report_context, + self.report_context.as_ref(), line, &mut self.current_event_type, ) @@ -939,7 +950,7 @@ fn postprocess_private_response_value(data: &mut Value, report_context: &Value) #[cfg(test)] mod tests { - use serde_json::json; + use serde_json::{json, Value}; use super::{ extract_provider_private_stream_error_body, maybe_build_provider_private_stream_normalizer, @@ -1116,6 +1127,44 @@ mod tests { assert!(text.contains(r#""content":"chunk""#)); } + #[test] + fn owned_handoff_preserves_private_binary_frame() { + let text = "frame".repeat(10_000); + let framed = connect_json_frame( + 0, + &serde_json::to_vec(&json!({ + "responseId":"ws-handoff", "response":{"text":text} + })) + .unwrap(), + ); + let split = 17_735; + let mut normalizer = { + let context = json!({"has_envelope":true, + "envelope_name":"windsurf:GetChatMessage", "provider_api_format":"openai:chat"}); + let mut normalizer = + maybe_build_provider_private_stream_normalizer(Some(&context)).unwrap(); + assert!(normalizer.push_chunk(&framed[..split]).unwrap().is_empty()); + normalizer.into_owned() + }; + let mut output = normalizer.push_chunk(&framed[split..]).unwrap(); + output.extend(normalizer.finish().unwrap()); + let output = String::from_utf8(output).unwrap(); + let events: Vec = output + .lines() + .filter_map(|l| l.strip_prefix("data: ")) + .filter(|p| *p != "[DONE]") + .map(|p| serde_json::from_str(p).unwrap()) + .collect(); + let recovered: String = events + .iter() + .filter_map(|e| { + e.pointer("/choices/0/delta/content") + .and_then(Value::as_str) + }) + .collect(); + assert_eq!(recovered, text); + } + #[test] fn unwraps_windsurf_connect_json_stream_frames() { let report_context = json!({ diff --git a/crates/aether-data/adapters/postgres/migrations/20260923000000_add_admin_wallet_batch_idempotency.sql b/crates/aether-data/adapters/postgres/migrations/20260923000000_add_admin_wallet_batch_idempotency.sql new file mode 100644 index 000000000..edd5cfdc8 --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260923000000_add_admin_wallet_batch_idempotency.sql @@ -0,0 +1,12 @@ +CREATE TABLE IF NOT EXISTS public.admin_user_wallet_balance_batches ( + admin_user_id character varying(64) NOT NULL, + idempotency_key character varying(128) NOT NULL, + request_fingerprint character varying(64) NOT NULL, + target_user_ids jsonb NOT NULL, + missing_user_ids jsonb NOT NULL, + warnings jsonb NOT NULL, + user_outcomes jsonb NOT NULL, + created_at_unix_secs bigint NOT NULL, + updated_at_unix_secs bigint NOT NULL, + CONSTRAINT admin_user_wallet_balance_batches_pkey PRIMARY KEY (admin_user_id, idempotency_key) +); diff --git a/crates/aether-data/adapters/postgres/src/candidate_selection.rs b/crates/aether-data/adapters/postgres/src/candidate_selection.rs index 0b1650263..dddec6d1b 100644 --- a/crates/aether-data/adapters/postgres/src/candidate_selection.rs +++ b/crates/aether-data/adapters/postgres/src/candidate_selection.rs @@ -102,6 +102,11 @@ INNER JOIN LATERAL ( AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') ) + OR ( + LOWER(BTRIM(p.provider_type)) = 'xai' + AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key') + AND LOWER($3) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') + ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -127,7 +132,8 @@ INNER JOIN LATERAL ( 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' ) @@ -187,6 +193,11 @@ WHERE p.is_active = TRUE AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') ) + OR ( + LOWER(BTRIM(p.provider_type)) = 'xai' + AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key') + AND LOWER($3) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') + ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -212,7 +223,8 @@ WHERE p.is_active = TRUE 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' ) @@ -365,6 +377,11 @@ INNER JOIN LATERAL ( AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') ) + OR ( + LOWER(BTRIM(p.provider_type)) = 'xai' + AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key') + AND LOWER($4) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') + ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -390,7 +407,8 @@ INNER JOIN LATERAL ( 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' ) @@ -451,6 +469,11 @@ WHERE p.is_active = TRUE AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') ) + OR ( + LOWER(BTRIM(p.provider_type)) = 'xai' + AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key') + AND LOWER($4) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') + ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -476,7 +499,8 @@ WHERE p.is_active = TRUE 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' ) @@ -632,11 +656,16 @@ WHERE p.is_active = TRUE ) ) ) - OR ( - LOWER(BTRIM(p.provider_type)) = 'grok' - AND LOWER(BTRIM(pak.auth_type)) = 'oauth' - AND LOWER($6) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') - ) + OR ( + LOWER(BTRIM(p.provider_type)) = 'grok' + AND LOWER(BTRIM(pak.auth_type)) = 'oauth' + AND LOWER($6) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') + ) + OR ( + LOWER(BTRIM(p.provider_type)) = 'xai' + AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key') + AND LOWER($6) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') + ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -662,7 +691,8 @@ WHERE p.is_active = TRUE 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' ) @@ -1717,6 +1747,24 @@ mod tests { } } + #[test] + fn candidate_selection_sql_allows_xai_oauth_responses_auth() { + let requested_model_sql = requested_model_selection_sql(); + for sql in [ + LIST_FOR_EXACT_API_FORMAT_SQL, + LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL, + LIST_POOL_KEYS_FOR_GROUP_SQL, + requested_model_sql.as_str(), + ] { + assert!(sql.contains("LOWER(BTRIM(p.provider_type)) = 'xai'")); + assert!(sql.contains("LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')")); + assert!(sql.contains( + "'openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video'" + )); + assert!(sql.contains("'xai'")); + } + } + #[test] fn candidate_selection_sql_allows_windsurf_openai_chat_managed_keys() { let requested_model_sql = requested_model_selection_sql(); diff --git a/crates/aether-data/adapters/postgres/src/usage/cleanup.rs b/crates/aether-data/adapters/postgres/src/usage/cleanup.rs index 21deef1b0..58e1c7d52 100644 --- a/crates/aether-data/adapters/postgres/src/usage/cleanup.rs +++ b/crates/aether-data/adapters/postgres/src/usage/cleanup.rs @@ -687,11 +687,31 @@ async fn cleanup_usage_raw_body_fields( Ok(total_cleaned) } +async fn truncate_usage_body_blobs_table(pool: &PostgresPool) -> Result<(), DataLayerError> { + let mut tx = pool.begin().await.map_err(postgres_error)?; + sqlx::query("SET LOCAL lock_timeout = '2s'") + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + sqlx::query("TRUNCATE TABLE usage_body_blobs") + .execute(&mut *tx) + .await + .map_err(postgres_error)?; + tx.commit().await.map_err(postgres_error)?; + Ok(()) +} + async fn cleanup_usage_compressed_body_fields( pool: &PostgresPool, cutoff_time: DateTime, batch_size: usize, ) -> Result { + if let Err(err) = truncate_usage_body_blobs_table(pool).await { + warn!( + error = %err, + "usage cleanup truncate usage_body_blobs table failed or timed out, falling back to batch deletion" + ); + } let mut total_cleaned = 0usize; loop { let rows = fetch_usage_body_cleanup_rows( diff --git a/crates/aether-data/adapters/postgres/src/usage/mod.rs b/crates/aether-data/adapters/postgres/src/usage/mod.rs index e1b9605b8..b1dddef95 100644 --- a/crates/aether-data/adapters/postgres/src/usage/mod.rs +++ b/crates/aether-data/adapters/postgres/src/usage/mod.rs @@ -1925,6 +1925,33 @@ fn usage_leaderboard_sql_fragments( } } +fn push_usage_user_scope( + builder: &mut QueryBuilder<'_, Postgres>, + column: &str, + user_id: Option<&str>, + user_ids: Option<&[String]>, +) { + if let Some(user_id) = user_id { + builder + .push(" AND ") + .push(column) + .push(" = ") + .push_bind(user_id.to_string()); + } + if let Some(user_ids) = user_ids { + if user_ids.is_empty() { + builder.push(" AND FALSE"); + } else { + builder.push(" AND ").push(column).push(" IN ("); + let mut separated = builder.separated(", "); + for user_id in user_ids { + separated.push_bind(user_id.clone()); + } + separated.push_unseparated(")"); + } + } +} + const LIST_RECENT_USAGE_AUDITS_PREFIX: &str = include_str!("queries/list_recent_usage_audits_prefix.sql"); @@ -3762,14 +3789,15 @@ OR (\"usage\".error_message IS NOT NULL AND BTRIM(\"usage\".error_message) <> '' start_day_utc: DateTime, end_day_utc: DateTime, user_id: Option<&str>, + user_ids: Option<&[String]>, ) -> Result { if start_day_utc >= end_day_utc { return Ok(StoredUsageAuditSummary::default()); } - let row = if let Some(user_id) = user_id { - sqlx::query( - r#" + let scoped_to_users = user_id.is_some() || user_ids.is_some(); + let mut builder = QueryBuilder::::new( + r#" SELECT COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests, COALESCE(SUM(input_tokens), 0)::BIGINT AS input_tokens, @@ -3794,68 +3822,41 @@ SELECT COALESCE(SUM(cache_read_cost), 0)::DOUBLE PRECISION AS cache_read_cost_usd, COALESCE(SUM(response_time_sum_ms), 0)::DOUBLE PRECISION AS total_response_time_ms, COALESCE(SUM(error_requests), 0)::BIGINT AS error_requests -FROM stats_user_daily -WHERE user_id = $1 - AND date >= $2 - AND date < $3 -"#, - ) - .bind(user_id) - .bind(start_day_utc) - .bind(end_day_utc) - .fetch_one(&self.pool) - .await - .map_postgres_err()? +FROM "#, + ); + builder.push(if scoped_to_users { + "stats_user_daily" } else { - sqlx::query( - r#" -SELECT - COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests, - COALESCE(SUM(input_tokens), 0)::BIGINT AS input_tokens, - COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens, - COALESCE(SUM( - CASE - WHEN effective_input_tokens = 0 AND total_input_context = 0 AND input_tokens > 0 - THEN input_tokens - ELSE effective_input_tokens - END - + output_tokens + cache_creation_tokens + cache_read_tokens - ), 0)::BIGINT AS recorded_total_tokens, - COALESCE(SUM(cache_creation_tokens), 0)::BIGINT AS cache_creation_tokens, - COALESCE(SUM(cache_creation_ephemeral_5m_tokens), 0)::BIGINT - AS cache_creation_ephemeral_5m_tokens, - COALESCE(SUM(cache_creation_ephemeral_1h_tokens), 0)::BIGINT - AS cache_creation_ephemeral_1h_tokens, - COALESCE(SUM(cache_read_tokens), 0)::BIGINT AS cache_read_tokens, - COALESCE(SUM(total_cost), 0)::DOUBLE PRECISION AS total_cost_usd, - COALESCE(SUM(actual_total_cost), 0)::DOUBLE PRECISION AS actual_total_cost_usd, - COALESCE(SUM(cache_creation_cost), 0)::DOUBLE PRECISION AS cache_creation_cost_usd, - COALESCE(SUM(cache_read_cost), 0)::DOUBLE PRECISION AS cache_read_cost_usd, - COALESCE(SUM(response_time_sum_ms), 0)::DOUBLE PRECISION AS total_response_time_ms, - COALESCE(SUM(error_requests), 0)::BIGINT AS error_requests -FROM stats_daily -WHERE date >= $1 - AND date < $2 -"#, - ) - .bind(start_day_utc) - .bind(end_day_utc) + "stats_daily" + }); + builder + .push(" WHERE date >= ") + .push_bind(start_day_utc) + .push(" AND date < ") + .push_bind(end_day_utc); + if scoped_to_users { + push_usage_user_scope(&mut builder, "user_id", user_id, user_ids); + } + + let row = builder + .build() .fetch_one(&self.pool) .await - .map_postgres_err()? - }; - + .map_postgres_err()?; decode_usage_audit_summary_row(&row) } async fn summarize_usage_audits_raw( &self, - created_from_unix_secs: u64, - created_until_unix_secs: u64, - user_id: Option<&str>, - provider_name: Option<&str>, - model: Option<&str>, + query: &UsageAuditSummaryQuery, ) -> Result { + let created_from_unix_secs = query.created_from_unix_secs; + let created_until_unix_secs = query.created_until_unix_secs; + let user_id = query.user_id.as_deref(); + let user_ids = query.user_ids.as_deref(); + let provider_names = query.provider_names.as_deref(); + let provider_name = query.provider_name.as_deref(); + let model = query.model.as_deref(); if created_from_unix_secs >= created_until_unix_secs { return Ok(StoredUsageAuditSummary::default()); } @@ -3907,11 +3908,12 @@ FROM usage_billing_facts AS "usage" .push("\"usage\".created_at < TO_TIMESTAMP(") .push_bind(created_until_unix_secs as f64) .push("::double precision)"); - if let Some(user_id) = user_id { - builder.push(if has_where { " AND " } else { " WHERE " }); + push_usage_user_scope(&mut builder, "\"usage\".user_id", user_id, user_ids); + if let Some(names) = provider_names { builder - .push("\"usage\".user_id = ") - .push_bind(user_id.to_string()); + .push(" AND \"usage\".provider_name = ANY(") + .push_bind(names.to_vec()) + .push("::text[])"); } if let Some(provider_name) = provider_name { builder.push(if has_where { " AND " } else { " WHERE " }); @@ -3939,55 +3941,32 @@ FROM usage_billing_facts AS "usage" &self, query: &UsageAuditSummaryQuery, ) -> Result { - if query.provider_name.is_some() || query.model.is_some() { - return self - .summarize_usage_audits_raw( - query.created_from_unix_secs, - query.created_until_unix_secs, - query.user_id.as_deref(), - query.provider_name.as_deref(), - query.model.as_deref(), - ) - .await; + // Provider rollups lack cache-cost/error detail required by this response. + // Keep these scopes on canonical facts rather than inventing missing daily totals. + if query.provider_names.is_some() || query.provider_name.is_some() || query.model.is_some() + { + return self.summarize_usage_audits_raw(query).await; } let Some(cutoff_utc) = self.read_stats_daily_cutoff_date().await? else { - return self - .summarize_usage_audits_raw( - query.created_from_unix_secs, - query.created_until_unix_secs, - query.user_id.as_deref(), - None, - None, - ) - .await; + return self.summarize_usage_audits_raw(query).await; }; let start_utc = dashboard_unix_secs_to_utc(query.created_from_unix_secs); let end_utc = dashboard_unix_secs_to_utc(query.created_until_unix_secs); let split = split_dashboard_daily_aggregate_range(start_utc, end_utc, cutoff_utc); let Some(_) = split.aggregate else { - return self - .summarize_usage_audits_raw( - query.created_from_unix_secs, - query.created_until_unix_secs, - query.user_id.as_deref(), - None, - None, - ) - .await; + return self.summarize_usage_audits_raw(query).await; }; let mut summary = StoredUsageAuditSummary::default(); if let Some((raw_start, raw_end)) = split.raw_leading { absorb_usage_audit_summary( &mut summary, - self.summarize_usage_audits_raw( - dashboard_utc_to_unix_secs(raw_start), - dashboard_utc_to_unix_secs(raw_end), - query.user_id.as_deref(), - None, - None, - ) + self.summarize_usage_audits_raw(&UsageAuditSummaryQuery { + created_from_unix_secs: dashboard_utc_to_unix_secs(raw_start), + created_until_unix_secs: dashboard_utc_to_unix_secs(raw_end), + ..query.clone() + }) .await?, ); } @@ -3998,6 +3977,7 @@ FROM usage_billing_facts AS "usage" aggregate_start, aggregate_end, query.user_id.as_deref(), + query.user_ids.as_deref(), ) .await?, ); @@ -4005,13 +3985,11 @@ FROM usage_billing_facts AS "usage" if let Some((raw_start, raw_end)) = split.raw_trailing { absorb_usage_audit_summary( &mut summary, - self.summarize_usage_audits_raw( - dashboard_utc_to_unix_secs(raw_start), - dashboard_utc_to_unix_secs(raw_end), - query.user_id.as_deref(), - None, - None, - ) + self.summarize_usage_audits_raw(&UsageAuditSummaryQuery { + created_from_unix_secs: dashboard_utc_to_unix_secs(raw_start), + created_until_unix_secs: dashboard_utc_to_unix_secs(raw_end), + ..query.clone() + }) .await?, ); } @@ -6952,12 +6930,17 @@ FROM usage_billing_facts AS "usage" .push("\"usage\".created_at < TO_TIMESTAMP(") .push_bind(query.created_until_unix_secs as f64) .push("::double precision)"); - if let Some(user_id) = query.user_id.as_deref() { - builder.push(if has_where { " AND " } else { " WHERE " }); - has_where = true; + push_usage_user_scope( + &mut builder, + "\"usage\".user_id", + query.user_id.as_deref(), + query.user_ids.as_deref(), + ); + if let Some(names) = query.provider_names.as_ref() { builder - .push("\"usage\".user_id = ") - .push_bind(user_id.to_string()); + .push(" AND \"usage\".provider_name = ANY(") + .push_bind(names.clone()) + .push("::text[])"); } if let Some(provider_name) = query.provider_name.as_deref() { builder.push(if has_where { " AND " } else { " WHERE " }); @@ -6998,62 +6981,54 @@ FROM usage_billing_facts AS "usage" start_day_utc: DateTime, end_day_utc: DateTime, user_id: Option<&str>, + user_ids: Option<&[String]>, + provider_names: Option<&[String]>, ) -> Result, DataLayerError> { if start_day_utc >= end_day_utc { return Ok(Vec::new()); } - let rows = if let Some(user_id) = user_id { - sqlx::query( - r#" + let scoped_to_users = user_id.is_some() || user_ids.is_some(); + let mut builder = QueryBuilder::::new( + r#" SELECT TO_CHAR(date, 'YYYY-MM-DD') AS bucket_key, - total_requests::BIGINT AS total_requests, - input_tokens::BIGINT AS input_tokens, - output_tokens::BIGINT AS output_tokens, - cache_creation_tokens::BIGINT AS cache_creation_tokens, - cache_read_tokens::BIGINT AS cache_read_tokens, - CAST(total_cost AS DOUBLE PRECISION) AS total_cost_usd, - CAST(response_time_sum_ms AS DOUBLE PRECISION) AS total_response_time_ms -FROM stats_user_daily -WHERE user_id = $1 - AND date >= $2 - AND date < $3 -ORDER BY date ASC -"#, - ) - .bind(user_id) - .bind(start_day_utc) - .bind(end_day_utc) - .fetch_all(&self.pool) - .await - .map_postgres_err()? + COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests, + COALESCE(SUM(input_tokens), 0)::BIGINT AS input_tokens, + COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens, + COALESCE(SUM(cache_creation_tokens), 0)::BIGINT AS cache_creation_tokens, + COALESCE(SUM(cache_read_tokens), 0)::BIGINT AS cache_read_tokens, + COALESCE(SUM(CAST(total_cost AS DOUBLE PRECISION)), 0) AS total_cost_usd, + COALESCE(SUM(CAST(response_time_sum_ms AS DOUBLE PRECISION)), 0) + AS total_response_time_ms +FROM "#, + ); + builder.push(if provider_names.is_some() { + "stats_user_daily_provider" + } else if scoped_to_users { + "stats_user_daily" } else { - sqlx::query( - r#" -SELECT - TO_CHAR(date, 'YYYY-MM-DD') AS bucket_key, - total_requests::BIGINT AS total_requests, - input_tokens::BIGINT AS input_tokens, - output_tokens::BIGINT AS output_tokens, - cache_creation_tokens::BIGINT AS cache_creation_tokens, - cache_read_tokens::BIGINT AS cache_read_tokens, - CAST(total_cost AS DOUBLE PRECISION) AS total_cost_usd, - CAST(response_time_sum_ms AS DOUBLE PRECISION) AS total_response_time_ms -FROM stats_daily -WHERE date >= $1 - AND date < $2 -ORDER BY date ASC -"#, - ) - .bind(start_day_utc) - .bind(end_day_utc) - .fetch_all(&self.pool) - .await - .map_postgres_err()? - }; + "stats_daily" + }); + builder + .push(" WHERE date >= ") + .push_bind(start_day_utc) + .push(" AND date < ") + .push_bind(end_day_utc); + if scoped_to_users { + push_usage_user_scope(&mut builder, "user_id", user_id, user_ids); + } + if let Some(names) = provider_names { + builder + .push(" AND provider_name = ANY(") + .push_bind(names.to_vec()) + .push("::text[])"); + } + builder.push(" GROUP BY date ORDER BY date ASC"); + + let mut rows = builder.build().fetch(&self.pool); let mut items = Vec::new(); - for row in rows { + while let Some(row) = rows.try_next().await.map_postgres_err()? { items.push(decode_usage_time_series_bucket_row(&row)?); } Ok(items) @@ -7134,11 +7109,13 @@ WHERE is_complete IS TRUE &mut grouped, self.summarize_usage_time_series_raw( &UsageTimeSeriesQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(raw_start), created_until_unix_secs: dashboard_utc_to_unix_secs(raw_end), granularity: UsageTimeSeriesGranularity::Day, tz_offset_minutes: 0, user_id: query.user_id.clone(), + user_ids: query.user_ids.clone(), provider_name: None, model: None, }, @@ -7153,6 +7130,8 @@ WHERE is_complete IS TRUE aggregate_start, aggregate_end, query.user_id.as_deref(), + query.user_ids.as_deref(), + query.provider_names.as_deref(), ) .await?, ); @@ -7160,11 +7139,13 @@ WHERE is_complete IS TRUE &mut grouped, self.summarize_usage_time_series_raw( &UsageTimeSeriesQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(aggregate_start), created_until_unix_secs: dashboard_utc_to_unix_secs(aggregate_end), granularity: UsageTimeSeriesGranularity::Day, tz_offset_minutes: 0, user_id: query.user_id.clone(), + user_ids: query.user_ids.clone(), provider_name: None, model: None, }, @@ -7177,11 +7158,13 @@ WHERE is_complete IS TRUE &mut grouped, self.summarize_usage_time_series_raw( &UsageTimeSeriesQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(raw_start), created_until_unix_secs: dashboard_utc_to_unix_secs(raw_end), granularity: UsageTimeSeriesGranularity::Day, tz_offset_minutes: 0, user_id: query.user_id.clone(), + user_ids: query.user_ids.clone(), provider_name: None, model: None, }, @@ -7195,7 +7178,11 @@ WHERE is_complete IS TRUE } } - if query.user_id.is_none() && query.tz_offset_minutes % 60 == 0 { + if query.provider_names.is_none() + && query.user_id.is_none() + && query.user_ids.is_none() + && query.tz_offset_minutes % 60 == 0 + { if let Some(cutoff_utc) = self.read_stats_hourly_cutoff().await? { let start_utc = dashboard_unix_secs_to_utc(query.created_from_unix_secs); let end_utc = dashboard_unix_secs_to_utc(query.created_until_unix_secs); @@ -7207,11 +7194,13 @@ WHERE is_complete IS TRUE &mut grouped, self.summarize_usage_time_series_raw( &UsageTimeSeriesQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(raw_start), created_until_unix_secs: dashboard_utc_to_unix_secs(raw_end), granularity: query.granularity, tz_offset_minutes: query.tz_offset_minutes, user_id: None, + user_ids: None, provider_name: None, model: None, }, @@ -7234,11 +7223,13 @@ WHERE is_complete IS TRUE &mut grouped, self.summarize_usage_time_series_raw( &UsageTimeSeriesQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(aggregate_start), created_until_unix_secs: dashboard_utc_to_unix_secs(aggregate_end), granularity: query.granularity, tz_offset_minutes: query.tz_offset_minutes, user_id: None, + user_ids: None, provider_name: None, model: None, }, @@ -7251,11 +7242,13 @@ WHERE is_complete IS TRUE &mut grouped, self.summarize_usage_time_series_raw( &UsageTimeSeriesQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(raw_start), created_until_unix_secs: dashboard_utc_to_unix_secs(raw_end), granularity: query.granularity, tz_offset_minutes: query.tz_offset_minutes, user_id: None, + user_ids: None, provider_name: None, model: None, }, @@ -7295,6 +7288,8 @@ WHERE "usage".created_at >= TO_TIMESTAMP($1::double precision) AND ($3::varchar IS NULL OR "usage".user_id = $3) AND ($4::varchar IS NULL OR "usage".provider_name = $4) AND ($5::varchar IS NULL OR "usage".model = $5) + AND ($6::text[] IS NULL OR "usage".user_id::text = ANY($6)) + AND ($7::text[] IS NULL OR "usage".provider_name = ANY($7)) GROUP BY group_key ORDER BY group_key ASC "#, @@ -7308,6 +7303,8 @@ ORDER BY group_key ASC .bind(query.user_id.as_deref()) .bind(query.provider_name.as_deref()) .bind(query.model.as_deref()) + .bind(query.user_ids.clone()) + .bind(query.provider_names.clone()) .fetch(&self.pool); let mut items = Vec::new(); while let Some(row) = rows.try_next().await.map_postgres_err()? { @@ -7322,6 +7319,51 @@ ORDER BY group_key ASC end_day_utc: DateTime, query: &UsageLeaderboardQuery, ) -> Result>, DataLayerError> { + if let Some(names) = query.provider_names.as_ref() { + let group_key = match query.group_by { + UsageLeaderboardGroupBy::User => "user_id", + UsageLeaderboardGroupBy::Model => "model", + UsageLeaderboardGroupBy::ApiKey => return Ok(None), + }; + let mut builder = QueryBuilder::::new(format!( + r#" +SELECT + {group_key} AS group_key, + MAX(NULLIF(BTRIM(username), '')) AS legacy_name, + COALESCE(SUM(total_requests), 0)::BIGINT AS request_count, + COALESCE(SUM(total_tokens), 0)::BIGINT AS total_tokens, + COALESCE(SUM(total_cost), 0)::DOUBLE PRECISION AS total_cost_usd +FROM stats_user_daily_model_provider +WHERE date >= +"# + )); + builder + .push_bind(start_day_utc) + .push(" AND date < ") + .push_bind(end_day_utc); + builder + .push(" AND provider_name = ANY(") + .push_bind(names.clone()) + .push("::text[])"); + push_usage_user_scope( + &mut builder, + "user_id", + query.user_id.as_deref(), + query.user_ids.as_deref(), + ); + if let Some(provider) = query.provider_name.as_ref() { + builder + .push(" AND provider_name = ") + .push_bind(provider.clone()); + } + if let Some(model) = query.model.as_ref() { + builder.push(" AND model = ").push_bind(model.clone()); + } + builder.push(format!(" GROUP BY {group_key} ORDER BY {group_key}")); + return fetch_usage_leaderboard_query(builder.build(), &self.pool) + .await + .map(Some); + } let items = match query.group_by { UsageLeaderboardGroupBy::Model => { let mut builder = if let Some(user_id) = query.user_id.as_deref() { @@ -7451,11 +7493,12 @@ WHERE date >= .push_bind(end_day_utc) .push(" AND provider_name = ") .push_bind(provider_name.to_string()); - if let Some(user_id) = query.user_id.as_deref() { - builder - .push(" AND user_id = ") - .push_bind(user_id.to_string()); - } + push_usage_user_scope( + &mut builder, + "user_id", + query.user_id.as_deref(), + query.user_ids.as_deref(), + ); builder.push(" GROUP BY user_id ORDER BY user_id ASC"); builder } else if let Some(model) = query.model.as_deref() { @@ -7477,11 +7520,12 @@ WHERE date >= .push_bind(end_day_utc) .push(" AND model = ") .push_bind(model.to_string()); - if let Some(user_id) = query.user_id.as_deref() { - builder - .push(" AND user_id = ") - .push_bind(user_id.to_string()); - } + push_usage_user_scope( + &mut builder, + "user_id", + query.user_id.as_deref(), + query.user_ids.as_deref(), + ); builder.push(" GROUP BY user_id ORDER BY user_id ASC"); builder } else { @@ -7505,11 +7549,12 @@ WHERE date >= .push(" AND date < ") .push_bind(end_day_utc) .push(" AND user_id IS NOT NULL"); - if let Some(user_id) = query.user_id.as_deref() { - builder - .push(" AND user_id = ") - .push_bind(user_id.to_string()); - } + push_usage_user_scope( + &mut builder, + "user_id", + query.user_id.as_deref(), + query.user_ids.as_deref(), + ); builder.push(" GROUP BY user_id ORDER BY user_id ASC"); builder }; @@ -7594,6 +7639,7 @@ WHERE stats_daily_api_key.date >= absorb_usage_leaderboard_rows( &mut grouped, self.summarize_usage_leaderboard_raw(&UsageLeaderboardQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(raw_start), created_until_unix_secs: dashboard_utc_to_unix_secs(raw_end), ..query.clone() @@ -7616,6 +7662,7 @@ WHERE stats_daily_api_key.date >= absorb_usage_leaderboard_rows( &mut grouped, self.summarize_usage_leaderboard_raw(&UsageLeaderboardQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(aggregate_start), created_until_unix_secs: dashboard_utc_to_unix_secs(aggregate_end), ..query.clone() @@ -7629,6 +7676,7 @@ WHERE stats_daily_api_key.date >= absorb_usage_leaderboard_rows( &mut grouped, self.summarize_usage_leaderboard_raw(&UsageLeaderboardQuery { + provider_names: query.provider_names.clone(), created_from_unix_secs: dashboard_utc_to_unix_secs(raw_start), created_until_unix_secs: dashboard_utc_to_unix_secs(raw_end), ..query.clone() diff --git a/crates/aether-data/adapters/postgres/src/usage/queries/find_by_id_sql.sql b/crates/aether-data/adapters/postgres/src/usage/queries/find_by_id_sql.sql index a5252794e..298d4909e 100644 --- a/crates/aether-data/adapters/postgres/src/usage/queries/find_by_id_sql.sql +++ b/crates/aether-data/adapters/postgres/src/usage/queries/find_by_id_sql.sql @@ -185,6 +185,10 @@ SELECT usage_http_audits.provider_request_body_ref AS http_provider_request_body_ref, usage_http_audits.response_body_ref AS http_response_body_ref, usage_http_audits.client_response_body_ref AS http_client_response_body_ref, + usage_http_audits.request_body_state AS http_request_body_state, + usage_http_audits.provider_request_body_state AS http_provider_request_body_state, + usage_http_audits.response_body_state AS http_response_body_state, + usage_http_audits.client_response_body_state AS http_client_response_body_state, usage_routing_snapshots.candidate_id AS routing_candidate_id, usage_routing_snapshots.candidate_index AS routing_candidate_index, usage_routing_snapshots.key_name AS routing_key_name, diff --git a/crates/aether-data/adapters/postgres/src/usage/queries/list_recent_usage_audits_prefix.sql b/crates/aether-data/adapters/postgres/src/usage/queries/list_recent_usage_audits_prefix.sql index d0909bcd2..d6d8a736d 100644 --- a/crates/aether-data/adapters/postgres/src/usage/queries/list_recent_usage_audits_prefix.sql +++ b/crates/aether-data/adapters/postgres/src/usage/queries/list_recent_usage_audits_prefix.sql @@ -183,6 +183,7 @@ SELECT OR NULLIF(BTRIM("usage".request_metadata->>'provider_reasoning_effort'), '') IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'provider_service_tier'), '') IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'provider_actual_service_tier'), '') IS NOT NULL + OR NULLIF(BTRIM("usage".request_metadata->>'provider_response_model'), '') IS NOT NULL OR ("usage".request_metadata->>'client_requested_stream') IN ('true', 'false') OR ("usage".request_metadata->>'upstream_is_stream') IN ('true', 'false') OR ("usage".request_metadata->>'websocket_mode') IN ('true', 'false') @@ -208,6 +209,8 @@ SELECT NULLIF(BTRIM("usage".request_metadata->>'provider_service_tier'), ''), 'provider_actual_service_tier', NULLIF(BTRIM("usage".request_metadata->>'provider_actual_service_tier'), ''), + 'provider_response_model', + NULLIF(BTRIM("usage".request_metadata->>'provider_response_model'), ''), 'client_requested_stream', CASE WHEN ("usage".request_metadata->>'client_requested_stream') IN ('true', 'false') diff --git a/crates/aether-data/adapters/postgres/src/usage/queries/list_usage_audits_prefix.sql b/crates/aether-data/adapters/postgres/src/usage/queries/list_usage_audits_prefix.sql index d0909bcd2..d6d8a736d 100644 --- a/crates/aether-data/adapters/postgres/src/usage/queries/list_usage_audits_prefix.sql +++ b/crates/aether-data/adapters/postgres/src/usage/queries/list_usage_audits_prefix.sql @@ -183,6 +183,7 @@ SELECT OR NULLIF(BTRIM("usage".request_metadata->>'provider_reasoning_effort'), '') IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'provider_service_tier'), '') IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'provider_actual_service_tier'), '') IS NOT NULL + OR NULLIF(BTRIM("usage".request_metadata->>'provider_response_model'), '') IS NOT NULL OR ("usage".request_metadata->>'client_requested_stream') IN ('true', 'false') OR ("usage".request_metadata->>'upstream_is_stream') IN ('true', 'false') OR ("usage".request_metadata->>'websocket_mode') IN ('true', 'false') @@ -208,6 +209,8 @@ SELECT NULLIF(BTRIM("usage".request_metadata->>'provider_service_tier'), ''), 'provider_actual_service_tier', NULLIF(BTRIM("usage".request_metadata->>'provider_actual_service_tier'), ''), + 'provider_response_model', + NULLIF(BTRIM("usage".request_metadata->>'provider_response_model'), ''), 'client_requested_stream', CASE WHEN ("usage".request_metadata->>'client_requested_stream') IN ('true', 'false') diff --git a/crates/aether-data/adapters/postgres/src/usage/tests.rs b/crates/aether-data/adapters/postgres/src/usage/tests.rs index 082765ae7..33a1b01e3 100644 --- a/crates/aether-data/adapters/postgres/src/usage/tests.rs +++ b/crates/aether-data/adapters/postgres/src/usage/tests.rs @@ -3251,13 +3251,15 @@ fn usage_sql_canonical_openai_cache_case_preserves_effective_and_total_tokens() aggregate_audit_summary .matches("WHEN effective_input_tokens = 0 AND total_input_context = 0") .count(), - 2 + 1, + "the shared daily aggregate query should define the legacy token fallback once" ); assert_eq!( aggregate_audit_summary .matches("+ output_tokens + cache_creation_tokens + cache_read_tokens") .count(), - 2 + 1, + "the shared daily aggregate query should define canonical total tokens once" ); assert!(!aggregate_audit_summary.contains("SUM(input_tokens + output_tokens)")); @@ -3467,6 +3469,18 @@ fn usage_sql_reads_http_audits_for_single_record_fetches() { assert!(super::FIND_BY_ID_SQL.contains("LEFT JOIN usage_http_audits")); assert!(super::FIND_BY_REQUEST_ID_SQL.contains("http_request_body_ref")); assert!(super::FIND_BY_ID_SQL.contains("http_client_response_body_ref")); + for sql in [super::FIND_BY_REQUEST_ID_SQL, super::FIND_BY_ID_SQL] { + for field in [ + "request_body", + "provider_request_body", + "response_body", + "client_response_body", + ] { + assert!(sql.contains(&format!( + "usage_http_audits.{field}_state AS http_{field}_state" + ))); + } + } } #[test] @@ -3597,6 +3611,8 @@ fn usage_sql_uses_json_null_placeholders_for_usage_payload_columns() { assert!(sql.contains("request_metadata->>'provider_reasoning_effort'")); assert!(sql.contains("request_metadata->>'provider_service_tier'")); assert!(sql.contains("request_metadata->>'provider_actual_service_tier'")); + assert!(sql.contains("request_metadata->>'provider_response_model'")); + assert!(sql.contains("'provider_response_model'")); assert!(sql.contains("request_metadata->>'websocket_mode'")); assert!(sql.contains("'websocket_mode'")); assert!(sql.contains("AS client_family")); diff --git a/crates/aether-data/adapters/postgres/src/wallet.rs b/crates/aether-data/adapters/postgres/src/wallet.rs index 0c9538690..72f86783c 100644 --- a/crates/aether-data/adapters/postgres/src/wallet.rs +++ b/crates/aether-data/adapters/postgres/src/wallet.rs @@ -24,8 +24,9 @@ use aether_data_contracts::repository::wallet::{ wallet_recharge_checkout_failed_response, wallet_recharge_checkout_uncertain_response, wallet_recharge_order_is_reclaimable_placeholder, wallet_recharge_replay_matches, wallet_recharge_response_is_checkout_placeholder, wallet_refund_proof_is_success, - AdjustWalletBalanceInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, - AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery, + AdjustWalletBalanceInBatchInput, AdjustWalletBalanceInput, AdminPaymentOrderListQuery, + AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminUserWalletBalanceBatchContext, + AdminUserWalletBalanceBatchUserOutcome, AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput, CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, @@ -34,21 +35,23 @@ use aether_data_contracts::repository::wallet::{ CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, - FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput, - ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, - RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, - StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, - StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, - StoredAdminRedeemCodePage, StoredAdminWalletLedgerItem, StoredAdminWalletLedgerPage, - StoredAdminWalletListItem, StoredAdminWalletListPage, StoredAdminWalletRefund, - StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem, - StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, - StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, + FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, + PrepareAdminUserWalletBalanceBatchInput, PrepareAdminUserWalletBalanceBatchOutcome, + ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, + StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, + StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, + StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminUserWalletBalanceBatch, + StoredAdminWalletLedgerItem, StoredAdminWalletLedgerPage, StoredAdminWalletListItem, + StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage, + StoredAdminWalletRefundRequestItem, StoredAdminWalletRefundRequestPage, + StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletWriteRepository, }; use aether_data_contracts::DataLayerError; +use std::collections::BTreeMap; use crate::{ error::{postgres_error, SqlxResultExt}, @@ -809,6 +812,44 @@ impl SqlxWalletRepository { } } +fn effective_wallet_adjustment_amount(input: &AdjustWalletBalanceInput, before_total: f64) -> f64 { + if input.clamp_deduction_to_available_balance && input.amount_usd < 0.0 { + if before_total < 0.0 { + // Clear legacy negative totals to zero and record the actual ledger delta. + -before_total + } else { + -(-input.amount_usd).min(before_total) + } + } else { + input.amount_usd + } +} + +async fn persist_admin_wallet_batch_outcomes( + connection: &mut sqlx::PgConnection, + context: &AdminUserWalletBalanceBatchContext, + outcomes: &BTreeMap, +) -> Result<(), DataLayerError> { + let outcomes = serde_json::to_value(outcomes).map_err(|error| { + DataLayerError::UnexpectedValue(format!("admin wallet batch outcomes are invalid: {error}")) + })?; + sqlx::query( + r#" +UPDATE admin_user_wallet_balance_batches +SET user_outcomes = $3, updated_at_unix_secs = $4 +WHERE admin_user_id = $1 AND idempotency_key = $2 + "#, + ) + .bind(&context.admin_user_id) + .bind(&context.idempotency_key) + .bind(outcomes) + .bind(Utc::now().timestamp().max(0)) + .execute(connection) + .await + .map_postgres_err()?; + Ok(()) +} + #[async_trait] impl WalletReadRepository for SqlxWalletRepository { async fn find( @@ -4266,7 +4307,8 @@ RETURNING async fn adjust_wallet_balance( &self, input: AdjustWalletBalanceInput, - ) -> Result, DataLayerError> { + ) -> Result)>, DataLayerError> + { if !input.amount_usd.is_finite() || input.amount_usd == 0.0 { return Err(DataLayerError::InvalidInput( "adjustment amount must be finite and non-zero".to_string(), @@ -4275,6 +4317,58 @@ RETURNING self.tx_runner .run_read_write(|tx| { Box::pin(async move { + let batch_context = input.batch_context.clone(); + let mut batch_user_outcomes = if let Some(context) = &batch_context { + let batch_row = sqlx::query( + r#" +SELECT target_user_ids, user_outcomes +FROM admin_user_wallet_balance_batches +WHERE admin_user_id = $1 AND idempotency_key = $2 +FOR UPDATE + "#, + ) + .bind(&context.admin_user_id) + .bind(&context.idempotency_key) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput( + "admin wallet batch was not prepared".to_string(), + ) + })?; + let target_user_ids: Vec = + serde_json::from_value(row_get(&batch_row, "target_user_ids")?) + .map_err(|error| { + DataLayerError::UnexpectedValue(format!( + "admin wallet batch target list is invalid: {error}" + )) + })?; + if !target_user_ids.iter().any(|id| id == &context.user_id) { + return Err(DataLayerError::InvalidInput( + "user is outside the prepared admin wallet batch".to_string(), + )); + } + let outcomes: BTreeMap = + serde_json::from_value(row_get(&batch_row, "user_outcomes")?).map_err( + |error| { + DataLayerError::UnexpectedValue(format!( + "admin wallet batch outcomes are invalid: {error}" + )) + }, + )?; + if let Some(outcome) = outcomes.get(&context.user_id) { + match outcome { + AdminUserWalletBalanceBatchUserOutcome::Succeeded => {} + AdminUserWalletBalanceBatchUserOutcome::Failed(_) => { + return Ok(None); + } + } + } + Some(outcomes) + } else { + None + }; let Some(row) = sqlx::query( r#" SELECT @@ -4289,13 +4383,20 @@ SELECT CAST(total_recharged AS DOUBLE PRECISION) AS total_recharged, CAST(total_consumed AS DOUBLE PRECISION) AS total_consumed, CAST(total_refunded AS DOUBLE PRECISION) AS total_refunded, - CAST(total_adjusted AS DOUBLE PRECISION) AS total_adjusted + CAST(total_adjusted AS DOUBLE PRECISION) AS total_adjusted, + CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs FROM wallets WHERE id = $1 + AND ($2::character varying IS NULL OR user_id = $2::character varying) FOR UPDATE "#, ) .bind(&input.wallet_id) + .bind( + batch_context + .as_ref() + .map(|context| context.user_id.as_str()), + ) .fetch_optional(&mut **tx) .await .map_postgres_err()? @@ -4316,17 +4417,40 @@ FOR UPDATE "wallet balance is invalid".to_string(), )); } + let wallet = map_wallet_row(&row)?; + let already_applied = batch_context.as_ref().is_some_and(|context| { + batch_user_outcomes.as_ref().is_some_and(|outcomes| { + outcomes.get(&context.user_id) + == Some(&AdminUserWalletBalanceBatchUserOutcome::Succeeded) + }) + }); + if already_applied { + return Ok(Some((wallet, None))); + } + let amount_usd = effective_wallet_adjustment_amount(&input, before_total); + if amount_usd == 0.0 { + if let (Some(context), Some(outcomes)) = + (batch_context.as_ref(), batch_user_outcomes.as_mut()) + { + outcomes.insert( + context.user_id.clone(), + AdminUserWalletBalanceBatchUserOutcome::Succeeded, + ); + persist_admin_wallet_batch_outcomes(tx, context, outcomes).await?; + } + return Ok(Some((wallet, None))); + } let mut after_recharge = before_recharge; let mut after_gift = before_gift; - if input.amount_usd > 0.0 { + if amount_usd > 0.0 { if input.balance_type.eq_ignore_ascii_case("gift") { - after_gift += input.amount_usd; + after_gift += amount_usd; } else { - after_recharge += input.amount_usd; + after_recharge += amount_usd; } } else { - let mut remaining = -input.amount_usd; + let mut remaining = -amount_usd; let consume_positive_bucket = |balance: &mut f64, to_consume: &mut f64| { if *to_consume <= 0.0 { return; @@ -4348,7 +4472,7 @@ FOR UPDATE } } let after_total = after_recharge + after_gift; - let after_total_adjusted = before_total_adjusted + input.amount_usd; + let after_total_adjusted = before_total_adjusted + amount_usd; if !after_recharge.is_finite() || !after_gift.is_finite() || !after_total.is_finite() @@ -4387,7 +4511,7 @@ RETURNING .bind(&input.wallet_id) .bind(after_recharge) .bind(after_gift) - .bind(input.amount_usd) + .bind(amount_usd) .fetch_one(&mut **tx) .await .map_postgres_err()?; @@ -4443,7 +4567,7 @@ VALUES ( ) .bind(&transaction_id) .bind(&input.wallet_id) - .bind(input.amount_usd) + .bind(amount_usd) .bind(before_total) .bind(after_total) .bind(before_recharge) @@ -4457,14 +4581,24 @@ VALUES ( .await .map_postgres_err()?; + if let (Some(context), Some(outcomes)) = + (batch_context.as_ref(), batch_user_outcomes.as_mut()) + { + outcomes.insert( + context.user_id.clone(), + AdminUserWalletBalanceBatchUserOutcome::Succeeded, + ); + persist_admin_wallet_batch_outcomes(tx, context, outcomes).await?; + } + Ok(Some(( wallet, - StoredAdminWalletTransaction { + Some(StoredAdminWalletTransaction { id: transaction_id, wallet_id: input.wallet_id, category: "adjust".to_string(), reason_code: "adjust_admin".to_string(), - amount: input.amount_usd, + amount: amount_usd, balance_before: before_total, balance_after: after_total, recharge_balance_before: before_recharge, @@ -4478,13 +4612,244 @@ VALUES ( operator_email: None, description: Some(description), created_at_unix_ms: Some(created_at), - }, + }), ))) }) }) .await } + async fn prepare_admin_user_wallet_balance_batch( + &self, + input: PrepareAdminUserWalletBalanceBatchInput, + ) -> Result { + let target_user_ids = serde_json::to_value(&input.target_user_ids).map_err(|error| { + DataLayerError::InvalidInput(format!("invalid batch target users: {error}")) + })?; + let missing_user_ids = serde_json::to_value(&input.missing_user_ids).map_err(|error| { + DataLayerError::InvalidInput(format!("invalid missing batch users: {error}")) + })?; + let warnings = serde_json::to_value(&input.warnings).map_err(|error| { + DataLayerError::InvalidInput(format!("invalid batch warnings: {error}")) + })?; + let now = Utc::now().timestamp().max(0); + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + sqlx::query( + r#" +INSERT INTO admin_user_wallet_balance_batches ( + admin_user_id, idempotency_key, request_fingerprint, target_user_ids, + missing_user_ids, warnings, user_outcomes, created_at_unix_secs, updated_at_unix_secs +) +VALUES ($1, $2, $3, $4, $5, $6, '{}'::jsonb, $7, $7) +ON CONFLICT (admin_user_id, idempotency_key) DO NOTHING + "#, + ) + .bind(&input.admin_user_id) + .bind(&input.idempotency_key) + .bind(&input.request_fingerprint) + .bind(target_user_ids) + .bind(missing_user_ids) + .bind(warnings) + .bind(now) + .execute(&mut **tx) + .await + .map_postgres_err()?; + + let row = sqlx::query( + r#" +SELECT admin_user_id, idempotency_key, request_fingerprint, target_user_ids, + missing_user_ids, warnings, user_outcomes +FROM admin_user_wallet_balance_batches +WHERE admin_user_id = $1 AND idempotency_key = $2 +FOR UPDATE + "#, + ) + .bind(&input.admin_user_id) + .bind(&input.idempotency_key) + .fetch_one(&mut **tx) + .await + .map_postgres_err()?; + let request_fingerprint: String = row_get(&row, "request_fingerprint")?; + if request_fingerprint != input.request_fingerprint { + return Ok(PrepareAdminUserWalletBalanceBatchOutcome::Conflict); + } + + let read_json = |column: &str| -> Result { + row_get(&row, column) + }; + let batch = StoredAdminUserWalletBalanceBatch { + admin_user_id: row_get(&row, "admin_user_id")?, + idempotency_key: row_get(&row, "idempotency_key")?, + request_fingerprint, + target_user_ids: serde_json::from_value(read_json("target_user_ids")?) + .map_err(|error| DataLayerError::UnexpectedValue(error.to_string()))?, + missing_user_ids: serde_json::from_value(read_json("missing_user_ids")?) + .map_err(|error| DataLayerError::UnexpectedValue(error.to_string()))?, + warnings: serde_json::from_value(read_json("warnings")?) + .map_err(|error| DataLayerError::UnexpectedValue(error.to_string()))?, + user_outcomes: serde_json::from_value(read_json("user_outcomes")?) + .map_err(|error| DataLayerError::UnexpectedValue(error.to_string()))?, + }; + Ok(PrepareAdminUserWalletBalanceBatchOutcome::Ready(batch)) + }) + }) + .await + } + + async fn get_admin_user_wallet_balance_batch( + &self, + admin_user_id: &str, + idempotency_key: &str, + expected_fingerprint: &str, + ) -> Result, DataLayerError> { + let row = sqlx::query( + r#" +SELECT admin_user_id, idempotency_key, request_fingerprint, target_user_ids, + missing_user_ids, warnings, user_outcomes +FROM admin_user_wallet_balance_batches +WHERE admin_user_id = $1 AND idempotency_key = $2 + "#, + ) + .bind(admin_user_id) + .bind(idempotency_key) + .fetch_optional(&self.pool) + .await + .map_postgres_err()?; + let Some(row) = row else { + return Ok(None); + }; + let request_fingerprint: String = row_get(&row, "request_fingerprint")?; + if request_fingerprint != expected_fingerprint { + return Ok(Some(PrepareAdminUserWalletBalanceBatchOutcome::Conflict)); + } + let read_json = + |column: &str| -> Result { row_get(&row, column) }; + let batch = StoredAdminUserWalletBalanceBatch { + admin_user_id: row_get(&row, "admin_user_id")?, + idempotency_key: row_get(&row, "idempotency_key")?, + request_fingerprint, + target_user_ids: serde_json::from_value(read_json("target_user_ids")?) + .map_err(|error| DataLayerError::UnexpectedValue(error.to_string()))?, + missing_user_ids: serde_json::from_value(read_json("missing_user_ids")?) + .map_err(|error| DataLayerError::UnexpectedValue(error.to_string()))?, + warnings: serde_json::from_value(read_json("warnings")?) + .map_err(|error| DataLayerError::UnexpectedValue(error.to_string()))?, + user_outcomes: serde_json::from_value(read_json("user_outcomes")?) + .map_err(|error| DataLayerError::UnexpectedValue(error.to_string()))?, + }; + Ok(Some(PrepareAdminUserWalletBalanceBatchOutcome::Ready( + batch, + ))) + } + + async fn adjust_admin_user_wallet_balance_batch_user( + &self, + input: AdjustWalletBalanceInBatchInput, + ) -> Result { + let context = AdminUserWalletBalanceBatchContext { + admin_user_id: input.admin_user_id.clone(), + idempotency_key: input.idempotency_key.clone(), + user_id: input.user_id.clone(), + }; + let mut adjustment = input.adjustment; + adjustment.batch_context = Some(context); + if self.adjust_wallet_balance(adjustment).await?.is_some() { + return Ok(AdminUserWalletBalanceBatchUserOutcome::Succeeded); + } + self.record_admin_user_wallet_balance_batch_failure( + &input.admin_user_id, + &input.idempotency_key, + &input.user_id, + "用户钱包不可用", + ) + .await + } + + async fn record_admin_user_wallet_balance_batch_failure( + &self, + admin_user_id: &str, + idempotency_key: &str, + user_id: &str, + reason: &str, + ) -> Result { + let admin_user_id = admin_user_id.to_string(); + let idempotency_key = idempotency_key.to_string(); + let user_id = user_id.to_string(); + let reason = reason.to_string(); + self.tx_runner + .run_read_write(|tx| { + Box::pin(async move { + let row = sqlx::query( + r#" +SELECT target_user_ids, user_outcomes +FROM admin_user_wallet_balance_batches +WHERE admin_user_id = $1 AND idempotency_key = $2 +FOR UPDATE + "#, + ) + .bind(&admin_user_id) + .bind(&idempotency_key) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + .ok_or_else(|| { + DataLayerError::InvalidInput( + "admin wallet batch was not prepared".to_string(), + ) + })?; + let target_user_ids: Vec = + serde_json::from_value(row_get(&row, "target_user_ids")?).map_err( + |error| { + DataLayerError::UnexpectedValue(format!( + "admin wallet batch target list is invalid: {error}" + )) + }, + )?; + if !target_user_ids.iter().any(|id| id == &user_id) { + return Err(DataLayerError::InvalidInput( + "user is outside the prepared admin wallet batch".to_string(), + )); + } + let mut outcomes: BTreeMap = + serde_json::from_value(row_get(&row, "user_outcomes")?).map_err( + |error| { + DataLayerError::UnexpectedValue(format!( + "admin wallet batch outcomes are invalid: {error}" + )) + }, + )?; + if let Some(outcome) = outcomes.get(&user_id) { + return Ok(outcome.clone()); + } + let outcome = AdminUserWalletBalanceBatchUserOutcome::Failed(reason); + outcomes.insert(user_id, outcome.clone()); + let value = serde_json::to_value(&outcomes).map_err(|error| { + DataLayerError::UnexpectedValue(format!( + "admin wallet batch outcomes are invalid: {error}" + )) + })?; + sqlx::query( + r#" +UPDATE admin_user_wallet_balance_batches +SET user_outcomes = $3, updated_at_unix_secs = $4 +WHERE admin_user_id = $1 AND idempotency_key = $2 + "#, + ) + .bind(&admin_user_id) + .bind(&idempotency_key) + .bind(value) + .bind(Utc::now().timestamp().max(0)) + .execute(&mut **tx) + .await + .map_postgres_err()?; + Ok(outcome) + }) + }) + .await + } + async fn create_manual_wallet_recharge( &self, mut input: CreateManualWalletRechargeInput, @@ -8493,13 +8858,16 @@ VALUES ($1, $2, 'gift', 'gift_initial', $3, 0, $3, 0, 0, 0, $3, 'system_task', $ #[cfg(test)] mod tests { use aether_data_contracts::repository::wallet::{ - CreateManualWalletRechargeInput, CreditAdminPaymentOrderInput, ProcessPaymentCallbackInput, + AdjustWalletBalanceInBatchInput, AdjustWalletBalanceInput, + AdminUserWalletBalanceBatchUserOutcome, CreateManualWalletRechargeInput, + CreditAdminPaymentOrderInput, PrepareAdminUserWalletBalanceBatchInput, + PrepareAdminUserWalletBalanceBatchOutcome, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, RedeemWalletCodeInput, RedeemWalletCodeOutcome, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletWriteRepository, }; use sqlx::Row; - use super::SqlxWalletRepository; + use super::{effective_wallet_adjustment_amount, SqlxWalletRepository}; use crate::{PostgresPoolConfig, PostgresPoolFactory}; #[test] @@ -8579,6 +8947,7 @@ mod tests { "payment_orders", "payment_callbacks", "wallet_transactions", + "admin_user_wallet_balance_batches", "user_plan_entitlements", "redeem_code_batches", "redeem_codes", @@ -9225,6 +9594,264 @@ mod tests { pool.close().await; } + #[test] + fn bulk_adjustment_clamp_is_opt_in_and_uses_available_total() { + let input = AdjustWalletBalanceInput { + wallet_id: "wallet-1".to_string(), + amount_usd: -100.0, + balance_type: "recharge".to_string(), + operator_id: None, + description: None, + clamp_deduction_to_available_balance: true, + batch_context: None, + }; + assert_eq!(effective_wallet_adjustment_amount(&input, 13.0), -13.0); + assert_eq!(effective_wallet_adjustment_amount(&input, 0.0), -0.0); + assert_eq!(effective_wallet_adjustment_amount(&input, -1.0), 1.0); + + let legacy_input = AdjustWalletBalanceInput { + clamp_deduction_to_available_balance: false, + ..input + }; + assert_eq!( + effective_wallet_adjustment_amount(&legacy_input, 13.0), + -100.0 + ); + } + + #[tokio::test] + #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL bootstrap schema"] + async fn live_bulk_wallet_adjustment_persists_actual_delta_and_skips_zero_ledger() { + let pool = isolated_wallet_test_pool().await; + let (wallet_id, _) = seed_wallet(&pool).await; + let repository = SqlxWalletRepository::new(pool.clone()); + + let (wallet, transaction) = repository + .adjust_wallet_balance(AdjustWalletBalanceInput { + wallet_id: wallet_id.clone(), + amount_usd: -100.0, + balance_type: "recharge".to_string(), + operator_id: Some("admin-user".to_string()), + description: Some("bulk deduction".to_string()), + clamp_deduction_to_available_balance: true, + batch_context: None, + }) + .await + .expect("bulk adjustment should succeed") + .expect("wallet should exist"); + let transaction = + transaction.expect("positive available balance should create a ledger row"); + assert_eq!(transaction.amount, -13.0); + assert_eq!(transaction.balance_before, 13.0); + assert_eq!(transaction.balance_after, 0.0); + assert_eq!(wallet.balance + wallet.gift_balance, 0.0); + let persisted_amount: f64 = sqlx::query_scalar( + "SELECT amount::double precision FROM wallet_transactions WHERE id = $1", + ) + .bind(&transaction.id) + .fetch_one(&pool) + .await + .expect("ledger should store the effective deduction"); + assert_eq!(persisted_amount, -13.0); + + let (wallet, transaction) = repository + .adjust_wallet_balance(AdjustWalletBalanceInput { + wallet_id: wallet_id.clone(), + amount_usd: -1.0, + balance_type: "recharge".to_string(), + operator_id: Some("admin-user".to_string()), + description: Some("bulk deduction at zero".to_string()), + clamp_deduction_to_available_balance: true, + batch_context: None, + }) + .await + .expect("zero-balance adjustment should succeed") + .expect("wallet should still exist"); + assert_eq!(wallet.balance + wallet.gift_balance, 0.0); + assert!( + transaction.is_none(), + "zero effective delta must not create a ledger row" + ); + let transaction_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = $1") + .bind(&wallet_id) + .fetch_one(&pool) + .await + .expect("ledger row count should be readable"); + assert_eq!(transaction_count, 1); + + sqlx::query("UPDATE wallets SET balance = -2, gift_balance = 1 WHERE id = $1") + .bind(&wallet_id) + .execute(&pool) + .await + .expect("legacy negative wallet balance should be seeded"); + let (wallet, transaction) = repository + .adjust_wallet_balance(AdjustWalletBalanceInput { + wallet_id: wallet_id.clone(), + amount_usd: -1.0, + balance_type: "recharge".to_string(), + operator_id: Some("admin-user".to_string()), + description: Some("bulk deduction floors a negative balance".to_string()), + clamp_deduction_to_available_balance: true, + batch_context: None, + }) + .await + .expect("legacy negative wallet should be floored at zero") + .expect("wallet should still exist"); + assert_eq!(wallet.balance + wallet.gift_balance, 0.0); + let transaction = transaction.expect("negative balance correction should be ledgered"); + assert_eq!(transaction.amount, 1.0); + assert_eq!(transaction.balance_before, -1.0); + assert_eq!(transaction.balance_after, 0.0); + let persisted_correction_amount: f64 = sqlx::query_scalar( + "SELECT amount::double precision FROM wallet_transactions WHERE id = $1", + ) + .bind(&transaction.id) + .fetch_one(&pool) + .await + .expect("ledger should store the negative-balance correction"); + assert_eq!(persisted_correction_amount, 1.0); + let transaction_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = $1") + .bind(&wallet_id) + .fetch_one(&pool) + .await + .expect("ledger row count should be readable"); + assert_eq!(transaction_count, 2); + pool.close().await; + } + + #[tokio::test] + #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL bootstrap schema"] + async fn live_admin_wallet_balance_batch_replays_committed_and_zero_delta_results_once() { + let pool = isolated_wallet_test_pool().await; + let (wallet_id, user_id) = seed_wallet(&pool).await; + let repository = SqlxWalletRepository::new(pool.clone()); + let admin_user_id = "admin-user".to_string(); + let make_adjustment = + |idempotency_key: &str, amount_usd: f64| AdjustWalletBalanceInBatchInput { + admin_user_id: admin_user_id.clone(), + idempotency_key: idempotency_key.to_string(), + user_id: user_id.clone(), + adjustment: AdjustWalletBalanceInput { + wallet_id: wallet_id.clone(), + amount_usd, + balance_type: "recharge".to_string(), + operator_id: Some(admin_user_id.clone()), + description: Some("idempotency integration test".to_string()), + clamp_deduction_to_available_balance: true, + batch_context: None, + }, + }; + let prepare_batch = |idempotency_key: &str, request_fingerprint: &str| { + PrepareAdminUserWalletBalanceBatchInput { + admin_user_id: admin_user_id.clone(), + idempotency_key: idempotency_key.to_string(), + request_fingerprint: request_fingerprint.to_string(), + target_user_ids: vec![user_id.clone()], + missing_user_ids: Vec::new(), + warnings: Vec::new(), + } + }; + + assert!(matches!( + repository + .prepare_admin_user_wallet_balance_batch(prepare_batch( + "deduct-key", + "deduct-fingerprint" + )) + .await + .unwrap(), + PrepareAdminUserWalletBalanceBatchOutcome::Ready(_) + )); + let deduct = make_adjustment("deduct-key", -50.0); + assert_eq!( + repository + .adjust_admin_user_wallet_balance_batch_user(deduct.clone()) + .await + .unwrap(), + AdminUserWalletBalanceBatchUserOutcome::Succeeded + ); + assert_eq!( + repository + .adjust_admin_user_wallet_balance_batch_user(deduct.clone()) + .await + .unwrap(), + AdminUserWalletBalanceBatchUserOutcome::Succeeded + ); + let wallet = repository + .find(WalletLookupKey::WalletId(&wallet_id)) + .await + .unwrap() + .unwrap(); + assert_eq!(wallet.balance + wallet.gift_balance, 0.0); + + assert!(matches!( + repository + .prepare_admin_user_wallet_balance_batch(prepare_batch( + "zero-key", + "zero-fingerprint" + )) + .await + .unwrap(), + PrepareAdminUserWalletBalanceBatchOutcome::Ready(_) + )); + let zero_delta = make_adjustment("zero-key", -5.0); + assert_eq!( + repository + .adjust_admin_user_wallet_balance_batch_user(zero_delta.clone()) + .await + .unwrap(), + AdminUserWalletBalanceBatchUserOutcome::Succeeded + ); + repository + .adjust_wallet_balance(AdjustWalletBalanceInput { + wallet_id: wallet_id.clone(), + amount_usd: 8.0, + balance_type: "recharge".to_string(), + operator_id: Some(admin_user_id.clone()), + description: Some("recharge after zero-delta batch".to_string()), + clamp_deduction_to_available_balance: true, + batch_context: None, + }) + .await + .unwrap() + .unwrap(); + for replay in [zero_delta, deduct] { + assert_eq!( + repository + .adjust_admin_user_wallet_balance_batch_user(replay) + .await + .unwrap(), + AdminUserWalletBalanceBatchUserOutcome::Succeeded + ); + } + assert_eq!( + repository + .prepare_admin_user_wallet_balance_batch(prepare_batch( + "zero-key", + "different-fingerprint" + )) + .await + .unwrap(), + PrepareAdminUserWalletBalanceBatchOutcome::Conflict + ); + let wallet = repository + .find(WalletLookupKey::WalletId(&wallet_id)) + .await + .unwrap() + .unwrap(); + assert_eq!(wallet.balance + wallet.gift_balance, 8.0); + let transaction_count: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = $1") + .bind(&wallet_id) + .fetch_one(&pool) + .await + .unwrap(); + assert_eq!(transaction_count, 2); + pool.close().await; + } + #[tokio::test] async fn repository_constructs_from_lazy_pool() { let factory = PostgresPoolFactory::new(PostgresPoolConfig { diff --git a/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs b/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs index 2e98472ec..f69b16e14 100644 --- a/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs +++ b/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs @@ -8,9 +8,10 @@ use serde_json::{Map, Value}; use crate::repository::candidates::sanitize_request_candidate_skip_reason; use super::{ - LIVE_SESSION_METADATA_KEY, PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY, - PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, - PROVIDER_REASONING_EFFORT_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, + normalize_provider_response_model, LIVE_SESSION_METADATA_KEY, + PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY, PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, + PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY, + PROVIDER_RESPONSE_MODEL_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, REALTIME_SESSION_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY, @@ -128,6 +129,12 @@ pub fn sanitize_usage_request_metadata_object(source: &Map) -> Op ] { insert_known_string(source, &mut target, key, sanitize_service_tier); } + insert_known_string( + source, + &mut target, + PROVIDER_RESPONSE_MODEL_METADATA_KEY, + normalize_provider_response_model, + ); insert_bounded_u64( source, &mut target, @@ -1334,6 +1341,23 @@ mod tests { } } + #[test] + fn persistence_projection_keeps_bounded_response_model_only_as_a_string() { + let metadata = sanitize_usage_request_metadata(Some(json!({ + "provider_response_model": " GPT-5.1 " + }))) + .expect("response model should remain"); + assert_eq!(metadata["provider_response_model"], "GPT-5.1"); + assert!(sanitize_usage_request_metadata(Some(json!({ + "provider_response_model": 42 + }))) + .is_none()); + assert!(sanitize_usage_request_metadata(Some(json!({ + "provider_response_model": "x".repeat(257) + }))) + .is_none()); + } + #[test] fn persistence_projection_keeps_only_bounded_settlement_facts() { let metadata = sanitize_usage_request_metadata(Some(json!({ diff --git a/crates/aether-data/contracts/src/repository/usage/mod.rs b/crates/aether-data/contracts/src/repository/usage/mod.rs index 2cd9f0bfa..b34dcba94 100644 --- a/crates/aether-data/contracts/src/repository/usage/mod.rs +++ b/crates/aether-data/contracts/src/repository/usage/mod.rs @@ -25,11 +25,12 @@ pub use policy::*; pub use types::{ canonical_usage_body_ref_for, extract_provider_actual_service_tier_from_response, extract_provider_cache_ttl_minutes_from_metadata, extract_provider_reasoning_effort_from_body, - extract_provider_service_tier_from_body, normalize_provider_service_tier, parse_usage_body_ref, + extract_provider_response_model_from_bodies, extract_provider_service_tier_from_body, + normalize_provider_response_model, normalize_provider_service_tier, parse_usage_body_ref, resolve_provider_cache_ttl_minutes, resolve_provider_service_tier_from_request_capture, - usage_body_ref, usage_request_metadata_client_family, ApiKeyLastUsedDelta, - ManagementTokenCounterDelta, PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest, - ProxyNodeCounterDelta, StoredProviderApiKeyUsageSummary, + usage_body_capture_is_authoritative, usage_body_ref, usage_request_metadata_client_family, + ApiKeyLastUsedDelta, ManagementTokenCounterDelta, PendingUsageCleanupSummary, + ProviderApiKeyWindowUsageRequest, ProxyNodeCounterDelta, StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredProviderUsageWindow, StoredRequestUsageAudit, StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBodyPayload, StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary, @@ -56,9 +57,9 @@ pub use types::{ UsageTimeSeriesQuery, UsageWriteRepository, LIVE_SESSION_METADATA_KEY, PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY, PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY, - PROVIDER_SERVICE_TIER_METADATA_KEY, REALTIME_SESSION_METADATA_KEY, - REQUESTED_REASONING_EFFORT_METADATA_KEY, ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, - ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY, - USAGE_PRICING_AVAILABLE_METADATA_KEY, WEBSOCKET_MODE_METADATA_KEY, - WEBSOCKET_TRANSPORT_METADATA_KEY, + PROVIDER_RESPONSE_MODEL_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, + REALTIME_SESSION_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, + ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY, + USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY, + WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY, }; diff --git a/crates/aether-data/contracts/src/repository/usage/types.rs b/crates/aether-data/contracts/src/repository/usage/types.rs index ff5d07643..ec9da244b 100644 --- a/crates/aether-data/contracts/src/repository/usage/types.rs +++ b/crates/aether-data/contracts/src/repository/usage/types.rs @@ -1,3 +1,4 @@ +use aether_ai_formats::normalize_api_format_alias; use async_trait::async_trait; use chrono::{DateTime, Utc}; use serde_json::Value; @@ -6,6 +7,7 @@ pub const PROVIDER_REASONING_EFFORT_METADATA_KEY: &str = "provider_reasoning_eff pub const REQUESTED_REASONING_EFFORT_METADATA_KEY: &str = "requested_reasoning_effort"; pub const PROVIDER_SERVICE_TIER_METADATA_KEY: &str = "provider_service_tier"; pub const PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY: &str = "provider_actual_service_tier"; +pub const PROVIDER_RESPONSE_MODEL_METADATA_KEY: &str = "provider_response_model"; pub const PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY: &str = "provider_cache_ttl_minutes"; pub const ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY: &str = "routing_candidate_skip_reason"; pub const ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY: &str = "routing_failure_diagnostic"; @@ -119,6 +121,141 @@ pub fn normalize_provider_service_tier(value: &str) -> Option { Some(value.to_ascii_lowercase()) } +/// 清洗模型名称,保留大小写,只去除首尾空白。 +pub fn normalize_provider_response_model(value: &str) -> Option { + let value = value.trim(); + if value.is_empty() || value.len() > 256 { + return None; + } + Some(value.to_string()) +} + +fn extract_model_at_paths(value: &Value, paths: &[&[&str]]) -> Option { + paths.iter().find_map(|path| { + let value = path + .iter() + .try_fold(value, |current, key| current.as_object()?.get(*key))?; + value.as_str().and_then(normalize_provider_response_model) + }) +} + +fn response_model_paths(provider_api_format: Option<&str>) -> &'static [&'static [&'static str]] { + match normalize_api_format_alias(provider_api_format.unwrap_or_default()).as_str() { + "gemini:generate_content" => { + // Gemini 原生响应使用 modelVersion;部分网关会改写为 model。 + &[&["modelVersion"], &["model_version"], &["model"]] + } + "gemini:embedding" => { + // Gemini Embedding 可能返回 model、modelVersion 或 Vertex 的 deployedModelId。 + &[ + &["model"], + &["modelVersion"], + &["model_version"], + &["deployedModelId"], + &["deployed_model_id"], + ] + } + "gemini:interactions" => { + // Interactions 请求既可能叫 model,也可能叫 agent;响应优先读取 model。 + &[ + &["model"], + &["modelVersion"], + &["model_version"], + &["agent"], + ] + } + _ => &[&["model"]], + } +} + +fn extract_model_from_known_response_wrappers( + response_body: &Value, + paths: &[&[&str]], +) -> Option { + // 只展开协议中已知的 response/chunks 包装,避免在候选内容、工具参数等任意嵌套 + // 对象中搜索同名字段,误把 role="model" 一类内容当成响应模型。 + extract_model_at_paths(response_body, paths) + .or_else(|| { + response_body + .get("response") + .and_then(|response| extract_model_at_paths(response, paths)) + }) + .or_else(|| { + response_body + .get("chunks") + .and_then(Value::as_array) + .and_then(|chunks| { + chunks.iter().rev().find_map(|chunk| { + extract_model_at_paths(chunk, paths).or_else(|| { + chunk + .get("response") + .and_then(|response| extract_model_at_paths(response, paths)) + }) + }) + }) + }) + .or_else(|| { + response_body + .get("response") + .and_then(|response| response.get("chunks")) + .and_then(Value::as_array) + .and_then(|chunks| { + chunks.iter().rev().find_map(|chunk| { + extract_model_at_paths(chunk, paths).or_else(|| { + chunk + .get("response") + .and_then(|response| extract_model_at_paths(response, paths)) + }) + }) + }) + }) +} + +fn extract_provider_model_from_response_body( + response_body: &Value, + provider_api_format: Option<&str>, +) -> Option { + extract_model_from_known_response_wrappers( + response_body, + response_model_paths(provider_api_format), + ) +} + +fn extract_provider_model_from_request_body( + request_body: &Value, + request_api_format: Option<&str>, +) -> Option { + let paths: &[&[&str]] = + match normalize_api_format_alias(request_api_format.unwrap_or_default()).as_str() { + "gemini:interactions" => &[&["model"], &["agent"]], + _ => &[&["model"]], + }; + extract_model_at_paths(request_body, paths) +} + +/// 只有请求体和响应体都可作为完整事实时,才计算响应模型,避免用截断内容猜测。 +pub fn extract_provider_response_model_from_bodies( + request_body: Option<&Value>, + request_body_state: Option, + request_api_format: Option<&str>, + response_body: Option<&Value>, + response_body_state: Option, + provider_api_format: Option<&str>, +) -> Option { + if !usage_body_capture_is_authoritative(request_body, request_body_state) + || !usage_body_capture_is_authoritative(response_body, response_body_state) + { + return None; + } + + let request_model = + extract_provider_model_from_request_body(request_body?, request_api_format)?; + let response_model = + extract_provider_model_from_response_body(response_body?, provider_api_format)?; + + (request_model != response_model).then_some(response_model) +} + /// Resolves a provider processing tier exclusively from the final upstream request. /// /// A complete captured body is authoritative, including when it contains no tier. The metadata @@ -129,7 +266,7 @@ pub fn resolve_provider_service_tier_from_request_capture( provider_request_body_state: Option, request_metadata: Option<&Value>, ) -> Option { - if request_body_capture_is_authoritative(provider_request_body, provider_request_body_state) { + if usage_body_capture_is_authoritative(provider_request_body, provider_request_body_state) { return extract_provider_service_tier_from_body(provider_request_body); } @@ -153,7 +290,7 @@ pub fn resolve_provider_service_tier_from_request_capture( .and_then(normalize_provider_service_tier) } -fn request_body_capture_is_authoritative( +pub fn usage_body_capture_is_authoritative( request_body: Option<&Value>, request_body_state: Option, ) -> bool { @@ -184,7 +321,7 @@ fn resolve_reasoning_effort_from_request_capture( request_metadata: Option<&Value>, metadata_key: &str, ) -> Option { - if request_body_capture_is_authoritative(request_body, request_body_state) { + if usage_body_capture_is_authoritative(request_body, request_body_state) { return extract_provider_reasoning_effort_from_body(request_body); } @@ -701,6 +838,11 @@ impl StoredRequestUsageAudit { }) } + pub fn provider_response_model(&self) -> Option { + self.request_metadata_string(PROVIDER_RESPONSE_MODEL_METADATA_KEY) + .and_then(normalize_provider_response_model) + } + pub fn provider_cache_ttl_minutes(&self) -> Option { resolve_provider_cache_ttl_minutes( self.endpoint_api_format @@ -1096,9 +1238,15 @@ pub struct UsageAuditAggregationQuery { #[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)] pub struct UsageAuditSummaryQuery { + /// Optional provider-name allowlist, intersected with provider_name; empty matches nothing. + #[serde(default)] + pub provider_names: Option>, pub created_from_unix_secs: u64, pub created_until_unix_secs: u64, pub user_id: Option, + /// Optional bulk user scope used by current user-group reporting. + /// An empty list intentionally matches no usage rows. + pub user_ids: Option>, pub provider_name: Option, pub model: Option, } @@ -1486,11 +1634,17 @@ pub enum UsageTimeSeriesGranularity { #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct UsageTimeSeriesQuery { + /// Optional provider-name allowlist, intersected with provider_name; empty matches nothing. + #[serde(default)] + pub provider_names: Option>, pub created_from_unix_secs: u64, pub created_until_unix_secs: u64, pub granularity: UsageTimeSeriesGranularity, pub tz_offset_minutes: i32, pub user_id: Option, + /// Optional bulk user scope used by current user-group reporting. + /// An empty list intentionally matches no usage rows. + pub user_ids: Option>, pub provider_name: Option, pub model: Option, } @@ -1517,10 +1671,16 @@ pub enum UsageLeaderboardGroupBy { #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct UsageLeaderboardQuery { + /// Optional provider-name allowlist, intersected with provider_name; empty matches nothing. + #[serde(default)] + pub provider_names: Option>, pub created_from_unix_secs: u64, pub created_until_unix_secs: u64, pub group_by: UsageLeaderboardGroupBy, pub user_id: Option, + /// Optional bulk user scope used by current user-group reporting. + /// An empty list intentionally matches no usage rows. + pub user_ids: Option>, pub provider_name: Option, pub model: Option, } @@ -2709,7 +2869,8 @@ fn parse_timestamp(value: i64, field_name: &str) -> Result, pub description: Option, + #[serde(default)] + pub clamp_deduction_to_available_balance: bool, + #[serde(default)] + pub batch_context: Option, +} + +#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct AdminUserWalletBalanceBatchContext { + pub admin_user_id: String, + pub idempotency_key: String, + pub user_id: String, +} + +#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct PrepareAdminUserWalletBalanceBatchInput { + pub admin_user_id: String, + pub idempotency_key: String, + pub request_fingerprint: String, + pub target_user_ids: Vec, + pub missing_user_ids: Vec, + pub warnings: Vec, +} + +#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct StoredAdminUserWalletBalanceBatch { + pub admin_user_id: String, + pub idempotency_key: String, + pub request_fingerprint: String, + pub target_user_ids: Vec, + pub missing_user_ids: Vec, + pub warnings: Vec, + pub user_outcomes: std::collections::BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub enum PrepareAdminUserWalletBalanceBatchOutcome { + Ready(StoredAdminUserWalletBalanceBatch), + Conflict, +} + +#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[serde(tag = "status", content = "reason", rename_all = "snake_case")] +pub enum AdminUserWalletBalanceBatchUserOutcome { + Succeeded, + Failed(String), +} + +#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +pub struct AdjustWalletBalanceInBatchInput { + pub admin_user_id: String, + pub idempotency_key: String, + pub user_id: String, + pub adjustment: AdjustWalletBalanceInput, } #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] @@ -2894,7 +2947,51 @@ pub trait WalletWriteRepository: Send + Sync { async fn adjust_wallet_balance( &self, input: AdjustWalletBalanceInput, - ) -> Result, crate::DataLayerError>; + ) -> Result< + Option<(StoredWalletSnapshot, Option)>, + crate::DataLayerError, + >; + + async fn prepare_admin_user_wallet_balance_batch( + &self, + _input: PrepareAdminUserWalletBalanceBatchInput, + ) -> Result { + Err(crate::DataLayerError::InvalidInput( + "idempotent admin wallet batches are not available".to_string(), + )) + } + + async fn get_admin_user_wallet_balance_batch( + &self, + _admin_user_id: &str, + _idempotency_key: &str, + _request_fingerprint: &str, + ) -> Result, crate::DataLayerError> { + Err(crate::DataLayerError::InvalidInput( + "idempotent admin wallet batches are not available".to_string(), + )) + } + + async fn adjust_admin_user_wallet_balance_batch_user( + &self, + _input: AdjustWalletBalanceInBatchInput, + ) -> Result { + Err(crate::DataLayerError::InvalidInput( + "idempotent admin wallet batches are not available".to_string(), + )) + } + + async fn record_admin_user_wallet_balance_batch_failure( + &self, + _admin_user_id: &str, + _idempotency_key: &str, + _user_id: &str, + _reason: &str, + ) -> Result { + Err(crate::DataLayerError::InvalidInput( + "idempotent admin wallet batches are not available".to_string(), + )) + } async fn create_manual_wallet_recharge( &self, diff --git a/crates/aether-data/runtime/schema/generated/postgres/baseline/005_wallet_billing.sql b/crates/aether-data/runtime/schema/generated/postgres/baseline/005_wallet_billing.sql index 680556a5d..212d7e96f 100644 --- a/crates/aether-data/runtime/schema/generated/postgres/baseline/005_wallet_billing.sql +++ b/crates/aether-data/runtime/schema/generated/postgres/baseline/005_wallet_billing.sql @@ -344,3 +344,17 @@ CREATE INDEX IF NOT EXISTS idx_redeem_codes_status ON public.redeem_codes USING CREATE INDEX IF NOT EXISTS idx_redeem_codes_redeemed_user ON public.redeem_codes USING btree (redeemed_by_user_id, redeemed_at); CREATE INDEX IF NOT EXISTS idx_redeem_codes_redeemed_order ON public.redeem_codes USING btree (redeemed_payment_order_id); +CREATE TABLE IF NOT EXISTS public.admin_user_wallet_balance_batches ( + admin_user_id character varying(64) NOT NULL, + idempotency_key character varying(128) NOT NULL, + request_fingerprint character varying(64) NOT NULL, + target_user_ids jsonb NOT NULL, + missing_user_ids jsonb NOT NULL, + warnings jsonb NOT NULL, + user_outcomes jsonb NOT NULL, + created_at_unix_secs bigint NOT NULL, + updated_at_unix_secs bigint NOT NULL +); + +ALTER TABLE ONLY public.admin_user_wallet_balance_batches ADD CONSTRAINT admin_user_wallet_balance_batches_pkey PRIMARY KEY (admin_user_id, idempotency_key); + diff --git a/crates/aether-data/runtime/schema/logical/005_wallet_billing.toml b/crates/aether-data/runtime/schema/logical/005_wallet_billing.toml index 967af8f37..f9d578e6c 100644 --- a/crates/aether-data/runtime/schema/logical/005_wallet_billing.toml +++ b/crates/aether-data/runtime/schema/logical/005_wallet_billing.toml @@ -1370,3 +1370,47 @@ columns = ["redeemed_by_user_id", "redeemed_at"] [[table.redeem_codes.indexes]] name = "idx_redeem_codes_redeemed_order" columns = ["redeemed_payment_order_id"] + +[table.admin_user_wallet_balance_batches] +domain = "wallet_billing" +order = 110 +primary_key = ["admin_user_id", "idempotency_key"] + +[[table.admin_user_wallet_balance_batches.columns]] +name = "admin_user_id" +type = "text_id" +length = 64 + +[[table.admin_user_wallet_balance_batches.columns]] +name = "idempotency_key" +type = "text" +length = 128 + +[[table.admin_user_wallet_balance_batches.columns]] +name = "request_fingerprint" +type = "text" +length = 64 + +[[table.admin_user_wallet_balance_batches.columns]] +name = "target_user_ids" +type = "json" + +[[table.admin_user_wallet_balance_batches.columns]] +name = "missing_user_ids" +type = "json" + +[[table.admin_user_wallet_balance_batches.columns]] +name = "warnings" +type = "json" + +[[table.admin_user_wallet_balance_batches.columns]] +name = "user_outcomes" +type = "json" + +[[table.admin_user_wallet_balance_batches.columns]] +name = "created_at_unix_secs" +type = "unix_seconds" + +[[table.admin_user_wallet_balance_batches.columns]] +name = "updated_at_unix_secs" +type = "unix_seconds" diff --git a/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs b/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs index 682a0d6db..edcb15b6c 100644 --- a/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs +++ b/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs @@ -1592,6 +1592,7 @@ fn pending_migrations_from_applied_skips_versions_already_applied() { 20260921010000, 20260921020000, 20260921020100, + 20260923000000, 20261001000000, ] ); @@ -2563,12 +2564,14 @@ INSERT INTO public.stats_daily_api_key ( .expect("API key daily aggregate fixtures should be inserted"); let leaderboard_query = UsageLeaderboardQuery { + provider_names: None, created_from_unix_secs: u64::try_from(stats_day.timestamp()) .expect("historical stats day should be nonnegative"), created_until_unix_secs: u64::try_from((stats_day + chrono::Duration::days(1)).timestamp()) .expect("historical stats end should be nonnegative"), group_by: UsageLeaderboardGroupBy::ApiKey, user_id: Some("leaderboard-owner".to_string()), + user_ids: None, provider_name: None, model: None, }; diff --git a/crates/aether-data/runtime/src/lifecycle/migrate/tests/legacy_overview_upgrade.rs b/crates/aether-data/runtime/src/lifecycle/migrate/tests/legacy_overview_upgrade.rs index 12e879f91..81f1211e1 100644 --- a/crates/aether-data/runtime/src/lifecycle/migrate/tests/legacy_overview_upgrade.rs +++ b/crates/aether-data/runtime/src/lifecycle/migrate/tests/legacy_overview_upgrade.rs @@ -107,7 +107,14 @@ WHERE version=20260919000000; .iter() .map(|migration| migration.version) .collect::>(), - vec![DIRTY_EVENTS, 20260921010000, 20260921020000, 20260921020100, 20261001000000] + vec![ + DIRTY_EVENTS, + 20260921010000, + 20260921020000, + 20260921020100, + 20260923000000, + 20261001000000, + ] ); assert_eq!( rows_snapshot(&pool, "_sqlx_migrations").await, diff --git a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs index 665803064..cf70cc1bb 100644 --- a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs +++ b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs @@ -346,6 +346,16 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format "openai:chat" | "openai:responses" | "claude:messages" | "openai:image" ) } + "xai" => { + matches!(auth_type.as_str(), "oauth" | "bearer" | "api_key") + && matches!( + api_format.as_str(), + "openai:responses" + | "openai:responses:compact" + | "openai:image" + | "openai:video" + ) + } "windsurf" => { matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer") && api_format == "openai:chat" @@ -591,6 +601,59 @@ mod tests { assert_eq!(rows[0].global_model_name, "grok-4.20-0309-non-reasoning"); } + #[tokio::test] + async fn includes_xai_oauth_rows_for_responses_models() { + let mut row = sample_row("provider-xai", "openai:responses", "grok-4", 10); + row.provider_type = "xai".to_string(); + row.provider_name = "xai".to_string(); + row.key_auth_type = "oauth".to_string(); + row.key_api_formats = Some(vec![ + "openai:responses".to_string(), + "openai:responses:compact".to_string(), + ]); + let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![row]); + + let rows = repository + .list_for_exact_api_format("openai:responses") + .await + .expect("list should succeed"); + + assert_eq!(rows.len(), 1); + assert_eq!(rows[0].provider_type, "xai"); + assert_eq!(rows[0].global_model_name, "grok-4"); + } + + #[tokio::test] + async fn includes_xai_oauth_rows_for_image_and_video_models() { + let mut image = sample_row("provider-xai", "openai:image", "grok-imagine-image", 10); + image.provider_type = "xai".to_string(); + image.provider_name = "xai".to_string(); + image.key_auth_type = "oauth".to_string(); + image.key_api_formats = Some(vec!["openai:image".to_string(), "openai:video".to_string()]); + + let mut video = image.clone(); + video.endpoint_id = "endpoint-video".to_string(); + video.endpoint_api_format = "openai:video".to_string(); + video.global_model_name = "grok-imagine-video".to_string(); + video.model_provider_model_name = "grok-imagine-video".to_string(); + + let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![image, video]); + + let image_rows = repository + .list_for_exact_api_format("openai:image") + .await + .expect("list should succeed"); + assert_eq!(image_rows.len(), 1); + assert_eq!(image_rows[0].global_model_name, "grok-imagine-image"); + + let video_rows = repository + .list_for_exact_api_format("openai:video") + .await + .expect("list should succeed"); + assert_eq!(video_rows.len(), 1); + assert_eq!(video_rows[0].global_model_name, "grok-imagine-video"); + } + #[tokio::test] async fn requested_model_filter_respects_endpoint_scoped_default_mapping() { let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10); diff --git a/crates/aether-data/runtime/src/repository/usage/memory.rs b/crates/aether-data/runtime/src/repository/usage/memory.rs index d10d9ddd8..ab921e9e7 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory.rs @@ -600,6 +600,22 @@ fn usage_matches_summary_query( return false; } } + if let Some(user_ids) = query.user_ids.as_deref() { + if !item + .user_id + .as_ref() + .is_some_and(|user_id| user_ids.contains(user_id)) + { + return false; + } + } + if query + .provider_names + .as_ref() + .is_some_and(|names| !names.contains(&item.provider_name)) + { + return false; + } if let Some(provider_name) = query.provider_name.as_deref() { if item.provider_name != provider_name { return false; @@ -627,6 +643,22 @@ fn usage_matches_time_series_query( return false; } } + if let Some(user_ids) = query.user_ids.as_deref() { + if !item + .user_id + .as_ref() + .is_some_and(|user_id| user_ids.contains(user_id)) + { + return false; + } + } + if query + .provider_names + .as_ref() + .is_some_and(|names| !names.contains(&item.provider_name)) + { + return false; + } if let Some(provider_name) = query.provider_name.as_deref() { if item.provider_name != provider_name { return false; @@ -998,6 +1030,22 @@ fn usage_matches_leaderboard_query( return false; } } + if let Some(user_ids) = query.user_ids.as_deref() { + if !item + .user_id + .as_ref() + .is_some_and(|user_id| user_ids.contains(user_id)) + { + return false; + } + } + if query + .provider_names + .as_ref() + .is_some_and(|names| !names.contains(&item.provider_name)) + { + return false; + } if let Some(provider_name) = query.provider_name.as_deref() { if item.provider_name != provider_name { return false; diff --git a/crates/aether-data/runtime/src/repository/usage/memory/tests.rs b/crates/aether-data/runtime/src/repository/usage/memory/tests.rs index c0e1ac629..3d8c0a2a9 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory/tests.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory/tests.rs @@ -18,7 +18,7 @@ use aether_data_contracts::repository::usage::{ UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery, UsageAuditListQuery, UsageAuditSummaryQuery, UsageBodyCaptureState, UsageBodyField, UsageDashboardSummaryQuery, UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageProviderPerformanceQuery, - UsageTimeSeriesGranularity, + UsageTimeSeriesGranularity, UsageTimeSeriesQuery, }; use serde_json::json; @@ -937,6 +937,7 @@ async fn unmetered_session_audit_counts_lifecycle_without_token_or_cost_contribu let summary = repository .summarize_usage_audits(&UsageAuditSummaryQuery { + provider_names: None, created_from_unix_secs: 0, created_until_unix_secs: 1_000, ..UsageAuditSummaryQuery::default() @@ -2693,10 +2694,12 @@ async fn dashboard_and_leaderboard_total_tokens_use_effective_cache_aware_tokens let leaderboard = repository .summarize_usage_leaderboard(&UsageLeaderboardQuery { + provider_names: None, created_from_unix_secs: 1_711_000_000, created_until_unix_secs: 1_711_000_001, group_by: UsageLeaderboardGroupBy::User, user_id: None, + user_ids: None, provider_name: None, model: None, }) @@ -2706,6 +2709,140 @@ async fn dashboard_and_leaderboard_total_tokens_use_effective_cache_aware_tokens assert_eq!(leaderboard[0].total_tokens, 120); } +#[tokio::test] +async fn usage_analytics_filters_by_multiple_user_ids() { + let user_one = sample_usage("req-user-1", 1_711_000_000); + let mut user_two = sample_usage("req-user-2", 1_711_000_000); + user_two.user_id = Some("user-2".to_string()); + let mut user_three = sample_usage("req-user-3", 1_711_000_000); + user_three.user_id = Some("user-3".to_string()); + let repository = InMemoryUsageReadRepository::seed(vec![user_one, user_two, user_three]); + let scoped_user_ids = vec!["user-1".to_string(), "user-2".to_string()]; + + let summary = repository + .summarize_usage_audits(&UsageAuditSummaryQuery { + provider_names: None, + created_from_unix_secs: 1_711_000_000, + created_until_unix_secs: 1_711_000_001, + user_ids: Some(scoped_user_ids.clone()), + ..Default::default() + }) + .await + .expect("summary should filter by multiple users"); + assert_eq!(summary.total_requests, 2); + + let buckets = repository + .summarize_usage_time_series(&UsageTimeSeriesQuery { + provider_names: None, + created_from_unix_secs: 1_711_000_000, + created_until_unix_secs: 1_711_000_001, + granularity: UsageTimeSeriesGranularity::Day, + tz_offset_minutes: 0, + user_id: None, + user_ids: Some(scoped_user_ids.clone()), + provider_name: None, + model: None, + }) + .await + .expect("time series should filter by multiple users"); + assert_eq!( + buckets + .iter() + .map(|bucket| bucket.total_requests) + .sum::(), + 2 + ); + + let leaderboard = repository + .summarize_usage_leaderboard(&UsageLeaderboardQuery { + provider_names: None, + created_from_unix_secs: 1_711_000_000, + created_until_unix_secs: 1_711_000_001, + group_by: UsageLeaderboardGroupBy::User, + user_id: None, + user_ids: Some(scoped_user_ids), + provider_name: None, + model: None, + }) + .await + .expect("leaderboard should filter by multiple users"); + assert_eq!(leaderboard.len(), 2); + assert!(leaderboard.iter().all(|item| item.group_key != "user-3")); +} + +#[tokio::test] +async fn usage_analytics_intersects_provider_allowlist_and_user_scope() { + let mut a = sample_usage("allowed", 1_711_000_000); + a.provider_name = "Gemini".to_string(); + a.user_id = Some("user-1".to_string()); + let mut b = a.clone(); + b.request_id = "other-provider".to_string(); + b.provider_name = "Other".to_string(); + let mut c = a.clone(); + c.request_id = "other-user".to_string(); + c.user_id = Some("user-2".to_string()); + let repository = InMemoryUsageReadRepository::seed(vec![a, b, c]); + for (names, provider, count) in [ + ( + Some(vec!["Gemini".to_string(), "Gemini".to_string()]), + None, + 1, + ), + (Some(vec![]), None, 0), + ( + Some(vec!["Gemini".to_string()]), + Some("Other".to_string()), + 0, + ), + (None, None, 2), + ] { + let summary = repository + .summarize_usage_audits(&UsageAuditSummaryQuery { + created_from_unix_secs: 1_711_000_000, + created_until_unix_secs: 1_711_000_001, + user_ids: Some(vec!["user-1".to_string()]), + provider_names: names.clone(), + provider_name: provider.clone(), + ..Default::default() + }) + .await + .unwrap(); + assert_eq!(summary.total_requests, count); + let buckets = repository + .summarize_usage_time_series(&UsageTimeSeriesQuery { + created_from_unix_secs: 1_711_000_000, + created_until_unix_secs: 1_711_000_001, + user_id: None, + user_ids: Some(vec!["user-1".to_string()]), + provider_names: names.clone(), + provider_name: provider.clone(), + model: None, + granularity: UsageTimeSeriesGranularity::Day, + tz_offset_minutes: 0, + }) + .await + .unwrap(); + assert_eq!( + buckets.iter().map(|row| row.total_requests).sum::(), + count + ); + let rows = repository + .summarize_usage_leaderboard(&UsageLeaderboardQuery { + created_from_unix_secs: 1_711_000_000, + created_until_unix_secs: 1_711_000_001, + user_id: None, + user_ids: Some(vec!["user-1".to_string()]), + provider_names: names, + provider_name: provider, + model: None, + group_by: UsageLeaderboardGroupBy::User, + }) + .await + .unwrap(); + assert_eq!(rows.iter().map(|row| row.request_count).sum::(), count); + } +} + #[tokio::test] async fn summarizes_provider_api_key_last_used_at_in_seconds() { let repository = InMemoryUsageReadRepository::seed(vec![ diff --git a/crates/aether-data/runtime/src/repository/wallet/memory.rs b/crates/aether-data/runtime/src/repository/wallet/memory.rs index fb2287ce4..8a1fe0056 100644 --- a/crates/aether-data/runtime/src/repository/wallet/memory.rs +++ b/crates/aether-data/runtime/src/repository/wallet/memory.rs @@ -2328,8 +2328,13 @@ impl WalletWriteRepository for InMemoryWalletRepository { async fn adjust_wallet_balance( &self, _input: AdjustWalletBalanceInput, - ) -> Result, DataLayerError> - { + ) -> Result< + Option<( + StoredWalletSnapshot, + Option, + )>, + DataLayerError, + > { Ok(None) } diff --git a/crates/aether-data/runtime/src/repository/wallet/mod.rs b/crates/aether-data/runtime/src/repository/wallet/mod.rs index 8029a4e77..387a8346f 100644 --- a/crates/aether-data/runtime/src/repository/wallet/mod.rs +++ b/crates/aether-data/runtime/src/repository/wallet/mod.rs @@ -16,27 +16,30 @@ pub use aether_data_contracts::repository::wallet::{ wallet_recharge_order_is_checkout_placeholder, wallet_recharge_order_is_reclaimable_placeholder, wallet_recharge_replay_matches, wallet_recharge_response_is_checkout_placeholder, wallet_refund_proof_is_success, - AdjustWalletBalanceInput, AdminPaymentCallbackRecord, AdminPaymentOrderListQuery, - AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminWalletLedgerQuery, - AdminWalletListQuery, AdminWalletPaymentOrderRecord, AdminWalletRefundRecord, - AdminWalletRefundRequestListQuery, AdminWalletTransactionRecord, CanonicalWalletRefundFields, - CompareAndSwapPaymentOrderStripeClientSecretInput, CompleteAdminWalletRefundInput, - CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, - CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, - CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome, - CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, - CreatedAdminRedeemCodePlaintext, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, + AdjustWalletBalanceInBatchInput, AdjustWalletBalanceInput, AdminPaymentCallbackRecord, + AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, + AdminUserWalletBalanceBatchContext, AdminUserWalletBalanceBatchUserOutcome, + AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletPaymentOrderRecord, + AdminWalletRefundRecord, AdminWalletRefundRequestListQuery, AdminWalletTransactionRecord, + CanonicalWalletRefundFields, CompareAndSwapPaymentOrderStripeClientSecretInput, + CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput, + CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput, + CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput, + CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput, + CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext, + CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, - FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput, - ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput, - RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback, - StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, - StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, - StoredAdminRedeemCodePage, StoredAdminWalletLedgerItem, StoredAdminWalletLedgerPage, - StoredAdminWalletListItem, StoredAdminWalletListPage, StoredAdminWalletRefund, - StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem, - StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, - StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, + FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, + PrepareAdminUserWalletBalanceBatchInput, PrepareAdminUserWalletBalanceBatchOutcome, + ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, + ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome, + StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, + StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, + StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminUserWalletBalanceBatch, + StoredAdminWalletLedgerItem, StoredAdminWalletLedgerPage, StoredAdminWalletListItem, + StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage, + StoredAdminWalletRefundRequestItem, StoredAdminWalletRefundRequestPage, + StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome, WalletReadRepository, WalletReadSeed, WalletReadSnapshot, WalletRepository, diff --git a/crates/aether-model-fetch/src/logic.rs b/crates/aether-model-fetch/src/logic.rs index 9b8523921..e0eb69759 100644 --- a/crates/aether-model-fetch/src/logic.rs +++ b/crates/aether-model-fetch/src/logic.rs @@ -546,7 +546,7 @@ pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool { pub fn provider_type_uses_preset_models(provider_type: &str) -> bool { matches!( provider_type.trim().to_ascii_lowercase().as_str(), - "claude_code" | "gemini_cli" | "grok" + "claude_code" | "gemini_cli" | "grok" | "xai" ) } @@ -604,6 +604,22 @@ pub fn preset_models_for_provider(provider_type: &str) -> Option> { preset_model("grok-imagine-image-pro", "xai", "Grok Imagine Image Pro", "openai:image"), preset_model("grok-imagine-image-edit", "xai", "Grok Imagine Image Edit", "openai:image"), ], + "xai" => vec![ + preset_model("grok-4.6", "xai", "Grok 4.6", "openai:responses"), + preset_model("grok-build-0.1", "xai", "Grok Build 0.1", "openai:responses"), + preset_model("grok-4.5", "xai", "Grok 4.5", "openai:responses"), + preset_model("grok-4.3", "xai", "Grok 4.3", "openai:responses"), + preset_model("grok-4.20-0309-reasoning", "xai", "Grok 4.20 0309 Reasoning", "openai:responses"), + preset_model("grok-4.20-0309-non-reasoning", "xai", "Grok 4.20 0309 Non-Reasoning", "openai:responses"), + preset_model("grok-4.20-multi-agent-0309", "xai", "Grok 4.20 Multi-Agent 0309", "openai:responses"), + preset_model("grok-3-mini", "xai", "Grok 3 Mini", "openai:responses"), + preset_model("grok-3-mini-fast", "xai", "Grok 3 Mini Fast", "openai:responses"), + preset_model("grok-composer-2.5-fast", "xai", "Grok Composer 2.5 Fast", "openai:responses"), + preset_model("grok-imagine-image", "xai", "Grok Imagine Image", "openai:image"), + preset_model("grok-imagine-image-quality", "xai", "Grok Imagine Image Quality", "openai:image"), + preset_model("grok-imagine-video", "xai", "Grok Imagine Video", "openai:video"), + preset_model("grok-imagine-video-1.5", "xai", "Grok Imagine Video 1.5", "openai:video"), + ], _ => return None, }; Some(models) @@ -977,7 +993,7 @@ fn build_codex_models_url(base_url: &str, client_version: Option<&str>) -> Optio } else if !has_client_version { query_parts.push(format!( "client_version={}", - aether_ai_formats::CODEX_CLIENT_VERSION + aether_ai_formats::codex_client_version() )); } if !query_parts.is_empty() { @@ -1338,6 +1354,7 @@ mod tests { #[test] fn build_models_fetch_url_uses_codex_backend_models_endpoint() { + let client_version = aether_ai_formats::codex_client_version(); assert_eq!( build_models_fetch_url( "codex", @@ -1347,7 +1364,7 @@ mod tests { Some(( format!( "https://chatgpt.com/backend-api/codex/models?client_version={}", - aether_ai_formats::CODEX_CLIENT_VERSION + client_version ), "openai:responses".to_string() )) @@ -1977,4 +1994,39 @@ mod tests { assert_eq!(models[15]["api_formats"], json!(["openai:image"])); assert_eq!(models[18]["api_formats"], json!(["openai:image"])); } + + #[test] + fn preset_models_cover_xai_cli_catalog() { + let models = preset_models_for_provider("xai").expect("preset models should exist"); + let model_ids = models + .iter() + .map(|model| model["id"].as_str().expect("model id")) + .collect::>(); + assert_eq!( + model_ids, + vec![ + "grok-4.6", + "grok-build-0.1", + "grok-4.5", + "grok-4.3", + "grok-4.20-0309-reasoning", + "grok-4.20-0309-non-reasoning", + "grok-4.20-multi-agent-0309", + "grok-3-mini", + "grok-3-mini-fast", + "grok-composer-2.5-fast", + "grok-imagine-image", + "grok-imagine-image-quality", + "grok-imagine-video", + "grok-imagine-video-1.5", + ] + ); + assert!(models.iter().all(|model| model["owned_by"] == json!("xai"))); + assert_eq!(models[0]["api_formats"], json!(["openai:responses"])); + assert_eq!(models[10]["api_formats"], json!(["openai:image"])); + assert_eq!(models[12]["api_formats"], json!(["openai:video"])); + assert!(models + .iter() + .any(|model| model["id"] == "grok-imagine-image")); + } } diff --git a/crates/aether-model-fetch/src/transport.rs b/crates/aether-model-fetch/src/transport.rs index dc1e29b88..cc1d165e7 100644 --- a/crates/aether-model-fetch/src/transport.rs +++ b/crates/aether-model-fetch/src/transport.rs @@ -627,28 +627,30 @@ fn standard_models_fetch_headers( let api_format = aether_ai_formats::normalize_api_format_alias(api_format); let provider_type = provider_type.trim().to_ascii_lowercase(); if provider_type == "codex" && api_format.starts_with("openai:") { + // 模型目录请求也必须使用当前动态画像,不能回退到编译时固定版本。 + let dynamic_client_version = aether_ai_formats::codex_client_version(); let client_version = codex_client_version .map(str::trim) .filter(|value| !value.is_empty()) - .unwrap_or(aether_ai_formats::CODEX_CLIENT_VERSION); + .unwrap_or(dynamic_client_version.as_str()); return BTreeMap::from([ ( "user-agent".to_string(), format!( "{}/{client_version}", - aether_ai_formats::CODEX_CLIENT_ORIGINATOR + aether_ai_formats::codex_client_originator() ), ), ( "originator".to_string(), - aether_ai_formats::CODEX_CLIENT_ORIGINATOR.to_string(), + aether_ai_formats::codex_client_originator(), ), ]); } match api_format.as_str() { "openai:responses" | "openai:responses:compact" => BTreeMap::from([( "user-agent".to_string(), - aether_ai_formats::CODEX_CLIENT_USER_AGENT.to_string(), + aether_ai_formats::codex_client_user_agent(), )]), "claude:messages" => { let mut headers = BTreeMap::from([( @@ -894,7 +896,7 @@ mod tests { assert_eq!(plan.url, "https://example.com/models"); assert_eq!( plan.headers.get("user-agent").map(String::as_str), - Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT) + Some(aether_ai_formats::codex_client_user_agent().as_str()) ); assert_eq!( plan.headers.get("authorization").map(String::as_str), @@ -993,7 +995,7 @@ mod tests { plan.url, format!( "https://chatgpt.com/backend-api/codex/models?client_version={}", - aether_ai_formats::CODEX_CLIENT_VERSION + aether_ai_formats::codex_client_version() ) ); assert_eq!( @@ -1018,7 +1020,7 @@ mod tests { ); assert_eq!( plan.headers.get("user-agent").map(String::as_str), - Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT) + Some(aether_ai_formats::codex_client_user_agent().as_str()) ); assert!(!plan.headers.contains_key("version")); } diff --git a/crates/aether-oauth/src/provider/providers/generic.rs b/crates/aether-oauth/src/provider/providers/generic.rs index d3f761573..cfe85b19d 100644 --- a/crates/aether-oauth/src/provider/providers/generic.rs +++ b/crates/aether-oauth/src/provider/providers/generic.rs @@ -150,6 +150,27 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ uses_json_payload: false, include_scope_in_token_request: true, }, + GenericProviderOAuthTemplate { + provider_type: "xai", + display_name: "xAI", + authorize_url: "https://auth.x.ai/oauth2/device/code", + token_url: "https://auth.x.ai/oauth2/token", + client_id: "b1a00492-073a-47ea-816f-4c329264a828", + client_id_env: None, + client_secret_env: None, + scopes: &[ + "openid", + "profile", + "email", + "offline_access", + "grok-cli:access", + "api:access", + ], + redirect_uri: "", + use_pkce: false, + uses_json_payload: false, + include_scope_in_token_request: false, + }, ]; #[derive(Clone)] @@ -212,6 +233,10 @@ impl GenericProviderOAuthAdapter { self } + pub(super) fn token_url_for_provider(&self) -> String { + self.token_url() + } + fn token_url(&self) -> String { self.token_url_override .clone() @@ -389,7 +414,10 @@ impl GenericProviderOAuthAdapter { self.token_set_from_payload(payload) } - fn token_set_from_payload(&self, payload: Value) -> Result { + pub(super) fn token_set_from_payload( + &self, + payload: Value, + ) -> Result { let token_set = OAuthTokenSet::from_token_payload(payload.clone()) .ok_or_else(|| OAuthError::invalid_response("token response missing access_token"))?; let mut auth_config = serde_json::Map::new(); @@ -945,6 +973,7 @@ mod tests { fn resolves_generic_provider_templates() { assert!(template_for_provider_type("codex").is_some()); assert!(template_for_provider_type("claude_code").is_some()); + assert!(template_for_provider_type("xai").is_some()); assert!(template_for_provider_type("kiro").is_none()); } diff --git a/crates/aether-oauth/src/provider/providers/mod.rs b/crates/aether-oauth/src/provider/providers/mod.rs index f6fe878a8..f0cfeea68 100644 --- a/crates/aether-oauth/src/provider/providers/mod.rs +++ b/crates/aether-oauth/src/provider/providers/mod.rs @@ -4,6 +4,7 @@ mod codex; mod generic; mod kiro; mod windsurf; +mod xai; pub use antigravity::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL}; pub use claude_code::{ @@ -27,3 +28,7 @@ pub use windsurf::{ WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE, WINDSURF_SHOW_AUTH_TOKEN_REDIRECT, WINDSURF_SIGNIN_URL, }; +pub use xai::{ + XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_GRANT_TYPE, + XAI_DEVICE_CODE_URL, XAI_OAUTH_SCOPES, XAI_PROVIDER_TYPE, XAI_TOKEN_URL, +}; diff --git a/crates/aether-oauth/src/provider/providers/xai.rs b/crates/aether-oauth/src/provider/providers/xai.rs new file mode 100644 index 000000000..f936bdd04 --- /dev/null +++ b/crates/aether-oauth/src/provider/providers/xai.rs @@ -0,0 +1,668 @@ +use super::generic::{template_for_provider_type, GenericProviderOAuthAdapter}; +use crate::core::{ + current_unix_secs, redacted_oauth_error_body_excerpt, OAuthDeviceAuthorization, OAuthError, +}; +use crate::network::{OAuthHttpExecutor, OAuthHttpRequest}; +use crate::provider::{ + ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthCapabilities, + ProviderOAuthImportInput, ProviderOAuthRequestAuth, ProviderOAuthTokenSet, + ProviderOAuthTransportContext, +}; +use async_trait::async_trait; +use serde_json::{json, Map, Value}; +use std::collections::BTreeMap; +use url::form_urlencoded; + +pub const XAI_PROVIDER_TYPE: &str = "xai"; +pub const XAI_DEVICE_CODE_URL: &str = "https://auth.x.ai/oauth2/device/code"; +pub const XAI_TOKEN_URL: &str = "https://auth.x.ai/oauth2/token"; +pub const XAI_CLIENT_ID: &str = "b1a00492-073a-47ea-816f-4c329264a828"; +pub const XAI_OAUTH_SCOPES: &[&str] = &[ + "openid", + "profile", + "email", + "offline_access", + "grok-cli:access", + "api:access", +]; +pub const XAI_DEVICE_CODE_GRANT_TYPE: &str = "urn:ietf:params:oauth:grant-type:device_code"; + +const DEFAULT_DEVICE_EXPIRES_IN_SECS: u64 = 600; +const DEFAULT_DEVICE_POLL_INTERVAL_SECS: u64 = 5; + +#[derive(Debug, Clone, PartialEq)] +pub enum XaiDevicePollOutcome { + Pending, + SlowDown, + Authorized(Box), +} + +#[derive(Clone)] +pub struct XaiProviderOAuthAdapter { + inner: GenericProviderOAuthAdapter, + device_url_override: Option, +} + +impl std::fmt::Debug for XaiProviderOAuthAdapter { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("XaiProviderOAuthAdapter") + .field( + "has_device_url_override", + &self.device_url_override.is_some(), + ) + .finish_non_exhaustive() + } +} + +impl Default for XaiProviderOAuthAdapter { + fn default() -> Self { + Self { + inner: GenericProviderOAuthAdapter::new( + template_for_provider_type(XAI_PROVIDER_TYPE).expect("xai template should exist"), + ), + device_url_override: None, + } + } +} + +impl XaiProviderOAuthAdapter { + pub fn with_endpoint_overrides( + mut self, + device_url: impl Into, + token_url: impl Into, + ) -> Self { + self.device_url_override = Some(device_url.into()); + self.inner = self.inner.with_token_url_override(token_url); + self + } + + fn device_url(&self) -> String { + self.device_url_override + .clone() + .unwrap_or_else(|| XAI_DEVICE_CODE_URL.to_string()) + } + + pub async fn start_device_flow( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + ) -> Result { + let form = form_urlencoded::Serializer::new(String::new()) + .append_pair("client_id", XAI_CLIENT_ID) + .append_pair("scope", &XAI_OAUTH_SCOPES.join(" ")) + .finish() + .into_bytes(); + let response = executor + .execute(OAuthHttpRequest { + request_id: "provider-oauth:xai-device-code".to_string(), + method: reqwest::Method::POST, + url: self.device_url(), + headers: form_headers(), + content_type: Some("application/x-www-form-urlencoded".to_string()), + json_body: None, + body_bytes: Some(form), + network: ctx.network.clone(), + transport_profile: None, + }) + .await?; + if !(200..300).contains(&response.status_code) { + return Err(OAuthError::HttpStatus { + status_code: response.status_code, + body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text), + }); + } + let payload = response_json(&response) + .ok_or_else(|| OAuthError::invalid_response("xAI device code response is not json"))?; + parse_device_authorization(&payload) + } + + pub async fn poll_device_token( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + device_code: &str, + ) -> Result { + let device_code = device_code.trim(); + if device_code.is_empty() { + return Err(OAuthError::invalid_request("xAI device_code is required")); + } + let form = form_urlencoded::Serializer::new(String::new()) + .append_pair("grant_type", XAI_DEVICE_CODE_GRANT_TYPE) + .append_pair("device_code", device_code) + .append_pair("client_id", XAI_CLIENT_ID) + .finish() + .into_bytes(); + let response = executor + .execute(OAuthHttpRequest { + request_id: "provider-oauth:xai-device-token".to_string(), + method: reqwest::Method::POST, + url: self.inner.token_url_for_provider(), + headers: form_headers(), + content_type: Some("application/x-www-form-urlencoded".to_string()), + json_body: None, + body_bytes: Some(form), + network: ctx.network.clone(), + transport_profile: None, + }) + .await?; + let payload = response_json(&response); + if let Some(error_code) = payload.as_ref().and_then(oauth_error_code) { + return match error_code.as_str() { + "authorization_pending" => Ok(XaiDevicePollOutcome::Pending), + "slow_down" => Ok(XaiDevicePollOutcome::SlowDown), + "expired_token" => Err(OAuthError::invalid_request("xAI device code expired")), + "access_denied" => Err(OAuthError::invalid_request( + "xAI device authorization denied", + )), + other => Err(OAuthError::invalid_response(format!( + "xAI device token error: {other}" + ))), + }; + } + if !(200..300).contains(&response.status_code) { + return Err(OAuthError::HttpStatus { + status_code: response.status_code, + body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text), + }); + } + let payload = payload + .ok_or_else(|| OAuthError::invalid_response("xAI device token response is not json"))?; + let mut token_set = self.inner.token_set_from_payload(payload)?; + let raw_payload = token_set.token_set.raw_payload.clone(); + mark_oauth_auth_config(&mut token_set.auth_config); + enrich_xai_identity(&mut token_set.auth_config, raw_payload.as_ref()); + Ok(XaiDevicePollOutcome::Authorized(Box::new(token_set))) + } + + async fn import_raw_api_key( + &self, + input: &ProviderOAuthImportInput, + api_key: &str, + ) -> Result { + let api_key = api_key.trim(); + if api_key.is_empty() { + return Err(OAuthError::invalid_request("xAI api_key is required")); + } + let mut auth_config = Map::new(); + auth_config.insert("provider_type".to_string(), json!(XAI_PROVIDER_TYPE)); + auth_config.insert("auth_method".to_string(), json!("api_key")); + auth_config.insert("using_api".to_string(), json!(true)); + auth_config.insert("updated_at".to_string(), json!(current_unix_secs())); + if let Some(name) = input + .name + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + { + auth_config.insert("name".to_string(), json!(name)); + } + Ok(ProviderOAuthTokenSet { + token_set: crate::core::OAuthTokenSet { + access_token: api_key.to_string(), + refresh_token: None, + token_type: Some("Bearer".to_string()), + scope: None, + expires_at_unix_secs: None, + raw_payload: None, + }, + auth_config: Value::Object(auth_config), + }) + } +} + +#[async_trait] +impl ProviderOAuthAdapter for XaiProviderOAuthAdapter { + fn provider_type(&self) -> &'static str { + XAI_PROVIDER_TYPE + } + + fn capabilities(&self) -> ProviderOAuthCapabilities { + ProviderOAuthCapabilities { + supports_authorization_code: false, + supports_cookie_authorization: false, + supports_refresh_token_import: true, + supports_batch_import: true, + supports_device_flow: true, + supports_account_probe: false, + rotates_refresh_token: true, + } + } + + async fn import_credentials( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + input: ProviderOAuthImportInput, + ) -> Result { + if let Some(api_key) = + raw_credential_string(input.raw_credentials.as_ref(), &["api_key", "apiKey"]) + { + return self.import_raw_api_key(&input, &api_key).await; + } + let refresh_token = input + .refresh_token + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .or_else(|| { + raw_credential_string( + input.raw_credentials.as_ref(), + &["refresh_token", "refreshToken"], + ) + }); + if let Some(refresh_token) = refresh_token { + let mut imported = self + .inner + .import_credentials( + executor, + ctx, + ProviderOAuthImportInput { + refresh_token: Some(refresh_token), + ..input + }, + ) + .await?; + let raw_payload = imported.token_set.raw_payload.clone(); + mark_oauth_auth_config(&mut imported.auth_config); + enrich_xai_identity(&mut imported.auth_config, raw_payload.as_ref()); + return Ok(imported); + } + if let Some(access_token) = raw_credential_string( + input.raw_credentials.as_ref(), + &["access_token", "accessToken"], + ) { + return self.import_raw_api_key(&input, &access_token).await; + } + Err(OAuthError::invalid_request( + "xAI credentials require api_key, access_token, or refresh_token", + )) + } + + async fn refresh( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + account: &ProviderOAuthAccount, + ) -> Result { + let mut refreshed = self.inner.refresh(executor, ctx, account).await?; + let raw_payload = refreshed.token_set.raw_payload.clone(); + mark_oauth_auth_config(&mut refreshed.auth_config); + enrich_xai_identity(&mut refreshed.auth_config, raw_payload.as_ref()); + Ok(refreshed) + } + + fn resolve_request_auth( + &self, + account: &ProviderOAuthAccount, + ) -> Result { + self.inner.resolve_request_auth(account) + } + + fn account_fingerprint(&self, account: &ProviderOAuthAccount) -> Option { + self.inner.account_fingerprint(account) + } +} + +fn mark_oauth_auth_config(auth_config: &mut Value) { + let Some(object) = auth_config.as_object_mut() else { + return; + }; + object.insert("provider_type".to_string(), json!(XAI_PROVIDER_TYPE)); + object.insert("auth_method".to_string(), json!("oauth")); + object.insert("using_api".to_string(), json!(false)); +} + +fn enrich_xai_identity(auth_config: &mut Value, raw_payload: Option<&Value>) { + let Some(object) = auth_config.as_object_mut() else { + return; + }; + let id_token = raw_payload + .and_then(|payload| payload.get("id_token")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()); + if let Some(id_token) = id_token { + object + .entry("id_token".to_string()) + .or_insert_with(|| json!(id_token)); + if let Some(claims) = decode_jwt_claims(id_token) { + if !object.contains_key("email") { + if let Some(email) = claims.get("email").and_then(Value::as_str) { + let email = email.trim(); + if !email.is_empty() { + object.insert("email".to_string(), json!(email)); + } + } + } + if !object.contains_key("sub") { + if let Some(sub) = claims.get("sub").and_then(Value::as_str) { + let sub = sub.trim(); + if !sub.is_empty() { + object.insert("sub".to_string(), json!(sub)); + } + } + } + } + } +} + +fn parse_device_authorization(payload: &Value) -> Result { + let device_code = + json_non_empty_string(payload, &["device_code", "deviceCode"]).ok_or_else(|| { + OAuthError::invalid_response("xAI device code response missing device_code") + })?; + let user_code = + json_non_empty_string(payload, &["user_code", "userCode"]).ok_or_else(|| { + OAuthError::invalid_response("xAI device code response missing user_code") + })?; + let verification_uri = json_non_empty_string( + payload, + &["verification_uri", "verificationUri", "verification_url"], + ) + .unwrap_or_default(); + let verification_uri_complete = json_non_empty_string( + payload, + &[ + "verification_uri_complete", + "verificationUriComplete", + "verification_url_complete", + ], + ) + .unwrap_or_else(|| verification_uri.clone()); + if verification_uri.is_empty() && verification_uri_complete.is_empty() { + return Err(OAuthError::invalid_response( + "xAI device code response missing verification URI", + )); + } + Ok(OAuthDeviceAuthorization { + device_code, + user_code, + verification_uri: if verification_uri.is_empty() { + verification_uri_complete.clone() + } else { + verification_uri + }, + verification_uri_complete, + expires_in: json_u64(payload, &["expires_in", "expiresIn"]) + .unwrap_or(DEFAULT_DEVICE_EXPIRES_IN_SECS), + interval: json_u64(payload, &["interval"]).unwrap_or(DEFAULT_DEVICE_POLL_INTERVAL_SECS), + }) +} + +fn raw_credential_string(raw: Option<&Value>, keys: &[&str]) -> Option { + let object = raw?.as_object()?; + keys.iter().find_map(|key| { + object + .get(*key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }) +} + +fn response_json(response: &crate::network::OAuthHttpResponse) -> Option { + response + .json_body + .clone() + .or_else(|| serde_json::from_str::(&response.body_text).ok()) +} + +fn oauth_error_code(payload: &Value) -> Option { + payload + .get("error") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + +fn json_non_empty_string(payload: &Value, keys: &[&str]) -> Option { + keys.iter().find_map(|key| { + payload + .get(*key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }) +} + +fn json_u64(payload: &Value, keys: &[&str]) -> Option { + keys.iter().find_map(|key| match payload.get(*key)? { + Value::Number(number) => number.as_u64(), + Value::String(string) => string.trim().parse::().ok(), + _ => None, + }) +} + +fn form_headers() -> BTreeMap { + BTreeMap::from([ + ( + "content-type".to_string(), + "application/x-www-form-urlencoded".to_string(), + ), + ("accept".to_string(), "application/json".to_string()), + ]) +} + +fn decode_jwt_claims(token: &str) -> Option> { + use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; + const MAX_UNVERIFIED_JWT_CLAIMS_BYTES: usize = 64 * 1024; + + let payload = token.split('.').nth(1)?; + let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES + .saturating_add(2) + .checked_div(3) + .unwrap_or(usize::MAX) + .saturating_mul(4); + if payload.len() > max_encoded_len { + return None; + } + let bytes = URL_SAFE_NO_PAD.decode(payload.as_bytes()).ok()?; + if bytes.len() > MAX_UNVERIFIED_JWT_CLAIMS_BYTES { + return None; + } + serde_json::from_slice::(&bytes) + .ok()? + .as_object() + .cloned() +} + +#[cfg(test)] +mod tests { + use super::{ + XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_GRANT_TYPE, + XAI_OAUTH_SCOPES, XAI_PROVIDER_TYPE, + }; + use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse}; + use crate::provider::{ + ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthImportInput, + ProviderOAuthTransportContext, + }; + use async_trait::async_trait; + use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; + use serde_json::{json, Value}; + use std::collections::BTreeMap; + use std::sync::{Arc, Mutex}; + + #[derive(Clone)] + struct ScriptedExecutor { + seen_request: Arc>>, + status_code: u16, + payload: Value, + } + + #[async_trait] + impl OAuthHttpExecutor for ScriptedExecutor { + async fn execute( + &self, + request: OAuthHttpRequest, + ) -> Result { + *self.seen_request.lock().expect("mutex should lock") = Some(request); + Ok(OAuthHttpResponse { + status_code: self.status_code, + body_text: self.payload.to_string(), + json_body: Some(self.payload.clone()), + }) + } + } + + fn transport_context() -> ProviderOAuthTransportContext { + ProviderOAuthTransportContext { + provider_id: "provider-xai".to_string(), + provider_type: XAI_PROVIDER_TYPE.to_string(), + endpoint_id: None, + key_id: None, + auth_type: Some("oauth".to_string()), + decrypted_api_key: None, + decrypted_auth_config: None, + provider_config: None, + endpoint_config: None, + key_config: None, + network: crate::network::OAuthNetworkContext::provider_operation(None), + } + } + + fn encoded_jwt(claims: &Value) -> String { + format!( + "header.{}.signature", + URL_SAFE_NO_PAD.encode(serde_json::to_vec(claims).expect("claims should encode")) + ) + } + + #[tokio::test] + async fn imports_api_key_as_official_api_credential() { + let adapter = XaiProviderOAuthAdapter::default(); + let executor = ScriptedExecutor { + seen_request: Arc::new(Mutex::new(None)), + status_code: 200, + payload: json!({}), + }; + let result = adapter + .import_credentials( + &executor, + &transport_context(), + ProviderOAuthImportInput { + provider_type: XAI_PROVIDER_TYPE.to_string(), + name: Some("work".to_string()), + refresh_token: None, + raw_credentials: Some(json!({"api_key": "xai-key-123"})), + network: crate::network::OAuthNetworkContext::provider_operation(None), + }, + ) + .await + .expect("api key import should succeed"); + + assert_eq!(result.token_set.access_token, "xai-key-123"); + assert_eq!(result.auth_config["using_api"], json!(true)); + assert_eq!(result.auth_config["auth_method"], json!("api_key")); + assert!(executor.seen_request.lock().expect("lock").is_none()); + } + + #[tokio::test] + async fn device_poll_treats_authorization_pending_as_pending() { + let adapter = XaiProviderOAuthAdapter::default(); + let seen = Arc::new(Mutex::new(None)); + let executor = ScriptedExecutor { + seen_request: Arc::clone(&seen), + status_code: 400, + payload: json!({"error": "authorization_pending"}), + }; + let outcome = adapter + .poll_device_token(&executor, &transport_context(), "device-code") + .await + .expect("pending should not be fatal"); + assert_eq!(outcome, XaiDevicePollOutcome::Pending); + + let request = seen.lock().expect("lock").clone().expect("request"); + let body = request.body_bytes.expect("body"); + let fields = url::form_urlencoded::parse(&body) + .into_owned() + .collect::>(); + assert_eq!(fields["grant_type"], XAI_DEVICE_CODE_GRANT_TYPE); + assert_eq!(fields["device_code"], "device-code"); + assert_eq!(fields["client_id"], XAI_CLIENT_ID); + } + + #[tokio::test] + async fn refresh_posts_client_id_and_refresh_token_without_scope() { + let adapter = XaiProviderOAuthAdapter::default(); + let seen = Arc::new(Mutex::new(None)); + let id_token = encoded_jwt(&json!({"email": "user@x.ai", "sub": "subject-1"})); + let executor = ScriptedExecutor { + seen_request: Arc::clone(&seen), + status_code: 200, + payload: json!({ + "access_token": "new-access", + "refresh_token": "new-refresh", + "id_token": id_token, + "expires_in": 3600 + }), + }; + let account = ProviderOAuthAccount { + provider_type: XAI_PROVIDER_TYPE.to_string(), + access_token: "old-access".to_string(), + auth_config: json!({ + "provider_type": XAI_PROVIDER_TYPE, + "refresh_token": "old-refresh", + "using_api": false, + }), + expires_at_unix_secs: None, + identity: BTreeMap::new(), + }; + let result = adapter + .refresh(&executor, &transport_context(), &account) + .await + .expect("refresh should succeed"); + assert_eq!(result.token_set.access_token, "new-access"); + assert_eq!(result.auth_config["using_api"], json!(false)); + assert_eq!(result.auth_config["email"], json!("user@x.ai")); + assert_eq!(result.auth_config["sub"], json!("subject-1")); + + let request = seen.lock().expect("lock").clone().expect("request"); + let body = request.body_bytes.expect("body"); + let fields = url::form_urlencoded::parse(&body) + .into_owned() + .collect::>(); + assert_eq!(fields["grant_type"], "refresh_token"); + assert_eq!(fields["client_id"], XAI_CLIENT_ID); + assert_eq!(fields["refresh_token"], "old-refresh"); + assert!(!fields.contains_key("scope")); + assert!(XAI_OAUTH_SCOPES.join(" ").contains("grok-cli:access")); + } + + #[tokio::test] + async fn start_device_flow_posts_client_id_and_scope() { + let adapter = XaiProviderOAuthAdapter::default(); + let seen = Arc::new(Mutex::new(None)); + let executor = ScriptedExecutor { + seen_request: Arc::clone(&seen), + status_code: 200, + payload: json!({ + "device_code": "dc-1", + "user_code": "ABCD-EFGH", + "verification_uri": "https://auth.x.ai/device", + "verification_uri_complete": "https://auth.x.ai/device?user_code=ABCD-EFGH", + "expires_in": 600, + "interval": 5 + }), + }; + let authorization = adapter + .start_device_flow(&executor, &transport_context()) + .await + .expect("device start should succeed"); + assert_eq!(authorization.user_code, "ABCD-EFGH"); + assert_eq!(authorization.device_code, "dc-1"); + + let request = seen.lock().expect("lock").clone().expect("request"); + let body = request.body_bytes.expect("body"); + let fields = url::form_urlencoded::parse(&body) + .into_owned() + .collect::>(); + assert_eq!(fields["client_id"], XAI_CLIENT_ID); + assert_eq!(fields["scope"], XAI_OAUTH_SCOPES.join(" ")); + } +} diff --git a/crates/aether-oauth/src/provider/service.rs b/crates/aether-oauth/src/provider/service.rs index 8b121f5de..1c5edbf4d 100644 --- a/crates/aether-oauth/src/provider/service.rs +++ b/crates/aether-oauth/src/provider/service.rs @@ -21,7 +21,7 @@ impl ProviderOAuthService { use super::providers::{ AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter, CodexProviderOAuthAdapter, GenericProviderOAuthAdapter, KiroProviderOAuthAdapter, - WindsurfProviderOAuthAdapter, + WindsurfProviderOAuthAdapter, XaiProviderOAuthAdapter, }; let mut service = Self::new() @@ -29,7 +29,8 @@ impl ProviderOAuthService { .with_adapter(Arc::new(ClaudeCodeProviderOAuthAdapter::default())) .with_adapter(Arc::new(CodexProviderOAuthAdapter::default())) .with_adapter(Arc::new(AntigravityProviderOAuthAdapter::default())) - .with_adapter(Arc::new(WindsurfProviderOAuthAdapter)); + .with_adapter(Arc::new(WindsurfProviderOAuthAdapter)) + .with_adapter(Arc::new(XaiProviderOAuthAdapter::default())); for provider_type in ["chatgpt_web", "gemini_cli"] { if let Some(adapter) = GenericProviderOAuthAdapter::for_provider_type(provider_type) { service = service.with_adapter(Arc::new(adapter)); @@ -144,6 +145,7 @@ mod tests { "antigravity", "kiro", "windsurf", + "xai", ] { assert!( service.adapter(provider_type).is_ok(), diff --git a/crates/aether-provider/pool/Cargo.toml b/crates/aether-provider/pool/Cargo.toml index f08919450..ac389e6d9 100644 --- a/crates/aether-provider/pool/Cargo.toml +++ b/crates/aether-provider/pool/Cargo.toml @@ -11,6 +11,7 @@ aether-contracts.workspace = true aether-data-contracts.workspace = true aether-pool-core.workspace = true aether-provider-transport.workspace = true +chrono.workspace = true serde_json.workspace = true url.workspace = true uuid.workspace = true diff --git a/crates/aether-provider/pool/src/lib.rs b/crates/aether-provider/pool/src/lib.rs index 2eb0f0529..6b2c59adf 100644 --- a/crates/aether-provider/pool/src/lib.rs +++ b/crates/aether-provider/pool/src/lib.rs @@ -16,27 +16,30 @@ pub use presets::{ pub use provider::{ProviderPoolAdapter, ProviderPoolMemberInput}; pub use providers::{ build_antigravity_pool_quota_request, build_antigravity_pool_quota_summary_request, - build_chatgpt_web_pool_quota_request, build_codex_pool_quota_request, - build_codex_pool_reset_credit_consume_request, build_codex_pool_reset_credits_request, - build_gemini_cli_pool_quota_request, build_kiro_pool_quota_request, - build_windsurf_pool_model_configs_request, + build_chatgpt_web_pool_quota_request, build_claude_code_pool_quota_request, + build_codex_pool_quota_request, build_codex_pool_reset_credit_consume_request, + build_codex_pool_reset_credits_request, build_gemini_cli_pool_quota_request, + build_kiro_pool_quota_request, build_windsurf_pool_model_configs_request, build_windsurf_pool_model_configs_request_with_base_url, build_windsurf_pool_quota_request, build_windsurf_pool_quota_request_with_base_url, build_windsurf_pool_rate_limit_request, - build_windsurf_pool_rate_limit_request_with_base_url, enrich_chatgpt_web_quota_metadata, - grok_mode_id_for_model, grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model, + build_windsurf_pool_rate_limit_request_with_base_url, build_xai_pool_billing_request, + build_xai_pool_user_request, enrich_chatgpt_web_quota_metadata, grok_mode_id_for_model, + grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model, grok_supported_quota_windows_for_tier, normalize_chatgpt_web_image_quota_limit, - AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter, - DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter, - KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, UnsupportedQuotaProviderPoolAdapter, + AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, ClaudeCodeProviderPoolAdapter, + CodexProviderPoolAdapter, DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, + GrokProviderPoolAdapter, KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, + UnsupportedQuotaProviderPoolAdapter, XaiProviderPoolAdapter, ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH, ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, CHATGPT_WEB_CONVERSATION_INIT_PATH, CHATGPT_WEB_DEFAULT_BASE_URL, CODEX_WHAM_RESET_CREDITS_CONSUME_URL, CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL, GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, - WINDSURF_USER_STATUS_PATH, + WINDSURF_USER_STATUS_PATH, XAI_BILLING_PATH, XAI_USER_PATH, }; pub use quota::{ - provider_pool_key_account_quota_exhausted, provider_pool_key_model_quota_exhausted, + provider_pool_codex_metadata_has_account_quota, provider_pool_key_account_quota_exhausted, + provider_pool_key_minimum_quota_reached, provider_pool_key_model_quota_exhausted, provider_pool_key_model_quota_hard_blocked, provider_pool_key_quota_hard_blocked, provider_pool_key_scheduling_label, provider_pool_member_quota_snapshot, provider_pool_quota_metadata_provider_type, provider_pool_quota_metadata_updated_at, @@ -81,7 +84,8 @@ mod tests { "grok", "kiro", "vertex_ai", - "windsurf" + "windsurf", + "xai" ] ); assert!(service @@ -100,21 +104,37 @@ mod tests { [ "antigravity", "chatgpt_web", + "claude_code", "codex", "gemini_cli", "grok", "kiro", - "windsurf" + "windsurf", + "xai" ] ); + assert!(service.supports_quota_refresh("claude_code")); assert!(service.supports_quota_refresh("codex")); assert!(service.supports_quota_refresh("antigravity")); assert!(service.supports_quota_refresh("grok")); assert!(service.supports_quota_refresh("gemini_cli")); assert!(service.supports_quota_refresh("windsurf")); + assert!(service.supports_quota_refresh("xai")); + let claude_spec = build_claude_code_pool_quota_request( + "key-1", + ("authorization".to_string(), "Bearer access".to_string()), + ); + assert_eq!(claude_spec.method, "GET"); assert_eq!( - service.quota_refresh_unsupported_message("claude_code"), - "Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口" + claude_spec.url, + "https://api.anthropic.com/api/oauth/usage?cedar_ember=1&skip_spend=1" + ); + assert_eq!( + claude_spec + .headers + .get("anthropic-beta") + .map(String::as_str), + Some("oauth-2025-04-20") ); assert_eq!( service.quota_refresh_unsupported_message("vertex_ai"), @@ -642,11 +662,11 @@ mod tests { assert_eq!( free_first["providers"], - json!(["codex", "grok", "kiro", "windsurf"]) + json!(["codex", "grok", "kiro", "windsurf", "xai"]) ); assert_eq!( recent_refresh["providers"], - json!(["codex", "grok", "kiro", "windsurf"]) + json!(["codex", "grok", "kiro", "windsurf", "xai"]) ); assert_eq!(free_first["default_enabled"], json!(false)); assert_eq!(recent_refresh["default_enabled"], json!(false)); @@ -722,6 +742,275 @@ mod tests { assert!(unsupported.is_empty()); } + #[test] + fn codex_minimum_quota_respects_boundary_and_provider() { + for (used_percent, expected) in [ + (98.9, false), + (98.99999, false), + (99.0, true), + (99.5, true), + (100.0, true), + ] { + let key = sample_key(Some(json!({ + "codex": { "primary_used_percent": used_percent } + }))); + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", None), + expected, + "used_percent={used_percent}" + ); + assert!(!provider_pool_key_minimum_quota_reached(&key, "kiro", None)); + } + assert!(!provider_pool_key_minimum_quota_reached( + &sample_key(None), + "codex", + None + )); + } + + #[test] + fn codex_minimum_quota_short_model_names_still_use_account_windows() { + for (used_percent, expected) in [(98.0, false), (99.0, true)] { + let key = sample_key(Some(json!({ + "codex": { "primary_used_percent": used_percent } + }))); + for model in ["o1", "o3", "", " "] { + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", Some(model)), + expected, + "model={model:?}, used_percent={used_percent}" + ); + } + } + } + + #[test] + fn codex_minimum_quota_respects_windows_and_reset() { + let now = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("system time should be after unix epoch") + .as_secs(); + for (reset_at, expected) in [(now + 3600, true), (now - 60, false)] { + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "updated_at": now - 600, + "windows": [ + { "code": "5h", "used_ratio": 0.2 }, + { "code": "weekly", "used_ratio": 0.99, "reset_at": reset_at } + ] + } + })); + for model in [None, Some("gpt-5.4")] { + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", model), + expected + ); + } + } + } + + #[test] + fn codex_minimum_quota_checks_either_window_and_relative_resets() { + let now = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("system time should be after unix epoch") + .as_secs(); + for prefix in ["primary", "secondary"] { + for (reset_seconds, expected) in [(3600, true), (60, false)] { + let key = sample_key(Some(json!({ + "codex": { + "updated_at": now - 600, + "allowed": true, + "limit_reached": false, + format!("{prefix}_used_percent"): 99.0, + format!("{prefix}_reset_after_seconds"): reset_seconds + } + }))); + for model in [None, Some("gpt-5.4")] { + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", model), + expected, + "prefix={prefix}, reset_seconds={reset_seconds}, model={model:?}" + ); + } + assert!(!provider_pool_key_account_quota_exhausted(&key, "codex")); + } + } + } + + #[test] + fn codex_minimum_quota_uses_newest_source_with_or_without_model() { + for (snapshot_at, metadata_at, metadata_wins) in [ + (Some(200), Some(100), false), + (Some(100), Some(200), true), + (None, None, false), + (Some(100), None, false), + (None, Some(100), true), + ] { + for snapshot_reached in [false, true] { + let mut key = sample_key(Some(json!({ + "codex": { + "updated_at": metadata_at, + "primary_used_percent": if snapshot_reached { 20.0 } else { 99.0 } + } + }))); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "observed_at": snapshot_at, + "windows": [{ + "code": "5h", + "used_ratio": if snapshot_reached { 0.99 } else { 0.2 }, + "is_exhausted": false + }] + } + })); + for model in [None, Some("gpt-5.4")] { + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", model), + snapshot_reached != metadata_wins, + "snapshot_at={snapshot_at:?}, metadata_at={metadata_at:?}, model={model:?}" + ); + } + } + } + } + + #[test] + fn codex_minimum_quota_isolates_model_buckets_and_checks_each_model_window() { + for explicit_model in [false, true] { + let mut key = sample_key(None); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "windows": [ + { "code": "5h", "used_ratio": 0.2 }, + { + "code": "spark_5h", "used_ratio": 0.99, + "model": if explicit_model { Some("spark") } else { None }, + "is_exhausted": false + }, + { + "code": "spark_weekly", "used_ratio": 0.2, + "model": if explicit_model { Some("spark") } else { None }, + "is_exhausted": false + } + ] + } + })); + assert!(provider_pool_key_minimum_quota_reached( + &key, + "codex", + Some("gpt-5.3-codex-spark") + )); + for model in [None, Some("gpt-5.4")] { + assert!(!provider_pool_key_minimum_quota_reached( + &key, "codex", model + )); + } + assert_eq!( + provider_pool_key_model_quota_exhausted(&key, "codex", "gpt-5.3-codex-spark"), + Some(false) + ); + } + for (account_percent, spark_percent) in [(99.0, 20.0), (20.0, 99.0)] { + let key = sample_key(Some(json!({ + "codex": { + "primary_used_percent": account_percent, + "spark_primary_used_percent": spark_percent, + "spark_secondary_used_percent": 20.0 + } + }))); + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", Some("gpt-5.3-codex-spark")), + spark_percent == 99.0 + ); + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", Some("gpt-5.4")), + account_percent == 99.0 + ); + } + } + + #[test] + fn codex_minimum_quota_keeps_model_bucket_when_account_metadata_is_newer() { + for (snapshot_ratio, account_percent) in [(0.2, 99.0), (0.99, 20.0)] { + let mut key = sample_key(Some(json!({ + "codex": { "updated_at": 200, "primary_used_percent": account_percent } + }))); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "observed_at": 100, + "windows": [{ "code": "spark_5h", "used_ratio": snapshot_ratio }] + } + })); + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", Some("gpt-5.3-codex-spark")), + snapshot_ratio == 0.99 + ); + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", Some("gpt-5.4")), + account_percent == 99.0 + ); + + key.upstream_metadata.as_mut().unwrap()["codex"]["spark_primary_used_percent"] = + json!(account_percent); + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", Some("gpt-5.3-codex-spark")), + account_percent == 99.0, + "the newer observation for the same model bucket must win" + ); + } + } + + #[test] + fn codex_minimum_quota_supports_remaining_values_and_ignores_unknown_data() { + for (metrics, expected) in [ + (json!({ "remaining_ratio": 0.01 }), true), + (json!({ "remaining_percent": "1" }), true), + (json!({ "remaining": 1, "limit": 100 }), true), + (json!({ "used_percent": "99" }), true), + (json!({ "used_ratio": null }), false), + (json!({ "used_ratio": "NaN" }), false), + (json!({ "remaining": 0, "limit": 0 }), false), + (json!({ "remaining_ratio": 0.01001 }), false), + ] { + let key = sample_key(Some(json!({ + "codex": { "quota_by_model": { "spark": metrics } } + }))); + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", Some("gpt-5.3-codex-spark")), + expected, + "metrics={metrics}" + ); + assert!(!provider_pool_key_minimum_quota_reached( + &key, "codex", None + )); + } + let disabled = sample_key(Some(json!({ + "codex": { "primary_used_percent": 99.0, "primary_window_minutes": 0 } + }))); + assert!(!provider_pool_key_minimum_quota_reached( + &disabled, "codex", None + )); + + let mut mismatched = sample_key(None); + mismatched.status_snapshot = Some(json!({ + "quota": { + "provider_type": "kiro", + "windows": [{ "code": "weekly", "used_ratio": 0.99 }] + } + })); + assert!(!provider_pool_key_minimum_quota_reached( + &mismatched, + "codex", + None + )); + } + #[test] fn provider_quota_exhaustion_is_adapter_owned() { assert!(provider_pool_key_account_quota_exhausted( @@ -1326,6 +1615,245 @@ mod tests { assert!(!signals.quota_exhausted); } + #[test] + fn codex_newer_flat_quota_metadata_clears_stale_exhausted_snapshot() { + let now = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("system time should be after unix epoch") + .as_secs(); + let service = ProviderPoolService::with_builtin_adapters(); + // Headers and quota refreshes persist the flat metadata shape while the + // derived snapshot may still describe the previous quota observation. + for used_percent in [83.0, 99.0, 100.0] { + let mut key = sample_key(Some(json!({ + "codex": { + "updated_at": now, + "primary_used_percent": used_percent, + "primary_reset_at": now + 3600 + } + }))); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "observed_at": now - 60, + "updated_at": now - 60, + "code": "exhausted", + "exhausted": true, + "allowed": false, + "limit_reached": true, + "reset_at": now + 3600, + "windows": [{ + "code": "weekly", + "scope": "account", + "used_ratio": 1.0, + "remaining_ratio": 0.0, + "is_exhausted": true, + "reset_at": now + 3600 + }] + } + })); + for model in [None, Some("gpt-5.4"), Some("o3")] { + let signals = service.member_signals("codex", &key, None, model); + assert_eq!(signals.quota_exhausted, used_percent >= 100.0); + assert!(!signals.quota_hard_blocked); + assert_eq!( + provider_pool_key_minimum_quota_reached(&key, "codex", model), + used_percent >= 99.0 + ); + } + } + } + + #[test] + fn codex_account_and_model_quota_agree_on_explicit_allow_and_deny_flags() { + let service = ProviderPoolService::with_builtin_adapters(); + for (flags, exhausted) in [ + (json!({ "allowed": true }), false), + (json!({ "limit_reached": false }), false), + (json!({ "allowed": true, "limit_reached": true }), true), + (json!({ "allowed": false, "limit_reached": false }), true), + ] { + for use_windows in [false, true] { + let mut metadata = flags.clone(); + metadata["updated_at"] = json!(300); + if use_windows { + metadata["windows"] = json!([{ "code": "weekly", "used_ratio": 1.0 }]); + } else { + metadata["primary_used_percent"] = json!(100.0); + } + let key = sample_key(Some(json!({ "codex": metadata }))); + for model in [None, Some("gpt-5.4"), Some("o3")] { + assert_eq!( + service + .member_signals("codex", &key, None, model) + .quota_exhausted, + exhausted, + "flags={flags}, use_windows={use_windows}, model={model:?}" + ); + // The opt-in reserve still protects a numerically full + // window even when the upstream reports it as allowed. + assert!(provider_pool_key_minimum_quota_reached( + &key, "codex", model + )); + } + } + } + } + + #[test] + fn codex_model_quota_honors_latest_account_refusal_without_blocking_spark() { + let service = ProviderPoolService::with_builtin_adapters(); + for include_window in [false, true] { + let mut metadata = json!({ + "updated_at": 300, + "allowed": false, + "limit_reached": true, + "spark_primary_used_percent": 17.0 + }); + if include_window { + metadata["primary_used_percent"] = json!(83.0); + } + let mut key = sample_key(Some(json!({ "codex": metadata }))); + for include_snapshot in [false, true] { + if include_snapshot { + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "observed_at": 200, + "exhausted": false, + "windows": [{ "code": "weekly", "used_ratio": 0.83 }] + } + })); + } + assert!( + service + .member_signals("codex", &key, None, Some("gpt-5.4")) + .quota_exhausted + ); + assert!( + !service + .member_signals("codex", &key, None, Some("gpt-5.3-codex-spark")) + .quota_exhausted + ); + } + } + } + + #[test] + fn codex_quota_source_freshness_supports_all_timestamp_formats() { + let service = ProviderPoolService::with_builtin_adapters(); + for observed_at in [ + json!(1_700_000_200_u64), + json!(1_700_000_200_000_u64), + json!("1700000200000"), + json!("2023-11-14T22:16:40Z"), + ] { + for use_windows in [false, true] { + let metadata = if use_windows { + json!({ + "updated_at": observed_at, + "windows": [{ "code": "weekly", "used_ratio": 0.83 }] + }) + } else { + json!({ "updated_at": observed_at, "primary_used_percent": 83.0 }) + }; + let mut key = sample_key(Some(json!({ "codex": metadata }))); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "observed_at": "2023-11-14T22:15:00Z", + "exhausted": true, + "allowed": false, + "windows": [{ "code": "weekly", "used_ratio": 1.0 }] + } + })); + for model in [None, Some("gpt-5.4")] { + let signals = service.member_signals("codex", &key, None, model); + assert!(!signals.quota_exhausted, "timestamp={observed_at}"); + assert!(!signals.quota_hard_blocked, "timestamp={observed_at}"); + assert!(!provider_pool_key_minimum_quota_reached( + &key, "codex", model + )); + } + + // Swap freshness while retaining the same observations. An + // older usable bucket cannot erase a later exhausted snapshot. + key.status_snapshot.as_mut().unwrap()["quota"]["observed_at"] = + json!("2023-11-14T22:18:20Z"); + assert!(provider_pool_key_account_quota_exhausted(&key, "codex")); + assert!(provider_pool_key_quota_hard_blocked(&key, "codex")); + assert_eq!( + provider_pool_key_model_quota_exhausted(&key, "codex", "gpt-5.4"), + Some(true) + ); + } + } + } + + #[test] + fn codex_account_quota_recovery_requires_new_account_observation() { + for metadata in [ + json!({ "updated_at": 100, "primary_used_percent": 83.0 }), + json!({ "primary_used_percent": 83.0 }), + json!({ "updated_at": 300, "plan_type": "plus" }), + json!({ "updated_at": 300, "credits_unlimited": false }), + json!({ "updated_at": 300, "spark_primary_used_percent": 83.0 }), + json!({ "updated_at": 300, "windows": [{ "code": "weekly" }] }), + ] { + let mut key = sample_key(Some(json!({ "codex": metadata }))); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "observed_at": 200, + "exhausted": true, + "allowed": false, + "limit_reached": true, + "windows": [{ "code": "weekly", "used_ratio": 1.0 }] + } + })); + assert!(provider_pool_key_account_quota_exhausted(&key, "codex")); + assert!(provider_pool_key_quota_hard_blocked(&key, "codex")); + assert_eq!( + provider_pool_key_model_quota_exhausted(&key, "codex", "gpt-5.4"), + Some(true) + ); + } + + let mut key = sample_key(Some(json!({ + "codex": { + "updated_at": 300, + "primary_used_percent": 83.0, + "allowed": false, + "limit_reached": true + } + }))); + key.status_snapshot = Some(json!({ + "quota": { "provider_type": "codex", "updated_at": 200, "exhausted": false } + })); + assert!(provider_pool_key_account_quota_exhausted(&key, "codex")); + assert!(provider_pool_key_quota_hard_blocked(&key, "codex")); + } + + #[test] + fn codex_account_updates_do_not_override_independent_model_quota() { + for (spark_ratio, account_percent) in [(0.83, 100.0), (1.0, 83.0)] { + let mut key = sample_key(Some(json!({ + "codex": { "updated_at": 300, "primary_used_percent": account_percent } + }))); + key.status_snapshot = Some(json!({ + "quota": { + "provider_type": "codex", + "updated_at": 200, + "windows": [{ "code": "spark_5h", "used_ratio": spark_ratio }] + } + })); + assert_eq!( + provider_pool_key_model_quota_exhausted(&key, "codex", "gpt-5.3-codex-spark"), + Some(spark_ratio >= 1.0) + ); + } + } + #[test] fn codex_explicit_quota_block_is_hard_until_reset() { let now = std::time::SystemTime::now() diff --git a/crates/aether-provider/pool/src/providers/claude_code.rs b/crates/aether-provider/pool/src/providers/claude_code.rs new file mode 100644 index 000000000..195dd40cb --- /dev/null +++ b/crates/aether-provider/pool/src/providers/claude_code.rs @@ -0,0 +1,82 @@ +use std::collections::BTreeMap; + +use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint; + +use crate::capability::ProviderPoolCapabilities; +use crate::provider::{ + provider_pool_endpoint_format_matches, provider_pool_matching_endpoint, ProviderPoolAdapter, +}; +use crate::quota_refresh::ProviderPoolQuotaRequestSpec; + +pub const CLAUDE_CODE_OAUTH_USAGE_URL: &str = + "https://api.anthropic.com/api/oauth/usage?cedar_ember=1&skip_spend=1"; +pub const CLAUDE_CODE_OAUTH_BETA: &str = "oauth-2025-04-20"; +pub const CLAUDE_CODE_USAGE_USER_AGENT: &str = "claude-cli/2.1.284 (external, cli)"; + +#[derive(Debug, Clone, Default)] +pub struct ClaudeCodeProviderPoolAdapter; + +impl ProviderPoolAdapter for ClaudeCodeProviderPoolAdapter { + fn provider_type(&self) -> &'static str { + "claude_code" + } + + fn capabilities(&self) -> ProviderPoolCapabilities { + ProviderPoolCapabilities { + quota_refresh: true, + ..ProviderPoolCapabilities::default() + } + } + + fn quota_refresh_endpoint( + &self, + endpoints: &[StoredProviderCatalogEndpoint], + include_inactive: bool, + ) -> Option { + provider_pool_matching_endpoint(endpoints, include_inactive, |endpoint| { + provider_pool_endpoint_format_matches(endpoint, "claude:messages") + }) + } + + fn quota_refresh_missing_endpoint_message(&self) -> String { + "找不到有效的 claude:messages 端点".to_string() + } +} + +/// Builds the `GET /api/oauth/usage` request that reports the account's 5h / 7d +/// utilization windows (and, via `cedar_ember=1`, its quota reset credits). The URL is fixed (the origin allowlist only accepts +/// `api.anthropic.com`), independent of the inference endpoint's base URL. +pub fn build_claude_code_pool_quota_request( + key_id: &str, + authorization: (String, String), +) -> ProviderPoolQuotaRequestSpec { + let headers = BTreeMap::from([ + ("authorization".to_string(), authorization.1), + ("accept".to_string(), "application/json".to_string()), + ("content-type".to_string(), "application/json".to_string()), + ( + "anthropic-beta".to_string(), + CLAUDE_CODE_OAUTH_BETA.to_string(), + ), + // The reset-credit (`cedar_ember`) block is only returned to first-party CLI callers. + ("x-app".to_string(), "cli".to_string()), + ( + "user-agent".to_string(), + CLAUDE_CODE_USAGE_USER_AGENT.to_string(), + ), + ]); + + ProviderPoolQuotaRequestSpec { + request_id: format!("claude-code-quota:{key_id}"), + provider_name: "claude_code".to_string(), + quota_kind: "claude_code".to_string(), + method: "GET".to_string(), + url: CLAUDE_CODE_OAUTH_USAGE_URL.to_string(), + headers, + content_type: None, + json_body: None, + client_api_format: "claude:messages".to_string(), + provider_api_format: "claude_code:oauth_usage".to_string(), + model_name: Some("oauth_usage".to_string()), + } +} diff --git a/crates/aether-provider/pool/src/providers/codex.rs b/crates/aether-provider/pool/src/providers/codex.rs index f7133d281..86145f350 100644 --- a/crates/aether-provider/pool/src/providers/codex.rs +++ b/crates/aether-provider/pool/src/providers/codex.rs @@ -10,10 +10,11 @@ use crate::provider::{ ProviderPoolMemberInput, }; use crate::quota::{ - provider_pool_current_unix_secs, provider_pool_json_bool, provider_pool_json_f64, - provider_pool_member_quota_snapshot, provider_pool_metadata_bucket, - provider_pool_model_quota_exhausted, provider_pool_quota_snapshot_exhausted_decision, - provider_pool_reset_deadline_elapsed, provider_pool_timestamp_unix_secs, + provider_pool_codex_metadata_has_account_quota, provider_pool_current_unix_secs, + provider_pool_json_bool, provider_pool_json_f64, provider_pool_member_quota_snapshot, + provider_pool_metadata_bucket, provider_pool_model_quota_exhausted, + provider_pool_quota_snapshot_exhausted_decision, provider_pool_reset_deadline_elapsed, + provider_pool_source_account_quota_exhausted, provider_pool_timestamp_unix_secs, }; use crate::quota_refresh::ProviderPoolQuotaRequestSpec; @@ -54,6 +55,9 @@ impl ProviderPoolAdapter for CodexProviderPoolAdapter { }) { return exhausted; } + if let Some(bucket) = codex_newer_account_quota_metadata(input.key, input.provider_type) { + return quota_exhausted_from_bucket(bucket); + } if let Some(quota_snapshot) = provider_pool_member_quota_snapshot(input.key, input.provider_type) { @@ -108,15 +112,12 @@ fn codex_explicit_quota_block_active( key: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey, provider_type: &str, ) -> bool { + if let Some(bucket) = codex_newer_account_quota_metadata(key, provider_type) { + return codex_explicit_quota_block_from_bucket(bucket); + } let Some(quota_snapshot) = provider_pool_member_quota_snapshot(key, provider_type) else { return provider_pool_metadata_bucket(key.upstream_metadata.as_ref(), provider_type) - .is_some_and(|bucket| { - (provider_pool_json_bool(bucket.get("allowed")) == Some(false) - || provider_pool_json_bool(bucket.get("limit_reached")) == Some(true)) - && !["primary", "secondary"] - .into_iter() - .any(|prefix| codex_window_reset_elapsed(bucket, prefix)) - }); + .is_some_and(codex_explicit_quota_block_from_bucket); }; let explicitly_blocked = provider_pool_json_bool(quota_snapshot.get("allowed")) == Some(false) || provider_pool_json_bool(quota_snapshot.get("limit_reached")) == Some(true); @@ -133,6 +134,35 @@ fn codex_explicit_quota_block_active( }) } +fn codex_explicit_quota_block_from_bucket(bucket: &Map) -> bool { + (provider_pool_json_bool(bucket.get("allowed")) == Some(false) + || provider_pool_json_bool(bucket.get("limit_reached")) == Some(true)) + && !["primary", "secondary"] + .into_iter() + .any(|prefix| codex_window_reset_elapsed(bucket, prefix)) +} + +/// A successful refresh can update raw metadata before the status snapshot. +/// Do not keep an old account-level block once a newer quota observation exists. +/// Identity-only or model-only updates cannot clear an account quota decision. +fn codex_newer_account_quota_metadata<'a>( + key: &'a aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey, + provider_type: &str, +) -> Option<&'a Map> { + let snapshot = provider_pool_member_quota_snapshot(key, provider_type)?; + let metadata = provider_pool_metadata_bucket(key.upstream_metadata.as_ref(), provider_type)?; + if !provider_pool_codex_metadata_has_account_quota(metadata) { + return None; + } + let metadata_observed_at = provider_pool_timestamp_unix_secs(metadata.get("observed_at")) + .or_else(|| provider_pool_timestamp_unix_secs(metadata.get("updated_at")))?; + let snapshot_observed_at = provider_pool_timestamp_unix_secs(snapshot.get("observed_at")) + .or_else(|| provider_pool_timestamp_unix_secs(snapshot.get("updated_at"))); + snapshot_observed_at + .is_none_or(|observed_at| metadata_observed_at >= observed_at) + .then_some(metadata) +} + fn build_codex_wham_headers( resolved_oauth_auth: Option<(String, String)>, decrypted_api_key: Option<&str>, @@ -322,11 +352,14 @@ pub(crate) fn quota_exhausted_from_bucket(bucket: &Map) -> bool { if provider_pool_json_bool(bucket.get("credits_unlimited")) == Some(true) { return false; } + let account_windows_exhausted = provider_pool_source_account_quota_exhausted(bucket); let has_window_data = provider_pool_json_f64(bucket.get("primary_used_percent")).is_some() - || provider_pool_json_f64(bucket.get("secondary_used_percent")).is_some(); + || provider_pool_json_f64(bucket.get("secondary_used_percent")).is_some() + || account_windows_exhausted.is_some(); if !has_window_data && provider_pool_json_bool(bucket.get("has_credits")) == Some(false) { return true; } - codex_window_used_percent_exhausted(bucket, "primary") + account_windows_exhausted == Some(true) + || codex_window_used_percent_exhausted(bucket, "primary") || codex_window_used_percent_exhausted(bucket, "secondary") } diff --git a/crates/aether-provider/pool/src/providers/mod.rs b/crates/aether-provider/pool/src/providers/mod.rs index 72734c81e..fe166a912 100644 --- a/crates/aether-provider/pool/src/providers/mod.rs +++ b/crates/aether-provider/pool/src/providers/mod.rs @@ -1,5 +1,6 @@ pub mod antigravity; pub mod chatgpt_web; +pub mod claude_code; pub mod codex; pub mod default; pub mod gemini_cli; @@ -7,6 +8,7 @@ pub mod grok; pub mod kiro; pub mod unsupported; pub mod windsurf; +pub mod xai; pub use antigravity::AntigravityProviderPoolAdapter; pub use antigravity::{ @@ -19,6 +21,11 @@ pub use chatgpt_web::{ normalize_chatgpt_web_image_quota_limit, CHATGPT_WEB_CONVERSATION_INIT_PATH, CHATGPT_WEB_DEFAULT_BASE_URL, }; +pub use claude_code::ClaudeCodeProviderPoolAdapter; +pub use claude_code::{ + build_claude_code_pool_quota_request, CLAUDE_CODE_OAUTH_BETA, CLAUDE_CODE_OAUTH_USAGE_URL, + CLAUDE_CODE_USAGE_USER_AGENT, +}; pub use codex::CodexProviderPoolAdapter; pub use codex::{ build_codex_pool_quota_request, build_codex_pool_reset_credit_consume_request, @@ -39,10 +46,7 @@ pub use kiro::{ build_kiro_pool_quota_request, KiroPoolQuotaAuthInput, KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION, }; -pub use unsupported::{ - UnsupportedQuotaProviderPoolAdapter, CLAUDE_CODE_PROVIDER_POOL_ADAPTER, - VERTEX_AI_PROVIDER_POOL_ADAPTER, -}; +pub use unsupported::{UnsupportedQuotaProviderPoolAdapter, VERTEX_AI_PROVIDER_POOL_ADAPTER}; pub use windsurf::{ build_windsurf_pool_model_configs_request, build_windsurf_pool_model_configs_request_with_base_url, build_windsurf_pool_quota_request, @@ -51,3 +55,7 @@ pub use windsurf::{ WINDSURF_DEFAULT_BASE_URL, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, WINDSURF_USER_STATUS_PATH, }; +pub use xai::{ + build_xai_pool_billing_request, build_xai_pool_user_request, XaiProviderPoolAdapter, + XAI_BILLING_PATH, XAI_USER_PATH, +}; diff --git a/crates/aether-provider/pool/src/providers/unsupported.rs b/crates/aether-provider/pool/src/providers/unsupported.rs index 0a9dbaf1f..5b7fabb05 100644 --- a/crates/aether-provider/pool/src/providers/unsupported.rs +++ b/crates/aether-provider/pool/src/providers/unsupported.rs @@ -28,12 +28,6 @@ impl ProviderPoolAdapter for UnsupportedQuotaProviderPoolAdapter { } } -pub const CLAUDE_CODE_PROVIDER_POOL_ADAPTER: UnsupportedQuotaProviderPoolAdapter = - UnsupportedQuotaProviderPoolAdapter::new( - "claude_code", - "Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口", - ); - pub const VERTEX_AI_PROVIDER_POOL_ADAPTER: UnsupportedQuotaProviderPoolAdapter = UnsupportedQuotaProviderPoolAdapter::new( "vertex_ai", diff --git a/crates/aether-provider/pool/src/providers/xai.rs b/crates/aether-provider/pool/src/providers/xai.rs new file mode 100644 index 000000000..4817c7d16 --- /dev/null +++ b/crates/aether-provider/pool/src/providers/xai.rs @@ -0,0 +1,255 @@ +use std::collections::BTreeMap; + +use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint; +use aether_provider_transport::xai::{ + insert_cli_identity_headers, XAI_CHAT_PROXY_BASE_URL, XAI_PROVIDER_TYPE, +}; +use serde_json::{Map, Value}; + +use crate::capability::ProviderPoolCapabilities; +use crate::provider::{ + provider_pool_endpoint_format_matches, provider_pool_matching_endpoint, ProviderPoolAdapter, + ProviderPoolMemberInput, +}; +use crate::quota::{ + provider_pool_current_unix_secs, provider_pool_json_bool, provider_pool_json_f64, + provider_pool_metadata_bucket, provider_pool_model_quota_exhausted, + provider_pool_quota_snapshot_exhausted_decision, provider_pool_reset_deadline_elapsed, + provider_pool_timestamp_unix_secs, +}; +use crate::quota_refresh::ProviderPoolQuotaRequestSpec; + +pub const XAI_USER_PATH: &str = "/user"; +pub const XAI_BILLING_PATH: &str = "/billing?format=credits"; + +#[derive(Debug, Clone, Default)] +pub struct XaiProviderPoolAdapter; + +impl ProviderPoolAdapter for XaiProviderPoolAdapter { + fn provider_type(&self) -> &'static str { + XAI_PROVIDER_TYPE + } + + fn capabilities(&self) -> ProviderPoolCapabilities { + ProviderPoolCapabilities { + plan_tier: true, + quota_reset: true, + quota_refresh: true, + } + } + + fn quota_exhausted(&self, input: &ProviderPoolMemberInput<'_>) -> bool { + if let Some(exhausted) = input.provider_model_name.and_then(|model| { + provider_pool_model_quota_exhausted(input.key, input.provider_type, model) + }) { + return exhausted; + } + if let Some(exhausted) = + provider_pool_quota_snapshot_exhausted_decision(input.key, input.provider_type) + { + return exhausted; + } + provider_pool_metadata_bucket(input.key.upstream_metadata.as_ref(), input.provider_type) + .is_some_and(quota_exhausted_from_bucket) + } + + fn quota_refresh_endpoint( + &self, + endpoints: &[StoredProviderCatalogEndpoint], + include_inactive: bool, + ) -> Option { + provider_pool_matching_endpoint(endpoints, include_inactive, |endpoint| { + provider_pool_endpoint_format_matches(endpoint, "openai:responses") + }) + .or_else(|| provider_pool_matching_endpoint(endpoints, include_inactive, |_| true)) + } + + fn quota_refresh_missing_endpoint_message(&self) -> String { + "找不到有效的 openai:responses 端点".to_string() + } +} + +pub fn build_xai_pool_user_request( + key_id: &str, + authorization: (String, String), +) -> ProviderPoolQuotaRequestSpec { + build_xai_pool_request( + format!("xai-user:{key_id}"), + "xai:user", + "user", + XAI_USER_PATH, + authorization, + None, + ) +} + +pub fn build_xai_pool_billing_request( + key_id: &str, + authorization: (String, String), + user_id: Option<&str>, +) -> ProviderPoolQuotaRequestSpec { + build_xai_pool_request( + format!("xai-billing:{key_id}"), + "xai:billing", + "billing", + XAI_BILLING_PATH, + authorization, + user_id, + ) +} + +fn build_xai_pool_request( + request_id: String, + provider_api_format: &str, + model_name: &str, + path: &str, + authorization: (String, String), + user_id: Option<&str>, +) -> ProviderPoolQuotaRequestSpec { + let mut headers = BTreeMap::from([ + (authorization.0, authorization.1), + ("accept".to_string(), "application/json".to_string()), + ]); + insert_cli_identity_headers(&mut headers); + if let Some(user_id) = user_id.map(str::trim).filter(|value| !value.is_empty()) { + headers.insert("x-userid".to_string(), user_id.to_string()); + } + + ProviderPoolQuotaRequestSpec { + request_id, + provider_name: XAI_PROVIDER_TYPE.to_string(), + quota_kind: XAI_PROVIDER_TYPE.to_string(), + method: "GET".to_string(), + url: format!("{}{path}", XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/')), + headers, + content_type: None, + json_body: None, + client_api_format: "openai:responses".to_string(), + provider_api_format: provider_api_format.to_string(), + model_name: Some(model_name.to_string()), + } +} + +pub(crate) fn quota_exhausted_from_bucket(bucket: &Map) -> bool { + if provider_pool_current_unix_secs().is_some_and(|now| { + provider_pool_reset_deadline_elapsed( + bucket, + provider_pool_timestamp_unix_secs(bucket.get("updated_at")), + now, + ) + }) { + return false; + } + + let usage_exhausted = provider_pool_json_f64(bucket.get("remaining")) + .is_some_and(|value| value <= 0.0) + || provider_pool_json_f64(bucket.get("usage_percentage")) + .is_some_and(|value| value >= 100.0 - 1e-6) + || match ( + provider_pool_json_f64(bucket.get("usage_limit")), + provider_pool_json_f64(bucket.get("current_usage")), + ) { + (Some(limit), Some(current)) if limit > 0.0 => current >= limit, + _ => false, + }; + if !usage_exhausted { + return false; + } + + let prepaid_available = + provider_pool_json_f64(bucket.get("prepaid_balance")).is_some_and(|value| value > 0.0); + if prepaid_available { + return false; + } + + let on_demand_enabled = provider_pool_json_bool(bucket.get("on_demand_enabled")) != Some(false); + let on_demand_cap = provider_pool_json_f64(bucket.get("on_demand_cap")).unwrap_or(0.0); + let on_demand_used = provider_pool_json_f64(bucket.get("on_demand_used")).unwrap_or(0.0); + if on_demand_enabled && on_demand_cap > 0.0 && on_demand_used < on_demand_cap { + return false; + } + + true +} + +#[cfg(test)] +mod tests { + use super::{ + build_xai_pool_billing_request, build_xai_pool_user_request, quota_exhausted_from_bucket, + }; + use aether_provider_transport::xai::{ + XAI_CHAT_PROXY_BASE_URL, XAI_CLIENT_IDENTIFIER_VALUE, XAI_TOKEN_AUTH_VALUE, + }; + use serde_json::{json, Map}; + + fn bucket(value: serde_json::Value) -> Map { + value.as_object().cloned().expect("bucket should be object") + } + + #[test] + fn user_and_billing_requests_pin_cli_chat_proxy_and_identity_headers() { + let authorization = ("authorization".to_string(), "Bearer xai-access".to_string()); + let user = build_xai_pool_user_request("key-1", authorization.clone()); + let billing = build_xai_pool_billing_request("key-1", authorization, Some("user-42")); + + assert_eq!( + user.url, + format!("{}/user", XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/')) + ); + assert_eq!( + billing.url, + format!( + "{}/billing?format=credits", + XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/') + ) + ); + assert_eq!( + user.headers.get("x-xai-token-auth").map(String::as_str), + Some(XAI_TOKEN_AUTH_VALUE) + ); + assert_eq!( + user.headers + .get("x-grok-client-identifier") + .map(String::as_str), + Some(XAI_CLIENT_IDENTIFIER_VALUE) + ); + assert!(!user.headers.contains_key("x-userid")); + assert_eq!( + billing.headers.get("x-userid").map(String::as_str), + Some("user-42") + ); + assert_eq!( + billing.headers.get("authorization").map(String::as_str), + Some("Bearer xai-access") + ); + } + + #[test] + fn percent_exhausted_without_prepaid_or_on_demand_is_exhausted() { + assert!(quota_exhausted_from_bucket(&bucket(json!({ + "usage_percentage": 100.0, + "prepaid_balance": 0.0, + "on_demand_cap": 0.0, + "on_demand_used": 0.0 + })))); + } + + #[test] + fn unified_billing_zero_on_demand_cap_is_not_exhausted_when_percent_remains() { + assert!(!quota_exhausted_from_bucket(&bucket(json!({ + "usage_percentage": 46.0, + "prepaid_balance": 0.0, + "on_demand_cap": 0.0, + "on_demand_used": 0.0 + })))); + } + + #[test] + fn prepaid_balance_keeps_account_available_after_weekly_pool_hits_100() { + assert!(!quota_exhausted_from_bucket(&bucket(json!({ + "usage_percentage": 100.0, + "prepaid_balance": 12.5, + "on_demand_cap": 0.0 + })))); + } +} diff --git a/crates/aether-provider/pool/src/quota.rs b/crates/aether-provider/pool/src/quota.rs index 9d18d506c..af49f3f63 100644 --- a/crates/aether-provider/pool/src/quota.rs +++ b/crates/aether-provider/pool/src/quota.rs @@ -78,8 +78,18 @@ pub(crate) fn provider_pool_model_quota_exhausted( provider_type: &str, provider_model_name: &str, ) -> Option { - let requested = provider_pool_identifier_tokens(provider_model_name); - if requested.is_empty() { + provider_pool_quota_reaches_reserve(key, provider_type, Some(provider_model_name), 0.0) +} + +fn provider_pool_quota_reaches_reserve( + key: &StoredProviderCatalogKey, + provider_type: &str, + provider_model_name: Option<&str>, + reserve_ratio: f64, +) -> Option { + let is_codex = provider_type.trim().eq_ignore_ascii_case("codex"); + let requested = provider_model_name.map(provider_pool_identifier_tokens); + if reserve_ratio <= 0.0 && requested.as_ref().is_some_and(|tokens| tokens.is_empty()) { return None; } @@ -92,29 +102,84 @@ pub(crate) fn provider_pool_model_quota_exhausted( provider_pool_metadata_bucket(key.upstream_metadata.as_ref(), provider_type), ]; let mut resolved = None::<(Option, bool)>; + let mut resolved_specific_bucket = false; for source in sources.into_iter().flatten() { - let windows = provider_pool_collect_quota_windows(source); + let observed_at = provider_pool_timestamp_unix_secs(source.get("observed_at")) + .or_else(|| provider_pool_timestamp_unix_secs(source.get("updated_at"))); + let account_signal = (is_codex && reserve_ratio <= 0.0) + .then(|| provider_pool_codex_account_quota_signal(source, observed_at)) + .flatten(); + let mut windows = provider_pool_collect_quota_windows(source); + if is_codex { + // Raw refresh metadata can be newer than the materialized windows. + // Include all four legacy slots before selecting the model bucket. + for prefix in ["primary", "secondary", "spark_primary", "spark_secondary"] { + let Some(used_percent) = + provider_pool_json_f64(source.get(&format!("{prefix}_used_percent"))) + else { + continue; + }; + if provider_pool_json_f64(source.get(&format!("{prefix}_window_minutes"))) + == Some(0.0) + { + continue; + } + let mut window = Map::from_iter([ + ("code".to_string(), json!(prefix)), + ("used_percent".to_string(), json!(used_percent)), + ]); + for field in [ + "reset_at", + "next_reset_at", + "reset_seconds", + "reset_after_seconds", + ] { + if let Some(value) = source.get(&format!("{prefix}_{field}")) { + window.insert(field.to_string(), value.clone()); + } + } + windows.push(window); + } + windows.retain(provider_pool_quota_window_has_observation); + if let Some(exhausted) = account_signal { + if !windows.iter().any(provider_pool_window_is_generic) { + // A flags-only quota refresh is still a newer account + // observation, but must not override a model-specific bucket. + windows.push(Map::from_iter([ + ("code".to_string(), json!("account")), + ("is_exhausted".to_string(), json!(exhausted)), + ])); + } + } + } if windows.is_empty() { continue; } - let observed_at = provider_pool_timestamp_unix_secs(source.get("observed_at")) - .or_else(|| provider_pool_timestamp_unix_secs(source.get("updated_at"))); let model_matches = windows .iter() .filter(|window| { - provider_pool_window_explicitly_matches_model(window, provider_model_name) + provider_model_name.is_some_and(|model| { + provider_pool_window_explicitly_matches_model(window, model) + }) }) .collect::>(); if !model_matches.is_empty() { - let exhausted = - provider_pool_explicit_model_windows_exhausted(model_matches, observed_at); + let exhausted = if reserve_ratio > 0.0 { + // Reserving quota protects every applicable rate-limit window, + // including short and weekly limits for an explicit model. + provider_pool_any_window_exhausted(model_matches, observed_at, reserve_ratio) + } else { + provider_pool_explicit_model_windows_exhausted(model_matches, observed_at) + }; if resolved.is_none() + || (is_codex && !resolved_specific_bucket) || provider_pool_should_replace_model_quota_resolution( resolved.as_ref().and_then(|(observed_at, _)| *observed_at), observed_at, ) { resolved = Some((observed_at, exhausted)); + resolved_specific_bucket = true; } continue; } @@ -125,21 +190,34 @@ pub(crate) fn provider_pool_model_quota_exhausted( // windows isolated without baking in names such as "spark". let family_matches = windows .iter() - .filter(|window| provider_pool_window_family_matches_model(window, &requested)) + .filter(|window| { + requested + .as_ref() + .is_some_and(|tokens| provider_pool_window_family_matches_model(window, tokens)) + }) .collect::>(); if !family_matches.is_empty() { - let exhausted = provider_pool_any_window_exhausted(family_matches, observed_at); + let exhausted = + provider_pool_any_window_exhausted(family_matches, observed_at, reserve_ratio); if resolved.is_none() + || (is_codex && !resolved_specific_bucket) || provider_pool_should_replace_model_quota_resolution( resolved.as_ref().and_then(|(observed_at, _)| *observed_at), observed_at, ) { resolved = Some((observed_at, exhausted)); + resolved_specific_bucket = true; } continue; } + // A newer account observation does not update an independent model + // bucket. Only compare freshness between applicable model sources. + if is_codex && resolved_specific_bucket { + continue; + } + // Account-scoped windows (for example the ordinary weekly and short // windows emitted by Codex) apply to every model that has no more // specific family. Restrict this fallback to well-known structural @@ -150,7 +228,9 @@ pub(crate) fn provider_pool_model_quota_exhausted( .filter(|window| provider_pool_window_is_generic(window)) .collect::>(); if !generic_matches.is_empty() { - let exhausted = provider_pool_any_window_exhausted(generic_matches, observed_at); + let exhausted = account_signal.unwrap_or_else(|| { + provider_pool_any_window_exhausted(generic_matches, observed_at, reserve_ratio) + }); if resolved.is_none() || provider_pool_should_replace_model_quota_resolution( resolved.as_ref().and_then(|(observed_at, _)| *observed_at), @@ -165,6 +245,37 @@ pub(crate) fn provider_pool_model_quota_exhausted( resolved.map(|(_, exhausted)| exhausted) } +fn provider_pool_codex_account_quota_signal( + source: &Map, + observed_at: Option, +) -> Option { + let allowed = provider_pool_json_bool(source.get("allowed")); + let limit_reached = provider_pool_json_bool(source.get("limit_reached")); + if allowed == Some(false) || limit_reached == Some(true) { + let reset_elapsed = provider_pool_current_unix_secs().is_some_and(|now| { + provider_pool_reset_deadline_elapsed(source, observed_at, now) + || ["primary", "secondary"].into_iter().any(|prefix| { + let window = [ + "reset_at", + "next_reset_at", + "reset_seconds", + "reset_after_seconds", + ] + .into_iter() + .filter_map(|field| { + source + .get(&format!("{prefix}_{field}")) + .map(|value| (field.to_string(), value.clone())) + }) + .collect::>(); + provider_pool_reset_deadline_elapsed(&window, observed_at, now) + }) + }); + return (!reset_elapsed).then_some(true); + } + (allowed == Some(true) || limit_reached == Some(false)).then_some(false) +} + fn provider_pool_should_replace_model_quota_resolution( previous_observed_at: Option, next_observed_at: Option, @@ -209,16 +320,132 @@ fn provider_pool_explicit_model_windows_exhausted( fn provider_pool_any_window_exhausted( windows: Vec<&Map>, snapshot_observed_at: Option, + reserve_ratio: f64, ) -> bool { let now_unix_secs = provider_pool_current_unix_secs(); windows.iter().any(|window| { - provider_pool_quota_window_is_exhausted(window) + provider_pool_quota_window_reaches_reserve(window, reserve_ratio) && !now_unix_secs.is_some_and(|now| { provider_pool_reset_deadline_elapsed(window, snapshot_observed_at, now) }) }) } +fn provider_pool_quota_window_has_observation(window: &Map) -> bool { + ["is_exhausted", "exhausted"] + .into_iter() + .any(|field| provider_pool_json_bool(window.get(field)).is_some()) + || [ + "used_ratio", + "usage_ratio", + "used_percent", + "remaining_ratio", + "remaining_fraction", + "remaining_percent", + ] + .into_iter() + .any(|field| provider_pool_json_f64(window.get(field)).is_some()) + || (provider_pool_json_f64( + window + .get("remaining") + .or_else(|| window.get("remaining_value")), + ) + .is_some() + && provider_pool_json_f64( + window + .get("limit") + .or_else(|| window.get("limit_value")) + .or_else(|| window.get("total")), + ) + .is_some_and(|limit| limit > 0.0)) +} + +pub(crate) fn provider_pool_source_account_quota_exhausted( + source: &Map, +) -> Option { + let windows = provider_pool_collect_quota_windows(source); + let account_windows = windows + .iter() + .filter(|window| provider_pool_window_is_generic(window)) + .filter(|window| provider_pool_quota_window_has_observation(window)) + .collect::>(); + if account_windows.is_empty() { + return None; + } + let observed_at = provider_pool_timestamp_unix_secs(source.get("observed_at")) + .or_else(|| provider_pool_timestamp_unix_secs(source.get("updated_at"))); + Some(provider_pool_any_window_exhausted( + account_windows, + observed_at, + 0.0, + )) +} + +/// Whether Codex metadata contains an account quota observation rather than an +/// identity-only update or an independent model quota bucket. +pub fn provider_pool_codex_metadata_has_account_quota(source: &Map) -> bool { + ["primary_used_percent", "secondary_used_percent"] + .into_iter() + .any(|field| provider_pool_json_f64(source.get(field)).is_some()) + || provider_pool_source_account_quota_exhausted(source).is_some() + || ["allowed", "limit_reached", "has_credits"] + .into_iter() + .any(|field| provider_pool_json_bool(source.get(field)).is_some()) + || provider_pool_json_bool(source.get("credits_unlimited")) == Some(true) +} + +/// Whether an applicable Codex window has at most 1% remaining. Missing quota +/// data and windows whose reset has elapsed do not trigger this opt-in guard. +pub fn provider_pool_key_minimum_quota_reached( + key: &StoredProviderCatalogKey, + provider_type: &str, + provider_model_name: Option<&str>, +) -> bool { + if !provider_type.trim().eq_ignore_ascii_case("codex") { + return false; + } + provider_pool_quota_reaches_reserve(key, provider_type, provider_model_name, 0.01) + .unwrap_or(false) +} + +fn provider_pool_quota_window_reaches_reserve( + window: &Map, + reserve_ratio: f64, +) -> bool { + if provider_pool_quota_window_is_exhausted(window) { + return true; + } + if reserve_ratio <= 0.0 { + return false; + } + let used_ratio = provider_pool_json_f64(window.get("used_ratio")) + .or_else(|| provider_pool_json_f64(window.get("usage_ratio"))) + .or_else(|| provider_pool_json_f64(window.get("used_percent")).map(|value| value / 100.0)) + .or_else(|| { + provider_pool_json_f64(window.get("remaining_ratio")) + .or_else(|| provider_pool_json_f64(window.get("remaining_fraction"))) + .map(|value| 1.0 - value) + }) + .or_else(|| { + provider_pool_json_f64(window.get("remaining_percent")).map(|value| 1.0 - value / 100.0) + }) + .or_else(|| { + let remaining = provider_pool_json_f64( + window + .get("remaining") + .or_else(|| window.get("remaining_value")), + )?; + let limit = provider_pool_json_f64( + window + .get("limit") + .or_else(|| window.get("limit_value")) + .or_else(|| window.get("total")), + )?; + (limit > 0.0).then_some(1.0 - remaining / limit) + }); + used_ratio.is_some_and(|used_ratio| used_ratio >= 1.0 - reserve_ratio) +} + fn provider_pool_window_is_generic(window: &Map) -> bool { if provider_pool_window_has_explicit_model(window) || window @@ -585,7 +812,16 @@ pub(crate) fn provider_pool_json_f64(value: Option<&Value>) -> Option { } pub(crate) fn provider_pool_timestamp_unix_secs(value: Option<&Value>) -> Option { - let mut timestamp = provider_pool_json_f64(value)?; + let mut timestamp = match provider_pool_json_f64(value) { + Some(timestamp) => timestamp, + None => { + return chrono::DateTime::parse_from_rfc3339(value?.as_str()?.trim()) + .ok()? + .timestamp() + .try_into() + .ok(); + } + }; if timestamp <= 0.0 { return None; } diff --git a/crates/aether-provider/pool/src/service.rs b/crates/aether-provider/pool/src/service.rs index f61b07a7f..f86d68a13 100644 --- a/crates/aether-provider/pool/src/service.rs +++ b/crates/aether-provider/pool/src/service.rs @@ -11,10 +11,10 @@ use crate::capability::ProviderPoolCapability; use crate::presets::normalize_provider_scheduling_presets; use crate::provider::{ProviderPoolAdapter, ProviderPoolMemberInput}; use crate::providers::{ - AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter, - DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter, - KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, CLAUDE_CODE_PROVIDER_POOL_ADAPTER, - VERTEX_AI_PROVIDER_POOL_ADAPTER, + AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, ClaudeCodeProviderPoolAdapter, + CodexProviderPoolAdapter, DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, + GrokProviderPoolAdapter, KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, + XaiProviderPoolAdapter, VERTEX_AI_PROVIDER_POOL_ADAPTER, }; #[derive(Clone)] @@ -48,13 +48,14 @@ impl ProviderPoolService { pub fn with_builtin_adapters() -> Self { Self::new() .with_adapter(Arc::new(AntigravityProviderPoolAdapter)) - .with_adapter(Arc::new(CLAUDE_CODE_PROVIDER_POOL_ADAPTER)) + .with_adapter(Arc::new(ClaudeCodeProviderPoolAdapter)) .with_adapter(Arc::new(CodexProviderPoolAdapter)) .with_adapter(Arc::new(GeminiCliProviderPoolAdapter)) .with_adapter(Arc::new(GrokProviderPoolAdapter)) .with_adapter(Arc::new(KiroProviderPoolAdapter)) .with_adapter(Arc::new(ChatGptWebProviderPoolAdapter)) .with_adapter(Arc::new(WindsurfProviderPoolAdapter)) + .with_adapter(Arc::new(XaiProviderPoolAdapter)) .with_adapter(Arc::new(VERTEX_AI_PROVIDER_POOL_ADAPTER)) } diff --git a/crates/aether-provider/transport/src/agent_identity.rs b/crates/aether-provider/transport/src/agent_identity.rs index abdb9d806..d7b7ee3ac 100644 --- a/crates/aether-provider/transport/src/agent_identity.rs +++ b/crates/aether-provider/transport/src/agent_identity.rs @@ -780,11 +780,11 @@ async fn register_codex_agent_identity_from_access_token_with_auth_api_base_url( ), ( "user-agent".to_string(), - aether_ai_formats::CODEX_CLIENT_USER_AGENT.to_string(), + aether_ai_formats::codex_client_user_agent(), ), ( "originator".to_string(), - aether_ai_formats::CODEX_CLIENT_ORIGINATOR.to_string(), + aether_ai_formats::codex_client_originator(), ), ]); if options.is_fedramp_account { @@ -799,7 +799,7 @@ async fn register_codex_agent_identity_from_access_token_with_auth_api_base_url( content_type: Some("application/json".to_string()), json_body: Some(json!({ "abom": { - "agent_version": aether_ai_formats::CODEX_CLIENT_VERSION, + "agent_version": aether_ai_formats::codex_client_version(), "agent_harness_id": CODEX_AGENT_IDENTITY_AGENT_HARNESS_ID, "running_location": format!("cli-{}", std::env::consts::OS), }, diff --git a/crates/aether-provider/transport/src/antigravity/fabric_exec_schema.json b/crates/aether-provider/transport/src/antigravity/fabric_exec_schema.json new file mode 100644 index 000000000..1b98d673e --- /dev/null +++ b/crates/aether-provider/transport/src/antigravity/fabric_exec_schema.json @@ -0,0 +1,24 @@ +{ + "type": "object", + "required": ["code"], + "properties": { + "code": {"type": "string", "description": "TypeScript function body."}, + "payloads": {"type": "object", "patternProperties": {"^.*$": {"type": "string"}}}, + "resultFormat": {"anyOf": [ + {"type": "string", "const": "auto"}, + {"type": "string", "const": "yaml"}, + {"type": "string", "const": "json"}, + {"type": "string", "const": "text"} + ]}, + "tokenBudget": {"type": "number", "minimum": 1}, + "agentBudget": {"type": "number", "minimum": 1}, + "timeoutMs": {"type": "number", "minimum": 1}, + "display": {"anyOf": [ + {"type": "object", "properties": { + "name": {"type": "string"}, + "description": {"type": "string"} + }}, + {"type": "string"} + ]} + } +} diff --git a/crates/aether-provider/transport/src/antigravity/mod.rs b/crates/aether-provider/transport/src/antigravity/mod.rs index 4d09239e2..ac1c91ffb 100644 --- a/crates/aether-provider/transport/src/antigravity/mod.rs +++ b/crates/aether-provider/transport/src/antigravity/mod.rs @@ -1,6 +1,7 @@ mod auth; mod policy; mod request; +mod schema; mod url; pub use auth::{ diff --git a/crates/aether-provider/transport/src/antigravity/request.rs b/crates/aether-provider/transport/src/antigravity/request.rs index b0ddc5b9b..2f009b7f3 100644 --- a/crates/aether-provider/transport/src/antigravity/request.rs +++ b/crates/aether-provider/transport/src/antigravity/request.rs @@ -1,6 +1,7 @@ use serde_json::{Map, Value}; use super::auth::{AntigravityRequestAuth, ANTIGRAVITY_REQUEST_USER_AGENT}; +use super::schema::{normalize_claude_unions, normalize_tool_parameters, SchemaBudget}; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum AntigravityEnvelopeRequestType { @@ -31,6 +32,7 @@ pub enum AntigravityRequestEnvelopeUnsupportedReason { MissingContents, MissingRequestId, MissingModel, + ToolSchemaBudgetExceeded, } pub fn classify_antigravity_safe_request_body( @@ -78,46 +80,71 @@ pub fn build_antigravity_safe_v1internal_request( inner_request.remove("model"); inner_request.remove("safetySettings"); inner_request.remove("safety_settings"); + normalize_antigravity_claude_thought_history(&mut inner_request, model); normalize_antigravity_builtin_tool_names(&mut inner_request); - normalize_antigravity_function_declaration_parameters(&mut inner_request); + if normalize_antigravity_function_declaration_parameters(&mut inner_request, model).is_err() + { + return AntigravityRequestEnvelopeSupport::Unsupported( + AntigravityRequestEnvelopeUnsupportedReason::ToolSchemaBudgetExceeded, + ); + } let request_id = non_empty_string_field(source, "requestId").unwrap_or(request_id); let user_agent = non_empty_string_field(source, "userAgent").unwrap_or(ANTIGRAVITY_REQUEST_USER_AGENT); - let request_type = - existing_v1internal_request_type(source).unwrap_or_else(|| request_type.as_str()); + let existing_request_type = existing_v1internal_request_type(source); - return AntigravityRequestEnvelopeSupport::Supported(serde_json::json!({ + let mut envelope = serde_json::json!({ "project": auth.project_id, "requestId": request_id, "request": Value::Object(inner_request), "model": model, "userAgent": user_agent, - "requestType": request_type, - })); + }); + if let Some(existing_request_type) = existing_request_type { + envelope["requestType"] = Value::String(existing_request_type.to_string()); + } else if request_type != AntigravityEnvelopeRequestType::Agent { + envelope["requestType"] = Value::String(request_type.as_str().to_string()); + } + return AntigravityRequestEnvelopeSupport::Supported(envelope); } let mut inner_request: Map = source.clone(); inner_request.remove("model"); inner_request.remove("safetySettings"); inner_request.remove("safety_settings"); + normalize_antigravity_claude_thought_history(&mut inner_request, model); normalize_antigravity_builtin_tool_names(&mut inner_request); - normalize_antigravity_function_declaration_parameters(&mut inner_request); + if normalize_antigravity_function_declaration_parameters(&mut inner_request, model).is_err() { + return AntigravityRequestEnvelopeSupport::Unsupported( + AntigravityRequestEnvelopeUnsupportedReason::ToolSchemaBudgetExceeded, + ); + } - AntigravityRequestEnvelopeSupport::Supported(serde_json::json!({ + let mut envelope = serde_json::json!({ "project": auth.project_id, "requestId": request_id, "request": Value::Object(inner_request), "model": model, "userAgent": ANTIGRAVITY_REQUEST_USER_AGENT, - "requestType": request_type.as_str(), - })) + }); + if request_type != AntigravityEnvelopeRequestType::Agent { + envelope["requestType"] = Value::String(request_type.as_str().to_string()); + } + AntigravityRequestEnvelopeSupport::Supported(envelope) } -/// Antigravity's private v1internal Gemini surface still uses the legacy -/// `googleSearchRetrieval` spelling. The public Gemini converter emits the -/// newer `googleSearch` spelling, which the private backend rejects when it is -/// combined with function declarations. Normalize only at this transport -/// boundary so public Gemini requests retain their native shape. +/// Antigravity's private v1internal Gemini surface takes the same +/// `googleSearch` grounding tool as the public one. Only the snake_case alias +/// needs folding into the canonical camelCase key. +/// +/// This used to rewrite `googleSearch` into the Gemini 1.5-era +/// `googleSearchRetrieval` spelling. Gemini 3 rejects that: the model emits a +/// `google_search` call the backend cannot bind to any declared tool, and the +/// turn dies with `MALFORMED_FUNCTION_CALL`, e.g. +/// `Malformed function call: call:google_search{query:current UTC date}` +/// observed against `daily-cloudcode-pa.googleapis.com` with +/// `tools: [{"googleSearchRetrieval": {}}]` and no function declarations. +/// CLIProxyAPI sends `googleSearch` to the same v1internal surface. fn normalize_antigravity_builtin_tool_names(request: &mut Map) { let Some(tools) = request.get_mut("tools").and_then(Value::as_array_mut) else { return; @@ -128,23 +155,55 @@ fn normalize_antigravity_builtin_tool_names(request: &mut Map) { continue; }; - if let Some(payload) = tool_object.remove("googleSearch") { - tool_object - .entry("googleSearchRetrieval".to_string()) - .or_insert(payload); - } if let Some(payload) = tool_object.remove("google_search") { tool_object - .entry("googleSearchRetrieval".to_string()) + .entry("googleSearch".to_string()) .or_insert(payload); } } } -fn normalize_antigravity_function_declaration_parameters(request: &mut Map) { - let Some(tools) = request.get_mut("tools").and_then(Value::as_array_mut) else { +/// Claude requires a replayable signature on historical thinking blocks. In +/// particular, Responses reasoning summaries are not signed thinking. Omit that +/// non-replayable metadata instead of inventing a signature or promoting private +/// reasoning into ordinary assistant text. Keep Gemini's native policy unchanged. +fn normalize_antigravity_claude_thought_history(request: &mut Map, model: &str) { + if !model.trim().to_ascii_lowercase().starts_with("claude-") { + return; + } + let Some(contents) = request.get_mut("contents").and_then(Value::as_array_mut) else { return; }; + contents.retain_mut(|message| { + if message.get("role").and_then(Value::as_str) != Some("model") { + return true; + } + let Some(parts) = message.get_mut("parts").and_then(Value::as_array_mut) else { + return true; + }; + let previous_len = parts.len(); + parts.retain(|part| { + part.get("thought").and_then(Value::as_bool) != Some(true) + || ["thoughtSignature", "thought_signature"].iter().any(|key| { + part.get(*key) + .and_then(Value::as_str) + .is_some_and(|signature| !signature.trim().is_empty()) + }) + }); + // Do not introduce empty messages when a turn contained only a summary. + !parts.is_empty() || parts.len() == previous_len + }); +} + +fn normalize_antigravity_function_declaration_parameters( + request: &mut Map, + model: &str, +) -> Result<(), ()> { + let Some(tools) = request.get_mut("tools").and_then(Value::as_array_mut) else { + return Ok(()); + }; + let mut budget = SchemaBudget::default(); + let claude = model.trim().to_ascii_lowercase().starts_with("claude-"); for tool in tools { let Some(tool_object) = tool.as_object_mut() else { @@ -168,9 +227,16 @@ fn normalize_antigravity_function_declaration_parameters(request: &mut Map) -> Option<&Map> { @@ -207,6 +273,282 @@ mod tests { }; use crate::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT; + #[test] + fn antigravity_claude_fabric_union_regression_across_client_formats() { + use aether_ai_formats::formats::shared::standard_matrix::build_standard_request_body; + let schema: serde_json::Value = + serde_json::from_str(include_str!("fabric_exec_schema.json")).unwrap(); + let clients = [ + ( + "gemini:generate_content", + json!({"contents": [{"role":"user","parts":[{"text":"hi"}]}], + "tools":[{"functionDeclarations":[{"name":"fabric_exec","parametersJsonSchema":schema}]}]}), + ), + ( + "claude:messages", + json!({"messages":[{"role":"user","content":"hi"}],"max_tokens":128, + "tools":[{"name":"fabric_exec","input_schema":schema}]}), + ), + ( + "openai:chat", + json!({"messages":[{"role":"user","content":"hi"}], + "tools":[{"type":"function","function":{"name":"fabric_exec","parameters":schema}}]}), + ), + ( + "openai:responses", + json!({"input":"hi", + "tools":[{"type":"function","name":"fabric_exec","parameters":schema}]}), + ), + ]; + for (format, body) in clients { + for model in ["claude-opus-4-6-thinking", "gemini-test"] { + let converted = build_standard_request_body( + &body, + format, + model, + "antigravity", + "gemini:generate_content", + "", + true, + None, + None, + ) + .unwrap(); + for request in [converted.clone(), json!({"request":converted})] { + let AntigravityRequestEnvelopeSupport::Supported(output) = + build_antigravity_safe_v1internal_request( + &sample_auth(), + "union-regression", + model, + &request, + AntigravityEnvelopeRequestType::Agent, + ) + else { + panic!("failed {format} {model}"); + }; + let s = &output["request"]["tools"][0]["functionDeclarations"][0]["parameters"]; + assert_eq!(s["required"], json!(["code"])); + assert_eq!( + s["properties"]["payloads"]["additionalProperties"], + json!({"type":"string"}) + ); + assert_eq!(s["properties"]["tokenBudget"]["minimum"], 1); + if model.starts_with("claude-") { + assert_eq!( + s["properties"]["resultFormat"], + json!({"type":"string","enum":["auto","json","text","yaml"]}) + ); + let display = &s["properties"]["display"]; + assert!( + display.get("type").is_none(), + "do not select one union branch" + ); + assert!(display.get("anyOf").is_none()); + assert!(display["description"].as_str().unwrap().contains("object")); + assert!(display["description"].as_str().unwrap().contains("string")); + } else { + assert!(s["properties"]["resultFormat"]["anyOf"].is_array()); + assert!(s["properties"]["display"]["anyOf"].is_array()); + } + let AntigravityRequestEnvelopeSupport::Supported(twice) = + build_antigravity_safe_v1internal_request( + &sample_auth(), + "union-regression", + model, + &output, + AntigravityEnvelopeRequestType::Agent, + ) + else { + panic!("idempotence"); + }; + assert_eq!(twice, output); + } + } + } + } + + #[test] + fn antigravity_claude_omits_only_unsigned_thought_parts() { + let signed = json!({"text":"signed plan","thought":true,"thoughtSignature":"signed-value"}); + let signed_alias = + json!({"text":"signed alias","thought":true,"thought_signature":"alias-value"}); + let call = json!({"functionCall":{"id":"call_1","name":"lookup","args":{}},"thoughtSignature":"skip_thought_signature_validator"}); + let result = json!({"role":"user","parts":[{"functionResponse":{"id":"call_1","name":"lookup","response":{"result":"ok"}}}]}); + let body = json!({ + "contents":[ + {"role":"user","parts":[{"text":"hello"}]}, + {"role":"model","parts":[{"text":"unsigned-only summary","thought":true}]}, + {"role":"model","parts":[ + {"text":"unsigned summary","thought":true}, + {"text":"empty signature","thought":true,"thoughtSignature":""}, + {"text":"blank signature","thought":true,"thoughtSignature":" "}, + {"text":"non-string signature","thought":true,"thoughtSignature":12}, + signed,signed_alias,{"text":"visible answer"},call + ]}, + result + ], + "generationConfig":{"maxOutputTokens":64000,"thinkingConfig":{"includeThoughts":true,"thinkingBudget":4096}} + }); + for model in [ + "claude-sonnet-4-6", + "claude-opus-4-6-thinking", + "gemini-3.8-flash-high", + ] { + for wrapped in [false, true] { + let input = if wrapped { + json!({"request":body}) + } else { + body.clone() + }; + let original = input.clone(); + let AntigravityRequestEnvelopeSupport::Supported(output) = + build_antigravity_safe_v1internal_request( + &sample_auth(), + "thought-test", + model, + &input, + AntigravityEnvelopeRequestType::Agent, + ) + else { + panic!("expected supported envelope"); + }; + let expected = if model.starts_with("claude-") { + json!([ + {"role":"user","parts":[{"text":"hello"}]}, + {"role":"model","parts":[signed,signed_alias,{"text":"visible answer"},call]}, + result + ]) + } else { + body["contents"].clone() + }; + assert_eq!( + output["request"]["contents"], expected, + "{model} wrapped={wrapped}" + ); + assert_eq!( + output["request"]["generationConfig"], + body["generationConfig"] + ); + assert_eq!(input, original, "do not mutate caller-owned input"); + let rebuilt = build_antigravity_safe_v1internal_request( + &sample_auth(), + "thought-test", + model, + &output, + AntigravityEnvelopeRequestType::Agent, + ); + assert_eq!( + rebuilt, + AntigravityRequestEnvelopeSupport::Supported(output) + ); + } + } + } + + #[test] + fn antigravity_claude_cross_format_unsigned_reasoning_history() { + use aether_ai_formats::formats::shared::standard_matrix::build_standard_request_body; + let clients = [ + ( + "openai:responses", + json!({ + "max_output_tokens":64000, + "input":[ + {"role":"user","content":"hello"}, + {"type":"reasoning","id":"rs_history","status":"completed","summary":[{"type":"summary_text","text":"historical summary"}],"content":[]}, + {"role":"assistant","content":[{"type":"output_text","text":"visible answer"}]}, + {"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}"}, + {"type":"function_call_output","call_id":"call_1","output":"ok"}, + {"role":"user","content":"continue"} + ] + }), + ), + ( + "claude:messages", + json!({ + "max_tokens":64000, + "messages":[ + {"role":"user","content":"hello"}, + {"role":"assistant","content":[ + {"type":"thinking","thinking":"historical summary"}, + {"type":"text","text":"visible answer"}, + {"type":"tool_use","id":"call_1","name":"lookup","input":{}} + ]}, + {"role":"user","content":[{"type":"tool_result","tool_use_id":"call_1","content":"ok"},{"type":"text","text":"continue"}]} + ] + }), + ), + ]; + for (format, body) in clients { + for model in [ + "claude-sonnet-4-6", + "claude-opus-4-6-thinking", + "gemini-3.8-flash-high", + ] { + let converted = build_standard_request_body( + &body, + format, + model, + "antigravity", + "gemini:generate_content", + "", + true, + None, + None, + ) + .unwrap(); + assert!( + converted["contents"] + .as_array() + .unwrap() + .iter() + .flat_map(|m| m["parts"].as_array().unwrap()) + .any(|p| p["thought"] == true), + "fixture must exercise unsigned thoughts: {format}" + ); + let AntigravityRequestEnvelopeSupport::Supported(output) = + build_antigravity_safe_v1internal_request( + &sample_auth(), + "reasoning-history", + model, + &converted, + AntigravityEnvelopeRequestType::Agent, + ) + else { + panic!("expected supported envelope"); + }; + let messages = output["request"]["contents"].as_array().unwrap(); + assert!(messages + .iter() + .all(|m| !m["parts"].as_array().unwrap().is_empty())); + let parts: Vec<_> = messages + .iter() + .flat_map(|m| m["parts"].as_array().unwrap()) + .collect(); + assert_eq!( + parts.iter().any(|p| p["thought"] == true), + !model.starts_with("claude-") + ); + assert!(parts.iter().any(|p| p["text"] == "visible answer")); + assert!(parts.iter().any(|p| p["text"] == "continue")); + assert!(parts.iter().any(|p| p["functionCall"]["id"] == "call_1" + && p["functionCall"]["name"] == "lookup")); + assert!(parts.iter().any(|p| p["functionResponse"]["id"] == "call_1" + && p["functionResponse"]["name"] == "lookup")); + if model.starts_with("claude-") { + assert!( + !parts.iter().any(|p| p["text"] == "historical summary"), + "do not promote private reasoning to visible text" + ); + } + assert_eq!( + output["request"]["generationConfig"]["maxOutputTokens"], + 64000 + ); + } + } + } + fn sample_auth() -> AntigravityRequestAuth { AntigravityRequestAuth { project_id: "project-ant-123".to_string(), @@ -215,6 +557,236 @@ mod tests { } } + #[test] + fn antigravity_combined_client_conversion_preserves_schemas_until_transport() { + use aether_ai_formats::formats::shared::standard_matrix::build_standard_request_body; + let schema = json!({"type": "object", "properties": { + "mode": {"const": "fast"}, + "payloads": {"type": "object", "patternProperties": {"^.*$": {"type": "string"}}}, + "name": {"type": "string", "minLength": 1} + }, "required": ["mode"]}); + let clients = [ + ( + "claude:messages", + json!({"model": "client-model", "max_tokens": 128, + "messages": [{"role": "user", "content": "hi"}], + "tools": [{"name": "probe", "input_schema": schema}]}), + ), + ( + "openai:chat", + json!({"model": "client-model", + "messages": [{"role": "user", "content": "hi"}], + "tools": [{"type": "function", "function": {"name": "probe", "parameters": schema}}]}), + ), + ( + "openai:responses", + json!({"model": "client-model", "input": "hi", + "tools": [{"type": "function", "name": "probe", "parameters": schema}]}), + ), + ( + "gemini:generate_content", + json!({"model": "client-model", + "contents": [{"role": "user", "parts": [{"text": "hi"}]}], + "tools": [{"functionDeclarations": [{"name": "probe", "parameters": schema}]}]}), + ), + ]; + for (source, original) in clients { + for model in ["claude-sonnet-test", "gemini-test"] { + let converted = build_standard_request_body( + &original, + source, + model, + " AnTiGrAvItY ", + "gemini:generate_content", + "", + true, + None, + None, + ) + .expect("Antigravity conversion"); + assert_eq!( + converted["tools"][0]["functionDeclarations"][0]["parameters"], schema, + "{source}" + ); + let AntigravityRequestEnvelopeSupport::Supported(envelope) = + build_antigravity_safe_v1internal_request( + &sample_auth(), + "review-test", + model, + &converted, + AntigravityEnvelopeRequestType::Agent, + ) + else { + panic!("schema should fit budget"); + }; + let parameters = + &envelope["request"]["tools"][0]["functionDeclarations"][0]["parameters"]; + assert_eq!( + parameters["properties"]["mode"], + json!({"type": "string", "enum": ["fast"]}) + ); + assert_eq!( + parameters["properties"]["payloads"]["additionalProperties"], + json!({"type": "string"}) + ); + assert_eq!(parameters["properties"]["name"]["minLength"], 1); + assert_eq!(envelope["model"], model); + } + // The default public Gemini policy must remain unchanged. + let public = build_standard_request_body( + &original, + source, + "gemini-test", + "gemini", + "gemini:generate_content", + "", + true, + None, + None, + ) + .unwrap(); + let parameters = &public["tools"][0]["functionDeclarations"][0]["parameters"]; + if source == "gemini:generate_content" { + assert_eq!(parameters, &schema); + } else { + assert!(parameters["properties"]["mode"].get("const").is_none()); + assert_eq!(parameters["properties"]["name"]["minLength"], "1"); + } + } + } + + #[test] + fn antigravity_responses_conversion_keeps_scoped_previous_response_history() { + use aether_ai_formats::{ + api::record_converted_response_history, + formats::shared::standard_matrix::build_standard_request_body, + }; + let schema = json!({"type": "object", "properties": {"mode": {"const": "fast"}}}); + let response_id = "resp_antigravity_schema_history_test"; + let scope = "antigravity-schema-history"; + record_converted_response_history(&json!({ + "needs_conversion": true, "client_api_format": "openai:responses", + "provider_api_format": "openai:chat", "api_key_id": scope, + "original_request_body": {"model": "client", "input": "first"} + }), &json!({"id": response_id, "status": "completed", "output": [{ + "type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "remembered"}] + }]})).expect("seed scoped history"); + let input = json!({"model": "client", "previous_response_id": response_id, + "input": "second", "tools": [{"type": "function", "name": "probe", "parameters": schema}]}); + let output = build_standard_request_body( + &input, + "openai:responses", + "claude-test", + "antigravity", + "gemini:generate_content", + "", + true, + None, + Some(scope), + ) + .expect("expand history"); + assert_eq!(output["contents"][0]["parts"][0]["text"], "first"); + assert_eq!(output["contents"][1]["parts"][0]["text"], "remembered"); + assert_eq!(output["contents"][2]["parts"][0]["text"], "second"); + assert_eq!( + output["tools"][0]["functionDeclarations"][0]["parameters"], + schema + ); + assert!(build_standard_request_body( + &input, + "openai:responses", + "claude-test", + "antigravity", + "gemini:generate_content", + "", + true, + None, + Some("different-key") + ) + .is_none()); + } + + #[test] + fn antigravity_rejects_shared_schema_budget_exhaustion_on_both_envelope_paths() { + use super::AntigravityRequestEnvelopeUnsupportedReason; + let declaration = json!({"name": "probe", "parameters": { + "type": "object", "description": "x".repeat(600_000) + }}); + let mut body = json!({"contents": [], "tools": [{"functionDeclarations": [declaration]}]}); + assert!(matches!( + build_antigravity_safe_v1internal_request( + &sample_auth(), + "test", + "claude-test", + &body, + AntigravityEnvelopeRequestType::Agent, + ), + AntigravityRequestEnvelopeSupport::Supported(_) + )); + body["tools"][0]["functionDeclarations"] + .as_array_mut() + .unwrap() + .push(declaration); + for wrapped in [false, true] { + let input = if wrapped { + json!({"request": body}) + } else { + body.clone() + }; + let snapshot = input.clone(); + assert_eq!( + build_antigravity_safe_v1internal_request( + &sample_auth(), + "test", + "claude-test", + &input, + AntigravityEnvelopeRequestType::Agent, + ), + AntigravityRequestEnvelopeSupport::Unsupported( + AntigravityRequestEnvelopeUnsupportedReason::ToolSchemaBudgetExceeded + ) + ); + assert_eq!(input, snapshot); + } + } + + #[test] + fn search_only_request_keeps_the_modern_google_search_spelling() { + // Reproduces the live failure: a grounding-only request (no function + // declarations) that went out as `googleSearchRetrieval` came back as + // `Malformed function call: call:google_search{query:current UTC date}` + // from daily-cloudcode-pa.googleapis.com. + let request_body = json!({ + "contents": [ + { "role": "user", "parts": [{ "text": "today's UTC date?" }] } + ], + "tools": [{ "googleSearch": {} }] + }); + + let envelope = match build_antigravity_safe_v1internal_request( + &sample_auth(), + "request-ant-search-1", + "gemini-3.8-flash-high", + &request_body, + AntigravityEnvelopeRequestType::Agent, + ) { + AntigravityRequestEnvelopeSupport::Supported(envelope) => envelope, + AntigravityRequestEnvelopeSupport::Unsupported(reason) => { + panic!("search-only envelope should be supported: {reason:?}") + } + }; + + let tools = envelope["request"]["tools"] + .as_array() + .expect("tools should survive"); + assert_eq!(tools.len(), 1, "{tools:?}"); + assert_eq!(tools[0]["googleSearch"], json!({})); + assert!( + tools[0].get("googleSearchRetrieval").is_none(), + "the Gemini 1.5 spelling must not be reintroduced: {tools:?}" + ); + } + #[test] fn real_agent_request_preserves_antigravity_agent_fields() { let request_body = json!({ @@ -298,7 +870,7 @@ mod tests { assert_eq!(envelope["requestId"], "request-ant-agent-123"); assert_eq!(envelope["model"], "gemini-3.5-flash-low"); assert_eq!(envelope["userAgent"], ANTIGRAVITY_REQUEST_USER_AGENT); - assert_eq!(envelope["requestType"], "agent"); + assert!(envelope.get("requestType").is_none()); assert!(envelope["request"].get("model").is_none()); assert!(envelope["request"].get("safetySettings").is_none()); assert_eq!( @@ -321,12 +893,9 @@ mod tests { .get("include_server_side_tool_invocations") .is_none()); assert!(envelope["request"]["tools"][0] - .get("googleSearch") + .get("googleSearchRetrieval") .is_none()); - assert_eq!( - envelope["request"]["tools"][0]["googleSearchRetrieval"], - json!({}) - ); + assert_eq!(envelope["request"]["tools"][0]["googleSearch"], json!({})); assert_eq!( envelope["request"]["tools"][1]["functionDeclarations"][0]["name"], "run_command" @@ -464,7 +1033,7 @@ mod tests { .get("google_search") .is_none()); assert_eq!( - envelope["request"]["tools"][0]["googleSearchRetrieval"], + envelope["request"]["tools"][0]["googleSearch"], json!({ "dynamicRetrievalConfig": { "mode": "MODE_UNSPECIFIED" @@ -473,6 +1042,162 @@ mod tests { ); } + #[test] + fn antigravity_envelope_downgrades_fabric_schema_on_all_input_paths() { + let schema = json!({ + "$schema": "http://json-schema.org/draft-07/schema#", + "type": "object", + "properties": { + "code": { "type": "string", "description": "TypeScript function body" }, + "payloads": { + "type": "object", + "patternProperties": { "^.*$": { "type": "string" } } + }, + "resultFormat": { + "anyOf": [ + { "type": "string", "const": "auto" }, + { "type": "string", "const": "yaml" }, + { "type": "string", "const": "json" }, + { "type": "string", "const": "text" } + ] + }, + "display": { + "anyOf": [ + { "type": "object", "properties": { "name": { "type": "string" } } }, + { "type": "string" } + ] + }, + "tokenBudget": { "type": "number", "minimum": 1 } + }, + "required": ["code"], + "minProperties": "1" + }); + let mut expected = schema.clone(); + expected.as_object_mut().unwrap().remove("$schema"); + expected["minProperties"] = json!(1); + expected["properties"]["payloads"] = json!({ + "type": "object", "additionalProperties": { "type": "string" } + }); + for (index, format) in ["auto", "yaml", "json", "text"].iter().enumerate() { + expected["properties"]["resultFormat"]["anyOf"][index] = + json!({ "type": "string", "enum": [format] }); + } + + for declarations_key in ["functionDeclarations", "function_declarations"] { + for parameters_key in [ + "parametersJsonSchema", + "parameters_json_schema", + "parameters", + ] { + for wrapped in [false, true] { + let mut body = json!({ + "contents": [{ "role": "user", "parts": [{ "text": "hello" }] }], + "tools": [{ "googleSearch": {} }, {}], + "generationConfig": { "maxOutputTokens": 4096 }, + "labels": { "const": "not a schema" } + }); + body["tools"][1][declarations_key] = json!([ + { "name": "fabric_exec", "description": "Execute TypeScript" }, + { "name": "other", "parameters": { "type": "object" } } + ]); + body["tools"][1][declarations_key][0][parameters_key] = schema.clone(); + if wrapped { + body = json!({ "request": body, "requestId": "client-id" }); + } + let original = body.clone(); + let AntigravityRequestEnvelopeSupport::Supported(envelope) = + build_antigravity_safe_v1internal_request( + &sample_auth(), + "trace-id", + "gemini-test", + &body, + AntigravityEnvelopeRequestType::Agent, + ) + else { + panic!("Fabric request should be supported"); + }; + let declaration = &envelope["request"]["tools"][1][declarations_key][0]; + assert_eq!( + declaration["parameters"], expected, + "{declarations_key}/{parameters_key}/wrapped={wrapped}" + ); + assert_eq!(declaration["name"], "fabric_exec"); + assert!(declaration.get("parametersJsonSchema").is_none()); + assert!(declaration.get("parameters_json_schema").is_none()); + assert_eq!( + envelope["request"]["generationConfig"]["maxOutputTokens"], + 4096 + ); + assert_eq!(envelope["request"]["labels"]["const"], "not a schema"); + assert_eq!( + envelope["request"]["tools"][0], + json!({ "googleSearch": {} }) + ); + assert_eq!( + envelope["requestId"], + if wrapped { "client-id" } else { "trace-id" } + ); + assert_eq!(envelope["project"], "project-ant-123"); + assert_eq!(body, original, "caller-owned input must remain unchanged"); + + let AntigravityRequestEnvelopeSupport::Supported(rebuilt) = + build_antigravity_safe_v1internal_request( + &sample_auth(), + "trace-id", + "gemini-test", + &envelope, + AntigravityEnvelopeRequestType::Agent, + ) + else { + panic!("converted envelope should be supported"); + }; + assert_eq!(rebuilt, envelope, "normalization must be idempotent"); + } + } + } + } + + #[test] + fn antigravity_schema_aliases_do_not_override_existing_parameters() { + let body = json!({ + "contents": [], + "tools": [{ "functionDeclarations": [{ + "name": "select", + "parameters": { "type": "string", "const": "existing" }, + "parametersJsonSchema": { "type": "string", "const": "camel" }, + "parameters_json_schema": { "type": "string", "const": "snake" } + }, { + "name": "select_alias", + "parametersJsonSchema": { "type": "string", "const": "camel" }, + "parameters_json_schema": { "type": "string", "const": "snake" } + }] }] + }); + let AntigravityRequestEnvelopeSupport::Supported(envelope) = + build_antigravity_safe_v1internal_request( + &sample_auth(), + "trace-id", + "gemini-test", + &body, + AntigravityEnvelopeRequestType::Agent, + ) + else { + panic!("request should be supported"); + }; + let declarations = &envelope["request"]["tools"][0]["functionDeclarations"]; + assert_eq!( + declarations[0]["parameters"], + json!({ "type": "string", "enum": ["existing"] }) + ); + assert_eq!( + declarations[1]["parameters"], + json!({ "type": "string", "enum": ["camel"] }) + ); + for declaration in declarations.as_array().unwrap() { + assert!(declaration.get("parametersJsonSchema").is_none()); + assert!(declaration.get("parameters_json_schema").is_none()); + } + } + #[test] fn antigravity_envelope_normalizes_json_schema_parameter_spellings() { let request_body = json!({ diff --git a/crates/aether-provider/transport/src/antigravity/schema.rs b/crates/aether-provider/transport/src/antigravity/schema.rs new file mode 100644 index 000000000..fee8790f3 --- /dev/null +++ b/crates/aether-provider/transport/src/antigravity/schema.rs @@ -0,0 +1,789 @@ +use std::collections::BTreeSet; + +use serde_json::{json, Map, Value}; + +/// Cloud Code's `parameters` uses a protobuf-shaped schema, while its Claude +/// backend validates the translated `tools.*.custom.input_schema` as JSON Schema +/// draft 2020-12. Keep this lowering at the Antigravity boundary: it preserves the +/// common subset and converts protobuf JSON int64 strings back to JSON numbers. +/// Callers must still validate tool arguments against their original schema. +pub(super) fn normalize_tool_parameters( + parameters: &mut Value, + budget: &mut SchemaBudget, +) -> Result<(), ()> { + let lowered = lower_schema(parameters, parameters, &mut BTreeSet::new(), 0, budget); + if budget.exhausted { + return Err(()); + } + *parameters = lowered; + Ok(()) +} + +/// Cloud Code accepts these unions but its Claude bridge rejects typed anyOf +/// branches (verified with the real fabric_exec schema). Do not apply this to +/// Gemini models. Fold string literal alternatives exactly; otherwise relax the +/// union and retain its constraints as guidance, never choose an arbitrary branch. +/// The caller must validate generated arguments against the original schema. +pub(super) fn normalize_claude_unions( + schema: &mut Value, + budget: &mut SchemaBudget, +) -> Result<(), ()> { + lower_claude_unions(schema, budget) +} + +fn lower_claude_unions(value: &mut Value, budget: &mut SchemaBudget) -> Result<(), ()> { + let Some(schema) = value.as_object_mut() else { + return Ok(()); + }; + if let Some(Value::Array(branches)) = schema.remove("anyOf") { + // Only intersect an existing sibling enum when both sides are strings; + // an empty intersection cannot be represented by protobuf enum (omitted). + let literals = string_union_literals(&branches); + let folded = literals.and_then(|mut literals| { + if let Some(Value::Array(existing)) = schema.get("enum") { + literals.retain(|literal| existing.contains(&Value::String(literal.clone()))); + } + (!literals.is_empty()).then_some(literals) + }); + if let Some(literals) = folded { + schema.entry("type").or_insert(json!("string")); + schema.insert("enum".into(), json!(literals)); + } else { + let guidance = serde_json::to_string(&branches).expect("JSON value serializes"); + let description = schema.entry("description").or_insert(json!("")); + let original = description.as_str().unwrap_or_default(); + *description = json!(format!("{original}\nAccepted alternatives (validate against the original tool schema): {guidance}").trim()); + if !budget.charge(description) { + return Err(()); + } + } + } + // Traverse schema positions only: descriptions/defaults/examples and property + // names such as `anyOf` are data, not keywords to rewrite. + if let Some(Value::Object(properties)) = schema.get_mut("properties") { + for child in properties.values_mut() { + lower_claude_unions(child, budget)?; + } + } + for key in ["items", "additionalProperties"] { + if let Some(child) = schema.get_mut(key) { + lower_claude_unions(child, budget)?; + } + } + Ok(()) +} + +fn string_union_literals(branches: &[Value]) -> Option> { + if branches.is_empty() { + return None; + } + let mut literals = BTreeSet::new(); + for branch in branches { + let branch = branch.as_object()?; + if branch + .keys() + .any(|key| !matches!(key.as_str(), "type" | "enum" | "description" | "title")) + || branch.get("type").is_some_and(|ty| ty != "string") + { + return None; + } + let values = branch.get("enum")?.as_array()?; + if values.is_empty() { + return None; + } + for literal in values { + literals.insert(literal.as_str()?.to_owned()); + } + } + Some(literals) +} + +/// Shared across all tool schemas in one request. Counting serialized input at +/// every expansion conservatively bounds cloning work, including literal data +/// and repeated acyclic references, without allocating serialized copies. +pub(super) struct SchemaBudget { + nodes_left: usize, + bytes_left: usize, + exhausted: bool, +} + +impl Default for SchemaBudget { + fn default() -> Self { + Self { + nodes_left: 4096, + bytes_left: 1024 * 1024, + exhausted: false, + } + } +} + +impl SchemaBudget { + fn charge(&mut self, value: &Value) -> bool { + if self.exhausted || self.nodes_left == 0 { + self.exhausted = true; + return false; + } + self.nodes_left -= 1; + if serde_json::to_writer(&mut *self, value).is_err() { + self.exhausted = true; + return false; + } + true + } +} + +impl std::io::Write for SchemaBudget { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + if bytes.len() > self.bytes_left { + return Err(std::io::Error::other( + "tool schema expansion budget exceeded", + )); + } + self.bytes_left -= bytes.len(); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +/// JSON Schema $ref siblings are conjunctive, not an object-spread override. +/// Merge properties and required sets; retain the referenced constraint when +/// two validation keywords cannot be intersected in the wire subset. This +/// relaxes validation rather than manufacturing a contradictory schema. +fn merge_ref_siblings(mut target: Value, siblings: Value) -> Value { + let target_object = target.as_object_mut().expect("lowered schema is an object"); + let Value::Object(siblings) = siblings else { + unreachable!("lowered schema is an object") + }; + for (key, value) in siblings { + match (key.as_str(), target_object.get_mut(&key), value) { + ("properties", Some(Value::Object(properties)), Value::Object(children)) => { + for (name, child) in children { + if let Some(existing) = properties.get_mut(&name) { + *existing = merge_ref_siblings(existing.take(), child); + } else { + properties.insert(name, child); + } + } + } + ("required", Some(Value::Array(required)), Value::Array(names)) => { + let mut seen: BTreeSet = required + .iter() + .filter_map(Value::as_str) + .map(str::to_owned) + .collect(); + for name in names { + if seen.insert(name.as_str().expect("required contains strings").to_owned()) { + required.push(name); + } + } + } + ("title" | "description" | "default" | "example", _, value) | (_, None, value) => { + target_object.insert(key, value); + } + _ => {} + } + } + target +} + +fn lower_schema( + value: &Value, + root: &Value, + resolving: &mut BTreeSet, + depth: usize, + budget: &mut SchemaBudget, +) -> Value { + // Bound both recursive references and expansion of deeply nested schemas. + if !budget.charge(value) || depth >= 64 { + return json!({}); + } + let Some(source) = value.as_object() else { + // Boolean schemas have no protobuf equivalent (false is relaxed). + return json!({}); + }; + let mut source = source.clone(); + if let Some(Value::String(reference)) = source.remove("$ref") { + if let Some(target) = reference + .strip_prefix('#') + .and_then(|pointer| root.pointer(pointer)) + .filter(|target| target.is_object()) + { + if resolving.insert(reference.clone()) { + let target = lower_schema(target, root, resolving, depth + 1, budget); + resolving.remove(&reference); + let siblings = + lower_schema(&Value::Object(source), root, resolving, depth + 1, budget); + return merge_ref_siblings(target, siblings); + } + } + // Unresolved, external or cyclic refs retain only their sibling fields. + } + + let mut schema = Map::new(); + // Keep only a typed JSON-Schema subset. Gemini's protobuf JSON mapping renders + // int64 constraints as strings, but Claude rejects those at custom.input_schema. + match source.get("type") { + Some(Value::String(schema_type)) => { + if let Some(schema_type) = json_schema_type(schema_type) { + schema.insert("type".to_string(), Value::String(schema_type.to_string())); + } + } + Some(Value::Array(types)) => { + schema.insert("type".to_string(), Value::Array(types.clone())); + } + _ => {} + } + for key in ["format", "title", "description", "pattern"] { + if let Some(value) = source.get(key).filter(|value| value.is_string()) { + schema.insert(key.to_string(), value.clone()); + } + } + if let Some(value) = source.get("nullable").filter(|value| value.is_boolean()) { + schema.insert("nullable".to_string(), value.clone()); + } + if let Some(values) = json_schema_string_array(source.get("enum"), true) { + schema.insert("enum".to_string(), values); + } + for key in ["minimum", "maximum"] { + if let Some(value) = source.get(key).filter(|value| value.is_number()) { + schema.insert(key.to_string(), value.clone()); + } + } + for key in [ + "minItems", + "maxItems", + "minLength", + "maxLength", + "minProperties", + "maxProperties", + ] { + if let Some(value) = source.get(key).and_then(json_schema_nonnegative_integer) { + schema.insert(key.to_string(), value); + } + } + if let Some(values) = json_schema_string_array(source.get("required"), false) { + schema.insert("required".to_string(), values); + } + if let Some(values) = json_schema_string_array(source.get("propertyOrdering"), false) { + schema.insert("propertyOrdering".to_string(), values); + } + for key in ["default", "example"] { + if let Some(value) = source.get(key) { + schema.insert(key.to_string(), value.clone()); + } + } + + if let Some(constant) = source.get("const") { + if constant.is_string() { + schema.insert("enum".to_string(), json!([constant])); + schema.entry("type").or_insert(json!("string")); + } else { + // The protobuf enum field is repeated string. Never stringify numeric + // or boolean literals into enums: that changes the argument's type. + let description = schema.entry("description").or_insert(json!("")); + let prefix = description.as_str().unwrap_or_default(); + *description = json!(format!("{prefix}\nMust equal: {constant}").trim()); + } + } + if let Some(values) = schema.get_mut("enum").and_then(Value::as_array_mut) { + // Mixed/non-string enums cannot be represented without changing types. + if !values.iter().all(Value::is_string) { + schema.remove("enum"); + } + } + + if let Some(properties) = source.get("properties").and_then(Value::as_object) { + schema.insert( + "properties".to_string(), + Value::Object( + properties + .iter() + .map(|(name, child)| { + ( + name.clone(), + lower_schema(child, root, resolving, depth + 1, budget), + ) + }) + .collect(), + ), + ); + } + if let Some(items) = source.get("items") { + schema.insert( + "items".to_string(), + lower_schema(items, root, resolving, depth + 1, budget), + ); + } + // oneOf's exclusivity is not supported; anyOf retains the alternatives. + if let Some(branches) = source + .get("anyOf") + .or_else(|| source.get("any_of")) + .or_else(|| source.get("oneOf")) + .and_then(Value::as_array) + { + if !branches.is_empty() { + schema.insert( + "anyOf".to_string(), + Value::Array( + branches + .iter() + .map(|child| lower_schema(child, root, resolving, depth + 1, budget)) + .collect(), + ), + ); + } + } + + let patterns = source.get("patternProperties").and_then(Value::as_object); + let wildcard = patterns + .filter(|patterns| patterns.len() == 1) + .and_then(|patterns| { + patterns + .iter() + .next() + .filter(|(pattern, _)| matches!(pattern.as_str(), ".*" | "^.*$")) + }) + .map(|(_, child)| child); + // For a catch-all, JSON Schema's additionalProperties applies to no unmatched + // keys, so even an explicit `false` must not suppress the dictionary values. + let additional = wildcard.or_else(|| { + // Without regex matching, an additional-properties constraint could + // incorrectly reject keys formerly accepted by a pattern. Relax it too. + patterns + .is_none_or(Map::is_empty) + .then(|| source.get("additionalProperties")) + .flatten() + }); + if let Some(additional) = additional { + if additional.is_boolean() { + // Dropping non-wildcard patterns must not turn their allowed keys into + // forbidden additional properties. General regex maps are relaxed. + if additional != &Value::Bool(false) || patterns.is_none_or(Map::is_empty) { + schema.insert("additionalProperties".to_string(), additional.clone()); + } + } else { + schema.insert( + "additionalProperties".to_string(), + lower_schema(additional, root, resolving, depth + 1, budget), + ); + } + } + + if let Some(Value::Array(types)) = schema.get("type").cloned() { + schema.remove("type"); + let nullable = types.iter().any(|value| value.as_str() == Some("null")); + let types: BTreeSet<_> = types + .iter() + .filter_map(Value::as_str) + .filter_map(json_schema_type) + .filter(|value| *value != "null") + .collect(); + if types.len() == 1 { + schema.insert("type".to_string(), json!(types.first().unwrap())); + } else if !types.is_empty() { + schema.entry("anyOf").or_insert_with(|| { + Value::Array(types.iter().map(|ty| json!({ "type": ty })).collect()) + }); + } else if nullable { + schema.insert("type".to_string(), json!("null")); + } + if nullable && !types.is_empty() { + schema.insert("nullable".to_string(), Value::Bool(true)); + } + } + Value::Object(schema) +} + +fn json_schema_type(value: &str) -> Option<&'static str> { + match value.trim().to_ascii_lowercase().as_str() { + "array" => Some("array"), + "boolean" => Some("boolean"), + "integer" => Some("integer"), + "null" => Some("null"), + "number" => Some("number"), + "object" => Some("object"), + "string" => Some("string"), + _ => None, + } +} + +fn json_schema_nonnegative_integer(value: &Value) -> Option { + let value = match value { + Value::Number(value) => value.as_u64(), + // protobuf JSON encodes int64 fields as decimal strings. + Value::String(value) => value.parse::().ok(), + _ => None, + }?; + Some(Value::from(value)) +} + +fn json_schema_string_array(value: Option<&Value>, require_non_empty: bool) -> Option { + let values = value?.as_array()?; + let mut seen = BTreeSet::new(); + let values = values + .iter() + .map(Value::as_str) + .collect::>>()? + .into_iter() + .filter(|value| seen.insert(*value)) + .map(|value| Value::String(value.to_string())) + .collect::>(); + (!require_non_empty || !values.is_empty()).then_some(Value::Array(values)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn claude_unions_keep_siblings_literals_and_nested_schema_positions() { + let union = json!({"anyOf":[{"type":"string"},{"type":"object"}]}); + let mut schema = json!({"type":"object", "default":union, "properties":{ + "anyOf":{"type":"array","items":union}, + "map":{"type":"object","additionalProperties":union}, + "literal":{"enum":["b"],"anyOf":[{"type":"string","enum":["a"]},{"type":"string","enum":["b"]}]}, + "constrained":{"description":"Original","type":"string","minLength":2,"anyOf":[{"type":"string","pattern":"a+"},{"type":"number"}]} + }}); + normalize_claude_unions(&mut schema, &mut SchemaBudget::default()).unwrap(); + assert_eq!( + schema["default"], union, + "literal data must not be rewritten" + ); + assert!(schema["properties"]["anyOf"]["items"]["description"].is_string()); + assert!(schema["properties"]["map"]["additionalProperties"]["description"].is_string()); + assert_eq!( + schema["properties"]["literal"], + json!({"type":"string","enum":["b"]}) + ); + let constrained = &schema["properties"]["constrained"]; + assert_eq!(constrained["minLength"], 2); + assert_eq!(constrained["type"], "string"); + assert!(constrained["description"] + .as_str() + .unwrap() + .starts_with("Original")); + assert!(constrained["description"] + .as_str() + .unwrap() + .contains("pattern")); + let once = schema.clone(); + normalize_claude_unions(&mut schema, &mut SchemaBudget::default()).unwrap(); + assert_eq!(schema, once); + } + + #[test] + fn claude_union_guidance_is_charged_to_shared_budget() { + let mut schema = json!({"anyOf":[{"type":"string"},{"type":"number"}]}); + let mut budget = SchemaBudget { + nodes_left: 4096, + bytes_left: 8, + exhausted: false, + }; + assert!(normalize_claude_unions(&mut schema, &mut budget).is_err()); + assert!(budget.exhausted); + } + + #[test] + fn claude_union_folding_does_not_drop_branch_constraints() { + let mut schema = json!({"anyOf":[ + {"type":"string","enum":["a"],"minLength":2}, + {"type":"string","enum":["b"]} + ]}); + normalize_claude_unions(&mut schema, &mut SchemaBudget::default()).unwrap(); + assert!(schema.get("enum").is_none()); + assert!(schema["description"] + .as_str() + .unwrap() + .contains("minLength")); + } + + fn lowered(mut schema: Value) -> Value { + normalize_tool_parameters(&mut schema, &mut SchemaBudget::default()).unwrap(); + let once = schema.clone(); + normalize_tool_parameters(&mut schema, &mut SchemaBudget::default()).unwrap(); + assert_eq!(schema, once, "lowering must be idempotent"); + schema + } + + #[test] + fn reference_siblings_preserve_properties_and_required_without_contradictions() { + let result = lowered(json!({ + "$defs": {"Base": { + "type": "object", "properties": {"a": {"type": "string", "minLength": 1}}, + "required": ["a"], "additionalProperties": false + }}, + "$ref": "#/$defs/Base", + "properties": {"a": {"maxLength": 8}, "b": {"type": "string"}}, + "required": ["a"] + })); + assert_eq!( + result, + json!({ + "type": "object", "properties": { + "a": {"type": "string", "minLength": 1, "maxLength": 8}, + "b": {"type": "string"} + }, "required": ["a"], "additionalProperties": false + }) + ); + // A property schema named "additionalProperties" is data, not a keyword. + let merged = merge_ref_siblings( + json!({"properties": {"additionalProperties": {"type": "string"}}, "required": ["a"]}), + json!({"properties": {"b": {"type": "number"}}, "required": ["b", "a"]}), + ); + assert_eq!(merged["required"], json!(["a", "b"])); + assert_eq!( + merged["properties"]["additionalProperties"]["type"], + "string" + ); + } + + #[test] + fn sibling_reference_to_completed_target_is_not_a_cycle() { + let result = lowered(json!({ + "$defs": {"Base": {"type": "object", "properties": {"mode": {"const": "fast"}}}}, + "$ref": "#/$defs/Base", "properties": {"nested": {"$ref": "#/$defs/Base"}} + })); + assert_eq!( + result["properties"]["nested"]["properties"]["mode"], + json!({"type": "string", "enum": ["fast"]}) + ); + } + + #[test] + fn acyclic_branching_references_exhaust_budget_without_mutating_input() { + let mut schema = json!({"$defs": {"D0": {"type": "string"}}, "$ref": "#/$defs/D24"}); + for index in 1..=24 { + let reference = format!("#/$defs/D{}", index - 1); + schema["$defs"][format!("D{index}")] = json!({"type": "object", "properties": { + "left": {"$ref": reference}, "right": {"$ref": reference} + }}); + } + let original = schema.clone(); + let mut budget = SchemaBudget::default(); + assert!(normalize_tool_parameters(&mut schema, &mut budget).is_err()); + assert!(budget.exhausted); + assert_eq!(schema, original); + } + + #[test] + fn limits_nodes_and_literal_bytes_not_only_reference_depth() { + let mut wide = json!({"type": "object", "properties": {}}); + for index in 0..5000 { + wide["properties"][format!("p{index}")] = json!({}); + } + let mut budget = SchemaBudget::default(); + assert!(normalize_tool_parameters(&mut wide, &mut budget).is_err()); + assert_eq!(budget.nodes_left, 0); + let mut literal = json!({"type": "object", "default": "x".repeat(1024 * 1024)}); + assert!(normalize_tool_parameters(&mut literal, &mut SchemaBudget::default()).is_err()); + } + + #[test] + fn recursively_lowers_schema_nodes_without_touching_property_names_or_literal_data() { + let literal = json!({ "const": "data", "patternProperties": { "^.*$": 1 } }); + let result = lowered(json!({ + "$schema": "draft", "$id": "id", "x-custom": true, + "type": "object", "additionalProperties": false, + "properties": { + "const": { "const": "value", "readOnly": true }, + "patternProperties": { + "type": "array", "uniqueItems": true, + "items": { "oneOf": [{ "const": "a" }, { "const": "b" }] } + }, + "dictionary": { + "type": "object", + "additionalProperties": { + "any_of": [{ "const": "nested" }], "deprecated": true + } + } + }, + "default": literal, "example": literal, + "required": ["const"], "propertyOrdering": ["const", "patternProperties"] + })); + assert_eq!( + result, + json!({ + "type": "object", "additionalProperties": false, + "properties": { + "const": { "type": "string", "enum": ["value"] }, + "patternProperties": { + "type": "array", + "items": { "anyOf": [ + { "type": "string", "enum": ["a"] }, + { "type": "string", "enum": ["b"] } + ] } + }, + "dictionary": { + "type": "object", "additionalProperties": { + "anyOf": [{ "type": "string", "enum": ["nested"] }] + } + } + }, + "default": literal, "example": literal, + "required": ["const"], "propertyOrdering": ["const", "patternProperties"] + }) + ); + } + + #[test] + fn catch_all_patterns_preserve_dictionary_values_even_with_additional_properties_false() { + for pattern in [".*", "^.*$"] { + for additional in [json!(false), json!(true), json!({ "type": "number" })] { + let mut schema = json!({ "type": "object", "additionalProperties": additional }); + schema["patternProperties"][pattern] = json!({ "const": "value" }); + assert_eq!( + lowered(schema), + json!({ + "type": "object", "additionalProperties": { "type": "string", "enum": ["value"] } + }) + ); + } + } + } + + #[test] + fn general_patterns_are_not_mistaken_for_catch_all_dictionaries() { + for patterns in [ + json!({ "^x-": { "type": "string" } }), + json!({ "^.*$": { "type": "string" }, "^x-": { "maxLength": 5 } }), + ] { + for additional in [json!(false), json!({ "type": "number" })] { + assert_eq!( + lowered(json!({ + "type": "object", "patternProperties": patterns, + "additionalProperties": additional + })), + json!({ "type": "object" }) + ); + } + } + assert_eq!( + lowered(json!({ "type": "object", "additionalProperties": true })), + json!({ "type": "object", "additionalProperties": true }) + ); + } + + #[test] + fn resolves_local_references_with_siblings_and_terminates_cycles() { + let result = lowered(json!({ + "$defs": { + "Mode": { "const": "fast", "description": "original" }, + "Node": { "type": "object", "properties": { "next": { "$ref": "#/$defs/Node" } } } + }, + "type": "object", + "properties": { + "mode": { "$ref": "#/$defs/Mode", "description": "override" }, + "again": { "$ref": "#/$defs/Mode" }, + "node": { "$ref": "#/$defs/Node" }, + "missing": { "$ref": "#/$defs/Missing", "type": "string" }, + "external": { "$ref": "https://example.test/schema", "description": "external" } + } + })); + assert_eq!( + result["properties"]["mode"], + json!({ "type": "string", "enum": ["fast"], "description": "override" }) + ); + assert_eq!(result["properties"]["again"]["enum"], json!(["fast"])); + assert_eq!( + result["properties"]["node"]["properties"]["next"], + json!({}) + ); + assert_eq!(result["properties"]["missing"], json!({ "type": "string" })); + assert_eq!( + result["properties"]["external"], + json!({ "description": "external" }) + ); + assert!(result.get("$defs").is_none()); + } + + #[test] + fn lowers_nullable_type_unions_and_non_string_constants_without_invalid_enums() { + assert_eq!( + lowered(json!({ "type": ["string", "null", "string"] })), + json!({ "type": "string", "nullable": true }) + ); + assert_eq!( + lowered(json!({ "type": ["string", "number"] })), + json!({ "anyOf": [{ "type": "number" }, { "type": "string" }] }) + ); + for (ty, value) in [ + ("integer", json!(42)), + ("boolean", json!(true)), + ("null", Value::Null), + ] { + let result = lowered(json!({ "type": ty, "const": value })); + assert_eq!(result["type"], ty); + assert_eq!(result["description"], format!("Must equal: {value}")); + assert!(result.get("enum").is_none()); + assert!(result.get("const").is_none()); + } + assert_eq!( + lowered(json!({ "type": "integer", "enum": [1, 2] })), + json!({ "type": "integer" }) + ); + } + + #[test] + fn converts_protobuf_integer_strings_to_draft_2020_numbers() { + let result = lowered(json!({ + "type": ["OBJECT", "not-a-json-schema-type"], + "minProperties": "1", + "maxProperties": "invalid", + "required": ["args", "args"], + "enum": [], + "anyOf": [], + "properties": { + "args": { + "type": "ARRAY", + "minItems": "2", + "maxItems": "3", + "minLength": "4", + "maxLength": 8, + "minimum": "invalid", + "maximum": 1, + "items": { "type": "STRING", "minLength": "0" } + } + } + })); + assert_eq!( + result, + json!({ + "type": "object", + "minProperties": 1, + "required": ["args"], + "properties": { + "args": { + "type": "array", + "minItems": 2, + "maxItems": 3, + "minLength": 4, + "maxLength": 8, + "maximum": 1, + "items": { "type": "string", "minLength": 0 } + } + } + }) + ); + } + + #[test] + fn handles_boolean_and_deep_schemas_without_panicking() { + for input in [json!(true), json!(false), Value::Null] { + assert_eq!(lowered(input), json!({})); + } + let mut nested = json!({ "const": "deep" }); + for _ in 0..70 { + nested = json!({ "type": "array", "items": nested }); + } + let result = lowered(nested); + let mut child = &result; + for _ in 0..64 { + assert_eq!(child["type"], "array"); + child = &child["items"]; + } + assert_eq!(child, &json!({})); + } +} diff --git a/crates/aether-provider/transport/src/claude_code/fingerprint.rs b/crates/aether-provider/transport/src/claude_code/fingerprint.rs index c72e58e93..6796cdb9d 100644 --- a/crates/aether-provider/transport/src/claude_code/fingerprint.rs +++ b/crates/aether-provider/transport/src/claude_code/fingerprint.rs @@ -117,11 +117,11 @@ mod tests { let map = header_fingerprint_from_fingerprint(&fp).expect("header fingerprint"); assert_eq!(map["identity_profile_version"], "2026-04"); - assert_eq!(map["cli_version"], "2.1.161"); - assert_eq!(map["billing_cli_version"], "2.1.161"); - assert_eq!(map["stainless_package_version"], "0.94.0"); - assert_eq!(map["stainless_runtime_version"], "v24.3.0"); - assert_eq!(map["user_agent"], "claude-cli/2.1.161 (external, cli)"); + assert_eq!(map["cli_version"], "2.1.284"); + assert_eq!(map["billing_cli_version"], "2.1.284"); + assert_eq!(map["stainless_package_version"], "0.112.1"); + assert_eq!(map["stainless_runtime_version"], "v26.3.0"); + assert_eq!(map["user_agent"], "claude-cli/2.1.284 (external, cli)"); assert_eq!( fp["transport_profile"]["extra"]["claude_code_identity_profile_version"], "2026-04" @@ -144,9 +144,9 @@ mod tests { let sanitized = sanitize_fingerprint(&raw, "test-key"); let map = header_fingerprint_from_fingerprint(&sanitized).expect("header fingerprint"); - assert_eq!(map["stainless_package_version"], "0.94.0"); - assert_eq!(map["stainless_runtime_version"], "v24.3.0"); - assert_eq!(map["user_agent"], "claude-cli/2.1.161 (external, cli)"); + assert_eq!(map["stainless_package_version"], "0.112.1"); + assert_eq!(map["stainless_runtime_version"], "v26.3.0"); + assert_eq!(map["user_agent"], "claude-cli/2.1.284 (external, cli)"); assert_eq!(map["vscode_session_id"], "existing-session"); } diff --git a/crates/aether-provider/transport/src/claude_code/mimicry.rs b/crates/aether-provider/transport/src/claude_code/mimicry.rs new file mode 100644 index 000000000..4ead9b97a --- /dev/null +++ b/crates/aether-provider/transport/src/claude_code/mimicry.rs @@ -0,0 +1,674 @@ +use serde_json::{json, Map, Value}; +use sha2::{Digest, Sha256}; + +use super::profile::current_claude_code_transport_identity_profile; +use crate::snapshot::GatewayProviderTransportSnapshot; + +const BILLING_HEADER_PREFIX: &str = "x-anthropic-billing-header"; +const BILLING_ENTRYPOINT_MARKER: &str = "cc_entrypoint="; +const CLAUDE_CODE_SYSTEM_PROMPT: &str = "You are Claude Code, Anthropic's official CLI for Claude."; +const FINGERPRINT_SALT: &str = "59cf53e54c78"; +const FINGERPRINT_CHAR_INDICES: [usize; 3] = [4, 7, 20]; +const DEFAULT_MAX_TOKENS: u64 = 128_000; +const MAX_CACHE_CONTROL_BLOCKS: usize = 4; +const SYSTEM_INSTRUCTIONS_ACK: &str = "Understood. I will follow these instructions."; +const DEVICE_ID_SEED_NAMESPACE: &str = "aether-claude-code-device"; + +const CLAUDE_CODE_PROMPT_PREFIXES: [&str; 4] = [ + "You are Claude Code, Anthropic's official CLI for Claude", + "You are a Claude agent, built on Anthropic's Claude Agent SDK", + "You are a file search specialist for Claude Code", + "You are a helpful AI assistant tasked with summarizing conversations", +]; + +/// Tool-agnostic part of the real Claude Code system prompt. It brings the system block +/// count and size close to genuine CLI traffic without injecting tool-specific instructions. +const CLAUDE_CODE_SYSTEM_PROMPT_EXPANSION: &str = r#"You are an interactive agent that helps users with software engineering tasks. Use the instructions below and the tools available to you to assist the user. + +IMPORTANT: Assist with authorized security testing, defensive security, CTF challenges, and educational contexts. Refuse requests for destructive techniques, DoS attacks, mass targeting, supply chain compromise, or detection evasion for malicious purposes. Dual-use security tools (C2 frameworks, credential testing, exploit development) require clear authorization context: pentesting engagements, CTF competitions, security research, or defensive use cases. +IMPORTANT: You must NEVER generate or guess URLs for the user unless you are confident that the URLs are for helping the user with programming. You may use URLs provided by the user in their messages or local files. + +# Tone and style + - Only use emojis if the user explicitly requests it. Avoid using emojis in all communication unless asked. + - Your responses should be short and concise. + - When referencing specific functions or pieces of code include the pattern file_path:line_number to allow the user to easily navigate to the source code location. + - When referencing GitHub issues or pull requests, use the owner/repo#123 format (e.g. anthropics/claude-code#100) so they render as clickable links. + - Do not use a colon before tool calls. Your tool calls may not be shown directly in the output, so text like "Let me read the file:" followed by a read tool call should just be "Let me read the file." with a period."#; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct ClaudeCodeBodyMimicryContext<'a> { + pub key_id: &'a str, + pub device_id: Option<&'a str>, + pub account_uuid: Option<&'a str>, +} + +/// Applies Claude Code body mimicry for a claude_code provider transport (claude:messages only). +/// Returns whether the body was changed. +pub fn apply_claude_code_body_mimicry_for_transport( + provider_request_body: &mut Value, + transport: &GatewayProviderTransportSnapshot, + provider_api_format: &str, +) -> bool { + if !transport + .provider + .provider_type + .trim() + .eq_ignore_ascii_case("claude_code") + || !aether_ai_formats::normalize_api_format_alias(provider_api_format) + .eq_ignore_ascii_case("claude:messages") + { + return false; + } + + let auth_config = transport + .key + .decrypted_auth_config + .as_deref() + .and_then(|raw| serde_json::from_str::(raw).ok()); + let auth_config_str = |field: &str| { + auth_config + .as_ref() + .and_then(|config| config.get(field)) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }; + let device_id = auth_config_str("device_id"); + let account_uuid = auth_config_str("account_uuid"); + + apply_claude_code_body_mimicry( + provider_request_body, + ClaudeCodeBodyMimicryContext { + key_id: transport.key.id.as_str(), + device_id: device_id.as_deref(), + account_uuid: account_uuid.as_deref(), + }, + ) +} + +/// Makes a non-Claude-Code request body look like genuine Claude Code traffic so it is +/// accepted together with the Claude Code identity headers on OAuth credentials. +/// The operation is idempotent: a body that already carries the billing block and +/// `metadata.user_id` is treated as genuine and left untouched. +pub fn apply_claude_code_body_mimicry( + body: &mut Value, + context: ClaudeCodeBodyMimicryContext<'_>, +) -> bool { + let before = body.clone(); + let Some(object) = body.as_object_mut() else { + return false; + }; + if is_genuine_claude_code_body(object) { + return false; + } + + let profile = *current_claude_code_transport_identity_profile(); + let model = object + .get("model") + .and_then(Value::as_str) + .unwrap_or_default() + .to_ascii_lowercase(); + let first_user_text = first_user_text(object); + + rewrite_system_and_migrate_instructions( + object, + profile.billing_cli_version(), + &first_user_text, + model.contains("fable"), + ); + inject_metadata_user_id(object, &context, &first_user_text); + fill_fixed_fields(object, &model); + enforce_cache_control_limit(object); + + *body != before +} + +fn is_genuine_claude_code_body(body: &Map) -> bool { + let has_user_id = body + .get("metadata") + .and_then(|metadata| metadata.get("user_id")) + .and_then(Value::as_str) + .is_some_and(|value| !value.trim().is_empty()); + has_user_id + && body + .get("system") + .and_then(Value::as_array) + .is_some_and(|system| { + system.iter().any(|block| { + block + .get("text") + .and_then(Value::as_str) + .is_some_and(|text| { + text.starts_with(BILLING_HEADER_PREFIX) + && text.contains(BILLING_ENTRYPOINT_MARKER) + }) + }) + }) +} + +fn has_claude_code_prefix(text: &str) -> bool { + let text = text.trim_start(); + CLAUDE_CODE_PROMPT_PREFIXES + .iter() + .any(|prefix| text.starts_with(prefix)) +} + +fn first_user_text(body: &Map) -> String { + let Some(message) = body + .get("messages") + .and_then(Value::as_array) + .and_then(|messages| { + messages + .iter() + .find(|message| message.get("role").and_then(Value::as_str) == Some("user")) + }) + else { + return String::new(); + }; + match message.get("content") { + Some(Value::String(text)) => text.clone(), + Some(Value::Array(blocks)) => blocks + .iter() + .find(|block| block.get("type").and_then(Value::as_str) == Some("text")) + .and_then(|block| block.get("text")) + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(), + _ => String::new(), + } +} + +fn hex(bytes: &[u8]) -> String { + bytes.iter().map(|byte| format!("{byte:02x}")).collect() +} + +/// Three hex chars derived from the first user text, matching the Claude Code CLI +/// billing attribution fingerprint. Indexing is by byte, out-of-range yields `'0'`. +fn claude_code_fingerprint(first_user_text: &str, cli_version: &str) -> String { + let bytes = first_user_text.as_bytes(); + let picked = FINGERPRINT_CHAR_INDICES + .iter() + .map(|index| bytes.get(*index).copied().unwrap_or(b'0')) + .collect::>(); + let mut hasher = Sha256::new(); + hasher.update(FINGERPRINT_SALT.as_bytes()); + hasher.update(&picked); + hasher.update(cli_version.as_bytes()); + hex(&hasher.finalize())[..3].to_string() +} + +fn billing_header_text(first_user_text: &str, cli_version: &str) -> String { + format!( + "{BILLING_HEADER_PREFIX}: cc_version={cli_version}.{}; cc_entrypoint=cli;", + claude_code_fingerprint(first_user_text, cli_version) + ) +} + +/// Returns the joined original system text plus the cache_control of its last block that has one. +fn extract_system_text(system: Option<&Value>) -> (String, Option) { + match system { + Some(Value::String(text)) => (text.clone(), None), + Some(Value::Array(blocks)) => { + let mut parts = Vec::new(); + let mut cache_control = None; + for block in blocks { + if let Some(text) = block.get("text").and_then(Value::as_str) { + if !text.trim().is_empty() { + parts.push(text.to_string()); + } + } + if let Some(value) = block.get("cache_control").filter(|value| !value.is_null()) { + cache_control = Some(value.clone()); + } + } + (parts.join("\n\n"), cache_control) + } + _ => (String::new(), None), + } +} + +fn rewrite_system_and_migrate_instructions( + body: &mut Map, + cli_version: &str, + first_user_text: &str, + identity_only: bool, +) { + let (original_text, original_cache_control) = extract_system_text(body.get("system")); + + let mut system = vec![ + json!({"type": "text", "text": billing_header_text(first_user_text, cli_version)}), + json!({"type": "text", "text": CLAUDE_CODE_SYSTEM_PROMPT}), + ]; + if !identity_only { + system.push(json!({ + "type": "text", + "text": CLAUDE_CODE_SYSTEM_PROMPT_EXPANSION, + "cache_control": {"type": "ephemeral", "ttl": "5m"}, + })); + } + body.insert("system".to_string(), Value::Array(system)); + + let original_text = original_text.trim(); + if original_text.is_empty() + || original_text == CLAUDE_CODE_SYSTEM_PROMPT + || has_claude_code_prefix(original_text) + { + return; + } + + let mut instruction_block = json!({ + "type": "text", + "text": format!("[System Instructions]\n{original_text}"), + }); + if let Some(cache_control) = original_cache_control { + instruction_block["cache_control"] = cache_control; + } + let mut messages = vec![ + json!({"role": "user", "content": [instruction_block]}), + json!({"role": "assistant", "content": [{"type": "text", "text": SYSTEM_INSTRUCTIONS_ACK}]}), + ]; + if let Some(Value::Array(original)) = body.remove("messages") { + messages.extend(original); + } + body.insert("messages".to_string(), Value::Array(messages)); +} + +fn stable_device_id(context: &ClaudeCodeBodyMimicryContext<'_>) -> String { + if let Some(device_id) = context.device_id.filter(|value| !value.is_empty()) { + return device_id.to_string(); + } + let mut hasher = Sha256::new(); + hasher.update(format!("{DEVICE_ID_SEED_NAMESPACE}::{}", context.key_id).as_bytes()); + hex(&hasher.finalize()) +} + +/// Session id that stays stable while a conversation grows: seeded by key and first user text. +fn stable_session_id(key_id: &str, first_user_text: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(format!("{key_id}::{first_user_text}").as_bytes()); + let digest = hasher.finalize(); + let mut bytes = [0u8; 16]; + bytes.copy_from_slice(&digest[..16]); + bytes[6] = (bytes[6] & 0x0f) | 0x40; + bytes[8] = (bytes[8] & 0x3f) | 0x80; + format!( + "{}-{}-{}-{}-{}", + hex(&bytes[0..4]), + hex(&bytes[4..6]), + hex(&bytes[6..8]), + hex(&bytes[8..10]), + hex(&bytes[10..16]), + ) +} + +fn inject_metadata_user_id( + body: &mut Map, + context: &ClaudeCodeBodyMimicryContext<'_>, + first_user_text: &str, +) { + let has_user_id = body + .get("metadata") + .and_then(|metadata| metadata.get("user_id")) + .and_then(Value::as_str) + .is_some_and(|value| !value.trim().is_empty()); + if has_user_id { + return; + } + // Built by hand to keep the key order of the real CLI (serde_json sorts keys). + let json_string = |value: &str| Value::String(value.to_string()).to_string(); + let user_id = format!( + "{{\"device_id\":{},\"account_uuid\":{},\"session_id\":{}}}", + json_string(&stable_device_id(context)), + json_string(context.account_uuid.unwrap_or_default()), + json_string(&stable_session_id(context.key_id, first_user_text)), + ); + match body.get_mut("metadata").and_then(Value::as_object_mut) { + Some(metadata) => { + metadata.insert("user_id".to_string(), Value::String(user_id)); + } + None => { + body.insert("metadata".to_string(), json!({"user_id": user_id})); + } + } +} + +fn is_opus_5_5(model: &str) -> bool { + model.contains("opus-5-5") || model.contains("opus-5.5") +} + +fn is_signed_thinking_5_5(model: &str) -> bool { + is_opus_5_5(model) || model.contains("sonnet-5-5") || model.contains("sonnet-5.5") +} + +fn fill_fixed_fields(body: &mut Map, model: &str) { + body.entry("tools".to_string()).or_insert_with(|| json!([])); + if !body.contains_key("temperature") && !is_opus_5_5(model) { + body.insert("temperature".to_string(), json!(1)); + } + body.entry("max_tokens".to_string()) + .or_insert_with(|| json!(DEFAULT_MAX_TOKENS)); + + let tools_empty = body + .get("tools") + .and_then(Value::as_array) + .is_none_or(Vec::is_empty); + if tools_empty && !is_signed_thinking_5_5(model) { + body.remove("tool_choice"); + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum CacheControlSlot { + Tool(usize), + Message(usize, usize), + System(usize), +} + +fn cache_control_slots(body: &Map) -> Vec { + let has = |value: &Value| value.get("cache_control").is_some_and(|cc| !cc.is_null()); + let mut slots = Vec::new(); + let collect = |key: &str, make: &dyn Fn(usize) -> CacheControlSlot, slots: &mut Vec<_>| { + if let Some(items) = body.get(key).and_then(Value::as_array) { + for (index, item) in items.iter().enumerate() { + if has(item) { + slots.push(make(index)); + } + } + } + }; + collect("tools", &CacheControlSlot::Tool, &mut slots); + if let Some(messages) = body.get("messages").and_then(Value::as_array) { + for (message_index, message) in messages.iter().enumerate() { + if let Some(blocks) = message.get("content").and_then(Value::as_array) { + for (block_index, block) in blocks.iter().enumerate() { + if has(block) { + slots.push(CacheControlSlot::Message(message_index, block_index)); + } + } + } + } + } + collect("system", &CacheControlSlot::System, &mut slots); + slots +} + +fn remove_cache_control(body: &mut Map, slot: &CacheControlSlot) { + let target = match slot { + CacheControlSlot::Tool(index) => body + .get_mut("tools") + .and_then(|tools| tools.get_mut(*index)), + CacheControlSlot::Message(message, block) => body + .get_mut("messages") + .and_then(|messages| messages.get_mut(*message)) + .and_then(|message| message.get_mut("content")) + .and_then(|content| content.get_mut(*block)), + CacheControlSlot::System(index) => body + .get_mut("system") + .and_then(|system| system.get_mut(*index)), + }; + if let Some(object) = target.and_then(Value::as_object_mut) { + object.remove("cache_control"); + } +} + +/// Anthropic accepts at most 4 cache breakpoints. Drop the excess in the same order the +/// upstream-friendly proxy does: tools (last first), messages (first first), system (last first). +fn enforce_cache_control_limit(body: &mut Map) { + let slots = cache_control_slots(body); + if slots.len() <= MAX_CACHE_CONTROL_BLOCKS { + return; + } + let mut excess = slots.len() - MAX_CACHE_CONTROL_BLOCKS; + let tools = slots + .iter() + .filter(|slot| matches!(slot, CacheControlSlot::Tool(_))) + .rev(); + let messages = slots + .iter() + .filter(|slot| matches!(slot, CacheControlSlot::Message(..))); + let system = slots + .iter() + .filter(|slot| matches!(slot, CacheControlSlot::System(_))) + .rev(); + for slot in tools.chain(messages).chain(system) { + if excess == 0 { + break; + } + remove_cache_control(body, slot); + excess -= 1; + } +} + +#[cfg(test)] +mod tests { + use serde_json::{json, Value}; + + use super::{ + apply_claude_code_body_mimicry, claude_code_fingerprint, stable_session_id, + ClaudeCodeBodyMimicryContext, CLAUDE_CODE_SYSTEM_PROMPT, + }; + + const CONTEXT: ClaudeCodeBodyMimicryContext<'static> = ClaudeCodeBodyMimicryContext { + key_id: "key-1", + device_id: None, + account_uuid: Some("acct-1"), + }; + + fn pi_body() -> Value { + json!({ + "model": "claude-sonnet-5-5", + "max_tokens": 4096, + "stream": true, + "system": [{ + "type": "text", + "text": "You are an expert coding assistant operating inside pi.", + "cache_control": {"type": "ephemeral"} + }], + "messages": [{"role": "user", "content": "hello world, please help"}], + "tools": [{"name": "fabric_exec", "input_schema": {"type": "object"}}], + "thinking": {"type": "adaptive"} + }) + } + + #[test] + fn rewrites_system_into_claude_code_blocks_and_migrates_original_instructions() { + let mut body = pi_body(); + assert!(apply_claude_code_body_mimicry(&mut body, CONTEXT)); + + let system = body["system"].as_array().unwrap(); + assert_eq!(system.len(), 3); + let billing = system[0]["text"].as_str().unwrap(); + assert!(billing.starts_with("x-anthropic-billing-header: cc_version=2.1.284.")); + assert!(billing.ends_with("; cc_entrypoint=cli;")); + assert!(system[0].get("cache_control").is_none()); + assert_eq!(system[1]["text"], CLAUDE_CODE_SYSTEM_PROMPT); + assert!(system[1].get("cache_control").is_none()); + assert_eq!(system[2]["cache_control"]["type"], "ephemeral"); + assert_eq!(system[2]["cache_control"]["ttl"], "5m"); + + let messages = body["messages"].as_array().unwrap(); + assert_eq!(messages.len(), 3); + assert_eq!( + messages[0]["content"][0]["text"], + "[System Instructions]\nYou are an expert coding assistant operating inside pi." + ); + assert_eq!( + messages[0]["content"][0]["cache_control"]["type"], + "ephemeral" + ); + assert_eq!(messages[1]["role"], "assistant"); + assert_eq!(messages[2]["content"], "hello world, please help"); + } + + #[test] + fn fingerprint_uses_original_first_user_text_not_migrated_instructions() { + let mut body = pi_body(); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + let expected = claude_code_fingerprint("hello world, please help", "2.1.284"); + assert!(body["system"][0]["text"] + .as_str() + .unwrap() + .contains(&format!("cc_version=2.1.284.{expected};"))); + } + + #[test] + fn fingerprint_is_three_hex_chars_and_pads_short_text_with_zero() { + let fp = claude_code_fingerprint("", "2.1.284"); + assert_eq!(fp.len(), 3); + assert!(fp.chars().all(|c| c.is_ascii_hexdigit())); + assert_eq!(fp, claude_code_fingerprint("abc", "2.1.284")); + assert_ne!( + claude_code_fingerprint("0123456789012345678901234", "2.1.284"), + claude_code_fingerprint("", "2.1.284") + ); + } + + #[test] + fn injects_metadata_user_id_in_cli_key_order_with_stable_ids() { + let mut first = pi_body(); + let mut second = pi_body(); + apply_claude_code_body_mimicry(&mut first, CONTEXT); + apply_claude_code_body_mimicry(&mut second, CONTEXT); + + let user_id = first["metadata"]["user_id"].as_str().unwrap(); + assert!(user_id.starts_with("{\"device_id\":\"")); + let parsed: Value = serde_json::from_str(user_id).unwrap(); + assert_eq!(parsed["device_id"].as_str().unwrap().len(), 64); + assert_eq!(parsed["account_uuid"], "acct-1"); + assert_eq!(parsed["session_id"].as_str().unwrap().len(), 36); + assert_eq!(first["metadata"], second["metadata"]); + assert_eq!( + parsed["session_id"], + stable_session_id("key-1", "hello world, please help") + ); + } + + #[test] + fn keeps_client_metadata_user_id() { + let mut body = pi_body(); + body["metadata"] = json!({"user_id": "client-provided"}); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + assert_eq!(body["metadata"]["user_id"], "client-provided"); + } + + #[test] + fn is_idempotent_and_leaves_genuine_claude_code_bodies_untouched() { + let mut body = pi_body(); + assert!(apply_claude_code_body_mimicry(&mut body, CONTEXT)); + let once = body.clone(); + assert!(!apply_claude_code_body_mimicry(&mut body, CONTEXT)); + assert_eq!(body, once); + + let mut genuine = json!({ + "model": "claude-sonnet-5-5", + "system": [ + {"type": "text", "text": "x-anthropic-billing-header: cc_version=2.1.284.abc; cc_entrypoint=cli;"}, + {"type": "text", "text": CLAUDE_CODE_SYSTEM_PROMPT} + ], + "metadata": {"user_id": "{\"device_id\":\"d\",\"account_uuid\":\"\",\"session_id\":\"s\"}"}, + "messages": [{"role": "user", "content": "hi"}] + }); + let before = genuine.clone(); + assert!(!apply_claude_code_body_mimicry(&mut genuine, CONTEXT)); + assert_eq!(genuine, before); + } + + #[test] + fn does_not_duplicate_system_that_already_carries_claude_code_identity() { + let mut body = pi_body(); + body["system"] = json!(CLAUDE_CODE_SYSTEM_PROMPT); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + assert_eq!(body["messages"].as_array().unwrap().len(), 1); + assert_eq!(body["system"].as_array().unwrap().len(), 3); + } + + #[test] + fn fable_models_only_get_billing_and_identity_blocks() { + let mut body = pi_body(); + body["model"] = json!("claude-fable-5-1"); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + assert_eq!(body["system"].as_array().unwrap().len(), 2); + } + + #[test] + fn fills_fixed_fields_only_when_missing() { + let mut body = json!({ + "model": "claude-sonnet-5-5", + "messages": [{"role": "user", "content": "hi"}], + "tool_choice": {"type": "auto"} + }); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + assert_eq!(body["tools"], json!([])); + assert_eq!(body["temperature"], 1); + assert_eq!(body["max_tokens"], 128000); + // Sonnet 5.5 keeps tool_choice even without tools. + assert!(body.get("tool_choice").is_some()); + + let mut body = json!({ + "model": "claude-opus-5-5", + "temperature": 0.2, + "max_tokens": 10, + "messages": [{"role": "user", "content": "hi"}] + }); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + assert_eq!(body["temperature"], 0.2); + assert_eq!(body["max_tokens"], 10); + + let mut body = json!({ + "model": "claude-opus-5-5", + "messages": [{"role": "user", "content": "hi"}] + }); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + assert!(body.get("temperature").is_none()); + + let mut body = json!({ + "model": "claude-opus-4-6", + "messages": [{"role": "user", "content": "hi"}], + "tool_choice": {"type": "auto"} + }); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + assert!(body.get("tool_choice").is_none()); + } + + #[test] + fn caps_cache_control_breakpoints_at_four() { + let mut body = pi_body(); + body["tools"] = json!([ + {"name": "a", "cache_control": {"type": "ephemeral"}}, + {"name": "b", "cache_control": {"type": "ephemeral"}} + ]); + body["messages"] = json!([ + {"role": "user", "content": [{"type": "text", "text": "hello world, please help", "cache_control": {"type": "ephemeral"}}]}, + {"role": "user", "content": [{"type": "text", "text": "again", "cache_control": {"type": "ephemeral"}}]} + ]); + apply_claude_code_body_mimicry(&mut body, CONTEXT); + + let count = |value: &Value| value.get("cache_control").is_some() as usize; + let total: usize = body["tools"] + .as_array() + .unwrap() + .iter() + .map(count) + .sum::() + + body["system"] + .as_array() + .unwrap() + .iter() + .map(count) + .sum::() + + body["messages"] + .as_array() + .unwrap() + .iter() + .flat_map(|message| message["content"].as_array().cloned().unwrap_or_default()) + .map(|block| count(&block)) + .sum::(); + assert_eq!(total, 4); + // 6 breakpoints (2 tools, 3 messages incl. migrated instructions, 1 system): the two + // excess ones come from the tools, last first, so both tool breakpoints go. + assert!(body["tools"][0].get("cache_control").is_none()); + assert!(body["tools"][1].get("cache_control").is_none()); + assert!(body["system"][2].get("cache_control").is_some()); + } +} diff --git a/crates/aether-provider/transport/src/claude_code/mod.rs b/crates/aether-provider/transport/src/claude_code/mod.rs index 7317f54a0..28df0f09e 100644 --- a/crates/aether-provider/transport/src/claude_code/mod.rs +++ b/crates/aether-provider/transport/src/claude_code/mod.rs @@ -1,5 +1,6 @@ mod auth; mod fingerprint; +mod mimicry; mod policy; mod profile; mod request; @@ -10,6 +11,10 @@ pub use fingerprint::{ generate_fingerprint, generate_random_fingerprint, header_fingerprint_from_fingerprint, sanitize_fingerprint, }; +pub use mimicry::{ + apply_claude_code_body_mimicry, apply_claude_code_body_mimicry_for_transport, + ClaudeCodeBodyMimicryContext, +}; pub use policy::{ local_claude_code_transport_unsupported_reason_with_network, supports_local_claude_code_transport_with_network, diff --git a/crates/aether-provider/transport/src/claude_code/profile.rs b/crates/aether-provider/transport/src/claude_code/profile.rs index 45d15d340..8778ea8a2 100644 --- a/crates/aether-provider/transport/src/claude_code/profile.rs +++ b/crates/aether-provider/transport/src/claude_code/profile.rs @@ -82,13 +82,13 @@ pub const CLAUDE_CODE_TRANSPORT_IDENTITY_2026_04: ClaudeCodeTransportIdentityPro version: ClaudeCodeTransportIdentityProfileVersion::V2026_04, transport_profile_id: "claude_code_nodejs", anthropic_version: "2023-06-01", - cli_version: "2.1.161", + cli_version: "2.1.284", stainless_lang: "js", - stainless_package_version: "0.94.0", + stainless_package_version: "0.112.1", stainless_os: "Linux", stainless_arch: "arm64", stainless_runtime: "node", - stainless_runtime_version: "v24.3.0", + stainless_runtime_version: "v26.3.0", stainless_retry_count: "0", stainless_timeout: "600", message_required_betas: MESSAGE_BETAS_2026_04, @@ -286,14 +286,14 @@ mod tests { let profile = *current_claude_code_transport_identity_profile(); assert_eq!(profile.version().as_str(), "2026-04"); - assert_eq!(profile.cli_version(), "2.1.161"); + assert_eq!(profile.cli_version(), "2.1.284"); assert_eq!(profile.billing_cli_version(), profile.cli_version()); assert_eq!( profile.user_agent(), format!("claude-cli/{} (external, cli)", profile.cli_version()) ); - assert_eq!(profile.stainless_package_version(), "0.94.0"); - assert_eq!(profile.stainless_runtime_version(), "v24.3.0"); + assert_eq!(profile.stainless_package_version(), "0.112.1"); + assert_eq!(profile.stainless_runtime_version(), "v26.3.0"); } #[test] diff --git a/crates/aether-provider/transport/src/claude_code/request.rs b/crates/aether-provider/transport/src/claude_code/request.rs index 8e07352c7..8e2e15c6d 100644 --- a/crates/aether-provider/transport/src/claude_code/request.rs +++ b/crates/aether-provider/transport/src/claude_code/request.rs @@ -252,11 +252,11 @@ mod tests { assert_eq!(built.get("x-app").map(String::as_str), Some("cli")); assert_eq!( built.get("x-stainless-package-version").map(String::as_str), - Some("0.94.0") + Some("0.112.1") ); assert_eq!( built.get("x-stainless-runtime-version").map(String::as_str), - Some("v24.3.0") + Some("v26.3.0") ); assert_eq!( built.get("x-stainless-timeout").map(String::as_str), @@ -264,7 +264,7 @@ mod tests { ); assert_eq!( built.get("user-agent").map(String::as_str), - Some("claude-cli/2.1.161 (external, cli)") + Some("claude-cli/2.1.284 (external, cli)") ); assert_eq!( built.get("authorization").map(String::as_str), @@ -367,7 +367,7 @@ mod tests { assert_eq!( body["system"][0]["text"], - "x-anthropic-billing-header: cc_version=2.1.161.abc; cc_entrypoint=cli;" + "x-anthropic-billing-header: cc_version=2.1.284.abc; cc_entrypoint=cli;" ); } } diff --git a/crates/aether-provider/transport/src/conversion.rs b/crates/aether-provider/transport/src/conversion.rs index 90d02cd75..421f82557 100644 --- a/crates/aether-provider/transport/src/conversion.rs +++ b/crates/aether-provider/transport/src/conversion.rs @@ -754,6 +754,90 @@ mod tests { )); } + #[test] + fn xai_responses_transport_converts_standard_client_protocols() { + let transport = transport_snapshot("xai", "openai:responses", "oauth", true, None); + + for client_api_format in ["openai:chat", "claude:messages", "gemini:generate_content"] { + assert!( + request_pair_allowed_for_transport( + &transport, + client_api_format, + "openai:responses" + ), + "{client_api_format} should convert onto xAI Responses" + ); + assert_eq!( + candidate_transport_pair_skip_reason(&transport, client_api_format), + None + ); + } + assert!(request_conversion_transport_supported( + &transport, + RequestConversionKind::ToOpenAiResponses + )); + assert!(!request_pair_allowed_for_transport( + &transport, + "openai:image", + "openai:responses" + )); + assert!(!request_pair_allowed_for_transport( + &transport, + "openai:video", + "openai:responses" + )); + for isolated in ["openai:responses:compact", "openai:image", "openai:video"] { + assert!( + !request_pair_allowed_for_transport(&transport, isolated, "openai:responses"), + "{isolated} must not convert onto xAI Responses" + ); + } + } + + #[test] + fn xai_compact_and_media_endpoints_are_same_format_only() { + let compact = transport_snapshot("xai", "openai:responses:compact", "oauth", true, None); + assert!(request_pair_allowed_for_transport( + &compact, + "openai:responses:compact", + "openai:responses:compact" + )); + for client_api_format in [ + "openai:chat", + "openai:responses", + "claude:messages", + "gemini:generate_content", + ] { + assert!( + !request_pair_allowed_for_transport( + &compact, + client_api_format, + "openai:responses:compact" + ), + "{client_api_format} must not convert onto xAI compact" + ); + } + + for api_format in ["openai:image", "openai:video"] { + let transport = transport_snapshot("xai", api_format, "oauth", true, None); + assert!( + request_pair_allowed_for_transport(&transport, api_format, api_format), + "{api_format} same-format transport should be allowed" + ); + for client_api_format in [ + "openai:chat", + "openai:responses", + "claude:messages", + "gemini:generate_content", + ] { + assert!( + !request_pair_allowed_for_transport(&transport, client_api_format, api_format), + "{client_api_format} must not convert onto {api_format}" + ); + } + } + } + #[test] fn windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() { let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None); diff --git a/crates/aether-provider/transport/src/lib.rs b/crates/aether-provider/transport/src/lib.rs index 9fa3ab2ee..21411f369 100644 --- a/crates/aether-provider/transport/src/lib.rs +++ b/crates/aether-provider/transport/src/lib.rs @@ -30,6 +30,7 @@ pub mod url; pub mod vertex; mod video; pub mod windsurf; +pub mod xai; pub use aether_oauth as oauth; pub use agent_identity::{ @@ -195,3 +196,10 @@ pub use windsurf::{ local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH, WINDSURF_ENVELOPE_NAME, }; +pub use xai::{ + extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, + insert_cli_identity_headers, insert_cli_identity_headers_if_needed, is_xai_provider_transport, + resolved_xai_request_base_url, resolved_xai_upstream_base_url, + should_attach_cli_identity_headers, xai_auth_uses_api, xai_uses_official_api, XAI_API_BASE_URL, + XAI_CHAT_PROXY_BASE_URL, XAI_PROVIDER_TYPE, +}; diff --git a/crates/aether-provider/transport/src/openai_image/mod.rs b/crates/aether-provider/transport/src/openai_image/mod.rs index 2170611c6..2390d274d 100644 --- a/crates/aether-provider/transport/src/openai_image/mod.rs +++ b/crates/aether-provider/transport/src/openai_image/mod.rs @@ -84,6 +84,11 @@ fn is_dedicated_openai_image_provider(transport: &GatewayProviderTransportSnapsh .trim() .eq_ignore_ascii_case("codex") || is_grok_provider_transport(transport) + || transport + .provider + .provider_type + .trim() + .eq_ignore_ascii_case("xai") } pub fn resolve_openai_image_auth( @@ -92,7 +97,10 @@ pub fn resolve_openai_image_auth( if is_grok_provider_transport(transport) { return resolve_grok_session_auth(transport); } - resolve_local_openai_bearer_auth(transport) + resolve_local_openai_bearer_auth(transport).or_else(|| { + crate::generic_oauth::resolve_local_generic_oauth_transport_authorization(transport) + .map(|value| ("authorization".to_string(), value)) + }) } pub fn build_openai_image_upstream_url( @@ -100,7 +108,11 @@ pub fn build_openai_image_upstream_url( request_path: Option<&str>, request_query: Option<&str>, ) -> String { - build_openai_image_url(&transport.endpoint.base_url, request_path, request_query) + build_openai_image_url( + &crate::xai::resolved_xai_request_base_url(transport, "openai:image"), + request_path, + request_query, + ) } pub fn build_openai_image_headers( @@ -113,6 +125,11 @@ pub fn build_openai_image_headers( &BTreeMap::new(), ); provider_request_headers.insert("content-type".to_string(), "application/json".to_string()); + crate::xai::insert_cli_identity_headers_if_needed( + input.transport, + "openai:image", + &mut provider_request_headers, + ); if let Some(accept) = input.accept { provider_request_headers.insert("accept".to_string(), accept.to_string()); } else { @@ -280,6 +297,48 @@ mod tests { ); } + #[test] + fn xai_oauth_image_uses_cli_proxy() { + let mut transport = sample_transport(); + transport.provider.provider_type = "xai".to_string(); + transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".to_string(); + transport.key.auth_type = "oauth".to_string(); + transport.key.decrypted_auth_config = + Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string()); + + assert_eq!( + openai_image_transport_unsupported_reason(&transport, "openai:image"), + None + ); + assert_eq!( + build_openai_image_upstream_url(&transport, Some("/v1/images/generations"), None), + "https://cli-chat-proxy.grok.com/v1/images/generations" + ); + assert_eq!( + build_openai_image_upstream_url(&transport, Some("/v1/images/edits"), None), + "https://cli-chat-proxy.grok.com/v1/images/edits" + ); + let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput { + transport: &transport, + headers: &HeaderMap::new(), + auth_header: "authorization", + auth_value: "Bearer test-token", + accept: None, + header_rules: None, + provider_request_body: &json!({"prompt": "A cat"}), + original_request_body: &json!({"prompt": "A cat"}), + }) + .unwrap(); + assert_eq!( + headers.get("x-xai-token-auth").map(String::as_str), + Some("xai-grok-cli") + ); + assert_eq!( + headers.get("authorization").map(String::as_str), + Some("Bearer test-token") + ); + } + #[test] fn codex_is_supported_by_dedicated_openai_image_transport_policy() { let mut transport = sample_transport(); diff --git a/crates/aether-provider/transport/src/provider_types.rs b/crates/aether-provider/transport/src/provider_types.rs index 4e9045c8d..cac278079 100644 --- a/crates/aether-provider/transport/src/provider_types.rs +++ b/crates/aether-provider/transport/src/provider_types.rs @@ -275,6 +275,17 @@ const WINDSURF_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy { ..STANDARD_RUNTIME_POLICY }; +const XAI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy { + fixed_provider: true, + api_format_inheritance: ProviderApiFormatInheritance::OAuthOrBearer, + enable_format_conversion_by_default: true, + oauth_is_bearer_like: true, + supports_model_fetch: false, + supports_local_openai_chat_transport: false, + supports_local_same_format_transport: true, + ..STANDARD_RUNTIME_POLICY +}; + const CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate { provider_type: "claude_code", version: 2, @@ -446,6 +457,39 @@ const WINDSURF_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTem runtime_policy: WINDSURF_RUNTIME_POLICY, }; +const XAI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate { + provider_type: "xai", + version: 2, + base_url: crate::xai::XAI_CHAT_PROXY_BASE_URL, + endpoints: &[ + FixedProviderEndpointTemplate { + item_key: "openai:responses", + api_format: "openai:responses", + custom_path: None, + config_defaults: FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS, + }, + FixedProviderEndpointTemplate { + item_key: "openai:responses:compact", + api_format: "openai:responses:compact", + custom_path: None, + config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS, + }, + FixedProviderEndpointTemplate { + item_key: "openai:image", + api_format: "openai:image", + custom_path: None, + config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS, + }, + FixedProviderEndpointTemplate { + item_key: "openai:video", + api_format: "openai:video", + custom_path: None, + config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS, + }, + ], + runtime_policy: XAI_RUNTIME_POLICY, +}; + pub fn provider_type_is_fixed(provider_type: &str) -> bool { provider_runtime_policy(provider_type).fixed_provider } @@ -498,6 +542,7 @@ pub fn fixed_provider_template(provider_type: &str) -> Option<&'static FixedProv "vertex_ai" => Some(&VERTEX_AI_FIXED_PROVIDER_TEMPLATE), "antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE), "windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE), + "xai" => Some(&XAI_FIXED_PROVIDER_TEMPLATE), _ => None, } } @@ -613,6 +658,16 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option Some(ProviderOAuthTemplate { + provider_type: "xai", + display_name: "xAI", + authorize_url: aether_oauth::provider::providers::XAI_DEVICE_CODE_URL, + token_url: aether_oauth::provider::providers::XAI_TOKEN_URL, + client_id: aether_oauth::provider::providers::XAI_CLIENT_ID, + scopes: aether_oauth::provider::providers::XAI_OAUTH_SCOPES, + redirect_uri: "", + use_pkce: false, + }), _ => None, } } @@ -825,6 +880,50 @@ mod tests { assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf")); } + #[test] + fn xai_fixed_provider_template_exposes_responses_media_endpoints() { + let template = fixed_provider_template("xai").expect("xai template should exist"); + assert_eq!(template.provider_type, "xai"); + assert_eq!(template.base_url, crate::xai::XAI_CHAT_PROXY_BASE_URL); + assert_eq!(template.version, 2); + assert_eq!( + template + .endpoints + .iter() + .map(|item| item.api_format) + .collect::>(), + vec![ + "openai:responses", + "openai:responses:compact", + "openai:image", + "openai:video" + ] + ); + + let policy = provider_runtime_policy("xai"); + assert!(policy.fixed_provider); + assert!(policy.enable_format_conversion_by_default); + assert!(policy.oauth_is_bearer_like); + assert!(!policy.supports_model_fetch); + assert!(policy.supports_local_same_format_transport); + assert!(!policy.supports_local_openai_chat_transport); + assert!(fixed_provider_key_inherits_api_formats( + "xai", "oauth", None + )); + assert!(fixed_provider_key_inherits_api_formats( + "xai", "bearer", None + )); + + let template = provider_type_admin_oauth_template("xai").expect("xai oauth template"); + assert_eq!(template.provider_type, "xai"); + assert_eq!(template.display_name, "xAI"); + assert_eq!( + template.token_url, + aether_oauth::provider::providers::XAI_TOKEN_URL + ); + assert!(!ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"xai")); + } + #[test] fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() { assert!(fixed_provider_key_inherits_api_formats( diff --git a/crates/aether-provider/transport/src/request_body.rs b/crates/aether-provider/transport/src/request_body.rs index 3fda480cb..9994f0563 100644 --- a/crates/aether-provider/transport/src/request_body.rs +++ b/crates/aether-provider/transport/src/request_body.rs @@ -1,6 +1,8 @@ use serde_json::{Map, Value}; -use crate::claude_code::sanitize_claude_code_request_body; +use crate::claude_code::{ + apply_claude_code_body_mimicry_for_transport, sanitize_claude_code_request_body, +}; use crate::snapshot::GatewayProviderTransportSnapshot; use crate::vertex::is_vertex_transport_context; @@ -40,8 +42,18 @@ pub fn apply_transport_request_body_semantics( .trim() .eq_ignore_ascii_case("claude_code") { + apply_claude_code_body_mimicry_for_transport( + provider_request_body, + transport, + provider_api_format.as_str(), + ); sanitize_claude_code_request_body(provider_request_body); } + aether_ai_formats::apply_xai_upstream_payload_edits( + provider_request_body, + transport.provider.provider_type.as_str(), + provider_api_format.as_str(), + ); if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) { apply_vertex_gemini_embedding_body_semantics(provider_request_body)?; } @@ -446,4 +458,37 @@ mod tests { assert!(error.message().contains("cannot be mapped")); assert!(body.get("input").is_some()); } + + #[test] + fn claude_code_body_mimicry_is_applied_by_default_to_claude_code_only() { + let pi_body = || { + json!({ + "model": "claude-sonnet-5-5", + "max_tokens": 1024, + "system": "You are an expert coding assistant operating inside pi.", + "messages": [{"role": "user", "content": "hello world, please help"}] + }) + }; + + let transport = sample_transport("claude_code", "https://api.anthropic.com"); + let mut body = pi_body(); + apply_transport_request_body_semantics(&mut body, &transport, "claude:messages") + .expect("semantics should apply"); + assert_eq!(body["system"].as_array().map(Vec::len), Some(3)); + assert!(body["metadata"]["user_id"].as_str().is_some()); + assert_eq!(body["messages"].as_array().map(Vec::len), Some(3)); + + // Running the semantics twice (cross-format planners do) must not stack the rewrite. + let once = body.clone(); + apply_transport_request_body_semantics(&mut body, &transport, "claude:messages") + .expect("semantics should apply"); + assert_eq!(body, once); + + // Non-claude_code providers are never touched. + let other = sample_transport("custom", "https://example.com"); + let mut untouched = pi_body(); + apply_transport_request_body_semantics(&mut untouched, &other, "claude:messages") + .expect("semantics should apply"); + assert_eq!(untouched, pi_body()); + } } diff --git a/crates/aether-provider/transport/src/request_url/mod.rs b/crates/aether-provider/transport/src/request_url/mod.rs index 22dc905c6..43168999c 100644 --- a/crates/aether-provider/transport/src/request_url/mod.rs +++ b/crates/aether-provider/transport/src/request_url/mod.rs @@ -120,6 +120,12 @@ fn build_transport_request_url_inner( return Some(url); } + let xai_base = + crate::xai::resolved_xai_upstream_base_url(transport, &normalized_provider_api_format); + let request_base_url = xai_base + .as_deref() + .unwrap_or(transport.endpoint.base_url.as_str()); + let custom_path_template = transport .endpoint .custom_path @@ -164,7 +170,7 @@ fn build_transport_request_url_inner( path.to_string() }; let mut url = build_passthrough_path_url( - &transport.endpoint.base_url, + request_base_url, normalized_path.as_str(), params.request_query, blocked_keys, @@ -190,75 +196,68 @@ fn build_transport_request_url_inner( let url = match normalized_provider_api_format.as_str() { "openai:chat" => Some(build_openai_chat_url( - &transport.endpoint.base_url, + request_base_url, params.request_query, )), "openai:responses" => Some(build_openai_responses_url( - &transport.endpoint.base_url, + request_base_url, params.request_query, false, )), "openai:responses:compact" => Some(build_openai_responses_url( - &transport.endpoint.base_url, + request_base_url, params.request_query, true, )), "openai:search" => Some(build_openai_search_url( - &transport.endpoint.base_url, + request_base_url, params.request_query, )), "openai:realtime" => build_passthrough_path_url( - &transport.endpoint.base_url, + request_base_url, "/v1/realtime", params.request_query, GATEWAY_CREDENTIAL_QUERY_KEYS, ) .and_then(|url| replace_realtime_model_query(url, params.mapped_model?)), "codex:live" => build_passthrough_path_url( - &transport.endpoint.base_url, + request_base_url, "/live", params.request_query, GATEWAY_CREDENTIAL_QUERY_KEYS, ), "openai:embedding" | "jina:embedding" => { - build_provider_embedding_v1_url(&transport.endpoint.base_url, params.request_query) + build_provider_embedding_v1_url(request_base_url, params.request_query) + } + "aliyun:multimodal_embedding" => { + build_aliyun_multimodal_embedding_url(request_base_url, params.request_query) } - "aliyun:multimodal_embedding" => build_aliyun_multimodal_embedding_url( - &transport.endpoint.base_url, - params.request_query, - ), "openai:rerank" | "jina:rerank" => { - build_provider_rerank_v1_url(&transport.endpoint.base_url, params.request_query) + build_provider_rerank_v1_url(request_base_url, params.request_query) } "claude:messages" => Some(if is_claude_count_tokens { - build_default_claude_count_tokens_url( - &transport.endpoint.base_url, - params.request_query, - ) + build_default_claude_count_tokens_url(request_base_url, params.request_query) } else { - build_claude_messages_url(&transport.endpoint.base_url, params.request_query) + build_claude_messages_url(request_base_url, params.request_query) }), "gemini:generate_content" => build_gemini_content_url( - &transport.endpoint.base_url, + request_base_url, params.mapped_model?, params.upstream_is_stream, params.request_query, ), "gemini:embedding" => build_gemini_embedding_url( - &transport.endpoint.base_url, + request_base_url, params.mapped_model?, params.request_query, gemini_embedding_batch, ), "gemini:interactions" => { - build_gemini_interactions_url(&transport.endpoint.base_url, params.request_query) + build_gemini_interactions_url(request_base_url, params.request_query) + } + "doubao:embedding" => { + build_passthrough_path_url(request_base_url, "/embeddings", params.request_query, &[]) } - "doubao:embedding" => build_passthrough_path_url( - &transport.endpoint.base_url, - "/embeddings", - params.request_query, - &[], - ), _ => None, }?; @@ -2417,4 +2416,82 @@ mod tests { "https://api.example.com/v1/messages?model=claude%26admin%3Dtrue%23fragment" ); } + + #[test] + fn xai_oauth_responses_use_cli_chat_proxy() { + let mut transport = sample_transport( + "xai", + "openai:responses", + "https://cli-chat-proxy.grok.com/v1", + None, + ); + transport.key.auth_type = "oauth".to_string(); + transport.key.decrypted_auth_config = + Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string()); + + let url = build_transport_request_url( + &transport, + TransportRequestUrlParams { + provider_api_format: "openai:responses", + mapped_model: Some("grok-4"), + upstream_is_stream: true, + request_query: None, + kiro_api_region: None, + api_operation: None, + }, + ) + .expect("xai oauth responses URL"); + + assert_eq!(url, "https://cli-chat-proxy.grok.com/v1/responses"); + } + + #[test] + fn xai_compact_and_using_api_use_official_api() { + let mut oauth = sample_transport( + "xai", + "openai:responses:compact", + "https://cli-chat-proxy.grok.com/v1", + None, + ); + oauth.key.auth_type = "oauth".to_string(); + oauth.key.decrypted_auth_config = + Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string()); + + let compact = build_transport_request_url( + &oauth, + TransportRequestUrlParams { + provider_api_format: "openai:responses:compact", + mapped_model: Some("grok-4"), + upstream_is_stream: false, + request_query: None, + kiro_api_region: None, + api_operation: None, + }, + ) + .expect("xai compact URL"); + assert_eq!(compact, "https://api.x.ai/v1/responses/compact"); + + let mut api_key = sample_transport( + "xai", + "openai:responses", + "https://cli-chat-proxy.grok.com/v1", + None, + ); + api_key.key.auth_type = "oauth".to_string(); + api_key.key.decrypted_auth_config = Some(r#"{"using_api":true}"#.to_string()); + + let official = build_transport_request_url( + &api_key, + TransportRequestUrlParams { + provider_api_format: "openai:responses", + mapped_model: Some("grok-4"), + upstream_is_stream: true, + request_query: None, + kiro_api_region: None, + api_operation: None, + }, + ) + .expect("xai api key URL"); + assert_eq!(official, "https://api.x.ai/v1/responses"); + } } diff --git a/crates/aether-provider/transport/src/same_format_provider/mod.rs b/crates/aether-provider/transport/src/same_format_provider/mod.rs index d84b3d519..3f3fa6358 100644 --- a/crates/aether-provider/transport/src/same_format_provider/mod.rs +++ b/crates/aether-provider/transport/src/same_format_provider/mod.rs @@ -1189,7 +1189,7 @@ mod tests { ); assert_eq!( build_body(compat_behavior)["system"][0]["text"], - "x-anthropic-billing-header: cc_version=2.1.161.abc; cc_entrypoint=cli;" + "x-anthropic-billing-header: cc_version=2.1.284.abc; cc_entrypoint=cli;" ); let mut legacy = sample_transport("claude_code"); diff --git a/crates/aether-provider/transport/src/standard/mod.rs b/crates/aether-provider/transport/src/standard/mod.rs index 56b5c0402..c5ed29970 100644 --- a/crates/aether-provider/transport/src/standard/mod.rs +++ b/crates/aether-provider/transport/src/standard/mod.rs @@ -396,6 +396,12 @@ pub fn build_standard_provider_request_headers( force_identity_accept_encoding(&mut headers); } + crate::xai::insert_cli_identity_headers_if_needed( + input.transport, + input.provider_api_format, + &mut headers, + ); + let declared_connection_headers = crate::headers::declared_connection_header_names(input.headers, input.extra_headers); crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers); diff --git a/crates/aether-provider/transport/src/video/mod.rs b/crates/aether-provider/transport/src/video/mod.rs index 85dfa0cd9..69af710bd 100644 --- a/crates/aether-provider/transport/src/video/mod.rs +++ b/crates/aether-provider/transport/src/video/mod.rs @@ -1,6 +1,7 @@ use std::collections::BTreeMap; use std::fmt; +use aether_contracts::ProxySnapshot; use aether_data_contracts::repository::video_tasks::StoredVideoTask; use aether_video_tasks_core::{ LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput, @@ -12,11 +13,13 @@ use super::auth::{ build_passthrough_headers_with_auth, resolve_local_gemini_auth, resolve_local_openai_bearer_auth, }; -use super::network::{resolve_transport_execution_timeouts, resolve_transport_profile}; +use super::network::{ + resolve_transport_execution_timeouts, resolve_transport_profile, + resolve_transport_proxy_snapshot, +}; use super::policy::{ local_gemini_transport_unsupported_reason_with_network, - local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport, - supports_local_standard_transport, + local_standard_transport_unsupported_reason_with_network, }; use super::rules::{ apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers, @@ -32,6 +35,7 @@ pub enum ProviderVideoCreateFamily { #[derive(Clone, Copy)] pub struct ProviderVideoCreateHeadersInput<'a> { + pub transport: &'a GatewayProviderTransportSnapshot, pub headers: &'a http::HeaderMap, pub auth_header: &'a str, pub auth_value: &'a str, @@ -79,6 +83,13 @@ pub trait VideoTaskTransportSnapshotLookup: Send + Sync { endpoint_id: &str, key_id: &str, ) -> Result, String>; + + async fn resolve_video_task_proxy( + &self, + transport: &GatewayProviderTransportSnapshot, + ) -> Option { + resolve_transport_proxy_snapshot(transport) + } } pub fn resolve_local_video_task_transport( @@ -89,13 +100,17 @@ pub fn resolve_local_video_task_transport( let api_format = api_format.trim(); let (auth_header, auth_value) = match api_format { "openai:video" => { - if !supports_local_standard_transport(transport, api_format) { + if local_standard_transport_unsupported_reason_with_network(transport, api_format) + .is_some() + { return None; } - resolve_local_openai_bearer_auth(transport)? + resolve_openai_compatible_video_auth(transport)? } "gemini:video" => { - if !supports_local_gemini_transport(transport, api_format) { + if local_gemini_transport_unsupported_reason_with_network(transport, api_format) + .is_some() + { return None; } resolve_local_gemini_auth(transport)? @@ -103,9 +118,9 @@ pub fn resolve_local_video_task_transport( _ => return None, }; - Some(LocalVideoTaskTransport::from_bridge_input( - LocalVideoTaskTransportBridgeInput { - upstream_base_url: transport.endpoint.base_url.clone(), + let mut resolved = + LocalVideoTaskTransport::from_bridge_input(LocalVideoTaskTransportBridgeInput { + upstream_base_url: crate::xai::resolved_xai_request_base_url(transport, api_format), provider_name: Some(transport.provider.name.clone()), provider_id: transport.provider.id.clone(), endpoint_id: transport.endpoint.id.clone(), @@ -114,11 +129,12 @@ pub fn resolve_local_video_task_transport( auth_value, content_type: Some("application/json".to_string()), model_name, - proxy: None, + proxy: resolve_transport_proxy_snapshot(transport), transport_profile: resolve_transport_profile(transport), timeouts: resolve_transport_execution_timeouts(transport), - }, - )) + }); + crate::xai::insert_cli_identity_headers_if_needed(transport, api_format, &mut resolved.headers); + Some(resolved) } pub fn video_create_transport_unsupported_reason( @@ -141,7 +157,7 @@ pub fn resolve_video_create_auth( family: ProviderVideoCreateFamily, ) -> Option<(String, String)> { match family { - ProviderVideoCreateFamily::OpenAi => resolve_local_openai_bearer_auth(transport), + ProviderVideoCreateFamily::OpenAi => resolve_openai_compatible_video_auth(transport), ProviderVideoCreateFamily::Gemini => resolve_local_gemini_auth(transport), } } @@ -173,6 +189,15 @@ pub fn build_video_create_request_body( Some(provider_request_body) } +fn resolve_openai_compatible_video_auth( + transport: &GatewayProviderTransportSnapshot, +) -> Option<(String, String)> { + resolve_local_openai_bearer_auth(transport).or_else(|| { + crate::generic_oauth::resolve_local_generic_oauth_transport_authorization(transport) + .map(|value| ("authorization".to_string(), value)) + }) +} + pub fn build_video_create_upstream_url( transport: &GatewayProviderTransportSnapshot, request_path: &str, @@ -193,7 +218,13 @@ pub fn build_video_create_upstream_url( ProviderVideoCreateFamily::Gemini => &["key"][..], }; return build_passthrough_path_url( - &transport.endpoint.base_url, + &crate::xai::resolved_xai_request_base_url( + transport, + match family { + ProviderVideoCreateFamily::OpenAi => "openai:video", + ProviderVideoCreateFamily::Gemini => "gemini:video", + }, + ), path, request_query, blocked_keys, @@ -202,8 +233,14 @@ pub fn build_video_create_upstream_url( match family { ProviderVideoCreateFamily::OpenAi => build_passthrough_path_url( - &transport.endpoint.base_url, - openai_video_api_root_request_path(request_path), + &crate::xai::resolved_xai_request_base_url(transport, "openai:video"), + if crate::xai::is_xai_provider_transport(transport) + && matches!(request_path, "/v1/videos" | "/openai/v1/videos") + { + "/videos/generations" + } else { + openai_video_api_root_request_path(request_path) + }, request_query, &[], ), @@ -216,6 +253,7 @@ pub fn build_video_create_upstream_url( } fn openai_video_api_root_request_path(request_path: &str) -> &str { + let request_path = request_path.strip_prefix("/openai").unwrap_or(request_path); if request_path.starts_with("/v1/") { &request_path[3..] } else { @@ -232,6 +270,11 @@ pub fn build_video_create_headers( input.auth_value, &BTreeMap::new(), ); + crate::xai::insert_cli_identity_headers_if_needed( + input.transport, + "openai:video", + &mut provider_request_headers, + ); if !apply_local_header_rules_with_request_headers( &mut provider_request_headers, input.header_rules, @@ -281,16 +324,22 @@ pub async fn reconstruct_local_video_task_snapshot( return Ok(None); }; - let Some(local_transport) = + let Some(mut local_transport) = resolve_local_video_task_transport(&transport, provider_api_format, task.model.clone()) else { return Ok(None); }; - Ok(LocalVideoTaskSnapshot::from_stored_task_with_transport( - task, - local_transport, - )) + // Resolve deployment-managed nodes, system defaults and tunnel affinity just as + // creation does; serialized task metadata intentionally contains no credentials. + local_transport.proxy = lookup.resolve_video_task_proxy(&transport).await; + + let mut snapshot = + LocalVideoTaskSnapshot::from_stored_task_with_transport(task, local_transport); + if let Some(LocalVideoTaskSnapshot::OpenAi(seed)) = &mut snapshot { + seed.xai_provider = crate::xai::is_xai_provider_transport(&transport); + } + Ok(snapshot) } #[cfg(test)] @@ -441,6 +490,46 @@ mod tests { assert_eq!(transport.provider_id, "provider-1"); } + #[tokio::test] + async fn reconstructs_video_with_configured_proxy_and_profile() { + let mut transport = sample_transport("openai:video", "oauth"); + transport.provider.provider_type = "xai".into(); + transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".into(); + transport.provider.proxy = Some(json!({"enabled":true,"url":"http://127.0.0.1:9876"})); + transport.provider.config = Some(json!({"fingerprint":{"transport_profile":{ + "profile_id":"test-video","backend":"reqwest_rustls","http_mode":"auto","pool_scope":"key" + }}})); + transport.key.decrypted_auth_config = Some(r#"{"using_api":false}"#.into()); + let lookup = TestLookup(Some(transport)); + let snapshot = reconstruct_local_video_task_snapshot(&lookup, &sample_stored_video_task()) + .await + .unwrap() + .expect("proxied video must resume after restart"); + let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else { + panic!("expected OpenAI video") + }; + assert!(seed.xai_provider); + assert_eq!( + seed.transport.proxy.as_ref().unwrap().url.as_deref(), + Some("http://127.0.0.1:9876/") + ); + assert_eq!( + seed.transport + .transport_profile + .as_ref() + .unwrap() + .profile_id, + "test-video" + ); + assert_eq!( + seed.transport + .headers + .get("x-xai-token-auth") + .map(String::as_str), + Some("xai-grok-cli") + ); + } + #[test] fn resolves_gemini_video_transport() { let transport = resolve_local_video_task_transport( @@ -489,6 +578,120 @@ mod tests { assert_eq!(url, "https://api.openai.example/v1/videos?trace=1"); } + #[test] + fn xai_video_create_paths_preserve_auth_hosts_and_custom_endpoints() { + for (auth, base) in [ + ("oauth", "https://cli-chat-proxy.grok.com/v1"), + ("api_key", "https://api.x.ai/v1"), + ] { + let mut transport = sample_transport("openai:video", auth); + transport.provider.provider_type = "xai".into(); + transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".into(); + transport.key.decrypted_auth_config = + (auth == "oauth").then(|| r#"{"using_api":false}"#.into()); + for path in ["/v1/videos", "/openai/v1/videos", "/v1/videos/generations"] { + assert_eq!( + build_video_create_upstream_url( + &transport, + path, + Some("trace=1"), + "grok-imagine-video", + ProviderVideoCreateFamily::OpenAi + ) + .unwrap(), + format!("{base}/videos/generations?trace=1") + ); + } + transport.endpoint.base_url = "https://gateway.example/prefix/v1".into(); + assert_eq!( + build_video_create_upstream_url( + &transport, + "/openai/v1/videos", + None, + "grok-imagine-video", + ProviderVideoCreateFamily::OpenAi + ) + .unwrap(), + "https://gateway.example/prefix/v1/videos/generations" + ); + transport.endpoint.custom_path = Some("/custom/videos/generations".into()); + let url = build_video_create_upstream_url( + &transport, + "/openai/v1/videos", + None, + "grok-imagine-video", + ProviderVideoCreateFamily::OpenAi, + ) + .unwrap(); + assert!(url.ends_with("/custom/videos/generations"), "{url}"); + } + let transport = sample_transport("openai:video", "api_key"); + assert_eq!( + build_video_create_upstream_url( + &transport, + "/openai/v1/videos", + None, + "sora", + ProviderVideoCreateFamily::OpenAi + ), + build_video_create_upstream_url( + &transport, + "/v1/videos", + None, + "sora", + ProviderVideoCreateFamily::OpenAi + ) + ); + } + + #[test] + fn xai_oauth_video_uses_cli_proxy() { + let mut transport = sample_transport("openai:video", "oauth"); + transport.provider.provider_type = "xai".to_string(); + transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".to_string(); + transport.key.decrypted_auth_config = + Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string()); + let url = build_video_create_upstream_url( + &transport, + "/v1/videos/generations", + None, + "grok-imagine-video", + ProviderVideoCreateFamily::OpenAi, + ) + .expect("url should build"); + + assert_eq!(url, "https://cli-chat-proxy.grok.com/v1/videos/generations"); + let headers = build_video_create_headers(ProviderVideoCreateHeadersInput { + transport: &transport, + headers: &http::HeaderMap::new(), + auth_header: "authorization", + auth_value: "Bearer test-token", + header_rules: None, + provider_request_body: &json!({"prompt": "A cat"}), + original_request_body: &json!({"prompt": "A cat"}), + }) + .unwrap(); + assert_eq!( + headers.get("x-xai-token-auth").map(String::as_str), + Some("xai-grok-cli") + ); + + let reconstructed = super::resolve_local_video_task_transport( + &transport, + "openai:video", + Some("grok-imagine-video".into()), + ) + .unwrap(); + assert_eq!( + reconstructed.upstream_base_url, + "https://cli-chat-proxy.grok.com/v1" + ); + assert_eq!( + reconstructed.headers.get("x-xai-token-auth"), + headers.get("x-xai-token-auth") + ); + } + #[test] fn builds_gemini_video_create_url_and_removes_client_key_query() { let transport = sample_transport("gemini:video", "api_key"); @@ -512,6 +715,7 @@ mod tests { let provider_request_body = json!({"prompt": "make a clip"}); let original_request_body = provider_request_body.clone(); let headers = build_video_create_headers(ProviderVideoCreateHeadersInput { + transport: &sample_transport("openai:video", "bearer"), headers: &http::HeaderMap::new(), auth_header: "authorization", auth_value: "Bearer secret", diff --git a/crates/aether-provider/transport/src/xai.rs b/crates/aether-provider/transport/src/xai.rs new file mode 100644 index 000000000..5dbeae020 --- /dev/null +++ b/crates/aether-provider/transport/src/xai.rs @@ -0,0 +1,444 @@ +pub mod video; + +use std::collections::BTreeMap; + +use aether_ai_formats::normalize_api_format_alias; +use serde_json::Value; + +use crate::snapshot::GatewayProviderTransportSnapshot; + +pub const XAI_PROVIDER_TYPE: &str = "xai"; +pub const XAI_CHAT_PROXY_BASE_URL: &str = "https://cli-chat-proxy.grok.com/v1"; +pub const XAI_API_BASE_URL: &str = "https://api.x.ai/v1"; +pub const XAI_CLIENT_VERSION: &str = "0.2.120"; +pub const XAI_TOKEN_AUTH_HEADER: &str = "x-xai-token-auth"; +pub const XAI_TOKEN_AUTH_VALUE: &str = "xai-grok-cli"; +pub const XAI_CLIENT_VERSION_HEADER: &str = "x-grok-client-version"; +pub const XAI_CLIENT_IDENTIFIER_HEADER: &str = "x-grok-client-identifier"; +pub const XAI_CLIENT_IDENTIFIER_VALUE: &str = "grok-shell"; +pub const XAI_AUTHENTICATE_RESPONSE_HEADER: &str = "x-authenticateresponse"; +pub const XAI_AUTHENTICATE_RESPONSE_VALUE: &str = "authenticate-response"; + +pub fn xai_cli_user_agent() -> String { + format!("xai-grok-workspace/{XAI_CLIENT_VERSION}") +} + +pub fn is_xai_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool { + transport + .provider + .provider_type + .trim() + .eq_ignore_ascii_case(XAI_PROVIDER_TYPE) +} + +pub fn xai_uses_official_api(api_format: &str) -> bool { + matches!( + normalize_api_format_alias(api_format).as_str(), + "openai:responses:compact" + ) +} + +pub fn resolved_xai_upstream_base_url( + transport: &GatewayProviderTransportSnapshot, + api_format: &str, +) -> Option { + if !is_xai_provider_transport(transport) { + return None; + } + let stored = transport.endpoint.base_url.trim(); + if xai_uses_official_api(api_format) { + if stored.is_empty() + || is_cli_chat_proxy_base_url(stored) + || is_official_api_base_url(stored) + { + return Some(XAI_API_BASE_URL.to_string()); + } + return Some(trim_base_url(stored)); + } + if xai_using_api(transport) { + if stored.is_empty() || is_cli_chat_proxy_base_url(stored) { + return Some(XAI_API_BASE_URL.to_string()); + } + return Some(trim_base_url(stored)); + } + if stored.is_empty() || is_official_api_base_url(stored) { + return Some(XAI_CHAT_PROXY_BASE_URL.to_string()); + } + Some(trim_base_url(stored)) +} + +pub fn resolved_xai_request_base_url( + transport: &GatewayProviderTransportSnapshot, + api_format: &str, +) -> String { + resolved_xai_upstream_base_url(transport, api_format) + .unwrap_or_else(|| trim_base_url(&transport.endpoint.base_url)) +} + +pub fn should_attach_cli_identity_headers( + transport: &GatewayProviderTransportSnapshot, + api_format: &str, +) -> bool { + if !is_xai_provider_transport(transport) { + return false; + } + if xai_uses_official_api(api_format) { + return false; + } + resolved_xai_upstream_base_url(transport, api_format) + .as_deref() + .is_some_and(is_cli_chat_proxy_base_url) +} + +pub fn insert_cli_identity_headers(headers: &mut BTreeMap) { + let user_agent = xai_cli_user_agent(); + for (name, value) in [ + (XAI_TOKEN_AUTH_HEADER, XAI_TOKEN_AUTH_VALUE), + (XAI_CLIENT_VERSION_HEADER, XAI_CLIENT_VERSION), + ("user-agent", user_agent.as_str()), + (XAI_CLIENT_IDENTIFIER_HEADER, XAI_CLIENT_IDENTIFIER_VALUE), + ( + XAI_AUTHENTICATE_RESPONSE_HEADER, + XAI_AUTHENTICATE_RESPONSE_VALUE, + ), + ] { + if !headers + .keys() + .any(|existing| existing.eq_ignore_ascii_case(name)) + { + headers.insert(name.to_string(), value.to_string()); + } + } +} + +pub fn insert_cli_identity_headers_if_needed( + transport: &GatewayProviderTransportSnapshot, + api_format: &str, + headers: &mut BTreeMap, +) { + if should_attach_cli_identity_headers(transport, api_format) { + insert_cli_identity_headers(headers); + } +} + +pub fn xai_auth_uses_api(auth_type: &str, decrypted_auth_config: Option<&str>) -> bool { + if let Some(value) = auth_config_using_api(decrypted_auth_config) { + return value; + } + let auth_type = auth_type.trim().to_ascii_lowercase(); + if auth_type == "oauth" || auth_config_has_refresh_token(decrypted_auth_config) { + return false; + } + matches!(auth_type.as_str(), "api_key" | "bearer" | "apikey") +} + +pub fn extract_xai_user_id_from_auth_config(raw_auth_config: Option<&str>) -> Option { + let value = parse_auth_config(raw_auth_config)?; + extract_xai_user_id_from_value(&value) +} + +pub fn extract_xai_user_id_from_value(value: &Value) -> Option { + const PATHS: &[&[&str]] = &[ + &["userId"], + &["user_id"], + &["id"], + &["sub"], + &["user", "userId"], + &["user", "id"], + &["user", "user_id"], + &["user", "sub"], + ]; + PATHS.iter().find_map(|path| { + let mut current = value; + for key in *path { + current = current.get(*key)?; + } + coerce_xai_id(current) + }) +} + +fn xai_using_api(transport: &GatewayProviderTransportSnapshot) -> bool { + xai_auth_uses_api( + transport.key.auth_type.as_str(), + transport.key.decrypted_auth_config.as_deref(), + ) +} + +fn coerce_xai_id(value: &Value) -> Option { + match value { + Value::String(text) => { + let trimmed = text.trim(); + (!trimmed.is_empty()).then(|| trimmed.to_string()) + } + Value::Number(number) => { + let rendered = number.to_string(); + (!rendered.is_empty()).then_some(rendered) + } + _ => None, + } +} + +fn auth_config_using_api(raw_auth_config: Option<&str>) -> Option { + let value = parse_auth_config(raw_auth_config)?; + let using_api = value.get("using_api")?; + match using_api { + Value::Bool(value) => Some(*value), + Value::String(value) => value.trim().parse::().ok(), + _ => None, + } +} + +fn auth_config_has_refresh_token(raw_auth_config: Option<&str>) -> bool { + let value = match parse_auth_config(raw_auth_config) { + Some(value) => value, + None => return false, + }; + ["refresh_token", "refreshToken"] + .iter() + .find_map(|field| value.get(*field).and_then(Value::as_str)) + .map(str::trim) + .is_some_and(|value| !value.is_empty()) +} + +fn parse_auth_config(raw_auth_config: Option<&str>) -> Option { + raw_auth_config + .map(str::trim) + .filter(|value| !value.is_empty()) + .and_then(|value| serde_json::from_str::(value).ok()) +} + +fn trim_base_url(url: &str) -> String { + url.trim().trim_end_matches('/').to_string() +} + +fn normalize_base_url(url: &str) -> String { + trim_base_url(url).to_ascii_lowercase() +} + +fn is_official_api_base_url(url: &str) -> bool { + normalize_base_url(url) == normalize_base_url(XAI_API_BASE_URL) +} + +fn is_cli_chat_proxy_base_url(url: &str) -> bool { + normalize_base_url(url) == normalize_base_url(XAI_CHAT_PROXY_BASE_URL) +} + +#[cfg(test)] +mod tests { + use super::{ + insert_cli_identity_headers_if_needed, is_xai_provider_transport, + resolved_xai_upstream_base_url, should_attach_cli_identity_headers, XAI_API_BASE_URL, + XAI_CHAT_PROXY_BASE_URL, XAI_CLIENT_IDENTIFIER_VALUE, XAI_TOKEN_AUTH_VALUE, + }; + use crate::snapshot::{ + GatewayProviderTransportEndpoint, GatewayProviderTransportKey, + GatewayProviderTransportProvider, GatewayProviderTransportSnapshot, + }; + use std::collections::BTreeMap; + + fn sample_transport( + auth_type: &str, + auth_config: Option<&str>, + base_url: &str, + ) -> GatewayProviderTransportSnapshot { + GatewayProviderTransportSnapshot { + provider: GatewayProviderTransportProvider { + id: "provider-xai".to_string(), + name: "xAI".to_string(), + provider_type: "xai".to_string(), + website: None, + is_active: true, + keep_priority_on_conversion: false, + enable_format_conversion: true, + concurrent_limit: None, + max_retries: None, + proxy: None, + request_timeout_secs: None, + stream_first_byte_timeout_secs: None, + config: None, + }, + endpoint: GatewayProviderTransportEndpoint { + id: "endpoint-xai".to_string(), + provider_id: "provider-xai".to_string(), + api_format: "openai:responses".to_string(), + api_family: None, + endpoint_kind: None, + is_active: true, + base_url: base_url.to_string(), + header_rules: None, + body_rules: None, + max_retries: None, + custom_path: None, + config: None, + format_acceptance_config: None, + proxy: None, + }, + key: GatewayProviderTransportKey { + id: "key-xai".to_string(), + provider_id: "provider-xai".to_string(), + name: "key".to_string(), + auth_type: auth_type.to_string(), + is_active: true, + api_formats: None, + auth_type_by_format: None, + allow_auth_channel_mismatch_formats: None, + allowed_models: None, + capabilities: None, + rate_multipliers: None, + global_priority_by_format: None, + expires_at_unix_secs: None, + proxy: None, + fingerprint: None, + upstream_metadata: None, + decrypted_api_key: "access-token".to_string(), + decrypted_auth_config: auth_config.map(ToOwned::to_owned), + }, + } + } + + #[test] + fn oauth_defaults_to_cli_chat_proxy_for_responses() { + let transport = sample_transport( + "oauth", + Some(r#"{"refresh_token":"rt","using_api":false}"#), + XAI_CHAT_PROXY_BASE_URL, + ); + assert!(is_xai_provider_transport(&transport)); + assert_eq!( + resolved_xai_upstream_base_url(&transport, "openai:responses").as_deref(), + Some(XAI_CHAT_PROXY_BASE_URL) + ); + assert!(should_attach_cli_identity_headers( + &transport, + "openai:responses" + )); + } + + #[test] + fn compact_and_using_api_stay_on_official_api() { + let oauth = sample_transport( + "oauth", + Some(r#"{"refresh_token":"rt","using_api":false}"#), + XAI_CHAT_PROXY_BASE_URL, + ); + assert_eq!( + resolved_xai_upstream_base_url(&oauth, "openai:responses:compact").as_deref(), + Some(XAI_API_BASE_URL) + ); + assert!(!should_attach_cli_identity_headers( + &oauth, + "openai:responses:compact" + )); + + let api_key = sample_transport( + "oauth", + Some(r#"{"using_api":true}"#), + XAI_CHAT_PROXY_BASE_URL, + ); + assert_eq!( + resolved_xai_upstream_base_url(&api_key, "openai:responses").as_deref(), + Some(XAI_API_BASE_URL) + ); + assert!(!should_attach_cli_identity_headers( + &api_key, + "openai:responses" + )); + } + + #[test] + fn media_routing_and_cli_headers_follow_auth_and_base_url() { + for api_format in ["openai:image", "openai:video"] { + for stored in ["", XAI_API_BASE_URL, XAI_CHAT_PROXY_BASE_URL] { + for (auth_type, config, expected) in [ + ( + "oauth", + Some(r#"{"refresh_token":"rt","using_api":false}"#), + XAI_CHAT_PROXY_BASE_URL, + ), + ("oauth", Some(r#"{"using_api":true}"#), XAI_API_BASE_URL), + ("bearer", None, XAI_API_BASE_URL), + ] { + let transport = sample_transport(auth_type, config, stored); + assert_eq!( + resolved_xai_upstream_base_url(&transport, api_format).as_deref(), + Some(expected) + ); + assert_eq!( + should_attach_cli_identity_headers(&transport, api_format), + expected == XAI_CHAT_PROXY_BASE_URL + ); + } + } + let custom = sample_transport("oauth", None, "https://custom.example/v1"); + assert_eq!( + resolved_xai_upstream_base_url(&custom, api_format).as_deref(), + Some("https://custom.example/v1") + ); + assert!(!should_attach_cli_identity_headers(&custom, api_format)); + } + } + + #[test] + fn bearer_without_refresh_uses_official_api() { + let transport = sample_transport("bearer", None, XAI_CHAT_PROXY_BASE_URL); + assert_eq!( + resolved_xai_upstream_base_url(&transport, "openai:responses").as_deref(), + Some(XAI_API_BASE_URL) + ); + } + + #[test] + fn cli_headers_do_not_override_existing_values() { + let transport = sample_transport( + "oauth", + Some(r#"{"refresh_token":"rt"}"#), + XAI_CHAT_PROXY_BASE_URL, + ); + let mut headers = BTreeMap::from([( + "x-grok-client-identifier".to_string(), + "custom-client".to_string(), + )]); + insert_cli_identity_headers_if_needed(&transport, "openai:responses", &mut headers); + assert_eq!( + headers.get("x-grok-client-identifier").map(String::as_str), + Some("custom-client") + ); + assert_eq!( + headers.get("x-xai-token-auth").map(String::as_str), + Some(XAI_TOKEN_AUTH_VALUE) + ); + assert_eq!( + headers.get("x-authenticateresponse").map(String::as_str), + Some("authenticate-response") + ); + assert_ne!( + headers.get("x-grok-client-identifier").map(String::as_str), + Some(XAI_CLIENT_IDENTIFIER_VALUE) + ); + } + + #[test] + fn extracts_user_id_from_user_payload_and_auth_config_sub() { + use super::{ + extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, xai_auth_uses_api, + }; + use serde_json::json; + + assert_eq!( + extract_xai_user_id_from_value(&json!({"userId": "user-42"})).as_deref(), + Some("user-42") + ); + assert_eq!( + extract_xai_user_id_from_auth_config(Some(r#"{"sub":"subject-1"}"#)).as_deref(), + Some("subject-1") + ); + assert!(!xai_auth_uses_api( + "oauth", + Some(r#"{"refresh_token":"rt","using_api":false}"#) + )); + assert!(xai_auth_uses_api( + "bearer", + Some(r#"{"api_key":"xai-key","using_api":true}"#) + )); + } +} diff --git a/crates/aether-provider/transport/src/xai/video.rs b/crates/aether-provider/transport/src/xai/video.rs new file mode 100644 index 000000000..b715952d2 --- /dev/null +++ b/crates/aether-provider/transport/src/xai/video.rs @@ -0,0 +1,147 @@ +use serde_json::{json, Value}; + +/// Native xAI video requests live under /v1; the OpenAI-compatible adapter under /openai/v1. +pub fn is_native_video_request(provider_type: &str, path: &str) -> bool { + provider_type.trim().eq_ignore_ascii_case("xai") + && matches!( + path, + "/v1/videos" | "/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions" + ) +} + +pub fn is_explicit_native_video_path(path: &str) -> bool { + matches!( + path, + "/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions" + ) +} + +/// Convert the OpenAI video request contract to xAI's native contract. +/// Native requests bypass this adapter so provider-specific fields remain intact. +pub fn convert_openai_video_request(body: &Value) -> Result { + let prompt = text(&body["prompt"]).ok_or("prompt is required")?; + let seconds = match &body["seconds"] { + Value::Null => 4, + Value::String(value) if value.trim().is_empty() => 4, + Value::String(value) => value + .trim() + .parse::() + .map_err(|_| "seconds must be an integer")?, + value => value.as_i64().ok_or("seconds must be an integer")?, + } + .clamp(1, 15); + let size = text(&body["size"]).unwrap_or("720x1280"); + let default_ratio = match size { + "720x1280" | "1024x1792" => "9:16", + "1280x720" | "1792x1024" => "16:9", + _ => return Err("size must be one of 720x1280, 1280x720, 1024x1792, or 1792x1024"), + }; + let ratio = match text(&body["aspect_ratio"]) + .unwrap_or("") + .to_ascii_lowercase() + .as_str() + { + "square" | "1:1" => "1:1", + "landscape" | "16:9" => "16:9", + "portrait" | "9:16" => "9:16", + "4:3" => "4:3", + "3:4" => "3:4", + "3:2" => "3:2", + "2:3" => "2:3", + _ => default_ratio, + }; + let resolution = if text(&body["resolution"]).is_some_and(|v| v.eq_ignore_ascii_case("480p")) { + "480p" + } else { + "720p" + }; + if text(&body["input_reference"]["file_id"]).is_some() { + return Err("input_reference.file_id is not supported for xAI video generation; use input_reference.image_url"); + } + let image = text(&body["input_reference"]["image_url"]) + .or_else(|| image_url(&body["image"])) + .or_else(|| text(&body["image_url"])); + let references: Vec<_> = ["reference_images", "reference_image_urls"] + .into_iter() + .filter_map(|key| body[key].as_array()) + .flatten() + .filter_map(image_url) + .map(|url| json!({"url":url})) + .collect(); + if references.len() > 7 { + return Err("reference_images supports at most 7 images on xAI"); + } + if image.is_some() && !references.is_empty() { + return Err("image and reference_images cannot be combined on xAI"); + } + let mut result = json!({"model":body["model"], "prompt":prompt, "duration":seconds, "aspect_ratio":ratio, "resolution":resolution}); + if let Some(url) = image { + result["image"] = json!({"url":url}); + } + if !references.is_empty() { + result["reference_images"] = json!(references); + } + Ok(result) +} + +fn text(value: &Value) -> Option<&str> { + value + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +fn image_url(value: &Value) -> Option<&str> { + text(value) + .or_else(|| text(&value["url"])) + .or_else(|| text(&value["image_url"])) + .or_else(|| text(&value["image_url"]["url"])) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn xai_video_compatibility_maps_duration_size_and_references() { + let converted = convert_openai_video_request(&json!({ + "model":"grok-imagine-video", "prompt":"A cat", "seconds":"8", "size":"1280x720", + "reference_images":[{"image_url":{"url":"https://example.com/a.png"}}], + "reference_image_urls":["https://example.com/b.png"] + })) + .unwrap(); + assert_eq!( + converted, + json!({"model":"grok-imagine-video", "prompt":"A cat", "duration":8, + "aspect_ratio":"16:9", "resolution":"720p", "reference_images":[{"url":"https://example.com/a.png"},{"url":"https://example.com/b.png"}]}) + ); + let defaults = convert_openai_video_request(&json!({"prompt":"A cat"})).unwrap(); + assert_eq!(defaults["duration"], 4); + assert_eq!(defaults["aspect_ratio"], "9:16"); + for (seconds, expected) in [(-1, 1), (30, 15)] { + assert_eq!( + convert_openai_video_request(&json!({"prompt":"A cat", "seconds":seconds})) + .unwrap()["duration"], + expected + ); + } + } + + #[test] + fn xai_video_compatibility_validates_requests_and_maps_image_input() { + for invalid in [ + json!({}), + json!({"prompt":"cat","seconds":"1.5"}), + json!({"prompt":"cat","size":"foo"}), + json!({"prompt":"cat","input_reference":{"file_id":"file-1"}}), + json!({"prompt":"cat","image":"https://example.com/a.png","reference_images":["https://example.com/b.png"]}), + json!({"prompt":"cat","reference_images":vec!["https://example.com/a.png";8]}), + ] { + assert!(convert_openai_video_request(&invalid).is_err(), "{invalid}"); + } + let body = convert_openai_video_request(&json!({"prompt":"cat","input_reference":{"image_url":"https://example.com/a.png"},"aspect_ratio":"square","resolution":"480p"})).unwrap(); + assert_eq!(body["image"]["url"], "https://example.com/a.png"); + assert_eq!(body["aspect_ratio"], "1:1"); + assert_eq!(body["resolution"], "480p"); + } +} diff --git a/crates/aether-testing/integration/Cargo.toml b/crates/aether-testing/integration/Cargo.toml index 1bc86bb8b..b0d164e8a 100644 --- a/crates/aether-testing/integration/Cargo.toml +++ b/crates/aether-testing/integration/Cargo.toml @@ -15,11 +15,14 @@ aether-data-contracts.workspace = true aether-gateway = { workspace = true, features = ["testkit"] } aether-runtime.workspace = true aether-runtime-state.workspace = true +aether-tunnel.workspace = true aether-testkit = { workspace = true, features = ["gateway", "postgres"] } +arc-swap = "1" axum.workspace = true futures-util.workspace = true http.workspace = true reqwest.workspace = true +rustls.workspace = true serde.workspace = true serde_json.workspace = true sha2.workspace = true diff --git a/crates/aether-testing/integration/tests/tunnel_runtime_e2e.rs b/crates/aether-testing/integration/tests/tunnel_runtime_e2e.rs new file mode 100644 index 000000000..baca13281 --- /dev/null +++ b/crates/aether-testing/integration/tests/tunnel_runtime_e2e.rs @@ -0,0 +1,527 @@ +//! Gateway-backed tunnel end-to-end regressions. +//! +//! 这两个用例需要真实 Gateway 路由和 Tunnel 进程状态,因此放在独立 +//! integration package,避免 Workspace Rest 的普通目标编译 Gateway。 + +use std::future::Future; +use std::pin::Pin; +use std::sync::atomic::AtomicU64; +use std::sync::{Arc, Once}; +use std::task::{Context, Poll}; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use aether_contracts::tunnel::{ + sign_tunnel_relay_request, tunnel_relay_payload_digest, TUNNEL_RELAY_AUTH_NONCE_HEADER, + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, TUNNEL_RELAY_AUTH_SENDER_HEADER, + TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, +}; +use aether_gateway::{build_router_with_state, AppState as GatewayAppState}; +use aether_tunnel::config::Config; +use aether_tunnel::registration::client::AetherClient; +use aether_tunnel::runtime::DynamicConfig; +use aether_tunnel::state::{ + AppState as TunnelAppState, ServerContext, TunnelMetrics, TunnelRequestMetrics, +}; +use aether_tunnel::target_filter::DnsCache; +use aether_tunnel::tunnel::protocol; +use aether_tunnel::tunnel::run; +use aether_tunnel::upstream_client; +use arc_swap::ArcSwap; +use axum::Router; +use reqwest::StatusCode; +use tokio::sync::watch; + +struct SessionTask(tokio::task::JoinHandle); + +impl SessionTask { + fn new(handle: tokio::task::JoinHandle) -> Self { + Self(handle) + } +} + +impl Future for SessionTask { + type Output = Result; + + fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll { + Pin::new(&mut self.0).poll(context) + } +} + +impl Drop for SessionTask { + fn drop(&mut self) { + self.0.abort(); + } +} + +#[tokio::test] +async fn tunnel_reconnects_after_gateway_restart() { + ensure_rustls_provider(); + + let gateway_port = reserve_local_port().expect("gateway port should reserve"); + let gateway_base_url = format!("http://127.0.0.1:{gateway_port}"); + let (gateway_state, mut gateway_handle) = start_gateway_on_port(gateway_port) + .await + .expect("gateway should start"); + + let mut tunnel_config = sample_config(&gateway_base_url); + tunnel_config.tunnel_security = aether_tunnel::config::TunnelSecurity::NonTlsRequired; + tunnel_config.tunnel_encryption_key = + Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".to_string()); + let state = sample_state(tunnel_config); + let server = sample_server(&state, "node-recovery"); + let (shutdown_tx, shutdown_rx) = watch::channel(false); + let tunnel_task = tokio::spawn({ + let state = Arc::clone(&state); + let server = Arc::clone(&server); + let (_drain_tx, drain_rx) = watch::channel(false); + async move { + run(&state, &server, 0, shutdown_rx, drain_rx).await; + } + }); + + wait_until_relay_status( + &gateway_base_url, + "node-recovery", + StatusCode::GATEWAY_TIMEOUT, + ) + .await; + + gateway_handle.abort(); + let _ = (&mut gateway_handle).await; + assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1); + + let (_restarted_gateway_state, restarted_gateway_handle) = + start_gateway_on_port_retry(gateway_port) + .await + .expect("gateway should restart on fixed port"); + gateway_handle = restarted_gateway_handle; + + wait_until_relay_status( + &gateway_base_url, + "node-recovery", + StatusCode::GATEWAY_TIMEOUT, + ) + .await; + + assert!(server.tunnel_metrics.snapshot().connect_successes >= 2); + let _ = shutdown_tx.send(true); + tokio::time::timeout(Duration::from_secs(5), tunnel_task) + .await + .expect("tunnel task should stop") + .expect("tunnel task should join"); + gateway_handle.abort(); +} + +async fn wait_until_relay_status(gateway_base_url: &str, node_id: &str, expected: StatusCode) { + let deadline = tokio::time::Instant::now() + Duration::from_secs(10); + let mut last_observed = None::; + loop { + if let Some((status, body)) = probe_relay_status(gateway_base_url, node_id).await { + last_observed = Some(format!("{status} body={body}")); + if status == expected { + return; + } + } + assert!( + tokio::time::Instant::now() < deadline, + "relay status did not become {expected} within timeout; last={:?}", + last_observed + ); + tokio::time::sleep(Duration::from_millis(25)).await; + } +} + +async fn probe_relay_status(gateway_base_url: &str, node_id: &str) -> Option<(StatusCode, String)> { + let response = relay_response(gateway_base_url, node_id, relay_probe_envelope()).await?; + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + Some((status, body)) +} + +async fn relay_response( + gateway_base_url: &str, + node_id: &str, + payload: Vec, +) -> Option { + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("test clock should be after epoch") + .as_secs(); + let nonce = uuid::Uuid::new_v4().simple().to_string(); + let digest = tunnel_relay_payload_digest(&payload, &[]); + let signature = sign_tunnel_relay_request( + b"tunnel-reconnect-test-secret-at-least-32-bytes", + "tunnel-reconnect-test-client", + "tunnel-reconnect-test-gateway", + node_id, + "", + false, + timestamp, + &nonce, + &digest, + ); + reqwest::Client::new() + .post(format!( + "{gateway_base_url}/api/internal/tunnel/relay/{node_id}" + )) + .header("content-type", "application/octet-stream") + .header( + TUNNEL_RELAY_AUTH_SENDER_HEADER, + "tunnel-reconnect-test-client", + ) + .header( + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + "tunnel-reconnect-test-gateway", + ) + .header(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, timestamp) + .header(TUNNEL_RELAY_AUTH_NONCE_HEADER, nonce) + .header( + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + digest.encode_header_value(), + ) + .header(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, signature) + .body(payload) + .send() + .await + .ok() +} + +fn relay_probe_envelope() -> Vec { + let meta = protocol::RequestMeta { + provider_id: None, + endpoint_id: None, + key_id: None, + method: "GET".to_string(), + url: "http://127.0.0.1:80/blocked".to_string(), + headers: std::collections::HashMap::new(), + stream: false, + request_timeout_ms: None, + stream_first_byte_timeout_ms: None, + timeout: 5, + follow_redirects: None, + http1_only: false, + transport_profile: None, + }; + let meta_json = + serde_json::to_vec(&meta).expect("tunnel relay probe metadata should serialize"); + let mut envelope = Vec::with_capacity(4 + meta_json.len()); + envelope.extend_from_slice(&(meta_json.len() as u32).to_be_bytes()); + envelope.extend_from_slice(&meta_json); + envelope +} + +async fn start_gateway_on_port( + port: u16, +) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> { + // The embedded gateway now fails closed when relay authentication is + // not configured. Keep this integration fixture explicitly authenticated. + static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(()); + let state = { + let _guard = ENV_LOCK.lock().unwrap(); + let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET"); + let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID"); + std::env::set_var( + "AETHER_TUNNEL_RELAY_AUTH_SECRET", + "tunnel-reconnect-test-secret-at-least-32-bytes", + ); + std::env::set_var( + "AETHER_GATEWAY_INSTANCE_ID", + "tunnel-reconnect-test-gateway", + ); + let mut state = GatewayAppState::new().expect("gateway test state should build"); + aether_gateway::configure_test_tunnel_security( + &mut state, + "node-recovery", + "test-generation-1", + "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=", + ); + restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret); + restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance); + state + }; + let router = build_router_with_state(state.clone()); + let handle = spawn_router_on_port(port, router).await?; + Ok((state, handle)) +} + +#[tokio::test] +async fn negotiated_small_window_streams_large_responses_and_cancels_idle_upstream() { + use axum::body::{Body, Bytes}; + use axum::routing::get; + use futures_util::StreamExt; + + ensure_rustls_provider(); + let upstream_port = reserve_local_port().unwrap(); + let upstream = Router::new() + .route( + "/large", + get(|| async { Body::from(vec![b'x'; 2 * 1024 * 1024]) }), + ) + .route( + "/idle", + get(|| async { + let first = futures_util::stream::once(async { + Ok::<_, std::io::Error>(Bytes::from_static(b"data: started\n\n")) + }); + ( + [("content-type", "text/event-stream")], + Body::from_stream(first.chain(futures_util::stream::pending())), + ) + }), + ); + let upstream_task = + SessionTask::new(spawn_router_on_port(upstream_port, upstream).await.unwrap()); + let gateway_port = reserve_local_port().unwrap(); + let gateway_url = format!("http://127.0.0.1:{gateway_port}"); + let (_, gateway_task) = start_gateway_on_port(gateway_port).await.unwrap(); + let gateway_task = SessionTask::new(gateway_task); + let mut config = sample_config(&gateway_url); + config.tunnel_security = aether_tunnel::config::TunnelSecurity::NonTlsRequired; + config.tunnel_encryption_key = Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".into()); + config.tunnel_stream_initial_window_bytes = 512 * 1024; + config.tunnel_drain_deadline_ms = 100; + config.allow_private_targets = true; + config.allowed_ports.push(upstream_port); + let state = sample_state(config); + let server = sample_server(&state, "node-recovery"); + let (shutdown_tx, shutdown_rx) = watch::channel(false); + let (_drain_tx, drain_rx) = watch::channel(false); + let tunnel_task = SessionTask::new(tokio::spawn({ + let state = Arc::clone(&state); + let server = Arc::clone(&server); + async move { + run(&state, &server, 0, shutdown_rx, drain_rx).await; + } + })); + wait_until_relay_status(&gateway_url, "node-recovery", StatusCode::GATEWAY_TIMEOUT).await; + + let envelope = |path: &str| { + let mut meta: protocol::RequestMeta = + serde_json::from_slice(&relay_probe_envelope()[4..]).unwrap(); + meta.url = format!("http://127.0.0.1:{upstream_port}/{path}"); + meta.stream = true; + meta.timeout = 10; + meta.stream_first_byte_timeout_ms = Some(10_000); + let encoded = serde_json::to_vec(&meta).unwrap(); + let mut result = (encoded.len() as u32).to_be_bytes().to_vec(); + result.extend(encoded); + result + }; + let response = relay_response(&gateway_url, "node-recovery", envelope("large")) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let body = tokio::time::timeout(Duration::from_secs(10), response.bytes()) + .await + .unwrap() + .unwrap(); + assert_eq!(body.len(), 2 * 1024 * 1024); + assert!(body.iter().all(|byte| *byte == b'x')); + + let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle")) + .await + .unwrap(); + assert_eq!( + response.chunk().await.unwrap().unwrap(), + "data: started\n\n" + ); + drop(response); + tokio::time::timeout(Duration::from_secs(3), async { + while server + .active_connections + .load(std::sync::atomic::Ordering::Acquire) + != 0 + { + tokio::task::yield_now().await; + } + }) + .await + .expect("cancelled SSE must release the upstream handler"); + + let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle")) + .await + .unwrap(); + assert!(response.chunk().await.unwrap().is_some()); + shutdown_tx.send(true).unwrap(); + tokio::time::timeout(Duration::from_secs(3), tunnel_task) + .await + .unwrap() + .unwrap(); + assert_eq!( + server + .active_connections + .load(std::sync::atomic::Ordering::Acquire), + 0 + ); + drop(response); + drop(gateway_task); + drop(upstream_task); +} + +fn restore_test_env(key: &str, value: Option) { + if let Some(value) = value { + std::env::set_var(key, value); + } else { + std::env::remove_var(key); + } +} + +async fn start_gateway_on_port_retry( + port: u16, +) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> { + let mut attempts = 0usize; + loop { + match start_gateway_on_port(port).await { + Ok(server) => return Ok(server), + Err(err) => { + attempts += 1; + if attempts >= 20 { + return Err(err); + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + } + } +} + +async fn spawn_router_on_port( + port: u16, + app: Router, +) -> Result, std::io::Error> { + let listener = tokio::net::TcpListener::bind(("127.0.0.1", port)).await?; + Ok(tokio::spawn(async move { + axum::serve( + listener, + app.into_make_service_with_connect_info::(), + ) + .await + .expect("gateway test server should run"); + })) +} + +fn reserve_local_port() -> Result { + let listener = std::net::TcpListener::bind("127.0.0.1:0")?; + let port = listener.local_addr()?.port(); + drop(listener); + Ok(port) +} + +fn sample_state(config: Config) -> Arc { + let config = Arc::new(config); + let dns_cache = Arc::new(DnsCache::new(Duration::from_secs(60), 128)); + let upstream_client_pool = + upstream_client::UpstreamClientPool::new(Arc::clone(&config), Arc::clone(&dns_cache)); + Arc::new(TunnelAppState { + config, + dns_cache, + upstream_client_pool, + tunnel_tls_config: Arc::new(aether_tunnel::tunnel::client::build_tls_config()), + resource_monitor: Arc::new(aether_tunnel::hardware::RuntimeResourceMonitor::new()), + stream_gate: None, + distributed_stream_gate: None, + }) +} + +fn sample_server(state: &Arc, node_id: &str) -> Arc { + let config = Arc::clone(&state.config); + Arc::new(ServerContext { + server_label: "gateway-owned-tunnel".to_string(), + aether_url: config.aether_url.clone(), + management_token: config.management_token.clone(), + tunnel_security: config.tunnel_security, + tunnel_encryption_key: config.tunnel_encryption_key.clone(), + node_name: config.node_name.clone(), + node_id: Arc::new(std::sync::RwLock::new(node_id.to_string())), + tunnel_generation: "test-generation-1".to_string(), + aether_client: Arc::new(AetherClient::new( + &config, + &config.aether_url, + &config.management_token, + )), + dynamic: Arc::new(ArcSwap::from_pointee(DynamicConfig::from_config(&config))), + active_connections: Arc::new(AtomicU64::new(0)), + metrics: Arc::new(TunnelRequestMetrics::new()), + tunnel_metrics: Arc::new(TunnelMetrics::new()), + }) +} + +fn sample_config(aether_url: &str) -> Config { + Config { + aether_url: aether_url.to_string(), + management_token: "token".to_string(), + public_ip: None, + node_name: "tunnel-test".to_string(), + tunnel_security: aether_tunnel::config::TunnelSecurity::Off, + tunnel_encryption_key: None, + node_region: None, + heartbeat_interval: 1, + allowed_ports: vec![80, 443], + allow_private_targets: false, + aether_request_timeout_secs: 10, + aether_connect_timeout_secs: 2, + aether_pool_max_idle_per_host: 8, + aether_pool_idle_timeout_secs: 90, + aether_tcp_keepalive_secs: 60, + aether_tcp_nodelay: true, + aether_http2: true, + aether_outbound_proxy_url: None, + aether_retry_max_attempts: 1, + aether_retry_base_delay_ms: 50, + aether_retry_max_delay_ms: 100, + diagnostics_bind: None, + max_concurrent_connections: None, + max_in_flight_streams: None, + distributed_stream_limit: None, + distributed_stream_redis_url: None, + distributed_stream_redis_key_prefix: None, + distributed_stream_lease_ttl_ms: 30_000, + distributed_stream_renew_interval_ms: 10_000, + distributed_stream_command_timeout_ms: 1_000, + dns_cache_ttl_secs: 60, + dns_cache_capacity: 128, + upstream_connect_timeout_secs: 30, + upstream_pool_max_idle_per_host: 4, + upstream_pool_idle_timeout_secs: 60, + upstream_client_pool_capacity: aether_tunnel::config::DEFAULT_UPSTREAM_CLIENT_POOL_CAPACITY, + upstream_tcp_keepalive_secs: 60, + upstream_tcp_nodelay: true, + upstream_proxy_url: None, + upstream_proxy_remote_dns: false, + legacy_redirect_replay_budget_bytes_ignored: None, + emit_proxy_timing_header: true, + log_level: "info".to_string(), + log_destination: aether_tunnel::config::TunnelLogDestinationArg::Stdout, + log_dir: None, + log_rotation: aether_tunnel::config::TunnelLogRotationArg::Daily, + log_retention_days: 7, + log_max_files: 30, + tunnel_reconnect_base_ms: 50, + tunnel_reconnect_max_ms: 250, + tunnel_ping_interval_ms: 1_000, + tunnel_max_streams: Some(8), + tunnel_profile: aether_tunnel::config::TunnelProfileArg::Lite, + tunnel_stream_initial_window_bytes: + aether_tunnel::config::DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES, + tunnel_drain_deadline_ms: aether_tunnel::config::DEFAULT_TUNNEL_DRAIN_DEADLINE_MS, + tunnel_connect_timeout_ms: 2_000, + tunnel_ipv4_only: false, + tunnel_ipv6_only: false, + tunnel_tcp_keepalive_secs: 30, + tunnel_tcp_nodelay: true, + tunnel_stale_timeout_ms: 5_000, + tunnel_connections: Some(1), + tunnel_connections_max: Some(1), + tunnel_scale_check_interval_ms: 1_000, + tunnel_scale_up_threshold_percent: 70, + tunnel_scale_down_threshold_percent: 35, + tunnel_scale_down_grace_secs: 15, + } +} + +fn ensure_rustls_provider() { + static INIT: Once = Once::new(); + INIT.call_once(|| { + let _ = rustls::crypto::ring::default_provider().install_default(); + }); +} diff --git a/crates/aether-testing/testkit/src/postgres.rs b/crates/aether-testing/testkit/src/postgres.rs index 957301f6f..76876b5b4 100644 --- a/crates/aether-testing/testkit/src/postgres.rs +++ b/crates/aether-testing/testkit/src/postgres.rs @@ -1,5 +1,8 @@ use std::path::PathBuf; use std::process::{Child, Command, Stdio}; +use std::sync::atomic::{AtomicU64, Ordering}; + +static POSTGRES_WORKDIR_SEQ: AtomicU64 = AtomicU64::new(0); use aether_data::driver::postgres::PostgresPoolConfig; use aether_data::{DataBackends, DataLayerConfig}; @@ -21,10 +24,20 @@ pub struct ManagedPostgresServer { impl ManagedPostgresServer { pub async fn start() -> Result> { let port = reserve_local_port()?; + // pid+port is not unique: cargo test shares one PID, and ephemeral ports + // are reused after the listener is dropped. Parallel e2e tests then hit + // create_dir AlreadyExists. + let seq = POSTGRES_WORKDIR_SEQ.fetch_add(1, Ordering::Relaxed); + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or(0); let workdir = std::env::temp_dir().join(format!( - "aether-postgres-baseline-{}-{}", + "aether-postgres-baseline-{}-{}-{}-{}", std::process::id(), - port + port, + seq, + nanos )); let data_dir = workdir.join("data"); std::fs::create_dir(&workdir)?; diff --git a/crates/aether-usage/runtime/src/event_wire.rs b/crates/aether-usage/runtime/src/event_wire.rs index 477603c0a..60427df40 100644 --- a/crates/aether-usage/runtime/src/event_wire.rs +++ b/crates/aether-usage/runtime/src/event_wire.rs @@ -15,9 +15,9 @@ use super::{BorrowedUsageEventEnvelope, UsageEvent, UsageEventData, USAGE_EVENT_ use crate::body_capture::mark_usage_event_capture_truncated; use crate::request_metadata::{ attach_client_request_body_metadata, attach_provider_request_body_metadata, - attach_provider_response_body_metadata, clear_client_request_body_metadata, - clear_provider_request_body_metadata, request_body_derived_facts_action, - RequestBodyDerivedFactsAction, + attach_provider_response_body_metadata, attach_provider_response_model_metadata, + clear_client_request_body_metadata, clear_provider_request_body_metadata, + request_body_derived_facts_action, RequestBodyDerivedFactsAction, }; const DIAGNOSTIC_FIELDS: [&str; 8] = [ @@ -238,6 +238,15 @@ impl WireOverrides { RequestBodyDerivedFactsAction::Clear | RequestBodyDerivedFactsAction::Preserve => {} } metadata = attach_provider_response_body_metadata(metadata, data.response_body.as_ref()); + metadata = attach_provider_response_model_metadata( + metadata, + request_body, + data.request_body_state, + data.api_format.as_deref(), + data.response_body.as_ref(), + data.response_body_state, + data.endpoint_api_format.as_deref(), + ); // Billing reads raw-body TTL before metadata regardless of capture state. // Preserve that precedence independently of reasoning and tier authority. if let Some(cache_ttl) = body_cache_ttl { diff --git a/crates/aether-usage/runtime/src/queue.rs b/crates/aether-usage/runtime/src/queue.rs index 23b6dadb6..04d8fe6fe 100644 --- a/crates/aether-usage/runtime/src/queue.rs +++ b/crates/aether-usage/runtime/src/queue.rs @@ -104,6 +104,13 @@ impl UsageQueue { pub async fn enqueue(&self, event: &UsageEvent) -> Result { let encoded = self.encode_event(event)?; + self.enqueue_encoded(encoded).await + } + + pub(crate) async fn enqueue_encoded( + &self, + encoded: EncodedUsageEvent, + ) -> Result { self.runner .append_fields_with_maxlen( &self.stream, @@ -117,7 +124,10 @@ impl UsageQueue { self.encode_event(event).map(|_| ()) } - fn encode_event(&self, event: &UsageEvent) -> Result { + pub(crate) fn encode_event( + &self, + event: &UsageEvent, + ) -> Result { let encoded = match event.to_bounded_stream_fields(self.config.queue_payload_max_bytes) { Ok(encoded) => encoded, Err(error) => { diff --git a/crates/aether-usage/runtime/src/record.rs b/crates/aether-usage/runtime/src/record.rs index ac46cc482..b7436e7b0 100644 --- a/crates/aether-usage/runtime/src/record.rs +++ b/crates/aether-usage/runtime/src/record.rs @@ -5,9 +5,9 @@ use aether_data_contracts::DataLayerError; use crate::request_metadata::{ attach_client_request_body_metadata, attach_provider_request_body_metadata, - clear_client_request_body_metadata, clear_provider_request_body_metadata, - request_body_derived_facts_action, sanitize_usage_request_metadata, - RequestBodyDerivedFactsAction, + attach_provider_response_model_metadata, clear_client_request_body_metadata, + clear_provider_request_body_metadata, request_body_derived_facts_action, + sanitize_usage_request_metadata, RequestBodyDerivedFactsAction, }; use crate::{UsageEvent, UsageEventType}; @@ -85,6 +85,16 @@ pub fn build_upsert_usage_record_from_event( } RequestBodyDerivedFactsAction::Preserve => {} } + // 响应模型必须在 body 被裁剪前从客户端请求体和上游响应体共同派生;缺少权威 body 时保留已派生事实。 + data.request_metadata = attach_provider_response_model_metadata( + data.request_metadata, + data.request_body.as_ref(), + data.request_body_state, + data.api_format.as_deref(), + data.response_body.as_ref(), + data.response_body_state, + data.endpoint_api_format.as_deref(), + ); let now_unix_secs = event.timestamp_ms / 1_000; Ok(UpsertUsageRecord { diff --git a/crates/aether-usage/runtime/src/request_metadata.rs b/crates/aether-usage/runtime/src/request_metadata.rs index 33f7b5e0f..3d61866e8 100644 --- a/crates/aether-usage/runtime/src/request_metadata.rs +++ b/crates/aether-usage/runtime/src/request_metadata.rs @@ -1,13 +1,15 @@ use aether_contracts::ExecutionPlan; use aether_data_contracts::repository::usage::{ extract_provider_actual_service_tier_from_response, - extract_provider_reasoning_effort_from_body, extract_provider_service_tier_from_body, - normalize_provider_service_tier, resolve_provider_cache_ttl_minutes, + extract_provider_reasoning_effort_from_body, extract_provider_response_model_from_bodies, + extract_provider_service_tier_from_body, normalize_provider_service_tier, + resolve_provider_cache_ttl_minutes, sanitize_usage_request_metadata as project_usage_request_metadata, sanitize_usage_request_metadata_object as project_usage_request_metadata_object, sanitize_usage_request_metadata_ref as project_usage_request_metadata_ref, - UsageBodyCaptureState, PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, - PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY, + usage_body_capture_is_authoritative, UsageBodyCaptureState, + PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, + PROVIDER_REASONING_EFFORT_METADATA_KEY, PROVIDER_RESPONSE_MODEL_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, }; use serde_json::{Map, Value}; @@ -249,6 +251,83 @@ pub(crate) fn attach_provider_response_body_metadata( attach_provider_actual_service_tier_metadata(metadata, actual_service_tier.as_deref()) } +pub(crate) fn attach_provider_response_model_metadata( + metadata: Option, + request_body: Option<&Value>, + request_body_state: Option, + request_api_format: Option<&str>, + response_body: Option<&Value>, + response_body_state: Option, + provider_api_format: Option<&str>, +) -> Option { + let both_bodies_are_authoritative = + usage_body_capture_is_authoritative(request_body, request_body_state) + && usage_body_capture_is_authoritative(response_body, response_body_state); + let response_model = extract_provider_response_model_from_bodies( + request_body, + request_body_state, + request_api_format, + response_body, + response_body_state, + provider_api_format, + ); + if !both_bodies_are_authoritative && response_model.is_none() { + return metadata; + } + + let mut object = match metadata { + Some(Value::Object(object)) => object, + _ => Map::new(), + }; + // 完整终态 body 是最终候选的权威事实;相同、无效或缺失模型都要清除旧候选值。 + if both_bodies_are_authoritative { + object.remove(PROVIDER_RESPONSE_MODEL_METADATA_KEY); + } + if let Some(response_model) = response_model { + object.insert( + PROVIDER_RESPONSE_MODEL_METADATA_KEY.to_string(), + Value::String(response_model), + ); + } + (!object.is_empty()).then_some(Value::Object(object)) +} + +/// 终态候选无法完成比较时,显式清除旧响应模型,避免重试/故障转移残留。 +pub(crate) fn refresh_provider_response_model_metadata( + metadata: Option, + request_body: Option<&Value>, + request_body_state: Option, + request_api_format: Option<&str>, + response_body: Option<&Value>, + response_body_state: Option, + provider_api_format: Option<&str>, +) -> Option { + let mut object = match metadata { + Some(Value::Object(object)) => object, + _ => Map::new(), + }; + object.remove(PROVIDER_RESPONSE_MODEL_METADATA_KEY); + + if usage_body_capture_is_authoritative(request_body, request_body_state) + && usage_body_capture_is_authoritative(response_body, response_body_state) + { + if let Some(response_model) = extract_provider_response_model_from_bodies( + request_body, + request_body_state, + request_api_format, + response_body, + response_body_state, + provider_api_format, + ) { + object.insert( + PROVIDER_RESPONSE_MODEL_METADATA_KEY.to_string(), + Value::String(response_model), + ); + } + } + (!object.is_empty()).then_some(Value::Object(object)) +} + /// Refreshes the response-derived tier for a terminal snapshot. Complete response objects are /// authoritative even when they contain no tier (which clears a stale candidate value). Capture /// placeholders/absent bodies are not authoritative, so a terminal summary already present in @@ -312,6 +391,7 @@ pub(crate) fn attach_provider_actual_service_tier_metadata( #[cfg(test)] mod tests { use aether_contracts::{ExecutionPlan, RequestBody}; + use aether_data_contracts::repository::usage::UsageBodyCaptureState; use serde_json::{json, Value}; use std::collections::BTreeMap; @@ -323,8 +403,9 @@ mod tests { use super::{ attach_client_request_body_metadata, attach_provider_actual_service_tier_metadata, attach_provider_request_body_metadata, attach_provider_response_body_metadata, - build_usage_request_metadata_seed, merge_usage_request_metadata, - merge_usage_request_metadata_owned, refresh_provider_response_body_metadata, + attach_provider_response_model_metadata, build_usage_request_metadata_seed, + merge_usage_request_metadata, merge_usage_request_metadata_owned, + refresh_provider_response_body_metadata, refresh_provider_response_model_metadata, retain_first_byte_request_metadata, sanitize_usage_request_metadata, sanitize_usage_request_metadata_ref, }; @@ -802,6 +883,44 @@ mod tests { ); } + #[test] + fn response_model_metadata_is_independent_from_mapping_and_clears_stale_values() { + let metadata = attach_provider_response_model_metadata( + Some(json!({"provider_response_model": "old-model", "trace_id": "trace-1"})), + Some(&json!({"model": "gpt-5"})), + Some(UsageBodyCaptureState::Inline), + Some("openai:chat"), + Some(&json!({"model": "gpt-5.1"})), + Some(UsageBodyCaptureState::Inline), + Some("openai:chat"), + ) + .expect("response model should be attached"); + assert_eq!(metadata["provider_response_model"], "gpt-5.1"); + assert_eq!(metadata["trace_id"], "trace-1"); + + let metadata = refresh_provider_response_model_metadata( + Some(json!({"provider_response_model": "gpt-5.1"})), + Some(&json!({"model": "gpt-5"})), + Some(UsageBodyCaptureState::Inline), + Some("openai:chat"), + Some(&json!({"model": "gpt-5"})), + Some(UsageBodyCaptureState::Inline), + Some("openai:chat"), + ); + assert!(metadata.is_none()); + + let metadata = refresh_provider_response_model_metadata( + Some(json!({"provider_response_model": "gpt-5.1"})), + None, + Some(UsageBodyCaptureState::Disabled), + Some("openai:chat"), + Some(&json!({"model": "gpt-5.2"})), + Some(UsageBodyCaptureState::Inline), + Some("openai:chat"), + ); + assert!(metadata.is_none()); + } + #[test] fn terminal_response_refresh_replaces_stale_actual_tier() { let metadata = refresh_provider_response_body_metadata( diff --git a/crates/aether-usage/runtime/src/runtime.rs b/crates/aether-usage/runtime/src/runtime.rs index cb02c868f..48c1354f9 100644 --- a/crates/aether-usage/runtime/src/runtime.rs +++ b/crates/aether-usage/runtime/src/runtime.rs @@ -21,9 +21,10 @@ use crate::executor::spawn_on_usage_background_runtime; use crate::queue::is_permanent_enqueue_error; use crate::request_metadata::{ attach_client_request_body_metadata, attach_provider_request_body_metadata, - attach_provider_response_body_metadata, clear_client_request_body_metadata, - clear_provider_request_body_metadata, request_body_derived_facts_action, - retain_first_byte_request_metadata, RequestBodyDerivedFactsAction, + attach_provider_response_body_metadata, attach_provider_response_model_metadata, + clear_client_request_body_metadata, clear_provider_request_body_metadata, + request_body_derived_facts_action, retain_first_byte_request_metadata, + RequestBodyDerivedFactsAction, }; use crate::settlement::{ reconcile_usage_policy_cost_for_event_with_result, settle_usage_with_reconciled_cost, @@ -4919,7 +4920,7 @@ impl UsageRuntime { &self, data: &T, queue: UsageQueue, - event: UsageEvent, + mut event: UsageEvent, ) -> TerminalPersistenceOutcome where T: UsageRuntimeAccess, @@ -4952,7 +4953,20 @@ impl UsageRuntime { .await; }; - if let Err(err) = queue.enqueue(&event).await { + let enqueue_result = match queue.encode_event(&event) { + Ok(encoded) => { + if encoded.diagnostics_omitted + && self + .try_write_terminal_direct_fallback(data, &mut event, "queue_wire_limit") + .await + { + return TerminalPersistenceOutcome::PersistedDirectly; + } + queue.enqueue_encoded(encoded).await + } + Err(err) => Err(err), + }; + if let Err(err) = enqueue_result { drop(_guard); if is_permanent_enqueue_error(&err) { return self @@ -5284,8 +5298,17 @@ fn preserve_request_facts_with_legacy_missing( fn preserve_provider_response_facts(event: &mut UsageEvent) { let metadata = event.data.request_metadata.take(); - event.data.request_metadata = + let metadata = attach_provider_response_body_metadata(metadata, event.data.response_body.as_ref()); + event.data.request_metadata = attach_provider_response_model_metadata( + metadata, + event.data.request_body.as_ref(), + event.data.request_body_state, + event.data.api_format.as_deref(), + event.data.response_body.as_ref(), + event.data.response_body_state, + event.data.endpoint_api_format.as_deref(), + ); } impl UsageQueueHealthSnapshot { @@ -6649,7 +6672,7 @@ mod tests { } use std::collections::BTreeMap; - use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; use std::time::Instant; @@ -7375,9 +7398,27 @@ mod tests { queue: Arc, policy_started: Arc, release_policy: Arc, + policy_released: Arc, policy_reads: Arc, } + impl BlockingPolicyQueueConfiguredUsageStore { + fn new(queue: Arc) -> Self { + Self { + queue, + policy_started: Arc::new(tokio::sync::Notify::new()), + release_policy: Arc::new(tokio::sync::Notify::new()), + policy_released: Arc::new(AtomicBool::new(false)), + policy_reads: Arc::new(AtomicUsize::new(0)), + } + } + + fn release_blocked_policy(&self) { + self.policy_released.store(true, Ordering::Release); + self.release_policy.notify_waiters(); + } + } + #[derive(Default)] struct FailingPolicyUsageStore { inner: NoRedisUsageStore, @@ -8409,7 +8450,18 @@ mod tests { async fn body_capture_policy(&self) -> Result { self.policy_reads.fetch_add(1, Ordering::AcqRel); self.policy_started.notify_one(); - self.release_policy.notified().await; + // Latch the gate: Notify is edge-triggered, and later policy reads + // (or a waiter that subscribed after a single notify) must not hang. + loop { + if self.policy_released.load(Ordering::Acquire) { + break; + } + let notified = self.release_policy.notified(); + if self.policy_released.load(Ordering::Acquire) { + break; + } + notified.await; + } Ok(UsageBodyCapturePolicy::default()) } } @@ -9043,15 +9095,39 @@ mod tests { .await .expect("a duplicate first-byte marker must release the terminal barrier"); - let records = store.records.lock().expect("records lock"); - assert_eq!( - records.len(), - 2, - "the duplicate first byte must be coalesced" - ); - assert_eq!(records[0].status, "streaming"); - assert_eq!(records[1].status, "completed"); - drop(records); + { + let records = store.records.lock().expect("records lock"); + assert_eq!( + records.len(), + 2, + "the duplicate first byte must be coalesced" + ); + assert_eq!(records[0].status, "streaming"); + assert_eq!(records[1].status, "completed"); + } + + // The terminal persistence notification can arrive before the submission + // dispatcher accounts for its completed task and releases admission. + timeout(Duration::from_secs(1), async { + loop { + let snapshot = runtime.metrics_snapshot(); + if snapshot.lifecycle_submission_pending == 0 + && snapshot.first_byte_persistence_pending == 0 + && snapshot.ordered_lifecycle_pending == 0 + && runtime + .lifecycle_submission + .state + .admission + .available_permits() + == CAPACITY + { + break; + } + sleep(Duration::from_millis(1)).await; + } + }) + .await + .expect("duplicate first-byte submission accounting should drain"); let snapshot = runtime.metrics_snapshot(); assert_eq!(snapshot.lifecycle_submission_pending, 0); @@ -9979,6 +10055,24 @@ mod tests { 2, "the later direct caller should receive its own bounded write attempt" ); + // Direct persistence can finish before the submission worker joins the + // barrier handoff and accounts for its completed slot. + timeout(Duration::from_secs(1), async { + loop { + let snapshot = runtime.metrics_snapshot(); + let submission = &runtime.lifecycle_submission.state; + if snapshot.terminal_submission_pending == 0 + && snapshot.ordered_lifecycle_pending == 0 + && snapshot.lifecycle_submission_pending == 0 + && submission.admission.available_permits() == submission.capacity + { + break; + } + sleep(Duration::from_millis(1)).await; + } + }) + .await + .expect("failed terminal submission accounting and admission should drain"); let snapshot = runtime.metrics_snapshot(); assert_eq!(snapshot.terminal_submission_pending, 0); assert_eq!(snapshot.ordered_lifecycle_pending, 0); @@ -10086,24 +10180,43 @@ mod tests { .await; assert_eq!(remaining_policy_panics.load(Ordering::Acquire), 0); - let records = records.lock().expect("records lock"); - assert_eq!( - records - .iter() - .filter(|record| record.request_id == healthy_request_id) - .count(), - 2, - "the same terminal shard should continue processing healthy requests" - ); - assert!( - records - .iter() - .filter(|record| record.request_id == failed_request_id) - .count() - == 1, - "only the later healthy attempt should persist for the panicked request" - ); - drop(records); + { + let records = records.lock().expect("records lock"); + assert_eq!( + records + .iter() + .filter(|record| record.request_id == healthy_request_id) + .count(), + 2, + "the same terminal shard should continue processing healthy requests" + ); + assert!( + records + .iter() + .filter(|record| record.request_id == failed_request_id) + .count() + == 1, + "only the later healthy attempt should persist for the panicked request" + ); + } + // The final direct attempt also submits a barrier whose worker may + // account for completion after the persistence call has returned. + timeout(Duration::from_secs(1), async { + loop { + let snapshot = runtime.metrics_snapshot(); + let submission = &runtime.lifecycle_submission.state; + if snapshot.terminal_submission_pending == 0 + && snapshot.ordered_lifecycle_pending == 0 + && snapshot.lifecycle_submission_pending == 0 + && submission.admission.available_permits() == submission.capacity + { + break; + } + sleep(Duration::from_millis(1)).await; + } + }) + .await + .expect("panicked terminal submission accounting and admission should drain"); let snapshot = runtime.metrics_snapshot(); assert_eq!(snapshot.terminal_submission_pending, 0); assert_eq!(snapshot.ordered_lifecycle_pending, 0); @@ -12325,12 +12438,9 @@ mod tests { async fn event_capture_budget_bounds_blocked_policy_waiters_and_releases_on_cancel_or_basic() { for limit in [0, 64 * 1024] { let runtime = UsageRuntime::new(UsageRuntimeConfig::default()).expect("runtime"); - let store = BlockingPolicyQueueConfiguredUsageStore { - queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())), - policy_started: Arc::new(tokio::sync::Notify::new()), - release_policy: Arc::new(tokio::sync::Notify::new()), - policy_reads: Arc::new(AtomicUsize::new(0)), - }; + let store = BlockingPolicyQueueConfiguredUsageStore::new(Arc::new( + RuntimeState::memory(MemoryRuntimeStateConfig::default()), + )); let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new( limit, )); @@ -12391,7 +12501,7 @@ mod tests { .await .expect("replacement policy read starts"); assert_eq!(budget.retained_bytes(), retained); - store.release_policy.notify_one(); + store.release_blocked_policy(); let event = timeout(Duration::from_secs(2), completing) .await .expect("Basic policy completes") @@ -12675,6 +12785,159 @@ mod tests { assert_eq!(queued.data.total_cost_usd, None); } + fn oversized_full_terminal_event(max_bytes: usize) -> (serde_json::Value, UsageEvent) { + let body = json!({"content": "full-body".repeat(max_bytes / 16)}); + let mut event = UsageEvent::new( + UsageEventType::Completed, + "oversized-full-capture", + UsageEventData { + provider_name: "openai".to_string(), + model: "gpt-5".to_string(), + status_code: Some(200), + total_tokens: Some(12), + request_body: Some(body.clone()), + provider_request_body: Some(body.clone()), + response_body: Some(body.clone()), + client_response_body: Some(body.clone()), + ..UsageEventData::default() + }, + ); + apply_usage_body_capture_policy_to_event( + UsageBodyCapturePolicy { + record_level: UsageRequestRecordLevel::Full, + }, + &mut event, + ); + (body, event) + } + + #[tokio::test] + async fn oversized_full_terminal_capture_is_persisted_without_queue_truncation() { + let config = UsageRuntimeConfig { + enabled: true, + queue_terminal_events: true, + consumer_block_ms: 1, + ..UsageRuntimeConfig::default() + }; + let queue_runner: Arc = + Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let queue = UsageQueue::new(Arc::clone(&queue_runner), config.clone()).unwrap(); + let store = EnrichmentCountingQueueStore { + records: Mutex::new(Vec::new()), + queue: queue_runner, + enrich_calls: AtomicUsize::new(0), + }; + let runtime = UsageRuntime::new(config.clone()).unwrap(); + let (body, event) = oversized_full_terminal_event(config.queue_payload_max_bytes); + assert!( + event + .to_bounded_stream_fields(config.queue_payload_max_bytes) + .unwrap() + .diagnostics_omitted + ); + + let outcome = runtime.enqueue_or_write_terminal(&store, event).await; + + assert_eq!( + outcome, + super::TerminalPersistenceOutcome::PersistedDirectly + ); + assert_eq!(store.enrich_calls.load(Ordering::Acquire), 1); + assert_eq!(queue.stats().await.unwrap().stream_length, 0); + let records = store.records.lock().unwrap(); + assert_eq!(records.len(), 1); + for captured in [ + &records[0].request_body, + &records[0].provider_request_body, + &records[0].response_body, + &records[0].client_response_body, + ] { + assert_eq!(captured.as_ref(), Some(&body)); + } + } + + #[tokio::test] + async fn oversized_full_terminal_capture_keeps_bounded_queue_fallback() { + for unavailable in ["writer", "write_failure", "worker_gate", "fallback_gate"] { + let config = UsageRuntimeConfig { + enabled: true, + queue_terminal_events: true, + consumer_block_ms: 1, + worker_record_concurrency_limit: Some(1), + ..UsageRuntimeConfig::default() + }; + let queue_runner: Arc = + Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let queue = UsageQueue::new(Arc::clone(&queue_runner), config.clone()).unwrap(); + queue.ensure_consumer_group().await.unwrap(); + let store = FailingWriteQueueConfiguredUsageStore { + queue: Arc::clone(&queue_runner), + upsert_attempts: Arc::new(AtomicUsize::new(0)), + }; + let queue_only = QueueOnlyUsageStore { + queue: queue_runner, + upsert_attempts: Arc::clone(&store.upsert_attempts), + }; + let runtime = UsageRuntime::new(config.clone()).unwrap(); + let worker_permit = (unavailable == "worker_gate").then(|| { + runtime + .worker_record_gate + .as_ref() + .unwrap() + .try_acquire() + .unwrap() + }); + let fallback_permit = (unavailable == "fallback_gate").then(|| { + runtime + .terminal_direct_fallback_state + .try_acquire() + .unwrap() + }); + let (_, event) = oversized_full_terminal_event(config.queue_payload_max_bytes); + + let outcome = if unavailable == "writer" { + runtime.enqueue_or_write_terminal(&queue_only, event).await + } else { + runtime.enqueue_or_write_terminal(&store, event).await + }; + + assert_eq!( + outcome, + super::TerminalPersistenceOutcome::Queued, + "{unavailable}" + ); + assert_eq!( + store.upsert_attempts.load(Ordering::Acquire), + usize::from(unavailable == "write_failure") + ); + let entries = queue + .read_group("oversized-capture-consumer") + .await + .unwrap(); + assert_eq!(entries.len(), 1); + assert!(entries[0].fields["payload"].len() <= config.queue_payload_max_bytes); + let queued = UsageEvent::from_stream_fields(&entries[0].fields).unwrap(); + assert_eq!(queued.data.total_tokens, Some(12)); + for (body, state) in [ + (&queued.data.request_body, queued.data.request_body_state), + ( + &queued.data.provider_request_body, + queued.data.provider_request_body_state, + ), + (&queued.data.response_body, queued.data.response_body_state), + ( + &queued.data.client_response_body, + queued.data.client_response_body_state, + ), + ] { + assert!(body.is_none()); + assert_eq!(state, Some(UsageBodyCaptureState::Truncated)); + } + assert_eq!(runtime.metrics_snapshot().terminal_enqueue_failed_total, 0); + drop((worker_permit, fallback_permit)); + } + } + #[tokio::test] async fn terminal_enqueue_failure_uses_bounded_direct_database_fallback() { let config = UsageRuntimeConfig { @@ -13398,12 +13661,7 @@ mod tests { Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let queue: Arc = tracked_queue.clone(); - let store = BlockingPolicyQueueConfiguredUsageStore { - queue, - policy_started: Arc::new(tokio::sync::Notify::new()), - release_policy: Arc::new(tokio::sync::Notify::new()), - policy_reads: Arc::new(AtomicUsize::new(0)), - }; + let store = BlockingPolicyQueueConfiguredUsageStore::new(queue); let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let request_id = "req-terminal-seed-waits-for-turn"; let plan = terminal_test_plan(request_id); @@ -13425,9 +13683,10 @@ mod tests { assert_eq!(blocked_snapshot.terminal_submission_in_flight, 0); assert!(blocked_snapshot.lifecycle_submission_pending >= 2); - store.release_policy.notify_waiters(); + store.release_blocked_policy(); timeout(Duration::from_secs(2), async { loop { + store.release_blocked_policy(); let snapshot = runtime.metrics_snapshot(); if tracked_queue.successful_appends.load(Ordering::Acquire) == 1 && snapshot.lifecycle_submission_pending == 0 @@ -13435,7 +13694,7 @@ mod tests { { break; } - tokio::task::yield_now().await; + sleep(Duration::from_millis(1)).await; } }) .await @@ -13468,12 +13727,7 @@ mod tests { Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let queue: Arc = tracked_queue.clone(); - let store = BlockingPolicyQueueConfiguredUsageStore { - queue, - policy_started: Arc::new(tokio::sync::Notify::new()), - release_policy: Arc::new(tokio::sync::Notify::new()), - policy_reads: Arc::new(AtomicUsize::new(0)), - }; + let store = BlockingPolicyQueueConfiguredUsageStore::new(queue); let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let policy_started = store.policy_started.notified(); @@ -13520,9 +13774,10 @@ mod tests { assert_eq!(blocked_snapshot.terminal_submission_in_flight, 1); assert!(blocked_snapshot.lifecycle_submission_pending <= BACKLOG + 1); - store.release_policy.notify_waiters(); + store.release_blocked_policy(); timeout(Duration::from_secs(5), async { loop { + store.release_blocked_policy(); let snapshot = runtime.metrics_snapshot(); if tracked_queue.successful_appends.load(Ordering::Acquire) == BACKLOG + 1 && snapshot.lifecycle_submission_pending == 0 @@ -13530,7 +13785,7 @@ mod tests { { break; } - tokio::task::yield_now().await; + sleep(Duration::from_millis(1)).await; } }) .await @@ -13701,7 +13956,7 @@ mod tests { .iter() .any(|record| record.request_id == blocked_request_id)); - release_build.notify_waiters(); + release_build.notify_one(); timeout(Duration::from_secs(2), async { loop { let snapshot = runtime.metrics_snapshot(); @@ -13828,12 +14083,7 @@ mod tests { Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let queue: Arc = tracked_queue.clone(); - let store = BlockingPolicyQueueConfiguredUsageStore { - queue, - policy_started: Arc::new(tokio::sync::Notify::new()), - release_policy: Arc::new(tokio::sync::Notify::new()), - policy_reads: Arc::new(AtomicUsize::new(0)), - }; + let store = BlockingPolicyQueueConfiguredUsageStore::new(queue); let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let policy_started = store.policy_started.notified(); runtime @@ -13892,10 +14142,10 @@ mod tests { .expect("terminal submissions should reach the execution backlog"); let saturated_snapshot = runtime.metrics_snapshot(); - store.release_policy.notify_waiters(); + store.release_blocked_policy(); let all_completed = timeout(Duration::from_secs(2), async { loop { - store.release_policy.notify_waiters(); + store.release_blocked_policy(); if tracked_queue.successful_appends.load(Ordering::Acquire) == EXCESS_SUBMISSIONS + 1 && runtime.metrics_snapshot().terminal_submission_in_flight == 0 diff --git a/crates/aether-usage/runtime/src/write.rs b/crates/aether-usage/runtime/src/write.rs index 72fbaaede..3e2956352 100644 --- a/crates/aether-usage/runtime/src/write.rs +++ b/crates/aether-usage/runtime/src/write.rs @@ -17,9 +17,10 @@ use crate::body_capture::{ }; use crate::request_metadata::{ attach_client_request_body_metadata, attach_provider_actual_service_tier_metadata, - attach_provider_request_body_metadata, build_usage_request_metadata_seed, - merge_usage_request_metadata, merge_usage_request_metadata_owned, - refresh_provider_response_body_metadata, sanitize_usage_request_metadata, + attach_provider_request_body_metadata, attach_provider_response_model_metadata, + build_usage_request_metadata_seed, merge_usage_request_metadata, + merge_usage_request_metadata_owned, refresh_provider_response_body_metadata, + refresh_provider_response_model_metadata, sanitize_usage_request_metadata, sanitize_usage_request_metadata_ref, }; use crate::{ @@ -727,6 +728,15 @@ fn build_terminal_usage_event_from_seed_impl( Some(model.as_str()), provider_request.as_ref(), ); + let request_metadata = attach_provider_response_model_metadata( + request_metadata, + request_body.as_ref(), + body_states.request_body_state, + Some(client_contract.as_str()), + provider_response.as_ref(), + body_states.response_body_state, + Some(provider_contract.as_str()), + ); let mut data = UsageEventData { user_id, @@ -1073,6 +1083,15 @@ pub fn build_sync_terminal_usage_seed( context_seed.request_metadata, provider_response_full.as_ref(), ); + let request_metadata = refresh_provider_response_model_metadata( + request_metadata, + context_seed.request_body.as_ref(), + context_seed.body_states.request_body_state, + Some(context_seed.client_contract.as_str()), + provider_response_full.as_ref(), + provider_response_body_state, + Some(context_seed.provider_contract.as_str()), + ); TerminalUsageSeed { token_measurement_source, @@ -1252,6 +1271,15 @@ pub fn build_stream_terminal_usage_seed( context_seed.request_metadata, provider_response_full.as_ref(), ); + let request_metadata = refresh_provider_response_model_metadata( + request_metadata, + context_seed.request_body.as_ref(), + context_seed.body_states.request_body_state, + Some(context_seed.client_contract.as_str()), + provider_response_full.as_ref(), + provider_response_body_state, + Some(context_seed.provider_contract.as_str()), + ); // The parser's terminal summary is authoritative when a response body is truncated or the // body and summary disagree; attach it after the body refresh so it wins. let request_metadata = attach_provider_actual_service_tier_metadata( @@ -3080,14 +3108,26 @@ fn parse_sse_body_for_storage(text: &str) -> Option { let mut chunks = Vec::new(); let mut total_chunks = 0_u64; let mut saw_done = false; + let mut first_parse_error = None; for_each_sse_payload(text, |payload| { if payload == "[DONE]" { saw_done = true; return; } total_chunks += 1; - if let Ok(json_body) = serde_json::from_str::(payload) { - chunks.push(json_body); + match serde_json::from_str::(payload) { + Ok(json_body) => chunks.push(json_body), + Err(error) if first_parse_error.is_none() => { + // A later valid event must not hide an earlier broken one. + // Store diagnostics only, without duplicating raw user content. + first_parse_error = Some(json!({ + "chunk_index": total_chunks - 1, + "line": error.line(), + "column": error.column(), + "message": error.to_string(), + })); + } + Err(_) => {} } }); if total_chunks == 0 && !saw_done { @@ -3104,6 +3144,11 @@ fn parse_sse_body_for_storage(text: &str) -> Option { if saw_done { metadata.insert("has_completion".to_string(), Value::Bool(true)); } + if let Some(error) = first_parse_error { + // Capture truncation can also cause a parse error; this describes the + // captured payload, not an assertion that the provider sent bad JSON. + metadata.insert("first_parse_error".to_string(), error); + } if stored_chunks < total_chunks { metadata.insert( "dropped_chunks".to_string(), @@ -7221,6 +7266,20 @@ mod tests { ); } + #[test] + fn parse_sse_body_for_storage_reports_bad_event_before_valid_terminal() { + let body = concat!( + "data: {\"tools\":[}\n\n", + "data: {\"type\":\"response.completed\"}\n\n", + ); + let parsed = parse_sse_body_for_storage(body).unwrap(); + assert_eq!(parsed["metadata"]["dropped_chunks"], 1); + assert_eq!(parsed["metadata"]["first_parse_error"]["chunk_index"], 0); + assert!(parsed["metadata"]["first_parse_error"]["message"].is_string()); + assert_eq!(parsed["chunks"][0]["type"], "response.completed"); + assert!(parsed.get("raw_response").is_none()); + } + #[test] fn extract_token_counts_from_value_handles_crlf_and_cr_sse_text() { let sse_body = concat!( diff --git a/crates/aether-video-tasks-core/Cargo.toml b/crates/aether-video-tasks-core/Cargo.toml index 25f284b52..f50899619 100644 --- a/crates/aether-video-tasks-core/Cargo.toml +++ b/crates/aether-video-tasks-core/Cargo.toml @@ -12,5 +12,6 @@ aether-data-contracts.workspace = true async-trait.workspace = true serde.workspace = true serde_json.workspace = true +sha2.workspace = true url.workspace = true uuid.workspace = true diff --git a/crates/aether-video-tasks-core/src/openai.rs b/crates/aether-video-tasks-core/src/openai.rs index 8b0d657cb..5738f07d1 100644 --- a/crates/aether-video-tasks-core/src/openai.rs +++ b/crates/aether-video-tasks-core/src/openai.rs @@ -43,6 +43,27 @@ pub fn map_openai_stored_task_to_read_response( } fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus) -> Value { + if task.client_api_format.as_deref() == Some("xai:video") { + let mut body = json!({"status":match status { + VideoTaskStatus::Completed => "done", + VideoTaskStatus::Expired => "expired", + VideoTaskStatus::Failed | VideoTaskStatus::Cancelled | VideoTaskStatus::Deleted => "failed", + _ => "pending", + }}); + if let Some(model) = task.model { + body["model"] = json!(model); + } + if let Some(url) = task.video_url { + body["video"] = json!({"url":url}); + if let Some(duration) = task.duration_seconds { + body["video"]["duration"] = json!(duration); + } + } + if status == VideoTaskStatus::Failed { + body["error"] = json!({"code":sanitize_video_task_error_code(task.error_code).unwrap_or_else(|| "unknown".into()),"message":"Video generation failed"}); + } + return body; + } let mut body = json!({ "id": task.id, "object": "video", @@ -57,6 +78,9 @@ fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus) if let Some(prompt) = task.prompt { body["prompt"] = Value::String(prompt); } + if let Some(seconds) = task.duration_seconds { + body["seconds"] = json!(seconds.to_string()); + } if let Some(size) = task.size { body["size"] = Value::String(size); } @@ -91,21 +115,97 @@ fn map_openai_stored_task_status(status: VideoTaskStatus) -> &'static str { } impl OpenAiVideoTaskSeed { + pub fn uses_xai_provider(&self) -> bool { + self.xai_provider || self.is_xai_native() + } + + pub fn is_xai_native(&self) -> bool { + self.persistence.client_api_format == "xai:video" + } + + pub fn native_create_body_json(&self) -> Value { + let mut body = self.native_response.clone().unwrap_or_else(|| json!({})); + body["request_id"] = json!(self.local_task_id); + if body.get("id").is_some() { + body["id"] = json!(self.local_task_id); + } + body + } + + fn native_read_body_json(&self) -> Value { + if let Some(mut body) = self.native_response.clone().filter(|body| { + body.get("status").is_some() + || body.get("error").is_some() + || body.get("code").is_some() + }) { + if body.get("request_id").is_some() { + body["request_id"] = json!(self.local_task_id); + } + if body.get("id").is_some() { + body["id"] = json!(self.local_task_id); + } + return body; + } + let mut body = json!({"status":match self.status { + LocalVideoTaskStatus::Completed => "done", + LocalVideoTaskStatus::Expired => "expired", + LocalVideoTaskStatus::Failed | LocalVideoTaskStatus::Cancelled | LocalVideoTaskStatus::Deleted => "failed", + _ => "pending", + }}); + if let Some(model) = &self.model { + body["model"] = json!(model); + } + if let Some(url) = &self.video_url { + body["video"] = json!({"url":url}); + if let Some(duration) = self.seconds.as_deref().and_then(|v| v.parse::().ok()) { + body["video"]["duration"] = json!(duration); + } + } + if self.error_code.is_some() { + body["error"] = json!({"code":self.error_code,"message":"Video generation failed"}); + } + body + } + pub fn apply_provider_body(&mut self, provider_body: &Map) { + if self.uses_xai_provider() { + self.native_response = Some(Value::Object(provider_body.clone())); + } + let raw_status = provider_body .get("status") .and_then(Value::as_str) .map(str::trim) .unwrap_or_default(); - self.status = match raw_status { - "queued" => LocalVideoTaskStatus::Queued, - "processing" => LocalVideoTaskStatus::Processing, - "completed" => LocalVideoTaskStatus::Completed, - "failed" => LocalVideoTaskStatus::Failed, - "cancelled" => LocalVideoTaskStatus::Cancelled, + // Accept xAI's native lifecycle vocabulary alongside OpenAI's fields. + self.status = match raw_status.to_ascii_lowercase().as_str() { + "queued" | "pending" => LocalVideoTaskStatus::Queued, + "processing" | "in_progress" | "running" => LocalVideoTaskStatus::Processing, + "completed" | "done" | "succeeded" | "success" => LocalVideoTaskStatus::Completed, + "failed" | "error" => LocalVideoTaskStatus::Failed, + "cancelled" | "canceled" => LocalVideoTaskStatus::Cancelled, "expired" => LocalVideoTaskStatus::Expired, _ => LocalVideoTaskStatus::Submitted, }; + let error = provider_body.get("error").filter(|value| !value.is_null()); + let error_code = provider_body + .get("code") + .and_then(Value::as_str) + .filter(|value| !value.trim().is_empty()) + .or_else(|| { + error + .and_then(|value| value.get("code")) + .and_then(Value::as_str) + }); + // xAI may report a failed job as a 200 response with code/error only. + if (error.is_some() || error_code.is_some()) + && !matches!( + self.status, + LocalVideoTaskStatus::Cancelled | LocalVideoTaskStatus::Expired + ) + { + self.status = LocalVideoTaskStatus::Failed; + } self.progress_percent = provider_body .get("progress") .and_then(Value::as_u64) @@ -117,20 +217,35 @@ impl OpenAiVideoTaskSeed { }); self.completed_at_unix_secs = provider_body.get("completed_at").and_then(Value::as_u64); self.expires_at_unix_secs = provider_body.get("expires_at").and_then(Value::as_u64); - let error = provider_body.get("error").and_then(Value::as_object); - self.error_code = sanitize_video_task_error_code( - error - .and_then(|value| value.get("code")) - .and_then(Value::as_str) - .map(str::to_string), - ); + self.error_code = sanitize_video_task_error_code(error_code.map(str::to_string)); self.error_message = None; self.video_url = provider_body .get("video_url") .or_else(|| provider_body.get("url")) .or_else(|| provider_body.get("result_url")) + .or_else(|| { + provider_body + .get("video") + .and_then(|video| video.get("url")) + }) .and_then(Value::as_str) .map(str::to_string); + if let Some(seconds) = provider_body + .get("seconds") + .or_else(|| { + provider_body + .get("video") + .and_then(|video| video.get("duration")) + }) + .filter(|value| value.is_string() || value.is_number()) + { + self.seconds = Some( + seconds + .as_str() + .map(str::to_string) + .unwrap_or_else(|| seconds.to_string()), + ); + } } pub fn build_content_stream_action( @@ -239,6 +354,9 @@ impl OpenAiVideoTaskSeed { } pub fn client_body_json(&self) -> Value { + if self.is_xai_native() { + return self.native_read_body_json(); + } let mut body = json!({ "id": self.local_task_id, "object": "video", @@ -259,6 +377,9 @@ impl OpenAiVideoTaskSeed { if let Some(seconds) = &self.seconds { body["seconds"] = Value::String(seconds.clone()); } + if let Some(video_url) = &self.video_url { + body["video_url"] = Value::String(video_url.clone()); + } if let Some(remixed_from_video_id) = &self.remixed_from_video_id { body["remixed_from_video_id"] = Value::String(remixed_from_video_id.clone()); } @@ -357,12 +478,20 @@ impl OpenAiVideoTaskSeed { } pub fn build_get_follow_up_plan(&self, trace_id: &str) -> Option { - if !matches!( + let refreshable = matches!( self.status, LocalVideoTaskStatus::Submitted | LocalVideoTaskStatus::Queued | LocalVideoTaskStatus::Processing - ) { + ) || (self.uses_xai_provider() + && self.native_response.is_none() + && matches!( + self.status, + LocalVideoTaskStatus::Completed + | LocalVideoTaskStatus::Failed + | LocalVideoTaskStatus::Expired + )); + if !refreshable { return None; } @@ -573,7 +702,12 @@ impl OpenAiVideoTaskSeed { }; let mut record = UpsertVideoTask { id: self.local_task_id.clone(), - short_id: None, + // The production schema requires a unique, non-null short_id (at most 16 chars). + // Derive it deterministically so repeated capture and legacy snapshot reloads agree. + short_id: Some(self.local_short_id.clone().unwrap_or_else(|| { + use sha2::{Digest, Sha256}; + format!("{:x}", Sha256::digest(self.local_task_id.as_bytes()))[..16].to_string() + })), request_id: self.persistence.request_id.clone(), user_id: self.user_id.clone(), api_key_id: self.api_key_id.clone(), @@ -589,7 +723,11 @@ impl OpenAiVideoTaskSeed { model: self.model.clone().or_else(|| Some(String::new())), prompt: self.prompt.clone().or_else(|| Some(String::new())), original_request_body: None, - duration_seconds: request_body_u32(&self.persistence.original_request_body, "seconds"), + duration_seconds: self + .seconds + .as_deref() + .and_then(|value| value.parse().ok()) + .or_else(|| request_body_u32(&self.persistence.original_request_body, "seconds")), resolution: request_body_string(&self.persistence.original_request_body, "resolution"), aspect_ratio: request_body_string( &self.persistence.original_request_body, @@ -697,6 +835,9 @@ mod tests { #[test] fn builds_minimal_openai_persistence_record_without_sensitive_snapshot() { let seed = OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-openai-sensitive".to_string(), upstream_task_id: "upstream-openai-sensitive".to_string(), created_at_unix_ms: 1_712_345_678, @@ -747,6 +888,12 @@ mod tests { let record = seed.to_upsert_record(); + let short_id = record + .short_id + .as_deref() + .expect("database short_id is required"); + assert_eq!(short_id.len(), 16); + assert_eq!(seed.to_upsert_record().short_id, record.short_id); assert_eq!(record.error_code.as_deref(), Some("provider_error")); assert!(record.original_request_body.is_none()); assert!(record.progress_message.is_none()); @@ -759,6 +906,8 @@ mod tests { let mut stored = record.into_stored(); stored.status = VideoTaskStatus::Completed; + // Migrated tasks can already have a short ID unrelated to the derived ID. + stored.short_id = Some("legacy-short-id".to_string()); let snapshot = LocalVideoTaskSnapshot::from_stored_task_with_transport(&stored, seed.transport) .expect("stored task should reconstruct with current transport"); @@ -766,6 +915,21 @@ mod tests { panic!("expected OpenAI snapshot"); }; assert_eq!(restored.prompt, stored.prompt); + assert_eq!(restored.to_upsert_record().short_id, stored.short_id); + let mut embedded = stored.clone(); + let mut legacy_snapshot = + serde_json::to_value(LocalVideoTaskSnapshot::OpenAi(restored.clone())).unwrap(); + legacy_snapshot["OpenAi"] + .as_object_mut() + .unwrap() + .remove("local_short_id"); + embedded.request_metadata = Some(json!({"rust_local_snapshot": legacy_snapshot})); + let embedded_snapshot = LocalVideoTaskSnapshot::from_stored_task(&embedded) + .expect("legacy embedded snapshot should hydrate"); + assert_eq!( + embedded_snapshot.to_upsert_record().short_id, + stored.short_id + ); assert_eq!(restored.to_upsert_record().video_url, stored.video_url); let Some(LocalVideoTaskContentAction::StreamPlan(plan)) = restored.build_content_stream_action(None, "trace-download") diff --git a/crates/aether-video-tasks-core/src/path.rs b/crates/aether-video-tasks-core/src/path.rs index c266fd4e0..f04251770 100644 --- a/crates/aether-video-tasks-core/src/path.rs +++ b/crates/aether-video-tasks-core/src/path.rs @@ -8,7 +8,9 @@ use uuid::Uuid; use crate::{LocalVideoTaskRegistryMutation, LocalVideoTaskStatus, VideoTaskTruthSourceMode}; pub fn extract_openai_task_id_from_path(path: &str) -> Option<&str> { - let suffix = path.strip_prefix("/v1/videos/")?; + let suffix = path + .strip_prefix("/v1/videos/") + .or_else(|| path.strip_prefix("/openai/v1/videos/"))?; if suffix.is_empty() || suffix.contains('/') || suffix.ends_with(":cancel") @@ -29,21 +31,27 @@ pub fn extract_gemini_short_id_from_path(path: &str) -> Option<&str> { } pub fn extract_openai_task_id_from_cancel_path(path: &str) -> Option<&str> { - let suffix = path.strip_prefix("/v1/videos/")?; + let suffix = path + .strip_prefix("/v1/videos/") + .or_else(|| path.strip_prefix("/openai/v1/videos/"))?; suffix .strip_suffix("/cancel") .filter(|value| !value.is_empty()) } pub fn extract_openai_task_id_from_remix_path(path: &str) -> Option<&str> { - let suffix = path.strip_prefix("/v1/videos/")?; + let suffix = path + .strip_prefix("/v1/videos/") + .or_else(|| path.strip_prefix("/openai/v1/videos/"))?; suffix .strip_suffix("/remix") .filter(|value| !value.is_empty()) } pub fn extract_openai_task_id_from_content_path(path: &str) -> Option<&str> { - let suffix = path.strip_prefix("/v1/videos/")?; + let suffix = path + .strip_prefix("/v1/videos/") + .or_else(|| path.strip_prefix("/openai/v1/videos/"))?; suffix .strip_suffix("/content") .filter(|value| !value.is_empty()) diff --git a/crates/aether-video-tasks-core/src/read_side.rs b/crates/aether-video-tasks-core/src/read_side.rs index fd8749ce1..003dfcdc1 100644 --- a/crates/aether-video-tasks-core/src/read_side.rs +++ b/crates/aether-video-tasks-core/src/read_side.rs @@ -73,7 +73,7 @@ async fn read_openai_video_task_response( } None => state.find_stored_video_task(lookup).await?, }; - let Some(task) = task else { + let Some(mut task) = task else { return Ok(None); }; @@ -81,6 +81,9 @@ async fn read_openai_video_task_response( return Ok(None); } + if request_path.starts_with("/openai/v1/videos/") { + task.client_api_format = Some("openai:video".into()); + } Ok(Some(map_openai_stored_task_to_read_response(task))) } diff --git a/crates/aether-video-tasks-core/src/service.rs b/crates/aether-video-tasks-core/src/service.rs index c94fc00db..a4ba50545 100644 --- a/crates/aether-video-tasks-core/src/service.rs +++ b/crates/aether-video-tasks-core/src/service.rs @@ -105,13 +105,8 @@ impl VideoTaskService { if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative { return None; } - match route_family { - Some("openai") => extract_openai_task_id_from_path(request_path) - .and_then(|task_id| self.store.read_openai(task_id)), - Some("gemini") => extract_gemini_short_id_from_path(request_path) - .and_then(|short_id| self.store.read_gemini(short_id)), - _ => None, - } + self.snapshot_for_route(route_family, request_path) + .map(|snapshot| snapshot.read_response_for_path(request_path)) } pub fn read_response_for_user( @@ -126,7 +121,7 @@ impl VideoTaskService { let snapshot = self.snapshot_for_route(route_family, request_path)?; snapshot .belongs_to_user(user_id) - .then(|| snapshot.read_response()) + .then(|| snapshot.read_response_for_path(request_path)) } pub fn snapshot_for_route( diff --git a/crates/aether-video-tasks-core/src/snapshot.rs b/crates/aether-video-tasks-core/src/snapshot.rs index 6aea1cb21..7979dc051 100644 --- a/crates/aether-video-tasks-core/src/snapshot.rs +++ b/crates/aether-video-tasks-core/src/snapshot.rs @@ -29,6 +29,7 @@ impl LocalVideoTaskSnapshot { // contain stale identity fields after a task import or repair. match &mut snapshot { Self::OpenAi(seed) => { + seed.local_short_id = task.short_id.clone(); seed.user_id = task.user_id.clone(); seed.api_key_id = task.api_key_id.clone(); } @@ -51,6 +52,9 @@ impl LocalVideoTaskSnapshot { "openai:video" => { let upstream_task_id = non_empty_owned(task.external_task_id.as_ref())?; Some(Self::OpenAi(OpenAiVideoTaskSeed { + local_short_id: task.short_id.clone(), + native_response: None, + xai_provider: persistence.client_api_format == "xai:video", local_task_id: task.id.clone(), upstream_task_id, created_at_unix_ms: task.created_at_unix_ms, @@ -142,6 +146,19 @@ impl LocalVideoTaskSnapshot { } } + pub fn read_response_for_path(&self, path: &str) -> LocalVideoTaskReadResponse { + if let Self::OpenAi(seed) = self { + let mut seed = seed.clone(); + if path.starts_with("/openai/v1/videos/") { + seed.persistence.client_api_format = "openai:video".to_string(); + } else if path.starts_with("/v1/videos/") && seed.uses_xai_provider() { + seed.persistence.client_api_format = "xai:video".to_string(); + } + return Self::OpenAi(seed).read_response(); + } + self.read_response() + } + pub fn read_response(&self) -> LocalVideoTaskReadResponse { match self { Self::OpenAi(seed) => match seed.status { diff --git a/crates/aether-video-tasks-core/src/sync.rs b/crates/aether-video-tasks-core/src/sync.rs index 2bf49bd57..11ef4b8ce 100644 --- a/crates/aether-video-tasks-core/src/sync.rs +++ b/crates/aether-video-tasks-core/src/sync.rs @@ -19,14 +19,17 @@ impl LocalVideoTaskSeed { ) -> Option { let transport = LocalVideoTaskTransport::from_plan(plan)?; let persistence = LocalVideoTaskPersistence::from_report_context(report_context, plan); - match report_kind { + let mut seed = match report_kind { "openai_video_create_sync_finalize" => { - let upstream_id = provider_body.get("id").and_then(Value::as_str)?.trim(); - if upstream_id.is_empty() { - return None; - } + let upstream_id = openai_video_provider_task_id(provider_body)?; Some(Self::OpenAiCreate(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: report_context + .get("video_provider_xai") + .and_then(Value::as_bool) + .unwrap_or(false), local_task_id: context_text(report_context, "local_task_id") .unwrap_or_else(|| Uuid::new_v4().to_string()), upstream_task_id: upstream_id.to_string(), @@ -37,8 +40,12 @@ impl LocalVideoTaskSeed { model: context_text(report_context, "model") .or_else(|| request_body_text(report_context, "model")), prompt: request_body_text(report_context, "prompt"), - size: request_body_text(report_context, "size"), - seconds: request_body_text(report_context, "seconds"), + size: context_text(report_context, "video_size") + .or_else(|| request_body_text(report_context, "size")), + seconds: context_u64(report_context, "video_duration") + .map(|v| v.to_string()) + .or_else(|| request_body_text(report_context, "seconds")) + .or_else(|| request_body_text(report_context, "duration")), remixed_from_video_id: None, status: LocalVideoTaskStatus::Submitted, progress_percent: 0, @@ -52,12 +59,15 @@ impl LocalVideoTaskSeed { })) } "openai_video_remix_sync_finalize" => { - let upstream_id = provider_body.get("id").and_then(Value::as_str)?.trim(); - if upstream_id.is_empty() { - return None; - } + let upstream_id = openai_video_provider_task_id(provider_body)?; Some(Self::OpenAiRemix(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: report_context + .get("video_provider_xai") + .and_then(Value::as_bool) + .unwrap_or(false), local_task_id: context_text(report_context, "local_task_id") .unwrap_or_else(|| Uuid::new_v4().to_string()), upstream_task_id: upstream_id.to_string(), @@ -68,8 +78,12 @@ impl LocalVideoTaskSeed { model: context_text(report_context, "model") .or_else(|| request_body_text(report_context, "model")), prompt: request_body_text(report_context, "prompt"), - size: request_body_text(report_context, "size"), - seconds: request_body_text(report_context, "seconds"), + size: context_text(report_context, "video_size") + .or_else(|| request_body_text(report_context, "size")), + seconds: context_u64(report_context, "video_duration") + .map(|v| v.to_string()) + .or_else(|| request_body_text(report_context, "seconds")) + .or_else(|| request_body_text(report_context, "duration")), remixed_from_video_id: context_text(report_context, "task_id") .or_else(|| request_body_text(report_context, "remix_video_id")), status: LocalVideoTaskStatus::Submitted, @@ -110,7 +124,11 @@ impl LocalVideoTaskSeed { })) } _ => None, + }?; + if let Self::OpenAiCreate(task) | Self::OpenAiRemix(task) = &mut seed { + task.apply_provider_body(provider_body); } + Some(seed) } pub fn success_report_kind(&self) -> &'static str { @@ -144,12 +162,28 @@ impl LocalVideoTaskSeed { pub fn client_body_json(&self) -> Value { match self { - Self::OpenAiCreate(seed) | Self::OpenAiRemix(seed) => seed.client_body_json(), + Self::OpenAiCreate(seed) | Self::OpenAiRemix(seed) => { + if seed.is_xai_native() { + seed.native_create_body_json() + } else { + seed.client_body_json() + } + } Self::GeminiCreate(seed) => seed.client_body_json(), } } } +fn openai_video_provider_task_id(body: &Map) -> Option<&str> { + // xAI's OpenAI-compatible video creation returns request_id instead of id. + ["id", "request_id"].into_iter().find_map(|field| { + body.get(field) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + }) +} + impl VideoTaskTruthSourceMode { pub fn prepare_sync_success( self, @@ -353,6 +387,234 @@ mod tests { resolve_local_sync_success_background_report_kind, }; + #[test] + fn xai_native_video_protocol_survives_persistence_and_preserves_provider_fields() { + use crate::{ + LocalVideoTaskContentAction, LocalVideoTaskSnapshot, VideoTaskService, + VideoTaskTruthSourceMode, + }; + let mut plan = + build_internal_finalize_video_plan("native-create", "openai:video", None).unwrap(); + plan.url = "https://api.x.ai/v1/videos/generations".into(); + plan.headers + .insert("authorization".into(), "Bearer test-key".into()); + let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); + let context = json!({"local_task_id":"native-local", "user_id":"owner", "model":"grok-imagine-video", "video_client_protocol":"xai", "video_duration":6}); + let success = service + .prepare_sync_success( + "openai_video_create_sync_finalize", + json!({"request_id":"native-upstream", "future_field":true}) + .as_object() + .unwrap(), + context.as_object().unwrap(), + &plan, + ) + .unwrap(); + assert_eq!( + success.client_body_json(), + json!({"request_id":"native-local","future_field":true}) + ); + let mut snapshot = success.to_snapshot(); + let body = json!({"status":"done","video":{"url":"https://vidgen.x.ai/video.mp4","duration":6,"respect_moderation":true},"future_field":[1,2]}); + snapshot.apply_provider_body(body.as_object().unwrap()); + assert_eq!(snapshot.read_response().body_json, body); + assert_eq!( + snapshot + .read_response_for_path("/openai/v1/videos/native-local") + .body_json["status"], + "completed" + ); + let LocalVideoTaskSnapshot::OpenAi(seed) = &snapshot else { + panic!("openai task expected") + }; + let Some(LocalVideoTaskContentAction::StreamPlan(download)) = + seed.build_content_stream_action(None, "download") + else { + panic!("download expected") + }; + assert_eq!(download.url, "https://vidgen.x.ai/video.mp4"); + assert!(download.headers.is_empty()); + let stored = snapshot.to_upsert_record().into_stored(); + assert!(stored.request_metadata.is_none()); + assert_eq!(stored.client_api_format.as_deref(), Some("xai:video")); + let restored = LocalVideoTaskSnapshot::from_stored_task_with_transport( + &stored, + seed.transport.clone(), + ) + .unwrap(); + assert_eq!(restored.read_response().body_json["status"], "done"); + service.record_snapshot(restored); + let poll = service + .prepare_read_refresh_sync_plan_for_user( + Some("openai"), + "/v1/videos/native-local", + "owner", + "poll", + ) + .unwrap(); + assert_eq!(poll.plan.url, "https://api.x.ai/v1/videos/native-upstream"); + assert!(service + .prepare_read_refresh_sync_plan_for_user( + Some("openai"), + "/v1/videos/native-local", + "foreign", + "poll" + ) + .is_none()); + assert!(service.apply_read_refresh_projection(&poll, body.as_object().unwrap())); + assert_eq!( + service + .read_response_for_user(Some("openai"), "/v1/videos/native-local", "owner") + .unwrap() + .body_json, + body + ); + } + + #[test] + fn xai_video_lifecycle_creates_polls_persists_and_downloads() { + use crate::{ + LocalVideoTaskContentAction, LocalVideoTaskSnapshot, VideoTaskService, + VideoTaskTruthSourceMode, + }; + for api_root in ["https://cli-chat-proxy.grok.com/v1", "https://api.x.ai/v1"] { + let mut plan = + build_internal_finalize_video_plan("xai-create", "openai:video", None).unwrap(); + plan.url = format!("{api_root}/videos/generations"); + plan.headers + .insert("authorization".into(), "Bearer test-token".into()); + let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); + let context = json!({"local_task_id": "local-video", "model": "grok-imagine-video", "original_request_body": {"prompt": "A cat", "seconds": "6"}}); + let success = service + .prepare_sync_success( + "openai_video_create_sync_finalize", + json!({"request_id": "xai-request"}).as_object().unwrap(), + context.as_object().unwrap(), + &plan, + ) + .unwrap(); + assert_eq!(success.client_body_json()["id"], "local-video"); + assert_eq!(success.client_body_json()["status"], "queued"); + let snapshot = success.to_snapshot(); + assert_eq!( + snapshot.to_upsert_record().external_task_id.as_deref(), + Some("xai-request") + ); + service.record_snapshot(snapshot.clone()); + let poll = service + .prepare_poll_refresh_plan_for_snapshot(snapshot, "xai-poll") + .unwrap(); + assert_eq!(poll.plan.method, "GET"); + assert_eq!(poll.plan.url, format!("{api_root}/videos/xai-request")); + assert_eq!( + poll.plan.headers.get("authorization"), + plan.headers.get("authorization") + ); + assert!(service.apply_read_refresh_projection( + &poll, + json!({"status": "pending"}).as_object().unwrap() + )); + assert_eq!( + service + .read_response(Some("openai"), "/v1/videos/local-video") + .unwrap() + .body_json["status"], + "queued" + ); + assert!(service.apply_read_refresh_projection(&poll, json!({ + "status": "done", "video": {"url": "https://vidgen.x.ai/result.mp4", "duration": 6} + }).as_object().unwrap())); + let snapshot = service + .snapshot_for_route(Some("openai"), "/v1/videos/local-video") + .unwrap(); + assert!(!snapshot.is_active_for_refresh()); + let record = snapshot.to_upsert_record(); + assert_eq!( + record.status, + aether_data_contracts::repository::video_tasks::VideoTaskStatus::Completed + ); + assert_eq!( + record.video_url.as_deref(), + Some("https://vidgen.x.ai/result.mp4") + ); + assert_eq!(record.duration_seconds, Some(6)); + let response = snapshot.read_response(); + assert_eq!(response.body_json["status"], "completed"); + assert_eq!(response.body_json["progress"], 100); + assert_eq!( + response.body_json["video_url"], + "https://vidgen.x.ai/result.mp4" + ); + let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else { + panic!("OpenAI video expected") + }; + let Some(LocalVideoTaskContentAction::StreamPlan(download)) = + seed.build_content_stream_action(None, "download") + else { + panic!("download expected") + }; + assert_eq!(download.url, "https://vidgen.x.ai/result.mp4"); + assert!( + download.headers.is_empty(), + "provider credentials must not be sent to the media CDN" + ); + } + } + + #[test] + fn xai_video_errors_are_terminal_even_without_a_status() { + use crate::{LocalVideoTaskSnapshot, VideoTaskTruthSourceMode}; + let mut plan = + build_internal_finalize_video_plan("xai-create", "openai:video", None).unwrap(); + plan.url = "https://cli-chat-proxy.grok.com/v1/videos/generations".into(); + for body in [ + json!({"code": "content_policy_violation", "error": "Rejected"}), + json!({"error": {"code": "content_policy_violation", "message": "Rejected"}}), + json!({"status": "failed", "error": "Rejected"}), + ] { + let mut snapshot = VideoTaskTruthSourceMode::RustAuthoritative + .prepare_sync_success( + "openai_video_create_sync_finalize", + json!({"request_id": "xai-request"}).as_object().unwrap(), + &Default::default(), + &plan, + ) + .unwrap() + .to_snapshot(); + snapshot.apply_provider_body(body.as_object().unwrap()); + assert!(!snapshot.is_active_for_refresh()); + assert_eq!(snapshot.read_response().body_json["status"], "failed"); + let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else { + panic!("OpenAI video expected") + }; + assert!(seed.error_message.is_none()); + } + } + + #[test] + fn openai_video_id_takes_precedence_over_xai_alias() { + assert_eq!( + super::openai_video_provider_task_id( + json!({"id": "openai-id", "request_id": "trace-id"}) + .as_object() + .unwrap() + ), + Some("openai-id") + ); + assert_eq!( + super::openai_video_provider_task_id( + json!({"id": " ", "request_id": "xai-id"}) + .as_object() + .unwrap() + ), + Some("xai-id") + ); + assert_eq!( + super::openai_video_provider_task_id(json!({"request_id": " "}).as_object().unwrap()), + None + ); + } + #[test] fn builds_local_sync_finalize_read_response_for_supported_video_finalize_kinds() { let delete_response = build_local_sync_finalize_read_response( diff --git a/crates/aether-video-tasks-core/src/transport_domain.rs b/crates/aether-video-tasks-core/src/transport_domain.rs index 9503e0087..9f3b267ae 100644 --- a/crates/aether-video-tasks-core/src/transport_domain.rs +++ b/crates/aether-video-tasks-core/src/transport_domain.rs @@ -71,8 +71,16 @@ impl LocalVideoTaskPersistence { .unwrap_or_else(|| plan.request_id.clone()), username: context_text(report_context, "username"), api_key_name: context_text(report_context, "api_key_name"), - client_api_format: context_text(report_context, "client_api_format") - .unwrap_or_else(|| plan.client_api_format.clone()), + client_api_format: if report_context + .get("video_client_protocol") + .and_then(Value::as_str) + == Some("xai") + { + "xai:video".to_string() + } else { + context_text(report_context, "client_api_format") + .unwrap_or_else(|| plan.client_api_format.clone()) + }, provider_api_format: context_text(report_context, "provider_api_format") .unwrap_or_else(|| plan.provider_api_format.clone()), original_request_body: report_context diff --git a/crates/aether-video-tasks-core/src/types.rs b/crates/aether-video-tasks-core/src/types.rs index 81d8a29f1..69e6b255a 100644 --- a/crates/aether-video-tasks-core/src/types.rs +++ b/crates/aether-video-tasks-core/src/types.rs @@ -201,6 +201,13 @@ pub struct LocalVideoTaskPersistence { #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct OpenAiVideoTaskSeed { + /// Preserve existing database identity; older snapshots derive it from the local task ID. + #[serde(default)] + pub local_short_id: Option, + #[serde(default)] + pub native_response: Option, + #[serde(default)] + pub xai_provider: bool, pub local_task_id: String, pub upstream_task_id: String, pub created_at_unix_ms: u64, diff --git a/docs/operations/gateway-ci-timeout-reduction-plan.md b/docs/operations/gateway-ci-timeout-reduction-plan.md new file mode 100644 index 000000000..18be03ccf --- /dev/null +++ b/docs/operations/gateway-ci-timeout-reduction-plan.md @@ -0,0 +1,262 @@ +# Test (Gateway) CI 耗时减负方案 + +- 文档日期:2026-09-23(Asia/Shanghai) +- 目标 job:`.github/workflows/rust-ci.yml` 中的 `test_gateway`(显示名 `Test (Gateway)`) +- 优化目标:**缩短 CI 耗时**,不减少安全/回归断言 +- 关联审计:`docs/operations/system-slimming-audit-2026-09-08.md` +- **第一批(A1+A2+A4)已实施**,待 CI 前后对照确认分钟数。 + +本文只新增方案文档,不修改业务代码、测试或工作流。文中收益均为基于历史日志与源码结构的**估计值**,每批改动落地后须用同一提交做前后对照实测。 + +--- + +## 一、结论 + +`Test (Gateway)` 的耗时大头是**巨型单一 lib 测试目标的编译**(约 60%),其次是 **5300+ 测试的执行**(约 40%)。单纯减少“测试项数”对总分钟数帮助有限;应优先砍掉**重复编译指纹浪费**和**少数重型夹具**。 + +- 单 job 实测约 **11~17 分钟**,是普通 Rust CI 的关键路径。 +- lib 测试约 5300 项、bins 约 78 项,全部 `#[cfg(test)]` 代码编进**同一个** rustc 调用,单进程峰值约 **7GB**,曾出现 OOM。 +- 执行段历史波动 **214s~390s**;其中 4 个慢用例合计约 **54 秒**(约占快样本执行时间的 25%)。 +- 仓库已有瘦身审计的立场不变:**减编译与运行重量,不先减安全断言**。 + +### 不做清单(铁律) + +- 不删除 OAuth、配额、安全头、并发门禁、备份密码学、PII/断开结算等断言。 +- 不下调 `PBKDF2_ITERATIONS`(`crates/aether-crypto/src/python_fernet.rs`)。 +- 不把关键测试标 `ignore`,不跨测试共享可变 `AppState`。 +- 暂不做 nextest 多 runner 分片(每个 runner 各自重编 7GB 巨型目标会净亏;须先复用同一次构建产物再评估)。 + +--- + +## 二、现状基线 + +### 2.1 Job 做什么 + +`.github/workflows/rust-ci.yml` 中 `test_gateway` 关键步骤: + +| 步骤 | 行号 | 命令 / 配置 | +| --- | --- | --- | +| Rust cache | 247-251 | `shared-key: rust-ci-${{ runner.os }}`(与 Data/Rest/Integration 等共用) | +| Setup mold | 256-257 | **仅此 job** 安装 mold | +| Test lib | 265-271 | `cargo nextest run -p aether-gateway --lib` | +| Test bins | 273-279 | `cargo nextest run -p aether-gateway --bins` | + +两个测试 step 的环境变量:`RUSTC_WRAPPER=sccache`、`RUST_MIN_STACK=16777216`、`RUSTFLAGS: "-C link-arg=-fuse-ld=mold"`(**RUSTFLAGS 只在 step 级,未提到 job 级**)。 + +Workflow 级(72-76 行):`CARGO_INCREMENTAL=0`、`CARGO_PROFILE_DEV_DEBUG=0`、`CARGO_PROFILE_TEST_DEBUG=0`。 + +**未发现** `.config/nextest.toml` 或任何 nextest 配置文件;CI 使用 nextest 默认并行度。 + +**未执行**:`apps/aether-gateway/tests/admin_unsigned_identity_headers.rs`(1 个安全集成用例)——两条 nextest 命令只覆盖 `--lib` 与 `--bins`。 + +### 2.2 耗时拆分(历史日志) + +来源:`docs/operations/system-slimming-audit-2026-09-08.md`。 + +| 运行 | Gateway lib 编译 | Gateway lib 执行 | Test lib 步骤合计 | +| --- | --- | --- | --- | +| `34174603131` | 4:57(297s) | **213.807s**(5,139 项) | 514s | +| `34153166516` | 6:22(382s) | **390.039s** | — | + +同轮 bins:编译约 2:27~3:05,执行仅 **0.27~0.33s**(78 项)。 + +- 编译 : 执行 ≈ **57:43**(快样本)~ **49:51**(慢样本)。 +- 合并 `--lib --bins` 为一条命令**不能**省掉普通库与 `cfg(test)` 测试库的两种构建。 + +### 2.3 测试规模(静态统计) + +范围:`apps/aether-gateway`。 + +| 分区 | 约 `#[test]` / `#[tokio::test]` 数量 | +| --- | --- | +| `src/handlers` 内联 | 1224 | +| `src/tests/` 测试树(含 control 727、frontdoor 217、architecture 208 等) | ~1358 | +| `src/execution_runtime` 内联 | ~591 | +| `src/control` 内联 | 459 | +| `src/ai_serving` 内联 | 302 | +| `src/main.rs` + `src/bin/*`(bins 目标) | ~77 | +| 其余分散内联 | ~1500+ | +| **合计** | **约 5380~5400** | + +与 CI 实测(lib 5139 + bins 78)同量级;差异来自 cfg 门控与统计口径。 + +架构守卫:已迁至 `apps/aether-gateway/tests/architecture/`(入口 `tests/architecture_guard.rs`),共 13 文件、约 **14,925 行、208 个 `#[test]`**,断言全部是 `fs::read_to_string` + 字符串/`Cargo.toml` 规则检查,**不依赖 gateway 私有类型**。CI 由 `Test integration targets` 步骤执行(`cargo nextest run -p aether-gateway --tests`)。 + +### 2.4 编译为何是单点瓶颈 + +- `apps/aether-gateway/src/lib.rs:235-236` 把整棵 `tests/` 树挂进**同一个** lib test binary。 +- `src/tests/` 约 130 个文件、17 万行;连同业务代码,单 rustc 调用编译约 53 万行。 +- 历史记录(`openai-responses-websocket-plan.md`):单进程 rustc 峰值约 **7GB**,极端时 RSS 8.9GB 并触发 `oom_kill`;`-j` / `--test-threads` **不影响**该单进程峰值。 +- 对照实验:临时裁掉无关测试模块后,编译可降至约 3 分钟内且零 OOM——证明“编译面”而非“并行度”是主因。 +- 本地 `RUSTC_WRAPPER=sccache` 命中率曾低至约 1.5%;CI 样本 sccache 命中约 95% 仍要数分钟——剩余问题在 **crate-type 不可缓存调用与 RUSTFLAGS 指纹不一致**,不是“再加一层缓存”。 + +### 2.5 执行段慢点(已实测或源码可证) + +| 优先级 | 模式 | 位置 | 估计可省 | +| --- | --- | --- | --- | +| 1 | 池调度超大夹具(1700/2048 账号,每 key 真实 `seal_provider_catalog_key_api_key`) | `src/dispatch/pool_scheduler.rs`(`large_pool_fixture` 约 :5036;慢用例约 :4048、:3967、:4641) | **30~50s**(Top4 中 3 项) | +| 2 | 备份 v1 兼容 + 17 个 historical 密钥各走 10 万次 PBKDF2 | `src/backup/executor.rs:1222` 一带;`python_fernet.rs` | **10~15s**(与上表 14.3s 重叠) | +| 3 | 每测试独立进程重复付 DEVELOPMENT_ENCRYPTION_KEY 的 PBKDF2(nextest 无法跨用例共享缓存) | 全包大量 `encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, …)` | **10~25s**(非密码学夹具改直接密钥后) | +| 4 | 真实 `start_server` 起服极多(tests 树约 1794 次)+ `AppState::new()`(tests 树约 816 次) | `src/tests/mod.rs:33-47` 等 | **15~40s**(改 oneshot 仅限纯路由断言) | +| 5 | 固定 `sleep(100ms)` 扇出、个别 mock 内 5s/30s sleep | `tests/ai_execute/**`、`stream_pump.rs` 等 | **5~20s** | +| 6 | 架构守卫重复全树扫描(迁出后不再计入 lib) | `tests/architecture/*` | 执行 **~3-10s** + 编译见阶段 A | + +**执行优化硬顶**:即使执行砍掉一半,关键路径大约只省 **40~80 秒**;样本间 214s vs 390s 的抖动本身就 ±176s。因此**编译侧(阶段 A/C)才是分钟级收益来源**。 + +--- + +## 三、分阶段方案 + +### 阶段 A:低风险、优先实施(预计单 job 省 1~3 分钟) + +#### A1. 架构守卫迁出独立测试目标【已实施】 + +- **改动**:`src/tests/architecture/**`(208 项 / 约 1.5 万行)已迁至 `apps/aether-gateway/tests/architecture/`,入口 `apps/aether-gateway/tests/architecture_guard.rs`;已从 `src/tests/mod.rs` 移除 `mod architecture;`。 +- **为什么省**:从巨型 lib `cfg(test)` 编译单元削掉约 15k 行,压低 rustc 7GB 峰值,降低 OOM 与编译墙钟;对照实验表明裁测试面可明显缩短编译。 +- **收益估计**:编译 **-20~60s**,执行 **-数秒**;OOM 风险显著下降。 +- **风险**:低。断言逻辑一字不改,仅更换宿主;helper 已改为 `pub(crate)`。 +- **CI**:`test_gateway` 新增 `Test integration targets`:`cargo nextest run -p aether-gateway --tests`(同时覆盖 A4 的 `admin_unsigned_identity_headers`)。 +- **验收**:架构 208 项 + 安全 1 项在新 target 全绿(本地实测 209 passed);lib 测试数减少约 208;lib 编译时间与 rustc 峰值待 CI 对照。 + +#### A2. 统一构建指纹(RUSTFLAGS + 工具链提到 job 级)【已实施】 + +- **改动**: + 1. `test_gateway` 的 mold `RUSTFLAGS`、`RUST_MIN_STACK`、`RUSTC_WRAPPER` 已上移至 **job 级 `env`**,lib/bins/integration 三步共用同一指纹。 + 2. `rust-ci.yml` 中**全部** `Install Rust toolchain` 步骤已钉住 `toolchain: 1.95.0`(与 `rust-toolchain.toml`、fmt/clippy 一致),消除浮动 stable 漂移。 + 3. `test_gateway` 的 `shared-key` 改为 `rust-ci-gateway-test-${{ runner.os }}`,避免 mold 指纹与无 mold 的 job 互相污染共享缓存(一次冷缓存成本)。 +- **为什么省**:原先 step 级 RUSTFLAGS + 浮动 stable + 十余个 job 共用同一 cache key,造成“看似命中、实际重编”,放大 4:57 vs 6:22 的波动;stable 漂移还会触发偶发全量重编。 +- **收益估计**:热缓存编译 **-1~2 分钟波动收窄**;避免偶发 **-5~10 分钟** 尖峰;gateway 独立 cache key 后与 Rest/Data 不再争抢/覆盖。 +- **风险**:低。mold 本就在用;独立 key 首次为冷缓存。 +- **验收**:连续多次 run 的 lib 编译时间方差下降;全 workflow 无未钉 toolchain 步骤。 + +#### A3.(可选,紧随 A2)按用途拆分 rust-cache key + +- **改动**:Gateway **test** 与 lint/check 类 job 不再共用同一 `shared-key`;或评估 `cache-workspace-crates: true`。 +- **收益估计**:编译 **-30~90s**(估计)。 +- **风险**:中。缓存体积上升;勿为每个细碎目标无限新建 key,勿直接缓存完整巨型 `target/`。 + +#### A4. 补跑漏掉的安全集成测试【已实施】 + +- **改动**:`test_gateway` 增加 `Test integration targets`:`cargo nextest run -p aether-gateway --tests`,覆盖 `apps/aether-gateway/tests/` 下全部目标(`architecture_guard` + `admin_unsigned_identity_headers`)。 +- **收益**:时长 **+10~30s**,换取已确认的覆盖缺口(安全优先,与 A1 同 PR)。 +- **风险**:低。该测试走公开 API。本地已实测 1/1 passed。 + +### 阶段 B:只动测试代码(预计执行段省 40~70 秒) + +#### B1. 池调度大夹具轻量化 + +- **改动**:`large_pool_fixture` 对非“规模/扫描预算语义”的用例,改为轻量 repository/credential fixture 或预构造行;**保留**: + - 扫描预算、跳过计数、分页/游标语义断言; + - **至少 1~2 条** 1700/2048 规模边界用例(可保留真实加密封装)。 +- **收益估计**:执行 **-30~50s**。 +- **风险**:中。不得把池规模缩到失去原回归条件;不得删断言。 + +#### B2. PBKDF2 夹具去重 + +- **改动**:对**不测密码学语义**的夹具,改用直接 32 字节 base64 密钥(`decode_direct_fernet_key` 路径)或预制密文,绕开每进程 10 万次迭代。 +- **必须保留**:备份 historical 密钥兼容、生产强度派生、显式 PBKDF2 行为测试(如 `python_fernet` 相关用例)。 +- **收益估计**:执行 **-10~25s**。 +- **风险**:中。**绝不降低迭代次数**。 + +#### B3.(可选,后置)起服与 helper 瘦身 + +- 纯鉴权/路由断言:`start_server` → `tower` oneshot(已有 `send_request`);**涉及超时、连接、完整中间件链的用例保留真实 server**。 +- 合并 8+ 份同构 `run_*_test` 大栈 helper 为单一 helper;OAuth/keys/quota 同构用例可数据驱动,**表内逐行保留断言**。 +- 固定 `sleep` 改为 channel/notify 事件驱动,**不删除等待本身**(防 flaky)。 +- **收益估计**:执行 **-15~40s**;改动面大于 B1/B2,单独 PR。 + +#### B4. nextest 稳定性配置(非直接减分钟) + +- 新增 `.config/nextest.toml`:显式 `test-threads`(对齐 runner vCPU,避免规格漂移)、`slow-timeout`(防单测卡死拖满 job)。 +- **不要**为提速下调 `RUST_MIN_STACK`(16MB 为深栈/管理面用例正确性所需)。 + +### 阶段 C:流水线级(主要缩短累计 runner,间接稳定 Gateway 缓存) + +| # | 改动 | 收益 | 风险 | +| --- | --- | --- | --- | +| C1 | **解除 Tunnel → Gateway dev-dependency**(`apps/aether-tunnel/Cargo.toml:48-49`;端到端场景迁到 integration package) | Rest/Clippy 不再重复编 Gateway → 流水线累计约 **-8~10 分钟 runner**;缓解共享缓存污染 | 中:须迁移 `src/tunnel/mod.rs` 中依赖 `AppState`/`build_router_with_state` 的用例并比对测试清单 | +| C2 | **路径过滤分层**(`rust-ci.yml` push/PR paths):README、安装脚本、Compose、部分 `tests/*.sh` 不再触发全量 Rust 编译 | 非 Rust 变更 **整段跳过 Test (Gateway)**;须保留稳定 gate 防止 required check 永久 pending;补 `rust-toolchain.toml`、`.cargo/**` 触发项 | 中 | +| C3 | Nightly 空 doctest、重复 adapter/feature job 治理(见既有瘦身审计) | 流水线约 **-4 分钟**,**不在**本 job 关键路径 | 低-中 | + +--- + +## 四、实施顺序与验收 + +| 批次 | 内容 | 预期(估计) | 验收标准 | +| --- | --- | --- | --- | +| **第一批 PR** | A1 架构守卫迁出 + A2 指纹统一 + A4 补安全集成测试【已实施】 | 单 job **-1~2 分钟** + 覆盖补齐 | 架构 208 + 安全 1 本地 209 全绿;CI 对照编译时间下降 | +| **第二批 PR** | B1 池调度夹具 + B2 PBKDF2 去重 | 执行 **-40~70s** | 断言集合不减;4 个历史慢用例计时明显下降;密码学用例仍为生产强度 | +| **第三批 PR** | A3 拆缓存 key + C1 Tunnel 解耦 + C2 路径过滤 | 缓存更稳 + 流水线累计大幅下降 | Rest 依赖闭包不再含 Gateway;无关路径 PR 不再拉起全量编译;Gateway 编译方差下降 | +| **可选** | B3 起服/helper、B4 nextest.toml | 再 **-15~40s** + 稳定性 | 无新增 flaky;慢测试超时有告警 | + +### 对照方法(每批必做) + +1. 固定同一 SHA、runner 规格、工具链与 feature 集合。 +2. 记录:`Test lib` / `Test bins` 的**编译墙钟**、**执行墙钟**、测试总数、失败数。 +3. 记录 rustc 峰值内存(如有)与 sccache 命中率。 +4. 区分**冷缓存 / 热缓存**,区分 **job 耗时 / 流水线总耗时**。 +5. **测试总数只允许**因 A1 迁移而在 lib 与新 target 之间搬家;禁止静默减少断言。 + +### 成功指标 + +- **首要**:普通 PR 上 `Test (Gateway)` 墙钟时间下降且结果稳定(方差收窄)。 +- **次要**:全流水线累计 runner 分钟下降(阶段 C)。 +- **禁止**把“删除测试数量”当作成功标准。 + +--- + +## 五、可重复的只读核查命令 + +```sh +# 测试属性数量(lib 树 / bins) +rg -c '#\[(tokio::)?test\]' apps/aether-gateway/src -g '*.rs' | awk -F: '{s+=$2} END {print s}' +rg -c '#\[(tokio::)?test\]' apps/aether-gateway/src/main.rs apps/aether-gateway/src/bin -g '*.rs' + +# 架构守卫规模(迁移后) +rg -c '#\[(tokio::)?test\]' apps/aether-gateway/tests/architecture -g '*.rs' +wc -l apps/aether-gateway/tests/architecture/*.rs + +# CI 集成测试步骤与指纹 +rg -n 'Test integration targets|rust-ci-gateway-test|toolchain: 1.95.0|RUSTFLAGS' .github/workflows/rust-ci.yml + +# 是否存在 nextest 配置 +ls .config/nextest.toml 2>/dev/null || echo 'no nextest.toml' + +# Tunnel 反向依赖 +rg -n 'aether-gateway' apps/aether-tunnel/Cargo.toml + +# CI 中 mold/RUSTFLAGS/缓存 key +rg -n 'mold|RUSTFLAGS|shared-key|nextest run -p aether-gateway' .github/workflows/rust-ci.yml +``` + +历史耗时与慢用例计时以 `docs/operations/system-slimming-audit-2026-09-08.md` 及对应 GitHub Actions 运行 ID 为准;临时 API JSON 不入库。 + +--- + +## 六、风险与回滚 + +| 风险 | 缓解 | +| --- | --- | +| 迁出架构守卫后漏挂模块 | 迁移前后对比 208 项清单;CI 显式跑新 target | +| RUSTFLAGS 上移后某 job 链接失败 | 先在 `test_gateway` job 级验证,再推广到其它 job;保留 step 级回滚 diff | +| 夹具轻量化导致规模回归失效 | 强制保留 ≥1 条大规模边界用例;PR 中 diff 审查断言 | +| 路径过滤导致 required check 永久 pending | 增加始终运行的轻量 gate job | +| Tunnel 解耦丢失端到端场景 | 迁移前后测试清单比对;场景迁入 integration job | + +回滚单位按 **PR 批次**:每批独立可 revert,不把 A/B/C 混在同一提交。 + +--- + +## 七、附录:明确不能减的测试(摘录) + +| 类别 | 主要位置(约) | 说明 | +| --- | --- | --- | +| OAuth 导入/刷新/吊销 | `src/tests/control/admin/oauth.rs` 等 | 账户安全核心 | +| 配额与失效删 key | `src/tests/control/admin/endpoints/quota.rs` 等 | 计费与授权正确性 | +| 安全头 / 管理面访问 | `security.rs`、`health_access.rs`、`operational_auth` | 未签名身份头等 | +| 并发门禁 | `src/tests/concurrency.rs` | 过载拒绝与独立准入 | +| 备份历史兼容与密码学 | `src/backup/executor.rs` | 保留生产强度 PBKDF2 | +| 池调度扫描预算语义 | `src/dispatch/pool_scheduler.rs` | 可改夹具实现,不可删语义断言 | +| AI 断开结算 / PII | `src/tests/ai_execute/...` | 计费完整性与隐私 | + +> 与 `system-slimming-audit-2026-09-08.md` 一致:首要结果是 **PR 更快得到正确反馈**,不是测试条数变少。 diff --git a/docs/operations/xai-provider.md b/docs/operations/xai-provider.md new file mode 100644 index 000000000..71aee7f72 --- /dev/null +++ b/docs/operations/xai-provider.md @@ -0,0 +1,147 @@ +# xAI provider behavior + +The following rules preserve the provider-specific behavior of the `xai` provider +across Aether's request and transport layers. + +## Responses and tools + +- HTTP requests drop `previous_response_id`. Clients must supply conversation + history; this provider does not add an HTTP response-ID history store. +- `metadata.user_id` is removed. Claude clients copy it onto converted Responses + bodies and xAI rejects the field. +- Preserve requested `reasoning.encrypted_content`. On a native Responses-to-Responses + hop, keep provider-owned input items instead of rebuilding them through the canonical + format. xAI encrypted reasoning may have IDs that do not use OpenAI's `rs` prefix. + Aether's Gemini signature carriers remain excluded from xAI replay. +- The replay policy is selected from the configured provider type. A model called + `grok-*` on another provider does not opt into that policy. WebSocket continuation + metadata retains the selected policy across reconnects. +- A regular client function called `web_search` remains a function. Claude hosted + search choices are resolved against the original typed tool declaration, including + declarations with a different name. +- When only `image_generation` is allowed, keep only that tool and retain the requested + `auto` or `required` mode. For mixed allowed-tool lists, remove the image choice while + preserving the other allowed entries, as required by xAI's tool-choice schema. +- Reasoning effort is stripped for models that do not accept it. +- OpenAI-style image reference aliases in a request body are rewritten to xAI's + shape without touching chat message parts. + +## Routing and credentials + +OAuth requests default to `https://cli-chat-proxy.grok.com/v1`; API-key or +`using_api=true` requests default to `https://api.x.ai/v1`. Explicit custom gateways +are preserved. Compact remains on the official endpoint. CLI identity headers are +applied where the selected upstream requires them. + +Account binding uses the xAI device code flow: the gateway requests a device code, +the operator authorizes it out of band, and the gateway polls for the token set. +There is no local callback listener, so headless deployments can bind accounts. +Refresh tokens can also be imported individually or in batches, and are rotated +on refresh. + +Quota refresh reads `/user` and `/billing?format=credits` and stores a structured +usage snapshot. A prepaid balance keeps an account selectable after the weekly +allowance is exhausted. API-key accounts skip the subscription billing surface. + +## Images and videos + +OAuth media requests default to `https://cli-chat-proxy.grok.com/v1`; API-key or +`using_api=true` requests default to `https://api.x.ai/v1`. Explicit custom gateways +are preserved. Compact remains on the official endpoint. CLI identity headers are +applied to media requests and restored when a persisted video task's polling transport +is reconstructed. + +Aether's OpenAI-compatible task parser accepts xAI's `request_id` creation field, +status aliases such as `pending` and `done`, nested `video.url` and `video.duration`, +and failure payloads containing `code` / `error` without a status. Existing OpenAI +`id` takes precedence. The client receives Aether's local task ID; polling uses the +upstream task ID and selected credential. Completed video downloads use the returned +media URL without forwarding provider authentication headers to the media host. + +### Public video protocols + +The xAI provider supports two video surfaces: + +| Operation | xAI native | OpenAI compatible | +| --- | --- | --- | +| Create | `POST /v1/videos/generations` | `POST /openai/v1/videos` | +| Edit / extend | `POST /v1/videos/edits`, `POST /v1/videos/extensions` | — | +| Retrieve | `GET /v1/videos/{request_id}` | `GET /openai/v1/videos/{id}` | +| Download | use the returned `video.url` | `GET /openai/v1/videos/{id}/content` | + +For xAI, `POST /v1/videos` is a native creation alias. Other providers retain +Aether's existing OpenAI-compatible `/v1/videos` behavior. xAI callers using +OpenAI `seconds` / `size` parameters must use `/openai/v1/videos`. The adapter +maps these to numeric `duration`, `aspect_ratio`, and `resolution`; it also adapts +image references. This implementation defaults to 4 seconds, portrait, and 720p, +clamps `duration` to 1-15, and validates inputs. Explicit native requests retain +native parameters and additional provider fields. + +Default xAI creation targets `/videos/generations` on the selected upstream host. +Explicit custom endpoint paths still take precedence. Native generation, editing, +and extension paths only select xAI provider candidates. + +Native creation returns `request_id`; native retrieval preserves `done`, nested +`video.url`, and provider fields such as `respect_moderation`. The identifier is +an opaque Aether task ID so queries remain scoped to the owning user and pinned +to the original upstream task and credential. The explicit `/openai/v1/videos` +surface projects `id`, `completed`, and `video_url`. + +The task row records the native client protocol as `xai:video`, while its provider +transport remains `openai:video`. This survives restart without storing request +bodies or credentials. Raw native responses are cached only in memory; after +reconstruction the gateway refreshes from the original provider to recover its +response fields, including for completed tasks. If refreshing is unavailable, +the stored task still provides the native status and media URL projection. + +OpenAI/xAI task persistence supplies a stable 16-character `short_id`, as required +by the PostgreSQL schema. Existing rows retain their original short ID across +reconstruction, including legacy embedded snapshots. This internal identifier is +separate from the opaque local task ID returned to clients; no schema change or +historical row rewrite is needed. + +Task retrieval and content downloads are admitted by the production GET execution +gate. Reconstructed tasks resolve proxy nodes, system proxy defaults, tunnel affinity, +and transport profiles through the same deployment resolver used for creation; +configured proxy routes must not silently turn into direct requests after restart. + +### Runtime configuration + +Standalone Rust deployments must set +`AETHER_GATEWAY_VIDEO_TASK_TRUTH_SOURCE_MODE=rust-authoritative` and restart the +gateway to enable video task retrieval, polling, and content downloads. The CLI's +legacy default is `python-sync-report`: creation can return a task ID in that mode, +but the local task read/refresh paths are disabled and may return HTTP 503. + +When the gateway also serves the frontend, `/openai/v1/videos` and its subpaths +must bypass the static SPA handler and be mounted as API routes. Otherwise a +successful-looking HTTP 200 response to a video query may contain `text/html` +instead of the task's JSON response. The lifecycle regression includes the static +frontend to cover this production configuration. + +## Regression coverage + +The format tests cover client and hosted search choices, image-only and mixed tool +restrictions, encrypted reasoning replay, image reference rewriting, and unchanged +OpenAI replay restrictions. Transport tests cover OAuth/API-key/custom routing and +media identity headers. Video-task tests exercise creation, polling, terminal +projection, persistence fields, content-download planning, and status-less errors +using local fixtures. They do not make paid generation requests. + +The HTTP regression exercises all native creation paths and the compatibility +prefix through the public router and candidate planner, then checks polling, +cross-user denial, persistence, retrieval from a fresh gateway instance, and downloads +through both prefixes without leaking authorization to the media host. It uses the +real HTTP executor and a managed proxy node backed by a local test server, with no +execution-runtime override. The background poller also has a real HTTP proxy-node +regression, so production method guards and transport reconstruction are exercised. +CI also runs the same HTTP lifecycle with the PostgreSQL repository and the +production column constraints/indexes in an isolated temporary table. This catches +persistence failures that the in-memory repository cannot expose. The test uses +local `initdb`, `postgres`, and `pg_ctl` (already provided by the gateway CI job), +or an explicit `AETHER_TEST_DATABASE_URL` pointing to an isolated test database. + +```sh +cargo test -p aether-ai-formats -p aether-provider-transport -p aether-video-tasks-core --lib +cargo test -p aether-gateway --lib xai +``` diff --git a/frontend/src/api/__tests__/admin-analytics-cache.spec.ts b/frontend/src/api/__tests__/admin-analytics-cache.spec.ts index 776812243..a4a14c82b 100644 --- a/frontend/src/api/__tests__/admin-analytics-cache.spec.ts +++ b/frontend/src/api/__tests__/admin-analytics-cache.spec.ts @@ -82,4 +82,35 @@ describe('adminApi analytics cache options', () => { }) expect(getMock).toHaveBeenNthCalledWith(4, '/api/admin/stats/errors/distribution', { params }) }) + + it('requests the user group leaderboard with scoped cache parameters', async () => { + const groupParams = { + ...params, + metric: 'cost' as const, + offset: 10, + limit: 10, + include_inactive: true, + } + getMock.mockResolvedValueOnce({ + data: { items: [], total: 0, metric: 'cost', attribution: 'current_membership' }, + }) + + await expect(adminApi.getLeaderboardUserGroups(groupParams)).resolves.toMatchObject({ + attribution: 'current_membership', + }) + + expect(buildCacheKeyMock).toHaveBeenCalledWith( + 'admin:stats:leaderboard:user-groups', + groupParams + ) + expect(cachedRequestMock).toHaveBeenCalledWith( + 'admin:stats:leaderboard:user-groups', + expect.any(Function), + 20 * 1000 + ) + expect(getMock).toHaveBeenCalledWith('/api/admin/stats/leaderboard/user-groups', { + params: groupParams, + }) + }) + }) diff --git a/frontend/src/api/__tests__/providers-pool-cache.spec.ts b/frontend/src/api/__tests__/providers-pool-cache.spec.ts new file mode 100644 index 000000000..745444bc5 --- /dev/null +++ b/frontend/src/api/__tests__/providers-pool-cache.spec.ts @@ -0,0 +1,76 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { getMock, patchMock } = vi.hoisted(() => ({ getMock: vi.fn(), patchMock: vi.fn() })) + +vi.mock('@/api/client', () => ({ default: { get: getMock, patch: patchMock } })) + +import { getPoolOverview, listPoolKeys, listPoolScores } from '@/api/endpoints/pool' +import { getProvider, updateProvider } from '@/api/endpoints/providers' +import { cache } from '@/utils/cache' + +const options = { cacheTtlMs: 30_000 } +const provider = { id: 'codex', pool_advanced: { reserve_minimum_quota: false } } + +function deferred() { + let resolve!: (value: T) => void + const promise = new Promise((resolvePromise) => { resolve = resolvePromise }) + return { promise, resolve } +} + +beforeEach(() => { + cache.clear() + getMock.mockReset() + patchMock.mockReset() + patchMock.mockResolvedValue({ data: provider }) +}) + +describe('provider pool settings cache invalidation', () => { + it('refreshes every cached key page, score page and overview after changing pool settings', async () => { + getMock.mockResolvedValue({ data: { version: 'old' } }) + await getPoolOverview(options) + await listPoolKeys('codex', {}, options) + await listPoolKeys('codex', { page: 2, status: 'quota_exhausted' }, options) + await listPoolScores('codex', {}, options) + await listPoolKeys('codex-other', {}, options) + await updateProvider('codex', { pool_advanced: { reserve_minimum_quota: false } }) + getMock.mockResolvedValue({ data: { version: 'new' } }) + + await expect(getPoolOverview(options)).resolves.toEqual({ version: 'new' }) + await expect(listPoolKeys('codex', {}, options)).resolves.toEqual({ version: 'new' }) + await expect(listPoolKeys('codex', { page: 2, status: 'quota_exhausted' }, options)) + .resolves.toEqual({ version: 'new' }) + await expect(listPoolScores('codex', {}, options)).resolves.toEqual({ version: 'new' }) + await expect(listPoolKeys('codex-other', {}, options)).resolves.toEqual({ version: 'old' }) + }) + + it('does not reuse or cache a key request started before the settings were saved', async () => { + const oldResponse = deferred<{ data: { version: string } }>() + const newResponse = deferred<{ data: { version: string } }>() + getMock.mockReturnValueOnce(oldResponse.promise).mockReturnValueOnce(newResponse.promise) + const oldRequest = listPoolKeys('codex', { page: 2 }, options) + await updateProvider('codex', { pool_advanced: { reserve_minimum_quota: false } }) + const newRequest = listPoolKeys('codex', { page: 2 }, options) + expect(getMock).toHaveBeenCalledTimes(2) + + oldResponse.resolve({ data: { version: 'old' } }) + await oldRequest + const deduped = listPoolKeys('codex', { page: 2 }, options) + expect(getMock).toHaveBeenCalledTimes(2) + newResponse.resolve({ data: { version: 'new' } }) + await expect(newRequest).resolves.toEqual({ version: 'new' }) + await expect(deduped).resolves.toEqual({ version: 'new' }) + await expect(listPoolKeys('codex', { page: 2 }, options)).resolves.toEqual({ version: 'new' }) + }) + + it('does not reuse a provider detail request started before a successful save', async () => { + const oldResponse = deferred<{ data: typeof provider }>() + getMock.mockReturnValueOnce(oldResponse.promise).mockResolvedValueOnce({ data: provider }) + const oldRequest = getProvider('codex') + await updateProvider('codex', { pool_advanced: { reserve_minimum_quota: false } }) + const newRequest = getProvider('codex') + expect(getMock).toHaveBeenCalledTimes(2) + oldResponse.resolve({ data: { ...provider, pool_advanced: { reserve_minimum_quota: true } } }) + await oldRequest + await expect(newRequest).resolves.toMatchObject(provider) + }) +}) diff --git a/frontend/src/api/__tests__/users.spec.ts b/frontend/src/api/__tests__/users.spec.ts index 3f2f8d486..f4dc72e52 100644 --- a/frontend/src/api/__tests__/users.spec.ts +++ b/frontend/src/api/__tests__/users.spec.ts @@ -17,7 +17,7 @@ vi.mock('@/utils/cache', () => ({ cachedRequest: cachedRequestMock, })) -import { usersApi } from '@/api/users' +import { buildUserBatchBalanceAdjustmentPayload, usersApi } from '@/api/users' describe('usersApi admin list query', () => { beforeEach(() => { @@ -104,3 +104,25 @@ describe('usersApi admin list query', () => { expect(getMock).toHaveBeenCalledWith('/api/admin/users/target-user/api-keys') }) }) + +describe('user batch wallet balance payload', () => { + it('builds an addition payload from a valid positive amount', () => { + expect(buildUserBatchBalanceAdjustmentPayload('add', '12.5')).toEqual({ + operation: 'add', + amount: 12.5, + }) + }) + + it('builds a deduction payload from a valid positive amount', () => { + expect(buildUserBatchBalanceAdjustmentPayload('deduct', 4)).toEqual({ + operation: 'deduct', + amount: 4, + }) + }) + + it('rejects zero, blank, and non-finite amounts', () => { + expect(buildUserBatchBalanceAdjustmentPayload('add', '0')).toBeNull() + expect(buildUserBatchBalanceAdjustmentPayload('deduct', '')).toBeNull() + expect(buildUserBatchBalanceAdjustmentPayload('deduct', '1e999')).toBeNull() + }) +}) diff --git a/frontend/src/api/admin.ts b/frontend/src/api/admin.ts index 86e8dc382..1010f67c0 100644 --- a/frontend/src/api/admin.ts +++ b/frontend/src/api/admin.ts @@ -709,6 +709,8 @@ export interface LeaderboardItem { requests: number tokens: number cost: number + member_count?: number + active_member_count?: number } export interface LeaderboardResponse { @@ -717,6 +719,7 @@ export interface LeaderboardResponse { metric: string start_date?: string | null end_date?: string | null + attribution?: 'current_membership' } export interface CostForecastResponse { @@ -1298,6 +1301,7 @@ export const adminApi = { model?: string include_inactive?: boolean exclude_admin?: boolean + user_group_id?: string }): Promise { const cacheKey = buildCacheKey('admin:stats:leaderboard:users', params) return cachedRequest( @@ -1312,6 +1316,35 @@ export const adminApi = { ) }, + async getLeaderboardUserGroups(params?: { + start_date?: string + end_date?: string + preset?: string + timezone?: string + tz_offset_minutes?: number + metric?: 'requests' | 'tokens' | 'cost' + order?: 'asc' | 'desc' + limit?: number + offset?: number + provider_name?: string + model?: string + include_inactive?: boolean + exclude_admin?: boolean + }): Promise { + const cacheKey = buildCacheKey('admin:stats:leaderboard:user-groups', params) + return cachedRequest( + cacheKey, + async () => { + const response = await apiClient.get( + '/api/admin/stats/leaderboard/user-groups', + { params } + ) + return response.data + }, + 20 * 1000 + ) + }, + async getLeaderboardApiKeys(params?: { start_date?: string end_date?: string @@ -1595,6 +1628,7 @@ export const adminApi = { timezone?: string tz_offset_minutes?: number user_id?: string + user_group_id?: string model?: string provider_name?: string }, diff --git a/frontend/src/api/dashboard.ts b/frontend/src/api/dashboard.ts index 44e6867eb..4d36ba1f8 100644 --- a/frontend/src/api/dashboard.ts +++ b/frontend/src/api/dashboard.ts @@ -219,6 +219,7 @@ export interface RequestDetail { has_format_conversion?: boolean | null model: string target_model?: string | null // 映射后的目标模型名 + response_model?: string | null // 上游响应体实际返回的模型名 requested_reasoning_effort?: string | null reasoning_effort?: string | null service_tier?: string | null diff --git a/frontend/src/api/endpoints/provider_oauth.ts b/frontend/src/api/endpoints/provider_oauth.ts index 5b0c1c68e..048113212 100644 --- a/frontend/src/api/endpoints/provider_oauth.ts +++ b/frontend/src/api/endpoints/provider_oauth.ts @@ -376,7 +376,7 @@ function jsonValueContainsAgentIdentity(value: unknown): boolean { export interface DeviceAuthorizeRequest { start_url?: string region?: string - auth_type?: 'builder_id' | 'identity_center' | 'google' | 'github' | 'browser' + auth_type?: 'builder_id' | 'identity_center' | 'google' | 'github' | 'browser' | 'device' login_option?: 'google' | 'github' | 'default' redirect_uri?: string proxy_node_id?: string diff --git a/frontend/src/api/endpoints/providers.ts b/frontend/src/api/endpoints/providers.ts index ab70f598b..d73206ebf 100644 --- a/frontend/src/api/endpoints/providers.ts +++ b/frontend/src/api/endpoints/providers.ts @@ -1,5 +1,5 @@ import client from '../client' -import { buildCacheKey, cachedRequest, dedupedRequest } from '@/utils/cache' +import { buildCacheKey, cache, cachedRequest, dedupedRequest } from '@/utils/cache' import type { ClaudeCodeAdvancedConfig, FailoverRulesConfig, @@ -142,6 +142,15 @@ export async function updateProvider( requestOptions?: ProviderRequestOptions, ): Promise { const response = await client.patch(`/api/admin/providers/${providerId}`, data, requestOptions) + cache.delete(`providers:detail:${providerId}`) + if ('pool_advanced' in data) { + cache.delete('pool:overview') + for (const kind of ['keys', 'scores']) { + const prefix = `pool:${kind}:${providerId}` + cache.delete(prefix) + cache.deleteByPrefix(`${prefix}:`) + } + } return normalizeProviderSummary(response.data) } diff --git a/frontend/src/api/endpoints/types/provider.ts b/frontend/src/api/endpoints/types/provider.ts index 0d55c80c1..ee6f53fd2 100644 --- a/frontend/src/api/endpoints/types/provider.ts +++ b/frontend/src/api/endpoints/types/provider.ts @@ -462,6 +462,23 @@ export interface GrokUpstreamMetadata { account_user_id?: string | null } +export interface XaiUpstreamMetadata { + updated_at?: number + subscription_title?: string + usage_percentage?: number + remaining_percentage?: number + usage_label?: string + usage_limit?: number + current_usage?: number + remaining?: number + next_reset_at?: number + prepaid_balance?: number + on_demand_cap?: number + on_demand_used?: number + on_demand_remaining?: number + period_type?: string +} + export interface GeminiCliTierMetadata { id?: string | null tierType?: string | null @@ -512,14 +529,29 @@ export interface GeminiCliUpstreamMetadata { quota_by_model?: Record | null } +export interface ClaudeCodeUpstreamMetadata { + updated_at?: number + five_hour_used_percent?: number + five_hour_reset_at?: number + seven_day_used_percent?: number + seven_day_reset_at?: number + seven_day_sonnet_used_percent?: number + seven_day_sonnet_reset_at?: number + seven_day_fable_used_percent?: number + seven_day_fable_reset_at?: number + reset_credits?: QuotaResetCreditsSnapshot +} + export interface UpstreamMetadata { codex?: CodexUpstreamMetadata + claude_code?: ClaudeCodeUpstreamMetadata antigravity?: AntigravityUpstreamMetadata kiro?: KiroUpstreamMetadata windsurf?: WindsurfUpstreamMetadata chatgpt_web?: ChatGPTWebUpstreamMetadata grok?: GrokUpstreamMetadata gemini_cli?: GeminiCliUpstreamMetadata + xai?: XaiUpstreamMetadata } // 按格式的健康度数据 @@ -758,7 +790,7 @@ export interface HealthRelatedMonitorResponse { related_providers: HealthRelatedMonitor[] } -export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'chatgpt_web' | 'gemini_cli' | 'antigravity' | 'kiro' | 'grok' | 'windsurf' | 'vertex_ai' +export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'chatgpt_web' | 'gemini_cli' | 'antigravity' | 'kiro' | 'grok' | 'xai' | 'windsurf' | 'vertex_ai' export interface ClaudeCodeAdvancedConfig { // 会话数量控制:null/undefined 表示不限制 @@ -804,6 +836,8 @@ export interface PoolAdvancedConfig { sticky_session_ttl_seconds?: number | null load_threshold_percent?: number | null skip_exhausted_accounts?: boolean | null + // Codex only: treat remaining quota <= 1% as exhausted (default false). + reserve_minimum_quota?: boolean // 旧字段(兼容读取) lru_enabled?: boolean scheduling_mode?: 'lru' | 'multi_score' | null diff --git a/frontend/src/api/me.ts b/frontend/src/api/me.ts index a816c9875..b1509c05f 100644 --- a/frontend/src/api/me.ts +++ b/frontend/src/api/me.ts @@ -60,6 +60,7 @@ export interface UsageRecordDetail { reasoning_effort?: string | null service_tier?: string | null actual_service_tier?: string | null + response_model?: string | null input_tokens: number effective_input_tokens?: number output_tokens: number @@ -392,6 +393,7 @@ export const meApi = { has_format_conversion?: boolean | null has_fallback?: boolean | null target_model?: string | null + response_model?: string | null request_type?: string | null requested_reasoning_effort?: string | null reasoning_effort?: string | null @@ -438,6 +440,7 @@ export const meApi = { has_format_conversion?: boolean | null has_fallback?: boolean | null target_model?: string | null + response_model?: string | null request_type?: string | null requested_reasoning_effort?: string | null reasoning_effort?: string | null diff --git a/frontend/src/api/usage.ts b/frontend/src/api/usage.ts index 122781503..e6bd58359 100644 --- a/frontend/src/api/usage.ts +++ b/frontend/src/api/usage.ts @@ -20,6 +20,7 @@ export interface UsageRecord { reasoning_effort?: string | null service_tier?: string | null actual_service_tier?: string | null + response_model?: string | null input_tokens: number effective_input_tokens?: number output_tokens: number @@ -136,6 +137,7 @@ export interface UsageFilters { from?: string to?: string user_id?: string // UUID + user_group_id?: string // UUID provider_id?: string // UUID model?: string search?: string @@ -641,6 +643,7 @@ export const usageApi = { has_format_conversion?: boolean | null has_fallback?: boolean | null target_model?: string | null + response_model?: string | null request_type?: string | null requested_reasoning_effort?: string | null reasoning_effort?: string | null @@ -712,6 +715,7 @@ export const usageApi = { has_format_conversion?: boolean | null has_fallback?: boolean | null target_model?: string | null + response_model?: string | null request_type?: string | null requested_reasoning_effort?: string | null reasoning_effort?: string | null diff --git a/frontend/src/api/usageRecords.ts b/frontend/src/api/usageRecords.ts index 8ef0b4150..24afc8c38 100644 --- a/frontend/src/api/usageRecords.ts +++ b/frontend/src/api/usageRecords.ts @@ -19,6 +19,7 @@ export interface UsageRecord { model: string target_model?: string | null // 映射后的目标模型名(若无映射则为空) model_version?: string | null // Provider 返回的实际模型版本(列表轻量字段) + response_model?: string | null // 上游响应体实际返回的模型名 request_type?: string | null // 由请求语义识别出的操作类型 requested_reasoning_effort?: string | null // 用户请求侧 reasoning 级别,用于展示转换关系 reasoning_effort?: string | null // 从发送给 Provider 的请求体提取的 reasoning 级别 @@ -65,5 +66,13 @@ export interface UsageRecord { response_time_updated_at?: string | null has_fallback?: boolean has_retry?: boolean + /** + * 是否存在被调度跳过的候选(候选在调度阶段即被判定不可用,从未向上游发起请求)。 + * 与 has_fallback 的区别:has_fallback 代表"更靠前的候选真的失败了", + * 本字段代表"更靠前的候选压根没被发出去",用于解释"无报错却换了提供商"。 + */ + has_skipped_candidate?: boolean + /** 被跳过候选的原因列表(后端已按候选顺序去重) */ + skipped_candidate_reasons?: string[] image_progress?: ImageProgress | null } diff --git a/frontend/src/api/users.ts b/frontend/src/api/users.ts index de41cdb49..19062939c 100644 --- a/frontend/src/api/users.ts +++ b/frontend/src/api/users.ts @@ -120,9 +120,28 @@ export interface UserBatchRolePayload { role: UserRole } -export type UserBatchAction = 'enable' | 'disable' | 'update_access_control' | 'update_role' +export type UserBatchBalanceOperation = 'add' | 'deduct' -export type UserBatchActionPayload = UserBatchAccessControlPayload | UserBatchRolePayload +export interface UserBatchBalanceAdjustmentPayload { + operation: UserBatchBalanceOperation + amount: number +} + +export function buildUserBatchBalanceAdjustmentPayload( + operation: UserBatchBalanceOperation, + amountInput: string | number, +): UserBatchBalanceAdjustmentPayload | null { + const amount = typeof amountInput === 'number' ? amountInput : Number(amountInput.trim()) + if (!Number.isFinite(amount) || amount <= 0) return null + return { operation, amount } +} + +export type UserBatchAction = 'enable' | 'disable' | 'update_access_control' | 'update_role' | 'adjust_wallet_balance' + +export type UserBatchActionPayload = + | UserBatchAccessControlPayload + | UserBatchRolePayload + | UserBatchBalanceAdjustmentPayload export interface UserBatchToggleActionRequest { selection: UserBatchSelection @@ -142,10 +161,18 @@ export interface UserBatchRoleActionRequest { payload: UserBatchRolePayload } +export interface UserBatchBalanceActionRequest { + selection: UserBatchSelection + action: 'adjust_wallet_balance' + payload: UserBatchBalanceAdjustmentPayload + idempotency_key: string +} + export type UserBatchActionRequest = | UserBatchToggleActionRequest | UserBatchAccessControlActionRequest | UserBatchRoleActionRequest + | UserBatchBalanceActionRequest export interface UserBatchActionFailure { user_id: string @@ -157,6 +184,10 @@ export interface UserBatchActionResponse { success: number failed: number failures: UserBatchActionFailure[] + interrupted?: boolean + completed_user_ids?: string[] + uncertain_user_ids?: string[] + unprocessed_user_ids?: string[] warnings?: UserBatchSelectionWarning[] action?: string modified_fields?: string[] diff --git a/frontend/src/components/common/HelpHint.vue b/frontend/src/components/common/HelpHint.vue index e02a33041..48c7050d1 100644 --- a/frontend/src/components/common/HelpHint.vue +++ b/frontend/src/components/common/HelpHint.vue @@ -24,7 +24,7 @@ const open = ref(false)